feat(tool): tool schema semplification

This commit is contained in:
afjcjsbx
2026-04-27 21:10:30 +02:00
parent 4eeb69688e
commit cd7717bc15
23 changed files with 654 additions and 136 deletions
+1
View File
@@ -41,6 +41,7 @@
},
{
"model_name": "gemini",
"_comment": "Optional: set \"tool_schema_transform\": \"simple\" for providers that reject complex tool JSON Schema.",
"model": "antigravity/gemini-2.0-flash",
"auth_method": "oauth"
},
+40 -16
View File
@@ -116,23 +116,47 @@ This design also enables **multi-agent support** with flexible provider selectio
#### `model_list` Entry Fields
| Field | Type | Required | Description |
|-------|------|----------|-------------|
| `model_name` | string | Yes | Unique name used to reference this model in agent config |
| `provider` | string | No | Preferred provider identifier. When present, PicoClaw sends `model` unchanged to that provider |
| `model` | string | Yes | Native model ID when `provider` is set. If `provider` is omitted, the legacy `provider/model` form is still supported |
| `api_keys` | string[] | Yes* | API key(s) for authentication. Multiple keys enable per-request rotation. Not required for local providers (Ollama, LM Studio, VLLM) |
| `api_base` | string | No | Override the default API endpoint URL |
| `proxy` | string | No | HTTP proxy URL for this model entry |
| `user_agent` | string | No | Custom `User-Agent` header sent with API requests (supported by OpenAI-compatible, Gemini, Anthropic, and Azure providers) |
| `request_timeout` | int | No | Request timeout in seconds (default varies by provider) |
| `max_tokens_field` | string | No | Override the max tokens field name in request body (e.g., `max_completion_tokens` for o1 models) |
| `thinking_level` | string | No | Extended thinking level: `off`, `low`, `medium`, `high`, `xhigh`, or `adaptive` |
| `extra_body` | object | No | Additional fields to inject into every request body |
| Field | Type | Required | Description |
|-------|------|----------|---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| `model_name` | string | Yes | Unique name used to reference this model in agent config |
| `provider` | string | No | Preferred provider identifier. When present, PicoClaw sends `model` unchanged to that provider |
| `model` | string | Yes | Native model ID when `provider` is set. If `provider` is omitted, the legacy `provider/model` form is still supported |
| `api_keys` | string[] | Yes* | API key(s) for authentication. Multiple keys enable per-request rotation. Not required for local providers (Ollama, LM Studio, VLLM) |
| `api_base` | string | No | Override the default API endpoint URL |
| `proxy` | string | No | HTTP proxy URL for this model entry |
| `user_agent` | string | No | Custom `User-Agent` header sent with API requests (supported by OpenAI-compatible, Gemini, Anthropic, and Azure providers) |
| `request_timeout` | int | No | Request timeout in seconds (default varies by provider) |
| `max_tokens_field` | string | No | Override the max tokens field name in request body (e.g., `max_completion_tokens` for o1 models) |
| `thinking_level` | string | No | Extended thinking level: `off`, `low`, `medium`, `high`, `xhigh`, or `adaptive` |
| `tool_schema_transform` | string | No | Optional compatibility transform for tool parameter schemas. Default: disabled. Supported values: `simple`. |
| `extra_body` | object | No | Additional fields to inject into every request body |
| `custom_headers` | object | No | Additional HTTP headers to inject into every request (e.g., `{"X-Source":"coding-plan"}`). If a key matches a built-in header, the custom value overrides the built-in one (e.g., `Authorization`, `User-Agent`, `Content-Type`, `Accept`). |
| `rpm` | int | No | Per-minute request rate limit |
| `fallbacks` | string[] | No | Fallback model names for automatic failover |
| `enabled` | bool | No | Whether this model entry is active (default: `true`) |
| `rpm` | int | No | Per-minute request rate limit |
| `fallbacks` | string[] | No | Fallback model names for automatic failover |
| `enabled` | bool | No | Whether this model entry is active (default: `true`) |
#### Tool Schema Compatibility
By default, PicoClaw now forwards tool JSON Schemas unchanged.
Some providers reject advanced JSON Schema features such as `$ref`, `$defs`, `anyOf`, `oneOf`, `allOf`, `pattern`, or numeric/string constraints inside tool declarations. For those models, you can opt into a compatibility transform per model entry with `tool_schema_transform`.
Use `simple` when the upstream provider expects the conservative style function schema subset:
```json
{
"model_name": "gemini-2.5-flash-safe-tools",
"provider": "gemini",
"model": "gemini-2.5-flash",
"api_keys": ["your-gemini-key"],
"tool_schema_transform": "simple"
}
```
Notes:
- Default behavior is disabled. If you omit `tool_schema_transform`, PicoClaw sends the original tool schema.
- The setting is per model entry, so you can enable it only for the providers that need it.
#### Provider / Model Resolution
+46 -39
View File
@@ -17,6 +17,7 @@ import (
"github.com/sipeed/picoclaw/pkg"
"github.com/sipeed/picoclaw/pkg/fileutil"
"github.com/sipeed/picoclaw/pkg/logger"
providercommon "github.com/sipeed/picoclaw/pkg/providers/common"
)
// rrCounter is a global counter for round-robin load balancing across models.
@@ -553,12 +554,13 @@ type ModelConfig struct {
Workspace string `json:"workspace,omitempty"` // Workspace path for CLI-based providers
// Optional optimizations
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")
RequestTimeout int `json:"request_timeout,omitempty"`
ThinkingLevel string `json:"thinking_level,omitempty"` // Extended thinking: off|low|medium|high|xhigh|adaptive
ExtraBody map[string]any `json:"extra_body,omitempty"` // Additional fields to inject into request body
CustomHeaders map[string]string `json:"custom_headers,omitempty"` // Additional headers to inject into every HTTP request
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")
RequestTimeout int `json:"request_timeout,omitempty"`
ThinkingLevel string `json:"thinking_level,omitempty"` // Extended thinking: off|low|medium|high|xhigh|adaptive
ToolSchemaTransform string `json:"tool_schema_transform,omitempty"` // Optional tool schema compatibility transform (e.g. "simple")
ExtraBody map[string]any `json:"extra_body,omitempty"` // Additional fields to inject into request body
CustomHeaders map[string]string `json:"custom_headers,omitempty"` // Additional headers to inject into every HTTP request
APIKeys SecureStrings `json:"api_keys,omitzero" yaml:"api_keys,omitempty"` // API authentication keys (multiple keys for failover)
@@ -595,6 +597,9 @@ func (c *ModelConfig) Validate() error {
if c.Model == "" {
return fmt.Errorf("model is required")
}
if _, err := providercommon.NormalizeToolSchemaTransform(c.ToolSchemaTransform); err != nil {
return err
}
return nil
}
@@ -1419,23 +1424,24 @@ func expandMultiKeyModels(models []*ModelConfig) []*ModelConfig {
// Create a copy for the additional key
additionalEntry := &ModelConfig{
ModelName: expandedName,
Provider: m.Provider,
Model: m.Model,
APIBase: m.APIBase,
APIKeys: SimpleSecureStrings(keys[i]),
Proxy: m.Proxy,
AuthMethod: m.AuthMethod,
ConnectMode: m.ConnectMode,
Workspace: m.Workspace,
RPM: m.RPM,
MaxTokensField: m.MaxTokensField,
RequestTimeout: m.RequestTimeout,
ThinkingLevel: m.ThinkingLevel,
ExtraBody: m.ExtraBody,
CustomHeaders: m.CustomHeaders,
UserAgent: m.UserAgent,
isVirtual: true,
ModelName: expandedName,
Provider: m.Provider,
Model: m.Model,
APIBase: m.APIBase,
APIKeys: SimpleSecureStrings(keys[i]),
Proxy: m.Proxy,
AuthMethod: m.AuthMethod,
ConnectMode: m.ConnectMode,
Workspace: m.Workspace,
RPM: m.RPM,
MaxTokensField: m.MaxTokensField,
RequestTimeout: m.RequestTimeout,
ThinkingLevel: m.ThinkingLevel,
ToolSchemaTransform: m.ToolSchemaTransform,
ExtraBody: m.ExtraBody,
CustomHeaders: m.CustomHeaders,
UserAgent: m.UserAgent,
isVirtual: true,
}
expanded = append(expanded, additionalEntry)
fallbackNames = append(fallbackNames, expandedName)
@@ -1443,22 +1449,23 @@ func expandMultiKeyModels(models []*ModelConfig) []*ModelConfig {
// Create the primary entry with first key and fallbacks
primaryEntry := &ModelConfig{
ModelName: originalName,
Provider: m.Provider,
Model: m.Model,
APIBase: m.APIBase,
Proxy: m.Proxy,
AuthMethod: m.AuthMethod,
ConnectMode: m.ConnectMode,
Workspace: m.Workspace,
RPM: m.RPM,
MaxTokensField: m.MaxTokensField,
RequestTimeout: m.RequestTimeout,
ThinkingLevel: m.ThinkingLevel,
ExtraBody: m.ExtraBody,
CustomHeaders: m.CustomHeaders,
UserAgent: m.UserAgent,
APIKeys: SimpleSecureStrings(keys[0]),
ModelName: originalName,
Provider: m.Provider,
Model: m.Model,
APIBase: m.APIBase,
Proxy: m.Proxy,
AuthMethod: m.AuthMethod,
ConnectMode: m.ConnectMode,
Workspace: m.Workspace,
RPM: m.RPM,
MaxTokensField: m.MaxTokensField,
RequestTimeout: m.RequestTimeout,
ThinkingLevel: m.ThinkingLevel,
ToolSchemaTransform: m.ToolSchemaTransform,
ExtraBody: m.ExtraBody,
CustomHeaders: m.CustomHeaders,
UserAgent: m.UserAgent,
APIKeys: SimpleSecureStrings(keys[0]),
}
// Prepend new fallbacks to existing ones
+30
View File
@@ -1945,6 +1945,36 @@ func TestModelConfig_CustomHeadersRoundTrip(t *testing.T) {
}
}
func TestModelConfig_ToolSchemaTransformRoundTrip(t *testing.T) {
dir := t.TempDir()
cfgPath := filepath.Join(dir, "config.json")
cfg := &Config{
Version: CurrentVersion,
ModelList: []*ModelConfig{
{
ModelName: "test-model",
Model: "openai/test",
APIKeys: SimpleSecureStrings("sk-test"),
ToolSchemaTransform: "simple",
},
},
}
if err := SaveConfig(cfgPath, cfg); err != nil {
t.Fatalf("SaveConfig error: %v", err)
}
loaded, err := LoadConfig(cfgPath)
if err != nil {
t.Fatalf("LoadConfig error: %v", err)
}
if got := loaded.ModelList[0].ToolSchemaTransform; got != "simple" {
t.Fatalf("ToolSchemaTransform = %q, want %q", got, "simple")
}
}
func TestDefaultConfig_MinimaxExtraBody(t *testing.T) {
cfg := DefaultConfig()
+18
View File
@@ -158,6 +158,15 @@ func TestModelConfig_Validate(t *testing.T) {
},
wantErr: false,
},
{
name: "valid tool schema transform",
config: ModelConfig{
ModelName: "test",
Model: "openai/gpt-4o",
ToolSchemaTransform: "simple",
},
wantErr: false,
},
{
name: "missing model_name",
config: ModelConfig{
@@ -177,6 +186,15 @@ func TestModelConfig_Validate(t *testing.T) {
config: ModelConfig{},
wantErr: true,
},
{
name: "invalid tool schema transform",
config: ModelConfig{
ModelName: "test",
Model: "openai/gpt-4o",
ToolSchemaTransform: "invalid",
},
wantErr: true,
},
}
for _, tt := range tests {
+16 -9
View File
@@ -187,15 +187,16 @@ func TestExpandMultiKeyModels_Deduplication(t *testing.T) {
func TestExpandMultiKeyModels_PreservesOtherFields(t *testing.T) {
modelCfg := &ModelConfig{
ModelName: "gpt-4",
Provider: "openrouter",
Model: "openai/gpt-4o",
APIBase: "https://api.example.com",
Proxy: "http://proxy:8080",
RPM: 60,
MaxTokensField: "max_completion_tokens",
RequestTimeout: 30,
ThinkingLevel: "high",
ModelName: "gpt-4",
Provider: "openrouter",
Model: "openai/gpt-4o",
APIBase: "https://api.example.com",
Proxy: "http://proxy:8080",
RPM: 60,
MaxTokensField: "max_completion_tokens",
RequestTimeout: 30,
ThinkingLevel: "high",
ToolSchemaTransform: "simple",
}
modelCfg.APIKeys = SimpleSecureStrings("key0", "key1") // Use internal field for multi-key testing
models := []*ModelConfig{modelCfg}
@@ -225,6 +226,9 @@ func TestExpandMultiKeyModels_PreservesOtherFields(t *testing.T) {
if primary.ThinkingLevel != "high" {
t.Errorf("expected thinking_level preserved, got %q", primary.ThinkingLevel)
}
if primary.ToolSchemaTransform != "simple" {
t.Errorf("expected tool_schema_transform preserved, got %q", primary.ToolSchemaTransform)
}
// Check additional entry also preserves fields
additional := result[0]
@@ -237,6 +241,9 @@ func TestExpandMultiKeyModels_PreservesOtherFields(t *testing.T) {
if additional.RPM != 60 {
t.Errorf("expected additional rpm preserved, got %d", additional.RPM)
}
if additional.ToolSchemaTransform != "simple" {
t.Errorf("expected additional tool_schema_transform preserved, got %q", additional.ToolSchemaTransform)
}
}
func TestExpandMultiKeyModels_IsVirtualFlag(t *testing.T) {
+11 -5
View File
@@ -16,11 +16,11 @@ var geminiSupportedTypes = map[string]bool{
"string": true,
}
// SanitizeSchemaForGemini reduces a JSON Schema to the conservative subset
// accepted by Gemini-style function declarations. It resolves local refs,
// collapses composition keywords like anyOf/oneOf/allOf, and strips advanced
// keywords that Gemini rejects.
func SanitizeSchemaForGemini(schema map[string]any) map[string]any {
// SanitizeSchemaForGoogle reduces a JSON Schema to the conservative subset
// accepted by Google/Gemini-style function declarations. It resolves local
// refs, collapses composition keywords like anyOf/oneOf/allOf, and strips
// advanced keywords that Gemini-compatible backends often reject.
func SanitizeSchemaForGoogle(schema map[string]any) map[string]any {
if schema == nil {
return nil
}
@@ -39,6 +39,12 @@ func SanitizeSchemaForGemini(schema map[string]any) map[string]any {
return result
}
// SanitizeSchemaForGemini is kept as a compatibility alias for the original
// Google/Gemini sanitizer name.
func SanitizeSchemaForGemini(schema map[string]any) map[string]any {
return SanitizeSchemaForGoogle(schema)
}
type geminiSchemaSanitizer struct {
root map[string]any
}
@@ -0,0 +1,59 @@
package common
import (
"fmt"
"strings"
)
const (
ToolSchemaTransformOff = ""
ToolSchemaTransformSimple = "simple"
)
// NormalizeToolSchemaTransform resolves user-facing aliases to a canonical
// transform mode. Empty values and explicit "off"-style values disable schema
// transformation.
func NormalizeToolSchemaTransform(raw string) (string, error) {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "", "off", "none", "native":
return ToolSchemaTransformOff, nil
case "simple", "basic", "strict", "flat":
return ToolSchemaTransformSimple, nil
default:
return "", fmt.Errorf("unsupported tool_schema_transform %q (supported: off, simple)", raw)
}
}
// TransformToolDefinitions clones tool definitions and applies the configured
// schema transform to function parameter schemas. When the transform is off, the
// original slice is returned unchanged.
func TransformToolDefinitions(tools []ToolDefinition, transform string) ([]ToolDefinition, error) {
transform, err := NormalizeToolSchemaTransform(transform)
if err != nil {
return nil, err
}
if transform == ToolSchemaTransformOff || len(tools) == 0 {
return tools, nil
}
out := make([]ToolDefinition, len(tools))
for i, tool := range tools {
out[i] = tool
if tool.Type != "function" {
continue
}
out[i].Function = tool.Function
out[i].Function.Parameters = transformToolSchema(tool.Function.Parameters, transform)
}
return out, nil
}
func transformToolSchema(schema map[string]any, transform string) map[string]any {
switch transform {
case ToolSchemaTransformSimple:
return SanitizeSchemaForGoogle(schema)
default:
return cloneGeminiSchemaMap(schema)
}
}
+31 -19
View File
@@ -168,7 +168,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if err != nil {
return nil, "", err
}
return provider, modelID, nil
return finalizeProviderFromConfig(provider, modelID, cfg)
}
// OpenAI with API key
if cfg.APIKey() == "" && cfg.APIBase == "" {
@@ -189,7 +189,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.CustomHeaders,
)
provider.SetProviderName(protocol)
return provider, modelID, nil
return finalizeProviderFromConfig(provider, modelID, cfg)
case "azure", "azure-openai":
// Azure OpenAI uses deployment-based URLs, api-key header auth,
@@ -202,13 +202,13 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
"api_base is required for azure protocol (e.g., https://your-resource.openai.azure.com)",
)
}
return azure.NewProviderWithTimeout(
return finalizeProviderFromConfig(azure.NewProviderWithTimeout(
cfg.APIKey(),
cfg.APIBase,
cfg.Proxy,
userAgent,
cfg.RequestTimeout,
), modelID, nil
), modelID, cfg)
case "bedrock":
// AWS Bedrock uses AWS SDK credentials (env vars, profiles, IAM roles, etc.)
@@ -244,7 +244,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if err != nil {
return nil, "", fmt.Errorf("creating bedrock provider: %w", err)
}
return provider, modelID, nil
return finalizeProviderFromConfig(provider, modelID, cfg)
case "litellm", "lmstudio", "openrouter", "groq", "zhipu", "nvidia", "venice",
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
@@ -270,7 +270,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.CustomHeaders,
)
provider.SetProviderName(protocol)
return provider, modelID, nil
return finalizeProviderFromConfig(provider, modelID, cfg)
case "gemini":
if cfg.APIKey() == "" && cfg.APIBase == "" {
@@ -280,7 +280,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if apiBase == "" {
apiBase = getDefaultAPIBase(protocol)
}
return NewGeminiProvider(
return finalizeProviderFromConfig(NewGeminiProvider(
cfg.APIKey(),
apiBase,
cfg.Proxy,
@@ -288,7 +288,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.RequestTimeout,
cfg.ExtraBody,
cfg.CustomHeaders,
), modelID, nil
), modelID, cfg)
case "minimax":
// Minimax requires reasoning_split: true in the request body
@@ -317,7 +317,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.CustomHeaders,
)
provider.SetProviderName(protocol)
return provider, modelID, nil
return finalizeProviderFromConfig(provider, modelID, cfg)
case "anthropic":
if cfg.AuthMethod == "oauth" || cfg.AuthMethod == "token" {
@@ -326,7 +326,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if err != nil {
return nil, "", err
}
return provider, modelID, nil
return finalizeProviderFromConfig(provider, modelID, cfg)
}
// Use API key with HTTP API
apiBase := cfg.APIBase
@@ -347,7 +347,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
cfg.CustomHeaders,
)
provider.SetProviderName(protocol)
return provider, modelID, nil
return finalizeProviderFromConfig(provider, modelID, cfg)
case "anthropic-messages":
// Anthropic Messages API with native format (HTTP-based, no SDK)
@@ -358,12 +358,12 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if cfg.APIKey() == "" {
return nil, "", fmt.Errorf("api_key is required for anthropic-messages protocol (model: %s)", cfg.Model)
}
return anthropicmessages.NewProviderWithTimeout(
return finalizeProviderFromConfig(anthropicmessages.NewProviderWithTimeout(
cfg.APIKey(),
apiBase,
userAgent,
cfg.RequestTimeout,
), modelID, nil
), modelID, cfg)
case "coding-plan-anthropic", "alibaba-coding-anthropic":
// Alibaba Coding Plan with Anthropic-compatible API
@@ -374,29 +374,29 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if cfg.APIKey() == "" {
return nil, "", fmt.Errorf("api_key is required for %q protocol (model: %s)", protocol, cfg.Model)
}
return anthropicmessages.NewProviderWithTimeout(
return finalizeProviderFromConfig(anthropicmessages.NewProviderWithTimeout(
cfg.APIKey(),
apiBase,
userAgent,
cfg.RequestTimeout,
), modelID, nil
), modelID, cfg)
case "antigravity":
return NewAntigravityProvider(), modelID, nil
return finalizeProviderFromConfig(NewAntigravityProvider(), modelID, cfg)
case "claude-cli", "claudecli":
workspace := cfg.Workspace
if workspace == "" {
workspace = "."
}
return NewClaudeCliProvider(workspace), modelID, nil
return finalizeProviderFromConfig(NewClaudeCliProvider(workspace), modelID, cfg)
case "codex-cli", "codexcli":
workspace := cfg.Workspace
if workspace == "" {
workspace = "."
}
return NewCodexCliProvider(workspace), modelID, nil
return finalizeProviderFromConfig(NewCodexCliProvider(workspace), modelID, cfg)
case "github-copilot", "copilot":
apiBase := cfg.APIBase
@@ -411,13 +411,25 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
if err != nil {
return nil, "", err
}
return provider, modelID, nil
return finalizeProviderFromConfig(provider, modelID, cfg)
default:
return nil, "", fmt.Errorf("unknown protocol %q in model %q", protocol, cfg.Model)
}
}
func finalizeProviderFromConfig(
provider LLMProvider,
modelID string,
cfg *config.ModelConfig,
) (LLMProvider, string, error) {
wrapped, err := wrapProviderWithToolSchemaTransform(provider, cfg.ToolSchemaTransform)
if err != nil {
return nil, "", err
}
return wrapped, modelID, nil
}
func isEmptyAPIKeyAllowed(protocol string) bool {
meta, ok := protocolMetaByName[protocol]
return ok && meta.emptyAPIKeyAllowed
+39
View File
@@ -1202,3 +1202,42 @@ func TestCreateProviderFromConfig_BedrockWithEndpointURL(t *testing.T) {
// Unexpected error - fail the test
t.Errorf("unexpected error from bedrock provider: %v", err)
}
func TestCreateProviderFromConfig_ToolSchemaTransformWrapsProvider(t *testing.T) {
cfg := &config.ModelConfig{
ModelName: "claude-cli-test",
Provider: "claude-cli",
Model: "claude-sonnet-4.6",
Workspace: t.TempDir(),
ToolSchemaTransform: "simple",
}
provider, modelID, err := CreateProviderFromConfig(cfg)
if err != nil {
t.Fatalf("CreateProviderFromConfig() error = %v", err)
}
if modelID != "claude-sonnet-4.6" {
t.Fatalf("modelID = %q, want %q", modelID, "claude-sonnet-4.6")
}
if _, ok := provider.(*toolSchemaTransformProvider); !ok {
t.Fatalf("provider = %T, want *toolSchemaTransformProvider", provider)
}
}
func TestCreateProviderFromConfig_InvalidToolSchemaTransform(t *testing.T) {
cfg := &config.ModelConfig{
ModelName: "claude-cli-test",
Provider: "claude-cli",
Model: "claude-sonnet-4.6",
Workspace: t.TempDir(),
ToolSchemaTransform: "invalid",
}
_, _, err := CreateProviderFromConfig(cfg)
if err == nil {
t.Fatal("CreateProviderFromConfig() expected error for invalid tool_schema_transform")
}
if !strings.Contains(err.Error(), "tool_schema_transform") {
t.Fatalf("error = %v, want mention tool_schema_transform", err)
}
}
+1 -1
View File
@@ -264,7 +264,7 @@ func (p *GeminiProvider) buildRequestBody(
funcDecls = append(funcDecls, geminiFunctionDeclaration{
Name: t.Function.Name,
Description: t.Function.Description,
Parameters: common.SanitizeSchemaForGemini(t.Function.Parameters),
Parameters: t.Function.Parameters,
})
}
if len(funcDecls) > 0 {
@@ -5,11 +5,8 @@ import (
"fmt"
"net/http"
"net/http/httptest"
"reflect"
"strings"
"testing"
providercommon "github.com/sipeed/picoclaw/pkg/providers/common"
)
func TestGeminiProvider_ChatSeparatesThoughtAndToolCall(t *testing.T) {
@@ -262,7 +259,7 @@ func TestGeminiProvider_ChatStreamSkipsEmptyDataFrames(t *testing.T) {
}
}
func TestGeminiProvider_BuildRequestBody_SanitizesComplexToolSchemas(t *testing.T) {
func TestGeminiProvider_BuildRequestBody_PreservesComplexToolSchemasByDefault(t *testing.T) {
provider := NewGeminiProvider("test-key", "https://example.com/v1beta", "", "", 0, nil, nil)
schema := map[string]any{
"type": "object",
@@ -315,9 +312,8 @@ func TestGeminiProvider_BuildRequestBody_SanitizesComplexToolSchemas(t *testing.
t.Fatalf("parameters = %#v, want map", tools[0].FunctionDeclarations[0].Parameters)
}
want := providercommon.SanitizeSchemaForGemini(schema)
if !reflect.DeepEqual(got, want) {
t.Fatalf("sanitized parameters mismatch\n got: %#v\nwant: %#v", got, want)
if got["$defs"] == nil {
t.Fatalf("parameters = %#v, want raw schema with $defs preserved by default", got)
}
}
+2 -3
View File
@@ -291,18 +291,17 @@ func (p *AntigravityProvider) buildRequest(
}
}
// Build tools (sanitize schemas for Gemini compatibility)
// Build tools
if len(tools) > 0 {
var funcDecls []antigravityFuncDecl
for _, t := range tools {
if t.Type != "function" {
continue
}
params := common.SanitizeSchemaForGemini(t.Function.Parameters)
funcDecls = append(funcDecls, antigravityFuncDecl{
Name: t.Function.Name,
Description: t.Function.Description,
Parameters: params,
Parameters: t.Function.Parameters,
})
}
if len(funcDecls) > 0 {
@@ -1,10 +1,7 @@
package oauthprovider
import (
"reflect"
"testing"
providercommon "github.com/sipeed/picoclaw/pkg/providers/common"
)
func TestBuildRequestUsesFunctionFieldsWhenToolCallNameMissing(t *testing.T) {
@@ -77,7 +74,7 @@ func TestParseSSEResponse_SplitsThoughtAndVisibleContent(t *testing.T) {
}
}
func TestBuildRequest_SanitizesComplexToolSchemas(t *testing.T) {
func TestBuildRequest_PreservesComplexToolSchemasByDefault(t *testing.T) {
p := &AntigravityProvider{}
schema := map[string]any{
"type": "object",
@@ -135,9 +132,11 @@ func TestBuildRequest_SanitizesComplexToolSchemas(t *testing.T) {
t.Fatalf("request tools = %#v, want one function declaration", req.Tools)
}
got := req.Tools[0].FunctionDeclarations[0].Parameters
want := providercommon.SanitizeSchemaForGemini(schema)
if !reflect.DeepEqual(got, want) {
t.Fatalf("sanitized parameters mismatch\n got: %#v\nwant: %#v", got, want)
got, ok := req.Tools[0].FunctionDeclarations[0].Parameters.(map[string]any)
if !ok {
t.Fatalf("parameters = %#v, want map", req.Tools[0].FunctionDeclarations[0].Parameters)
}
if got["$defs"] == nil {
t.Fatalf("parameters = %#v, want raw schema with $defs preserved by default", got)
}
}
+84
View File
@@ -0,0 +1,84 @@
package providers
import (
"context"
"github.com/sipeed/picoclaw/pkg/providers/common"
)
type toolSchemaTransformProvider struct {
delegate LLMProvider
transform string
}
type toolSchemaStreamingProvider struct {
*toolSchemaTransformProvider
}
func wrapProviderWithToolSchemaTransform(delegate LLMProvider, transform string) (LLMProvider, error) {
transform, err := common.NormalizeToolSchemaTransform(transform)
if err != nil {
return nil, err
}
if transform == common.ToolSchemaTransformOff || delegate == nil {
return delegate, nil
}
base := &toolSchemaTransformProvider{
delegate: delegate,
transform: transform,
}
if _, ok := delegate.(StreamingProvider); ok {
return &toolSchemaStreamingProvider{toolSchemaTransformProvider: base}, nil
}
return base, nil
}
func (p *toolSchemaTransformProvider) Chat(
ctx context.Context,
messages []Message,
tools []ToolDefinition,
model string,
options map[string]any,
) (*LLMResponse, error) {
transformed, err := common.TransformToolDefinitions(tools, p.transform)
if err != nil {
return nil, err
}
return p.delegate.Chat(ctx, messages, transformed, model, options)
}
func (p *toolSchemaTransformProvider) GetDefaultModel() string {
return p.delegate.GetDefaultModel()
}
func (p *toolSchemaStreamingProvider) ChatStream(
ctx context.Context,
messages []Message,
tools []ToolDefinition,
model string,
options map[string]any,
onChunk func(accumulated string),
) (*LLMResponse, error) {
streaming := p.delegate.(StreamingProvider)
transformed, err := common.TransformToolDefinitions(tools, p.transform)
if err != nil {
return nil, err
}
return streaming.ChatStream(ctx, messages, transformed, model, options, onChunk)
}
func (p *toolSchemaTransformProvider) SupportsThinking() bool {
tc, ok := p.delegate.(ThinkingCapable)
return ok && tc.SupportsThinking()
}
func (p *toolSchemaTransformProvider) SupportsNativeSearch() bool {
ns, ok := p.delegate.(NativeSearchCapable)
return ok && ns.SupportsNativeSearch()
}
func (p *toolSchemaTransformProvider) Close() {
if stateful, ok := p.delegate.(StatefulProvider); ok {
stateful.Close()
}
}
+104
View File
@@ -0,0 +1,104 @@
package providers
import (
"context"
"reflect"
"testing"
providercommon "github.com/sipeed/picoclaw/pkg/providers/common"
)
type toolCaptureProvider struct {
lastTools []ToolDefinition
}
func (p *toolCaptureProvider) Chat(
_ context.Context,
_ []Message,
tools []ToolDefinition,
_ string,
_ map[string]any,
) (*LLMResponse, error) {
p.lastTools = tools
return &LLMResponse{Content: "ok"}, nil
}
func (p *toolCaptureProvider) GetDefaultModel() string {
return "test"
}
func TestWrapProviderWithToolSchemaTransform_DisabledPassesToolsThrough(t *testing.T) {
capture := &toolCaptureProvider{}
wrapped, err := wrapProviderWithToolSchemaTransform(capture, "")
if err != nil {
t.Fatalf("wrapProviderWithToolSchemaTransform() error = %v", err)
}
tools := []ToolDefinition{{
Type: "function",
Function: ToolFunctionDefinition{
Name: "noop",
Parameters: map[string]any{"type": "object"},
},
}}
_, err = wrapped.Chat(t.Context(), nil, tools, "test", nil)
if err != nil {
t.Fatalf("Chat() error = %v", err)
}
if !reflect.DeepEqual(capture.lastTools, tools) {
t.Fatalf("tools mutated with transform off\n got: %#v\nwant: %#v", capture.lastTools, tools)
}
}
func TestWrapProviderWithToolSchemaTransform_GoogleSanitizesSchemas(t *testing.T) {
capture := &toolCaptureProvider{}
wrapped, err := wrapProviderWithToolSchemaTransform(capture, "google")
if err != nil {
t.Fatalf("wrapProviderWithToolSchemaTransform() error = %v", err)
}
schema := map[string]any{
"type": "object",
"properties": map[string]any{
"parent": map[string]any{
"anyOf": []any{
map[string]any{"$ref": "#/$defs/pageParent"},
map[string]any{"$ref": "#/$defs/databaseParent"},
},
},
},
"$defs": map[string]any{
"pageParent": map[string]any{
"type": "object",
"properties": map[string]any{
"page_id": map[string]any{"type": "string"},
},
},
"databaseParent": map[string]any{
"type": "object",
"properties": map[string]any{
"database_id": map[string]any{"type": "string"},
},
},
},
}
tools := []ToolDefinition{{
Type: "function",
Function: ToolFunctionDefinition{
Name: "mcp_notion_create",
Parameters: schema,
},
}}
_, err = wrapped.Chat(t.Context(), nil, tools, "test", nil)
if err != nil {
t.Fatalf("Chat() error = %v", err)
}
want := providercommon.SanitizeSchemaForGoogle(schema)
got := capture.lastTools[0].Function.Parameters
if !reflect.DeepEqual(got, want) {
t.Fatalf("sanitized parameters mismatch\n got: %#v\nwant: %#v", got, want)
}
}
+34 -29
View File
@@ -35,14 +35,15 @@ type modelResponse struct {
Proxy string `json:"proxy,omitempty"`
AuthMethod string `json:"auth_method,omitempty"`
// Advanced fields
ConnectMode string `json:"connect_mode,omitempty"`
Workspace string `json:"workspace,omitempty"`
RPM int `json:"rpm,omitempty"`
MaxTokensField string `json:"max_tokens_field,omitempty"`
RequestTimeout int `json:"request_timeout,omitempty"`
ThinkingLevel string `json:"thinking_level,omitempty"`
ExtraBody map[string]any `json:"extra_body,omitempty"`
CustomHeaders map[string]string `json:"custom_headers,omitempty"`
ConnectMode string `json:"connect_mode,omitempty"`
Workspace string `json:"workspace,omitempty"`
RPM int `json:"rpm,omitempty"`
MaxTokensField string `json:"max_tokens_field,omitempty"`
RequestTimeout int `json:"request_timeout,omitempty"`
ThinkingLevel string `json:"thinking_level,omitempty"`
ToolSchemaTransform string `json:"tool_schema_transform,omitempty"`
ExtraBody map[string]any `json:"extra_body,omitempty"`
CustomHeaders map[string]string `json:"custom_headers,omitempty"`
// Meta
Enabled bool `json:"enabled"`
Available bool `json:"available"`
@@ -78,27 +79,28 @@ func (h *Handler) handleListModels(w http.ResponseWriter, r *http.Request) {
for i, m := range cfg.ModelList {
provider, modelID := providers.ExtractProtocol(m)
models = append(models, modelResponse{
Index: i,
ModelName: m.ModelName,
Provider: provider,
Model: modelID,
APIBase: m.APIBase,
APIKey: maskAPIKey(m.APIKey()),
Proxy: m.Proxy,
AuthMethod: m.AuthMethod,
ConnectMode: m.ConnectMode,
Workspace: m.Workspace,
RPM: m.RPM,
MaxTokensField: m.MaxTokensField,
RequestTimeout: m.RequestTimeout,
ThinkingLevel: m.ThinkingLevel,
ExtraBody: m.ExtraBody,
CustomHeaders: m.CustomHeaders,
Enabled: m.Enabled,
Available: modelStatuses[i].Available,
Status: modelStatuses[i].Status,
IsDefault: m.ModelName == defaultModel,
IsVirtual: m.IsVirtual(),
Index: i,
ModelName: m.ModelName,
Provider: provider,
Model: modelID,
APIBase: m.APIBase,
APIKey: maskAPIKey(m.APIKey()),
Proxy: m.Proxy,
AuthMethod: m.AuthMethod,
ConnectMode: m.ConnectMode,
Workspace: m.Workspace,
RPM: m.RPM,
MaxTokensField: m.MaxTokensField,
RequestTimeout: m.RequestTimeout,
ThinkingLevel: m.ThinkingLevel,
ToolSchemaTransform: m.ToolSchemaTransform,
ExtraBody: m.ExtraBody,
CustomHeaders: m.CustomHeaders,
Enabled: m.Enabled,
Available: modelStatuses[i].Available,
Status: modelStatuses[i].Status,
IsDefault: m.ModelName == defaultModel,
IsVirtual: m.IsVirtual(),
})
}
@@ -237,6 +239,9 @@ func (h *Handler) handleUpdateModel(w http.ResponseWriter, r *http.Request) {
} else if len(mc.CustomHeaders) == 0 {
mc.CustomHeaders = nil
}
if _, ok := rawFields["tool_schema_transform"]; !ok {
mc.ToolSchemaTransform = cfg.ModelList[idx].ToolSchemaTransform
}
// Preserve the existing Provider when the caller omits it. This keeps the
// update API backward-compatible for clients that haven't started sending
// the new field yet, while still allowing explicit clearing via "".
+94
View File
@@ -584,6 +584,37 @@ func TestHandleAddModel_PersistsCustomHeaders(t *testing.T) {
}
}
func TestHandleAddModel_PersistsToolSchemaTransform(t *testing.T) {
configPath, cleanup := setupOAuthTestEnv(t)
defer cleanup()
h := NewHandler(configPath)
mux := http.NewServeMux()
h.RegisterRoutes(mux)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/models", bytes.NewBufferString(`{
"model_name":"new-model-transform",
"model":"openai/gpt-4o-mini",
"tool_schema_transform":"simple"
}`))
req.Header.Set("Content-Type", "application/json")
mux.ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
}
cfg, err := config.LoadConfig(configPath)
if err != nil {
t.Fatalf("LoadConfig() error = %v", err)
}
added := cfg.ModelList[len(cfg.ModelList)-1]
if got := added.ToolSchemaTransform; got != "simple" {
t.Fatalf("tool_schema_transform = %q, want %q", got, "simple")
}
}
func TestHandleUpdateModel_CustomHeadersPreserveAndClear(t *testing.T) {
configPath, cleanup := setupOAuthTestEnv(t)
defer cleanup()
@@ -649,6 +680,69 @@ func TestHandleUpdateModel_CustomHeadersPreserveAndClear(t *testing.T) {
}
}
func TestHandleUpdateModel_ToolSchemaTransformPreserveAndClear(t *testing.T) {
configPath, cleanup := setupOAuthTestEnv(t)
defer cleanup()
cfg, err := config.LoadConfig(configPath)
if err != nil {
t.Fatalf("LoadConfig() error = %v", err)
}
cfg.ModelList = []*config.ModelConfig{{
ModelName: "editable",
Model: "openai/gpt-4o-mini",
APIKeys: config.SimpleSecureStrings("sk-existing"),
ToolSchemaTransform: "google",
}}
err = config.SaveConfig(configPath, cfg)
if err != nil {
t.Fatalf("SaveConfig() error = %v", err)
}
h := NewHandler(configPath)
mux := http.NewServeMux()
h.RegisterRoutes(mux)
recPreserve := httptest.NewRecorder()
reqPreserve := httptest.NewRequest(http.MethodPut, "/api/models/0", bytes.NewBufferString(`{
"model_name":"editable",
"model":"openai/gpt-4o-mini"
}`))
reqPreserve.Header.Set("Content-Type", "application/json")
mux.ServeHTTP(recPreserve, reqPreserve)
if recPreserve.Code != http.StatusOK {
t.Fatalf("preserve status = %d, want %d, body=%s", recPreserve.Code, http.StatusOK, recPreserve.Body.String())
}
afterPreserve, err := config.LoadConfig(configPath)
if err != nil {
t.Fatalf("LoadConfig() after preserve error = %v", err)
}
if got := afterPreserve.ModelList[0].ToolSchemaTransform; got != "google" {
t.Fatalf("preserved tool_schema_transform = %q, want %q", got, "google")
}
recClear := httptest.NewRecorder()
reqClear := httptest.NewRequest(http.MethodPut, "/api/models/0", bytes.NewBufferString(`{
"model_name":"editable",
"model":"openai/gpt-4o-mini",
"tool_schema_transform":""
}`))
reqClear.Header.Set("Content-Type", "application/json")
mux.ServeHTTP(recClear, reqClear)
if recClear.Code != http.StatusOK {
t.Fatalf("clear status = %d, want %d, body=%s", recClear.Code, http.StatusOK, recClear.Body.String())
}
afterClear, err := config.LoadConfig(configPath)
if err != nil {
t.Fatalf("LoadConfig() after clear error = %v", err)
}
if afterClear.ModelList[0].ToolSchemaTransform != "" {
t.Fatalf("tool_schema_transform = %q, want empty", afterClear.ModelList[0].ToolSchemaTransform)
}
}
func TestHandleUpdateModel_PersistsProvider(t *testing.T) {
configPath, cleanup := setupOAuthTestEnv(t)
defer cleanup()
+1
View File
@@ -19,6 +19,7 @@ export interface ModelInfo {
max_tokens_field?: string
request_timeout?: number
thinking_level?: string
tool_schema_transform?: string
extra_body?: Record<string, unknown>
custom_headers?: Record<string, string>
// Meta
@@ -36,6 +36,7 @@ interface AddForm {
maxTokensField: string
requestTimeout: string
thinkingLevel: string
toolSchemaTransform: string
extraBody: string
customHeaders: string
}
@@ -54,6 +55,7 @@ const EMPTY_ADD_FORM: AddForm = {
maxTokensField: "",
requestTimeout: "",
thinkingLevel: "",
toolSchemaTransform: "",
extraBody: "",
customHeaders: "",
}
@@ -139,6 +141,7 @@ export function AddModelSheet({
? Number(form.requestTimeout)
: undefined,
thinking_level: form.thinkingLevel.trim() || undefined,
tool_schema_transform: form.toolSchemaTransform.trim() || undefined,
extra_body: form.extraBody.trim()
? JSON.parse(form.extraBody.trim())
: undefined,
@@ -333,6 +336,17 @@ export function AddModelSheet({
/>
</Field>
<Field
label={t("models.field.toolSchemaTransform")}
hint={t("models.field.toolSchemaTransformHint")}
>
<Input
value={form.toolSchemaTransform}
onChange={setField("toolSchemaTransform")}
placeholder="google"
/>
</Field>
<Field
label={t("models.field.extraBody")}
hint={t("models.field.extraBodyHint")}
@@ -35,6 +35,7 @@ interface EditForm {
maxTokensField: string
requestTimeout: string
thinkingLevel: string
toolSchemaTransform: string
extraBody: string
customHeaders: string
}
@@ -66,6 +67,7 @@ export function EditModelSheet({
maxTokensField: "",
requestTimeout: "",
thinkingLevel: "",
toolSchemaTransform: "",
extraBody: "",
customHeaders: "",
})
@@ -90,6 +92,7 @@ export function EditModelSheet({
? String(model.request_timeout)
: "",
thinkingLevel: model.thinking_level ?? "",
toolSchemaTransform: model.tool_schema_transform ?? "",
extraBody: model.extra_body
? JSON.stringify(model.extra_body, null, 2)
: "",
@@ -132,6 +135,7 @@ export function EditModelSheet({
? Number(form.requestTimeout)
: undefined,
thinking_level: form.thinkingLevel || undefined,
tool_schema_transform: form.toolSchemaTransform.trim() || undefined,
extra_body: form.extraBody.trim()
? JSON.parse(form.extraBody.trim())
: {},
@@ -325,6 +329,17 @@ export function EditModelSheet({
/>
</Field>
<Field
label={t("models.field.toolSchemaTransform")}
hint={t("models.field.toolSchemaTransformHint")}
>
<Input
value={form.toolSchemaTransform}
onChange={setField("toolSchemaTransform")}
placeholder="google"
/>
</Field>
<Field
label={t("models.field.extraBody")}
hint={t("models.field.extraBodyHint")}
+2
View File
@@ -287,6 +287,8 @@
"thinkingLevelHint": "Extended thinking budget: off, low, medium, high, xhigh, adaptive.",
"maxTokensField": "Max Tokens Field",
"maxTokensFieldHint": "Override the request field name for max tokens, e.g. max_completion_tokens.",
"toolSchemaTransform": "Tool Schema Transform",
"toolSchemaTransformHint": "Optional compatibility transform for tool JSON schemas. Leave blank for native behavior. Supported values: simple.",
"extraBody": "Extra Body",
"extraBodyHint": "Additional JSON fields to inject into the request body, e.g. {\"reasoning_split\": true}.",
"customHeaders": "Custom Headers",
+2
View File
@@ -287,6 +287,8 @@
"thinkingLevelHint": "扩展思考预算:off、low、medium、high、xhigh、adaptive。",
"maxTokensField": "Max Tokens 字段名",
"maxTokensFieldHint": "覆盖请求中 max_tokens 的字段名,例如 max_completion_tokens。",
"toolSchemaTransform": "工具 Schema 转换",
"toolSchemaTransformHint": "可选的工具 JSON Schema 兼容性转换。留空表示保持原生行为。当前支持值:simple。",
"extraBody": "Extra Body",
"extraBodyHint": "要注入到请求体中的额外 JSON 字段,例如 {\"reasoning_split\": true}。",
"customHeaders": "Custom Headers",