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 -}