200 lines
7.0 KiB
Go
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
|
|
}
|