#!/usr/bin/env python3
"""NULL-6: explicit finite SQL contract, Python 3 standard library only.

Run: python3 foundations-sql-null-checker.py
This is an executable teaching model, not a general SQL engine. Occurrence
indices are diagnostic identities, not assumed keys in the user tables.
"""
from collections import Counter
from enum import Enum
from itertools import product
import json
import sqlite3

class TV(Enum):
    F=0; U=1; T=2

def sql_eq(x,y):
    return TV.U if x is None or y is None else (TV.T if x==y else TV.F)
def sql_not(x):
    return {TV.T:TV.F, TV.F:TV.T, TV.U:TV.U}[x]
def sql_and(x,y):
    if TV.F in (x,y): return TV.F
    if TV.U in (x,y): return TV.U
    return TV.T
def sql_or(x,y):
    if TV.T in (x,y): return TV.T
    if TV.U in (x,y): return TV.U
    return TV.F
def identity(x,y):
    return (x is None and y is None) or (x is not None and y is not None and x==y)
def sequence_bags(alphabet,max_len):
    return [xs for n in range(max_len+1) for xs in product(alphabet,repeat=n)]

CUSTOMERS=[('A',1),('A',1),('B',2),('C',3),('D',None),('E',4)]
ORDERS=[(1,10,'ok'),(1,10,'ok'),(1,None,'ok'),(2,20,'hold'),
        (2,None,'ok'),(None,99,'ok'),(4,0,'ok'),(4,5,'ok')]
BLOCKED=[2,None]

class CapacityExceeded(Exception): pass
class NotReady(Exception): pass
class CardinalityError(Exception): pass

class ExactMembership:
    """Bounded exact chained table; force_collision separates proof from cost."""
    def __init__(self,capacity,force_collision=False):
        if type(capacity) is not int or capacity<0:
            raise ValueError('capacity must be a nonnegative integer; bool is not accepted')
        self.capacity=capacity; self.force_collision=force_collision
        self.buckets={}; self.size=0; self.empty=True; self.has_null=False
        self.phase='BUILD'; self.comparisons=0; self.reads=0
    def _bucket(self,x): return 0 if self.force_collision else hash(x)
    def _contains(self,x):
        for y in self.buckets.get(self._bucket(x),[]):
            self.comparisons+=1
            if x==y: return True
        return False
    def add(self,x):
        if self.phase!='BUILD': raise NotReady(self.phase)
        self.reads+=1; self.empty=False
        if x is None: self.has_null=True; return
        if self._contains(x): return
        if self.size==self.capacity:
            self.phase='FAILED'; raise CapacityExceeded(self.capacity)
        self.buckets.setdefault(self._bucket(x),[]).append(x); self.size+=1
    def finish(self):
        if self.phase!='BUILD': raise NotReady(self.phase)
        self.phase='READY'
    def build(self,xs):
        try:
            for x in xs: self.add(x)
            self.finish()
        except Exception:
            self.phase='FAILED'
            raise
        return self
    def _ready(self):
        if self.phase!='READY': raise NotReady(self.phase)
    def in_value(self,x):
        self._ready()
        if self.empty: return TV.F
        if x is None: return TV.U
        if self._contains(x): return TV.T
        return TV.U if self.has_null else TV.F
    def not_exists(self,x):
        self._ready()
        return not (x is not None and self._contains(x))

def reference_in(x,ys):
    result=TV.F
    for y in ys: result=sql_or(result,sql_eq(x,y))
    return result

def outer_loop(left,right,predicate):
    """(left index, right index or None, left value, right value or None)."""
    out=[]
    for i,r in enumerate(left):
        matched=False
        for j,s in enumerate(right):
            if predicate(r,s) is TV.T:
                matched=True; out.append((i,j,r,s))
        if not matched: out.append((i,None,r,None))
    return out

def outer_spec(left,right,predicate):
    # Independent occurrence-set definition, not the loop's matched state.
    pairs=[(i,j,r,s) for i,r in enumerate(left) for j,s in enumerate(right)
           if predicate(r,s) is TV.T]
    matched={i for i,j,r,s in pairs}
    return pairs+[(i,None,r,None) for i,r in enumerate(left) if i not in matched]

def finish_aggregate(st):
    n,c,total=st
    return (n,c,total if c else None)
def summary(orders,capacity=None):
    if capacity is not None and (type(capacity) is not int or capacity<0):
        raise ValueError('capacity must be None or a nonnegative integer; bool is not accepted')
    out={}; updates=0
    for k,amount,status in orders:
        if status!='ok' or k is None: continue
        if k not in out:
            if capacity is not None and len(out)>=capacity: raise CapacityExceeded(capacity)
            out[k]=[0,0,0]
        st=out[k]; st[0]+=1
        if amount is not None: st[1]+=1; st[2]+=amount
        updates+=1
    return {k:finish_aggregate(st) for k,st in out.items()},updates

def aggregate_reference(customers,orders):
    out=[]
    for label,k in customers:
        selected=[o[1] for o in orders if sql_eq(k,o[0]) is TV.T and o[2]=='ok']
        nonnull=[v for v in selected if v is not None]
        out.append((label,k,len(selected),len(nonnull),sum(nonnull) if nonnull else None))
    return out

def aggregate_plan(customers,orders):
    g,_=summary(orders)
    return [(label,k,*g.get(k,(0,0,None))) for label,k in customers]

def domain_plan(customers,orders):
    domain=[]
    for _,k in customers:
        if not any(identity(k,d) for d in domain): domain.append(k)
    calculated=[(d,aggregate_reference([('_',d)],orders)[0][2:]) for d in domain]
    return [(label,k,*v) for label,k in customers for d,v in calculated if identity(k,d)]

def composed_plan(customers,orders,blocked):
    h=ExactMembership(len(set(x for x in blocked if x is not None))).build(blocked)
    g,updates=summary(orders)
    result=[]; excluded=0; block_probes=0; group_probes=0
    for label,k in customers:
        if k is not None: block_probes+=1
        if not h.not_exists(k): excluded+=1; continue
        if k is not None: group_probes+=1
        result.append((label,k,*g.get(k,(0,0,None))))
    return result,dict(block_reads=len(blocked),order_reads=len(orders),
        customer_reads=len(customers),aggregate_updates=updates,
        aggregate_groups=len(g),block_members=h.size,block_probes=block_probes,
        group_probes=group_probes,excluded=excluded,output_occurrences=len(result))

def scalar_reference(customers,orders):
    out=[]
    for label,k in customers:
        values=[a for kk,a,status in orders if sql_eq(k,kk) is TV.T and status=='ok']
        if len(values)>1: raise CardinalityError(k)
        out.append((label,k,values[0] if values else None))
    return out

def scalar_plan(customers,orders):
    states={}
    # Store an error marker, but do not raise for unrequested parameter keys.
    for k,a,status in orders:
        if k is None or status!='ok': continue
        count,first=states.get(k,(0,None))
        states[k]=(min(2,count+1),a if count==0 else first)
    out=[]
    for label,k in customers:
        count,first=states.get(k,(0,None))
        if count==2: raise CardinalityError(k)
        out.append((label,k,first if count else None))
    return out

def observed(f,*args):
    try: return ('bag',Counter(f(*args)))
    except CardinalityError: return ('CardinalityError',)

def run():
    if not __debug__:
        raise RuntimeError('Run without -O: this teaching checker requires active assertions')
    for invalid in (True, False, -1, 1.5, '2'):
        try: ExactMembership(invalid); raise AssertionError('invalid membership capacity accepted')
        except ValueError: pass
        try: summary([],invalid); raise AssertionError('invalid aggregation capacity accepted')
        except ValueError: pass
    truth_checks=0
    for p,q in product(TV,repeat=2):
        assert (sql_and(p,q) is TV.T)==(p is TV.T and q is TV.T)
        assert sql_not(sql_and(p,q))==sql_or(sql_not(p),sql_not(q))
        assert sql_not(sql_or(p,q))==sql_and(sql_not(p),sql_not(q))
        truth_checks+=1
    assert sql_or(TV.U,sql_not(TV.U)) is TV.U
    # Different truth values can be WHERE equivalent, but not under NOT.
    assert (TV.U is TV.T)==(TV.F is TV.T)
    assert sql_not(TV.U)!=sql_not(TV.F)

    bags=sequence_bags((None,0,1),3)
    membership_cases=0
    for ys in bags:
        for collision in (False,True):
            h=ExactMembership(2,collision).build(ys)
            for x in (None,0,1):
                assert h.in_value(x)==reference_in(x,ys)
                assert h.not_exists(x)==not_any_true(x,ys)
                membership_cases+=1
    assert sql_not(ExactMembership(0).build([]).in_value(None)) is TV.T
    h=ExactMembership(1); h.add(2); h.add(None)
    try: h.not_exists(5); raise AssertionError('incomplete table accepted')
    except NotReady: pass
    try: h.add(5); raise AssertionError('capacity overflow accepted')
    except CapacityExceeded: pass
    try: h.not_exists(5); raise AssertionError('failed table accepted')
    except NotReady: pass

    outer_cases=0
    for left in bags:
        for right in bags:
            pred=sql_eq
            got=outer_loop(left,right,pred)
            assert Counter(got)==Counter(outer_spec(left,right,pred))
            # WHERE s=1 is null-rejecting, so outer-to-inner rewrite holds.
            filtered=[z for z in got if sql_eq(z[3],1) is TV.T]
            inner=[(i,j,r,s) for i,r in enumerate(left) for j,s in enumerate(right)
                   if sql_and(sql_eq(r,s),sql_eq(s,1)) is TV.T]
            assert Counter(filtered)==Counter(inner)
            outer_cases+=1

    by_key=lambda c,o:sql_eq(c[1],o[0])
    by_ok=lambda c,o:sql_and(by_key(c,o),sql_eq(o[2],'ok'))
    joins=outer_loop(CUSTOMERS,ORDERS,by_key)
    on_ok=outer_loop(CUSTOMERS,ORDERS,by_ok)
    where_ok=[z for z in joins if sql_eq(None if z[3] is None else z[3][2],'ok') is TV.T]
    assert [len(joins),len(on_ok),len(where_ok)]==[12,11,9]
    semi=[c for c in CUSTOMERS if any(by_ok(c,o) is TV.T for o in ORDERS)]
    anti=[c for c in CUSTOMERS if not any(by_ok(c,o) is TV.T for o in ORDERS)]
    assert Counter(semi+anti)==Counter(CUSTOMERS)
    assert [len(semi),len(anti)]==[4,2]
    kept=[c for c in CUSTOMERS if not_any_true(c[1],BLOCKED)]
    reference=aggregate_reference(kept,ORDERS)
    result,cost=composed_plan(CUSTOMERS,ORDERS,BLOCKED)
    expected=[('A',1,3,2,20),('A',1,3,2,20),('C',3,0,0,None),
              ('D',None,0,0,None),('E',4,2,2,5)]
    assert Counter(result)==Counter(reference)==Counter(expected)
    assert Counter(aggregate_plan(CUSTOMERS,ORDERS))==Counter(domain_plan(CUSTOMERS,ORDERS))
    assert not [c for c in CUSTOMERS if sql_not(reference_in(c[1],BLOCKED)) is TV.T]
    assert cost==dict(block_reads=2,order_reads=8,customer_reads=6,aggregate_updates=6,
        aggregate_groups=3,block_members=1,block_probes=5,group_probes=4,excluded=1,output_occurrences=5)

    aggregate_cases=0; scalar_cases=0
    variants=((None,5,'ok'),(0,None,'ok'),(0,0,'ok'),(0,2,'ok'),(1,3,'hold'),(1,4,'ok'))
    order_bags=sequence_bags(variants,3)
    customer_bags=[[(str(i%2),k) for i,k in enumerate(xs)] for xs in bags]
    for orders in order_bags:
        for customers in customer_bags:
            ref=Counter(aggregate_reference(customers,orders))
            assert ref==Counter(aggregate_plan(customers,orders))==Counter(domain_plan(customers,orders))
            aggregate_cases+=1
            assert observed(scalar_reference,customers,orders)==observed(scalar_plan,customers,orders)
            scalar_cases+=1
    def broken_input():
        yield 2
        raise OSError('injected input read failure before key 5')
    broken=ExactMembership(2)
    try: broken.build(broken_input()); raise AssertionError('read failure swallowed')
    except OSError: pass
    assert broken.phase=='FAILED'
    try: broken.finish(); raise AssertionError('failed read finalized')
    except NotReady: pass
    try: broken.not_exists(5); raise AssertionError('partial negative accepted')
    except NotReady: pass
    assert scalar_plan([],[(1,10,'ok'),(1,10,'ok')])==[]
    assert scalar_plan([('C',3)],[(1,10,'ok'),(1,10,'ok')])==[('C',3,None)]
    assert observed(scalar_plan,[('A',1)],[(1,10,'ok'),(1,10,'ok')])==('CardinalityError',)
    try: summary(ORDERS,2); raise AssertionError('aggregate overflow accepted')
    except CapacityExceeded: pass

    # Structural transfer: all free parameters, including threshold, must be kept.
    transfer_customers=[('A',1,5),('A',1,5),('B',1,15),('C',None,0)]
    transfer_orders=[(1,10),(1,20),(1,None)]
    def amount_gt(a,t):
        return TV.U if a is None or t is None else (TV.T if a>t else TV.F)
    def evaluate_parameter(k,threshold):
        xs=[a for kk,a in transfer_orders
            if sql_and(sql_eq(k,kk),amount_gt(a,threshold)) is TV.T]
        return (len(xs),len(xs),sum(xs) if xs else None)
    transfer_direct=[(label,k,t,*evaluate_parameter(k,t)) for label,k,t in transfer_customers]
    domain=list(dict.fromkeys((k,t) for label,k,t in transfer_customers))
    e=[(k,t,evaluate_parameter(k,t)) for k,t in domain]
    transfer_decorrelated=[(label,k,t,*v) for label,k,t in transfer_customers for kk,tt,v in e
                          if identity(k,kk) and identity(t,tt)]
    transfer_expected=[('A',1,5,2,2,30),('A',1,5,2,2,30),('B',1,15,1,1,20),('C',None,0,0,0,None)]
    assert Counter(transfer_direct)==Counter(transfer_decorrelated)==Counter(transfer_expected)
    assert len(transfer_customers)*len(transfer_orders)==12
    assert len(domain)*len(transfer_orders)==9
    assert evaluate_parameter(1,None)==(0,0,None)

    # Deliberately bad rewrite witnesses, asserted to differ rather than pass.
    wrong_not_in_customers=[c for c in CUSTOMERS if sql_not(reference_in(c[1],BLOCKED)) is TV.T]
    wrong_not_in_bag=aggregate_reference(wrong_not_in_customers,ORDERS)
    wrong_coalesce_sum=[(*row[:-1],0 if row[-1] is None else row[-1]) for row in result]
    mutants={
      'deduplicate_outer':Counter(aggregate_plan(list(dict.fromkeys(CUSTOMERS)),ORDERS))!=Counter(aggregate_reference(CUSTOMERS,ORDERS)),
      'not_in_for_not_exists':Counter(wrong_not_in_bag)!=Counter(result),
      'where_to_on':len(where_ok)!=len(on_ok),
      'nullable_payload_as_match_flag':any(z[2][0]=='B' and z[1] is not None and z[3][1] is None for z in on_ok),
      'sum_null_to_zero':Counter(wrong_coalesce_sum)!=Counter(result),
      'discard_scalar_duplicates':observed(scalar_plan,[('A',1)],[(1,10,'ok')])!=observed(scalar_plan,[('A',1)],[(1,10,'ok'),(1,10,'ok')])
    }
    assert all(mutants.values())

    # A second execution implementation checks supported SQL bag/NULL rules.
    # SQLite permits multiple scalar rows, so it is NOT a scalar-error oracle.
    con=sqlite3.connect(':memory:')
    con.executescript('CREATE TABLE Customer(label TEXT,k INTEGER); CREATE TABLE Orders(k INTEGER,amount INTEGER,status TEXT); CREATE TABLE Blocked(k INTEGER);')
    con.executemany('INSERT INTO Customer VALUES (?,?)',CUSTOMERS)
    con.executemany('INSERT INTO Orders VALUES (?,?,?)',ORDERS)
    con.executemany('INSERT INTO Blocked VALUES (?)',[(k,) for k in BLOCKED])
    sql='''SELECT c.label,c.k,
      (SELECT COUNT(*) FROM Orders o WHERE o.k=c.k AND o.status='ok'),
      (SELECT COUNT(o.amount) FROM Orders o WHERE o.k=c.k AND o.status='ok'),
      (SELECT SUM(o.amount) FROM Orders o WHERE o.k=c.k AND o.status='ok')
      FROM Customer c WHERE NOT EXISTS(SELECT 1 FROM Blocked b WHERE b.k=c.k)'''
    sqlite_result=con.execute(sql).fetchall()
    assert Counter(sqlite_result)==Counter(expected)
    rewritten='''SELECT c.label,c.k,COALESCE(g.n,0),COALESCE(g.c,0),g.s
      FROM Customer c LEFT JOIN (SELECT k,COUNT(*) n,COUNT(amount) c,SUM(amount) s
      FROM Orders WHERE status='ok' AND k IS NOT NULL GROUP BY k) g ON c.k=g.k
      WHERE NOT EXISTS(SELECT 1 FROM Blocked b WHERE b.k=c.k)'''
    assert Counter(con.execute(rewritten).fetchall())==Counter(expected)
    con.close()
    return dict(model='NULL-6, pure finite bag SQL subset',truth_table_pairs=truth_checks,
       membership_cases=membership_cases,outer_occurrence_cases=outer_cases,
       aggregate_cases=aggregate_cases,scalar_success_or_error_cases=scalar_cases,
       fixed_output=expected,cost=cost,mutant_witnesses=mutants,
       input_read_failure='FAILED; finish and probe remain rejected',
       structural_transfer=dict(parameters=['k','threshold'],domain_size=len(domain),direct_candidate_checks=12,domain_candidate_checks=9,null_threshold_result=evaluate_parameter(1,None),output=transfer_expected),
       sql_engine=dict(name='SQLite',version=sqlite3.sqlite_version,queries_checked=2,
         limitation='No scalar cardinality-error cross-check: SQLite has different behavior'),
       proof_scope='Bounded exhaustive tests check implementation; page invariants prove stated general rules')

def not_any_true(x,ys):
    return not any(sql_eq(x,y) is TV.T for y in ys)

if __name__=='__main__':
    print(json.dumps(run(),ensure_ascii=False,indent=2))
