AI 网关切换 OpenAI 兼容面并移除 chat 端点,新增模型黑白名单
This commit is contained in:
@@ -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 }
|
||||
@@ -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
@@ -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)。
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 }
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user