package server import ( "context" "encoding/json" "fmt" "net/http" "strings" "time" "github.com/edgeai/gateway/internal/adapter" "github.com/edgeai/gateway/internal/auth" "github.com/edgeai/gateway/internal/config" ctxasm "github.com/edgeai/gateway/internal/context" "github.com/edgeai/gateway/internal/handler" "github.com/edgeai/gateway/internal/middleware" "github.com/edgeai/gateway/internal/observability" "github.com/edgeai/gateway/internal/router" "github.com/edgeai/gateway/internal/session" "github.com/edgeai/gateway/internal/task" "github.com/edgeai/gateway/pkg/api" "github.com/google/uuid" ) func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) return } requestID := middleware.GetRequestID(r.Context()) identity := auth.GetAppIdentityFromRequest(r) var req api.ChatRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "invalid JSON body", requestID)) return } // Validate required fields if req.Model == "" { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "model is required", requestID)) return } if len(req.Messages) == 0 { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "messages is required", requestID)) return } // Check model permission if identity != nil && !auth.CheckModelPermission(identity, req.Model) { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrPermissionDenied, "model not allowed for this application", requestID)) return } // Resolve logical model target, err := s.modelMap.Resolve(req.Model) if err != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, err.Error(), requestID)) return } // Get adapter adapterInst, err := s.registry.Get(target.Provider) if err != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, err.Error(), requestID)) return } // Parse priority priority := config.ParsePriority(req.Priority) if priority == 0 && identity != nil && !auth.CheckPriorityPermission(identity, 0) { priority = int(task.PriorityNormal) // downgrade to P2 if not allowed P0 } // Backpressure check: reject low-priority requests under high load s.backpressure.Update(s.scheduler.RunningCount(), s.scheduler.QueueLength()) if !s.backpressure.ShouldAccept(priority) { reason := s.backpressure.RejectReason(priority) s.metrics.IncRequest("backpressure_rejected") handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrQueueFull, reason, requestID)) return } // Circuit breaker check if !s.breaker.AllowRequest() { s.metrics.IncRequest("circuit_open") handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, "circuit breaker open, please retry later", requestID)) return } // Create task taskID := uuid.New().String() tk := task.NewTask(taskID, requestID, identity.AppID, identity.TenantID, req.Model, task.TaskPriority(priority), req.Stream) tk.SessionID = req.SessionID // Submit to scheduler if err := s.scheduler.Submit(tk); err != nil { s.metrics.IncRequest("queue_full") handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrQueueFull, "queue is full, please retry later", requestID)) return } s.metrics.SetQueueLength(s.scheduler.QueueLength()) // Wait for task to be dequeued queueStart := time.Now() ctx, cancel := context.WithTimeout(r.Context(), time.Duration(s.cfg.Timeouts.DefaultQueueMs)*time.Millisecond) defer cancel() dequeued, err := s.scheduler.GetNext(ctx) if err != nil { s.scheduler.Complete(tk.ID) s.metrics.IncRequest("queue_timeout") tk.Transition(task.StateTimedOut) handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrQueueTimeout, "queue timeout", requestID)) return } // Record queue time dequeued.QueueMs = int(time.Since(queueStart) / time.Millisecond) // Transition to RUNNING dequeued.Transition(task.StateRunning) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) // Build adapter request — assemble context if session_id provided messages := req.Messages if req.SessionID != "" { sess, err := s.sessions.Get(req.SessionID) if err == nil && sess != nil { // Load session history and assemble context history := s.loadSessionHistory(sess) if len(history) > 0 { assembler := ctxasm.NewAssembler(&s.cfg.Context) result := assembler.Assemble(history, req.Messages, target.ContextWindow, target.MaxOutputTokens, req.ContextPolicy) messages = result.Messages } } } adapterReq := &adapter.ChatRequest{ RequestID: requestID, Model: target.ActualModel, Messages: messages, MaxTokens: target.MaxOutputTokens, Temperature: req.Temperature, TopP: req.TopP, Stream: req.Stream, CancelCh: dequeued.Cancelled(), } if req.MaxOutputTokens > 0 { adapterReq.MaxTokens = req.MaxOutputTokens } if req.Stream { s.handleStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model) } else { s.handleNonStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model) } } func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, _ *router.ModelTarget, logicalModel string) { sse := handler.NewSSEWriter(w) if sse == nil { s.scheduler.Complete(tk.ID) handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, "streaming not supported", requestID)) return } tk.Transition(task.StateStreaming) inferenceStart := time.Now() ch, err := adapterInst.ChatCompletionStream(r.Context(), req) if err != nil { s.scheduler.Complete(tk.ID) tk.Transition(task.StateFailed) s.metrics.IncTask("failed") handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, err.Error(), requestID)) return } inputTokens, outputTokens, err := handler.StreamChatCompletion(sse, ch, requestID, tk.ID, logicalModel) if err != nil { s.logger.Error("streaming error", observability.F().Event("stream_error").TaskID(tk.ID).Reason(err.Error())) tk.Transition(task.StateFailed) s.metrics.IncTask("failed") s.breaker.RecordError() } else { tk.Transition(task.StateCompleted) s.metrics.IncTask("completed") s.breaker.RecordSuccess() } tk.InferenceMs = int(time.Since(inferenceStart) / time.Millisecond) tk.TotalMs = int(time.Since(tk.CreatedAt) / time.Millisecond) s.metrics.AddTokens(inputTokens, outputTokens) s.scheduler.Complete(tk.ID) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) s.metrics.SetQueueLength(s.scheduler.QueueLength()) s.metrics.IncRequest("stream_ok") } func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, target *router.ModelTarget, logicalModel string) { ctx, cancel := context.WithTimeout(r.Context(), time.Duration(s.cfg.Timeouts.DefaultInferenceMs)*time.Millisecond) defer cancel() inferenceStart := time.Now() resp, err := adapterInst.ChatCompletion(ctx, req) if err != nil { // Try overload fallback strategies fallbackResp, fbErr := s.tryOverloadFallback(r.Context(), req, target, logicalModel) if fbErr == nil && fallbackResp != nil { tk.Degraded = true tk.InferenceMs = int(time.Since(inferenceStart) / time.Millisecond) tk.TotalMs = int(time.Since(tk.CreatedAt) / time.Millisecond) s.scheduler.Complete(tk.ID) tk.Transition(task.StateCompleted) s.metrics.IncTask("completed") s.metrics.IncRequest("ok") s.metrics.IncRequest("overload_fallback") s.metrics.AddTokens(fallbackResp.InputTokens, fallbackResp.OutputTokens) s.breaker.RecordSuccess() s.scheduler.Complete(tk.ID) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) s.metrics.SetQueueLength(s.scheduler.QueueLength()) chatResp := handler.BuildChatResponse(requestID, tk.ID, logicalModel, fallbackResp) chatResp.Degraded = true chatResp.Timing = s.buildTiming(tk) handler.WriteJSON(w, http.StatusOK, chatResp) return } s.scheduler.Complete(tk.ID) tk.Transition(task.StateFailed) s.metrics.IncTask("failed") s.metrics.IncRequest("error") s.breaker.RecordError() if strings.Contains(err.Error(), "timeout") || ctx.Err() != nil { tk.Transition(task.StateTimedOut) handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInferenceTimeout, "inference timeout", requestID)) } else { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, err.Error(), requestID)) } return } tk.InferenceMs = int(time.Since(inferenceStart) / time.Millisecond) tk.TotalMs = int(time.Since(tk.CreatedAt) / time.Millisecond) tk.Transition(task.StateCompleted) s.metrics.IncTask("completed") s.metrics.IncRequest("ok") s.metrics.AddTokens(resp.InputTokens, resp.OutputTokens) s.breaker.RecordSuccess() s.scheduler.Complete(tk.ID) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) s.metrics.SetQueueLength(s.scheduler.QueueLength()) chatResp := handler.BuildChatResponse(requestID, tk.ID, logicalModel, resp) chatResp.Timing = s.buildTiming(tk) handler.WriteJSON(w, http.StatusOK, chatResp) } // tryOverloadFallback attempts fallback strategies when the primary model fails. func (s *Server) tryOverloadFallback(ctx context.Context, req *adapter.ChatRequest, target *router.ModelTarget, logicalModel string) (*adapter.ChatResponse, error) { exclude := map[string]bool{target.Endpoint: true} result := s.overload.Resolve(logicalModel, exclude) if result.Rejected { return nil, fmt.Errorf("overload: %s", result.Reason) } fallbackAdapter, err := s.registry.Get(result.Target.Provider) if err != nil { return nil, fmt.Errorf("overload: adapter not found: %w", err) } fallbackReq := &adapter.ChatRequest{ RequestID: req.RequestID, Model: result.Target.ActualModel, Messages: req.Messages, MaxTokens: result.Target.MaxOutputTokens, Temperature: req.Temperature, TopP: req.TopP, Stream: false, CancelCh: req.CancelCh, } s.logger.Info("overload fallback", observability.F(). Event("overload_fallback"). Set("strategy", router.StrategyName(result.Strategy)). Set("from", logicalModel). Set("to", result.Target.LogicalModel)) fbCtx, fbCancel := context.WithTimeout(ctx, time.Duration(s.cfg.Timeouts.DefaultInferenceMs)*time.Millisecond) defer fbCancel() return fallbackAdapter.ChatCompletion(fbCtx, fallbackReq) } func (s *Server) handleModels(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet { handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) return } models := s.modelMap.List() data := make([]api.ModelInfo, len(models)) for i, m := range models { data[i] = api.ModelInfo{ ID: m, Object: "model", OwnedBy: "edgeai-gateway", } } resp := api.ModelListResponse{ Object: "list", Data: data, } handler.WriteJSON(w, http.StatusOK, resp) } func (s *Server) handleSessions(w http.ResponseWriter, r *http.Request) { requestID := middleware.GetRequestID(r.Context()) identity := auth.GetAppIdentityFromRequest(r) switch r.Method { case http.MethodPost: var req api.SessionRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "invalid JSON body", requestID)) return } if req.ApplicationID == "" { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "application_id is required", requestID)) return } sessionID := uuid.New().String() tenantID := "" if identity != nil { tenantID = identity.TenantID } sess, err := s.sessions.Create(sessionID, req.ApplicationID, tenantID, req.UserID, req.Config) if err != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID)) return } resp := api.SessionResponse{ SessionID: sess.ID, ApplicationID: sess.ApplicationID, UserID: sess.UserID, CreatedAt: sess.CreatedAt.Format(time.RFC3339), LastActive: sess.LastActive.Format(time.RFC3339), } handler.WriteJSON(w, http.StatusCreated, resp) case http.MethodGet: // List sessions (simplified: return empty for now) handler.WriteJSON(w, http.StatusOK, map[string]any{"sessions": []any{}}) default: handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) } } func (s *Server) handleSessionByID(w http.ResponseWriter, r *http.Request) { requestID := middleware.GetRequestID(r.Context()) sessionID := strings.TrimPrefix(r.URL.Path, "/v1/sessions/") switch r.Method { case http.MethodGet: sess, err := s.sessions.Get(sessionID) if err != nil || sess == nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "session not found", requestID)) return } handler.WriteJSON(w, http.StatusOK, sess) case http.MethodDelete: if err := s.sessions.Delete(sessionID); err != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID)) return } handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "deleted"}) default: handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) } } // loadSessionHistory returns the message history from a session. func (s *Server) loadSessionHistory(sess *session.Session) []api.Message { if sess == nil { return nil } return sess.Messages } // buildTiming constructs Timing metadata from a task. func (s *Server) buildTiming(tk *task.Task) *api.Timing { return &api.Timing{ QueueMs: tk.QueueMs, FirstTokenMs: tk.FirstTokenMs, InferenceMs: tk.InferenceMs, TotalMs: tk.TotalMs, } } func (s *Server) handleTasks(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet { handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) return } handler.WriteJSON(w, http.StatusOK, map[string]any{"tasks": []any{}}) } func (s *Server) handleTaskByID(w http.ResponseWriter, r *http.Request) { requestID := middleware.GetRequestID(r.Context()) taskID := strings.TrimPrefix(r.URL.Path, "/v1/tasks/") switch r.Method { case http.MethodGet: tk, ok := s.scheduler.GetTask(taskID) if !ok { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "task not found", requestID)) return } resp := map[string]any{ "task_id": tk.ID, "request_id": tk.RequestID, "session_id": tk.SessionID, "state": string(tk.GetState()), "priority": int(tk.Priority), "logical_model": tk.LogicalModel, "stream": tk.Stream, "created_at": tk.CreatedAt.Format(time.RFC3339), "degraded": tk.Degraded, } if tk.StartedAt != nil { resp["started_at"] = tk.StartedAt.Format(time.RFC3339) } if tk.CompletedAt != nil { resp["completed_at"] = tk.CompletedAt.Format(time.RFC3339) } if tk.CancelReason != "" { resp["cancel_reason"] = tk.CancelReason } if tk.ErrorMessage != "" { resp["error_message"] = tk.ErrorMessage } resp["timing"] = s.buildTiming(tk) handler.WriteJSON(w, http.StatusOK, resp) case http.MethodDelete: tk, ok := s.scheduler.GetTask(taskID) if !ok { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "task not found", requestID)) return } if tk.IsTerminal() { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "task already in terminal state", requestID)) return } tk.Cancel("client requested cancellation") s.scheduler.Complete(tk.ID) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) s.metrics.SetQueueLength(s.scheduler.QueueLength()) handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "cancelled"}) default: handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) } }