package server import ( "context" "encoding/json" "fmt" "net/http" "strconv" "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/connector" 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 } // 参数范围校验 if req.Temperature != nil && (*req.Temperature < 0 || *req.Temperature > 2) { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "temperature must be between 0 and 2", requestID)) return } if req.TopP != nil && (*req.TopP < 0 || *req.TopP > 1) { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "top_p must be between 0 and 1", requestID)) return } if req.MaxOutputTokens < 0 { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "max_output_tokens must be non-negative", requestID)) return } if len(req.Messages) > s.cfg.Context.MaxSessionMessages { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrContextTooLarge, fmt.Sprintf("messages count exceeds max %d", s.cfg.Context.MaxSessionMessages), requestID)) return } // 幂等性检查(仅非流式请求支持幂等) if !req.Stream && req.IdempotencyKey != "" && identity != nil { if s.checkIdempotency(w, req.IdempotencyKey, identity.AppID) { 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 } // Per-app 限流由全局 RateLimit 中间件处理,此处不再重复检查 // Resolve logical model target, err := s.modelMap.Resolve(req.Model) if err != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, err.Error(), requestID)) return } // 校验 max_output_tokens 不超过模型配置上限 if req.MaxOutputTokens > 0 && target.MaxOutputTokens > 0 && req.MaxOutputTokens > target.MaxOutputTokens { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, fmt.Sprintf("max_output_tokens %d exceeds model limit %d", req.MaxOutputTokens, target.MaxOutputTokens), 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()) w.Header().Set("X-Backpressure-Level", fmt.Sprintf("%d", s.backpressure.Level())) 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") s.metrics.SetBreakerState(1) // open handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, "circuit breaker open, please retry later", requestID)) return } // 同步熔断器状态到 Prometheus 指标 switch s.breaker.State() { case "open": s.metrics.SetBreakerState(1) case "half_open": s.metrics.SetBreakerState(2) default: s.metrics.SetBreakerState(0) } // 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()) // 解析 per-request 超时覆盖 var overrides map[string]int if req.Timeouts != nil { overrides = map[string]int{ "queue_ms": req.Timeouts.QueueMs, "first_token_ms": req.Timeouts.FirstTokenMs, "inference_ms": req.Timeouts.InferenceMs, "total_ms": req.Timeouts.TotalMs, } } resolvedTimeouts := s.timeoutMgr.ResolveTimeouts(&s.cfg.Timeouts, overrides) // Wait for task to be dequeued queueStart := time.Now() ctx, cancel := context.WithTimeout(r.Context(), time.Duration(resolvedTimeouts.QueueMs)*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 != "" { // 会话访问隔离:校验 session 属于当前 app var sess *session.Session var sessErr error if identity != nil { sess, sessErr = s.sessions.GetForApp(req.SessionID, identity.AppID) } else { sess, sessErr = s.sessions.Get(req.SessionID) } if sessErr == nil && sess != nil { // Load session history and assemble context history := s.loadSessionHistory(sess) if len(history) > 0 { assembler := ctxasm.NewAssembler(&s.cfg.Context) // 注入摘要器,启用真实 LLM 摘要(受配置控制) if s.cfg.Context.EnableLLMSummary { summaryGen := ctxasm.NewAdapterSummaryGenerator(adapterInst, target.ActualModel, s.cfg.Context.SummaryMaxTokens) summarizer := ctxasm.NewSummarizer(summaryGen, time.Duration(s.cfg.Context.SummaryTimeoutSeconds)*time.Second) assembler.SetSummarizer(summarizer) } 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, req.SessionID, req.Messages, identity.AppID, resolvedTimeouts) } else { // 非流式请求:如果带有幂等键,捕获响应用于缓存 if req.IdempotencyKey != "" && identity != nil { crw := &captureResponseWriter{ResponseWriter: w} s.handleNonStreaming(crw, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model, req.SessionID, req.Messages, identity.AppID, resolvedTimeouts) s.storeIdempotencyResult(req.IdempotencyKey, identity.AppID, crw.statusCode, crw.body) } else { s.handleNonStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model, req.SessionID, req.Messages, identity.AppID, resolvedTimeouts) } } } 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, sessionID string, originalMessages []api.Message, appID string, tc *connector.TimeoutConfig) { 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() // 流式超时控制:使用 per-request 解析后的 inference timeout streamCtx, streamCancel := context.WithTimeout(r.Context(), time.Duration(tc.InferenceMs)*time.Millisecond) defer streamCancel() ch, err := adapterInst.ChatCompletionStream(streamCtx, req) if err != nil { s.scheduler.Complete(tk.ID) tk.Transition(task.StateFailed) s.metrics.IncTask("failed") // 发送 SSE 错误事件 sse.WriteChunk(map[string]any{ "id": requestID, "object": "chat.completion.chunk", "model": logicalModel, "choices": []map[string]any{ { "index": 0, "delta": map[string]any{}, "finish_reason": "error", }, }, "error": map[string]string{ "code": handler.ErrModelUnavailable, "message": err.Error(), }, }) sse.WriteDone() handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, err.Error(), requestID)) return } inputTokens, outputTokens, fullContent, err := handler.StreamChatCompletion(sse, ch, requestID, tk.ID, logicalModel, func() { tk.FirstTokenMs = int(time.Since(inferenceStart) / time.Millisecond) s.metrics.ObserveFirstTokenLatency(int64(tk.FirstTokenMs)) }) if err != nil { // 检查是否客户端断开 if r.Context().Err() != nil { s.logger.Info("client disconnected during streaming", observability.F().Event("client_disconnect").TaskID(tk.ID)) tk.Cancel("client disconnected") s.metrics.IncCancellation() } else { s.logger.Error("streaming error", observability.F().Event("stream_error").TaskID(tk.ID).Reason(err.Error())) // 尝试非流式 fallback 降级 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.metrics.ObserveGatewayLatency(int64(tk.TotalMs)) tk.Transition(task.StateCompleted) s.metrics.IncTask("completed") s.metrics.IncRequest("stream_fallback") s.metrics.IncDegraded() 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()) s.usageTracker.Record(appID, fallbackResp.InputTokens, fallbackResp.OutputTokens, false) // 将降级结果作为单个 SSE chunk 发送 sse.WriteChunk(map[string]any{ "id": requestID, "object": "chat.completion.chunk", "model": logicalModel, "choices": []map[string]any{ { "index": 0, "delta": map[string]any{ "content": fallbackResp.Content, }, "finish_reason": fallbackResp.FinishReason, }, }, "usage": map[string]int{ "input_tokens": fallbackResp.InputTokens, "output_tokens": fallbackResp.OutputTokens, "total_tokens": fallbackResp.InputTokens + fallbackResp.OutputTokens, }, "degraded": true, }) sse.WriteDone() if sessionID != "" && fallbackResp.Content != "" { s.saveSessionMessages(sessionID, originalMessages, fallbackResp.Content) } return } // fallback 也失败,发送 SSE 错误事件 sse.WriteChunk(map[string]any{ "id": requestID, "object": "chat.completion.chunk", "model": logicalModel, "choices": []map[string]any{ { "index": 0, "delta": map[string]any{}, "finish_reason": "error", }, }, "error": map[string]string{ "code": handler.ErrInferenceTimeout, "message": err.Error(), }, }) sse.WriteDone() 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.ObserveGatewayLatency(int64(tk.TotalMs)) s.metrics.AddTokens(inputTokens, outputTokens) s.scheduler.Complete(tk.ID) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) s.metrics.SetQueueLength(s.scheduler.QueueLength()) if err != nil { s.metrics.IncRequest("stream_error") } else { s.metrics.IncRequest("stream_ok") } // 记录 per-app 用量 s.usageTracker.Record(appID, inputTokens, outputTokens, err != nil) // 保存对话到 session if sessionID != "" && fullContent != "" { s.saveSessionMessages(sessionID, originalMessages, fullContent) } } 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, sessionID string, originalMessages []api.Message, appID string, tc *connector.TimeoutConfig) { ctx, cancel := context.WithTimeout(r.Context(), time.Duration(tc.InferenceMs)*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()) // 记录 per-app 用量(降级请求) s.usageTracker.Record(appID, fallbackResp.InputTokens, fallbackResp.OutputTokens, false) chatResp := handler.BuildChatResponse(requestID, tk.ID, logicalModel, fallbackResp) chatResp.Degraded = true chatResp.Timing = s.buildTiming(tk) // 通过 HTTP 头暴露关键 timing 指标 w.Header().Set("X-Timing-Queue-Ms", strconv.Itoa(tk.QueueMs)) w.Header().Set("X-Timing-Inference-Ms", strconv.Itoa(tk.InferenceMs)) w.Header().Set("X-Timing-Total-Ms", strconv.Itoa(tk.TotalMs)) 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) // 记录延迟指标 s.metrics.ObserveGatewayLatency(int64(tk.TotalMs)) 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()) // 记录 per-app 用量 s.usageTracker.Record(appID, resp.InputTokens, resp.OutputTokens, false) chatResp := handler.BuildChatResponse(requestID, tk.ID, logicalModel, resp) chatResp.Timing = s.buildTiming(tk) // 通过 HTTP 头暴露关键 timing 指标,便于客户端和监控系统采集 w.Header().Set("X-Timing-Queue-Ms", strconv.Itoa(tk.QueueMs)) w.Header().Set("X-Timing-Inference-Ms", strconv.Itoa(tk.InferenceMs)) w.Header().Set("X-Timing-Total-Ms", strconv.Itoa(tk.TotalMs)) handler.WriteJSON(w, http.StatusOK, chatResp) // 保存对话到 session if sessionID != "" { s.saveSessionMessages(sessionID, originalMessages, resp.Content) } } // 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 { target, err := s.modelMap.Resolve(m) if err != nil { data[i] = api.ModelInfo{ID: m, Object: "model", OwnedBy: "edgeai-gateway"} continue } data[i] = api.ModelInfo{ ID: m, Object: "model", OwnedBy: "edgeai-gateway", Provider: target.Provider, ContextWindow: target.ContextWindow, MaxOutputTokens: target.MaxOutputTokens, } } 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: // 列出当前 app 的会话(支持分页) appID := "" if identity != nil { appID = identity.AppID } page := 1 pageSize := 20 if v := r.URL.Query().Get("page"); v != "" { if n, err := strconv.Atoi(v); err == nil && n > 0 { page = n } } if v := r.URL.Query().Get("page_size"); v != "" { if n, err := strconv.Atoi(v); err == nil && n > 0 && n <= 100 { pageSize = n } } offset := (page - 1) * pageSize sessions, err := s.sessions.ListByApp(appID, pageSize, offset) if err != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID)) return } // 脱敏:不返回 messages 全文,只返回元信息 items := make([]map[string]any, 0, len(sessions)) for _, sess := range sessions { items = append(items, map[string]any{ "session_id": sess.ID, "application_id": sess.ApplicationID, "user_id": sess.UserID, "message_count": len(sess.Messages), "created_at": sess.CreatedAt.Format(time.RFC3339), "last_active": sess.LastActive.Format(time.RFC3339), }) } handler.WriteJSON(w, http.StatusOK, map[string]any{ "sessions": items, "page": page, "page_size": pageSize, }) 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()) identity := auth.GetAppIdentityFromRequest(r) sessionID := strings.TrimPrefix(r.URL.Path, "/v1/sessions/") switch r.Method { case http.MethodGet: // 会话访问隔离:校验 session 属于当前 app var sess *session.Session var err error if identity != nil { sess, err = s.sessions.GetForApp(sessionID, identity.AppID) } else { 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: // 会话访问隔离:校验 session 属于当前 app if identity != nil { if err := s.sessions.DeleteForApp(sessionID, identity.AppID); err != nil { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, err.Error(), requestID)) return } } else { 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 } // saveSessionMessages 将用户消息和助手回复保存到会话历史。 // 当配置了 MaxSessionMessages 时,自动裁剪最旧消息以防止无限增长。 func (s *Server) saveSessionMessages(sessionID string, userMessages []api.Message, assistantContent string) { if sessionID == "" { return } maxMsgs := s.cfg.Context.MaxSessionMessages for _, msg := range userMessages { if err := s.sessions.AddMessageWithLimit(sessionID, msg, maxMsgs); err != nil { s.logger.Error("failed to save user message to session", observability.F().Event("session_save_error").Reason(err.Error())) } } if assistantContent != "" { assistantMsg := api.Message{Role: "assistant", Content: assistantContent} if err := s.sessions.AddMessageWithLimit(sessionID, assistantMsg, maxMsgs); err != nil { s.logger.Error("failed to save assistant message to session", observability.F().Event("session_save_error").Reason(err.Error())) } } } // 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 } // 分页参数 page := 1 pageSize := 20 if v := r.URL.Query().Get("page"); v != "" { if n, err := strconv.Atoi(v); err == nil && n > 0 { page = n } } if v := r.URL.Query().Get("page_size"); v != "" { if n, err := strconv.Atoi(v); err == nil && n > 0 && n <= 100 { pageSize = n } } offset := (page - 1) * pageSize // 状态过滤 status := r.URL.Query().Get("status") tasks, total := s.scheduler.ListTasks(status, pageSize, offset) items := make([]map[string]any, 0, len(tasks)) for _, tk := range tasks { item := 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 { item["started_at"] = tk.StartedAt.Format(time.RFC3339) } if tk.CompletedAt != nil { item["completed_at"] = tk.CompletedAt.Format(time.RFC3339) } if tk.ErrorMessage != "" { item["error_message"] = tk.ErrorMessage } item["timing"] = s.buildTiming(tk) items = append(items, item) } handler.WriteJSON(w, http.StatusOK, map[string]any{ "tasks": items, "total": total, "page": page, "page_size": pageSize, }) } 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")) } }