#[cfg(feature = "websocket")]
use serde::{Deserialize, Serialize};
#[cfg(feature = "websocket")]
use serde_json::Value;
#[cfg(feature = "websocket")]
pub const MAX_MESSAGE_SIZE: usize = 1_048_576;
#[cfg(feature = "websocket")]
pub const MAX_JSON_DEPTH: usize = 16;
#[cfg(feature = "websocket")]
pub const MAX_STRING_LENGTH: usize = 64 * 1024;
#[cfg(feature = "websocket")]
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum WebSocketMessage {
#[serde(rename = "request")]
Request {
id: String,
method: String,
params: serde_json::Value,
},
#[serde(rename = "response")]
Response {
id: String,
result: serde_json::Value,
},
#[serde(rename = "error")]
Error {
id: String,
error: String,
},
#[serde(rename = "notification")]
Notification {
event: String,
data: serde_json::Value,
},
}
#[cfg(feature = "websocket")]
pub fn parse_websocket_message(text: &str) -> Result<WebSocketMessage, String> {
if text.len() > MAX_MESSAGE_SIZE {
return Err(format!(
"Message too large: {} bytes (max: {} bytes)",
text.len(),
MAX_MESSAGE_SIZE
));
}
use serde_json::Deserializer;
let mut max_depth = 0;
let mut current_depth = 0;
let deserializer = Deserializer::from_str(text);
for result in deserializer.into_iter::<Value>() {
match result {
Ok(value) => {
let depth = calculate_value_depth(&value, &mut current_depth);
max_depth = max_depth.max(depth);
if max_depth > MAX_JSON_DEPTH {
return Err(format!(
"JSON nesting too deep: depth {} (max: {})",
max_depth, MAX_JSON_DEPTH
));
}
}
Err(e) => {
return Err(format!("Invalid JSON: {}", e));
}
}
}
let msg = serde_json::from_str::<WebSocketMessage>(text)
.map_err(|e| format!("Invalid JSON: {}", e))?;
validate_string_limits(&msg)?;
Ok(msg)
}
pub fn calculate_value_depth(value: &serde_json::Value, current_depth: &mut usize) -> usize {
match value {
Value::Object(map) => {
*current_depth += 1;
let level = *current_depth;
let max_child_depth = map
.values()
.map(|v| calculate_value_depth(v, current_depth))
.max()
.unwrap_or(0);
*current_depth -= 1;
max_child_depth.max(level)
}
Value::Array(arr) => {
*current_depth += 1;
let level = *current_depth;
let max_child_depth = arr
.iter()
.map(|v| calculate_value_depth(v, current_depth))
.max()
.unwrap_or(0);
*current_depth -= 1;
max_child_depth.max(level)
}
_ => *current_depth,
}
}
fn validate_string_limits(msg: &WebSocketMessage) -> Result<(), String> {
let too_long = |field: &str, len: usize| -> String {
format!(
"String field '{}' too long: {} bytes (max: {})",
field, len, MAX_STRING_LENGTH
)
};
match msg {
WebSocketMessage::Request { id, method, .. } => {
if id.len() > MAX_STRING_LENGTH {
return Err(too_long("id", id.len()));
}
if method.len() > MAX_STRING_LENGTH {
return Err(too_long("method", method.len()));
}
}
WebSocketMessage::Response { id, .. } => {
if id.len() > MAX_STRING_LENGTH {
return Err(too_long("id", id.len()));
}
}
WebSocketMessage::Error { id, error } => {
if id.len() > MAX_STRING_LENGTH {
return Err(too_long("id", id.len()));
}
if error.len() > MAX_STRING_LENGTH {
return Err(too_long("error", error.len()));
}
}
WebSocketMessage::Notification { event, .. } => {
if event.len() > MAX_STRING_LENGTH {
return Err(too_long("event", event.len()));
}
}
}
Ok(())
}
#[cfg(test)]
pub fn calculate_json_depth(text: &str) -> usize {
let mut depth = 0;
let mut max_depth = 0;
let mut in_string = false;
let mut escaped = false;
for c in text.chars() {
if in_string {
if escaped {
escaped = false;
} else if c == '\\' {
escaped = true;
} else if c == '"' {
in_string = false;
}
} else if c == '"' {
in_string = true;
escaped = false;
} else if c == '{' || c == '[' {
depth += 1;
max_depth = max_depth.max(depth);
} else if (c == '}' || c == ']') && depth > 0 {
depth -= 1;
}
}
max_depth
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn calculate_value_depth_primitive_at_nonzero_depth() {
let value = serde_json::json!(42);
let mut depth = 5;
let result = calculate_value_depth(&value, &mut depth);
assert_eq!(result, 5);
assert_eq!(depth, 5);
}
#[test]
fn calculate_value_depth_object_with_empty_child() {
let value = serde_json::json!({"a": {}, "b": 1});
let mut depth = 0;
let result = calculate_value_depth(&value, &mut depth);
assert_eq!(result, 2);
assert_eq!(depth, 0);
}
#[test]
fn calculate_value_depth_array_with_empty_child() {
let value = serde_json::json!([[], 1]);
let mut depth = 0;
let result = calculate_value_depth(&value, &mut depth);
assert_eq!(result, 2);
assert_eq!(depth, 0);
}
#[test]
fn parse_websocket_message_rejects_deep_empty_containers() {
let mut deep_json = String::from("[]");
for _ in 0..=MAX_JSON_DEPTH {
deep_json = format!("[{}]", deep_json);
}
let result = parse_websocket_message(&deep_json);
assert!(result.is_err());
assert!(
result.unwrap_err().contains("nesting too deep"),
"deep empty-container nesting must be rejected"
);
}
#[test]
fn parse_websocket_message_rejects_oversized_string_fields() {
let long_id = "x".repeat(MAX_STRING_LENGTH + 1);
let json = serde_json::json!({
"type": "request",
"id": long_id,
"method": "m",
"params": {}
})
.to_string();
let result = parse_websocket_message(&json);
assert!(result.is_err());
assert!(
result.unwrap_err().contains("too long"),
"oversized string field must be rejected"
);
let boundary_id = "x".repeat(MAX_STRING_LENGTH);
let json = serde_json::json!({
"type": "request",
"id": boundary_id,
"method": "m",
"params": {}
})
.to_string();
assert!(parse_websocket_message(&json).is_ok());
}
#[test]
fn calculate_json_depth_unmatched_closing_brace() {
assert_eq!(calculate_json_depth("}"), 0);
assert_eq!(calculate_json_depth("]"), 0);
assert_eq!(calculate_json_depth("}]}]"), 0);
}
#[test]
fn calculate_json_depth_with_escape_sequence_in_string() {
let text = r#"{"a":"b\\c"}"#;
let depth = calculate_json_depth(text);
assert_eq!(
depth, 1,
"Single-level JSON with escape should return depth 1"
);
}
#[test]
fn parse_websocket_message_multiple_top_level_values() {
let json = r#"{"type":"request","id":"1","method":"m","params":{}}{"extra":true}"#;
let result = parse_websocket_message(json);
assert!(result.is_err());
}
}