package service import ( "context" "crypto/rand" "crypto/sha256" "encoding/hex" "encoding/json" "errors" "fmt" "log" "sort" "strconv" "strings" "sync" "sync/atomic" "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 // 内容日志(红线例外)约束:开启必须限时(上限 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 // onChannelsChanged 在渠道增删后触发,由 main 装配为探测任务同步钩子 onChannelsChanged func(context.Context) // filterDeprecated 是「过滤弃用模型」开关(内存镜像,持久化在 settings 表) filterDeprecated atomic.Bool // streamGuard* 是 Responses 流式保险丝(instructions+tools 合计超阈值改非流式) streamGuardEnabled atomic.Bool streamGuardKB atomic.Int64 // grokWebSearch / grokXSearch 是 xai. 模型服务端搜索工具默认注入开关 grokWebSearch atomic.Bool grokXSearch atomic.Bool // upstreamWaitSec 是 responses 直通的上游无响应预算(秒):非流式为单次尝试 // 总超时,流式为等待响应头预算;multi-agent/搜索类模型远超 SDK 默认 60s upstreamWaitSec atomic.Int64 } // NewAiGatewayService 组装依赖;调用 StartCleanup 后开始调用日志周期清理。 func NewAiGatewayService(db *gorm.DB, configs *OciConfigService, client oci.Client) *AiGatewayService { s := &AiGatewayService{db: db, configs: configs, client: client, lastTouch: map[uint]time.Time{}} s.filterDeprecated.Store(loadBoolSetting(db, settingAiFilterDeprecated, false)) s.streamGuardEnabled.Store(loadBoolSetting(db, settingAiStreamGuardEnabled, true)) s.streamGuardKB.Store(int64(loadIntSetting(db, settingAiStreamGuardKB, defaultStreamGuardKB))) s.grokWebSearch.Store(loadBoolSetting(db, settingAiGrokWebSearch, true)) s.grokXSearch.Store(loadBoolSetting(db, settingAiGrokXSearch, true)) s.upstreamWaitSec.Store(int64(loadIntSetting(db, settingAiUpstreamWaitSec, defaultUpstreamWaitSec))) return s } // AI 网关运行时设置的配置键;bool 值存 "1"/"0"。 const ( // settingAiFilterDeprecated 是「过滤弃用模型」开关,缺省关闭。 settingAiFilterDeprecated = "ai_filter_deprecated" // settingAiStreamGuardEnabled / settingAiStreamGuardKB 是流式保险丝开关与 // 阈值(KB),缺省开、60(上游对 instructions+tools >≈64.5KB 流式静默断流)。 settingAiStreamGuardEnabled = "ai_stream_guard_enabled" settingAiStreamGuardKB = "ai_stream_guard_kb" // settingAiGrokWebSearch / settingAiGrokXSearch 是 grok 搜索工具默认注入 // 开关,缺省开。 settingAiGrokWebSearch = "ai_grok_web_search" settingAiGrokXSearch = "ai_grok_x_search" // settingAiUpstreamWaitSec 是 responses 直通的上游无响应预算(秒)。 settingAiUpstreamWaitSec = "ai_upstream_wait_seconds" ) // defaultStreamGuardKB 是保险丝阈值缺省值,低于实测断流边界留余量。 const defaultStreamGuardKB = 60 // defaultUpstreamWaitSec 是上游无响应预算缺省值(秒):multi-agent 非流式 // 实测 100~180s 才回响应头,给足余量;上下限见 SetUpstreamWait。 const defaultUpstreamWaitSec = 300 // loadBoolSetting 读 settings 表布尔键,无行或值非法时返回缺省。 func loadBoolSetting(db *gorm.DB, key string, def bool) bool { var row model.Setting if err := db.Where("key = ?", key).First(&row).Error; err != nil { return def } return row.Value == "1" } // loadIntSetting 读 settings 表整数键,无行或解析失败时返回缺省。 func loadIntSetting(db *gorm.DB, key string, def int) int { var row model.Setting if err := db.Where("key = ?", key).First(&row).Error; err != nil { return def } n, err := strconv.Atoi(row.Value) if err != nil { return def } return n } // saveBoolSetting 持久化布尔键。 func (s *AiGatewayService) saveBoolSetting(ctx context.Context, key string, on bool) error { value := "0" if on { value = "1" } return s.db.WithContext(ctx).Save(&model.Setting{Key: key, Value: value}).Error } // FilterDeprecated 返回「过滤弃用模型」开关状态。 func (s *AiGatewayService) FilterDeprecated() bool { return s.filterDeprecated.Load() } // SetFilterDeprecated 持久化并即时生效开关:开启后已宣布弃用 // (deprecated_at 非空,即使未退役)的模型从列表与路由中排除。 func (s *AiGatewayService) SetFilterDeprecated(ctx context.Context, on bool) error { if err := s.saveBoolSetting(ctx, settingAiFilterDeprecated, on); err != nil { return fmt.Errorf("保存过滤弃用模型开关: %w", err) } s.filterDeprecated.Store(on) return nil } // StreamGuard 返回流式保险丝开关与阈值(KB)。 func (s *AiGatewayService) StreamGuard() (bool, int) { return s.streamGuardEnabled.Load(), int(s.streamGuardKB.Load()) } // SetStreamGuard 持久化并即时生效流式保险丝;kb 限定 1..1024。 func (s *AiGatewayService) SetStreamGuard(ctx context.Context, on bool, kb int) error { if kb < 1 || kb > 1024 { return fmt.Errorf("流式保险丝阈值须在 1..1024 KB, 收到 %d", kb) } if err := s.saveBoolSetting(ctx, settingAiStreamGuardEnabled, on); err != nil { return fmt.Errorf("保存流式保险丝开关: %w", err) } err := s.db.WithContext(ctx). Save(&model.Setting{Key: settingAiStreamGuardKB, Value: strconv.Itoa(kb)}).Error if err != nil { return fmt.Errorf("保存流式保险丝阈值: %w", err) } s.streamGuardEnabled.Store(on) s.streamGuardKB.Store(int64(kb)) return nil } // UpstreamWait 返回 responses 直通的上游无响应预算。 func (s *AiGatewayService) UpstreamWait() time.Duration { return time.Duration(s.upstreamWaitSec.Load()) * time.Second } // SetUpstreamWait 持久化并即时生效上游无响应预算;sec 限定 30..900。 func (s *AiGatewayService) SetUpstreamWait(ctx context.Context, sec int) error { if sec < 30 || sec > 900 { return fmt.Errorf("上游无响应预算须在 30..900 秒, 收到 %d", sec) } err := s.db.WithContext(ctx). Save(&model.Setting{Key: settingAiUpstreamWaitSec, Value: strconv.Itoa(sec)}).Error if err != nil { return fmt.Errorf("保存上游无响应预算: %w", err) } s.upstreamWaitSec.Store(int64(sec)) return nil } // GrokSearch 返回 grok 服务端搜索工具默认注入开关(web_search, x_search)。 func (s *AiGatewayService) GrokSearch() (bool, bool) { return s.grokWebSearch.Load(), s.grokXSearch.Load() } // SetGrokSearch 持久化并即时生效 grok 搜索工具默认注入开关。 func (s *AiGatewayService) SetGrokSearch(ctx context.Context, web, x bool) error { if err := s.saveBoolSetting(ctx, settingAiGrokWebSearch, web); err != nil { return fmt.Errorf("保存 grok web_search 开关: %w", err) } if err := s.saveBoolSetting(ctx, settingAiGrokXSearch, x); err != nil { return fmt.Errorf("保存 grok x_search 开关: %w", err) } s.grokWebSearch.Store(web) s.grokXSearch.Store(x) return nil } // 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, models []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), Models: normalizeKeyModels(models), Enabled: true} if err := s.db.WithContext(ctx).Create(key).Error; err != nil { return "", nil, fmt.Errorf("密钥名称或取值与现有密钥重复") } return raw, key, nil } // normalizeKeyModels 规范化模型白名单:trim、剔空串、保序去重;结果为空返回 nil(= 不限)。 func normalizeKeyModels(models []string) []string { var out []string seen := map[string]bool{} for _, m := range models { m = strings.TrimSpace(m) if m == "" || seen[m] { continue } seen[m] = true out = append(out, m) } return out } 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 修改密钥名称 / 启用状态 / 分组 / 模型白名单(指针非空即覆盖,可置空)。 func (s *AiGatewayService) UpdateKey(ctx context.Context, id uint, name string, enabled *bool, group *string, models *[]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 models != nil { // 手动序列化走 map 更新,不依赖 GORM map 路径对 serializer 的支持 b, err := json.Marshal(normalizeKeyModels(*models)) if err != nil { return fmt.Errorf("serialize models: %w", err) } updates["models"] = string(b) } 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 if err := s.db.WithContext(ctx).Order("priority ASC, id ASC").Find(&chs).Error; err != nil { return nil, err } var rows []struct { ChannelID uint N int64 } q := s.db.WithContext(ctx).Model(&model.AiModelCache{}). Select("channel_id, COUNT(*) AS n"). Where("name NOT IN (SELECT name FROM ai_model_blacklists)") if s.FilterDeprecated() { q = q.Where("deprecated_at IS NULL") } err := q.Group("channel_id").Scan(&rows).Error if err != nil { return nil, err } counts := make(map[uint]int64, len(rows)) for _, r := range rows { counts[r.ChannelID] = r.N } for i := range chs { chs[i].ModelCount = counts[chs[i].ID] } return chs, nil } // 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 探测渠道可用性:服务可见性 → 模型同步 → 极小 max_tokens 配额试调。 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)) } models = supportedGatewayModels(models) models, err = s.withoutBlacklisted(ctx, models) if err != nil { return "error", truncateErr(err.Error()) } 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()) } return s.probeChat(ctx, cred, ch, models) } // probeChat 按候选顺序试调(上限 8,已验证的探测模型置首位):遇「模型不可按需 // 调用」(微调基座 400 / 实体不存在 404)换下一个候选,错误信息带模型名供用户加入 // 黑名单;401/403 与鉴权类 404 可能只是模型级无权限,记录后继续换候选,全部候选 // 失败且出现过鉴权拒绝才定论 no_quota;其余错误累计 3 次止损。 func (s *AiGatewayService) probeChat(ctx context.Context, cred oci.Credentials, ch *model.AiChannel, models []oci.GenAiModel) (string, string) { status, detail := "error", "无可试调对话模型" quotaDetail := "" errBudget := 3 for _, m := range probeCandidates(ch.ProbeModel, 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): status, detail = "error", truncateErr(fmt.Sprintf("%s: 不可按需调用,建议加入模型黑名单", m.Name)) case code == 401 || code == 403 || code == 404: quotaDetail = truncateErr(fmt.Sprintf("%s: %s", m.Name, oci.CompactError(err))) default: status, detail = "error", truncateErr(fmt.Sprintf("%s: %s", m.Name, oci.CompactError(err))) if errBudget--; errBudget == 0 { return status, detail } } } if quotaDetail != "" { return "no_quota", quotaDetail } return status, detail } // probeCandidateCap 是单次探测的候选上限:放宽到 8 让一次探测有机会越过整批 // 不可按需调用的坏模型找到可用者;其他错误另有 3 次止损预算。 const probeCandidateCap = 8 // probeCandidates 只取对话模型,按可靠度排序后跨厂商取候选:用户已验证的 // probeModel 固定放首位(不做能力过滤,测试通过即有效),其余主流文本模型优先; // voice 等负分形态(元数据标 CHAT 但实际不可对话)直接排除,不浪费试调预算; // 每厂商先取最高分再按分数补位——部分区域某厂商全为微调基座(调用必失败), // 不能让单一厂商占满候选名额拖垮整个渠道的探测结论。 func probeCandidates(probeModel string, models []oci.GenAiModel) []oci.GenAiModel { var pinned, sorted []oci.GenAiModel for _, m := range models { if probeModel != "" && m.Name == probeModel { pinned = append(pinned, m) continue } 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 append(pinned, 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 重新拉取渠道区域的模型列表并覆盖缓存,黑名单中的模型不入库。 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)) } models = supportedGatewayModels(models) if models, err = s.withoutBlacklisted(ctx, models); err != nil { return nil, err } if err := s.replaceModels(ctx, id, models); err != nil { return nil, err } return s.channelModels(ctx, id) } // replaceModels 以事务整组覆盖渠道模型缓存。 func (s *AiGatewayService) replaceModels(ctx context.Context, channelID uint, models []oci.GenAiModel) error { rows := make([]model.AiModelCache, 0, len(models)) now := time.Now() for _, m := range models { rows = append(rows, model.AiModelCache{ChannelID: channelID, ModelOcid: m.Ocid, Name: m.Name, Vendor: m.Vendor, Capability: m.Capability, SyncedAt: now, DeprecatedAt: m.Deprecated, RetiredAt: m.Retired}) } 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 }) } // channelModels 列出渠道模型缓存;黑名单模型查询层兜底过滤 // (拉黑即删缓存,正常不会残留,防御旧数据 / 并发窗口); // 「过滤弃用模型」开关开启时同样剔除已宣布弃用者(数据保留,展示口径过滤)。 func (s *AiGatewayService) channelModels(ctx context.Context, channelID uint) ([]model.AiModelCache, error) { q := s.db.WithContext(ctx).Where("channel_id = ?", channelID). Where("name NOT IN (SELECT name FROM ai_model_blacklists)") if s.FilterDeprecated() { q = q.Where("deprecated_at IS NULL") } var rows []model.AiModelCache err := q.Order("name ASC").Find(&rows).Error return rows, err } // ChannelModels 列出渠道的模型缓存(名称排序),渠道不存在时报错。 func (s *AiGatewayService) ChannelModels(ctx context.Context, id uint) ([]model.AiModelCache, error) { var n int64 if err := s.db.WithContext(ctx).Model(&model.AiChannel{}).Where("id = ?", id).Count(&n).Error; err != nil { return nil, err } if n == 0 { return nil, fmt.Errorf("渠道不存在") } return s.channelModels(ctx, id) } // TestChannelModel 对渠道缓存中的指定模型发极小试调;通过时把该模型设 // 为渠道探测验证模型(此后探测置于候选首位),渠道探测状态不为 ok 时顺带置 ok 并 // 复位熔断;未通过仅返回错误,不改动渠道状态。 func (s *AiGatewayService) TestChannelModel(ctx context.Context, id uint, name string) (*model.AiChannel, error) { var ch model.AiChannel if err := s.db.WithContext(ctx).First(&ch, id).Error; err != nil { return nil, fmt.Errorf("渠道不存在") } var mc model.AiModelCache err := s.db.WithContext(ctx).Where("channel_id = ? AND name = ?", id, name).First(&mc).Error if err != nil { return nil, fmt.Errorf("模型不在该渠道缓存中,请先同步模型") } cred, err := s.configs.credentialsByID(ctx, ch.OciConfigID) if err != nil { return nil, err } code, err := s.client.GenAiProbeChat(ctx, cred, ch.Region, mc.ModelOcid, mc.Name) if code != 200 && code != 429 { msg := fmt.Sprintf("HTTP %d", code) if err != nil { msg = oci.CompactError(err) } return nil, fmt.Errorf("测试未通过:%s", truncateErr(msg)) } return s.adoptProbeModel(ctx, &ch, name) } // adoptProbeModel 记录探测验证模型并返回更新后的渠道; // 状态不为 ok 时一并置 ok 并复位熔断。 func (s *AiGatewayService) adoptProbeModel(ctx context.Context, ch *model.AiChannel, name string) (*model.AiChannel, error) { updates := map[string]any{"probe_model": name} if ch.ProbeStatus != "ok" { updates["probe_status"] = "ok" updates["probe_error"] = "" updates["last_probe_at"] = time.Now() updates["fail_count"] = 0 updates["disabled_until"] = gorm.Expr("NULL") } if err := s.db.WithContext(ctx).Model(&model.AiChannel{}).Where("id = ?", ch.ID).Updates(updates).Error; err != nil { return nil, err } // 重读用新变量:gorm 扫描 NULL 列到已有值的结构体时会保留旧值 var fresh model.AiChannel err := s.db.WithContext(ctx).First(&fresh, ch.ID).Error return &fresh, 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) if group != "" { q = q.Where("ai_channels.channel_group = ?", group) } if s.FilterDeprecated() { q = q.Where("ai_model_caches.deprecated_at IS NULL") } 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 } // AggregatedModel 是聚合模型目录条目(设置页黑名单添加弹窗用)。 type AggregatedModel struct { Name string `json:"name"` Capability string `json:"capability"` } // AggregatedModels 返回启用渠道去重后的模型目录(含能力);空能力归一为 CHAT。 // 与模型列表口径一致:「过滤弃用」开启时弃用模型不出现在目录中。 func (s *AiGatewayService) AggregatedModels(ctx context.Context) ([]AggregatedModel, error) { q := s.db.WithContext(ctx). Joins("JOIN ai_channels ON ai_channels.id = ai_model_caches.channel_id AND ai_channels.enabled = ?", true) if s.FilterDeprecated() { q = q.Where("ai_model_caches.deprecated_at IS NULL") } var rows []model.AiModelCache if err := q.Order("ai_model_caches.name ASC").Find(&rows).Error; err != nil { return nil, fmt.Errorf("聚合模型目录: %w", err) } seen := map[string]bool{} out := []AggregatedModel{} for _, r := range rows { if seen[r.Name] { continue } seen[r.Name] = true cap := r.Capability if cap == "" { cap = "CHAT" } out = append(out, AggregatedModel{Name: r.Name, Capability: cap}) } return out, 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). 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]++ } 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 } // ---- 模型黑名单 ---- // Blacklist 列出全部黑名单模型(按名称排序)。 func (s *AiGatewayService) Blacklist(ctx context.Context) ([]model.AiModelBlacklist, error) { var rows []model.AiModelBlacklist err := s.db.WithContext(ctx).Order("name ASC").Find(&rows).Error return rows, err } // AddBlacklist 把模型名加入黑名单并删除全部渠道缓存中的同名条目; // 该模型此后同步 / 探测均被过滤,直到移出黑名单后重新同步。 func (s *AiGatewayService) AddBlacklist(ctx context.Context, name string) (*model.AiModelBlacklist, error) { name = strings.TrimSpace(name) if name == "" { return nil, fmt.Errorf("模型名不能为空") } row := model.AiModelBlacklist{Name: name} err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { var n int64 if err := tx.Model(&model.AiModelBlacklist{}).Where("name = ?", name).Count(&n).Error; err != nil { return err } if n > 0 { return fmt.Errorf("模型已在黑名单中") } if err := tx.Create(&row).Error; err != nil { return err } return tx.Where("name = ?", name).Delete(&model.AiModelCache{}).Error }) if err != nil { return nil, err } return &row, nil } // RemoveBlacklist 把模型移出黑名单;缓存不回填,下次同步 / 探测自然恢复入池。 func (s *AiGatewayService) RemoveBlacklist(ctx context.Context, id uint) error { res := s.db.WithContext(ctx).Delete(&model.AiModelBlacklist{}, id) if res.Error != nil { return res.Error } if res.RowsAffected == 0 { return fmt.Errorf("黑名单条目不存在") } return nil } // withoutBlacklisted 过滤掉黑名单中的模型,同步与探测入库前统一经此收口。 // supportedGatewayModels 过滤模型目录:对话模型仅保留实测支持 OpenAI 兼容面的 // 厂商(xai / meta / openai)——typed chat 面已剔除,google / cohere 对话模型无 // 上游通路,不入目录(不出现在列表、路由与探测候选);EMBEDDING 模型不受影响。 func supportedGatewayModels(models []oci.GenAiModel) []oci.GenAiModel { out := make([]oci.GenAiModel, 0, len(models)) for _, m := range models { if m.Capability == "CHAT" && !compatChatVendor(m.Name) { continue } out = append(out, m) } return out } func compatChatVendor(model string) bool { for _, prefix := range []string{"xai.", "meta.", "openai."} { if strings.HasPrefix(model, prefix) { return true } } return false } func (s *AiGatewayService) withoutBlacklisted(ctx context.Context, models []oci.GenAiModel) ([]oci.GenAiModel, error) { var names []string if err := s.db.WithContext(ctx).Model(&model.AiModelBlacklist{}).Pluck("name", &names).Error; err != nil { return nil, fmt.Errorf("读取模型黑名单: %w", err) } if len(names) == 0 { return models, nil } banned := make(map[string]bool, len(names)) for _, n := range names { banned[n] = true } out := make([]oci.GenAiModel, 0, len(models)) for _, m := range models { if !banned[m.Name] { out = append(out, m) } } return out, 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() }