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