import copy
import datetime as dt
import http.server
import json
from pathlib import Path
import socket
import ssl
import tempfile
import threading
import time
import unittest
from unittest import mock

from odexa_ref import crypto
from odexa_ref.fixture import make_tls
from profiles import delegation
from profiles.authority_transport import AuthorityTransport, ScopedCredential, TransportError, utc_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):
        received = self.rfile.read(int(self.headers.get('Content-Length', '0')))
        self.server.requests.append((self.command, self.path, dict(self.headers), received))
        route = self.server.routes.get(self.path, {'status':404,'body':b'{}'})
        time.sleep(route.get('delay', 0))
        try:
            self.send_response_only(route.get('status', 200))
            for name, value in route.get('headers', [('Content-Type', 'application/json'),
                                                    ('Content-Length', str(len(route['body']))),
                                                    ('Cache-Control', 'no-store')]):
                self.send_header(name, value)
            self.send_header('Connection', 'close'); self.end_headers()
            if route.get('trickle'):
                for byte in route['body']:
                    self.wfile.write(bytes([byte])); self.wfile.flush(); time.sleep(route['trickle'])
            else:
                self.wfile.write(route['body']); self.wfile.flush()
        except (OSError, ValueError): pass
        self.close_connection = True


class TransportTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.directory = tempfile.TemporaryDirectory()
        cls.path = Path(cls.directory.name); make_tls(cls.path)
        context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
        context.load_cert_chain(cls.path/'tls-cert.pem', cls.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 = context.wrap_socket(server.socket, server_side=True)
            server.routes = {}; server.requests = []
            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()
        cls.directory.cleanup()

    def setUp(self):
        self.server, self.other = [s for s, _ in self.servers]
        for server in (self.server, self.other): server.routes = {}; server.requests = []
        self.transport = AuthorityTransport(ca_file=str(self.path/'ca.pem'), allowed_ips=['127.0.0.1'])
        corpus = json.loads((Path(__file__).parents[1]/'profiles/fixtures/delegation-cases.json').read_text())
        self.policy = copy.deepcopy(corpus['policy_document'])
        self.authority = copy.deepcopy(next(c['value'] for c in corpus['cases']
                                           if c['definition'] == 'authority' and c['structural_valid']))
        self.op = copy.deepcopy(next(c['value'] for c in corpus['cases']
                                    if c['definition'] == 'operation' and c['structural_valid']))
        def update(value):
            if isinstance(value, dict): return {k:update(v) for k,v in value.items()}
            if isinstance(value, list): return [update(v) for v in value]
            if isinstance(value, str):
                value = value.replace('https://publisher.example', self.origin).replace('https://provider.example', self.provider)
                if value == '2026-09-16T00:00:00Z': return (dt.datetime.now(dt.timezone.utc)-dt.timedelta(days=1)).strftime('%Y-%m-%dT%H:%M:%SZ')
                if value == '2026-09-18T00:00:00Z': return (dt.datetime.now(dt.timezone.utc)+dt.timedelta(days=1)).strftime('%Y-%m-%dT%H:%M:%SZ')
            return value
        self.policy, self.authority, self.op = map(update, (self.policy, self.authority, self.op))
        self.server.routes['/odexa-service.json'] = {'body':crypto.json_bytes(self.authority)}
        self.server.routes['/odexa.json'] = {'body':crypto.json_bytes(self.policy)}

    def fetch_fails(self):
        with self.assertRaises(TransportError): self.transport.fetch(self.origin)

    def test_rejected_status_closes_unread_response_stream(self):
        self.server.routes['/odexa-service.json']['status']=403
        original=http.client.HTTPConnection.getresponse;received=[]
        def capture(connection,*args,**kwargs):
            response=original(connection,*args,**kwargs);received.append(response);return response
        with mock.patch.object(http.client.HTTPConnection,'getresponse',capture):self.fetch_fails()
        self.assertEqual(len(received),1)
        self.assertTrue(received[0].closed)
        self.assertIsNone(received[0].fp)

    def test_real_two_origin_fetch_exact_bytes_and_authority_evaluation(self):
        snapshot = self.transport.fetch(self.origin)
        self.assertEqual(snapshot.authority_bytes, self.server.routes['/odexa-service.json']['body'])
        self.assertEqual(snapshot.policy_bytes, self.server.routes['/odexa.json']['body'])
        self.assertEqual(snapshot.observation['authority_digest'], crypto.digest(snapshot.authority_bytes))
        result = delegation.evaluate_authority(snapshot.authority_bytes, snapshot.policy_bytes, self.op,
                                                now=utc_now(), observation=snapshot.observation)
        self.assertEqual(result['decision'], 'allow', result)
        self.assertNotEqual(self.origin, self.provider)
        self.assertEqual(len(self.server.requests), 2)
        self.assertFalse(self.other.requests)
        for method, path, headers, body in self.server.requests:
            self.assertEqual(method, 'GET'); self.assertEqual(body, b'')
            self.assertNotIn('Authorization', headers); self.assertNotIn('Cookie', headers)
            self.assertEqual(headers['Cache-Control'], 'no-cache, no-store')

    def test_default_rejects_loopback_before_connect(self):
        with self.assertRaises(TransportError): AuthorityTransport(ca_file=str(self.path/'ca.pem')).fetch(self.origin)
        self.assertFalse(self.server.requests)

    def test_untrusted_ca_and_wrong_hostname_fail(self):
        with self.assertRaises(TransportError): AuthorityTransport(allowed_ips=['127.0.0.1']).fetch(self.origin)
        original = socket.getaddrinfo
        with mock.patch('profiles.authority_transport.socket.getaddrinfo',
                        side_effect=lambda host, port, **kw: original('127.0.0.1',port,**kw)):
            with self.assertRaises(TransportError): self.transport.fetch(f'https://wrong.example:{self.server.server_port}')
        self.assertFalse(self.server.requests)

    def test_resolve_once_and_pin_both_connections(self):
        original = socket.getaddrinfo
        with mock.patch('profiles.authority_transport.socket.getaddrinfo', wraps=original) as lookup:
            self.transport.fetch(self.origin)
        self.assertEqual(lookup.call_count, 1)

    def test_mixed_public_private_dns_rejected_without_connect(self):
        records = [(socket.AF_INET,socket.SOCK_STREAM,6,'',('8.8.8.8',443)),
                   (socket.AF_INET,socket.SOCK_STREAM,6,'',('10.1.2.3',443))]
        with mock.patch('profiles.authority_transport.socket.getaddrinfo',return_value=records):
            with self.assertRaises(TransportError): self.transport.fetch('https://publisher.example')
        self.assertFalse(self.server.requests)

    def test_redirect_never_contacts_other_origin(self):
        self.server.routes['/odexa-service.json'] = {'status':302,'body':b'{}',
            'headers':[('Location',self.provider+'/odexa-service.json'),('Content-Length','2')]}
        self.fetch_fails(); self.assertFalse(self.other.requests)

    def test_no_stale_fallback_after_prior_success_and_second_document_failure(self):
        self.transport.fetch(self.origin)
        self.server.routes['/odexa.json']['status'] = 503
        self.fetch_fails()
        self.assertEqual([r[1] for r in self.server.requests],
                         ['/odexa-service.json','/odexa.json','/odexa-service.json','/odexa.json'])

    def test_no_store_and_age_requirements(self):
        raw = self.server.routes['/odexa-service.json']['body']
        base = [('Content-Type','application/json'),('Content-Length',str(len(raw)))]
        for headers in [base, base+[('Cache-Control','max-age=60')],
                        base+[('Cache-Control','no-store'),('Age','1')],
                        base+[('Cache-Control','no-store'),('Warning','110 stale')]]:
            with self.subTest(headers=headers):
                self.server.routes['/odexa-service.json']['headers'] = headers
                self.fetch_fails()

    def test_strict_media_framing_and_size(self):
        body = self.server.routes['/odexa-service.json']['body']
        for special in [[('Content-Type','text/html')], [('Content-Type','application/json;charset=latin1')],
                        [('Content-Type','application/json'),('Content-Type','application/json')],
                        [('Content-Type','application/json'),('Content-Encoding','gzip')],
                        [('Content-Type','application/json'),('Transfer-Encoding','chunked')]]:
            with self.subTest(special=special):
                self.server.routes['/odexa-service.json']['headers'] = special + [('Content-Length',str(len(body))),('Cache-Control','no-store')]
                self.fetch_fails()
        for length in ['131073','0002','2, 2']:
            with self.subTest(length=length):
                self.server.routes['/odexa-service.json']['headers'] = [('Content-Type','application/json'),('Content-Length',length),('Cache-Control','no-store')]
                self.fetch_fails()
        self.server.routes['/odexa-service.json']['headers'] = [('Content-Type','application/json'),('Content-Length',str(len(body))),('Content-Length',str(len(body))),('Cache-Control','no-store')]
        self.fetch_fails()

    def test_invalid_duplicate_or_non_utf8_json(self):
        for body in [b'{"x":1,"x":2}', b'{"x":1.0}', b'{"x":-0}', b'[]', b'\xff', b'{"x":"\\ud800"}']:
            with self.subTest(body=body):
                self.server.routes['/odexa-service.json'] = {'body':body}
                self.fetch_fails()

    def test_truncated_body_and_total_trickle_deadline(self):
        self.server.routes['/odexa-service.json'] = {'body':b'{}', 'headers':[
            ('Content-Type','application/json'),('Content-Length','3'),('Cache-Control','no-store')]}
        self.fetch_fails()
        self.transport = AuthorityTransport(ca_file=str(self.path/'ca.pem'), allowed_ips=['127.0.0.1'], deadline_seconds=0.15)
        self.server.routes['/odexa-service.json'] = {'body':b'{"long":"abcdefghij"}', 'trickle':0.04}
        start = time.monotonic(); self.fetch_fails()
        self.assertLess(time.monotonic()-start, 0.7)

    def test_combined_fetch_budget_and_start_time(self):
        self.transport = AuthorityTransport(ca_file=str(self.path/'ca.pem'), allowed_ips=['127.0.0.1'], deadline_seconds=0.2)
        self.server.routes['/odexa-service.json']['delay'] = 0.12
        self.server.routes['/odexa.json']['delay'] = 0.12
        start = time.monotonic(); self.fetch_fails()
        self.assertLess(time.monotonic()-start, 0.7)
        self.server.routes['/odexa-service.json']['delay'] = 0
        self.server.routes['/odexa.json']['delay'] = 0
        with mock.patch('profiles.authority_transport.utc_now',return_value='2026-09-17T01:02:03Z') as clock:
            snapshot = self.transport.fetch(self.origin)
        self.assertEqual(snapshot.observation['checked_at'],'2026-09-17T01:02:03Z')
        self.assertEqual(clock.call_count,1)

    def test_dns_wait_has_deadline_without_late_connection(self):
        original = socket.getaddrinfo
        def slow(*args,**kwargs):
            time.sleep(0.25); return original(*args,**kwargs)
        self.transport = AuthorityTransport(ca_file=str(self.path/'ca.pem'), allowed_ips=['127.0.0.1'], deadline_seconds=0.05)
        start = time.monotonic()
        with mock.patch('profiles.authority_transport.socket.getaddrinfo',side_effect=slow): self.fetch_fails()
        self.assertLess(time.monotonic()-start, 0.2)
        time.sleep(0.3); self.assertFalse(self.server.requests)

    def test_scoped_post_jose_no_ambient_credentials_or_proxy(self):
        body=b'{"check":"one"}'; signed=b'header.payload.signature'
        self.other.routes['/tenant-a/api/verify'] = {'body':signed,'headers':[
            ('Content-Type','application/jose'),('Content-Length',str(len(signed)))]}
        credential=ScopedCredential(self.provider,'/tenant-a/api/','Basic synthetic-only')
        self.assertNotIn('synthetic-only', repr(credential))
        with mock.patch.dict('os.environ',{'HTTPS_PROXY':'https://invalid.invalid','HTTP_PROXY':'http://invalid.invalid'}):
            result=self.transport.request(self.provider,'/tenant-a/api/verify',method='POST',body=body,
                credential=credential,response_type='application/jose',max_bytes=262144)
        self.assertEqual(result.body,signed)
        self.assertEqual(self.other.requests[-1][2]['Authorization'],'Basic synthetic-only')
        self.assertEqual(self.other.requests[-1][3],body)
        for origin,path in [(self.origin,'/tenant-a/api/verify'),(self.provider,'/tenant-b/api/verify'),
                            (self.provider,'/tenant-a/api2/verify')]:
            with self.subTest(origin=origin,path=path):
                with self.assertRaises(TransportError):
                    self.transport.request(origin,path,method='POST',body=body,credential=credential)
        self.assertEqual(len(self.other.requests),1); self.assertFalse(self.server.requests)

    def test_request_rejects_unsafe_path_body_and_header_before_network(self):
        cases=[{'path':'/a/../b'}, {'path':'/a%2Fb'}, {'path':'/a?x=1'}, {'path':'/a#x'},
               {'path':'/a','method':'POST','body':b' '*131073},
               {'path':'/a','method':'GET','body':b'{}'},
               {'path':'/a','method':'POST','body':b'{"x":1,"x":2}'},
               {'path':'/a','credential':ScopedCredential(self.origin,'/','Basic x\r\nInjected: y')}]
        for case in cases:
            with self.subTest(case=case):
                with self.assertRaises(ValueError): self.transport.request(self.origin,**case)
        self.assertFalse(self.server.requests)

    def test_before_send_denial_after_tls_emits_no_http_or_credential(self):
        calls = []
        credential = ScopedCredential(self.provider, '/api/', 'Basic never-send-this')
        for decision in (False, None, 1):
            def guard():
                calls.append(decision)
                return decision
            with self.subTest(decision=decision):
                with self.assertRaises(TransportError):
                    self.transport.request(self.provider, '/api/verify', method='POST', body=b'{}',
                                           credential=credential, before_send=guard)
        self.assertEqual(calls, [False, None, 1])
        self.assertFalse(self.other.requests)

    def test_signed_jose_post_preserves_exact_bytes_and_rejects_other_media(self):
        from cryptography.hazmat.primitives.asymmetric import ec
        raw = b'{"sample":"exact whitespace", "count":1}'
        signed = crypto.sign_jws(raw, ec.derive_private_key(7, ec.SECP256R1()),
                                 'https://agent.example/keys/one', 'odexa-acceptance+jws').encode()
        self.other.routes['/api/agreements'] = {'body': b'{}'}
        result = self.transport.request(self.provider, '/api/agreements', method='POST',
                                         body=signed, content_type='application/jose')
        self.assertEqual(result.body, b'{}')
        self.assertEqual(self.other.requests[0][3], signed)
        self.assertEqual(self.other.requests[0][2]['Content-Type'], 'application/jose')
        for content_type, body in [('text/plain', signed), ('application/jose;charset=utf-8', signed),
                                   ('application/jose', b'not.a.jws')]:
            with self.subTest(content_type=content_type,body=body):
                with self.assertRaises(TransportError):
                    self.transport.request(self.provider, '/api/agreements', method='POST',
                                           body=body, content_type=content_type)
        self.assertEqual(len(self.other.requests), 1)
        key = ec.derive_private_key(7, ec.SECP256R1())
        boundary_payload = b'{"x":"' + b'a' * (131072-8) + b'"}'
        boundary = crypto.sign_jws(boundary_payload, key, 'https://agent.example/keys/one', 'odexa-event+jws').encode()
        self.assertGreater(len(boundary), 131072)
        self.transport.request(self.provider, '/api/agreements', method='POST', body=boundary,
                               content_type='application/jose')
        self.assertEqual(self.other.requests[-1][3], boundary)
        too_large = crypto.sign_jws(boundary_payload[:-2]+b'a"}', key, 'https://agent.example/keys/one', 'odexa-event+jws').encode()
        for body in (too_large, b'x'*262145):
            with self.subTest(length=len(body)):
                with self.assertRaises(TransportError):
                    self.transport.request(self.provider, '/api/agreements', method='POST', body=body,
                                           content_type='application/jose')
        self.assertEqual(len(self.other.requests), 2)

    def test_explicit_created_or_retry_statuses_without_redirect_acceptance(self):
        route = self.other.routes['/api/events'] = {'body':b'{}','status':201}
        for status in (201,200):
            route['status']=status
            result=self.transport.request(self.provider,'/api/events',method='POST',body=b'{}',expected_status=(200,201))
            self.assertEqual(result.status,status)
        route['status']=202
        with self.assertRaises(TransportError):
            self.transport.request(self.provider,'/api/events',method='POST',body=b'{}',expected_status=(200,201))
        count=len(self.other.requests)
        for statuses in [(200,302),(),(200,200),[200,201],True]:
            with self.subTest(statuses=statuses):
                with self.assertRaises(TransportError):
                    self.transport.request(self.provider,'/api/events',expected_status=statuses)
        self.assertEqual(len(self.other.requests),count)

    def test_explicit_trusted_paid_request_type_does_not_widen_crypto_defaults(self):
        from cryptography.hazmat.primitives.asymmetric import ec
        typ='odexa-payment-mandate+jws'; types=frozenset({typ}); before=set(crypto.TYPES)
        body=crypto.sign_jws(b'{"sample":"synthetic mandate"}',ec.derive_private_key(7,ec.SECP256R1()),
            'https://agent.example/keys/one',typ,allowed_types=types).encode()
        self.other.routes['/api/mandates']={'body':b'{}'}
        with self.assertRaises(TransportError):
            self.transport.request(self.provider,'/api/mandates',method='POST',body=body,content_type='application/jose')
        self.assertFalse(self.other.requests)
        self.transport.request(self.provider,'/api/mandates',method='POST',body=body,
            content_type='application/jose',allowed_jws_types=types)
        self.assertEqual(self.other.requests[-1][3],body);self.assertEqual(set(crypto.TYPES),before)
        with self.assertRaises(ValueError):crypto.split_jws(body)


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