"""Bounded streaming PBF parser backed by pyosmium."""

from __future__ import annotations

import hashlib
import os
import stat
import zlib
from collections.abc import Callable
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, BinaryIO, Iterator, Protocol, runtime_checkable

import osmium

from ..domain.enums import OsmObjectType
from ..domain.models import OsmIdentity
from ..errors import (
    ContractViolation,
    PbfParseError,
    PbfResourceLimitExceeded,
    PbfStageContractError,
    PbfTransactionError,
    PbfValidationError,
)
from .models import (
    ParsedOsmObject,
    PbfObjectCounts,
    PbfParseLimits,
    PbfParseSummary,
    RelationMember,
    VerifiedPbfInput,
)

_CHUNK_SIZE = 1024 * 1024
_RELATION_TYPE = {
    "n": OsmObjectType.NODE,
    "w": OsmObjectType.WAY,
    "r": OsmObjectType.RELATION,
}


def _read_varint(data: bytes, offset: int) -> tuple[int, int]:
    value = 0
    for shift in range(0, 70, 7):
        if offset >= len(data):
            raise PbfValidationError("PBF protobuf varint is truncated")
        byte = data[offset]
        offset += 1
        value |= (byte & 0x7F) << shift
        if byte < 0x80:
            return value, offset
    raise PbfValidationError("PBF protobuf varint is too long")


def _protobuf_fields(data: bytes) -> Iterator[tuple[int, int, int | bytes]]:
    offset = 0
    while offset < len(data):
        key, offset = _read_varint(data, offset)
        field_number = key >> 3
        wire_type = key & 0x07
        if field_number == 0:
            raise PbfValidationError("PBF protobuf field number is invalid")
        if wire_type == 0:
            value, offset = _read_varint(data, offset)
            yield field_number, wire_type, value
        elif wire_type == 1:
            end = offset + 8
            if end > len(data):
                raise PbfValidationError("PBF protobuf fixed64 field is truncated")
            yield field_number, wire_type, data[offset:end]
            offset = end
        elif wire_type == 2:
            length, offset = _read_varint(data, offset)
            end = offset + length
            if end > len(data):
                raise PbfValidationError("PBF protobuf bytes field is truncated")
            yield field_number, wire_type, data[offset:end]
            offset = end
        elif wire_type == 5:
            end = offset + 4
            if end > len(data):
                raise PbfValidationError("PBF protobuf fixed32 field is truncated")
            yield field_number, wire_type, data[offset:end]
            offset = end
        else:
            raise PbfValidationError("PBF protobuf wire type is unsupported")


def _parse_blob_header(data: bytes) -> tuple[str, int]:
    blob_type: str | None = None
    data_size: int | None = None
    for field_number, wire_type, value in _protobuf_fields(data):
        if field_number == 1:
            if wire_type != 2 or not isinstance(value, bytes):
                raise PbfValidationError("PBF BlobHeader type field is invalid")
            if blob_type is not None:
                raise PbfValidationError("PBF BlobHeader type field is duplicated")
            try:
                blob_type = value.decode("ascii")
            except UnicodeDecodeError:
                raise PbfValidationError("PBF BlobHeader type is not ASCII") from None
        elif field_number == 3:
            if wire_type != 0 or not isinstance(value, int):
                raise PbfValidationError("PBF BlobHeader datasize field is invalid")
            if data_size is not None:
                raise PbfValidationError("PBF BlobHeader datasize field is duplicated")
            data_size = value
    if blob_type is None or data_size is None or data_size <= 0:
        raise PbfValidationError("PBF BlobHeader is incomplete")
    return blob_type, data_size


def _validate_blob(data: bytes, limits: PbfParseLimits) -> int:
    raw_size: int | None = None
    payload_type: int | None = None
    payload: bytes | None = None
    for field_number, wire_type, value in _protobuf_fields(data):
        if field_number == 2:
            if wire_type != 0 or not isinstance(value, int):
                raise PbfValidationError("PBF Blob raw_size field is invalid")
            if raw_size is not None:
                raise PbfValidationError("PBF Blob raw_size field is duplicated")
            raw_size = value
        elif field_number in {1, 3}:
            if wire_type != 2 or not isinstance(value, bytes):
                raise PbfValidationError("PBF Blob payload field is invalid")
            if payload is not None:
                raise PbfValidationError("PBF Blob must contain exactly one approved payload")
            payload_type = field_number
            payload = value
        elif field_number in {4, 5, 6, 7}:
            raise PbfValidationError("PBF Blob compression is not approved in Phase 2A.3")

    if payload_type is None or payload is None:
        raise PbfValidationError("PBF Blob must contain exactly one approved payload")
    if payload_type == 1:
        if raw_size is not None and raw_size != len(payload):
            raise PbfValidationError("raw PBF Blob size metadata is inconsistent")
        uncompressed_size = len(payload)
    else:
        if raw_size is None or raw_size <= 0:
            raise PbfValidationError("compressed PBF Blob must declare raw_size")
        if raw_size > limits.max_uncompressed_blob_bytes:
            raise PbfResourceLimitExceeded("PBF Blob exceeds max_uncompressed_blob_bytes")
        try:
            decompressor = zlib.decompressobj()
            uncompressed = decompressor.decompress(payload, limits.max_uncompressed_blob_bytes + 1)
            if (
                decompressor.unconsumed_tail
                or len(uncompressed) > limits.max_uncompressed_blob_bytes
            ):
                raise PbfResourceLimitExceeded("PBF Blob exceeds max_uncompressed_blob_bytes")
            remaining = limits.max_uncompressed_blob_bytes + 1 - len(uncompressed)
            uncompressed += decompressor.flush(remaining)
        except zlib.error:
            raise PbfValidationError("compressed PBF Blob is invalid") from None
        if (
            len(uncompressed) > limits.max_uncompressed_blob_bytes
            or not decompressor.eof
            or decompressor.unused_data
        ):
            raise PbfValidationError("compressed PBF Blob framing is invalid")
        if len(uncompressed) != raw_size:
            raise PbfValidationError("compressed PBF Blob raw_size is inconsistent")
        uncompressed_size = len(uncompressed)
    if uncompressed_size > limits.max_uncompressed_blob_bytes:
        raise PbfResourceLimitExceeded("PBF Blob exceeds max_uncompressed_blob_bytes")
    return uncompressed_size


def _read_exact(handle: BinaryIO, size: int, description: str) -> bytes:
    data = handle.read(size)
    if len(data) != size:
        raise PbfValidationError(f"{description} is truncated")
    return data


def _validate_pbf_envelope(path: Path, limits: PbfParseLimits) -> None:
    seen_header = False
    data_blocks = 0
    blob_count = 0
    total_uncompressed_bytes = 0
    try:
        with path.open("rb") as handle:
            while True:
                prefix = handle.read(4)
                if prefix == b"":
                    break
                blob_count += 1
                if blob_count > limits.max_blob_count:
                    raise PbfResourceLimitExceeded("PBF Blob count exceeds max_blob_count")
                if len(prefix) != 4:
                    raise PbfValidationError("PBF BlobHeader length prefix is truncated")
                header_size = int.from_bytes(prefix, byteorder="big", signed=False)
                if header_size <= 0 or header_size > limits.max_blob_header_bytes:
                    raise PbfResourceLimitExceeded("PBF BlobHeader exceeds bounded size")
                header = _read_exact(handle, header_size, "PBF BlobHeader")
                blob_type, blob_size = _parse_blob_header(header)
                if blob_size > limits.max_blob_size_bytes:
                    raise PbfResourceLimitExceeded("PBF Blob exceeds max_blob_size_bytes")
                blob = _read_exact(handle, blob_size, "PBF Blob")
                total_uncompressed_bytes += _validate_blob(blob, limits)
                if total_uncompressed_bytes > limits.max_total_uncompressed_bytes:
                    raise PbfResourceLimitExceeded(
                        "PBF cumulative expansion exceeds max_total_uncompressed_bytes"
                    )

                if blob_type == "OSMHeader":
                    if seen_header or data_blocks:
                        raise PbfValidationError("PBF OSMHeader position is invalid")
                    seen_header = True
                elif blob_type == "OSMData":
                    if not seen_header:
                        raise PbfValidationError("PBF OSMData appears before OSMHeader")
                    data_blocks += 1
                else:
                    raise PbfValidationError("PBF contains an unsupported Blob type")
    except OSError as exc:
        raise PbfValidationError("PBF envelope cannot be read") from exc
    if not seen_header or data_blocks == 0:
        raise PbfValidationError("PBF must contain OSMHeader and OSMData blocks")


@runtime_checkable
class PbfObjectSink(Protocol):
    """Transactional streaming sink for detached parsed objects."""

    def accept(self, parsed_object: ParsedOsmObject) -> None:
        """Stage one detached node, way, or relation."""

    def finish(self, summary: PbfParseSummary) -> None:
        """Finalize staged objects only after a complete, stable parse."""

    def abort(self) -> None:
        """Discard staged objects after any parsing or sink failure."""


def _abort_sink_or_raise(sink: PbfObjectSink, failure: Exception) -> None:
    try:
        sink.abort()
    except Exception as abort_error:
        failures = ExceptionGroup(
            "PBF operation and sink abort both failed",
            [failure, abort_error],
        )
        raise PbfTransactionError("PBF operation failed and sink abort failed") from failures


def _utc_now() -> datetime:
    return datetime.now(timezone.utc)


def _regular_file_stat(path: Path) -> os.stat_result:
    try:
        result = path.lstat()
    except OSError as exc:
        raise PbfValidationError("PBF path cannot be inspected") from exc
    if stat.S_ISLNK(result.st_mode):
        raise PbfValidationError("PBF path cannot be a symbolic link")
    if not stat.S_ISREG(result.st_mode):
        raise PbfValidationError("PBF path must be a regular file")
    return result


def _stat_identity(result: os.stat_result) -> tuple[int, int, int, int, int]:
    return (
        result.st_dev,
        result.st_ino,
        result.st_size,
        result.st_mtime_ns,
        result.st_ctime_ns,
    )


def _hash_file(path: Path, max_size_bytes: int) -> tuple[str, int, os.stat_result]:
    before = _regular_file_stat(path)
    if before.st_size <= 0:
        raise PbfValidationError("PBF file must not be empty")
    if before.st_size > max_size_bytes:
        raise PbfResourceLimitExceeded("PBF file exceeds max_file_size_bytes")

    digest = hashlib.sha256()
    size = 0
    try:
        with path.open("rb") as handle:
            while chunk := handle.read(_CHUNK_SIZE):
                size += len(chunk)
                if size > max_size_bytes:
                    raise PbfResourceLimitExceeded("PBF file exceeds max_file_size_bytes")
                digest.update(chunk)
    except OSError as exc:
        raise PbfValidationError("PBF file cannot be read") from exc

    after = _regular_file_stat(path)
    if _stat_identity(before) != _stat_identity(after):
        raise PbfValidationError("PBF file changed during validation")
    if size != before.st_size:
        raise PbfValidationError("PBF byte count does not match file metadata")
    return digest.hexdigest(), size, after


class _BoundedHandler(osmium.SimpleHandler):
    def __init__(self, sink: PbfObjectSink, limits: PbfParseLimits) -> None:
        super().__init__()
        self._sink = sink
        self._limits = limits
        self.nodes = 0
        self.ways = 0
        self.relations = 0
        self.total_tags = 0

    @property
    def total_objects(self) -> int:
        return self.nodes + self.ways + self.relations

    def _check_object_slot(self) -> None:
        if self.total_objects >= self._limits.max_objects:
            raise PbfResourceLimitExceeded("PBF object count exceeds max_objects")

    def _tags(self, obj: Any) -> tuple[tuple[str, str], ...]:
        tags: list[tuple[str, str]] = []
        for tag in obj.tags:
            if len(tags) >= self._limits.max_tags_per_object:
                raise PbfResourceLimitExceeded("PBF object exceeds max_tags_per_object")
            if self.total_tags + len(tags) + 1 > self._limits.max_total_tags:
                raise PbfResourceLimitExceeded("PBF tag count exceeds max_total_tags")
            tags.append((str(tag.k), str(tag.v)))
        return tuple(tags)

    def _accept(self, parsed_object: ParsedOsmObject) -> None:
        try:
            self._sink.accept(parsed_object)
        except ContractViolation as exc:
            raise PbfStageContractError("sink_accept", exc) from exc

    def node(self, node: Any) -> None:
        self._check_object_slot()
        tags = self._tags(node)
        try:
            longitude_e7 = round(float(node.location.lon) * 10_000_000)
            latitude_e7 = round(float(node.location.lat) * 10_000_000)
        except (AttributeError, RuntimeError, ValueError) as exc:
            raise PbfParseError("node has no valid location") from exc
        try:
            parsed = ParsedOsmObject(
                identity=OsmIdentity(osm_type=OsmObjectType.NODE, osm_id=int(node.id)),
                version=int(node.version),
                tags=tags,
                longitude_e7=longitude_e7,
                latitude_e7=latitude_e7,
            )
        except ContractViolation as exc:
            raise PbfStageContractError("parsed_object", exc) from exc
        self._accept(parsed)
        self.nodes += 1
        self.total_tags += len(tags)

    def way(self, way: Any) -> None:
        self._check_object_slot()
        tags = self._tags(way)
        refs: list[int] = []
        for node in way.nodes:
            if len(refs) >= self._limits.max_way_nodes:
                raise PbfResourceLimitExceeded("way exceeds max_way_nodes")
            refs.append(int(node.ref))
        try:
            parsed = ParsedOsmObject(
                identity=OsmIdentity(osm_type=OsmObjectType.WAY, osm_id=int(way.id)),
                version=int(way.version),
                tags=tags,
                node_refs=tuple(refs),
            )
        except ContractViolation as exc:
            raise PbfStageContractError("parsed_object", exc) from exc
        self._accept(parsed)
        self.ways += 1
        self.total_tags += len(tags)

    def relation(self, relation: Any) -> None:
        self._check_object_slot()
        tags = self._tags(relation)
        members: list[RelationMember] = []
        try:
            for member in relation.members:
                if len(members) >= self._limits.max_relation_members:
                    raise PbfResourceLimitExceeded("relation exceeds max_relation_members")
                member_type = _RELATION_TYPE.get(str(member.type))
                if member_type is None:
                    raise PbfParseError("relation contains an unsupported member type")
                members.append(
                    RelationMember(
                        member_type=member_type,
                        ref=int(member.ref),
                        role=str(member.role),
                    )
                )
            parsed = ParsedOsmObject(
                identity=OsmIdentity(osm_type=OsmObjectType.RELATION, osm_id=int(relation.id)),
                version=int(relation.version),
                tags=tags,
                members=tuple(members),
            )
        except ContractViolation as exc:
            raise PbfStageContractError("parsed_object", exc) from exc
        self._accept(parsed)
        self.relations += 1
        self.total_tags += len(tags)


class OsmiumPbfParser:
    """Parse a verified local PBF through a bounded streaming sink."""

    def __init__(
        self,
        limits: PbfParseLimits | None = None,
        *,
        clock: Callable[[], datetime] = _utc_now,
    ) -> None:
        self._limits = PbfParseLimits() if limits is None else limits
        self._clock = clock

    def parse(self, approved_input: VerifiedPbfInput, sink: PbfObjectSink) -> PbfParseSummary:
        if not isinstance(approved_input, VerifiedPbfInput):
            raise PbfValidationError("approved_input must be VerifiedPbfInput")
        if not isinstance(sink, PbfObjectSink):
            raise PbfValidationError("sink must implement PbfObjectSink")

        path = approved_input.local_path
        content_sha256, size_bytes, validated_stat = _hash_file(
            path, self._limits.max_file_size_bytes
        )
        if size_bytes != approved_input.size_bytes:
            raise PbfValidationError("PBF size does not match approved input evidence")
        if content_sha256 != approved_input.content_sha256:
            raise PbfValidationError("PBF checksum does not match approved input evidence")
        _validate_pbf_envelope(path, self._limits)
        started_at = self._clock()
        handler = _BoundedHandler(sink, self._limits)
        try:
            handler.apply_file(str(path), locations=False)
            completed_at = self._clock()

            after = _regular_file_stat(path)
            if _stat_identity(validated_stat) != _stat_identity(after):
                raise PbfValidationError("PBF file changed during parsing")

            try:
                summary = PbfParseSummary(
                    content_sha256=content_sha256,
                    size_bytes=size_bytes,
                    counts=PbfObjectCounts(
                        nodes=handler.nodes,
                        ways=handler.ways,
                        relations=handler.relations,
                    ),
                    total_tags=handler.total_tags,
                    parser_started_at=started_at,
                    parser_completed_at=completed_at,
                )
            except ContractViolation as exc:
                raise PbfStageContractError("parse_summary", exc) from exc
            try:
                sink.finish(summary)
            except ContractViolation as exc:
                raise PbfStageContractError("sink_finish", exc) from exc
            return summary
        except (PbfParseError, PbfResourceLimitExceeded, PbfValidationError) as exc:
            _abort_sink_or_raise(sink, exc)
            raise
        except Exception as exc:
            _abort_sink_or_raise(sink, exc)
            raise PbfParseError("PBF parsing failed") from exc
