"""A later job resumes retained public SDK bytes without bypassing digest or account checks."""
import hashlib
import os
from pathlib import Path
import shutil
import tempfile
import unittest
from unittest.mock import patch

from flutter_manifest_fixture import tls_origin
import test_flutter_archive as archive_fixture

ARCHIVE, ROOT = archive_fixture.ARCHIVE, archive_fixture.ROOT
DIGEST = hashlib.sha256(ARCHIVE).hexdigest()


def invoke(root, server, cert, **kwargs):
    with patch.dict(os.environ, {"GOWALK_FLUTTER_ARCHIVE_CACHE": str(root / "gowalk-flutter-archives")}):
        return archive_fixture.ArchiveRecoveryTests().invoke(root, server, cert, **kwargs)


class ArchiveCacheTests(unittest.TestCase):
    def test_later_job_resumes_after_the_previous_job_exhausted_its_retries(self):
        with tempfile.TemporaryDirectory() as first, tempfile.TemporaryDirectory() as second:
            root, resumed = Path(first), Path(second)
            with tls_origin(root, ["partial", "range_partial", "range_partial", "range_partial"],
                           ARCHIVE) as (server, cert):
                result = invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout), (18, b""))
            partial = root / "gowalk-flutter-archives" / DIGEST / "download"
            self.assertEqual(partial.read_bytes(), ARCHIVE[:124])
            shutil.copytree(root / "gowalk-flutter-archives", resumed / "gowalk-flutter-archives")
            with tls_origin(resumed, ["range"], ARCHIVE) as (server, cert):
                result = invoke(resumed, server, cert)
                self.assertEqual((result.returncode, result.stdout), (0, ARCHIVE))
                self.assertEqual(server.ranges, ["bytes=124-"])
            self.assertFalse((resumed / "gowalk-flutter-archives" / DIGEST / "download").exists())

    def test_complete_cached_archive_is_verified_without_another_archive_request(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            partial = root / "gowalk-flutter-archives" / DIGEST / "download"
            partial.parent.mkdir(parents=True)
            partial.write_bytes(ARCHIVE)
            with tls_origin(root, [], ARCHIVE) as (server, cert):
                result = invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout, server.connects), (0, ARCHIVE, 0))

    def test_changed_digest_cannot_reuse_another_release_partial(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            foreign = root / "gowalk-flutter-archives" / ("a" * 64) / "download"
            foreign.parent.mkdir(parents=True)
            foreign.write_bytes(b"unrelated release")
            with tls_origin(root, ["ok"], ARCHIVE) as (server, cert):
                result = invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout), (0, ARCHIVE))
                self.assertEqual(server.ranges, [""])
            self.assertEqual(foreign.read_bytes(), b"unrelated release")

    def test_corrupted_restored_partial_fails_verification_and_is_not_cached_again(self):
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            partial = root / "gowalk-flutter-archives" / DIGEST / "download"
            partial.parent.mkdir(parents=True)
            partial.write_bytes(b"wrong bytes")
            with tls_origin(root, ["range"], ARCHIVE) as (server, cert):
                result = invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout), (1, b""))
                self.assertEqual(server.ranges, ["bytes=11-"])
            self.assertFalse(partial.exists())

    def test_symlinked_cache_does_not_read_or_write_another_directory(self):
        with tempfile.TemporaryDirectory() as directory, tempfile.TemporaryDirectory() as outside:
            root = Path(directory)
            (root / "gowalk-flutter-archives").symlink_to(outside, target_is_directory=True)
            with tls_origin(root, [], ARCHIVE) as (server, cert):
                result = invoke(root, server, cert)
                self.assertEqual((result.returncode, result.stdout, server.connects), (1, b"", 0))
            self.assertEqual(list(Path(outside).iterdir()), [])


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