diff --git a/pkg/httpclient/client.go b/pkg/httpclient/client.go index b0f4ae720..2bbd4ff3e 100644 --- a/pkg/httpclient/client.go +++ b/pkg/httpclient/client.go @@ -19,6 +19,10 @@ type HTTPOptions struct { Header http.Header Query url.Values + // dropSSEKeepaliveEvents enables keepalive-frame dropping in the SSE + // filter transport; see WithSSEKeepaliveFilter. + dropSSEKeepaliveEvents bool + // cagentID resolves the persistent install UUID stamped as // `X-Cagent-Id` on gateway-bound requests. It defaults to // [userid.Get]; tests inject their own source via @@ -52,7 +56,10 @@ func NewHTTPClient(ctx context.Context, opts ...Opt) *http.Client { var wrapped http.RoundTripper = &userAgentTransport{ httpOptions: httpOptions, - rt: &sseFilterTransport{base: rt}, + rt: &sseFilterTransport{ + base: rt, + dropKeepaliveEvents: httpOptions.dropSSEKeepaliveEvents, + }, } if httpOptions.refreshAuth != nil { // Outermost, so a replayed request goes through the whole chain again. @@ -187,6 +194,19 @@ func WithQuery(query url.Values) Opt { } } +// WithSSEKeepaliveFilter makes the SSE filter transport also drop whole +// `event: keepalive` frames whose data carries no payload (`data: {}` or +// empty). The Docker AI Gateway emits such frames during long generations +// (e.g. Gemini image output), and google.golang.org/genai's SSE parser +// treats any `event:` line as a fatal invalid chunk. Only enable this for +// clients whose SDK cannot tolerate `event:` lines — providers like +// Anthropic rely on `event:` headers as meaningful framing. +func WithSSEKeepaliveFilter() Opt { + return func(o *HTTPOptions) { + o.dropSSEKeepaliveEvents = true + } +} + // newTransport returns an HTTP transport with automatic gzip compression disabled and Docker Desktop PAC support. func newTransport(_ context.Context) http.RoundTripper { rt := newAllowPrivateIPsTransport() diff --git a/pkg/httpclient/sse_filter.go b/pkg/httpclient/sse_filter.go index e588ee8d5..093a47dd9 100644 --- a/pkg/httpclient/sse_filter.go +++ b/pkg/httpclient/sse_filter.go @@ -30,8 +30,17 @@ import ( // (comment-only events, or events bearing only `event:` / `id:` headers) // never reach the SDK. Well-formed events pass through verbatim, and the // filter is a no-op on non-SSE responses. +// +// dropKeepaliveEvents additionally drops whole `event: keepalive` frames +// whose data carries no payload (`data: {}` or empty). The Docker AI +// Gateway emits such frames during long generations; the genai SDK's SSE +// parser hard-fails on ANY `event:` line, so they must never reach it. The +// mode is opt-in (see WithSSEKeepaliveFilter) because other providers — +// Anthropic in particular — use `event:` headers as meaningful framing that +// must pass through untouched. type sseFilterTransport struct { - base http.RoundTripper + base http.RoundTripper + dropKeepaliveEvents bool } func (t *sseFilterTransport) RoundTrip(req *http.Request) (*http.Response, error) { @@ -42,7 +51,7 @@ func (t *sseFilterTransport) RoundTrip(req *http.Request) (*http.Response, error // Match the prefix so charset suffixes (e.g. "text/event-stream; // charset=utf-8") still trigger filtering. if strings.HasPrefix(strings.ToLower(res.Header.Get("Content-Type")), "text/event-stream") { - res.Body = newSSEFilterReader(res.Body) + res.Body = newSSEFilterReader(res.Body, t.dropKeepaliveEvents) } return res, err } @@ -58,15 +67,19 @@ type sseFilterReader struct { out bytes.Buffer // bytes ready to hand back to the caller pending bytes.Buffer // accumulated lines for the current event hasData bool // saw at least one `data:` line in `pending` + + dropKeepaliveEvents bool // see sseFilterTransport + isKeepalive bool // current event is named `keepalive` + hasMeaningfulData bool // saw a `data:` line whose payload isn't empty or `{}` } -func newSSEFilterReader(src io.ReadCloser) *sseFilterReader { +func newSSEFilterReader(src io.ReadCloser, dropKeepaliveEvents bool) *sseFilterReader { scn := bufio.NewScanner(src) // SSE events can be large (long completion tokens, image URLs, …). Match // the buffer size used by openai-go's own SSE decoder so we don't trip // `bufio.ErrTooLong` on payloads it would happily accept. scn.Buffer(make([]byte, 0, 64*1024), bufio.MaxScanTokenSize<<9) - return &sseFilterReader{src: src, scn: scn} + return &sseFilterReader{src: src, scn: scn, dropKeepaliveEvents: dropKeepaliveEvents} } func (r *sseFilterReader) Read(p []byte) (int, error) { @@ -85,24 +98,47 @@ func (r *sseFilterReader) Read(p []byte) (int, error) { func (r *sseFilterReader) consumeLine(line []byte) { switch { case len(line) == 0: - // Event boundary: emit the buffered event iff it had data. - if r.hasData { + // Event boundary: emit the buffered event iff it had data and is + // not a payload-free keepalive frame in keepalive-dropping mode. + if r.hasData && (!r.isKeepalive || r.hasMeaningfulData) { r.out.Write(r.pending.Bytes()) r.out.WriteByte('\n') } r.pending.Reset() r.hasData = false + r.isKeepalive = false + r.hasMeaningfulData = false case line[0] == ':': // SSE comment — drop entirely. default: r.pending.Write(line) r.pending.WriteByte('\n') - if bytes.HasPrefix(line, []byte("data:")) { + if value, ok := fieldValue(line, "data"); ok { r.hasData = true + if r.dropKeepaliveEvents { + if payload := bytes.TrimSpace(value); len(payload) > 0 && !bytes.Equal(payload, []byte("{}")) { + r.hasMeaningfulData = true + } + } + } else if r.dropKeepaliveEvents { + if value, ok := fieldValue(line, "event"); ok && string(bytes.TrimSpace(value)) == "keepalive" { + r.isKeepalive = true + } } } } +// fieldValue returns the value of an SSE line whose field name is `name`, +// with the single optional leading space the SSE grammar allows already +// removed. +func fieldValue(line []byte, name string) ([]byte, bool) { + value, ok := bytes.CutPrefix(line, []byte(name+":")) + if !ok { + return nil, false + } + return bytes.TrimPrefix(value, []byte(" ")), true +} + func (r *sseFilterReader) Close() error { return r.src.Close() } diff --git a/pkg/httpclient/sse_filter_test.go b/pkg/httpclient/sse_filter_test.go index 11eb3f7ff..a5508891d 100644 --- a/pkg/httpclient/sse_filter_test.go +++ b/pkg/httpclient/sse_filter_test.go @@ -153,7 +153,7 @@ func TestSSEFilter_LargeEvent(t *testing.T) { t.Parallel() largeData := "data: " + strings.Repeat("x", 256*1024) + "\n\n" - r := newSSEFilterReader(io.NopCloser(strings.NewReader(largeData))) + r := newSSEFilterReader(io.NopCloser(strings.NewReader(largeData)), false) output, err := io.ReadAll(r) require.NoError(t, err) @@ -167,7 +167,7 @@ func TestSSEFilter_PartialReads(t *testing.T) { t.Parallel() input := "data: test1\n\ndata: test2\n\n" - r := newSSEFilterReader(io.NopCloser(strings.NewReader(input))) + r := newSSEFilterReader(io.NopCloser(strings.NewReader(input)), false) var output []byte buf := make([]byte, 5) @@ -192,7 +192,7 @@ func TestSSEFilter_IncompleteEventAtEOF(t *testing.T) { t.Parallel() input := "data: complete\n\ndata: incomplete" - r := newSSEFilterReader(io.NopCloser(strings.NewReader(input))) + r := newSSEFilterReader(io.NopCloser(strings.NewReader(input)), false) output, err := io.ReadAll(r) require.NoError(t, err) @@ -204,7 +204,7 @@ func TestSSEFilter_IncompleteEventAtEOF(t *testing.T) { func TestSSEFilter_EmptyInput(t *testing.T) { t.Parallel() - r := newSSEFilterReader(io.NopCloser(strings.NewReader(""))) + r := newSSEFilterReader(io.NopCloser(strings.NewReader("")), false) output, err := io.ReadAll(r) require.NoError(t, err) @@ -218,7 +218,7 @@ func TestSSEFilter_OnlyComments(t *testing.T) { t.Parallel() input := ": comment1\n\n: comment2\n\n" - r := newSSEFilterReader(io.NopCloser(strings.NewReader(input))) + r := newSSEFilterReader(io.NopCloser(strings.NewReader(input)), false) output, err := io.ReadAll(r) require.NoError(t, err) @@ -230,7 +230,7 @@ func TestSSEFilter_OnlyComments(t *testing.T) { func TestSSEFilter_ScannerError(t *testing.T) { t.Parallel() - r := newSSEFilterReader(io.NopCloser(&errorReader{err: io.ErrUnexpectedEOF})) + r := newSSEFilterReader(io.NopCloser(&errorReader{err: io.ErrUnexpectedEOF}), false) _, err := io.ReadAll(r) assert.ErrorIs(t, err, io.ErrUnexpectedEOF) @@ -256,7 +256,7 @@ func TestSSEFilter_CloseWithoutRead(t *testing.T) { onClose: func() { closed = true }, } - r := newSSEFilterReader(tracker) + r := newSSEFilterReader(tracker, false) require.NoError(t, r.Close()) assert.True(t, closed, "underlying reader should be closed") } @@ -344,3 +344,154 @@ func fetchThroughFilter(t *testing.T, url string) string { require.NoError(t, err) return string(body) } + +// Gemini-shaped data chunks used by the keepalive tests: a text delta and a +// media (inlineData) delta of the kind an image-output model streams. +const ( + geminiTextChunk = `data: {"candidates":[{"content":{"parts":[{"text":"hi"}],"role":"model"}}]}` + "\n\n" + geminiMediaChunk = `data: {"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"aGVsbG8="}}],"role":"model"},"finishReason":"STOP"}]}` + "\n\n" + keepaliveFrame = "event: keepalive\ndata: {}\n\n" +) + +// TestSSEFilter_KeepaliveMode covers the opt-in keepalive-dropping mode used +// by the Gemini gateway client: payload-free `event: keepalive` frames are +// removed while every other frame — including named events with meaningful +// data — passes through verbatim. +func TestSSEFilter_KeepaliveMode(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + in string + want string + }{ + { + // The gateway scenario: keepalive frames interleaved with + // real text and media chunks during a long image generation. + name: "drops keepalive frames interleaved with data and media chunks", + in: keepaliveFrame + geminiTextChunk + keepaliveFrame + keepaliveFrame + geminiMediaChunk, + want: geminiTextChunk + geminiMediaChunk, + }, + { + name: "drops keepalive without a space after data:", + in: "event: keepalive\ndata:{}\n\n" + geminiTextChunk, + want: geminiTextChunk, + }, + { + name: "drops keepalive with an empty data payload", + in: "event: keepalive\ndata:\n\n" + geminiTextChunk, + want: geminiTextChunk, + }, + { + // Anthropic-style framing: a named event with meaningful data + // must never be touched, even in keepalive mode. + name: "preserves named events with meaningful data", + in: "event: content_block_delta\ndata: {\"delta\":{\"text\":\"hi\"}}\n\n", + want: "event: content_block_delta\ndata: {\"delta\":{\"text\":\"hi\"}}\n\n", + }, + { + // Conservative: only payload-free keepalives are dropped. A + // keepalive-named event carrying real data is preserved. + name: "preserves keepalive-named event with meaningful data", + in: "event: keepalive\ndata: {\"note\":\"x\"}\n\n", + want: "event: keepalive\ndata: {\"note\":\"x\"}\n\n", + }, + { + // The base filter's behavior is unchanged by keepalive mode. + name: "still drops comment-only and no-data event frames", + in: ": ping\n\nevent: ping\nid: abc\n\n" + geminiTextChunk, + want: geminiTextChunk, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + r := newSSEFilterReader(io.NopCloser(strings.NewReader(tt.in)), true) + out, err := io.ReadAll(r) + require.NoError(t, err) + assert.Equal(t, tt.want, string(out)) + }) + } +} + +// TestSSEFilter_KeepaliveMode_OutputParsableByGenaiStyleParser feeds a +// keepalive-interleaved Gemini stream through the keepalive-mode filter and +// verifies the result against the constraint that made the fix necessary: +// genai's iterateResponseStream (google.golang.org/genai api_client.go) +// treats ANY non-blank line without a `data:` prefix as a fatal invalid +// chunk. Every data payload must survive, in order. +func TestSSEFilter_KeepaliveMode_OutputParsableByGenaiStyleParser(t *testing.T) { + t.Parallel() + + in := keepaliveFrame + geminiTextChunk + keepaliveFrame + geminiMediaChunk + keepaliveFrame + r := newSSEFilterReader(io.NopCloser(strings.NewReader(in)), true) + out, err := io.ReadAll(r) + require.NoError(t, err) + + var payloads []string + for line := range strings.Lines(string(out)) { + line = strings.TrimSuffix(line, "\n") + if line == "" { + continue + } + require.True(t, strings.HasPrefix(line, "data:"), "genai would reject this line as an invalid stream chunk: %q", line) + payloads = append(payloads, strings.TrimPrefix(line, "data: ")) + } + + require.Len(t, payloads, 2) + assert.Contains(t, payloads[0], `"text":"hi"`) + assert.Contains(t, payloads[1], `"inlineData"`) +} + +// TestSSEFilter_SharedPathKeepsKeepaliveFrames pins that the shared default +// filter (used by every other provider) does NOT gain keepalive dropping: +// an `event: keepalive` frame has a data line, so it passes through +// verbatim, exactly like Anthropic's meaningful named events. +func TestSSEFilter_SharedPathKeepsKeepaliveFrames(t *testing.T) { + t.Parallel() + + in := keepaliveFrame + + "event: content_block_delta\ndata: {\"delta\":{\"text\":\"hi\"}}\n\n" + + assert.Equal(t, in, fetchSSE(t, in)) +} + +// TestNewHTTPClient_SSEKeepaliveFilterOptIn verifies the option wiring end +// to end through NewHTTPClient: keepalive frames are dropped only when +// WithSSEKeepaliveFilter is passed, and the default client leaves them in. +func TestNewHTTPClient_SSEKeepaliveFilterOptIn(t *testing.T) { + t.Parallel() + + in := keepaliveFrame + geminiTextChunk + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + _, _ = io.WriteString(w, in) + })) + t.Cleanup(srv.Close) + + fetch := func(t *testing.T, client *http.Client) string { + t.Helper() + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, srv.URL, http.NoBody) + require.NoError(t, err) + res, err := client.Do(req) + require.NoError(t, err) + defer func() { _ = res.Body.Close() }() + body, err := io.ReadAll(res.Body) + require.NoError(t, err) + return string(body) + } + + t.Run("opted in drops keepalive frames", func(t *testing.T) { + t.Parallel() + client := NewHTTPClient(t.Context(), WithSSEKeepaliveFilter()) + assert.Equal(t, geminiTextChunk, fetch(t, client)) + }) + + t.Run("default keeps keepalive frames", func(t *testing.T) { + t.Parallel() + client := NewHTTPClient(t.Context()) + assert.Equal(t, in, fetch(t, client)) + }) +} diff --git a/pkg/model/provider/gemini/client.go b/pkg/model/provider/gemini/client.go index 29d50987a..4acefa9a3 100644 --- a/pkg/model/provider/gemini/client.go +++ b/pkg/model/provider/gemini/client.go @@ -179,6 +179,12 @@ func NewClient(ctx context.Context, cfg *latest.ModelConfig, env environment.Pro } } + // The gateway keeps long generations alive with `event: keepalive` + // + `data: {}` frames, which genai's SSE parser rejects as fatal + // invalid chunks. Drop them here, on the gateway path only — direct + // Gemini/Vertex clients never receive them. + httpOptions = append(httpOptions, httpclient.WithSSEKeepaliveFilter()) + gatewayHTTPClient := httpclient.NewHTTPClient(ctx, httpOptions...) globalOptions.WrapTransport(ctx, gatewayHTTPClient) diff --git a/pkg/model/provider/gemini/gateway_sse_keepalive_test.go b/pkg/model/provider/gemini/gateway_sse_keepalive_test.go new file mode 100644 index 000000000..67548872f --- /dev/null +++ b/pkg/model/provider/gemini/gateway_sse_keepalive_test.go @@ -0,0 +1,110 @@ +package gemini + +import ( + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/config/latest" + "github.com/docker/docker-agent/pkg/environment" + "github.com/docker/docker-agent/pkg/model/provider/options" +) + +// writeGeminiSSEResponseWithKeepalives replays what the Docker AI Gateway +// sends during a long image generation: `event: keepalive` + `data: {}` +// frames interleaved with real text and inlineData chunks. +func writeGeminiSSEResponseWithKeepalives(w http.ResponseWriter) { + w.Header().Set("Content-Type", "text/event-stream") + _, _ = io.WriteString(w, + "event: keepalive\ndata: {}\n\n"+ + `data: {"candidates":[{"content":{"parts":[{"text":"here it comes"}],"role":"model"}}]}`+"\n\n"+ + "event: keepalive\ndata: {}\n\n"+ + "event: keepalive\ndata: {}\n\n"+ + `data: {"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"aGVsbG8="}}],"role":"model"},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":1,"candidatesTokenCount":1,"totalTokenCount":2}}`+"\n\n") +} + +// collectStream drains a stream, concatenating text deltas and collecting +// text deltas, and returns the first non-EOF error (nil on clean EOF). +func collectStream(t *testing.T, stream chat.MessageStream) (string, error) { + t.Helper() + defer stream.Close() + + var text strings.Builder + for { + resp, err := stream.Recv() + if errors.Is(err, io.EOF) { + return text.String(), nil + } + if err != nil { + return text.String(), err + } + for _, choice := range resp.Choices { + text.WriteString(choice.Delta.Content) + } + } +} + +// TestCreateChatCompletionStream_GatewaySurvivesKeepaliveFrames pins the +// keepalive fix end to end: a gateway stream interleaved with keepalive +// frames completes without error and delivers every text delta. +// Without the gateway-scoped filter, genai's SSE parser fails the whole +// stream with "invalid stream chunk: event: keepalive". +func TestCreateChatCompletionStream_GatewaySurvivesKeepaliveFrames(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + writeGeminiSSEResponseWithKeepalives(w) + })) + t.Cleanup(server.Close) + + cfg := &latest.ModelConfig{Provider: "google", Model: "gemini-3-pro-image-preview"} + env := environment.NewMapEnvProvider(map[string]string{ + environment.DockerDesktopTokenEnv: "test-dd-token", + }) + client, err := NewClient(t.Context(), cfg, env, options.WithGateway(server.URL)) + require.NoError(t, err) + + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{ + {Role: chat.MessageRoleUser, Content: "generate an image of a red panda"}, + }, nil) + require.NoError(t, err) + + text, err := collectStream(t, stream) + require.NoError(t, err, "keepalive frames must never reach the genai SSE parser on the gateway path") + assert.Equal(t, "here it comes", text) +} + +// TestCreateChatCompletionStream_DirectPathKeepaliveUnfiltered pins the +// scoping: the direct (non-gateway) Gemini API path does NOT get keepalive +// filtering, so the same frames still surface genai's invalid-chunk error. +// Direct Gemini never emits these frames — this test only guards against +// the filter accidentally widening beyond the gateway client. +func TestCreateChatCompletionStream_DirectPathKeepaliveUnfiltered(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + writeGeminiSSEResponseWithKeepalives(w) + })) + t.Cleanup(server.Close) + + cfg := &latest.ModelConfig{Provider: "google", Model: "gemini-3-pro-image-preview", BaseURL: server.URL} + env := environment.NewMapEnvProvider(map[string]string{"GOOGLE_API_KEY": "test-key"}) + client, err := NewClient(t.Context(), cfg, env) + require.NoError(t, err) + + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{ + {Role: chat.MessageRoleUser, Content: "hello"}, + }, nil) + require.NoError(t, err) + + _, err = collectStream(t, stream) + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid stream chunk") +} diff --git a/pkg/model/provider/gemini/image_response_modalities_test.go b/pkg/model/provider/gemini/image_response_modalities_test.go new file mode 100644 index 000000000..4380d2c7f --- /dev/null +++ b/pkg/model/provider/gemini/image_response_modalities_test.go @@ -0,0 +1,405 @@ +package gemini + +import ( + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/config/latest" + "github.com/docker/docker-agent/pkg/environment" + "github.com/docker/docker-agent/pkg/model/provider/base" + "github.com/docker/docker-agent/pkg/model/provider/options" + "github.com/docker/docker-agent/pkg/rag/types" + "github.com/docker/docker-agent/pkg/tools" +) + +// TestWantsImageResponseModalities exhaustively covers the predicate gating +// TEXT+IMAGE response modalities: it must be true only for supported Gemini +// surfaces, only for a model explicitly declared image-output-capable via +// output_capabilities.image, and only outside title-generation/compaction +// utility calls. +func TestWantsImageResponseModalities(t *testing.T) { + t.Parallel() + + declaredTrue := &latest.OutputCapabilitiesConfig{Image: latest.Bool(true)} + declaredFalse := &latest.OutputCapabilitiesConfig{Image: latest.Bool(false)} + + tests := []struct { + name string + apiSurface string + declared *latest.OutputCapabilitiesConfig + opts []options.Opt + want bool + }{ + {name: "gateway, declared true, ordinary chat: wants modalities", apiSurface: apiSurfaceGateway, declared: declaredTrue, want: true}, + {name: "direct Gemini API, declared true, ordinary chat: wants modalities", apiSurface: apiSurfaceGeminiAPI, declared: declaredTrue, want: true}, + {name: "Vertex AI, declared true, ordinary chat: wants modalities", apiSurface: apiSurfaceVertexAI, declared: declaredTrue, want: true}, + {name: "gateway, declared false: never", apiSurface: apiSurfaceGateway, declared: declaredFalse, want: false}, + {name: "direct Gemini API, declaration missing: never", apiSurface: apiSurfaceGeminiAPI, declared: nil, want: false}, + { + name: "gateway, declared true, generating title: never", apiSurface: apiSurfaceGateway, declared: declaredTrue, + opts: []options.Opt{options.WithGeneratingTitle()}, want: false, + }, + { + name: "gateway, declared true, compacting: never", apiSurface: apiSurfaceGateway, declared: declaredTrue, + opts: []options.Opt{options.WithCompacting()}, want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + client := &Client{ + Config: base.Config{ + ModelConfig: latest.ModelConfig{ + Provider: "google", + Model: "gemini-2.5-flash-image", + OutputCapabilities: tt.declared, + }, + ModelOptions: options.Apply(tt.opts...), + }, + apiSurface: tt.apiSurface, + } + + assert.Equal(t, tt.want, client.wantsImageResponseModalities(tt.declared != nil && tt.declared.Image != nil && *tt.declared.Image)) + }) + } +} + +// capturedRequests is a mutex-guarded log of raw request bodies received by +// a [newBodyCapturingGeminiServer], letting tests assert exactly what was +// (or, cheaply, was not — an empty log) serialized onto the wire. +type capturedRequests struct { + mu sync.Mutex + bodies [][]byte +} + +func (c *capturedRequests) add(b []byte) { + c.mu.Lock() + defer c.mu.Unlock() + c.bodies = append(c.bodies, b) +} + +func (c *capturedRequests) all() [][]byte { + c.mu.Lock() + defer c.mu.Unlock() + return append([][]byte(nil), c.bodies...) +} + +// newBodyCapturingGeminiServer starts an httptest server that records the +// raw request body of every call it receives before responding via +// respond, so tests can decode the exact generationConfig sent on the wire. +func newBodyCapturingGeminiServer(t *testing.T, respond func(w http.ResponseWriter)) (*httptest.Server, *capturedRequests) { + t.Helper() + captured := &capturedRequests{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + captured.add(body) + respond(w) + })) + t.Cleanup(server.Close) + return server, captured +} + +// writeGeminiGenerateContentJSONResponse writes a minimal, non-streaming +// Gemini generateContent JSON response whose sole text part is text (used +// for Rerank, which does not use the SSE streaming endpoint). +func writeGeminiGenerateContentJSONResponse(w http.ResponseWriter, text string) { + w.Header().Set("Content-Type", "application/json") + payload, _ := json.Marshal(map[string]any{ + "candidates": []map[string]any{{ + "content": map[string]any{ + "role": "model", + "parts": []map[string]any{{"text": text}}, + }, + }}, + }) + _, _ = w.Write(payload) +} + +// responseModalitiesInBody decodes body's generationConfig.responseModalities +// (see genai's generateContentConfigToMldev, which nests the serialized +// GenerateContentConfig under a top-level "generationConfig" key for the +// Gemini Developer API), returning nil when either key is absent. +func responseModalitiesInBody(t *testing.T, body []byte) []string { + t.Helper() + + var req map[string]any + require.NoError(t, json.Unmarshal(body, &req)) + + genCfg, ok := req["generationConfig"].(map[string]any) + if !ok { + return nil + } + raw, ok := genCfg["responseModalities"].([]any) + if !ok { + return nil + } + out := make([]string, len(raw)) + for i, v := range raw { + out[i], _ = v.(string) + } + return out +} + +func drainStream(t *testing.T, stream chat.MessageStream) { + t.Helper() + defer stream.Close() + for { + if _, err := stream.Recv(); err != nil { + break + } + } +} + +// TestCreateChatCompletionStream_ImageResponseModalities_PositiveRoutes pins +// that supported Gemini surfaces request TEXT+IMAGE output in that order. +func TestCreateChatCompletionStream_ImageResponseModalities_PositiveRoutes(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + cfg func(serverURL string) *latest.ModelConfig + env map[string]string + gateway bool + }{ + { + name: "gateway", + cfg: func(string) *latest.ModelConfig { + return &latest.ModelConfig{Provider: "google", Model: "gemini-2.5-flash-image", OutputCapabilities: &latest.OutputCapabilitiesConfig{Image: latest.Bool(true)}} + }, + env: map[string]string{environment.DockerDesktopTokenEnv: "test-dd-token"}, + gateway: true, + }, + { + name: "direct Gemini API", + cfg: func(serverURL string) *latest.ModelConfig { + return &latest.ModelConfig{Provider: "google", Model: "gemini-2.5-flash-image", BaseURL: serverURL, OutputCapabilities: &latest.OutputCapabilitiesConfig{Image: latest.Bool(true)}} + }, + env: map[string]string{"GOOGLE_API_KEY": "test-key"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + server, captured := newBodyCapturingGeminiServer(t, writeGeminiSSEResponse) + var opts []options.Opt + if tt.gateway { + opts = append(opts, options.WithGateway(server.URL)) + } + client, err := NewClient(t.Context(), tt.cfg(server.URL), environment.NewMapEnvProvider(tt.env), opts...) + require.NoError(t, err) + + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{{Role: chat.MessageRoleUser, Content: "generate an image of a red panda"}}, nil) + require.NoError(t, err) + drainStream(t, stream) + + bodies := captured.all() + require.Len(t, bodies, 1) + assert.Equal(t, []string{"TEXT", "IMAGE"}, responseModalitiesInBody(t, bodies[0])) + }) + } +} + +// TestCreateChatCompletionStream_ImageResponseModalities_AbsentOnOtherRoutes +// pins that undeclared image output and utility calls send no response +// modalities. +func TestCreateChatCompletionStream_ImageResponseModalities_AbsentOnOtherRoutes(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + cfg func(serverURL string) *latest.ModelConfig + env map[string]string + gateway bool + opts []options.Opt + }{ + { + name: "direct Gemini API, declaration missing: absent", + cfg: func(serverURL string) *latest.ModelConfig { + return &latest.ModelConfig{Provider: "google", Model: "gemini-2.5-flash-image", BaseURL: serverURL} + }, + env: map[string]string{"GOOGLE_API_KEY": "test-key"}, + }, + { + name: "gateway, declared false: absent", + cfg: func(string) *latest.ModelConfig { + return &latest.ModelConfig{ + Provider: "google", Model: "gemini-2.5-flash-image", + OutputCapabilities: &latest.OutputCapabilitiesConfig{Image: latest.Bool(false)}, + } + }, + env: map[string]string{environment.DockerDesktopTokenEnv: "test-dd-token"}, + gateway: true, + }, + { + name: "gateway, declaration missing: absent", + cfg: func(string) *latest.ModelConfig { + return &latest.ModelConfig{Provider: "google", Model: "gemini-2.5-flash-image"} + }, + env: map[string]string{environment.DockerDesktopTokenEnv: "test-dd-token"}, + gateway: true, + }, + { + name: "gateway, declared true, generating title: absent", + cfg: func(string) *latest.ModelConfig { + return &latest.ModelConfig{ + Provider: "google", Model: "gemini-2.5-flash-image", + OutputCapabilities: &latest.OutputCapabilitiesConfig{Image: latest.Bool(true)}, + } + }, + env: map[string]string{environment.DockerDesktopTokenEnv: "test-dd-token"}, + gateway: true, + opts: []options.Opt{options.WithGeneratingTitle()}, + }, + { + name: "gateway, declared true, compacting: absent", + cfg: func(string) *latest.ModelConfig { + return &latest.ModelConfig{ + Provider: "google", Model: "gemini-2.5-flash-image", + OutputCapabilities: &latest.OutputCapabilitiesConfig{Image: latest.Bool(true)}, + } + }, + env: map[string]string{environment.DockerDesktopTokenEnv: "test-dd-token"}, + gateway: true, + opts: []options.Opt{options.WithCompacting()}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + server, captured := newBodyCapturingGeminiServer(t, writeGeminiSSEResponse) + + cfg := tt.cfg(server.URL) + env := environment.NewMapEnvProvider(tt.env) + opts := tt.opts + if tt.gateway { + opts = append(opts, options.WithGateway(server.URL)) + } + client, err := NewClient(t.Context(), cfg, env, opts...) + require.NoError(t, err) + + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{ + {Role: chat.MessageRoleUser, Content: "hello"}, + }, nil) + require.NoError(t, err) + drainStream(t, stream) + + bodies := captured.all() + require.Len(t, bodies, 1) + assert.Nil(t, responseModalitiesInBody(t, bodies[0]), "response modalities must be absent on this route") + }) + } +} + +// TestRerank_NeverSetsResponseModalities pins that Rerank — which shares +// buildConfig with CreateChatCompletionStream but must retain its own +// structured-output-only config — never gains response modalities, even on +// a model explicitly declared image-output-capable. +func TestRerank_NeverSetsResponseModalities(t *testing.T) { + t.Parallel() + + server, captured := newBodyCapturingGeminiServer(t, func(w http.ResponseWriter) { + writeGeminiGenerateContentJSONResponse(w, `{"scores":[1]}`) + }) + + cfg := &latest.ModelConfig{ + Provider: "google", + Model: "gemini-2.5-flash-image", + OutputCapabilities: &latest.OutputCapabilitiesConfig{Image: latest.Bool(true)}, + } + env := environment.NewMapEnvProvider(map[string]string{ + environment.DockerDesktopTokenEnv: "test-dd-token", + }) + client, err := NewClient(t.Context(), cfg, env, options.WithGateway(server.URL)) + require.NoError(t, err) + + scores, err := client.Rerank(t.Context(), "query", []types.Document{{Content: "doc1"}}, "") + require.NoError(t, err) + require.Len(t, scores, 1) + + bodies := captured.all() + require.Len(t, bodies, 1) + assert.Nil(t, responseModalitiesInBody(t, bodies[0]), "Rerank must never request response modalities") +} + +// TestCreateChatCompletionStream_ImageResponseModalities_GuardRejectedRoutesNeverDispatch +// preserves guard precedence: on every declared-image route, a request shape +// the guard rejects (custom function tools, a built-in tool, or structured +// output) must still make zero provider calls, so nothing — including response +// modalities — is ever serialized onto the wire. +func TestCreateChatCompletionStream_ImageResponseModalities_GuardRejectedRoutesNeverDispatch(t *testing.T) { + t.Parallel() + + newRejectedClient := func(t *testing.T, serverURL string, extraOpts ...options.Opt) *Client { + t.Helper() + cfg := &latest.ModelConfig{ + Provider: "google", + Model: "gemini-2.5-flash-image", + OutputCapabilities: &latest.OutputCapabilitiesConfig{Image: latest.Bool(true)}, + } + env := environment.NewMapEnvProvider(map[string]string{ + environment.DockerDesktopTokenEnv: "test-dd-token", + }) + opts := append([]options.Opt{options.WithGateway(serverURL)}, extraOpts...) + client, err := NewClient(t.Context(), cfg, env, opts...) + require.NoError(t, err) + return client + } + + assertRejectedWithNoDispatch := func(t *testing.T, err error, stream chat.MessageStream, captured *capturedRequests) { + t.Helper() + require.Nil(t, stream) + var incompatible *ImageOutputRequestIncompatibleError + require.ErrorAs(t, err, &incompatible) + assert.Empty(t, captured.all(), "guard rejection must dispatch nothing, so no modalities are ever serialized") + } + + t.Run("custom function tools rejected", func(t *testing.T) { + t.Parallel() + server, captured := newBodyCapturingGeminiServer(t, writeGeminiSSEResponse) + client := newRejectedClient(t, server.URL) + + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{ + {Role: chat.MessageRoleUser, Content: "hello"}, + }, []tools.Tool{{Name: "read_file", Description: "reads a file", Parameters: map[string]any{"type": "object"}}}) + + assertRejectedWithNoDispatch(t, err, stream, captured) + }) + + t.Run("built-in tool rejected", func(t *testing.T) { + t.Parallel() + server, captured := newBodyCapturingGeminiServer(t, writeGeminiSSEResponse) + client := newRejectedClient(t, server.URL) + client.ModelConfig.ProviderOpts = map[string]any{"google_search": true} + + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{ + {Role: chat.MessageRoleUser, Content: "hello"}, + }, nil) + + assertRejectedWithNoDispatch(t, err, stream, captured) + }) + + t.Run("structured output rejected", func(t *testing.T) { + t.Parallel() + server, captured := newBodyCapturingGeminiServer(t, writeGeminiSSEResponse) + client := newRejectedClient(t, server.URL, options.WithStructuredOutput(&latest.StructuredOutput{Schema: map[string]any{"type": "object"}})) + + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{ + {Role: chat.MessageRoleUser, Content: "hello"}, + }, nil) + + assertRejectedWithNoDispatch(t, err, stream, captured) + }) +}