"""Free draft-3 lifecycle using real current-authority HTTPS acquisition."""
import base64
import copy
import datetime as dt
import http.server
from pathlib import Path
import ssl
import tempfile
import threading
import unittest

from cryptography.hazmat.primitives.asymmetric import ec
from odexa_ref import crypto
from odexa_ref.fixture import make_tls
from profiles.authority_history import AuthorityHistory
from profiles.authority_transport import AuthorityTransport, ScopedCredential, utc_now
from profiles import free_contracts as c
from profiles.free_store import FreeStore
from profiles.free_service import FreeAgreementService, uid, time, later
from test_delegation_profile import fixture, ORIGIN, PROVIDER, NOW


class Handler(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):
        raw = self.rfile.read(int(self.headers.get('Content-Length', '0')))
        self.server.requests.append((self.command, self.path, dict(self.headers), raw))
        routed = self.server.route(self.command, self.path, dict(self.headers), raw)
        status, headers, body = routed[:3]
        callbacks = routed[3] if len(routed) == 4 else None
        if callbacks:
            try: callbacks['begin']()
            except ValueError:
                status, headers, body, callbacks = 403, {'Content-Type': 'application/json', 'Cache-Control': 'no-store'}, b'{}', None
        written, complete = 0, False
        try:
            self.send_response_only(status)
            for name, value in headers.items(): self.send_header(name, value)
            self.send_header('Content-Length', str(len(body))); self.send_header('Connection', 'close')
            self.end_headers(); written = self.wfile.write(body); self.wfile.flush(); complete = True
        except OSError: pass
        finally:
            if callbacks: callbacks['finish'](written, complete)
        self.close_connection = True


class FreeServiceTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.tls = tempfile.TemporaryDirectory(); cls.tls_path = Path(cls.tls.name); make_tls(cls.tls_path)
        ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
        ctx.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), Handler)
            server.daemon_threads = True; server.socket = ctx.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:{server.server_port}' for server, _ 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.temp = tempfile.TemporaryDirectory(); self.path = Path(self.temp.name)
        self.offset = 0; self.origin_status = 200
        self.make_service()

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

    def now(self): return later(utc_now(), self.offset)

    def make_service(self, external=True, reuse=False):
        authority, pol, _ = fixture(external)
        anchor = time(utc_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(ORIGIN, self.origin).replace(PROVIDER, self.provider)
                if len(value) == 20 and value.endswith('Z'):
                    try: return (anchor + (time(value) - time(NOW))).strftime('%Y-%m-%dT%H:%M:%SZ')
                    except ValueError: pass
            return value
        self.authority, self.policy = remap(authority), remap(pol)
        svc = self.authority['services'][0]
        svc['capabilities'].append('receive_events'); svc['signing_keys'][0]['uses'].append('odexa-event-record+jws')
        if external: self.authority['delegations'][0]['capabilities'].append('receive_events')
        self.binding = dict(protocol_version=c.VERSION, origin=self.origin, service_id=svc['id'], issuer=svc['issuer'],
            delegation_id=self.authority['delegations'][0]['id'] if external else None)
        self.base = svc['base_url']; self.host = self.provider if external else self.origin
        self.transport = AuthorityTransport(ca_file=str(self.tls_path / 'ca.pem'), allowed_ips=['127.0.0.1'])
        self.history = AuthorityHistory(self.path / 'authority.sqlite3')
        self.store = FreeStore(self.path / 'free.sqlite3', {k: self.binding[k] for k in ('origin', 'service_id', 'issuer')})
        self.signing_key = ec.derive_private_key(1, ec.SECP256R1())
        self.service_args = dict(binding=self.binding, base_url=self.base, key_id=svc['signing_keys'][0]['kid'],
            private_key=self.signing_key, terms_bytes=b'Free use subject to the originator policy.\n',
            transport=self.transport, history=self.history, clock=self.now)
        self.service = FreeAgreementService(self.store, **self.service_args)
        self.agent_key = ec.derive_private_key(3, ec.SECP256R1()); self.gateway_key = ec.derive_private_key(4, ec.SECP256R1())
        self.principal = self.origin + '/principals/test'
        self.clients = {}
        for role, key in [('agent', self.agent_key), ('gateway', self.gateway_key), ('admin', None)]:
            secret = role[0] * 43
            registered = self.service.provision_client(role, secret, [self.principal], role=role,
                public_jwk=crypto.public_jwk(key) if key else None)
            registered['authorization'] = 'Basic ' + base64.b64encode((role + ':' + secret).encode()).decode()
            self.clients[role] = registered
        self.request = dict(url=self.origin + '/about', actions=['retrieve'], purposes=['public_retrieval'], supported_obligations=['attribution'])
        self.origin_server, self.provider_server = [s for s, _ in self.servers]
        for server in (self.origin_server, self.provider_server): server.requests = []
        self.origin_server.route = self.origin_route
        self.provider_server.route = self.service.dispatch

    def origin_route(self, method, path, headers, body):
        if method == 'GET' and path in {'/odexa.json', '/odexa-service.json'}:
            data = self.policy if path == '/odexa.json' else self.authority
            return self.origin_status, {'Content-Type': 'application/json', 'Cache-Control': 'no-store'}, crypto.json_bytes(data)
        if self.host == self.origin: return self.service.dispatch(method, path, headers, body)
        return 404, {'Content-Type': 'application/json', 'Cache-Control': 'no-store'}, b'{}'

    def call(self, route, value=None, role='agent', *, method='POST', wire=False, status=None):
        body = value if isinstance(value, bytes) else crypto.json_bytes(value) if value is not None else b''
        media = 'application/jose' if isinstance(value, bytes) else 'application/json'
        path = self.base[len(self.host):] + route
        if wire:
            result = self.transport.request(self.host, path, method=method, body=body,
                credential=ScopedCredential(self.host, self.base[len(self.host):], self.clients[role]['authorization']),
                content_type=media, response_type='application/jose' if route == 'agreements' or route.endswith(('/status', '/receipt', '/revoke')) or route == 'events' else 'application/json',
                expected_status=status or 200, max_bytes=262144)
            return result.status, result.body
        actual, _, raw = self.service.dispatch(method, path, {'Authorization': self.clients[role]['authorization'], 'Content-Type': media}, body)
        if status is not None: self.assertEqual(actual, status, raw)
        return actual, raw

    def offer(self, wire=False):
        _, raw = self.call('offers', dict(self.binding, principal_id=self.principal, request=self.request, access_seconds=120, use_seconds=240), wire=wire, status=201)
        self.offer_raw, self.offer_doc = raw, crypto.strict_json(raw)
        return self.offer_doc

    def acceptance(self, offer=None):
        offer = offer or self.offer_doc
        self.accept_doc = dict(self.binding, offer_id=offer['offer_id'], offer_digest=crypto.digest(self.offer_raw),
            client_id='agent', principal_id=self.principal, intent='accept', nonce=offer['nonce'], accepted_at=self.now(),
            idempotency_key='acceptance-retry-key')
        self.accept_jws = crypto.sign_jws(crypto.json_bytes(self.accept_doc), self.agent_key, self.clients['agent']['key_id'], 'odexa-acceptance+jws').encode()
        return self.accept_jws

    def agree(self, wire=False):
        self.offer(wire); body = self.acceptance()
        _, receipt = self.call('agreements', body, wire=wire, status=201)
        raw, _ = crypto.verify_jws(receipt, crypto.public_jwk(self.signing_key), 'odexa-receipt+jws', self.service.key_id)
        self.receipt_jws, self.receipt = receipt, crypto.strict_json(raw)
        self.agreement_id = self.receipt['agreement_id']
        return self.receipt

    def token(self, wire=False):
        _, raw = self.call('tokens', dict(self.binding, agreement_id=self.agreement_id), wire=wire, status=200)
        self.token_doc = crypto.strict_json(raw); return self.token_doc['access_token']

    def introspect(self, token, wire=False, **overrides):
        request = dict(self.binding, token=token, resource_url=self.request['url'], method='GET',
            actions=self.request['actions'], purposes=self.request['purposes'], request_id=uid())
        request.update(overrides)
        _, raw = self.call('introspect', request, 'gateway', wire=wire, status=200)
        return crypto.strict_json(raw)

    def test_real_two_origin_offer_accept_retry_token_status_revoke(self):
        self.agree(wire=True)
        _, retry = self.call('agreements', self.accept_jws, wire=True, status=200)
        self.assertEqual(retry, self.receipt_jws)
        token = self.token(wire=True)
        result = self.introspect(token, wire=True); self.assertTrue(result['permitted'])
        self.assertEqual(result['admission']['principal_id'], self.principal)
        _, raw = self.call(f'agreements/{self.agreement_id}/status', method='GET', wire=True, status=200)
        status = crypto.strict_json(crypto.verify_jws(raw, crypto.public_jwk(self.signing_key), 'odexa-status+jws')[0])
        self.assertEqual(status['state'], 'active')
        value = dict(self.binding, agreement_id=self.agreement_id, reason='owner requested', idempotency_key='revoke-retry-key-1')
        _, revoked = self.call(f'agreements/{self.agreement_id}/revoke', value, wire=True, status=200)
        self.assertEqual(crypto.strict_json(crypto.verify_jws(revoked, crypto.public_jwk(self.signing_key), 'odexa-status+jws')[0])['state'], 'revoked')
        self.assertFalse(self.introspect(token, wire=True)['permitted'])
        for method, path, headers, body in self.origin_server.requests:
            self.assertNotIn('Authorization', headers); self.assertEqual(body, b'')
        self.assertTrue(self.provider_server.requests)

    def test_provider_absent_same_origin_same_contract(self):
        self.temp.cleanup(); self.temp = tempfile.TemporaryDirectory(); self.path = Path(self.temp.name)
        self.make_service(external=False)
        self.agree(wire=True); token = self.token(wire=True)
        self.assertTrue(self.introspect(token, wire=True)['permitted'])
        self.assertFalse(self.provider_server.requests)
        self.assertIsNone(self.receipt['delegation_id'])

    def test_policy_or_authority_change_prevents_new_acceptance(self):
        self.offer(); body = self.acceptance()
        self.authority['revision'] = 2
        self.call('agreements', body, status=409)
        with self.store.connect() as db: self.assertEqual(db.execute('SELECT count(*) FROM agreements').fetchone()[0], 0)

    def test_withdrawal_blocks_new_access_but_original_receipt_survives(self):
        self.agree(); token = self.token()
        self.authority.update(revision=2, delegations=[])
        self.call('tokens', dict(self.binding, agreement_id=self.agreement_id), status=503)
        _, retry = self.call('agreements', self.accept_jws, status=200)
        self.assertEqual(retry, self.receipt_jws)
        self.origin_status = 503
        _, fetched = self.call(f'agreements/{self.agreement_id}/receipt', method='GET', status=200)
        self.assertEqual(fetched, self.receipt_jws)

    def test_wrong_tenant_issuer_key_principal_and_scope(self):
        data = dict(self.binding, principal_id=self.principal, request=self.request, access_seconds=120, use_seconds=240)
        for field, value in [('issuer', self.provider+'/attacker'), ('principal_id', self.origin+'/other')]:
            self.call('offers', dict(data, **{field: value}), status=403)
        self.authority['services'][0]['base_url'] = self.provider + '/wrong-tenant/'
        self.call('offers', data, status=503)
        self.authority['services'][0]['base_url'] = self.base
        self.service.private_key = self.agent_key; self.service.public_jwk = crypto.public_jwk(self.agent_key)
        self.call('offers', data, status=503)

    def test_scope_expansion_wrong_role_and_duplicate_admission_denied(self):
        self.agree(); token = self.token()
        self.assertFalse(self.introspect(token, purposes=['internal_knowledge'])['permitted'])
        request = dict(self.binding, token=token, resource_url=self.request['url'], method='GET', actions=['retrieve'], purposes=['public_retrieval'], request_id=uid())
        self.call('introspect', request, 'agent', status=403)
        self.call('introspect', request, 'gateway', status=200)
        self.call('introspect', request, 'gateway', status=409)

    def test_origin_outage_and_commit_time_delay_fail_closed(self):
        self.origin_status = 503
        data = dict(self.binding, principal_id=self.principal, request=self.request, access_seconds=120, use_seconds=240)
        self.call('offers', data, status=503)
        self.origin_status = 200
        original = self.history.require_current
        def delay(snapshot):
            result = original(snapshot); self.offset = 6; return result
        self.history.require_current = delay
        self.call('offers', data, status=503)
        with self.store.connect() as db: self.assertEqual(db.execute('SELECT count(*) FROM offers').fetchone()[0], 0)

    def test_exact_acceptance_restart_after_commit_and_conflict(self):
        self.agree()
        self.service = FreeAgreementService(self.store, **self.service_args)
        self.provider_server.route = self.service.dispatch
        _, raw = self.call('agreements', self.accept_jws, status=200)
        self.assertEqual(raw, self.receipt_jws)
        self.accept_doc['nonce'] = crypto.b64u(b'\0'*32)
        changed = crypto.sign_jws(crypto.json_bytes(self.accept_doc), self.agent_key, self.clients['agent']['key_id'], 'odexa-acceptance+jws').encode()
        self.call('agreements', changed, status=409)

    def test_http_auth_scheme_and_media_type_are_case_insensitive(self):
        data = dict(self.binding, principal_id=self.principal, request=self.request, access_seconds=120, use_seconds=240)
        headers = {'Authorization': self.clients['agent']['authorization'].replace('Basic ', 'bAsIc '), 'Content-Type': ' Application/JSON ; charset=utf-8'}
        status, _, raw = self.service.dispatch('POST', self.base[len(self.host):]+'offers', headers, crypto.json_bytes(data))
        self.assertEqual(status, 201, raw)

    def test_every_emitted_status_statement_is_retained_without_false_state_transition(self):
        self.agree()
        route = f'agreements/{self.agreement_id}/status'
        _, first = self.call(route, method='GET', status=200)
        _, second = self.call(route, method='GET', status=200)
        with self.store.connect() as db:
            rows = db.execute('SELECT jws,version FROM status_attestations').fetchall()
            self.assertEqual({r['jws'].encode() for r in rows}, {first, second})
            self.assertEqual({r['version'] for r in rows}, {1})
            self.assertEqual(db.execute('SELECT count(*) FROM status_history').fetchone()[0], 1)

    def test_gateway_report_is_bound_to_admission_and_retained_on_exact_retry(self):
        self.agree(); token = self.token(); response = self.introspect(token)
        admitted = response['admission']; now = self.now()
        event = dict(self.binding, event_id=uid(), event_type='delivery.completed', occurred_at=now,
            reporter_id=self.clients['gateway']['reporter_id'], source='origin_observed', resource_url=self.request['url'],
            policy_id=admitted['policy_id'], policy_revision=admitted['policy_revision'], actions=['retrieve'], purposes=['public_retrieval'],
            agreement_id=self.agreement_id, asset_ref=None, operation=None, quantity=None, unit=None, related_events=[], derived_from=[],
            http=dict(method='GET', status=200, kind='full', delivery_id=response['request_id'], hop_id=uid(),
                ingress_id=self.origin+'/ingress/one', hop_role='end_client', boundary_id=self.origin+'/boundaries/one',
                cache_status='bypass', content_codings=[], content_bytes=4, content_digest=crypto.digest(b'test'),
                decoded_bytes=4, decoded_digest=crypto.digest(b'test'), range=None,
                representation_metadata={'media_type':'text/plain','media_parameters':{},'languages':[]}))
        def sign(value):
            return crypto.sign_jws(crypto.json_bytes(value), self.gateway_key, self.clients['gateway']['key_id'], 'odexa-event+jws').encode()
        old = dict(event, occurred_at=later(admitted['checked_at'], -1))
        self.call('events', sign(old), role='gateway', status=400)
        _, intake = self.call('events', sign(event), role='gateway', status=201)
        _, retry = self.call('events', sign(event), role='gateway', status=200)
        self.assertEqual(intake, retry)
        payload = crypto.strict_json(crypto.verify_jws(intake, crypto.public_jwk(self.signing_key), 'odexa-event-record+jws')[0])
        self.assertEqual(payload['assurance'], 'origin_key_verified')
        changed = dict(event, event_id=uid())
        self.call('events', sign(changed), role='gateway', status=409)


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