use std::collections::HashSet;
use crate::catalog::providers::DatabaseProvider;
use crate::catalog::{DatabaseId, Error as CatalogError, NamespaceId};
use crate::expr::function_facts::FunctionFacts;
use crate::kvs::Transaction;
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) enum WriteSource {
Body,
Function(String),
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub(crate) struct ResolvedMutability {
pub(crate) write: Option<WriteSource>,
pub(crate) opaque_effects: bool,
pub(crate) undefined_callee: bool,
}
impl ResolvedMutability {
pub(crate) fn possibly_writes(&self) -> bool {
self.write.is_some() || self.opaque_effects || self.undefined_callee
}
}
pub(crate) async fn resolve_mutability(
txn: &Transaction,
ns: NamespaceId,
db: DatabaseId,
root: &FunctionFacts,
) -> anyhow::Result<ResolvedMutability> {
let mut resolved = ResolvedMutability {
write: root.direct_writes.then_some(WriteSource::Body),
opaque_effects: root.opaque_effects,
undefined_callee: false,
};
let mut visited: HashSet<String> = root.calls.iter().cloned().collect();
let mut pending: Vec<String> = root.calls.iter().cloned().collect();
while resolved.write.is_none() {
let Some(name) = pending.pop() else {
break;
};
let facts = match txn.get_db_function(ns, db, &name, None).await {
Ok(def) => def.block.function_facts(),
Err(e) => {
if matches!(e.downcast_ref::<CatalogError>(), Some(CatalogError::FcNotFound { .. }))
{
resolved.undefined_callee = true;
continue;
}
return Err(e);
}
};
if facts.direct_writes {
resolved.write = Some(WriteSource::Function(name));
break;
}
resolved.opaque_effects |= facts.opaque_effects;
for call in facts.calls {
if visited.insert(call.clone()) {
pending.push(call);
}
}
}
Ok(resolved)
}
#[cfg(all(test, feature = "kv-mem"))]
#[allow(clippy::unwrap_used)]
mod tests {
use std::sync::Arc;
use super::*;
use crate::dbs::Session;
use crate::expr::Expr;
use crate::kvs::{Datastore, TransactionType};
async fn datastore_with(setup: &str) -> Arc<Datastore> {
let ds = Datastore::new("memory").await.unwrap();
let sess = Session::owner().with_ns("test").with_db("test");
ds.execute("DEFINE NAMESPACE test; DEFINE DATABASE test", &sess, None).await.unwrap();
for res in ds.execute(setup, &sess, None).await.unwrap() {
res.result.unwrap();
}
ds
}
async fn resolve(ds: &Datastore, expression: &str) -> ResolvedMutability {
let expr: Expr = crate::syn::expr(expression).unwrap().into();
let facts = expr.function_facts();
let txn = ds.transaction(TransactionType::Read).await.unwrap();
let db = txn.get_db_by_name("test", "test", None).await.unwrap().unwrap();
let resolved =
resolve_mutability(&txn, db.namespace_id, db.database_id, &facts).await.unwrap();
txn.cancel().await.unwrap();
resolved
}
#[test]
fn facts_classify_each_callable_kind() {
let facts = |expression: &str| {
let expr: Expr = crate::syn::expr(expression).unwrap().into();
expr.function_facts()
};
assert_eq!(facts("math::abs(-1)"), FunctionFacts::default());
let f = facts("fn::a(fn::b(1))");
assert!(!f.direct_writes && !f.opaque_effects);
assert_eq!(f.calls.len(), 2);
for expression in ["api::invoke('/x')", "eval::surql('RETURN 1')", "$fn(1)"] {
let f = facts(expression);
assert!(f.opaque_effects && !f.direct_writes, "{expression}");
}
for expression in ["|| { CREATE log }", "(|| { CREATE log })()"] {
let f = facts(expression);
assert!(f.direct_writes && !f.opaque_effects, "{expression}");
}
let f = facts("{ CREATE log; fn::after(); }");
assert!(f.direct_writes);
assert!(f.calls.contains("after"));
}
#[tokio::test]
async fn a_direct_write_is_the_body_source() {
let ds = datastore_with("").await;
let resolved = resolve(&ds, "(CREATE log)").await;
assert_eq!(resolved.write, Some(WriteSource::Body));
}
#[tokio::test]
async fn a_write_two_calls_deep_names_the_writing_function() {
let ds = datastore_with(
"DEFINE FUNCTION fn::sink() { CREATE log; RETURN 1; };
DEFINE FUNCTION fn::relay() { RETURN fn::sink(); };",
)
.await;
let resolved = resolve(&ds, "fn::relay()").await;
assert_eq!(resolved.write, Some(WriteSource::Function("sink".to_owned())));
}
#[tokio::test]
async fn a_write_free_cycle_resolves_clean() {
let ds = datastore_with(
"DEFINE FUNCTION fn::ping($n: number) { RETURN IF $n > 0 { fn::pong($n - 1) } ELSE { 0 }; };
DEFINE FUNCTION fn::pong($n: number) { RETURN IF $n > 0 { fn::ping($n - 1) } ELSE { 0 }; };",
)
.await;
let resolved = resolve(&ds, "fn::ping(3)").await;
assert_eq!(resolved.write, None);
assert!(!resolved.possibly_writes());
}
#[tokio::test]
async fn a_cycle_carrying_a_write_names_its_writer() {
let ds = datastore_with(
"DEFINE FUNCTION fn::a($n: number) { RETURN IF $n > 0 { fn::b($n - 1) } ELSE { 0 }; };
DEFINE FUNCTION fn::b($n: number) { CREATE log; RETURN fn::a($n); };",
)
.await;
let resolved = resolve(&ds, "fn::a(3)").await;
assert_eq!(resolved.write, Some(WriteSource::Function("b".to_owned())));
}
#[tokio::test]
async fn an_undefined_callee_is_possible_but_not_provable() {
let ds = datastore_with("").await;
let resolved = resolve(&ds, "fn::ghost()").await;
assert_eq!(resolved.write, None);
assert!(resolved.undefined_callee);
assert!(resolved.possibly_writes());
}
#[tokio::test]
async fn opaque_callables_are_possible_but_not_provable() {
let ds = datastore_with("").await;
for expression in ["eval::surql('RETURN 1')", "$fn(1)"] {
let resolved = resolve(&ds, expression).await;
assert_eq!(resolved.write, None, "{expression}");
assert!(resolved.opaque_effects, "{expression}");
assert!(resolved.possibly_writes(), "{expression}");
}
}
#[tokio::test]
async fn a_write_in_a_call_argument_is_direct() {
let ds = datastore_with("DEFINE FUNCTION fn::id($x: any) { RETURN $x; };").await;
for expression in ["fn::id((CREATE log).id)", "$fn((CREATE log).id)"] {
let resolved = resolve(&ds, expression).await;
assert_eq!(resolved.write, Some(WriteSource::Body), "{expression}");
}
}
}