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

359 lines
14 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}