#!/usr/bin/env python3
"""Finite positive Datalog: query-scoped magic rewriting and delete/rederive.

Variables start with '?'; other strings are constants. A fact/atom is
(predicate, tuple(terms)); a rule is (head, tuple(body_atoms)).
Standard library only; stdout only. Explicit checks survive python -O.
"""
from collections import deque
from itertools import product
import json
import random
import re


def require(ok, message):
    if not ok:
        raise ValueError(message)


def variable(t):
    return type(t) is str and t.startswith('?')


def atom(name, *terms):
    return name, tuple(terms)


def variables(a):
    return {t for t in a[1] if variable(t)}


def validate(rules, edb_schema):
    require(type(rules) in (tuple, list) and type(edb_schema) is dict, 'bad program')
    require(all(type(p) is str and p and type(a) is int and a >= 0 for p, a in edb_schema.items()), 'bad EDB schema')
    arities, heads = dict(edb_schema), set()
    for rule in rules:
        require(type(rule) is tuple and len(rule) == 2 and type(rule[1]) is tuple, 'bad rule')
        for a in (rule[0],) + rule[1]:
            require(type(a) is tuple and len(a) == 2 and type(a[0]) is str and a[0] and type(a[1]) is tuple, 'bad atom')
            require(all(type(t) is str and (not variable(t) or len(t) > 1) for t in a[1]), 'bad term')
            require(a[0] not in arities or arities[a[0]] == len(a[1]), 'arity mismatch')
            arities[a[0]] = len(a[1])
        head, body = rule
        require(head[0] not in edb_schema, 'EDB cannot be a rule head')
        require(variables(head) <= set().union(*(variables(a) for a in body)), 'unsafe head variable')
        heads.add(head[0])
    for _, body in rules:
        require(all(a[0] in heads or a[0] in edb_schema for a in body), 'undefined body predicate')
    return arities, heads


def check_facts(facts, schema):
    require(type(facts) in (set, frozenset), 'facts must have set semantics')
    for fact in facts:
        require(type(fact) is tuple and len(fact) == 2, 'bad fact')
        p, row = fact
        require(p in schema and type(row) is tuple and len(row) == schema[p], 'fact schema mismatch')
        require(all(type(v) is str and not variable(v) for v in row), 'nonconstant fact value')


def instantiate(a, env):
    return a[0], tuple(env[t] if variable(t) else t for t in a[1])


def matches(rule, known, stats=None):
    """Left-to-right relational matching, counting tested relation rows.

    Index construction and sorting are additional work, not row_tests.
    """
    index = {}
    for p, row in known:
        index.setdefault(p, []).append(row)
    for rows in index.values():
        rows.sort()
    head, body = rule
    current = [({}, ())]
    for a in body:
        following = []
        for env, support in current:
            for row in index.get(a[0], ()):
                if stats is not None:
                    stats['row_tests'] += 1
                new, okay = dict(env), True
                for term, value in zip(a[1], row):
                    if variable(term):
                        if term in new and new[term] != value:
                            okay = False
                            break
                        new[term] = value
                    elif term != value:
                        okay = False
                        break
                if okay:
                    following.append((new, support + ((a[0], row),)))
        current = following
    return [(instantiate(head, env), support) for env, support in current]


def materialize(rules, edb_schema, base):
    validate(rules, edb_schema)
    check_facts(base, edb_schema)
    known, trace = set(base), []
    stats = {'row_tests': 0, 'successful_instances': 0}
    while True:
        new = set()
        for rule in rules:
            instances = matches(rule, known, stats)
            stats['successful_instances'] += len(instances)
            new.update(h for h, _ in instances if h not in known)
        trace.append(set(new))
        if not new:
            break
        known.update(new)
    return {'idb': known - base, 'trace': trace, 'stats': stats}


def decorated(kind, p, pattern):
    return '@' + kind + ':' + p + ':' + pattern


def magic_rewrite(rules, edb_schema, predicate, pattern, values):
    arities, heads = validate(rules, edb_schema)
    require(all(re.fullmatch(r'[A-Za-z][A-Za-z0-9_]*', p) for p in arities), 'source names reserve @ for generated symbols')
    require(predicate in heads and type(pattern) is str and len(pattern) == arities[predicate] and set(pattern) <= {'b', 'f'}, 'bad query pattern')
    require(type(values) is tuple and len(values) == pattern.count('b') and all(type(v) is str and not variable(v) for v in values), 'bad query bindings')
    generated, seen_rules, pending, visited = [], set(), deque([(predicate, pattern)]), set()
    adornment_trace = []
    def add(rule):
        if rule not in seen_rules:
            seen_rules.add(rule)
            generated.append(rule)
    add((atom(decorated('m', predicate, pattern), *values), ()))
    while pending:
        p, pat = pending.popleft()
        if (p, pat) in visited:
            continue
        visited.add((p, pat))
        for source_index, (head, body) in enumerate(rules):
            if head[0] != p:
                continue
            selected = tuple(t for t, flag in zip(head[1], pat) if flag == 'b')
            guard = atom(decorated('m', p, pat), *selected)
            bound = {t for t in selected if variable(t)}
            prefix, rewritten, calls = [guard], [], []
            for a in body:
                if a[0] in heads:
                    call_pattern = ''.join('b' if not variable(t) or t in bound else 'f' for t in a[1])
                    call_args = tuple(t for t, flag in zip(a[1], call_pattern) if flag == 'b')
                    magic_head = atom(decorated('m', a[0], call_pattern), *call_args)
                    if tuple(prefix) != (magic_head,):
                        add((magic_head, tuple(prefix)))
                    actual = atom(decorated('a', a[0], call_pattern), *a[1])
                    pending.append((a[0], call_pattern))
                    calls.append({'predicate': a[0], 'pattern': call_pattern, 'bound_before': sorted(bound), 'prefix': tuple(prefix)})
                else:
                    actual = a
                rewritten.append(actual)
                prefix.append(actual)
                bound.update(variables(a))
            add((atom(decorated('a', p, pat), *head[1]), (guard,) + tuple(rewritten)))
            adornment_trace.append({'source_rule': source_index, 'head_pattern': pat, 'calls': calls})
    validate(generated, edb_schema)
    return {'rules': tuple(generated), 'predicate': predicate, 'pattern': pattern, 'values': values,
            'answer_predicate': decorated('a', predicate, pattern), 'adornments': sorted(visited), 'adornment_trace': adornment_trace}


def select_answers(facts, predicate, pattern, values):
    positions = [i for i, flag in enumerate(pattern) if flag == 'b']
    return {row for p, row in facts if p == predicate and tuple(row[i] for i in positions) == values}


def dred(rules, edb_schema, old_base, old_idb, deletions, insertions=frozenset()):
    """Requires an exact old least materialization under fixed rules.

    No full recomputation oracle is invoked inside this update algorithm.
    The caller tests against a fresh oracle outside it.
    """
    arities, heads = validate(rules, edb_schema)
    for facts in (old_base, deletions, insertions):
        check_facts(facts, edb_schema)
    check_facts(old_idb, {p: arities[p] for p in heads})
    new_base = (old_base - deletions) | insertions
    minus, plus = old_base - new_base, new_base - old_base
    stats = {'row_tests': 0, 'old_instances': 0, 'rederive_body_tests': 0, 'insert_instances': 0}
    if not minus and not plus:
        return {'base': set(new_base), 'idb': set(old_idb), 'actual_delete': set(), 'actual_insert': set(),
                'overdeleted': set(), 'rederived': set(), 'inserted_idb': set(),
                'overdelete_trace': [], 'rederive_trace': [], 'insert_trace': [], 'stats': stats}
    old = old_base | old_idb
    instances = [instance for rule in rules for instance in matches(rule, old, stats)]
    stats['old_instances'] = len(instances)
    watchers = {}
    for head, body in instances:
        for fact in body:
            watchers.setdefault(fact, []).append(head)
    deleted, frontier, delete_trace = set(minus), set(minus), []
    while frontier:
        delete_trace.append(set(frontier))
        following = {head for fact in frontier for head in watchers.get(fact, ())} - deleted
        deleted.update(following)
        frontier = following
    overdeleted = deleted & old_idb
    current = old - deleted
    rederived, rederive_trace = set(), []
    while True:
        new = set()
        for head, body in instances:
            if head not in overdeleted or head in current:
                continue
            okay = True
            for fact in body:
                stats['rederive_body_tests'] += 1
                if fact not in current:
                    okay = False
                    break
            if okay:
                new.add(head)
        rederive_trace.append(set(new))
        if not new:
            break
        current.update(new)
        rederived.update(new)
    # Existing Datalog delta principle: at least one witness fact is new.
    current.update(plus)
    frontier, insert_trace, inserted_idb = set(plus), [], set()
    while frontier:
        new = set()
        changed_predicates = {p for p, _ in frontier}
        for rule in rules:
            if not any(a[0] in changed_predicates for a in rule[1]):
                continue
            for head, body in matches(rule, current, stats):
                if any(f in frontier for f in body):
                    stats['insert_instances'] += 1
                    if head not in current:
                        new.add(head)
        insert_trace.append(set(new))
        current.update(new)
        inserted_idb.update(new)
        frontier = new
    return {'base': set(new_base), 'idb': {f for f in current if f[0] in heads},
            'actual_delete': set(minus), 'actual_insert': set(plus), 'overdeleted': overdeleted,
            'rederived': rederived, 'inserted_idb': inserted_idb,
            'overdelete_trace': delete_trace, 'rederive_trace': rederive_trace, 'insert_trace': insert_trace, 'stats': stats}


def unsupported_peeling_fault(rules, base, old_idb):
    """Wrong deletion: a surviving local support can be circular, with no proof."""
    current = base | old_idb
    while True:
        supported = {h for rule in rules for h, _ in matches(rule, current)}
        following = base | (old_idb & supported)
        if following == current:
            return following - base
        current = following


def ground_oracle(rules, schema, base):
    """Separate finite-domain enumeration, used only for validation."""
    values = {v for _, row in base for v in row}
    values.update(t for h, body in rules for a in (h,) + body for t in a[1] if not variable(t))
    known, ground = set(base), []
    for head, body in rules:
        names = sorted(set().union(*(variables(a) for a in (head,) + body)))
        for vals in product(sorted(values), repeat=len(names)):
            env = dict(zip(names, vals))
            ground.append((instantiate(head, env), tuple(instantiate(a, env) for a in body)))
    while True:
        new = {h for h, body in ground if all(a in known for a in body)} - known
        if not new:
            return known - base
        known.update(new)


def reach_program():
    return ((atom('R', '?X', '?Y'), (atom('E', '?X', '?Y'),)),
            (atom('R', '?X', '?Z'), (atom('R', '?X', '?Y'), atom('R', '?Y', '?Z'))))


def show_atom(a):
    return a[0] + '(' + ','.join(a[1]) + ')'


def show_rule(rule):
    h, body = rule
    return show_atom(h) + ' <- ' + ', '.join(show_atom(a) for a in body)


def encode(value):
    if type(value) in (set, frozenset):
        return [encode(v) for v in sorted(value)]
    if type(value) is dict:
        return {k: encode(v) for k, v in value.items()}
    if type(value) in (tuple, list):
        return [encode(v) for v in value]
    return value


def self_test():
    rules, schema = reach_program(), {'E': 2}
    base = {atom('E', *edge) for edge in [('a','b'),('b','c'),('c','b'),('c','d'),('a','e'),('e','d'),('u','v'),('v','w'),('w','u')]}
    full = materialize(rules, schema, base)
    magic = magic_rewrite(rules, schema, 'R', 'bf', ('a',))
    partial = materialize(magic['rules'], schema, base)
    answers = select_answers(partial['idb'], magic['answer_predicate'], 'bf', ('a',))
    require(answers == {('a',y) for y in 'bcde'}, 'initial query')
    require(len(full['idb']) == 20 and len(partial['idb']) == 16, 'original/magic counts')
    delete = {atom('E','a','b')}
    updated = dred(rules, schema, base, full['idb'], delete)
    updated_magic = dred(magic['rules'], schema, base, partial['idb'], delete)
    require(updated['idb'] == ground_oracle(rules, schema, base-delete), 'DRed ordinary')
    require(updated_magic['idb'] == ground_oracle(magic['rules'], schema, base-delete), 'DRed transformed')
    require(updated['overdeleted'] == {atom('R','a',y) for y in 'bcd'} and updated['rederived'] == {atom('R','a','d')}, 'ordinary intermediate sets')
    new_answers = select_answers(updated_magic['idb'], magic['answer_predicate'], 'bf', ('a',))
    require(new_answers == {('a','d'),('a','e')}, 'new query')
    require(len(updated_magic['overdeleted']) == 12 and len(updated_magic['rederived']) == 2 and len(updated_magic['idb']) == 6, 'magic intermediate counts')
    fault = unsupported_peeling_fault(rules, base-delete, full['idb'])
    require(atom('R','a','b') in fault and atom('R','a','c') in fault and fault != updated['idb'], 'circular self-support fault')
    restored = dred(magic['rules'], schema, base-delete, updated_magic['idb'], set(), delete)
    require(restored['idb'] == partial['idb'], 'insertion after deletion')
    cases = [
        (((atom('Ready'), ()),), {}, set(), 'Ready', '', ()),
        (((atom('P','?X','?X'), (atom('E','?X','?Y'),)),), {'E':2}, {atom('E','a','b')}, 'P','bf',('a',)),
        (((atom('P','a'), ()),), {}, set(), 'P','b',('z',)),
        (((atom('P','?X','?Z'), (atom('E','?X','?Y'),atom('Q','?Z','?Y'))),
          (atom('Q','?U','?V'), (atom('F','?U','?V'),))), {'E':2,'F':2},
          {atom('E','a','b'),atom('F','c','b'),atom('F','d','z')}, 'P','bf',('a',)),
    ]
    for program, edb, inp, pred, pat, vals in cases:
        transformed = magic_rewrite(program, edb, pred, pat, vals)
        got = materialize(transformed['rules'], edb, inp)['idb']
        expected = ground_oracle(program, edb, inp)
        require(select_answers(got, transformed['answer_predicate'], pat, vals) == select_answers(expected, pred, pat, vals), 'binding boundary')
    rng = random.Random(20020)
    query_checks, updates = 0, 0
    for n in range(6):
        names = [str(i) for i in range(n)]
        for case in range(30):
            inp = {atom('E',x,y) for x in names for y in names if rng.random()<.25}
            original = ground_oracle(rules, schema, inp)
            for pat in ['bf','fb','bb','ff']:
                vals = tuple(rng.choice(names+['outside']) for _ in range(pat.count('b')))
                transformed = magic_rewrite(rules, schema, 'R', pat, vals)
                actual = materialize(transformed['rules'], schema, inp)['idb']
                require(select_answers(actual, transformed['answer_predicate'], pat, vals) == select_answers(original, 'R', pat, vals), 'random magic query')
                query_checks += 1
            transformed = magic_rewrite(rules, schema, 'R', 'bf', (names[0] if names else 'outside',))
            state = materialize(transformed['rules'], schema, inp)['idb']
            for step in range(5):
                deletes = {f for f in inp if rng.random()<.3}
                adds = {atom('E',x,y) for x in names for y in names if rng.random()<.15}
                result = dred(transformed['rules'], schema, inp, state, deletes, adds)
                inp = (inp-deletes)|adds
                expected = ground_oracle(transformed['rules'], schema, inp)
                require(result['idb']==expected, 'incremental full-recomputation oracle')
                source = ground_oracle(rules, schema, inp)
                require(select_answers(result['idb'], transformed['answer_predicate'], 'bf', transformed['values']) == select_answers(source, 'R', 'bf', transformed['values']), 'updated source-query equality')
                state = result['idb']
                updates += 1
    extra = {atom('E','d','fresh'), atom('E','fresh','a')}
    replaced = dred(rules, schema, base, full['idb'], delete, extra)
    require(replaced['idb'] == ground_oracle(rules,schema,(base-delete)|extra), 'new active-domain values')
    noop = dred(rules,schema,base,full['idb'],delete,delete)
    require(noop['idb']==full['idb'] and noop['stats']['old_instances']==0, 'cancelled update')
    return {'status':'PASS','source_rules':[show_rule(r) for r in rules], 'base':base,
            'magic_rules':[show_rule(r) for r in magic['rules']], 'adornment_trace':magic['adornment_trace'],
            'original':full,'magic':partial,'query_before':answers,
            'ordinary_delete':updated,'magic_delete':updated_magic,'query_after':new_answers,
            'wrong_peeling_extra':fault-updated['idb'],'reinsertion':restored,
            'new_values_batch':replaced,'cancelled_update':noop,
            'regressions':{'binding_boundaries':len(cases),'query_comparisons':query_checks,'incremental_batches':updates}}


if __name__ == '__main__':
    print(json.dumps(encode(self_test()), ensure_ascii=False, indent=2, sort_keys=True))
