use serde::{Deserialize, Serialize};
use super::channels::{
Channel, FilterKey, SubscribableChannel, filter_requirement, join_filter_keys,
};
use super::payloads::Payload;
use crate::error::{RadionError, Result};
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct ChannelFilters {
#[serde(skip_serializing_if = "Option::is_none")]
pub wallets: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub market_ids: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub token_ids: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub min_usd: Option<f64>,
}
impl ChannelFilters {
fn has(&self, key: FilterKey) -> bool {
match key {
FilterKey::Wallets => self.wallets.as_ref().is_some_and(|v| !v.is_empty()),
FilterKey::MarketIds => self.market_ids.as_ref().is_some_and(|v| !v.is_empty()),
FilterKey::TokenIds => self.token_ids.as_ref().is_some_and(|v| !v.is_empty()),
FilterKey::MinUsd => self.min_usd.is_some(),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Subscription {
pub id: String,
pub channel: SubscribableChannel,
pub filters: Option<ChannelFilters>,
}
impl Subscription {
pub fn new(id: impl Into<String>, channel: impl Into<SubscribableChannel>) -> Self {
Self {
id: id.into(),
channel: channel.into(),
filters: None,
}
}
#[must_use]
pub fn with_filters(mut self, filters: ChannelFilters) -> Self {
self.filters = Some(filters);
self
}
pub fn validate(&self) -> Result<()> {
let confirmed = self.channel.confirmed();
let Some(requirement) = filter_requirement(confirmed) else {
return Ok(());
};
let satisfied = requirement
.required_any_of
.iter()
.any(|key| self.filters.as_ref().is_some_and(|f| f.has(*key)));
if satisfied {
return Ok(());
}
Err(RadionError::connection(format!(
"channel \"{}\" requires a {} filter",
self.channel,
join_filter_keys(requirement.required_any_of),
)))
}
}
#[derive(Debug, Clone, Serialize)]
#[serde(tag = "action", rename_all = "snake_case")]
pub(crate) enum OutboundFrame {
Subscribe {
id: String,
channel: SubscribableChannel,
#[serde(skip_serializing_if = "Option::is_none")]
filters: Option<ChannelFilters>,
},
Unsubscribe {
id: String,
},
Ping,
}
impl OutboundFrame {
pub(crate) fn subscribe(subscription: &Subscription) -> Self {
Self::Subscribe {
id: subscription.id.clone(),
channel: subscription.channel,
filters: subscription.filters.clone(),
}
}
}
#[derive(Debug, Clone)]
pub struct ChannelEvent {
pub id: String,
pub channel: String,
pub data: Payload,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub(crate) enum InboundFrame {
Event {
id: String,
channel: String,
data: serde_json::Value,
},
Subscribed {
#[allow(dead_code)]
id: String,
#[allow(dead_code)]
channel: Option<String>,
},
Unsubscribed {
#[allow(dead_code)]
id: String,
#[allow(dead_code)]
channel: Option<String>,
},
Pong,
Error {
message: String,
code: Option<String>,
id: Option<String>,
channel: Option<String>,
#[allow(dead_code)]
skipped: Option<u64>,
},
}
pub(crate) fn parse_inbound_frame(raw: &str) -> Option<InboundFrame> {
serde_json::from_str(raw).ok()
}
impl InboundFrame {
pub(crate) fn into_channel_event(self) -> Option<ChannelEvent> {
let Self::Event { id, channel, data } = self else {
return None;
};
let confirmed = channel
.strip_prefix("mempool.")
.unwrap_or(&channel)
.parse::<Channel>()
.ok();
let payload = match confirmed {
Some(channel) => Payload::from_channel(channel, data),
None => Payload::Other(data),
};
Some(ChannelEvent {
id,
channel,
data: payload,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::realtime::payloads::{Payload, TradingEventType};
#[test]
fn validates_required_filters() {
assert!(Subscription::new("w", Channel::Wallets).validate().is_err());
let ok = Subscription::new("w", Channel::Wallets).with_filters(ChannelFilters {
wallets: Some(vec!["0x1".into()]),
..Default::default()
});
assert!(ok.validate().is_ok());
let markets = Subscription::new("m", Channel::Markets).with_filters(ChannelFilters {
token_ids: Some(vec!["1".into()]),
..Default::default()
});
assert!(markets.validate().is_ok());
assert!(Subscription::new("t", Channel::Trading).validate().is_ok());
}
#[test]
fn serializes_outbound_frames() {
let ping = serde_json::to_string(&OutboundFrame::Ping).unwrap();
assert_eq!(ping, r#"{"action":"ping"}"#);
let unsub = serde_json::to_string(&OutboundFrame::Unsubscribe { id: "x".into() }).unwrap();
assert_eq!(unsub, r#"{"action":"unsubscribe","id":"x"}"#);
let sub = OutboundFrame::subscribe(&Subscription::new("trading", Channel::Trading));
let json: serde_json::Value =
serde_json::from_str(&serde_json::to_string(&sub).unwrap()).unwrap();
assert_eq!(json["action"], "subscribe");
assert_eq!(json["channel"], "trading");
assert!(json.get("filters").is_none());
}
#[test]
fn parses_and_types_event_frames() {
let raw = r#"{"type":"event","id":"t","channel":"trading","data":{"type":"order_filled_v2","side":1,"tokenId":"0xabc"}}"#;
let frame = parse_inbound_frame(raw).expect("valid frame");
let event = frame.into_channel_event().expect("event");
assert_eq!(event.id, "t");
match event.data {
Payload::Trading(trade) => {
assert_eq!(trade.kind, TradingEventType::OrderFilledV2);
assert_eq!(trade.side, Some(1));
assert_eq!(trade.token_id.as_deref(), Some("0xabc"));
}
other => panic!("expected trading payload, got {other:?}"),
}
}
#[test]
fn unknown_channel_falls_back_to_other() {
let raw = r#"{"type":"event","id":"m","channel":"mempool.unknownz","data":{"foo":1}}"#;
let event = parse_inbound_frame(raw)
.unwrap()
.into_channel_event()
.unwrap();
assert!(matches!(event.data, Payload::Other(_)));
}
#[test]
fn drops_malformed_frames() {
assert!(parse_inbound_frame("not json").is_none());
assert!(parse_inbound_frame(r#"{"type":"mystery"}"#).is_none());
}
#[test]
fn parses_error_frame() {
let raw = r#"{"type":"error","message":"boom","code":"bad","id":"x"}"#;
assert!(matches!(
parse_inbound_frame(raw),
Some(InboundFrame::Error { .. })
));
}
}