summaryrefslogtreecommitdiffhomepage
path: root/internal/llm/provider
diff options
context:
space:
mode:
Diffstat (limited to 'internal/llm/provider')
-rw-r--r--internal/llm/provider/openai.go32
-rw-r--r--internal/llm/provider/provider.go12
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")