package server import ( "context" "encoding/json" "net/http" "strings" "time" "github.com/edgeai/gateway/internal/adapter" "github.com/edgeai/gateway/internal/auth" "github.com/edgeai/gateway/internal/config" "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/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 } // Create task taskID := uuid.New().String() tk := task.NewTask(taskID, requestID, identity.AppID, identity.TenantID, req.Model, task.TaskPriority(priority), req.Stream) // 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 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") handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrQueueTimeout, "queue timeout", requestID)) return } // Transition to RUNNING dequeued.Transition(task.StateRunning) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) // Build adapter request adapterReq := &adapter.ChatRequest{ RequestID: requestID, Model: target.ActualModel, Messages: req.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, target *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) 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") } else { tk.Transition(task.StateCompleted) s.metrics.IncTask("completed") } 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() resp, err := adapterInst.ChatCompletion(ctx, req) if err != nil { s.scheduler.Complete(tk.ID) tk.Transition(task.StateFailed) s.metrics.IncTask("failed") s.metrics.IncRequest("error") if strings.Contains(err.Error(), "timeout") || ctx.Err() != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInferenceTimeout, "inference timeout", requestID)) } else { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, err.Error(), requestID)) } return } tk.Transition(task.StateCompleted) s.metrics.IncTask("completed") s.metrics.IncRequest("ok") s.metrics.AddTokens(resp.InputTokens, resp.OutputTokens) 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) handler.WriteJSON(w, http.StatusOK, chatResp) } 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")) } }