prism_mcp_rs/utils/
uri.rs1use crate::core::error::{McpError, McpResult};
7use std::collections::HashMap;
8use url::Url;
9
10pub fn parse_uri_with_params(uri: &str) -> McpResult<(String, HashMap<String, String>)> {
12 if uri.starts_with("file:///") || uri.contains("://") {
13 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 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 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
49pub 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
71pub 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
101pub 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
122pub 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 if uri.contains("://") {
130 Url::parse(uri).map_err(|e| McpError::InvalidUri(format!("Invalid URI '{uri}': {e}")))?;
132 } else if uri.starts_with('/') {
133 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 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
151pub fn normalize_uri(uri: &str) -> McpResult<String> {
153 validate_uri(uri)?;
154
155 if uri.contains("://") {
156 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 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 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 let mut normalized = uri.to_string();
191
192 while normalized.contains("//") {
194 normalized = normalized.replace("//", "/");
195 }
196
197 if normalized.len() > 1 && normalized.ends_with('/') {
199 normalized.pop();
200 }
201
202 Ok(normalized)
203 }
204}
205
206pub fn join_uri(base: &str, relative: &str) -> McpResult<String> {
208 if relative.contains("://") {
209 return Ok(relative.to_string());
211 }
212
213 if relative.starts_with('/') {
214 return Ok(relative.to_string());
216 }
217
218 if base.contains("://") {
219 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 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
237pub 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
258pub 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}