Skip to main content

relay_knowledge/net/http/
qos_client.rs

1use std::{error::Error, fmt};
2
3use serde::de::DeserializeOwned;
4
5use crate::net::{
6    http::qos_request_context_active,
7    qos::{QosPermit, QosPolicy, QosRuntime, RejectReason},
8};
9
10/// Error raised by QoS-gated outbound reqwest calls.
11#[derive(Debug)]
12pub enum QosHttpClientError {
13    QosRejected(RejectReason),
14    Transport(reqwest::Error),
15}
16
17impl QosHttpClientError {
18    /// Returns whether the transport layer reported a timeout.
19    pub fn is_timeout(&self) -> bool {
20        matches!(self, Self::Transport(error) if error.is_timeout())
21    }
22
23    /// Returns whether outbound admission rejected the request before I/O.
24    pub fn is_qos_rejected(&self) -> bool {
25        matches!(self, Self::QosRejected(_))
26    }
27}
28
29impl fmt::Display for QosHttpClientError {
30    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
31        match self {
32            Self::QosRejected(reason) => {
33                write!(formatter, "request rejected by QoS: {}", reason.as_str())
34            }
35            Self::Transport(error) => error.fmt(formatter),
36        }
37    }
38}
39
40impl Error for QosHttpClientError {}
41
42/// Reqwest response that keeps the QoS request permit until the body is consumed.
43pub struct QosHttpResponse {
44    inner: Option<reqwest::Response>,
45    qos: QosRuntime,
46    _permit: Option<QosPermit>,
47    cancellation: CancellationGuard,
48}
49
50impl QosHttpResponse {
51    fn with_permit(
52        inner: reqwest::Response,
53        qos: QosRuntime,
54        permit: QosPermit,
55        cancellation: CancellationGuard,
56    ) -> Self {
57        Self {
58            inner: Some(inner),
59            qos,
60            _permit: Some(permit),
61            cancellation,
62        }
63    }
64
65    fn without_permit(
66        inner: reqwest::Response,
67        qos: QosRuntime,
68        cancellation: CancellationGuard,
69    ) -> Self {
70        Self {
71            inner: Some(inner),
72            qos,
73            _permit: None,
74            cancellation,
75        }
76    }
77
78    pub fn status(&self) -> reqwest::StatusCode {
79        self.inner.as_ref().expect("response is available").status()
80    }
81
82    pub fn content_length(&self) -> Option<u64> {
83        self.inner
84            .as_ref()
85            .expect("response is available")
86            .content_length()
87    }
88
89    pub async fn json<T>(mut self) -> Result<T, reqwest::Error>
90    where
91        T: DeserializeOwned,
92    {
93        let result = self
94            .inner
95            .take()
96            .expect("response is available")
97            .json::<T>()
98            .await;
99        self.cancellation.complete();
100        record_body_timeout(&self.qos, result)
101    }
102
103    pub async fn text(mut self) -> Result<String, reqwest::Error> {
104        let result = self
105            .inner
106            .take()
107            .expect("response is available")
108            .text()
109            .await;
110        self.cancellation.complete();
111        record_body_timeout(&self.qos, result)
112    }
113
114    pub async fn bytes(mut self) -> Result<Vec<u8>, reqwest::Error> {
115        let result = self
116            .inner
117            .take()
118            .expect("response is available")
119            .bytes()
120            .await
121            .map(|bytes| bytes.to_vec());
122        self.cancellation.complete();
123        record_body_timeout(&self.qos, result)
124    }
125
126    pub async fn chunk(&mut self) -> Result<Option<Vec<u8>>, reqwest::Error> {
127        let result = self
128            .inner
129            .as_mut()
130            .expect("response is available")
131            .chunk()
132            .await
133            .map(|chunk| chunk.map(|bytes| bytes.to_vec()));
134        if !matches!(&result, Ok(Some(_))) {
135            self.cancellation.complete();
136        }
137        record_body_timeout(&self.qos, result)
138    }
139}
140
141struct CancellationGuard {
142    qos: QosRuntime,
143    completed: bool,
144}
145
146impl CancellationGuard {
147    fn new(qos: QosRuntime) -> Self {
148        Self {
149            qos,
150            completed: false,
151        }
152    }
153
154    fn complete(&mut self) {
155        self.completed = true;
156    }
157}
158
159impl Drop for CancellationGuard {
160    fn drop(&mut self) {
161        if !self.completed {
162            self.qos.record_cancelled();
163        }
164    }
165}
166
167/// Sends an outbound reqwest request after acquiring a QoS request permit.
168pub async fn send_request_with_qos(
169    qos: &QosRuntime,
170    policy: &QosPolicy,
171    request: reqwest::RequestBuilder,
172) -> Result<QosHttpResponse, QosHttpClientError> {
173    if qos_request_context_active() {
174        return send_request_without_new_permit(qos, request).await;
175    }
176
177    let permit = qos
178        .admit_request(policy)
179        .map_err(QosHttpClientError::QosRejected)?;
180    let mut cancellation = CancellationGuard::new(qos.clone());
181    match request.send().await {
182        Ok(response) => Ok(QosHttpResponse::with_permit(
183            response,
184            qos.clone(),
185            permit,
186            cancellation,
187        )),
188        Err(error) => {
189            cancellation.complete();
190            if error.is_timeout() {
191                qos.record_timed_out();
192            }
193            Err(QosHttpClientError::Transport(error))
194        }
195    }
196}
197
198async fn send_request_without_new_permit(
199    qos: &QosRuntime,
200    request: reqwest::RequestBuilder,
201) -> Result<QosHttpResponse, QosHttpClientError> {
202    let mut cancellation = CancellationGuard::new(qos.clone());
203    match request.send().await {
204        Ok(response) => Ok(QosHttpResponse::without_permit(
205            response,
206            qos.clone(),
207            cancellation,
208        )),
209        Err(error) => {
210            cancellation.complete();
211            if error.is_timeout() {
212                qos.record_timed_out();
213            }
214            Err(QosHttpClientError::Transport(error))
215        }
216    }
217}
218
219fn record_body_timeout<T>(
220    qos: &QosRuntime,
221    result: Result<T, reqwest::Error>,
222) -> Result<T, reqwest::Error> {
223    if matches!(&result, Err(error) if error.is_timeout()) {
224        qos.record_timed_out();
225    }
226    result
227}
228
229#[cfg(test)]
230#[path = "qos_client_tests.rs"]
231mod qos_client_tests;