address review feedback: handle non-native strict mode edge case and add unit tests
This commit is contained in:
parent
5675cf31c8
commit
0986e2e4b5
2 changed files with 107 additions and 3 deletions
|
|
@ -545,13 +545,23 @@ func (p *Provider) finalizeTools(tools []ToolDefinition, options map[string]any)
|
|||
return tools
|
||||
}
|
||||
|
||||
// 3. Sanitization / Explicit Force Path
|
||||
// Build compatible map based on 'useStrict' decision
|
||||
// 3. Sanitization / Explicit Force Path / Edge Case Handling
|
||||
// Decide whether to use strict mode
|
||||
useStrict := isNative
|
||||
if hasForce {
|
||||
if forceStrict && !isNative {
|
||||
// Non-native providers don't support strict mode, ignore user setting
|
||||
log.Printf(
|
||||
"openai_compat: strict_mode=true ignored for non-OpenAI provider %q (unsupported field)",
|
||||
p.apiBase,
|
||||
)
|
||||
useStrict = false
|
||||
} else {
|
||||
useStrict = forceStrict
|
||||
}
|
||||
}
|
||||
|
||||
// Build compatible map based on 'useStrict' decision
|
||||
out := make([]any, 0, len(tools))
|
||||
for _, t := range tools {
|
||||
toolMap := map[string]any{
|
||||
|
|
|
|||
|
|
@ -841,3 +841,97 @@ func TestSerializeMessages_StripsSystemParts(t *testing.T) {
|
|||
t.Fatal("system_parts should not appear in serialized output")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinalizeTools_StrictMode(t *testing.T) {
|
||||
tools := []ToolDefinition{
|
||||
{
|
||||
Type: "function",
|
||||
Function: ToolFunctionDefinition{
|
||||
Name: "test_func",
|
||||
Description: "testing",
|
||||
Parameters: map[string]any{"type": "object"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
apiBase string
|
||||
options map[string]any
|
||||
wantStrict bool
|
||||
wantSame bool // whether it should be a direct pass-through
|
||||
}{
|
||||
{
|
||||
name: "native openai - default",
|
||||
apiBase: "https://api.openai.com/v1",
|
||||
options: nil,
|
||||
wantStrict: false,
|
||||
wantSame: true,
|
||||
},
|
||||
{
|
||||
name: "native openai - force strict",
|
||||
apiBase: "https://api.openai.com/v1",
|
||||
options: map[string]any{"strict_mode": true},
|
||||
wantStrict: true,
|
||||
wantSame: false,
|
||||
},
|
||||
{
|
||||
name: "native openai - force off",
|
||||
apiBase: "https://api.openai.com/v1",
|
||||
options: map[string]any{"strict_mode": false},
|
||||
wantStrict: false,
|
||||
wantSame: false,
|
||||
},
|
||||
{
|
||||
name: "non-native - default (stripped)",
|
||||
apiBase: "https://api.deepseek.com/v1",
|
||||
options: nil,
|
||||
wantStrict: false,
|
||||
wantSame: false,
|
||||
},
|
||||
{
|
||||
name: "non-native - force strict (ignored)",
|
||||
apiBase: "https://api.groq.com/openai/v1",
|
||||
options: map[string]any{"strict_mode": true},
|
||||
wantStrict: false,
|
||||
wantSame: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
p := &Provider{apiBase: tt.apiBase}
|
||||
res := p.finalizeTools(tools, tt.options)
|
||||
|
||||
if tt.wantSame {
|
||||
if fmt.Sprintf("%p", res) != fmt.Sprintf("%p", tools) {
|
||||
t.Errorf("expected pass-through, got new slice")
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Verify if it's a slice of maps with/without strict
|
||||
resSlice, ok := res.([]any)
|
||||
if !ok {
|
||||
t.Fatalf("expected []any, got %T", res)
|
||||
}
|
||||
if len(resSlice) != 1 {
|
||||
t.Fatalf("expected 1 tool, got %d", len(resSlice))
|
||||
}
|
||||
|
||||
toolMap := resSlice[0].(map[string]any)
|
||||
fnMap := toolMap["function"].(map[string]any)
|
||||
val, hasStrict := fnMap["strict"]
|
||||
|
||||
if tt.wantStrict {
|
||||
if !hasStrict || val != true {
|
||||
t.Errorf("expected strict: true, got %v", val)
|
||||
}
|
||||
} else {
|
||||
if hasStrict {
|
||||
t.Errorf("did not expect strict field, got %v", val)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue