use core::fmt;
use crate::json::JsonValue;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum CompareOp {
Eq,
Ne,
Co,
Sw,
Ew,
Gt,
Ge,
Lt,
Le,
}
#[derive(Debug, Clone, PartialEq)]
pub enum FilterValue {
Str(String),
Bool(bool),
Num(f64),
Null,
}
#[derive(Debug, Clone, PartialEq)]
pub enum Filter {
Present(String),
Compare(String, CompareOp, FilterValue),
And(Box<Filter>, Box<Filter>),
Or(Box<Filter>, Box<Filter>),
Not(Box<Filter>),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ScimFilterError {
message: &'static str,
}
impl ScimFilterError {
#[must_use]
pub fn message(&self) -> &str {
self.message
}
}
impl fmt::Display for ScimFilterError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "invalid SCIM filter: {}", self.message)
}
}
impl std::error::Error for ScimFilterError {}
fn err(message: &'static str) -> ScimFilterError {
ScimFilterError { message }
}
#[derive(Debug, Clone, PartialEq)]
enum Token {
Ident(String),
Str(String),
LParen,
RParen,
}
fn tokenize(input: &str) -> Result<Vec<Token>, ScimFilterError> {
let mut tokens = Vec::new();
let chars: Vec<char> = input.chars().collect();
let mut i = 0;
while i < chars.len() {
let c = chars[i];
if c.is_whitespace() {
i += 1;
} else if c == '(' {
tokens.push(Token::LParen);
i += 1;
} else if c == ')' {
tokens.push(Token::RParen);
i += 1;
} else if c == '"' {
let mut s = String::new();
i += 1;
loop {
let Some(&ch) = chars.get(i) else {
return Err(err("unterminated string"));
};
match ch {
'"' => {
i += 1;
break;
}
'\\' => {
let Some(&next) = chars.get(i + 1) else {
return Err(err("dangling escape"));
};
s.push(next);
i += 2;
}
_ => {
s.push(ch);
i += 1;
}
}
}
tokens.push(Token::Str(s));
} else {
let start = i;
while i < chars.len() && !chars[i].is_whitespace() && chars[i] != '(' && chars[i] != ')'
{
i += 1;
}
let word: String = chars[start..i].iter().collect();
tokens.push(Token::Ident(word));
}
}
Ok(tokens)
}
const MAX_FILTER_DEPTH: usize = 32;
const MAX_FILTER_NODES: usize = 512;
struct Parser {
tokens: Vec<Token>,
pos: usize,
depth: usize,
nodes: usize,
}
impl Parser {
fn peek(&self) -> Option<&Token> {
self.tokens.get(self.pos)
}
fn next(&mut self) -> Option<Token> {
let t = self.tokens.get(self.pos).cloned();
if t.is_some() {
self.pos += 1;
}
t
}
fn bump(&mut self) -> Result<(), ScimFilterError> {
self.nodes += 1;
if self.nodes > MAX_FILTER_NODES {
return Err(err("filter too large"));
}
Ok(())
}
fn peek_keyword(&self, kw: &str) -> bool {
matches!(self.peek(), Some(Token::Ident(w)) if w.eq_ignore_ascii_case(kw))
}
fn parse_or(&mut self) -> Result<Filter, ScimFilterError> {
let mut left = self.parse_and()?;
while self.peek_keyword("or") {
self.pos += 1;
let right = self.parse_and()?;
self.bump()?;
left = Filter::Or(Box::new(left), Box::new(right));
}
Ok(left)
}
fn parse_and(&mut self) -> Result<Filter, ScimFilterError> {
let mut left = self.parse_unary()?;
while self.peek_keyword("and") {
self.pos += 1;
let right = self.parse_unary()?;
self.bump()?;
left = Filter::And(Box::new(left), Box::new(right));
}
Ok(left)
}
fn parse_unary(&mut self) -> Result<Filter, ScimFilterError> {
if self.peek_keyword("not") {
self.pos += 1;
if !matches!(self.next(), Some(Token::LParen)) {
return Err(err("expected '(' after not"));
}
let inner = self.parse_group()?;
if !matches!(self.next(), Some(Token::RParen)) {
return Err(err("expected ')'"));
}
self.bump()?;
return Ok(Filter::Not(Box::new(inner)));
}
if matches!(self.peek(), Some(Token::LParen)) {
self.pos += 1;
let inner = self.parse_group()?;
if !matches!(self.next(), Some(Token::RParen)) {
return Err(err("expected ')'"));
}
return Ok(inner);
}
self.parse_attr_expr()
}
fn parse_group(&mut self) -> Result<Filter, ScimFilterError> {
self.depth += 1;
if self.depth > MAX_FILTER_DEPTH {
return Err(err("filter nesting too deep"));
}
let inner = self.parse_or()?;
self.depth -= 1;
Ok(inner)
}
fn parse_attr_expr(&mut self) -> Result<Filter, ScimFilterError> {
let Some(Token::Ident(attr)) = self.next() else {
return Err(err("expected attribute path"));
};
let Some(Token::Ident(op)) = self.next() else {
return Err(err("expected operator"));
};
if op.eq_ignore_ascii_case("pr") {
self.bump()?;
return Ok(Filter::Present(attr));
}
let cmp = match op.to_ascii_lowercase().as_str() {
"eq" => CompareOp::Eq,
"ne" => CompareOp::Ne,
"co" => CompareOp::Co,
"sw" => CompareOp::Sw,
"ew" => CompareOp::Ew,
"gt" => CompareOp::Gt,
"ge" => CompareOp::Ge,
"lt" => CompareOp::Lt,
"le" => CompareOp::Le,
_ => return Err(err("unknown operator")),
};
let value = match self.next() {
Some(Token::Str(s)) => FilterValue::Str(s),
Some(Token::Ident(w)) => literal_from_word(&w),
_ => return Err(err("expected comparison value")),
};
self.bump()?;
Ok(Filter::Compare(attr, cmp, value))
}
}
fn literal_from_word(w: &str) -> FilterValue {
match w {
"true" => FilterValue::Bool(true),
"false" => FilterValue::Bool(false),
"null" => FilterValue::Null,
_ => w
.parse::<f64>()
.map_or_else(|_| FilterValue::Str(w.to_string()), FilterValue::Num),
}
}
pub fn parse_filter(input: &str) -> Result<Filter, ScimFilterError> {
let tokens = tokenize(input)?;
if tokens.is_empty() {
return Err(err("empty filter"));
}
let mut parser = Parser {
tokens,
pos: 0,
depth: 0,
nodes: 0,
};
let filter = parser.parse_or()?;
if parser.pos != parser.tokens.len() {
return Err(err("trailing tokens"));
}
Ok(filter)
}
impl Filter {
#[must_use]
pub fn matches(&self, resource: &JsonValue) -> bool {
match self {
Filter::Present(path) => resolve_path(resource, path).iter().any(|v| !v.is_null()),
Filter::Compare(path, op, value) => {
let resolved = resolve_path(resource, path);
if resolved.is_empty() {
return matches!(op, CompareOp::Ne);
}
resolved.iter().any(|v| compare(v, *op, value))
}
Filter::And(a, b) => a.matches(resource) && b.matches(resource),
Filter::Or(a, b) => a.matches(resource) || b.matches(resource),
Filter::Not(inner) => !inner.matches(resource),
}
}
}
fn resolve_path<'a>(resource: &'a JsonValue, path: &str) -> Vec<&'a JsonValue> {
let mut current: Vec<&JsonValue> = vec![resource];
for segment in path.split('.') {
let mut next: Vec<&JsonValue> = Vec::new();
for node in current {
match node {
JsonValue::Array(items) => {
for item in items {
if let Some(obj) = item.as_object() {
if let Some((_, v)) =
obj.iter().find(|(k, _)| k.eq_ignore_ascii_case(segment))
{
next.push(v);
}
}
}
}
other => {
if let Some(obj) = other.as_object() {
if let Some((_, v)) =
obj.iter().find(|(k, _)| k.eq_ignore_ascii_case(segment))
{
next.push(v);
}
}
}
}
}
if next.is_empty() {
return Vec::new();
}
current = next;
}
current
.into_iter()
.flat_map(|v| match v {
JsonValue::Array(items) => items.iter().collect::<Vec<_>>(),
other => vec![other],
})
.collect()
}
fn compare(actual: &JsonValue, op: CompareOp, expected: &FilterValue) -> bool {
if let (Some(a), FilterValue::Str(e)) = (actual.as_str(), expected) {
let (al, el) = (a.to_lowercase(), e.to_lowercase());
return match op {
CompareOp::Eq => al == el,
CompareOp::Ne => al != el,
CompareOp::Co => al.contains(&el),
CompareOp::Sw => al.starts_with(&el),
CompareOp::Ew => al.ends_with(&el),
CompareOp::Gt => al > el,
CompareOp::Ge => al >= el,
CompareOp::Lt => al < el,
CompareOp::Le => al <= el,
};
}
if let (Some(a), FilterValue::Num(e)) = (actual.as_f64(), expected) {
return match op {
CompareOp::Eq => (a - e).abs() < f64::EPSILON,
CompareOp::Ne => (a - e).abs() >= f64::EPSILON,
CompareOp::Gt => a > *e,
CompareOp::Ge => a >= *e,
CompareOp::Lt => a < *e,
CompareOp::Le => a <= *e,
CompareOp::Co | CompareOp::Sw | CompareOp::Ew => false,
};
}
if let (Some(a), FilterValue::Bool(e)) = (actual.as_bool(), expected) {
return match op {
CompareOp::Eq => a == *e,
CompareOp::Ne => a != *e,
_ => false,
};
}
if let FilterValue::Null = expected {
return match op {
CompareOp::Eq => actual.is_null(),
CompareOp::Ne => !actual.is_null(),
_ => false,
};
}
false
}
#[cfg(test)]
mod tests {
use super::*;
fn user() -> JsonValue {
JsonValue::parse(
r#"{"userName":"BJensen","active":true,"name":{"familyName":"Jensen"},"age":30}"#,
)
.unwrap()
}
fn m(filter: &str) -> bool {
parse_filter(filter).unwrap().matches(&user())
}
#[test]
fn eq_is_case_insensitive() {
assert!(m(r#"userName eq "bjensen""#));
assert!(m(r#"userName eq "BJENSEN""#));
assert!(!m(r#"userName eq "other""#));
}
#[test]
fn present_contains_starts_ends() {
assert!(m("userName pr"));
assert!(!m("nickName pr"));
assert!(m(r#"userName co "jen""#));
assert!(m(r#"userName sw "bj""#));
assert!(m(r#"userName ew "sen""#));
}
#[test]
fn dotted_path_and_bool_and_number() {
assert!(m(r#"name.familyName eq "jensen""#));
assert!(m("active eq true"));
assert!(!m("active eq false"));
assert!(m("age gt 20"));
assert!(m("age le 30"));
assert!(!m("age gt 30"));
}
#[test]
fn logical_and_or_not_grouping() {
assert!(m(r#"userName eq "bjensen" and active eq true"#));
assert!(!m(r#"userName eq "nope" and active eq true"#));
assert!(m(r#"userName eq "nope" or active eq true"#));
assert!(m(r#"not (userName eq "nope")"#));
assert!(m(
r#"(userName eq "bjensen" or userName eq "x") and active eq true"#
));
}
#[test]
fn precedence_or_binds_loosest() {
assert!(
parse_filter(r#"userName eq "nope" and active eq true or age eq 30"#)
.unwrap()
.matches(&user())
);
}
#[test]
fn rejects_malformed() {
assert!(parse_filter("").is_err());
assert!(parse_filter("userName").is_err());
assert!(parse_filter("userName eq").is_err());
assert!(parse_filter(r#"(userName eq "x""#).is_err());
assert!(parse_filter(r#"userName zz "x""#).is_err());
}
#[test]
fn rejects_over_long_flat_chain() {
let chain = "userName pr or ".repeat(MAX_FILTER_NODES + 10) + "userName pr";
assert!(parse_filter(&chain).is_err());
let ok = "userName pr or active eq true or age gt 10";
assert!(parse_filter(ok).is_ok());
}
#[test]
fn rejects_over_deep_nesting() {
let deep = format!("{}userName pr", "(".repeat(MAX_FILTER_DEPTH + 5));
assert!(parse_filter(&deep).is_err());
let ok = format!("{}userName pr{}", "(".repeat(4), ")".repeat(4));
assert!(parse_filter(&ok).is_ok());
}
#[test]
fn multi_valued_paths_resolve() {
let user = JsonValue::parse(
r#"{"userName":"bj","emails":[{"value":"a@x.com","type":"work"},
{"value":"b@y.com","type":"home"}]}"#,
)
.unwrap();
assert!(
parse_filter(r#"emails.value eq "b@y.com""#)
.unwrap()
.matches(&user)
);
assert!(
parse_filter(r#"emails.type eq "work""#)
.unwrap()
.matches(&user)
);
assert!(
!parse_filter(r#"emails.value eq "nope@z.com""#)
.unwrap()
.matches(&user)
);
assert!(parse_filter("emails.value pr").unwrap().matches(&user));
}
#[test]
fn ne_against_an_absent_attribute_matches() {
let user = JsonValue::parse(r#"{"userName":"bj"}"#).unwrap();
assert!(parse_filter(r#"nickName ne "bob""#).unwrap().matches(&user));
assert!(!parse_filter(r#"nickName eq "bob""#).unwrap().matches(&user));
assert!(
parse_filter(r#"userName ne "other""#)
.unwrap()
.matches(&user)
);
assert!(!parse_filter(r#"userName ne "bj""#).unwrap().matches(&user));
}
}