summaryrefslogtreecommitdiffhomepage
path: root/internal/llm/provider
diff options
context:
space:
mode:
Diffstat (limited to 'internal/llm/provider')
-rw-r--r--internal/llm/provider/anthropic.go531
-rw-r--r--internal/llm/provider/bedrock.go101
-rw-r--r--internal/llm/provider/gemini.go533
-rw-r--r--internal/llm/provider/openai.go401
-rw-r--r--internal/llm/provider/provider.go169
5 files changed, 1053 insertions, 682 deletions
diff --git a/internal/llm/provider/anthropic.go b/internal/llm/provider/anthropic.go
index 93c4308ad..c3a4efc49 100644
--- a/internal/llm/provider/anthropic.go
+++ b/internal/llm/provider/anthropic.go
@@ -12,187 +12,257 @@ import (
"github.com/anthropics/anthropic-sdk-go"
"github.com/anthropics/anthropic-sdk-go/bedrock"
"github.com/anthropics/anthropic-sdk-go/option"
- "github.com/kujtimiihoxha/termai/internal/llm/models"
+ "github.com/kujtimiihoxha/termai/internal/config"
"github.com/kujtimiihoxha/termai/internal/llm/tools"
+ "github.com/kujtimiihoxha/termai/internal/logging"
"github.com/kujtimiihoxha/termai/internal/message"
)
-type anthropicProvider struct {
- client anthropic.Client
- model models.Model
- maxTokens int64
- apiKey string
- systemMessage string
- useBedrock bool
- disableCache bool
+type anthropicOptions struct {
+ useBedrock bool
+ disableCache bool
+ shouldThink func(userMessage string) bool
}
-type AnthropicOption func(*anthropicProvider)
+type AnthropicOption func(*anthropicOptions)
-func WithAnthropicSystemMessage(message string) AnthropicOption {
- return func(a *anthropicProvider) {
- a.systemMessage = message
- }
+type anthropicClient struct {
+ providerOptions providerClientOptions
+ options anthropicOptions
+ client anthropic.Client
}
-func WithAnthropicMaxTokens(maxTokens int64) AnthropicOption {
- return func(a *anthropicProvider) {
- a.maxTokens = maxTokens
- }
-}
+type AnthropicClient ProviderClient
-func WithAnthropicModel(model models.Model) AnthropicOption {
- return func(a *anthropicProvider) {
- a.model = model
+func newAnthropicClient(opts providerClientOptions) AnthropicClient {
+ anthropicOpts := anthropicOptions{}
+ for _, o := range opts.anthropicOptions {
+ o(&anthropicOpts)
}
-}
-func WithAnthropicKey(apiKey string) AnthropicOption {
- return func(a *anthropicProvider) {
- a.apiKey = apiKey
+ anthropicClientOptions := []option.RequestOption{}
+ if opts.apiKey != "" {
+ anthropicClientOptions = append(anthropicClientOptions, option.WithAPIKey(opts.apiKey))
}
-}
-
-func WithAnthropicBedrock() AnthropicOption {
- return func(a *anthropicProvider) {
- a.useBedrock = true
+ if anthropicOpts.useBedrock {
+ anthropicClientOptions = append(anthropicClientOptions, bedrock.WithLoadDefaultConfig(context.Background()))
}
-}
-func WithAnthropicDisableCache() AnthropicOption {
- return func(a *anthropicProvider) {
- a.disableCache = true
+ client := anthropic.NewClient(anthropicClientOptions...)
+ return &anthropicClient{
+ providerOptions: opts,
+ options: anthropicOpts,
+ client: client,
}
}
-func NewAnthropicProvider(opts ...AnthropicOption) (Provider, error) {
- provider := &anthropicProvider{
- maxTokens: 1024,
- }
+func (a *anthropicClient) convertMessages(messages []message.Message) (anthropicMessages []anthropic.MessageParam) {
+ cachedBlocks := 0
+ for _, msg := range messages {
+ switch msg.Role {
+ case message.User:
+ content := anthropic.NewTextBlock(msg.Content().String())
+ if cachedBlocks < 2 && !a.options.disableCache {
+ content.OfRequestTextBlock.CacheControl = anthropic.CacheControlEphemeralParam{
+ Type: "ephemeral",
+ }
+ cachedBlocks++
+ }
+ anthropicMessages = append(anthropicMessages, anthropic.NewUserMessage(content))
- for _, opt := range opts {
- opt(provider)
- }
+ case message.Assistant:
+ blocks := []anthropic.ContentBlockParamUnion{}
+ if msg.Content().String() != "" {
+ content := anthropic.NewTextBlock(msg.Content().String())
+ if cachedBlocks < 2 && !a.options.disableCache {
+ content.OfRequestTextBlock.CacheControl = anthropic.CacheControlEphemeralParam{
+ Type: "ephemeral",
+ }
+ cachedBlocks++
+ }
+ blocks = append(blocks, content)
+ }
- if provider.systemMessage == "" {
- return nil, errors.New("system message is required")
- }
+ 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))
+ }
- anthropicOptions := []option.RequestOption{}
+ if len(blocks) == 0 {
+ logging.Warn("There is a message without content, investigate")
+ // This should never happend but we log this because we might have a bug in our cleanup method
+ continue
+ }
+ anthropicMessages = append(anthropicMessages, anthropic.NewAssistantMessage(blocks...))
- if provider.apiKey != "" {
- anthropicOptions = append(anthropicOptions, option.WithAPIKey(provider.apiKey))
- }
- if provider.useBedrock {
- anthropicOptions = append(anthropicOptions, bedrock.WithLoadDefaultConfig(context.Background()))
+ 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...))
+ }
}
-
- provider.client = anthropic.NewClient(anthropicOptions...)
- return provider, nil
+ return
}
-func (a *anthropicProvider) SendMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) {
- messages = cleanupMessages(messages)
- anthropicMessages := a.convertToAnthropicMessages(messages)
- anthropicTools := a.convertToAnthropicTools(tools)
-
- response, err := a.client.Messages.New(
- ctx,
- anthropic.MessageNewParams{
- Model: anthropic.Model(a.model.APIModel),
- MaxTokens: a.maxTokens,
- Temperature: anthropic.Float(0),
- Messages: anthropicMessages,
- Tools: anthropicTools,
- System: []anthropic.TextBlockParam{
- {
- Text: a.systemMessage,
- CacheControl: anthropic.CacheControlEphemeralParam{
- Type: "ephemeral",
- },
- },
+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 err != nil {
- return nil, err
- }
+ }
- content := ""
- for _, block := range response.Content {
- if text, ok := block.AsAny().(anthropic.TextBlock); ok {
- content += text.Text
+ if i == len(tools)-1 && !a.options.disableCache {
+ toolParam.CacheControl = anthropic.CacheControlEphemeralParam{
+ Type: "ephemeral",
+ }
}
- }
- toolCalls := a.extractToolCalls(response.Content)
- tokenUsage := a.extractTokenUsage(response.Usage)
+ anthropicTools[i] = anthropic.ToolUnionParam{OfTool: &toolParam}
+ }
- return &ProviderResponse{
- Content: content,
- ToolCalls: toolCalls,
- Usage: tokenUsage,
- }, nil
+ return anthropicTools
}
-func (a *anthropicProvider) StreamResponse(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (<-chan ProviderEvent, error) {
- messages = cleanupMessages(messages)
- anthropicMessages := a.convertToAnthropicMessages(messages)
- anthropicTools := a.convertToAnthropicTools(tools)
+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 lastMessage.Role == message.User && strings.Contains(strings.ToLower(lastMessage.Content().String()), "think") {
- thinkingParam = anthropic.ThinkingConfigParamUnion{
- OfThinkingConfigEnabled: &anthropic.ThinkingConfigEnabledParam{
- BudgetTokens: int64(float64(a.maxTokens) * 0.8),
- Type: "enabled",
- },
+ 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)
}
- temperature = anthropic.Float(1)
}
- eventChan := make(chan ProviderEvent)
+ 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",
+ },
+ },
+ },
+ }
+}
- go func() {
- defer close(eventChan)
+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)
+ logging.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 {
+ retry, after, retryErr := a.shouldRetry(attempts, err)
+ if retryErr != nil {
+ return nil, retryErr
+ }
+ if retry {
+ logging.WarnPersist("Retrying due to rate limit... attempt %d of %d", logging.PersistTimeArg, time.Millisecond*time.Duration(after+100))
+ select {
+ case <-ctx.Done():
+ return nil, ctx.Err()
+ case <-time.After(time.Duration(after) * time.Millisecond):
+ continue
+ }
+ }
+ return nil, retryErr
+ }
- const maxRetries = 8
- attempts := 0
+ content := ""
+ for _, block := range anthropicResponse.Content {
+ if text, ok := block.AsAny().(anthropic.TextBlock); ok {
+ content += text.Text
+ }
+ }
- for {
+ 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)
+ logging.Debug("Prepared messages", "messages", string(jsonData))
+ }
+ attempts := 0
+ eventChan := make(chan ProviderEvent)
+ go func() {
+ for {
attempts++
-
- stream := a.client.Messages.NewStreaming(
+ anthropicStream := a.client.Messages.NewStreaming(
ctx,
- anthropic.MessageNewParams{
- Model: anthropic.Model(a.model.APIModel),
- MaxTokens: a.maxTokens,
- Temperature: temperature,
- Messages: anthropicMessages,
- Tools: anthropicTools,
- Thinking: thinkingParam,
- System: []anthropic.TextBlockParam{
- {
- Text: a.systemMessage,
- CacheControl: anthropic.CacheControlEphemeralParam{
- Type: "ephemeral",
- },
- },
- },
- },
+ preparedMessages,
)
-
accumulatedMessage := anthropic.Message{}
- for stream.Next() {
- event := stream.Current()
+ for anthropicStream.Next() {
+ event := anthropicStream.Current()
err := accumulatedMessage.Accumulate(event)
if err != nil {
eventChan <- ProviderEvent{Type: EventError, Error: err}
- return // Don't retry on accumulation errors
+ continue
}
switch event := event.AsAny().(type) {
@@ -211,6 +281,7 @@ func (a *anthropicProvider) StreamResponse(ctx context.Context, messages []messa
Content: event.Delta.Text,
}
}
+ // TODO: check if we can somehow stream tool calls
case anthropic.ContentBlockStopEvent:
eventChan <- ProviderEvent{Type: EventContentStop}
@@ -223,84 +294,87 @@ func (a *anthropicProvider) StreamResponse(ctx context.Context, messages []messa
}
}
- toolCalls := a.extractToolCalls(accumulatedMessage.Content)
- tokenUsage := a.extractTokenUsage(accumulatedMessage.Usage)
-
eventChan <- ProviderEvent{
Type: EventComplete,
Response: &ProviderResponse{
Content: content,
- ToolCalls: toolCalls,
- Usage: tokenUsage,
- FinishReason: string(accumulatedMessage.StopReason),
+ ToolCalls: a.toolCalls(accumulatedMessage),
+ Usage: a.usage(accumulatedMessage),
+ FinishReason: a.finishReason(string(accumulatedMessage.StopReason)),
},
}
}
}
- err := stream.Err()
+ err := anthropicStream.Err()
if err == nil || errors.Is(err, io.EOF) {
+ close(eventChan)
return
}
-
- var apierr *anthropic.Error
- if !errors.As(err, &apierr) {
- eventChan <- ProviderEvent{Type: EventError, Error: err}
- return
- }
-
- if apierr.StatusCode != 429 && apierr.StatusCode != 529 {
- eventChan <- ProviderEvent{Type: EventError, Error: err}
+ // If there is an error we are going to see if we can retry the call
+ retry, after, retryErr := a.shouldRetry(attempts, err)
+ if retryErr != nil {
+ eventChan <- ProviderEvent{Type: EventError, Error: retryErr}
+ close(eventChan)
return
}
-
- if attempts > maxRetries {
- eventChan <- ProviderEvent{
- Type: EventError,
- Error: errors.New("maximum retry attempts reached for rate limit (429)"),
- }
- return
- }
-
- retryMs := 0
- retryAfterValues := apierr.Response.Header.Values("Retry-After")
- if len(retryAfterValues) > 0 {
- var retryAfterSec int
- if _, err := fmt.Sscanf(retryAfterValues[0], "%d", &retryAfterSec); err == nil {
- retryMs = retryAfterSec * 1000
- eventChan <- ProviderEvent{
- Type: EventWarning,
- Info: fmt.Sprintf("[Rate limited: waiting %d seconds as specified by API]", retryAfterSec),
+ if retry {
+ logging.WarnPersist("Retrying due to rate limit... attempt %d of %d", logging.PersistTimeArg, time.Millisecond*time.Duration(after+100))
+ select {
+ case <-ctx.Done():
+ // context cancelled
+ if ctx.Err() != nil {
+ eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()}
}
+ close(eventChan)
+ return
+ case <-time.After(time.Duration(after) * time.Millisecond):
+ continue
}
- } else {
- eventChan <- ProviderEvent{
- Type: EventWarning,
- Info: fmt.Sprintf("[Retrying due to rate limit... attempt %d of %d]", attempts, maxRetries),
- }
-
- backoffMs := 2000 * (1 << (attempts - 1))
- jitterMs := int(float64(backoffMs) * 0.2)
- retryMs = backoffMs + jitterMs
}
- select {
- case <-ctx.Done():
+ if ctx.Err() != nil {
eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()}
- return
- case <-time.After(time.Duration(retryMs) * time.Millisecond):
- continue
}
+ close(eventChan)
+ return
}
}()
+ return eventChan
+}
- return eventChan, nil
+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 *anthropicProvider) extractToolCalls(content []anthropic.ContentBlockUnion) []message.ToolCall {
+func (a *anthropicClient) toolCalls(msg anthropic.Message) []message.ToolCall {
var toolCalls []message.ToolCall
- for _, block := range content {
+ for _, block := range msg.Content {
switch variant := block.AsAny().(type) {
case anthropic.ToolUseBlock:
toolCall := message.ToolCall{
@@ -316,90 +390,33 @@ func (a *anthropicProvider) extractToolCalls(content []anthropic.ContentBlockUni
return toolCalls
}
-func (a *anthropicProvider) extractTokenUsage(usage anthropic.Usage) TokenUsage {
+func (a *anthropicClient) usage(msg anthropic.Message) TokenUsage {
return TokenUsage{
- InputTokens: usage.InputTokens,
- OutputTokens: usage.OutputTokens,
- CacheCreationTokens: usage.CacheCreationInputTokens,
- CacheReadTokens: usage.CacheReadInputTokens,
+ InputTokens: msg.Usage.InputTokens,
+ OutputTokens: msg.Usage.OutputTokens,
+ CacheCreationTokens: msg.Usage.CacheCreationInputTokens,
+ CacheReadTokens: msg.Usage.CacheReadInputTokens,
}
}
-func (a *anthropicProvider) convertToAnthropicTools(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,
- },
- }
-
- if i == len(tools)-1 && !a.disableCache {
- toolParam.CacheControl = anthropic.CacheControlEphemeralParam{
- Type: "ephemeral",
- }
- }
-
- anthropicTools[i] = anthropic.ToolUnionParam{OfTool: &toolParam}
+func WithAnthropicBedrock(useBedrock bool) AnthropicOption {
+ return func(options *anthropicOptions) {
+ options.useBedrock = useBedrock
}
-
- return anthropicTools
}
-func (a *anthropicProvider) convertToAnthropicMessages(messages []message.Message) []anthropic.MessageParam {
- anthropicMessages := make([]anthropic.MessageParam, 0, len(messages))
- cachedBlocks := 0
-
- for _, msg := range messages {
- switch msg.Role {
- case message.User:
- content := anthropic.NewTextBlock(msg.Content().String())
- if cachedBlocks < 2 && !a.disableCache {
- content.OfRequestTextBlock.CacheControl = anthropic.CacheControlEphemeralParam{
- Type: "ephemeral",
- }
- cachedBlocks++
- }
- anthropicMessages = append(anthropicMessages, anthropic.NewUserMessage(content))
-
- case message.Assistant:
- blocks := []anthropic.ContentBlockParamUnion{}
- if msg.Content().String() != "" {
- content := anthropic.NewTextBlock(msg.Content().String())
- if cachedBlocks < 2 && !a.disableCache {
- content.OfRequestTextBlock.CacheControl = anthropic.CacheControlEphemeralParam{
- Type: "ephemeral",
- }
- cachedBlocks++
- }
- blocks = append(blocks, content)
- }
-
- 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))
- }
+func WithAnthropicDisableCache() AnthropicOption {
+ return func(options *anthropicOptions) {
+ options.disableCache = true
+ }
+}
- if len(blocks) > 0 {
- anthropicMessages = append(anthropicMessages, anthropic.NewAssistantMessage(blocks...))
- }
+func DefaultShouldThinkFn(s string) bool {
+ return strings.Contains(strings.ToLower(s), "think")
+}
- 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...))
- }
+func WithAnthropicShouldThinkFn(fn func(string) bool) AnthropicOption {
+ return func(options *anthropicOptions) {
+ options.shouldThink = fn
}
-
- return anthropicMessages
}
diff --git a/internal/llm/provider/bedrock.go b/internal/llm/provider/bedrock.go
index 677f4676b..d76925ad1 100644
--- a/internal/llm/provider/bedrock.go
+++ b/internal/llm/provider/bedrock.go
@@ -7,33 +7,29 @@ import (
"os"
"strings"
- "github.com/kujtimiihoxha/termai/internal/llm/models"
"github.com/kujtimiihoxha/termai/internal/llm/tools"
"github.com/kujtimiihoxha/termai/internal/message"
)
-type bedrockProvider struct {
- childProvider Provider
- model models.Model
- maxTokens int64
- systemMessage string
+type bedrockOptions struct {
+ // Bedrock specific options can be added here
}
-func (b *bedrockProvider) SendMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) {
- return b.childProvider.SendMessages(ctx, messages, tools)
-}
+type BedrockOption func(*bedrockOptions)
-func (b *bedrockProvider) StreamResponse(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (<-chan ProviderEvent, error) {
- return b.childProvider.StreamResponse(ctx, messages, tools)
+type bedrockClient struct {
+ providerOptions providerClientOptions
+ options bedrockOptions
+ childProvider ProviderClient
}
-func NewBedrockProvider(opts ...BedrockOption) (Provider, error) {
- provider := &bedrockProvider{}
- for _, opt := range opts {
- opt(provider)
- }
+type BedrockClient ProviderClient
+
+func newBedrockClient(opts providerClientOptions) BedrockClient {
+ bedrockOpts := bedrockOptions{}
+ // Apply bedrock specific options if they are added in the future
- // based on the AWS region prefix the model name with, us, eu, ap, sa, etc.
+ // Get AWS region from environment
region := os.Getenv("AWS_REGION")
if region == "" {
region = os.Getenv("AWS_DEFAULT_REGION")
@@ -43,45 +39,62 @@ func NewBedrockProvider(opts ...BedrockOption) (Provider, error) {
region = "us-east-1" // default region
}
if len(region) < 2 {
- return nil, errors.New("AWS_REGION or AWS_DEFAULT_REGION environment variable is invalid")
+ return &bedrockClient{
+ providerOptions: opts,
+ options: bedrockOpts,
+ childProvider: nil, // Will cause an error when used
+ }
}
+
+ // Prefix the model name with region
regionPrefix := region[:2]
- provider.model.APIModel = fmt.Sprintf("%s.%s", regionPrefix, provider.model.APIModel)
+ modelName := opts.model.APIModel
+ opts.model.APIModel = fmt.Sprintf("%s.%s", regionPrefix, modelName)
- if strings.Contains(string(provider.model.APIModel), "anthropic") {
- anthropic, err := NewAnthropicProvider(
- WithAnthropicModel(provider.model),
- WithAnthropicMaxTokens(provider.maxTokens),
- WithAnthropicSystemMessage(provider.systemMessage),
- WithAnthropicBedrock(),
+ // 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(),
)
- provider.childProvider = anthropic
- if err != nil {
- return nil, err
+ return &bedrockClient{
+ providerOptions: opts,
+ options: bedrockOpts,
+ childProvider: newAnthropicClient(anthropicOpts),
}
- } else {
- return nil, errors.New("unsupported model for bedrock provider")
}
- return provider, nil
-}
-
-type BedrockOption func(*bedrockProvider)
-func WithBedrockSystemMessage(message string) BedrockOption {
- return func(a *bedrockProvider) {
- a.systemMessage = message
+ // 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 WithBedrockMaxTokens(maxTokens int64) BedrockOption {
- return func(a *bedrockProvider) {
- a.maxTokens = maxTokens
+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 WithBedrockModel(model models.Model) BedrockOption {
- return func(a *bedrockProvider) {
- a.model = model
+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)
+} \ No newline at end of file
diff --git a/internal/llm/provider/gemini.go b/internal/llm/provider/gemini.go
index 2d1db2b64..804baea28 100644
--- a/internal/llm/provider/gemini.go
+++ b/internal/llm/provider/gemini.go
@@ -4,80 +4,68 @@ import (
"context"
"encoding/json"
"errors"
+ "fmt"
+ "io"
+ "strings"
+ "time"
"github.com/google/generative-ai-go/genai"
"github.com/google/uuid"
- "github.com/kujtimiihoxha/termai/internal/llm/models"
+ "github.com/kujtimiihoxha/termai/internal/config"
"github.com/kujtimiihoxha/termai/internal/llm/tools"
+ "github.com/kujtimiihoxha/termai/internal/logging"
"github.com/kujtimiihoxha/termai/internal/message"
"google.golang.org/api/iterator"
"google.golang.org/api/option"
)
-type geminiProvider struct {
- client *genai.Client
- model models.Model
- maxTokens int32
- apiKey string
- systemMessage string
+type geminiOptions struct {
+ disableCache bool
}
-type GeminiOption func(*geminiProvider)
+type GeminiOption func(*geminiOptions)
-func NewGeminiProvider(ctx context.Context, opts ...GeminiOption) (Provider, error) {
- provider := &geminiProvider{
- maxTokens: 5000,
- }
+type geminiClient struct {
+ providerOptions providerClientOptions
+ options geminiOptions
+ client *genai.Client
+}
- for _, opt := range opts {
- opt(provider)
- }
+type GeminiClient ProviderClient
- if provider.systemMessage == "" {
- return nil, errors.New("system message is required")
+func newGeminiClient(opts providerClientOptions) GeminiClient {
+ geminiOpts := geminiOptions{}
+ for _, o := range opts.geminiOptions {
+ o(&geminiOpts)
}
- client, err := genai.NewClient(ctx, option.WithAPIKey(provider.apiKey))
+ client, err := genai.NewClient(context.Background(), option.WithAPIKey(opts.apiKey))
if err != nil {
- return nil, err
- }
- provider.client = client
-
- return provider, nil
-}
-
-func WithGeminiSystemMessage(message string) GeminiOption {
- return func(p *geminiProvider) {
- p.systemMessage = message
+ logging.Error("Failed to create Gemini client", "error", err)
+ return nil
}
-}
-func WithGeminiMaxTokens(maxTokens int32) GeminiOption {
- return func(p *geminiProvider) {
- p.maxTokens = maxTokens
+ return &geminiClient{
+ providerOptions: opts,
+ options: geminiOpts,
+ client: client,
}
}
-func WithGeminiModel(model models.Model) GeminiOption {
- return func(p *geminiProvider) {
- p.model = model
- }
-}
-
-func WithGeminiKey(apiKey string) GeminiOption {
- return func(p *geminiProvider) {
- p.apiKey = apiKey
- }
-}
+func (g *geminiClient) convertMessages(messages []message.Message) []*genai.Content {
+ var history []*genai.Content
-func (p *geminiProvider) Close() {
- if p.client != nil {
- p.client.Close()
- }
-}
+ // Add system message first
+ history = append(history, &genai.Content{
+ Parts: []genai.Part{genai.Text(g.providerOptions.systemMessage)},
+ Role: "user",
+ })
-func (p *geminiProvider) convertToGeminiHistory(messages []message.Message) []*genai.Content {
- var history []*genai.Content
+ // Add a system response to acknowledge the system message
+ history = append(history, &genai.Content{
+ Parts: []genai.Part{genai.Text("I'll help you with that.")},
+ Role: "model",
+ })
for _, msg := range messages {
switch msg.Role {
@@ -86,6 +74,7 @@ func (p *geminiProvider) convertToGeminiHistory(messages []message.Message) []*g
Parts: []genai.Part{genai.Text(msg.Content().String())},
Role: "user",
})
+
case message.Assistant:
content := &genai.Content{
Role: "model",
@@ -107,6 +96,7 @@ func (p *geminiProvider) convertToGeminiHistory(messages []message.Message) []*g
}
history = append(history, content)
+
case message.Tool:
for _, result := range msg.ToolResults() {
response := map[string]interface{}{"result": result.Content}
@@ -114,10 +104,11 @@ func (p *geminiProvider) convertToGeminiHistory(messages []message.Message) []*g
if err == nil {
response = parsed
}
+
var toolCall message.ToolCall
- for _, msg := range messages {
- if msg.Role == message.Assistant {
- for _, call := range msg.ToolCalls() {
+ for _, m := range messages {
+ if m.Role == message.Assistant {
+ for _, call := range m.ToolCalls() {
if call.ID == result.ToolCallID {
toolCall = call
break
@@ -140,186 +131,358 @@ func (p *geminiProvider) convertToGeminiHistory(messages []message.Message) []*g
return history
}
-func (p *geminiProvider) extractTokenUsage(resp *genai.GenerateContentResponse) TokenUsage {
- if resp == nil || resp.UsageMetadata == nil {
- return TokenUsage{}
- }
+func (g *geminiClient) convertTools(tools []tools.BaseTool) []*genai.Tool {
+ geminiTools := make([]*genai.Tool, 0, len(tools))
- return TokenUsage{
- InputTokens: int64(resp.UsageMetadata.PromptTokenCount),
- OutputTokens: int64(resp.UsageMetadata.CandidatesTokenCount),
- CacheCreationTokens: 0, // Not directly provided by Gemini
- CacheReadTokens: int64(resp.UsageMetadata.CachedContentTokenCount),
+ 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,
+ },
+ }
+
+ geminiTools = append(geminiTools, &genai.Tool{
+ FunctionDeclarations: []*genai.FunctionDeclaration{declaration},
+ })
}
+
+ return geminiTools
}
-func (p *geminiProvider) SendMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) {
- messages = cleanupMessages(messages)
- model := p.client.GenerativeModel(p.model.APIModel)
- model.SetMaxOutputTokens(p.maxTokens)
+func (g *geminiClient) finishReason(reason genai.FinishReason) message.FinishReason {
+ reasonStr := reason.String()
+ switch {
+ case reasonStr == "STOP":
+ return message.FinishReasonEndTurn
+ case reasonStr == "MAX_TOKENS":
+ return message.FinishReasonMaxTokens
+ case strings.Contains(reasonStr, "FUNCTION") || strings.Contains(reasonStr, "TOOL"):
+ return message.FinishReasonToolUse
+ default:
+ return message.FinishReasonUnknown
+ }
+}
- model.SystemInstruction = genai.NewUserContent(genai.Text(p.systemMessage))
+func (g *geminiClient) send(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) {
+ model := g.client.GenerativeModel(g.providerOptions.model.APIModel)
+ model.SetMaxOutputTokens(int32(g.providerOptions.maxTokens))
+ // Convert tools
if len(tools) > 0 {
- declarations := p.convertToolsToGeminiFunctionDeclarations(tools)
- for _, declaration := range declarations {
- model.Tools = append(model.Tools, &genai.Tool{FunctionDeclarations: []*genai.FunctionDeclaration{declaration}})
- }
+ model.Tools = g.convertTools(tools)
}
- chat := model.StartChat()
- chat.History = p.convertToGeminiHistory(messages[:len(messages)-1]) // Exclude last message
+ // Convert messages
+ geminiMessages := g.convertMessages(messages)
- lastUserMsg := messages[len(messages)-1]
- resp, err := chat.SendMessage(ctx, genai.Text(lastUserMsg.Content().String()))
- if err != nil {
- return nil, err
+ cfg := config.Get()
+ if cfg.Debug {
+ jsonData, _ := json.Marshal(geminiMessages)
+ logging.Debug("Prepared messages", "messages", string(jsonData))
}
- var content string
- var toolCalls []message.ToolCall
+ attempts := 0
+ for {
+ attempts++
+ chat := model.StartChat()
+ chat.History = geminiMessages[:len(geminiMessages)-1] // All but last message
+
+ lastMsg := geminiMessages[len(geminiMessages)-1]
+ var lastText string
+ for _, part := range lastMsg.Parts {
+ if text, ok := part.(genai.Text); ok {
+ lastText = string(text)
+ break
+ }
+ }
- if len(resp.Candidates) > 0 && resp.Candidates[0].Content != nil {
- for _, part := range resp.Candidates[0].Content.Parts {
- switch p := part.(type) {
- case genai.Text:
- content = string(p)
- case genai.FunctionCall:
- id := "call_" + uuid.New().String()
- args, _ := json.Marshal(p.Args)
- toolCalls = append(toolCalls, message.ToolCall{
- ID: id,
- Name: p.Name,
- Input: string(args),
- Type: "function",
- })
+ resp, err := chat.SendMessage(ctx, genai.Text(lastText))
+ // 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)
+ if retryErr != nil {
+ return nil, retryErr
}
+ if retry {
+ logging.WarnPersist("Retrying due to rate limit... attempt %d of %d", logging.PersistTimeArg, time.Millisecond*time.Duration(after+100))
+ select {
+ case <-ctx.Done():
+ return nil, ctx.Err()
+ case <-time.After(time.Duration(after) * time.Millisecond):
+ continue
+ }
+ }
+ return nil, retryErr
}
- }
- tokenUsage := p.extractTokenUsage(resp)
+ content := ""
+ var toolCalls []message.ToolCall
+
+ if len(resp.Candidates) > 0 && resp.Candidates[0].Content != nil {
+ for _, part := range resp.Candidates[0].Content.Parts {
+ switch p := part.(type) {
+ case genai.Text:
+ content = string(p)
+ case genai.FunctionCall:
+ id := "call_" + uuid.New().String()
+ args, _ := json.Marshal(p.Args)
+ toolCalls = append(toolCalls, message.ToolCall{
+ ID: id,
+ Name: p.Name,
+ Input: string(args),
+ Type: "function",
+ })
+ }
+ }
+ }
- return &ProviderResponse{
- Content: content,
- ToolCalls: toolCalls,
- Usage: tokenUsage,
- }, nil
+ return &ProviderResponse{
+ Content: content,
+ ToolCalls: toolCalls,
+ Usage: g.usage(resp),
+ FinishReason: g.finishReason(resp.Candidates[0].FinishReason),
+ }, nil
+ }
}
-func (p *geminiProvider) StreamResponse(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (<-chan ProviderEvent, error) {
- messages = cleanupMessages(messages)
- model := p.client.GenerativeModel(p.model.APIModel)
- model.SetMaxOutputTokens(p.maxTokens)
-
- model.SystemInstruction = genai.NewUserContent(genai.Text(p.systemMessage))
+func (g *geminiClient) stream(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent {
+ model := g.client.GenerativeModel(g.providerOptions.model.APIModel)
+ model.SetMaxOutputTokens(int32(g.providerOptions.maxTokens))
+ // Convert tools
if len(tools) > 0 {
- declarations := p.convertToolsToGeminiFunctionDeclarations(tools)
- for _, declaration := range declarations {
- model.Tools = append(model.Tools, &genai.Tool{FunctionDeclarations: []*genai.FunctionDeclaration{declaration}})
- }
+ model.Tools = g.convertTools(tools)
}
- chat := model.StartChat()
- chat.History = p.convertToGeminiHistory(messages[:len(messages)-1]) // Exclude last message
+ // Convert messages
+ geminiMessages := g.convertMessages(messages)
- lastUserMsg := messages[len(messages)-1]
-
- iter := chat.SendMessageStream(ctx, genai.Text(lastUserMsg.Content().String()))
+ cfg := config.Get()
+ if cfg.Debug {
+ jsonData, _ := json.Marshal(geminiMessages)
+ logging.Debug("Prepared messages", "messages", string(jsonData))
+ }
+ attempts := 0
eventChan := make(chan ProviderEvent)
go func() {
defer close(eventChan)
- var finalResp *genai.GenerateContentResponse
- currentContent := ""
- toolCalls := []message.ToolCall{}
-
for {
- resp, err := iter.Next()
- if err == iterator.Done {
- break
- }
- if err != nil {
- eventChan <- ProviderEvent{
- Type: EventError,
- Error: err,
+ attempts++
+ chat := model.StartChat()
+ chat.History = geminiMessages[:len(geminiMessages)-1] // All but last message
+
+ lastMsg := geminiMessages[len(geminiMessages)-1]
+ var lastText string
+ for _, part := range lastMsg.Parts {
+ if text, ok := part.(genai.Text); ok {
+ lastText = string(text)
+ break
}
- return
}
- finalResp = resp
+ iter := chat.SendMessageStream(ctx, genai.Text(lastText))
- if len(resp.Candidates) > 0 && resp.Candidates[0].Content != nil {
- for _, part := range resp.Candidates[0].Content.Parts {
- switch p := part.(type) {
- case genai.Text:
- newText := string(p)
- eventChan <- ProviderEvent{
- Type: EventContentDelta,
- Content: newText,
- }
- currentContent += newText
- case genai.FunctionCall:
- id := "call_" + uuid.New().String()
- args, _ := json.Marshal(p.Args)
- newCall := message.ToolCall{
- ID: id,
- Name: p.Name,
- Input: string(args),
- Type: "function",
- }
+ currentContent := ""
+ toolCalls := []message.ToolCall{}
+ var finalResp *genai.GenerateContentResponse
- isNew := true
- for _, existing := range toolCalls {
- if existing.Name == newCall.Name && existing.Input == newCall.Input {
- isNew = false
- break
+ eventChan <- ProviderEvent{Type: EventContentStart}
+
+ for {
+ resp, err := iter.Next()
+ if err == iterator.Done {
+ break
+ }
+ if err != nil {
+ retry, after, retryErr := g.shouldRetry(attempts, err)
+ if retryErr != nil {
+ eventChan <- ProviderEvent{Type: EventError, Error: retryErr}
+ return
+ }
+ if retry {
+ logging.WarnPersist("Retrying due to rate limit... attempt %d of %d", logging.PersistTimeArg, time.Millisecond*time.Duration(after+100))
+ select {
+ case <-ctx.Done():
+ if ctx.Err() != nil {
+ eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()}
}
+
+ return
+ case <-time.After(time.Duration(after) * time.Millisecond):
+ 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 p := part.(type) {
+ case genai.Text:
+ newText := string(p)
+ delta := newText[len(currentContent):]
+ if delta != "" {
+ eventChan <- ProviderEvent{
+ Type: EventContentDelta,
+ Content: delta,
+ }
+ currentContent = newText
+ }
+ case genai.FunctionCall:
+ id := "call_" + uuid.New().String()
+ args, _ := json.Marshal(p.Args)
+ newCall := message.ToolCall{
+ ID: id,
+ Name: p.Name,
+ Input: string(args),
+ Type: "function",
+ }
- if isNew {
- toolCalls = append(toolCalls, newCall)
+ isNew := true
+ for _, existing := range toolCalls {
+ if existing.Name == newCall.Name && existing.Input == newCall.Input {
+ isNew = false
+ break
+ }
+ }
+
+ if isNew {
+ toolCalls = append(toolCalls, newCall)
+ }
}
}
}
}
- }
- tokenUsage := p.extractTokenUsage(finalResp)
+ eventChan <- ProviderEvent{Type: EventContentStop}
- eventChan <- ProviderEvent{
- Type: EventComplete,
- Response: &ProviderResponse{
- Content: currentContent,
- ToolCalls: toolCalls,
- Usage: tokenUsage,
- FinishReason: string(finalResp.Candidates[0].FinishReason.String()),
- },
+ if finalResp != nil {
+ eventChan <- ProviderEvent{
+ Type: EventComplete,
+ Response: &ProviderResponse{
+ Content: currentContent,
+ ToolCalls: toolCalls,
+ Usage: g.usage(finalResp),
+ FinishReason: g.finishReason(finalResp.Candidates[0].FinishReason),
+ },
+ }
+ return
+ }
+
+ // If we get here, we need to retry
+ if attempts > maxRetries {
+ eventChan <- ProviderEvent{
+ Type: EventError,
+ Error: fmt.Errorf("maximum retry attempts reached: %d retries", maxRetries),
+ }
+ return
+ }
+
+ // Wait before retrying
+ select {
+ case <-ctx.Done():
+ if ctx.Err() != nil {
+ eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()}
+ }
+ return
+ case <-time.After(time.Duration(2000*(1<<(attempts-1))) * time.Millisecond):
+ continue
+ }
}
}()
- return eventChan, nil
+ return eventChan
}
-func (p *geminiProvider) convertToolsToGeminiFunctionDeclarations(tools []tools.BaseTool) []*genai.FunctionDeclaration {
- declarations := make([]*genai.FunctionDeclaration, len(tools))
+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)
+ }
- for i, tool := range tools {
- info := tool.Info()
- declarations[i] = &genai.FunctionDeclaration{
- Name: info.Name,
- Description: info.Description,
- Parameters: &genai.Schema{
- Type: genai.TypeObject,
- Properties: convertSchemaProperties(info.Parameters),
- Required: info.Required,
- },
+ // 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 funcCall, ok := part.(genai.FunctionCall); ok {
+ id := "call_" + uuid.New().String()
+ args, _ := json.Marshal(funcCall.Args)
+ toolCalls = append(toolCalls, message.ToolCall{
+ ID: id,
+ Name: funcCall.Name,
+ Input: string(args),
+ Type: "function",
+ })
+ }
}
}
- return declarations
+ 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 {
@@ -396,8 +559,12 @@ func mapJSONTypeToGenAI(jsonType string) genai.Type {
}
}
-func parseJsonToMap(jsonStr string) (map[string]interface{}, error) {
- var result map[string]interface{}
- err := json.Unmarshal([]byte(jsonStr), &result)
- return result, err
+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
index dbfde3fa8..9c2ad2012 100644
--- a/internal/llm/provider/openai.go
+++ b/internal/llm/provider/openai.go
@@ -2,89 +2,65 @@ package provider
import (
"context"
+ "encoding/json"
"errors"
+ "fmt"
+ "io"
+ "time"
- "github.com/kujtimiihoxha/termai/internal/llm/models"
+ "github.com/kujtimiihoxha/termai/internal/config"
"github.com/kujtimiihoxha/termai/internal/llm/tools"
+ "github.com/kujtimiihoxha/termai/internal/logging"
"github.com/kujtimiihoxha/termai/internal/message"
"github.com/openai/openai-go"
"github.com/openai/openai-go/option"
)
-type openaiProvider struct {
- client openai.Client
- model models.Model
- maxTokens int64
- baseURL string
- apiKey string
- systemMessage string
+type openaiOptions struct {
+ baseURL string
+ disableCache bool
}
-type OpenAIOption func(*openaiProvider)
+type OpenAIOption func(*openaiOptions)
-func NewOpenAIProvider(opts ...OpenAIOption) (Provider, error) {
- provider := &openaiProvider{
- maxTokens: 5000,
- }
-
- for _, opt := range opts {
- opt(provider)
- }
-
- clientOpts := []option.RequestOption{
- option.WithAPIKey(provider.apiKey),
- }
- if provider.baseURL != "" {
- clientOpts = append(clientOpts, option.WithBaseURL(provider.baseURL))
- }
-
- provider.client = openai.NewClient(clientOpts...)
- if provider.systemMessage == "" {
- return nil, errors.New("system message is required")
- }
-
- return provider, nil
+type openaiClient struct {
+ providerOptions providerClientOptions
+ options openaiOptions
+ client openai.Client
}
-func WithOpenAISystemMessage(message string) OpenAIOption {
- return func(p *openaiProvider) {
- p.systemMessage = message
- }
-}
+type OpenAIClient ProviderClient
-func WithOpenAIMaxTokens(maxTokens int64) OpenAIOption {
- return func(p *openaiProvider) {
- p.maxTokens = maxTokens
+func newOpenAIClient(opts providerClientOptions) OpenAIClient {
+ openaiOpts := openaiOptions{}
+ for _, o := range opts.openaiOptions {
+ o(&openaiOpts)
}
-}
-func WithOpenAIModel(model models.Model) OpenAIOption {
- return func(p *openaiProvider) {
- p.model = model
+ openaiClientOptions := []option.RequestOption{}
+ if opts.apiKey != "" {
+ openaiClientOptions = append(openaiClientOptions, option.WithAPIKey(opts.apiKey))
}
-}
-
-func WithOpenAIBaseURL(baseURL string) OpenAIOption {
- return func(p *openaiProvider) {
- p.baseURL = baseURL
+ if openaiOpts.baseURL != "" {
+ openaiClientOptions = append(openaiClientOptions, option.WithBaseURL(openaiOpts.baseURL))
}
-}
-func WithOpenAIKey(apiKey string) OpenAIOption {
- return func(p *openaiProvider) {
- p.apiKey = apiKey
+ client := openai.NewClient(openaiClientOptions...)
+ return &openaiClient{
+ providerOptions: opts,
+ options: openaiOpts,
+ client: client,
}
}
-func (p *openaiProvider) convertToOpenAIMessages(messages []message.Message) []openai.ChatCompletionMessageParamUnion {
- var chatMessages []openai.ChatCompletionMessageParamUnion
-
- chatMessages = append(chatMessages, openai.SystemMessage(p.systemMessage))
+func (o *openaiClient) convertMessages(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:
- chatMessages = append(chatMessages, openai.UserMessage(msg.Content().String()))
+ openaiMessages = append(openaiMessages, openai.UserMessage(msg.Content().String()))
case message.Assistant:
assistantMsg := openai.ChatCompletionAssistantMessageParam{
@@ -111,23 +87,23 @@ func (p *openaiProvider) convertToOpenAIMessages(messages []message.Message) []o
}
}
- chatMessages = append(chatMessages, openai.ChatCompletionMessageParamUnion{
+ openaiMessages = append(openaiMessages, openai.ChatCompletionMessageParamUnion{
OfAssistant: &assistantMsg,
})
case message.Tool:
for _, result := range msg.ToolResults() {
- chatMessages = append(chatMessages,
+ openaiMessages = append(openaiMessages,
openai.ToolMessage(result.Content, result.ToolCallID),
)
}
}
}
- return chatMessages
+ return
}
-func (p *openaiProvider) convertToOpenAITools(tools []tools.BaseTool) []openai.ChatCompletionToolParam {
+func (o *openaiClient) convertTools(tools []tools.BaseTool) []openai.ChatCompletionToolParam {
openaiTools := make([]openai.ChatCompletionToolParam, len(tools))
for i, tool := range tools {
@@ -148,133 +124,238 @@ func (p *openaiProvider) convertToOpenAITools(tools []tools.BaseTool) []openai.C
return openaiTools
}
-func (p *openaiProvider) extractTokenUsage(usage openai.CompletionUsage) TokenUsage {
- cachedTokens := int64(0)
-
- cachedTokens = usage.PromptTokensDetails.CachedTokens
- inputTokens := usage.PromptTokens - cachedTokens
-
- return TokenUsage{
- InputTokens: inputTokens,
- OutputTokens: usage.CompletionTokens,
- CacheCreationTokens: 0, // OpenAI doesn't provide this directly
- CacheReadTokens: cachedTokens,
+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 (p *openaiProvider) SendMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) {
- messages = cleanupMessages(messages)
- chatMessages := p.convertToOpenAIMessages(messages)
- openaiTools := p.convertToOpenAITools(tools)
-
- params := openai.ChatCompletionNewParams{
- Model: openai.ChatModel(p.model.APIModel),
- Messages: chatMessages,
- MaxTokens: openai.Int(p.maxTokens),
- Tools: openaiTools,
- }
-
- response, err := p.client.Chat.Completions.New(ctx, params)
- if err != nil {
- return nil, err
+func (o *openaiClient) preparedParams(messages []openai.ChatCompletionMessageParamUnion, tools []openai.ChatCompletionToolParam) openai.ChatCompletionNewParams {
+ return openai.ChatCompletionNewParams{
+ Model: openai.ChatModel(o.providerOptions.model.APIModel),
+ Messages: messages,
+ MaxTokens: openai.Int(o.providerOptions.maxTokens),
+ Tools: tools,
}
+}
- content := ""
- if response.Choices[0].Message.Content != "" {
- content = response.Choices[0].Message.Content
+func (o *openaiClient) send(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (response *ProviderResponse, err error) {
+ params := o.preparedParams(o.convertMessages(messages), o.convertTools(tools))
+ cfg := config.Get()
+ if cfg.Debug {
+ jsonData, _ := json.Marshal(params)
+ logging.Debug("Prepared messages", "messages", string(jsonData))
}
-
- var toolCalls []message.ToolCall
- if len(response.Choices[0].Message.ToolCalls) > 0 {
- toolCalls = make([]message.ToolCall, len(response.Choices[0].Message.ToolCalls))
- for i, call := range response.Choices[0].Message.ToolCalls {
- toolCalls[i] = message.ToolCall{
- ID: call.ID,
- Name: call.Function.Name,
- Input: call.Function.Arguments,
- Type: "function",
+ 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)
+ if retryErr != nil {
+ return nil, retryErr
}
+ if retry {
+ logging.WarnPersist("Retrying due to rate limit... attempt %d of %d", logging.PersistTimeArg, time.Millisecond*time.Duration(after+100))
+ select {
+ case <-ctx.Done():
+ return nil, ctx.Err()
+ case <-time.After(time.Duration(after) * time.Millisecond):
+ continue
+ }
+ }
+ return nil, retryErr
}
- }
- tokenUsage := p.extractTokenUsage(response.Usage)
+ content := ""
+ if openaiResponse.Choices[0].Message.Content != "" {
+ content = openaiResponse.Choices[0].Message.Content
+ }
- return &ProviderResponse{
- Content: content,
- ToolCalls: toolCalls,
- Usage: tokenUsage,
- }, nil
+ return &ProviderResponse{
+ Content: content,
+ ToolCalls: o.toolCalls(*openaiResponse),
+ Usage: o.usage(*openaiResponse),
+ FinishReason: o.finishReason(string(openaiResponse.Choices[0].FinishReason)),
+ }, nil
+ }
}
-func (p *openaiProvider) StreamResponse(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (<-chan ProviderEvent, error) {
- messages = cleanupMessages(messages)
- chatMessages := p.convertToOpenAIMessages(messages)
- openaiTools := p.convertToOpenAITools(tools)
-
- params := openai.ChatCompletionNewParams{
- Model: openai.ChatModel(p.model.APIModel),
- Messages: chatMessages,
- MaxTokens: openai.Int(p.maxTokens),
- Tools: openaiTools,
- StreamOptions: openai.ChatCompletionStreamOptionsParam{
- IncludeUsage: openai.Bool(true),
- },
+func (o *openaiClient) stream(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent {
+ params := o.preparedParams(o.convertMessages(messages), o.convertTools(tools))
+ params.StreamOptions = openai.ChatCompletionStreamOptionsParam{
+ IncludeUsage: openai.Bool(true),
}
- stream := p.client.Chat.Completions.NewStreaming(ctx, params)
+ cfg := config.Get()
+ if cfg.Debug {
+ jsonData, _ := json.Marshal(params)
+ logging.Debug("Prepared messages", "messages", string(jsonData))
+ }
+ attempts := 0
eventChan := make(chan ProviderEvent)
- toolCalls := make([]message.ToolCall, 0)
go func() {
- defer close(eventChan)
-
- acc := openai.ChatCompletionAccumulator{}
- currentContent := ""
-
- for stream.Next() {
- chunk := stream.Current()
- acc.AddChunk(chunk)
-
- if tool, ok := acc.JustFinishedToolCall(); ok {
- toolCalls = append(toolCalls, message.ToolCall{
- ID: tool.Id,
- Name: tool.Name,
- Input: tool.Arguments,
- Type: "function",
- })
- }
+ 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)
+
+ if tool, ok := acc.JustFinishedToolCall(); ok {
+ toolCalls = append(toolCalls, message.ToolCall{
+ ID: tool.Id,
+ Name: tool.Name,
+ Input: tool.Arguments,
+ Type: "function",
+ })
+ }
- for _, choice := range chunk.Choices {
- if choice.Delta.Content != "" {
- eventChan <- ProviderEvent{
- Type: EventContentDelta,
- Content: choice.Delta.Content,
+ for _, choice := range chunk.Choices {
+ if choice.Delta.Content != "" {
+ eventChan <- ProviderEvent{
+ Type: EventContentDelta,
+ Content: choice.Delta.Content,
+ }
+ currentContent += choice.Delta.Content
}
- currentContent += choice.Delta.Content
}
}
- }
- if err := stream.Err(); err != nil {
- eventChan <- ProviderEvent{
- Type: EventError,
- Error: err,
+ err := openaiStream.Err()
+ if err == nil || errors.Is(err, io.EOF) {
+ // Stream completed successfully
+ eventChan <- ProviderEvent{
+ Type: EventComplete,
+ Response: &ProviderResponse{
+ Content: currentContent,
+ ToolCalls: toolCalls,
+ Usage: o.usage(acc.ChatCompletion),
+ FinishReason: o.finishReason(string(acc.ChatCompletion.Choices[0].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)
+ if retryErr != nil {
+ eventChan <- ProviderEvent{Type: EventError, Error: retryErr}
+ close(eventChan)
+ return
+ }
+ if retry {
+ logging.WarnPersist("Retrying due to rate limit... attempt %d of %d", logging.PersistTimeArg, time.Millisecond*time.Duration(after+100))
+ select {
+ case <-ctx.Done():
+ // context cancelled
+ if ctx.Err() == nil {
+ eventChan <- ProviderEvent{Type: EventError, Error: ctx.Err()}
+ }
+ close(eventChan)
+ return
+ case <-time.After(time.Duration(after) * time.Millisecond):
+ continue
+ }
+ }
+ eventChan <- ProviderEvent{Type: EventError, Error: retryErr}
+ close(eventChan)
return
}
+ }()
- tokenUsage := p.extractTokenUsage(acc.Usage)
+ return eventChan
+}
- eventChan <- ProviderEvent{
- Type: EventComplete,
- Response: &ProviderResponse{
- Content: currentContent,
- ToolCalls: toolCalls,
- Usage: tokenUsage,
- },
+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
+}
- return eventChan, nil
+func (o *openaiClient) toolCalls(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",
+ }
+ 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,
+ }
+}
+
+func WithOpenAIBaseURL(baseURL string) OpenAIOption {
+ return func(options *openaiOptions) {
+ options.baseURL = baseURL
+ }
+}
+
+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 34d91f2b7..1a5b3dc8a 100644
--- a/internal/llm/provider/provider.go
+++ b/internal/llm/provider/provider.go
@@ -2,14 +2,17 @@ package provider
import (
"context"
+ "fmt"
+ "github.com/kujtimiihoxha/termai/internal/llm/models"
"github.com/kujtimiihoxha/termai/internal/llm/tools"
"github.com/kujtimiihoxha/termai/internal/message"
)
-// EventType represents the type of streaming event
type EventType string
+const maxRetries = 8
+
const (
EventContentStart EventType = "content_start"
EventContentDelta EventType = "content_delta"
@@ -18,7 +21,6 @@ const (
EventComplete EventType = "complete"
EventError EventType = "error"
EventWarning EventType = "warning"
- EventInfo EventType = "info"
)
type TokenUsage struct {
@@ -32,61 +34,152 @@ type ProviderResponse struct {
Content string
ToolCalls []message.ToolCall
Usage TokenUsage
- FinishReason string
+ FinishReason message.FinishReason
}
type ProviderEvent struct {
- Type EventType
+ Type EventType
+
Content string
Thinking string
- ToolCall *message.ToolCall
- Error error
Response *ProviderResponse
- // Used for giving users info on e.x retry
- Info string
+ 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, error)
+ StreamResponse(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent
+
+ Model() models.Model
+}
+
+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.ProviderMock:
+ // TODO: implement mock client for test
+ panic("not implemented")
+ }
+ return nil, fmt.Errorf("provider not supported: %s", providerName)
}
-func cleanupMessages(messages []message.Message) []message.Message {
- // First pass: filter out canceled messages
- var cleanedMessages []message.Message
+func (p *baseProvider[C]) cleanMessages(messages []message.Message) (cleaned []message.Message) {
for _, msg := range messages {
- if msg.FinishReason() != "canceled" || len(msg.ToolCalls()) > 0 {
- // if there are toolCalls this means we want to return it to the LLM telling it that those tools have been
- // cancelled
- cleanedMessages = append(cleanedMessages, msg)
+ // The message has no content
+ if len(msg.Parts) == 0 {
+ continue
}
+ cleaned = append(cleaned, msg)
}
+ return
+}
- // Second pass: filter out tool messages without a corresponding tool call
- var result []message.Message
- toolMessageIDs := make(map[string]bool)
+func (p *baseProvider[C]) SendMessages(ctx context.Context, messages []message.Message, tools []tools.BaseTool) (*ProviderResponse, error) {
+ messages = p.cleanMessages(messages)
+ return p.client.send(ctx, messages, tools)
+}
- for _, msg := range cleanedMessages {
- if msg.Role == message.Assistant {
- for _, toolCall := range msg.ToolCalls() {
- toolMessageIDs[toolCall.ID] = true // Mark as referenced
- }
- }
+func (p *baseProvider[C]) Model() models.Model {
+ return p.options.model
+}
+
+func (p *baseProvider[C]) StreamResponse(ctx context.Context, messages []message.Message, tools []tools.BaseTool) <-chan ProviderEvent {
+ messages = p.cleanMessages(messages)
+ return p.client.stream(ctx, messages, tools)
+}
+
+func WithAPIKey(apiKey string) ProviderClientOption {
+ return func(options *providerClientOptions) {
+ options.apiKey = apiKey
}
+}
- // Keep only messages that aren't unreferenced tool messages
- for _, msg := range cleanedMessages {
- if msg.Role == message.Tool {
- for _, toolCall := range msg.ToolResults() {
- if referenced, exists := toolMessageIDs[toolCall.ToolCallID]; exists && referenced {
- result = append(result, msg)
- }
- }
- } else {
- result = append(result, msg)
- }
+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
}
- return result
}