"""Paid archive replay against actual TLS agreement/payer/verifier/report flows."""
import copy
import json
import sqlite3
import sys
import unittest
from unittest import mock

from odexa_ref import crypto
from profiles import portable_evidence as p, paid_evidence as paid, asset_evidence as assets
from profiles import paid_contracts as pc, storage_contracts as sc
from profiles.authority_transport import ScopedCredential
from profiles.authority_history import AuthorityHistory
from profiles.asset_catalog import AssetCatalog
from profiles.free_gateway import FreeGateway
from profiles.gateway_reports import GatewayReporter
from profiles.free_service import uid, later
import test_paid_service as fixture
from test_storage_service import configure, agree as storage_agree, action, session, signed
from test_storage_evidence import enable_export
from test_asset_runtime import publish, asset_trust
from test_asset_evidence import ref


class PaidEvidenceTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls): fixture.PaidServiceTests.setUpClass()
    @classmethod
    def tearDownClass(cls): fixture.PaidServiceTests.tearDownClass()
    def setUp(self):
        self.t = fixture.PaidServiceTests(); self.t.setUp(); self.f = self.t.f
        enable_export(self.f)
    def tearDown(self): self.t.tearDown()

    def agree(self, *, access=120, storage=False):
        t, f = self.t, self.f
        if storage:
            configure(f, t); storage_agree(f, t)
        else:
            _, f.offer_raw = f.call('offers', dict(f.binding, principal_id=f.principal, payer_id=t.quote_payer,
                request=f.request, access_seconds=access, use_seconds=240), wire=True, status=201)
            f.offer_doc = crypto.strict_json(f.offer_raw); t.accept_jws = f.acceptance()
            _, f.receipt_jws = f.call('agreements', t.accept_jws, wire=True, status=201)
            f.receipt = crypto.strict_json(crypto.verify_jws(f.receipt_jws, crypto.public_jwk(f.signing_key), 'odexa-receipt+jws')[0])
            f.agreement_id = f.receipt['agreement_id']
            pc.validate_receipt(f.receipt, offer_bytes=f.offer_raw, acceptance_bytes=crypto.split_jws(t.accept_jws)[3])
        with f.store.connect() as db:
            t.quote_raw = bytes(db.execute('SELECT body FROM documents WHERE digest=?', (f.offer_doc['payment']['quote_digest'],)).fetchone()[0])
        t.q = (sc.PAID if storage else pc).bind_quote(f.offer_doc, t.quote_raw)

    def check(self, state=None, check_id=None):
        if state is not None: self.t.provider_state = state
        f = self.f
        return f.call(f'agreements/{f.agreement_id}/payment-check', dict(f.binding,
            agreement_id=f.agreement_id, check_id=check_id or uid()), wire=True, status=200)

    def export(self, *, storage=False, asset_bundle=None, atrust=None):
        f = self.f; snap = f.transport.fetch(f.origin)
        operation = dict(operation_id=uid(), payment_request_digest=None, kind='export_evidence',
            **{k:f.binding[k] for k in ('origin','service_id','issuer','delegation_id')},
            endpoint=f.base+'exports', key_id=f.service.key_id, key_use=p.TYPE, request=copy.deepcopy(f.request),
            offer_seconds=0, access_seconds=0, use_seconds=0, payment_mode='none', asset_id=None,
            version_id=None, agreement_id=f.agreement_id)
        pins = {p.digest(x):x for x in (snap.authority_bytes,snap.policy_bytes)}
        for observed in f.history.observations(f.origin, limit=1000):
            row = f.history.history(f.origin, observed['authority_revision'], observed['policy_revision'])
            for name in ('authority_bytes','policy_bytes'):
                value = bytes(row[name]); pins[p.digest(value)] = value
        with f.store.connect() as db:
            clients = {r['key_id']:dict(p._key_descriptor(r),revoked=False) for r in db.execute('SELECT * FROM clients WHERE jwk IS NOT NULL')}
            payers = {r['key_id']:dict(payer_client_id=r['client_id'],payer_id=r['payer_id'],key_id=r['key_id'],
                public_jwk=p.load(bytes(r['jwk'])),revoked=False) for r in db.execute('SELECT * FROM payers')}
        trust = paid.PaidExternalTrust(snap,pins,clients,payers,atrust)
        extra = dict(assets=True,asset_bundle_bytes=asset_bundle,asset_trust=atrust) if asset_bundle is not None else {}
        body = paid.build_bundle(f.store,storage=storage,agreement_id=f.agreement_id,snapshot=snap,history=f.history,
            export_operation=operation,private_key=f.signing_key,key_id=f.service.key_id,now=f.now,**extra)
        return body, trust

    def verify(self, body=None, trust=None, **selection):
        if body is None: body,trust = self.export(**selection)
        return paid.verify_bundle(body,trust=trust,now=self.f.now,**selection)

    def edit(self, body, callback):
        envelope = p.load(body); manifest = p.load(crypto.split_jws(envelope['manifest_jws'],allowed_types={p.TYPE})[3])
        blobs = {e['digest']:crypto.unb64u(e['body_b64url']) for e in envelope['blobs']}
        items = [(e,blobs[e['digest']]) for e in manifest['inventory']]
        callback(items,manifest)
        new_blobs = {}
        for entry,raw in items:
            entry.update(digest=p.digest(raw),size=len(raw));new_blobs[entry['digest']] = raw
        manifest['inventory'] = [e for e,_ in items]
        envelope['blobs'] = [dict(digest=sha,body_b64url=crypto.b64u(raw)) for sha,raw in sorted(new_blobs.items())]
        envelope['manifest_jws'] = crypto.sign_jws(p.raw(manifest),self.f.signing_key,self.f.service.key_id,p.TYPE,allowed_types={p.TYPE})
        return p.raw(envelope)

    def alter(self, body, kind, change, *, select=lambda e,v:True):
        def edit(items,manifest):
            index = next(i for i,(entry,raw) in enumerate(items) if entry['kind']==kind and select(entry,p.load(raw)))
            entry,raw = items[index]; value = p.load(raw); change(value)
            items[index] = entry,p.raw(value)
        return self.edit(body,edit)

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

    def test_pending_without_mandate_import_restart_and_free_profile_is_closed(self):
        self.agree();body,trust = self.export(); report=self.verify(body,trust)
        self.assertEqual(report['payment']['payment_state'],'pending')
        self.assertFalse(report['payment']['mandate_authenticated'])
        self.assertEqual(report['lifecycle']['indexed_state'],'pending_payment')
        path=self.f.path/'archive'/'evidence.sqlite3'
        first=paid.PaidEvidenceArchive(path).import_bundle(body,trust=trust,now=self.f.now)
        again=paid.PaidEvidenceArchive(path).import_bundle(body,trust=trust,now=self.f.now)
        self.assertEqual(first,again)
        self.assertEqual(paid.PaidEvidenceArchive(path).get(first['origin'],first['bundle_id'])['bundle_bytes'],body)
        self.assertFalse(first['reactivates_access']);self.assertFalse(first['payment']['bank_settlement_verified'])
        with self.assertRaises(p.PortabilityError):p.verify_bundle(body,trust=trust,now=self.f.now)
        with self.assertRaises(p.PortabilityError):paid.verify_bundle(body,trust=trust,storage=True,now=self.f.now)

    def test_separate_payer_mandate_retained_pending_is_not_a_payment(self):
        self.agree();self.t.mandate();report=self.verify()
        self.assertTrue(report['payment']['mandate_authenticated']);self.assertEqual(report['payment']['verification_count'],0)
        self.assertEqual(report['lifecycle']['indexed_status_version'],1)

    def test_failed_pending_confirmed_exact_retry_and_original_receipt(self):
        self.agree();self.t.mandate();self.check('failed');self.check('pending'); check=uid();self.check('confirmed',check)
        self.check(check_id=check)
        _,status=self.f.call(f'agreements/{self.f.agreement_id}/status',method='GET',wire=True,status=200)
        report=self.verify();self.assertEqual(report['payment']['payment_state'],'confirmed')
        self.assertEqual(report['payment']['verification_count'],4)
        self.assertEqual(report['lifecycle']['indexed_status_version'],2)
        self.assertEqual(self.f.receipt['state_at_issue'],'pending_payment')

    def test_reversal_without_status_does_not_require_invented_signed_status(self):
        self.agree();self.t.mandate();self.check('confirmed');self.check('reversed')
        report=self.verify();self.assertEqual(report['payment']['payment_state'],'reversed')
        self.assertEqual(report['lifecycle']['indexed_state'],'revoked')
        self.assertEqual(report['lifecycle']['indexed_status_version'],3)
        self.assertEqual(report['lifecycle']['signed_status_versions'],[])

    def test_late_confirmation_and_lazy_expiry_never_restore_access(self):
        self.agree(access=2);self.t.mandate();self.f.offset=2
        with mock.patch('profiles.authority_transport.utc_now',side_effect=self.f.now):
            self.check('confirmed');report=self.verify()
        self.assertEqual(report['payment']['payment_state'],'confirmed')
        self.assertEqual(report['payment']['access_at_snapshot'],'expired')
        self.assertEqual(report['lifecycle']['indexed_state'],'expired')

    def test_explicit_revocation_requires_retained_original_signed_status(self):
        self.agree();self.t.mandate();self.check('confirmed'); f=self.f
        f.call(f'agreements/{f.agreement_id}/revoke',dict(f.binding,agreement_id=f.agreement_id,reason='owner requested',idempotency_key=uid()),wire=True,status=200)
        body,trust=self.export();self.assertEqual(self.verify(body,trust)['lifecycle']['indexed_state'],'revoked')
        def remove(items,m):
            items[:]=[x for x in items if x[0]['kind']!='status_jws'];m['cut'].update(status_history=[],status_attestations_count=0)
        with self.assertRaises(p.PortabilityError):self.verify(self.edit(body,remove),trust)

    def test_actual_paid_delivery_report_and_admission_need_confirmed_timeline(self):
        self.reporting_gateway();self.agree();self.t.mandate();self.t.provider_state='confirmed'
        self.assertEqual(self.t.resource(self.f.token(wire=True)),self.t.asset)
        self.assertEqual(self.t.gateway.publish_reports()[0]['state'],'acknowledged')
        body,trust=self.export();self.assertEqual(self.verify(body,trust)['manifest']['cut']['records_count'],1)
        def change(v):
            response=p.load(crypto.unb64u(v['response_b64url']));response['admission']['status_version']=1
            v['response_b64url']=crypto.b64u(p.raw(response))
        with self.assertRaises(p.PortabilityError):self.verify(self.alter(body,'admission',change),trust)

    def test_independent_payer_pin_is_required_and_known_compromise_is_not_hidden(self):
        self.agree();self.t.mandate();self.check('confirmed');body,trust=self.export()
        for pins in ({},{k:dict(v,revoked=True) for k,v in trust.payer_keys.items()}):
            bad=paid.PaidExternalTrust(trust.current_snapshot,trust.pinned_documents,trust.client_keys,pins)
            with self.assertRaises(p.PortabilityError):self.verify(body,bad)
        def corrupt(v):v['payload_b64url']=crypto.b64u(p.raw(dict(self.t.mandate_doc,total_minor='126')))
        with self.assertRaises(p.PortabilityError):self.verify(self.alter(body,'payer_proof',corrupt),trust)

    def test_audit_missing_reordered_changed_state_or_cut_fails_after_exporter_resigns(self):
        self.agree();self.t.mandate();self.check('failed');self.check('confirmed');body,trust=self.export()
        variants=[self.alter(body,'payment_state',lambda v:v.update(access_state='pending_payment')),
                  self.alter(body,'payment_audit',lambda v:v.update(state_digest='sha256:'+'0'*64)),
                  self.alter(body,'payment_audit',lambda v:v.update(recorded_at=later(self.f.now(),1))),
                  self.alter(body,'payment_audit',lambda v:v.update(kind='accepted'),select=lambda e,v:v['kind']=='verification'),
                  self.alter(body,'agreement',lambda v:v.update(status_version=99))]
        def remove(items,m):
            index=next(i for i,(e,raw) in enumerate(items) if e['kind']=='payment_audit' and p.load(raw)['kind']=='verification')
            items.pop(index);m['cut']['payment_audit_count']-=1
        variants.append(self.edit(body,remove))
        def reorder(items,m):
            selected=[i for i,(e,raw) in enumerate(items) if e['kind']=='payment_audit' and p.load(raw)['kind']=='verification']
            i,j=selected;one=p.load(items[i][1]);two=p.load(items[j][1])
            one['sequence'],two['sequence']=two['sequence'],one['sequence']
            for index,value in ((i,one),(j,two)):
                entry=items[index][0];entry['id']=str(value['sequence']);items[index]=entry,p.raw(value)
        variants.append(self.edit(body,reorder))
        variants.append(self.edit(body,lambda items,m:m['cut'].update(payment_audit_count=0)))
        for variant in variants:
            with self.subTest(digest=p.digest(variant)),self.assertRaises(p.PortabilityError):self.verify(variant,trust)

    def test_signed_provider_proof_and_authority_are_both_required(self):
        self.agree();self.t.mandate();self.check('confirmed');body,trust=self.export()
        def corrupt(v):
            context=p.load(crypto.unb64u(v['context_b64url']));context['response_jws']=self.t.mandate_jws.decode()
            v['context_b64url']=crypto.b64u(p.raw(context))
        bad=self.alter(body,'payment_audit',corrupt,select=lambda e,v:v['kind']=='verification')
        with self.assertRaises(p.PortabilityError):self.verify(bad,trust)
        self.f.authority['revision']=2;self.t.verifier['signing_keys'][0].update(state='retired',retired_at=self.f.now())
        body,trust=self.export();self.assertEqual(self.verify(body,trust)['payment']['payment_state'],'confirmed')
        self.f.authority['revision']=3;self.t.verifier['signing_keys'][0].update(state='revoked',revoked_at=self.f.now())
        body,trust=self.export()
        with self.assertRaises(p.PortabilityError):self.verify(body,trust)

    def test_valid_provider_signature_cannot_change_amount_or_authority_request(self):
        self.agree();self.t.mandate();self.check('confirmed');body,trust=self.export()
        def amount(v):
            value=p.load(crypto.unb64u(v['input_b64url']));value['total_minor']='126';payload=p.raw(value)
            context=p.load(crypto.unb64u(v['context_b64url']))
            context['response_jws']=crypto.sign_jws(payload,self.t.payment_key,self.t.payment_kid,'odexa-payment-status+jws',allowed_types={'odexa-payment-status+jws'})
            v.update(input_b64url=crypto.b64u(payload),context_b64url=crypto.b64u(p.raw(context)))
        def scope(v):
            context=p.load(crypto.unb64u(v['context_b64url']));context['request_digest']='sha256:'+'0'*64
            v['context_b64url']=crypto.b64u(p.raw(context))
        for change in (amount,scope):
            with self.subTest(change=change.__name__),self.assertRaises(p.PortabilityError):
                self.verify(self.alter(body,'payment_audit',change,select=lambda e,v:v['kind']=='verification'),trust)

    def test_pending_reversal_is_terminal_version_two(self):
        self.agree(access=2);self.t.mandate();self.check('reversed')
        report=self.verify();self.assertEqual(report['lifecycle']['indexed_status_version'],2)
        self.assertEqual(report['lifecycle']['indexed_state'],'revoked')
        body,trust=self.export()
        for change in (dict(state='active',reason='payment_confirmed'),dict(state='pending_payment',status_version=1,reason='accepted')):
            with self.subTest(change=change),self.assertRaises(p.PortabilityError):
                self.verify(self.alter(body,'agreement',lambda v:v.update(change)),trust)

    def test_pending_agreement_can_expire_without_payer_or_provider_call(self):
        self.agree(access=2);self.f.offset=2
        with mock.patch('profiles.authority_transport.utc_now',side_effect=self.f.now):
            self.f.call(f'agreements/{self.f.agreement_id}/status',method='GET',wire=True,status=200)
            report=self.verify()
        self.assertEqual(report['payment']['payment_state'],'pending');self.assertEqual(report['payment']['verification_count'],0)
        self.assertEqual(report['lifecycle']['indexed_status_version'],2);self.assertEqual(self.t.verify_calls,0)

    def test_source_restart_preserves_exact_payment_proofs_without_credentials(self):
        self.agree();self.t.mandate();self.t.provider_state='confirmed';token=self.f.token(wire=True)
        before,trust=self.export()
        self.f.store=type(self.f.store)(self.f.store.path,self.f.store.identity)
        after,trust2=self.export();self.verify(after,trust2)
        def parts(body):
            env=p.load(body);m=p.load(crypto.split_jws(env['manifest_jws'],allowed_types={p.TYPE})[3])
            blobs={x['digest']:crypto.unb64u(x['body_b64url']) for x in env['blobs']}
            return {(x['kind'],x['id']):blobs[x['digest']] for x in m['inventory'] if x['kind'] in p.PAYMENT_KINDS}
        self.assertEqual(parts(before),parts(after))
        for blob in p.load(after)['blobs']:
            raw=crypto.unb64u(blob['body_b64url'])
            self.assertNotIn(token.encode(),raw)
            for secret in (b'p'*43,b'a'*43,b'g'*43):self.assertNotIn(secret,raw)
            self.assertNotIn(b'secret_hash',raw);self.assertNotIn(b'PRIVATE KEY',raw)

    def test_selected_paid_storage_has_signed_copy_parent_closure(self):
        self.agree(storage=True);self.t.mandate();self.check('confirmed')
        w=action(self.f);start=session(w)
        for e in (w,start):self.f.call('events',signed(self.f,e),wire=True,status=201)
        body,trust=self.export(storage=True);report=self.verify(body,trust,storage=True)
        self.assertEqual(report['storage_metrics']['sessions_known'],1)
        self.assertEqual(report['payment']['payment_state'],'confirmed')

    def asset_archive(self,storage):
        f=self.f;service=f.authority['services'][0];grant=f.authority['delegations'][0]
        service['capabilities'].append('publish_assets');service['signing_keys'][0]['uses'].append(assets.TYPE)
        for scope in (service['scope'],grant):scope.update(asset_id_prefixes=[f.origin+'/assets/'],version_id_prefixes=[f.origin+'/versions/'])
        grant['capabilities'].append('publish_assets')
        self.agree(storage=storage);self.t.mandate();self.check('confirmed')
        catalog=AssetCatalog(f.path/'catalog.sqlite3');m1=publish(f,catalog,1,b'earlier')
        m2=publish(f,catalog,2,self.t.asset,parent=ref(m1))
        if storage:
            w=action(f);w['asset_ref']=ref(m2);start=session(w)
            for e in (w,start):f.call('events',signed(f,e),wire=True,status=201)
        else:
            self.reporting_gateway();original=self.t.gateway.finish_delivery
            rep=m2['representations'][0]
            def finish(request_id,**kwargs):
                return original(request_id,**kwargs,asset_ref=ref(m2),decoded_bytes=len(self.t.asset),decoded_digest=p.digest(self.t.asset),
                    representation_metadata={k:rep[k] for k in ('media_type','media_parameters','languages')})
            self.t.gateway.finish_delivery=finish
            self.t.resource(f.token(wire=True));self.t.gateway.publish_reports()
        atrust=asset_trust(f);closure=catalog.bundle_for([ref(m2)])
        body,trust=self.export(storage=storage,asset_bundle=closure,atrust=atrust)
        report=self.verify(body,trust,storage=storage,assets=True)
        self.assertTrue(report['asset_correlations']);self.assertFalse(report['downstream_use_verified'])
        archive=paid.PaidEvidenceArchive(f.path/'archive'/'assets.sqlite3',storage=storage,assets=True)
        self.assertFalse(archive.import_bundle(body,trust=trust,now=f.now)['reactivates_access'])
        with self.assertRaises(p.PortabilityError):self.verify(body,trust,storage=storage)
        def missing_parent(v):v['manifests'].pop(0)
        with self.assertRaises(p.PortabilityError):self.verify(self.alter(body,'asset_bundle',missing_parent),trust,storage=storage,assets=True)
        return body

    def test_paid_asset_graph_and_actual_delivery_archive(self):self.asset_archive(False)
    def test_paid_asset_storage_graph_archive(self):self.asset_archive(True)

    def test_expired_archive_lock_verification_leaves_no_import(self):
        self.agree();body,trust=self.export();instant=self.f.now();calls=0
        def clock():
            nonlocal calls
            calls+=1
            return instant if calls<4 else later(instant,6)
        path=self.f.path/'archive'/'late.sqlite3';archive=paid.PaidEvidenceArchive(path)
        with self.assertRaises(p.PortabilityError):archive.import_bundle(body,trust=trust,now=clock)
        with sqlite3.connect(path) as db:self.assertEqual(db.execute('SELECT COUNT(*) FROM archives').fetchone()[0],0)


def structural_cases():
    cases=[];PaidEvidenceTests.setUpClass()
    try:
        for storage,asset,filename in ((False,False,'paid-portable'),(True,False,'storage-paid-portable'),
                                       (False,True,'assets-paid-portable'),(True,True,'assets-storage-paid-portable')):
            t=PaidEvidenceTests();t.setUp()
            try:
                if asset:body=t.asset_archive(storage)
                else:
                    t.agree(storage=storage);t.t.mandate();t.check('confirmed');body,_=t.export(storage=storage)
                envelope=p.load(body);manifest=p.load(crypto.split_jws(envelope['manifest_jws'],allowed_types={p.TYPE})[3])
                blobs={x['digest']:crypto.unb64u(x['body_b64url']) for x in envelope['blobs']}
                values={'bundle':envelope,'manifest':manifest}
                for kind in ('agreement','payment_quote','payment_accepted','payment_state','payer_proof','payment_audit'):
                    entry=next(e for e in manifest['inventory'] if e['kind']==kind)
                    values[kind]=p.load(blobs[entry['digest']])
                for definition,value in values.items():
                    cases.append(dict(name=filename+'_'+definition,schema=filename,definition=definition,value=value,structural_valid=True))
                changes=[('bundle',lambda v:v.update(profile='odexa-free-evidence-snapshot-1')),
                    ('manifest',lambda v:v['cut'].pop('payment_audit_count')),
                    ('payment_state',lambda v:v.update(unknown=True)),
                    ('payment_audit',lambda v:v.update(sequence=-1)),
                    ('payer_proof',lambda v:v.update(secret='not-a-wire-field')),
                    ('agreement',lambda v:v.update(state='confirmed'))]
                for index,(definition,change) in enumerate(changes):
                    value=copy.deepcopy(values[definition]);change(value)
                    cases.append(dict(name=filename+'_invalid_'+str(index),schema=filename,definition=definition,value=value,structural_valid=False))
            finally:t.tearDown()
    finally:PaidEvidenceTests.tearDownClass()
    return cases


if __name__=='__main__':
    if '--cases' in sys.argv:print(json.dumps(structural_cases(),indent=2))
    else:unittest.main()
