Bound and strictly decode provider responses

This commit is contained in:
2026-08-11 23:46:16 +00:00
parent 3a43550f70
commit 2b6a7f83c4
4 changed files with 348 additions and 12 deletions

View File

@@ -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 == "" {