tiny-datalog 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,1455 @@
1
+ #!/usr/bin/env python3
2
+ """
3
+ datalog.py — a small Datalog engine with semi-naive evaluation and
4
+ stratified negation. Pure standard-library Python.
5
+
6
+ Syntax
7
+ ------
8
+ fact(a, b). % ground facts
9
+ edge(a, b) @ 3. % facts may carry a numeric weight
10
+ % (ignored here; used by semiring.py)
11
+ head(X) :- body(X, Y), not q(Y). % rules; `not` is stratified negation
12
+ % and # start line comments
13
+
14
+ Constants are lowercase identifiers, numbers (int or float), or quoted
15
+ strings. Compound terms like s(N) are *parsed* but rejected by
16
+ validation — banning function symbols is precisely the restriction that
17
+ makes Datalog terminate. For Horn clauses with function symbols, see the
18
+ top-down interpreter in prolog.py.
19
+ Variables start with an uppercase letter or underscore ('_' is anonymous).
20
+ `not` is reserved for negation.
21
+
22
+ Semantics
23
+ ---------
24
+ * Safety: every variable in a rule head, and every variable in a negated
25
+ body literal, must also appear in a positive body literal of that rule.
26
+ * Stratified negation: IDB predicates are partitioned into strata so that
27
+ no predicate depends (directly or transitively) on its own negation.
28
+ If negation occurs inside a recursive cycle, the program is rejected
29
+ and the offending cycle is reported.
30
+ * Semi-naive evaluation: each stratum is evaluated to fixpoint; after the
31
+ first round, recursive rules are re-evaluated only with the previous
32
+ round's new facts (the "delta") substituted into each recursive body
33
+ position in turn, instead of recomputing every join from scratch.
34
+ * Magic sets: with --magic, each query is answered by first rewriting the
35
+ program (adornments + magic predicates, left-to-right sideways
36
+ information passing) so that bottom-up evaluation only derives facts
37
+ relevant to the query's bound arguments — goal-directed evaluation
38
+ without giving up semi-naive. Negated subgoals are not specialised:
39
+ their predicates are included untransformed and computed in full, which
40
+ keeps the rewriting stratified whenever the original program is.
41
+ (Implementation: magic.py.)
42
+
43
+ Stratifiability is a *syntactic* condition; rejection by the stratified
44
+ engine does not by itself mean a program is semantically paradoxical.
45
+ For small programs, `--models` grounds the program and reports the
46
+ semantic story: all stable models (by exhaustive search) and the
47
+ well-founded (three-valued) model. (Implementation: semantics.py.)
48
+
49
+ This file is the core: AST, parser, safety validation, stratification,
50
+ and the semi-naive evaluator, plus the CLI. The Under the hood sections of
51
+ lessons 1-3 are a guided tour of how it all works.
52
+
53
+ CLI
54
+ ---
55
+ python3 datalog.py program.dl # print derived relations
56
+ python3 datalog.py --trace program.dl # + strata and per-round deltas
57
+ python3 datalog.py -q 'eats_in_cafe(X)' program.dl
58
+ python3 datalog.py --magic -q 'path(n5, X)' program.dl # goal-directed
59
+ python3 datalog.py --models program.dl # stable + well-founded models
60
+ """
61
+
62
+ from __future__ import annotations
63
+
64
+ import argparse
65
+ import math
66
+ import re
67
+ import sys
68
+ import time
69
+ from collections import defaultdict, deque
70
+ from dataclasses import dataclass
71
+
72
+
73
+ # ---------------------------------------------------------------------------
74
+ # AST
75
+ # ---------------------------------------------------------------------------
76
+
77
+ @dataclass(frozen=True)
78
+ class Var:
79
+ name: str
80
+
81
+ def __str__(self):
82
+ return "_" if self.anonymous else self.name
83
+
84
+ @property
85
+ def anonymous(self):
86
+ """True for a renamed `_`. The parser names each one `_#1`, `_#2`,
87
+ ...: '#' starts a comment, so no variable a user writes can have
88
+ that name, and none can collide with it. (prolog.py's renamings
89
+ append '#n', so a renamed `_` still starts '_#' and a renamed X,
90
+ as 'X#3', still prints as itself.)"""
91
+ return self.name.startswith("_#")
92
+
93
+
94
+ @dataclass(frozen=True)
95
+ class Const:
96
+ value: object # str or int
97
+
98
+ def __str__(self):
99
+ return _format_value(self.value)
100
+
101
+
102
+ @dataclass(frozen=True)
103
+ class Struct:
104
+ """A compound term like s(N) or cons(H, T). Parsed for prolog.py's
105
+ benefit; Datalog validation rejects it (the function-symbol ban)."""
106
+ functor: str
107
+ args: tuple
108
+
109
+ def __str__(self):
110
+ return "%s(%s)" % (self.functor, ", ".join(map(str, self.args)))
111
+
112
+
113
+ @dataclass(frozen=True)
114
+ class Atom:
115
+ pred: str
116
+ args: tuple
117
+
118
+ def __str__(self):
119
+ if not self.args:
120
+ return self.pred
121
+ return "%s(%s)" % (self.pred, ", ".join(map(str, self.args)))
122
+
123
+
124
+ @dataclass(frozen=True)
125
+ class Literal:
126
+ atom: Atom
127
+ negated: bool = False
128
+
129
+ def __str__(self):
130
+ return ("not " if self.negated else "") + str(self.atom)
131
+
132
+
133
+ @dataclass(frozen=True)
134
+ class Rule:
135
+ head: Atom
136
+ body: tuple # tuple of Literal; empty tuple => fact
137
+ weight: object = None # numeric fact annotation `@ w`; facts only
138
+ retract: bool = False # `fact~.` — an update for incremental.py
139
+
140
+ def __str__(self):
141
+ if not self.body:
142
+ if self.retract:
143
+ return "%s~." % self.head
144
+ if self.weight is not None:
145
+ return "%s @ %s." % (self.head, self.weight)
146
+ return "%s." % self.head
147
+ return "%s :- %s." % (self.head, ", ".join(map(str, self.body)))
148
+
149
+
150
+ class DatalogError(Exception):
151
+ pass
152
+
153
+
154
+ class ParseError(DatalogError):
155
+ pass
156
+
157
+
158
+ class SafetyError(DatalogError):
159
+ pass
160
+
161
+
162
+ class StratificationError(DatalogError):
163
+ def __init__(self, message, cycle=None):
164
+ super().__init__(message)
165
+ self.cycle = cycle or []
166
+
167
+
168
+ # ---------------------------------------------------------------------------
169
+ # Parser
170
+ # ---------------------------------------------------------------------------
171
+
172
+ # One regex, alternatives tried in order, each wrapped in a named group —
173
+ # whichever group matched tells us the token kind. Two orderings matter:
174
+ # `:-` must be tried somewhere `:` alone can't shadow it (there is no
175
+ # lone-colon token, so it's safe), and the number alternative must come
176
+ # before `dot`, so that in `edge(a, b) @ 3.5.` the "3.5" is one float
177
+ # token and the final "." still terminates the clause.
178
+ _TOKEN = re.compile(
179
+ r"""
180
+ (?P<ws>\s+)
181
+ | (?P<comment>[%\#][^\n]*)
182
+ | (?P<implies>:-)
183
+ | (?P<lparen>\() | (?P<rparen>\)) | (?P<comma>,) | (?P<at>@)
184
+ | (?P<retract>~)
185
+ | (?P<number>-?[0-9]+(?:\.[0-9]+)?(?:[eE][-+]?[0-9]+)?)
186
+ | (?P<dot>\.)
187
+ | (?P<string>"[^"\n]*"|'[^'\n]*')
188
+ | (?P<var>[A-Z_][A-Za-z0-9_]*)
189
+ | (?P<ident>[a-z][A-Za-z0-9_]*)
190
+ """,
191
+ re.VERBOSE,
192
+ )
193
+
194
+
195
+ def _num(text, line):
196
+ n = float(text) if any(c in text for c in ".eE") else int(text)
197
+ if isinstance(n, float) and not math.isfinite(n):
198
+ # 1e400 would become inf, which prints as a constant named inf
199
+ raise ParseError("line %d: number %s is too large" % (line, text))
200
+ return n
201
+
202
+
203
+ def _shown(tok):
204
+ """A token as an error message quotes it."""
205
+ return "end of input" if tok[0] == "eof" else repr(tok[1])
206
+
207
+
208
+ def _tokenize(text):
209
+ pos, line = 0, 1
210
+ tokens = []
211
+ while pos < len(text):
212
+ m = _TOKEN.match(text, pos)
213
+ if not m:
214
+ raise ParseError("line %d: unexpected character %r" % (line, text[pos]))
215
+ kind = m.lastgroup
216
+ value = m.group()
217
+ if kind not in ("ws", "comment"):
218
+ tokens.append((kind, value, line))
219
+ line += value.count("\n")
220
+ pos = m.end()
221
+ tokens.append(("eof", "", line))
222
+ return tokens
223
+
224
+
225
+ class _Parser:
226
+ """Recursive descent over the token stream — one method per grammar
227
+ rule, reading top to bottom:
228
+
229
+ program := clause*
230
+ clause := atom [ '@' number ] '.' | atom ':-' literal (',' literal)* '.'
231
+ literal := [ 'not' ] atom
232
+ atom := IDENT [ '(' term (',' term)* ')' ]
233
+ term := VARIABLE | NUMBER | STRING | IDENT [ '(' term... ')' ]
234
+
235
+ The last alternative of `term` (an identifier with arguments) is a
236
+ compound term like s(N) — parsed here so prolog.py can share this
237
+ parser, but rejected later by Datalog validation."""
238
+
239
+ def __init__(self, text):
240
+ self.tokens = _tokenize(text)
241
+ self.i = 0
242
+ self.fresh = 0 # counter for renaming each `_` to a fresh variable
243
+
244
+ def _peek(self):
245
+ return self.tokens[self.i]
246
+
247
+ def _next(self):
248
+ tok = self.tokens[self.i]
249
+ self.i += 1
250
+ return tok
251
+
252
+ def _expect(self, kind):
253
+ tok = self._next()
254
+ if tok[0] != kind:
255
+ raise ParseError("line %d: expected %s, got %s"
256
+ % (tok[2], kind, _shown(tok)))
257
+ return tok
258
+
259
+ def parse_program(self):
260
+ clauses = []
261
+ while self._peek()[0] != "eof":
262
+ clauses.append(self._parse_clause())
263
+ return clauses
264
+
265
+ def _parse_clause(self):
266
+ head = self._parse_atom()
267
+ body = ()
268
+ weight = None
269
+ retract = False
270
+ kind = self._peek()[0]
271
+ if kind == "at":
272
+ self._next()
273
+ tok = self._expect("number")
274
+ weight = _num(tok[1], tok[2])
275
+ elif kind == "retract":
276
+ self._next()
277
+ retract = True
278
+ elif kind == "implies":
279
+ self._next()
280
+ lits = [self._parse_literal()]
281
+ while self._peek()[0] == "comma":
282
+ self._next()
283
+ lits.append(self._parse_literal())
284
+ body = tuple(lits)
285
+ self._expect("dot")
286
+ return Rule(head, body, weight, retract)
287
+
288
+ def _parse_literal(self):
289
+ kind, value, line = self._peek()
290
+ negated = False
291
+ if kind == "ident" and value == "not":
292
+ self._next()
293
+ negated = True
294
+ if self._peek()[0] == "lparen":
295
+ raise ParseError("line %d: `not` is negation, not a predicate "
296
+ "— write `not p(X)`" % line)
297
+ return Literal(self._parse_atom(), negated)
298
+
299
+ def _parse_atom(self):
300
+ tok = self._expect("ident")
301
+ pred = tok[1]
302
+ if pred == "not":
303
+ raise ParseError("line %d: `not` is reserved for negation and "
304
+ "cannot name a predicate" % tok[2])
305
+ args = []
306
+ if self._peek()[0] == "lparen":
307
+ self._next()
308
+ args.append(self._parse_term())
309
+ while self._peek()[0] == "comma":
310
+ self._next()
311
+ args.append(self._parse_term())
312
+ self._expect("rparen")
313
+ return Atom(pred, tuple(args))
314
+
315
+ def _parse_term(self):
316
+ tok = self._next()
317
+ kind, value, line = tok
318
+ if kind == "var":
319
+ if value == "_":
320
+ self.fresh += 1
321
+ return Var("_#%d" % self.fresh) # see Var.anonymous
322
+ return Var(value)
323
+ if kind == "ident":
324
+ if value == "not":
325
+ raise ParseError("line %d: `not` is reserved for negation; "
326
+ "write the constant as \"not\"" % line)
327
+ if self._peek()[0] == "lparen":
328
+ self._next()
329
+ args = [self._parse_term()]
330
+ while self._peek()[0] == "comma":
331
+ self._next()
332
+ args.append(self._parse_term())
333
+ self._expect("rparen")
334
+ return Struct(value, tuple(args))
335
+ return Const(value)
336
+ if kind == "number":
337
+ return Const(_num(value, line))
338
+ if kind == "string":
339
+ return Const(value[1:-1])
340
+ raise ParseError("line %d: expected a term, got %s"
341
+ % (line, _shown(tok)))
342
+
343
+
344
+ def parse(text):
345
+ """Parse a Datalog program into a list of Rule (facts have empty body)."""
346
+ return _Parser(text).parse_program()
347
+
348
+
349
+ # ---------------------------------------------------------------------------
350
+ # Validation: arity consistency, groundness of facts, rule safety
351
+ # ---------------------------------------------------------------------------
352
+
353
+ AGGREGATES = {"count", "sum", "min", "max"}
354
+
355
+
356
+ def _aggregate_of(atom):
357
+ """The (index, functor, variable) of an aggregate term like sum(V) in
358
+ a rule head, or None. At most one aggregate per head."""
359
+ found = None
360
+ for i, a in enumerate(atom.args):
361
+ if isinstance(a, Struct) and a.functor in AGGREGATES \
362
+ and len(a.args) == 1 and isinstance(a.args[0], Var):
363
+ if found is not None:
364
+ raise SafetyError("at most one aggregate per head: %s" % atom)
365
+ found = (i, a.functor, a.args[0])
366
+ return found
367
+
368
+
369
+ def validate(clauses, arity=None):
370
+ """Check arities, ground facts, and safety. Returns {pred: arity}.
371
+ An `arity` seed map lets a caller check new clauses against an
372
+ already-loaded program's signature (incremental.py does this).
373
+ "Safety" is range restriction, and it is what makes every relation
374
+ finite: a variable may appear in a rule head, or under `not`, only if
375
+ a positive body literal also binds it. Without it, p(X) :- q(a)
376
+ would assert p of *everything*, and `not r(X)` with X unbound would
377
+ quantify over an open universe. The compound-term check is the
378
+ Datalog boundary itself — see the module docstring and prolog.py."""
379
+ arity = dict(arity) if arity is not None else {}
380
+
381
+ def check_arity(atom):
382
+ n = arity.setdefault(atom.pred, len(atom.args))
383
+ if n != len(atom.args):
384
+ raise SafetyError(
385
+ "predicate %s used with arity %d and %d" % (atom.pred, len(atom.args), n))
386
+
387
+ def check_term(a, rule):
388
+ if isinstance(a, Struct):
389
+ raise SafetyError(
390
+ "function symbols are not Datalog: term %s in %s. "
391
+ "Datalog bans compound terms so that bottom-up "
392
+ "evaluation always terminates; for Horn clauses with "
393
+ "function symbols use the top-down engine (prolog.py)."
394
+ % (a, rule))
395
+
396
+ for rule in clauses:
397
+ check_arity(rule.head)
398
+ # heads may carry one aggregate term, e.g. total(P, sum(A)); any
399
+ # other compound term is the function-symbol boundary
400
+ agg = _aggregate_of(rule.head)
401
+ for i, a in enumerate(rule.head.args):
402
+ if not (agg and i == agg[0]):
403
+ check_term(a, rule)
404
+ if agg and not rule.body:
405
+ raise SafetyError("an aggregate needs a rule body: %s" % rule)
406
+ # unreachable from parse() — the grammar's `@ weight` and `:- body`
407
+ # are exclusive branches — but validate() is the checkpoint for
408
+ # clauses a caller assembled directly, and semiring.py silently
409
+ # skips weighted rules rather than failing on them
410
+ if rule.weight is not None and rule.body:
411
+ raise SafetyError("only facts may carry an @ weight: %s" % rule)
412
+ for lit in rule.body:
413
+ check_arity(lit.atom)
414
+ for a in lit.atom.args:
415
+ check_term(a, rule)
416
+ if not rule.body:
417
+ if any(isinstance(a, Var) for a in rule.head.args):
418
+ raise SafetyError("fact is not ground: %s" % rule)
419
+ continue
420
+ positive_vars = {a.name for lit in rule.body if not lit.negated
421
+ for a in lit.atom.args if isinstance(a, Var)}
422
+ head_vars = [a for a in rule.head.args if isinstance(a, Var)]
423
+ if agg:
424
+ head_vars.append(agg[2]) # the aggregated variable
425
+ for a in head_vars:
426
+ if a.name not in positive_vars:
427
+ raise SafetyError(
428
+ "unsafe rule: head variable %s is not bound by a positive "
429
+ "body literal in: %s" % (a, rule))
430
+ for lit in rule.body:
431
+ if lit.negated:
432
+ for a in lit.atom.args:
433
+ if isinstance(a, Var) and a.name not in positive_vars:
434
+ raise SafetyError(
435
+ "unsafe rule: variable %s of negated literal %s is not "
436
+ "bound by a positive literal in: %s" % (a, lit, rule))
437
+ return arity
438
+
439
+
440
+ # ---------------------------------------------------------------------------
441
+ # Stratification
442
+ # ---------------------------------------------------------------------------
443
+
444
+ def _tarjan(nodes, edges):
445
+ """Strongly connected components; returns {node: scc_id}.
446
+
447
+ Why SCCs? A program is stratifiable exactly when no *cycle* of
448
+ dependencies contains a negative edge, and every cycle lives inside
449
+ one SCC — so the whole check reduces to: does any negative edge have
450
+ both endpoints in the same component? This is Tarjan's algorithm in
451
+ its iterative form (an explicit frame stack instead of recursion, so
452
+ a long dependency chain can't hit Python's recursion limit)."""
453
+ adj = defaultdict(list)
454
+ for u, v, _neg in edges:
455
+ adj[u].append(v)
456
+ index, low, scc = {}, {}, {}
457
+ stack, on_stack = [], set()
458
+ counter = 0
459
+ scc_id = 0
460
+ for root in sorted(nodes):
461
+ if root in index:
462
+ continue
463
+ index[root] = low[root] = counter
464
+ counter += 1
465
+ stack.append(root)
466
+ on_stack.add(root)
467
+ frames = [(root, iter(adj[root]))]
468
+ while frames:
469
+ node, it = frames[-1]
470
+ advanced = False
471
+ for child in it:
472
+ if child not in index:
473
+ index[child] = low[child] = counter
474
+ counter += 1
475
+ stack.append(child)
476
+ on_stack.add(child)
477
+ frames.append((child, iter(adj[child])))
478
+ advanced = True
479
+ break
480
+ elif child in on_stack:
481
+ low[node] = min(low[node], index[child])
482
+ if advanced:
483
+ continue
484
+ frames.pop()
485
+ if frames:
486
+ parent = frames[-1][0]
487
+ low[parent] = min(low[parent], low[node])
488
+ if low[node] == index[node]:
489
+ sid = scc_id
490
+ scc_id += 1
491
+ while True:
492
+ w = stack.pop()
493
+ on_stack.discard(w)
494
+ scc[w] = sid
495
+ if w == node:
496
+ break
497
+ return scc
498
+
499
+
500
+ def _find_cycle(u, v, edges, sccs, first_kind):
501
+ """Given a strict edge u -> v inside one SCC, return a cycle
502
+ [(from, to, kind), ...] from u back to u through v."""
503
+ sid = sccs[u]
504
+ adj = defaultdict(list)
505
+ for a, b, neg in edges:
506
+ if sccs.get(a) == sid and sccs.get(b) == sid:
507
+ adj[a].append((b, neg))
508
+ prev = {v: None}
509
+ queue = deque([v])
510
+ while queue:
511
+ n = queue.popleft()
512
+ if n == u:
513
+ break
514
+ for b, neg in adj[n]:
515
+ if b not in prev:
516
+ prev[b] = (n, neg)
517
+ queue.append(b)
518
+ path = []
519
+ n = u
520
+ while prev[n] is not None:
521
+ p, kind = prev[n]
522
+ path.append((p, n, kind))
523
+ n = p
524
+ path.reverse()
525
+ return [(u, v, first_kind)] + path
526
+
527
+
528
+ def _format_cycle(cycle):
529
+ parts = [cycle[0][0]]
530
+ for _u, v, kind in cycle:
531
+ parts.append(" --> " if kind == "+" else " --%s--> " % kind)
532
+ parts.append(v)
533
+ return "".join(parts)
534
+
535
+
536
+ def stratify(clauses):
537
+ """Assign a stratum (1-based int) to each IDB predicate.
538
+
539
+ Raises StratificationError, with the offending cycle attached, if
540
+ negation occurs inside a recursive cycle.
541
+ """
542
+ rules = [r for r in clauses if r.body]
543
+ idb = {r.head.pred for r in rules}
544
+ # Edges are labelled: "+" ordinary, "not" through negation, "agg"
545
+ # into an aggregating rule. Negation and aggregation both demand
546
+ # "finish that relation completely before I look" — so both are
547
+ # strict, and both are forbidden inside a cycle.
548
+ edges = set() # (head_pred, body_pred, kind): head depends on body
549
+ for r in rules:
550
+ aggregating = _aggregate_of(r.head) is not None
551
+ for lit in r.body:
552
+ if lit.atom.pred in idb:
553
+ kind = ("not" if lit.negated
554
+ else "agg" if aggregating else "+")
555
+ edges.add((r.head.pred, lit.atom.pred, kind))
556
+
557
+ sccs = _tarjan(idb, edges)
558
+ for (u, v, kind) in sorted(edges):
559
+ if kind != "+" and sccs.get(u) == sccs.get(v):
560
+ cycle = _find_cycle(u, v, edges, sccs, kind)
561
+ what = ("aggregation" if any(k == "agg" for _a, _b, k in cycle)
562
+ else "negation")
563
+ raise StratificationError(
564
+ "program is not stratifiable — %s occurs inside a "
565
+ "recursive cycle: %s. No stratum assignment exists, so the "
566
+ "program has no stratified model." % (what,
567
+ _format_cycle(cycle)),
568
+ cycle=cycle)
569
+
570
+ # Assign stratum numbers by relaxation: a predicate must sit at least
571
+ # as high as anything it depends on, and *strictly* higher than
572
+ # anything it depends on through negation ("compute that completely
573
+ # before I ask what's not in it"). The SCC check above guarantees no
574
+ # negative cycle, so these constraints have a finite solution and the
575
+ # loop terminates at the least one.
576
+ stratum = {p: 1 for p in idb}
577
+ changed = True
578
+ while changed:
579
+ changed = False
580
+ for (u, v, kind) in edges:
581
+ need = stratum[v] + (0 if kind == "+" else 1)
582
+ if stratum[u] < need:
583
+ stratum[u] = need
584
+ changed = True
585
+ return stratum
586
+
587
+
588
+ # ---------------------------------------------------------------------------
589
+ # Evaluation
590
+ # ---------------------------------------------------------------------------
591
+
592
+ _MISSING = object()
593
+ _EMPTY = frozenset()
594
+
595
+
596
+ def _match(args, tup, subst):
597
+ """Extend subst so that args == tup, or return None.
598
+
599
+ This is one-way unification (pattern matching): `tup` is always
600
+ ground, so a variable either takes the tuple's value or must agree
601
+ with its earlier binding, and a constant simply has to be equal.
602
+ Joins fall out for free — matching path(X, Y) then edge(Y, Z) under
603
+ one growing substitution *is* the join on Y."""
604
+ s = dict(subst)
605
+ for a, v in zip(args, tup):
606
+ if isinstance(a, Const):
607
+ if a.value != v:
608
+ return None
609
+ else:
610
+ bound = s.get(a.name, _MISSING)
611
+ if bound is _MISSING:
612
+ s[a.name] = v
613
+ elif bound != v:
614
+ return None
615
+ return s
616
+
617
+
618
+ class Program:
619
+ def __init__(self, clauses):
620
+ for c in clauses:
621
+ if c.retract:
622
+ raise SafetyError(
623
+ "retraction (%s) is an update, not a statement — a "
624
+ "static program simply wouldn't assert the fact. "
625
+ "Apply it to a live materialisation via incremental.py."
626
+ % c)
627
+ self.arity = validate(clauses)
628
+ self.facts = [r for r in clauses if not r.body]
629
+ self.rules = [r for r in clauses if r.body]
630
+ self.idb = {r.head.pred for r in self.rules}
631
+ self.strata = stratify(clauses)
632
+
633
+
634
+ class Engine:
635
+ """Bottom-up, stratum-by-stratum semi-naive evaluator.
636
+
637
+ Data representation, in full: `rels` maps each predicate name to a
638
+ Python set of ground tuples — path -> {("a","b"), ("a","c")}. That's
639
+ the whole database. Rules never delete (Datalog is monotone within a
640
+ stratum), so evaluation is: grow these sets until one full pass adds
641
+ nothing. Strata are computed once, then processed in order, so by
642
+ the time a negated literal is consulted its relation is finished."""
643
+
644
+ def __init__(self, program, naive=False):
645
+ self.program = program
646
+ self.naive = naive # True: skip the delta discipline
647
+ self.rels = defaultdict(set) # pred -> set of ground tuples
648
+ self.stats = [] # per-stratum iteration statistics
649
+ # derivation-order stamps: base facts 0, then one tick per
650
+ # absorbed round — --explain uses these to build well-founded
651
+ # derivation trees (a fact's premises always carry lower stamps)
652
+ self.first_seen = {}
653
+ self._stamp = 0
654
+ # cumulative evaluation seconds per rule (--trace reports the
655
+ # hottest ones -- the join that is eating your run)
656
+ self.rule_time = defaultdict(float)
657
+
658
+ def run(self):
659
+ for fact in self.program.facts:
660
+ tup = tuple(a.value for a in fact.head.args)
661
+ self.rels[fact.head.pred].add(tup)
662
+ self.first_seen.setdefault((fact.head.pred, tup), 0)
663
+ by_stratum = defaultdict(list)
664
+ for rule in self.program.rules:
665
+ by_stratum[self.program.strata[rule.head.pred]].append(rule)
666
+ for level in sorted(by_stratum):
667
+ self._eval_stratum(level, by_stratum[level])
668
+ return self.rels
669
+
670
+ def _eval_stratum(self, level, rules):
671
+ preds = {r.head.pred for r in rules}
672
+ stat = {"stratum": level, "preds": sorted(preds), "iterations": []}
673
+ self.stats.append(stat)
674
+ if self.naive:
675
+ self._eval_stratum_naive(rules, stat)
676
+ return
677
+
678
+ # Round 1: evaluate every rule of the stratum against the full db.
679
+ delta = defaultdict(set)
680
+ for rule in rules:
681
+ t0 = time.perf_counter()
682
+ for tup in self._produce(rule):
683
+ if tup not in self.rels[rule.head.pred]:
684
+ delta[rule.head.pred].add(tup)
685
+ self.rule_time[rule] += time.perf_counter() - t0
686
+ self._absorb(delta, stat)
687
+
688
+ # Recursive rules: a positive body literal names a stratum predicate.
689
+ recursive = []
690
+ for rule in rules:
691
+ occs = [i for i, lit in enumerate(rule.body)
692
+ if not lit.negated and lit.atom.pred in preds]
693
+ if occs:
694
+ recursive.append((rule, occs))
695
+
696
+ # Semi-naive rounds: substitute the previous round's delta into each
697
+ # recursive position in turn; every other literal reads the full
698
+ # (already-updated) relations, so no new derivation is missed and
699
+ # nothing is recomputed from only-old facts.
700
+ while delta:
701
+ new_delta = defaultdict(set)
702
+ for rule, occs in recursive:
703
+ head = rule.head.pred
704
+ t0 = time.perf_counter()
705
+ for i in occs:
706
+ if not delta.get(rule.body[i].atom.pred):
707
+ continue
708
+ for tup in self._eval_rule(rule, delta_occ=i, delta=delta):
709
+ if tup not in self.rels[head]:
710
+ new_delta[head].add(tup)
711
+ self.rule_time[rule] += time.perf_counter() - t0
712
+ delta = new_delta
713
+ self._absorb(delta, stat)
714
+
715
+ def _eval_stratum_naive(self, rules, stat):
716
+ """Naive evaluation: every rule against the whole database, every
717
+ round, until nothing new appears — no delta discipline, so every
718
+ already-known fact is re-derived every round. Deliberately
719
+ wasteful: run --naive --trace beside the default to watch
720
+ semi-naive earn its name (Lesson 2)."""
721
+ stat["produced"] = [] # total tuples derived per round
722
+ while True:
723
+ delta = defaultdict(set)
724
+ produced = 0
725
+ for rule in rules:
726
+ t0 = time.perf_counter()
727
+ for tup in self._produce(rule):
728
+ produced += 1
729
+ if tup not in self.rels[rule.head.pred]:
730
+ delta[rule.head.pred].add(tup)
731
+ self.rule_time[rule] += time.perf_counter() - t0
732
+ stat["produced"].append(produced)
733
+ self._absorb(delta, stat)
734
+ if not delta:
735
+ return
736
+
737
+ def _produce(self, rule):
738
+ """All head tuples one rule derives right now (aggregate-aware)."""
739
+ if _aggregate_of(rule.head):
740
+ return self._eval_aggregate(rule)
741
+ return self._eval_rule(rule)
742
+
743
+ def _absorb(self, delta, stat):
744
+ self._stamp += 1
745
+ for pred, tuples in delta.items():
746
+ self.rels[pred] |= tuples
747
+ for t in tuples:
748
+ self.first_seen.setdefault((pred, t), self._stamp)
749
+ stat["iterations"].append(
750
+ {p: len(ts) for p, ts in delta.items() if ts})
751
+
752
+ def _rule_substitutions(self, rule, delta_occ=None, delta=None,
753
+ seed=None):
754
+ """Every substitution satisfying the rule body (positives joined
755
+ first — they bind; negatives filter afterwards, against fully
756
+ computed lower strata). If delta_occ is given, the positive
757
+ literal at that body index reads from `delta` instead of the
758
+ full relations — the semi-naive restriction. A `seed`
759
+ substitution pre-binds variables (--explain uses this)."""
760
+ # Positives first (they bind variables), negatives filter afterwards.
761
+ ordered = sorted(range(len(rule.body)),
762
+ key=lambda i: rule.body[i].negated)
763
+ substs = [dict(seed) if seed else {}]
764
+ for i in ordered:
765
+ lit = rule.body[i]
766
+ if not substs:
767
+ return []
768
+ if lit.negated:
769
+ rel = self.rels.get(lit.atom.pred, _EMPTY)
770
+ substs = [s for s in substs
771
+ if self._instantiate(lit.atom, s) not in rel]
772
+ else:
773
+ if delta_occ is not None and i == delta_occ:
774
+ rel = delta.get(lit.atom.pred, _EMPTY)
775
+ else:
776
+ rel = self.rels.get(lit.atom.pred, _EMPTY)
777
+ args = lit.atom.args
778
+ new = []
779
+ for s in substs:
780
+ for tup in rel:
781
+ m = _match(args, tup, s)
782
+ if m is not None:
783
+ new.append(m)
784
+ substs = new
785
+ return substs
786
+
787
+ def _eval_rule(self, rule, delta_occ=None, delta=None):
788
+ """Yield head tuples derivable from one rule."""
789
+ for s in self._rule_substitutions(rule, delta_occ, delta):
790
+ yield self._instantiate(rule.head, s)
791
+
792
+ def _eval_aggregate(self, rule):
793
+ """Aggregate rules — total(P, sum(A)) :- charge(P, C, A). — group
794
+ the body's distinct *solutions* by the plain head arguments and
795
+ fold the aggregate over each group. Set semantics applies to
796
+ solutions (rows), not to the aggregated values: two different
797
+ charges of 50 sum to 100, and count gives the same answer
798
+ whichever bound variable you name — matching SQL and Soufflé.
799
+ Stratification has already guaranteed the body relations are
800
+ complete (aggregation edges are strict, like negation), so one
801
+ evaluation suffices."""
802
+ idx, func, _var = _aggregate_of(rule.head)
803
+ for key, values in self._aggregate_groups(rule).items():
804
+ out = list(key)
805
+ out.insert(idx, _fold(func, values, rule))
806
+ yield tuple(out)
807
+
808
+ def _aggregate_groups(self, rule):
809
+ """{group key: [aggregated value per distinct body solution]} for
810
+ an aggregate rule; the key is the head's plain arguments."""
811
+ idx, _func, var = _aggregate_of(rule.head)
812
+ groups = defaultdict(list)
813
+ seen = set()
814
+ for s in self._rule_substitutions(rule):
815
+ witness = tuple(sorted(s.items()))
816
+ if witness in seen:
817
+ continue
818
+ seen.add(witness)
819
+ key = tuple(a.value if isinstance(a, Const) else s[a.name]
820
+ for j, a in enumerate(rule.head.args) if j != idx)
821
+ groups[key].append(s[var.name])
822
+ return groups
823
+
824
+ @staticmethod
825
+ def _instantiate(atom, subst):
826
+ return tuple(a.value if isinstance(a, Const) else subst[a.name]
827
+ for a in atom.args)
828
+
829
+
830
+ def _fold(func, values, rule):
831
+ """Fold one group's values with count, sum, min, or max."""
832
+ try:
833
+ if func == "count":
834
+ agg = len(values)
835
+ elif func == "sum":
836
+ agg = sum(values)
837
+ elif func == "min":
838
+ agg = min(values)
839
+ else:
840
+ agg = max(values)
841
+ except TypeError:
842
+ raise DatalogError(
843
+ "cannot %s over mixed or non-numeric values in: %s"
844
+ % (func, rule))
845
+ except OverflowError: # e.g. a float added to an int beyond its range
846
+ raise DatalogError("%s overflowed in: %s" % (func, rule))
847
+ if isinstance(agg, float) and not math.isfinite(agg):
848
+ # a float sum can overflow; inf would print as a constant named inf
849
+ raise DatalogError("%s gives %s, which is not a finite number, in: %s"
850
+ % (func, agg, rule))
851
+ return agg
852
+
853
+
854
+ def run_program(text):
855
+ """Parse, stratify, and evaluate a program; return the Engine."""
856
+ engine = Engine(Program(parse(text)))
857
+ engine.run()
858
+ return engine
859
+
860
+
861
+ # ---------------------------------------------------------------------------
862
+ # CLI
863
+ # ---------------------------------------------------------------------------
864
+
865
+ def read_program(path):
866
+ """Read a program file, reporting the everyday mistakes — a typo in
867
+ the name, a directory, an unreadable file — as a DatalogError, which
868
+ every CLI here already knows how to print. Without this they escape
869
+ as a traceback, which tells a reader nothing they can act on."""
870
+ try:
871
+ with open(path) as fh:
872
+ return fh.read()
873
+ except OSError as exc:
874
+ raise DatalogError("cannot read %s: %s"
875
+ % (path, exc.strerror or exc))
876
+ except UnicodeDecodeError:
877
+ raise DatalogError("cannot read %s: not text (a binary file?)"
878
+ % path)
879
+
880
+
881
+ def _sort_key(tup):
882
+ # numbers sort numerically (and before strings); strings sort as text
883
+ return tuple((0, v) if isinstance(v, (int, float)) else (1, str(v))
884
+ for v in tup)
885
+
886
+
887
+ def _format_value(v):
888
+ """A value as source text that parses back to it: bare identifiers
889
+ stay bare, anything else is quoted — with single quotes if it
890
+ contains a double one. (The lexer has no escapes, so a string with
891
+ both kinds of quote cannot be written, and so never needs printing.)"""
892
+ if isinstance(v, str) and re.fullmatch(r"[a-z][A-Za-z0-9_]*", v) \
893
+ and v != "not":
894
+ return v
895
+ if isinstance(v, (int, float)):
896
+ return str(v)
897
+ return "'%s'" % v if '"' in v else '"%s"' % v
898
+
899
+
900
+ def format_atom(pred, tup):
901
+ if not tup:
902
+ return pred
903
+ return "%s(%s)" % (pred, ", ".join(_format_value(v) for v in tup))
904
+
905
+
906
+ def format_fact(pred, tup):
907
+ return format_atom(pred, tup) + "."
908
+
909
+
910
+ def _print_strata(program):
911
+ levels = defaultdict(list)
912
+ for pred, level in program.strata.items():
913
+ levels[level].append("%s/%d" % (pred, program.arity[pred]))
914
+ print("Stratification:")
915
+ for level in sorted(levels):
916
+ print(" stratum %d: %s" % (level, ", ".join(sorted(levels[level]))))
917
+
918
+
919
+ def _print_stats(engine):
920
+ print("Naive evaluation:" if engine.naive else "Semi-naive evaluation:")
921
+ for stat in engine.stats:
922
+ print(" stratum %d (%s):" % (stat["stratum"], ", ".join(stat["preds"])))
923
+ produced = stat.get("produced")
924
+ for n, round_ in enumerate(stat["iterations"], 1):
925
+ extra = ""
926
+ if produced and n <= len(produced):
927
+ extra = " (%d tuples derived)" % produced[n - 1]
928
+ if round_:
929
+ deltas = ", ".join("+%d %s" % (c, p)
930
+ for p, c in sorted(round_.items()))
931
+ print(" round %d: %s%s" % (n, deltas, extra))
932
+ else:
933
+ print(" round %d: no new facts — fixpoint%s" % (n, extra))
934
+ sizes = ", ".join("%s=%d" % (p, len(engine.rels.get(p, ())))
935
+ for p in stat["preds"])
936
+ print(" sizes: %s" % sizes)
937
+ hot = sorted(engine.rule_time.items(), key=lambda kv: -kv[1])[:3]
938
+ if hot and hot[0][1] >= 0.01:
939
+ print(" hottest rules:")
940
+ total = sum(engine.rule_time.values()) or 1.0
941
+ for rule, secs in hot:
942
+ print(" %5.2fs (%2.0f%%) %s"
943
+ % (secs, 100 * secs / total, rule))
944
+
945
+
946
+ def _atom_sort_key(atom):
947
+ pred, args = atom
948
+ return (pred, _sort_key(args))
949
+
950
+
951
+ def _format_atoms(atoms):
952
+ return " ".join(format_fact(p, t)
953
+ for p, t in sorted(atoms, key=_atom_sort_key))
954
+
955
+
956
+ def _print_models(clauses):
957
+ """Report the semantic story: stable models and the well-founded model."""
958
+ from tiny_datalog.semantics import (
959
+ ground_program, stable_models, well_founded)
960
+ try:
961
+ stratify(clauses)
962
+ print("Syntactic check: stratifiable.")
963
+ except StratificationError as exc:
964
+ print("Syntactic check: not stratifiable (%s)." % _format_cycle(exc.cycle))
965
+ print(" (Syntactic only — an unstratifiable program may still have "
966
+ "stable models.)")
967
+ grounding = ground_program(clauses)
968
+ facts = grounding[0]
969
+ try:
970
+ models = stable_models(clauses, grounding=grounding)
971
+ except DatalogError as exc:
972
+ # too big for exhaustive search -- but the well-founded model
973
+ # below takes polynomial time, so it is still worth reporting
974
+ print("Stable models: search skipped (%s)." % exc)
975
+ models = None
976
+ if models == []:
977
+ print("Stable models: none — no consistent two-valued model exists.")
978
+ elif models:
979
+ print("Stable models: %d" % len(models))
980
+ for i, m in enumerate(
981
+ sorted(models, key=lambda m: _format_atoms(m - facts)), 1):
982
+ print(" model %d: %s" % (i, _format_atoms(m - facts)
983
+ or "(EDB facts only)"))
984
+ true, undef = well_founded(clauses, grounding=grounding)
985
+ print("Well-founded model (three-valued):")
986
+ print(" true: %s" % (_format_atoms(true - facts) or "(EDB facts only)"))
987
+ print(" undefined: %s" % (_format_atoms(undef) or "(none)"))
988
+ return 0
989
+
990
+
991
+ def parse_goal(q):
992
+ """Parse a query/goal string into a single atom (no validation —
993
+ prolog.py uses this too, and its goals may carry compound terms).
994
+ The final '.' is optional; it is added as a token, not as text, so
995
+ that a trailing `% comment` cannot swallow it."""
996
+ parser = _Parser(q)
997
+ tokens = parser.tokens
998
+ if len(tokens) > 1 and tokens[-2][0] != "dot":
999
+ tokens.insert(-1, ("dot", ".", tokens[-1][2]))
1000
+ clauses = parser.parse_program()
1001
+ if len(clauses) != 1 or clauses[0].body:
1002
+ raise ParseError("query must be a single atom: %r" % q)
1003
+ if clauses[0].weight is not None or clauses[0].retract:
1004
+ raise ParseError("a query is a plain atom — no @ weight or ~ "
1005
+ "retraction: %r" % q)
1006
+ return clauses[0].head
1007
+
1008
+
1009
+ def check_query_atom(atom, arity=None):
1010
+ """Datalog-side validation of a query atom: no compound terms, and
1011
+ arity agreement with the program when known. The single home for
1012
+ these checks — the CLI, magic.py, and semiring.py all route here."""
1013
+ for a in atom.args:
1014
+ if isinstance(a, Struct):
1015
+ raise SafetyError(
1016
+ "function symbols are not Datalog: term %s in query %s "
1017
+ "(see prolog.py)" % (a, atom))
1018
+ if arity is not None and atom.pred in arity \
1019
+ and arity[atom.pred] != len(atom.args):
1020
+ raise SafetyError(
1021
+ "query %s has arity %d but %s is used with arity %d"
1022
+ % (atom, len(atom.args), atom.pred, arity[atom.pred]))
1023
+
1024
+
1025
+ def _parse_query_atom(q, arity=None):
1026
+ atom = parse_goal(q)
1027
+ check_query_atom(atom, arity)
1028
+ return atom
1029
+
1030
+
1031
+ def match_answers(atom, tuples):
1032
+ """The tuples matching a query atom: constants filter, variables bind."""
1033
+ return [tup for tup in tuples if _match(atom.args, tup, {}) is not None]
1034
+
1035
+
1036
+ def _print_answers(atom, tuples, suffix=""):
1037
+ print("?- %s%s" % (atom, suffix))
1038
+ answers = sorted(tuples, key=_sort_key)
1039
+ for tup in answers:
1040
+ print(" " + format_fact(atom.pred, tup))
1041
+ print(" (%d answer%s)" % (len(answers), "" if len(answers) == 1 else "s"))
1042
+
1043
+
1044
+ # ---------------------------------------------------------------------------
1045
+ # --explain: derivation trees
1046
+ # ---------------------------------------------------------------------------
1047
+ # Ask the engine WHY it believes a fact. The trick that keeps the tree
1048
+ # well-founded: every fact carries a derivation-order stamp (Engine
1049
+ # first_seen), and its first derivation necessarily used premises with
1050
+ # strictly smaller stamps — so searching for a rule instance whose
1051
+ # positive premises all precede the fact always succeeds and can never
1052
+ # justify a fact by itself.
1053
+
1054
+ def _derivation_of(engine, pred, tup, index=None):
1055
+ """A (rule, premises) justification for a derived fact, where every
1056
+ positive premise strictly precedes it in derivation order; None for
1057
+ base facts. premises is a list of (literal, ground_tuple). `index`
1058
+ is a lookup cache shared across one explanation (see _earlier_solution)."""
1059
+ stamp = engine.first_seen.get((pred, tup), 0)
1060
+ if stamp == 0:
1061
+ return None # stated in the program, whatever rules also say
1062
+ index = {} if index is None else index
1063
+ for rule in engine.program.rules:
1064
+ if rule.head.pred != pred:
1065
+ continue
1066
+ if _aggregate_of(rule.head):
1067
+ group = _aggregate_group(engine, rule, tup, index)
1068
+ if group is not None:
1069
+ return rule, group
1070
+ continue
1071
+ seed = _match(rule.head.args, tup, {})
1072
+ if seed is None:
1073
+ continue
1074
+ s = _earlier_solution(engine, rule, seed, stamp, index)
1075
+ if s is not None:
1076
+ return rule, [(lit, engine._instantiate(lit.atom, s))
1077
+ for lit in rule.body]
1078
+ return None
1079
+
1080
+
1081
+ def _earlier_solution(engine, rule, seed, stamp, index):
1082
+ """One body solution extending `seed` whose positive premises all
1083
+ carry stamps below `stamp`, or None. The evaluator's join scans
1084
+ whole relations, which is fine once per round but far too slow once
1085
+ per tree node — a 400-step chain would cost 400 full joins. So this
1086
+ join takes the head's bindings first, always picks next the positive
1087
+ literal with the most arguments already known, and looks its matches
1088
+ up in a hash index (built once per relation and pattern of known
1089
+ positions, then reused down the tree). Negatives filter at the end."""
1090
+ positives = [lit for lit in rule.body if not lit.negated]
1091
+ substs = [seed]
1092
+ known = set(seed)
1093
+
1094
+ def is_known(a):
1095
+ return isinstance(a, Const) or a.name in known
1096
+
1097
+ while positives and substs:
1098
+ lit = max(positives, key=lambda l: sum(map(is_known, l.atom.args)))
1099
+ positives.remove(lit)
1100
+ args, pred = lit.atom.args, lit.atom.pred
1101
+ cols = tuple(j for j, a in enumerate(args) if is_known(a))
1102
+ table = _index_on(engine, index, pred, cols)
1103
+ new = []
1104
+ for s in substs:
1105
+ key = tuple(args[j].value if isinstance(args[j], Const)
1106
+ else s[args[j].name] for j in cols)
1107
+ for t in table.get(key, ()):
1108
+ if engine.first_seen.get((pred, t), 0) < stamp:
1109
+ m = _match(args, t, s)
1110
+ if m is not None:
1111
+ new.append(m)
1112
+ substs = new
1113
+ known |= {a.name for a in args if isinstance(a, Var)}
1114
+ for s in substs:
1115
+ if all(engine._instantiate(lit.atom, s)
1116
+ not in engine.rels.get(lit.atom.pred, _EMPTY)
1117
+ for lit in rule.body if lit.negated):
1118
+ return s
1119
+ return None
1120
+
1121
+
1122
+ def _index_on(engine, index, pred, cols):
1123
+ """pred's tuples grouped by their values at positions `cols`."""
1124
+ if (pred, cols) not in index:
1125
+ table = defaultdict(list)
1126
+ for t in engine.rels.get(pred, _EMPTY):
1127
+ table[tuple(t[j] for j in cols)].append(t)
1128
+ index[pred, cols] = table
1129
+ return index[pred, cols]
1130
+
1131
+
1132
+ def _group_values(engine, rule, tup, index=None):
1133
+ """The aggregated values (one per distinct body solution) of the
1134
+ group that head tuple `tup` belongs to; [] if the group is empty.
1135
+ The rule's groups are computed once per `index` cache."""
1136
+ index = {} if index is None else index
1137
+ if rule not in index:
1138
+ index[rule] = engine._aggregate_groups(rule)
1139
+ idx = _aggregate_of(rule.head)[0]
1140
+ key = tuple(v for j, v in enumerate(tup) if j != idx)
1141
+ return index[rule].get(key, [])
1142
+
1143
+
1144
+ def _aggregate_group(engine, rule, tup, index=None):
1145
+ """If this aggregate rule computes exactly `tup` — its group exists
1146
+ *and* folds to tup's value — the contributing values, presented as a
1147
+ pseudo-premise; otherwise None."""
1148
+ idx, func, var = _aggregate_of(rule.head)
1149
+ values = _group_values(engine, rule, tup, index)
1150
+ if not values or _fold(func, values, rule) != tup[idx]:
1151
+ return None
1152
+ shown = ", ".join(_format_value(v) for v in
1153
+ sorted(values, key=lambda x:
1154
+ (0, x) if isinstance(x, (int, float))
1155
+ else (1, str(x))))
1156
+ return [("aggregate", "%s over %d body solution%s of %s: [%s]"
1157
+ % (func, len(values), "" if len(values) == 1 else "s", var,
1158
+ shown))]
1159
+
1160
+
1161
+ def explain(engine, pred, tup):
1162
+ """Build an indented derivation tree for one fact; returns the lines.
1163
+ Depth-first with an explicit stack rather than recursion, so a long
1164
+ derivation chain can't hit Python's recursion limit (the same reason
1165
+ _tarjan is iterative). Stack entries are (indent, pred, tup) for a
1166
+ fact still to explain, or (indent, None, text) for a finished line."""
1167
+ lines, shown, index = [], set(), {}
1168
+ stack = [(0, pred, tup)]
1169
+ while stack:
1170
+ indent, pred, tup = stack.pop()
1171
+ pad = " " * indent
1172
+ if pred is None:
1173
+ lines.append(pad + tup)
1174
+ continue
1175
+ label = format_atom(pred, tup)
1176
+ if (pred, tup) in shown:
1177
+ lines.append("%s%s (derivation shown above)" % (pad, label))
1178
+ continue
1179
+ derivation = _derivation_of(engine, pred, tup, index)
1180
+ if derivation is None:
1181
+ lines.append("%s%s (base fact)" % (pad, label))
1182
+ continue
1183
+ shown.add((pred, tup))
1184
+ rule, premises = derivation
1185
+ lines.append("%s%s [via %s]" % (pad, label, rule))
1186
+ children = []
1187
+ for item in premises:
1188
+ if item[0] == "aggregate":
1189
+ children.append((indent, None, " = %s" % item[1]))
1190
+ elif item[0].negated:
1191
+ children.append((indent, None,
1192
+ " not %s (absent from its completed "
1193
+ "stratum)" % format_atom(item[0].atom.pred,
1194
+ item[1])))
1195
+ else:
1196
+ children.append((indent + 1, item[0].atom.pred, item[1]))
1197
+ stack.extend(reversed(children)) # so the first premise pops first
1198
+ return lines
1199
+
1200
+
1201
+ def whynot(engine, pred, tup, lines=None):
1202
+ """Why is this ground fact NOT derived? For each rule that could
1203
+ head it, walk the body in evaluation order and report the first
1204
+ literal the join dies at. When the blocker is a negated literal,
1205
+ the negated atom *holds* -- so its positive derivation is the
1206
+ culprit, and it is explained inline (the complement's why)."""
1207
+ lines = [] if lines is None else lines
1208
+ label = format_atom(pred, tup)
1209
+ rules = [r for r in engine.program.rules if r.head.pred == pred]
1210
+ if not rules:
1211
+ lines.append("%s is not a stated fact, and no rule derives %s."
1212
+ % (label, pred))
1213
+ return lines
1214
+ lines.append("%s is not derived. Per rule:" % label)
1215
+ headless = 0
1216
+ for rule in rules:
1217
+ if _aggregate_of(rule.head):
1218
+ idx, func, _var = _aggregate_of(rule.head)
1219
+ values = _group_values(engine, rule, tup)
1220
+ lines.append(" via %s" % rule)
1221
+ if not values:
1222
+ lines.append(" blocked: no body solutions produce this "
1223
+ "group (an empty group yields no fact)")
1224
+ else:
1225
+ lines.append(" blocked: the group exists, but its %s is "
1226
+ "%s, not %s" % (func,
1227
+ _format_value(_fold(func, values, rule)),
1228
+ _format_value(tup[idx])))
1229
+ continue
1230
+ seed = _match(rule.head.args, tup, {})
1231
+ if seed is None:
1232
+ headless += 1
1233
+ continue
1234
+ lines.append(" via %s" % rule)
1235
+ substs = [dict(seed)]
1236
+ ordered = sorted(range(len(rule.body)),
1237
+ key=lambda i: rule.body[i].negated)
1238
+ blocked = None
1239
+ for i in ordered:
1240
+ lit = rule.body[i]
1241
+ if lit.negated:
1242
+ survivors = [s for s in substs
1243
+ if engine._instantiate(lit.atom, s)
1244
+ not in engine.rels.get(lit.atom.pred, _EMPTY)]
1245
+ else:
1246
+ rel = engine.rels.get(lit.atom.pred, _EMPTY)
1247
+ survivors = []
1248
+ for s in substs:
1249
+ for t in rel:
1250
+ m = _match(lit.atom.args, t, s)
1251
+ if m is not None:
1252
+ survivors.append(m)
1253
+ if not survivors:
1254
+ blocked = (lit, substs[0])
1255
+ break
1256
+ substs = survivors
1257
+ if blocked is None:
1258
+ lines.append(" (this rule does derive it -- "
1259
+ "the fact should exist; please report)")
1260
+ continue
1261
+ lit, s = blocked
1262
+ shown = Atom(lit.atom.pred,
1263
+ tuple(Const(s[a.name]) if isinstance(a, Var)
1264
+ and a.name in s else a for a in lit.atom.args))
1265
+ if lit.negated:
1266
+ inst = engine._instantiate(lit.atom, s)
1267
+ lines.append(" blocked at: not %s -- %s holds:"
1268
+ % (shown, format_atom(lit.atom.pred, inst)))
1269
+ for l in explain(engine, lit.atom.pred, inst):
1270
+ lines.append(" " + l)
1271
+ else:
1272
+ lines.append(" blocked at: %s -- no matching fact" % shown)
1273
+ if headless:
1274
+ lines.append(" (%d rule%s for %s cannot match this head and "
1275
+ "%s skipped)" % (headless, "" if headless == 1 else "s",
1276
+ pred,
1277
+ "was" if headless == 1 else "were"))
1278
+ return lines
1279
+
1280
+
1281
+ def _run_explain(q, engine):
1282
+ atom = _parse_query_atom(q, engine.program.arity)
1283
+ matches = sorted(match_answers(atom, engine.rels.get(atom.pred, ())),
1284
+ key=_sort_key)
1285
+ print("?- explain %s" % atom)
1286
+ if not matches:
1287
+ if all(isinstance(a, Const) for a in atom.args):
1288
+ tup = tuple(a.value for a in atom.args)
1289
+ for line in whynot(engine, atom.pred, tup):
1290
+ print(" " + line)
1291
+ else:
1292
+ print(" (no matching facts)")
1293
+ return
1294
+ for tup in matches:
1295
+ for line in explain(engine, atom.pred, tup):
1296
+ print(" " + line)
1297
+
1298
+
1299
+ def _run_query(q, engine):
1300
+ atom = _parse_query_atom(q, engine.program.arity)
1301
+ _print_answers(atom, match_answers(atom, engine.rels.get(atom.pred, ())))
1302
+
1303
+
1304
+ def _run_magic_query(q, clauses, trace):
1305
+ from tiny_datalog.magic import magic_transform
1306
+ atom = _parse_query_atom(q)
1307
+ transformed, answer_pred = magic_transform(clauses, atom)
1308
+ mprog = Program(transformed)
1309
+ mengine = Engine(mprog)
1310
+ if trace:
1311
+ print("Magic-sets rewriting (answer predicate %s):" % answer_pred)
1312
+ for c in transformed:
1313
+ print(" %s" % c)
1314
+ print()
1315
+ _print_strata(mprog)
1316
+ print()
1317
+ mengine.run()
1318
+ if trace:
1319
+ _print_stats(mengine)
1320
+ magic_total = sum(len(mengine.rels.get(p, ())) for p in mprog.idb)
1321
+ try:
1322
+ fengine = Engine(Program(clauses))
1323
+ fengine.run()
1324
+ full_total = sum(len(fengine.rels.get(p, ()))
1325
+ for p in fengine.program.idb)
1326
+ print("[magic] %d IDB facts derived vs %d under full evaluation"
1327
+ % (magic_total, full_total))
1328
+ except StratificationError:
1329
+ print("[magic] %d IDB facts derived (no full-evaluation baseline: "
1330
+ "the original program is not stratifiable)" % magic_total)
1331
+ except DatalogError as exc:
1332
+ print("[magic] %d IDB facts derived (no full-evaluation baseline: "
1333
+ "%s)" % (magic_total, exc))
1334
+ print()
1335
+ _print_answers(atom, match_answers(atom, mengine.rels.get(answer_pred, ())),
1336
+ suffix=" [magic]")
1337
+
1338
+
1339
+ def main(argv=None):
1340
+ ap = argparse.ArgumentParser(
1341
+ description="A small Datalog engine with semi-naive evaluation "
1342
+ "and stratified negation.")
1343
+ ap.add_argument("file", help="Datalog program (.dl)")
1344
+ ap.add_argument("-t", "--trace", action="store_true",
1345
+ help="print stratification and per-round delta statistics")
1346
+ ap.add_argument("-a", "--all", action="store_true",
1347
+ help="print EDB (input) relations too, not just derived ones")
1348
+ ap.add_argument("-q", "--query", action="append", default=[], metavar="ATOM",
1349
+ help="query, e.g. 'eats_in_cafe(X)' (repeatable)")
1350
+ ap.add_argument("-m", "--models", action="store_true",
1351
+ help="skip stratified evaluation; instead ground the "
1352
+ "program and report all stable models (exhaustive "
1353
+ "search, small programs only) and the well-founded "
1354
+ "three-valued model")
1355
+ ap.add_argument("--naive", action="store_true",
1356
+ help="evaluate naively (no delta discipline); with "
1357
+ "--trace, prints tuples-derived per round so the "
1358
+ "semi-naive comparison is measurable")
1359
+ ap.add_argument("-e", "--explain", action="append", default=[],
1360
+ metavar="ATOM",
1361
+ help="print a derivation tree for every fact matching "
1362
+ "the atom (repeatable)")
1363
+ ap.add_argument("-M", "--magic", action="store_true",
1364
+ help="answer each -q query via the magic-sets rewriting "
1365
+ "(goal-directed: only facts relevant to the query's "
1366
+ "bound arguments are derived); with --trace, also "
1367
+ "print the rewritten program and derivation counts")
1368
+ args = ap.parse_args(argv)
1369
+
1370
+ try:
1371
+ # every mode re-validates via Program / magic_transform /
1372
+ # ground_program, so parsing is all that must happen up front
1373
+ clauses = parse(read_program(args.file))
1374
+ except DatalogError as exc:
1375
+ print("error: %s" % exc, file=sys.stderr)
1376
+ return 1
1377
+
1378
+ if args.models:
1379
+ try:
1380
+ return _print_models(clauses)
1381
+ except DatalogError as exc:
1382
+ print("error: %s" % exc, file=sys.stderr)
1383
+ return 1
1384
+
1385
+ if args.magic:
1386
+ if not args.query:
1387
+ print("error: --magic requires at least one -q/--query",
1388
+ file=sys.stderr)
1389
+ return 1
1390
+ for q in args.query:
1391
+ try:
1392
+ _run_magic_query(q, clauses, args.trace)
1393
+ except StratificationError as exc:
1394
+ print("REJECTED: %s" % exc, file=sys.stderr)
1395
+ return 2
1396
+ except DatalogError as exc:
1397
+ print("error: %s" % exc, file=sys.stderr)
1398
+ return 1
1399
+ return 0
1400
+
1401
+ try:
1402
+ program = Program(clauses)
1403
+ except StratificationError as exc:
1404
+ print("REJECTED: %s" % exc, file=sys.stderr)
1405
+ print("(This is a syntactic verdict. Run with --models for the "
1406
+ "semantic one: stable models and the well-founded model.)",
1407
+ file=sys.stderr)
1408
+ return 2
1409
+ except DatalogError as exc:
1410
+ print("error: %s" % exc, file=sys.stderr)
1411
+ return 1
1412
+
1413
+ engine = Engine(program, naive=args.naive)
1414
+ if args.trace:
1415
+ _print_strata(program)
1416
+ print()
1417
+ try:
1418
+ engine.run()
1419
+ except DatalogError as exc: # e.g. sum over a non-number
1420
+ print("error: %s" % exc, file=sys.stderr)
1421
+ return 1
1422
+ if args.trace:
1423
+ _print_stats(engine)
1424
+ print()
1425
+
1426
+ if args.query or args.explain:
1427
+ for q in args.query:
1428
+ try:
1429
+ _run_query(q, engine)
1430
+ except DatalogError as exc:
1431
+ print("error: %s" % exc, file=sys.stderr)
1432
+ return 1
1433
+ for q in args.explain:
1434
+ try:
1435
+ _run_explain(q, engine)
1436
+ except DatalogError as exc:
1437
+ print("error: %s" % exc, file=sys.stderr)
1438
+ return 1
1439
+ return 0
1440
+
1441
+ preds = sorted(set(program.arity) if args.all else program.idb)
1442
+ for pred in preds:
1443
+ tuples = engine.rels.get(pred, set())
1444
+ kind = "derived" if pred in program.idb else "input"
1445
+ print("%% %s/%d (%s) — %d fact%s" %
1446
+ (pred, program.arity[pred], kind, len(tuples),
1447
+ "" if len(tuples) == 1 else "s"))
1448
+ for tup in sorted(tuples, key=_sort_key):
1449
+ print(format_fact(pred, tup))
1450
+ print()
1451
+ return 0
1452
+
1453
+
1454
+ if __name__ == "__main__":
1455
+ sys.exit(main())