use fraiseql_core::runtime::protocol::{ClientMessage, ServerMessage};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum WsProtocol {
GraphqlTransportWs,
GraphqlWs,
}
impl WsProtocol {
#[must_use]
pub fn from_header(header: Option<&str>) -> Option<Self> {
let header = header?;
for token in header.split(',') {
match token.trim() {
"graphql-transport-ws" => return Some(Self::GraphqlTransportWs),
"graphql-ws" => return Some(Self::GraphqlWs),
_ => {},
}
}
None
}
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::GraphqlTransportWs => "graphql-transport-ws",
Self::GraphqlWs => "graphql-ws",
}
}
}
pub struct ProtocolCodec {
protocol: WsProtocol,
}
impl ProtocolCodec {
#[must_use]
pub const fn new(protocol: WsProtocol) -> Self {
Self { protocol }
}
#[must_use]
pub const fn protocol(&self) -> WsProtocol {
self.protocol
}
pub fn decode(&self, raw: &str) -> Result<ClientMessage, ProtocolError> {
match self.protocol {
WsProtocol::GraphqlTransportWs => {
serde_json::from_str(raw).map_err(|e| ProtocolError::InvalidJson(e.to_string()))
},
WsProtocol::GraphqlWs => {
let mut msg: ClientMessage = serde_json::from_str(raw)
.map_err(|e| ProtocolError::InvalidJson(e.to_string()))?;
msg.message_type = translate_legacy_client_type(&msg.message_type).to_string();
Ok(msg)
},
}
}
pub fn encode(&self, msg: &ServerMessage) -> Result<Option<String>, ProtocolError> {
match self.protocol {
WsProtocol::GraphqlTransportWs => {
let json =
msg.to_json().map_err(|e| ProtocolError::SerializationFailed(e.to_string()))?;
Ok(Some(json))
},
WsProtocol::GraphqlWs => {
let wire_type = translate_legacy_server_type(&msg.message_type);
if wire_type.is_none() {
return Ok(None);
}
let wire_type = wire_type.expect("wire_type is Some; None was returned above");
if wire_type == "ka" {
let ka = serde_json::json!({"type": "ka"});
return Ok(Some(ka.to_string()));
}
let mut value = serde_json::to_value(msg)
.map_err(|e| ProtocolError::SerializationFailed(e.to_string()))?;
if let Some(obj) = value.as_object_mut() {
obj.insert(
"type".to_string(),
serde_json::Value::String(wire_type.to_string()),
);
}
let json = serde_json::to_string(&value)
.map_err(|e| ProtocolError::SerializationFailed(e.to_string()))?;
Ok(Some(json))
},
}
}
#[must_use]
pub fn uses_keepalive(&self) -> bool {
self.protocol == WsProtocol::GraphqlWs
}
}
fn translate_legacy_client_type(legacy: &str) -> &str {
match legacy {
"start" => "subscribe",
"stop" => "complete",
other => other,
}
}
fn translate_legacy_server_type(modern: &str) -> Option<&str> {
match modern {
"next" => Some("data"),
"ping" => Some("ka"),
"pong" => None,
other => Some(other),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum ProtocolError {
InvalidJson(String),
SerializationFailed(String),
}
impl std::fmt::Display for ProtocolError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::InvalidJson(e) => write!(f, "invalid JSON: {e}"),
Self::SerializationFailed(e) => write!(f, "serialization failed: {e}"),
}
}
}
impl std::error::Error for ProtocolError {}