feat(provider): add gemini google-generative-ai compatibility
This commit is contained in:
parent
9536f55a05
commit
c1a7296a85
1 changed files with 311 additions and 0 deletions
|
|
@ -70,6 +70,10 @@ func (p *HTTPProvider) Chat(ctx context.Context, messages []Message, tools []Too
|
||||||
return nil, fmt.Errorf("API base not configured")
|
return nil, fmt.Errorf("API base not configured")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if isGoogleGenerativeAI(p.apiMode) {
|
||||||
|
return p.chatWithGoogleGenerativeAI(ctx, messages, tools, model, options)
|
||||||
|
}
|
||||||
|
|
||||||
useResponses := shouldPreferResponses(model, p.apiMode)
|
useResponses := shouldPreferResponses(model, p.apiMode)
|
||||||
if useResponses {
|
if useResponses {
|
||||||
resp, err := p.chatWithResponses(ctx, messages, tools, model, options)
|
resp, err := p.chatWithResponses(ctx, messages, tools, model, options)
|
||||||
|
|
@ -302,6 +306,49 @@ func (p *HTTPProvider) chatWithResponses(ctx context.Context, messages []Message
|
||||||
return parseCodexResponse(&apiResponse), nil
|
return parseCodexResponse(&apiResponse), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (p *HTTPProvider) chatWithGoogleGenerativeAI(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error) {
|
||||||
|
model = normalizeModelForGemini(model)
|
||||||
|
requestBody := buildGeminiRequest(messages, tools, options)
|
||||||
|
|
||||||
|
jsonData, err := json.Marshal(requestBody)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal gemini request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
endpoint := fmt.Sprintf("%s/models/%s:generateContent", strings.TrimRight(p.apiBase, "/"), url.PathEscape(model))
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "POST", endpoint, bytes.NewReader(jsonData))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||||
|
}
|
||||||
|
logger.DebugCF("provider", "HTTP request", map[string]interface{}{
|
||||||
|
"url": req.URL.String(),
|
||||||
|
"method": req.Method,
|
||||||
|
})
|
||||||
|
|
||||||
|
p.applyHeaders(req)
|
||||||
|
|
||||||
|
resp, err := p.httpClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to send request: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
logger.DebugCF("provider", "HTTP response error", map[string]interface{}{
|
||||||
|
"status": resp.StatusCode,
|
||||||
|
"body": utils.Truncate(string(body), 500),
|
||||||
|
})
|
||||||
|
return nil, &httpProviderError{statusCode: resp.StatusCode, body: string(body), url: req.URL.String()}
|
||||||
|
}
|
||||||
|
|
||||||
|
return parseGeminiResponse(body)
|
||||||
|
}
|
||||||
|
|
||||||
func normalizeModelForHTTP(model string) string {
|
func normalizeModelForHTTP(model string) string {
|
||||||
// Strip provider prefix from model name (e.g., moonshot/kimi-k2.5 -> kimi-k2.5)
|
// Strip provider prefix from model name (e.g., moonshot/kimi-k2.5 -> kimi-k2.5)
|
||||||
if idx := strings.Index(model, "/"); idx != -1 {
|
if idx := strings.Index(model, "/"); idx != -1 {
|
||||||
|
|
@ -313,6 +360,255 @@ func normalizeModelForHTTP(model string) string {
|
||||||
return model
|
return model
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func normalizeModelForGemini(model string) string {
|
||||||
|
if idx := strings.Index(model, "/"); idx != -1 {
|
||||||
|
prefix := strings.ToLower(model[:idx])
|
||||||
|
if prefix == "gemini" || prefix == "google" {
|
||||||
|
return model[idx+1:]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildGeminiRequest(messages []Message, tools []ToolDefinition, options map[string]interface{}) map[string]interface{} {
|
||||||
|
requestBody := map[string]interface{}{}
|
||||||
|
contents := make([]map[string]interface{}, 0, len(messages))
|
||||||
|
systemParts := make([]map[string]interface{}, 0, 1)
|
||||||
|
callNameByID := map[string]string{}
|
||||||
|
|
||||||
|
for _, msg := range messages {
|
||||||
|
switch msg.Role {
|
||||||
|
case "system":
|
||||||
|
if msg.Content != "" {
|
||||||
|
systemParts = append(systemParts, map[string]interface{}{"text": msg.Content})
|
||||||
|
}
|
||||||
|
|
||||||
|
case "user":
|
||||||
|
if msg.ToolCallID != "" {
|
||||||
|
name := callNameByID[msg.ToolCallID]
|
||||||
|
if name == "" {
|
||||||
|
name = "tool_result"
|
||||||
|
}
|
||||||
|
contents = append(contents, map[string]interface{}{
|
||||||
|
"role": "user",
|
||||||
|
"parts": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"functionResponse": map[string]interface{}{
|
||||||
|
"name": name,
|
||||||
|
"response": map[string]interface{}{
|
||||||
|
"content": msg.Content,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if msg.Content != "" {
|
||||||
|
contents = append(contents, map[string]interface{}{
|
||||||
|
"role": "user",
|
||||||
|
"parts": []map[string]interface{}{
|
||||||
|
{"text": msg.Content},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
case "assistant":
|
||||||
|
parts := make([]map[string]interface{}, 0, 1+len(msg.ToolCalls))
|
||||||
|
if msg.Content != "" {
|
||||||
|
parts = append(parts, map[string]interface{}{"text": msg.Content})
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range msg.ToolCalls {
|
||||||
|
name := tc.Name
|
||||||
|
if name == "" && tc.Function != nil {
|
||||||
|
name = tc.Function.Name
|
||||||
|
}
|
||||||
|
if name == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
args := tc.Arguments
|
||||||
|
if args == nil {
|
||||||
|
args = map[string]interface{}{}
|
||||||
|
}
|
||||||
|
if tc.ID != "" {
|
||||||
|
callNameByID[tc.ID] = name
|
||||||
|
}
|
||||||
|
|
||||||
|
parts = append(parts, map[string]interface{}{
|
||||||
|
"functionCall": map[string]interface{}{
|
||||||
|
"name": name,
|
||||||
|
"args": args,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(parts) > 0 {
|
||||||
|
contents = append(contents, map[string]interface{}{
|
||||||
|
"role": "model",
|
||||||
|
"parts": parts,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
case "tool":
|
||||||
|
name := callNameByID[msg.ToolCallID]
|
||||||
|
if name == "" {
|
||||||
|
name = "tool"
|
||||||
|
}
|
||||||
|
contents = append(contents, map[string]interface{}{
|
||||||
|
"role": "user",
|
||||||
|
"parts": []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"functionResponse": map[string]interface{}{
|
||||||
|
"name": name,
|
||||||
|
"response": map[string]interface{}{
|
||||||
|
"content": msg.Content,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(contents) > 0 {
|
||||||
|
requestBody["contents"] = contents
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(systemParts) > 0 {
|
||||||
|
requestBody["system_instruction"] = map[string]interface{}{
|
||||||
|
"parts": systemParts,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(tools) > 0 {
|
||||||
|
declarations := make([]map[string]interface{}, 0, len(tools))
|
||||||
|
for _, t := range tools {
|
||||||
|
declarations = append(declarations, map[string]interface{}{
|
||||||
|
"name": t.Function.Name,
|
||||||
|
"description": t.Function.Description,
|
||||||
|
"parameters": t.Function.Parameters,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
requestBody["tools"] = []map[string]interface{}{
|
||||||
|
{"functionDeclarations": declarations},
|
||||||
|
}
|
||||||
|
requestBody["toolConfig"] = map[string]interface{}{
|
||||||
|
"functionCallingConfig": map[string]interface{}{
|
||||||
|
"mode": "AUTO",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
generationConfig := map[string]interface{}{}
|
||||||
|
if maxTokens, ok := options["max_tokens"].(int); ok {
|
||||||
|
generationConfig["maxOutputTokens"] = maxTokens
|
||||||
|
}
|
||||||
|
if temperature, ok := options["temperature"].(float64); ok {
|
||||||
|
generationConfig["temperature"] = temperature
|
||||||
|
}
|
||||||
|
if len(generationConfig) > 0 {
|
||||||
|
requestBody["generationConfig"] = generationConfig
|
||||||
|
}
|
||||||
|
|
||||||
|
return requestBody
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseGeminiResponse(body []byte) (*LLMResponse, error) {
|
||||||
|
var resp struct {
|
||||||
|
Candidates []struct {
|
||||||
|
Content struct {
|
||||||
|
Parts []struct {
|
||||||
|
Text string `json:"text"`
|
||||||
|
FunctionCall *struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Args map[string]interface{} `json:"args"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
} `json:"functionCall"`
|
||||||
|
} `json:"parts"`
|
||||||
|
} `json:"content"`
|
||||||
|
FinishReason string `json:"finishReason"`
|
||||||
|
} `json:"candidates"`
|
||||||
|
UsageMetadata struct {
|
||||||
|
PromptTokenCount int `json:"promptTokenCount"`
|
||||||
|
CandidatesTokenCount int `json:"candidatesTokenCount"`
|
||||||
|
TotalTokenCount int `json:"totalTokenCount"`
|
||||||
|
} `json:"usageMetadata"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(body, &resp); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal gemini response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(resp.Candidates) == 0 {
|
||||||
|
return &LLMResponse{
|
||||||
|
Content: "",
|
||||||
|
FinishReason: "stop",
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
candidate := resp.Candidates[0]
|
||||||
|
var content strings.Builder
|
||||||
|
toolCalls := make([]ToolCall, 0)
|
||||||
|
toolCallCount := 0
|
||||||
|
for _, part := range candidate.Content.Parts {
|
||||||
|
if part.Text != "" {
|
||||||
|
content.WriteString(part.Text)
|
||||||
|
}
|
||||||
|
if part.FunctionCall != nil && part.FunctionCall.Name != "" {
|
||||||
|
toolCallCount++
|
||||||
|
callID := part.FunctionCall.ID
|
||||||
|
if callID == "" {
|
||||||
|
callID = fmt.Sprintf("gemini_call_%d", toolCallCount)
|
||||||
|
}
|
||||||
|
args := part.FunctionCall.Args
|
||||||
|
if args == nil {
|
||||||
|
args = map[string]interface{}{}
|
||||||
|
}
|
||||||
|
toolCalls = append(toolCalls, ToolCall{
|
||||||
|
ID: callID,
|
||||||
|
Name: part.FunctionCall.Name,
|
||||||
|
Arguments: args,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
finishReason := mapGeminiFinishReason(candidate.FinishReason)
|
||||||
|
if len(toolCalls) > 0 {
|
||||||
|
finishReason = "tool_calls"
|
||||||
|
}
|
||||||
|
|
||||||
|
usage := &UsageInfo{
|
||||||
|
PromptTokens: resp.UsageMetadata.PromptTokenCount,
|
||||||
|
CompletionTokens: resp.UsageMetadata.CandidatesTokenCount,
|
||||||
|
TotalTokens: resp.UsageMetadata.TotalTokenCount,
|
||||||
|
}
|
||||||
|
if usage.PromptTokens == 0 && usage.CompletionTokens == 0 && usage.TotalTokens == 0 {
|
||||||
|
usage = nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return &LLMResponse{
|
||||||
|
Content: content.String(),
|
||||||
|
ToolCalls: toolCalls,
|
||||||
|
FinishReason: finishReason,
|
||||||
|
Usage: usage,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func mapGeminiFinishReason(reason string) string {
|
||||||
|
switch strings.ToUpper(reason) {
|
||||||
|
case "MAX_TOKENS":
|
||||||
|
return "length"
|
||||||
|
case "STOP", "":
|
||||||
|
return "stop"
|
||||||
|
case "SAFETY", "PROHIBITED_CONTENT", "RECITATION":
|
||||||
|
return "content_filter"
|
||||||
|
default:
|
||||||
|
return "stop"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func stripTemperature(options map[string]interface{}) map[string]interface{} {
|
func stripTemperature(options map[string]interface{}) map[string]interface{} {
|
||||||
if options == nil {
|
if options == nil {
|
||||||
return nil
|
return nil
|
||||||
|
|
@ -343,6 +639,15 @@ func shouldPreferResponses(model, apiMode string) bool {
|
||||||
return strings.Contains(lower, "gpt-5") || strings.Contains(lower, "codex") || strings.Contains(lower, "o1")
|
return strings.Contains(lower, "gpt-5") || strings.Contains(lower, "codex") || strings.Contains(lower, "o1")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func isGoogleGenerativeAI(apiMode string) bool {
|
||||||
|
switch strings.ToLower(apiMode) {
|
||||||
|
case "google-generative-ai", "google", "gemini":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func shouldFallbackFromResponses(err error) bool {
|
func shouldFallbackFromResponses(err error) bool {
|
||||||
var httpErr *httpProviderError
|
var httpErr *httpProviderError
|
||||||
if errors.As(err, &httpErr) {
|
if errors.As(err, &httpErr) {
|
||||||
|
|
@ -464,6 +769,9 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) {
|
||||||
apiKey = cfg.Providers.Gemini.APIKey
|
apiKey = cfg.Providers.Gemini.APIKey
|
||||||
apiBase = cfg.Providers.Gemini.APIBase
|
apiBase = cfg.Providers.Gemini.APIBase
|
||||||
apiMode = cfg.Providers.Gemini.API
|
apiMode = cfg.Providers.Gemini.API
|
||||||
|
if apiMode == "" {
|
||||||
|
apiMode = "google-generative-ai"
|
||||||
|
}
|
||||||
headers = cfg.Providers.Gemini.Headers
|
headers = cfg.Providers.Gemini.Headers
|
||||||
if apiBase == "" {
|
if apiBase == "" {
|
||||||
apiBase = "https://generativelanguage.googleapis.com/v1beta"
|
apiBase = "https://generativelanguage.googleapis.com/v1beta"
|
||||||
|
|
@ -540,6 +848,9 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) {
|
||||||
apiBase = cfg.Providers.Gemini.APIBase
|
apiBase = cfg.Providers.Gemini.APIBase
|
||||||
proxy = cfg.Providers.Gemini.Proxy
|
proxy = cfg.Providers.Gemini.Proxy
|
||||||
apiMode = cfg.Providers.Gemini.API
|
apiMode = cfg.Providers.Gemini.API
|
||||||
|
if apiMode == "" {
|
||||||
|
apiMode = "google-generative-ai"
|
||||||
|
}
|
||||||
headers = cfg.Providers.Gemini.Headers
|
headers = cfg.Providers.Gemini.Headers
|
||||||
if apiBase == "" {
|
if apiBase == "" {
|
||||||
apiBase = "https://generativelanguage.googleapis.com/v1beta"
|
apiBase = "https://generativelanguage.googleapis.com/v1beta"
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue