1use std::sync::Arc;
7use std::time::{Duration, Instant};
8
9use bytes::Bytes;
10use futures_util::{SinkExt, StreamExt};
11use helix_core::effect::TransportId;
12use helix_core::ports::FrameSender;
13use helix_core::tick::InboundBytes;
14use helix_core::{PortError, Tick};
15use reqwest::header::USER_AGENT;
16use tokio::sync::{mpsc, oneshot, watch, Mutex};
17use tokio_tungstenite::tungstenite::client::IntoClientRequest;
18use tokio_tungstenite::tungstenite::Message;
19
20use crate::metrics::{
21 AsyncMetricSink, LabelKey, MetricEvent, MetricId, MetricLabels, NoopMetricSink,
22};
23use crate::tick_ingress::TickIngressSender;
24use crate::trace::TraceCarrier;
25
26#[path = "network_util.rs"]
28mod network_util;
29use network_util::{header_value, headers_to_strings, map_ws_err, validate_url};
30
31mod config;
32pub use config::{HostHeaderRegistry, HostNetworkConfig};
33
34mod http;
35pub use http::{PreparedHttpRequest, SharedHttpClient};
36
37const WS_FRAME_BUFFER: usize = 64;
38
39#[derive(Clone)]
40pub(crate) enum InboundTickSender {
41 Raw(mpsc::Sender<Tick>),
42 Stamped(TickIngressSender),
43}
44
45impl InboundTickSender {
46 async fn send(&self, tick: Tick) -> Result<(), Tick> {
47 self.send_with_trace(tick, None).await
48 }
49
50 async fn send_with_trace(&self, tick: Tick, carrier: Option<TraceCarrier>) -> Result<(), Tick> {
52 match self {
53 Self::Raw(tx) => tx.send(tick).await.map_err(|error| error.0),
54 Self::Stamped(tx) => tx.send_with_trace(tick, carrier).await,
55 }
56 }
57}
58
59type WsSink = futures_util::stream::SplitSink<
60 tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>>,
61 Message,
62>;
63
64#[derive(Clone, Copy, Debug, PartialEq, Eq)]
65enum WsLifecyclePhase {
66 Prepared,
67 Active,
68 Finished,
69}
70
71pub(crate) struct WsConnectionLifecycle {
72 phase: Mutex<WsLifecyclePhase>,
73 inbound: Option<(TransportId, InboundTickSender)>,
74}
75
76impl WsConnectionLifecycle {
77 pub(crate) fn new(inbound: Option<(TransportId, InboundTickSender)>) -> Self {
78 Self {
79 phase: Mutex::new(WsLifecyclePhase::Prepared),
80 inbound,
81 }
82 }
83
84 pub(crate) async fn activate(&self) -> Result<bool, PortError> {
85 let mut phase = self.phase.lock().await;
86 if *phase != WsLifecyclePhase::Prepared {
87 return Ok(false);
88 }
89
90 if let Some((id, tick_tx)) = &self.inbound {
91 tick_tx
92 .send(Tick::Connected(*id))
93 .await
94 .map_err(|_| PortError::Transport("lifecycle tick channel closed".to_string()))?;
95 }
96 *phase = WsLifecyclePhase::Active;
97 Ok(true)
98 }
99
100 pub(crate) async fn finish_once(&self) {
101 let mut phase = self.phase.lock().await;
102 match *phase {
103 WsLifecyclePhase::Prepared => {
104 *phase = WsLifecyclePhase::Finished;
105 }
106 WsLifecyclePhase::Active => {
107 if let Some((id, tick_tx)) = &self.inbound {
108 tick_tx.send(Tick::Disconnected(*id)).await.ok();
109 }
110 *phase = WsLifecyclePhase::Finished;
113 }
114 WsLifecyclePhase::Finished => {}
115 }
116 }
117}
118
119#[must_use = "sender 注册完成后必须调用 activate,才能发布 Connected 并放行 reader"]
124pub struct WsConnectionActivation {
125 reader_gate: oneshot::Sender<()>,
126 lifecycle: Arc<WsConnectionLifecycle>,
127}
128
129impl WsConnectionActivation {
130 pub async fn activate(self) -> Result<(), PortError> {
131 if !self.lifecycle.activate().await? {
132 return Ok(());
133 }
134
135 self.reader_gate
136 .send(())
137 .map_err(|_| PortError::Transport("websocket reader stopped before activation".into()))
138 }
139}
140
141struct WsCloseCompletion {
142 result: watch::Sender<Option<Result<(), String>>>,
143}
144
145impl WsCloseCompletion {
146 fn new() -> Self {
147 let (result, _) = watch::channel(None);
148 Self { result }
149 }
150
151 fn complete(&self, result: Result<(), String>) {
152 self.result.send_replace(Some(result));
153 }
154
155 async fn wait(&self) -> Result<(), PortError> {
156 let mut result_rx = self.result.subscribe();
157 loop {
158 let result = result_rx.borrow().clone();
159 if let Some(result) = result {
160 return result.map_err(PortError::Transport);
161 }
162 result_rx.changed().await.map_err(|_| {
163 PortError::Transport("websocket cleanup completion channel closed".to_string())
164 })?;
165 }
166 }
167}
168
169enum WsState {
170 Disconnected,
171 Closing {
172 completion: Arc<WsCloseCompletion>,
173 },
174 Connected {
175 sink: WsSink,
176 frame_rx: mpsc::Receiver<Result<Bytes, PortError>>,
177 reader: tokio::task::JoinHandle<()>,
178 lifecycle: Arc<WsConnectionLifecycle>,
179 },
180}
181
182pub struct SharedWsClient {
190 config: HostNetworkConfig,
191 headers: HostHeaderRegistry,
192 state: Arc<Mutex<WsState>>,
193 inbound: Option<(TransportId, InboundTickSender)>,
195 metrics: Arc<dyn AsyncMetricSink>,
196}
197
198impl SharedWsClient {
199 pub fn new(config: HostNetworkConfig) -> Result<Self, PortError> {
200 Self::with_registry(config, HostHeaderRegistry::default())
201 }
202
203 pub fn with_registry(
204 config: HostNetworkConfig,
205 headers: HostHeaderRegistry,
206 ) -> Result<Self, PortError> {
207 validate_url(&config.ws_url, "ws_url")?;
208 Ok(Self {
209 config,
210 headers,
211 state: Arc::new(Mutex::new(WsState::Disconnected)),
212 inbound: None,
213 metrics: Arc::new(NoopMetricSink),
214 })
215 }
216
217 pub fn headers(&self) -> HostHeaderRegistry {
218 self.headers.clone()
219 }
220
221 pub fn with_inbound_tick(mut self, id: TransportId, tick_tx: mpsc::Sender<Tick>) -> Self {
226 self.inbound = Some((id, InboundTickSender::Raw(tick_tx)));
227 self
228 }
229
230 pub fn with_stamped_inbound_tick(
231 mut self,
232 id: TransportId,
233 tick_tx: TickIngressSender,
234 ) -> Self {
235 self.inbound = Some((id, InboundTickSender::Stamped(tick_tx)));
236 self
237 }
238
239 pub fn with_metric_sink(mut self, metrics: Arc<dyn AsyncMetricSink>) -> Self {
240 self.metrics = metrics;
241 self
242 }
243}
244
245impl SharedWsClient {
246 pub async fn connect(&mut self) -> Result<WsConnectionActivation, PortError> {
251 let is_disconnected = {
252 let state = self.state.lock().await;
253 matches!(&*state, WsState::Disconnected)
254 };
255 if !is_disconnected {
256 return Err(PortError::Transport(
257 "connect on active or closing transport".to_string(),
258 ));
259 }
260
261 let mut request = self
262 .config
263 .ws_url
264 .as_str()
265 .into_client_request()
266 .map_err(|e| PortError::Transport(format!("invalid ws request: {e}")))?;
267 for (name, value) in self.headers.snapshot().await.iter() {
268 request.headers_mut().insert(name.clone(), value.clone());
269 }
270 if let Some(user_agent) = &self.config.user_agent {
271 request
272 .headers_mut()
273 .insert(USER_AGENT, header_value("user-agent", user_agent)?);
274 }
275 let handshake_headers = headers_to_strings(request.headers());
276 crate::network_debug::dump_ws_handshake(&self.config.ws_url, &handshake_headers);
277
278 let connect_started = Instant::now();
279 let connect_result = tokio_tungstenite::connect_async(request).await;
280 let connect_status = if connect_result.is_ok() {
281 "ok"
282 } else {
283 "error"
284 };
285 if self.metrics.is_enabled() {
286 let _ = self.metrics.try_record(MetricEvent::histogram(
287 MetricId::WsConnectDurationSeconds,
288 connect_started.elapsed().as_secs_f64(),
289 MetricLabels::one(LabelKey::Stage, "ws").with(LabelKey::Status, connect_status),
290 ));
291 }
292 let (ws, _resp) = connect_result.map_err(map_ws_err)?;
293 let (sink, mut stream) = ws.split();
294 let (frame_tx, frame_rx) = mpsc::channel::<Result<Bytes, PortError>>(WS_FRAME_BUFFER);
295 let (reader_gate, reader_gate_rx) = oneshot::channel();
296 let (reader_ready_tx, reader_ready_rx) = oneshot::channel();
297
298 let dump_frames = std::env::var("HELIX_DUMP_FRAMES").is_ok();
300 let inbound = self.inbound.clone();
301 let metrics = Arc::clone(&self.metrics);
302 let lifecycle = Arc::new(WsConnectionLifecycle::new(inbound.clone()));
303 let reader_lifecycle = Arc::clone(&lifecycle);
304
305 let reader = tokio::spawn(async move {
306 reader_ready_tx.send(()).ok();
307 if reader_gate_rx.await.is_err() {
308 return;
309 }
310
311 while let Some(msg) = stream.next().await {
312 let frame = match msg {
314 Ok(Message::Binary(data)) => Ok(Bytes::from(data)),
315 Ok(Message::Text(text)) => Ok(Bytes::from(text.into_bytes())),
316 Ok(Message::Close(_)) => break,
317 Ok(Message::Ping(_)) | Ok(Message::Pong(_)) | Ok(Message::Frame(_)) => continue,
318 Err(e) => Err(map_ws_err(e)),
319 };
320
321 if dump_frames {
322 if let Ok(bytes) = &frame {
323 tracing::info!(
324 target: "helix_ws_frames",
325 len = bytes.len(),
326 frame = %String::from_utf8_lossy(bytes),
327 "RAW WS inbound"
328 );
329 }
330 }
331 if let Ok(bytes) = &frame {
332 crate::network_debug::dump_ws_inbound(bytes);
333 if metrics.is_enabled() {
334 let _ = metrics.try_record(MetricEvent::counter(
335 MetricId::OperationsTotal,
336 1.0,
337 MetricLabels::one(LabelKey::Stage, "ws")
338 .with(LabelKey::Protocol, "ws")
339 .with(LabelKey::Direction, "inbound")
340 .with(LabelKey::Status, "ok"),
341 ));
342 }
343 }
344
345 match &inbound {
346 Some((_, tick_tx)) => match frame {
349 Ok(bytes) => {
350 let carrier = TraceCarrier::from_ws_frame(&bytes);
351 if tick_tx
352 .send_with_trace(Tick::Inbound(InboundBytes(bytes)), carrier)
353 .await
354 .is_err()
355 {
356 break; }
358 }
359 Err(_) => break,
361 },
362 None => {
364 if frame_tx.send(frame).await.is_err() {
365 break; }
367 }
368 }
369 }
370
371 reader_lifecycle.finish_once().await;
373 });
374
375 if reader_ready_rx.await.is_err() {
376 reader.abort();
377 return Err(PortError::Transport(
378 "websocket reader stopped before reaching activation gate".to_string(),
379 ));
380 }
381
382 *self.state.lock().await = WsState::Connected {
383 sink,
384 frame_rx,
385 reader,
386 lifecycle: Arc::clone(&lifecycle),
387 };
388 Ok(WsConnectionActivation {
389 reader_gate,
390 lifecycle,
391 })
392 }
393
394 async fn send_frame(&self, frame: Bytes) -> Result<(), PortError> {
401 crate::network_debug::dump_ws_outbound(&frame);
402 let mut state = self.state.lock().await;
403 let result = match &mut *state {
404 WsState::Connected { sink, .. } => {
405 let text = String::from_utf8(frame.to_vec())
406 .map_err(|e| PortError::Transport(format!("non-utf8 ws frame: {e}")))?;
407 sink.send(Message::Text(text)).await.map_err(map_ws_err)
408 }
409 WsState::Disconnected => Err(PortError::Transport(
410 "send on disconnected transport".to_string(),
411 )),
412 WsState::Closing { .. } => Err(PortError::Transport(
413 "send on closing transport".to_string(),
414 )),
415 };
416 drop(state);
417 if self.metrics.is_enabled() {
418 let status = if result.is_ok() { "ok" } else { "error" };
419 let labels = MetricLabels::one(LabelKey::Stage, "ws")
420 .with(LabelKey::Protocol, "ws")
421 .with(LabelKey::Direction, "outbound")
422 .with(LabelKey::Status, status);
423 let _ = self.metrics.try_record(MetricEvent::counter(
424 MetricId::OperationsTotal,
425 1.0,
426 labels,
427 ));
428 if result.is_err() {
429 let _ = self.metrics.try_record(MetricEvent::counter(
430 MetricId::ErrorsTotal,
431 1.0,
432 labels.with(LabelKey::ErrorKind, "ws_send_failed"),
433 ));
434 }
435 }
436 result
437 }
438
439 pub async fn recv(&self) -> Result<Option<Bytes>, PortError> {
440 let mut state = self.state.lock().await;
441 match &mut *state {
442 WsState::Connected { frame_rx, .. } => match frame_rx.recv().await {
443 Some(Ok(bytes)) => Ok(Some(bytes)),
444 Some(Err(e)) => Err(e),
445 None => Ok(None),
446 },
447 WsState::Disconnected | WsState::Closing { .. } => Ok(None),
448 }
449 }
450
451 pub async fn close(&self) -> Result<(), PortError> {
452 enum CloseAction {
453 Done,
454 Wait(Arc<WsCloseCompletion>),
455 Start {
456 sink: WsSink,
457 reader: tokio::task::JoinHandle<()>,
458 lifecycle: Arc<WsConnectionLifecycle>,
459 completion: Arc<WsCloseCompletion>,
460 close_timeout: Duration,
461 },
462 }
463
464 let action = {
465 let mut state = self.state.lock().await;
466 match std::mem::replace(&mut *state, WsState::Disconnected) {
467 WsState::Disconnected => CloseAction::Done,
468 WsState::Closing { completion } => {
469 *state = WsState::Closing {
470 completion: Arc::clone(&completion),
471 };
472 CloseAction::Wait(completion)
473 }
474 WsState::Connected {
475 sink,
476 reader,
477 lifecycle,
478 ..
479 } => {
480 let completion = Arc::new(WsCloseCompletion::new());
481 *state = WsState::Closing {
482 completion: Arc::clone(&completion),
483 };
484 CloseAction::Start {
485 sink,
486 reader,
487 lifecycle,
488 completion,
489 close_timeout: self.config.timeout,
490 }
491 }
492 }
493 };
494
495 match action {
496 CloseAction::Done => Ok(()),
497 CloseAction::Wait(completion) => completion.wait().await,
498 CloseAction::Start {
499 mut sink,
500 reader,
501 lifecycle,
502 completion,
503 close_timeout,
504 } => {
505 let state = Arc::clone(&self.state);
506 let completion_for_cleanup = Arc::clone(&completion);
507 let cleanup = tokio::spawn(async move {
508 let close_result =
511 match tokio::time::timeout(close_timeout, sink.send(Message::Close(None)))
512 .await
513 {
514 Ok(result) => result.map_err(|error| error.to_string()),
515 Err(_) => Err(format!(
516 "websocket Close write timed out after {}ms",
517 close_timeout.as_millis()
518 )),
519 };
520 reader.abort();
521 let _ = reader.await;
524 lifecycle.finish_once().await;
525 let mut state = state.lock().await;
526 completion_for_cleanup.complete(close_result);
527 *state = WsState::Disconnected;
528 });
529
530 cleanup.await.map_err(|error| {
531 PortError::Transport(format!("websocket cleanup task failed: {error}"))
532 })?;
533 completion.wait().await
535 }
536 }
537 }
538}
539
540#[async_trait::async_trait]
541impl FrameSender for SharedWsClient {
542 async fn send(&self, frame: Bytes) -> Result<(), PortError> {
543 self.send_frame(frame).await
544 }
545}