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.
This commit is contained in:
erio
2026-04-19 23:25:31 +08:00
parent 0ff3bfe5fa
commit 40976d5f08
3 changed files with 393 additions and 145 deletions
+131 -113
View File
@@ -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)
}
@@ -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 }
@@ -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
}