feat: support custom proxy endpoints and chat stream for Vertex/Gemini
This adds support for providing custom API endpoints to the vertex provider
(such as Google AI Studio URLs or custom proxy servers) via the `api_base`
configuration option.
When a custom endpoint is provided, the provider correctly:
1. Formats the URL to `{apiBase}/{model_id}:{action}`
2. Injects the API key into the query parameters (`?key=...`) rather than
attempting to send a standard GCP OAuth `Authorization: Bearer` header.
Standard Vertex AI configurations without a custom `api_base` retain the
traditional behavior (GCP base URL + Bearer authentication).
Additionally, this patch implements `ChatStream` for the `vertex` provider,
allowing responses to stream over SSE (`alt=sse`) endpoints.
Resolves the user's request for Gemini API interoperability.
Co-authored-by: TanLuong <28281768+TanLuong@users.noreply.github.com>
This commit is contained in:
parent
6b503ca745
commit
80a2ac949b
3 changed files with 332 additions and 17 deletions
22
README.md
22
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).
|
||||
|
||||
</details>
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue