use apollo_router::plugin::Plugin;
use apollo_router::plugin::PluginInit;
use apollo_router::register_plugin;
use apollo_router::services::supergraph;
use schemars::JsonSchema;
use serde::Deserialize;
use tower::BoxError;
use tower::ServiceBuilder;
use tower::ServiceExt;

const CLIENT_NAME_MISSING_CONTEXT_KEY: &str = "client_name_missing";

#[derive(Debug)]
struct RequireApolloClientName {
    configuration: Configuration,
}

#[derive(Debug, Default, Deserialize, JsonSchema)]
struct Configuration {
    enabled: bool,
}

#[async_trait::async_trait]
impl Plugin for RequireApolloClientName {
    type Config = Configuration;

    async fn new(init: PluginInit<Self::Config>) -> Result<Self, BoxError> {
        tracing::info!("require_apollo_client_name started");
        Ok(RequireApolloClientName {
            configuration: init.config,
        })
    }

    fn supergraph_service(&self, service: supergraph::BoxService) -> supergraph::BoxService {
        let enabled = self.configuration.enabled;

        ServiceBuilder::new()
            .map_request(move |request: supergraph::Request| {
                if enabled
                    && !request
                        .supergraph_request
                        .headers()
                        .contains_key("Apollographql-Client-Name")
                {
                    if let Err(e) = request
                        .context
                        .insert(CLIENT_NAME_MISSING_CONTEXT_KEY, true)
                    {
                        tracing::error!(
                            "Failed to insert {} context value: {}",
                            CLIENT_NAME_MISSING_CONTEXT_KEY,
                            e
                        );
                    }
                }
                request
            })
            .map_response(|response: supergraph::Response| {
                let context = response.context.clone();
                if context
                    .get::<_, bool>(CLIENT_NAME_MISSING_CONTEXT_KEY)
                    .ok()
                    .flatten()
                    == Some(true)
                {
                    return supergraph::Response::error_builder()
                        .error(
                            apollo_router::graphql::Error::builder()
                                .message("Apollographql-Client-Name header must be provided.")
                                .extension_code("BAD_REQUEST")
                                .build(),
                        )
                        .status_code(http::StatusCode::BAD_REQUEST)
                        .context(context)
                        .build()
                        .unwrap_or_else(|e| {
                            tracing::error!(
                                "Failed to build error response for missing client name: {}",
                                e
                            );
                            response
                        });
                }
                response
            })
            .service(service)
            .boxed()
    }
}

register_plugin!(
    "theorchard",
    "require_apollo_client_name",
    RequireApolloClientName
);
