AI 网关切换 OpenAI 兼容面并移除 chat 端点,新增模型黑白名单
CI / test (push) Successful in 31s
Release / release (push) Successful in 54s

This commit is contained in:
2026-07-12 17:48:28 +08:00
parent 7706f59549
commit 489cb49cb3
34 changed files with 2602 additions and 2901 deletions
-332
View File
@@ -1,332 +0,0 @@
package service
import (
"encoding/json"
"fmt"
"strings"
"oci-portal/internal/aiwire"
)
// ErrAiUnsupportedBlock 表示请求含网关无法承接的内容块(仅支持文本与图片)。
var ErrAiUnsupportedBlock = fmt.Errorf("暂不支持文本与图片以外的内容块")
// AnthropicToIR 把 Anthropic Messages 请求转为 IR(OpenAI 线格式)。
// tool_result 块拆为独立 tool 角色消息(一拆多),顶层 system 变首条 system 消息。
func AnthropicToIR(req aiwire.MessagesRequest) (aiwire.ChatRequest, error) {
ir := aiwire.ChatRequest{
Model: req.Model,
Temperature: req.Temperature,
TopP: req.TopP,
TopK: req.TopK,
Stop: req.StopSequences,
Stream: req.Stream,
Tools: anthTools(req.Tools),
ToolChoice: anthToolChoice(req.ToolChoice),
}
if req.MaxTokens > 0 {
mt := req.MaxTokens
ir.MaxTokens = &mt
}
if sys := req.SystemText(); sys != "" {
ir.Messages = append(ir.Messages, aiwire.ChatMessage{Role: "system", Content: aiwire.NewTextContent(sys)})
}
for _, m := range req.Messages {
msgs, err := anthMessageToIR(m)
if err != nil {
return ir, err
}
ir.Messages = append(ir.Messages, msgs...)
}
return ir, nil
}
// anthMessageToIR 拆解单条 Anthropic 消息;user 消息里的 tool_result 前置为独立 tool 消息,
// image 块转为 IR image_url 部件。
func anthMessageToIR(m aiwire.AnthMessage) ([]aiwire.ChatMessage, error) {
var out []aiwire.ChatMessage
var parts []aiwire.ContentPart
var toolCalls []aiwire.ToolCall
hasImage := false
for _, b := range m.Content.AllBlocks() {
switch b.Type {
case "text":
parts = append(parts, aiwire.ContentPart{Type: "text", Text: b.Text})
case "image":
p, err := anthImagePart(b.Source)
if err != nil {
return nil, err
}
parts, hasImage = append(parts, p), true
case "tool_use":
toolCalls = append(toolCalls, aiwire.ToolCall{ID: b.ID, Type: "function",
Function: aiwire.FunctionCall{Name: b.Name, Arguments: string(b.Input)}})
case "tool_result":
out = append(out, aiwire.ChatMessage{Role: "tool", ToolCallID: b.ToolUseID,
Content: aiwire.NewTextContent(b.ResultText())})
case "thinking", "redacted_thinking":
// 一期忽略 thinking 块
default:
return nil, ErrAiUnsupportedBlock
}
}
content := partsContent(parts, hasImage)
if hasImage || content.JoinText() != "" || len(toolCalls) > 0 {
out = append(out, aiwire.ChatMessage{Role: m.Role, Content: content, ToolCalls: toolCalls})
}
return out, nil
}
// anthImagePart 把 Anthropic image 块转为 IR image_url 部件(base64 → data URI)。
func anthImagePart(source json.RawMessage) (aiwire.ContentPart, error) {
var src struct {
Type string `json:"type"`
MediaType string `json:"media_type"`
Data string `json:"data"`
URL string `json:"url"`
}
if err := json.Unmarshal(source, &src); err != nil {
return aiwire.ContentPart{}, fmt.Errorf("image source 解析失败: %w", err)
}
switch src.Type {
case "base64":
if src.MediaType == "" || src.Data == "" {
return aiwire.ContentPart{}, fmt.Errorf("image source 缺少 media_type 或 data")
}
url := "data:" + src.MediaType + ";base64," + src.Data
return aiwire.ContentPart{Type: "image_url", ImageURL: &aiwire.ImageURL{URL: url}}, nil
case "url":
if src.URL == "" {
return aiwire.ContentPart{}, fmt.Errorf("image source 缺少 url")
}
return aiwire.ContentPart{Type: "image_url", ImageURL: &aiwire.ImageURL{URL: src.URL}}, nil
}
return aiwire.ContentPart{}, fmt.Errorf("不支持的 image source 类型 %q", src.Type)
}
// partsContent 无图片时退回字符串形态(与上游线格式习惯一致),含图片时保留块数组。
func partsContent(parts []aiwire.ContentPart, hasImage bool) aiwire.Content {
if !hasImage {
var sb strings.Builder
for _, p := range parts {
sb.WriteString(p.Text)
}
return aiwire.NewTextContent(sb.String())
}
return aiwire.NewPartsContent(parts)
}
func anthTools(tools []aiwire.AnthTool) []aiwire.Tool {
if len(tools) == 0 {
return nil
}
out := make([]aiwire.Tool, 0, len(tools))
for _, t := range tools {
out = append(out, aiwire.Tool{Type: "function", Function: aiwire.FunctionDef{
Name: t.Name, Description: t.Description, Parameters: t.InputSchema}})
}
return out
}
// anthToolChoice 映射 {type:auto|any|tool,name} → OpenAI 形态。
func anthToolChoice(raw json.RawMessage) json.RawMessage {
if len(raw) == 0 {
return nil
}
var tc struct {
Type string `json:"type"`
Name string `json:"name"`
}
if json.Unmarshal(raw, &tc) != nil {
return nil
}
switch tc.Type {
case "auto":
return json.RawMessage(`"auto"`)
case "any":
return json.RawMessage(`"required"`)
case "tool":
b, _ := json.Marshal(map[string]any{"type": "function", "function": map[string]string{"name": tc.Name}})
return b
case "none":
return json.RawMessage(`"none"`)
}
return nil
}
// IRRespToAnthropic 把 IR 非流式响应转为 Anthropic Messages 响应。
func IRRespToAnthropic(resp *aiwire.ChatResponse, id string) aiwire.MessagesResponse {
out := aiwire.MessagesResponse{ID: id, Type: "message", Role: "assistant", Model: resp.Model, Content: []aiwire.AnthBlock{}}
if len(resp.Choices) > 0 {
choice := resp.Choices[0]
if text := choice.Message.Content.JoinText(); text != "" {
out.Content = append(out.Content, aiwire.AnthBlock{Type: "text", Text: text})
}
for _, tc := range choice.Message.ToolCalls {
out.Content = append(out.Content, aiwire.AnthBlock{Type: "tool_use", ID: tc.ID,
Name: tc.Function.Name, Input: argsToJSON(tc.Function.Arguments)})
}
out.StopReason = anthStopReason(choice.FinishReason)
}
if resp.Usage != nil {
out.Usage = aiwire.AnthUsage{InputTokens: resp.Usage.PromptTokens, OutputTokens: resp.Usage.CompletionTokens,
CacheReadInputTokens: resp.Usage.CachedTokens()}
}
return out
}
// argsToJSON 保证 tool_use.input 是合法 JSON 对象(模型可能产出非法片段)。
func argsToJSON(args string) json.RawMessage {
trimmed := strings.TrimSpace(args)
if trimmed == "" {
return json.RawMessage(`{}`)
}
if json.Valid([]byte(trimmed)) {
return json.RawMessage(trimmed)
}
b, _ := json.Marshal(map[string]string{"_raw": args})
return b
}
func anthStopReason(finish string) string {
switch finish {
case "length":
return "max_tokens"
case "tool_calls":
return "tool_use"
default:
return "end_turn"
}
}
// ---- Anthropic 流式状态机 ----
// AnthEvent 是一条待写出的 Anthropic SSE 事件。
type AnthEvent struct {
Event string
Data any
}
// AnthStream 把 IR chunk 流聚合为 Anthropic 事件序列:
// message_start → content_block_start/delta/stop(text 与 tool_use 分块)→ message_delta → message_stop。
type AnthStream struct {
id, model string
started bool
blockOpen bool
blockIsTool bool
toolID string
blockIndex int
stopReason string
usage aiwire.AnthUsage
}
// NewAnthStream 构造状态机;id 为响应消息 ID。
func NewAnthStream(id, model string) *AnthStream {
return &AnthStream{id: id, model: model, blockIndex: -1, stopReason: "end_turn"}
}
// Feed 消费一个 IR chunk,返回应立即写出的事件。
func (st *AnthStream) Feed(chunk aiwire.ChatChunk) []AnthEvent {
var events []AnthEvent
if !st.started {
st.started = true
events = append(events, st.startEvent())
}
if chunk.Usage != nil {
st.usage = aiwire.AnthUsage{InputTokens: chunk.Usage.PromptTokens, OutputTokens: chunk.Usage.CompletionTokens,
CacheReadInputTokens: chunk.Usage.CachedTokens()}
}
for _, choice := range chunk.Choices {
events = append(events, st.feedDelta(choice.Delta)...)
if choice.FinishReason != nil && *choice.FinishReason != "" {
st.stopReason = anthStopReason(*choice.FinishReason)
}
}
return events
}
func (st *AnthStream) startEvent() AnthEvent {
return AnthEvent{Event: "message_start", Data: map[string]any{
"type": "message_start",
"message": map[string]any{
"id": st.id, "type": "message", "role": "assistant", "model": st.model,
"content": []any{}, "stop_reason": nil,
"usage": map[string]int{"input_tokens": 0, "output_tokens": 0},
},
}}
}
// feedDelta 处理文本与工具调用增量,必要时切块。
func (st *AnthStream) feedDelta(d aiwire.Delta) []AnthEvent {
var events []AnthEvent
if d.Content != "" {
if !st.blockOpen || st.blockIsTool {
events = append(events, st.openBlock(false, "", "")...)
}
events = append(events, AnthEvent{Event: "content_block_delta", Data: map[string]any{
"type": "content_block_delta", "index": st.blockIndex,
"delta": map[string]string{"type": "text_delta", "text": d.Content},
}})
}
for _, tc := range d.ToolCalls {
if tc.ID != "" && (!st.blockOpen || !st.blockIsTool || st.toolID != tc.ID) {
events = append(events, st.openBlock(true, tc.ID, tc.Function.Name)...)
}
if tc.Function.Arguments != "" && st.blockOpen && st.blockIsTool {
events = append(events, AnthEvent{Event: "content_block_delta", Data: map[string]any{
"type": "content_block_delta", "index": st.blockIndex,
"delta": map[string]string{"type": "input_json_delta", "partial_json": tc.Function.Arguments},
}})
}
}
return events
}
// openBlock 关闭当前块并打开新块(text 或 tool_use)。
func (st *AnthStream) openBlock(isTool bool, toolID, toolName string) []AnthEvent {
var events []AnthEvent
if st.blockOpen {
events = append(events, st.closeBlockEvent())
}
st.blockOpen, st.blockIsTool, st.toolID = true, isTool, toolID
st.blockIndex++
block := map[string]any{"type": "text", "text": ""}
if isTool {
block = map[string]any{"type": "tool_use", "id": toolID, "name": toolName, "input": map[string]any{}}
}
events = append(events, AnthEvent{Event: "content_block_start", Data: map[string]any{
"type": "content_block_start", "index": st.blockIndex, "content_block": block,
}})
return events
}
func (st *AnthStream) closeBlockEvent() AnthEvent {
return AnthEvent{Event: "content_block_stop", Data: map[string]any{
"type": "content_block_stop", "index": st.blockIndex,
}}
}
// Finish 在上游流结束后收尾:关块 → message_delta(stop_reason+usage)→ message_stop。
func (st *AnthStream) Finish() []AnthEvent {
var events []AnthEvent
if !st.started {
events = append(events, st.startEvent())
st.started = true
}
if st.blockOpen {
events = append(events, st.closeBlockEvent())
st.blockOpen = false
}
events = append(events,
AnthEvent{Event: "message_delta", Data: map[string]any{
"type": "message_delta",
"delta": map[string]any{"stop_reason": st.stopReason, "stop_sequence": nil},
"usage": map[string]int{"output_tokens": st.usage.OutputTokens},
}},
AnthEvent{Event: "message_stop", Data: map[string]any{"type": "message_stop"}},
)
return events
}
// Usage 返回聚合到的用量(供调用日志)。
func (st *AnthStream) Usage() aiwire.AnthUsage { return st.usage }
-136
View File
@@ -1,136 +0,0 @@
package service
import (
"encoding/json"
"strings"
"testing"
"oci-portal/internal/aiwire"
)
func mustAnthReq(t *testing.T, body string) aiwire.MessagesRequest {
t.Helper()
var req aiwire.MessagesRequest
if err := json.Unmarshal([]byte(body), &req); err != nil {
t.Fatalf("unmarshal anthropic request: %v", err)
}
return req
}
func TestAnthropicToIR(t *testing.T) {
req := mustAnthReq(t, `{
"model": "meta.llama-3.3-70b-instruct", "max_tokens": 128, "system": "你是助手",
"messages": [
{"role": "user", "content": "东京天气?"},
{"role": "assistant", "content": [
{"type": "text", "text": "查询中"},
{"type": "tool_use", "id": "t1", "name": "get_weather", "input": {"city": "东京"}}
]},
{"role": "user", "content": [
{"type": "tool_result", "tool_use_id": "t1", "content": "晴 25 度"},
{"type": "text", "text": "继续"}
]}
],
"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}],
"tool_choice": {"type": "any"}
}`)
ir, err := AnthropicToIR(req)
if err != nil {
t.Fatalf("AnthropicToIR: %v", err)
}
roles := make([]string, 0, len(ir.Messages))
for _, m := range ir.Messages {
roles = append(roles, m.Role)
}
want := []string{"system", "user", "assistant", "tool", "user"}
if strings.Join(roles, ",") != strings.Join(want, ",") {
t.Fatalf("roles = %v, want %v", roles, want)
}
if ir.Messages[2].ToolCalls[0].Function.Name != "get_weather" {
t.Errorf("tool_use 未映射: %+v", ir.Messages[2])
}
if ir.Messages[3].ToolCallID != "t1" || ir.Messages[3].Content.JoinText() != "晴 25 度" {
t.Errorf("tool_result 未拆为 tool 消息: %+v", ir.Messages[3])
}
if ir.MaxTokens == nil || *ir.MaxTokens != 128 || len(ir.Tools) != 1 {
t.Errorf("max_tokens/tools 未直通")
}
if string(ir.ToolChoice) != `"required"` {
t.Errorf("tool_choice any → %s, want required", ir.ToolChoice)
}
// image 块拒绝
bad := mustAnthReq(t, `{"model":"m","max_tokens":1,"messages":[{"role":"user","content":[{"type":"image","source":{}}]}]}`)
if _, err := AnthropicToIR(bad); err == nil {
t.Error("image 块应被拒绝")
}
}
func TestIRRespToAnthropic(t *testing.T) {
resp := &aiwire.ChatResponse{
Model: "m1",
Choices: []aiwire.Choice{{
Message: aiwire.ChatMessage{
Role: "assistant",
Content: aiwire.NewTextContent("查到了"),
ToolCalls: []aiwire.ToolCall{{ID: "t1", Type: "function",
Function: aiwire.FunctionCall{Name: "get_weather", Arguments: `{"city":"东京"}`}}},
},
FinishReason: "tool_calls",
}},
Usage: &aiwire.Usage{PromptTokens: 10, CompletionTokens: 5},
}
out := IRRespToAnthropic(resp, "msg_1")
if len(out.Content) != 2 || out.Content[0].Type != "text" || out.Content[1].Type != "tool_use" {
t.Fatalf("content 块 = %+v", out.Content)
}
if out.StopReason != "tool_use" || out.Usage.InputTokens != 10 || out.Usage.OutputTokens != 5 {
t.Errorf("stop_reason/usage 未映射: %+v", out)
}
if string(out.Content[1].Input) != `{"city":"东京"}` {
t.Errorf("tool input = %s", out.Content[1].Input)
}
}
// eventTypes 提取事件类型序列便于断言。
func eventTypes(events []AnthEvent) string {
types := make([]string, 0, len(events))
for _, e := range events {
types = append(types, e.Event)
}
return strings.Join(types, ",")
}
func TestAnthStreamTextAndTool(t *testing.T) {
st := NewAnthStream("msg_1", "m1")
fr := "tool_calls"
var events []AnthEvent
events = append(events, st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{Content: "你"}}}})...)
events = append(events, st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{Content: "好"}}}})...)
events = append(events, st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{
ToolCalls: []aiwire.ToolCallDelta{{ID: "t1", Function: aiwire.FunctionCallDelta{Name: "get_weather", Arguments: `{"ci`}}},
}}}})...)
events = append(events, st.Feed(aiwire.ChatChunk{
Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{ToolCalls: []aiwire.ToolCallDelta{{ID: "t1", Function: aiwire.FunctionCallDelta{Arguments: `ty":"东京"}`}}}}, FinishReason: &fr}},
Usage: &aiwire.Usage{PromptTokens: 8, CompletionTokens: 4},
})...)
events = append(events, st.Finish()...)
got := eventTypes(events)
want := "message_start,content_block_start,content_block_delta,content_block_delta," +
"content_block_stop,content_block_start,content_block_delta,content_block_delta," +
"content_block_stop,message_delta,message_stop"
if got != want {
t.Fatalf("事件序列:\n got %s\nwant %s", got, want)
}
if st.Usage().OutputTokens != 4 {
t.Errorf("usage 未聚合: %+v", st.Usage())
}
}
func TestAnthStreamEmpty(t *testing.T) {
st := NewAnthStream("msg_1", "m1")
events := st.Finish()
if got := eventTypes(events); got != "message_start,message_delta,message_stop" {
t.Errorf("空流事件序列 = %s", got)
}
}
+147 -177
View File
@@ -5,6 +5,7 @@ import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"log"
@@ -30,10 +31,6 @@ const (
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
@@ -59,17 +56,13 @@ type AiGatewayService struct {
// 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{}}
return &AiGatewayService{db: db, configs: configs, client: client, lastTouch: map[uint]time.Time{}}
}
// SetOnChannelsChanged 注册渠道数量变化钩子(渠道创建/删除成功后调用)。
@@ -87,7 +80,7 @@ func (s *AiGatewayService) fireChannelsChanged(ctx context.Context) {
// CreateKey 生成网关密钥并返回明文(仅此一次);customValue 非空时使用给定值,
// group 非空时该密钥只在同分组渠道内路由。
func (s *AiGatewayService) CreateKey(ctx context.Context, name, customValue, group string) (string, *model.AiKey, error) {
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("密钥名称不能为空")
@@ -103,13 +96,28 @@ func (s *AiGatewayService) CreateKey(ctx context.Context, name, customValue, gro
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}
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[:])
@@ -122,8 +130,8 @@ func (s *AiGatewayService) Keys(ctx context.Context) ([]model.AiKey, error) {
return keys, err
}
// UpdateKey 修改密钥名称 / 启用状态 / 分组(group 指针非空即覆盖,可置空)。
func (s *AiGatewayService) UpdateKey(ctx context.Context, id uint, name string, enabled *bool, group *string) error {
// 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
@@ -134,6 +142,14 @@ func (s *AiGatewayService) UpdateKey(ctx context.Context, id uint, name string,
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
}
@@ -298,12 +314,17 @@ func (s *AiGatewayService) ProbeChannel(ctx context.Context, id uint) (*model.Ai
return &fresh, nil
}
// probe 执行探测并返回 (状态, 错误摘要);同时完成模型缓存同步(保留不可用标记)。
// 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 服务不可用或未开放)"
@@ -311,31 +332,12 @@ func (s *AiGatewayService) probe(ctx context.Context, cred oci.Credentials, ch *
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)
return s.probeChat(ctx, cred, ch, models)
}
// 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 次止损。
// probeChat 按候选顺序试调(上限 8):遇「模型不可按需调用」(微调基座 400 / 实体
// 不存在 404)换下一个候选,错误信息带模型名供用户加入黑名单;401/403 与鉴权类 404
// 属租户级直接定论 no_quota;其余错误(元数据标 CHAT 但实际不可对话等)累计 3 次止损。
func (s *AiGatewayService) probeChat(ctx context.Context, cred oci.Credentials, ch *model.AiChannel, models []oci.GenAiModel) (string, string) {
status, detail := "error", "无可试调对话模型"
errBudget := 3
@@ -345,8 +347,7 @@ func (s *AiGatewayService) probeChat(ctx context.Context, cred oci.Credentials,
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))
status, detail = "error", truncateErr(fmt.Sprintf("%s: 不可按需调用,建议加入模型黑名单", m.Name))
case code == 401 || code == 403 || code == 404:
return "no_quota", truncateErr(oci.CompactError(err))
default:
@@ -359,106 +360,8 @@ func (s *AiGatewayService) probeChat(ctx context.Context, cred oci.Credentials,
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 次止损预算。
// probeCandidateCap 是单次探测的候选上限:放宽到 8 让一次探测有机会越过整批
// 不可按需调用的坏模型找到可用者;其他错误另有 3 次止损预算。
const probeCandidateCap = 8
// probeCandidates 只取对话模型,按可靠度排序后跨厂商取候选:主流文本模型优先;
@@ -551,8 +454,7 @@ func truncateErr(msg string) string {
return msg
}
// SyncModels 重新拉取渠道区域的模型列表并覆盖缓存(标记按 OCID 结转),
// 随后触发后台验证:新模型逐个试调,不可按需调用的数十秒内从池中剔除。
// 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 {
@@ -566,29 +468,23 @@ func (s *AiGatewayService) SyncModels(ctx context.Context, id uint) ([]model.AiM
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
}
s.validateModelsAsync(ctx, id)
return s.channelModels(ctx, id)
}
// replaceModels 以事务整组覆盖渠道模型缓存,按 OCID 结转不可用标记与验证时间
// (OCID 变化视为新条目,自然回到待验证状态)。
// replaceModels 以事务整组覆盖渠道模型缓存
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)
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 {
@@ -601,33 +497,17 @@ func (s *AiGatewayService) replaceModels(ctx context.Context, channelID uint, mo
})
}
// 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 非空时仅聚合该分组渠道(与密钥分组路由口径一致)。
// 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)
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)
}
@@ -658,7 +538,6 @@ func (s *AiGatewayService) DeprecatingModels(ctx context.Context, within time.Du
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
@@ -679,7 +558,7 @@ func (s *AiGatewayService) DeprecatingModels(ctx context.Context, within time.Du
return out, nil
}
// ProbeAll 逐个探测全部渠道并顺带验证模型池,返回状态汇总;供 AI 探测后台任务调用。
// ProbeAll 逐个探测全部渠道,返回状态汇总;供 AI 探测后台任务调用。
func (s *AiGatewayService) ProbeAll(ctx context.Context) (string, error) {
chs, err := s.Channels(ctx)
if err != nil {
@@ -695,8 +574,6 @@ func (s *AiGatewayService) ProbeAll(ctx context.Context) (string, 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"])
@@ -706,6 +583,99 @@ func (s *AiGatewayService) ProbeAll(ctx context.Context) (string, error) {
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)。
+25 -25
View File
@@ -3,6 +3,7 @@ package service
import (
"context"
"errors"
"io"
"math/rand"
"time"
@@ -26,24 +27,25 @@ type aiCandidate struct {
modelOcid string
}
// Chat 编排非流式调用:选渠道 → 调用 → 可重试错误换渠道(整请求上限 3 次)。
// group 非空时只在同分组渠道内路由(取自调用密钥)
func (s *AiGatewayService) Chat(ctx context.Context, ir aiwire.ChatRequest, group string) (*aiwire.ChatResponse, ChatMeta, error) {
// RespPassthrough 编排一次非流式直通调用:选渠道(priority→加权随机)→ 调用 →
// 可重试错误换渠道(整请求上限 3 次)并维护熔断;group 非空时只在同分组渠道内路由。
// 上游为 OpenAI-compatible /actions/v1/responses(实测可用,无 Oracle 文档合同)。
func (s *AiGatewayService) RespPassthrough(ctx context.Context, raw []byte, modelName, group string) ([]byte, ChatMeta, error) {
meta := ChatMeta{}
excluded := map[uint]bool{}
var lastErr error
for attempt := 0; attempt < 3; attempt++ {
cand, err := s.pick(ctx, ir.Model, group, "CHAT", excluded)
cand, err := s.pick(ctx, modelName, group, "CHAT", excluded)
if err != nil {
return nil, meta, firstErr(lastErr, err)
}
meta.ChannelID, meta.ChannelName = cand.ch.ID, cand.ch.Name
resp, err := s.callOnce(ctx, cand, ir)
payload, err := s.passthroughOnce(ctx, cand, raw)
if err == nil {
s.markSuccess(ctx, cand.ch.ID)
return resp, meta, nil
return payload, meta, nil
}
retry, penalize := s.noteCallErr(ctx, cand, err)
retry, penalize := switchable(err)
if !retry {
return nil, meta, err
}
@@ -57,22 +59,22 @@ func (s *AiGatewayService) Chat(ctx context.Context, ir aiwire.ChatRequest, grou
return nil, meta, lastErr
}
func (s *AiGatewayService) callOnce(ctx context.Context, cand *aiCandidate, ir aiwire.ChatRequest) (*aiwire.ChatResponse, error) {
func (s *AiGatewayService) passthroughOnce(ctx context.Context, cand *aiCandidate, raw []byte) ([]byte, error) {
cred, err := s.configs.credentialsByID(ctx, cand.ch.OciConfigID)
if err != nil {
return nil, err
}
return s.client.GenAiChat(ctx, cred, cand.ch.Region, cand.modelOcid, ir)
return s.client.GenAiCompatResponses(ctx, cred, cand.ch.Region, raw)
}
// OpenStream 编排流式调用:流建立成功即绑定渠道,建立失败可换渠道重试。
// group 语义与 Chat 相同
func (s *AiGatewayService) OpenStream(ctx context.Context, ir aiwire.ChatRequest, group string) (oci.GenAiStream, ChatMeta, error) {
// RespPassthroughStream 编排流式直通:流建立成功即绑定渠道,建立失败按 switchable
// 换渠道重试;建立后的中断不重试、不计熔断(与 OpenStream 语义一致)
func (s *AiGatewayService) RespPassthroughStream(ctx context.Context, raw []byte, modelName, group string) (io.ReadCloser, ChatMeta, error) {
meta := ChatMeta{}
excluded := map[uint]bool{}
var lastErr error
for attempt := 0; attempt < 3; attempt++ {
cand, err := s.pick(ctx, ir.Model, group, "CHAT", excluded)
cand, err := s.pick(ctx, modelName, group, "CHAT", excluded)
if err != nil {
return nil, meta, firstErr(lastErr, err)
}
@@ -81,12 +83,12 @@ func (s *AiGatewayService) OpenStream(ctx context.Context, ir aiwire.ChatRequest
if err != nil {
return nil, meta, err
}
stream, err := s.client.GenAiChatStream(ctx, cred, cand.ch.Region, cand.modelOcid, ir)
stream, err := s.client.GenAiCompatResponsesStream(ctx, cred, cand.ch.Region, raw)
if err == nil {
s.markSuccess(ctx, cand.ch.ID)
return stream, meta, nil
}
retry, penalize := s.noteCallErr(ctx, cand, err)
retry, penalize := switchable(err)
if !retry {
return nil, meta, err
}
@@ -108,12 +110,11 @@ func firstErr(lastErr, pickErr error) error {
return pickErr
}
// noteCallErr 汇总一次调用失败:模型级不可用(微调基座 400 / 实体不存在 404)先标记
// 剔除该 (渠道, 模型),换渠道重试且不计熔断——这是模型×区域供给问题而非渠道健康问题;
// 429 / 5xx / 网络错误换渠道并计熔断;其余 4xx 直接透传。
func (s *AiGatewayService) noteCallErr(ctx context.Context, cand *aiCandidate, err error) (retry, penalize bool) {
// switchable 判断调用失败是否换渠道重试、是否计入熔断:模型级不可用(微调基座
// 400 / 实体不存在 404)换渠道但不计熔断——这是模型×区域供给问题而非渠道健康
// 问题,持久解决靠加入模型黑名单;429 / 5xx / 网络错误换渠道并计熔断;其余 4xx 直接透传。
func switchable(err error) (retry, penalize bool) {
if oci.IsModelUnavailable(err) {
s.markModelUnusable(ctx, cand.ch.ID, cand.modelOcid, oci.CompactError(err))
return true, false
}
if status, ok := oci.ServiceStatus(err); ok {
@@ -154,11 +155,10 @@ func (s *AiGatewayService) pick(ctx context.Context, modelName, group, capabilit
return &aiCandidate{ch: chosen, modelOcid: ocids[chosen.ID]}, nil
}
// modelChannels 查出提供该模型的渠道 ID 及各自的模型 OCID(剔除不可按需调用标记);
// capability=CHAT 时兼容存量空串(加列前只同步对话模型);unusable 需兼容 NULL
// (AutoMigrate 加列后、首次重同步前的存量行)。
// modelChannels 查出提供该模型的渠道 ID 及各自的模型 OCID;
// capability=CHAT 时兼容存量空串(加列前只同步对话模型)
func (s *AiGatewayService) modelChannels(ctx context.Context, modelName, capability string) (map[uint]string, []uint, error) {
q := s.db.WithContext(ctx).Where("name = ? AND (unusable = ? OR unusable IS NULL)", modelName, false)
q := s.db.WithContext(ctx).Where("name = ?", modelName)
if capability == "CHAT" {
q = q.Where("capability IN ?", []string{"CHAT", ""})
} else {
@@ -234,7 +234,7 @@ func (s *AiGatewayService) Embeddings(ctx context.Context, req aiwire.Embeddings
s.markSuccess(ctx, cand.ch.ID)
return resp, meta, nil
}
retry, penalize := s.noteCallErr(ctx, cand, err)
retry, penalize := switchable(err)
if !retry {
return nil, meta, err
}
+29
View File
@@ -0,0 +1,29 @@
package service
import (
"strings"
"testing"
"oci-portal/internal/oci"
)
// TestSupportedGatewayModels 断言 CHAT 目录仅保留兼容面厂商,EMBEDDING 不受影响。
func TestSupportedGatewayModels(t *testing.T) {
in := []oci.GenAiModel{
{Name: "xai.grok-4.3", Capability: "CHAT"},
{Name: "meta.llama-4-maverick-17b-128e", Capability: "CHAT"},
{Name: "openai.gpt-oss-120b", Capability: "CHAT"},
{Name: "google.gemini-2.5-flash", Capability: "CHAT"},
{Name: "cohere.command-a-03-2025", Capability: "CHAT"},
{Name: "cohere.embed-v4.0", Capability: "EMBEDDING"},
}
out := supportedGatewayModels(in)
names := make([]string, 0, len(out))
for _, m := range out {
names = append(names, m.Name)
}
want := "xai.grok-4.3,meta.llama-4-maverick-17b-128e,openai.gpt-oss-120b,cohere.embed-v4.0"
if got := strings.Join(names, ","); got != want {
t.Fatalf("过滤结果 = %s, want %s", got, want)
}
}
+267 -229
View File
@@ -1,9 +1,11 @@
package service
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"strings"
"testing"
"time"
@@ -45,15 +47,46 @@ type gatewayStubClient struct {
probeCode int
probeErr error
// probeSeq 非空时逐次弹出,模拟按候选依次试调;弹尽后回落 probeCode/probeErr
probeSeq []probeResult
chatResp *aiwire.ChatResponse
chatErrs []error
chatCalls int
regions []string
probeSeq []probeResult
// probedNames 记录试调过的模型名,供断言候选过滤
probedNames []string
chatCalls int
regions []string
embedVecs [][]float32
embedUsage *aiwire.Usage
embedErr error
passPayload []byte
passErrs []error
passCalls int
passRegions []string
}
func (f *gatewayStubClient) GenAiCompatResponses(ctx context.Context, cred oci.Credentials, region string, body []byte) ([]byte, error) {
f.passCalls++
f.passRegions = append(f.passRegions, region)
if len(f.passErrs) > 0 {
err := f.passErrs[0]
f.passErrs = f.passErrs[1:]
if err != nil {
return nil, err
}
}
return f.passPayload, nil
}
func (f *gatewayStubClient) GenAiCompatResponsesStream(ctx context.Context, cred oci.Credentials, region string, body []byte) (io.ReadCloser, error) {
f.passCalls++
f.passRegions = append(f.passRegions, region)
if len(f.passErrs) > 0 {
err := f.passErrs[0]
f.passErrs = f.passErrs[1:]
if err != nil {
return nil, err
}
}
return io.NopCloser(bytes.NewReader(f.passPayload)), nil
}
func (f *gatewayStubClient) GenAiEmbed(ctx context.Context, cred oci.Credentials, region, modelOcid string, inputs []string, dimensions *int) ([][]float32, *aiwire.Usage, error) {
@@ -71,6 +104,7 @@ type probeResult struct {
}
func (f *gatewayStubClient) GenAiProbeChat(ctx context.Context, cred oci.Credentials, region, modelOcid, modelName string) (int, error) {
f.probedNames = append(f.probedNames, modelName)
if len(f.probeSeq) > 0 {
r := f.probeSeq[0]
f.probeSeq = f.probeSeq[1:]
@@ -79,23 +113,10 @@ func (f *gatewayStubClient) GenAiProbeChat(ctx context.Context, cred oci.Credent
return f.probeCode, f.probeErr
}
func (f *gatewayStubClient) GenAiChat(ctx context.Context, cred oci.Credentials, region, modelOcid string, ir aiwire.ChatRequest) (*aiwire.ChatResponse, error) {
f.chatCalls++
f.regions = append(f.regions, region)
if len(f.chatErrs) > 0 {
err := f.chatErrs[0]
f.chatErrs = f.chatErrs[1:]
if err != nil {
return nil, err
}
}
return f.chatResp, nil
}
func newTestGateway(t *testing.T, client oci.Client) (*AiGatewayService, *OciConfigService) {
t.Helper()
svc := newTestService(t, client)
if err := svc.db.AutoMigrate(&model.AiKey{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}); err != nil {
if err := svc.db.AutoMigrate(&model.AiKey{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiModelBlacklist{}, &model.AiCallLog{}); err != nil {
t.Fatalf("auto migrate ai tables: %v", err)
}
return NewAiGatewayService(svc.db, svc, client), svc
@@ -105,14 +126,14 @@ func TestAiKeyLifecycle(t *testing.T) {
gw, _ := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}})
ctx := context.Background()
raw, key, err := gw.CreateKey(ctx, "免费-ai-api-key", "test-api-key", "")
raw, key, err := gw.CreateKey(ctx, "免费-ai-api-key", "test-api-key", "", nil)
if err != nil || raw != "test-api-key" || key.Tail != "-key" {
t.Fatalf("CreateKey 自定义值 = %q %+v, %v", raw, key, err)
}
if _, _, err := gw.CreateKey(ctx, "short", "abc", ""); err == nil {
if _, _, err := gw.CreateKey(ctx, "short", "abc", "", nil); err == nil {
t.Error("过短自定义密钥应被拒绝")
}
auto, _, err := gw.CreateKey(ctx, "auto", "", "")
auto, _, err := gw.CreateKey(ctx, "auto", "", "", nil)
if err != nil || len(auto) < 40 || auto[:3] != "sk-" {
t.Fatalf("CreateKey 随机值 = %q, %v", auto, err)
}
@@ -124,7 +145,7 @@ func TestAiKeyLifecycle(t *testing.T) {
t.Errorf("错误密钥 err = %v", err)
}
off := false
_ = gw.UpdateKey(ctx, key.ID, "", &off, nil)
_ = gw.UpdateKey(ctx, key.ID, "", &off, nil, nil)
if _, err := gw.VerifyKey(ctx, "test-api-key"); !errors.Is(err, ErrAiKeyInvalid) {
t.Errorf("禁用后 VerifyKey err = %v", err)
}
@@ -187,51 +208,6 @@ func seedChannel(t *testing.T, gw *AiGatewayService, cfgID uint, region string,
return ch
}
func TestAiChatRetrySwitchesChannel(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
chatResp: &aiwire.ChatResponse{Model: "meta.llama-3.3-70b-instruct", Choices: []aiwire.Choice{{Message: aiwire.ChatMessage{Role: "assistant", Content: aiwire.NewTextContent("hi")}, FinishReason: "stop"}}},
chatErrs: []error{stubServiceError{status: 429}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
// 两个同优先级渠道
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
ctx := context.Background()
resp, meta, err := gw.Chat(ctx, aiwire.ChatRequest{Model: "meta.llama-3.3-70b-instruct", Messages: []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent("你好")}}}, "")
if err != nil || resp == nil {
t.Fatalf("Chat = %v, %v", resp, err)
}
if meta.Retries != 1 || client.chatCalls != 2 {
t.Errorf("应换渠道重试一次: retries=%d calls=%d", meta.Retries, client.chatCalls)
}
if len(client.regions) != 2 && client.regions[0] == client.regions[1] {
t.Errorf("重试未换渠道: %v", client.regions)
}
// 未知模型
if _, _, err := gw.Chat(ctx, aiwire.ChatRequest{Model: "no-such-model"}, ""); !errors.Is(err, ErrAiUnknownModel) {
t.Errorf("未知模型 err = %v", err)
}
}
func TestAiChatNonRetryablePassThrough(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
chatErrs: []error{stubServiceError{status: 400}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
_, meta, err := gw.Chat(context.Background(), aiwire.ChatRequest{Model: "meta.llama-3.3-70b-instruct"}, "")
if err == nil || meta.Retries != 0 || client.chatCalls != 1 {
t.Errorf("400 应直接透传不重试: err=%v retries=%d calls=%d", err, meta.Retries, client.chatCalls)
}
}
func TestPickPriorityAndBreaker(t *testing.T) {
gw, svc := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}})
cfg := importAliveConfig(t, svc)
@@ -282,8 +258,8 @@ func TestMarkFailureBackoff(t *testing.T) {
func TestAiGroupRouting(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
chatResp: &aiwire.ChatResponse{Model: "meta.llama-3.3-70b-instruct", Choices: []aiwire.Choice{{Message: aiwire.ChatMessage{Role: "assistant", Content: aiwire.NewTextContent("hi")}, FinishReason: "stop"}}},
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
passPayload: []byte(`{"id":"resp_1","usage":{"input_tokens":3,"output_tokens":1}}`),
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
@@ -292,17 +268,17 @@ func TestAiGroupRouting(t *testing.T) {
gw.db.Model(vip).Update("channel_group", "vip")
gw.db.Create(&model.AiModelCache{ChannelID: other.ID, ModelOcid: "ocid1..cr", Name: "cohere.command-r", Vendor: "cohere", SyncedAt: time.Now()})
ctx := context.Background()
req := aiwire.ChatRequest{Model: "meta.llama-3.3-70b-instruct", Messages: []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent("你好")}}}
body := []byte(`{"model":"meta.llama-3.3-70b-instruct","input":"你好"}`)
// 分组密钥只落同分组渠道
for i := 0; i < 5; i++ {
_, meta, err := gw.Chat(ctx, req, "vip")
_, meta, err := gw.RespPassthrough(ctx, body, "meta.llama-3.3-70b-instruct", "vip")
if err != nil || meta.ChannelID != vip.ID {
t.Fatalf("vip 组第 %d 次落点 = %d, err %v, want %d", i, meta.ChannelID, err, vip.ID)
}
}
// 分组内无渠道 → 无可用渠道
if _, _, err := gw.Chat(ctx, req, "nope"); !errors.Is(err, ErrAiNoChannel) {
if _, _, err := gw.RespPassthrough(ctx, body, "meta.llama-3.3-70b-instruct", "nope"); !errors.Is(err, ErrAiNoChannel) {
t.Errorf("空分组 err = %v, want ErrAiNoChannel", err)
}
// 模型列表按分组过滤
@@ -315,12 +291,12 @@ func TestAiGroupRouting(t *testing.T) {
t.Errorf("不限组模型数 = %d, want 2", len(all.Data))
}
// 密钥分组落库
_, key, err := gw.CreateKey(ctx, "vip-key", "vip-secret-1234", "vip")
_, key, err := gw.CreateKey(ctx, "vip-key", "vip-secret-1234", "vip", nil)
if err != nil || key.Group != "vip" {
t.Fatalf("CreateKey group = %+v, %v", key, err)
}
empty := ""
_ = gw.UpdateKey(ctx, key.ID, "", nil, &empty)
_ = gw.UpdateKey(ctx, key.ID, "", nil, &empty, nil)
var fresh model.AiKey
gw.db.First(&fresh, key.ID)
if fresh.Group != "" {
@@ -465,7 +441,8 @@ func entityNotFoundErr() stubServiceError {
}
func TestProbeSkipsUnavailableModels(t *testing.T) {
// 首选 llama 不可按需调用(基座 400 / 实体 404)→ 标记剔除 → 换候选成功 → 渠道判可用
// 首选 llama 不可按需调用(基座 400 / 实体 404)→ 换候选成功 → 渠道判可用;
// 模型全量入库不再自动标记,持久剔除交由用户手动拉黑
tests := []struct {
name string
bad probeResult
@@ -496,112 +473,16 @@ func TestProbeSkipsUnavailableModels(t *testing.T) {
if err != nil || probed.ProbeStatus != "ok" {
t.Fatalf("坏候选后应换候选并判可用: %+v, %v", probed, err)
}
var row model.AiModelCache
if err := gw.db.Where("channel_id = ? AND model_ocid = ?", ch.ID, "m1").First(&row).Error; err != nil || !row.Unusable {
t.Fatalf("首个坏候选应被标记不可用: %+v, %v", row, err)
}
// 再次探测:同步保留标记,m1 不再进候选(probeSeq 只需一次 200)
client.probeSeq = []probeResult{{200, nil}}
probed, err = gw.ProbeChannel(ctx, ch.ID)
if err != nil || probed.ProbeStatus != "ok" {
t.Fatalf("复测应跳过已标记模型: %+v, %v", probed, err)
}
var again model.AiModelCache
gw.db.Where("channel_id = ? AND model_ocid = ?", ch.ID, "m1").First(&again)
if !again.Unusable {
t.Error("探测触发的同步应保留不可用标记")
}
// 手动同步同样保留标记(坏模型不随重新同步复活);近期已检的标记不被后台验证翻转
client.probeCode = 200
if _, err := gw.SyncModels(ctx, ch.ID); err != nil {
t.Fatalf("SyncModels: %v", err)
}
gw.Wait()
var kept model.AiModelCache
gw.db.Where("channel_id = ? AND model_ocid = ?", ch.ID, "m1").First(&kept)
if !kept.Unusable {
t.Error("手动同步不应清除不可用标记")
rows, _ := gw.channelModels(ctx, ch.ID)
if len(rows) != 3 {
t.Errorf("模型应全量入库, got %d", len(rows))
}
})
}
}
func TestValidateModelsAfterSync(t *testing.T) {
// 同步后后台验证:坏模型标记剔除、好模型记录已检、其他 4xx 不改状态、非对话模型不试调
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
models: []oci.GenAiModel{
{Ocid: "v1", Name: "meta.llama-4-maverick", Vendor: "meta", Capability: "CHAT"},
{Ocid: "v2", Name: "xai.grok-4", Vendor: "xai", Capability: "CHAT"},
{Ocid: "v3", Name: "xai.grok-voice-agent", Vendor: "xai", Capability: "CHAT"},
{Ocid: "v4", Name: "cohere.embed-v4.0", Vendor: "cohere", Capability: "EMBEDDING"},
},
probeSeq: []probeResult{{404, entityNotFoundErr()}, {200, nil}, {400, stubServiceError{status: 400}}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
ctx := context.Background()
ch, err := gw.CreateChannel(ctx, ChannelInput{OciConfigID: cfg.ID, Region: "us-ashburn-1"})
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
if _, err := gw.SyncModels(ctx, ch.ID); err != nil {
t.Fatalf("SyncModels: %v", err)
}
gw.Wait()
want := map[string]struct {
unusable bool
checked bool
}{
"v1": {true, true}, // 实体 404 → 标记剔除
"v2": {false, true}, // 200 → 可用已检
"v3": {false, true}, // 普通 400 → 已检不标记
"v4": {false, false}, // EMBEDDING 不试调
}
rows, _ := gw.channelModels(ctx, ch.ID)
for _, r := range rows {
w := want[r.ModelOcid]
if r.Unusable != w.unusable || (r.CheckedAt != nil) != w.checked {
t.Errorf("%s: unusable=%v checked=%v, want %+v", r.ModelOcid, r.Unusable, r.CheckedAt != nil, w)
}
}
list, err := gw.GatewayModels(ctx, "")
if err != nil || len(list.Data) != 3 {
t.Errorf("网关列表应只剔除被标记的坏模型(余 v2/v3/v4), got %+v, %v", list.Data, err)
}
}
func TestValidateRecheckUnmarksRecovered(t *testing.T) {
// 已标记模型超过复检间隔后重验:恢复供给(200)自动解除标记
client := &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}, probeCode: 200}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
ch := seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
ctx := context.Background()
old := time.Now().Add(-25 * time.Hour)
gw.db.Model(&model.AiModelCache{}).Where("channel_id = ?", ch.ID).
Updates(map[string]any{"unusable": true, "unusable_reason": "x", "checked_at": old})
gw.validateChannelModels(ctx, ch.ID)
var row model.AiModelCache
gw.db.Where("channel_id = ?", ch.ID).First(&row)
if row.Unusable || row.UnusableReason != "" || row.CheckedAt == nil || !row.CheckedAt.After(old) {
t.Errorf("超期复检应解除标记并刷新已检时间: %+v", row)
}
// 未超期的标记不复检(probeSeq 为空、fallback 200 也不会被消费)
fresh := time.Now()
gw.db.Model(&model.AiModelCache{}).Where("channel_id = ?", ch.ID).
Updates(map[string]any{"unusable": true, "checked_at": fresh})
gw.validateChannelModels(ctx, ch.ID)
gw.db.Where("channel_id = ?", ch.ID).First(&row)
if !row.Unusable {
t.Error("未超期的标记不应被复检翻转")
}
}
func TestProbeAuth404StillNoQuota(t *testing.T) {
// 鉴权类 404(NotAuthorizedOrNotFound)仍属租户级,直接定论 no_quota 且不标记模型
// 鉴权类 404(NotAuthorizedOrNotFound)仍属租户级,直接定论 no_quota
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
models: []oci.GenAiModel{{Ocid: "m1", Name: "meta.llama-3.3-70b-instruct", Vendor: "meta"}},
@@ -617,31 +498,26 @@ func TestProbeAuth404StillNoQuota(t *testing.T) {
if probed.ProbeStatus != "no_quota" {
t.Errorf("鉴权 404 status = %q, want no_quota", probed.ProbeStatus)
}
var marked int64
gw.db.Model(&model.AiModelCache{}).Where("unusable = ?", true).Count(&marked)
if marked != 0 {
t.Errorf("鉴权 404 不应标记模型, marked=%d", marked)
}
}
func TestAiChatFinetuneSwitchesChannelWithoutPenalty(t *testing.T) {
// 微调基座 400 换渠道重试成功,且不计入熔断失败
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
chatResp: &aiwire.ChatResponse{Model: "meta.llama-3.3-70b-instruct", Choices: []aiwire.Choice{{Message: aiwire.ChatMessage{Role: "assistant", Content: aiwire.NewTextContent("hi")}, FinishReason: "stop"}}},
chatErrs: []error{finetuneBaseErr()},
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
passPayload: []byte(`{"id":"resp_1","usage":{"input_tokens":3,"output_tokens":1}}`),
passErrs: []error{finetuneBaseErr()},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
resp, meta, err := gw.Chat(context.Background(), aiwire.ChatRequest{Model: "meta.llama-3.3-70b-instruct", Messages: []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent("你好")}}}, "")
resp, meta, err := gw.RespPassthrough(context.Background(), []byte(`{}`), "meta.llama-3.3-70b-instruct", "")
if err != nil || resp == nil {
t.Fatalf("Chat = %v, %v", resp, err)
t.Fatalf("RespPassthrough = %v, %v", resp, err)
}
if meta.Retries != 1 || client.chatCalls != 2 {
t.Errorf("应换渠道重试一次: retries=%d calls=%d", meta.Retries, client.chatCalls)
if meta.Retries != 1 || client.passCalls != 2 {
t.Errorf("应换渠道重试一次: retries=%d calls=%d", meta.Retries, client.passCalls)
}
var chs []model.AiChannel
gw.db.Find(&chs)
@@ -650,38 +526,93 @@ func TestAiChatFinetuneSwitchesChannelWithoutPenalty(t *testing.T) {
t.Errorf("微调基座 400 不应计入熔断: 渠道 %s failCount=%d", ch.Name, ch.FailCount)
}
}
// 失败渠道的该模型被标记,不再参与路由;成功渠道不受影响
region := client.regions[0]
var row model.AiModelCache
gw.db.Where("model_ocid = ?", "ocid1..m-"+region).First(&row)
if !row.Unusable {
t.Errorf("失败渠道的模型应被标记不可用: %+v", row)
}
func TestBlacklistLifecycle(t *testing.T) {
// 拉黑:全渠道同名缓存删除、列表与路由立即不可见;重复/空名被拒;
// 移出黑名单后重新同步恢复入池
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
models: []oci.GenAiModel{{Ocid: "ocid1..m-eu-frankfurt-1", Name: "meta.llama-3.3-70b-instruct", Vendor: "meta"}},
}
var usable int64
gw.db.Model(&model.AiModelCache{}).Where("unusable = ?", false).Count(&usable)
if usable != 1 {
t.Errorf("成功渠道模型应保持可用, usable=%d", usable)
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
ch := seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
ctx := context.Background()
item, err := gw.AddBlacklist(ctx, " meta.llama-3.3-70b-instruct ")
if err != nil || item.Name != "meta.llama-3.3-70b-instruct" {
t.Fatalf("AddBlacklist = %+v, %v", item, err)
}
var left int64
gw.db.Model(&model.AiModelCache{}).Count(&left)
if left != 0 {
t.Errorf("拉黑应删除全部渠道的同名缓存, 剩 %d", left)
}
if _, _, err := gw.RespPassthrough(ctx, []byte(`{}`), "meta.llama-3.3-70b-instruct", ""); !errors.Is(err, ErrAiUnknownModel) {
t.Errorf("拉黑后路由应按未知模型拒绝: %v", err)
}
if _, err := gw.AddBlacklist(ctx, "meta.llama-3.3-70b-instruct"); err == nil {
t.Error("重复拉黑应被拒绝")
}
if _, err := gw.AddBlacklist(ctx, " "); err == nil {
t.Error("空模型名应被拒绝")
}
rows, err := gw.Blacklist(ctx)
if err != nil || len(rows) != 1 {
t.Fatalf("Blacklist = %+v, %v", rows, err)
}
if err := gw.RemoveBlacklist(ctx, rows[0].ID); err != nil {
t.Fatalf("RemoveBlacklist: %v", err)
}
if err := gw.RemoveBlacklist(ctx, rows[0].ID); err == nil {
t.Error("移除不存在的条目应报错")
}
// 移出后重新同步:模型恢复入池
if _, err := gw.SyncModels(ctx, ch.ID); err != nil {
t.Fatalf("SyncModels: %v", err)
}
list, _ := gw.GatewayModels(ctx, "")
if len(list.Data) != 1 {
t.Errorf("移出黑名单并同步后应恢复入池: %+v", list.Data)
}
}
func TestUnusableModelExcludedFromPoolAndRouting(t *testing.T) {
// 唯一渠道的模型被标记后:网关列表不再展示,路由按未知模型拒绝,渠道详情仍可见标记
gw, svc := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}})
func TestSyncAndProbeFilterBlacklisted(t *testing.T) {
// 同步与探测都过滤黑名单模型:不入库、不进试调候选
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
models: []oci.GenAiModel{
{Ocid: "m1", Name: "meta.llama-3-70b-instruct", Vendor: "meta"},
{Ocid: "m2", Name: "cohere.command-a-03-2025", Vendor: "cohere"},
},
probeCode: 200,
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
ch := seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
ctx := context.Background()
gw.markModelUnusable(ctx, ch.ID, "ocid1..m-eu-frankfurt-1", "Entity with key … not found")
list, err := gw.GatewayModels(ctx, "")
if err != nil || len(list.Data) != 0 {
t.Errorf("已标记模型不应出现在网关列表: %+v, %v", list.Data, err)
if _, err := gw.AddBlacklist(ctx, "meta.llama-3-70b-instruct"); err != nil {
t.Fatalf("AddBlacklist: %v", err)
}
if _, _, err := gw.Chat(ctx, aiwire.ChatRequest{Model: "meta.llama-3.3-70b-instruct"}, ""); !errors.Is(err, ErrAiUnknownModel) {
t.Errorf("已标记模型路由应拒绝: %v", err)
ch, err := gw.CreateChannel(ctx, ChannelInput{OciConfigID: cfg.ID, Region: "eu-frankfurt-1"})
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
rows, err := gw.channelModels(ctx, ch.ID)
if err != nil || len(rows) != 1 || !rows[0].Unusable || rows[0].UnusableReason == "" {
t.Errorf("渠道详情应保留标记行便于排查: %+v, %v", rows, err)
probed, err := gw.ProbeChannel(ctx, ch.ID)
if err != nil || probed.ProbeStatus != "ok" {
t.Fatalf("ProbeChannel = %+v, %v", probed, err)
}
rows, _ := gw.channelModels(ctx, ch.ID)
if len(rows) != 1 || rows[0].Name != "cohere.command-a-03-2025" {
t.Errorf("探测同步应过滤黑名单模型: %+v", rows)
}
if len(client.probedNames) != 1 || client.probedNames[0] != "cohere.command-a-03-2025" {
t.Errorf("试调候选不应包含黑名单模型: %v", client.probedNames)
}
models, err := gw.SyncModels(ctx, ch.ID)
if err != nil || len(models) != 1 || models[0].Name != "cohere.command-a-03-2025" {
t.Errorf("SyncModels 应过滤黑名单模型: %+v, %v", models, err)
}
}
@@ -745,7 +676,7 @@ func TestAiEmbeddings(t *testing.T) {
t.Errorf("chat 模型走 embeddings err = %v, want ErrAiUnknownModel", err)
}
// embedding 模型名打 chat:同样未知模型
if _, _, err := gw.Chat(ctx, aiwire.ChatRequest{Model: "cohere.embed-v4.0", Messages: []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent("x")}}}, ""); !errors.Is(err, ErrAiUnknownModel) {
if _, _, err := gw.RespPassthrough(ctx, []byte(`{}`), "cohere.embed-v4.0", ""); !errors.Is(err, ErrAiUnknownModel) {
t.Errorf("embedding 模型走 chat err = %v, want ErrAiUnknownModel", err)
}
}
@@ -756,7 +687,7 @@ func TestAiContentLogSwitch(t *testing.T) {
t.Fatalf("migrate content log: %v", err)
}
ctx := context.Background()
_, key, err := gw.CreateKey(ctx, "k1", "content-key-1234", "")
_, key, err := gw.CreateKey(ctx, "k1", "content-key-1234", "", nil)
if err != nil {
t.Fatalf("CreateKey: %v", err)
}
@@ -793,23 +724,130 @@ func TestAiContentLogSwitch(t *testing.T) {
}
}
func TestLegacyNullUnusableRowsStayVisible(t *testing.T) {
// 升级路径回归:AutoMigrate 加列后、首次重同步前,存量行 unusable 为 NULL,
// 网关列表与路由必须照常包含这些行,不能因 unusable = false 过滤而整池消失
gw, svc := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}})
func TestRespPassthroughSwitchesChannel(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
passPayload: []byte(`{"id":"resp_1","usage":{"input_tokens":9,"output_tokens":3}}`),
passErrs: []error{stubServiceError{status: 429}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
ch := seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
ctx := context.Background()
if err := gw.db.Exec("UPDATE ai_model_caches SET unusable = NULL WHERE channel_id = ?", ch.ID).Error; err != nil {
t.Fatalf("set legacy null: %v", err)
}
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
list, err := gw.GatewayModels(ctx, "")
if err != nil || len(list.Data) != 1 {
t.Errorf("存量 NULL 行应仍在网关列表: %+v, %v", list.Data, err)
payload, meta, err := gw.RespPassthrough(context.Background(), []byte(`{}`), "meta.llama-3.3-70b-instruct", "")
if err != nil || len(payload) == 0 {
t.Fatalf("RespPassthrough = %v, %v", payload, err)
}
_, ids, err := gw.modelChannels(ctx, "meta.llama-3.3-70b-instruct", "CHAT")
if err != nil || len(ids) != 1 {
t.Errorf("存量 NULL 行应仍参与路由: %v, %v", ids, err)
if meta.Retries != 1 || client.passCalls != 2 {
t.Errorf("应换渠道重试一次: retries=%d calls=%d", meta.Retries, client.passCalls)
}
if len(client.passRegions) == 2 && client.passRegions[0] == client.passRegions[1] {
t.Errorf("重试未换渠道: %v", client.passRegions)
}
}
func TestRespPassthroughNonRetryable(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
passErrs: []error{stubServiceError{status: 400}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
_, meta, err := gw.RespPassthrough(context.Background(), []byte(`{}`), "meta.llama-3.3-70b-instruct", "")
if err == nil || meta.Retries != 0 || client.passCalls != 1 {
t.Errorf("400 应直接透传不重试: err=%v retries=%d calls=%d", err, meta.Retries, client.passCalls)
}
}
func TestNormalizeKeyModels(t *testing.T) {
tests := []struct {
name string
in []string
want []string
}{
{"nil 输入", nil, nil},
{"全空串", []string{"", " "}, nil},
{"trim 与保序去重", []string{" a ", "b", "a", "", "b "}, []string{"a", "b"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := normalizeKeyModels(tt.in)
if len(got) != len(tt.want) {
t.Fatalf("normalizeKeyModels(%v) = %v, want %v", tt.in, got, tt.want)
}
for i := range tt.want {
if got[i] != tt.want[i] {
t.Fatalf("normalizeKeyModels(%v) = %v, want %v", tt.in, got, tt.want)
}
}
})
}
}
// TestAiKeyModelsPersistence 钉住 serializer:json 字段经 Create 与 Updates(map) 两条路径的往返。
func TestAiKeyModelsPersistence(t *testing.T) {
gw, _ := newTestGateway(t, &gatewayStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}})
ctx := context.Background()
_, key, err := gw.CreateKey(ctx, "m-key", "model-key-1234", "", []string{" x-model ", "x-model", "y-model", ""})
if err != nil {
t.Fatalf("CreateKey: %v", err)
}
fresh, err := gw.VerifyKey(ctx, "model-key-1234")
if err != nil || len(fresh.Models) != 2 || fresh.Models[0] != "x-model" || fresh.Models[1] != "y-model" {
t.Fatalf("创建后读回 models = %v, %v", fresh.Models, err)
}
set := []string{"z-model"}
if err := gw.UpdateKey(ctx, key.ID, "", nil, nil, &set); err != nil {
t.Fatalf("UpdateKey 覆盖: %v", err)
}
fresh, _ = gw.VerifyKey(ctx, "model-key-1234")
if len(fresh.Models) != 1 || fresh.Models[0] != "z-model" {
t.Fatalf("覆盖后 models = %v", fresh.Models)
}
if err := gw.UpdateKey(ctx, key.ID, "renamed", nil, nil, nil); err != nil {
t.Fatalf("UpdateKey 不传 models: %v", err)
}
fresh, _ = gw.VerifyKey(ctx, "model-key-1234")
if len(fresh.Models) != 1 {
t.Fatalf("未传 models 却被改动: %v", fresh.Models)
}
empty := []string{}
if err := gw.UpdateKey(ctx, key.ID, "", nil, nil, &empty); err != nil {
t.Fatalf("UpdateKey 清空: %v", err)
}
fresh, _ = gw.VerifyKey(ctx, "model-key-1234")
if len(fresh.Models) != 0 {
t.Fatalf("清空后 models = %v, want 空", fresh.Models)
}
}
// TestRespPassthroughStreamSwitchesChannel 断言流式直通建立失败按 switchable 换渠道。
func TestRespPassthroughStreamSwitchesChannel(t *testing.T) {
client := &gatewayStubClient{
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
passPayload: []byte("data: {\"type\":\"response.completed\"}\n\n"),
passErrs: []error{stubServiceError{status: 503}},
}
gw, svc := newTestGateway(t, client)
cfg := importAliveConfig(t, svc)
seedChannel(t, gw, cfg.ID, "eu-frankfurt-1", 1, 1)
seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1)
stream, meta, err := gw.RespPassthroughStream(context.Background(), []byte(`{}`), "meta.llama-3.3-70b-instruct", "")
if err != nil || stream == nil {
t.Fatalf("RespPassthroughStream = %v", err)
}
defer stream.Close()
if meta.Retries != 1 || client.passCalls != 2 {
t.Errorf("应换渠道重试一次: retries=%d calls=%d", meta.Retries, client.passCalls)
}
payload, _ := io.ReadAll(stream)
if !bytes.Contains(payload, []byte("response.completed")) {
t.Errorf("流内容未透传: %s", payload)
}
}
+61 -206
View File
@@ -8,38 +8,8 @@ import (
"oci-portal/internal/aiwire"
)
// ResponsesToIR 把 Responses API 请求转换为网关 IR(OpenAI Chat 形态)。
// 有状态特性与内置工具不支持,直接报错(API 层 400)。
func ResponsesToIR(req aiwire.RespRequest) (aiwire.ChatRequest, error) {
if err := respRejectUnsupported(req); err != nil {
return aiwire.ChatRequest{}, err
}
ir := aiwire.ChatRequest{
Model: req.Model,
MaxTokens: req.MaxOutputTokens,
Temperature: req.Temperature,
TopP: req.TopP,
Stream: req.Stream,
ToolChoice: respToolChoice(req.ToolChoice),
}
if req.Reasoning != nil {
ir.ReasoningEffort = req.Reasoning.Effort
}
if req.Instructions != "" {
ir.Messages = append(ir.Messages, aiwire.ChatMessage{Role: "system", Content: aiwire.NewTextContent(req.Instructions)})
}
msgs, err := respInputToMessages(req.Input)
if err != nil {
return aiwire.ChatRequest{}, err
}
ir.Messages = append(ir.Messages, msgs...)
ir.Tools = respTools(req.Tools)
ir.ResponseFormat = respFormat(req.Text)
return ir, nil
}
// respRejectUnsupported 拒绝无状态网关无法承接的请求特性。
func respRejectUnsupported(req aiwire.RespRequest) error {
// respRejectStateful 拒绝有状态特性(网关无状态)。
func respRejectStateful(req aiwire.RespRequest) error {
if req.PreviousResponseID != "" {
return fmt.Errorf("previous_response_id 不支持:网关不保存历史响应,请在 input 中自带完整上下文(store:false 模式)")
}
@@ -49,195 +19,80 @@ func respRejectUnsupported(req aiwire.RespRequest) error {
if req.Background != nil && *req.Background {
return fmt.Errorf("background 模式不支持")
}
return nil
}
// RespServerTools 报告工具列表是否含 xAI 服务端工具(web_search / x_search)。
func RespServerTools(tools []aiwire.RespTool) bool {
for _, t := range tools {
if t.Type == "web_search" || t.Type == "x_search" {
return true
}
}
return false
}
// RespPassthroughValidate 校验直通请求:模型必填,有状态特性不支持,工具类型
// 只放行 function 与已实测的 web_search / x_search;流式仅在含服务端工具时拒绝
// (工具流式事件形态未实测,不放开)。
func RespPassthroughValidate(req aiwire.RespRequest) error {
if strings.TrimSpace(req.Model) == "" {
return fmt.Errorf("model 不能为空")
}
if req.Stream && RespServerTools(req.Tools) {
return fmt.Errorf("服务端工具暂不支持流式:请去掉 stream 或改用 function 工具")
}
if err := respRejectStateful(req); err != nil {
return err
}
for _, t := range req.Tools {
if t.Type != "function" {
return fmt.Errorf("不支持的工具类型 %q:仅支持 function 工具", t.Type)
switch t.Type {
case "function", "web_search", "x_search":
default:
return fmt.Errorf("不支持的工具类型 %q:服务端工具仅支持 web_search / x_search", t.Type)
}
}
return nil
}
// respInputToMessages 把 input(string 或 item 数组)展开为 IR 消息序列。
func respInputToMessages(in aiwire.RespInput) ([]aiwire.ChatMessage, error) {
if !in.IsArray {
if strings.TrimSpace(in.Text) == "" {
return nil, fmt.Errorf("input 不能为空")
}
return []aiwire.ChatMessage{{Role: "user", Content: aiwire.NewTextContent(in.Text)}}, nil
// RespPassthroughBody 以原始请求体为基构造上游 body:强制 store:false(禁上游
// 存态),stream 原样保留(流式直通);用 json.Number 保真未知字段与数值。
func RespPassthroughBody(raw []byte) ([]byte, error) {
dec := json.NewDecoder(strings.NewReader(string(raw)))
dec.UseNumber()
var body map[string]any
if err := dec.Decode(&body); err != nil {
return nil, fmt.Errorf("解析请求体: %w", err)
}
var msgs []aiwire.ChatMessage
for _, it := range in.Items {
out, err := respItemToMessages(it, msgs)
if err != nil {
return nil, err
}
msgs = out
}
return msgs, nil
body["store"] = false
return json.Marshal(body)
}
// respItemToMessages 追加一个 input item;连续 function_call 合并进同一 assistant 消息
func respItemToMessages(it aiwire.RespItem, msgs []aiwire.ChatMessage) ([]aiwire.ChatMessage, error) {
switch it.Type {
case "", "message":
if it.Content.HasUnsupported() {
return nil, ErrAiUnsupportedBlock
}
role := respRole(it.Role)
return append(msgs, aiwire.ChatMessage{Role: role, Content: respContentToIR(it.Content)}), nil
case "function_call":
call := aiwire.ToolCall{ID: it.CallID, Type: "function",
Function: aiwire.FunctionCall{Name: it.Name, Arguments: it.Arguments}}
if n := len(msgs); n > 0 && msgs[n-1].Role == "assistant" && len(msgs[n-1].ToolCalls) > 0 {
msgs[n-1].ToolCalls = append(msgs[n-1].ToolCalls, call)
return msgs, nil
}
return append(msgs, aiwire.ChatMessage{Role: "assistant", ToolCalls: []aiwire.ToolCall{call}}), nil
case "function_call_output":
return append(msgs, aiwire.ChatMessage{Role: "tool", ToolCallID: it.CallID,
Content: aiwire.NewTextContent(it.OutputText())}), nil
case "reasoning":
return msgs, nil // 推理块不回灌上游
default:
return nil, fmt.Errorf("不支持的 input item 类型 %q", it.Type)
// RespPassthroughUsage 从直通响应提取用量;缺失时返回 nil(日志记零)
func RespPassthroughUsage(payload []byte) *aiwire.Usage {
var root struct {
Usage *aiwire.RespUsage `json:"usage"`
}
}
// respRole 归一 item 角色;developer 视为 system,缺省 user。
func respRole(role string) string {
switch role {
case "system", "developer":
return "system"
case "assistant":
return "assistant"
default:
return "user"
}
}
// respContentToIR 把 Responses 内容转 IR;纯文本拍平,含图片时保留部件顺序。
func respContentToIR(c aiwire.RespContent) aiwire.Content {
var parts []aiwire.ContentPart
hasImage := false
for _, p := range c.Parts {
switch p.Type {
case "input_image":
parts = append(parts, aiwire.ContentPart{Type: "image_url",
ImageURL: &aiwire.ImageURL{URL: p.ImageURL, Detail: p.Detail}})
hasImage = true
default:
if p.Text != "" {
parts = append(parts, aiwire.ContentPart{Type: "text", Text: p.Text})
}
}
}
if !c.IsArray || !hasImage {
return aiwire.NewTextContent(c.JoinText())
}
return aiwire.NewPartsContent(parts)
}
// respTools 把扁平工具定义还原为嵌套 Chat 形态;strict 无 OCI 对应,忽略。
func respTools(tools []aiwire.RespTool) []aiwire.Tool {
if len(tools) == 0 {
if json.Unmarshal(payload, &root) != nil || root.Usage == nil {
return nil
}
out := make([]aiwire.Tool, 0, len(tools))
for _, t := range tools {
out = append(out, aiwire.Tool{Type: "function",
Function: aiwire.FunctionDef{Name: t.Name, Description: t.Description, Parameters: t.Parameters}})
usage := &aiwire.Usage{PromptTokens: root.Usage.InputTokens,
CompletionTokens: root.Usage.OutputTokens, TotalTokens: root.Usage.TotalTokens}
if cached := root.Usage.InputTokensDetails.CachedTokens; cached > 0 {
usage.PromptTokensDetails = &aiwire.PromptTokensDetails{CachedTokens: cached}
}
return out
return usage
}
// respToolChoice 转换 tool_choice:字符串直通;{type:function,name} 转嵌套;
// allowed_tools 等对象形态降级 auto
func respToolChoice(raw json.RawMessage) json.RawMessage {
if len(raw) == 0 {
// RespStreamCompletedUsage 从一行 SSE data JSON 中提取 response.completed 事件的
// usage;非 completed 事件或解析失败返回 nil。流式直通逐行喂入,最后一次非 nil 生效
func RespStreamCompletedUsage(data []byte) *aiwire.Usage {
var ev struct {
Type string `json:"type"`
Response json.RawMessage `json:"response"`
}
if json.Unmarshal(data, &ev) != nil || ev.Type != "response.completed" || len(ev.Response) == 0 {
return nil
}
t := strings.TrimSpace(string(raw))
if strings.HasPrefix(t, "\"") {
return raw
}
var obj struct {
Type string `json:"type"`
Name string `json:"name"`
}
if err := json.Unmarshal(raw, &obj); err != nil || obj.Type != "function" || obj.Name == "" {
return json.RawMessage(`"auto"`)
}
out, _ := json.Marshal(map[string]any{"type": "function", "function": map[string]string{"name": obj.Name}})
return out
}
// respFormat 把 text.format 转成 response_format;text 形态无需显式指定。
func respFormat(text *aiwire.RespText) *aiwire.ResponseFormat {
if text == nil || text.Format == nil {
return nil
}
switch text.Format.Type {
case "json_object":
return &aiwire.ResponseFormat{Type: "json_object"}
case "json_schema":
spec, _ := json.Marshal(map[string]any{
"name": text.Format.Name, "schema": json.RawMessage(text.Format.Schema), "strict": text.Format.Strict,
})
return &aiwire.ResponseFormat{Type: "json_schema", JSONSchema: spec}
default:
return nil
}
}
// IRRespToResponses 把 IR 非流式响应装配为 Response 对象。
func IRRespToResponses(resp *aiwire.ChatResponse, id string, created int64) aiwire.Response {
out := aiwire.Response{ID: id, Object: "response", CreatedAt: created,
Status: "completed", Model: resp.Model, Output: []aiwire.RespOutItem{}, Store: false}
if len(resp.Choices) == 0 {
return out
}
choice := resp.Choices[0]
if text := choice.Message.Content.JoinText(); text != "" {
out.Output = append(out.Output, respMessageItem(id, 0, text))
}
for i, tc := range choice.Message.ToolCalls {
out.Output = append(out.Output, aiwire.RespOutItem{Type: "function_call",
ID: respItemID(id, "fc", len(out.Output)+i), Status: "completed",
CallID: tc.ID, Name: tc.Function.Name, Arguments: respArgs(tc.Function.Arguments)})
}
if choice.FinishReason == "length" {
out.Status = "incomplete"
out.IncompleteDetails = &aiwire.RespIncomplete{Reason: "max_output_tokens"}
}
out.Usage = respUsage(resp.Usage)
return out
}
func respMessageItem(respID string, idx int, text string) aiwire.RespOutItem {
return aiwire.RespOutItem{Type: "message", ID: respItemID(respID, "msg", idx),
Status: "completed", Role: "assistant",
Content: []aiwire.RespOutPart{{Type: "output_text", Text: text, Annotations: []any{}}}}
}
// respItemID 从响应 ID 派生确定性 item ID。
func respItemID(respID, kind string, idx int) string {
return fmt.Sprintf("%s_%s_%d", kind, strings.TrimPrefix(respID, "resp_"), idx)
}
// respArgs 保证 arguments 是合法 JSON 字符串(空实参回退 {})。
func respArgs(args string) string {
if strings.TrimSpace(args) == "" {
return "{}"
}
return args
}
// respUsage 把 IR 用量改写为 Responses 命名口径。
func respUsage(u *aiwire.Usage) *aiwire.RespUsage {
if u == nil {
return nil
}
return &aiwire.RespUsage{InputTokens: u.PromptTokens, OutputTokens: u.CompletionTokens,
TotalTokens: u.TotalTokens,
InputTokensDetails: aiwire.RespInDetails{CachedTokens: u.CachedTokens()}}
return RespPassthroughUsage(ev.Response)
}
-199
View File
@@ -1,199 +0,0 @@
package service
import (
"strings"
"oci-portal/internal/aiwire"
)
// RespEvent 是一条 Responses SSE 语义事件(event 名 + data 载荷)。
type RespEvent struct {
Event string
Data map[string]any
}
// RespStream 把 IR chunk 流聚合为 Responses 语义事件序列。
// 与 AnthStream 同构,但终态事件须携带完整 Response 快照,故全程缓冲文本与实参。
type RespStream struct {
id string
model string
created int64
seq int
started bool
items []aiwire.RespOutItem
msgOpen bool
text strings.Builder
tool aiwire.RespOutItem
toolIdx int
tOpen bool
args strings.Builder
finish string
usage *aiwire.Usage
}
// NewRespStream 构造状态机;id 形如 resp_*,created 为响应时间戳。
func NewRespStream(id, model string, created int64) *RespStream {
return &RespStream{id: id, model: model, created: created}
}
// Feed 消化一个上游 chunk,返回应立即下发的事件。
func (s *RespStream) Feed(chunk aiwire.ChatChunk) []RespEvent {
var evs []RespEvent
if !s.started {
s.started = true
evs = append(evs, s.respEvent("response.created", "in_progress"),
s.respEvent("response.in_progress", "in_progress"))
}
if chunk.Usage != nil {
s.usage = chunk.Usage
}
if len(chunk.Choices) == 0 {
return evs
}
choice := chunk.Choices[0]
if choice.Delta.Content != "" {
evs = append(evs, s.feedText(choice.Delta.Content)...)
}
for _, tc := range choice.Delta.ToolCalls {
evs = append(evs, s.feedTool(tc)...)
}
if choice.FinishReason != nil {
s.finish = *choice.FinishReason
}
return evs
}
// Finish 关闭未闭合的块并产出终态事件(带完整 Response 快照)。
func (s *RespStream) Finish() []RespEvent {
var evs []RespEvent
if !s.started {
s.started = true
evs = append(evs, s.respEvent("response.created", "in_progress"))
}
evs = append(evs, s.closeMsg()...)
evs = append(evs, s.closeTool()...)
status := "completed"
if s.finish == "length" {
status = "incomplete"
}
return append(evs, s.respEvent("response."+status, status))
}
// Usage 返回聚合到的用量(可能为 nil),供调用日志。
func (s *RespStream) Usage() *aiwire.Usage { return s.usage }
// ev 构造带自增 sequence_number 的事件。
func (s *RespStream) ev(typ string, kv map[string]any) RespEvent {
s.seq++
data := map[string]any{"type": typ, "sequence_number": s.seq}
for k, v := range kv {
data[k] = v
}
return RespEvent{Event: typ, Data: data}
}
// respEvent 构造携带 Response 快照的生命周期事件。
func (s *RespStream) respEvent(typ, status string) RespEvent {
return s.ev(typ, map[string]any{"response": s.snapshot(status)})
}
// snapshot 组装当前累计状态的 Response 对象。
func (s *RespStream) snapshot(status string) aiwire.Response {
resp := aiwire.Response{ID: s.id, Object: "response", CreatedAt: s.created,
Status: status, Model: s.model, Store: false,
Output: append([]aiwire.RespOutItem{}, s.items...)}
resp.Usage = respUsage(s.usage)
if status == "incomplete" {
resp.IncompleteDetails = &aiwire.RespIncomplete{Reason: "max_output_tokens"}
}
return resp
}
// feedText 处理文本增量:必要时开 message item 与 content part。
func (s *RespStream) feedText(delta string) []RespEvent {
evs := s.closeTool()
if !s.msgOpen {
s.msgOpen = true
s.text.Reset()
item := aiwire.RespOutItem{Type: "message", ID: s.itemID("msg"), Status: "in_progress",
Role: "assistant", Content: []aiwire.RespOutPart{}}
evs = append(evs, s.ev("response.output_item.added", map[string]any{
"output_index": len(s.items), "item": item}))
evs = append(evs, s.ev("response.content_part.added", map[string]any{
"item_id": item.ID, "output_index": len(s.items), "content_index": 0,
"part": aiwire.RespOutPart{Type: "output_text", Text: "", Annotations: []any{}}}))
}
s.text.WriteString(delta)
return append(evs, s.ev("response.output_text.delta", map[string]any{
"item_id": s.itemID("msg"), "output_index": len(s.items), "content_index": 0, "delta": delta}))
}
// closeMsg 闭合当前 message item(text done → part done → item done)。
func (s *RespStream) closeMsg() []RespEvent {
if !s.msgOpen {
return nil
}
s.msgOpen = false
id, idx, full := s.itemID("msg"), len(s.items), s.text.String()
part := aiwire.RespOutPart{Type: "output_text", Text: full, Annotations: []any{}}
item := aiwire.RespOutItem{Type: "message", ID: id, Status: "completed",
Role: "assistant", Content: []aiwire.RespOutPart{part}}
evs := []RespEvent{
s.ev("response.output_text.done", map[string]any{"item_id": id, "output_index": idx, "content_index": 0, "text": full}),
s.ev("response.content_part.done", map[string]any{"item_id": id, "output_index": idx, "content_index": 0, "part": part}),
s.ev("response.output_item.done", map[string]any{"output_index": idx, "item": item}),
}
s.items = append(s.items, item)
return evs
}
// feedTool 处理工具调用增量:新 Index 先闭合旧调用再开新 item。
func (s *RespStream) feedTool(tc aiwire.ToolCallDelta) []RespEvent {
evs := s.closeMsg()
if s.tOpen && tc.Index != s.toolIdx {
evs = append(evs, s.closeTool()...)
}
if !s.tOpen {
s.tOpen, s.toolIdx = true, tc.Index
s.args.Reset()
s.tool = aiwire.RespOutItem{Type: "function_call", ID: s.itemID("fc"), Status: "in_progress",
CallID: tc.ID, Name: tc.Function.Name}
evs = append(evs, s.ev("response.output_item.added", map[string]any{
"output_index": len(s.items), "item": s.tool}))
}
if tc.ID != "" && s.tool.CallID == "" {
s.tool.CallID = tc.ID
}
if tc.Function.Name != "" && s.tool.Name == "" {
s.tool.Name = tc.Function.Name
}
if tc.Function.Arguments == "" {
return evs
}
s.args.WriteString(tc.Function.Arguments)
return append(evs, s.ev("response.function_call_arguments.delta", map[string]any{
"item_id": s.tool.ID, "output_index": len(s.items), "delta": tc.Function.Arguments}))
}
// closeTool 闭合当前 function_call item(arguments done → item done)。
func (s *RespStream) closeTool() []RespEvent {
if !s.tOpen {
return nil
}
s.tOpen = false
s.tool.Arguments = respArgs(s.args.String())
s.tool.Status = "completed"
idx := len(s.items)
evs := []RespEvent{
s.ev("response.function_call_arguments.done", map[string]any{
"item_id": s.tool.ID, "output_index": idx, "arguments": s.tool.Arguments}),
s.ev("response.output_item.done", map[string]any{"output_index": idx, "item": s.tool}),
}
s.items = append(s.items, s.tool)
return evs
}
// itemID 按当前 output 序号派生确定性 item ID。
func (s *RespStream) itemID(kind string) string {
return respItemID(s.id, kind, len(s.items))
}
+106 -135
View File
@@ -17,47 +17,6 @@ func respReq(t *testing.T, raw string) aiwire.RespRequest {
return req
}
func TestResponsesToIR(t *testing.T) {
req := respReq(t, `{
"model": "meta.llama-3.3-70b-instruct",
"instructions": "你是助手",
"max_output_tokens": 128,
"input": [
{"role": "user", "content": "查天气"},
{"type": "function_call", "call_id": "call_1", "name": "get_weather", "arguments": "{\"city\":\"上海\"}"},
{"type": "function_call_output", "call_id": "call_1", "output": "{\"temp\":31}"},
{"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "31 度"}]}
],
"tools": [{"type": "function", "name": "get_weather", "description": "查天气", "parameters": {"type": "object"}, "strict": true}],
"tool_choice": {"type": "function", "name": "get_weather"},
"text": {"format": {"type": "json_schema", "name": "out", "schema": {"type": "object"}}}
}`)
ir, err := ResponsesToIR(req)
if err != nil {
t.Fatalf("ResponsesToIR: %v", err)
}
roles := make([]string, 0, len(ir.Messages))
for _, m := range ir.Messages {
roles = append(roles, m.Role)
}
want := []string{"system", "user", "assistant", "tool", "assistant"}
if strings.Join(roles, ",") != strings.Join(want, ",") {
t.Errorf("roles = %v, want %v", roles, want)
}
if ir.Messages[2].ToolCalls[0].ID != "call_1" || ir.Messages[3].ToolCallID != "call_1" {
t.Errorf("tool call 链接错误: %+v", ir.Messages)
}
if deref(ir.MaxTokens) != 128 || len(ir.Tools) != 1 || ir.Tools[0].Function.Name != "get_weather" {
t.Errorf("参数映射错误: max=%v tools=%+v", ir.MaxTokens, ir.Tools)
}
if !strings.Contains(string(ir.ToolChoice), `"function"`) || !strings.Contains(string(ir.ToolChoice), "get_weather") {
t.Errorf("tool_choice = %s", ir.ToolChoice)
}
if ir.ResponseFormat == nil || ir.ResponseFormat.Type != "json_schema" {
t.Errorf("response_format = %+v", ir.ResponseFormat)
}
}
func deref(p *int) int {
if p == nil {
return 0
@@ -65,112 +24,124 @@ func deref(p *int) int {
return *p
}
func TestResponsesToIRRejects(t *testing.T) {
cases := []struct{ name, raw string }{
{"previous_response_id", `{"model":"m","input":"hi","previous_response_id":"resp_x"}`},
{"conversation", `{"model":"m","input":"hi","conversation":{"id":"conv_1"}}`},
{"background", `{"model":"m","input":"hi","background":true}`},
{"builtin tool", `{"model":"m","input":"hi","tools":[{"type":"web_search"}]}`},
{"file_id image", `{"model":"m","input":[{"role":"user","content":[{"type":"input_image","file_id":"file_1"}]}]}`},
{"unknown item", `{"model":"m","input":[{"type":"item_reference","id":"x"}]}`},
func TestRespServerTools(t *testing.T) {
tests := []struct {
name string
tools []aiwire.RespTool
want bool
}{
{"web_search", []aiwire.RespTool{{Type: "web_search"}}, true},
{"x_search混用", []aiwire.RespTool{{Type: "function", Name: "f"}, {Type: "x_search"}}, true},
{"仅function", []aiwire.RespTool{{Type: "function", Name: "f"}}, false},
{"空", nil, false},
{"其他类型", []aiwire.RespTool{{Type: "code_interpreter"}}, false},
}
for _, c := range cases {
if _, err := ResponsesToIR(respReq(t, c.raw)); err == nil {
t.Errorf("%s 应被拒绝", c.name)
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if got := RespServerTools(test.tools); got != test.want {
t.Fatalf("got %v, want %v", got, test.want)
}
})
}
// input 字符串形态 + store/reasoning 忽略项不报错
req := respReq(t, `{"model":"m","input":"你好","store":false,"reasoning":{"effort":"low"},"parallel_tool_calls":true}`)
ir, err := ResponsesToIR(req)
if err != nil || len(ir.Messages) != 1 || ir.Messages[0].Role != "user" || ir.ReasoningEffort != "low" {
t.Errorf("字符串 input = %+v, %v", ir.Messages, err)
}
func TestRespPassthroughValidate(t *testing.T) {
prev := "resp_1"
bg := true
tests := []struct {
name string
req aiwire.RespRequest
wantErr bool
}{
{"合法", aiwire.RespRequest{Model: "m", Tools: []aiwire.RespTool{{Type: "web_search"}}}, false},
{"混用function", aiwire.RespRequest{Model: "m", Tools: []aiwire.RespTool{{Type: "web_search"}, {Type: "function", Name: "f"}}}, false},
{"缺model", aiwire.RespRequest{Tools: []aiwire.RespTool{{Type: "web_search"}}}, true},
{"工具加流式拒绝", aiwire.RespRequest{Model: "m", Stream: true, Tools: []aiwire.RespTool{{Type: "web_search"}}}, true},
{"无工具流式放行", aiwire.RespRequest{Model: "m", Stream: true}, false},
{"function工具流式放行", aiwire.RespRequest{Model: "m", Stream: true, Tools: []aiwire.RespTool{{Type: "function", Name: "f"}}}, false},
{"有状态拒绝", aiwire.RespRequest{Model: "m", PreviousResponseID: prev}, true},
{"background拒绝", aiwire.RespRequest{Model: "m", Background: &bg}, true},
{"未知工具拒绝", aiwire.RespRequest{Model: "m", Tools: []aiwire.RespTool{{Type: "web_search"}, {Type: "mcp"}}}, true},
}
// input_image(url 形态)放行并保留图文顺序
req2 := respReq(t, `{"model":"m","input":[{"role":"user","content":[{"type":"input_text","text":"看图"},{"type":"input_image","image_url":"https://x/1.png","detail":"low"}]}]}`)
ir2, err := ResponsesToIR(req2)
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
err := RespPassthroughValidate(test.req)
if (err != nil) != test.wantErr {
t.Fatalf("err = %v, wantErr %v", err, test.wantErr)
}
})
}
}
func TestRespPassthroughBody(t *testing.T) {
raw := []byte(`{"model":"m","input":"hi","stream":true,"store":true,"max_output_tokens":128,"custom_field":{"a":1.5}}`)
out, err := RespPassthroughBody(raw)
if err != nil {
t.Fatalf("input_image 应放行: %v", err)
t.Fatalf("RespPassthroughBody: %v", err)
}
parts := ir2.Messages[0].Content.Parts
if len(parts) != 2 || parts[1].Type != "image_url" || parts[1].ImageURL == nil || parts[1].ImageURL.URL != "https://x/1.png" {
t.Errorf("parts = %+v", parts)
var body map[string]any
if err := json.Unmarshal(out, &body); err != nil {
t.Fatalf("unmarshal out: %v", err)
}
if body["store"] != false {
t.Errorf("store 应强制 false, got %v", body["store"])
}
if body["stream"] != true {
t.Errorf("stream 应原样保留, got %v", body["stream"])
}
if string(out) == "" || !strings.Contains(string(out), `"max_output_tokens":128`) {
t.Errorf("数值字段应保真: %s", out)
}
if !strings.Contains(string(out), `"custom_field"`) {
t.Errorf("未知字段应保留: %s", out)
}
if _, err := RespPassthroughBody([]byte("not-json")); err == nil {
t.Error("非法 JSON 应报错")
}
}
func TestIRRespToResponses(t *testing.T) {
resp := &aiwire.ChatResponse{Model: "m", Choices: []aiwire.Choice{{
Message: aiwire.ChatMessage{Role: "assistant", Content: aiwire.NewTextContent("你好"),
ToolCalls: []aiwire.ToolCall{{ID: "call_9", Type: "function",
Function: aiwire.FunctionCall{Name: "f", Arguments: ""}}}},
FinishReason: "length",
}}, Usage: &aiwire.Usage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15}}
out := IRRespToResponses(resp, "resp_abc", 1751966400)
if out.Status != "incomplete" || out.IncompleteDetails == nil || out.IncompleteDetails.Reason != "max_output_tokens" {
t.Errorf("length 应映射 incomplete: %+v", out)
func TestRespPassthroughUsage(t *testing.T) {
tests := []struct {
name string
payload string
wantNil bool
prompt, cached int
}{
{"完整", `{"usage":{"input_tokens":9,"output_tokens":3,"total_tokens":12,"input_tokens_details":{"cached_tokens":5}}}`, false, 9, 5},
{"无细分", `{"usage":{"input_tokens":9,"output_tokens":3,"total_tokens":12}}`, false, 9, 0},
{"无usage", `{"id":"resp_1"}`, true, 0, 0},
{"非法JSON", `xx`, true, 0, 0},
}
if len(out.Output) != 2 || out.Output[0].Type != "message" || out.Output[1].Type != "function_call" {
t.Fatalf("output = %+v", out.Output)
}
if out.Output[0].Content[0].Text != "你好" || out.Output[1].CallID != "call_9" || out.Output[1].Arguments != "{}" {
t.Errorf("item 装配错误: %+v", out.Output)
}
if out.Usage.InputTokens != 10 || out.Usage.OutputTokens != 5 || out.Store {
t.Errorf("usage/store 错误: %+v", out)
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
usage := RespPassthroughUsage([]byte(test.payload))
if (usage == nil) != test.wantNil {
t.Fatalf("usage = %v, wantNil %v", usage, test.wantNil)
}
if usage == nil {
return
}
if usage.PromptTokens != test.prompt || usage.CachedTokens() != test.cached {
t.Fatalf("prompt=%d cached=%d, want %d/%d", usage.PromptTokens, usage.CachedTokens(), test.prompt, test.cached)
}
})
}
}
func chunkText(text string) aiwire.ChatChunk {
return aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{Content: text}}}}
}
func TestRespStreamTextAndTool(t *testing.T) {
st := NewRespStream("resp_x", "m", 1751966400)
var types []string
collect := func(evs []RespEvent) {
for _, ev := range evs {
types = append(types, ev.Event)
// TestRespStreamCompletedUsage 断言流式 usage 只从 completed 事件提取。
func TestRespStreamCompletedUsage(t *testing.T) {
completed := []byte(`{"type":"response.completed","response":{"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15,"input_tokens_details":{"cached_tokens":4}}}}`)
if u := RespStreamCompletedUsage(completed); u == nil || u.PromptTokens != 10 || u.CompletionTokens != 5 ||
u.PromptTokensDetails == nil || u.PromptTokensDetails.CachedTokens != 4 {
t.Fatalf("completed 事件 usage 解析失败: %+v", u)
}
for _, data := range []string{
`{"type":"response.output_text.delta","delta":"hi"}`,
`{"type":"response.completed"}`,
`not-json`,
} {
if u := RespStreamCompletedUsage([]byte(data)); u != nil {
t.Fatalf("非 completed 或缺 response 应返回 nil: %s", data)
}
}
collect(st.Feed(chunkText("你")))
collect(st.Feed(chunkText("好")))
collect(st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{
ToolCalls: []aiwire.ToolCallDelta{{Index: 0, ID: "call_1", Function: aiwire.FunctionCallDelta{Name: "f", Arguments: "{\"a\""}}}}}}}))
collect(st.Feed(aiwire.ChatChunk{Choices: []aiwire.ChunkChoice{{Delta: aiwire.Delta{
ToolCalls: []aiwire.ToolCallDelta{{Index: 0, Function: aiwire.FunctionCallDelta{Arguments: ":1}"}}}}}},
Usage: &aiwire.Usage{PromptTokens: 3, CompletionTokens: 7, TotalTokens: 10}}))
finish := st.Finish()
collect(finish)
want := []string{
"response.created", "response.in_progress",
"response.output_item.added", "response.content_part.added", "response.output_text.delta",
"response.output_text.delta",
"response.output_text.done", "response.content_part.done", "response.output_item.done",
"response.output_item.added", "response.function_call_arguments.delta",
"response.function_call_arguments.delta",
"response.function_call_arguments.done", "response.output_item.done",
"response.completed",
}
if strings.Join(types, "\n") != strings.Join(want, "\n") {
t.Errorf("事件序列 =\n%s\nwant\n%s", strings.Join(types, "\n"), strings.Join(want, "\n"))
}
last := finish[len(finish)-1].Data
resp := last["response"].(aiwire.Response)
if len(resp.Output) != 2 || resp.Output[0].Content[0].Text != "你好" || resp.Output[1].Arguments != "{\"a\":1}" {
t.Errorf("终态快照 = %+v", resp.Output)
}
if resp.Usage == nil || resp.Usage.TotalTokens != 10 {
t.Errorf("终态 usage = %+v", resp.Usage)
}
// sequence_number 严格递增
if seq, ok := last["sequence_number"].(int); !ok || seq != len(want) {
t.Errorf("最终 sequence_number = %v, want %d", last["sequence_number"], len(want))
}
}
func TestRespStreamEmpty(t *testing.T) {
st := NewRespStream("resp_e", "m", 1)
evs := st.Finish()
if len(evs) != 2 || evs[0].Event != "response.created" || evs[1].Event != "response.completed" {
t.Errorf("空流事件 = %+v", evs)
}
}
+411
View File
@@ -0,0 +1,411 @@
package service
import (
"encoding/json"
"fmt"
"strings"
"oci-portal/internal/aiwire"
)
// Anthropic Messages ↔ OCI OpenAI 兼容面(/actions/v1/responses)直通转换。
// typed chat 面剔除后 Messages 入口的唯一上游通路。语义损失(README 已披露):
// stop_sequences / top_k / metadata / thinking 无对应字段,忽略;上游 reasoning
// 输出项与增量事件丢弃(Anthropic thinking 块含签名语义,不伪造)。
// ErrAiUnsupportedBlock 表示请求含网关无法承接的内容块(仅支持文本/图片/工具块)。
var ErrAiUnsupportedBlock = fmt.Errorf("暂不支持文本与图片以外的内容块")
// AnthropicToResponsesBody 把 Messages 请求转为直通 body(强制 store:false)。
func AnthropicToResponsesBody(req aiwire.MessagesRequest) ([]byte, error) {
input, err := anthInputItems(req.Messages)
if err != nil {
return nil, err
}
body := map[string]any{"model": req.Model, "max_output_tokens": req.MaxTokens,
"input": input, "store": false}
if sys := req.SystemText(); sys != "" {
body["instructions"] = sys
}
if req.Temperature != nil {
body["temperature"] = *req.Temperature
}
if req.TopP != nil {
body["top_p"] = *req.TopP
}
if tools := anthRespTools(req.Tools); tools != nil {
body["tools"] = tools
}
if tc := anthRespToolChoice(req.ToolChoice); tc != nil {
body["tool_choice"] = tc
}
if req.OutputConfig != nil && req.OutputConfig.Effort != "" {
body["reasoning"] = map[string]string{"effort": strings.ToLower(req.OutputConfig.Effort)}
}
if req.Stream {
body["stream"] = true
}
return json.Marshal(body)
}
// anthInputItems 把消息序列展开为 Responses input 项:tool_use / tool_result 为
// 独立 function_call / function_call_output 项,其余聚合为 message 项;遇独立项时
// 先冲刷已聚合部件,保持块间相对顺序。
func anthInputItems(messages []aiwire.AnthMessage) ([]any, error) {
var items []any
for _, m := range messages {
var parts []map[string]any
flush := func() {
if len(parts) > 0 {
items = append(items, map[string]any{"role": m.Role, "content": parts})
parts = nil
}
}
for _, b := range m.Content.AllBlocks() {
item, part, err := anthBlockItem(m.Role, b)
if err != nil {
return nil, err
}
if part != nil {
parts = append(parts, part)
}
if item != nil {
flush()
items = append(items, item)
}
}
flush()
}
return items, nil
}
// anthBlockItem 把单个内容块转为独立项或消息部件(thinking 忽略,未知块拒绝)。
func anthBlockItem(role string, b aiwire.AnthBlock) (item, part map[string]any, err error) {
switch b.Type {
case "text":
return nil, anthTextPart(role, b.Text), nil
case "image":
part, err = anthImageInput(b.Source)
return nil, part, err
case "tool_use":
return map[string]any{"type": "function_call", "call_id": b.ID,
"name": b.Name, "arguments": string(b.Input)}, nil, nil
case "tool_result":
return map[string]any{"type": "function_call_output",
"call_id": b.ToolUseID, "output": b.ResultText()}, nil, nil
case "thinking", "redacted_thinking":
return nil, nil, nil
default:
return nil, nil, ErrAiUnsupportedBlock
}
}
// anthTextPart 按角色选择 Responses 文本部件类型(assistant 历史为 output_text)。
func anthTextPart(role, text string) map[string]any {
if role == "assistant" {
return map[string]any{"type": "output_text", "text": text}
}
return map[string]any{"type": "input_text", "text": text}
}
// anthImageInput 把 Anthropic image source 转为 Responses input_image 部件。
func anthImageInput(source json.RawMessage) (map[string]any, error) {
var src struct {
Type string `json:"type"`
MediaType string `json:"media_type"`
Data string `json:"data"`
URL string `json:"url"`
}
if err := json.Unmarshal(source, &src); err != nil {
return nil, fmt.Errorf("image source 解析失败: %w", err)
}
switch src.Type {
case "base64":
if src.MediaType == "" || src.Data == "" {
return nil, fmt.Errorf("image source 缺少 media_type 或 data")
}
return map[string]any{"type": "input_image",
"image_url": "data:" + src.MediaType + ";base64," + src.Data}, nil
case "url":
if src.URL == "" {
return nil, fmt.Errorf("image source 缺少 url")
}
return map[string]any{"type": "input_image", "image_url": src.URL}, nil
}
return nil, fmt.Errorf("不支持的 image source 类型 %q", src.Type)
}
func anthRespTools(tools []aiwire.AnthTool) []map[string]any {
if len(tools) == 0 {
return nil
}
out := make([]map[string]any, 0, len(tools))
for _, t := range tools {
out = append(out, map[string]any{"type": "function", "name": t.Name,
"description": t.Description, "parameters": t.InputSchema})
}
return out
}
// anthRespToolChoice 映射 {type:auto|any|tool|none,name} → Responses 形态。
func anthRespToolChoice(raw json.RawMessage) any {
if len(raw) == 0 {
return nil
}
var tc struct {
Type string `json:"type"`
Name string `json:"name"`
}
if json.Unmarshal(raw, &tc) != nil {
return nil
}
switch tc.Type {
case "auto":
return "auto"
case "any":
return "required"
case "none":
return "none"
case "tool":
return map[string]string{"type": "function", "name": tc.Name}
}
return nil
}
// respPayload 是直通响应中本转换关心的子集(未知字段忽略)。
type respPayload struct {
Model string `json:"model"`
Status string `json:"status"`
IncompleteDetails *respIncomplete `json:"incomplete_details"`
Output []respOutputItem `json:"output"`
Usage *aiwire.RespUsage `json:"usage"`
Error *map[string]string `json:"error"`
}
type respIncomplete struct {
Reason string `json:"reason"`
}
type respOutputItem struct {
Type string `json:"type"`
CallID string `json:"call_id"`
Name string `json:"name"`
Arguments string `json:"arguments"`
Content []struct {
Type string `json:"type"`
Text string `json:"text"`
} `json:"content"`
}
// ResponsesToAnthropic 把直通非流式响应转为 Anthropic Messages 响应。
func ResponsesToAnthropic(payload []byte, msgID string) (*aiwire.MessagesResponse, error) {
var resp respPayload
if err := json.Unmarshal(payload, &resp); err != nil {
return nil, fmt.Errorf("解析上游响应: %w", err)
}
out := &aiwire.MessagesResponse{ID: msgID, Type: "message", Role: "assistant",
Model: resp.Model, Content: []aiwire.AnthBlock{}, StopReason: "end_turn"}
for _, item := range resp.Output {
switch item.Type {
case "message":
for _, part := range item.Content {
if part.Type == "output_text" && part.Text != "" {
out.Content = append(out.Content, aiwire.AnthBlock{Type: "text", Text: part.Text})
}
}
case "function_call":
out.Content = append(out.Content, aiwire.AnthBlock{Type: "tool_use", ID: item.CallID,
Name: item.Name, Input: argsToJSON(item.Arguments)})
out.StopReason = "tool_use"
}
}
if resp.Status == "incomplete" && resp.IncompleteDetails != nil &&
resp.IncompleteDetails.Reason == "max_output_tokens" {
out.StopReason = "max_tokens"
}
out.Usage = anthUsageFromResp(resp.Usage)
return out, nil
}
func anthUsageFromResp(u *aiwire.RespUsage) aiwire.AnthUsage {
if u == nil {
return aiwire.AnthUsage{}
}
return aiwire.AnthUsage{InputTokens: u.InputTokens, OutputTokens: u.OutputTokens,
CacheReadInputTokens: u.InputTokensDetails.CachedTokens}
}
// argsToJSON 保证 tool_use.input 是合法 JSON 对象(模型可能产出非法片段)。
func argsToJSON(args string) json.RawMessage {
trimmed := strings.TrimSpace(args)
if trimmed == "" {
return json.RawMessage(`{}`)
}
if json.Valid([]byte(trimmed)) {
return json.RawMessage(trimmed)
}
b, _ := json.Marshal(map[string]string{"_raw": args})
return b
}
// ---- Anthropic 流式桥:Responses SSE 事件 → Anthropic 事件序列 ----
// AnthEvent 是一条待写出的 Anthropic SSE 事件。
type AnthEvent struct {
Event string
Data any
}
// AnthRespBridge 把直通 SSE 事件流桥接为 Anthropic 事件序列:
// message_start → content_block_start/delta/stop(text 与 tool_use 分块)→ message_delta → message_stop。
// reasoning 系列事件丢弃。
type AnthRespBridge struct {
id, model string
started bool
blockOpen bool
blockIsTool bool
blockIndex int
stopReason string
usage aiwire.AnthUsage
}
// NewAnthRespBridge 构造桥;id 为响应消息 ID。
func NewAnthRespBridge(id, model string) *AnthRespBridge {
return &AnthRespBridge{id: id, model: model, blockIndex: -1, stopReason: "end_turn"}
}
// respStreamEvent 是直通 SSE data JSON 中桥关心的子集。
type respStreamEvent struct {
Type string `json:"type"`
Delta string `json:"delta"`
Item *respOutputItem `json:"item"`
Response *respPayload `json:"response"`
}
// Feed 消费一行 SSE data JSON,返回应立即写出的事件。
func (st *AnthRespBridge) Feed(data []byte) []AnthEvent {
var ev respStreamEvent
if json.Unmarshal(data, &ev) != nil {
return nil
}
var events []AnthEvent
if !st.started {
st.started = true
events = append(events, st.startEvent())
}
switch ev.Type {
case "response.output_item.added":
if ev.Item != nil && ev.Item.Type == "function_call" {
events = append(events, st.openBlock(true, ev.Item.CallID, ev.Item.Name)...)
st.stopReason = "tool_use"
}
case "response.output_text.delta":
events = append(events, st.textDelta(ev.Delta)...)
case "response.function_call_arguments.delta":
events = append(events, st.argsDelta(ev.Delta)...)
case "response.completed", "response.incomplete", "response.failed":
st.finishFrom(ev.Response)
}
return events
}
func (st *AnthRespBridge) textDelta(delta string) []AnthEvent {
var events []AnthEvent
if !st.blockOpen || st.blockIsTool {
events = append(events, st.openBlock(false, "", "")...)
}
events = append(events, AnthEvent{Event: "content_block_delta", Data: map[string]any{
"type": "content_block_delta", "index": st.blockIndex,
"delta": map[string]string{"type": "text_delta", "text": delta},
}})
return events
}
func (st *AnthRespBridge) argsDelta(delta string) []AnthEvent {
if !st.blockOpen || !st.blockIsTool {
return nil
}
return []AnthEvent{{Event: "content_block_delta", Data: map[string]any{
"type": "content_block_delta", "index": st.blockIndex,
"delta": map[string]string{"type": "input_json_delta", "partial_json": delta},
}}}
}
// finishFrom 记录终态:usage 与 stop_reason(max_output_tokens 截断 → max_tokens)。
func (st *AnthRespBridge) finishFrom(resp *respPayload) {
if resp == nil {
return
}
st.usage = anthUsageFromResp(resp.Usage)
if resp.Status == "incomplete" && resp.IncompleteDetails != nil &&
resp.IncompleteDetails.Reason == "max_output_tokens" {
st.stopReason = "max_tokens"
}
}
func (st *AnthRespBridge) startEvent() AnthEvent {
return AnthEvent{Event: "message_start", Data: map[string]any{
"type": "message_start",
"message": map[string]any{
"id": st.id, "type": "message", "role": "assistant", "model": st.model,
"content": []any{}, "stop_reason": nil,
"usage": map[string]int{"input_tokens": 0, "output_tokens": 0},
},
}}
}
// openBlock 关闭当前块并打开新块(text 或 tool_use)。
func (st *AnthRespBridge) openBlock(isTool bool, toolID, toolName string) []AnthEvent {
var events []AnthEvent
if st.blockOpen {
events = append(events, st.closeBlockEvent())
}
st.blockOpen, st.blockIsTool = true, isTool
st.blockIndex++
block := map[string]any{"type": "text", "text": ""}
if isTool {
block = map[string]any{"type": "tool_use", "id": toolID, "name": toolName, "input": map[string]any{}}
}
events = append(events, AnthEvent{Event: "content_block_start", Data: map[string]any{
"type": "content_block_start", "index": st.blockIndex, "content_block": block,
}})
return events
}
func (st *AnthRespBridge) closeBlockEvent() AnthEvent {
return AnthEvent{Event: "content_block_stop", Data: map[string]any{
"type": "content_block_stop", "index": st.blockIndex,
}}
}
// Finish 在上游流结束后收尾:关块 → message_delta(stop_reason+usage)→ message_stop。
func (st *AnthRespBridge) Finish() []AnthEvent {
var events []AnthEvent
if !st.started {
st.started = true
events = append(events, st.startEvent())
}
if st.blockOpen {
events = append(events, st.closeBlockEvent())
st.blockOpen = false
}
usage := map[string]int{"output_tokens": st.usage.OutputTokens}
if st.usage.InputTokens > 0 {
usage["input_tokens"] = st.usage.InputTokens
}
if st.usage.CacheReadInputTokens > 0 {
usage["cache_read_input_tokens"] = st.usage.CacheReadInputTokens
}
events = append(events,
AnthEvent{Event: "message_delta", Data: map[string]any{
"type": "message_delta",
"delta": map[string]any{"stop_reason": st.stopReason, "stop_sequence": nil},
"usage": usage,
}},
AnthEvent{Event: "message_stop", Data: map[string]any{"type": "message_stop"}},
)
return events
}
// Usage 返回聚合到的用量(供调用日志)。
func (st *AnthRespBridge) Usage() aiwire.AnthUsage { return st.usage }
+173
View File
@@ -0,0 +1,173 @@
package service
import (
"encoding/json"
"strings"
"testing"
"oci-portal/internal/aiwire"
)
func mustAnthReq(t *testing.T, raw string) aiwire.MessagesRequest {
t.Helper()
var req aiwire.MessagesRequest
if err := json.Unmarshal([]byte(raw), &req); err != nil {
t.Fatalf("解析请求: %v", err)
}
return req
}
// TestAnthropicToResponsesBody 断言 system/消息/工具/effort 的直通装配与 store 强制。
func TestAnthropicToResponsesBody(t *testing.T) {
req := mustAnthReq(t, `{
"model": "xai.grok-4.3", "max_tokens": 128, "system": "你是助手", "stream": true,
"temperature": 0.5,
"messages": [
{"role": "user", "content": "东京天气?"},
{"role": "assistant", "content": [
{"type": "text", "text": "查询中"},
{"type": "tool_use", "id": "t1", "name": "get_weather", "input": {"city": "东京"}}
]},
{"role": "user", "content": [
{"type": "tool_result", "tool_use_id": "t1", "content": "晴 25 度"},
{"type": "text", "text": "继续"}
]}
],
"tools": [{"name": "get_weather", "description": "查天气", "input_schema": {"type": "object"}}],
"tool_choice": {"type": "any"},
"output_config": {"effort": "HIGH"}
}`)
payload, err := AnthropicToResponsesBody(req)
if err != nil {
t.Fatalf("AnthropicToResponsesBody: %v", err)
}
var body map[string]any
if err := json.Unmarshal(payload, &body); err != nil {
t.Fatalf("unmarshal: %v", err)
}
if body["model"] != "xai.grok-4.3" || body["max_output_tokens"] != float64(128) ||
body["instructions"] != "你是助手" || body["store"] != false || body["stream"] != true {
t.Fatalf("顶层字段装配错误: %v", body)
}
if body["tool_choice"] != "required" {
t.Fatalf("tool_choice = %v, want required", body["tool_choice"])
}
reasoning, _ := body["reasoning"].(map[string]any)
if reasoning["effort"] != "high" {
t.Fatalf("effort = %v, want high(小写透传)", reasoning)
}
input, _ := body["input"].([]any)
// user 消息、assistant 文本消息、function_call、function_call_output、末条 user 文本
if len(input) != 5 {
t.Fatalf("input 项数 = %d, want 5: %s", len(input), payload)
}
kinds := make([]string, 0, len(input))
for _, it := range input {
m := it.(map[string]any)
if ty, ok := m["type"].(string); ok {
kinds = append(kinds, ty)
} else {
kinds = append(kinds, "message/"+m["role"].(string))
}
}
want := "message/user,message/assistant,function_call,function_call_output,message/user"
if got := strings.Join(kinds, ","); got != want {
t.Fatalf("input 顺序 = %s, want %s", got, want)
}
if !strings.Contains(string(payload), `"output_text"`) {
t.Fatalf("assistant 历史文本应为 output_text 部件: %s", payload)
}
}
// TestAnthropicToResponsesBodyRejects 断言不支持的内容块拒绝与图片装配。
func TestAnthropicToResponsesBodyRejects(t *testing.T) {
bad := mustAnthReq(t, `{"model":"m","max_tokens":8,"messages":[
{"role":"user","content":[{"type":"document","source":{}}]}]}`)
if _, err := AnthropicToResponsesBody(bad); err == nil {
t.Fatal("document 块应拒绝")
}
img := mustAnthReq(t, `{"model":"m","max_tokens":8,"messages":[
{"role":"user","content":[{"type":"image","source":{"type":"base64","media_type":"image/png","data":"QUJD"}}]}]}`)
payload, err := AnthropicToResponsesBody(img)
if err != nil || !strings.Contains(string(payload), "data:image/png;base64,QUJD") {
t.Fatalf("图片应转 data URI: %v %s", err, payload)
}
}
// TestResponsesToAnthropic 断言输出块、stop_reason 与 usage 的回转。
func TestResponsesToAnthropic(t *testing.T) {
payload := []byte(`{"model":"xai.grok-4.3","status":"completed","output":[
{"type":"reasoning","summary":[]},
{"type":"message","content":[{"type":"output_text","text":"你好"}]},
{"type":"function_call","call_id":"c1","name":"get_weather","arguments":"{\"city\":\"东京\"}"}],
"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15,"input_tokens_details":{"cached_tokens":4}}}`)
out, err := ResponsesToAnthropic(payload, "msg_1")
if err != nil {
t.Fatalf("ResponsesToAnthropic: %v", err)
}
if len(out.Content) != 2 || out.Content[0].Type != "text" || out.Content[0].Text != "你好" {
t.Fatalf("content 装配错误: %+v", out.Content)
}
if out.Content[1].Type != "tool_use" || out.Content[1].ID != "c1" || string(out.Content[1].Input) != `{"city":"东京"}` {
t.Fatalf("tool_use 装配错误: %+v", out.Content[1])
}
if out.StopReason != "tool_use" {
t.Fatalf("stop_reason = %s, want tool_use", out.StopReason)
}
if out.Usage.InputTokens != 10 || out.Usage.OutputTokens != 5 || out.Usage.CacheReadInputTokens != 4 {
t.Fatalf("usage = %+v", out.Usage)
}
trunc := []byte(`{"model":"m","status":"incomplete","incomplete_details":{"reason":"max_output_tokens"},
"output":[{"type":"message","content":[{"type":"output_text","text":"半"}]}]}`)
out2, err := ResponsesToAnthropic(trunc, "msg_2")
if err != nil || out2.StopReason != "max_tokens" {
t.Fatalf("截断 stop_reason = %v, %v", out2, err)
}
}
func bridgeEventTypes(events []AnthEvent) string {
kinds := make([]string, 0, len(events))
for _, ev := range events {
kinds = append(kinds, ev.Event)
}
return strings.Join(kinds, ",")
}
// TestAnthRespBridge 断言流桥:文本增量、工具调用切块、reasoning 丢弃与收尾事件。
func TestAnthRespBridge(t *testing.T) {
st := NewAnthRespBridge("msg_1", "m1")
var events []AnthEvent
feed := func(lines ...string) {
for _, l := range lines {
events = append(events, st.Feed([]byte(l))...)
}
}
feed(`{"type":"response.created","response":{"model":"m1"}}`,
`{"type":"response.reasoning_text.delta","delta":"思考中"}`,
`{"type":"response.output_text.delta","delta":"你"}`,
`{"type":"response.output_text.delta","delta":"好"}`,
`{"type":"response.output_item.added","item":{"type":"function_call","call_id":"c1","name":"f"}}`,
`{"type":"response.function_call_arguments.delta","delta":"{\"a\":1}"}`,
`{"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":8,"output_tokens":4}}}`)
events = append(events, st.Finish()...)
got := bridgeEventTypes(events)
want := "message_start,content_block_start,content_block_delta,content_block_delta," +
"content_block_stop,content_block_start,content_block_delta,content_block_stop," +
"message_delta,message_stop"
if got != want {
t.Fatalf("事件序列:\n got %s\nwant %s", got, want)
}
if st.Usage().InputTokens != 8 || st.Usage().OutputTokens != 4 {
t.Fatalf("usage = %+v", st.Usage())
}
}
// TestAnthRespBridgeEmpty 断言空流也产出完整事件骨架。
func TestAnthRespBridgeEmpty(t *testing.T) {
st := NewAnthRespBridge("msg_1", "m1")
if got := bridgeEventTypes(st.Finish()); got != "message_start,message_delta,message_stop" {
t.Errorf("空流事件序列 = %s", got)
}
}