diff options
Diffstat (limited to 'internal/llm/provider')
| -rw-r--r-- | internal/llm/provider/anthropic.go | 472 | ||||
| -rw-r--r-- | internal/llm/provider/azure.go | 47 | ||||
| -rw-r--r-- | internal/llm/provider/bedrock.go | 100 | ||||
| -rw-r--r-- | internal/llm/provider/gemini.go | 555 | ||||
| -rw-r--r-- | internal/llm/provider/openai.go | 149 | ||||
| -rw-r--r-- | internal/llm/provider/openai_completion.go | 317 | ||||
| -rw-r--r-- | internal/llm/provider/openai_response.go | 393 | ||||
| -rw-r--r-- | internal/llm/provider/provider.go | 269 | ||||
| -rw-r--r-- | internal/llm/provider/vertexai.go | 34 |
9 files changed, 0 insertions, 2336 deletions
diff --git a/internal/llm/provider/anthropic.go b/internal/llm/provider/anthropic.go deleted file mode 100644 index 24bcb48fb..000000000 --- a/internal/llm/provider/anthropic.go +++ /dev/null @@ -1,472 +0,0 @@ -package provider - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "io" - "strings" - "time" - - "github.com/anthropics/anthropic-sdk-go" - "github.com/anthropics/anthropic-sdk-go/bedrock" - "github.com/anthropics/anthropic-sdk-go/option" - "github.com/sst/opencode/internal/config" - "github.com/sst/opencode/internal/llm/models" - "github.com/sst/opencode/internal/llm/tools" - "github.com/sst/opencode/internal/message" - "github.com/sst/opencode/internal/status" - "log/slog" -) - -type anthropicOptions struct { - useBedrock bool - disableCache bool - shouldThink func(userMessage string) bool -} - -type AnthropicOption func(*anthropicOptions) - -type anthropicClient struct { - providerOptions providerClientOptions - options anthropicOptions - client anthropic.Client -} - -type AnthropicClient ProviderClient - -func newAnthropicClient(opts providerClientOptions) AnthropicClient { - anthropicOpts := anthropicOptions{} - for _, o := range opts.anthropicOptions { - o(&anthropicOpts) - } - - anthropicClientOptions := []option.RequestOption{} - if opts.apiKey != "" { - anthropicClientOptions = append(anthropicClientOptions, option.WithAPIKey(opts.apiKey)) - } - if anthropicOpts.useBedrock { - anthropicClientOptions = append(anthropicClientOptions, bedrock.WithLoadDefaultConfig(context.Background())) - } - - client := anthropic.NewClient(anthropicClientOptions...) - return &anthropicClient{ - providerOptions: opts, - options: anthropicOpts, - client: client, - } -} - -func (a *anthropicClient) convertMessages(messages []message.Message) (anthropicMessages []anthropic.MessageParam) { - for i, msg := range messages { - cache := false - if i > len(messages)-3 { - cache = true - } - switch msg.Role { - case message.User: - content := anthropic.NewTextBlock(msg.Content().String()) - if cache && !a.options.disableCache { - content.OfRequestTextBlock.CacheControl = anthropic.CacheControlEphemeralParam{ - Type: "ephemeral", - } - } - var contentBlocks []anthropic.ContentBlockParamUnion - contentBlocks = append(contentBlocks, content) - for _, binaryContent := range msg.BinaryContent() { - base64Image := binaryContent.String(models.ProviderAnthropic) - imageBlock := anthropic.NewImageBlockBase64(binaryContent.MIMEType, base64Image) - contentBlocks = append(contentBlocks, imageBlock) - } - anthropicMessages = append(anthropicMessages, anthropic.NewUserMessage(contentBlocks...)) - - case message.Assistant: - blocks := []anthropic.ContentBlockParamUnion{} - - if msg.Content() != nil { - content := msg.Content().String() - if strings.TrimSpace(content) != "" { - block := anthropic.NewTextBlock(content) - if cache && !a.options.disableCache { - block.OfRequestTextBlock.CacheControl = anthropic.CacheControlEphemeralParam{ - Type: "ephemeral", - } - } - blocks = append(blocks, block) - } - } - - for _, toolCall := range msg.ToolCalls() { - var inputMap map[string]any - err := json.Unmarshal([]byte(toolCall.Input), &inputMap) - if err != nil { - continue - } - blocks = append(blocks, anthropic.ContentBlockParamOfRequestToolUseBlock(toolCall.ID, inputMap, toolCall.Name)) - } - - if len(blocks) == 0 { - slog.Warn("There is a message without content, investigate, this should not happen") - continue - } - anthropicMessages = append(anthropicMessages, anthropic.NewAssistantMessage(blocks...)) - - case message.Tool: - results := make([]anthropic.ContentBlockParamUnion, len(msg.ToolResults())) - for i, toolResult := range msg.ToolResults() { - results[i] = anthropic.NewToolResultBlock(toolResult.ToolCallID, toolResult.Content, toolResult.IsError) - } - anthropicMessages = append(anthropicMessages, anthropic.NewUserMessage(results...)) - } - } - return -} - -func (a *anthropicClient) convertTools(tools []tools.BaseTool) []anthropic.ToolUnionParam { - anthropicTools := make([]anthropic.ToolUnionParam, len(tools)) - - for i, tool := range tools { - info := tool.Info() - toolParam := anthropic.ToolParam{ - Name: info.Name, - Description: anthropic.String(info.Description), - InputSchema: anthropic.ToolInputSchemaParam{ - Properties: info.Parameters, - // TODO: figure out how we can tell claude the required fields? - }, - } - - if i == len(tools)-1 && !a.options.disableCache { - toolParam.CacheControl = anthropic.CacheControlEphemeralParam{ - Type: "ephemeral", - } - } - - anthropicTools[i] = anthropic.ToolUnionParam{OfTool: &toolParam} - } - - return anthropicTools -} - -func (a *anthropicClient) finishReason(reason string) message.FinishReason { - switch reason { - case "end_turn": - return message.FinishReasonEndTurn - case "max_tokens": - return message.FinishReasonMaxTokens - case "tool_use": - return message.FinishReasonToolUse - case "stop_sequence": - return message.FinishReasonEndTurn - default: - return message.FinishReasonUnknown - } -} - -func (a *anthropicClient) preparedMessages(messages []anthropic.MessageParam, tools []anthropic.ToolUnionParam) anthropic.MessageNewParams { - var thinkingParam anthropic.ThinkingConfigParamUnion - lastMessage := messages[len(messages)-1] - isUser := lastMessage.Role == anthropic.MessageParamRoleUser - messageContent := "" - temperature := anthropic.Float(0) - if isUser { - for _, m := range lastMessage.Content { - if m.OfRequestTextBlock != nil && m.OfRequestTextBlock.Text != "" { - messageContent = m.OfRequestTextBlock.Text - } - } - if messageContent != "" && a.options.shouldThink != nil && a.options.shouldThink(messageContent) { - thinkingParam = anthropic.ThinkingConfigParamUnion{ - OfThinkingConfigEnabled: &anthropic.ThinkingConfigEnabledParam{ - BudgetTokens: int64(float64(a.providerOptions.maxTokens) * 0.8), - Type: "enabled", - }, - } - temperature = anthropic.Float(1) - } - } - - return anthropic.MessageNewParams{ - Model: anthropic.Model(a.providerOptions.model.APIModel), - MaxTokens: a.providerOptions.maxTokens, - Temperature: temperature, - Messages: messages, - Tools: tools, - Thinking: thinkingParam, - System: []anthropic.TextBlockParam{ - { - Text: a.providerOptions.systemMessage, - CacheControl: anthropic.CacheControlEphemeralParam{ - Type: "ephemeral", - }, - }, - }, - } -} - -func (a *anthropicClient) send(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (resposne *ProviderResponse, err error) { - preparedMessages := a.preparedMessages(a.convertMessages(messages), a.convertTools(tools)) - cfg := config.Get() - if cfg.Debug { - jsonData, _ := json.Marshal(preparedMessages) - slog.Debug("Prepared messages", "messages", string(jsonData)) - } - - attempts := 0 - for { - attempts++ - anthropicResponse, err := a.client.Messages.New( - ctx, - preparedMessages, - ) - // If there is an error we are going to see if we can retry the call - if err != nil { - slog.Error("Error in Anthropic API call", "error", err) - retry, after, retryErr := a.shouldRetry(attempts, err) - duration := time.Duration(after) * time.Millisecond - if retryErr != nil { - return nil, retryErr - } - if retry { - status.Warn(fmt.Sprintf("Retrying due to rate limit... attempt %d of %d", attempts, maxRetries), status.WithDuration(duration)) - select { - case <-ctx.Done(): - return nil, ctx.Err() - case <-time.After(duration): - continue - } - } - return nil, retryErr - } - - content := "" - for _, block := range anthropicResponse.Content { - if text, ok := block.AsAny().(anthropic.TextBlock); ok { - content += text.Text - } - } - - return &ProviderResponse{ - Content: content, - ToolCalls: a.toolCalls(*anthropicResponse), - Usage: a.usage(*anthropicResponse), - }, nil - } -} - -func (a *anthropicClient) stream(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent { - preparedMessages := a.preparedMessages(a.convertMessages(messages), a.convertTools(tools)) - cfg := config.Get() - if cfg.Debug { - jsonData, _ := json.Marshal(preparedMessages) - slog.Debug("Prepared messages", "messages", string(jsonData)) - } - attempts := 0 - eventChan := make(chan ProviderEvent) - go func() { - for { - attempts++ - anthropicStream := a.client.Messages.NewStreaming( - ctx, - preparedMessages, - ) - accumulatedMessage := anthropic.Message{} - - currentToolCallID := "" - for anthropicStream.Next() { - event := anthropicStream.Current() - err := accumulatedMessage.Accumulate(event) - if err != nil { - slog.Warn("Error accumulating message", "error", err) - continue - } - - switch event := event.AsAny().(type) { - case anthropic.ContentBlockStartEvent: - if event.ContentBlock.Type == "text" { - eventChan <- ProviderEvent{Type: EventContentStart} - } else if event.ContentBlock.Type == "tool_use" { - currentToolCallID = event.ContentBlock.ID - eventChan <- ProviderEvent{ - Type: EventToolUseStart, - ToolCall: &message.ToolCall{ - ID: event.ContentBlock.ID, - Name: event.ContentBlock.Name, - Finished: false, - }, - } - } - - case anthropic.ContentBlockDeltaEvent: - if event.Delta.Type == "thinking_delta" && event.Delta.Thinking != "" { - eventChan <- ProviderEvent{ - Type: EventThinkingDelta, - Thinking: event.Delta.Thinking, - } - } else if event.Delta.Type == "text_delta" && event.Delta.Text != "" { - eventChan <- ProviderEvent{ - Type: EventContentDelta, - Content: event.Delta.Text, - } - } else if event.Delta.Type == "input_json_delta" { - if currentToolCallID != "" { - eventChan <- ProviderEvent{ - Type: EventToolUseDelta, - ToolCall: &message.ToolCall{ - ID: currentToolCallID, - Finished: false, - Input: event.Delta.JSON.PartialJSON.Raw(), - }, - } - } - } - case anthropic.ContentBlockStopEvent: - if currentToolCallID != "" { - eventChan <- ProviderEvent{ - Type: EventToolUseStop, - ToolCall: &message.ToolCall{ - ID: currentToolCallID, - }, - } - currentToolCallID = "" - } else { - eventChan <- ProviderEvent{Type: EventContentStop} - } - - case anthropic.MessageStopEvent: - content := "" - for _, block := range accumulatedMessage.Content { - if text, ok := block.AsAny().(anthropic.TextBlock); ok { - content += text.Text - } - } - - eventChan <- ProviderEvent{ - Type: EventComplete, - Response: &ProviderResponse{ - Content: content, - ToolCalls: a.toolCalls(accumulatedMessage), - Usage: a.usage(accumulatedMessage), - FinishReason: a.finishReason(string(accumulatedMessage.StopReason)), - }, - } - } - } - - err := anthropicStream.Err() - if err == nil || errors.Is(err, io.EOF) { - close(eventChan) - return - } - // If there is an error we are going to see if we can retry the call - retry, after, retryErr := a.shouldRetry(attempts, err) - duration := time.Duration(after) * time.Millisecond - if retryErr != nil { - eventChan <- ProviderEvent{Type: EventError, Error: retryErr} - close(eventChan) - return - } - if retry { - status.Warn(fmt.Sprintf("Retrying due to rate limit... attempt %d of %d", attempts, maxRetries), status.WithDuration(duration)) - select { - case <-ctx.Done(): - // context cancelled - if ctx.Err() != nil { - eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()} - } - close(eventChan) - return - case <-time.After(duration): - continue - } - } - if ctx.Err() != nil { - eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()} - } - - close(eventChan) - return - } - }() - return eventChan -} - -func (a *anthropicClient) shouldRetry(attempts int, err error) (bool, int64, error) { - var apierr *anthropic.Error - if !errors.As(err, &apierr) { - return false, 0, err - } - - if apierr.StatusCode != 429 && apierr.StatusCode != 529 { - return false, 0, err - } - - if attempts > maxRetries { - return false, 0, fmt.Errorf("maximum retry attempts reached for rate limit: %d retries", maxRetries) - } - - retryMs := 0 - retryAfterValues := apierr.Response.Header.Values("Retry-After") - - backoffMs := 2000 * (1 << (attempts - 1)) - jitterMs := int(float64(backoffMs) * 0.2) - retryMs = backoffMs + jitterMs - if len(retryAfterValues) > 0 { - if _, err := fmt.Sscanf(retryAfterValues[0], "%d", &retryMs); err == nil { - retryMs = retryMs * 1000 - } - } - return true, int64(retryMs), nil -} - -func (a *anthropicClient) toolCalls(msg anthropic.Message) []message.ToolCall { - var toolCalls []message.ToolCall - - for _, block := range msg.Content { - switch variant := block.AsAny().(type) { - case anthropic.ToolUseBlock: - toolCall := message.ToolCall{ - ID: variant.ID, - Name: variant.Name, - Input: string(variant.Input), - Type: string(variant.Type), - Finished: true, - } - toolCalls = append(toolCalls, toolCall) - } - } - - return toolCalls -} - -func (a *anthropicClient) usage(msg anthropic.Message) TokenUsage { - return TokenUsage{ - InputTokens: msg.Usage.InputTokens, - OutputTokens: msg.Usage.OutputTokens, - CacheCreationTokens: msg.Usage.CacheCreationInputTokens, - CacheReadTokens: msg.Usage.CacheReadInputTokens, - } -} - -func WithAnthropicBedrock(useBedrock bool) AnthropicOption { - return func(options *anthropicOptions) { - options.useBedrock = useBedrock - } -} - -func WithAnthropicDisableCache() AnthropicOption { - return func(options *anthropicOptions) { - options.disableCache = true - } -} - -func DefaultShouldThinkFn(s string) bool { - return strings.Contains(strings.ToLower(s), "think") -} - -func WithAnthropicShouldThinkFn(fn func(string) bool) AnthropicOption { - return func(options *anthropicOptions) { - options.shouldThink = fn - } -} diff --git a/internal/llm/provider/azure.go b/internal/llm/provider/azure.go deleted file mode 100644 index 6368a181c..000000000 --- a/internal/llm/provider/azure.go +++ /dev/null @@ -1,47 +0,0 @@ -package provider - -import ( - "os" - - "github.com/Azure/azure-sdk-for-go/sdk/azidentity" - "github.com/openai/openai-go" - "github.com/openai/openai-go/azure" - "github.com/openai/openai-go/option" -) - -type azureClient struct { - *openaiClient -} - -type AzureClient ProviderClient - -func newAzureClient(opts providerClientOptions) AzureClient { - - endpoint := os.Getenv("AZURE_OPENAI_ENDPOINT") // ex: https://foo.openai.azure.com - apiVersion := os.Getenv("AZURE_OPENAI_API_VERSION") // ex: 2025-04-01-preview - - if endpoint == "" || apiVersion == "" { - return &azureClient{openaiClient: newOpenAIClient(opts).(*openaiClient)} - } - - reqOpts := []option.RequestOption{ - azure.WithEndpoint(endpoint, apiVersion), - } - - if opts.apiKey != "" || os.Getenv("AZURE_OPENAI_API_KEY") != "" { - key := opts.apiKey - if key == "" { - key = os.Getenv("AZURE_OPENAI_API_KEY") - } - reqOpts = append(reqOpts, azure.WithAPIKey(key)) - } else if cred, err := azidentity.NewDefaultAzureCredential(nil); err == nil { - reqOpts = append(reqOpts, azure.WithTokenCredential(cred)) - } - - base := &openaiClient{ - providerOptions: opts, - client: openai.NewClient(reqOpts...), - } - - return &azureClient{openaiClient: base} -} diff --git a/internal/llm/provider/bedrock.go b/internal/llm/provider/bedrock.go deleted file mode 100644 index 4622ae5ff..000000000 --- a/internal/llm/provider/bedrock.go +++ /dev/null @@ -1,100 +0,0 @@ -package provider - -import ( - "context" - "errors" - "fmt" - "os" - "strings" - - "github.com/sst/opencode/internal/llm/tools" - "github.com/sst/opencode/internal/message" -) - -type bedrockOptions struct { - // Bedrock specific options can be added here -} - -type BedrockOption func(*bedrockOptions) - -type bedrockClient struct { - providerOptions providerClientOptions - options bedrockOptions - childProvider ProviderClient -} - -type BedrockClient ProviderClient - -func newBedrockClient(opts providerClientOptions) BedrockClient { - bedrockOpts := bedrockOptions{} - // Apply bedrock specific options if they are added in the future - - // Get AWS region from environment - region := os.Getenv("AWS_REGION") - if region == "" { - region = os.Getenv("AWS_DEFAULT_REGION") - } - - if region == "" { - region = "us-east-1" // default region - } - if len(region) < 2 { - return &bedrockClient{ - providerOptions: opts, - options: bedrockOpts, - childProvider: nil, // Will cause an error when used - } - } - - // Prefix the model name with region - regionPrefix := region[:2] - modelName := opts.model.APIModel - opts.model.APIModel = fmt.Sprintf("%s.%s", regionPrefix, modelName) - - // Determine which provider to use based on the model - if strings.Contains(string(opts.model.APIModel), "anthropic") { - // Create Anthropic client with Bedrock configuration - anthropicOpts := opts - anthropicOpts.anthropicOptions = append(anthropicOpts.anthropicOptions, - WithAnthropicBedrock(true), - WithAnthropicDisableCache(), - ) - return &bedrockClient{ - providerOptions: opts, - options: bedrockOpts, - childProvider: newAnthropicClient(anthropicOpts), - } - } - - // Return client with nil childProvider if model is not supported - // This will cause an error when used - return &bedrockClient{ - providerOptions: opts, - options: bedrockOpts, - childProvider: nil, - } -} - -func (b *bedrockClient) send(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) { - if b.childProvider == nil { - return nil, errors.New("unsupported model for bedrock provider") - } - return b.childProvider.send(ctx, messages, tools) -} - -func (b *bedrockClient) stream(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent { - eventChan := make(chan ProviderEvent) - - if b.childProvider == nil { - go func() { - eventChan <- ProviderEvent{ - Type: EventError, - Error: errors.New("unsupported model for bedrock provider"), - } - close(eventChan) - }() - return eventChan - } - - return b.childProvider.stream(ctx, messages, tools) -} diff --git a/internal/llm/provider/gemini.go b/internal/llm/provider/gemini.go deleted file mode 100644 index 8b8e33698..000000000 --- a/internal/llm/provider/gemini.go +++ /dev/null @@ -1,555 +0,0 @@ -package provider - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "io" - "strings" - "time" - - "github.com/google/uuid" - "github.com/sst/opencode/internal/config" - "github.com/sst/opencode/internal/llm/tools" - "github.com/sst/opencode/internal/message" - "github.com/sst/opencode/internal/status" - "google.golang.org/genai" - "log/slog" -) - -type geminiOptions struct { - disableCache bool -} - -type GeminiOption func(*geminiOptions) - -type geminiClient struct { - providerOptions providerClientOptions - options geminiOptions - client *genai.Client -} - -type GeminiClient ProviderClient - -func newGeminiClient(opts providerClientOptions) GeminiClient { - geminiOpts := geminiOptions{} - for _, o := range opts.geminiOptions { - o(&geminiOpts) - } - - client, err := genai.NewClient(context.Background(), &genai.ClientConfig{APIKey: opts.apiKey, Backend: genai.BackendGeminiAPI}) - if err != nil { - slog.Error("Failed to create Gemini client", "error", err) - return nil - } - - return &geminiClient{ - providerOptions: opts, - options: geminiOpts, - client: client, - } -} - -func (g *geminiClient) convertMessages(messages []message.Message) []*genai.Content { - var history []*genai.Content - for _, msg := range messages { - switch msg.Role { - case message.User: - var parts []*genai.Part - parts = append(parts, &genai.Part{Text: msg.Content().String()}) - for _, binaryContent := range msg.BinaryContent() { - imageFormat := strings.Split(binaryContent.MIMEType, "/") - parts = append(parts, &genai.Part{InlineData: &genai.Blob{ - MIMEType: imageFormat[1], - Data: binaryContent.Data, - }}) - } - history = append(history, &genai.Content{ - Parts: parts, - Role: "user", - }) - case message.Assistant: - content := &genai.Content{ - Role: "model", - Parts: []*genai.Part{}, - } - - if msg.Content().String() != "" { - content.Parts = append(content.Parts, &genai.Part{Text: msg.Content().String()}) - } - - if len(msg.ToolCalls()) > 0 { - for _, call := range msg.ToolCalls() { - args, _ := parseJsonToMap(call.Input) - content.Parts = append(content.Parts, &genai.Part{ - FunctionCall: &genai.FunctionCall{ - Name: call.Name, - Args: args, - }, - }) - } - } - - history = append(history, content) - - case message.Tool: - for _, result := range msg.ToolResults() { - response := map[string]interface{}{"result": result.Content} - parsed, err := parseJsonToMap(result.Content) - if err == nil { - response = parsed - } - - var toolCall message.ToolCall - for _, m := range messages { - if m.Role == message.Assistant { - for _, call := range m.ToolCalls() { - if call.ID == result.ToolCallID { - toolCall = call - break - } - } - } - } - - history = append(history, &genai.Content{ - Parts: []*genai.Part{ - { - FunctionResponse: &genai.FunctionResponse{ - Name: toolCall.Name, - Response: response, - }, - }, - }, - Role: "function", - }) - } - } - } - - return history -} - -func (g *geminiClient) convertTools(tools []tools.BaseTool) []*genai.Tool { - geminiTool := &genai.Tool{} - geminiTool.FunctionDeclarations = make([]*genai.FunctionDeclaration, 0, len(tools)) - - for _, tool := range tools { - info := tool.Info() - declaration := &genai.FunctionDeclaration{ - Name: info.Name, - Description: info.Description, - Parameters: &genai.Schema{ - Type: genai.TypeObject, - Properties: convertSchemaProperties(info.Parameters), - Required: info.Required, - }, - } - - geminiTool.FunctionDeclarations = append(geminiTool.FunctionDeclarations, declaration) - } - - return []*genai.Tool{geminiTool} -} - -func (g *geminiClient) finishReason(reason genai.FinishReason) message.FinishReason { - switch { - case reason == genai.FinishReasonStop: - return message.FinishReasonEndTurn - case reason == genai.FinishReasonMaxTokens: - return message.FinishReasonMaxTokens - default: - return message.FinishReasonUnknown - } -} - -func (g *geminiClient) send(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) { - // Convert messages - geminiMessages := g.convertMessages(messages) - - cfg := config.Get() - if cfg.Debug { - jsonData, _ := json.Marshal(geminiMessages) - slog.Debug("Prepared messages", "messages", string(jsonData)) - } - - history := geminiMessages[:len(geminiMessages)-1] // All but last message - lastMsg := geminiMessages[len(geminiMessages)-1] - config := &genai.GenerateContentConfig{ - MaxOutputTokens: int32(g.providerOptions.maxTokens), - SystemInstruction: &genai.Content{ - Parts: []*genai.Part{{Text: g.providerOptions.systemMessage}}, - }, - } - if len(tools) > 0 { - config.Tools = g.convertTools(tools) - } - chat, _ := g.client.Chats.Create(ctx, g.providerOptions.model.APIModel, config, history) - - attempts := 0 - for { - attempts++ - var toolCalls []message.ToolCall - - var lastMsgParts []genai.Part - for _, part := range lastMsg.Parts { - lastMsgParts = append(lastMsgParts, *part) - } - resp, err := chat.SendMessage(ctx, lastMsgParts...) - // If there is an error we are going to see if we can retry the call - if err != nil { - retry, after, retryErr := g.shouldRetry(attempts, err) - duration := time.Duration(after) * time.Millisecond - if retryErr != nil { - return nil, retryErr - } - if retry { - status.Warn(fmt.Sprintf("Retrying due to rate limit... attempt %d of %d", attempts, maxRetries), status.WithDuration(duration)) - select { - case <-ctx.Done(): - return nil, ctx.Err() - case <-time.After(duration): - continue - } - } - return nil, retryErr - } - - content := "" - - if len(resp.Candidates) > 0 && resp.Candidates[0].Content != nil { - for _, part := range resp.Candidates[0].Content.Parts { - switch { - case part.Text != "": - content = string(part.Text) - case part.FunctionCall != nil: - id := "call_" + uuid.New().String() - args, _ := json.Marshal(part.FunctionCall.Args) - toolCalls = append(toolCalls, message.ToolCall{ - ID: id, - Name: part.FunctionCall.Name, - Input: string(args), - Type: "function", - Finished: true, - }) - } - } - } - finishReason := message.FinishReasonEndTurn - if len(resp.Candidates) > 0 { - finishReason = g.finishReason(resp.Candidates[0].FinishReason) - } - if len(toolCalls) > 0 { - finishReason = message.FinishReasonToolUse - } - - return &ProviderResponse{ - Content: content, - ToolCalls: toolCalls, - Usage: g.usage(resp), - FinishReason: finishReason, - }, nil - } -} - -func (g *geminiClient) stream(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent { - // Convert messages - geminiMessages := g.convertMessages(messages) - - cfg := config.Get() - if cfg.Debug { - jsonData, _ := json.Marshal(geminiMessages) - slog.Debug("Prepared messages", "messages", string(jsonData)) - } - - history := geminiMessages[:len(geminiMessages)-1] // All but last message - lastMsg := geminiMessages[len(geminiMessages)-1] - config := &genai.GenerateContentConfig{ - MaxOutputTokens: int32(g.providerOptions.maxTokens), - SystemInstruction: &genai.Content{ - Parts: []*genai.Part{{Text: g.providerOptions.systemMessage}}, - }, - } - if len(tools) > 0 { - config.Tools = g.convertTools(tools) - } - chat, _ := g.client.Chats.Create(ctx, g.providerOptions.model.APIModel, config, history) - - attempts := 0 - eventChan := make(chan ProviderEvent) - - go func() { - defer close(eventChan) - - for { - attempts++ - - currentContent := "" - toolCalls := []message.ToolCall{} - var finalResp *genai.GenerateContentResponse - - eventChan <- ProviderEvent{Type: EventContentStart} - - var lastMsgParts []genai.Part - - for _, part := range lastMsg.Parts { - lastMsgParts = append(lastMsgParts, *part) - } - for resp, err := range chat.SendMessageStream(ctx, lastMsgParts...) { - if err != nil { - retry, after, retryErr := g.shouldRetry(attempts, err) - duration := time.Duration(after) * time.Millisecond - if retryErr != nil { - eventChan <- ProviderEvent{Type: EventError, Error: retryErr} - return - } - if retry { - status.Warn(fmt.Sprintf("Retrying due to rate limit... attempt %d of %d", attempts, maxRetries), status.WithDuration(duration)) - select { - case <-ctx.Done(): - if ctx.Err() != nil { - eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()} - } - - return - case <-time.After(duration): - break - } - } else { - eventChan <- ProviderEvent{Type: EventError, Error: err} - return - } - } - - finalResp = resp - - if len(resp.Candidates) > 0 && resp.Candidates[0].Content != nil { - for _, part := range resp.Candidates[0].Content.Parts { - switch { - case part.Text != "": - delta := string(part.Text) - if delta != "" { - eventChan <- ProviderEvent{ - Type: EventContentDelta, - Content: delta, - } - currentContent += delta - } - case part.FunctionCall != nil: - id := "call_" + uuid.New().String() - args, _ := json.Marshal(part.FunctionCall.Args) - newCall := message.ToolCall{ - ID: id, - Name: part.FunctionCall.Name, - Input: string(args), - Type: "function", - Finished: true, - } - - isNew := true - for _, existing := range toolCalls { - if existing.Name == newCall.Name && existing.Input == newCall.Input { - isNew = false - break - } - } - - if isNew { - toolCalls = append(toolCalls, newCall) - } - } - } - } - } - - eventChan <- ProviderEvent{Type: EventContentStop} - - if finalResp != nil { - - finishReason := message.FinishReasonEndTurn - if len(finalResp.Candidates) > 0 { - finishReason = g.finishReason(finalResp.Candidates[0].FinishReason) - } - if len(toolCalls) > 0 { - finishReason = message.FinishReasonToolUse - } - eventChan <- ProviderEvent{ - Type: EventComplete, - Response: &ProviderResponse{ - Content: currentContent, - ToolCalls: toolCalls, - Usage: g.usage(finalResp), - FinishReason: finishReason, - }, - } - return - } - - } - }() - - return eventChan -} - -func (g *geminiClient) shouldRetry(attempts int, err error) (bool, int64, error) { - // Check if error is a rate limit error - if attempts > maxRetries { - return false, 0, fmt.Errorf("maximum retry attempts reached for rate limit: %d retries", maxRetries) - } - - // Gemini doesn't have a standard error type we can check against - // So we'll check the error message for rate limit indicators - if errors.Is(err, io.EOF) { - return false, 0, err - } - - errMsg := err.Error() - isRateLimit := false - - // Check for common rate limit error messages - if contains(errMsg, "rate limit", "quota exceeded", "too many requests") { - isRateLimit = true - } - - if !isRateLimit { - return false, 0, err - } - - // Calculate backoff with jitter - backoffMs := 2000 * (1 << (attempts - 1)) - jitterMs := int(float64(backoffMs) * 0.2) - retryMs := backoffMs + jitterMs - - return true, int64(retryMs), nil -} - -func (g *geminiClient) toolCalls(resp *genai.GenerateContentResponse) []message.ToolCall { - var toolCalls []message.ToolCall - - if len(resp.Candidates) > 0 && resp.Candidates[0].Content != nil { - for _, part := range resp.Candidates[0].Content.Parts { - if part.FunctionCall != nil { - id := "call_" + uuid.New().String() - args, _ := json.Marshal(part.FunctionCall.Args) - toolCalls = append(toolCalls, message.ToolCall{ - ID: id, - Name: part.FunctionCall.Name, - Input: string(args), - Type: "function", - }) - } - } - } - - return toolCalls -} - -func (g *geminiClient) usage(resp *genai.GenerateContentResponse) TokenUsage { - if resp == nil || resp.UsageMetadata == nil { - return TokenUsage{} - } - - return TokenUsage{ - InputTokens: int64(resp.UsageMetadata.PromptTokenCount), - OutputTokens: int64(resp.UsageMetadata.CandidatesTokenCount), - CacheCreationTokens: 0, // Not directly provided by Gemini - CacheReadTokens: int64(resp.UsageMetadata.CachedContentTokenCount), - } -} - -func WithGeminiDisableCache() GeminiOption { - return func(options *geminiOptions) { - options.disableCache = true - } -} - -// Helper functions -func parseJsonToMap(jsonStr string) (map[string]interface{}, error) { - var result map[string]interface{} - err := json.Unmarshal([]byte(jsonStr), &result) - return result, err -} - -func convertSchemaProperties(parameters map[string]interface{}) map[string]*genai.Schema { - properties := make(map[string]*genai.Schema) - - for name, param := range parameters { - properties[name] = convertToSchema(param) - } - - return properties -} - -func convertToSchema(param interface{}) *genai.Schema { - schema := &genai.Schema{Type: genai.TypeString} - - paramMap, ok := param.(map[string]interface{}) - if !ok { - return schema - } - - if desc, ok := paramMap["description"].(string); ok { - schema.Description = desc - } - - typeVal, hasType := paramMap["type"] - if !hasType { - return schema - } - - typeStr, ok := typeVal.(string) - if !ok { - return schema - } - - schema.Type = mapJSONTypeToGenAI(typeStr) - - switch typeStr { - case "array": - schema.Items = processArrayItems(paramMap) - case "object": - if props, ok := paramMap["properties"].(map[string]interface{}); ok { - schema.Properties = convertSchemaProperties(props) - } - } - - return schema -} - -func processArrayItems(paramMap map[string]interface{}) *genai.Schema { - items, ok := paramMap["items"].(map[string]interface{}) - if !ok { - return nil - } - - return convertToSchema(items) -} - -func mapJSONTypeToGenAI(jsonType string) genai.Type { - switch jsonType { - case "string": - return genai.TypeString - case "number": - return genai.TypeNumber - case "integer": - return genai.TypeInteger - case "boolean": - return genai.TypeBoolean - case "array": - return genai.TypeArray - case "object": - return genai.TypeObject - default: - return genai.TypeString // Default to string for unknown types - } -} - -func contains(s string, substrs ...string) bool { - for _, substr := range substrs { - if strings.Contains(strings.ToLower(s), strings.ToLower(substr)) { - return true - } - } - return false -} diff --git a/internal/llm/provider/openai.go b/internal/llm/provider/openai.go deleted file mode 100644 index db77a3844..000000000 --- a/internal/llm/provider/openai.go +++ /dev/null @@ -1,149 +0,0 @@ -package provider - -import ( - "context" - "errors" - "fmt" - "log/slog" - "github.com/openai/openai-go" - "github.com/openai/openai-go/option" - "github.com/sst/opencode/internal/llm/models" - "github.com/sst/opencode/internal/llm/tools" - "github.com/sst/opencode/internal/message" -) - -type openaiOptions struct { - baseURL string - disableCache bool - reasoningEffort string - extraHeaders map[string]string -} - -type OpenAIOption func(*openaiOptions) - -type openaiClient struct { - providerOptions providerClientOptions - options openaiOptions - client openai.Client -} - -type OpenAIClient ProviderClient - -func newOpenAIClient(opts providerClientOptions) OpenAIClient { - openaiOpts := openaiOptions{ - reasoningEffort: "medium", - } - for _, o := range opts.openaiOptions { - o(&openaiOpts) - } - - openaiClientOptions := []option.RequestOption{} - if opts.apiKey != "" { - openaiClientOptions = append(openaiClientOptions, option.WithAPIKey(opts.apiKey)) - } - if openaiOpts.baseURL != "" { - openaiClientOptions = append(openaiClientOptions, option.WithBaseURL(openaiOpts.baseURL)) - } - - if openaiOpts.extraHeaders != nil { - for key, value := range openaiOpts.extraHeaders { - openaiClientOptions = append(openaiClientOptions, option.WithHeader(key, value)) - } - } - - client := openai.NewClient(openaiClientOptions...) - return &openaiClient{ - providerOptions: opts, - options: openaiOpts, - client: client, - } -} - -func (o *openaiClient) send(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (response *ProviderResponse, err error) { - if o.providerOptions.model.ID == models.OpenAIModels[models.CodexMini].ID || o.providerOptions.model.ID == models.OpenAIModels[models.O1Pro].ID { - return o.sendResponseMessages(ctx, messages, tools) - } - return o.sendChatcompletionMessage(ctx, messages, tools) -} - -func (o *openaiClient) stream(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent { - if o.providerOptions.model.ID == models.OpenAIModels[models.CodexMini].ID || o.providerOptions.model.ID == models.OpenAIModels[models.O1Pro].ID { - return o.streamResponseMessages(ctx, messages, tools) - } - return o.streamChatCompletionMessages(ctx, messages, tools) -} - - -func (o *openaiClient) finishReason(reason string) message.FinishReason { - switch reason { - case "stop": - return message.FinishReasonEndTurn - case "length": - return message.FinishReasonMaxTokens - case "tool_calls": - return message.FinishReasonToolUse - default: - return message.FinishReasonUnknown - } -} - - -func (o *openaiClient) shouldRetry(attempts int, err error) (bool, int64, error) { - var apierr *openai.Error - if !errors.As(err, &apierr) { - return false, 0, err - } - - if apierr.StatusCode != 429 && apierr.StatusCode != 500 { - return false, 0, err - } - - if attempts > maxRetries { - return false, 0, fmt.Errorf("maximum retry attempts reached for rate limit: %d retries", maxRetries) - } - - retryMs := 0 - retryAfterValues := apierr.Response.Header.Values("Retry-After") - - backoffMs := 2000 * (1 << (attempts - 1)) - jitterMs := int(float64(backoffMs) * 0.2) - retryMs = backoffMs + jitterMs - if len(retryAfterValues) > 0 { - if _, err := fmt.Sscanf(retryAfterValues[0], "%d", &retryMs); err == nil { - retryMs = retryMs * 1000 - } - } - return true, int64(retryMs), nil -} - - -func WithOpenAIBaseURL(baseURL string) OpenAIOption { - return func(options *openaiOptions) { - options.baseURL = baseURL - } -} - -func WithOpenAIExtraHeaders(headers map[string]string) OpenAIOption { - return func(options *openaiOptions) { - options.extraHeaders = headers - } -} - -func WithOpenAIDisableCache() OpenAIOption { - return func(options *openaiOptions) { - options.disableCache = true - } -} - -func WithReasoningEffort(effort string) OpenAIOption { - return func(options *openaiOptions) { - defaultReasoningEffort := "medium" - switch effort { - case "low", "medium", "high": - defaultReasoningEffort = effort - default: - slog.Warn("Invalid reasoning effort, using default: medium") - } - options.reasoningEffort = defaultReasoningEffort - } -} diff --git a/internal/llm/provider/openai_completion.go b/internal/llm/provider/openai_completion.go deleted file mode 100644 index e3b837231..000000000 --- a/internal/llm/provider/openai_completion.go +++ /dev/null @@ -1,317 +0,0 @@ -package provider - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "io" - "log/slog" - "time" - - "github.com/openai/openai-go" - "github.com/openai/openai-go/shared" - "github.com/sst/opencode/internal/config" - "github.com/sst/opencode/internal/llm/models" - "github.com/sst/opencode/internal/llm/tools" - "github.com/sst/opencode/internal/message" - "github.com/sst/opencode/internal/status" -) - -func (o *openaiClient) convertMessagesToChatCompletionMessages(messages []message.Message) (openaiMessages []openai.ChatCompletionMessageParamUnion) { - // Add system message first - openaiMessages = append(openaiMessages, openai.SystemMessage(o.providerOptions.systemMessage)) - - for _, msg := range messages { - switch msg.Role { - case message.User: - var content []openai.ChatCompletionContentPartUnionParam - textBlock := openai.ChatCompletionContentPartTextParam{Text: msg.Content().String()} - content = append(content, openai.ChatCompletionContentPartUnionParam{OfText: &textBlock}) - for _, binaryContent := range msg.BinaryContent() { - imageURL := openai.ChatCompletionContentPartImageImageURLParam{URL: binaryContent.String(models.ProviderOpenAI)} - imageBlock := openai.ChatCompletionContentPartImageParam{ImageURL: imageURL} - - content = append(content, openai.ChatCompletionContentPartUnionParam{OfImageURL: &imageBlock}) - } - - openaiMessages = append(openaiMessages, openai.UserMessage(content)) - - case message.Assistant: - assistantMsg := openai.ChatCompletionAssistantMessageParam{ - Role: "assistant", - } - - if msg.Content().String() != "" { - assistantMsg.Content = openai.ChatCompletionAssistantMessageParamContentUnion{ - OfString: openai.String(msg.Content().String()), - } - } - - if len(msg.ToolCalls()) > 0 { - assistantMsg.ToolCalls = make([]openai.ChatCompletionMessageToolCallParam, len(msg.ToolCalls())) - for i, call := range msg.ToolCalls() { - assistantMsg.ToolCalls[i] = openai.ChatCompletionMessageToolCallParam{ - ID: call.ID, - Type: "function", - Function: openai.ChatCompletionMessageToolCallFunctionParam{ - Name: call.Name, - Arguments: call.Input, - }, - } - } - } - - openaiMessages = append(openaiMessages, openai.ChatCompletionMessageParamUnion{ - OfAssistant: &assistantMsg, - }) - - case message.Tool: - for _, result := range msg.ToolResults() { - openaiMessages = append(openaiMessages, - openai.ToolMessage(result.Content, result.ToolCallID), - ) - } - } - } - - return -} - -func (o *openaiClient) convertToChatCompletionTools(tools []tools.BaseTool) []openai.ChatCompletionToolParam { - openaiTools := make([]openai.ChatCompletionToolParam, len(tools)) - - for i, tool := range tools { - info := tool.Info() - openaiTools[i] = openai.ChatCompletionToolParam{ - Function: openai.FunctionDefinitionParam{ - Name: info.Name, - Description: openai.String(info.Description), - Parameters: openai.FunctionParameters{ - "type": "object", - "properties": info.Parameters, - "required": info.Required, - }, - }, - } - } - - return openaiTools -} - -func (o *openaiClient) preparedChatCompletionParams(messages []openai.ChatCompletionMessageParamUnion, tools []openai.ChatCompletionToolParam) openai.ChatCompletionNewParams { - params := openai.ChatCompletionNewParams{ - Model: openai.ChatModel(o.providerOptions.model.APIModel), - Messages: messages, - Tools: tools, - } - if o.providerOptions.model.CanReason == true { - params.MaxCompletionTokens = openai.Int(o.providerOptions.maxTokens) - switch o.options.reasoningEffort { - case "low": - params.ReasoningEffort = shared.ReasoningEffortLow - case "medium": - params.ReasoningEffort = shared.ReasoningEffortMedium - case "high": - params.ReasoningEffort = shared.ReasoningEffortHigh - default: - params.ReasoningEffort = shared.ReasoningEffortMedium - } - } else { - params.MaxTokens = openai.Int(o.providerOptions.maxTokens) - } - - if o.providerOptions.model.Provider == models.ProviderOpenRouter { - params.WithExtraFields(map[string]any{ - "provider": map[string]any{ - "require_parameters": true, - }, - }) - } - - return params -} - -func (o *openaiClient) sendChatcompletionMessage(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (response *ProviderResponse, err error) { - params := o.preparedChatCompletionParams(o.convertMessagesToChatCompletionMessages(messages), o.convertToChatCompletionTools(tools)) - cfg := config.Get() - if cfg.Debug { - jsonData, _ := json.Marshal(params) - slog.Debug("Prepared messages", "messages", string(jsonData)) - } - attempts := 0 - for { - attempts++ - openaiResponse, err := o.client.Chat.Completions.New( - ctx, - params, - ) - // If there is an error we are going to see if we can retry the call - if err != nil { - retry, after, retryErr := o.shouldRetry(attempts, err) - duration := time.Duration(after) * time.Millisecond - if retryErr != nil { - return nil, retryErr - } - if retry { - status.Warn(fmt.Sprintf("Retrying due to rate limit... attempt %d of %d", attempts, maxRetries), status.WithDuration(duration)) - select { - case <-ctx.Done(): - return nil, ctx.Err() - case <-time.After(duration): - continue - } - } - return nil, retryErr - } - - content := "" - if openaiResponse.Choices[0].Message.Content != "" { - content = openaiResponse.Choices[0].Message.Content - } - - toolCalls := o.chatCompletionToolCalls(*openaiResponse) - finishReason := o.finishReason(string(openaiResponse.Choices[0].FinishReason)) - - if len(toolCalls) > 0 { - finishReason = message.FinishReasonToolUse - } - - return &ProviderResponse{ - Content: content, - ToolCalls: toolCalls, - Usage: o.usage(*openaiResponse), - FinishReason: finishReason, - }, nil - } -} - -func (o *openaiClient) streamChatCompletionMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent { - params := o.preparedChatCompletionParams(o.convertMessagesToChatCompletionMessages(messages), o.convertToChatCompletionTools(tools)) - params.StreamOptions = openai.ChatCompletionStreamOptionsParam{ - IncludeUsage: openai.Bool(true), - } - - cfg := config.Get() - if cfg.Debug { - jsonData, _ := json.Marshal(params) - slog.Debug("Prepared messages", "messages", string(jsonData)) - } - - attempts := 0 - eventChan := make(chan ProviderEvent) - - go func() { - for { - attempts++ - openaiStream := o.client.Chat.Completions.NewStreaming( - ctx, - params, - ) - - acc := openai.ChatCompletionAccumulator{} - currentContent := "" - toolCalls := make([]message.ToolCall, 0) - - for openaiStream.Next() { - chunk := openaiStream.Current() - acc.AddChunk(chunk) - - for _, choice := range chunk.Choices { - if choice.Delta.Content != "" { - eventChan <- ProviderEvent{ - Type: EventContentDelta, - Content: choice.Delta.Content, - } - currentContent += choice.Delta.Content - } - } - } - - err := openaiStream.Err() - if err == nil || errors.Is(err, io.EOF) { - // Stream completed successfully - finishReason := o.finishReason(string(acc.ChatCompletion.Choices[0].FinishReason)) - if len(acc.ChatCompletion.Choices[0].Message.ToolCalls) > 0 { - toolCalls = append(toolCalls, o.chatCompletionToolCalls(acc.ChatCompletion)...) - } - if len(toolCalls) > 0 { - finishReason = message.FinishReasonToolUse - } - - eventChan <- ProviderEvent{ - Type: EventComplete, - Response: &ProviderResponse{ - Content: currentContent, - ToolCalls: toolCalls, - Usage: o.usage(acc.ChatCompletion), - FinishReason: finishReason, - }, - } - close(eventChan) - return - } - - // If there is an error we are going to see if we can retry the call - retry, after, retryErr := o.shouldRetry(attempts, err) - duration := time.Duration(after) * time.Millisecond - if retryErr != nil { - eventChan <- ProviderEvent{Type: EventError, Error: retryErr} - close(eventChan) - return - } - if retry { - status.Warn(fmt.Sprintf("Retrying due to rate limit... attempt %d of %d", attempts, maxRetries), status.WithDuration(duration)) - select { - case <-ctx.Done(): - // context cancelled - if ctx.Err() == nil { - eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()} - } - close(eventChan) - return - case <-time.After(duration): - continue - } - } - eventChan <- ProviderEvent{Type: EventError, Error: retryErr} - close(eventChan) - return - } - }() - - return eventChan -} - - -func (o *openaiClient) chatCompletionToolCalls(completion openai.ChatCompletion) []message.ToolCall { - var toolCalls []message.ToolCall - - if len(completion.Choices) > 0 && len(completion.Choices[0].Message.ToolCalls) > 0 { - for _, call := range completion.Choices[0].Message.ToolCalls { - toolCall := message.ToolCall{ - ID: call.ID, - Name: call.Function.Name, - Input: call.Function.Arguments, - Type: "function", - Finished: true, - } - toolCalls = append(toolCalls, toolCall) - } - } - - return toolCalls -} - -func (o *openaiClient) usage(completion openai.ChatCompletion) TokenUsage { - cachedTokens := completion.Usage.PromptTokensDetails.CachedTokens - inputTokens := completion.Usage.PromptTokens - cachedTokens - - return TokenUsage{ - InputTokens: inputTokens, - OutputTokens: completion.Usage.CompletionTokens, - CacheCreationTokens: 0, // OpenAI doesn't provide this directly - CacheReadTokens: cachedTokens, - } -} - diff --git a/internal/llm/provider/openai_response.go b/internal/llm/provider/openai_response.go deleted file mode 100644 index 96a61c4db..000000000 --- a/internal/llm/provider/openai_response.go +++ /dev/null @@ -1,393 +0,0 @@ -package provider - - -import ( - "github.com/openai/openai-go" - "github.com/openai/openai-go/responses" - "github.com/sst/opencode/internal/llm/models" - "github.com/sst/opencode/internal/llm/tools" - "github.com/sst/opencode/internal/message" - "context" - "encoding/json" - "errors" - "fmt" - "io" - "time" - - "log/slog" - - "github.com/openai/openai-go/shared" - "github.com/sst/opencode/internal/config" - "github.com/sst/opencode/internal/status" -) - -func (o *openaiClient) convertMessagesToResponseParams(messages []message.Message) responses.ResponseInputParam { - inputItems := responses.ResponseInputParam{} - - inputItems = append(inputItems, responses.ResponseInputItemUnionParam{ - OfMessage: &responses.EasyInputMessageParam{ - Content: responses.EasyInputMessageContentUnionParam{OfString: openai.String(o.providerOptions.systemMessage)}, - Role: responses.EasyInputMessageRoleSystem, - }, - }) - - for _, msg := range messages { - switch msg.Role { - case message.User: - inputItemContentList := responses.ResponseInputMessageContentListParam{ - responses.ResponseInputContentUnionParam{ - OfInputText: &responses.ResponseInputTextParam{ - Text: msg.Content().String(), - }, - }, - } - - for _, binaryContent := range msg.BinaryContent() { - inputItemContentList = append(inputItemContentList, responses.ResponseInputContentUnionParam{ - OfInputImage: &responses.ResponseInputImageParam{ - ImageURL: openai.String(binaryContent.String(models.ProviderOpenAI)), - }, - }) - } - - userMsg := responses.ResponseInputItemUnionParam{ - OfInputMessage: &responses.ResponseInputItemMessageParam{ - Content: inputItemContentList, - Role: string(responses.ResponseInputMessageItemRoleUser), - }, - } - inputItems = append(inputItems, userMsg) - - case message.Assistant: - if msg.Content().String() != "" { - assistantMsg := responses.ResponseInputItemUnionParam{ - OfOutputMessage: &responses.ResponseOutputMessageParam{ - Content: []responses.ResponseOutputMessageContentUnionParam{{ - OfOutputText: &responses.ResponseOutputTextParam{ - Text: msg.Content().String(), - }, - }}, - }, - } - inputItems = append(inputItems, assistantMsg) - } - - if len(msg.ToolCalls()) > 0 { - for _, call := range msg.ToolCalls() { - toolMsg := responses.ResponseInputItemUnionParam{ - OfFunctionCall: &responses.ResponseFunctionToolCallParam{ - CallID: call.ID, - Name: call.Name, - Arguments: call.Input, - }, - } - inputItems = append(inputItems, toolMsg) - } - } - - case message.Tool: - for _, result := range msg.ToolResults() { - toolMsg := responses.ResponseInputItemUnionParam{ - OfFunctionCallOutput: &responses.ResponseInputItemFunctionCallOutputParam{ - Output: result.Content, - CallID: result.ToolCallID, - }, - } - inputItems = append(inputItems, toolMsg) - } - } - } - - return inputItems -} - -func (o *openaiClient) convertToResponseTools(tools []tools.BaseTool) []responses.ToolUnionParam { - outputTools := make([]responses.ToolUnionParam, len(tools)) - - for i, tool := range tools { - info := tool.Info() - outputTools[i] = responses.ToolUnionParam{ - OfFunction: &responses.FunctionToolParam{ - Name: info.Name, - Description: openai.String(info.Description), - Parameters: map[string]any{ - "type": "object", - "properties": info.Parameters, - "required": info.Required, - }, - }, - } - } - - return outputTools -} - - -func (o *openaiClient) preparedResponseParams(input responses.ResponseInputParam, tools []responses.ToolUnionParam) responses.ResponseNewParams { - params := responses.ResponseNewParams{ - Model: shared.ResponsesModel(o.providerOptions.model.APIModel), - Input: responses.ResponseNewParamsInputUnion{OfInputItemList: input}, - Tools: tools, - } - - params.MaxOutputTokens = openai.Int(o.providerOptions.maxTokens) - - if o.providerOptions.model.CanReason == true { - switch o.options.reasoningEffort { - case "low": - params.Reasoning.Effort = shared.ReasoningEffortLow - case "medium": - params.Reasoning.Effort = shared.ReasoningEffortMedium - case "high": - params.Reasoning.Effort = shared.ReasoningEffortHigh - default: - params.Reasoning.Effort = shared.ReasoningEffortMedium - } - } - - if o.providerOptions.model.Provider == models.ProviderOpenRouter { - params.WithExtraFields(map[string]any{ - "provider": map[string]any{ - "require_parameters": true, - }, - }) - } - - return params -} - -func (o *openaiClient) sendResponseMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (response *ProviderResponse, err error) { - params := o.preparedResponseParams(o.convertMessagesToResponseParams(messages), o.convertToResponseTools(tools)) - cfg := config.Get() - if cfg.Debug { - jsonData, _ := json.Marshal(params) - slog.Debug("Prepared messages", "messages", string(jsonData)) - } - attempts := 0 - for { - attempts++ - openaiResponse, err := o.client.Responses.New( - ctx, - params, - ) - // If there is an error we are going to see if we can retry the call - if err != nil { - retry, after, retryErr := o.shouldRetry(attempts, err) - duration := time.Duration(after) * time.Millisecond - if retryErr != nil { - return nil, retryErr - } - if retry { - status.Warn(fmt.Sprintf("Retrying due to rate limit... attempt %d of %d", attempts, maxRetries), status.WithDuration(duration)) - select { - case <-ctx.Done(): - return nil, ctx.Err() - case <-time.After(duration): - continue - } - } - return nil, retryErr - } - - content := "" - if openaiResponse.OutputText() != "" { - content = openaiResponse.OutputText() - } - - toolCalls := o.responseToolCalls(*openaiResponse) - finishReason := o.finishReason("stop") - - if len(toolCalls) > 0 { - finishReason = message.FinishReasonToolUse - } - - return &ProviderResponse{ - Content: content, - ToolCalls: toolCalls, - Usage: o.responseUsage(*openaiResponse), - FinishReason: finishReason, - }, nil - } -} - -func (o *openaiClient) streamResponseMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent { - eventChan := make(chan ProviderEvent) - - params := o.preparedResponseParams(o.convertMessagesToResponseParams(messages), o.convertToResponseTools(tools)) - - cfg := config.Get() - if cfg.Debug { - jsonData, _ := json.Marshal(params) - slog.Debug("Prepared messages", "messages", string(jsonData)) - } - - attempts := 0 - - go func() { - for { - attempts++ - stream := o.client.Responses.NewStreaming(ctx, params) - outputText := "" - currentToolCallID := "" - for stream.Next() { - event := stream.Current() - - switch event := event.AsAny().(type) { - case responses.ResponseCompletedEvent: - toolCalls := o.responseToolCalls(event.Response) - finishReason := o.finishReason("stop") - - if len(toolCalls) > 0 { - finishReason = message.FinishReasonToolUse - } - - eventChan <- ProviderEvent{ - Type: EventComplete, - Response: &ProviderResponse{ - Content: outputText, - ToolCalls: toolCalls, - Usage: o.responseUsage(event.Response), - FinishReason: finishReason, - }, - } - close(eventChan) - return - - case responses.ResponseTextDeltaEvent: - outputText += event.Delta - eventChan <- ProviderEvent{ - Type: EventContentDelta, - Content: event.Delta, - } - - case responses.ResponseTextDoneEvent: - eventChan <- ProviderEvent{ - Type: EventContentStop, - Content: outputText, - } - close(eventChan) - return - - case responses.ResponseOutputItemAddedEvent: - if event.Item.Type == "function_call" { - currentToolCallID = event.Item.ID - eventChan <- ProviderEvent{ - Type: EventToolUseStart, - ToolCall: &message.ToolCall{ - ID: event.Item.ID, - Name: event.Item.Name, - Finished: false, - }, - } - } - - case responses.ResponseFunctionCallArgumentsDeltaEvent: - if event.ItemID == currentToolCallID { - eventChan <- ProviderEvent{ - Type: EventToolUseDelta, - ToolCall: &message.ToolCall{ - ID: currentToolCallID, - Finished: false, - Input: event.Delta, - }, - } - } - - case responses.ResponseFunctionCallArgumentsDoneEvent: - if event.ItemID == currentToolCallID { - eventChan <- ProviderEvent{ - Type: EventToolUseStop, - ToolCall: &message.ToolCall{ - ID: currentToolCallID, - Input: event.Arguments, - }, - } - currentToolCallID = "" - } - - case responses.ResponseOutputItemDoneEvent: - if event.Item.Type == "function_call" { - eventChan <- ProviderEvent{ - Type: EventToolUseStop, - ToolCall: &message.ToolCall{ - ID: event.Item.ID, - Name: event.Item.Name, - Input: event.Item.Arguments, - Finished: true, - }, - } - currentToolCallID = "" - } - - } - } - - err := stream.Err() - if err == nil || errors.Is(err, io.EOF) { - close(eventChan) - return - } - - // If there is an error we are going to see if we can retry the call - retry, after, retryErr := o.shouldRetry(attempts, err) - duration := time.Duration(after) * time.Millisecond - if retryErr != nil { - eventChan <- ProviderEvent{Type: EventError, Error: retryErr} - close(eventChan) - return - } - if retry { - status.Warn(fmt.Sprintf("Retrying due to rate limit... attempt %d of %d", attempts, maxRetries), status.WithDuration(duration)) - select { - case <-ctx.Done(): - // context cancelled - if ctx.Err() == nil { - eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()} - } - close(eventChan) - return - case <-time.After(duration): - continue - } - } - eventChan <- ProviderEvent{Type: EventError, Error: retryErr} - close(eventChan) - return - } - }() - - return eventChan -} - - -func (o *openaiClient) responseToolCalls(response responses.Response) []message.ToolCall { - var toolCalls []message.ToolCall - - for _, output := range response.Output { - if output.Type == "function_call" { - call := output.AsFunctionCall() - toolCall := message.ToolCall{ - ID: call.ID, - Name: call.Name, - Input: call.Arguments, - Type: "function", - Finished: true, - } - toolCalls = append(toolCalls, toolCall) - } - } - - return toolCalls -} - -func (o *openaiClient) responseUsage(response responses.Response) TokenUsage { - cachedTokens := response.Usage.InputTokensDetails.CachedTokens - inputTokens := response.Usage.InputTokens - cachedTokens - - return TokenUsage{ - InputTokens: inputTokens, - OutputTokens: response.Usage.OutputTokens, - CacheCreationTokens: 0, // OpenAI doesn't provide this directly - CacheReadTokens: cachedTokens, - } -} diff --git a/internal/llm/provider/provider.go b/internal/llm/provider/provider.go deleted file mode 100644 index adcbfdbf7..000000000 --- a/internal/llm/provider/provider.go +++ /dev/null @@ -1,269 +0,0 @@ -package provider - -import ( - "context" - "fmt" - "log/slog" - - "github.com/sst/opencode/internal/llm/models" - "github.com/sst/opencode/internal/llm/tools" - "github.com/sst/opencode/internal/message" -) - -type EventType string - -const maxRetries = 8 - -const ( - EventContentStart EventType = "content_start" - EventToolUseStart EventType = "tool_use_start" - EventToolUseDelta EventType = "tool_use_delta" - EventToolUseStop EventType = "tool_use_stop" - EventContentDelta EventType = "content_delta" - EventThinkingDelta EventType = "thinking_delta" - EventContentStop EventType = "content_stop" - EventComplete EventType = "complete" - EventError EventType = "error" - EventWarning EventType = "warning" -) - -type TokenUsage struct { - InputTokens int64 - OutputTokens int64 - CacheCreationTokens int64 - CacheReadTokens int64 -} - -type ProviderResponse struct { - Content string - ToolCalls []message.ToolCall - Usage TokenUsage - FinishReason message.FinishReason -} - -type ProviderEvent struct { - Type EventType - - Content string - Thinking string - Response *ProviderResponse - ToolCall *message.ToolCall - Error error -} -type Provider interface { - SendMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) - - StreamResponse(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent - - Model() models.Model - - MaxTokens() int64 -} - -type providerClientOptions struct { - apiKey string - model models.Model - maxTokens int64 - systemMessage string - - anthropicOptions []AnthropicOption - openaiOptions []OpenAIOption - geminiOptions []GeminiOption - bedrockOptions []BedrockOption -} - -type ProviderClientOption func(*providerClientOptions) - -type ProviderClient interface { - send(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) - stream(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent -} - -type baseProvider[C ProviderClient] struct { - options providerClientOptions - client C -} - -func NewProvider(providerName models.ModelProvider, opts ...ProviderClientOption) (Provider, error) { - clientOptions := providerClientOptions{} - for _, o := range opts { - o(&clientOptions) - } - switch providerName { - case models.ProviderAnthropic: - return &baseProvider[AnthropicClient]{ - options: clientOptions, - client: newAnthropicClient(clientOptions), - }, nil - case models.ProviderOpenAI: - return &baseProvider[OpenAIClient]{ - options: clientOptions, - client: newOpenAIClient(clientOptions), - }, nil - case models.ProviderGemini: - return &baseProvider[GeminiClient]{ - options: clientOptions, - client: newGeminiClient(clientOptions), - }, nil - case models.ProviderBedrock: - return &baseProvider[BedrockClient]{ - options: clientOptions, - client: newBedrockClient(clientOptions), - }, nil - case models.ProviderGROQ: - clientOptions.openaiOptions = append(clientOptions.openaiOptions, - WithOpenAIBaseURL("https://api.groq.com/openai/v1"), - ) - return &baseProvider[OpenAIClient]{ - options: clientOptions, - client: newOpenAIClient(clientOptions), - }, nil - case models.ProviderAzure: - return &baseProvider[AzureClient]{ - options: clientOptions, - client: newAzureClient(clientOptions), - }, nil - case models.ProviderVertexAI: - return &baseProvider[VertexAIClient]{ - options: clientOptions, - client: newVertexAIClient(clientOptions), - }, nil - case models.ProviderOpenRouter: - clientOptions.openaiOptions = append(clientOptions.openaiOptions, - WithOpenAIBaseURL("https://openrouter.ai/api/v1"), - WithOpenAIExtraHeaders(map[string]string{ - "HTTP-Referer": "opencode.ai", - "X-Title": "OpenCode", - }), - ) - return &baseProvider[OpenAIClient]{ - options: clientOptions, - client: newOpenAIClient(clientOptions), - }, nil - case models.ProviderXAI: - clientOptions.openaiOptions = append(clientOptions.openaiOptions, - WithOpenAIBaseURL("https://api.x.ai/v1"), - ) - return &baseProvider[OpenAIClient]{ - options: clientOptions, - client: newOpenAIClient(clientOptions), - }, nil - - case models.ProviderMock: - // TODO: implement mock client for test - panic("not implemented") - } - return nil, fmt.Errorf("provider not supported: %s", providerName) -} - -func (p *baseProvider[C]) cleanMessages(messages []message.Message) (cleaned []message.Message) { - for _, msg := range messages { - // The message has no content - if len(msg.Parts) == 0 { - continue - } - cleaned = append(cleaned, msg) - } - return -} - -func (p *baseProvider[C]) SendMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) { - messages = p.cleanMessages(messages) - response, err := p.client.send(ctx, messages, tools) - if err == nil && response != nil { - slog.Debug("API request token usage", - "model", p.options.model.Name, - "input_tokens", response.Usage.InputTokens, - "output_tokens", response.Usage.OutputTokens, - "cache_creation_tokens", response.Usage.CacheCreationTokens, - "cache_read_tokens", response.Usage.CacheReadTokens, - "total_tokens", response.Usage.InputTokens+response.Usage.OutputTokens) - } - return response, err -} - -func (p *baseProvider[C]) Model() models.Model { - return p.options.model -} - -func (p *baseProvider[C]) MaxTokens() int64 { - return p.options.maxTokens -} - -func (p *baseProvider[C]) StreamResponse(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent { - messages = p.cleanMessages(messages) - eventChan := p.client.stream(ctx, messages, tools) - - // Create a new channel to intercept events - wrappedChan := make(chan ProviderEvent) - - go func() { - defer close(wrappedChan) - - for event := range eventChan { - // Pass the event through - wrappedChan <- event - - // Log token usage when we get the complete event - if event.Type == EventComplete && event.Response != nil { - slog.Debug("API streaming request token usage", - "model", p.options.model.Name, - "input_tokens", event.Response.Usage.InputTokens, - "output_tokens", event.Response.Usage.OutputTokens, - "cache_creation_tokens", event.Response.Usage.CacheCreationTokens, - "cache_read_tokens", event.Response.Usage.CacheReadTokens, - "total_tokens", event.Response.Usage.InputTokens+event.Response.Usage.OutputTokens) - } - } - }() - - return wrappedChan -} - -func WithAPIKey(apiKey string) ProviderClientOption { - return func(options *providerClientOptions) { - options.apiKey = apiKey - } -} - -func WithModel(model models.Model) ProviderClientOption { - return func(options *providerClientOptions) { - options.model = model - } -} - -func WithMaxTokens(maxTokens int64) ProviderClientOption { - return func(options *providerClientOptions) { - options.maxTokens = maxTokens - } -} - -func WithSystemMessage(systemMessage string) ProviderClientOption { - return func(options *providerClientOptions) { - options.systemMessage = systemMessage - } -} - -func WithAnthropicOptions(anthropicOptions ...AnthropicOption) ProviderClientOption { - return func(options *providerClientOptions) { - options.anthropicOptions = anthropicOptions - } -} - -func WithOpenAIOptions(openaiOptions ...OpenAIOption) ProviderClientOption { - return func(options *providerClientOptions) { - options.openaiOptions = openaiOptions - } -} - -func WithGeminiOptions(geminiOptions ...GeminiOption) ProviderClientOption { - return func(options *providerClientOptions) { - options.geminiOptions = geminiOptions - } -} - -func WithBedrockOptions(bedrockOptions ...BedrockOption) ProviderClientOption { - return func(options *providerClientOptions) { - options.bedrockOptions = bedrockOptions - } -} diff --git a/internal/llm/provider/vertexai.go b/internal/llm/provider/vertexai.go deleted file mode 100644 index 328d213fe..000000000 --- a/internal/llm/provider/vertexai.go +++ /dev/null @@ -1,34 +0,0 @@ -package provider - -import ( - "context" - "log/slog" - "os" - - "google.golang.org/genai" -) - -type VertexAIClient ProviderClient - -func newVertexAIClient(opts providerClientOptions) VertexAIClient { - geminiOpts := geminiOptions{} - for _, o := range opts.geminiOptions { - o(&geminiOpts) - } - - client, err := genai.NewClient(context.Background(), &genai.ClientConfig{ - Project: os.Getenv("VERTEXAI_PROJECT"), - Location: os.Getenv("VERTEXAI_LOCATION"), - Backend: genai.BackendVertexAI, - }) - if err != nil { - slog.Error("Failed to create VertexAI client", "error", err) - return nil - } - - return &geminiClient{ - providerOptions: opts, - options: geminiOpts, - client: client, - } -} |
