use std::fmt;
use std::str::FromStr;
use serde::de::{self, SeqAccess, Visitor};
use serde::ser::{SerializeSeq, Serializer};
use serde::{Deserialize, Deserializer, Serialize};
use thiserror::Error;
use super::subscription_id::SubscriptionId;
use crate::event::{Event, EventId};
const TAG_EVENT: &str = "EVENT";
const TAG_OK: &str = "OK";
const TAG_EOSE: &str = "EOSE";
const TAG_CLOSED: &str = "CLOSED";
const TAG_NOTICE: &str = "NOTICE";
const TAG_AUTH: &str = "AUTH";
const TAG_COUNT: &str = "COUNT";
const TAG_NEG_MSG: &str = "NEG-MSG";
const TAG_NEG_ERR: &str = "NEG-ERR";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
#[non_exhaustive]
pub enum MachineReadablePrefixError {
#[error("unknown machine-readable prefix")]
Unknown,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
#[non_exhaustive]
pub enum MachineReadablePrefix {
Duplicate,
Pow,
Blocked,
RateLimited,
Invalid,
Error,
Restricted,
Mute,
AuthRequired,
PaymentRequired,
}
impl MachineReadablePrefix {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Duplicate => "duplicate",
Self::Pow => "pow",
Self::Blocked => "blocked",
Self::RateLimited => "rate-limited",
Self::Invalid => "invalid",
Self::Error => "error",
Self::Restricted => "restricted",
Self::Mute => "mute",
Self::AuthRequired => "auth-required",
Self::PaymentRequired => "payment-required",
}
}
#[must_use]
pub fn from_reason(reason: &str) -> Option<Self> {
let (prefix, _rest) = reason.split_once(':')?;
prefix.parse().ok()
}
}
impl fmt::Display for MachineReadablePrefix {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
impl FromStr for MachineReadablePrefix {
type Err = MachineReadablePrefixError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let value = match s {
"duplicate" => Self::Duplicate,
"pow" => Self::Pow,
"blocked" => Self::Blocked,
"rate-limited" => Self::RateLimited,
"invalid" => Self::Invalid,
"error" => Self::Error,
"restricted" => Self::Restricted,
"mute" => Self::Mute,
"auth-required" => Self::AuthRequired,
"payment-required" => Self::PaymentRequired,
_ => return Err(MachineReadablePrefixError::Unknown),
};
Ok(value)
}
}
#[derive(Debug, Clone, Error)]
#[non_exhaustive]
pub enum RelayMessageError {
#[error("relay message must not be empty")]
Empty,
#[error("unknown relay message tag `{0}`")]
UnknownTag(String),
#[error("malformed `{tag}` message: {reason}")]
Malformed {
tag: &'static str,
reason: String,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
#[allow(
clippy::large_enum_variant,
reason = "EVENT inherently carries a full Event while the control variants \
(OK/EOSE/CLOSED/NOTICE/AUTH) are small; boxing it would add \
allocation churn on the relay-send and pool-receive hot paths, \
where EVENT is by far the most common message and is moved \
straight into/out of the enum. rust-nostr makes the same \
trade-off via a Cow<Event> EVENT variant."
)]
pub enum RelayMessage {
Event {
subscription_id: SubscriptionId,
event: Event,
},
Ok {
event_id: EventId,
accepted: bool,
message: String,
},
EndOfStoredEvents(SubscriptionId),
Closed {
subscription_id: SubscriptionId,
message: String,
},
Notice(String),
Auth(String),
Count {
subscription_id: SubscriptionId,
count: u64,
},
NegMsg {
subscription_id: SubscriptionId,
message: String,
},
NegErr {
subscription_id: SubscriptionId,
message: String,
},
}
impl Serialize for RelayMessage {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
match self {
Self::Event {
subscription_id,
event,
} => {
let mut seq = serializer.serialize_seq(Some(3))?;
seq.serialize_element(TAG_EVENT)?;
seq.serialize_element(subscription_id)?;
seq.serialize_element(event)?;
seq.end()
}
Self::Ok {
event_id,
accepted,
message,
} => {
let mut seq = serializer.serialize_seq(Some(4))?;
seq.serialize_element(TAG_OK)?;
seq.serialize_element(event_id)?;
seq.serialize_element(accepted)?;
seq.serialize_element(message)?;
seq.end()
}
Self::EndOfStoredEvents(id) => {
let mut seq = serializer.serialize_seq(Some(2))?;
seq.serialize_element(TAG_EOSE)?;
seq.serialize_element(id)?;
seq.end()
}
Self::Closed {
subscription_id,
message,
} => {
let mut seq = serializer.serialize_seq(Some(3))?;
seq.serialize_element(TAG_CLOSED)?;
seq.serialize_element(subscription_id)?;
seq.serialize_element(message)?;
seq.end()
}
Self::Notice(message) => {
let mut seq = serializer.serialize_seq(Some(2))?;
seq.serialize_element(TAG_NOTICE)?;
seq.serialize_element(message)?;
seq.end()
}
Self::Auth(challenge) => {
let mut seq = serializer.serialize_seq(Some(2))?;
seq.serialize_element(TAG_AUTH)?;
seq.serialize_element(challenge)?;
seq.end()
}
Self::Count {
subscription_id,
count,
} => {
#[derive(Serialize)]
struct CountPayload {
count: u64,
}
let mut seq = serializer.serialize_seq(Some(3))?;
seq.serialize_element(TAG_COUNT)?;
seq.serialize_element(subscription_id)?;
seq.serialize_element(&CountPayload { count: *count })?;
seq.end()
}
Self::NegMsg {
subscription_id,
message,
} => {
let mut seq = serializer.serialize_seq(Some(3))?;
seq.serialize_element(TAG_NEG_MSG)?;
seq.serialize_element(subscription_id)?;
seq.serialize_element(message)?;
seq.end()
}
Self::NegErr {
subscription_id,
message,
} => {
let mut seq = serializer.serialize_seq(Some(3))?;
seq.serialize_element(TAG_NEG_ERR)?;
seq.serialize_element(subscription_id)?;
seq.serialize_element(message)?;
seq.end()
}
}
}
}
impl<'de> Deserialize<'de> for RelayMessage {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
struct RelayVisitor;
impl<'de> Visitor<'de> for RelayVisitor {
type Value = RelayMessage;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a Nostr relay message array")
}
fn visit_seq<A>(self, mut seq: A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
let tag: String = seq
.next_element()?
.ok_or_else(|| de::Error::custom(RelayMessageError::Empty))?;
match tag.as_str() {
TAG_EVENT => decode_event(&mut seq),
TAG_OK => decode_ok(&mut seq),
TAG_EOSE => decode_eose(&mut seq),
TAG_CLOSED => decode_closed(&mut seq),
TAG_NOTICE => decode_notice(&mut seq),
TAG_AUTH => decode_auth(&mut seq),
TAG_COUNT => decode_count(&mut seq),
TAG_NEG_MSG => decode_neg_msg(&mut seq),
TAG_NEG_ERR => decode_neg_err(&mut seq),
other => Err(de::Error::custom(RelayMessageError::UnknownTag(
other.to_owned(),
))),
}
}
}
deserializer.deserialize_seq(RelayVisitor)
}
}
fn malformed<E: de::Error>(tag: &'static str, reason: &str) -> E {
E::custom(RelayMessageError::Malformed {
tag,
reason: reason.to_owned(),
})
}
fn decode_event<'de, A>(seq: &mut A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
let subscription_id: SubscriptionId = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_EVENT, "missing subscription id"))?;
let event: Event = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_EVENT, "missing event"))?;
Ok(RelayMessage::Event {
subscription_id,
event,
})
}
fn decode_ok<'de, A>(seq: &mut A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
let event_id: EventId = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_OK, "missing event id"))?;
let accepted: bool = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_OK, "missing accepted flag"))?;
let message: String = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_OK, "missing message"))?;
Ok(RelayMessage::Ok {
event_id,
accepted,
message,
})
}
fn decode_eose<'de, A>(seq: &mut A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
let id: SubscriptionId = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_EOSE, "missing subscription id"))?;
Ok(RelayMessage::EndOfStoredEvents(id))
}
fn decode_closed<'de, A>(seq: &mut A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
let subscription_id: SubscriptionId = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_CLOSED, "missing subscription id"))?;
let message: String = seq.next_element()?.unwrap_or_default();
Ok(RelayMessage::Closed {
subscription_id,
message,
})
}
fn decode_notice<'de, A>(seq: &mut A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
let message: String = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_NOTICE, "missing message"))?;
Ok(RelayMessage::Notice(message))
}
fn decode_auth<'de, A>(seq: &mut A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
let challenge: String = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_AUTH, "missing challenge"))?;
Ok(RelayMessage::Auth(challenge))
}
fn decode_count<'de, A>(seq: &mut A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
#[derive(Deserialize)]
struct CountPayload {
count: u64,
}
let subscription_id: SubscriptionId = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_COUNT, "missing subscription id"))?;
let payload: CountPayload = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_COUNT, "missing count payload"))?;
Ok(RelayMessage::Count {
subscription_id,
count: payload.count,
})
}
fn decode_neg_msg<'de, A>(seq: &mut A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
let subscription_id: SubscriptionId = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_NEG_MSG, "missing subscription id"))?;
let message: String = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_NEG_MSG, "missing message"))?;
Ok(RelayMessage::NegMsg {
subscription_id,
message,
})
}
fn decode_neg_err<'de, A>(seq: &mut A) -> Result<RelayMessage, A::Error>
where
A: SeqAccess<'de>,
{
let subscription_id: SubscriptionId = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_NEG_ERR, "missing subscription id"))?;
let message: String = seq
.next_element()?
.ok_or_else(|| malformed::<A::Error>(TAG_NEG_ERR, "missing message"))?;
Ok(RelayMessage::NegErr {
subscription_id,
message,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Keys;
use crate::event::EventBuilder;
fn keys() -> Keys {
Keys::parse("0000000000000000000000000000000000000000000000000000000000000003").unwrap()
}
fn signed_event() -> Event {
EventBuilder::text_note("hello")
.sign_with_keys(&keys())
.unwrap()
}
fn sub() -> SubscriptionId {
SubscriptionId::new("sub-1").unwrap()
}
#[test]
fn event_round_trip() {
let msg = RelayMessage::Event {
subscription_id: sub(),
event: signed_event(),
};
let json = serde_json::to_string(&msg).unwrap();
assert!(json.starts_with("[\"EVENT\",\"sub-1\","));
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn ok_round_trip_with_message() {
let msg = RelayMessage::Ok {
event_id: signed_event().id,
accepted: false,
message: "blocked: spam".to_owned(),
};
let json = serde_json::to_string(&msg).unwrap();
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn ok_with_empty_message_round_trip() {
let msg = RelayMessage::Ok {
event_id: signed_event().id,
accepted: true,
message: String::new(),
};
let json = serde_json::to_string(&msg).unwrap();
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn eose_round_trip() {
let msg = RelayMessage::EndOfStoredEvents(sub());
let json = serde_json::to_string(&msg).unwrap();
assert_eq!(json, "[\"EOSE\",\"sub-1\"]");
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn closed_round_trip() {
let msg = RelayMessage::Closed {
subscription_id: sub(),
message: "auth-required: please authenticate".to_owned(),
};
let json = serde_json::to_string(&msg).unwrap();
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn notice_round_trip() {
let msg = RelayMessage::Notice("welcome".to_owned());
let json = serde_json::to_string(&msg).unwrap();
assert_eq!(json, "[\"NOTICE\",\"welcome\"]");
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn auth_challenge_round_trip() {
let msg = RelayMessage::Auth("challenge-string".to_owned());
let json = serde_json::to_string(&msg).unwrap();
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn count_round_trip() {
let msg = RelayMessage::Count {
subscription_id: sub(),
count: 42,
};
let json = serde_json::to_string(&msg).unwrap();
assert_eq!(json, "[\"COUNT\",\"sub-1\",{\"count\":42}]");
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn machine_readable_prefix_parses() {
assert_eq!(
MachineReadablePrefix::from_reason("blocked: spam"),
Some(MachineReadablePrefix::Blocked)
);
assert_eq!(
MachineReadablePrefix::from_reason("auth-required: please"),
Some(MachineReadablePrefix::AuthRequired)
);
assert_eq!(
MachineReadablePrefix::from_reason("mute: nobody listening"),
Some(MachineReadablePrefix::Mute)
);
assert!(MachineReadablePrefix::from_reason("no prefix").is_none());
assert!(MachineReadablePrefix::from_reason("unknown: thing").is_none());
}
#[test]
fn ok_round_trips_mute_prefix() {
let msg = RelayMessage::Ok {
event_id: signed_event().id,
accepted: false,
message: "mute: nobody was listening".to_owned(),
};
let json = serde_json::to_string(&msg).unwrap();
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn ok_rejects_missing_message_per_nip01() {
let json = format!(r#"["OK","{}",true]"#, signed_event().id.to_hex());
let err = serde_json::from_str::<RelayMessage>(&json).unwrap_err();
assert!(err.to_string().contains("missing message"));
}
#[test]
fn unknown_tag_rejected() {
let json = "[\"WAT\",\"x\"]";
let err = serde_json::from_str::<RelayMessage>(json).unwrap_err();
assert!(err.to_string().contains("unknown relay message tag"));
}
#[test]
fn neg_msg_round_trip() {
let msg = RelayMessage::NegMsg {
subscription_id: sub(),
message: "deadbeef".to_owned(),
};
let json = serde_json::to_string(&msg).unwrap();
assert_eq!(json, "[\"NEG-MSG\",\"sub-1\",\"deadbeef\"]");
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn neg_err_round_trip() {
let msg = RelayMessage::NegErr {
subscription_id: sub(),
message: "blocked: spam".to_owned(),
};
let json = serde_json::to_string(&msg).unwrap();
let parsed: RelayMessage = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, msg);
}
#[test]
fn neg_msg_missing_payload_rejected() {
let json = "[\"NEG-MSG\",\"sub-1\"]";
let err = serde_json::from_str::<RelayMessage>(json).unwrap_err();
assert!(err.to_string().contains("missing message"));
}
}