diff --git a/internal/ai/current.go b/internal/ai/current.go new file mode 100644 index 0000000..9f9aa48 --- /dev/null +++ b/internal/ai/current.go @@ -0,0 +1,233 @@ +package ai + +import ( + "encoding/json" + "errors" + "fmt" + "math" + "net/url" + "time" + + "git.ipao.vip/rogee/go-sip/internal/configread" + "git.ipao.vip/rogee/go-sip/internal/contract" + doubaospeech "github.com/GizClaw/doubao-speech-go" +) + +// CurrentBound is a per-call, immutable binding of the approved task and +// provider settings. It must be built before dispatching any side effects. +// Credential values are never included in errors or logs. +type CurrentBound struct { + Mode string + ASR CurrentASR + LLM *CurrentLLM + TTS *CurrentTTS + Prompt string + Opening string + HangupKeywords []string +} + +type CurrentASR struct { + Provider configread.CurrentProvider + Request doubaospeech.ASRV2Config + Timeout time.Duration +} + +type CurrentLLM struct { + Provider configread.CurrentProvider + Model string + Temperature *float64 + MaxTokens *int64 + Timeout time.Duration +} + +type CurrentTTS struct { + Provider configread.CurrentProvider + Request doubaospeech.TTSV2Request + Timeout time.Duration +} + +type currentAgentSettings struct { + Mode string `json:"mode"` + ASR struct { + ProviderRef string `json:"provider_ref"` + Model string `json:"model"` + Language string `json:"language"` + Interim *bool `json:"interim"` + TimeoutMS *int64 `json:"timeout_ms"` + Input struct { + Encoding string `json:"encoding"` + SampleRateHz int `json:"sample_rate_hz"` + Channels int `json:"channels"` + SampleWidthBytes int `json:"sample_width_bytes"` + } `json:"input"` + } `json:"asr"` + LLM *struct { + ProviderRef string `json:"provider_ref"` + Model string `json:"model"` + Temperature *float64 `json:"temperature"` + MaxTokens *int64 `json:"max_tokens"` + TimeoutMS *int64 `json:"timeout_ms"` + } `json:"llm"` + TTS *struct { + ProviderRef string `json:"provider_ref"` + Model string `json:"model"` + Voice string `json:"voice"` + Speed *float64 `json:"speed"` + TimeoutMS *int64 `json:"timeout_ms"` + Format struct { + Encoding string `json:"encoding"` + SampleRateHz int `json:"sample_rate_hz"` + Channels int `json:"channels"` + } `json:"format"` + } `json:"tts"` + Prompt *struct { + Text string `json:"text"` + } `json:"prompt"` + Conversation *struct { + Opening string `json:"opening"` + HangupKeywords []string `json:"hangup_keywords"` + } `json:"conversation"` +} + +// BindCurrent rejects schema-valid settings which the selected SDK cannot +// express. In particular, the published TTS schema is intentionally not +// silently narrowed to the SDK's speed/format capabilities. +func BindCurrent(task configread.CurrentTask, providers map[string]configread.CurrentProvider) (CurrentBound, error) { + if len(task.Raw) == 0 { + return CurrentBound{}, errors.New("approved task snapshot is missing") + } + if err := contract.ValidateCurrent("config-read", task.Raw); err != nil { + return CurrentBound{}, fmt.Errorf("task snapshot: %w", err) + } + var frozen configread.CurrentTask + if err := json.Unmarshal(task.Raw, &frozen); err != nil { + return CurrentBound{}, fmt.Errorf("decode task snapshot: %w", err) + } + if frozen.Resource != "task_config" || frozen.TaskID != task.TaskID || frozen.TenantID != task.TenantID || frozen.DispatcherID != task.DispatcherID { + return CurrentBound{}, errors.New("approved task snapshot identity mismatch") + } + var settings currentAgentSettings + if err := json.Unmarshal(frozen.Agent.Raw, &settings); err != nil { + return CurrentBound{}, fmt.Errorf("decode immutable AI settings: %w", err) + } + bound := CurrentBound{Mode: settings.Mode} + asrProvider, err := currentProvider(providers, settings.ASR.ProviderRef, "asr", "volcengine_asr") + if err != nil { + return CurrentBound{}, err + } + if settings.ASR.Input.Encoding != "pcm_s16le" || settings.ASR.Input.Channels != 1 || settings.ASR.Input.SampleWidthBytes != 2 { + return CurrentBound{}, errors.New("ASR input format is unsupported by selected SDK") + } + sampleRate, err := currentSampleRate(settings.ASR.Input.SampleRateHz) + if err != nil { + return CurrentBound{}, fmt.Errorf("ASR input: %w", err) + } + lang := doubaospeech.Language(settings.ASR.Language) + switch lang { + case doubaospeech.LanguageZhCN, doubaospeech.LanguageEnUS, doubaospeech.LanguageJaJP, doubaospeech.LanguageKoKR: + default: + return CurrentBound{}, errors.New("ASR language is unsupported by selected SDK") + } + asrTimeout, err := currentTimeout(settings.ASR.TimeoutMS) + if err != nil { + return CurrentBound{}, fmt.Errorf("ASR timeout: %w", err) + } + asrRequest := doubaospeech.ASRV2Config{ + Format: doubaospeech.FormatPCMS16LE, SampleRate: sampleRate, + Channel: 1, Bits: 16, Language: lang, + Request: &doubaospeech.ASRV2RequestConfig{ModelName: settings.ASR.Model}, + } + if settings.ASR.Interim != nil { + asrRequest.Request.EnableNonstream = new(bool) + *asrRequest.Request.EnableNonstream = !*settings.ASR.Interim + if *settings.ASR.Interim { + asrRequest.ResultType = "full" + } else { + asrRequest.ResultType = "single" + } + asrRequest.Request.ResultType = asrRequest.ResultType + } + bound.ASR = CurrentASR{Provider: asrProvider, Request: asrRequest, Timeout: asrTimeout} + + switch settings.Mode { + case "asr_only": + return bound, nil + case "full_ai": + if settings.LLM == nil || settings.TTS == nil || settings.Prompt == nil || settings.Conversation == nil { + return CurrentBound{}, errors.New("full AI settings are incomplete") + } + default: + return CurrentBound{}, errors.New("AI mode is unsupported") + } + llmProvider, err := currentProvider(providers, settings.LLM.ProviderRef, "llm", "openai_compatible") + if err != nil { + return CurrentBound{}, err + } + llmTimeout, err := currentTimeout(settings.LLM.TimeoutMS) + if err != nil { + return CurrentBound{}, fmt.Errorf("LLM timeout: %w", err) + } + bound.LLM = &CurrentLLM{Provider: llmProvider, Model: settings.LLM.Model, Temperature: settings.LLM.Temperature, MaxTokens: settings.LLM.MaxTokens, Timeout: llmTimeout} + ttsProvider, err := currentProvider(providers, settings.TTS.ProviderRef, "tts", "volcengine_tts") + if err != nil { + return CurrentBound{}, err + } + if settings.TTS.Format.Encoding != "pcm_s16le" || settings.TTS.Format.Channels != 1 { + return CurrentBound{}, errors.New("TTS output encoding/channel is unsupported by selected SDK") + } + ttsSampleRate, err := currentSampleRate(settings.TTS.Format.SampleRateHz) + if err != nil { + return CurrentBound{}, fmt.Errorf("TTS output: %w", err) + } + ttsRate := 0 + if settings.TTS.Speed != nil { + floatRate := (*settings.TTS.Speed - 1) * 100 + if floatRate < -50 || floatRate > 100 || math.Abs(floatRate-math.Round(floatRate)) > 1e-9 { + return CurrentBound{}, errors.New("TTS speed is unsupported by selected SDK") + } + ttsRate = int(math.Round(floatRate)) + } + ttsTimeout, err := currentTimeout(settings.TTS.TimeoutMS) + if err != nil { + return CurrentBound{}, fmt.Errorf("TTS timeout: %w", err) + } + bound.TTS = &CurrentTTS{Provider: ttsProvider, Request: doubaospeech.TTSV2Request{ + Speaker: settings.TTS.Voice, ResourceID: settings.TTS.Model, + Format: doubaospeech.FormatPCMS16LE, SampleRate: ttsSampleRate, SpeechRate: ttsRate, + }, Timeout: ttsTimeout} + bound.Prompt = settings.Prompt.Text + bound.Opening = settings.Conversation.Opening + bound.HangupKeywords = append([]string(nil), settings.Conversation.HangupKeywords...) + return bound, nil +} + +func currentProvider(providers map[string]configread.CurrentProvider, ref, role, adapter string) (configread.CurrentProvider, error) { + p, found := providers[ref] + if !found || ref == "" || p.ProviderRef != ref || !p.Enabled || p.Role != role || p.Adapter != adapter || p.Credential == "" { + return configread.CurrentProvider{}, fmt.Errorf("%s provider is missing, disabled, or incompatible", role) + } + u, err := url.Parse(p.Endpoint) + if err != nil || u.Host == "" || u.User != nil || (u.Scheme != "https" && u.Scheme != "http" && u.Scheme != "wss" && u.Scheme != "ws") { + return configread.CurrentProvider{}, fmt.Errorf("%s provider endpoint is invalid", role) + } + return p, nil +} + +func currentSampleRate(rate int) (doubaospeech.SampleRate, error) { + switch rate { + case 8000, 16000, 22050, 24000, 32000, 44100, 48000: + return doubaospeech.SampleRate(rate), nil + default: + return 0, errors.New("sample rate is unsupported by selected SDK") + } +} + +func currentTimeout(ms *int64) (time.Duration, error) { + if ms == nil { + return 0, nil + } + if *ms <= 0 || *ms > int64(math.MaxInt64/int64(time.Millisecond)) { + return 0, errors.New("timeout is outside representable range") + } + return time.Duration(*ms) * time.Millisecond, nil +} diff --git a/internal/ai/current_pipeline.go b/internal/ai/current_pipeline.go new file mode 100644 index 0000000..f546285 --- /dev/null +++ b/internal/ai/current_pipeline.go @@ -0,0 +1,206 @@ +package ai + +import ( + "context" + "errors" + "fmt" + "net/url" + "strings" + "time" + + "github.com/GizClaw/doubao-speech-go" + "github.com/openai/openai-go/v3" + "github.com/openai/openai-go/v3/option" + "github.com/openai/openai-go/v3/packages/param" + "github.com/openai/openai-go/v3/shared" +) + +// CurrentCall owns the keyword action for exactly one call. A failed or +// uncertain hangup is never attempted again for that call. +type CurrentCall struct { + bound CurrentBound + keyword *KeywordHangup +} + +func NewCurrentCall(bound CurrentBound, hangup func(context.Context) error) (*CurrentCall, error) { + keyword, err := NewKeywordHangup(bound.HangupKeywords, hangup) + if err != nil { + return nil, err + } + return &CurrentCall{bound: bound, keyword: keyword}, nil +} + +func (c *CurrentCall) HandleFinalASR(ctx context.Context, segment ASRSegment) (bool, error) { + if c == nil { + return false, errors.New("AI call is unavailable") + } + return c.keyword.Handle(ctx, segment) +} + +// RunTurn performs one utterance. No LLM or TTS request is sent after a +// keyword hangup, and only a final user recognition may enter the LLM. +func (c *CurrentCall) RunTurn(ctx context.Context, pcm16 []byte) (TurnResult, error) { + if c == nil { + return TurnResult{}, errors.New("AI call is unavailable") + } + text, err := c.bound.Recognize(ctx, pcm16) + if err != nil { + return TurnResult{}, fmt.Errorf("ASR failed: %w", err) + } + stopped, err := c.HandleFinalASR(ctx, ASRSegment{Source: "user", Text: text, Final: true}) + if err != nil { + return TurnResult{}, err + } + if stopped { + return TurnResult{Transcript: text, EndedByKeyword: true}, nil + } + if c.bound.Mode == "asr_only" { + return TurnResult{Transcript: text}, nil + } + reply, err := c.bound.Complete(ctx, text) + if err != nil { + return TurnResult{}, fmt.Errorf("LLM failed: %w", err) + } + audio, err := c.bound.Synthesize(ctx, reply) + if err != nil { + return TurnResult{}, fmt.Errorf("TTS failed: %w", err) + } + return TurnResult{Transcript: text, Reply: reply, AudioPCM16: audio}, nil +} + +func (b CurrentBound) Recognize(ctx context.Context, pcm16 []byte) (string, error) { + if len(pcm16) == 0 || len(pcm16)%2 != 0 { + return "", errors.New("ASR requires nonempty signed 16-bit PCM") + } + ctx, cancel := currentDeadline(ctx, b.ASR.Timeout) + defer cancel() + u, err := url.Parse(b.ASR.Provider.Endpoint) + if err != nil { + return "", errors.New("ASR endpoint is invalid") + } + if u.Scheme == "https" { + u.Scheme = "wss" + } else if u.Scheme == "http" { + u.Scheme = "ws" + } + client := doubaospeech.NewClient("", + doubaospeech.WithAPIKey(b.ASR.Provider.Credential), + doubaospeech.WithWebSocketURL(u.String()), + ) + request := b.ASR.Request + session, err := client.ASRV2.OpenStreamSession(ctx, &request) + if err != nil { + return "", err + } + defer session.Close() + if err := session.SendAudio(ctx, pcm16, true); err != nil { + return "", err + } + var final finalASRAccumulator + for result, recvErr := range session.Recv() { + if recvErr != nil { + return "", recvErr + } + if result == nil { + continue + } + final.Add(result.Text, result.IsFinal) + if result.IsFinal { + break + } + } + return final.Result() +} + +func (b CurrentBound) Complete(ctx context.Context, finalUserText string) (string, error) { + if b.Mode != "full_ai" || b.LLM == nil { + return "", errors.New("ASR-only mode does not call LLM") + } + if strings.TrimSpace(finalUserText) == "" { + return "", errors.New("LLM requires final user text") + } + ctx, cancel := currentDeadline(ctx, b.LLM.Timeout) + defer cancel() + client := openai.NewClient( + option.WithAPIKey(b.LLM.Provider.Credential), + option.WithBaseURL(strings.TrimRight(b.LLM.Provider.Endpoint, "/")), + option.WithMaxRetries(0), + ) + messages := []openai.ChatCompletionMessageParamUnion{openai.SystemMessage(b.Prompt), openai.UserMessage(finalUserText)} + params := openai.ChatCompletionNewParams{Model: shared.ChatModel(b.LLM.Model), Messages: messages} + if b.LLM.Temperature != nil { + params.Temperature = param.NewOpt(*b.LLM.Temperature) + } + if b.LLM.MaxTokens != nil { + params.MaxTokens = param.NewOpt(*b.LLM.MaxTokens) + } + result, err := client.Chat.Completions.New(ctx, params) + if err != nil { + return "", err + } + if len(result.Choices) == 0 || strings.TrimSpace(result.Choices[0].Message.Content) == "" { + return "", errors.New("LLM returned no reply") + } + return result.Choices[0].Message.Content, nil +} + +func (b CurrentBound) Synthesize(ctx context.Context, text string) ([]byte, error) { + if b.Mode != "full_ai" || b.TTS == nil { + return nil, errors.New("ASR-only mode does not call TTS") + } + if strings.TrimSpace(text) == "" { + return nil, errors.New("TTS requires nonempty reply") + } + ctx, cancel := currentDeadline(ctx, b.TTS.Timeout) + defer cancel() + request := b.TTS.Request // per-call copy: concurrent calls never share mutable SDK parameters + request.Text = text + client := doubaospeech.NewClient("", + doubaospeech.WithAPIKey(b.TTS.Provider.Credential), + doubaospeech.WithBaseURL(b.TTS.Provider.Endpoint), + ) + var audio []byte + completed := false + for chunk, err := range client.TTSV2.Stream(ctx, &request) { + if err != nil { + return nil, err + } + if chunk == nil { + return nil, errors.New("TTS returned an empty stream chunk") + } + audio = append(audio, chunk.Audio...) + if chunk.IsLast { + completed = true + break + } + } + if !completed || len(audio) == 0 || len(audio)%2 != 0 { + return nil, errors.New("TTS stream ended without complete PCM16 audio") + } + return audio, nil +} + +func currentDeadline(ctx context.Context, timeout time.Duration) (context.Context, context.CancelFunc) { + if timeout > 0 { + return context.WithTimeout(ctx, timeout) + } + return context.WithCancel(ctx) +} + +type finalASRAccumulator struct { + text string + final bool +} + +func (a *finalASRAccumulator) Add(text string, final bool) { + if !a.final && final { + a.text, a.final = text, true + } +} + +func (a *finalASRAccumulator) Result() (string, error) { + if !a.final || strings.TrimSpace(a.text) == "" { + return "", errors.New("ASR returned no final user transcript") + } + return a.text, nil +} diff --git a/internal/ai/current_pipeline_test.go b/internal/ai/current_pipeline_test.go new file mode 100644 index 0000000..5b72a9e --- /dev/null +++ b/internal/ai/current_pipeline_test.go @@ -0,0 +1,135 @@ +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) + } +} diff --git a/internal/ai/current_test.go b/internal/ai/current_test.go new file mode 100644 index 0000000..814551f --- /dev/null +++ b/internal/ai/current_test.go @@ -0,0 +1,157 @@ +package ai + +import ( + "encoding/json" + "os" + "strings" + "testing" + "time" + + "git.ipao.vip/rogee/go-sip/internal/configread" + doubaospeech "github.com/GizClaw/doubao-speech-go" +) + +func currentFixture(t *testing.T, mode string) (configread.CurrentTask, map[string]configread.CurrentProvider) { + t.Helper() + name := "config-read-task-full.json" + if mode == "asr_only" { + name = "config-read-task-asr.json" + } + raw, err := os.ReadFile("../../contracts/local/examples/" + name) + if err != nil { + t.Fatal(err) + } + var task configread.CurrentTask + if err := json.Unmarshal(raw, &task); err != nil { + t.Fatal(err) + } + raw, err = os.ReadFile("../../contracts/local/examples/config-read-providers.json") + if err != nil { + t.Fatal(err) + } + var list struct { + Providers []configread.CurrentProvider `json:"providers"` + } + if err := json.Unmarshal(raw, &list); err != nil { + t.Fatal(err) + } + providers := make(map[string]configread.CurrentProvider, len(list.Providers)) + for _, p := range list.Providers { + providers[p.ProviderRef] = p + } + return task, providers +} + +func changeCurrentAgent(t *testing.T, task configread.CurrentTask, change func(map[string]any)) configread.CurrentTask { + t.Helper() + var body map[string]any + if err := json.Unmarshal(task.Raw, &body); err != nil { + t.Fatal(err) + } + agent := body["agent"].(map[string]any) + change(agent) + raw, err := json.Marshal(body) + if err != nil { + t.Fatal(err) + } + var updated configread.CurrentTask + if err := json.Unmarshal(raw, &updated); err != nil { + t.Fatal(err) + } + return updated +} + +func TestBindCurrentFullAIUsesApprovedSDKFields(t *testing.T) { + task, providers := currentFixture(t, "full_ai") + bound, err := BindCurrent(task, providers) + if err != nil { + t.Fatal(err) + } + if bound.Mode != "full_ai" || bound.LLM == nil || bound.TTS == nil { + t.Fatalf("full mode needs all three providers: mode=%q LLM=%t TTS=%t", bound.Mode, bound.LLM != nil, bound.TTS != nil) + } + if bound.ASR.Provider.Credential != providers["asr-example"].Credential || bound.ASR.Request.Format != doubaospeech.FormatPCMS16LE || bound.ASR.Request.SampleRate != 16000 || bound.ASR.Request.Channel != 1 || bound.ASR.Request.Bits != 16 || bound.ASR.Request.Language != doubaospeech.LanguageZhCN || bound.ASR.Request.ResultType != "full" || bound.ASR.Timeout != 5*time.Second { + t.Fatal("ASR approved input/interim/credential/timeout not bound to SDK request") + } + if bound.LLM.Provider.Credential != providers["llm-example"].Credential || bound.LLM.Model != "example-chat" || bound.LLM.Temperature == nil || *bound.LLM.Temperature != 0 || bound.LLM.MaxTokens == nil || *bound.LLM.MaxTokens != 256 || bound.LLM.Timeout != 5*time.Second { + t.Fatal("LLM model/explicit zero/limit/credential/timeout not bound") + } + if bound.TTS.Provider.Credential != providers["tts-example"].Credential || bound.TTS.Request.ResourceID != "example-tts" || bound.TTS.Request.Speaker != "example-neutral" || bound.TTS.Request.Format != doubaospeech.FormatPCMS16LE || bound.TTS.Request.SampleRate != 16000 || bound.TTS.Request.SpeechRate != 0 || bound.TTS.Timeout != 5*time.Second { + t.Fatal("TTS model/voice/speed/format/credential/timeout not bound to SDK request") + } + if len(bound.HangupKeywords) != 1 || bound.HangupKeywords[0] != "不用了" || bound.Prompt != "Example only" || bound.Opening != "Example greeting" { + t.Fatal("immutable prompt and keyword behavior not bound") + } +} + +func TestBindCurrentASROnlyDoesNotBindOtherProviders(t *testing.T) { + task, providers := currentFixture(t, "asr_only") + delete(providers, "llm-example") + delete(providers, "tts-example") + bound, err := BindCurrent(task, providers) + if err != nil { + t.Fatal(err) + } + if bound.Mode != "asr_only" || bound.LLM != nil || bound.TTS != nil || bound.ASR.Request.ResultType != "single" { + t.Fatal("ASR-only mode must not inherit LLM/TTS configuration or interim results") + } +} + +func TestBindCurrentRejectsSDKUnsupportedTTSWithoutChangingSchema(t *testing.T) { + for _, tc := range []struct { + name string + edit func(map[string]any) + }{ + {"pcma", func(tts map[string]any) { tts["format"].(map[string]any)["encoding"] = "pcma" }}, + {"speed-below", func(tts map[string]any) { tts["speed"] = 0.25 }}, + {"speed-above", func(tts map[string]any) { tts["speed"] = 3.0 }}, + {"speed-unrepresentable", func(tts map[string]any) { tts["speed"] = 1.005 }}, + {"sample-rate", func(tts map[string]any) { tts["format"].(map[string]any)["sample_rate_hz"] = 12345 }}, + } { + t.Run(tc.name, func(t *testing.T) { + task, providers := currentFixture(t, "full_ai") + task = changeCurrentAgent(t, task, func(agent map[string]any) { tc.edit(agent["tts"].(map[string]any)) }) + _, err := BindCurrent(task, providers) + if err == nil || !strings.Contains(err.Error(), "TTS") || strings.Contains(err.Error(), providers["tts-example"].Credential) { + t.Fatalf("expected explicit non-secret TTS capability error, got %v", err) + } + }) + } +} + +func TestBindCurrentRejectsUnauthorizedProviderBeforeCall(t *testing.T) { + for _, tc := range []struct { + name string + mutate func(map[string]configread.CurrentProvider) + }{ + {"missing", func(ps map[string]configread.CurrentProvider) { delete(ps, "asr-example") }}, + {"disabled", func(ps map[string]configread.CurrentProvider) { + p := ps["asr-example"] + p.Enabled = false + ps[p.ProviderRef] = p + }}, + {"wrong-role", func(ps map[string]configread.CurrentProvider) { + p := ps["asr-example"] + p.Role = "tts" + ps[p.ProviderRef] = p + }}, + {"wrong-adapter", func(ps map[string]configread.CurrentProvider) { + p := ps["asr-example"] + p.Adapter = "unknown" + ps[p.ProviderRef] = p + }}, + {"missing-credential", func(ps map[string]configread.CurrentProvider) { + p := ps["asr-example"] + p.Credential = "" + ps[p.ProviderRef] = p + }}, + } { + t.Run(tc.name, func(t *testing.T) { + task, providers := currentFixture(t, "full_ai") + tc.mutate(providers) + if _, err := BindCurrent(task, providers); err == nil { + t.Fatal("unavailable AI provider cannot authorize execution") + } + }) + } +} diff --git a/internal/ai/keyword_hangup.go b/internal/ai/keyword_hangup.go new file mode 100644 index 0000000..f9d14ee --- /dev/null +++ b/internal/ai/keyword_hangup.go @@ -0,0 +1,77 @@ +package ai + +import ( + "context" + "errors" + "fmt" + "strings" + "sync" + "unicode/utf8" +) + +// ASRSegment carries the source and finality of one text observation. +// Assistant speech or interim ASR never authorizes keyword termination. +type ASRSegment struct { + Source string + Text string + Final bool +} + +// KeywordHangup is private to one call. Once a hangup is attempted its +// outcome may be unknown, so repeated ASR notifications cannot hang up again. +type KeywordHangup struct { + keywords []string + hangup func(context.Context) error + mu sync.Mutex + requested bool +} + +func NewKeywordHangup(keywords []string, hangup func(context.Context) error) (*KeywordHangup, error) { + if hangup == nil { + return nil, errors.New("keyword termination requires a real hangup action") + } + frozen := append([]string(nil), keywords...) + for _, keyword := range frozen { + if keyword == "" || !utf8.ValidString(keyword) { + return nil, errors.New("keyword termination requires nonempty UTF-8 literals") + } + } + return &KeywordHangup{keywords: frozen, hangup: hangup}, nil +} + +func (k *KeywordHangup) Handle(ctx context.Context, segment ASRSegment) (bool, error) { + if k == nil { + return false, errors.New("keyword terminator is unavailable") + } + if segment.Source != "user" && segment.Source != "assistant" { + return false, errors.New("ASR text source is unknown") + } + if !segment.Final || segment.Source != "user" { + return false, nil + } + if !utf8.ValidString(segment.Text) { + return false, errors.New("final user ASR text is not UTF-8") + } + k.mu.Lock() + if k.requested { + k.mu.Unlock() + return false, nil + } + matched := false + for _, keyword := range k.keywords { + if strings.Contains(segment.Text, keyword) { + matched = true + break + } + } + if !matched { + k.mu.Unlock() + return false, nil + } + k.requested = true + k.mu.Unlock() + if err := k.hangup(ctx); err != nil { + return true, fmt.Errorf("keyword-triggered hangup outcome unknown: %w", err) + } + return true, nil +} diff --git a/internal/ai/keyword_hangup_test.go b/internal/ai/keyword_hangup_test.go new file mode 100644 index 0000000..47899d4 --- /dev/null +++ b/internal/ai/keyword_hangup_test.go @@ -0,0 +1,99 @@ +package ai + +import ( + "context" + "errors" + "strings" + "sync" + "sync/atomic" + "testing" +) + +func TestKeywordHangupOnlyFinalUserLiteralAndOnlyOnce(t *testing.T) { + var calls atomic.Int64 + stopper, err := NewKeywordHangup([]string{"停止通话", "STOP"}, func(context.Context) error { calls.Add(1); return nil }) + if err != nil { + t.Fatal(err) + } + cases := []struct { + event ASRSegment + expected bool + }{ + {ASRSegment{Source: "user", Text: "请停止通话", Final: false}, false}, + {ASRSegment{Source: "assistant", Text: "请停止通话", Final: true}, false}, + {ASRSegment{Source: "user", Text: "请停通话", Final: true}, false}, + {ASRSegment{Source: "user", Text: "stop", Final: true}, false}, + {ASRSegment{Source: "user", Text: "现在停止通话", Final: true}, true}, + {ASRSegment{Source: "user", Text: "现在停止通话", Final: true}, false}, + {ASRSegment{Source: "user", Text: "STOP", Final: true}, false}, + } + for i, tc := range cases { + triggered, err := stopper.Handle(context.Background(), tc.event) + if err != nil || triggered != tc.expected { + t.Fatalf("event %d: triggered=%v want=%v err=%v", i, triggered, tc.expected, err) + } + } + if calls.Load() != 1 { + t.Fatalf("duplicate or interim hangup: calls=%d", calls.Load()) + } +} + +func TestKeywordHangupFailureIsVisibleButUnknownActionNeverRepeated(t *testing.T) { + var calls atomic.Int64 + stopper, err := NewKeywordHangup([]string{"终止"}, func(context.Context) error { calls.Add(1); return errors.New("mock ARI hangup outcome unknown") }) + if err != nil { + t.Fatal(err) + } + result := ASRSegment{Source: "user", Text: "请终止", Final: true} + if triggered, err := stopper.Handle(context.Background(), result); !triggered || err == nil || !strings.Contains(err.Error(), "outcome unknown") { + t.Fatalf("hangup failure hidden: %v %v", triggered, err) + } + if triggered, err := stopper.Handle(context.Background(), result); triggered || err != nil || calls.Load() != 1 { + t.Fatalf("unknown hangup retried automatically: %v %v calls=%d", triggered, err, calls.Load()) + } +} + +func TestKeywordHangupConcurrentDuplicateFinalTriggersOnce(t *testing.T) { + var calls atomic.Int64 + stopper, err := NewKeywordHangup([]string{"停机"}, func(context.Context) error { calls.Add(1); return nil }) + if err != nil { + t.Fatal(err) + } + var wg sync.WaitGroup + for i := 0; i < 32; i++ { + wg.Add(1) + go func() { + defer wg.Done() + _, _ = stopper.Handle(context.Background(), ASRSegment{Source: "user", Text: "停机", Final: true}) + }() + } + wg.Wait() + if calls.Load() != 1 { + t.Fatalf("concurrent duplicate final ASR hung up %d times", calls.Load()) + } +} + +func TestKeywordHangupRejectsInvalidConfigurationAndUnknownSource(t *testing.T) { + callback := func(context.Context) error { return nil } + if _, err := NewKeywordHangup([]string{""}, callback); err == nil { + t.Fatal("accepted empty keyword matching every transcript") + } + if _, err := NewKeywordHangup([]string{string([]byte{0xff})}, callback); err == nil { + t.Fatal("accepted invalid UTF-8 keyword") + } + if _, err := NewKeywordHangup([]string{"终止"}, nil); err == nil { + t.Fatal("accepted missing actual hangup action") + } + words := []string{"终止"} + stopper, err := NewKeywordHangup(words, callback) + if err != nil { + t.Fatal(err) + } + words[0] = "改动" + if triggered, err := stopper.Handle(context.Background(), ASRSegment{Source: "user", Text: "终止", Final: true}); err != nil || !triggered { + t.Fatalf("keyword list was mutable after approval: %v %v", triggered, err) + } + if _, err := stopper.Handle(context.Background(), ASRSegment{Source: "other", Text: "终止", Final: true}); err == nil { + t.Fatal("unknown text source was silently accepted") + } +} diff --git a/internal/ai/provider_pipeline.go b/internal/ai/provider_pipeline.go index 67b2c8c..9b4c955 100644 --- a/internal/ai/provider_pipeline.go +++ b/internal/ai/provider_pipeline.go @@ -120,11 +120,12 @@ type providerSnapshotConfig struct { // TurnResult contains only bounded facts and audio bytes needed by the caller. // Callers must persist a hash/length, not the transcript or prompt. type TurnResult struct { - Transcript string - Reply string - AudioPCM16 []byte - InvalidCall bool - InvalidReason string + Transcript string + Reply string + AudioPCM16 []byte + EndedByKeyword bool + InvalidCall bool + InvalidReason string } // Synthesize turns a bounded reply into signed 16-bit little-endian 16kHz