use std::pin::Pin;
use std::task::{Context, Poll};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use bytes::Bytes;
use futures_core::Stream;
use reqwest::multipart;
use crate::error::{Error, Result};
use super::{Client, MultipartOptions, RequestOptions};
pub const WIRE_TARGET: &str = "xdk::wire";
impl Client {
pub async fn send_request(&self, options: &RequestOptions) -> Result<serde_json::Value> {
self.send_request_with(options, self.request_timeout())
.await
}
pub(crate) async fn send_request_with(
&self,
options: &RequestOptions,
timeout: std::time::Duration,
) -> Result<serde_json::Value> {
match self.send_request_once(options, timeout).await {
Err(error) => match self.rate_limit_pause(&error) {
Some(pause) => {
tokio::time::sleep(pause).await;
self.send_request_once(options, timeout).await
}
None => Err(error),
},
sent => sent,
}
}
async fn send_request_once(
&self,
options: &RequestOptions,
timeout: std::time::Duration,
) -> Result<serde_json::Value> {
let method = options.method.to_uppercase();
let method = if method.is_empty() { "GET" } else { &method };
let url = self.build_url(&options.target)?;
let req_method = reqwest::Method::from_bytes(method.as_bytes())
.map_err(|_| Error::InvalidMethod(method.to_string()))?;
let mut builder = self.http().request(req_method, &url).timeout(timeout);
let sends_body =
!options.data.is_empty() && matches!(method, "POST" | "PUT" | "PATCH" | "DELETE");
if sends_body {
builder = with_body(builder, options);
}
let builder = self.with_headers(builder, options, sends_body).await?;
trace_request(method, &url);
let resp = builder.send().await?;
self.record_rate_limit(resp.headers());
trace_response(resp.status(), resp.headers());
read_response(resp).await
}
pub async fn send_multipart_request(
&self,
options: &MultipartOptions,
) -> Result<serde_json::Value> {
match self.send_multipart_once(options).await {
Err(error) => match self.rate_limit_pause(&error) {
Some(pause) => {
tokio::time::sleep(pause).await;
self.send_multipart_once(options).await
}
None => Err(error),
},
sent => sent,
}
}
async fn send_multipart_once(&self, options: &MultipartOptions) -> Result<serde_json::Value> {
let method = options.request.method.to_uppercase();
let method = if method.is_empty() { "POST" } else { &method };
let url = self.build_url(&options.request.target)?;
let req_method = reqwest::Method::from_bytes(method.as_bytes())
.map_err(|_| Error::InvalidMethod(method.to_string()))?;
let mut form = multipart::Form::new();
if !options.file_field.is_empty() && !options.file_path.is_empty() {
let part = multipart::Part::file(&options.file_path)
.await
.map_err(|e| Error::io(format!("error opening file: {e}")).with_source(e))?;
form = form.part(options.file_field.clone(), part);
} else if !options.file_field.is_empty() && !options.file_data.is_empty() {
let part = multipart::Part::bytes(options.file_data.clone())
.file_name(options.file_name.clone());
form = form.part(options.file_field.clone(), part);
}
for (key, value) in &options.form_fields {
form = form.text(key.clone(), value.clone());
}
let builder = self
.http()
.request(req_method, &url)
.timeout(self.request_timeout())
.multipart(form);
let builder = self.with_headers(builder, &options.request, false).await?;
trace_request(method, &url);
let resp = builder.send().await?;
self.record_rate_limit(resp.headers());
read_response(resp).await
}
pub async fn stream_request(&self, options: &RequestOptions) -> Result<StreamLines> {
let method = options.method.to_uppercase();
let method = if method.is_empty() { "GET" } else { &method };
let url = self.build_url(&options.target)?;
let req_method = reqwest::Method::from_bytes(method.as_bytes())
.map_err(|_| Error::InvalidMethod(method.to_string()))?;
let mut builder = self.http().request(req_method, &url);
let sends_body = !options.data.is_empty();
if sends_body {
builder = with_body(builder, options);
}
let builder = self.with_headers(builder, options, sends_body).await?;
trace_request(method, &url);
let resp = builder.send().await?;
self.record_rate_limit(resp.headers());
trace_response(resp.status(), resp.headers());
if resp.status().as_u16() >= 400 {
return Err(api_error(resp).await?);
}
Ok(StreamLines {
body: Box::pin(resp.bytes_stream()),
buf: Vec::new(),
done: false,
})
}
async fn with_headers(
&self,
mut builder: reqwest::RequestBuilder,
options: &RequestOptions,
sets_content_type: bool,
) -> Result<reqwest::RequestBuilder> {
for header in &options.headers {
if let Some((key, value)) = header.split_once(':') {
builder = builder.header(key.trim(), value.trim());
}
}
if !options.no_auth && !user_supplied_header(&options.headers, "Authorization") {
let auth_header = self.get_auth_header(options).await?;
builder = builder.header("Authorization", auth_header);
}
if !user_supplied_header(&options.headers, "User-Agent") {
builder = builder.header("User-Agent", self.inner.user_agent.as_str());
}
if options.trace && !user_supplied_header(&options.headers, "X-B3-Flags") {
builder = builder.header("X-B3-Flags", "1");
}
note_header_overrides(
&options.headers,
sets_content_type,
!options.no_auth,
options.trace,
);
Ok(builder)
}
fn rate_limit_pause(&self, error: &Error) -> Option<Duration> {
let max_wait = self.inner.rate_limit_max_wait?;
let Error::Api {
status: 429,
reset_at: Some(reset_at),
..
} = error
else {
return None;
};
let reset = UNIX_EPOCH + Duration::from_secs(*reset_at);
let wait = reset.duration_since(SystemTime::now()).unwrap_or_default();
(wait <= max_wait).then_some(wait)
}
}
fn with_body(
mut builder: reqwest::RequestBuilder,
options: &RequestOptions,
) -> reqwest::RequestBuilder {
if !user_supplied_header(&options.headers, "Content-Type") {
let content_type = if serde_json::from_str::<serde_json::Value>(&options.data).is_ok() {
"application/json"
} else {
"application/x-www-form-urlencoded"
};
builder = builder.header("Content-Type", content_type);
}
builder.body(options.data.clone())
}
async fn read_response(resp: reqwest::Response) -> Result<serde_json::Value> {
if resp.status().as_u16() >= 400 {
return Err(api_error(resp).await?);
}
let body = resp.text().await?;
if body.is_empty() {
return Ok(serde_json::json!({}));
}
Ok(serde_json::from_str(&body).unwrap_or(serde_json::Value::String(body)))
}
async fn api_error(resp: reqwest::Response) -> Result<Error> {
let status = resp.status().as_u16();
let reset_at = (status == 429)
.then(|| super::RateLimit::from_headers(resp.headers()))
.flatten()
.and_then(|window| window.reset_at);
let body = resp.text().await?;
let body = match serde_json::from_str::<serde_json::Value>(&body) {
Ok(json) => json.to_string(),
Err(_) if body.is_empty() => "{}".to_string(),
Err(_) => body,
};
Ok(Error::Api {
status,
body,
reset_at,
})
}
type ByteStream = Pin<Box<dyn Stream<Item = reqwest::Result<Bytes>> + Send>>;
pub struct StreamLines {
body: ByteStream,
buf: Vec<u8>,
done: bool,
}
impl StreamLines {
pub async fn next_line(&mut self) -> Result<Option<String>> {
std::future::poll_fn(|cx| Pin::new(&mut *self).poll_next(cx))
.await
.transpose()
}
fn take_line(&mut self) -> Option<String> {
let end = self.buf.iter().position(|&b| b == b'\n')?;
let mut line: Vec<u8> = self.buf.drain(..=end).collect();
line.pop();
if line.last() == Some(&b'\r') {
line.pop();
}
Some(String::from_utf8_lossy(&line).into_owned())
}
}
impl std::fmt::Debug for StreamLines {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StreamLines").finish_non_exhaustive()
}
}
impl Stream for StreamLines {
type Item = Result<String>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
if let Some(line) = self.take_line() {
if line.is_empty() {
continue;
}
return Poll::Ready(Some(Ok(line)));
}
if self.done {
let rest = std::mem::take(&mut self.buf);
let line = String::from_utf8_lossy(&rest);
let line = line.trim_end_matches('\r').to_string();
return Poll::Ready((!line.is_empty()).then_some(Ok(line)));
}
match self.body.as_mut().poll_next(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(None) => self.done = true,
Poll::Ready(Some(Ok(chunk))) => self.buf.extend_from_slice(&chunk),
Poll::Ready(Some(Err(e))) => {
self.done = true;
return Poll::Ready(Some(Err(Error::io(e.to_string()).with_source(e))));
}
}
}
}
}
fn user_supplied_header(headers: &[String], name: &str) -> bool {
headers
.iter()
.filter_map(|h| h.split_once(':'))
.any(|(key, _)| key.trim().eq_ignore_ascii_case(name))
}
fn note_header_overrides(
headers: &[String],
would_set_content_type: bool,
would_set_auth: bool,
would_set_trace: bool,
) {
let candidates: [(&str, bool); 4] = [
("Content-Type", would_set_content_type),
("Authorization", would_set_auth),
("User-Agent", true),
("X-B3-Flags", would_set_trace),
];
for (name, xurl_wanted) in candidates {
if xurl_wanted && user_supplied_header(headers, name) {
tracing::debug!(
target: WIRE_TARGET,
kind = "note",
header = name,
"user-supplied header detected; the default is not appended"
);
}
}
}
fn trace_request(method: &str, url: &str) {
tracing::debug!(target: WIRE_TARGET, kind = "request", method, url);
}
fn trace_response(status: reqwest::StatusCode, headers: &reqwest::header::HeaderMap) {
tracing::debug!(target: WIRE_TARGET, kind = "status", status = %status);
for (key, value) in headers {
tracing::debug!(
target: WIRE_TARGET,
kind = "header",
name = %key,
value = value.to_str().unwrap_or("")
);
}
tracing::debug!(target: WIRE_TARGET, kind = "end");
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn user_supplied_header_detects_canonical_case() {
let headers = vec!["Authorization: Bearer foo".to_string()];
assert!(user_supplied_header(&headers, "Authorization"));
}
#[test]
fn user_supplied_header_is_case_insensitive_on_input_key() {
for raw in [
"authorization: Bearer foo",
"AUTHORIZATION: Bearer foo",
"aUtHoRiZaTiOn: Bearer foo",
] {
assert!(
user_supplied_header(&[raw.to_string()], "Authorization"),
"did not detect: {raw}"
);
}
}
#[test]
fn user_supplied_header_is_case_insensitive_on_query_name() {
let headers = vec!["Authorization: Bearer foo".to_string()];
for query in ["authorization", "AUTHORIZATION", "aUtHoRiZaTiOn"] {
assert!(
user_supplied_header(&headers, query),
"did not detect with query: {query}"
);
}
}
#[test]
fn user_supplied_header_ignores_surrounding_whitespace_on_key() {
let headers = vec![" Authorization : Bearer foo".to_string()];
assert!(user_supplied_header(&headers, "Authorization"));
}
#[test]
fn user_supplied_header_false_for_empty_list() {
let headers: Vec<String> = Vec::new();
assert!(!user_supplied_header(&headers, "Authorization"));
assert!(!user_supplied_header(&headers, "User-Agent"));
}
#[test]
fn user_supplied_header_does_not_match_substring_keys() {
let headers = vec![
"Cookie: session=abc".to_string(),
"X-Authorization-Hint: ignored".to_string(),
];
assert!(!user_supplied_header(&headers, "Authorization"));
}
#[test]
fn user_supplied_header_false_for_unparseable_entry() {
let headers = vec!["malformed-no-colon".to_string()];
assert!(!user_supplied_header(&headers, "Authorization"));
}
#[test]
fn user_supplied_header_detects_each_xurl_added_header() {
let headers = vec![
"Content-Type: application/xml".to_string(),
"Authorization: Bearer foo".to_string(),
"User-Agent: custom/1.0".to_string(),
"X-B3-Flags: 0".to_string(),
];
assert!(user_supplied_header(&headers, "Content-Type"));
assert!(user_supplied_header(&headers, "Authorization"));
assert!(user_supplied_header(&headers, "User-Agent"));
assert!(user_supplied_header(&headers, "X-B3-Flags"));
}
}