Files
go-sip/internal/ai/current.go
T

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
}