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.
- tiny_datalog/__init__.py +42 -0
- tiny_datalog/containment.py +228 -0
- tiny_datalog/datalog.py +1455 -0
- tiny_datalog/incremental.py +424 -0
- tiny_datalog/magic.py +175 -0
- tiny_datalog/prolog.py +335 -0
- tiny_datalog/semantics.py +182 -0
- tiny_datalog/semiring.py +383 -0
- tiny_datalog/subsumption.py +475 -0
- tiny_datalog/tabling.py +220 -0
- tiny_datalog-0.1.0.dist-info/METADATA +360 -0
- tiny_datalog-0.1.0.dist-info/RECORD +16 -0
- tiny_datalog-0.1.0.dist-info/WHEEL +5 -0
- tiny_datalog-0.1.0.dist-info/entry_points.txt +8 -0
- tiny_datalog-0.1.0.dist-info/licenses/LICENSE +21 -0
- tiny_datalog-0.1.0.dist-info/top_level.txt +1 -0
tiny_datalog/datalog.py
ADDED
|
@@ -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())
|