package service import ( "context" "encoding/json" "errors" "strings" "testing" "time" "oci-portal/internal/aiwire" "oci-portal/internal/model" "oci-portal/internal/oci" ) // TestSpeechBodyNormalize 断言 TTS 请求体校验与 language 缺省注入。 func TestSpeechBodyNormalize(t *testing.T) { tests := []struct { name string raw string wantErr bool wantLang string }{ {"缺 model 拒绝", `{"input":"你好"}`, true, ""}, {"缺 input 拒绝", `{"model":"xai.grok-tts"}`, true, ""}, {"language 缺省注入 auto", `{"model":"xai.grok-tts","input":"你好","voice":"ara"}`, false, "auto"}, {"language 已有保留", `{"model":"xai.grok-tts","input":"你好","language":"zh"}`, false, "zh"}, {"非 JSON 拒绝", ``, true, ""}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { modelName, body, err := SpeechBodyNormalize([]byte(tt.raw)) if (err != nil) != tt.wantErr { t.Fatalf("err = %v, wantErr %v", err, tt.wantErr) } if err != nil { return } var out map[string]any _ = json.Unmarshal(body, &out) if modelName != "xai.grok-tts" || out["language"] != tt.wantLang { t.Fatalf("model=%s language=%v, want %s", modelName, out["language"], tt.wantLang) } }) } } // TestTtsBodyConvert 断言 xAI 官方 TTS 格式到 OpenAI 兼容形态的转换。 func TestTtsBodyConvert(t *testing.T) { tests := []struct { name string raw string wantErr bool wantModel string }{ {"缺 text 拒绝", `{"language":"zh"}`, true, ""}, {"缺 language 拒绝", `{"text":"你好"}`, true, ""}, {"缺省注入默认模型", `{"text":"你好","language":"zh"}`, false, "xai.grok-tts"}, {"model 扩展字段可覆盖", `{"model":"xai.other-tts","text":"你好","language":"auto"}`, false, "xai.other-tts"}, {"非 JSON 拒绝", ``, true, ""}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { modelName, _, err := TtsBodyConvert([]byte(tt.raw)) if (err != nil) != tt.wantErr { t.Fatalf("err = %v, wantErr %v", err, tt.wantErr) } if err != nil { return } if modelName != tt.wantModel { t.Fatalf("model = %s, want %s", modelName, tt.wantModel) } }) } } // TestTtsBodyConvertMapping 断言字段映射与未知字段保留。 func TestTtsBodyConvertMapping(t *testing.T) { raw := `{"text":"你好","language":"zh","voice_id":"ara","speed":1.2,` + `"output_format":{"codec":"mp3","sample_rate":44100}}` _, body, err := TtsBodyConvert([]byte(raw)) if err != nil { t.Fatalf("err = %v", err) } var out map[string]any _ = json.Unmarshal(body, &out) if out["input"] != "你好" || out["voice"] != "ara" { t.Fatalf("input/voice 映射错误: %v", out) } if _, ok := out["text"]; ok { t.Fatal("text 字段应被移除") } if _, ok := out["voice_id"]; ok { t.Fatal("voice_id 字段应被移除") } of, _ := out["output_format"].(map[string]any) if of == nil || of["codec"] != "mp3" { t.Fatalf("output_format 应原样保留: %v", out["output_format"]) } if out["language"] != "zh" || out["speed"] == nil { t.Fatalf("language/speed 应保留: %v", out) } } // TestModerationInputs 断言 input 的 string / []string 解析与边界校验。 func TestModerationInputs(t *testing.T) { tests := []struct { name string raw string wantN int wantErr bool }{ {"单字符串", `"hello"`, 1, false}, {"数组", `["a","b"]`, 2, false}, {"空字符串拒绝", `""`, 0, true}, {"空数组拒绝", `[]`, 0, true}, {"含空条目拒绝", `["a",""]`, 0, true}, {"超上限拒绝", `["1","2","3","4","5","6","7","8","9"]`, 0, true}, {"非法类型拒绝", `{"x":1}`, 0, true}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { got, err := ModerationInputs(json.RawMessage(tt.raw)) if (err != nil) != tt.wantErr || len(got) != tt.wantN { t.Fatalf("got %v (err=%v), want n=%d wantErr=%v", got, err, tt.wantN, tt.wantErr) } }) } } // TestModerationResultMapping 断言 guardrails 结果到 OpenAI 外壳的映射与 flagged 判定。 func TestModerationResultMapping(t *testing.T) { one := 1.0 zero := 0.0 tests := []struct { name string outcome oci.GuardrailsOutcome wantFlagged bool wantPii int }{ {"内容审核命中", oci.GuardrailsOutcome{Categories: []oci.GuardrailCategory{{Name: "OVERALL", Score: 1}}}, true, 0}, {"提示注入命中", oci.GuardrailsOutcome{PromptInjectionScore: &one}, true, 0}, {"仅 PII 不 flag", oci.GuardrailsOutcome{Categories: []oci.GuardrailCategory{{Name: "OVERALL", Score: 0}}, PromptInjectionScore: &zero, Pii: []oci.GuardrailPiiHit{{Text: "Jane", Label: "PERSON", Score: 0.99}}}, false, 1}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { res := moderationResult(&tt.outcome) if res.Flagged != tt.wantFlagged || len(res.Pii) != tt.wantPii { t.Fatalf("res = %+v, want flagged=%v pii=%d", res, tt.wantFlagged, tt.wantPii) } }) } } // TestAiRerank 断言重排编排:能力路由、top_n 透传、return_documents 回填与越界防御。 func TestAiRerank(t *testing.T) { client := &gatewayStubClient{ fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}, rerankRanks: []oci.RerankRank{{Index: 1, Score: 0.9}, {Index: 0, Score: 0.4}, {Index: 9, Score: 0.1}}, } gw, svc := newTestGateway(t, client) cfg := importAliveConfig(t, svc) ch := seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1) gw.db.Create(&model.AiModelCache{ChannelID: ch.ID, ModelOcid: "ocid1..rr", Name: "cohere.rerank-v4.0-fast", Vendor: "cohere", Capability: "RERANK", SyncedAt: time.Now()}) ctx := context.Background() yes := true req := aiwire.RerankRequest{Model: "cohere.rerank-v4.0-fast", Query: "q", Documents: []string{"d0", "d1"}, ReturnDocuments: &yes} resp, _, err := gw.Rerank(ctx, req, "") if err != nil || len(resp.Results) != 2 { t.Fatalf("Rerank = %+v, %v(越界 index 应被丢弃)", resp, err) } if resp.Results[0].Index != 1 || resp.Results[0].Document == nil || resp.Results[0].Document.Text != "d1" { t.Fatalf("results[0] = %+v", resp.Results[0]) } // 对话模型名打 rerank:能力不匹配 → 未知模型 if _, _, err := gw.Rerank(ctx, aiwire.RerankRequest{Model: "meta.llama-3.3-70b-instruct", Query: "q", Documents: []string{"d"}}, ""); !errors.Is(err, ErrAiUnknownModel) { t.Errorf("chat 模型走 rerank err = %v, want ErrAiUnknownModel", err) } } // TestAiSpeech 断言 TTS 编排走 TTS 能力路由并透传音频与 Content-Type。 func TestAiSpeech(t *testing.T) { client := &gatewayStubClient{ fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}, speechAudio: []byte{0xFF, 0xF3}, speechCT: "audio/mpeg", } gw, svc := newTestGateway(t, client) cfg := importAliveConfig(t, svc) ch := seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1) gw.db.Create(&model.AiModelCache{ChannelID: ch.ID, ModelOcid: "ocid1..tts", Name: "xai.grok-tts", Vendor: "xai", Capability: "TTS", SyncedAt: time.Now()}) audio, ct, _, err := gw.Speech(context.Background(), "xai.grok-tts", []byte(`{"model":"xai.grok-tts","input":"你好","language":"auto"}`), "") if err != nil || ct != "audio/mpeg" || len(audio) != 2 { t.Fatalf("Speech = %d bytes, ct=%s, %v", len(audio), ct, err) } } // TestAiModerations 断言审核编排:无模型维度按分组选渠道,多条输入逐条聚合。 func TestAiModerations(t *testing.T) { one := 1.0 client := &gatewayStubClient{ fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}, guardOutcome: &oci.GuardrailsOutcome{Categories: []oci.GuardrailCategory{{Name: "OVERALL", Score: 1}}, PromptInjectionScore: &one}, } gw, svc := newTestGateway(t, client) cfg := importAliveConfig(t, svc) seedChannel(t, gw, cfg.ID, "us-chicago-1", 1, 1) resp, meta, err := gw.Moderations(context.Background(), "modr_test", []string{"a", "b"}, "") if err != nil || len(resp.Results) != 2 || !resp.Results[0].Flagged || resp.ID != "modr_test" { t.Fatalf("Moderations = %+v, meta=%+v, %v", resp, meta, err) } if !strings.Contains(resp.Model, "guardrails") { t.Errorf("model = %s", resp.Model) } // 分组不匹配 → 无可用渠道 if _, _, err := gw.Moderations(context.Background(), "modr_x", []string{"a"}, "ghost-group"); !errors.Is(err, ErrAiNoChannel) { t.Errorf("ghost 分组 err = %v, want ErrAiNoChannel", err) } }