use alloc::borrow::Cow;
use alloc::string::String;
use alloc::vec::Vec;
use core::fmt;
use serde::de::{self, SeqAccess, Visitor};
use serde::ser::SerializeSeq;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use super::SubscriptionId;
use crate::event::Event;
use crate::filter::Filter;
use crate::util::impl_json_methods;
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum ClientMessage<'a> {
Event(Cow<'a, Event>),
Req {
subscription_id: Cow<'a, SubscriptionId>,
filters: Vec<Cow<'a, Filter>>,
},
Count {
subscription_id: Cow<'a, SubscriptionId>,
filter: Cow<'a, Filter>,
},
Close(Cow<'a, SubscriptionId>),
Auth(Cow<'a, Event>),
NegOpen {
subscription_id: Cow<'a, SubscriptionId>,
filter: Cow<'a, Filter>,
initial_message: Cow<'a, str>,
},
NegMsg {
subscription_id: Cow<'a, SubscriptionId>,
message: Cow<'a, str>,
},
NegClose {
subscription_id: Cow<'a, SubscriptionId>,
},
}
impl ClientMessage<'_> {
#[inline]
pub fn event(event: Event) -> Self {
Self::Event(Cow::Owned(event))
}
#[inline]
pub fn req<T>(subscription_id: SubscriptionId, filters: T) -> Self
where
T: Into<Vec<Filter>>,
{
Self::Req {
subscription_id: Cow::Owned(subscription_id),
filters: filters.into().into_iter().map(Cow::Owned).collect(),
}
}
#[inline]
pub fn count(subscription_id: SubscriptionId, filter: Filter) -> Self {
Self::Count {
subscription_id: Cow::Owned(subscription_id),
filter: Cow::Owned(filter),
}
}
#[inline]
pub fn close(subscription_id: SubscriptionId) -> Self {
Self::Close(Cow::Owned(subscription_id))
}
#[inline]
pub fn auth(event: Event) -> Self {
Self::Auth(Cow::Owned(event))
}
pub fn neg_open(
subscription_id: SubscriptionId,
filter: Filter,
initial_message: String,
) -> Self {
Self::NegOpen {
subscription_id: Cow::Owned(subscription_id),
filter: Cow::Owned(filter),
initial_message: Cow::Owned(initial_message),
}
}
#[inline]
pub fn is_event(&self) -> bool {
matches!(self, ClientMessage::Event(_))
}
#[inline]
pub fn is_req(&self) -> bool {
matches!(self, ClientMessage::Req { .. })
}
#[inline]
pub fn is_close(&self) -> bool {
matches!(self, ClientMessage::Close(_))
}
#[inline]
pub fn is_auth(&self) -> bool {
matches!(self, ClientMessage::Auth(_))
}
fn len(&self) -> usize {
match self {
Self::Event(..) | Self::Close(..) | Self::Auth(..) | Self::NegClose { .. } => 2,
Self::Count { .. } | Self::NegMsg { .. } => 3,
Self::Req { filters, .. } => 2 + filters.len(),
Self::NegOpen { .. } => 4,
}
}
}
impl Serialize for ClientMessage<'_> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut seq = serializer.serialize_seq(Some(self.len()))?;
match self {
Self::Event(event) => {
seq.serialize_element("EVENT")?;
seq.serialize_element(event)?;
}
Self::Req {
subscription_id,
filters,
} => {
seq.serialize_element("REQ")?;
seq.serialize_element(subscription_id)?;
for filter in filters {
seq.serialize_element(filter)?;
}
}
Self::Count {
subscription_id,
filter,
} => {
seq.serialize_element("COUNT")?;
seq.serialize_element(subscription_id)?;
seq.serialize_element(filter)?;
}
Self::Close(subscription_id) => {
seq.serialize_element("CLOSE")?;
seq.serialize_element(subscription_id)?;
}
Self::Auth(event) => {
seq.serialize_element("AUTH")?;
seq.serialize_element(event)?;
}
Self::NegOpen {
subscription_id,
filter,
initial_message,
} => {
seq.serialize_element("NEG-OPEN")?;
seq.serialize_element(subscription_id)?;
seq.serialize_element(filter)?;
seq.serialize_element(initial_message)?;
}
Self::NegMsg {
subscription_id,
message,
} => {
seq.serialize_element("NEG-MSG")?;
seq.serialize_element(subscription_id)?;
seq.serialize_element(message)?;
}
Self::NegClose { subscription_id } => {
seq.serialize_element("NEG-CLOSE")?;
seq.serialize_element(subscription_id)?;
}
}
seq.end()
}
}
impl<'de> Deserialize<'de> for ClientMessage<'_> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_seq(ClientMessageVisitor)
}
}
struct ClientMessageVisitor;
impl<'de> Visitor<'de> for ClientMessageVisitor {
type Value = ClientMessage<'static>;
fn expecting(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.write_str("a client message array")
}
fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
where
A: SeqAccess<'de>,
{
fn malformed<E>() -> E
where
E: de::Error,
{
E::custom("invalid message format")
}
macro_rules! next {
() => {
seq.next_element()?.ok_or_else(malformed)?
};
}
let message_type: String = next!();
let message: ClientMessage<'static> = match message_type.as_str() {
"EVENT" => ClientMessage::Event(Cow::Owned(next!())),
"REQ" => {
let subscription_id: SubscriptionId = next!();
let mut filters: Vec<Cow<'static, Filter>> = Vec::new();
while let Some(filter) = seq.next_element::<Filter>()? {
filters.push(Cow::Owned(filter));
}
if filters.is_empty() {
return Err(malformed());
}
return Ok(ClientMessage::Req {
subscription_id: Cow::Owned(subscription_id),
filters,
});
}
"COUNT" => ClientMessage::Count {
subscription_id: Cow::Owned(next!()),
filter: Cow::Owned(next!()),
},
"CLOSE" => ClientMessage::Close(Cow::Owned(next!())),
"AUTH" => ClientMessage::Auth(Cow::Owned(next!())),
"NEG-OPEN" => ClientMessage::NegOpen {
subscription_id: Cow::Owned(next!()),
filter: Cow::Owned(next!()),
initial_message: Cow::Owned(next!()),
},
"NEG-MSG" => ClientMessage::NegMsg {
subscription_id: Cow::Owned(next!()),
message: Cow::Owned(next!()),
},
"NEG-CLOSE" => ClientMessage::NegClose {
subscription_id: Cow::Owned(next!()),
},
_ => return Err(malformed()),
};
while seq.next_element::<de::IgnoredAny>()?.is_some() {}
Ok(message)
}
}
impl_json_methods!(ClientMessage<'_>);
#[cfg(test)]
mod tests {
use core::str::FromStr;
use super::*;
use crate::error::ErrorKind;
use crate::event::Kind;
use crate::key::PublicKey;
const EVENT_JSON: &str = r#"{"id":"70b10f70c1318967eddf12527799411b1a9780ad9c43858f5e5fcd45486a13a5","pubkey":"379e863e8357163b5bce5d2688dc4f1dcc2d505222fb8d74db600f30535dfdfe","created_at":1612809991,"kind":1,"tags":[],"content":"test","sig":"273a9cd5d11455590f4359500bccb7a89428262b96b3ea87a756b770964472f8c3e87f5d5e64d8d2e859a71462a3f477b554565c4f2f326cb01dd7620db71502"}"#;
#[test]
fn test_client_message_req() {
let pk =
PublicKey::from_str("379e863e8357163b5bce5d2688dc4f1dcc2d505222fb8d74db600f30535dfdfe")
.unwrap();
let client_req = ClientMessage::req(SubscriptionId::new("test"), Filter::new().pubkey(pk));
assert_eq!(
client_req.as_json(),
r##"["REQ","test",{"#p":["379e863e8357163b5bce5d2688dc4f1dcc2d505222fb8d74db600f30535dfdfe"]}]"##
);
}
#[test]
fn test_client_message_custom_kind() {
let client_req = ClientMessage::req(
SubscriptionId::new("test"),
Filter::new().kind(Kind::Custom(22)),
);
assert_eq!(client_req.as_json(), r##"["REQ","test",{"kinds":[22]}]"##);
}
#[test]
fn parse_trailing_elements() {
let cases: [(&str, ClientMessage); 3] = [
(
r#"["COUNT","sub",{"kinds":[1]},"extra"]"#,
ClientMessage::count(
SubscriptionId::new("sub"),
Filter::new().kind(Kind::TextNote),
),
),
(
r#"["CLOSE","sub",{"extra":true}]"#,
ClientMessage::close(SubscriptionId::new("sub")),
),
(
r#"["NEG-MSG","sub","deadbeef",1,2]"#,
ClientMessage::NegMsg {
subscription_id: Cow::Owned(SubscriptionId::new("sub")),
message: Cow::Borrowed("deadbeef"),
},
),
];
for (json, expected) in cases {
assert_eq!(ClientMessage::from_json(json).unwrap(), expected, "{json}");
}
}
#[test]
fn round_trip_every_variant() {
let event: Event = Event::from_json(EVENT_JSON).unwrap();
let sub = || SubscriptionId::new("sub");
let filter = || Filter::new().kind(Kind::TextNote);
let messages: [ClientMessage; 8] = [
ClientMessage::event(event.clone()),
ClientMessage::req(sub(), vec![filter(), Filter::new().author(event.pubkey)]),
ClientMessage::count(sub(), filter()),
ClientMessage::close(sub()),
ClientMessage::auth(event),
ClientMessage::neg_open(sub(), filter(), String::from("deadbeef")),
ClientMessage::NegMsg {
subscription_id: Cow::Owned(sub()),
message: Cow::Borrowed("deadbeef"),
},
ClientMessage::NegClose {
subscription_id: Cow::Owned(sub()),
},
];
for message in messages {
let json: String = message.as_json();
assert_eq!(ClientMessage::from_json(&json).unwrap(), message, "{json}");
}
}
#[test]
fn parse_rejects_unknown_type_and_non_array() {
for json in [
r#"["NOT-A-REAL-TYPE","x"]"#,
r#"{"type":"EVENT"}"#,
r#""EVENT""#,
r#"[]"#,
r#"["REQ","sub"]"#,
r#"["NEG-OPEN","sub",{},16,"deadbeef"]"#,
] {
let err = ClientMessage::from_json(json).unwrap_err();
assert_eq!(err.kind(), ErrorKind::Malformed, "{json}");
}
}
#[test]
fn event_message_embeds_canonical_event() {
let event: Event = Event::from_json(EVENT_JSON).unwrap();
let message: ClientMessage = ClientMessage::event(event.clone());
let json: String = message.as_json();
assert!(
json.contains(&event.as_json()),
"event was not embedded verbatim: {json}"
);
}
}