Bind approved AI snapshots to provider SDK calls

This commit is contained in:
2026-09-29 20:41:00 +08:00
parent c350c483a5
commit 5064b7765a
7 changed files with 913 additions and 5 deletions
+233
View File
@@ -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
}
+206
View File
@@ -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
}
+135
View File
@@ -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)
}
}
+157
View File
@@ -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")
}
})
}
}
+77
View File
@@ -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
}
+99
View File
@@ -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")
}
}
+6 -5
View File
@@ -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