194 lines
8.1 KiB
Go
194 lines
8.1 KiB
Go
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)
|
|
}
|
|
}
|