use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use apollo_opentelemetry::metrics::Clock;
use bytes::Bytes;
use http_body::Body;
use http_body_util::BodyExt as _;
use http_body_util::combinators::UnsyncBoxBody;
use opentelemetry::{Context as OtelContext, KeyValue};
use opentelemetry_semantic_conventions::attribute as semconv;
use tower::util::BoxCloneSyncService;
use tower::{BoxError, Service};
use crate::builder::HttpClientBuilder;
use crate::config::HttpClientConfig;
use crate::error::HttpClientError;
use crate::metrics::{HttpMetrics, RequestMetrics};
use crate::protocol::{HostSender, HttpBody, Origin, ProtocolVersionCell};
use crate::spans::HttpSpans;
use apollo_http_shared::body::{BodySpanAttr, CountingBody};
pub(crate) struct HttpClientState {
pub(crate) pools: HashMap<Origin, HostSender>,
pub(crate) make_sender: Box<dyn Fn(Origin) -> HostSender + Send + Sync>,
}
impl HttpClientState {
fn get_or_create(&mut self, origin: Origin) -> HostSender {
self.pools
.entry(origin)
.or_insert_with_key(|key| (self.make_sender)(key.clone()))
.clone()
}
}
fn extract_pool_key(uri: &http::Uri) -> Result<Origin, HttpClientError> {
match (uri.scheme().cloned(), uri.authority().cloned()) {
(Some(scheme), Some(authority)) => Ok(Origin { scheme, authority }),
_ => {
let err: http::uri::InvalidUri = http::Uri::try_from("").unwrap_err();
Err(HttpClientError::InvalidUri { source: err.into() })
}
}
}
#[cfg(not(unix))]
pub(crate) fn unsupported_scheme_sender() -> HostSender {
Arc::new(|req, _| {
let scheme = req.uri().scheme_str().unwrap_or("").to_owned();
Box::pin(async move { Err(HttpClientError::UnsupportedScheme { scheme }) })
})
}
struct ClientGuards {
_span_cx: OtelContext,
_req_metrics: RequestMetrics,
}
#[derive(Clone)]
pub(crate) struct HttpClientDispatch {
pub(crate) state: Arc<Mutex<HttpClientState>>,
pub(crate) metrics: HttpMetrics,
pub(crate) spans: HttpSpans,
pub(crate) configured_protocol_version: http::Version,
}
type HttpClientDispatchFuture = Pin<
Box<
dyn Future<Output = Result<http::Response<UnsyncBoxBody<Bytes, BoxError>>, HttpClientError>>
+ Send,
>,
>;
impl Service<http::Request<HttpBody>> for HttpClientDispatch {
type Response = http::Response<UnsyncBoxBody<Bytes, BoxError>>;
type Error = HttpClientError;
type Future = HttpClientDispatchFuture;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, req: http::Request<HttpBody>) -> Self::Future {
let state = self.state.clone();
let metrics = self.metrics.clone();
let spans = self.spans.clone();
let protocol_version = ProtocolVersionCell::new(self.configured_protocol_version);
Box::pin(async move {
let span_cx = OtelContext::current();
let key = extract_pool_key(req.uri())?;
let per_host = state.lock().unwrap().get_or_create(key.clone());
let method = req.method().clone();
let (mut req_metrics, req_body_size) =
metrics.begin_request(&method, &key, protocol_version.clone());
let req_span_attr = spans.request_body_size().then(|| {
let attr = BodySpanAttr::new(semconv::HTTP_REQUEST_BODY_SIZE);
attr.bind(span_cx.clone());
attr
});
let req = req.map(|body| {
CountingBody::new(
body,
req_body_size.clone(),
(),
protocol_version.clone(),
req_span_attr,
)
.boxed()
});
match per_host(req, protocol_version.clone()).await {
Ok(resp) => {
let status = resp.status();
req_metrics.on_success(status);
req_body_size.set(KeyValue::new(
semconv::HTTP_RESPONSE_STATUS_CODE,
i64::from(status.as_u16()),
));
spans.on_response(protocol_version.get(), &key, status, resp.headers());
let resp_body_size = metrics.begin_response(&method, &key, status);
let resp_span_attr = spans.response_body_size().then(|| {
let attr = BodySpanAttr::new(semconv::HTTP_RESPONSE_BODY_SIZE);
attr.bind(span_cx.clone());
attr
});
let guards = ClientGuards {
_span_cx: span_cx,
_req_metrics: req_metrics,
};
let response = resp.map(|body| {
CountingBody::new(
body,
resp_body_size,
guards,
protocol_version.clone(),
resp_span_attr,
)
.boxed_unsync()
});
Ok(response)
}
Err(e) => {
req_metrics.on_error();
req_body_size.set(KeyValue::new(semconv::ERROR_TYPE, "_OTHER"));
spans.on_error(&key, &e.to_string());
Err(e)
}
}
})
}
}
pub struct HttpClient {
inner: BoxCloneSyncService<
http::Request<HttpBody>,
http::Response<UnsyncBoxBody<Bytes, BoxError>>,
HttpClientError,
>,
}
impl Clone for HttpClient {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
impl HttpClient {
pub fn new(config: &HttpClientConfig) -> Result<Self, HttpClientError> {
HttpClientBuilder::new(config.clone()).build()
}
#[doc(hidden)]
pub fn with_clock(config: &HttpClientConfig, clock: Clock) -> Result<Self, HttpClientError> {
HttpClientBuilder::new(config.clone()).build_with_clock(clock)
}
pub fn builder(config: HttpClientConfig) -> HttpClientBuilder {
HttpClientBuilder::new(config)
}
pub(crate) fn from_service(
inner: BoxCloneSyncService<
http::Request<HttpBody>,
http::Response<UnsyncBoxBody<Bytes, BoxError>>,
HttpClientError,
>,
) -> Self {
Self { inner }
}
}
impl<B> Service<http::Request<B>> for HttpClient
where
B: Body<Data = Bytes> + Send + Sync + 'static,
B::Error: Into<BoxError> + Send + Sync + 'static,
{
type Response = http::Response<UnsyncBoxBody<Bytes, BoxError>>;
type Error = HttpClientError;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: http::Request<B>) -> Self::Future {
let req = req.map(|body| body.map_err(Into::into).boxed());
self.inner.call(req)
}
}
const _: fn() = || {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<HttpClient>();
};
#[cfg(test)]
mod tests {
use super::*;
use apollo_opentelemetry_test::{TelemetryContext, assert_metrics_snapshot};
use bytes::Bytes;
use http_body_util::Full;
use tower::ServiceExt as _;
#[test]
fn extract_pool_key_valid_uri() {
let uri: http::Uri = "https://example.com/path".parse().unwrap();
let Origin { scheme, authority } = extract_pool_key(&uri).unwrap();
assert_eq!(scheme, http::uri::Scheme::HTTPS);
assert_eq!(authority.host(), "example.com");
}
#[test]
fn extract_pool_key_relative_uri_returns_invalid_uri_error() {
let uri: http::Uri = "/path/only".parse().unwrap();
let err = extract_pool_key(&uri).unwrap_err();
assert!(matches!(err, HttpClientError::InvalidUri { .. }));
}
#[test]
fn new_with_default_config_succeeds() {
let config = HttpClientConfig::default();
assert!(HttpClient::new(&config).is_ok());
}
#[test]
fn new_returns_invalid_config_when_proxy_url_skips_validation() {
use crate::config::{ProxyConfig, ProxyUrl};
let url: apollo_configuration::types::Url =
"file:///not-a-network-url".parse().expect("parse");
let config = HttpClientConfig {
proxy: Some(ProxyConfig {
url: ProxyUrl::new(url),
}),
..HttpClientConfig::default()
};
let err = HttpClient::new(&config)
.err()
.expect("hostless proxy URL must be rejected");
assert!(
matches!(err, HttpClientError::InvalidConfig { .. }),
"expected InvalidConfig, got {err:?}"
);
}
#[tokio::test]
async fn records_error_metric_on_failed_request() {
let ctx = TelemetryContext::new();
let config: HttpClientConfig =
apollo_configuration::parse_yaml("connect_timeout: 1s", &Default::default())
.expect("valid config");
let (clock, _mock) = Clock::mock();
let client = HttpClient::with_clock(&config, clock).expect("valid config");
let req = http::Request::builder()
.method(http::Method::GET)
.uri("http://127.0.0.1:1")
.body(
Full::new(Bytes::new())
.map_err(|never: std::convert::Infallible| match never {})
.boxed(),
)
.unwrap();
let _ = client.oneshot(req).await;
assert_metrics_snapshot!(ctx, @r#"
- name: http.client.active_requests
description: Number of HTTP requests currently in flight
unit: "{request}"
data:
type: Sum
data_points:
- attributes:
http.request.method: GET
server.address: 127.0.0.1
server.port: "1"
value: 0
is_monotonic: false
temporality: Cumulative
- name: http.client.request.duration
description: Duration of HTTP client requests
unit: s
data:
type: Histogram
data_points:
- attributes:
error.type: _OTHER
http.request.method: GET
network.protocol.version: "1.1"
server.address: 127.0.0.1
server.port: "1"
count: 1
sum: 0
min: 0
max: 0
bounds:
- 0.005
- 0.01
- 0.025
- 0.05
- 0.075
- 0.1
- 0.25
- 0.5
- 0.75
- 1
- 2.5
- 5
- 7.5
- 10
bucket_counts:
- 1
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
temporality: Cumulative
"#);
}
}