880 lines
30 KiB
Go
880 lines
30 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"errors"
|
|
"fmt"
|
|
"log"
|
|
"sort"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"gorm.io/gorm"
|
|
"gorm.io/gorm/clause"
|
|
|
|
"oci-portal/internal/aiwire"
|
|
"oci-portal/internal/model"
|
|
"oci-portal/internal/oci"
|
|
)
|
|
|
|
const (
|
|
aiLogRetention = 90 * 24 * time.Hour
|
|
aiLogMaxRows = 50000
|
|
aiLogCleanupTick = 24 * time.Hour
|
|
// aiFailThreshold 起连续失败次数触发熔断,退避 2^(n-阈值) 分钟,封顶 30 分钟
|
|
aiFailThreshold = 5
|
|
aiBackoffCap = 30 * time.Minute
|
|
// aiKeyTouchGap 是 LastUsedAt 的最小写库间隔,避免高频调用刷库
|
|
aiKeyTouchGap = time.Minute
|
|
// modelRecheckGap 是已标记不可用模型的复检间隔(恢复供给自动解除标记);
|
|
// validateBatchCap 限制单渠道单轮验证的试调次数
|
|
modelRecheckGap = 20 * time.Hour
|
|
validateBatchCap = 32
|
|
// 内容日志(红线例外)约束:开启必须限时(上限 7 天),正文截断,短保留
|
|
aiContentLogMaxHours = 168
|
|
aiContentLogRetention = 7 * 24 * time.Hour
|
|
aiContentLogMaxRows = 10000
|
|
aiContentBodyLimit = 64 * 1024
|
|
)
|
|
|
|
var (
|
|
// ErrAiKeyInvalid 表示网关密钥不存在或已禁用。
|
|
ErrAiKeyInvalid = errors.New("无效或已禁用的 API 密钥")
|
|
// ErrAiUnknownModel 表示没有任何渠道支持请求的模型。
|
|
ErrAiUnknownModel = errors.New("未知模型:没有渠道提供该模型")
|
|
// ErrAiNoChannel 表示模型有渠道支持但当前全部不可用(禁用/熔断)。
|
|
ErrAiNoChannel = errors.New("暂无可用渠道,请稍后重试")
|
|
)
|
|
|
|
// AiGatewayService 是 AI 网关核心:密钥、渠道号池、模型缓存与调用编排。
|
|
type AiGatewayService struct {
|
|
db *gorm.DB
|
|
configs *OciConfigService
|
|
client oci.Client
|
|
wg sync.WaitGroup
|
|
// touchMu 保护各密钥的最近触达时间(内存节流,不追求跨实例精确)
|
|
touchMu sync.Mutex
|
|
lastTouch map[uint]time.Time
|
|
// validateMu 保护 validating:同一渠道的模型验证不并发
|
|
validateMu sync.Mutex
|
|
validating map[uint]bool
|
|
// onChannelsChanged 在渠道增删后触发,由 main 装配为探测任务同步钩子
|
|
onChannelsChanged func(context.Context)
|
|
}
|
|
|
|
// NewAiGatewayService 组装依赖;调用 StartCleanup 后开始调用日志周期清理。
|
|
func NewAiGatewayService(db *gorm.DB, configs *OciConfigService, client oci.Client) *AiGatewayService {
|
|
return &AiGatewayService{db: db, configs: configs, client: client,
|
|
lastTouch: map[uint]time.Time{}, validating: map[uint]bool{}}
|
|
}
|
|
|
|
// SetOnChannelsChanged 注册渠道数量变化钩子(渠道创建/删除成功后调用)。
|
|
func (s *AiGatewayService) SetOnChannelsChanged(fn func(context.Context)) {
|
|
s.onChannelsChanged = fn
|
|
}
|
|
|
|
func (s *AiGatewayService) fireChannelsChanged(ctx context.Context) {
|
|
if s.onChannelsChanged != nil {
|
|
s.onChannelsChanged(ctx)
|
|
}
|
|
}
|
|
|
|
// ---- 网关密钥 ----
|
|
|
|
// CreateKey 生成网关密钥并返回明文(仅此一次);customValue 非空时使用给定值,
|
|
// group 非空时该密钥只在同分组渠道内路由。
|
|
func (s *AiGatewayService) CreateKey(ctx context.Context, name, customValue, group string) (string, *model.AiKey, error) {
|
|
name = strings.TrimSpace(name)
|
|
if name == "" {
|
|
return "", nil, fmt.Errorf("密钥名称不能为空")
|
|
}
|
|
raw := strings.TrimSpace(customValue)
|
|
if raw == "" {
|
|
buf := make([]byte, 24)
|
|
if _, err := rand.Read(buf); err != nil {
|
|
return "", nil, fmt.Errorf("generate key: %w", err)
|
|
}
|
|
raw = "sk-" + hex.EncodeToString(buf)
|
|
}
|
|
if len(raw) < 8 {
|
|
return "", nil, fmt.Errorf("自定义密钥至少 8 个字符")
|
|
}
|
|
key := &model.AiKey{Name: name, KeyHash: hashKey(raw), Tail: raw[len(raw)-4:], Group: strings.TrimSpace(group), Enabled: true}
|
|
if err := s.db.WithContext(ctx).Create(key).Error; err != nil {
|
|
return "", nil, fmt.Errorf("密钥名称或取值与现有密钥重复")
|
|
}
|
|
return raw, key, nil
|
|
}
|
|
|
|
func hashKey(raw string) string {
|
|
sum := sha256.Sum256([]byte(raw))
|
|
return hex.EncodeToString(sum[:])
|
|
}
|
|
|
|
// Keys 列出全部密钥(不含任何明文信息)。
|
|
func (s *AiGatewayService) Keys(ctx context.Context) ([]model.AiKey, error) {
|
|
var keys []model.AiKey
|
|
err := s.db.WithContext(ctx).Order("id DESC").Find(&keys).Error
|
|
return keys, err
|
|
}
|
|
|
|
// UpdateKey 修改密钥名称 / 启用状态 / 分组(group 指针非空即覆盖,可置空)。
|
|
func (s *AiGatewayService) UpdateKey(ctx context.Context, id uint, name string, enabled *bool, group *string) error {
|
|
updates := map[string]any{}
|
|
if name = strings.TrimSpace(name); name != "" {
|
|
updates["name"] = name
|
|
}
|
|
if enabled != nil {
|
|
updates["enabled"] = *enabled
|
|
}
|
|
if group != nil {
|
|
updates["key_group"] = strings.TrimSpace(*group)
|
|
}
|
|
if len(updates) == 0 {
|
|
return nil
|
|
}
|
|
return s.db.WithContext(ctx).Model(&model.AiKey{}).Where("id = ?", id).Updates(updates).Error
|
|
}
|
|
|
|
// DeleteKey 删除密钥,立即使其失效。
|
|
func (s *AiGatewayService) DeleteKey(ctx context.Context, id uint) error {
|
|
return s.db.WithContext(ctx).Delete(&model.AiKey{}, id).Error
|
|
}
|
|
|
|
// VerifyKey 校验请求携带的密钥;通过后节流更新 LastUsedAt。
|
|
func (s *AiGatewayService) VerifyKey(ctx context.Context, raw string) (*model.AiKey, error) {
|
|
if raw == "" {
|
|
return nil, ErrAiKeyInvalid
|
|
}
|
|
var key model.AiKey
|
|
err := s.db.WithContext(ctx).Where("key_hash = ?", hashKey(raw)).First(&key).Error
|
|
if err != nil || !key.Enabled {
|
|
return nil, ErrAiKeyInvalid
|
|
}
|
|
s.touchKey(ctx, key.ID)
|
|
return &key, nil
|
|
}
|
|
|
|
// touchKey 更新最近使用时间,间隔小于 aiKeyTouchGap 时跳过写库。
|
|
func (s *AiGatewayService) touchKey(ctx context.Context, id uint) {
|
|
now := time.Now()
|
|
s.touchMu.Lock()
|
|
last, ok := s.lastTouch[id]
|
|
if ok && now.Sub(last) < aiKeyTouchGap {
|
|
s.touchMu.Unlock()
|
|
return
|
|
}
|
|
s.lastTouch[id] = now
|
|
s.touchMu.Unlock()
|
|
s.db.WithContext(ctx).Model(&model.AiKey{}).Where("id = ?", id).Update("last_used_at", now)
|
|
}
|
|
|
|
// ---- 渠道 ----
|
|
|
|
// ChannelInput 是创建 / 更新渠道的输入;Group 指针非空即覆盖分组(可置空)。
|
|
type ChannelInput struct {
|
|
OciConfigID uint `json:"ociConfigId"`
|
|
Region string `json:"region"`
|
|
Name string `json:"name"`
|
|
Group *string `json:"group"`
|
|
Enabled *bool `json:"enabled"`
|
|
Priority *int `json:"priority"`
|
|
Weight *int `json:"weight"`
|
|
}
|
|
|
|
// CreateChannel 新建渠道(租户×区域唯一);探测由前端随后显式触发。
|
|
func (s *AiGatewayService) CreateChannel(ctx context.Context, in ChannelInput) (*model.AiChannel, error) {
|
|
if in.OciConfigID == 0 || strings.TrimSpace(in.Region) == "" {
|
|
return nil, fmt.Errorf("渠道需要指定租户配置与区域")
|
|
}
|
|
cfg, err := s.configs.Get(ctx, in.OciConfigID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
ch := &model.AiChannel{
|
|
Name: strings.TrimSpace(in.Name),
|
|
OciConfigID: in.OciConfigID,
|
|
Region: strings.TrimSpace(in.Region),
|
|
Enabled: true,
|
|
Priority: valueOr(in.Priority, 1),
|
|
Weight: valueOr(in.Weight, 1),
|
|
}
|
|
if in.Group != nil {
|
|
ch.Group = strings.TrimSpace(*in.Group)
|
|
}
|
|
if ch.Name == "" {
|
|
ch.Name = fmt.Sprintf("%s·%s", cfg.Alias, ch.Region)
|
|
}
|
|
if err := s.db.WithContext(ctx).Create(ch).Error; err != nil {
|
|
return nil, fmt.Errorf("该租户与区域的渠道已存在")
|
|
}
|
|
s.fireChannelsChanged(ctx)
|
|
return ch, nil
|
|
}
|
|
|
|
func valueOr(p *int, def int) int {
|
|
if p != nil {
|
|
return *p
|
|
}
|
|
return def
|
|
}
|
|
|
|
// Channels 列出全部渠道。
|
|
func (s *AiGatewayService) Channels(ctx context.Context) ([]model.AiChannel, error) {
|
|
var chs []model.AiChannel
|
|
err := s.db.WithContext(ctx).Order("priority ASC, id ASC").Find(&chs).Error
|
|
return chs, err
|
|
}
|
|
|
|
// UpdateChannel 修改渠道名称 / 分组 / 启停 / 优先级 / 权重。
|
|
func (s *AiGatewayService) UpdateChannel(ctx context.Context, id uint, in ChannelInput) error {
|
|
updates := map[string]any{}
|
|
if name := strings.TrimSpace(in.Name); name != "" {
|
|
updates["name"] = name
|
|
}
|
|
if in.Group != nil {
|
|
updates["channel_group"] = strings.TrimSpace(*in.Group)
|
|
}
|
|
if in.Enabled != nil {
|
|
updates["enabled"] = *in.Enabled
|
|
}
|
|
if in.Priority != nil {
|
|
updates["priority"] = *in.Priority
|
|
}
|
|
if in.Weight != nil {
|
|
updates["weight"] = *in.Weight
|
|
}
|
|
if len(updates) == 0 {
|
|
return nil
|
|
}
|
|
return s.db.WithContext(ctx).Model(&model.AiChannel{}).Where("id = ?", id).Updates(updates).Error
|
|
}
|
|
|
|
// DeleteChannel 删除渠道并清空其模型缓存。
|
|
func (s *AiGatewayService) DeleteChannel(ctx context.Context, id uint) error {
|
|
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
|
if err := tx.Where("channel_id = ?", id).Delete(&model.AiModelCache{}).Error; err != nil {
|
|
return err
|
|
}
|
|
return tx.Delete(&model.AiChannel{}, id).Error
|
|
})
|
|
if err == nil {
|
|
s.fireChannelsChanged(ctx)
|
|
}
|
|
return err
|
|
}
|
|
|
|
// ---- 探测与模型同步 ----
|
|
|
|
// ProbeChannel 探测渠道可用性:服务可见性 → 模型同步 → maxTokens=1 配额试调。
|
|
func (s *AiGatewayService) ProbeChannel(ctx context.Context, id uint) (*model.AiChannel, error) {
|
|
var ch model.AiChannel
|
|
if err := s.db.WithContext(ctx).First(&ch, id).Error; err != nil {
|
|
return nil, fmt.Errorf("渠道不存在")
|
|
}
|
|
cred, err := s.configs.credentialsByID(ctx, ch.OciConfigID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
status, probeErr := s.probe(ctx, cred, &ch)
|
|
now := time.Now()
|
|
updates := map[string]any{"probe_status": status, "probe_error": probeErr, "last_probe_at": now}
|
|
if status == "ok" {
|
|
updates["fail_count"] = 0
|
|
updates["disabled_until"] = gorm.Expr("NULL")
|
|
}
|
|
if err := s.db.WithContext(ctx).Model(&ch).Updates(updates).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
// 重读用新变量:gorm 扫描 NULL 列到已有值的结构体时会保留旧值
|
|
var fresh model.AiChannel
|
|
if err := s.db.WithContext(ctx).First(&fresh, id).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
return &fresh, nil
|
|
}
|
|
|
|
// probe 执行探测并返回 (状态, 错误摘要);同时完成模型缓存同步(保留不可用标记)。
|
|
func (s *AiGatewayService) probe(ctx context.Context, cred oci.Credentials, ch *model.AiChannel) (string, string) {
|
|
models, err := s.client.ListGenAiModels(ctx, cred, ch.Region)
|
|
if err != nil {
|
|
return classifyProbeErr(err), truncateErr(oci.CompactError(err))
|
|
}
|
|
if len(models) == 0 {
|
|
_ = s.replaceModels(ctx, ch.ID, nil)
|
|
return "no_service", "区域无可用模型(GenAI 服务不可用或未开放)"
|
|
}
|
|
if err := s.replaceModels(ctx, ch.ID, models); err != nil {
|
|
return "error", truncateErr(err.Error())
|
|
}
|
|
usable, err := s.usableModels(ctx, ch.ID, models)
|
|
if err != nil {
|
|
return "error", truncateErr(err.Error())
|
|
}
|
|
return s.probeChat(ctx, cred, ch, usable)
|
|
}
|
|
|
|
// usableModels 过滤掉缓存中已标记不可按需调用的模型,探测候选不再反复踩坑。
|
|
func (s *AiGatewayService) usableModels(ctx context.Context, channelID uint, models []oci.GenAiModel) ([]oci.GenAiModel, error) {
|
|
marks, err := s.loadModelMarks(ctx, channelID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
out := make([]oci.GenAiModel, 0, len(models))
|
|
for _, m := range models {
|
|
if !marks[m.Ocid].Unusable {
|
|
out = append(out, m)
|
|
}
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
// probeChat 按候选顺序试调(上限 8):遇「模型不可按需调用」标记剔除并换下一个;
|
|
// 401/403 与鉴权类 404 属租户级直接定论 no_quota;其余错误(部分模型元数据标 CHAT
|
|
// 但实际不可对话,如 voice agent)累计 3 次止损。
|
|
func (s *AiGatewayService) probeChat(ctx context.Context, cred oci.Credentials, ch *model.AiChannel, models []oci.GenAiModel) (string, string) {
|
|
status, detail := "error", "无可试调对话模型"
|
|
errBudget := 3
|
|
for _, m := range probeCandidates(models) {
|
|
code, err := s.client.GenAiProbeChat(ctx, cred, ch.Region, m.Ocid, m.Name)
|
|
switch {
|
|
case code == 200 || code == 429:
|
|
return "ok", ""
|
|
case oci.IsModelUnavailable(err):
|
|
s.markModelUnusable(ctx, ch.ID, m.Ocid, oci.CompactError(err))
|
|
status, detail = "error", truncateErr(fmt.Sprintf("%s: 不可按需调用,已从模型池剔除", m.Name))
|
|
case code == 401 || code == 403 || code == 404:
|
|
return "no_quota", truncateErr(oci.CompactError(err))
|
|
default:
|
|
status, detail = "error", truncateErr(fmt.Sprintf("%s: %s", m.Name, oci.CompactError(err)))
|
|
if errBudget--; errBudget == 0 {
|
|
return status, detail
|
|
}
|
|
}
|
|
}
|
|
return status, detail
|
|
}
|
|
|
|
// markModelUnusable 把 (渠道, 模型OCID) 标记为不可按需调用;失败仅记日志不阻断主流程。
|
|
func (s *AiGatewayService) markModelUnusable(ctx context.Context, channelID uint, ocid, reason string) {
|
|
err := s.db.WithContext(ctx).Model(&model.AiModelCache{}).
|
|
Where("channel_id = ? AND model_ocid = ?", channelID, ocid).
|
|
Updates(map[string]any{"unusable": true, "unusable_reason": shortReason(reason), "checked_at": time.Now()}).Error
|
|
if err != nil {
|
|
log.Printf("mark model unusable: %v", err)
|
|
}
|
|
}
|
|
|
|
func shortReason(reason string) string {
|
|
if len(reason) > 200 {
|
|
return reason[:200]
|
|
}
|
|
return reason
|
|
}
|
|
|
|
// beginValidate 抢占渠道的验证执行权,同渠道同时只跑一轮。
|
|
func (s *AiGatewayService) beginValidate(id uint) bool {
|
|
s.validateMu.Lock()
|
|
defer s.validateMu.Unlock()
|
|
if s.validating[id] {
|
|
return false
|
|
}
|
|
s.validating[id] = true
|
|
return true
|
|
}
|
|
|
|
func (s *AiGatewayService) endValidate(id uint) {
|
|
s.validateMu.Lock()
|
|
delete(s.validating, id)
|
|
s.validateMu.Unlock()
|
|
}
|
|
|
|
// validateTargets 取待验证行:从未验证的新模型,以及标记超过复检间隔的模型(仅对话能力,
|
|
// 兼容存量空串);单轮上限 validateBatchCap 控制调用量。
|
|
func (s *AiGatewayService) validateTargets(ctx context.Context, channelID uint) ([]model.AiModelCache, error) {
|
|
stale := time.Now().Add(-modelRecheckGap)
|
|
var rows []model.AiModelCache
|
|
err := s.db.WithContext(ctx).
|
|
Where("channel_id = ? AND capability IN ?", channelID, []string{"CHAT", ""}).
|
|
Where("checked_at IS NULL OR (unusable = ? AND checked_at < ?)", true, stale).
|
|
Order("id ASC").Limit(validateBatchCap).Find(&rows).Error
|
|
return rows, err
|
|
}
|
|
|
|
// validateChannelModels 逐个试调渠道内待验证模型并落结论,把不可按需调用的模型
|
|
// 从池中剔除、把恢复供给的解除标记;ListModels 无字段可事先判别,只能试调习得。
|
|
func (s *AiGatewayService) validateChannelModels(ctx context.Context, channelID uint) {
|
|
if !s.beginValidate(channelID) {
|
|
return
|
|
}
|
|
defer s.endValidate(channelID)
|
|
var ch model.AiChannel
|
|
if err := s.db.WithContext(ctx).First(&ch, channelID).Error; err != nil {
|
|
return
|
|
}
|
|
cred, err := s.configs.credentialsByID(ctx, ch.OciConfigID)
|
|
if err != nil {
|
|
return
|
|
}
|
|
rows, err := s.validateTargets(ctx, channelID)
|
|
if err != nil {
|
|
log.Printf("validate channel %d models: %v", channelID, err)
|
|
return
|
|
}
|
|
for _, r := range rows {
|
|
s.validateOne(ctx, cred, &ch, r)
|
|
}
|
|
}
|
|
|
|
// validateOne 试调单个模型并落库:可用(200/429)清除标记;模型级不可用标记剔除;
|
|
// 其余 4xx 记录已检不改状态;5xx/网络类瞬态不落 checked_at,下轮再验。
|
|
func (s *AiGatewayService) validateOne(ctx context.Context, cred oci.Credentials, ch *model.AiChannel, r model.AiModelCache) {
|
|
code, err := s.client.GenAiProbeChat(ctx, cred, ch.Region, r.ModelOcid, r.Name)
|
|
updates := map[string]any{"checked_at": time.Now()}
|
|
switch {
|
|
case code == 200 || code == 429:
|
|
updates["unusable"], updates["unusable_reason"] = false, ""
|
|
case oci.IsModelUnavailable(err):
|
|
updates["unusable"], updates["unusable_reason"] = true, shortReason(oci.CompactError(err))
|
|
case code == 0 || code >= 500:
|
|
return
|
|
}
|
|
s.db.WithContext(ctx).Model(&model.AiModelCache{}).Where("id = ?", r.ID).Updates(updates)
|
|
}
|
|
|
|
// validateModelsAsync 后台验证渠道模型池(手动同步后触发,不阻塞接口响应)。
|
|
func (s *AiGatewayService) validateModelsAsync(ctx context.Context, channelID uint) {
|
|
s.wg.Add(1)
|
|
go func() {
|
|
defer s.wg.Done()
|
|
vctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 3*time.Minute)
|
|
defer cancel()
|
|
s.validateChannelModels(vctx, channelID)
|
|
}()
|
|
}
|
|
|
|
// probeCandidateCap 是单次探测的候选上限:模型级不可用会被标记剔除不再重试,
|
|
// 放宽到 8 让一次探测有机会越过整批坏模型找到可用者;其他错误另有 3 次止损预算。
|
|
const probeCandidateCap = 8
|
|
|
|
// probeCandidates 只取对话模型,按可靠度排序后跨厂商取候选:主流文本模型优先;
|
|
// voice 等负分形态(元数据标 CHAT 但实际不可对话)直接排除,不浪费试调预算;
|
|
// 每厂商先取最高分再按分数补位——部分区域某厂商全为微调基座(调用必失败),
|
|
// 不能让单一厂商占满候选名额拖垮整个渠道的探测结论。
|
|
func probeCandidates(models []oci.GenAiModel) []oci.GenAiModel {
|
|
var sorted []oci.GenAiModel
|
|
for _, m := range models {
|
|
if (m.Capability == "" || m.Capability == "CHAT") && probeScore(m.Name) >= 0 {
|
|
sorted = append(sorted, m)
|
|
}
|
|
}
|
|
sort.SliceStable(sorted, func(i, j int) bool {
|
|
return probeScore(sorted[i].Name) > probeScore(sorted[j].Name)
|
|
})
|
|
return diversifyByVendor(sorted, probeCandidateCap)
|
|
}
|
|
|
|
// diversifyByVendor 从已排序列表先每厂商各取一个,不足 limit 再按原序补位。
|
|
func diversifyByVendor(sorted []oci.GenAiModel, limit int) []oci.GenAiModel {
|
|
picked := make([]oci.GenAiModel, 0, limit)
|
|
used := make(map[int]bool)
|
|
seenVendor := make(map[string]bool)
|
|
for i, m := range sorted {
|
|
if len(picked) >= limit {
|
|
break
|
|
}
|
|
if v := modelVendor(m); !seenVendor[v] {
|
|
seenVendor[v] = true
|
|
used[i] = true
|
|
picked = append(picked, m)
|
|
}
|
|
}
|
|
for i, m := range sorted {
|
|
if len(picked) >= limit {
|
|
break
|
|
}
|
|
if !used[i] {
|
|
picked = append(picked, m)
|
|
}
|
|
}
|
|
return picked
|
|
}
|
|
|
|
// modelVendor 取厂商标识;OCI 未回填 vendor 时退化为模型名「.」前缀。
|
|
func modelVendor(m oci.GenAiModel) string {
|
|
if m.Vendor != "" {
|
|
return strings.ToLower(m.Vendor)
|
|
}
|
|
name := strings.ToLower(m.Name)
|
|
if i := strings.Index(name, "."); i > 0 {
|
|
return name[:i]
|
|
}
|
|
return name
|
|
}
|
|
|
|
func probeScore(name string) int {
|
|
n := strings.ToLower(name)
|
|
switch {
|
|
case strings.Contains(n, "voice") || strings.Contains(n, "embed") || strings.Contains(n, "rerank"):
|
|
return -10
|
|
case strings.Contains(n, "llama") && !strings.Contains(n, "vision"):
|
|
return 5
|
|
case strings.Contains(n, "gemini") || strings.Contains(n, "gpt-oss"):
|
|
return 4
|
|
case strings.Contains(n, "command"):
|
|
return 3
|
|
case strings.Contains(n, "grok") && !strings.Contains(n, "multi-agent"):
|
|
return 2
|
|
default:
|
|
return 0
|
|
}
|
|
}
|
|
|
|
// classifyProbeErr 区分「区域无服务端点」与其他错误。
|
|
func classifyProbeErr(err error) string {
|
|
msg := strings.ToLower(err.Error())
|
|
if strings.Contains(msg, "no such host") || strings.Contains(msg, "timeout") ||
|
|
strings.Contains(msg, "connection refused") || strings.Contains(msg, "dial tcp") {
|
|
return "no_service"
|
|
}
|
|
return "error"
|
|
}
|
|
|
|
func truncateErr(msg string) string {
|
|
if len(msg) > 500 {
|
|
return msg[:500]
|
|
}
|
|
return msg
|
|
}
|
|
|
|
// SyncModels 重新拉取渠道区域的模型列表并覆盖缓存(标记按 OCID 结转),
|
|
// 随后触发后台验证:新模型逐个试调,不可按需调用的数十秒内从池中剔除。
|
|
func (s *AiGatewayService) SyncModels(ctx context.Context, id uint) ([]model.AiModelCache, error) {
|
|
var ch model.AiChannel
|
|
if err := s.db.WithContext(ctx).First(&ch, id).Error; err != nil {
|
|
return nil, fmt.Errorf("渠道不存在")
|
|
}
|
|
cred, err := s.configs.credentialsByID(ctx, ch.OciConfigID)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
models, err := s.client.ListGenAiModels(ctx, cred, ch.Region)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("同步模型失败:%s", oci.CompactError(err))
|
|
}
|
|
if err := s.replaceModels(ctx, id, models); err != nil {
|
|
return nil, err
|
|
}
|
|
s.validateModelsAsync(ctx, id)
|
|
return s.channelModels(ctx, id)
|
|
}
|
|
|
|
// replaceModels 以事务整组覆盖渠道模型缓存,按 OCID 结转不可用标记与验证时间
|
|
// (OCID 变化视为新条目,自然回到待验证状态)。
|
|
func (s *AiGatewayService) replaceModels(ctx context.Context, channelID uint, models []oci.GenAiModel) error {
|
|
marks, err := s.loadModelMarks(ctx, channelID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
rows := make([]model.AiModelCache, 0, len(models))
|
|
now := time.Now()
|
|
for _, m := range models {
|
|
row := model.AiModelCache{ChannelID: channelID, ModelOcid: m.Ocid, Name: m.Name, Vendor: m.Vendor,
|
|
Capability: m.Capability, SyncedAt: now, DeprecatedAt: m.Deprecated, RetiredAt: m.Retired}
|
|
if prev, ok := marks[m.Ocid]; ok {
|
|
row.Unusable, row.UnusableReason, row.CheckedAt = prev.Unusable, prev.UnusableReason, prev.CheckedAt
|
|
}
|
|
rows = append(rows, row)
|
|
}
|
|
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
|
if err := tx.Where("channel_id = ?", channelID).Delete(&model.AiModelCache{}).Error; err != nil {
|
|
return err
|
|
}
|
|
if len(rows) == 0 {
|
|
return nil
|
|
}
|
|
return tx.Create(&rows).Error
|
|
})
|
|
}
|
|
|
|
// loadModelMarks 取渠道内带标记或已验证过的行(OCID → 行),供同步结转与候选过滤。
|
|
func (s *AiGatewayService) loadModelMarks(ctx context.Context, channelID uint) (map[string]model.AiModelCache, error) {
|
|
var rows []model.AiModelCache
|
|
err := s.db.WithContext(ctx).Select("model_ocid", "unusable", "unusable_reason", "checked_at").
|
|
Where("channel_id = ? AND (unusable = ? OR checked_at IS NOT NULL)", channelID, true).Find(&rows).Error
|
|
if err != nil {
|
|
return nil, fmt.Errorf("load model marks: %w", err)
|
|
}
|
|
marks := make(map[string]model.AiModelCache, len(rows))
|
|
for _, r := range rows {
|
|
marks[r.ModelOcid] = r
|
|
}
|
|
return marks, nil
|
|
}
|
|
|
|
func (s *AiGatewayService) channelModels(ctx context.Context, channelID uint) ([]model.AiModelCache, error) {
|
|
var rows []model.AiModelCache
|
|
err := s.db.WithContext(ctx).Where("channel_id = ?", channelID).Order("name ASC").Find(&rows).Error
|
|
return rows, err
|
|
}
|
|
|
|
// GatewayModels 聚合启用渠道的可用模型(按名称去重,剔除不可按需调用标记),
|
|
// 供 /ai/v1/models;group 非空时仅聚合该分组渠道(与密钥分组路由口径一致)。
|
|
func (s *AiGatewayService) GatewayModels(ctx context.Context, group string) (aiwire.ModelList, error) {
|
|
q := s.db.WithContext(ctx).
|
|
Joins("JOIN ai_channels ON ai_channels.id = ai_model_caches.channel_id AND ai_channels.enabled = ?", true).
|
|
Where("(ai_model_caches.unusable = ? OR ai_model_caches.unusable IS NULL)", false)
|
|
if group != "" {
|
|
q = q.Where("ai_channels.channel_group = ?", group)
|
|
}
|
|
var rows []model.AiModelCache
|
|
err := q.Order("ai_model_caches.name ASC").Find(&rows).Error
|
|
list := aiwire.ModelList{Object: "list", Data: []aiwire.Model{}}
|
|
if err != nil {
|
|
return list, err
|
|
}
|
|
seen := map[string]bool{}
|
|
for _, r := range rows {
|
|
if seen[r.Name] {
|
|
continue
|
|
}
|
|
seen[r.Name] = true
|
|
list.Data = append(list.Data, aiwire.Model{ID: r.Name, Object: "model", Created: r.SyncedAt.Unix(), OwnedBy: r.Vendor})
|
|
}
|
|
return list, nil
|
|
}
|
|
|
|
// DeprecatingModels 返回 within 窗口内即将退役或即将弃用的在池模型(按名称去重):
|
|
// 退役(TimeOnDemandRetired)才导致不可调用,单独标注;已过弃用日但未到退役日的
|
|
// 模型仍可正常调用,不再反复告警;已过退役日的在同步层剔除,不会出现在池中。
|
|
func (s *AiGatewayService) DeprecatingModels(ctx context.Context, within time.Duration) ([]string, error) {
|
|
now := time.Now()
|
|
deadline := now.Add(within)
|
|
var rows []model.AiModelCache
|
|
err := s.db.WithContext(ctx).
|
|
Where("(retired_at IS NOT NULL AND retired_at > ? AND retired_at <= ?) OR (deprecated_at IS NOT NULL AND deprecated_at >= ? AND deprecated_at <= ?)",
|
|
now, deadline, now, deadline).
|
|
Where("(unusable = ? OR unusable IS NULL)", false).
|
|
Order("name ASC").Find(&rows).Error
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
seen := map[string]bool{}
|
|
var out []string
|
|
for _, r := range rows {
|
|
if seen[r.Name] {
|
|
continue
|
|
}
|
|
seen[r.Name] = true
|
|
if r.RetiredAt != nil && r.RetiredAt.After(now) && !r.RetiredAt.After(deadline) {
|
|
out = append(out, fmt.Sprintf("%s(%s 退役,届时无法调用)", r.Name, r.RetiredAt.Format("2006-01-02")))
|
|
continue
|
|
}
|
|
out = append(out, fmt.Sprintf("%s(%s 宣布弃用,退役前仍可调用)", r.Name, r.DeprecatedAt.Format("2006-01-02")))
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
// ProbeAll 逐个探测全部渠道并顺带验证模型池,返回状态汇总;供 AI 探测后台任务调用。
|
|
func (s *AiGatewayService) ProbeAll(ctx context.Context) (string, error) {
|
|
chs, err := s.Channels(ctx)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
counts := map[string]int{}
|
|
var failures []string
|
|
for _, ch := range chs {
|
|
fresh, err := s.ProbeChannel(ctx, ch.ID)
|
|
if err != nil {
|
|
counts["error"]++
|
|
failures = append(failures, fmt.Sprintf("#%d %s", ch.ID, truncateErr(err.Error())))
|
|
continue
|
|
}
|
|
counts[fresh.ProbeStatus]++
|
|
// 后台任务里同步执行:新模型验证入池、坏模型定期复检,零人工收敛
|
|
s.validateChannelModels(ctx, ch.ID)
|
|
}
|
|
msg := fmt.Sprintf("probed %d: %d ok, %d no_service, %d no_quota, %d error",
|
|
len(chs), counts["ok"], counts["no_service"], counts["no_quota"], counts["error"])
|
|
if len(failures) > 0 {
|
|
msg += "; " + strings.Join(failures, "; ")
|
|
}
|
|
return msg, nil
|
|
}
|
|
|
|
// ---- 调用日志 ----
|
|
|
|
// LogCall 落一条调用日志(仅元数据与用量,永不含请求 / 响应正文),返回落库 ID 供内容日志关联(失败为 0)。
|
|
func (s *AiGatewayService) LogCall(entry model.AiCallLog) uint {
|
|
entry.ErrMsg = truncateErr(entry.ErrMsg)
|
|
stored := false
|
|
err := s.db.Transaction(func(tx *gorm.DB) error {
|
|
ok, err := lockAiLogParent(tx, &model.AiChannel{}, entry.ChannelID)
|
|
if err != nil || !ok {
|
|
return err
|
|
}
|
|
stored = true
|
|
return tx.Create(&entry).Error
|
|
})
|
|
if err != nil {
|
|
log.Printf("ai call log: %v", err)
|
|
return 0
|
|
}
|
|
if !stored {
|
|
return 0
|
|
}
|
|
return entry.ID
|
|
}
|
|
|
|
func lockAiLogParent(tx *gorm.DB, value any, id uint) (bool, error) {
|
|
if id == 0 {
|
|
return true, nil
|
|
}
|
|
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(value, id).Error
|
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return false, nil
|
|
}
|
|
return err == nil, err
|
|
}
|
|
|
|
// CallLogs 分页查询调用日志。
|
|
func (s *AiGatewayService) CallLogs(ctx context.Context, page, size int) ([]model.AiCallLog, int64, error) {
|
|
if page < 1 {
|
|
page = 1
|
|
}
|
|
if size < 1 || size > 200 {
|
|
size = 50
|
|
}
|
|
var total int64
|
|
q := s.db.WithContext(ctx).Model(&model.AiCallLog{})
|
|
if err := q.Count(&total).Error; err != nil {
|
|
return nil, 0, err
|
|
}
|
|
var rows []model.AiCallLog
|
|
err := q.Order("id DESC").Offset((page - 1) * size).Limit(size).Find(&rows).Error
|
|
return rows, total, err
|
|
}
|
|
|
|
// ---- 内容日志(红线例外,按密钥显式限时开启) ----
|
|
|
|
// UpdateKeyContentLog 设置密钥内容日志窗口:hours=0 立即关闭,>0 从现在起开启 N 小时(上限 7 天)。
|
|
func (s *AiGatewayService) UpdateKeyContentLog(ctx context.Context, id uint, hours int) (*model.AiKey, error) {
|
|
if hours < 0 || hours > aiContentLogMaxHours {
|
|
return nil, fmt.Errorf("内容日志时长需在 0-%d 小时之间", aiContentLogMaxHours)
|
|
}
|
|
updates := map[string]any{"content_log_until": gorm.Expr("NULL")}
|
|
if hours > 0 {
|
|
updates["content_log_until"] = time.Now().Add(time.Duration(hours) * time.Hour)
|
|
}
|
|
if err := s.db.WithContext(ctx).Model(&model.AiKey{}).Where("id = ?", id).Updates(updates).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
// 重读用新变量:gorm 扫描 NULL 列到已有值的结构体时会保留旧值
|
|
var fresh model.AiKey
|
|
if err := s.db.WithContext(ctx).First(&fresh, id).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
return &fresh, nil
|
|
}
|
|
|
|
// LogContent 写一条内容日志(调用方已确认密钥开启且未过期);正文截断至 64KB。
|
|
func (s *AiGatewayService) LogContent(entry model.AiContentLog) {
|
|
if entry.CallLogID == 0 {
|
|
return
|
|
}
|
|
entry.RequestBody = truncateBody(entry.RequestBody)
|
|
entry.ResponseBody = truncateBody(entry.ResponseBody)
|
|
err := s.db.Transaction(func(tx *gorm.DB) error {
|
|
ok, err := lockAiLogParent(tx, &model.AiCallLog{}, entry.CallLogID)
|
|
if err != nil || !ok {
|
|
return err
|
|
}
|
|
return tx.Create(&entry).Error
|
|
})
|
|
if err != nil {
|
|
log.Printf("ai content log: %v", err)
|
|
}
|
|
}
|
|
|
|
func truncateBody(s string) string {
|
|
if len(s) > aiContentBodyLimit {
|
|
return s[:aiContentBodyLimit]
|
|
}
|
|
return s
|
|
}
|
|
|
|
// ContentLogs 分页查询内容日志(keyID / callLogID 为 0 时不过滤)。
|
|
func (s *AiGatewayService) ContentLogs(ctx context.Context, keyID, callLogID uint, page, size int) ([]model.AiContentLog, int64, error) {
|
|
if page < 1 {
|
|
page = 1
|
|
}
|
|
if size < 1 || size > 100 {
|
|
size = 20
|
|
}
|
|
q := s.db.WithContext(ctx).Model(&model.AiContentLog{})
|
|
if keyID > 0 {
|
|
q = q.Where("key_id = ?", keyID)
|
|
}
|
|
if callLogID > 0 {
|
|
q = q.Where("call_log_id = ?", callLogID)
|
|
}
|
|
var total int64
|
|
if err := q.Count(&total).Error; err != nil {
|
|
return nil, 0, err
|
|
}
|
|
var rows []model.AiContentLog
|
|
err := q.Order("id DESC").Offset((page - 1) * size).Limit(size).Find(&rows).Error
|
|
return rows, total, err
|
|
}
|
|
|
|
// StartCleanup 启动调用日志周期清理:启动即清一次,之后每 24h 一次。
|
|
func (s *AiGatewayService) StartCleanup(ctx context.Context) {
|
|
s.wg.Add(1)
|
|
go func() {
|
|
defer s.wg.Done()
|
|
s.cleanupOnce(ctx)
|
|
ticker := time.NewTicker(aiLogCleanupTick)
|
|
defer ticker.Stop()
|
|
for {
|
|
select {
|
|
case <-ctx.Done():
|
|
return
|
|
case <-ticker.C:
|
|
s.cleanupOnce(ctx)
|
|
}
|
|
}
|
|
}()
|
|
}
|
|
|
|
func (s *AiGatewayService) cleanupOnce(ctx context.Context) {
|
|
s.cleanupTable(ctx, &model.AiCallLog{}, aiLogRetention, aiLogMaxRows, "ai log")
|
|
s.cleanupTable(ctx, &model.AiContentLog{}, aiContentLogRetention, aiContentLogMaxRows, "ai content log")
|
|
}
|
|
|
|
// cleanupTable 按保留期与行数上限清理日志表(超限删最旧)。
|
|
func (s *AiGatewayService) cleanupTable(ctx context.Context, m any, retention time.Duration, maxRows int, tag string) {
|
|
cutoff := time.Now().Add(-retention)
|
|
if err := s.db.WithContext(ctx).Where("created_at < ?", cutoff).Delete(m).Error; err != nil {
|
|
log.Printf("%s cleanup: %v", tag, err)
|
|
return
|
|
}
|
|
var total int64
|
|
if err := s.db.WithContext(ctx).Model(m).Count(&total).Error; err != nil {
|
|
return
|
|
}
|
|
if overflow := int(total) - maxRows; overflow > 0 {
|
|
var ids []uint
|
|
s.db.WithContext(ctx).Model(m).Order("id ASC").Limit(overflow).Pluck("id", &ids)
|
|
if len(ids) > 0 {
|
|
s.db.WithContext(ctx).Delete(m, ids)
|
|
}
|
|
}
|
|
}
|
|
|
|
// Wait 等待后台清理 goroutine 退出。
|
|
func (s *AiGatewayService) Wait() { s.wg.Wait() }
|