import concurrent.futures
import copy
import json
import os
from pathlib import Path
import sqlite3
import subprocess
import sys
import tempfile
import threading
import time
import unittest

from cryptography.hazmat.primitives.asymmetric import ec
from odexa_ref import crypto
from profiles.authority_history import AuthorityHistory, HistoryError
from profiles.authority_transport import AuthoritySnapshot
from test_delegation_profile import fixture, observation, NOW, ORIGIN, BEFORE, AFTER


class AuthorityHistoryTests(unittest.TestCase):
    def setUp(self):
        self.temp = tempfile.TemporaryDirectory()
        self.path = Path(self.temp.name) / 'history' / 'authority.sqlite3'
        self.store = AuthorityHistory(self.path, clock=lambda: NOW)
        self.authority, self.pol, _ = fixture()

    def tearDown(self):
        self.temp.cleanup()

    def snapshot(self, authority=None, pol=None, **obs_changes):
        a = self.authority if authority is None else authority
        p = self.pol if pol is None else pol
        obs = observation(a, p)
        obs['origin'] = a['origin']
        obs['authority_url'] = a['origin'] + '/odexa-service.json'
        obs['policy_url'] = a['origin'] + '/odexa.json'
        obs.update(obs_changes)
        return AuthoritySnapshot(crypto.json_bytes(a), crypto.json_bytes(p), obs)

    def observe(self, **kwargs):
        return self.store.observe(self.snapshot(), **kwargs)

    def bump(self):
        self.authority['revision'] += 1

    def new_key(self, scalar=2):
        key = copy.deepcopy(self.authority['services'][0]['signing_keys'][0])
        key.update(kid='https://provider.example/keys/' + str(scalar),
                   public_jwk=crypto.public_jwk(ec.derive_private_key(scalar, ec.SECP256R1())),
                   state='active', retired_at=None, revoked_at=None)
        return key

    def test_restart_retains_exact_bytes_and_qualified_reference(self):
        snapshot = self.snapshot()
        result = self.store.observe(snapshot)
        reopened = AuthorityHistory(self.path, clock=lambda: NOW)
        current = reopened.current(ORIGIN)
        self.assertEqual(current['authority_bytes'], snapshot.authority_bytes)
        self.assertEqual(current['policy_bytes'], snapshot.policy_bytes)
        self.assertEqual(current['observation'], snapshot.observation)
        self.assertEqual(current['snapshot_digest'], result['snapshot_digest'])
        self.assertFalse(current['resource_permission_granted'])
        self.assertEqual(current['authority_status'], 'highest_observed')
        self.assertEqual(reopened.history(ORIGIN, 1, 1)['authority_status'], 'historical_only')
        self.assertEqual(reopened.require_current(snapshot)['snapshot_digest'], result['snapshot_digest'])

    def test_identical_refresh_retains_documents_and_new_observation(self):
        first = self.observe()
        later = '2026-09-17T00:00:01Z'
        second = self.store.observe(self.snapshot(checked_at=later), now=lambda: later)
        self.assertEqual(first['snapshot_digest'], second['snapshot_digest'])
        self.assertGreater(second['observation_sequence'], first['observation_sequence'])
        self.assertEqual(len(self.store.observations(ORIGIN)), 2)

    def test_authority_and_policy_revisions_progress_independently(self):
        self.observe()
        self.pol['revision'] = 2
        self.observe()
        self.bump()
        self.observe()
        self.assertEqual(self.store.current(ORIGIN)['policy_revision'], 2)
        self.assertEqual(self.store.current(ORIGIN)['authority_revision'], 2)
        self.assertEqual(self.store.history(ORIGIN, 1, 2)['policy_revision'], 2)
        with self.assertRaises(HistoryError):
            self.store.history(ORIGIN, 2, 1)  # Never synthesize an unobserved pair.

    def test_both_revision_rollbacks_rejected_after_restart(self):
        self.observe(); self.bump(); self.pol['revision'] = 2; self.observe()
        reopened = AuthorityHistory(self.path, clock=lambda: NOW)
        for a, p in [(1, 2), (2, 1), (1, 1)]:
            with self.subTest(authority=a, policy=p):
                authority, pol = copy.deepcopy(self.authority), copy.deepcopy(self.pol)
                authority['revision'], pol['revision'] = a, p
                with self.assertRaises(HistoryError):
                    reopened.observe(self.snapshot(authority, pol))
        self.assertEqual(len(reopened.observations(ORIGIN)), 2)

    def test_same_revision_whitespace_conflict_is_not_canonicalized_away(self):
        self.observe()
        snap = self.snapshot()
        for kind in ('authority', 'policy'):
            with self.subTest(kind=kind):
                ar, pr, obs = snap.authority_bytes, snap.policy_bytes, copy.deepcopy(snap.observation)
                if kind == 'authority':
                    ar = json.dumps(self.authority, indent=2).encode(); obs['authority_digest'] = crypto.digest(ar)
                else:
                    pr = json.dumps(self.pol, indent=2).encode(); obs['policy_digest'] = crypto.digest(pr)
                with self.assertRaises(HistoryError):
                    self.store.observe(AuthoritySnapshot(ar, pr, obs))
        self.assertEqual(len(self.store.observations(ORIGIN)), 1)

    def test_same_revision_semantic_change_rejected(self):
        self.observe()
        self.authority['services'][0]['base_url'] = 'https://new-provider.example/api/'
        with self.assertRaises(HistoryError): self.observe()
        self.assertEqual(self.store.current(ORIGIN)['authority_revision'], 1)

    def test_policy_lineage_switch_is_not_a_revision_reset(self):
        self.observe(); self.bump(); self.pol['revision'] = 2
        self.authority['policy_id'] = self.pol['policy_id'] = ORIGIN + '/policies/replacement'
        with self.assertRaises(HistoryError): self.observe()

    def test_authenticated_higher_revision_endpoint_migration_requires_no_manual_allowlist(self):
        first = self.snapshot(); self.store.observe(first)
        self.bump()
        new = 'https://new-provider.example/tenant-c/api/'
        self.authority['services'][0]['base_url'] = new
        result = self.observe()
        self.assertEqual(result['endpoint_changes'], [{'service_id': ORIGIN + '/services/one',
            'issuer': 'https://provider.example/operator',
            'previous_base_url': 'https://provider.example/tenant-a/api/', 'base_url': new}])
        self.assertEqual(self.store.history(ORIGIN, 1, 1)['authority_bytes'], first.authority_bytes)
        with self.assertRaises(HistoryError): self.store.require_current(first)

    def test_service_id_cannot_redefine_issuer(self):
        self.observe(); self.bump()
        self.authority['services'][0]['issuer'] = 'https://new-provider.example/operator'
        with self.assertRaises(HistoryError): self.observe()

    def test_new_provider_has_new_service_and_key_identity(self):
        self.observe(); self.bump()
        service = self.authority['services'][0]
        service.update(id=ORIGIN + '/services/two', issuer='https://new-provider.example/operator',
                       base_url='https://new-provider.example/api/', signing_keys=[self.new_key()])
        self.authority['delegations'][0]['delegate_service_id'] = service['id']
        result = self.observe()
        self.assertEqual(result['authority_revision'], 2)
        self.assertEqual(self.store.history(ORIGIN, 1, 1)['authority_revision'], 1)

    def test_withdrawn_service_id_cannot_silently_resume(self):
        self.observe()
        old_service = copy.deepcopy(self.authority['services'][0])
        self.bump()
        service = self.authority['services'][0]
        service['id'] = ORIGIN + '/services/two'  # Same issuer and still-published key allowed.
        self.authority['delegations'][0]['delegate_service_id'] = service['id']
        self.observe(); self.bump()
        self.authority['services'] = [old_service]
        self.authority['delegations'][0]['delegate_service_id'] = old_service['id']
        with self.assertRaises(HistoryError): self.observe()

    def test_key_rotation_retirement_then_revocation_keeps_original_receipts(self):
        original = self.snapshot(); self.store.observe(original)
        self.bump(); old = self.authority['services'][0]['signing_keys'][0]
        old.update(state='retired', retired_at=NOW)
        self.authority['services'][0]['signing_keys'].append(self.new_key())
        self.observe(); self.bump(); old.update(state='revoked', revoked_at=NOW)
        self.observe()
        reopened = AuthorityHistory(self.path, clock=lambda: NOW)
        self.assertEqual(reopened.history(ORIGIN, 1, 1)['authority_bytes'], original.authority_bytes)
        self.assertEqual(reopened.current(ORIGIN)['authority_revision'], 3)

    def test_key_material_validity_uses_and_issuer_redefinitions_rejected(self):
        self.observe()
        base = copy.deepcopy(self.authority)
        for change in ('material', 'validity', 'uses', 'issuer'):
            with self.subTest(change=change):
                self.authority = copy.deepcopy(base); self.bump()
                key = self.authority['services'][0]['signing_keys'][0]
                if change == 'material': key['public_jwk'] = self.new_key()['public_jwk']
                if change == 'validity': key['not_after'] = '2026-09-19T00:00:00Z'
                if change == 'uses': key['uses'] = ['odexa-receipt+jws']
                if change == 'issuer':
                    service = self.authority['services'][0]
                    service.update(id=ORIGIN + '/services/new', issuer='https://provider.example/new-operator')
                    self.authority['delegations'][0]['delegate_service_id'] = service['id']
                with self.assertRaises(HistoryError): self.observe()
        self.assertEqual(len(self.store.observations(ORIGIN)), 1)

    def test_retired_and_revoked_keys_cannot_become_active_again(self):
        self.observe(); self.bump()
        key = self.authority['services'][0]['signing_keys'][0]
        key.update(state='retired', retired_at=NOW); self.observe(); self.bump()
        key.update(state='active', retired_at=None)
        with self.assertRaises(HistoryError): self.observe()
        key.update(state='revoked', retired_at=NOW, revoked_at=NOW); self.observe(); self.bump()
        key.update(state='retired', revoked_at=None)
        with self.assertRaises(HistoryError): self.observe()

    def test_terminal_timestamps_cannot_change(self):
        self.observe(); self.bump()
        key = self.authority['services'][0]['signing_keys'][0]
        key.update(state='retired', retired_at=BEFORE); self.observe(); self.bump()
        key['retired_at'] = NOW
        with self.assertRaises(HistoryError): self.observe()
        key.update(retired_at=BEFORE, state='revoked', revoked_at=NOW); self.observe(); self.bump()
        key['revoked_at'] = BEFORE
        with self.assertRaises(HistoryError): self.observe()

    def test_removed_key_is_withdrawn_even_if_original_metadata_calls_it_active(self):
        self.observe(); self.bump()
        old = self.authority['services'][0]['signing_keys'][0]
        self.authority['services'][0]['signing_keys'] = [self.new_key()]
        self.observe(); self.bump(); self.authority['services'][0]['signing_keys'].append(old)
        with self.assertRaises(HistoryError): self.observe()

    def test_revoked_material_cannot_return_under_new_kid_after_removal_and_restart(self):
        self.observe(); self.bump()
        key = self.authority['services'][0]['signing_keys'][0]
        key.update(state='revoked', revoked_at=NOW); self.observe()
        old = copy.deepcopy(key)
        self.bump(); self.authority['services'][0]['signing_keys'] = [self.new_key()]; self.observe()
        self.store = AuthorityHistory(self.path, clock=lambda: NOW)
        for state in ('active', 'retired'):
            with self.subTest(state=state):
                alias = copy.deepcopy(old)
                alias.update(kid='https://provider.example/keys/alias', state=state, revoked_at=None,
                             retired_at=NOW if state == 'retired' else None)
                self.authority['revision'] = 4
                self.authority['services'][0]['signing_keys'].append(alias)
                with self.assertRaises(HistoryError): self.observe()
                self.authority['services'][0]['signing_keys'].pop()
        self.assertEqual(self.store.current(ORIGIN)['authority_revision'], 3)

    def test_first_snapshot_cannot_call_same_material_revoked_and_active_or_retired(self):
        key = self.authority['services'][0]['signing_keys'][0]
        key.update(state='revoked', revoked_at=NOW)
        for state in ('active', 'retired'):
            with self.subTest(state=state):
                alias = copy.deepcopy(key)
                alias.update(kid='https://provider.example/keys/alias', state=state, revoked_at=None,
                             retired_at=NOW if state == 'retired' else None)
                self.authority['services'][0]['signing_keys'].append(alias)
                with self.assertRaises(HistoryError): self.observe()
                self.authority['services'][0]['signing_keys'].pop()
        with self.assertRaises(HistoryError): self.store.current(ORIGIN)

    def test_retirement_alone_does_not_blacklist_material_under_new_kid(self):
        self.observe(); self.bump()
        key = self.authority['services'][0]['signing_keys'][0]
        key.update(state='retired', retired_at=NOW); self.observe(); self.bump()
        alias = copy.deepcopy(key)
        alias.update(kid='https://provider.example/keys/renewed', state='active', retired_at=None)
        self.authority['services'][0]['signing_keys'].append(alias)
        self.observe()
        self.assertEqual(self.store.current(ORIGIN)['authority_revision'], 3)

    def test_revoked_material_memory_does_not_cross_origin_boundary(self):
        self.authority['services'][0]['signing_keys'][0].update(state='revoked', revoked_at=NOW)
        self.observe()
        self.authority['services'][0]['signing_keys'][0].update(state='active', revoked_at=None)
        raw = crypto.json_bytes(self.authority).replace(ORIGIN.encode(), b'https://second.example')
        praw = crypto.json_bytes(self.pol).replace(ORIGIN.encode(), b'https://second.example')
        self.store.observe(self.snapshot(crypto.strict_json(raw), crypto.strict_json(praw)))
        self.assertEqual(self.store.current('https://second.example')['authority_revision'], 1)

    def test_conflicting_key_states_between_services_rejected_atomically(self):
        self.observe(); self.bump()
        other = copy.deepcopy(self.authority['services'][0]); other['id'] = ORIGIN + '/services/two'
        other['signing_keys'][0].update(state='retired', retired_at=NOW)
        self.authority['services'].append(other)
        with self.assertRaises(HistoryError): self.observe()
        self.assertEqual(self.store.current(ORIGIN)['authority_revision'], 1)

    def test_invalid_observation_bindings_do_not_create_origin(self):
        for changes in [{'available': False}, {'revalidated': False}, {'redirected': True},
                        {'tls_verified': False}, {'source': 'cache'}, {'available': 1},
                        {'policy_digest': 'sha256:' + '0' * 64}, {'authority_url': ORIGIN + '/elsewhere'},
                        {'origin': 'https://attacker.example'}]:
            with self.subTest(changes=changes):
                with self.assertRaises(HistoryError): self.store.observe(self.snapshot(**changes))
        with self.assertRaises(HistoryError): self.store.current(ORIGIN)

    def test_age_boundary_future_expiry_and_future_revocation(self):
        for checked in ['2026-09-16T23:59:54Z', '2026-09-17T00:00:01Z']:
            with self.subTest(checked=checked):
                with self.assertRaises(HistoryError): self.store.observe(self.snapshot(checked_at=checked))
        self.store.observe(self.snapshot(checked_at='2026-09-16T23:59:55Z'))
        for mutate in [lambda: self.authority.update(expires_at=NOW), lambda: self.pol.update(expires_at=NOW),
                       lambda: self.authority['services'][0]['signing_keys'][0].update(state='revoked', revoked_at=AFTER)]:
            a, p = copy.deepcopy(self.authority), copy.deepcopy(self.pol)
            self.bump(); mutate()
            with self.assertRaises(HistoryError): self.observe()
            self.authority, self.pol = a, p

    def test_history_does_not_reactivate_expired_snapshot(self):
        snap = self.snapshot(); self.store.observe(snap)
        later = '2026-09-18T00:00:00Z'
        reopened = AuthorityHistory(self.path, clock=lambda: later)
        self.assertEqual(reopened.history(ORIGIN, 1, 1)['authority_status'], 'historical_only')
        # Guard asserts only local high-water equality, deliberately not freshness.
        self.assertFalse(reopened.require_current(snap)['resource_permission_granted'])
        with self.assertRaises(HistoryError): reopened.observe(snap)

    def test_require_current_checks_each_revision_and_exact_bytes(self):
        snap = self.snapshot(); self.store.observe(snap)
        self.pol['revision'] = 2; self.observe()
        with self.assertRaises(HistoryError): self.store.require_current(snap)
        fresh = self.snapshot(); self.store.require_current(fresh)
        self.pol['rules'][0]['effect'] = 'permit'
        with self.assertRaises(HistoryError): self.store.require_current(self.snapshot())

    def test_clock_rollback_rejected_and_scalar_time_not_accepted(self):
        self.observe()
        early = '2026-09-16T23:59:59Z'
        with self.assertRaises(HistoryError):
            self.store.observe(self.snapshot(checked_at=early), now=lambda: early)
        with self.assertRaises(HistoryError): self.observe(now=NOW)
        self.assertEqual(len(self.store.observations(ORIGIN)), 1)

    def test_lock_wait_resamples_clock_and_rejects_stale_observation(self):
        self.observe()
        blocker = sqlite3.connect(self.path, isolation_level=None); blocker.execute('BEGIN IMMEDIATE')
        value = [NOW]
        started = threading.Event()
        def writer():
            started.set()
            return self.observe(now=lambda: value[0])
        try:
            with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
                pending = pool.submit(writer)
                self.assertTrue(started.wait(2)); time.sleep(0.05)
                self.assertFalse(pending.done())
                value[0] = '2026-09-17T00:00:06Z'
                blocker.rollback()
                with self.assertRaises(HistoryError): pending.result(timeout=5)
        finally:
            blocker.close()
        self.assertEqual(len(self.store.observations(ORIGIN)), 1)

    def test_concurrent_same_revision_conflicts_have_one_winner(self):
        self.observe()
        self.bump()
        first = self.snapshot()
        self.authority['services'][0]['base_url'] = 'https://provider.example/new/api/'
        second = self.snapshot()
        barrier = threading.Barrier(2)
        def writer(snapshot):
            separate = AuthorityHistory(self.path, clock=lambda: NOW)
            barrier.wait()
            try:
                separate.observe(snapshot); return True
            except HistoryError:
                return False
        with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool:
            results = list(pool.map(writer, (first, second)))
        self.assertEqual(sorted(results), [False, True])
        self.assertEqual(len(self.store.observations(ORIGIN)), 2)
        self.assertIn(self.store.current(ORIGIN)['authority_bytes'], (first.authority_bytes, second.authority_bytes))

    def test_validation_failure_after_pin_changes_rolls_back_whole_transaction(self):
        self.observe(); self.bump()
        service = self.authority['services'][0]
        service['base_url'] = 'https://provider.example/new/api/'
        service['signing_keys'][0]['public_jwk'] = self.new_key()['public_jwk']
        with self.assertRaises(HistoryError): self.observe()
        service['signing_keys'][0]['public_jwk'] = fixture()[0]['services'][0]['signing_keys'][0]['public_jwk']
        result = self.observe()
        self.assertEqual(result['endpoint_changes'][0]['previous_base_url'], 'https://provider.example/tenant-a/api/')

    def test_process_crash_before_and_after_commit(self):
        self.observe(); self.bump()
        snapshot = self.snapshot()
        payload = json.dumps({'authority': self.authority, 'policy': self.pol, 'observation': snapshot.observation})
        script = '''
import json,os,sys
from types import SimpleNamespace
from odexa_ref import crypto
from profiles.authority_history import AuthorityHistory
value=json.loads(sys.stdin.read())
snapshot=SimpleNamespace(authority_bytes=crypto.json_bytes(value['authority']),policy_bytes=crypto.json_bytes(value['policy']),observation=value['observation'])
def fault(stage):
    if stage==sys.argv[2]: os._exit(81)
store=AuthorityHistory(sys.argv[1],clock=lambda:'2026-09-17T00:00:00Z',fault_hook=fault)
store.observe(snapshot)
'''
        for stage, expected in [('before_commit', 1), ('after_commit', 2)]:
            with self.subTest(stage=stage):
                run = subprocess.run([sys.executable, '-c', script, str(self.path), stage], input=payload,
                                     text=True, capture_output=True, timeout=10)
                self.assertEqual(run.returncode, 81, run.stderr)
                reopened = AuthorityHistory(self.path, clock=lambda: NOW)
                self.assertEqual(reopened.current(ORIGIN)['authority_revision'], expected)
                self.assertEqual(len(reopened.observations(ORIGIN)), expected)

    def test_origins_are_isolated_and_audit_pages_bounded(self):
        self.observe()
        self.bump(); self.observe()
        raw = crypto.json_bytes(self.authority).replace(ORIGIN.encode(), b'https://second.example')
        praw = crypto.json_bytes(self.pol).replace(ORIGIN.encode(), b'https://second.example')
        a, p = crypto.strict_json(raw), crypto.strict_json(praw)
        self.store.observe(self.snapshot(a, p))
        self.assertEqual(self.store.current('https://second.example')['authority_revision'], 2)
        first = self.store.observations(ORIGIN, limit=1)
        rest = self.store.observations(ORIGIN, after_sequence=first[0]['sequence'], limit=1)
        self.assertEqual(len(first), 1); self.assertEqual(len(rest), 1)
        self.assertGreater(rest[0]['sequence'], first[0]['sequence'])
        with self.assertRaises(HistoryError): self.store.observations(ORIGIN, limit=1001)

    def test_private_files_and_symlink_boundaries(self):
        self.assertEqual(self.path.stat().st_mode & 0o777, 0o600)
        self.assertEqual(self.path.parent.stat().st_mode & 0o777, 0o700)
        public = Path(self.temp.name) / 'public'; public.mkdir(mode=0o755)
        with self.assertRaises(HistoryError): AuthorityHistory(public / 'db')
        alias = self.path.parent / 'alias'; alias.symlink_to(self.path)
        with self.assertRaises(OSError): AuthorityHistory(alias)

    def test_strict_wire_duplicates_and_fractional_revision_rejected(self):
        snap = self.snapshot()
        for ar in [snap.authority_bytes.replace(b'"revision":1', b'"revision":1.0'),
                   snap.authority_bytes.replace(b'"revision":1', b'"revision":1,"revision":1')]:
            obs = dict(snap.observation, authority_digest=crypto.digest(ar))
            with self.assertRaises(HistoryError): self.store.observe(AuthoritySnapshot(ar, snap.policy_bytes, obs))


if __name__ == '__main__':
    unittest.main()
