Files

225 lines
6.0 KiB
Go

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