use std::collections::HashMap;
use std::sync::Arc;
use futures::future::BoxFuture;
use http::StatusCode;
use tower::BoxError;
use tower::Service;
use crate::graphql::Error;
use crate::plugins::authorization::AUTHENTICATION_REQUIRED_KEY;
use crate::plugins::authorization::AuthorizationPlugin;
use crate::plugins::authorization::CacheKeyMetadata;
use crate::plugins::authorization::REQUIRED_POLICIES_KEY;
use crate::plugins::authorization::REQUIRED_SCOPES_KEY;
use crate::services::query_parsing::ParsedDocument;
use crate::services::supergraph;
use crate::spec::Schema;
pub(crate) struct ExtractAuthorizationChecksLayer {
schema: Arc<Schema>,
}
impl ExtractAuthorizationChecksLayer {
pub(crate) fn new(schema: Arc<Schema>) -> Self {
Self { schema }
}
}
impl<S> tower::Layer<S> for ExtractAuthorizationChecksLayer {
type Service = ExtractAuthorizationChecksService<S>;
fn layer(&self, inner: S) -> Self::Service {
ExtractAuthorizationChecksService {
inner,
schema: self.schema.clone(),
}
}
}
#[derive(Clone)]
pub(crate) struct ExtractAuthorizationChecksService<S> {
inner: S,
schema: Arc<Schema>,
}
impl<S> Service<supergraph::Request> for ExtractAuthorizationChecksService<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>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: supergraph::Request) -> Self::Future {
let inner = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, inner);
let schema = self.schema.clone();
Box::pin(async move {
let doc = req
.context
.extensions()
.with_lock(|lock| lock.get::<ParsedDocument>().cloned());
let doc = match doc {
Some(doc) => doc,
None => {
let errors = vec![
Error::builder()
.message("Cannot find executable document".to_string())
.extension_code("MISSING_EXECUTABLE_DOCUMENT")
.build(),
];
return Ok(supergraph::Response::builder()
.errors(errors)
.status_code(StatusCode::INTERNAL_SERVER_ERROR)
.context(req.context)
.build()
.expect("response is valid"));
}
};
let operation_name = req.supergraph_request.body().operation_name.as_deref();
let CacheKeyMetadata {
is_authenticated,
scopes,
policies,
} = AuthorizationPlugin::generate_cache_metadata(
&doc.executable,
operation_name,
schema.supergraph_schema(),
false,
);
if is_authenticated {
req.context
.insert(AUTHENTICATION_REQUIRED_KEY, true)
.unwrap();
}
if !scopes.is_empty() {
req.context.insert(REQUIRED_SCOPES_KEY, scopes).unwrap();
}
if !policies.is_empty() {
let policies: HashMap<String, Option<bool>> =
policies.into_iter().map(|policy| (policy, None)).collect();
req.context.insert(REQUIRED_POLICIES_KEY, policies).unwrap();
}
inner.call(req).await
})
}
}
#[cfg(test)]
mod tests {
use std::collections::HashSet;
use std::sync::Arc;
use tower::ServiceBuilder;
use tower::ServiceExt as _;
use super::*;
use crate::Configuration;
use crate::Context;
use crate::plugins::authorization::REQUIRED_SCOPES_KEY;
use crate::spec::Query;
const REQUIRES_SCOPES_SCHEMA: &str =
include_str!("../../../tests/fixtures/supergraph-auth.graphql");
const POLICY_SCHEMA: &str =
include_str!("../../../tests/fixtures/directives/policy/policy_basic_schema.graphql");
const AUTHENTICATED_SCHEMA: &str =
include_str!("../../../tests/integration/fixtures/authenticated_directive.graphql");
#[tokio::test]
async fn extracts_scopes() {
let query = "query { me { id name } }";
let config = Configuration::default();
let schema = Arc::new(Schema::parse(REQUIRES_SCOPES_SCHEMA, &config).unwrap());
let doc = Query::parse_document(query, None, &schema, &config).unwrap();
let (mock, mut handle) =
tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let driver = tokio::spawn(async move {
let (req, responder) = handle.next_request().await.unwrap();
responder.send_response(
supergraph::Response::fake_builder()
.context(req.context)
.build()
.unwrap(),
);
});
let mut service = ServiceBuilder::new()
.layer(ExtractAuthorizationChecksLayer::new(schema))
.service(mock);
let context = Context::new();
context
.extensions()
.with_lock(|lock| lock.insert::<ParsedDocument>(doc));
let req = supergraph::Request::fake_builder()
.query(query)
.context(context)
.build()
.unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.context
.get::<_, HashSet<String>>(REQUIRED_SCOPES_KEY)
.unwrap(),
Some(HashSet::from([
"profile".to_string(),
"read:name".to_string(),
"read:user".to_string()
])),
"required scopes should have been inserted into context"
);
crate::plugin::test::await_mock_driver(driver).await;
}
#[tokio::test]
async fn extracts_authenticated() {
let query = "query { products(limit: 1) { price } }";
let config = Configuration::default();
let schema = Arc::new(Schema::parse(AUTHENTICATED_SCHEMA, &config).unwrap());
let doc = Query::parse_document(query, None, &schema, &config).unwrap();
let (mock, mut handle) =
tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let driver = tokio::spawn(async move {
let (req, responder) = handle.next_request().await.unwrap();
responder.send_response(
supergraph::Response::fake_builder()
.context(req.context)
.build()
.unwrap(),
);
});
let mut service = ServiceBuilder::new()
.layer(ExtractAuthorizationChecksLayer::new(schema))
.service(mock);
let context = Context::new();
context
.extensions()
.with_lock(|lock| lock.insert::<ParsedDocument>(doc));
let req = supergraph::Request::fake_builder()
.query(query)
.context(context)
.build()
.unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.context
.get_json_value(AUTHENTICATION_REQUIRED_KEY)
.unwrap(),
serde_json_bytes::json!(true),
"required scopes should have been inserted into context"
);
crate::plugin::test::await_mock_driver(driver).await;
}
#[tokio::test]
async fn extracts_policies() {
let query = "query { private { id } }";
let config = Configuration::default();
let schema = Arc::new(Schema::parse(POLICY_SCHEMA, &config).unwrap());
let doc = Query::parse_document(query, None, &schema, &config).unwrap();
let (mock, mut handle) =
tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let driver = tokio::spawn(async move {
let (req, responder) = handle.next_request().await.unwrap();
responder.send_response(
supergraph::Response::fake_builder()
.context(req.context)
.build()
.unwrap(),
);
});
let mut service = ServiceBuilder::new()
.layer(ExtractAuthorizationChecksLayer::new(schema))
.service(mock);
let context = Context::new();
context
.extensions()
.with_lock(|lock| lock.insert::<ParsedDocument>(doc));
let req = supergraph::Request::fake_builder()
.query(query)
.context(context)
.build()
.unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.context.get_json_value(REQUIRED_POLICIES_KEY).unwrap(),
serde_json_bytes::json!({ "admin": null }),
"required policies should have been inserted into context"
);
crate::plugin::test::await_mock_driver(driver).await;
}
#[tokio::test]
async fn errors_without_document() {
let config = Configuration::default();
let schema = Arc::new(Schema::parse(REQUIRES_SCOPES_SCHEMA, &config).unwrap());
let (mock, handle) = tower_test::mock::pair::<supergraph::Request, supergraph::Response>();
let mut service = ServiceBuilder::new()
.layer(ExtractAuthorizationChecksLayer::new(schema))
.service(mock);
let mut response = service
.ready()
.await
.unwrap()
.call(
supergraph::Request::fake_builder()
.query("query { me { id name } }")
.build()
.unwrap(),
)
.await
.unwrap();
assert_eq!(
StatusCode::INTERNAL_SERVER_ERROR,
response.response.status()
);
let graphql_response = response.next_response().await.unwrap();
assert!(graphql_response.contains_error_code("MISSING_EXECUTABLE_DOCUMENT"));
crate::plugin::test::assert_no_mock_calls(handle).await;
}
}