From defc841eb57d3b437923bc2d8593a056ecb27e21 Mon Sep 17 00:00:00 2001 From: Marketen Date: Mon, 28 Sep 2026 23:04:31 +0200 Subject: [PATCH 1/7] fix: never drop or corrupt streamed tool calls The OpenAI-compatible stream parser forwarded only the first tool call of each chunk, so providers that pack parallel calls lost the rest (clients saw "malformed tool call"). The handler also discarded output after the first finish_reason, turned mid-stream provider errors into a clean [DONE], and data: lines without a space were ignored. Normalize every stream to the OpenAI shape: all calls forwarded with contiguous indexes, one id and name per call (synthesized id if missing), "{}" for calls without arguments, text kept alongside calls, a finish reason always present, usage-only chunks no longer ending the stream, and empty or error streams reported as provider errors so fallback runs. Co-Authored-By: Claude Opus 5.5 --- .../http/handlers/chat_completions_handler.go | 58 +++ .../handlers/chat_completions_stream_test.go | 38 ++ .../adapters/providers/openai/stream.go | 355 ++++++++++++++---- .../openai/stream_diagnostics_test.go | 3 +- .../adapters/providers/openai/stream_test.go | 14 +- .../providers/openai/stream_toolcalls_test.go | 172 +++++++++ 6 files changed, 552 insertions(+), 88 deletions(-) create mode 100644 apps/gateway/internal/adapters/providers/openai/stream_toolcalls_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 ad078bf..437d061 100644 --- a/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go +++ b/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go @@ -144,12 +144,43 @@ func (h *ChatCompletionsHandler) writeStream( if event.ProviderResponseID != "" { responseID = event.ProviderResponseID } + if event.Type == domain.StreamEventError { + // A provider error after output started must not look like a + // complete response: no finish chunk and no [DONE]. + message := "the provider stream failed" + if event.Error != nil { + message = event.Error.Error() + } + writeMu.Lock() + sw.WriteData(map[string]any{ + "error": map[string]any{"type": "provider_error", "message": message}, + }) + writeMu.Unlock() + return + } + write := func(event domain.StreamEvent) { + chunk, _ := mapper.DomainStreamEventToChatChunk(event, responseModelID, responseID, createdAt) + if chunk == nil { + return + } + writeMu.Lock() + 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) + } + } // 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 { + // Output the finishing chunk carried goes out now; only the finish + // signal waits. + for _, delta := range splitDeltas(&event) { + write(delta) + } requestID := middleware.GetRequestID(r.Context()) providerID := responseID timer := time.AfterFunc(streamUsageDrainTimeout, func() { @@ -167,6 +198,10 @@ func (h *ChatCompletionsHandler) writeStream( if tail.FinishReason != nil { event.FinishReason = tail.FinishReason } + // Output after the finish reason is still the response. + for _, delta := range splitDeltas(&tail) { + write(delta) + } } timer.Stop() } @@ -191,3 +226,26 @@ func (h *ChatCompletionsHandler) writeStream( } } } + +// splitDeltas moves the text, reasoning, and tool-call output out of an event +// into delta events of their own, leaving the rest (finish, usage) in place. +func splitDeltas(event *domain.StreamEvent) []domain.StreamEvent { + var out []domain.StreamEvent + content := event.ContentDelta != nil && *event.ContentDelta != "" + reasoning := event.ReasoningDelta != nil && *event.ReasoningDelta != "" + if content || reasoning { + delta := domain.StreamEvent{Type: domain.StreamEventOutputTextDelta} + if content { + delta.ContentDelta = event.ContentDelta + } + if reasoning { + delta.ReasoningDelta = event.ReasoningDelta + } + out = append(out, delta) + } + if event.ToolCallDelta != nil { + out = append(out, domain.StreamEvent{Type: domain.StreamEventToolCallDelta, ToolCallDelta: event.ToolCallDelta}) + } + event.ContentDelta, event.ReasoningDelta, event.ToolCallDelta = nil, nil, nil + return out +} 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 72726d3..321a41c 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 @@ -158,3 +158,41 @@ type finalizationLogRecorder struct { } func (l *finalizationLogRecorder) Warn(message string, _ ...any) { l.warnings <- message } + +func TestChatCompletionsHandler_ForwardsOutputAroundTheFinishReason(t *testing.T) { + finish, early, late, args := "tool_calls", "Reading.", "late text", "{}" + id, name := "call_1", "read" + stream := &handlerTestStream{events: []domain.StreamEvent{ + {Type: domain.StreamEventCompleted, FinishReason: &finish, ContentDelta: &early}, + {Type: domain.StreamEventToolCallDelta, ToolCallDelta: &domain.ToolCallDelta{Index: 0, ID: &id, Name: &name, ArgumentsDelta: &args}}, + {Type: domain.StreamEventOutputTextDelta, ContentDelta: &late}, + }} + recorder := httptest.NewRecorder() + handler := &ChatCompletionsHandler{logger: &confidentialTestLogger{}} + handler.writeStream(recorder, httptest.NewRequest("POST", "/v1/chat/completions", nil), domain.GenerateRequest{PublicModelID: "test"}, stream, nil, func() {}) + body := recorder.Body.String() + order := []string{`"content":"Reading."`, `"id":"call_1"`, `"content":"late text"`, `"finish_reason":"tool_calls"`, "[DONE]"} + at := 0 + for _, want := range order { + i := strings.Index(body[at:], want) + if i < 0 { + t.Fatalf("missing %s in order: %s", want, body) + } + at += i + } +} + +func TestChatCompletionsHandler_ProviderErrorEventIsNotCompletion(t *testing.T) { + content := "partial" + stream := &handlerTestStream{events: []domain.StreamEvent{ + {Type: domain.StreamEventOutputTextDelta, ContentDelta: &content}, + {Type: domain.StreamEventError, Error: domain.ErrProviderError(502, "overloaded")}, + }} + recorder := httptest.NewRecorder() + handler := &ChatCompletionsHandler{logger: &confidentialTestLogger{}} + handler.writeStream(recorder, httptest.NewRequest("POST", "/v1/chat/completions", nil), domain.GenerateRequest{PublicModelID: "test"}, stream, nil, func() {}) + body := recorder.Body.String() + if strings.Contains(body, "[DONE]") || strings.Contains(body, `"finish_reason":"`) || !strings.Contains(body, `"error"`) { + t.Fatalf("error event reported as completion: %s", body) + } +} diff --git a/apps/gateway/internal/adapters/providers/openai/stream.go b/apps/gateway/internal/adapters/providers/openai/stream.go index 50afd4c..cd99a60 100644 --- a/apps/gateway/internal/adapters/providers/openai/stream.go +++ b/apps/gateway/internal/adapters/providers/openai/stream.go @@ -2,6 +2,8 @@ package openai import ( "bufio" + "crypto/rand" + "encoding/hex" "encoding/json" "io" "net/http" @@ -10,19 +12,38 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) -// Stream reads SSE events from an OpenAI-compatible streaming response. +// maxSSELine bounds one SSE line. Providers that do not stream tool +// arguments send a whole call (for example a large file) in one line. +const maxSSELine = 16 << 20 + +// Stream reads SSE events from an OpenAI-compatible streaming response and +// normalizes them, so every client sees the same well-formed shape whatever +// the provider sends: +// +// - every tool call in a chunk is forwarded, not only the first; +// - each tool call gets contiguous indexes, one id (synthesized if the +// provider omits it), and its id and name only once; +// - a call that received no arguments gets "{}"; +// - text in a chunk that also carries tool calls is kept; +// - a usage-only chunk before the finish reason does not end the stream; +// - a stream that ends without a finish reason gets one, while a stream with +// no chunk at all or an error payload is reported as a provider error. type Stream struct { diagnostics *streamDiagnostics resp *http.Response scanner *bufio.Scanner - done bool + done bool // upstream fully read includeReasoningContent bool - deferredCompleted *domain.StreamEvent // stashed when tool-call delta + finish_reason arrive in one chunk + pending []domain.StreamEvent + tools toolCallNormalizer + completed bool // a finish reason was forwarded + sawChunk bool + lastUsage *domain.Usage } func NewStream(resp *http.Response, providerName ...string) *Stream { scanner := bufio.NewScanner(resp.Body) - scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) + scanner.Buffer(make([]byte, 0, 64*1024), maxSSELine) includeReasoningContent := len(providerName) > 0 && providerName[0] == "deepseek" return &Stream{ resp: resp, @@ -32,38 +53,43 @@ func NewStream(resp *http.Response, providerName ...string) *Stream { } func (s *Stream) Recv() (domain.StreamEvent, error) { - if s.done { - return domain.StreamEvent{}, io.EOF - } - - // Return a stashed completion event from a previous chunk that carried - // both a tool-call delta and a finish_reason. - if s.deferredCompleted != nil { - event := *s.deferredCompleted - s.deferredCompleted = nil - return event, nil + for len(s.pending) == 0 { + if s.done { + return domain.StreamEvent{}, io.EOF + } + events, err := s.read() + if err != nil { + return domain.StreamEvent{}, err + } + s.pending = events } + event := s.pending[0] + s.pending = s.pending[1:] + return event, nil +} +// read consumes upstream lines until they produce events or the stream ends. +func (s *Stream) read() ([]domain.StreamEvent, error) { for s.scanner.Scan() { line := s.scanner.Text() - if line == "" { continue } - - if !strings.HasPrefix(line, "data: ") { - if s.diagnostics != nil && strings.HasPrefix(line, "data:") { + // SSE allows "data:" with or without a following space. + data, ok := strings.CutPrefix(line, "data:") + if !ok { + if s.diagnostics != nil && !isSSEField(line) { s.diagnostics.unsupported++ } continue } - - data := strings.TrimPrefix(line, "data: ") - + data = strings.TrimPrefix(data, " ") + if data == "" { + continue + } if data == "[DONE]" { s.diagnostics.end("done_marker", nil) - s.done = true - return domain.StreamEvent{}, io.EOF + return s.end() } if s.diagnostics != nil { @@ -76,39 +102,98 @@ func (s *Stream) Recv() (domain.StreamEvent, error) { } continue } + s.sawChunk = true 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 + return nil, s.fail(chunk, err) + } + if err := providerStreamError(chunk.Error); err != nil { + return nil, s.fail(chunk, err) } - events := mapChunkToStreamEvents(chunk, s.includeReasoningContent) + events := s.normalize(mapChunkToStreamEvents(chunk, s.includeReasoningContent)) if s.diagnostics != nil { s.diagnostics.observe(chunk, events) } - if len(events) == 0 { - continue - } for i := range events { events[i].ProviderResponseID = chunk.ID } - // If the mapper produced two events (tool-call delta + completed), - // return the first now and stash the second for the next Recv(). - if len(events) > 1 { - s.deferredCompleted = &events[1] + if len(events) > 0 { + return events, nil } - return events[0], nil } if err := s.scanner.Err(); err != nil { s.diagnostics.end("read_error", err) - return domain.StreamEvent{}, err + return nil, err } - s.diagnostics.end("eof", nil) + return s.end() +} + +// isSSEField reports SSE lines other than data: comments, event, id, retry. +func isSSEField(line string) bool { + return strings.HasPrefix(line, ":") || strings.HasPrefix(line, "event:") || + strings.HasPrefix(line, "id:") || strings.HasPrefix(line, "retry:") +} + +func (s *Stream) fail(chunk chatCompletionChunk, err error) error { + if s.diagnostics != nil { + s.diagnostics.observe(chunk, nil) + } + s.diagnostics.end("provider_error", err) s.done = true - return domain.StreamEvent{}, io.EOF + return err +} + +// end finishes a stream the provider closed. A response without any chunk is +// a provider failure (so fallback can run); a response without a finish +// reason gets one, so clients always see a complete message. +func (s *Stream) end() ([]domain.StreamEvent, error) { + s.done = true + if !s.sawChunk { + return nil, domain.ErrProviderError(http.StatusBadGateway, "provider returned an empty stream") + } + if s.completed { + return nil, nil + } + finish := "stop" + return s.normalize([]domain.StreamEvent{{Type: domain.StreamEventCompleted, FinishReason: &finish, Usage: s.lastUsage}}), nil +} + +// normalize applies the per-stream rules described on Stream. +func (s *Stream) normalize(events []domain.StreamEvent) []domain.StreamEvent { + out := make([]domain.StreamEvent, 0, len(events)+1) + for _, event := range events { + if event.Usage != nil { + s.lastUsage = event.Usage + } + switch { + case event.ToolCallDelta != nil: + event.ToolCallDelta = s.tools.normalize(event.ToolCallDelta) + if event.ToolCallDelta == nil { + if event.Usage == nil { + continue + } + event.Type = domain.StreamEventOutputMessageDelta + } + case event.Type == domain.StreamEventCompleted && s.completed: + // Trailing usage after the finish reason. + case event.Type == domain.StreamEventCompleted && event.FinishReason == nil: + // Usage-only chunk before the model finished: not the end. + event.Type = domain.StreamEventOutputMessageDelta + case event.Type == domain.StreamEventCompleted: + out = append(out, s.tools.fillEmptyArguments()...) + if event.Usage == nil { + event.Usage = s.lastUsage // Usage sent before the finish reason. + } + if s.tools.seen() && *event.FinishReason == "stop" { + finish := "tool_calls" + event.FinishReason = &finish + } + s.completed = true + } + out = append(out, event) + } + return out } func (s *Stream) Close() error { @@ -117,6 +202,97 @@ func (s *Stream) Close() error { return s.resp.Body.Close() } +// toolCallNormalizer turns provider tool-call fragments into the OpenAI shape: +// indexes 0..n-1 in order of appearance, the id, type, and name on the first +// fragment only, and arguments as they stream. Providers differ: some pack +// several calls in one chunk, some repeat the id and name on every fragment, +// some reuse index 0 for every call with distinct ids, and some omit the id. +type toolCallNormalizer struct { + byIndex map[int]*toolCallSlot + slots []*toolCallSlot +} + +type toolCallSlot struct { + index int + id, name string + arguments int +} + +func (n *toolCallNormalizer) seen() bool { return len(n.slots) > 0 } + +func (n *toolCallNormalizer) normalize(in *domain.ToolCallDelta) *domain.ToolCallDelta { + if n.byIndex == nil { + n.byIndex = map[int]*toolCallSlot{} + } + id := "" + if in.ID != nil { + id = *in.ID + } + slot := n.byIndex[in.Index] + if slot != nil && id != "" && id != slot.id { + slot = nil // A new call that reuses the index. + } + out := &domain.ToolCallDelta{} + if slot == nil { + slot = &toolCallSlot{index: len(n.slots), id: id} + if slot.id == "" { + slot.id = newToolCallID() + } + n.slots = append(n.slots, slot) + n.byIndex[in.Index] = slot + out.ID = &slot.id + } + out.Index = slot.index + if in.Name != nil && *in.Name != "" && *in.Name != slot.name { + name := *in.Name + if slot.name != "" { + if !strings.HasPrefix(name, slot.name) { + slot.name += name // Streamed in pieces. + } else { + name = strings.TrimPrefix(name, slot.name) // Repeated with more text. + slot.name += name + } + } else { + slot.name = name + } + if name != "" { + out.Name = &name + } + } + if in.ArgumentsDelta != nil && *in.ArgumentsDelta != "" { + arguments := *in.ArgumentsDelta + slot.arguments += len(arguments) + out.ArgumentsDelta = &arguments + } + if out.ID == nil && out.Name == nil && out.ArgumentsDelta == nil { + return nil + } + return out +} + +// fillEmptyArguments gives calls that received no arguments "{}", which is +// what OpenAI sends and what strict clients parse. +func (n *toolCallNormalizer) fillEmptyArguments() []domain.StreamEvent { + var events []domain.StreamEvent + for _, slot := range n.slots { + if slot.arguments == 0 { + slot.arguments = 2 + empty := "{}" + events = append(events, domain.StreamEvent{ + Type: domain.StreamEventToolCallDelta, + ToolCallDelta: &domain.ToolCallDelta{Index: slot.index, ArgumentsDelta: &empty}, + }) + } + } + return events +} + +func newToolCallID() string { + b := make([]byte, 12) + _, _ = rand.Read(b) + return "call_" + hex.EncodeToString(b) +} + type chatCompletionChunk struct { ID string `json:"id"` Object string `json:"object"` @@ -150,6 +326,30 @@ type chatCompletionChunk struct { } `json:"prompt_tokens_details,omitempty"` } `json:"usage,omitempty"` BaseResp *providerBaseResponse `json:"base_resp,omitempty"` + Error *providerStreamErr `json:"error,omitempty"` +} + +// providerStreamErr is an error some providers send as a data line after the +// stream started (overload, context length, moderation). +type providerStreamErr struct { + Message string `json:"message"` + Type string `json:"type"` + Code any `json:"code"` +} + +func providerStreamError(e *providerStreamErr) error { + if e == nil { + return nil + } + message := strings.TrimSpace(e.Message) + if message == "" { + message = "provider stream failed" + } + return domain.ErrProviderError(http.StatusBadGateway, message).WithMeta( + "upstream_error_type", e.Type, + "upstream_code", e.Code, + "upstream_error", message, + ) } type providerBaseResponse struct { @@ -213,16 +413,11 @@ func mapChunkToStreamEvents(chunk chatCompletionChunk, includeReasoningContent . } } - // When a chunk carries both a tool-call delta AND a finish_reason (some - // providers, e.g. MiniMax, pack the final argument fragment and the - // finish signal into one chunk), we must emit the tool-call delta - // FIRST so the client receives the complete JSON arguments before the - // stream is marked done. - if choice.FinishReason != nil && len(choice.Delta.ToolCalls) > 0 { - tc := choice.Delta.ToolCalls[0] - tcd := &domain.ToolCallDelta{ - Index: tc.Index, - } + // Every tool call in the chunk, in order. Some providers pack several + // calls, or a whole call, into one chunk. + toolEvents := make([]domain.StreamEvent, 0, len(choice.Delta.ToolCalls)) + for _, tc := range choice.Delta.ToolCalls { + tcd := &domain.ToolCallDelta{Index: tc.Index} if tc.ID != "" { tcd.ID = &tc.ID } @@ -232,26 +427,39 @@ func mapChunkToStreamEvents(chunk chatCompletionChunk, includeReasoningContent . if tc.Function.Arguments != "" { tcd.ArgumentsDelta = &tc.Function.Arguments } + toolEvents = append(toolEvents, domain.StreamEvent{Type: domain.StreamEventToolCallDelta, ToolCallDelta: tcd}) + } + // Text that shares a chunk with tool calls comes first, as generated. + var text []domain.StreamEvent + if len(toolEvents) > 0 { + content := choice.Delta.Content != nil && *choice.Delta.Content != "" + reasoning := keepReasoning && choice.Delta.ReasoningContent != nil && *choice.Delta.ReasoningContent != "" + if content || reasoning { + event := domain.StreamEvent{Type: domain.StreamEventOutputTextDelta} + if content { + event.ContentDelta = choice.Delta.Content + } + if reasoning { + event.ReasoningDelta = choice.Delta.ReasoningContent + } + text = append(text, event) + } + } + + // When a chunk carries tool-call deltas AND a finish_reason (e.g. + // MiniMax packs the final argument fragment and the finish signal into + // one chunk), the deltas go first so the client has complete arguments + // before the stream is marked done. + if choice.FinishReason != nil && len(toolEvents) > 0 { completedEvent := domain.StreamEvent{ Type: domain.StreamEventCompleted, FinishReason: choice.FinishReason, } - if choice.Delta.Content != nil && *choice.Delta.Content != "" { - completedEvent.ContentDelta = choice.Delta.Content - } - if keepReasoning && choice.Delta.ReasoningContent != nil && *choice.Delta.ReasoningContent != "" { - completedEvent.ReasoningDelta = choice.Delta.ReasoningContent - } if chunk.Usage != nil { completedEvent.Usage = chunkUsageToDomain(chunk.Usage) } - return []domain.StreamEvent{ - { - Type: domain.StreamEventToolCallDelta, - ToolCallDelta: tcd, - }, - completedEvent, - } + events = append(text, toolEvents...) + return append(events, completedEvent) } // Check for finish reason -> completed event @@ -272,25 +480,8 @@ func mapChunkToStreamEvents(chunk chatCompletionChunk, includeReasoningContent . return []domain.StreamEvent{event} } - // Tool call delta - if len(choice.Delta.ToolCalls) > 0 { - tc := choice.Delta.ToolCalls[0] - tcd := &domain.ToolCallDelta{ - Index: tc.Index, - } - if tc.ID != "" { - tcd.ID = &tc.ID - } - if tc.Function.Name != "" { - tcd.Name = &tc.Function.Name - } - if tc.Function.Arguments != "" { - tcd.ArgumentsDelta = &tc.Function.Arguments - } - return []domain.StreamEvent{{ - Type: domain.StreamEventToolCallDelta, - ToolCallDelta: tcd, - }} + if len(toolEvents) > 0 { + return append(text, toolEvents...) } // Text content delta diff --git a/apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go b/apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go index bf8e9ea..9dd78ba 100644 --- a/apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go +++ b/apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go @@ -48,7 +48,8 @@ func TestStreamDiagnostics_TerminationAndUsage(t *testing.T) { {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: "unsupported framing", data: finish + "{SECRET-FRAMING}\n", end: "eof", level: "warn", unsupported: 1}, + {name: "data without space", data: finish + "data:{SECRET-FRAMING}\n", end: "eof", level: "warn", malformed: 1}, {name: "canceled", data: finish, end: "canceled", level: "warn", readErr: context.Canceled}, {name: "timeout", data: finish, end: "deadline_exceeded", level: "warn", readErr: context.DeadlineExceeded}, } { diff --git a/apps/gateway/internal/adapters/providers/openai/stream_test.go b/apps/gateway/internal/adapters/providers/openai/stream_test.go index 061b28b..8cdb605 100644 --- a/apps/gateway/internal/adapters/providers/openai/stream_test.go +++ b/apps/gateway/internal/adapters/providers/openai/stream_test.go @@ -214,11 +214,15 @@ func TestStream_ReasoningContentIgnoredForNonDeepSeekProviders(t *testing.T) { func TestStream_MapsDeepSeekCacheHitTokens(t *testing.T) { stream := newTestStream("data: {\"choices\":[],\"usage\":{\"prompt_tokens\":17,\"completion_tokens\":9,\"total_tokens\":26,\"prompt_cache_hit_tokens\":7}}\n\n") - event, err := stream.Recv() - if err != nil { - t.Fatalf("Recv returned error: %v", err) - } - if event.Type != domain.StreamEventCompleted || event.Usage == nil { + // A usage-only chunk is not the end; the stream's end is, with that usage. + var event domain.StreamEvent + for event.Type != domain.StreamEventCompleted { + var err error + if event, err = stream.Recv(); err != nil { + t.Fatalf("Recv returned error: %v", err) + } + } + if event.Usage == nil || event.FinishReason == nil || *event.FinishReason != "stop" { t.Fatalf("event = %+v, want completed event with usage", event) } if event.Usage.CacheReadTokens != 7 { diff --git a/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go b/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go new file mode 100644 index 0000000..b19e25b --- /dev/null +++ b/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go @@ -0,0 +1,172 @@ +package openai + +import ( + "encoding/json" + "errors" + "io" + "strings" + "testing" + + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" +) + +// strictCall is what a client that concatenates every fragment by index (as +// openai-python and most SDKs do) reassembles. +type strictCall struct{ id, name, args string } + +func drain(t *testing.T, sse string) (calls []strictCall, text string, finish string, err error) { + t.Helper() + stream := newTestStream(sse) + for { + event, recvErr := stream.Recv() + if recvErr == io.EOF { + return calls, text, finish, nil + } + if recvErr != nil { + return calls, text, finish, recvErr + } + if event.ContentDelta != nil { + text += *event.ContentDelta + } + if event.FinishReason != nil { + finish = *event.FinishReason + } + if d := event.ToolCallDelta; d != nil { + if d.Index > len(calls) { + t.Fatalf("index %d skips ahead of %d calls", d.Index, len(calls)) + } + if d.Index == len(calls) { + calls = append(calls, strictCall{}) + } + c := &calls[d.Index] + if d.ID != nil { + c.id += *d.ID + } + if d.Name != nil { + c.name += *d.Name + } + if d.ArgumentsDelta != nil { + c.args += *d.ArgumentsDelta + } + } + } +} + +func sse(chunks ...string) string { + var b strings.Builder + for _, c := range chunks { + b.WriteString("data: " + c + "\n\n") + } + return b.String() + "data: [DONE]\n\n" +} + +func wellFormed(t *testing.T, calls []strictCall, names ...string) { + t.Helper() + if len(calls) != len(names) { + t.Fatalf("got %d calls %+v, want %d", len(calls), calls, len(names)) + } + ids := map[string]bool{} + for i, c := range calls { + if c.name != names[i] || c.id == "" || ids[c.id] || !json.Valid([]byte(c.args)) { + t.Fatalf("call %d malformed: %+v", i, c) + } + ids[c.id] = true + } +} + +func TestStream_ToolCallProviderShapes(t *testing.T) { + for name, tc := range map[string]struct { + sse string + names []string + text string + }{ + "several calls packed in one chunk": {sse: sse( + `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"read","arguments":"{\"p\":1}"}},{"index":1,"id":"b","type":"function","function":{"name":"list","arguments":"{}"}}]}}]}`, + `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), + names: []string{"read", "list"}}, + "several calls packed with the finish reason": {sse: sse( + `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}},{"index":1,"id":"b","function":{"name":"read","arguments":"{}"}},{"index":2,"id":"c","function":{"name":"read","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}`), + names: []string{"read", "read", "read"}}, + "id and name repeated on every fragment": {sse: sse( + `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"read","arguments":"{\"p\""}}]}}]}`, + `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"read","arguments":":1}"}}]}}]}`, + `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), + names: []string{"read"}}, + "index reused with distinct ids": {sse: sse( + `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}}]}}]}`, + `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"b","function":{"name":"list","arguments":"{"}}]}}]}`, + `{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"}"}}]}}]}`, + `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), + names: []string{"read", "list"}}, + "missing id and sparse index": {sse: sse( + `{"choices":[{"delta":{"tool_calls":[{"index":3,"function":{"name":"read","arguments":"{}"}}]}}]}`, + `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), + names: []string{"read"}}, + "no arguments and a stop finish reason": {sse: sse( + `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"list"}}]}}]}`, + `{"choices":[{"delta":{},"finish_reason":"stop"}]}`), + names: []string{"list"}}, + "text in the same chunk as a call": {sse: sse( + `{"choices":[{"delta":{"content":"Reading.","tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}}]}}]}`, + `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), + names: []string{"read"}, text: "Reading."}, + "no finish reason at all": {sse: sse( + `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}}]}}]}`), + names: []string{"read"}}, + } { + t.Run(name, func(t *testing.T) { + calls, text, finish, err := drain(t, tc.sse) + if err != nil { + t.Fatal(err) + } + wellFormed(t, calls, tc.names...) + if text != tc.text || finish != "tool_calls" { + t.Fatalf("text %q finish %q", text, finish) + } + }) + } +} + +func TestStream_UsageBeforeFinishDoesNotEndTheStream(t *testing.T) { + stream := newTestStream(sse( + `{"choices":[],"usage":{"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}}`, + `{"choices":[{"delta":{"content":"late"}}]}`, + `{"choices":[{"delta":{},"finish_reason":"stop"}]}`)) + var got []domain.StreamEvent + for { + event, err := stream.Recv() + if err != nil { + break + } + got = append(got, event) + } + if len(got) != 3 || got[0].Type == domain.StreamEventCompleted || got[1].ContentDelta == nil || got[2].Type != domain.StreamEventCompleted || got[2].Usage == nil { + t.Fatalf("events %+v", got) + } +} + +func TestStream_ProviderFailures(t *testing.T) { + for name, data := range map[string]string{ + "error payload mid-stream": "data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\n\ndata: {\"error\":{\"message\":\"overloaded\",\"type\":\"server_error\"}}\n\n", + "empty stream": "data: [DONE]\n\n", + "nothing at all": "", + } { + t.Run(name, func(t *testing.T) { + _, _, _, err := drain(t, data) + var gwErr *domain.GatewayError + if !errors.As(err, &gwErr) { + t.Fatalf("err = %v, want a provider error", err) + } + }) + } +} + +func TestStream_LargeSingleLineToolCall(t *testing.T) { + big := strings.Repeat("x", 3<<20) + calls, _, _, err := drain(t, sse( + `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"write","arguments":"{\"text\":\"`+big+`\"}"}}]},"finish_reason":"tool_calls"}]}`)) + if err != nil { + t.Fatal(err) + } + wellFormed(t, calls, "write") +} From d1247102bd1bde5978351c78eff667bd34e708a4 Mon Sep 17 00:00:00 2001 From: Marketen Date: Tue, 29 Sep 2026 10:33:16 +0200 Subject: [PATCH 2/7] feat: proxy OpenAI-compatible requests instead of translating them Requests and stream chunks were parsed into a narrow domain model and rebuilt, losing whatever it didn't cover: tool calls after the first in a chunk, output after the first finish chunk, images, reasoning from every provider but DeepSeek, and unknown request fields. For OpenAI-compatible providers (all of production) the client's body now goes upstream with a few documented edits (model name, usage for billing, token limit clamp, provider compatibility shims), and the provider's response comes back byte for byte except the model field. The gateway reads usage, finish reason, and errors from the passing stream for metering and fallback. The translating path remains only for PII masking with the filter on and for Anthropic, keeps its fidelity fixes, and no longer rewrites tool calls. Co-Authored-By: Claude Opus 5.5 --- .../http/handlers/chat_completions_handler.go | 110 +++- .../handlers/chat_completions_stream_test.go | 57 ++ .../internal/adapters/http/sse/writer.go | 9 + .../adapters/providers/openai/proxy.go | 227 ++++++++ .../adapters/providers/openai/proxy_test.go | 94 ++++ .../adapters/providers/openai/stream.go | 126 +---- .../providers/openai/stream_toolcalls_test.go | 37 +- .../adapters/providers/tinfoil/adapter.go | 43 ++ .../internal/application/ports/provider.go | 20 + .../services/chat_completions_service.go | 6 + .../internal/application/services/proxy.go | 507 ++++++++++++++++++ .../application/services/proxy_test.go | 262 +++++++++ 12 files changed, 1347 insertions(+), 151 deletions(-) create mode 100644 apps/gateway/internal/adapters/providers/openai/proxy.go create mode 100644 apps/gateway/internal/adapters/providers/openai/proxy_test.go create mode 100644 apps/gateway/internal/application/services/proxy.go create mode 100644 apps/gateway/internal/application/services/proxy_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 437d061..8c69eea 100644 --- a/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go +++ b/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go @@ -2,9 +2,11 @@ package handlers import ( "context" + "errors" "io" "net/http" "sync" + "sync/atomic" "time" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/mapper" @@ -43,6 +45,45 @@ func (h *ChatCompletionsHandler) Handle(w http.ResponseWriter, r *http.Request) return } + genReq, err := mapper.ChatCompletionRequestToDomain(body) + if err != nil { + WriteErrorWithLog(w, r, h.logger, err) + return + } + + // The gateway is a proxy: the request goes upstream as the client sent it + // and the provider's response comes back as it came. Only requests that + // need content rewritten (PII masking) or another wire format take the + // translating path below. + upstream, cancelUpstream := context.WithCancel(context.WithoutCancel(r.Context())) + defer cancelUpstream() + var finished atomic.Bool + stopWatch := context.AfterFunc(r.Context(), func() { + // A client gone after the finish reason still gets its usage metered: + // the upstream is drained briefly instead of cut off. + if finished.Load() { + time.AfterFunc(streamUsageDrainTimeout, cancelUpstream) + return + } + cancelUpstream() + }) + defer stopWatch() + call, err := h.service.Proxy(upstream, body, genReq, token) + if err == nil { + if call.Stream != nil { + h.relay(w, r, call.Stream, &finished) + return + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write(call.Body) + return + } + if !errors.Is(err, services.ErrNotProxyable) { + WriteErrorWithLog(w, r, h.logger, err) + return + } + if fields := mapper.UnknownChatCompletionFields(body); len(fields) > 0 { h.logger.Warn("chat completion request ignored unknown fields", "request_id", middleware.GetRequestID(r.Context()), @@ -51,12 +92,6 @@ func (h *ChatCompletionsHandler) Handle(w http.ResponseWriter, r *http.Request) ) } - genReq, err := mapper.ChatCompletionRequestToDomain(body) - if err != nil { - WriteErrorWithLog(w, r, h.logger, err) - return - } - if genReq.Stream { h.handleStream(w, r, genReq, token) return @@ -72,6 +107,69 @@ func (h *ChatCompletionsHandler) Handle(w http.ResponseWriter, r *http.Request) WriteJSON(w, http.StatusOK, resp) } +// relay sends a provider's stream to the client line by line, as it came. +// Keep-alive comments fill long silences (slow reasoning), and a stream the +// provider finished without [DONE] gets one, so every client sees an end. +func (h *ChatCompletionsHandler) relay(w http.ResponseWriter, r *http.Request, stream *services.ProxyStream, finished *atomic.Bool) { + defer stream.Close() + sw, err := sse.NewWriter(w) + if err != nil { + WriteError(w, domain.ErrInternal("an internal error occurred")) + return + } + var writeMu sync.Mutex + lastWrite := time.Now() + write := func(line []byte) { + writeMu.Lock() + defer writeMu.Unlock() + lastWrite = time.Now() + // A client that left keeps the loop draining the upstream for usage. + _ = sw.WriteLine(line) + } + stopKeepAlive := make(chan struct{}) + defer close(stopKeepAlive) + go func() { + ticker := time.NewTicker(5 * time.Second) + defer ticker.Stop() + for { + select { + case <-stopKeepAlive: + return + case <-ticker.C: + writeMu.Lock() + if time.Since(lastWrite) >= 15*time.Second { + lastWrite = time.Now() + _ = sw.WriteComment("keepalive") + } + writeMu.Unlock() + } + } + }() + + var readErr error + for { + line, err := stream.Next() + if err != nil { + readErr = err + break + } + if stream.Finished() { + finished.Store(true) + } + write(line) + } + switch { + case stream.Complete() && !stream.SawDone(): + write([]byte("data: [DONE]")) + write(nil) + case readErr != io.EOF && r.Context().Err() == nil: + // The upstream broke mid-response: say so rather than end quietly. + h.logger.Warn("proxied stream interrupted", "request_id", middleware.GetRequestID(r.Context())) + write([]byte(`data: {"error":{"type":"provider_error","code":"stream_interrupted","message":"The provider stream ended unexpectedly."}}`)) + write(nil) + } +} + func (h *ChatCompletionsHandler) handleStream(w http.ResponseWriter, r *http.Request, genReq domain.GenerateRequest, token string) { ctx, cancelUpstream := context.WithCancel(r.Context()) defer cancelUpstream() 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 321a41c..25eb290 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 @@ -6,9 +6,12 @@ import ( "io" "net/http/httptest" "strings" + "sync/atomic" "testing" "time" + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/services" + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) @@ -196,3 +199,57 @@ func TestChatCompletionsHandler_ProviderErrorEventIsNotCompletion(t *testing.T) t.Fatalf("error event reported as completion: %s", body) } } + +type brokenBody struct{ io.Reader } + +func (brokenBody) Close() error { return nil } + +type failAfter struct { + r io.Reader + err error +} + +func (f *failAfter) Read(p []byte) (int, error) { + n, err := f.r.Read(p) + if err == io.EOF { + return n, f.err + } + return n, err +} + +func TestRelay_EndsEveryStream(t *testing.T) { + finishedNoDone := "data: {\"choices\":[{\"delta\":{\"content\":\"hi\"},\"finish_reason\":\"stop\"}]}\n\n" + for name, tc := range map[string]struct { + body io.Reader + want []string + not []string + }{ + "relayed as sent": {body: strings.NewReader(finishedNoDone + "data: [DONE]\n\n"), want: []string{`"content":"hi"`}}, + "finish without [DONE] gets one": {body: strings.NewReader(finishedNoDone), want: []string{"data: [DONE]"}}, + "broken mid-response says so": { + body: &failAfter{r: strings.NewReader("data: {\"choices\":[{\"delta\":{\"content\":\"par\"}}]}\n\n"), err: errors.New("connection reset")}, + want: []string{"stream_interrupted"}, not: []string{"[DONE]"}, + }, + } { + t.Run(name, func(t *testing.T) { + recorder := httptest.NewRecorder() + handler := &ChatCompletionsHandler{logger: &confidentialTestLogger{}} + var finished atomic.Bool + handler.relay(recorder, httptest.NewRequest("POST", "/v1/chat/completions", nil), services.NewProxyStream(brokenBody{tc.body}, "m"), &finished) + body := recorder.Body.String() + for _, want := range tc.want { + if !strings.Contains(body, want) { + t.Fatalf("missing %q in %q", want, body) + } + } + for _, not := range tc.not { + if strings.Contains(body, not) { + t.Fatalf("unexpected %q in %q", not, body) + } + } + if strings.Count(body, "data: [DONE]") > 1 { + t.Fatalf("[DONE] twice: %q", body) + } + }) + } +} diff --git a/apps/gateway/internal/adapters/http/sse/writer.go b/apps/gateway/internal/adapters/http/sse/writer.go index e83eaed..d9b5f3d 100644 --- a/apps/gateway/internal/adapters/http/sse/writer.go +++ b/apps/gateway/internal/adapters/http/sse/writer.go @@ -79,3 +79,12 @@ func (sw *Writer) WriteComment(text string) error { sw.flusher.Flush() return nil } + +// WriteLine relays one line of an upstream SSE stream as it came. +func (sw *Writer) WriteLine(line []byte) error { + if _, err := sw.w.Write(append(line, '\n')); err != nil { + return err + } + sw.flusher.Flush() + return nil +} diff --git a/apps/gateway/internal/adapters/providers/openai/proxy.go b/apps/gateway/internal/adapters/providers/openai/proxy.go new file mode 100644 index 0000000..b8f836d --- /dev/null +++ b/apps/gateway/internal/adapters/providers/openai/proxy.go @@ -0,0 +1,227 @@ +package openai + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + + "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" +) + +// gatewayOnlyFields are request fields for the gateway itself; providers +// never see them. +var gatewayOnlyFields = []string{"provider_options"} + +// PrepareProxyBody turns the client's request body into the provider's. The +// gateway is a proxy: every field the client sent is forwarded unchanged, +// including ones the gateway doesn't know (images, reasoning options, +// provider extensions). The only edits are: +// +// - model: the provider's name for the public model; +// - stream_options.include_usage: on for streams, so usage can be metered; +// - the output-token limit: clamped to the model's maximum, and sent as +// max_tokens to the providers that only document that name (Novita, +// DeepSeek); +// - provider compatibility for valid OpenAI requests the provider would +// otherwise reject: developer role as system, content:null on assistant +// tool-call messages, and reasoning_content on DeepSeek tool-call turns; +// - gateway-only fields are removed. +func PrepareProxyBody(raw []byte, model domain.PublicModel, stream bool) (map[string]any, []string, error) { + dec := json.NewDecoder(bytes.NewReader(raw)) + dec.UseNumber() // Numbers are forwarded exactly as sent. + var body map[string]any + if err := dec.Decode(&body); err != nil || body == nil { + return nil, nil, domain.ErrInvalidField("invalid JSON body") + } + policy := policyForProvider(model.ProviderConfig.ProviderName) + var transforms []string + + body["model"] = model.UpstreamModelName + for _, field := range gatewayOnlyFields { + delete(body, field) + } + if stream { + body["stream"] = true + options, _ := body["stream_options"].(map[string]any) + if options == nil { + options = map[string]any{} + } + options["include_usage"] = true + body["stream_options"] = options + } else { + delete(body, "stream_options") // Only valid on streams. + } + + if limit, field, ok := tokenLimit(body); ok { + if model.MaxOutputTokens > 0 && limit > int64(model.MaxOutputTokens) { + limit = int64(model.MaxOutputTokens) + transforms = append(transforms, "token_limit=clamped") + } + delete(body, "max_tokens") + delete(body, "max_completion_tokens") + target := field // Keep the client's field, unless the provider only knows max_tokens. + if policy.useMaxTokens { + target = "max_tokens" + } + if target != field { + transforms = append(transforms, "token_limit_field="+target) + } + body[target] = limit + } + + messages := asMaps(body["messages"]) + for _, msg := range messages { + role, _ := msg["role"].(string) + if role == "developer" && !policy.forwardDeveloperRole { + msg["role"] = "system" + transforms = append(transforms, "developer_role=system") + } + if role != "assistant" { + continue + } + toolCalls := asMaps(msg["tool_calls"]) + if len(toolCalls) > 0 { + msg["tool_calls"] = toolCalls + } + if _, has := msg["content"]; !has && len(toolCalls) > 0 && policy.explicitNullAssistantToolContent { + msg["content"] = nil + transforms = append(transforms, "assistant_tool_content=null") + } + if _, has := msg["reasoning_content"]; !has && len(toolCalls) > 0 && policy.requireToolReasoningContent { + msg["reasoning_content"] = "" + transforms = append(transforms, "assistant_tool_reasoning_content=empty") + } + } + if messages != nil { + body["messages"] = messages + } + if tools := asMaps(body["tools"]); tools != nil { + body["tools"] = tools + } + return body, dedupe(transforms), nil +} + +// tokenLimit reads the client's output-token limit and which field held it. +func tokenLimit(body map[string]any) (int64, string, bool) { + for _, field := range []string{"max_completion_tokens", "max_tokens"} { + if n, ok := body[field].(json.Number); ok { + if v, err := n.Int64(); err == nil { + return v, field, true + } + } + } + return 0, "", false +} + +// asMaps views a JSON array of objects as maps (shared with the body, so +// edits apply in place). It returns nil for anything else. +func asMaps(v any) []map[string]any { + switch list := v.(type) { + case []map[string]any: + return list + case []any: + out := make([]map[string]any, 0, len(list)) + for _, item := range list { + m, ok := item.(map[string]any) + if !ok { + return nil + } + out = append(out, m) + } + return out + } + return nil +} + +func dedupe(list []string) []string { + seen := map[string]bool{} + out := list[:0] + for _, s := range list { + if !seen[s] { + seen[s] = true + out = append(out, s) + } + } + return out +} + +// ProxyStream forwards a streaming request, with the same retries as the +// translating path, and returns the provider's SSE body untouched. +func (a *Adapter) ProxyStream(ctx context.Context, raw []byte, model domain.PublicModel) (ports.ProxyResponse, error) { + body, _, err := PrepareProxyBody(raw, model, true) + if err != nil { + return ports.ProxyResponse{}, err + } + resp, err := a.proxy(ctx, body, model, func(ctx context.Context, apiKey string, body map[string]any) (*http.Response, error) { + return a.client.DoStream(ctx, model.ProviderConfig.BaseURL, apiKey, body) + }) + if err != nil { + return ports.ProxyResponse{}, err + } + return ports.ProxyResponse{Body: resp.Body}, nil +} + +// ProxyJSON forwards a non-streaming request and returns the provider's JSON. +func (a *Adapter) ProxyJSON(ctx context.Context, raw []byte, model domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) { + body, _, err := PrepareProxyBody(raw, model, false) + if err != nil { + return nil, nil, err + } + var out []byte + _, err = a.proxy(ctx, body, model, func(ctx context.Context, apiKey string, body map[string]any) (*http.Response, error) { + var err error + out, err = a.client.Do(ctx, model.ProviderConfig.BaseURL, apiKey, body) + if err != nil { + return nil, err + } + return &http.Response{Body: io.NopCloser(bytes.NewReader(nil))}, nil + }) + return out, nil, err +} + +// proxy sends body with the provider's retry policy (Novita's transient +// rejections), and maps errors like the translating path. +func (a *Adapter) proxy(ctx context.Context, body map[string]any, model domain.PublicModel, send func(context.Context, string, map[string]any) (*http.Response, error)) (*http.Response, error) { + apiKey := resolveAPIKey(model.ProviderConfig.APIKeySecretRef) + if apiKey == "" { + return nil, missingProviderCredentialError(model.ProviderConfig.ProviderName) + } + built := builtProviderRequest{Body: body, Policy: policyForProvider(model.ProviderConfig.ProviderName).name} + active := built + var retryReason string + invalidTraceRetries, overloadRetries := 0, 0 + downgradeRetried := false + for attempt := 1; ; attempt++ { + a.logProviderRequest(ctx, model, active, attempt, retryReason) + resp, err := send(ctx, apiKey, active.Body) + if err == nil { + return resp, nil + } + if retry := maybeBuildNovitaSameBodyRetry(model, err, active.Body); retry.CanRetry && + canSpendSameBodyRetry(retry.RetryReason, &invalidTraceRetries, &overloadRetries) { + retryReason = retry.RetryReason + metrics.ProviderRetries.WithLabelValues(model.ProviderConfig.ProviderName, retryReason).Inc() + if !sleepBeforeProviderRetry(ctx, retryReason) { + return nil, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), active, attempt, retryReason) + } + continue + } + if !downgradeRetried { + if retry := maybeBuildNovitaDowngradeRetry(model, err, active.Body); retry.CanRetry { + downgradeRetried = true + retryReason = retry.RetryReason + metrics.ProviderRetries.WithLabelValues(model.ProviderConfig.ProviderName, retryReason).Inc() + active = builtProviderRequest{Body: retry.Body, Policy: built.Policy, Omitted: retry.Omitted} + if !sleepBeforeProviderRetry(ctx, retryReason) { + return nil, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), active, attempt, retryReason) + } + continue + } + } + return nil, withProviderPolicyMeta(mapProviderErrorWithCompatibilityContext(err, model, active.Body), active, attempt, retryReason) + } +} diff --git a/apps/gateway/internal/adapters/providers/openai/proxy_test.go b/apps/gateway/internal/adapters/providers/openai/proxy_test.go new file mode 100644 index 0000000..148f0aa --- /dev/null +++ b/apps/gateway/internal/adapters/providers/openai/proxy_test.go @@ -0,0 +1,94 @@ +package openai + +import ( + "encoding/json" + "reflect" + "testing" + + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" +) + +func prepared(t *testing.T, raw, provider string, stream bool) map[string]any { + t.Helper() + body, _, err := PrepareProxyBody([]byte(raw), domain.PublicModel{ + UpstreamModelName: "upstream-name", MaxOutputTokens: 8000, + ProviderConfig: domain.ProviderConfig{ProviderName: provider}, + }, stream) + if err != nil { + t.Fatal(err) + } + // Compare as JSON the provider will receive. + b, _ := json.Marshal(body) + var out map[string]any + json.Unmarshal(b, &out) + return out +} + +func decode(t *testing.T, raw string) map[string]any { + t.Helper() + var out map[string]any + if err := json.Unmarshal([]byte(raw), &out); err != nil { + t.Fatal(err) + } + return out +} + +// Everything the client sends reaches the provider, including what the +// gateway doesn't model: images, provider extensions, reasoning options. +func TestPrepareProxyBody_ForwardsTheClientRequest(t *testing.T) { + raw := `{"model":"public/x","messages":[{"role":"user","content":[{"type":"text","text":"What is this?"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AAAA","detail":"high"}}]}, + {"role":"assistant","content":"Let me look.","reasoning_content":"thinking","tool_calls":[{"id":"a","type":"function","function":{"name":"read","arguments":"{}"}}]}, + {"role":"tool","tool_call_id":"a","content":[{"type":"text","text":"file"}]}], + "tools":[{"type":"function","function":{"name":"read","description":"Read","parameters":{"type":"object","properties":{}},"strict":false}}], + "temperature":0.7,"seed":9007199254740993,"top_k":40,"chat_template_kwargs":{"enable_thinking":false},"reasoning":{"effort":"high"},"response_format":{"type":"json_object"}, + "n":1,"logprobs":true,"provider_options":{"x":1}}` + got := prepared(t, raw, "phala", false) + want := decode(t, raw) + want["model"] = "upstream-name" + delete(want, "provider_options") + if !reflect.DeepEqual(got, want) { + t.Fatalf("request changed:\n got %v\nwant %v", got, want) + } + // Large integers are forwarded exactly. + body, _, _ := PrepareProxyBody([]byte(raw), domain.PublicModel{UpstreamModelName: "m"}, false) + if b, _ := json.Marshal(body["seed"]); string(b) != "9007199254740993" { + t.Fatalf("seed = %s", b) + } +} + +func TestPrepareProxyBody_Edits(t *testing.T) { + t.Run("streams always report usage", func(t *testing.T) { + got := prepared(t, `{"messages":[],"stream":true,"stream_options":{"include_usage":false,"continuous_usage_stats":true}}`, "novita", true) + options := got["stream_options"].(map[string]any) + if got["stream"] != true || options["include_usage"] != true || options["continuous_usage_stats"] != true { + t.Fatalf("stream_options %v", options) + } + }) + t.Run("output limit clamped and named per provider", func(t *testing.T) { + for provider, field := range map[string]string{"novita": "max_tokens", "deepseek": "max_tokens", "tinfoil": "max_completion_tokens", "phala": "max_completion_tokens"} { + got := prepared(t, `{"messages":[],"max_completion_tokens":50000}`, provider, false) + if got[field] != float64(8000) || len(got) != 3 { + t.Fatalf("%s: %v", provider, got) + } + } + // Other providers get the client's own field. + if got := prepared(t, `{"messages":[],"max_tokens":100}`, "tinfoil", false); got["max_tokens"] != float64(100) || len(got) != 3 { + t.Fatalf("field renamed: %v", got) + } + }) + t.Run("provider compatibility for valid OpenAI requests", func(t *testing.T) { + raw := `{"messages":[{"role":"developer","content":"be brief"},{"role":"assistant","tool_calls":[{"id":"a","type":"function","function":{"name":"f","arguments":"{}"}}]}]}` + deepseek := prepared(t, raw, "deepseek", false)["messages"].([]any) + if deepseek[0].(map[string]any)["role"] != "system" { + t.Fatal("developer role not mapped") + } + assistant := deepseek[1].(map[string]any) + if v, ok := assistant["content"]; !ok || v != nil || assistant["reasoning_content"] != "" { + t.Fatalf("deepseek assistant %v", assistant) + } + other := prepared(t, raw, "tinfoil", false)["messages"].([]any)[1].(map[string]any) + if _, ok := other["reasoning_content"]; ok { + t.Fatal("reasoning_content added for a provider that doesn't need it") + } + }) +} diff --git a/apps/gateway/internal/adapters/providers/openai/stream.go b/apps/gateway/internal/adapters/providers/openai/stream.go index cd99a60..8316008 100644 --- a/apps/gateway/internal/adapters/providers/openai/stream.go +++ b/apps/gateway/internal/adapters/providers/openai/stream.go @@ -2,8 +2,6 @@ package openai import ( "bufio" - "crypto/rand" - "encoding/hex" "encoding/json" "io" "net/http" @@ -16,18 +14,15 @@ import ( // arguments send a whole call (for example a large file) in one line. const maxSSELine = 16 << 20 -// Stream reads SSE events from an OpenAI-compatible streaming response and -// normalizes them, so every client sees the same well-formed shape whatever -// the provider sends: +// Stream reads SSE events from an OpenAI-compatible streaming response for +// the translating path (PII masking, non-OpenAI clients of the domain model). +// Proxied requests don't use it. It keeps everything the provider sent: // -// - every tool call in a chunk is forwarded, not only the first; -// - each tool call gets contiguous indexes, one id (synthesized if the -// provider omits it), and its id and name only once; -// - a call that received no arguments gets "{}"; -// - text in a chunk that also carries tool calls is kept; -// - a usage-only chunk before the finish reason does not end the stream; -// - a stream that ends without a finish reason gets one, while a stream with -// no chunk at all or an error payload is reported as a provider error. +// - every tool call in a chunk, not only the first; +// - text in a chunk that also carries tool calls; +// - output after an early usage-only chunk (which doesn't end the stream); +// - a finish reason at the end, and a provider error for a stream with no +// chunk at all or an error payload. type Stream struct { diagnostics *streamDiagnostics resp *http.Response @@ -35,7 +30,6 @@ type Stream struct { done bool // upstream fully read includeReasoningContent bool pending []domain.StreamEvent - tools toolCallNormalizer completed bool // a finish reason was forwarded sawChunk bool lastUsage *domain.Usage @@ -167,28 +161,15 @@ func (s *Stream) normalize(events []domain.StreamEvent) []domain.StreamEvent { s.lastUsage = event.Usage } switch { - case event.ToolCallDelta != nil: - event.ToolCallDelta = s.tools.normalize(event.ToolCallDelta) - if event.ToolCallDelta == nil { - if event.Usage == nil { - continue - } - event.Type = domain.StreamEventOutputMessageDelta - } case event.Type == domain.StreamEventCompleted && s.completed: // Trailing usage after the finish reason. case event.Type == domain.StreamEventCompleted && event.FinishReason == nil: // Usage-only chunk before the model finished: not the end. event.Type = domain.StreamEventOutputMessageDelta case event.Type == domain.StreamEventCompleted: - out = append(out, s.tools.fillEmptyArguments()...) if event.Usage == nil { event.Usage = s.lastUsage // Usage sent before the finish reason. } - if s.tools.seen() && *event.FinishReason == "stop" { - finish := "tool_calls" - event.FinishReason = &finish - } s.completed = true } out = append(out, event) @@ -202,97 +183,6 @@ func (s *Stream) Close() error { return s.resp.Body.Close() } -// toolCallNormalizer turns provider tool-call fragments into the OpenAI shape: -// indexes 0..n-1 in order of appearance, the id, type, and name on the first -// fragment only, and arguments as they stream. Providers differ: some pack -// several calls in one chunk, some repeat the id and name on every fragment, -// some reuse index 0 for every call with distinct ids, and some omit the id. -type toolCallNormalizer struct { - byIndex map[int]*toolCallSlot - slots []*toolCallSlot -} - -type toolCallSlot struct { - index int - id, name string - arguments int -} - -func (n *toolCallNormalizer) seen() bool { return len(n.slots) > 0 } - -func (n *toolCallNormalizer) normalize(in *domain.ToolCallDelta) *domain.ToolCallDelta { - if n.byIndex == nil { - n.byIndex = map[int]*toolCallSlot{} - } - id := "" - if in.ID != nil { - id = *in.ID - } - slot := n.byIndex[in.Index] - if slot != nil && id != "" && id != slot.id { - slot = nil // A new call that reuses the index. - } - out := &domain.ToolCallDelta{} - if slot == nil { - slot = &toolCallSlot{index: len(n.slots), id: id} - if slot.id == "" { - slot.id = newToolCallID() - } - n.slots = append(n.slots, slot) - n.byIndex[in.Index] = slot - out.ID = &slot.id - } - out.Index = slot.index - if in.Name != nil && *in.Name != "" && *in.Name != slot.name { - name := *in.Name - if slot.name != "" { - if !strings.HasPrefix(name, slot.name) { - slot.name += name // Streamed in pieces. - } else { - name = strings.TrimPrefix(name, slot.name) // Repeated with more text. - slot.name += name - } - } else { - slot.name = name - } - if name != "" { - out.Name = &name - } - } - if in.ArgumentsDelta != nil && *in.ArgumentsDelta != "" { - arguments := *in.ArgumentsDelta - slot.arguments += len(arguments) - out.ArgumentsDelta = &arguments - } - if out.ID == nil && out.Name == nil && out.ArgumentsDelta == nil { - return nil - } - return out -} - -// fillEmptyArguments gives calls that received no arguments "{}", which is -// what OpenAI sends and what strict clients parse. -func (n *toolCallNormalizer) fillEmptyArguments() []domain.StreamEvent { - var events []domain.StreamEvent - for _, slot := range n.slots { - if slot.arguments == 0 { - slot.arguments = 2 - empty := "{}" - events = append(events, domain.StreamEvent{ - Type: domain.StreamEventToolCallDelta, - ToolCallDelta: &domain.ToolCallDelta{Index: slot.index, ArgumentsDelta: &empty}, - }) - } - } - return events -} - -func newToolCallID() string { - b := make([]byte, 12) - _, _ = rand.Read(b) - return "call_" + hex.EncodeToString(b) -} - type chatCompletionChunk struct { ID string `json:"id"` Object string `json:"object"` diff --git a/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go b/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go index b19e25b..080fb38 100644 --- a/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go +++ b/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go @@ -75,44 +75,27 @@ func wellFormed(t *testing.T, calls []strictCall, names ...string) { } func TestStream_ToolCallProviderShapes(t *testing.T) { + // The translating path keeps what providers send; it doesn't repair it. for name, tc := range map[string]struct { - sse string - names []string - text string + sse string + names []string + text string + finish string }{ "several calls packed in one chunk": {sse: sse( `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"read","arguments":"{\"p\":1}"}},{"index":1,"id":"b","type":"function","function":{"name":"list","arguments":"{}"}}]}}]}`, `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), - names: []string{"read", "list"}}, + names: []string{"read", "list"}, finish: "tool_calls"}, "several calls packed with the finish reason": {sse: sse( `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}},{"index":1,"id":"b","function":{"name":"read","arguments":"{}"}},{"index":2,"id":"c","function":{"name":"read","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}`), - names: []string{"read", "read", "read"}}, - "id and name repeated on every fragment": {sse: sse( - `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"read","arguments":"{\"p\""}}]}}]}`, - `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"read","arguments":":1}"}}]}}]}`, - `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), - names: []string{"read"}}, - "index reused with distinct ids": {sse: sse( - `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}}]}}]}`, - `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"b","function":{"name":"list","arguments":"{"}}]}}]}`, - `{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"}"}}]}}]}`, - `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), - names: []string{"read", "list"}}, - "missing id and sparse index": {sse: sse( - `{"choices":[{"delta":{"tool_calls":[{"index":3,"function":{"name":"read","arguments":"{}"}}]}}]}`, - `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), - names: []string{"read"}}, - "no arguments and a stop finish reason": {sse: sse( - `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"list"}}]}}]}`, - `{"choices":[{"delta":{},"finish_reason":"stop"}]}`), - names: []string{"list"}}, + names: []string{"read", "read", "read"}, finish: "tool_calls"}, "text in the same chunk as a call": {sse: sse( `{"choices":[{"delta":{"content":"Reading.","tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}}]}}]}`, `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), - names: []string{"read"}, text: "Reading."}, + names: []string{"read"}, text: "Reading.", finish: "tool_calls"}, "no finish reason at all": {sse: sse( `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}}]}}]}`), - names: []string{"read"}}, + names: []string{"read"}, finish: "stop"}, } { t.Run(name, func(t *testing.T) { calls, text, finish, err := drain(t, tc.sse) @@ -120,7 +103,7 @@ func TestStream_ToolCallProviderShapes(t *testing.T) { t.Fatal(err) } wellFormed(t, calls, tc.names...) - if text != tc.text || finish != "tool_calls" { + if text != tc.text || finish != tc.finish { t.Fatalf("text %q finish %q", text, finish) } }) diff --git a/apps/gateway/internal/adapters/providers/tinfoil/adapter.go b/apps/gateway/internal/adapters/providers/tinfoil/adapter.go index f927d84..2a32440 100644 --- a/apps/gateway/internal/adapters/providers/tinfoil/adapter.go +++ b/apps/gateway/internal/adapters/providers/tinfoil/adapter.go @@ -340,3 +340,46 @@ func (c *sdkVerifiedClient) TransportMode() string { func (c *sdkVerifiedClient) GroundTruth() *verifierclient.GroundTruth { return c.groundTruth } + +// ProxyStream forwards a streaming request over the attested transport and +// returns the SSE body untouched, with the transport proof. +func (a *Adapter) ProxyStream(ctx context.Context, raw []byte, model domain.PublicModel) (ports.ProxyResponse, error) { + apiKey := os.Getenv(model.ProviderConfig.APIKeySecretRef) + if apiKey == "" { + return ports.ProxyResponse{}, missingProviderCredentialError() + } + body, _, err := openai.PrepareProxyBody(raw, model, true) + if err != nil { + return ports.ProxyResponse{}, err + } + verified, proof, err := a.newVerifiedClient(ctx, model) + if err != nil { + return ports.ProxyResponse{}, err + } + resp, err := a.doStream(ctx, verified, apiKey, body) + if err != nil { + return ports.ProxyResponse{}, openai.MapProviderErrorWithCompatibilityContext(err, model, body) + } + return ports.ProxyResponse{Body: resp.Body, Proof: proof}, nil +} + +// ProxyJSON forwards a non-streaming request over the attested transport. +func (a *Adapter) ProxyJSON(ctx context.Context, raw []byte, model domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) { + apiKey := os.Getenv(model.ProviderConfig.APIKeySecretRef) + if apiKey == "" { + return nil, nil, missingProviderCredentialError() + } + body, _, err := openai.PrepareProxyBody(raw, model, false) + if err != nil { + return nil, nil, err + } + verified, proof, err := a.newVerifiedClient(ctx, model) + if err != nil { + return nil, nil, err + } + out, err := a.do(ctx, verified, apiKey, body) + if err != nil { + return nil, nil, openai.MapProviderErrorWithCompatibilityContext(err, model, body) + } + return out, proof, nil +} diff --git a/apps/gateway/internal/application/ports/provider.go b/apps/gateway/internal/application/ports/provider.go index 5503197..50b1185 100644 --- a/apps/gateway/internal/application/ports/provider.go +++ b/apps/gateway/internal/application/ports/provider.go @@ -2,6 +2,7 @@ package ports import ( "context" + "io" "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) @@ -24,3 +25,22 @@ type GenerationStream interface { type VerifiedTransportProofProvider interface { VerifiedTransportProof() *domain.TinfoilTransportProof } + +// ProxyProvider forwards the client's OpenAI chat-completions request body +// upstream (with the model and the few documented edits the adapter makes) +// and returns the provider's response as it came. Providers that speak the OpenAI +// wire format implement it; the gateway then acts as a proxy instead of +// translating through canonical events. +type ProxyProvider interface { + // ProxyStream sends a streaming request and returns the SSE body. + ProxyStream(ctx context.Context, raw []byte, model domain.PublicModel) (ProxyResponse, error) + // ProxyJSON sends a non-streaming request and returns the JSON body. + ProxyJSON(ctx context.Context, raw []byte, model domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) +} + +// ProxyResponse is an upstream streaming response. +type ProxyResponse struct { + Body io.ReadCloser + // Proof is set when the transport was verified before content was sent. + Proof *domain.TinfoilTransportProof +} diff --git a/apps/gateway/internal/application/services/chat_completions_service.go b/apps/gateway/internal/application/services/chat_completions_service.go index 7e5a9a6..97f4c2c 100644 --- a/apps/gateway/internal/application/services/chat_completions_service.go +++ b/apps/gateway/internal/application/services/chat_completions_service.go @@ -28,3 +28,9 @@ func (s *ChatCompletionsService) ExecuteStream(ctx context.Context, req domain.G stream, _, model, err := s.generate.ExecuteStream(ctx, domain.EndpointChatCompletions, req, bearerToken) return stream, model, err } + +// Proxy runs a /v1/chat/completions request as a proxy; it returns +// ErrNotProxyable when the request needs the translating path. +func (s *ChatCompletionsService) Proxy(ctx context.Context, raw []byte, req domain.GenerateRequest, bearerToken string) (*ProxyCall, error) { + return s.generate.Proxy(ctx, domain.EndpointChatCompletions, raw, req, bearerToken) +} diff --git a/apps/gateway/internal/application/services/proxy.go b/apps/gateway/internal/application/services/proxy.go new file mode 100644 index 0000000..b26a440 --- /dev/null +++ b/apps/gateway/internal/application/services/proxy.go @@ -0,0 +1,507 @@ +package services + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "time" + + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/middleware" + "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" +) + +// ErrNotProxyable means a request needs the translating path: PII masking +// rewrites content, and providers without the OpenAI wire format (Anthropic) +// need translation. It is returned before anything is reserved or sent. +var ErrNotProxyable = errors.New("request needs the translating path") + +// maxSSELine bounds one SSE line (a whole tool call can arrive in one). +const maxSSELine = 16 << 20 + +// ProxyCall is a request the gateway proxies: the provider's response goes to +// the client as it came, except the model field, which names the public model. +type ProxyCall struct { + Model domain.PublicModel + // Stream is set for streaming requests. + Stream *ProxyStream + // Body is the response of a non-streaming request. + Body []byte +} + +// Proxy runs a chat-completions request as a proxy. The gateway authenticates, +// routes, validates, reserves credit, and meters usage it reads from the +// response; it never rebuilds the request or the response. +func (s *GenerateService) Proxy(ctx context.Context, endpoint string, raw []byte, req domain.GenerateRequest, bearerToken string) (*ProxyCall, error) { + start := time.Now() + requestID := middleware.GetRequestID(ctx) + + authCtx, err := s.auth.AuthenticateAPIKey(ctx, bearerToken) + if err != nil { + s.recordTerminalOutcome(metrics.OutcomeError, endpoint, req, nil, time.Since(start).Milliseconds()) + return nil, err + } + if s.masksPII(authCtx.APIKey.PIIMode) { + return nil, ErrNotProxyable + } + model, execReq, err := s.resolveModel(ctx, req) + if err != nil { + s.recordTerminalOutcome(metrics.OutcomeError, endpoint, req, nil, time.Since(start).Milliseconds()) + return nil, err + } + if _, ok := s.proxyProvider(model); !ok { + return nil, ErrNotProxyable + } + if model.Fallback != nil { + if _, ok := s.proxyProvider(withProviderTarget(model, *model.Fallback)); !ok { + model.Fallback = nil // Only fall back to a provider that can proxy. + } + } + if err := s.validateRequest(endpoint, execReq, model); err != nil { + s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) + s.recordFailure(ctx, nil, &authCtx, endpoint, &execReq, &model, err, nil, time.Since(start).Milliseconds()) + return nil, err + } + reservationID, err := s.metering.Reserve(ctx, authCtx, endpoint, execReq, model, uuid.NewString()) + if err != nil { + s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) + return nil, err + } + finish := func(target domain.PublicModel, o outcome) { + s.finishProxy(ctx, authCtx, endpoint, execReq, target, requestID, reservationID, start, o) + } + + attempt := func(target domain.PublicModel) (*ProxyCall, error) { + provider, _ := s.proxyProvider(target) + upstreamStart := time.Now() + if !execReq.Stream { + body, proof, err := provider.ProxyJSON(ctx, raw, target) + recordUpstreamLatency(execReq.PublicModelID, target.ProviderConfig.ProviderName, upstreamStart, err) + if err != nil { + return nil, err + } + rewritten, o, err := observeJSON(body, target.PublicModelID) + if err != nil { + return nil, err + } + o.proof = proof + finish(target, o) + return &ProxyCall{Model: target, Body: rewritten}, nil + } + resp, err := provider.ProxyStream(ctx, raw, target) + recordUpstreamLatency(execReq.PublicModelID, target.ProviderConfig.ProviderName, upstreamStart, err) + if err != nil { + return nil, err + } + stream := newProxyStream(resp.Body, target.PublicModelID) + stream.proof = resp.Proof + stream.hideUsage = !wantsUsage(raw) + // Nothing has reached the client yet: an upstream that fails before its + // first output can still fall back. + if err := stream.prime(); err != nil { + stream.body.Close() + return nil, err + } + stream.onEnd = func(o outcome) { finish(target, o) } + return &ProxyCall{Model: target, Stream: stream}, nil + } + + call, err := attempt(model) + executed := model + if err != nil && shouldTryFallback(ctx, err, model.Fallback) { + executed = withProviderTarget(model, *model.Fallback) + s.logFallback(requestID, model, executed, err) + call, err = attempt(executed) + } + if err != nil { + fields := s.buildErrorLogFields(ctx, requestID, &authCtx, endpoint, execReq.PublicModelID, executed, err, time.Since(start).Milliseconds()) + s.logger.Error("provider proxy failed", fields...) + s.recordGeneration(metrics.OutcomeError, endpoint, execReq, executed, time.Since(start).Milliseconds()) + s.recordFailure(ctx, &reservationID, &authCtx, endpoint, &execReq, &executed, err, nil, time.Since(start).Milliseconds()) + return nil, err + } + return call, nil +} + +// masksPII reports whether this key's requests have content masked: only +// when the key asks for it and the filter is on (it is off in production). +func (s *GenerateService) masksPII(piiMode string) bool { + mode, ok := domain.NormalizeAPIKeyPIIMode(piiMode) + return (!ok || mode != domain.APIKeyPIIModeOff) && s.pii != nil && s.pii.Enabled() +} + +func (s *GenerateService) proxyProvider(model domain.PublicModel) (ports.ProxyProvider, bool) { + provider, err := s.registry.GetProvider(model.ProviderConfig.ProviderName) + if err != nil { + return nil, false + } + proxy, ok := provider.(ports.ProxyProvider) + return proxy, ok +} + +// outcome is what the gateway read from a proxied response. +type outcome struct { + providerID string + finish *string + usage *domain.Usage + proof *domain.TinfoilTransportProof + err error // Set when the response failed. + end string // How a stream ended, for logs. +} + +// finishProxy meters and logs a proxied response, like the translating path. +func (s *GenerateService) finishProxy(ctx context.Context, authCtx domain.AuthContext, endpoint string, req domain.GenerateRequest, model domain.PublicModel, requestID, reservationID string, start time.Time, o outcome) { + ctx = context.WithoutCancel(ctx) // Accounting outlives the client. + latencyMs := time.Since(start).Milliseconds() + fields := []any{ + "request_id", requestID, "reservation_id", reservationID, "provider_request_id", o.providerID, + "account_id", authCtx.Account.ID, "endpoint", endpoint, "model", req.PublicModelID, + "provider", model.ProviderConfig.ProviderName, "provider_model", model.UpstreamModelName, + "stream", req.Stream, "stream_end", o.end, "finish_reason", logfields.FinishReason(o.finish), + "usage_received", o.usage != nil, "usage", o.usage, "latency_ms", latencyMs, "proxy", true, + } + if o.err != nil { + s.logger.Error("proxied generation failed", append(fields, "error_type", fmt.Sprintf("%T", o.err))...) + s.recordGeneration(metrics.OutcomeError, endpoint, req, model, latencyMs) + s.recordFailure(ctx, &reservationID, &authCtx, endpoint, &req, &model, o.err, o.usage, latencyMs) + return + } + result := domain.GenerateResult{ + ID: o.providerID, PublicModelID: model.PublicModelID, ProviderName: model.ProviderConfig.ProviderName, + ProviderModelID: model.ProviderModelID, FinishReason: o.finish, Usage: o.usage, TinfoilProof: o.proof, + } + if result.TinfoilProof != nil && result.TinfoilProof.ProviderResponseID == "" { + result.TinfoilProof.ProviderResponseID = o.providerID + } + s.storeTinfoilProof(ctx, authCtx, model, result) + metrics.RecordUsage(o.usage, req.PublicModelID, model.ProviderConfig.ProviderName) + if o.usage == nil { + s.logger.Warn("proxied generation completed without usage", fields...) + } + s.recordGeneration(metrics.OutcomeSuccess, endpoint, req, model, latencyMs) + s.logger.Info("generation completed", append(fields, "gateway_status", 200, "upstream_status", 200)...) + if err := s.metering.RecordSuccess(ctx, reservationID, authCtx, endpoint, req, result, model, latencyMs); err != nil { + s.logger.Error("failed to record usage", "request_id", requestID, "reservation_id", reservationID, "error_type", fmt.Sprintf("%T", err)) + } +} + +// chunk is what the gateway reads from a provider's JSON: never content. +type chunk struct { + ID string `json:"id"` + Model *string `json:"model"` + Choices []struct { + Delta *struct { + Content *string `json:"content"` + ReasoningContent *string `json:"reasoning_content"` + Reasoning *string `json:"reasoning"` + ToolCalls []json.RawMessage `json:"tool_calls"` + } `json:"delta"` + FinishReason *string `json:"finish_reason"` + } `json:"choices"` + Usage *usageJSON `json:"usage"` + Error json.RawMessage `json:"error"` + BaseResp *struct { + StatusCode int `json:"status_code"` + StatusMsg string `json:"status_msg"` + } `json:"base_resp"` +} + +type usageJSON struct { + PromptTokens int64 `json:"prompt_tokens"` + CompletionTokens int64 `json:"completion_tokens"` + TotalTokens int64 `json:"total_tokens"` + PromptCacheHitTokens int64 `json:"prompt_cache_hit_tokens"` + PromptTokensDetails *struct { + CachedTokens int64 `json:"cached_tokens"` + } `json:"prompt_tokens_details"` +} + +func (u *usageJSON) domain() *domain.Usage { + if u == nil { + return nil + } + usage := &domain.Usage{PromptTokens: u.PromptTokens, CompletionTokens: u.CompletionTokens, TotalTokens: u.TotalTokens, CacheReadTokens: u.PromptCacheHitTokens} + if u.PromptTokensDetails != nil && u.PromptTokensDetails.CachedTokens > 0 { + usage.CacheReadTokens = u.PromptTokensDetails.CachedTokens + } + return usage +} + +// failure reports an error a provider sent in a 200 body or mid-stream. +func (c *chunk) failure() error { + if c.BaseResp != nil && c.BaseResp.StatusCode != 0 { + return domain.ErrProviderError(http.StatusBadGateway, "provider returned an unsuccessful response").WithMeta("upstream_code", c.BaseResp.StatusCode) + } + if len(c.Error) > 0 && string(c.Error) != "null" { + return domain.ErrProviderError(http.StatusBadGateway, "provider stream failed") + } + return nil +} + +// commits reports whether a chunk carries output, after which the response +// belongs to the client and can no longer fall back. +func (c *chunk) commits() bool { + for _, choice := range c.Choices { + if choice.FinishReason != nil { + return true + } + if d := choice.Delta; d != nil && (len(d.ToolCalls) > 0 || + (d.Content != nil && *d.Content != "") || + (d.ReasoningContent != nil && *d.ReasoningContent != "") || + (d.Reasoning != nil && *d.Reasoning != "")) { + return true + } + } + return false +} + +func (c *chunk) finishReason() *string { + for _, choice := range c.Choices { + if choice.FinishReason != nil && *choice.FinishReason != "" { + return choice.FinishReason + } + } + return nil +} + +// rewriteModel names the public model in a JSON object. Only the value of +// the top-level "model" field changes; every other byte stays as sent. +func rewriteModel(raw []byte, c *chunk, public string) []byte { + if c.Model == nil || *c.Model == public { + return raw + } + dec := json.NewDecoder(bytes.NewReader(raw)) + if tok, err := dec.Token(); err != nil || tok != json.Delim('{') { + return raw + } + for dec.More() { + key, err := dec.Token() + if err != nil { + return raw + } + afterKey := int(dec.InputOffset()) + var value json.RawMessage + if err := dec.Decode(&value); err != nil { + return raw + } + if key != "model" { + continue + } + end := int(dec.InputOffset()) + start := afterKey + for start < end && (raw[start] == ':' || raw[start] == ' ' || raw[start] == '\t' || raw[start] == '\n' || raw[start] == '\r') { + start++ + } + name, _ := json.Marshal(public) + out := make([]byte, 0, len(raw)+len(name)) + out = append(out, raw[:start]...) + out = append(out, name...) + return append(out, raw[end:]...) + } + return raw +} + +// observeJSON reads a non-streaming response. +func observeJSON(body []byte, public string) ([]byte, outcome, error) { + var c chunk + if err := json.Unmarshal(body, &c); err != nil { + return nil, outcome{}, domain.ErrProviderError(http.StatusBadGateway, "provider returned an invalid response") + } + if err := c.failure(); err != nil { + return nil, outcome{}, err + } + return rewriteModel(body, &c, public), outcome{providerID: c.ID, finish: c.finishReason(), usage: c.Usage.domain(), end: "complete"}, nil +} + +// ProxyStream relays a provider's SSE stream line by line, reading usage and +// the finish reason as they pass. Lines are forwarded exactly, except the +// model field of JSON events. +type ProxyStream struct { + body io.ReadCloser + scanner *bufio.Scanner + public string + pending [][]byte + proof *domain.TinfoilTransportProof + onEnd func(outcome) + + // hideUsage drops usage-only events the client didn't ask for (the + // gateway requests them to meter); usage is still read from them. + hideUsage bool + skipBlank bool + + o outcome + sawDone bool + finished bool // Metered. + chunks int +} + +// wantsUsage reports whether the client asked for usage in its stream. +func wantsUsage(raw []byte) bool { + var req struct { + StreamOptions *struct { + IncludeUsage bool `json:"include_usage"` + } `json:"stream_options"` + } + return json.Unmarshal(raw, &req) == nil && req.StreamOptions != nil && req.StreamOptions.IncludeUsage +} + +// NewProxyStream relays body, naming the public model. The service creates +// streams; it is exported for handler tests. +func NewProxyStream(body io.ReadCloser, public string) *ProxyStream { + return newProxyStream(body, public) +} + +func newProxyStream(body io.ReadCloser, public string) *ProxyStream { + scanner := bufio.NewScanner(body) + scanner.Buffer(make([]byte, 0, 64<<10), maxSSELine) + return &ProxyStream{body: body, scanner: scanner, public: public} +} + +// prime reads until the first output, so an upstream that fails first can +// fall back without the client seeing anything. +func (p *ProxyStream) prime() error { + for p.scanner.Scan() { + line, c := p.observe(p.scanner.Bytes()) + if p.keep(line, c) { + p.pending = append(p.pending, line) + } + if p.o.err != nil { + return p.o.err + } + if p.sawDone || (c != nil && c.commits()) { + return nil + } + } + if err := p.scanner.Err(); err != nil { + return domain.ErrProviderError(http.StatusBadGateway, "provider stream failed before output") + } + if p.chunks == 0 { + return domain.ErrProviderError(http.StatusBadGateway, "provider returned an empty stream") + } + return nil +} + +// observe reads one line; data events are parsed for metering and get the +// public model name. +func (p *ProxyStream) observe(raw []byte) ([]byte, *chunk) { + line := append([]byte(nil), raw...) + data, ok := bytes.CutPrefix(line, []byte("data:")) + if !ok { + return line, nil + } + data = bytes.TrimPrefix(data, []byte(" ")) + if string(bytes.TrimSpace(data)) == "[DONE]" { + p.sawDone = true + return line, nil + } + var c chunk + if json.Unmarshal(data, &c) != nil { + return line, nil + } + p.chunks++ + if c.ID != "" { + p.o.providerID = c.ID + } + if u := c.Usage.domain(); u != nil { + p.o.usage = u + } + if f := c.finishReason(); f != nil { + p.o.finish = f + } + if err := c.failure(); err != nil && p.o.err == nil { + p.o.err = err + } + prefix := line[:len(line)-len(data)] + return append(append([]byte(nil), prefix...), rewriteModel(data, &c, p.public)...), &c +} + +// keep reports whether a line goes to the client: everything does, except a +// usage-only event the client didn't ask for, with its blank separator. +func (p *ProxyStream) keep(line []byte, c *chunk) bool { + if p.skipBlank && len(line) == 0 { + p.skipBlank = false + return false + } + p.skipBlank = false + if p.hideUsage && c != nil && len(c.Choices) == 0 && c.Usage != nil && len(c.Error) == 0 { + p.skipBlank = true + return false + } + return true +} + +// Next returns the next line to send, without its newline. At the end it +// returns io.EOF (clean) or the read error. +func (p *ProxyStream) Next() ([]byte, error) { + if len(p.pending) > 0 { + line := p.pending[0] + p.pending = p.pending[1:] + return line, nil + } + for p.scanner.Scan() { + if line, c := p.observe(p.scanner.Bytes()); p.keep(line, c) { + return line, nil + } + } + err := p.scanner.Err() + switch { + case p.o.err != nil: + p.end("provider_error") + case err == nil && (p.sawDone || p.o.finish != nil): + p.end("complete") + case err == nil: + // The provider closed without a finish reason or [DONE]: the output + // may be cut short; bill what it reported, and log it. + p.end("eof_without_finish") + case p.o.finish != nil: + p.end("read_error_after_finish") // Complete; only trailing bytes were lost. + default: + p.o.err = domain.ErrProviderError(http.StatusBadGateway, "provider stream failed") + p.end("read_error") + } + if err == nil { + err = io.EOF + } + return nil, err +} + +// Complete reports whether the provider finished the response cleanly. +func (p *ProxyStream) Complete() bool { return p.sawDone || p.o.finish != nil } + +// Finished reports whether a finish reason has passed. +func (p *ProxyStream) Finished() bool { return p.o.finish != nil } + +// SawDone reports whether the provider sent [DONE]. +func (p *ProxyStream) SawDone() bool { return p.sawDone } + +// Close releases the upstream. Closing before the end is a cancellation, +// unless the response had already finished. +func (p *ProxyStream) Close() error { + if !p.finished { + if p.o.finish == nil && p.o.err == nil { + p.o.err = domain.ErrClientCanceled() + } + p.end("closed_before_eof") + } + return p.body.Close() +} + +func (p *ProxyStream) end(how string) { + if p.finished { + return + } + p.finished = true + p.o.end = how + p.o.proof = p.proof + if p.onEnd != nil { + p.onEnd(p.o) + } +} diff --git a/apps/gateway/internal/application/services/proxy_test.go b/apps/gateway/internal/application/services/proxy_test.go new file mode 100644 index 0000000..ee0b71b --- /dev/null +++ b/apps/gateway/internal/application/services/proxy_test.go @@ -0,0 +1,262 @@ +package services + +import ( + "context" + "errors" + "io" + "strings" + "testing" + + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" +) + +// proxyStub is an OpenAI-compatible provider that returns canned bodies. +type proxyStub struct { + stubProvider + sse string + json string + err error + calls int + lastRaw string +} + +func (p *proxyStub) ProxyStream(_ context.Context, raw []byte, _ domain.PublicModel) (ports.ProxyResponse, error) { + p.calls++ + p.lastRaw = string(raw) + if p.err != nil { + return ports.ProxyResponse{}, p.err + } + return ports.ProxyResponse{Body: io.NopCloser(strings.NewReader(p.sse))}, nil +} + +func (p *proxyStub) ProxyJSON(_ context.Context, raw []byte, _ domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) { + p.calls++ + p.lastRaw = string(raw) + return []byte(p.json), nil, p.err +} + +// resultMeter records what was metered. +type resultMeter struct { + stubUsageMeter + result *domain.GenerateResult + failed error +} + +func (m *resultMeter) RecordSuccess(ctx context.Context, id string, a domain.AuthContext, e string, req domain.GenerateRequest, res domain.GenerateResult, model domain.PublicModel, l int64) error { + m.result = &res + return m.stubUsageMeter.RecordSuccess(ctx, id, a, e, req, res, model, l) +} + +func (m *resultMeter) RecordFailure(ctx context.Context, id *string, a *domain.AuthContext, e string, req *domain.GenerateRequest, model *domain.PublicModel, err error, u *domain.Usage, l int64) error { + m.failed = err + return m.stubUsageMeter.RecordFailure(ctx, id, a, e, req, model, err, u, l) +} + +func proxyService(meter *resultMeter, pii string, providers map[string]ports.GenerationProvider, fallback bool) *GenerateService { + model := domain.PublicModel{ + PublicModelID: "zai-org/glm-5.3-flash", ProviderModelID: "pm", UpstreamModelName: "glm-5-3-flash", + ProviderConfig: domain.ProviderConfig{ProviderName: "primary"}, + SupportsChatCompletions: true, SupportsChatCompletionsStream: true, SupportsTools: true, SupportsParallelToolCalls: true, + MaxContextWindow: 1000, MaxOutputTokens: 100, + } + if fallback { + model.Fallback = &domain.ProviderTarget{ProviderModelID: "fm", UpstreamModelName: "glm-fallback", ProviderConfig: domain.ProviderConfig{ProviderName: "fallback"}} + } + var filter ports.PIIFilter + if pii != "" { + filter = &fakePIIFilter{enabled: true} + } + return NewGenerateService( + &stubAuthService{authCtx: domain.AuthContext{Account: domain.Account{ID: "acc1", Status: domain.AccountStatusActive}, APIKey: domain.APIKey{ID: "key1", Active: true, PIIMode: pii}}}, + &stubModelCatalog{model: model}, nil, &stubProviderRegistry{providers: providers}, meter, filter, stubLogger{}, + ) +} + +func relayAll(t *testing.T, s *ProxyStream) (string, error) { + t.Helper() + var b strings.Builder + for { + line, err := s.Next() + if err != nil { + if err == io.EOF { + err = nil + } + return b.String(), err + } + b.Write(line) + b.WriteByte('\n') + } +} + +func streamRequest() domain.GenerateRequest { + return domain.GenerateRequest{PublicModelID: "zai-org/glm-5.3-flash", Stream: true} +} + +// The provider's stream reaches the client byte for byte (the shapes that +// broke the translating path included), except the model name; usage and +// the finish reason are metered from it. +func TestProxyStream_ForwardsProviderBytes(t *testing.T) { + upstream := ": OPENROUTER PROCESSING\n\n" + + `data: {"id":"c1","model":"glm-5-3-flash","choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"Think &"},"logprobs":null,"finish_reason":null}],"system_fingerprint":"x"}` + "\n\n" + + `data:{"id":"c1","model":"glm-5-3-flash","choices":[{"index":0,"delta":{"content":"Reading.","tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"read","arguments":"{\"p\":1}"}},{"index":1,"id":"b","type":"function","function":{"name":"list","arguments":""}}]},"finish_reason":"tool_calls"}]}` + "\n\n" + + `data: {"id":"c1","model":"glm-5-3-flash","choices":[],"usage":{"prompt_tokens":10,"completion_tokens":5,"total_tokens":15,"prompt_tokens_details":{"cached_tokens":4}}}` + "\n\n" + + "data: [DONE]\n\n" + primary := &proxyStub{sse: upstream} + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary}, false) + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{"model":"zai-org/glm-5.3-flash","stream_options":{"include_usage":true}}`), streamRequest(), "key") + if err != nil { + t.Fatal(err) + } + got, err := relayAll(t, call.Stream) + if err != nil { + t.Fatal(err) + } + want := strings.ReplaceAll(upstream, `"model":"glm-5-3-flash"`, `"model":"zai-org/glm-5.3-flash"`) + if got != want { + t.Fatalf("relayed stream differs:\n got %q\nwant %q", got, want) + } + call.Stream.Close() + if meter.result == nil || meter.result.Usage == nil || meter.result.Usage.PromptTokens != 10 || meter.result.Usage.CacheReadTokens != 4 || + meter.result.FinishReason == nil || *meter.result.FinishReason != "tool_calls" || meter.result.ID != "c1" || meter.successCalls != 1 || meter.failureCalls != 0 { + t.Fatalf("metering: %+v success=%d failure=%d", meter.result, meter.successCalls, meter.failureCalls) + } +} + +func TestProxyStream_FallsBackBeforeFirstOutput(t *testing.T) { + primary := &proxyStub{sse: "data: {\"id\":\"x\",\"choices\":[{\"delta\":{\"role\":\"assistant\"}}]}\n\ndata: {\"error\":{\"message\":\"overloaded\"}}\n\n"} + fallback := &proxyStub{sse: "data: {\"id\":\"y\",\"model\":\"glm-fallback\",\"choices\":[{\"delta\":{\"content\":\"hi\"},\"finish_reason\":\"stop\"}]}\n\ndata: [DONE]\n\n"} + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary, "fallback": fallback}, true) + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key") + if err != nil { + t.Fatal(err) + } + got, _ := relayAll(t, call.Stream) + if strings.Contains(got, "overloaded") || !strings.Contains(got, `"content":"hi"`) || call.Model.ProviderConfig.ProviderName != "fallback" { + t.Fatalf("fallback not used cleanly: %q", got) + } + if meter.successCalls != 1 { + t.Fatalf("success=%d", meter.successCalls) + } +} + +func TestProxyStream_ErrorAfterOutputIsForwardedAndNotBilledAsSuccess(t *testing.T) { + primary := &proxyStub{sse: "data: {\"choices\":[{\"delta\":{\"content\":\"partial\"}}]}\n\ndata: {\"error\":{\"message\":\"overloaded\"}}\n\n"} + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary}, false) + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key") + if err != nil { + t.Fatal(err) + } + got, _ := relayAll(t, call.Stream) + if !strings.Contains(got, "overloaded") || call.Stream.Complete() { + t.Fatalf("error not relayed: %q", got) + } + if meter.failed == nil || meter.successCalls != 0 { + t.Fatalf("failure not recorded") + } +} + +func TestProxyStream_EmptyStreamIsAProviderError(t *testing.T) { + primary := &proxyStub{sse: ""} + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary}, false) + if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key"); err == nil { + t.Fatal("empty stream accepted") + } + if meter.failureCalls != 1 { + t.Fatalf("reservation not released") + } +} + +func TestProxyStream_ClientLeaving(t *testing.T) { + body := "data: {\"choices\":[{\"delta\":{\"content\":\"a\"}}]}\n\ndata: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":1}}\n\ndata: [DONE]\n\n" + for name, lines := range map[string]int{"before finish": 1, "after finish": 3} { + t.Run(name, func(t *testing.T) { + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": &proxyStub{sse: body}}, false) + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key") + if err != nil { + t.Fatal(err) + } + for range lines { + call.Stream.Next() + } + call.Stream.Close() + if name == "before finish" && (meter.failed == nil || meter.successCalls != 0) { + t.Fatal("abandoned stream billed as success") + } + if name == "after finish" && (meter.successCalls != 1 || meter.result.Usage == nil) { + t.Fatal("finished stream not billed") + } + }) + } +} + +func TestProxyJSON(t *testing.T) { + primary := &proxyStub{json: `{"id":"r1","model":"glm-5-3-flash","choices":[{"message":{"role":"assistant","content":"hi","reasoning_content":"why"},"finish_reason":"stop"}],"usage":{"prompt_tokens":2,"completion_tokens":1}}`} + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary}, false) + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{PublicModelID: "zai-org/glm-5.3-flash"}, "key") + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(call.Body), `"model":"zai-org/glm-5.3-flash"`) || !strings.Contains(string(call.Body), `"reasoning_content":"why"`) { + t.Fatalf("body %s", call.Body) + } + if meter.result == nil || meter.result.Usage.PromptTokens != 2 { + t.Fatal("usage not metered") + } + primary.json = `{"base_resp":{"status_code":1008,"status_msg":"insufficient balance"}}` + if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{PublicModelID: "zai-org/glm-5.3-flash"}, "key"); err == nil { + t.Fatal("error body returned as success") + } +} + +func TestProxy_PIIKeysAndOtherWireFormatsTranslate(t *testing.T) { + meter := &resultMeter{} + svc := proxyService(meter, "balanced", map[string]ports.GenerationProvider{"primary": &proxyStub{}}, false) + if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key"); !errors.Is(err, ErrNotProxyable) { + t.Fatalf("PII key proxied: %v", err) + } + svc = proxyService(meter, "", map[string]ports.GenerationProvider{"primary": &stubProvider{}}, false) + if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key"); !errors.Is(err, ErrNotProxyable) { + t.Fatalf("non-OpenAI provider proxied: %v", err) + } + if meter.reserveCalls != 0 { + t.Fatal("reserved before choosing the path") + } +} + +// With the PII filter off (as in production), a key's PII mode masks +// nothing, so its requests are proxied too. +func TestProxy_PIIModeWithFilterOffProxies(t *testing.T) { + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": &proxyStub{json: `{"id":"r","choices":[]}`}}, false) + svc.auth = &stubAuthService{authCtx: domain.AuthContext{Account: domain.Account{ID: "acc1", Status: domain.AccountStatusActive}, APIKey: domain.APIKey{ID: "key1", Active: true, PIIMode: "high"}}} + if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{PublicModelID: "zai-org/glm-5.3-flash"}, "key"); err != nil { + t.Fatalf("not proxied: %v", err) + } +} + +// The gateway asks every provider for usage, to meter it; a client that +// didn't ask gets the stream it would get from OpenAI, without that event. +func TestProxyStream_UnrequestedUsageIsMeteredNotForwarded(t *testing.T) { + body := "data: {\"choices\":[{\"delta\":{\"content\":\"a\"},\"finish_reason\":\"stop\"}]}\n\ndata: {\"choices\":[],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":1}}\n\ndata: [DONE]\n\n" + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": &proxyStub{sse: body}}, false) + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{"stream":true}`), streamRequest(), "key") + if err != nil { + t.Fatal(err) + } + got, _ := relayAll(t, call.Stream) + want := "data: {\"choices\":[{\"delta\":{\"content\":\"a\"},\"finish_reason\":\"stop\"}]}\n\ndata: [DONE]\n\n" + if got != want { + t.Fatalf("got %q", got) + } + if meter.result == nil || meter.result.Usage == nil || meter.result.Usage.PromptTokens != 3 { + t.Fatal("usage not metered") + } +} From 20c981b309226174145d631da6d221f0e746ade1 Mon Sep 17 00:00:00 2001 From: Marketen Date: Tue, 29 Sep 2026 10:35:52 +0200 Subject: [PATCH 3/7] fix: build the rewritten event without size arithmetic Co-Authored-By: Claude Opus 5.5 --- apps/gateway/internal/application/services/proxy.go | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/apps/gateway/internal/application/services/proxy.go b/apps/gateway/internal/application/services/proxy.go index b26a440..136218e 100644 --- a/apps/gateway/internal/application/services/proxy.go +++ b/apps/gateway/internal/application/services/proxy.go @@ -301,10 +301,7 @@ func rewriteModel(raw []byte, c *chunk, public string) []byte { start++ } name, _ := json.Marshal(public) - out := make([]byte, 0, len(raw)+len(name)) - out = append(out, raw[:start]...) - out = append(out, name...) - return append(out, raw[end:]...) + return bytes.Join([][]byte{raw[:start], name, raw[end:]}, nil) } return raw } From 0293742785d9b86738b646694745f60dc1947c04 Mon Sep 17 00:00:00 2001 From: Marketen Date: Tue, 29 Sep 2026 11:09:52 +0200 Subject: [PATCH 4/7] refactor: one request path, and PII masking on the proxy Every request now takes the proxy path: the translating path (request rebuild, canonical stream events, per-chunk re-encoding) is removed. PII masking, the reason that path was kept, now works on the raw OpenAI JSON: the LLM-bound text of the request is masked in place, and placeholders in the response are restored as events pass, holding back text that may end in a split placeholder. Also removed: the unused Anthropic adapter, the response DTOs and mapper, the canonical stream/output/tool types, and control-plane leftovers in pkg/domain (promos, usage queries, catalog/pricing types). The request mapper now only validates and reads what routing, feature checks and metering need. Production code: 13.2k -> 9.8k lines. Co-Authored-By: Claude Opus 5.5 --- .../http/dto/chat_completion_request.go | 75 -- .../http/dto/chat_completion_response.go | 66 -- .../http/handlers/chat_completions_handler.go | 221 +---- .../handlers/chat_completions_stream_test.go | 189 ---- .../adapters/http/mapper/chat_mapper_test.go | 359 ++------ .../adapters/http/mapper/chat_to_domain.go | 356 +++----- .../adapters/http/mapper/domain_to_chat.go | 158 ---- .../internal/adapters/http/sse/writer.go | 43 - .../adapters/providers/anthropic/adapter.go | 184 ---- .../providers/anthropic/adapter_error_test.go | 38 - .../adapters/providers/anthropic/client.go | 138 --- .../adapters/providers/anthropic/mapper.go | 152 ---- .../adapters/providers/anthropic/stream.go | 218 ----- .../adapters/providers/openai/adapter.go | 214 ----- .../providers/openai/adapter_error_test.go | 6 +- .../providers/openai/adapter_usage_test.go | 39 - .../adapters/providers/openai/mapper.go | 208 ----- .../adapters/providers/openai/policy_test.go | 336 +------ .../adapters/providers/openai/proxy.go | 32 +- .../adapters/providers/openai/proxy_test.go | 18 + .../adapters/providers/openai/stream.go | 438 --------- .../providers/openai/stream_diagnostics.go | 95 -- .../openai/stream_diagnostics_test.go | 133 --- .../adapters/providers/openai/stream_test.go | 231 ----- .../providers/openai/stream_toolcalls_test.go | 155 ---- .../adapters/providers/registry/registry.go | 12 +- .../adapters/providers/tinfoil/adapter.go | 102 +-- .../providers/tinfoil/adapter_error_test.go | 6 +- .../providers/tinfoil/adapter_test.go | 99 +- .../internal/application/ports/provider.go | 42 +- .../application/ports/provider_registry.go | 2 +- .../services/chat_completions_service.go | 18 +- .../application/services/generate_service.go | 431 +-------- .../services/generate_service_test.go | 213 ++--- .../internal/application/services/pii.go | 698 ++++++++++++++ .../application/services/pii_masking.go | 798 ---------------- .../application/services/pii_masking_test.go | 851 ------------------ .../internal/application/services/pii_test.go | 308 +++++++ .../internal/application/services/proxy.go | 147 +-- .../application/services/proxy_test.go | 59 +- .../services/stream_logging_test.go | 119 --- apps/gateway/main.go | 6 +- pkg/domain/account.go | 14 - pkg/domain/domain_test.go | 4 - pkg/domain/error.go | 44 +- pkg/domain/generation.go | 67 +- pkg/domain/input.go | 17 - pkg/domain/output.go | 16 - pkg/domain/pricing.go | 14 - pkg/domain/promo.go | 127 --- pkg/domain/promo_test.go | 95 -- pkg/domain/provider_catalog.go | 16 - pkg/domain/provider_model.go | 12 - pkg/domain/public_model_cp.go | 29 - pkg/domain/router.go | 18 - pkg/domain/stream.go | 24 - pkg/domain/tee.go | 16 - pkg/domain/tool.go | 31 +- pkg/domain/usage_query.go | 80 -- 59 files changed, 1492 insertions(+), 7145 deletions(-) delete mode 100644 apps/gateway/internal/adapters/http/dto/chat_completion_request.go delete mode 100644 apps/gateway/internal/adapters/http/dto/chat_completion_response.go delete mode 100644 apps/gateway/internal/adapters/http/mapper/domain_to_chat.go delete mode 100644 apps/gateway/internal/adapters/providers/anthropic/adapter.go delete mode 100644 apps/gateway/internal/adapters/providers/anthropic/adapter_error_test.go delete mode 100644 apps/gateway/internal/adapters/providers/anthropic/client.go delete mode 100644 apps/gateway/internal/adapters/providers/anthropic/mapper.go delete mode 100644 apps/gateway/internal/adapters/providers/anthropic/stream.go delete mode 100644 apps/gateway/internal/adapters/providers/openai/adapter_usage_test.go delete mode 100644 apps/gateway/internal/adapters/providers/openai/stream.go delete mode 100644 apps/gateway/internal/adapters/providers/openai/stream_diagnostics.go delete mode 100644 apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go delete mode 100644 apps/gateway/internal/adapters/providers/openai/stream_test.go delete mode 100644 apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go create mode 100644 apps/gateway/internal/application/services/pii.go delete mode 100644 apps/gateway/internal/application/services/pii_masking.go delete mode 100644 apps/gateway/internal/application/services/pii_masking_test.go create mode 100644 apps/gateway/internal/application/services/pii_test.go delete mode 100644 apps/gateway/internal/application/services/stream_logging_test.go delete mode 100644 pkg/domain/input.go delete mode 100644 pkg/domain/output.go delete mode 100644 pkg/domain/pricing.go delete mode 100644 pkg/domain/promo.go delete mode 100644 pkg/domain/promo_test.go delete mode 100644 pkg/domain/provider_catalog.go delete mode 100644 pkg/domain/provider_model.go delete mode 100644 pkg/domain/public_model_cp.go delete mode 100644 pkg/domain/stream.go delete mode 100644 pkg/domain/usage_query.go diff --git a/apps/gateway/internal/adapters/http/dto/chat_completion_request.go b/apps/gateway/internal/adapters/http/dto/chat_completion_request.go deleted file mode 100644 index 2f06762..0000000 --- a/apps/gateway/internal/adapters/http/dto/chat_completion_request.go +++ /dev/null @@ -1,75 +0,0 @@ -package dto - -import "encoding/json" - -// ChatCompletionRequest is the DTO for POST /v1/chat/completions. -type ChatCompletionRequest struct { - Model string `json:"model"` - Messages []ChatMessage `json:"messages"` - Stream *bool `json:"stream,omitempty"` - MaxTokens *int `json:"max_tokens,omitempty"` - MaxCompletionTokens *int `json:"max_completion_tokens,omitempty"` - N *int `json:"n,omitempty"` - Temperature *float64 `json:"temperature,omitempty"` - ReasoningEffort *string `json:"reasoning_effort,omitempty"` - TopP *float64 `json:"top_p,omitempty"` - Stop json.RawMessage `json:"stop,omitempty"` - Tools []ToolDefinition `json:"tools,omitempty"` - ToolChoice json.RawMessage `json:"tool_choice,omitempty"` - ParallelToolCalls *bool `json:"parallel_tool_calls,omitempty"` - User *string `json:"user,omitempty"` - Metadata map[string]any `json:"metadata,omitempty"` - ResponseFormat *ResponseFormat `json:"response_format,omitempty"` - ProviderOptions map[string]any `json:"provider_options,omitempty"` - PresencePenalty *float64 `json:"presence_penalty,omitempty"` - FrequencyPenalty *float64 `json:"frequency_penalty,omitempty"` - LogitBias map[string]int `json:"logit_bias,omitempty"` - Seed *int `json:"seed,omitempty"` - Logprobs *bool `json:"logprobs,omitempty"` - TopLogprobs *int `json:"top_logprobs,omitempty"` - Store *bool `json:"store,omitempty"` - ServiceTier *string `json:"service_tier,omitempty"` -} - -// ChatMessage is a message in the chat completions format. -type ChatMessage struct { - Role string `json:"role"` - Content json.RawMessage `json:"content,omitempty"` - ReasoningContent *string `json:"reasoning_content,omitempty"` - ToolCalls []ChatToolCall `json:"tool_calls,omitempty"` - ToolCallID *string `json:"tool_call_id,omitempty"` - Name *string `json:"name,omitempty"` -} - -// ChatToolCall is a tool call in assistant messages. -type ChatToolCall struct { - ID string `json:"id"` - Type string `json:"type"` - Function ChatFunctionCall `json:"function"` -} - -// ChatFunctionCall holds the function call details. -type ChatFunctionCall struct { - Name string `json:"name"` - Arguments string `json:"arguments"` -} - -// ResponseFormat for structured outputs. -type ResponseFormat struct { - Type string `json:"type"` - JSONSchema map[string]any `json:"json_schema,omitempty"` -} - -// ToolDefinition describes a tool in the chat completions format. -type ToolDefinition struct { - Type string `json:"type"` - Function ToolFunction `json:"function"` -} - -// ToolFunction holds the function details inside a ToolDefinition. -type ToolFunction struct { - Name string `json:"name"` - Description string `json:"description,omitempty"` - Parameters map[string]any `json:"parameters,omitempty"` - Strict *bool `json:"strict,omitempty"` -} diff --git a/apps/gateway/internal/adapters/http/dto/chat_completion_response.go b/apps/gateway/internal/adapters/http/dto/chat_completion_response.go deleted file mode 100644 index 80183f0..0000000 --- a/apps/gateway/internal/adapters/http/dto/chat_completion_response.go +++ /dev/null @@ -1,66 +0,0 @@ -package dto - -// ChatCompletionResponse is the non-streaming response for POST /v1/chat/completions. -type ChatCompletionResponse struct { - ID string `json:"id"` - Object string `json:"object"` - Created int64 `json:"created"` - Model string `json:"model"` - Choices []ChatCompletionChoice `json:"choices"` - Usage *ChatCompletionUsage `json:"usage,omitempty"` -} - -// ChatCompletionChoice is a single choice in the response. -type ChatCompletionChoice struct { - Index int `json:"index"` - Message ChatChoiceMessage `json:"message"` - FinishReason *string `json:"finish_reason"` -} - -// ChatChoiceMessage is the assistant message in a choice. -type ChatChoiceMessage struct { - Role string `json:"role"` - Content *string `json:"content"` - ReasoningContent *string `json:"reasoning_content,omitempty"` - ToolCalls []ChatToolCall `json:"tool_calls,omitempty"` -} - -// ChatCompletionUsage holds token usage for chat completions. -type ChatCompletionUsage struct { - PromptTokens int64 `json:"prompt_tokens"` - CompletionTokens int64 `json:"completion_tokens"` - TotalTokens int64 `json:"total_tokens"` -} - -// ChatCompletionChunk is a streaming chunk for chat completions. -type ChatCompletionChunk struct { - ID string `json:"id"` - Object string `json:"object"` - Created int64 `json:"created"` - Model string `json:"model"` - Choices []ChatCompletionChunkChoice `json:"choices"` - Usage *ChatCompletionUsage `json:"usage,omitempty"` -} - -// ChatCompletionChunkChoice is a streaming choice delta. -type ChatCompletionChunkChoice struct { - Index int `json:"index"` - Delta ChatChunkDelta `json:"delta"` - FinishReason *string `json:"finish_reason"` -} - -// ChatChunkDelta holds the delta content in a streaming chunk. -type ChatChunkDelta struct { - Role string `json:"role,omitempty"` - Content *string `json:"content,omitempty"` - ReasoningContent *string `json:"reasoning_content,omitempty"` - ToolCalls []ChatToolCallChunk `json:"tool_calls,omitempty"` -} - -// ChatToolCallChunk is a partial tool call in a stream chunk. -type ChatToolCallChunk struct { - Index int `json:"index"` - ID string `json:"id,omitempty"` - Type string `json:"type,omitempty"` - Function ChatFunctionCall `json:"function"` -} 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 8c69eea..2ad27f9 100644 --- a/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go +++ b/apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go @@ -2,7 +2,6 @@ package handlers import ( "context" - "errors" "io" "net/http" "sync" @@ -15,7 +14,6 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/services" "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" - "github.com/google/uuid" ) // Bound the wait for usage after generation has finished, without extending @@ -52,9 +50,7 @@ func (h *ChatCompletionsHandler) Handle(w http.ResponseWriter, r *http.Request) } // The gateway is a proxy: the request goes upstream as the client sent it - // and the provider's response comes back as it came. Only requests that - // need content rewritten (PII masking) or another wire format take the - // translating path below. + // and the provider's response comes back as it came. upstream, cancelUpstream := context.WithCancel(context.WithoutCancel(r.Context())) defer cancelUpstream() var finished atomic.Bool @@ -69,42 +65,17 @@ func (h *ChatCompletionsHandler) Handle(w http.ResponseWriter, r *http.Request) }) defer stopWatch() call, err := h.service.Proxy(upstream, body, genReq, token) - if err == nil { - if call.Stream != nil { - h.relay(w, r, call.Stream, &finished) - return - } - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusOK) - _, _ = w.Write(call.Body) - return - } - if !errors.Is(err, services.ErrNotProxyable) { + if err != nil { WriteErrorWithLog(w, r, h.logger, err) return } - - if fields := mapper.UnknownChatCompletionFields(body); len(fields) > 0 { - h.logger.Warn("chat completion request ignored unknown fields", - "request_id", middleware.GetRequestID(r.Context()), - "path", r.URL.Path, - "fields", fields, - ) - } - - if genReq.Stream { - h.handleStream(w, r, genReq, token) + if call.Stream != nil { + h.relay(w, r, call.Stream, &finished) return } - - result, err := h.service.Execute(r.Context(), genReq, token) - if err != nil { - WriteErrorWithLog(w, r, h.logger, err) - return - } - - resp := mapper.DomainToChatCompletionResponse(result) - WriteJSON(w, http.StatusOK, resp) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write(call.Body) } // relay sends a provider's stream to the client line by line, as it came. @@ -169,181 +140,3 @@ func (h *ChatCompletionsHandler) relay(w http.ResponseWriter, r *http.Request, s write(nil) } } - -func (h *ChatCompletionsHandler) handleStream(w http.ResponseWriter, r *http.Request, genReq domain.GenerateRequest, token string) { - 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, cancelUpstream) -} - -func (h *ChatCompletionsHandler) writeStream( - w http.ResponseWriter, - r *http.Request, - genReq domain.GenerateRequest, - stream ports.GenerationStream, - model *domain.PublicModel, - cancelUpstream context.CancelFunc, -) { - responseModelID := genReq.PublicModelID - if model != nil { - responseModelID = model.PublicModelID - } - - sw, err := sse.NewWriter(w) - if err != nil { - WriteError(w, domain.ErrInternal("an internal error occurred")) - return - } - - responseID := uuid.New().String()[:12] - createdAt := time.Now().Unix() - - // Keep-alive: write SSE comments periodically to prevent proxy timeouts. - var writeMu sync.Mutex - stopKeepAlive := make(chan struct{}) - defer close(stopKeepAlive) - go func() { - ticker := time.NewTicker(15 * time.Second) - defer ticker.Stop() - for { - select { - case <-stopKeepAlive: - return - case <-ticker.C: - writeMu.Lock() - sw.WriteComment("keepalive") - writeMu.Unlock() - } - } - }() - - for { - event, err := stream.Recv() - if err != nil { - if err == io.EOF { - writeMu.Lock() - sw.WriteDone() - writeMu.Unlock() - break - } - writeMu.Lock() - sw.WriteData(map[string]any{ - "error": map[string]any{"type": "internal_error", "message": err.Error()}, - }) - writeMu.Unlock() - return - } - if event.ProviderResponseID != "" { - responseID = event.ProviderResponseID - } - if event.Type == domain.StreamEventError { - // A provider error after output started must not look like a - // complete response: no finish chunk and no [DONE]. - message := "the provider stream failed" - if event.Error != nil { - message = event.Error.Error() - } - writeMu.Lock() - sw.WriteData(map[string]any{ - "error": map[string]any{"type": "provider_error", "message": message}, - }) - writeMu.Unlock() - return - } - write := func(event domain.StreamEvent) { - chunk, _ := mapper.DomainStreamEventToChatChunk(event, responseModelID, responseID, createdAt) - if chunk == nil { - return - } - writeMu.Lock() - 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) - } - } - - // 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 { - // Output the finishing chunk carried goes out now; only the finish - // signal waits. - for _, delta := range splitDeltas(&event) { - write(delta) - } - 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 - } - // Output after the finish reason is still the response. - for _, delta := range splitDeltas(&tail) { - write(delta) - } - } - timer.Stop() - } - chunk, done := mapper.DomainStreamEventToChatChunk(event, responseModelID, responseID, createdAt) - if chunk != nil { - writeMu.Lock() - 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() - writeErr := sw.WriteDone() - writeMu.Unlock() - 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 - } - } -} - -// splitDeltas moves the text, reasoning, and tool-call output out of an event -// into delta events of their own, leaving the rest (finish, usage) in place. -func splitDeltas(event *domain.StreamEvent) []domain.StreamEvent { - var out []domain.StreamEvent - content := event.ContentDelta != nil && *event.ContentDelta != "" - reasoning := event.ReasoningDelta != nil && *event.ReasoningDelta != "" - if content || reasoning { - delta := domain.StreamEvent{Type: domain.StreamEventOutputTextDelta} - if content { - delta.ContentDelta = event.ContentDelta - } - if reasoning { - delta.ReasoningDelta = event.ReasoningDelta - } - out = append(out, delta) - } - if event.ToolCallDelta != nil { - out = append(out, domain.StreamEvent{Type: domain.StreamEventToolCallDelta, ToolCallDelta: event.ToolCallDelta}) - } - event.ContentDelta, event.ReasoningDelta, event.ToolCallDelta = nil, nil, nil - return out -} 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 25eb290..d0e14d1 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,205 +1,16 @@ package handlers import ( - "context" "errors" "io" "net/http/httptest" "strings" "sync/atomic" "testing" - "time" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/services" - - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) -type handlerTestStream struct { - 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++ - return event, nil - } - if s.err != nil { - return domain.StreamEvent{}, s.err - } - 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) { - content := "hello" - finishReason := "stop" - tests := []struct { - name string - events []domain.StreamEvent - }{ - { - name: "upstream done marker becomes EOF", - events: []domain.StreamEvent{{ - Type: domain.StreamEventOutputTextDelta, - ContentDelta: &content, - }}, - }, - { - name: "completed event followed by EOF", - events: []domain.StreamEvent{{ - Type: domain.StreamEventCompleted, - FinishReason: &finishReason, - }}, - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - handler := &ChatCompletionsHandler{logger: &confidentialTestLogger{}} - recorder := httptest.NewRecorder() - request := httptest.NewRequest("POST", "/v1/chat/completions", 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()) - } - }) - } -} - -func TestChatCompletionsHandler_StreamErrorDoesNotClaimCompletion(t *testing.T) { - handler := &ChatCompletionsHandler{logger: &confidentialTestLogger{}} - 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, 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 } - -func TestChatCompletionsHandler_ForwardsOutputAroundTheFinishReason(t *testing.T) { - finish, early, late, args := "tool_calls", "Reading.", "late text", "{}" - id, name := "call_1", "read" - stream := &handlerTestStream{events: []domain.StreamEvent{ - {Type: domain.StreamEventCompleted, FinishReason: &finish, ContentDelta: &early}, - {Type: domain.StreamEventToolCallDelta, ToolCallDelta: &domain.ToolCallDelta{Index: 0, ID: &id, Name: &name, ArgumentsDelta: &args}}, - {Type: domain.StreamEventOutputTextDelta, ContentDelta: &late}, - }} - recorder := httptest.NewRecorder() - handler := &ChatCompletionsHandler{logger: &confidentialTestLogger{}} - handler.writeStream(recorder, httptest.NewRequest("POST", "/v1/chat/completions", nil), domain.GenerateRequest{PublicModelID: "test"}, stream, nil, func() {}) - body := recorder.Body.String() - order := []string{`"content":"Reading."`, `"id":"call_1"`, `"content":"late text"`, `"finish_reason":"tool_calls"`, "[DONE]"} - at := 0 - for _, want := range order { - i := strings.Index(body[at:], want) - if i < 0 { - t.Fatalf("missing %s in order: %s", want, body) - } - at += i - } -} - -func TestChatCompletionsHandler_ProviderErrorEventIsNotCompletion(t *testing.T) { - content := "partial" - stream := &handlerTestStream{events: []domain.StreamEvent{ - {Type: domain.StreamEventOutputTextDelta, ContentDelta: &content}, - {Type: domain.StreamEventError, Error: domain.ErrProviderError(502, "overloaded")}, - }} - recorder := httptest.NewRecorder() - handler := &ChatCompletionsHandler{logger: &confidentialTestLogger{}} - handler.writeStream(recorder, httptest.NewRequest("POST", "/v1/chat/completions", nil), domain.GenerateRequest{PublicModelID: "test"}, stream, nil, func() {}) - body := recorder.Body.String() - if strings.Contains(body, "[DONE]") || strings.Contains(body, `"finish_reason":"`) || !strings.Contains(body, `"error"`) { - t.Fatalf("error event reported as completion: %s", body) - } -} - type brokenBody struct{ io.Reader } func (brokenBody) Close() error { return nil } diff --git a/apps/gateway/internal/adapters/http/mapper/chat_mapper_test.go b/apps/gateway/internal/adapters/http/mapper/chat_mapper_test.go index ca9d14a..1dfc9e2 100644 --- a/apps/gateway/internal/adapters/http/mapper/chat_mapper_test.go +++ b/apps/gateway/internal/adapters/http/mapper/chat_mapper_test.go @@ -1,318 +1,65 @@ -package mapper_test +package mapper import ( - "encoding/json" + "errors" + "strings" "testing" - "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/mapper" "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) -func TestChatCompletionRequestToDomain_Basic(t *testing.T) { - raw := json.RawMessage(`{ - "model": "openai/gpt-4.1-mini", - "messages": [ - {"role": "user", "content": "Hello"} - ] - }`) - - req, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - if req.PublicModelID != "openai/gpt-4.1-mini" { - t.Errorf("model = %s, want openai/gpt-4.1-mini", req.PublicModelID) - } - if len(req.Input) != 1 { - t.Fatalf("input len = %d, want 1", len(req.Input)) - } - if *req.Input[0].Content != "Hello" { - t.Errorf("content = %s, want Hello", *req.Input[0].Content) - } - if *req.Input[0].Role != "user" { - t.Errorf("role = %s, want user", *req.Input[0].Role) - } -} - -func TestChatCompletionRequestToDomain_DeveloperRole(t *testing.T) { - raw := json.RawMessage(`{ - "model": "openai/gpt-5", - "messages": [ - {"role": "developer", "content": "Follow these instructions"} - ] - }`) - - req, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(req.Input) != 1 { - t.Fatalf("input len = %d, want 1", len(req.Input)) - } - if *req.Input[0].Role != "developer" { - t.Fatalf("role = %s, want developer", *req.Input[0].Role) - } - if *req.Input[0].Content != "Follow these instructions" { - t.Fatalf("content = %s, want developer content", *req.Input[0].Content) - } -} - -func TestChatCompletionRequestToDomain_EmptyMessages(t *testing.T) { - raw := json.RawMessage(`{"model": "test", "messages": []}`) - _, err := mapper.ChatCompletionRequestToDomain(raw) - if err == nil { - t.Fatal("expected error for empty messages") - } -} - -func TestChatCompletionRequestToDomain_BothMaxTokens(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "max_tokens": 100, - "max_completion_tokens": 200 - }`) - _, err := mapper.ChatCompletionRequestToDomain(raw) - if err == nil { - t.Fatal("expected error when both max_tokens and max_completion_tokens are set") - } -} - -func TestChatCompletionRequestToDomain_MaxTokensOnly(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "max_tokens": 100 - }`) - req, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if req.MaxOutputTokens == nil || *req.MaxOutputTokens != 100 { - t.Errorf("max_output_tokens = %v, want 100", req.MaxOutputTokens) - } -} - -func TestChatCompletionRequestToDomain_MaxCompletionTokensOnly(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "max_completion_tokens": 200 - }`) - req, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if req.MaxOutputTokens == nil || *req.MaxOutputTokens != 200 { - t.Errorf("max_output_tokens = %v, want 200", req.MaxOutputTokens) - } -} - -func TestChatCompletionRequestToDomain_ReasoningEffort(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "reasoning_effort": "low" - }`) - req, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if req.ReasoningEffort == nil || *req.ReasoningEffort != "low" { - t.Errorf("reasoning_effort = %v, want low", req.ReasoningEffort) - } - if fields := mapper.UnknownChatCompletionFields(raw); len(fields) != 0 { - t.Fatalf("unknown fields = %v, want none", fields) - } -} - -func TestChatCompletionRequestToDomain_UnknownField(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "foobar": true - }`) - _, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unknown fields should be silently ignored, got: %v", err) - } -} - -func TestUnknownChatCompletionFields(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "zeta": true, - "n": 1, - "best_of": 2, - "alpha": "ignored" - }`) - fields := mapper.UnknownChatCompletionFields(raw) - if len(fields) != 2 || fields[0] != "alpha" || fields[1] != "zeta" { - t.Fatalf("fields = %v, want [alpha zeta]", fields) - } -} - -func TestChatCompletionRequestToDomain_NOne(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "n": 1 - }`) - _, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unexpected error for n=1: %v", err) - } -} - -func TestChatCompletionRequestToDomain_NMultipleUnsupported(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "n": 2 - }`) - _, err := mapper.ChatCompletionRequestToDomain(raw) - if err == nil { - t.Fatal("expected error for n > 1") - } -} - -func TestChatCompletionRequestToDomain_ToolMessage(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [ - {"role": "user", "content": "What is the weather?"}, - {"role": "assistant", "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "get_weather", "arguments": "{\"location\":\"NYC\"}"}}]}, - {"role": "tool", "tool_call_id": "call_1", "content": "72F sunny"} - ] - }`) - - req, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(req.Input) != 3 { - t.Fatalf("input len = %d, want 3", len(req.Input)) - } - if len(req.Input[1].ToolCalls) != 1 { - t.Fatalf("assistant tool_calls len = %d, want 1", len(req.Input[1].ToolCalls)) - } - if *req.Input[2].ToolCallID != "call_1" { - t.Errorf("tool_call_id = %s, want call_1", *req.Input[2].ToolCallID) - } -} - -func TestChatCompletionRequestToDomain_AssistantReasoningContent(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [ - {"role": "assistant", "reasoning_content": "thoughts", "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}]} - ] - }`) - - req, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if req.Input[0].ReasoningContent == nil || *req.Input[0].ReasoningContent != "thoughts" { - t.Fatalf("reasoning_content = %v, want thoughts", req.Input[0].ReasoningContent) - } +func read(t *testing.T, body string) (domain.GenerateRequest, error) { + t.Helper() + return ChatCompletionRequestToDomain([]byte(body)) } -func TestChatCompletionRequestToDomain_ToolMessageMissingID(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [ - {"role": "tool", "content": "result"} - ] - }`) - _, err := mapper.ChatCompletionRequestToDomain(raw) - if err == nil { - t.Fatal("expected error for tool message without tool_call_id") - } -} - -func TestChatCompletionRequestToDomain_SystemWithToolCalls(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [ - {"role": "system", "content": "You are helpful", "tool_calls": [{"id": "x", "type": "function", "function": {"name": "f", "arguments": "{}"}}]} - ] - }`) - _, err := mapper.ChatCompletionRequestToDomain(raw) - if err == nil { - t.Fatal("expected error for system message with tool_calls") - } -} - -func TestChatCompletionRequestToDomain_Stop(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "stop": ["END", "STOP"] - }`) - req, err := mapper.ChatCompletionRequestToDomain(raw) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(req.Stop) != 2 { - t.Errorf("stop len = %d, want 2", len(req.Stop)) - } -} - -func TestChatCompletionRequestToDomain_StopString(t *testing.T) { - raw := json.RawMessage(`{ - "model": "test", - "messages": [{"role": "user", "content": "hi"}], - "stop": "END" - }`) - req, err := mapper.ChatCompletionRequestToDomain(raw) +func TestChatCompletionRequestToDomain_ReadsTheSummary(t *testing.T) { + req, err := read(t, `{"model":"m","stream":true,"max_completion_tokens":50,"parallel_tool_calls":false, + "response_format":{"type":"json_object"},"temperature":0.2,"unknown_extension":{"x":1}, + "messages":[{"role":"developer","content":"be brief"}, + {"role":"user","content":[{"type":"text","text":"What is "},{"type":"image_url","image_url":{"url":"data:image/png;base64,AA"}},{"type":"text","text":"this?"}]}, + {"role":"assistant","content":null,"reasoning_content":"r","tool_calls":[{"id":"a","type":"function","function":{"name":"f","arguments":"{}"}}]}, + {"role":"tool","tool_call_id":"a","content":"42"}], + "tools":[{"type":"function","function":{"name":"f","parameters":{"type":"object"}}}]}`) if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(req.Stop) != 1 || req.Stop[0] != "END" { - t.Errorf("stop = %v, want [END]", req.Stop) - } -} - -func TestDomainToChatCompletionResponse(t *testing.T) { - content := "Hello!" - reasoning := "thinking" - role := "assistant" - finishReason := "stop" - result := domain.GenerateResult{ - ID: "test-id-123", - CreatedUnix: 1700000000, - PublicModelID: "openai/gpt-4.1-mini", - Output: []domain.OutputItem{ - {Type: "message", Role: &role, Content: &content, ReasoningContent: &reasoning}, - }, - FinishReason: &finishReason, - Usage: &domain.Usage{ - PromptTokens: 10, - CompletionTokens: 5, - TotalTokens: 15, - }, - } - - resp := mapper.DomainToChatCompletionResponse(result) - if resp.Object != "chat.completion" { - t.Errorf("object = %s, want chat.completion", resp.Object) - } - if len(resp.Choices) != 1 { - t.Fatalf("choices len = %d, want 1", len(resp.Choices)) - } - if *resp.Choices[0].Message.Content != "Hello!" { - t.Errorf("content = %s, want Hello!", *resp.Choices[0].Message.Content) - } - if resp.Choices[0].Message.ReasoningContent == nil || *resp.Choices[0].Message.ReasoningContent != "thinking" { - t.Errorf("reasoning_content = %v, want thinking", resp.Choices[0].Message.ReasoningContent) - } - if *resp.Choices[0].FinishReason != "stop" { - t.Errorf("finish_reason = %s, want stop", *resp.Choices[0].FinishReason) - } - if resp.Usage.TotalTokens != 15 { - t.Errorf("total_tokens = %d, want 15", resp.Usage.TotalTokens) + t.Fatal(err) + } + if req.PublicModelID != "m" || !req.Stream || *req.MaxOutputTokens != 50 || *req.ParallelToolCalls || !req.StructuredOutput { + t.Fatalf("summary %+v", req) + } + if len(req.Input) != 4 || *req.Input[0].Role != "developer" || *req.Input[1].Content != "What is this?" || req.Input[2].Content != nil || *req.Input[3].Content != "42" { + t.Fatalf("router view %+v", req.Input) + } + if len(req.Tools) != 1 || req.Tools[0].Name != "f" { + t.Fatalf("tools %+v", req.Tools) + } + if req, _ := read(t, `{"model":"m","max_tokens":7,"messages":[{"role":"user","content":"hi"}],"response_format":{"type":"text"}}`); *req.MaxOutputTokens != 7 || req.StructuredOutput { + t.Fatalf("max_tokens / text format: %+v", req) + } +} + +func TestChatCompletionRequestToDomain_Rejects(t *testing.T) { + msg := `"messages":[{"role":"user","content":"hi"}]` + for body, want := range map[string]string{ + `{not json`: "invalid JSON body", + `{` + msg + `}`: "model is required", + `{"model":"m","messages":[]}`: "messages is required", + `{"model":"m",` + msg + `,"max_tokens":1,"max_completion_tokens":1}`: "both max_tokens", + `{"model":"m",` + msg + `,"n":2}`: "'n' only supports", + `{"model":"m",` + msg + `,"functions":[]}`: "'functions' is not supported", + `{"model":"m","messages":[{"role":"tool","content":"x"}]}`: "requires tool_call_id", + `{"model":"m","messages":[{"role":"system","content":"x","tool_calls":[]}]}`: "must not contain tool_calls", + `{"model":"m","messages":[{"role":"user"}]}`: "content is required", + `{"model":"m","messages":[{"role":"robot","content":"x"}]}`: "unsupported role", + `{"model":"m",` + msg + `,"tools":[{"type":"web","function":{"name":"f"}}]}`: "unsupported tool type", + `{"model":"m",` + msg + `,"tools":[{"type":"function","function":{}}]}`: "name is required", + `{"model":"m",` + msg + `,"tool_choice":"sometimes"}`: "invalid tool_choice", + `{"model":"m",` + msg + `,"tool_choice":{"type":"function","function":{}}}`: "function name is required", + } { + _, err := read(t, body) + var gwErr *domain.GatewayError + if !errors.As(err, &gwErr) || !strings.Contains(err.Error(), want) { + t.Errorf("%s: err = %v, want %q", body, err, want) + } } } diff --git a/apps/gateway/internal/adapters/http/mapper/chat_to_domain.go b/apps/gateway/internal/adapters/http/mapper/chat_to_domain.go index 971c2b8..3d2c23b 100644 --- a/apps/gateway/internal/adapters/http/mapper/chat_to_domain.go +++ b/apps/gateway/internal/adapters/http/mapper/chat_to_domain.go @@ -3,60 +3,69 @@ package mapper import ( "encoding/json" "fmt" - "sort" - "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/dto" "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) -// unsupportedChatFields are known top-level fields that must be rejected -// because they change semantics in ways the gateway cannot support. -var unsupportedChatFields = map[string]bool{ - "best_of": true, - "function_call": true, "functions": true, +// unsupportedChatFields change semantics in ways the gateway cannot support. +var unsupportedChatFields = []string{"best_of", "function_call", "functions"} + +// chatRequest is what the gateway reads from a chat-completions body. The +// body itself is forwarded as sent; this only validates it and extracts what +// routing, feature checks, and metering need. +type chatRequest struct { + Model string `json:"model"` + Messages []chatMessage `json:"messages"` + Stream *bool `json:"stream"` + MaxTokens *int `json:"max_tokens"` + MaxCompletionTokens *int `json:"max_completion_tokens"` + N *int `json:"n"` + Tools []chatTool `json:"tools"` + ToolChoice json.RawMessage `json:"tool_choice"` + ParallelToolCalls *bool `json:"parallel_tool_calls"` + ResponseFormat *struct { + Type string `json:"type"` + } `json:"response_format"` +} + +type chatMessage struct { + Role string `json:"role"` + Content json.RawMessage `json:"content"` + ToolCalls json.RawMessage `json:"tool_calls"` + ToolCallID *string `json:"tool_call_id"` } -// knownChatFields are allowed top-level fields. -var knownChatFields = map[string]bool{ - "model": true, "messages": true, "stream": true, - "max_tokens": true, "max_completion_tokens": true, "n": true, - "temperature": true, "reasoning_effort": true, "top_p": true, "stop": true, - "tools": true, "tool_choice": true, "parallel_tool_calls": true, - "user": true, "metadata": true, "response_format": true, - "provider_options": true, "stream_options": true, - // Passed through / silently ignored: - "presence_penalty": true, "frequency_penalty": true, "logit_bias": true, - "seed": true, "logprobs": true, "top_logprobs": true, - "suffix": true, "echo": true, "service_tier": true, "store": true, +type chatTool struct { + Type string `json:"type"` + Function struct { + Name string `json:"name"` + } `json:"function"` } -// ChatCompletionRequestToDomain maps a /v1/chat/completions request DTO to the canonical GenerateRequest. +// ChatCompletionRequestToDomain validates a /v1/chat/completions body and +// reads the request summary the gateway works with: the model, streaming, +// the output limit, the features used, and the text of each message and the +// tool names for the router. func ChatCompletionRequestToDomain(raw json.RawMessage) (domain.GenerateRequest, error) { - var rawMap map[string]json.RawMessage - if err := json.Unmarshal(raw, &rawMap); err != nil { + var fields map[string]json.RawMessage + if err := json.Unmarshal(raw, &fields); err != nil { return domain.GenerateRequest{}, domain.ErrInvalidField("invalid JSON body") } - - for key := range rawMap { - if unsupportedChatFields[key] { + for _, key := range unsupportedChatFields { + if _, ok := fields[key]; ok { return domain.GenerateRequest{}, domain.ErrInvalidField(fmt.Sprintf("field '%s' is not supported on /v1/chat/completions in this version", key)) } - // Unknown fields are silently ignored for client compatibility } - - var req dto.ChatCompletionRequest + var req chatRequest if err := json.Unmarshal(raw, &req); err != nil { return domain.GenerateRequest{}, domain.ErrInvalidField("invalid request body: " + err.Error()) } - if req.Model == "" { return domain.GenerateRequest{}, domain.ErrInvalidField("model is required") } - if len(req.Messages) == 0 { return domain.GenerateRequest{}, domain.ErrInvalidField("messages is required and cannot be empty") } - if req.MaxTokens != nil && req.MaxCompletionTokens != nil { return domain.GenerateRequest{}, domain.ErrInvalidField("cannot provide both max_tokens and max_completion_tokens") } @@ -64,271 +73,122 @@ func ChatCompletionRequestToDomain(raw json.RawMessage) (domain.GenerateRequest, return domain.GenerateRequest{}, domain.ErrInvalidField("field 'n' only supports value 1 on /v1/chat/completions in this version") } - input, err := mapChatMessages(req.Messages) - if err != nil { - return domain.GenerateRequest{}, err - } - gen := domain.GenerateRequest{ - PublicModelID: req.Model, - Input: input, - Temperature: req.Temperature, - ReasoningEffort: req.ReasoningEffort, - TopP: req.TopP, - Stream: req.Stream != nil && *req.Stream, - User: req.User, - Metadata: req.Metadata, - ProviderOptions: req.ProviderOptions, - PresencePenalty: req.PresencePenalty, - FrequencyPenalty: req.FrequencyPenalty, - LogitBias: req.LogitBias, - Seed: req.Seed, - Logprobs: req.Logprobs, - TopLogprobs: req.TopLogprobs, - Store: req.Store, - ServiceTier: req.ServiceTier, + PublicModelID: req.Model, + Stream: req.Stream != nil && *req.Stream, + MaxOutputTokens: req.MaxTokens, + ParallelToolCalls: req.ParallelToolCalls, + StructuredOutput: req.ResponseFormat != nil && (req.ResponseFormat.Type == "json_object" || req.ResponseFormat.Type == "json_schema"), } - - // Normalize max tokens - if req.MaxTokens != nil { - gen.MaxOutputTokens = req.MaxTokens - } else if req.MaxCompletionTokens != nil { + if gen.MaxOutputTokens == nil { gen.MaxOutputTokens = req.MaxCompletionTokens } - - if req.ParallelToolCalls != nil { - gen.ParallelToolCalls = req.ParallelToolCalls - } - - // Parse stop - if len(req.Stop) > 0 { - stops, err := parseStop(req.Stop) + for i, msg := range req.Messages { + item, err := readMessage(i, msg) if err != nil { return domain.GenerateRequest{}, err } - gen.Stop = stops + gen.Input = append(gen.Input, item) } - - // Map tools for _, t := range req.Tools { - td, err := mapToolDefinition(t) - if err != nil { - return domain.GenerateRequest{}, err + if t.Type != "" && t.Type != "function" { + return domain.GenerateRequest{}, domain.ErrInvalidField(fmt.Sprintf("unsupported tool type: %s", t.Type)) } - gen.Tools = append(gen.Tools, td) + if t.Function.Name == "" { + return domain.GenerateRequest{}, domain.ErrInvalidField("tool function name is required") + } + gen.Tools = append(gen.Tools, domain.ToolDefinition{Name: t.Function.Name}) } - - // Map tool_choice if len(req.ToolChoice) > 0 { - tc, err := parseToolChoice(req.ToolChoice) - if err != nil { + if err := checkToolChoice(req.ToolChoice); err != nil { return domain.GenerateRequest{}, err } - gen.ToolChoice = tc - } - - // Map response_format -> TextConfig - if req.ResponseFormat != nil { - gen.TextConfig = &domain.ResponseTextConfig{ - FormatType: &req.ResponseFormat.Type, - JSONSchema: req.ResponseFormat.JSONSchema, - } } - return gen, nil } -// UnknownChatCompletionFields returns top-level request fields the gateway does -// not currently understand. The mapper ignores these fields for client -// compatibility, but handlers can log them for visibility. -func UnknownChatCompletionFields(raw json.RawMessage) []string { - var rawMap map[string]json.RawMessage - if err := json.Unmarshal(raw, &rawMap); err != nil { - return nil - } - - fields := make([]string, 0) - for key := range rawMap { - if knownChatFields[key] || unsupportedChatFields[key] { - continue +// readMessage checks a message's shape and reads its role and text. +func readMessage(i int, msg chatMessage) (domain.InputItem, error) { + hasToolCalls := len(msg.ToolCalls) > 0 && string(msg.ToolCalls) != "null" + role := msg.Role + item := domain.InputItem{Role: &role} + switch msg.Role { + case "system", "developer", "user", "tool": + if hasToolCalls { + return item, domain.ErrInvalidField(fmt.Sprintf("message[%d]: %s message must not contain tool_calls", i, msg.Role)) } - fields = append(fields, key) - } - sort.Strings(fields) - return fields -} - -func mapChatMessages(messages []dto.ChatMessage) ([]domain.InputItem, error) { - var result []domain.InputItem - - for i, msg := range messages { - switch msg.Role { - case "system", "developer", "user": - if msg.ToolCalls != nil { - return nil, domain.ErrInvalidField(fmt.Sprintf("message[%d]: %s message must not contain tool_calls", i, msg.Role)) - } - if msg.ToolCallID != nil { - return nil, domain.ErrInvalidField(fmt.Sprintf("message[%d]: %s message must not contain tool_call_id", i, msg.Role)) - } - content, err := extractStringContent(msg.Content) - if err != nil { - return nil, domain.ErrInvalidField(fmt.Sprintf("message[%d]: %s", i, err.Error())) - } - role := msg.Role - result = append(result, domain.InputItem{ - Type: domain.InputItemTypeMessage, - Role: &role, - Content: &content, - }) - - case "assistant": - if msg.ToolCallID != nil { - return nil, domain.ErrInvalidField(fmt.Sprintf("message[%d]: assistant message must not contain tool_call_id", i)) - } - role := msg.Role - item := domain.InputItem{ - Type: domain.InputItemTypeMessage, - Role: &role, - } - // May have content - if len(msg.Content) > 0 { - content, err := extractStringContent(msg.Content) - if err == nil { - item.Content = &content - } - } - if msg.ReasoningContent != nil { - item.ReasoningContent = msg.ReasoningContent - } - // May have tool_calls - for _, tc := range msg.ToolCalls { - item.ToolCalls = append(item.ToolCalls, domain.ToolCall{ - ID: tc.ID, - Name: tc.Function.Name, - ArgumentsJSON: tc.Function.Arguments, - }) - } - result = append(result, item) - - case "tool": - if msg.ToolCallID == nil || *msg.ToolCallID == "" { - return nil, domain.ErrToolMessageInvalid(fmt.Sprintf("message[%d]: tool message requires tool_call_id", i)) - } - if msg.ToolCalls != nil { - return nil, domain.ErrInvalidField(fmt.Sprintf("message[%d]: tool message must not contain tool_calls", i)) - } - content, err := extractStringContent(msg.Content) - if err != nil { - return nil, domain.ErrInvalidField(fmt.Sprintf("message[%d]: %s", i, err.Error())) - } - role := msg.Role - result = append(result, domain.InputItem{ - Type: domain.InputItemTypeMessage, - Role: &role, - Content: &content, - ToolCallID: msg.ToolCallID, - }) - - default: - return nil, domain.ErrInvalidField(fmt.Sprintf("message[%d]: unsupported role '%s'", i, msg.Role)) + if msg.Role == "tool" && (msg.ToolCallID == nil || *msg.ToolCallID == "") { + return item, domain.ErrToolMessageInvalid(fmt.Sprintf("message[%d]: tool message requires tool_call_id", i)) + } + if msg.Role != "tool" && msg.ToolCallID != nil { + return item, domain.ErrInvalidField(fmt.Sprintf("message[%d]: %s message must not contain tool_call_id", i, msg.Role)) + } + text, err := textContent(msg.Content) + if err != nil { + return item, domain.ErrInvalidField(fmt.Sprintf("message[%d]: %s", i, err.Error())) + } + item.Content = &text + case "assistant": + if msg.ToolCallID != nil { + return item, domain.ErrInvalidField(fmt.Sprintf("message[%d]: assistant message must not contain tool_call_id", i)) } + if text, err := textContent(msg.Content); err == nil { + item.Content = &text + } + default: + return item, domain.ErrInvalidField(fmt.Sprintf("message[%d]: unsupported role '%s'", i, msg.Role)) } - - return result, nil + return item, nil } -func extractStringContent(raw json.RawMessage) (string, error) { - if len(raw) == 0 { +// textContent reads a message's text: a string, or the text parts of an +// array (images and other parts are forwarded, but carry no text). +func textContent(raw json.RawMessage) (string, error) { + if len(raw) == 0 || string(raw) == "null" { return "", fmt.Errorf("content is required") } - var s string if err := json.Unmarshal(raw, &s); err == nil { return s, nil } - - // Try array of content parts var parts []struct { Type string `json:"type"` Text string `json:"text"` } if err := json.Unmarshal(raw, &parts); err != nil { - return "", fmt.Errorf("content must be a string or array of text content parts") + return "", fmt.Errorf("content must be a string or array of content parts") } - - var combined string + var text string for _, p := range parts { - switch p.Type { - case "text", "input_text", "output_text": - combined += p.Text - default: - // Skip non-text content parts for now + if p.Type == "text" || p.Type == "input_text" || p.Type == "output_text" { + text += p.Text } } - return combined, nil + return text, nil } -func parseStop(raw json.RawMessage) ([]string, error) { - var s string - if err := json.Unmarshal(raw, &s); err == nil { - return []string{s}, nil - } - - var arr []string - if err := json.Unmarshal(raw, &arr); err != nil { - return nil, domain.ErrInvalidField("stop must be a string or array of strings") - } - return arr, nil -} - -func mapToolDefinition(t dto.ToolDefinition) (domain.ToolDefinition, error) { - if t.Type != "" && t.Type != "function" { - return domain.ToolDefinition{}, domain.ErrInvalidField(fmt.Sprintf("unsupported tool type: %s", t.Type)) - } - if t.Function.Name == "" { - return domain.ToolDefinition{}, domain.ErrInvalidField("tool function name is required") - } - strict := false - if t.Function.Strict != nil { - strict = *t.Function.Strict - } - return domain.ToolDefinition{ - Name: t.Function.Name, - Description: t.Function.Description, - Parameters: t.Function.Parameters, - Strict: strict, - }, nil -} - -func parseToolChoice(raw json.RawMessage) (*domain.ToolChoice, error) { - var s string - if err := json.Unmarshal(raw, &s); err == nil { - switch s { - case "none", "auto", "required": - return &domain.ToolChoice{Mode: s}, nil - default: - return nil, domain.ErrInvalidField(fmt.Sprintf("invalid tool_choice value: %s", s)) +func checkToolChoice(raw json.RawMessage) error { + var mode string + if err := json.Unmarshal(raw, &mode); err == nil { + if mode == "none" || mode == "auto" || mode == "required" { + return nil } + return domain.ErrInvalidField(fmt.Sprintf("invalid tool_choice value: %s", mode)) } - - var obj struct { + var named struct { Type string `json:"type"` Function struct { Name string `json:"name"` } `json:"function"` } - if err := json.Unmarshal(raw, &obj); err != nil { - return nil, domain.ErrInvalidField("tool_choice must be a string or object") + if err := json.Unmarshal(raw, &named); err != nil { + return domain.ErrInvalidField("tool_choice must be a string or object") } - - if obj.Type != "function" { - return nil, domain.ErrInvalidField(fmt.Sprintf("unsupported tool_choice type: %s", obj.Type)) + if named.Type != "function" { + return domain.ErrInvalidField(fmt.Sprintf("unsupported tool_choice type: %s", named.Type)) } - if obj.Function.Name == "" { - return nil, domain.ErrInvalidField("tool_choice function name is required") + if named.Function.Name == "" { + return domain.ErrInvalidField("tool_choice function name is required") } - return &domain.ToolChoice{ - Mode: domain.ToolChoiceFunction, - FunctionName: &obj.Function.Name, - }, nil + return nil } diff --git a/apps/gateway/internal/adapters/http/mapper/domain_to_chat.go b/apps/gateway/internal/adapters/http/mapper/domain_to_chat.go deleted file mode 100644 index 1c1dccd..0000000 --- a/apps/gateway/internal/adapters/http/mapper/domain_to_chat.go +++ /dev/null @@ -1,158 +0,0 @@ -package mapper - -import ( - "strings" - - "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/dto" - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" -) - -// DomainToChatCompletionResponse maps a canonical GenerateResult to a chat completions response DTO. -func DomainToChatCompletionResponse(result domain.GenerateResult) dto.ChatCompletionResponse { - resp := dto.ChatCompletionResponse{ - ID: chatCompletionID(result.ID), - Object: "chat.completion", - Created: result.CreatedUnix, - Model: result.PublicModelID, - } - - message := dto.ChatChoiceMessage{ - Role: "assistant", - } - - for _, out := range result.Output { - if out.Content != nil && *out.Content != "" { - message.Content = out.Content - } - if out.ReasoningContent != nil && *out.ReasoningContent != "" { - message.ReasoningContent = out.ReasoningContent - } - for _, tc := range out.ToolCalls { - message.ToolCalls = append(message.ToolCalls, dto.ChatToolCall{ - ID: tc.ID, - Type: "function", - Function: dto.ChatFunctionCall{ - Name: tc.Name, - Arguments: tc.ArgumentsJSON, - }, - }) - } - } - - choice := dto.ChatCompletionChoice{ - Index: 0, - Message: message, - FinishReason: result.FinishReason, - } - - resp.Choices = []dto.ChatCompletionChoice{choice} - - if result.Usage != nil { - resp.Usage = &dto.ChatCompletionUsage{ - PromptTokens: result.Usage.PromptTokens, - CompletionTokens: result.Usage.CompletionTokens, - TotalTokens: result.Usage.TotalTokens, - } - } - - return resp -} - -// DomainStreamEventToChatChunk maps a canonical StreamEvent to a chat completions streaming chunk. -func DomainStreamEventToChatChunk(event domain.StreamEvent, model string, responseID string, createdAt int64) (data *dto.ChatCompletionChunk, done bool) { - chunk := &dto.ChatCompletionChunk{ - ID: chatCompletionID(responseID), - Object: "chat.completion.chunk", - Created: createdAt, - Model: model, - } - - switch event.Type { - case domain.StreamEventOutputTextDelta: - choice := dto.ChatCompletionChunkChoice{ - Index: 0, - Delta: dto.ChatChunkDelta{ - Content: event.ContentDelta, - ReasoningContent: event.ReasoningDelta, - }, - } - if event.Role != nil { - choice.Delta.Role = *event.Role - } - chunk.Choices = []dto.ChatCompletionChunkChoice{choice} - return chunk, false - - case domain.StreamEventOutputMessageDelta: - choice := dto.ChatCompletionChunkChoice{ - Index: 0, - Delta: dto.ChatChunkDelta{}, - } - if event.Role != nil { - choice.Delta.Role = *event.Role - } - chunk.Choices = []dto.ChatCompletionChunkChoice{choice} - return chunk, false - - case domain.StreamEventToolCallDelta: - choice := dto.ChatCompletionChunkChoice{ - Index: 0, - Delta: dto.ChatChunkDelta{}, - } - if event.ToolCallDelta != nil { - tc := dto.ChatToolCallChunk{ - Index: event.ToolCallDelta.Index, - } - if event.ToolCallDelta.ID != nil { - tc.ID = *event.ToolCallDelta.ID - tc.Type = "function" - } - fn := dto.ChatFunctionCall{} - if event.ToolCallDelta.Name != nil { - fn.Name = *event.ToolCallDelta.Name - } - if event.ToolCallDelta.ArgumentsDelta != nil { - fn.Arguments = *event.ToolCallDelta.ArgumentsDelta - } - tc.Function = fn - choice.Delta.ToolCalls = []dto.ChatToolCallChunk{tc} - } - chunk.Choices = []dto.ChatCompletionChunkChoice{choice} - return chunk, false - - case domain.StreamEventCompleted: - delta := dto.ChatChunkDelta{} - if event.ContentDelta != nil { - delta.Content = event.ContentDelta - } - if event.ReasoningDelta != nil { - delta.ReasoningContent = event.ReasoningDelta - } - choice := dto.ChatCompletionChunkChoice{ - Index: 0, - Delta: delta, - FinishReason: event.FinishReason, - } - chunk.Choices = []dto.ChatCompletionChunkChoice{choice} - if event.Usage != nil { - chunk.Usage = &dto.ChatCompletionUsage{ - PromptTokens: event.Usage.PromptTokens, - CompletionTokens: event.Usage.CompletionTokens, - TotalTokens: event.Usage.TotalTokens, - } - } - return chunk, true - - case domain.StreamEventError: - return nil, true - - default: - return nil, false - } -} - -func chatCompletionID(id string) string { - if strings.HasPrefix(id, "chatcmpl-") { - return id - } - return "chatcmpl-" + id -} diff --git a/apps/gateway/internal/adapters/http/sse/writer.go b/apps/gateway/internal/adapters/http/sse/writer.go index d9b5f3d..0ab129f 100644 --- a/apps/gateway/internal/adapters/http/sse/writer.go +++ b/apps/gateway/internal/adapters/http/sse/writer.go @@ -1,7 +1,6 @@ package sse import ( - "encoding/json" "fmt" "net/http" ) @@ -29,48 +28,6 @@ func NewWriter(w http.ResponseWriter) (*Writer, error) { return &Writer{w: w, flusher: flusher}, nil } -// WriteEvent writes a named SSE event with JSON data. -func (sw *Writer) WriteEvent(eventType string, data any) error { - jsonData, err := json.Marshal(data) - if err != nil { - return fmt.Errorf("failed to marshal SSE data: %w", err) - } - - if eventType != "" { - if _, err := fmt.Fprintf(sw.w, "event: %s\n", eventType); err != nil { - return err - } - } - if _, err := fmt.Fprintf(sw.w, "data: %s\n\n", jsonData); err != nil { - return err - } - sw.flusher.Flush() - return nil -} - -// WriteData writes a data-only SSE event with JSON content. -func (sw *Writer) WriteData(data any) error { - jsonData, err := json.Marshal(data) - if err != nil { - return fmt.Errorf("failed to marshal SSE data: %w", err) - } - - if _, err := fmt.Fprintf(sw.w, "data: %s\n\n", jsonData); err != nil { - return err - } - sw.flusher.Flush() - return nil -} - -// WriteDone writes the [DONE] terminal marker. -func (sw *Writer) WriteDone() error { - if _, err := fmt.Fprint(sw.w, "data: [DONE]\n\n"); err != nil { - return err - } - sw.flusher.Flush() - return nil -} - // WriteComment writes an SSE comment (ignored by clients). Useful as a keep-alive. func (sw *Writer) WriteComment(text string) error { if _, err := fmt.Fprintf(sw.w, ": %s\n\n", text); err != nil { diff --git a/apps/gateway/internal/adapters/providers/anthropic/adapter.go b/apps/gateway/internal/adapters/providers/anthropic/adapter.go deleted file mode 100644 index 564e6a1..0000000 --- a/apps/gateway/internal/adapters/providers/anthropic/adapter.go +++ /dev/null @@ -1,184 +0,0 @@ -package anthropic - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "os" - "strings" - "time" - - "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" - "github.com/google/uuid" -) - -// Adapter is the Anthropic provider adapter. -type Adapter struct { - client *Client -} - -func NewAdapter(timeout time.Duration) *Adapter { - return &Adapter{client: NewClient(timeout)} -} - -func (a *Adapter) Generate(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) { - apiKey := resolveAPIKey(model.ProviderConfig.APIKeySecretRef) - if apiKey == "" { - return domain.GenerateResult{}, missingProviderCredentialError() - } - - body := buildRequestBody(req, model) - delete(body, "stream") - - respBody, err := a.client.Do(ctx, model.ProviderConfig.BaseURL, apiKey, body) - if err != nil { - return domain.GenerateResult{}, mapProviderError(err) - } - - return parseResponse(respBody, req, model) -} - -func (a *Adapter) StreamGenerate(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (ports.GenerationStream, error) { - apiKey := resolveAPIKey(model.ProviderConfig.APIKeySecretRef) - if apiKey == "" { - return nil, missingProviderCredentialError() - } - - body := buildRequestBody(req, model) - - resp, err := a.client.DoStream(ctx, model.ProviderConfig.BaseURL, apiKey, body) - if err != nil { - return nil, mapProviderError(err) - } - - return NewStream(resp), nil -} - -func missingProviderCredentialError() *domain.GatewayError { - return domain.ErrInternal("an internal error occurred").WithMeta( - "provider", "anthropic", - "reason", "provider API key is not configured", - ) -} - -func resolveAPIKey(secretRef string) string { - return os.Getenv(secretRef) -} - -func parseResponse(data json.RawMessage, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) { - var resp struct { - ID string `json:"id"` - Type string `json:"type"` - Role string `json:"role"` - Content []struct { - Type string `json:"type"` - Text string `json:"text,omitempty"` - ID string `json:"id,omitempty"` - Name string `json:"name,omitempty"` - Input any `json:"input,omitempty"` - } `json:"content"` - StopReason string `json:"stop_reason"` - Usage struct { - InputTokens int64 `json:"input_tokens"` - OutputTokens int64 `json:"output_tokens"` - CacheCreationInputTokens int64 `json:"cache_creation_input_tokens"` - CacheReadInputTokens int64 `json:"cache_read_input_tokens"` - } `json:"usage"` - } - - if err := json.Unmarshal(data, &resp); err != nil { - return domain.GenerateResult{}, fmt.Errorf("failed to parse provider response: %w", err) - } - - result := domain.GenerateResult{ - ID: resp.ID, - CreatedUnix: time.Now().Unix(), - PublicModelID: req.PublicModelID, - ProviderName: model.ProviderConfig.ProviderName, - ProviderModelID: model.ProviderModelID, - } - - if result.ID == "" { - result.ID = uuid.New().String() - } - - role := resp.Role - if role == "" { - role = "assistant" - } - - out := domain.OutputItem{ - Type: domain.OutputItemTypeMessage, - Role: &role, - } - - for _, block := range resp.Content { - switch block.Type { - case "text": - text := block.Text - out.Content = &text - case "tool_use": - argsJSON, _ := json.Marshal(block.Input) - out.ToolCalls = append(out.ToolCalls, domain.ToolCall{ - ID: block.ID, - Name: block.Name, - ArgumentsJSON: string(argsJSON), - }) - } - } - - result.Output = []domain.OutputItem{out} - - finishReason := mapStopReason(resp.StopReason) - result.FinishReason = &finishReason - - // Anthropic input_tokens = only uncached input. Normalize PromptTokens to total input. - totalInput := resp.Usage.InputTokens + resp.Usage.CacheCreationInputTokens + resp.Usage.CacheReadInputTokens - result.Usage = &domain.Usage{ - PromptTokens: totalInput, - CompletionTokens: resp.Usage.OutputTokens, - TotalTokens: totalInput + resp.Usage.OutputTokens, - CacheCreationTokens: resp.Usage.CacheCreationInputTokens, - CacheReadTokens: resp.Usage.CacheReadInputTokens, - } - - return result, nil -} - -func mapProviderError(err error) error { - if err == nil { - return nil - } - - var httpErr *ProviderHTTPError - if errors.As(err, &httpErr) { - var gwErr *domain.GatewayError - switch { - case httpErr.StatusCode == 401 || httpErr.StatusCode == 403: - gwErr = domain.ErrProviderError(502, fmt.Sprintf("provider auth/permission error: %s", httpErr.Body)) - case httpErr.StatusCode == 429: - gwErr = domain.ErrProviderError(429, "provider rate limited: "+httpErr.Body) - case httpErr.StatusCode == 503: - gwErr = domain.ErrProviderUnavailable("anthropic") - case httpErr.StatusCode >= 500: - gwErr = domain.ErrProviderError(502, fmt.Sprintf("provider server error: %s", httpErr.Body)) - default: - gwErr = domain.ErrProviderError(502, fmt.Sprintf("provider rejected request: %s", httpErr.Body)) - } - return gwErr.WithMeta( - "upstream_status", httpErr.StatusCode, - "upstream_error", httpErr.Body, - ) - } - - msg := err.Error() - if strings.Contains(msg, "timeout") || strings.Contains(msg, "deadline") { - return domain.ErrProviderTimeout("anthropic").WithMeta("upstream_error", msg) - } - if strings.Contains(msg, "connection refused") || strings.Contains(msg, "no such host") { - return domain.ErrProviderUnavailable("anthropic").WithMeta("upstream_error", msg) - } - return domain.ErrProviderError(502, msg).WithMeta("upstream_error", msg) -} diff --git a/apps/gateway/internal/adapters/providers/anthropic/adapter_error_test.go b/apps/gateway/internal/adapters/providers/anthropic/adapter_error_test.go deleted file mode 100644 index 01a5765..0000000 --- a/apps/gateway/internal/adapters/providers/anthropic/adapter_error_test.go +++ /dev/null @@ -1,38 +0,0 @@ -package anthropic - -import ( - "context" - "net/http" - "testing" - - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" -) - -func TestAdapterMissingCredentialReturnsInternalError(t *testing.T) { - adapter := NewAdapter(0) - model := domain.PublicModel{} - calls := []struct { - name string - call func() error - }{ - {name: "generate", call: func() error { - _, err := adapter.Generate(context.Background(), domain.GenerateRequest{}, model) - return err - }}, - {name: "stream", call: func() error { - _, err := adapter.StreamGenerate(context.Background(), domain.GenerateRequest{}, model) - return err - }}, - } - for _, tt := range calls { - t.Run(tt.name, func(t *testing.T) { - gatewayErr, ok := tt.call().(*domain.GatewayError) - if !ok || gatewayErr.HTTPStatus != http.StatusInternalServerError || gatewayErr.Type != domain.ErrTypeInternal || gatewayErr.Code != domain.ErrCodeInternalError { - t.Fatalf("error = %#v, want generic internal error", gatewayErr) - } - if gatewayErr.Message != "an internal error occurred" || gatewayErr.Metadata["provider"] != "anthropic" { - t.Fatalf("error = %#v, want generic response with provider log metadata", gatewayErr) - } - }) - } -} diff --git a/apps/gateway/internal/adapters/providers/anthropic/client.go b/apps/gateway/internal/adapters/providers/anthropic/client.go deleted file mode 100644 index 9121464..0000000 --- a/apps/gateway/internal/adapters/providers/anthropic/client.go +++ /dev/null @@ -1,138 +0,0 @@ -package anthropic - -import ( - "bytes" - "context" - "crypto/tls" - "encoding/json" - "fmt" - "io" - "net" - "net/http" - "time" - - enclavenetwork "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/enclave/network" -) - -// Client handles HTTP communication with the Anthropic API. -type Client struct { - httpClient *http.Client - responseTimeout time.Duration -} - -func NewClient(timeout time.Duration) *Client { - transport := &http.Transport{ - DialContext: (&net.Dialer{ - Timeout: 10 * time.Second, - KeepAlive: 30 * time.Second, - }).DialContext, - ForceAttemptHTTP2: true, - TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS12}, - TLSHandshakeTimeout: 10 * time.Second, - ResponseHeaderTimeout: timeout, - MaxIdleConns: 100, - MaxIdleConnsPerHost: 20, - IdleConnTimeout: 90 * time.Second, - ExpectContinueTimeout: 1 * time.Second, - } - enclavenetwork.ConfigureProviderTransport(transport) - return &Client{ - httpClient: &http.Client{Transport: transport}, - responseTimeout: timeout, - } -} - -// jsonUnmarshalBytes is used by mapper.go -func jsonUnmarshalBytes(data []byte, v any) error { - return json.Unmarshal(data, v) -} - -func (c *Client) Do(ctx context.Context, baseURL, apiKey string, body map[string]any) (json.RawMessage, error) { - ctx, cancel := context.WithTimeout(ctx, c.responseTimeout) - defer cancel() - - jsonBody, err := json.Marshal(body) - if err != nil { - return nil, fmt.Errorf("failed to marshal request: %w", err) - } - - endpoint := baseURL + "/v1/messages" - req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(jsonBody)) - if err != nil { - return nil, fmt.Errorf("failed to create request: %w", err) - } - req.Header.Set("Content-Type", "application/json") - req.Header.Set("x-api-key", apiKey) - req.Header.Set("anthropic-version", "2023-06-01") - - resp, err := c.httpClient.Do(req) - if err != nil { - return nil, fmt.Errorf("provider request failed: %w", err) - } - defer resp.Body.Close() - - respBody, err := io.ReadAll(io.LimitReader(resp.Body, 10*1024*1024)) - if err != nil { - return nil, fmt.Errorf("failed to read provider response: %w", err) - } - - if resp.StatusCode >= 400 { - return nil, parseProviderError(resp.StatusCode, respBody) - } - - return respBody, nil -} - -func (c *Client) DoStream(ctx context.Context, baseURL, apiKey string, body map[string]any) (*http.Response, error) { - jsonBody, err := json.Marshal(body) - if err != nil { - return nil, fmt.Errorf("failed to marshal request: %w", err) - } - - endpoint := baseURL + "/v1/messages" - req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(jsonBody)) - if err != nil { - return nil, fmt.Errorf("failed to create request: %w", err) - } - req.Header.Set("Content-Type", "application/json") - req.Header.Set("x-api-key", apiKey) - req.Header.Set("anthropic-version", "2023-06-01") - - resp, err := c.httpClient.Do(req) - if err != nil { - return nil, fmt.Errorf("provider stream request failed: %w", err) - } - - if resp.StatusCode >= 400 { - body, _ := io.ReadAll(io.LimitReader(resp.Body, 10*1024*1024)) - resp.Body.Close() - return nil, parseProviderError(resp.StatusCode, body) - } - - return resp, nil -} - -// ProviderHTTPError carries the upstream HTTP status code. -type ProviderHTTPError struct { - StatusCode int - Body string -} - -func (e *ProviderHTTPError) Error() string { - return fmt.Sprintf("provider error (%d): %s", e.StatusCode, e.Body) -} - -func parseProviderError(statusCode int, body []byte) error { - var errResp struct { - Error struct { - Type string `json:"type"` - Message string `json:"message"` - } `json:"error"` - } - msg := string(body) - if json.Unmarshal(body, &errResp) == nil && errResp.Error.Message != "" { - msg = errResp.Error.Message - } - - return &ProviderHTTPError{StatusCode: statusCode, Body: msg} -} diff --git a/apps/gateway/internal/adapters/providers/anthropic/mapper.go b/apps/gateway/internal/adapters/providers/anthropic/mapper.go deleted file mode 100644 index 9a77615..0000000 --- a/apps/gateway/internal/adapters/providers/anthropic/mapper.go +++ /dev/null @@ -1,152 +0,0 @@ -package anthropic - -import ( - "encoding/json" - - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" -) - -func buildRequestBody(req domain.GenerateRequest, model domain.PublicModel) map[string]any { - body := map[string]any{ - "model": model.UpstreamModelName, - } - - messages, system := buildMessages(req) - body["messages"] = messages - if system != "" { - body["system"] = system - } - - if req.MaxOutputTokens != nil { - v := *req.MaxOutputTokens - if model.MaxOutputTokens > 0 && v > model.MaxOutputTokens { - v = model.MaxOutputTokens - } - body["max_tokens"] = v - } else { - body["max_tokens"] = model.MaxOutputTokens - } - - if req.Temperature != nil { - body["temperature"] = *req.Temperature - } - if req.TopP != nil { - body["top_p"] = *req.TopP - } - if len(req.Stop) > 0 { - body["stop_sequences"] = req.Stop - } - if req.Stream { - body["stream"] = true - } - - if len(req.Tools) > 0 { - tools := make([]map[string]any, 0, len(req.Tools)) - for _, t := range req.Tools { - tool := map[string]any{ - "name": t.Name, - "description": t.Description, - "input_schema": t.Parameters, - } - tools = append(tools, tool) - } - body["tools"] = tools - } - - if req.ToolChoice != nil { - switch req.ToolChoice.Mode { - case domain.ToolChoiceAuto: - body["tool_choice"] = map[string]any{"type": "auto"} - case domain.ToolChoiceRequired: - body["tool_choice"] = map[string]any{"type": "any"} - case domain.ToolChoiceFunction: - body["tool_choice"] = map[string]any{ - "type": "tool", - "name": *req.ToolChoice.FunctionName, - } - case domain.ToolChoiceNone: - delete(body, "tools") - } - } - - return body -} - -func buildMessages(req domain.GenerateRequest) ([]map[string]any, string) { - var messages []map[string]any - var system string - - if req.Instructions != nil && *req.Instructions != "" { - system = *req.Instructions - } - - for _, item := range req.Input { - role := "user" - if item.Role != nil { - role = *item.Role - } - - if role == "system" || role == "developer" { - if item.Content != nil { - if system != "" { - system += "\n" - } - system += *item.Content - } - continue - } - - msg := map[string]any{ - "role": role, - } - - if role == "tool" && item.ToolCallID != nil { - msg["role"] = "user" - msg["content"] = []map[string]any{ - { - "type": "tool_result", - "tool_use_id": *item.ToolCallID, - "content": safeContent(item.Content), - }, - } - } else if len(item.ToolCalls) > 0 { - content := make([]map[string]any, 0) - if item.Content != nil && *item.Content != "" { - content = append(content, map[string]any{ - "type": "text", - "text": *item.Content, - }) - } - for _, tc := range item.ToolCalls { - content = append(content, map[string]any{ - "type": "tool_use", - "id": tc.ID, - "name": tc.Name, - "input": parseJSONOrString(tc.ArgumentsJSON), - }) - } - msg["content"] = content - } else if item.Content != nil { - msg["content"] = *item.Content - } - - messages = append(messages, msg) - } - - return messages, system -} - -func safeContent(s *string) string { - if s == nil { - return "" - } - return *s -} - -func parseJSONOrString(s string) any { - var v any - if err := json.Unmarshal([]byte(s), &v); err != nil { - return s - } - return v -} diff --git a/apps/gateway/internal/adapters/providers/anthropic/stream.go b/apps/gateway/internal/adapters/providers/anthropic/stream.go deleted file mode 100644 index 30e6044..0000000 --- a/apps/gateway/internal/adapters/providers/anthropic/stream.go +++ /dev/null @@ -1,218 +0,0 @@ -package anthropic - -import ( - "bufio" - "encoding/json" - "io" - "net/http" - "strings" - - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" -) - -// Stream reads SSE events from an Anthropic streaming response. -type Stream struct { - resp *http.Response - scanner *bufio.Scanner - done bool - usage *domain.Usage - toolIdx int -} - -func NewStream(resp *http.Response) *Stream { - scanner := bufio.NewScanner(resp.Body) - scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) - return &Stream{resp: resp, scanner: scanner} -} - -func (s *Stream) Recv() (domain.StreamEvent, error) { - if s.done { - return domain.StreamEvent{}, io.EOF - } - - var currentEvent string - - for s.scanner.Scan() { - line := s.scanner.Text() - - if line == "" { - continue - } - - if strings.HasPrefix(line, "event: ") { - currentEvent = strings.TrimPrefix(line, "event: ") - continue - } - - if !strings.HasPrefix(line, "data: ") { - continue - } - - data := strings.TrimPrefix(line, "data: ") - - event := s.processEvent(currentEvent, []byte(data)) - if event != nil { - return *event, nil - } - } - - if err := s.scanner.Err(); err != nil { - return domain.StreamEvent{}, err - } - - s.done = true - return domain.StreamEvent{}, io.EOF -} - -func (s *Stream) Close() error { - s.done = true - return s.resp.Body.Close() -} - -func (s *Stream) processEvent(eventType string, data []byte) *domain.StreamEvent { - switch eventType { - case "content_block_delta": - var delta struct { - Index int `json:"index"` - Delta struct { - Type string `json:"type"` - Text string `json:"text,omitempty"` - PartialJSON string `json:"partial_json,omitempty"` - } `json:"delta"` - } - if json.Unmarshal(data, &delta) != nil { - return nil - } - - if delta.Delta.Type == "text_delta" { - return &domain.StreamEvent{ - Type: domain.StreamEventOutputTextDelta, - ContentDelta: &delta.Delta.Text, - } - } - if delta.Delta.Type == "input_json_delta" { - return &domain.StreamEvent{ - Type: domain.StreamEventToolCallDelta, - ToolCallDelta: &domain.ToolCallDelta{ - Index: delta.Index, - ArgumentsDelta: &delta.Delta.PartialJSON, - }, - } - } - - case "content_block_start": - var block struct { - Index int `json:"index"` - ContentBlock struct { - Type string `json:"type"` - ID string `json:"id,omitempty"` - Name string `json:"name,omitempty"` - } `json:"content_block"` - } - if json.Unmarshal(data, &block) != nil { - return nil - } - if block.ContentBlock.Type == "tool_use" { - return &domain.StreamEvent{ - Type: domain.StreamEventToolCallDelta, - ToolCallDelta: &domain.ToolCallDelta{ - Index: block.Index, - ID: &block.ContentBlock.ID, - Name: &block.ContentBlock.Name, - }, - } - } - - case "message_start": - var msg struct { - Message struct { - Usage struct { - InputTokens int64 `json:"input_tokens"` - CacheCreationInputTokens int64 `json:"cache_creation_input_tokens"` - CacheReadInputTokens int64 `json:"cache_read_input_tokens"` - } `json:"usage"` - } `json:"message"` - } - if json.Unmarshal(data, &msg) == nil { - // Normalize PromptTokens to total input (uncached + cache_creation + cache_read) - totalInput := msg.Message.Usage.InputTokens + msg.Message.Usage.CacheCreationInputTokens + msg.Message.Usage.CacheReadInputTokens - s.usage = &domain.Usage{ - PromptTokens: totalInput, - CacheCreationTokens: msg.Message.Usage.CacheCreationInputTokens, - CacheReadTokens: msg.Message.Usage.CacheReadInputTokens, - } - } - role := "assistant" - return &domain.StreamEvent{ - Type: domain.StreamEventOutputMessageDelta, - Role: &role, - } - - case "message_delta": - var delta struct { - Delta struct { - StopReason string `json:"stop_reason"` - } `json:"delta"` - Usage struct { - OutputTokens int64 `json:"output_tokens"` - } `json:"usage"` - } - if json.Unmarshal(data, &delta) != nil { - return nil - } - - if s.usage != nil { - s.usage.CompletionTokens = delta.Usage.OutputTokens - s.usage.TotalTokens = s.usage.PromptTokens + s.usage.CompletionTokens - } - - finishReason := mapStopReason(delta.Delta.StopReason) - return &domain.StreamEvent{ - Type: domain.StreamEventCompleted, - FinishReason: &finishReason, - Usage: s.usage, - } - - case "message_stop": - s.done = true - return nil - - case "error": - var errData struct { - Error struct { - Type string `json:"type"` - Message string `json:"message"` - } `json:"error"` - } - msg := "unknown error" - if json.Unmarshal(data, &errData) == nil { - msg = errData.Error.Message - } - return &domain.StreamEvent{ - Type: domain.StreamEventError, - Error: &domain.GatewayError{ - HTTPStatus: 502, - Type: domain.ErrTypeProvider, - Code: domain.ErrCodeProviderUnavailable, - Message: msg, - }, - } - } - - return nil -} - -func mapStopReason(reason string) string { - switch reason { - case "end_turn": - return "stop" - case "max_tokens": - return "length" - case "tool_use": - return "tool_calls" - case "stop_sequence": - return "stop" - default: - return reason - } -} diff --git a/apps/gateway/internal/adapters/providers/openai/adapter.go b/apps/gateway/internal/adapters/providers/openai/adapter.go index 63d5ab9..03ac127 100644 --- a/apps/gateway/internal/adapters/providers/openai/adapter.go +++ b/apps/gateway/internal/adapters/providers/openai/adapter.go @@ -2,7 +2,6 @@ package openai import ( "context" - "encoding/json" "errors" "fmt" "net/http" @@ -10,10 +9,8 @@ import ( "strings" "time" - "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/google/uuid" ) // Adapter is the OpenAI-compatible provider adapter. @@ -39,122 +36,6 @@ func NewAdapter(timeout time.Duration, logger ...ports.Logger) *Adapter { return &Adapter{client: NewClient(timeout), logger: l} } -func (a *Adapter) Generate(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) { - apiKey := resolveAPIKey(model.ProviderConfig.APIKeySecretRef) - if apiKey == "" { - return domain.GenerateResult{}, missingProviderCredentialError(model.ProviderConfig.ProviderName) - } - - built := buildProviderRequest(req, model) - built.Body["stream"] = false - delete(built.Body, "stream_options") - - activeBuilt := built - var rawResp []byte - var err error - var retryReason string - invalidTraceSameBodyRetries := 0 - serverOverloadRetries := 0 - downgradeRetried := false - for attempt := 1; ; attempt++ { - a.logProviderRequest(ctx, model, activeBuilt, attempt, retryReason) - rawResp, err = a.client.Do(ctx, model.ProviderConfig.BaseURL, apiKey, activeBuilt.Body) - if err == nil { - break - } - if retry := maybeBuildNovitaSameBodyRetry(model, err, activeBuilt.Body); retry.CanRetry { - if canSpendSameBodyRetry(retry.RetryReason, &invalidTraceSameBodyRetries, &serverOverloadRetries) { - retryReason = retry.RetryReason - metrics.ProviderRetries.WithLabelValues(model.ProviderConfig.ProviderName, retryReason).Inc() - if !sleepBeforeProviderRetry(ctx, retryReason) { - return domain.GenerateResult{}, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), activeBuilt, attempt, retryReason) - } - continue - } - } - if !downgradeRetried { - if retry := maybeBuildNovitaDowngradeRetry(model, err, activeBuilt.Body); retry.CanRetry { - downgradeRetried = true - retryReason = retry.RetryReason - metrics.ProviderRetries.WithLabelValues(model.ProviderConfig.ProviderName, retryReason).Inc() - activeBuilt = builtProviderRequest{ - Body: retry.Body, - Policy: built.Policy, - Transforms: built.Transforms, - Omitted: retry.Omitted, - } - if !sleepBeforeProviderRetry(ctx, retryReason) { - return domain.GenerateResult{}, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), activeBuilt, attempt, retryReason) - } - continue - } - } - return domain.GenerateResult{}, withProviderPolicyMeta(mapProviderErrorWithCompatibilityContext(err, model, activeBuilt.Body), activeBuilt, attempt, retryReason) - } - - result, err := parseResponse(rawResp, req, model) - if err != nil { - return domain.GenerateResult{}, err - } - return result, nil -} - -func (a *Adapter) StreamGenerate(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (ports.GenerationStream, error) { - apiKey := resolveAPIKey(model.ProviderConfig.APIKeySecretRef) - if apiKey == "" { - return nil, missingProviderCredentialError(model.ProviderConfig.ProviderName) - } - - built := buildProviderRequest(req, model) - - activeBuilt := built - var streamResp *http.Response - var err error - var retryReason string - 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 { - break - } - if retry := maybeBuildNovitaSameBodyRetry(model, err, activeBuilt.Body); retry.CanRetry { - if canSpendSameBodyRetry(retry.RetryReason, &invalidTraceSameBodyRetries, &serverOverloadRetries) { - retryReason = retry.RetryReason - metrics.ProviderRetries.WithLabelValues(model.ProviderConfig.ProviderName, retryReason).Inc() - if !sleepBeforeProviderRetry(ctx, retryReason) { - return nil, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), activeBuilt, attempt, retryReason) - } - continue - } - } - if !downgradeRetried { - if retry := maybeBuildNovitaDowngradeRetry(model, err, activeBuilt.Body); retry.CanRetry { - downgradeRetried = true - retryReason = retry.RetryReason - metrics.ProviderRetries.WithLabelValues(model.ProviderConfig.ProviderName, retryReason).Inc() - activeBuilt = builtProviderRequest{ - Body: retry.Body, - Policy: built.Policy, - Transforms: built.Transforms, - Omitted: retry.Omitted, - } - if !sleepBeforeProviderRetry(ctx, retryReason) { - return nil, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), activeBuilt, attempt, retryReason) - } - continue - } - } - return nil, withProviderPolicyMeta(mapProviderErrorWithCompatibilityContext(err, model, activeBuilt.Body), activeBuilt, attempt, retryReason) - } - - return NewStream(streamResp, model.ProviderConfig.ProviderName).WithDiagnostics(ctx, a.logger, model, attempts), nil -} - func missingProviderCredentialError(providerName string) *domain.GatewayError { return domain.ErrInternal("an internal error occurred").WithMeta( "provider", providerName, @@ -166,101 +47,6 @@ func resolveAPIKey(secretRef string) string { return os.Getenv(secretRef) } -func parseResponse(data json.RawMessage, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) { - var resp struct { - ID string `json:"id"` - Created int64 `json:"created"` - Choices []struct { - Index int `json:"index"` - Message struct { - Role string `json:"role"` - Content *string `json:"content"` - ReasoningContent *string `json:"reasoning_content,omitempty"` - ToolCalls []struct { - ID string `json:"id"` - Type string `json:"type"` - Function struct { - Name string `json:"name"` - Arguments string `json:"arguments"` - } `json:"function"` - } `json:"tool_calls,omitempty"` - } `json:"message"` - FinishReason *string `json:"finish_reason"` - } `json:"choices"` - Usage *struct { - PromptTokens int64 `json:"prompt_tokens"` - CompletionTokens int64 `json:"completion_tokens"` - TotalTokens int64 `json:"total_tokens"` - PromptCacheHitTokens int64 `json:"prompt_cache_hit_tokens"` - PromptTokensDetails *struct { - CachedTokens int64 `json:"cached_tokens"` - } `json:"prompt_tokens_details,omitempty"` - } `json:"usage"` - BaseResp *providerBaseResponse `json:"base_resp,omitempty"` - } - - if err := json.Unmarshal(data, &resp); err != nil { - return domain.GenerateResult{}, fmt.Errorf("failed to parse provider response: %w", err) - } - if err := providerBaseResponseError(resp.BaseResp); err != nil { - return domain.GenerateResult{}, err - } - - result := domain.GenerateResult{ - ID: resp.ID, - CreatedUnix: resp.Created, - PublicModelID: req.PublicModelID, - ProviderName: model.ProviderConfig.ProviderName, - ProviderModelID: model.ProviderModelID, - } - - if result.ID == "" { - result.ID = uuid.New().String() - } - if result.CreatedUnix == 0 { - result.CreatedUnix = time.Now().Unix() - } - - for _, choice := range resp.Choices { - role := choice.Message.Role - out := domain.OutputItem{ - Type: domain.OutputItemTypeMessage, - Role: &role, - Content: choice.Message.Content, - } - if model.ProviderConfig.ProviderName == "deepseek" { - out.ReasoningContent = choice.Message.ReasoningContent - } - for _, tc := range choice.Message.ToolCalls { - out.ToolCalls = append(out.ToolCalls, domain.ToolCall{ - ID: tc.ID, - Name: tc.Function.Name, - ArgumentsJSON: tc.Function.Arguments, - }) - } - result.Output = append(result.Output, out) - result.FinishReason = choice.FinishReason - } - - if resp.Usage != nil { - result.Usage = &domain.Usage{ - PromptTokens: resp.Usage.PromptTokens, - CompletionTokens: resp.Usage.CompletionTokens, - TotalTokens: resp.Usage.TotalTokens, - CacheReadTokens: resp.Usage.PromptCacheHitTokens, - } - if resp.Usage.PromptTokensDetails != nil && resp.Usage.PromptTokensDetails.CachedTokens > 0 { - result.Usage.CacheReadTokens = resp.Usage.PromptTokensDetails.CachedTokens - } - } - - return result, nil -} - -func ParseResponse(data []byte, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) { - return parseResponse(data, req, model) -} - func mapProviderError(err error, providerName string) error { if err == nil { return nil diff --git a/apps/gateway/internal/adapters/providers/openai/adapter_error_test.go b/apps/gateway/internal/adapters/providers/openai/adapter_error_test.go index 8b52d1a..91d95f1 100644 --- a/apps/gateway/internal/adapters/providers/openai/adapter_error_test.go +++ b/apps/gateway/internal/adapters/providers/openai/adapter_error_test.go @@ -15,12 +15,12 @@ func TestAdapterMissingCredentialReturnsInternalError(t *testing.T) { name string call func() error }{ - {name: "generate", call: func() error { - _, err := adapter.Generate(context.Background(), domain.GenerateRequest{}, model) + {name: "complete", call: func() error { + _, _, err := adapter.Complete(context.Background(), []byte(`{"messages":[]}`), model) return err }}, {name: "stream", call: func() error { - _, err := adapter.StreamGenerate(context.Background(), domain.GenerateRequest{}, model) + _, err := adapter.Stream(context.Background(), []byte(`{"messages":[]}`), model) return err }}, } diff --git a/apps/gateway/internal/adapters/providers/openai/adapter_usage_test.go b/apps/gateway/internal/adapters/providers/openai/adapter_usage_test.go deleted file mode 100644 index a0fea53..0000000 --- a/apps/gateway/internal/adapters/providers/openai/adapter_usage_test.go +++ /dev/null @@ -1,39 +0,0 @@ -package openai - -import ( - "errors" - "net/http" - "testing" - - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" -) - -func TestParseResponse_MapsDeepSeekCacheHitTokens(t *testing.T) { - data := []byte(`{"usage":{"prompt_tokens":17,"completion_tokens":9,"total_tokens":26,"prompt_cache_hit_tokens":7}}`) - result, err := ParseResponse(data, domain.GenerateRequest{}, domain.PublicModel{}) - if err != nil { - t.Fatalf("ParseResponse returned error: %v", err) - } - if result.Usage == nil || result.Usage.CacheReadTokens != 7 { - t.Fatalf("cache-read tokens = %+v, want 7", result.Usage) - } -} - -func TestParseResponse_RejectsNonzeroBaseResponse(t *testing.T) { - _, err := ParseResponse([]byte(`{"base_resp":{"status_code":17,"status_msg":"provider error"}}`), domain.GenerateRequest{}, domain.PublicModel{}) - var gatewayErr *domain.GatewayError - if !errors.As(err, &gatewayErr) { - t.Fatalf("error = %v, want GatewayError", err) - } - if gatewayErr.HTTPStatus != http.StatusBadGateway || gatewayErr.Metadata["upstream_code"] != 17 { - t.Fatalf("gateway error = %#v", gatewayErr) - } -} - -func TestParseResponse_AllowsAbsentOrZeroBaseResponse(t *testing.T) { - for _, data := range []string{`{}`, `{"base_resp":{"status_code":0,"status_msg":""}}`} { - if _, err := ParseResponse([]byte(data), domain.GenerateRequest{}, domain.PublicModel{}); err != nil { - t.Fatalf("ParseResponse(%s): %v", data, err) - } - } -} diff --git a/apps/gateway/internal/adapters/providers/openai/mapper.go b/apps/gateway/internal/adapters/providers/openai/mapper.go index ab2c41c..957ef27 100644 --- a/apps/gateway/internal/adapters/providers/openai/mapper.go +++ b/apps/gateway/internal/adapters/providers/openai/mapper.go @@ -8,7 +8,6 @@ type providerPolicy struct { name string useMaxTokens bool explicitNullAssistantToolContent bool - forwardAssistantReasoningContent bool requireToolReasoningContent bool forwardDeveloperRole bool } @@ -27,147 +26,6 @@ type retryBuildResult struct { CanRetry bool } -func buildRequestBody(req domain.GenerateRequest, model domain.PublicModel) map[string]any { - return buildProviderRequest(req, model).Body -} - -func BuildRequestBody(req domain.GenerateRequest, model domain.PublicModel) map[string]any { - return buildRequestBody(req, model) -} - -func buildProviderRequest(req domain.GenerateRequest, model domain.PublicModel) builtProviderRequest { - policy := policyForProvider(model.ProviderConfig.ProviderName) - body := map[string]any{ - "model": model.UpstreamModelName, - } - transforms := make([]string, 0, 2) - - messages, messageTransforms := buildMessages(req, policy) - body["messages"] = messages - transforms = append(transforms, messageTransforms...) - - if req.MaxOutputTokens != nil { - v := *req.MaxOutputTokens - if model.MaxOutputTokens > 0 && v > model.MaxOutputTokens { - v = model.MaxOutputTokens - } - if policy.useMaxTokens { - body["max_tokens"] = v - transforms = append(transforms, "token_limit_field=max_tokens") - } else { - body["max_completion_tokens"] = v - } - } - if req.Temperature != nil { - body["temperature"] = *req.Temperature - } - if req.ReasoningEffort != nil && *req.ReasoningEffort != "" { - body["reasoning_effort"] = *req.ReasoningEffort - } - if req.TopP != nil { - body["top_p"] = *req.TopP - } - if len(req.Stop) > 0 { - if len(req.Stop) == 1 { - body["stop"] = req.Stop[0] - } else { - body["stop"] = req.Stop - } - } - if req.Stream { - body["stream"] = true - body["stream_options"] = map[string]any{"include_usage": true} - } - if req.User != nil { - body["user"] = *req.User - } - - if len(req.Tools) > 0 { - tools := make([]map[string]any, 0, len(req.Tools)) - for _, t := range req.Tools { - fn := map[string]any{ - "name": t.Name, - "description": t.Description, - "parameters": t.Parameters, - } - // Only include "strict" when explicitly true; many providers - // reject this OpenAI-specific field. - if t.Strict { - fn["strict"] = true - } - tool := map[string]any{ - "type": "function", - "function": fn, - } - tools = append(tools, tool) - } - body["tools"] = tools - } - - if req.ToolChoice != nil { - switch req.ToolChoice.Mode { - case domain.ToolChoiceNone, domain.ToolChoiceAuto, domain.ToolChoiceRequired: - body["tool_choice"] = req.ToolChoice.Mode - case domain.ToolChoiceFunction: - body["tool_choice"] = map[string]any{ - "type": "function", - "function": map[string]any{ - "name": *req.ToolChoice.FunctionName, - }, - } - } - } - - if req.ParallelToolCalls != nil { - body["parallel_tool_calls"] = *req.ParallelToolCalls - } - - // Pass-through parameters - if req.PresencePenalty != nil { - body["presence_penalty"] = *req.PresencePenalty - } - if req.FrequencyPenalty != nil { - body["frequency_penalty"] = *req.FrequencyPenalty - } - if len(req.LogitBias) > 0 { - body["logit_bias"] = req.LogitBias - } - if req.Seed != nil { - body["seed"] = *req.Seed - } - if req.Logprobs != nil { - body["logprobs"] = *req.Logprobs - } - if req.TopLogprobs != nil { - body["top_logprobs"] = *req.TopLogprobs - } - if req.Store != nil { - body["store"] = *req.Store - } - if req.ServiceTier != nil { - body["service_tier"] = *req.ServiceTier - } - - if req.TextConfig != nil && req.TextConfig.FormatType != nil { - switch *req.TextConfig.FormatType { - case "json_object": - body["response_format"] = map[string]any{"type": "json_object"} - case "json_schema": - rf := map[string]any{"type": "json_schema"} - if req.TextConfig.JSONSchema != nil { - rf["json_schema"] = req.TextConfig.JSONSchema - } - body["response_format"] = rf - } - } - - return builtProviderRequest{ - Body: body, - Policy: policy.name, - Transforms: transforms, - } -} - func policyForProvider(providerName string) providerPolicy { switch providerName { case "deepseek": @@ -178,7 +36,6 @@ func policyForProvider(providerName string) providerPolicy { name: "deepseek", useMaxTokens: true, explicitNullAssistantToolContent: true, - forwardAssistantReasoningContent: true, requireToolReasoningContent: true, } case "novita": @@ -201,71 +58,6 @@ func policyForProvider(providerName string) providerPolicy { } } -func buildMessages(req domain.GenerateRequest, policy providerPolicy) ([]map[string]any, []string) { - var messages []map[string]any - var transforms []string - - if req.Instructions != nil && *req.Instructions != "" { - messages = append(messages, map[string]any{ - "role": "system", - "content": *req.Instructions, - }) - } - - for _, item := range req.Input { - msg := map[string]any{} - - role := "user" - if item.Role != nil { - role = *item.Role - } - if role == "developer" && !policy.forwardDeveloperRole { - role = "system" - transforms = append(transforms, "developer_role=system") - } - msg["role"] = role - - if item.Content != nil { - msg["content"] = *item.Content - } else if policy.explicitNullAssistantToolContent && role == "assistant" && len(item.ToolCalls) > 0 { - msg["content"] = nil - transforms = append(transforms, "assistant_tool_content=null") - } - - if item.ToolCallID != nil { - msg["tool_call_id"] = *item.ToolCallID - } - - if role == "assistant" && policy.forwardAssistantReasoningContent { - if item.ReasoningContent != nil { - msg["reasoning_content"] = *item.ReasoningContent - } else if policy.requireToolReasoningContent && len(item.ToolCalls) > 0 { - msg["reasoning_content"] = "" - transforms = append(transforms, "assistant_tool_reasoning_content=empty") - } - } - - if len(item.ToolCalls) > 0 { - tcs := make([]map[string]any, 0, len(item.ToolCalls)) - for _, tc := range item.ToolCalls { - tcs = append(tcs, map[string]any{ - "id": tc.ID, - "type": "function", - "function": map[string]any{ - "name": tc.Name, - "arguments": tc.ArgumentsJSON, - }, - }) - } - msg["tool_calls"] = tcs - } - - messages = append(messages, msg) - } - - return messages, transforms -} - func buildNovitaRetryRequest(body map[string]any) retryBuildResult { retryBody := cloneBody(body) omitted := make([]string, 0, 6) diff --git a/apps/gateway/internal/adapters/providers/openai/policy_test.go b/apps/gateway/internal/adapters/providers/openai/policy_test.go index 3d36278..ec2f254 100644 --- a/apps/gateway/internal/adapters/providers/openai/policy_test.go +++ b/apps/gateway/internal/adapters/providers/openai/policy_test.go @@ -13,231 +13,6 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) -func TestBuildProviderRequest_DefaultPolicyTokenLimit(t *testing.T) { - maxTokens := 32 - req := domain.GenerateRequest{ - PublicModelID: "openai/test", - Input: []domain.InputItem{message("user", "hi")}, - MaxOutputTokens: &maxTokens, - } - model := domain.PublicModel{ - UpstreamModelName: "test-model", - ProviderConfig: domain.ProviderConfig{ProviderName: "openai"}, - MaxOutputTokens: 128, - } - - built := buildProviderRequest(req, model) - if built.Policy != "openai-compatible" { - t.Fatalf("policy = %q, want openai-compatible", built.Policy) - } - if got := built.Body["max_completion_tokens"]; got != 32 { - t.Fatalf("max_completion_tokens = %v, want 32", got) - } - if _, ok := built.Body["max_tokens"]; ok { - t.Fatal("default policy must not send max_tokens") - } -} - -func TestBuildProviderRequest_ForwardsReasoningEffort(t *testing.T) { - effort := "low" - maxTokens := 1024 - req := domain.GenerateRequest{ - PublicModelID: "phala/gpt-oss-20b", - Input: []domain.InputItem{message("user", "hi")}, - ReasoningEffort: &effort, - MaxOutputTokens: &maxTokens, - } - model := domain.PublicModel{ - UpstreamModelName: "phala/gpt-oss-20b", - ProviderConfig: domain.ProviderConfig{ProviderName: "phala"}, - MaxOutputTokens: 1024, - } - - built := buildProviderRequest(req, model) - if got := built.Body["reasoning_effort"]; got != "low" { - t.Fatalf("reasoning_effort = %v, want low", got) - } -} - -func TestBuildProviderRequest_DeveloperRolePolicy(t *testing.T) { - req := domain.GenerateRequest{ - PublicModelID: "test/developer-role", - Input: []domain.InputItem{message("developer", "follow instructions")}, - } - - compatible := buildProviderRequest(req, domain.PublicModel{ - UpstreamModelName: "compatible-model", - ProviderConfig: domain.ProviderConfig{ProviderName: "phala"}, - }) - messages := compatible.Body["messages"].([]map[string]any) - if messages[0]["role"] != "system" { - t.Fatalf("compatible role = %v, want system", messages[0]["role"]) - } - if !containsString(compatible.Transforms, "developer_role=system") { - t.Fatalf("transforms = %v, want developer_role=system", compatible.Transforms) - } - - openaiBuilt := buildProviderRequest(req, domain.PublicModel{ - UpstreamModelName: "gpt-5", - ProviderConfig: domain.ProviderConfig{ProviderName: "openai"}, - }) - messages = openaiBuilt.Body["messages"].([]map[string]any) - if messages[0]["role"] != "developer" { - t.Fatalf("openai role = %v, want developer", messages[0]["role"]) - } -} - -func TestBuildProviderRequest_OmittedTokenLimitStaysOmitted(t *testing.T) { - req := domain.GenerateRequest{ - PublicModelID: "openai/test", - Input: []domain.InputItem{message("user", "hi")}, - } - model := domain.PublicModel{ - UpstreamModelName: "test-model", - ProviderConfig: domain.ProviderConfig{ProviderName: "openai"}, - MaxOutputTokens: 128, - } - - built := buildProviderRequest(req, model) - if _, ok := built.Body["max_completion_tokens"]; ok { - t.Fatal("max_completion_tokens must stay omitted when the client omits a token limit") - } - if _, ok := built.Body["max_tokens"]; ok { - t.Fatal("max_tokens must stay omitted when the client omits a token limit") - } -} - -func TestBuildProviderRequest_NovitaPolicy(t *testing.T) { - maxTokens := 32 - req := domain.GenerateRequest{ - PublicModelID: "moonshotai/kimi-k2.6", - Input: []domain.InputItem{message("user", "hi")}, - MaxOutputTokens: &maxTokens, - } - model := novitaModel("http://example.test") - - built := buildProviderRequest(req, model) - if built.Policy != "novita" { - t.Fatalf("policy = %q, want novita", built.Policy) - } - if got := built.Body["max_tokens"]; got != 32 { - t.Fatalf("max_tokens = %v, want 32", got) - } - if _, ok := built.Body["max_completion_tokens"]; ok { - t.Fatal("Novita policy must not send max_completion_tokens") - } - if !containsString(built.Transforms, "token_limit_field=max_tokens") { - t.Fatalf("transforms = %v, want token_limit_field=max_tokens", built.Transforms) - } -} - -func TestBuildProviderRequest_DeepSeekPolicy(t *testing.T) { - maxTokens := 32 - req := domain.GenerateRequest{ - PublicModelID: "deepseek/deepseek-v4-pro", - Input: []domain.InputItem{message("user", "hi")}, - MaxOutputTokens: &maxTokens, - } - model := deepseekModel("http://example.test") - - built := buildProviderRequest(req, model) - if built.Policy != "deepseek" { - t.Fatalf("policy = %q, want deepseek", built.Policy) - } - if got := built.Body["max_tokens"]; got != 32 { - t.Fatalf("max_tokens = %v, want 32", got) - } - if _, ok := built.Body["max_completion_tokens"]; ok { - t.Fatal("DeepSeek policy must not send max_completion_tokens") - } - if !containsString(built.Transforms, "token_limit_field=max_tokens") { - t.Fatalf("transforms = %v, want token_limit_field=max_tokens", built.Transforms) - } -} - -func TestBuildProviderRequest_DeepSeekAssistantToolContentNull(t *testing.T) { - req := domain.GenerateRequest{ - PublicModelID: "deepseek/deepseek-v4-pro", - Input: []domain.InputItem{ - { - Type: domain.InputItemTypeMessage, - Role: stringPtr("assistant"), - ToolCalls: []domain.ToolCall{{ID: "call_1", Name: "noop", ArgumentsJSON: "{}"}}, - }, - }, - } - - built := buildProviderRequest(req, deepseekModel("http://example.test")) - messages := built.Body["messages"].([]map[string]any) - if _, ok := messages[0]["content"]; !ok { - t.Fatal("DeepSeek assistant tool-call message must include content key") - } - if messages[0]["content"] != nil { - t.Fatalf("content = %v, want nil", messages[0]["content"]) - } - if !containsString(built.Transforms, "assistant_tool_content=null") { - t.Fatalf("transforms = %v, want assistant_tool_content=null", built.Transforms) - } - if got, ok := messages[0]["reasoning_content"].(string); !ok || got != "" { - t.Fatalf("reasoning_content = %#v, want empty string", messages[0]["reasoning_content"]) - } - if !containsString(built.Transforms, "assistant_tool_reasoning_content=empty") { - t.Fatalf("transforms = %v, want assistant_tool_reasoning_content=empty", built.Transforms) - } -} - -func TestBuildProviderRequest_DeepSeekPreservesAssistantReasoningContent(t *testing.T) { - reasoning := "used the search result" - req := domain.GenerateRequest{ - PublicModelID: "deepseek/deepseek-v4-pro", - Input: []domain.InputItem{ - { - Type: domain.InputItemTypeMessage, - Role: stringPtr("assistant"), - ReasoningContent: &reasoning, - ToolCalls: []domain.ToolCall{{ID: "call_1", Name: "noop", ArgumentsJSON: "{}"}}, - }, - }, - } - - built := buildProviderRequest(req, deepseekModel("http://example.test")) - messages := built.Body["messages"].([]map[string]any) - if got := messages[0]["reasoning_content"]; got != reasoning { - t.Fatalf("reasoning_content = %#v, want %q", got, reasoning) - } - if containsString(built.Transforms, "assistant_tool_reasoning_content=empty") { - t.Fatalf("transforms = %v, must not synthesize empty reasoning when client supplied it", built.Transforms) - } -} - -func TestBuildProviderRequest_NovitaAssistantToolContentNull(t *testing.T) { - req := domain.GenerateRequest{ - PublicModelID: "moonshotai/kimi-k2.6", - Input: []domain.InputItem{ - { - Type: domain.InputItemTypeMessage, - Role: stringPtr("assistant"), - ToolCalls: []domain.ToolCall{{ID: "call_1", Name: "noop", ArgumentsJSON: "{}"}}, - }, - }, - } - - built := buildProviderRequest(req, novitaModel("http://example.test")) - messages := built.Body["messages"].([]map[string]any) - if _, ok := messages[0]["content"]; !ok { - t.Fatal("Novita assistant tool-call message must include content key") - } - if messages[0]["content"] != nil { - t.Fatalf("content = %v, want nil", messages[0]["content"]) - } - if !containsString(built.Transforms, "assistant_tool_content=null") { - t.Fatalf("transforms = %v, want assistant_tool_content=null", built.Transforms) - } - if _, ok := messages[0]["reasoning_content"]; ok { - t.Fatal("Novita assistant tool-call message must not include DeepSeek reasoning_content compatibility field") - } -} - func TestBuildNovitaRetryRequest_GuardedDowngrade(t *testing.T) { body := map[string]any{ "model": "moonshotai/kimi-k2.6", @@ -298,7 +73,7 @@ func TestBuildNovitaRetryRequest_NamedToolChoiceIsNotDowngraded(t *testing.T) { } } -func TestAdapterGenerate_NovitaRetriesSafeDowngrade(t *testing.T) { +func TestAdapterComplete_NovitaRetriesSafeDowngrade(t *testing.T) { t.Setenv("NOVITA_TEST_KEY", "test-key") var attempts int var bodies []map[string]any @@ -320,26 +95,14 @@ func TestAdapterGenerate_NovitaRetriesSafeDowngrade(t *testing.T) { })) defer server.Close() - maxTokens := 8 - parallel := false - req := domain.GenerateRequest{ - PublicModelID: "moonshotai/kimi-k2.6", - Input: []domain.InputItem{message("user", "hi")}, - MaxOutputTokens: &maxTokens, - ParallelToolCalls: ¶llel, - Store: boolPtr(true), - ServiceTier: stringPtr("auto"), - User: stringPtr("user-1"), - ToolChoice: &domain.ToolChoice{Mode: domain.ToolChoiceAuto}, - Tools: []domain.ToolDefinition{{Name: "noop", Parameters: map[string]any{"type": "object"}}}, - } + raw := []byte(`{"model":"moonshotai/kimi-k2.6","messages":[{"role":"user","content":"hi"}],"max_tokens":8,"parallel_tool_calls":false,"store":true,"service_tier":"auto","user":"user-1","tool_choice":"auto","tools":[{"type":"function","function":{"name":"noop","parameters":{"type":"object"}}}]}`) adapter := &Adapter{ client: &Client{httpClient: server.Client(), responseTimeout: time.Second}, } - _, err := adapter.Generate(context.Background(), req, novitaModel(server.URL)) + _, _, err := adapter.Complete(context.Background(), raw, novitaModel(server.URL)) if err != nil { - t.Fatalf("Generate returned error: %v", err) + t.Fatalf("Complete returned error: %v", err) } if attempts != 4 { t.Fatalf("attempts = %d, want 4", attempts) @@ -364,7 +127,7 @@ func TestAdapterGenerate_NovitaRetriesSafeDowngrade(t *testing.T) { } } -func TestAdapterGenerate_NovitaNamedToolChoiceFailsClearlyAfterSameBodyRetry(t *testing.T) { +func TestAdapterComplete_NovitaNamedToolChoiceFailsClearlyAfterSameBodyRetry(t *testing.T) { t.Setenv("NOVITA_TEST_KEY", "test-key") var attempts int server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -374,22 +137,12 @@ func TestAdapterGenerate_NovitaNamedToolChoiceFailsClearlyAfterSameBodyRetry(t * })) defer server.Close() - maxTokens := 8 - req := domain.GenerateRequest{ - PublicModelID: "moonshotai/kimi-k2.6", - Input: []domain.InputItem{message("user", "hi")}, - MaxOutputTokens: &maxTokens, - ToolChoice: &domain.ToolChoice{ - Mode: domain.ToolChoiceFunction, - FunctionName: stringPtr("noop"), - }, - Tools: []domain.ToolDefinition{{Name: "noop", Parameters: map[string]any{"type": "object"}}}, - } + raw := []byte(`{"model":"moonshotai/kimi-k2.6","messages":[{"role":"user","content":"hi"}],"max_tokens":8,"tool_choice":{"type":"function","function":{"name":"noop"}},"tools":[{"type":"function","function":{"name":"noop","parameters":{"type":"object"}}}]}`) adapter := &Adapter{ client: &Client{httpClient: server.Client(), responseTimeout: time.Second}, } - _, err := adapter.Generate(context.Background(), req, novitaModel(server.URL)) + _, _, err := adapter.Complete(context.Background(), raw, novitaModel(server.URL)) if err == nil { t.Fatal("expected error") } @@ -401,7 +154,7 @@ func TestAdapterGenerate_NovitaNamedToolChoiceFailsClearlyAfterSameBodyRetry(t * } } -func TestAdapterGenerate_NovitaSafeRetryStillReportsNamedToolChoiceIncompatibility(t *testing.T) { +func TestAdapterComplete_NovitaSafeRetryStillReportsNamedToolChoiceIncompatibility(t *testing.T) { t.Setenv("NOVITA_TEST_KEY", "test-key") var attempts int server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -411,24 +164,12 @@ func TestAdapterGenerate_NovitaSafeRetryStillReportsNamedToolChoiceIncompatibili })) defer server.Close() - maxTokens := 8 - parallel := false - req := domain.GenerateRequest{ - PublicModelID: "moonshotai/kimi-k2.6", - Input: []domain.InputItem{message("user", "hi")}, - MaxOutputTokens: &maxTokens, - ParallelToolCalls: ¶llel, - ToolChoice: &domain.ToolChoice{ - Mode: domain.ToolChoiceFunction, - FunctionName: stringPtr("noop"), - }, - Tools: []domain.ToolDefinition{{Name: "noop", Parameters: map[string]any{"type": "object"}}}, - } + raw := []byte(`{"model":"moonshotai/kimi-k2.6","messages":[{"role":"user","content":"hi"}],"max_tokens":8,"parallel_tool_calls":false,"tool_choice":{"type":"function","function":{"name":"noop"}},"tools":[{"type":"function","function":{"name":"noop","parameters":{"type":"object"}}}]}`) adapter := &Adapter{ client: &Client{httpClient: server.Client(), responseTimeout: time.Second}, } - _, err := adapter.Generate(context.Background(), req, novitaModel(server.URL)) + _, _, err := adapter.Complete(context.Background(), raw, novitaModel(server.URL)) if err == nil { t.Fatal("expected error") } @@ -440,7 +181,7 @@ func TestAdapterGenerate_NovitaSafeRetryStillReportsNamedToolChoiceIncompatibili } } -func TestAdapterGenerate_NovitaFailedSameBodyRetryIncludesProviderParams(t *testing.T) { +func TestAdapterComplete_NovitaFailedSameBodyRetryIncludesProviderParams(t *testing.T) { t.Setenv("NOVITA_TEST_KEY", "test-key") var attempts int server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -450,18 +191,12 @@ func TestAdapterGenerate_NovitaFailedSameBodyRetryIncludesProviderParams(t *test })) defer server.Close() - maxTokens := 8 - req := domain.GenerateRequest{ - PublicModelID: "moonshotai/kimi-k2.6", - Input: []domain.InputItem{message("user", "hi")}, - MaxOutputTokens: &maxTokens, - Tools: []domain.ToolDefinition{{Name: "noop", Parameters: map[string]any{"type": "object"}}}, - } + raw := []byte(`{"model":"moonshotai/kimi-k2.6","messages":[{"role":"user","content":"hi"}],"max_tokens":8,"tools":[{"type":"function","function":{"name":"noop","parameters":{"type":"object"}}}]}`) adapter := &Adapter{ client: &Client{httpClient: server.Client(), responseTimeout: time.Second}, } - _, err := adapter.Generate(context.Background(), req, novitaModel(server.URL)) + _, _, err := adapter.Complete(context.Background(), raw, novitaModel(server.URL)) if err == nil { t.Fatal("expected error") } @@ -488,7 +223,7 @@ func TestAdapterGenerate_NovitaFailedSameBodyRetryIncludesProviderParams(t *test } } -func TestAdapterGenerate_NovitaServerOverloadRetriesSameBody(t *testing.T) { +func TestAdapterComplete_NovitaServerOverloadRetriesSameBody(t *testing.T) { t.Setenv("NOVITA_TEST_KEY", "test-key") var attempts int var firstBody map[string]any @@ -512,22 +247,14 @@ func TestAdapterGenerate_NovitaServerOverloadRetriesSameBody(t *testing.T) { })) defer server.Close() - maxTokens := 8 - parallel := false - req := domain.GenerateRequest{ - PublicModelID: "moonshotai/kimi-k2.6", - Input: []domain.InputItem{message("user", "hi")}, - MaxOutputTokens: &maxTokens, - ParallelToolCalls: ¶llel, - Tools: []domain.ToolDefinition{{Name: "noop", Parameters: map[string]any{"type": "object"}}}, - } + raw := []byte(`{"model":"moonshotai/kimi-k2.6","messages":[{"role":"user","content":"hi"}],"max_tokens":8,"parallel_tool_calls":false,"tools":[{"type":"function","function":{"name":"noop","parameters":{"type":"object"}}}]}`) adapter := &Adapter{ client: &Client{httpClient: server.Client(), responseTimeout: time.Second}, } - _, err := adapter.Generate(context.Background(), req, novitaModel(server.URL)) + _, _, err := adapter.Complete(context.Background(), raw, novitaModel(server.URL)) if err != nil { - t.Fatalf("Generate returned error: %v", err) + t.Fatalf("Complete returned error: %v", err) } if attempts != 2 { t.Fatalf("attempts = %d, want 2", attempts) @@ -594,35 +321,6 @@ func novitaModel(baseURL string) domain.PublicModel { } } -func deepseekModel(baseURL string) domain.PublicModel { - return domain.PublicModel{ - PublicModelID: "deepseek/deepseek-v4-pro", - UpstreamModelName: "deepseek-v4-pro", - MaxOutputTokens: 8192, - ProviderConfig: domain.ProviderConfig{ - ProviderName: "deepseek", - BaseURL: baseURL, - APIKeySecretRef: "DEEPSEEK_TEST_KEY", - }, - } -} - -func message(role, content string) domain.InputItem { - return domain.InputItem{ - Type: domain.InputItemTypeMessage, - Role: stringPtr(role), - Content: stringPtr(content), - } -} - -func stringPtr(v string) *string { - return &v -} - -func boolPtr(v bool) *bool { - return &v -} - func containsString(values []string, want string) bool { for _, v := range values { if v == want { diff --git a/apps/gateway/internal/adapters/providers/openai/proxy.go b/apps/gateway/internal/adapters/providers/openai/proxy.go index b8f836d..fcfe844 100644 --- a/apps/gateway/internal/adapters/providers/openai/proxy.go +++ b/apps/gateway/internal/adapters/providers/openai/proxy.go @@ -149,30 +149,30 @@ func dedupe(list []string) []string { return out } -// ProxyStream forwards a streaming request, with the same retries as the -// translating path, and returns the provider's SSE body untouched. -func (a *Adapter) ProxyStream(ctx context.Context, raw []byte, model domain.PublicModel) (ports.ProxyResponse, error) { - body, _, err := PrepareProxyBody(raw, model, true) +// Stream forwards a streaming request, with the provider's retry policy, +// and returns the provider's SSE body untouched. +func (a *Adapter) Stream(ctx context.Context, raw []byte, model domain.PublicModel) (ports.ProviderStream, error) { + body, transforms, err := PrepareProxyBody(raw, model, true) if err != nil { - return ports.ProxyResponse{}, err + return ports.ProviderStream{}, err } - resp, err := a.proxy(ctx, body, model, func(ctx context.Context, apiKey string, body map[string]any) (*http.Response, error) { + resp, err := a.proxy(ctx, body, transforms, model, func(ctx context.Context, apiKey string, body map[string]any) (*http.Response, error) { return a.client.DoStream(ctx, model.ProviderConfig.BaseURL, apiKey, body) }) if err != nil { - return ports.ProxyResponse{}, err + return ports.ProviderStream{}, err } - return ports.ProxyResponse{Body: resp.Body}, nil + return ports.ProviderStream{Body: resp.Body}, nil } -// ProxyJSON forwards a non-streaming request and returns the provider's JSON. -func (a *Adapter) ProxyJSON(ctx context.Context, raw []byte, model domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) { - body, _, err := PrepareProxyBody(raw, model, false) +// Complete forwards a non-streaming request and returns the provider's JSON. +func (a *Adapter) Complete(ctx context.Context, raw []byte, model domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) { + body, transforms, err := PrepareProxyBody(raw, model, false) if err != nil { return nil, nil, err } var out []byte - _, err = a.proxy(ctx, body, model, func(ctx context.Context, apiKey string, body map[string]any) (*http.Response, error) { + _, err = a.proxy(ctx, body, transforms, model, func(ctx context.Context, apiKey string, body map[string]any) (*http.Response, error) { var err error out, err = a.client.Do(ctx, model.ProviderConfig.BaseURL, apiKey, body) if err != nil { @@ -184,13 +184,13 @@ func (a *Adapter) ProxyJSON(ctx context.Context, raw []byte, model domain.Public } // proxy sends body with the provider's retry policy (Novita's transient -// rejections), and maps errors like the translating path. -func (a *Adapter) proxy(ctx context.Context, body map[string]any, model domain.PublicModel, send func(context.Context, string, map[string]any) (*http.Response, error)) (*http.Response, error) { +// rejections), and maps provider errors to gateway errors. +func (a *Adapter) proxy(ctx context.Context, body map[string]any, transforms []string, model domain.PublicModel, send func(context.Context, string, map[string]any) (*http.Response, error)) (*http.Response, error) { apiKey := resolveAPIKey(model.ProviderConfig.APIKeySecretRef) if apiKey == "" { return nil, missingProviderCredentialError(model.ProviderConfig.ProviderName) } - built := builtProviderRequest{Body: body, Policy: policyForProvider(model.ProviderConfig.ProviderName).name} + built := builtProviderRequest{Body: body, Policy: policyForProvider(model.ProviderConfig.ProviderName).name, Transforms: transforms} active := built var retryReason string invalidTraceRetries, overloadRetries := 0, 0 @@ -215,7 +215,7 @@ func (a *Adapter) proxy(ctx context.Context, body map[string]any, model domain.P downgradeRetried = true retryReason = retry.RetryReason metrics.ProviderRetries.WithLabelValues(model.ProviderConfig.ProviderName, retryReason).Inc() - active = builtProviderRequest{Body: retry.Body, Policy: built.Policy, Omitted: retry.Omitted} + active = builtProviderRequest{Body: retry.Body, Policy: built.Policy, Transforms: transforms, Omitted: retry.Omitted} if !sleepBeforeProviderRetry(ctx, retryReason) { return nil, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), active, attempt, retryReason) } diff --git a/apps/gateway/internal/adapters/providers/openai/proxy_test.go b/apps/gateway/internal/adapters/providers/openai/proxy_test.go index 148f0aa..9f34562 100644 --- a/apps/gateway/internal/adapters/providers/openai/proxy_test.go +++ b/apps/gateway/internal/adapters/providers/openai/proxy_test.go @@ -92,3 +92,21 @@ func TestPrepareProxyBody_Edits(t *testing.T) { } }) } + +// The provider policies, for requests the old translating path also covered. +func TestPrepareProxyBody_ProviderPolicies(t *testing.T) { + // No limit from the client means no limit sent. + if got := prepared(t, `{"messages":[]}`, "novita", false); len(got) != 2 { + t.Fatalf("limit invented: %v", got) + } + // DeepSeek keeps the reasoning the client sent back. + raw := `{"messages":[{"role":"assistant","content":"x","reasoning_content":"because","tool_calls":[{"id":"a","type":"function","function":{"name":"f","arguments":"{}"}}]}]}` + msg := prepared(t, raw, "deepseek", false)["messages"].([]any)[0].(map[string]any) + if msg["reasoning_content"] != "because" || msg["content"] != "x" { + t.Fatalf("assistant message changed: %v", msg) + } + // OpenAI itself gets the developer role. + if role := prepared(t, `{"messages":[{"role":"developer","content":"x"}]}`, "openai", false)["messages"].([]any)[0].(map[string]any)["role"]; role != "developer" { + t.Fatalf("role = %v", role) + } +} diff --git a/apps/gateway/internal/adapters/providers/openai/stream.go b/apps/gateway/internal/adapters/providers/openai/stream.go deleted file mode 100644 index 8316008..0000000 --- a/apps/gateway/internal/adapters/providers/openai/stream.go +++ /dev/null @@ -1,438 +0,0 @@ -package openai - -import ( - "bufio" - "encoding/json" - "io" - "net/http" - "strings" - - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" -) - -// maxSSELine bounds one SSE line. Providers that do not stream tool -// arguments send a whole call (for example a large file) in one line. -const maxSSELine = 16 << 20 - -// Stream reads SSE events from an OpenAI-compatible streaming response for -// the translating path (PII masking, non-OpenAI clients of the domain model). -// Proxied requests don't use it. It keeps everything the provider sent: -// -// - every tool call in a chunk, not only the first; -// - text in a chunk that also carries tool calls; -// - output after an early usage-only chunk (which doesn't end the stream); -// - a finish reason at the end, and a provider error for a stream with no -// chunk at all or an error payload. -type Stream struct { - diagnostics *streamDiagnostics - resp *http.Response - scanner *bufio.Scanner - done bool // upstream fully read - includeReasoningContent bool - pending []domain.StreamEvent - completed bool // a finish reason was forwarded - sawChunk bool - lastUsage *domain.Usage -} - -func NewStream(resp *http.Response, providerName ...string) *Stream { - scanner := bufio.NewScanner(resp.Body) - scanner.Buffer(make([]byte, 0, 64*1024), maxSSELine) - includeReasoningContent := len(providerName) > 0 && providerName[0] == "deepseek" - return &Stream{ - resp: resp, - scanner: scanner, - includeReasoningContent: includeReasoningContent, - } -} - -func (s *Stream) Recv() (domain.StreamEvent, error) { - for len(s.pending) == 0 { - if s.done { - return domain.StreamEvent{}, io.EOF - } - events, err := s.read() - if err != nil { - return domain.StreamEvent{}, err - } - s.pending = events - } - event := s.pending[0] - s.pending = s.pending[1:] - return event, nil -} - -// read consumes upstream lines until they produce events or the stream ends. -func (s *Stream) read() ([]domain.StreamEvent, error) { - for s.scanner.Scan() { - line := s.scanner.Text() - if line == "" { - continue - } - // SSE allows "data:" with or without a following space. - data, ok := strings.CutPrefix(line, "data:") - if !ok { - if s.diagnostics != nil && !isSSEField(line) { - s.diagnostics.unsupported++ - } - continue - } - data = strings.TrimPrefix(data, " ") - if data == "" { - continue - } - if data == "[DONE]" { - s.diagnostics.end("done_marker", nil) - return s.end() - } - - 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 - } - s.sawChunk = true - if err := providerBaseResponseError(chunk.BaseResp); err != nil { - return nil, s.fail(chunk, err) - } - if err := providerStreamError(chunk.Error); err != nil { - return nil, s.fail(chunk, err) - } - events := s.normalize(mapChunkToStreamEvents(chunk, s.includeReasoningContent)) - if s.diagnostics != nil { - s.diagnostics.observe(chunk, events) - } - for i := range events { - events[i].ProviderResponseID = chunk.ID - } - if len(events) > 0 { - return events, nil - } - } - - if err := s.scanner.Err(); err != nil { - s.diagnostics.end("read_error", err) - return nil, err - } - s.diagnostics.end("eof", nil) - return s.end() -} - -// isSSEField reports SSE lines other than data: comments, event, id, retry. -func isSSEField(line string) bool { - return strings.HasPrefix(line, ":") || strings.HasPrefix(line, "event:") || - strings.HasPrefix(line, "id:") || strings.HasPrefix(line, "retry:") -} - -func (s *Stream) fail(chunk chatCompletionChunk, err error) error { - if s.diagnostics != nil { - s.diagnostics.observe(chunk, nil) - } - s.diagnostics.end("provider_error", err) - s.done = true - return err -} - -// end finishes a stream the provider closed. A response without any chunk is -// a provider failure (so fallback can run); a response without a finish -// reason gets one, so clients always see a complete message. -func (s *Stream) end() ([]domain.StreamEvent, error) { - s.done = true - if !s.sawChunk { - return nil, domain.ErrProviderError(http.StatusBadGateway, "provider returned an empty stream") - } - if s.completed { - return nil, nil - } - finish := "stop" - return s.normalize([]domain.StreamEvent{{Type: domain.StreamEventCompleted, FinishReason: &finish, Usage: s.lastUsage}}), nil -} - -// normalize applies the per-stream rules described on Stream. -func (s *Stream) normalize(events []domain.StreamEvent) []domain.StreamEvent { - out := make([]domain.StreamEvent, 0, len(events)+1) - for _, event := range events { - if event.Usage != nil { - s.lastUsage = event.Usage - } - switch { - case event.Type == domain.StreamEventCompleted && s.completed: - // Trailing usage after the finish reason. - case event.Type == domain.StreamEventCompleted && event.FinishReason == nil: - // Usage-only chunk before the model finished: not the end. - event.Type = domain.StreamEventOutputMessageDelta - case event.Type == domain.StreamEventCompleted: - if event.Usage == nil { - event.Usage = s.lastUsage // Usage sent before the finish reason. - } - s.completed = true - } - out = append(out, event) - } - return out -} - -func (s *Stream) Close() error { - s.diagnostics.end("closed", nil) - s.done = true - return s.resp.Body.Close() -} - -type chatCompletionChunk struct { - ID string `json:"id"` - Object string `json:"object"` - Created int64 `json:"created"` - Model string `json:"model"` - Choices []struct { - Index int `json:"index"` - Delta struct { - Role string `json:"role,omitempty"` - Content *string `json:"content,omitempty"` - ReasoningContent *string `json:"reasoning_content,omitempty"` - ToolCalls []struct { - Index int `json:"index"` - ID string `json:"id,omitempty"` - Type string `json:"type,omitempty"` - Function struct { - Name string `json:"name,omitempty"` - Arguments string `json:"arguments,omitempty"` - } `json:"function"` - } `json:"tool_calls,omitempty"` - } `json:"delta"` - FinishReason *string `json:"finish_reason"` - } `json:"choices"` - Usage *struct { - PromptTokens int64 `json:"prompt_tokens"` - CompletionTokens int64 `json:"completion_tokens"` - TotalTokens int64 `json:"total_tokens"` - PromptCacheHitTokens int64 `json:"prompt_cache_hit_tokens"` - PromptTokensDetails *struct { - CachedTokens int64 `json:"cached_tokens"` - } `json:"prompt_tokens_details,omitempty"` - } `json:"usage,omitempty"` - BaseResp *providerBaseResponse `json:"base_resp,omitempty"` - Error *providerStreamErr `json:"error,omitempty"` -} - -// providerStreamErr is an error some providers send as a data line after the -// stream started (overload, context length, moderation). -type providerStreamErr struct { - Message string `json:"message"` - Type string `json:"type"` - Code any `json:"code"` -} - -func providerStreamError(e *providerStreamErr) error { - if e == nil { - return nil - } - message := strings.TrimSpace(e.Message) - if message == "" { - message = "provider stream failed" - } - return domain.ErrProviderError(http.StatusBadGateway, message).WithMeta( - "upstream_error_type", e.Type, - "upstream_code", e.Code, - "upstream_error", message, - ) -} - -type providerBaseResponse struct { - StatusCode int `json:"status_code"` - StatusMsg string `json:"status_msg"` -} - -func providerBaseResponseError(resp *providerBaseResponse) error { - if resp == nil || resp.StatusCode == 0 { - return nil - } - message := strings.TrimSpace(resp.StatusMsg) - if message == "" { - message = "provider returned an unsuccessful response" - } - return domain.ErrProviderError(http.StatusBadGateway, message).WithMeta( - "upstream_code", resp.StatusCode, - "upstream_error", message, - ) -} - -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 { - return []domain.StreamEvent{{ - Type: domain.StreamEventCompleted, - Usage: chunkUsageToDomain(chunk.Usage), - }} - } - - if len(chunk.Choices) == 0 { - return nil - } - - choice := chunk.Choices[0] - - // Some providers (e.g. Novita/MiniMax) send reasoning_content chunks - // with finish_reason set before the actual content chunk. If we honour - // that finish_reason the stream closes before real content arrives. - // Neutralise it so the subsequent content chunk carries the real signal. - // The stream still terminates via upstream "data: [DONE]" / EOF even if - // the later chunk happens to lack finish_reason. - if choice.FinishReason != nil && choice.Delta.ReasoningContent != nil { - hasContent := choice.Delta.Content != nil && *choice.Delta.Content != "" - if !hasContent && len(choice.Delta.ToolCalls) == 0 { - choice.FinishReason = nil - } - } - - // Every tool call in the chunk, in order. Some providers pack several - // calls, or a whole call, into one chunk. - toolEvents := make([]domain.StreamEvent, 0, len(choice.Delta.ToolCalls)) - for _, tc := range choice.Delta.ToolCalls { - tcd := &domain.ToolCallDelta{Index: tc.Index} - if tc.ID != "" { - tcd.ID = &tc.ID - } - if tc.Function.Name != "" { - tcd.Name = &tc.Function.Name - } - if tc.Function.Arguments != "" { - tcd.ArgumentsDelta = &tc.Function.Arguments - } - toolEvents = append(toolEvents, domain.StreamEvent{Type: domain.StreamEventToolCallDelta, ToolCallDelta: tcd}) - } - // Text that shares a chunk with tool calls comes first, as generated. - var text []domain.StreamEvent - if len(toolEvents) > 0 { - content := choice.Delta.Content != nil && *choice.Delta.Content != "" - reasoning := keepReasoning && choice.Delta.ReasoningContent != nil && *choice.Delta.ReasoningContent != "" - if content || reasoning { - event := domain.StreamEvent{Type: domain.StreamEventOutputTextDelta} - if content { - event.ContentDelta = choice.Delta.Content - } - if reasoning { - event.ReasoningDelta = choice.Delta.ReasoningContent - } - text = append(text, event) - } - } - - // When a chunk carries tool-call deltas AND a finish_reason (e.g. - // MiniMax packs the final argument fragment and the finish signal into - // one chunk), the deltas go first so the client has complete arguments - // before the stream is marked done. - if choice.FinishReason != nil && len(toolEvents) > 0 { - completedEvent := domain.StreamEvent{ - Type: domain.StreamEventCompleted, - FinishReason: choice.FinishReason, - } - if chunk.Usage != nil { - completedEvent.Usage = chunkUsageToDomain(chunk.Usage) - } - events = append(text, toolEvents...) - return append(events, completedEvent) - } - - // Check for finish reason -> completed event - if choice.FinishReason != nil { - event := domain.StreamEvent{ - Type: domain.StreamEventCompleted, - FinishReason: choice.FinishReason, - } - if choice.Delta.Content != nil && *choice.Delta.Content != "" { - event.ContentDelta = choice.Delta.Content - } - if keepReasoning && choice.Delta.ReasoningContent != nil && *choice.Delta.ReasoningContent != "" { - event.ReasoningDelta = choice.Delta.ReasoningContent - } - if chunk.Usage != nil { - event.Usage = chunkUsageToDomain(chunk.Usage) - } - return []domain.StreamEvent{event} - } - - if len(toolEvents) > 0 { - return append(text, toolEvents...) - } - - // Text content delta - if choice.Delta.Content != nil { - event := domain.StreamEvent{ - Type: domain.StreamEventOutputTextDelta, - ContentDelta: choice.Delta.Content, - } - if keepReasoning { - event.ReasoningDelta = choice.Delta.ReasoningContent - } - if choice.Delta.Role != "" { - event.Role = &choice.Delta.Role - } - return []domain.StreamEvent{event} - } - - // Reasoning content delta - if keepReasoning && choice.Delta.ReasoningContent != nil { - event := domain.StreamEvent{ - Type: domain.StreamEventOutputTextDelta, - ReasoningDelta: choice.Delta.ReasoningContent, - } - if choice.Delta.Role != "" { - event.Role = &choice.Delta.Role - } - return []domain.StreamEvent{event} - } - - // Role-only delta (first chunk often) - if choice.Delta.Role != "" { - role := choice.Delta.Role - return []domain.StreamEvent{{ - Type: domain.StreamEventOutputMessageDelta, - Role: &role, - }} - } - - return nil -} - -func chunkUsageToDomain(u *struct { - PromptTokens int64 `json:"prompt_tokens"` - CompletionTokens int64 `json:"completion_tokens"` - TotalTokens int64 `json:"total_tokens"` - PromptCacheHitTokens int64 `json:"prompt_cache_hit_tokens"` - PromptTokensDetails *struct { - CachedTokens int64 `json:"cached_tokens"` - } `json:"prompt_tokens_details,omitempty"` -}) *domain.Usage { - if u == nil { - return nil - } - usage := &domain.Usage{ - PromptTokens: u.PromptTokens, - CompletionTokens: u.CompletionTokens, - TotalTokens: u.TotalTokens, - CacheReadTokens: u.PromptCacheHitTokens, - } - if u.PromptTokensDetails != nil && u.PromptTokensDetails.CachedTokens > 0 { - usage.CacheReadTokens = u.PromptTokensDetails.CachedTokens - } - return usage -} diff --git a/apps/gateway/internal/adapters/providers/openai/stream_diagnostics.go b/apps/gateway/internal/adapters/providers/openai/stream_diagnostics.go deleted file mode 100644 index cc6f3ac..0000000 --- a/apps/gateway/internal/adapters/providers/openai/stream_diagnostics.go +++ /dev/null @@ -1,95 +0,0 @@ -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 deleted file mode 100644 index 9dd78ba..0000000 --- a/apps/gateway/internal/adapters/providers/openai/stream_diagnostics_test.go +++ /dev/null @@ -1,133 +0,0 @@ -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 + "{SECRET-FRAMING}\n", end: "eof", level: "warn", unsupported: 1}, - {name: "data without space", data: finish + "data:{SECRET-FRAMING}\n", end: "eof", level: "warn", malformed: 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 deleted file mode 100644 index 8cdb605..0000000 --- a/apps/gateway/internal/adapters/providers/openai/stream_test.go +++ /dev/null @@ -1,231 +0,0 @@ -package openai - -import ( - "encoding/json" - "io" - "net/http" - "strings" - "testing" - - "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 -} - -func (fakeBody) Close() error { return nil } - -func newTestStream(sseData string) *Stream { - body := fakeBody{strings.NewReader(sseData)} - resp := &http.Response{Body: body} - return NewStream(resp) -} - -func TestStream_ToolCallDeltaAndFinishReasonInSameChunk(t *testing.T) { - sseData := "data: {\"id\":\"c1\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"c1\",\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":0,\"id\":\"call_1\",\"type\":\"function\",\"function\":{\"name\":\"write_file\",\"arguments\":\"\"}}]},\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"c1\",\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":0,\"function\":{\"arguments\":\"{\\\"path\\\":\\\"f.txt\\\"\"}}]},\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"c1\",\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":0,\"function\":{\"arguments\":\"}\"}}]},\"finish_reason\":\"tool_calls\"}]}\n\n" + - "data: [DONE]\n" - - body := fakeBody{strings.NewReader(sseData)} - stream := NewStream(&http.Response{Body: body}, "deepseek") - - ev1, err := stream.Recv() - if err != nil { - t.Fatalf("ev1: %v", err) - } - if ev1.Type != domain.StreamEventOutputMessageDelta { - t.Fatalf("ev1 type = %v, want OutputMessageDelta", ev1.Type) - } - - ev2, err := stream.Recv() - if err != nil { - t.Fatalf("ev2: %v", err) - } - if ev2.Type != domain.StreamEventToolCallDelta { - t.Fatalf("ev2 type = %v, want ToolCallDelta", ev2.Type) - } - - ev3, err := stream.Recv() - if err != nil { - t.Fatalf("ev3: %v", err) - } - if ev3.Type != domain.StreamEventToolCallDelta { - t.Fatalf("ev3 type = %v, want ToolCallDelta", ev3.Type) - } - - // The combined chunk with tool_calls + finish_reason must emit - // the tool call delta FIRST, then the completed event. - ev4, err := stream.Recv() - if err != nil { - t.Fatalf("ev4: %v", err) - } - if ev4.Type != domain.StreamEventToolCallDelta { - t.Fatalf("ev4 type = %v, want ToolCallDelta (deferred finish)", ev4.Type) - } - if ev4.ToolCallDelta == nil || ev4.ToolCallDelta.ArgumentsDelta == nil { - t.Fatal("ev4: missing arguments delta") - } - if *ev4.ToolCallDelta.ArgumentsDelta != "}" { - t.Fatalf("ev4 args = %q, want %q", *ev4.ToolCallDelta.ArgumentsDelta, "}") - } - - ev5, err := stream.Recv() - if err != nil { - t.Fatalf("ev5: %v", err) - } - if ev5.Type != domain.StreamEventCompleted { - t.Fatalf("ev5 type = %v, want Completed", ev5.Type) - } - if ev5.FinishReason == nil || *ev5.FinishReason != "tool_calls" { - t.Fatalf("ev5 finish_reason = %v, want tool_calls", ev5.FinishReason) - } - - _, err = stream.Recv() - if err != io.EOF { - t.Fatalf("expected EOF, got %v", err) - } -} - -func TestStream_FinishReasonWithoutToolCalls(t *testing.T) { - sseData := "data: {\"id\":\"c2\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"Hi\"},\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"c2\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n" + - "data: [DONE]\n" - - body := fakeBody{strings.NewReader(sseData)} - stream := NewStream(&http.Response{Body: body}, "deepseek") - - ev1, err := stream.Recv() - if err != nil { - t.Fatalf("ev1: %v", err) - } - if ev1.Type != domain.StreamEventOutputTextDelta { - t.Fatalf("ev1 type = %v, want OutputTextDelta", ev1.Type) - } - - ev2, err := stream.Recv() - if err != nil { - t.Fatalf("ev2: %v", err) - } - if ev2.Type != domain.StreamEventCompleted { - t.Fatalf("ev2 type = %v, want Completed", ev2.Type) - } - if ev2.FinishReason == nil || *ev2.FinishReason != "stop" { - t.Fatalf("ev2 finish_reason = %v, want stop", ev2.FinishReason) - } - - _, err = stream.Recv() - if err != io.EOF { - t.Fatalf("expected EOF, got %v", err) - } -} - -func TestStream_RejectsNonzeroBaseResponse(t *testing.T) { - stream := newTestStream("data: {\"base_resp\":{\"status_code\":17,\"status_msg\":\"provider error\"}}\n\n") - _, err := stream.Recv() - if err == nil { - t.Fatal("expected provider envelope error") - } - gatewayErr, ok := err.(*domain.GatewayError) - if !ok || gatewayErr.HTTPStatus != http.StatusBadGateway { - t.Fatalf("error = %#v", err) - } -} - -func TestStream_ReasoningContentDelta(t *testing.T) { - sseData := "data: {\"id\":\"c3\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"reasoning_content\":\"thinking\"},\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"c3\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"answer\"},\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"c3\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n" + - "data: [DONE]\n" - - body := fakeBody{strings.NewReader(sseData)} - stream := NewStream(&http.Response{Body: body}, "deepseek") - - ev1, err := stream.Recv() - if err != nil { - t.Fatalf("ev1: %v", err) - } - if ev1.Type != domain.StreamEventOutputTextDelta { - t.Fatalf("ev1 type = %v, want OutputTextDelta", ev1.Type) - } - if ev1.ReasoningDelta == nil || *ev1.ReasoningDelta != "thinking" { - t.Fatalf("ev1 reasoning = %v, want thinking", ev1.ReasoningDelta) - } - - ev2, err := stream.Recv() - if err != nil { - t.Fatalf("ev2: %v", err) - } - if ev2.ContentDelta == nil || *ev2.ContentDelta != "answer" { - t.Fatalf("ev2 content = %v, want answer", ev2.ContentDelta) - } -} - -func TestStream_ReasoningContentIgnoredForNonDeepSeekProviders(t *testing.T) { - sseData := "data: {\"id\":\"c4\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"reasoning_content\":\"thinking\"},\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"c4\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"answer\"},\"finish_reason\":null}]}\n\n" + - "data: [DONE]\n" - - stream := newTestStream(sseData) - - ev1, err := stream.Recv() - if err != nil { - t.Fatalf("ev1: %v", err) - } - if ev1.ReasoningDelta != nil { - t.Fatalf("ev1 reasoning = %v, want nil for non-DeepSeek provider", *ev1.ReasoningDelta) - } - - ev2, err := stream.Recv() - if err != nil { - t.Fatalf("ev2: %v", err) - } - if ev2.ContentDelta == nil || *ev2.ContentDelta != "answer" { - t.Fatalf("ev2 content = %v, want answer", ev2.ContentDelta) - } -} - -func TestStream_MapsDeepSeekCacheHitTokens(t *testing.T) { - stream := newTestStream("data: {\"choices\":[],\"usage\":{\"prompt_tokens\":17,\"completion_tokens\":9,\"total_tokens\":26,\"prompt_cache_hit_tokens\":7}}\n\n") - - // A usage-only chunk is not the end; the stream's end is, with that usage. - var event domain.StreamEvent - for event.Type != domain.StreamEventCompleted { - var err error - if event, err = stream.Recv(); err != nil { - t.Fatalf("Recv returned error: %v", err) - } - } - if event.Usage == nil || event.FinishReason == nil || *event.FinishReason != "stop" { - t.Fatalf("event = %+v, want completed event with usage", event) - } - if event.Usage.CacheReadTokens != 7 { - t.Fatalf("cache-read tokens = %d, want 7", event.Usage.CacheReadTokens) - } -} diff --git a/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go b/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go deleted file mode 100644 index 080fb38..0000000 --- a/apps/gateway/internal/adapters/providers/openai/stream_toolcalls_test.go +++ /dev/null @@ -1,155 +0,0 @@ -package openai - -import ( - "encoding/json" - "errors" - "io" - "strings" - "testing" - - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" -) - -// strictCall is what a client that concatenates every fragment by index (as -// openai-python and most SDKs do) reassembles. -type strictCall struct{ id, name, args string } - -func drain(t *testing.T, sse string) (calls []strictCall, text string, finish string, err error) { - t.Helper() - stream := newTestStream(sse) - for { - event, recvErr := stream.Recv() - if recvErr == io.EOF { - return calls, text, finish, nil - } - if recvErr != nil { - return calls, text, finish, recvErr - } - if event.ContentDelta != nil { - text += *event.ContentDelta - } - if event.FinishReason != nil { - finish = *event.FinishReason - } - if d := event.ToolCallDelta; d != nil { - if d.Index > len(calls) { - t.Fatalf("index %d skips ahead of %d calls", d.Index, len(calls)) - } - if d.Index == len(calls) { - calls = append(calls, strictCall{}) - } - c := &calls[d.Index] - if d.ID != nil { - c.id += *d.ID - } - if d.Name != nil { - c.name += *d.Name - } - if d.ArgumentsDelta != nil { - c.args += *d.ArgumentsDelta - } - } - } -} - -func sse(chunks ...string) string { - var b strings.Builder - for _, c := range chunks { - b.WriteString("data: " + c + "\n\n") - } - return b.String() + "data: [DONE]\n\n" -} - -func wellFormed(t *testing.T, calls []strictCall, names ...string) { - t.Helper() - if len(calls) != len(names) { - t.Fatalf("got %d calls %+v, want %d", len(calls), calls, len(names)) - } - ids := map[string]bool{} - for i, c := range calls { - if c.name != names[i] || c.id == "" || ids[c.id] || !json.Valid([]byte(c.args)) { - t.Fatalf("call %d malformed: %+v", i, c) - } - ids[c.id] = true - } -} - -func TestStream_ToolCallProviderShapes(t *testing.T) { - // The translating path keeps what providers send; it doesn't repair it. - for name, tc := range map[string]struct { - sse string - names []string - text string - finish string - }{ - "several calls packed in one chunk": {sse: sse( - `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"read","arguments":"{\"p\":1}"}},{"index":1,"id":"b","type":"function","function":{"name":"list","arguments":"{}"}}]}}]}`, - `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), - names: []string{"read", "list"}, finish: "tool_calls"}, - "several calls packed with the finish reason": {sse: sse( - `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}},{"index":1,"id":"b","function":{"name":"read","arguments":"{}"}},{"index":2,"id":"c","function":{"name":"read","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}`), - names: []string{"read", "read", "read"}, finish: "tool_calls"}, - "text in the same chunk as a call": {sse: sse( - `{"choices":[{"delta":{"content":"Reading.","tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}}]}}]}`, - `{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`), - names: []string{"read"}, text: "Reading.", finish: "tool_calls"}, - "no finish reason at all": {sse: sse( - `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"read","arguments":"{}"}}]}}]}`), - names: []string{"read"}, finish: "stop"}, - } { - t.Run(name, func(t *testing.T) { - calls, text, finish, err := drain(t, tc.sse) - if err != nil { - t.Fatal(err) - } - wellFormed(t, calls, tc.names...) - if text != tc.text || finish != tc.finish { - t.Fatalf("text %q finish %q", text, finish) - } - }) - } -} - -func TestStream_UsageBeforeFinishDoesNotEndTheStream(t *testing.T) { - stream := newTestStream(sse( - `{"choices":[],"usage":{"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}}`, - `{"choices":[{"delta":{"content":"late"}}]}`, - `{"choices":[{"delta":{},"finish_reason":"stop"}]}`)) - var got []domain.StreamEvent - for { - event, err := stream.Recv() - if err != nil { - break - } - got = append(got, event) - } - if len(got) != 3 || got[0].Type == domain.StreamEventCompleted || got[1].ContentDelta == nil || got[2].Type != domain.StreamEventCompleted || got[2].Usage == nil { - t.Fatalf("events %+v", got) - } -} - -func TestStream_ProviderFailures(t *testing.T) { - for name, data := range map[string]string{ - "error payload mid-stream": "data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\n\ndata: {\"error\":{\"message\":\"overloaded\",\"type\":\"server_error\"}}\n\n", - "empty stream": "data: [DONE]\n\n", - "nothing at all": "", - } { - t.Run(name, func(t *testing.T) { - _, _, _, err := drain(t, data) - var gwErr *domain.GatewayError - if !errors.As(err, &gwErr) { - t.Fatalf("err = %v, want a provider error", err) - } - }) - } -} - -func TestStream_LargeSingleLineToolCall(t *testing.T) { - big := strings.Repeat("x", 3<<20) - calls, _, _, err := drain(t, sse( - `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"a","function":{"name":"write","arguments":"{\"text\":\"`+big+`\"}"}}]},"finish_reason":"tool_calls"}]}`)) - if err != nil { - t.Fatal(err) - } - wellFormed(t, calls, "write") -} diff --git a/apps/gateway/internal/adapters/providers/registry/registry.go b/apps/gateway/internal/adapters/providers/registry/registry.go index 69ef0fe..4e363d7 100644 --- a/apps/gateway/internal/adapters/providers/registry/registry.go +++ b/apps/gateway/internal/adapters/providers/registry/registry.go @@ -8,25 +8,25 @@ import ( // Registry selects provider adapters by name. type Registry struct { - providers map[string]ports.GenerationProvider - defaultProvider ports.GenerationProvider + providers map[string]ports.Provider + defaultProvider ports.Provider } func NewRegistry() *Registry { - return &Registry{providers: make(map[string]ports.GenerationProvider)} + return &Registry{providers: make(map[string]ports.Provider)} } -func (r *Registry) Register(name string, provider ports.GenerationProvider) { +func (r *Registry) Register(name string, provider ports.Provider) { r.providers[name] = provider } // SetDefault sets a fallback adapter returned when no provider is explicitly // registered under the requested name. -func (r *Registry) SetDefault(provider ports.GenerationProvider) { +func (r *Registry) SetDefault(provider ports.Provider) { r.defaultProvider = provider } -func (r *Registry) GetProvider(providerName string) (ports.GenerationProvider, error) { +func (r *Registry) GetProvider(providerName string) (ports.Provider, error) { if p, ok := r.providers[providerName]; ok { return p, nil } diff --git a/apps/gateway/internal/adapters/providers/tinfoil/adapter.go b/apps/gateway/internal/adapters/providers/tinfoil/adapter.go index 2a32440..3b57a71 100644 --- a/apps/gateway/internal/adapters/providers/tinfoil/adapter.go +++ b/apps/gateway/internal/adapters/providers/tinfoil/adapter.go @@ -37,9 +37,9 @@ type VerifiedClientFactory interface { NewVerifiedClient(ctx context.Context, model domain.PublicModel) (VerifiedClient, error) } -// Adapter is the Tinfoil provider adapter. It reuses the OpenAI-compatible -// request/response mapper, but all HTTP traffic goes through Tinfoil's -// attested EHBP client instead of the generic OpenAI adapter. +// Adapter is the Tinfoil provider adapter. It prepares bodies like the +// OpenAI-compatible adapter, but all HTTP traffic goes through Tinfoil's +// attested EHBP client. type Adapter struct { timeout time.Duration factory VerifiedClientFactory @@ -61,59 +61,6 @@ func NewAdapterWithFactory(timeout time.Duration, factory VerifiedClientFactory, return &Adapter{timeout: timeout, factory: factory, logger: l} } -func (a *Adapter) Generate(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) { - apiKey := os.Getenv(model.ProviderConfig.APIKeySecretRef) - if apiKey == "" { - return domain.GenerateResult{}, missingProviderCredentialError() - } - - body := openai.BuildRequestBody(req, model) - body["stream"] = false - delete(body, "stream_options") - - verified, proof, err := a.newVerifiedClient(ctx, model) - if err != nil { - return domain.GenerateResult{}, err - } - - respBody, err := a.do(ctx, verified, apiKey, body) - if err != nil { - return domain.GenerateResult{}, openai.MapProviderErrorWithCompatibilityContext(err, model, body) - } - - result, err := openai.ParseResponse(respBody, req, model) - if err != nil { - return domain.GenerateResult{}, err - } - proof.ProviderResponseID = result.ID - result.TinfoilProof = proof - return result, nil -} - -func (a *Adapter) StreamGenerate(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (ports.GenerationStream, error) { - apiKey := os.Getenv(model.ProviderConfig.APIKeySecretRef) - if apiKey == "" { - return nil, missingProviderCredentialError() - } - - body := openai.BuildRequestBody(req, model) - - verified, proof, err := a.newVerifiedClient(ctx, model) - if err != nil { - return nil, err - } - - resp, err := a.doStream(ctx, verified, apiKey, body) - if err != nil { - return nil, openai.MapProviderErrorWithCompatibilityContext(err, model, body) - } - - return &Stream{ - inner: openai.NewStream(resp, providerName).WithDiagnostics(ctx, a.logger, model, 1), - proof: proof, - }, nil -} - func missingProviderCredentialError() *domain.GatewayError { return domain.ErrInternal("an internal error occurred").WithMeta( "provider", providerName, @@ -256,31 +203,6 @@ func tinfoilSDKVersion() string { return "github.com/tinfoilsh/tinfoil-go" } -// Stream decorates the OpenAI-compatible stream with Tinfoil proof evidence. -type Stream struct { - inner ports.GenerationStream - proof *domain.TinfoilTransportProof -} - -func (s *Stream) Recv() (domain.StreamEvent, error) { - return s.inner.Recv() -} - -func (s *Stream) Close() error { - return s.inner.Close() -} - -func (s *Stream) VerifiedTransportProof() *domain.TinfoilTransportProof { - if s.proof == nil { - return nil - } - cp := *s.proof - if len(s.proof.VerificationEvidenceJSON) > 0 { - cp.VerificationEvidenceJSON = json.RawMessage(append([]byte(nil), s.proof.VerificationEvidenceJSON...)) - } - return &cp -} - // SDKClientFactory is the production Tinfoil SDK factory. type SDKClientFactory struct{} @@ -341,30 +263,30 @@ func (c *sdkVerifiedClient) GroundTruth() *verifierclient.GroundTruth { return c.groundTruth } -// ProxyStream forwards a streaming request over the attested transport and +// Stream forwards a streaming request over the attested transport and // returns the SSE body untouched, with the transport proof. -func (a *Adapter) ProxyStream(ctx context.Context, raw []byte, model domain.PublicModel) (ports.ProxyResponse, error) { +func (a *Adapter) Stream(ctx context.Context, raw []byte, model domain.PublicModel) (ports.ProviderStream, error) { apiKey := os.Getenv(model.ProviderConfig.APIKeySecretRef) if apiKey == "" { - return ports.ProxyResponse{}, missingProviderCredentialError() + return ports.ProviderStream{}, missingProviderCredentialError() } body, _, err := openai.PrepareProxyBody(raw, model, true) if err != nil { - return ports.ProxyResponse{}, err + return ports.ProviderStream{}, err } verified, proof, err := a.newVerifiedClient(ctx, model) if err != nil { - return ports.ProxyResponse{}, err + return ports.ProviderStream{}, err } resp, err := a.doStream(ctx, verified, apiKey, body) if err != nil { - return ports.ProxyResponse{}, openai.MapProviderErrorWithCompatibilityContext(err, model, body) + return ports.ProviderStream{}, openai.MapProviderErrorWithCompatibilityContext(err, model, body) } - return ports.ProxyResponse{Body: resp.Body, Proof: proof}, nil + return ports.ProviderStream{Body: resp.Body, Proof: proof}, nil } -// ProxyJSON forwards a non-streaming request over the attested transport. -func (a *Adapter) ProxyJSON(ctx context.Context, raw []byte, model domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) { +// Complete forwards a non-streaming request over the attested transport. +func (a *Adapter) Complete(ctx context.Context, raw []byte, model domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) { apiKey := os.Getenv(model.ProviderConfig.APIKeySecretRef) if apiKey == "" { return nil, nil, missingProviderCredentialError() diff --git a/apps/gateway/internal/adapters/providers/tinfoil/adapter_error_test.go b/apps/gateway/internal/adapters/providers/tinfoil/adapter_error_test.go index 269001f..26aa66b 100644 --- a/apps/gateway/internal/adapters/providers/tinfoil/adapter_error_test.go +++ b/apps/gateway/internal/adapters/providers/tinfoil/adapter_error_test.go @@ -15,12 +15,12 @@ func TestAdapterMissingCredentialReturnsInternalError(t *testing.T) { name string call func() error }{ - {name: "generate", call: func() error { - _, err := adapter.Generate(context.Background(), domain.GenerateRequest{}, model) + {name: "complete", call: func() error { + _, _, err := adapter.Complete(context.Background(), []byte(`{"messages":[]}`), model) return err }}, {name: "stream", call: func() error { - _, err := adapter.StreamGenerate(context.Background(), domain.GenerateRequest{}, model) + _, err := adapter.Stream(context.Background(), []byte(`{"messages":[]}`), model) return err }}, } diff --git a/apps/gateway/internal/adapters/providers/tinfoil/adapter_test.go b/apps/gateway/internal/adapters/providers/tinfoil/adapter_test.go index cbf1c5b..df565e5 100644 --- a/apps/gateway/internal/adapters/providers/tinfoil/adapter_test.go +++ b/apps/gateway/internal/adapters/providers/tinfoil/adapter_test.go @@ -8,6 +8,7 @@ import ( "io" "net/http" "net/http/httptest" + "strings" "testing" "time" @@ -15,7 +16,7 @@ import ( verifierclient "github.com/tinfoilsh/tinfoil-go/verifier/client" ) -func TestGenerateUsesVerifiedClientAndStoresProofEvidence(t *testing.T) { +func TestCompleteUsesVerifiedClientAndReturnsProofEvidence(t *testing.T) { t.Setenv("TINFOIL_API_KEY", "test-key") var seenRequest bool @@ -31,7 +32,7 @@ func TestGenerateUsesVerifiedClientAndStoresProofEvidence(t *testing.T) { if err := json.NewDecoder(r.Body).Decode(&body); err != nil { t.Fatalf("decode request: %v", err) } - if body["model"] != "kimi-k2-6" || body["stream"] != false { + if body["model"] != "kimi-k2-6" || body["stream"] == true { t.Fatalf("unexpected request body: %#v", body) } w.Header().Set("Content-Type", "application/json") @@ -46,31 +47,28 @@ func TestGenerateUsesVerifiedClientAndStoresProofEvidence(t *testing.T) { t.Setenv("TINFOIL_PROXY_BASE_URL", server.URL) factory := &fakeFactory{client: fakeClient(server.Client())} - result, err := NewAdapterWithFactory(time.Second, factory).Generate(context.Background(), simpleRequest(false), simpleModel(server.URL)) + body, proof, err := NewAdapterWithFactory(time.Second, factory).Complete(context.Background(), simpleRequest(), simpleModel(server.URL)) if err != nil { - t.Fatalf("Generate returned error: %v", err) + t.Fatalf("Complete returned error: %v", err) } if !seenRequest { t.Fatalf("provider request was not sent") } - if result.ID != "cmpl-tinfoil-1" || result.TinfoilProof == nil { - t.Fatalf("unexpected result/proof: id=%q proof=%#v", result.ID, result.TinfoilProof) + if !strings.Contains(string(body), `"id":"cmpl-tinfoil-1"`) || proof == nil { + t.Fatalf("unexpected body/proof: body=%s proof=%#v", body, proof) } - if result.TinfoilProof.ProviderResponseID != "cmpl-tinfoil-1" { - t.Fatalf("proof response id was not filled: %#v", result.TinfoilProof.ProviderResponseID) + if proof.EnclaveHost == nil || *proof.EnclaveHost != "inference.tinfoil.sh" { + t.Fatalf("unexpected enclave host in proof: %#v", proof.EnclaveHost) } - if result.TinfoilProof.EnclaveHost == nil || *result.TinfoilProof.EnclaveHost != "inference.tinfoil.sh" { - t.Fatalf("unexpected enclave host in proof: %#v", result.TinfoilProof.EnclaveHost) + if proof.TransportMode == nil || *proof.TransportMode != "ehbp" { + t.Fatalf("unexpected transport mode in proof: %#v", proof.TransportMode) } - if result.TinfoilProof.TransportMode == nil || *result.TinfoilProof.TransportMode != "ehbp" { - t.Fatalf("unexpected transport mode in proof: %#v", result.TinfoilProof.TransportMode) - } - if len(result.TinfoilProof.VerificationEvidenceJSON) == 0 { + if len(proof.VerificationEvidenceJSON) == 0 { t.Fatalf("expected verification evidence JSON") } } -func TestStreamGenerateReturnsVerifiedTransportProof(t *testing.T) { +func TestStreamReturnsBodyAndVerifiedTransportProof(t *testing.T) { t.Setenv("TINFOIL_API_KEY", "test-key") server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -82,54 +80,30 @@ func TestStreamGenerateReturnsVerifiedTransportProof(t *testing.T) { defer server.Close() t.Setenv("TINFOIL_PROXY_BASE_URL", server.URL) - stream, err := NewAdapterWithFactory(time.Second, &fakeFactory{client: fakeClient(server.Client())}).StreamGenerate( + stream, err := NewAdapterWithFactory(time.Second, &fakeFactory{client: fakeClient(server.Client())}).Stream( context.Background(), - simpleRequest(true), + simpleRequest(), simpleModel(server.URL), ) if err != nil { - t.Fatalf("StreamGenerate returned error: %v", err) + t.Fatalf("Stream returned error: %v", err) } - defer stream.Close() - - var text string - var completed bool - for { - event, recvErr := stream.Recv() - if recvErr == io.EOF { - break - } - if recvErr != nil { - t.Fatalf("Recv returned error: %v", recvErr) - } - if event.ContentDelta != nil { - text += *event.ContentDelta - } - if event.Type == domain.StreamEventCompleted { - completed = true - } - } - if text != "hi" || !completed { - t.Fatalf("unexpected stream events: text=%q completed=%v", text, completed) + defer stream.Body.Close() + body, _ := io.ReadAll(stream.Body) + if !strings.Contains(string(body), `"content":"hi"`) || !strings.Contains(string(body), "[DONE]") { + t.Fatalf("stream not returned as sent: %s", body) } - proofProvider, ok := stream.(interface { - VerifiedTransportProof() *domain.TinfoilTransportProof - }) - if !ok { - t.Fatalf("stream does not expose proof evidence") - } - proof := proofProvider.VerifiedTransportProof() - if proof == nil || proof.TransportMode == nil || *proof.TransportMode != "ehbp" { - t.Fatalf("unexpected stream proof: %#v", proof) + if stream.Proof == nil || stream.Proof.TransportMode == nil || *stream.Proof.TransportMode != "ehbp" { + t.Fatalf("unexpected stream proof: %#v", stream.Proof) } } -func TestGenerateFailsClosedWhenAttestationFails(t *testing.T) { +func TestFailsClosedWhenAttestationFails(t *testing.T) { t.Setenv("TINFOIL_API_KEY", "test-key") var called bool factory := &fakeFactory{err: errors.New("attestation failed"), onCall: func() { called = true }} - _, err := NewAdapterWithFactory(time.Second, factory).Generate(context.Background(), simpleRequest(false), simpleModel("https://example.invalid/v1")) + _, _, err := NewAdapterWithFactory(time.Second, factory).Complete(context.Background(), simpleRequest(), simpleModel("https://example.invalid/v1")) if err == nil { t.Fatalf("expected error") } @@ -142,12 +116,12 @@ func TestGenerateFailsClosedWhenAttestationFails(t *testing.T) { } } -func TestGenerateDoesNotVerifyWithoutAPIKey(t *testing.T) { +func TestDoesNotVerifyWithoutAPIKey(t *testing.T) { t.Setenv("TINFOIL_API_KEY", "") var called bool factory := &fakeFactory{client: fakeClient(http.DefaultClient), onCall: func() { called = true }} - _, err := NewAdapterWithFactory(time.Second, factory).Generate(context.Background(), simpleRequest(false), simpleModel("https://example.invalid/v1")) + _, _, err := NewAdapterWithFactory(time.Second, factory).Complete(context.Background(), simpleRequest(), simpleModel("https://example.invalid/v1")) if err == nil { t.Fatalf("expected error") } @@ -156,7 +130,7 @@ func TestGenerateDoesNotVerifyWithoutAPIKey(t *testing.T) { } } -func TestGenerateMapsProviderErrors(t *testing.T) { +func TestMapsProviderErrors(t *testing.T) { t.Setenv("TINFOIL_API_KEY", "test-key") server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -166,9 +140,9 @@ func TestGenerateMapsProviderErrors(t *testing.T) { defer server.Close() t.Setenv("TINFOIL_PROXY_BASE_URL", server.URL) - _, err := NewAdapterWithFactory(time.Second, &fakeFactory{client: fakeClient(server.Client())}).Generate( + _, _, err := NewAdapterWithFactory(time.Second, &fakeFactory{client: fakeClient(server.Client())}).Complete( context.Background(), - simpleRequest(false), + simpleRequest(), simpleModel(server.URL), ) var gwErr *domain.GatewayError @@ -201,18 +175,9 @@ func TestRequestURLUsesExplicitProxyBaseURL(t *testing.T) { } } -func simpleRequest(stream bool) domain.GenerateRequest { - role := "user" - content := "hello" - return domain.GenerateRequest{ - PublicModelID: "tinfoil/kimi-k2-6", - Stream: stream, - Input: []domain.InputItem{{ - Type: domain.InputItemTypeMessage, - Role: &role, - Content: &content, - }}, - } +// simpleRequest is a client's chat request body. +func simpleRequest() []byte { + return []byte(`{"model":"private/kimi-k3","messages":[{"role":"user","content":"hello"}]}`) } func simpleModel(baseURL string) domain.PublicModel { diff --git a/apps/gateway/internal/application/ports/provider.go b/apps/gateway/internal/application/ports/provider.go index 50b1185..eab2dd1 100644 --- a/apps/gateway/internal/application/ports/provider.go +++ b/apps/gateway/internal/application/ports/provider.go @@ -7,39 +7,19 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) -// GenerationProvider translates canonical requests to upstream provider calls. -type GenerationProvider interface { - Generate(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) - StreamGenerate(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (GenerationStream, error) +// Provider forwards a client's OpenAI chat-completions request body upstream +// and returns the provider's response as it came. The adapter changes only +// what its PrepareProxyBody documents (the model name, usage for metering, +// and provider compatibility for valid OpenAI requests). +type Provider interface { + // Stream sends a streaming request and returns the SSE body. + Stream(ctx context.Context, raw []byte, model domain.PublicModel) (ProviderStream, error) + // Complete sends a non-streaming request and returns the JSON body. + Complete(ctx context.Context, raw []byte, model domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) } -// GenerationStream reads canonical stream events from a provider. -type GenerationStream interface { - Recv() (domain.StreamEvent, error) - Close() error -} - -// VerifiedTransportProofProvider exposes safe proof evidence produced by a -// provider adapter whose request transport is verified before user content is -// sent upstream. -type VerifiedTransportProofProvider interface { - VerifiedTransportProof() *domain.TinfoilTransportProof -} - -// ProxyProvider forwards the client's OpenAI chat-completions request body -// upstream (with the model and the few documented edits the adapter makes) -// and returns the provider's response as it came. Providers that speak the OpenAI -// wire format implement it; the gateway then acts as a proxy instead of -// translating through canonical events. -type ProxyProvider interface { - // ProxyStream sends a streaming request and returns the SSE body. - ProxyStream(ctx context.Context, raw []byte, model domain.PublicModel) (ProxyResponse, error) - // ProxyJSON sends a non-streaming request and returns the JSON body. - ProxyJSON(ctx context.Context, raw []byte, model domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) -} - -// ProxyResponse is an upstream streaming response. -type ProxyResponse struct { +// ProviderStream is an upstream streaming response. +type ProviderStream struct { Body io.ReadCloser // Proof is set when the transport was verified before content was sent. Proof *domain.TinfoilTransportProof diff --git a/apps/gateway/internal/application/ports/provider_registry.go b/apps/gateway/internal/application/ports/provider_registry.go index ac4ab96..faf81f1 100644 --- a/apps/gateway/internal/application/ports/provider_registry.go +++ b/apps/gateway/internal/application/ports/provider_registry.go @@ -2,5 +2,5 @@ package ports // ProviderRegistry selects a provider adapter by name. type ProviderRegistry interface { - GetProvider(providerName string) (GenerationProvider, error) + GetProvider(providerName string) (Provider, error) } diff --git a/apps/gateway/internal/application/services/chat_completions_service.go b/apps/gateway/internal/application/services/chat_completions_service.go index 97f4c2c..f7d7625 100644 --- a/apps/gateway/internal/application/services/chat_completions_service.go +++ b/apps/gateway/internal/application/services/chat_completions_service.go @@ -7,7 +7,7 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) -// ChatCompletionsService handles /v1/chat/completions compatibility mapping. +// ChatCompletionsService handles /v1/chat/completions. type ChatCompletionsService struct { generate *GenerateService logger ports.Logger @@ -17,20 +17,8 @@ func NewChatCompletionsService(generate *GenerateService, logger ports.Logger) * return &ChatCompletionsService{generate: generate, logger: logger} } -// Execute processes a non-streaming /v1/chat/completions request. -func (s *ChatCompletionsService) Execute(ctx context.Context, req domain.GenerateRequest, bearerToken string) (domain.GenerateResult, error) { - result, _, err := s.generate.Execute(ctx, domain.EndpointChatCompletions, req, bearerToken) - return result, err -} - -// ExecuteStream processes a streaming /v1/chat/completions request. -func (s *ChatCompletionsService) ExecuteStream(ctx context.Context, req domain.GenerateRequest, bearerToken string) (ports.GenerationStream, *domain.PublicModel, error) { - stream, _, model, err := s.generate.ExecuteStream(ctx, domain.EndpointChatCompletions, req, bearerToken) - return stream, model, err -} - -// Proxy runs a /v1/chat/completions request as a proxy; it returns -// ErrNotProxyable when the request needs the translating path. +// Proxy forwards a request; raw is the client's body, req what the gateway +// read from it. func (s *ChatCompletionsService) Proxy(ctx context.Context, raw []byte, req domain.GenerateRequest, bearerToken string) (*ProxyCall, error) { return s.generate.Proxy(ctx, domain.EndpointChatCompletions, raw, req, bearerToken) } diff --git a/apps/gateway/internal/application/services/generate_service.go b/apps/gateway/internal/application/services/generate_service.go index ed11777..4fc891d 100644 --- a/apps/gateway/internal/application/services/generate_service.go +++ b/apps/gateway/internal/application/services/generate_service.go @@ -4,20 +4,17 @@ import ( "context" "errors" "fmt" - "io" "strconv" "strings" "time" - "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/middleware" "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" ) -// GenerateService is the single canonical execution path for all generation endpoints. +// GenerateService runs chat-completions requests: it authenticates, routes, +// validates, masks PII, reserves credit, forwards to the provider, and meters. type GenerateService struct { auth ports.AuthService catalog ports.ModelCatalog @@ -65,230 +62,6 @@ func (s *GenerateService) SetTinfoilProofRepository(proofs ports.TinfoilTranspor s.tinfoilProofs = proofs } -// Execute runs a non-streaming generation request through the canonical flow. -func (s *GenerateService) Execute(ctx context.Context, endpoint string, req domain.GenerateRequest, bearerToken string) (domain.GenerateResult, domain.AuthContext, error) { - start := time.Now() - requestID := middleware.GetRequestID(ctx) - - authCtx, err := s.auth.AuthenticateAPIKey(ctx, bearerToken) - if err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, req, nil, time.Since(start).Milliseconds()) - return domain.GenerateResult{}, domain.AuthContext{}, err - } - - model, execReq, err := s.resolveModel(ctx, req) - if err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, req, nil, time.Since(start).Milliseconds()) - return domain.GenerateResult{}, authCtx, err - } - - if err := s.validateRequest(endpoint, execReq, model); err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) - s.recordFailure(ctx, nil, &authCtx, endpoint, &execReq, &model, err, nil, time.Since(start).Milliseconds()) - return domain.GenerateResult{}, authCtx, err - } - - piiMapping, err := s.maskRequest(ctx, &execReq, authCtx.APIKey.PIIMode) - if err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) - s.recordFailure(ctx, nil, &authCtx, endpoint, &execReq, &model, err, nil, time.Since(start).Milliseconds()) - return domain.GenerateResult{}, authCtx, err - } - - reservationRequestID := uuid.NewString() - reservationID, err := s.metering.Reserve(ctx, authCtx, endpoint, execReq, model, reservationRequestID) - if err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) - return domain.GenerateResult{}, authCtx, err - } - - result, executionModel, err := s.generateWithFallback(ctx, execReq, model, requestID, piiMapping) - if err != nil { - err = sanitizeErrorWithPIIMapping(err, piiMapping) - fields := s.buildErrorLogFields(ctx, requestID, &authCtx, endpoint, execReq.PublicModelID, executionModel, err, time.Since(start).Milliseconds()) - s.logger.Error("provider generation failed", fields...) - s.recordGeneration(metrics.OutcomeError, endpoint, execReq, executionModel, time.Since(start).Milliseconds()) - s.recordFailure(ctx, &reservationID, &authCtx, endpoint, &execReq, &executionModel, err, nil, time.Since(start).Milliseconds()) - return domain.GenerateResult{}, authCtx, err - } - - s.storeTinfoilProof(ctx, authCtx, executionModel, result) - - latencyMs := time.Since(start).Milliseconds() - metrics.RecordUsage(result.Usage, execReq.PublicModelID, executionModel.ProviderConfig.ProviderName) - s.recordGeneration(metrics.OutcomeSuccess, endpoint, execReq, executionModel, latencyMs) - s.logger.Info("generation completed", - "request_id", requestID, - "account_id", authCtx.Account.ID, - "endpoint", endpoint, - "model", execReq.PublicModelID, - "provider", executionModel.ProviderConfig.ProviderName, - "latency_ms", latencyMs, - "finish_reason", logfields.FinishReason(result.FinishReason), - "gateway_status", 200, - "upstream_status", 200, - ) - - // Provider execution may finish at the same moment the client disconnects. - // Finalizing the reservation must outlive request cancellation so a - // successful upstream call is still charged and its hold is released. - if err := s.metering.RecordSuccess(context.WithoutCancel(ctx), reservationID, authCtx, endpoint, execReq, result, executionModel, latencyMs); err != nil { - s.logger.Error("failed to record usage", - "request_id", requestID, - "reservation_id", reservationID, - "account_id", authCtx.Account.ID, - "error_type", fmt.Sprintf("%T", err), - ) - } - - unmaskResult(&result, piiMapping, s.logger, - "request_id", requestID, - "endpoint", endpoint, - "model", execReq.PublicModelID, - "provider", executionModel.ProviderConfig.ProviderName, - "pii_mode", authCtx.APIKey.PIIMode, - ) - return result, authCtx, nil -} - -// ExecuteStream runs a streaming generation request through the canonical flow. -func (s *GenerateService) ExecuteStream(ctx context.Context, endpoint string, req domain.GenerateRequest, bearerToken string) (ports.GenerationStream, domain.AuthContext, *domain.PublicModel, error) { - start := time.Now() - requestID := middleware.GetRequestID(ctx) - - authCtx, err := s.auth.AuthenticateAPIKey(ctx, bearerToken) - if err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, req, nil, time.Since(start).Milliseconds()) - return nil, domain.AuthContext{}, nil, err - } - - model, execReq, err := s.resolveModel(ctx, req) - if err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, req, nil, time.Since(start).Milliseconds()) - return nil, authCtx, nil, err - } - - if err := s.validateRequest(endpoint, execReq, model); err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) - s.recordFailure(ctx, nil, &authCtx, endpoint, &execReq, &model, err, nil, time.Since(start).Milliseconds()) - return nil, authCtx, &model, err - } - - piiMapping, err := s.maskRequest(ctx, &execReq, authCtx.APIKey.PIIMode) - if err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) - s.recordFailure(ctx, nil, &authCtx, endpoint, &execReq, &model, err, nil, time.Since(start).Milliseconds()) - return nil, authCtx, &model, err - } - - reservationRequestID := uuid.NewString() - reservationID, err := s.metering.Reserve(ctx, authCtx, endpoint, execReq, model, reservationRequestID) - if err != nil { - s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) - return nil, authCtx, &model, err - } - - stream, executionModel, err := s.streamWithFallback(ctx, execReq, model, requestID, piiMapping) - if err != nil { - err = sanitizeErrorWithPIIMapping(err, piiMapping) - fields := s.buildErrorLogFields(ctx, requestID, &authCtx, endpoint, execReq.PublicModelID, executionModel, err, time.Since(start).Milliseconds()) - s.logger.Error("provider stream generation failed", fields...) - s.recordGeneration(metrics.OutcomeError, endpoint, execReq, executionModel, time.Since(start).Milliseconds()) - s.recordFailure(ctx, &reservationID, &authCtx, endpoint, &execReq, &executionModel, err, nil, time.Since(start).Milliseconds()) - return nil, authCtx, &executionModel, err - } - - wrapped := &usageTrackingStream{ - inner: stream, - service: s, - ctx: ctx, - authCtx: authCtx, - endpoint: endpoint, - req: execReq, - model: executionModel, - requestID: requestID, - reservationID: reservationID, - start: start, - piiMapping: piiMapping, - } - - if piiMapping != nil && piiMapping.Len() > 0 { - return newPIIUnmaskingStreamWithLogger(wrapped, piiMapping, s.logger, - "request_id", requestID, - "endpoint", endpoint, - "model", execReq.PublicModelID, - "provider", executionModel.ProviderConfig.ProviderName, - "pii_mode", authCtx.APIKey.PIIMode, - ), authCtx, &executionModel, nil - } - return wrapped, authCtx, &executionModel, nil -} - -func (s *GenerateService) generateWithFallback( - ctx context.Context, - req domain.GenerateRequest, - model domain.PublicModel, - requestID string, - piiMapping *domain.PIIMapping, -) (domain.GenerateResult, domain.PublicModel, error) { - // ponytail: one fallback attempt is intentional; add chains only when a real routing need appears. - result, err := s.generateOnce(ctx, req, model) - if err == nil || !shouldTryFallback(ctx, err, model.Fallback) { - return result, model, err - } - - fallback := withProviderTarget(model, *model.Fallback) - s.logFallback(requestID, model, fallback, sanitizeErrorWithPIIMapping(err, piiMapping)) - result, err = s.generateOnce(ctx, req, fallback) - return result, fallback, err -} - -func (s *GenerateService) generateOnce(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) { - provider, err := s.registry.GetProvider(model.ProviderConfig.ProviderName) - if err != nil { - return domain.GenerateResult{}, err - } - upstreamStart := time.Now() - result, err := provider.Generate(ctx, req, model) - recordUpstreamLatency(req.PublicModelID, model.ProviderConfig.ProviderName, upstreamStart, err) - return result, err -} - -func (s *GenerateService) streamWithFallback( - ctx context.Context, - req domain.GenerateRequest, - model domain.PublicModel, - requestID string, - piiMapping *domain.PIIMapping, -) (ports.GenerationStream, domain.PublicModel, error) { - stream, err := s.streamOnce(ctx, req, model) - if err == nil { - stream, err = primeStream(stream) - } - if err == nil || !shouldTryFallback(ctx, err, model.Fallback) { - return stream, model, err - } - - fallback := withProviderTarget(model, *model.Fallback) - s.logFallback(requestID, model, fallback, sanitizeErrorWithPIIMapping(err, piiMapping)) - stream, err = s.streamOnce(ctx, req, fallback) - if err == nil { - stream, err = primeStream(stream) - } - return stream, fallback, err -} - -func (s *GenerateService) streamOnce(ctx context.Context, req domain.GenerateRequest, model domain.PublicModel) (ports.GenerationStream, error) { - provider, err := s.registry.GetProvider(model.ProviderConfig.ProviderName) - if err != nil { - return nil, err - } - upstreamStart := time.Now() - stream, err := provider.StreamGenerate(ctx, req, model) - recordUpstreamLatency(req.PublicModelID, model.ProviderConfig.ProviderName, upstreamStart, err) - return stream, err -} - func recordUpstreamLatency(publicModelID, providerName string, start time.Time, err error) { outcome := metrics.OutcomeSuccess if err != nil { @@ -327,63 +100,6 @@ func (s *GenerateService) logFallback(requestID string, primary, fallback domain ) } -type bufferedGenerationStream struct { - inner ports.GenerationStream - events []domain.StreamEvent -} - -func primeStream(stream ports.GenerationStream) (ports.GenerationStream, error) { - buffered := &bufferedGenerationStream{inner: stream} - for { - event, err := stream.Recv() - if err != nil { - if err == io.EOF { - return buffered, nil - } - _ = stream.Close() - return nil, err - } - if event.Type == domain.StreamEventError { - _ = stream.Close() - if event.Error != nil { - return nil, event.Error - } - return nil, domain.ErrProviderError(502, "provider stream failed before output") - } - buffered.events = append(buffered.events, event) - if streamEventCommitsResponse(event) { - return buffered, nil - } - } -} - -func streamEventCommitsResponse(event domain.StreamEvent) bool { - return event.Type == domain.StreamEventCompleted || - (event.ContentDelta != nil && *event.ContentDelta != "") || - (event.ReasoningDelta != nil && *event.ReasoningDelta != "") || - event.ToolCallDelta != nil -} - -func (s *bufferedGenerationStream) Recv() (domain.StreamEvent, error) { - if len(s.events) == 0 { - return s.inner.Recv() - } - event := s.events[0] - s.events = s.events[1:] - return event, nil -} - -func (s *bufferedGenerationStream) Close() error { - return s.inner.Close() -} - -func (s *bufferedGenerationStream) VerifiedTransportProof() *domain.TinfoilTransportProof { - if proof, ok := s.inner.(ports.VerifiedTransportProofProvider); ok { - return proof.VerifiedTransportProof() - } - return nil -} - func (s *GenerateService) resolveModel(ctx context.Context, req domain.GenerateRequest) (domain.PublicModel, domain.GenerateRequest, error) { requestedModelID := req.PublicModelID if req.RequestedModelID == "" { @@ -482,7 +198,7 @@ func (s *GenerateService) validateRequest(endpoint string, req domain.GenerateRe return domain.ErrUnsupportedFeature("parallel_tool_calls") } - if req.TextConfig != nil && !model.SupportsStructuredOutput { + if req.StructuredOutput && !model.SupportsStructuredOutput { return domain.ErrUnsupportedFeature("structured_output") } @@ -494,147 +210,6 @@ func (s *GenerateService) validateRequest(endpoint string, req domain.GenerateRe return nil } -// usageTrackingStream wraps a GenerationStream to record usage on completion. -type usageTrackingStream struct { - inner ports.GenerationStream - service *GenerateService - ctx context.Context - authCtx domain.AuthContext - endpoint string - req domain.GenerateRequest - model domain.PublicModel - requestID string - reservationID string - start time.Time - piiMapping *domain.PIIMapping - lastUsage *domain.Usage - 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 - } - s.recordCompletion(err) - return event, err - } - - if event.Usage != nil { - s.lastUsage = event.Usage - } - if event.ProviderResponseID != "" { - s.providerResponseID = event.ProviderResponseID - } - - if event.Type == domain.StreamEventCompleted { - if event.FinishReason != nil { - s.finishReason = event.FinishReason - } - } - - if event.Type == domain.StreamEventError { - s.endReason = "provider_error_event" - var gwErr error - if event.Error != nil { - gwErr = event.Error - } - s.recordCompletion(gwErr) - } - - return event, nil -} - -func (s *usageTrackingStream) Close() error { - if !s.finished { - s.endReason = "closed_before_eof" - s.recordCompletion(context.Canceled) - } - return s.inner.Close() -} - -func (s *usageTrackingStream) recordCompletion(err error) { - if s.finished { - return - } - s.finished = true - ctx := context.WithoutCancel(s.ctx) - latencyMs := time.Since(s.start).Milliseconds() - 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) - return - } - // Stream ended normally (io.EOF) — record success with accumulated usage. - result := domain.GenerateResult{ - ID: s.providerResponseID, - PublicModelID: s.model.PublicModelID, - ProviderName: s.model.ProviderConfig.ProviderName, - ProviderModelID: s.model.ProviderModelID, - FinishReason: s.finishReason, - Usage: s.lastUsage, - } - if verified, ok := s.inner.(ports.VerifiedTransportProofProvider); ok { - result.TinfoilProof = verified.VerifiedTransportProof() - } - 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) - 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 { - 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(), - } -} - func (s *GenerateService) recordFailure(ctx context.Context, reservationID *string, auth *domain.AuthContext, endpoint string, req *domain.GenerateRequest, model *domain.PublicModel, err error, partialUsage *domain.Usage, latencyMs int64) { if s.metering == nil { return diff --git a/apps/gateway/internal/application/services/generate_service_test.go b/apps/gateway/internal/application/services/generate_service_test.go index 95c8572..81266cc 100644 --- a/apps/gateway/internal/application/services/generate_service_test.go +++ b/apps/gateway/internal/application/services/generate_service_test.go @@ -3,6 +3,7 @@ package services import ( "context" "io" + "strings" "testing" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/middleware" @@ -113,11 +114,11 @@ func (s *stubUsageMeter) RecordFailure(ctx context.Context, reservationID *strin type stubProviderRegistry struct { provider *stubProvider - providers map[string]ports.GenerationProvider + providers map[string]ports.Provider err error } -func (s *stubProviderRegistry) GetProvider(name string) (ports.GenerationProvider, error) { +func (s *stubProviderRegistry) GetProvider(name string) (ports.Provider, error) { if s.err != nil { return nil, s.err } @@ -131,41 +132,52 @@ func (s *stubProviderRegistry) GetProvider(name string) (ports.GenerationProvide return s.provider, nil } +// stubProvider returns canned provider bodies. type stubProvider struct { - beforeReturn func() - err error - stream ports.GenerationStream - streamErr error - generateCalls int - streamCalls int + beforeReturn func() + err error + json string // Complete's body; a short answer when empty. + sse string // Stream's body. + calls int + lastRaw string } -func (s *stubProvider) Generate(_ context.Context, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) { - s.generateCalls++ +const stubAnswer = `{"id":"provider-response","model":"upstream","choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}` + +func (s *stubProvider) Complete(_ context.Context, raw []byte, _ domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) { + s.calls++ + s.lastRaw = string(raw) if s.beforeReturn != nil { s.beforeReturn() } if s.err != nil { - return domain.GenerateResult{}, s.err - } - return domain.GenerateResult{ - ID: "provider-response", - PublicModelID: req.PublicModelID, - ProviderName: model.ProviderConfig.ProviderName, - ProviderModelID: model.ProviderModelID, - }, nil + return nil, nil, s.err + } + if s.json == "" { + return []byte(stubAnswer), nil, nil + } + return []byte(s.json), nil, nil } -func TestGenerateServiceExecute_FinalizesSuccessAfterRequestCancellation(t *testing.T) { +func (s *stubProvider) Stream(_ context.Context, raw []byte, _ domain.PublicModel) (ports.ProviderStream, error) { + s.calls++ + s.lastRaw = string(raw) + if s.err != nil { + return ports.ProviderStream{}, s.err + } + return ports.ProviderStream{Body: io.NopCloser(strings.NewReader(s.sse))}, nil +} + +func TestProxy_FinalizesSuccessAfterRequestCancellation(t *testing.T) { ctx, cancel := context.WithCancel(context.WithValue(context.Background(), middleware.RequestIDKey, "client-request-id")) meter := &stubUsageMeter{} svc := newDirectModelGenerateService(meter, &stubProvider{beforeReturn: cancel}) - _, _, err := svc.Execute(ctx, domain.EndpointChatCompletions, domain.GenerateRequest{ + _, err := svc.Proxy(ctx, domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{ PublicModelID: "model", }, "sk-test") if err != nil { - t.Fatalf("Execute returned error: %v", err) + t.Fatalf("Proxy returned error: %v", err) } if meter.successCalls != 1 || meter.lastReservationID != "reservation-1" { t.Fatalf("success finalization = %d for %q, want one call for reservation-1", meter.successCalls, meter.lastReservationID) @@ -181,7 +193,7 @@ func TestGenerateServiceExecute_FinalizesSuccessAfterRequestCancellation(t *test } } -func TestGenerateServiceExecute_ReleasesReservationAfterRequestCancellation(t *testing.T) { +func TestProxy_ReleasesReservationAfterRequestCancellation(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) meter := &stubUsageMeter{} svc := newDirectModelGenerateService(meter, &stubProvider{ @@ -189,7 +201,7 @@ func TestGenerateServiceExecute_ReleasesReservationAfterRequestCancellation(t *t err: domain.ErrProviderUnavailable("openai"), }) - _, _, err := svc.Execute(ctx, domain.EndpointChatCompletions, domain.GenerateRequest{ + _, err := svc.Proxy(ctx, domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{ PublicModelID: "model", }, "sk-test") if err == nil { @@ -203,136 +215,64 @@ func TestGenerateServiceExecute_ReleasesReservationAfterRequestCancellation(t *t } } -func TestGenerateServiceExecute_UsesConfiguredFallbackOnce(t *testing.T) { +func TestProxy_UsesConfiguredFallbackOnce(t *testing.T) { meter := &stubUsageMeter{} primary := &stubProvider{err: domain.ErrProviderUnavailable("primary")} fallback := &stubProvider{} svc := newFallbackGenerateService(meter, primary, fallback) - result, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{ PublicModelID: "model", }, "sk-test") if err != nil { - t.Fatalf("Execute returned error: %v", err) + t.Fatalf("Proxy returned error: %v", err) } - if primary.generateCalls != 1 || fallback.generateCalls != 1 { - t.Fatalf("generate calls = primary %d, fallback %d; want 1 each", primary.generateCalls, fallback.generateCalls) + if primary.calls != 1 || fallback.calls != 1 { + t.Fatalf("generate calls = primary %d, fallback %d; want 1 each", primary.calls, fallback.calls) } if meter.reserveCalls != 1 || meter.successCalls != 1 || meter.failureCalls != 0 { t.Fatalf("metering calls = reserve %d, success %d, failure %d", meter.reserveCalls, meter.successCalls, meter.failureCalls) } - if result.ProviderName != "fallback" || result.ProviderModelID != "fallback-model" { - t.Fatalf("result target = %s/%s, want fallback/fallback-model", result.ProviderName, result.ProviderModelID) + if call.Model.ProviderConfig.ProviderName != "fallback" || call.Model.ProviderModelID != "fallback-model" { + t.Fatalf("result target = %s/%s, want fallback/fallback-model", call.Model.ProviderConfig.ProviderName, call.Model.ProviderModelID) } if meter.lastSuccessModel.ProviderConfig.ProviderName != "fallback" { t.Fatalf("metered provider = %q, want fallback", meter.lastSuccessModel.ProviderConfig.ProviderName) } } -func TestGenerateServiceExecute_ReturnsFallbackFailure(t *testing.T) { +func TestProxy_ReturnsFallbackFailure(t *testing.T) { meter := &stubUsageMeter{} primary := &stubProvider{err: domain.ErrProviderUnavailable("primary")} fallbackErr := domain.ErrProviderTimeout("fallback") fallback := &stubProvider{err: fallbackErr} svc := newFallbackGenerateService(meter, primary, fallback) - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ + _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{ PublicModelID: "model", }, "sk-test") if err != fallbackErr { t.Fatalf("error = %v, want fallback error %v", err, fallbackErr) } - if primary.generateCalls != 1 || fallback.generateCalls != 1 || meter.failureCalls != 1 { - t.Fatalf("calls = primary %d, fallback %d, failures %d", primary.generateCalls, fallback.generateCalls, meter.failureCalls) + if primary.calls != 1 || fallback.calls != 1 || meter.failureCalls != 1 { + t.Fatalf("calls = primary %d, fallback %d, failures %d", primary.calls, fallback.calls, meter.failureCalls) } } -func TestGenerateServiceExecute_DoesNotFallbackAfterCancellation(t *testing.T) { +func TestProxy_DoesNotFallbackAfterCancellation(t *testing.T) { meter := &stubUsageMeter{} primary := &stubProvider{err: domain.ErrClientCanceled()} fallback := &stubProvider{} svc := newFallbackGenerateService(meter, primary, fallback) - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ + _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{ PublicModelID: "model", }, "sk-test") if err == nil { t.Fatal("expected cancellation error") } - if fallback.generateCalls != 0 { - t.Fatalf("fallback calls = %d, want 0", fallback.generateCalls) - } -} - -func TestGenerateServiceExecuteStream_FallsBackBeforeVisibleOutput(t *testing.T) { - role := "assistant" - content := "fallback answer" - finish := "stop" - primaryStream := &stubStream{steps: []streamStep{ - {event: domain.StreamEvent{Type: domain.StreamEventOutputMessageDelta, Role: &role}}, - {err: domain.ErrProviderUnavailable("primary")}, - }} - fallbackStream := &stubStream{steps: []streamStep{ - {event: domain.StreamEvent{Type: domain.StreamEventOutputTextDelta, ContentDelta: &content}}, - {event: domain.StreamEvent{Type: domain.StreamEventCompleted, FinishReason: &finish}}, - }} - meter := &stubUsageMeter{} - primary := &stubProvider{stream: primaryStream} - fallback := &stubProvider{stream: fallbackStream} - svc := newFallbackGenerateService(meter, primary, fallback) - - stream, _, model, err := svc.ExecuteStream(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "model", - Stream: true, - }, "sk-test") - if err != nil { - t.Fatalf("ExecuteStream returned error: %v", err) - } - if model.ProviderConfig.ProviderName != "fallback" || !primaryStream.closed { - t.Fatalf("selected provider = %q, primary closed = %v", model.ProviderConfig.ProviderName, primaryStream.closed) - } - event, err := stream.Recv() - if err != nil || event.ContentDelta == nil || *event.ContentDelta != content { - t.Fatalf("first event = %#v, err = %v; want fallback content", event, err) - } - for { - if _, err := stream.Recv(); err != nil { - if err != io.EOF { - t.Fatalf("drain stream: %v", err) - } - break - } - } - if primary.streamCalls != 1 || fallback.streamCalls != 1 || meter.successCalls != 1 { - t.Fatalf("calls = primary %d, fallback %d, success %d", primary.streamCalls, fallback.streamCalls, meter.successCalls) - } -} - -func TestGenerateServiceExecuteStream_DoesNotFallbackAfterVisibleOutput(t *testing.T) { - content := "primary answer" - primary := &stubProvider{stream: &stubStream{steps: []streamStep{ - {event: domain.StreamEvent{Type: domain.StreamEventOutputTextDelta, ContentDelta: &content}}, - {err: domain.ErrProviderUnavailable("primary")}, - }}} - fallback := &stubProvider{stream: &stubStream{}} - meter := &stubUsageMeter{} - svc := newFallbackGenerateService(meter, primary, fallback) - - stream, _, _, err := svc.ExecuteStream(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "model", - Stream: true, - }, "sk-test") - if err != nil { - t.Fatalf("ExecuteStream returned error: %v", err) - } - if _, err := stream.Recv(); err != nil { - t.Fatalf("first Recv: %v", err) - } - if _, err := stream.Recv(); err == nil { - t.Fatal("expected primary stream error") - } - if fallback.streamCalls != 0 || meter.failureCalls != 1 { - t.Fatalf("fallback calls = %d, failure calls = %d; want 0 and 1", fallback.streamCalls, meter.failureCalls) + if fallback.calls != 0 { + t.Fatalf("fallback calls = %d, want 0", fallback.calls) } } @@ -376,7 +316,7 @@ func newFallbackGenerateService(meter *stubUsageMeter, primary, fallback *stubPr MaxOutputTokens: 100, }}, nil, - &stubProviderRegistry{providers: map[string]ports.GenerationProvider{ + &stubProviderRegistry{providers: map[string]ports.Provider{ "primary": primary, "fallback": fallback, }}, @@ -386,35 +326,6 @@ func newFallbackGenerateService(meter *stubUsageMeter, primary, fallback *stubPr ) } -func (s *stubProvider) StreamGenerate(_ context.Context, _ domain.GenerateRequest, _ domain.PublicModel) (ports.GenerationStream, error) { - s.streamCalls++ - return s.stream, s.streamErr -} - -type streamStep struct { - event domain.StreamEvent - err error -} - -type stubStream struct { - steps []streamStep - closed bool -} - -func (s *stubStream) Recv() (domain.StreamEvent, error) { - if len(s.steps) == 0 { - return domain.StreamEvent{}, io.EOF - } - step := s.steps[0] - s.steps = s.steps[1:] - return step.event, step.err -} - -func (s *stubStream) Close() error { - s.closed = true - return nil -} - type stubLogger struct{} func (stubLogger) Debug(string, ...any) {} @@ -422,7 +333,7 @@ func (stubLogger) Info(string, ...any) {} func (stubLogger) Warn(string, ...any) {} func (stubLogger) Error(string, ...any) {} -func TestGenerateServiceExecute_RejectsWhenBalanceIsEmpty(t *testing.T) { +func TestProxy_RejectsWhenBalanceIsEmpty(t *testing.T) { usage := &stubUsageMeter{reserveErr: domain.ErrInsufficientBalance()} svc := NewGenerateService( &stubAuthService{ @@ -448,7 +359,7 @@ func TestGenerateServiceExecute_RejectsWhenBalanceIsEmpty(t *testing.T) { stubLogger{}, ) - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ + _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{ PublicModelID: "openai/gpt-4.1-mini", }, "sk-test") if err == nil { @@ -466,7 +377,7 @@ func TestGenerateServiceExecute_RejectsWhenBalanceIsEmpty(t *testing.T) { } } -func TestGenerateServiceExecute_DoesNotRouteUnknownModel(t *testing.T) { +func TestProxy_DoesNotRouteUnknownModel(t *testing.T) { router := &stubRouterClient{decision: domain.RouteDecision{PublicModelID: "minimax/minimax-m2.7"}} svc := NewGenerateService( &stubAuthService{authCtx: domain.AuthContext{ @@ -481,7 +392,7 @@ func TestGenerateServiceExecute_DoesNotRouteUnknownModel(t *testing.T) { stubLogger{}, ) - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ + _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{ PublicModelID: "unknown/model", }, "sk-test") if err == nil { @@ -496,7 +407,7 @@ func TestGenerateServiceExecute_DoesNotRouteUnknownModel(t *testing.T) { } } -func TestGenerateServiceExecute_RoutesExplicitRouterToSelectedModel(t *testing.T) { +func TestProxy_RoutesExplicitRouterToSelectedModel(t *testing.T) { usage := &stubUsageMeter{} category := "long-context" reason := "embedding_matched" @@ -545,17 +456,17 @@ func TestGenerateServiceExecute_RoutesExplicitRouterToSelectedModel(t *testing.T stubLogger{}, ) - result, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{ PublicModelID: "dappnode/router", }, "sk-test") if err != nil { - t.Fatalf("Execute returned error: %v", err) + t.Fatalf("Proxy returned error: %v", err) } if router.calls != 1 { t.Fatalf("router calls = %d, want 1", router.calls) } - if result.PublicModelID != "minimax/minimax-m2.7" { - t.Fatalf("result model = %q, want routed model", result.PublicModelID) + if call.Model.PublicModelID != "minimax/minimax-m2.7" || !strings.Contains(string(call.Body), `"model":"minimax/minimax-m2.7"`) { + t.Fatalf("result model = %q, want routed model", call.Model.PublicModelID) } if usage.lastSuccessReq.RequestedModelID != "dappnode/router" { t.Fatalf("requested model = %q, want router id", usage.lastSuccessReq.RequestedModelID) @@ -592,7 +503,7 @@ func TestGenerateServiceExecute_RoutesExplicitRouterToSelectedModel(t *testing.T } } -func TestGenerateServiceExecute_ReturnsRouterErrorForExplicitRouterOutage(t *testing.T) { +func TestProxy_ReturnsRouterErrorForExplicitRouterOutage(t *testing.T) { router := &stubRouterClient{err: domain.ErrInternal("an internal error occurred").WithMeta("dependency", "router")} svc := NewGenerateService( &stubAuthService{authCtx: domain.AuthContext{ @@ -612,7 +523,7 @@ func TestGenerateServiceExecute_ReturnsRouterErrorForExplicitRouterOutage(t *tes stubLogger{}, ) - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ + _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{ PublicModelID: "dappnode/router", }, "sk-test") if err == nil { diff --git a/apps/gateway/internal/application/services/pii.go b/apps/gateway/internal/application/services/pii.go new file mode 100644 index 0000000..8f51eb8 --- /dev/null +++ b/apps/gateway/internal/application/services/pii.go @@ -0,0 +1,698 @@ +package services + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "sort" + "strconv" + "strings" + + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" +) + +// PII masking works on the raw OpenAI JSON. Requests have the text of their +// messages masked before they go upstream; responses have the placeholders +// restored as they pass. Only the LLM-bound text changes; every other field +// is forwarded as sent. + +// maskBody masks PII in a chat request body: message content (tool results +// as JSON, masking only string values), reasoning sent back, tool-call +// arguments (as JSON when valid), and the user field. Tool definitions and +// other static fields are not scanned. It returns the body unchanged and no +// mapping when the key doesn't mask, the filter is off, or nothing was found. +// On filter failure it fails closed, unless configured to fail open. +func (s *GenerateService) maskBody(ctx context.Context, raw []byte, piiMode string) ([]byte, *domain.PIIMapping, error) { + mode, ok := domain.NormalizeAPIKeyPIIMode(piiMode) + if !ok || mode == domain.APIKeyPIIModeOff || s.pii == nil || !s.pii.Enabled() { + return raw, nil, nil + } + body, err := decodeObject(raw) + if err != nil { + return raw, nil, nil // Invalid bodies were rejected before this. + } + masker := &requestPIIMasker{ + service: s, + ctx: ctx, + piiMode: mode, + mapping: domain.NewPIIMapping(), + cache: make(map[string]string), + entityCounts: make(map[string]int), + surfaceCounts: make(map[string]int), + } + if err := masker.maskBody(body); err != nil { + s.logger.Warn("pii filter error", "error", err, "fail_open", s.piiFailOpen) + if s.piiFailOpen { + return raw, nil, nil + } + return nil, nil, domain.ErrInternal("an internal error occurred") + } + if masker.mapping.Len() == 0 { + return raw, nil, nil + } + s.logger.Debug("pii masked", + "pii_mode", mode, + "tokens", masker.mapping.Len(), + "entity_counts", masker.entityCounts, + "surface_counts", masker.surfaceCounts, + ) + masked, err := marshalNoEscape(body) + if err != nil { + return nil, nil, domain.ErrInternal("an internal error occurred") + } + return masked, masker.mapping, nil +} + +func (m *requestPIIMasker) maskBody(body map[string]any) error { + if user, ok := body["user"].(string); ok { + masked, err := m.maskText("user", user) + if err != nil { + return err + } + body["user"] = masked + } + messages, _ := body["messages"].([]any) + for _, item := range messages { + msg, ok := item.(map[string]any) + if !ok { + continue + } + surface, mask := "message_content", m.maskText + if role, _ := msg["role"].(string); role == "tool" { + surface, mask = "tool_result_content", m.maskJSONAwareString + } + switch content := msg["content"].(type) { + case string: + masked, err := mask(surface, content) + if err != nil { + return err + } + msg["content"] = masked + case []any: + for _, p := range content { + part, ok := p.(map[string]any) + if !ok { + continue + } + if text, ok := part["text"].(string); ok { + masked, err := mask(surface, text) + if err != nil { + return err + } + part["text"] = masked + } + } + } + for _, field := range []string{"reasoning_content", "reasoning"} { + if text, ok := msg[field].(string); ok { + masked, err := m.maskText("assistant_reasoning_content", text) + if err != nil { + return err + } + msg[field] = masked + } + } + calls, _ := msg["tool_calls"].([]any) + for _, c := range calls { + call, _ := c.(map[string]any) + fn, _ := call["function"].(map[string]any) + if args, ok := fn["arguments"].(string); ok { + masked, err := m.maskJSONAwareString("assistant_tool_call_arguments", args) + if err != nil { + return err + } + fn["arguments"] = masked + } + } + } + return nil +} + +// decodeObject decodes a JSON object, keeping numbers exactly as sent. +func decodeObject(raw []byte) (map[string]any, error) { + dec := json.NewDecoder(bytes.NewReader(raw)) + dec.UseNumber() + var obj map[string]any + if err := dec.Decode(&obj); err != nil || obj == nil { + return nil, errors.New("not a JSON object") + } + return obj, nil +} + +func marshalNoEscape(value any) ([]byte, error) { + var b bytes.Buffer + enc := json.NewEncoder(&b) + enc.SetEscapeHTML(false) + if err := enc.Encode(value); err != nil { + return nil, err + } + return bytes.TrimSuffix(b.Bytes(), []byte("\n")), nil +} + +// generated names the fields of a message the model writes, with the surface +// they're reported under. +var generated = []struct{ field, surface string }{ + {"content", "assistant_content"}, + {"reasoning_content", "assistant_reasoning_content"}, + {"reasoning", "assistant_reasoning_content"}, +} + +const toolArgsSurface = "assistant_tool_call_arguments" + +// unmasker restores PII placeholders in a proxied response. Streamed text is +// held back while it ends in what may be the start of a placeholder, so one +// split across events is still restored; the rest is released with the +// choice's finish reason, or in a final event. +type unmasker struct { + mapping *domain.PIIMapping + held map[string]*heldChoice // By choice index. + restored map[string]*strings.Builder + logger ports.Logger + fields []any +} + +type heldChoice struct { + text map[string]*strings.Builder // By field. + tools map[string]*strings.Builder // Arguments, by tool-call index. +} + +// newUnmasker returns nil when nothing was masked. +func newUnmasker(mapping *domain.PIIMapping, logger ports.Logger, fields ...any) *unmasker { + if mapping == nil || mapping.Len() == 0 { + return nil + } + return &unmasker{mapping: mapping, held: map[string]*heldChoice{}, restored: map[string]*strings.Builder{}, logger: logger, fields: fields} +} + +func (u *unmasker) choice(index string) *heldChoice { + h, ok := u.held[index] + if !ok { + h = &heldChoice{text: map[string]*strings.Builder{}, tools: map[string]*strings.Builder{}} + u.held[index] = h + } + return h +} + +func buffer(m map[string]*strings.Builder, key string) *strings.Builder { + b, ok := m[key] + if !ok { + b = &strings.Builder{} + m[key] = b + } + return b +} + +func (u *unmasker) track(surface, text string) { + buffer(u.restored, surface).WriteString(text) +} + +// event restores one streamed JSON event. Events without generated text +// pass unchanged. +func (u *unmasker) event(data []byte) []byte { + obj, err := decodeObject(data) + if err != nil { + return data + } + choices, _ := obj["choices"].([]any) + changed := false + for pos, item := range choices { + choice, ok := item.(map[string]any) + if !ok { + continue + } + index := choiceIndex(choice["index"], pos) + held := u.choice(index) + delta, _ := choice["delta"].(map[string]any) + for _, g := range generated { + if text, ok := delta[g.field].(string); ok && text != "" { + out := feedPIIStreamBuffer(buffer(held.text, g.field), text, u.mapping) + u.track(g.surface, out) + delta[g.field] = out + changed = true + } + } + for _, call := range toolCallObjects(delta) { + fn, _ := call["function"].(map[string]any) + if args, ok := fn["arguments"].(string); ok && args != "" { + out := feedPIIStreamBuffer(buffer(held.tools, choiceIndex(call["index"], 0)), args, u.mapping) + u.track(toolArgsSurface, out) + fn["arguments"] = out + changed = true + } + } + if choice["finish_reason"] != nil && u.release(index, choice) { + changed = true + } + } + if !changed { + return data + } + out, err := marshalNoEscape(obj) + if err != nil { + return data + } + return out +} + +// release adds what a choice still holds to its delta. +func (u *unmasker) release(index string, choice map[string]any) bool { + held, ok := u.held[index] + if !ok { + return false + } + delete(u.held, index) + delta, _ := choice["delta"].(map[string]any) + if delta == nil { + delta = map[string]any{} + } + released := false + for _, g := range generated { + if b, ok := held.text[g.field]; ok { + if tail := flushPIIStreamBuffer(b, u.mapping); tail != "" { + existing, _ := delta[g.field].(string) + delta[g.field] = existing + tail + u.track(g.surface, tail) + released = true + } + } + } + keys := make([]string, 0, len(held.tools)) + for key := range held.tools { + keys = append(keys, key) + } + sort.Strings(keys) + for _, key := range keys { + tail := flushPIIStreamBuffer(held.tools[key], u.mapping) + if tail == "" { + continue + } + u.track(toolArgsSurface, tail) + released = true + calls, _ := delta["tool_calls"].([]any) + appended := false + for _, call := range toolCallObjects(delta) { + if choiceIndex(call["index"], 0) == key { + fn, _ := call["function"].(map[string]any) + if fn == nil { + fn = map[string]any{} + call["function"] = fn + } + existing, _ := fn["arguments"].(string) + fn["arguments"] = existing + tail + appended = true + } + } + if !appended { + delta["tool_calls"] = append(calls, map[string]any{"index": json.Number(key), "function": map[string]any{"arguments": tail}}) + } + } + if released { + choice["delta"] = delta + } + return released +} + +// tail returns an event with the text still held for choices the provider +// never finished, or nil. +func (u *unmasker) tail(id, model string) []byte { + indexes := make([]string, 0, len(u.held)) + for index := range u.held { + indexes = append(indexes, index) + } + sort.Strings(indexes) + var choices []any + for _, index := range indexes { + choice := map[string]any{"index": json.Number(index), "delta": map[string]any{}, "finish_reason": nil} + if u.release(index, choice) { + choices = append(choices, choice) + } + } + if len(choices) == 0 { + return nil + } + out, err := marshalNoEscape(map[string]any{"id": id, "object": "chat.completion.chunk", "model": model, "choices": choices}) + if err != nil { + return nil + } + return out +} + +// body restores a non-streaming response. +func (u *unmasker) body(data []byte) []byte { + obj, err := decodeObject(data) + if err != nil { + return data + } + choices, _ := obj["choices"].([]any) + for _, item := range choices { + choice, _ := item.(map[string]any) + msg, _ := choice["message"].(map[string]any) + for _, g := range generated { + if text, ok := msg[g.field].(string); ok && text != "" { + msg[g.field] = domain.Unmask(text, u.mapping) + u.track(g.surface, msg[g.field].(string)) + } + } + for _, call := range toolCallObjects(msg) { + fn, _ := call["function"].(map[string]any) + if args, ok := fn["arguments"].(string); ok && args != "" { + fn["arguments"] = domain.Unmask(args, u.mapping) + u.track(toolArgsSurface, fn["arguments"].(string)) + } + } + } + out, err := marshalNoEscape(obj) + if err != nil { + return data + } + return out +} + +// report logs placeholders the model changed so they could not be restored. +func (u *unmasker) report(stream bool) { + groups := map[string][]string{} + for surface, text := range u.restored { + trackUnresolvedPIITokens(groups, u.mapping, surface, text.String()) + } + logUnresolvedPIITokenGroups(u.logger, groups, append(u.fields, "stream", stream)...) +} + +func toolCallObjects(parent map[string]any) []map[string]any { + calls, _ := parent["tool_calls"].([]any) + out := make([]map[string]any, 0, len(calls)) + for _, c := range calls { + if call, ok := c.(map[string]any); ok { + out = append(out, call) + } + } + return out +} + +// choiceIndex reads a JSON index, or uses the position when there is none. +func choiceIndex(v any, pos int) string { + if n, ok := v.(json.Number); ok { + return n.String() + } + return strconv.Itoa(pos) +} + +type requestPIIMasker struct { + service *GenerateService + ctx context.Context + piiMode string + mapping *domain.PIIMapping + cache map[string]string + entityCounts map[string]int + surfaceCounts map[string]int +} + +func (m *requestPIIMasker) maskText(surface, text string) (string, error) { + if text == "" { + return text, nil + } + if masked, ok := m.cache[text]; ok { + return masked, nil + } + entities, err := m.service.pii.Analyze(m.ctx, text, ports.PIIAnalyzeOptions{ + Language: m.service.piiLang, + Mode: m.piiMode, + }) + if err != nil { + return "", err + } + m.trackEntities(surface, entities) + masked := domain.ApplyMask(text, entities, m.mapping) + m.cache[text] = masked + return masked, nil +} + +func (m *requestPIIMasker) trackEntities(surface string, entities []domain.PIIEntity) { + for _, entity := range entities { + m.entityCounts[entity.Type]++ + if surface != "" { + m.surfaceCounts[surface]++ + } + } +} + +func (m *requestPIIMasker) maskJSONAwareString(surface, raw string) (string, error) { + if raw == "" { + return raw, nil + } + + var value any + dec := json.NewDecoder(strings.NewReader(raw)) + dec.UseNumber() + if err := dec.Decode(&value); err != nil { + return m.maskText(surface, raw) + } + if dec.Decode(&struct{}{}) != io.EOF { + return m.maskText(surface, raw) + } + + masked, changed, err := m.maskAnyWithChanged(surface, value) + if err != nil { + return "", err + } + if !changed { + return raw, nil + } + return marshalJSONNoEscape(masked) +} + +func (m *requestPIIMasker) maskAnyWithChanged(surface string, value any) (any, bool, error) { + switch v := value.(type) { + case string: + masked, err := m.maskText(surface, v) + return masked, masked != v, err + case []any: + out := make([]any, len(v)) + changed := false + for i := range v { + masked, itemChanged, err := m.maskAnyWithChanged(surface, v[i]) + if err != nil { + return nil, false, err + } + out[i] = masked + changed = changed || itemChanged + } + return out, changed, nil + case map[string]any: + out := make(map[string]any, len(v)) + changed := false + for key, val := range v { + masked, itemChanged, err := m.maskAnyWithChanged(surface, val) + if err != nil { + return nil, false, err + } + out[key] = masked + changed = changed || itemChanged + } + return out, changed, nil + default: + return value, false, nil + } +} + +func marshalJSONNoEscape(value any) (string, error) { + var b strings.Builder + enc := json.NewEncoder(&b) + enc.SetEscapeHTML(false) + if err := enc.Encode(value); err != nil { + return "", err + } + return strings.TrimSuffix(b.String(), "\n"), nil +} + +// feedPIIStreamBuffer appends `delta` to the internal buffer and returns the largest prefix +// that contains only complete (or no) placeholder tokens, with those tokens +// already replaced by their original values. +// +// The buffer holds back from the last unmatched '[' onward so we never split +// a token across two emitted chunks. +func feedPIIStreamBuffer(buf *strings.Builder, delta string, mapping *domain.PIIMapping) string { + buf.WriteString(delta) + full := buf.String() + + // Find the last '[' that has no matching ']' after it. Everything up to + // that index is safe to emit; everything from it onward stays buffered. + cut := len(full) + for i := len(full) - 1; i >= 0; i-- { + if full[i] == '[' { + if !strings.ContainsRune(full[i:], ']') { + cut = i + break + } + // Has a closing bracket — entire string is safe. + break + } + if full[i] == ']' { + // We hit a closing bracket before any opener — safe. + break + } + } + if bareCut := bareAliasHoldStart(full, mapping); bareCut >= 0 && bareCut < cut { + cut = bareCut + } + + safe := full[:cut] + tail := full[cut:] + + buf.Reset() + buf.WriteString(tail) + + return domain.Unmask(safe, mapping) +} + +func bareAliasHoldStart(text string, mapping *domain.PIIMapping) int { + if mapping == nil || mapping.Len() == 0 || text == "" { + return -1 + } + aliases := mapping.BareTokenAliases() + if len(aliases) == 0 { + return -1 + } + maxLen := 0 + for _, alias := range aliases { + if len(alias) > maxLen { + maxLen = len(alias) + } + } + startAt := len(text) - maxLen + if startAt < 0 { + startAt = 0 + } + for start := len(text) - 1; start >= startAt; start-- { + if !isBareTokenBoundary(text, start-1) { + continue + } + suffix := text[start:] + for _, alias := range aliases { + if len(suffix) <= len(alias) && strings.EqualFold(alias[:len(suffix)], suffix) { + return start + } + } + } + return -1 +} + +func isBareTokenBoundary(text string, index int) bool { + if index < 0 || index >= len(text) { + return true + } + c := text[index] + return !((c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '_') +} + +// flushPIIStreamBuffer returns any buffered text, unmasking whatever placeholders are +// complete and leaving partial ones intact. +func flushPIIStreamBuffer(buf *strings.Builder, mapping *domain.PIIMapping) string { + if buf.Len() == 0 { + return "" + } + tail := buf.String() + buf.Reset() + return domain.Unmask(tail, mapping) +} + +func sanitizeErrorWithPIIMapping(err error, mapping *domain.PIIMapping) error { + if err == nil || mapping == nil || mapping.Len() == 0 { + return err + } + var gwErr *domain.GatewayError + if !errors.As(err, &gwErr) { + return errors.New(mapping.MaskKnownOriginals(err.Error())) + } + cp := *gwErr + cp.Message = mapping.MaskKnownOriginals(cp.Message) + if len(gwErr.Metadata) > 0 { + cp.Metadata = make(map[string]any, len(gwErr.Metadata)) + for key, value := range gwErr.Metadata { + cp.Metadata[key] = sanitizePIILogValue(value, mapping) + } + } + return &cp +} + +func trackUnresolvedPIITokens(groups map[string][]string, mapping *domain.PIIMapping, surface, text string) { + if groups == nil || mapping == nil || mapping.Len() == 0 || text == "" { + return + } + tokens := mapping.UnresolvedTokens(text) + if len(tokens) == 0 { + return + } + groups[surface] = mergeStringSets(groups[surface], tokens) +} + +func logUnresolvedPIITokenGroups(logger ports.Logger, groups map[string][]string, fields ...any) { + if logger == nil || len(groups) == 0 { + return + } + surfaces := make([]string, 0, len(groups)) + for surface := range groups { + surfaces = append(surfaces, surface) + } + sort.Strings(surfaces) + for _, surface := range surfaces { + tokens := groups[surface] + sort.Strings(tokens) + logFields := append([]any(nil), fields...) + logFields = append(logFields, + "surface", surface, + "unresolved_token_count", len(tokens), + "unresolved_tokens", tokens, + ) + logger.Warn("pii restoration unresolved tokens", logFields...) + } +} + +func mergeStringSets(existing []string, incoming []string) []string { + seen := make(map[string]struct{}, len(existing)+len(incoming)) + out := make([]string, 0, len(existing)+len(incoming)) + for _, token := range existing { + if _, ok := seen[token]; ok { + continue + } + seen[token] = struct{}{} + out = append(out, token) + } + for _, token := range incoming { + if _, ok := seen[token]; ok { + continue + } + seen[token] = struct{}{} + out = append(out, token) + } + return out +} + +func sanitizePIILogValue(value any, mapping *domain.PIIMapping) any { + switch v := value.(type) { + case string: + return mapping.MaskKnownOriginals(v) + case []any: + out := make([]any, len(v)) + for i := range v { + out[i] = sanitizePIILogValue(v[i], mapping) + } + return out + case []string: + out := make([]string, len(v)) + for i := range v { + out[i] = mapping.MaskKnownOriginals(v[i]) + } + return out + case map[string]any: + out := make(map[string]any, len(v)) + for key, item := range v { + out[key] = sanitizePIILogValue(item, mapping) + } + return out + default: + return value + } +} diff --git a/apps/gateway/internal/application/services/pii_masking.go b/apps/gateway/internal/application/services/pii_masking.go deleted file mode 100644 index 45c0d62..0000000 --- a/apps/gateway/internal/application/services/pii_masking.go +++ /dev/null @@ -1,798 +0,0 @@ -package services - -import ( - "context" - "encoding/json" - "errors" - "io" - "sort" - "strings" - - "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" -) - -// maskRequest detects PII in the dynamic LLM-bound fields of `req` and -// rewrites those fields in place with deterministic placeholder tokens. -// It covers instructions, input content, assistant reasoning content, -// assistant tool-call arguments, and the OpenAI-compatible user field. -// The returned mapping must be passed to unmaskResult / piiUnmaskingStream so -// the upstream model's response can be restored to its original form. -// -// Returns (nil, nil) when the filter is disabled or there is nothing to mask. -// On filter failure, behavior depends on s.piiFailOpen: fail-closed (default) -// returns an error; fail-open logs and returns (nil, nil) so the request -// proceeds with the original text. -func (s *GenerateService) maskRequest(ctx context.Context, req *domain.GenerateRequest, piiMode string) (*domain.PIIMapping, error) { - mode, ok := domain.NormalizeAPIKeyPIIMode(piiMode) - if !ok || mode == domain.APIKeyPIIModeOff { - return nil, nil - } - if s.pii == nil || !s.pii.Enabled() { - return nil, nil - } - - original := cloneGenerateRequest(*req) - masking := &requestPIIMasker{ - service: s, - ctx: ctx, - piiMode: mode, - mapping: domain.NewPIIMapping(), - cache: make(map[string]string), - entityCounts: make(map[string]int), - surfaceCounts: make(map[string]int), - } - - maskPtr := func(surface string, ptr **string) error { - if ptr == nil || *ptr == nil || **ptr == "" { - return nil - } - masked, err := masking.maskText(surface, **ptr) - if err != nil { - return err - } - **ptr = masked - return nil - } - - maskJSONPtr := func(surface string, ptr *string) (string, error) { - if ptr == nil || *ptr == "" { - return "", nil - } - return masking.maskJSONAwareString(surface, *ptr) - } - - if err := maskPtr("instructions", &req.Instructions); err != nil { - return s.handlePIIMaskError(req, original, err) - } - if err := maskPtr("user", &req.User); err != nil { - return s.handlePIIMaskError(req, original, err) - } - for i := range req.Input { - item := &req.Input[i] - if item.Content != nil && *item.Content != "" { - var ( - masked string - err error - ) - if item.Role != nil && *item.Role == "tool" { - masked, err = masking.maskJSONAwareString("tool_result_content", *item.Content) - } else { - masked, err = masking.maskText("message_content", *item.Content) - } - if err != nil { - return s.handlePIIMaskError(req, original, err) - } - *item.Content = masked - } - if err := maskPtr("assistant_reasoning_content", &item.ReasoningContent); err != nil { - return s.handlePIIMaskError(req, original, err) - } - for j := range item.ToolCalls { - masked, err := maskJSONPtr("assistant_tool_call_arguments", &item.ToolCalls[j].ArgumentsJSON) - if err != nil { - return s.handlePIIMaskError(req, original, err) - } - if masked != "" { - item.ToolCalls[j].ArgumentsJSON = masked - } - } - } - - if masking.mapping.Len() == 0 { - return nil, nil - } - s.logger.Debug("pii masked", - "pii_mode", mode, - "tokens", masking.mapping.Len(), - "entity_counts", masking.entityCounts, - "surface_counts", masking.surfaceCounts, - ) - return masking.mapping, nil -} - -func (s *GenerateService) handlePIIMaskError(req *domain.GenerateRequest, original domain.GenerateRequest, err error) (*domain.PIIMapping, error) { - s.logger.Warn("pii filter error", "error", err, "fail_open", s.piiFailOpen) - if s.piiFailOpen { - *req = original - return nil, nil - } - return nil, domain.ErrInternal("an internal error occurred") -} - -type requestPIIMasker struct { - service *GenerateService - ctx context.Context - piiMode string - mapping *domain.PIIMapping - cache map[string]string - entityCounts map[string]int - surfaceCounts map[string]int -} - -func (m *requestPIIMasker) maskText(surface, text string) (string, error) { - if text == "" { - return text, nil - } - if masked, ok := m.cache[text]; ok { - return masked, nil - } - entities, err := m.service.pii.Analyze(m.ctx, text, ports.PIIAnalyzeOptions{ - Language: m.service.piiLang, - Mode: m.piiMode, - }) - if err != nil { - return "", err - } - m.trackEntities(surface, entities) - masked := domain.ApplyMask(text, entities, m.mapping) - m.cache[text] = masked - return masked, nil -} - -func (m *requestPIIMasker) trackEntities(surface string, entities []domain.PIIEntity) { - for _, entity := range entities { - m.entityCounts[entity.Type]++ - if surface != "" { - m.surfaceCounts[surface]++ - } - } -} - -func (m *requestPIIMasker) maskJSONAwareString(surface, raw string) (string, error) { - if raw == "" { - return raw, nil - } - - var value any - dec := json.NewDecoder(strings.NewReader(raw)) - dec.UseNumber() - if err := dec.Decode(&value); err != nil { - return m.maskText(surface, raw) - } - if dec.Decode(&struct{}{}) != io.EOF { - return m.maskText(surface, raw) - } - - masked, changed, err := m.maskAnyWithChanged(surface, value) - if err != nil { - return "", err - } - if !changed { - return raw, nil - } - return marshalJSONNoEscape(masked) -} - -func (m *requestPIIMasker) maskAnyWithChanged(surface string, value any) (any, bool, error) { - switch v := value.(type) { - case string: - masked, err := m.maskText(surface, v) - return masked, masked != v, err - case []any: - out := make([]any, len(v)) - changed := false - for i := range v { - masked, itemChanged, err := m.maskAnyWithChanged(surface, v[i]) - if err != nil { - return nil, false, err - } - out[i] = masked - changed = changed || itemChanged - } - return out, changed, nil - case map[string]any: - out := make(map[string]any, len(v)) - changed := false - for key, val := range v { - masked, itemChanged, err := m.maskAnyWithChanged(surface, val) - if err != nil { - return nil, false, err - } - out[key] = masked - changed = changed || itemChanged - } - return out, changed, nil - default: - return value, false, nil - } -} - -func marshalJSONNoEscape(value any) (string, error) { - var b strings.Builder - enc := json.NewEncoder(&b) - enc.SetEscapeHTML(false) - if err := enc.Encode(value); err != nil { - return "", err - } - return strings.TrimSuffix(b.String(), "\n"), nil -} - -// unmaskResult walks every output text and tool-call argument field of a -// non-streaming generation result and restores PII using `mapping`. Safe to -// call with a nil mapping. -func unmaskResult(result *domain.GenerateResult, mapping *domain.PIIMapping, logger ports.Logger, logFields ...any) { - if result == nil || mapping == nil || mapping.Len() == 0 { - return - } - unresolved := make(map[string][]string) - for i := range result.Output { - item := &result.Output[i] - if item.Content != nil && *item.Content != "" { - restored := domain.Unmask(*item.Content, mapping) - trackUnresolvedPIITokens(unresolved, mapping, "assistant_content", restored) - item.Content = &restored - } - if item.ReasoningContent != nil && *item.ReasoningContent != "" { - restored := domain.Unmask(*item.ReasoningContent, mapping) - trackUnresolvedPIITokens(unresolved, mapping, "assistant_reasoning_content", restored) - item.ReasoningContent = &restored - } - for j := range item.ToolCalls { - if item.ToolCalls[j].ArgumentsJSON != "" { - item.ToolCalls[j].ArgumentsJSON = domain.Unmask(item.ToolCalls[j].ArgumentsJSON, mapping) - trackUnresolvedPIITokens(unresolved, mapping, "assistant_tool_call_arguments", item.ToolCalls[j].ArgumentsJSON) - } - } - } - logUnresolvedPIITokenGroups(logger, unresolved, append(logFields, "stream", false)...) -} - -// piiUnmaskingStream wraps a GenerationStream and restores PII tokens in -// output content, reasoning, and tool-call argument deltas before they reach -// the client. It buffers partial tokens that span chunk boundaries -// (e.g. "[PER" + "SON_1]") so substitution always sees a complete token. -type piiUnmaskingStream struct { - inner ports.GenerationStream - mapping *domain.PIIMapping - logger ports.Logger - logFields []any - textBuf strings.Builder - reasonBuf strings.Builder - toolArgBufs map[int]*strings.Builder - restoredText strings.Builder - restoredReason strings.Builder - restoredToolArgs map[int]*strings.Builder - pending []domain.StreamEvent - eofPending bool - reported bool - closed bool -} - -func newPIIUnmaskingStream(inner ports.GenerationStream, mapping *domain.PIIMapping) *piiUnmaskingStream { - return newPIIUnmaskingStreamWithLogger(inner, mapping, nil) -} - -func newPIIUnmaskingStreamWithLogger(inner ports.GenerationStream, mapping *domain.PIIMapping, logger ports.Logger, logFields ...any) *piiUnmaskingStream { - return &piiUnmaskingStream{ - inner: inner, - mapping: mapping, - logger: logger, - logFields: append([]any(nil), logFields...), - toolArgBufs: make(map[int]*strings.Builder), - restoredToolArgs: make(map[int]*strings.Builder), - } -} - -// Recv pulls the next event from the wrapped stream and rewrites any text -// delta. Non-text events pass through untouched. -func (s *piiUnmaskingStream) Recv() (domain.StreamEvent, error) { - if len(s.pending) > 0 { - event := s.pending[0] - s.pending = s.pending[1:] - return event, nil - } - if s.eofPending { - return domain.StreamEvent{}, io.EOF - } - - event, err := s.inner.Recv() - - // When the upstream stream finishes, flush any text we held back so the - // client doesn't lose a trailing fragment. The tail is returned with a nil - // error first because the HTTP handler ignores events returned with EOF. - if err == io.EOF { - if tails := s.flushBufferedEvents(); len(tails) > 0 { - s.reportUnresolvedToolArgs() - s.eofPending = true - s.pending = append(s.pending, tails[1:]...) - return tails[0], nil - } - s.reportUnresolvedToolArgs() - return event, err - } - - if err != nil { - return event, err - } - - if event.Type == domain.StreamEventCompleted { - return s.rewriteCompleted(event), nil - } - - if event.Type == domain.StreamEventError && event.Error != nil { - if sanitized, ok := sanitizeErrorWithPIIMapping(event.Error, s.mapping).(*domain.GatewayError); ok { - event.Error = sanitized - } - return event, nil - } - - if event.ContentDelta != nil && *event.ContentDelta != "" { - emit := feedPIIStreamBuffer(&s.textBuf, *event.ContentDelta, s.mapping) - s.trackRestoredText(emit) - event.ContentDelta = &emit - } - if event.ReasoningDelta != nil && *event.ReasoningDelta != "" { - emit := feedPIIStreamBuffer(&s.reasonBuf, *event.ReasoningDelta, s.mapping) - s.trackRestoredReasoning(emit) - event.ReasoningDelta = &emit - } - if event.ToolCallDelta != nil && event.ToolCallDelta.ArgumentsDelta != nil && *event.ToolCallDelta.ArgumentsDelta != "" { - buf := s.toolArgBuffer(event.ToolCallDelta.Index) - emit := feedPIIStreamBuffer(buf, *event.ToolCallDelta.ArgumentsDelta, s.mapping) - s.trackRestoredToolArg(event.ToolCallDelta.Index, emit) - event.ToolCallDelta.ArgumentsDelta = &emit - } - return event, nil -} - -func (s *piiUnmaskingStream) rewriteCompleted(event domain.StreamEvent) domain.StreamEvent { - if event.ContentDelta != nil && *event.ContentDelta != "" { - content := feedPIIStreamBuffer(&s.textBuf, *event.ContentDelta, s.mapping) - s.trackRestoredText(content) - event.ContentDelta = &content - } - if event.ReasoningDelta != nil && *event.ReasoningDelta != "" { - reasoning := feedPIIStreamBuffer(&s.reasonBuf, *event.ReasoningDelta, s.mapping) - s.trackRestoredReasoning(reasoning) - event.ReasoningDelta = &reasoning - } - if event.ToolCallDelta != nil && event.ToolCallDelta.ArgumentsDelta != nil && *event.ToolCallDelta.ArgumentsDelta != "" { - buf := s.toolArgBuffer(event.ToolCallDelta.Index) - args := feedPIIStreamBuffer(buf, *event.ToolCallDelta.ArgumentsDelta, s.mapping) - s.trackRestoredToolArg(event.ToolCallDelta.Index, args) - event.ToolCallDelta.ArgumentsDelta = &args - } - - prefixes := completedDeltaPrefixEvents(event) - tails := s.flushBufferedEvents() - s.reportUnresolvedToolArgs() - if len(prefixes) == 0 && len(tails) == 0 { - return event - } - - event.ContentDelta = nil - event.ReasoningDelta = nil - event.ToolCallDelta = nil - - events := make([]domain.StreamEvent, 0, len(prefixes)+len(tails)+1) - events = append(events, prefixes...) - events = append(events, tails...) - events = append(events, event) - s.pending = append(s.pending, events[1:]...) - return events[0] -} - -// Close releases the wrapped stream. Buffered text is discarded — Recv already -// returns the tail on EOF, and a Close before EOF means the client gave up. -func (s *piiUnmaskingStream) Close() error { - s.closed = true - s.textBuf.Reset() - s.reasonBuf.Reset() - for _, buf := range s.toolArgBufs { - buf.Reset() - } - return s.inner.Close() -} - -func (s *piiUnmaskingStream) toolArgBuffer(index int) *strings.Builder { - if buf, ok := s.toolArgBufs[index]; ok { - return buf - } - buf := &strings.Builder{} - s.toolArgBufs[index] = buf - return buf -} - -func (s *piiUnmaskingStream) flushBufferedEvents() []domain.StreamEvent { - events := make([]domain.StreamEvent, 0, 2+len(s.toolArgBufs)) - if tail := flushPIIStreamBuffer(&s.textBuf, s.mapping); tail != "" { - s.trackRestoredText(tail) - events = append(events, contentDeltaEvent(tail)) - } - if tail := flushPIIStreamBuffer(&s.reasonBuf, s.mapping); tail != "" { - s.trackRestoredReasoning(tail) - events = append(events, reasoningDeltaEvent(tail)) - } - indexes := make([]int, 0, len(s.toolArgBufs)) - for index := range s.toolArgBufs { - indexes = append(indexes, index) - } - sort.Ints(indexes) - for _, index := range indexes { - buf := s.toolArgBufs[index] - if tail := flushPIIStreamBuffer(buf, s.mapping); tail != "" { - s.trackRestoredToolArg(index, tail) - events = append(events, toolCallArgumentDeltaEvent(index, tail)) - } - } - return events -} - -func (s *piiUnmaskingStream) trackRestoredToolArg(index int, delta string) { - if delta == "" { - return - } - buf, ok := s.restoredToolArgs[index] - if !ok { - buf = &strings.Builder{} - s.restoredToolArgs[index] = buf - } - buf.WriteString(delta) -} - -func (s *piiUnmaskingStream) trackRestoredText(delta string) { - if delta != "" { - s.restoredText.WriteString(delta) - } -} - -func (s *piiUnmaskingStream) trackRestoredReasoning(delta string) { - if delta != "" { - s.restoredReason.WriteString(delta) - } -} - -func (s *piiUnmaskingStream) reportUnresolvedToolArgs() { - if s.reported { - return - } - s.reported = true - unresolved := make(map[string][]string) - trackUnresolvedPIITokens(unresolved, s.mapping, "assistant_content", s.restoredText.String()) - trackUnresolvedPIITokens(unresolved, s.mapping, "assistant_reasoning_content", s.restoredReason.String()) - for _, buf := range s.restoredToolArgs { - trackUnresolvedPIITokens(unresolved, s.mapping, "assistant_tool_call_arguments", buf.String()) - } - logUnresolvedPIITokenGroups(s.logger, unresolved, append(s.logFields, "stream", true)...) -} - -// feedPIIStreamBuffer appends `delta` to the internal buffer and returns the largest prefix -// that contains only complete (or no) placeholder tokens, with those tokens -// already replaced by their original values. -// -// The buffer holds back from the last unmatched '[' onward so we never split -// a token across two emitted chunks. -func feedPIIStreamBuffer(buf *strings.Builder, delta string, mapping *domain.PIIMapping) string { - buf.WriteString(delta) - full := buf.String() - - // Find the last '[' that has no matching ']' after it. Everything up to - // that index is safe to emit; everything from it onward stays buffered. - cut := len(full) - for i := len(full) - 1; i >= 0; i-- { - if full[i] == '[' { - if !strings.ContainsRune(full[i:], ']') { - cut = i - break - } - // Has a closing bracket — entire string is safe. - break - } - if full[i] == ']' { - // We hit a closing bracket before any opener — safe. - break - } - } - if bareCut := bareAliasHoldStart(full, mapping); bareCut >= 0 && bareCut < cut { - cut = bareCut - } - - safe := full[:cut] - tail := full[cut:] - - buf.Reset() - buf.WriteString(tail) - - return domain.Unmask(safe, mapping) -} - -func bareAliasHoldStart(text string, mapping *domain.PIIMapping) int { - if mapping == nil || mapping.Len() == 0 || text == "" { - return -1 - } - aliases := mapping.BareTokenAliases() - if len(aliases) == 0 { - return -1 - } - maxLen := 0 - for _, alias := range aliases { - if len(alias) > maxLen { - maxLen = len(alias) - } - } - startAt := len(text) - maxLen - if startAt < 0 { - startAt = 0 - } - for start := len(text) - 1; start >= startAt; start-- { - if !isBareTokenBoundary(text, start-1) { - continue - } - suffix := text[start:] - for _, alias := range aliases { - if len(suffix) <= len(alias) && strings.EqualFold(alias[:len(suffix)], suffix) { - return start - } - } - } - return -1 -} - -func isBareTokenBoundary(text string, index int) bool { - if index < 0 || index >= len(text) { - return true - } - c := text[index] - return !((c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '_') -} - -// flushPIIStreamBuffer returns any buffered text, unmasking whatever placeholders are -// complete and leaving partial ones intact. -func flushPIIStreamBuffer(buf *strings.Builder, mapping *domain.PIIMapping) string { - if buf.Len() == 0 { - return "" - } - tail := buf.String() - buf.Reset() - return domain.Unmask(tail, mapping) -} - -func completedDeltaPrefixEvents(event domain.StreamEvent) []domain.StreamEvent { - events := make([]domain.StreamEvent, 0, 2) - if (event.ContentDelta != nil && *event.ContentDelta != "") || (event.ReasoningDelta != nil && *event.ReasoningDelta != "") { - events = append(events, domain.StreamEvent{ - Type: domain.StreamEventOutputTextDelta, - ContentDelta: event.ContentDelta, - ReasoningDelta: event.ReasoningDelta, - }) - } - if event.ToolCallDelta != nil && event.ToolCallDelta.ArgumentsDelta != nil && *event.ToolCallDelta.ArgumentsDelta != "" { - events = append(events, domain.StreamEvent{ - Type: domain.StreamEventToolCallDelta, - ToolCallDelta: event.ToolCallDelta, - }) - } - return events -} - -func contentDeltaEvent(delta string) domain.StreamEvent { - return domain.StreamEvent{ - Type: domain.StreamEventOutputTextDelta, - ContentDelta: &delta, - } -} - -func reasoningDeltaEvent(delta string) domain.StreamEvent { - return domain.StreamEvent{ - Type: domain.StreamEventOutputTextDelta, - ReasoningDelta: &delta, - } -} - -func toolCallArgumentDeltaEvent(index int, delta string) domain.StreamEvent { - return domain.StreamEvent{ - Type: domain.StreamEventToolCallDelta, - ToolCallDelta: &domain.ToolCallDelta{ - Index: index, - ArgumentsDelta: &delta, - }, - } -} - -func sanitizeErrorWithPIIMapping(err error, mapping *domain.PIIMapping) error { - if err == nil || mapping == nil || mapping.Len() == 0 { - return err - } - var gwErr *domain.GatewayError - if !errors.As(err, &gwErr) { - return errors.New(mapping.MaskKnownOriginals(err.Error())) - } - cp := *gwErr - cp.Message = mapping.MaskKnownOriginals(cp.Message) - if len(gwErr.Metadata) > 0 { - cp.Metadata = make(map[string]any, len(gwErr.Metadata)) - for key, value := range gwErr.Metadata { - cp.Metadata[key] = sanitizePIILogValue(value, mapping) - } - } - return &cp -} - -func trackUnresolvedPIITokens(groups map[string][]string, mapping *domain.PIIMapping, surface, text string) { - if groups == nil || mapping == nil || mapping.Len() == 0 || text == "" { - return - } - tokens := mapping.UnresolvedTokens(text) - if len(tokens) == 0 { - return - } - groups[surface] = mergeStringSets(groups[surface], tokens) -} - -func logUnresolvedPIITokenGroups(logger ports.Logger, groups map[string][]string, fields ...any) { - if logger == nil || len(groups) == 0 { - return - } - surfaces := make([]string, 0, len(groups)) - for surface := range groups { - surfaces = append(surfaces, surface) - } - sort.Strings(surfaces) - for _, surface := range surfaces { - tokens := groups[surface] - sort.Strings(tokens) - logFields := append([]any(nil), fields...) - logFields = append(logFields, - "surface", surface, - "unresolved_token_count", len(tokens), - "unresolved_tokens", tokens, - ) - logger.Warn("pii restoration unresolved tokens", logFields...) - } -} - -func mergeStringSets(existing []string, incoming []string) []string { - seen := make(map[string]struct{}, len(existing)+len(incoming)) - out := make([]string, 0, len(existing)+len(incoming)) - for _, token := range existing { - if _, ok := seen[token]; ok { - continue - } - seen[token] = struct{}{} - out = append(out, token) - } - for _, token := range incoming { - if _, ok := seen[token]; ok { - continue - } - seen[token] = struct{}{} - out = append(out, token) - } - return out -} - -func sanitizePIILogValue(value any, mapping *domain.PIIMapping) any { - switch v := value.(type) { - case string: - return mapping.MaskKnownOriginals(v) - case []any: - out := make([]any, len(v)) - for i := range v { - out[i] = sanitizePIILogValue(v[i], mapping) - } - return out - case []string: - out := make([]string, len(v)) - for i := range v { - out[i] = mapping.MaskKnownOriginals(v[i]) - } - return out - case map[string]any: - out := make(map[string]any, len(v)) - for key, item := range v { - out[key] = sanitizePIILogValue(item, mapping) - } - return out - default: - return value - } -} - -func cloneGenerateRequest(req domain.GenerateRequest) domain.GenerateRequest { - clone := req - clone.RouterID = cloneStringPtr(req.RouterID) - clone.RoutedPublicModelID = cloneStringPtr(req.RoutedPublicModelID) - clone.MatchedCategory = cloneStringPtr(req.MatchedCategory) - clone.DecisionReason = cloneStringPtr(req.DecisionReason) - clone.Instructions = cloneStringPtr(req.Instructions) - clone.User = cloneStringPtr(req.User) - clone.ServiceTier = cloneStringPtr(req.ServiceTier) - if req.Input != nil { - clone.Input = make([]domain.InputItem, len(req.Input)) - for i := range req.Input { - clone.Input[i] = req.Input[i] - clone.Input[i].Role = cloneStringPtr(req.Input[i].Role) - clone.Input[i].Content = cloneStringPtr(req.Input[i].Content) - clone.Input[i].ReasoningContent = cloneStringPtr(req.Input[i].ReasoningContent) - clone.Input[i].ToolCallID = cloneStringPtr(req.Input[i].ToolCallID) - if req.Input[i].ToolCalls != nil { - clone.Input[i].ToolCalls = append([]domain.ToolCall(nil), req.Input[i].ToolCalls...) - } - } - } - if req.Tools != nil { - clone.Tools = make([]domain.ToolDefinition, len(req.Tools)) - for i := range req.Tools { - clone.Tools[i] = req.Tools[i] - clone.Tools[i].Parameters = cloneMap(req.Tools[i].Parameters) - } - } - if req.Stop != nil { - clone.Stop = append([]string(nil), req.Stop...) - } - if req.RoutingCategoryScores != nil { - clone.RoutingCategoryScores = append([]domain.RoutingCategoryScore(nil), req.RoutingCategoryScores...) - } - clone.Metadata = cloneMap(req.Metadata) - clone.ProviderOptions = cloneMap(req.ProviderOptions) - if req.TextConfig != nil { - tc := *req.TextConfig - tc.FormatType = cloneStringPtr(req.TextConfig.FormatType) - tc.JSONSchema = cloneMap(req.TextConfig.JSONSchema) - clone.TextConfig = &tc - } - if req.ToolChoice != nil { - tc := *req.ToolChoice - tc.FunctionName = cloneStringPtr(req.ToolChoice.FunctionName) - clone.ToolChoice = &tc - } - if req.LogitBias != nil { - clone.LogitBias = make(map[string]int, len(req.LogitBias)) - for key, value := range req.LogitBias { - clone.LogitBias[key] = value - } - } - return clone -} - -func cloneStringPtr(in *string) *string { - if in == nil { - return nil - } - out := *in - return &out -} - -func cloneMap(in map[string]any) map[string]any { - if in == nil { - return nil - } - raw, err := json.Marshal(in) - if err != nil { - out := make(map[string]any, len(in)) - for key, value := range in { - out[key] = value - } - return out - } - var out map[string]any - if err := json.Unmarshal(raw, &out); err != nil { - shallow := make(map[string]any, len(in)) - for key, value := range in { - shallow[key] = value - } - return shallow - } - return out -} diff --git a/apps/gateway/internal/application/services/pii_masking_test.go b/apps/gateway/internal/application/services/pii_masking_test.go deleted file mode 100644 index ecceda4..0000000 --- a/apps/gateway/internal/application/services/pii_masking_test.go +++ /dev/null @@ -1,851 +0,0 @@ -package services - -import ( - "context" - "io" - "strings" - "testing" - - "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" - "github.com/shopspring/decimal" -) - -// fakePIIFilter is a test double for ports.PIIFilter. It returns the -// pre-configured entities when called and tracks invocation counts so tests -// can assert the masking pipeline ran. -type fakePIIFilter struct { - enabled bool - err error - calls int - lastText string - lastOpts ports.PIIAnalyzeOptions - // byText optionally maps an exact input string to entity spans, so tests - // can return different detections for different request fields. - byText map[string][]domain.PIIEntity - // entities is returned when byText does not match. - entities []domain.PIIEntity -} - -type loggedWarning struct { - msg string - fields []any -} - -type captureLogger struct { - warnings []loggedWarning -} - -func (l *captureLogger) Debug(string, ...any) {} -func (l *captureLogger) Info(string, ...any) {} -func (l *captureLogger) Error(string, ...any) {} -func (l *captureLogger) Warn(msg string, fields ...any) { - l.warnings = append(l.warnings, loggedWarning{msg: msg, fields: append([]any(nil), fields...)}) -} - -func (f *fakePIIFilter) Enabled() bool { return f.enabled } - -func (f *fakePIIFilter) Analyze(_ context.Context, text string, opts ports.PIIAnalyzeOptions) ([]domain.PIIEntity, error) { - f.calls++ - f.lastText = text - f.lastOpts = opts - if f.err != nil { - return nil, f.err - } - if f.byText != nil { - if v, ok := f.byText[text]; ok { - return v, nil - } - return nil, nil - } - return f.entities, nil -} - -// captureProvider records the exact GenerateRequest the service forwards so -// tests can assert that user-supplied PII never reaches the upstream call. -type captureProvider struct { - lastReq domain.GenerateRequest - respondMsg string - respondOutput []domain.OutputItem -} - -func (p *captureProvider) Generate(_ context.Context, req domain.GenerateRequest, model domain.PublicModel) (domain.GenerateResult, error) { - p.lastReq = req - msg := p.respondMsg - role := "assistant" - output := p.respondOutput - if output == nil { - output = []domain.OutputItem{ - {Type: domain.OutputItemTypeMessage, Role: &role, Content: &msg}, - } - } - return domain.GenerateResult{ - ID: "resp-1", - PublicModelID: req.PublicModelID, - ProviderName: model.ProviderConfig.ProviderName, - ProviderModelID: model.ProviderModelID, - Output: output, - }, nil -} - -func (p *captureProvider) StreamGenerate(context.Context, domain.GenerateRequest, domain.PublicModel) (ports.GenerationStream, error) { - return nil, nil -} - -func buildPIITestService(filter ports.PIIFilter, prov *captureProvider) (*GenerateService, *stubUsageMeter) { - return buildPIITestServiceWithPIIMode(filter, prov, domain.APIKeyPIIModeBalanced) -} - -func buildPIITestServiceWithPIIMode(filter ports.PIIFilter, prov *captureProvider, piiMode string) (*GenerateService, *stubUsageMeter) { - usage := &stubUsageMeter{} - selected := domain.PublicModel{ - PublicModelID: "openai/gpt-4.1-mini", - ProviderModelID: "gpt-4.1-mini", - ProviderConfig: domain.ProviderConfig{ProviderName: "openai"}, - SupportsChatCompletions: true, - SupportsTools: true, - MaxContextWindow: 100000, - MaxOutputTokens: 16384, - InputPricePerMillion: decimal.NewFromFloat(0.75), - OutputPricePerMillion: decimal.NewFromFloat(4.50), - } - svc := NewGenerateService( - &stubAuthService{authCtx: domain.AuthContext{ - Account: domain.Account{ID: "acc1", Status: domain.AccountStatusActive}, - APIKey: domain.APIKey{ID: "key1", Active: true, PIIMode: piiMode}, - }}, - &stubModelCatalog{model: selected}, - nil, - &captureProviderRegistry{p: prov}, - usage, - filter, - stubLogger{}, - ) - return svc, usage -} - -// captureProviderRegistry returns the captureProvider as a GenerationProvider. -type captureProviderRegistry struct { - p *captureProvider -} - -func (r *captureProviderRegistry) GetProvider(string) (ports.GenerationProvider, error) { - return r.p, nil -} - -func newMessageInput(text string) []domain.InputItem { - role := "user" - t := text - return []domain.InputItem{ - {Type: domain.InputItemTypeMessage, Role: &role, Content: &t}, - } -} - -// captureProvider implements ports.GenerationProvider directly so we can -// observe the exact request the service forwards upstream. - -func TestGenerateServiceExecute_MasksRequestAndUnmasksResponse(t *testing.T) { - prompt := "My name is John Smith, email john@x.com" - // Pre-computed byte offsets for "John Smith" (11..21) and "john@x.com" (29..39). - filter := &fakePIIFilter{ - enabled: true, - entities: []domain.PIIEntity{ - {Type: "PERSON", Start: 11, End: 21, Score: 0.99}, - {Type: "EMAIL", Start: 29, End: 39, Score: 0.99}, - }, - } - prov := &captureProvider{respondMsg: "Hi [PERSON_1], I see your email [EMAIL_1]."} - - svc, _ := buildPIITestService(filter, prov) - - result, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Input: newMessageInput(prompt), - }, "sk-test") - if err != nil { - t.Fatalf("Execute: %v", err) - } - - // The provider must have seen only placeholders. - gotContent := *prov.lastReq.Input[0].Content - if strings.Contains(gotContent, "John Smith") || strings.Contains(gotContent, "john@x.com") { - t.Fatalf("upstream prompt still contains PII: %q", gotContent) - } - if !strings.Contains(gotContent, "[PERSON_1]") || !strings.Contains(gotContent, "[EMAIL_1]") { - t.Fatalf("upstream prompt missing placeholders: %q", gotContent) - } - - // The client-facing result must have the original PII restored. - if got := *result.Output[0].Content; got != "Hi John Smith, I see your email john@x.com." { - t.Fatalf("response not unmasked: %q", got) - } - if filter.calls != 1 { - t.Fatalf("filter calls = %d, want 1", filter.calls) - } -} - -func TestGenerateServiceExecute_MasksInstructionsContentReasoningAndToolArgs(t *testing.T) { - instructions := "Call Jane at jane@example.com" - prompt := "Email Jane at jane@example.com" - toolArgs := `{"email":"jane@example.com","note":"Call Jane"}` - role := "user" - assistantInputRole := "assistant" - assistantRole := "assistant" - reasoning := "I should contact Jane" - filter := &fakePIIFilter{ - enabled: true, - byText: map[string][]domain.PIIEntity{ - instructions: { - {Type: "PERSON", Start: 5, End: 9, Score: 0.99}, - {Type: "EMAIL", Start: 13, End: 29, Score: 0.99}, - }, - prompt: { - {Type: "PERSON", Start: 6, End: 10, Score: 0.99}, - {Type: "EMAIL", Start: 14, End: 30, Score: 0.99}, - }, - reasoning: { - {Type: "PERSON", Start: 17, End: 21, Score: 0.99}, - }, - "jane@example.com": { - {Type: "EMAIL", Start: 0, End: 16, Score: 0.99}, - }, - "Call Jane": { - {Type: "PERSON", Start: 5, End: 9, Score: 0.99}, - }, - }, - } - prov := &captureProvider{ - respondOutput: []domain.OutputItem{ - { - Type: domain.OutputItemTypeMessage, - Role: &assistantRole, - Content: strPtr("Message [PERSON_1] at [EMAIL_1]."), - ToolCalls: []domain.ToolCall{ - {ID: "call_1", Name: "send_email", ArgumentsJSON: `{"email":"[EMAIL_1]"}`}, - }, - }, - }, - } - svc, _ := buildPIITestService(filter, prov) - - result, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Instructions: &instructions, - Input: []domain.InputItem{ - { - Type: domain.InputItemTypeMessage, - Role: &role, - Content: strPtr(prompt), - }, - { - Type: domain.InputItemTypeMessage, - Role: &assistantInputRole, - ReasoningContent: &reasoning, - ToolCalls: []domain.ToolCall{ - {ID: "call_1", Name: "send_email", ArgumentsJSON: toolArgs}, - }, - }, - }, - }, "sk-test") - if err != nil { - t.Fatalf("Execute: %v", err) - } - - if got := *prov.lastReq.Instructions; strings.Contains(got, "Jane") || strings.Contains(got, "jane@example.com") { - t.Fatalf("upstream instructions still contain PII: %q", got) - } - if got := *prov.lastReq.Input[0].Content; strings.Contains(got, "Jane") || strings.Contains(got, "jane@example.com") { - t.Fatalf("upstream content still contains PII: %q", got) - } - if got := *prov.lastReq.Input[1].ReasoningContent; strings.Contains(got, "Jane") { - t.Fatalf("upstream reasoning content still contains PII: %q", got) - } - wantArgs := `{"email":"[EMAIL_1]","note":"Call [PERSON_1]"}` - if got := prov.lastReq.Input[1].ToolCalls[0].ArgumentsJSON; got != wantArgs { - t.Fatalf("tool-call arguments not masked:\n got = %q\nwant = %q", got, wantArgs) - } - if filter.calls != 5 { - t.Fatalf("filter calls = %d, want 5 for instructions, content, reasoning, and JSON values", filter.calls) - } - - if got := *result.Output[0].Content; got != "Message Jane at jane@example.com." { - t.Fatalf("response content not restored: %q", got) - } - if got := result.Output[0].ToolCalls[0].ArgumentsJSON; got != `{"email":"jane@example.com"}` { - t.Fatalf("response tool-call arguments not restored, got %q", got) - } -} - -func TestGenerateServiceExecute_MasksToolResultJSONAndUserField(t *testing.T) { - toolRole := "tool" - user := "alice@example.com" - toolResult := `{"customer":{"name":"Alice","email":"alice@example.com"},"paid":false}` - filter := &fakePIIFilter{ - enabled: true, - byText: map[string][]domain.PIIEntity{ - "user prompt": { - {Type: "PERSON", Start: 0, End: 0, Score: 0.99}, // invalid span proves no accidental fallback masking - }, - "alice@example.com": { - {Type: "EMAIL", Start: 0, End: 17, Score: 0.99}, - }, - "Alice": { - {Type: "PERSON", Start: 0, End: 5, Score: 0.99}, - }, - }, - } - prov := &captureProvider{respondMsg: "ok"} - svc, _ := buildPIITestService(filter, prov) - - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - User: &user, - Input: []domain.InputItem{ - { - Type: domain.InputItemTypeMessage, - Role: &toolRole, - Content: &toolResult, - ToolCallID: strPtr("call_1"), - }, - }, - }, "sk-test") - if err != nil { - t.Fatalf("Execute: %v", err) - } - - if got := *prov.lastReq.User; got != "[EMAIL_1]" { - t.Fatalf("user field not masked: %q", got) - } - gotContent := *prov.lastReq.Input[0].Content - if strings.Contains(gotContent, "Alice") || strings.Contains(gotContent, "alice@example.com") { - t.Fatalf("tool result still contains PII: %q", gotContent) - } - if !strings.Contains(gotContent, "[PERSON_1]") || !strings.Contains(gotContent, "[EMAIL_1]") { - t.Fatalf("tool result missing placeholders: %q", gotContent) - } -} - -func TestGenerateServiceExecute_MasksInvalidJSONToolArgumentsAsPlainText(t *testing.T) { - assistantRole := "assistant" - args := `email=jane@example.com note=Call Jane` - filter := &fakePIIFilter{ - enabled: true, - byText: map[string][]domain.PIIEntity{ - args: { - {Type: "EMAIL", Start: 6, End: 22, Score: 0.99}, - {Type: "PERSON", Start: 33, End: 37, Score: 0.99}, - }, - }, - } - prov := &captureProvider{respondMsg: "ok"} - svc, _ := buildPIITestService(filter, prov) - - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Input: []domain.InputItem{ - { - Type: domain.InputItemTypeMessage, - Role: &assistantRole, - ToolCalls: []domain.ToolCall{ - {ID: "call_1", Name: "send_email", ArgumentsJSON: args}, - }, - }, - }, - }, "sk-test") - if err != nil { - t.Fatalf("Execute: %v", err) - } - - want := `email=[EMAIL_1] note=Call [PERSON_1]` - if got := prov.lastReq.Input[0].ToolCalls[0].ArgumentsJSON; got != want { - t.Fatalf("invalid JSON tool args not masked as plain text:\n got = %q\nwant = %q", got, want) - } -} - -func TestGenerateServiceExecute_StaticFieldsAreNotScanned(t *testing.T) { - description := "Send email to Jane" - filter := &fakePIIFilter{ - enabled: true, - byText: map[string][]domain.PIIEntity{ - description: { - {Type: "PERSON", Start: 14, End: 18, Score: 0.99}, - }, - }, - } - - prov := &captureProvider{respondMsg: "ok"} - svc, _ := buildPIITestService(filter, prov) - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Input: newMessageInput("hello"), - Tools: []domain.ToolDefinition{ - {Name: "send_email", Description: description, Parameters: map[string]any{"type": "object"}}, - }, - }, "sk-test") - if err != nil { - t.Fatalf("Execute: %v", err) - } - if got := prov.lastReq.Tools[0].Description; got != description { - t.Fatalf("static description should not be masked, got %q", got) - } -} - -func TestGenerateServiceExecute_FailsClosedWhenFilterErrors(t *testing.T) { - filter := &fakePIIFilter{enabled: true, err: io.ErrUnexpectedEOF} - prov := &captureProvider{} - svc, usage := buildPIITestService(filter, prov) - - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Input: newMessageInput("hello"), - }, "sk-test") - if err == nil { - t.Fatal("expected fail-closed error when filter errors") - } - if usage.failureCalls == 0 { - t.Fatalf("expected failure recorded") - } - if prov.lastReq.PublicModelID != "" { - t.Fatalf("provider should not have been called") - } -} - -func TestGenerateServiceExecute_FailOpenContinuesWhenFilterErrors(t *testing.T) { - filter := &fakePIIFilter{enabled: true, err: io.ErrUnexpectedEOF} - prov := &captureProvider{respondMsg: "ok"} - svc, _ := buildPIITestService(filter, prov) - svc.SetPIIOptions("en", true) - - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Input: newMessageInput("hello"), - }, "sk-test") - if err != nil { - t.Fatalf("Execute (fail-open): %v", err) - } - if got := *prov.lastReq.Input[0].Content; got != "hello" { - t.Fatalf("prompt should pass through on fail-open, got %q", got) - } -} - -func TestGenerateServiceExecute_BypassesWhenFilterDisabled(t *testing.T) { - filter := &fakePIIFilter{enabled: false} - prov := &captureProvider{respondMsg: "ok"} - svc, _ := buildPIITestService(filter, prov) - - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Input: newMessageInput("My name is John Smith"), - }, "sk-test") - if err != nil { - t.Fatalf("Execute: %v", err) - } - if filter.calls != 0 { - t.Fatalf("filter should not be called when disabled") - } - if got := *prov.lastReq.Input[0].Content; got != "My name is John Smith" { - t.Fatalf("prompt should pass through, got %q", got) - } -} - -func TestGenerateServiceExecute_BypassesWhenAPIKeyPIIModeOff(t *testing.T) { - filter := &fakePIIFilter{enabled: true, entities: []domain.PIIEntity{{Type: "PERSON", Start: 11, End: 21, Score: 0.99}}} - prov := &captureProvider{respondMsg: "ok"} - svc, _ := buildPIITestServiceWithPIIMode(filter, prov, domain.APIKeyPIIModeOff) - - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Input: newMessageInput("My name is John Smith"), - }, "sk-test") - if err != nil { - t.Fatalf("Execute: %v", err) - } - if filter.calls != 0 { - t.Fatalf("filter should not be called when pii_mode=off, calls=%d", filter.calls) - } - if got := *prov.lastReq.Input[0].Content; got != "My name is John Smith" { - t.Fatalf("prompt should pass through, got %q", got) - } -} - -func TestGenerateServiceExecute_PassesAPIKeyPIIModeToFilter(t *testing.T) { - filter := &fakePIIFilter{enabled: true, entities: []domain.PIIEntity{{Type: "EMAIL_ADDRESS", Start: 6, End: 22, Score: 0.99}}} - prov := &captureProvider{respondMsg: "Email [EMAIL_ADDRESS_1]"} - svc, _ := buildPIITestServiceWithPIIMode(filter, prov, domain.APIKeyPIIModeLow) - - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Input: newMessageInput("email jane@example.com"), - }, "sk-test") - if err != nil { - t.Fatalf("Execute: %v", err) - } - if filter.lastOpts.Mode != domain.APIKeyPIIModeLow { - t.Fatalf("filter mode = %q, want low", filter.lastOpts.Mode) - } -} - -func TestGenerateServiceExecute_SkipsEmptyPIIFields(t *testing.T) { - instructions := "" - content := "" - filter := &fakePIIFilter{enabled: true, entities: []domain.PIIEntity{{Type: "PERSON", Start: 0, End: 4}}} - prov := &captureProvider{respondMsg: "ok"} - svc, _ := buildPIITestService(filter, prov) - - _, _, err := svc.Execute(context.Background(), domain.EndpointChatCompletions, domain.GenerateRequest{ - PublicModelID: "openai/gpt-4.1-mini", - Instructions: &instructions, - Input: []domain.InputItem{ - {Type: domain.InputItemTypeMessage, Content: &content}, - }, - }, "sk-test") - if err != nil { - t.Fatalf("Execute: %v", err) - } - if filter.calls != 0 { - t.Fatalf("filter calls = %d, want 0 for empty fields", filter.calls) - } - if got := *prov.lastReq.Input[0].Content; got != "" { - t.Fatalf("empty content should pass through, got %q", got) - } -} - -// --- streaming unmask wrapper tests --- - -// scriptedStream replays a fixed sequence of events to the wrapper under test. -type scriptedStream struct { - events []domain.StreamEvent - errs []error - idx int - closed bool -} - -func (s *scriptedStream) Recv() (domain.StreamEvent, error) { - if s.idx >= len(s.events) { - return domain.StreamEvent{}, io.EOF - } - ev := s.events[s.idx] - var err error - if s.idx < len(s.errs) { - err = s.errs[s.idx] - } - s.idx++ - return ev, err -} - -func (s *scriptedStream) Close() error { s.closed = true; return nil } - -func textDelta(s string) domain.StreamEvent { - v := s - return domain.StreamEvent{Type: domain.StreamEventOutputTextDelta, ContentDelta: &v} -} - -func completedEvent() domain.StreamEvent { - reason := "stop" - return domain.StreamEvent{Type: domain.StreamEventCompleted, FinishReason: &reason} -} - -func toolCallDelta(args string) domain.StreamEvent { - return domain.StreamEvent{ - Type: domain.StreamEventToolCallDelta, - ToolCallDelta: &domain.ToolCallDelta{ - Index: 0, - ArgumentsDelta: &args, - }, - } -} - -func collectStream(t *testing.T, st ports.GenerationStream) string { - t.Helper() - var b strings.Builder - for { - ev, err := st.Recv() - if err == io.EOF { - return b.String() - } - if err != nil { - t.Fatalf("Recv error: %v", err) - } - if ev.ContentDelta != nil { - b.WriteString(*ev.ContentDelta) - } - } -} - -func collectStreamUntilCompleted(t *testing.T, st ports.GenerationStream) string { - t.Helper() - var b strings.Builder - for { - ev, err := st.Recv() - if err == io.EOF { - return b.String() - } - if err != nil { - t.Fatalf("Recv error: %v", err) - } - if ev.ContentDelta != nil { - b.WriteString(*ev.ContentDelta) - } - if ev.Type == domain.StreamEventCompleted { - for { - if _, err := st.Recv(); err != nil { - return b.String() - } - } - } - } -} - -func TestPIIUnmaskingStream_PassthroughWithoutTokens(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("PERSON", "John") - inner := &scriptedStream{events: []domain.StreamEvent{textDelta("hello "), textDelta("world")}} - got := collectStream(t, newPIIUnmaskingStream(inner, m)) - if got != "hello world" { - t.Fatalf("got %q, want hello world", got) - } -} - -func TestPIIUnmaskingStream_RestoresCompleteToken(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("PERSON", "John") - inner := &scriptedStream{events: []domain.StreamEvent{textDelta("hi [PERSON_1] there")}} - got := collectStream(t, newPIIUnmaskingStream(inner, m)) - if got != "hi John there" { - t.Fatalf("got %q", got) - } -} - -func TestPIIUnmaskingStream_BuffersTokenSplitAcrossChunks(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("PERSON", "John") - inner := &scriptedStream{events: []domain.StreamEvent{ - textDelta("hi [PER"), - textDelta("SON_1] there"), - }} - got := collectStream(t, newPIIUnmaskingStream(inner, m)) - if got != "hi John there" { - t.Fatalf("got %q", got) - } -} - -func TestPIIUnmaskingStream_KeepsUnknownTokensVerbatim(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("PERSON", "John") - inner := &scriptedStream{events: []domain.StreamEvent{textDelta("see [GHOST_5] and [PERSON_1]")}} - got := collectStream(t, newPIIUnmaskingStream(inner, m)) - if got != "see [GHOST_5] and John" { - t.Fatalf("got %q", got) - } -} - -func TestPIIUnmaskingStream_FlushesTailOnEOF(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("PERSON", "John") - // Last chunk leaves an unterminated `[PER` in the buffer. - inner := &scriptedStream{events: []domain.StreamEvent{textDelta("end with [PER")}} - got := collectStream(t, newPIIUnmaskingStream(inner, m)) - if got != "end with [PER" { - t.Fatalf("got %q", got) - } -} - -func TestPIIUnmaskingStream_FlushesTailBeforeCompleted(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("PERSON", "John") - inner := &scriptedStream{events: []domain.StreamEvent{ - textDelta("hi [PER"), - completedEvent(), - }} - got := collectStreamUntilCompleted(t, newPIIUnmaskingStream(inner, m)) - if got != "hi [PER" { - t.Fatalf("got %q", got) - } -} - -func TestPIIUnmaskingStream_RestoresToolCallArgumentDeltas(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("EMAIL", "jane@example.com") - inner := &scriptedStream{events: []domain.StreamEvent{ - toolCallDelta(`{"email":"[EMAIL_1]"}`), - }} - st := newPIIUnmaskingStream(inner, m) - - ev, err := st.Recv() - if err != nil { - t.Fatalf("Recv: %v", err) - } - if ev.ToolCallDelta == nil || ev.ToolCallDelta.ArgumentsDelta == nil { - t.Fatalf("expected tool-call argument delta, got %#v", ev) - } - if got := *ev.ToolCallDelta.ArgumentsDelta; got != `{"email":"jane@example.com"}` { - t.Fatalf("tool-call argument delta not restored, got %q", got) - } -} - -func TestPIIUnmaskingStream_RestoresBracketlessToolCallArgumentAlias(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("EMAIL_ADDRESS", "jane@example.com") - inner := &scriptedStream{events: []domain.StreamEvent{ - toolCallDelta(`{"email":"EMAIL_ADDRESS_1"}`), - }} - st := newPIIUnmaskingStream(inner, m) - - ev, err := st.Recv() - if err != nil { - t.Fatalf("Recv: %v", err) - } - if got := *ev.ToolCallDelta.ArgumentsDelta; got != `{"email":"jane@example.com"}` { - t.Fatalf("tool-call argument delta not restored, got %q", got) - } -} - -func TestPIIUnmaskingStream_BuffersBracketlessAliasSplitAcrossToolCallArgumentDeltas(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("EMAIL_ADDRESS", "jane@example.com") - inner := &scriptedStream{events: []domain.StreamEvent{ - toolCallDelta(`{"email":"EMAIL_ADD`), - toolCallDelta(`RESS_1"}`), - }} - st := newPIIUnmaskingStream(inner, m) - - ev, err := st.Recv() - if err != nil { - t.Fatalf("Recv first: %v", err) - } - if got := *ev.ToolCallDelta.ArgumentsDelta; got != `{"email":"` { - t.Fatalf("first tool-call argument delta = %q", got) - } - - ev, err = st.Recv() - if err != nil { - t.Fatalf("Recv second: %v", err) - } - if got := *ev.ToolCallDelta.ArgumentsDelta; got != `jane@example.com"}` { - t.Fatalf("second tool-call argument delta = %q", got) - } -} - -func TestPIIUnmaskingStream_BuffersTokenSplitAcrossToolCallArgumentDeltas(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("EMAIL", "jane@example.com") - inner := &scriptedStream{events: []domain.StreamEvent{ - toolCallDelta(`{"email":"[EMA`), - toolCallDelta(`IL_1]"}`), - }} - st := newPIIUnmaskingStream(inner, m) - - ev, err := st.Recv() - if err != nil { - t.Fatalf("Recv first: %v", err) - } - if got := *ev.ToolCallDelta.ArgumentsDelta; got != `{"email":"` { - t.Fatalf("first tool-call argument delta = %q", got) - } - - ev, err = st.Recv() - if err != nil { - t.Fatalf("Recv second: %v", err) - } - if got := *ev.ToolCallDelta.ArgumentsDelta; got != `jane@example.com"}` { - t.Fatalf("second tool-call argument delta = %q", got) - } -} - -func TestUnmaskResult_LogsUnresolvedToolCallArgumentTokens(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("EMAIL_ADDRESS", "jane@example.com") - logger := &captureLogger{} - role := "assistant" - result := domain.GenerateResult{ - Output: []domain.OutputItem{ - { - Type: domain.OutputItemTypeMessage, - Role: &role, - ToolCalls: []domain.ToolCall{ - {ID: "call_1", Name: "send_email", ArgumentsJSON: `{"email":"EMAIL_ADDRESS_2"}`}, - }, - }, - }, - } - - unmaskResult(&result, m, logger, "request_id", "req-1") - - if len(logger.warnings) != 1 { - t.Fatalf("warnings = %d, want 1", len(logger.warnings)) - } - if logger.warnings[0].msg != "pii restoration unresolved tokens" { - t.Fatalf("warning msg = %q", logger.warnings[0].msg) - } - if got := result.Output[0].ToolCalls[0].ArgumentsJSON; got != `{"email":"EMAIL_ADDRESS_2"}` { - t.Fatalf("arguments changed unexpectedly: %q", got) - } -} - -func TestPIIUnmaskingStream_LogsUnresolvedToolCallArgumentTokensOnCompleted(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("EMAIL_ADDRESS", "jane@example.com") - logger := &captureLogger{} - inner := &scriptedStream{events: []domain.StreamEvent{ - toolCallDelta(`{"email":"EMAIL_ADDRESS_2"}`), - completedEvent(), - }} - st := newPIIUnmaskingStreamWithLogger(inner, m, logger, "request_id", "req-1") - - for { - _, err := st.Recv() - if err == io.EOF { - break - } - if err != nil { - t.Fatalf("Recv: %v", err) - } - } - - if len(logger.warnings) != 1 { - t.Fatalf("warnings = %d, want 1", len(logger.warnings)) - } - if logger.warnings[0].msg != "pii restoration unresolved tokens" { - t.Fatalf("warning msg = %q", logger.warnings[0].msg) - } -} - -func TestSanitizeErrorWithPIIMapping_MasksKnownOriginalsInMessageAndMetadata(t *testing.T) { - m := domain.NewPIIMapping() - m.Token("EMAIL", "jane@example.com") - err := domain.ErrProviderError(502, "provider echoed jane@example.com").WithMeta( - "upstream_error", `bad request for jane@example.com`, - "nested", map[string]any{"body": "jane@example.com"}, - ) - - sanitized, ok := sanitizeErrorWithPIIMapping(err, m).(*domain.GatewayError) - if !ok { - t.Fatalf("expected GatewayError") - } - if strings.Contains(sanitized.Message, "jane@example.com") { - t.Fatalf("message still contains PII: %q", sanitized.Message) - } - if got := sanitized.Metadata["upstream_error"]; got != `bad request for [EMAIL_1]` { - t.Fatalf("upstream_error = %#v", got) - } - nested := sanitized.Metadata["nested"].(map[string]any) - if got := nested["body"]; got != "[EMAIL_1]" { - t.Fatalf("nested body = %#v", got) - } -} - -func TestPIIUnmaskingStream_CloseDelegates(t *testing.T) { - m := domain.NewPIIMapping() - inner := &scriptedStream{} - st := newPIIUnmaskingStream(inner, m) - if err := st.Close(); err != nil { - t.Fatalf("Close: %v", err) - } - if !inner.closed { - t.Fatal("inner stream not closed") - } -} - -func strPtr(v string) *string { - return &v -} diff --git a/apps/gateway/internal/application/services/pii_test.go b/apps/gateway/internal/application/services/pii_test.go new file mode 100644 index 0000000..eb5e846 --- /dev/null +++ b/apps/gateway/internal/application/services/pii_test.go @@ -0,0 +1,308 @@ +package services + +import ( + "context" + "encoding/json" + "io" + "strings" + "testing" + + "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" +) + +// fakePIIFilter returns configured entities per exact text and counts calls. +type fakePIIFilter struct { + enabled bool + err error + calls int + lastOpts ports.PIIAnalyzeOptions + byText map[string][]domain.PIIEntity +} + +func (f *fakePIIFilter) Enabled() bool { return f.enabled } + +func (f *fakePIIFilter) Analyze(_ context.Context, text string, opts ports.PIIAnalyzeOptions) ([]domain.PIIEntity, error) { + f.calls++ + f.lastOpts = opts + if f.err != nil { + return nil, f.err + } + return f.byText[text], nil +} + +type captureLogger struct { + stubLogger + warnings []string +} + +func (l *captureLogger) Warn(msg string, _ ...any) { l.warnings = append(l.warnings, msg) } + +func piiService(filter ports.PIIFilter, mode string, provider *stubProvider) (*GenerateService, *resultMeter) { + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.Provider{"primary": provider}, false) + svc.pii = filter + svc.auth = &stubAuthService{authCtx: domain.AuthContext{Account: domain.Account{ID: "acc1", Status: domain.AccountStatusActive}, APIKey: domain.APIKey{ID: "key1", Active: true, PIIMode: mode}}} + return svc, meter +} + +func entity(kind string, text, in string) domain.PIIEntity { + start := strings.Index(in, text) + return domain.PIIEntity{Type: kind, Start: start, End: start + len(text), Score: 0.99} +} + +// The LLM-bound text of the request is masked and the response restored; +// everything else is forwarded as sent. +func TestPII_MasksRequestAndRestoresResponse(t *testing.T) { + prompt := "Email Jane at jane@example.com" + reasoning := "I should contact Jane" + filter := &fakePIIFilter{enabled: true, byText: map[string][]domain.PIIEntity{ + prompt: {entity("PERSON", "Jane", prompt), entity("EMAIL", "jane@example.com", prompt)}, + reasoning: {entity("PERSON", "Jane", reasoning)}, + "jane@example.com": {entity("EMAIL", "jane@example.com", "jane@example.com")}, + "Call Jane": {entity("PERSON", "Jane", "Call Jane")}, + "Alice": {entity("PERSON", "Alice", "Alice")}, + "Look at Alice": {entity("PERSON", "Alice", "Look at Alice")}, + }} + provider := &stubProvider{json: `{"id":"r1","model":"upstream","choices":[{"index":0,"message":{"role":"assistant","content":"Message [PERSON_1] at [EMAIL_1].","tool_calls":[{"id":"c","type":"function","function":{"name":"send","arguments":"{\"email\":\"[EMAIL_1]\"}"}}]},"finish_reason":"tool_calls"}],"usage":{"prompt_tokens":1,"completion_tokens":1},"system_fingerprint":"fp"}`} + svc, _ := piiService(filter, domain.APIKeyPIIModeBalanced, provider) + raw := `{"model":"m","user":"jane@example.com","messages":[ + {"role":"user","content":"` + prompt + `"}, + {"role":"user","content":[{"type":"text","text":"Look at Alice"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AAAA"}}]}, + {"role":"assistant","reasoning_content":"` + reasoning + `","tool_calls":[{"id":"c1","type":"function","function":{"name":"send","arguments":"{\"email\":\"jane@example.com\",\"note\":\"Call Jane\"}"}}]}, + {"role":"tool","tool_call_id":"c1","content":"{\"customer\":{\"name\":\"Alice\"},\"paid\":false,\"id\":9007199254740993}"}], + "tools":[{"type":"function","function":{"name":"send","description":"Send email to Jane","parameters":{"type":"object"}}}],"temperature":0.2}` + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(raw), domain.GenerateRequest{PublicModelID: "zai-org/glm-5.3-flash"}, "key") + if err != nil { + t.Fatal(err) + } + sent := provider.lastRaw + for _, leaked := range []string{"jane@example.com", "Contact Jane", "Call Jane", `"name":"Alice"`, "Look at Alice", "Email Jane"} { + if strings.Contains(sent, leaked) { + t.Fatalf("upstream body leaks %q: %s", leaked, sent) + } + } + for _, kept := range []string{"[EMAIL_1]", "[PERSON_1]", `"Send email to Jane"`, `data:image/png;base64,AAAA`, `9007199254740993`, `"temperature":0.2`, `\"paid\":false`} { + if !strings.Contains(sent, kept) { + t.Fatalf("upstream body missing %s: %s", kept, sent) + } + } + var resp map[string]any + if err := json.Unmarshal(call.Body, &resp); err != nil { + t.Fatal(err) + } + msg := resp["choices"].([]any)[0].(map[string]any)["message"].(map[string]any) + if msg["content"] != "Message Jane at jane@example.com." { + t.Fatalf("content = %v", msg["content"]) + } + args := msg["tool_calls"].([]any)[0].(map[string]any)["function"].(map[string]any)["arguments"] + if args != `{"email":"jane@example.com"}` || resp["system_fingerprint"] != "fp" || resp["model"] != "zai-org/glm-5.3-flash" { + t.Fatalf("response %s", call.Body) + } + if filter.lastOpts.Mode != domain.APIKeyPIIModeBalanced { + t.Fatalf("mode = %q", filter.lastOpts.Mode) + } +} + +func TestPII_InvalidJSONToolArgumentsMaskedAsText(t *testing.T) { + args := `email=jane@example.com note=Call Jane` + filter := &fakePIIFilter{enabled: true, byText: map[string][]domain.PIIEntity{ + args: {entity("EMAIL", "jane@example.com", args), entity("PERSON", "Jane", args)}, + }} + provider := &stubProvider{} + svc, _ := piiService(filter, domain.APIKeyPIIModeHigh, provider) + raw := `{"messages":[{"role":"assistant","tool_calls":[{"id":"c","type":"function","function":{"name":"f","arguments":"` + args + `"}}]}]}` + if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(raw), domain.GenerateRequest{PublicModelID: "zai-org/glm-5.3-flash"}, "key"); err != nil { + t.Fatal(err) + } + if !strings.Contains(provider.lastRaw, `"arguments":"email=[EMAIL_1] note=Call [PERSON_1]"`) { + t.Fatalf("sent %s", provider.lastRaw) + } +} + +func TestPII_BypassAndFailureModes(t *testing.T) { + raw := `{"messages":[{"role":"user","content":"hi Jane"}],"seed":9007199254740993}` + detect := map[string][]domain.PIIEntity{"hi Jane": {entity("PERSON", "Jane", "hi Jane")}} + for name, tc := range map[string]struct { + filter *fakePIIFilter + mode string + failOpen bool + wantErr bool + wantSent string // "" means sent unchanged + }{ + "key without masking": {filter: &fakePIIFilter{enabled: true, byText: detect}, mode: domain.APIKeyPIIModeOff}, + "filter off": {filter: &fakePIIFilter{enabled: false, byText: detect}, mode: domain.APIKeyPIIModeHigh}, + "nothing found": {filter: &fakePIIFilter{enabled: true}, mode: domain.APIKeyPIIModeHigh}, + "filter fails closed": {filter: &fakePIIFilter{enabled: true, err: io.ErrUnexpectedEOF}, mode: domain.APIKeyPIIModeHigh, wantErr: true}, + "filter fails open": {filter: &fakePIIFilter{enabled: true, err: io.ErrUnexpectedEOF}, mode: domain.APIKeyPIIModeHigh, failOpen: true}, + "masking": {filter: &fakePIIFilter{enabled: true, byText: detect}, mode: domain.APIKeyPIIModeLow, wantSent: `hi [PERSON_1]`}, + } { + t.Run(name, func(t *testing.T) { + provider := &stubProvider{} + svc, meter := piiService(tc.filter, tc.mode, provider) + svc.piiFailOpen = tc.failOpen + _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(raw), domain.GenerateRequest{PublicModelID: "zai-org/glm-5.3-flash"}, "key") + if tc.wantErr { + if err == nil || provider.calls != 0 || meter.reserveCalls != 0 { + t.Fatalf("err=%v calls=%d reserve=%d", err, provider.calls, meter.reserveCalls) + } + return + } + if err != nil { + t.Fatal(err) + } + if tc.wantSent == "" && provider.lastRaw != raw { + t.Fatalf("body changed: %s", provider.lastRaw) + } + if tc.wantSent != "" && (!strings.Contains(provider.lastRaw, tc.wantSent) || !strings.Contains(provider.lastRaw, "9007199254740993")) { + t.Fatalf("sent %s", provider.lastRaw) + } + }) + } +} + +// streamed collects what a client reassembles from a proxied stream. +type streamed struct { + content, reasoning string + args map[int]string + finish string + raw string +} + +func runPIIStream(t *testing.T, sse string, logger ports.Logger) streamed { + t.Helper() + filter := &fakePIIFilter{enabled: true, byText: map[string][]domain.PIIEntity{ + "Email jane@example.com and John": {entity("EMAIL", "jane@example.com", "Email jane@example.com and John"), entity("PERSON", "John", "Email jane@example.com and John")}, + }} + provider := &stubProvider{sse: sse} + svc, _ := piiService(filter, domain.APIKeyPIIModeBalanced, provider) + if logger != nil { + svc.logger = logger + } + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{"stream":true,"stream_options":{"include_usage":true},"messages":[{"role":"user","content":"Email jane@example.com and John"}]}`), streamRequest(), "key") + if err != nil { + t.Fatal(err) + } + got, err := relayAll(t, call.Stream) + if err != nil { + t.Fatal(err) + } + call.Stream.Close() + out := streamed{args: map[int]string{}, raw: got} + for _, line := range strings.Split(got, "\n") { + data, ok := strings.CutPrefix(line, "data: ") + if !ok || data == "[DONE]" { + continue + } + var c struct { + Choices []struct { + Delta struct { + Content string `json:"content"` + ReasoningContent string `json:"reasoning_content"` + ToolCalls []struct { + Index int `json:"index"` + Function struct { + Arguments string `json:"arguments"` + } `json:"function"` + } `json:"tool_calls"` + } `json:"delta"` + FinishReason *string `json:"finish_reason"` + } `json:"choices"` + } + if err := json.Unmarshal([]byte(data), &c); err != nil { + t.Fatalf("invalid event %q: %v", data, err) + } + for _, ch := range c.Choices { + out.content += ch.Delta.Content + out.reasoning += ch.Delta.ReasoningContent + for _, tc := range ch.Delta.ToolCalls { + out.args[tc.Index] += tc.Function.Arguments + } + if ch.FinishReason != nil { + out.finish = *ch.FinishReason + } + } + } + return out +} + +func ev(delta string, finish string) string { + f := "null" + if finish != "" { + f = `"` + finish + `"` + } + return `data: {"id":"s1","model":"upstream","choices":[{"index":0,"delta":` + delta + `,"finish_reason":` + f + `}]}` + "\n\n" +} + +func TestPII_StreamRestoresPlaceholders(t *testing.T) { + sse := ev(`{"role":"assistant"}`, "") + + ev(`{"reasoning_content":"Write to [PER"}`, "") + + ev(`{"reasoning_content":"SON_1]"}`, "") + + ev(`{"content":"Hi [PERSON_1], see [GHOST_5] and [EMA"}`, "") + + ev(`{"content":"IL_1]"}`, "") + + ev(`{"tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"send","arguments":"{\"to\":\"EMA"}}]}`, "") + + ev(`{"tool_calls":[{"index":0,"function":{"arguments":"IL_1\"}"}}]}`, "") + + ev(`{"content":" bye [PER"}`, "stop") + + `data: {"id":"s1","choices":[],"usage":{"prompt_tokens":5,"completion_tokens":5}}` + "\n\n" + + "data: [DONE]\n\n" + got := runPIIStream(t, sse, nil) + if got.content != "Hi John, see [GHOST_5] and jane@example.com bye [PER" || got.reasoning != "Write to John" || got.finish != "stop" { + t.Fatalf("content %q reasoning %q finish %q", got.content, got.reasoning, got.finish) + } + // A placeholder the model wrote without brackets, split across events. + if got.args[0] != `{"to":"jane@example.com"}` { + t.Fatalf("args %q", got.args[0]) + } + // The event with no text passes byte for byte. + if !strings.Contains(got.raw, `data: {"id":"s1","model":"zai-org/glm-5.3-flash","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}`) { + t.Fatalf("role event changed: %s", got.raw) + } +} + +func TestPII_StreamReleasesHeldTextWhenTheProviderNeverFinishes(t *testing.T) { + got := runPIIStream(t, ev(`{"content":"Hi [PERSON_1] and [EMA"}`, "")+ev(`{"content":"IL_1] then [PER"}`, "")+"data: [DONE]\n\n", nil) + if got.content != "Hi John and jane@example.com then [PER" { + t.Fatalf("content %q", got.content) + } + if strings.Count(got.raw, "[DONE]") != 1 || strings.Index(got.raw, "then [PER") > strings.Index(got.raw, "[DONE]") { + t.Fatalf("held text not released before [DONE]: %s", got.raw) + } +} + +func TestPII_StreamToolArgumentsSplitAcrossEvents(t *testing.T) { + sse := ev(`{"tool_calls":[{"index":0,"id":"a","type":"function","function":{"name":"send","arguments":"{\"to\":\"[EMA"}},{"index":1,"id":"b","type":"function","function":{"name":"send","arguments":"{\"to\":\"[PERSON_1]\"}"}}]}`, "") + + ev(`{"tool_calls":[{"index":0,"function":{"arguments":"IL_1]\"}"}}]}`, "tool_calls") + "data: [DONE]\n\n" + got := runPIIStream(t, sse, nil) + if got.args[0] != `{"to":"jane@example.com"}` || got.args[1] != `{"to":"John"}` || got.finish != "tool_calls" { + t.Fatalf("args %v finish %q", got.args, got.finish) + } +} + +func TestPII_LogsPlaceholdersTheModelChanged(t *testing.T) { + logger := &captureLogger{} + runPIIStream(t, ev(`{"content":"Hi [PERSON_7]"}`, "stop")+"data: [DONE]\n\n", logger) + found := false + for _, w := range logger.warnings { + found = found || w == "pii restoration unresolved tokens" + } + if !found { + t.Fatalf("warnings %v", logger.warnings) + } +} + +func TestSanitizeErrorWithPIIMapping_MasksKnownOriginalsInMessageAndMetadata(t *testing.T) { + m := domain.NewPIIMapping() + m.Token("EMAIL", "jane@example.com") + err := domain.ErrProviderError(502, "provider echoed jane@example.com").WithMeta( + "upstream_error", `bad request for jane@example.com`, + "nested", map[string]any{"body": "jane@example.com"}, + ) + sanitized, ok := sanitizeErrorWithPIIMapping(err, m).(*domain.GatewayError) + if !ok { + t.Fatalf("expected GatewayError") + } + if strings.Contains(sanitized.Message, "jane@example.com") || sanitized.Metadata["upstream_error"] != `bad request for [EMAIL_1]` || + sanitized.Metadata["nested"].(map[string]any)["body"] != "[EMAIL_1]" { + t.Fatalf("not sanitized: %+v", sanitized) + } +} diff --git a/apps/gateway/internal/application/services/proxy.go b/apps/gateway/internal/application/services/proxy.go index 136218e..5690641 100644 --- a/apps/gateway/internal/application/services/proxy.go +++ b/apps/gateway/internal/application/services/proxy.go @@ -5,7 +5,6 @@ import ( "bytes" "context" "encoding/json" - "errors" "fmt" "io" "net/http" @@ -13,22 +12,17 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/middleware" "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" ) -// ErrNotProxyable means a request needs the translating path: PII masking -// rewrites content, and providers without the OpenAI wire format (Anthropic) -// need translation. It is returned before anything is reserved or sent. -var ErrNotProxyable = errors.New("request needs the translating path") - // maxSSELine bounds one SSE line (a whole tool call can arrive in one). const maxSSELine = 16 << 20 -// ProxyCall is a request the gateway proxies: the provider's response goes to -// the client as it came, except the model field, which names the public model. +// ProxyCall is a proxied request: the provider's response goes to the client +// as it came, except the model field (the public model) and, for keys that +// mask PII, the restored text. type ProxyCall struct { Model domain.PublicModel // Stream is set for streaming requests. @@ -37,9 +31,10 @@ type ProxyCall struct { Body []byte } -// Proxy runs a chat-completions request as a proxy. The gateway authenticates, -// routes, validates, reserves credit, and meters usage it reads from the -// response; it never rebuilds the request or the response. +// Proxy runs a chat-completions request. The gateway authenticates, routes, +// validates, masks PII for keys that ask for it, reserves credit, and meters +// the usage it reads from the response; it never rebuilds the request or the +// response. func (s *GenerateService) Proxy(ctx context.Context, endpoint string, raw []byte, req domain.GenerateRequest, bearerToken string) (*ProxyCall, error) { start := time.Now() requestID := middleware.GetRequestID(ctx) @@ -49,46 +44,49 @@ func (s *GenerateService) Proxy(ctx context.Context, endpoint string, raw []byte s.recordTerminalOutcome(metrics.OutcomeError, endpoint, req, nil, time.Since(start).Milliseconds()) return nil, err } - if s.masksPII(authCtx.APIKey.PIIMode) { - return nil, ErrNotProxyable - } model, execReq, err := s.resolveModel(ctx, req) if err != nil { s.recordTerminalOutcome(metrics.OutcomeError, endpoint, req, nil, time.Since(start).Milliseconds()) return nil, err } - if _, ok := s.proxyProvider(model); !ok { - return nil, ErrNotProxyable - } - if model.Fallback != nil { - if _, ok := s.proxyProvider(withProviderTarget(model, *model.Fallback)); !ok { - model.Fallback = nil // Only fall back to a provider that can proxy. - } - } if err := s.validateRequest(endpoint, execReq, model); err != nil { s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) s.recordFailure(ctx, nil, &authCtx, endpoint, &execReq, &model, err, nil, time.Since(start).Milliseconds()) return nil, err } + masked, mapping, err := s.maskBody(ctx, raw, authCtx.APIKey.PIIMode) + if err != nil { + s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) + s.recordFailure(ctx, nil, &authCtx, endpoint, &execReq, &model, err, nil, time.Since(start).Milliseconds()) + return nil, err + } reservationID, err := s.metering.Reserve(ctx, authCtx, endpoint, execReq, model, uuid.NewString()) if err != nil { s.recordTerminalOutcome(metrics.OutcomeError, endpoint, execReq, &model, time.Since(start).Milliseconds()) return nil, err } finish := func(target domain.PublicModel, o outcome) { + o.err = sanitizeErrorWithPIIMapping(o.err, mapping) s.finishProxy(ctx, authCtx, endpoint, execReq, target, requestID, reservationID, start, o) } + unmask := func(target domain.PublicModel) *unmasker { + return newUnmasker(mapping, s.logger, "request_id", requestID, "endpoint", endpoint, "model", execReq.PublicModelID, + "provider", target.ProviderConfig.ProviderName, "pii_mode", authCtx.APIKey.PIIMode) + } attempt := func(target domain.PublicModel) (*ProxyCall, error) { - provider, _ := s.proxyProvider(target) + provider, err := s.registry.GetProvider(target.ProviderConfig.ProviderName) + if err != nil { + return nil, domain.ErrProviderUnavailable(target.ProviderConfig.ProviderName) + } upstreamStart := time.Now() if !execReq.Stream { - body, proof, err := provider.ProxyJSON(ctx, raw, target) + body, proof, err := provider.Complete(ctx, masked, target) recordUpstreamLatency(execReq.PublicModelID, target.ProviderConfig.ProviderName, upstreamStart, err) if err != nil { return nil, err } - rewritten, o, err := observeJSON(body, target.PublicModelID) + rewritten, o, err := observeJSON(body, target.PublicModelID, unmask(target)) if err != nil { return nil, err } @@ -96,7 +94,7 @@ func (s *GenerateService) Proxy(ctx context.Context, endpoint string, raw []byte finish(target, o) return &ProxyCall{Model: target, Body: rewritten}, nil } - resp, err := provider.ProxyStream(ctx, raw, target) + resp, err := provider.Stream(ctx, masked, target) recordUpstreamLatency(execReq.PublicModelID, target.ProviderConfig.ProviderName, upstreamStart, err) if err != nil { return nil, err @@ -104,6 +102,7 @@ func (s *GenerateService) Proxy(ctx context.Context, endpoint string, raw []byte stream := newProxyStream(resp.Body, target.PublicModelID) stream.proof = resp.Proof stream.hideUsage = !wantsUsage(raw) + stream.unmask = unmask(target) // Nothing has reached the client yet: an upstream that fails before its // first output can still fall back. if err := stream.prime(); err != nil { @@ -118,12 +117,13 @@ func (s *GenerateService) Proxy(ctx context.Context, endpoint string, raw []byte executed := model if err != nil && shouldTryFallback(ctx, err, model.Fallback) { executed = withProviderTarget(model, *model.Fallback) - s.logFallback(requestID, model, executed, err) + s.logFallback(requestID, model, executed, sanitizeErrorWithPIIMapping(err, mapping)) call, err = attempt(executed) } if err != nil { + err = sanitizeErrorWithPIIMapping(err, mapping) fields := s.buildErrorLogFields(ctx, requestID, &authCtx, endpoint, execReq.PublicModelID, executed, err, time.Since(start).Milliseconds()) - s.logger.Error("provider proxy failed", fields...) + s.logger.Error("provider request failed", fields...) s.recordGeneration(metrics.OutcomeError, endpoint, execReq, executed, time.Since(start).Milliseconds()) s.recordFailure(ctx, &reservationID, &authCtx, endpoint, &execReq, &executed, err, nil, time.Since(start).Milliseconds()) return nil, err @@ -131,22 +131,6 @@ func (s *GenerateService) Proxy(ctx context.Context, endpoint string, raw []byte return call, nil } -// masksPII reports whether this key's requests have content masked: only -// when the key asks for it and the filter is on (it is off in production). -func (s *GenerateService) masksPII(piiMode string) bool { - mode, ok := domain.NormalizeAPIKeyPIIMode(piiMode) - return (!ok || mode != domain.APIKeyPIIModeOff) && s.pii != nil && s.pii.Enabled() -} - -func (s *GenerateService) proxyProvider(model domain.PublicModel) (ports.ProxyProvider, bool) { - provider, err := s.registry.GetProvider(model.ProviderConfig.ProviderName) - if err != nil { - return nil, false - } - proxy, ok := provider.(ports.ProxyProvider) - return proxy, ok -} - // outcome is what the gateway read from a proxied response. type outcome struct { providerID string @@ -157,7 +141,7 @@ type outcome struct { end string // How a stream ended, for logs. } -// finishProxy meters and logs a proxied response, like the translating path. +// finishProxy meters and logs a proxied response. func (s *GenerateService) finishProxy(ctx context.Context, authCtx domain.AuthContext, endpoint string, req domain.GenerateRequest, model domain.PublicModel, requestID, reservationID string, start time.Time, o outcome) { ctx = context.WithoutCancel(ctx) // Accounting outlives the client. latencyMs := time.Since(start).Milliseconds() @@ -166,10 +150,10 @@ func (s *GenerateService) finishProxy(ctx context.Context, authCtx domain.AuthCo "account_id", authCtx.Account.ID, "endpoint", endpoint, "model", req.PublicModelID, "provider", model.ProviderConfig.ProviderName, "provider_model", model.UpstreamModelName, "stream", req.Stream, "stream_end", o.end, "finish_reason", logfields.FinishReason(o.finish), - "usage_received", o.usage != nil, "usage", o.usage, "latency_ms", latencyMs, "proxy", true, + "usage_received", o.usage != nil, "usage", o.usage, "latency_ms", latencyMs, } if o.err != nil { - s.logger.Error("proxied generation failed", append(fields, "error_type", fmt.Sprintf("%T", o.err))...) + s.logger.Error("generation failed", append(fields, "error_type", fmt.Sprintf("%T", o.err))...) s.recordGeneration(metrics.OutcomeError, endpoint, req, model, latencyMs) s.recordFailure(ctx, &reservationID, &authCtx, endpoint, &req, &model, o.err, o.usage, latencyMs) return @@ -184,7 +168,7 @@ func (s *GenerateService) finishProxy(ctx context.Context, authCtx domain.AuthCo s.storeTinfoilProof(ctx, authCtx, model, result) metrics.RecordUsage(o.usage, req.PublicModelID, model.ProviderConfig.ProviderName) if o.usage == nil { - s.logger.Warn("proxied generation completed without usage", fields...) + s.logger.Warn("generation completed without usage", fields...) } s.recordGeneration(metrics.OutcomeSuccess, endpoint, req, model, latencyMs) s.logger.Info("generation completed", append(fields, "gateway_status", 200, "upstream_status", 200)...) @@ -307,7 +291,7 @@ func rewriteModel(raw []byte, c *chunk, public string) []byte { } // observeJSON reads a non-streaming response. -func observeJSON(body []byte, public string) ([]byte, outcome, error) { +func observeJSON(body []byte, public string, unmask *unmasker) ([]byte, outcome, error) { var c chunk if err := json.Unmarshal(body, &c); err != nil { return nil, outcome{}, domain.ErrProviderError(http.StatusBadGateway, "provider returned an invalid response") @@ -315,6 +299,10 @@ func observeJSON(body []byte, public string) ([]byte, outcome, error) { if err := c.failure(); err != nil { return nil, outcome{}, err } + if unmask != nil { + body = unmask.body(body) + unmask.report(false) + } return rewriteModel(body, &c, public), outcome{providerID: c.ID, finish: c.finishReason(), usage: c.Usage.domain(), end: "complete"}, nil } @@ -333,6 +321,8 @@ type ProxyStream struct { // gateway requests them to meter); usage is still read from them. hideUsage bool skipBlank bool + // unmask restores PII placeholders, for keys that mask. + unmask *unmasker o outcome sawDone bool @@ -366,9 +356,11 @@ func newProxyStream(body io.ReadCloser, public string) *ProxyStream { // fall back without the client seeing anything. func (p *ProxyStream) prime() error { for p.scanner.Scan() { - line, c := p.observe(p.scanner.Bytes()) - if p.keep(line, c) { - p.pending = append(p.pending, line) + lines, c := p.observe(p.scanner.Bytes()) + for _, line := range lines { + if p.keep(line, c) { + p.pending = append(p.pending, line) + } } if p.o.err != nil { return p.o.err @@ -386,22 +378,26 @@ func (p *ProxyStream) prime() error { return nil } -// observe reads one line; data events are parsed for metering and get the -// public model name. -func (p *ProxyStream) observe(raw []byte) ([]byte, *chunk) { +// observe reads one line; data events are parsed for metering, get the +// public model name, and have PII restored for keys that mask. It returns +// the lines to send: usually the one it read. +func (p *ProxyStream) observe(raw []byte) ([][]byte, *chunk) { line := append([]byte(nil), raw...) data, ok := bytes.CutPrefix(line, []byte("data:")) if !ok { - return line, nil + return [][]byte{line}, nil } data = bytes.TrimPrefix(data, []byte(" ")) if string(bytes.TrimSpace(data)) == "[DONE]" { p.sawDone = true - return line, nil + if tail := p.heldTail(); tail != nil { + return [][]byte{tail, {}, line}, nil + } + return [][]byte{line}, nil } var c chunk if json.Unmarshal(data, &c) != nil { - return line, nil + return [][]byte{line}, nil } p.chunks++ if c.ID != "" { @@ -417,7 +413,21 @@ func (p *ProxyStream) observe(raw []byte) ([]byte, *chunk) { p.o.err = err } prefix := line[:len(line)-len(data)] - return append(append([]byte(nil), prefix...), rewriteModel(data, &c, p.public)...), &c + if p.unmask != nil && len(c.Choices) > 0 { + data = p.unmask.event(data) + } + return [][]byte{append(append([]byte(nil), prefix...), rewriteModel(data, &c, p.public)...)}, &c +} + +// heldTail returns an event with PII-restored text still held back, if any. +func (p *ProxyStream) heldTail() []byte { + if p.unmask == nil { + return nil + } + if tail := p.unmask.tail(p.o.providerID, p.public); tail != nil { + return append([]byte("data: "), tail...) + } + return nil } // keep reports whether a line goes to the client: everything does, except a @@ -444,10 +454,20 @@ func (p *ProxyStream) Next() ([]byte, error) { return line, nil } for p.scanner.Scan() { - if line, c := p.observe(p.scanner.Bytes()); p.keep(line, c) { - return line, nil + lines, c := p.observe(p.scanner.Bytes()) + for _, line := range lines { + if p.keep(line, c) { + p.pending = append(p.pending, line) + } + } + if len(p.pending) > 0 { + return p.Next() } } + if tail := p.heldTail(); tail != nil { + p.pending = append(p.pending, tail, []byte{}) + return p.Next() + } err := p.scanner.Err() switch { case p.o.err != nil: @@ -498,6 +518,9 @@ func (p *ProxyStream) end(how string) { p.finished = true p.o.end = how p.o.proof = p.proof + if p.unmask != nil { + p.unmask.report(true) + } if p.onEnd != nil { p.onEnd(p.o) } diff --git a/apps/gateway/internal/application/services/proxy_test.go b/apps/gateway/internal/application/services/proxy_test.go index ee0b71b..2b6b4a5 100644 --- a/apps/gateway/internal/application/services/proxy_test.go +++ b/apps/gateway/internal/application/services/proxy_test.go @@ -2,7 +2,6 @@ package services import ( "context" - "errors" "io" "strings" "testing" @@ -11,30 +10,7 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) -// proxyStub is an OpenAI-compatible provider that returns canned bodies. -type proxyStub struct { - stubProvider - sse string - json string - err error - calls int - lastRaw string -} - -func (p *proxyStub) ProxyStream(_ context.Context, raw []byte, _ domain.PublicModel) (ports.ProxyResponse, error) { - p.calls++ - p.lastRaw = string(raw) - if p.err != nil { - return ports.ProxyResponse{}, p.err - } - return ports.ProxyResponse{Body: io.NopCloser(strings.NewReader(p.sse))}, nil -} - -func (p *proxyStub) ProxyJSON(_ context.Context, raw []byte, _ domain.PublicModel) ([]byte, *domain.TinfoilTransportProof, error) { - p.calls++ - p.lastRaw = string(raw) - return []byte(p.json), nil, p.err -} +type proxyStub = stubProvider // resultMeter records what was metered. type resultMeter struct { @@ -53,7 +29,7 @@ func (m *resultMeter) RecordFailure(ctx context.Context, id *string, a *domain.A return m.stubUsageMeter.RecordFailure(ctx, id, a, e, req, model, err, u, l) } -func proxyService(meter *resultMeter, pii string, providers map[string]ports.GenerationProvider, fallback bool) *GenerateService { +func proxyService(meter *resultMeter, pii string, providers map[string]ports.Provider, fallback bool) *GenerateService { model := domain.PublicModel{ PublicModelID: "zai-org/glm-5.3-flash", ProviderModelID: "pm", UpstreamModelName: "glm-5-3-flash", ProviderConfig: domain.ProviderConfig{ProviderName: "primary"}, @@ -104,7 +80,7 @@ func TestProxyStream_ForwardsProviderBytes(t *testing.T) { "data: [DONE]\n\n" primary := &proxyStub{sse: upstream} meter := &resultMeter{} - svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary}, false) + svc := proxyService(meter, "", map[string]ports.Provider{"primary": primary}, false) call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{"model":"zai-org/glm-5.3-flash","stream_options":{"include_usage":true}}`), streamRequest(), "key") if err != nil { t.Fatal(err) @@ -128,7 +104,7 @@ func TestProxyStream_FallsBackBeforeFirstOutput(t *testing.T) { primary := &proxyStub{sse: "data: {\"id\":\"x\",\"choices\":[{\"delta\":{\"role\":\"assistant\"}}]}\n\ndata: {\"error\":{\"message\":\"overloaded\"}}\n\n"} fallback := &proxyStub{sse: "data: {\"id\":\"y\",\"model\":\"glm-fallback\",\"choices\":[{\"delta\":{\"content\":\"hi\"},\"finish_reason\":\"stop\"}]}\n\ndata: [DONE]\n\n"} meter := &resultMeter{} - svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary, "fallback": fallback}, true) + svc := proxyService(meter, "", map[string]ports.Provider{"primary": primary, "fallback": fallback}, true) call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key") if err != nil { t.Fatal(err) @@ -145,7 +121,7 @@ func TestProxyStream_FallsBackBeforeFirstOutput(t *testing.T) { func TestProxyStream_ErrorAfterOutputIsForwardedAndNotBilledAsSuccess(t *testing.T) { primary := &proxyStub{sse: "data: {\"choices\":[{\"delta\":{\"content\":\"partial\"}}]}\n\ndata: {\"error\":{\"message\":\"overloaded\"}}\n\n"} meter := &resultMeter{} - svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary}, false) + svc := proxyService(meter, "", map[string]ports.Provider{"primary": primary}, false) call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key") if err != nil { t.Fatal(err) @@ -162,7 +138,7 @@ func TestProxyStream_ErrorAfterOutputIsForwardedAndNotBilledAsSuccess(t *testing func TestProxyStream_EmptyStreamIsAProviderError(t *testing.T) { primary := &proxyStub{sse: ""} meter := &resultMeter{} - svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary}, false) + svc := proxyService(meter, "", map[string]ports.Provider{"primary": primary}, false) if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key"); err == nil { t.Fatal("empty stream accepted") } @@ -176,7 +152,7 @@ func TestProxyStream_ClientLeaving(t *testing.T) { for name, lines := range map[string]int{"before finish": 1, "after finish": 3} { t.Run(name, func(t *testing.T) { meter := &resultMeter{} - svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": &proxyStub{sse: body}}, false) + svc := proxyService(meter, "", map[string]ports.Provider{"primary": &proxyStub{sse: body}}, false) call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key") if err != nil { t.Fatal(err) @@ -198,7 +174,7 @@ func TestProxyStream_ClientLeaving(t *testing.T) { func TestProxyJSON(t *testing.T) { primary := &proxyStub{json: `{"id":"r1","model":"glm-5-3-flash","choices":[{"message":{"role":"assistant","content":"hi","reasoning_content":"why"},"finish_reason":"stop"}],"usage":{"prompt_tokens":2,"completion_tokens":1}}`} meter := &resultMeter{} - svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": primary}, false) + svc := proxyService(meter, "", map[string]ports.Provider{"primary": primary}, false) call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{PublicModelID: "zai-org/glm-5.3-flash"}, "key") if err != nil { t.Fatal(err) @@ -215,26 +191,11 @@ func TestProxyJSON(t *testing.T) { } } -func TestProxy_PIIKeysAndOtherWireFormatsTranslate(t *testing.T) { - meter := &resultMeter{} - svc := proxyService(meter, "balanced", map[string]ports.GenerationProvider{"primary": &proxyStub{}}, false) - if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key"); !errors.Is(err, ErrNotProxyable) { - t.Fatalf("PII key proxied: %v", err) - } - svc = proxyService(meter, "", map[string]ports.GenerationProvider{"primary": &stubProvider{}}, false) - if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key"); !errors.Is(err, ErrNotProxyable) { - t.Fatalf("non-OpenAI provider proxied: %v", err) - } - if meter.reserveCalls != 0 { - t.Fatal("reserved before choosing the path") - } -} - // With the PII filter off (as in production), a key's PII mode masks // nothing, so its requests are proxied too. func TestProxy_PIIModeWithFilterOffProxies(t *testing.T) { meter := &resultMeter{} - svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": &proxyStub{json: `{"id":"r","choices":[]}`}}, false) + svc := proxyService(meter, "", map[string]ports.Provider{"primary": &proxyStub{json: `{"id":"r","choices":[]}`}}, false) svc.auth = &stubAuthService{authCtx: domain.AuthContext{Account: domain.Account{ID: "acc1", Status: domain.AccountStatusActive}, APIKey: domain.APIKey{ID: "key1", Active: true, PIIMode: "high"}}} if _, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), domain.GenerateRequest{PublicModelID: "zai-org/glm-5.3-flash"}, "key"); err != nil { t.Fatalf("not proxied: %v", err) @@ -246,7 +207,7 @@ func TestProxy_PIIModeWithFilterOffProxies(t *testing.T) { func TestProxyStream_UnrequestedUsageIsMeteredNotForwarded(t *testing.T) { body := "data: {\"choices\":[{\"delta\":{\"content\":\"a\"},\"finish_reason\":\"stop\"}]}\n\ndata: {\"choices\":[],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":1}}\n\ndata: [DONE]\n\n" meter := &resultMeter{} - svc := proxyService(meter, "", map[string]ports.GenerationProvider{"primary": &proxyStub{sse: body}}, false) + svc := proxyService(meter, "", map[string]ports.Provider{"primary": &proxyStub{sse: body}}, false) call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{"stream":true}`), streamRequest(), "key") if err != nil { t.Fatal(err) diff --git a/apps/gateway/internal/application/services/stream_logging_test.go b/apps/gateway/internal/application/services/stream_logging_test.go deleted file mode 100644 index d071104..0000000 --- a/apps/gateway/internal/application/services/stream_logging_test.go +++ /dev/null @@ -1,119 +0,0 @@ -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/apps/gateway/main.go b/apps/gateway/main.go index eb50bbd..f5fc99d 100644 --- a/apps/gateway/main.go +++ b/apps/gateway/main.go @@ -18,7 +18,6 @@ import ( meteringadapter "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/metering" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/observability/metrics" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/pii/presidio" - "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/providers/anthropic" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/providers/openai" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/providers/registry" tinfoilprovider "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/providers/tinfoil" @@ -68,10 +67,9 @@ func main() { ctx := context.Background() providerRegistry := registry.NewRegistry() - providerRegistry.Register("anthropic", anthropic.NewAdapter(providerTimeout)) providerRegistry.Register("tinfoil", tinfoilprovider.NewAdapter(providerTimeout, zapLogger)) - // Any provider not explicitly registered falls back to the OpenAI-compatible adapter. - // This allows adding new providers (e.g. novita, mistral) via DB only — no code changes. + // Every other provider speaks the OpenAI wire format, so new ones (e.g. + // novita, mistral) are added in the database only. providerRegistry.SetDefault(openai.NewAdapter(providerTimeout, zapLogger)) meteringClient := meteringadapter.NewClient(meteringURL, meteringToken, 5*time.Second) diff --git a/pkg/domain/account.go b/pkg/domain/account.go index 045c728..3a50b99 100644 --- a/pkg/domain/account.go +++ b/pkg/domain/account.go @@ -96,20 +96,6 @@ func NormalizeAPIKeyPIIMode(mode string) (string, bool) { } } -// CreateAPIKeyParams holds parameters for creating a new API key. -type CreateAPIKeyParams struct { - AccountID string - Name *string - PIIMode string - ExpiresAt *time.Time -} - -// CreateAPIKeyResult is returned after key creation, including the raw key shown only once. -type CreateAPIKeyResult struct { - APIKey APIKey - RawKey string -} - // AuthContext carries the authenticated account and API key for a request. type AuthContext struct { Account Account diff --git a/pkg/domain/domain_test.go b/pkg/domain/domain_test.go index af09df7..1479622 100644 --- a/pkg/domain/domain_test.go +++ b/pkg/domain/domain_test.go @@ -31,16 +31,12 @@ func TestErrorConstructors(t *testing.T) { wantCode string }{ {"InvalidAPIKey", domain.ErrInvalidAPIKey("bad"), 401, domain.ErrTypeAuthentication, domain.ErrCodeInvalidAPIKey}, - {"InactiveAPIKey", domain.ErrInactiveAPIKey(), 401, domain.ErrTypeAuthentication, domain.ErrCodeInactiveAPIKey}, - {"InactiveAccount", domain.ErrInactiveAccount(), 403, domain.ErrTypePermission, domain.ErrCodeInactiveAccount}, {"UnsupportedModel", domain.ErrUnsupportedModel("foo"), 404, domain.ErrTypeInvalidRequest, domain.ErrCodeUnsupportedModel}, {"UnsupportedEndpoint", domain.ErrUnsupportedEndpoint("m", "e"), 422, domain.ErrTypeInvalidRequest, domain.ErrCodeUnsupportedEndpoint}, {"UnsupportedFeature", domain.ErrUnsupportedFeature("f"), 422, domain.ErrTypeInvalidRequest, domain.ErrCodeUnsupportedFeature}, {"ProviderUnavailable", domain.ErrProviderUnavailable("p"), 503, domain.ErrTypeProvider, domain.ErrCodeProviderUnavailable}, {"ProviderTimeout", domain.ErrProviderTimeout("p"), 504, domain.ErrTypeProvider, domain.ErrCodeProviderTimeout}, {"InvalidField", domain.ErrInvalidField("f"), 400, domain.ErrTypeInvalidRequest, domain.ErrCodeInvalidField}, - {"InsufficientBalance", domain.ErrInsufficientBalance(), 402, domain.ErrTypePermission, domain.ErrCodeInsufficientBalance}, - {"UnknownField", domain.ErrUnknownField("f"), 400, domain.ErrTypeInvalidRequest, domain.ErrCodeUnknownField}, {"Internal", domain.ErrInternal("i"), 500, domain.ErrTypeInternal, domain.ErrCodeInternalError}, } diff --git a/pkg/domain/error.go b/pkg/domain/error.go index 048a3b0..0382d9a 100644 --- a/pkg/domain/error.go +++ b/pkg/domain/error.go @@ -1,7 +1,6 @@ package domain import ( - "errors" "fmt" ) @@ -86,14 +85,6 @@ func ErrInvalidAPIKey(msg string) *GatewayError { return &GatewayError{HTTPStatus: 401, Type: ErrTypeAuthentication, Code: ErrCodeInvalidAPIKey, Message: msg} } -func ErrInactiveAPIKey() *GatewayError { - return &GatewayError{HTTPStatus: 401, Type: ErrTypeAuthentication, Code: ErrCodeInactiveAPIKey, Message: "API key is inactive"} -} - -func ErrInactiveAccount() *GatewayError { - return &GatewayError{HTTPStatus: 403, Type: ErrTypePermission, Code: ErrCodeInactiveAccount, Message: "Account is inactive"} -} - func ErrUnsupportedModel(model string) *GatewayError { return &GatewayError{HTTPStatus: 404, Type: ErrTypeInvalidRequest, Code: ErrCodeUnsupportedModel, Message: fmt.Sprintf("Model '%s' is not available", model)} } @@ -122,18 +113,6 @@ func ErrInvalidField(msg string) *GatewayError { return &GatewayError{HTTPStatus: 400, Type: ErrTypeInvalidRequest, Code: ErrCodeInvalidField, Message: msg} } -func ErrInsufficientBalance() *GatewayError { - return &GatewayError{HTTPStatus: 402, Type: ErrTypePermission, Code: ErrCodeInsufficientBalance, Message: "Account has no spendable balance"} -} - -func ErrUnknownField(field string) *GatewayError { - return &GatewayError{HTTPStatus: 400, Type: ErrTypeInvalidRequest, Code: ErrCodeUnknownField, Message: fmt.Sprintf("Unknown field: '%s'", field)} -} - -func ErrToolSchemaInvalid(msg string) *GatewayError { - return &GatewayError{HTTPStatus: 400, Type: ErrTypeInvalidRequest, Code: ErrCodeToolSchemaInvalid, Message: msg} -} - func ErrToolMessageInvalid(msg string) *GatewayError { return &GatewayError{HTTPStatus: 400, Type: ErrTypeInvalidRequest, Code: ErrCodeToolMessageInvalid, Message: msg} } @@ -154,27 +133,12 @@ const ( ErrCodeAlreadyExists = "already_exists" ) -func ErrNotFound(resource, id string) *GatewayError { - return &GatewayError{HTTPStatus: 404, Type: ErrTypeInvalidRequest, Code: ErrCodeNotFound, Message: fmt.Sprintf("%s '%s' not found", resource, id)} -} - -func ErrForbidden(msg string) *GatewayError { - return &GatewayError{HTTPStatus: 403, Type: ErrTypePermission, Code: ErrCodeForbidden, Message: msg} -} - func ErrConflict(msg string) *GatewayError { return &GatewayError{HTTPStatus: 409, Type: ErrTypeInvalidRequest, Code: ErrCodeConflict, Message: msg} } - -func ErrAlreadyExists(resource, id string) *GatewayError { - return &GatewayError{HTTPStatus: 409, Type: ErrTypeInvalidRequest, Code: ErrCodeAlreadyExists, Message: fmt.Sprintf("%s '%s' already exists", resource, id)} +func ErrNotFound(resource, id string) *GatewayError { + return &GatewayError{HTTPStatus: 404, Type: ErrTypeInvalidRequest, Code: ErrCodeNotFound, Message: fmt.Sprintf("%s '%s' not found", resource, id)} } - -// IsAlreadyExists returns true if the error is a GatewayError with code already_exists. -func IsAlreadyExists(err error) bool { - var gwErr *GatewayError - if errors.As(err, &gwErr) { - return gwErr.Code == ErrCodeAlreadyExists - } - return false +func ErrInsufficientBalance() *GatewayError { + return &GatewayError{HTTPStatus: 402, Type: ErrTypePermission, Code: ErrCodeInsufficientBalance, Message: "Account has no spendable balance"} } diff --git a/pkg/domain/generation.go b/pkg/domain/generation.go index 6b84845..84e02d5 100644 --- a/pkg/domain/generation.go +++ b/pkg/domain/generation.go @@ -5,60 +5,51 @@ const ( EndpointChatCompletions = "chat_completions" ) -// ResponseTextConfig configures structured text output. -type ResponseTextConfig struct { - FormatType *string - JSONSchema map[string]any -} - -// GenerateRequest is the canonical internal generation request. +// GenerateRequest is what the gateway reads from a client's request: enough +// to route, check the model's features, and meter. The request body itself is +// forwarded to the provider as sent. type GenerateRequest struct { - PublicModelID string - RequestedModelID string - RouterID *string - RoutedPublicModelID *string + PublicModelID string + RequestedModelID string // Routing decision metadata. Set when the request was resolved via a // router; nil for direct model requests. + RouterID *string + RoutedPublicModelID *string MatchedCategory *string RoutingScore *float32 RoutingCategoryScores []RoutingCategoryScore DecisionReason *string FallbackUsed *bool - Input []InputItem - Instructions *string - MaxOutputTokens *int - Temperature *float64 - ReasoningEffort *string - TopP *float64 - Stop []string - Stream bool - Tools []ToolDefinition - ToolChoice *ToolChoice - ParallelToolCalls *bool - User *string - Metadata map[string]any - TextConfig *ResponseTextConfig - ProviderOptions map[string]any - // Pass-through parameters (forwarded to providers that support them) - PresencePenalty *float64 `json:"presence_penalty,omitempty"` - FrequencyPenalty *float64 `json:"frequency_penalty,omitempty"` - LogitBias map[string]int `json:"logit_bias,omitempty"` - Seed *int `json:"seed,omitempty"` - Logprobs *bool `json:"logprobs,omitempty"` - TopLogprobs *int `json:"top_logprobs,omitempty"` - Store *bool `json:"store,omitempty"` - ServiceTier *string `json:"service_tier,omitempty"` + // Input and Tools are what routers read: each message's role and text, + // and the tool names. + Input []InputItem + Tools []ToolDefinition + + Stream bool + MaxOutputTokens *int + ParallelToolCalls *bool + // StructuredOutput is set when the client asks for JSON output. + StructuredOutput bool +} + +// InputItem is one message: its role and text. +type InputItem struct { + Role *string + Content *string +} + +// ToolDefinition names a function tool offered to the model. +type ToolDefinition struct { + Name string } -// GenerateResult is the canonical internal generation result. +// GenerateResult is what the gateway read from a provider's response. type GenerateResult struct { ID string - CreatedUnix int64 PublicModelID string ProviderName string ProviderModelID string - Output []OutputItem FinishReason *string Usage *Usage TinfoilProof *TinfoilTransportProof diff --git a/pkg/domain/input.go b/pkg/domain/input.go deleted file mode 100644 index b31478b..0000000 --- a/pkg/domain/input.go +++ /dev/null @@ -1,17 +0,0 @@ -package domain - -// InputItemType constants. -const ( - InputItemTypeMessage = "message" - InputItemTypeText = "text" -) - -// InputItem is a canonical input item for generation requests. -type InputItem struct { - Type string // "message" or "text" - Role *string - Content *string - ReasoningContent *string - ToolCalls []ToolCall - ToolCallID *string -} diff --git a/pkg/domain/output.go b/pkg/domain/output.go deleted file mode 100644 index 2f87b72..0000000 --- a/pkg/domain/output.go +++ /dev/null @@ -1,16 +0,0 @@ -package domain - -// OutputItemType constants. -const ( - OutputItemTypeMessage = "message" - OutputItemTypeText = "text" -) - -// OutputItem is a canonical output item from a generation result. -type OutputItem struct { - Type string // "message" or "text" - Role *string - Content *string - ReasoningContent *string - ToolCalls []ToolCall -} diff --git a/pkg/domain/pricing.go b/pkg/domain/pricing.go deleted file mode 100644 index a03a20d..0000000 --- a/pkg/domain/pricing.go +++ /dev/null @@ -1,14 +0,0 @@ -package domain - -import "time" - -// ModelPricing represents pricing for a provider model at a point in time. -type ModelPricing struct { - ID string - ProviderModelID string - InputPricePer1MTokensMicrocents int64 - OutputPricePer1MTokensMicrocents int64 - CacheReadPricePer1MTokensMicrocents *int64 - CacheWritePricePer1MTokensMicrocents *int64 - EffectiveFrom time.Time -} diff --git a/pkg/domain/promo.go b/pkg/domain/promo.go deleted file mode 100644 index 51d143e..0000000 --- a/pkg/domain/promo.go +++ /dev/null @@ -1,127 +0,0 @@ -package domain - -import ( - "regexp" - "strings" - "time" -) - -// Promo code configuration constants. -const ( - PromoCodeMinLength = 4 - PromoCodeMaxLength = 64 - - // NewUserGraceWindow is how long after account creation a "new users only" - // promo can still be redeemed. The auto-redeem path fires immediately after - // Authgear signup (account is auto-provisioned moments before), but a slack - // window keeps the experience resilient to clock skew, slow UI loads, and - // manual "redeem on Billing" attempts by a brand-new user. - NewUserGraceWindow = 24 * time.Hour -) - -// promoCodePattern matches normalized promo codes: uppercase letters, digits, -// underscores and hyphens. -var promoCodePattern = regexp.MustCompile(`^[A-Z0-9_-]+$`) - -// PromoCode represents an operator-issued code that grants prepaid credit. -type PromoCode struct { - ID string - Code string - AmountCents int64 - Currency string - MaxRedemptions *int - NewUsersOnly bool - Active bool - ExpiresAt *time.Time - Description *string - CreatedAt time.Time - UpdatedAt time.Time - // RedemptionCount is the number of redemptions recorded for this code. - // Populated only by list/get queries that aggregate redemptions; zero elsewhere. - RedemptionCount int -} - -// RemainingRedemptions reports how many more redemptions are allowed, or nil -// when MaxRedemptions is unset (unlimited). -func (p PromoCode) RemainingRedemptions() *int { - if p.MaxRedemptions == nil { - return nil - } - remaining := *p.MaxRedemptions - p.RedemptionCount - if remaining < 0 { - remaining = 0 - } - return &remaining -} - -// IsExpiredAt reports whether the promo code is past its expiry at the given time. -func (p PromoCode) IsExpiredAt(now time.Time) bool { - return p.ExpiresAt != nil && !now.Before(*p.ExpiresAt) -} - -// PromoRedemption records a single successful redemption of a promo code. -type PromoRedemption struct { - ID string - PromoCodeID string - AccountID string - AccountName *string - AmountMicrocents int64 - CreatedAt time.Time -} - -// RedeemResult is returned by the redemption flow. Applied is false when the -// account had already redeemed this code (idempotent success) — in that case -// Redemption holds the original redemption row. -type RedeemResult struct { - Applied bool - Redemption PromoRedemption - PromoCode PromoCode -} - -// NormalizePromoCode uppercases and trims the input and validates its shape. -// Returns the normalized code and whether it is valid. -func NormalizePromoCode(raw string) (string, bool) { - code := strings.ToUpper(strings.TrimSpace(raw)) - if len(code) < PromoCodeMinLength || len(code) > PromoCodeMaxLength { - return "", false - } - if !promoCodePattern.MatchString(code) { - return "", false - } - return code, true -} - -// IsNewUserAt reports whether accountCreatedAt falls inside the new-user grace -// window ending at now. Used to enforce PromoCode.NewUsersOnly. -func IsNewUserAt(accountCreatedAt, now time.Time) bool { - return now.Before(accountCreatedAt.Add(NewUserGraceWindow)) -} - -// Promo error constructors. These map to the control-plane error codes. -// 410 Gone signals "this code no longer accepts redemptions" (deactivated or -// expired); 409 Conflict signals "you already redeemed this code" or "you are -// not eligible"; 422 signals the cap was hit. - -func ErrPromoNotFound(code string) *GatewayError { - return &GatewayError{HTTPStatus: 404, Type: ErrTypeInvalidRequest, Code: ErrCodeNotFound, Message: "promo code '" + code + "' not found"} -} - -func ErrPromoInactive(code string) *GatewayError { - return &GatewayError{HTTPStatus: 410, Type: ErrTypeInvalidRequest, Code: ErrCodeConflict, Message: "promo code '" + code + "' is no longer active"} -} - -func ErrPromoExpired(code string) *GatewayError { - return &GatewayError{HTTPStatus: 410, Type: ErrTypeInvalidRequest, Code: ErrCodeConflict, Message: "promo code '" + code + "' has expired"} -} - -func ErrPromoMaxRedemptionsReached(code string) *GatewayError { - return &GatewayError{HTTPStatus: 422, Type: ErrTypeInvalidRequest, Code: ErrCodeConflict, Message: "promo code '" + code + "' has reached its maximum redemptions"} -} - -func ErrPromoNotEligibleNewUsersOnly(code string) *GatewayError { - return &GatewayError{HTTPStatus: 409, Type: ErrTypeInvalidRequest, Code: ErrCodeConflict, Message: "promo code '" + code + "' is only available to new users"} -} - -func ErrPromoAlreadyRedeemed(code string) *GatewayError { - return &GatewayError{HTTPStatus: 409, Type: ErrTypeInvalidRequest, Code: ErrCodeConflict, Message: "promo code '" + code + "' has already been redeemed"} -} diff --git a/pkg/domain/promo_test.go b/pkg/domain/promo_test.go deleted file mode 100644 index 95f5692..0000000 --- a/pkg/domain/promo_test.go +++ /dev/null @@ -1,95 +0,0 @@ -package domain - -import ( - "testing" - "time" -) - -func TestNormalizePromoCode(t *testing.T) { - cases := []struct { - in string - want string - wantOk bool - }{ - {"WELCOME5", "WELCOME5", true}, - {" welcome5 ", "WELCOME5", true}, - {"Welcome-5", "WELCOME-5", true}, - {"BONUS_CODE-3", "BONUS_CODE-3", true}, - {"ABC", "", false}, // too short - {"", "", false}, // empty - {"WITH SPACE", "", false}, // space not allowed - {"BAD!CODE", "", false}, // invalid char - {"LOWER1234", "LOWER1234", true}, - {"a_b-c1234", "A_B-C1234", true}, - } - for _, c := range cases { - got, ok := NormalizePromoCode(c.in) - if got != c.want || ok != c.wantOk { - t.Errorf("NormalizePromoCode(%q) = (%q,%v), want (%q,%v)", c.in, got, ok, c.want, c.wantOk) - } - } -} - -func TestNormalizePromoCode_TooLong(t *testing.T) { - long := "" - for i := 0; i < PromoCodeMaxLength+1; i++ { - long += "A" - } - if _, ok := NormalizePromoCode(long); ok { - t.Errorf("NormalizePromoCode accepted a code longer than %d chars", PromoCodeMaxLength) - } -} - -func TestPromoCodeRemainingRedemptions(t *testing.T) { - // unlimited - unlimited := PromoCode{MaxRedemptions: nil, RedemptionCount: 5} - if r := unlimited.RemainingRedemptions(); r != nil { - t.Errorf("unlimited code returned remaining %v, want nil", *r) - } - - // 100 / 5 used → 95 - max := 100 - p := PromoCode{MaxRedemptions: &max, RedemptionCount: 5} - r := p.RemainingRedemptions() - if r == nil || *r != 95 { - t.Errorf("remaining = %v, want 95", r) - } - - // exhausted clamps to 0 - exhausted := PromoCode{MaxRedemptions: &max, RedemptionCount: 150} - r = exhausted.RemainingRedemptions() - if r == nil || *r != 0 { - t.Errorf("remaining = %v, want 0", r) - } -} - -func TestPromoCodeIsExpiredAt(t *testing.T) { - expiry := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) - p := PromoCode{ExpiresAt: &expiry} - if p.IsExpiredAt(expiry) != true { - t.Error("code should be expired at expiry time") - } - if p.IsExpiredAt(expiry.Add(-time.Second)) != false { - t.Error("code should not be expired one second before expiry") - } - noExpiry := PromoCode{} - if noExpiry.IsExpiredAt(time.Now()) != false { - t.Error("code with nil expiry should never be expired") - } -} - -func TestIsNewUserAt(t *testing.T) { - created := time.Date(2026, 6, 22, 12, 0, 0, 0, time.UTC) - // Within window - if !IsNewUserAt(created, created.Add(time.Hour)) { - t.Error("1h after creation should be new user") - } - // Exactly at window boundary → no longer new (exclusive) - if IsNewUserAt(created, created.Add(NewUserGraceWindow)) { - t.Error("at grace boundary should not be new user") - } - // Well after window - if IsNewUserAt(created, created.Add(48*time.Hour)) { - t.Error("48h after creation should not be new user") - } -} diff --git a/pkg/domain/provider_catalog.go b/pkg/domain/provider_catalog.go deleted file mode 100644 index bd4cf27..0000000 --- a/pkg/domain/provider_catalog.go +++ /dev/null @@ -1,16 +0,0 @@ -package domain - -// ProviderCatalogEntry is one model offered by a provider, as discovered from -// the provider's own model-listing API. Prices are the provider's raw rate in -// USD per 1M tokens, before any FX conversion or markup. -type ProviderCatalogEntry struct { - ProviderModelName string - Title string - Description string - ContextSize int64 - MaxOutputTokens int64 - InputUSDPer1M float64 - OutputUSDPer1M float64 - SupportsTools bool - SupportsReasoning bool -} diff --git a/pkg/domain/provider_model.go b/pkg/domain/provider_model.go deleted file mode 100644 index 18b64f3..0000000 --- a/pkg/domain/provider_model.go +++ /dev/null @@ -1,12 +0,0 @@ -package domain - -import "time" - -// ProviderModel represents an upstream provider's model configuration. -type ProviderModel struct { - ID string - ProviderName string - ProviderModelName string - Active bool - CreatedAt time.Time -} diff --git a/pkg/domain/public_model_cp.go b/pkg/domain/public_model_cp.go deleted file mode 100644 index 266a70b..0000000 --- a/pkg/domain/public_model_cp.go +++ /dev/null @@ -1,29 +0,0 @@ -package domain - -import "time" - -// PublicModelEntry represents a public model as managed by the control plane. -type PublicModelEntry struct { - ID string - DisplayName string - ProviderModelID string - Active bool - Description *string - MaxContextWindow int - MaxOutputTokens int - SupportsChatCompletions bool - SupportsChatCompletionsStream bool - SupportsTools bool - SupportsParallelToolCalls bool - SupportsStructuredOutput bool - SupportsReasoning bool - ProofMode string - CreatedAt time.Time -} - -func (m PublicModelEntry) EffectiveProofMode() string { - if m.ProofMode != "" && m.ProofMode != ProofModeNone { - return m.ProofMode - } - return ProofModeNone -} diff --git a/pkg/domain/router.go b/pkg/domain/router.go index aa4b499..41260b6 100644 --- a/pkg/domain/router.go +++ b/pkg/domain/router.go @@ -74,21 +74,3 @@ type ModelCatalogEntry struct { Router *RouterEntry EURToUSDRate float64 } - -// RouterCategory is a strategy-specific routing target managed by the router -// service. The control plane treats categories as opaque metadata it can -// proxy on behalf of the admin UI. -type RouterCategory struct { - ID string - RouterID string - Name string - PublicModelID string - Threshold float32 -} - -// RouterCategoryInput is the create/update payload for a router category. -type RouterCategoryInput struct { - Name string - PublicModelID string - Threshold float32 -} diff --git a/pkg/domain/stream.go b/pkg/domain/stream.go deleted file mode 100644 index 062fc7c..0000000 --- a/pkg/domain/stream.go +++ /dev/null @@ -1,24 +0,0 @@ -package domain - -// StreamEventType constants. -const ( - StreamEventOutputTextDelta = "output_text_delta" - StreamEventOutputMessageDelta = "output_message_delta" - StreamEventToolCallDelta = "tool_call_delta" - StreamEventCompleted = "completed" - StreamEventError = "error" -) - -// StreamEvent is a canonical stream event normalized from provider stream formats. -type StreamEvent struct { - Type string // output_text_delta, output_message_delta, tool_call_delta, completed, error - ProviderResponseID string - ChoiceIndex *int - Role *string - ContentDelta *string - ReasoningDelta *string - ToolCallDelta *ToolCallDelta - FinishReason *string - Usage *Usage - Error *GatewayError -} diff --git a/pkg/domain/tee.go b/pkg/domain/tee.go index cbfb084..d46c537 100644 --- a/pkg/domain/tee.go +++ b/pkg/domain/tee.go @@ -38,19 +38,3 @@ type TinfoilTransportProof struct { CreatedAt time.Time VerifiedAt *time.Time } - -// TinfoilProofListParams describes a user-facing proof history query. -type TinfoilProofListParams struct { - Offset int - Limit int - Status string - Query string -} - -// TinfoilTransportProofRecord enriches a proof with safe API key context for -// dashboard history views. It never contains raw API key material. -type TinfoilTransportProofRecord struct { - Proof TinfoilTransportProof - APIKeyName *string - APIKeyPrefix *string -} diff --git a/pkg/domain/tool.go b/pkg/domain/tool.go index d5ca835..fd4d72e 100644 --- a/pkg/domain/tool.go +++ b/pkg/domain/tool.go @@ -1,38 +1,9 @@ package domain -// ToolDefinition describes a function tool available to the model. -type ToolDefinition struct { - Name string - Description string - Parameters map[string]any - Strict bool -} - -// ToolCall represents a tool invocation returned by the model. -type ToolCall struct { - ID string - Name string - ArgumentsJSON string -} - -// ToolCallDelta represents a partial tool call in a stream event. -type ToolCallDelta struct { - Index int - ID *string - Name *string - ArgumentsDelta *string -} - -// ToolChoiceMode constants. +// tool_choice modes. const ( ToolChoiceNone = "none" ToolChoiceAuto = "auto" ToolChoiceRequired = "required" ToolChoiceFunction = "function" ) - -// ToolChoice represents the tool_choice parameter. -type ToolChoice struct { - Mode string // "none", "auto", "required", "function" - FunctionName *string // set only when Mode == "function" -} diff --git a/pkg/domain/usage_query.go b/pkg/domain/usage_query.go deleted file mode 100644 index bbe2a84..0000000 --- a/pkg/domain/usage_query.go +++ /dev/null @@ -1,80 +0,0 @@ -package domain - -import "time" - -// UsageEvent represents a single usage event as stored in the database. -type UsageEvent struct { - ID string - RequestID *string - AccountID string - APIKeyID string - PublicModelID *string - ProviderModelID *string - RouterID *string - RoutedPublicModelID *string - PublicModelName *string - ProviderName *string - ProviderRequestID *string - InputTokens *int64 - OutputTokens *int64 - CacheCreationTokens *int64 - CacheReadTokens *int64 - InputPricePer1MTokensMicrocents *int64 - OutputPricePer1MTokensMicrocents *int64 - CacheReadPricePer1MTokensMicrocents *int64 - CacheWritePricePer1MTokensMicrocents *int64 - CostMicrocents *int64 - ChargedMonthlyMicrocents *int64 - ChargedPrepaidMicrocents *int64 - LatencyMs *int64 - CreatedAt time.Time - // AccountEmail is populated by list queries that JOIN accounts; nil elsewhere. - AccountEmail *string -} - -// UsageSummary holds aggregated usage data for an account or globally. -type UsageSummary struct { - TotalRequests int64 - TotalInputTokens int64 - TotalOutputTokens int64 - TotalCostMicrocents int64 - AvgLatencyMs float64 - From time.Time - To time.Time -} - -// UsageTimeBucket holds usage data aggregated for a single time bucket. -type UsageTimeBucket struct { - Bucket time.Time - TotalRequests int64 - TotalInputTokens int64 - TotalOutputTokens int64 - TotalCostMicrocents int64 -} - -// UsageEventRoutingDetails is the minimal projection used to surface the -// router decision behind a single usage event. RoutingScore is snapshotted -// at decision time. Threshold is the current configured threshold of the -// matched category (looked up at read time, not snapshotted), and is nil -// when no category matched or the category was since deleted. -type UsageEventRoutingDetails struct { - EventID string - RouterID *string - MatchedCategory *string - RoutingScore *float32 - CategoryScores []RoutingCategoryScore - MatchedThreshold *float32 - DecisionReason *string - FallbackUsed *bool -} - -// UsageByDimension holds usage data aggregated by an arbitrary grouping key -// (e.g. model name, provider name, or account ID). -type UsageByDimension struct { - Key string - Label string - TotalRequests int64 - TotalInputTokens int64 - TotalOutputTokens int64 - TotalCostMicrocents int64 -} From 2d4909dff2980a5a91845cc335f825ebcb6fe769 Mon Sep 17 00:00:00 2001 From: Marketen Date: Tue, 29 Sep 2026 11:28:55 +0200 Subject: [PATCH 5/7] fix: drop provider workarounds the providers no longer need Verified against each provider's API on 2026-09-29: - Novita honors max_completion_tokens on every production model, so the gateway no longer renames it to max_tokens (DeepSeek still needs it: it silently ignores max_completion_tokens). - Novita and DeepSeek accept assistant tool-call messages without content, so content:null is no longer added. - No Novita model rejects parallel_tool_calls, store, service_tier, user, or tool_choice=auto, so the downgrade retry that dropped them is removed. The Novita same-body retries (generic 400 with trace id, 429 overload), the DeepSeek reasoning_content fill, and developer-as-system stay. Co-Authored-By: Claude Opus 5.5 --- .../adapters/providers/openai/adapter.go | 16 +-- .../adapters/providers/openai/logging.go | 3 - .../adapters/providers/openai/mapper.go | 85 +++------------- .../adapters/providers/openai/policy_test.go | 97 +++---------------- .../adapters/providers/openai/proxy.go | 28 +----- .../adapters/providers/openai/proxy_test.go | 4 +- 6 files changed, 36 insertions(+), 197 deletions(-) diff --git a/apps/gateway/internal/adapters/providers/openai/adapter.go b/apps/gateway/internal/adapters/providers/openai/adapter.go index 03ac127..8a87290 100644 --- a/apps/gateway/internal/adapters/providers/openai/adapter.go +++ b/apps/gateway/internal/adapters/providers/openai/adapter.go @@ -122,13 +122,6 @@ func maybeBuildNovitaSameBodyRetry(model domain.PublicModel, err error, body map return retryBuildResult{} } -func maybeBuildNovitaDowngradeRetry(model domain.PublicModel, err error, body map[string]any) retryBuildResult { - if model.ProviderConfig.ProviderName != "novita" || !isInvalidRequestHTTPError(err) { - return retryBuildResult{} - } - return buildNovitaRetryRequest(body) -} - func isInvalidRequestHTTPError(err error) bool { var httpErr *ProviderHTTPError if !errors.As(err, &httpErr) || httpErr.StatusCode != http.StatusBadRequest { @@ -223,7 +216,7 @@ func mapProviderErrorWithCompatibilityContext(err error, model domain.PublicMode "upstream_code", httpErr.Code, "upstream_reason", httpErr.Reason, "upstream_trace_id", httpErr.TraceID, - "compatibility_note", "tool_choice_required_or_named_not_downgraded", + "compatibility_note", "tool_choice_required_or_named", ) } } @@ -258,16 +251,13 @@ func withProviderPolicyMeta(err error, built builtProviderRequest, attempt int, } if retryReason != "" { fields = append(fields, "retry_reason", retryReason) - if built.Policy == "novita" && retryReason == "novita_invalid_request_same_body_retry" && len(built.Omitted) == 0 { - fields = append(fields, "retry_outcome", "same_body_failed_no_safe_downgrade") + if retryReason == "novita_invalid_request_same_body_retry" { + fields = append(fields, "retry_outcome", "same_body_failed") } } if len(built.Transforms) > 0 { fields = append(fields, "transforms", built.Transforms) } - if len(built.Omitted) > 0 { - fields = append(fields, "omitted_fields", built.Omitted) - } return gwErr.WithMeta(fields...) } diff --git a/apps/gateway/internal/adapters/providers/openai/logging.go b/apps/gateway/internal/adapters/providers/openai/logging.go index 856fbbb..e8c0eb9 100644 --- a/apps/gateway/internal/adapters/providers/openai/logging.go +++ b/apps/gateway/internal/adapters/providers/openai/logging.go @@ -26,9 +26,6 @@ func (a *Adapter) logProviderRequest(ctx context.Context, model domain.PublicMod if len(built.Transforms) > 0 { fields = append(fields, "transforms", built.Transforms) } - if len(built.Omitted) > 0 { - fields = append(fields, "omitted_fields", built.Omitted) - } a.logger.Info("provider request", fields...) } diff --git a/apps/gateway/internal/adapters/providers/openai/mapper.go b/apps/gateway/internal/adapters/providers/openai/mapper.go index 957ef27..5a92ee1 100644 --- a/apps/gateway/internal/adapters/providers/openai/mapper.go +++ b/apps/gateway/internal/adapters/providers/openai/mapper.go @@ -1,27 +1,20 @@ package openai -import ( - "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" -) - type providerPolicy struct { - name string - useMaxTokens bool - explicitNullAssistantToolContent bool - requireToolReasoningContent bool - forwardDeveloperRole bool + name string + useMaxTokens bool + requireToolReasoningContent bool + forwardDeveloperRole bool } type builtProviderRequest struct { Body map[string]any Policy string Transforms []string - Omitted []string } type retryBuildResult struct { Body map[string]any - Omitted []string RetryReason string CanRetry bool } @@ -29,73 +22,21 @@ type retryBuildResult struct { func policyForProvider(providerName string) providerPolicy { switch providerName { case "deepseek": - // DeepSeek exposes an OpenAI-compatible chat API at /chat/completions, - // but currently documents the legacy `max_tokens` field and nullable - // assistant tool-call content. - return providerPolicy{ - name: "deepseek", - useMaxTokens: true, - explicitNullAssistantToolContent: true, - requireToolReasoningContent: true, - } - case "novita": - // Novita exposes an OpenAI-compatible API, but its public Chat - // Completions docs currently document `max_tokens` rather than - // OpenAI's newer `max_completion_tokens`, and require a `content` - // field that may be null for assistant tool-call messages. Keep these - // as Novita-only wire-shape translations; do not leak them into the - // public OpenAI-compatible gateway API. + // Verified against DeepSeek's API (2026-09-29): it silently ignores + // max_completion_tokens, and in thinking mode rejects assistant + // tool-call turns sent back without reasoning_content. return providerPolicy{ - name: "novita", - useMaxTokens: true, - explicitNullAssistantToolContent: true, + name: "deepseek", + useMaxTokens: true, + requireToolReasoningContent: true, } default: + // Every provider but OpenAI gets developer messages as system ones: + // DeepSeek and Kimi K3 reject the developer role, and GLM 5.3 on + // Novita ignores it (verified 2026-09-29). return providerPolicy{ name: "openai-compatible", forwardDeveloperRole: providerName == "openai", } } } - -func buildNovitaRetryRequest(body map[string]any) retryBuildResult { - retryBody := cloneBody(body) - omitted := make([]string, 0, 6) - - for _, field := range []string{"parallel_tool_calls", "store", "service_tier", "user"} { - if _, ok := retryBody[field]; ok { - delete(retryBody, field) - omitted = append(omitted, field) - } - } - - if toolChoice, ok := retryBody["tool_choice"]; ok { - switch v := toolChoice.(type) { - case string: - switch v { - case domain.ToolChoiceAuto: - delete(retryBody, "tool_choice") - omitted = append(omitted, "tool_choice=auto") - case domain.ToolChoiceNone: - delete(retryBody, "tool_choice") - delete(retryBody, "tools") - omitted = append(omitted, "tool_choice=none", "tools") - } - } - } - - return retryBuildResult{ - Body: retryBody, - Omitted: omitted, - RetryReason: "novita_invalid_request_guarded_downgrade", - CanRetry: len(omitted) > 0, - } -} - -func cloneBody(body map[string]any) map[string]any { - clone := make(map[string]any, len(body)) - for k, v := range body { - clone[k] = v - } - return clone -} diff --git a/apps/gateway/internal/adapters/providers/openai/policy_test.go b/apps/gateway/internal/adapters/providers/openai/policy_test.go index ec2f254..887f28b 100644 --- a/apps/gateway/internal/adapters/providers/openai/policy_test.go +++ b/apps/gateway/internal/adapters/providers/openai/policy_test.go @@ -13,67 +13,7 @@ import ( "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) -func TestBuildNovitaRetryRequest_GuardedDowngrade(t *testing.T) { - body := map[string]any{ - "model": "moonshotai/kimi-k2.6", - "parallel_tool_calls": false, - "store": true, - "service_tier": "auto", - "user": "user-1", - "tool_choice": "auto", - } - - retry := buildNovitaRetryRequest(body) - if !retry.CanRetry { - t.Fatal("expected retry to be allowed") - } - for _, field := range []string{"parallel_tool_calls", "store", "service_tier", "user", "tool_choice"} { - if _, ok := retry.Body[field]; ok { - t.Fatalf("retry body still contains %s", field) - } - } - for _, omitted := range []string{"parallel_tool_calls", "store", "service_tier", "user", "tool_choice=auto"} { - if !containsString(retry.Omitted, omitted) { - t.Fatalf("omitted = %v, missing %s", retry.Omitted, omitted) - } - } -} - -func TestBuildNovitaRetryRequest_ToolChoiceNoneRemovesTools(t *testing.T) { - body := map[string]any{ - "model": "moonshotai/kimi-k2.6", - "tool_choice": "none", - "tools": []map[string]any{{"type": "function"}}, - } - - retry := buildNovitaRetryRequest(body) - if !retry.CanRetry { - t.Fatal("expected retry to be allowed") - } - if _, ok := retry.Body["tool_choice"]; ok { - t.Fatal("retry body still contains tool_choice") - } - if _, ok := retry.Body["tools"]; ok { - t.Fatal("retry body still contains tools") - } -} - -func TestBuildNovitaRetryRequest_NamedToolChoiceIsNotDowngraded(t *testing.T) { - body := map[string]any{ - "model": "moonshotai/kimi-k2.6", - "tool_choice": map[string]any{ - "type": "function", - "function": map[string]any{"name": "noop"}, - }, - } - - retry := buildNovitaRetryRequest(body) - if retry.CanRetry { - t.Fatalf("named tool_choice must not be downgraded, omitted = %v", retry.Omitted) - } -} - -func TestAdapterComplete_NovitaRetriesSafeDowngrade(t *testing.T) { +func TestAdapterComplete_NovitaRetriesTheSameBody(t *testing.T) { t.Setenv("NOVITA_TEST_KEY", "test-key") var attempts int var bodies []map[string]any @@ -85,7 +25,7 @@ func TestAdapterComplete_NovitaRetriesSafeDowngrade(t *testing.T) { t.Fatalf("decode request: %v", err) } bodies = append(bodies, body) - if attempts <= 3 { + if attempts <= 2 { w.WriteHeader(http.StatusBadRequest) w.Write([]byte(`{"message":"invalid request error trace_id: testtrace","type":"invalid_request_error"}`)) return @@ -104,27 +44,16 @@ func TestAdapterComplete_NovitaRetriesSafeDowngrade(t *testing.T) { if err != nil { t.Fatalf("Complete returned error: %v", err) } - if attempts != 4 { - t.Fatalf("attempts = %d, want 4", attempts) - } - for _, field := range []string{"parallel_tool_calls", "store", "service_tier", "user", "tool_choice"} { - if _, ok := bodies[1][field]; !ok { - t.Fatalf("same-body retry should still contain %s", field) - } - } - for _, field := range []string{"parallel_tool_calls", "store", "service_tier", "user", "tool_choice"} { - if _, ok := bodies[2][field]; !ok { - t.Fatalf("second same-body retry should still contain %s", field) - } + if attempts != 3 { + t.Fatalf("attempts = %d, want 3", attempts) } - for _, field := range []string{"parallel_tool_calls", "store", "service_tier", "user", "tool_choice"} { - if _, ok := bodies[3][field]; ok { - t.Fatalf("downgrade retry body still contains %s", field) + for i, body := range bodies { + for _, field := range []string{"parallel_tool_calls", "store", "service_tier", "user", "tool_choice", "tools"} { + if _, ok := body[field]; !ok { + t.Fatalf("attempt %d changed the body: missing %s", i+1, field) + } } } - if _, ok := bodies[3]["tools"]; !ok { - t.Fatal("tool_choice=auto retry should keep tools") - } } func TestAdapterComplete_NovitaNamedToolChoiceFailsClearlyAfterSameBodyRetry(t *testing.T) { @@ -173,8 +102,8 @@ func TestAdapterComplete_NovitaSafeRetryStillReportsNamedToolChoiceIncompatibili if err == nil { t.Fatal("expected error") } - if attempts != 4 { - t.Fatalf("attempts = %d, want 4", attempts) + if attempts != 3 { + t.Fatalf("attempts = %d, want 3", attempts) } if !strings.Contains(err.Error(), "tool_choice") { t.Fatalf("error = %q, want clear tool_choice context", err.Error()) @@ -208,8 +137,8 @@ func TestAdapterComplete_NovitaFailedSameBodyRetryIncludesProviderParams(t *test if !errors.As(err, &gwErr) { t.Fatalf("error = %T, want *domain.GatewayError", err) } - if got := gwErr.Metadata["retry_outcome"]; got != "same_body_failed_no_safe_downgrade" { - t.Fatalf("retry_outcome = %v, want same_body_failed_no_safe_downgrade", got) + if got := gwErr.Metadata["retry_outcome"]; got != "same_body_failed" { + t.Fatalf("retry_outcome = %v, want same_body_failed", got) } params, ok := gwErr.Metadata["provider_params"].(map[string]any) if !ok { diff --git a/apps/gateway/internal/adapters/providers/openai/proxy.go b/apps/gateway/internal/adapters/providers/openai/proxy.go index fcfe844..22bfcdb 100644 --- a/apps/gateway/internal/adapters/providers/openai/proxy.go +++ b/apps/gateway/internal/adapters/providers/openai/proxy.go @@ -87,10 +87,6 @@ func PrepareProxyBody(raw []byte, model domain.PublicModel, stream bool) (map[st if len(toolCalls) > 0 { msg["tool_calls"] = toolCalls } - if _, has := msg["content"]; !has && len(toolCalls) > 0 && policy.explicitNullAssistantToolContent { - msg["content"] = nil - transforms = append(transforms, "assistant_tool_content=null") - } if _, has := msg["reasoning_content"]; !has && len(toolCalls) > 0 && policy.requireToolReasoningContent { msg["reasoning_content"] = "" transforms = append(transforms, "assistant_tool_reasoning_content=empty") @@ -191,37 +187,23 @@ func (a *Adapter) proxy(ctx context.Context, body map[string]any, transforms []s return nil, missingProviderCredentialError(model.ProviderConfig.ProviderName) } built := builtProviderRequest{Body: body, Policy: policyForProvider(model.ProviderConfig.ProviderName).name, Transforms: transforms} - active := built var retryReason string invalidTraceRetries, overloadRetries := 0, 0 - downgradeRetried := false for attempt := 1; ; attempt++ { - a.logProviderRequest(ctx, model, active, attempt, retryReason) - resp, err := send(ctx, apiKey, active.Body) + a.logProviderRequest(ctx, model, built, attempt, retryReason) + resp, err := send(ctx, apiKey, built.Body) if err == nil { return resp, nil } - if retry := maybeBuildNovitaSameBodyRetry(model, err, active.Body); retry.CanRetry && + if retry := maybeBuildNovitaSameBodyRetry(model, err, built.Body); retry.CanRetry && canSpendSameBodyRetry(retry.RetryReason, &invalidTraceRetries, &overloadRetries) { retryReason = retry.RetryReason metrics.ProviderRetries.WithLabelValues(model.ProviderConfig.ProviderName, retryReason).Inc() if !sleepBeforeProviderRetry(ctx, retryReason) { - return nil, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), active, attempt, retryReason) + return nil, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), built, attempt, retryReason) } continue } - if !downgradeRetried { - if retry := maybeBuildNovitaDowngradeRetry(model, err, active.Body); retry.CanRetry { - downgradeRetried = true - retryReason = retry.RetryReason - metrics.ProviderRetries.WithLabelValues(model.ProviderConfig.ProviderName, retryReason).Inc() - active = builtProviderRequest{Body: retry.Body, Policy: built.Policy, Transforms: transforms, Omitted: retry.Omitted} - if !sleepBeforeProviderRetry(ctx, retryReason) { - return nil, withProviderPolicyMeta(mapProviderError(context.Cause(ctx), model.ProviderConfig.ProviderName), active, attempt, retryReason) - } - continue - } - } - return nil, withProviderPolicyMeta(mapProviderErrorWithCompatibilityContext(err, model, active.Body), active, attempt, retryReason) + return nil, withProviderPolicyMeta(mapProviderErrorWithCompatibilityContext(err, model, built.Body), built, attempt, retryReason) } } diff --git a/apps/gateway/internal/adapters/providers/openai/proxy_test.go b/apps/gateway/internal/adapters/providers/openai/proxy_test.go index 9f34562..c759548 100644 --- a/apps/gateway/internal/adapters/providers/openai/proxy_test.go +++ b/apps/gateway/internal/adapters/providers/openai/proxy_test.go @@ -65,7 +65,7 @@ func TestPrepareProxyBody_Edits(t *testing.T) { } }) t.Run("output limit clamped and named per provider", func(t *testing.T) { - for provider, field := range map[string]string{"novita": "max_tokens", "deepseek": "max_tokens", "tinfoil": "max_completion_tokens", "phala": "max_completion_tokens"} { + for provider, field := range map[string]string{"deepseek": "max_tokens", "novita": "max_completion_tokens", "tinfoil": "max_completion_tokens", "phala": "max_completion_tokens"} { got := prepared(t, `{"messages":[],"max_completion_tokens":50000}`, provider, false) if got[field] != float64(8000) || len(got) != 3 { t.Fatalf("%s: %v", provider, got) @@ -83,7 +83,7 @@ func TestPrepareProxyBody_Edits(t *testing.T) { t.Fatal("developer role not mapped") } assistant := deepseek[1].(map[string]any) - if v, ok := assistant["content"]; !ok || v != nil || assistant["reasoning_content"] != "" { + if _, ok := assistant["content"]; ok || assistant["reasoning_content"] != "" { t.Fatalf("deepseek assistant %v", assistant) } other := prepared(t, raw, "tinfoil", false)["messages"].([]any)[1].(map[string]any) From 1185ac124751d1802fc855198967ccc55d43c35c Mon Sep 17 00:00:00 2001 From: Marketen Date: Tue, 29 Sep 2026 15:29:19 +0200 Subject: [PATCH 6/7] feat: log chunk and malformed-chunk counts when a stream ends The removed per-stream diagnostics logged these; keep them on the one end-of-request log line. Co-Authored-By: Claude Opus 5.5 --- apps/gateway/internal/application/services/proxy.go | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/apps/gateway/internal/application/services/proxy.go b/apps/gateway/internal/application/services/proxy.go index 5690641..5d9b460 100644 --- a/apps/gateway/internal/application/services/proxy.go +++ b/apps/gateway/internal/application/services/proxy.go @@ -139,6 +139,8 @@ type outcome struct { proof *domain.TinfoilTransportProof err error // Set when the response failed. end string // How a stream ended, for logs. + // Stream diagnostics: data events read, and ones that were not JSON. + chunks, malformed int } // finishProxy meters and logs a proxied response. @@ -151,6 +153,7 @@ func (s *GenerateService) finishProxy(ctx context.Context, authCtx domain.AuthCo "provider", model.ProviderConfig.ProviderName, "provider_model", model.UpstreamModelName, "stream", req.Stream, "stream_end", o.end, "finish_reason", logfields.FinishReason(o.finish), "usage_received", o.usage != nil, "usage", o.usage, "latency_ms", latencyMs, + "chunks_received", o.chunks, "malformed_chunks", o.malformed, } if o.err != nil { s.logger.Error("generation failed", append(fields, "error_type", fmt.Sprintf("%T", o.err))...) @@ -397,9 +400,11 @@ func (p *ProxyStream) observe(raw []byte) ([][]byte, *chunk) { } var c chunk if json.Unmarshal(data, &c) != nil { + p.o.malformed++ return [][]byte{line}, nil } p.chunks++ + p.o.chunks = p.chunks if c.ID != "" { p.o.providerID = c.ID } From a3b87337c15980c93472ef4061f642ead20c1c6b Mon Sep 17 00:00:00 2001 From: Marketen Date: Tue, 29 Sep 2026 16:04:39 +0200 Subject: [PATCH 7/7] fix: send each choice's role once in streams Novita repeats "role":"assistant" on every chunk for some models (MiniMax M2.7/M3 always, GLM 5.3-flash on tool calls, GLM 5.3 sometimes). OpenAI sends it once, and the OpenAI SDKs' stream helpers join repeated fields, so agents replayed the message with role "assistantassistant..." and the next request was rejected. Later repeats are dropped; other events pass unchanged. Co-Authored-By: Claude Opus 5.5 --- .../internal/application/services/proxy.go | 54 ++++++++++++++++++- .../application/services/proxy_test.go | 20 +++++++ 2 files changed, 72 insertions(+), 2 deletions(-) diff --git a/apps/gateway/internal/application/services/proxy.go b/apps/gateway/internal/application/services/proxy.go index 5d9b460..eaea5f4 100644 --- a/apps/gateway/internal/application/services/proxy.go +++ b/apps/gateway/internal/application/services/proxy.go @@ -8,6 +8,7 @@ import ( "fmt" "io" "net/http" + "strconv" "time" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/middleware" @@ -185,7 +186,9 @@ type chunk struct { ID string `json:"id"` Model *string `json:"model"` Choices []struct { + Index *int `json:"index"` Delta *struct { + Role *string `json:"role"` Content *string `json:"content"` ReasoningContent *string `json:"reasoning_content"` Reasoning *string `json:"reasoning"` @@ -327,8 +330,12 @@ type ProxyStream struct { // unmask restores PII placeholders, for keys that mask. unmask *unmasker - o outcome - sawDone bool + o outcome + sawDone bool + // roleSent records choices whose role was sent: some providers repeat + // it on every chunk, which clients that join fields (the OpenAI SDKs' + // stream helpers) turn into "assistantassistant...". + roleSent map[int]bool finished bool // Metered. chunks int } @@ -418,12 +425,55 @@ func (p *ProxyStream) observe(raw []byte) ([][]byte, *chunk) { p.o.err = err } prefix := line[:len(line)-len(data)] + data = p.dropRepeatedRole(data, &c) if p.unmask != nil && len(c.Choices) > 0 { data = p.unmask.event(data) } return [][]byte{append(append([]byte(nil), prefix...), rewriteModel(data, &c, p.public)...)}, &c } +// dropRepeatedRole removes a role the stream already sent for a choice, as +// OpenAI sends it once. Events without a repeated role pass unchanged. +func (p *ProxyStream) dropRepeatedRole(data []byte, c *chunk) []byte { + drop := map[string]bool{} + for pos, choice := range c.Choices { + if choice.Delta == nil || choice.Delta.Role == nil { + continue + } + index := pos + if choice.Index != nil { + index = *choice.Index + } + if p.roleSent[index] { + drop[strconv.Itoa(index)] = true + continue + } + if p.roleSent == nil { + p.roleSent = map[int]bool{} + } + p.roleSent[index] = true + } + if len(drop) == 0 { + return data + } + obj, err := decodeObject(data) + if err != nil { + return data + } + choices, _ := obj["choices"].([]any) + for pos, item := range choices { + choice, _ := item.(map[string]any) + if delta, ok := choice["delta"].(map[string]any); ok && drop[choiceIndex(choice["index"], pos)] { + delete(delta, "role") + } + } + out, err := marshalNoEscape(obj) + if err != nil { + return data + } + return out +} + // heldTail returns an event with PII-restored text still held back, if any. func (p *ProxyStream) heldTail() []byte { if p.unmask == nil { diff --git a/apps/gateway/internal/application/services/proxy_test.go b/apps/gateway/internal/application/services/proxy_test.go index 2b6b4a5..44853a0 100644 --- a/apps/gateway/internal/application/services/proxy_test.go +++ b/apps/gateway/internal/application/services/proxy_test.go @@ -221,3 +221,23 @@ func TestProxyStream_UnrequestedUsageIsMeteredNotForwarded(t *testing.T) { t.Fatal("usage not metered") } } + +// Some providers repeat "role" on every chunk; OpenAI sends it once, and the +// OpenAI SDKs' stream helpers join repeated fields into "assistantassistant". +func TestProxyStream_SendsRoleOncePerChoice(t *testing.T) { + first := `data: {"id":"c","choices":[{"index":0,"delta":{"role":"assistant","content":"a"},"finish_reason":null}]}` + body := first + "\n\n" + + `data: {"id":"c","choices":[{"index":0,"delta":{"role":"assistant","content":"b"},"finish_reason":null}]}` + "\n\n" + + `data: {"id":"c","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"id":"t","type":"function","function":{"name":"f","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}` + "\n\n" + + "data: [DONE]\n\n" + meter := &resultMeter{} + svc := proxyService(meter, "", map[string]ports.Provider{"primary": &proxyStub{sse: body}}, false) + call, err := svc.Proxy(context.Background(), domain.EndpointChatCompletions, []byte(`{}`), streamRequest(), "key") + if err != nil { + t.Fatal(err) + } + got, _ := relayAll(t, call.Stream) + if strings.Count(got, `"role"`) != 1 || !strings.HasPrefix(got, first+"\n") || !strings.Contains(got, `"content":"b"`) || !strings.Contains(got, `"name":"f"`) { + t.Fatalf("stream %s", got) + } +}