import copy
import concurrent.futures
import sqlite3
import tempfile
import threading
import unittest

from cryptography.hazmat.primitives.asymmetric import ec
from odexa_ref import contracts as c, crypto
from profiles import delegation as d, payment_adapter as adapter, payments as p
from profiles.payment_ledger import PaymentLedger
import test_profile_integration as integration
from test_payment_profile import NOW, U


class PaymentAdapterTests(unittest.TestCase):
    def setUp(self):
        self.fixture = integration.ProfileIntegrationTests()
        self.fixture.setUp()
        self.f = self.fixture
        self.temp = tempfile.TemporaryDirectory()
        self.ledger = PaymentLedger(self.temp.name)
        self.ledger.start(c.json_bytes(self.f.q), c.json_bytes(self.f.a), now=NOW)
        self.ledger.authorize(self.f.a['agreement_id'], c.json_bytes(self.f.m),
                               authenticated_payer_id=self.f.q['payer_id'], now=NOW)
        self.key = ec.derive_private_key(1, ec.SECP256R1())

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

    def observation(self):
        origin = self.f.q['origin']
        return dict(source='https', origin=origin, authority_url=origin+'/odexa-service.json',
            policy_url=origin+'/odexa.json', authority_digest=d.digest(c.json_bytes(self.f.authority)),
            policy_digest=d.digest(c.json_bytes(self.f.policy)), checked_at=NOW,
            available=True, tls_verified=True, redirected=False, revalidated=True)

    def sign(self, value=None, key=None, typ='odexa-payment-status+jws', kid=None):
        return crypto.sign_jws(c.json_bytes(self.f.v if value is None else value), key or self.key,
            kid or self.f.operation['key_id'], typ, allowed_types=crypto.TYPES | {'odexa-payment-status+jws'})

    def commit(self, compact=None, observation=None):
        return adapter.commit_payment_response(self.ledger, self.f.a['agreement_id'], compact or self.sign(),
            c.json_bytes(self.f.authority), c.json_bytes(self.f.policy), self.f.operation,
            now=NOW, observation=observation or self.observation())

    def test_valid_signature_current_authority_and_exact_settlement_commit_atomically(self):
        compact = self.sign()
        state = self.commit(compact)
        self.assertEqual((state.payment_state, state.access_state), ('confirmed', 'active'))
        context = c.load_json(self.ledger.records(self.f.a['agreement_id'])['audit'][-1]['context'])
        self.assertEqual(context['response_jws'], compact)
        self.assertEqual(self.commit(compact), state)
        self.assertEqual(len(self.ledger.records(self.f.a['agreement_id'])['audit']), 3)

    def test_wrong_signer_key_role_and_exact_bytes_fail_without_commit(self):
        compact = self.sign()
        header, body, signature = compact.split('.')
        bad_body = crypto.b64u(c.json_bytes(dict(self.f.v, state='reversed')))
        candidates = [self.sign(key=ec.derive_private_key(2, ec.SECP256R1())),
                      self.sign(kid='https://provider.example/keys/attacker'),
                      self.sign(typ='odexa-status+jws'), '.'.join([header, bad_body, signature])]
        for candidate in candidates:
            with self.subTest(candidate=candidate[:35]):
                with self.assertRaises((crypto.CryptoError, p.PaymentError)): self.commit(candidate)
        self.assertEqual(len(self.ledger.records(self.f.a['agreement_id'])['audit']), 2)

    def test_signed_wrong_amount_or_check_is_not_authorized_settlement(self):
        for response in [dict(self.f.v, total_minor='126'), dict(self.f.v, check_id=U(99))]:
            with self.assertRaises(p.PaymentError): self.commit(self.sign(response))
        self.assertEqual(self.ledger.get(self.f.a['agreement_id'], now=NOW).payment_state, 'pending')

    def test_removed_delegation_and_revoked_key_fail_even_with_good_signature(self):
        compact = self.sign()
        saved = copy.deepcopy(self.f.authority)
        self.f.authority['delegations'] = []
        with self.assertRaises(p.PaymentError): self.commit(compact)
        self.f.authority = saved
        self.f.authority['services'][0]['signing_keys'][0].update(state='revoked', revoked_at=NOW)
        with self.assertRaises(p.PaymentError): self.commit(compact)

    def test_unavailable_unverified_or_wrong_exact_origin_documents_fail(self):
        for field, value in [('available', False), ('tls_verified', False), ('authority_digest', 'sha256:'+'0'*64)]:
            observation = self.observation(); observation[field] = value
            with self.assertRaises(p.PaymentError): self.commit(observation=observation)

    def test_payment_role_does_not_expand_default_free_wire_vocabulary(self):
        with self.assertRaises(crypto.CryptoError):
            crypto.sign_jws(c.json_bytes(self.f.v), self.key, self.f.operation['key_id'], 'odexa-payment-status+jws')
        with self.assertRaises(crypto.CryptoError): crypto.split_jws(self.sign())

    def delayed_commit(self, later):
        clock_value = [NOW]
        first_clock_read = threading.Event()
        def clock():
            value = clock_value[0]
            first_clock_read.set()
            return value
        lock = sqlite3.connect(self.ledger.path, isolation_level=None)
        lock.execute('BEGIN IMMEDIATE')
        def commit():
            return adapter.commit_payment_response(self.ledger, self.f.a['agreement_id'], self.sign(),
                c.json_bytes(self.f.authority), c.json_bytes(self.f.policy), self.f.operation,
                now=clock, observation=self.observation())
        with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor:
            future = executor.submit(commit)
            try:
                self.assertTrue(first_clock_read.wait(3))
                clock_value[0] = later
            finally:
                lock.rollback(); lock.close()
            return future.result(timeout=5)

    def test_write_lock_wait_rechecks_stale_authority(self):
        with self.assertRaises(p.PaymentError): self.delayed_commit('2026-09-17T00:02:06Z')
        self.assertEqual(len(self.ledger.records(self.f.a['agreement_id'])['audit']), 2)

    def test_write_lock_wait_rechecks_key_expiry_within_freshness_window(self):
        self.f.authority['services'][0]['signing_keys'][0]['not_after'] = '2026-09-17T00:02:01Z'
        with self.assertRaises(p.PaymentError): self.delayed_commit('2026-09-17T00:02:02Z')
        self.assertEqual(len(self.ledger.records(self.f.a['agreement_id'])['audit']), 2)

    def test_write_lock_wait_applies_access_expiry_before_confirmation(self):
        self.f.a['access_expires_at'] = '2026-09-17T00:02:01Z'
        self.ledger = PaymentLedger(self.temp.name+'/short-window')
        self.ledger.start(c.json_bytes(self.f.q), c.json_bytes(self.f.a), now=NOW)
        self.ledger.authorize(self.f.a['agreement_id'], c.json_bytes(self.f.m),
                               authenticated_payer_id=self.f.q['payer_id'], now=NOW)
        state = self.delayed_commit('2026-09-17T00:02:02Z')
        self.assertEqual((state.payment_state, state.access_state), ('confirmed', 'expired'))
        self.assertEqual(self.ledger.records(self.f.a['agreement_id'])['audit'][-1]['recorded_at'], '2026-09-17T00:02:02Z')


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