2026-08-14 09:49:11 +08:00
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 ))
2026-08-14 11:15:50 +08:00
if err != nil || parsed . Host == "" {
return "" , & validationError { field : "base_url" , message : "Base URL 必须是完整的 HTTP 或 HTTPS 地址" }
}
parsed . Scheme = strings . ToLower ( parsed . Scheme )
if parsed . Scheme != "http" && parsed . Scheme != "https" {
return "" , & validationError { field : "base_url" , message : "Base URL 只支持 HTTP 或 HTTPS" }
2026-08-14 09:49:11 +08:00
}
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
}