Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package handlers

import (
"context"
"io"
"net/http"
"sync"
Expand All @@ -15,6 +16,10 @@ import (
"github.com/google/uuid"
)

// Bound the wait for usage after generation has finished, without extending
// the lifetime of requests canceled before completion.
const streamUsageDrainTimeout = services.StreamFinalizationTimeout

// ChatCompletionsHandler handles POST /v1/chat/completions.
type ChatCompletionsHandler struct {
service *services.ChatCompletionsService
Expand Down Expand Up @@ -68,13 +73,15 @@ func (h *ChatCompletionsHandler) Handle(w http.ResponseWriter, r *http.Request)
}

func (h *ChatCompletionsHandler) handleStream(w http.ResponseWriter, r *http.Request, genReq domain.GenerateRequest, token string) {
stream, model, err := h.service.ExecuteStream(r.Context(), genReq, token)
ctx, cancelUpstream := context.WithCancel(r.Context())
defer cancelUpstream()
stream, model, err := h.service.ExecuteStream(ctx, genReq, token)
if err != nil {
WriteErrorWithLog(w, r, h.logger, err)
return
}
defer stream.Close()
h.writeStream(w, r, genReq, stream, model)
h.writeStream(w, r, genReq, stream, model, cancelUpstream)
}

func (h *ChatCompletionsHandler) writeStream(
Expand All @@ -83,6 +90,7 @@ func (h *ChatCompletionsHandler) writeStream(
genReq domain.GenerateRequest,
stream ports.GenerationStream,
model *domain.PublicModel,
cancelUpstream context.CancelFunc,
) {
responseModelID := genReq.PublicModelID
if model != nil {
Expand Down Expand Up @@ -137,23 +145,48 @@ func (h *ChatCompletionsHandler) writeStream(
responseID = event.ProviderResponseID
}

// Clients may close on finish_reason as well as [DONE]. Keep both
// signals back until metering has consumed trailing usage and settled
// the request. Canceling the upstream context safely unblocks reads
// if a provider never terminates its post-completion stream.
if event.Type == domain.StreamEventCompleted {
requestID := middleware.GetRequestID(r.Context())
providerID := responseID
timer := time.AfterFunc(streamUsageDrainTimeout, func() {
h.logger.Warn("stream finalization timeout", "request_id", requestID, "provider_request_id", providerID, "model", responseModelID, "timeout_ms", streamUsageDrainTimeout.Milliseconds())
cancelUpstream()
})
for {
tail, err := stream.Recv()
if err != nil {
break
}
if tail.Usage != nil {
event.Usage = tail.Usage
}
if tail.FinishReason != nil {
event.FinishReason = tail.FinishReason
}
}
timer.Stop()
}
chunk, done := mapper.DomainStreamEventToChatChunk(event, responseModelID, responseID, createdAt)
if chunk != nil {
writeMu.Lock()
sw.WriteData(chunk)
writeErr := sw.WriteData(chunk)
writeMu.Unlock()
if writeErr != nil {
h.logger.Warn("stream client write failed", "request_id", middleware.GetRequestID(r.Context()), "provider_request_id", responseID, "phase", "chunk", "write_failed", true)
}
}
if done {
writeMu.Lock()
sw.WriteDone()
writeErr := sw.WriteDone()
writeMu.Unlock()
// Keep draining so the usage-tracking wrapper can accumulate
// the final usage chunk before io.EOF triggers recording.
for {
if _, err := stream.Recv(); err != nil {
break
}
if writeErr != nil {
h.logger.Warn("stream client write failed", "request_id", middleware.GetRequestID(r.Context()), "provider_request_id", responseID, "phase", "done", "write_failed", true)
}
h.logger.Debug("stream completion sent", "request_id", middleware.GetRequestID(r.Context()), "provider_request_id", responseID, "done_write_failed", writeErr != nil, "client_context_canceled", r.Context().Err() != nil)
break
}
}
Expand Down
Original file line number Diff line number Diff line change
@@ -1,22 +1,28 @@
package handlers

import (
"context"
"errors"
"io"
"net/http/httptest"
"strings"
"testing"
"time"

"github.com/dappnode/dappnode-nexus-gateway/pkg/domain"
)

type handlerTestStream struct {
events []domain.StreamEvent
err error
index int
events []domain.StreamEvent
err error
index int
beforeRecv func()
}

func (s *handlerTestStream) Recv() (domain.StreamEvent, error) {
if s.beforeRecv != nil {
s.beforeRecv()
}
if s.index < len(s.events) {
event := s.events[s.index]
s.index++
Expand All @@ -28,6 +34,71 @@ func (s *handlerTestStream) Recv() (domain.StreamEvent, error) {
return domain.StreamEvent{}, io.EOF
}

func TestChatCompletionsHandler_SettlesBeforeEmittingFinish(t *testing.T) {
finish := "tool_calls"
recorder := httptest.NewRecorder()
stream := &handlerTestStream{
events: []domain.StreamEvent{
{Type: domain.StreamEventCompleted, FinishReason: &finish},
{Type: domain.StreamEventCompleted, Usage: &domain.Usage{PromptTokens: 100, CompletionTokens: 20, TotalTokens: 120}},
},
beforeRecv: func() {
if strings.Contains(recorder.Body.String(), `"finish_reason"`) || strings.Contains(recorder.Body.String(), "[DONE]") {
t.Error("client received a completion signal before upstream EOF")
}
},
}
handler := &ChatCompletionsHandler{logger: &confidentialTestLogger{}}
handler.writeStream(recorder, httptest.NewRequest("POST", "/v1/chat/completions", nil), domain.GenerateRequest{PublicModelID: "test"}, stream, nil, func() {})
if !strings.Contains(recorder.Body.String(), `"prompt_tokens":100`) || !strings.Contains(recorder.Body.String(), `"finish_reason":"tool_calls"`) {
t.Fatalf("missing final usage or finish: %s", recorder.Body.String())
}
if strings.Count(recorder.Body.String(), "data: [DONE]") != 1 {
t.Fatalf("expected exactly one DONE: %s", recorder.Body.String())
}
}

type stalledCompletionStream struct {
handlerTestStream
ctx context.Context
}

func (s *stalledCompletionStream) Recv() (domain.StreamEvent, error) {
if s.index < len(s.events) {
return s.handlerTestStream.Recv()
}
<-s.ctx.Done()
return domain.StreamEvent{}, s.ctx.Err()
}

func TestChatCompletionsHandler_BoundsTrailingUsageWait(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), streamUsageDrainTimeout+2*time.Second)
defer cancel()
finish := "tool_calls"
stream := &stalledCompletionStream{
handlerTestStream: handlerTestStream{events: []domain.StreamEvent{{Type: domain.StreamEventCompleted, FinishReason: &finish}}},
ctx: ctx,
}
recorder := httptest.NewRecorder()
log := &finalizationLogRecorder{warnings: make(chan string, 1)}
handler := &ChatCompletionsHandler{logger: log}
handler.writeStream(recorder, httptest.NewRequest("POST", "/v1/chat/completions", nil), domain.GenerateRequest{PublicModelID: "test"}, stream, nil, cancel)
select {
case message := <-log.warnings:
if message != "stream finalization timeout" {
t.Fatalf("unexpected warning: %s", message)
}
default:
t.Fatal("finalization timeout was not logged")
}
if ctx.Err() != context.Canceled {
t.Fatal("stalled upstream was not canceled")
}
if !strings.Contains(recorder.Body.String(), `"finish_reason":"tool_calls"`) || strings.Contains(recorder.Body.String(), `"usage"`) {
t.Fatalf("expected finished generation without fabricated usage: %s", recorder.Body.String())
}
}

func (*handlerTestStream) Close() error { return nil }

func TestChatCompletionsHandler_StreamEmitsExactlyOneDoneOnCleanEOF(t *testing.T) {
Expand Down Expand Up @@ -59,7 +130,7 @@ func TestChatCompletionsHandler_StreamEmitsExactlyOneDoneOnCleanEOF(t *testing.T
recorder := httptest.NewRecorder()
request := httptest.NewRequest("POST", "/v1/chat/completions", nil)

handler.writeStream(recorder, request, domain.GenerateRequest{PublicModelID: "test"}, &handlerTestStream{events: test.events}, nil)
handler.writeStream(recorder, request, domain.GenerateRequest{PublicModelID: "test"}, &handlerTestStream{events: test.events}, nil, func() {})

if count := strings.Count(recorder.Body.String(), "data: [DONE]\n\n"); count != 1 {
t.Fatalf("DONE marker count = %d, want 1; stream = %q", count, recorder.Body.String())
Expand All @@ -73,9 +144,17 @@ func TestChatCompletionsHandler_StreamErrorDoesNotClaimCompletion(t *testing.T)
recorder := httptest.NewRecorder()
request := httptest.NewRequest("POST", "/v1/chat/completions", nil)

handler.writeStream(recorder, request, domain.GenerateRequest{PublicModelID: "test"}, &handlerTestStream{err: errors.New("upstream failed")}, nil)
handler.writeStream(recorder, request, domain.GenerateRequest{PublicModelID: "test"}, &handlerTestStream{err: errors.New("upstream failed")}, nil, func() {})

if strings.Contains(recorder.Body.String(), "data: [DONE]\n\n") {
t.Fatalf("errored stream claimed completion: %q", recorder.Body.String())
}
}

// A channel keeps the timer's logging callback safe to observe from the test.
type finalizationLogRecorder struct {
confidentialTestLogger
warnings chan string
}

func (l *finalizationLogRecorder) Warn(message string, _ ...any) { l.warnings <- message }
10 changes: 6 additions & 4 deletions apps/gateway/internal/adapters/http/handlers/common.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package handlers

import (
"encoding/json"
"fmt"
"net/http"

"github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/dto"
Expand Down Expand Up @@ -41,7 +42,7 @@ func WriteErrorWithLog(w http.ResponseWriter, r *http.Request, logger ports.Logg
logger.Error("untyped internal error",
"request_id", middleware.GetRequestID(r.Context()),
"path", r.URL.Path,
"original_error", err.Error(),
"error_type", fmt.Sprintf("%T", err),
)
gwErr = domain.ErrInternal("an internal error occurred")
}
Expand All @@ -53,17 +54,18 @@ func WriteErrorWithLog(w http.ResponseWriter, r *http.Request, logger ports.Logg
"status", gwErr.HTTPStatus,
"gateway_status", gwErr.HTTPStatus,
"error_code", gwErr.Code,
"error", gwErr.Message,
}
fields = append(fields, gwErr.LogFields()...)
// Provider errors can echo prompts; only retain numeric upstream status.
if status, ok := gwErr.Metadata["upstream_status"].(int); ok {
fields = append(fields, "upstream_status", status)
}
logger.Error("request failed", fields...)
} else if gwErr.HTTPStatus >= 400 {
logger.Warn("request error",
"request_id", middleware.GetRequestID(r.Context()),
"status", gwErr.HTTPStatus,
"gateway_status", gwErr.HTTPStatus,
"error_code", gwErr.Code,
"error", gwErr.Message,
)
}

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
package handlers

import (
"errors"
"fmt"
"net/http/httptest"
"strings"
"testing"

"github.com/dappnode/dappnode-nexus-gateway/pkg/domain"
)

type privacyLogger struct{ output strings.Builder }

func (l *privacyLogger) Debug(m string, f ...any) { fmt.Fprint(&l.output, m, f) }
func (l *privacyLogger) Info(m string, f ...any) { l.Debug(m, f...) }
func (l *privacyLogger) Warn(m string, f ...any) { l.Debug(m, f...) }
func (l *privacyLogger) Error(m string, f ...any) { l.Debug(m, f...) }

func TestErrorLogsDoNotEchoRequestContents(t *testing.T) {
const secret = "PRIVATE-PROMPT-CANARY"
for _, err := range []error{errors.New(secret), domain.ErrInvalidField(secret), domain.ErrProviderError(502, secret).WithMeta("upstream_error", secret, "upstream_status", 400)} {
logger := &privacyLogger{}
WriteErrorWithLog(httptest.NewRecorder(), httptest.NewRequest("POST", "/v1/chat/completions", nil), logger, err)
if strings.Contains(logger.output.String(), secret) {
t.Fatal("error log leaked private content")
}
}
}
Loading