"""Real TLS transfers resume private SDK bytes and verify the official manifest digest."""
import hashlib
import json
import os
from pathlib import Path
import subprocess
import tempfile
import unittest

from flutter_manifest_fixture import tls_origin

ROOT = Path(__file__).parent
ARCHIVE = bytes(range(256)) * 512


class ArchiveRecoveryTests(unittest.TestCase):
    def invoke(self, root, server, cert, hooked=True, trusted=True, digest=None):
        env = {**os.environ, "HOME": str(root), "RUNNER_TEMP": str(root),
               "VERSION_MANIFEST": json.dumps({"sha256": hashlib.sha256(ARCHIVE).hexdigest()
                                                if digest is None else digest})}
        for key in ("BASH_ENV", "CURL_CA_BUNDLE", "SSL_CERT_FILE", "SSL_CERT_DIR",
                    "HTTPS_PROXY", "https_proxy", "HTTP_PROXY", "http_proxy", "ALL_PROXY", "all_proxy"):
            env.pop(key, None)
        if trusted:
            env["CURL_CA_BUNDLE"] = str(cert)
        if hooked:
            env.update(BASH_ENV=str(ROOT / "flutter_manifest_env.sh"),
                       GOWALK_FLUTTER_MANIFEST_HELPER=str(ROOT / "flutter_manifest.py"))
            # A local curl preference must never disable the helper's certificate checks.
            (root / ".curlrc").write_text('insecure\n')
        url = server.url("/flutter_infra_release/releases/stable/linux/flutter_linux_3.44.0-stable.tar.xz")
        result = subprocess.run(["bash", "-c", 'curl --connect-timeout 15 --retry 5 "$1"', "bash", url],
                                env=env, capture_output=True, timeout=20)
        self.assertFalse(list(root.glob("flutter-manifest-*")))
        return result

    def test_original_partial_exit18_recovers_to_only_the_complete_archive(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            with tls_origin(root, ["partial"], ARCHIVE) as (server, cert):
                original = self.invoke(root, server, cert, hooked=False)
                self.assertEqual(original.returncode, 18)
            with tls_origin(root, ["partial", "range"], ARCHIVE) as (server, cert):
                result = self.invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout), (0, ARCHIVE))
                self.assertEqual(server.connects, 2)
                self.assertEqual(server.ranges, ["", "bytes=31-"])
                self.assertIn(b"flutter_archive_recovered attempts=2", result.stderr)

    def test_archive_transport_never_retries_auth_or_untrusted_certificate(self):
        for mode, trusted in [("origin_auth", True), ("ok", False)]:
            with self.subTest(mode=mode), tempfile.TemporaryDirectory() as directory:
                root = Path(directory)
                with tls_origin(root, [mode], ARCHIVE) as (server, cert):
                    result = self.invoke(root, server, cert, trusted=trusted)
                    self.assertNotEqual(result.returncode, 0)
                    self.assertEqual((result.stdout, server.connects), (b"", 1))

    def test_authentication_refusal_after_partial_transfer_stops_without_publishing(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            with tls_origin(root, ["partial", "origin_auth"], ARCHIVE) as (server, cert):
                result = self.invoke(root, server, cert)
                self.assertNotEqual(result.returncode, 0)
                self.assertEqual((result.stdout, server.connects), (b"", 2))

    def test_exhausted_archive_retry_never_emits_partial_bytes(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            modes = ["partial", "range_partial", "range_partial", "range_partial"]
            with tls_origin(root, modes, ARCHIVE) as (server, cert):
                result = self.invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout, server.connects), (18, b"", 4))
                self.assertEqual(server.ranges, ["", "bytes=31-", "bytes=62-", "bytes=93-"])
                self.assertIn(b'"retained_bytes":124', result.stderr)

    def test_repeated_disconnects_resume_until_the_verified_archive_is_complete(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            with tls_origin(root, ["partial", "range_partial", "range"], ARCHIVE) as (server, cert):
                result = self.invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout), (0, ARCHIVE))
                self.assertEqual(server.ranges, ["", "bytes=31-", "bytes=62-"])

    def test_origin_ignoring_range_restarts_without_appending_another_full_copy(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            with tls_origin(root, ["partial", "ok", "ok"], ARCHIVE) as (server, cert):
                result = self.invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout), (0, ARCHIVE))
                self.assertEqual(server.ranges, ["", "bytes=31-", ""])

    def test_changed_resumed_bytes_fail_checksum_before_any_output(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            with tls_origin(root, ["partial", "range_corrupt"], ARCHIVE) as (server, cert):
                result = self.invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout, server.connects), (1, b"", 2))

    def test_missing_manifest_digest_refuses_before_download(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            with tls_origin(root, [], ARCHIVE) as (server, cert):
                result = self.invoke(root, server, cert, digest="")
                self.assertEqual((result.returncode, result.stdout, server.connects), (1, b"", 0))


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