287 lines
11 KiB
Go
287 lines
11 KiB
Go
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
|
|
PromptMaxBytes int
|
|
AllowedVariables []string
|
|
Opening string
|
|
HangupKeywords []string
|
|
Conversation CurrentConversation
|
|
}
|
|
|
|
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
|
|
}
|
|
|
|
// CurrentConversation preserves the task's explicit dialogue limits. An
|
|
// absent positive limit remains zero; an explicit false remains non-nil.
|
|
type CurrentConversation 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_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"`
|
|
AllowedVariables []string `json:"allowed_variables"`
|
|
MaxBytes *int `json:"max_bytes"`
|
|
} `json:"prompt"`
|
|
Conversation *struct {
|
|
Opening string `json:"opening"`
|
|
HangupKeywords []string `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"`
|
|
}
|
|
|
|
// 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)
|
|
}
|
|
if sampleRate != doubaospeech.SampleRate(16000) {
|
|
return CurrentBound{}, errors.New("ASR media requires 16000 Hz PCM16")
|
|
}
|
|
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)
|
|
}
|
|
if ttsSampleRate != doubaospeech.SampleRate(16000) {
|
|
return CurrentBound{}, errors.New("TTS media requires 16000 Hz PCM16")
|
|
}
|
|
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.AllowedVariables = append([]string(nil), settings.Prompt.AllowedVariables...)
|
|
if settings.Prompt.MaxBytes != nil {
|
|
bound.PromptMaxBytes = *settings.Prompt.MaxBytes
|
|
if len(bound.Prompt) > bound.PromptMaxBytes {
|
|
return CurrentBound{}, errors.New("prompt exceeds its approved byte limit")
|
|
}
|
|
}
|
|
if settings.Conversation.AllowInterrupt != nil && *settings.Conversation.AllowInterrupt {
|
|
return CurrentBound{}, errors.New("conversation interrupt is unsupported by current media controller")
|
|
}
|
|
silence, err := currentTimeout(settings.Conversation.SilenceTimeoutMS)
|
|
if err != nil {
|
|
return CurrentBound{}, fmt.Errorf("conversation silence timeout: %w", err)
|
|
}
|
|
maxDuration, err := currentTimeout(settings.Conversation.MaxDurationMS)
|
|
if err != nil {
|
|
return CurrentBound{}, fmt.Errorf("conversation duration: %w", err)
|
|
}
|
|
bound.Conversation = CurrentConversation{
|
|
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 = 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
|
|
}
|