Files
AIRouter/internal/server/admin_handlers.go
T
selfrelease da9c8334d8
CI / lint (push) Has been cancelled
CI / test (push) Has been cancelled
CI / build (push) Has been cancelled
CI / security-scan (push) Has been cancelled
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 全链路传播
2026-08-03 15:43:11 +08:00

334 lines
11 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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/keysGET 列出所有 KeyPOST 创建新 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/usageGET 返回所有 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-breakerGET 返回熔断器状态)。
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/schedulerGET 返回调度器状态)。
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/reloadPOST 触发配置热重载)。
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": "配置已热重载,模型映射已更新",
})
}