use std::collections::{BTreeMap, HashMap};
use crate::cypher::record::NamedRecord;
use crate::Value;
use anyhow::{anyhow, bail, Result};
fn gherkin_unescape(s: &str) -> String {
let mut out = String::with_capacity(s.len());
let mut chars = s.chars().peekable();
while let Some(c) = chars.next() {
if c == '\\' {
match chars.peek() {
Some('\\') => {
out.push('\\');
chars.next();
}
Some('|') => {
out.push('|');
chars.next();
}
Some('n') => {
out.push('\n');
chars.next();
}
_ => out.push(c),
}
} else {
out.push(c);
}
}
out
}
pub fn parse_expected(cell: &str) -> Result<Value> {
let unescaped = gherkin_unescape(cell.trim());
let mut p = Parser::new(&unescaped);
let v = p.parse_value()?;
p.skip_ws();
if !p.at_end() {
bail!("trailing characters in expected cell: {cell:?}");
}
Ok(v)
}
pub fn compare_result(
actual: &[NamedRecord],
columns: &[String],
expected_rows: &[Vec<Value>],
ordered: bool,
) -> Result<()> {
let actual_rows = project_rows(actual, columns);
if actual_rows.len() != expected_rows.len() {
bail!(
"row count mismatch: expected {}, got {}\nexpected={:?}\nactual={:?}",
expected_rows.len(),
actual_rows.len(),
expected_rows,
actual_rows
);
}
if ordered {
for (i, (got, want)) in actual_rows.iter().zip(expected_rows.iter()).enumerate() {
if !rows_equal(got, want) {
bail!(
"row {i} mismatch\n expected: {:?}\n actual: {:?}",
want,
got
);
}
}
} else {
let mut used = vec![false; actual_rows.len()];
for want in expected_rows {
let mut found = false;
for (i, got) in actual_rows.iter().enumerate() {
if !used[i] && rows_equal(got, want) {
used[i] = true;
found = true;
break;
}
}
if !found {
bail!(
"expected row not found in actual:\n expected: {:?}\n actual: {:?}",
want,
actual_rows
);
}
}
}
Ok(())
}
fn project_rows(records: &[NamedRecord], columns: &[String]) -> Vec<Vec<Value>> {
records
.iter()
.map(|rec| {
columns
.iter()
.map(|col| rec.get(col).cloned().unwrap_or(Value::Null))
.collect()
})
.collect()
}
fn rows_equal(got: &[Value], want: &[Value]) -> bool {
if got.len() != want.len() {
return false;
}
got.iter().zip(want.iter()).all(|(g, w)| value_equal(g, w))
}
fn value_equal(a: &Value, b: &Value) -> bool {
match (a, b) {
(Value::Node(na), Value::Node(nb)) => {
na.labels == nb.labels && properties_equal(&na.properties, &nb.properties)
}
(Value::Edge(ea), Value::Edge(eb)) => {
ea.label == eb.label && properties_equal(&ea.properties, &eb.properties)
}
(Value::Path(pa), Value::Path(pb)) => {
pa.nodes.len() == pb.nodes.len()
&& pa.edges.len() == pb.edges.len()
&& pa
.nodes
.iter()
.zip(pb.nodes.iter())
.all(|(a, b)| value_equal(&Value::Node(a.clone()), &Value::Node(b.clone())))
&& pa
.edges
.iter()
.zip(pb.edges.iter())
.all(|(a, b)| value_equal(&Value::Edge(a.clone()), &Value::Edge(b.clone())))
}
(Value::List(la), Value::List(lb)) => {
la.len() == lb.len() && la.iter().zip(lb.iter()).all(|(x, y)| value_equal(x, y))
}
(Value::Map(ma), Value::Map(mb)) => {
ma.len() == mb.len()
&& ma
.iter()
.all(|(k, v)| mb.get(k).is_some_and(|w| value_equal(v, w)))
}
(Value::Date(d), Value::String(s)) | (Value::String(s), Value::Date(d)) => {
d.to_string() == *s
}
(Value::LocalTime(t), Value::String(s)) | (Value::String(s), Value::LocalTime(t)) => {
t.to_string() == *s
}
(Value::Time(t), Value::String(s)) | (Value::String(s), Value::Time(t)) => {
t.to_string() == *s
}
(Value::LocalDateTime(dt), Value::String(s))
| (Value::String(s), Value::LocalDateTime(dt)) => dt.to_string() == *s,
(Value::DateTime(dt), Value::String(s)) | (Value::String(s), Value::DateTime(dt)) => {
dt.to_string() == *s
}
(Value::Duration(d), Value::String(s)) | (Value::String(s), Value::Duration(d)) => {
d.to_string() == *s
}
(Value::F64(fa), Value::F64(fb)) if fa.is_nan() && fb.is_nan() => true,
_ => a == b,
}
}
pub fn compare_result_ignore_list_order(
actual: &[NamedRecord],
columns: &[String],
expected_rows: &[Vec<Value>],
ordered: bool,
) -> Result<()> {
let actual_rows = project_rows(actual, columns);
if actual_rows.len() != expected_rows.len() {
bail!(
"row count mismatch: expected {}, got {}\nexpected={:?}\nactual={:?}",
expected_rows.len(),
actual_rows.len(),
expected_rows,
actual_rows
);
}
if ordered {
for (i, (got, want)) in actual_rows.iter().zip(expected_rows.iter()).enumerate() {
if !rows_equal_ignore_list_order(got, want) {
bail!(
"row {i} mismatch\n expected: {:?}\n actual: {:?}",
want,
got
);
}
}
} else {
let mut used = vec![false; actual_rows.len()];
for want in expected_rows {
let mut found = false;
for (i, got) in actual_rows.iter().enumerate() {
if !used[i] && rows_equal_ignore_list_order(got, want) {
used[i] = true;
found = true;
break;
}
}
if !found {
bail!(
"expected row not found in actual:\n expected: {:?}\n actual: {:?}",
want,
actual_rows
);
}
}
}
Ok(())
}
fn rows_equal_ignore_list_order(got: &[Value], want: &[Value]) -> bool {
if got.len() != want.len() {
return false;
}
got.iter()
.zip(want.iter())
.all(|(g, w)| value_equal_ignore_list_order(g, w))
}
fn value_equal_ignore_list_order(a: &Value, b: &Value) -> bool {
match (a, b) {
(Value::List(la), Value::List(lb)) => {
if la.len() != lb.len() {
return false;
}
let mut used = vec![false; lb.len()];
for x in la {
let mut matched = false;
for (j, y) in lb.iter().enumerate() {
if !used[j] && value_equal_ignore_list_order(x, y) {
used[j] = true;
matched = true;
break;
}
}
if !matched {
return false;
}
}
true
}
_ => value_equal(a, b),
}
}
fn properties_equal(a: &HashMap<String, Value>, b: &HashMap<String, Value>) -> bool {
a.len() == b.len()
&& a.iter()
.all(|(k, v)| b.get(k).is_some_and(|w| value_equal(v, w)))
}
struct Parser<'a> {
src: &'a str,
pos: usize,
}
impl<'a> Parser<'a> {
fn new(src: &'a str) -> Self {
Self { src, pos: 0 }
}
fn at_end(&self) -> bool {
self.pos >= self.src.len()
}
fn rest(&self) -> &'a str {
&self.src[self.pos..]
}
fn peek(&self) -> Option<char> {
self.rest().chars().next()
}
fn advance(&mut self, n: usize) {
self.pos += n;
}
fn skip_ws(&mut self) {
while let Some(c) = self.peek() {
if c.is_whitespace() {
self.advance(c.len_utf8());
} else {
break;
}
}
}
fn consume(&mut self, tag: &str) -> bool {
if self.rest().starts_with(tag) {
self.advance(tag.len());
true
} else {
false
}
}
fn expect(&mut self, tag: &str) -> Result<()> {
if self.consume(tag) {
Ok(())
} else {
bail!(
"expected {tag:?} at position {} (rest={:?})",
self.pos,
self.rest()
)
}
}
fn parse_value(&mut self) -> Result<Value> {
self.skip_ws();
match self.peek() {
None => Err(anyhow!("unexpected end of expected cell")),
Some('n') if self.rest().starts_with("null") => {
self.advance(4);
Ok(Value::Null)
}
Some('N') if self.rest().starts_with("NaN") => {
self.advance(3);
Ok(Value::F64(f64::NAN))
}
Some('t') if self.rest().starts_with("true") => {
self.advance(4);
Ok(Value::Bool(true))
}
Some('f') if self.rest().starts_with("false") => {
self.advance(5);
Ok(Value::Bool(false))
}
Some('\'') => self.parse_string(),
Some('[') => {
let after_bracket = self.rest()[1..].trim_start();
if after_bracket.starts_with(':') {
self.parse_edge()
} else {
self.parse_list()
}
}
Some('{') => self.parse_map(),
Some('(') => self.parse_node(),
Some('<') => self.parse_path(),
Some(c) if c == '-' || c.is_ascii_digit() => self.parse_number(),
Some(c) => Err(anyhow!(
"unexpected character {c:?} at position {}",
self.pos
)),
}
}
fn parse_string(&mut self) -> Result<Value> {
self.expect("'")?;
let mut out = String::new();
while let Some(c) = self.peek() {
if c == '\'' {
self.advance(1);
return Ok(Value::String(out));
}
if c == '\\' {
self.advance(1);
match self.peek() {
Some('n') => {
out.push('\n');
self.advance(1);
}
Some('t') => {
out.push('\t');
self.advance(1);
}
Some('\'') => {
out.push('\'');
self.advance(1);
}
Some('\\') => {
out.push('\\');
self.advance(1);
}
Some(other) => {
out.push(other);
self.advance(other.len_utf8());
}
None => bail!("unterminated escape in string literal"),
}
continue;
}
out.push(c);
self.advance(c.len_utf8());
}
bail!("unterminated string literal")
}
fn parse_number(&mut self) -> Result<Value> {
let start = self.pos;
if self.consume("-") {
}
while let Some(c) = self.peek() {
if c.is_ascii_digit() {
self.advance(1);
} else {
break;
}
}
let mut is_float = false;
if self.peek() == Some('.') {
is_float = true;
self.advance(1);
while let Some(c) = self.peek() {
if c.is_ascii_digit() {
self.advance(1);
} else {
break;
}
}
}
if matches!(self.peek(), Some('e') | Some('E')) {
is_float = true;
self.advance(1);
if matches!(self.peek(), Some('+') | Some('-')) {
self.advance(1);
}
while let Some(c) = self.peek() {
if c.is_ascii_digit() {
self.advance(1);
} else {
break;
}
}
}
let text = &self.src[start..self.pos];
if is_float {
Ok(Value::F64(
text.parse().map_err(|e| anyhow!("float parse: {e}"))?,
))
} else {
Ok(Value::I64(
text.parse().map_err(|e| anyhow!("int parse: {e}"))?,
))
}
}
fn parse_list(&mut self) -> Result<Value> {
self.expect("[")?;
let mut items = Vec::new();
self.skip_ws();
if self.consume("]") {
return Ok(Value::List(items));
}
loop {
items.push(self.parse_value()?);
self.skip_ws();
if self.consume("]") {
return Ok(Value::List(items));
}
self.expect(",")?;
}
}
fn parse_map(&mut self) -> Result<Value> {
self.expect("{")?;
let mut entries: BTreeMap<String, Value> = BTreeMap::new();
self.skip_ws();
if self.consume("}") {
return Ok(Value::Map(entries));
}
loop {
self.skip_ws();
let key = self.parse_ident()?;
self.skip_ws();
self.expect(":")?;
let val = self.parse_value()?;
entries.insert(key, val);
self.skip_ws();
if self.consume("}") {
return Ok(Value::Map(entries));
}
self.expect(",")?;
}
}
fn parse_ident(&mut self) -> Result<String> {
let start = self.pos;
while let Some(c) = self.peek() {
if c.is_ascii_alphanumeric() || c == '_' {
self.advance(c.len_utf8());
} else {
break;
}
}
if start == self.pos {
bail!("expected identifier at position {}", self.pos);
}
Ok(self.src[start..self.pos].to_string())
}
fn parse_node(&mut self) -> Result<Value> {
use crate::{Node, NodeId};
self.expect("(")?;
self.skip_ws();
while let Some(c) = self.peek() {
if c == ':' || c == ')' || c == '{' {
break;
}
self.advance(c.len_utf8());
}
let mut labels = Vec::new();
while self.consume(":") {
labels.push(self.parse_ident()?);
self.skip_ws();
}
labels.sort();
let mut properties = HashMap::new();
if self.peek() == Some('{') {
if let Value::Map(m) = self.parse_map()? {
for (k, v) in m {
properties.insert(k, v);
}
}
self.skip_ws();
}
self.expect(")")?;
Ok(Value::Node(Node {
id: NodeId(0), labels,
properties,
}))
}
fn parse_edge(&mut self) -> Result<Value> {
use crate::{Edge, NodeId};
self.expect("[")?;
self.skip_ws();
while let Some(c) = self.peek() {
if c == ':' || c == ']' || c == '{' {
break;
}
self.advance(c.len_utf8());
}
let mut label = String::new();
if self.consume(":") {
label = self.parse_ident()?;
self.skip_ws();
}
let mut properties = HashMap::new();
if self.peek() == Some('{') {
if let Value::Map(m) = self.parse_map()? {
for (k, v) in m {
properties.insert(k, v);
}
}
self.skip_ws();
}
self.expect("]")?;
Ok(Value::Edge(Edge {
src: NodeId(0), dst: NodeId(0),
label,
properties,
}))
}
fn parse_path(&mut self) -> Result<Value> {
use crate::PathValue;
self.expect("<")?;
let mut nodes = Vec::new();
let mut edges = Vec::new();
self.skip_ws();
if self.peek() == Some('(') {
if let Value::Node(n) = self.parse_node()? {
nodes.push(n);
}
}
loop {
self.skip_ws();
if self.peek() == Some('>') && !self.rest().starts_with(">-") {
break;
}
if self.at_end() {
break;
}
if self.consume("-") {
if let Value::Edge(e) = self.parse_edge()? {
edges.push(e);
}
self.expect("->")?;
} else if self.consume("<-") {
if let Value::Edge(e) = self.parse_edge()? {
edges.push(e);
}
self.expect("-")?;
} else {
break;
}
self.skip_ws();
if self.peek() == Some('(') {
if let Value::Node(n) = self.parse_node()? {
nodes.push(n);
}
}
}
self.expect(">")?;
Ok(Value::Path(PathValue { nodes, edges }))
}
}
#[cfg(test)]
mod tests {
#[allow(unused_imports)]
use super::parse_expected;
use crate::Value;
#[test]
#[allow(clippy::approx_constant)]
fn scalars() {
assert_eq!(parse_expected("null").unwrap(), Value::Null);
assert_eq!(parse_expected("true").unwrap(), Value::Bool(true));
assert_eq!(parse_expected("42").unwrap(), Value::I64(42));
assert_eq!(parse_expected("-7").unwrap(), Value::I64(-7));
assert_eq!(parse_expected("3.14").unwrap(), Value::F64(3.14));
assert_eq!(
parse_expected("'hello'").unwrap(),
Value::String("hello".into())
);
}
#[test]
fn lists_and_maps() {
assert_eq!(
parse_expected("[1, 2, 3]").unwrap(),
Value::List(vec![Value::I64(1), Value::I64(2), Value::I64(3)])
);
assert_eq!(parse_expected("[]").unwrap(), Value::List(vec![]));
let m = parse_expected("{a: 1, b: 'x'}").unwrap();
if let Value::Map(map) = m {
assert_eq!(map.len(), 2);
assert_eq!(map.get("a"), Some(&Value::I64(1)));
} else {
panic!("expected map");
}
}
#[test]
fn node_pattern() {
let n = parse_expected("(:Person {name: 'Alice', age: 30})").unwrap();
if let Value::Node(node) = n {
assert_eq!(node.labels, vec!["Person".to_string()]);
assert_eq!(
node.properties.get("name"),
Some(&Value::String("Alice".into()))
);
} else {
panic!("expected node");
}
}
#[test]
fn multi_label_node() {
let n = parse_expected("(:A:B)").unwrap();
if let Value::Node(node) = n {
assert_eq!(node.labels, vec!["A".to_string(), "B".to_string()]);
} else {
panic!("expected node");
}
}
#[test]
fn path_literal() {
let p = parse_expected("<(:A)-[:R]->(:B)>").unwrap();
if let Value::Path(path) = p {
assert_eq!(path.nodes.len(), 2);
assert_eq!(path.edges.len(), 1);
assert_eq!(path.nodes[0].labels, vec!["A".to_string()]);
assert_eq!(path.nodes[1].labels, vec!["B".to_string()]);
assert_eq!(path.edges[0].label, "R");
} else {
panic!("expected path");
}
}
}