use std::collections::{HashMap, HashSet};
use std::fmt::Display;
use std::sync::Arc;
use async_graphql::dynamic::indexmap::IndexMap;
use async_graphql::dynamic::{
Enum, Field, FieldFuture, FieldValue, InputObject, InputValue, Object, ResolverContext, Type,
TypeRef,
};
use async_graphql::{Name, Value as GraphqlValue};
use surrealdb_types::ToSql;
use super::error::{GraphqlError, resolver_error};
use super::relations::{RelationDirection, RelationInfo};
#[derive(Clone, Debug)]
pub(crate) struct RelationFieldInfo {
pub field_name: String,
pub relation_table: TableName,
pub dir: Dir,
}
use super::schema::{
SchemaContext, graphql_to_sql_kind, graphql_to_sql_kind_with_scope, sql_value_to_graphql_value,
sql_value_to_graphql_value_with_kind,
};
use crate::catalog::providers::TableProvider;
use crate::catalog::{FieldDefinition, TableDefinition};
use crate::dbs::Session;
use crate::expr::field::{Field as SelectField, Selector};
use crate::expr::group::{Group, Groups};
use crate::expr::lookup::{Lookup, LookupKind, LookupSubject};
use crate::expr::order::{OrderList, Ordering};
use crate::expr::part::Part;
use crate::expr::statements::SelectStatement;
use crate::expr::{
self, BinaryOperator, Cond, Dir, Expr, Fields, Function, FunctionCall, Idiom, Kind,
KindLiteral, Limit, Literal, LogicalPlan, Start, TopLevelExpr,
};
use crate::graphql::error::internal_error;
use crate::graphql::schema::{
filter_type_name, geometry_graphql_type_name, kind_to_type, kind_to_type_with_enum_prefix,
unwrap_type,
};
use crate::graphql::utils::{GraphqlValueUtils, execute_plan};
use crate::kvs::Datastore;
use crate::val::{Array as SurArray, Datetime, Object as SurObject, RecordId, TableName, Value};
const MAX_ID_IN_LIST: usize = 1000;
pub(crate) fn field_graphql_name(fd: &FieldDefinition) -> String {
if let Some(ref alias) = fd.graphql_alias
&& is_valid_graphql_identifier(alias)
{
return alias.clone();
}
idiom_to_graphql_name(&fd.name)
}
pub(crate) fn is_valid_graphql_identifier_pub(s: &str) -> bool {
is_valid_graphql_identifier(s)
}
fn is_valid_graphql_identifier(s: &str) -> bool {
let mut chars = s.chars();
let Some(first) = chars.next() else {
return false;
};
if !(first.is_ascii_alphabetic() || first == '_') {
return false;
}
chars.all(|c| c.is_ascii_alphanumeric() || c == '_')
}
pub(crate) fn idiom_to_graphql_name(idiom: &Idiom) -> String {
if idiom.0.len() == 1
&& let Part::Field(name) = &idiom.0[0]
{
return name.as_str().to_owned();
}
let raw = idiom.to_sql();
let mut out = String::with_capacity(raw.len());
for (i, c) in raw.chars().enumerate() {
let ok = if i == 0 {
c.is_ascii_alphabetic() || c == '_'
} else {
c.is_ascii_alphanumeric() || c == '_'
};
out.push(if ok {
c
} else {
'_'
});
}
if out.is_empty() {
out.push('_');
}
out
}
fn order_asc(field_name: String) -> expr::Order {
expr::Order {
value: Idiom::field(field_name),
direction: true,
..Default::default()
}
}
fn order_desc(field_name: String) -> expr::Order {
expr::Order {
value: Idiom::field(field_name),
..expr::Order::default()
}
}
#[derive(Clone, Debug)]
pub(crate) struct VersionedRecord {
pub rid: RecordId,
pub version: Option<Datetime>,
}
#[derive(Clone, Debug)]
pub(crate) struct CachedRecord {
pub rid: RecordId,
pub version: Option<Datetime>,
pub data: SurObject,
}
fn version_to_expr(version: &Option<Datetime>) -> Expr {
match version {
Some(dt) => Expr::Literal(Literal::Datetime(*dt)),
None => Expr::Literal(Literal::None),
}
}
fn parse_version_arg(
args: &IndexMap<Name, GraphqlValue>,
) -> Result<Option<Datetime>, GraphqlError> {
match args.get("version") {
Some(GraphqlValue::String(s)) => {
let dt = crate::syn::datetime(s)
.map_err(|_| resolver_error(format!("Invalid version datetime: {s}")))?;
Ok(Some(dt.into()))
}
Some(GraphqlValue::Null) | None => Ok(None),
Some(_) => Err(resolver_error("version must be a datetime string")),
}
}
fn parse_start_arg(args: &IndexMap<Name, GraphqlValue>) -> Option<Start> {
args.get("start").and_then(|v| v.as_i64()).map(|s| Start(Expr::Literal(Literal::Integer(s))))
}
fn parse_limit_arg(args: &IndexMap<Name, GraphqlValue>) -> Option<Limit> {
args.get("limit").and_then(|v| v.as_i64()).map(|l| Limit(Expr::Literal(Literal::Integer(l))))
}
fn parse_order_arg(
args: &IndexMap<Name, GraphqlValue>,
fds: &[FieldDefinition],
) -> Result<Option<Ordering>, GraphqlError> {
let order = args.get("order");
let to_lookup = |name: &str| -> String {
fds.iter()
.find(|fd| field_graphql_name(fd) == name)
.map(|fd| idiom_to_graphql_name(&fd.name))
.unwrap_or_else(|| name.to_string())
};
match order {
Some(GraphqlValue::Object(o)) => {
let mut orders = vec![];
let mut current = o;
loop {
let asc = current.get("asc");
let desc = current.get("desc");
match (asc, desc) {
(Some(_), Some(_)) => {
return Err(resolver_error("Found both ASC and DESC in order"));
}
(Some(GraphqlValue::Enum(a)), None) => {
orders.push(order_asc(to_lookup(a.as_str())))
}
(None, Some(GraphqlValue::Enum(d))) => {
orders.push(order_desc(to_lookup(d.as_str())))
}
(_, _) => break,
}
if let Some(GraphqlValue::Object(next)) = current.get("then") {
current = next;
} else {
break;
}
}
Ok(Some(Ordering::Order(OrderList(orders))))
}
_ => Ok(None),
}
}
pub(crate) fn parse_filter_arg(
args: &IndexMap<Name, GraphqlValue>,
fds: &[FieldDefinition],
tb_name: &str,
relations: &[RelationFieldInfo],
) -> Result<Option<Cond>, GraphqlError> {
let filter = args.get("filter").or_else(|| args.get("where"));
match filter {
Some(GraphqlValue::Object(o)) => Ok(Some(cond_from_filter(o, fds, tb_name, relations)?)),
Some(f) => {
error!(
"Found filter {f}, which should be object and should have \
been rejected by async graphql."
);
Err(resolver_error("Value in cond doesn't fit schema"))
}
None => Ok(None),
}
}
fn select_all_from_record(rid: &RecordId, version: &Option<Datetime>) -> SelectStatement {
SelectStatement {
what: vec![Value::RecordId(rid.clone()).into_literal()],
fields: Fields::all(),
only: true,
version: version_to_expr(version),
timeout: Expr::Literal(Literal::None),
omit: vec![],
with: None,
cond: None,
split: None,
group: None,
order: None,
limit: None,
start: None,
fetch: None,
explain: None,
tempfiles: false,
}
}
fn select_field_from_record(
rid: &RecordId,
field_name: &str,
version: &Option<Datetime>,
) -> SelectStatement {
SelectStatement {
what: vec![Value::RecordId(rid.clone()).into_literal()],
fields: Fields::Value(Box::new(Selector {
expr: Expr::Idiom(Idiom::field(field_name.to_string())),
alias: None,
})),
only: true,
version: version_to_expr(version),
timeout: Expr::Literal(Literal::None),
omit: vec![],
with: None,
cond: None,
split: None,
group: None,
order: None,
limit: None,
start: None,
fetch: None,
explain: None,
tempfiles: false,
}
}
fn select_all_from_table(
what: Expr,
cond: Option<Cond>,
order: Option<Ordering>,
limit: Option<Limit>,
start: Option<Start>,
version: &Option<Datetime>,
) -> SelectStatement {
SelectStatement {
what: vec![what],
fields: Fields::all(),
order,
cond,
limit,
start,
version: version_to_expr(version),
timeout: Expr::Literal(Literal::None),
omit: vec![],
only: false,
with: None,
split: None,
group: None,
fetch: None,
explain: None,
tempfiles: false,
}
}
async fn execute_select(
ds: &Datastore,
sess: &Session,
stmt: SelectStatement,
) -> Result<Value, GraphqlError> {
let plan = LogicalPlan {
expressions: vec![TopLevelExpr::Expr(Expr::Select(Box::new(stmt)))],
};
execute_plan(ds, sess, plan).await
}
struct NestedSubField {
name: String,
kind: Option<Kind>,
comment: Option<String>,
}
struct NestedObjectInfo {
graphql_type_name: String,
is_array: bool,
optional: bool,
sub_fields: Vec<NestedSubField>,
}
fn detect_nested_objects(
table_name: &str,
fds: &[FieldDefinition],
) -> HashMap<String, NestedObjectInfo> {
let mut children_by_parent: HashMap<String, Vec<NestedSubField>> = HashMap::new();
let mut parent_has_wildcard: HashMap<String, bool> = HashMap::new();
for fd in fds.iter() {
let parts = &fd.name.0;
if parts.len() < 2 {
continue; }
let parent_name = match &parts[0] {
Part::Field(name) => name.as_str().to_owned(),
_ => continue,
};
let child_name = match parts.last() {
Some(Part::Field(name)) => name.as_str().to_owned(),
_ => continue,
};
let has_wildcard = parts[1..parts.len() - 1].iter().any(|p| matches!(p, Part::All));
let expected_len = if has_wildcard {
3
} else {
2
};
if parts.len() != expected_len {
continue; }
parent_has_wildcard.entry(parent_name.clone()).or_insert(has_wildcard);
children_by_parent.entry(parent_name.clone()).or_default().push(NestedSubField {
name: child_name,
kind: fd.field_kind.clone(),
comment: fd.comment.clone(),
});
}
let mut result = HashMap::new();
for (parent_name, sub_fields) in children_by_parent {
let parent_fd = fds.iter().find(|fd| {
fd.name.0.len() == 1 && matches!(&fd.name.0[0], Part::Field(n) if n == &parent_name)
});
let is_array = parent_has_wildcard.get(&parent_name).copied().unwrap_or(false);
let parent_kind = parent_fd.and_then(|fd| fd.field_kind.as_ref());
let (kind_ok, optional) = match parent_kind {
Some(Kind::Object) if !is_array => (true, false),
Some(Kind::Array(inner, _)) if is_array => (matches!(**inner, Kind::Object), false),
Some(Kind::Either(ks)) => {
let has_none = ks.iter().any(|k| matches!(k, Kind::None));
if !is_array {
let has_object = ks.iter().any(|k| matches!(k, Kind::Object));
(has_none && has_object, has_none)
} else {
let has_array_obj = ks.iter().any(
|k| matches!(k, Kind::Array(inner, _) if matches!(**inner, Kind::Object)),
);
(has_none && has_array_obj, has_none)
}
}
None => (true, true),
_ => (false, false),
};
if !kind_ok {
continue;
}
let graphql_type_name = format!("{table_name}_{parent_name}");
result.insert(
parent_name,
NestedObjectInfo {
graphql_type_name,
is_array,
optional,
sub_fields,
},
);
}
for fd in fds.iter() {
if fd.name.0.len() != 1 {
continue;
}
let Part::Field(name) = &fd.name.0[0] else {
continue;
};
let parent_name = name.as_str().to_owned();
if result.contains_key(&parent_name) {
continue;
}
let Some(kind) = fd.field_kind.as_ref() else {
continue;
};
let (literal_map, is_array, optional) = match extract_literal_object(kind) {
Some(x) => x,
None => continue,
};
let sub_fields: Vec<NestedSubField> = literal_map
.iter()
.map(|(k, v)| NestedSubField {
name: k.as_str().to_owned(),
kind: Some(v.clone()),
comment: None,
})
.collect();
if sub_fields.is_empty() {
continue;
}
let graphql_type_name = format!("{table_name}_{parent_name}");
result.insert(
parent_name,
NestedObjectInfo {
graphql_type_name,
is_array,
optional,
sub_fields,
},
);
}
result
}
fn extract_literal_object(
kind: &Kind,
) -> Option<(&std::collections::BTreeMap<surrealdb_strand::Strand, Kind>, bool, bool)> {
match kind {
Kind::Literal(KindLiteral::Object(map)) => Some((map, false, false)),
Kind::Array(inner, _) => match inner.as_ref() {
Kind::Literal(KindLiteral::Object(map)) => Some((map, true, false)),
_ => None,
},
Kind::Either(ks) => {
let has_none = ks.iter().any(|k| matches!(k, Kind::None | Kind::Null));
let non_none: Vec<&Kind> =
ks.iter().filter(|k| !matches!(k, Kind::None | Kind::Null)).collect();
if non_none.len() != 1 {
return None;
}
match non_none[0] {
Kind::Literal(KindLiteral::Object(map)) => Some((map, false, has_none)),
Kind::Array(inner, _) => match inner.as_ref() {
Kind::Literal(KindLiteral::Object(map)) => Some((map, true, has_none)),
_ => None,
},
_ => None,
}
}
_ => None,
}
}
fn make_nested_object_type(
type_name: &str,
sub_fields: &[NestedSubField],
types: &mut Vec<Type>,
) -> Result<Object, GraphqlError> {
let mut obj = Object::new(type_name);
for sf in sub_fields {
let Some(ref kind) = sf.kind else {
continue;
};
let enum_scope = format!("{type_name}_{}", sf.name);
let fd_type = kind_to_type_with_enum_prefix(kind.clone(), types, false, Some(&enum_scope))?;
let resolver = make_sub_field_resolver(sf.name.clone(), sf.kind.clone(), Some(enum_scope));
let mut field = Field::new(&sf.name, fd_type, resolver);
if let Some(ref comment) = sf.comment {
field = field.description(comment.clone());
}
obj = obj.field(field);
}
Ok(obj)
}
fn make_sub_field_resolver(
field_name: String,
kind: Option<Kind>,
enum_scope: Option<String>,
) -> impl for<'a> Fn(ResolverContext<'a>) -> FieldFuture<'a> + Send + Sync + 'static {
move |ctx: ResolverContext| {
let field_name = field_name.clone();
let field_kind = kind.clone();
let enum_scope = enum_scope.clone();
FieldFuture::new(async move {
let obj = ctx.parent_value.try_downcast_ref::<SurObject>()?;
match obj.get(&field_name) {
Some(val) => match val {
Value::None | Value::Null => Ok(None),
Value::RecordId(rid) => {
let field_val = FieldValue::owned_any(VersionedRecord {
rid: rid.clone(),
version: None,
});
let field_val = match field_kind {
Some(Kind::Record(ref ts)) if ts.is_empty() || ts.len() > 1 => {
field_val.with_type(rid.table.clone())
}
_ => field_val,
};
Ok(Some(field_val))
}
Value::Geometry(g) => {
let type_name = geometry_graphql_type_name(g);
let field_val = FieldValue::owned_any(g.clone());
let field_val = match &field_kind {
Some(Kind::Geometry(ks)) if ks.is_empty() || ks.len() > 1 => {
field_val.with_type(type_name)
}
_ => field_val,
};
Ok(Some(field_val))
}
v => {
let graphql_val = sql_value_to_graphql_value_with_kind(
v.clone(),
field_kind.as_ref(),
enum_scope.as_deref(),
)
.map_err(async_graphql::Error::from)?;
Ok(Some(FieldValue::value(graphql_val)))
}
},
None => Ok(None),
}
})
}
}
fn make_nested_object_field_resolver(
fd_name: impl Into<String>,
is_array: bool,
) -> impl for<'a> Fn(ResolverContext<'a>) -> FieldFuture<'a> + Send + Sync + 'static {
let fd_name = fd_name.into();
move |ctx: ResolverContext| {
let fd_name = fd_name.clone();
FieldFuture::new(async move {
if let Ok(cached) = ctx.parent_value.try_downcast_ref::<CachedRecord>() {
let val = cached.data.get(&fd_name).cloned().unwrap_or(Value::None);
return resolve_nested_object_value(val, is_array);
}
let ds = ctx.data::<Arc<Datastore>>()?;
let sess = ctx.data::<Arc<Session>>()?;
let (rid, version) = match ctx.parent_value.try_downcast_ref::<VersionedRecord>() {
Ok(vr) => (vr.rid.clone(), vr.version),
Err(_) => {
let rid = ctx.parent_value.try_downcast_ref::<RecordId>()?;
(rid.clone(), None)
}
};
let stmt = select_field_from_record(&rid, &fd_name, &version);
let val = execute_select(ds, sess, stmt).await?;
resolve_nested_object_value(val, is_array)
})
}
}
fn resolve_nested_object_value(
val: Value,
is_array: bool,
) -> Result<Option<FieldValue<'static>>, async_graphql::Error> {
if is_array {
match val {
Value::Array(arr) => {
let items: Vec<FieldValue> = arr
.0
.into_iter()
.filter_map(|v| match v {
Value::Object(obj) => Some(FieldValue::owned_any(obj)),
_ => None,
})
.collect();
Ok(Some(FieldValue::list(items)))
}
Value::None | Value::Null => Ok(None),
_ => Ok(None),
}
} else {
match val {
Value::Object(obj) => Ok(Some(FieldValue::owned_any(obj))),
Value::None | Value::Null => Ok(None),
_ => {
let out = sql_value_to_graphql_value(val).map_err(async_graphql::Error::from)?;
Ok(Some(FieldValue::value(out)))
}
}
}
}
pub(crate) fn filter_name_from_table(tb_name: impl Display) -> String {
format!("_filter_{tb_name}")
}
fn objects_to_cached_records(
arr: SurArray,
version: Option<Datetime>,
) -> Result<Option<FieldValue<'static>>, async_graphql::Error> {
let out: Result<Vec<FieldValue>, GraphqlError> = arr
.0
.into_iter()
.map(|v| match v {
Value::Object(obj) => {
let rid = match obj.get("id") {
Some(Value::RecordId(rid)) => rid.clone(),
_ => {
error!("Object missing 'id' field or id is not a RecordId: {obj:?}");
return Err(internal_error("Record missing 'id' field"));
}
};
Ok(FieldValue::owned_any(CachedRecord {
rid,
version,
data: obj,
}))
}
_ => {
error!("Expected object in result, found: {v:?}");
Err(internal_error("Expected object in result"))
}
})
.collect();
match out {
Ok(l) => Ok(Some(FieldValue::list(l))),
Err(e) => Err(e.into()),
}
}
fn make_table_list_field(
tb: &TableDefinition,
fds: Arc<[FieldDefinition]>,
rel_filters: Arc<[RelationFieldInfo]>,
kvs: Arc<Datastore>,
) -> Field {
let tb_name = tb.name.clone();
let tb_name_str = tb_name.as_str().to_string();
let table_order_name = format!("_order_{tb_name}");
let table_filter_name = filter_name_from_table(&tb_name);
let field_name = super::naming::list_field_name(tb);
Field::new(field_name, TypeRef::named_nn_list_nn(&tb_name_str), move |ctx| {
let tb_name = tb_name.clone();
let fds = Arc::clone(&fds);
let rel_filters = Arc::clone(&rel_filters);
let kvs = Arc::clone(&kvs);
FieldFuture::new(async move {
let sess = ctx.data::<Arc<Session>>()?;
let args = ctx.args.as_index_map();
trace!("received request with args: {args:?}");
let start = parse_start_arg(args);
let limit = parse_limit_arg(args);
let version = parse_version_arg(args)?;
let order = parse_order_arg(args, &fds)?;
let tb_name_str_ref = tb_name.as_str();
let cond = parse_filter_arg(args, &fds, tb_name_str_ref, &rel_filters)?;
trace!("parsed order: {order:?}");
trace!("parsed filter: {cond:?}");
let stmt =
select_all_from_table(Expr::Table(tb_name), cond, order, limit, start, &version);
let res = execute_select(&kvs, sess, stmt).await?;
match res {
Value::Array(a) => objects_to_cached_records(a, version),
v => {
error!("Found top level value, in result which should be array: {v:?}");
Err(internal_error("Unexpected result type from table query").into())
}
}
})
})
.description(
super::naming::description_with_deprecation(
tb.comment.as_deref(),
tb.graphql_deprecated.as_deref(),
)
.unwrap_or_else(|| {
format!("Generated from table `{}`\nallows querying a table with filters", tb.name)
}),
)
.argument(InputValue::new("limit", TypeRef::named(TypeRef::INT)))
.argument(InputValue::new("start", TypeRef::named(TypeRef::INT)))
.argument(InputValue::new("order", TypeRef::named(&table_order_name)))
.argument(InputValue::new("filter", TypeRef::named(&table_filter_name)))
.argument(InputValue::new("where", TypeRef::named(&table_filter_name)))
.argument(InputValue::new("version", TypeRef::named(TypeRef::STRING)))
}
fn make_table_get_field(tb: &TableDefinition, kvs: Arc<Datastore>) -> Field {
let tb_name = tb.name.clone();
let tb_name_str = tb_name.as_str().to_string();
let field_name = super::naming::get_field_name(tb);
Field::new(field_name, TypeRef::named(&tb_name_str), move |ctx| {
let tb_name = tb_name.clone();
let kvs = Arc::clone(&kvs);
FieldFuture::new(async move {
let sess = ctx.data::<Arc<Session>>()?;
let args = ctx.args.as_index_map();
let id = match args.get("id").and_then(GraphqlValueUtils::as_string) {
Some(i) => i,
None => {
return Err(
internal_error("Schema validation failed: No id found in _get_").into()
);
}
};
let version = parse_version_arg(args)?;
let rid_str = format!("{tb_name}:{id}");
let record_id: RecordId = match crate::syn::record_id(&rid_str) {
Ok(x) => x.into(),
Err(_) => RecordId::new(tb_name, id),
};
let stmt = select_all_from_record(&record_id, &version);
let res = execute_select(&kvs, sess, stmt).await?;
match res {
Value::Object(obj) => {
let rid = match obj.get("id") {
Some(Value::RecordId(rid)) => rid.clone(),
_ => return Ok(None),
};
Ok(Some(FieldValue::owned_any(CachedRecord {
rid,
version,
data: obj,
})))
}
_ => Ok(None),
}
})
})
.description(
super::naming::description_with_deprecation(
tb.comment.as_deref(),
tb.graphql_deprecated.as_deref(),
)
.unwrap_or_else(|| {
format!(
"Generated from table `{}`\nallows querying a single record in a table by ID",
tb.name
)
}),
)
.argument(InputValue::new("id", TypeRef::named_nn(TypeRef::ID)))
.argument(InputValue::new("version", TypeRef::named(TypeRef::STRING)))
}
fn make_generic_get_field(kvs: Arc<Datastore>) -> Field {
Field::new("_get", TypeRef::named("record"), move |ctx| {
let kvs = Arc::clone(&kvs);
FieldFuture::new(async move {
let sess = ctx.data::<Arc<Session>>()?;
let args = ctx.args.as_index_map();
let id = match args.get("id").and_then(GraphqlValueUtils::as_string) {
Some(i) => i,
None => {
return Err(
internal_error("Schema validation failed: No id found in _get").into()
);
}
};
let version = parse_version_arg(args)?;
let record_id: RecordId = match crate::syn::record_id(&id) {
Ok(x) => x.into(),
Err(_) => {
return Err(resolver_error("Invalid record ID format").into());
}
};
let stmt = select_all_from_record(&record_id, &version);
let res = execute_select(&kvs, sess, stmt).await?;
match res {
Value::Object(obj) => {
let rid = match obj.get("id") {
Some(Value::RecordId(rid)) => rid.clone(),
_ => return Ok(None),
};
let table_name = rid.table.clone();
Ok(Some(
FieldValue::owned_any(CachedRecord {
rid,
version,
data: obj,
})
.with_type(table_name),
))
}
_ => Ok(None),
}
})
})
.description("Allows fetching arbitrary records".to_string())
.argument(InputValue::new("id", TypeRef::named_nn(TypeRef::ID)))
.argument(InputValue::new("version", TypeRef::named(TypeRef::STRING)))
}
struct TableGraphQLTypes {
ty_obj: Object,
orderable: Enum,
order: InputObject,
filter: InputObject,
rel_filters: Vec<RelationFieldInfo>,
}
fn build_table_type(
tb: &TableDefinition,
fds: &[FieldDefinition],
relations: &[RelationInfo],
exposed_table_names: &HashSet<TableName>,
relation_table_fds: &HashMap<TableName, Arc<[FieldDefinition]>>,
types: &mut Vec<Type>,
) -> Result<TableGraphQLTypes, GraphqlError> {
let tb_name = &tb.name;
let tb_name_str = tb_name.as_str().to_string();
let table_orderable_name = format!("_orderable_{tb_name}");
let table_order_name = format!("_order_{tb_name}");
let table_filter_name = filter_name_from_table(tb_name);
let mut orderable = Enum::new(&table_orderable_name).item("id").description(format!(
"Generated from `{tb_name}` the fields which a query can be ordered by"
));
let order = InputObject::new(&table_order_name)
.description(format!("Generated from `{tb_name}` an object representing a query ordering"))
.field(InputValue::new("asc", TypeRef::named(&table_orderable_name)))
.field(InputValue::new("desc", TypeRef::named(&table_orderable_name)))
.field(InputValue::new("then", TypeRef::named(&table_order_name)));
let mut filter = InputObject::new(&table_filter_name)
.field(InputValue::new("id", TypeRef::named("_filter_id")))
.field(InputValue::new("and", TypeRef::named_nn_list(&table_filter_name)))
.field(InputValue::new("or", TypeRef::named_nn_list(&table_filter_name)))
.field(InputValue::new("not", TypeRef::named(&table_filter_name)));
let mut ty_obj = Object::new(&tb_name_str)
.field(Field::new(
"id",
TypeRef::named_nn(TypeRef::ID),
make_table_field_resolver("id", Some(Kind::Record(vec![tb_name.clone()])), None),
))
.implement("record");
let mut existing_field_names: HashSet<String> = HashSet::new();
existing_field_names.insert("id".to_string());
let nested_objects = detect_nested_objects(&tb_name_str, fds);
for fd in fds.iter() {
let Some(ref kind) = fd.field_kind else {
continue;
};
if fd.name.is_id() {
continue;
}
if fd.name.0.len() > 1 {
continue;
}
let lookup_name = idiom_to_graphql_name(&fd.name);
let fd_name = Name::new(super::naming::field_graphql_name(fd));
existing_field_names.insert(fd_name.to_string());
if let Some(nested) = nested_objects.get(lookup_name.as_str()) {
let nested_type =
make_nested_object_type(&nested.graphql_type_name, &nested.sub_fields, types)?;
types.push(Type::Object(nested_type));
let fd_type = if nested.is_array {
let list = TypeRef::List(Box::new(TypeRef::named_nn(&nested.graphql_type_name)));
if nested.optional {
list
} else {
TypeRef::NonNull(Box::new(list))
}
} else if nested.optional {
TypeRef::named(&nested.graphql_type_name)
} else {
TypeRef::named_nn(&nested.graphql_type_name)
};
orderable = orderable.item(fd_name.to_string());
let mut field = Field::new(
fd_name.as_str(),
fd_type,
make_nested_object_field_resolver(lookup_name.clone(), nested.is_array),
);
field = field.description(if let Some(ref c) = fd.comment {
c.clone()
} else {
format!("Nested object field `{}`", fd_name.as_str())
});
ty_obj = ty_obj.field(field);
continue;
}
let enum_scope = format!("{}_{}", tb_name_str, fd_name);
let fd_type = kind_to_type_with_enum_prefix(kind.clone(), types, false, Some(&enum_scope))?;
orderable = orderable.item(fd_name.to_string());
let type_filter_name = format!("_filter_{}", filter_type_name(&fd_type));
let filter_already_exists = types.iter().any(|t| match t {
Type::InputObject(io) => io.type_name() == type_filter_name,
_ => false,
});
if !filter_already_exists {
let type_filter = Type::InputObject(filter_from_type(
kind,
type_filter_name.clone(),
types,
Some(&enum_scope),
)?);
trace!("\n{type_filter:?}\n");
types.push(type_filter);
}
filter = filter.field(InputValue::new(fd_name.as_str(), TypeRef::named(type_filter_name)));
let mut field = Field::new(
fd_name.as_str(),
fd_type,
make_table_field_resolver(&lookup_name, fd.field_kind.clone(), Some(enum_scope)),
);
if let Some(desc) = super::naming::description_with_deprecation(
fd.comment.as_deref(),
fd.graphql_deprecated.as_deref(),
) {
field = field.description(desc);
}
ty_obj = ty_obj.field(field);
}
let mut rel_filters: Vec<RelationFieldInfo> = Vec::new();
for rel in relations.iter() {
if !exposed_table_names.contains(&rel.table_name) {
continue;
}
let rel_table_str = rel.table_name.as_str().to_owned();
let rel_fds = relation_table_fds.get(&rel.table_name).cloned();
if rel.from_tables.iter().any(|n| n.as_str() == tb_name_str.as_str()) {
let field_name = rel_table_str.clone();
if !existing_field_names.contains(&field_name) {
existing_field_names.insert(field_name.clone());
ty_obj = ty_obj.field(make_relation_field(
&field_name,
&rel_table_str,
rel.table_name.clone(),
RelationDirection::Outgoing,
rel_fds.clone(),
));
let rel_filter_name =
register_relation_filter(&tb_name_str, &field_name, "out", types);
filter = filter
.field(InputValue::new(field_name.clone(), TypeRef::named(rel_filter_name)));
rel_filters.push(RelationFieldInfo {
field_name,
relation_table: rel.table_name.clone(),
dir: Dir::Out,
});
} else {
trace!(
"Skipping outgoing relation field '{}' on table '{}': \
conflicts with existing field",
field_name, tb_name_str
);
}
}
if rel.to_tables.iter().any(|n| n.as_str() == tb_name_str.as_str()) {
let field_name = format!("{}_in", rel_table_str);
if !existing_field_names.contains(&field_name) {
existing_field_names.insert(field_name.clone());
ty_obj = ty_obj.field(make_relation_field(
&field_name,
&rel_table_str,
rel.table_name.clone(),
RelationDirection::Incoming,
rel_fds.clone(),
));
let rel_filter_name =
register_relation_filter(&tb_name_str, &field_name, "in", types);
filter = filter
.field(InputValue::new(field_name.clone(), TypeRef::named(rel_filter_name)));
rel_filters.push(RelationFieldInfo {
field_name,
relation_table: rel.table_name.clone(),
dir: Dir::In,
});
} else {
trace!(
"Skipping incoming relation field '{}' on table '{}': \
conflicts with existing field",
field_name, tb_name_str
);
}
}
}
Ok(TableGraphQLTypes {
ty_obj,
orderable,
order,
filter,
rel_filters,
})
}
fn register_relation_filter(
tb_name: &str,
field_name: &str,
dir_token: &str,
types: &mut Vec<Type>,
) -> String {
let name = format!("_relation_filter_{tb_name}_{field_name}_{dir_token}");
if types.iter().any(|t| match t {
Type::InputObject(io) => io.type_name() == name,
_ => false,
}) {
return name;
}
let io = InputObject::new(&name)
.description(format!(
"Filter predicates evaluated against the `{field_name}` relation traversal of `{tb_name}`. \
Only `count` is supported."
))
.field(InputValue::new("count", TypeRef::named(COUNT_FILTER_INPUT)));
types.push(Type::InputObject(io));
name
}
pub async fn process_tbs(
tbs: Arc<[TableDefinition]>,
mut query: Object,
types: &mut Vec<Type>,
ctx: &SchemaContext<'_>,
relations: &[RelationInfo],
table_fields: &mut HashMap<TableName, Arc<[FieldDefinition]>>,
) -> Result<Object, GraphqlError> {
let mut relation_table_fds: HashMap<TableName, Arc<[FieldDefinition]>> = HashMap::new();
for rel in relations.iter() {
if let std::collections::hash_map::Entry::Vacant(e) =
relation_table_fds.entry(rel.table_name.clone())
{
let fds = ctx.tx.all_tb_fields(ctx.ns, ctx.db, &rel.table_name, None).await?;
e.insert(fds);
}
}
let exposed_table_names: HashSet<TableName> = tbs.iter().map(|t| t.name.clone()).collect();
{
const BUILTIN: &str = "<built-in>";
let mut seen: HashMap<String, String> = HashMap::new();
let mut seen_types: HashMap<String, String> = HashMap::new();
for reserved in ["__schema", "__type", "__typename"] {
seen.insert(reserved.to_owned(), BUILTIN.to_owned());
}
for reserved in [
"Query",
"Mutation",
"Subscription",
PAGE_INFO_TYPE,
ID_RANGE_INPUT,
COUNT_FILTER_INPUT,
VECTOR_DISTANCE_ENUM,
NUM_OP_ENUM,
KNN_INPUT,
SIMILARITY_INPUT,
MATCHES_INPUT,
CALL_INPUT,
"_filter_id",
] {
seen_types.insert(reserved.to_owned(), BUILTIN.to_owned());
}
for tb in tbs.iter() {
let list = super::naming::list_field_name(tb);
let get = super::naming::get_field_name(tb);
let conn_field = format!("{list}Connection");
let aggregate = format!("{}_aggregate", tb.name.as_str());
for name in [list, get, conn_field, aggregate] {
if let Some(prior) = seen.get(&name) {
let prior_desc = if prior == BUILTIN {
format!("built-in query field `{name}`")
} else {
format!("table `{prior}`")
};
return Err(super::error::schema_error(format!(
"GraphQL naming collision on `{}` — {} and table `{}` produce the \
same query field. Set an explicit `GRAPHQL_ALIAS` on the table.",
name, prior_desc, tb.name
)));
}
seen.insert(name, tb.name.as_str().to_owned());
}
let tb_ty = tb.name.as_str().to_owned();
let conn_ty = connection_type_name(tb);
let edge_ty = edge_type_name(tb);
for ty_name in [tb_ty, conn_ty, edge_ty] {
if let Some(prior) = seen_types.get(&ty_name) {
let prior_desc = if prior == BUILTIN {
format!("built-in type `{ty_name}`")
} else {
format!("table `{prior}`")
};
return Err(super::error::schema_error(format!(
"GraphQL naming collision on type `{}` — {} and table `{}` produce \
the same type. Set an explicit `GRAPHQL_ALIAS` on the table.",
ty_name, prior_desc, tb.name
)));
}
seen_types.insert(ty_name, tb.name.as_str().to_owned());
}
}
}
for tb in tbs.iter() {
trace!("Adding table: {}", tb.name);
let fds = ctx.tx.all_tb_fields(ctx.ns, ctx.db, &tb.name, None).await?;
table_fields.insert(tb.name.clone(), Arc::clone(&fds));
let tt = build_table_type(
tb,
&fds,
relations,
&exposed_table_names,
&relation_table_fds,
types,
)?;
let rel_filters: Arc<[RelationFieldInfo]> = tt.rel_filters.into();
types.push(Type::Object(tt.ty_obj));
types.push(tt.order.into());
types.push(Type::Enum(tt.orderable));
types.push(Type::InputObject(tt.filter));
query = query.field(make_table_list_field(
tb,
Arc::clone(&fds),
Arc::clone(&rel_filters),
Arc::clone(ctx.datastore),
));
query = query.field(make_table_get_field(tb, Arc::clone(ctx.datastore)));
let (agg_obj, agg_enum) = build_aggregate_type(tb.name.as_str(), &fds, types);
types.push(Type::Object(agg_obj));
types.push(Type::Enum(agg_enum));
query = query.field(make_table_aggregate_field(
tb,
Arc::clone(&fds),
Arc::clone(&rel_filters),
Arc::clone(ctx.datastore),
));
build_connection_types(tb, types);
query = query.field(make_table_connection_field(
tb,
Arc::clone(&fds),
Arc::clone(&rel_filters),
Arc::clone(ctx.datastore),
));
}
query = query.field(make_generic_get_field(Arc::clone(ctx.datastore)));
Ok(query)
}
fn make_table_field_resolver(
fd_name: impl Into<String>,
kind: Option<Kind>,
enum_scope: Option<String>,
) -> impl for<'a> Fn(ResolverContext<'a>) -> FieldFuture<'a> + Send + Sync + 'static {
let fd_name = fd_name.into();
move |ctx: ResolverContext| {
let fd_name = fd_name.clone();
let field_kind = kind.clone();
let enum_scope = enum_scope.clone();
FieldFuture::new({
async move {
if let Ok(cached) = ctx.parent_value.try_downcast_ref::<CachedRecord>() {
return resolve_field_from_cached_record(
&ctx,
cached,
&fd_name,
&field_kind,
enum_scope.as_deref(),
)
.await;
}
let ds = ctx.data::<Arc<Datastore>>()?;
let sess = ctx.data::<Arc<Session>>()?;
let (rid, version) = match ctx.parent_value.try_downcast_ref::<VersionedRecord>() {
Ok(vr) => (vr.rid.clone(), vr.version),
Err(_) => {
let rid = ctx.parent_value.try_downcast_ref::<RecordId>()?;
(rid.clone(), None)
}
};
let stmt = select_field_from_record(&rid, &fd_name, &version);
let val = execute_select(ds, sess, stmt).await?;
resolve_field_value(
&ctx,
val,
&fd_name,
&field_kind,
&version,
enum_scope.as_deref(),
)
.await
}
})
}
}
async fn resolve_field_value(
ctx: &ResolverContext<'_>,
val: Value,
fd_name: &str,
field_kind: &Option<Kind>,
version: &Option<Datetime>,
enum_scope: Option<&str>,
) -> Result<Option<FieldValue<'static>>, async_graphql::Error> {
match val {
Value::RecordId(target_rid) if fd_name != "id" => {
let ds = ctx.data::<Arc<Datastore>>()?;
let sess = ctx.data::<Arc<Session>>()?;
let stmt = select_all_from_record(&target_rid, version);
let target_val = execute_select(ds, sess, stmt).await?;
match target_val {
Value::Object(obj) => {
let field_val = FieldValue::owned_any(CachedRecord {
rid: target_rid.clone(),
version: *version,
data: obj,
});
let field_val = match field_kind {
Some(Kind::Record(ts)) if ts.is_empty() || ts.len() > 1 => {
field_val.with_type(target_rid.table)
}
_ => field_val,
};
Ok(Some(field_val))
}
Value::None | Value::Null => Ok(None),
_ => Ok(None),
}
}
Value::Geometry(g) => {
let type_name = geometry_graphql_type_name(&g);
let field_val = FieldValue::owned_any(g);
let field_val = match field_kind {
Some(Kind::Geometry(ks)) if ks.is_empty() || ks.len() > 1 => {
field_val.with_type(type_name)
}
_ => field_val,
};
Ok(Some(field_val))
}
Value::None | Value::Null => Ok(None),
v => {
let out = sql_value_to_graphql_value_with_kind(v, field_kind.as_ref(), enum_scope)
.map_err(async_graphql::Error::from)?;
Ok(Some(FieldValue::value(out)))
}
}
}
async fn resolve_field_from_cached_record(
ctx: &ResolverContext<'_>,
cached: &CachedRecord,
fd_name: &str,
field_kind: &Option<Kind>,
enum_scope: Option<&str>,
) -> Result<Option<FieldValue<'static>>, async_graphql::Error> {
let val = cached.data.get(fd_name).cloned().unwrap_or(Value::None);
resolve_field_value(ctx, val, fd_name, field_kind, &cached.version, enum_scope).await
}
fn make_relation_field(
field_name: &str,
rel_table_type_name: &str,
rel_table_name: TableName,
direction: RelationDirection,
rel_fds: Option<Arc<[FieldDefinition]>>,
) -> Field {
let table_filter_name = filter_name_from_table(rel_table_type_name);
let table_order_name = format!("_order_{}", rel_table_type_name);
let desc = match direction {
RelationDirection::Outgoing => {
format!("Outgoing `{}` relations from this record", rel_table_type_name)
}
RelationDirection::Incoming => {
format!("Incoming `{}` relations to this record", rel_table_type_name)
}
};
Field::new(
field_name,
TypeRef::named_nn_list_nn(rel_table_type_name),
make_relation_field_resolver(rel_table_name, direction, rel_fds),
)
.description(desc)
.argument(InputValue::new("limit", TypeRef::named(TypeRef::INT)))
.argument(InputValue::new("start", TypeRef::named(TypeRef::INT)))
.argument(InputValue::new("order", TypeRef::named(&table_order_name)))
.argument(InputValue::new("filter", TypeRef::named(&table_filter_name)))
.argument(InputValue::new("where", TypeRef::named(&table_filter_name)))
}
fn make_relation_field_resolver(
relation_table_name: TableName,
direction: RelationDirection,
rel_fds: Option<Arc<[FieldDefinition]>>,
) -> impl for<'a> Fn(ResolverContext<'a>) -> FieldFuture<'a> + Send + Sync + 'static {
move |ctx: ResolverContext| {
let relation_table = relation_table_name.clone();
let fds = rel_fds.clone();
FieldFuture::new(async move {
let ds = ctx.data::<Arc<Datastore>>()?;
let sess = ctx.data::<Arc<Session>>()?;
let (rid, version) =
if let Ok(cached) = ctx.parent_value.try_downcast_ref::<CachedRecord>() {
(cached.rid.clone(), cached.version)
} else if let Ok(vr) = ctx.parent_value.try_downcast_ref::<VersionedRecord>() {
(vr.rid.clone(), vr.version)
} else {
let rid = ctx.parent_value.try_downcast_ref::<RecordId>()?;
(rid.clone(), None)
};
let args = ctx.args.as_index_map();
let start = parse_start_arg(args);
let limit = parse_limit_arg(args);
let order = parse_order_arg(args, fds.as_deref().unwrap_or(&[]))?;
let filter_field = match direction {
RelationDirection::Outgoing => "in",
RelationDirection::Incoming => "out",
};
let mut base_cond = Expr::Binary {
left: Box::new(Expr::Idiom(Idiom::field(filter_field.to_string()))),
op: BinaryOperator::Equal,
right: Box::new(Value::RecordId(rid.clone()).into_literal()),
};
if let Some(ref fds) = fds
&& let Some(user_cond) = parse_filter_arg(args, fds, relation_table.as_str(), &[])?
{
base_cond = Expr::Binary {
left: Box::new(base_cond),
op: BinaryOperator::And,
right: Box::new(user_cond.0),
};
}
let cond = Some(Cond(base_cond));
let stmt = select_all_from_table(
Expr::Table(relation_table),
cond,
order,
limit,
start,
&version,
);
let res = execute_select(ds, sess, stmt).await?;
match res {
Value::Array(a) => objects_to_cached_records(a, version),
v => {
error!("Expected array result for relation query, found: {v:?}");
Err(internal_error("Unexpected result type for relation query").into())
}
}
})
}
}
macro_rules! filter_impl {
($filter:ident, $ty:ident, $name:expr_2021) => {
$filter = $filter.field(InputValue::new($name, $ty.clone()));
};
}
fn filter_id() -> InputObject {
let mut filter = InputObject::new("_filter_id");
let ty = TypeRef::named(TypeRef::ID);
filter_impl!(filter, ty, "eq");
filter_impl!(filter, ty, "ne");
filter_impl!(filter, ty, "gt");
filter_impl!(filter, ty, "gte");
filter_impl!(filter, ty, "lt");
filter_impl!(filter, ty, "lte");
let list_ty = TypeRef::named_nn_list(TypeRef::ID);
filter_impl!(filter, list_ty, "in");
let range_ty = TypeRef::named(ID_RANGE_INPUT);
filter_impl!(filter, range_ty, "range");
filter
}
fn filter_from_type(
kind: &Kind,
filter_name: String,
types: &mut Vec<Type>,
enum_scope: Option<&str>,
) -> Result<InputObject, GraphqlError> {
let effective_kind = match kind {
Kind::Either(ks) => {
let non_none: Vec<&Kind> =
ks.iter().filter(|k| !matches!(k, Kind::None | Kind::Null)).collect();
if non_none.len() == 1 {
non_none[0].clone()
} else {
kind.clone()
}
}
_ => kind.clone(),
};
let (eq_ne_ty, supports_eq_ne) = match &effective_kind {
Kind::Record(ts) => match ts.len() {
1 => (
TypeRef::named(filter_name_from_table(
ts.first().expect("ts should have exactly one element").as_str(),
)),
true,
),
_ => (TypeRef::named(TypeRef::ID), true),
},
Kind::Array(inner, _) if matches!(**inner, Kind::Record(_)) => {
(TypeRef::named(TypeRef::ID), false)
}
k => {
(unwrap_type(kind_to_type_with_enum_prefix(k.clone(), types, true, enum_scope)?), true)
}
};
let mut filter = InputObject::new(filter_name);
if supports_eq_ne {
filter_impl!(filter, eq_ne_ty, "eq");
filter_impl!(filter, eq_ne_ty, "ne");
}
let call_ty = TypeRef::named(CALL_INPUT);
filter_impl!(filter, call_ty, "call");
if numeric_array_inner(&effective_kind).is_some() {
let knn_ty = TypeRef::named(KNN_INPUT);
filter_impl!(filter, knn_ty, "nearest");
let sim_ty = TypeRef::named(SIMILARITY_INPUT);
filter_impl!(filter, sim_ty, "similarity");
}
match effective_kind {
Kind::String => {
let str_ty = TypeRef::named(TypeRef::STRING);
filter_impl!(filter, str_ty, "contains");
filter_impl!(filter, str_ty, "startsWith");
filter_impl!(filter, str_ty, "endsWith");
filter_impl!(filter, str_ty, "regex");
let list_ty = TypeRef::named_nn_list(TypeRef::STRING);
filter_impl!(filter, list_ty, "in");
let matches_ty = TypeRef::named(MATCHES_INPUT);
filter_impl!(filter, matches_ty, "matches");
}
Kind::Int => {
let num_ty = TypeRef::named(TypeRef::INT);
filter_impl!(filter, num_ty, "gt");
filter_impl!(filter, num_ty, "gte");
filter_impl!(filter, num_ty, "lt");
filter_impl!(filter, num_ty, "lte");
let list_ty = TypeRef::named_nn_list(TypeRef::INT);
filter_impl!(filter, list_ty, "in");
}
Kind::Float => {
let num_ty = TypeRef::named(TypeRef::FLOAT);
filter_impl!(filter, num_ty, "gt");
filter_impl!(filter, num_ty, "gte");
filter_impl!(filter, num_ty, "lt");
filter_impl!(filter, num_ty, "lte");
let list_ty = TypeRef::named_nn_list(TypeRef::FLOAT);
filter_impl!(filter, list_ty, "in");
}
Kind::Number => {
let num_ty = TypeRef::named("number");
filter_impl!(filter, num_ty, "gt");
filter_impl!(filter, num_ty, "gte");
filter_impl!(filter, num_ty, "lt");
filter_impl!(filter, num_ty, "lte");
let list_ty = TypeRef::named_nn_list("number");
filter_impl!(filter, list_ty, "in");
}
Kind::Decimal => {
let num_ty = TypeRef::named("decimal");
filter_impl!(filter, num_ty, "gt");
filter_impl!(filter, num_ty, "gte");
filter_impl!(filter, num_ty, "lt");
filter_impl!(filter, num_ty, "lte");
let list_ty = TypeRef::named_nn_list("decimal");
filter_impl!(filter, list_ty, "in");
}
Kind::Datetime => {
let dt_ty = TypeRef::named("datetime");
filter_impl!(filter, dt_ty, "gt");
filter_impl!(filter, dt_ty, "gte");
filter_impl!(filter, dt_ty, "lt");
filter_impl!(filter, dt_ty, "lte");
}
Kind::Record(_) => {
let list_ty = TypeRef::named_nn_list(TypeRef::ID);
filter_impl!(filter, list_ty, "in");
}
Kind::Array(ref inner, _) if matches!(**inner, Kind::Record(_)) => {
let id_ty = TypeRef::named(TypeRef::ID);
filter_impl!(filter, id_ty, "contains");
}
Kind::Any
| Kind::None
| Kind::Null
| Kind::Bool
| Kind::Bytes
| Kind::Duration
| Kind::Object
| Kind::Uuid
| Kind::Regex
| Kind::Table(_)
| Kind::Geometry(_)
| Kind::Either(_)
| Kind::Set(_, _)
| Kind::Array(_, _)
| Kind::Function(_, _)
| Kind::Range
| Kind::Literal(_)
| Kind::File(_) => {}
};
Ok(filter)
}
pub(super) fn cond_from_filter(
filter: &IndexMap<Name, GraphqlValue>,
fds: &[FieldDefinition],
tb_name: &str,
relations: &[RelationFieldInfo],
) -> Result<Cond, GraphqlError> {
val_from_filter(filter, fds, tb_name, relations).map(Cond)
}
fn val_from_filter(
filter: &IndexMap<Name, GraphqlValue>,
fds: &[FieldDefinition],
tb_name: &str,
relations: &[RelationFieldInfo],
) -> Result<Expr, GraphqlError> {
if filter.is_empty() {
return Err(resolver_error("Table filter must have at least one item"));
}
if filter.len() == 1 {
let (k, v) = filter.iter().next().expect("filter has exactly one item");
return match k.as_str().to_lowercase().as_str() {
"or" => aggregate(v, AggregateOp::Or, fds, tb_name, relations),
"and" => aggregate(v, AggregateOp::And, fds, tb_name, relations),
"not" => negate(v, fds, tb_name, relations),
_ => binop(k.as_str(), v, fds, tb_name, relations),
};
}
let mut exprs = Vec::with_capacity(filter.len());
for (k, v) in filter.iter() {
let expr = match k.as_str().to_lowercase().as_str() {
"or" => aggregate(v, AggregateOp::Or, fds, tb_name, relations)?,
"and" => aggregate(v, AggregateOp::And, fds, tb_name, relations)?,
"not" => negate(v, fds, tb_name, relations)?,
_ => binop(k.as_str(), v, fds, tb_name, relations)?,
};
exprs.push(expr);
}
let mut iter = exprs.into_iter();
let mut combined = iter.next().expect("at least one filter entry");
for next_expr in iter {
combined = Expr::Binary {
left: Box::new(combined),
op: BinaryOperator::And,
right: Box::new(next_expr),
};
}
Ok(combined)
}
fn parse_binary_op(name: &str) -> Option<expr::BinaryOperator> {
match name {
"eq" => Some(expr::BinaryOperator::Equal),
"ne" => Some(expr::BinaryOperator::NotEqual),
"gt" => Some(expr::BinaryOperator::MoreThan),
"gte" => Some(expr::BinaryOperator::MoreThanEqual),
"lt" => Some(expr::BinaryOperator::LessThan),
"lte" => Some(expr::BinaryOperator::LessThanEqual),
"in" => Some(expr::BinaryOperator::Inside),
_ => None,
}
}
fn parse_function_op(name: &str) -> Option<&'static str> {
match name {
"contains" => Some("string::contains"),
"startsWith" => Some("string::starts_with"),
"endsWith" => Some("string::ends_with"),
"regex" => Some("string::matches"),
_ => None,
}
}
fn negate(
filter: &GraphqlValue,
fds: &[FieldDefinition],
tb_name: &str,
relations: &[RelationFieldInfo],
) -> Result<Expr, GraphqlError> {
let obj = filter.as_object().ok_or(resolver_error("Value of NOT must be object"))?;
let inner_cond = val_from_filter(obj, fds, tb_name, relations)?;
Ok(Expr::Prefix {
op: expr::PrefixOperator::Not,
expr: Box::new(inner_cond),
})
}
#[derive(Clone, Copy)]
enum AggregateOp {
And,
Or,
}
fn aggregate(
filter: &GraphqlValue,
op: AggregateOp,
fds: &[FieldDefinition],
tb_name: &str,
relations: &[RelationFieldInfo],
) -> Result<Expr, GraphqlError> {
let op_str = match op {
AggregateOp::And => "AND",
AggregateOp::Or => "OR",
};
let op = match op {
AggregateOp::And => BinaryOperator::And,
AggregateOp::Or => BinaryOperator::Or,
};
let list =
filter.as_list().ok_or(resolver_error(format!("Value of {op_str} should be a list")))?;
let filter_arr = list
.iter()
.map(|v| v.as_object().map(|o| val_from_filter(o, fds, tb_name, relations)))
.collect::<Option<Result<Vec<Expr>, GraphqlError>>>()
.ok_or(resolver_error(format!("List of {op_str} should contain objects")))??;
let mut iter = filter_arr.into_iter();
let mut cond = iter
.next()
.ok_or(resolver_error(format!("List of {op_str} should contain at least one object")))?;
for clause in iter {
cond = Expr::Binary {
left: Box::new(clause),
op: op.clone(),
right: Box::new(cond),
}
}
Ok(cond)
}
fn binop(
field_name: &str,
val: &GraphqlValue,
fds: &[FieldDefinition],
tb_name: &str,
relations: &[RelationFieldInfo],
) -> Result<Expr, GraphqlError> {
let obj = val.as_object().ok_or(resolver_error("Field filter should be object"))?;
let Some(fd) = fds.iter().find(|fd| field_graphql_name(fd) == field_name) else {
if field_name == "id" {
return binop_for_id(obj);
}
if let Some(rel) = relations.iter().find(|r| r.field_name == field_name) {
return binop_for_relation(rel, obj);
}
return Err(resolver_error(format!("Field `{field_name}` not found")));
};
if obj.is_empty() {
return Err(resolver_error("Field filter must have at least one operator"));
}
let lookup_name = idiom_to_graphql_name(&fd.name);
let field_kind = fd.field_kind.clone().unwrap_or(Kind::Any);
let enum_scope = format!("{tb_name}_{field_name}");
let mut exprs = Vec::with_capacity(obj.len());
for (k, v) in obj.iter() {
let op_name = k.as_str();
let lhs = Expr::Idiom(Idiom::field(lookup_name.clone()));
if let Some(binary_op) = parse_binary_op(op_name) {
let rhs_kind = if op_name == "in" {
Kind::Array(Box::new(field_kind.clone()), None)
} else {
field_kind.clone()
};
let rhs = graphql_to_sql_kind_with_scope(v, rhs_kind, Some(&enum_scope))?;
exprs.push(Expr::Binary {
left: Box::new(lhs),
op: binary_op,
right: Box::new(rhs.into_literal()),
});
} else if op_name == "contains"
&& matches!(strip_option(&field_kind), Kind::Array(inner, _) if matches!(*inner, Kind::Record(_)))
{
let rhs = graphql_to_sql_kind(v, Kind::Record(vec![]))?;
exprs.push(Expr::Binary {
left: Box::new(lhs),
op: BinaryOperator::Contain,
right: Box::new(rhs.into_literal()),
});
} else if let Some(fn_name) = parse_function_op(op_name) {
let rhs = graphql_to_sql_kind(v, Kind::String)?;
exprs.push(Expr::FunctionCall(Box::new(FunctionCall {
receiver: Function::Normal(fn_name.to_string()),
arguments: vec![lhs, rhs.into_literal()],
})));
} else {
match op_name {
"nearest" => exprs.push(translate_nearest(lhs, v)?),
"similarity" => exprs.push(translate_similarity(lhs, v)?),
"matches" => exprs.push(translate_matches(lhs, v)?),
"call" => exprs.push(translate_call(lhs, v)?),
_ => {
return Err(resolver_error(format!("Unsupported filter operator: {op_name}")));
}
}
}
}
let mut iter = exprs.into_iter();
let mut combined = iter.next().expect("at least one operator");
for next_expr in iter {
combined = Expr::Binary {
left: Box::new(combined),
op: BinaryOperator::And,
right: Box::new(next_expr),
};
}
Ok(combined)
}
fn binop_for_relation(
rel: &RelationFieldInfo,
obj: &IndexMap<Name, GraphqlValue>,
) -> Result<Expr, GraphqlError> {
if obj.is_empty() {
return Err(resolver_error("Relation filter must have at least one operator"));
}
let lookup = Lookup {
kind: LookupKind::Graph(rel.dir),
what: vec![LookupSubject::Table {
table: rel.relation_table.clone(),
referencing_field: None,
}],
..Default::default()
};
let count_lhs = Expr::FunctionCall(Box::new(FunctionCall {
receiver: Function::Normal("count".to_string()),
arguments: vec![Expr::Idiom(Idiom(vec![Part::Lookup(Box::new(lookup))]))],
}));
let mut exprs: Vec<Expr> = Vec::new();
for (k, v) in obj.iter() {
match k.as_str() {
"count" => {
let count_obj = v
.as_object()
.ok_or(resolver_error("`count` must be an object of operators"))?;
if count_obj.is_empty() {
return Err(resolver_error("`count` filter must have at least one operator"));
}
for (op_name, op_val) in count_obj.iter() {
let Some(binary_op) = parse_binary_op(op_name.as_str()) else {
return Err(resolver_error(format!(
"Unsupported count operator: {op_name}"
)));
};
if matches!(binary_op, BinaryOperator::Inside) {
return Err(resolver_error(
"`in` is not supported on a relation `count` filter",
));
}
let rhs = graphql_to_sql_kind(op_val, Kind::Int)?;
exprs.push(Expr::Binary {
left: Box::new(count_lhs.clone()),
op: binary_op,
right: Box::new(rhs.into_literal()),
});
}
}
other => {
return Err(resolver_error(format!(
"Unsupported relation filter operator: {other}"
)));
}
}
}
let mut iter = exprs.into_iter();
let mut combined = iter.next().expect("at least one relation operator");
for next in iter {
combined = Expr::Binary {
left: Box::new(combined),
op: BinaryOperator::And,
right: Box::new(next),
};
}
Ok(combined)
}
fn translate_id_range(lhs: Expr, val: &GraphqlValue) -> Result<Expr, GraphqlError> {
let obj = val.as_object().ok_or(resolver_error("Value of `range` must be an object"))?;
let from = obj.get("from").filter(|v| !matches!(v, GraphqlValue::Null));
let to = obj.get("to").filter(|v| !matches!(v, GraphqlValue::Null));
let inclusive = obj
.get("inclusive")
.and_then(|v| match v {
GraphqlValue::Boolean(b) => Some(*b),
_ => None,
})
.unwrap_or(false);
if from.is_none() && to.is_none() {
return Err(resolver_error("`range` requires at least one of `from` or `to`"));
}
let mut clauses: Vec<Expr> = Vec::with_capacity(2);
if let Some(f) = from {
let rhs = graphql_to_sql_kind(f, Kind::Record(vec![]))?;
clauses.push(Expr::Binary {
left: Box::new(lhs.clone()),
op: BinaryOperator::MoreThanEqual,
right: Box::new(rhs.into_literal()),
});
}
if let Some(t) = to {
let rhs = graphql_to_sql_kind(t, Kind::Record(vec![]))?;
let op = if inclusive {
BinaryOperator::LessThanEqual
} else {
BinaryOperator::LessThan
};
clauses.push(Expr::Binary {
left: Box::new(lhs),
op,
right: Box::new(rhs.into_literal()),
});
}
let mut iter = clauses.into_iter();
let mut combined = iter.next().expect("at least one range bound");
for next in iter {
combined = Expr::Binary {
left: Box::new(combined),
op: BinaryOperator::And,
right: Box::new(next),
};
}
Ok(combined)
}
fn binop_for_id(obj: &IndexMap<Name, GraphqlValue>) -> Result<Expr, GraphqlError> {
if obj.is_empty() {
return Err(resolver_error("ID filter must have at least one operator"));
}
let mut exprs = Vec::with_capacity(obj.len());
for (k, v) in obj.iter() {
let op_name = k.as_str();
let lhs = Expr::Idiom(Idiom::field("id".to_string()));
if op_name == "in" {
let rhs = graphql_to_sql_kind(v, Kind::Array(Box::new(Kind::Record(vec![])), None))?;
let ids = match rhs {
Value::Array(arr) => arr.0,
_ => return Err(resolver_error("`id.in` expected an array of IDs")),
};
if ids.is_empty() {
return Err(resolver_error("`id.in` must contain at least one ID"));
}
if ids.len() > MAX_ID_IN_LIST {
return Err(resolver_error(format!(
"`id.in` accepts at most {MAX_ID_IN_LIST} IDs (got {})",
ids.len()
)));
}
let mut iter = ids.into_iter();
let first = iter.next().expect("non-empty");
let mut combined = Expr::Binary {
left: Box::new(lhs.clone()),
op: BinaryOperator::Equal,
right: Box::new(first.into_literal()),
};
for next in iter {
combined = Expr::Binary {
left: Box::new(combined),
op: BinaryOperator::Or,
right: Box::new(Expr::Binary {
left: Box::new(lhs.clone()),
op: BinaryOperator::Equal,
right: Box::new(next.into_literal()),
}),
};
}
exprs.push(combined);
} else if let Some(binary_op) = parse_binary_op(op_name) {
let rhs = graphql_to_sql_kind(v, Kind::Record(vec![]))?;
exprs.push(Expr::Binary {
left: Box::new(lhs),
op: binary_op,
right: Box::new(rhs.into_literal()),
});
} else if op_name == "range" {
exprs.push(translate_id_range(lhs, v)?);
} else {
return Err(resolver_error(format!("Unsupported ID filter operator: {op_name}")));
}
}
let mut iter = exprs.into_iter();
let mut combined = iter.next().expect("at least one operator");
for next_expr in iter {
combined = Expr::Binary {
left: Box::new(combined),
op: BinaryOperator::And,
right: Box::new(next_expr),
};
}
Ok(combined)
}
const VECTOR_DISTANCE_ENUM: &str = "_VectorDistance";
const NUM_OP_ENUM: &str = "_NumOp";
const KNN_INPUT: &str = "_KnnInput";
const SIMILARITY_INPUT: &str = "_SimilarityInput";
const MATCHES_INPUT: &str = "_MatchesInput";
const CALL_INPUT: &str = "_CallInput";
const ID_RANGE_INPUT: &str = "_IdRangeInput";
const COUNT_FILTER_INPUT: &str = "_CountFilterInput";
const PAGE_INFO_TYPE: &str = "PageInfo";
pub(crate) fn register_filter_helper_types(types: &mut Vec<Type>) {
types.push(Type::Enum(
Enum::new(VECTOR_DISTANCE_ENUM)
.description(
"Vector distance / similarity metric for the `nearest` and `similarity` filter \
operators.",
)
.item("COSINE")
.item("EUCLIDEAN")
.item("MANHATTAN")
.item("HAMMING")
.item("JACCARD")
.item("CHEBYSHEV")
.item("PEARSON"),
));
types.push(Type::Enum(
Enum::new(NUM_OP_ENUM)
.description("Comparison operator for the `similarity` and `call` filter operators.")
.item("eq")
.item("ne")
.item("gt")
.item("gte")
.item("lt")
.item("lte"),
));
types.push(Type::InputObject(
InputObject::new(KNN_INPUT)
.description(
"K-nearest-neighbour predicate. Translates to SurrealQL `field <|k,distance|> to`.",
)
.field(InputValue::new("to", TypeRef::named_nn_list_nn(TypeRef::FLOAT)))
.field(InputValue::new("k", TypeRef::named_nn(TypeRef::INT)))
.field(InputValue::new("distance", TypeRef::named_nn(VECTOR_DISTANCE_ENUM))),
));
types.push(Type::InputObject(
InputObject::new(SIMILARITY_INPUT)
.description(
"Vector similarity / distance predicate. Calls the matching `vector::*` function \
on the field and `to`, then compares the result against `value` using `op`.",
)
.field(InputValue::new("to", TypeRef::named_nn_list_nn(TypeRef::FLOAT)))
.field(InputValue::new("distance", TypeRef::named_nn(VECTOR_DISTANCE_ENUM)))
.field(InputValue::new("op", TypeRef::named_nn(NUM_OP_ENUM)))
.field(InputValue::new("value", TypeRef::named_nn(TypeRef::FLOAT))),
));
types.push(Type::InputObject(
InputObject::new(MATCHES_INPUT)
.description(
"Full-text-search predicate. Translates to SurrealQL `field @@ query`. Requires \
a `DEFINE INDEX … SEARCH ANALYZER …` on the field.",
)
.field(InputValue::new("query", TypeRef::named_nn(TypeRef::STRING))),
));
types.push(Type::InputObject(
InputObject::new(CALL_INPUT)
.description(
"Generic function-call predicate. Translates to SurrealQL `fn(field, ...args) op \
value`. Function permissions are enforced at execution time.",
)
.field(InputValue::new("fn", TypeRef::named_nn(TypeRef::STRING)))
.field(InputValue::new("args", TypeRef::named_list("JSON")))
.field(InputValue::new("op", TypeRef::named_nn(NUM_OP_ENUM)))
.field(InputValue::new("value", TypeRef::named_nn("JSON"))),
));
types.push(Type::InputObject(
InputObject::new(ID_RANGE_INPUT)
.description(
"Record ID range predicate. Omitted bounds are unbounded. `inclusive` selects \
`..=` (inclusive end) vs the default `..` (exclusive end). At least one of \
`from` or `to` must be supplied.",
)
.field(InputValue::new("from", TypeRef::named(TypeRef::ID)))
.field(InputValue::new("to", TypeRef::named(TypeRef::ID)))
.field(InputValue::new("inclusive", TypeRef::named(TypeRef::BOOLEAN))),
));
types.push(Type::InputObject(filter_id()));
types.push(Type::InputObject(
InputObject::new(COUNT_FILTER_INPUT)
.description(
"Numeric comparison applied to the count of a relation (graph) traversal in a \
WHERE clause. Multiple operators in one object combine with implicit AND.",
)
.field(InputValue::new("eq", TypeRef::named(TypeRef::INT)))
.field(InputValue::new("ne", TypeRef::named(TypeRef::INT)))
.field(InputValue::new("gt", TypeRef::named(TypeRef::INT)))
.field(InputValue::new("gte", TypeRef::named(TypeRef::INT)))
.field(InputValue::new("lt", TypeRef::named(TypeRef::INT)))
.field(InputValue::new("lte", TypeRef::named(TypeRef::INT))),
));
types.push(Type::Object(
Object::new(PAGE_INFO_TYPE)
.description(
"Cursor pagination metadata. `hasNextPage` / `hasPreviousPage` are computed \
from the over-fetch on the requested direction; the opposite-direction flag \
runs a small probe query when actually selected by the client.",
)
.field(Field::new("hasNextPage", TypeRef::named_nn(TypeRef::BOOLEAN), |ctx| {
FieldFuture::new(async move {
let p = ctx.parent_value.try_downcast_ref::<PageInfoValue>()?;
Ok(Some(FieldValue::value(GraphqlValue::Boolean(p.has_next_page))))
})
}))
.field(Field::new("hasPreviousPage", TypeRef::named_nn(TypeRef::BOOLEAN), |ctx| {
FieldFuture::new(async move {
let p = ctx.parent_value.try_downcast_ref::<PageInfoValue>()?;
Ok(Some(FieldValue::value(GraphqlValue::Boolean(p.has_previous_page))))
})
}))
.field(Field::new("startCursor", TypeRef::named(TypeRef::STRING), |ctx| {
FieldFuture::new(async move {
let p = ctx.parent_value.try_downcast_ref::<PageInfoValue>()?;
Ok(Some(match p.start_cursor.as_deref() {
Some(s) => FieldValue::value(GraphqlValue::String(s.to_owned())),
None => FieldValue::value(GraphqlValue::Null),
}))
})
}))
.field(Field::new("endCursor", TypeRef::named(TypeRef::STRING), |ctx| {
FieldFuture::new(async move {
let p = ctx.parent_value.try_downcast_ref::<PageInfoValue>()?;
Ok(Some(match p.end_cursor.as_deref() {
Some(s) => FieldValue::value(GraphqlValue::String(s.to_owned())),
None => FieldValue::value(GraphqlValue::Null),
}))
})
})),
));
}
#[derive(Clone, Debug, Default)]
struct PageInfoValue {
has_next_page: bool,
has_previous_page: bool,
start_cursor: Option<String>,
end_cursor: Option<String>,
}
fn strip_option(kind: &Kind) -> Kind {
match kind {
Kind::Either(ks) => {
let non_none: Vec<&Kind> =
ks.iter().filter(|k| !matches!(k, Kind::None | Kind::Null)).collect();
if non_none.len() == 1 {
non_none[0].clone()
} else {
kind.clone()
}
}
_ => kind.clone(),
}
}
fn numeric_array_inner(kind: &Kind) -> Option<Kind> {
match kind {
Kind::Array(inner, _) => match inner.as_ref() {
Kind::Float | Kind::Int | Kind::Number | Kind::Decimal => Some(*inner.clone()),
Kind::Either(ks) => {
let non_none: Vec<&Kind> =
ks.iter().filter(|k| !matches!(k, Kind::None | Kind::Null)).collect();
if non_none.len() == 1
&& matches!(non_none[0], Kind::Float | Kind::Int | Kind::Number | Kind::Decimal)
{
Some(non_none[0].clone())
} else {
None
}
}
_ => None,
},
Kind::Either(ks) => {
let non_none: Vec<&Kind> =
ks.iter().filter(|k| !matches!(k, Kind::None | Kind::Null)).collect();
if non_none.len() == 1 {
numeric_array_inner(non_none[0])
} else {
None
}
}
_ => None,
}
}
fn distance_variant(name: &str) -> Option<crate::catalog::Distance> {
use crate::catalog::Distance;
match name {
"COSINE" => Some(Distance::Cosine),
"EUCLIDEAN" => Some(Distance::Euclidean),
"MANHATTAN" => Some(Distance::Manhattan),
"HAMMING" => Some(Distance::Hamming),
"JACCARD" => Some(Distance::Jaccard),
"CHEBYSHEV" => Some(Distance::Chebyshev),
"PEARSON" => Some(Distance::Pearson),
_ => None,
}
}
fn distance_function_name(name: &str) -> Option<&'static str> {
match name {
"COSINE" => Some("vector::similarity::cosine"),
"JACCARD" => Some("vector::similarity::jaccard"),
"PEARSON" => Some("vector::similarity::pearson"),
"EUCLIDEAN" => Some("vector::distance::euclidean"),
"MANHATTAN" => Some("vector::distance::manhattan"),
"HAMMING" => Some("vector::distance::hamming"),
"CHEBYSHEV" => Some("vector::distance::chebyshev"),
_ => None,
}
}
fn num_op_to_binop(name: &str) -> Option<BinaryOperator> {
match name {
"eq" => Some(BinaryOperator::Equal),
"ne" => Some(BinaryOperator::NotEqual),
"gt" => Some(BinaryOperator::MoreThan),
"gte" => Some(BinaryOperator::MoreThanEqual),
"lt" => Some(BinaryOperator::LessThan),
"lte" => Some(BinaryOperator::LessThanEqual),
_ => None,
}
}
fn take_enum<'a>(
obj: &'a IndexMap<Name, GraphqlValue>,
key: &str,
) -> Result<&'a str, GraphqlError> {
let Some(v) = obj.get(key) else {
return Err(resolver_error(format!("missing `{key}` in filter operator input")));
};
match v {
GraphqlValue::Enum(n) => Ok(n.as_str()),
GraphqlValue::String(s) => Ok(s.as_str()),
_ => Err(resolver_error(format!("`{key}` must be an enum or string"))),
}
}
fn translate_nearest(field: Expr, val: &GraphqlValue) -> Result<Expr, GraphqlError> {
use crate::expr::operator::NearestNeighbor;
let obj = val.as_object().ok_or(resolver_error("`nearest` value must be an object"))?;
let k = obj
.get("k")
.and_then(|v| v.as_i64())
.ok_or(resolver_error("`nearest.k` must be an integer"))?;
let k: u32 =
u32::try_from(k.max(0)).map_err(|_| resolver_error("`nearest.k` does not fit in a u32"))?;
let dist_name = take_enum(obj, "distance")?;
let dist = distance_variant(dist_name)
.ok_or_else(|| resolver_error(format!("Unknown distance metric: {dist_name}")))?;
let to_val = obj.get("to").ok_or(resolver_error("`nearest.to` is required"))?;
let to = graphql_to_sql_kind(to_val, Kind::Array(Box::new(Kind::Float), None))?;
Ok(Expr::Binary {
left: Box::new(field),
op: BinaryOperator::NearestNeighbor(Box::new(NearestNeighbor::K(k, dist))),
right: Box::new(to.into_literal()),
})
}
fn translate_similarity(field: Expr, val: &GraphqlValue) -> Result<Expr, GraphqlError> {
let obj = val.as_object().ok_or(resolver_error("`similarity` value must be an object"))?;
let dist_name = take_enum(obj, "distance")?;
let fn_name = distance_function_name(dist_name)
.ok_or_else(|| resolver_error(format!("Unknown distance metric: {dist_name}")))?;
let op_name = take_enum(obj, "op")?;
let op = num_op_to_binop(op_name)
.ok_or_else(|| resolver_error(format!("Unknown comparison op: {op_name}")))?;
let value = obj.get("value").ok_or(resolver_error("`similarity.value` is required"))?;
let value = graphql_to_sql_kind(value, Kind::Float)?;
let to_val = obj.get("to").ok_or(resolver_error("`similarity.to` is required"))?;
let to = graphql_to_sql_kind(to_val, Kind::Array(Box::new(Kind::Float), None))?;
let call = Expr::FunctionCall(Box::new(FunctionCall {
receiver: Function::Normal(fn_name.to_string()),
arguments: vec![field, to.into_literal()],
}));
Ok(Expr::Binary {
left: Box::new(call),
op,
right: Box::new(value.into_literal()),
})
}
fn translate_matches(field: Expr, val: &GraphqlValue) -> Result<Expr, GraphqlError> {
use crate::expr::operator::{BooleanOperator, MatchesOperator};
let obj = val.as_object().ok_or(resolver_error("`matches` value must be an object"))?;
let q = obj
.get("query")
.and_then(|v| match v {
GraphqlValue::String(s) => Some(s.as_str()),
_ => None,
})
.ok_or(resolver_error("`matches.query` must be a string"))?;
let query = graphql_to_sql_kind(&GraphqlValue::String(q.to_string()), Kind::String)?;
Ok(Expr::Binary {
left: Box::new(field),
op: BinaryOperator::Matches(MatchesOperator {
rf: None,
operator: BooleanOperator::And,
}),
right: Box::new(query.into_literal()),
})
}
fn translate_call(field: Expr, val: &GraphqlValue) -> Result<Expr, GraphqlError> {
let obj = val.as_object().ok_or(resolver_error("`call` value must be an object"))?;
let fn_name = obj
.get("fn")
.and_then(|v| match v {
GraphqlValue::String(s) => Some(s.as_str()),
_ => None,
})
.ok_or(resolver_error("`call.fn` must be a string"))?;
let mut args: Vec<Expr> = vec![field];
if let Some(args_val) = obj.get("args")
&& !matches!(args_val, GraphqlValue::Null)
{
let list = args_val.as_list().ok_or(resolver_error("`call.args` must be a list"))?;
for a in list {
let v = graphql_to_sql_kind(a, Kind::Any)?;
args.push(v.into_literal());
}
}
let op_name = take_enum(obj, "op")?;
let op = num_op_to_binop(op_name)
.ok_or_else(|| resolver_error(format!("Unknown comparison op: {op_name}")))?;
let value = obj.get("value").ok_or(resolver_error("`call.value` is required"))?;
let value = graphql_to_sql_kind(value, Kind::Any)?;
let receiver = if let Some(custom) = fn_name.strip_prefix("fn::") {
Function::Custom(custom.to_string())
} else {
Function::Normal(fn_name.to_string())
};
let call = Expr::FunctionCall(Box::new(FunctionCall {
receiver,
arguments: args,
}));
Ok(Expr::Binary {
left: Box::new(call),
op,
right: Box::new(value.into_literal()),
})
}
#[derive(Clone)]
struct AggregateRow(SurObject);
fn is_numeric_kind(kind: &Kind) -> bool {
match kind {
Kind::Float | Kind::Int | Kind::Number | Kind::Decimal => true,
Kind::Either(ks) => {
let non_none: Vec<&Kind> =
ks.iter().filter(|k| !matches!(k, Kind::None | Kind::Null)).collect();
non_none.len() == 1
&& matches!(non_none[0], Kind::Float | Kind::Int | Kind::Number | Kind::Decimal)
}
_ => false,
}
}
fn numeric_kind(kind: &Kind) -> Kind {
match kind {
Kind::Float | Kind::Int | Kind::Number | Kind::Decimal => kind.clone(),
Kind::Either(ks) => {
let non_none: Vec<Kind> =
ks.iter().filter(|k| !matches!(k, Kind::None | Kind::Null)).cloned().collect();
if non_none.len() == 1 {
non_none.into_iter().next().expect("len == 1")
} else {
Kind::Number
}
}
_ => Kind::Number,
}
}
fn aggregate_object_type_name(tb: &str) -> String {
format!("{tb}_aggregate_row")
}
fn aggregate_field_name(tb: &str) -> String {
format!("{tb}_aggregate")
}
fn aggregate_groupable_enum_name(tb: &str) -> String {
format!("_groupable_{tb}")
}
fn build_aggregate_type(
tb_name: &str,
fds: &[FieldDefinition],
types: &mut Vec<Type>,
) -> (Object, Enum) {
let obj_name = aggregate_object_type_name(tb_name);
let mut obj = Object::new(&obj_name)
.description(format!("Aggregation row for `{tb_name}`. `count` is always set; numeric `{{field}}_min/max/sum/avg` are filled per numeric field; group-key columns hold the value of each requested `groupBy` field for the row."));
obj = obj.field(Field::new("count", TypeRef::named_nn(TypeRef::INT), |ctx| {
FieldFuture::new(async move {
let row = ctx.parent_value.try_downcast_ref::<AggregateRow>()?;
let v = row.0.get("count").cloned().unwrap_or(Value::None);
Ok(Some(FieldValue::value(sql_value_to_graphql_value(v)?)))
})
}));
let mut groupable_items: Vec<String> = Vec::new();
for fd in fds {
let fname = idiom_to_graphql_name(&fd.name);
let kind = fd.field_kind.clone().unwrap_or(Kind::Any);
if is_numeric_kind(&kind) {
let inner = numeric_kind(&kind);
let ty_min_max = match kind_to_type(inner.clone(), types, false) {
Ok(t) => unwrap_type(t),
Err(_) => TypeRef::named("number"),
};
let ty_sum = TypeRef::named("number");
let ty_avg = TypeRef::named("number");
for (suffix, ty) in
[("min", ty_min_max.clone()), ("max", ty_min_max), ("sum", ty_sum), ("avg", ty_avg)]
{
let key = format!("{fname}_{suffix}");
let key_for_resolver = key.clone();
obj = obj.field(Field::new(key, ty, move |ctx| {
let k = key_for_resolver.clone();
FieldFuture::new(async move {
let row = ctx.parent_value.try_downcast_ref::<AggregateRow>()?;
let v = row.0.get(k.as_str()).cloned().unwrap_or(Value::None);
if matches!(v, Value::None | Value::Null) {
Ok(None)
} else {
Ok(Some(FieldValue::value(sql_value_to_graphql_value(v)?)))
}
})
}));
}
} else {
let ty = match kind_to_type_with_enum_prefix(
kind.clone(),
types,
false,
Some(&format!("{tb_name}_{fname}")),
) {
Ok(t) => unwrap_type(t),
Err(_) => TypeRef::named("any"),
};
let key_for_resolver = fname.clone();
obj = obj.field(Field::new(fname.clone(), ty, move |ctx| {
let k = key_for_resolver.clone();
FieldFuture::new(async move {
let row = ctx.parent_value.try_downcast_ref::<AggregateRow>()?;
let v = row.0.get(k.as_str()).cloned().unwrap_or(Value::None);
if matches!(v, Value::None | Value::Null) {
Ok(None)
} else {
Ok(Some(FieldValue::value(sql_value_to_graphql_value(v)?)))
}
})
}));
groupable_items.push(fname);
}
}
let enum_name = aggregate_groupable_enum_name(tb_name);
let mut groupable = Enum::new(&enum_name).description(format!(
"Fields of `{tb_name}` that can be used as `groupBy` keys in the aggregate query."
));
if groupable_items.is_empty() {
groupable = groupable.item("_NONE");
} else {
for item in &groupable_items {
groupable = groupable.item(item);
}
}
(obj, groupable)
}
fn make_table_aggregate_field(
tb: &TableDefinition,
fds: Arc<[FieldDefinition]>,
rel_filters: Arc<[RelationFieldInfo]>,
kvs: Arc<Datastore>,
) -> Field {
let tb_name = tb.name.clone();
let tb_name_str = tb_name.as_str().to_string();
let table_filter_name = filter_name_from_table(&tb_name);
let aggregate_row_name = aggregate_object_type_name(&tb_name_str);
let groupable_enum_name = aggregate_groupable_enum_name(&tb_name_str);
let field_name = aggregate_field_name(&super::naming::list_field_name(tb));
let numeric_fields: Vec<String> = fds
.iter()
.filter(|fd| fd.field_kind.as_ref().is_some_and(is_numeric_kind))
.map(|fd| idiom_to_graphql_name(&fd.name))
.collect();
Field::new(field_name, TypeRef::named_nn_list_nn(&aggregate_row_name), move |ctx| {
let tb_name = tb_name.clone();
let fds = Arc::clone(&fds);
let rel_filters = Arc::clone(&rel_filters);
let kvs = Arc::clone(&kvs);
let numeric_fields = numeric_fields.clone();
FieldFuture::new(async move {
let sess = ctx.data::<Arc<Session>>()?;
let args = ctx.args.as_index_map();
let tb_name_str_ref = tb_name.as_str();
let cond = parse_filter_arg(args, &fds, tb_name_str_ref, &rel_filters)?;
let mut group_keys: Vec<String> = Vec::new();
if let Some(gb) = args.get("groupBy")
&& !matches!(gb, GraphqlValue::Null)
{
let list = gb
.as_list()
.ok_or(resolver_error("`groupBy` must be a list"))?;
for v in list {
let name = match v {
GraphqlValue::Enum(n) => n.as_str(),
GraphqlValue::String(s) => s.as_str(),
_ => {
return Err(resolver_error(
"`groupBy` items must be enum values",
)
.into());
}
};
if name != "_NONE" {
group_keys.push(name.to_string());
}
}
}
let mut select_fields: Vec<SelectField> = Vec::new();
select_fields.push(SelectField::Single(Selector {
expr: Expr::FunctionCall(Box::new(FunctionCall {
receiver: Function::Normal("count".to_string()),
arguments: vec![],
})),
alias: Some(Idiom::field("count".to_string())),
}));
for fname in &numeric_fields {
for (suffix, fn_name) in [
("min", "math::min"),
("max", "math::max"),
("sum", "math::sum"),
("avg", "math::mean"),
] {
select_fields.push(SelectField::Single(Selector {
expr: Expr::FunctionCall(Box::new(FunctionCall {
receiver: Function::Normal(fn_name.to_string()),
arguments: vec![Expr::Idiom(Idiom::field(fname.clone()))],
})),
alias: Some(Idiom::field(format!("{fname}_{suffix}"))),
}));
}
}
for gk in &group_keys {
select_fields.push(SelectField::Single(Selector {
expr: Expr::Idiom(Idiom::field(gk.clone())),
alias: None,
}));
}
let group = Some(Groups(
group_keys.iter().map(|gk| Group(Idiom::field(gk.clone()))).collect(),
));
let stmt = SelectStatement {
what: vec![Expr::Table(tb_name)],
fields: Fields::Select(select_fields),
cond,
group,
order: None,
limit: None,
start: None,
version: version_to_expr(&None),
timeout: Expr::Literal(Literal::None),
omit: vec![],
only: false,
with: None,
split: None,
fetch: None,
explain: None,
tempfiles: false,
};
let res = execute_select(&kvs, sess, stmt).await?;
let arr = match res {
Value::Array(a) => a,
v => SurArray::from(vec![v]),
};
let items: Vec<FieldValue> = arr
.iter()
.map(|v| {
let obj = match v {
Value::Object(o) => o.clone(),
_ => SurObject::default(),
};
FieldValue::owned_any(AggregateRow(obj))
})
.collect();
Ok(Some(FieldValue::list(items)))
})
})
.description(format!(
"Aggregation query over `{tb_name_str}`. Returns `count` plus per-numeric-field `min/max/sum/avg`. Group rows by one or more non-numeric fields via the `groupBy` argument."
))
.argument(InputValue::new("filter", TypeRef::named(&table_filter_name)))
.argument(InputValue::new("groupBy", TypeRef::named_list(&groupable_enum_name)))
}
#[derive(Clone, Debug)]
struct EdgeValue {
cursor: String,
node: CachedRecord,
}
fn connection_type_name(tb: &TableDefinition) -> String {
format!("{}Connection", super::naming::to_pascal_case(super::naming::table_base_name(tb)))
}
fn edge_type_name(tb: &TableDefinition) -> String {
format!("{}Edge", super::naming::to_pascal_case(super::naming::table_base_name(tb)))
}
fn build_connection_types(tb: &TableDefinition, types: &mut Vec<Type>) {
let tb_name_str = tb.name.as_str().to_string();
let connection_name = connection_type_name(tb);
let edge_name = edge_type_name(tb);
let node_type = tb_name_str.clone();
let edge_obj = Object::new(&edge_name)
.description(format!("A single edge in a `{tb_name_str}` cursor-paginated connection."))
.field(Field::new("cursor", TypeRef::named_nn(TypeRef::STRING), |ctx| {
FieldFuture::new(async move {
let e = ctx.parent_value.try_downcast_ref::<EdgeValue>()?;
Ok(Some(FieldValue::value(GraphqlValue::String(e.cursor.clone()))))
})
}))
.field(Field::new("node", TypeRef::named_nn(node_type), move |ctx| {
FieldFuture::new(async move {
let e = ctx.parent_value.try_downcast_ref::<EdgeValue>()?;
Ok(Some(FieldValue::owned_any(e.node.clone())))
})
}));
types.push(Type::Object(edge_obj));
let conn_obj = Object::new(&connection_name)
.description(format!(
"Cursor-paginated `{tb_name_str}` records. Forward: pass `after` and read \
`pageInfo.endCursor`. Backward: pass `before` and read `pageInfo.startCursor`."
))
.field(Field::new("edges", TypeRef::named_nn_list_nn(&edge_name), |ctx| {
FieldFuture::new(async move {
let c = ctx.parent_value.try_downcast_ref::<ConnectionValue>()?;
let items: Vec<FieldValue> =
c.edges.iter().cloned().map(FieldValue::owned_any).collect();
Ok(Some(FieldValue::list(items)))
})
}))
.field(Field::new("pageInfo", TypeRef::named_nn(PAGE_INFO_TYPE), |ctx| {
FieldFuture::new(async move {
let c = ctx.parent_value.try_downcast_ref::<ConnectionValue>()?;
Ok(Some(FieldValue::owned_any(c.page_info.clone())))
})
}))
.field(Field::new("totalCount", TypeRef::named(TypeRef::INT), |ctx| {
FieldFuture::new(async move {
let c = ctx.parent_value.try_downcast_ref::<ConnectionValue>()?;
let sess = ctx.data::<Arc<Session>>()?;
let count = run_connection_total_count(&c.total_count_query, sess).await?;
Ok(Some(FieldValue::value(GraphqlValue::Number(count.into()))))
})
}));
types.push(Type::Object(conn_obj));
}
#[derive(Clone)]
struct ConnectionValue {
edges: Vec<EdgeValue>,
page_info: PageInfoValue,
total_count_query: TotalCountQuery,
}
#[derive(Clone)]
struct TotalCountQuery {
kvs: Arc<Datastore>,
tb_name: TableName,
cond: Option<Cond>,
version: Option<Datetime>,
}
fn encode_cursor(rid: &RecordId) -> String {
use base64::Engine;
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(rid.to_sql())
}
fn connection_selects_opposite_page_flag(
ctx: &async_graphql::dynamic::ResolverContext<'_>,
backward: bool,
) -> bool {
let target = if backward {
"hasNextPage"
} else {
"hasPreviousPage"
};
for top in ctx.field().selection_set() {
if top.name() == "pageInfo" {
for inner in top.selection_set() {
if inner.name() == target {
return true;
}
}
}
}
false
}
fn decode_cursor(s: &str) -> Option<RecordId> {
use base64::Engine;
let bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(s).ok()?;
let id_str = String::from_utf8(bytes).ok()?;
crate::syn::record_id(&id_str).ok().map(Into::into)
}
async fn run_connection_total_count(
q: &TotalCountQuery,
sess: &Arc<Session>,
) -> Result<i64, async_graphql::Error> {
let stmt = SelectStatement {
what: vec![Expr::Table(q.tb_name.clone())],
fields: Fields::Select(vec![SelectField::Single(Selector {
expr: Expr::FunctionCall(Box::new(FunctionCall {
receiver: Function::Normal("count".to_string()),
arguments: vec![],
})),
alias: Some(Idiom::field("count".to_string())),
})]),
cond: q.cond.clone(),
group: Some(Groups(Vec::new())),
version: version_to_expr(&q.version),
timeout: Expr::Literal(Literal::None),
omit: vec![],
only: false,
with: None,
split: None,
order: None,
limit: None,
start: None,
fetch: None,
explain: None,
tempfiles: false,
};
let res = execute_select(&q.kvs, sess, stmt).await?;
let arr = match res {
Value::Array(a) => a,
_ => return Ok(0),
};
let Some(first) = arr.0.into_iter().next() else {
return Ok(0);
};
let Value::Object(obj) = first else {
return Ok(0);
};
match obj.get("count") {
Some(Value::Number(n)) => Ok(n.to_int()),
_ => Ok(0),
}
}
fn make_table_connection_field(
tb: &TableDefinition,
fds: Arc<[FieldDefinition]>,
rel_filters: Arc<[RelationFieldInfo]>,
kvs: Arc<Datastore>,
) -> Field {
let tb_name = tb.name.clone();
let tb_name_str = tb_name.as_str().to_string();
let connection_name = connection_type_name(tb);
let table_filter_name = filter_name_from_table(&tb_name);
let field_name = format!("{}Connection", super::naming::list_field_name(tb));
Field::new(field_name, TypeRef::named_nn(connection_name), move |ctx| {
let tb_name = tb_name.clone();
let fds = Arc::clone(&fds);
let rel_filters = Arc::clone(&rel_filters);
let kvs = Arc::clone(&kvs);
FieldFuture::new(async move {
let sess = ctx.data::<Arc<Session>>()?;
let args = ctx.args.as_index_map();
let version = parse_version_arg(args)?;
let mut cond = parse_filter_arg(args, &fds, tb_name.as_str(), &rel_filters)?;
let first = args.get("first").and_then(|v| v.as_i64());
let last = args.get("last").and_then(|v| v.as_i64());
let after = match args.get("after") {
Some(GraphqlValue::String(s)) => Some(s.as_str().to_string()),
_ => None,
};
let before = match args.get("before") {
Some(GraphqlValue::String(s)) => Some(s.as_str().to_string()),
_ => None,
};
if first.is_some() && last.is_some() {
return Err(resolver_error("Pass either `first` or `last`, not both").into());
}
if after.is_some() && before.is_some() {
return Err(resolver_error("Pass either `after` or `before`, not both").into());
}
let backward = last.is_some() || before.is_some();
let page_size = if backward {
last.unwrap_or(20).clamp(1, 1000)
} else {
first.unwrap_or(20).clamp(1, 1000)
};
let cursor_str = if backward {
before.as_deref()
} else {
after.as_deref()
};
let cond_without_cursor = cond.clone();
let decoded_cursor: Option<RecordId> = match cursor_str {
Some(cs) => {
let rid = decode_cursor(cs).ok_or_else(|| resolver_error("invalid cursor"))?;
if rid.table.as_str() != tb_name.as_str() {
return Err(resolver_error(format!(
"cursor table mismatch: cursor decodes to table `{}`, expected `{}`",
rid.table.as_str(),
tb_name.as_str()
))
.into());
}
let op = if backward {
BinaryOperator::LessThan
} else {
BinaryOperator::MoreThan
};
let extra = Expr::Binary {
left: Box::new(Expr::Idiom(Idiom::field("id".to_string()))),
op,
right: Box::new(Value::RecordId(rid.clone()).into_literal()),
};
cond = Some(match cond {
Some(Cond(prev)) => Cond(Expr::Binary {
left: Box::new(prev),
op: BinaryOperator::And,
right: Box::new(extra),
}),
None => Cond(extra),
});
Some(rid)
}
None => None,
};
let total_count_query = TotalCountQuery {
kvs: Arc::clone(&kvs),
tb_name: tb_name.clone(),
cond: cond.clone(),
version,
};
let order_for_query = if backward {
Some(Ordering::Order(OrderList(vec![expr::Order {
value: Idiom::field("id".to_string()),
..expr::Order::default()
}])))
} else {
Some(Ordering::Order(OrderList(vec![order_asc("id".to_string())])))
};
let limit = Some(Limit(Expr::Literal(Literal::Integer(page_size.saturating_add(1)))));
let stmt = select_all_from_table(
Expr::Table(tb_name.clone()),
cond,
order_for_query,
limit,
None,
&version,
);
let res = execute_select(&kvs, sess, stmt).await?;
let arr = match res {
Value::Array(a) => a.0,
v => {
error!("connection query returned non-array: {v:?}");
return Err(internal_error("connection query result not an array").into());
}
};
let returned = arr.len() as i64;
let has_more = returned > page_size;
let mut edges: Vec<EdgeValue> = Vec::with_capacity(page_size as usize);
for v in arr.into_iter().take(page_size as usize) {
let Value::Object(obj) = v else {
continue;
};
let rid = match obj.get("id") {
Some(Value::RecordId(rid)) => rid.clone(),
_ => continue,
};
let cursor = encode_cursor(&rid);
edges.push(EdgeValue {
cursor,
node: CachedRecord {
rid,
version,
data: obj,
},
});
}
if backward {
edges.reverse();
}
let start_cursor = edges.first().map(|e| e.cursor.clone());
let end_cursor = edges.last().map(|e| e.cursor.clone());
let needs_opposite_probe =
decoded_cursor.is_some() && connection_selects_opposite_page_flag(&ctx, backward);
let opposite_has_records = if needs_opposite_probe {
let rid = decoded_cursor.expect("guarded by needs_opposite_probe");
let probe_op = if backward {
BinaryOperator::MoreThanEqual
} else {
BinaryOperator::LessThanEqual
};
let probe_extra = Expr::Binary {
left: Box::new(Expr::Idiom(Idiom::field("id".to_string()))),
op: probe_op,
right: Box::new(Value::RecordId(rid).into_literal()),
};
let probe_cond = Some(match cond_without_cursor {
Some(Cond(prev)) => Cond(Expr::Binary {
left: Box::new(prev),
op: BinaryOperator::And,
right: Box::new(probe_extra),
}),
None => Cond(probe_extra),
});
let probe_limit = Some(Limit(Expr::Literal(Literal::Integer(1))));
let probe_stmt = select_all_from_table(
Expr::Table(tb_name.clone()),
probe_cond,
None,
probe_limit,
None,
&version,
);
let probe_res = execute_select(&kvs, sess, probe_stmt).await?;
matches!(&probe_res, Value::Array(a) if !a.0.is_empty())
} else {
false
};
let (has_next_page, has_previous_page) = if backward {
(opposite_has_records, has_more)
} else {
(has_more, opposite_has_records)
};
let conn = ConnectionValue {
edges,
page_info: PageInfoValue {
has_next_page,
has_previous_page,
start_cursor,
end_cursor,
},
total_count_query,
};
Ok(Some(FieldValue::owned_any(conn)))
})
})
.description(format!(
"Cursor-paginated `{tb_name_str}` list. Forward: pass `first` and `after` (from \
`pageInfo.endCursor`). Backward: pass `last` and `before` (from `pageInfo.startCursor`). \
Iterates by record id; use the non-connection list query for a custom sort key. \
`totalCount` runs a separate `count()` query on demand."
))
.argument(InputValue::new("first", TypeRef::named(TypeRef::INT)))
.argument(InputValue::new("after", TypeRef::named(TypeRef::STRING)))
.argument(InputValue::new("last", TypeRef::named(TypeRef::INT)))
.argument(InputValue::new("before", TypeRef::named(TypeRef::STRING)))
.argument(InputValue::new("filter", TypeRef::named(&table_filter_name)))
.argument(InputValue::new("where", TypeRef::named(&table_filter_name)))
.argument(InputValue::new("version", TypeRef::named(TypeRef::STRING)))
}