package oci import ( "encoding/json" "testing" "github.com/oracle/oci-go-sdk/v65/generativeaiinference" "oci-portal/internal/aiwire" ) // TestIrToSDK 断言 IR→GENERIC 的角色拆装、采样参数与工具映射。 func TestIrToSDK(t *testing.T) { temp, mt := 0.7, 100 ir := aiwire.ChatRequest{ Model: "meta.llama-3.3-70b-instruct", Messages: []aiwire.ChatMessage{ {Role: "system", Content: aiwire.NewTextContent("你是助手")}, {Role: "user", Content: aiwire.NewTextContent("你好")}, {Role: "assistant", ToolCalls: []aiwire.ToolCall{{ID: "c1", Type: "function", Function: aiwire.FunctionCall{Name: "get_weather", Arguments: `{"city":"东京"}`}}}}, {Role: "tool", ToolCallID: "c1", Content: aiwire.NewTextContent("晴")}, }, Temperature: &temp, MaxTokens: &mt, Stop: aiwire.StringList{"END"}, Tools: []aiwire.Tool{{Type: "function", Function: aiwire.FunctionDef{Name: "get_weather", Parameters: json.RawMessage(`{"type":"object"}`)}}}, ToolChoice: json.RawMessage(`"auto"`), } req := irToSDK(ir, true) if len(req.Messages) != 4 { t.Fatalf("messages = %d, want 4", len(req.Messages)) } if _, ok := req.Messages[0].(generativeaiinference.SystemMessage); !ok { t.Fatalf("messages[0] = %T, want SystemMessage", req.Messages[0]) } am, ok := req.Messages[2].(generativeaiinference.AssistantMessage) if !ok || len(am.ToolCalls) != 1 { t.Fatalf("assistant tool calls 未映射: %T %+v", req.Messages[2], am) } tm, ok := req.Messages[3].(generativeaiinference.ToolMessage) if !ok || deref(tm.ToolCallId) != "c1" { t.Fatalf("tool message 未映射 toolCallId: %+v", tm) } if req.Temperature == nil || *req.Temperature != 0.7 { t.Fatalf("temperature 未直通") } if req.MaxTokens == nil || *req.MaxTokens != 100 || len(req.Stop) != 1 { t.Fatalf("maxTokens/stop 未直通") } if _, ok := req.ToolChoice.(generativeaiinference.ToolChoiceAuto); !ok { t.Fatalf("toolChoice = %T, want auto", req.ToolChoice) } if len(req.Tools) != 1 || req.IsStream == nil || !*req.IsStream || req.StreamOptions == nil { t.Fatalf("tools/stream 选项未映射") } } // TestSdkToIR 断言响应文本、工具调用与用量的反向映射。 func TestSdkToIR(t *testing.T) { text, fr := "东京晴", "tool_calls" idx, pt, ct, tt := 0, 10, 5, 15 id, name, args := "c1", "get_weather", `{"city":"东京"}` resp := generativeaiinference.GenericChatResponse{ Choices: []generativeaiinference.ChatChoice{{ Index: &idx, Message: generativeaiinference.AssistantMessage{ Content: []generativeaiinference.ChatContent{generativeaiinference.TextContent{Text: &text}}, ToolCalls: []generativeaiinference.ToolCall{generativeaiinference.FunctionCall{Id: &id, Name: &name, Arguments: &args}}, }, FinishReason: &fr, }}, Usage: &generativeaiinference.Usage{PromptTokens: &pt, CompletionTokens: &ct, TotalTokens: &tt}, } out := sdkToIR(resp, "m1") if len(out.Choices) != 1 || out.Choices[0].Message.Content.JoinText() != "东京晴" { t.Fatalf("文本未映射: %+v", out) } if out.Choices[0].FinishReason != "tool_calls" || len(out.Choices[0].Message.ToolCalls) != 1 { t.Fatalf("finishReason/toolCalls 未映射: %+v", out.Choices[0]) } if out.Usage == nil || out.Usage.TotalTokens != 15 { t.Fatalf("usage 未映射: %+v", out.Usage) } } // TestParseGenAiEvent 断言流事件宽容解析:文本增量、结束原因与 usage 事件。 func TestParseGenAiEvent(t *testing.T) { chunk, ok := parseGenAiEvent([]byte(`{"index":0,"message":{"role":"ASSISTANT","content":[{"type":"TEXT","text":"你"}]}}`), "m1") if !ok || len(chunk.Choices) != 1 || chunk.Choices[0].Delta.Content != "你" { t.Fatalf("文本增量解析失败: %+v", chunk) } chunk, ok = parseGenAiEvent([]byte(`{"finishReason":"stop","usage":{"promptTokens":3,"completionTokens":2,"totalTokens":5}}`), "m1") if !ok || chunk.Choices[0].FinishReason == nil || *chunk.Choices[0].FinishReason != "stop" { t.Fatalf("finishReason 解析失败: %+v", chunk) } if chunk.Usage == nil || chunk.Usage.TotalTokens != 5 { t.Fatalf("usage 解析失败: %+v", chunk.Usage) } if _, ok := parseGenAiEvent([]byte(`{}`), "m1"); ok { t.Fatal("空事件应返回 ok=false") } } func TestSdkUsageCachedTokens(t *testing.T) { p, c, tot, cached := 10, 5, 15, 8 u := sdkUsageToIR(&generativeaiinference.Usage{PromptTokens: &p, CompletionTokens: &c, TotalTokens: &tot, PromptTokensDetails: &generativeaiinference.PromptTokensDetails{CachedTokens: &cached}}) if u.CachedTokens() != 8 { t.Errorf("CachedTokens = %d, want 8", u.CachedTokens()) } if sdkUsageToIR(&generativeaiinference.Usage{PromptTokens: &p}).CachedTokens() != 0 { t.Error("无细分时 CachedTokens 应为 0") } } func TestIrToCohereSDK(t *testing.T) { ir := aiwire.ChatRequest{ Model: "cohere.command-r-plus", Messages: []aiwire.ChatMessage{ {Role: "system", Content: aiwire.NewTextContent("你是助手")}, {Role: "user", Content: aiwire.NewTextContent("第一问")}, {Role: "assistant", Content: aiwire.NewTextContent("第一答")}, {Role: "user", Content: aiwire.NewTextContent("第二问")}, }, } req, err := irToCohereSDK(ir, true) if err != nil { t.Fatalf("irToCohereSDK: %v", err) } if *req.Message != "第二问" || *req.PreambleOverride != "你是助手" || len(req.ChatHistory) != 2 { t.Errorf("拆装错误: msg=%q preamble=%v history=%d", *req.Message, req.PreambleOverride, len(req.ChatHistory)) } if _, ok := req.ChatHistory[0].(generativeaiinference.CohereUserMessage); !ok { t.Errorf("history[0] 应为 user: %T", req.ChatHistory[0]) } if _, ok := req.ChatHistory[1].(generativeaiinference.CohereChatBotMessage); !ok { t.Errorf("history[1] 应为 chatbot: %T", req.ChatHistory[1]) } if req.IsStream == nil || !*req.IsStream { t.Error("IsStream 未设置") } if req.StreamOptions == nil || req.StreamOptions.IsIncludeUsage == nil || !*req.StreamOptions.IsIncludeUsage { t.Error("流式应默认开启 usage 回传(StreamOptions.IsIncludeUsage)") } // 工具定义降级为扁平参数表(仅取 JSON Schema 顶层 properties) ir.Tools = []aiwire.Tool{{Type: "function", Function: aiwire.FunctionDef{ Name: "get_weather", Parameters: json.RawMessage(`{"type":"object","properties":{"city":{"type":"string","description":"城市"}},"required":["city"]}`)}}} req2, err := irToCohereSDK(ir, false) if err != nil { t.Fatalf("工具请求应支持: %v", err) } tool := req2.Tools[0] if *tool.Name != "get_weather" || *tool.Description != "get_weather" { t.Errorf("tool = %+v", tool) } if pd, ok := tool.ParameterDefinitions["city"]; !ok || *pd.Type != "string" || pd.IsRequired == nil || !*pd.IsRequired { t.Errorf("param city = %+v", pd) } // 图片输入仍拒绝 ir.Tools = nil ir.Messages = append(ir.Messages, aiwire.ChatMessage{Role: "user", Content: aiwire.NewPartsContent([]aiwire.ContentPart{ {Type: "image_url", ImageURL: &aiwire.ImageURL{URL: "data:image/png;base64,xx"}}})}) if _, err := irToCohereSDK(ir, false); err == nil { t.Error("图片输入应被拒绝") } // 无 user 消息被拒 if _, err := irToCohereSDK(aiwire.ChatRequest{Model: "cohere.x", Messages: []aiwire.ChatMessage{{Role: "system", Content: aiwire.NewTextContent("s")}}}, false); err == nil { t.Error("无 user 消息应被拒绝") } } func TestCohereSDKToIR(t *testing.T) { text := "回答" resp := generativeaiinference.CohereChatResponse{ Text: &text, FinishReason: generativeaiinference.CohereChatResponseFinishReasonComplete, } out := cohereSDKToIR(resp, "cohere.command-r-plus") if out.Choices[0].Message.Content.JoinText() != "回答" || out.Choices[0].FinishReason != "stop" { t.Errorf("cohereSDKToIR = %+v", out.Choices[0]) } } func TestParseGenAiEventCohere(t *testing.T) { chunk, ok := parseGenAiEvent([]byte(`{"apiFormat":"COHERE","text":"你好"}`), "cohere.command-r") if !ok || chunk.Choices[0].Delta.Content != "你好" { t.Errorf("cohere text 事件解析 = %+v, %v", chunk, ok) } chunk, ok = parseGenAiEvent([]byte(`{"apiFormat":"COHERE","finishReason":"COMPLETE","usage":{"promptTokens":3,"completionTokens":5,"totalTokens":8}}`), "cohere.command-r") if !ok || *chunk.Choices[0].FinishReason != "stop" || chunk.Usage.TotalTokens != 8 { t.Errorf("cohere 终帧解析 = %+v, %v", chunk, ok) } }