5759c1862e
M0-M7 已完成:核心网关(身份/RBAC/TOTP/OIDC/SAML/Provider/配额/路由/内容策略/审计/定价)+ 资源市场(MCP/Skills/数字员工)。 含 22 个 PostgreSQL 迁移、管理端/门户端前端源码、OpenAPI 契约、部署 compose。 Co-Authored-By: Claude <noreply@anthropic.com>
499 lines
17 KiB
Go
499 lines
17 KiB
Go
package provider
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"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
|
|
}
|
|
|
|
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(),
|
|
}
|
|
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("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) 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) {
|
|
plaintext, err := h.cipher.Decrypt(record.EncryptedCredentials, record.CredentialKEKVersion)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
var credentials Credentials
|
|
if err := json.Unmarshal(plaintext, &credentials); err != nil {
|
|
return nil, err
|
|
}
|
|
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": credentials.APIKey != "", "api_key_masked": maskSecret(credentials.APIKey),
|
|
}, 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, "无法从上游供应商获取模型信息")
|
|
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:]
|
|
}
|