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) 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    /// Raw transport primitive. Production callers must enter through a typed
44    /// command-family adapter or `ConnectionLifecycleRuntime`.
45    #[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    /// Raw reserved-header transport primitive; kept crate-private so command
51    /// identity, channel, auth and retry policy cannot be bypassed externally.
52    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(&registration_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(&registration_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                        &registration_identity,
224                        &termination,
225                    ) {
226                        emit_diagnostic(&self.diagnostic_sink, &facts);
227                        return Err(FutuError::Timeout);
228                    }
229
230                    // A different terminal observer won the generation. Keep
231                    // this waiter registered and return that winner's facts.
232                    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            // Admission publication and terminal drain share this mutex. If
246            // admission wins, facts cannot observe writer_admitted=false or
247            // drain before the command is queued. If termination wins, the
248            // lifecycle CAS rejects this generation before any publication.
249            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                            &registration_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    /// Raw detached-response transport primitive. Typed backend owners decide
414    /// which commands may legitimately proceed without awaiting the response.
415    ///
416    /// The serial remains registered until a response, connection failure, or
417    /// bounded timeout consumes it. This matches C++ `SendTCPProto_ProtoBuf`:
418    /// the business caller proceeds immediately, while the transport still
419    /// owns and drains the eventual response instead of misclassifying it as
420    /// an unmatched push.
421    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(&registration_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(&registration_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    /// Test-only transport probe used by the duplex mock-backend canary.
535    ///
536    /// This deliberately has a different name from the production primitive,
537    /// and is absent unless `test-util` is explicitly enabled.
538    #[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}