use std::collections::HashMap;
use std::time::{Duration, Instant};
use uuid::Uuid;
const PAIRING_TTL: Duration = Duration::from_secs(600);
const SCANNED_TTL: Duration = Duration::from_secs(60);
const CONFIRMED_TTL: Duration = Duration::from_secs(86_400);
const VTOKEN_CLAIM_WINDOW: Duration = Duration::from_secs(120);
pub const MAX_PAIRING_SESSIONS: usize = 1024;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PairingStatus {
Wait,
Scanned,
Confirmed,
Expired,
}
#[derive(Debug, Clone)]
pub struct PairingSession {
pub code: String,
pub created_at: Instant,
pub scanned_at: Option<Instant>,
pub status: PairingStatus,
pub vtoken: Option<String>,
pub client_name: Option<String>,
pub client_label: Option<String>,
pub csrf: Option<String>,
}
impl PairingSession {
fn is_expired(&self) -> bool {
match self.status {
PairingStatus::Confirmed => false,
PairingStatus::Scanned => self
.scanned_at
.map(|t| t.elapsed() > SCANNED_TTL)
.unwrap_or_else(|| self.created_at.elapsed() > PAIRING_TTL),
_ => self.created_at.elapsed() > PAIRING_TTL,
}
}
fn should_evict(&self) -> bool {
match self.status {
PairingStatus::Confirmed => self.created_at.elapsed() > CONFIRMED_TTL,
PairingStatus::Scanned => self
.scanned_at
.map(|t| t.elapsed() > SCANNED_TTL)
.unwrap_or_else(|| self.created_at.elapsed() > PAIRING_TTL),
_ => self.created_at.elapsed() > PAIRING_TTL,
}
}
pub fn public_status(&self) -> PairingStatus {
if self.is_expired() {
PairingStatus::Expired
} else {
self.status.clone()
}
}
pub fn status_str(&self) -> &'static str {
match self.public_status() {
PairingStatus::Wait => "wait",
PairingStatus::Scanned => "scaned",
PairingStatus::Confirmed => "confirmed",
PairingStatus::Expired => "expired",
}
}
}
#[derive(Debug, Default)]
pub struct PairingRegistry {
sessions: HashMap<String, PairingSession>,
confirmed_sessions: HashMap<String, (PairingSession, Instant)>,
}
impl PairingRegistry {
pub fn new() -> Self {
Self::default()
}
fn purge_expired(&mut self) {
self.sessions.retain(|_, s| !s.should_evict());
self.confirmed_sessions
.retain(|_, (_, confirmed_at)| confirmed_at.elapsed() < CONFIRMED_TTL);
}
pub fn create(&mut self) -> Result<String, PairingError> {
self.purge_expired();
if self.sessions.len() + self.confirmed_sessions.len() >= MAX_PAIRING_SESSIONS {
return Err(PairingError::TooManySessions);
}
let code = format!("pair_{}", Uuid::new_v4().simple());
self.sessions.insert(
code.clone(),
PairingSession {
code: code.clone(),
created_at: Instant::now(),
scanned_at: None,
status: PairingStatus::Wait,
vtoken: None,
client_name: None,
client_label: None,
csrf: None,
},
);
Ok(code)
}
pub fn get(&self, code: &str) -> Option<PairingSession> {
if let Some(session) = self.sessions.get(code) {
return Some(session.clone());
}
if let Some((session, _)) = self.confirmed_sessions.get(code) {
return Some(session.clone());
}
None
}
pub fn claim_confirmed_vtoken(
&mut self,
code: &str,
) -> Option<(PairingSession, Option<String>)> {
self.purge_expired();
if let Some((session, confirmed_at)) = self.confirmed_sessions.get_mut(code) {
if session.vtoken.is_none() {
return Some((session.clone(), None));
}
if confirmed_at.elapsed() < VTOKEN_CLAIM_WINDOW {
let token = session.vtoken.clone();
return Some((session.clone(), token));
}
let _ = session.vtoken.take();
return Some((session.clone(), None));
}
self.sessions
.get(code)
.map(|session| (session.clone(), None))
}
pub fn mark_scanned(&mut self, code: &str) -> bool {
self.purge_expired();
if let Some(session) = self.sessions.get_mut(code) {
if session.is_expired() {
session.status = PairingStatus::Expired;
return false;
}
if session.status == PairingStatus::Wait {
session.status = PairingStatus::Scanned;
session.scanned_at = Some(Instant::now());
}
if session.csrf.is_none() {
session.csrf = Some(generate_csrf());
}
return true;
}
false
}
pub fn pre_check_confirm(&mut self, code: &str, csrf_header: &str) -> Result<(), PairingError> {
self.purge_expired();
if self.confirmed_sessions.contains_key(code) {
return Err(PairingError::AlreadyConfirmed);
}
let session = self.sessions.get(code).ok_or(PairingError::NotFound)?;
if session.is_expired() {
return Err(PairingError::Expired);
}
if session.status == PairingStatus::Confirmed {
return Err(PairingError::AlreadyConfirmed);
}
match session.csrf.as_deref() {
Some(token) if constant_time_eq(token.as_bytes(), csrf_header.as_bytes()) => {}
_ => return Err(PairingError::CsrfMismatch),
}
if session.status != PairingStatus::Scanned {
return Err(PairingError::NotScanned);
}
Ok(())
}
pub fn confirm(
&mut self,
code: &str,
client_name: String,
client_label: Option<String>,
vtoken: String,
csrf_header: &str,
) -> Result<(), PairingError> {
self.purge_expired();
if self.confirmed_sessions.contains_key(code) {
return Err(PairingError::AlreadyConfirmed);
}
let mut session = self.sessions.remove(code).ok_or(PairingError::NotFound)?;
if session.is_expired() {
return Err(PairingError::Expired);
}
if session.status == PairingStatus::Confirmed {
return Err(PairingError::AlreadyConfirmed);
}
match session.csrf.as_deref() {
Some(token) if constant_time_eq(token.as_bytes(), csrf_header.as_bytes()) => {
session.csrf = None;
}
_ => return Err(PairingError::CsrfMismatch),
}
if session.status != PairingStatus::Scanned {
return Err(PairingError::NotScanned);
}
session.status = PairingStatus::Confirmed;
session.vtoken = Some(vtoken);
session.client_name = Some(client_name);
session.client_label = client_label;
self.confirmed_sessions
.insert(code.to_string(), (session, Instant::now()));
Ok(())
}
pub fn remove_confirmed(&mut self, code: &str) {
self.confirmed_sessions.remove(code);
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum PairingError {
NotFound,
Expired,
AlreadyConfirmed,
NotScanned,
CsrfMismatch,
TooManySessions,
NameCollision,
}
fn generate_csrf() -> String {
use rand::RngCore;
let mut bytes = [0u8; 16];
rand::rng().fill_bytes(&mut bytes);
bytes.iter().map(|b| format!("{b:02x}")).collect()
}
fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
let mut diff: u8 = 0;
for (x, y) in a.iter().zip(b.iter()) {
diff |= x ^ y;
}
diff == 0
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn create_and_confirm_pairing() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let csrf = reg.get(&code).unwrap().csrf.clone().unwrap();
reg.confirm(
&code,
"openclaw-test".to_string(),
Some("Test".to_string()),
"vhub_abc".to_string(),
&csrf,
)
.unwrap();
let session = reg.get(&code).unwrap();
assert_eq!(session.status_str(), "confirmed");
assert_eq!(session.vtoken.as_deref(), Some("vhub_abc"));
assert!(session.csrf.is_none(), "csrf must be consumed on confirm");
let (_snap, first) = reg.claim_confirmed_vtoken(&code).unwrap();
assert_eq!(first.as_deref(), Some("vhub_abc"));
let (_snap, second) = reg.claim_confirmed_vtoken(&code).unwrap();
assert_eq!(second.as_deref(), Some("vhub_abc"));
assert_eq!(reg.get(&code).unwrap().status_str(), "confirmed");
}
#[test]
fn expired_pairing_rejected() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
let session = reg.sessions.get_mut(&code).unwrap();
session.created_at = Instant::now() - Duration::from_secs(700);
let csrf = "0".repeat(32);
assert_eq!(reg.get(&code).unwrap().status_str(), "expired");
assert!(reg
.confirm(&code, "x".into(), None, "vhub_x".into(), &csrf,)
.is_err());
}
#[test]
fn confirm_rejected_when_status_is_wait() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
let err = reg
.confirm(
&code,
"x".into(),
None,
"vhub_x".into(),
"0".repeat(32).as_str(),
)
.unwrap_err();
assert_eq!(err, PairingError::CsrfMismatch);
}
#[test]
fn confirm_after_concurrent_attempt_returns_only_one_winner() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let csrf = reg.get(&code).unwrap().csrf.clone().unwrap();
reg.confirm(&code, "first".into(), None, "vhub_1".into(), &csrf)
.unwrap();
let err = reg
.confirm(&code, "second".into(), None, "vhub_2".into(), &csrf)
.unwrap_err();
assert_eq!(err, PairingError::AlreadyConfirmed);
}
#[test]
fn csrf_token_consumed_after_confirm() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let csrf = reg.get(&code).unwrap().csrf.clone().unwrap();
reg.confirm(&code, "client".into(), None, "vhub_x".into(), &csrf)
.unwrap();
let err = reg
.confirm(&code, "attacker".into(), None, "vhub_y".into(), &csrf)
.unwrap_err();
assert_eq!(err, PairingError::AlreadyConfirmed);
}
#[test]
fn scanned_session_expires_after_scanned_ttl_not_pairing_ttl() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let session = reg.sessions.get_mut(&code).unwrap();
session.scanned_at = Some(Instant::now() - Duration::from_secs(SCANNED_TTL.as_secs() + 5));
assert_eq!(
reg.get(&code).unwrap().status_str(),
"expired",
"scanned session must expire after SCANNED_TTL, not PAIRING_TTL"
);
}
#[test]
fn generate_csrf_is_unique_and_hex() {
let a = generate_csrf();
let b = generate_csrf();
assert_eq!(a.len(), 32, "csrf must be 32 hex chars");
assert!(a.chars().all(|c| c.is_ascii_hexdigit()));
assert_ne!(a, b, "two consecutive csrf tokens must differ");
}
#[test]
fn too_many_sessions_returns_error() {
let mut reg = PairingRegistry::new();
for _ in 0..MAX_PAIRING_SESSIONS {
reg.create().unwrap();
}
let err = reg.create().unwrap_err();
assert_eq!(err, PairingError::TooManySessions);
}
#[test]
fn confirmed_sessions_are_evicted_after_confirmed_ttl() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let csrf = reg.get(&code).unwrap().csrf.clone().unwrap();
reg.confirm(&code, "client".into(), None, "vhub_x".into(), &csrf)
.unwrap();
assert_eq!(reg.sessions.len(), 0);
assert_eq!(
reg.get(&code).unwrap().status_str(),
"confirmed",
"freshly confirmed session must be visible via confirmed_sessions"
);
reg.confirmed_sessions.get_mut(&code).unwrap().1 =
Instant::now() - Duration::from_secs(86_400 + 60);
reg.purge_expired();
assert!(
reg.get(&code).is_none(),
"Confirmed session must be evicted after CONFIRMED_TTL"
);
let code2 = reg.create().unwrap();
assert!(
reg.get(&code2).is_some(),
"create must succeed once the immortal Confirmed entry is evicted"
);
}
#[test]
fn pre_check_confirm_rejects_wrong_csrf() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let err = reg
.pre_check_confirm(&code, "definitely-wrong-csrf")
.unwrap_err();
assert_eq!(
err,
PairingError::CsrfMismatch,
"wrong CSRF must be rejected with CsrfMismatch"
);
}
#[test]
fn pre_check_confirm_accepts_correct_csrf() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let csrf = reg.get(&code).unwrap().csrf.clone().unwrap();
reg.pre_check_confirm(&code, &csrf)
.expect("correct CSRF must be accepted");
}
#[test]
fn confirmed_session_is_not_expired_after_pairing_ttl() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let csrf = reg.get(&code).unwrap().csrf.clone().unwrap();
reg.confirm(&code, "client".into(), None, "vhub_ok".into(), &csrf)
.unwrap();
reg.confirmed_sessions.get_mut(&code).unwrap().0.created_at =
Instant::now() - Duration::from_secs(PAIRING_TTL.as_secs() + 5);
assert_eq!(
reg.get(&code).unwrap().status_str(),
"confirmed",
"Confirmed session must not appear expired after PAIRING_TTL"
);
}
#[test]
fn purge_expired_removes_stale_wait_sessions() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.sessions.get_mut(&code).unwrap().created_at =
Instant::now() - Duration::from_secs(PAIRING_TTL.as_secs() + 5);
reg.purge_expired();
assert!(
reg.get(&code).is_none(),
"stale Wait session must be removed by purge_expired"
);
}
#[test]
fn purge_expired_removes_stale_scanned_sessions() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let session = reg.sessions.get_mut(&code).unwrap();
session.scanned_at = Some(Instant::now() - Duration::from_secs(SCANNED_TTL.as_secs() + 5));
reg.purge_expired();
assert!(
reg.get(&code).is_none(),
"stale Scanned session must be removed by purge_expired"
);
}
#[test]
fn session_cap_includes_confirmed_sessions() {
let mut reg = PairingRegistry::new();
for _ in 0..(MAX_PAIRING_SESSIONS - 1) {
reg.create().unwrap();
}
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let csrf = reg.get(&code).unwrap().csrf.clone().unwrap();
reg.confirm(&code, "c".into(), None, "vhub_cap".into(), &csrf)
.unwrap();
let err = reg.create().unwrap_err();
assert_eq!(
err,
PairingError::TooManySessions,
"cap must account for both pending and confirmed sessions"
);
}
#[test]
fn claim_confirmed_vtoken_reclaimable_within_window() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let csrf = reg.get(&code).unwrap().csrf.clone().unwrap();
reg.confirm(&code, "client".into(), None, "vhub_once".into(), &csrf)
.unwrap();
let (session1, token1) = reg.claim_confirmed_vtoken(&code).unwrap();
assert_eq!(session1.status_str(), "confirmed");
assert_eq!(token1.as_deref(), Some("vhub_once"));
let (session2, token2) = reg.claim_confirmed_vtoken(&code).unwrap();
assert_eq!(
session2.status_str(),
"confirmed",
"confirmed stub must remain after claim"
);
assert_eq!(
token2.as_deref(),
Some("vhub_once"),
"second claim within window must re-issue the same vtoken"
);
assert_eq!(
reg.get(&code).unwrap().vtoken.as_deref(),
Some("vhub_once"),
"stored session must keep vtoken during claim window"
);
}
#[test]
fn claim_confirmed_vtoken_clears_after_claim_window() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
reg.mark_scanned(&code);
let csrf = reg.get(&code).unwrap().csrf.clone().unwrap();
reg.confirm(&code, "client".into(), None, "vhub_window".into(), &csrf)
.unwrap();
let (_, token_fresh) = reg.claim_confirmed_vtoken(&code).unwrap();
assert_eq!(token_fresh.as_deref(), Some("vhub_window"));
reg.confirmed_sessions.get_mut(&code).unwrap().1 =
Instant::now() - VTOKEN_CLAIM_WINDOW - Duration::from_secs(1);
let (session, token) = reg.claim_confirmed_vtoken(&code).unwrap();
assert_eq!(session.status_str(), "confirmed");
assert!(
token.is_none(),
"claim after window must not return the vtoken"
);
assert!(
reg.get(&code).unwrap().vtoken.is_none(),
"stored session must keep vtoken cleared after window"
);
let (_, token_again) = reg.claim_confirmed_vtoken(&code).unwrap();
assert!(
token_again.is_none(),
"subsequent claims must keep seeing None after clear"
);
}
#[test]
fn claim_confirmed_vtoken_on_wait_returns_none_token() {
let mut reg = PairingRegistry::new();
let code = reg.create().unwrap();
let (session, token) = reg.claim_confirmed_vtoken(&code).unwrap();
assert_eq!(session.status_str(), "wait");
assert!(token.is_none());
reg.mark_scanned(&code);
let (session, token) = reg.claim_confirmed_vtoken(&code).unwrap();
assert_eq!(session.status_str(), "scaned");
assert!(token.is_none());
}
#[test]
fn claim_confirmed_vtoken_unknown_code_returns_none() {
let mut reg = PairingRegistry::new();
assert!(reg.claim_confirmed_vtoken("pair_missing").is_none());
}
}