"""Approval and reconciliation regressions with trusted local identity fixtures."""
import json
import multiprocessing
from pathlib import Path
import sqlite3
import tempfile
import unittest
from unittest.mock import patch
from approval_runtime import ApprovalRuntime, Principal
from cli import DOCS
from memory import scope_key, put, forget
from runtime import InjectedCrash, encode


def decision_race(path, scope, binding_hash, action, barrier, output):
    app=ApprovalRuntime(path,clock=lambda:1000)
    actor=Principal('reviewer-user',frozenset({'report_reviewer'}),frozenset({scope}))
    barrier.wait(timeout=10)
    try:
        output.put(app.decide('j',actor,binding_hash,action))
    except ValueError:
        output.put('conflict')
    finally:
        app.db.close()


class ApprovalTests(unittest.TestCase):
    def setUp(self):
        self.tmp=tempfile.TemporaryDirectory()
        self.path=str(Path(self.tmp.name)/'approval.sqlite')
        self.now=1000.0
        self.app=ApprovalRuntime(self.path,clock=lambda:self.now,lease_seconds=30)
        self.scope=scope_key('t1','requester','p1')
        self.actor=Principal('reviewer',frozenset({'report_reviewer'}),frozenset({self.scope}))
        self.spec={'version':'approval-report/v1','provider':'fixture','model':'','endpoint':'',
                   'docs':DOCS,'report_target':'local-report-store','approval_ttl':120}
        self.app.submit('j',self.spec,self.scope)

    def tearDown(self):
        self.app.db.close()
        self.tmp.cleanup()

    def wait_for_approval(self):
        self.assertEqual(self.app.run('j'),'waiting_approval')
        return self.app.approval('j')['binding_hash']

    def approve(self):
        h=self.wait_for_approval()
        self.assertEqual(self.app.decide('j',self.actor,h,'approve'),'approved')
        return h

    def receipt_count(self):
        path=self.path+'.publisher.sqlite'
        if not Path(path).exists(): return 0
        with sqlite3.connect(path) as db:
            return db.execute('SELECT COUNT(*) FROM receipts').fetchone()[0]

    def test_pause_has_no_side_effect_and_keeps_checkpoints(self):
        self.wait_for_approval()
        self.assertEqual(set(self.app.snapshot('j')['checkpoints']),{'collect','draft','verify'})
        self.assertEqual(self.receipt_count(),0)
        self.assertEqual(self.app.run('j'),'waiting_approval')
        self.assertEqual(self.app.job('j')['attempts'],1)

    def test_approved_resume_does_not_redraft(self):
        self.approve()
        with patch('approval_runtime.draft',side_effect=AssertionError('must not redraft')):
            self.assertEqual(self.app.run('j',model_fn=lambda *_:self.fail('redrafted')),'succeeded')
        self.assertEqual(self.app.approval('j')['status'],'completed')
        self.assertEqual(self.receipt_count(),1)

    def test_requester_cannot_approve_own_action(self):
        h=self.wait_for_approval()
        actor=Principal('requester',self.actor.roles,self.actor.allowed_scopes)
        with self.assertRaises(PermissionError): self.app.decide('j',actor,h,'approve')
        self.assertEqual(self.app.approval('j')['status'],'pending')

    def test_outside_scope_rejected(self):
        h=self.wait_for_approval()
        actor=Principal('outsider',self.actor.roles,frozenset())
        with self.assertRaises(PermissionError): self.app.decide('j',actor,h,'approve')
        self.assertEqual(self.receipt_count(),0)

    def test_missing_reviewer_role_rejected(self):
        h=self.wait_for_approval()
        actor=Principal('viewer',frozenset(),self.actor.allowed_scopes)
        with self.assertRaises(PermissionError): self.app.decide('j',actor,h,'approve')

    def test_hash_must_match_reviewed_action(self):
        self.wait_for_approval()
        with self.assertRaises(ValueError): self.app.decide('j',self.actor,'another-hash','approve')
        self.assertEqual(self.app.job('j')['status'],'waiting_approval')

    def test_duplicate_approval_is_idempotent(self):
        h=self.approve()
        self.app.decide('j',self.actor,h,'approve')
        events=self.app.snapshot('j')['events']
        self.assertEqual(sum(e['kind']=='approval_approved' for e in events),1)
        self.assertEqual(self.app.run('j'),'succeeded')
        self.assertEqual(self.receipt_count(),1)

    def test_rejection_is_terminal(self):
        h=self.wait_for_approval()
        self.assertEqual(self.app.decide('j',self.actor,h,'reject'),'rejected')
        self.assertEqual(self.app.run('j'),'failed')
        with self.assertRaises(ValueError): self.app.decide('j',self.actor,h,'approve')
        self.assertEqual(self.receipt_count(),0)

    def test_pending_approval_expires(self):
        h=self.wait_for_approval()
        self.now+=121
        self.assertEqual(self.app.decide('j',self.actor,h,'approve'),'expired')
        self.assertEqual(self.app.run('j'),'failed')
        self.assertEqual(self.receipt_count(),0)

    def test_approval_expires_before_dispatch(self):
        self.approve()
        self.now+=121
        self.assertEqual(self.app.run('j'),'failed')
        self.assertEqual(self.app.approval('j')['status'],'expired')
        self.assertEqual(self.receipt_count(),0)

    def test_draft_change_after_approval_invalidates_binding(self):
        self.approve()
        self.app.db.execute('UPDATE checkpoints SET payload=? WHERE job_id=? AND step=?',
                            (encode({'summary':'changed','citations':['s1']}),'j','draft'))
        self.assertEqual(self.app.run('j'),'failed')
        self.assertEqual(self.app.approval('j')['status'],'invalidated')
        self.assertEqual(self.receipt_count(),0)

    def test_memory_change_while_waiting_invalidates_decision(self):
        h=self.wait_for_approval()
        put(self.app.db,self.scope,'language','en','message-new',2000)
        self.assertEqual(self.app.decide('j',self.actor,h,'approve'),'invalidated')
        self.assertEqual(self.app.run('j'),'failed')
        self.assertEqual(self.receipt_count(),0)

    def test_cancel_before_dispatch_prevents_publish(self):
        self.approve()
        self.app.cancel('j')
        self.assertEqual(self.app.run('j'),'cancelled')
        self.assertEqual(self.app.approval('j')['status'],'cancelled')
        self.assertEqual(self.receipt_count(),0)

    def test_crash_after_effect_and_expiry_reconciles_without_write(self):
        self.approve()
        with self.assertRaises(InjectedCrash): self.app.run('j','after_effect')
        self.now+=121
        with patch.object(self.app,'publish',side_effect=AssertionError('must not resend')):
            self.assertEqual(self.app.run('j'),'effect_confirmed')
        self.assertEqual(self.receipt_count(),1)
        self.assertFalse(self.app.saved('j','publish')['fresh_write_performed'])
        self.assertEqual(self.app.run('j'),'effect_confirmed')

    def test_no_receipt_after_dispatch_stays_unknown(self):
        self.approve()
        with self.assertRaises(InjectedCrash): self.app.run('j','before_effect')
        self.now+=31
        with patch.object(self.app,'publish',side_effect=AssertionError('must not resend')):
            self.assertEqual(self.app.run('j'),'reconciling')
            self.assertEqual(self.app.run('j'),'reconciling')
        self.assertEqual(self.receipt_count(),0)
        self.assertIsNone(self.app.saved('j','publish'))

    def test_cancel_after_effect_does_not_claim_rollback(self):
        self.approve()
        with self.assertRaises(InjectedCrash): self.app.run('j','after_effect')
        self.app.cancel('j')
        self.assertEqual(self.app.job('j')['status'],'reconciling')
        self.assertEqual(self.app.reconcile('j'),'effect_confirmed')
        self.assertEqual(self.app.job('j')['error'],'cancel_requested_after_dispatch')
        self.assertEqual(self.receipt_count(),1)

    def test_memory_revoke_after_effect_still_allows_receipt_lookup(self):
        self.approve()
        with self.assertRaises(InjectedCrash): self.app.run('j','after_effect')
        forget(self.app.db,self.scope,'language')
        self.now+=31
        self.assertEqual(self.app.run('j'),'effect_confirmed')
        self.assertEqual(self.receipt_count(),1)

    def test_concurrent_approve_reject_one_decision_wins(self):
        h=self.wait_for_approval()
        ctx=multiprocessing.get_context('spawn')
        barrier,q=ctx.Barrier(2),ctx.Queue()
        processes=[ctx.Process(target=decision_race,args=(self.path,self.scope,h,a,barrier,q))
                   for a in ('approve','reject')]
        for proc in processes: proc.start()
        results=[q.get(timeout=15),q.get(timeout=15)]
        for proc in processes:
            proc.join(timeout=10)
            self.assertEqual(proc.exitcode,0)
        self.assertEqual(results.count('conflict'),1)
        self.assertEqual(sum(r in ('approved','rejected') for r in results),1)
        q.close()


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