136 lines
5.2 KiB
Go
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)
|
|
}
|
|
}
|