184 lines
7.1 KiB
Go
184 lines
7.1 KiB
Go
package repository
|
|
|
|
import (
|
|
"database/sql"
|
|
"errors"
|
|
"fmt"
|
|
|
|
"cmautobuy/admin/model"
|
|
"github.com/go-sql-driver/mysql"
|
|
)
|
|
|
|
var ErrAIProviderNameExists = errors.New("AI 服务商名称已经存在")
|
|
|
|
func ListAIProviders(q Execer) ([]model.AIProviderConfig, error) {
|
|
rows, err := q.Query(`SELECT provider_id,name,base_url,model,timeout_seconds,max_concurrency,
|
|
confidence_threshold_bps,enabled,last_test_status,last_test_message,last_tested_at,
|
|
last_test_fingerprint,created_by_user_id,updated_by_user_id,created_at,updated_at
|
|
FROM ai_provider_configs ORDER BY enabled DESC,name,provider_id`)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("查询 AI 服务商配置失败: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
var result []model.AIProviderConfig
|
|
for rows.Next() {
|
|
item, err := scanAIProvider(rows.Scan)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
result = append(result, item)
|
|
}
|
|
return result, rows.Err()
|
|
}
|
|
|
|
func GetAIProvider(q Execer, providerID string) (*model.AIProviderConfig, error) {
|
|
item, err := scanAIProvider(func(dest ...any) error {
|
|
return q.QueryRow(`SELECT provider_id,name,base_url,model,timeout_seconds,max_concurrency,
|
|
confidence_threshold_bps,enabled,last_test_status,last_test_message,last_tested_at,
|
|
last_test_fingerprint,created_by_user_id,updated_by_user_id,created_at,updated_at
|
|
FROM ai_provider_configs WHERE provider_id=?`, providerID).Scan(dest...)
|
|
})
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
if err != nil {
|
|
return nil, fmt.Errorf("查询 AI 服务商配置失败: %w", err)
|
|
}
|
|
return &item, nil
|
|
}
|
|
|
|
func GetEnabledAIProvider(q Execer) (*model.AIProviderConfig, error) {
|
|
item, err := scanAIProvider(func(dest ...any) error {
|
|
return q.QueryRow(`SELECT provider_id,name,base_url,model,timeout_seconds,max_concurrency,
|
|
confidence_threshold_bps,enabled,last_test_status,last_test_message,last_tested_at,
|
|
last_test_fingerprint,created_by_user_id,updated_by_user_id,created_at,updated_at
|
|
FROM ai_provider_configs WHERE enabled=1`).Scan(dest...)
|
|
})
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
if err != nil {
|
|
return nil, fmt.Errorf("查询当前 AI 服务商失败: %w", err)
|
|
}
|
|
return &item, nil
|
|
}
|
|
|
|
type scanValues func(dest ...any) error
|
|
|
|
func scanAIProvider(scan scanValues) (model.AIProviderConfig, error) {
|
|
var item model.AIProviderConfig
|
|
var enabled int
|
|
var lastMessage, lastTestedAt, fingerprint sql.NullString
|
|
err := scan(&item.ProviderID, &item.Name, &item.BaseURL, &item.Model, &item.TimeoutSeconds,
|
|
&item.MaxConcurrency, &item.ConfidenceThresholdBPS, &enabled, &item.LastTestStatus,
|
|
&lastMessage, &lastTestedAt, &fingerprint, &item.CreatedByUserID, &item.UpdatedByUserID,
|
|
&item.CreatedAt, &item.UpdatedAt)
|
|
item.Enabled = enabled == 1
|
|
item.LastTestMessage = lastMessage.String
|
|
item.LastTestedAt = lastTestedAt.String
|
|
item.LastTestFingerprint = fingerprint.String
|
|
return item, err
|
|
}
|
|
|
|
func InsertAIProvider(q Execer, item model.AIProviderConfig) error {
|
|
_, err := q.Exec(`INSERT INTO ai_provider_configs
|
|
(provider_id,name,base_url,model,timeout_seconds,max_concurrency,confidence_threshold_bps,
|
|
enabled,last_test_status,created_by_user_id,updated_by_user_id,created_at,updated_at)
|
|
VALUES(?,?,?,?,?,?,?,0,'pending',?,?,?,?)`, item.ProviderID, item.Name, item.BaseURL,
|
|
item.Model, item.TimeoutSeconds, item.MaxConcurrency, item.ConfidenceThresholdBPS,
|
|
item.CreatedByUserID, item.UpdatedByUserID, item.CreatedAt, item.UpdatedAt)
|
|
return aiProviderWriteError("新增 AI 服务商配置", err)
|
|
}
|
|
|
|
func UpdateAIProvider(q Execer, item model.AIProviderConfig) (bool, error) {
|
|
result, err := q.Exec(`UPDATE ai_provider_configs SET name=?,base_url=?,model=?,timeout_seconds=?,
|
|
max_concurrency=?,confidence_threshold_bps=?,last_test_status='pending',last_test_message=NULL,
|
|
last_tested_at=NULL,last_test_fingerprint=NULL,enabled=0,updated_by_user_id=?,updated_at=? WHERE provider_id=?`,
|
|
item.Name, item.BaseURL, item.Model, item.TimeoutSeconds, item.MaxConcurrency,
|
|
item.ConfidenceThresholdBPS, item.UpdatedByUserID, item.UpdatedAt, item.ProviderID)
|
|
if err != nil {
|
|
return false, aiProviderWriteError("更新 AI 服务商配置", err)
|
|
}
|
|
n, err := result.RowsAffected()
|
|
return n == 1, err
|
|
}
|
|
|
|
func RecordAIProviderTest(q Execer, providerID, status, message, testedAt, fingerprint, actorID string) (bool, error) {
|
|
result, err := q.Exec(`UPDATE ai_provider_configs SET last_test_status=?,last_test_message=?,
|
|
last_tested_at=?,last_test_fingerprint=?,updated_by_user_id=?,updated_at=? WHERE provider_id=?`,
|
|
status, nullableText(message), testedAt, nullableText(fingerprint), actorID, testedAt, providerID)
|
|
if err != nil {
|
|
return false, fmt.Errorf("记录 AI 连接测试失败: %w", err)
|
|
}
|
|
n, err := result.RowsAffected()
|
|
return n == 1, err
|
|
}
|
|
|
|
func InvalidateAIProviderTest(q Execer, providerID, actorID, updatedAt string) (bool, error) {
|
|
result, err := q.Exec(`UPDATE ai_provider_configs SET last_test_status='pending',last_test_message=NULL,
|
|
last_tested_at=NULL,last_test_fingerprint=NULL,enabled=0,updated_by_user_id=?,updated_at=? WHERE provider_id=?`,
|
|
actorID, updatedAt, providerID)
|
|
if err != nil {
|
|
return false, fmt.Errorf("重置 AI 连接测试状态失败: %w", err)
|
|
}
|
|
n, err := result.RowsAffected()
|
|
return n == 1, err
|
|
}
|
|
|
|
func LockAIProviders(q Execer) error {
|
|
rows, err := q.Query(`SELECT provider_id FROM ai_provider_configs FOR UPDATE`)
|
|
if err != nil {
|
|
return fmt.Errorf("锁定 AI 服务商配置失败: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
for rows.Next() {
|
|
var providerID string
|
|
if err := rows.Scan(&providerID); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
return rows.Err()
|
|
}
|
|
|
|
func EnableAIProvider(q Execer, providerID, actorID, updatedAt string) (bool, error) {
|
|
if _, err := q.Exec(`UPDATE ai_provider_configs SET enabled=0,updated_by_user_id=?,updated_at=? WHERE enabled=1`, actorID, updatedAt); err != nil {
|
|
return false, fmt.Errorf("停用旧 AI 服务商失败: %w", err)
|
|
}
|
|
result, err := q.Exec(`UPDATE ai_provider_configs SET enabled=1,updated_by_user_id=?,updated_at=? WHERE provider_id=?`, actorID, updatedAt, providerID)
|
|
if err != nil {
|
|
return false, fmt.Errorf("启用 AI 服务商失败: %w", err)
|
|
}
|
|
n, err := result.RowsAffected()
|
|
return n == 1, err
|
|
}
|
|
|
|
func DisableAIProvider(q Execer, providerID, actorID, updatedAt string) (bool, error) {
|
|
result, err := q.Exec(`UPDATE ai_provider_configs SET enabled=0,updated_by_user_id=?,updated_at=? WHERE provider_id=?`, actorID, updatedAt, providerID)
|
|
if err != nil {
|
|
return false, fmt.Errorf("停用 AI 服务商失败: %w", err)
|
|
}
|
|
n, err := result.RowsAffected()
|
|
return n == 1, err
|
|
}
|
|
|
|
func InsertAIProviderAudit(q Execer, audit model.AIProviderAudit) error {
|
|
_, err := q.Exec(`INSERT INTO ai_provider_audits(provider_id,action,details_json,actor_user_id,created_at)
|
|
VALUES(?,?,?,?,?)`, nullableText(audit.ProviderID), audit.Action, audit.DetailsJSON,
|
|
audit.ActorUserID, audit.CreatedAt)
|
|
if err != nil {
|
|
return fmt.Errorf("写入 AI 配置审计失败: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func aiProviderWriteError(action string, err error) error {
|
|
if err == nil {
|
|
return nil
|
|
}
|
|
var mysqlErr *mysql.MySQLError
|
|
if errors.As(err, &mysqlErr) && mysqlErr.Number == 1062 {
|
|
return ErrAIProviderNameExists
|
|
}
|
|
return fmt.Errorf("%s失败: %w", action, err)
|
|
}
|