use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use anyhow::Result;
use axum::{
body::Body,
extract::{
connect_info::ConnectInfo,
ws::{Message, WebSocket, WebSocketUpgrade},
Path, State,
},
http::{HeaderMap, StatusCode},
response::{IntoResponse, Response},
routing::get,
Router,
};
use dashmap::DashMap;
use futures_util::{SinkExt, StreamExt};
use tokio::sync::{oneshot, RwLock};
use tracing::{info, warn};
use uuid::Uuid;
use super::auth::{verify_register, verifying_key_from_b64};
use super::device::validate_device_id;
use super::protocol::RelayMessage;
use super::ratelimit::RateLimiter;
const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const REGISTER_TIMEOUT: Duration = Duration::from_secs(15);
const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(60);
const HUB_OUTBOUND_CAPACITY: usize = 256;
const MAX_REGISTERED_HUBS: usize = 1024;
const WS_RATE_MAX: usize = 120;
const WS_RATE_WINDOW_SECS: u64 = 600;
const PAIR_RATE_MAX: usize = 30;
const PAIR_RATE_WINDOW_SECS: u64 = 60;
#[derive(Clone)]
pub struct RelayState {
hubs: Arc<RwLock<HashMap<String, HubHandle>>>,
pending: Arc<DashMap<String, oneshot::Sender<RelayMessage>>>,
device_keys: Arc<RwLock<HashMap<String, [u8; 32]>>>,
ws_rate_limiter: Arc<RateLimiter>,
pair_rate_limiter: Arc<RateLimiter>,
used_register_nonces: Arc<DashMap<(String, i64), Instant>>,
}
static NEXT_CONN_ID: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
#[derive(Clone)]
struct HubHandle {
tx: tokio::sync::mpsc::Sender<RelayMessage>,
conn_id: u64,
}
pub fn build_relay_router() -> Router {
let state = RelayState {
hubs: Arc::new(RwLock::new(HashMap::new())),
pending: Arc::new(DashMap::new()),
device_keys: Arc::new(RwLock::new(HashMap::new())),
ws_rate_limiter: Arc::new(RateLimiter::new(WS_RATE_MAX, WS_RATE_WINDOW_SECS)),
pair_rate_limiter: Arc::new(RateLimiter::new(PAIR_RATE_MAX, PAIR_RATE_WINDOW_SECS)),
used_register_nonces: Arc::new(DashMap::new()),
};
Router::new()
.route("/health", get(health))
.route("/ws/pairing", get(ws_pairing))
.route("/pair/{device_id}/{code}", get(pair_page))
.route(
"/pair/{device_id}/{code}/confirm",
axum::routing::post(pair_confirm),
)
.with_state(state)
}
async fn health() -> impl IntoResponse {
(
StatusCode::OK,
r#"{"status":"healthy","service":"ilink-relay"}"#,
)
}
async fn ws_pairing(
ws: WebSocketUpgrade,
State(state): State<RelayState>,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
) -> impl IntoResponse {
if !state.ws_rate_limiter.allow(&addr.ip().to_string()) {
warn!(ip = %addr.ip(), "ws pairing rate limited");
return StatusCode::TOO_MANY_REQUESTS.into_response();
}
ws.on_upgrade(move |socket| handle_hub_socket(state, socket))
}
async fn accept_registration(
state: &RelayState,
device_id: &str,
public_key: &str,
timestamp: i64,
signature: &str,
) -> Result<(), String> {
if !validate_device_id(device_id) {
return Err("invalid device_id".into());
}
let verifying_key = verifying_key_from_b64(public_key).map_err(|e| e.to_string())?;
verify_register(&verifying_key, device_id, timestamp, signature, unix_now())
.map_err(|e| e.to_string())?;
let nonce_key = (device_id.to_string(), timestamp);
let skew = Duration::from_secs(super::auth::REGISTER_MAX_SKEW_SECS as u64);
state
.used_register_nonces
.retain(|_, seen_at| seen_at.elapsed() < skew);
if state.used_register_nonces.contains_key(&nonce_key) {
return Err("registration nonce already used; replay detected".into());
}
state.used_register_nonces.insert(nonce_key, Instant::now());
let key_bytes = verifying_key.to_bytes();
let mut keys = state.device_keys.write().await;
if let Some(existing) = keys.get(device_id) {
if existing != &key_bytes {
return Err("device_id already bound to another key".into());
}
} else {
keys.insert(device_id.to_string(), key_bytes);
}
Ok(())
}
async fn handle_hub_socket(state: RelayState, socket: WebSocket) {
let conn_id = NEXT_CONN_ID.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let (mut write, mut read) = socket.split();
let (out_tx, mut out_rx) = tokio::sync::mpsc::channel::<RelayMessage>(HUB_OUTBOUND_CAPACITY);
let mut device_id: Option<String> = None;
let register_deadline = tokio::time::sleep(REGISTER_TIMEOUT);
tokio::pin!(register_deadline);
let mut keepalive = tokio::time::interval(KEEPALIVE_INTERVAL);
keepalive.tick().await;
loop {
tokio::select! {
_ = &mut register_deadline, if device_id.is_none() => {
warn!("hub registration timeout");
break;
}
_ = keepalive.tick(), if device_id.is_some() => {
if write.send(Message::Ping(vec![].into())).await.is_err() {
break;
}
}
incoming = read.next() => {
let Some(msg) = incoming else { break };
let Ok(msg) = msg else { break };
let Message::Text(text) = msg else { continue };
let Ok(parsed) = RelayMessage::from_json(text.as_ref()) else { continue };
match parsed {
RelayMessage::Register {
device_id: id,
public_key,
timestamp,
signature,
} if device_id.is_none() => {
{
let hubs = state.hubs.read().await;
if !hubs.contains_key(&id) && hubs.len() >= MAX_REGISTERED_HUBS {
let reason = format!(
"relay at capacity ({MAX_REGISTERED_HUBS} Hubs); \
retry later"
);
warn!(device_id = %id, "hub registration rejected: cap reached");
let err = RelayMessage::Registered {
ok: false,
error: Some(reason),
};
let _ = write
.send(Message::Text(err.to_json().unwrap_or_default().into()))
.await;
break;
}
}
match accept_registration(&state, &id, &public_key, timestamp, &signature).await {
Ok(()) => {
state.hubs.write().await.insert(
id.clone(),
HubHandle { tx: out_tx.clone(), conn_id },
);
device_id = Some(id.clone());
info!(device_id = %id, "hub registered");
let ok = RelayMessage::Registered { ok: true, error: None };
let _ = write.send(Message::Text(ok.to_json().unwrap_or_default().into())).await;
}
Err(reason) => {
warn!(device_id = %id, reason = %reason, "hub registration rejected");
let err = RelayMessage::Registered {
ok: false,
error: Some(reason),
};
let _ = write.send(Message::Text(err.to_json().unwrap_or_default().into())).await;
break;
}
}
}
RelayMessage::Response { id, status, headers, body } => {
let tx = state.pending.remove(&id).map(|(_, tx)| tx);
if let Some(tx) = tx {
let _ = tx.send(RelayMessage::Response { id, status, headers, body });
}
}
_ => {}
}
}
Some(outgoing) = out_rx.recv() => {
if let Ok(json) = outgoing.to_json() {
if write.send(Message::Text(json.into())).await.is_err() {
break;
}
}
}
}
}
if let Some(id) = device_id {
let mut hubs = state.hubs.write().await;
if hubs.get(&id).is_some_and(|h| h.conn_id == conn_id) {
hubs.remove(&id);
}
info!(device_id = %id, conn_id, "hub disconnected");
}
}
struct PendingRequestGuard {
pending: Arc<DashMap<String, oneshot::Sender<RelayMessage>>>,
request_id: String,
}
impl Drop for PendingRequestGuard {
fn drop(&mut self) {
self.pending.remove(&self.request_id);
}
}
async fn forward_to_hub(
state: &RelayState,
device_id: &str,
method: &str,
path: &str,
headers: HeaderMap,
body: Option<String>,
) -> Result<RelayMessage, StatusCode> {
let hub = {
let hubs = state.hubs.read().await;
hubs.get(device_id).cloned()
};
let hub = hub.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
let request_id = Uuid::new_v4().to_string();
let (tx, rx) = oneshot::channel();
state.pending.insert(request_id.clone(), tx);
let _guard = PendingRequestGuard {
pending: state.pending.clone(),
request_id: request_id.clone(),
};
let mut hdr_map = HashMap::new();
for (k, v) in headers.iter() {
if let Ok(s) = v.to_str() {
hdr_map.insert(k.to_string(), s.to_string());
}
}
let req = RelayMessage::Request {
id: request_id.clone(),
method: method.to_string(),
path: path.to_string(),
headers: hdr_map,
body,
};
if hub.tx.try_send(req).is_err() {
return Err(StatusCode::SERVICE_UNAVAILABLE);
}
match tokio::time::timeout(REQUEST_TIMEOUT, rx).await {
Ok(Ok(resp)) => Ok(resp),
_ => Err(StatusCode::GATEWAY_TIMEOUT),
}
}
fn relay_response_to_http(resp: RelayMessage) -> Response {
let RelayMessage::Response {
status,
headers,
body,
..
} = resp
else {
return (StatusCode::BAD_GATEWAY, "invalid relay response").into_response();
};
let mut builder = Response::builder().status(status);
for (k, v) in headers {
if k.eq_ignore_ascii_case("transfer-encoding") {
continue;
}
builder = builder.header(k, v);
}
let body = body.unwrap_or_default();
builder.body(Body::from(body)).unwrap_or_else(|_| {
(StatusCode::INTERNAL_SERVER_ERROR, "response build error").into_response()
})
}
fn check_pair_rate(state: &RelayState, addr: &SocketAddr) -> Result<(), StatusCode> {
if state.pair_rate_limiter.allow(&addr.ip().to_string()) {
Ok(())
} else {
warn!(ip = %addr.ip(), "pairing HTTP rate limited");
Err(StatusCode::TOO_MANY_REQUESTS)
}
}
async fn pair_page(
State(state): State<RelayState>,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
Path((device_id, code)): Path<(String, String)>,
) -> Response {
if let Err(status) = check_pair_rate(&state, &addr) {
return (status, "too many pairing requests").into_response();
}
match forward_to_hub(
&state,
&device_id,
"GET",
&format!("/hub/pair/{code}"),
HeaderMap::new(),
None,
)
.await
{
Ok(resp) => relay_response_to_http(resp),
Err(status) => (
status,
format!("Hub 未在线,请确认本机 ilink-hub serve 正在运行 (device: {device_id})"),
)
.into_response(),
}
}
async fn pair_confirm(
State(state): State<RelayState>,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
Path((device_id, code)): Path<(String, String)>,
mut headers: HeaderMap,
body: String,
) -> Response {
if let Err(status) = check_pair_rate(&state, &addr) {
return (status, "too many pairing requests").into_response();
}
if let Ok(v) = addr.ip().to_string().parse() {
headers.insert("x-forwarded-for", v);
}
match forward_to_hub(
&state,
&device_id,
"POST",
&format!("/hub/pair/{code}/confirm"),
headers,
Some(body),
)
.await
{
Ok(resp) => relay_response_to_http(resp),
Err(status) => (
status,
format!("Hub 未在线,请确认本机 ilink-hub serve 正在运行 (device: {device_id})"),
)
.into_response(),
}
}
pub async fn serve(addr: &str) -> Result<()> {
let router = build_relay_router();
let listener = tokio::net::TcpListener::bind(addr).await?;
info!(%addr, "ilink-relay listening");
axum::serve(
listener,
router.into_make_service_with_connect_info::<SocketAddr>(),
)
.with_graceful_shutdown(shutdown_signal())
.await?;
Ok(())
}
async fn shutdown_signal() {
let ctrl_c = async {
tokio::signal::ctrl_c()
.await
.expect("failed to install Ctrl+C handler");
};
#[cfg(unix)]
let terminate = async {
tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
.expect("failed to install SIGTERM handler")
.recv()
.await;
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => {},
_ = terminate => {},
}
}
fn unix_now() -> i64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or(Duration::ZERO)
.as_secs() as i64
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_relay_pending_request_no_leak_on_cancel() {
let state = RelayState {
hubs: Arc::new(RwLock::new(HashMap::new())),
pending: Arc::new(DashMap::new()),
device_keys: Arc::new(RwLock::new(HashMap::new())),
ws_rate_limiter: Arc::new(RateLimiter::new(10, 60)),
pair_rate_limiter: Arc::new(RateLimiter::new(10, 60)),
used_register_nonces: Arc::new(DashMap::new()),
};
let device_id = "test-device".to_string();
let (ws_tx, mut ws_rx) = tokio::sync::mpsc::channel(HUB_OUTBOUND_CAPACITY);
state.hubs.write().await.insert(
device_id.clone(),
HubHandle {
tx: ws_tx,
conn_id: 0,
},
);
let state_clone = state.clone();
let handle = tokio::spawn(async move {
let _ = forward_to_hub(
&state_clone,
"test-device",
"GET",
"/test-path",
HeaderMap::new(),
None,
)
.await;
});
let msg = tokio::time::timeout(Duration::from_millis(500), ws_rx.recv())
.await
.expect("should receive request on ws channel")
.expect("msg is some");
let request_id = match msg {
RelayMessage::Request { id, .. } => id,
_ => panic!("expected RelayMessage::Request"),
};
{
assert!(
state.pending.contains_key(&request_id),
"request must be in pending map"
);
}
handle.abort();
let _ = handle.await;
tokio::time::sleep(Duration::from_millis(50)).await;
{
assert!(
!state.pending.contains_key(&request_id),
"pending request must be cleaned up on cancel!"
);
}
}
#[tokio::test]
async fn test_relay_disconnect_does_not_erase_device_key() {
let state = RelayState {
hubs: Arc::new(RwLock::new(HashMap::new())),
pending: Arc::new(DashMap::new()),
device_keys: Arc::new(RwLock::new(HashMap::new())),
ws_rate_limiter: Arc::new(RateLimiter::new(10, 60)),
pair_rate_limiter: Arc::new(RateLimiter::new(10, 60)),
used_register_nonces: Arc::new(DashMap::new()),
};
let device_id = "device_123".to_string();
let pub_key = [7u8; 32];
state
.device_keys
.write()
.await
.insert(device_id.clone(), pub_key);
let (ws_tx, _ws_rx) = tokio::sync::mpsc::channel(HUB_OUTBOUND_CAPACITY);
state.hubs.write().await.insert(
device_id.clone(),
HubHandle {
tx: ws_tx,
conn_id: 1,
},
);
assert!(state.device_keys.read().await.contains_key(&device_id));
if let Some(id) = Some(device_id.clone()) {
let conn_id = 1u64;
let mut hubs = state.hubs.write().await;
if hubs.get(&id).is_some_and(|h| h.conn_id == conn_id) {
hubs.remove(&id);
}
}
assert!(
state.device_keys.read().await.contains_key(&device_id),
"device key must be preserved after hub disconnects (SEC-M4-001)"
);
}
#[tokio::test]
async fn test_relay_disconnect_prevents_hijacking() {
let state = RelayState {
hubs: Arc::new(RwLock::new(HashMap::new())),
pending: Arc::new(DashMap::new()),
device_keys: Arc::new(RwLock::new(HashMap::new())),
ws_rate_limiter: Arc::new(RateLimiter::new(10, 60)),
pair_rate_limiter: Arc::new(RateLimiter::new(10, 60)),
used_register_nonces: Arc::new(DashMap::new()),
};
let device_id = "device_123".to_string();
let pub_key_a = [7u8; 32];
let pub_key_b = [8u8; 32];
state
.device_keys
.write()
.await
.insert(device_id.clone(), pub_key_a);
let (ws_tx, _ws_rx) = tokio::sync::mpsc::channel(HUB_OUTBOUND_CAPACITY);
state.hubs.write().await.insert(
device_id.clone(),
HubHandle {
tx: ws_tx,
conn_id: 1,
},
);
assert!(state.device_keys.read().await.contains_key(&device_id));
state.hubs.write().await.remove(&device_id);
let accept_hijack = {
let keys = state.device_keys.read().await;
if let Some(existing) = keys.get(&device_id) {
if existing != &pub_key_b {
Err("device_id already bound to another key".to_string())
} else {
Ok(())
}
} else {
Ok(())
}
};
assert!(accept_hijack.is_err());
assert_eq!(
accept_hijack.unwrap_err(),
"device_id already bound to another key"
);
let accept_legit = {
let keys = state.device_keys.read().await;
if let Some(existing) = keys.get(&device_id) {
if existing != &pub_key_a {
Err("device_id already bound to another key".to_string())
} else {
Ok(())
}
} else {
Ok(())
}
};
assert!(accept_legit.is_ok());
}
#[tokio::test]
async fn test_relay_stale_cleanup_does_not_evict_new_connection() {
let state = RelayState {
hubs: Arc::new(RwLock::new(HashMap::new())),
pending: Arc::new(DashMap::new()),
device_keys: Arc::new(RwLock::new(HashMap::new())),
ws_rate_limiter: Arc::new(RateLimiter::new(10, 60)),
pair_rate_limiter: Arc::new(RateLimiter::new(10, 60)),
used_register_nonces: Arc::new(DashMap::new()),
};
let device_id = "device_reconnect".to_string();
let (old_tx, _old_rx) = tokio::sync::mpsc::channel(HUB_OUTBOUND_CAPACITY);
state.hubs.write().await.insert(
device_id.clone(),
HubHandle {
tx: old_tx,
conn_id: 1,
},
);
let (new_tx, _new_rx) = tokio::sync::mpsc::channel(HUB_OUTBOUND_CAPACITY);
state.hubs.write().await.insert(
device_id.clone(),
HubHandle {
tx: new_tx,
conn_id: 2,
},
);
let stale_conn_id = 1u64;
{
let mut hubs = state.hubs.write().await;
if hubs
.get(&device_id)
.is_some_and(|h| h.conn_id == stale_conn_id)
{
hubs.remove(&device_id);
}
}
assert!(
state.hubs.read().await.contains_key(&device_id),
"new connection must remain registered after stale cleanup"
);
assert_eq!(
state.hubs.read().await.get(&device_id).map(|h| h.conn_id),
Some(2),
"conn_id of remaining entry must be the new connection's"
);
}
#[test]
fn relay_response_to_http_non_response_returns_bad_gateway() {
let resp = relay_response_to_http(RelayMessage::Ping);
assert_eq!(resp.status(), StatusCode::BAD_GATEWAY);
}
#[tokio::test]
async fn relay_response_to_http_maps_status_and_body() {
let msg = RelayMessage::Response {
id: "r1".into(),
status: 200,
headers: HashMap::new(),
body: Some("hello".into()),
};
let resp = relay_response_to_http(msg);
assert_eq!(resp.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(&bytes[..], b"hello");
}
#[tokio::test]
async fn relay_response_to_http_strips_transfer_encoding_header() {
let mut headers = HashMap::new();
headers.insert("transfer-encoding".to_string(), "chunked".to_string());
headers.insert("x-custom".to_string(), "value".to_string());
let msg = RelayMessage::Response {
id: "r1".into(),
status: 200,
headers,
body: None,
};
let resp = relay_response_to_http(msg);
assert!(
resp.headers().get("transfer-encoding").is_none(),
"transfer-encoding must be stripped"
);
assert!(
resp.headers().get("x-custom").is_some(),
"x-custom must pass through"
);
}
#[tokio::test]
async fn relay_response_to_http_empty_body_returns_empty_bytes() {
let msg = RelayMessage::Response {
id: "r1".into(),
status: 204,
headers: HashMap::new(),
body: None,
};
let resp = relay_response_to_http(msg);
assert_eq!(resp.status(), 204);
let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX)
.await
.unwrap();
assert!(bytes.is_empty(), "None body must produce empty bytes");
}
#[test]
fn unix_now_returns_plausible_timestamp() {
let ts = unix_now();
let year_2020: i64 = 1_577_836_800;
let year_2100: i64 = 4_102_444_800;
assert!(
ts > year_2020 && ts < year_2100,
"unix_now() = {ts} outside plausible range"
);
}
}