225 lines
6.0 KiB
Go
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) }
|