use std::sync::Arc;
use futures::future::BoxFuture;
use http::StatusCode;
use tower::BoxError;
use tower::Service;
use crate::apollo_studio_interop::UsageReporting;
use crate::compute_job::MaybeBackPressureError;
use crate::context::OPERATION_KIND;
use crate::context::OPERATION_NAME;
use crate::error::Error as RouterError;
use crate::graphql::ErrorExtension;
use crate::graphql::IntoGraphQLErrors;
use crate::query_planner::OperationKind;
use crate::services::query_parsing;
use crate::services::query_parsing::ParsedDocument;
use crate::services::supergraph;
use crate::spec::SpecError;
pub(crate) struct ParseQueryLayer {
query_parsing_service: query_parsing::BoxCloneService,
redact_query_validation_errors: bool,
}
impl ParseQueryLayer {
pub(crate) fn new(
query_parsing_service: query_parsing::BoxCloneService,
redact_query_validation_errors: bool,
) -> Self {
Self {
query_parsing_service,
redact_query_validation_errors,
}
}
}
impl<S> tower::Layer<S> for ParseQueryLayer {
type Service = ParseQueryService<S>;
fn layer(&self, inner: S) -> Self::Service {
ParseQueryService {
inner,
query_parsing_service: self.query_parsing_service.clone(),
redact_query_validation_errors: self.redact_query_validation_errors,
}
}
}
#[derive(Clone)]
pub(crate) struct ParseQueryService<S> {
inner: S,
query_parsing_service: query_parsing::BoxCloneService,
redact_query_validation_errors: bool,
}
impl<S> Service<supergraph::Request> for ParseQueryService<S>
where
S: Service<supergraph::Request, Response = supergraph::Response, Error = BoxError>
+ Clone
+ Send
+ 'static,
S::Future: Send + 'static,
{
type Response = supergraph::Response;
type Error = BoxError;
type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
fn poll_ready(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::task::ready!(self.query_parsing_service.poll_ready(cx)).map_err(|err| match err {
MaybeBackPressureError::PermanentError(err) => Box::new(err) as BoxError,
MaybeBackPressureError::TemporaryError(err) => Box::new(err) as BoxError,
})?;
self.inner.poll_ready(cx)
}
fn call(&mut self, req: supergraph::Request) -> Self::Future {
let query_parsing_service = self.query_parsing_service.clone();
let mut query_parsing_service =
std::mem::replace(&mut self.query_parsing_service, query_parsing_service);
let inner = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, inner);
let redact_query_validation_errors = self.redact_query_validation_errors;
Box::pin(async move {
let query = req.supergraph_request.body().query.as_ref();
if query.is_none() || query.unwrap().trim().is_empty() {
let errors = vec![
RouterError::builder()
.message("Must provide query string.".to_string())
.extension_code("MISSING_QUERY_STRING")
.build(),
];
return Ok(supergraph::Response::builder()
.errors(errors)
.status_code(StatusCode::BAD_REQUEST)
.context(req.context)
.build()
.expect("response is valid"));
}
let operation_name = req.supergraph_request.body().operation_name.clone();
let query = req
.supergraph_request
.body()
.query
.clone()
.expect("query presence was already checked");
match query_parsing_service
.call(query_parsing::Request::new(query, operation_name.clone()))
.await
{
Ok(doc) => {
req.context
.insert(OPERATION_NAME, doc.operation.name.clone())
.expect("cannot insert operation name into context; this is a bug");
let operation_kind = OperationKind::from(doc.operation.operation_type);
req.context
.insert(OPERATION_KIND, operation_kind)
.expect("cannot insert operation kind in the context; this is a bug");
req.context
.extensions()
.with_lock(|lock| lock.insert::<ParsedDocument>(doc));
inner.call(req).await
}
Err(MaybeBackPressureError::PermanentError(errors)) => {
let errors = if redact_query_validation_errors
&& matches!(errors, SpecError::ValidationError(_))
{
SpecError::Redacted
} else {
errors
};
req.context.extensions().with_lock(|lock| {
lock.insert(Arc::new(UsageReporting::Error(
errors.get_error_key().to_string(),
)))
});
let errors = match errors.into_graphql_errors() {
Ok(v) => v,
Err(errors) => vec![
crate::graphql::Error::builder()
.message(errors.to_string())
.extension_code(errors.extension_code())
.build(),
],
};
Ok(supergraph::Response::builder()
.errors(errors)
.status_code(StatusCode::BAD_REQUEST)
.context(req.context)
.build()
.expect("response is valid"))
}
Err(MaybeBackPressureError::TemporaryError(error)) => {
req.context.extensions().with_lock(|lock| {
let error_key =
SpecError::ValidationError(crate::error::ValidationErrors {
errors: vec![],
})
.get_error_key();
lock.insert(Arc::new(UsageReporting::Error(error_key.to_string())))
});
Ok(supergraph::Response::builder()
.error(error.to_graphql_error())
.status_code(StatusCode::SERVICE_UNAVAILABLE)
.context(req.context)
.build()
.expect("response is valid"))
}
}
})
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use http::StatusCode;
use tower::Service as _;
use tower::ServiceBuilder;
use tower::ServiceExt as _;
use super::ParseQueryLayer;
use crate::Configuration;
use crate::compute_job::MaybeBackPressureError;
use crate::context::OPERATION_KIND;
use crate::context::OPERATION_NAME;
use crate::services::OperationKind;
use crate::services::query_parsing;
use crate::services::supergraph;
const SCHEMA: &str = include_str!("../../../testing_schema.graphql");
fn downcast_mock_err(err: tower::BoxError) -> query_parsing::ServiceError {
*err.downcast()
.expect("mock should only return ServiceErrors")
}
async fn mock_parser(
mut handle: tower_test::mock::Handle<query_parsing::Request, query_parsing::ParsedDocument>,
schema: Arc<crate::spec::Schema>,
config: Arc<Configuration>,
) {
while let Some((req, responder)) = handle.next_request().await {
match crate::spec::Query::parse_document(
&req.query,
req.operation_name.as_deref(),
&schema,
&config,
) {
Ok(document) => responder.send_response(document),
Err(err) => responder.send_error(MaybeBackPressureError::PermanentError(err)),
}
}
}
#[tokio::test]
async fn it_accepts_valid_query() {
let (query_parsing_service, query_parsing_handle) =
tower_test::mock::pair::<query_parsing::Request, query_parsing::ParsedDocument>();
let query_parsing_service = ServiceBuilder::new()
.map_err(downcast_mock_err)
.service(query_parsing_service)
.boxed_clone();
let config = Arc::new(Configuration::default());
let schema = Arc::new(crate::spec::Schema::parse(SCHEMA, &config).unwrap());
let query_parsing_driver = tokio::spawn(mock_parser(query_parsing_handle, schema, config));
let (mock, mut handle) =
tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let inner_driver = tokio::spawn(async move {
let (req, responder) = handle.next_request().await.unwrap();
assert!(
req.context
.extensions()
.with_lock(|lock| lock.contains_key::<query_parsing::ParsedDocument>())
);
assert!(
req.context
.get::<_, Option<String>>(OPERATION_NAME)
.unwrap()
.is_some()
);
assert!(
req.context
.get::<_, OperationKind>(OPERATION_KIND)
.unwrap()
.is_some()
);
responder.send_response(supergraph::Response::fake_builder().build().unwrap());
});
let mut service = ServiceBuilder::new()
.layer(ParseQueryLayer::new(query_parsing_service, false))
.service(mock);
let response = service
.ready()
.await
.unwrap()
.call(
supergraph::Request::fake_builder()
.query("query { me { id } }")
.build()
.unwrap(),
)
.await
.unwrap();
assert_eq!(http::StatusCode::OK, response.response.status());
drop(service);
crate::plugin::test::await_mock_driver(query_parsing_driver).await;
crate::plugin::test::await_mock_driver(inner_driver).await;
}
#[tokio::test]
async fn it_rejects_backpressure() {
let (query_parsing_service, mut query_parsing_handle) =
tower_test::mock::pair::<query_parsing::Request, query_parsing::ParsedDocument>();
let query_parsing_service = ServiceBuilder::new()
.map_err(downcast_mock_err)
.service(query_parsing_service)
.boxed_clone();
let (mock, handle) = tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let query_parsing_driver = tokio::task::spawn(async move {
let (_req, responder) = query_parsing_handle.next_request().await.unwrap();
responder.send_error(MaybeBackPressureError::TemporaryError(
crate::compute_job::ComputeBackPressureError,
) as query_parsing::ServiceError);
});
let mut service = ServiceBuilder::new()
.layer(ParseQueryLayer::new(query_parsing_service, false))
.service(mock);
let mut response = service
.ready()
.await
.unwrap()
.call(
supergraph::Request::fake_builder()
.query("query { me { id } }")
.build()
.unwrap(),
)
.await
.unwrap();
assert_eq!(StatusCode::SERVICE_UNAVAILABLE, response.response.status());
let graphql_response = response.next_response().await.unwrap();
assert!(graphql_response.contains_error_code("REQUEST_CONCURRENCY_LIMITED"));
drop(service);
crate::plugin::test::await_mock_driver(query_parsing_driver).await;
crate::plugin::test::assert_no_mock_calls(handle).await;
}
#[tokio::test]
async fn it_rejects_missing_query() {
let (query_parsing_service, query_parsing_handle) =
tower_test::mock::pair::<query_parsing::Request, query_parsing::ParsedDocument>();
let query_parsing_service = ServiceBuilder::new()
.map_err(downcast_mock_err)
.service(query_parsing_service)
.boxed_clone();
let (mock, handle) = tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let mut service = ServiceBuilder::new()
.layer(ParseQueryLayer::new(query_parsing_service, false))
.service(mock);
let mut response = service
.ready()
.await
.unwrap()
.call(supergraph::Request::fake_builder().build().unwrap())
.await
.unwrap();
assert_eq!(StatusCode::BAD_REQUEST, response.response.status());
let graphql_response = response.next_response().await.unwrap();
assert!(graphql_response.contains_error_code("MISSING_QUERY_STRING"));
crate::plugin::test::assert_no_mock_calls(query_parsing_handle).await;
crate::plugin::test::assert_no_mock_calls(handle).await;
}
#[tokio::test]
async fn it_rejects_empty_query() {
let (query_parsing_service, query_parsing_handle) =
tower_test::mock::pair::<query_parsing::Request, query_parsing::ParsedDocument>();
let query_parsing_service = ServiceBuilder::new()
.map_err(downcast_mock_err)
.service(query_parsing_service)
.boxed_clone();
let (mock, handle) = tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let mut service = ServiceBuilder::new()
.layer(ParseQueryLayer::new(query_parsing_service, false))
.service(mock);
let mut response = service
.ready()
.await
.unwrap()
.call(
supergraph::Request::fake_builder()
.query("")
.build()
.unwrap(),
)
.await
.unwrap();
assert_eq!(StatusCode::BAD_REQUEST, response.response.status());
let graphql_response = response.next_response().await.unwrap();
assert!(graphql_response.contains_error_code("MISSING_QUERY_STRING"));
crate::plugin::test::assert_no_mock_calls(query_parsing_handle).await;
crate::plugin::test::assert_no_mock_calls(handle).await;
}
#[tokio::test]
async fn it_rejects_invalid_query() {
let (query_parsing_service, query_parsing_handle) =
tower_test::mock::pair::<query_parsing::Request, query_parsing::ParsedDocument>();
let query_parsing_service = ServiceBuilder::new()
.map_err(downcast_mock_err)
.service(query_parsing_service)
.boxed_clone();
let config = Arc::new(Configuration::default());
let schema = Arc::new(crate::spec::Schema::parse(SCHEMA, &config).unwrap());
let query_parsing_driver = tokio::spawn(mock_parser(query_parsing_handle, schema, config));
let (mock, handle) = tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let mut service = ServiceBuilder::new()
.layer(ParseQueryLayer::new(query_parsing_service, false))
.service(mock);
let mut response = service
.ready()
.await
.unwrap()
.call(
supergraph::Request::fake_builder()
.query("query Missing { doesNotExist }")
.build()
.unwrap(),
)
.await
.unwrap();
assert_eq!(StatusCode::BAD_REQUEST, response.response.status());
let graphql_response = response.next_response().await.unwrap();
assert!(graphql_response.contains_error_code("GRAPHQL_VALIDATION_FAILED"));
drop(service);
crate::plugin::test::await_mock_driver(query_parsing_driver).await;
crate::plugin::test::assert_no_mock_calls(handle).await;
}
#[tokio::test]
async fn it_redacts_validation_error() {
let (query_parsing_service, query_parsing_handle) =
tower_test::mock::pair::<query_parsing::Request, query_parsing::ParsedDocument>();
let query_parsing_service = ServiceBuilder::new()
.map_err(downcast_mock_err)
.service(query_parsing_service)
.boxed_clone();
let config = Arc::new(Configuration::default());
let schema = Arc::new(crate::spec::Schema::parse(SCHEMA, &config).unwrap());
let query_parsing_driver = tokio::spawn(mock_parser(query_parsing_handle, schema, config));
let (mock, handle) = tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let mut service = ServiceBuilder::new()
.layer(ParseQueryLayer::new(query_parsing_service, true))
.service(mock);
let mut response = service
.ready()
.await
.unwrap()
.call(
supergraph::Request::fake_builder()
.query("query Missing { doesNotExist }")
.build()
.unwrap(),
)
.await
.unwrap();
assert_eq!(StatusCode::BAD_REQUEST, response.response.status());
let graphql_response = response.next_response().await.unwrap();
assert!(graphql_response.contains_error_code("UNKNOWN_ERROR"));
assert!(
!serde_json::to_string(&graphql_response)
.unwrap()
.contains("doesNotExist")
);
drop(service);
crate::plugin::test::await_mock_driver(query_parsing_driver).await;
crate::plugin::test::assert_no_mock_calls(handle).await;
}
}