1use crate::error::{Result, WebmcpError};
2use std::collections::{HashMap, HashSet};
3use std::sync::{Arc, Mutex};
4use std::time::{Duration, Instant};
5use uuid::Uuid;
6
7const PAIRING_CODE_HEX_DIGITS: usize = 12;
8const MAX_FAILED_PAIRING_ATTEMPTS: u8 = 5;
9
10pub(crate) fn is_valid_origin(origin: &str) -> bool {
11 let Ok(parsed) = url::Url::parse(origin) else {
12 return false;
13 };
14 origin == origin.trim()
15 && !origin.chars().any(char::is_whitespace)
16 && !origin.contains('*')
17 && matches!(parsed.scheme(), "http" | "https")
18 && parsed.host_str().is_some_and(|host| !host.is_empty())
19 && parsed.username().is_empty()
20 && parsed.password().is_none()
21 && (parsed.path().is_empty() || (parsed.path() == "/" && !origin.ends_with('/')))
22 && parsed.query().is_none()
23 && parsed.fragment().is_none()
24}
25
26#[derive(Debug)]
27struct PendingPairing {
28 code: String,
29 origin: Option<String>,
30 expires_at: Instant,
31 used: bool,
32 failed_attempts: u8,
33}
34
35#[derive(Debug)]
36struct PairingState {
37 pending: Option<PendingPairing>,
38 sessions: HashMap<String, PairingSessionState>,
39}
40
41#[derive(Debug)]
42struct PairingSessionState {
43 origin: String,
44 expires_at: Instant,
45}
46
47#[derive(Clone)]
49pub struct PairingDisplay {
50 code: String,
51 expires_at: Instant,
52}
53
54impl PairingDisplay {
55 pub fn code(&self) -> &str {
57 &self.code
58 }
59
60 pub fn expires_in(&self) -> Duration {
62 self.expires_at.saturating_duration_since(Instant::now())
63 }
64}
65
66impl std::fmt::Debug for PairingDisplay {
67 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
68 formatter
69 .debug_struct("PairingDisplay")
70 .field("expires_in", &self.expires_in())
71 .finish_non_exhaustive()
72 }
73}
74
75#[derive(Clone)]
77pub struct PairingSession {
78 token: String,
79 origin: String,
80 expires_at: Instant,
81}
82
83impl PairingSession {
84 pub fn token(&self) -> &str {
86 &self.token
87 }
88
89 pub fn origin(&self) -> &str {
91 &self.origin
92 }
93
94 pub fn expires_in(&self) -> Duration {
96 self.expires_at.saturating_duration_since(Instant::now())
97 }
98}
99
100impl std::fmt::Debug for PairingSession {
101 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
102 formatter
103 .debug_struct("PairingSession")
104 .field("origin", &self.origin)
105 .field("expires_in", &self.expires_in())
106 .finish_non_exhaustive()
107 }
108}
109
110#[derive(Clone)]
112pub struct PairingManager {
113 allowed_origins: Arc<HashSet<String>>,
114 ttl: Duration,
115 state: Arc<Mutex<PairingState>>,
116}
117
118impl PairingManager {
119 pub fn new<I, S>(allowed_origins: I, ttl: Duration) -> Result<Self>
121 where
122 I: IntoIterator<Item = S>,
123 S: Into<String>,
124 {
125 if ttl.is_zero() || ttl > Duration::from_secs(3600) {
126 return Err(WebmcpError::InvalidRequest("pairing ttl must be between 1 and 3600 seconds".to_string()));
127 }
128
129 let origins = allowed_origins.into_iter().map(Into::into).collect::<Vec<String>>();
130 if origins.iter().any(|origin| !is_valid_origin(origin)) {
131 return Err(WebmcpError::InvalidRequest("WebMCP origins must be exact non-wildcard origins".to_string()));
132 }
133
134 Ok(Self {
135 allowed_origins: Arc::new(origins.into_iter().collect()),
136 ttl,
137 state: Arc::new(Mutex::new(PairingState { pending: None, sessions: HashMap::new() })),
138 })
139 }
140
141 pub fn begin_pairing(&self) -> PairingDisplay {
143 self.begin_pairing_inner(None, false)
144 }
145
146 fn begin_pairing_inner(&self, origin: Option<String>, revoke_sessions: bool) -> PairingDisplay {
147 let code = Uuid::new_v4()
148 .simple()
149 .to_string()
150 .chars()
151 .take(PAIRING_CODE_HEX_DIGITS)
152 .collect::<String>()
153 .to_ascii_uppercase();
154 let expires_at = Instant::now() + self.ttl;
155 let mut state = self.state.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
156 if revoke_sessions {
157 state.sessions.clear();
158 }
159 state.pending = Some(PendingPairing {
160 code: code.clone(),
161 origin,
162 expires_at,
163 used: false,
164 failed_attempts: 0,
165 });
166 PairingDisplay { code, expires_at }
167 }
168
169 pub fn begin_pairing_for_origin(&self, origin: impl Into<String>) -> Result<PairingDisplay> {
171 let origin = origin.into();
172 self.ensure_origin_allowed(&origin)?;
173 Ok(self.begin_pairing_inner(Some(origin), false))
174 }
175
176 pub fn replace_pairing_for_origin(&self, origin: impl Into<String>) -> Result<PairingDisplay> {
178 let origin = origin.into();
179 self.ensure_origin_allowed(&origin)?;
180 Ok(self.begin_pairing_inner(Some(origin), true))
181 }
182
183 pub fn pair(&self, code: &str, origin: &str) -> Result<PairingSession> {
185 self.ensure_origin_allowed(origin)?;
186 let now = Instant::now();
187 let mut state = self.state.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
188 let pending = state.pending.as_mut().ok_or(WebmcpError::PairingExpired)?;
189 if pending.expires_at <= now {
190 state.pending = None;
191 return Err(WebmcpError::PairingExpired);
192 }
193 if pending.used {
194 return Err(WebmcpError::PairingUsed);
195 }
196 if pending.code != code {
197 pending.failed_attempts = pending.failed_attempts.saturating_add(1);
198 if pending.failed_attempts >= MAX_FAILED_PAIRING_ATTEMPTS {
199 state.pending = None;
200 }
201 return Err(WebmcpError::PairingExpired);
202 }
203 if pending.origin.as_deref().is_some_and(|expected| expected != origin) {
204 return Err(WebmcpError::OriginRejected(origin.to_string()));
205 }
206
207 pending.used = true;
208 let token = Uuid::new_v4().simple().to_string();
209 let expires_at = now + self.ttl;
210 drop(
211 state
212 .sessions
213 .insert(token.clone(), PairingSessionState { origin: origin.to_string(), expires_at }),
214 );
215 Ok(PairingSession { token, origin: origin.to_string(), expires_at })
216 }
217
218 pub fn validate(&self, token: &str, origin: &str) -> Result<()> {
220 self.validate_session(token, origin, false)
221 }
222
223 pub fn refresh(&self, token: &str, origin: &str) -> Result<()> {
229 self.validate_session(token, origin, true)
230 }
231
232 fn validate_session(&self, token: &str, origin: &str, refresh: bool) -> Result<()> {
233 self.ensure_origin_allowed(origin)?;
234 let now = Instant::now();
235 let mut state = self.state.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
236 let valid = state
237 .sessions
238 .get(token)
239 .is_some_and(|session| session.expires_at > now && session.origin == origin);
240 if !valid {
241 drop(state.sessions.remove(token));
242 return Err(WebmcpError::Unauthorized);
243 }
244 if refresh && let Some(session) = state.sessions.get_mut(token) {
245 session.expires_at = now + self.ttl;
246 }
247 Ok(())
248 }
249
250 pub fn resume(&self, token: &str, origin: &str) -> Result<PairingSession> {
252 self.refresh(token, origin)?;
253 let state = self.state.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
254 let session = state.sessions.get(token).ok_or(WebmcpError::Unauthorized)?;
255 Ok(PairingSession {
256 token: token.to_string(),
257 origin: session.origin.clone(),
258 expires_at: session.expires_at,
259 })
260 }
261
262 pub fn revoke(&self, token: &str) -> bool {
264 let mut state = self.state.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
265 state.sessions.remove(token).is_some()
266 }
267
268 pub fn revoke_all(&self) {
270 let mut state = self.state.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
271 state.sessions.clear();
272 state.pending = None;
273 }
274
275 pub fn is_origin_allowed(&self, origin: &str) -> bool {
277 self.allowed_origins.contains(origin)
278 }
279
280 fn ensure_origin_allowed(&self, origin: &str) -> Result<()> {
281 if origin.is_empty() || !self.allowed_origins.contains(origin) {
282 return Err(WebmcpError::OriginRejected(origin.to_string()));
283 }
284 Ok(())
285 }
286}
287
288#[cfg(test)]
289mod tests {
290 use super::*;
291
292 #[test]
293 fn pairing_is_one_time_and_origin_bound() {
294 let manager = PairingManager::new(["https://example.test"], Duration::from_secs(60)).expect("manager");
295 let display = manager.begin_pairing();
296 assert_eq!(display.code().len(), 12);
297 let session = manager.pair(display.code(), "https://example.test").expect("pair");
298 assert!(manager.validate(session.token(), "https://example.test").is_ok());
299 assert!(matches!(manager.pair(display.code(), "https://example.test"), Err(WebmcpError::PairingUsed)));
300 assert!(matches!(
301 manager.validate(session.token(), "https://other.test"),
302 Err(WebmcpError::OriginRejected(_))
303 ));
304 }
305
306 #[test]
307 fn configured_origins_can_keep_independent_sessions_active() {
308 let manager = PairingManager::new(
309 ["https://vtcode.vinhnx.chatgpt.site", "https://vinhnx.github.io"],
310 Duration::from_secs(60),
311 )
312 .expect("manager");
313
314 let chatgpt_pairing = manager
315 .begin_pairing_for_origin("https://vtcode.vinhnx.chatgpt.site")
316 .expect("ChatGPT Site pairing");
317 let chatgpt_session = manager
318 .pair(chatgpt_pairing.code(), "https://vtcode.vinhnx.chatgpt.site")
319 .expect("ChatGPT Site session");
320
321 let github_pairing = manager
322 .begin_pairing_for_origin("https://vinhnx.github.io")
323 .expect("GitHub Pages pairing");
324 let github_session = manager
325 .pair(github_pairing.code(), "https://vinhnx.github.io")
326 .expect("GitHub Pages session");
327
328 assert!(
329 manager
330 .validate(chatgpt_session.token(), "https://vtcode.vinhnx.chatgpt.site")
331 .is_ok()
332 );
333 assert!(manager.validate(github_session.token(), "https://vinhnx.github.io").is_ok());
334 assert!(matches!(
335 manager.validate(chatgpt_session.token(), "https://vinhnx.github.io"),
336 Err(WebmcpError::Unauthorized)
337 ));
338 }
339
340 #[test]
341 fn expired_pairing_cannot_be_reused() {
342 let manager = PairingManager::new(["https://example.test"], Duration::from_millis(1)).expect("manager");
343 let display = manager.begin_pairing();
344 std::thread::sleep(Duration::from_millis(3));
345 assert!(matches!(manager.pair(display.code(), "https://example.test"), Err(WebmcpError::PairingExpired)));
346 }
347
348 #[test]
349 fn revocation_invalidates_session() {
350 let manager = PairingManager::new(["https://example.test"], Duration::from_secs(60)).expect("manager");
351 let display = manager.begin_pairing();
352 let session = manager.pair(display.code(), "https://example.test").expect("pair");
353 assert!(manager.revoke(session.token()));
354 assert!(matches!(manager.validate(session.token(), "https://example.test"), Err(WebmcpError::Unauthorized)));
355 }
356
357 #[test]
358 fn replacing_pairing_revokes_sessions_and_preserves_origin_validation() {
359 let manager = PairingManager::new(["https://example.test"], Duration::from_secs(60)).expect("manager");
360 let display = manager.begin_pairing();
361 let session = manager.pair(display.code(), "https://example.test").expect("pair");
362
363 let replacement = manager
364 .replace_pairing_for_origin("https://example.test")
365 .expect("replacement pairing");
366 assert!(matches!(manager.validate(session.token(), "https://example.test"), Err(WebmcpError::Unauthorized)));
367 assert!(matches!(manager.pair(display.code(), "https://example.test"), Err(WebmcpError::PairingExpired)));
368 assert!(manager.pair(replacement.code(), "https://example.test").is_ok());
369 }
370
371 #[test]
372 fn rejected_replacement_does_not_revoke_current_session() {
373 let manager = PairingManager::new(["https://example.test"], Duration::from_secs(60)).expect("manager");
374 let display = manager.begin_pairing();
375 let session = manager.pair(display.code(), "https://example.test").expect("pair");
376
377 assert!(matches!(
378 manager.replace_pairing_for_origin("https://other.test"),
379 Err(WebmcpError::OriginRejected(_))
380 ));
381 assert!(manager.validate(session.token(), "https://example.test").is_ok());
382 }
383
384 #[test]
385 fn reconnect_can_resume_without_reusing_the_pairing_code() {
386 let manager = PairingManager::new(["https://example.test"], Duration::from_secs(60)).expect("manager");
387 let display = manager.begin_pairing();
388 let session = manager.pair(display.code(), "https://example.test").expect("pair");
389 let resumed = manager.resume(session.token(), "https://example.test").expect("resume");
390 assert_eq!(resumed.token(), session.token());
391 assert!(matches!(manager.pair(display.code(), "https://example.test"), Err(WebmcpError::PairingUsed)));
392 }
393
394 #[test]
395 fn wildcard_and_malformed_origins_are_rejected() {
396 for origin in [
397 "*",
398 "https://*",
399 " https://example.test ",
400 "localhost",
401 "ftp://example.test",
402 "https://example.test/path",
403 "https://user@example.test",
404 ] {
405 assert!(PairingManager::new([origin], Duration::from_secs(60)).is_err(), "accepted {origin}");
406 }
407 }
408
409 #[test]
410 fn repeated_invalid_codes_expire_the_pending_code() {
411 let manager = PairingManager::new(["https://example.test"], Duration::from_secs(60)).expect("manager");
412 let display = manager.begin_pairing();
413 for _ in 0..5 {
414 assert!(matches!(manager.pair("000000000000", "https://example.test"), Err(WebmcpError::PairingExpired)));
415 }
416 assert!(matches!(manager.pair(display.code(), "https://example.test"), Err(WebmcpError::PairingExpired)));
417 }
418
419 #[test]
420 fn refreshing_an_active_session_extends_only_its_inactivity_deadline() {
421 let manager = PairingManager::new(["https://example.test"], Duration::from_millis(400)).expect("manager");
422 let display = manager.begin_pairing();
423 let session = manager.pair(display.code(), "https://example.test").expect("pair");
424
425 std::thread::sleep(Duration::from_millis(150));
426 manager.refresh(session.token(), "https://example.test").expect("refresh");
427 std::thread::sleep(Duration::from_millis(150));
428 assert!(manager.validate(session.token(), "https://example.test").is_ok());
429
430 std::thread::sleep(Duration::from_millis(450));
431 assert!(matches!(manager.validate(session.token(), "https://example.test"), Err(WebmcpError::Unauthorized)));
432 }
433}