"""Two distinct real TLS origins, real signatures and durable local history/state."""
import copy
import datetime as dt
import http.server
import json
from pathlib import Path
import re
import ssl
import tempfile
import threading
import unittest

from cryptography.hazmat.primitives.asymmetric import ec
from odexa_ref import contracts as c, crypto
from odexa_ref.fixture import make_tls
from profiles.authority_history import AuthorityHistory, HistoryError
from profiles.authority_transport import AuthorityTransport, ScopedCredential, TransportError
from profiles.network_payments import NetworkPaymentVerifier
from profiles.payment_ledger import PaymentLedger, utc_now
from profiles import payments as p
import test_profile_integration as integration
from test_payment_profile import NOW, U


class PaymentHandler(http.server.BaseHTTPRequestHandler):
    protocol_version = 'HTTP/1.1'
    def log_message(self, *_): pass
    def do_GET(self): self.respond()
    def do_POST(self): self.respond()
    def respond(self):
        body = self.rfile.read(int(self.headers.get('Content-Length', '0')))
        self.server.requests.append((self.command, self.path, dict(self.headers), body))
        status, media, raw = self.server.route(self.command, self.path, self.headers, body)
        try:
            self.send_response_only(status)
            self.send_header('Content-Type', media)
            self.send_header('Content-Length', str(len(raw)))
            self.send_header('Cache-Control', 'no-store')
            self.send_header('Connection', 'close')
            self.end_headers(); self.wfile.write(raw); self.wfile.flush()
        except OSError: pass
        self.close_connection = True


class NetworkPaymentTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.tls = tempfile.TemporaryDirectory()
        cls.tls_path = Path(cls.tls.name); make_tls(cls.tls_path)
        tls = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
        tls.load_cert_chain(cls.tls_path/'tls-cert.pem', cls.tls_path/'tls-key.pem')
        cls.servers = []
        for _ in range(2):
            server = http.server.ThreadingHTTPServer(('127.0.0.1', 0), PaymentHandler)
            server.daemon_threads = True
            server.socket = tls.wrap_socket(server.socket, server_side=True)
            thread = threading.Thread(target=server.serve_forever, daemon=True); thread.start()
            cls.servers.append((server, thread))
        cls.origin, cls.provider = [f'https://127.0.0.1:{s.server_port}' for s, _ in cls.servers]

    @classmethod
    def tearDownClass(cls):
        for server, thread in cls.servers:
            server.shutdown(); server.server_close(); thread.join(timeout=2)
        cls.tls.cleanup()

    def setUp(self):
        self.origin_server, self.provider_server = [s for s, _ in self.servers]
        for server in (self.origin_server, self.provider_server): server.requests = []
        base = integration.ProfileIntegrationTests(); base.setUp()
        anchor = p.time(utc_now()); old_anchor = p.time(NOW)
        def remap(value):
            if isinstance(value, dict): return {k: remap(v) for k, v in value.items()}
            if isinstance(value, list): return [remap(v) for v in value]
            if isinstance(value, str):
                value = value.replace('https://origin.example', self.origin).replace('https://provider.example', self.provider)
                if re.fullmatch(r'\d{4}-\d\d-\d\dT\d\d:\d\d:\d\dZ', value):
                    return (anchor + (p.time(value)-old_anchor)).strftime('%Y-%m-%dT%H:%M:%SZ')
            return value
        self.q, self.a, self.m, self.v, self.authority, self.policy = map(remap,
            (base.q, base.a, base.m, base.v, base.authority, base.policy))
        self.a['quote_digest'] = c.digest(c.json_bytes(self.q))
        self.m['quote_digest'] = self.a['quote_digest']
        self.v['quote_digest'] = self.a['quote_digest']
        self.v['request_digest'] = c.digest(c.json_bytes(self.a['request']))
        self.signing_key = ec.derive_private_key(1, ec.SECP256R1())
        self.signing_kid = self.authority['services'][0]['signing_keys'][0]['kid']
        self.temp = tempfile.TemporaryDirectory()
        self.ledger = PaymentLedger(Path(self.temp.name)/'payment')
        self.history = AuthorityHistory(Path(self.temp.name)/'authority.sqlite3')
        self.ledger.start(c.json_bytes(self.q), c.json_bytes(self.a))
        self.ledger.authorize(self.a['agreement_id'], c.json_bytes(self.m), authenticated_payer_id=self.q['payer_id'])
        self.transport = AuthorityTransport(ca_file=str(self.tls_path/'ca.pem'), allowed_ips=['127.0.0.1'])
        self.verifier = NetworkPaymentVerifier(self.transport, self.history)
        self.credential = ScopedCredential(self.provider, '/api/', 'Bearer synthetic-provider-only')
        self.authority_status = 200
        self.after_provider = None
        self.response_cache = {}
        self.origin_server.route = self.origin_route
        self.provider_server.route = self.provider_route

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

    def origin_route(self, method, path, headers, body):
        if method == 'GET' and path == '/odexa-service.json':
            return self.authority_status, 'application/json', c.json_bytes(self.authority)
        if method == 'GET' and path == '/odexa.json': return 200, 'application/json', c.json_bytes(self.policy)
        return 404, 'application/json', b'{}'

    def provider_route(self, method, path, headers, body):
        if method != 'POST' or path != '/api/verify': return 404, 'application/json', b'{}'
        if headers.get('Authorization') != self.credential.authorization: return 401, 'application/json', b'{}'
        request = c.load_json(body)
        try: p.validate_check(request, self.q, self.a, self.m, self.a['quote_digest'])
        except p.PaymentError: return 400, 'application/json', b'{}'
        key = request['check_id']
        if key not in self.response_cache:
            response = dict(self.v, check_id=key, checked_at=utc_now())
            raw = c.json_bytes(response)
            compact = crypto.sign_jws(raw, self.signing_key, self.signing_kid, 'odexa-payment-status+jws',
                                      allowed_types={'odexa-payment-status+jws'}).encode()
            self.response_cache[key] = compact
        if self.after_provider: self.after_provider()
        return 200, 'application/jose', self.response_cache[key]

    def check(self, check_id=U(5)):
        return self.verifier.check(self.ledger, self.a['agreement_id'], check_id=check_id,
            delegation_id=self.origin+'/delegations/payment', credential=self.credential)

    def assert_pending(self):
        state = self.ledger.get(self.a['agreement_id'])
        self.assertEqual((state.payment_state, state.access_state), ('pending', 'pending_payment'))
        self.assertEqual(len(self.ledger.records(self.a['agreement_id'])['audit']), 2)

    def test_two_origin_signed_confirmation_retains_exact_proof_and_retries_read_only(self):
        self.assertNotEqual(self.origin, self.provider)
        state = self.check()
        self.assertEqual((state.payment_state, state.access_state), ('confirmed', 'active'))
        self.assertEqual(self.check(), state)
        records = self.ledger.records(self.a['agreement_id'])
        self.assertEqual(len(records['audit']), 3)
        proof = c.load_json(records['audit'][-1]['context'])['response_jws']
        self.assertEqual(proof.encode(), self.response_cache[U(5)])
        self.assertEqual([r[1] for r in self.provider_server.requests], ['/api/verify', '/api/verify'])
        self.assertEqual(len(self.origin_server.requests), 8)
        for method, path, headers, body in self.origin_server.requests:
            self.assertEqual(method, 'GET'); self.assertNotIn('Authorization', headers)
            self.assertNotIn('Cookie', headers); self.assertEqual(body, b'')
        for method, path, headers, body in self.provider_server.requests:
            self.assertEqual(method, 'POST'); self.assertEqual(headers['Authorization'], self.credential.authorization)
            self.assertNotIn('Cookie', headers)

    def test_withdrawal_during_provider_reply_blocks_commit(self):
        def withdraw(): self.authority.update(revision=2, delegations=[])
        self.after_provider = withdraw
        with self.assertRaises(p.PaymentError): self.check()
        self.assertEqual(len(self.provider_server.requests), 1)
        self.assert_pending()

    def test_origin_outage_after_reply_has_no_stale_fallback(self):
        self.after_provider = lambda: setattr(self, 'authority_status', 503)
        with self.assertRaises(TransportError): self.check()
        self.assert_pending()
        self.after_provider = None; self.authority_status = 200
        self.assertEqual(self.check().payment_state, 'confirmed')
        self.assertEqual(len(self.response_cache), 1)

    def test_wrong_tenant_is_denied_before_provider_contact(self):
        self.authority['services'][0]['base_url'] = self.provider+'/another-tenant/'
        with self.assertRaises(p.PaymentError): self.check()
        self.assertFalse(self.provider_server.requests); self.assert_pending()

    def test_wrong_signature_and_signed_wrong_amount_do_not_commit(self):
        self.signing_key = ec.derive_private_key(2, ec.SECP256R1())
        with self.assertRaises(crypto.CryptoError): self.check()
        self.assert_pending()
        self.response_cache.clear(); self.signing_key = ec.derive_private_key(1, ec.SECP256R1())
        self.v['total_minor'] = '126'
        with self.assertRaises(p.PaymentError): self.check()
        self.assert_pending()

    def test_revoked_response_key_blocks_commit(self):
        def revoke():
            self.authority['revision'] = 2
            self.authority['services'][0]['signing_keys'][0].update(state='revoked', revoked_at=utc_now())
        self.after_provider = revoke
        with self.assertRaises(p.PaymentError): self.check()
        self.assert_pending()

    def test_key_rotation_uses_new_authorized_key_and_preserves_old_signed_record(self):
        self.v.update(state='pending', provider_reference=None)
        self.assertEqual(self.check().payment_state, 'pending')
        old_proof = self.ledger.records(self.a['agreement_id'])['audit'][-1]['context']
        key = self.authority['services'][0]['signing_keys'][0]
        key.update(state='retired', retired_at=utc_now())
        self.signing_key = ec.derive_private_key(2, ec.SECP256R1())
        self.signing_kid = self.provider+'/keys/two'
        new_key = dict(key, kid=self.signing_kid, public_jwk=crypto.public_jwk(self.signing_key),
                       state='active', retired_at=None, revoked_at=None)
        self.authority['services'][0]['signing_keys'].append(new_key)
        self.authority['revision'] = 2
        self.v.update(state='confirmed', provider_reference='txn-1', sequence=2)
        self.assertEqual(self.check(U(6)).payment_state, 'confirmed')
        records = self.ledger.records(self.a['agreement_id'])['audit']
        self.assertEqual(records[-2]['context'], old_proof)
        self.assertEqual(c.load_json(records[-1]['context'])['key_id'], self.signing_kid)

    def test_metadata_rollback_after_reply_is_rejected_by_retained_history(self):
        self.authority['revision'] = 2
        self.after_provider = lambda: self.authority.update(revision=1)
        with self.assertRaises(HistoryError): self.check()
        self.assert_pending()

    def test_missing_payer_mandate_never_contacts_provider(self):
        self.ledger = PaymentLedger(Path(self.temp.name)/'unmandated')
        self.ledger.start(c.json_bytes(self.q), c.json_bytes(self.a))
        with self.assertRaises(p.PaymentError): self.check()
        self.assertFalse(self.origin_server.requests); self.assertFalse(self.provider_server.requests)

    def test_pretransmission_wait_cannot_send_context_using_expired_authority(self):
        offset = [0]
        self.verifier.clock = lambda: (p.time(utc_now()) + dt.timedelta(seconds=offset[0])).strftime('%Y-%m-%dT%H:%M:%SZ')
        original = self.history.require_current
        count = [0]
        def delayed_guard(snapshot):
            result = original(snapshot)
            count[0] += 1
            if count[0] == 2: offset[0] = 6
            return result
        self.history.require_current = delayed_guard
        with self.assertRaises(TransportError): self.check()
        self.assertEqual(count[0], 2)
        self.assertFalse(self.provider_server.requests)
        self.assert_pending()


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