use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use std::time::Duration;
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use tokio::sync::{broadcast, mpsc, oneshot};
use tokio::task::JoinHandle;
use tokio_stream::wrappers::BroadcastStream;
use tokio_stream::wrappers::errors::BroadcastStreamRecvError;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::header::{HeaderName, HeaderValue};
use tokio_tungstenite::tungstenite::{Error as WsError, Message};
use super::auth::{TokenProvider, build_auth_query_url};
#[cfg(feature = "compression")]
use super::compression::{inflate, with_compress_query};
use super::protocol::{
ChannelEvent, InboundFrame, OutboundFrame, Subscription, parse_inbound_frame,
};
use super::reconnect::{ReconnectOptions, ReconnectPolicy};
use super::subscription::SubscriptionManager;
use crate::config::DEFAULT_WS_URL;
use crate::error::{RadionError, Result};
const EVENT_BUFFER: usize = 1024;
const LIFECYCLE_BUFFER: usize = 64;
#[derive(Debug, Clone, Copy)]
pub struct HeartbeatOptions {
pub interval: Duration,
pub timeout: Duration,
}
impl Default for HeartbeatOptions {
fn default() -> Self {
Self {
interval: Duration::from_secs(15),
timeout: Duration::from_secs(10),
}
}
}
#[derive(Debug, Clone)]
pub struct RealtimeOptions {
pub api_key: String,
pub url: String,
pub reconnect: Option<ReconnectOptions>,
pub heartbeat: Option<HeartbeatOptions>,
pub token_provider: Option<TokenProvider>,
pub auth_in_query: bool,
#[cfg(feature = "compression")]
#[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
pub compression: bool,
}
impl RealtimeOptions {
pub fn new(api_key: impl Into<String>) -> Self {
Self {
api_key: api_key.into(),
url: DEFAULT_WS_URL.to_string(),
reconnect: Some(ReconnectOptions::default()),
heartbeat: Some(HeartbeatOptions::default()),
token_provider: None,
auth_in_query: false,
#[cfg(feature = "compression")]
compression: false,
}
}
#[must_use]
pub fn url(mut self, url: impl Into<String>) -> Self {
self.url = url.into();
self
}
#[must_use]
pub fn reconnect(mut self, options: ReconnectOptions) -> Self {
self.reconnect = Some(options);
self
}
#[must_use]
pub fn disable_reconnect(mut self) -> Self {
self.reconnect = None;
self
}
#[must_use]
pub fn heartbeat(mut self, options: HeartbeatOptions) -> Self {
self.heartbeat = Some(options);
self
}
#[must_use]
pub fn disable_heartbeat(mut self) -> Self {
self.heartbeat = None;
self
}
#[must_use]
pub fn token(mut self, token: impl Into<String>) -> Self {
self.token_provider = Some(TokenProvider::from_static(token));
self
}
#[must_use]
pub fn token_provider(mut self, provider: TokenProvider) -> Self {
self.token_provider = Some(provider);
self
}
#[must_use]
pub fn auth_in_query(mut self, enabled: bool) -> Self {
self.auth_in_query = enabled;
self
}
#[cfg(feature = "compression")]
#[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
#[must_use]
pub fn compression(mut self, enabled: bool) -> Self {
self.compression = enabled;
self
}
fn decode_binary(&self, bytes: &[u8]) -> Result<String> {
#[cfg(feature = "compression")]
{
if self.compression {
return inflate(bytes);
}
}
std::str::from_utf8(bytes)
.map(str::to_owned)
.map_err(RadionError::transport)
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum LifecycleEvent {
Open,
Close {
code: u16,
reason: String,
},
Reconnect {
attempt: u32,
delay: Duration,
},
Warning {
code: String,
id: Option<String>,
message: String,
},
Error(RadionError),
}
enum Command {
Subscribe(Subscription),
Unsubscribe(String),
Close { code: u16, reason: String },
}
#[derive(Debug)]
pub struct RealtimeClient {
options: RealtimeOptions,
cmd_tx: mpsc::UnboundedSender<Command>,
cmd_rx: Mutex<Option<mpsc::UnboundedReceiver<Command>>>,
events_tx: broadcast::Sender<ChannelEvent>,
lifecycle_tx: broadcast::Sender<LifecycleEvent>,
connected: Arc<AtomicBool>,
task: Mutex<Option<JoinHandle<()>>>,
}
impl RealtimeClient {
pub fn new(options: RealtimeOptions) -> Self {
let (cmd_tx, cmd_rx) = mpsc::unbounded_channel();
let (events_tx, _) = broadcast::channel(EVENT_BUFFER);
let (lifecycle_tx, _) = broadcast::channel(LIFECYCLE_BUFFER);
Self {
options,
cmd_tx,
cmd_rx: Mutex::new(Some(cmd_rx)),
events_tx,
lifecycle_tx,
connected: Arc::new(AtomicBool::new(false)),
task: Mutex::new(None),
}
}
pub fn connected(&self) -> bool {
self.connected.load(Ordering::SeqCst)
}
pub async fn connect(&self) -> Result<()> {
if self.connected() {
return Ok(());
}
let Some(cmd_rx) = self.cmd_rx.lock().expect("cmd_rx mutex poisoned").take() else {
return Ok(());
};
let (ready_tx, ready_rx) = oneshot::channel();
let task = tokio::spawn(run(
self.options.clone(),
cmd_rx,
self.events_tx.clone(),
self.lifecycle_tx.clone(),
Arc::clone(&self.connected),
ready_tx,
));
*self.task.lock().expect("task mutex poisoned") = Some(task);
match ready_rx.await {
Ok(result) => result,
Err(_) => Err(RadionError::connection(
"connection task ended before connecting",
)),
}
}
pub async fn subscribe(&self, subscription: Subscription) -> Result<ChannelEventStream> {
subscription.validate()?;
let id = subscription.id.clone();
let rx = self.events_tx.subscribe();
self.cmd_tx
.send(Command::Subscribe(subscription))
.map_err(|_| RadionError::connection("client has been closed"))?;
Ok(ChannelEventStream {
inner: BroadcastStream::new(rx),
filter_id: Some(id),
})
}
pub async fn unsubscribe(&self, id: impl Into<String>) -> Result<()> {
self.cmd_tx
.send(Command::Unsubscribe(id.into()))
.map_err(|_| RadionError::connection("client has been closed"))
}
pub fn events(&self) -> ChannelEventStream {
ChannelEventStream {
inner: BroadcastStream::new(self.events_tx.subscribe()),
filter_id: None,
}
}
pub fn lifecycle(&self) -> LifecycleStream {
LifecycleStream {
inner: BroadcastStream::new(self.lifecycle_tx.subscribe()),
}
}
pub async fn close(&self, code: u16, reason: impl Into<String>) {
let _ = self.cmd_tx.send(Command::Close {
code,
reason: reason.into(),
});
let handle = self.task.lock().expect("task mutex poisoned").take();
if let Some(handle) = handle {
let _ = handle.await;
}
}
}
pub struct ChannelEventStream {
inner: BroadcastStream<ChannelEvent>,
filter_id: Option<String>,
}
impl Stream for ChannelEventStream {
type Item = ChannelEvent;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
match self.inner.poll_next_unpin(cx) {
Poll::Ready(Some(Ok(event))) => {
if self.filter_id.as_ref().is_none_or(|id| *id == event.id) {
return Poll::Ready(Some(event));
}
}
Poll::Ready(Some(Err(BroadcastStreamRecvError::Lagged(_)))) => {}
Poll::Ready(None) => return Poll::Ready(None),
Poll::Pending => return Poll::Pending,
}
}
}
}
pub struct LifecycleStream {
inner: BroadcastStream<LifecycleEvent>,
}
impl Stream for LifecycleStream {
type Item = LifecycleEvent;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
match self.inner.poll_next_unpin(cx) {
Poll::Ready(Some(Ok(event))) => return Poll::Ready(Some(event)),
Poll::Ready(Some(Err(BroadcastStreamRecvError::Lagged(_)))) => {}
Poll::Ready(None) => return Poll::Ready(None),
Poll::Pending => return Poll::Pending,
}
}
}
}
enum SessionOutcome {
Shutdown { code: u16, reason: String },
Disconnected { code: u16, reason: String },
}
async fn run(
options: RealtimeOptions,
mut cmd_rx: mpsc::UnboundedReceiver<Command>,
events_tx: broadcast::Sender<ChannelEvent>,
lifecycle_tx: broadcast::Sender<LifecycleEvent>,
connected: Arc<AtomicBool>,
ready_tx: oneshot::Sender<Result<()>>,
) {
let mut ready = Some(ready_tx);
let mut policy = ReconnectPolicy::new(options.reconnect.unwrap_or_default());
let mut subscriptions = SubscriptionManager::default();
loop {
match connect_ws(&options).await {
Ok(ws) => {
connected.store(true, Ordering::SeqCst);
policy.reset();
if let Some(tx) = ready.take() {
let _ = tx.send(Ok(()));
}
let _ = lifecycle_tx.send(LifecycleEvent::Open);
#[cfg(feature = "tracing")]
tracing::debug!(url = %options.url, "radion realtime connected");
let outcome = session(
ws,
&options,
&mut cmd_rx,
&events_tx,
&lifecycle_tx,
&mut subscriptions,
)
.await;
connected.store(false, Ordering::SeqCst);
match outcome {
SessionOutcome::Shutdown { code, reason } => {
let _ = lifecycle_tx.send(LifecycleEvent::Close { code, reason });
return;
}
SessionOutcome::Disconnected { code, reason } => {
let _ = lifecycle_tx.send(LifecycleEvent::Close { code, reason });
if options.reconnect.is_none() {
return;
}
}
}
}
Err(error) => {
if let Some(tx) = ready.take() {
let _ = tx.send(Err(error));
return;
}
let _ = lifecycle_tx.send(LifecycleEvent::Error(error));
if options.reconnect.is_none() {
return;
}
}
}
let delay = policy.next_delay();
let _ = lifecycle_tx.send(LifecycleEvent::Reconnect {
attempt: policy.attempts(),
delay,
});
#[cfg(feature = "tracing")]
tracing::debug!(
?delay,
attempt = policy.attempts(),
"radion realtime reconnecting"
);
tokio::select! {
() = tokio::time::sleep(delay) => {}
cmd = cmd_rx.recv() => match cmd {
Some(Command::Subscribe(subscription)) => {
subscriptions.add(subscription);
}
Some(Command::Unsubscribe(id)) => {
subscriptions.remove(&id);
}
Some(Command::Close { .. }) | None => return,
},
}
}
}
async fn connect_ws(
options: &RealtimeOptions,
) -> Result<
impl Stream<Item = std::result::Result<Message, WsError>> + Sink<Message, Error = WsError> + Unpin,
> {
let token = match &options.token_provider {
Some(provider) => Some(provider.fetch().await?),
None => None,
};
let base_url = connect_url(options);
let (ws, _response) = if options.auth_in_query {
let url = build_auth_query_url(&base_url, &options.api_key, token.as_deref());
let request = url.into_client_request().map_err(RadionError::transport)?;
tokio_tungstenite::connect_async(request)
.await
.map_err(RadionError::transport)?
} else {
let mut request = base_url
.as_str()
.into_client_request()
.map_err(RadionError::transport)?;
let api_key = HeaderValue::from_str(&options.api_key).map_err(RadionError::transport)?;
request
.headers_mut()
.insert(HeaderName::from_static("x-api-key"), api_key);
if let Some(token) = &token {
let bearer = HeaderValue::from_str(&format!("Bearer {token}"))
.map_err(RadionError::transport)?;
request
.headers_mut()
.insert(HeaderName::from_static("authorization"), bearer);
}
tokio_tungstenite::connect_async(request)
.await
.map_err(RadionError::transport)?
};
Ok(ws)
}
fn connect_url(options: &RealtimeOptions) -> String {
#[cfg(feature = "compression")]
{
if options.compression {
return with_compress_query(&options.url);
}
}
options.url.clone()
}
async fn session<S>(
mut ws: S,
options: &RealtimeOptions,
cmd_rx: &mut mpsc::UnboundedReceiver<Command>,
events_tx: &broadcast::Sender<ChannelEvent>,
lifecycle_tx: &broadcast::Sender<LifecycleEvent>,
subscriptions: &mut SubscriptionManager,
) -> SessionOutcome
where
S: Stream<Item = std::result::Result<Message, WsError>>
+ Sink<Message, Error = WsError>
+ Unpin,
{
let replay: Vec<_> = subscriptions
.desired()
.map(OutboundFrame::subscribe)
.collect();
for frame in replay {
send(&mut ws, frame).await;
}
let mut ping = options
.heartbeat
.map(|hb| tokio::time::interval(hb.interval));
let mut stale_deadline: Option<tokio::time::Instant> = None;
loop {
let stale = async {
match stale_deadline {
Some(deadline) => tokio::time::sleep_until(deadline).await,
None => std::future::pending().await,
}
};
tokio::select! {
message = ws.next() => match message {
Some(Ok(message)) => {
stale_deadline = None;
if let Some(outcome) = handle_message(&message, options, events_tx, lifecycle_tx) {
return outcome;
}
}
Some(Err(error)) => {
let _ = lifecycle_tx.send(LifecycleEvent::Error(RadionError::transport(error)));
return SessionOutcome::Disconnected { code: 1006, reason: String::new() };
}
None => return SessionOutcome::Disconnected { code: 1006, reason: String::new() },
},
command = cmd_rx.recv() => match command {
Some(Command::Subscribe(subscription)) => {
if subscriptions.add(subscription.clone()) {
send(&mut ws, OutboundFrame::subscribe(&subscription)).await;
}
}
Some(Command::Unsubscribe(id)) => {
if subscriptions.remove(&id) {
send(&mut ws, OutboundFrame::Unsubscribe { id }).await;
}
}
Some(Command::Close { code, reason }) => {
let _ = ws.close().await;
return SessionOutcome::Shutdown { code, reason };
}
None => {
let _ = ws.close().await;
return SessionOutcome::Shutdown { code: 1000, reason: String::from("client dropped") };
}
},
() = next_ping(&mut ping) => {
send(&mut ws, OutboundFrame::Ping).await;
if stale_deadline.is_none() {
if let Some(hb) = options.heartbeat {
stale_deadline = Some(tokio::time::Instant::now() + hb.timeout);
}
}
}
() = stale => {
let _ = lifecycle_tx.send(LifecycleEvent::Error(RadionError::connection("stale connection")));
return SessionOutcome::Disconnected { code: 1006, reason: String::from("stale connection") };
}
}
}
}
async fn next_ping(ping: &mut Option<tokio::time::Interval>) {
match ping {
Some(interval) => {
interval.tick().await;
}
None => std::future::pending().await,
}
}
fn handle_message(
message: &Message,
options: &RealtimeOptions,
events_tx: &broadcast::Sender<ChannelEvent>,
lifecycle_tx: &broadcast::Sender<LifecycleEvent>,
) -> Option<SessionOutcome> {
match message {
Message::Text(text) => {
route_text(text, events_tx, lifecycle_tx);
None
}
Message::Binary(bytes) => {
match options.decode_binary(bytes) {
Ok(text) => route_text(&text, events_tx, lifecycle_tx),
Err(error) => {
let _ = lifecycle_tx.send(LifecycleEvent::Error(error));
}
}
None
}
Message::Close(frame) => {
let (code, reason) = frame
.as_ref()
.map(|frame| (u16::from(frame.code), frame.reason.to_string()))
.unwrap_or((1005, String::new()));
Some(SessionOutcome::Disconnected { code, reason })
}
Message::Ping(_) | Message::Pong(_) | Message::Frame(_) => None,
}
}
fn route_text(
text: &str,
events_tx: &broadcast::Sender<ChannelEvent>,
lifecycle_tx: &broadcast::Sender<LifecycleEvent>,
) {
let Some(frame) = parse_inbound_frame(text) else {
return;
};
match frame {
frame @ InboundFrame::Event { .. } => {
if let Some(event) = frame.into_channel_event() {
let _ = events_tx.send(event);
}
}
InboundFrame::Warning { code, id, message } => {
let _ = lifecycle_tx.send(LifecycleEvent::Warning { code, id, message });
}
InboundFrame::Error {
message,
code,
id,
channel,
..
} => {
let _ = lifecycle_tx.send(LifecycleEvent::Error(RadionError::Server {
message,
code,
channel,
id,
}));
}
InboundFrame::Pong
| InboundFrame::Subscribed { .. }
| InboundFrame::Unsubscribed { .. } => {}
}
}
async fn send<S>(ws: &mut S, frame: OutboundFrame)
where
S: Sink<Message, Error = WsError> + Unpin,
{
if let Ok(text) = serde_json::to_string(&frame) {
let _ = ws.send(Message::text(text)).await;
}
}
#[cfg(test)]
mod auth_wiring_tests {
use super::*;
#[test]
fn defaults_have_no_token_and_header_mode() {
let options = RealtimeOptions::new("k");
assert!(options.token_provider.is_none());
assert!(!options.auth_in_query);
}
#[tokio::test]
async fn static_token_builder_sets_provider() {
let options = RealtimeOptions::new("k").token("jwt");
let provider = options.token_provider.expect("provider set");
assert_eq!(provider.fetch().await.unwrap(), "jwt");
}
#[test]
fn auth_in_query_builder_flips_flag() {
assert!(RealtimeOptions::new("k").auth_in_query(true).auth_in_query);
}
#[test]
fn accepts_async_provider() {
let _ = RealtimeOptions::new("k")
.token_provider(TokenProvider::new(|| async { Ok("x".into()) }));
}
}
#[cfg(all(test, feature = "compression"))]
mod compression_wiring_tests {
use std::io::Write;
use flate2::Compression;
use flate2::write::ZlibEncoder;
use super::*;
const PONG: &str = r#"{"type":"pong"}"#;
const EVENT: &str = r#"{"type":"event","id":"t","channel":"trading","confirmed":true,"seq":1,"sent_at_ms":1721818200123,"data":{"type":"order_cancelled"}}"#;
fn deflate(text: &str) -> Vec<u8> {
let mut encoder = ZlibEncoder::new(Vec::new(), Compression::default());
encoder.write_all(text.as_bytes()).expect("writes");
encoder.finish().expect("finishes")
}
#[test]
fn compression_is_off_by_default() {
assert!(!RealtimeOptions::new("k").compression);
}
#[test]
fn builder_flips_the_flag() {
assert!(RealtimeOptions::new("k").compression(true).compression);
}
#[test]
fn connect_url_is_untouched_when_compression_is_off() {
let options = RealtimeOptions::new("k").url("wss://example.test/ws");
assert_eq!(connect_url(&options), "wss://example.test/ws");
}
#[test]
fn connect_url_asks_for_zlib_when_compression_is_on() {
let options = RealtimeOptions::new("k")
.url("wss://example.test/ws")
.compression(true);
assert_eq!(connect_url(&options), "wss://example.test/ws?compress=zlib");
}
#[test]
fn connect_url_keeps_an_existing_query() {
let options = RealtimeOptions::new("k")
.url("wss://example.test/ws?v=1")
.compression(true);
assert_eq!(
connect_url(&options),
"wss://example.test/ws?v=1&compress=zlib"
);
}
#[test]
fn binary_frames_inflate_when_compression_is_on() {
let options = RealtimeOptions::new("k").compression(true);
assert_eq!(options.decode_binary(&deflate(PONG)).unwrap(), PONG);
}
#[test]
fn binary_frames_stay_plain_when_compression_is_off() {
let options = RealtimeOptions::new("k");
assert_eq!(options.decode_binary(PONG.as_bytes()).unwrap(), PONG);
}
#[test]
fn inflate_failure_surfaces_on_the_lifecycle_stream() {
let options = RealtimeOptions::new("k").compression(true);
let (events_tx, _events_rx) = broadcast::channel(EVENT_BUFFER);
let (lifecycle_tx, mut lifecycle_rx) = broadcast::channel(LIFECYCLE_BUFFER);
let outcome = handle_message(
&Message::binary(b"not zlib at all".to_vec()),
&options,
&events_tx,
&lifecycle_tx,
);
assert!(outcome.is_none());
match lifecycle_rx.try_recv().expect("lifecycle event") {
LifecycleEvent::Error(RadionError::Decompression(_)) => {}
other => panic!("expected a decompression error, got {other:?}"),
}
}
#[test]
fn text_and_compressed_binary_both_deliver_events() {
let options = RealtimeOptions::new("k").compression(true);
let (events_tx, mut events_rx) = broadcast::channel(EVENT_BUFFER);
let (lifecycle_tx, _lifecycle_rx) = broadcast::channel(LIFECYCLE_BUFFER);
handle_message(&Message::text(EVENT), &options, &events_tx, &lifecycle_tx);
handle_message(
&Message::binary(deflate(EVENT)),
&options,
&events_tx,
&lifecycle_tx,
);
assert_eq!(events_rx.try_recv().expect("text event").channel, "trading");
assert_eq!(
events_rx.try_recv().expect("binary event").channel,
"trading"
);
}
}