"""Independent saved-artifact audit: imports no simulator or policy functions."""
import json
import math
from pathlib import Path

ROOT = Path(__file__).resolve().parent / 'results'
D = ((0, -1), (1, 0), (0, 1), (-1, 0))


def close(a, b):
    assert math.isclose(a, b, abs_tol=1e-10), (a, b)


def distance(rows, source, target):
    # Repeated relaxation, independently of generator's BFS route implementation.
    cells = {(x, y) for y, row in enumerate(rows) for x, v in enumerate(row) if v == '.'}
    values = {p: 10000 for p in cells}
    values[source] = 0
    for _ in range(len(cells)):
        old = values.copy()
        for x, y in cells:
            values[(x, y)] = min([old[(x, y)]] +
                [old[(x+dx, y+dy)]+1 for dx, dy in D if (x+dx, y+dy) in cells])
        if values == old:
            break
    return values[target]


def check_obs(obs, rows, pos, goal):
    assert set(obs) == {'position', 'goal', 'patch'}
    assert obs['position'] == list(pos) and obs['goal'] == list(goal)
    assert len(obs['patch']) == 3 and all(len(r) == 3 for r in obs['patch'])
    for iy in range(3):
        for ix in range(3):
            p = (pos[0]+ix-1, pos[1]+iy-1)
            expected = 1 if rows[p[1]][p[0]] == '#' else 2 if p == goal else 0
            assert obs['patch'][iy][ix] == expected


def main():
    data = json.loads((ROOT/'maps.json').read_text())
    specs = {s['seed']: s for s in data['test']}
    hashes = [s['map_sha256'] for s in data['dev']+data['test']]
    assert len(set(hashes)) == len(hashes) == 35
    runs = [json.loads(line) for line in (ROOT/'episodes.jsonl').read_text().splitlines()]
    keys = [(r['seed'], r['policy'], r['horizon']) for r in runs]
    assert len(set(keys)) == len(keys) == 360
    rebuilt, steps, loops = {}, 0, []
    for r in runs:
        s = specs[r['seed']]
        pos, goal, rows = tuple(s['start']), tuple(s['goal']), s['rows']
        shortest = distance(rows, pos, goal)
        assert shortest == r['shortest'] and r['map_sha256'] == s['map_sha256']
        rewards, moves, collisions = [], 0, 0
        visited = {pos}
        revisits = 0
        for t, step in enumerate(r['trace'], 1):
            assert step['t'] == t
            check_obs(step['observation'], rows, pos, goal)
            a = step['action']
            assert type(a) is int and a in range(4)
            q = tuple(pos[k] + D[a][k] for k in range(2))
            collision = rows[q[1]][q[0]] == '#'
            if not collision:
                pos = q
                moves += 1
                revisits += pos in visited
                visited.add(pos)
            collisions += collision
            term, trunc = pos == goal, t >= r['horizon']
            reward = float(term) - .01 - .05 * collision
            assert step['collision'] == collision
            assert step['terminated'] == term and step['truncated'] == trunc
            if t < len(r['trace']):
                assert not (term or trunc)
            else:
                assert term or trunc
            close(step['reward'], reward)
            rewards.append(reward)
            check_obs(step['next_observation'], rows, pos, goal)
        steps += t
        success = pos == goal
        vals = dict(success=float(success), steps=t, moves=moves,
                    collision_rate=collisions/t, return_=sum(rewards),
                    spl=success*shortest/max(shortest, moves),
                    action_efficiency=success*shortest/max(shortest, t))
        assert collisions == r['collisions']
        for k, v in vals.items():
            close(v, r[k])
        rebuilt.setdefault((r['horizon'], r['policy']), []).append(vals)
        if r['horizon'] == 32 and r['policy'] in ('memory_reset', 'memory'):
            loops.append(dict(seed=r['seed'], policy=r['policy'], success=success,
                              steps=t, shortest=shortest, revisits=revisits,
                              final_position=pos))
    summary = json.loads((ROOT/'summary.json').read_text())
    for (h, p), rows in rebuilt.items():
        for k in rows[0]:
            close(sum(r[k] for r in rows)/len(rows), summary[str(h)][p][k])
    indexed = {(r['seed'], r['policy'], r['horizon']): r for r in runs}
    for seed in specs:
        for p in ('random', 'memory_reset', 'memory', 'oracle'):
            for small, large in ((16,32),(32,64)):
                a, b = indexed[(seed,p,small)]['trace'], indexed[(seed,p,large)]['trace']
                for x, y in zip(a, b):
                    # Only budget flag can differ on the shared prefix.
                    assert {k:v for k,v in x.items() if k!='truncated'} == {
                           k:v for k,v in y.items() if k!='truncated'}
    # Metric fixture: same one-cell successful displacement, two blocked attempts.
    metric_fixture = dict(shortest=1, moves=1, actions=3, spl=1., action_efficiency=1/3)
    failures = [r for r in loops if not r['success']]
    result = dict(status='passed', maps=35, episodes=360, replayed_steps=steps,
                  observations_checked=2*steps, horizon_prefix_pairs=240,
                  metrics_recomputed=12*7, duplicate_seeds_rejected=data['duplicate_seeds_rejected'],
                  metric_fixture=metric_fixture, diagnostic_rows=loops, failures=failures)
    (ROOT/'audit.json').write_text(json.dumps(result, indent=2)+'\n')
    print(json.dumps({k:v for k,v in result.items() if k not in ('diagnostic_rows','failures')},indent=2))
    print('H32 memory failures:', [r for r in failures if r['policy']=='memory'])
    print('H32 memory_reset failures:',len([r for r in failures if r['policy']=='memory_reset']))


if __name__ == '__main__':
    main()
