//! df-relay 中继服务核心实现 //! //! 设计依据:设计文档「Layer2」—— axum WS Server + 广播中继。 //! 接受两类连接:小程序(device_id 鉴权)+ 桌面端(token 配对),按 device_id //! 配对转发(非全局广播),纯转发无业务逻辑。 //! //! 协议(简单握手): //! 1. 客户端建立 WS 后,首条消息发 JSON `Hello { kind, device_id, token }` 宣告身份。 //! 2. relay 校验 token(MVP:env `DF_RELAY_TOKEN` 或硬编码常量;生产级鉴权留 Phase3)。 //! 3. 校验通过 → 注册连接、进入收发循环;失败 → 发 Error 帧 + Close。 //! 4. 后续消息按 kind 路由:Event(device→miniapp)/ Command(miniapp→device)/ Control。 //! //! AiChatEvent JSON 透传:relay 不解析 payload,只按 device_id + 方向转发。 use std::net::SocketAddr; use async_trait::async_trait; use axum::{ extract::{ ws::{Message, WebSocket, WebSocketUpgrade}, State, }, response::IntoResponse, routing::get, Router, }; use futures_util::{SinkExt, StreamExt}; use serde::{Deserialize, Serialize}; use tokio::sync::mpsc; use crate::broadcast::{BroadcastMessage, ClientKind, MessageKind}; use crate::conn::{next_conn_id, ConnHandle, ConnId, RelayState}; use crate::error::{RelayError, Result}; /// 读取期望 token(必需:env `DF_RELAY_TOKEN` 必须设置,未设置时 panic)。 /// 生产级鉴权(每 device 独立 token + 过期刷新)留 Phase3。 fn expected_token() -> String { std::env::var("DF_RELAY_TOKEN").unwrap_or_else(|_| { panic!("必须设置环境变量 DF_RELAY_TOKEN") }) } /// 客户端首消息:身份宣告(简单协议) #[derive(Debug, Clone, Serialize, Deserialize)] pub struct Hello { /// 客户端类型("device" / "miniapp") pub kind: ClientKindWire, /// 配对绑定的设备 ID pub device_id: String, /// 配对 token pub token: String, } /// Hello.kind 的传输表示(serde 字符串,与 ClientKind 解耦避免 rename 歧义) #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "snake_case")] pub enum ClientKindWire { Device, Miniapp, } impl From for ClientKind { fn from(w: ClientKindWire) -> Self { match w { ClientKindWire::Device => ClientKind::Device, ClientKindWire::Miniapp => ClientKind::Miniapp, } } } /// 中继服务抽象 /// /// 设计为 trait 便于测试 mock + 未来替换实现(如换 tonic gRPC 网关)。 #[async_trait] pub trait RelayServer: Send + Sync { /// 启动 HTTP/WS 服务监听指定地址 async fn start(&self, addr: &str) -> Result<()>; /// 广播消息给指定 device_id 绑定的对端 async fn broadcast(&self, msg: BroadcastMessage) -> Result<()>; /// 查询 device 是否有在线连接(离线降级判断用) fn is_device_online(&self, device_id: &str) -> bool; } /// 默认中继服务(持有共享 RelayState) pub struct DefaultRelayServer { state: RelayState, } impl DefaultRelayServer { pub fn new() -> Self { Self { state: RelayState::new(), } } /// 从既有 RelayState 构造(测试 / 外部复用) pub fn with_state(state: RelayState) -> Self { Self { state } } /// 暴露共享状态(外部可读连接数等) pub fn state(&self) -> RelayState { self.state.clone() } } impl Default for DefaultRelayServer { fn default() -> Self { Self::new() } } #[async_trait] impl RelayServer for DefaultRelayServer { async fn start(&self, addr: &str) -> Result<()> { let socket_addr: SocketAddr = addr .parse() .map_err(|e| RelayError::Start(format!("地址解析失败 {addr}: {e}")))?; let app = build_router(self.state.clone()); let listener = tokio::net::TcpListener::bind(&socket_addr) .await .map_err(|e| RelayError::Start(format!("监听绑定失败 {addr}: {e}")))?; tracing::info!(%addr, "df-relay WS 服务已启动"); axum::serve(listener, app) .await .map_err(|e| RelayError::Start(format!("axum::serve 失败: {e}")))?; Ok(()) } async fn broadcast(&self, msg: BroadcastMessage) -> Result<()> { let delivered = self.state.route(&msg).await; if delivered == 0 { // 对端离线不算硬错误(MVP 返回 Ok,离线降级由调用方据 is_device_online 判断) tracing::debug!( device_id = %msg.device_id, kind = ?msg.kind, "广播无对端在线(消息丢弃)" ); } Ok(()) } fn is_device_online(&self, device_id: &str) -> bool { // trait 同步签名:tokio Mutex 用 try_lock 快照,失败保守返回 false match self.state.inner().try_lock() { Ok(g) => g.is_device_online(device_id), Err(_) => false, } } } /// axum WS 路由构造 /// /// 暴露 `/ws/device`(桌面端连入)与 `/ws/miniapp`(小程序连入)两个端点, /// 共享 RelayState。端点仅决定「期望的客户端类型」,真正的身份宣告在首消息 /// Hello 中再次校验(防误连/误用)。 pub fn build_router(state: RelayState) -> Router { Router::new() .route("/ws/device", get(device_ws_handler)) .route("/ws/miniapp", get(miniapp_ws_handler)) .with_state(state) } /// 桌面端 WS upgrade handler async fn device_ws_handler( ws: WebSocketUpgrade, State(state): State, ) -> impl IntoResponse { tracing::debug!("桌面端 WS 连接接入"); ws.on_upgrade(move |socket| handle_connection(socket, state, ClientKindWire::Device)) } /// 小程序 WS upgrade handler async fn miniapp_ws_handler( ws: WebSocketUpgrade, State(state): State, ) -> impl IntoResponse { tracing::debug!("小程序 WS 连接接入"); ws.on_upgrade(move |socket| handle_connection(socket, state, ClientKindWire::Miniapp)) } /// WS 连接生命周期(握手 → 收发循环 → 注销) /// /// 步骤: /// 1. 等待首条 Hello 文本帧,校验 kind 与 token。 /// 2. 校验通过:分配 conn_id + mpsc,注册 ConnHandle,派发广播读取任务。 /// 3. 主循环:从 socket recv 文本帧 → 构造 BroadcastMessage → route 投递。 /// 4. 同时读取 mpsc 广播队列 → 写回 socket(双任务用 split sink/stream)。 /// 5. 任一端断开 → 注销连接、关闭 mpsc。 async fn handle_connection(socket: WebSocket, state: RelayState, expected: ClientKindWire) { // 握手阶段:等待首条 Hello let (mut socket_tx, mut socket_rx) = socket.split(); let hello = match recv_hello(&mut socket_rx).await { Ok(h) => h, Err(e) => { tracing::warn!(error = %e, "握手失败:未收到合法 Hello"); let _ = send_text( &mut socket_tx, r#"{"kind":"control","error":"handshake_failed"}"#, ) .await; let _ = socket_tx.close().await; return; } }; // 身份 + token 双因子校验 if hello.kind != expected { tracing::warn!( ?hello.kind, ?expected, "握手失败:客户端类型与端点不匹配" ); let _ = send_text( &mut socket_tx, r#"{"kind":"control","error":"kind_mismatch"}"#, ) .await; let _ = socket_tx.close().await; return; } if hello.token != expected_token() { tracing::warn!( device_id = %hello.device_id, "握手失败:token 校验不通过" ); let _ = send_text( &mut socket_tx, r#"{"kind":"control","error":"auth_failed"}"#, ) .await; let _ = socket_tx.close().await; return; } let conn_id = next_conn_id(); let kind: ClientKind = hello.kind.into(); let device_id = hello.device_id.clone(); tracing::info!( conn_id = conn_id.0, ?kind, device_id = %device_id, "连接握手通过,进入收发循环" ); // 握手通过:立即发 ack 控制帧给客户端。 // 客户端据此判定握手成功(首条非 error 消息即 handshaked),不依赖等待对端首条业务消息 —— // 否则单端连入时(device 离线)relay 静默,客户端永卡 handshaking,send 被 handshaked 守卫拦截。 let _ = send_text( &mut socket_tx, r#"{"kind":"control","payload":{"control_kind":"hello_ack"}}"#, ) .await; // 建立广播投递 mpsc(连接读取任务消费 → 写回 socket) let (bc_tx, bc_rx) = mpsc::unbounded_channel::(); let handle = ConnHandle::new(conn_id, kind, device_id.clone(), bc_tx); state.add_conn(handle).await; // 派发广播读取任务:从 bc_rx 取消息 → 序列化 → 写 socket let mut bc_task = tokio::spawn(broadcast_pump(bc_rx, socket_tx)); // 主循环:从 socket recv → 构造 BroadcastMessage → route loop { tokio::select! { // socket 入帧 maybe_msg = socket_rx.next() => { match maybe_msg { Some(Ok(Message::Text(text))) => { if let Err(e) = handle_inbound_text(&state, conn_id, kind, &device_id, &text).await { tracing::warn!(conn_id = conn_id.0, error = %e, "入站消息处理失败,忽略"); } } Some(Ok(Message::Binary(_))) => { // MVP 仅支持文本帧;二进制帧忽略(协议层可后续扩展) tracing::debug!(conn_id = conn_id.0, "收到二进制帧,忽略"); } Some(Ok(Message::Ping(_))) | Some(Ok(Message::Pong(_))) => { // axum/tungstenite 协议层 Ping/Pong 自动处理,这里仅记录 tracing::trace!(conn_id = conn_id.0, "协议层 Ping/Pong"); } Some(Ok(Message::Close(_))) | None => { tracing::info!(conn_id = conn_id.0, "客户端主动关闭连接"); break; } Some(Err(e)) => { tracing::warn!(conn_id = conn_id.0, error = %e, "socket 接收错误,断开"); break; } } } // 广播 pump 任务结束(socket_tx 关闭或 mpsc 关闭) res = &mut bc_task => { match res { Ok(()) => { tracing::debug!(conn_id = conn_id.0, "广播 pump 任务正常结束"); } Err(e) => { tracing::warn!(conn_id = conn_id.0, error = %e, "广播 pump 任务 panic"); } } break; } } } // 注销连接 if let Some(d) = state.remove_conn(conn_id).await { tracing::info!(conn_id = conn_id.0, device_id = %d, "连接已注销"); } // 结束 pump 任务(若仍在运行) bc_task.abort(); } /// 接收并解析首条 Hello 文本帧 async fn recv_hello(rx: &mut futures_util::stream::SplitStream) -> Result { let deadline = tokio::time::Duration::from_secs(10); let next = tokio::time::timeout(deadline, rx.next()) .await .map_err(|_| RelayError::Client("握手超时(10s 未收到 Hello)".into()))?; let msg = next .ok_or_else(|| RelayError::Client("握手阶段连接关闭".into()))? .map_err(|e| RelayError::WebSocket(format!("握手 recv 失败: {e}")))?; let text = match msg { Message::Text(t) => t, Message::Binary(_) => { return Err(RelayError::Client("握手首帧必须为文本".into())); } _ => return Err(RelayError::Client("握手首帧类型非法".into())), }; let hello: Hello = serde_json::from_str(&text).map_err(|e| RelayError::Client(format!("Hello 解析失败: {e}")))?; Ok(hello) } /// 处理入站文本帧(构造 BroadcastMessage → route) async fn handle_inbound_text( state: &RelayState, conn_id: ConnId, kind: ClientKind, device_id: &str, raw: &str, ) -> Result<()> { // 入站文本即业务 payload(relay 不解析),包成 BroadcastMessage // payload 直接用原始 JSON 值;若客户端发非 JSON 文本,则包成字符串值 let payload: serde_json::Value = serde_json::from_str(raw).unwrap_or(serde_json::Value::String(raw.to_string())); let now = now_ms(); let (msg_kind, from) = match kind { ClientKind::Device => (MessageKind::Event, ClientKind::Device), ClientKind::Miniapp => (MessageKind::Command, ClientKind::Miniapp), }; let msg = BroadcastMessage { device_id: device_id.to_string(), kind: msg_kind, source: conn_id, from, payload, ts: now, }; let delivered = state.route(&msg).await; tracing::debug!( conn_id = conn_id.0, ?msg_kind, device_id = %device_id, delivered, "入站消息已路由" ); Ok(()) } /// 广播 pump:从 mpsc 取消息,序列化后写回 socket sink /// /// 任务退出条件:bc_rx 关闭(对端 handle 全部 drop)/ socket_tx 关闭出错。 async fn broadcast_pump( mut bc_rx: mpsc::UnboundedReceiver, mut socket_tx: futures_util::stream::SplitSink, ) { while let Some(msg) = bc_rx.recv().await { let text = match serde_json::to_string(&msg) { Ok(t) => t, Err(e) => { tracing::warn!(error = %e, "广播消息序列化失败,跳过"); continue; } }; if let Err(e) = socket_tx.send(Message::Text(text)).await { tracing::warn!(error = %e, "广播写回 socket 失败,pump 退出"); break; } } } /// 便捷发送文本帧 async fn send_text( tx: &mut futures_util::stream::SplitSink, text: &str, ) -> Result<()> { tx.send(Message::Text(text.to_string())) .await .map_err(|e| RelayError::WebSocket(format!("发送失败: {e}"))) } /// 当前毫秒时间戳(避开 chrono workspace 依赖,直接用 std + SystemTime) fn now_ms() -> i64 { use std::time::{SystemTime, UNIX_EPOCH}; SystemTime::now() .duration_since(UNIX_EPOCH) .map(|d| d.as_millis() as i64) .unwrap_or(0) }