"""Publisher delivery boundary and independent Node/Python TLS exchange."""
import json
import os
from pathlib import Path
import subprocess
import unittest

from odexa_ref import crypto
from profiles.authority_history import AuthorityHistory
from profiles.authority_transport import ScopedCredential
from profiles.free_gateway import FreeGateway, GatewayError
import test_free_service as fixture


class FreeGatewayTests(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(); self.configure_gateway()

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

    def configure_gateway(self):
        f = self.f
        self.args = dict(binding=f.binding, base_url=f.base, transport=f.transport,
            history=AuthorityHistory(f.path/'gateway-history.sqlite3'),
            credential=ScopedCredential(f.host, f.base[len(f.host):], f.clients['gateway']['authorization']), clock=f.now)
        self.gateway = FreeGateway(f.path/'gateway.sqlite3', **self.args)
        self.original_route = f.origin_route
        self.asset = b'Originator-owned test article.\n'
        f.origin_server.route = self.route

    def route(self, method, path, headers, body):
        if method != 'GET' or path != '/about': return self.original_route(method, path, headers, body)
        try:
            auth = headers.get('Authorization', '')
            if not auth.startswith('Bearer '): raise GatewayError('no token')
            request_id = self.gateway.admit(auth[7:], resource_url=self.f.origin+'/about',
                actions=headers.get('Odexa-Actions', '').split(','), purposes=headers.get('Odexa-Purposes', '').split(','))
        except ValueError:
            return 403, {'Content-Type': 'application/json', 'Cache-Control': 'no-store'}, b'{}'
        return 200, {'Content-Type': 'text/plain', 'Cache-Control': 'no-store'}, self.asset, {
            'begin': lambda: self.gateway.begin_delivery(request_id),
            'finish': lambda written, complete: self.gateway.finish_delivery(request_id,
                bytes_written=written, representation_digest=crypto.digest(self.asset), complete=complete)}

    def admit(self, token):
        return self.gateway.admit(token, resource_url=self.f.request['url'], actions=['retrieve'], purposes=['public_retrieval'])

    def test_admission_delivery_and_unresolved_restart_are_distinct(self):
        f = self.f; f.agree(); token = f.token()
        first = self.admit(token)
        self.gateway = FreeGateway(f.path/'gateway.sqlite3', **self.args)
        self.assertEqual(self.gateway.records()[0]['state'], 'admitted')
        response = self.gateway.begin_delivery(first)
        self.assertEqual(response['admission']['principal_id'], f.principal)
        with self.assertRaises(GatewayError): self.gateway.begin_delivery(first)
        self.gateway.finish_delivery(first, bytes_written=len(self.asset), representation_digest=crypto.digest(self.asset), complete=True)
        self.assertEqual(self.gateway.records()[0]['state'], 'sent')
        second = self.admit(token); self.gateway.begin_delivery(second)
        self.gateway = FreeGateway(f.path/'gateway.sqlite3', **self.args)
        self.assertEqual(self.gateway.records()[1]['state'], 'in_flight')
        with self.assertRaises(GatewayError): self.gateway.begin_delivery(second)
        retained = b''.join(bytes(row['response']) for row in self.gateway.records())
        self.assertNotIn(token.encode(), retained)

    def test_withdrawal_outage_expiry_and_cached_body_cannot_reuse_admission(self):
        f = self.f; f.agree(); token = f.token(); first = self.admit(token)
        f.offset = 6
        with self.assertRaises(GatewayError): self.gateway.begin_delivery(first)
        f.offset = 0; f.authority.update(revision=2, delegations=[])
        with self.assertRaises(GatewayError): self.admit(token)
        # Even if a caller holds the body in its cache, it needs a new admission.
        self.assertEqual(len(self.gateway.records()), 1)
        f.origin_status = 503
        with self.assertRaises(ValueError): self.admit(token)

    def test_change_after_provider_reply_blocks_local_admission(self):
        f = self.f; f.agree(); token = f.token()
        original = f.provider_server.route
        def route(*args):
            result = original(*args)
            if args[1].endswith('/introspect'): f.authority.update(revision=2, delegations=[])
            return result
        f.provider_server.route = route
        with self.assertRaises(GatewayError): self.admit(token)
        self.assertEqual(self.gateway.records(), [])

    def test_completion_cannot_claim_invalid_time_or_byte_count(self):
        f = self.f; f.agree(); token = f.token(); request_id = self.admit(token)
        self.gateway.begin_delivery(request_id)
        for count in [True, -1, 9007199254740992]:
            with self.assertRaises(GatewayError):
                self.gateway.finish_delivery(request_id, bytes_written=count, representation_digest=crypto.digest(self.asset), complete=True)
        f.offset = -1
        with self.assertRaises(GatewayError):
            self.gateway.finish_delivery(request_id, bytes_written=1, representation_digest=crypto.digest(self.asset), complete=False)
        self.assertEqual(self.gateway.records()[0]['state'], 'in_flight')

    def test_pre_send_expiry_does_not_disclose_token_or_gateway_credential(self):
        f = self.f; f.agree(); token = f.token()
        old = self.args['history'].require_current; count = 0
        def slow(snapshot):
            nonlocal count
            result = old(snapshot); count += 1
            if count == 2: f.offset = 6
            return result
        self.args['history'].require_current = slow
        before = len(f.provider_server.requests)
        with self.assertRaises(ValueError): self.admit(token)
        self.assertEqual(len(f.provider_server.requests), before)
        self.assertEqual(self.gateway.records(), [])

    def run_independent(self, name):
        f = self.f
        node = os.environ.get('ODEXA_NODE', 'node')
        private = dict(crypto.public_jwk(f.agent_key), d=crypto.b64u(f.agent_key.private_numbers().private_value.to_bytes(32, 'big')))
        config = dict(publisherOrigin=f.origin, serviceId=f.binding['service_id'], baseUrl=f.base,
            delegationId=f.binding['delegation_id'], caFile=str(f.tls_path/'ca.pem'),
            agent=dict(clientId='agent', principalId=f.principal, password='a'*43, keyId=f.clients['agent']['key_id'],
                reporterId=f.clients['agent']['reporter_id'], privateJwk=private),
            request=f.request, accessSeconds=120, useSeconds=240,
            expectedTermsDigest=crypto.digest(f.service.terms_bytes), includeResource=True)
        config_path = f.path/'node-client-private.json'; config_path.write_bytes(crypto.json_bytes(config)); config_path.chmod(0o600)
        out = Path(os.environ.get('ODEXA_FREE_REPORT_DIR', str(f.path/'reports')))/name
        result = subprocess.run([node, 'verification/free-client.mjs', '--config', str(config_path), '--outdir', str(out)],
            capture_output=True, text=True, timeout=30)
        report = json.loads((out/'free-client-report.json').read_text())
        self.assertEqual(result.returncode, 0, report)
        self.assertEqual(report['failed'], 0, report)
        self.assertTrue(report['retrieval_reported'])
        records = self.gateway.records()
        self.assertEqual(len(records), 1); self.assertEqual(records[0]['state'], 'sent')
        self.assertEqual(records[0]['bytes_written'], len(self.asset))
        # Each credential is restricted to its intended selected service or asset.
        for method, path, headers, body in f.origin_server.requests:
            if path in {'/odexa.json', '/odexa-service.json'}: self.assertNotIn('Authorization', headers)
        with f.store.connect() as db:
            self.assertEqual(db.execute('SELECT count(*) FROM records').fetchone()[0], 1)
            self.assertEqual(db.execute("SELECT count(*) FROM agreements WHERE state='revoked'").fetchone()[0], 1)

    def test_independent_node_two_origin_exchange_with_real_delivery_and_report(self):
        self.run_independent('delegated')

    def test_independent_node_direct_exchange_without_provider(self):
        f = self.f; f.temp.cleanup()
        import tempfile
        f.temp = tempfile.TemporaryDirectory(); f.path = Path(f.temp.name)
        f.make_service(external=False); self.configure_gateway()
        self.run_independent('direct')
        self.assertFalse(f.provider_server.requests)


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