Bind approved AI snapshots to provider SDK calls
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user