summaryrefslogtreecommitdiffhomepage
path: root/internal/llm/provider/gemini.go
diff options
context:
space:
mode:
authorKujtim Hoxha <[email protected]>2025-04-16 20:06:23 +0200
committerKujtim Hoxha <[email protected]>2025-04-21 13:42:00 +0200
commitbbfa60c787f2ec459f1689b9a650ddbec9693ed9 (patch)
treef7f2aa31c460c8cc22ec40cc299c386277152241 /internal/llm/provider/gemini.go
parent76b4065f17b87a63092acfd98c997bab53700b35 (diff)
downloadopencode-bbfa60c787f2ec459f1689b9a650ddbec9693ed9.tar.gz
opencode-bbfa60c787f2ec459f1689b9a650ddbec9693ed9.zip
reimplement agent,provider and add file history
Diffstat (limited to 'internal/llm/provider/gemini.go')
-rw-r--r--internal/llm/provider/gemini.go533
1 files changed, 350 insertions, 183 deletions
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
}
+