"""Durable transport reporting, including actual selected-collector HTTPS."""
import contextlib
import copy
import os
from pathlib import Path
import sqlite3
import tempfile
import unittest
from unittest import mock

from cryptography.hazmat.primitives.asymmetric import ec
from odexa_ref import crypto
from profiles.authority_history import AuthorityHistory
from profiles.authority_transport import ScopedCredential, HTTPSResult
from profiles.free_gateway import FreeGateway, GatewayError
from profiles.gateway_reports import GatewayReporter, delivery_report
from test_free_contracts import fixtures, NOW, uid
import test_free_service as tls_fixture


class GatewayReportJournalTests(unittest.TestCase):
    def setUp(self):
        self.temp = tempfile.TemporaryDirectory(); self.path = Path(self.temp.name)
        values, context = fixtures(); self.binding = values['binding']; self.response = values['introspection_response']
        self.request_id = self.response['request_id']; self.key = ec.derive_private_key(4, ec.SECP256R1())
        self.base = 'https://provider.example/tenant-a/'
        self.reporter = GatewayReporter(self.base+'gateway', self.base+'gateway/key', self.key,
            ScopedCredential('https://provider.example', '/tenant-a/', 'Basic reporter-only'))
        self.args = dict(binding=self.binding, base_url=self.base, transport=None, history=None,
            credential=ScopedCredential('https://provider.example', '/tenant-a/', 'Basic checker-only'),
            reporting=self.reporter, clock=lambda: NOW)
        self.gateway = FreeGateway(self.path/'gateway.sqlite3', **self.args)
        with self.gateway._tx() as db:
            db.execute('INSERT INTO admissions VALUES(?,?,?,?,?,?,?,?,?,?,?,?)',
                (self.request_id, crypto.json_bytes(self.response), context[1], context[0], b'{}', b'{}',
                 'in_flight', NOW, NOW, None, None, None))

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

    def finish(self, **changes):
        args = dict(bytes_written=3, representation_digest=crypto.digest(b'abc'), complete=True)
        args.update(changes); self.gateway.finish_delivery(self.request_id, **args)
        return crypto.strict_json(bytes(self.gateway.report_records()[0]['payload']))

    def test_terminal_and_exact_signed_event_persist_across_restart(self):
        event = self.finish(representation_metadata=dict(media_type='text/plain', media_parameters={}, languages=['en']))
        row = self.gateway.report_records()[0]
        raw, _ = crypto.verify_jws(bytes(row['jws']), crypto.public_jwk(self.key), 'odexa-event+jws', self.reporter.key_id)
        self.assertEqual(raw, bytes(row['payload'])); self.assertEqual(event['source'], 'origin_observed')
        self.assertEqual(event['http']['content_digest'], crypto.digest(b'abc'))
        self.assertEqual(event['http']['delivery_id'], self.request_id)
        fresh = FreeGateway(self.path/'gateway.sqlite3', **self.args)
        self.assertEqual(fresh.report_records(), self.gateway.report_records())
        self.assertEqual(fresh.records()[0]['state'], 'sent')
        with self.assertRaises(ValueError): self.finish()
        self.assertEqual(len(fresh.report_records()), 1)
        self.assertNotIn(b'Basic', bytes(row['payload'])+bytes(row['jws']))

    def test_unknown_failed_prefix_never_uses_intended_full_digest(self):
        event = self.finish(bytes_written=1, complete=False)
        self.assertEqual(event['event_type'], 'delivery.failed')
        self.assertEqual(event['http']['content_bytes'], 1)
        self.assertIsNone(event['http']['content_digest']); self.assertIsNone(event['http']['decoded_digest'])
        self.assertIsNone(self.gateway.records()[0]['representation_digest'])
        self.assertEqual(self.gateway.records()[0]['state'], 'partial')

    def test_known_failed_prefix_and_range_use_only_observed_bytes(self):
        args = dict(binding=self.binding, admission=self.response, reporter_id=self.reporter.reporter_id,
            event_id=uid(20), hop_id=uid(21), started_at=NOW, occurred_at=NOW,
            bytes_written=1, representation_digest=crypto.digest(b'abc'), complete=False,
            observed_digest=crypto.digest(b'a'))
        event = delivery_report(**args)
        self.assertEqual(event['http']['content_digest'], crypto.digest(b'a'))
        args.update(status=206, complete=True, byte_range=dict(first=0,last=0,complete_length=3), observed_digest=None)
        event = delivery_report(**args)
        self.assertEqual(event['http']['kind'], 'partial'); self.assertIsNone(event['http']['content_digest'])
        args['observed_digest'] = crypto.digest(b'a')
        self.assertEqual(delivery_report(**args)['http']['decoded_digest'], crypto.digest(b'a'))

    def test_head_not_modified_coded_and_unsupported_partial_remain_distinct(self):
        for method, status, kind in [('HEAD',200,'head'),('GET',304,'not_modified'),('GET',204,'no_content')]:
            response = copy.deepcopy(self.response); response['admission']['method'] = method
            event = delivery_report(self.binding,response,self.reporter.reporter_id,event_id=uid(20),hop_id=uid(21),
                started_at=NOW,occurred_at=NOW,bytes_written=0,representation_digest=crypto.digest(b'abc'),complete=True,status=status)
            self.assertEqual(event['event_type'],'response.completed'); self.assertEqual(event['http']['kind'],kind)
            self.assertIsNone(event['http']['content_digest'])
        args = dict(event_id=uid(20),hop_id=uid(21),started_at=NOW,occurred_at=NOW,
                    bytes_written=3,representation_digest=crypto.digest(b'abc'),complete=True,content_codings=['gzip'])
        event = delivery_report(self.binding,self.response,self.reporter.reporter_id,**args)
        self.assertIsNone(event['http']['decoded_digest']); self.assertIsNone(event['http']['decoded_bytes'])
        event = delivery_report(self.binding,self.response,self.reporter.reporter_id,status=206,**args)
        self.assertEqual(event['http']['kind'],'unsupported_partial'); self.assertIsNone(event['http']['content_digest'])

    def test_signer_or_sql_failure_rolls_back_terminal_and_outbox_together(self):
        with mock.patch('profiles.free_gateway.crypto.sign_jws',side_effect=ValueError('signer unavailable')):
            with self.assertRaises(ValueError): self.finish()
        self.assertEqual(self.gateway.records()[0]['state'],'in_flight'); self.assertEqual(self.gateway.report_records(),[])
        with self.gateway._tx() as db:
            db.execute("CREATE TRIGGER reject_terminal BEFORE UPDATE ON admissions BEGIN SELECT RAISE(ABORT,'synthetic disk failure'); END")
        with self.assertRaises(sqlite3.Error): self.finish()
        self.assertEqual(self.gateway.records()[0]['state'],'in_flight'); self.assertEqual(self.gateway.report_records(),[])

    @unittest.skipUnless(hasattr(os,'fork'), 'POSIX crash recovery test')
    def test_process_exit_after_outbox_insert_before_terminal_commit_recovers_unresolved(self):
        child = os.fork()
        if child == 0:
            old = self.gateway._tx
            @contextlib.contextmanager
            def dying_transaction():
                with old() as db:
                    db.create_function('synthetic_crash',0,lambda:os._exit(85))
                    db.execute('CREATE TEMP TRIGGER crash BEFORE UPDATE ON admissions BEGIN SELECT synthetic_crash(); END')
                    yield db
            self.gateway._tx = dying_transaction
            self.finish(); os._exit(86)
        _, status = os.waitpid(child,0); self.assertEqual(os.WEXITSTATUS(status),85)
        recovered = FreeGateway(self.path/'gateway.sqlite3', **self.args)
        self.assertEqual(recovered.records()[0]['state'],'in_flight'); self.assertEqual(recovered.report_records(),[])

    def test_reporter_switch_or_disabled_mode_cannot_rewrite_journal_identity(self):
        self.finish()
        for reporting in [None, GatewayReporter(self.base+'other',self.reporter.key_id,self.key,self.reporter.credential)]:
            with self.assertRaises(ValueError): FreeGateway(self.path/'gateway.sqlite3', **dict(self.args, reporting=reporting))
        with self.assertRaises(ValueError):
            FreeGateway(self.path/'other.sqlite3', **dict(self.args, reporting=GatewayReporter(
                self.reporter.reporter_id,self.reporter.key_id,self.key,ScopedCredential('https://attacker.example','/','Basic x'))))


class GatewayReportTLSTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls): tls_fixture.FreeServiceTests.setUpClass()
    @classmethod
    def tearDownClass(cls): tls_fixture.FreeServiceTests.tearDownClass()
    def setUp(self):
        self.f = tls_fixture.FreeServiceTests(); self.f.setUp(); self.configure()
    def tearDown(self): self.f.tearDown()

    def configure(self):
        f = self.f
        credential = ScopedCredential(f.host,f.base[len(f.host):],f.clients['gateway']['authorization'])
        self.args = dict(binding=f.binding,base_url=f.base,transport=f.transport,
            history=AuthorityHistory(f.path/'gateway-history.sqlite3'),credential=credential,clock=f.now,
            reporting=GatewayReporter(f.clients['gateway']['reporter_id'],f.clients['gateway']['key_id'],f.gateway_key,credential))
        self.gateway = FreeGateway(f.path/'gateway.sqlite3',**self.args)

    def queue(self, *, complete=True, method='GET', status=200):
        f=self.f; f.agree(); token=f.token()
        rid=self.gateway.admit(token,resource_url=f.request['url'],method=method,actions=['retrieve'],purposes=['public_retrieval'])
        self.gateway.begin_delivery(rid)
        self.gateway.finish_delivery(rid,bytes_written=0 if method=='HEAD' or status==304 else 3,
            representation_digest=crypto.digest(b'abc'),complete=complete,status=status,
            representation_metadata=dict(media_type='text/plain',media_parameters={},languages=[]))
        return rid

    def test_real_two_origin_publication_exact_ack_and_restart(self):
        rid=self.queue(); queued=self.gateway.report_records()[0]
        self.assertEqual(self.gateway.publish_reports(),[dict(request_id=rid,state='acknowledged')])
        row=self.gateway.report_records()[0]
        raw,_=crypto.verify_jws(bytes(row['acknowledgement']),crypto.public_jwk(self.f.signing_key),'odexa-event-record+jws')
        intake=crypto.strict_json(raw)
        self.assertEqual(intake['reporter_jws'].encode(),bytes(queued['jws']))
        self.assertEqual(intake['assurance'],'origin_key_verified')
        fresh=FreeGateway(self.f.path/'gateway.sqlite3',**self.args)
        self.assertEqual(fresh.publish_reports(),[]); self.assertEqual(fresh.report_records(),self.gateway.report_records())
        for _,path,headers,_ in self.f.origin_server.requests:
            if path in {'/odexa.json','/odexa-service.json'}: self.assertNotIn('Authorization',headers)
        with self.f.store.connect() as db: self.assertEqual(db.execute('SELECT count(*) FROM records').fetchone()[0],1)

    def test_direct_same_origin_publication_without_provider(self):
        self.f.temp.cleanup(); self.f.temp=tempfile.TemporaryDirectory(); self.f.path=Path(self.f.temp.name)
        self.f.make_service(external=False); self.configure(); self.queue(method='HEAD')
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'acknowledged')
        self.assertFalse(self.f.provider_server.requests)

    def test_lost_ack_exact_retry_after_restart_retains_original_receipt(self):
        self.queue(complete=False); original=self.f.transport.request; calls=[]; first_ack=[]
        def lose(origin,path,**kwargs):
            result=original(origin,path,**kwargs)
            if path.endswith('/events'):
                calls.append(kwargs['body'])
                if len(calls)==1:
                    first_ack.append(result.body); raise OSError('synthetic lost response')
            return result
        self.f.transport.request=lose
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'queued')
        self.gateway=FreeGateway(self.f.path/'gateway.sqlite3',**self.args)
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'acknowledged')
        row=self.gateway.report_records()[0]
        self.assertEqual(calls[0],calls[1]); self.assertEqual(bytes(row['acknowledgement']),first_ack[0])
        self.assertEqual(row['attempts'],2)
        with self.f.store.connect() as db: self.assertEqual(db.execute('SELECT count(*) FROM records').fetchone()[0],1)

    def test_withdrawal_or_outage_never_sends_report_credential(self):
        self.queue(); before=len(self.f.provider_server.requests)
        self.f.origin_status=503
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'queued')
        self.assertEqual(len(self.f.provider_server.requests),before)
        self.f.origin_status=200; self.f.authority.update(revision=2,delegations=[])
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'queued')
        self.assertEqual(len(self.f.provider_server.requests),before)

    def test_clock_delay_before_send_blocks_disclosure(self):
        self.queue(); old=self.args['history'].require_current; count=0
        def slow(snapshot):
            nonlocal count
            result=old(snapshot);count+=1
            if count==2:self.f.offset=6
            return result
        self.args['history'].require_current=slow;before=len(self.f.provider_server.requests)
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'queued')
        self.assertEqual(len(self.f.provider_server.requests),before)

    def test_withdrawal_after_collector_commit_keeps_original_event_queued(self):
        self.queue(); original=self.f.provider_server.route
        def withdraw(*args):
            result=original(*args)
            if args[1].endswith('/events'):self.f.authority.update(revision=2,delegations=[])
            return result
        self.f.provider_server.route=withdraw
        queued=self.gateway.report_records()[0]
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'queued')
        self.assertEqual(self.gateway.report_records()[0]['jws'],queued['jws'])
        self.assertIsNone(self.gateway.report_records()[0]['acknowledgement'])
        with self.f.store.connect() as db:self.assertEqual(db.execute('SELECT count(*) FROM records').fetchone()[0],1)

    def test_signed_wrong_report_or_bad_signature_never_acknowledges(self):
        self.queue(); original=self.f.transport.request
        def tamper(origin,path,**kwargs):
            result=original(origin,path,**kwargs)
            if path.endswith('/events'):
                raw,header=crypto.verify_jws(result.body,crypto.public_jwk(self.f.signing_key),'odexa-event-record+jws')
                intake=crypto.strict_json(raw);intake['authenticated_reporter_id']=self.f.origin+'/wrong'
                wrong=crypto.sign_jws(crypto.json_bytes(intake),self.f.signing_key,header['kid'],'odexa-event-record+jws').encode()
                return HTTPSResult(result.status,result.headers,wrong,result.peer_ip)
            return result
        self.f.transport.request=tamper
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'queued')
        self.assertIsNone(self.gateway.report_records()[0]['acknowledgement'])
        def bad_signature(origin,path,**kwargs):
            result=original(origin,path,**kwargs)
            if path.endswith('/events'):
                raw,header=crypto.verify_jws(result.body,crypto.public_jwk(self.f.signing_key),'odexa-event-record+jws')
                wrong=crypto.sign_jws(raw,ec.derive_private_key(99,ec.SECP256R1()),header['kid'],'odexa-event-record+jws').encode()
                return HTTPSResult(result.status,result.headers,wrong,result.peer_ip)
            return result
        self.f.transport.request=bad_signature
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'queued')
        self.f.transport.request=original
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'acknowledged')

    def test_selected_endpoint_change_is_not_followed_silently(self):
        self.queue(); before=len(self.f.provider_server.requests)
        self.f.authority['revision']=2;self.f.authority['services'][0]['base_url']=self.f.provider+'/other-tenant/'
        self.assertEqual(self.gateway.publish_reports()[0]['state'],'queued')
        self.assertEqual(len(self.f.provider_server.requests),before)


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