use serde::de::{Error as DeError, MapAccess, Visitor};
use serde::ser::{Error as SerError, SerializeMap};
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use crate::version::ProtocolVersion;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct OperationId(pub String);
impl From<String> for OperationId {
fn from(value: String) -> Self {
Self(value)
}
}
impl From<&str> for OperationId {
fn from(value: &str) -> Self {
Self(value.to_string())
}
}
impl std::fmt::Display for OperationId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
impl Serialize for OperationId {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(&self.0)
}
}
impl<'de> Deserialize<'de> for OperationId {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let value = String::deserialize(deserializer)?;
if value.is_empty() {
return Err(D::Error::custom("operation id must be a non-empty string"));
}
Ok(Self(value))
}
}
pub type Cursor = u64;
pub const FRAME_KINDS: &[&str] = &[
"handshake",
"handshake_ack",
"request",
"response",
"error",
"cancel",
"subscribe",
"subscribe_ack",
"unsubscribe",
"unsubscribe_ack",
"event",
];
pub const CLIENT_TO_SERVER_KINDS: &[&str] =
&["handshake", "request", "cancel", "subscribe", "unsubscribe"];
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct HandshakePayload {
pub version: ProtocolVersion,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct HandshakeAckPayload {
pub version: ProtocolVersion,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct RequestPayload {
pub id: OperationId,
pub ops: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub deadline_ms: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub namespace: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub actor_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub visible_namespaces: Option<Vec<String>>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ResponsePayload {
pub id: OperationId,
pub result: serde_json::Value,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ErrorPayload {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub id: Option<OperationId>,
pub code: crate::error::WireErrorCode,
pub message: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct CancelPayload {
pub id: OperationId,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct SubscribePayload {
pub id: OperationId,
pub topic: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub resume_cursor: Option<Cursor>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct SubscribeAckPayload {
pub id: OperationId,
pub topic: String,
pub start_cursor: Cursor,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct UnsubscribePayload {
pub id: OperationId,
pub topic: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct UnsubscribeAckPayload {
pub id: OperationId,
pub topic: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct EventPayload {
pub topic: String,
pub cursor: Cursor,
pub occurred_at: String,
pub payload: serde_json::Value,
}
#[derive(Debug, Clone, PartialEq)]
pub enum Frame {
Handshake {
version: ProtocolVersion,
},
HandshakeAck {
version: ProtocolVersion,
},
Request {
id: OperationId,
ops: String,
deadline_ms: Option<u64>,
namespace: Option<String>,
actor_id: Option<String>,
visible_namespaces: Option<Vec<String>>,
},
Response {
id: OperationId,
result: serde_json::Value,
},
Error {
id: Option<OperationId>,
code: crate::error::WireErrorCode,
message: String,
unrecognized_code: Option<String>,
},
Cancel {
id: OperationId,
},
Subscribe {
id: OperationId,
topic: String,
resume_cursor: Option<Cursor>,
},
SubscribeAck {
id: OperationId,
topic: String,
start_cursor: Cursor,
},
Unsubscribe {
id: OperationId,
topic: String,
},
UnsubscribeAck {
id: OperationId,
topic: String,
},
Event {
topic: String,
cursor: Cursor,
occurred_at: String,
payload: serde_json::Value,
},
}
impl Frame {
pub const fn kind(&self) -> &'static str {
match self {
Frame::Handshake { .. } => "handshake",
Frame::HandshakeAck { .. } => "handshake_ack",
Frame::Request { .. } => "request",
Frame::Response { .. } => "response",
Frame::Error { .. } => "error",
Frame::Cancel { .. } => "cancel",
Frame::Subscribe { .. } => "subscribe",
Frame::SubscribeAck { .. } => "subscribe_ack",
Frame::Unsubscribe { .. } => "unsubscribe",
Frame::UnsubscribeAck { .. } => "unsubscribe_ack",
Frame::Event { .. } => "event",
}
}
}
fn validate_frame_for_serialize(frame: &Frame) -> Result<(), String> {
use crate::error::TerminalScope;
fn check_id(kind: &str, id: &OperationId) -> Result<(), String> {
if id.0.is_empty() {
return Err(format!(
"frame kind {kind:?}: operation id must be a non-empty string"
));
}
Ok(())
}
fn check_version(kind: &str, version: ProtocolVersion) -> Result<(), String> {
if version.get() == 0 {
return Err(format!(
"frame kind {kind:?}: protocol version 0 does not exist"
));
}
Ok(())
}
match frame {
Frame::Handshake { version } => check_version("handshake", *version)?,
Frame::HandshakeAck { version } => check_version("handshake_ack", *version)?,
Frame::Request { id, .. } => check_id("request", id)?,
Frame::Response { id, .. } => check_id("response", id)?,
Frame::Cancel { id } => check_id("cancel", id)?,
Frame::Subscribe { id, .. } => check_id("subscribe", id)?,
Frame::SubscribeAck { id, .. } => check_id("subscribe_ack", id)?,
Frame::Unsubscribe { id, .. } => check_id("unsubscribe", id)?,
Frame::UnsubscribeAck { id, .. } => check_id("unsubscribe_ack", id)?,
Frame::Event { .. } => {}
Frame::Error {
id,
code,
unrecognized_code,
..
} => {
if let Some(raw_code) = unrecognized_code {
return Err(format!(
"fallback error frame (unrecognized code {raw_code:?}) is not re-encodable: \
re-encoding would emit \"internal\" and discard the newer wire code"
));
}
if let Some(id) = id {
check_id("error", id)?;
}
match (code.terminal_scope(), id) {
(TerminalScope::Connection, Some(id)) => {
return Err(format!(
"error frame violates the id/scope rule: connection-terminal code {code} \
must not carry an operation id, got {id}"
));
}
(TerminalScope::Request, None) => {
return Err(format!(
"error frame violates the id/scope rule: request-terminal code {code} \
must echo the operation id it terminates"
));
}
(TerminalScope::Connection, None) | (TerminalScope::Request, Some(_)) => {}
}
}
}
Ok(())
}
impl Serialize for Frame {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
if let Err(detail) = validate_frame_for_serialize(self) {
return Err(S::Error::custom(detail));
}
let size_hint = match self {
Frame::Handshake { .. } | Frame::HandshakeAck { .. } | Frame::Cancel { .. } => 2,
Frame::Subscribe { resume_cursor, .. } => {
if resume_cursor.is_some() {
4
} else {
3
}
}
Frame::Request {
deadline_ms,
namespace,
actor_id,
visible_namespaces,
..
} => {
3 + usize::from(deadline_ms.is_some())
+ usize::from(namespace.is_some())
+ usize::from(actor_id.is_some())
+ usize::from(visible_namespaces.is_some())
}
Frame::Error { id, .. } => 3 + usize::from(id.is_some()),
Frame::Response { .. } | Frame::Unsubscribe { .. } | Frame::UnsubscribeAck { .. } => 3,
Frame::SubscribeAck { .. } => 4,
Frame::Event { .. } => 5,
};
let mut map = serializer.serialize_map(Some(size_hint))?;
map.serialize_entry("kind", self.kind())?;
match self {
Frame::Handshake { version } => {
map.serialize_entry("version", version)?;
}
Frame::HandshakeAck { version } => {
map.serialize_entry("version", version)?;
}
Frame::Request {
id,
ops,
deadline_ms,
namespace,
actor_id,
visible_namespaces,
} => {
map.serialize_entry("id", id)?;
map.serialize_entry("ops", ops)?;
if let Some(deadline_ms) = deadline_ms {
map.serialize_entry("deadline_ms", deadline_ms)?;
}
if let Some(namespace) = namespace {
map.serialize_entry("namespace", namespace)?;
}
if let Some(actor_id) = actor_id {
map.serialize_entry("actor_id", actor_id)?;
}
if let Some(visible_namespaces) = visible_namespaces {
map.serialize_entry("visible_namespaces", visible_namespaces)?;
}
}
Frame::Response { id, result } => {
map.serialize_entry("id", id)?;
map.serialize_entry("result", result)?;
}
Frame::Error {
id, code, message, ..
} => {
if let Some(id) = id {
map.serialize_entry("id", id)?;
}
map.serialize_entry("code", code)?;
map.serialize_entry("message", message)?;
}
Frame::Cancel { id } => {
map.serialize_entry("id", id)?;
}
Frame::Subscribe {
id,
topic,
resume_cursor,
} => {
map.serialize_entry("id", id)?;
map.serialize_entry("topic", topic)?;
if let Some(resume_cursor) = resume_cursor {
map.serialize_entry("resume_cursor", resume_cursor)?;
}
}
Frame::SubscribeAck {
id,
topic,
start_cursor,
} => {
map.serialize_entry("id", id)?;
map.serialize_entry("topic", topic)?;
map.serialize_entry("start_cursor", start_cursor)?;
}
Frame::Unsubscribe { id, topic } => {
map.serialize_entry("id", id)?;
map.serialize_entry("topic", topic)?;
}
Frame::UnsubscribeAck { id, topic } => {
map.serialize_entry("id", id)?;
map.serialize_entry("topic", topic)?;
}
Frame::Event {
topic,
cursor,
occurred_at,
payload,
} => {
map.serialize_entry("topic", topic)?;
map.serialize_entry("cursor", cursor)?;
map.serialize_entry("occurred_at", occurred_at)?;
map.serialize_entry("payload", payload)?;
}
}
map.end()
}
}
impl<'de> Deserialize<'de> for Frame {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct FrameVisitor;
impl<'de> Visitor<'de> for FrameVisitor {
type Value = Frame;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a JSON object with a string \"kind\" discriminant field")
}
fn visit_map<A: MapAccess<'de>>(self, map: A) -> Result<Frame, A::Error> {
let mut object: serde_json::Map<String, serde_json::Value> =
Deserialize::deserialize(serde::de::value::MapAccessDeserializer::new(map))?;
let kind_value = object
.get("kind")
.ok_or_else(|| A::Error::custom("missing field `kind`"))?;
let kind = kind_value
.as_str()
.ok_or_else(|| A::Error::custom("field `kind` must be a string"))?;
let kind = kind.to_string();
object.remove("kind");
fn parse<T: serde::de::DeserializeOwned>(
object: &serde_json::Map<String, serde_json::Value>,
) -> Result<T, serde_json::Error> {
serde_json::from_value(serde_json::Value::Object(object.clone()))
}
match kind.as_str() {
"handshake" => {
let payload: HandshakePayload = parse(&object).map_err(A::Error::custom)?;
Ok(Frame::Handshake {
version: payload.version,
})
}
"handshake_ack" => {
let payload: HandshakeAckPayload =
parse(&object).map_err(A::Error::custom)?;
Ok(Frame::HandshakeAck {
version: payload.version,
})
}
"request" => {
let payload: RequestPayload = parse(&object).map_err(A::Error::custom)?;
Ok(Frame::Request {
id: payload.id,
ops: payload.ops,
deadline_ms: payload.deadline_ms,
namespace: payload.namespace,
actor_id: payload.actor_id,
visible_namespaces: payload.visible_namespaces,
})
}
"response" => {
let payload: ResponsePayload = parse(&object).map_err(A::Error::custom)?;
Ok(Frame::Response {
id: payload.id,
result: payload.result,
})
}
"error" => {
let raw_code = object
.get("code")
.and_then(|c| c.as_str())
.map(str::to_string);
let payload: ErrorPayload = parse(&object).map_err(A::Error::custom)?;
let code_in_closed_set = raw_code
.as_deref()
.is_some_and(|c| crate::error::WIRE_ERROR_CODES.contains(&c));
if code_in_closed_set {
match (payload.code.terminal_scope(), payload.id.as_ref()) {
(crate::error::TerminalScope::Connection, Some(id)) => {
return Err(A::Error::custom(format!(
"{}connection-terminal code {code} must not carry an operation id, got {id}",
crate::codec::INCONSISTENT_SCOPE_ERROR_PREFIX,
code = payload.code
)));
}
(crate::error::TerminalScope::Request, None) => {
return Err(A::Error::custom(format!(
"{}request-terminal code {code} must echo the operation id it terminates",
crate::codec::INCONSISTENT_SCOPE_ERROR_PREFIX,
code = payload.code
)));
}
(crate::error::TerminalScope::Connection, None)
| (crate::error::TerminalScope::Request, Some(_)) => {}
}
}
Ok(Frame::Error {
id: payload.id,
code: payload.code,
message: payload.message,
unrecognized_code: if code_in_closed_set { None } else { raw_code },
})
}
"cancel" => {
let payload: CancelPayload = parse(&object).map_err(A::Error::custom)?;
Ok(Frame::Cancel { id: payload.id })
}
"subscribe" => {
let payload: SubscribePayload = parse(&object).map_err(A::Error::custom)?;
Ok(Frame::Subscribe {
id: payload.id,
topic: payload.topic,
resume_cursor: payload.resume_cursor,
})
}
"subscribe_ack" => {
let payload: SubscribeAckPayload =
parse(&object).map_err(A::Error::custom)?;
Ok(Frame::SubscribeAck {
id: payload.id,
topic: payload.topic,
start_cursor: payload.start_cursor,
})
}
"unsubscribe" => {
let payload: UnsubscribePayload =
parse(&object).map_err(A::Error::custom)?;
Ok(Frame::Unsubscribe {
id: payload.id,
topic: payload.topic,
})
}
"unsubscribe_ack" => {
let payload: UnsubscribeAckPayload =
parse(&object).map_err(A::Error::custom)?;
Ok(Frame::UnsubscribeAck {
id: payload.id,
topic: payload.topic,
})
}
"event" => {
let payload: EventPayload = parse(&object).map_err(A::Error::custom)?;
Ok(Frame::Event {
topic: payload.topic,
cursor: payload.cursor,
occurred_at: payload.occurred_at,
payload: payload.payload,
})
}
other => Err(A::Error::custom(format!("unknown frame kind: {other:?}"))),
}
}
}
deserializer.deserialize_map(FrameVisitor)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_operation_id_is_rejected_at_deserialization() {
let err = serde_json::from_str::<OperationId>(r#""""#).unwrap_err();
assert!(
err.to_string().contains("non-empty"),
"unexpected error: {err}"
);
assert_eq!(OperationId::from(""), OperationId("".to_string()));
}
#[test]
fn non_empty_operation_id_round_trips() {
let id: OperationId = serde_json::from_str(r#""op-1""#).unwrap();
assert_eq!(id, OperationId::from("op-1"));
assert_eq!(serde_json::to_string(&id).unwrap(), r#""op-1""#);
}
#[test]
fn direct_serde_rejects_every_invalid_wire_shape() {
let frames = [
Frame::Handshake {
version: ProtocolVersion::new(0),
},
Frame::HandshakeAck {
version: ProtocolVersion::new(0),
},
Frame::Request {
id: OperationId::from(""),
ops: "stats()".to_string(),
deadline_ms: None,
namespace: None,
actor_id: None,
visible_namespaces: None,
},
Frame::Response {
id: OperationId::from(""),
result: serde_json::json!({}),
},
Frame::Error {
id: Some(OperationId::from("")),
code: crate::error::WireErrorCode::Internal,
message: "failure".to_string(),
unrecognized_code: None,
},
Frame::Cancel {
id: OperationId::from(""),
},
Frame::Subscribe {
id: OperationId::from(""),
topic: "a.b".to_string(),
resume_cursor: None,
},
Frame::SubscribeAck {
id: OperationId::from(""),
topic: "a.b".to_string(),
start_cursor: 0,
},
Frame::Unsubscribe {
id: OperationId::from(""),
topic: "a.b".to_string(),
},
Frame::UnsubscribeAck {
id: OperationId::from(""),
topic: "a.b".to_string(),
},
Frame::Error {
id: Some(OperationId::from("op-1")),
code: crate::error::WireErrorCode::FrameTooLarge,
message: "too big".to_string(),
unrecognized_code: None,
},
Frame::Error {
id: None,
code: crate::error::WireErrorCode::Cancelled,
message: "cancelled".to_string(),
unrecognized_code: None,
},
Frame::Error {
id: None,
code: crate::error::WireErrorCode::Internal,
message: "future".to_string(),
unrecognized_code: Some("future_code_xyz".to_string()),
},
];
for frame in frames {
assert!(
serde_json::to_vec(&frame).is_err(),
"invalid frame {:?} serialized successfully",
frame.kind()
);
}
}
#[test]
fn direct_serde_matches_encode_payload_for_a_valid_frame() {
let frame = Frame::Request {
id: OperationId::from("op-1"),
ops: "stats()".to_string(),
deadline_ms: Some(5000),
namespace: Some("default".to_string()),
actor_id: Some("actor".to_string()),
visible_namespaces: Some(vec!["default".to_string()]),
};
let direct = serde_json::to_vec(&frame).unwrap();
let encoded = crate::codec::encode_frame(&frame).unwrap();
assert_eq!(direct, &encoded[crate::codec::LENGTH_PREFIX_BYTES..]);
assert_eq!(serde_json::from_slice::<Frame>(&direct).unwrap(), frame);
}
#[test]
fn unknown_kind_through_direct_serde_is_rejected() {
let err = serde_json::from_str::<Frame>(r#"{"kind":"ping"}"#).unwrap_err();
assert!(err.to_string().contains("unknown frame kind"));
}
#[test]
fn client_to_server_kinds_are_all_frame_kinds() {
for kind in CLIENT_TO_SERVER_KINDS {
assert!(
FRAME_KINDS.contains(kind),
"client-to-server kind {kind:?} is missing from FRAME_KINDS"
);
}
}
#[test]
fn direct_serde_rejects_connection_terminal_error_carrying_an_id() {
let err = serde_json::from_str::<Frame>(
r#"{"kind":"error","id":"op-1","code":"frame_too_large","message":"too big"}"#,
)
.unwrap_err();
let message = err.to_string();
assert!(
message.contains("id/scope rule"),
"unexpected error: {message}"
);
assert!(message.contains("frame_too_large"), "error: {message}");
assert!(message.contains("op-1"), "error: {message}");
}
#[test]
fn direct_serde_rejects_request_terminal_error_without_an_id() {
let err = serde_json::from_str::<Frame>(
r#"{"kind":"error","code":"cancelled","message":"cancelled"}"#,
)
.unwrap_err();
let message = err.to_string();
assert!(
message.contains("id/scope rule"),
"unexpected error: {message}"
);
assert!(message.contains("cancelled"), "error: {message}");
}
#[test]
fn direct_serde_accepts_both_consistent_error_scopes() {
let connection_terminal: Frame = serde_json::from_str(
r#"{"kind":"error","code":"unsupported_version","message":"no common version"}"#,
)
.unwrap();
assert!(matches!(connection_terminal, Frame::Error { id: None, .. }));
let request_terminal: Frame = serde_json::from_str(
r#"{"kind":"error","id":"op-9","code":"deadline_exceeded","message":"too slow"}"#,
)
.unwrap();
match request_terminal {
Frame::Error {
id,
code,
unrecognized_code,
..
} => {
assert_eq!(id, Some(OperationId::from("op-9")));
assert_eq!(code, crate::error::WireErrorCode::DeadlineExceeded);
assert!(unrecognized_code.is_none());
}
other => panic!("expected an error frame, got {other:?}"),
}
}
#[test]
fn direct_serde_fills_unrecognized_code_for_an_unknown_wire_code() {
for json in [
r#"{"kind":"error","id":"op-7","code":"future_code_xyz","message":"from newer peer"}"#,
r#"{"kind":"error","code":"future_code_xyz","message":"from newer peer"}"#,
] {
let frame: Frame = serde_json::from_str(json).unwrap();
match frame {
Frame::Error {
code,
unrecognized_code,
..
} => {
assert_eq!(code, crate::error::WireErrorCode::Internal);
assert_eq!(unrecognized_code.as_deref(), Some("future_code_xyz"));
}
other => panic!("expected an error frame, got {other:?}"),
}
}
}
#[test]
fn direct_serde_leaves_unrecognized_code_none_for_a_literal_internal() {
let frame: Frame = serde_json::from_str(
r#"{"kind":"error","id":"op-1","code":"internal","message":"boom"}"#,
)
.unwrap();
match frame {
Frame::Error {
unrecognized_code, ..
} => assert!(unrecognized_code.is_none()),
other => panic!("expected an error frame, got {other:?}"),
}
}
}