use std::fmt;
use std::ops::ControlFlow;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use futures::StreamExt;
use futures::future::ready;
use futures::stream::once;
use http::StatusCode;
use parking_lot::Mutex;
use rhai::AST;
use rhai::Dynamic;
use rhai::Engine;
use rhai::EvalAltResult;
use rhai::FnPtr;
use rhai::FuncArgs;
use rhai::Instant;
use rhai::Scope;
use rhai::Shared;
use schemars::JsonSchema;
use serde::Deserialize;
use tower::BoxError;
use tower::ServiceBuilder;
use tower::ServiceExt;
use self::engine::RhaiService;
use self::engine::SharedMut;
use crate::error::Error;
use crate::layers::ServiceBuilderExt;
use crate::plugin::Plugin;
use crate::plugin::PluginInit;
use crate::plugins::rhai::engine::OptionDance;
use crate::services::PipelineStep;
mod engine;
pub(crate) const RHAI_SPAN_NAME: &str = "rhai_plugin";
mod execution;
mod router;
mod subgraph;
mod supergraph;
struct Rhai {
ast: AST,
engine: Arc<Engine>,
scope: Arc<Mutex<Scope<'static>>>,
}
fn default_intern_strings() -> bool {
true
}
#[derive(Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
#[schemars(rename = "RhaiConfig")]
pub(crate) struct Conf {
scripts: Option<PathBuf>,
main: Option<String>,
#[serde(default = "default_intern_strings")]
intern_strings: bool,
}
#[async_trait::async_trait]
impl Plugin for Rhai {
type Config = Conf;
async fn new(init: PluginInit<Self::Config>) -> Result<Self, BoxError> {
let sdl = init.supergraph_sdl.clone();
let scripts_path = match init.config.scripts {
Some(path) => path,
None => "rhai".into(),
};
let main_file = match init.config.main {
Some(main) => main,
None => "main.rhai".to_string(),
};
let main = scripts_path.join(main_file);
let engine = Arc::new(Rhai::new_rhai_engine(
Some(scripts_path),
sdl.to_string(),
main.clone(),
init.config.intern_strings,
));
let ast = engine
.compile_file(main.clone())
.map_err(|err| format!("in Rhai script {}: {}", main.display(), err))?;
let mut scope = Scope::new();
scope.push_constant("apollo_sdl", sdl.to_string());
scope.push_constant("apollo_start", Instant::now());
engine.run_ast_with_scope(&mut scope, &ast)?;
Ok(Self {
ast,
engine,
scope: Arc::new(Mutex::new(scope)),
})
}
fn router_service(&self, service: router::BoxCloneService) -> router::BoxCloneService {
const FUNCTION_NAME_SERVICE: &str = "router_service";
if !self.ast_has_function(FUNCTION_NAME_SERVICE) {
return service;
}
tracing::debug!("router_service function found");
let shared_service = Arc::new(Mutex::new(Some(service)));
if let Err(error) = self.run_rhai_service(
FUNCTION_NAME_SERVICE,
None,
ServiceStep::Router(shared_service.clone()),
self.scope.clone(),
) {
tracing::error!(
service = "RouterService",
"service callback failed: {error}"
);
}
shared_service.take_unwrap()
}
fn supergraph_service(
&self,
service: supergraph::BoxCloneService,
) -> supergraph::BoxCloneService {
const FUNCTION_NAME_SERVICE: &str = "supergraph_service";
if !self.ast_has_function(FUNCTION_NAME_SERVICE) {
return service;
}
tracing::debug!("supergraph_service function found");
let shared_service = Arc::new(Mutex::new(Some(service)));
if let Err(error) = self.run_rhai_service(
FUNCTION_NAME_SERVICE,
None,
ServiceStep::Supergraph(shared_service.clone()),
self.scope.clone(),
) {
tracing::error!(
service = "SupergraphService",
"service callback failed: {error}"
);
}
shared_service.take_unwrap()
}
fn execution_service(&self, service: execution::BoxCloneService) -> execution::BoxCloneService {
const FUNCTION_NAME_SERVICE: &str = "execution_service";
if !self.ast_has_function(FUNCTION_NAME_SERVICE) {
return service;
}
tracing::debug!("execution_service function found");
let shared_service = Arc::new(Mutex::new(Some(service)));
if let Err(error) = self.run_rhai_service(
FUNCTION_NAME_SERVICE,
None,
ServiceStep::Execution(shared_service.clone()),
self.scope.clone(),
) {
tracing::error!(
service = "ExecutionService",
"service callback failed: {error}"
);
}
shared_service.take_unwrap()
}
fn subgraph_service(
&self,
name: &str,
service: subgraph::BoxCloneService,
) -> subgraph::BoxCloneService {
const FUNCTION_NAME_SERVICE: &str = "subgraph_service";
if !self.ast_has_function(FUNCTION_NAME_SERVICE) {
return service;
}
tracing::debug!("subgraph_service function found");
let shared_service = Arc::new(Mutex::new(Some(service)));
if let Err(error) = self.run_rhai_service(
FUNCTION_NAME_SERVICE,
Some(name),
ServiceStep::Subgraph(shared_service.clone()),
self.scope.clone(),
) {
tracing::error!(
service = "SubgraphService",
subgraph = name,
"service callback failed: {error}"
);
}
shared_service.take_unwrap()
}
}
#[derive(Clone, Debug)]
pub(crate) enum ServiceStep {
Router(SharedMut<router::BoxCloneService>),
Supergraph(SharedMut<supergraph::BoxCloneService>),
Execution(SharedMut<execution::BoxCloneService>),
Subgraph(SharedMut<subgraph::BoxCloneService>),
}
macro_rules! gen_map_request {
($base: ident, $borrow: ident, $rhai_service: ident, $callback: ident, $stage: expr) => {
$borrow.replace(|service| {
fn rhai_service_span() -> impl Fn(&$base::Request) -> tracing::Span + Clone {
move |_request: &$base::Request| {
tracing::info_span!(
RHAI_SPAN_NAME,
"rhai service" = stringify!($base::Request),
"otel.kind" = "INTERNAL"
)
}
}
ServiceBuilder::new()
.instrument(rhai_service_span())
.checkpoint_async(move |request: $base::Request| {
let rhai_service = $rhai_service.clone();
let callback = $callback.clone();
async move {
let shared_request = Shared::new(Mutex::new(Some(request)));
let result: Result<Dynamic, Box<EvalAltResult>> =
execute(&rhai_service, $stage, &callback, (shared_request.clone(),))
.await;
if let Err(error) = result {
let error_details = process_error(error);
if error_details.body.is_none() {
tracing::error!("map_request callback failed: {error_details:#?}");
}
let mut guard = shared_request.lock();
let request_opt = guard.take();
return $base::request_failure(
request_opt.unwrap().context,
error_details,
);
}
let mut guard = shared_request.lock();
let request_opt = guard.take();
Ok(ControlFlow::Continue(request_opt.unwrap()))
}
})
.service(service)
.boxed_clone()
})
};
}
macro_rules! gen_map_router_deferred_request {
($base: ident, $borrow: ident, $rhai_service: ident, $callback: ident, $stage: expr) => {
$borrow.replace(|service| {
fn rhai_service_span() -> impl Fn(&$base::Request) -> tracing::Span + Clone {
move |_request: &$base::Request| {
tracing::info_span!(
RHAI_SPAN_NAME,
"rhai service" = stringify!($base::Request),
"otel.kind" = "INTERNAL"
)
}
}
ServiceBuilder::new()
.instrument(rhai_service_span())
.checkpoint_async(move |chunked_request: $base::Request| {
let rhai_service = $rhai_service.clone();
let callback = $callback.clone();
async move {
let $base::Request { router_request, context } = chunked_request;
let (parts, stream) = router_request.into_parts();
let request = $base::FirstRequest {
context,
request: http::Request::from_parts(
parts,
(),
),
};
let shared_request = Shared::new(Mutex::new(Some(request)));
let result = execute(&rhai_service, $stage, &callback, (shared_request.clone(),)).await;
if let Err(error) = result {
let error_details = process_error(error);
if error_details.body.is_none() {
tracing::error!("map_request callback failed: {error_details:#?}");
}
let mut guard = shared_request.lock();
let request_opt = guard.take();
return $base::request_failure(request_opt.unwrap().context, error_details);
}
let request_opt = shared_request.lock().take();
let $base::FirstRequest { context, request } =
request_opt.unwrap();
let (parts, _body) = http::Request::from(request).into_parts();
Ok(ControlFlow::Continue($base::Request {
context,
router_request: http::Request::from_parts(parts, stream),
}))
}
})
.service(service)
.boxed_clone()
})
};
}
macro_rules! gen_map_response {
($base: ident, $borrow: ident, $rhai_service: ident, $callback: ident, $stage: expr) => {
$borrow.replace(|service| {
service
.and_then(move |response: $base::Response| async move {
let shared_response = Shared::new(Mutex::new(Some(response)));
let result: Result<Dynamic, Box<EvalAltResult>> = execute(
&$rhai_service,
$stage,
&$callback,
(shared_response.clone(),),
)
.await;
if let Err(error) = result {
let error_details = process_error(error);
if error_details.body.is_none() {
tracing::error!("map_request callback failed: {error_details:#?}");
}
let mut guard = shared_response.lock();
let response_opt = guard.take();
return Ok($base::response_failure(
response_opt.unwrap().context,
error_details,
));
}
let mut guard = shared_response.lock();
let response_opt = guard.take();
Ok(response_opt.unwrap())
})
.boxed_clone()
})
};
}
macro_rules! gen_map_router_deferred_response {
($base: ident, $borrow: ident, $rhai_service: ident, $callback: ident, $stage: expr) => {
$borrow.replace(|service| {
service.and_then(
|mapped_response: $base::Response| async move {
let $base::Response { response, context } = mapped_response;
let (parts, stream) = response.into_parts();
let response = $base::FirstResponse {
context,
response: http::Response::from_parts(
parts,
(),
)
.into(),
};
let shared_response = Shared::new(Mutex::new(Some(response)));
let result = execute(
&$rhai_service,
$stage,
&$callback,
(shared_response.clone(),),
)
.await;
if let Err(error) = result {
let error_details = process_error(error);
if error_details.body.is_none() {
tracing::error!("map_request callback failed: {error_details:#?}");
}
let response_opt = shared_response.lock().take();
return Ok($base::response_failure(
response_opt.unwrap().context,
error_details
));
}
let response_opt = shared_response.lock().take();
let $base::FirstResponse { context, response } =
response_opt.unwrap();
let (parts, _body) = http::Response::from(response).into_parts();
Ok($base::Response {
context,
response: http::Response::from_parts(parts, stream),
})
},
).boxed_clone()
})
};
}
macro_rules! gen_map_deferred_response {
($base: ident, $borrow: ident, $rhai_service: ident, $callback: ident, $stage: expr) => {
$borrow.replace(|service| {
service.and_then(
|mapped_response: $base::Response| async move {
let $base::Response { response, context } = mapped_response;
let (parts, stream) = response.into_parts();
let (first, rest) = StreamExt::into_future(stream).await;
if first.is_none() {
let error_details = ErrorDetails {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: Some("rhai execution error: empty response".to_string()),
position: None,
body: None
};
return Ok($base::response_failure(
context,
error_details
));
}
let response = $base::FirstResponse {
context,
response: http::Response::from_parts(
parts,
first.expect("already checked"),
)
.into(),
};
let shared_response = Shared::new(Mutex::new(Some(response)));
let result = execute(
&$rhai_service,
$stage,
&$callback,
(shared_response.clone(),),
)
.await;
if let Err(error) = result {
let error_details = process_error(error);
if error_details.body.is_none() {
tracing::error!("map_request callback failed: {error_details:#?}");
}
let mut guard = shared_response.lock();
let response_opt = guard.take();
return Ok($base::response_failure(
response_opt.unwrap().context,
error_details
));
}
let mut guard = shared_response.lock();
let response_opt = guard.take();
let $base::FirstResponse { context, response } =
response_opt.unwrap();
let (parts, body) = http::Response::from(response).into_parts();
let ctx = context.clone();
let mapped_stream = rest.filter_map(move |deferred_response| {
let rhai_service = $rhai_service.clone();
let context = context.clone();
let callback = $callback.clone();
async move {
let response = $base::DeferredResponse {
context,
response: deferred_response,
};
let shared_response = Shared::new(Mutex::new(Some(response)));
let result = execute(
&rhai_service,
$stage,
&callback,
(shared_response.clone(),),
)
.await;
if let Err(error) = result {
let error_details = process_error(error);
if error_details.body.is_none() {
tracing::error!("map_request callback failed: {error_details:#?}");
}
let mut guard = shared_response.lock();
let response_opt = guard.take();
let $base::DeferredResponse { mut response, .. } = response_opt.unwrap();
let error = Error::builder()
.message(error_details.message.unwrap_or_default())
.build();
response.errors = vec![error];
return Some(response);
}
let mut guard = shared_response.lock();
let response_opt = guard.take();
let $base::DeferredResponse { response, .. } =
response_opt.unwrap();
Some(response)
}
});
let response = http::Response::from_parts(
parts,
once(ready(body)).chain(mapped_stream).boxed(),
)
.into();
Ok($base::Response {
context: ctx,
response,
})
},
).boxed_clone()
})
};
}
impl ServiceStep {
fn map_request(&mut self, rhai_service: RhaiService, callback: FnPtr) {
match self {
ServiceStep::Router(service) => {
gen_map_router_deferred_request!(
router,
service,
rhai_service,
callback,
PipelineStep::RouterRequest
);
}
ServiceStep::Supergraph(service) => {
gen_map_request!(
supergraph,
service,
rhai_service,
callback,
PipelineStep::SupergraphRequest
);
}
ServiceStep::Execution(service) => {
gen_map_request!(
execution,
service,
rhai_service,
callback,
PipelineStep::ExecutionRequest
);
}
ServiceStep::Subgraph(service) => {
gen_map_request!(
subgraph,
service,
rhai_service,
callback,
PipelineStep::SubgraphRequest
);
}
}
}
fn map_response(&mut self, rhai_service: RhaiService, callback: FnPtr) {
match self {
ServiceStep::Router(service) => {
gen_map_router_deferred_response!(
router,
service,
rhai_service,
callback,
PipelineStep::RouterResponse
);
}
ServiceStep::Supergraph(service) => {
gen_map_deferred_response!(
supergraph,
service,
rhai_service,
callback,
PipelineStep::SupergraphResponse
);
}
ServiceStep::Execution(service) => {
gen_map_deferred_response!(
execution,
service,
rhai_service,
callback,
PipelineStep::ExecutionResponse
);
}
ServiceStep::Subgraph(service) => {
gen_map_response!(
subgraph,
service,
rhai_service,
callback,
PipelineStep::SubgraphResponse
);
}
}
}
}
#[derive(Deserialize, Debug)]
struct Position {
line: Option<usize>,
pos: Option<usize>,
}
impl fmt::Display for Position {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if let Some((line, pos)) = self.line.zip(self.pos) {
write!(f, "line {line}, position {pos}")
} else {
write!(f, "none")
}
}
}
impl From<&rhai::Position> for Position {
fn from(value: &rhai::Position) -> Self {
Self {
line: value.line(),
pos: value.position(),
}
}
}
#[derive(Deserialize, Debug)]
struct ErrorDetails {
#[serde(
with = "http_serde::status_code",
default = "default_thrown_status_code"
)]
status: StatusCode,
message: Option<String>,
position: Option<Position>,
body: Option<crate::graphql::Response>,
}
fn default_thrown_status_code() -> StatusCode {
StatusCode::INTERNAL_SERVER_ERROR
}
fn process_error(error: Box<EvalAltResult>) -> ErrorDetails {
let mut error_details = ErrorDetails {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: Some(format!("rhai execution error: '{error}'")),
position: None,
body: None,
};
let inner_error = error.unwrap_inner();
if let EvalAltResult::ErrorRuntime(obj, pos) = inner_error {
if let Ok(temp_error_details) = rhai::serde::from_dynamic::<ErrorDetails>(obj) {
if temp_error_details.message.is_some() || temp_error_details.body.is_some() {
error_details = temp_error_details;
} else {
error_details.status = temp_error_details.status;
}
}
error_details.position = Some(pos.into());
}
error_details
}
async fn execute(
rhai_service: &RhaiService,
stage: PipelineStep,
callback: &FnPtr,
args: impl FuncArgs + Send + 'static,
) -> Result<Dynamic, Box<EvalAltResult>> {
let rhai_service = rhai_service.clone();
let callback = callback.clone();
let start = Instant::now();
let (result, duration) = match tokio::task::spawn_blocking(move || {
let result = if callback.is_curried() {
callback.call(&rhai_service.engine, &rhai_service.ast, args)
} else {
let mut scope = rhai_service.scope.lock().clone();
rhai_service
.engine
.call_fn(&mut scope, &rhai_service.ast, callback.fn_name(), args)
};
(result, start.elapsed())
})
.await
{
Ok(result) => result,
Err(join_error) => (
Err(Box::new(EvalAltResult::ErrorSystem(
"rhai script execution task did not complete".to_string(),
Box::new(join_error),
))),
Duration::default(),
),
};
record_rhai_execution(stage, duration, result.is_ok());
result
}
fn record_rhai_execution(stage: PipelineStep, duration: Duration, succeeded: bool) {
let duration = duration.as_secs_f64();
let stage = stage.to_string();
f64_histogram_with_unit!(
"apollo.router.operations.rhai.duration",
"Time spent executing a Rhai script callback, in seconds",
"s",
duration,
"rhai.stage" = stage,
"rhai.succeeded" = succeeded
);
}
register_plugin!("apollo", "rhai", Rhai);
#[cfg(test)]
mod tests;