Skip to main content

futu_backend/conn/
request.rs

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    /// Raw transport primitive. Production callers must enter through a typed
40    /// command-family adapter or `ConnectionLifecycleRuntime`.
41    #[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    /// Raw reserved-header transport primitive; kept crate-private so command
47    /// identity, channel, auth and retry policy cannot be bypassed externally.
48    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(&registration_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(&registration_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                        &registration_identity,
220                        &termination,
221                    ) {
222                        emit_diagnostic(&self.diagnostic_sink, &facts);
223                        return Err(FutuError::Timeout);
224                    }
225
226                    // A different terminal observer won the generation. Keep
227                    // this waiter registered and return that winner's facts.
228                    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            // Admission publication and terminal drain share this mutex. If
242            // admission wins, facts cannot observe writer_admitted=false or
243            // drain before the command is queued. If termination wins, the
244            // lifecycle CAS rejects this generation before any publication.
245            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                            &registration_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    /// Raw detached-response transport primitive. Typed backend owners decide
379    /// which commands may legitimately proceed without awaiting the response.
380    ///
381    /// The serial remains registered until a response, connection failure, or
382    /// bounded timeout consumes it. This matches C++ `SendTCPProto_ProtoBuf`:
383    /// the business caller proceeds immediately, while the transport still
384    /// owns and drains the eventual response instead of misclassifying it as
385    /// an unmatched push.
386    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(&registration_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(&registration_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    /// Test-only transport probe used by the duplex mock-backend canary.
500    ///
501    /// This deliberately has a different name from the production primitive,
502    /// and is absent unless `test-util` is explicitly enabled.
503    #[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}