use std::fmt;
use serde::de::{DeserializeOwned, Error as _};
use serde::ser::SerializeStruct;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use serde_json::{Map, Value};
use crate::auth::AccessToken;
use crate::error::{codes, ApiError};
use crate::ids::UserId;
use crate::kinds;
use crate::version::PROTOCOL_VERSION;
pub const MAX_MESSAGE_BYTES: usize = 1024 * 1024;
pub const AUTH_TIMEOUT_SECS: u64 = 5;
pub trait WsCall: Serialize + DeserializeOwned {
type Response: Serialize + DeserializeOwned;
const KIND: &'static str;
}
pub trait ServerPush: Serialize + DeserializeOwned {
const KIND: &'static str;
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize)]
#[non_exhaustive]
pub struct Ack {}
impl Ack {
pub const fn new() -> Self {
Self {}
}
}
impl<'de> Deserialize<'de> for Ack {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
serde::de::IgnoredAny::deserialize(deserializer)?;
Ok(Ack {})
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct CloseCode(pub u16);
impl CloseCode {
pub const NORMAL: CloseCode = CloseCode(1000);
pub const GOING_AWAY: CloseCode = CloseCode(1001);
pub const POLICY_VIOLATION: CloseCode = CloseCode(1008);
pub const MESSAGE_TOO_BIG: CloseCode = CloseCode(1009);
pub const INTERNAL_ERROR: CloseCode = CloseCode(1011);
pub const TRY_AGAIN_LATER: CloseCode = CloseCode(1013);
pub const UNAUTHORIZED: CloseCode = CloseCode(4001);
pub const BANNED: CloseCode = CloseCode(4003);
pub const REPLACED: CloseCode = CloseCode(4009);
pub const UNSUPPORTED_PROTOCOL: CloseCode = CloseCode(4010);
pub const ALL: &'static [CloseCode] = &[
Self::NORMAL,
Self::GOING_AWAY,
Self::POLICY_VIOLATION,
Self::MESSAGE_TOO_BIG,
Self::INTERNAL_ERROR,
Self::TRY_AGAIN_LATER,
Self::UNAUTHORIZED,
Self::BANNED,
Self::REPLACED,
Self::UNSUPPORTED_PROTOCOL,
];
pub const fn get(self) -> u16 {
self.0
}
pub const fn is_permanent(self) -> bool {
self.0 >= 4000 && self.0 < 4100
}
}
impl fmt::Display for CloseCode {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&self.0, f)
}
}
impl From<u16> for CloseCode {
fn from(code: u16) -> Self {
CloseCode(code)
}
}
impl From<CloseCode> for u16 {
fn from(code: CloseCode) -> Self {
code.0
}
}
#[derive(Serialize)]
struct TypedFrame<'a, D: Serialize> {
#[serde(rename = "type")]
kind: &'a str,
data: &'a D,
}
fn error_from(value: Option<Value>, fallback_code: &str) -> ApiError {
match value {
Some(value) => match ApiError::deserialize(&value) {
Ok(error) => error,
Err(_) => ApiError::new(fallback_code, "the error payload is not an API error").with_details(value),
},
None => ApiError::new(fallback_code, ""),
}
}
fn response_id(object: &Map<String, Value>) -> Option<u64> {
let id = object.get("id").and_then(Value::as_u64)?;
(object.contains_key("ok") || object.get("type").and_then(Value::as_str).is_none()).then_some(id)
}
#[derive(Clone, Debug, PartialEq, Serialize)]
#[non_exhaustive]
pub struct WsRequestFrame<T = Value> {
pub id: u64,
#[serde(rename = "type")]
pub kind: String,
pub data: T,
}
impl<T> WsRequestFrame<T> {
pub fn new(id: u64, kind: impl Into<String>, data: T) -> Self {
Self { id, kind: kind.into(), data }
}
}
impl<C: WsCall> WsRequestFrame<C> {
pub fn call(id: u64, call: C) -> Self {
Self { id, kind: C::KIND.to_string(), data: call }
}
}
impl WsRequestFrame<Value> {
pub fn data_as<T: DeserializeOwned>(&self) -> Result<T, serde_json::Error> {
T::deserialize(&self.data)
}
}
#[derive(Deserialize)]
struct RawRequest {
id: u64,
#[serde(rename = "type")]
kind: String,
#[serde(default)]
data: Value,
}
impl<'de, T: DeserializeOwned> Deserialize<'de> for WsRequestFrame<T> {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let raw = RawRequest::deserialize(deserializer)?;
let data = T::deserialize(raw.data).map_err(D::Error::custom)?;
Ok(Self { id: raw.id, kind: raw.kind, data })
}
}
#[derive(Clone, Debug, PartialEq)]
#[non_exhaustive]
pub struct WsResponseFrame<T = Value> {
pub id: u64,
pub result: Result<T, ApiError>,
}
impl<T> WsResponseFrame<T> {
pub fn ok(id: u64, data: T) -> Self {
Self { id, result: Ok(data) }
}
pub fn error(id: u64, error: ApiError) -> Self {
Self { id, result: Err(error) }
}
pub fn is_ok(&self) -> bool {
self.result.is_ok()
}
}
impl WsResponseFrame<Value> {
pub fn ok_serialize<T: Serialize + ?Sized>(id: u64, data: &T) -> Result<Self, serde_json::Error> {
Ok(Self::ok(id, serde_json::to_value(data)?))
}
pub fn decode<T: DeserializeOwned>(&self) -> Result<Result<T, ApiError>, serde_json::Error> {
match &self.result {
Ok(data) => T::deserialize(data).map(Ok),
Err(error) => Ok(Err(error.clone())),
}
}
}
impl<T: Serialize> Serialize for WsResponseFrame<T> {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let mut frame = serializer.serialize_struct("WsResponseFrame", 3)?;
frame.serialize_field("id", &self.id)?;
match &self.result {
Ok(data) => {
frame.serialize_field("ok", &true)?;
frame.serialize_field("data", data)?;
}
Err(error) => {
frame.serialize_field("ok", &false)?;
frame.serialize_field("error", error)?;
}
}
frame.end()
}
}
fn response_from_object<T: DeserializeOwned>(mut object: Map<String, Value>) -> Result<WsResponseFrame<T>, String> {
let id = response_id(&object).ok_or("not an answer: it needs a numeric `id` and an `ok` field or no `type`")?;
let ok = object.get("ok").and_then(Value::as_bool).unwrap_or(true);
let result = if ok {
let data = object.remove("data").unwrap_or(Value::Null);
Ok(T::deserialize(data).map_err(|e| e.to_string())?)
} else {
Err(error_from(object.remove("error"), codes::INTERNAL))
};
Ok(WsResponseFrame { id, result })
}
impl<'de, T: DeserializeOwned> Deserialize<'de> for WsResponseFrame<T> {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let object = Map::<String, Value>::deserialize(deserializer)?;
response_from_object(object).map_err(D::Error::custom)
}
}
#[derive(Clone, Debug, PartialEq)]
#[non_exhaustive]
pub struct WsPushFrame<T = Value> {
pub kind: String,
pub data: T,
}
impl<T> WsPushFrame<T> {
pub fn new(kind: impl Into<String>, data: T) -> Self {
Self { kind: kind.into(), data }
}
}
impl<T> WsPushFrame<T> {
pub fn checked(kind: impl Into<String>, data: T) -> Option<Self> {
let kind = kind.into();
(!kind.is_empty() && !kinds::is_reserved(&kind)).then_some(Self { kind, data })
}
}
impl<P: ServerPush> WsPushFrame<P> {
pub fn push(data: P) -> Self {
Self { kind: P::KIND.to_string(), data }
}
}
impl WsPushFrame<Value> {
pub fn data_as<T: DeserializeOwned>(&self) -> Result<T, serde_json::Error> {
T::deserialize(&self.data)
}
}
impl<T: Serialize> Serialize for WsPushFrame<T> {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
TypedFrame { kind: &self.kind, data: &self.data }.serialize(serializer)
}
}
fn push_from_object<T: DeserializeOwned>(mut object: Map<String, Value>) -> Result<WsPushFrame<T>, String> {
if response_id(&object).is_some() {
return Err("not a push: it is an answer".into());
}
let kind = match object.get("type").and_then(Value::as_str) {
Some(kinds::AUTH_OK | kinds::AUTH_FAILED) => return Err("not a push: it is an auth result".into()),
Some(kind) => kind.to_string(),
None => return Err("not a push: no `type`".into()),
};
let data = T::deserialize(object.remove("data").unwrap_or(Value::Null)).map_err(|e| e.to_string())?;
Ok(WsPushFrame { kind, data })
}
impl<'de, T: DeserializeOwned> Deserialize<'de> for WsPushFrame<T> {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let object = Map::<String, Value>::deserialize(deserializer)?;
push_from_object(object).map_err(D::Error::custom)
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct WsAuth {
pub token: AccessToken,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub protocol: Option<u32>,
}
impl WsAuth {
pub fn new(token: impl Into<AccessToken>) -> Self {
Self { token: token.into(), protocol: Some(PROTOCOL_VERSION) }
}
pub fn to_message(&self) -> String {
serde_json::to_string(&TypedFrame { kind: kinds::AUTH, data: self }).unwrap_or_default()
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct WsAuthOk {
pub user_id: UserId,
pub protocol: u32,
}
impl WsAuthOk {
pub fn new(user_id: UserId) -> Self {
Self { user_id, protocol: PROTOCOL_VERSION }
}
}
#[derive(Clone, Debug, PartialEq)]
#[non_exhaustive]
pub enum WsServerFrame {
Response(WsResponseFrame),
Push(WsPushFrame),
AuthOk(Option<WsAuthOk>),
AuthFailed(ApiError),
}
impl WsServerFrame {
pub fn parse(text: &str) -> Result<Self, serde_json::Error> {
serde_json::from_str(text)
}
pub fn to_json(&self) -> Result<String, serde_json::Error> {
serde_json::to_string(self)
}
}
impl Serialize for WsServerFrame {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
match self {
WsServerFrame::Response(frame) => frame.serialize(serializer),
WsServerFrame::Push(frame) => frame.serialize(serializer),
WsServerFrame::AuthOk(Some(ok)) => TypedFrame { kind: kinds::AUTH_OK, data: ok }.serialize(serializer),
WsServerFrame::AuthOk(None) => {
let mut frame = serializer.serialize_struct("WsServerFrame", 1)?;
frame.serialize_field("type", kinds::AUTH_OK)?;
frame.end()
}
WsServerFrame::AuthFailed(error) => {
let mut frame = serializer.serialize_struct("WsServerFrame", 2)?;
frame.serialize_field("type", kinds::AUTH_FAILED)?;
frame.serialize_field("error", error)?;
frame.end()
}
}
}
}
impl<'de> Deserialize<'de> for WsServerFrame {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let mut object = Map::<String, Value>::deserialize(deserializer)?;
if response_id(&object).is_some() {
return response_from_object(object).map(WsServerFrame::Response).map_err(D::Error::custom);
}
match object.get("type").and_then(Value::as_str) {
Some(kinds::AUTH_OK) => Ok(WsServerFrame::AuthOk(object.remove("data").and_then(|data| WsAuthOk::deserialize(data).ok()))),
Some(kinds::AUTH_FAILED) => Ok(WsServerFrame::AuthFailed(error_from(object.remove("error"), codes::UNAUTHORIZED))),
Some(_) => push_from_object(object).map(WsServerFrame::Push).map_err(D::Error::custom),
None => Err(D::Error::custom("not a protocol frame: neither an answer nor a `type`")),
}
}
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub enum WsClientFrame {
Request(WsRequestFrame),
Auth(WsAuth),
}
impl WsClientFrame {
pub fn parse(text: &str) -> Result<Self, FrameError> {
let value: Value = serde_json::from_str(text).map_err(|e| FrameError { id: None, message: e.to_string() })?;
let id = value.as_object().and_then(|object| object.get("id")).and_then(Value::as_u64);
WsClientFrame::deserialize(value).map_err(|e| FrameError { id, message: e.to_string() })
}
pub fn request_id(text: &str) -> Option<u64> {
let value: Value = serde_json::from_str(text).ok()?;
value.as_object()?.get("id")?.as_u64()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct FrameError {
pub id: Option<u64>,
pub message: String,
}
impl FrameError {
pub fn answer(&self) -> Option<WsResponseFrame> {
self.id.map(|id| WsResponseFrame::error(id, ApiError::new(codes::BAD_REQUEST, "the request is malformed")))
}
}
impl fmt::Display for FrameError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.id {
Some(id) => write!(f, "malformed request {id}: {}", self.message),
None => write!(f, "not a request: {}", self.message),
}
}
}
impl std::error::Error for FrameError {}
impl Serialize for WsClientFrame {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
match self {
WsClientFrame::Request(frame) => frame.serialize(serializer),
WsClientFrame::Auth(auth) => TypedFrame { kind: kinds::AUTH, data: auth }.serialize(serializer),
}
}
}
impl<'de> Deserialize<'de> for WsClientFrame {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let mut object = Map::<String, Value>::deserialize(deserializer)?;
if object.contains_key("id") {
return WsRequestFrame::deserialize(Value::Object(object)).map(WsClientFrame::Request).map_err(D::Error::custom);
}
match object.get("type").and_then(Value::as_str) {
Some(kinds::AUTH) => {
let data = object.remove("data").unwrap_or(Value::Null);
WsAuth::deserialize(data).map(WsClientFrame::Auth).map_err(D::Error::custom)
}
_ => Err(D::Error::custom("not a request: no `id`")),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn close_codes() {
assert!(CloseCode::UNAUTHORIZED.is_permanent() && CloseCode(4099).is_permanent());
assert!(!CloseCode::GOING_AWAY.is_permanent() && !CloseCode(4100).is_permanent() && !CloseCode(3999).is_permanent());
assert_eq!(u16::from(CloseCode::TRY_AGAIN_LATER), 1013);
assert_eq!(CloseCode::from(1000), CloseCode::NORMAL);
assert_eq!(CloseCode::REPLACED.to_string(), "4009");
assert_eq!(CloseCode::BANNED.get(), 4003);
}
#[test]
fn ack_accepts_anything() {
for json in ["{}", "null", r#"{"later":1}"#, "true"] {
assert_eq!(serde_json::from_str::<Ack>(json).ok(), Some(Ack::new()), "{json}");
}
assert_eq!(serde_json::to_string(&Ack::new()).ok().as_deref(), Some("{}"));
}
#[test]
fn error_payload_fallbacks() {
let frame = WsServerFrame::parse(r#"{"id":1,"ok":false,"error":"nope"}"#).ok();
let Some(WsServerFrame::Response(WsResponseFrame { result: Err(error), .. })) = frame else {
panic!("not an error answer: {frame:?}");
};
assert!(error.is(codes::INTERNAL));
assert_eq!(error.details, Some(Value::String("nope".into())));
let frame = WsServerFrame::parse(r#"{"type":"auth.failed"}"#).ok();
assert!(matches!(frame, Some(WsServerFrame::AuthFailed(e)) if e.is(codes::UNAUTHORIZED)));
}
}