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

200 lines
7.0 KiB
Go

package ai
import (
"encoding/json"
"errors"
"fmt"
"net/url"
"strings"
"time"
"git.ipao.vip/rogee/go-sip/internal/configread"
"git.ipao.vip/rogee/go-sip/internal/contract"
)
// Binding freezes the task, credentials and raw scalar parameters before any side effect.
// Errors never contain credentials, prompt text or parameter values.
type Binding struct {
Mode string
ASR ASRConfig
LLM *LLMConfig
TTS *TTSConfig
Prompt string
AllowedVariables []string
Opening string
HangupKeywords []HangupKeyword
Conversation ConversationConfig
}
type ASRConfig struct {
Provider configread.Provider
Model string
Params map[string]json.RawMessage
}
type LLMConfig struct {
Provider configread.Provider
Model string
Params map[string]json.RawMessage
}
type TTSConfig struct {
Provider configread.Provider
Protocol string
Model string
Voice string
Params map[string]json.RawMessage
}
type ConversationConfig struct {
AllowInterrupt *bool
SilenceTimeout time.Duration
MaxDuration time.Duration
MaxTurns int
SentenceMaxChars int
MaxPendingAudioChunks int
}
type modelSettings struct {
ProviderRef string `json:"provider_ref"`
Model string `json:"model"`
Params map[string]json.RawMessage `json:"params"`
Voice string `json:"voice,omitempty"`
}
type currentAgentSettings struct {
Mode string `json:"mode"`
ASR modelSettings `json:"asr"`
LLM *modelSettings `json:"llm"`
TTS *modelSettings `json:"tts"`
Prompt *struct {
Text string `json:"text"`
AllowedVariables []string `json:"allowed_variables"`
} `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"`
}
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{}, errors.New("decode task snapshot")
}
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 s currentAgentSettings
if err := json.Unmarshal(frozen.Agent.Raw, &s); err != nil {
return Binding{}, errors.New("decode immutable AI settings")
}
b := Binding{Mode: s.Mode}
p, err := currentProvider(providers, s.ASR.ProviderRef, "asr")
if err != nil {
return Binding{}, err
}
if err := reservedParams(s.ASR.Params, "reqid", "sequence", "model_name"); err != nil {
return Binding{}, err
}
b.ASR = ASRConfig{Provider: p, Model: s.ASR.Model, Params: s.ASR.Params}
if s.Mode == "asr_only" {
return b, nil
}
if s.Mode != "full_ai" || s.LLM == nil || s.TTS == nil || s.Prompt == nil {
return Binding{}, errors.New("full AI settings are incomplete")
}
p, err = currentProvider(providers, s.LLM.ProviderRef, "llm")
if err != nil {
return Binding{}, err
}
if err := reservedParams(s.LLM.Params, "model", "messages"); err != nil {
return Binding{}, err
}
b.LLM = &LLMConfig{Provider: p, Model: s.LLM.Model, Params: s.LLM.Params}
p, err = currentProvider(providers, s.TTS.ProviderRef, "tts")
if err != nil {
return Binding{}, err
}
if err := reservedParams(s.TTS.Params, "voice"); err != nil {
return Binding{}, err
}
b.TTS = &TTSConfig{Provider: p, Protocol: TTSProtocolDashScopeTask, Model: s.TTS.Model, Voice: s.TTS.Voice, Params: s.TTS.Params}
b.Prompt = s.Prompt.Text
b.AllowedVariables = append([]string(nil), s.Prompt.AllowedVariables...)
if s.Conversation != nil {
c := s.Conversation
b.Opening = c.Opening
b.HangupKeywords = cloneHangupKeywords(c.HangupKeywords)
if err := validateHangupKeywords(b.HangupKeywords); err != nil {
return Binding{}, err
}
b.Conversation = ConversationConfig{AllowInterrupt: c.AllowInterrupt, MaxTurns: c.MaxTurns, SentenceMaxChars: c.SentenceMaxChars, MaxPendingAudioChunks: c.MaxPendingAudioChunks}
if c.AllowInterrupt != nil && *c.AllowInterrupt {
return Binding{}, errors.New("allow_interrupt=true is unsupported by current media controller")
}
if c.SilenceTimeoutMS != nil {
b.Conversation.SilenceTimeout = time.Duration(*c.SilenceTimeoutMS) * time.Millisecond
}
if c.MaxDurationMS != nil {
b.Conversation.MaxDuration = time.Duration(*c.MaxDurationMS) * time.Millisecond
}
}
return b, nil
}
// Protocol identity belongs to the task/session, not vendor business parameters.
// Reject collisions explicitly instead of overwriting or dropping supplied params.
func reservedParams(params map[string]json.RawMessage, keys ...string) error {
for _, key := range keys {
if _, found := params[key]; found {
return fmt.Errorf("params conflicts with protocol field %q", key)
}
}
return nil
}
func currentProvider(providers map[string]configread.Provider, ref, role string) (configread.Provider, error) {
p, ok := providers[ref]
if !ok || p.Code != ref || p.Credential == "" {
return configread.Provider{}, fmt.Errorf("%s provider connection is missing or incomplete", role)
}
switch role {
case "asr":
if p.Code != "volcengine" && p.Code != "ali_bailian" {
return configread.Provider{}, errors.New("ASR provider is unsupported")
}
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")
}
p.Endpoint = p.WSEndpoint
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 != "" || strings.ContainsAny(p.Credential, "\r\n") {
return configread.Provider{}, fmt.Errorf("%s provider endpoint or credential is invalid", role)
}
if role == "llm" {
if u.Scheme != "https" && u.Scheme != "http" {
return configread.Provider{}, errors.New("LLM requires an HTTP connection")
}
} else if u.Scheme != "wss" && u.Scheme != "ws" {
return configread.Provider{}, errors.New("speech model requires a WebSocket connection")
}
return p, nil
}