Skip to main content

tdm_server_rust/web/
collaboration_hub_controller.rs

1//! 在线翻译协作 Hub 接口 (Collaboration Hub Controller)
2//!
3//! 兼容 `@microsoft/signalr` 前端:
4//!
5//! | 方法 | 路径 | 说明 |
6//! |------|------|------|
7//! | POST | `/hubs/translator-collaboration/negotiate` | SignalR 协商,返回 connectionId |
8//! | GET  | `/hubs/translator-collaboration` | WebSocket 升级(SignalR JSON Hub 协议) |
9//!
10//! ## 鉴权
11//!
12//! WS 不走 `/api` 的 `auth_middleware`:浏览器无法为 WS 设置请求头,token 由 SignalR 以
13//! 查询参数 `access_token` 透传(也兼容 `token` / `Authorization` 头)。鉴权失败拒绝升级。
14
15use crate::{
16    app::AppState,
17    collaboration::{hub, signalr},
18    utils::jwt::JwtUtil,
19};
20use axum::{
21    extract::{
22        ws::{Message, WebSocket, WebSocketUpgrade},
23        Query, State,
24    },
25    http::{HeaderMap, StatusCode},
26    response::{IntoResponse, Response},
27    routing::{get, post},
28    Json, Router,
29};
30use futures::{SinkExt, StreamExt};
31use serde::Deserialize;
32use serde_json::json;
33use std::time::Duration;
34
35/// 服务端心跳间隔(秒),低于 SignalR 默认 serverTimeout(30s)
36const KEEPALIVE_INTERVAL_SECS: u64 = 15;
37
38/// WS 连接查询参数
39#[derive(Debug, Deserialize)]
40struct WsQuery {
41    /// SignalR 透传的访问令牌
42    #[serde(default)]
43    access_token: Option<String>,
44}
45
46/// 协作 Hub 路由(挂载于顶层,路径含 `/hubs` 前缀)
47pub fn routes() -> Router<AppState> {
48    Router::new()
49        .route("/hubs/translator-collaboration/negotiate", post(negotiate))
50        .route("/hubs/translator-collaboration", get(ws_entry))
51}
52
53/// SignalR 协商:返回连接 ID 与可用传输(仅 WebSockets)
54#[tracing::instrument(skip_all, level = "debug")]
55async fn negotiate() -> impl IntoResponse {
56    let connection_id = uuid::Uuid::new_v4().to_string();
57    Json(json!({
58        "connectionId": connection_id,
59        "availableTransports": [
60            { "transport": "WebSockets", "transferFormats": ["Text", "Binary"] }
61        ]
62    }))
63}
64
65/// WebSocket 入口:鉴权后升级连接
66#[tracing::instrument(skip_all, level = "info")]
67async fn ws_entry(
68    State(state): State<AppState>,
69    Query(q): Query<WsQuery>,
70    headers: HeaderMap,
71    ws: WebSocketUpgrade,
72) -> Response {
73    let token = q
74        .access_token
75        .clone()
76        .or_else(|| extract_header_token(&headers));
77    let Some(user_id) = token.and_then(|t| parse_user_id(&state, &t)) else {
78        return (StatusCode::UNAUTHORIZED, "实时协作需要登录喵……").into_response();
79    };
80    ws.on_upgrade(move |socket| handle_socket(state, socket, user_id))
81}
82
83/// 从请求头提取 token(兼容非浏览器客户端)
84fn extract_header_token(headers: &HeaderMap) -> Option<String> {
85    headers
86        .get("token")
87        .or_else(|| headers.get("Authorization"))
88        .and_then(|v| v.to_str().ok())
89        .map(|s| s.trim_start_matches("Bearer ").to_string())
90}
91
92/// 解析 JWT 得到用户(组员)ID 字符串
93fn parse_user_id(state: &AppState, token: &str) -> Option<String> {
94    JwtUtil::new(&state.config.jwt.sign_key, state.config.jwt.expire_ms)
95        .parse(token)
96        .ok()
97        .map(|c| c.id.to_string())
98}
99
100/// 处理单个 WS 连接的完整生命周期
101async fn handle_socket(state: AppState, socket: WebSocket, user_id: String) {
102    let conn_id = uuid::Uuid::new_v4().to_string();
103    let (mut sink, mut stream) = socket.split();
104    let mut rx = state.collaboration.register(conn_id.clone(), user_id);
105    let out_tx = state.collaboration.sender(&conn_id);
106
107    // 出站任务:把通道里的文本帧写到 WS
108    let writer = tokio::spawn(async move {
109        while let Some(frame) = rx.recv().await {
110            if sink.send(Message::Text(frame)).await.is_err() {
111                break;
112            }
113        }
114    });
115
116    // 心跳任务:周期性向通道推 Ping,维持 SignalR 连接
117    let ping_state = state.clone();
118    let ping_conn = conn_id.clone();
119    let pinger = tokio::spawn(async move {
120        let mut interval = tokio::time::interval(Duration::from_secs(KEEPALIVE_INTERVAL_SECS));
121        interval.tick().await;
122        loop {
123            interval.tick().await;
124            match ping_state.collaboration.sender(&ping_conn) {
125                Some(tx) => {
126                    if tx.send(signalr::ping_frame()).is_err() {
127                        break;
128                    }
129                }
130                None => break,
131            }
132        }
133    });
134
135    let mut handshaken = false;
136    while let Some(Ok(msg)) = stream.next().await {
137        match msg {
138            Message::Text(text) => {
139                // 首帧为 SignalR 握手请求,回 `{}\x1e`
140                if !handshaken && signalr::is_handshake_request(&text) {
141                    if let Some(tx) = &out_tx {
142                        let _ = tx.send(signalr::handshake_response());
143                    }
144                    handshaken = true;
145                    continue;
146                }
147                let messages = match signalr::decode_messages(&text) {
148                    Ok(m) => m,
149                    Err(_) => continue,
150                };
151                for message in messages {
152                    match message {
153                        signalr::HubMessage::Invocation(inv) => {
154                            let result = hub::handle_invocation(
155                                &state,
156                                &conn_id,
157                                &inv.target,
158                                inv.arguments,
159                            )
160                            .await;
161                            // 仅当客户端使用 invoke(带 invocationId)时回 Completion
162                            if let Some(inv_id) = inv.invocation_id {
163                                let completion = match result {
164                                    Ok(r) => signalr::HubMessage::completion(inv_id, r),
165                                    Err(e) => signalr::HubMessage::completion_error(inv_id, e),
166                                };
167                                if let (Some(tx), Ok(frame)) =
168                                    (&out_tx, signalr::encode_message(&completion))
169                                {
170                                    let _ = tx.send(frame);
171                                }
172                            }
173                        }
174                        signalr::HubMessage::Ping => {
175                            if let Some(tx) = &out_tx {
176                                let _ = tx.send(signalr::ping_frame());
177                            }
178                        }
179                        signalr::HubMessage::Completion(_) => {}
180                    }
181                }
182            }
183            Message::Close(_) => break,
184            _ => {}
185        }
186    }
187
188    // 清理:释放锁、移除在线状态、广播
189    hub::disconnect(&state, &conn_id).await;
190    writer.abort();
191    pinger.abort();
192}