Skip to main content

prism_mcp_rs/utils/
uri.rs

1//! URI handling utilities
2//!
3//! Module provides utilities for parsing, validating, and manipulating URIs
4//! used in the MCP protocol for resources and other operations.
5
6use crate::core::error::{McpError, McpResult};
7use std::collections::HashMap;
8use url::Url;
9
10/// Parse a URI and extract query parameters
11pub fn parse_uri_with_params(uri: &str) -> McpResult<(String, HashMap<String, String>)> {
12    if uri.starts_with("file:///") || uri.contains("://") {
13        // Full URI
14        let parsed = Url::parse(uri)
15            .map_err(|e| McpError::InvalidUri(format!("Invalid URI '{uri}': {e}")))?;
16
17        let base_uri = format!(
18            "{}://{}{}",
19            parsed.scheme(),
20            parsed.host_str().unwrap_or(""),
21            parsed.path()
22        );
23
24        let mut params = HashMap::new();
25        for (key, value) in parsed.query_pairs() {
26            params.insert(key.to_string(), value.to_string());
27        }
28
29        Ok((base_uri, params))
30    } else if uri.starts_with('/') {
31        // Absolute path
32        if let Some((path, query)) = uri.split_once('?') {
33            let params = parse_query_string(query)?;
34            Ok((path.to_string(), params))
35        } else {
36            Ok((uri.to_string(), HashMap::new()))
37        }
38    } else {
39        // Relative path or simple identifier
40        if let Some((path, query)) = uri.split_once('?') {
41            let params = parse_query_string(query)?;
42            Ok((path.to_string(), params))
43        } else {
44            Ok((uri.to_string(), HashMap::new()))
45        }
46    }
47}
48
49/// Parse a query string into parameters
50pub fn parse_query_string(query: &str) -> McpResult<HashMap<String, String>> {
51    let mut params = HashMap::new();
52
53    for pair in query.split('&') {
54        if pair.is_empty() {
55            continue;
56        }
57
58        if let Some((key, value)) = pair.split_once('=') {
59            let decoded_key = percent_decode(key)?;
60            let decoded_value = percent_decode(value)?;
61            params.insert(decoded_key, decoded_value);
62        } else {
63            let decoded_key = percent_decode(pair)?;
64            params.insert(decoded_key, String::new());
65        }
66    }
67
68    Ok(params)
69}
70
71/// Simple percent decoding for URI components
72pub fn percent_decode(s: &str) -> McpResult<String> {
73    let mut result = String::new();
74    let mut chars = s.chars().peekable();
75
76    while let Some(ch) = chars.next() {
77        if ch == '%' {
78            let hex1 = chars
79                .next()
80                .ok_or_else(|| McpError::InvalidUri("Incomplete percent encoding".to_string()))?;
81            let hex2 = chars
82                .next()
83                .ok_or_else(|| McpError::InvalidUri("Incomplete percent encoding".to_string()))?;
84
85            let hex_str = format!("{hex1}{hex2}");
86            let byte = u8::from_str_radix(&hex_str, 16).map_err(|_| {
87                McpError::InvalidUri(format!("Invalid hex in percent encoding: {hex_str}"))
88            })?;
89
90            result.push(byte as char);
91        } else if ch == '+' {
92            result.push(' ');
93        } else {
94            result.push(ch);
95        }
96    }
97
98    Ok(result)
99}
100
101/// Simple percent encoding for URI components
102pub fn percent_encode(s: &str) -> String {
103    let mut result = String::new();
104
105    for byte in s.bytes() {
106        match byte {
107            b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
108                result.push(byte as char);
109            }
110            b' ' => {
111                result.push('+');
112            }
113            _ => {
114                result.push_str(&format!("%{byte:02X}"));
115            }
116        }
117    }
118
119    result
120}
121
122/// Validate that a string is a valid URI
123pub fn validate_uri(uri: &str) -> McpResult<()> {
124    if uri.is_empty() {
125        return Err(McpError::InvalidUri("URI cannot be empty".to_string()));
126    }
127
128    // Check for basic URI patterns
129    if uri.contains("://") {
130        // Full URI - try to parse with url crate
131        Url::parse(uri).map_err(|e| McpError::InvalidUri(format!("Invalid URI '{uri}': {e}")))?;
132    } else if uri.starts_with('/') {
133        // Absolute path - basic validation
134        if uri.contains('\0') || uri.contains('\n') || uri.contains('\r') {
135            return Err(McpError::InvalidUri(
136                "URI contains invalid characters".to_string(),
137            ));
138        }
139    } else {
140        // Relative path or identifier - allow most characters
141        if uri.contains('\0') || uri.contains('\n') || uri.contains('\r') {
142            return Err(McpError::InvalidUri(
143                "URI contains invalid characters".to_string(),
144            ));
145        }
146    }
147
148    Ok(())
149}
150
151/// Normalize a URI to a standard form
152pub fn normalize_uri(uri: &str) -> McpResult<String> {
153    validate_uri(uri)?;
154
155    if uri.contains("://") {
156        // Full URI - normalize with url crate
157        let parsed = Url::parse(uri)
158            .map_err(|e| McpError::InvalidUri(format!("Invalid URI '{uri}': {e}")))?;
159        let mut normalized = parsed.to_string();
160
161        // Remove duplicate slashes in path
162        if let Ok(mut url) = Url::parse(&normalized) {
163            let path = url.path();
164            let clean_path = path.replace("//", "/");
165            url.set_path(&clean_path);
166            normalized = url.to_string();
167        }
168
169        // Remove trailing slash unless it's the root
170        if normalized.ends_with('/') && !normalized.ends_with("://") {
171            let path_start = match normalized.find("://") {
172                Some(pos) => pos + 3,
173                None => {
174                    return Err(McpError::InvalidUri(format!(
175                        "URI '{normalized}' missing protocol separator"
176                    )));
177                }
178            };
179            if let Some(path_start_slash) = normalized[path_start..].find('/') {
180                let full_path_start = path_start + path_start_slash;
181                if full_path_start + 1 < normalized.len() {
182                    normalized.pop();
183                }
184            }
185        }
186
187        Ok(normalized)
188    } else {
189        // Path - basic normalization
190        let mut normalized = uri.to_string();
191
192        // Remove duplicate slashes
193        while normalized.contains("//") {
194            normalized = normalized.replace("//", "/");
195        }
196
197        // Remove trailing slash unless it's the root
198        if normalized.len() > 1 && normalized.ends_with('/') {
199            normalized.pop();
200        }
201
202        Ok(normalized)
203    }
204}
205
206/// Join a base URI with a relative path
207pub fn join_uri(base: &str, relative: &str) -> McpResult<String> {
208    if relative.contains("://") {
209        // Relative is actually absolute
210        return Ok(relative.to_string());
211    }
212
213    if relative.starts_with('/') {
214        // Relative path is absolute, return it as-is
215        return Ok(relative.to_string());
216    }
217
218    if base.contains("://") {
219        // Full URI base
220        let base_url = Url::parse(base)
221            .map_err(|e| McpError::InvalidUri(format!("Invalid base URI '{base}': {e}")))?;
222        let joined = base_url.join(relative).map_err(|e| {
223            McpError::InvalidUri(format!("Cannot join '{relative}' to '{base}': {e}"))
224        })?;
225        Ok(joined.to_string())
226    } else {
227        // Path base
228        let mut result = base.to_string();
229        if !result.ends_with('/') && !relative.starts_with('/') {
230            result.push('/');
231        }
232        result.push_str(relative);
233        normalize_uri(&result)
234    }
235}
236
237/// Extract the file extension from a URI
238pub fn get_uri_extension(uri: &str) -> Option<String> {
239    let path = if uri.contains("://") {
240        Url::parse(uri).ok()?.path().to_string()
241    } else {
242        uri.to_string()
243    };
244
245    if let Some(dot_pos) = path.rfind('.') {
246        if let Some(slash_pos) = path.rfind('/') {
247            if dot_pos > slash_pos {
248                return Some(path[dot_pos + 1..].to_lowercase());
249            }
250        } else {
251            return Some(path[dot_pos + 1..].to_lowercase());
252        }
253    }
254
255    None
256}
257
258/// Guess MIME type from URI extension
259pub fn guess_mime_type(uri: &str) -> Option<String> {
260    match get_uri_extension(uri)?.as_str() {
261        "txt" => Some("text/plain".to_string()),
262        "html" | "htm" => Some("text/html".to_string()),
263        "css" => Some("text/css".to_string()),
264        "js" => Some("application/javascript".to_string()),
265        "json" => Some("application/json".to_string()),
266        "xml" => Some("application/xml".to_string()),
267        "pdf" => Some("application/pdf".to_string()),
268        "zip" => Some("application/zip".to_string()),
269        "png" => Some("image/png".to_string()),
270        "jpg" | "jpeg" => Some("image/jpeg".to_string()),
271        "gif" => Some("image/gif".to_string()),
272        "webp" => Some("image/webp".to_string()),
273        "svg" => Some("image/svg+xml".to_string()),
274        "mp3" => Some("audio/mpeg".to_string()),
275        "wav" => Some("audio/wav".to_string()),
276        "mp4" => Some("video/mp4".to_string()),
277        "webm" => Some("video/webm".to_string()),
278        "csv" => Some("text/csv".to_string()),
279        "md" => Some("text/markdown".to_string()),
280        "yaml" | "yml" => Some("application/x-yaml".to_string()),
281        "toml" => Some("application/toml".to_string()),
282        _ => None,
283    }
284}
285
286#[cfg(test)]
287mod tests {
288    use super::*;
289
290    #[test]
291    fn test_parse_uri_with_params() {
292        let (uri, params) =
293            parse_uri_with_params("https://example.com/path?key=value&foo=bar").unwrap();
294        assert_eq!(uri, "https://example.com/path");
295        assert_eq!(params.get("key"), Some(&"value".to_string()));
296        assert_eq!(params.get("foo"), Some(&"bar".to_string()));
297    }
298
299    #[test]
300    fn test_parse_query_string() {
301        let params = parse_query_string("key=value&foo=bar&empty=").unwrap();
302        assert_eq!(params.get("key"), Some(&"value".to_string()));
303        assert_eq!(params.get("foo"), Some(&"bar".to_string()));
304        assert_eq!(params.get("empty"), Some(&"".to_string()));
305    }
306
307    #[test]
308    fn test_percent_encode_decode() {
309        let original = "hello world!@#$%";
310        let encoded = percent_encode(original);
311        let decoded = percent_decode(&encoded).unwrap();
312        assert_eq!(decoded, original);
313    }
314
315    #[test]
316    fn test_validate_uri() {
317        assert!(validate_uri("https://example.com").is_ok());
318        assert!(validate_uri("/absolute/path").is_ok());
319        assert!(validate_uri("relative/path").is_ok());
320        assert!(validate_uri("").is_err());
321        assert!(validate_uri("invalid\0uri").is_err());
322    }
323
324    #[test]
325    fn test_normalize_uri() {
326        assert_eq!(
327            normalize_uri("https://example.com//path//").unwrap(),
328            "https://example.com/path"
329        );
330        assert_eq!(normalize_uri("/path//to//file/").unwrap(), "/path/to/file");
331        assert_eq!(normalize_uri("/").unwrap(), "/");
332    }
333
334    #[test]
335    fn test_join_uri() {
336        assert_eq!(
337            join_uri("https://example.com", "path/to/file").unwrap(),
338            "https://example.com/path/to/file"
339        );
340        assert_eq!(
341            join_uri("/base", "relative/path").unwrap(),
342            "/base/relative/path"
343        );
344        assert_eq!(join_uri("/base/", "/absolute").unwrap(), "/absolute");
345    }
346
347    #[test]
348    fn test_get_uri_extension() {
349        assert_eq!(get_uri_extension("file.txt"), Some("txt".to_string()));
350        assert_eq!(
351            get_uri_extension("https://example.com/file.JSON"),
352            Some("json".to_string())
353        );
354        assert_eq!(
355            get_uri_extension("/path/to/file.tar.gz"),
356            Some("gz".to_string())
357        );
358        assert_eq!(get_uri_extension("no-extension"), None);
359    }
360
361    #[test]
362    fn test_guess_mime_type() {
363        assert_eq!(
364            guess_mime_type("file.json"),
365            Some("application/json".to_string())
366        );
367        assert_eq!(guess_mime_type("image.PNG"), Some("image/png".to_string()));
368        assert_eq!(
369            guess_mime_type("document.pdf"),
370            Some("application/pdf".to_string())
371        );
372        assert_eq!(guess_mime_type("unknown.xyz"), None);
373    }
374}