Skip to main content

xidl_parser/rest_hir/semantics/
upgrade.rs

1use 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/// WebSocket-mode configuration parsed from `@upgrade(protocol = "websocket", ...)`.
21#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
22pub struct WebSocketConfig {
23    pub codec: WebSocketCodec,
24    pub subprotocol: Option<String>,
25    /// Active ping interval in milliseconds. `None` or `0` means no active ping.
26    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    // RFC 6455 token: 1*( %x21 / %x23-2B / %x2D-3A / %x3C-5B / %x5D-7E )
146    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
161/// Parses duration strings such as `"500ms"`, `"20s"`, `"1m"`, or `"0"`.
162pub 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    // bare integer is treated as milliseconds
181    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}