mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
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:
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user