Files
go-sip/internal/ai/current_pipeline_test.go
T

136 lines
5.2 KiB
Go

package ai
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestCurrentTTSPassesApprovedParametersToSDK(t *testing.T) {
task, providers := currentFixture(t, "full_ai")
task = changeCurrentAgent(t, task, func(agent map[string]any) { agent["tts"].(map[string]any)["speed"] = 1.3 })
var captured map[string]any
var key, resource string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
key, resource = r.Header.Get("X-Api-Key"), r.Header.Get("X-Api-Resource-Id")
if r.Method != http.MethodPost || r.URL.Path != "/api/v3/tts/unidirectional" {
http.Error(w, "unexpected SDK endpoint", http.StatusBadRequest)
return
}
if err := json.NewDecoder(r.Body).Decode(&captured); err != nil {
http.Error(w, "bad SDK request", http.StatusBadRequest)
return
}
_, _ = fmt.Fprintf(w, `{"code":0,"data":%q}`+"\n", base64.StdEncoding.EncodeToString([]byte{1, 0, 2, 0}))
_, _ = fmt.Fprintln(w, `{"code":20000000,"message":"ok","data":null}`)
}))
defer server.Close()
p := providers["tts-example"]
p.Endpoint = server.URL
providers[p.ProviderRef] = p
bound, err := BindCurrent(task, providers)
if err != nil {
t.Fatal(err)
}
audio, err := bound.Synthesize(context.Background(), "批准的回复")
if err != nil || string(audio) != string([]byte{1, 0, 2, 0}) {
t.Fatalf("SDK TTS response: length=%d err=%v", len(audio), err)
}
if key != p.Credential || resource != "example-tts" {
t.Fatal("provider credential/resource not passed to official request")
}
params := captured["req_params"].(map[string]any)
format := params["audio_params"].(map[string]any)
if params["text"] != "批准的回复" || params["speaker"] != "example-neutral" || format["format"] != "pcm_s16le" || format["sample_rate"] != float64(16000) || format["speech_rate"] != float64(30) {
t.Fatalf("SDK TTS approved speed/voice/format not preserved: %v", format)
}
}
func TestCurrentLLMPassesExplicitZeroAndDoesNotRetry(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 := BindCurrent(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 TestCurrentKeywordHangupUsesFinalUserTextOnly(t *testing.T) {
task, providers := currentFixture(t, "full_ai")
bound, err := BindCurrent(task, providers)
if err != nil {
t.Fatal(err)
}
calls := 0
call, err := NewCurrentCall(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 TestCurrentASRRequiresFinalResult(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)
}
}