Merge pull request #1 from TanLuong/add-vertex-ai-provider-3635156555573898008

feat: implement Google Vertex AI provider
This commit is contained in:
Nhat Tan 2026-03-24 16:29:48 +07:00 committed by GitHub
commit 2c26acb902
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 599 additions and 1 deletions

View file

@ -955,6 +955,10 @@ type ModelConfig struct {
ConnectMode string `json:"connect_mode,omitempty"` // Connection mode: stdio, grpc ConnectMode string `json:"connect_mode,omitempty"` // Connection mode: stdio, grpc
Workspace string `json:"workspace,omitempty"` // Workspace path for CLI-based providers Workspace string `json:"workspace,omitempty"` // Workspace path for CLI-based providers
// Provider specific fields
ProjectID string `json:"project_id,omitempty"` // Project ID (e.g. for Google Vertex AI)
Region string `json:"region,omitempty"` // Region (e.g. for Google Vertex AI)
// Optional optimizations // Optional optimizations
RPM int `json:"rpm,omitempty"` // Requests per minute limit RPM int `json:"rpm,omitempty"` // Requests per minute limit
MaxTokensField string `json:"max_tokens_field,omitempty"` // Field name for max tokens (e.g., "max_completion_tokens") MaxTokensField string `json:"max_tokens_field,omitempty"` // Field name for max tokens (e.g., "max_completion_tokens")
@ -2149,6 +2153,8 @@ func expandMultiKeyModels(models []*ModelConfig) []*ModelConfig {
AuthMethod: m.AuthMethod, AuthMethod: m.AuthMethod,
ConnectMode: m.ConnectMode, ConnectMode: m.ConnectMode,
Workspace: m.Workspace, Workspace: m.Workspace,
ProjectID: m.ProjectID,
Region: m.Region,
RPM: m.RPM, RPM: m.RPM,
MaxTokensField: m.MaxTokensField, MaxTokensField: m.MaxTokensField,
RequestTimeout: m.RequestTimeout, RequestTimeout: m.RequestTimeout,
@ -2168,6 +2174,8 @@ func expandMultiKeyModels(models []*ModelConfig) []*ModelConfig {
AuthMethod: m.AuthMethod, AuthMethod: m.AuthMethod,
ConnectMode: m.ConnectMode, ConnectMode: m.ConnectMode,
Workspace: m.Workspace, Workspace: m.Workspace,
ProjectID: m.ProjectID,
Region: m.Region,
RPM: m.RPM, RPM: m.RPM,
MaxTokensField: m.MaxTokensField, MaxTokensField: m.MaxTokensField,
RequestTimeout: m.RequestTimeout, RequestTimeout: m.RequestTimeout,

View file

@ -15,6 +15,7 @@ import (
anthropicmessages "github.com/sipeed/picoclaw/pkg/providers/anthropic_messages" anthropicmessages "github.com/sipeed/picoclaw/pkg/providers/anthropic_messages"
"github.com/sipeed/picoclaw/pkg/providers/azure" "github.com/sipeed/picoclaw/pkg/providers/azure"
"github.com/sipeed/picoclaw/pkg/providers/bedrock" "github.com/sipeed/picoclaw/pkg/providers/bedrock"
"github.com/sipeed/picoclaw/pkg/providers/vertex"
) )
// createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store. // createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store.
@ -118,6 +119,24 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.RequestTimeout, cfg.RequestTimeout,
), modelID, nil ), modelID, nil
case "vertex", "vertex_ai", "vertex-ai":
if cfg.APIKey() == "" {
return nil, "", fmt.Errorf("api_key is required for vertex protocol")
}
if cfg.APIBase == "" && cfg.ProjectID == "" {
return nil, "", fmt.Errorf("either api_base or project_id is required for vertex protocol")
}
provider := vertex.NewProvider(
cfg.APIKey(),
cfg.APIBase,
cfg.Proxy,
cfg.ProjectID,
cfg.Region,
vertex.WithRequestTimeout(time.Duration(cfg.RequestTimeout)*time.Second),
)
return provider, modelID, nil
case "bedrock": case "bedrock":
// AWS Bedrock uses AWS SDK credentials (env vars, profiles, IAM roles, etc.) // AWS Bedrock uses AWS SDK credentials (env vars, profiles, IAM roles, etc.)
// api_base can be: // api_base can be:

View file

@ -0,0 +1,393 @@
// PicoClaw - Ultra-lightweight personal AI agent
// License: MIT
//
// Copyright (c) 2026 PicoClaw contributors
package vertex
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
"github.com/sipeed/picoclaw/pkg/providers/common"
"github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
)
type (
LLMResponse = protocoltypes.LLMResponse
Message = protocoltypes.Message
ToolDefinition = protocoltypes.ToolDefinition
ToolCall = protocoltypes.ToolCall
FunctionCall = protocoltypes.FunctionCall
)
const (
defaultRequestTimeout = common.DefaultRequestTimeout
)
// Provider implements the LLM provider interface for Google Vertex AI.
// It uses the standard Vertex AI REST API for Gemini models.
type Provider struct {
apiKey string
apiBase string // If provided, overrides the default construction
projectID string
region string
httpClient *http.Client
}
// Option configures the Vertex Provider.
type Option func(*Provider)
// WithRequestTimeout sets the HTTP request timeout.
func WithRequestTimeout(timeout time.Duration) Option {
return func(p *Provider) {
if timeout > 0 {
p.httpClient.Timeout = timeout
}
}
}
// NewProvider creates a new Vertex AI provider.
func NewProvider(apiKey, apiBase, proxy, projectID, region string, opts ...Option) *Provider {
p := &Provider{
apiKey: apiKey,
apiBase: strings.TrimRight(apiBase, "/"),
projectID: projectID,
region: region,
httpClient: common.NewHTTPClient(proxy),
}
for _, opt := range opts {
if opt != nil {
opt(p)
}
}
return p
}
// 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)
}
region := p.region
if region == "" {
region = "us-central1"
}
return fmt.Sprintf("https://%s-aiplatform.googleapis.com/v1/projects/%s/locations/%s/publishers/google/models/%s:generateContent", region, p.projectID, region, model)
}
// parseMediaData converts base64 media data into the Vertex AI inlineData format.
// It tries to detect mime type from the data URI scheme if present.
func parseMediaData(mediaData string) map[string]any {
mimeType := "image/jpeg"
data := mediaData
if strings.HasPrefix(mediaData, "data:") {
idx := strings.Index(mediaData, ";base64,")
if idx != -1 {
mimeType = mediaData[5:idx]
data = mediaData[idx+8:]
}
}
return map[string]any{
"inlineData": map[string]any{
"mimeType": mimeType,
"data": data,
},
}
}
// buildRequestBody formats the standard messages and tools into the Vertex AI (Gemini) REST payload format.
func (p *Provider) buildRequestBody(messages []Message, tools []ToolDefinition, options map[string]any) (map[string]any, error) {
req := make(map[string]any)
var contents []map[string]any
var systemInstruction *map[string]any
var currentContent map[string]any
for _, msg := range messages {
switch msg.Role {
case "system":
systemInstruction = &map[string]any{
"role": "system",
"parts": []map[string]any{
{"text": msg.Content},
},
}
case "user":
if currentContent != nil && currentContent["role"] == "user" {
// We need to group consecutive user messages (like tool responses)
} else {
if currentContent != nil {
contents = append(contents, currentContent)
}
currentContent = map[string]any{
"role": "user",
"parts": []map[string]any{},
}
}
parts := currentContent["parts"].([]map[string]any)
if msg.ToolCallID != "" {
// Tool response
parts = append(parts, map[string]any{
"functionResponse": map[string]any{
"name": msg.ToolCallID,
"response": map[string]any{
"result": msg.Content,
},
},
})
} else {
if msg.Content != "" {
parts = append(parts, map[string]any{"text": msg.Content})
}
for _, media := range msg.Media {
parts = append(parts, parseMediaData(media))
}
}
currentContent["parts"] = parts
case "assistant":
if currentContent != nil {
contents = append(contents, currentContent)
}
currentContent = map[string]any{
"role": "model",
"parts": []map[string]any{},
}
parts := currentContent["parts"].([]map[string]any)
if msg.Content != "" {
parts = append(parts, map[string]any{"text": msg.Content})
}
for _, tc := range msg.ToolCalls {
parts = append(parts, map[string]any{
"functionCall": map[string]any{
"name": tc.Name,
"args": tc.Arguments,
},
})
}
currentContent["parts"] = parts
case "tool":
if currentContent != nil && currentContent["role"] == "user" {
// Group tool response
} else {
if currentContent != nil {
contents = append(contents, currentContent)
}
currentContent = map[string]any{
"role": "user",
"parts": []map[string]any{},
}
}
parts := currentContent["parts"].([]map[string]any)
// Try to handle tool responses that might be strings instead of objects
// if they are just basic strings. But Gemini API expects an object.
responseObj := map[string]any{"result": msg.Content}
parts = append(parts, map[string]any{
"functionResponse": map[string]any{
"name": msg.ToolCallID,
"response": responseObj,
},
})
currentContent["parts"] = parts
}
}
if currentContent != nil {
contents = append(contents, currentContent)
}
req["contents"] = contents
if systemInstruction != nil {
req["systemInstruction"] = *systemInstruction
}
if len(tools) > 0 {
var funcDecls []map[string]any
for _, t := range tools {
if t.Type != "function" {
continue
}
decl := map[string]any{
"name": t.Function.Name,
"description": t.Function.Description,
}
if t.Function.Parameters != nil {
decl["parameters"] = t.Function.Parameters
}
funcDecls = append(funcDecls, decl)
}
if len(funcDecls) > 0 {
req["tools"] = []map[string]any{
{
"functionDeclarations": funcDecls,
},
}
}
}
generationConfig := make(map[string]any)
if val, ok := options["max_tokens"]; ok {
if maxTokens, ok := common.AsInt(val); ok {
generationConfig["maxOutputTokens"] = maxTokens
}
}
if temp, ok := common.AsFloat(options["temperature"]); ok {
generationConfig["temperature"] = temp
}
if len(generationConfig) > 0 {
req["generationConfig"] = generationConfig
}
return req, nil
}
func (p *Provider) Chat(
ctx context.Context,
messages []Message,
tools []ToolDefinition,
model string,
options map[string]any,
) (*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)
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 != "" {
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")
}
bodyBytes, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("failed to read response body: %w", err)
}
return p.parseResponse(bodyBytes)
}
func (p *Provider) parseResponse(body []byte) (*LLMResponse, error) {
var vResp 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"`
} `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, &vResp); err != nil {
return nil, fmt.Errorf("failed to decode response: %w", err)
}
if len(vResp.Candidates) == 0 {
return nil, fmt.Errorf("no candidates in response")
}
candidate := vResp.Candidates[0]
var content string
var toolCalls []ToolCall
for _, part := range candidate.Content.Parts {
if part.Text != "" {
content += part.Text
}
if part.FunctionCall != nil {
argsJSON, _ := json.Marshal(part.FunctionCall.Args)
toolCalls = append(toolCalls, 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),
},
})
}
}
finishReason := candidate.FinishReason
if finishReason == "STOP" {
finishReason = "stop"
} else if len(toolCalls) > 0 {
finishReason = "tool_calls"
}
return &LLMResponse{
Content: content,
ToolCalls: toolCalls,
FinishReason: finishReason,
Usage: &protocoltypes.UsageInfo{
PromptTokens: vResp.UsageMetadata.PromptTokenCount,
CompletionTokens: vResp.UsageMetadata.CandidatesTokenCount,
TotalTokens: vResp.UsageMetadata.TotalTokenCount,
},
}, nil
}
func (p *Provider) GetDefaultModel() string {
return "gemini-1.5-pro-preview-0409"
}

View file

@ -0,0 +1,178 @@
package vertex
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/sipeed/picoclaw/pkg/providers/protocoltypes"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestProvider_buildURL(t *testing.T) {
tests := []struct {
name string
apiBase string
projectID string
region string
model string
expected string
}{
{
name: "Default construction",
projectID: "my-project",
region: "us-central1",
model: "gemini-1.5-pro",
expected: "https://us-central1-aiplatform.googleapis.com/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-1.5-pro:generateContent",
},
{
name: "Default region",
projectID: "my-project",
model: "gemini-1.5-flash",
expected: "https://us-central1-aiplatform.googleapis.com/v1/projects/my-project/locations/us-central1/publishers/google/models/gemini-1.5-flash:generateContent",
},
{
name: "Override with base URL without method",
apiBase: "http://localhost:8080/v1",
model: "gemini-1.0-pro",
expected: "http://localhost:8080/v1/models/gemini-1.0-pro:generateContent",
},
{
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",
},
}
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)
assert.Equal(t, tt.expected, actual)
})
}
}
func TestProvider_buildRequestBody(t *testing.T) {
p := NewProvider("key", "", "", "proj", "us-central1")
messages := []protocoltypes.Message{
{Role: "system", Content: "You are a helpful assistant."},
{Role: "user", Content: "Hello!", Media: []string{"data:image/png;base64,iVBORw0KGgo"}},
{Role: "assistant", ToolCalls: []protocoltypes.ToolCall{{Name: "get_weather", Arguments: map[string]any{"location": "Tokyo"}}}},
{Role: "tool", ToolCallID: "get_weather", Content: "Sunny"},
{Role: "assistant", ToolCalls: []protocoltypes.ToolCall{{Name: "get_time", Arguments: map[string]any{"location": "Tokyo"}}}},
{Role: "tool", ToolCallID: "get_time", Content: "12:00 PM"},
}
tools := []protocoltypes.ToolDefinition{
{
Type: "function",
Function: protocoltypes.ToolFunctionDefinition{
Name: "get_weather",
Description: "Get the current weather",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"location": map[string]any{"type": "string"},
},
"required": []any{"location"},
},
},
},
}
options := map[string]any{
"temperature": 0.5,
"max_tokens": 1000,
}
req, err := p.buildRequestBody(messages, tools, options)
require.NoError(t, err)
// Check generation config
genCfg, ok := req["generationConfig"].(map[string]any)
require.True(t, ok)
assert.Equal(t, 0.5, genCfg["temperature"])
assert.Equal(t, 1000, genCfg["maxOutputTokens"])
// Check system instruction
sysInstr, ok := req["systemInstruction"].(map[string]any)
require.True(t, ok)
assert.Equal(t, "system", sysInstr["role"])
parts := sysInstr["parts"].([]map[string]any)
assert.Equal(t, "You are a helpful assistant.", parts[0]["text"])
// Check contents grouping (should combine the two tool responses into one user message)
contents, ok := req["contents"].([]map[string]any)
require.True(t, ok)
require.Len(t, contents, 5) // user, assistant, user(tool), assistant, user(tool)
assert.Equal(t, "user", contents[0]["role"])
userParts := contents[0]["parts"].([]map[string]any)
assert.Equal(t, "Hello!", userParts[0]["text"])
assert.Equal(t, "image/png", userParts[1]["inlineData"].(map[string]any)["mimeType"])
assert.Equal(t, "model", contents[1]["role"])
assert.Equal(t, "user", contents[2]["role"])
toolParts := contents[2]["parts"].([]map[string]any)
assert.Equal(t, "get_weather", toolParts[0]["functionResponse"].(map[string]any)["name"])
assert.Equal(t, "model", contents[3]["role"])
}
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"))
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"
}
],
"usageMetadata": {
"promptTokenCount": 10,
"candidatesTokenCount": 5,
"totalTokenCount": 15
}
}`
w.Header().Set("Content-Type", "application/json")
w.Write([]byte(mockResp))
}))
defer ts.Close()
p := NewProvider("test-key", ts.URL, "", "", "")
messages := []protocoltypes.Message{
{Role: "user", Content: "Hi"},
}
resp, err := p.Chat(context.Background(), messages, nil, "gemini-1.5-pro", nil)
require.NoError(t, err)
assert.Equal(t, "Hello, world!", resp.Content)
assert.Equal(t, "stop", resp.FinishReason)
assert.Equal(t, 10, resp.Usage.PromptTokens)
assert.Equal(t, 5, resp.Usage.CompletionTokens)
assert.Equal(t, 15, resp.Usage.TotalTokens)
}

View file

@ -170,7 +170,7 @@ func setupPicoEnabledEnv(t *testing.T) (string, func()) {
ModelList: map[string]config.ModelSecurityEntry{ ModelList: map[string]config.ModelSecurityEntry{
"custom-default": {APIKeys: []string{"sk-default"}}, "custom-default": {APIKeys: []string{"sk-default"}},
}, },
Channels: config.ChannelsSecurity{ Channels: &config.ChannelsSecurity{
Pico: &config.PicoSecurity{Token: "test-pico-token"}, Pico: &config.PicoSecurity{Token: "test-pico-token"},
}, },
}) })