use std::collections::HashMap;
use futures_util::{SinkExt, StreamExt};
use tokio::sync::{mpsc, watch};
use crate::error::{AriError, Result};
#[derive(Debug, Clone, serde::Deserialize)]
#[serde(tag = "event")]
#[non_exhaustive]
pub enum MediaEvent {
#[serde(rename = "MEDIA_START")]
MediaStart {
connection_id: String,
channel: String,
channel_id: String,
format: String,
optimal_frame_size: u32,
ptime: u32,
#[serde(default)]
channel_variables: HashMap<String, String>,
},
#[serde(rename = "DTMF_END")]
DtmfEnd { digit: String, duration_ms: u32 },
#[serde(rename = "MEDIA_XOFF")]
MediaXoff,
#[serde(rename = "MEDIA_XON")]
MediaXon,
#[serde(rename = "STATUS")]
Status {
channel: String,
format: String,
queue_size: u32,
buffering_active: bool,
media_paused: bool,
},
#[serde(rename = "MEDIA_BUFFERING_COMPLETED")]
MediaBufferingCompleted {
#[serde(default)]
correlation_id: Option<String>,
},
#[serde(rename = "MEDIA_MARK_PROCESSED")]
MediaMarkProcessed,
#[serde(rename = "QUEUE_DRAINED")]
QueueDrained,
}
#[derive(Debug, Clone, serde::Serialize)]
#[serde(tag = "command")]
#[non_exhaustive]
pub enum MediaCommand {
#[serde(rename = "ANSWER")]
Answer,
#[serde(rename = "HANGUP")]
Hangup {
#[serde(skip_serializing_if = "Option::is_none")]
cause: Option<u32>,
},
#[serde(rename = "START_MEDIA_BUFFERING")]
StartMediaBuffering,
#[serde(rename = "STOP_MEDIA_BUFFERING")]
StopMediaBuffering {
#[serde(skip_serializing_if = "Option::is_none")]
correlation_id: Option<String>,
},
#[serde(rename = "FLUSH_MEDIA")]
FlushMedia,
#[serde(rename = "PAUSE_MEDIA")]
PauseMedia,
#[serde(rename = "CONTINUE_MEDIA")]
ContinueMedia,
#[serde(rename = "MARK_MEDIA")]
MarkMedia,
#[serde(rename = "GET_STATUS")]
GetStatus,
#[serde(rename = "REPORT_QUEUE_DRAINED")]
ReportQueueDrained,
}
enum InternalCmd {
Audio(Vec<u8>),
Command(String),
}
pub struct MediaChannel {
event_rx: mpsc::Receiver<MediaEvent>,
audio_rx: mpsc::Receiver<Vec<u8>>,
command_tx: mpsc::Sender<InternalCmd>,
shutdown_tx: watch::Sender<bool>,
task_handle: tokio::task::JoinHandle<()>,
}
impl std::fmt::Debug for MediaChannel {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MediaChannel")
.field("connected", &!self.command_tx.is_closed())
.finish()
}
}
type OutboundWsStream =
tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>>;
type AcceptedWsStream = tokio_tungstenite::WebSocketStream<tokio::net::TcpStream>;
impl MediaChannel {
pub async fn connect(url: &str) -> Result<Self> {
let (ws_stream, _) = tokio_tungstenite::connect_async(url)
.await
.map_err(|e| AriError::WebSocket(e.to_string()))?;
Ok(Self::spawn_outbound(ws_stream))
}
pub fn from_accepted(ws_stream: AcceptedWsStream) -> Self {
let (event_tx, event_rx) = mpsc::channel(64);
let (audio_tx, audio_rx) = mpsc::channel(256);
let (command_tx, command_rx) = mpsc::channel(64);
let (shutdown_tx, shutdown_rx) = watch::channel(false);
let task_handle = tokio::spawn(media_loop(
ws_stream,
event_tx,
audio_tx,
command_rx,
shutdown_rx,
));
Self {
event_rx,
audio_rx,
command_tx,
shutdown_tx,
task_handle,
}
}
fn spawn_outbound(ws_stream: OutboundWsStream) -> Self {
let (event_tx, event_rx) = mpsc::channel(64);
let (audio_tx, audio_rx) = mpsc::channel(256);
let (command_tx, command_rx) = mpsc::channel(64);
let (shutdown_tx, shutdown_rx) = watch::channel(false);
let task_handle = tokio::spawn(media_loop(
ws_stream,
event_tx,
audio_tx,
command_rx,
shutdown_rx,
));
Self {
event_rx,
audio_rx,
command_tx,
shutdown_tx,
task_handle,
}
}
pub async fn recv_event(&mut self) -> Option<MediaEvent> {
self.event_rx.recv().await
}
pub async fn recv_audio(&mut self) -> Option<Vec<u8>> {
self.audio_rx.recv().await
}
pub async fn send_command(&self, cmd: MediaCommand) -> Result<()> {
let json = serde_json::to_string(&cmd).map_err(AriError::Json)?;
self.command_tx
.send(InternalCmd::Command(json))
.await
.map_err(|_| AriError::Disconnected)
}
pub async fn send_audio(&self, data: Vec<u8>) -> Result<()> {
if data.len() > 65500 {
return Err(AriError::WebSocket(format!(
"audio frame too large: {} bytes (max 65500)",
data.len()
)));
}
self.command_tx
.send(InternalCmd::Audio(data))
.await
.map_err(|_| AriError::Disconnected)
}
pub async fn answer(&self) -> Result<()> {
self.send_command(MediaCommand::Answer).await
}
pub async fn hangup(&self, cause: Option<u32>) -> Result<()> {
self.send_command(MediaCommand::Hangup { cause }).await
}
pub async fn start_buffering(&self) -> Result<()> {
self.send_command(MediaCommand::StartMediaBuffering).await
}
pub async fn stop_buffering(&self, correlation_id: Option<String>) -> Result<()> {
self.send_command(MediaCommand::StopMediaBuffering { correlation_id })
.await
}
pub async fn flush(&self) -> Result<()> {
self.send_command(MediaCommand::FlushMedia).await
}
pub async fn pause(&self) -> Result<()> {
self.send_command(MediaCommand::PauseMedia).await
}
pub async fn resume(&self) -> Result<()> {
self.send_command(MediaCommand::ContinueMedia).await
}
pub async fn mark(&self) -> Result<()> {
self.send_command(MediaCommand::MarkMedia).await
}
pub async fn get_status(&self) -> Result<()> {
self.send_command(MediaCommand::GetStatus).await
}
pub async fn report_queue_drained(&self) -> Result<()> {
self.send_command(MediaCommand::ReportQueueDrained).await
}
pub fn disconnect(&self) {
let _ = self.shutdown_tx.send(true);
self.task_handle.abort();
}
}
impl Drop for MediaChannel {
fn drop(&mut self) {
self.disconnect();
}
}
fn is_critical_event(event: &MediaEvent) -> bool {
matches!(
event,
MediaEvent::MediaStart { .. } | MediaEvent::MediaBufferingCompleted { .. }
)
}
async fn media_loop<S>(
ws_stream: tokio_tungstenite::WebSocketStream<S>,
event_tx: mpsc::Sender<MediaEvent>,
audio_tx: mpsc::Sender<Vec<u8>>,
mut command_rx: mpsc::Receiver<InternalCmd>,
mut shutdown_rx: watch::Receiver<bool>,
) where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
use tokio_tungstenite::tungstenite::Message;
let (mut write, mut read) = ws_stream.split();
loop {
tokio::select! {
msg = read.next() => {
match msg {
Some(Ok(Message::Text(text))) => {
match serde_json::from_str::<MediaEvent>(&text) {
Ok(event) => {
if is_critical_event(&event) {
if event_tx.send(event).await.is_err() {
return;
}
} else {
match event_tx.try_send(event) {
Ok(()) => {}
Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => {
tracing::warn!("media event channel full, dropping event");
}
Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => {
return;
}
}
}
}
Err(e) => {
tracing::warn!(
error = %e,
"failed to parse media event"
);
}
}
}
Some(Ok(Message::Binary(data))) => {
match audio_tx.try_send(data) {
Ok(()) => {}
Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => {
tracing::debug!("audio channel full, dropping frame");
}
Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => {
return;
}
}
}
Some(Ok(Message::Close(_))) => {
tracing::debug!("media websocket closed by peer");
return;
}
Some(Err(e)) => {
tracing::warn!(error = %e, "media websocket read error");
return;
}
None => return,
_ => {}
}
}
cmd = command_rx.recv() => {
match cmd {
Some(InternalCmd::Audio(data)) => {
if let Err(e) = write.send(Message::Binary(data)).await {
tracing::warn!(error = %e, "failed to send audio frame");
return;
}
}
Some(InternalCmd::Command(json)) => {
if let Err(e) = write.send(Message::Text(json)).await {
tracing::warn!(error = %e, "failed to send media command");
return;
}
}
None => return,
}
}
_ = shutdown_rx.changed() => {
if *shutdown_rx.borrow() {
tracing::debug!("media channel shutdown requested");
return;
}
}
}
}
}