359 lines
14 KiB
Go
359 lines
14 KiB
Go
package ai
|
||
|
||
import (
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"math"
|
||
"net/url"
|
||
"strings"
|
||
"time"
|
||
|
||
"git.ipao.vip/rogee/go-sip/internal/configread"
|
||
"git.ipao.vip/rogee/go-sip/internal/contract"
|
||
doubaospeech "github.com/GizClaw/doubao-speech-go"
|
||
)
|
||
|
||
// Binding 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 Binding struct {
|
||
Mode string
|
||
ASR ASRConfig
|
||
LLM *LLMConfig
|
||
TTS *TTSConfig
|
||
Prompt string
|
||
PromptMaxBytes int
|
||
AllowedVariables []string
|
||
Opening string
|
||
HangupKeywords []HangupKeyword
|
||
Conversation ConversationConfig
|
||
}
|
||
|
||
type ASRConfig struct {
|
||
Provider configread.Provider
|
||
Request doubaospeech.ASRV2Config
|
||
Model string
|
||
Language string
|
||
Timeout time.Duration
|
||
}
|
||
|
||
type LLMConfig struct {
|
||
Provider configread.Provider
|
||
Model string
|
||
Temperature *float64
|
||
MaxTokens *int64
|
||
Timeout time.Duration
|
||
}
|
||
|
||
type TTSConfig struct {
|
||
Provider configread.Provider
|
||
Protocol string
|
||
Model string
|
||
Voice string
|
||
LanguageType string
|
||
Speed float64
|
||
SampleRate int
|
||
Timeout time.Duration
|
||
}
|
||
|
||
// ConversationConfig preserves the task's explicit dialogue limits. An
|
||
// absent positive limit remains zero; an explicit false remains non-nil.
|
||
type ConversationConfig struct {
|
||
AllowInterrupt *bool
|
||
SilenceTimeout time.Duration
|
||
MaxDuration time.Duration
|
||
MaxTurns int
|
||
SentenceMaxChars int
|
||
MaxPendingAudioChunks int
|
||
}
|
||
|
||
type currentAgentSettings struct {
|
||
Mode string `json:"mode"`
|
||
ASR struct {
|
||
ProviderRef string `json:"provider_id"`
|
||
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_id"`
|
||
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_id"`
|
||
Protocol string `json:"protocol"`
|
||
Model string `json:"model"`
|
||
Voice string `json:"voice"`
|
||
LanguageType string `json:"language_type"`
|
||
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"`
|
||
AllowedVariables []string `json:"allowed_variables"`
|
||
MaxBytes *int `json:"max_bytes"`
|
||
} `json:"prompt"`
|
||
Conversation *struct {
|
||
Opening string `json:"opening"`
|
||
HangupKeywords []HangupKeyword `json:"hangup_keywords"`
|
||
AllowInterrupt *bool `json:"allow_interrupt"`
|
||
SilenceTimeoutMS *int64 `json:"silence_timeout_ms"`
|
||
MaxDurationMS *int64 `json:"max_duration_ms"`
|
||
MaxTurns int `json:"max_turns"`
|
||
SentenceMaxChars int `json:"sentence_max_chars"`
|
||
MaxPendingAudioChunks int `json:"max_pending_audio_chunks"`
|
||
} `json:"conversation"`
|
||
}
|
||
|
||
// Bind rejects schema-valid settings which the approved provider cannot
|
||
// express; neither defaults nor lossy format/speed conversions are permitted.
|
||
func Bind(task configread.Task, providers map[string]configread.Provider) (Binding, error) {
|
||
if len(task.Raw) == 0 {
|
||
return Binding{}, errors.New("approved task snapshot is missing")
|
||
}
|
||
if err := contract.ValidateCurrent("config-read", task.Raw); err != nil {
|
||
return Binding{}, fmt.Errorf("task snapshot: %w", err)
|
||
}
|
||
var frozen configread.Task
|
||
if err := json.Unmarshal(task.Raw, &frozen); err != nil {
|
||
return Binding{}, 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 Binding{}, errors.New("approved task snapshot identity mismatch")
|
||
}
|
||
var settings currentAgentSettings
|
||
if err := json.Unmarshal(frozen.Agent.Raw, &settings); err != nil {
|
||
return Binding{}, fmt.Errorf("decode immutable AI settings: %w", err)
|
||
}
|
||
bound := Binding{Mode: settings.Mode}
|
||
asrProvider, err := currentProvider(providers, settings.ASR.ProviderRef, "asr", "")
|
||
if err != nil {
|
||
return Binding{}, err
|
||
}
|
||
if settings.ASR.Input.Encoding != "pcm_s16le" || settings.ASR.Input.Channels != 1 || settings.ASR.Input.SampleWidthBytes != 2 {
|
||
return Binding{}, errors.New("ASR input format is unsupported by selected SDK")
|
||
}
|
||
sampleRate, err := currentSampleRate(settings.ASR.Input.SampleRateHz)
|
||
if err != nil {
|
||
return Binding{}, fmt.Errorf("ASR input: %w", err)
|
||
}
|
||
lang := doubaospeech.Language(settings.ASR.Language)
|
||
if asrProvider.Code == "ali_bailian" {
|
||
if settings.ASR.Interim != nil {
|
||
return Binding{}, errors.New("Bailian task ASR protocol cannot express interim selection")
|
||
}
|
||
if strings.TrimSpace(settings.ASR.Model) == "" || (sampleRate != 8000 && sampleRate != 16000) {
|
||
return Binding{}, errors.New("Bailian task ASR requires an explicit model and 8000/16000 Hz PCM16")
|
||
}
|
||
if _, err := bailianLanguageHint(settings.ASR.Language); err != nil {
|
||
return Binding{}, fmt.Errorf("ASR language: %w", err)
|
||
}
|
||
} else {
|
||
if sampleRate != 16000 {
|
||
return Binding{}, errors.New("ASR media requires 16000 Hz PCM16")
|
||
}
|
||
switch lang {
|
||
case doubaospeech.LanguageZhCN, doubaospeech.LanguageEnUS, doubaospeech.LanguageJaJP, doubaospeech.LanguageKoKR:
|
||
default:
|
||
return Binding{}, errors.New("ASR language is unsupported by selected SDK")
|
||
}
|
||
}
|
||
asrTimeout, err := currentTimeout(settings.ASR.TimeoutMS)
|
||
if err != nil {
|
||
return Binding{}, 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 = ASRConfig{Provider: asrProvider, Request: asrRequest, Model: settings.ASR.Model, Language: settings.ASR.Language, 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 Binding{}, errors.New("full AI settings are incomplete")
|
||
}
|
||
default:
|
||
return Binding{}, errors.New("AI mode is unsupported")
|
||
}
|
||
llmProvider, err := currentProvider(providers, settings.LLM.ProviderRef, "llm", "")
|
||
if err != nil {
|
||
return Binding{}, err
|
||
}
|
||
llmTimeout, err := currentTimeout(settings.LLM.TimeoutMS)
|
||
if err != nil {
|
||
return Binding{}, fmt.Errorf("LLM timeout: %w", err)
|
||
}
|
||
bound.LLM = &LLMConfig{Provider: llmProvider, Model: settings.LLM.Model, Temperature: settings.LLM.Temperature, MaxTokens: settings.LLM.MaxTokens, Timeout: llmTimeout}
|
||
ttsProvider, err := currentProvider(providers, settings.TTS.ProviderRef, "tts", settings.TTS.Protocol)
|
||
if err != nil {
|
||
return Binding{}, err
|
||
}
|
||
if settings.TTS.Format.Encoding != "pcm_s16le" || settings.TTS.Format.Channels != 1 {
|
||
return Binding{}, errors.New("TTS output encoding/channel is unsupported by selected SDK")
|
||
}
|
||
ttsSampleRate, err := currentSampleRate(settings.TTS.Format.SampleRateHz)
|
||
if err != nil {
|
||
return Binding{}, fmt.Errorf("TTS output: %w", err)
|
||
}
|
||
if ttsSampleRate != doubaospeech.SampleRate(16000) {
|
||
return Binding{}, errors.New("TTS media requires 16000 Hz PCM16")
|
||
}
|
||
if strings.TrimSpace(settings.TTS.Model) == "" || strings.TrimSpace(settings.TTS.Voice) == "" || strings.TrimSpace(settings.TTS.LanguageType) == "" || settings.TTS.Speed == nil {
|
||
return Binding{}, errors.New("TTS model, voice, language and speed must be explicit")
|
||
}
|
||
switch settings.TTS.Protocol {
|
||
case TTSProtocolDashScopeHTTP:
|
||
if *settings.TTS.Speed != 1 {
|
||
return Binding{}, errors.New("TTS HTTP protocol cannot express the requested speed")
|
||
}
|
||
case TTSProtocolDashScopeTask:
|
||
if *settings.TTS.Speed < 0.5 || *settings.TTS.Speed > 2 {
|
||
return Binding{}, errors.New("TTS task protocol speed is outside 0.5–2")
|
||
}
|
||
if _, err := bailianTTSLanguageHints(settings.TTS.LanguageType); err != nil {
|
||
return Binding{}, fmt.Errorf("TTS language: %w", err)
|
||
}
|
||
default:
|
||
return Binding{}, errors.New("TTS protocol is unsupported")
|
||
}
|
||
ttsTimeout, err := currentTimeout(settings.TTS.TimeoutMS)
|
||
if err != nil {
|
||
return Binding{}, fmt.Errorf("TTS timeout: %w", err)
|
||
}
|
||
bound.TTS = &TTSConfig{
|
||
Provider: ttsProvider, Protocol: settings.TTS.Protocol, Model: settings.TTS.Model, Voice: settings.TTS.Voice, LanguageType: settings.TTS.LanguageType,
|
||
Speed: *settings.TTS.Speed, SampleRate: int(ttsSampleRate), Timeout: ttsTimeout,
|
||
}
|
||
bound.Prompt = settings.Prompt.Text
|
||
bound.AllowedVariables = append([]string(nil), settings.Prompt.AllowedVariables...)
|
||
if settings.Prompt.MaxBytes != nil {
|
||
bound.PromptMaxBytes = *settings.Prompt.MaxBytes
|
||
if len(bound.Prompt) > bound.PromptMaxBytes {
|
||
return Binding{}, errors.New("prompt exceeds its approved byte limit")
|
||
}
|
||
}
|
||
if settings.Conversation.AllowInterrupt != nil && *settings.Conversation.AllowInterrupt {
|
||
return Binding{}, errors.New("conversation interrupt is unsupported by current media controller")
|
||
}
|
||
silence, err := currentTimeout(settings.Conversation.SilenceTimeoutMS)
|
||
if err != nil {
|
||
return Binding{}, fmt.Errorf("conversation silence timeout: %w", err)
|
||
}
|
||
maxDuration, err := currentTimeout(settings.Conversation.MaxDurationMS)
|
||
if err != nil {
|
||
return Binding{}, fmt.Errorf("conversation duration: %w", err)
|
||
}
|
||
bound.Conversation = ConversationConfig{
|
||
AllowInterrupt: settings.Conversation.AllowInterrupt,
|
||
SilenceTimeout: silence, MaxDuration: maxDuration,
|
||
MaxTurns: settings.Conversation.MaxTurns,
|
||
SentenceMaxChars: settings.Conversation.SentenceMaxChars,
|
||
MaxPendingAudioChunks: settings.Conversation.MaxPendingAudioChunks,
|
||
}
|
||
bound.Opening = settings.Conversation.Opening
|
||
bound.HangupKeywords = cloneHangupKeywords(settings.Conversation.HangupKeywords)
|
||
if err := validateHangupKeywords(bound.HangupKeywords); err != nil {
|
||
return Binding{}, err
|
||
}
|
||
return bound, nil
|
||
}
|
||
|
||
func currentProvider(providers map[string]configread.Provider, ref, role, protocol string) (configread.Provider, error) {
|
||
p, found := providers[ref]
|
||
if !found || ref == "" || p.ProviderRef != ref || p.Credential == "" {
|
||
return configread.Provider{}, fmt.Errorf("%s provider connection is missing or incomplete", role)
|
||
}
|
||
if len(p.ExtraConfig) != 0 {
|
||
var extra map[string]json.RawMessage
|
||
if err := json.Unmarshal(p.ExtraConfig, &extra); err != nil || len(extra) != 0 {
|
||
return configread.Provider{}, fmt.Errorf("%s provider has unsupported connection parameters", role)
|
||
}
|
||
}
|
||
switch role {
|
||
case "asr":
|
||
if p.Code != "volcengine" && p.Code != "ali_bailian" {
|
||
return configread.Provider{}, errors.New("ASR provider is unsupported")
|
||
}
|
||
if p.Code == "ali_bailian" {
|
||
p.Endpoint = p.WSEndpoint
|
||
}
|
||
case "llm":
|
||
if p.Code != "openai_compatible" && p.Code != "ali_bailian" {
|
||
return configread.Provider{}, errors.New("LLM provider is unsupported")
|
||
}
|
||
case "tts":
|
||
if p.Code != "ali_bailian" {
|
||
return configread.Provider{}, errors.New("TTS provider is unsupported")
|
||
}
|
||
switch protocol {
|
||
case TTSProtocolDashScopeTask:
|
||
p.Endpoint = p.WSEndpoint
|
||
case TTSProtocolDashScopeHTTP:
|
||
default:
|
||
return configread.Provider{}, errors.New("TTS protocol is unsupported")
|
||
}
|
||
default:
|
||
return configread.Provider{}, errors.New("AI purpose is unsupported")
|
||
}
|
||
u, err := url.Parse(p.Endpoint)
|
||
if err != nil || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" || (u.Scheme != "https" && u.Scheme != "http" && u.Scheme != "wss" && u.Scheme != "ws") || strings.ContainsAny(p.Credential, "\r\n") {
|
||
return configread.Provider{}, fmt.Errorf("%s provider endpoint or credential is invalid", role)
|
||
}
|
||
if (role == "asr" && p.Code == "ali_bailian" || role == "tts" && protocol == TTSProtocolDashScopeTask) && u.Scheme != "wss" && u.Scheme != "ws" {
|
||
return configread.Provider{}, errors.New("selected speech model requires a WebSocket connection")
|
||
}
|
||
if (role == "llm" || role == "tts" && protocol == TTSProtocolDashScopeHTTP) && u.Scheme != "http" && u.Scheme != "https" {
|
||
return configread.Provider{}, errors.New("selected protocol requires an HTTP connection")
|
||
}
|
||
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
|
||
}
|