Files
ai-operator/internal/ai/pipeline/streaming_provider.go
T

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"
}