From 2b6a7f83c4b85590c9994c3798c45cfd8d359bfb Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Tue, 11 Aug 2026 23:46:16 +0000 Subject: [PATCH] Bound and strictly decode provider responses --- docs/integrations/openai-compatible-chat.md | 23 +- docs/internal/llm.md | 16 +- internal/llm/openai_compatible_client.go | 51 +++- internal/llm/openai_compatible_client_test.go | 270 +++++++++++++++++- 4 files changed, 348 insertions(+), 12 deletions(-) diff --git a/docs/integrations/openai-compatible-chat.md b/docs/integrations/openai-compatible-chat.md index aa4ca75..4da1e64 100644 --- a/docs/integrations/openai-compatible-chat.md +++ b/docs/integrations/openai-compatible-chat.md @@ -88,13 +88,24 @@ request fields. ## Response Handling -Any 2xx response is decoded as an OpenAI-compatible chat response. The client -returns the first choice's non-empty message content and maps prompt, -completion, total, cached, and cache-write token counts. +Any 2xx response body is limited to 16 MiB (16,777,216 bytes). A larger +declared `Content-Length` is rejected before the body is read, and streamed, +chunked, or underreported bodies are read through the same bound with at most +one additional byte used to detect overflow. A body exactly at the limit is +allowed. The body is closed on every outcome and an oversized stream is not +drained. -Invalid JSON, absent choices, and empty first-choice content are malformed -responses. For a non-2xx status, the error includes the status code but never -the provider response body. +The bounded body must contain exactly one OpenAI-compatible JSON response +object followed only by JSON whitespace and EOF. The client returns the first +choice's non-empty message content and maps prompt, completion, total, cached, +and cache-write token counts. Invalid or truncated JSON, trailing non-whitespace +data, a second JSON value, absent choices, empty first-choice content, and size +overflow are malformed responses and return no partial result. + +For a non-2xx status, the error includes the status code but never the provider +response body. Promptkit does not yet parse provider error envelopes; bounded +non-success parsing belongs to the +[structured-generation-error roadmap](../roadmap/structured-generation-errors.md). An outbound `http.Client.Do` failure retains both Promptkit's request-failure identity and the exact transport error for `errors.Is` and `errors.As` checks. diff --git a/docs/internal/llm.md b/docs/internal/llm.md index 301f306..d6ca67e 100644 --- a/docs/internal/llm.md +++ b/docs/internal/llm.md @@ -37,7 +37,8 @@ resolved request target may supply the endpoint. Generation then: 4. composes `/chat/completions` through parsed URL path operations; 5. resolves authentication; 6. performs the outbound request under the applicable deadlines; and -7. decodes the first response choice and token usage. +7. decodes one strictly framed, size-bounded response object and maps its first + choice and token usage. `internal/llm` owns the set of reserved OpenAI-compatible request fields used when validating extra parameters. Backend registration consumes the same rule @@ -72,6 +73,16 @@ Invalid nonempty configured endpoints are configuration failures. A missing or invalid final selected endpoint is an invalid generation request and is rejected before transport. +Successful response bodies have a fixed 16 MiB limit enforced by declared +length and by reading at most one byte beyond the boundary. The decoder accepts +exactly one JSON object plus trailing whitespace and EOF. Size overflow, +truncation, malformed JSON, trailing data, and a second value are malformed +responses with no partial result or provider content in the error. Every body +is closed, and an unbounded oversized stream is not drained. Non-success +responses remain status-only; bounded provider error-envelope parsing belongs +to the +[structured-generation-error roadmap](../roadmap/structured-generation-errors.md). + An `http.Client.Do` failure is represented by a redacting multi-cause error: the package request-failure sentinel and the exact returned transport error are both available through `errors.Is` and `errors.As`, while the rendered text @@ -88,7 +99,8 @@ The own configuration, client cloning, deterministic deadline precedence, authentication, request and response mapping, malformed data, error identity, cancellation, endpoint selection and composition, pre-transport rejection, and -response-body suppression. The root +bounded single-document response framing, closure, and response-body +suppression. The root transport contract tests also verify that resolved backend settings reach this client without serializing backend identity and that ordinary-run cancellation retains its public generation and context identities. All use local test diff --git a/internal/llm/openai_compatible_client.go b/internal/llm/openai_compatible_client.go index 48ab417..3071487 100644 --- a/internal/llm/openai_compatible_client.go +++ b/internal/llm/openai_compatible_client.go @@ -25,6 +25,8 @@ var ( ErrMalformedResponse = errors.New("malformed llm response") ) +const maxOpenAIChatResponseBytes int64 = 16 << 20 + type requestFailedError struct { cause error } @@ -156,10 +158,13 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera _, _ = io.Copy(io.Discard, io.LimitReader(httpResp.Body, 4096)) return nil, fmt.Errorf("%w: status=%d", ErrUnexpectedStatus, httpResp.StatusCode) } + if httpResp.ContentLength > maxOpenAIChatResponseBytes { + return nil, openAIChatResponseTooLargeError() + } - var wireResp openAIChatResponse - if err := json.NewDecoder(httpResp.Body).Decode(&wireResp); err != nil { - return nil, fmt.Errorf("%w: failed to decode response: %v", ErrMalformedResponse, err) + wireResp, err := decodeOpenAIChatResponse(httpResp.Body) + if err != nil { + return nil, err } if len(wireResp.Choices) == 0 { @@ -182,6 +187,46 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera }, nil } +func decodeOpenAIChatResponse(body io.Reader) (openAIChatResponse, error) { + limited := &io.LimitedReader{ + R: body, + N: maxOpenAIChatResponseBytes + 1, + } + decoder := json.NewDecoder(limited) + + var response openAIChatResponse + if err := decoder.Decode(&response); err != nil { + if limited.N == 0 { + return openAIChatResponse{}, openAIChatResponseTooLargeError() + } + return openAIChatResponse{}, fmt.Errorf("%w: failed to decode response", ErrMalformedResponse) + } + if limited.N == 0 { + return openAIChatResponse{}, openAIChatResponseTooLargeError() + } + + var trailing any + if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) { + if limited.N == 0 { + return openAIChatResponse{}, openAIChatResponseTooLargeError() + } + return openAIChatResponse{}, fmt.Errorf("%w: response contains trailing data", ErrMalformedResponse) + } + if limited.N == 0 { + return openAIChatResponse{}, openAIChatResponseTooLargeError() + } + + return response, nil +} + +func openAIChatResponseTooLargeError() error { + return fmt.Errorf( + "%w: response exceeds %d-byte limit", + ErrMalformedResponse, + maxOpenAIChatResponseBytes, + ) +} + func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultModel string) (openAIChatRequest, error) { model := strings.TrimSpace(req.Target.Model) if model == "" { diff --git a/internal/llm/openai_compatible_client_test.go b/internal/llm/openai_compatible_client_test.go index 25f9378..d70ca37 100644 --- a/internal/llm/openai_compatible_client_test.go +++ b/internal/llm/openai_compatible_client_test.go @@ -17,7 +17,10 @@ import ( "gitea.maximumdirect.net/eric/promptkit/internal/domain" ) -var errTransportStopped = errors.New("transport stopped after request inspection") +var ( + errTransportStopped = errors.New("transport stopped after request inspection") + errResponseReadPastLimit = errors.New("response reader was read past the allowed boundary") +) type deadlineCapturingTransport struct { deadline time.Time @@ -52,6 +55,72 @@ func (waitingContextTransport) RoundTrip(req *http.Request) (*http.Response, err return nil, req.Context().Err() } +type countingReadCloser struct { + reader io.Reader + bytesRead int64 + closed bool +} + +func (r *countingReadCloser) Read(p []byte) (int, error) { + n, err := r.reader.Read(p) + r.bytesRead += int64(n) + return n, err +} + +func (r *countingReadCloser) Close() error { + r.closed = true + return nil +} + +type repeatingByteReader byte + +func (r repeatingByteReader) Read(p []byte) (int, error) { + for i := range p { + p[i] = byte(r) + } + return len(p), nil +} + +type guardedRepeatingReader struct { + value byte + remaining int64 + readPastLimit bool +} + +func (r *guardedRepeatingReader) Read(p []byte) (int, error) { + if r.remaining == 0 { + r.readPastLimit = true + return 0, errResponseReadPastLimit + } + if int64(len(p)) > r.remaining { + p = p[:r.remaining] + } + for i := range p { + p[i] = r.value + } + r.remaining -= int64(len(p)) + return len(p), nil +} + +const ( + successResponsePrefix = `{"choices":[{"message":{"content":"` + successResponseSuffix = `"}}]}` + responseContentMarker = "provider-secret-fragment" +) + +func sizedSuccessResponseBody(size int64) *countingReadCloser { + contentBytes := size - int64(len(successResponsePrefix)+len(successResponseSuffix)) + if contentBytes < int64(len(responseContentMarker)) { + panic("successful response size is too small") + } + return &countingReadCloser{reader: io.MultiReader( + strings.NewReader(successResponsePrefix), + strings.NewReader(responseContentMarker), + io.LimitReader(repeatingByteReader('x'), contentBytes-int64(len(responseContentMarker))), + strings.NewReader(successResponseSuffix), + )} +} + func assertDeadlineNear(t *testing.T, deadline, before, after time.Time, duration time.Duration) { t.Helper() @@ -1055,6 +1124,205 @@ func TestOpenAICompatibleClientMalformedResponseMissingChoices(t *testing.T) { } } +func TestOpenAICompatibleClientBoundsSuccessfulResponseBodies(t *testing.T) { + tests := []struct { + name string + size int64 + contentLength int64 + wantErr bool + }{ + {name: "just below limit with content length", size: maxOpenAIChatResponseBytes - 1, contentLength: maxOpenAIChatResponseBytes - 1}, + {name: "exact limit with content length", size: maxOpenAIChatResponseBytes, contentLength: maxOpenAIChatResponseBytes}, + {name: "one byte over with content length", size: maxOpenAIChatResponseBytes + 1, contentLength: maxOpenAIChatResponseBytes + 1, wantErr: true}, + {name: "just below limit without content length", size: maxOpenAIChatResponseBytes - 1, contentLength: -1}, + {name: "exact limit without content length", size: maxOpenAIChatResponseBytes, contentLength: -1}, + {name: "one byte over without content length", size: maxOpenAIChatResponseBytes + 1, contentLength: -1, wantErr: true}, + {name: "one byte over with underreported content length", size: maxOpenAIChatResponseBytes + 1, contentLength: maxOpenAIChatResponseBytes - 1, wantErr: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + body := sizedSuccessResponseBody(tc.size) + client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ + BaseURL: "https://provider.example/v1", + Model: "m", + HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: body, + ContentLength: tc.contentLength, + Request: req, + }, nil + })}, + }) + if err != nil { + t.Fatalf("construct client: %v", err) + } + + response, err := client.Generate(context.Background(), domain.GenerateRequest{ + Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, + }) + if tc.wantErr { + if response != nil { + t.Fatalf("expected no partial response, got %+v", response) + } + if !errors.Is(err, ErrMalformedResponse) { + t.Fatalf("expected ErrMalformedResponse, got %v", err) + } + if strings.Contains(err.Error(), responseContentMarker) { + t.Fatalf("error exposed provider content: %v", err) + } + } else { + if err != nil { + t.Fatalf("generate: %v", err) + } + wantContentBytes := tc.size - int64(len(successResponsePrefix)+len(successResponseSuffix)) + if int64(len(response.Content)) != wantContentBytes || !strings.HasPrefix(response.Content, responseContentMarker) { + t.Fatalf("unexpected response content length or prefix") + } + } + if body.bytesRead > maxOpenAIChatResponseBytes+1 { + t.Fatalf("read %d bytes, limit is %d", body.bytesRead, maxOpenAIChatResponseBytes+1) + } + if tc.wantErr && tc.contentLength > maxOpenAIChatResponseBytes && body.bytesRead != 0 { + t.Fatalf("read %d bytes despite oversized Content-Length", body.bytesRead) + } + if tc.wantErr && tc.contentLength <= maxOpenAIChatResponseBytes && body.bytesRead != maxOpenAIChatResponseBytes+1 { + t.Fatalf("read %d bytes, want one byte beyond the limit", body.bytesRead) + } + if !tc.wantErr && body.bytesRead != tc.size { + t.Fatalf("read %d bytes, want complete %d-byte response", body.bytesRead, tc.size) + } + if !body.closed { + t.Fatal("response body was not closed") + } + }) + } +} + +func TestOpenAICompatibleClientRejectsContinuingOversizedResponse(t *testing.T) { + prefix := successResponsePrefix + responseContentMarker + continuation := &guardedRepeatingReader{ + value: 'x', + remaining: maxOpenAIChatResponseBytes + 1 - int64(len(prefix)), + } + body := &countingReadCloser{reader: io.MultiReader( + strings.NewReader(prefix), + continuation, + )} + client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ + BaseURL: "https://provider.example/v1", + Model: "m", + HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: body, + ContentLength: -1, + Request: req, + }, nil + })}, + }) + if err != nil { + t.Fatalf("construct client: %v", err) + } + + type outcome struct { + response *domain.GenerateResponse + err error + } + done := make(chan outcome, 1) + go func() { + response, generateErr := client.Generate(context.Background(), domain.GenerateRequest{ + Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, + }) + done <- outcome{response: response, err: generateErr} + }() + + select { + case result := <-done: + if result.response != nil { + t.Fatalf("expected no partial response, got %+v", result.response) + } + if !errors.Is(result.err, ErrMalformedResponse) { + t.Fatalf("expected ErrMalformedResponse, got %v", result.err) + } + if strings.Contains(result.err.Error(), responseContentMarker) { + t.Fatalf("error exposed provider content: %v", result.err) + } + case <-time.After(5 * time.Second): + t.Fatal("timed out rejecting continuing oversized response") + } + if body.bytesRead != maxOpenAIChatResponseBytes+1 { + t.Fatalf("read %d bytes, want %d", body.bytesRead, maxOpenAIChatResponseBytes+1) + } + if continuation.readPastLimit { + t.Fatal("response reader was read past one byte beyond the limit") + } + if !body.closed { + t.Fatal("response body was not closed") + } +} + +func TestOpenAICompatibleClientRequiresSingleResponseDocument(t *testing.T) { + validResponse := `{"choices":[{"message":{"content":"ok"}}]}` + tests := []struct { + name string + body string + wantErr bool + }{ + {name: "one document", body: validResponse}, + {name: "trailing whitespace", body: validResponse + " \n\t\r "}, + {name: "trailing garbage", body: validResponse + " " + responseContentMarker, wantErr: true}, + {name: "second JSON value", body: validResponse + ` {"detail":"` + responseContentMarker + `"}`, wantErr: true}, + {name: "truncated document", body: validResponse[:len(validResponse)-2], wantErr: true}, + {name: "missing choices", body: `{"choices":[]}`, wantErr: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + body := &countingReadCloser{reader: strings.NewReader(tc.body)} + client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ + BaseURL: "https://provider.example/v1", + Model: "m", + HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: body, + ContentLength: int64(len(tc.body)), + Request: req, + }, nil + })}, + }) + if err != nil { + t.Fatalf("construct client: %v", err) + } + + response, err := client.Generate(context.Background(), domain.GenerateRequest{ + Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, + }) + if tc.wantErr { + if response != nil { + t.Fatalf("expected no partial response, got %+v", response) + } + if !errors.Is(err, ErrMalformedResponse) { + t.Fatalf("expected ErrMalformedResponse, got %v", err) + } + if strings.Contains(err.Error(), responseContentMarker) { + t.Fatalf("error exposed provider content: %v", err) + } + } else if err != nil || response == nil || response.Content != "ok" { + t.Fatalf("response = %+v, error = %v", response, err) + } + if !body.closed { + t.Fatal("response body was not closed") + } + }) + } +} + func TestOpenAICompatibleClientGenerationTimeoutSetsEarlierDeadline(t *testing.T) { transport := &deadlineCapturingTransport{err: context.DeadlineExceeded} generationTimeout := 2 * time.Second