#!/usr/bin/env python3
"""Bounded proper-clause SAT stream evidence; no limiting-threshold estimator."""
import argparse
import hashlib
import json
import math
import platform
import random
import time
from fractions import Fraction
from pathlib import Path

LAW = 'iid_proper_with_replacement'
MAX_BYTES = 1048576
MAX_CLAUSES = 2048
MAX_WORK = 5000000


def integer(value, low, high, name):
    if type(value) is not int or not low <= value <= high:
        raise ValueError(f'{name} must be an integer in [{low},{high}]')
    return value


def validate(data):
    if type(data) is not dict or set(data) != {'schema', 'n', 'k', 'law', 'clauses'}:
        raise ValueError('expected exactly schema, n, k, law, clauses')
    if data['schema'] != 'random-sat-stream-v1' or data['law'] != LAW:
        raise ValueError('unsupported schema or declared clause law')
    n = integer(data['n'], 3, 128, 'n')
    k = integer(data['k'], 3, min(8, n), 'k')
    if type(data['clauses']) is not list or len(data['clauses']) > MAX_CLAUSES:
        raise ValueError('clauses must be a list of at most 2048 clauses')
    for clause in data['clauses']:
        if type(clause) is not list or len(clause) != k:
            raise ValueError('each clause must contain exactly k literals')
        for literal in clause:
            integer(literal, -n, n, 'literal')
            if literal == 0:
                raise ValueError('literal zero is invalid')
        variables = [abs(x) for x in clause]
        if variables != sorted(set(variables)):
            raise ValueError('clause variables must be distinct and in increasing order')
    return data


def no_duplicates(pairs):
    result = {}
    for key, value in pairs:
        if key in result:
            raise ValueError('duplicate JSON key')
        result[key] = value
    return result


def load(raw):
    if len(raw) > MAX_BYTES:
        raise ValueError('input exceeds one MiB')
    def bad_constant(value):
        raise ValueError('nonfinite JSON is invalid')
    return validate(json.loads(raw, object_pairs_hook=no_duplicates,
                               parse_constant=bad_constant))


def clause_from_rank(n, k, rank):
    integer(n, 3, 128, 'n'); integer(k, 3, min(8, n), 'k')
    integer(rank, 0, math.comb(n, k) * 2**k - 1, 'rank')
    combination, signs = divmod(rank, 2**k)
    variables = []
    start = 1
    for remaining in range(k, 0, -1):
        for variable in range(start, n - remaining + 2):
            count = math.comb(n - variable, remaining - 1)
            if combination < count:
                variables.append(variable)
                start = variable + 1
                break
            combination -= count
    return [v if signs & (1 << i) else -v for i, v in enumerate(variables)]


def generate(n, k, count, seed):
    integer(n, 3, 128, 'n'); integer(k, 3, min(8, n), 'k')
    integer(count, 0, MAX_CLAUSES, 'count'); integer(seed, 0, 2**63 - 1, 'seed')
    rng = random.Random(seed)
    size = math.comb(n, k) * 2**k
    return dict(schema='random-sat-stream-v1', n=n, k=k, law=LAW,
                clauses=[clause_from_rank(n, k, rng.randrange(size)) for _ in range(count)])


def satisfies(clauses, assignment):
    return all(any(assignment[abs(x)-1] == (1 if x > 0 else 0) for x in c) for c in clauses)


def exact_observations(data, work_limit=MAX_WORK):
    validate(data); integer(work_limit, 0, MAX_WORK, 'work limit')
    n, k, clauses = data['n'], data['k'], data['clauses']
    records = [dict(prefix=0, status='SAT', basis='checked_assignment', assignment=[0]*n)]
    if n > 16:
        return records, dict(stop='exact_variable_limit', work=0)
    assignments = 2**n
    words = (assignments + 63)//64
    work = words*n
    if work > work_limit:
        return records, dict(stop='work_limit', work=0)
    full = (1 << assignments)-1
    positive = [(full//((1 << (1 << i))+1)) << (1 << i) for i in range(n)]
    alive = full
    for m, clause in enumerate(clauses, 1):
        charge = words*(k+1)
        if work+charge > work_limit:
            records.append(dict(prefix=m, status='UNKNOWN', basis='work_limit'))
            return records, dict(stop='work_limit', work=work)
        work += charge
        mask = 0
        for literal in clause:
            bitset = positive[abs(literal)-1]
            mask |= bitset if literal > 0 else full ^ bitset
        alive &= mask
        if not alive:
            records.append(dict(prefix=m, status='UNSAT', basis='exhaustive_assignment_bitset'))
            return records, dict(stop='first_unsat', work=work)
        witness = (alive & -alive).bit_length()-1
        assignment = [(witness >> i) & 1 for i in range(n)]
        records.append(dict(prefix=m, status='SAT', basis='checked_assignment', assignment=assignment))
    return records, dict(stop='stream_end', work=work)


def solver_observations(data, seconds=0.05, total_seconds=10.0):
    """Existing CP-SAT; SAT witnesses checked, UNSAT statuses not proof-checked."""
    validate(data)
    if type(seconds) not in (int, float) or not math.isfinite(seconds) or not 0 < seconds <= 1:
        raise ValueError('solver seconds must be positive and at most one')
    if type(total_seconds) not in (int, float) or not math.isfinite(total_seconds) or not 0 < total_seconds <= 30:
        raise ValueError('total solver seconds must be positive and at most thirty')
    import ortools
    from ortools.sat.python import cp_model
    n, clauses = data['n'], data['clauses']
    model = cp_model.CpModel()
    variables = [model.new_bool_var(f'x{i+1}') for i in range(n)]
    assignment = [0]*n
    records = [dict(prefix=0, status='SAT', basis='checked_assignment', assignment=assignment)]
    start = time.monotonic()
    calls = 0
    stop = 'stream_end'
    for m, clause in enumerate(clauses, 1):
        if time.monotonic()-start >= total_seconds:
            records.append(dict(prefix=m, status='UNKNOWN', basis='between_call_time_limit'))
            stop = 'time_limit'; break
        model.add_bool_or([variables[abs(x)-1] if x > 0 else ~variables[abs(x)-1] for x in clause])
        if satisfies([clause], assignment):
            records.append(dict(prefix=m, status='SAT', basis='checked_assignment_reused', assignment=assignment))
            continue
        solver = cp_model.CpSolver()
        solver.parameters.max_time_in_seconds = min(seconds, max(1e-9, total_seconds-(time.monotonic()-start)))
        solver.parameters.num_search_workers = 1
        solver.parameters.random_seed = 0
        status = solver.solve(model); calls += 1
        record = dict(prefix=m, solver_status=solver.status_name(status), wall_seconds=solver.wall_time)
        if status in (cp_model.OPTIMAL, cp_model.FEASIBLE):
            assignment = [int(solver.value(v)) for v in variables]
            if not satisfies(clauses[:m], assignment):
                raise ValueError('solver returned an invalid SAT witness')
            record.update(status='SAT', basis='checked_assignment', assignment=assignment)
        elif status == cp_model.INFEASIBLE:
            record.update(status='UNSAT', basis='solver_status_only')
            stop = 'solver_unsat'
        else:
            record.update(status='UNKNOWN', basis='solver_status_only')
            stop = 'solver_unknown'
        records.append(record)
        if record['status'] != 'SAT':
            break
    return records, dict(stop=stop, solver='OR-Tools CP-SAT', version=ortools.__version__,
                         calls=calls, per_call_seconds=seconds, total_seconds=total_seconds,
                         elapsed_seconds=time.monotonic()-start, independent_unsat_proof_checked=False,
                         note='Model retained, but each solve is a fresh CP-SAT call; no retained learned state. Time limits are not a hard process watchdog.')


def summarize(data, records):
    """Internal observations only; supplied statuses are not an evidence API."""
    m = len(data['clauses'])
    lower = max(r['prefix'] for r in records if r['status'] == 'SAT')+1
    verified = [r['prefix'] for r in records if r['status'] == 'UNSAT' and r['basis'] == 'exhaustive_assignment_bitset']
    claimed = [r['prefix'] for r in records if r['status'] == 'UNSAT']
    upper = min(verified) if verified else None
    solver_upper = min(claimed) if claimed else None
    exact_h = upper if upper == lower else None
    clipped = exact_h if exact_h is not None and exact_h < m else (m if lower >= m else None)
    return dict(independent_H_lower=lower, independent_H_upper=upper,
                solver_claimed_H_upper=solver_upper, independently_exact_H=exact_h,
                independently_exact_last_sat=exact_h-1 if exact_h is not None else None,
                independently_exact_clipped_H=clipped, cap=m,
                observed_survival_through_stream=(lower > m),
                right_censored_after_stream=(lower > m and upper is None),
                unknown_prefixes=[r['prefix'] for r in records if r['status'] == 'UNKNOWN'],
                first_failure_determined=exact_h is not None)


def coupon_control(cap):
    """Exact n=k=3 law via occupancy counts; an independent finite control."""
    integer(cap, 0, 256, 'coupon cap')
    ways = [1]+[0]*8
    survival = []
    for m in range(cap+1):
        survival.append(Fraction(sum(ways[:8]), 8**m))
        updated = [0]*9
        for occupied, count in enumerate(ways):
            updated[occupied] += count*occupied
            if occupied < 8:
                updated[occupied+1] += count*(8-occupied)
        ways = updated
    mean = sum(survival[:cap], Fraction())
    second = sum(((2*m+1)*survival[m] for m in range(cap)), Fraction())
    uncapped = 8*sum((Fraction(1, i) for i in range(1, 9)), Fraction())
    variance = 64*sum((Fraction(1, i*i) for i in range(1, 9)), Fraction())-uncapped
    return dict(n=3, k=3, cap=cap, survival=[str(p) for p in survival],
                clipped_mean=str(mean), clipped_variance=str(second-mean*mean),
                uncapped_mean=str(uncapped), uncapped_variance=str(variance),
                no_replacement_H=8, no_replacement_variance=0)


def audit(data, backend='exact', work_limit=MAX_WORK, seconds=0.05, total_seconds=10):
    validate(data)
    if backend == 'exact':
        observations, run = exact_observations(data, work_limit)
    elif backend == 'ortools':
        observations, run = solver_observations(data, seconds, total_seconds)
    else:
        raise ValueError('backend must be exact or ortools')
    n, k, clauses = data['n'], data['k'], data['clauses']
    bound = min(Fraction(1), 2**n * Fraction(2**k-1, 2**k)**len(clauses))
    return dict(schema='random-sat-audit-v1', backend=backend, n=n, k=k,
                declared_law=LAW, actual_randomness_or_iid_verified=False,
                source_proof_verified=False, asymptotic_threshold_estimated=False,
                duplicate_clause_count=len(clauses)-len(set(map(tuple, clauses))),
                clause_type_count=math.comb(n, k)*2**k,
                conditional_first_moment_survival_upper=str(bound),
                summary=summarize(data, observations), observations=observations, run=run,
                limits=dict(max_bytes=MAX_BYTES, max_clauses=MAX_CLAUSES, max_n=128,
                            max_exact_n=16, max_work=MAX_WORK),
                python=platform.python_version())


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', required=True, type=Path)
    parser.add_argument('--backend', choices=['exact', 'ortools'], default='exact')
    parser.add_argument('--work-limit', type=int, default=MAX_WORK)
    parser.add_argument('--seconds', type=float, default=0.05)
    parser.add_argument('--total-seconds', type=float, default=10)
    args = parser.parse_args()
    if args.output.exists():
        raise SystemExit('Refusing to replace an existing output directory')
    with args.input.open('rb') as f:
        raw = f.read(MAX_BYTES+1)
    report = audit(load(raw), args.backend, args.work_limit, args.seconds, args.total_seconds)
    report['input_sha256'] = hashlib.sha256(raw).hexdigest()
    report['checker_sha256'] = hashlib.sha256(Path(__file__).read_bytes()).hexdigest()
    args.output.mkdir(parents=True, exist_ok=False)
    with (args.output/'input.json').open('xb') as f:
        f.write(raw)
    with (args.output/'report.json').open('x') as f:
        json.dump(report, f, indent=2); f.write('\n')
    print(json.dumps(report['summary']))


if __name__ == '__main__':
    main()
