summaryrefslogtreecommitdiffhomepage
path: root/internal/llm/provider
diff options
context:
space:
mode:
Diffstat (limited to 'internal/llm/provider')
-rw-r--r--internal/llm/provider/anthropic.go472
-rw-r--r--internal/llm/provider/azure.go47
-rw-r--r--internal/llm/provider/bedrock.go100
-rw-r--r--internal/llm/provider/gemini.go555
-rw-r--r--internal/llm/provider/openai.go149
-rw-r--r--internal/llm/provider/openai_completion.go317
-rw-r--r--internal/llm/provider/openai_response.go393
-rw-r--r--internal/llm/provider/provider.go269
-rw-r--r--internal/llm/provider/vertexai.go34
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,
- }
-}