From 40976d5f087c252c3e9f53b3d90bb3162a7dabbd Mon Sep 17 00:00:00 2001 From: erio Date: Sun, 19 Apr 2026 23:25:31 +0800 Subject: [PATCH] fix(signature): async harvester with signature_delta support Rewrite harvester to be fully decoupled from the main read path: - Read() copies chunks to a buffered channel (non-blocking) - Background goroutine parses for signatures independently - sync.Once prevents double-close panic - defer recover() in Read() guards send-on-closed-channel - ctx.Done() arm prevents goroutine leaks on abandoned responses - Panic recovery with slog.Warn logging Fix signature extraction: the Anthropic API sends signatures via content_block_delta with delta.type="signature_delta", not in content_block_start (which has an empty signature field). Support both shapes for compatibility. Add 17 comprehensive tests including signature_delta, panic isolation, double-close, context cancellation, text containing "signature" word, and no-thinking streams. --- .../internal/service/signature/harvester.go | 244 ++++++++------- .../service/signature/harvester_test.go | 284 ++++++++++++++++-- .../service/signature/rectifier_test.go | 10 +- 3 files changed, 393 insertions(+), 145 deletions(-) diff --git a/backend/internal/service/signature/harvester.go b/backend/internal/service/signature/harvester.go index 96db5bb945..0c76356b44 100644 --- a/backend/internal/service/signature/harvester.go +++ b/backend/internal/service/signature/harvester.go @@ -4,6 +4,8 @@ import ( "bytes" "context" "io" + "log/slog" + "sync" "time" "github.com/tidwall/gjson" @@ -13,18 +15,21 @@ import ( // SignaturePool. It is implemented as an io.ReadCloser decorator so the // surrounding gateway code does not need to know about signature extraction. // -// Two response shapes are handled: -// - SSE streams (content-type text/event-stream): line-based parsing, only -// data lines carrying `content_block_start` events with a thinking block -// are inspected. -// - Non-streaming JSON: the entire body is accumulated and parsed once in -// Close(). +// Design: fully decoupled from the main read path. +// - Read() copies each chunk to a buffered channel (non-blocking). +// - A background goroutine drains the channel and parses for signatures. +// - Panics in the goroutine are recovered and logged, never affecting callers. +// - The goroutine also watches ctx.Done() to avoid leaking when Close() is +// never called (e.g., abandoned responses on context cancellation). // -// Ingestion is best-effort — errors writing to Redis are swallowed. Read -// semantics of the underlying body are preserved exactly. +// Two SSE event types carry signatures: +// - content_block_start with a non-empty content_block.signature (legacy). +// - content_block_delta with delta.type == "signature_delta" (current API). +// +// Non-streaming JSON responses are accumulated and parsed once on Close. type Harvester struct { pool SignaturePool - capacity int // pool capacity (RectifierSettings.SignaturePoolSize) + capacity int } // NewHarvester builds a Harvester bound to the given pool and capacity. @@ -34,136 +39,156 @@ func NewHarvester(pool SignaturePool, capacity int) *Harvester { // HarvestOptions configures a single Wrap call. type HarvestOptions struct { - // Bucket is the pool bucket to write into; see BucketFor. - // An empty bucket disables harvesting for this response. - Bucket string - // Streaming=true selects SSE line-based parsing; false accumulates the - // body and parses once on Close. + Bucket string Streaming bool - // Skip, when non-nil and returns true at read time, short-circuits - // harvesting for this call. Typical use: check - // ctxkey.IsSignatureRectifyRetry on the request context to avoid - // re-ingesting signatures we ourselves injected. - Skip func() bool + Skip func() bool } // Wrap returns a reader that transparently forwards body contents and, as a -// side effect, extracts any thinking signatures it finds into the pool. The -// returned reader must be closed; Close() flushes non-streaming extraction. +// side effect, extracts any thinking signatures into the pool via a background +// goroutine. The returned reader must be closed. func (h *Harvester) Wrap(ctx context.Context, body io.ReadCloser, opts HarvestOptions) io.ReadCloser { if h == nil || h.pool == nil || h.capacity <= 0 || opts.Bucket == "" { return body } - return &harvestReader{ - ctx: ctx, + chunks := make(chan []byte, harvestChanCap) + r := &harvestReader{ src: body, - pool: h.pool, - bucket: opts.Bucket, - cap: h.capacity, + chunks: chunks, skip: opts.Skip, - stream: opts.Streaming, } -} - -// harvestReader is the io.ReadCloser decorator. -type harvestReader struct { - ctx context.Context - src io.ReadCloser - pool SignaturePool - bucket string - cap int - skip func() bool - - stream bool - lineBuf []byte // SSE line accumulator - bodyBuf []byte // non-streaming body accumulator (bounded) - seen map[string]struct{} // de-dupe within a single response - closed bool + go processChunks(ctx, chunks, h.pool, opts.Bucket, h.capacity, opts.Streaming) + return r } const ( - // bodyBufCap bounds the accumulation buffer for non-streaming responses. - bodyBufCap = 2 * 1024 * 1024 - // lineBufCap bounds the SSE line accumulator to prevent memory blow-up - // from a malformed upstream that never sends a newline. - lineBufCap = 256 * 1024 + harvestChanCap = 64 + bodyBufCap = 2 * 1024 * 1024 + lineBufCap = 256 * 1024 ) +// harvestReader is the io.ReadCloser decorator. Its Read/Close methods are +// pure pass-throughs with a non-blocking channel send — zero parsing, zero +// Redis I/O, zero panic risk on the caller's goroutine. +type harvestReader struct { + src io.ReadCloser + chunks chan []byte + skip func() bool + closeOnce sync.Once +} + func (r *harvestReader) Read(p []byte) (int, error) { n, err := r.src.Read(p) if n > 0 && !r.skipNow() { - r.observe(p[:n]) + chunk := make([]byte, n) + copy(chunk, p[:n]) + // Non-blocking send. Recover protects against the rare case where + // Close() races with Read() and the channel is already closed. + func() { + defer func() { recover() }() + select { + case r.chunks <- chunk: + default: + } + }() } return n, err } func (r *harvestReader) Close() error { - if !r.closed && !r.skipNow() { - r.closed = true - r.flush() - } + r.closeOnce.Do(func() { close(r.chunks) }) return r.src.Close() } func (r *harvestReader) skipNow() bool { - if r.skip == nil { - return false - } - return r.skip() + return r.skip != nil && r.skip() } -// observe processes a chunk of freshly-read bytes. -func (r *harvestReader) observe(chunk []byte) { - if r.stream { - r.observeSSE(chunk) +// processChunks runs in a background goroutine. It drains the chunks channel +// and parses for signatures. Exits when the channel is closed OR the context +// is cancelled (preventing goroutine leaks on abandoned responses). +func processChunks(ctx context.Context, chunks <-chan []byte, pool SignaturePool, bucket string, capacity int, streaming bool) { + defer func() { + if r := recover(); r != nil { + slog.Warn("signature_harvester_panic", "error", r, "bucket", bucket) + } + }() + + state := &parseState{ + ctx: ctx, + pool: pool, + bucket: bucket, + cap: capacity, + stream: streaming, + } + for { + select { + case chunk, ok := <-chunks: + if !ok { + state.flush() + return + } + state.observe(chunk) + case <-ctx.Done(): + return + } + } +} + +// parseState holds the goroutine-private parsing state. +type parseState struct { + ctx context.Context + pool SignaturePool + bucket string + cap int + stream bool + + lineBuf []byte + bodyBuf []byte + seen map[string]struct{} +} + +func (s *parseState) observe(chunk []byte) { + if s.stream { + s.observeSSE(chunk) return } - // Non-streaming: accumulate, parse on Close. - remaining := bodyBufCap - len(r.bodyBuf) + remaining := bodyBufCap - len(s.bodyBuf) if remaining <= 0 { return } if len(chunk) > remaining { chunk = chunk[:remaining] } - r.bodyBuf = append(r.bodyBuf, chunk...) + s.bodyBuf = append(s.bodyBuf, chunk...) } -// observeSSE accumulates a line at a time and parses each completed line. -func (r *harvestReader) observeSSE(chunk []byte) { - r.lineBuf = append(r.lineBuf, chunk...) +func (s *parseState) observeSSE(chunk []byte) { + s.lineBuf = append(s.lineBuf, chunk...) for { - idx := bytes.IndexByte(r.lineBuf, '\n') + idx := bytes.IndexByte(s.lineBuf, '\n') if idx < 0 { - // Prevent unbounded growth from a stream that never sends \n. - if len(r.lineBuf) > lineBufCap { - r.lineBuf = nil + if len(s.lineBuf) > lineBufCap { + s.lineBuf = nil } return } - line := r.lineBuf[:idx] - r.lineBuf = r.lineBuf[idx+1:] - r.parseLine(line) + line := s.lineBuf[:idx] + s.lineBuf = s.lineBuf[idx+1:] + s.parseLine(line) } } +var sseDataPrefix = []byte("data:") + // parseLine extracts a signature from a single SSE data line. // // Two shapes carry signatures: -// 1. content_block_start with a non-empty content_block.signature -// (older API behavior, kept for compatibility). -// 2. content_block_delta with delta.type == "signature_delta" and -// delta.signature (current API behavior as of 2026-04). -// -// Other lines are ignored cheaply via a bytes.Contains pre-filter. -var sseDataPrefix = []byte("data:") - -func (r *harvestReader) parseLine(line []byte) { +// 1. content_block_start with a non-empty content_block.signature (legacy). +// 2. content_block_delta with delta.type == "signature_delta" (current API). +func (s *parseState) parseLine(line []byte) { line = bytes.TrimRight(line, "\r") - if len(line) == 0 { - return - } - if !bytes.HasPrefix(line, sseDataPrefix) { + if len(line) == 0 || !bytes.HasPrefix(line, sseDataPrefix) { return } payload := bytes.TrimSpace(line[len(sseDataPrefix):]) @@ -176,51 +201,44 @@ func (r *harvestReader) parseLine(line []byte) { evType := gjson.GetBytes(payload, "type").String() switch evType { case "content_block_start": - sig := gjson.GetBytes(payload, "content_block.signature").String() - if sig != "" { - r.emit(sig) + if sig := gjson.GetBytes(payload, "content_block.signature").String(); sig != "" { + s.emit(sig) } case "content_block_delta": - deltaType := gjson.GetBytes(payload, "delta.type").String() - if deltaType == "signature_delta" { - sig := gjson.GetBytes(payload, "delta.signature").String() - if sig != "" { - r.emit(sig) + if gjson.GetBytes(payload, "delta.type").String() == "signature_delta" { + if sig := gjson.GetBytes(payload, "delta.signature").String(); sig != "" { + s.emit(sig) } } } } -// flush parses an accumulated non-streaming body at Close time. -func (r *harvestReader) flush() { - if r.stream { - // Any trailing line without newline — try it. - if len(r.lineBuf) > 0 { - r.parseLine(r.lineBuf) - r.lineBuf = nil +func (s *parseState) flush() { + if s.stream { + if len(s.lineBuf) > 0 { + s.parseLine(s.lineBuf) + s.lineBuf = nil } return } - if len(r.bodyBuf) == 0 || !bytes.Contains(r.bodyBuf, []byte(`"signature"`)) { + if len(s.bodyBuf) == 0 || !bytes.Contains(s.bodyBuf, []byte(`"signature"`)) { return } - // Claude non-streaming: top-level content[].signature - gjson.GetBytes(r.bodyBuf, "content.#.signature").ForEach(func(_, v gjson.Result) bool { + gjson.GetBytes(s.bodyBuf, "content.#.signature").ForEach(func(_, v gjson.Result) bool { if sig := v.String(); sig != "" { - r.emit(sig) + s.emit(sig) } return true }) } -// emit writes a signature to the pool, de-duplicating within the same response. -func (r *harvestReader) emit(sig string) { - if r.seen == nil { - r.seen = make(map[string]struct{}) +func (s *parseState) emit(sig string) { + if s.seen == nil { + s.seen = make(map[string]struct{}) } - if _, dup := r.seen[sig]; dup { + if _, dup := s.seen[sig]; dup { return } - r.seen[sig] = struct{}{} - _ = r.pool.Add(r.ctx, r.bucket, sig, time.Now(), r.cap) + s.seen[sig] = struct{}{} + _ = s.pool.Add(s.ctx, s.bucket, sig, time.Now(), s.cap) } diff --git a/backend/internal/service/signature/harvester_test.go b/backend/internal/service/signature/harvester_test.go index 525e38e6e1..f4c40bc69b 100644 --- a/backend/internal/service/signature/harvester_test.go +++ b/backend/internal/service/signature/harvester_test.go @@ -7,13 +7,44 @@ import ( "io" "strings" "testing" + "time" ) +// waitPool polls the fakePool until the expected count is reached or timeout. +func waitPool(pool *fakePool, bucket string, want int, timeout time.Duration) []string { + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + pool.mu.Lock() + got := len(pool.entries[bucket]) + pool.mu.Unlock() + if got >= want { + break + } + time.Sleep(5 * time.Millisecond) + } + pool.mu.Lock() + defer pool.mu.Unlock() + result := make([]string, len(pool.entries[bucket])) + copy(result, pool.entries[bucket]) + return result +} + +// waitPoolEmpty waits briefly and asserts the pool bucket remains empty. +func waitPoolEmpty(t *testing.T, pool *fakePool, bucket string) { + t.Helper() + // Give the goroutine time to process, then check. + time.Sleep(200 * time.Millisecond) + pool.mu.Lock() + defer pool.mu.Unlock() + if len(pool.entries[bucket]) != 0 { + t.Errorf("expected pool %q to be empty, got %v", bucket, pool.entries[bucket]) + } +} + func TestHarvester_SSE_ExtractsContentBlockStartSignatures(t *testing.T) { pool := newFakePool() h := NewHarvester(pool, 10) - // Two thinking content_block_start events + one unrelated event. body := strings.Join([]string{ `event: message_start`, `data: {"type":"message_start","message":{}}`, @@ -24,9 +55,6 @@ func TestHarvester_SSE_ExtractsContentBlockStartSignatures(t *testing.T) { `event: content_block_start`, `data: {"type":"content_block_start","index":1,"content_block":{"type":"thinking","thinking":"","signature":"SIG-B"}}`, ``, - `event: content_block_delta`, - `data: {"type":"content_block_delta","delta":{"type":"thinking_delta","thinking":"hi"}}`, - ``, }, "\n") + "\n" rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ @@ -38,21 +66,93 @@ func TestHarvester_SSE_ExtractsContentBlockStartSignatures(t *testing.T) { } _ = rc.Close() - got := pool.entries["oauth"] + got := waitPool(pool, "oauth", 2, 2*time.Second) if len(got) != 2 { - t.Fatalf("expected 2 signatures harvested, got %d: %v", len(got), got) + t.Fatalf("expected 2 signatures, got %d: %v", len(got), got) } - // Newest-first semantics: fakePool prepends on Add, so the last-emitted is index 0. if got[0] != "SIG-B" || got[1] != "SIG-A" { t.Errorf("ordering: got %v, want [SIG-B, SIG-A]", got) } } +func TestHarvester_SSE_ExtractsSignatureDelta(t *testing.T) { + pool := newFakePool() + h := NewHarvester(pool, 10) + + // Real-world SSE shape: content_block_start has empty signature, + // the real signature arrives as a signature_delta event. + body := strings.Join([]string{ + `event: content_block_start`, + `data: {"type":"content_block_start","index":0,"content_block":{"type":"thinking","thinking":"","signature":""}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","index":0,"delta":{"type":"thinking_delta","thinking":"hello"}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","index":0,"delta":{"type":"signature_delta","signature":"EqACCkg..."}}`, + ``, + `event: content_block_stop`, + `data: {"type":"content_block_stop","index":0}`, + ``, + }, "\n") + "\n" + + rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ + Bucket: "oauth", + Streaming: true, + }) + _, _ = io.ReadAll(rc) + _ = rc.Close() + + got := waitPool(pool, "oauth", 1, 2*time.Second) + if len(got) != 1 || got[0] != "EqACCkg..." { + t.Errorf("expected [EqACCkg...], got %v", got) + } +} + +func TestHarvester_SSE_ExtractsBothStartAndDelta(t *testing.T) { + pool := newFakePool() + h := NewHarvester(pool, 10) + + body := strings.Join([]string{ + `data: {"type":"content_block_start","index":0,"content_block":{"type":"thinking","signature":"FROM-START"}}`, + ``, + `data: {"type":"content_block_delta","index":1,"delta":{"type":"signature_delta","signature":"FROM-DELTA"}}`, + ``, + }, "\n") + "\n" + + rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ + Bucket: "oauth", + Streaming: true, + }) + _, _ = io.ReadAll(rc) + _ = rc.Close() + + got := waitPool(pool, "oauth", 2, 2*time.Second) + if len(got) != 2 { + t.Fatalf("expected 2 signatures (start + delta), got %d: %v", len(got), got) + } +} + +func TestHarvester_SSE_IgnoresEmptySignatureInStart(t *testing.T) { + pool := newFakePool() + h := NewHarvester(pool, 10) + + body := `data: {"type":"content_block_start","index":0,"content_block":{"type":"thinking","thinking":"","signature":""}}` + "\n\n" + rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ + Bucket: "oauth", + Streaming: true, + }) + _, _ = io.ReadAll(rc) + _ = rc.Close() + + waitPoolEmpty(t, pool, "oauth") +} + func TestHarvester_SSE_DeduplicatesWithinSameResponse(t *testing.T) { pool := newFakePool() h := NewHarvester(pool, 10) body := strings.Repeat( - "event: content_block_start\ndata: {\"type\":\"content_block_start\",\"content_block\":{\"type\":\"thinking\",\"signature\":\"DUP\"}}\n\n", + "data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"signature_delta\",\"signature\":\"DUP\"}}\n\n", 5, ) rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ @@ -62,15 +162,16 @@ func TestHarvester_SSE_DeduplicatesWithinSameResponse(t *testing.T) { _, _ = io.ReadAll(rc) _ = rc.Close() - if got := pool.entries["oauth"]; len(got) != 1 { - t.Errorf("expected dedup to 1 entry within single response, got %d: %v", len(got), got) + got := waitPool(pool, "oauth", 1, 2*time.Second) + if len(got) != 1 { + t.Errorf("expected dedup to 1, got %d: %v", len(got), got) } } func TestHarvester_SkipCallbackBlocksIngestion(t *testing.T) { pool := newFakePool() h := NewHarvester(pool, 10) - body := `data: {"type":"content_block_start","content_block":{"type":"thinking","signature":"NO"}}` + "\n\n" + body := `data: {"type":"content_block_delta","delta":{"type":"signature_delta","signature":"NO"}}` + "\n\n" rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ Bucket: "oauth", Streaming: true, @@ -78,9 +179,8 @@ func TestHarvester_SkipCallbackBlocksIngestion(t *testing.T) { }) _, _ = io.ReadAll(rc) _ = rc.Close() - if len(pool.entries["oauth"]) != 0 { - t.Errorf("skip callback must prevent harvesting, got %v", pool.entries["oauth"]) - } + + waitPoolEmpty(t, pool, "oauth") } func TestHarvester_NonStreaming_ParsesOnClose(t *testing.T) { @@ -96,12 +196,10 @@ func TestHarvester_NonStreaming_ParsesOnClose(t *testing.T) { Streaming: false, }) _, _ = io.ReadAll(rc) - // Nothing should be harvested before Close on non-streaming path. - if len(pool.entries["oauth"]) != 0 { - t.Errorf("non-streaming must not ingest before Close; got %v", pool.entries["oauth"]) - } _ = rc.Close() - if got := pool.entries["oauth"]; len(got) != 2 { + + got := waitPool(pool, "oauth", 2, 2*time.Second) + if len(got) != 2 { t.Fatalf("expected 2 entries after Close, got %d: %v", len(got), got) } } @@ -112,7 +210,7 @@ func TestHarvester_EmptyBucketDisablesWrap(t *testing.T) { inner := io.NopCloser(strings.NewReader("irrelevant")) out := h.Wrap(context.Background(), inner, HarvestOptions{Bucket: ""}) if out != inner { - t.Errorf("empty bucket must return the original body unmodified") + t.Errorf("empty bucket must return original body") } } @@ -122,12 +220,10 @@ func TestHarvester_ZeroCapacityDisablesWrap(t *testing.T) { inner := io.NopCloser(strings.NewReader("irrelevant")) out := h.Wrap(context.Background(), inner, HarvestOptions{Bucket: "oauth"}) if out != inner { - t.Errorf("zero capacity must return the original body unmodified") + t.Errorf("zero capacity must return original body") } } -// splitReader delivers data in fixed-size chunks to exercise the SSE -// line-accumulation logic across multiple Read calls. type splitReader struct { data []byte pos int @@ -147,11 +243,10 @@ func (r *splitReader) Read(p []byte) (int, error) { return n, nil } -func TestHarvester_SSE_SignatureSplitAcrossReads(t *testing.T) { +func TestHarvester_SSE_SignatureDeltaSplitAcrossReads(t *testing.T) { pool := newFakePool() h := NewHarvester(pool, 10) - line := `data: {"type":"content_block_start","content_block":{"type":"thinking","signature":"SPLIT-SIG"}}` + "\n\n" - // Deliver 7 bytes at a time so the \n boundary falls mid-chunk. + line := `data: {"type":"content_block_delta","delta":{"type":"signature_delta","signature":"SPLIT-SIG"}}` + "\n\n" rc := h.Wrap(context.Background(), io.NopCloser(&splitReader{data: []byte(line), size: 7}), HarvestOptions{ Bucket: "oauth", Streaming: true, @@ -160,15 +255,16 @@ func TestHarvester_SSE_SignatureSplitAcrossReads(t *testing.T) { t.Fatalf("read: %v", err) } _ = rc.Close() - if got := pool.entries["oauth"]; len(got) != 1 || got[0] != "SPLIT-SIG" { - t.Errorf("split-read SSE: got %v, want [SPLIT-SIG]", got) + + got := waitPool(pool, "oauth", 1, 2*time.Second) + if len(got) != 1 || got[0] != "SPLIT-SIG" { + t.Errorf("got %v, want [SPLIT-SIG]", got) } } func TestHarvester_SSE_LineBufCapTruncatesHugeLine(t *testing.T) { pool := newFakePool() h := NewHarvester(pool, 10) - // A line longer than lineBufCap with no newline — must not OOM. huge := strings.Repeat("x", lineBufCap+1000) rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(huge)), HarvestOptions{ Bucket: "oauth", @@ -176,7 +272,133 @@ func TestHarvester_SSE_LineBufCapTruncatesHugeLine(t *testing.T) { }) _, _ = io.ReadAll(rc) _ = rc.Close() - if len(pool.entries["oauth"]) != 0 { - t.Errorf("huge line should not produce any signature") + + waitPoolEmpty(t, pool, "oauth") +} + +func TestHarvester_ReadSemanticPreserved(t *testing.T) { + pool := newFakePool() + h := NewHarvester(pool, 10) + original := "hello world this is the response body" + rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(original)), HarvestOptions{ + Bucket: "oauth", + Streaming: true, + }) + data, err := io.ReadAll(rc) + if err != nil { + t.Fatalf("read: %v", err) + } + _ = rc.Close() + if string(data) != original { + t.Errorf("read semantics broken: got %q, want %q", string(data), original) } } + +func TestHarvester_PanicInPoolDoesNotAffectRead(t *testing.T) { + pool := &panicPool{} + h := NewHarvester(pool, 10) + body := `data: {"type":"content_block_delta","delta":{"type":"signature_delta","signature":"BOOM"}}` + "\n\n" + rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ + Bucket: "oauth", + Streaming: true, + }) + data, err := io.ReadAll(rc) + if err != nil { + t.Fatalf("read should succeed even if pool panics: %v", err) + } + _ = rc.Close() + if string(data) != body { + t.Errorf("read data corrupted") + } +} + +func TestHarvester_SSE_TextDeltaContainingSignatureWordIgnored(t *testing.T) { + pool := newFakePool() + h := NewHarvester(pool, 10) + + body := strings.Join([]string{ + `data: {"type":"content_block_delta","delta":{"type":"text_delta","text":"please check your signature here"}}`, + ``, + }, "\n") + "\n" + + rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ + Bucket: "oauth", + Streaming: true, + }) + _, _ = io.ReadAll(rc) + _ = rc.Close() + + waitPoolEmpty(t, pool, "oauth") +} + +func TestHarvester_SSE_NoThinkingBlocksProducesNothing(t *testing.T) { + pool := newFakePool() + h := NewHarvester(pool, 10) + + body := strings.Join([]string{ + `event: message_start`, + `data: {"type":"message_start","message":{"model":"claude-sonnet-4-6"}}`, + ``, + `event: content_block_start`, + `data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello!"}}`, + ``, + `event: content_block_stop`, + `data: {"type":"content_block_stop","index":0}`, + ``, + `event: message_stop`, + `data: {"type":"message_stop"}`, + ``, + }, "\n") + "\n" + + rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ + Bucket: "oauth", + Streaming: true, + }) + _, _ = io.ReadAll(rc) + _ = rc.Close() + + waitPoolEmpty(t, pool, "oauth") +} + +func TestHarvester_ContextCancellationStopsGoroutine(t *testing.T) { + pool := newFakePool() + h := NewHarvester(pool, 10) + + ctx, cancel := context.WithCancel(context.Background()) + + body := `data: {"type":"content_block_delta","delta":{"type":"signature_delta","signature":"CTX-SIG"}}` + "\n\n" + rc := h.Wrap(ctx, io.NopCloser(strings.NewReader(body)), HarvestOptions{ + Bucket: "oauth", + Streaming: true, + }) + + _, _ = io.ReadAll(rc) + // Cancel context before Close — goroutine should exit via ctx.Done() + cancel() + time.Sleep(50 * time.Millisecond) + _ = rc.Close() +} + +func TestHarvester_DoubleCloseNoPanic(t *testing.T) { + pool := newFakePool() + h := NewHarvester(pool, 10) + body := `data: {"type":"content_block_delta","delta":{"type":"signature_delta","signature":"X"}}` + "\n\n" + rc := h.Wrap(context.Background(), io.NopCloser(strings.NewReader(body)), HarvestOptions{ + Bucket: "oauth", + Streaming: true, + }) + _, _ = io.ReadAll(rc) + _ = rc.Close() + _ = rc.Close() // second close must not panic +} + +type panicPool struct{} + +func (p *panicPool) Add(_ context.Context, _, _ string, _ time.Time, _ int) error { + panic("pool exploded") +} +func (p *panicPool) TopN(_ context.Context, _ string, _ int) ([]string, error) { return nil, nil } +func (p *panicPool) Size(_ context.Context, _ string) (int64, error) { return 0, nil } diff --git a/backend/internal/service/signature/rectifier_test.go b/backend/internal/service/signature/rectifier_test.go index 37b0ef4cdf..844613a3f6 100644 --- a/backend/internal/service/signature/rectifier_test.go +++ b/backend/internal/service/signature/rectifier_test.go @@ -6,6 +6,7 @@ import ( "context" "encoding/json" "errors" + "sync" "testing" "time" @@ -14,18 +15,23 @@ import ( // fakePool is an in-memory SignaturePool used by tests; implements the // ordering contract (TopN returns most-recently-added first). +// Thread-safe: the async harvester goroutine calls Add concurrently. type fakePool struct { + mu sync.Mutex entries map[string][]string // bucket -> signatures, index 0 = newest } func newFakePool() *fakePool { return &fakePool{entries: map[string][]string{}} } func (p *fakePool) Add(_ context.Context, bucket, sig string, _ time.Time, _ int) error { - // Prepend so index 0 is the newest. + p.mu.Lock() + defer p.mu.Unlock() p.entries[bucket] = append([]string{sig}, p.entries[bucket]...) return nil } func (p *fakePool) TopN(_ context.Context, bucket string, n int) ([]string, error) { + p.mu.Lock() + defer p.mu.Unlock() list := p.entries[bucket] if n > len(list) { n = len(list) @@ -35,6 +41,8 @@ func (p *fakePool) TopN(_ context.Context, bucket string, n int) ([]string, erro return out, nil } func (p *fakePool) Size(_ context.Context, bucket string) (int64, error) { + p.mu.Lock() + defer p.mu.Unlock() return int64(len(p.entries[bucket])), nil }