diff --git a/README.md b/README.md index a60637a7c..dfca38bbe 100644 --- a/README.md +++ b/README.md @@ -383,6 +383,7 @@ PicoClaw supports 30+ LLM providers through the `model_list` configuration. Use | [Azure OpenAI](https://portal.azure.com/) | `azure/` | Required | Enterprise Azure deployment | | [GitHub Copilot](https://github.com/features/copilot) | `github-copilot/` | OAuth | Device code login | | [Antigravity](https://console.cloud.google.com/) | `antigravity/` | OAuth | Google Cloud AI | +| [Google Vertex AI](https://cloud.google.com/vertex-ai) | `vertex/` | Required / OAuth | GCP Vertex AI | | [AWS Bedrock](https://console.aws.amazon.com/bedrock)* | `bedrock/` | AWS credentials | Claude, Llama, Mistral on AWS | > \* AWS Bedrock requires build tag: `go build -tags bedrock`. Set `api_base` to a region name (e.g., `us-east-1`) for automatic endpoint resolution across all AWS partitions (aws, aws-cn, aws-us-gov). When using a full endpoint URL instead, you must also configure `AWS_REGION` via environment variable or AWS config/profile. @@ -416,6 +417,27 @@ PicoClaw supports 30+ LLM providers through the `model_list` configuration. Use } ``` +**Google AI Studio / Vertex AI (Custom Endpoints):** +You can use the `vertex` provider to query Google AI Studio endpoints or custom Vertex AI URLs. This supports using an `api_key` as a query parameter (`?key=...`). +```json +{ + "model_list": [ + { + "model_name": "gemini-pro-studio", + "model": "vertex/gemini-1.5-pro", + "api_base": "https://generativelanguage.googleapis.com/v1beta/models", + "api_key": "YOUR_GEMINI_API_KEY" + }, + { + "model_name": "gemini-pro-vertex", + "model": "vertex/gemini-3-pro-preview", + "api_base": "https://us-central1-aiplatform.googleapis.com/v1/publishers/google/models", + "api_key": "YOUR_VERTEX_API_KEY" + } + ] +} +``` + For full provider configuration details, see [Providers & Models](docs/providers.md). diff --git a/pkg/providers/vertex/provider.go b/pkg/providers/vertex/provider.go index 45b67285a..59a3060b2 100644 --- a/pkg/providers/vertex/provider.go +++ b/pkg/providers/vertex/provider.go @@ -6,6 +6,7 @@ package vertex import ( + "bufio" "bytes" "context" "encoding/json" @@ -73,19 +74,44 @@ func NewProvider(apiKey, apiBase, proxy, projectID, region string, opts ...Optio } // buildURL constructs the Vertex AI REST endpoint URL. -func (p *Provider) buildURL(model string) string { - if p.apiBase != "" { - if strings.Contains(p.apiBase, "generateContent") { - return p.apiBase - } - return fmt.Sprintf("%s/models/%s:generateContent", p.apiBase, model) +func (p *Provider) buildURL(model string, action string) string { + if action == "" { + action = "generateContent" } - region := p.region - if region == "" { - region = "us-central1" + var baseURL string + if p.apiBase != "" { + if strings.Contains(p.apiBase, "generateContent") { + baseURL = p.apiBase + } else { + baseURL = fmt.Sprintf("%s/%s:%s", p.apiBase, model, action) + } + } else { + region := p.region + if region == "" { + region = "us-central1" + } + baseURL = fmt.Sprintf("https://%s-aiplatform.googleapis.com/v1/projects/%s/locations/%s/publishers/google/models/%s:%s", region, p.projectID, region, model, action) } - return fmt.Sprintf("https://%s-aiplatform.googleapis.com/v1/projects/%s/locations/%s/publishers/google/models/%s:generateContent", region, p.projectID, region, model) + + // Only append ?key= for custom apiBase endpoints + if p.apiBase != "" && p.apiKey != "" && !strings.Contains(baseURL, "key=") { + if strings.Contains(baseURL, "?") { + baseURL = fmt.Sprintf("%s&key=%s", baseURL, p.apiKey) + } else { + baseURL = fmt.Sprintf("%s?key=%s", baseURL, p.apiKey) + } + } + + if action == "streamGenerateContent" && !strings.Contains(baseURL, "alt=sse") { + if strings.Contains(baseURL, "?") { + baseURL = fmt.Sprintf("%s&alt=sse", baseURL) + } else { + baseURL = fmt.Sprintf("%s?alt=sse", baseURL) + } + } + + return baseURL } @@ -287,7 +313,7 @@ func (p *Provider) Chat( return nil, fmt.Errorf("failed to marshal request: %w", err) } - requestURL := p.buildURL(model) + requestURL := p.buildURL(model, "generateContent") req, err := http.NewRequestWithContext(ctx, "POST", requestURL, bytes.NewReader(jsonData)) if err != nil { @@ -295,7 +321,7 @@ func (p *Provider) Chat( } req.Header.Set("Content-Type", "application/json") - if p.apiKey != "" { + if p.apiKey != "" && !strings.Contains(requestURL, "key=") { req.Header.Set("Authorization", "Bearer "+p.apiKey) } @@ -317,6 +343,161 @@ func (p *Provider) Chat( return p.parseResponse(bodyBytes) } +func (p *Provider) ChatStream( + ctx context.Context, + messages []Message, + tools []ToolDefinition, + model string, + options map[string]any, + onChunk func(accumulated string), +) (*LLMResponse, error) { + if p.apiBase == "" && p.projectID == "" { + return nil, fmt.Errorf("Vertex AI requires either an api_base or a project_id") + } + + requestBody, err := p.buildRequestBody(messages, tools, options) + if err != nil { + return nil, fmt.Errorf("failed to build request body: %w", err) + } + + jsonData, err := json.Marshal(requestBody) + if err != nil { + return nil, fmt.Errorf("failed to marshal request: %w", err) + } + + requestURL := p.buildURL(model, "streamGenerateContent") + + req, err := http.NewRequestWithContext(ctx, "POST", requestURL, bytes.NewReader(jsonData)) + if err != nil { + return nil, fmt.Errorf("failed to create request: %w", err) + } + + req.Header.Set("Content-Type", "application/json") + if p.apiKey != "" && !strings.Contains(requestURL, "key=") { + req.Header.Set("Authorization", "Bearer "+p.apiKey) + } + + resp, err := p.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("failed to send request: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, common.HandleErrorResponse(resp, "vertex") + } + + var accumulatedText string + var allToolCalls []ToolCall + var finalResponse *LLMResponse + + scanner := bufio.NewScanner(resp.Body) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" { + continue + } + + if strings.HasPrefix(line, "data: ") { + line = strings.TrimPrefix(line, "data: ") + } else if line == "[" || line == "]" || line == "," { + continue + } + + var chunk struct { + Candidates []struct { + Content struct { + Parts []struct { + Text string `json:"text"` + FunctionCall *struct { + Name string `json:"name"` + Args map[string]any `json:"args"` + } `json:"functionCall,omitempty"` + } `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,omitempty"` + } + + if err := json.Unmarshal([]byte(line), &chunk); err != nil { + continue + } + + if len(chunk.Candidates) > 0 { + candidate := chunk.Candidates[0] + for _, part := range candidate.Content.Parts { + if part.Text != "" { + accumulatedText += part.Text + if onChunk != nil { + onChunk(accumulatedText) + } + } + if part.FunctionCall != nil { + argsJSON, _ := json.Marshal(part.FunctionCall.Args) + toolCall := ToolCall{ + ID: fmt.Sprintf("call_%s_%d", part.FunctionCall.Name, time.Now().UnixNano()), + Name: part.FunctionCall.Name, + Arguments: part.FunctionCall.Args, + Function: &FunctionCall{ + Name: part.FunctionCall.Name, + Arguments: string(argsJSON), + }, + } + allToolCalls = append(allToolCalls, toolCall) + } + } + + if candidate.FinishReason != "" && finalResponse == nil { + finishReason := candidate.FinishReason + if finishReason == "STOP" { + finishReason = "stop" + } else if len(allToolCalls) > 0 { + finishReason = "tool_calls" + } + + finalResponse = &LLMResponse{ + Content: accumulatedText, + ToolCalls: allToolCalls, + FinishReason: finishReason, + } + } + } + + if chunk.UsageMetadata != nil { + if finalResponse == nil { + finalResponse = &LLMResponse{ + Content: accumulatedText, + ToolCalls: allToolCalls, + } + } + finalResponse.Usage = &protocoltypes.UsageInfo{ + PromptTokens: chunk.UsageMetadata.PromptTokenCount, + CompletionTokens: chunk.UsageMetadata.CandidatesTokenCount, + TotalTokens: chunk.UsageMetadata.TotalTokenCount, + } + } + } + + if finalResponse == nil { + finishReason := "stop" + if len(allToolCalls) > 0 { + finishReason = "tool_calls" + } + finalResponse = &LLMResponse{ + Content: accumulatedText, + ToolCalls: allToolCalls, + FinishReason: finishReason, + } + } + + return finalResponse, nil +} + func (p *Provider) parseResponse(body []byte) (*LLMResponse, error) { var vResp struct { Candidates []struct { diff --git a/pkg/providers/vertex/provider_test.go b/pkg/providers/vertex/provider_test.go index 67a96d84b..65555083b 100644 --- a/pkg/providers/vertex/provider_test.go +++ b/pkg/providers/vertex/provider_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "net/http" "net/http/httptest" + "net/url" "testing" @@ -37,22 +38,22 @@ func TestProvider_buildURL(t *testing.T) { }, { name: "Override with base URL without method", - apiBase: "http://localhost:8080/v1", + apiBase: "http://localhost:8080/v1/models", model: "gemini-1.0-pro", - expected: "http://localhost:8080/v1/models/gemini-1.0-pro:generateContent", + expected: "http://localhost:8080/v1/models/gemini-1.0-pro:generateContent?key=key", }, { name: "Override with full endpoint URL", apiBase: "https://my-custom-proxy.com/my-endpoint:generateContent", model: "gemini-1.5-pro", - expected: "https://my-custom-proxy.com/my-endpoint:generateContent", + expected: "https://my-custom-proxy.com/my-endpoint:generateContent?key=key", }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { p := NewProvider("key", tt.apiBase, "", tt.projectID, tt.region) - actual := p.buildURL(tt.model) + actual := p.buildURL(tt.model, "generateContent") assert.Equal(t, tt.expected, actual) }) } @@ -132,7 +133,9 @@ func TestProvider_Chat(t *testing.T) { // Create a mock server ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { assert.Equal(t, "POST", r.Method) - assert.Equal(t, "Bearer test-key", r.Header.Get("Authorization")) + // Because we're using the query 'key=test-key' we no longer have Bearer authentication + // assert.Equal(t, "Bearer test-key", r.Header.Get("Authorization")) + assert.Contains(t, r.URL.String(), "key=test-key") var reqBody map[string]any err := json.NewDecoder(r.Body).Decode(&reqBody) @@ -176,3 +179,112 @@ func TestProvider_Chat(t *testing.T) { assert.Equal(t, 5, resp.Usage.CompletionTokens) assert.Equal(t, 15, resp.Usage.TotalTokens) } + +func TestProvider_ChatStream(t *testing.T) { + // Create a mock server + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + assert.Contains(t, r.URL.String(), "key=test-key") + assert.Contains(t, r.URL.String(), "alt=sse") + + var reqBody map[string]any + err := json.NewDecoder(r.Body).Decode(&reqBody) + require.NoError(t, err) + + w.Header().Set("Content-Type", "text/event-stream") + // Write mock chunks + w.Write([]byte(`data: {"candidates":[{"content":{"parts":[{"text":"Hello"}]}}]}` + "\n\n")) + w.Write([]byte(`data: {"candidates":[{"content":{"parts":[{"text":", world!"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":10,"candidatesTokenCount":5,"totalTokenCount":15}}` + "\n\n")) + })) + defer ts.Close() + + p := NewProvider("test-key", ts.URL, "", "my-project", "us-central1") + opts := make(map[string]any) + + var chunks []string + resp, err := p.ChatStream( + context.Background(), + []protocoltypes.Message{{Role: "user", Content: "Say hello!"}}, + nil, + "gemini-1.5-pro", + opts, + func(accumulated string) { + chunks = append(chunks, accumulated) + }, + ) + + require.NoError(t, err) + require.NotNil(t, resp) + + assert.Equal(t, "Hello, world!", resp.Content) + assert.Equal(t, "stop", resp.FinishReason) + assert.NotNil(t, resp.Usage) + if resp.Usage != nil { + assert.Equal(t, 10, resp.Usage.PromptTokens) + assert.Equal(t, 5, resp.Usage.CompletionTokens) + assert.Equal(t, 15, resp.Usage.TotalTokens) + } + + assert.Equal(t, []string{"Hello", "Hello, world!"}, chunks) +} + +func TestProvider_Chat_Standard(t *testing.T) { + // Create a mock server + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + // Standard Vertex without apiBase should use Bearer authentication + assert.Equal(t, "Bearer test-key", r.Header.Get("Authorization")) + + var reqBody map[string]any + err := json.NewDecoder(r.Body).Decode(&reqBody) + require.NoError(t, err) + + // Return a mock response + mockResp := `{ + "candidates": [ + { + "content": { + "parts": [ + {"text": "Hello, world!"} + ] + }, + "finishReason": "STOP" + } + ] + }` + w.Header().Set("Content-Type", "application/json") + w.Write([]byte(mockResp)) + })) + defer ts.Close() + + // Use an empty apiBase so it builds standard Vertex URLs + p := NewProvider("test-key", "", "", "my-project", "us-central1") + // Since buildURL will use aiplatform.googleapis.com, we override the httpClient Transport + // to redirect requests to our mock server for this test by swapping the base URL out in a custom RoundTripper. + p.httpClient.Transport = roundTripperFunc(func(req *http.Request) (*http.Response, error) { + req.URL, _ = url.Parse(ts.URL) + return http.DefaultTransport.RoundTrip(req) + }) + + opts := make(map[string]any) + + resp, err := p.Chat( + context.Background(), + []protocoltypes.Message{{Role: "user", Content: "Say hello!"}}, + nil, + "gemini-1.5-pro", + opts, + ) + + require.NoError(t, err) + require.NotNil(t, resp) + + assert.Equal(t, "Hello, world!", resp.Content) + assert.Equal(t, "stop", resp.FinishReason) +} + +type roundTripperFunc func(*http.Request) (*http.Response, error) + +func (f roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +}