470 lines
14 KiB
Go
470 lines
14 KiB
Go
package pipeline
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"log/slog"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"ai-operator/internal/ai"
|
|
"ai-operator/internal/ai/llm"
|
|
"ai-operator/internal/ai/stt"
|
|
"ai-operator/internal/ai/tts"
|
|
"ai-operator/internal/config"
|
|
"ai-operator/internal/media"
|
|
"ai-operator/internal/tools"
|
|
)
|
|
|
|
type TTSFactory func(language string) tts.StreamingTTS
|
|
|
|
type StreamingProvider struct {
|
|
cfg config.Config
|
|
logger *slog.Logger
|
|
stt stt.StreamingProvider
|
|
llm llm.StreamingLLM
|
|
ttsFactory TTSFactory
|
|
|
|
events chan ai.VoiceEvent
|
|
mu sync.Mutex
|
|
sttMu sync.Mutex
|
|
stats ai.VoiceProviderStats
|
|
callID string
|
|
prompt string
|
|
closed bool
|
|
sttReq stt.StreamRequest
|
|
|
|
pending map[string]chan ai.ToolResult
|
|
turnSeq int
|
|
ttsCancel context.CancelFunc
|
|
sttRestartWindow time.Time
|
|
sttRestartCount int
|
|
sttDisabledUntil time.Time
|
|
}
|
|
|
|
func NewStreamingProvider(cfg config.Config, logger *slog.Logger) *StreamingProvider {
|
|
return NewStreamingProviderWithDeps(cfg, logger, stt.NewElevenLabsRealtime(cfg), llm.NewOpenAIStreaming(cfg), func(language string) tts.StreamingTTS {
|
|
return tts.NewElevenLabsWS(cfg)
|
|
})
|
|
}
|
|
|
|
func NewStreamingProviderWithDeps(cfg config.Config, logger *slog.Logger, sttProvider stt.StreamingProvider, llmProvider llm.StreamingLLM, ttsFactory TTSFactory) *StreamingProvider {
|
|
return &StreamingProvider{cfg: cfg, logger: logger, stt: sttProvider, llm: llmProvider, ttsFactory: ttsFactory, events: make(chan ai.VoiceEvent, 128), pending: map[string]chan ai.ToolResult{}}
|
|
}
|
|
|
|
func (p *StreamingProvider) StartSession(ctx context.Context, sc ai.VoiceSessionConfig) error {
|
|
if p.stt == nil || p.llm == nil || p.ttsFactory == nil {
|
|
return errors.New("pipeline streaming provider dependencies are not configured")
|
|
}
|
|
p.mu.Lock()
|
|
p.events = make(chan ai.VoiceEvent, 128)
|
|
p.pending = map[string]chan ai.ToolResult{}
|
|
p.closed = false
|
|
p.callID = sc.CallID
|
|
p.prompt = sc.SystemPrompt
|
|
p.stats = ai.VoiceProviderStats{}
|
|
now := time.Now().UTC()
|
|
p.stats.StartedAt = &now
|
|
p.mu.Unlock()
|
|
req := stt.StreamRequest{CallID: sc.CallID, Model: p.cfg.STT.Model, LanguageCode: p.cfg.STT.LanguageCode, LanguageAuto: p.cfg.STT.LanguageAuto, SampleRate: p.cfg.STT.SampleRate, InputFormat: p.cfg.STT.InputFormat}
|
|
p.sttReq = req
|
|
if err := p.stt.Start(ctx, req); err != nil {
|
|
return err
|
|
}
|
|
p.emit(ai.VoiceEvent{Type: ai.VoiceEventSessionStarted, CallID: sc.CallID, At: now})
|
|
go p.sttLoop()
|
|
if p.cfg.Pipeline.InitialGreeting {
|
|
go p.speakText(context.Background(), "greeting", "ru", "Здравствуйте, меня зовут Жанна. Я AI-оператор QazAimaqGas. Чем могу помочь?", nil)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (p *StreamingProvider) SendAudio(ctx context.Context, chunk media.AudioChunk) error {
|
|
p.mu.Lock()
|
|
closed := p.closed
|
|
sttDisabled := time.Now().Before(p.sttDisabledUntil)
|
|
p.stats.InputAudioFrames++
|
|
p.stats.InputAudioBytes += int64(len(chunk.Data))
|
|
p.mu.Unlock()
|
|
if closed || sttDisabled {
|
|
return nil
|
|
}
|
|
err := p.stt.SendAudio(ctx, chunk)
|
|
if err == nil {
|
|
return nil
|
|
}
|
|
if isClosedWebSocketError(err) {
|
|
if restartErr := p.restartSTT(ctx); restartErr != nil {
|
|
p.log("pipeline stt restart failed", "channel_id", p.callID, "error", restartErr)
|
|
return nil
|
|
}
|
|
if retryErr := p.stt.SendAudio(ctx, chunk); retryErr != nil {
|
|
p.log("pipeline stt audio send failed after restart", "channel_id", p.callID, "error", retryErr)
|
|
p.disableSTTBriefly()
|
|
return nil
|
|
}
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
|
|
func (p *StreamingProvider) SendToolResult(ctx context.Context, result ai.ToolResult) error {
|
|
p.mu.Lock()
|
|
ch := p.pending[result.ToolCallID]
|
|
p.mu.Unlock()
|
|
if ch == nil {
|
|
return nil
|
|
}
|
|
select {
|
|
case ch <- result:
|
|
case <-ctx.Done():
|
|
return ctx.Err()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (p *StreamingProvider) Close(ctx context.Context) error {
|
|
p.cancelTTS()
|
|
if p.stt != nil {
|
|
_ = p.stt.Close(ctx)
|
|
}
|
|
p.mu.Lock()
|
|
now := time.Now().UTC()
|
|
p.stats.ClosedAt = &now
|
|
p.mu.Unlock()
|
|
p.emit(ai.VoiceEvent{Type: ai.VoiceEventClosed, CallID: p.callID, At: now})
|
|
p.closeEvents()
|
|
return nil
|
|
}
|
|
|
|
func (p *StreamingProvider) Events() <-chan ai.VoiceEvent { return p.events }
|
|
|
|
func (p *StreamingProvider) Stats() ai.VoiceProviderStats {
|
|
p.mu.Lock()
|
|
defer p.mu.Unlock()
|
|
return p.stats
|
|
}
|
|
|
|
func (p *StreamingProvider) sttLoop() {
|
|
for ev := range p.stt.Events() {
|
|
switch ev.Type {
|
|
case stt.EventPartialTranscript:
|
|
p.emit(ai.VoiceEvent{Type: ai.VoiceEventUserTranscriptDelta, CallID: p.callID, Text: ev.Text, At: ev.At, Metadata: map[string]any{"language": ev.Language}})
|
|
if p.cfg.Pipeline.BargeIn && strings.TrimSpace(ev.Text) != "" {
|
|
p.interrupt()
|
|
}
|
|
case stt.EventSpeechStarted:
|
|
if p.cfg.Pipeline.BargeIn {
|
|
p.interrupt()
|
|
}
|
|
case stt.EventFinalTranscript, stt.EventCommittedTranscript:
|
|
text := strings.TrimSpace(ev.Text)
|
|
if text == "" {
|
|
continue
|
|
}
|
|
p.emit(ai.VoiceEvent{Type: ai.VoiceEventUserTranscriptDone, CallID: p.callID, Text: text, At: ev.At, Metadata: map[string]any{"language": ev.Language}})
|
|
go p.handleTurn(text, ev.Language)
|
|
case stt.EventError:
|
|
p.emit(ai.VoiceEvent{Type: ai.VoiceEventError, CallID: p.callID, Error: ev.Error, At: ev.At})
|
|
case stt.EventClosed:
|
|
if ev.Error != "" {
|
|
p.log("pipeline stt websocket closed", "channel_id", p.callID, "reason", ev.Error)
|
|
}
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
func (p *StreamingProvider) handleTurn(userText, language string) {
|
|
if language == "" {
|
|
language = inferLanguage(userText)
|
|
}
|
|
p.mu.Lock()
|
|
p.turnSeq++
|
|
turnID := fmt.Sprintf("turn-%d", p.turnSeq)
|
|
p.mu.Unlock()
|
|
metrics := &LatencyMetrics{CallID: p.callID, TurnID: turnID, TurnStartedAt: time.Now().UTC(), STTCommittedAt: time.Now().UTC(), STTChars: len([]rune(userText))}
|
|
|
|
toolID := turnID + "-kb"
|
|
ch := make(chan ai.ToolResult, 1)
|
|
p.mu.Lock()
|
|
p.pending[toolID] = ch
|
|
p.mu.Unlock()
|
|
p.emit(ai.VoiceEvent{Type: ai.VoiceEventToolCall, CallID: p.callID, ToolCall: &ai.ToolCall{ID: toolID, Name: tools.SearchKnowledgeBase, Arguments: map[string]any{"query": userText, "limit": 5}}, At: time.Now().UTC()})
|
|
var toolResult ai.ToolResult
|
|
select {
|
|
case toolResult = <-ch:
|
|
case <-time.After(time.Duration(p.cfg.Pipeline.MaxTurnSeconds) * time.Second):
|
|
toolResult = ai.ToolResult{CallID: p.callID, ToolCallID: toolID, Error: "tool_timeout", Result: map[string]any{"message": "База знаний временно недоступна."}}
|
|
}
|
|
p.mu.Lock()
|
|
delete(p.pending, toolID)
|
|
p.mu.Unlock()
|
|
answer := answerFromToolResult(toolResult, language)
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Duration(p.cfg.Pipeline.MaxTurnSeconds)*time.Second)
|
|
defer cancel()
|
|
llmEvents, err := p.llm.StreamGenerate(ctx, llm.GenerateRequest{CallID: p.callID, SystemPrompt: p.prompt, Messages: []llm.Message{{Role: "user", Content: llm.BuildAnswerPrompt(userText, answer)}}, MaxTokens: p.cfg.LLM.MaxOutputTokens, Temperature: p.cfg.LLM.Temperature})
|
|
if err != nil {
|
|
_ = p.speakText(ctx, turnID, language, answer, metrics)
|
|
return
|
|
}
|
|
_ = p.speakLLMStream(ctx, turnID, language, llmEvents, metrics)
|
|
}
|
|
|
|
func (p *StreamingProvider) speakLLMStream(ctx context.Context, turnID, language string, events <-chan llm.Event, metrics *LatencyMetrics) error {
|
|
ctx, cancel := context.WithCancel(ctx)
|
|
p.setTTSCancel(cancel)
|
|
defer p.clearTTSCancel(cancel)
|
|
ttsStream := p.ttsFactory(language)
|
|
if err := ttsStream.Start(ctx, tts.TTSStreamRequest{CallID: p.callID, Language: language, VoiceID: p.voiceID(language), ModelID: p.cfg.Eleven.TTSModelID, OutputFormat: p.cfg.Eleven.TTSOutputFormat, SampleRate: p.cfg.Eleven.TTSSampleRate}); err != nil {
|
|
return err
|
|
}
|
|
defer ttsStream.Close(context.Background())
|
|
audioDone := p.forwardTTSAudio(ctx, ttsStream, metrics)
|
|
chunker := NewTextChunker(p.cfg.LLM.StreamChunkMinChars, p.cfg.LLM.StreamChunkMaxChars, p.cfg.Pipeline.TTSStartAfterChars, time.Duration(p.cfg.LLM.FirstChunkTimeoutMS)*time.Millisecond, p.cfg.Pipeline.TTSStartAfterPunctuation)
|
|
for ev := range events {
|
|
switch ev.Type {
|
|
case llm.EventTextDelta:
|
|
if metrics.LLMFirstTokenAt.IsZero() {
|
|
metrics.LLMFirstTokenAt = time.Now().UTC()
|
|
}
|
|
metrics.LLMChars += len([]rune(ev.Text))
|
|
for _, chunk := range chunker.Add(ev.Text, false) {
|
|
if err := ttsStream.SendText(ctx, NaturalizeForVoice(chunk, language, p.cfg.Natural), true); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
case llm.EventError:
|
|
return errors.New(ev.Error)
|
|
case llm.EventDone:
|
|
for _, chunk := range chunker.Flush() {
|
|
if err := ttsStream.SendText(ctx, NaturalizeForVoice(chunk, language, p.cfg.Natural), true); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
_ = ttsStream.SendText(ctx, "", true)
|
|
select {
|
|
case <-audioDone:
|
|
case <-time.After(ttsFinalDrain(p.cfg.Pipeline.MaxTurnSeconds)):
|
|
case <-ctx.Done():
|
|
}
|
|
metrics.CompletedAt = time.Now().UTC()
|
|
p.log("pipeline turn latency", "metrics", metrics.Snapshot())
|
|
return nil
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (p *StreamingProvider) speakText(ctx context.Context, turnID, language, text string, metrics *LatencyMetrics) error {
|
|
ch := make(chan llm.Event, 4)
|
|
ch <- llm.Event{Type: llm.EventTextDelta, Text: text}
|
|
ch <- llm.Event{Type: llm.EventDone}
|
|
close(ch)
|
|
if metrics == nil {
|
|
metrics = &LatencyMetrics{CallID: p.callID, TurnID: turnID, TurnStartedAt: time.Now().UTC(), STTCommittedAt: time.Now().UTC()}
|
|
}
|
|
return p.speakLLMStream(ctx, turnID, language, ch, metrics)
|
|
}
|
|
|
|
func (p *StreamingProvider) forwardTTSAudio(ctx context.Context, stream tts.StreamingTTS, metrics *LatencyMetrics) <-chan struct{} {
|
|
done := make(chan struct{})
|
|
go func() {
|
|
defer close(done)
|
|
for audio := range stream.Audio() {
|
|
pcm := audio.Data
|
|
if len(pcm)%2 != 0 {
|
|
p.log("trimmed odd trailing byte from tts pcm", "channel_id", p.callID, "bytes", len(pcm))
|
|
pcm = pcm[:len(pcm)-1]
|
|
}
|
|
if len(pcm) == 0 {
|
|
continue
|
|
}
|
|
now := time.Now().UTC()
|
|
if metrics.TTSFirstAudioAt.IsZero() {
|
|
metrics.TTSFirstAudioAt = now
|
|
metrics.FirstAudibleAudioAt = now
|
|
}
|
|
metrics.TTSAudioBytes += len(pcm)
|
|
p.mu.Lock()
|
|
p.stats.OutputAudioFrames++
|
|
p.stats.OutputAudioBytes += int64(len(pcm))
|
|
p.mu.Unlock()
|
|
p.emit(ai.VoiceEvent{Type: ai.VoiceEventAssistantAudioDelta, CallID: p.callID, Audio: pcm, At: now})
|
|
select {
|
|
case <-ctx.Done():
|
|
return
|
|
default:
|
|
}
|
|
}
|
|
p.emit(ai.VoiceEvent{Type: ai.VoiceEventAssistantAudioDone, CallID: p.callID, At: time.Now().UTC()})
|
|
}()
|
|
return done
|
|
}
|
|
|
|
func (p *StreamingProvider) restartSTT(ctx context.Context) error {
|
|
p.sttMu.Lock()
|
|
defer p.sttMu.Unlock()
|
|
now := time.Now()
|
|
p.mu.Lock()
|
|
if now.Before(p.sttDisabledUntil) {
|
|
p.mu.Unlock()
|
|
return nil
|
|
}
|
|
if p.sttRestartWindow.IsZero() || now.Sub(p.sttRestartWindow) > 5*time.Second {
|
|
p.sttRestartWindow = now
|
|
p.sttRestartCount = 0
|
|
}
|
|
p.sttRestartCount++
|
|
if p.sttRestartCount > 3 {
|
|
p.sttDisabledUntil = now.Add(10 * time.Second)
|
|
p.mu.Unlock()
|
|
p.log("pipeline stt restart suppressed", "channel_id", p.callID, "cooldown", "10s")
|
|
return errors.New("stt restart suppressed")
|
|
}
|
|
p.mu.Unlock()
|
|
p.mu.Lock()
|
|
if p.closed {
|
|
p.mu.Unlock()
|
|
return nil
|
|
}
|
|
req := p.sttReq
|
|
p.mu.Unlock()
|
|
_ = p.stt.Close(context.Background())
|
|
if err := p.stt.Start(ctx, req); err != nil {
|
|
return err
|
|
}
|
|
go p.sttLoop()
|
|
p.log("pipeline stt websocket restarted", "channel_id", p.callID)
|
|
return nil
|
|
}
|
|
|
|
func (p *StreamingProvider) disableSTTBriefly() {
|
|
p.mu.Lock()
|
|
p.sttDisabledUntil = time.Now().Add(10 * time.Second)
|
|
p.mu.Unlock()
|
|
p.log("pipeline stt disabled after repeated websocket close", "channel_id", p.callID, "cooldown", "10s")
|
|
}
|
|
|
|
func isClosedWebSocketError(err error) bool {
|
|
if err == nil {
|
|
return false
|
|
}
|
|
s := strings.ToLower(err.Error())
|
|
return strings.Contains(s, "websocket: close") || strings.Contains(s, "close sent") || strings.Contains(s, "not connected") || strings.Contains(s, "closed network connection")
|
|
}
|
|
|
|
func ttsFinalDrain(maxTurnSeconds int) time.Duration {
|
|
if maxTurnSeconds <= 0 || maxTurnSeconds > 8 {
|
|
return 8 * time.Second
|
|
}
|
|
return time.Duration(maxTurnSeconds) * time.Second
|
|
}
|
|
|
|
func (p *StreamingProvider) interrupt() {
|
|
p.cancelTTS()
|
|
p.emit(ai.VoiceEvent{Type: ai.VoiceEventInterruption, CallID: p.callID, At: time.Now().UTC()})
|
|
}
|
|
|
|
func (p *StreamingProvider) setTTSCancel(cancel context.CancelFunc) {
|
|
p.mu.Lock()
|
|
p.ttsCancel = cancel
|
|
p.mu.Unlock()
|
|
}
|
|
|
|
func (p *StreamingProvider) clearTTSCancel(cancel context.CancelFunc) {
|
|
p.mu.Lock()
|
|
if fmt.Sprintf("%p", p.ttsCancel) == fmt.Sprintf("%p", cancel) {
|
|
p.ttsCancel = nil
|
|
}
|
|
p.mu.Unlock()
|
|
}
|
|
|
|
func (p *StreamingProvider) cancelTTS() {
|
|
p.mu.Lock()
|
|
cancel := p.ttsCancel
|
|
p.ttsCancel = nil
|
|
p.mu.Unlock()
|
|
if cancel != nil {
|
|
cancel()
|
|
}
|
|
}
|
|
|
|
func (p *StreamingProvider) voiceID(language string) string {
|
|
if language == "kk" && p.cfg.Eleven.VoiceIDKK != "" {
|
|
return p.cfg.Eleven.VoiceIDKK
|
|
}
|
|
if language == "kk" && p.cfg.Eleven.VoiceIDKK == "" {
|
|
p.log("kk voice id missing, using ru voice fallback")
|
|
}
|
|
return p.cfg.Eleven.VoiceIDRU
|
|
}
|
|
|
|
func (p *StreamingProvider) emit(ev ai.VoiceEvent) {
|
|
p.mu.Lock()
|
|
defer p.mu.Unlock()
|
|
p.stats.EventsSent++
|
|
if p.closed {
|
|
return
|
|
}
|
|
select {
|
|
case p.events <- ev:
|
|
default:
|
|
}
|
|
}
|
|
|
|
func (p *StreamingProvider) closeEvents() {
|
|
p.mu.Lock()
|
|
defer p.mu.Unlock()
|
|
if !p.closed {
|
|
close(p.events)
|
|
p.closed = true
|
|
}
|
|
}
|
|
|
|
func (p *StreamingProvider) log(msg string, args ...any) {
|
|
if p.logger != nil {
|
|
p.logger.Info(msg, args...)
|
|
}
|
|
}
|
|
|
|
func answerFromToolResult(result ai.ToolResult, language string) string {
|
|
if result.Result == nil {
|
|
return noAnswer(language)
|
|
}
|
|
b, _ := json.Marshal(result.Result)
|
|
var m map[string]any
|
|
_ = json.Unmarshal(b, &m)
|
|
for _, key := range []string{"answer_text", "message"} {
|
|
if s, ok := m[key].(string); ok && strings.TrimSpace(s) != "" {
|
|
return s
|
|
}
|
|
}
|
|
if result.Error == "region_required_for_question" {
|
|
if language == "kk" {
|
|
return "Қалаңызды немесе облысыңызды нақтылап жіберіңізші."
|
|
}
|
|
return "Подскажите, пожалуйста, ваш город или область?"
|
|
}
|
|
return noAnswer(language)
|
|
}
|
|
|
|
func noAnswer(language string) string {
|
|
if language == "kk" {
|
|
return "Бұл сұрақ бойынша білім базасында нақты ақпарат жоқ."
|
|
}
|
|
return "В базе знаний нет точной информации по этому вопросу."
|
|
}
|
|
|
|
func inferLanguage(text string) string {
|
|
for _, r := range text {
|
|
if strings.ContainsRune("әғқңөұүіһӘҒҚҢӨҰҮІҺ", r) {
|
|
return "kk"
|
|
}
|
|
}
|
|
return "ru"
|
|
}
|