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:] }