use livekit_protocol as proto;
use livekit_runtime::JoinHandle;
use prost::Message as ProtoMessage;
use std::{sync::Arc, time::Duration};
use tokio::sync::{mpsc, oneshot};
use super::{SignalError, SignalResult};
#[derive(Debug)]
enum InternalMessage {
Signal {
signal: proto::signal_request::Message,
response_chn: oneshot::Sender<SignalResult<()>>,
},
Close,
}
#[derive(Debug)]
pub(super) struct SignalStream {
internal_tx: mpsc::Sender<InternalMessage>,
read_handle: JoinHandle<()>,
write_handle: JoinHandle<()>,
}
impl SignalStream {
pub async fn connect(
url: url::Url,
token: &str,
connect_timeout: Duration,
) -> SignalResult<(Self, mpsc::UnboundedReceiver<Box<proto::signal_response::Message>>)> {
log::info!("connecting to {}", livekit_net::redact_url(&url));
super::check_token_format(token)?;
let transport = super::require_ws_client()?;
let headers = super::bearer_headers(token);
let conn = livekit_runtime::timeout(
connect_timeout,
transport.connect(url.to_string(), headers, connect_timeout.as_millis() as u64),
)
.await
.map_err(|_| SignalError::Timeout("signal connection timed out".into()))??
.connection;
let (emitter, events) = mpsc::unbounded_channel();
let (internal_tx, internal_rx) = mpsc::channel::<InternalMessage>(8);
let write_handle = livekit_runtime::spawn(Self::write_task(internal_rx, conn.clone()));
let read_handle =
livekit_runtime::spawn(Self::read_task(internal_tx.clone(), conn, emitter));
Ok((Self { internal_tx, read_handle, write_handle }, events))
}
pub async fn close(self, notify_close: bool) {
if notify_close {
let _ = self.internal_tx.send(InternalMessage::Close).await;
}
let _ = self.write_handle.await;
let _ = self.read_handle.await;
}
pub async fn send(&self, signal: proto::signal_request::Message) -> SignalResult<()> {
let (send, recv) = oneshot::channel();
let msg = InternalMessage::Signal { signal, response_chn: send };
let _ = self.internal_tx.send(msg).await;
recv.await.map_err(|_| SignalError::SendError)?
}
async fn write_task(
mut internal_rx: mpsc::Receiver<InternalMessage>,
conn: Arc<dyn livekit_net::WsConnection>,
) {
while let Some(msg) = internal_rx.recv().await {
match msg {
InternalMessage::Signal { signal, response_chn } => {
let data = proto::SignalRequest { message: Some(signal) }.encode_to_vec();
if let Err(err) = conn.send(data).await {
let _ = response_chn.send(Err(err.into()));
break;
}
let _ = response_chn.send(Ok(()));
}
InternalMessage::Close => break,
}
}
conn.close().await;
}
async fn read_task(
internal_tx: mpsc::Sender<InternalMessage>,
conn: Arc<dyn livekit_net::WsConnection>,
emitter: mpsc::UnboundedSender<Box<proto::signal_response::Message>>,
) {
loop {
match conn.recv().await {
Ok(Some(bytes)) => {
match proto::SignalResponse::decode(bytes.as_slice()) {
Ok(res) => {
if let Some(msg) = res.message {
let _ = emitter.send(Box::new(msg));
}
}
Err(e) => {
log::error!("failed to decode SignalResponse: {:?}", e);
}
}
}
Ok(None) => {
let _ = internal_tx.send(InternalMessage::Close).await;
break;
}
Err(e) => {
log::error!("websocket recv error: {:?}", e);
let _ = internal_tx.send(InternalMessage::Close).await;
break;
}
}
}
}
}