diff --git a/builtin/builtin.go b/builtin/builtin.go index 93622fb..d462b68 100644 --- a/builtin/builtin.go +++ b/builtin/builtin.go @@ -36,6 +36,7 @@ var ( symbols.Contains: {ast.ArgModeInput, ast.ArgModeInput}, symbols.Filter: {ast.ArgModeInput}, symbols.Lt: {ast.ArgModeInput, ast.ArgModeInput}, + symbols.Ne: {ast.ArgModeInput, ast.ArgModeInput}, symbols.Le: {ast.ArgModeInput, ast.ArgModeInput}, symbols.Gt: {ast.ArgModeInput, ast.ArgModeInput}, symbols.Ge: {ast.ArgModeInput, ast.ArgModeInput}, @@ -271,6 +272,11 @@ func Decide(atom ast.Atom, subst *unionfind.UnionFind) (bool, []*unionfind.Union } return false, nil, nil + case symbols.Ne.Symbol: + if len(atom.Args) != 2 { + return false, nil, fmt.Errorf("wrong number of arguments for built-in predicate ':ne': %v", atom.Args) + } + return !atom.Args[0].Equals(atom.Args[1]), []*unionfind.UnionFind{subst}, nil case symbols.Lt.Symbol: if len(atom.Args) != 2 { return false, nil, fmt.Errorf("wrong number of arguments for built-in predicate '<': %v", atom.Args) diff --git a/symbols/symbols.go b/symbols/symbols.go index 6151350..393f0f9 100644 --- a/symbols/symbols.go +++ b/symbols/symbols.go @@ -42,6 +42,8 @@ var ( // Lt is the less-than relation on numbers. Lt = ast.PredicateSym{":lt", 2} + // Ne is inequality over any two evaluated constants. + Ne = ast.PredicateSym{":ne", 2} // Le is the less-than-or-equal relation on numbers. Le = ast.PredicateSym{":le", 2} @@ -369,6 +371,7 @@ var ( Filter: NewRelType(BoolType()), // TODO: support float64 Lt: NewRelType(ast.NumberBound, ast.NumberBound), + Ne: NewRelType(ast.AnyBound, ast.AnyBound), Le: NewRelType(ast.NumberBound, ast.NumberBound), Gt: NewRelType(ast.NumberBound, ast.NumberBound), Ge: NewRelType(ast.NumberBound, ast.NumberBound),