293 lines
11 KiB
Go
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, "密钥已删后不得写入")
|
|
}
|