Skip to main content

vtcode_webmcp/
pairing.rs

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/// A terminal-displayable, short-lived pairing code.
48#[derive(Clone)]
49pub struct PairingDisplay {
50    code: String,
51    expires_at: Instant,
52}
53
54impl PairingDisplay {
55    /// Returns the code that must be entered by the browser.
56    pub fn code(&self) -> &str {
57        &self.code
58    }
59
60    /// Returns the remaining lifetime of the code.
61    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/// An authenticated, in-memory browser session.
76#[derive(Clone)]
77pub struct PairingSession {
78    token: String,
79    origin: String,
80    expires_at: Instant,
81}
82
83impl PairingSession {
84    /// Returns the bearer token. Callers should keep it in memory only.
85    pub fn token(&self) -> &str {
86        &self.token
87    }
88
89    /// Returns the origin bound to this session.
90    pub fn origin(&self) -> &str {
91        &self.origin
92    }
93
94    /// Returns the remaining inactivity-lease lifetime.
95    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/// Manages one-time pairing and revocable in-memory sessions.
111#[derive(Clone)]
112pub struct PairingManager {
113    allowed_origins: Arc<HashSet<String>>,
114    ttl: Duration,
115    state: Arc<Mutex<PairingState>>,
116}
117
118impl PairingManager {
119    /// Create a manager with an explicit origin allowlist and inactivity lease.
120    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    /// Starts a fresh pairing code, invalidating any older unconsumed code.
142    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    /// Starts a code bound to a specific allowed origin.
170    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    /// Revoke all browser sessions and issue a fresh code in one state update.
177    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    /// Consumes a code and creates a session token bound to the browser origin.
184    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    /// Validates a session token and its origin binding without extending it.
219    pub fn validate(&self, token: &str, origin: &str) -> Result<()> {
220        self.validate_session(token, origin, false)
221    }
222
223    /// Validates a session and extends its inactivity deadline.
224    ///
225    /// Pairing codes remain one-time and expire according to the configured
226    /// TTL. An authenticated browser session may remain connected longer than
227    /// that TTL as long as it continues making authenticated requests.
228    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    /// Rehydrates an existing in-memory session for a reconnecting socket.
251    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    /// Revokes one session token.
263    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    /// Revokes every active browser session and pending code.
269    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    /// Returns whether the origin is explicitly allowed.
276    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}