mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: preserve gemini thought signatures (#25933)
AI Bridge reserializes OpenAI chat-completions requests before sending them upstream. For Gemini OpenAI-compatible routes, that OpenAI typed-parameter round trip drops `tool_calls[].extra_content.google.thought_signature`, so Google rejects tool-result continuations with `Function call is missing a thought_signature`. This PR: - patches the AI Bridge upstream serialization boundary for Gemini OpenAI-compatible chat completions - shares the Gemini thought-signature patching helpers with chatd's OpenAI-compatible transport patch to keep behavior consistent - treats direct Google OpenAI-compatible upstream endpoints as Gemini-scoped even when the request model is an alias - adds the Google fallback thought signature to every assistant tool call in the active turn, including parallel tool calls - covers the regression that `extra_content` is dropped before the upstream body is patched > Mux updated this PR description on behalf of Mike. --------- Co-authored-by: Susana Cardoso Ferreira <susana@coder.com>
This commit is contained in:
co-authored by
Susana Cardoso Ferreira
parent
a9fb2619e4
commit
c349ea6b78
@@ -0,0 +1,163 @@
|
||||
// Package googleopenai contains compatibility helpers for Google's
|
||||
// OpenAI-compatible Gemini APIs.
|
||||
package googleopenai
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// DummyThoughtSignature is Google's documented last-resort bypass for callers
|
||||
// that cannot preserve a real Gemini thought signature through OpenAI-compatible
|
||||
// serialization. See https://ai.google.dev/gemini-api/docs/thought-signatures.
|
||||
const DummyThoughtSignature = "skip_thought_signature_validator"
|
||||
|
||||
// ShouldPatchOpenAICompatRequest reports whether a client-side
|
||||
// OpenAI-compatible request should carry Gemini thought signatures.
|
||||
func ShouldPatchOpenAICompatRequest(baseURL string, modelID string) bool {
|
||||
// Direct Google endpoints are already provider-scoped. Patch them even when
|
||||
// the configured model ID is an alias without a Gemini prefix.
|
||||
if isDirectGeminiOpenAIEndpoint(baseURL) {
|
||||
return true
|
||||
}
|
||||
return isCoderAIBridgeEndpoint(baseURL) && isGeminiModelID(modelID)
|
||||
}
|
||||
|
||||
// ShouldPatchGoogleUpstreamRequest reports whether an AI Bridge upstream
|
||||
// OpenAI-compatible request should carry Gemini thought signatures.
|
||||
func ShouldPatchGoogleUpstreamRequest(baseURL string) bool {
|
||||
return isDirectGeminiOpenAIEndpoint(baseURL)
|
||||
}
|
||||
|
||||
// Vertex AI has different hosts and paths. Add it here only with a fixture that
|
||||
// confirms it accepts the same thought-signature fallback shape.
|
||||
func isDirectGeminiOpenAIEndpoint(baseURL string) bool {
|
||||
parsed, ok := parseBaseURL(baseURL)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
host := strings.ToLower(parsed.Hostname())
|
||||
path := strings.ToLower(parsed.EscapedPath())
|
||||
return host == "generativelanguage.googleapis.com" && strings.Contains(path, "/openai")
|
||||
}
|
||||
|
||||
func isCoderAIBridgeEndpoint(baseURL string) bool {
|
||||
parsed, ok := parseBaseURL(baseURL)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return strings.ToLower(parsed.Hostname()) == "coder-aibridge"
|
||||
}
|
||||
|
||||
// parseBaseURL parses a provider base URL, handling bare hostnames without
|
||||
// a scheme by prepending "https://".
|
||||
func parseBaseURL(baseURL string) (*url.URL, bool) {
|
||||
baseURL = strings.TrimSpace(baseURL)
|
||||
if baseURL == "" {
|
||||
return nil, false
|
||||
}
|
||||
parsed, err := url.Parse(baseURL)
|
||||
if err == nil && parsed.Hostname() == "" && !strings.Contains(baseURL, "://") {
|
||||
parsed, err = url.Parse("https://" + baseURL)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, false
|
||||
}
|
||||
return parsed, true
|
||||
}
|
||||
|
||||
func isGeminiModelID(modelID string) bool {
|
||||
modelID = strings.ToLower(strings.TrimSpace(modelID))
|
||||
return strings.HasPrefix(modelID, "gemini-") || strings.Contains(modelID, "/gemini-")
|
||||
}
|
||||
|
||||
// PatchThoughtSignatures adds fallback thought signatures to Gemini tool-call
|
||||
// history in body. It returns changed=false when no patch is needed.
|
||||
func PatchThoughtSignatures(body []byte) ([]byte, bool, error) {
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
if !AddThoughtSignaturesToLatestTurn(payload) {
|
||||
return body, false, nil
|
||||
}
|
||||
patched, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
return patched, true, nil
|
||||
}
|
||||
|
||||
// AddThoughtSignaturesToLatestTurn patches only the current turn because
|
||||
// completed tool-call/result pairs from earlier turns are not validated by
|
||||
// Google as active function calls.
|
||||
func AddThoughtSignaturesToLatestTurn(payload map[string]any) bool {
|
||||
messages, ok := payload["messages"].([]any)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
|
||||
currentTurnStart := -1
|
||||
for i, raw := range messages {
|
||||
message, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if role, _ := message["role"].(string); role == "user" {
|
||||
currentTurnStart = i
|
||||
}
|
||||
}
|
||||
if currentTurnStart == -1 {
|
||||
return false
|
||||
}
|
||||
|
||||
changed := false
|
||||
for _, raw := range messages[currentTurnStart+1:] {
|
||||
message, ok := raw.(map[string]any)
|
||||
if !ok || !isAssistantRole(message["role"]) {
|
||||
continue
|
||||
}
|
||||
toolCalls, ok := message["tool_calls"].([]any)
|
||||
if !ok || len(toolCalls) == 0 {
|
||||
continue
|
||||
}
|
||||
// Every tool call in parallel batches needs a signature,
|
||||
// not just the first one.
|
||||
for _, rawToolCall := range toolCalls {
|
||||
toolCall, ok := rawToolCall.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if ensureThoughtSignature(toolCall) {
|
||||
changed = true
|
||||
}
|
||||
}
|
||||
}
|
||||
return changed
|
||||
}
|
||||
|
||||
// Gemini can serialize assistant messages with its native "model" role.
|
||||
func isAssistantRole(role any) bool {
|
||||
roleValue, _ := role.(string)
|
||||
return roleValue == "assistant" || roleValue == "model"
|
||||
}
|
||||
|
||||
// Real provider signatures are preserved when present.
|
||||
func ensureThoughtSignature(toolCall map[string]any) bool {
|
||||
extraContent, _ := toolCall["extra_content"].(map[string]any)
|
||||
google, _ := extraContent["google"].(map[string]any)
|
||||
if signature, _ := google["thought_signature"].(string); signature != "" {
|
||||
return false
|
||||
}
|
||||
if extraContent == nil {
|
||||
extraContent = map[string]any{}
|
||||
toolCall["extra_content"] = extraContent
|
||||
}
|
||||
if google == nil {
|
||||
google = map[string]any{}
|
||||
extraContent["google"] = google
|
||||
}
|
||||
google["thought_signature"] = DummyThoughtSignature
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
package googleopenai_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/internal/googleopenai"
|
||||
)
|
||||
|
||||
func TestShouldPatchOpenAICompatRequest(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
baseURL string
|
||||
modelID string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "direct endpoint with gemini model",
|
||||
baseURL: "https://generativelanguage.googleapis.com/v1beta/openai/",
|
||||
modelID: "gemini-3.5-flash",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "direct endpoint does not require gemini model name",
|
||||
baseURL: "https://generativelanguage.googleapis.com/v1beta/openai/",
|
||||
modelID: "gpt-4o",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "coder aibridge gemini route",
|
||||
baseURL: "http://coder-aibridge/v1",
|
||||
modelID: "gemini-3.5-flash",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "aibridge endpoint requires gemini model",
|
||||
baseURL: "http://coder-aibridge/v1",
|
||||
modelID: "gpt-4o",
|
||||
},
|
||||
{
|
||||
name: "other gateway unchanged",
|
||||
baseURL: "https://gateway.vercel.ai/v1",
|
||||
modelID: "google/gemini-3.5-flash",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.Equal(t, tt.want, googleopenai.ShouldPatchOpenAICompatRequest(tt.baseURL, tt.modelID))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestShouldPatchGoogleUpstreamRequest(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
baseURL string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "gemini api openai endpoint",
|
||||
baseURL: "https://generativelanguage.googleapis.com/v1beta/openai/",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "openai endpoint",
|
||||
baseURL: "https://api.openai.com/v1/",
|
||||
},
|
||||
{
|
||||
name: "vertex endpoint not enabled without fixture",
|
||||
baseURL: "https://us-central1-aiplatform.googleapis.com/v1/",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.Equal(t, tt.want, googleopenai.ShouldPatchGoogleUpstreamRequest(tt.baseURL))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddThoughtSignaturesToLatestTurn(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
payload := decodePayload(t, []byte(`{
|
||||
"messages":[
|
||||
{"role":"user","content":"previous turn"},
|
||||
{
|
||||
"role":"assistant",
|
||||
"tool_calls":[{"id":"old-call","type":"function","function":{"name":"old","arguments":"{}"}}]
|
||||
},
|
||||
{"role":"tool","tool_call_id":"old-call","content":"{}"},
|
||||
{"role":"user","content":"current turn"},
|
||||
{
|
||||
"role":"model",
|
||||
"tool_calls":[
|
||||
{"id":"call-1","type":"function","function":{"name":"list_templates","arguments":"{}"}},
|
||||
{"id":"call-2","type":"function","function":{"name":"read_template","arguments":"{}"}}
|
||||
]
|
||||
}
|
||||
]
|
||||
}`))
|
||||
|
||||
require.True(t, googleopenai.AddThoughtSignaturesToLatestTurn(payload))
|
||||
require.Empty(t, thoughtSignature(t, payload, 1, 0), "previous turns should stay unchanged")
|
||||
require.Equal(t, googleopenai.DummyThoughtSignature, thoughtSignature(t, payload, 4, 0))
|
||||
require.Equal(t, googleopenai.DummyThoughtSignature, thoughtSignature(t, payload, 4, 1))
|
||||
}
|
||||
|
||||
func TestAddThoughtSignaturesToLatestTurnPreservesRealSignature(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
payload := decodePayload(t, []byte(`{
|
||||
"messages":[
|
||||
{"role":"user","content":"current turn"},
|
||||
{
|
||||
"role":"assistant",
|
||||
"tool_calls":[{
|
||||
"id":"call-1",
|
||||
"type":"function",
|
||||
"function":{"name":"list_templates","arguments":"{}"},
|
||||
"extra_content":{"google":{"thought_signature":"real-signature"}}
|
||||
}]
|
||||
}
|
||||
]
|
||||
}`))
|
||||
|
||||
require.False(t, googleopenai.AddThoughtSignaturesToLatestTurn(payload))
|
||||
require.Equal(t, "real-signature", thoughtSignature(t, payload, 1, 0))
|
||||
}
|
||||
|
||||
func decodePayload(t *testing.T, body []byte) map[string]any {
|
||||
t.Helper()
|
||||
|
||||
var payload map[string]any
|
||||
require.NoError(t, json.Unmarshal(body, &payload))
|
||||
return payload
|
||||
}
|
||||
|
||||
func thoughtSignature(t *testing.T, payload map[string]any, messageIndex int, toolCallIndex int) string {
|
||||
t.Helper()
|
||||
|
||||
messages, ok := payload["messages"].([]any)
|
||||
require.True(t, ok)
|
||||
require.Greater(t, len(messages), messageIndex)
|
||||
message, ok := messages[messageIndex].(map[string]any)
|
||||
require.True(t, ok)
|
||||
toolCalls, ok := message["tool_calls"].([]any)
|
||||
require.True(t, ok)
|
||||
require.Greater(t, len(toolCalls), toolCallIndex)
|
||||
toolCall, ok := toolCalls[toolCallIndex].(map[string]any)
|
||||
require.True(t, ok)
|
||||
extraContent, _ := toolCall["extra_content"].(map[string]any)
|
||||
google, _ := extraContent["google"].(map[string]any)
|
||||
signature, _ := google["thought_signature"].(string)
|
||||
return signature
|
||||
}
|
||||
Reference in New Issue
Block a user