Skip to content
75 changes: 0 additions & 75 deletions apps/gateway/internal/adapters/http/dto/chat_completion_request.go

This file was deleted.

This file was deleted.

165 changes: 57 additions & 108 deletions apps/gateway/internal/adapters/http/handlers/chat_completions_handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"io"
"net/http"
"sync"
"sync/atomic"
"time"

"github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/http/mapper"
Expand All @@ -13,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
Expand Down Expand Up @@ -43,151 +43,100 @@ func (h *ChatCompletionsHandler) Handle(w http.ResponseWriter, r *http.Request)
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,
)
}

genReq, err := mapper.ChatCompletionRequestToDomain(body)
if err != nil {
WriteErrorWithLog(w, r, h.logger, err)
return
}

if genReq.Stream {
h.handleStream(w, r, genReq, token)
return
}

result, err := h.service.Execute(r.Context(), genReq, token)
// The gateway is a proxy: the request goes upstream as the client sent it
// and the provider's response comes back as it came.
upstream, cancelUpstream := context.WithCancel(context.WithoutCancel(r.Context()))
defer cancelUpstream()
var finished atomic.Bool
stopWatch := context.AfterFunc(r.Context(), func() {
// A client gone after the finish reason still gets its usage metered:
// the upstream is drained briefly instead of cut off.
if finished.Load() {
time.AfterFunc(streamUsageDrainTimeout, cancelUpstream)
return
}
cancelUpstream()
})
defer stopWatch()
call, err := h.service.Proxy(upstream, body, genReq, token)
if err != nil {
WriteErrorWithLog(w, r, h.logger, err)
return
}

resp := mapper.DomainToChatCompletionResponse(result)
WriteJSON(w, http.StatusOK, resp)
}

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)
if call.Stream != nil {
h.relay(w, r, call.Stream, &finished)
return
}
defer stream.Close()
h.writeStream(w, r, genReq, stream, model, cancelUpstream)
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_, _ = w.Write(call.Body)
}

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
}

// relay sends a provider's stream to the client line by line, as it came.
// Keep-alive comments fill long silences (slow reasoning), and a stream the
// provider finished without [DONE] gets one, so every client sees an end.
func (h *ChatCompletionsHandler) relay(w http.ResponseWriter, r *http.Request, stream *services.ProxyStream, finished *atomic.Bool) {
defer stream.Close()
sw, err := sse.NewWriter(w)
if err != nil {
WriteError(w, domain.ErrInternal("an internal error occurred"))
return
}

responseID := uuid.New().String()[:12]
createdAt := time.Now().Unix()

// Keep-alive: write SSE comments periodically to prevent proxy timeouts.
var writeMu sync.Mutex
lastWrite := time.Now()
write := func(line []byte) {
writeMu.Lock()
defer writeMu.Unlock()
lastWrite = time.Now()
// A client that left keeps the loop draining the upstream for usage.
_ = sw.WriteLine(line)
}
stopKeepAlive := make(chan struct{})
defer close(stopKeepAlive)
go func() {
ticker := time.NewTicker(15 * time.Second)
ticker := time.NewTicker(5 * time.Second)
defer ticker.Stop()
for {
select {
case <-stopKeepAlive:
return
case <-ticker.C:
writeMu.Lock()
sw.WriteComment("keepalive")
if time.Since(lastWrite) >= 15*time.Second {
lastWrite = time.Now()
_ = sw.WriteComment("keepalive")
}
writeMu.Unlock()
}
}
}()

var readErr error
for {
event, err := stream.Recv()
line, err := stream.Next()
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
}

// 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()
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)
readErr = err
break
}
if stream.Finished() {
finished.Store(true)
}
write(line)
}
switch {
case stream.Complete() && !stream.SawDone():
write([]byte("data: [DONE]"))
write(nil)
case readErr != io.EOF && r.Context().Err() == nil:
// The upstream broke mid-response: say so rather than end quietly.
h.logger.Warn("proxied stream interrupted", "request_id", middleware.GetRequestID(r.Context()))
write([]byte(`data: {"error":{"type":"provider_error","code":"stream_interrupted","message":"The provider stream ended unexpectedly."}}`))
write(nil)
}
}
Loading
Loading