Files
oci-portal/internal/service/aigateway_extras_test.go
2026-07-22 16:51:23 +08:00

293 lines
11 KiB
Go

package service
import (
"context"
"encoding/json"
"errors"
"strings"
"testing"
"time"
"oci-portal/internal/aiwire"
"oci-portal/internal/model"
"oci-portal/internal/oci"
"gorm.io/gorm"
)
// 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 拒绝", `<html>`, 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 拒绝", `<html>`, 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)
}
}
// TestCleanupTableBatchesOverflow 锁定超限清理的固定批次行为:
// 溢出量大于单批上限时分多轮删完,且始终删最旧的行。
func TestCleanupTableBatchesOverflow(t *testing.T) {
gw, _ := newTestGateway(t, &fakeClient{})
old := cleanupBatch
cleanupBatch = 3
t.Cleanup(func() { cleanupBatch = old })
for i := 0; i < 10; i++ {
if err := gw.db.Create(&model.AiCallLog{Endpoint: "chat"}).Error; err != nil {
t.Fatalf("seed %d: %v", i, err)
}
}
// maxRows=2:溢出 8 行,需 3 轮批次(3+3+2)
gw.cleanupTable(context.Background(), &model.AiCallLog{}, 24*365*time.Hour, 2, "test")
var rest []model.AiCallLog
if err := gw.db.Order("id").Find(&rest).Error; err != nil {
t.Fatalf("load rest: %v", err)
}
if len(rest) != 2 {
t.Fatalf("remaining = %d, want 2", len(rest))
}
for _, row := range rest {
if row.ID <= 8 {
t.Errorf("row %d survived, want oldest deleted first", row.ID)
}
}
}
// TestLogContentRechecksKeyInsideTx 锁定「立即关闭」语义:写入事务内回读密钥,
// 窗口已关、密钥停用或已删时,鉴权期的旧快照不得再落敏感正文。
func TestLogContentRechecksKeyInsideTx(t *testing.T) {
gw, _ := newTestGateway(t, &fakeClient{})
if err := gw.db.AutoMigrate(&model.AiContentLog{}); err != nil {
t.Fatalf("migrate content log: %v", err)
}
future := time.Now().Add(time.Hour)
key := model.AiKey{Name: "k", KeyHash: "h", Enabled: true, ContentLogUntil: &future}
call := model.AiCallLog{Endpoint: "chat"}
if err := gw.db.Create(&key).Error; err != nil {
t.Fatalf("seed key: %v", err)
}
if err := gw.db.Create(&call).Error; err != nil {
t.Fatalf("seed call: %v", err)
}
countIs := func(want int64, note string) {
t.Helper()
var n int64
if err := gw.db.Model(&model.AiContentLog{}).Count(&n).Error; err != nil || n != want {
t.Fatalf("%s: count=%d (%v), want %d", note, n, err, want)
}
}
gw.LogContent(model.AiContentLog{CallLogID: call.ID, KeyID: key.ID, RequestBody: "prompt"})
countIs(1, "窗口开启时应写入")
// 管理员立即关闭窗口:在途请求携带的旧快照不得再写
if err := gw.db.Model(&model.AiKey{}).Where("id = ?", key.ID).
Update("content_log_until", gorm.Expr("NULL")).Error; err != nil {
t.Fatalf("close window: %v", err)
}
gw.LogContent(model.AiContentLog{CallLogID: call.ID, KeyID: key.ID, RequestBody: "late"})
countIs(1, "窗口关闭后不得写入")
// 密钥删除后同样拒写
if err := gw.db.Delete(&model.AiKey{}, key.ID).Error; err != nil {
t.Fatalf("delete key: %v", err)
}
gw.LogContent(model.AiContentLog{CallLogID: call.ID, KeyID: key.ID, RequestBody: "orphan"})
countIs(1, "密钥已删后不得写入")
}