use std::{sync::Arc, time::Duration};
use fraiseql_core::{
error::{FraiseQLError, Result},
runtime::{QueryFunctionRequest, QueryFunctionResolver},
};
use crate::subsystems::BeforeMutationHooks;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct QueryFunctionBudget(Duration);
impl QueryFunctionBudget {
pub const DEFAULT_MS: u64 = 5_000;
pub const ENV: &'static str = "FRAISEQL_FUNCTIONS_REQUEST_QUERY_BUDGET_MS";
#[must_use]
pub const fn from_millis(millis: u64) -> Self {
Self(Duration::from_millis(millis))
}
#[must_use]
pub fn from_getter(get: impl Fn(&str) -> Option<String>) -> Self {
Self::from_millis(
get(Self::ENV).and_then(|value| value.parse().ok()).unwrap_or(Self::DEFAULT_MS),
)
}
#[must_use]
pub fn from_env() -> Self {
Self::from_getter(|key| std::env::var(key).ok())
}
#[must_use]
pub const fn duration(self) -> Duration {
self.0
}
#[must_use]
pub const fn is_enforced(self) -> bool {
!self.0.is_zero()
}
#[must_use]
pub fn for_function(self, declared_ms: Option<u64>) -> Self {
declared_ms.map_or(self, Self::from_millis)
}
}
impl Default for QueryFunctionBudget {
fn default() -> Self {
Self::from_millis(Self::DEFAULT_MS)
}
}
pub struct FunctionQueryResolver {
hooks: Arc<BeforeMutationHooks>,
budget: QueryFunctionBudget,
}
impl FunctionQueryResolver {
#[must_use]
pub fn new(hooks: Arc<BeforeMutationHooks>) -> Self {
Self {
hooks,
budget: QueryFunctionBudget::default(),
}
}
#[must_use]
pub const fn with_budget(mut self, budget: QueryFunctionBudget) -> Self {
self.budget = budget;
self
}
}
impl QueryFunctionResolver for FunctionQueryResolver {
fn resolve<'a>(
&'a self,
request: QueryFunctionRequest<'a>,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<serde_json::Value>> + Send + 'a>>
{
Box::pin(async move {
let field = request.field.to_string();
let budget = self.budget.for_function(self.declared_timeout_ms(request.function));
run_within_budget(&field, budget, self.invoke(request)).await
})
}
}
impl FunctionQueryResolver {
fn declared_timeout_ms(&self, function: &str) -> Option<u64> {
self.hooks.request_query_timeouts.get(function).copied().flatten()
}
#[cfg(feature = "functions-runtime")]
async fn invoke(&self, request: QueryFunctionRequest<'_>) -> Result<serde_json::Value> {
use fraiseql_functions::host::request_query::{
RequestQueryHost, interpret_query_answer, request_query_payload,
};
let module = self.hooks.module_registry.get(request.function).ok_or_else(|| {
FraiseQLError::Validation {
message: format!(
"`{}` is backed by the function `{}`, which is not in the module registry — \
the module was not loaded from `[functions] module_dir` at boot",
request.field, request.function
),
path: Some(request.field.to_string()),
}
})?;
let payload = request_query_payload(request.field, request.arguments.clone());
let host = RequestQueryHost::new(
payload.clone(),
Some(Arc::clone(&request.reader)),
request.principal,
super::after_mutation::host_context_config(),
);
let limits = fraiseql_functions::ResourceLimits {
max_duration: self
.budget
.for_function(self.declared_timeout_ms(request.function))
.duration(),
..fraiseql_functions::ResourceLimits::default()
};
let result = self
.hooks
.observer
.invoke_with_context(module, payload, Arc::new(host), limits)
.await
.map_err(|error| attribute(request.field, request.function, &error))?;
Ok(interpret_query_answer(result.value))
}
#[cfg(not(feature = "functions-runtime"))]
#[allow(clippy::unused_async)] async fn invoke(&self, request: QueryFunctionRequest<'_>) -> Result<serde_json::Value> {
Err(FraiseQLError::Unsupported {
message: format!(
"`{}` is backed by the function `{}`, but this build has no function runtime \
(feature `functions-runtime`) — the field is refused rather than answered with \
nothing",
request.field, request.function
),
})
}
}
#[cfg(feature = "functions-runtime")]
fn attribute(field: &str, function: &str, error: &FraiseQLError) -> FraiseQLError {
FraiseQLError::Validation {
message: format!("`{field}` (function `{function}`): {error}"),
path: Some(field.to_string()),
}
}
async fn run_within_budget<F>(
field: &str,
budget: QueryFunctionBudget,
invocation: F,
) -> Result<serde_json::Value>
where
F: std::future::Future<Output = Result<serde_json::Value>>,
{
if !budget.is_enforced() {
return invocation.await;
}
match tokio::time::timeout(budget.duration(), invocation).await {
Ok(answer) => answer,
Err(_) => Err(FraiseQLError::Timeout {
timeout_ms: u64::try_from(budget.duration().as_millis()).unwrap_or(u64::MAX),
query: Some(field.to_string()),
}),
}
}
#[cfg(test)]
mod tests;