"""Actual TLS paid exchange with a separately authenticated payer and verifier."""
import base64
import copy
import json
import http.client
import ssl
import threading
import unittest
from urllib.parse import urlsplit
from cryptography.hazmat.primitives.asymmetric import ec

from odexa_ref import crypto
from profiles import payments as p, paid_contracts as pc
from profiles.authority_transport import ScopedCredential
from profiles.paid_store import PaidStore, TransactionPayments
from profiles.paid_service import PaidAgreementService
from profiles.free_service import uid, later
from profiles.free_gateway import FreeGateway
from profiles.authority_history import AuthorityHistory
import test_free_service as fixture


class PaidServiceTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls): fixture.FreeServiceTests.setUpClass()
    @classmethod
    def tearDownClass(cls): fixture.FreeServiceTests.tearDownClass()

    def setUp(self):
        self.f=fixture.FreeServiceTests(); self.f.setUp(); f=self.f
        self.quote_payer=f.origin+'/payers/one'; self.payer_key=ec.derive_private_key(5,ec.SECP256R1())
        self.payer_kid=f.provider+'/payer-key'; self.payment_key=ec.derive_private_key(6,ec.SECP256R1())
        self.payment_kid=f.provider+'/payment-key'
        service=f.authority['services'][0]; grant=f.authority['delegations'][0]
        service['limits']['payment_mode']=grant['payment_mode']='external'
        verifier=copy.deepcopy(service)
        verifier.update(id=f.origin+'/services/payment',base_url=f.provider+'/payment/',issuer=f.provider+'/payment-operator',capabilities=['verify_payments'],limits=None)
        verifier['signing_keys']=[dict(service['signing_keys'][0],kid=self.payment_kid,public_jwk=crypto.public_jwk(self.payment_key),uses=['odexa-payment-status+jws'])]
        self.verifier_id=verifier['id']; self.verifier=verifier
        f.authority['services'].append(verifier)
        vg=copy.deepcopy(grant); vg.update(id=f.origin+'/delegations/payment',delegate_service_id=verifier['id'],capabilities=['verify_payments'],max_access_seconds=0,max_use_seconds=0)
        f.authority['delegations'].append(vg); self.verifier_grant=vg
        f.store=PaidStore(f.path/'paid.sqlite3',{k:f.binding[k] for k in ('origin','service_id','issuer')})
        self.verify_credential=ScopedCredential(f.provider,'/payment/','Bearer synthetic-read-only-verifier')
        f.service=PaidAgreementService(f.store,**f.service_args,quote_factory=self.quote,
            verifier_delegation_id=vg['id'],verifier_credential=self.verify_credential)
        for role,key in [('agent',f.agent_key),('gateway',f.gateway_key),('admin',None)]:
            f.service.provision_client(role,role[0]*43,[f.principal],role=role,public_jwk=crypto.public_jwk(key) if key else None)
        f.service.provision_payer('payer','p'*43,self.quote_payer,key_id=self.payer_kid,public_jwk=crypto.public_jwk(self.payer_key))
        self.payer_auth='Basic '+base64.b64encode(b'payer:'+b'p'*43).decode()
        self.provider_state='pending'; self.sequence=0; self.responses={}; self.verify_calls=0; self.after_verify=None
        self.wrong_amount=False; self.wrong_signature=False
        f.provider_server.route=self.provider_route
        self.gateway=FreeGateway(f.path/'paid-gateway.sqlite3',binding=f.binding,base_url=f.base,transport=f.transport,
            history=AuthorityHistory(f.path/'paid-gateway-authority.sqlite3'),clock=f.now,
            credential=ScopedCredential(f.provider,f.base[len(f.provider):],f.clients['gateway']['authorization']))
        self.original_origin=f.origin_route; f.origin_server.route=self.origin_route
        self.asset=b'{"content":"Licensed originator content."}'
        self.delivery_finished=threading.Event(); self.delivery_error=None

    def tearDown(self): self.f.tearDown()

    def quote(self, offer, payer_id):
        f=self.f
        return crypto.json_bytes(dict(protocol_version=p.VERSION,profile=p.PROFILE,record_type='payment.quote',quote_id=uid(),
            origin=f.origin,service_id=f.binding['service_id'],verification_service_id=self.verifier_id,
            provider_id=self.verifier['issuer'],verification_endpoint=self.verifier['base_url']+'verify',payer_id=payer_id,
            offer_id=offer['offer_id'],currency='USD',minor_unit_scale=2,total_minor='125',
            items=[dict(item_id='resource',kind='resource_license',payee_id=f.origin+'/payee',amount_minor='100',description='Resource licence'),
                   dict(item_id='service',kind='provider_service',payee_id=f.provider+'/payee',amount_minor='25',description='Payment management')],
            issued_at=offer['issued_at'],expires_at=offer['expires_at']))

    def provider_route(self, method, path, headers, body):
        if method!='POST' or path!='/payment/verify': return self.f.service.dispatch(method,path,headers,body)
        self.verify_calls+=1
        if headers.get('Authorization')!=self.verify_credential.authorization:
            return 401,{'Content-Type':'application/json','Cache-Control':'no-store'},b'{}'
        check=crypto.strict_json(body)
        if check['check_id'] not in self.responses:
            self.sequence+=1
            value=dict(check,record_type='payment.verification',sequence=self.sequence,state=self.provider_state,
                provider_reference='reference-transaction-1' if self.provider_state in {'confirmed','reversed'} else None,
                effective_at=self.f.now(),checked_at=self.f.now())
            if self.wrong_amount: value['total_minor']='126'
            signer=self.payer_key if self.wrong_signature else self.payment_key
            self.responses[check['check_id']]=crypto.sign_jws(crypto.json_bytes(value),signer,self.payment_kid,'odexa-payment-status+jws',allowed_types={'odexa-payment-status+jws'}).encode()
        if self.after_verify: self.after_verify()
        return 200,{'Content-Type':'application/jose','Cache-Control':'no-store'},self.responses[check['check_id']]

    def origin_route(self, method, path, headers, body):
        if method!='GET' or path!='/about': return self.original_origin(method,path,headers,body)
        try:
            auth=headers.get('Authorization','')
            if not auth.startswith('Bearer '): raise ValueError('no token')
            rid=self.gateway.admit(auth[7:],resource_url=self.f.request['url'],actions=['retrieve'],purposes=['public_retrieval'])
        except ValueError: return 403,{'Content-Type':'application/json','Cache-Control':'no-store'},b'{}'
        return 200,{'Content-Type':'application/json','Cache-Control':'no-store'},self.asset,{
            'begin':lambda:self.gateway.begin_delivery(rid),
            'finish':lambda count,complete:self.finish_delivery(rid,count,complete)}

    def finish_delivery(self, rid, count, complete):
        # The HTTP client can read the last byte before this handler commits its
        # delivery outbox. Tests wait for that actual boundary, never a sleep.
        try:
            self.gateway.finish_delivery(rid,bytes_written=count,representation_digest=crypto.digest(self.asset),complete=complete)
        except Exception as error:
            self.delivery_error=error
        finally:
            self.delivery_finished.set()

    def offer(self,access=120):
        f=self.f
        _,f.offer_raw=f.call('offers',dict(f.binding,principal_id=f.principal,payer_id=self.quote_payer,request=f.request,access_seconds=access,use_seconds=240),status=201)
        f.offer_doc=crypto.strict_json(f.offer_raw)
        with f.store.connect() as db: self.quote_raw=bytes(db.execute('SELECT body FROM documents WHERE digest=?',(f.offer_doc['payment']['quote_digest'],)).fetchone()['body'])
        self.q=pc.bind_quote(f.offer_doc,self.quote_raw)

    def agree(self,access=120):
        f=self.f; self.offer(access); self.accept_jws=f.acceptance()
        _,receipt=f.call('agreements',self.accept_jws,wire=True,status=201)
        f.receipt_jws=receipt; f.receipt=crypto.strict_json(crypto.verify_jws(receipt,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(self.accept_jws)[3])
        self.assertEqual(f.receipt['state_at_issue'],'pending_payment')

    def mandate(self, *, payer=None, status=201, wire=True):
        f=self.f
        self.mandate_doc=dict(protocol_version=p.VERSION,profile=p.PROFILE,record_type='payment.mandate',
            mandate_id=uid(),agreement_id=f.agreement_id,quote_id=self.q['quote_id'],quote_digest=crypto.digest(self.quote_raw),
            origin=f.origin,service_id=f.binding['service_id'],provider_id=self.verifier['issuer'],payer_id=payer or self.quote_payer,
            currency='USD',minor_unit_scale=2,total_minor='125',issued_at=f.now(),expires_at=later(f.now(),120),authorization='pay_exact_quote_once')
        self.mandate_jws=crypto.sign_jws(crypto.json_bytes(self.mandate_doc),self.payer_key,self.payer_kid,pc.MANDATE_TYPE,allowed_types={pc.MANDATE_TYPE}).encode()
        route=f.base[len(f.provider):]+f'agreements/{f.agreement_id}/payment-mandate'
        if wire:
            return f.transport.request(f.provider,route,method='POST',body=self.mandate_jws,
                credential=ScopedCredential(f.provider,f.base[len(f.provider):],self.payer_auth),content_type='application/jose',
                allowed_jws_types={pc.MANDATE_TYPE},expected_status=status).body
        code,_,raw=f.service.dispatch('POST',route,{'Authorization':self.payer_auth,'Content-Type':'application/jose'},self.mandate_jws)
        self.assertEqual(code,status,raw); return raw

    def check(self, check_id=None,status=200):
        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()),status=status)

    def resource(self,token,status=200):
        self.delivery_finished.clear(); self.delivery_error=None
        origin=urlsplit(self.f.origin)
        connection=http.client.HTTPSConnection(origin.hostname,origin.port,context=ssl.create_default_context(cafile=str(self.f.tls_path/'ca.pem')),timeout=5)
        try:
            connection.request('GET','/about',headers={'Authorization':'Bearer '+token})
            response=connection.getresponse(); raw=response.read()
            self.assertEqual(response.status,status,raw)
            if status==200:
                self.assertTrue(self.delivery_finished.wait(5),'gateway completion callback did not finish')
                if self.delivery_error is not None: raise self.delivery_error
            return raw
        finally: connection.close()

    def test_assent_alone_and_agent_credential_cannot_authorize_payment(self):
        self.agree(); f=self.f
        f.call('tokens',dict(f.binding,agreement_id=f.agreement_id),status=403)
        self.assertEqual(self.verify_calls,0)
        self.mandate(wire=False)
        route=f.base[len(f.provider):]+f'agreements/{f.agreement_id}/payment-mandate'
        code,_,_=f.service.dispatch('POST',route,{'Authorization':f.clients['agent']['authorization'],'Content-Type':'application/jose'},self.mandate_jws)
        self.assertEqual(code,401)

    def test_actual_paid_pending_confirmed_delivery_reversed_denial(self):
        self.agree(); f=self.f; self.mandate()
        f.call('tokens',dict(f.binding,agreement_id=f.agreement_id),status=403)
        self.assertEqual(self.verify_calls,1)
        self.provider_state='confirmed'; token=f.token(wire=True)
        self.assertEqual(self.resource(token),self.asset)
        self.provider_state='reversed'
        self.resource(token,status=403)
        f.call('tokens',dict(f.binding,agreement_id=f.agreement_id),status=403)
        with f.store.connect() as db:
            row=db.execute('SELECT * FROM agreements').fetchone()
            self.assertEqual(row['state'],'revoked'); self.assertEqual(row['receipt'].encode(),f.receipt_jws)
            self.assertEqual(TransactionPayments(db)._read(row['id']).payment_state,'reversed')
        self.assertEqual(len(self.gateway.records()),1)

    def test_wrong_payer_amount_signature_and_withdrawal_do_not_activate(self):
        self.agree(); self.mandate(payer=self.f.origin+'/payers/other',status=400,wire=False)
        self.mandate(); self.provider_state='confirmed'; self.wrong_amount=True
        self.check(status=400)
        self.wrong_amount=False; self.wrong_signature=True; self.check(status=400)
        self.wrong_signature=False
        def withdraw():
            self.f.authority['revision']=2
            self.f.authority['delegations']=[d for d in self.f.authority['delegations'] if d['id']!=self.verifier_grant['id']]
        self.after_verify=withdraw; self.check(status=400)
        with self.f.store.connect() as db:
            self.assertEqual(db.execute('SELECT state FROM agreements').fetchone()['state'],'pending_payment')
            self.assertEqual(db.execute("SELECT count(*) FROM payment_audit WHERE kind='verification'").fetchone()[0],0)

    def test_exact_verification_retry_never_reactivates_after_reversal(self):
        self.agree(); self.mandate(); self.provider_state='confirmed'; first=uid()
        _,raw=self.check(first); self.assertEqual(crypto.strict_json(raw)['access_state'],'active')
        self.provider_state='reversed'; self.check()
        _,raw=self.check(first); self.assertEqual(crypto.strict_json(raw)['access_state'],'revoked')
        with self.f.store.connect() as db:
            self.assertEqual(db.execute("SELECT count(*) FROM payment_audit WHERE kind='verification'").fetchone()[0],2)

    def test_paid_acceptance_and_payment_state_share_atomic_commit(self):
        f=self.f; self.offer(); acceptance=f.acceptance()
        def crash(stage):
            if stage=='before_commit': raise RuntimeError('simulated precommit interruption')
        f.store.fault_hook=crash
        with self.assertRaises(RuntimeError): f.call('agreements',acceptance)
        f.store.fault_hook=None
        with f.store.connect() as db:
            self.assertEqual(db.execute('SELECT count(*) FROM agreements').fetchone()[0],0)
            self.assertEqual(db.execute('SELECT count(*) FROM payment_states').fetchone()[0],0)
        _,original=f.call('agreements',acceptance,status=201)
        _,retry=f.call('agreements',acceptance,status=200); self.assertEqual(original,retry)

    def test_late_financial_confirmation_cannot_extend_expired_access(self):
        self.agree(access=2); self.mandate(); self.f.offset=2; self.provider_state='confirmed'
        _,raw=self.check(); result=crypto.strict_json(raw)
        self.assertEqual((result['payment_state'],result['access_state']),('confirmed','expired'))
        self.f.call('tokens',dict(self.f.binding,agreement_id=self.f.agreement_id),status=403)

    def test_original_paid_receipt_available_after_origin_outage(self):
        self.agree(); self.f.origin_status=503
        _,receipt=self.f.call('agreements',self.accept_jws,status=200)
        self.assertEqual(receipt,self.f.receipt_jws)
        self.f.call('tokens',dict(self.f.binding,agreement_id=self.f.agreement_id),status=503)


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