use std::collections::HashSet;
use parse_rust_core::{classify, ErrorCode, ParseError, ParseValue};
use parse_rust_storage::{Comparison, Constraint};
use serde_json::Value as Json;
#[derive(Debug, Clone, Default)]
pub struct ParsedWhere {
pub clauses: Vec<ParsedClause>,
}
#[derive(Debug, Clone)]
pub enum ParsedClause {
Field(Constraint),
RelatedTo {
class_name: String,
object_id: String,
key: String,
},
Or(Vec<ParsedWhere>),
And(Vec<ParsedWhere>),
Nor(Vec<ParsedWhere>),
}
impl ParsedWhere {
pub fn is_empty(&self) -> bool {
self.clauses.is_empty()
}
pub fn field_keys(&self) -> Vec<String> {
let mut out = Vec::new();
self.collect_field_keys(&mut out);
out
}
fn collect_field_keys(&self, out: &mut Vec<String>) {
for clause in &self.clauses {
match clause {
ParsedClause::Field(c) => out.push(c.field.clone()),
ParsedClause::RelatedTo { .. } => {}
ParsedClause::Or(branches)
| ParsedClause::And(branches)
| ParsedClause::Nor(branches) => {
for branch in branches {
branch.collect_field_keys(out);
}
}
}
}
}
pub fn pinned_object_id(&self) -> Option<&str> {
self.clauses.iter().find_map(|c| match c {
ParsedClause::Field(Constraint {
field,
comparison: Comparison::Equal(ParseValue::String(id)),
}) if field == "objectId" => Some(id.as_str()),
_ => None,
})
}
pub fn push(&mut self, clause: ParsedClause) {
self.clauses.push(clause);
}
}
pub fn parse_where(where_json: &Json) -> Result<ParsedWhere, ParseError> {
parse_where_at(where_json, true)
}
fn parse_where_at(where_json: &Json, top_level: bool) -> Result<ParsedWhere, ParseError> {
let Json::Object(map) = where_json else {
return Err(ParseError::invalid_query(
"where must be an object".to_string(),
));
};
let mut out = ParsedWhere::default();
for (field, value) in map {
if field == "ACL" {
return Err(ParseError::invalid_query(
"Cannot query on ACL.".to_string(),
));
}
if let Some(clause) = parse_query_level_key(field, value)? {
out.push(clause);
continue;
}
match value {
Json::Object(inner) if is_operator_document(inner) => {
for constraint in parse_operators(field, inner)? {
out.push(ParsedClause::Field(constraint));
}
}
Json::Object(inner) if top_level && is_mixed_document(inner) => {
let mut rewritten = serde_json::Map::new();
let mut equal_to = serde_json::Map::new();
for (key, v) in inner {
if key.starts_with('$') {
rewritten.insert(key.clone(), v.clone());
} else {
equal_to.insert(key.clone(), v.clone());
}
}
rewritten.insert("$eq".to_string(), Json::Object(equal_to));
for constraint in parse_operators(field, &rewritten)? {
out.push(ParsedClause::Field(constraint));
}
}
literal => out.push(ParsedClause::Field(Constraint {
field: field.clone(),
comparison: Comparison::Equal(parse_rust_core::classify_raw(literal.clone())?),
})),
}
}
Ok(out)
}
fn parse_query_level_key(field: &str, value: &Json) -> Result<Option<ParsedClause>, ParseError> {
if !field.starts_with('$') {
return Ok(None);
}
let clause = match field {
"$or" | "$and" | "$nor" => {
let branches = match value {
Json::Array(items) => items
.iter()
.map(|item| parse_where_at(item, false))
.collect::<Result<Vec<_>, _>>()?,
_ => {
return Err(ParseError::invalid_query(if field == "$nor" {
"Bad $nor format - use an array of at least 1 value.".to_string()
} else {
format!("Bad {field} format - use an array value.")
}))
}
};
if field == "$nor" && branches.is_empty() {
return Err(ParseError::invalid_query(
"Bad $nor format - use an array of at least 1 value.".to_string(),
));
}
match field {
"$or" => ParsedClause::Or(branches),
"$and" => ParsedClause::And(branches),
_ => ParsedClause::Nor(branches),
}
}
"$relatedTo" => parse_related_to(value)?,
other => {
return Err(ParseError::invalid_query(format!(
"unsupported query operator: {other}"
)))
}
};
Ok(Some(clause))
}
fn parse_related_to(value: &Json) -> Result<ParsedClause, ParseError> {
let bad = || ParseError::invalid_query("improper usage of $relatedTo".to_string());
let Json::Object(map) = value else {
return Err(bad());
};
let key = match map.get("key") {
Some(Json::String(k)) => k.clone(),
_ => return Err(bad()),
};
let object = map.get("object").ok_or_else(bad)?;
match classify(object.clone())? {
ParseValue::Pointer {
class_name,
object_id,
} => Ok(ParsedClause::RelatedTo {
class_name,
object_id,
key,
}),
_ => Err(bad()),
}
}
fn parse_operators(
field: &str,
inner: &serde_json::Map<String, Json>,
) -> Result<Vec<Constraint>, ParseError> {
let mut out = Vec::new();
let regex = inner.get("$regex");
let options = inner.get("$options");
if regex.is_none() && options.is_some() {
return Err(ParseError::invalid_query(
"$options is only valid with $regex".to_string(),
));
}
if let Some(regex) = regex {
let Json::String(pattern) = regex else {
return Err(ParseError::invalid_query(
"$regex value must be a string".to_string(),
));
};
let options = match options {
None => None,
Some(Json::String(o)) => {
if !o.chars().all(|c| matches!(c, 'i' | 'm' | 'x' | 's' | 'u')) || o.is_empty() {
return Err(ParseError::invalid_query(format!(
"Bad $options value for query: {o}"
)));
}
Some(o.clone())
}
Some(_) => {
return Err(ParseError::invalid_query(
"$options value must be a string".to_string(),
))
}
};
out.push(Constraint {
field: field.to_string(),
comparison: Comparison::Regex {
pattern: pattern.clone(),
options,
},
});
}
for (op, operand) in inner {
if op == "$regex" || op == "$options" {
continue;
}
let operand = parse_rust_core::classify_raw(operand.clone())?;
out.push(Constraint {
field: field.to_string(),
comparison: Comparison::from_operator(op, operand)?,
});
}
Ok(out)
}
fn is_operator_document(map: &serde_json::Map<String, Json>) -> bool {
!map.is_empty() && map.keys().all(|k| k.starts_with('$'))
}
fn is_mixed_document(map: &serde_json::Map<String, Json>) -> bool {
map.keys().any(|k| k.starts_with('$')) && map.keys().any(|k| !k.starts_with('$'))
}
pub const CLIENT_QUERYABLE_INTERNAL_FIELDS: [&str; 2] = ["_rperm", "_wperm"];
pub const MASTER_QUERYABLE_INTERNAL_FIELDS: [&str; 10] = [
"_email_verify_token",
"_perishable_token",
"_perishable_token_expires_at",
"_email_verify_token_expires_at",
"_failed_login_count",
"_account_lockout_expires_at",
"_password_changed_at",
"_password_history",
"_tombstone",
"_session_token",
];
pub fn validate_query_keys(where_: &ParsedWhere, is_master: bool) -> Result<(), ParseError> {
for key in where_.field_keys() {
if matches_query_key_regex(&key)
|| CLIENT_QUERYABLE_INTERNAL_FIELDS.contains(&key.as_str())
|| (is_master && MASTER_QUERYABLE_INTERNAL_FIELDS.contains(&key.as_str()))
{
continue;
}
return Err(ParseError::invalid_key_name(format!(
"Invalid key name: {key}"
)));
}
Ok(())
}
fn matches_query_key_regex(key: &str) -> bool {
let mut chars = key.chars();
match chars.next() {
Some(c) if c.is_ascii_alphabetic() => {}
_ => return false,
}
chars.all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '.')
}
pub const MAX_INCLUDE_DEPTH: usize = 20;
pub const MAX_INCLUDE_PATHS: usize = 500;
pub fn parse_include(include: &str) -> Result<Vec<Vec<String>>, ParseError> {
let mut paths: Vec<Vec<String>> = Vec::new();
let mut seen: HashSet<&str> = HashSet::new();
for raw in include.split(',').map(str::trim).filter(|s| !s.is_empty()) {
if raw == "*" {
return Err(ParseError::new(
ErrorCode::CommandUnavailable,
"include=* is not supported yet.",
));
}
let parts: Vec<&str> = raw.split('.').collect();
if parts.len() > MAX_INCLUDE_DEPTH {
return Err(ParseError::invalid_query(format!(
"include path is too deep: at most {MAX_INCLUDE_DEPTH} components."
)));
}
for depth in 1..=parts.len() {
let end = parts[..depth].iter().map(|p| p.len()).sum::<usize>() + depth - 1;
if !seen.insert(&raw[..end]) {
continue;
}
if paths.len() == MAX_INCLUDE_PATHS {
return Err(ParseError::invalid_query(format!(
"too many include paths: at most {MAX_INCLUDE_PATHS}."
)));
}
paths.push(parts[..depth].iter().map(|s| s.to_string()).collect());
}
}
paths.sort_by_key(Vec::len);
Ok(paths)
}
#[cfg(test)]
mod tests {
use super::*;
fn j(s: &str) -> Json {
serde_json::from_str(s).expect("test literal")
}
fn fields(w: &ParsedWhere) -> Vec<&Constraint> {
w.clauses
.iter()
.filter_map(|c| match c {
ParsedClause::Field(f) => Some(f),
_ => None,
})
.collect()
}
#[test]
fn a_bare_value_is_equality() {
let w = parse_where(&j(r#"{"title":"hello"}"#)).expect("parse");
let c = fields(&w);
assert_eq!(c.len(), 1);
assert_eq!(c[0].field, "title");
assert!(matches!(
&c[0].comparison,
Comparison::Equal(ParseValue::String(s)) if s == "hello"
));
}
#[test]
fn several_operators_on_one_field_become_several_constraints() {
let w = parse_where(&j(r#"{"views":{"$gt":1,"$lt":9}}"#)).expect("parse");
assert_eq!(fields(&w).len(), 2);
assert!(fields(&w).iter().all(|x| x.field == "views"));
}
#[test]
fn a_tagged_value_is_a_literal_and_reaches_the_backend_undecoded() {
let w = parse_where(&j(
r#"{"author":{"__type":"Pointer","className":"_User","objectId":"u1","extra":7}}"#,
))
.expect("parse");
let Comparison::Equal(ParseValue::Object(map)) = &fields(&w)[0].comparison else {
panic!("expected a raw object operand, got {:?}", fields(&w)[0]);
};
assert_eq!(map.len(), 4, "{map:?}");
assert!(matches!(map.get("objectId"), Some(ParseValue::String(s)) if s == "u1"));
assert!(matches!(map.get("extra"), Some(ParseValue::Number(n)) if *n == 7.0));
}
#[test]
fn an_unsupported_operator_is_refused() {
for op in [
"$inQuery",
"$notInQuery",
"$select",
"$dontSelect",
"$text",
"$nearSphere",
"$containedBy",
"$geoWithin",
] {
let src = format!(r#"{{"title":{{"{op}":1}}}}"#);
let e = parse_where(&j(&src)).unwrap_err();
assert_eq!(e.code, ErrorCode::InvalidQuery, "{op}");
assert!(e.message.contains(op), "{op}: {}", e.message);
}
}
#[test]
fn an_unknown_query_level_operator_is_refused() {
let e = parse_where(&j(r#"{"$nope":[]}"#)).unwrap_err();
assert_eq!(e.code, ErrorCode::InvalidQuery);
assert!(e.message.contains("$nope"));
}
#[test]
fn logical_operators_parse_recursively() {
let w = parse_where(&j(
r#"{"$or":[{"a":1},{"$and":[{"b":2},{"c":3}]}],"$nor":[{"d":4}]}"#,
))
.expect("parse");
assert_eq!(w.clauses.len(), 2);
match &w.clauses[0] {
ParsedClause::Or(branches) => {
assert_eq!(branches.len(), 2);
assert!(matches!(branches[1].clauses[0], ParsedClause::And(_)));
}
other => panic!("expected Or, got {other:?}"),
}
assert!(matches!(w.clauses[1], ParsedClause::Nor(_)));
}
#[test]
fn a_non_array_logical_operator_is_invalid() {
for src in [r#"{"$or":{"a":1}}"#, r#"{"$and":3}"#, r#"{"$nor":[]}"#] {
let e = parse_where(&j(src)).unwrap_err();
assert_eq!(e.code, ErrorCode::InvalidQuery, "{src}");
}
}
#[test]
fn regex_folds_its_options_in() {
let w = parse_where(&j(r#"{"title":{"$regex":"^a","$options":"im"}}"#)).expect("parse");
let c = fields(&w);
assert_eq!(c.len(), 1, "two keys make one comparison");
assert!(matches!(
&c[0].comparison,
Comparison::Regex { pattern, options } if pattern == "^a" && options.as_deref() == Some("im")
));
}
#[test]
fn regex_rejects_a_non_string_pattern_and_bad_options() {
assert!(parse_where(&j(r#"{"title":{"$regex":3}}"#)).is_err());
let e = parse_where(&j(r#"{"title":{"$regex":"a","$options":"z"}}"#)).unwrap_err();
assert_eq!(e.code, ErrorCode::InvalidQuery);
assert!(e.message.contains("Bad $options value for query: z"));
assert!(parse_where(&j(r#"{"title":{"$options":"i"}}"#)).is_err());
}
#[test]
fn all_parses() {
let w = parse_where(&j(r#"{"tags":{"$all":["a","b"]}}"#)).expect("parse");
assert!(matches!(&fields(&w)[0].comparison, Comparison::All(v) if v.len() == 2));
}
#[test]
fn related_to_parses_into_its_own_clause() {
let w = parse_where(&j(
r#"{"$relatedTo":{"object":{"__type":"Pointer","className":"_Role","objectId":"r1"},"key":"users"}}"#,
))
.expect("parse");
match &w.clauses[0] {
ParsedClause::RelatedTo {
class_name,
object_id,
key,
} => {
assert_eq!(class_name, "_Role");
assert_eq!(object_id, "r1");
assert_eq!(key, "users");
}
other => panic!("expected RelatedTo, got {other:?}"),
}
assert!(parse_where(&j(r#"{"$relatedTo":{"key":"users"}}"#)).is_err());
assert!(parse_where(&j(r#"{"$relatedTo":{"object":3,"key":"u"}}"#)).is_err());
}
#[test]
fn querying_on_acl_is_refused() {
let e = parse_where(&j(r#"{"ACL":{"*":{"read":true}}}"#)).unwrap_err();
assert_eq!(e.code, ErrorCode::InvalidQuery);
assert_eq!(e.message, "Cannot query on ACL.");
}
#[test]
fn in_and_exists_parse() {
let w =
parse_where(&j(r#"{"tag":{"$in":["a","b"]},"x":{"$exists":true}}"#)).expect("parse");
assert_eq!(fields(&w).len(), 2);
}
#[test]
fn an_empty_object_is_an_empty_query_not_an_operator_document() {
assert!(parse_where(&j("{}")).expect("parse").is_empty());
let w = parse_where(&j(r#"{"meta":{}}"#)).expect("parse");
assert!(matches!(
&fields(&w)[0].comparison,
Comparison::Equal(ParseValue::Object(_))
));
}
#[test]
fn where_must_be_an_object() {
assert!(parse_where(&j("[]")).is_err());
assert!(parse_where(&j("3")).is_err());
}
#[test]
fn field_keys_reach_into_logical_clauses() {
let w = parse_where(&j(r#"{"a":1,"$or":[{"b":2},{"$and":[{"c":3}]}]}"#)).expect("parse");
let mut keys = w.field_keys();
keys.sort();
assert_eq!(keys, vec!["a", "b", "c"]);
}
#[test]
fn pinned_object_id_only_reads_a_top_level_equality() {
let w = parse_where(&j(r#"{"objectId":"abc"}"#)).expect("parse");
assert_eq!(w.pinned_object_id(), Some("abc"));
let w = parse_where(&j(r#"{"$or":[{"objectId":"abc"}]}"#)).expect("parse");
assert_eq!(w.pinned_object_id(), None);
}
#[test]
fn include_paths_materialize_prefixes_and_sort_by_depth() {
assert_eq!(
parse_include("a.b.c,d").expect("parse"),
vec![
vec!["a".to_string()],
vec!["d".to_string()],
vec!["a".to_string(), "b".to_string()],
vec!["a".to_string(), "b".to_string(), "c".to_string()],
]
);
assert!(parse_include("").expect("parse").is_empty());
}
#[test]
fn a_hostile_include_is_refused_rather_than_expanded() {
let deep = vec!["a"; MAX_INCLUDE_DEPTH + 1].join(".");
let err = parse_include(&deep).expect_err("over the depth limit");
assert_eq!(err.code, ErrorCode::InvalidQuery);
let huge = vec!["a"; 24_000].join(".");
assert_eq!(
parse_include(&huge).expect_err("over the depth limit").code,
ErrorCode::InvalidQuery
);
let wide = (0..MAX_INCLUDE_PATHS + 1)
.map(|i| format!("f{i}"))
.collect::<Vec<_>>()
.join(",");
assert_eq!(
parse_include(&wide).expect_err("over the path limit").code,
ErrorCode::InvalidQuery
);
assert_eq!(
parse_include("author.company.owner").expect("parse").len(),
3
);
let at_depth = vec!["a"; MAX_INCLUDE_DEPTH].join(".");
assert_eq!(
parse_include(&at_depth)
.expect("exactly at the limit")
.len(),
MAX_INCLUDE_DEPTH
);
}
#[test]
fn a_client_cannot_name_an_internal_column_in_a_query() {
let w = parse_where(&j(r#"{"_hashed_password":{"$regex":"^a"}}"#)).expect("parse");
for is_master in [false, true] {
let e = validate_query_keys(&w, is_master).unwrap_err();
assert_eq!(e.code, ErrorCode::InvalidKeyName);
assert_eq!(e.message, "Invalid key name: _hashed_password");
}
let w = parse_where(&j(r#"{"_rperm":"u1"}"#)).expect("parse");
assert!(validate_query_keys(&w, false).is_ok());
let w = parse_where(&j(r#"{"_session_token":"r:t"}"#)).expect("parse");
assert!(validate_query_keys(&w, false).is_err());
assert!(validate_query_keys(&w, true).is_ok());
let w = parse_where(&j(r#"{"$or":[{"_session_token":"r:t"}]}"#)).expect("parse");
assert!(validate_query_keys(&w, false).is_err());
}
#[test]
fn ordinary_and_dotted_field_names_pass() {
for src in [r#"{"title":"x"}"#, r#"{"meta.a_b":1}"#] {
let w = parse_where(&j(src)).expect("parse");
assert!(validate_query_keys(&w, false).is_ok(), "{src}");
}
let w = parse_where(&j(r#"{"1bad":1}"#)).expect("parse");
assert!(validate_query_keys(&w, false).is_err());
}
#[test]
fn include_all_is_an_explicit_error() {
let e = parse_include("*").unwrap_err();
assert_eq!(e.code, ErrorCode::CommandUnavailable);
assert!(e.message.contains("include=*"));
}
}
#[cfg(test)]
mod mixed_constraint_tests {
use super::*;
fn parse(json: &str) -> Result<ParsedWhere, ParseError> {
parse_where(&serde_json::from_str(json).expect("test literal"))
}
#[test]
fn a_mixed_constraint_becomes_an_eq_plus_the_operators() {
let parsed = parse(r#"{"meta": {"foo": 1, "$gt": 0}}"#).expect("parses");
let mut equals = 0;
let mut greater = 0;
for clause in &parsed.clauses {
let ParsedClause::Field(c) = clause else {
panic!("expected field clauses")
};
assert_eq!(c.field, "meta");
match &c.comparison {
Comparison::EqualOperator(ParseValue::Object(map)) => {
assert_eq!(map.len(), 1, "only the direct keys: {map:?}");
assert!(map.contains_key("foo"));
equals += 1;
}
Comparison::GreaterThan(_) => greater += 1,
other => panic!("unexpected comparison: {other:?}"),
}
}
assert_eq!((equals, greater), (1, 1));
}
#[test]
fn a_mixed_constraint_is_not_one_literal_equality() {
let parsed = parse(r#"{"meta": {"foo": 1, "$gt": 0}}"#).expect("parses");
assert_eq!(
parsed.clauses.len(),
2,
"one clause means it was read as a literal"
);
}
#[test]
fn unmixed_documents_are_unchanged() {
let ops = parse(r#"{"n": {"$gt": 0, "$lt": 9}}"#).expect("parses");
assert_eq!(ops.clauses.len(), 2);
let literal = parse(r#"{"p": {"__type": "Pointer", "className": "C", "objectId": "x"}}"#)
.expect("parses");
assert_eq!(literal.clauses.len(), 1);
}
#[test]
fn the_rewrite_does_not_reach_inside_or() {
let parsed = parse(r#"{"$or": [{"meta": {"foo": 1, "$gt": 0}}]}"#).expect("parses");
let [ParsedClause::Or(branches)] = parsed.clauses.as_slice() else {
panic!("expected one $or clause")
};
assert_eq!(branches.len(), 1);
assert_eq!(
branches[0].clauses.len(),
1,
"inside $or the mixed document stays one literal equality"
);
}
#[test]
fn an_explicit_eq_operator_is_accepted() {
let parsed = parse(r#"{"n": {"$eq": 5}}"#).expect("parses");
assert_eq!(parsed.clauses.len(), 1);
let ParsedClause::Field(c) = &parsed.clauses[0] else {
panic!("expected a field clause")
};
assert!(matches!(c.comparison, Comparison::EqualOperator(_)));
}
}