sync: migrate ai-operator to Gitea (2026-08-10)
This commit is contained in:
@@ -0,0 +1,87 @@
|
||||
package fake
|
||||
|
||||
import (
|
||||
"ai-operator/internal/ai"
|
||||
"ai-operator/internal/media"
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Provider struct {
|
||||
events chan ai.VoiceEvent
|
||||
mu sync.Mutex
|
||||
stats ai.VoiceProviderStats
|
||||
callID string
|
||||
toolResults []ai.ToolResult
|
||||
closed bool
|
||||
}
|
||||
|
||||
func New() *Provider { return &Provider{events: make(chan ai.VoiceEvent, 16)} }
|
||||
func (p *Provider) StartSession(ctx context.Context, cfg ai.VoiceSessionConfig) error {
|
||||
p.mu.Lock()
|
||||
p.events = make(chan ai.VoiceEvent, 16)
|
||||
p.closed = false
|
||||
p.stats = ai.VoiceProviderStats{}
|
||||
now := time.Now().UTC()
|
||||
p.stats.StartedAt = &now
|
||||
p.callID = cfg.CallID
|
||||
p.mu.Unlock()
|
||||
p.emit(ai.VoiceEvent{Type: ai.VoiceEventSessionStarted, CallID: cfg.CallID, At: now})
|
||||
p.emit(ai.VoiceEvent{Type: ai.VoiceEventAssistantTranscriptDelta, CallID: cfg.CallID, Text: "test provider connected", At: now})
|
||||
return nil
|
||||
}
|
||||
func (p *Provider) SendAudio(ctx context.Context, ch media.AudioChunk) error {
|
||||
p.mu.Lock()
|
||||
p.stats.InputAudioFrames++
|
||||
p.stats.InputAudioBytes += int64(len(ch.Data))
|
||||
p.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
func (p *Provider) SendToolResult(ctx context.Context, r ai.ToolResult) error {
|
||||
p.mu.Lock()
|
||||
p.toolResults = append(p.toolResults, r)
|
||||
p.stats.EventsSent++
|
||||
p.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
func (p *Provider) Close(ctx context.Context) error {
|
||||
p.mu.Lock()
|
||||
now := time.Now().UTC()
|
||||
p.stats.ClosedAt = &now
|
||||
callID := p.callID
|
||||
p.mu.Unlock()
|
||||
p.emit(ai.VoiceEvent{Type: ai.VoiceEventClosed, CallID: callID, At: now})
|
||||
p.closeEvents()
|
||||
return nil
|
||||
}
|
||||
func (p *Provider) Events() <-chan ai.VoiceEvent { return p.events }
|
||||
func (p *Provider) Stats() ai.VoiceProviderStats { p.mu.Lock(); defer p.mu.Unlock(); return p.stats }
|
||||
func (p *Provider) ToolResults() []ai.ToolResult {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
out := make([]ai.ToolResult, len(p.toolResults))
|
||||
copy(out, p.toolResults)
|
||||
return out
|
||||
}
|
||||
func (p *Provider) EmitForTest(e ai.VoiceEvent) { p.emit(e) }
|
||||
func (p *Provider) emit(e ai.VoiceEvent) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.stats.EventsSent++
|
||||
if p.closed {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case p.events <- e:
|
||||
default:
|
||||
}
|
||||
}
|
||||
func (p *Provider) closeEvents() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if !p.closed {
|
||||
close(p.events)
|
||||
p.closed = true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package fake
|
||||
|
||||
import (
|
||||
"ai-operator/internal/ai"
|
||||
"ai-operator/internal/media"
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestFakeProvider(t *testing.T) {
|
||||
p := New()
|
||||
ctx := context.Background()
|
||||
if err := p.StartSession(ctx, ai.VoiceSessionConfig{CallID: "c1"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := p.SendAudio(ctx, media.AudioChunk{CallID: "c1", Data: []byte{0, 0}, Timestamp: time.Now()}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if st := p.Stats(); st.InputAudioFrames != 1 || st.InputAudioBytes != 2 {
|
||||
t.Fatalf("stats=%+v", st)
|
||||
}
|
||||
_ = p.Close(ctx)
|
||||
}
|
||||
|
||||
func TestFakeProviderCanStartAfterClose(t *testing.T) {
|
||||
p := New()
|
||||
ctx := context.Background()
|
||||
for _, callID := range []string{"c1", "c2"} {
|
||||
if err := p.StartSession(ctx, ai.VoiceSessionConfig{CallID: callID}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := p.Close(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for range p.Events() {
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"ai-operator/internal/config"
|
||||
)
|
||||
|
||||
type OpenAIStreaming struct {
|
||||
cfg config.Config
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
func NewOpenAIStreaming(cfg config.Config) *OpenAIStreaming {
|
||||
return &OpenAIStreaming{cfg: cfg, client: &http.Client{Timeout: cfg.LLM.Timeout}}
|
||||
}
|
||||
|
||||
func (p *OpenAIStreaming) StreamGenerate(ctx context.Context, req GenerateRequest) (<-chan Event, error) {
|
||||
if p.cfg.OpenAI.APIKey == "" {
|
||||
return nil, errors.New("OPENAI_API_KEY is required for streaming LLM")
|
||||
}
|
||||
out := make(chan Event, 64)
|
||||
go p.run(ctx, req, out)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (p *OpenAIStreaming) run(ctx context.Context, req GenerateRequest, out chan<- Event) {
|
||||
defer close(out)
|
||||
payload := map[string]any{
|
||||
"model": p.cfg.LLM.Model,
|
||||
"temperature": p.cfg.LLM.Temperature,
|
||||
"max_tokens": p.cfg.LLM.MaxOutputTokens,
|
||||
"stream": true,
|
||||
}
|
||||
msgs := []map[string]string{{"role": "system", "content": req.SystemPrompt}}
|
||||
for _, m := range req.Messages {
|
||||
if strings.TrimSpace(m.Content) != "" {
|
||||
msgs = append(msgs, map[string]string{"role": m.Role, "content": m.Content})
|
||||
}
|
||||
}
|
||||
payload["messages"] = msgs
|
||||
body, _ := json.Marshal(payload)
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, "https://api.openai.com/v1/chat/completions", bytes.NewReader(body))
|
||||
if err != nil {
|
||||
out <- Event{Type: EventError, Error: err.Error()}
|
||||
return
|
||||
}
|
||||
httpReq.Header.Set("Authorization", "Bearer "+p.cfg.OpenAI.APIKey)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
resp, err := p.client.Do(httpReq)
|
||||
if err != nil {
|
||||
out <- Event{Type: EventError, Error: err.Error()}
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < 200 || resp.StatusCode > 299 {
|
||||
out <- Event{Type: EventError, Error: "openai streaming llm returned non-2xx"}
|
||||
return
|
||||
}
|
||||
sc := bufio.NewScanner(resp.Body)
|
||||
for sc.Scan() {
|
||||
line := strings.TrimSpace(sc.Text())
|
||||
if !strings.HasPrefix(line, "data:") {
|
||||
continue
|
||||
}
|
||||
data := strings.TrimSpace(strings.TrimPrefix(line, "data:"))
|
||||
if data == "[DONE]" {
|
||||
out <- Event{Type: EventDone}
|
||||
return
|
||||
}
|
||||
var chunk struct {
|
||||
Choices []struct {
|
||||
Delta struct {
|
||||
Content string `json:"content"`
|
||||
} `json:"delta"`
|
||||
} `json:"choices"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(data), &chunk); err != nil {
|
||||
continue
|
||||
}
|
||||
for _, c := range chunk.Choices {
|
||||
if c.Delta.Content != "" {
|
||||
out <- Event{Type: EventTextDelta, Text: c.Delta.Content}
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := sc.Err(); err != nil {
|
||||
out <- Event{Type: EventError, Error: err.Error()}
|
||||
return
|
||||
}
|
||||
out <- Event{Type: EventDone}
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
package llm
|
||||
|
||||
import "strings"
|
||||
|
||||
func BuildAnswerPrompt(userText string, toolAnswer string) string {
|
||||
toolAnswer = strings.TrimSpace(toolAnswer)
|
||||
if toolAnswer == "" {
|
||||
toolAnswer = "В базе знаний нет точной информации по этому вопросу."
|
||||
}
|
||||
return "Сформулируй короткий голосовой ответ Жанны только по этому содержанию. Ответ 1-3 предложения, без упоминания JSON/tools/chunks. Вопрос клиента: " + userText + "\nСодержимое KB/tool result: " + toolAnswer
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
package llm
|
||||
|
||||
import "context"
|
||||
|
||||
type EventType string
|
||||
|
||||
const (
|
||||
EventTextDelta EventType = "text_delta"
|
||||
EventToolCallDelta EventType = "tool_call_delta"
|
||||
EventToolCallDone EventType = "tool_call_done"
|
||||
EventDone EventType = "done"
|
||||
EventError EventType = "error"
|
||||
)
|
||||
|
||||
type Message struct {
|
||||
Role string
|
||||
Content string
|
||||
}
|
||||
|
||||
type GenerateRequest struct {
|
||||
CallID string
|
||||
SystemPrompt string
|
||||
Messages []Message
|
||||
MaxTokens int
|
||||
Temperature float64
|
||||
}
|
||||
|
||||
type Event struct {
|
||||
Type EventType
|
||||
Text string
|
||||
Error string
|
||||
}
|
||||
|
||||
type StreamingLLM interface {
|
||||
StreamGenerate(ctx context.Context, req GenerateRequest) (<-chan Event, error)
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
package realtime
|
||||
|
||||
import (
|
||||
"ai-operator/internal/ai"
|
||||
"ai-operator/internal/audio"
|
||||
"ai-operator/internal/config"
|
||||
"encoding/json"
|
||||
)
|
||||
|
||||
func BuildSessionUpdate(cfg config.OpenAIConfig, instructions string) ([]byte, error) {
|
||||
session := map[string]any{"type": "realtime", "model": cfg.RealtimeModel, "instructions": instructions, "audio": map[string]any{"input": map[string]any{"format": map[string]any{"type": "audio/pcm", "rate": cfg.RealtimeInputSampleRate}, "turn_detection": map[string]any{"type": cfg.RealtimeTurnDetection}}, "output": map[string]any{"format": map[string]any{"type": "audio/pcm", "rate": cfg.RealtimeOutputSampleRate}}}, "reasoning": map[string]any{"effort": cfg.RealtimeReasoningEffort}}
|
||||
if cfg.RealtimeVoice != "" {
|
||||
session["audio"].(map[string]any)["output"].(map[string]any)["voice"] = cfg.RealtimeVoice
|
||||
}
|
||||
return json.Marshal(map[string]any{"type": "session.update", "session": session})
|
||||
}
|
||||
func BuildAudioAppend(pcm []byte) ([]byte, error) {
|
||||
return json.Marshal(map[string]any{"type": "input_audio_buffer.append", "audio": audio.Base64Encode(pcm)})
|
||||
}
|
||||
func BuildCommit() ([]byte, error) {
|
||||
return json.Marshal(map[string]any{"type": "input_audio_buffer.commit"})
|
||||
}
|
||||
func BuildResponseCreate() ([]byte, error) {
|
||||
return json.Marshal(map[string]any{"type": "response.create"})
|
||||
}
|
||||
func BuildResponseCreateWithInstructions(instructions string) ([]byte, error) {
|
||||
if instructions == "" {
|
||||
return BuildResponseCreate()
|
||||
}
|
||||
return json.Marshal(map[string]any{"type": "response.create", "response": map[string]any{"instructions": instructions}})
|
||||
}
|
||||
func BuildResponseCancel() ([]byte, error) {
|
||||
return json.Marshal(map[string]any{"type": "response.cancel"})
|
||||
}
|
||||
func BuildInputAudioClear() ([]byte, error) {
|
||||
return json.Marshal(map[string]any{"type": "input_audio_buffer.clear"})
|
||||
}
|
||||
func BuildToolResult(result ai.ToolResult) ([]byte, error) {
|
||||
b, _ := json.Marshal(result.Result)
|
||||
if result.Error != "" {
|
||||
b = []byte(result.Error)
|
||||
}
|
||||
return json.Marshal(map[string]any{"type": "conversation.item.create", "item": map[string]any{"type": "function_call_output", "call_id": result.ToolCallID, "output": string(b)}})
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package realtime
|
||||
|
||||
import (
|
||||
"ai-operator/internal/ai"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"time"
|
||||
)
|
||||
|
||||
type GenericRealtimeEvent struct {
|
||||
Type string
|
||||
Raw map[string]any
|
||||
}
|
||||
|
||||
func ParseServerEvent(data []byte) (ai.VoiceEvent, GenericRealtimeEvent, error) {
|
||||
var base map[string]any
|
||||
if err := json.Unmarshal(data, &base); err != nil {
|
||||
return ai.VoiceEvent{}, GenericRealtimeEvent{}, err
|
||||
}
|
||||
typ, _ := base["type"].(string)
|
||||
ev := ai.VoiceEvent{At: time.Now().UTC(), Metadata: map[string]any{"openai_type": typ}}
|
||||
switch typ {
|
||||
case "session.created":
|
||||
ev.Type = ai.VoiceEventSessionStarted
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
case "session.updated":
|
||||
ev.Type = ai.VoiceEventSessionUpdated
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
case "error":
|
||||
ev.Type = ai.VoiceEventError
|
||||
if e, ok := base["error"].(map[string]any); ok {
|
||||
ev.Error = toStr(e["message"])
|
||||
} else {
|
||||
ev.Error = "openai error"
|
||||
}
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
case "input_audio_buffer.speech_started":
|
||||
ev.Type = ai.VoiceEventInterruption
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
case "response.output_audio.delta":
|
||||
ev.Type = ai.VoiceEventAssistantAudioDelta
|
||||
b, err := base64.StdEncoding.DecodeString(toStr(base["delta"]))
|
||||
if err != nil {
|
||||
return ai.VoiceEvent{}, GenericRealtimeEvent{}, err
|
||||
}
|
||||
ev.Audio = b
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
case "response.output_audio.done":
|
||||
ev.Type = ai.VoiceEventAssistantAudioDone
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
case "response.output_audio_transcript.delta", "response.output_text.delta":
|
||||
ev.Type = ai.VoiceEventAssistantTranscriptDelta
|
||||
ev.Text = toStr(base["delta"])
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
case "response.output_audio_transcript.done", "response.output_text.done":
|
||||
ev.Type = ai.VoiceEventAssistantTranscriptDone
|
||||
ev.Text = toStr(base["transcript"]) + toStr(base["text"])
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
case "rate_limits.updated":
|
||||
ev.Type = ai.VoiceEventRateLimitsUpdated
|
||||
ev.Metadata = base
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
case "response.done":
|
||||
if tc := extractToolCall(base); tc != nil {
|
||||
ev.Type = ai.VoiceEventToolCall
|
||||
ev.ToolCall = tc
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
}
|
||||
ev.Type = ai.VoiceEventAssistantAudioDone
|
||||
return ev, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
default:
|
||||
return ai.VoiceEvent{}, GenericRealtimeEvent{Type: typ, Raw: base}, nil
|
||||
}
|
||||
}
|
||||
func extractToolCall(m map[string]any) *ai.ToolCall {
|
||||
resp, ok := m["response"].(map[string]any)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
outs, ok := resp["output"].([]any)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
for _, o := range outs {
|
||||
item, ok := o.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if toStr(item["type"]) == "function_call" {
|
||||
raw := toStr(item["arguments"])
|
||||
args := map[string]any{}
|
||||
_ = json.Unmarshal([]byte(raw), &args)
|
||||
return &ai.ToolCall{ID: toStr(item["call_id"]), Name: toStr(item["name"]), Arguments: args, RawArguments: raw}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func toStr(v any) string {
|
||||
if s, ok := v.(string); ok {
|
||||
return s
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,224 @@
|
||||
package realtime
|
||||
|
||||
import (
|
||||
"ai-operator/internal/ai"
|
||||
"ai-operator/internal/audio"
|
||||
"ai-operator/internal/config"
|
||||
"ai-operator/internal/media"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"github.com/gorilla/websocket"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Provider struct {
|
||||
cfg config.Config
|
||||
logger *slog.Logger
|
||||
conn *websocket.Conn
|
||||
events chan ai.VoiceEvent
|
||||
mu sync.Mutex
|
||||
stats ai.VoiceProviderStats
|
||||
inputBytes, outputBytes int64
|
||||
eventsClosed bool
|
||||
}
|
||||
|
||||
func NewProvider(cfg config.Config, logger *slog.Logger) *Provider {
|
||||
return &Provider{cfg: cfg, logger: logger, events: make(chan ai.VoiceEvent, 64)}
|
||||
}
|
||||
func (p *Provider) StartSession(ctx context.Context, sc ai.VoiceSessionConfig) error {
|
||||
if p.cfg.OpenAI.APIKey == "" {
|
||||
return errors.New("OPENAI_API_KEY is required for OpenAI Realtime")
|
||||
}
|
||||
p.mu.Lock()
|
||||
p.events = make(chan ai.VoiceEvent, 64)
|
||||
p.eventsClosed = false
|
||||
p.stats = ai.VoiceProviderStats{}
|
||||
p.inputBytes = 0
|
||||
p.outputBytes = 0
|
||||
p.mu.Unlock()
|
||||
u, _, err := BuildURL(p.cfg.OpenAI)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
h := http.Header{}
|
||||
h.Set("Authorization", "Bearer "+p.cfg.OpenAI.APIKey)
|
||||
h.Set("OpenAI-Safety-Identifier", safetyID(sc.CallID))
|
||||
d := websocket.Dialer{HandshakeTimeout: p.cfg.OpenAI.RealtimeConnectTimeout}
|
||||
conn, _, err := d.DialContext(ctx, u, h)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.conn = conn
|
||||
now := time.Now().UTC()
|
||||
p.mu.Lock()
|
||||
p.stats.StartedAt = &now
|
||||
p.mu.Unlock()
|
||||
if err := p.waitFor(ctx, "session.created"); err != nil {
|
||||
return err
|
||||
}
|
||||
instr := sc.SystemPrompt
|
||||
if instr == "" {
|
||||
instr = "You are a test voice provider adapter for an AI operator. For now, keep responses very short. Do not claim to access customer data. Do not answer business questions yet."
|
||||
}
|
||||
msg, err := BuildSessionUpdate(p.cfg.OpenAI, instr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := conn.WriteMessage(websocket.TextMessage, msg); err != nil {
|
||||
return err
|
||||
}
|
||||
_ = p.waitFor(ctx, "session.updated")
|
||||
go p.readLoop(sc.CallID)
|
||||
if p.cfg.OpenAI.RealtimeInitialGreeting {
|
||||
greeting := "Say exactly this short natural greeting in a calm, natural contact-center voice, without asking to choose language or region: Здравствуйте, меня зовут Жанна. Я AI-оператор QazAimaqGas. Чем могу помочь?"
|
||||
rc, err := BuildResponseCreateWithInstructions(greeting)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := conn.WriteMessage(websocket.TextMessage, rc); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (p *Provider) waitFor(ctx context.Context, want string) error {
|
||||
deadline := time.After(p.cfg.OpenAI.RealtimeConnectTimeout)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-deadline:
|
||||
return errors.New("timeout waiting for " + want)
|
||||
default:
|
||||
_, data, err := p.conn.ReadMessage()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var m map[string]any
|
||||
_ = jsonUnmarshal(data, &m)
|
||||
if toStr(m["type"]) == want {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
func (p *Provider) readLoop(callID string) {
|
||||
defer p.closeEvents()
|
||||
for {
|
||||
_, data, err := p.conn.ReadMessage()
|
||||
if err != nil {
|
||||
p.emit(ai.VoiceEvent{Type: ai.VoiceEventClosed, CallID: callID, At: time.Now().UTC()})
|
||||
return
|
||||
}
|
||||
ev, _, err := ParseServerEvent(data)
|
||||
if err != nil {
|
||||
p.emit(ai.VoiceEvent{Type: ai.VoiceEventError, CallID: callID, Error: err.Error(), At: time.Now().UTC()})
|
||||
continue
|
||||
}
|
||||
if ev.Type != "" {
|
||||
ev.CallID = callID
|
||||
p.emit(ev)
|
||||
}
|
||||
}
|
||||
}
|
||||
func (p *Provider) SendAudio(ctx context.Context, ch media.AudioChunk) error {
|
||||
if p.conn == nil {
|
||||
return errors.New("openai realtime websocket not connected")
|
||||
}
|
||||
if !audio.IsPCM16Aligned(ch.Data) {
|
||||
return errors.New("pcm16 audio is not aligned")
|
||||
}
|
||||
p.mu.Lock()
|
||||
p.inputBytes += int64(len(ch.Data))
|
||||
if p.inputBytes > int64(p.cfg.OpenAI.RealtimeMaxInputAudioBytes) {
|
||||
p.mu.Unlock()
|
||||
return errors.New("max input audio bytes exceeded")
|
||||
}
|
||||
p.stats.InputAudioFrames++
|
||||
p.stats.InputAudioBytes += int64(len(ch.Data))
|
||||
p.mu.Unlock()
|
||||
chunks := audio.ChunkBytes(ch.Data, p.cfg.OpenAI.RealtimeMaxAudioChunkBytes)
|
||||
for _, c := range chunks {
|
||||
msg, err := BuildAudioAppend(c)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := p.conn.WriteMessage(websocket.TextMessage, msg); err != nil {
|
||||
return err
|
||||
}
|
||||
p.mu.Lock()
|
||||
p.stats.EventsSent++
|
||||
p.mu.Unlock()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (p *Provider) SendToolResult(ctx context.Context, r ai.ToolResult) error {
|
||||
if p.conn == nil {
|
||||
return errors.New("openai realtime websocket not connected")
|
||||
}
|
||||
msg, err := BuildToolResult(r)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := p.conn.WriteMessage(websocket.TextMessage, msg); err != nil {
|
||||
return err
|
||||
}
|
||||
rc, _ := BuildResponseCreate()
|
||||
return p.conn.WriteMessage(websocket.TextMessage, rc)
|
||||
}
|
||||
func (p *Provider) Close(ctx context.Context) error {
|
||||
p.mu.Lock()
|
||||
now := time.Now().UTC()
|
||||
p.stats.ClosedAt = &now
|
||||
p.mu.Unlock()
|
||||
if p.conn != nil {
|
||||
_ = p.conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""))
|
||||
return p.conn.Close()
|
||||
}
|
||||
p.closeEvents()
|
||||
return nil
|
||||
}
|
||||
func (p *Provider) Events() <-chan ai.VoiceEvent { return p.events }
|
||||
func (p *Provider) Stats() ai.VoiceProviderStats { p.mu.Lock(); defer p.mu.Unlock(); return p.stats }
|
||||
func (p *Provider) emit(e ai.VoiceEvent) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.stats.EventsReceived++
|
||||
if e.Type == ai.VoiceEventAssistantAudioDelta {
|
||||
p.outputBytes += int64(len(e.Audio))
|
||||
p.stats.OutputAudioFrames++
|
||||
p.stats.OutputAudioBytes += int64(len(e.Audio))
|
||||
}
|
||||
if e.Type == ai.VoiceEventToolCall {
|
||||
p.stats.ToolCallsReceived++
|
||||
}
|
||||
if e.Type == ai.VoiceEventError {
|
||||
p.stats.Errors++
|
||||
}
|
||||
if p.eventsClosed {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case p.events <- e:
|
||||
default:
|
||||
}
|
||||
}
|
||||
func (p *Provider) closeEvents() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if !p.eventsClosed {
|
||||
close(p.events)
|
||||
p.eventsClosed = true
|
||||
}
|
||||
}
|
||||
func safetyID(callID string) string {
|
||||
sum := sha256.Sum256([]byte(callID))
|
||||
return "aiop-" + hex.EncodeToString(sum[:])[:16]
|
||||
}
|
||||
func jsonUnmarshal(b []byte, v any) error { return json.Unmarshal(b, v) }
|
||||
@@ -0,0 +1,161 @@
|
||||
package realtime
|
||||
|
||||
import (
|
||||
"ai-operator/internal/ai"
|
||||
"ai-operator/internal/config"
|
||||
"ai-operator/internal/media"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"github.com/gorilla/websocket"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestURLBuildersParsers(t *testing.T) {
|
||||
cfg := config.OpenAIConfig{RealtimeURL: "wss://api.openai.com/v1/realtime", RealtimeModel: "gpt realtime", RealtimeVoice: "marin", RealtimeInputSampleRate: 24000, RealtimeOutputSampleRate: 24000, RealtimeTurnDetection: "server_vad", RealtimeReasoningEffort: "low"}
|
||||
u, _, err := BuildURL(cfg)
|
||||
if err != nil || !strings.Contains(u, "model=gpt+realtime") {
|
||||
t.Fatalf("url=%s err=%v", u, err)
|
||||
}
|
||||
b, err := BuildSessionUpdate(cfg, "test")
|
||||
if err != nil || strings.Contains(string(b), "sk-") {
|
||||
t.Fatal(string(b))
|
||||
}
|
||||
if !strings.Contains(string(b), `"voice":"marin"`) {
|
||||
t.Fatalf("session.update missing configured voice: %s", b)
|
||||
}
|
||||
for _, fn := range []func() ([]byte, error){func() ([]byte, error) { return BuildAudioAppend([]byte{1, 2}) }, BuildCommit, BuildResponseCreate, BuildResponseCancel, BuildInputAudioClear, func() ([]byte, error) {
|
||||
return BuildToolResult(ai.ToolResult{ToolCallID: "tc", Result: map[string]string{"ok": "true"}})
|
||||
}} {
|
||||
x, err := fn()
|
||||
if err != nil || !json.Valid(x) {
|
||||
t.Fatalf("bad event %s %v", x, err)
|
||||
}
|
||||
}
|
||||
ev, _, err := ParseServerEvent([]byte(`{"type":"response.output_audio.delta","delta":"AQI="}`))
|
||||
if err != nil || ev.Type != ai.VoiceEventAssistantAudioDelta || len(ev.Audio) != 2 {
|
||||
t.Fatalf("ev=%+v err=%v", ev, err)
|
||||
}
|
||||
ev, _, _ = ParseServerEvent([]byte(`{"type":"input_audio_buffer.speech_started"}`))
|
||||
if ev.Type != ai.VoiceEventInterruption {
|
||||
t.Fatal(ev.Type)
|
||||
}
|
||||
ev, _, _ = ParseServerEvent([]byte(`{"type":"response.done","response":{"output":[{"type":"function_call","call_id":"c","name":"set_language","arguments":"{\"language\":\"ru\"}"}]}}`))
|
||||
if ev.ToolCall == nil || ev.ToolCall.Name != "set_language" {
|
||||
t.Fatalf("tool=%+v", ev.ToolCall)
|
||||
}
|
||||
if _, _, err := ParseServerEvent([]byte(`bad`)); err == nil {
|
||||
t.Fatal("want err")
|
||||
}
|
||||
}
|
||||
func TestProviderFakeWebSocket(t *testing.T) {
|
||||
up := websocket.Upgrader{}
|
||||
var auth string
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
auth = r.Header.Get("Authorization")
|
||||
c, err := up.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer c.Close()
|
||||
_ = c.WriteJSON(map[string]any{"type": "session.created"})
|
||||
_, msg, _ := c.ReadMessage()
|
||||
if !strings.Contains(string(msg), "session.update") {
|
||||
t.Errorf("msg=%s", msg)
|
||||
}
|
||||
_ = c.WriteJSON(map[string]any{"type": "session.updated"})
|
||||
_, _, _ = c.ReadMessage()
|
||||
_ = c.WriteJSON(map[string]any{"type": "response.output_audio.delta", "delta": "AQI="})
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}))
|
||||
defer srv.Close()
|
||||
cfg := config.Config{OpenAI: config.OpenAIConfig{APIKey: "sk-test", RealtimeURL: "ws" + strings.TrimPrefix(srv.URL, "http"), RealtimeModel: "m", RealtimeConnectTimeout: time.Second, RealtimeMaxAudioChunkBytes: 100, RealtimeMaxInputAudioBytes: 1000, RealtimeInputSampleRate: 24000, RealtimeOutputSampleRate: 24000, RealtimeTurnDetection: "server_vad", RealtimeReasoningEffort: "low"}}
|
||||
p := NewProvider(cfg, nil)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
if err := p.StartSession(ctx, ai.VoiceSessionConfig{CallID: "c1"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if auth != "Bearer sk-test" {
|
||||
t.Fatal("auth missing")
|
||||
}
|
||||
if err := p.SendAudio(ctx, media.AudioChunk{Data: []byte{0, 0}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
select {
|
||||
case ev := <-p.Events():
|
||||
_ = ev
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("no event")
|
||||
}
|
||||
_ = p.Close(ctx)
|
||||
}
|
||||
func TestProviderCanStartAfterClose(t *testing.T) {
|
||||
up := websocket.Upgrader{}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
c, err := up.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer c.Close()
|
||||
_ = c.WriteJSON(map[string]any{"type": "session.created"})
|
||||
_, msg, _ := c.ReadMessage()
|
||||
if !strings.Contains(string(msg), "session.update") {
|
||||
t.Errorf("msg=%s", msg)
|
||||
}
|
||||
_ = c.WriteJSON(map[string]any{"type": "session.updated"})
|
||||
_ = c.WriteJSON(map[string]any{"type": "response.output_audio.delta", "delta": "AQI="})
|
||||
for {
|
||||
if _, _, err := c.ReadMessage(); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
cfg := config.Config{OpenAI: config.OpenAIConfig{APIKey: "sk-test", RealtimeURL: "ws" + strings.TrimPrefix(srv.URL, "http"), RealtimeModel: "m", RealtimeConnectTimeout: time.Second, RealtimeMaxAudioChunkBytes: 100, RealtimeMaxInputAudioBytes: 1000, RealtimeInputSampleRate: 24000, RealtimeOutputSampleRate: 24000, RealtimeTurnDetection: "server_vad", RealtimeReasoningEffort: "low"}}
|
||||
p := NewProvider(cfg, nil)
|
||||
for _, callID := range []string{"c1", "c2"} {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
if err := p.StartSession(ctx, ai.VoiceSessionConfig{CallID: callID}); err != nil {
|
||||
cancel()
|
||||
t.Fatal(err)
|
||||
}
|
||||
select {
|
||||
case ev := <-p.Events():
|
||||
if ev.CallID != callID {
|
||||
cancel()
|
||||
t.Fatalf("call_id=%s want %s", ev.CallID, callID)
|
||||
}
|
||||
case <-ctx.Done():
|
||||
cancel()
|
||||
t.Fatal("no event")
|
||||
}
|
||||
if err := p.Close(ctx); err != nil {
|
||||
cancel()
|
||||
t.Fatal(err)
|
||||
}
|
||||
for {
|
||||
select {
|
||||
case _, ok := <-p.Events():
|
||||
if !ok {
|
||||
cancel()
|
||||
goto next
|
||||
}
|
||||
case <-ctx.Done():
|
||||
cancel()
|
||||
t.Fatal("events channel did not close")
|
||||
}
|
||||
}
|
||||
next:
|
||||
}
|
||||
}
|
||||
func TestCostGuard(t *testing.T) {
|
||||
p := NewProvider(config.Config{OpenAI: config.OpenAIConfig{RealtimeMaxInputAudioBytes: 1}}, nil)
|
||||
p.conn = &websocket.Conn{}
|
||||
if err := p.SendAudio(context.Background(), media.AudioChunk{Data: []byte{0, 0}}); err == nil {
|
||||
t.Fatal("want cost err")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package realtime
|
||||
|
||||
import (
|
||||
"ai-operator/internal/config"
|
||||
"net/url"
|
||||
)
|
||||
|
||||
func BuildURL(cfg config.OpenAIConfig) (string, string, error) {
|
||||
u, err := url.Parse(cfg.RealtimeURL)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("model", cfg.RealtimeModel)
|
||||
u.RawQuery = q.Encode()
|
||||
return u.String(), u.String(), nil
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
package pipeline
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"ai-operator/internal/ai/llm"
|
||||
"ai-operator/internal/ai/stt"
|
||||
"ai-operator/internal/ai/tts"
|
||||
"ai-operator/internal/media"
|
||||
)
|
||||
|
||||
type FakeSTT struct {
|
||||
events chan stt.Event
|
||||
audio int
|
||||
}
|
||||
|
||||
func NewFakeSTT() *FakeSTT { return &FakeSTT{events: make(chan stt.Event, 16)} }
|
||||
|
||||
func (f *FakeSTT) Start(ctx context.Context, req stt.StreamRequest) error { return nil }
|
||||
|
||||
func (f *FakeSTT) SendAudio(ctx context.Context, chunk media.AudioChunk) error {
|
||||
f.audio += len(chunk.Data)
|
||||
f.events <- stt.Event{Type: stt.EventPartialTranscript, Text: "Сколько стоит", At: time.Now().UTC()}
|
||||
f.events <- stt.Event{Type: stt.EventCommittedTranscript, Text: "Сколько стоит первичное подключение газа?", Language: "ru", At: time.Now().UTC()}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *FakeSTT) Events() <-chan stt.Event { return f.events }
|
||||
|
||||
func (f *FakeSTT) Close(ctx context.Context) error {
|
||||
close(f.events)
|
||||
return nil
|
||||
}
|
||||
|
||||
type FakeLLM struct {
|
||||
Text string
|
||||
}
|
||||
|
||||
func (f FakeLLM) StreamGenerate(ctx context.Context, req llm.GenerateRequest) (<-chan llm.Event, error) {
|
||||
out := make(chan llm.Event, 8)
|
||||
text := f.Text
|
||||
if text == "" {
|
||||
text = "Первичное подключение газа к газовому оборудованию осуществляется бесплатно."
|
||||
}
|
||||
go func() {
|
||||
defer close(out)
|
||||
parts := strings.SplitAfter(text, " ")
|
||||
for _, p := range parts {
|
||||
if p != "" {
|
||||
out <- llm.Event{Type: llm.EventTextDelta, Text: p}
|
||||
}
|
||||
}
|
||||
out <- llm.Event{Type: llm.EventDone}
|
||||
}()
|
||||
return out, nil
|
||||
}
|
||||
|
||||
type FakeTTSFactory struct {
|
||||
mu sync.Mutex
|
||||
Texts []string
|
||||
Starts int
|
||||
Cancels int
|
||||
}
|
||||
|
||||
func (f *FakeTTSFactory) New(language string) tts.StreamingTTS {
|
||||
return &FakeTTS{factory: f, audio: make(chan tts.AudioChunk, 16)}
|
||||
}
|
||||
|
||||
type FakeTTS struct {
|
||||
factory *FakeTTSFactory
|
||||
audio chan tts.AudioChunk
|
||||
closed bool
|
||||
}
|
||||
|
||||
func (f *FakeTTS) Start(ctx context.Context, req tts.TTSStreamRequest) error {
|
||||
f.factory.mu.Lock()
|
||||
f.factory.Starts++
|
||||
f.factory.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *FakeTTS) SendText(ctx context.Context, text string, flush bool) error {
|
||||
f.factory.mu.Lock()
|
||||
if text != "" {
|
||||
f.factory.Texts = append(f.factory.Texts, text)
|
||||
}
|
||||
f.factory.mu.Unlock()
|
||||
if text != "" {
|
||||
select {
|
||||
case f.audio <- tts.AudioChunk{Data: []byte{0, 1, 0, 1}, Timestamp: time.Now().UTC()}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
if flush && text == "" {
|
||||
f.Close(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *FakeTTS) Audio() <-chan tts.AudioChunk { return f.audio }
|
||||
|
||||
func (f *FakeTTS) Close(ctx context.Context) error {
|
||||
if !f.closed {
|
||||
close(f.audio)
|
||||
f.closed = true
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package pipeline
|
||||
|
||||
import "time"
|
||||
|
||||
type LatencyMetrics struct {
|
||||
CallID string
|
||||
TurnID string
|
||||
TurnStartedAt time.Time
|
||||
STTFirstPartialAt time.Time
|
||||
STTCommittedAt time.Time
|
||||
LLMFirstTokenAt time.Time
|
||||
TTSFirstAudioAt time.Time
|
||||
FirstAudibleAudioAt time.Time
|
||||
CompletedAt time.Time
|
||||
STTChars int
|
||||
LLMChars int
|
||||
TTSAudioBytes int
|
||||
Interruptions int
|
||||
}
|
||||
|
||||
func (m LatencyMetrics) Snapshot() map[string]any {
|
||||
return map[string]any{
|
||||
"call_id": m.CallID,
|
||||
"turn_id": m.TurnID,
|
||||
"stt_first_partial_ms": sinceMS(m.TurnStartedAt, m.STTFirstPartialAt),
|
||||
"stt_committed_ms": sinceMS(m.TurnStartedAt, m.STTCommittedAt),
|
||||
"llm_first_token_ms": sinceMS(m.STTCommittedAt, m.LLMFirstTokenAt),
|
||||
"tts_first_audio_ms": sinceMS(m.LLMFirstTokenAt, m.TTSFirstAudioAt),
|
||||
"first_audible_audio_ms": sinceMS(m.STTCommittedAt, m.FirstAudibleAudioAt),
|
||||
"total_turn_ms": sinceMS(m.TurnStartedAt, m.CompletedAt),
|
||||
"stt_chars": m.STTChars,
|
||||
"llm_chars": m.LLMChars,
|
||||
"tts_audio_bytes": m.TTSAudioBytes,
|
||||
"interruptions_count": m.Interruptions,
|
||||
}
|
||||
}
|
||||
|
||||
func sinceMS(start, end time.Time) any {
|
||||
if start.IsZero() || end.IsZero() {
|
||||
return nil
|
||||
}
|
||||
return end.Sub(start).Milliseconds()
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
package pipeline
|
||||
|
||||
import (
|
||||
"ai-operator/internal/ai/tts"
|
||||
"ai-operator/internal/config"
|
||||
)
|
||||
|
||||
func NaturalizeForVoice(text string, language string, cfg config.NaturalnessConfig) string {
|
||||
return tts.NaturalizeForVoice(text, language, cfg)
|
||||
}
|
||||
|
||||
func RemoveAudioTags(text string) string {
|
||||
return tts.RemoveAudioTags(text)
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package pipeline
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"ai-operator/internal/ai"
|
||||
"ai-operator/internal/config"
|
||||
"ai-operator/internal/media"
|
||||
)
|
||||
|
||||
func TestTextChunker(t *testing.T) {
|
||||
c := NewTextChunker(10, 80, 20, time.Millisecond, true)
|
||||
chunks := c.Add("Первое предложение. Второе", false)
|
||||
if len(chunks) != 1 || chunks[0] != "Первое предложение." {
|
||||
t.Fatalf("chunks=%v", chunks)
|
||||
}
|
||||
flush := c.Flush()
|
||||
if len(flush) != 1 || flush[0] != "Второе" {
|
||||
t.Fatalf("flush=%v", flush)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNaturalizer(t *testing.T) {
|
||||
cfg := config.NaturalnessConfig{Enabled: true, AudioTagsEnabled: true, AllowNonverbalTags: true, MaxAudioTagsPerResponse: 2}
|
||||
got := NaturalizeForVoice("Здравствуйте, меня зовут Жанна. Чем могу помочь?", "ru", cfg)
|
||||
if got == "" || got == "Здравствуйте, меня зовут Жанна. Чем могу помочь?" {
|
||||
t.Fatalf("not naturalized: %q", got)
|
||||
}
|
||||
if stripped := RemoveAudioTags(got); stripped == "" || stripped == got {
|
||||
t.Fatalf("tags not stripped: %q", stripped)
|
||||
}
|
||||
cfg.AllowCough = false
|
||||
if got := NaturalizeForVoice("Первичное подключение газа бесплатно.", "ru", cfg); containsRune(got, "cough") {
|
||||
t.Fatalf("cough added unexpectedly: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamingProviderFakeEndToEnd(t *testing.T) {
|
||||
cfg := config.Config{
|
||||
STT: config.STTConfig{Model: "scribe_realtime_v2", SampleRate: 16000, InputFormat: "pcm_16000"},
|
||||
LLM: config.LLMConfig{StreamChunkMinChars: 10, StreamChunkMaxChars: 160, FirstChunkTimeoutMS: 1, MaxOutputTokens: 100, Temperature: 0.2},
|
||||
Eleven: config.ElevenLabsConfig{VoiceIDRU: "voice", TTSModelID: "eleven_flash_v2_5", TTSOutputFormat: "pcm_16000", TTSSampleRate: 16000},
|
||||
Natural: config.NaturalnessConfig{Enabled: false},
|
||||
Pipeline: config.PipelineConfig{InitialGreeting: false, BargeIn: true, MaxTurnSeconds: 2, TTSStartAfterChars: 20, TTSStartAfterPunctuation: true},
|
||||
}
|
||||
factory := &FakeTTSFactory{}
|
||||
p := NewStreamingProviderWithDeps(cfg, nil, NewFakeSTT(), FakeLLM{}, factory.New)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
if err := p.StartSession(ctx, ai.VoiceSessionConfig{CallID: "c1", SystemPrompt: "prompt"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := p.SendAudio(ctx, media.AudioChunk{Data: []byte{0, 0}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tool := false
|
||||
audio := false
|
||||
for !audio {
|
||||
select {
|
||||
case ev := <-p.Events():
|
||||
if ev.Type == ai.VoiceEventToolCall {
|
||||
tool = true
|
||||
_ = p.SendToolResult(ctx, ai.ToolResult{CallID: ev.CallID, ToolCallID: ev.ToolCall.ID, Result: map[string]any{"answer_text": "Первичное подключение бесплатно."}})
|
||||
}
|
||||
if ev.Type == ai.VoiceEventAssistantAudioDelta {
|
||||
audio = true
|
||||
}
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timeout")
|
||||
}
|
||||
}
|
||||
if !tool || !audio || len(factory.Texts) == 0 {
|
||||
t.Fatalf("tool=%t audio=%t texts=%v", tool, audio, factory.Texts)
|
||||
}
|
||||
}
|
||||
|
||||
func containsRune(s, sub string) bool {
|
||||
for i := 0; i+len(sub) <= len(s); i++ {
|
||||
if s[i:i+len(sub)] == sub {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,469 @@
|
||||
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"
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
package pipeline
|
||||
|
||||
// Streaming session orchestration lives in streaming_provider.go. This file is
|
||||
// intentionally kept as the package boundary for future per-call session state.
|
||||
@@ -0,0 +1,120 @@
|
||||
package pipeline
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
type TextChunker struct {
|
||||
MinChars int
|
||||
MaxChars int
|
||||
StartAfter int
|
||||
FirstTimeout time.Duration
|
||||
Punctuation bool
|
||||
buf strings.Builder
|
||||
firstAt time.Time
|
||||
}
|
||||
|
||||
func NewTextChunker(minChars, maxChars, startAfter int, firstTimeout time.Duration, punctuation bool) *TextChunker {
|
||||
if minChars <= 0 {
|
||||
minChars = 30
|
||||
}
|
||||
if maxChars < minChars {
|
||||
maxChars = 160
|
||||
}
|
||||
if startAfter <= 0 {
|
||||
startAfter = 60
|
||||
}
|
||||
if firstTimeout <= 0 {
|
||||
firstTimeout = 1200 * time.Millisecond
|
||||
}
|
||||
return &TextChunker{MinChars: minChars, MaxChars: maxChars, StartAfter: startAfter, FirstTimeout: firstTimeout, Punctuation: punctuation}
|
||||
}
|
||||
|
||||
func (c *TextChunker) Add(delta string, final bool) []string {
|
||||
delta = normalizeForSpeech(delta)
|
||||
if delta != "" {
|
||||
if c.buf.Len() == 0 {
|
||||
c.firstAt = time.Now()
|
||||
}
|
||||
c.buf.WriteString(delta)
|
||||
}
|
||||
var out []string
|
||||
for {
|
||||
s := strings.TrimSpace(c.buf.String())
|
||||
if s == "" {
|
||||
c.buf.Reset()
|
||||
return out
|
||||
}
|
||||
emitAt := c.emitIndex(s, final)
|
||||
if emitAt <= 0 {
|
||||
return out
|
||||
}
|
||||
chunk := strings.TrimSpace(s[:emitAt])
|
||||
out = append(out, chunk)
|
||||
rest := strings.TrimSpace(s[emitAt:])
|
||||
c.buf.Reset()
|
||||
c.buf.WriteString(rest)
|
||||
if rest == "" {
|
||||
return out
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *TextChunker) Flush() []string { return c.Add("", true) }
|
||||
|
||||
func (c *TextChunker) emitIndex(s string, final bool) int {
|
||||
if final {
|
||||
return len(s)
|
||||
}
|
||||
if c.Punctuation && utf8.RuneCountInString(s) >= c.MinChars {
|
||||
if idx := lastSentenceBoundary(s, c.MaxChars); idx > 0 {
|
||||
return idx
|
||||
}
|
||||
}
|
||||
if utf8.RuneCountInString(s) >= c.MaxChars {
|
||||
return byteIndexByRunes(s, c.MaxChars)
|
||||
}
|
||||
if utf8.RuneCountInString(s) >= c.StartAfter && time.Since(c.firstAt) >= c.FirstTimeout {
|
||||
return len(s)
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func lastSentenceBoundary(s string, maxRunes int) int {
|
||||
limit := byteIndexByRunes(s, maxRunes)
|
||||
if limit <= 0 || limit > len(s) {
|
||||
limit = len(s)
|
||||
}
|
||||
last := -1
|
||||
for i, r := range s[:limit] {
|
||||
if r == '.' || r == '?' || r == '!' || r == '…' {
|
||||
last = i + len(string(r))
|
||||
}
|
||||
}
|
||||
return last
|
||||
}
|
||||
|
||||
func byteIndexByRunes(s string, n int) int {
|
||||
if n <= 0 {
|
||||
return 0
|
||||
}
|
||||
i := 0
|
||||
for pos := range s {
|
||||
if i == n {
|
||||
return pos
|
||||
}
|
||||
i++
|
||||
}
|
||||
return len(s)
|
||||
}
|
||||
|
||||
var markdownPattern = regexp.MustCompile(`[*_` + "`" + `#>\[\]]`)
|
||||
|
||||
func normalizeForSpeech(s string) string {
|
||||
s = markdownPattern.ReplaceAllString(s, "")
|
||||
s = strings.ReplaceAll(s, "\n", " ")
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package provider
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strings"
|
||||
|
||||
"ai-operator/internal/ai"
|
||||
"ai-operator/internal/ai/fake"
|
||||
realtime "ai-operator/internal/ai/openai/realtime"
|
||||
"ai-operator/internal/ai/pipeline"
|
||||
"ai-operator/internal/config"
|
||||
)
|
||||
|
||||
func NewVoiceProvider(cfg config.Config, logger *slog.Logger) (ai.VoiceProvider, error) {
|
||||
switch strings.TrimSpace(cfg.Voice.Provider) {
|
||||
case "", "fake":
|
||||
return fake.New(), nil
|
||||
case "openai_realtime":
|
||||
if cfg.OpenAI.APIKey == "" {
|
||||
return nil, errors.New("OPENAI_API_KEY is required for openai_realtime provider")
|
||||
}
|
||||
return realtime.NewProvider(cfg, logger), nil
|
||||
case "pipeline_elevenlabs", "pipeline_elevenlabs_streaming":
|
||||
if cfg.Eleven.APIKey == "" {
|
||||
return nil, errors.New("ELEVENLABS_API_KEY is required for pipeline_elevenlabs_streaming provider")
|
||||
}
|
||||
if cfg.Eleven.VoiceIDRU == "" {
|
||||
return nil, errors.New("ELEVENLABS_VOICE_ID_RU is required for pipeline_elevenlabs_streaming provider")
|
||||
}
|
||||
if cfg.OpenAI.APIKey == "" {
|
||||
return nil, errors.New("OPENAI_API_KEY is required for pipeline_elevenlabs_streaming provider")
|
||||
}
|
||||
return pipeline.NewStreamingProvider(cfg, logger), nil
|
||||
default:
|
||||
return nil, errors.New("unknown voice provider")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package provider
|
||||
|
||||
import (
|
||||
"ai-operator/internal/config"
|
||||
"io"
|
||||
"log/slog"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestFactory(t *testing.T) {
|
||||
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
if p, err := NewVoiceProvider(config.Config{Voice: config.VoiceConfig{Provider: "fake"}}, logger); err != nil || p == nil {
|
||||
t.Fatalf("fake %v", err)
|
||||
}
|
||||
if _, err := NewVoiceProvider(config.Config{Voice: config.VoiceConfig{Provider: "openai_realtime"}}, logger); err == nil {
|
||||
t.Fatal("want key error")
|
||||
}
|
||||
if _, err := NewVoiceProvider(config.Config{Voice: config.VoiceConfig{Provider: "bad"}}, logger); err == nil {
|
||||
t.Fatal("want unknown error")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package stt
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
func ParseElevenLabsEvent(data []byte) (Event, error) {
|
||||
var m map[string]any
|
||||
if err := json.Unmarshal(data, &m); err != nil {
|
||||
return Event{}, err
|
||||
}
|
||||
typ := strings.ToLower(toString(m["type"]))
|
||||
text := firstString(m, "text", "transcript", "partial", "final")
|
||||
lang := firstString(m, "language", "language_code")
|
||||
ev := Event{Text: text, Language: lang, At: time.Now().UTC(), Metadata: map[string]any{}}
|
||||
switch typ {
|
||||
case "partial_transcript", "partial", "transcript.partial":
|
||||
ev.Type = EventPartialTranscript
|
||||
case "final_transcript", "final", "transcript.final":
|
||||
ev.Type = EventFinalTranscript
|
||||
case "committed_transcript", "committed", "transcript.committed":
|
||||
ev.Type = EventCommittedTranscript
|
||||
case "speech_started", "speech.start", "vad.speech_started":
|
||||
ev.Type = EventSpeechStarted
|
||||
case "speech_ended", "speech.end", "vad.speech_ended":
|
||||
ev.Type = EventSpeechEnded
|
||||
case "error":
|
||||
ev.Type = EventError
|
||||
ev.Error = firstString(m, "error", "message")
|
||||
case "closed":
|
||||
ev.Type = EventClosed
|
||||
default:
|
||||
if final, _ := m["is_final"].(bool); final && text != "" {
|
||||
ev.Type = EventCommittedTranscript
|
||||
} else if text != "" {
|
||||
ev.Type = EventPartialTranscript
|
||||
}
|
||||
}
|
||||
return ev, nil
|
||||
}
|
||||
|
||||
func firstString(m map[string]any, keys ...string) string {
|
||||
for _, k := range keys {
|
||||
if s := toString(m[k]); s != "" {
|
||||
return s
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func toString(v any) string {
|
||||
if s, ok := v.(string); ok {
|
||||
return s
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package stt
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseElevenLabsEvent(t *testing.T) {
|
||||
ev, err := ParseElevenLabsEvent([]byte(`{"type":"partial_transcript","text":"Сколько стоит"}`))
|
||||
if err != nil || ev.Type != EventPartialTranscript || ev.Text == "" {
|
||||
t.Fatalf("ev=%+v err=%v", ev, err)
|
||||
}
|
||||
ev, err = ParseElevenLabsEvent([]byte(`{"type":"transcript.final","text":"готово"}`))
|
||||
if err != nil || ev.Type != EventFinalTranscript {
|
||||
t.Fatalf("ev=%+v err=%v", ev, err)
|
||||
}
|
||||
ev, err = ParseElevenLabsEvent([]byte(`{"is_final":true,"transcript":"готово"}`))
|
||||
if err != nil || ev.Type != EventCommittedTranscript {
|
||||
t.Fatalf("ev=%+v err=%v", ev, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
package stt
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"ai-operator/internal/config"
|
||||
"ai-operator/internal/media"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
type ElevenLabsRealtime struct {
|
||||
cfg config.Config
|
||||
conn *websocket.Conn
|
||||
events chan Event
|
||||
mu sync.Mutex
|
||||
closed bool
|
||||
}
|
||||
|
||||
func NewElevenLabsRealtime(cfg config.Config) *ElevenLabsRealtime {
|
||||
return &ElevenLabsRealtime{cfg: cfg, events: make(chan Event, 64)}
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) Start(ctx context.Context, req StreamRequest) error {
|
||||
if p.cfg.Eleven.APIKey == "" {
|
||||
return errors.New("ELEVENLABS_API_KEY is required for ElevenLabs realtime STT")
|
||||
}
|
||||
u, err := url.Parse(p.cfg.Eleven.STTURL)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("model_id", nonEmpty(req.Model, p.cfg.STT.Model))
|
||||
q.Set("audio_format", nonEmpty(req.InputFormat, p.cfg.STT.InputFormat))
|
||||
q.Set("commit_strategy", "manual")
|
||||
if req.LanguageCode != "" {
|
||||
q.Set("language_code", req.LanguageCode)
|
||||
}
|
||||
u.RawQuery = q.Encode()
|
||||
h := http.Header{}
|
||||
h.Set("xi-api-key", p.cfg.Eleven.APIKey)
|
||||
d := websocket.Dialer{HandshakeTimeout: p.cfg.STT.Timeout}
|
||||
conn, _, err := d.DialContext(ctx, u.String(), h)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.mu.Lock()
|
||||
p.conn = conn
|
||||
p.events = make(chan Event, 64)
|
||||
p.closed = false
|
||||
p.mu.Unlock()
|
||||
go p.readLoop()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) SendAudio(ctx context.Context, chunk media.AudioChunk) error {
|
||||
p.mu.Lock()
|
||||
conn := p.conn
|
||||
p.mu.Unlock()
|
||||
if conn == nil {
|
||||
return errors.New("elevenlabs stt websocket not connected")
|
||||
}
|
||||
return conn.WriteMessage(websocket.BinaryMessage, chunk.Data)
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) Events() <-chan Event { return p.events }
|
||||
|
||||
func (p *ElevenLabsRealtime) Close(ctx context.Context) error {
|
||||
p.mu.Lock()
|
||||
conn := p.conn
|
||||
p.conn = nil
|
||||
p.mu.Unlock()
|
||||
if conn != nil {
|
||||
_ = conn.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""), time.Now().Add(time.Second))
|
||||
_ = conn.Close()
|
||||
}
|
||||
p.closeEvents()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) readLoop() {
|
||||
for {
|
||||
p.mu.Lock()
|
||||
conn := p.conn
|
||||
p.mu.Unlock()
|
||||
if conn == nil {
|
||||
p.closeEvents()
|
||||
return
|
||||
}
|
||||
_, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
p.emit(Event{Type: EventClosed, Error: sanitizeCloseError(err), At: time.Now().UTC()})
|
||||
p.closeEvents()
|
||||
return
|
||||
}
|
||||
ev, err := ParseElevenLabsEvent(data)
|
||||
if err != nil {
|
||||
p.emit(Event{Type: EventError, Error: err.Error(), At: time.Now().UTC()})
|
||||
continue
|
||||
}
|
||||
if ev.Type != "" {
|
||||
p.emit(ev)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func sanitizeCloseError(err error) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
return err.Error()
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) emit(ev Event) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if p.closed {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case p.events <- ev:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) closeEvents() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if !p.closed {
|
||||
close(p.events)
|
||||
p.closed = true
|
||||
}
|
||||
}
|
||||
|
||||
func nonEmpty(v, fallback string) string {
|
||||
if v != "" {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func nonZero(v, fallback int) int {
|
||||
if v != 0 {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
package stt
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"ai-operator/internal/config"
|
||||
"ai-operator/internal/media"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
type ElevenLabsRealtime struct {
|
||||
cfg config.Config
|
||||
conn *websocket.Conn
|
||||
events chan Event
|
||||
mu sync.Mutex
|
||||
closed bool
|
||||
}
|
||||
|
||||
func NewElevenLabsRealtime(cfg config.Config) *ElevenLabsRealtime {
|
||||
return &ElevenLabsRealtime{cfg: cfg, events: make(chan Event, 64)}
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) Start(ctx context.Context, req StreamRequest) error {
|
||||
if p.cfg.Eleven.APIKey == "" {
|
||||
return errors.New("ELEVENLABS_API_KEY is required for ElevenLabs realtime STT")
|
||||
}
|
||||
u, err := url.Parse(p.cfg.Eleven.STTURL)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("model_id", nonEmpty(req.Model, p.cfg.STT.Model))
|
||||
q.Set("sample_rate", strconv.Itoa(nonZero(req.SampleRate, p.cfg.STT.SampleRate)))
|
||||
q.Set("input_format", nonEmpty(req.InputFormat, p.cfg.STT.InputFormat))
|
||||
if req.LanguageCode != "" {
|
||||
q.Set("language_code", req.LanguageCode)
|
||||
}
|
||||
u.RawQuery = q.Encode()
|
||||
h := http.Header{}
|
||||
h.Set("xi-api-key", p.cfg.Eleven.APIKey)
|
||||
d := websocket.Dialer{HandshakeTimeout: p.cfg.STT.Timeout}
|
||||
conn, _, err := d.DialContext(ctx, u.String(), h)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.mu.Lock()
|
||||
p.conn = conn
|
||||
p.events = make(chan Event, 64)
|
||||
p.closed = false
|
||||
p.mu.Unlock()
|
||||
init := map[string]any{
|
||||
"type": "start",
|
||||
"model_id": nonEmpty(req.Model, p.cfg.STT.Model),
|
||||
"sample_rate": nonZero(req.SampleRate, p.cfg.STT.SampleRate),
|
||||
"input_format": nonEmpty(req.InputFormat, p.cfg.STT.InputFormat),
|
||||
"partial_transcripts": p.cfg.STT.PartialEnabled,
|
||||
"committed_only_for_llm": p.cfg.STT.CommittedOnlyForLLM,
|
||||
"language_detection_enabled": req.LanguageAuto,
|
||||
}
|
||||
if req.LanguageCode != "" {
|
||||
init["language_code"] = req.LanguageCode
|
||||
}
|
||||
if err := conn.WriteJSON(init); err != nil {
|
||||
_ = conn.Close()
|
||||
return err
|
||||
}
|
||||
go p.readLoop()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) SendAudio(ctx context.Context, chunk media.AudioChunk) error {
|
||||
p.mu.Lock()
|
||||
conn := p.conn
|
||||
p.mu.Unlock()
|
||||
if conn == nil {
|
||||
return errors.New("elevenlabs stt websocket not connected")
|
||||
}
|
||||
return conn.WriteMessage(websocket.BinaryMessage, chunk.Data)
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) Events() <-chan Event { return p.events }
|
||||
|
||||
func (p *ElevenLabsRealtime) Close(ctx context.Context) error {
|
||||
p.mu.Lock()
|
||||
conn := p.conn
|
||||
p.conn = nil
|
||||
p.mu.Unlock()
|
||||
if conn != nil {
|
||||
_ = conn.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""), time.Now().Add(time.Second))
|
||||
_ = conn.Close()
|
||||
}
|
||||
p.closeEvents()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) readLoop() {
|
||||
for {
|
||||
p.mu.Lock()
|
||||
conn := p.conn
|
||||
p.mu.Unlock()
|
||||
if conn == nil {
|
||||
p.closeEvents()
|
||||
return
|
||||
}
|
||||
_, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
p.emit(Event{Type: EventClosed, Error: sanitizeCloseError(err), At: time.Now().UTC()})
|
||||
p.closeEvents()
|
||||
return
|
||||
}
|
||||
ev, err := ParseElevenLabsEvent(data)
|
||||
if err != nil {
|
||||
p.emit(Event{Type: EventError, Error: err.Error(), At: time.Now().UTC()})
|
||||
continue
|
||||
}
|
||||
if ev.Type != "" {
|
||||
p.emit(ev)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func sanitizeCloseError(err error) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
return err.Error()
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) emit(ev Event) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if p.closed {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case p.events <- ev:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (p *ElevenLabsRealtime) closeEvents() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if !p.closed {
|
||||
close(p.events)
|
||||
p.closed = true
|
||||
}
|
||||
}
|
||||
|
||||
func nonEmpty(v, fallback string) string {
|
||||
if v != "" {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func nonZero(v, fallback int) int {
|
||||
if v != 0 {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package stt
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"ai-operator/internal/media"
|
||||
)
|
||||
|
||||
type EventType string
|
||||
|
||||
const (
|
||||
EventPartialTranscript EventType = "partial_transcript"
|
||||
EventFinalTranscript EventType = "final_transcript"
|
||||
EventCommittedTranscript EventType = "committed_transcript"
|
||||
EventSpeechStarted EventType = "speech_started"
|
||||
EventSpeechEnded EventType = "speech_ended"
|
||||
EventError EventType = "error"
|
||||
EventClosed EventType = "closed"
|
||||
)
|
||||
|
||||
type StreamRequest struct {
|
||||
CallID string
|
||||
Model string
|
||||
LanguageCode string
|
||||
LanguageAuto bool
|
||||
SampleRate int
|
||||
InputFormat string
|
||||
}
|
||||
|
||||
type Event struct {
|
||||
Type EventType
|
||||
Text string
|
||||
Language string
|
||||
Error string
|
||||
At time.Time
|
||||
Metadata map[string]any
|
||||
}
|
||||
|
||||
type StreamingProvider interface {
|
||||
Start(ctx context.Context, req StreamRequest) error
|
||||
SendAudio(ctx context.Context, chunk media.AudioChunk) error
|
||||
Events() <-chan Event
|
||||
Close(ctx context.Context) error
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package tts
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"ai-operator/internal/config"
|
||||
)
|
||||
|
||||
var tagPattern = regexp.MustCompile(`\[[^\]]+\]`)
|
||||
|
||||
func NaturalizeForVoice(text string, language string, cfg config.NaturalnessConfig) string {
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" || !cfg.Enabled || !cfg.AudioTagsEnabled || !cfg.AllowNonverbalTags || cfg.MaxAudioTagsPerResponse <= 0 {
|
||||
return text
|
||||
}
|
||||
lower := strings.ToLower(text)
|
||||
if containsAny(lower, []string{"тариф", "оплат", "безопас", "документ", "адрес", "телефон", "газ", "құжат", "төлем", "қауіпсіз", "мекенжай"}) {
|
||||
return "[calmly] " + text
|
||||
}
|
||||
return "[warmly] " + insertBriefPause(text, cfg.MaxAudioTagsPerResponse)
|
||||
}
|
||||
|
||||
func RemoveAudioTags(text string) string {
|
||||
return strings.TrimSpace(tagPattern.ReplaceAllString(text, ""))
|
||||
}
|
||||
|
||||
func insertBriefPause(text string, maxTags int) string {
|
||||
if maxTags < 2 {
|
||||
return text
|
||||
}
|
||||
if i := strings.Index(text, ". "); i > 0 {
|
||||
return text[:i+1] + " [brief pause] " + text[i+2:]
|
||||
}
|
||||
return text
|
||||
}
|
||||
|
||||
func containsAny(s string, needles []string) bool {
|
||||
for _, n := range needles {
|
||||
if strings.Contains(s, n) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
package tts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"ai-operator/internal/config"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
type ElevenLabsWS struct {
|
||||
cfg config.Config
|
||||
conn *websocket.Conn
|
||||
audio chan AudioChunk
|
||||
mu sync.Mutex
|
||||
closed bool
|
||||
}
|
||||
|
||||
func NewElevenLabsWS(cfg config.Config) *ElevenLabsWS {
|
||||
return &ElevenLabsWS{cfg: cfg, audio: make(chan AudioChunk, 64)}
|
||||
}
|
||||
|
||||
func (p *ElevenLabsWS) Start(ctx context.Context, req TTSStreamRequest) error {
|
||||
if p.cfg.Eleven.APIKey == "" {
|
||||
return errors.New("ELEVENLABS_API_KEY is required for ElevenLabs streaming TTS")
|
||||
}
|
||||
if req.VoiceID == "" {
|
||||
return errors.New("ElevenLabs voice id is required for streaming TTS")
|
||||
}
|
||||
base := strings.TrimRight(p.cfg.Eleven.TTSURL, "/") + "/" + url.PathEscape(req.VoiceID) + "/stream-input"
|
||||
u, err := url.Parse(base)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("model_id", nonEmpty(req.ModelID, p.cfg.Eleven.TTSModelID))
|
||||
q.Set("output_format", nonEmpty(req.OutputFormat, p.cfg.Eleven.TTSOutputFormat))
|
||||
q.Set("optimize_streaming_latency", strconv.Itoa(p.cfg.Eleven.TTSOptimizeStreamingLatency))
|
||||
u.RawQuery = q.Encode()
|
||||
h := http.Header{}
|
||||
h.Set("xi-api-key", p.cfg.Eleven.APIKey)
|
||||
d := websocket.Dialer{HandshakeTimeout: p.cfg.Eleven.TTSTimeout}
|
||||
conn, _, err := d.DialContext(ctx, u.String(), h)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.mu.Lock()
|
||||
p.conn = conn
|
||||
p.audio = make(chan AudioChunk, 64)
|
||||
p.closed = false
|
||||
p.mu.Unlock()
|
||||
bos := map[string]any{
|
||||
"text": " ",
|
||||
"voice_settings": map[string]any{
|
||||
"stability": p.cfg.Eleven.TTSStability,
|
||||
"similarity_boost": p.cfg.Eleven.TTSSimilarityBoost,
|
||||
"style": p.cfg.Eleven.TTSStyle,
|
||||
"use_speaker_boost": p.cfg.Eleven.TTSUseSpeakerBoost,
|
||||
},
|
||||
}
|
||||
if err := conn.WriteJSON(bos); err != nil {
|
||||
return err
|
||||
}
|
||||
go p.readLoop()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *ElevenLabsWS) SendText(ctx context.Context, text string, flush bool) error {
|
||||
p.mu.Lock()
|
||||
conn := p.conn
|
||||
p.mu.Unlock()
|
||||
if conn == nil {
|
||||
return errors.New("elevenlabs tts websocket not connected")
|
||||
}
|
||||
msg := map[string]any{"text": text, "try_trigger_generation": flush}
|
||||
if flush {
|
||||
msg["flush"] = true
|
||||
}
|
||||
return conn.WriteJSON(msg)
|
||||
}
|
||||
|
||||
func (p *ElevenLabsWS) Audio() <-chan AudioChunk { return p.audio }
|
||||
|
||||
func (p *ElevenLabsWS) Close(ctx context.Context) error {
|
||||
p.mu.Lock()
|
||||
conn := p.conn
|
||||
p.conn = nil
|
||||
p.mu.Unlock()
|
||||
if conn != nil {
|
||||
_ = conn.WriteJSON(map[string]any{"text": ""})
|
||||
_ = conn.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""), time.Now().Add(time.Second))
|
||||
_ = conn.Close()
|
||||
}
|
||||
p.closeAudio()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *ElevenLabsWS) readLoop() {
|
||||
for {
|
||||
p.mu.Lock()
|
||||
conn := p.conn
|
||||
p.mu.Unlock()
|
||||
if conn == nil {
|
||||
p.closeAudio()
|
||||
return
|
||||
}
|
||||
_, data, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
p.closeAudio()
|
||||
return
|
||||
}
|
||||
audio := parseAudio(data)
|
||||
if len(audio) == 0 {
|
||||
continue
|
||||
}
|
||||
p.emit(AudioChunk{Data: audio, Timestamp: time.Now().UTC()})
|
||||
}
|
||||
}
|
||||
|
||||
func parseAudio(data []byte) []byte {
|
||||
var m map[string]any
|
||||
if json.Unmarshal(data, &m) == nil {
|
||||
for _, k := range []string{"audio", "audio_base64"} {
|
||||
if s, ok := m[k].(string); ok && s != "" {
|
||||
b, _ := base64.StdEncoding.DecodeString(s)
|
||||
return b
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
func (p *ElevenLabsWS) emit(ch AudioChunk) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if p.closed {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case p.audio <- ch:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (p *ElevenLabsWS) closeAudio() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if !p.closed {
|
||||
close(p.audio)
|
||||
p.closed = true
|
||||
}
|
||||
}
|
||||
|
||||
func nonEmpty(v, fallback string) string {
|
||||
if v != "" {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
package tts
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseAudioIgnoresControlJSON(t *testing.T) {
|
||||
if got := parseAudio([]byte(`{"isFinal":true}`)); len(got) != 0 {
|
||||
t.Fatalf("control json parsed as audio: %d bytes", len(got))
|
||||
}
|
||||
if got := parseAudio([]byte(`{"audio":"AQI="}`)); len(got) != 2 {
|
||||
t.Fatalf("audio json not decoded: %d bytes", len(got))
|
||||
}
|
||||
if got := parseAudio([]byte{0, 1, 0, 1}); len(got) != 4 {
|
||||
t.Fatalf("binary audio not passed through: %d bytes", len(got))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package tts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
type AudioChunk struct {
|
||||
Data []byte
|
||||
Timestamp time.Time
|
||||
}
|
||||
|
||||
type TTSStreamRequest struct {
|
||||
CallID string
|
||||
Language string
|
||||
VoiceID string
|
||||
ModelID string
|
||||
OutputFormat string
|
||||
SampleRate int
|
||||
}
|
||||
|
||||
type StreamingTTS interface {
|
||||
Start(ctx context.Context, req TTSStreamRequest) error
|
||||
SendText(ctx context.Context, text string, flush bool) error
|
||||
Audio() <-chan AudioChunk
|
||||
Close(ctx context.Context) error
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"ai-operator/internal/media"
|
||||
)
|
||||
|
||||
type VoiceProvider interface {
|
||||
StartSession(ctx context.Context, config VoiceSessionConfig) error
|
||||
SendAudio(ctx context.Context, chunk media.AudioChunk) error
|
||||
SendToolResult(ctx context.Context, result ToolResult) error
|
||||
Close(ctx context.Context) error
|
||||
Events() <-chan VoiceEvent
|
||||
Stats() VoiceProviderStats
|
||||
}
|
||||
|
||||
type VoiceSessionConfig struct {
|
||||
CallID string
|
||||
ProviderSessionID string
|
||||
Language string
|
||||
RegionCode string
|
||||
SystemPrompt string
|
||||
InputAudioFormat string
|
||||
OutputAudioFormat string
|
||||
InputSampleRate int
|
||||
OutputSampleRate int
|
||||
Metadata map[string]string
|
||||
}
|
||||
|
||||
type VoiceEventType string
|
||||
|
||||
const (
|
||||
VoiceEventSessionStarted VoiceEventType = "session.started"
|
||||
VoiceEventSessionUpdated VoiceEventType = "session.updated"
|
||||
VoiceEventUserTranscriptDelta VoiceEventType = "user.transcript.delta"
|
||||
VoiceEventUserTranscriptDone VoiceEventType = "user.transcript.done"
|
||||
VoiceEventAssistantTranscriptDelta VoiceEventType = "assistant.transcript.delta"
|
||||
VoiceEventAssistantTranscriptDone VoiceEventType = "assistant.transcript.done"
|
||||
VoiceEventAssistantAudioDelta VoiceEventType = "assistant.audio.delta"
|
||||
VoiceEventAssistantAudioDone VoiceEventType = "assistant.audio.done"
|
||||
VoiceEventToolCall VoiceEventType = "tool.call"
|
||||
VoiceEventInterruption VoiceEventType = "interruption"
|
||||
VoiceEventRateLimitsUpdated VoiceEventType = "rate_limits.updated"
|
||||
VoiceEventError VoiceEventType = "error"
|
||||
VoiceEventClosed VoiceEventType = "closed"
|
||||
)
|
||||
|
||||
type VoiceEvent struct {
|
||||
Type VoiceEventType
|
||||
CallID string
|
||||
Text string
|
||||
Audio []byte
|
||||
ToolCall *ToolCall
|
||||
Error string
|
||||
Metadata map[string]any
|
||||
At time.Time
|
||||
}
|
||||
|
||||
type ToolCall struct {
|
||||
ID string
|
||||
Name string
|
||||
Arguments map[string]any
|
||||
RawArguments string
|
||||
}
|
||||
type ToolResult struct {
|
||||
CallID string
|
||||
ToolCallID string
|
||||
Result any
|
||||
Error string
|
||||
}
|
||||
|
||||
type VoiceProviderStats struct {
|
||||
InputAudioFrames int64
|
||||
InputAudioBytes int64
|
||||
OutputAudioFrames int64
|
||||
OutputAudioBytes int64
|
||||
EventsReceived int64
|
||||
EventsSent int64
|
||||
ToolCallsReceived int64
|
||||
Errors int64
|
||||
StartedAt *time.Time
|
||||
ClosedAt *time.Time
|
||||
}
|
||||
Reference in New Issue
Block a user