feat: 十轮网关优化 - 安全加固/可观测性/性能/可靠性
- SSE Keepalive Ping (15s心跳防止代理断连) - Timing HTTP 头 (X-Timing-Queue/Inference/Total-Ms) - Adapter Request-ID 传播到后端 - Session 清理日志回调 - Server 安全加固 (ReadHeaderTimeout/MaxHeaderBytes 防 slowloris) - Usage Tracker 数据保留清理 (retentionDays + 定期清理) - Config Reload 后 Adapter Registry 更新 (RegisterIfAbsent + RWMutex) - Rate Limiter 空闲 Bucket 清理 (30分钟过期) - Shutdown Drain 超时可配置 (ShutdownDrainSeconds) - Config 模型字段校验增强 (provider/endpoint/actual_model) - Auth 过期 Key 自动清理 (5分钟扫描) - Admin API Rate Limiting - Adapter Health Check 独立超时 (每个 adapter 3s) - TCP 连接阶段超时 (DialContext 5s + KeepAlive 30s) - 幂等键缓存、审计日志、Gzip 中间件、CORS Expose Headers - Backpressure 响应头、熔断器 Prometheus 指标 - 连接池优化、Trace-ID 全链路传播
This commit is contained in:
@@ -0,0 +1,333 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"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/handler"
|
||||
"github.com/edgeai/gateway/internal/middleware"
|
||||
"github.com/edgeai/gateway/internal/observability"
|
||||
)
|
||||
|
||||
// handleAdminKeys 处理 /v1/admin/keys(GET 列出所有 Key,POST 创建新 Key)。
|
||||
func (s *Server) handleAdminKeys(w http.ResponseWriter, r *http.Request) {
|
||||
requestID := middleware.GetRequestID(r.Context())
|
||||
|
||||
switch r.Method {
|
||||
case http.MethodGet:
|
||||
keys, err := s.auth.ListKeys()
|
||||
if err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID))
|
||||
return
|
||||
}
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]any{"keys": keys})
|
||||
|
||||
case http.MethodPost:
|
||||
var req struct {
|
||||
AppID string `json:"app_id"`
|
||||
TenantID string `json:"tenant_id"`
|
||||
Name string `json:"name"`
|
||||
AllowedModels []string `json:"allowed_models"`
|
||||
AllowedPriorities []int `json:"allowed_priorities"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
ExpiresAt string `json:"expires_at"` // RFC3339 格式,空表示永不过期
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "invalid JSON body", requestID))
|
||||
return
|
||||
}
|
||||
if req.AppID == "" {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "app_id is required", requestID))
|
||||
return
|
||||
}
|
||||
if req.Name == "" {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "name is required", requestID))
|
||||
return
|
||||
}
|
||||
|
||||
apiKey := auth.GenerateAPIKey()
|
||||
identity := &auth.AppIdentity{
|
||||
AppID: req.AppID,
|
||||
TenantID: req.TenantID,
|
||||
Name: req.Name,
|
||||
AllowedModels: req.AllowedModels,
|
||||
AllowedPriorities: req.AllowedPriorities,
|
||||
IsAdmin: req.IsAdmin,
|
||||
}
|
||||
if req.ExpiresAt != "" {
|
||||
if t, err := time.Parse(time.RFC3339, req.ExpiresAt); err == nil {
|
||||
identity.ExpiresAt = &t
|
||||
}
|
||||
}
|
||||
|
||||
if err := s.auth.AddKey(apiKey, identity); err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID))
|
||||
return
|
||||
}
|
||||
|
||||
// 审计日志
|
||||
actor := ""
|
||||
if id := auth.GetAppIdentityFromRequest(r); id != nil {
|
||||
actor = id.AppID
|
||||
}
|
||||
s.audit.RecordFromRequest(r, actor, "create_key", "api_key", req.AppID, http.StatusCreated, map[string]interface{}{
|
||||
"app_id": req.AppID,
|
||||
"name": req.Name,
|
||||
"is_admin": req.IsAdmin,
|
||||
"allowed_models": req.AllowedModels,
|
||||
})
|
||||
|
||||
handler.WriteJSON(w, http.StatusCreated, map[string]any{
|
||||
"api_key": apiKey,
|
||||
"app_id": req.AppID,
|
||||
"name": req.Name,
|
||||
"is_admin": req.IsAdmin,
|
||||
"allowed_models": req.AllowedModels,
|
||||
"expires_at": req.ExpiresAt,
|
||||
"message": "请妥善保存此 API Key,之后将无法再次查看",
|
||||
})
|
||||
|
||||
default:
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
|
||||
}
|
||||
}
|
||||
|
||||
// handleAdminKeyByID 处理 /v1/admin/keys/{id}(DELETE 删除,PATCH 禁用,PUT 轮换)。
|
||||
func (s *Server) handleAdminKeyByID(w http.ResponseWriter, r *http.Request) {
|
||||
requestID := middleware.GetRequestID(r.Context())
|
||||
idStr := strings.TrimPrefix(r.URL.Path, "/v1/admin/keys/")
|
||||
id, err := strconv.ParseInt(idStr, 10, 64)
|
||||
if err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "invalid key id", requestID))
|
||||
return
|
||||
}
|
||||
|
||||
switch r.Method {
|
||||
case http.MethodDelete:
|
||||
if err := s.auth.DeleteKey(id); err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, err.Error(), requestID))
|
||||
return
|
||||
}
|
||||
// 审计日志
|
||||
actor := ""
|
||||
if ident := auth.GetAppIdentityFromRequest(r); ident != nil {
|
||||
actor = ident.AppID
|
||||
}
|
||||
s.audit.RecordFromRequest(r, actor, "delete_key", "api_key", strconv.FormatInt(id, 10), http.StatusOK, nil)
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "deleted"})
|
||||
|
||||
case http.MethodPatch:
|
||||
if err := s.auth.DisableKey(id); err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, err.Error(), requestID))
|
||||
return
|
||||
}
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "disabled"})
|
||||
|
||||
case http.MethodPut:
|
||||
// Key 轮换:生成新 Key,禁用旧 Key
|
||||
newKey, err := s.auth.RotateKey(id)
|
||||
if err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, err.Error(), requestID))
|
||||
return
|
||||
}
|
||||
// 审计日志
|
||||
actor := ""
|
||||
if ident := auth.GetAppIdentityFromRequest(r); ident != nil {
|
||||
actor = ident.AppID
|
||||
}
|
||||
s.audit.RecordFromRequest(r, actor, "rotate_key", "api_key", strconv.FormatInt(id, 10), http.StatusOK, nil)
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]any{
|
||||
"status": "rotated",
|
||||
"api_key": newKey,
|
||||
"message": "旧 Key 已禁用,请妥善保存新 Key",
|
||||
})
|
||||
|
||||
default:
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
|
||||
}
|
||||
}
|
||||
|
||||
// handleAdminUsage 处理 /v1/admin/usage(GET 返回所有 app 的用量统计)。
|
||||
func (s *Server) handleAdminUsage(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
|
||||
return
|
||||
}
|
||||
|
||||
usage := s.usageTracker.GetAll()
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]any{"usage": usage})
|
||||
}
|
||||
|
||||
// handleAdminUsageByApp 处理 /v1/admin/usage/{app_id}(GET 返回指定 app 的用量统计,支持 ?history=hours 查询历史)。
|
||||
func (s *Server) handleAdminUsageByApp(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
|
||||
return
|
||||
}
|
||||
|
||||
appID := strings.TrimPrefix(r.URL.Path, "/v1/admin/usage/")
|
||||
if appID == "" {
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "app_id is required"))
|
||||
return
|
||||
}
|
||||
|
||||
// 支持 ?history=24 查询历史用量
|
||||
if historyHours := r.URL.Query().Get("history"); historyHours != "" {
|
||||
hours := 24
|
||||
fmt.Sscanf(historyHours, "%d", &hours)
|
||||
history, err := s.usageTracker.GetHistory(appID, hours)
|
||||
if err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInternalError, err.Error()))
|
||||
return
|
||||
}
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]any{
|
||||
"app_id": appID,
|
||||
"hours": hours,
|
||||
"history": history,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
usage := s.usageTracker.Get(appID)
|
||||
handler.WriteJSON(w, http.StatusOK, usage)
|
||||
}
|
||||
|
||||
// handleAdminRateLimit 处理 /v1/admin/ratelimit/{app_id}(GET 查看限流,PUT 设置自定义限流)。
|
||||
func (s *Server) handleAdminRateLimit(w http.ResponseWriter, r *http.Request) {
|
||||
requestID := middleware.GetRequestID(r.Context())
|
||||
appID := strings.TrimPrefix(r.URL.Path, "/v1/admin/ratelimit/")
|
||||
if appID == "" {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "app_id is required", requestID))
|
||||
return
|
||||
}
|
||||
|
||||
switch r.Method {
|
||||
case http.MethodGet:
|
||||
burst, ratePerMin := s.rateLimiter.GetLimit(appID)
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]any{
|
||||
"app_id": appID,
|
||||
"burst": burst,
|
||||
"rate_per_minute": ratePerMin,
|
||||
})
|
||||
|
||||
case http.MethodPut:
|
||||
var req struct {
|
||||
Burst int `json:"burst"`
|
||||
RatePerMinute int `json:"rate_per_minute"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "invalid JSON body", requestID))
|
||||
return
|
||||
}
|
||||
if req.Burst <= 0 || req.RatePerMinute <= 0 {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "burst and rate_per_minute must be positive", requestID))
|
||||
return
|
||||
}
|
||||
refillPerSec := float64(req.RatePerMinute) / 60.0
|
||||
s.rateLimiter.SetLimit(appID, float64(req.Burst), refillPerSec)
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]any{
|
||||
"app_id": appID,
|
||||
"burst": req.Burst,
|
||||
"rate_per_minute": req.RatePerMinute,
|
||||
"status": "updated",
|
||||
})
|
||||
|
||||
default:
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
|
||||
}
|
||||
}
|
||||
|
||||
// handleAdminCircuitBreaker 处理 /v1/admin/circuit-breaker(GET 返回熔断器状态)。
|
||||
func (s *Server) handleAdminCircuitBreaker(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
|
||||
return
|
||||
}
|
||||
|
||||
stats := s.breaker.Stats()
|
||||
handler.WriteJSON(w, http.StatusOK, stats)
|
||||
}
|
||||
|
||||
// handleAdminScheduler 处理 /v1/admin/scheduler(GET 返回调度器状态)。
|
||||
func (s *Server) handleAdminScheduler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
|
||||
return
|
||||
}
|
||||
|
||||
stats := map[string]any{
|
||||
"running_tasks": s.scheduler.RunningCount(),
|
||||
"queued_tasks": s.scheduler.QueueLength(),
|
||||
"max_running": s.cfg.Scheduler.MaxRunningTasks,
|
||||
"max_queued": s.cfg.Scheduler.MaxQueuedTasks,
|
||||
"priority_aging_seconds": s.cfg.Scheduler.PriorityAgingSeconds,
|
||||
"fairness": s.cfg.Scheduler.Fairness,
|
||||
}
|
||||
handler.WriteJSON(w, http.StatusOK, stats)
|
||||
}
|
||||
|
||||
// handleAdminConfigReload 处理 /v1/admin/config/reload(POST 触发配置热重载)。
|
||||
func (s *Server) handleAdminConfigReload(w http.ResponseWriter, r *http.Request) {
|
||||
requestID := middleware.GetRequestID(r.Context())
|
||||
if r.Method != http.MethodPost {
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
|
||||
return
|
||||
}
|
||||
|
||||
// 重新加载配置文件
|
||||
cfgPath := config.ConfigPath()
|
||||
newCfg, err := config.Load(cfgPath)
|
||||
if err != nil {
|
||||
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, fmt.Sprintf("reload config failed: %v", err), requestID))
|
||||
return
|
||||
}
|
||||
|
||||
// 热更新模型映射
|
||||
s.modelMap.Update(newCfg)
|
||||
|
||||
// 注册新增 adapter(已有 adapter 不重复注册)
|
||||
for _, mc := range newCfg.Models {
|
||||
provider := mc.Provider
|
||||
endpoint := mc.Endpoint
|
||||
switch provider {
|
||||
case "ollama":
|
||||
s.registry.RegisterIfAbsent(provider, func() adapter.ModelAdapter {
|
||||
return adapter.NewOllamaAdapter(endpoint)
|
||||
})
|
||||
case "vllm":
|
||||
s.registry.RegisterIfAbsent(provider, func() adapter.ModelAdapter {
|
||||
return adapter.NewVLLMAdapter(endpoint)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 更新 Server 持有的配置引用
|
||||
s.cfg = newCfg
|
||||
|
||||
s.logger.Info("config hot-reloaded",
|
||||
observability.F().
|
||||
Event("config_reload").
|
||||
Set("config_path", cfgPath))
|
||||
|
||||
// 审计日志
|
||||
actor := ""
|
||||
if ident := auth.GetAppIdentityFromRequest(r); ident != nil {
|
||||
actor = ident.AppID
|
||||
}
|
||||
s.audit.RecordFromRequest(r, actor, "config_reload", "config", cfgPath, http.StatusOK, map[string]interface{}{
|
||||
"models": len(newCfg.Models),
|
||||
})
|
||||
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]any{
|
||||
"status": "reloaded",
|
||||
"models": len(newCfg.Models),
|
||||
"message": "配置已热重载,模型映射已更新",
|
||||
})
|
||||
}
|
||||
+376
-25
@@ -5,12 +5,14 @@ import (
|
||||
"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"
|
||||
@@ -47,12 +49,39 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
|
||||
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 {
|
||||
@@ -60,6 +89,13 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
|
||||
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 {
|
||||
@@ -75,6 +111,7 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
// 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")
|
||||
@@ -85,9 +122,19 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
|
||||
// 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()
|
||||
@@ -103,9 +150,21 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
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(s.cfg.Timeouts.DefaultQueueMs)*time.Millisecond)
|
||||
ctx, cancel := context.WithTimeout(r.Context(), time.Duration(resolvedTimeouts.QueueMs)*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
dequeued, err := s.scheduler.GetNext(ctx)
|
||||
@@ -127,12 +186,25 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
|
||||
// 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 {
|
||||
// 会话访问隔离:校验 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
|
||||
}
|
||||
@@ -155,13 +227,20 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
if req.Stream {
|
||||
s.handleStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model)
|
||||
s.handleStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model, req.SessionID, req.Messages, identity.AppID, resolvedTimeouts)
|
||||
} else {
|
||||
s.handleNonStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model)
|
||||
// 非流式请求:如果带有幂等键,捕获响应用于缓存
|
||||
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, _ *router.ModelTarget, logicalModel string) {
|
||||
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)
|
||||
@@ -172,21 +251,120 @@ func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapter
|
||||
tk.Transition(task.StateStreaming)
|
||||
inferenceStart := time.Now()
|
||||
|
||||
ch, err := adapterInst.ChatCompletionStream(r.Context(), req)
|
||||
// 流式超时控制:使用 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, err := handler.StreamChatCompletion(sse, ch, requestID, tk.ID, logicalModel)
|
||||
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 {
|
||||
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()
|
||||
// 检查是否客户端断开
|
||||
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")
|
||||
@@ -196,15 +374,30 @@ func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapter
|
||||
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())
|
||||
s.metrics.IncRequest("stream_ok")
|
||||
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) {
|
||||
ctx, cancel := context.WithTimeout(r.Context(), time.Duration(s.cfg.Timeouts.DefaultInferenceMs)*time.Millisecond)
|
||||
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()
|
||||
@@ -228,9 +421,16 @@ func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adap
|
||||
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
|
||||
}
|
||||
@@ -253,6 +453,9 @@ func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adap
|
||||
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")
|
||||
@@ -262,9 +465,21 @@ func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adap
|
||||
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.
|
||||
@@ -314,10 +529,18 @@ func (s *Server) handleModels(w http.ResponseWriter, r *http.Request) {
|
||||
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",
|
||||
ID: m,
|
||||
Object: "model",
|
||||
OwnedBy: "edgeai-gateway",
|
||||
Provider: target.Provider,
|
||||
ContextWindow: target.ContextWindow,
|
||||
MaxOutputTokens: target.MaxOutputTokens,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -366,8 +589,46 @@ func (s *Server) handleSessions(w http.ResponseWriter, r *http.Request) {
|
||||
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{}})
|
||||
// 列出当前 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"))
|
||||
@@ -376,11 +637,19 @@ func (s *Server) handleSessions(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
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:
|
||||
sess, err := s.sessions.Get(sessionID)
|
||||
// 会话访问隔离:校验 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
|
||||
@@ -388,9 +657,17 @@ func (s *Server) handleSessionByID(w http.ResponseWriter, r *http.Request) {
|
||||
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
|
||||
// 会话访问隔离:校验 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"})
|
||||
|
||||
@@ -407,6 +684,28 @@ func (s *Server) loadSessionHistory(sess *session.Session) []api.Message {
|
||||
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{
|
||||
@@ -422,7 +721,59 @@ func (s *Server) handleTasks(w http.ResponseWriter, r *http.Request) {
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
|
||||
return
|
||||
}
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]any{"tasks": []any{}})
|
||||
|
||||
// 分页参数
|
||||
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) {
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/edgeai/gateway/internal/handler"
|
||||
)
|
||||
|
||||
// idempotencyEntry 幂等键缓存条目。
|
||||
// 记录请求处理状态和响应,用于在重复请求时返回缓存结果。
|
||||
type idempotencyEntry struct {
|
||||
status int
|
||||
response []byte
|
||||
createdAt time.Time
|
||||
inFlight bool // 正在处理中
|
||||
}
|
||||
|
||||
const (
|
||||
// idempotencyTTL 幂等键缓存存活时间。
|
||||
idempotencyTTL = 10 * time.Minute
|
||||
// idempotencyCleanupInterval 清理间隔。
|
||||
idempotencyCleanupInterval = 5 * time.Minute
|
||||
)
|
||||
|
||||
// checkIdempotency 检查幂等键,如果重复请求则返回缓存的响应。
|
||||
// 如果是首次请求,返回 nil 表示可以继续处理。
|
||||
// 如果是正在处理中的重复请求,返回 409 Conflict。
|
||||
func (s *Server) checkIdempotency(w http.ResponseWriter, key, appID string) bool {
|
||||
if key == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
cacheKey := appID + ":" + key
|
||||
|
||||
s.idempotencyMu.Lock()
|
||||
entry, exists := s.idempotencyCache[cacheKey]
|
||||
if exists {
|
||||
if entry.inFlight {
|
||||
// 正在处理中,返回 409
|
||||
s.idempotencyMu.Unlock()
|
||||
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "duplicate idempotency key, request already in progress"))
|
||||
return true
|
||||
}
|
||||
|
||||
// 检查是否过期
|
||||
if time.Since(entry.createdAt) > idempotencyTTL {
|
||||
delete(s.idempotencyCache, cacheKey)
|
||||
} else {
|
||||
// 返回缓存的响应
|
||||
s.idempotencyMu.Unlock()
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Header().Set("X-Idempotent-Replay", "true")
|
||||
w.WriteHeader(entry.status)
|
||||
w.Write(entry.response)
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// 标记为正在处理中
|
||||
s.idempotencyCache[cacheKey] = &idempotencyEntry{
|
||||
inFlight: true,
|
||||
createdAt: time.Now(),
|
||||
}
|
||||
s.idempotencyMu.Unlock()
|
||||
return false
|
||||
}
|
||||
|
||||
// storeIdempotencyResult 存储幂等请求的响应结果。
|
||||
func (s *Server) storeIdempotencyResult(key, appID string, status int, response []byte) {
|
||||
if key == "" {
|
||||
return
|
||||
}
|
||||
|
||||
cacheKey := appID + ":" + key
|
||||
|
||||
s.idempotencyMu.Lock()
|
||||
s.idempotencyCache[cacheKey] = &idempotencyEntry{
|
||||
status: status,
|
||||
response: response,
|
||||
createdAt: time.Now(),
|
||||
inFlight: false,
|
||||
}
|
||||
s.idempotencyMu.Unlock()
|
||||
}
|
||||
|
||||
// cleanupIdempotencyCache 清理过期的幂等键缓存。
|
||||
func (s *Server) cleanupIdempotencyCache() {
|
||||
s.idempotencyMu.Lock()
|
||||
defer s.idempotencyMu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
for k, entry := range s.idempotencyCache {
|
||||
if now.Sub(entry.createdAt) > idempotencyTTL {
|
||||
delete(s.idempotencyCache, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// startIdempotencyCleanup 启动后台清理任务,定期清理过期的幂等键。
|
||||
// 返回停止函数。
|
||||
func (s *Server) startIdempotencyCleanup() func() {
|
||||
ticker := time.NewTicker(idempotencyCleanupInterval)
|
||||
stopCh := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
s.cleanupIdempotencyCache()
|
||||
case <-stopCh:
|
||||
ticker.Stop()
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
return func() {
|
||||
close(stopCh)
|
||||
}
|
||||
}
|
||||
|
||||
// captureResponseWriter 捕获响应状态和响应体,用于幂等性缓存。
|
||||
type captureResponseWriter struct {
|
||||
http.ResponseWriter
|
||||
statusCode int
|
||||
body []byte
|
||||
}
|
||||
|
||||
func (crw *captureResponseWriter) WriteHeader(code int) {
|
||||
crw.statusCode = code
|
||||
crw.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
|
||||
func (crw *captureResponseWriter) Write(b []byte) (int, error) {
|
||||
if crw.statusCode == 0 {
|
||||
crw.statusCode = 200
|
||||
}
|
||||
crw.body = append(crw.body, b...)
|
||||
return crw.ResponseWriter.Write(b)
|
||||
}
|
||||
|
||||
func (crw *captureResponseWriter) Flush() {
|
||||
if f, ok := crw.ResponseWriter.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
}
|
||||
+296
-36
@@ -7,33 +7,50 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"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"
|
||||
"github.com/edgeai/gateway/internal/handler"
|
||||
"github.com/edgeai/gateway/internal/middleware"
|
||||
"github.com/edgeai/gateway/internal/observability"
|
||||
"github.com/edgeai/gateway/internal/ratelimit"
|
||||
"github.com/edgeai/gateway/internal/router"
|
||||
"github.com/edgeai/gateway/internal/scheduler"
|
||||
"github.com/edgeai/gateway/internal/session"
|
||||
"github.com/edgeai/gateway/internal/usage"
|
||||
)
|
||||
|
||||
// Server is the main HTTP server for the AI gateway.
|
||||
type Server struct {
|
||||
cfg *config.Config
|
||||
logger *observability.Logger
|
||||
metrics *observability.Metrics
|
||||
HTTPSrv *http.Server
|
||||
auth *auth.Authenticator
|
||||
registry *adapter.Registry
|
||||
modelMap *router.LogicalModelMapping
|
||||
scheduler *scheduler.Scheduler
|
||||
sessions *session.Store
|
||||
breaker *scheduler.CircuitBreaker
|
||||
backpressure *scheduler.BackpressureManager
|
||||
overload *router.OverloadResolver
|
||||
cfg *config.Config
|
||||
logger *observability.Logger
|
||||
metrics *observability.Metrics
|
||||
HTTPSrv *http.Server
|
||||
auth *auth.Authenticator
|
||||
registry *adapter.Registry
|
||||
modelMap *router.LogicalModelMapping
|
||||
scheduler *scheduler.Scheduler
|
||||
sessions *session.Store
|
||||
breaker *scheduler.CircuitBreaker
|
||||
backpressure *scheduler.BackpressureManager
|
||||
overload *router.OverloadResolver
|
||||
rateLimiter *ratelimit.Limiter
|
||||
usageTracker *usage.Tracker
|
||||
audit *observability.AuditLogger
|
||||
timeoutMgr *connector.TimeoutManager
|
||||
sessionCleanup func() // 停止会话清理任务的函数
|
||||
adminSrv *http.Server // Admin API 独立端口服务器
|
||||
|
||||
healthCache map[string]bool // adapter 健康缓存
|
||||
healthCacheTime time.Time // 缓存时间
|
||||
healthCacheMu sync.RWMutex // 缓存锁
|
||||
|
||||
idempotencyCache map[string]*idempotencyEntry // 幂等键缓存
|
||||
idempotencyMu sync.Mutex // 幂等缓存锁
|
||||
}
|
||||
|
||||
// New creates a new Server instance with all components wired.
|
||||
@@ -105,36 +122,115 @@ func New(cfg *config.Config, logger *observability.Logger) (*Server, error) {
|
||||
// Initialize metrics
|
||||
metrics := observability.NewMetrics()
|
||||
|
||||
s := &Server{
|
||||
cfg: cfg,
|
||||
logger: logger,
|
||||
metrics: metrics,
|
||||
auth: authenticator,
|
||||
registry: registry,
|
||||
modelMap: modelMap,
|
||||
scheduler: sched,
|
||||
sessions: sessionStore,
|
||||
breaker: breaker,
|
||||
backpressure: bp,
|
||||
overload: overloadResolver,
|
||||
// Initialize per-app rate limiter (令牌桶)
|
||||
refillPerSec := float64(cfg.Auth.RateLimitPerMinute) / 60.0
|
||||
rateLimiter := ratelimit.NewLimiter(float64(cfg.Auth.RateLimitBurst), refillPerSec)
|
||||
|
||||
// Initialize per-app usage tracker (SQLite 持久化)
|
||||
usageDBPath := filepath.Join(filepath.Dir(dbPath), "usage.db")
|
||||
usageTracker, err := usage.NewTracker(time.Duration(cfg.Auth.UsageWindowMinutes)*time.Minute, usageDBPath, cfg.Observability.AuditRetentionDays)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("init usage tracker: %w", err)
|
||||
}
|
||||
|
||||
timeoutMgr := connector.NewTimeoutManager(&cfg.Timeouts)
|
||||
|
||||
auditLogger := observability.NewAuditLogger(filepath.Dir(dbPath), logger)
|
||||
|
||||
s := &Server{
|
||||
cfg: cfg,
|
||||
logger: logger,
|
||||
metrics: metrics,
|
||||
auth: authenticator,
|
||||
registry: registry,
|
||||
modelMap: modelMap,
|
||||
scheduler: sched,
|
||||
sessions: sessionStore,
|
||||
breaker: breaker,
|
||||
backpressure: bp,
|
||||
overload: overloadResolver,
|
||||
rateLimiter: rateLimiter,
|
||||
usageTracker: usageTracker,
|
||||
audit: auditLogger,
|
||||
timeoutMgr: timeoutMgr,
|
||||
healthCache: make(map[string]bool),
|
||||
idempotencyCache: make(map[string]*idempotencyEntry),
|
||||
}
|
||||
|
||||
// 启动会话 TTL 自动清理任务(带日志回调)
|
||||
s.sessionCleanup = sessionStore.StartCleanupTask(cfg.Context.SessionIdleTTLMinutes, 10, func(deleted int64) {
|
||||
logger.Info("session cleanup completed",
|
||||
observability.F().
|
||||
Event("session_cleanup").
|
||||
Set("deleted_sessions", deleted))
|
||||
})
|
||||
|
||||
// 启动幂等键缓存清理任务
|
||||
idempotencyCleanup := s.startIdempotencyCleanup()
|
||||
_ = idempotencyCleanup // 进程退出时自动回收
|
||||
|
||||
mux := http.NewServeMux()
|
||||
s.registerRoutes(mux)
|
||||
|
||||
// Apply middleware chain (order: Recovery → Logging → RequestID → BodyLimit → Auth → handler)
|
||||
// Apply middleware chain (order: Recovery → Logging → CORS → Auth → RateLimit → RequestID → BodyLimit → handler)
|
||||
h := middleware.RequestID(mux)
|
||||
h = middleware.BodyLimit(cfg.Server.MaxRequestBodyMB)(h)
|
||||
h = middleware.RateLimit(func(appID string) (bool, int, int) {
|
||||
info := s.rateLimiter.AllowWithInfo(appID)
|
||||
return info.Allowed, info.Limit, info.Remaining
|
||||
}, func(r *http.Request) string {
|
||||
if id := auth.GetAppIdentityFromRequest(r); id != nil {
|
||||
return id.AppID
|
||||
}
|
||||
return ""
|
||||
}, map[string]bool{"/health": true, "/ready": true, "/metrics": true})(h)
|
||||
h = s.auth.Middleware(h)
|
||||
h = middleware.CORS(cfg.Server.CORSAllowedOrigins)(h)
|
||||
h = middleware.Gzip(h)
|
||||
h = middleware.Logging(logger)(h)
|
||||
h = middleware.Recovery(logger)(h)
|
||||
|
||||
s.HTTPSrv = &http.Server{
|
||||
Addr: fmt.Sprintf("%s:%d", cfg.Server.Host, cfg.Server.Port),
|
||||
Handler: h,
|
||||
ReadTimeout: 30 * time.Second,
|
||||
WriteTimeout: 0, // no write timeout for SSE
|
||||
IdleTimeout: 120 * time.Second,
|
||||
Addr: fmt.Sprintf("%s:%d", cfg.Server.Host, cfg.Server.Port),
|
||||
Handler: h,
|
||||
ReadTimeout: 30 * time.Second,
|
||||
ReadHeaderTimeout: 10 * time.Second, // 防止 slowloris 攻击
|
||||
WriteTimeout: 0, // no write timeout for SSE
|
||||
IdleTimeout: 120 * time.Second,
|
||||
MaxHeaderBytes: 1 << 20, // 1MB 限制请求头大小
|
||||
}
|
||||
|
||||
// Admin API 独立端口(如果配置了 admin_port)
|
||||
if cfg.Server.AdminPort > 0 && cfg.Server.AdminPort != cfg.Server.Port {
|
||||
adminMux := http.NewServeMux()
|
||||
s.registerAdminRoutes(adminMux)
|
||||
|
||||
// Admin 中间件链(与主服务器相同,但额外加 CORS)
|
||||
adminH := middleware.RequestID(adminMux)
|
||||
adminH = middleware.BodyLimit(cfg.Server.MaxRequestBodyMB)(adminH)
|
||||
adminH = middleware.RateLimit(func(appID string) (bool, int, int) {
|
||||
info := s.rateLimiter.AllowWithInfo(appID)
|
||||
return info.Allowed, info.Limit, info.Remaining
|
||||
}, func(r *http.Request) string {
|
||||
if id := auth.GetAppIdentityFromRequest(r); id != nil {
|
||||
return id.AppID
|
||||
}
|
||||
return ""
|
||||
}, nil)(adminH)
|
||||
adminH = s.auth.Middleware(adminH)
|
||||
adminH = middleware.CORS(cfg.Server.CORSAllowedOrigins)(adminH)
|
||||
adminH = middleware.Logging(logger)(adminH)
|
||||
adminH = middleware.Recovery(logger)(adminH)
|
||||
|
||||
s.adminSrv = &http.Server{
|
||||
Addr: fmt.Sprintf("%s:%d", cfg.Server.Host, cfg.Server.AdminPort),
|
||||
Handler: adminH,
|
||||
ReadTimeout: 30 * time.Second,
|
||||
ReadHeaderTimeout: 10 * time.Second, // 防止 slowloris 攻击
|
||||
WriteTimeout: 30 * time.Second,
|
||||
IdleTimeout: 120 * time.Second,
|
||||
MaxHeaderBytes: 1 << 20, // 1MB 限制请求头大小
|
||||
}
|
||||
}
|
||||
|
||||
return s, nil
|
||||
@@ -161,53 +257,182 @@ func (s *Server) registerRoutes(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/v1/tasks/", s.handleTaskByID)
|
||||
}
|
||||
|
||||
// registerAdminRoutes 注册管理 API 路由到独立的 mux(运行在 admin_port)。
|
||||
func (s *Server) registerAdminRoutes(mux *http.ServeMux) {
|
||||
// Admin API: API Key 管理(需要 admin 权限)
|
||||
mux.Handle("/v1/admin/keys", auth.RequireAdmin(http.HandlerFunc(s.handleAdminKeys)))
|
||||
mux.Handle("/v1/admin/keys/", auth.RequireAdmin(http.HandlerFunc(s.handleAdminKeyByID)))
|
||||
|
||||
// Admin API: 用量统计(需要 admin 权限)
|
||||
mux.Handle("/v1/admin/usage", auth.RequireAdmin(http.HandlerFunc(s.handleAdminUsage)))
|
||||
mux.Handle("/v1/admin/usage/", auth.RequireAdmin(http.HandlerFunc(s.handleAdminUsageByApp)))
|
||||
|
||||
// Admin API: per-app 限流管理(需要 admin 权限)
|
||||
mux.Handle("/v1/admin/ratelimit/", auth.RequireAdmin(http.HandlerFunc(s.handleAdminRateLimit)))
|
||||
|
||||
// Admin API: 熔断器状态(需要 admin 权限)
|
||||
mux.Handle("/v1/admin/circuit-breaker", auth.RequireAdmin(http.HandlerFunc(s.handleAdminCircuitBreaker)))
|
||||
|
||||
// Admin API: 调度器状态(需要 admin 权限)
|
||||
mux.Handle("/v1/admin/scheduler", auth.RequireAdmin(http.HandlerFunc(s.handleAdminScheduler)))
|
||||
|
||||
// Admin API: 配置热重载(需要 admin 权限)
|
||||
mux.Handle("/v1/admin/config/reload", auth.RequireAdmin(http.HandlerFunc(s.handleAdminConfigReload)))
|
||||
|
||||
// Admin API: 健康检查(复用主服务器逻辑)
|
||||
mux.HandleFunc("/health", s.handleHealth)
|
||||
mux.HandleFunc("/ready", s.handleReady)
|
||||
}
|
||||
|
||||
// Authenticator returns the authenticator instance (for testing/management).
|
||||
func (s *Server) Authenticator() *auth.Authenticator {
|
||||
return s.auth
|
||||
}
|
||||
|
||||
// Start begins listening for HTTP requests.
|
||||
// 如果配置了 admin_port,同时启动 Admin API 服务器。
|
||||
// 启动时对所有已注册 adapter 执行健康预检,不可达的 adapter 记录警告但不阻止启动。
|
||||
func (s *Server) Start() error {
|
||||
s.logger.Info("http server starting", observability.F().
|
||||
Event("server_start").
|
||||
Set("addr", s.HTTPSrv.Addr))
|
||||
|
||||
// Adapter 启动健康预检
|
||||
s.checkAdaptersHealth()
|
||||
|
||||
// 启动 Admin API 服务器(独立 goroutine)
|
||||
if s.adminSrv != nil {
|
||||
s.logger.Info("admin api server starting", observability.F().
|
||||
Event("admin_server_start").
|
||||
Set("addr", s.adminSrv.Addr))
|
||||
go func() {
|
||||
if err := s.adminSrv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
||||
s.logger.Error("admin api server error", observability.F().
|
||||
Event("admin_server_error").
|
||||
Reason(err.Error()))
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
return s.HTTPSrv.ListenAndServe()
|
||||
}
|
||||
|
||||
// checkAdaptersHealth 对所有已注册 adapter 执行健康预检。
|
||||
// 不可达的 adapter 记录警告但不阻止启动,允许部分降级运行。
|
||||
// 每个 adapter 使用独立超时,避免单个 adapter 阻塞拖累其他。
|
||||
func (s *Server) checkAdaptersHealth() {
|
||||
for _, name := range s.registry.Names() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
a, _ := s.registry.Get(name)
|
||||
if err := a.HealthCheck(ctx); err != nil {
|
||||
s.logger.Warn("adapter health check failed on startup",
|
||||
observability.F().
|
||||
Event("adapter_health_check_failed").
|
||||
Set("adapter", name).
|
||||
Reason(err.Error()))
|
||||
} else {
|
||||
s.logger.Info("adapter health check passed",
|
||||
observability.F().
|
||||
Event("adapter_health_check_ok").
|
||||
Set("adapter", name))
|
||||
}
|
||||
cancel()
|
||||
}
|
||||
}
|
||||
|
||||
// Shutdown gracefully shuts down the server.
|
||||
// 等待最多 30 秒让进行中的请求完成,记录未完成任务状态。
|
||||
func (s *Server) Shutdown() error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
// 记录关机时的任务状态
|
||||
runningCount := s.scheduler.RunningCount()
|
||||
queueLength := s.scheduler.QueueLength()
|
||||
s.logger.Info("server shutting down",
|
||||
observability.F().
|
||||
Event("server_shutdown").
|
||||
Set("running_tasks", runningCount).
|
||||
Set("queued_tasks", queueLength))
|
||||
|
||||
// 等待运行中任务完成(超时时间从配置读取)
|
||||
drainSeconds := s.cfg.Timeouts.ShutdownDrainSeconds
|
||||
if drainSeconds <= 0 {
|
||||
drainSeconds = 10
|
||||
}
|
||||
if runningCount > 0 {
|
||||
s.logger.Info("waiting for running tasks to complete",
|
||||
observability.F().
|
||||
Event("shutdown_drain").
|
||||
Set("running_tasks", runningCount).
|
||||
Set("drain_timeout_seconds", drainSeconds))
|
||||
for i := 0; i < drainSeconds*10 && s.scheduler.RunningCount() > 0; i++ {
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
remaining := s.scheduler.RunningCount()
|
||||
if remaining > 0 {
|
||||
s.logger.Warn("shutdown: tasks still running after drain timeout",
|
||||
observability.F().
|
||||
Event("shutdown_drain_timeout").
|
||||
Set("remaining_tasks", remaining))
|
||||
} else {
|
||||
s.logger.Info("all running tasks completed",
|
||||
observability.F().Event("shutdown_drain_complete"))
|
||||
}
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
s.scheduler.Stop()
|
||||
if s.sessionCleanup != nil {
|
||||
s.sessionCleanup()
|
||||
}
|
||||
if s.sessions != nil {
|
||||
s.sessions.Close()
|
||||
}
|
||||
if s.auth != nil {
|
||||
s.auth.Close()
|
||||
}
|
||||
if s.usageTracker != nil {
|
||||
s.usageTracker.Close()
|
||||
}
|
||||
if s.audit != nil {
|
||||
s.audit.Close()
|
||||
}
|
||||
|
||||
if s.adminSrv != nil {
|
||||
s.adminSrv.Shutdown(ctx)
|
||||
}
|
||||
|
||||
return s.HTTPSrv.Shutdown(ctx)
|
||||
}
|
||||
|
||||
func (s *Server) handleHealth(w http.ResponseWriter, r *http.Request) {
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "ok"})
|
||||
// Liveness 探针:进程存活即返回 200
|
||||
handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "alive"})
|
||||
}
|
||||
|
||||
func (s *Server) handleReady(w http.ResponseWriter, r *http.Request) {
|
||||
// Readiness 探针:检查 adapter 健康状态 + DB 连通性
|
||||
// 使用 5 秒 TTL 缓存,避免高频探针请求打满 adapter
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
ready := true
|
||||
reasons := []string{}
|
||||
|
||||
for _, name := range s.registry.Names() {
|
||||
a, _ := s.registry.Get(name)
|
||||
if err := a.HealthCheck(r.Context()); err != nil {
|
||||
// 检查所有 adapter 健康状态(带缓存)
|
||||
adapterHealth := s.getCachedAdapterHealth(r.Context())
|
||||
for name, ok := range adapterHealth {
|
||||
if !ok {
|
||||
ready = false
|
||||
reasons = append(reasons, fmt.Sprintf("%s: %v", name, err))
|
||||
reasons = append(reasons, fmt.Sprintf("adapter %s: unhealthy", name))
|
||||
}
|
||||
}
|
||||
|
||||
// 检查 scheduler 是否正常
|
||||
if s.scheduler.RunningCount() >= s.cfg.Scheduler.MaxRunningTasks {
|
||||
ready = false
|
||||
reasons = append(reasons, "scheduler: at max capacity")
|
||||
}
|
||||
|
||||
if ready {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte(`{"status":"ready"}`))
|
||||
@@ -217,6 +442,41 @@ func (s *Server) handleReady(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
|
||||
// getCachedAdapterHealth 返回 adapter 健康状态,使用 5 秒 TTL 缓存。
|
||||
// 缓存过期后异步刷新,首次请求或缓存过期时同步调用 adapter HealthCheck。
|
||||
func (s *Server) getCachedAdapterHealth(ctx context.Context) map[string]bool {
|
||||
const healthCacheTTL = 5 * time.Second
|
||||
|
||||
s.healthCacheMu.RLock()
|
||||
if time.Since(s.healthCacheTime) < healthCacheTTL && len(s.healthCache) > 0 {
|
||||
result := make(map[string]bool, len(s.healthCache))
|
||||
for k, v := range s.healthCache {
|
||||
result[k] = v
|
||||
}
|
||||
s.healthCacheMu.RUnlock()
|
||||
return result
|
||||
}
|
||||
s.healthCacheMu.RUnlock()
|
||||
|
||||
// 缓存过期,同步执行健康检查
|
||||
result := make(map[string]bool)
|
||||
for _, name := range s.registry.Names() {
|
||||
a, _ := s.registry.Get(name)
|
||||
if err := a.HealthCheck(ctx); err != nil {
|
||||
result[name] = false
|
||||
} else {
|
||||
result[name] = true
|
||||
}
|
||||
}
|
||||
|
||||
s.healthCacheMu.Lock()
|
||||
s.healthCache = result
|
||||
s.healthCacheTime = time.Now()
|
||||
s.healthCacheMu.Unlock()
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func extractDBPath(connStr string) string {
|
||||
if strings.HasPrefix(connStr, "sqlite://") {
|
||||
return strings.TrimPrefix(connStr, "sqlite://")
|
||||
|
||||
Reference in New Issue
Block a user