From 0121da07601ac0a91d4dc759cdfd3e03d4c4bb4b Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 18 Mar 2026 20:24:34 +0100 Subject: [PATCH 01/11] Share SSE parsing across audit and usage observers --- internal/auditlog/stream_observer.go | 152 ++++++++++ internal/auditlog/stream_wrapper.go | 224 +------------- .../server/translated_inference_service.go | 11 +- internal/streaming/observed_sse_stream.go | 128 ++++++++ .../streaming/observed_sse_stream_test.go | 86 ++++++ internal/usage/stream_observer.go | 140 +++++++++ internal/usage/stream_wrapper.go | 278 +----------------- tests/perf/hotpath_test.go | 34 +-- 8 files changed, 543 insertions(+), 510 deletions(-) create mode 100644 internal/auditlog/stream_observer.go create mode 100644 internal/streaming/observed_sse_stream.go create mode 100644 internal/streaming/observed_sse_stream_test.go create mode 100644 internal/usage/stream_observer.go diff --git a/internal/auditlog/stream_observer.go b/internal/auditlog/stream_observer.go new file mode 100644 index 000000000..82218c6ea --- /dev/null +++ b/internal/auditlog/stream_observer.go @@ -0,0 +1,152 @@ +package auditlog + +import ( + "strings" + "time" +) + +// StreamLogObserver reconstructs stream metadata and optional response bodies +// from parsed SSE JSON payloads. +type StreamLogObserver struct { + logger LoggerInterface + entry *LogEntry + builder *streamResponseBuilder + logBodies bool + closed bool + startTime time.Time +} + +func NewStreamLogObserver(logger LoggerInterface, entry *LogEntry, path string) *StreamLogObserver { + if logger == nil || entry == nil { + return nil + } + + logBodies := logger.Config().LogBodies + var builder *streamResponseBuilder + if logBodies { + builder = &streamResponseBuilder{ + IsResponsesAPI: strings.HasPrefix(path, "/v1/responses"), + } + } + + return &StreamLogObserver{ + logger: logger, + entry: entry, + builder: builder, + logBodies: logBodies, + startTime: entry.Timestamp, + } +} + +func (o *StreamLogObserver) OnJSONEvent(event map[string]interface{}) { + if !o.logBodies || o.builder == nil { + return + } + if o.builder.IsResponsesAPI { + o.parseResponsesAPIEvent(event) + return + } + o.parseChatCompletionEvent(event) +} + +func (o *StreamLogObserver) OnStreamClose() { + if o.closed { + return + } + o.closed = true + + if o.entry != nil && !o.startTime.IsZero() { + o.entry.DurationNs = time.Since(o.startTime).Nanoseconds() + } + + if o.logBodies && o.builder != nil && o.entry != nil && o.entry.Data != nil { + if o.builder.IsResponsesAPI { + o.entry.Data.ResponseBody = o.builder.buildResponsesAPIResponse() + } else { + o.entry.Data.ResponseBody = o.builder.buildChatCompletionResponse() + } + o.entry.Data.ResponseBodyTooBigToHandle = o.builder.truncated + } + + if o.logger != nil && o.entry != nil { + o.logger.Write(o.entry) + } +} + +func (o *StreamLogObserver) parseChatCompletionEvent(event map[string]interface{}) { + if o.builder == nil { + return + } + + if o.builder.ID == "" { + if id, ok := event["id"].(string); ok { + o.builder.ID = id + } + if model, ok := event["model"].(string); ok { + o.builder.Model = model + } + if created, ok := event["created"].(float64); ok { + o.builder.Created = int64(created) + } + } + + if choices, ok := event["choices"].([]interface{}); ok && len(choices) > 0 { + if choice, ok := choices[0].(map[string]interface{}); ok { + if fr, ok := choice["finish_reason"].(string); ok && fr != "" { + o.builder.FinishReason = fr + } + + if delta, ok := choice["delta"].(map[string]interface{}); ok { + if role, ok := delta["role"].(string); ok { + o.builder.Role = role + } + if content, ok := delta["content"].(string); ok && content != "" { + o.appendContent(content) + } + } + } + } +} + +func (o *StreamLogObserver) parseResponsesAPIEvent(event map[string]interface{}) { + if o.builder == nil { + return + } + + eventType, _ := event["type"].(string) + switch eventType { + case "response.created", "response.completed", "response.done": + if resp, ok := event["response"].(map[string]interface{}); ok { + if id, ok := resp["id"].(string); ok { + o.builder.ResponseID = id + } + if status, ok := resp["status"].(string); ok { + o.builder.Status = status + } + if model, ok := resp["model"].(string); ok { + o.builder.Model = model + } + if createdAt, ok := resp["created_at"].(float64); ok { + o.builder.CreatedAt = int64(createdAt) + } + } + case "response.output_text.delta": + if delta, ok := event["delta"].(string); ok && delta != "" { + o.appendContent(delta) + } + } +} + +func (o *StreamLogObserver) appendContent(content string) { + if o.builder == nil || o.builder.truncated || o.builder.contentLen >= MaxContentCapture { + return + } + + remaining := MaxContentCapture - o.builder.contentLen + if len(content) > remaining { + content = content[:remaining] + o.builder.truncated = true + } + o.builder.Content.WriteString(content) + o.builder.contentLen += len(content) +} diff --git a/internal/auditlog/stream_wrapper.go b/internal/auditlog/stream_wrapper.go index ef6ec7fb7..233bc3bb1 100644 --- a/internal/auditlog/stream_wrapper.go +++ b/internal/auditlog/stream_wrapper.go @@ -1,11 +1,10 @@ package auditlog import ( - "bytes" - "encoding/json" "io" "strings" - "time" + + "gomodel/internal/streaming" ) // Note: MaxContentCapture and LogEntryStreamingKey constants are defined in constants.go @@ -35,232 +34,21 @@ type streamResponseBuilder struct { // for audit logging. type StreamLogWrapper struct { io.ReadCloser - logger LoggerInterface - entry *LogEntry - builder *streamResponseBuilder - logBodies bool - closed bool - startTime time.Time - pending []byte // pending partial SSE data between reads } // NewStreamLogWrapper creates a wrapper around a stream for audit logging. // When the stream is closed, it logs the accumulated entry. // The path parameter is used to detect whether this is a ChatCompletion or Responses API request. func NewStreamLogWrapper(stream io.ReadCloser, logger LoggerInterface, entry *LogEntry, path string) *StreamLogWrapper { - // Use entry's timestamp as start time for duration calculation - var startTime time.Time - if entry != nil { - startTime = entry.Timestamp - } - - // Check if body logging is enabled - logBodies := false - if logger != nil { - logBodies = logger.Config().LogBodies - } - - // Initialize builder if body logging is enabled - var builder *streamResponseBuilder - if logBodies { - builder = &streamResponseBuilder{ - IsResponsesAPI: strings.HasPrefix(path, "/v1/responses"), - } + observer := NewStreamLogObserver(logger, entry, path) + if observer == nil { + return &StreamLogWrapper{ReadCloser: stream} } - return &StreamLogWrapper{ - ReadCloser: stream, - logger: logger, - entry: entry, - startTime: startTime, - logBodies: logBodies, - builder: builder, + ReadCloser: streaming.NewObservedSSEStream(stream, observer), } } -// Read implements io.Reader and incrementally processes SSE chunks for audit capture. -func (w *StreamLogWrapper) Read(p []byte) (n int, err error) { - n, err = w.ReadCloser.Read(p) - if n > 0 { - // Parse SSE events and accumulate content if body logging is enabled - if w.logBodies && w.builder != nil { - w.processSSEData(p[:n]) - } - } - return n, err -} - -// processSSEData parses SSE events from the data chunk and accumulates content -func (w *StreamLogWrapper) processSSEData(data []byte) { - // Prepend any pending data from previous read - if len(w.pending) > 0 { - data = append(w.pending, data...) - w.pending = nil - } - - // Split on double newline (SSE event separator) - for { - idx := bytes.Index(data, []byte("\n\n")) - if idx == -1 { - // No complete event, save as pending - if len(data) > 0 { - w.pending = make([]byte, len(data)) - copy(w.pending, data) - } - return - } - - event := data[:idx] - data = data[idx+2:] - - w.processSSEEvent(event) - } -} - -// processSSEEvent processes a single SSE event -func (w *StreamLogWrapper) processSSEEvent(event []byte) { - // Find the data line - lines := bytes.Split(event, []byte("\n")) - for _, line := range lines { - if bytes.HasPrefix(line, []byte("data: ")) { - jsonData := bytes.TrimPrefix(line, []byte("data: ")) - // Skip [DONE] marker - if bytes.Equal(jsonData, []byte("[DONE]")) { - continue - } - w.parseEventJSON(jsonData) - } - } -} - -// parseEventJSON parses the JSON from an SSE event and accumulates data -func (w *StreamLogWrapper) parseEventJSON(data []byte) { - var event map[string]interface{} - if err := json.Unmarshal(data, &event); err != nil { - return - } - - if w.builder.IsResponsesAPI { - w.parseResponsesAPIEvent(event) - } else { - w.parseChatCompletionEvent(event) - } -} - -// parseChatCompletionEvent extracts data from a ChatCompletion streaming chunk -func (w *StreamLogWrapper) parseChatCompletionEvent(event map[string]interface{}) { - // Extract metadata from first event - if w.builder.ID == "" { - if id, ok := event["id"].(string); ok { - w.builder.ID = id - } - if model, ok := event["model"].(string); ok { - w.builder.Model = model - } - if created, ok := event["created"].(float64); ok { - w.builder.Created = int64(created) - } - } - - // Extract delta content from choices - if choices, ok := event["choices"].([]interface{}); ok && len(choices) > 0 { - if choice, ok := choices[0].(map[string]interface{}); ok { - // Extract finish_reason - if fr, ok := choice["finish_reason"].(string); ok && fr != "" { - w.builder.FinishReason = fr - } - - // Extract delta - if delta, ok := choice["delta"].(map[string]interface{}); ok { - // Extract role (usually in first chunk) - if role, ok := delta["role"].(string); ok { - w.builder.Role = role - } - // Extract and accumulate content - if content, ok := delta["content"].(string); ok && content != "" { - if !w.builder.truncated && w.builder.contentLen < MaxContentCapture { - remaining := MaxContentCapture - w.builder.contentLen - if len(content) > remaining { - content = content[:remaining] - w.builder.truncated = true - } - w.builder.Content.WriteString(content) - w.builder.contentLen += len(content) - } - } - } - } - } -} - -// parseResponsesAPIEvent extracts data from a Responses API streaming event -func (w *StreamLogWrapper) parseResponsesAPIEvent(event map[string]interface{}) { - eventType, _ := event["type"].(string) - - switch eventType { - case "response.created", "response.completed", "response.done": - // Extract response metadata - if resp, ok := event["response"].(map[string]interface{}); ok { - if id, ok := resp["id"].(string); ok { - w.builder.ResponseID = id - } - if status, ok := resp["status"].(string); ok { - w.builder.Status = status - } - if model, ok := resp["model"].(string); ok { - w.builder.Model = model - } - if createdAt, ok := resp["created_at"].(float64); ok { - w.builder.CreatedAt = int64(createdAt) - } - } - - case "response.output_text.delta": - // Accumulate text delta - if delta, ok := event["delta"].(string); ok && delta != "" { - if !w.builder.truncated && w.builder.contentLen < MaxContentCapture { - remaining := MaxContentCapture - w.builder.contentLen - if len(delta) > remaining { - delta = delta[:remaining] - w.builder.truncated = true - } - w.builder.Content.WriteString(delta) - w.builder.contentLen += len(delta) - } - } - } -} - -// Close implements io.Closer and logs the entry. -func (w *StreamLogWrapper) Close() error { - if w.closed { - return nil - } - w.closed = true - - // Calculate duration from start time - if w.entry != nil && !w.startTime.IsZero() { - w.entry.DurationNs = time.Since(w.startTime).Nanoseconds() - } - - // Build and store reconstructed response body if enabled - if w.logBodies && w.builder != nil && w.entry != nil && w.entry.Data != nil { - if w.builder.IsResponsesAPI { - w.entry.Data.ResponseBody = w.builder.buildResponsesAPIResponse() - } else { - w.entry.Data.ResponseBody = w.builder.buildChatCompletionResponse() - } - w.entry.Data.ResponseBodyTooBigToHandle = w.builder.truncated - } - - // Write log entry - if w.logger != nil && w.entry != nil { - w.logger.Write(w.entry) - } - - return w.ReadCloser.Close() -} - // buildChatCompletionResponse constructs a ChatCompletion response from accumulated data func (b *streamResponseBuilder) buildChatCompletionResponse() map[string]interface{} { return map[string]interface{}{ diff --git a/internal/server/translated_inference_service.go b/internal/server/translated_inference_service.go index f8a533c51..91730b393 100644 --- a/internal/server/translated_inference_service.go +++ b/internal/server/translated_inference_service.go @@ -10,6 +10,7 @@ import ( "gomodel/internal/auditlog" "gomodel/internal/core" + "gomodel/internal/streaming" "gomodel/internal/usage" ) @@ -143,11 +144,17 @@ func (s *translatedInferenceService) handleStreamingResponse(c *echo.Context, mo if streamEntry != nil { streamEntry.StatusCode = http.StatusOK } - wrappedStream := auditlog.WrapStreamForLogging(stream, s.logger, streamEntry, c.Request().URL.Path) requestID := requestIDFromContextOrHeader(c.Request()) endpoint := c.Request().URL.Path - wrappedStream = usage.WrapStreamForUsage(wrappedStream, s.usageLogger, model, provider, requestID, endpoint, s.pricingResolver) + observers := make([]streaming.Observer, 0, 2) + if s.logger != nil && s.logger.Config().Enabled && streamEntry != nil { + observers = append(observers, auditlog.NewStreamLogObserver(s.logger, streamEntry, endpoint)) + } + if s.usageLogger != nil && s.usageLogger.Config().Enabled { + observers = append(observers, usage.NewStreamUsageObserver(s.usageLogger, model, provider, requestID, endpoint, s.pricingResolver)) + } + wrappedStream := streaming.NewObservedSSEStream(stream, observers...) defer func() { _ = wrappedStream.Close() //nolint:errcheck diff --git a/internal/streaming/observed_sse_stream.go b/internal/streaming/observed_sse_stream.go new file mode 100644 index 000000000..2ec534eda --- /dev/null +++ b/internal/streaming/observed_sse_stream.go @@ -0,0 +1,128 @@ +package streaming + +import ( + "bytes" + "encoding/json" + "io" +) + +const maxPendingEventBytes = 256 * 1024 + +// Observer receives parsed JSON SSE payloads in stream order. +// Implementations must treat the payload as read-only. +type Observer interface { + OnJSONEvent(payload map[string]interface{}) + OnStreamClose() +} + +// ObservedSSEStream proxies bytes unchanged while parsing SSE JSON events once +// and fanning them out to observers. +type ObservedSSEStream struct { + io.ReadCloser + observers []Observer + pending []byte + closed bool +} + +// NewObservedSSEStream returns the original stream when there are no observers. +func NewObservedSSEStream(stream io.ReadCloser, observers ...Observer) io.ReadCloser { + filtered := make([]Observer, 0, len(observers)) + for _, observer := range observers { + if observer != nil { + filtered = append(filtered, observer) + } + } + if len(filtered) == 0 { + return stream + } + return &ObservedSSEStream{ + ReadCloser: stream, + observers: filtered, + } +} + +func (s *ObservedSSEStream) Read(p []byte) (n int, err error) { + n, err = s.ReadCloser.Read(p) + if n > 0 { + s.processChunk(p[:n]) + } + return n, err +} + +func (s *ObservedSSEStream) Close() error { + if s.closed { + return nil + } + s.closed = true + + if len(s.pending) > 0 { + s.processBufferedEvents(s.pending) + s.pending = nil + } + + for _, observer := range s.observers { + observer.OnStreamClose() + } + return s.ReadCloser.Close() +} + +func (s *ObservedSSEStream) processChunk(data []byte) { + if len(s.pending) > 0 { + combined := make([]byte, len(s.pending)+len(data)) + copy(combined, s.pending) + copy(combined[len(s.pending):], data) + data = combined + s.pending = nil + } + + for { + idx := bytes.Index(data, []byte("\n\n")) + if idx == -1 { + s.savePending(data) + return + } + + s.processEvent(data[:idx]) + data = data[idx+2:] + } +} + +func (s *ObservedSSEStream) processBufferedEvents(data []byte) { + for _, event := range bytes.Split(data, []byte("\n\n")) { + if len(event) == 0 { + continue + } + s.processEvent(event) + } +} + +func (s *ObservedSSEStream) processEvent(event []byte) { + lines := bytes.Split(event, []byte("\n")) + for _, line := range lines { + if !bytes.HasPrefix(line, []byte("data: ")) { + continue + } + jsonData := bytes.TrimPrefix(line, []byte("data: ")) + if bytes.Equal(jsonData, []byte("[DONE]")) { + continue + } + + var payload map[string]interface{} + if err := json.Unmarshal(jsonData, &payload); err != nil { + continue + } + for _, observer := range s.observers { + observer.OnJSONEvent(payload) + } + } +} + +func (s *ObservedSSEStream) savePending(data []byte) { + if len(data) == 0 { + return + } + if len(data) > maxPendingEventBytes { + data = data[len(data)-maxPendingEventBytes:] + } + s.pending = append(s.pending[:0], data...) +} diff --git a/internal/streaming/observed_sse_stream_test.go b/internal/streaming/observed_sse_stream_test.go new file mode 100644 index 000000000..8e24e8b4d --- /dev/null +++ b/internal/streaming/observed_sse_stream_test.go @@ -0,0 +1,86 @@ +package streaming + +import ( + "io" + "strings" + "testing" +) + +type trackingObserver struct { + eventCount int + lastID string + closed bool +} + +func (o *trackingObserver) OnJSONEvent(payload map[string]interface{}) { + o.eventCount++ + if id, _ := payload["id"].(string); id != "" { + o.lastID = id + } +} + +func (o *trackingObserver) OnStreamClose() { + o.closed = true +} + +func TestObservedSSEStream_PassesThroughAndFansOut(t *testing.T) { + streamData := `data: {"id":"chatcmpl-1","choices":[{"delta":{"content":"hi"}}]} + +data: {"id":"chatcmpl-2","usage":{"total_tokens":3}} + +data: [DONE] + +` + first := &trackingObserver{} + second := &trackingObserver{} + stream := NewObservedSSEStream(io.NopCloser(strings.NewReader(streamData)), first, second) + + data, err := io.ReadAll(stream) + if err != nil { + t.Fatalf("ReadAll error: %v", err) + } + if string(data) != streamData { + t.Fatalf("stream passthrough mismatch") + } + if err := stream.Close(); err != nil { + t.Fatalf("Close error: %v", err) + } + + for i, observer := range []*trackingObserver{first, second} { + if observer.eventCount != 2 { + t.Fatalf("observer %d eventCount = %d, want 2", i, observer.eventCount) + } + if observer.lastID != "chatcmpl-2" { + t.Fatalf("observer %d lastID = %q, want chatcmpl-2", i, observer.lastID) + } + if !observer.closed { + t.Fatalf("observer %d was not closed", i) + } + } +} + +func TestObservedSSEStream_ParsesFragmentedFinalEventOnClose(t *testing.T) { + streamData := `data: {"id":"chatcmpl-frag","usage":{"total_tokens":8}}` + observer := &trackingObserver{} + stream := NewObservedSSEStream(io.NopCloser(strings.NewReader(streamData)), observer) + + data, err := io.ReadAll(stream) + if err != nil { + t.Fatalf("ReadAll error: %v", err) + } + if string(data) != streamData { + t.Fatalf("stream passthrough mismatch") + } + if err := stream.Close(); err != nil { + t.Fatalf("Close error: %v", err) + } + if observer.eventCount != 1 { + t.Fatalf("eventCount = %d, want 1", observer.eventCount) + } + if observer.lastID != "chatcmpl-frag" { + t.Fatalf("lastID = %q, want chatcmpl-frag", observer.lastID) + } + if !observer.closed { + t.Fatal("observer was not closed") + } +} diff --git a/internal/usage/stream_observer.go b/internal/usage/stream_observer.go new file mode 100644 index 000000000..7b813fb35 --- /dev/null +++ b/internal/usage/stream_observer.go @@ -0,0 +1,140 @@ +package usage + +import "gomodel/internal/core" + +// StreamUsageObserver extracts usage data from parsed SSE JSON payloads. +type StreamUsageObserver struct { + logger LoggerInterface + pricingResolver PricingResolver + cachedEntry *UsageEntry + model string + provider string + requestID string + endpoint string + closed bool +} + +func NewStreamUsageObserver(logger LoggerInterface, model, provider, requestID, endpoint string, pricingResolver PricingResolver) *StreamUsageObserver { + if logger == nil { + return nil + } + return &StreamUsageObserver{ + logger: logger, + pricingResolver: pricingResolver, + model: model, + provider: provider, + requestID: requestID, + endpoint: endpoint, + } +} + +func (o *StreamUsageObserver) OnJSONEvent(chunk map[string]interface{}) { + entry := o.extractUsageFromEvent(chunk) + if entry != nil { + o.cachedEntry = entry + } +} + +func (o *StreamUsageObserver) OnStreamClose() { + if o.closed { + return + } + o.closed = true + if o.cachedEntry != nil && o.logger != nil { + o.logger.Write(o.cachedEntry) + } +} + +func (o *StreamUsageObserver) extractUsageFromEvent(chunk map[string]interface{}) *UsageEntry { + providerID, _ := chunk["id"].(string) + + model := o.model + if m, ok := chunk["model"].(string); ok && m != "" { + model = m + } + + usageRaw, ok := chunk["usage"] + if !ok { + if eventType, _ := chunk["type"].(string); eventType == "response.completed" || eventType == "response.done" { + if response, respOK := chunk["response"].(map[string]interface{}); respOK { + usageRaw, ok = response["usage"] + if id, idOK := response["id"].(string); idOK && id != "" { + providerID = id + } + if m, modelOK := response["model"].(string); modelOK && m != "" { + model = m + } + } + } + } + if !ok { + return nil + } + + usageMap, ok := usageRaw.(map[string]interface{}) + if !ok { + return nil + } + + var inputTokens, outputTokens, totalTokens int + rawData := make(map[string]any) + + if v, ok := usageMap["prompt_tokens"].(float64); ok { + inputTokens = int(v) + } + if v, ok := usageMap["input_tokens"].(float64); ok { + inputTokens = int(v) + } + if v, ok := usageMap["completion_tokens"].(float64); ok { + outputTokens = int(v) + } + if v, ok := usageMap["output_tokens"].(float64); ok { + outputTokens = int(v) + } + if v, ok := usageMap["total_tokens"].(float64); ok { + totalTokens = int(v) + } + + for field := range extendedFieldSet { + if v, ok := usageMap[field].(float64); ok && v > 0 { + rawData[field] = int(v) + } + } + + if details, ok := usageMap["prompt_tokens_details"].(map[string]interface{}); ok { + for k, v := range details { + if fv, ok := v.(float64); ok && fv > 0 { + rawData["prompt_"+k] = int(fv) + } + } + } + if details, ok := usageMap["completion_tokens_details"].(map[string]interface{}); ok { + for k, v := range details { + if fv, ok := v.(float64); ok && fv > 0 { + rawData["completion_"+k] = int(fv) + } + } + } + + if inputTokens == 0 && outputTokens == 0 && totalTokens == 0 { + return nil + } + if len(rawData) == 0 { + rawData = nil + } + + var pricingArgs []*core.ModelPricing + if o.pricingResolver != nil { + if p := o.pricingResolver.ResolvePricing(model, o.provider); p != nil { + pricingArgs = append(pricingArgs, p) + } + } + + return ExtractFromSSEUsage( + providerID, + inputTokens, outputTokens, totalTokens, + rawData, + o.requestID, model, o.provider, o.endpoint, + pricingArgs..., + ) +} diff --git a/internal/usage/stream_wrapper.go b/internal/usage/stream_wrapper.go index b783b4044..2d4dfeff3 100644 --- a/internal/usage/stream_wrapper.go +++ b/internal/usage/stream_wrapper.go @@ -1,17 +1,11 @@ package usage import ( - "bytes" - "encoding/json" "io" - "gomodel/internal/core" + "gomodel/internal/streaming" ) -// maxEventBufferRemainder is a safety valve for the event buffer remainder. -// If an incomplete event exceeds this size, it's discarded to prevent unbounded memory growth. -const maxEventBufferRemainder = 256 * 1024 // 256KB - // StreamUsageWrapper wraps an io.ReadCloser to capture usage data from SSE streams. // It incrementally parses SSE events as they arrive (on each \n\n boundary), // extracting and caching usage data immediately when found. This handles @@ -19,278 +13,18 @@ const maxEventBufferRemainder = 256 * 1024 // 256KB // the full response object alongside usage data. type StreamUsageWrapper struct { io.ReadCloser - logger LoggerInterface - pricingResolver PricingResolver - eventBuffer bytes.Buffer // accumulates raw bytes until \n\n found - cachedEntry *UsageEntry // stores extracted usage from latest event containing it - model string - provider string - requestID string - endpoint string - closed bool } // NewStreamUsageWrapper creates a wrapper around a stream to capture usage data. // When the stream is closed, it logs the cached usage entry if one was found. func NewStreamUsageWrapper(stream io.ReadCloser, logger LoggerInterface, model, provider, requestID, endpoint string, pricingResolver PricingResolver) *StreamUsageWrapper { - return &StreamUsageWrapper{ - ReadCloser: stream, - logger: logger, - pricingResolver: pricingResolver, - model: model, - provider: provider, - requestID: requestID, - endpoint: endpoint, - } -} - -// Read implements io.Reader. It reads from the underlying stream and incrementally -// parses complete SSE events to extract usage data as they arrive. -func (w *StreamUsageWrapper) Read(p []byte) (n int, err error) { - n, err = w.ReadCloser.Read(p) - if n > 0 { - w.eventBuffer.Write(p[:n]) - w.processCompleteEvents() - } - return n, err -} - -// processCompleteEvents scans the event buffer for complete SSE events (delimited by \n\n), -// extracts usage from each, and keeps only the unprocessed remainder. -func (w *StreamUsageWrapper) processCompleteEvents() { - data := w.eventBuffer.Bytes() - - // Find the last complete event boundary - lastBoundary := bytes.LastIndex(data, []byte("\n\n")) - if lastBoundary < 0 { - // No complete event yet — apply safety valve on remainder size. - // Trim to keep only the newest bytes so partial events are preserved. - if w.eventBuffer.Len() > maxEventBufferRemainder { - tail := w.eventBuffer.Bytes() - start := len(tail) - maxEventBufferRemainder - // If the underlying capacity has grown too large, allocate a new buffer, - // otherwise reuse the existing backing array to avoid unnecessary allocs. - if cap(tail) > maxEventBufferRemainder*2 { - var newBuf bytes.Buffer - newBuf.Write(tail[start:]) - w.eventBuffer = newBuf - } else { - copy(tail[:maxEventBufferRemainder], tail[start:]) - w.eventBuffer.Reset() - w.eventBuffer.Write(tail[:maxEventBufferRemainder]) - } - } - return - } - - // Split into complete events and remainder - completeData := data[:lastBoundary] - remainder := data[lastBoundary+2:] // skip the \n\n - - // Process each complete event - events := bytes.Split(completeData, []byte("\n\n")) - for _, event := range events { - if len(event) == 0 || bytes.Contains(event, []byte("[DONE]")) { - continue - } - - // Find data line(s) in this event - lines := bytes.Split(event, []byte("\n")) - for _, line := range lines { - trimmed := line - // Skip "event:" lines - if bytes.HasPrefix(trimmed, []byte("event:")) { - continue - } - if bytes.HasPrefix(trimmed, []byte("data: ")) { - jsonData := bytes.TrimPrefix(trimmed, []byte("data: ")) - entry := w.extractUsageFromJSON(jsonData) - if entry != nil { - w.cachedEntry = entry - } - } - } - } - - // Keep only the remainder, preventing capacity leaks from oversized buffers. - // Write copies remainder into a fresh buffer, releasing the old oversized backing array. - if len(remainder) > maxEventBufferRemainder { - remainder = remainder[len(remainder)-maxEventBufferRemainder:] - } - if w.eventBuffer.Cap() > maxEventBufferRemainder*2 { - var newBuf bytes.Buffer - newBuf.Write(remainder) - w.eventBuffer = newBuf - } else { - w.eventBuffer.Reset() - w.eventBuffer.Write(remainder) - } -} - -// Close implements io.Closer. It processes any remaining buffer data, -// logs the cached usage entry if found, and closes the underlying stream. -func (w *StreamUsageWrapper) Close() error { - if w.closed { - return nil - } - w.closed = true - - // Process any remaining data in the buffer as a final attempt - if w.eventBuffer.Len() > 0 { - entry := w.parseRemainingBuffer(w.eventBuffer.Bytes()) - if entry != nil { - w.cachedEntry = entry - } - } - - // Log the cached entry - if w.cachedEntry != nil && w.logger != nil { - w.logger.Write(w.cachedEntry) - } - - return w.ReadCloser.Close() -} - -// parseRemainingBuffer is a fallback parser for any unterminated data left in the buffer -// at Close time. It splits on \n\n and searches for usage data. -func (w *StreamUsageWrapper) parseRemainingBuffer(data []byte) *UsageEntry { - events := bytes.Split(data, []byte("\n\n")) - - for i := len(events) - 1; i >= 0; i-- { - event := events[i] - if len(event) == 0 || bytes.Contains(event, []byte("[DONE]")) { - continue - } - - lines := bytes.Split(event, []byte("\n")) - for _, line := range lines { - if bytes.HasPrefix(line, []byte("data: ")) { - jsonData := bytes.TrimPrefix(line, []byte("data: ")) - entry := w.extractUsageFromJSON(jsonData) - if entry != nil { - return entry - } - } - } - } - - return nil -} - -// extractUsageFromJSON attempts to extract usage from a JSON chunk. -func (w *StreamUsageWrapper) extractUsageFromJSON(data []byte) *UsageEntry { - // Try to parse as a generic map - var chunk map[string]interface{} - if err := json.Unmarshal(data, &chunk); err != nil { - return nil - } - - // Get provider ID (response ID) - providerID, _ := chunk["id"].(string) - - // Get model if available in the chunk (may differ from request model) - model := w.model - if m, ok := chunk["model"].(string); ok && m != "" { - model = m - } - - // Look for usage field (OpenAI/ChatCompletion format) - usageRaw, ok := chunk["usage"] - - // If not found at top level, check for Responses API format: - // {"type": "response.completed", "response": {"id": "...", "usage": {...}}} - if !ok { - if eventType, _ := chunk["type"].(string); eventType == "response.completed" || eventType == "response.done" { - if response, respOk := chunk["response"].(map[string]interface{}); respOk { - usageRaw, ok = response["usage"] - // Extract provider ID and model from response object - if id, idOk := response["id"].(string); idOk && id != "" { - providerID = id - } - if m, mOk := response["model"].(string); mOk && m != "" { - model = m - } - } - } - } - - if !ok { - return nil - } - - usageMap, ok := usageRaw.(map[string]interface{}) - if !ok { - return nil - } - - var inputTokens, outputTokens, totalTokens int - rawData := make(map[string]any) - - // Extract standard fields - if v, ok := usageMap["prompt_tokens"].(float64); ok { - inputTokens = int(v) - } - if v, ok := usageMap["input_tokens"].(float64); ok { - inputTokens = int(v) - } - if v, ok := usageMap["completion_tokens"].(float64); ok { - outputTokens = int(v) - } - if v, ok := usageMap["output_tokens"].(float64); ok { - outputTokens = int(v) + observer := NewStreamUsageObserver(logger, model, provider, requestID, endpoint, pricingResolver) + if observer == nil { + return &StreamUsageWrapper{ReadCloser: stream} } - if v, ok := usageMap["total_tokens"].(float64); ok { - totalTokens = int(v) - } - - // Extract extended usage data (provider-specific) using the field set - // derived from providerMappings in cost.go (single source of truth). - for field := range extendedFieldSet { - if v, ok := usageMap[field].(float64); ok && v > 0 { - rawData[field] = int(v) - } - } - - // Also check for nested prompt_tokens_details and completion_tokens_details (OpenAI) - if details, ok := usageMap["prompt_tokens_details"].(map[string]interface{}); ok { - for k, v := range details { - if fv, ok := v.(float64); ok && fv > 0 { - rawData["prompt_"+k] = int(fv) - } - } - } - if details, ok := usageMap["completion_tokens_details"].(map[string]interface{}); ok { - for k, v := range details { - if fv, ok := v.(float64); ok && fv > 0 { - rawData["completion_"+k] = int(fv) - } - } - } - - // Only create entry if we found some usage data - if inputTokens > 0 || outputTokens > 0 || totalTokens > 0 { - if len(rawData) == 0 { - rawData = nil - } - - // Resolve pricing for cost calculation - var pricingArgs []*core.ModelPricing - if w.pricingResolver != nil { - if p := w.pricingResolver.ResolvePricing(model, w.provider); p != nil { - pricingArgs = append(pricingArgs, p) - } - } - - return ExtractFromSSEUsage( - providerID, - inputTokens, outputTokens, totalTokens, - rawData, - w.requestID, model, w.provider, w.endpoint, - pricingArgs..., - ) + return &StreamUsageWrapper{ + ReadCloser: streaming.NewObservedSSEStream(stream, observer), } - - return nil } // WrapStreamForUsage wraps a stream with usage tracking if enabled. diff --git a/tests/perf/hotpath_test.go b/tests/perf/hotpath_test.go index 44b592939..d59bd8984 100644 --- a/tests/perf/hotpath_test.go +++ b/tests/perf/hotpath_test.go @@ -16,6 +16,7 @@ import ( "gomodel/internal/core" "gomodel/internal/providers" "gomodel/internal/server" + "gomodel/internal/streaming" "gomodel/internal/usage" ) @@ -207,7 +208,7 @@ func BenchmarkOpenAIResponsesStreamConverter(b *testing.B) { } } -func BenchmarkStreamingAuditAndUsageWrappers(b *testing.B) { +func BenchmarkSharedStreamingAuditAndUsageObservers(b *testing.B) { auditLogger := benchAuditLogger{cfg: auditlog.Config{Enabled: true, LogBodies: true}} usageLogger := benchUsageLogger{cfg: usage.Config{Enabled: true}} @@ -224,20 +225,17 @@ func BenchmarkStreamingAuditAndUsageWrappers(b *testing.B) { Data: &auditlog.LogData{}, } - stream := auditlog.WrapStreamForLogging( + stream := streaming.NewObservedSSEStream( io.NopCloser(strings.NewReader(sampleChatStream)), - auditLogger, - entry, - "/v1/chat/completions", - ) - stream = usage.WrapStreamForUsage( - stream, - usageLogger, - "gpt-4o-mini", - "mock", - "req-bench", - "/v1/chat/completions", - nil, + auditlog.NewStreamLogObserver(auditLogger, entry, "/v1/chat/completions"), + usage.NewStreamUsageObserver( + usageLogger, + "gpt-4o-mini", + "mock", + "req-bench", + "/v1/chat/completions", + nil, + ), ) if _, err := io.Copy(io.Discard, stream); err != nil { @@ -274,10 +272,10 @@ func TestHotPathPerfGuard(t *testing.T) { maxBytes: 32 * 1024, }, { - name: "stacked_stream_audit_and_usage_wrappers", - bench: BenchmarkStreamingAuditAndUsageWrappers, - maxAllocs: 340, - maxBytes: 20 * 1024, + name: "shared_stream_audit_and_usage_observers", + bench: BenchmarkSharedStreamingAuditAndUsageObservers, + maxAllocs: 240, + maxBytes: 14 * 1024, }, } From 49ef8528c0bfaa19ccb40ad9c139b8a33428a86b Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 18 Mar 2026 20:35:06 +0100 Subject: [PATCH 02/11] Remove dead usage stream wrapper --- internal/usage/constants.go | 3 +- ...rapper_test.go => stream_observer_test.go} | 232 ++++++------------ internal/usage/stream_wrapper.go | 37 --- 3 files changed, 72 insertions(+), 200 deletions(-) rename internal/usage/{stream_wrapper_test.go => stream_observer_test.go} (56%) delete mode 100644 internal/usage/stream_wrapper.go diff --git a/internal/usage/constants.go b/internal/usage/constants.go index a8bb5acfb..acb502aa1 100644 --- a/internal/usage/constants.go +++ b/internal/usage/constants.go @@ -15,6 +15,7 @@ const ( UsageEntryKey contextKey = "usage_entry" // UsageEntryStreamingKey is the context key for marking a request as streaming. - // When true, the middleware skips logging (StreamUsageWrapper handles it instead). + // When true, the middleware skips logging because streaming usage is handled + // by the shared SSE observer path. UsageEntryStreamingKey contextKey = "usage_entry_streaming" ) diff --git a/internal/usage/stream_wrapper_test.go b/internal/usage/stream_observer_test.go similarity index 56% rename from internal/usage/stream_wrapper_test.go rename to internal/usage/stream_observer_test.go index ec957ba46..40d00f0d7 100644 --- a/internal/usage/stream_wrapper_test.go +++ b/internal/usage/stream_observer_test.go @@ -6,10 +6,10 @@ import ( "sync" "testing" - "gomodel/internal/core" + "gomodel/internal/streaming" ) -// trackingLogger tracks written entries for testing +// trackingLogger tracks written entries for testing. type trackingLogger struct { entries []*UsageEntry mu sync.Mutex @@ -38,8 +38,7 @@ func (l *trackingLogger) getEntries() []*UsageEntry { return result } -func TestStreamUsageWrapper(t *testing.T) { - // OpenAI-style SSE stream with usage in final event +func TestStreamUsageObserverChatCompletionStream(t *testing.T) { streamData := `data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"content":"!"},"finish_reason":"stop"}],"usage":{"prompt_tokens":10,"completion_tokens":5,"total_tokens":15}} @@ -48,31 +47,26 @@ data: [DONE] ` logger := &trackingLogger{enabled: true} - stream := io.NopCloser(strings.NewReader(streamData)) - wrapper := NewStreamUsageWrapper(stream, logger, "gpt-4", "openai", "req-123", "/v1/chat/completions", nil) + stream := streaming.NewObservedSSEStream( + io.NopCloser(strings.NewReader(streamData)), + NewStreamUsageObserver(logger, "gpt-4", "openai", "req-123", "/v1/chat/completions", nil), + ) - // Read all data - data, err := io.ReadAll(wrapper) + data, err := io.ReadAll(stream) if err != nil { t.Fatalf("ReadAll error: %v", err) } - - // Verify data passed through if string(data) != streamData { - t.Errorf("data mismatch: got %d bytes, want %d bytes", len(data), len(streamData)) + t.Fatalf("stream passthrough mismatch") } - - // Close wrapper to trigger usage extraction - if err := wrapper.Close(); err != nil { + if err := stream.Close(); err != nil { t.Fatalf("Close error: %v", err) } - // Verify usage was extracted entries := logger.getEntries() if len(entries) != 1 { t.Fatalf("expected 1 entry, got %d", len(entries)) } - entry := entries[0] if entry.InputTokens != 10 { t.Errorf("InputTokens = %d, want 10", entry.InputTokens) @@ -91,25 +85,30 @@ data: [DONE] } } -func TestStreamUsageWrapperWithExtendedUsage(t *testing.T) { - // OpenAI o-series with prompt_tokens_details and completion_tokens_details - streamData := `data: {"id":"chatcmpl-456","object":"chat.completion.chunk","model":"o1-preview","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":"stop"}],"usage":{"prompt_tokens":100,"completion_tokens":50,"total_tokens":150,"prompt_tokens_details":{"cached_tokens":20},"completion_tokens_details":{"reasoning_tokens":10}}} - -data: [DONE] - -` +func TestStreamUsageObserverWithExtendedUsage(t *testing.T) { logger := &trackingLogger{enabled: true} - stream := io.NopCloser(strings.NewReader(streamData)) - wrapper := NewStreamUsageWrapper(stream, logger, "o1-preview", "openai", "req-456", "/v1/chat/completions", nil) - - _, _ = io.ReadAll(wrapper) - _ = wrapper.Close() + observer := NewStreamUsageObserver(logger, "o1-preview", "openai", "req-456", "/v1/chat/completions", nil) + observer.OnJSONEvent(map[string]interface{}{ + "id": "chatcmpl-456", + "model": "o1-preview", + "usage": map[string]interface{}{ + "prompt_tokens": float64(100), + "completion_tokens": float64(50), + "total_tokens": float64(150), + "prompt_tokens_details": map[string]interface{}{ + "cached_tokens": float64(20), + }, + "completion_tokens_details": map[string]interface{}{ + "reasoning_tokens": float64(10), + }, + }, + }) + observer.OnStreamClose() entries := logger.getEntries() if len(entries) != 1 { t.Fatalf("expected 1 entry, got %d", len(entries)) } - entry := entries[0] if entry.InputTokens != 100 { t.Errorf("InputTokens = %d, want 100", entry.InputTokens) @@ -117,8 +116,6 @@ data: [DONE] if entry.OutputTokens != 50 { t.Errorf("OutputTokens = %d, want 50", entry.OutputTokens) } - - // Check extended data was captured if entry.RawData == nil { t.Fatal("expected RawData to be set") } @@ -130,118 +127,41 @@ data: [DONE] } } -func TestStreamUsageWrapperNoUsage(t *testing.T) { - // Stream without usage data +func TestStreamUsageObserverNoUsage(t *testing.T) { streamData := `data: {"id":"chatcmpl-789","object":"chat.completion.chunk","model":"gpt-4","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":"stop"}]} data: [DONE] ` logger := &trackingLogger{enabled: true} - stream := io.NopCloser(strings.NewReader(streamData)) - wrapper := NewStreamUsageWrapper(stream, logger, "gpt-4", "openai", "req-789", "/v1/chat/completions", nil) + stream := streaming.NewObservedSSEStream( + io.NopCloser(strings.NewReader(streamData)), + NewStreamUsageObserver(logger, "gpt-4", "openai", "req-789", "/v1/chat/completions", nil), + ) - _, _ = io.ReadAll(wrapper) - _ = wrapper.Close() + _, _ = io.ReadAll(stream) + _ = stream.Close() - // Should not log anything if no usage found entries := logger.getEntries() if len(entries) != 0 { t.Errorf("expected 0 entries (no usage), got %d", len(entries)) } } -func TestStreamUsageWrapperDisabled(t *testing.T) { - streamData := `data: {"id":"chatcmpl-123","usage":{"prompt_tokens":10,"completion_tokens":5,"total_tokens":15}} - -data: [DONE] - -` - logger := &trackingLogger{enabled: false} // disabled - stream := io.NopCloser(strings.NewReader(streamData)) - wrapper := NewStreamUsageWrapper(stream, logger, "gpt-4", "openai", "req-123", "/v1/chat/completions", nil) - - _, _ = io.ReadAll(wrapper) - _ = wrapper.Close() - - // Should still log even when config says disabled (because Write() is called) - // The WrapStreamForUsage function is what should check enabled status - entries := logger.getEntries() - if len(entries) != 1 { - t.Errorf("expected 1 entry, got %d", len(entries)) - } -} - -func TestWrapStreamForUsageDisabled(t *testing.T) { - streamData := "test data" - logger := &trackingLogger{enabled: false} // disabled - stream := io.NopCloser(strings.NewReader(streamData)) - - wrapped := WrapStreamForUsage(stream, logger, "gpt-4", "openai", "req-123", "/v1/chat/completions", nil) - - // When disabled, should return original stream (not wrapped) - // This is determined by checking if wrapped is the same as original - data, _ := io.ReadAll(wrapped) - if string(data) != streamData { - t.Errorf("data mismatch") - } -} - -func TestWrapStreamForUsageNilLogger(t *testing.T) { - streamData := "test data" - stream := io.NopCloser(strings.NewReader(streamData)) - - wrapped := WrapStreamForUsage(stream, nil, "gpt-4", "openai", "req-123", "/v1/chat/completions", nil) - - // When nil logger, should return original stream - data, _ := io.ReadAll(wrapped) - if string(data) != streamData { - t.Errorf("data mismatch") - } -} - -func TestIsModelInteractionPath(t *testing.T) { - tests := []struct { - path string - want bool - }{ - {"/v1/chat/completions", true}, - {"/v1/chat/completions?foo=bar", true}, - {"/v1/responses", true}, - {"/v1/responses/123", true}, - {"/v1/models", false}, - {"/health", false}, - {"/metrics", false}, - {"/admin", false}, - {"/", false}, - {"", false}, - } - - for _, tt := range tests { - t.Run(tt.path, func(t *testing.T) { - got := core.IsModelInteractionPath(tt.path) - if got != tt.want { - t.Errorf("core.IsModelInteractionPath(%q) = %v, want %v", tt.path, got, tt.want) - } - }) - } -} - -func TestStreamUsageWrapperDoubleClose(t *testing.T) { - streamData := `data: {"id":"chatcmpl-123","usage":{"prompt_tokens":10,"completion_tokens":5,"total_tokens":15}} - -data: [DONE] - -` +func TestStreamUsageObserverDoubleClose(t *testing.T) { logger := &trackingLogger{enabled: true} - stream := io.NopCloser(strings.NewReader(streamData)) - wrapper := NewStreamUsageWrapper(stream, logger, "gpt-4", "openai", "req-123", "/v1/chat/completions", nil) - - _, _ = io.ReadAll(wrapper) - - // Close twice should not panic or double-log - _ = wrapper.Close() - _ = wrapper.Close() + observer := NewStreamUsageObserver(logger, "gpt-4", "openai", "req-123", "/v1/chat/completions", nil) + observer.OnJSONEvent(map[string]interface{}{ + "id": "chatcmpl-123", + "usage": map[string]interface{}{ + "prompt_tokens": float64(10), + "completion_tokens": float64(5), + "total_tokens": float64(15), + }, + }) + + observer.OnStreamClose() + observer.OnStreamClose() entries := logger.getEntries() if len(entries) != 1 { @@ -249,8 +169,7 @@ data: [DONE] } } -func TestStreamUsageWrapperResponsesAPI(t *testing.T) { - // Responses API format with event: prefixes and response.completed containing nested response.usage +func TestStreamUsageObserverResponsesAPI(t *testing.T) { streamData := `event: response.created data: {"type":"response.created","response":{"id":"resp-123","object":"response","status":"in_progress","model":"gpt-5"}} @@ -267,19 +186,19 @@ data: [DONE] ` logger := &trackingLogger{enabled: true} - stream := io.NopCloser(strings.NewReader(streamData)) - wrapper := NewStreamUsageWrapper(stream, logger, "gpt-5", "openai", "req-resp-1", "/v1/responses", nil) + stream := streaming.NewObservedSSEStream( + io.NopCloser(strings.NewReader(streamData)), + NewStreamUsageObserver(logger, "gpt-5", "openai", "req-resp-1", "/v1/responses", nil), + ) - data, err := io.ReadAll(wrapper) + data, err := io.ReadAll(stream) if err != nil { t.Fatalf("ReadAll error: %v", err) } - if string(data) != streamData { t.Errorf("data mismatch: got %d bytes, want %d bytes", len(data), len(streamData)) } - - if err := wrapper.Close(); err != nil { + if err := stream.Close(); err != nil { t.Fatalf("Close error: %v", err) } @@ -287,7 +206,6 @@ data: [DONE] if len(entries) != 1 { t.Fatalf("expected 1 entry, got %d", len(entries)) } - entry := entries[0] if entry.InputTokens != 15 { t.Errorf("InputTokens = %d, want 15", entry.InputTokens) @@ -306,13 +224,8 @@ data: [DONE] } } -func TestStreamUsageWrapperLargeResponsesDone(t *testing.T) { - // Regression test: response.completed event >8KB should not lose usage data. - // The old rolling 8KB buffer would truncate the beginning of this event. - - // Build a large output content to push the response.completed event well over 8KB - largeText := strings.Repeat("This is a long response from the model. ", 300) // ~12KB of text - +func TestStreamUsageObserverLargeResponsesDone(t *testing.T) { + largeText := strings.Repeat("This is a long response from the model. ", 300) streamData := `event: response.created data: {"type":"response.created","response":{"id":"resp-large","object":"response","status":"in_progress","model":"gpt-5"}} @@ -325,8 +238,6 @@ data: {"type":"response.completed","response":{"id":"resp-large","object":"respo data: [DONE] ` - - // Verify the response.completed event is actually >8KB doneEventStart := strings.Index(streamData, `data: {"type":"response.completed"`) doneEventEnd := strings.Index(streamData[doneEventStart:], "\n\n") doneEventSize := doneEventEnd @@ -335,19 +246,19 @@ data: [DONE] } logger := &trackingLogger{enabled: true} - stream := io.NopCloser(strings.NewReader(streamData)) - wrapper := NewStreamUsageWrapper(stream, logger, "gpt-5", "openai", "req-large", "/v1/responses", nil) + stream := streaming.NewObservedSSEStream( + io.NopCloser(strings.NewReader(streamData)), + NewStreamUsageObserver(logger, "gpt-5", "openai", "req-large", "/v1/responses", nil), + ) - data, err := io.ReadAll(wrapper) + data, err := io.ReadAll(stream) if err != nil { t.Fatalf("ReadAll error: %v", err) } - if string(data) != streamData { t.Errorf("data mismatch: got %d bytes, want %d bytes", len(data), len(streamData)) } - - if err := wrapper.Close(); err != nil { + if err := stream.Close(); err != nil { t.Fatalf("Close error: %v", err) } @@ -355,7 +266,6 @@ data: [DONE] if len(entries) != 1 { t.Fatalf("expected 1 entry, got %d (usage was lost from large response.completed event)", len(entries)) } - entry := entries[0] if entry.InputTokens != 100 { t.Errorf("InputTokens = %d, want 100", entry.InputTokens) @@ -371,22 +281,22 @@ data: [DONE] } } -func TestStreamUsageWrapperSmallReads(t *testing.T) { - // Fragmented reads (7-byte chunks) to verify cross-boundary event detection +func TestStreamUsageObserverSmallReads(t *testing.T) { streamData := `data: {"id":"chatcmpl-frag","object":"chat.completion.chunk","model":"gpt-4","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":"stop"}],"usage":{"prompt_tokens":5,"completion_tokens":3,"total_tokens":8}} data: [DONE] ` logger := &trackingLogger{enabled: true} - stream := io.NopCloser(strings.NewReader(streamData)) - wrapper := NewStreamUsageWrapper(stream, logger, "gpt-4", "openai", "req-frag", "/v1/chat/completions", nil) + stream := streaming.NewObservedSSEStream( + io.NopCloser(strings.NewReader(streamData)), + NewStreamUsageObserver(logger, "gpt-4", "openai", "req-frag", "/v1/chat/completions", nil), + ) - // Read in small chunks of 7 bytes buf := make([]byte, 7) var allData []byte for { - n, err := wrapper.Read(buf) + n, err := stream.Read(buf) if n > 0 { allData = append(allData, buf[:n]...) } @@ -401,8 +311,7 @@ data: [DONE] if string(allData) != streamData { t.Errorf("data mismatch: got %d bytes, want %d bytes", len(allData), len(streamData)) } - - if err := wrapper.Close(); err != nil { + if err := stream.Close(); err != nil { t.Fatalf("Close error: %v", err) } @@ -410,7 +319,6 @@ data: [DONE] if len(entries) != 1 { t.Fatalf("expected 1 entry, got %d", len(entries)) } - entry := entries[0] if entry.InputTokens != 5 { t.Errorf("InputTokens = %d, want 5", entry.InputTokens) diff --git a/internal/usage/stream_wrapper.go b/internal/usage/stream_wrapper.go deleted file mode 100644 index 2d4dfeff3..000000000 --- a/internal/usage/stream_wrapper.go +++ /dev/null @@ -1,37 +0,0 @@ -package usage - -import ( - "io" - - "gomodel/internal/streaming" -) - -// StreamUsageWrapper wraps an io.ReadCloser to capture usage data from SSE streams. -// It incrementally parses SSE events as they arrive (on each \n\n boundary), -// extracting and caching usage data immediately when found. This handles -// arbitrarily large events like the Responses API's response.completed which includes -// the full response object alongside usage data. -type StreamUsageWrapper struct { - io.ReadCloser -} - -// NewStreamUsageWrapper creates a wrapper around a stream to capture usage data. -// When the stream is closed, it logs the cached usage entry if one was found. -func NewStreamUsageWrapper(stream io.ReadCloser, logger LoggerInterface, model, provider, requestID, endpoint string, pricingResolver PricingResolver) *StreamUsageWrapper { - observer := NewStreamUsageObserver(logger, model, provider, requestID, endpoint, pricingResolver) - if observer == nil { - return &StreamUsageWrapper{ReadCloser: stream} - } - return &StreamUsageWrapper{ - ReadCloser: streaming.NewObservedSSEStream(stream, observer), - } -} - -// WrapStreamForUsage wraps a stream with usage tracking if enabled. -// This is a convenience function for use in handlers. -func WrapStreamForUsage(stream io.ReadCloser, logger LoggerInterface, model, provider, requestID, endpoint string, pricingResolver PricingResolver) io.ReadCloser { - if logger == nil || !logger.Config().Enabled { - return stream - } - return NewStreamUsageWrapper(stream, logger, model, provider, requestID, endpoint, pricingResolver) -} From 1507b6644c80cfa1e258e09cfa1b82c4cc8b3d61 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 18 Mar 2026 20:54:53 +0100 Subject: [PATCH 03/11] Remove dead audit stream wrapper --- internal/auditlog/auditlog_test.go | 35 ++++++++++-------------- internal/auditlog/constants.go | 5 ++-- internal/auditlog/middleware.go | 4 +-- internal/auditlog/stream_wrapper.go | 38 ++------------------------ internal/server/passthrough_support.go | 7 ++++- 5 files changed, 29 insertions(+), 60 deletions(-) diff --git a/internal/auditlog/auditlog_test.go b/internal/auditlog/auditlog_test.go index 8c32a680b..8720e18fc 100644 --- a/internal/auditlog/auditlog_test.go +++ b/internal/auditlog/auditlog_test.go @@ -8,6 +8,7 @@ import ( "encoding/json" "fmt" "gomodel/internal/core" + "gomodel/internal/streaming" "io" "net/http" "net/http/httptest" @@ -705,7 +706,7 @@ func TestIsModelInteractionPath(t *testing.T) { } } -func TestStreamLogWrapper(t *testing.T) { +func TestStreamLogObserver(t *testing.T) { // Create a mock stream with content streamContent := `data: {"id":"chatcmpl-123","choices":[{"delta":{"content":"Hello"}}]} @@ -733,19 +734,21 @@ data: [DONE] Data: &LogData{}, } - // Wrap the stream - wrapper := NewStreamLogWrapper(stream, logger, entry, "/v1/chat/completions") + observedStream := streaming.NewObservedSSEStream( + stream, + NewStreamLogObserver(logger, entry, "/v1/chat/completions"), + ) // Read all content var buf bytes.Buffer - _, err := io.Copy(&buf, wrapper) + _, err := io.Copy(&buf, observedStream) if err != nil { t.Fatalf("failed to read stream: %v", err) } - // Close wrapper to trigger logging - if err := wrapper.Close(); err != nil { - t.Fatalf("failed to close wrapper: %v", err) + // Close stream to trigger logging + if err := observedStream.Close(); err != nil { + t.Fatalf("failed to close stream: %v", err) } // Wait for async write @@ -757,20 +760,12 @@ data: [DONE] } } -func TestWrapStreamForLogging(t *testing.T) { - stream := io.NopCloser(strings.NewReader("test")) - - // Test with nil logger - result := WrapStreamForLogging(stream, nil, nil, "/v1/chat/completions") - if result != stream { - t.Error("expected original stream with nil logger") +func TestNewStreamLogObserverNilInputs(t *testing.T) { + if observer := NewStreamLogObserver(nil, &LogEntry{}, "/v1/chat/completions"); observer != nil { + t.Error("expected nil observer with nil logger") } - - // Test with disabled logger - noopLogger := &NoopLogger{} - result = WrapStreamForLogging(stream, noopLogger, &LogEntry{}, "/v1/chat/completions") - if result != stream { - t.Error("expected original stream with disabled logger") + if observer := NewStreamLogObserver(&NoopLogger{}, nil, "/v1/chat/completions"); observer != nil { + t.Error("expected nil observer with nil entry") } } diff --git a/internal/auditlog/constants.go b/internal/auditlog/constants.go index 1f8193813..f29965675 100644 --- a/internal/auditlog/constants.go +++ b/internal/auditlog/constants.go @@ -7,7 +7,7 @@ const ( MaxBodyCapture = 1024 * 1024 // MaxContentCapture is the maximum size of accumulated streaming content (1MB). - // Used by StreamLogWrapper to limit reconstructed response body size. + // Used by the stream observer to limit reconstructed response body size. MaxContentCapture = 1024 * 1024 // BatchFlushThreshold is the number of entries that triggers an immediate flush. @@ -27,6 +27,7 @@ const ( LogEntryKey contextKey = "auditlog_entry" // LogEntryStreamingKey is the context key for marking a request as streaming. - // When true, the middleware skips logging (StreamLogWrapper handles it instead). + // When true, the middleware skips logging because the stream observer path + // handles streaming audit logging. LogEntryStreamingKey contextKey = "auditlog_entry_streaming" ) diff --git a/internal/auditlog/middleware.go b/internal/auditlog/middleware.go index d72b4aa97..10066773b 100644 --- a/internal/auditlog/middleware.go +++ b/internal/auditlog/middleware.go @@ -149,7 +149,7 @@ func Middleware(logger LoggerInterface) echo.MiddlewareFunc { } } - // Write log entry asynchronously (skip if streaming - StreamLogWrapper handles it) + // Write log entry asynchronously (skip if streaming - the stream observer path handles it) if !IsEntryMarkedAsStreaming(c) { logger.Write(entry) } @@ -262,7 +262,7 @@ type responseBodyCapture struct { body *bytes.Buffer truncated bool // shouldCapture allows middleware to stop buffering once the request is - // known to be streaming. Streaming responses are handled by StreamLogWrapper. + // known to be streaming. Streaming responses are handled by the stream observer path. shouldCapture func() bool } diff --git a/internal/auditlog/stream_wrapper.go b/internal/auditlog/stream_wrapper.go index 233bc3bb1..76c155110 100644 --- a/internal/auditlog/stream_wrapper.go +++ b/internal/auditlog/stream_wrapper.go @@ -1,10 +1,7 @@ package auditlog import ( - "io" "strings" - - "gomodel/internal/streaming" ) // Note: MaxContentCapture and LogEntryStreamingKey constants are defined in constants.go @@ -30,25 +27,6 @@ type streamResponseBuilder struct { truncated bool } -// StreamLogWrapper wraps an io.ReadCloser to reconstruct streamed response bodies -// for audit logging. -type StreamLogWrapper struct { - io.ReadCloser -} - -// NewStreamLogWrapper creates a wrapper around a stream for audit logging. -// When the stream is closed, it logs the accumulated entry. -// The path parameter is used to detect whether this is a ChatCompletion or Responses API request. -func NewStreamLogWrapper(stream io.ReadCloser, logger LoggerInterface, entry *LogEntry, path string) *StreamLogWrapper { - observer := NewStreamLogObserver(logger, entry, path) - if observer == nil { - return &StreamLogWrapper{ReadCloser: stream} - } - return &StreamLogWrapper{ - ReadCloser: streaming.NewObservedSSEStream(stream, observer), - } -} - // buildChatCompletionResponse constructs a ChatCompletion response from accumulated data func (b *streamResponseBuilder) buildChatCompletionResponse() map[string]interface{} { return map[string]interface{}{ @@ -92,16 +70,6 @@ func (b *streamResponseBuilder) buildResponsesAPIResponse() map[string]interface } } -// WrapStreamForLogging wraps a stream with logging if enabled. -// This is a convenience function for use in handlers. -// The path parameter is used to detect whether this is a ChatCompletion or Responses API request. -func WrapStreamForLogging(stream io.ReadCloser, logger LoggerInterface, entry *LogEntry, path string) io.ReadCloser { - if logger == nil || !logger.Config().Enabled || entry == nil { - return stream - } - return NewStreamLogWrapper(stream, logger, entry, path) -} - // CreateStreamEntry creates a new log entry for a streaming request. // This should be called before starting the stream. func CreateStreamEntry(baseEntry *LogEntry) *LogEntry { @@ -109,8 +77,8 @@ func CreateStreamEntry(baseEntry *LogEntry) *LogEntry { return nil } - // Create a copy of the entry for the stream - // The stream wrapper will complete and write it when the stream closes + // Create a copy of the entry for the stream. + // The stream observer will complete and write it when the stream closes. entryCopy := &LogEntry{ ID: baseEntry.ID, Timestamp: baseEntry.Timestamp, @@ -172,7 +140,7 @@ func GetStreamEntryFromContext(c interface{ Get(string) interface{} }) *LogEntry } // MarkEntryAsStreaming marks the entry as a streaming request so the middleware -// knows not to log it (the stream wrapper will handle logging). +// knows not to log it (the stream observer path will handle logging). func MarkEntryAsStreaming(c interface{ Set(string, interface{}) }, isStreaming bool) { c.Set(string(LogEntryStreamingKey), isStreaming) } diff --git a/internal/server/passthrough_support.go b/internal/server/passthrough_support.go index e045dfd7c..8a1b5c0b3 100644 --- a/internal/server/passthrough_support.go +++ b/internal/server/passthrough_support.go @@ -12,6 +12,7 @@ import ( "gomodel/internal/auditlog" "gomodel/internal/core" + "gomodel/internal/streaming" ) var defaultEnabledPassthroughProviders = []string{"openai", "anthropic"} @@ -249,7 +250,11 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT streamEntry.StatusCode = resp.StatusCode } - wrappedStream := auditlog.WrapStreamForLogging(resp.Body, s.logger, streamEntry, passthroughAuditPath(c, providerType, endpoint, info)) + observers := make([]streaming.Observer, 0, 1) + if observer := auditlog.NewStreamLogObserver(s.logger, streamEntry, passthroughAuditPath(c, providerType, endpoint, info)); observer != nil { + observers = append(observers, observer) + } + wrappedStream := streaming.NewObservedSSEStream(resp.Body, observers...) defer func() { _ = wrappedStream.Close() }() From 7cf0f629cb79b5fc3beb6646d607fe8f2e420f6e Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 18 Mar 2026 20:57:43 +0100 Subject: [PATCH 04/11] Guard SSE pending buffer allocations --- internal/streaming/observed_sse_stream.go | 3 +++ internal/streaming/observed_sse_stream_test.go | 17 +++++++++++++++++ 2 files changed, 20 insertions(+) diff --git a/internal/streaming/observed_sse_stream.go b/internal/streaming/observed_sse_stream.go index 2ec534eda..28bf8a10d 100644 --- a/internal/streaming/observed_sse_stream.go +++ b/internal/streaming/observed_sse_stream.go @@ -68,6 +68,9 @@ func (s *ObservedSSEStream) Close() error { func (s *ObservedSSEStream) processChunk(data []byte) { if len(s.pending) > 0 { + if len(data) > maxPendingEventBytes { + data = data[len(data)-maxPendingEventBytes:] + } combined := make([]byte, len(s.pending)+len(data)) copy(combined, s.pending) copy(combined[len(s.pending):], data) diff --git a/internal/streaming/observed_sse_stream_test.go b/internal/streaming/observed_sse_stream_test.go index 8e24e8b4d..260362c48 100644 --- a/internal/streaming/observed_sse_stream_test.go +++ b/internal/streaming/observed_sse_stream_test.go @@ -1,6 +1,7 @@ package streaming import ( + "bytes" "io" "strings" "testing" @@ -84,3 +85,19 @@ func TestObservedSSEStream_ParsesFragmentedFinalEventOnClose(t *testing.T) { t.Fatal("observer was not closed") } } + +func TestObservedSSEStream_CapsCombinedPendingData(t *testing.T) { + s := &ObservedSSEStream{ + pending: bytes.Repeat([]byte("a"), maxPendingEventBytes), + } + data := bytes.Repeat([]byte("b"), maxPendingEventBytes+1024) + + s.processChunk(data) + + if got := len(s.pending); got != maxPendingEventBytes { + t.Fatalf("pending length = %d, want %d", got, maxPendingEventBytes) + } + if !bytes.Equal(s.pending, data[len(data)-maxPendingEventBytes:]) { + t.Fatal("pending bytes do not match the capped suffix of the latest chunk") + } +} From 7c1f71740c11afae1e9592c0fd81f19e611c62da Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 18 Mar 2026 21:05:42 +0100 Subject: [PATCH 05/11] Tighten streaming parser compatibility --- internal/server/handlers.go | 7 --- internal/server/handlers_test.go | 4 +- internal/streaming/observed_sse_stream.go | 59 ++++++++++++++++--- .../streaming/observed_sse_stream_test.go | 52 ++++++++++++++++ 4 files changed, 104 insertions(+), 18 deletions(-) diff --git a/internal/server/handlers.go b/internal/server/handlers.go index 12af26242..03cd22019 100644 --- a/internal/server/handlers.go +++ b/internal/server/handlers.go @@ -2,7 +2,6 @@ package server import ( - "io" "net/http" "github.com/labstack/echo/v5" @@ -63,12 +62,6 @@ func (h *Handler) SetBatchStore(store batchstore.Store) { h.batchStore = store } -// handleStreamingResponse handles SSE streaming responses for both ChatCompletion and Responses endpoints. -// It wraps the stream with audit logging and usage tracking, and sets appropriate SSE headers. -func (h *Handler) handleStreamingResponse(c *echo.Context, model, provider string, streamFn func() (io.ReadCloser, error)) error { - return h.translatedInference().handleStreamingResponse(c, model, provider, streamFn) -} - func (h *Handler) translatedInference() *translatedInferenceService { return &translatedInferenceService{ provider: h.provider, diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go index aa304fcb8..75177672a 100644 --- a/internal/server/handlers_test.go +++ b/internal/server/handlers_test.go @@ -2001,7 +2001,7 @@ func TestHandleStreamingResponse_FlushesEachChunk(t *testing.T) { }, } - err := handler.handleStreamingResponse(c, "gpt-4o-mini", "openai", func() (io.ReadCloser, error) { + err := handler.translatedInference().handleStreamingResponse(c, "gpt-4o-mini", "openai", func() (io.ReadCloser, error) { return stream, nil }) if err != nil { @@ -2093,7 +2093,7 @@ func TestHandleStreamingResponse_RecordsStreamingError(t *testing.T) { Data: &auditlog.LogData{}, }) - err := handler.handleStreamingResponse(c, "gpt-4o-mini", "openai", func() (io.ReadCloser, error) { + err := handler.translatedInference().handleStreamingResponse(c, "gpt-4o-mini", "openai", func() (io.ReadCloser, error) { return &erroringReadCloser{ data: []byte("data: {\"id\":\"1\"}\n\n"), err: expectedErr, diff --git a/internal/streaming/observed_sse_stream.go b/internal/streaming/observed_sse_stream.go index 28bf8a10d..d97ac1dfe 100644 --- a/internal/streaming/observed_sse_stream.go +++ b/internal/streaming/observed_sse_stream.go @@ -8,6 +8,13 @@ import ( const maxPendingEventBytes = 256 * 1024 +var ( + lfEventBoundary = []byte("\n\n") + crlfEventBoundary = []byte("\r\n\r\n") + dataPrefix = []byte("data:") + donePayload = []byte("[DONE]") +) + // Observer receives parsed JSON SSE payloads in stream order. // Implementations must treat the payload as read-only. type Observer interface { @@ -79,34 +86,39 @@ func (s *ObservedSSEStream) processChunk(data []byte) { } for { - idx := bytes.Index(data, []byte("\n\n")) + idx, sepLen := nextEventBoundary(data) if idx == -1 { s.savePending(data) return } s.processEvent(data[:idx]) - data = data[idx+2:] + data = data[idx+sepLen:] } } func (s *ObservedSSEStream) processBufferedEvents(data []byte) { - for _, event := range bytes.Split(data, []byte("\n\n")) { - if len(event) == 0 { - continue + for len(data) > 0 { + idx, sepLen := nextEventBoundary(data) + if idx == -1 { + s.processEvent(data) + return + } + if idx > 0 { + s.processEvent(data[:idx]) } - s.processEvent(event) + data = data[idx+sepLen:] } } func (s *ObservedSSEStream) processEvent(event []byte) { lines := bytes.Split(event, []byte("\n")) for _, line := range lines { - if !bytes.HasPrefix(line, []byte("data: ")) { + jsonData, ok := parseDataLine(line) + if !ok { continue } - jsonData := bytes.TrimPrefix(line, []byte("data: ")) - if bytes.Equal(jsonData, []byte("[DONE]")) { + if bytes.Equal(jsonData, donePayload) { continue } @@ -120,6 +132,35 @@ func (s *ObservedSSEStream) processEvent(event []byte) { } } +func nextEventBoundary(data []byte) (idx int, sepLen int) { + lfIdx := bytes.Index(data, lfEventBoundary) + crlfIdx := bytes.Index(data, crlfEventBoundary) + + switch { + case lfIdx == -1: + if crlfIdx == -1 { + return -1, 0 + } + return crlfIdx, len(crlfEventBoundary) + case crlfIdx == -1 || lfIdx < crlfIdx: + return lfIdx, len(lfEventBoundary) + default: + return crlfIdx, len(crlfEventBoundary) + } +} + +func parseDataLine(line []byte) ([]byte, bool) { + line = bytes.TrimSuffix(line, []byte("\r")) + if !bytes.HasPrefix(line, dataPrefix) { + return nil, false + } + payload := bytes.TrimPrefix(line, dataPrefix) + if len(payload) > 0 && payload[0] == ' ' { + payload = payload[1:] + } + return payload, true +} + func (s *ObservedSSEStream) savePending(data []byte) { if len(data) == 0 { return diff --git a/internal/streaming/observed_sse_stream_test.go b/internal/streaming/observed_sse_stream_test.go index 260362c48..b628b16d8 100644 --- a/internal/streaming/observed_sse_stream_test.go +++ b/internal/streaming/observed_sse_stream_test.go @@ -101,3 +101,55 @@ func TestObservedSSEStream_CapsCombinedPendingData(t *testing.T) { t.Fatal("pending bytes do not match the capped suffix of the latest chunk") } } + +func TestObservedSSEStream_HandlesCRLFAndDataWithoutSpace(t *testing.T) { + streamData := "data:{\"id\":\"chatcmpl-1\"}\r\n\r\ndata: {\"id\":\"chatcmpl-2\"}\r\n\r\ndata:[DONE]\r\n\r\n" + observer := &trackingObserver{} + stream := NewObservedSSEStream(io.NopCloser(strings.NewReader(streamData)), observer) + + data, err := io.ReadAll(stream) + if err != nil { + t.Fatalf("ReadAll error: %v", err) + } + if string(data) != streamData { + t.Fatalf("stream passthrough mismatch") + } + if err := stream.Close(); err != nil { + t.Fatalf("Close error: %v", err) + } + if observer.eventCount != 2 { + t.Fatalf("eventCount = %d, want 2", observer.eventCount) + } + if observer.lastID != "chatcmpl-2" { + t.Fatalf("lastID = %q, want chatcmpl-2", observer.lastID) + } + if !observer.closed { + t.Fatal("observer was not closed") + } +} + +func TestObservedSSEStream_ParsesCRLFBufferedEventsOnClose(t *testing.T) { + streamData := "data:{\"id\":\"chatcmpl-1\"}\r\n\r\ndata:{\"id\":\"chatcmpl-2\"}" + observer := &trackingObserver{} + stream := NewObservedSSEStream(io.NopCloser(strings.NewReader(streamData)), observer) + + data, err := io.ReadAll(stream) + if err != nil { + t.Fatalf("ReadAll error: %v", err) + } + if string(data) != streamData { + t.Fatalf("stream passthrough mismatch") + } + if err := stream.Close(); err != nil { + t.Fatalf("Close error: %v", err) + } + if observer.eventCount != 2 { + t.Fatalf("eventCount = %d, want 2", observer.eventCount) + } + if observer.lastID != "chatcmpl-2" { + t.Fatalf("lastID = %q, want chatcmpl-2", observer.lastID) + } + if !observer.closed { + t.Fatal("observer was not closed") + } +} From 6c852c63500a24c3377e89318dade218a1acb76e Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 18 Mar 2026 21:13:57 +0100 Subject: [PATCH 06/11] Cap SSE pending buffer before combine --- internal/streaming/observed_sse_stream.go | 20 ++++++++++++---- .../streaming/observed_sse_stream_test.go | 23 +++++++++++++++++++ 2 files changed, 38 insertions(+), 5 deletions(-) diff --git a/internal/streaming/observed_sse_stream.go b/internal/streaming/observed_sse_stream.go index d97ac1dfe..17a280822 100644 --- a/internal/streaming/observed_sse_stream.go +++ b/internal/streaming/observed_sse_stream.go @@ -75,12 +75,22 @@ func (s *ObservedSSEStream) Close() error { func (s *ObservedSSEStream) processChunk(data []byte) { if len(s.pending) > 0 { - if len(data) > maxPendingEventBytes { - data = data[len(data)-maxPendingEventBytes:] + pending := s.pending + pendingLen := len(pending) + if pendingLen > maxPendingEventBytes { + pending = pending[pendingLen-maxPendingEventBytes:] + pendingLen = maxPendingEventBytes } - combined := make([]byte, len(s.pending)+len(data)) - copy(combined, s.pending) - copy(combined[len(s.pending):], data) + + dataLen := len(data) + if dataLen > maxPendingEventBytes { + data = data[dataLen-maxPendingEventBytes:] + dataLen = maxPendingEventBytes + } + + combined := make([]byte, pendingLen+dataLen) + copy(combined, pending) + copy(combined[pendingLen:], data) data = combined s.pending = nil } diff --git a/internal/streaming/observed_sse_stream_test.go b/internal/streaming/observed_sse_stream_test.go index b628b16d8..5891a1832 100644 --- a/internal/streaming/observed_sse_stream_test.go +++ b/internal/streaming/observed_sse_stream_test.go @@ -102,6 +102,29 @@ func TestObservedSSEStream_CapsCombinedPendingData(t *testing.T) { } } +func TestObservedSSEStream_DropsOversizedPendingPrefixBeforeCombining(t *testing.T) { + observer := &trackingObserver{} + s := &ObservedSSEStream{ + observers: []Observer{observer}, + pending: append( + []byte("data: {\"id\":\"stale\"}\n\n"), + bytes.Repeat([]byte("x"), maxPendingEventBytes)..., + ), + } + + s.processChunk([]byte("\n\ndata: {\"id\":\"fresh\"}\n\n")) + + if observer.eventCount != 1 { + t.Fatalf("eventCount = %d, want 1", observer.eventCount) + } + if observer.lastID != "fresh" { + t.Fatalf("lastID = %q, want fresh", observer.lastID) + } + if len(s.pending) != 0 { + t.Fatalf("pending length = %d, want 0", len(s.pending)) + } +} + func TestObservedSSEStream_HandlesCRLFAndDataWithoutSpace(t *testing.T) { streamData := "data:{\"id\":\"chatcmpl-1\"}\r\n\r\ndata: {\"id\":\"chatcmpl-2\"}\r\n\r\ndata:[DONE]\r\n\r\n" observer := &trackingObserver{} From f33fd456fe6d8fa2701ab0a7307c4664613a0c89 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 18 Mar 2026 23:45:02 +0100 Subject: [PATCH 07/11] Harden SSE observation and passthrough usage --- internal/auditlog/auditlog_test.go | 7 +- internal/server/handlers.go | 2 + internal/server/handlers_test.go | 52 ++++++++++ internal/server/passthrough_service.go | 3 + internal/server/passthrough_support.go | 20 +++- internal/streaming/observed_sse_stream.go | 98 +++++++++++++------ .../streaming/observed_sse_stream_test.go | 81 +++++++++++++-- 7 files changed, 217 insertions(+), 46 deletions(-) diff --git a/internal/auditlog/auditlog_test.go b/internal/auditlog/auditlog_test.go index 8720e18fc..9731ade1a 100644 --- a/internal/auditlog/auditlog_test.go +++ b/internal/auditlog/auditlog_test.go @@ -725,7 +725,6 @@ data: [DONE] FlushInterval: 100 * time.Millisecond, } logger := NewLogger(store, cfg) - defer logger.Close() entry := &LogEntry{ ID: "test-entry", @@ -750,9 +749,9 @@ data: [DONE] if err := observedStream.Close(); err != nil { t.Fatalf("failed to close stream: %v", err) } - - // Wait for async write - time.Sleep(200 * time.Millisecond) + if err := logger.Close(); err != nil { + t.Fatalf("failed to close logger: %v", err) + } // Verify entry was logged if len(store.getEntries()) != 1 { diff --git a/internal/server/handlers.go b/internal/server/handlers.go index 03cd22019..00efd6c6b 100644 --- a/internal/server/handlers.go +++ b/internal/server/handlers.go @@ -94,6 +94,8 @@ func (h *Handler) passthrough() *passthroughService { return &passthroughService{ provider: h.provider, logger: h.logger, + usageLogger: h.usageLogger, + pricingResolver: h.pricingResolver, normalizePassthroughV1Prefix: h.normalizePassthroughV1Prefix, enabledPassthroughProviders: h.enabledPassthroughProviders, } diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go index 75177672a..a33137961 100644 --- a/internal/server/handlers_test.go +++ b/internal/server/handlers_test.go @@ -5413,6 +5413,58 @@ func TestProviderPassthrough_AnthropicStream(t *testing.T) { } } +func TestProviderPassthrough_OpenAIStreamWritesUsageEntry(t *testing.T) { + provider := &mockProvider{ + passthroughResponse: &core.PassthroughResponse{ + StatusCode: http.StatusOK, + Headers: map[string][]string{ + "Content-Type": {"text/event-stream"}, + }, + Body: io.NopCloser(strings.NewReader( + "data: {\"id\":\"resp-123\",\"model\":\"gpt-5-mini\",\"usage\":{\"input_tokens\":7,\"output_tokens\":3,\"total_tokens\":10}}\n\n" + + "data: [DONE]\n\n", + )), + }, + } + usageLog := &collectingUsageLogger{ + config: usage.Config{Enabled: true}, + } + + e := echo.New() + handler := NewHandler(provider, nil, usageLog, nil) + e.POST("/p/:provider/*", handler.ProviderPassthrough) + + req := httptest.NewRequest(http.MethodPost, "/p/openai/responses", strings.NewReader(`{"model":"gpt-5-mini"}`)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-Request-ID", "req-pass-stream-usage") + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rec.Code) + } + if len(usageLog.entries) != 1 { + t.Fatalf("usage entries = %d, want 1", len(usageLog.entries)) + } + + entry := usageLog.entries[0] + if entry.Provider != "openai" { + t.Fatalf("Provider = %q, want openai", entry.Provider) + } + if entry.Endpoint != "/p/openai/responses" { + t.Fatalf("Endpoint = %q, want /p/openai/responses", entry.Endpoint) + } + if entry.Model != "gpt-5-mini" { + t.Fatalf("Model = %q, want gpt-5-mini", entry.Model) + } + if entry.TotalTokens != 10 { + t.Fatalf("TotalTokens = %d, want 10", entry.TotalTokens) + } + if entry.RequestID != "req-pass-stream-usage" { + t.Fatalf("RequestID = %q, want req-pass-stream-usage", entry.RequestID) + } +} + func TestPassthroughStreamAuditPath_NormalizesKnownEndpoints(t *testing.T) { tests := []struct { name string diff --git a/internal/server/passthrough_service.go b/internal/server/passthrough_service.go index fb0d84529..4a19a2d10 100644 --- a/internal/server/passthrough_service.go +++ b/internal/server/passthrough_service.go @@ -5,11 +5,14 @@ import ( "gomodel/internal/auditlog" "gomodel/internal/core" + "gomodel/internal/usage" ) type passthroughService struct { provider core.RoutableProvider logger auditlog.LoggerInterface + usageLogger usage.LoggerInterface + pricingResolver usage.PricingResolver normalizePassthroughV1Prefix bool enabledPassthroughProviders map[string]struct{} } diff --git a/internal/server/passthrough_support.go b/internal/server/passthrough_support.go index 8a1b5c0b3..9634455ae 100644 --- a/internal/server/passthrough_support.go +++ b/internal/server/passthrough_support.go @@ -13,6 +13,7 @@ import ( "gomodel/internal/auditlog" "gomodel/internal/core" "gomodel/internal/streaming" + "gomodel/internal/usage" ) var defaultEnabledPassthroughProviders = []string{"openai", "anthropic"} @@ -250,10 +251,23 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT streamEntry.StatusCode = resp.StatusCode } - observers := make([]streaming.Observer, 0, 1) - if observer := auditlog.NewStreamLogObserver(s.logger, streamEntry, passthroughAuditPath(c, providerType, endpoint, info)); observer != nil { + requestID := requestIDFromContextOrHeader(c.Request()) + auditPath := passthroughAuditPath(c, providerType, endpoint, info) + model := "" + if info != nil { + model = strings.TrimSpace(info.Model) + } + model = resolvedModelFromPlan(core.GetExecutionPlan(c.Request().Context()), model) + + observers := make([]streaming.Observer, 0, 2) + if observer := auditlog.NewStreamLogObserver(s.logger, streamEntry, auditPath); observer != nil { observers = append(observers, observer) } + if s.usageLogger != nil && s.usageLogger.Config().Enabled { + if observer := usage.NewStreamUsageObserver(s.usageLogger, model, providerType, requestID, auditPath, s.pricingResolver); observer != nil { + observers = append(observers, observer) + } + } wrappedStream := streaming.NewObservedSSEStream(resp.Body, observers...) defer func() { _ = wrappedStream.Close() @@ -261,7 +275,7 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT c.Response().WriteHeader(resp.StatusCode) if err := flushStream(c.Response(), wrappedStream); err != nil { - recordStreamingError(streamEntry, "", providerType, c.Request().URL.Path, requestIDFromContextOrHeader(c.Request()), err) + recordStreamingError(streamEntry, model, providerType, c.Request().URL.Path, requestID, err) return err } return nil diff --git a/internal/streaming/observed_sse_stream.go b/internal/streaming/observed_sse_stream.go index 17a280822..178c92407 100644 --- a/internal/streaming/observed_sse_stream.go +++ b/internal/streaming/observed_sse_stream.go @@ -26,9 +26,10 @@ type Observer interface { // and fanning them out to observers. type ObservedSSEStream struct { io.ReadCloser - observers []Observer - pending []byte - closed bool + observers []Observer + pending []byte + closed bool + discarding bool } // NewObservedSSEStream returns the original stream when there are no observers. @@ -75,34 +76,60 @@ func (s *ObservedSSEStream) Close() error { func (s *ObservedSSEStream) processChunk(data []byte) { if len(s.pending) > 0 { - pending := s.pending - pendingLen := len(pending) - if pendingLen > maxPendingEventBytes { - pending = pending[pendingLen-maxPendingEventBytes:] - pendingLen = maxPendingEventBytes + idx, sepLen := nextEventBoundary(data) + if idx == -1 { + if len(data) > maxPendingEventBytes || len(s.pending) > maxPendingEventBytes-len(data) { + s.pending = nil + s.discarding = true + return + } + + combinedLen := len(s.pending) + len(data) + combined := make([]byte, combinedLen) + copy(combined, s.pending) + copy(combined[len(s.pending):], data) + s.pending = combined + return } - dataLen := len(data) - if dataLen > maxPendingEventBytes { - data = data[dataLen-maxPendingEventBytes:] - dataLen = maxPendingEventBytes + if idx > maxPendingEventBytes || len(s.pending) > maxPendingEventBytes-idx { + s.pending = nil + data = data[idx+sepLen:] + } else { + combinedLen := len(s.pending) + idx + event := make([]byte, combinedLen) + copy(event, s.pending) + copy(event[len(s.pending):], data[:idx]) + s.pending = nil + s.processEvent(event) + data = data[idx+sepLen:] } - - combined := make([]byte, pendingLen+dataLen) - copy(combined, pending) - copy(combined[pendingLen:], data) - data = combined - s.pending = nil } - for { + for len(data) > 0 { + if s.discarding { + idx, sepLen := nextEventBoundary(data) + if idx == -1 { + return + } + data = data[idx+sepLen:] + s.discarding = false + continue + } + idx, sepLen := nextEventBoundary(data) if idx == -1 { s.savePending(data) return } - s.processEvent(data[:idx]) + if idx > maxPendingEventBytes { + data = data[idx+sepLen:] + continue + } + if idx > 0 { + s.processEvent(data[:idx]) + } data = data[idx+sepLen:] } } @@ -123,22 +150,29 @@ func (s *ObservedSSEStream) processBufferedEvents(data []byte) { func (s *ObservedSSEStream) processEvent(event []byte) { lines := bytes.Split(event, []byte("\n")) + payloadLines := make([][]byte, 0, len(lines)) for _, line := range lines { jsonData, ok := parseDataLine(line) if !ok { continue } - if bytes.Equal(jsonData, donePayload) { - continue - } + payloadLines = append(payloadLines, jsonData) + } + if len(payloadLines) == 0 { + return + } - var payload map[string]interface{} - if err := json.Unmarshal(jsonData, &payload); err != nil { - continue - } - for _, observer := range s.observers { - observer.OnJSONEvent(payload) - } + jsonData := bytes.Join(payloadLines, []byte("\n")) + if bytes.Equal(jsonData, donePayload) { + return + } + + var payload map[string]interface{} + if err := json.Unmarshal(jsonData, &payload); err != nil { + return + } + for _, observer := range s.observers { + observer.OnJSONEvent(payload) } } @@ -176,7 +210,9 @@ func (s *ObservedSSEStream) savePending(data []byte) { return } if len(data) > maxPendingEventBytes { - data = data[len(data)-maxPendingEventBytes:] + s.pending = nil + s.discarding = true + return } s.pending = append(s.pending[:0], data...) } diff --git a/internal/streaming/observed_sse_stream_test.go b/internal/streaming/observed_sse_stream_test.go index 5891a1832..0a8bf0c4a 100644 --- a/internal/streaming/observed_sse_stream_test.go +++ b/internal/streaming/observed_sse_stream_test.go @@ -8,13 +8,15 @@ import ( ) type trackingObserver struct { - eventCount int - lastID string - closed bool + eventCount int + lastID string + lastPayload map[string]interface{} + closed bool } func (o *trackingObserver) OnJSONEvent(payload map[string]interface{}) { o.eventCount++ + o.lastPayload = payload if id, _ := payload["id"].(string); id != "" { o.lastID = id } @@ -86,7 +88,40 @@ func TestObservedSSEStream_ParsesFragmentedFinalEventOnClose(t *testing.T) { } } -func TestObservedSSEStream_CapsCombinedPendingData(t *testing.T) { +func TestObservedSSEStream_ReassemblesMultilineDataEvent(t *testing.T) { + streamData := "data: {\"id\":\"chatcmpl-multiline\",\n" + + "data: \"usage\":{\"total_tokens\":3}}\n\n" + + "data: [DONE]\n\n" + observer := &trackingObserver{} + stream := NewObservedSSEStream(io.NopCloser(strings.NewReader(streamData)), observer) + + data, err := io.ReadAll(stream) + if err != nil { + t.Fatalf("ReadAll error: %v", err) + } + if string(data) != streamData { + t.Fatalf("stream passthrough mismatch") + } + if err := stream.Close(); err != nil { + t.Fatalf("Close error: %v", err) + } + + if observer.eventCount != 1 { + t.Fatalf("eventCount = %d, want 1", observer.eventCount) + } + if observer.lastID != "chatcmpl-multiline" { + t.Fatalf("lastID = %q, want chatcmpl-multiline", observer.lastID) + } + usage, ok := observer.lastPayload["usage"].(map[string]interface{}) + if !ok { + t.Fatalf("usage = %#v, want object", observer.lastPayload["usage"]) + } + if got := usage["total_tokens"]; got != float64(3) { + t.Fatalf("usage.total_tokens = %#v, want 3", got) + } +} + +func TestObservedSSEStream_DiscardsOversizedPendingDataWithoutTailCapping(t *testing.T) { s := &ObservedSSEStream{ pending: bytes.Repeat([]byte("a"), maxPendingEventBytes), } @@ -94,11 +129,41 @@ func TestObservedSSEStream_CapsCombinedPendingData(t *testing.T) { s.processChunk(data) - if got := len(s.pending); got != maxPendingEventBytes { - t.Fatalf("pending length = %d, want %d", got, maxPendingEventBytes) + if got := len(s.pending); got != 0 { + t.Fatalf("pending length = %d, want 0", got) + } + if !s.discarding { + t.Fatal("discarding = false, want true") + } +} + +func TestObservedSSEStream_DropsOversizedBufferedEventAndResumesWithinSameChunk(t *testing.T) { + observer := &trackingObserver{} + s := &ObservedSSEStream{ + observers: []Observer{observer}, + pending: append( + []byte("data: {\"id\":\"too-big\",\"payload\":\""), + bytes.Repeat([]byte("a"), maxPendingEventBytes/2)..., + ), + } + data := append( + append( + append( + bytes.Repeat([]byte("b"), maxPendingEventBytes/2+1), + []byte("\"}\n\ndata: {\"id\":\"fresh\"}\n\n")..., + ), + bytes.Repeat([]byte("c"), maxPendingEventBytes+1)..., + ), + []byte("ignored-trailer")..., + ) + + s.processChunk(data) + + if observer.eventCount != 1 { + t.Fatalf("eventCount = %d, want 1", observer.eventCount) } - if !bytes.Equal(s.pending, data[len(data)-maxPendingEventBytes:]) { - t.Fatal("pending bytes do not match the capped suffix of the latest chunk") + if observer.lastID != "fresh" { + t.Fatalf("lastID = %q, want fresh", observer.lastID) } } From 2ecfd02985d8bbda52c316d4c7dadb27a2305301 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 19 Mar 2026 13:33:54 +0100 Subject: [PATCH 08/11] Keep passthrough usage on client-visible routes --- internal/server/handlers_test.go | 49 ++++++++++++++++++++++++++ internal/server/passthrough_support.go | 6 +++- 2 files changed, 54 insertions(+), 1 deletion(-) diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go index a33137961..517a9efe0 100644 --- a/internal/server/handlers_test.go +++ b/internal/server/handlers_test.go @@ -5465,6 +5465,55 @@ func TestProviderPassthrough_OpenAIStreamWritesUsageEntry(t *testing.T) { } } +func TestProviderPassthrough_OpenAIStreamUsageKeepsClientVisibleRoute(t *testing.T) { + provider := &mockProvider{ + passthroughResponse: &core.PassthroughResponse{ + StatusCode: http.StatusOK, + Headers: map[string][]string{ + "Content-Type": {"text/event-stream"}, + }, + Body: io.NopCloser(strings.NewReader( + "data: {\"id\":\"resp-123\",\"model\":\"gpt-5-mini\",\"usage\":{\"input_tokens\":7,\"output_tokens\":3,\"total_tokens\":10}}\n\n" + + "data: [DONE]\n\n", + )), + }, + } + usageLog := &collectingUsageLogger{ + config: usage.Config{Enabled: true}, + } + + e := echo.New() + handler := NewHandler(provider, nil, usageLog, nil) + e.POST("/p/:provider/*", handler.ProviderPassthrough) + + req := httptest.NewRequest(http.MethodPost, "/p/openai/v1/responses", strings.NewReader(`{"model":"gpt-5-mini"}`)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-Request-ID", "req-pass-stream-visible-path") + req = req.WithContext(core.WithExecutionPlan(req.Context(), &core.ExecutionPlan{ + Mode: core.ExecutionModePassthrough, + ProviderType: "openai", + Passthrough: &core.PassthroughRouteInfo{ + Provider: "openai", + RawEndpoint: "v1/responses", + NormalizedEndpoint: "responses", + AuditPath: "/v1/responses", + Model: "gpt-5-mini", + }, + })) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rec.Code) + } + if len(usageLog.entries) != 1 { + t.Fatalf("usage entries = %d, want 1", len(usageLog.entries)) + } + if got := usageLog.entries[0].Endpoint; got != "/p/openai/v1/responses" { + t.Fatalf("Endpoint = %q, want /p/openai/v1/responses", got) + } +} + func TestPassthroughStreamAuditPath_NormalizesKnownEndpoints(t *testing.T) { tests := []struct { name string diff --git a/internal/server/passthrough_support.go b/internal/server/passthrough_support.go index 9634455ae..d3d3e3949 100644 --- a/internal/server/passthrough_support.go +++ b/internal/server/passthrough_support.go @@ -253,6 +253,10 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT requestID := requestIDFromContextOrHeader(c.Request()) auditPath := passthroughAuditPath(c, providerType, endpoint, info) + usagePath := auditPath + if requestPath := strings.TrimSpace(c.Request().URL.Path); requestPath != "" { + usagePath = requestPath + } model := "" if info != nil { model = strings.TrimSpace(info.Model) @@ -264,7 +268,7 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT observers = append(observers, observer) } if s.usageLogger != nil && s.usageLogger.Config().Enabled { - if observer := usage.NewStreamUsageObserver(s.usageLogger, model, providerType, requestID, auditPath, s.pricingResolver); observer != nil { + if observer := usage.NewStreamUsageObserver(s.usageLogger, model, providerType, requestID, usagePath, s.pricingResolver); observer != nil { observers = append(observers, observer) } } From be90ea63802c4233fac2315bb1dd74a10902b55c Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 19 Mar 2026 14:18:42 +0100 Subject: [PATCH 09/11] Handle split SSE delimiters across reads --- internal/streaming/observed_sse_stream.go | 157 +++++++++++++++--- .../streaming/observed_sse_stream_test.go | 48 ++++++ 2 files changed, 183 insertions(+), 22 deletions(-) diff --git a/internal/streaming/observed_sse_stream.go b/internal/streaming/observed_sse_stream.go index 178c92407..c256d916d 100644 --- a/internal/streaming/observed_sse_stream.go +++ b/internal/streaming/observed_sse_stream.go @@ -7,6 +7,7 @@ import ( ) const maxPendingEventBytes = 256 * 1024 +const maxBoundaryTailBytes = 3 var ( lfEventBoundary = []byte("\n\n") @@ -26,10 +27,11 @@ type Observer interface { // and fanning them out to observers. type ObservedSSEStream struct { io.ReadCloser - observers []Observer - pending []byte - closed bool - discarding bool + observers []Observer + pending []byte + discardTail []byte + closed bool + discarding bool } // NewObservedSSEStream returns the original stream when there are no observers. @@ -76,44 +78,57 @@ func (s *ObservedSSEStream) Close() error { func (s *ObservedSSEStream) processChunk(data []byte) { if len(s.pending) > 0 { - idx, sepLen := nextEventBoundary(data) + pendingLen := len(s.pending) + idx, sepLen := nextJoinedEventBoundary(s.pending, data) if idx == -1 { - if len(data) > maxPendingEventBytes || len(s.pending) > maxPendingEventBytes-len(data) { - s.pending = nil - s.discarding = true + if len(data) > maxPendingEventBytes || pendingLen > maxPendingEventBytes-len(data) { + s.startDiscarding(s.pending, data) return } - combinedLen := len(s.pending) + len(data) + combinedLen := pendingLen + len(data) combined := make([]byte, combinedLen) copy(combined, s.pending) - copy(combined[len(s.pending):], data) + copy(combined[pendingLen:], data) s.pending = combined return } - if idx > maxPendingEventBytes || len(s.pending) > maxPendingEventBytes-idx { + if idx > maxPendingEventBytes { s.pending = nil - data = data[idx+sepLen:] - } else { - combinedLen := len(s.pending) + idx - event := make([]byte, combinedLen) - copy(event, s.pending) - copy(event[len(s.pending):], data[:idx]) + data = data[dataOffsetAfterBoundary(pendingLen, idx, sepLen):] + } else if idx < pendingLen { + event := append([]byte(nil), s.pending[:idx]...) s.pending = nil s.processEvent(event) - data = data[idx+sepLen:] + data = data[dataOffsetAfterBoundary(pendingLen, idx, sepLen):] + } else { + dataIdx := idx - pendingLen + if pendingLen > maxPendingEventBytes-dataIdx { + s.pending = nil + data = data[dataIdx+sepLen:] + } else { + combinedLen := pendingLen + dataIdx + event := make([]byte, combinedLen) + copy(event, s.pending) + copy(event[pendingLen:], data[:dataIdx]) + s.pending = nil + s.processEvent(event) + data = data[dataIdx+sepLen:] + } } } for len(data) > 0 { if s.discarding { - idx, sepLen := nextEventBoundary(data) + idx, sepLen := nextJoinedEventBoundary(s.discardTail, data) if idx == -1 { + s.discardTail = joinedSuffix(s.discardTail, data, maxBoundaryTailBytes) return } - data = data[idx+sepLen:] + data = data[dataOffsetAfterBoundary(len(s.discardTail), idx, sepLen):] s.discarding = false + s.discardTail = nil continue } @@ -210,9 +225,107 @@ func (s *ObservedSSEStream) savePending(data []byte) { return } if len(data) > maxPendingEventBytes { - s.pending = nil - s.discarding = true + s.startDiscarding(nil, data) return } s.pending = append(s.pending[:0], data...) } + +func (s *ObservedSSEStream) startDiscarding(prefix, data []byte) { + s.pending = nil + s.discarding = true + s.discardTail = joinedSuffix(prefix, data, maxBoundaryTailBytes) +} + +func nextJoinedEventBoundary(prefix, data []byte) (idx int, sepLen int) { + idx = -1 + + crossIdx, crossSepLen := nextBoundaryAcrossJoin(prefix, data) + if crossIdx != -1 { + idx, sepLen = crossIdx, crossSepLen + } + + dataIdx, dataSepLen := nextEventBoundary(data) + if dataIdx != -1 { + combinedIdx := len(prefix) + dataIdx + if idx == -1 || combinedIdx < idx { + idx, sepLen = combinedIdx, dataSepLen + } + } + + return idx, sepLen +} + +func nextBoundaryAcrossJoin(prefix, data []byte) (idx int, sepLen int) { + idx = -1 + start := len(prefix) - maxBoundaryTailBytes + if start < 0 { + start = 0 + } + + for offset := start; offset < len(prefix); offset++ { + for _, boundary := range [][]byte{lfEventBoundary, crlfEventBoundary} { + if offset+len(boundary) <= len(prefix) { + continue + } + if joinedBytesMatch(prefix, data, offset, boundary) { + if idx == -1 || offset < idx { + idx = offset + sepLen = len(boundary) + } + } + } + } + + return idx, sepLen +} + +func joinedBytesMatch(prefix, data []byte, start int, boundary []byte) bool { + for i, want := range boundary { + pos := start + i + var got byte + switch { + case pos < len(prefix): + got = prefix[pos] + case pos-len(prefix) < len(data): + got = data[pos-len(prefix)] + default: + return false + } + if got != want { + return false + } + } + return true +} + +func dataOffsetAfterBoundary(prefixLen, idx, sepLen int) int { + if idx >= prefixLen { + return idx - prefixLen + sepLen + } + + prefixConsumed := prefixLen - idx + if prefixConsumed >= sepLen { + return 0 + } + return sepLen - prefixConsumed +} + +func joinedSuffix(prefix, data []byte, n int) []byte { + if n <= 0 { + return nil + } + if len(data) >= n { + return append([]byte(nil), data[len(data)-n:]...) + } + + needPrefix := n - len(data) + if needPrefix > len(prefix) { + needPrefix = len(prefix) + } + + result := make([]byte, needPrefix+len(data)) + copy(result, prefix[len(prefix)-needPrefix:]) + copy(result[needPrefix:], data) + return result +} diff --git a/internal/streaming/observed_sse_stream_test.go b/internal/streaming/observed_sse_stream_test.go index 0a8bf0c4a..c543b262a 100644 --- a/internal/streaming/observed_sse_stream_test.go +++ b/internal/streaming/observed_sse_stream_test.go @@ -121,6 +121,26 @@ func TestObservedSSEStream_ReassemblesMultilineDataEvent(t *testing.T) { } } +func TestObservedSSEStream_DetectsBoundarySplitAcrossReads(t *testing.T) { + observer := &trackingObserver{} + s := &ObservedSSEStream{ + observers: []Observer{observer}, + } + + s.processChunk([]byte("data:{\"id\":\"chatcmpl-1\"}\r\n\r")) + s.processChunk([]byte("\ndata:{\"id\":\"chatcmpl-2\"}\r\n\r\n")) + + if observer.eventCount != 2 { + t.Fatalf("eventCount = %d, want 2", observer.eventCount) + } + if observer.lastID != "chatcmpl-2" { + t.Fatalf("lastID = %q, want chatcmpl-2", observer.lastID) + } + if len(s.pending) != 0 { + t.Fatalf("pending length = %d, want 0", len(s.pending)) + } +} + func TestObservedSSEStream_DiscardsOversizedPendingDataWithoutTailCapping(t *testing.T) { s := &ObservedSSEStream{ pending: bytes.Repeat([]byte("a"), maxPendingEventBytes), @@ -167,6 +187,34 @@ func TestObservedSSEStream_DropsOversizedBufferedEventAndResumesWithinSameChunk( } } +func TestObservedSSEStream_ResumesAfterDiscardWhenBoundarySplitsAcrossReads(t *testing.T) { + observer := &trackingObserver{} + s := &ObservedSSEStream{ + observers: []Observer{observer}, + } + + oversized := append( + append( + []byte("data:{\"id\":\"too-big\",\"payload\":\""), + bytes.Repeat([]byte("x"), maxPendingEventBytes)..., + ), + []byte("\"}\r\n\r")..., + ) + + s.processChunk(oversized) + s.processChunk([]byte("\ndata:{\"id\":\"fresh\"}\r\n\r\n")) + + if observer.eventCount != 1 { + t.Fatalf("eventCount = %d, want 1", observer.eventCount) + } + if observer.lastID != "fresh" { + t.Fatalf("lastID = %q, want fresh", observer.lastID) + } + if s.discarding { + t.Fatal("discarding = true, want false") + } +} + func TestObservedSSEStream_DropsOversizedPendingPrefixBeforeCombining(t *testing.T) { observer := &trackingObserver{} s := &ObservedSSEStream{ From c33b70551ee356a2b8e5d3a983236be03ab3842f Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 19 Mar 2026 14:37:22 +0100 Subject: [PATCH 10/11] Avoid double-close in passthrough SSE --- internal/server/handlers_test.go | 50 ++++++++++++++++++++++++++ internal/server/passthrough_support.go | 8 +++-- 2 files changed, 55 insertions(+), 3 deletions(-) diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go index 517a9efe0..98a0bf7f9 100644 --- a/internal/server/handlers_test.go +++ b/internal/server/handlers_test.go @@ -209,6 +209,19 @@ func (r *erroringReadCloser) Close() error { return nil } +type closeCountingReadCloser struct { + io.ReadCloser + closes int +} + +func (r *closeCountingReadCloser) Close() error { + r.closes++ + if r.ReadCloser == nil { + return nil + } + return r.ReadCloser.Close() +} + func setPathParam(c *echo.Context, name, value string) { c.SetPathValues(echo.PathValues{{Name: name, Value: value}}) } @@ -5413,6 +5426,43 @@ func TestProviderPassthrough_AnthropicStream(t *testing.T) { } } +func TestProviderPassthrough_StreamWithoutObserversClosesUpstreamBodyOnce(t *testing.T) { + body := &closeCountingReadCloser{ + ReadCloser: &chunkedReadCloser{ + chunks: [][]byte{ + []byte("event: message_start\n"), + []byte("data: {\"type\":\"message_start\"}\n\n"), + }, + }, + } + provider := &mockProvider{ + passthroughResponse: &core.PassthroughResponse{ + StatusCode: http.StatusOK, + Headers: map[string][]string{ + "Content-Type": {"text/event-stream"}, + }, + Body: body, + }, + } + + e := echo.New() + handler := NewHandler(provider, nil, nil, nil) + e.POST("/p/:provider/*", handler.ProviderPassthrough) + + req := httptest.NewRequest(http.MethodPost, "/p/anthropic/messages", strings.NewReader(`{"model":"claude-sonnet-4-5"}`)) + req.Header.Set("Content-Type", "application/json") + + rec := &flushCountingRecorder{ResponseRecorder: httptest.NewRecorder()} + e.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rec.Code) + } + if body.closes != 1 { + t.Fatalf("Close calls = %d, want 1", body.closes) + } +} + func TestProviderPassthrough_OpenAIStreamWritesUsageEntry(t *testing.T) { provider := &mockProvider{ passthroughResponse: &core.PassthroughResponse{ diff --git a/internal/server/passthrough_support.go b/internal/server/passthrough_support.go index d3d3e3949..9e0e4b9e2 100644 --- a/internal/server/passthrough_support.go +++ b/internal/server/passthrough_support.go @@ -273,9 +273,11 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT } } wrappedStream := streaming.NewObservedSSEStream(resp.Body, observers...) - defer func() { - _ = wrappedStream.Close() - }() + if len(observers) > 0 { + defer func() { + _ = wrappedStream.Close() + }() + } c.Response().WriteHeader(resp.StatusCode) if err := flushStream(c.Response(), wrappedStream); err != nil { From 7876d97a9c2c2d141af712f63fc998be0290d47b Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 19 Mar 2026 14:50:36 +0100 Subject: [PATCH 11/11] docs(repo): note commit title format --- AGENTS.md | 10 ++++++++++ CLAUDE.md | 10 ++++++++++ 2 files changed, 20 insertions(+) diff --git a/AGENTS.md b/AGENTS.md index 156851e53..86670cafb 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -21,3 +21,13 @@ Backward compatibility is not a primary constraint in the current development st 1. Make small, focused changes. 2. Run format/lint/tests relevant to the change. + +## Commit Format + +Use Conventional Commit format for commit subjects and PR titles: + +`type(scope): short summary` + +Allowed types: feat, fix, perf, docs, refactor, test, build, ci, chore, revert + +Squash merges should preserve the PR title as the resulting commit subject. diff --git a/CLAUDE.md b/CLAUDE.md index 6a5cf3a03..9bab63c31 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -44,6 +44,16 @@ make swagger # Regenerate Swagger docs **Build tags:** E2E tests require `-tags=e2e`, integration tests require `-tags=integration`, contract tests require `-tags=contract`. The Makefile handles this automatically. +## Commit And PR Title Format + +Use Conventional Commit format for commit subjects and PR titles: + +`type(scope): short summary` + +Allowed types: feat, fix, perf, docs, refactor, test, build, ci, chore, revert + +Prefer squash-and-merge to keep the merged commit subject aligned with the PR title. + ## Error Handling - All errors returned to clients must be instances of `core.GatewayError`.