201 lines
4.6 KiB
Go
201 lines
4.6 KiB
Go
package call
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"log/slog"
|
|
"sync"
|
|
"time"
|
|
|
|
"ai-operator/internal/ai"
|
|
"ai-operator/internal/audio"
|
|
"ai-operator/internal/media"
|
|
)
|
|
|
|
type MediaClient interface {
|
|
Audio() <-chan media.AudioChunk
|
|
SendAudio(ctx context.Context, data []byte) error
|
|
FlushMedia(ctx context.Context) error
|
|
Close(ctx context.Context) error
|
|
}
|
|
|
|
type AudioPump struct {
|
|
CallID string
|
|
Media MediaClient
|
|
Provider ai.VoiceProvider
|
|
AsteriskCodec media.Codec
|
|
AsteriskSampleRate int
|
|
ProviderInputSampleRate int
|
|
ProviderOutputSampleRate int
|
|
Logger *slog.Logger
|
|
|
|
ctx context.Context
|
|
cancel context.CancelFunc
|
|
wg sync.WaitGroup
|
|
mu sync.Mutex
|
|
|
|
suppressInputUntil time.Time
|
|
}
|
|
|
|
const assistantEchoTailSuppression = 900 * time.Millisecond
|
|
|
|
func (p *AudioPump) Start(ctx context.Context) error {
|
|
if p.Media == nil {
|
|
return errors.New("audio pump media client is nil")
|
|
}
|
|
if p.Provider == nil {
|
|
return errors.New("audio pump voice provider is nil")
|
|
}
|
|
if p.AsteriskSampleRate == 0 {
|
|
p.AsteriskSampleRate = 16000
|
|
}
|
|
if p.AsteriskCodec == "" {
|
|
p.AsteriskCodec = media.CodecSLIN16
|
|
}
|
|
if p.ProviderInputSampleRate == 0 {
|
|
p.ProviderInputSampleRate = 24000
|
|
}
|
|
if p.ProviderOutputSampleRate == 0 {
|
|
p.ProviderOutputSampleRate = 24000
|
|
}
|
|
p.ctx, p.cancel = context.WithCancel(ctx)
|
|
p.wg.Add(1)
|
|
go p.asteriskToProvider()
|
|
return nil
|
|
}
|
|
|
|
func (p *AudioPump) Stop(ctx context.Context) error {
|
|
if p.cancel != nil {
|
|
p.cancel()
|
|
}
|
|
done := make(chan struct{})
|
|
go func() {
|
|
p.wg.Wait()
|
|
close(done)
|
|
}()
|
|
select {
|
|
case <-done:
|
|
return nil
|
|
case <-ctx.Done():
|
|
return ctx.Err()
|
|
}
|
|
}
|
|
|
|
func (p *AudioPump) HandleProviderEvent(ctx context.Context, event ai.VoiceEvent) error {
|
|
switch event.Type {
|
|
case ai.VoiceEventAssistantAudioDelta:
|
|
if len(event.Audio) == 0 {
|
|
return nil
|
|
}
|
|
pcm, err := audio.ResamplePCM16MonoLinear(event.Audio, p.ProviderOutputSampleRate, p.AsteriskSampleRate)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
p.suppressInputFor(pcmDuration(pcm, p.AsteriskSampleRate) + assistantEchoTailSuppression)
|
|
out := p.encodeForAsterisk(pcm)
|
|
if err := p.Media.SendAudio(ctx, out); err != nil {
|
|
return err
|
|
}
|
|
p.log("assistant audio forwarded", "channel_id", p.CallID, "bytes", len(out))
|
|
case ai.VoiceEventAssistantAudioDone:
|
|
p.suppressInputFor(assistantEchoTailSuppression)
|
|
case ai.VoiceEventInterruption:
|
|
if p.inputSuppressed() {
|
|
p.log("ignored interruption during assistant playback", "channel_id", p.CallID)
|
|
return nil
|
|
}
|
|
if err := p.Media.FlushMedia(ctx); err != nil {
|
|
return err
|
|
}
|
|
p.log("media flushed on interruption", "channel_id", p.CallID)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (p *AudioPump) asteriskToProvider() {
|
|
defer p.wg.Done()
|
|
for {
|
|
select {
|
|
case <-p.ctx.Done():
|
|
return
|
|
case chunk, ok := <-p.Media.Audio():
|
|
if !ok {
|
|
return
|
|
}
|
|
if len(chunk.Data) == 0 {
|
|
continue
|
|
}
|
|
if p.inputSuppressed() {
|
|
continue
|
|
}
|
|
pcm := p.decodeFromAsterisk(chunk)
|
|
in, err := audio.ResamplePCM16MonoLinear(pcm, p.AsteriskSampleRate, p.ProviderInputSampleRate)
|
|
if err != nil {
|
|
p.log("audio pump input resample failed", "channel_id", p.CallID, "error", err)
|
|
continue
|
|
}
|
|
sendCtx, cancel := context.WithTimeout(p.ctx, 2*time.Second)
|
|
err = p.Provider.SendAudio(sendCtx, media.AudioChunk{CallID: p.CallID, Data: in, Codec: media.CodecSLIN16, Timestamp: chunk.Timestamp})
|
|
cancel()
|
|
if err != nil {
|
|
p.log("audio pump provider send failed", "channel_id", p.CallID, "error", err)
|
|
return
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
func (p *AudioPump) suppressInputFor(d time.Duration) {
|
|
if d <= 0 {
|
|
return
|
|
}
|
|
until := time.Now().Add(d)
|
|
p.mu.Lock()
|
|
if until.After(p.suppressInputUntil) {
|
|
p.suppressInputUntil = until
|
|
}
|
|
p.mu.Unlock()
|
|
}
|
|
|
|
func (p *AudioPump) inputSuppressed() bool {
|
|
p.mu.Lock()
|
|
defer p.mu.Unlock()
|
|
return time.Now().Before(p.suppressInputUntil)
|
|
}
|
|
|
|
func pcmDuration(pcm []byte, sampleRate int) time.Duration {
|
|
if sampleRate <= 0 || len(pcm) == 0 {
|
|
return 0
|
|
}
|
|
samples := len(pcm) / 2
|
|
return time.Duration(samples) * time.Second / time.Duration(sampleRate)
|
|
}
|
|
|
|
func (p *AudioPump) decodeFromAsterisk(chunk media.AudioChunk) []byte {
|
|
codec := chunk.Codec
|
|
if codec == "" {
|
|
codec = p.AsteriskCodec
|
|
}
|
|
switch codec {
|
|
case media.CodecULaw:
|
|
return audio.DecodeULaw(chunk.Data)
|
|
default:
|
|
return chunk.Data
|
|
}
|
|
}
|
|
|
|
func (p *AudioPump) encodeForAsterisk(pcm []byte) []byte {
|
|
switch p.AsteriskCodec {
|
|
case media.CodecULaw:
|
|
return audio.EncodeULaw(pcm)
|
|
default:
|
|
return pcm
|
|
}
|
|
}
|
|
|
|
func (p *AudioPump) log(msg string, args ...any) {
|
|
if p.Logger != nil {
|
|
p.Logger.Info(msg, args...)
|
|
}
|
|
}
|