"""Exercise generated ingress with an owned nginx process and real HTTP, WebSocket and SSE."""
from http.client import HTTPResponse
import json
from pathlib import Path
import shutil
import socket
import subprocess
import tempfile
import threading
import time
import unittest
from urllib.error import HTTPError, URLError
from urllib.request import Request, urlopen

import edge_vhost
from nginx_test_fixture import BackendServer, LEGACY_VHOST, stop
from test_edge_inputs import configured


class Ingress(unittest.TestCase):
    def setUp(self):
        self.assertIsNotNone(shutil.which("nginx"), "nginx is required for the backend ingress contract")
        directory = tempfile.TemporaryDirectory()
        self.addCleanup(directory.cleanup)
        self.root = Path(directory.name)
        server = self.server = BackendServer()
        self.addCleanup(server.server_close)
        thread = threading.Thread(target=server.serve_forever, daemon=True)
        thread.start()
        self.addCleanup(server.shutdown)
        with socket.socket() as listener:
            listener.bind(("127.0.0.1", 0))
            self.port = listener.getsockname()[1]
        self.config = self.root / ".backend-edge.json"
        self.config.write_text(json.dumps(configured()))
        legacy = self.root / "backend-fixture.conf"
        legacy.write_text(LEGACY_VHOST.replace(":31000;", f":{server.server_port};"))
        files = edge_vhost.planned(self.config, "fixture", "edge.example", server.server_port, self.root)
        for path, (data, mode) in files.items():
            path.write_bytes(data.replace(b"listen 80;", f"listen 127.0.0.1:{self.port};".encode()))
            path.chmod(mode)
        nginx = (f"pid {self.root}/nginx.pid;\nerror_log {self.root}/error.log;\n"
                 "events { worker_connections 64; }\nhttp { access_log off;\n"
                 f"include {self.root}/backend-fixture.conf;\n"
                 f"include {self.root}/backend-fixture-edge.conf;\n}}\n")
        (self.root / "nginx.conf").write_text(nginx)
        process = subprocess.Popen(["nginx", "-p", str(self.root) + "/", "-c", "nginx.conf",
                                    "-g", "daemon off;"], stdout=subprocess.PIPE, stderr=subprocess.PIPE)
        self.addCleanup(stop, process)
        for _ in range(50):
            try:
                self.fetch("legacy.example")
                break
            except URLError:
                self.assertIsNone(process.poll(), "owned nginx did not start")
                time.sleep(0.05)
        else:
            self.fail("owned nginx did not become ready")

    def fetch(self, host, token=None):
        headers = {"Host": host}
        if token is not None:
            headers["X-App-Robot-Backend-Token"] = token
        request = Request(f"http://127.0.0.1:{self.port}/healthz?probe=owned", headers=headers)
        try:
            with urlopen(request, timeout=2) as response:
                return response.status, json.load(response)
        except HTTPError as error:
            with error:
                return error.code, None

    def test_real_nginx_rejects_missing_and_forged_tokens_and_retains_legacy_access(self):
        self.assertEqual(self.fetch("edge.example"), (403, None))
        self.assertEqual(self.fetch("edge.example", "wrong-token"), (403, None))
        for host, token in (("edge.example", configured()["ingress_token"]), ("legacy.example", None)):
            status, body = self.fetch(host, token)
            self.assertEqual(status, 200)
            self.assertEqual(body["build_sha"], "a" * 40)
            self.assertEqual(body["path"], "/healthz?probe=owned")
            self.assertFalse(body["ingress_received"])
            self.assertIsNone(body["upgrade"])
            self.assertEqual(body["connection"], "close")

    def test_real_nginx_exchanges_websocket_frames_and_enforces_ingress(self):
        cases = (("edge.example", None, 403), ("edge.example", "wrong-token", 403),
                 ("edge.example", configured()["ingress_token"], 101), ("legacy.example", None, 101))
        for host, token, status in cases:
            with self.subTest(host=host, status=status), socket.create_connection(("127.0.0.1", self.port), 2) as peer:
                request = (f"GET /ws HTTP/1.1\r\nHost: {host}\r\nUpgrade: websocket\r\n"
                           "Connection: Upgrade\r\nSec-WebSocket-Version: 13\r\n"
                           "Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n")
                if token is not None:
                    request += f"X-App-Robot-Backend-Token: {token}\r\n"
                peer.sendall((request + "\r\n").encode())
                with HTTPResponse(peer) as response:
                    response.begin()
                    self.assertEqual(response.status, status)
                    if status != 101:
                        continue
                    self.assertEqual(response.getheader("Upgrade").lower(), "websocket")
                    self.assertEqual(response.getheader("Sec-WebSocket-Accept"), "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=")
                    payload, mask = b"ping-through-nginx", b"\x01\x02\x03\x04"
                    masked = bytes(value ^ mask[index % 4] for index, value in enumerate(payload))
                    peer.sendall(bytes((0x81, 0x80 | len(payload))) + mask + masked)
                    self.assertEqual(response.fp.read(2), bytes((0x81, len(payload))))
                    self.assertEqual(response.fp.read(len(payload)), payload)

    def test_real_nginx_delivers_sse_before_upstream_completion(self):
        for host, token in (("edge.example", configured()["ingress_token"]), ("legacy.example", None)):
            with self.subTest(host=host):
                self.server.stream_release.clear()
                self.server.stream_finished.clear()
                headers = {"Host": host}
                if token is not None:
                    headers["X-App-Robot-Backend-Token"] = token
                request = Request(f"http://127.0.0.1:{self.port}/stream", headers=headers)
                try:
                    with urlopen(request, timeout=2) as response:
                        self.assertEqual(response.readline(), b"data: first\n")
                        self.assertFalse(self.server.stream_finished.is_set())
                        self.server.stream_release.set()
                        self.assertEqual(response.read(), b"\ndata: second\n\n")
                finally:
                    self.server.stream_release.set()
                    self.assertTrue(self.server.stream_finished.wait(2))
