diff --git a/pkg/keg/omegadsl/ast.go b/pkg/keg/omegadsl/ast.go new file mode 100644 index 0000000..db5b263 --- /dev/null +++ b/pkg/keg/omegadsl/ast.go @@ -0,0 +1,105 @@ +package omegadsl + +// The AST mirrors the grammar in .tapper/specs/omega-dsl.ebnf. All node types +// are unexported; the only public handle is *Program (see omegadsl.go). + +type letBinding struct { + name string + value expr +} + +// expr is any evaluable Omega DSL expression. +type expr interface{ isExpr() } + +// numberLit is a decimal literal, e.g. 0.25. +type numberLit struct{ val float64 } + +// boolLit is true / false. +type boolLit struct{ val bool } + +// stringLit is a double-quoted literal used for comparisons and enum keys. +type stringLit struct{ val string } + +// letRef is a bare identifier referring to an earlier let binding. +type letRef struct { + name string + line int +} + +// metaRef reads a metadata field: meta., or self.meta. in a +// predicate (self == true addresses the scored node rather than the target). +type metaRef struct { + name string + self bool +} + +// statRef reads a node stat: stat. / self.stat.. +type statRef struct { + name string + self bool +} + +// linkSetRef is a directional relation accessor: out. / in. / bi.. +type linkSetRef struct { + dir Direction + rel string + line int +} + +// unaryExpr is prefix '-' or 'not'. +type unaryExpr struct { + op string + x expr +} + +// binaryExpr covers arithmetic, comparison, and and/or. +type binaryExpr struct { + op string + l, r expr +} + +// matchExpr maps a subject value to a number by pattern. +type matchExpr struct { + subject expr + arms []matchArm +} + +type matchArm struct { + isElse bool + pat string // subject is compared to this text (identifier/string/number/bool) + result expr +} + +// weightedExpr is the normalized weighted-average combinator. +type weightedExpr struct{ terms []weightedTerm } + +type weightedTerm struct { + weight expr + contrib expr +} + +// callExpr is a built-in function invocation. +type callExpr struct { + name string + args []callArg + line int +} + +// callArg is a positional (name == "") or named (score's k=v) argument. +type callArg struct { + name string + value expr +} + +func (numberLit) isExpr() {} +func (boolLit) isExpr() {} +func (stringLit) isExpr() {} +func (letRef) isExpr() {} +func (metaRef) isExpr() {} +func (statRef) isExpr() {} +func (linkSetRef) isExpr() {} +func (unaryExpr) isExpr() {} +func (binaryExpr) isExpr() {} +func (matchExpr) isExpr() {} +func (weightedExpr) isExpr() {} +func (callExpr) isExpr() {} diff --git a/pkg/keg/omegadsl/eval.go b/pkg/keg/omegadsl/eval.go new file mode 100644 index 0000000..f6a2a90 --- /dev/null +++ b/pkg/keg/omegadsl/eval.go @@ -0,0 +1,396 @@ +package omegadsl + +import ( + "errors" + "fmt" + "strconv" +) + +var errEmptyProgram = errors.New("omegadsl: empty program") + +// --- runtime values ------------------------------------------------------- + +type valueKind int + +const ( + kNumber valueKind = iota + kBool + kString + kAbsent // a referenced field/stat that is not present + kSet // a link set; only valid as the first argument of a set function +) + +type value struct { + kind valueKind + num float64 + b bool + s string + set []Node +} + +func numberVal(f float64) value { return value{kind: kNumber, num: f} } +func boolVal(b bool) value { return value{kind: kBool, b: b} } +func stringVal(s string) value { return value{kind: kString, s: s} } +func absentVal() value { return value{kind: kAbsent} } +func setVal(n []Node) value { return value{kind: kSet, set: n} } + +// --- evaluation scope ----------------------------------------------------- + +// scope carries the state threaded through evaluation. root is always the node +// being scored (backing self.* and out./in./bi.); current is the node bare +// meta.*/stat.* read from (root, or a link target inside a predicate). +type scope struct { + root Env + current Node + lets map[string]value + atRoot bool +} + +func evalProgram(p *Program, env Env) (float64, error) { + if env == nil { + return 0, errors.New("omegadsl: nil env") + } + sc := &scope{root: env, current: env, lets: map[string]value{}, atRoot: true} + for _, lb := range p.lets { + v, err := eval(lb.value, sc) + if err != nil { + return 0, err + } + sc.lets[lb.name] = v + } + v, err := eval(p.omega, sc) + if err != nil { + return 0, err + } + n, err := asNumber(v) + if err != nil { + return 0, fmt.Errorf("omega must be a number: %w", err) + } + return clamp01(n), nil +} + +// child returns a scope for evaluating a predicate against a link target. +func (sc *scope) child(target Node) *scope { + return &scope{root: sc.root, current: target, lets: sc.lets, atRoot: false} +} + +// --- expression evaluation ------------------------------------------------ + +// eval walks the parsed AST. This is a pure, sandboxed interpreter over a fixed, +// closed set of arithmetic/logic/lookup operations — it is NOT host-code +// execution: there is no reflection, no I/O, and no access to anything beyond +// the supplied Env's scalar meta/stat/link accessors. It cannot escape that +// surface regardless of program input. +func eval(e expr, sc *scope) (value, error) { + switch n := e.(type) { + case numberLit: + return numberVal(n.val), nil + case boolLit: + return boolVal(n.val), nil + case stringLit: + return stringVal(n.val), nil + case letRef: + v, ok := sc.lets[n.name] + if !ok { + return value{}, fmt.Errorf("undefined binding %q", n.name) + } + return v, nil + case metaRef: + return fieldValue(sc, n.self, false, n.name), nil + case statRef: + return fieldValue(sc, n.self, true, n.name), nil + case linkSetRef: + return evalLinkSet(n, sc) + case unaryExpr: + return evalUnary(n, sc) + case binaryExpr: + return evalBinary(n, sc) + case matchExpr: + return evalMatch(n, sc) + case weightedExpr: + return evalWeighted(n, sc) + case callExpr: + return evalCall(n, sc) + default: + return value{}, fmt.Errorf("omegadsl: unhandled expression %T", e) + } +} + +func fieldValue(sc *scope, self, stat bool, name string) value { + node := sc.current + if self { + node = sc.root + } + var ( + s string + ok bool + ) + if stat { + s, ok = node.Stat(name) + } else { + s, ok = node.Meta(name) + } + if !ok { + return absentVal() + } + return stringVal(s) +} + +func evalLinkSet(n linkSetRef, sc *scope) (value, error) { + if !sc.atRoot { + return value{}, fmt.Errorf("%s.%s: link sets cannot be used inside a predicate (multi-hop traversal is not supported)", n.dir, n.rel) + } + nodes, err := sc.root.LinkSet(n.rel, n.dir) + if err != nil { + return value{}, err + } + return setVal(nodes), nil +} + +func evalUnary(n unaryExpr, sc *scope) (value, error) { + v, err := eval(n.x, sc) + if err != nil { + return value{}, err + } + switch n.op { + case "-": + f, err := asNumber(v) + if err != nil { + return value{}, err + } + return numberVal(-f), nil + case "not": + return boolVal(!truthy(v)), nil + default: + return value{}, fmt.Errorf("omegadsl: unknown unary operator %q", n.op) + } +} + +func evalBinary(n binaryExpr, sc *scope) (value, error) { + switch n.op { + case "and": + l, err := eval(n.l, sc) + if err != nil { + return value{}, err + } + if !truthy(l) { + return boolVal(false), nil + } + r, err := eval(n.r, sc) + if err != nil { + return value{}, err + } + return boolVal(truthy(r)), nil + case "or": + l, err := eval(n.l, sc) + if err != nil { + return value{}, err + } + if truthy(l) { + return boolVal(true), nil + } + r, err := eval(n.r, sc) + if err != nil { + return value{}, err + } + return boolVal(truthy(r)), nil + } + + l, err := eval(n.l, sc) + if err != nil { + return value{}, err + } + r, err := eval(n.r, sc) + if err != nil { + return value{}, err + } + + switch n.op { + case "==": + return boolVal(valuesEqual(l, r)), nil + case "!=": + return boolVal(!valuesEqual(l, r)), nil + case "<", "<=", ">", ">=": + lf, err := asNumber(l) + if err != nil { + return value{}, err + } + rf, err := asNumber(r) + if err != nil { + return value{}, err + } + return boolVal(compareNumbers(n.op, lf, rf)), nil + case "+", "-", "*", "/": + lf, err := asNumber(l) + if err != nil { + return value{}, err + } + rf, err := asNumber(r) + if err != nil { + return value{}, err + } + if n.op == "/" && rf == 0 { + return value{}, errors.New("division by zero") + } + return numberVal(arith(n.op, lf, rf)), nil + default: + return value{}, fmt.Errorf("omegadsl: unknown operator %q", n.op) + } +} + +func evalMatch(n matchExpr, sc *scope) (value, error) { + subj, err := eval(n.subject, sc) + if err != nil { + return value{}, err + } + key, matchable := scalarString(subj) + if matchable { + for _, arm := range n.arms { + if !arm.isElse && arm.pat == key { + return eval(arm.result, sc) + } + } + } + for _, arm := range n.arms { + if arm.isElse { + return eval(arm.result, sc) + } + } + return numberVal(0), nil +} + +func evalWeighted(n weightedExpr, sc *scope) (value, error) { + var acc, total float64 + for _, term := range n.terms { + wv, err := eval(term.weight, sc) + if err != nil { + return value{}, err + } + w, err := asNumber(wv) + if err != nil { + return value{}, fmt.Errorf("weighted term weight: %w", err) + } + cv, err := eval(term.contrib, sc) + if err != nil { + return value{}, err + } + c, err := asNumber(cv) + if err != nil { + return value{}, fmt.Errorf("weighted term contribution: %w", err) + } + acc += w * clamp01(c) + total += w + } + if total <= 0 { + return numberVal(0), nil + } + return numberVal(acc / total), nil +} + +// --- coercions & helpers -------------------------------------------------- + +func asNumber(v value) (float64, error) { + switch v.kind { + case kNumber: + return v.num, nil + case kBool: + if v.b { + return 1, nil + } + return 0, nil + case kString: + f, err := strconv.ParseFloat(v.s, 64) + if err != nil { + return 0, fmt.Errorf("%q is not a number", v.s) + } + return f, nil + case kAbsent: + return 0, nil + default: + return 0, errors.New("a link set is not a number") + } +} + +func truthy(v value) bool { + switch v.kind { + case kBool: + return v.b + case kNumber: + return v.num != 0 + case kString: + return v.s != "" && v.s != "false" && v.s != "0" + case kAbsent: + return false + case kSet: + return len(v.set) > 0 + default: + return false + } +} + +// scalarString renders a scalar value for match/equality comparison. It returns +// ok == false for absent values and link sets, which never match. +func scalarString(v value) (string, bool) { + switch v.kind { + case kString: + return v.s, true + case kNumber: + return strconv.FormatFloat(v.num, 'g', -1, 64), true + case kBool: + if v.b { + return "true", true + } + return "false", true + default: + return "", false + } +} + +func valuesEqual(l, r value) bool { + ls, lok := scalarString(l) + rs, rok := scalarString(r) + if !lok || !rok { + return false // absent never equals anything, including another absent + } + return ls == rs +} + +func compareNumbers(op string, l, r float64) bool { + switch op { + case "<": + return l < r + case "<=": + return l <= r + case ">": + return l > r + case ">=": + return l >= r + default: + return false + } +} + +func arith(op string, l, r float64) float64 { + switch op { + case "+": + return l + r + case "-": + return l - r + case "*": + return l * r + case "/": + return l / r + default: + return 0 + } +} + +func clamp01(f float64) float64 { + if f < 0 { + return 0 + } + if f > 1 { + return 1 + } + return f +} diff --git a/pkg/keg/omegadsl/funcs.go b/pkg/keg/omegadsl/funcs.go new file mode 100644 index 0000000..bf06446 --- /dev/null +++ b/pkg/keg/omegadsl/funcs.go @@ -0,0 +1,273 @@ +package omegadsl + +import "fmt" + +// evalCall dispatches a built-in function. Set functions (count/has/any/all/ +// frac/avg/sum) receive a link set as their first argument and evaluate their +// predicate/expression argument once per link target (in the target's scope); +// the other builtins evaluate their arguments normally. +func evalCall(n callExpr, sc *scope) (value, error) { + switch n.name { + case "score": + return evalScore(n, sc) + case "clamp": + return evalClamp(n, sc) + case "min", "max": + return evalMinMax(n, sc) + case "present", "missing": + return evalPresence(n, sc) + case "count", "has", "any", "all", "frac", "avg", "sum": + return evalSetFunc(n, sc) + default: + return value{}, fmt.Errorf("unknown function %q", n.name) + } +} + +func positionalArgs(n callExpr) ([]expr, error) { + out := make([]expr, 0, len(n.args)) + for _, a := range n.args { + if a.name != "" { + return nil, fmt.Errorf("%s does not take named arguments (%q)", n.name, a.name) + } + out = append(out, a.value) + } + return out, nil +} + +// score(field, k1=v1, ..., else=e) — the enum-scoring builtin. The single +// positional argument is the subject field; named arguments map subject values +// to numeric contributions, with `else` as the fallback. +func evalScore(n callExpr, sc *scope) (value, error) { + var subject expr + pos := 0 + for _, a := range n.args { + if a.name == "" { + pos++ + subject = a.value + } + } + if pos != 1 { + return value{}, fmt.Errorf("score expects exactly one positional field argument, got %d", pos) + } + subj, err := eval(subject, sc) + if err != nil { + return value{}, err + } + key, ok := scalarString(subj) + + var matched, elseExpr expr + for _, a := range n.args { + switch { + case a.name == "": + // subject, already handled + case a.name == "else": + elseExpr = a.value + case ok && a.name == key: + matched = a.value + } + } + switch { + case matched != nil: + return evalNumber(matched, sc) + case elseExpr != nil: + return evalNumber(elseExpr, sc) + default: + return numberVal(0), nil + } +} + +func evalClamp(n callExpr, sc *scope) (value, error) { + args, err := positionalArgs(n) + if err != nil { + return value{}, err + } + if len(args) != 3 { + return value{}, fmt.Errorf("clamp expects 3 arguments (x, lo, hi), got %d", len(args)) + } + x, err := evalToNumber(args[0], sc) + if err != nil { + return value{}, err + } + lo, err := evalToNumber(args[1], sc) + if err != nil { + return value{}, err + } + hi, err := evalToNumber(args[2], sc) + if err != nil { + return value{}, err + } + if x < lo { + x = lo + } + if x > hi { + x = hi + } + return numberVal(x), nil +} + +func evalMinMax(n callExpr, sc *scope) (value, error) { + args, err := positionalArgs(n) + if err != nil { + return value{}, err + } + if len(args) == 0 { + return value{}, fmt.Errorf("%s expects at least one argument", n.name) + } + acc, err := evalToNumber(args[0], sc) + if err != nil { + return value{}, err + } + for _, a := range args[1:] { + f, err := evalToNumber(a, sc) + if err != nil { + return value{}, err + } + if n.name == "min" && f < acc { + acc = f + } + if n.name == "max" && f > acc { + acc = f + } + } + return numberVal(acc), nil +} + +// present(ref) reports whether a metadata/stat reference is set; missing is its +// negation. It is defined for any argument (a literal is always "present"). +func evalPresence(n callExpr, sc *scope) (value, error) { + args, err := positionalArgs(n) + if err != nil { + return value{}, err + } + if len(args) != 1 { + return value{}, fmt.Errorf("%s expects exactly one argument, got %d", n.name, len(args)) + } + v, err := eval(args[0], sc) + if err != nil { + return value{}, err + } + present := v.kind != kAbsent + if n.name == "missing" { + return boolVal(!present), nil + } + return boolVal(present), nil +} + +// setFuncArity is the required positional argument count per set function. A +// value of -1 means "1 or 2". +var setFuncArity = map[string]int{ + "count": -1, "has": 1, + "any": 2, "all": 2, "frac": 2, "avg": 2, "sum": 2, +} + +func evalSetFunc(n callExpr, sc *scope) (value, error) { + args, err := positionalArgs(n) + if err != nil { + return value{}, err + } + want := setFuncArity[n.name] + switch { + case want == -1 && (len(args) < 1 || len(args) > 2): + return value{}, fmt.Errorf("%s expects 1 or 2 arguments, got %d", n.name, len(args)) + case want != -1 && len(args) != want: + return value{}, fmt.Errorf("%s expects %d arguments, got %d", n.name, want, len(args)) + } + + setValue, err := eval(args[0], sc) + if err != nil { + return value{}, err + } + if setValue.kind != kSet { + return value{}, fmt.Errorf("%s expects a link set (out./in./bi.) as its first argument", n.name) + } + nodes := setValue.set + + // Functions with no predicate. + switch n.name { + case "has": + return boolVal(len(nodes) > 0), nil + case "count": + if len(args) == 1 { + return numberVal(float64(len(nodes))), nil + } + } + + pred := args[1] + switch n.name { + case "count": + matches, err := countMatches(nodes, pred, sc) + if err != nil { + return value{}, err + } + return numberVal(float64(matches)), nil + case "any": + matches, err := countMatches(nodes, pred, sc) + if err != nil { + return value{}, err + } + return boolVal(matches > 0), nil + case "all": + matches, err := countMatches(nodes, pred, sc) + if err != nil { + return value{}, err + } + return boolVal(matches == len(nodes)), nil + case "frac": + if len(nodes) == 0 { + return numberVal(0), nil + } + matches, err := countMatches(nodes, pred, sc) + if err != nil { + return value{}, err + } + return numberVal(float64(matches) / float64(len(nodes))), nil + case "avg", "sum": + var total float64 + for _, target := range nodes { + f, err := evalToNumber(pred, sc.child(target)) + if err != nil { + return value{}, err + } + total += f + } + if n.name == "avg" { + if len(nodes) == 0 { + return numberVal(0), nil + } + return numberVal(total / float64(len(nodes))), nil + } + return numberVal(total), nil + } + return value{}, fmt.Errorf("omegadsl: unhandled set function %q", n.name) +} + +func countMatches(nodes []Node, pred expr, sc *scope) (int, error) { + matches := 0 + for _, target := range nodes { + v, err := eval(pred, sc.child(target)) + if err != nil { + return 0, err + } + if truthy(v) { + matches++ + } + } + return matches, nil +} + +// evalNumber evaluates e and coerces the result to a number value. +func evalNumber(e expr, sc *scope) (value, error) { + f, err := evalToNumber(e, sc) + if err != nil { + return value{}, err + } + return numberVal(f), nil +} + +func evalToNumber(e expr, sc *scope) (float64, error) { + v, err := eval(e, sc) + if err != nil { + return 0, err + } + return asNumber(v) +} diff --git a/pkg/keg/omegadsl/lexer.go b/pkg/keg/omegadsl/lexer.go new file mode 100644 index 0000000..797e6c6 --- /dev/null +++ b/pkg/keg/omegadsl/lexer.go @@ -0,0 +1,269 @@ +package omegadsl + +import ( + "fmt" + "strconv" + "strings" + "unicode" + "unicode/utf8" +) + +type tokKind int + +const ( + tEOF tokKind = iota + tNumber + tString + tIdent // bare word; the parser classifies keywords/namespaces by text + tQIdent // backtick-quoted name, always a plain name (never a keyword) + tLParen + tRParen + tLBrace + tRBrace + tComma + tColon + tDot + tArrow // => + tAssign // = + tEq // == + tNeq // != + tLt + tLte + tGt + tGte + tPlus + tMinus + tStar + tSlash +) + +type token struct { + kind tokKind + text string + num float64 + line int +} + +func (t token) String() string { + switch t.kind { + case tEOF: + return "end of input" + case tNumber: + return fmt.Sprintf("number %v", t.num) + case tString: + return fmt.Sprintf("string %q", t.text) + case tIdent, tQIdent: + return fmt.Sprintf("%q", t.text) + default: + return fmt.Sprintf("%q", t.text) + } +} + +// lex tokenizes src. Whitespace, newlines, and # comments are insignificant and +// dropped; the language is layout-insensitive. +func lex(src string) ([]token, error) { + var toks []token + line := 1 + i := 0 + for i < len(src) { + c := src[i] + switch { + case c == '\n': + line++ + i++ + case c == ' ' || c == '\t' || c == '\r': + i++ + case c == '#': + for i < len(src) && src[i] != '\n' { + i++ + } + case c == '(': + toks = append(toks, token{kind: tLParen, text: "(", line: line}) + i++ + case c == ')': + toks = append(toks, token{kind: tRParen, text: ")", line: line}) + i++ + case c == '{': + toks = append(toks, token{kind: tLBrace, text: "{", line: line}) + i++ + case c == '}': + toks = append(toks, token{kind: tRBrace, text: "}", line: line}) + i++ + case c == ',': + toks = append(toks, token{kind: tComma, text: ",", line: line}) + i++ + case c == ':': + toks = append(toks, token{kind: tColon, text: ":", line: line}) + i++ + case c == '.': + toks = append(toks, token{kind: tDot, text: ".", line: line}) + i++ + case c == '+': + toks = append(toks, token{kind: tPlus, text: "+", line: line}) + i++ + case c == '-': + toks = append(toks, token{kind: tMinus, text: "-", line: line}) + i++ + case c == '*': + toks = append(toks, token{kind: tStar, text: "*", line: line}) + i++ + case c == '/': + toks = append(toks, token{kind: tSlash, text: "/", line: line}) + i++ + case c == '=': + if i+1 < len(src) && src[i+1] == '>' { + toks = append(toks, token{kind: tArrow, text: "=>", line: line}) + i += 2 + } else if i+1 < len(src) && src[i+1] == '=' { + toks = append(toks, token{kind: tEq, text: "==", line: line}) + i += 2 + } else { + toks = append(toks, token{kind: tAssign, text: "=", line: line}) + i++ + } + case c == '!': + if i+1 < len(src) && src[i+1] == '=' { + toks = append(toks, token{kind: tNeq, text: "!=", line: line}) + i += 2 + } else { + return nil, fmt.Errorf("line %d: unexpected %q (did you mean \"!=\" or \"not\"?)", line, "!") + } + case c == '<': + if i+1 < len(src) && src[i+1] == '=' { + toks = append(toks, token{kind: tLte, text: "<=", line: line}) + i += 2 + } else { + toks = append(toks, token{kind: tLt, text: "<", line: line}) + i++ + } + case c == '>': + if i+1 < len(src) && src[i+1] == '=' { + toks = append(toks, token{kind: tGte, text: ">=", line: line}) + i += 2 + } else { + toks = append(toks, token{kind: tGt, text: ">", line: line}) + i++ + } + case c == '"': + text, n, err := lexString(src[i:], line) + if err != nil { + return nil, err + } + toks = append(toks, token{kind: tString, text: text, line: line}) + i += n + case c == '`': + text, n, err := lexQuotedName(src[i:], line) + if err != nil { + return nil, err + } + toks = append(toks, token{kind: tQIdent, text: text, line: line}) + i += n + case c >= '0' && c <= '9': + num, text, n, err := lexNumber(src[i:], line) + if err != nil { + return nil, err + } + toks = append(toks, token{kind: tNumber, text: text, num: num, line: line}) + i += n + case isIdentStart(rune(c)): + text, n := lexIdent(src[i:]) + toks = append(toks, token{kind: tIdent, text: text, line: line}) + i += n + default: + r, _ := utf8.DecodeRuneInString(src[i:]) + return nil, fmt.Errorf("line %d: unexpected character %q", line, r) + } + } + toks = append(toks, token{kind: tEOF, text: "", line: line}) + return toks, nil +} + +func isIdentStart(r rune) bool { + return r == '_' || unicode.IsLetter(r) +} + +func isIdentCont(r rune) bool { + return r == '_' || unicode.IsLetter(r) || unicode.IsDigit(r) +} + +func lexIdent(s string) (string, int) { + i := 0 + for i < len(s) { + r, w := utf8.DecodeRuneInString(s[i:]) + if !isIdentCont(r) { + break + } + i += w + } + return s[:i], i +} + +func lexNumber(s string, line int) (float64, string, int, error) { + i := 0 + for i < len(s) && s[i] >= '0' && s[i] <= '9' { + i++ + } + if i < len(s) && s[i] == '.' { + i++ + start := i + for i < len(s) && s[i] >= '0' && s[i] <= '9' { + i++ + } + if i == start { + return 0, "", 0, fmt.Errorf("line %d: malformed number %q", line, s[:i]) + } + } + // Reject numbers glued to a letter (e.g. 90d): duration literals are not yet + // supported, and a bare "90d" is otherwise a confusing two-token sequence. + if i < len(s) { + if r, _ := utf8.DecodeRuneInString(s[i:]); isIdentStart(r) { + return 0, "", 0, fmt.Errorf("line %d: unexpected %q after number %q (duration literals are not supported yet)", line, string(r), s[:i]) + } + } + text := s[:i] + num, err := strconv.ParseFloat(text, 64) + if err != nil { + return 0, "", 0, fmt.Errorf("line %d: invalid number %q", line, text) + } + return num, text, i, nil +} + +func lexString(s string, line int) (string, int, error) { + var b strings.Builder + i := 1 // skip opening quote + for i < len(s) { + c := s[i] + if c == '\\' && i+1 < len(s) { + b.WriteByte(s[i+1]) + i += 2 + continue + } + if c == '"' { + return b.String(), i + 1, nil + } + if c == '\n' { + break + } + b.WriteByte(c) + i++ + } + return "", 0, fmt.Errorf("line %d: unterminated string literal", line) +} + +func lexQuotedName(s string, line int) (string, int, error) { + i := 1 // skip opening backtick + for i < len(s) { + if s[i] == '`' { + name := s[1:i] + if name == "" { + return "", 0, fmt.Errorf("line %d: empty backtick-quoted name", line) + } + return name, i + 1, nil + } + if s[i] == '\n' { + break + } + i++ + } + return "", 0, fmt.Errorf("line %d: unterminated backtick-quoted name", line) +} diff --git a/pkg/keg/omegadsl/omegadsl.go b/pkg/keg/omegadsl/omegadsl.go new file mode 100644 index 0000000..767a0f2 --- /dev/null +++ b/pkg/keg/omegadsl/omegadsl.go @@ -0,0 +1,96 @@ +// Package omegadsl implements the Omega DSL — a small expression language for +// keg schema maturity calculations. A program computes a single node's omega +// score (a float64 clamped to [0,1]) from the node's metadata and its links. +// +// The language, grammar, and semantics are specified in the tapper-hub repo at +// .tapper/specs/omega-dsl.md and .tapper/specs/omega-dsl.ebnf. This package is +// the reference implementation of that spec, minus temporal features (age() and +// duration literals), which are deferred because they interact poorly with +// deterministic snapshot replay. +// +// The package is self-contained (standard library only). Callers supply an Env +// that resolves the current node's metadata and its directional link sets; the +// evaluator never touches the keg graph directly. +package omegadsl + +// Direction selects which edges a link-set accessor (out./in./bi.) traverses. +type Direction int + +const ( + // DirOut is a forward link: nodes the current node points to. + DirOut Direction = iota + // DirIn is a backlink: nodes that point to the current node. + DirIn + // DirBi is the deduped union of forward links and backlinks. + DirBi +) + +func (d Direction) String() string { + switch d { + case DirOut: + return "out" + case DirIn: + return "in" + case DirBi: + return "bi" + default: + return "unknown" + } +} + +// Node is the evaluator's view of a single node — the scored node or a linked +// one. It exposes the two scalar namespaces a program can read: metadata +// (author-defined schema fields, addressed as meta.) and stats (computed +// node stats such as accessCount or links, addressed as stat.). Values are +// scalars rendered as strings; absent fields report ok == false. +type Node interface { + // Meta returns a scalar metadata field value and whether it is present. + Meta(field string) (string, bool) + // Stat returns a scalar node-stat value and whether it is present. + Stat(field string) (string, bool) +} + +// Env is the evaluator's view of the node being scored. It is a Node (its own +// meta.* and stat.*) plus LinkSet, which resolves a named relation in a given +// direction (out./in./bi.) to the set of linked nodes. LinkSet returns an error +// when the relation name is not declared by the schema. +type Env interface { + Node + LinkSet(relation string, dir Direction) ([]Node, error) +} + +// Program is a parsed, reusable Omega DSL program. It is safe to Eval +// concurrently against different Envs; parsing is done once via Parse. +type Program struct { + lets []letBinding + omega expr + rels []string // distinct relation names referenced by out./in./bi. +} + +// Parse compiles Omega DSL source into a Program. It returns an error on any +// lexical or syntactic problem; it does not evaluate the program. +func Parse(src string) (*Program, error) { + return parse(src) +} + +// Relations lists the distinct relation names the program references through +// out./in./bi. accessors. The schema layer uses this to verify every referenced +// relation is actually declared. +func (p *Program) Relations() []string { + if p == nil { + return nil + } + out := make([]string, len(p.rels)) + copy(out, p.rels) + return out +} + +// Eval runs the program against env and returns the clamped omega score in +// [0,1]. It returns an error if evaluation hits a type or lookup failure (for +// example, arithmetic on a non-numeric field, or a link-set inside a link-set). +func (p *Program) Eval(env Env) (float64, error) { + if p == nil { + return 0, errEmptyProgram + } + return evalProgram(p, env) +} diff --git a/pkg/keg/omegadsl/omegadsl_test.go b/pkg/keg/omegadsl/omegadsl_test.go new file mode 100644 index 0000000..c24caf7 --- /dev/null +++ b/pkg/keg/omegadsl/omegadsl_test.go @@ -0,0 +1,279 @@ +package omegadsl + +import ( + "math" + "testing" +) + +// mockNode is a test Node with fixed meta/stat maps. +type mockNode struct { + meta map[string]string + stat map[string]string +} + +func (m mockNode) Meta(f string) (string, bool) { v, ok := m.meta[f]; return v, ok } +func (m mockNode) Stat(f string) (string, bool) { v, ok := m.stat[f]; return v, ok } + +// mockEnv is a test Env: the scored node plus its directional link sets keyed by +// relation name. +type mockEnv struct { + mockNode + out map[string][]Node + in map[string][]Node + rel map[string]bool // declared relations (for LinkSet error behavior) +} + +func (e mockEnv) LinkSet(rel string, dir Direction) ([]Node, error) { + if e.rel != nil && !e.rel[rel] { + return nil, errUnknownRelation(rel) + } + switch dir { + case DirOut: + return e.out[rel], nil + case DirIn: + return e.in[rel], nil + case DirBi: + return dedupeNodes(append(append([]Node{}, e.out[rel]...), e.in[rel]...)), nil + default: + return nil, nil + } +} + +type unknownRel string + +func (u unknownRel) Error() string { return "unknown relation " + string(u) } +func errUnknownRelation(r string) error { return unknownRel(r) } + +func dedupeNodes(ns []Node) []Node { + seen := map[Node]bool{} + out := make([]Node, 0, len(ns)) + for _, n := range ns { + if !seen[n] { + seen[n] = true + out = append(out, n) + } + } + return out +} + +// node returns a pointer so link targets are hashable by identity (the real Env +// dedupes bi. sets by NodeId; mockNode holds maps and cannot be a map key). +func node(meta map[string]string) Node { return &mockNode{meta: meta} } + +func evalSrc(t *testing.T, src string, env Env) float64 { + t.Helper() + prog, err := Parse(src) + if err != nil { + t.Fatalf("parse error: %v\nsource:\n%s", err, src) + } + got, err := prog.Eval(env) + if err != nil { + t.Fatalf("eval error: %v\nsource:\n%s", err, src) + } + return got +} + +func assertDelta(t *testing.T, got, want float64) { + t.Helper() + if math.Abs(got-want) > 1e-9 { + t.Fatalf("got %v, want %v", got, want) + } +} + +func TestParseErrors(t *testing.T) { + cases := map[string]string{ + "no omega": `let a = 1`, + "duplicate omega": "omega = 1\nomega = 0", + "unknown identifier": `omega = status`, + "unknown function": `omega = frobnicate(1)`, + "bare namespace": `omega = meta`, + "reserved let name": `let meta = 1`, + "empty match": `omega = match meta.x { }`, + "unterminated string": `omega = "abc`, + "duration literal": `omega = 90d`, + } + for name, src := range cases { + t.Run(name, func(t *testing.T) { + if _, err := Parse(src); err == nil { + t.Fatalf("expected parse error for %q", src) + } + }) + } +} + +func TestMatchAndScore(t *testing.T) { + env := mockEnv{mockNode: mockNode{meta: map[string]string{"status": "ready"}}} + + assertDelta(t, evalSrc(t, ` + omega = match meta.status { + done => 1.0 + ready => 0.6 + else => 0.0 + }`, env), 0.6) + + assertDelta(t, evalSrc(t, `omega = score(meta.status, done=1, ready=0.6, else=0)`, env), 0.6) + + // Absent field falls through to else. + empty := mockEnv{mockNode: mockNode{meta: map[string]string{}}} + assertDelta(t, evalSrc(t, `omega = score(meta.status, done=1, else=0.1)`, empty), 0.1) + // No else, no match -> 0. + assertDelta(t, evalSrc(t, `omega = score(meta.status, done=1)`, empty), 0) +} + +func TestArithmeticComparisonAndClamp(t *testing.T) { + env := mockEnv{mockNode: mockNode{ + meta: map[string]string{"count": "8"}, + stat: map[string]string{"accessCount": "25"}, + }} + assertDelta(t, evalSrc(t, `omega = clamp(meta.count / 10, 0, 1)`, env), 0.8) + assertDelta(t, evalSrc(t, `omega = clamp(stat.accessCount / 10, 0, 1)`, env), 1) + assertDelta(t, evalSrc(t, `omega = meta.count > 5`, env), 1) // bool -> 1.0 + assertDelta(t, evalSrc(t, `omega = meta.count < 5`, env), 0) // bool -> 0.0 + assertDelta(t, evalSrc(t, `omega = min(meta.count, 3) / 3`, env), 1) +} + +func TestStringEqualityRequiresQuotes(t *testing.T) { + env := mockEnv{mockNode: mockNode{meta: map[string]string{"status": "done"}}} + assertDelta(t, evalSrc(t, `omega = meta.status == "done"`, env), 1) + assertDelta(t, evalSrc(t, `omega = meta.status == "draft"`, env), 0) + // absent never equals anything + empty := mockEnv{mockNode: mockNode{meta: map[string]string{}}} + assertDelta(t, evalSrc(t, `omega = meta.status == "done"`, empty), 0) +} + +func TestLinkDirectionsAndSetFunctions(t *testing.T) { + children := []Node{ + node(map[string]string{"status": "done"}), + node(map[string]string{"status": "done"}), + node(map[string]string{"status": "draft"}), + node(map[string]string{"status": "done"}), + } + parents := []Node{node(map[string]string{"type": "task"})} + env := mockEnv{ + mockNode: mockNode{meta: map[string]string{"status": "ready"}}, + out: map[string][]Node{"children": children, "parent": parents}, + rel: map[string]bool{"children": true, "parent": true}, + } + + assertDelta(t, evalSrc(t, `omega = count(out.children) / 10`, env), 0.4) // omega clamps to [0,1] + assertDelta(t, evalSrc(t, `omega = has(out.parent)`, env), 1) + assertDelta(t, evalSrc(t, `omega = frac(out.children, meta.status == "done")`, env), 0.75) + assertDelta(t, evalSrc(t, `omega = any(out.children, meta.status == "draft")`, env), 1) + assertDelta(t, evalSrc(t, `omega = all(out.children, meta.status == "done")`, env), 0) + assertDelta(t, evalSrc(t, `omega = avg(out.children, score(meta.status, done=1, else=0))`, env), 0.75) +} + +func TestBidirectionalAndBacklinks(t *testing.T) { + a := node(map[string]string{"n": "a"}) + b := node(map[string]string{"n": "b"}) + env := mockEnv{ + mockNode: mockNode{meta: map[string]string{}}, + out: map[string][]Node{"related": {a}}, + in: map[string][]Node{"related": {b}, "parent": {a, b, node(nil)}}, + rel: map[string]bool{"related": true, "parent": true}, + } + assertDelta(t, evalSrc(t, `omega = count(in.parent) / 10`, env), 0.3) // backlinks (omega clamps) + assertDelta(t, evalSrc(t, `omega = count(bi.related) / 2`, env), 1) // union of {a} and {b} +} + +func TestSelfReferenceInPredicate(t *testing.T) { + kids := []Node{ + node(map[string]string{"priority": "high"}), + node(map[string]string{"priority": "low"}), + } + env := mockEnv{ + mockNode: mockNode{meta: map[string]string{"priority": "high"}}, + in: map[string][]Node{"parent": kids}, + rel: map[string]bool{"parent": true}, + } + // children whose priority matches THIS node's priority + assertDelta(t, evalSrc(t, `omega = frac(in.parent, meta.priority == self.meta.priority)`, env), 0.5) +} + +func TestUnknownRelationErrors(t *testing.T) { + env := mockEnv{mockNode: mockNode{meta: map[string]string{}}, rel: map[string]bool{"parent": true}} + prog, err := Parse(`omega = count(out.children)`) + if err != nil { + t.Fatalf("unexpected parse error: %v", err) + } + if _, err := prog.Eval(env); err == nil { + t.Fatal("expected eval error for undeclared relation") + } + // Relations() surfaces referenced relations for schema validation. + if got := prog.Relations(); len(got) != 1 || got[0] != "children" { + t.Fatalf("Relations() = %v, want [children]", got) + } +} + +func TestNestedLinkSetRejected(t *testing.T) { + env := mockEnv{ + mockNode: mockNode{meta: map[string]string{}}, + out: map[string][]Node{"children": {node(nil)}}, + rel: map[string]bool{"children": true}, + } + prog, err := Parse(`omega = frac(out.children, has(out.grandchildren))`) + if err != nil { + t.Fatalf("unexpected parse error: %v", err) + } + if _, err := prog.Eval(env); err == nil { + t.Fatal("expected error: link set inside a predicate") + } +} + +// TestSpecWorkedExamples reproduces the two hand-traced examples from +// .tapper/specs/omega-dsl.md §12 (task ≈ 0.757, project ≈ 0.833). +func TestSpecWorkedExamples(t *testing.T) { + t.Run("task", func(t *testing.T) { + children := []Node{ + node(map[string]string{"status": "done"}), + node(map[string]string{"status": "done"}), + node(map[string]string{"status": "done"}), + node(map[string]string{"status": "draft"}), + } + env := mockEnv{ + mockNode: mockNode{meta: map[string]string{"status": "ready"}}, + out: map[string][]Node{"parent": {node(nil)}, "children": children}, + rel: map[string]bool{"parent": true, "children": true}, + } + src := ` + let status_score = match meta.status { + done => 1.0 + ready => 0.6 + draft => 0.25 + else => 0.0 + } + omega = weighted { + 0.4 : status_score + 0.3 : has(out.parent) + 0.2 : frac(out.children, meta.status == "done") + }` + // (0.4*0.6 + 0.3*1 + 0.2*0.75) / 0.9 + want := (0.4*0.6 + 0.3*1.0 + 0.2*0.75) / 0.9 + assertDelta(t, evalSrc(t, src, env), want) + }) + + t.Run("project", func(t *testing.T) { + tasks := []Node{ + node(map[string]string{"status": "done"}), + node(map[string]string{"status": "done"}), + node(map[string]string{"status": "done"}), + node(map[string]string{"status": "done"}), + node(map[string]string{"status": "ready"}), + } + related := []Node{node(nil), node(nil)} + env := mockEnv{ + mockNode: mockNode{meta: map[string]string{}}, + in: map[string][]Node{"parent": tasks}, + out: map[string][]Node{"related": related}, + rel: map[string]bool{"parent": true, "related": true}, + } + src := ` + omega = weighted { + 0.5 : frac(in.parent, meta.status == "done") + 0.3 : clamp(count(in.parent) / 5, 0, 1) + 0.2 : min(count(bi.related), 3) / 3 + }` + want := (0.5*0.8 + 0.3*1.0 + 0.2*(2.0/3.0)) / 1.0 + assertDelta(t, evalSrc(t, src, env), want) + }) +} diff --git a/pkg/keg/omegadsl/parser.go b/pkg/keg/omegadsl/parser.go new file mode 100644 index 0000000..6acc9b3 --- /dev/null +++ b/pkg/keg/omegadsl/parser.go @@ -0,0 +1,507 @@ +package omegadsl + +import ( + "fmt" + "sort" +) + +// reserved words may not be used as let-binding names. +var reserved = map[string]bool{ + "let": true, "omega": true, "match": true, "else": true, "weighted": true, + "and": true, "or": true, "not": true, "true": true, "false": true, + "meta": true, "stat": true, "out": true, "in": true, "bi": true, "self": true, +} + +// knownFuncs is the built-in function set. setFuncs is the subset whose first +// argument must be a link set (out./in./bi.). +var knownFuncs = map[string]bool{ + "score": true, "clamp": true, "min": true, "max": true, + "present": true, "missing": true, + "count": true, "has": true, "any": true, "all": true, + "frac": true, "avg": true, "sum": true, +} + +var setFuncs = map[string]bool{ + "count": true, "has": true, "any": true, "all": true, + "frac": true, "avg": true, "sum": true, +} + +type parser struct { + toks []token + pos int + lets map[string]bool + rels map[string]bool +} + +func parse(src string) (*Program, error) { + toks, err := lex(src) + if err != nil { + return nil, err + } + p := &parser{toks: toks, lets: map[string]bool{}, rels: map[string]bool{}} + return p.parseProgram() +} + +// --- token cursor helpers ------------------------------------------------- + +func (p *parser) peek() token { return p.toks[p.pos] } +func (p *parser) peekAt(n int) token { + if p.pos+n >= len(p.toks) { + return p.toks[len(p.toks)-1] + } + return p.toks[p.pos+n] +} +func (p *parser) next() token { t := p.toks[p.pos]; p.pos++; return t } +func (p *parser) at(k tokKind) bool { return p.peek().kind == k } + +func (p *parser) isKeyword(word string) bool { + t := p.peek() + return t.kind == tIdent && t.text == word +} + +func (p *parser) expect(k tokKind, what string) (token, error) { + t := p.peek() + if t.kind != k { + return t, fmt.Errorf("line %d: expected %s, found %s", t.line, what, t) + } + return p.next(), nil +} + +// --- program & statements ------------------------------------------------- + +func (p *parser) parseProgram() (*Program, error) { + var lets []letBinding + var omega expr + sawOmega := false + + for !p.at(tEOF) { + switch { + case p.isKeyword("let"): + lb, err := p.parseLet() + if err != nil { + return nil, err + } + lets = append(lets, lb) + p.lets[lb.name] = true + case p.isKeyword("omega"): + if sawOmega { + return nil, fmt.Errorf("line %d: program assigns omega more than once", p.peek().line) + } + p.next() + if _, err := p.expect(tAssign, "'=' after omega"); err != nil { + return nil, err + } + e, err := p.parseExpr() + if err != nil { + return nil, err + } + omega = e + sawOmega = true + default: + t := p.peek() + return nil, fmt.Errorf("line %d: expected 'let' or 'omega', found %s", t.line, t) + } + } + if !sawOmega { + return nil, fmt.Errorf("program must assign omega (e.g. `omega = ...`)") + } + return &Program{lets: lets, omega: omega, rels: sortedKeys(p.rels)}, nil +} + +func (p *parser) parseLet() (letBinding, error) { + p.next() // 'let' + name := p.peek() + if name.kind != tIdent { + return letBinding{}, fmt.Errorf("line %d: expected a name after 'let', found %s", name.line, name) + } + if reserved[name.text] { + return letBinding{}, fmt.Errorf("line %d: %q is a reserved word and cannot be a binding name", name.line, name.text) + } + p.next() + if _, err := p.expect(tAssign, fmt.Sprintf("'=' after 'let %s'", name.text)); err != nil { + return letBinding{}, err + } + value, err := p.parseExpr() + if err != nil { + return letBinding{}, err + } + return letBinding{name: name.text, value: value}, nil +} + +// --- expressions (precedence climbing) ------------------------------------ + +func (p *parser) parseExpr() (expr, error) { return p.parseOr() } + +func (p *parser) parseOr() (expr, error) { + l, err := p.parseAnd() + if err != nil { + return nil, err + } + for p.isKeyword("or") { + p.next() + r, err := p.parseAnd() + if err != nil { + return nil, err + } + l = binaryExpr{op: "or", l: l, r: r} + } + return l, nil +} + +func (p *parser) parseAnd() (expr, error) { + l, err := p.parseNot() + if err != nil { + return nil, err + } + for p.isKeyword("and") { + p.next() + r, err := p.parseNot() + if err != nil { + return nil, err + } + l = binaryExpr{op: "and", l: l, r: r} + } + return l, nil +} + +func (p *parser) parseNot() (expr, error) { + if p.isKeyword("not") { + p.next() + x, err := p.parseNot() + if err != nil { + return nil, err + } + return unaryExpr{op: "not", x: x}, nil + } + return p.parseComparison() +} + +var compareOps = map[tokKind]string{ + tEq: "==", tNeq: "!=", tLt: "<", tLte: "<=", tGt: ">", tGte: ">=", +} + +func (p *parser) parseComparison() (expr, error) { + l, err := p.parseAdd() + if err != nil { + return nil, err + } + if op, ok := compareOps[p.peek().kind]; ok { + p.next() + r, err := p.parseAdd() + if err != nil { + return nil, err + } + return binaryExpr{op: op, l: l, r: r}, nil + } + return l, nil +} + +func (p *parser) parseAdd() (expr, error) { + l, err := p.parseMul() + if err != nil { + return nil, err + } + for p.at(tPlus) || p.at(tMinus) { + op := p.next().text + r, err := p.parseMul() + if err != nil { + return nil, err + } + l = binaryExpr{op: op, l: l, r: r} + } + return l, nil +} + +func (p *parser) parseMul() (expr, error) { + l, err := p.parseUnary() + if err != nil { + return nil, err + } + for p.at(tStar) || p.at(tSlash) { + op := p.next().text + r, err := p.parseUnary() + if err != nil { + return nil, err + } + l = binaryExpr{op: op, l: l, r: r} + } + return l, nil +} + +func (p *parser) parseUnary() (expr, error) { + if p.at(tMinus) { + p.next() + x, err := p.parseUnary() + if err != nil { + return nil, err + } + return unaryExpr{op: "-", x: x}, nil + } + return p.parsePrimary() +} + +func (p *parser) parsePrimary() (expr, error) { + t := p.peek() + switch t.kind { + case tNumber: + p.next() + return numberLit{val: t.num}, nil + case tString: + p.next() + return stringLit{val: t.text}, nil + case tLParen: + p.next() + e, err := p.parseExpr() + if err != nil { + return nil, err + } + if _, err := p.expect(tRParen, "')'"); err != nil { + return nil, err + } + return e, nil + case tIdent: + return p.parseIdentPrimary() + default: + return nil, fmt.Errorf("line %d: unexpected %s", t.line, t) + } +} + +func (p *parser) parseIdentPrimary() (expr, error) { + t := p.peek() + switch t.text { + case "true": + p.next() + return boolLit{val: true}, nil + case "false": + p.next() + return boolLit{val: false}, nil + case "match": + return p.parseMatch() + case "weighted": + return p.parseWeighted() + case "meta", "stat": + p.next() + name, err := p.parseNameAfterDot(t.text) + if err != nil { + return nil, err + } + return p.makeFieldRef(t.text, name, false), nil + case "out", "in", "bi": + p.next() + rel, err := p.parseNameAfterDot(t.text) + if err != nil { + return nil, err + } + p.rels[rel] = true + return linkSetRef{dir: directionFor(t.text), rel: rel, line: t.line}, nil + case "self": + return p.parseSelfRef() + case "and", "or", "not", "else", "let", "omega": + return nil, fmt.Errorf("line %d: unexpected keyword %q", t.line, t.text) + } + + // Function call? + if p.peekAt(1).kind == tLParen { + return p.parseCall() + } + + // Otherwise a bare identifier — must be a defined let binding. + if reserved[t.text] { + return nil, fmt.Errorf("line %d: %q must be followed by a field, e.g. %s.", t.line, t.text, t.text) + } + if !p.lets[t.text] { + return nil, fmt.Errorf("line %d: unknown identifier %q (use meta.%s for metadata, stat.%s for a stat, or %q for a string literal)", t.line, t.text, t.text, t.text, t.text) + } + p.next() + return letRef{name: t.text, line: t.line}, nil +} + +func (p *parser) parseSelfRef() (expr, error) { + self := p.next() // 'self' + if _, err := p.expect(tDot, "'.' after 'self'"); err != nil { + return nil, err + } + ns := p.peek() + if ns.kind != tIdent || (ns.text != "meta" && ns.text != "stat") { + return nil, fmt.Errorf("line %d: 'self.' must be followed by 'meta' or 'stat', found %s", ns.line, ns) + } + p.next() + name, err := p.parseNameAfterDot("self." + ns.text) + if err != nil { + return nil, err + } + _ = self + return p.makeFieldRef(ns.text, name, true), nil +} + +func (p *parser) makeFieldRef(ns, name string, self bool) expr { + if ns == "stat" { + return statRef{name: name, self: self} + } + return metaRef{name: name, self: self} +} + +// parseNameAfterDot consumes "." where name is an identifier or a +// backtick-quoted name (for hyphenated fields like meta.`due-date`). +func (p *parser) parseNameAfterDot(prefix string) (string, error) { + if _, err := p.expect(tDot, fmt.Sprintf("'.' after %q", prefix)); err != nil { + return "", err + } + name := p.peek() + if name.kind != tIdent && name.kind != tQIdent { + return "", fmt.Errorf("line %d: expected a name after %q, found %s", name.line, prefix+".", name) + } + p.next() + return name.text, nil +} + +func (p *parser) parseCall() (expr, error) { + name := p.next() // function name + if !knownFuncs[name.text] { + return nil, fmt.Errorf("line %d: unknown function %q", name.line, name.text) + } + if _, err := p.expect(tLParen, fmt.Sprintf("'(' after %q", name.text)); err != nil { + return nil, err + } + var args []callArg + if !p.at(tRParen) { + for { + arg, err := p.parseArg() + if err != nil { + return nil, err + } + args = append(args, arg) + if p.at(tComma) { + p.next() + continue + } + break + } + } + if _, err := p.expect(tRParen, fmt.Sprintf("')' to close %q", name.text)); err != nil { + return nil, err + } + return callExpr{name: name.text, args: args, line: name.line}, nil +} + +func (p *parser) parseArg() (callArg, error) { + // Named argument: (identifier | string) '=' expr (used by score). + first := p.peek() + if (first.kind == tIdent || first.kind == tString) && p.peekAt(1).kind == tAssign { + p.next() // name + p.next() // '=' + value, err := p.parseExpr() + if err != nil { + return callArg{}, err + } + return callArg{name: first.text, value: value}, nil + } + value, err := p.parseExpr() + if err != nil { + return callArg{}, err + } + return callArg{value: value}, nil +} + +func (p *parser) parseMatch() (expr, error) { + p.next() // 'match' + subject, err := p.parseExpr() + if err != nil { + return nil, err + } + if _, err := p.expect(tLBrace, "'{' after match subject"); err != nil { + return nil, err + } + var arms []matchArm + for !p.at(tRBrace) { + if p.at(tEOF) { + return nil, fmt.Errorf("line %d: unterminated match block", p.peek().line) + } + arm, err := p.parseMatchArm() + if err != nil { + return nil, err + } + arms = append(arms, arm) + } + p.next() // '}' + if len(arms) == 0 { + return nil, fmt.Errorf("match block needs at least one arm") + } + return matchExpr{subject: subject, arms: arms}, nil +} + +func (p *parser) parseMatchArm() (matchArm, error) { + t := p.peek() + var arm matchArm + if t.kind == tIdent && t.text == "else" { + arm.isElse = true + p.next() + } else { + switch t.kind { + case tIdent, tString, tQIdent, tNumber: + arm.pat = t.text + p.next() + default: + return matchArm{}, fmt.Errorf("line %d: match pattern must be an identifier, string, or number, found %s", t.line, t) + } + } + if _, err := p.expect(tArrow, "'=>' in match arm"); err != nil { + return matchArm{}, err + } + result, err := p.parseExpr() + if err != nil { + return matchArm{}, err + } + arm.result = result + return arm, nil +} + +func (p *parser) parseWeighted() (expr, error) { + p.next() // 'weighted' + if _, err := p.expect(tLBrace, "'{' after weighted"); err != nil { + return nil, err + } + var terms []weightedTerm + for !p.at(tRBrace) { + if p.at(tEOF) { + return nil, fmt.Errorf("line %d: unterminated weighted block", p.peek().line) + } + weight, err := p.parseExpr() + if err != nil { + return nil, err + } + if _, err := p.expect(tColon, "':' between weight and contribution"); err != nil { + return nil, err + } + contrib, err := p.parseExpr() + if err != nil { + return nil, err + } + terms = append(terms, weightedTerm{weight: weight, contrib: contrib}) + } + p.next() // '}' + if len(terms) == 0 { + return nil, fmt.Errorf("weighted block needs at least one term") + } + return weightedExpr{terms: terms}, nil +} + +func directionFor(word string) Direction { + switch word { + case "in": + return DirIn + case "bi": + return DirBi + default: + return DirOut + } +} + +func sortedKeys(m map[string]bool) []string { + out := make([]string, 0, len(m)) + for k := range m { + out = append(out, k) + } + sort.Strings(out) + return out +}