From 62a9424ac08b5f547f0d4a7857312fd6ba13e786 Mon Sep 17 00:00:00 2001 From: Hung Nguyen Date: Sat, 8 Aug 2026 10:53:06 +0700 Subject: [PATCH 1/2] feat(agent): classify truncated and content-blocked LLM responses ErrLLMTruncated / ErrLLMContentBlocked, carried by IncompleteResponseError, so a cut-off response stops surfacing as a decode error against the caller's schema. Blocked stops are non-retryable; the retry content gate now also covers the first attempt, so a truncation never replays its own stream. --- pkg/agent/errors.go | 83 ++++++++++++++++++++++++++++++++++++ pkg/agent/llm_call.go | 9 +++- pkg/agent/retry.go | 8 +++- pkg/agent/retry_hook_test.go | 65 ++++++++++++++++++++++++++++ 4 files changed, 161 insertions(+), 4 deletions(-) diff --git a/pkg/agent/errors.go b/pkg/agent/errors.go index 54358b0..7188e2d 100644 --- a/pkg/agent/errors.go +++ b/pkg/agent/errors.go @@ -66,6 +66,33 @@ var ( // } ErrLLMAuth = errors.New("agent: LLM provider authentication or configuration failed") + // ErrLLMTruncated is returned when the provider stopped generating + // before the response was finished — an output-token cap, typically. + // + // The danger is that a truncated response is a *valid prefix*: the text + // looks fine and only fails when something downstream parses it, so the + // operator reads "unexpected end of JSON input" and inspects a schema + // that was never wrong. Providers detect the cut at the source, so route + // on the sentinel instead of matching decode-error text. + // + // The cut is a function of this response's length rather than of the + // request, so isRetryable leaves it retryable — but a truncation that + // already streamed content is not retried, because re-running it would + // replay the same text into the consumer's stream. The real fix is a + // larger provider token cap or a smaller ask, both caller decisions. + ErrLLMTruncated = errors.New("agent: LLM response truncated before completion") + + // ErrLLMContentBlocked is returned when the provider stopped generating + // for a content policy: safety, recitation, a blocklist, or an + // unsupported language. + // + // Deliberately distinct from ErrLLMTruncated because the two demand + // opposite responses. A truncation is length-dependent and may pass on + // the next attempt; a policy stop is deterministic for a given prompt, + // so every retry reproduces it. isRetryable treats it as terminal — the + // caller must change the request or surface the block to the user. + ErrLLMContentBlocked = errors.New("agent: LLM stopped generating for a content policy") + // ErrContextCancelled is returned when the request context is cancelled mid-loop. ErrContextCancelled = errors.New("agent: operation cancelled") @@ -201,3 +228,59 @@ func (e *LLMFailureError) Is(target error) bool { func (e *LLMFailureError) Unwrap() error { return e.Cause } + +// IncompleteKind classifies why a provider stopped generating early. Every +// provider has its own stop-reason vocabulary ("MAX_TOKENS", "length", +// "max_tokens"); each translates into these three cases so consumers route +// on one contract instead of three. +type IncompleteKind string + +const ( + // IncompleteTruncated: an output-length cap cut the response short. + IncompleteTruncated IncompleteKind = "truncated" + // IncompleteBlocked: a content policy stopped the generation. + IncompleteBlocked IncompleteKind = "blocked" + // IncompleteOther: neither a length cap nor a policy — a malformed + // tool call, or a backend stop the provider does not classify. The + // response is still partial, but it matches no sentinel, so retry + // behavior stays at the default. + IncompleteOther IncompleteKind = "other" +) + +// IncompleteResponseError reports a generation that ended before the model +// finished. Every provider returns it for the same situations — an output +// cap fired, or a content filter stopped the stream — so an adopter writes +// one branch rather than one per vendor: +// +// if errors.Is(err, agent.ErrLLMTruncated) { +// // raise the provider's token cap, or ask for less +// } +// +// Whatever streamed before the stop still rides on the returned LLMResult +// (content and usage), because it is real output that cost real tokens — +// it is a prefix, though, never a complete answer. Reason carries the raw +// provider stop reason for logs; Kind is what routing should key on, via +// the sentinels above. +type IncompleteResponseError struct { + // Provider names the vendor package that produced the error, matching + // the error-message prefix convention ("anthropic", "openai", "gemini"). + Provider string + // Reason is the provider's own stop reason, verbatim. + Reason string + // Kind is the provider-neutral classification. + Kind IncompleteKind +} + +func (e *IncompleteResponseError) Error() string { + return fmt.Sprintf("%s: generation stopped early (%s): the response is a partial prefix, not a complete answer", e.Provider, e.Reason) +} + +func (e *IncompleteResponseError) Is(target error) bool { + switch e.Kind { + case IncompleteTruncated: + return target == ErrLLMTruncated + case IncompleteBlocked: + return target == ErrLLMContentBlocked + } + return false +} diff --git a/pkg/agent/llm_call.go b/pkg/agent/llm_call.go index 9132997..40d27cd 100644 --- a/pkg/agent/llm_call.go +++ b/pkg/agent/llm_call.go @@ -14,8 +14,13 @@ import ( // streaming (consumer has already seen partial output). Returns the // final content, result, and (possibly retry-exhausted) error. func (al *AgentLoop) callLLMWithRetry(ctx context.Context, st *iterationState, msgs []history.Message) (string, LLMResult, error) { - finalContent, result, _, err := al.callLLM(ctx, st, msgs) - if err == nil || al.Retry == nil || !isRetryable(err) { + finalContent, result, contentEmitted, err := al.callLLM(ctx, st, msgs) + // The content gate applies to the first attempt too, not just to + // retries: once a partial answer has reached the consumer, a retry + // replays the whole answer into the same stream. Truncation errors + // always arrive this way, so without the gate every capped response + // would both duplicate itself and pay for a second full generation. + if err == nil || al.Retry == nil || contentEmitted || !isRetryable(err) { return finalContent, result, err } for attempt := 0; attempt < al.Retry.MaxRetries; attempt++ { diff --git a/pkg/agent/retry.go b/pkg/agent/retry.go index 1afaa6b..69cb99c 100644 --- a/pkg/agent/retry.go +++ b/pkg/agent/retry.go @@ -56,10 +56,14 @@ func (r *RetryConfig) delay(attempt int) time.Duration { return d } -// isRetryable returns false for context-level errors that should not be retried. +// isRetryable returns false for errors that fail identically on every +// attempt: context-level cancellation, and a provider content-policy stop +// (deterministic for a given prompt — retrying only burns the budget). func isRetryable(err error) bool { if err == nil { return false } - return !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) + return !errors.Is(err, context.Canceled) && + !errors.Is(err, context.DeadlineExceeded) && + !errors.Is(err, ErrLLMContentBlocked) } diff --git a/pkg/agent/retry_hook_test.go b/pkg/agent/retry_hook_test.go index ec5b543..4ee1afe 100644 --- a/pkg/agent/retry_hook_test.go +++ b/pkg/agent/retry_hook_test.go @@ -3,6 +3,7 @@ package agent import ( "context" "errors" + "fmt" "sync" "testing" "time" @@ -114,3 +115,67 @@ func TestRetry_OnAttemptNilIsZeroCost(t *testing.T) { // Compile-time check that history is wired through; the import is used by // the flakyProvider receiver method. var _ = history.Message{} + +// truncatingProvider streams a partial answer and then reports it as +// truncated — the shape every provider returns when its output cap fires. +type truncatingProvider struct { + mu sync.Mutex + calls int +} + +func (p *truncatingProvider) GenerateStream(_ context.Context, _ []history.Message, _ *tools.Registry, ch chan<- StreamEvent) (LLMResult, error) { + p.mu.Lock() + p.calls++ + p.mu.Unlock() + ch <- Event(ContentEvent{Text: "partial"}) + return LLMResult{Content: "partial"}, &IncompleteResponseError{ + Provider: "test", + Reason: "max_tokens", + Kind: IncompleteTruncated, + } +} + +func TestRetry_TruncationAfterStreamedContentIsNotRetried(t *testing.T) { + // Retrying would replay the whole answer into a stream the consumer has + // already read, and pay for a second cap-sized generation to do it. + prov := &truncatingProvider{} + loop, _ := setup(prov) + loop.Retry = &RetryConfig{MaxRetries: 3, BaseDelay: time.Millisecond, MaxDelay: time.Millisecond} + + if _, err := loop.RunIteration(context.Background(), "s1", "go"); err == nil { + t.Fatal("a truncated response must surface as an error") + } + + prov.mu.Lock() + defer prov.mu.Unlock() + if prov.calls != 1 { + t.Fatalf("provider called %d times, want 1 (no retry after streamed content)", prov.calls) + } +} + +func TestIsRetryable_ProviderStopClassification(t *testing.T) { + tests := []struct { + name string + err error + want bool + }{ + {"nil", nil, false}, + {"cancelled", context.Canceled, false}, + {"deadline", context.DeadlineExceeded, false}, + {"generic provider error", errors.New("connection reset"), true}, + // A length cut depends on this response, not the request, so the + // next attempt may well fit. + {"truncated", ErrLLMTruncated, true}, + {"wrapped truncated", fmt.Errorf("gemini: %w", ErrLLMTruncated), true}, + // A policy stop is deterministic: every retry reproduces it. + {"content blocked", ErrLLMContentBlocked, false}, + {"wrapped content blocked", fmt.Errorf("gemini: %w", ErrLLMContentBlocked), false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isRetryable(tt.err); got != tt.want { + t.Fatalf("isRetryable(%v) = %v, want %v", tt.err, got, tt.want) + } + }) + } +} From c30b119c04de47d1f54acce03450e7df802695b8 Mon Sep 17 00:00:00 2001 From: Hung Nguyen Date: Sat, 8 Aug 2026 10:53:18 +0700 Subject: [PATCH 2/2] fix(llm): report incomplete responses from all three providers Gemini never read FinishReason and OpenAI never read finish_reason, so a cap or content filter returned a valid prefix as a success; Anthropic only emitted a cap event. All three now return IncompleteResponseError, with the partial content and usage still on the result. --- pkg/llm/anthropic/anthropic.go | 21 +++- pkg/llm/anthropic/errors.go | 25 ++++ pkg/llm/anthropic/stop_reason_test.go | 167 ++++++++++++++++++++++++ pkg/llm/gemini/errors.go | 36 ++++++ pkg/llm/gemini/finish_reason_test.go | 174 ++++++++++++++++++++++++++ pkg/llm/gemini/gemini.go | 28 ++++- pkg/llm/gemini/gemini_vision.go | 10 +- pkg/llm/openai/errors.go | 26 ++++ pkg/llm/openai/finish_reason_test.go | 174 ++++++++++++++++++++++++++ pkg/llm/openai/openai.go | 22 +++- 10 files changed, 674 insertions(+), 9 deletions(-) create mode 100644 pkg/llm/anthropic/stop_reason_test.go create mode 100644 pkg/llm/gemini/finish_reason_test.go create mode 100644 pkg/llm/openai/finish_reason_test.go diff --git a/pkg/llm/anthropic/anthropic.go b/pkg/llm/anthropic/anthropic.go index c4dcf47..558ee2b 100644 --- a/pkg/llm/anthropic/anthropic.go +++ b/pkg/llm/anthropic/anthropic.go @@ -275,8 +275,9 @@ func (p *Provider) GenerateStream(ctx context.Context, memory []history.Message, // Surface the per-call MaxTokens truncation as a typed cap event so // adopters can render "model truncated; raise MaxTokens" instead of // silently shipping half-rendered code blocks. Skipped when soft - // truncation already fired the same event above. - if !truncated && string(accumulated.StopReason) == "max_tokens" { + // truncation already fired the same event above. The typed error is + // returned after the content is extracted, below. + if !truncated && accumulated.StopReason == anthropic.StopReasonMaxTokens { streamChan <- agent.LimitExhaustedStreamEvent(agent.LimitKindProviderMaxTokens, int(p.MaxTokens), 0) } @@ -319,7 +320,21 @@ func (p *Provider) GenerateStream(ctx context.Context, memory []history.Message, } usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens - return agent.LLMResult{Content: finalContent, ToolCalls: pendingCalls, Usage: usage}, nil + result := agent.LLMResult{Content: finalContent, ToolCalls: pendingCalls, Usage: usage} + // A capped or refused response is a prefix, not an answer. Returning it + // as a success is what turns a truncation into a decode error several + // layers up, naming the caller's schema instead of the real cause. The + // partial rides on the result so a host can still show what arrived. + // Soft truncation reports max_tokens too: the SDK could not finalize a + // tool_use block precisely because the cap fired mid-JSON. + stopReason := accumulated.StopReason + if truncated { + stopReason = anthropic.StopReasonMaxTokens + } + if err := stopReasonErr(stopReason); err != nil { + return result, err + } + return result, nil } // synthesizeStructuredTool builds a fake tool whose InputSchema diff --git a/pkg/llm/anthropic/errors.go b/pkg/llm/anthropic/errors.go index 16758b3..3ee3307 100644 --- a/pkg/llm/anthropic/errors.go +++ b/pkg/llm/anthropic/errors.go @@ -38,3 +38,28 @@ func classifyErr(err error) error { } return err } + +// stopReasonErr returns an *agent.IncompleteResponseError unless the +// generation ended cleanly, mapping Anthropic's stop reasons onto the +// provider-neutral classification consumers route on. +// +// "end_turn", "stop_sequence", and "tool_use" are complete responses. +// "pause_turn" is also complete — a long-running server tool asked the +// caller to continue the turn, and the content so far is intact. An empty +// reason means the API never reported one, treated as a clean stop so the +// default path is unchanged. +func stopReasonErr(r anthropic.StopReason) error { + var kind agent.IncompleteKind + switch r { + case "", anthropic.StopReasonEndTurn, anthropic.StopReasonStopSequence, + anthropic.StopReasonToolUse, anthropic.StopReasonPauseTurn: + return nil + case anthropic.StopReasonMaxTokens: + kind = agent.IncompleteTruncated + case anthropic.StopReasonRefusal: + kind = agent.IncompleteBlocked + default: + kind = agent.IncompleteOther + } + return &agent.IncompleteResponseError{Provider: "anthropic", Reason: string(r), Kind: kind} +} diff --git a/pkg/llm/anthropic/stop_reason_test.go b/pkg/llm/anthropic/stop_reason_test.go new file mode 100644 index 0000000..b1681d8 --- /dev/null +++ b/pkg/llm/anthropic/stop_reason_test.go @@ -0,0 +1,167 @@ +package anthropic + +import ( + "context" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/anthropics/anthropic-sdk-go" + "github.com/anthropics/anthropic-sdk-go/option" + "github.com/hung12ct/gopheragent/pkg/agent" + "github.com/hung12ct/gopheragent/pkg/history" +) + +// sse renders one Messages-API stream event. +func sse(event, data string) string { + return fmt.Sprintf("event: %s\ndata: %s\n\n", event, data) +} + +// textStream builds the event sequence for a single text block that ends +// with the given stop reason. +func textStream(text, stopReason string) []string { + return []string{ + sse("message_start", `{"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","model":"claude-test","content":[],"stop_reason":null,"usage":{"input_tokens":10,"output_tokens":1}}}`), + sse("content_block_start", `{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}`), + sse("content_block_delta", fmt.Sprintf(`{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":%q}}`, text)), + sse("content_block_stop", `{"type":"content_block_stop","index":0}`), + sse("message_delta", fmt.Sprintf(`{"type":"message_delta","delta":{"stop_reason":%q},"usage":{"output_tokens":8}}`, stopReason)), + sse("message_stop", `{"type":"message_stop"}`), + } +} + +// streamProvider returns a Provider whose Messages endpoint replays the +// given raw SSE frames, so the accumulation loop can be exercised without +// a live backend. +func streamProvider(t *testing.T, frames ...string) *Provider { + t.Helper() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + for _, f := range frames { + fmt.Fprint(w, f) + } + })) + t.Cleanup(srv.Close) + + client := anthropic.NewClient( + option.WithAPIKey("test-key"), + option.WithBaseURL(srv.URL+"/"), + ) + return &Provider{client: &client, model: anthropic.Model("claude-test"), MaxTokens: 1024} +} + +// runStream drives one GenerateStream call and returns the result, every +// event the provider emitted, and the call error. +func runStream(t *testing.T, p *Provider) (agent.LLMResult, []agent.StreamEvent, error) { + t.Helper() + ch := make(chan agent.StreamEvent, 64) + res, err := p.GenerateStream(context.Background(), []history.Message{{Role: "user", Content: "hi"}}, nil, ch) + close(ch) + var events []agent.StreamEvent + for ev := range ch { + events = append(events, ev) + } + return res, events, err +} + +func TestGenerateStream_MaxTokensSurfacesTruncation(t *testing.T) { + // The payload is a valid JSON prefix — the shape that reaches a caller + // as "unexpected end of JSON input" when the cap goes unreported. + p := streamProvider(t, textStream(`{"name":"ab`, "max_tokens")...) + + res, events, err := runStream(t, p) + if !errors.Is(err, agent.ErrLLMTruncated) { + t.Fatalf("errors.Is(ErrLLMTruncated) = false, got %v", err) + } + if errors.Is(err, agent.ErrLLMContentBlocked) { + t.Fatalf("a length cut must not classify as a content block: %v", err) + } + var inc *agent.IncompleteResponseError + if !errors.As(err, &inc) || inc.Provider != "anthropic" || inc.Reason != "max_tokens" { + t.Fatalf("want *agent.IncompleteResponseError{anthropic, max_tokens}, got %#v", err) + } + // The partial rides on the result so a host can show what arrived. + if res.Content != `{"name":"ab` { + t.Fatalf("partial content: got %q", res.Content) + } + + var limits int + for _, ev := range events { + if l, ok := ev.Payload.(agent.LimitExhaustedEvent); ok { + limits++ + if l.Kind != agent.LimitKindProviderMaxTokens || l.Limit != 1024 { + t.Fatalf("limit event: got %+v", l) + } + } + } + if limits != 1 { + t.Fatalf("want exactly 1 LimitExhaustedEvent, got %d", limits) + } +} + +func TestGenerateStream_RefusalIsBlocked(t *testing.T) { + p := streamProvider(t, textStream("I can't", "refusal")...) + + _, _, err := runStream(t, p) + if !errors.Is(err, agent.ErrLLMContentBlocked) { + t.Fatalf("errors.Is(ErrLLMContentBlocked) = false, got %v", err) + } + if errors.Is(err, agent.ErrLLMTruncated) { + t.Fatalf("a refusal must not classify as truncation: %v", err) + } +} + +func TestGenerateStream_CleanStopUnchanged(t *testing.T) { + p := streamProvider(t, textStream("hello world", "end_turn")...) + + res, events, err := runStream(t, p) + if err != nil { + t.Fatalf("clean stop must not error: %v", err) + } + if res.Content != "hello world" { + t.Fatalf("content: got %q", res.Content) + } + for _, ev := range events { + if _, ok := ev.Payload.(agent.LimitExhaustedEvent); ok { + t.Fatal("clean stop must not emit a limit event") + } + } +} + +func TestStopReasonErr_Classification(t *testing.T) { + tests := []struct { + reason anthropic.StopReason + wantErr bool + target error // nil = matches neither sentinel + }{ + {"", false, nil}, + {anthropic.StopReasonEndTurn, false, nil}, + {anthropic.StopReasonStopSequence, false, nil}, + {anthropic.StopReasonToolUse, false, nil}, + // pause_turn is a complete response the caller is asked to continue, + // not a cut one. + {anthropic.StopReasonPauseTurn, false, nil}, + {anthropic.StopReasonMaxTokens, true, agent.ErrLLMTruncated}, + {anthropic.StopReasonRefusal, true, agent.ErrLLMContentBlocked}, + {"model_context_window_exceeded", true, nil}, + } + for _, tt := range tests { + t.Run(string(tt.reason), func(t *testing.T) { + err := stopReasonErr(tt.reason) + if (err != nil) != tt.wantErr { + t.Fatalf("err = %v, wantErr %v", err, tt.wantErr) + } + if err == nil { + return + } + for _, sentinel := range []error{agent.ErrLLMTruncated, agent.ErrLLMContentBlocked} { + want := sentinel == tt.target + if errors.Is(err, sentinel) != want { + t.Fatalf("errors.Is(%v) = %v, want %v", sentinel, !want, want) + } + } + }) + } +} diff --git a/pkg/llm/gemini/errors.go b/pkg/llm/gemini/errors.go index c2e5b51..dde0940 100644 --- a/pkg/llm/gemini/errors.go +++ b/pkg/llm/gemini/errors.go @@ -37,3 +37,39 @@ func classifyErr(err error) error { } return err } + +// finishReasonErr returns an *agent.IncompleteResponseError unless the +// generation ended cleanly, mapping Gemini's finish reasons onto the +// provider-neutral classification consumers route on. An empty reason +// means the API never reported one (single-chunk responses on some +// endpoints); it is treated as a clean stop so the default path keeps its +// existing behavior exactly. +func finishReasonErr(r genai.FinishReason) error { + if r == "" || r == genai.FinishReasonStop || r == genai.FinishReasonUnspecified { + return nil + } + return &agent.IncompleteResponseError{ + Provider: "gemini", + Reason: string(r), + Kind: incompleteKind(r), + } +} + +// incompleteKind splits Gemini's finish reasons into the length cut that +// may pass on a different attempt and the content policy that will not. +// Reasons that are neither (OTHER, MALFORMED_FUNCTION_CALL, …) stay +// unclassified — the response is still partial, but nothing about the +// reason tells a caller whether to retry. +func incompleteKind(r genai.FinishReason) agent.IncompleteKind { + switch r { + case genai.FinishReasonMaxTokens: + return agent.IncompleteTruncated + case genai.FinishReasonSafety, genai.FinishReasonRecitation, + genai.FinishReasonBlocklist, genai.FinishReasonProhibitedContent, + genai.FinishReasonSPII, genai.FinishReasonLanguage, + genai.FinishReasonImageSafety, genai.FinishReasonImageProhibitedContent, + genai.FinishReasonImageRecitation: + return agent.IncompleteBlocked + } + return agent.IncompleteOther +} diff --git a/pkg/llm/gemini/finish_reason_test.go b/pkg/llm/gemini/finish_reason_test.go new file mode 100644 index 0000000..c6a8a23 --- /dev/null +++ b/pkg/llm/gemini/finish_reason_test.go @@ -0,0 +1,174 @@ +package gemini + +import ( + "context" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/hung12ct/gopheragent/pkg/agent" + "github.com/hung12ct/gopheragent/pkg/history" + "google.golang.org/genai" +) + +// streamProvider returns a Provider whose streaming endpoint replays the +// given JSON chunks as server-sent events, so the accumulation loop can be +// exercised without a live Gemini backend. +func streamProvider(t *testing.T, chunks ...string) *Provider { + t.Helper() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + for _, c := range chunks { + fmt.Fprintf(w, "data: %s\r\n\r\n", c) + } + })) + t.Cleanup(srv.Close) + + client, err := genai.NewClient(context.Background(), &genai.ClientConfig{ + APIKey: "test-key", + Backend: genai.BackendGeminiAPI, + HTTPOptions: genai.HTTPOptions{BaseURL: srv.URL}, + }) + if err != nil { + t.Fatalf("new client: %v", err) + } + return &Provider{client: client, model: "gemini-test"} +} + +// runStream drives one GenerateStream call and returns the result, the +// error, and every event the provider emitted. +func runStream(t *testing.T, p *Provider) (agent.LLMResult, []agent.StreamEvent, error) { + t.Helper() + ch := make(chan agent.StreamEvent, 64) + res, err := p.GenerateStream(context.Background(), []history.Message{{Role: "user", Content: "hi"}}, nil, ch) + close(ch) + var events []agent.StreamEvent + for ev := range ch { + events = append(events, ev) + } + return res, events, err +} + +func TestGenerateStream_MaxTokensSurfacesTruncation(t *testing.T) { + // The payload is a valid JSON prefix — exactly the shape that used to + // reach the caller as a success and fail at json.Unmarshal upstream. + p := streamProvider(t, + `{"candidates":[{"content":{"role":"model","parts":[{"text":"{\"name\":\"ab"}]}}],"usageMetadata":{"promptTokenCount":10,"candidatesTokenCount":5,"totalTokenCount":15}}`, + `{"candidates":[{"finishReason":"MAX_TOKENS"}],"usageMetadata":{"promptTokenCount":10,"candidatesTokenCount":8,"totalTokenCount":18}}`, + ) + + res, events, err := runStream(t, p) + if err == nil { + t.Fatal("truncated stream must not return a nil error") + } + if !errors.Is(err, agent.ErrLLMTruncated) { + t.Fatalf("errors.Is(ErrLLMTruncated) = false, got %v", err) + } + if errors.Is(err, agent.ErrLLMContentBlocked) { + t.Fatalf("a length cut must not classify as a content block: %v", err) + } + var inc *agent.IncompleteResponseError + if !errors.As(err, &inc) || inc.Reason != "MAX_TOKENS" || inc.Provider != "gemini" { + t.Fatalf("want *agent.IncompleteResponseError{gemini, MAX_TOKENS}, got %#v", err) + } + // The partial rides on the result so adopters can show what arrived. + if res.Content != `{"name":"ab` { + t.Fatalf("partial content: got %q", res.Content) + } + if res.Usage.TotalTokens != 18 { + t.Fatalf("usage must survive the error path: got %+v", res.Usage) + } + + var limits int + for _, ev := range events { + if l, ok := ev.Payload.(agent.LimitExhaustedEvent); ok { + limits++ + if l.Kind != agent.LimitKindProviderMaxTokens { + t.Fatalf("limit kind: got %q", l.Kind) + } + if l.Used != 8 { + t.Fatalf("limit used: want 8, got %d", l.Used) + } + } + } + if limits != 1 { + t.Fatalf("want exactly 1 LimitExhaustedEvent, got %d", limits) + } +} + +func TestGenerateStream_SafetyStopWithoutContent(t *testing.T) { + // A content filter blanks the candidate's Content, so the reason must be + // read before the nil-Content guard skips the chunk. + p := streamProvider(t, + `{"candidates":[{"finishReason":"SAFETY"}]}`, + ) + + _, _, err := runStream(t, p) + if !errors.Is(err, agent.ErrLLMContentBlocked) { + t.Fatalf("errors.Is(ErrLLMContentBlocked) = false, got %v", err) + } + if errors.Is(err, agent.ErrLLMTruncated) { + t.Fatalf("a policy stop must not classify as truncation: %v", err) + } +} + +func TestGenerateStream_CleanStopUnchanged(t *testing.T) { + p := streamProvider(t, + `{"candidates":[{"content":{"role":"model","parts":[{"text":"hello "}]}}]}`, + `{"candidates":[{"content":{"role":"model","parts":[{"text":"world"}]},"finishReason":"STOP"}]}`, + ) + + res, events, err := runStream(t, p) + if err != nil { + t.Fatalf("clean stop must not error: %v", err) + } + if res.Content != "hello world" { + t.Fatalf("content: got %q", res.Content) + } + for _, ev := range events { + if _, ok := ev.Payload.(agent.LimitExhaustedEvent); ok { + t.Fatal("clean stop must not emit a limit event") + } + } +} + +func TestFinishReasonErr_Classification(t *testing.T) { + tests := []struct { + reason genai.FinishReason + wantErr bool + target error // nil = matches neither sentinel + }{ + {"", false, nil}, + {genai.FinishReasonStop, false, nil}, + {genai.FinishReasonUnspecified, false, nil}, + {genai.FinishReasonMaxTokens, true, agent.ErrLLMTruncated}, + {genai.FinishReasonSafety, true, agent.ErrLLMContentBlocked}, + {genai.FinishReasonRecitation, true, agent.ErrLLMContentBlocked}, + {genai.FinishReasonBlocklist, true, agent.ErrLLMContentBlocked}, + {genai.FinishReasonProhibitedContent, true, agent.ErrLLMContentBlocked}, + {genai.FinishReasonSPII, true, agent.ErrLLMContentBlocked}, + {genai.FinishReasonLanguage, true, agent.ErrLLMContentBlocked}, + {genai.FinishReasonImageSafety, true, agent.ErrLLMContentBlocked}, + {genai.FinishReasonMalformedFunctionCall, true, nil}, + {genai.FinishReasonOther, true, nil}, + } + for _, tt := range tests { + t.Run(string(tt.reason), func(t *testing.T) { + err := finishReasonErr(tt.reason) + if (err != nil) != tt.wantErr { + t.Fatalf("err = %v, wantErr %v", err, tt.wantErr) + } + if err == nil { + return + } + for _, sentinel := range []error{agent.ErrLLMTruncated, agent.ErrLLMContentBlocked} { + want := sentinel == tt.target + if errors.Is(err, sentinel) != want { + t.Fatalf("errors.Is(%v) = %v, want %v", sentinel, !want, want) + } + } + }) + } +} diff --git a/pkg/llm/gemini/gemini.go b/pkg/llm/gemini/gemini.go index 0767d75..59b5518 100644 --- a/pkg/llm/gemini/gemini.go +++ b/pkg/llm/gemini/gemini.go @@ -173,6 +173,7 @@ func (p *Provider) GenerateStream(ctx context.Context, memory []history.Message, var finalContent string var pendingCalls []agent.PendingToolCall var usage agent.TokenUsage + var finishReason genai.FinishReason for resp, err := range iter { if err != nil { @@ -189,7 +190,16 @@ func (p *Provider) GenerateStream(ctx context.Context, memory []history.Message, TotalTokens: int(resp.UsageMetadata.TotalTokenCount), } } - if len(resp.Candidates) == 0 || resp.Candidates[0].Content == nil { + if len(resp.Candidates) == 0 { + continue + } + // Read the reason before the Content guard below: a stream stopped + // by a content filter carries the reason on a candidate whose + // Content is nil, so skipping early would drop the only signal. + if r := resp.Candidates[0].FinishReason; r != "" { + finishReason = r + } + if resp.Candidates[0].Content == nil { continue } @@ -209,11 +219,23 @@ func (p *Provider) GenerateStream(ctx context.Context, memory []history.Message, } } - return agent.LLMResult{ + result := agent.LLMResult{ Content: finalContent, ToolCalls: pendingCalls, Usage: usage, - }, nil + } + // A non-STOP finish reason means the accumulated text is a prefix, not + // an answer. Returning it as a success is what turns a truncation into + // a decode error several layers up, naming the caller's schema instead + // of the real cause. The partial rides on the result for adopters that + // want to show what arrived. + if err := finishReasonErr(finishReason); err != nil { + if finishReason == genai.FinishReasonMaxTokens { + streamChan <- agent.LimitExhaustedStreamEvent(agent.LimitKindProviderMaxTokens, 0, usage.CompletionTokens) + } + return result, err + } + return result, nil } // applySampling stamps the configured temperature/top_p/seed onto config. diff --git a/pkg/llm/gemini/gemini_vision.go b/pkg/llm/gemini/gemini_vision.go index 2c79eeb..cbdb25d 100644 --- a/pkg/llm/gemini/gemini_vision.go +++ b/pkg/llm/gemini/gemini_vision.go @@ -61,7 +61,15 @@ func (a *MediaAnalyzer) Analyze(ctx context.Context, media, prompt string) (stri if err != nil { return "", fmt.Errorf("gemini: media: %w", err) } - if len(resp.Candidates) == 0 || resp.Candidates[0].Content == nil { + if len(resp.Candidates) == 0 { + return "", fmt.Errorf("gemini: media: no content returned") + } + // Same silent-prefix trap as the streaming path: a non-STOP reason + // means the parts below are a partial answer, not the whole one. + if reasonErr := finishReasonErr(resp.Candidates[0].FinishReason); reasonErr != nil { + return "", fmt.Errorf("gemini: media: %w", reasonErr) + } + if resp.Candidates[0].Content == nil { return "", fmt.Errorf("gemini: media: no content returned") } var sb strings.Builder diff --git a/pkg/llm/openai/errors.go b/pkg/llm/openai/errors.go index d2f3ac9..cb0351b 100644 --- a/pkg/llm/openai/errors.go +++ b/pkg/llm/openai/errors.go @@ -52,3 +52,29 @@ func classifyErr(err error) error { } return err } + +// finishReasonErr returns an *agent.IncompleteResponseError unless the +// generation ended cleanly, mapping OpenAI's finish reasons onto the +// provider-neutral classification consumers route on. +// +// "stop", "tool_calls", and "function_call" are complete responses; an +// empty or "null" reason means the API never reported one, which is +// treated as a clean stop so the default path is unchanged. +func finishReasonErr(r openai.FinishReason) error { + var kind agent.IncompleteKind + switch r { + case "", openai.FinishReasonNull, openai.FinishReasonStop, + openai.FinishReasonToolCalls, openai.FinishReasonFunctionCall: + return nil + case openai.FinishReasonLength: + kind = agent.IncompleteTruncated + case openai.FinishReasonContentFilter: + kind = agent.IncompleteBlocked + default: + // Compat backends (Ollama, Groq, vLLM) invent their own reasons. + // Unknown means the response is not a documented clean stop, so + // report it as partial rather than assume it finished. + kind = agent.IncompleteOther + } + return &agent.IncompleteResponseError{Provider: "openai", Reason: string(r), Kind: kind} +} diff --git a/pkg/llm/openai/finish_reason_test.go b/pkg/llm/openai/finish_reason_test.go new file mode 100644 index 0000000..17fdcec --- /dev/null +++ b/pkg/llm/openai/finish_reason_test.go @@ -0,0 +1,174 @@ +package openai + +import ( + "context" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/hung12ct/gopheragent/pkg/agent" + "github.com/hung12ct/gopheragent/pkg/history" + "github.com/sashabaranov/go-openai" +) + +// streamProvider returns a Provider whose chat-completions endpoint +// replays the given JSON chunks as server-sent events, so the accumulation +// loop can be exercised without a live backend. +func streamProvider(t *testing.T, chunks ...string) *Provider { + t.Helper() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + for _, c := range chunks { + fmt.Fprintf(w, "data: %s\n\n", c) + } + fmt.Fprint(w, "data: [DONE]\n\n") + })) + t.Cleanup(srv.Close) + + p, err := NewCompat("test-key", "gpt-test", srv.URL+"/v1") + if err != nil { + t.Fatalf("NewCompat: %v", err) + } + return p +} + +// runStream drives one GenerateStream call and returns the result, every +// event the provider emitted, and the call error. +func runStream(t *testing.T, p *Provider) (agent.LLMResult, []agent.StreamEvent, error) { + t.Helper() + ch := make(chan agent.StreamEvent, 64) + res, err := p.GenerateStream(context.Background(), []history.Message{{Role: "user", Content: "hi"}}, nil, ch) + close(ch) + var events []agent.StreamEvent + for ev := range ch { + events = append(events, ev) + } + return res, events, err +} + +func TestGenerateStream_LengthSurfacesTruncation(t *testing.T) { + // The payload is a valid JSON prefix — the shape that used to reach the + // caller as a success and fail at json.Unmarshal upstream. + p := streamProvider(t, + `{"id":"1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"{\"name\":\"ab"}}]}`, + `{"id":"1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{},"finish_reason":"length"}]}`, + ) + + res, events, err := runStream(t, p) + if !errors.Is(err, agent.ErrLLMTruncated) { + t.Fatalf("errors.Is(ErrLLMTruncated) = false, got %v", err) + } + if errors.Is(err, agent.ErrLLMContentBlocked) { + t.Fatalf("a length cut must not classify as a content block: %v", err) + } + var inc *agent.IncompleteResponseError + if !errors.As(err, &inc) || inc.Provider != "openai" || inc.Reason != "length" { + t.Fatalf("want *agent.IncompleteResponseError{openai, length}, got %#v", err) + } + if res.Content != `{"name":"ab` { + t.Fatalf("partial content: got %q", res.Content) + } + + var limits int + for _, ev := range events { + if l, ok := ev.Payload.(agent.LimitExhaustedEvent); ok { + limits++ + if l.Kind != agent.LimitKindProviderMaxTokens { + t.Fatalf("limit kind: got %q", l.Kind) + } + } + } + if limits != 1 { + t.Fatalf("want exactly 1 LimitExhaustedEvent, got %d", limits) + } +} + +func TestGenerateStream_ContentFilterIsBlocked(t *testing.T) { + // A filtered response carries an empty delta, so the reason must be read + // before the delta is consumed. + p := streamProvider(t, + `{"id":"1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{},"finish_reason":"content_filter"}]}`, + ) + + _, events, err := runStream(t, p) + if !errors.Is(err, agent.ErrLLMContentBlocked) { + t.Fatalf("errors.Is(ErrLLMContentBlocked) = false, got %v", err) + } + if errors.Is(err, agent.ErrLLMTruncated) { + t.Fatalf("a policy stop must not classify as truncation: %v", err) + } + for _, ev := range events { + if _, ok := ev.Payload.(agent.LimitExhaustedEvent); ok { + t.Fatal("a content filter is not a token cap — no limit event") + } + } +} + +func TestGenerateStream_CleanStopUnchanged(t *testing.T) { + p := streamProvider(t, + `{"id":"1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"hello "}}]}`, + `{"id":"1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"world"},"finish_reason":"stop"}]}`, + ) + + res, _, err := runStream(t, p) + if err != nil { + t.Fatalf("clean stop must not error: %v", err) + } + if res.Content != "hello world" { + t.Fatalf("content: got %q", res.Content) + } +} + +func TestGenerateStream_ToolCallsAreACleanStop(t *testing.T) { + // tool_calls terminates a complete response; erroring here would break + // every ReAct turn that dispatches a tool. + p := streamProvider(t, + `{"id":"1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_1","function":{"name":"search","arguments":"{}"}}]}}]}`, + `{"id":"1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}`, + ) + + res, _, err := runStream(t, p) + if err != nil { + t.Fatalf("tool_calls must not error: %v", err) + } + if len(res.ToolCalls) != 1 || res.ToolCalls[0].Name != "search" { + t.Fatalf("tool calls: got %+v", res.ToolCalls) + } +} + +func TestFinishReasonErr_Classification(t *testing.T) { + tests := []struct { + reason openai.FinishReason + wantErr bool + target error // nil = matches neither sentinel + }{ + {"", false, nil}, + {openai.FinishReasonNull, false, nil}, + {openai.FinishReasonStop, false, nil}, + {openai.FinishReasonToolCalls, false, nil}, + {openai.FinishReasonFunctionCall, false, nil}, + {openai.FinishReasonLength, true, agent.ErrLLMTruncated}, + {openai.FinishReasonContentFilter, true, agent.ErrLLMContentBlocked}, + // Compat backends invent reasons; unknown is reported as partial. + {"eos", true, nil}, + } + for _, tt := range tests { + t.Run(string(tt.reason), func(t *testing.T) { + err := finishReasonErr(tt.reason) + if (err != nil) != tt.wantErr { + t.Fatalf("err = %v, wantErr %v", err, tt.wantErr) + } + if err == nil { + return + } + for _, sentinel := range []error{agent.ErrLLMTruncated, agent.ErrLLMContentBlocked} { + want := sentinel == tt.target + if errors.Is(err, sentinel) != want { + t.Fatalf("errors.Is(%v) = %v, want %v", sentinel, !want, want) + } + } + }) + } +} diff --git a/pkg/llm/openai/openai.go b/pkg/llm/openai/openai.go index 19f48f5..e8139d0 100644 --- a/pkg/llm/openai/openai.go +++ b/pkg/llm/openai/openai.go @@ -194,6 +194,7 @@ func (p *Provider) GenerateStream(ctx context.Context, memory []history.Message, var finalContent string var usage agent.TokenUsage + var finishReason openai.FinishReason // Accumulate parallel tool calls by index (OpenAI streams them split across chunks) type toolCallAccum struct { @@ -226,6 +227,11 @@ func (p *Provider) GenerateStream(ctx context.Context, memory []history.Message, if len(response.Choices) == 0 { continue } + // The stop reason lands on the final chunk, whose delta is empty — + // read it before anything below can skip the chunk. + if r := response.Choices[0].FinishReason; r != "" && r != openai.FinishReasonNull { + finishReason = r + } delta := response.Choices[0].Delta if delta.Content != "" { @@ -266,11 +272,23 @@ func (p *Provider) GenerateStream(ctx context.Context, memory []history.Message, }) } - return agent.LLMResult{ + result := agent.LLMResult{ Content: finalContent, ToolCalls: pendingCalls, Usage: usage, - }, nil + } + // A response cut by the token cap or stopped by a content filter is a + // prefix, not an answer. Returning it as a success is what turns a + // truncation into a decode error several layers up, naming the + // caller's schema instead of the real cause. The partial rides on the + // result so a host can still show what arrived. + if err := finishReasonErr(finishReason); err != nil { + if finishReason == openai.FinishReasonLength { + streamChan <- agent.LimitExhaustedStreamEvent(agent.LimitKindProviderMaxTokens, req.MaxTokens, usage.CompletionTokens) + } + return result, err + } + return result, nil } // applySampling stamps the configured temperature/top_p/seed onto req.