diff options
Diffstat (limited to 'internal/llm/provider')
| -rw-r--r-- | internal/llm/provider/openai.go | 32 | ||||
| -rw-r--r-- | internal/llm/provider/provider.go | 12 |
2 files changed, 41 insertions, 3 deletions
diff --git a/internal/llm/provider/openai.go b/internal/llm/provider/openai.go index 4d45aebfa..b557df535 100644 --- a/internal/llm/provider/openai.go +++ b/internal/llm/provider/openai.go @@ -21,6 +21,7 @@ type openaiOptions struct { baseURL string disableCache bool reasoningEffort string + extraHeaders map[string]string } type OpenAIOption func(*openaiOptions) @@ -49,6 +50,12 @@ func newOpenAIClient(opts providerClientOptions) OpenAIClient { 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, @@ -204,11 +211,18 @@ func (o *openaiClient) send(ctx context.Context, messages []message.Message, too content = openaiResponse.Choices[0].Message.Content } + toolCalls := o.toolCalls(*openaiResponse) + finishReason := o.finishReason(string(openaiResponse.Choices[0].FinishReason)) + + if len(toolCalls) > 0 { + finishReason = message.FinishReasonToolUse + } + return &ProviderResponse{ Content: content, - ToolCalls: o.toolCalls(*openaiResponse), + ToolCalls: toolCalls, Usage: o.usage(*openaiResponse), - FinishReason: o.finishReason(string(openaiResponse.Choices[0].FinishReason)), + FinishReason: finishReason, }, nil } } @@ -267,13 +281,19 @@ func (o *openaiClient) stream(ctx context.Context, messages []message.Message, t 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(toolCalls) > 0 { + finishReason = message.FinishReasonToolUse + } + eventChan <- ProviderEvent{ Type: EventComplete, Response: &ProviderResponse{ Content: currentContent, ToolCalls: toolCalls, Usage: o.usage(acc.ChatCompletion), - FinishReason: o.finishReason(string(acc.ChatCompletion.Choices[0].FinishReason)), + FinishReason: finishReason, }, } close(eventChan) @@ -375,6 +395,12 @@ func WithOpenAIBaseURL(baseURL string) OpenAIOption { } } +func WithOpenAIExtraHeaders(headers map[string]string) OpenAIOption { + return func(options *openaiOptions) { + options.extraHeaders = headers + } +} + func WithOpenAIDisableCache() OpenAIOption { return func(options *openaiOptions) { options.disableCache = true diff --git a/internal/llm/provider/provider.go b/internal/llm/provider/provider.go index 737b6fb00..1545bc27a 100644 --- a/internal/llm/provider/provider.go +++ b/internal/llm/provider/provider.go @@ -120,6 +120,18 @@ func NewProvider(providerName models.ModelProvider, opts ...ProviderClientOption options: clientOptions, client: newAzureClient(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.ProviderMock: // TODO: implement mock client for test panic("not implemented") |
