181 lines
6.5 KiB
Go
181 lines
6.5 KiB
Go
package ai
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"os/exec"
|
|
"strings"
|
|
"testing"
|
|
|
|
"git.ipao.vip/rogee/go-sip/internal/media"
|
|
)
|
|
|
|
func TestTTSPassesApprovedParametersToSDK(t *testing.T) {
|
|
if _, err := exec.LookPath("ffmpeg"); err != nil {
|
|
t.Skip("Bailian audio conversion requires ffmpeg")
|
|
}
|
|
want := []byte{1, 0, 2, 0}
|
|
wav, _, err := media.EncodeMonoWAV(want, 1024)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
task, providers := currentFixture(t, "full_ai")
|
|
var captured map[string]any
|
|
var credential string
|
|
postCount, getCount := 0, 0
|
|
var server *httptest.Server
|
|
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
switch r.URL.Path {
|
|
case "/api/v1/services/aigc/multimodal-generation/generation":
|
|
postCount++
|
|
credential = r.Header.Get("Authorization")
|
|
if r.Method != http.MethodPost || json.NewDecoder(r.Body).Decode(&captured) != nil {
|
|
http.Error(w, "invalid generation request", http.StatusBadRequest)
|
|
return
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = fmt.Fprintf(w, `{"output":{"audio":{"url":%q}}}`, server.URL+"/audio?test-token=redacted")
|
|
case "/audio":
|
|
getCount++
|
|
_, _ = w.Write(wav)
|
|
default:
|
|
http.Error(w, "unexpected endpoint", http.StatusNotFound)
|
|
}
|
|
}))
|
|
defer server.Close()
|
|
p := providers["tts-example"]
|
|
p.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation"
|
|
providers[p.ProviderRef] = p
|
|
bound, err := Bind(task, providers)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
audio, err := bound.Synthesize(context.Background(), "批准的回复")
|
|
if err != nil || string(audio) != string(want) {
|
|
t.Fatalf("Bailian TTS response: length=%d err=%v", len(audio), err)
|
|
}
|
|
input, ok := captured["input"].(map[string]any)
|
|
if !ok || captured["model"] != "qwen3-tts-flash" || input["voice"] != "Cherry" || input["language_type"] != "Chinese" || input["text"] != "批准的回复" || credential != "Bearer "+p.Credential || postCount != 1 || getCount != 1 {
|
|
t.Fatal("approved Bailian model/voice/text/credential or single-request bound was lost")
|
|
}
|
|
}
|
|
|
|
func TestLLMPassesExplicitZeroAndDoesNotRetry(t *testing.T) {
|
|
task, providers := currentFixture(t, "full_ai")
|
|
calls := 0
|
|
var captured map[string]any
|
|
var authorization string
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
calls++
|
|
authorization = r.Header.Get("Authorization")
|
|
if r.Method != http.MethodPost || !strings.HasSuffix(r.URL.Path, "/chat/completions") {
|
|
http.Error(w, "unexpected LLM request", http.StatusBadRequest)
|
|
return
|
|
}
|
|
if err := json.NewDecoder(r.Body).Decode(&captured); err != nil {
|
|
http.Error(w, "invalid LLM request", http.StatusBadRequest)
|
|
return
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = fmt.Fprintln(w, `{"id":"mock","object":"chat.completion","created":0,"model":"example-chat","choices":[{"index":0,"message":{"role":"assistant","content":"收到"},"finish_reason":"stop"}]}`)
|
|
}))
|
|
defer server.Close()
|
|
p := providers["llm-example"]
|
|
p.Endpoint = server.URL + "/v1"
|
|
providers[p.ProviderRef] = p
|
|
bound, err := Bind(task, providers)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
reply, err := bound.Complete(context.Background(), "用户说话")
|
|
if err != nil || reply != "收到" || calls != 1 || authorization != "Bearer "+p.Credential {
|
|
t.Fatalf("LLM SDK result=%q requests=%d err=%v", reply, calls, err)
|
|
}
|
|
if captured["model"] != "example-chat" || captured["temperature"] != float64(0) || captured["max_tokens"] != float64(256) {
|
|
t.Fatalf("LLM business values were not transmitted: %v", captured)
|
|
}
|
|
messages := captured["messages"].([]any)
|
|
if len(messages) != 2 || messages[0].(map[string]any)["content"] != "Example only" || messages[1].(map[string]any)["content"] != "用户说话" {
|
|
t.Fatal("immutable prompt and final user text must reach LLM")
|
|
}
|
|
}
|
|
|
|
func TestLLMFailureAndEmptyChoicesAreNotRetried(t *testing.T) {
|
|
for _, tc := range []struct {
|
|
name string
|
|
code int
|
|
body string
|
|
}{
|
|
{name: "provider failure", code: http.StatusInternalServerError, body: `{"error":{"message":"injected failure","type":"server_error"}}`},
|
|
{name: "empty choices", code: http.StatusOK, body: `{"id":"mock","object":"chat.completion","created":0,"model":"example-chat","choices":[]}`},
|
|
} {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
task, providers := currentFixture(t, "full_ai")
|
|
calls := 0
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
calls++
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(tc.code)
|
|
_, _ = fmt.Fprintln(w, tc.body)
|
|
}))
|
|
defer server.Close()
|
|
provider := providers["llm-example"]
|
|
provider.Endpoint = server.URL + "/v1"
|
|
providers[provider.ProviderRef] = provider
|
|
bound, err := Bind(task, providers)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if reply, err := bound.Complete(context.Background(), "用户说话"); err == nil || calls != 1 || reply != "" {
|
|
t.Fatalf("LLM failure must be visible without retry: reply=%q calls=%d err=%v", reply, calls, err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestKeywordHangupUsesFinalUserTextOnly(t *testing.T) {
|
|
task, providers := currentFixture(t, "full_ai")
|
|
bound, err := Bind(task, providers)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
calls := 0
|
|
call, err := NewCall(bound, func(context.Context) error { calls++; return nil })
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for _, segment := range []ASRSegment{
|
|
{Source: "user", Text: "不用了", Final: false},
|
|
{Source: "assistant", Text: "不用了", Final: true},
|
|
{Source: "user", Text: "继续说", Final: true},
|
|
} {
|
|
stop, err := call.HandleFinalASR(context.Background(), segment)
|
|
if err != nil || stop || calls != 0 {
|
|
t.Fatalf("interim/assistant/unmatched text cannot hang up: stopped=%t err=%v calls=%d", stop, err, calls)
|
|
}
|
|
}
|
|
for i := 0; i < 2; i++ {
|
|
stop, err := call.HandleFinalASR(context.Background(), ASRSegment{Source: "user", Text: "我不用了", Final: true})
|
|
if err != nil || stop != (i == 0) || calls != 1 {
|
|
t.Fatalf("matched final user text hangs up at most once: stopped=%t err=%v calls=%d", stop, err, calls)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestASRRequiresFinalResult(t *testing.T) {
|
|
var text finalASRAccumulator
|
|
text.Add("不用了", false)
|
|
if _, err := text.Result(); err == nil {
|
|
t.Fatal("interim-only transcript must never trigger keyword hangup or LLM")
|
|
}
|
|
text.Add("继续", true)
|
|
text.Add("不用了", false)
|
|
result, err := text.Result()
|
|
if err != nil || result != "继续" {
|
|
t.Fatalf("only the first final user result may be used: %q %v", result, err)
|
|
}
|
|
}
|