"""Offline failure injection; real DB, deterministic clock, mocked model transport."""
import copy
import json
import multiprocessing
import os
from pathlib import Path
import sqlite3
import subprocess
import sys
import tempfile
import unittest
from unittest.mock import patch
import urllib.error
from cli import DOCS
from evaluate import grade
from memory import scope_key, put, forget, current
from provider import draft, TransientError, PermanentError
from runtime import Runtime, InjectedCrash, LostLease


def compete(path, barrier, output):
    r = Runtime(path, clock=lambda: 1000)
    barrier.wait(timeout=10)
    output.put(r.claim('j', str(os.getpid())))
    r.db.close()


class Lab(unittest.TestCase):
    def setUp(self):
        self.tmp = tempfile.TemporaryDirectory()
        self.path = str(Path(self.tmp.name) / 'run.sqlite')
        self.now = 1000.0
        self.r = Runtime(self.path, clock=lambda: self.now, lease_seconds=30)
        self.scope = scope_key('t1', 'u1', 'p1')
        self.spec = {'version': 'report/v1', 'provider': 'fixture', 'model': '', 'endpoint': '', 'docs': DOCS}
        self.r.submit('j', self.spec, self.scope)

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

    def test_normal_and_repeat_delivery(self):
        self.assertEqual(self.r.run('j'), 'succeeded')
        before = self.r.snapshot('j')
        self.assertEqual(self.r.run('j'), 'succeeded')
        self.assertEqual(before, self.r.snapshot('j'))
        self.assertTrue(grade(before)['passed'])

    def test_input_identity_conflict(self):
        other = {**self.spec, 'version': 'report/v2'}
        with self.assertRaises(ValueError):
            self.r.submit('j', other, self.scope)

    def test_resume_after_collect(self):
        with self.assertRaises(InjectedCrash):
            self.r.run('j', 'after_collect')
        self.assertEqual(set(self.r.snapshot('j')['checkpoints']), {'collect'})
        self.now += 31
        self.assertEqual(self.r.run('j'), 'succeeded')
        events = self.r.snapshot('j')['events']
        self.assertEqual(sum(e['kind'] == 'step_started' and e['detail'] == 'collect' for e in events), 1)

    def test_effect_before_checkpoint(self):
        with self.assertRaises(InjectedCrash):
            self.r.run('j', 'after_effect')
        self.assertIsNone(self.r.saved('j', 'publish'))
        self.now += 31
        self.assertEqual(self.r.run('j'), 'succeeded')
        self.assertTrue(self.r.saved('j', 'publish')['deduplicated'])
        with sqlite3.connect(self.path + '.publisher.sqlite') as db:
            self.assertEqual(db.execute('SELECT COUNT(*) FROM receipts').fetchone()[0], 1)

    def test_old_generation_cannot_commit(self):
        old = self.r.claim('j', 'old')
        self.now += 31
        new = self.r.claim('j', 'new')
        self.assertGreater(new, old)
        with self.assertRaises(LostLease):
            self.r.commit_step('j', 'old', old, 'collect', {})
        self.assertIsNone(self.r.saved('j', 'collect'))

    def test_two_processes_only_one_claim(self):
        ctx = multiprocessing.get_context('spawn')
        barrier, q = ctx.Barrier(2), ctx.Queue()
        processes = [ctx.Process(target=compete, args=(self.path, barrier, q)) for _ in range(2)]
        for p in processes: p.start()
        values = [q.get(timeout=15), q.get(timeout=15)]
        for p in processes:
            p.join(timeout=10)
            self.assertEqual(p.exitcode, 0)
        self.assertEqual(sum(v is not None for v in values), 1)
        q.close()

    def test_cancel_fences_running_worker(self):
        generation = self.r.claim('j', 'worker')
        self.assertTrue(self.r.cancel('j'))
        with self.assertRaises(LostLease):
            self.r.commit_step('j', 'worker', generation, 'collect', {})
        self.assertEqual(self.r.run('j'), 'cancelled')

    def test_transient_budget_and_backoff(self):
        def down(*_): raise TransientError('HTTP 429')
        self.assertEqual(self.r.run('j', model_fn=down), 'retry_wait')
        self.assertEqual(self.r.run('j', model_fn=down), 'retry_wait')
        self.assertEqual(self.r.job('j')['attempts'], 1)
        self.now += 2
        self.assertEqual(self.r.run('j', model_fn=down), 'retry_wait')
        self.now += 4
        self.assertEqual(self.r.run('j', model_fn=down), 'failed')
        self.assertEqual(self.r.job('j')['attempts'], 3)

    def test_unknown_citation_rejected(self):
        def bad(*_): return {'summary': 'looks plausible', 'citations': ['secret-doc']}
        self.assertEqual(self.r.run('j', model_fn=bad), 'failed')
        self.assertIsNone(self.r.saved('j', 'publish'))

    def test_memory_scope_expiry_and_confirmation(self):
        put(self.r.db, self.scope, 'a', 'yes', 'm1', 2000)
        put(self.r.db, self.scope, 'expired', 'no', 'm2', 999)
        put(self.r.db, self.scope, 'guess', 'no', 'm3', 2000, confirmed=False)
        for scope in [scope_key('t2','u1','p1'), scope_key('t1','u2','p1'), scope_key('t1','u1','p2')]:
            put(self.r.db, scope, 'a', 'private', 'm4', 2000)
        self.assertEqual([m['key'] for m in current(self.r.db, self.scope, self.now)], ['a'])

    def test_memory_revision_supersedes(self):
        self.assertEqual(put(self.r.db, self.scope, 'language', 'en', 'm1', 2000), 1)
        self.assertEqual(put(self.r.db, self.scope, 'language', 'zh-CN', 'm2', 2000), 2)
        rows = current(self.r.db, self.scope, self.now)
        self.assertEqual([(m['version'], m['value']) for m in rows], [(2, 'zh-CN')])

    def test_deleted_memory_invalidates_saved_draft(self):
        put(self.r.db, self.scope, 'language', 'zh-CN', 'm1', 2000)
        self.r.submit('with-memory', self.spec, self.scope)
        with self.assertRaises(InjectedCrash):
            self.r.run('with-memory', 'after_draft')
        forget(self.r.db, self.scope, 'language')
        self.now += 31
        self.assertEqual(self.r.run('with-memory'), 'failed')
        self.assertIsNone(self.r.saved('with-memory', 'publish'))

    def test_memory_expiry_invalidates_task(self):
        put(self.r.db, self.scope, 'language', 'zh-CN', 'm1', 1001)
        self.r.submit('expires', self.spec, self.scope)
        self.now += 2
        self.assertEqual(self.r.run('expires'), 'failed')

    def test_idempotency_payload_conflict(self):
        self.r.publish('j', {'content': 'original'})
        with self.assertRaises(PermanentError):
            self.r.publish('j', {'content': 'different'})

    def test_deadline_stops_claim(self):
        self.now += 601
        self.assertEqual(self.r.run('j'), 'failed')
        self.assertIsNone(self.r.saved('j', 'collect'))

    def test_grader_rejects_trace_mutation(self):
        self.r.run('j')
        snapshot = copy.deepcopy(self.r.snapshot('j'))
        snapshot['events'].append(next(e for e in snapshot['events'] if e['kind'] == 'step_committed'))
        self.assertFalse(grade(snapshot)['passed'])

    def test_cli_really_terminates_after_commit(self):
        script = str(Path(__file__).with_name('cli.py'))
        path = str(Path(self.tmp.name) / 'process.sqlite')
        subprocess.run([sys.executable, script, 'submit', '--db', path], check=True, capture_output=True)
        crashed = subprocess.run([sys.executable, script, 'run', '--db', path, '--fault', 'after_collect'], capture_output=True)
        self.assertEqual(crashed.returncode, 75)
        r = Runtime(path)
        self.assertIsNotNone(r.saved('report-001', 'collect'))
        self.assertEqual(r.job('report-001')['status'], 'running')
        r.db.close()

    def test_api_adapter_mocked_success_and_429(self):
        spec = {**self.spec, 'provider': 'api', 'endpoint': 'https://example.invalid/chat/completions', 'model': 'test'}
        class Response:
            def __enter__(self): return self
            def __exit__(self, *_): pass
            def read(self, _):
                return json.dumps({'choices':[{'message':{'content':json.dumps({'summary':'ok','citations':['s1']})}}]}).encode()
        with patch.dict(os.environ, {'LLM_API_KEY':'test-key'}):
            with patch('urllib.request.OpenerDirector.open', return_value=Response()):
                self.assertEqual(draft(spec, {'docs':DOCS})['summary'], 'ok')
            with patch('urllib.request.OpenerDirector.open', side_effect=urllib.error.HTTPError(spec['endpoint'], 429, 'rate limited', {}, None)):
                with self.assertRaises(TransientError): draft(spec, {'docs':DOCS})


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