from __future__ import annotations

import base64
import hashlib
import zlib
from datetime import datetime, timedelta, timezone
from pathlib import Path

import pytest

from osm_lead_source_service.domain.enums import OsmObjectType
from osm_lead_source_service.errors import (
    ContractViolation,
    PbfResourceLimitExceeded,
    PbfStageContractError,
    PbfValidationError,
)
from osm_lead_source_service.pbf import (
    OsmiumPbfParser,
    ParsedOsmObject,
    PbfParseLimits,
    VerifiedPbfInput,
)

_FIXTURE = Path(__file__).parents[2] / "fixtures" / "pbf" / "tiny.osm.pbf.b64"
_EXPECTED_SHA256 = "1f181fa63b838de88c5c7d8ecff2a014ca52a9349c6470fd2656ad3c9e0ddb8c"


class CollectingSink:
    def __init__(self) -> None:
        self.objects: list[ParsedOsmObject] = []
        self.summary: object | None = None
        self.aborted = False

    def accept(self, parsed_object: ParsedOsmObject) -> None:
        self.objects.append(parsed_object)

    def finish(self, summary: object) -> None:
        self.summary = summary

    def abort(self) -> None:
        self.objects.clear()
        self.aborted = True


def materialize_fixture(tmp_path: Path) -> Path:
    path = tmp_path / "tiny.osm.pbf"
    path.write_bytes(base64.b64decode(_FIXTURE.read_text(encoding="ascii").strip(), validate=True))
    return path


def approved_input(
    path: Path,
    *,
    content_sha256: str | None = None,
    size_bytes: int | None = None,
) -> VerifiedPbfInput:
    data = path.read_bytes()
    return VerifiedPbfInput(
        local_path=path,
        content_sha256=hashlib.sha256(data).hexdigest()
        if content_sha256 is None
        else content_sha256,
        size_bytes=len(data) if size_bytes is None else size_bytes,
        provenance="checked-in synthetic fixture",
    )


def test_fixture_parses_all_supported_types_deterministically(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)
    sink = CollectingSink()
    start = datetime(2026, 7, 11, 1, 0, tzinfo=timezone.utc)
    times = iter((start, start + timedelta(seconds=1)))
    summary = OsmiumPbfParser(clock=lambda: next(times)).parse(approved_input(path), sink)

    assert summary.content_sha256 == _EXPECTED_SHA256
    assert summary.size_bytes == 438
    assert summary.counts.nodes == 2
    assert summary.counts.ways == 1
    assert summary.counts.relations == 1
    assert summary.counts.total == 4
    assert summary.total_tags == 8
    assert sink.summary is summary
    assert sink.aborted is False
    assert [obj.identity.osm_type for obj in sink.objects] == [
        OsmObjectType.NODE,
        OsmObjectType.NODE,
        OsmObjectType.WAY,
        OsmObjectType.RELATION,
    ]

    cafe = sink.objects[0]
    assert cafe.identity.osm_id == 1
    assert cafe.version == 1
    assert cafe.tags == (("amenity", "cafe"), ("name", "Fixture Cafe"))
    assert cafe.longitude_e7 == 1_005_000_000
    assert cafe.latitude_e7 == 137_500_000

    way = sink.objects[2]
    assert way.identity.osm_id == 10
    assert way.node_refs == (1, 2)
    assert way.tags == (("amenity", "restaurant"), ("name", "Fixture Way"))

    relation = sink.objects[3]
    assert relation.identity.osm_id == 20
    assert relation.members[0].member_type is OsmObjectType.WAY
    assert relation.members[0].ref == 10
    assert relation.members[0].role == "outer"


def test_file_size_limit_is_checked_before_parsing(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)
    sink = CollectingSink()
    parser = OsmiumPbfParser(PbfParseLimits(max_file_size_bytes=437))
    with pytest.raises(PbfResourceLimitExceeded, match="max_file_size_bytes"):
        parser.parse(approved_input(path), sink)
    assert sink.objects == []


def test_object_limit_stops_streaming(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)
    sink = CollectingSink()
    parser = OsmiumPbfParser(PbfParseLimits(max_objects=3))
    with pytest.raises(PbfResourceLimitExceeded, match="max_objects"):
        parser.parse(approved_input(path), sink)
    assert sink.objects == []
    assert sink.aborted is True


def test_total_tag_limit_stops_streaming(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)
    sink = CollectingSink()
    parser = OsmiumPbfParser(PbfParseLimits(max_total_tags=4, max_tags_per_object=4))
    with pytest.raises(PbfResourceLimitExceeded, match="max_total_tags"):
        parser.parse(approved_input(path), sink)


def test_corrupt_payload_is_rejected(tmp_path: Path) -> None:
    path = tmp_path / "corrupt.osm.pbf"
    path.write_bytes(b"not-a-pbf")
    with pytest.raises((PbfValidationError, PbfResourceLimitExceeded)):
        OsmiumPbfParser().parse(approved_input(path), CollectingSink())


def test_unapproved_input_contract_is_rejected(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)
    with pytest.raises(PbfValidationError, match="approved_input"):
        OsmiumPbfParser().parse(path, CollectingSink())  # type: ignore[arg-type]


def test_approved_hash_and_size_are_verified_before_parsing(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)
    with pytest.raises(PbfValidationError, match="checksum"):
        OsmiumPbfParser().parse(
            approved_input(path, content_sha256="0" * 64),
            CollectingSink(),
        )
    with pytest.raises(PbfValidationError, match="size"):
        OsmiumPbfParser().parse(
            approved_input(path, size_bytes=437),
            CollectingSink(),
        )


def test_symlink_input_is_rejected(tmp_path: Path) -> None:
    target = materialize_fixture(tmp_path)
    symlink = tmp_path / "linked.osm.pbf"
    symlink.symlink_to(target)
    with pytest.raises(PbfValidationError, match="symbolic link"):
        OsmiumPbfParser().parse(approved_input(symlink), CollectingSink())


def test_sink_contract_is_required(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)
    with pytest.raises(PbfValidationError, match="sink"):
        OsmiumPbfParser().parse(approved_input(path), object())  # type: ignore[arg-type]


def _varint(value: int) -> bytes:
    encoded = bytearray()
    while True:
        byte = value & 0x7F
        value >>= 7
        if value:
            encoded.append(byte | 0x80)
        else:
            encoded.append(byte)
            return bytes(encoded)


def _bytes_field(field_number: int, value: bytes) -> bytes:
    return _varint((field_number << 3) | 2) + _varint(len(value)) + value


def _int_field(field_number: int, value: int) -> bytes:
    return _varint(field_number << 3) + _varint(value)


def _block(blob_type: str, blob: bytes) -> bytes:
    header = _bytes_field(1, blob_type.encode("ascii")) + _int_field(3, len(blob))
    return len(header).to_bytes(4, "big") + header + blob


def test_uncompressed_blob_bomb_is_rejected_before_pyosmium(tmp_path: Path) -> None:
    header_blob = _bytes_field(1, b"x")
    malicious_data_blob = _int_field(2, 33 * 1024 * 1024) + _bytes_field(3, b"x")
    path = tmp_path / "bomb.osm.pbf"
    path.write_bytes(_block("OSMHeader", header_blob) + _block("OSMData", malicious_data_blob))
    with pytest.raises(PbfResourceLimitExceeded, match="max_uncompressed_blob_bytes"):
        OsmiumPbfParser().parse(approved_input(path), CollectingSink())


def test_compressed_blob_cannot_lie_about_uncompressed_size(tmp_path: Path) -> None:
    header_payload = b"x"
    header_blob = _int_field(2, len(header_payload)) + _bytes_field(
        3, zlib.compress(header_payload)
    )
    expanded = b"x" * 4096
    dishonest_data_blob = _int_field(2, 1) + _bytes_field(3, zlib.compress(expanded))
    path = tmp_path / "dishonest.osm.pbf"
    path.write_bytes(_block("OSMHeader", header_blob) + _block("OSMData", dishonest_data_blob))
    with pytest.raises(PbfValidationError, match="raw_size"):
        OsmiumPbfParser().parse(approved_input(path), CollectingSink())


def test_compressed_blob_bomb_cannot_hide_behind_small_raw_size(tmp_path: Path) -> None:
    header_payload = b"x"
    header_blob = _int_field(2, len(header_payload)) + _bytes_field(
        3, zlib.compress(header_payload)
    )
    expanded = b"x" * 2048
    malicious_data_blob = _int_field(2, 1) + _bytes_field(3, zlib.compress(expanded))
    path = tmp_path / "compressed-bomb.osm.pbf"
    path.write_bytes(_block("OSMHeader", header_blob) + _block("OSMData", malicious_data_blob))
    parser = OsmiumPbfParser(PbfParseLimits(max_uncompressed_blob_bytes=1024))
    with pytest.raises(PbfResourceLimitExceeded, match="max_uncompressed_blob_bytes"):
        parser.parse(approved_input(path), CollectingSink())


def test_sink_is_aborted_when_streaming_fails(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)

    class FailingSink(CollectingSink):
        def accept(self, parsed_object: ParsedOsmObject) -> None:
            super().accept(parsed_object)
            if len(self.objects) == 2:
                raise RuntimeError("synthetic sink failure")

    sink = FailingSink()
    with pytest.raises(Exception, match="PBF parsing failed"):
        OsmiumPbfParser().parse(approved_input(path), sink)
    assert sink.aborted is True
    assert sink.objects == []
    assert sink.summary is None


def test_sink_contract_failure_has_secret_safe_stage_code(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)

    class ContractFailingSink(CollectingSink):
        def accept(self, parsed_object: ParsedOsmObject) -> None:
            del parsed_object
            raise ContractViolation("parsed object contains a duplicate tag key")

    sink = ContractFailingSink()
    with pytest.raises(PbfStageContractError) as raised:
        OsmiumPbfParser().parse(approved_input(path), sink)

    assert raised.value.diagnostic_code == (
        "PbfStageContractError:sink_accept:duplicate_normalized_tag_key"
    )
    assert sink.aborted is True
    assert raised.value.__cause__ is not None
    assert "duplicate tag key" not in str(raised.value)


def test_way_reference_and_tag_detachment_are_bounded(tmp_path: Path) -> None:
    path = materialize_fixture(tmp_path)
    with pytest.raises(PbfResourceLimitExceeded, match="max_way_nodes"):
        OsmiumPbfParser(PbfParseLimits(max_way_nodes=1)).parse(
            approved_input(path), CollectingSink()
        )
    with pytest.raises(PbfResourceLimitExceeded, match="max_tags_per_object"):
        OsmiumPbfParser(PbfParseLimits(max_tags_per_object=1)).parse(
            approved_input(path), CollectingSink()
        )
