xidl_parser/rest_hir/semantics/
upgrade.rs1use crate::hir;
2use serde::{Deserialize, Serialize};
3use std::collections::HashMap;
4
5use super::annotations::{annotation_name, annotation_params, normalize_annotation_params};
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
8pub enum UpgradeMode {
9 WebSocket,
10 Raw,
11}
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
14pub enum WebSocketCodec {
15 Json,
16 Msgpack,
17 Bytes,
18}
19
20#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
22pub struct WebSocketConfig {
23 pub codec: WebSocketCodec,
24 pub subprotocol: Option<String>,
25 pub heartbeat_ms: Option<u64>,
27 pub max_message_bytes: Option<u64>,
28}
29
30impl Default for WebSocketConfig {
31 fn default() -> Self {
32 Self {
33 codec: WebSocketCodec::Json,
34 subprotocol: None,
35 heartbeat_ms: None,
36 max_message_bytes: None,
37 }
38 }
39}
40
41pub fn classify_upgrade_protocol(protocol: &str) -> Result<UpgradeMode, String> {
42 let trimmed = protocol.trim();
43 if trimmed.is_empty() {
44 return Err("@upgrade protocol parameter cannot be empty".to_string());
45 }
46 let lower = trimmed.to_ascii_lowercase();
47 match lower.as_str() {
48 "websocket" => Ok(UpgradeMode::WebSocket),
49 "ws" | "wss" => Err(
50 "@upgrade protocol cannot be \"ws\" or \"wss\"; use \"websocket\" instead".to_string(),
51 ),
52 _ => Ok(UpgradeMode::Raw),
53 }
54}
55
56fn find_upgrade_params(annotations: &[hir::Annotation]) -> Option<HashMap<String, String>> {
57 let annotation = annotations.iter().find(|annotation| {
58 annotation_name(annotation)
59 .map(|name| name.eq_ignore_ascii_case("upgrade"))
60 .unwrap_or(false)
61 })?;
62 annotation_params(annotation).map(normalize_annotation_params)
63}
64
65pub fn parse_upgrade_protocol(annotations: &[hir::Annotation]) -> Result<Option<String>, String> {
66 let Some(params) = find_upgrade_params(annotations) else {
67 return Ok(None);
68 };
69 let Some(proto) = params.get("protocol").cloned() else {
70 return Err(
71 "@upgrade annotation requires a protocol parameter (e.g. @upgrade(protocol=\"xidl-raw\"))"
72 .to_string(),
73 );
74 };
75 Ok(Some(proto))
76}
77
78pub fn parse_websocket_config(
79 annotations: &[hir::Annotation],
80 mode: UpgradeMode,
81) -> Result<Option<WebSocketConfig>, String> {
82 let Some(params) = find_upgrade_params(annotations) else {
83 return Ok(None);
84 };
85 if mode != UpgradeMode::WebSocket {
86 for key in ["codec", "subprotocol", "heartbeat", "max_message"] {
87 if params.contains_key(key) {
88 return Err(format!(
89 "@upgrade parameter {key} is only valid with protocol = \"websocket\""
90 ));
91 }
92 }
93 for key in params.keys() {
94 if key != "protocol" {
95 return Err(format!("@upgrade unknown parameter {key}"));
96 }
97 }
98 return Ok(None);
99 }
100
101 let mut config = WebSocketConfig::default();
102 if let Some(codec) = params.get("codec") {
103 config.codec = match codec.to_ascii_lowercase().as_str() {
104 "json" => WebSocketCodec::Json,
105 "msgpack" => WebSocketCodec::Msgpack,
106 "bytes" => WebSocketCodec::Bytes,
107 other => {
108 return Err(format!(
109 "unsupported @upgrade codec {other}, expected json, msgpack, or bytes"
110 ));
111 }
112 };
113 }
114 if let Some(sub) = params.get("subprotocol") {
115 validate_websocket_subprotocol(sub)?;
116 config.subprotocol = Some(sub.clone());
117 }
118 if let Some(heartbeat) = params.get("heartbeat") {
119 config.heartbeat_ms = Some(parse_duration_ms(heartbeat)?);
120 }
121 if let Some(max_message) = params.get("max_message") {
122 let value = max_message.parse::<u64>().map_err(|_| {
123 format!("@upgrade max_message must be a positive integer, got {max_message}")
124 })?;
125 if value < 1 {
126 return Err("@upgrade max_message must be >= 1".to_string());
127 }
128 config.max_message_bytes = Some(value);
129 }
130 for key in params.keys() {
131 match key.as_str() {
132 "protocol" | "codec" | "subprotocol" | "heartbeat" | "max_message" => {}
133 other => {
134 return Err(format!("@upgrade unknown parameter {other}"));
135 }
136 }
137 }
138 Ok(Some(config))
139}
140
141pub fn validate_websocket_subprotocol(value: &str) -> Result<(), String> {
142 if value.is_empty() {
143 return Err("@upgrade subprotocol cannot be empty".to_string());
144 }
145 let ok = value.bytes().all(|b| {
147 b == 0x21
148 || (0x23..=0x2B).contains(&b)
149 || (0x2D..=0x3A).contains(&b)
150 || (0x3C..=0x5B).contains(&b)
151 || (0x5D..=0x7E).contains(&b)
152 });
153 if !ok {
154 return Err(format!(
155 "@upgrade subprotocol must be an RFC 6455 token, got {value}"
156 ));
157 }
158 Ok(())
159}
160
161pub fn parse_duration_ms(value: &str) -> Result<u64, String> {
163 let trimmed = value.trim();
164 if trimmed.is_empty() {
165 return Err("@upgrade heartbeat cannot be empty".to_string());
166 }
167 let parse_num = |num: &str| -> Result<u64, String> {
168 num.parse::<u64>()
169 .map_err(|_| format!("invalid @upgrade heartbeat duration {value}"))
170 };
171 if let Some(num) = trimmed.strip_suffix("ms") {
172 return parse_num(num);
173 }
174 if let Some(num) = trimmed.strip_suffix('s') {
175 return Ok(parse_num(num)?.saturating_mul(1000));
176 }
177 if let Some(num) = trimmed.strip_suffix('m') {
178 return Ok(parse_num(num)?.saturating_mul(60_000));
179 }
180 parse_num(trimmed)
182}
183
184#[cfg(test)]
185mod tests {
186 use super::*;
187
188 #[test]
189 fn classify_protocol_modes() {
190 assert_eq!(
191 classify_upgrade_protocol("websocket").unwrap(),
192 UpgradeMode::WebSocket
193 );
194 assert_eq!(
195 classify_upgrade_protocol("WebSocket").unwrap(),
196 UpgradeMode::WebSocket
197 );
198 assert_eq!(
199 classify_upgrade_protocol("fastnet").unwrap(),
200 UpgradeMode::Raw
201 );
202 assert!(classify_upgrade_protocol("ws").is_err());
203 assert!(classify_upgrade_protocol("wss").is_err());
204 assert!(classify_upgrade_protocol("").is_err());
205 }
206
207 #[test]
208 fn parse_duration_variants() {
209 assert_eq!(parse_duration_ms("500ms").unwrap(), 500);
210 assert_eq!(parse_duration_ms("20s").unwrap(), 20_000);
211 assert_eq!(parse_duration_ms("1m").unwrap(), 60_000);
212 assert_eq!(parse_duration_ms("0").unwrap(), 0);
213 assert!(parse_duration_ms("abc").is_err());
214 }
215
216 #[test]
217 fn subprotocol_token_rules() {
218 assert!(validate_websocket_subprotocol("fastnet.v1").is_ok());
219 assert!(validate_websocket_subprotocol("").is_err());
220 assert!(validate_websocket_subprotocol("bad token").is_err());
221 }
222}