feat: 增加管理员 AI 模型配置 (#200)
This commit is contained in:
@@ -0,0 +1,460 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"cmautobuy/admin/model"
|
||||
"cmautobuy/admin/repository"
|
||||
)
|
||||
|
||||
type AIProviderInput struct {
|
||||
ProviderID string
|
||||
Name string
|
||||
BaseURL string
|
||||
Model string
|
||||
TimeoutSeconds int
|
||||
MaxConcurrency int
|
||||
ConfidenceThresholdBPS int
|
||||
}
|
||||
|
||||
type AIProviderListResult struct {
|
||||
Items []model.AIProviderConfig
|
||||
SecretStoreError string
|
||||
}
|
||||
|
||||
type AIEndpointPolicy struct {
|
||||
allowedHosts map[string]bool
|
||||
resolver *net.Resolver
|
||||
}
|
||||
|
||||
func NewAIEndpointPolicy(allowedHosts []string) AIEndpointPolicy {
|
||||
policy := AIEndpointPolicy{allowedHosts: map[string]bool{}, resolver: net.DefaultResolver}
|
||||
for _, host := range allowedHosts {
|
||||
if value := strings.ToLower(strings.TrimSpace(host)); value != "" {
|
||||
policy.allowedHosts[value] = true
|
||||
}
|
||||
}
|
||||
return policy
|
||||
}
|
||||
|
||||
func ListAIProviderConfigs(db *sql.DB, actor *model.User, secrets AISecretStore) (AIProviderListResult, error) {
|
||||
if actor == nil || !actor.IsAdmin() {
|
||||
return AIProviderListResult{}, ErrAdminRequired
|
||||
}
|
||||
items, err := repository.ListAIProviders(db)
|
||||
if err != nil {
|
||||
return AIProviderListResult{}, err
|
||||
}
|
||||
result := AIProviderListResult{Items: items}
|
||||
for i := range result.Items {
|
||||
configured, suffix, statusErr := secrets.Status(result.Items[i].ProviderID)
|
||||
if statusErr != nil {
|
||||
result.SecretStoreError = statusErr.Error()
|
||||
break
|
||||
}
|
||||
result.Items[i].SecretConfigured = configured
|
||||
result.Items[i].SecretSuffix = suffix
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func SaveAIProviderConfig(db *sql.DB, actor *model.User, input AIProviderInput, policy AIEndpointPolicy, now time.Time) (string, error) {
|
||||
if actor == nil || !actor.IsAdmin() {
|
||||
return "", ErrAdminRequired
|
||||
}
|
||||
item, err := validateAIProviderInput(input, policy)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
item.ProviderID = strings.TrimSpace(input.ProviderID)
|
||||
if item.ProviderID == "" {
|
||||
item.ProviderID, err = randomID("AIP-", 16)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("生成 AI 服务商编号失败: %w", err)
|
||||
}
|
||||
}
|
||||
at := now.UTC().Format(model.TimeLayout)
|
||||
item.UpdatedByUserID, item.UpdatedAt = actor.UserID, at
|
||||
tx, err := db.Begin()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
action := "update"
|
||||
details := map[string]any{"fields": []string{"name", "base_url", "model", "timeout_seconds", "max_concurrency", "confidence_threshold_bps"}}
|
||||
if input.ProviderID == "" {
|
||||
action = "create"
|
||||
item.CreatedByUserID, item.CreatedAt = actor.UserID, at
|
||||
if err := repository.InsertAIProvider(tx, item); err != nil {
|
||||
return "", aiProviderValidationError(err)
|
||||
}
|
||||
} else {
|
||||
found, err := repository.UpdateAIProvider(tx, item)
|
||||
if err != nil {
|
||||
return "", aiProviderValidationError(err)
|
||||
}
|
||||
if !found {
|
||||
return "", &validationError{field: "provider_id", message: "AI 服务商不存在,请刷新页面后重试"}
|
||||
}
|
||||
}
|
||||
if err := insertAIConfigAudit(tx, item.ProviderID, action, details, actor.UserID, at); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return item.ProviderID, nil
|
||||
}
|
||||
|
||||
func SetAIProviderSecret(db *sql.DB, actor *model.User, secrets AISecretStore, providerID, apiKey string, now time.Time) error {
|
||||
if actor == nil || !actor.IsAdmin() {
|
||||
return ErrAdminRequired
|
||||
}
|
||||
providerID = strings.TrimSpace(providerID)
|
||||
if err := invalidateAIProviderForSecretChange(db, actor, providerID, now); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := secrets.Set(providerID, apiKey); err != nil {
|
||||
return err
|
||||
}
|
||||
return insertAIConfigAudit(db, providerID, "secret_replace", map[string]any{"secret": "configured"}, actor.UserID, now.UTC().Format(model.TimeLayout))
|
||||
}
|
||||
|
||||
func ClearAIProviderSecret(db *sql.DB, actor *model.User, secrets AISecretStore, providerID string, now time.Time) error {
|
||||
if actor == nil || !actor.IsAdmin() {
|
||||
return ErrAdminRequired
|
||||
}
|
||||
providerID = strings.TrimSpace(providerID)
|
||||
if err := invalidateAIProviderForSecretChange(db, actor, providerID, now); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := secrets.Clear(providerID); err != nil {
|
||||
return err
|
||||
}
|
||||
return insertAIConfigAudit(db, providerID, "secret_clear", map[string]any{"secret": "cleared"}, actor.UserID, now.UTC().Format(model.TimeLayout))
|
||||
}
|
||||
|
||||
func invalidateAIProviderForSecretChange(db *sql.DB, actor *model.User, providerID string, now time.Time) error {
|
||||
if providerID == "" {
|
||||
return &validationError{field: "provider_id", message: "AI 服务商编号不能为空"}
|
||||
}
|
||||
found, err := repository.InvalidateAIProviderTest(db, providerID, actor.UserID, now.UTC().Format(model.TimeLayout))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !found {
|
||||
return &validationError{field: "provider_id", message: "AI 服务商不存在,请刷新页面后重试"}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type AIHTTPDoer interface {
|
||||
Do(*http.Request) (*http.Response, error)
|
||||
}
|
||||
|
||||
func TestAIProviderConnection(ctx context.Context, db *sql.DB, actor *model.User, secrets AISecretStore,
|
||||
providerID string, policy AIEndpointPolicy, client AIHTTPDoer, now time.Time) error {
|
||||
if actor == nil || !actor.IsAdmin() {
|
||||
return ErrAdminRequired
|
||||
}
|
||||
provider, err := repository.GetAIProvider(db, strings.TrimSpace(providerID))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if provider == nil {
|
||||
return &validationError{field: "provider_id", message: "AI 服务商不存在,请刷新页面后重试"}
|
||||
}
|
||||
secret, err := secrets.Get(provider.ProviderID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if secret == "" {
|
||||
return &validationError{field: "api_key", message: "请先保存 API Key"}
|
||||
}
|
||||
if err := policy.ValidateResolved(ctx, provider.BaseURL); err != nil {
|
||||
return err
|
||||
}
|
||||
if client == nil {
|
||||
client = NewSafeAIHTTPClient(policy, time.Duration(provider.TimeoutSeconds)*time.Second)
|
||||
}
|
||||
testErr := callAIHealthCheck(ctx, client, *provider, secret)
|
||||
status, message, fingerprint := "succeeded", "连接正常", aiProviderFingerprint(*provider)
|
||||
if testErr != nil {
|
||||
status, message, fingerprint = "failed", safeAIError(testErr), ""
|
||||
}
|
||||
at := now.UTC().Format(model.TimeLayout)
|
||||
found, recordErr := repository.RecordAIProviderTest(db, provider.ProviderID, status, message, at, fingerprint, actor.UserID)
|
||||
if recordErr != nil {
|
||||
return recordErr
|
||||
}
|
||||
if !found {
|
||||
return &validationError{field: "provider_id", message: "AI 服务商不存在,请刷新页面后重试"}
|
||||
}
|
||||
if err := insertAIConfigAudit(db, provider.ProviderID, "test_"+status, map[string]any{"status": status}, actor.UserID, at); err != nil {
|
||||
return err
|
||||
}
|
||||
if testErr != nil {
|
||||
return &validationError{field: "connection", message: "连接测试失败:" + message}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func EnableAIProviderConfig(db *sql.DB, actor *model.User, secrets AISecretStore, providerID string, now time.Time) error {
|
||||
if actor == nil || !actor.IsAdmin() {
|
||||
return ErrAdminRequired
|
||||
}
|
||||
tx, err := db.Begin()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
if err := repository.LockAIProviders(tx); err != nil {
|
||||
return err
|
||||
}
|
||||
provider, err := repository.GetAIProvider(tx, strings.TrimSpace(providerID))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if provider == nil {
|
||||
return &validationError{field: "provider_id", message: "AI 服务商不存在,请刷新页面后重试"}
|
||||
}
|
||||
configured, _, err := secrets.Status(provider.ProviderID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !configured {
|
||||
return &validationError{field: "api_key", message: "请先保存 API Key 并测试连接"}
|
||||
}
|
||||
if provider.LastTestStatus != "succeeded" || provider.LastTestFingerprint != aiProviderFingerprint(*provider) {
|
||||
return &validationError{field: "connection", message: "当前配置尚未通过连接测试,不能启用"}
|
||||
}
|
||||
at := now.UTC().Format(model.TimeLayout)
|
||||
found, err := repository.EnableAIProvider(tx, provider.ProviderID, actor.UserID, at)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !found {
|
||||
return &validationError{field: "provider_id", message: "AI 服务商不存在,请刷新页面后重试"}
|
||||
}
|
||||
if err := insertAIConfigAudit(tx, provider.ProviderID, "enable", map[string]any{"enabled": true}, actor.UserID, at); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
func DisableAIProviderConfig(db *sql.DB, actor *model.User, providerID string, now time.Time) error {
|
||||
if actor == nil || !actor.IsAdmin() {
|
||||
return ErrAdminRequired
|
||||
}
|
||||
at := now.UTC().Format(model.TimeLayout)
|
||||
found, err := repository.DisableAIProvider(db, strings.TrimSpace(providerID), actor.UserID, at)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !found {
|
||||
return &validationError{field: "provider_id", message: "AI 服务商不存在,请刷新页面后重试"}
|
||||
}
|
||||
return insertAIConfigAudit(db, strings.TrimSpace(providerID), "disable", map[string]any{"enabled": false}, actor.UserID, at)
|
||||
}
|
||||
|
||||
func validateAIProviderInput(input AIProviderInput, policy AIEndpointPolicy) (model.AIProviderConfig, error) {
|
||||
item := model.AIProviderConfig{Name: strings.TrimSpace(input.Name), Model: strings.TrimSpace(input.Model),
|
||||
TimeoutSeconds: input.TimeoutSeconds, MaxConcurrency: input.MaxConcurrency,
|
||||
ConfidenceThresholdBPS: input.ConfidenceThresholdBPS}
|
||||
if item.Name == "" || utf8.RuneCountInString(item.Name) > 191 {
|
||||
return item, &validationError{field: "name", message: "服务商名称不能为空且最多 191 个字符"}
|
||||
}
|
||||
if item.Model == "" || utf8.RuneCountInString(item.Model) > 191 {
|
||||
return item, &validationError{field: "model", message: "模型名称不能为空且最多 191 个字符"}
|
||||
}
|
||||
baseURL, err := policy.ValidateSyntax(input.BaseURL)
|
||||
if err != nil {
|
||||
return item, err
|
||||
}
|
||||
item.BaseURL = baseURL
|
||||
if item.TimeoutSeconds < 1 || item.TimeoutSeconds > 120 {
|
||||
return item, &validationError{field: "timeout_seconds", message: "超时秒数必须在 1 到 120 之间"}
|
||||
}
|
||||
if item.MaxConcurrency < 1 || item.MaxConcurrency > 16 {
|
||||
return item, &validationError{field: "max_concurrency", message: "最大并发必须在 1 到 16 之间"}
|
||||
}
|
||||
if item.ConfidenceThresholdBPS < 0 || item.ConfidenceThresholdBPS > 10000 {
|
||||
return item, &validationError{field: "confidence_threshold", message: "自动写入置信度必须在 0% 到 100% 之间"}
|
||||
}
|
||||
return item, nil
|
||||
}
|
||||
|
||||
// ParseConfidenceThresholdBPS 把页面百分比精确转换成基点,避免数据库存浮点数。
|
||||
func ParseConfidenceThresholdBPS(raw string) (int, error) {
|
||||
raw = strings.TrimSpace(raw)
|
||||
parts := strings.Split(raw, ".")
|
||||
if raw == "" || len(parts) > 2 {
|
||||
return 0, &validationError{field: "confidence_threshold", message: "置信度必须是 0 到 100 的数字,最多两位小数"}
|
||||
}
|
||||
whole, err := strconv.Atoi(parts[0])
|
||||
if err != nil || whole < 0 || whole > 100 {
|
||||
return 0, &validationError{field: "confidence_threshold", message: "置信度必须在 0% 到 100% 之间"}
|
||||
}
|
||||
fraction := 0
|
||||
if len(parts) == 2 {
|
||||
if len(parts[1]) == 0 || len(parts[1]) > 2 {
|
||||
return 0, &validationError{field: "confidence_threshold", message: "置信度最多保留两位小数"}
|
||||
}
|
||||
fraction, err = strconv.Atoi(parts[1] + strings.Repeat("0", 2-len(parts[1])))
|
||||
if err != nil {
|
||||
return 0, &validationError{field: "confidence_threshold", message: "置信度格式不正确"}
|
||||
}
|
||||
}
|
||||
if whole == 100 && fraction != 0 {
|
||||
return 0, &validationError{field: "confidence_threshold", message: "置信度不能超过 100%"}
|
||||
}
|
||||
return whole*100 + fraction, nil
|
||||
}
|
||||
|
||||
func (p AIEndpointPolicy) ValidateSyntax(raw string) (string, error) {
|
||||
parsed, err := url.Parse(strings.TrimSpace(raw))
|
||||
if err != nil || parsed.Scheme != "https" || parsed.Host == "" {
|
||||
return "", &validationError{field: "base_url", message: "Base URL 必须是完整的 HTTPS 地址"}
|
||||
}
|
||||
if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" {
|
||||
return "", &validationError{field: "base_url", message: "Base URL 不能包含账号密码、查询参数或片段"}
|
||||
}
|
||||
host := strings.ToLower(parsed.Hostname())
|
||||
if host == "localhost" || host == "metadata.google.internal" {
|
||||
return "", &validationError{field: "base_url", message: "Base URL 指向了受保护的本机或元数据地址"}
|
||||
}
|
||||
if ip := net.ParseIP(host); ip != nil && !p.allowedHosts[host] && !isPublicAIIP(ip) {
|
||||
return "", &validationError{field: "base_url", message: "Base URL 指向私有或本机地址,必须由部署配置显式允许"}
|
||||
}
|
||||
parsed.Path = strings.TrimRight(parsed.Path, "/")
|
||||
parsed.RawPath = strings.TrimRight(parsed.RawPath, "/")
|
||||
return parsed.String(), nil
|
||||
}
|
||||
|
||||
func (p AIEndpointPolicy) ValidateResolved(ctx context.Context, raw string) error {
|
||||
normalized, err := p.ValidateSyntax(raw)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
parsed, _ := url.Parse(normalized)
|
||||
host := strings.ToLower(parsed.Hostname())
|
||||
if p.allowedHosts[host] {
|
||||
return nil
|
||||
}
|
||||
addresses, err := p.resolver.LookupIPAddr(ctx, host)
|
||||
if err != nil || len(addresses) == 0 {
|
||||
return &validationError{field: "base_url", message: "无法解析 AI 服务商地址"}
|
||||
}
|
||||
for _, address := range addresses {
|
||||
if !isPublicAIIP(address.IP) {
|
||||
return &validationError{field: "base_url", message: "AI 服务商域名解析到私有、本机或链路本地地址"}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isPublicAIIP(ip net.IP) bool {
|
||||
return ip != nil && !ip.IsLoopback() && !ip.IsPrivate() && !ip.IsLinkLocalUnicast() &&
|
||||
!ip.IsLinkLocalMulticast() && !ip.IsUnspecified() && !ip.IsMulticast()
|
||||
}
|
||||
|
||||
func NewSafeAIHTTPClient(policy AIEndpointPolicy, timeout time.Duration) *http.Client {
|
||||
dialer := &net.Dialer{Timeout: timeout, KeepAlive: 30 * time.Second}
|
||||
// 不使用环境代理。否则实际拨号只会校验代理地址,目标地址可能绕过 SSRF 拨号校验。
|
||||
transport := &http.Transport{ForceAttemptHTTP2: true,
|
||||
DialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
host, _, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("AI 服务商网络地址无效")
|
||||
}
|
||||
if !policy.allowedHosts[strings.ToLower(host)] {
|
||||
addresses, err := policy.resolver.LookupIPAddr(ctx, host)
|
||||
if err != nil || len(addresses) == 0 {
|
||||
return nil, fmt.Errorf("无法解析 AI 服务商地址")
|
||||
}
|
||||
for _, item := range addresses {
|
||||
if !isPublicAIIP(item.IP) {
|
||||
return nil, fmt.Errorf("AI 服务商地址不在允许范围")
|
||||
}
|
||||
}
|
||||
}
|
||||
return dialer.DialContext(ctx, network, address)
|
||||
}}
|
||||
return &http.Client{Transport: transport, Timeout: timeout, CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
if len(via) >= 3 {
|
||||
return fmt.Errorf("AI 服务商重定向次数过多")
|
||||
}
|
||||
return policy.ValidateResolved(req.Context(), req.URL.String())
|
||||
}}
|
||||
}
|
||||
|
||||
func callAIHealthCheck(ctx context.Context, client AIHTTPDoer, provider model.AIProviderConfig, secret string) error {
|
||||
payload, _ := json.Marshal(map[string]any{"model": provider.Model, "messages": []map[string]string{{"role": "user", "content": "Reply with OK."}}, "max_tokens": 1, "temperature": 0})
|
||||
endpoint := strings.TrimRight(provider.BaseURL, "/") + "/chat/completions"
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return fmt.Errorf("准备连接测试失败")
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+secret)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("请求失败")
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 32<<10))
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return fmt.Errorf("服务返回 HTTP %d", resp.StatusCode)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func aiProviderFingerprint(provider model.AIProviderConfig) string {
|
||||
value := strings.Join([]string{provider.BaseURL, provider.Model, strconv.Itoa(provider.TimeoutSeconds),
|
||||
strconv.Itoa(provider.MaxConcurrency), strconv.Itoa(provider.ConfidenceThresholdBPS)}, "\x00")
|
||||
sum := sha256.Sum256([]byte(value))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func insertAIConfigAudit(q repository.Execer, providerID, action string, details any, actorID, at string) error {
|
||||
raw, err := json.Marshal(details)
|
||||
if err != nil {
|
||||
return fmt.Errorf("准备 AI 配置审计失败: %w", err)
|
||||
}
|
||||
return repository.InsertAIProviderAudit(q, model.AIProviderAudit{ProviderID: providerID, Action: action,
|
||||
DetailsJSON: string(raw), ActorUserID: actorID, CreatedAt: at})
|
||||
}
|
||||
|
||||
func aiProviderValidationError(err error) error {
|
||||
if errors.Is(err, repository.ErrAIProviderNameExists) {
|
||||
return &validationError{field: "name", message: "服务商名称已经存在"}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func safeAIError(err error) string {
|
||||
message := strings.TrimSpace(err.Error())
|
||||
if message == "" {
|
||||
message = "未知错误"
|
||||
}
|
||||
runes := []rune(message)
|
||||
if len(runes) > 200 {
|
||||
message = string(runes[:200])
|
||||
}
|
||||
return message
|
||||
}
|
||||
Reference in New Issue
Block a user