use std::net::SocketAddr;
use std::sync::Arc;
use http::StatusCode;
use opentelemetry::global;
use opentelemetry::trace::TraceContextExt;
use tower::ServiceBuilder;
use tower::load_shed::error::Overloaded;
use tower_http::trace::MakeSpan;
use tracing::Span;
use crate::Context;
use crate::graphql;
use crate::layers::ServiceBuilderExt;
use crate::plugins::telemetry::consts::OTEL_STATUS_CODE;
use crate::plugins::telemetry::consts::OTEL_STATUS_CODE_ERROR;
use crate::plugins::telemetry::pipeline_bypass::record_bypassed_request;
use crate::plugins::telemetry::span_factory;
use crate::services::router;
use crate::uplink::license_enforcement::LICENSE_EXPIRED_SHORT_MESSAGE;
use crate::uplink::license_enforcement::LicenseState;
#[derive(Clone, Default)]
pub(crate) struct PropagatingMakeSpan {
pub(crate) license: Arc<LicenseState>,
}
impl<B> MakeSpan<B> for PropagatingMakeSpan {
fn make_span(&mut self, request: &http::Request<B>) -> Span {
let context = global::get_text_map_propagator(|propagator| {
propagator.extract(&opentelemetry_http::HeaderExtractor(request.headers()))
});
let span = if context.span().span_context().is_valid()
|| context.span().span_context().trace_id() != opentelemetry::trace::TraceId::INVALID
{
let _context_guard = context.attach();
span_factory::create_router(request)
} else {
span_factory::create_router(request)
};
if matches!(
&*self.license,
LicenseState::LicensedWarn { limits: _ } | LicenseState::LicensedHalt { limits: _ }
) {
span.record(OTEL_STATUS_CODE, OTEL_STATUS_CODE_ERROR);
span.record("apollo_router.license", LICENSE_EXPIRED_SHORT_MESSAGE);
}
span
}
}
#[derive(Clone)]
pub(crate) struct ConnectionInfo {
pub(crate) peer_address: Option<SocketAddr>,
pub(crate) server_address: Option<SocketAddr>,
}
pub(crate) type ConnectionRouterService =
tower::util::BoxCloneSyncService<router::Request, router::Response, tower::BoxError>;
pub(crate) fn connection_router_service(
service: router::BoxCloneService,
) -> ConnectionRouterService {
ConnectionRouterService::new(
ServiceBuilder::new()
.map_future_with_request_data(|req: &router::Request| req.context.clone(), shed_as_503)
.load_shed()
.buffered()
.service(service),
)
}
async fn shed_as_503(
context: Context,
future: impl Future<Output = Result<router::Response, tower::BoxError>>,
) -> Result<router::Response, tower::BoxError> {
match future.await {
Err(err) if err.is::<Overloaded>() => {
record_bypassed_request(
StatusCode::SERVICE_UNAVAILABLE.as_u16(),
context.created_at.elapsed(),
);
tracing::debug!(
code = "REQUEST_CONCURRENCY_LIMITED",
"the connection queue is full, shedding request",
);
Ok(router::Response::error_builder()
.status_code(StatusCode::SERVICE_UNAVAILABLE)
.error(
graphql::Error::builder()
.message("Your request has been concurrency limited waiting for the router")
.extension_code("REQUEST_CONCURRENCY_LIMITED")
.build(),
)
.context(context)
.build()
.expect("overloaded response should build"))
}
other => other,
}
}
#[cfg(test)]
#[derive(Clone)]
pub(crate) struct NeverReady;
#[cfg(test)]
impl tower::Service<router::Request> for NeverReady {
type Response = router::Response;
type Error = tower::BoxError;
type Future = std::future::Pending<Result<router::Response, tower::BoxError>>;
fn poll_ready(
&mut self,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::task::Poll::Pending
}
fn call(&mut self, _req: router::Request) -> Self::Future {
unreachable!("load_shed should short-circuit calls to a never-ready service")
}
}
#[cfg(test)]
mod tests {
use tower::Service;
use super::*;
use crate::metrics::FutureMetricsExt;
#[tokio::test]
async fn connection_router_service_always_reports_ready() {
let mut connection_service =
connection_router_service(router::BoxCloneService::new(NeverReady));
let ready = futures::future::poll_fn(|cx| {
std::task::Poll::Ready(connection_service.poll_ready(cx))
})
.await;
assert!(
matches!(ready, std::task::Poll::Ready(Ok(()))),
"poll_ready must never report Pending to hyper"
);
}
#[tokio::test]
async fn shed_requests_answer_with_a_503_graphql_error() {
let response = shed_as_503(
Context::new(),
std::future::ready(Err(Overloaded::new().into())),
)
.await
.expect("a shed request should answer, not error");
assert_eq!(response.response.status(), StatusCode::SERVICE_UNAVAILABLE);
let body = router::body::into_bytes(response.response.into_body())
.await
.expect("the body should be readable");
let body: graphql::Response =
serde_json::from_slice(&body).expect("the body should be a GraphQL response");
let error = body.errors.first().expect("a GraphQL error is expected");
assert_eq!(
error.extensions.get("code").and_then(|code| code.as_str()),
Some("REQUEST_CONCURRENCY_LIMITED")
);
assert_eq!(
error.message,
"Your request has been concurrency limited waiting for the router"
);
}
#[tokio::test]
async fn shed_requests_record_a_request_duration() {
async {
let _ = shed_as_503(
Context::new(),
std::future::ready(Err(Overloaded::new().into())),
)
.await;
assert_histogram_count!(
"http.server.request.duration",
1,
"http.response.status_code" = 503i64
);
}
.with_metrics()
.await;
}
#[tokio::test]
async fn other_errors_pass_through() {
let err = shed_as_503(
Context::new(),
std::future::ready(Err("something else".into())),
)
.await
.expect_err("a non-overload error should stay an error");
assert_eq!(err.to_string(), "something else");
}
}