tdm_server_rust/web/
collaboration_hub_controller.rs1use 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
35const KEEPALIVE_INTERVAL_SECS: u64 = 15;
37
38#[derive(Debug, Deserialize)]
40struct WsQuery {
41 #[serde(default)]
43 access_token: Option<String>,
44}
45
46pub 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#[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#[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
83fn 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
92fn 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
100async 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 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 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 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 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 hub::disconnect(&state, &conn_id).await;
190 writer.abort();
191 pinger.abort();
192}