1use std::sync::Arc;
2use std::sync::atomic::{AtomicBool, Ordering};
3
4use tokio::sync::oneshot;
5
6use futu_core::error::{FutuError, Result};
7
8use super::diagnostics::{
9 PendingFailureFacts, PendingFailureKind, PendingRegistrationIdentity, PendingResponseEntry,
10 TcpLoginTransportStage, trace_tcp_login_transport_stage,
11};
12use super::inbound::response_ex_head_error;
13use super::lifecycle::{
14 ConnectionTerminationFacts, PendingRegistrationGuard,
15 claim_response_timeout_and_fail_other_pending, emit_diagnostic, mark_disconnected,
16 pending_failure_facts,
17};
18use super::{BackendCmd, BackendConn, NNFrame};
19
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub(crate) enum RequestTimeoutPolicy {
22 Disconnect,
23 KeepConnection,
24}
25
26#[derive(Clone)]
27pub(crate) struct WriterAdmissionObserver {
28 pub(crate) try_admit: Arc<dyn Fn(BackendWriterAdmission) -> bool + Send + Sync>,
29 pub(crate) publish_terminal: Option<Arc<WriterAdmissionPublishTerminal>>,
30 pub(crate) on_accepted: Arc<dyn Fn(BackendWriterAdmission) + Send + Sync>,
31}
32
33pub(crate) type WriterAdmissionPublishTerminal =
34 dyn Fn(BackendWriterAdmission, &mut dyn FnMut()) -> bool + Send + Sync;
35
36#[derive(Clone, Copy, Debug, PartialEq, Eq)]
37pub(crate) struct BackendWriterAdmission {
38 pub(crate) connection_generation: u64,
39 pub(crate) serial_no: u32,
40}
41
42impl BackendConn {
43 #[cfg(any(test, feature = "test-util"))]
46 pub(crate) async fn request(&self, cmd_id: u16, body: Vec<u8>) -> Result<NNFrame> {
47 self.request_with_reserved(cmd_id, body, [0u8; 10]).await
48 }
49
50 pub(crate) async fn request_with_reserved(
53 &self,
54 cmd_id: u16,
55 body: Vec<u8>,
56 reserved: [u8; 10],
57 ) -> Result<NNFrame> {
58 self.request_with_reserved_timeout(
59 cmd_id,
60 body,
61 reserved,
62 std::time::Duration::from_secs(10),
63 )
64 .await
65 }
66
67 pub(crate) async fn request_with_reserved_timeout(
68 &self,
69 cmd_id: u16,
70 body: Vec<u8>,
71 reserved: [u8; 10],
72 timeout: std::time::Duration,
73 ) -> Result<NNFrame> {
74 self.request_with_reserved_timeout_policy(
75 cmd_id,
76 body,
77 reserved,
78 timeout,
79 RequestTimeoutPolicy::Disconnect,
80 )
81 .await
82 }
83
84 pub(crate) async fn request_with_reserved_timeout_policy(
85 &self,
86 cmd_id: u16,
87 body: Vec<u8>,
88 reserved: [u8; 10],
89 timeout: std::time::Duration,
90 timeout_policy: RequestTimeoutPolicy,
91 ) -> Result<NNFrame> {
92 self.request_with_reserved_timeout_policy_observed(
93 cmd_id,
94 body,
95 reserved,
96 timeout,
97 timeout_policy,
98 None,
99 )
100 .await
101 }
102
103 pub(crate) async fn request_with_reserved_timeout_policy_observed(
104 &self,
105 cmd_id: u16,
106 body: Vec<u8>,
107 reserved: [u8; 10],
108 timeout: std::time::Duration,
109 timeout_policy: RequestTimeoutPolicy,
110 observer: Option<WriterAdmissionObserver>,
111 ) -> Result<NNFrame> {
112 if !self.is_connected() {
113 return Err(FutuError::NotInitialized);
114 }
115
116 let deadline = tokio::time::Instant::now() + timeout;
117 let outbound_guard = tokio::time::timeout_at(deadline, self.outbound_order.lock())
118 .await
119 .map_err(|_elapsed| FutuError::Timeout)?;
120 if !self.is_connected() {
121 return Err(FutuError::NotInitialized);
122 }
123 let frame = self.build_outbound_frame(cmd_id, body, reserved)?;
124 let serial_no = frame.header.serial_no;
125
126 let (resp_tx, mut resp_rx) = oneshot::channel();
127 let writer_admitted = Arc::new(AtomicBool::new(false));
128 let registration_identity = PendingRegistrationIdentity::new();
129 {
130 let mut pending = self.pending.lock();
131 if !self.is_connected() {
132 return Err(FutuError::NotInitialized);
133 }
134 match pending.entry(serial_no) {
135 std::collections::hash_map::Entry::Vacant(slot) => {
136 slot.insert(PendingResponseEntry {
137 cmd_id,
138 serial_no,
139 registration_identity: Arc::clone(®istration_identity),
140 writer_admitted: Arc::clone(&writer_admitted),
141 tx: resp_tx,
142 });
143 }
144 std::collections::hash_map::Entry::Occupied(existing) => {
145 tracing::error!(
146 cmd_id,
147 serial_no,
148 existing_cmd_id = existing.get().cmd_id,
149 connection_generation = self.connection_generation,
150 "internal_pending_serial_collision"
151 );
152 return Err(FutuError::Codec("internal_pending_serial_collision".into()));
153 }
154 }
155 }
156 trace_tcp_login_transport_stage(
157 cmd_id,
158 serial_no,
159 TcpLoginTransportStage::PendingRegistered,
160 );
161 let _pending_registration = PendingRegistrationGuard {
162 pending: Arc::clone(&self.pending),
163 serial_no,
164 registration_identity: Arc::clone(®istration_identity),
165 };
166 let permit = tokio::select! {
167 response = &mut resp_rx => {
168 drop(outbound_guard);
169 let resp = self.map_pending_response(
170 cmd_id,
171 serial_no,
172 &writer_admitted,
173 response,
174 )?;
175 if let Some(err) = response_ex_head_error(&resp) {
176 return Err(err);
177 }
178 return Ok(resp);
179 }
180 reserve = self.cmd_tx.reserve() => match reserve {
181 Ok(permit) => permit,
182 Err(_closed) => {
183 drop(outbound_guard);
184 if self.termination_lifecycle.is_terminated() {
185 let response = resp_rx.await;
186 return self.map_pending_response(
187 cmd_id,
188 serial_no,
189 &writer_admitted,
190 response,
191 );
192 }
193 mark_disconnected(&self.connected, &self.connected_tx);
194 return Err(FutuError::NotInitialized);
195 }
196 },
197 _ = tokio::time::sleep_until(deadline) => {
198 let progress = self.inbound_progress.snapshot();
199 let facts = pending_failure_facts(
200 PendingFailureKind::ResponseTimeout,
201 cmd_id,
202 serial_no,
203 writer_admitted.load(Ordering::Acquire),
204 &self.endpoint_fingerprint,
205 self.connection_generation,
206 &progress,
207 );
208 if timeout_policy == RequestTimeoutPolicy::Disconnect {
209 let termination = ConnectionTerminationFacts {
210 kind: PendingFailureKind::ResponseTimeout,
211 endpoint_fingerprint: self.endpoint_fingerprint.clone(),
212 connection_generation: self.connection_generation,
213 progress,
214 diagnostic_sink: self.diagnostic_sink.clone(),
215 };
216 if claim_response_timeout_and_fail_other_pending(
217 &self.termination_lifecycle,
218 &self.connected,
219 &self.connected_tx,
220 &self.shutdown_tx,
221 &self.pending,
222 serial_no,
223 ®istration_identity,
224 &termination,
225 ) {
226 emit_diagnostic(&self.diagnostic_sink, &facts);
227 return Err(FutuError::Timeout);
228 }
229
230 drop(outbound_guard);
233 return self.map_pending_response(
234 cmd_id,
235 serial_no,
236 &writer_admitted,
237 resp_rx.await,
238 );
239 }
240 emit_diagnostic(&self.diagnostic_sink, &facts);
241 return Err(FutuError::Timeout);
242 }
243 };
244 {
245 let pending_guard = self.pending.lock();
250 let Some(admission_activity) = self.termination_lifecycle.try_begin_admission() else {
251 drop(pending_guard);
252 drop(permit);
253 drop(outbound_guard);
254 return self.map_pending_response(
255 cmd_id,
256 serial_no,
257 &writer_admitted,
258 resp_rx.await,
259 );
260 };
261 #[cfg(any(test, feature = "test-util"))]
262 self.termination_lifecycle
263 .pause_admission_publish_for_test();
264 let admission = BackendWriterAdmission {
265 connection_generation: self.connection_generation,
266 serial_no,
267 };
268 let mut permit = Some(permit);
269 let mut frame = Some(frame);
270 let mut published = false;
271 let mut publish = || {
272 let (Some(permit), Some(frame)) = (permit.take(), frame.take()) else {
273 tracing::error!(
274 cmd_id,
275 serial_no,
276 connection_generation = self.connection_generation,
277 "writer admission terminal attempted duplicate publication"
278 );
279 return;
280 };
281 writer_admitted.store(true, Ordering::Release);
282 permit.send(BackendCmd::Send {
283 frame,
284 writer_admitted: Some(Arc::clone(&writer_admitted)),
285 });
286 published = true;
287 };
288 let admitted = match observer.as_ref() {
289 Some(observer) => match observer.publish_terminal.as_ref() {
290 Some(terminal) => terminal(admission, &mut publish),
291 None if (observer.try_admit)(admission) => {
292 publish();
293 true
294 }
295 None => false,
296 },
297 None => {
298 publish();
299 true
300 }
301 };
302 if !admitted && !published {
303 drop(admission_activity);
304 drop(pending_guard);
305 drop(outbound_guard);
306 return Err(FutuError::WriterAdmissionRejected);
307 }
308 if !published {
309 drop(admission_activity);
310 drop(pending_guard);
311 drop(outbound_guard);
312 return Err(FutuError::Codec(
313 "writer admission terminal returned success without publishing".into(),
314 ));
315 }
316 if let Some(observer) = &observer {
317 (observer.on_accepted)(admission);
318 }
319 drop(admission_activity);
320 drop(pending_guard);
321 }
322 drop(outbound_guard);
323
324 let resp = crate::delay_stats::trace_backend_request(cmd_id, async {
325 match tokio::time::timeout_at(deadline, &mut resp_rx).await {
326 Ok(response) => {
327 self.map_pending_response(cmd_id, serial_no, &writer_admitted, response)
328 }
329 Err(_elapsed) => {
330 let progress = self.inbound_progress.snapshot();
331 let facts = pending_failure_facts(
332 PendingFailureKind::ResponseTimeout,
333 cmd_id,
334 serial_no,
335 writer_admitted.load(Ordering::Acquire),
336 &self.endpoint_fingerprint,
337 self.connection_generation,
338 &progress,
339 );
340 if timeout_policy == RequestTimeoutPolicy::Disconnect {
341 let termination = ConnectionTerminationFacts {
342 kind: PendingFailureKind::ResponseTimeout,
343 endpoint_fingerprint: self.endpoint_fingerprint.clone(),
344 connection_generation: self.connection_generation,
345 progress,
346 diagnostic_sink: self.diagnostic_sink.clone(),
347 };
348 if claim_response_timeout_and_fail_other_pending(
349 &self.termination_lifecycle,
350 &self.connected,
351 &self.connected_tx,
352 &self.shutdown_tx,
353 &self.pending,
354 serial_no,
355 ®istration_identity,
356 &termination,
357 ) {
358 emit_diagnostic(&self.diagnostic_sink, &facts);
359 return Err(FutuError::Timeout);
360 }
361
362 return self.map_pending_response(
363 cmd_id,
364 serial_no,
365 &writer_admitted,
366 resp_rx.await,
367 );
368 }
369 emit_diagnostic(&self.diagnostic_sink, &facts);
370 Err(FutuError::Timeout)
371 }
372 }
373 })
374 .await?;
375
376 if let Some(err) = response_ex_head_error(&resp) {
377 return Err(err);
378 }
379
380 Ok(resp)
381 }
382
383 fn map_pending_response(
384 &self,
385 cmd_id: u16,
386 serial_no: u32,
387 writer_admitted: &AtomicBool,
388 response: std::result::Result<
389 std::result::Result<NNFrame, PendingFailureFacts>,
390 oneshot::error::RecvError,
391 >,
392 ) -> Result<NNFrame> {
393 match response {
394 Ok(Ok(resp)) => Ok(resp),
395 Ok(Err(facts)) => Err(FutuError::TransportFailure {
396 reason: facts.transport_reason(),
397 detail: facts.to_string(),
398 }),
399 Err(_sender_dropped) => {
400 tracing::error!(
401 cmd_id,
402 serial_no,
403 writer_admitted = writer_admitted.load(Ordering::Acquire),
404 endpoint_fingerprint = %self.endpoint_fingerprint,
405 connection_generation = self.connection_generation,
406 "internal_sender_drop: pending response sender vanished without typed facts"
407 );
408 Err(FutuError::Codec("internal_sender_drop".into()))
409 }
410 }
411 }
412
413 pub(crate) async fn send_fire_and_forget(&self, cmd_id: u16, body: Vec<u8>) -> Result<()> {
422 if !self.is_connected() {
423 return Err(FutuError::NotInitialized);
424 }
425 let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(10);
426 let outbound_guard = tokio::time::timeout_at(deadline, self.outbound_order.lock())
427 .await
428 .map_err(|_elapsed| FutuError::Timeout)?;
429 if !self.is_connected() {
430 return Err(FutuError::NotInitialized);
431 }
432 let frame = self.build_outbound_frame(cmd_id, body, [0u8; 10])?;
433 let serial_no = frame.header.serial_no;
434 let (resp_tx, mut resp_rx) = oneshot::channel();
435 let writer_admitted = Arc::new(AtomicBool::new(false));
436 let registration_identity = PendingRegistrationIdentity::new();
437 {
438 let mut pending = self.pending.lock();
439 if !self.is_connected() {
440 return Err(FutuError::NotInitialized);
441 }
442 match pending.entry(serial_no) {
443 std::collections::hash_map::Entry::Vacant(slot) => {
444 slot.insert(PendingResponseEntry {
445 cmd_id,
446 serial_no,
447 registration_identity: Arc::clone(®istration_identity),
448 writer_admitted: Arc::clone(&writer_admitted),
449 tx: resp_tx,
450 });
451 }
452 std::collections::hash_map::Entry::Occupied(_) => {
453 return Err(FutuError::Codec("internal_pending_serial_collision".into()));
454 }
455 }
456 }
457 let pending_registration = PendingRegistrationGuard {
458 pending: Arc::clone(&self.pending),
459 serial_no,
460 registration_identity: Arc::clone(®istration_identity),
461 };
462 let permit = tokio::select! {
463 response = &mut resp_rx => {
464 drop(outbound_guard);
465 self.map_pending_response(cmd_id, serial_no, &writer_admitted, response)?;
466 return Ok(());
467 }
468 reserve = self.cmd_tx.reserve() => match reserve {
469 Ok(permit) => permit,
470 Err(_closed) => {
471 drop(outbound_guard);
472 if self.termination_lifecycle.is_terminated() {
473 self.map_pending_response(
474 cmd_id,
475 serial_no,
476 &writer_admitted,
477 resp_rx.await,
478 )?;
479 return Ok(());
480 }
481 mark_disconnected(&self.connected, &self.connected_tx);
482 return Err(FutuError::NotInitialized);
483 }
484 },
485 _ = tokio::time::sleep_until(deadline) => {
486 return Err(FutuError::Timeout);
487 }
488 };
489 {
490 let pending_guard = self.pending.lock();
491 let Some(admission_activity) = self.termination_lifecycle.try_begin_admission() else {
492 drop(pending_guard);
493 drop(permit);
494 drop(outbound_guard);
495 self.map_pending_response(cmd_id, serial_no, &writer_admitted, resp_rx.await)?;
496 return Ok(());
497 };
498 #[cfg(test)]
499 self.termination_lifecycle
500 .pause_admission_publish_for_test();
501 writer_admitted.store(true, Ordering::Release);
502 permit.send(BackendCmd::Send {
503 frame,
504 writer_admitted: Some(Arc::clone(&writer_admitted)),
505 });
506 drop(admission_activity);
507 drop(pending_guard);
508 }
509 drop(outbound_guard);
510
511 tokio::spawn(async move {
512 let _pending_registration = pending_registration;
513 match tokio::time::timeout_at(deadline, resp_rx).await {
514 Ok(Ok(Ok(_response))) => {}
515 Ok(Ok(Err(facts))) => {
516 tracing::debug!(cmd_id, serial_no, detail = %facts, "detached backend response failed");
517 }
518 Ok(Err(_sender_dropped)) => {
519 tracing::debug!(
520 cmd_id,
521 serial_no,
522 "detached backend response sender dropped"
523 );
524 }
525 Err(_timeout) => {
526 tracing::debug!(cmd_id, serial_no, "detached backend response timed out");
527 }
528 }
529 });
530
531 Ok(())
532 }
533
534 #[cfg(feature = "test-util")]
539 pub async fn request_for_test(&self, cmd_id: u16, body: Vec<u8>) -> Result<NNFrame> {
540 self.request(cmd_id, body).await
541 }
542}