Files
superidou 9501751792 0.10.1: 安全与业务逻辑加固、新品牌与部署加固
三轮审查修复(60+ 项),相对远端 main(b536672)的关键变更:
- 安全: 数据面 SSRF 拨号防护(防 DNS rebinding)/上游凭据剥离/登录防枚举
  与锁定态统一/可信代理(X-Forwarded-For)限流加固/会话版本失效机制/
  撤销即时传播/弱密钥拒绝启动/脱敏字节级重写(保签名契约)
- 业务逻辑: 裸 body 上传 panic/bootstrap 审计管线卡死/定价通配符优先级/
  全局工具可见性/调度器停机补跑/TOTP 挑战令牌消费顺序/熔断探针语义/
  >4MB 响应 token 计量/管理员重置密码作废会话 等
- 前端: 新 logo(语枢 AI 网关主题)/Provider 凭据异常警示/删除入口/
  后端错误消息透传/localStorage 敏感数据收敛
- 部署: CREDENTIAL_MASTER_KEY 持久化与弱值拒绝/Provider DELETE 接口/
  nginx 安全头/worker 内存限制
- 新增迁移 000029(key_hash 索引)/000030(usage_daily 币种维度)
2026-08-13 10:50:51 +08:00

530 lines
18 KiB
Go

package provider
import (
"context"
"encoding/json"
"errors"
"log/slog"
"net/http"
"regexp"
"strings"
"time"
"aigateway.local/core/internal/identity"
"aigateway.local/core/internal/platform/apiresponse"
)
var providerCodePattern = regexp.MustCompile(`^[a-z][a-z0-9_-]{1,63}$`)
type AdminHTTPHandler struct {
repository *Repository
cipher *CredentialCipher
identity *identity.Service
allowPrivate bool
changeHook func(context.Context) error
operations AdminOperations
mux *http.ServeMux
logger *slog.Logger
}
func (h *AdminHTTPHandler) SetChangeHook(hook func(context.Context) error) {
h.changeHook = hook
}
func (h *AdminHTTPHandler) SetOperations(operations AdminOperations) {
h.operations = operations
}
type providerInput struct {
Code string `json:"code"`
Adapter string `json:"adapter"`
BaseURL string `json:"base_url"`
APIKey *string `json:"api_key"`
Capabilities []string `json:"capabilities"`
Config json.RawMessage `json:"config"`
Enabled *bool `json:"enabled"`
}
type modelRouteInput struct {
Name string `json:"name"`
SourceModel string `json:"source_model"`
TargetModel string `json:"target_model"`
ProviderID string `json:"provider_id"`
Weight int `json:"weight"`
Priority int `json:"priority"`
Conditions json.RawMessage `json:"conditions"`
Enabled *bool `json:"enabled"`
}
type modelRouteConditions struct {
Endpoints []string `json:"endpoints"`
APIKeyIDs []string `json:"api_key_ids"`
TenantIDs []string `json:"tenant_ids"`
}
func NewAdminHTTPHandler(repository *Repository, cipher *CredentialCipher, identityService *identity.Service, allowPrivate bool) *AdminHTTPHandler {
handler := &AdminHTTPHandler{
repository: repository, cipher: cipher, identity: identityService,
allowPrivate: allowPrivate, mux: http.NewServeMux(), logger: slog.Default(),
}
handler.mux.HandleFunc("GET /api/v1/admin/providers", handler.list)
handler.mux.HandleFunc("POST /api/v1/admin/providers", handler.create)
handler.mux.HandleFunc("PUT /api/v1/admin/providers/{provider_id}", handler.update)
handler.mux.HandleFunc("DELETE /api/v1/admin/providers/{provider_id}", handler.delete)
handler.mux.HandleFunc("POST /api/v1/admin/providers/{provider_id}/test", handler.testConnection)
handler.mux.HandleFunc("GET /api/v1/admin/providers/{provider_id}/models", handler.listModels)
handler.mux.HandleFunc("POST /api/v1/admin/providers/{provider_id}/models/sync", handler.syncModels)
handler.mux.HandleFunc("POST /api/v1/admin/providers/credentials/rotate", handler.rotateCredentials)
handler.mux.HandleFunc("GET /api/v1/admin/model-routes", handler.listModelRoutes)
handler.mux.HandleFunc("POST /api/v1/admin/model-routes", handler.createModelRoute)
handler.mux.HandleFunc("PUT /api/v1/admin/model-routes/{model_route_id}", handler.updateModelRoute)
handler.mux.HandleFunc("DELETE /api/v1/admin/model-routes/{model_route_id}", handler.deleteModelRoute)
return handler
}
func (h *AdminHTTPHandler) ServeHTTP(writer http.ResponseWriter, request *http.Request) {
h.mux.ServeHTTP(writer, request)
}
func (h *AdminHTTPHandler) list(writer http.ResponseWriter, request *http.Request) {
if _, ok := h.requirePermission(writer, request, identity.PermissionProviderRead); !ok {
return
}
records, err := h.repository.List(request.Context())
if err != nil {
h.writeError(writer, err)
return
}
items := make([]map[string]any, 0, len(records))
for _, record := range records {
item, err := h.view(record)
if err != nil {
h.writeError(writer, err)
return
}
items = append(items, item)
}
apiresponse.OK(writer, items)
}
func (h *AdminHTTPHandler) create(writer http.ResponseWriter, request *http.Request) {
actor, ok := h.requirePermission(writer, request, identity.PermissionProviderManage)
if !ok {
return
}
input, record, err := h.decodeRecord(writer, request)
if err != nil {
h.writeError(writer, err)
return
}
apiKey := ""
if input.APIKey != nil {
apiKey = strings.TrimSpace(*input.APIKey)
}
if err := h.setCredentials(&record, apiKey); err != nil {
h.writeError(writer, err)
return
}
created, err := h.repository.Create(request.Context(), record, actor.ID)
if err != nil {
h.writeError(writer, err)
return
}
if h.changeHook != nil {
_ = h.changeHook(request.Context())
}
view, err := h.view(created)
if err != nil {
h.writeError(writer, err)
return
}
apiresponse.OK(writer, view)
}
func (h *AdminHTTPHandler) delete(writer http.ResponseWriter, request *http.Request) {
actor, ok := h.requirePermission(writer, request, identity.PermissionProviderManage)
if !ok {
return
}
if err := h.repository.Delete(request.Context(), request.PathValue("provider_id"), actor.ID); err != nil {
h.writeError(writer, err)
return
}
h.propagateChange(request.Context())
apiresponse.OK(writer, map[string]bool{"deleted": true})
}
func (h *AdminHTTPHandler) update(writer http.ResponseWriter, request *http.Request) {
actor, ok := h.requirePermission(writer, request, identity.PermissionProviderManage)
if !ok {
return
}
input, record, err := h.decodeRecord(writer, request)
if err != nil {
h.writeError(writer, err)
return
}
record.ID = request.PathValue("provider_id")
if record.ID == "" {
apiresponse.Error(writer, http.StatusBadRequest, "provider_id 不能为空")
return
}
replaceCredentials := input.APIKey != nil
if replaceCredentials {
if err := h.setCredentials(&record, strings.TrimSpace(*input.APIKey)); err != nil {
h.writeError(writer, err)
return
}
}
if _, err := h.repository.Update(request.Context(), record, actor.ID, replaceCredentials); err != nil {
h.writeError(writer, err)
return
}
updated, err := h.repository.Get(request.Context(), record.ID)
if err != nil {
h.writeError(writer, err)
return
}
if h.changeHook != nil {
_ = h.changeHook(request.Context())
}
view, err := h.view(updated)
if err != nil {
h.writeError(writer, err)
return
}
apiresponse.OK(writer, view)
}
func (h *AdminHTTPHandler) decodeRecord(writer http.ResponseWriter, request *http.Request) (providerInput, Record, error) {
var input providerInput
decoder := json.NewDecoder(http.MaxBytesReader(writer, request.Body, 1<<20))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&input); err != nil {
return input, Record{}, errors.New("请求格式无效")
}
input.Code = strings.ToLower(strings.TrimSpace(input.Code))
if !providerCodePattern.MatchString(input.Code) {
return input, Record{}, errors.New("code 必须以小写字母开头,且只能包含小写字母、数字、下划线和连字符")
}
input.Adapter = strings.TrimSpace(input.Adapter)
if input.Adapter == "" {
input.Adapter = "openai-compatible"
}
if input.Adapter != "openai-compatible" {
return input, Record{}, errors.New("当前只支持 openai-compatible adapter")
}
validationCtx, cancel := context.WithTimeout(request.Context(), 3*time.Second)
defer cancel()
baseURL, err := ValidateBaseURL(validationCtx, input.BaseURL, h.allowPrivate)
if err != nil {
return input, Record{}, err
}
if len(input.Capabilities) == 0 {
input.Capabilities = []string{"chat", "responses", "embeddings", "models"}
}
if len(input.Config) == 0 {
input.Config = json.RawMessage(`{}`)
}
var configObject map[string]any
if err := json.Unmarshal(input.Config, &configObject); err != nil || configObject == nil {
return input, Record{}, errors.New("config 必须是 JSON 对象")
}
if value, exists := configObject["default"]; exists {
if _, ok := value.(bool); !ok {
return input, Record{}, errors.New("config.default 必须是布尔值")
}
}
enabled := true
if input.Enabled != nil {
enabled = *input.Enabled
}
return input, Record{
Code: input.Code, Adapter: input.Adapter, BaseURL: baseURL,
Capabilities: input.Capabilities, Config: input.Config, Enabled: enabled,
}, nil
}
func (h *AdminHTTPHandler) setCredentials(record *Record, apiKey string) error {
payload, err := json.Marshal(Credentials{APIKey: apiKey})
if err != nil {
return err
}
encrypted, version, err := h.cipher.Encrypt(payload)
if err != nil {
return err
}
record.EncryptedCredentials = encrypted
record.CredentialKEKVersion = version
return nil
}
func (h *AdminHTTPHandler) view(record Record) (map[string]any, error) {
// 凭据解密失败(KEK 轮换后旧记录、数据损坏)不得让整个列表接口报错:
// 降级为"未确认"状态并附警示,管理端仍可编辑/删除该记录恢复。
keyConfigured := false
masked := ""
credentialError := ""
plaintext, err := h.cipher.Decrypt(record.EncryptedCredentials, record.CredentialKEKVersion)
if err != nil {
credentialError = "凭据无法解密(加密密钥不匹配或数据损坏),请重新保存凭据"
} else {
var credentials Credentials
if json.Unmarshal(plaintext, &credentials) != nil {
credentialError = "凭据数据格式无效,请重新保存凭据"
} else {
keyConfigured = credentials.APIKey != ""
masked = maskSecret(credentials.APIKey)
}
}
return map[string]any{
"id": record.ID, "code": record.Code, "adapter": record.Adapter,
"base_url": record.BaseURL, "capabilities": record.Capabilities,
"config": record.Config, "enabled": record.Enabled, "revision": record.Revision,
"credential_kek_version": record.CredentialKEKVersion,
"key_configured": keyConfigured, "api_key_masked": masked,
"credential_error": credentialError,
}, nil
}
func (h *AdminHTTPHandler) testConnection(writer http.ResponseWriter, request *http.Request) {
if _, ok := h.requirePermission(writer, request, identity.PermissionProviderManage); !ok {
return
}
if !h.requireOperations(writer) {
return
}
result, err := h.operations.TestConnection(request.Context(), request.PathValue("provider_id"))
if err != nil {
h.writeError(writer, err)
return
}
apiresponse.OK(writer, result)
}
func (h *AdminHTTPHandler) listModels(writer http.ResponseWriter, request *http.Request) {
if _, ok := h.requirePermission(writer, request, identity.PermissionProviderRead); !ok {
return
}
if !h.requireOperations(writer) {
return
}
models, err := h.operations.ListModels(request.Context(), request.PathValue("provider_id"))
if err != nil {
h.writeError(writer, err)
return
}
if models == nil {
models = []Model{}
}
apiresponse.OK(writer, models)
}
func (h *AdminHTTPHandler) syncModels(writer http.ResponseWriter, request *http.Request) {
actor, ok := h.requirePermission(writer, request, identity.PermissionProviderManage)
if !ok {
return
}
if !h.requireOperations(writer) {
return
}
result, err := h.operations.SyncModels(request.Context(), request.PathValue("provider_id"), actor.ID)
if err != nil {
h.writeError(writer, err)
return
}
apiresponse.OK(writer, result)
}
func (h *AdminHTTPHandler) rotateCredentials(writer http.ResponseWriter, request *http.Request) {
actor, ok := h.requirePermission(writer, request, identity.PermissionProviderManage)
if !ok {
return
}
if !h.requireOperations(writer) {
return
}
result, err := h.operations.RotateCredentials(request.Context(), actor.ID)
if err != nil {
h.writeError(writer, err)
return
}
if h.changeHook != nil && result.Rotated > 0 {
_ = h.changeHook(request.Context())
}
apiresponse.OK(writer, result)
}
func (h *AdminHTTPHandler) listModelRoutes(writer http.ResponseWriter, request *http.Request) {
if _, ok := h.requirePermission(writer, request, identity.PermissionProviderRead); !ok {
return
}
routes, err := h.repository.ListModelRoutes(request.Context())
if err != nil {
h.writeError(writer, err)
return
}
apiresponse.OK(writer, routes)
}
func (h *AdminHTTPHandler) createModelRoute(writer http.ResponseWriter, request *http.Request) {
actor, ok := h.requirePermission(writer, request, identity.PermissionProviderManage)
if !ok {
return
}
route, err := h.decodeModelRoute(writer, request)
if err != nil {
h.writeError(writer, err)
return
}
created, err := h.repository.CreateModelRoute(request.Context(), route, actor.ID)
if err != nil {
h.writeError(writer, err)
return
}
h.propagateChange(request.Context())
apiresponse.OK(writer, created)
}
func (h *AdminHTTPHandler) updateModelRoute(writer http.ResponseWriter, request *http.Request) {
actor, ok := h.requirePermission(writer, request, identity.PermissionProviderManage)
if !ok {
return
}
route, err := h.decodeModelRoute(writer, request)
if err != nil {
h.writeError(writer, err)
return
}
route.ID = request.PathValue("model_route_id")
updated, err := h.repository.UpdateModelRoute(request.Context(), route, actor.ID)
if err != nil {
h.writeError(writer, err)
return
}
h.propagateChange(request.Context())
apiresponse.OK(writer, updated)
}
func (h *AdminHTTPHandler) deleteModelRoute(writer http.ResponseWriter, request *http.Request) {
actor, ok := h.requirePermission(writer, request, identity.PermissionProviderManage)
if !ok {
return
}
if err := h.repository.DeleteModelRoute(request.Context(), request.PathValue("model_route_id"), actor.ID); err != nil {
h.writeError(writer, err)
return
}
h.propagateChange(request.Context())
apiresponse.OK(writer, map[string]bool{"deleted": true})
}
func (h *AdminHTTPHandler) decodeModelRoute(writer http.ResponseWriter, request *http.Request) (ModelRoute, error) {
var input modelRouteInput
decoder := json.NewDecoder(http.MaxBytesReader(writer, request.Body, 1<<20))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&input); err != nil {
return ModelRoute{}, errors.New("请求格式无效")
}
input.Name = strings.TrimSpace(input.Name)
input.SourceModel = strings.TrimSpace(input.SourceModel)
input.TargetModel = strings.TrimSpace(input.TargetModel)
input.ProviderID = strings.TrimSpace(input.ProviderID)
if input.Name == "" || len(input.Name) > 128 || input.SourceModel == "" || len(input.SourceModel) > 512 || input.TargetModel == "" || len(input.TargetModel) > 512 || input.ProviderID == "" {
return ModelRoute{}, errors.New("名称、模型别名、上游模型或供应商无效")
}
if input.Weight == 0 {
input.Weight = 100
}
if input.Weight < 1 || input.Weight > 10000 || input.Priority < -100000 || input.Priority > 100000 {
return ModelRoute{}, errors.New("权重或优先级超出允许范围")
}
if len(input.Conditions) == 0 || string(input.Conditions) == "null" {
input.Conditions = json.RawMessage(`{}`)
}
var conditions modelRouteConditions
if err := json.Unmarshal(input.Conditions, &conditions); err != nil {
return ModelRoute{}, errors.New("conditions 必须是 JSON 对象")
}
allowedEndpoints := map[string]bool{"/v1/chat/completions": true, "/v1/responses": true, "/v1/embeddings": true, "/v1/messages": true}
for _, endpoint := range conditions.Endpoints {
if !allowedEndpoints[endpoint] {
return ModelRoute{}, errors.New("conditions.endpoints 包含不支持的网关端点")
}
}
enabled := true
if input.Enabled != nil {
enabled = *input.Enabled
}
return ModelRoute{Name: input.Name, SourceModel: input.SourceModel, TargetModel: input.TargetModel, ProviderID: input.ProviderID, Weight: input.Weight, Priority: input.Priority, Conditions: input.Conditions, Enabled: enabled}, nil
}
func (h *AdminHTTPHandler) propagateChange(ctx context.Context) {
if h.changeHook != nil {
_ = h.changeHook(ctx)
}
}
func (h *AdminHTTPHandler) requireOperations(writer http.ResponseWriter) bool {
if h.operations == nil {
apiresponse.Error(writer, http.StatusServiceUnavailable, "供应商控制面服务暂不可用")
return false
}
return true
}
func (h *AdminHTTPHandler) requirePermission(writer http.ResponseWriter, request *http.Request, permission string) (identity.Account, bool) {
account, err := h.identity.Authenticate(request.Context(), identity.KindAdmin, request.Header.Get("Authorization"))
if err != nil {
if errors.Is(err, identity.ErrInvalidSession) || errors.Is(err, identity.ErrNotFound) {
apiresponse.Error(writer, http.StatusUnauthorized, "登录状态无效或已过期")
} else if errors.Is(err, identity.ErrAccountDisabled) {
apiresponse.Error(writer, http.StatusForbidden, "管理员账号已被停用")
} else {
apiresponse.Error(writer, http.StatusServiceUnavailable, "身份服务暂不可用")
}
return identity.Account{}, false
}
if !identity.HasPermission(account, permission) {
apiresponse.Error(writer, http.StatusForbidden, "缺少模型供应商操作权限")
return identity.Account{}, false
}
return account, true
}
func (h *AdminHTTPHandler) writeError(writer http.ResponseWriter, err error) {
switch {
case errors.Is(err, ErrProviderNotFound):
apiresponse.Error(writer, http.StatusNotFound, "供应商不存在")
case errors.Is(err, ErrProviderExists):
apiresponse.Error(writer, http.StatusConflict, "供应商 code 已存在")
case errors.Is(err, ErrMultipleDefaults):
apiresponse.Error(writer, http.StatusConflict, "只能启用一个默认供应商")
case errors.Is(err, ErrModelRouteNotFound):
apiresponse.Error(writer, http.StatusNotFound, "模型路由不存在")
case errors.Is(err, ErrModelRouteExists):
apiresponse.Error(writer, http.StatusConflict, "相同模型、供应商和上游模型的路由已存在")
case errors.Is(err, ErrProviderStore), errors.Is(err, ErrCredentialKeyUnavailable):
apiresponse.Error(writer, http.StatusServiceUnavailable, "供应商配置服务暂不可用")
case errors.Is(err, ErrProviderUpstream):
apiresponse.Error(writer, http.StatusBadGateway, "无法从上游供应商获取模型信息")
case errors.Is(err, ErrBlockedAddress):
// 原始错误含解析出的地址(如 "blocked address 10.0.0.1"),泄露内网
// 拓扑;细节只进服务端日志,客户端返回通用提示。
h.logger.Warn("provider URL blocked", "error", err)
apiresponse.Error(writer, http.StatusBadRequest, "供应商地址不允许访问内网或保留网段")
default:
apiresponse.Error(writer, http.StatusBadRequest, err.Error())
}
}
func maskSecret(secret string) string {
if secret == "" {
return ""
}
if len(secret) <= 10 {
return "••••••••"
}
return secret[:4] + "••••••" + secret[len(secret)-4:]
}