From b06d16a017b1059b6d8e5151f2b01a80c906021a Mon Sep 17 00:00:00 2001 From: hcastc00 Date: Wed, 9 Sep 2026 12:38:50 +0200 Subject: [PATCH] fix: Request end up prematurely because of client disconection --- .../http/handlers/chat_completions_handler.go | 53 ++++- .../handlers/chat_completions_stream_test.go | 89 +++++++- .../internal/adapters/http/handlers/common.go | 10 +- .../http/handlers/privacy_logging_test.go | 29 +++ .../http/handlers/stream_usage_repro_test.go | 200 ++++++++++++++++++ .../adapters/providers/openai/adapter.go | 4 +- .../adapters/providers/openai/logging.go | 99 ++------- .../adapters/providers/openai/policy_test.go | 26 ++- .../adapters/providers/openai/stream.go | 36 +++- .../providers/openai/stream_diagnostics.go | 95 +++++++++ .../openai/stream_diagnostics_test.go | 132 ++++++++++++ .../adapters/providers/openai/stream_test.go | 25 +++ .../adapters/providers/tinfoil/adapter.go | 2 +- .../application/services/generate_service.go | 78 ++++--- .../services/stream_logging_test.go | 119 +++++++++++ pkg/observability/logfields/finish.go | 16 ++ pkg/observability/logger/zap.go | 43 ++-- pkg/observability/logger/zap_test.go | 31 +++ 18 files changed, 938 insertions(+), 149 deletions(-) create mode 100644 apps/gateway/internal/adapters/http/handlers/privacy_logging_test.go create mode 100644 apps/gateway/internal/adapters/http/handlers/stream_usage_repro_test.go create mode 100644 apps/gateway/internal/adapters/providers/openai/stream_diagnostics.go create mode 100644 apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go create mode 100644 apps/gateway/internal/application/services/stream_logging_test.go create mode 100644 pkg/observability/logfields/finish.go create mode 100644 pkg/observability/logger/zap_test.go diff --git a/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go b/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go index 4576889..ad078bf 100644 --- a/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go +++ b/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go @@ -1,6 +1,7 @@ package handlers import ( + "context" "io" "net/http" "sync" @@ -15,6 +16,10 @@ import ( "github.com/google/uuid" ) +// Bound the wait for usage after generation has finished, without extending +// the lifetime of requests canceled before completion. +const streamUsageDrainTimeout = 5 * time.Second + // ChatCompletionsHandler handles POST /v1/chat/completions. type ChatCompletionsHandler struct { service *services.ChatCompletionsService @@ -68,13 +73,15 @@ func (h *ChatCompletionsHandler) Handle(w http.ResponseWriter, r *http.Request) } func (h *ChatCompletionsHandler) handleStream(w http.ResponseWriter, r *http.Request, genReq domain.GenerateRequest, token string) { - stream, model, err := h.service.ExecuteStream(r.Context(), genReq, token) + ctx, cancelUpstream := context.WithCancel(r.Context()) + defer cancelUpstream() + stream, model, err := h.service.ExecuteStream(ctx, genReq, token) if err != nil { WriteErrorWithLog(w, r, h.logger, err) return } defer stream.Close() - h.writeStream(w, r, genReq, stream, model) + h.writeStream(w, r, genReq, stream, model, cancelUpstream) } func (h *ChatCompletionsHandler) writeStream( @@ -83,6 +90,7 @@ func (h *ChatCompletionsHandler) writeStream( genReq domain.GenerateRequest, stream ports.GenerationStream, model *domain.PublicModel, + cancelUpstream context.CancelFunc, ) { responseModelID := genReq.PublicModelID if model != nil { @@ -137,23 +145,48 @@ func (h *ChatCompletionsHandler) writeStream( responseID = event.ProviderResponseID } + // Clients may close on finish_reason as well as [DONE]. Keep both + // signals back until metering has consumed trailing usage and settled + // the request. Canceling the upstream context safely unblocks reads + // if a provider never terminates its post-completion stream. + if event.Type == domain.StreamEventCompleted { + requestID := middleware.GetRequestID(r.Context()) + providerID := responseID + timer := time.AfterFunc(streamUsageDrainTimeout, func() { + h.logger.Warn("stream finalization timeout", "request_id", requestID, "provider_request_id", providerID, "model", responseModelID, "timeout_ms", streamUsageDrainTimeout.Milliseconds()) + cancelUpstream() + }) + for { + tail, err := stream.Recv() + if err != nil { + break + } + if tail.Usage != nil { + event.Usage = tail.Usage + } + if tail.FinishReason != nil { + event.FinishReason = tail.FinishReason + } + } + timer.Stop() + } chunk, done := mapper.DomainStreamEventToChatChunk(event, responseModelID, responseID, createdAt) if chunk != nil { writeMu.Lock() - sw.WriteData(chunk) + writeErr := sw.WriteData(chunk) writeMu.Unlock() + if writeErr != nil { + h.logger.Warn("stream client write failed", "request_id", middleware.GetRequestID(r.Context()), "provider_request_id", responseID, "phase", "chunk", "write_failed", true) + } } if done { writeMu.Lock() - sw.WriteDone() + writeErr := sw.WriteDone() writeMu.Unlock() - // Keep draining so the usage-tracking wrapper can accumulate - // the final usage chunk before io.EOF triggers recording. - for { - if _, err := stream.Recv(); err != nil { - break - } + if writeErr != nil { + h.logger.Warn("stream client write failed", "request_id", middleware.GetRequestID(r.Context()), "provider_request_id", responseID, "phase", "done", "write_failed", true) } + h.logger.Debug("stream completion sent", "request_id", middleware.GetRequestID(r.Context()), "provider_request_id", responseID, "done_write_failed", writeErr != nil, "client_context_canceled", r.Context().Err() != nil) break } } diff --git a/apps/gateway/internal/adapters/http/handlers/chat_completions_stream_test.go b/apps/gateway/internal/adapters/http/handlers/chat_completions_stream_test.go index 0db9d93..72726d3 100644 --- a/apps/gateway/internal/adapters/http/handlers/chat_completions_stream_test.go +++ b/apps/gateway/internal/adapters/http/handlers/chat_completions_stream_test.go @@ -1,22 +1,28 @@ package handlers import ( + "context" "errors" "io" "net/http/httptest" "strings" "testing" + "time" "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) type handlerTestStream struct { - events []domain.StreamEvent - err error - index int + events []domain.StreamEvent + err error + index int + beforeRecv func() } func (s *handlerTestStream) Recv() (domain.StreamEvent, error) { + if s.beforeRecv != nil { + s.beforeRecv() + } if s.index < len(s.events) { event := s.events[s.index] s.index++ @@ -28,6 +34,71 @@ func (s *handlerTestStream) Recv() (domain.StreamEvent, error) { return domain.StreamEvent{}, io.EOF } +func TestChatCompletionsHandler_SettlesBeforeEmittingFinish(t *testing.T) { + finish := "tool_calls" + recorder := httptest.NewRecorder() + stream := &handlerTestStream{ + events: []domain.StreamEvent{ + {Type: domain.StreamEventCompleted, FinishReason: &finish}, + {Type: domain.StreamEventCompleted, Usage: &domain.Usage{PromptTokens: 100, CompletionTokens: 20, TotalTokens: 120}}, + }, + beforeRecv: func() { + if strings.Contains(recorder.Body.String(), `"finish_reason"`) || strings.Contains(recorder.Body.String(), "[DONE]") { + t.Error("client received a completion signal before upstream EOF") + } + }, + } + handler := &ChatCompletionsHandler{logger: &confidentialTestLogger{}} + handler.writeStream(recorder, httptest.NewRequest("POST", "/v1/chat/completions", nil), domain.GenerateRequest{PublicModelID: "test"}, stream, nil, func() {}) + if !strings.Contains(recorder.Body.String(), `"prompt_tokens":100`) || !strings.Contains(recorder.Body.String(), `"finish_reason":"tool_calls"`) { + t.Fatalf("missing final usage or finish: %s", recorder.Body.String()) + } + if strings.Count(recorder.Body.String(), "data: [DONE]") != 1 { + t.Fatalf("expected exactly one DONE: %s", recorder.Body.String()) + } +} + +type stalledCompletionStream struct { + handlerTestStream + ctx context.Context +} + +func (s *stalledCompletionStream) Recv() (domain.StreamEvent, error) { + if s.index < len(s.events) { + return s.handlerTestStream.Recv() + } + <-s.ctx.Done() + return domain.StreamEvent{}, s.ctx.Err() +} + +func TestChatCompletionsHandler_BoundsTrailingUsageWait(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), streamUsageDrainTimeout+2*time.Second) + defer cancel() + finish := "tool_calls" + stream := &stalledCompletionStream{ + handlerTestStream: handlerTestStream{events: []domain.StreamEvent{{Type: domain.StreamEventCompleted, FinishReason: &finish}}}, + ctx: ctx, + } + recorder := httptest.NewRecorder() + log := &finalizationLogRecorder{warnings: make(chan string, 1)} + handler := &ChatCompletionsHandler{logger: log} + handler.writeStream(recorder, httptest.NewRequest("POST", "/v1/chat/completions", nil), domain.GenerateRequest{PublicModelID: "test"}, stream, nil, cancel) + select { + case message := <-log.warnings: + if message != "stream finalization timeout" { + t.Fatalf("unexpected warning: %s", message) + } + default: + t.Fatal("finalization timeout was not logged") + } + if ctx.Err() != context.Canceled { + t.Fatal("stalled upstream was not canceled") + } + if !strings.Contains(recorder.Body.String(), `"finish_reason":"tool_calls"`) || strings.Contains(recorder.Body.String(), `"usage"`) { + t.Fatalf("expected finished generation without fabricated usage: %s", recorder.Body.String()) + } +} + func (*handlerTestStream) Close() error { return nil } func TestChatCompletionsHandler_StreamEmitsExactlyOneDoneOnCleanEOF(t *testing.T) { @@ -59,7 +130,7 @@ func TestChatCompletionsHandler_StreamEmitsExactlyOneDoneOnCleanEOF(t *testing.T recorder := httptest.NewRecorder() request := httptest.NewRequest("POST", "/v1/chat/completions", nil) - handler.writeStream(recorder, request, domain.GenerateRequest{PublicModelID: "test"}, &handlerTestStream{events: test.events}, nil) + handler.writeStream(recorder, request, domain.GenerateRequest{PublicModelID: "test"}, &handlerTestStream{events: test.events}, nil, func() {}) if count := strings.Count(recorder.Body.String(), "data: [DONE]\n\n"); count != 1 { t.Fatalf("DONE marker count = %d, want 1; stream = %q", count, recorder.Body.String()) @@ -73,9 +144,17 @@ func TestChatCompletionsHandler_StreamErrorDoesNotClaimCompletion(t *testing.T) recorder := httptest.NewRecorder() request := httptest.NewRequest("POST", "/v1/chat/completions", nil) - handler.writeStream(recorder, request, domain.GenerateRequest{PublicModelID: "test"}, &handlerTestStream{err: errors.New("upstream failed")}, nil) + handler.writeStream(recorder, request, domain.GenerateRequest{PublicModelID: "test"}, &handlerTestStream{err: errors.New("upstream failed")}, nil, func() {}) if strings.Contains(recorder.Body.String(), "data: [DONE]\n\n") { t.Fatalf("errored stream claimed completion: %q", recorder.Body.String()) } } + +// A channel keeps the timer's logging callback safe to observe from the test. +type finalizationLogRecorder struct { + confidentialTestLogger + warnings chan string +} + +func (l *finalizationLogRecorder) Warn(message string, _ ...any) { l.warnings <- message } diff --git a/apps/gateway/internal/adapters/http/handlers/common.go b/apps/gateway/internal/adapters/http/handlers/common.go index b8e53a9..9563b5e 100644 --- a/apps/gateway/internal/adapters/http/handlers/common.go +++ b/apps/gateway/internal/adapters/http/handlers/common.go @@ -2,6 +2,7 @@ package handlers import ( "encoding/json" + "fmt" "net/http" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/dto" @@ -41,7 +42,7 @@ func WriteErrorWithLog(w http.ResponseWriter, r *http.Request, logger ports.Logg logger.Error("untyped internal error", "request_id", middleware.GetRequestID(r.Context()), "path", r.URL.Path, - "original_error", err.Error(), + "error_type", fmt.Sprintf("%T", err), ) gwErr = domain.ErrInternal("an internal error occurred") } @@ -53,9 +54,11 @@ func WriteErrorWithLog(w http.ResponseWriter, r *http.Request, logger ports.Logg "status", gwErr.HTTPStatus, "gateway_status", gwErr.HTTPStatus, "error_code", gwErr.Code, - "error", gwErr.Message, } - fields = append(fields, gwErr.LogFields()...) + // Provider errors can echo prompts; only retain numeric upstream status. + if status, ok := gwErr.Metadata["upstream_status"].(int); ok { + fields = append(fields, "upstream_status", status) + } logger.Error("request failed", fields...) } else if gwErr.HTTPStatus >= 400 { logger.Warn("request error", @@ -63,7 +66,6 @@ func WriteErrorWithLog(w http.ResponseWriter, r *http.Request, logger ports.Logg "status", gwErr.HTTPStatus, "gateway_status", gwErr.HTTPStatus, "error_code", gwErr.Code, - "error", gwErr.Message, ) } diff --git a/apps/gateway/internal/adapters/http/handlers/privacy_logging_test.go b/apps/gateway/internal/adapters/http/handlers/privacy_logging_test.go new file mode 100644 index 0000000..4092cb9 --- /dev/null +++ b/apps/gateway/internal/adapters/http/handlers/privacy_logging_test.go @@ -0,0 +1,29 @@ +package handlers + +import ( + "errors" + "fmt" + "net/http/httptest" + "strings" + "testing" + + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" +) + +type privacyLogger struct{ output strings.Builder } + +func (l *privacyLogger) Debug(m string, f ...any) { fmt.Fprint(&l.output, m, f) } +func (l *privacyLogger) Info(m string, f ...any) { l.Debug(m, f...) } +func (l *privacyLogger) Warn(m string, f ...any) { l.Debug(m, f...) } +func (l *privacyLogger) Error(m string, f ...any) { l.Debug(m, f...) } + +func TestErrorLogsDoNotEchoRequestContents(t *testing.T) { + const secret = "PRIVATE-PROMPT-CANARY" + for _, err := range []error{errors.New(secret), domain.ErrInvalidField(secret), domain.ErrProviderError(502, secret).WithMeta("upstream_error", secret, "upstream_status", 400)} { + logger := &privacyLogger{} + WriteErrorWithLog(httptest.NewRecorder(), httptest.NewRequest("POST", "/v1/chat/completions", nil), logger, err) + if strings.Contains(logger.output.String(), secret) { + t.Fatal("error log leaked private content") + } + } +} diff --git a/apps/gateway/internal/adapters/http/handlers/stream_usage_repro_test.go b/apps/gateway/internal/adapters/http/handlers/stream_usage_repro_test.go new file mode 100644 index 0000000..3ce5df0 --- /dev/null +++ b/apps/gateway/internal/adapters/http/handlers/stream_usage_repro_test.go @@ -0,0 +1,200 @@ +package handlers_test + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/handlers" + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/providers/openai" + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/providers/registry" + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/services" + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" +) + +type streamUsageReproResult struct { + result domain.GenerateResult + err error +} + +type streamUsageReproMeter struct { + mockUsageMeter + recorded chan streamUsageReproResult +} + +func (m *streamUsageReproMeter) RecordSuccess(_ context.Context, _ string, _ domain.AuthContext, _ string, _ domain.GenerateRequest, result domain.GenerateResult, _ domain.PublicModel, _ int64) error { + m.recorded <- streamUsageReproResult{result: result} + return nil +} + +func (m *streamUsageReproMeter) RecordFailure(_ context.Context, _ *string, _ *domain.AuthContext, _ string, _ *domain.GenerateRequest, _ *domain.PublicModel, err error, _ *domain.Usage, _ int64) error { + m.recorded <- streamUsageReproResult{err: err} + return nil +} + +// End-to-end smoke test through real HTTP, the production OpenAI adapter, +// handler, and usage tracker. Only auth, catalog, and metering are faked. +// No database, credentials, or paid inference is involved. +func TestStreamingUsageRetention(t *testing.T) { + const usage = `{"prompt_tokens":100,"completion_tokens":20,"total_tokens":120}` + const toolDelta = `{"tool_calls":[{"index":0,"id":"call-local","type":"function","function":{"name":"lookup","arguments":"{}"}}]}` + t.Setenv("NEXUS_STREAM_REPRO_PROVIDER_KEY", "local-test-only") + + for _, tc := range []struct { + name string + placement string + closeOnDone bool + closeOnDelta bool + wantUsage bool + }{ + {name: "trailing_usage", placement: "trailing", wantUsage: true}, + {name: "client_cancels_before_finish", placement: "unfinished", closeOnDelta: true}, + } { + t.Run(tc.name, func(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + upstreamCanceled := make(chan struct{}) + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var request struct { + Stream bool `json:"stream"` + StreamOptions struct { + IncludeUsage bool `json:"include_usage"` + } `json:"stream_options"` + } + if err := json.NewDecoder(r.Body).Decode(&request); err != nil || !request.Stream || !request.StreamOptions.IncludeUsage { + t.Error("gateway did not request streaming with include_usage=true") + http.Error(w, "invalid streaming request", http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "text/event-stream") + emit := func(delta, finish, tokenUsage string) { + fmt.Fprintf(w, "data: {\"id\":\"chatcmpl-local-repro\",\"choices\":[{\"index\":0,\"delta\":%s,\"finish_reason\":%s}],\"usage\":%s}\n\n", delta, finish, tokenUsage) + w.(http.Flusher).Flush() + } + emit(toolDelta, "null", "null") + if tc.placement == "unfinished" { + select { + case <-r.Context().Done(): + close(upstreamCanceled) + case <-ctx.Done(): + } + return + } + emit(`{}`, `"tool_calls"`, "null") + // Delay trailing usage so the handler's drain loop has work to do. + delay := time.NewTimer(25 * time.Millisecond) + defer delay.Stop() + select { + case <-delay.C: + fmt.Fprintf(w, "data: {\"id\":\"chatcmpl-local-repro\",\"choices\":[],\"usage\":%s}\n\n", usage) + case <-r.Context().Done(): + close(upstreamCanceled) + return + case <-ctx.Done(): + return + } + io.WriteString(w, "data: [DONE]\n\n") + w.(http.Flusher).Flush() + })) + defer upstream.Close() + + model := domain.PublicModel{ + PublicModelID: "local-repro", ProviderModelID: "local-repro", UpstreamModelName: "local-repro", Active: true, + SupportsChatCompletions: true, SupportsChatCompletionsStream: true, SupportsTools: true, + ProviderConfig: domain.ProviderConfig{ + ProviderName: "novita", BaseURL: upstream.URL, + APIKeySecretRef: "NEXUS_STREAM_REPRO_PROVIDER_KEY", + }, + } + logger := &mockLogger{} + providers := registry.NewRegistry() + providers.Register("novita", openai.NewAdapter(5*time.Second)) + meter := &streamUsageReproMeter{recorded: make(chan streamUsageReproResult, 2)} + generate := services.NewGenerateService(&mockAuthService{}, &mockModelCatalog{model: model}, nil, providers, meter, nil, logger) + handler := handlers.NewChatCompletionsHandler(services.NewChatCompletionsService(generate, logger), logger) + gateway := httptest.NewServer(http.HandlerFunc(handler.Handle)) + defer gateway.Close() + + request, err := http.NewRequestWithContext(ctx, http.MethodPost, gateway.URL+"/v1/chat/completions", strings.NewReader(`{"model":"local-repro","stream":true,"messages":[{"role":"user","content":"Call lookup"}],"tools":[{"type":"function","function":{"name":"lookup","parameters":{"type":"object","properties":{}}}}]}`)) + if err != nil { + t.Fatal(err) + } + request.Header.Set("Authorization", "Bearer local-test-only") + request.Header.Set("Content-Type", "application/json") + response, err := gateway.Client().Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + body, _ := io.ReadAll(response.Body) + t.Fatalf("gateway status=%d body=%s", response.StatusCode, body) + } + scanner := bufio.NewScanner(response.Body) + sawDone := false + sawFinish := false + for scanner.Scan() { + if tc.closeOnDelta && strings.Contains(scanner.Text(), `"tool_calls"`) { + response.Body.Close() + break + } + if strings.Contains(scanner.Text(), `"finish_reason":"tool_calls"`) { + sawFinish = true + } + if scanner.Text() != "data: [DONE]" { + continue + } + sawDone = true + if tc.closeOnDone { + response.Body.Close() + break + } + } + if !tc.closeOnDelta && (!sawDone || !sawFinish) { + t.Fatalf("downstream sawDone=%t read error=%v", sawDone, scanner.Err()) + } + select { + case recorded := <-meter.recorded: + if tc.closeOnDelta { + if recorded.err == nil { + t.Fatal("client cancellation before finish was recorded as success") + } + select { + case <-upstreamCanceled: + case <-ctx.Done(): + t.Fatal("upstream continued generating after early client cancellation") + } + return + } + if recorded.err != nil { + t.Fatalf("recorded failure: %v", recorded.err) + } + result := recorded.result + if result.ID != "chatcmpl-local-repro" || result.FinishReason == nil || *result.FinishReason != "tool_calls" { + t.Fatalf("unexpected completion: %+v", result) + } + if (result.Usage != nil) != tc.wantUsage { + t.Fatalf("usage=%+v, want usage present=%t", result.Usage, tc.wantUsage) + } + if result.Usage != nil && (result.Usage.PromptTokens != 100 || result.Usage.CompletionTokens != 20) { + t.Fatalf("incorrect usage: %+v", result.Usage) + } + select { + case <-upstreamCanceled: + t.Fatal("upstream canceled before trailing usage was sent") + default: + } + t.Logf("success=true finish_reason=tool_calls provider_request_id=%s usage=%+v", result.ID, result.Usage) + case <-ctx.Done(): + t.Fatal("metering did not complete before timeout") + } + }) + } +} diff --git a/apps/gateway/internal/adapters/providers/openai/adapter.go b/apps/gateway/internal/adapters/providers/openai/adapter.go index 8147fc7..63d5ab9 100644 --- a/apps/gateway/internal/adapters/providers/openai/adapter.go +++ b/apps/gateway/internal/adapters/providers/openai/adapter.go @@ -114,7 +114,9 @@ func (a *Adapter) StreamGenerate(ctx context.Context, req domain.GenerateRequest invalidTraceSameBodyRetries := 0 serverOverloadRetries := 0 downgradeRetried := false + attempts := 0 for attempt := 1; ; attempt++ { + attempts = attempt a.logProviderRequest(ctx, model, activeBuilt, attempt, retryReason) streamResp, err = a.client.DoStream(ctx, model.ProviderConfig.BaseURL, apiKey, activeBuilt.Body) if err == nil { @@ -150,7 +152,7 @@ func (a *Adapter) StreamGenerate(ctx context.Context, req domain.GenerateRequest return nil, withProviderPolicyMeta(mapProviderErrorWithCompatibilityContext(err, model, activeBuilt.Body), activeBuilt, attempt, retryReason) } - return NewStream(streamResp, model.ProviderConfig.ProviderName), nil + return NewStream(streamResp, model.ProviderConfig.ProviderName).WithDiagnostics(ctx, a.logger, model, attempts), nil } func missingProviderCredentialError(providerName string) *domain.GatewayError { diff --git a/apps/gateway/internal/adapters/providers/openai/logging.go b/apps/gateway/internal/adapters/providers/openai/logging.go index e98193a..856fbbb 100644 --- a/apps/gateway/internal/adapters/providers/openai/logging.go +++ b/apps/gateway/internal/adapters/providers/openai/logging.go @@ -2,7 +2,6 @@ package openai import ( "context" - "sort" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/middleware" "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" @@ -34,24 +33,14 @@ func (a *Adapter) logProviderRequest(ctx context.Context, model domain.PublicMod } func summarizeProviderBody(body map[string]any) map[string]any { - summary := map[string]any{ - "fields": sortedKeys(body), + // Only fixed keys and numeric/boolean settings: arbitrary strings can contain prompts. + summary := map[string]any{} + for _, key := range []string{"stream", "max_tokens", "max_completion_tokens", "temperature", "top_p", "presence_penalty", "frequency_penalty", "seed", "logprobs", "top_logprobs", "parallel_tool_calls", "store"} { + switch value := body[key].(type) { + case bool, int, int32, int64, float32, float64: + summary[key] = value + } } - copyScalar(summary, body, "model") - copyScalar(summary, body, "stream") - copyScalar(summary, body, "max_tokens") - copyScalar(summary, body, "max_completion_tokens") - copyScalar(summary, body, "temperature") - copyScalar(summary, body, "top_p") - copyScalar(summary, body, "stop") - copyScalar(summary, body, "presence_penalty") - copyScalar(summary, body, "frequency_penalty") - copyScalar(summary, body, "seed") - copyScalar(summary, body, "logprobs") - copyScalar(summary, body, "top_logprobs") - copyScalar(summary, body, "parallel_tool_calls") - copyScalar(summary, body, "store") - copyScalar(summary, body, "service_tier") if _, ok := body["user"]; ok { summary["user"] = "[redacted]" } @@ -62,37 +51,15 @@ func summarizeProviderBody(body map[string]any) map[string]any { } if tools, ok := body["tools"].([]map[string]any); ok { summary["tool_count"] = len(tools) - summary["tool_names"] = summarizeToolNames(tools) - } - if toolChoice, ok := body["tool_choice"]; ok { - summary["tool_choice"] = summarizeToolChoice(toolChoice) - } - if responseFormat, ok := body["response_format"].(map[string]any); ok { - if formatType, ok := responseFormat["type"].(string); ok { - summary["response_format"] = formatType - } } if streamOptions, ok := body["stream_options"].(map[string]any); ok { - summary["stream_options"] = streamOptions + if includeUsage, ok := streamOptions["include_usage"].(bool); ok { + summary["stream_options"] = map[string]any{"include_usage": includeUsage} + } } return summary } -func copyScalar(summary, body map[string]any, key string) { - if v, ok := body[key]; ok { - summary[key] = v - } -} - -func sortedKeys(body map[string]any) []string { - keys := make([]string, 0, len(body)) - for k := range body { - keys = append(keys, k) - } - sort.Strings(keys) - return keys -} - func summarizeMessages(messages []map[string]any) []map[string]any { const maxEdge = 3 total := len(messages) @@ -106,7 +73,12 @@ func summarizeMessages(messages []map[string]any) []map[string]any { } item := map[string]any{} if role, ok := msg["role"].(string); ok { - item["role"] = role + switch role { + case "system", "developer", "user", "assistant", "tool", "function": + item["role"] = role + default: + item["role"] = "unknown" + } } if content, ok := msg["content"].(string); ok { item["content_chars"] = len(content) @@ -115,7 +87,6 @@ func summarizeMessages(messages []map[string]any) []map[string]any { } if toolCalls, ok := msg["tool_calls"].([]map[string]any); ok { item["tool_call_count"] = len(toolCalls) - item["tool_call_names"] = summarizeToolCallNames(toolCalls) } if reasoningContent, ok := msg["reasoning_content"].(string); ok { item["reasoning_content_chars"] = len(reasoningContent) @@ -139,41 +110,3 @@ func totalMessageContentChars(messages []map[string]any) int { } return total } - -func summarizeToolNames(tools []map[string]any) []string { - names := make([]string, 0, len(tools)) - for _, tool := range tools { - fn, _ := tool["function"].(map[string]any) - if name, ok := fn["name"].(string); ok { - names = append(names, name) - } - } - return names -} - -func summarizeToolCallNames(toolCalls []map[string]any) []string { - names := make([]string, 0, len(toolCalls)) - for _, toolCall := range toolCalls { - fn, _ := toolCall["function"].(map[string]any) - if name, ok := fn["name"].(string); ok { - names = append(names, name) - } - } - return names -} - -func summarizeToolChoice(toolChoice any) any { - switch v := toolChoice.(type) { - case string: - return v - case map[string]any: - fn, _ := v["function"].(map[string]any) - name, _ := fn["name"].(string) - return map[string]any{ - "type": v["type"], - "function_name": name, - } - default: - return "[present]" - } -} diff --git a/apps/gateway/internal/adapters/providers/openai/policy_test.go b/apps/gateway/internal/adapters/providers/openai/policy_test.go index f4bdaeb..3d36278 100644 --- a/apps/gateway/internal/adapters/providers/openai/policy_test.go +++ b/apps/gateway/internal/adapters/providers/openai/policy_test.go @@ -576,8 +576,8 @@ func TestSummarizeProviderBodyRedactsPromptTextUserAndToolSchema(t *testing.T) { if !strings.Contains(got, `"content_chars":18`) { t.Fatalf("summary = %s, want content length", got) } - if !strings.Contains(got, `"tool_names":["lookup"]`) { - t.Fatalf("summary = %s, want tool name only", got) + if !strings.Contains(got, `"tool_count":1`) { + t.Fatalf("summary = %s, want tool count only", got) } } @@ -631,3 +631,25 @@ func containsString(values []string, want string) bool { } return false } + +func TestProviderSummaryOmitsFreeformFields(t *testing.T) { + const secret = "PRIVATE-PROMPT-CANARY" + body := map[string]any{ + secret: secret, "model": secret, "stop": []string{secret}, "temperature": secret, + "tool_choice": map[string]any{"function": map[string]any{"name": secret}}, + "response_format": map[string]any{"type": secret}, + "stream_options": map[string]any{"include_usage": true, secret: secret}, + "messages": []map[string]any{{"role": secret, "content": secret, "reasoning_content": secret, "tool_calls": []map[string]any{{"function": map[string]any{"name": secret, "arguments": secret}}}}}, + "tools": []map[string]any{{"function": map[string]any{"name": secret}}}, + } + encoded, err := json.Marshal(summarizeProviderBody(body)) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(encoded), secret) { + t.Fatalf("private data leaked: %s", encoded) + } + if !strings.Contains(string(encoded), `"include_usage":true`) { + t.Fatal("missing usage request metadata") + } +} diff --git a/apps/gateway/internal/adapters/providers/openai/stream.go b/apps/gateway/internal/adapters/providers/openai/stream.go index babeedd..50afd4c 100644 --- a/apps/gateway/internal/adapters/providers/openai/stream.go +++ b/apps/gateway/internal/adapters/providers/openai/stream.go @@ -12,6 +12,7 @@ import ( // Stream reads SSE events from an OpenAI-compatible streaming response. type Stream struct { + diagnostics *streamDiagnostics resp *http.Response scanner *bufio.Scanner done bool @@ -51,24 +52,41 @@ func (s *Stream) Recv() (domain.StreamEvent, error) { } if !strings.HasPrefix(line, "data: ") { + if s.diagnostics != nil && strings.HasPrefix(line, "data:") { + s.diagnostics.unsupported++ + } continue } data := strings.TrimPrefix(line, "data: ") if data == "[DONE]" { + s.diagnostics.end("done_marker", nil) s.done = true return domain.StreamEvent{}, io.EOF } + if s.diagnostics != nil { + s.diagnostics.chunks++ + } var chunk chatCompletionChunk if err := json.Unmarshal([]byte(data), &chunk); err != nil { + if s.diagnostics != nil { + s.diagnostics.malformed++ + } continue } if err := providerBaseResponseError(chunk.BaseResp); err != nil { + if s.diagnostics != nil { + s.diagnostics.observe(chunk, nil) + } + s.diagnostics.end("provider_error", err) return domain.StreamEvent{}, err } events := mapChunkToStreamEvents(chunk, s.includeReasoningContent) + if s.diagnostics != nil { + s.diagnostics.observe(chunk, events) + } if len(events) == 0 { continue } @@ -84,14 +102,17 @@ func (s *Stream) Recv() (domain.StreamEvent, error) { } if err := s.scanner.Err(); err != nil { + s.diagnostics.end("read_error", err) return domain.StreamEvent{}, err } + s.diagnostics.end("eof", nil) s.done = true return domain.StreamEvent{}, io.EOF } func (s *Stream) Close() error { + s.diagnostics.end("closed", nil) s.done = true return s.resp.Body.Close() } @@ -150,7 +171,20 @@ func providerBaseResponseError(resp *providerBaseResponse) error { ) } -func mapChunkToStreamEvents(chunk chatCompletionChunk, includeReasoningContent ...bool) []domain.StreamEvent { +func mapChunkToStreamEvents(chunk chatCompletionChunk, includeReasoningContent ...bool) (events []domain.StreamEvent) { + // Usage is independent of the delta shape. Some compatible providers put + // it on text, role, or tool deltas instead of the final usage-only chunk. + defer func() { + if chunk.Usage == nil { + return + } + if len(events) == 0 { + // Preserve usage even when the chunk has no visible delta. A role + // event carries it to metering without claiming generation is done. + events = []domain.StreamEvent{{Type: domain.StreamEventOutputMessageDelta}} + } + events[0].Usage = chunkUsageToDomain(chunk.Usage) + }() keepReasoning := len(includeReasoningContent) > 0 && includeReasoningContent[0] // Handle usage-only chunk (often last chunk with stream_options.include_usage) if len(chunk.Choices) == 0 && chunk.Usage != nil { diff --git a/apps/gateway/internal/adapters/providers/openai/stream_diagnostics.go b/apps/gateway/internal/adapters/providers/openai/stream_diagnostics.go new file mode 100644 index 0000000..cc6f3ac --- /dev/null +++ b/apps/gateway/internal/adapters/providers/openai/stream_diagnostics.go @@ -0,0 +1,95 @@ +package openai + +import ( + "context" + "errors" + "io" + + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/middleware" + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" + "github.com/dappnode/dappnode-nexus-gateway/pkg/observability/logfields" +) + +// Diagnostics contain metadata only: never SSE bodies, content, or arguments. +type streamDiagnostics struct { + ctx context.Context + logger ports.Logger + fields []any + providerID string + finishReason *string + chunks, malformed, unsupported, usageChunks, mappedUsageChunks int + usage *domain.Usage + logged bool +} + +func (s *Stream) WithDiagnostics(ctx context.Context, logger ports.Logger, model domain.PublicModel, attempt int) *Stream { + if logger != nil { + s.diagnostics = &streamDiagnostics{ctx: ctx, logger: logger, fields: []any{ + "request_id", middleware.GetRequestID(ctx), "provider", model.ProviderConfig.ProviderName, + "provider_model", model.UpstreamModelName, "model", model.PublicModelID, "attempt", attempt, + }} + } + return s +} + +func (d *streamDiagnostics) observe(chunk chatCompletionChunk, events []domain.StreamEvent) { + if chunk.ID != "" { + d.providerID = chunk.ID + } + hasTools := false + var finish *string + for _, choice := range chunk.Choices { + if choice.FinishReason != nil { + finish = choice.FinishReason + d.finishReason = finish + } + hasTools = hasTools || len(choice.Delta.ToolCalls) > 0 + } + mapped := false + for _, event := range events { + mapped = mapped || event.Usage != nil + } + if chunk.Usage != nil { + d.usageChunks++ + d.usage = chunkUsageToDomain(chunk.Usage) + } + if mapped { + d.mappedUsageChunks++ + } + if chunk.Usage != nil || finish != nil { + fields := append([]any{}, d.fields...) + fields = append(fields, "provider_request_id", d.providerID, "chunk_index", d.chunks, + "finish_reason", logfields.FinishReason(finish), "has_tool_calls", hasTools, "choices_count", len(chunk.Choices), + "usage_received", chunk.Usage != nil, "usage_forwarded", mapped, "usage", chunkUsageToDomain(chunk.Usage)) + d.logger.Debug("provider stream chunk metadata", fields...) + } +} + +func (d *streamDiagnostics) end(reason string, err error) { + if d == nil || d.logged { + return + } + d.logged = true + if errors.Is(err, context.Canceled) || errors.Is(d.ctx.Err(), context.Canceled) { + if reason == "read_error" { + reason = "canceled" + } + } else if errors.Is(err, context.DeadlineExceeded) || errors.Is(d.ctx.Err(), context.DeadlineExceeded) { + if reason == "read_error" { + reason = "deadline_exceeded" + } + } + fields := append([]any{}, d.fields...) + fields = append(fields, "provider_request_id", d.providerID, "upstream_end", reason, + "upstream_done_seen", reason == "done_marker", "chunks_received", d.chunks, + "malformed_chunks", d.malformed, "unsupported_data_lines", d.unsupported, + "usage_chunks", d.usageChunks, "mapped_usage_chunks", d.mappedUsageChunks, + "usage_received", d.usage != nil, "usage", d.usage, "finish_reason", logfields.FinishReason(d.finishReason), + "context_canceled", errors.Is(d.ctx.Err(), context.Canceled)) + if d.usage == nil || d.malformed > 0 || d.unsupported > 0 || (err != nil && err != io.EOF) || reason == "closed" { + d.logger.Warn("provider stream ended", fields...) + } else { + d.logger.Info("provider stream ended", fields...) + } +} diff --git a/apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go b/apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go new file mode 100644 index 0000000..bf8e9ea --- /dev/null +++ b/apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go @@ -0,0 +1,132 @@ +package openai + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "testing" + + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/middleware" + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" +) + +type diagnosticLog struct { + level, message string + fields map[string]any +} +type diagnosticLogger struct{ entries []diagnosticLog } + +func (l *diagnosticLogger) add(level, message string, fields ...any) { + m := map[string]any{} + for i := 0; i+1 < len(fields); i += 2 { + m[fields[i].(string)] = fields[i+1] + } + l.entries = append(l.entries, diagnosticLog{level, message, m}) +} +func (l *diagnosticLogger) Debug(m string, f ...any) { l.add("debug", m, f...) } +func (l *diagnosticLogger) Info(m string, f ...any) { l.add("info", m, f...) } +func (l *diagnosticLogger) Warn(m string, f ...any) { l.add("warn", m, f...) } +func (l *diagnosticLogger) Error(m string, f ...any) { l.add("error", m, f...) } + +type terminalErrorReader struct{ err error } + +func (r terminalErrorReader) Read([]byte) (int, error) { return 0, r.err } + +func TestStreamDiagnostics_TerminationAndUsage(t *testing.T) { + finish := `data: {"id":"provider-id","choices":[{"delta":{"content":"SECRET-CONTENT"},"finish_reason":"tool_calls"}]}` + "\n\n" + usage := `data: {"id":"provider-id","choices":[],"usage":{"prompt_tokens":100,"completion_tokens":20,"total_tokens":120}}` + "\n\n" + for _, tc := range []struct { + name, data, end, level string + usage bool + malformed, unsupported int + readErr error + }{ + {name: "trailing usage", data: finish + usage + "data: [DONE]\n", end: "done_marker", level: "info", usage: true}, + {name: "provider omission", data: finish + "data: [DONE]\n", end: "done_marker", level: "warn"}, + {name: "bare eof", data: finish, end: "eof", level: "warn"}, + {name: "malformed", data: finish + "data: {SECRET-MALFORMED}\n" + usage + "data: [DONE]\n", end: "done_marker", level: "warn", usage: true, malformed: 1}, + {name: "unsupported framing", data: finish + "data:{SECRET-FRAMING}\n", end: "eof", level: "warn", unsupported: 1}, + {name: "canceled", data: finish, end: "canceled", level: "warn", readErr: context.Canceled}, + {name: "timeout", data: finish, end: "deadline_exceeded", level: "warn", readErr: context.DeadlineExceeded}, + } { + t.Run(tc.name, func(t *testing.T) { + logger := &diagnosticLogger{} + reader := io.Reader(strings.NewReader(tc.data)) + if tc.readErr != nil { + reader = io.MultiReader(reader, terminalErrorReader{tc.readErr}) + } + ctx := context.WithValue(context.Background(), middleware.RequestIDKey, "gateway-id") + stream := NewStream(&http.Response{Body: io.NopCloser(reader)}).WithDiagnostics(ctx, logger, domain.PublicModel{PublicModelID: "model", ProviderConfig: domain.ProviderConfig{ProviderName: "novita"}}, 1) + for { + if _, err := stream.Recv(); err != nil { + break + } + } + stream.Close() + stream.Close() + summaries := 0 + for _, entry := range logger.entries { + if strings.Contains(fmt.Sprint(entry.fields), "SECRET-") { + t.Fatal("stream content was logged") + } + if entry.message != "provider stream ended" { + continue + } + summaries++ + if entry.level != tc.level || entry.fields["upstream_end"] != tc.end || entry.fields["usage_received"] != tc.usage || entry.fields["malformed_chunks"] != tc.malformed || entry.fields["unsupported_data_lines"] != tc.unsupported { + t.Fatalf("unexpected summary: %+v", entry) + } + if entry.fields["request_id"] != "gateway-id" || entry.fields["provider_request_id"] != "provider-id" { + t.Fatalf("missing correlation: %+v", entry.fields) + } + if tc.usage && (entry.fields["usage_chunks"] != 1 || entry.fields["mapped_usage_chunks"] != 1) { + t.Fatalf("usage counters: %+v", entry.fields) + } + } + if summaries != 1 { + t.Fatalf("summaries=%d want 1", summaries) + } + }) + } +} + +func TestStreamDiagnostics_UsageOnToolDelta(t *testing.T) { + logger := &diagnosticLogger{} + stream := newTestStream(`data: {"id":"provider-id","choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"SECRET-ARGUMENTS"}}]},"finish_reason":null}],"usage":{"prompt_tokens":3,"completion_tokens":2}}`+"\n\ndata: [DONE]\n").WithDiagnostics(context.Background(), logger, domain.PublicModel{}, 1) + for { + if _, err := stream.Recv(); err != nil { + break + } + } + for _, entry := range logger.entries { + encoded, _ := json.Marshal(entry.fields) + if strings.Contains(string(encoded), "SECRET-ARGUMENTS") { + t.Fatal("tool arguments leaked") + } + if entry.message == "provider stream chunk metadata" && (entry.fields["usage_forwarded"] != true || entry.fields["has_tool_calls"] != true) { + t.Fatalf("unexpected chunk metadata: %+v", entry) + } + } +} + +func TestStreamDiagnosticsUnknownFinishReasonIsRedacted(t *testing.T) { + logger := &diagnosticLogger{} + stream := newTestStream(`data: {"id":"provider-id","choices":[{"delta":{},"finish_reason":"PRIVATE-PROMPT-CANARY"}]}`+"\n\ndata: [DONE]\n").WithDiagnostics(context.Background(), logger, domain.PublicModel{}, 1) + for { + if _, err := stream.Recv(); err != nil { + break + } + } + for _, entry := range logger.entries { + encoded, err := json.Marshal(entry.fields) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(encoded), "PRIVATE-PROMPT-CANARY") { + t.Fatal("finish reason leaked private content") + } + } +} diff --git a/apps/gateway/internal/adapters/providers/openai/stream_test.go b/apps/gateway/internal/adapters/providers/openai/stream_test.go index 6dbad7c..061b28b 100644 --- a/apps/gateway/internal/adapters/providers/openai/stream_test.go +++ b/apps/gateway/internal/adapters/providers/openai/stream_test.go @@ -1,6 +1,7 @@ package openai import ( + "encoding/json" "io" "net/http" "strings" @@ -9,6 +10,30 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) +func TestStream_PreservesUsageOnEveryDeltaShape(t *testing.T) { + for name, delta := range map[string]string{ + "text": `{"content":"hello"}`, + "role": `{"role":"assistant"}`, + "tool": `{"tool_calls":[{"index":0,"function":{"arguments":"{}"}}]}`, + "hidden reasoning": `{"reasoning_content":"thinking"}`, + "empty": `{}`, + } { + t.Run(name, func(t *testing.T) { + var chunk chatCompletionChunk + if err := json.Unmarshal([]byte(`{"choices":[{"delta":`+delta+`,"finish_reason":null}],"usage":{"prompt_tokens":100,"completion_tokens":20,"total_tokens":120}}`), &chunk); err != nil { + t.Fatal(err) + } + events := mapChunkToStreamEvents(chunk) + if len(events) == 0 || events[0].Usage == nil || events[0].Usage.PromptTokens != 100 || events[0].Usage.CompletionTokens != 20 { + t.Fatalf("usage lost: %+v", events) + } + if events[0].Type == domain.StreamEventCompleted { + t.Fatal("usage on a non-final delta must not signal completion") + } + }) + } +} + type fakeBody struct { *strings.Reader } diff --git a/apps/gateway/internal/adapters/providers/tinfoil/adapter.go b/apps/gateway/internal/adapters/providers/tinfoil/adapter.go index d15b5ee..f927d84 100644 --- a/apps/gateway/internal/adapters/providers/tinfoil/adapter.go +++ b/apps/gateway/internal/adapters/providers/tinfoil/adapter.go @@ -109,7 +109,7 @@ func (a *Adapter) StreamGenerate(ctx context.Context, req domain.GenerateRequest } return &Stream{ - inner: openai.NewStream(resp, providerName), + inner: openai.NewStream(resp, providerName).WithDiagnostics(ctx, a.logger, model, 1), proof: proof, }, nil } diff --git a/apps/gateway/internal/application/services/generate_service.go b/apps/gateway/internal/application/services/generate_service.go index b5aa472..ed11777 100644 --- a/apps/gateway/internal/application/services/generate_service.go +++ b/apps/gateway/internal/application/services/generate_service.go @@ -3,6 +3,7 @@ package services import ( "context" "errors" + "fmt" "io" "strconv" "strings" @@ -12,6 +13,7 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/observability/metrics" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" + "github.com/dappnode/dappnode-nexus-gateway/pkg/observability/logfields" "github.com/google/uuid" ) @@ -122,7 +124,7 @@ func (s *GenerateService) Execute(ctx context.Context, endpoint string, req doma "model", execReq.PublicModelID, "provider", executionModel.ProviderConfig.ProviderName, "latency_ms", latencyMs, - "finish_reason", result.FinishReason, + "finish_reason", logfields.FinishReason(result.FinishReason), "gateway_status", 200, "upstream_status", 200, ) @@ -135,7 +137,7 @@ func (s *GenerateService) Execute(ctx context.Context, endpoint string, req doma "request_id", requestID, "reservation_id", reservationID, "account_id", authCtx.Account.ID, - "error", err, + "error_type", fmt.Sprintf("%T", err), ) } @@ -321,7 +323,7 @@ func (s *GenerateService) logFallback(requestID string, primary, fallback domain "provider_model", primary.UpstreamModelName, "fallback_provider", fallback.ProviderConfig.ProviderName, "fallback_provider_model", fallback.UpstreamModelName, - "error", err, + "error_type", fmt.Sprintf("%T", err), ) } @@ -509,18 +511,28 @@ type usageTrackingStream struct { finishReason *string providerResponseID string finished bool + endReason string } func (s *usageTrackingStream) Recv() (domain.StreamEvent, error) { event, err := s.inner.Recv() if err != nil { + s.endReason = "read_error" + if errors.Is(err, context.Canceled) { + s.endReason = "canceled" + } + if errors.Is(err, context.DeadlineExceeded) { + s.endReason = "deadline_exceeded" + } if err == io.EOF { + s.endReason = "eof" s.recordCompletion(nil) return event, err } // If the model already sent a finish reason, treat post-completion // errors (e.g. context canceled after client disconnect) as success. if s.finishReason != nil { + s.service.logger.Warn("stream interrupted after finish", s.logFields()...) s.recordCompletion(nil) return event, io.EOF } @@ -542,6 +554,7 @@ func (s *usageTrackingStream) Recv() (domain.StreamEvent, error) { } if event.Type == domain.StreamEventError { + s.endReason = "provider_error_event" var gwErr error if event.Error != nil { gwErr = event.Error @@ -554,6 +567,7 @@ func (s *usageTrackingStream) Recv() (domain.StreamEvent, error) { func (s *usageTrackingStream) Close() error { if !s.finished { + s.endReason = "closed_before_eof" s.recordCompletion(context.Canceled) } return s.inner.Close() @@ -569,6 +583,11 @@ func (s *usageTrackingStream) recordCompletion(err error) { if err != nil { err = sanitizeErrorWithPIIMapping(err, s.piiMapping) fields := s.service.buildErrorLogFields(ctx, s.requestID, &s.authCtx, s.endpoint, s.req.PublicModelID, s.model, err, latencyMs) + fields = append(fields, + "reservation_id", s.reservationID, "provider_request_id", s.providerResponseID, + "stream", true, "stream_end", s.endReason, "finish_reason", logfields.FinishReason(s.finishReason), + "usage_received", s.lastUsage != nil, "usage", s.lastUsage, + "context_canceled", errors.Is(s.ctx.Err(), context.Canceled)) s.service.logger.Error("stream error", fields...) s.service.recordGeneration(metrics.OutcomeError, s.endpoint, s.req, s.model, latencyMs) s.service.recordFailure(ctx, &s.reservationID, &s.authCtx, s.endpoint, &s.req, &s.model, err, s.lastUsage, latencyMs) @@ -588,25 +607,31 @@ func (s *usageTrackingStream) recordCompletion(err error) { } s.service.storeTinfoilProof(ctx, s.authCtx, s.model, result) metrics.RecordUsage(s.lastUsage, s.req.PublicModelID, s.model.ProviderConfig.ProviderName) + if s.lastUsage == nil { + s.service.logger.Warn("stream completed without usage", s.logFields()...) + } s.service.recordGeneration(metrics.OutcomeSuccess, s.endpoint, s.req, s.model, latencyMs) - s.service.logger.Info("generation completed", - "request_id", s.requestID, - "account_id", s.authCtx.Account.ID, - "endpoint", s.endpoint, - "model", s.req.PublicModelID, - "provider", s.model.ProviderConfig.ProviderName, - "latency_ms", latencyMs, - "finish_reason", s.finishReason, - "gateway_status", 200, - "upstream_status", 200, - ) + fields := append(s.logFields(), "gateway_status", 200, "upstream_status", 200) + s.service.logger.Info("generation completed", fields...) if recErr := s.service.metering.RecordSuccess(ctx, s.reservationID, s.authCtx, s.endpoint, s.req, result, s.model, latencyMs); recErr != nil { - s.service.logger.Error("failed to record stream usage", - "request_id", s.requestID, - "reservation_id", s.reservationID, - "account_id", s.authCtx.Account.ID, - "error", recErr, - ) + fields = append(s.logFields(), "metering_status", "failed", "error_type", fmt.Sprintf("%T", recErr)) + s.service.logger.Error("failed to record stream usage", fields...) + } else { + fields = append(s.logFields(), "metering_status", "accepted") + s.service.logger.Info("stream metering completed", fields...) + } +} + +// Keep unknown usage as null, rather than presenting it as a zero-token request. +func (s *usageTrackingStream) logFields() []any { + return []any{ + "request_id", s.requestID, "reservation_id", s.reservationID, "provider_request_id", s.providerResponseID, + "account_id", s.authCtx.Account.ID, "endpoint", s.endpoint, "model", s.req.PublicModelID, + "provider", s.model.ProviderConfig.ProviderName, "provider_model", s.model.UpstreamModelName, + "stream", true, "stream_end", s.endReason, "finish_reason", logfields.FinishReason(s.finishReason), + "usage_received", s.lastUsage != nil, "usage", s.lastUsage, + "context_canceled", errors.Is(s.ctx.Err(), context.Canceled), + "latency_ms", time.Since(s.start).Milliseconds(), } } @@ -620,7 +645,7 @@ func (s *GenerateService) recordFailure(ctx context.Context, reservationID *stri ctx = context.WithoutCancel(ctx) } if recErr := s.metering.RecordFailure(ctx, reservationID, auth, endpoint, req, model, err, partialUsage, latencyMs); recErr != nil && s.logger != nil { - fields := []any{"error", recErr} + fields := []any{"error_type", fmt.Sprintf("%T", recErr)} if auth != nil { fields = append(fields, "account_id", auth.Account.ID) } @@ -660,7 +685,7 @@ func (s *GenerateService) recordTerminalOutcome(outcome, endpoint string, req do } // buildErrorLogFields builds a structured log field slice for error conditions, -// including request context, provider details, and any upstream error metadata. +// including request context, provider details, and numeric upstream status, never error messages or arbitrary metadata. func (s *GenerateService) buildErrorLogFields(ctx context.Context, requestID string, authCtx *domain.AuthContext, endpoint, publicModelID string, model domain.PublicModel, err error, latencyMs int64) []any { fields := []any{ "request_id", requestID, @@ -669,7 +694,7 @@ func (s *GenerateService) buildErrorLogFields(ctx context.Context, requestID str "provider", model.ProviderConfig.ProviderName, "provider_model", model.UpstreamModelName, "latency_ms", latencyMs, - "error", err.Error(), + "error_type", fmt.Sprintf("%T", err), } if authCtx != nil { fields = append(fields, "account_id", authCtx.Account.ID) @@ -679,7 +704,10 @@ func (s *GenerateService) buildErrorLogFields(ctx context.Context, requestID str if errors.As(err, &gwErr) { fields = append(fields, "error_code", gwErr.Code) fields = append(fields, "gateway_status", gwErr.HTTPStatus) - fields = append(fields, gwErr.LogFields()...) + // Upstream messages and arbitrary metadata may echo request contents. + if status, ok := gwErr.Metadata["upstream_status"].(int); ok { + fields = append(fields, "upstream_status", status) + } } return fields } @@ -724,7 +752,7 @@ func (s *GenerateService) storeTinfoilProof(ctx context.Context, auth domain.Aut "provider", model.ProviderConfig.ProviderName, "model", model.PublicModelID, "provider_response_id", proof.ProviderResponseID, - "error", err, + "error_type", fmt.Sprintf("%T", err), ) } } diff --git a/apps/gateway/internal/application/services/stream_logging_test.go b/apps/gateway/internal/application/services/stream_logging_test.go new file mode 100644 index 0000000..d071104 --- /dev/null +++ b/apps/gateway/internal/application/services/stream_logging_test.go @@ -0,0 +1,119 @@ +package services + +import ( + "context" + "encoding/json" + "errors" + "io" + "strings" + "testing" + "time" + + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" +) + +type streamLogEntry struct { + level, message string + fields map[string]any +} +type streamLogRecorder struct{ entries []streamLogEntry } + +func (l *streamLogRecorder) add(level, message string, fields ...any) { + m := map[string]any{} + for i := 0; i+1 < len(fields); i += 2 { + m[fields[i].(string)] = fields[i+1] + } + l.entries = append(l.entries, streamLogEntry{level, message, m}) +} +func (l *streamLogRecorder) Debug(m string, f ...any) { l.add("debug", m, f...) } +func (l *streamLogRecorder) Info(m string, f ...any) { l.add("info", m, f...) } +func (l *streamLogRecorder) Warn(m string, f ...any) { l.add("warn", m, f...) } +func (l *streamLogRecorder) Error(m string, f ...any) { l.add("error", m, f...) } + +type failingLogMeter struct { + stubUsageMeter + err error +} + +func (m *failingLogMeter) RecordSuccess(ctx context.Context, id string, a domain.AuthContext, endpoint string, req domain.GenerateRequest, result domain.GenerateResult, model domain.PublicModel, latency int64) error { + m.stubUsageMeter.RecordSuccess(ctx, id, a, endpoint, req, result, model, latency) + return m.err +} + +func TestStreamLogging_UsageAndMeteringOutcomes(t *testing.T) { + for _, tc := range []struct { + name string + usage *domain.Usage + readErr, meterErr error + beforeFinish bool + wantMessage, level string + }{ + {name: "missing", wantMessage: "stream completed without usage", level: "warn"}, + {name: "known", usage: &domain.Usage{PromptTokens: 100, CompletionTokens: 20}, wantMessage: "stream metering completed", level: "info"}, + {name: "canceled after finish", readErr: context.Canceled, wantMessage: "stream interrupted after finish", level: "warn"}, + {name: "metering failed", usage: &domain.Usage{PromptTokens: 100}, meterErr: errors.New("metering unavailable"), wantMessage: "failed to record stream usage", level: "error"}, + {name: "failed before finish", readErr: context.Canceled, beforeFinish: true, wantMessage: "stream error", level: "error"}, + } { + t.Run(tc.name, func(t *testing.T) { + log := &streamLogRecorder{} + meter := &failingLogMeter{err: tc.meterErr} + finish := "tool_calls" + event := domain.StreamEvent{Type: domain.StreamEventCompleted, FinishReason: &finish, Usage: tc.usage, ProviderResponseID: "provider-id"} + if tc.beforeFinish { + event.Type = domain.StreamEventOutputTextDelta + event.FinishReason = nil + } + steps := []streamStep{{event: event}} + if tc.readErr != nil { + steps = append(steps, streamStep{err: tc.readErr}) + } + stream := &usageTrackingStream{inner: &stubStream{steps: steps}, service: &GenerateService{logger: log, metering: meter}, ctx: context.Background(), requestID: "gateway-id", reservationID: "reservation-id", start: time.Now(), req: domain.GenerateRequest{PublicModelID: "model"}, model: domain.PublicModel{ProviderConfig: domain.ProviderConfig{ProviderName: "novita"}}} + for { + if _, err := stream.Recv(); err != nil { + if tc.readErr == nil && err != io.EOF { + t.Fatal(err) + } + break + } + } + stream.Close() + found := 0 + for _, entry := range log.entries { + if entry.message != tc.wantMessage { + continue + } + found++ + if entry.level != tc.level || entry.fields["request_id"] != "gateway-id" || entry.fields["reservation_id"] != "reservation-id" || entry.fields["provider_request_id"] != "provider-id" || entry.fields["usage_received"] != (tc.usage != nil) { + t.Fatalf("incomplete log: %+v", entry) + } + } + if found != 1 { + t.Fatalf("found %d logs for %q: %+v", found, tc.wantMessage, log.entries) + } + if tc.beforeFinish { + if meter.failureCalls != 1 || meter.successCalls != 0 { + t.Fatal("logging changed failure accounting") + } + } else if meter.successCalls != 1 { + t.Fatal("logging changed completion accounting") + } + }) + } +} + +func TestGenerationErrorLogsOmitProviderEcho(t *testing.T) { + const secret = "PRIVATE-PROMPT-CANARY" + err := domain.ErrProviderError(502, secret).WithMeta("upstream_error", secret, "upstream_status", 400, secret, secret) + service := &GenerateService{} + fields := service.buildErrorLogFields(context.Background(), "request-id", nil, "chat", "model", domain.PublicModel{}, err, 10) + encoded, marshalErr := json.Marshal(fields) + if marshalErr != nil { + t.Fatal(marshalErr) + } + if strings.Contains(string(encoded), secret) { + t.Fatalf("private data leaked: %s", encoded) + } + if !strings.Contains(string(encoded), "upstream_status") { + t.Fatal("missing upstream status") + } +} diff --git a/pkg/observability/logfields/finish.go b/pkg/observability/logfields/finish.go new file mode 100644 index 0000000..1341eaf --- /dev/null +++ b/pkg/observability/logfields/finish.go @@ -0,0 +1,16 @@ +package logfields + +// FinishReason allows protocol values only; upstream strings are untrusted. +// This sanitizes logs without changing the response or stored usage event. +func FinishReason(reason *string) *string { + if reason == nil { + return nil + } + switch *reason { + case "stop", "length", "tool_calls", "function_call", "content_filter": + return reason + default: + unknown := "unknown" + return &unknown + } +} diff --git a/pkg/observability/logger/zap.go b/pkg/observability/logger/zap.go index ac2a60b..41371d6 100644 --- a/pkg/observability/logger/zap.go +++ b/pkg/observability/logger/zap.go @@ -14,23 +14,7 @@ type ZapLogger struct { } func NewZapLogger(level string) (*ZapLogger, error) { - cfg := zap.NewProductionConfig() - cfg.Level = zap.NewAtomicLevel() - verbose := false - - switch level { - case "debug": - cfg.Level.SetLevel(zap.DebugLevel) - verbose = true - case "info": - cfg.Level.SetLevel(zap.InfoLevel) - case "warn": - cfg.Level.SetLevel(zap.WarnLevel) - case "error": - cfg.Level.SetLevel(zap.ErrorLevel) - default: - cfg.Level.SetLevel(zap.InfoLevel) - } + cfg, verbose := requestLoggingConfig(level) l, err := cfg.Build(zap.AddCallerSkip(1)) if err != nil { @@ -74,3 +58,28 @@ func (l *ZapLogger) Verbose() bool { func (l *ZapLogger) Sync() { l.sugar.Sync() } + +func requestLoggingConfig(level string) (zap.Config, bool) { + cfg := zap.NewProductionConfig() + // Request diagnostics are accounting evidence. Zap's production sampler + // groups by message, so distinct request IDs can otherwise be dropped. + cfg.Sampling = nil + cfg.Level = zap.NewAtomicLevel() + verbose := false + + switch level { + case "debug": + cfg.Level.SetLevel(zap.DebugLevel) + verbose = true + case "info": + cfg.Level.SetLevel(zap.InfoLevel) + case "warn": + cfg.Level.SetLevel(zap.WarnLevel) + case "error": + cfg.Level.SetLevel(zap.ErrorLevel) + default: + cfg.Level.SetLevel(zap.InfoLevel) + } + + return cfg, verbose +} diff --git a/pkg/observability/logger/zap_test.go b/pkg/observability/logger/zap_test.go new file mode 100644 index 0000000..bdf6ff7 --- /dev/null +++ b/pkg/observability/logger/zap_test.go @@ -0,0 +1,31 @@ +package logger + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestRequestWarningsAreNotSampledDuringBurst(t *testing.T) { + path := filepath.Join(t.TempDir(), "requests.jsonl") + cfg, _ := requestLoggingConfig("info") + cfg.OutputPaths = []string{path} + log, err := cfg.Build() + if err != nil { + t.Fatal(err) + } + for i := range 250 { + log.Sugar().Warnw("stream completed without usage", "request_id", i) + } + if err = log.Sync(); err != nil { + t.Fatal(err) + } + data, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + if got := strings.Count(string(data), "stream completed without usage"); got != 250 { + t.Fatalf("retained %d request warnings, want 250", got) + } +}