from __future__ import annotations

import hashlib
import json
from dataclasses import dataclass
from typing import Any

from bson import BSON


EPICRISIS_CONTEXT_STORAGE_PART_MAX_CHARS = 512_000
EPICRISIS_PART_MARKER = "__epicrisis_context_part__"
_PARTITION_METADATA_KEYS = {
    "schema_version",
    "cache_schema_version",
    "storage_mode",
    "context_parts_generation_id",
    "context_parts_manifest",
}


@dataclass(frozen=True)
class PartitionedEpicrisisContext:
    main_context: dict[str, Any]
    parts: tuple[dict[str, Any], ...]
    generation_id: str
    main_bson_bytes: int
    partitioned_fields: tuple[str, ...]


def _json_default(value: Any) -> str:
    return str(value)


def _stable_json(value: Any) -> str:
    return json.dumps(
        value,
        ensure_ascii=False,
        sort_keys=True,
        default=_json_default,
        separators=(",", ":"),
    )


def _generation_id(context: dict[str, Any], *, username: str, case_key: str) -> str:
    payload = {
        "username": str(username or "").strip(),
        "case_key": str(case_key or "").strip(),
        "context": context,
    }
    return hashlib.sha256(_stable_json(payload).encode("utf-8")).hexdigest()[:32]


def _split_serialized_value(value: Any, *, max_chars: int) -> list[str]:
    serialized = _stable_json(value)
    chunk_size = max(1, int(max_chars))
    return [serialized[index : index + chunk_size] for index in range(0, len(serialized), chunk_size)] or ["null"]


def _part_record(
    *,
    username: str,
    case_key: str,
    generation_id: str,
    field: str,
    index: int,
    total: int,
    payload: str,
) -> dict[str, Any]:
    payload_hash = hashlib.sha256(payload.encode("utf-8")).hexdigest()
    return {
        "usuario": username,
        "case_key": case_key,
        "generation_id": generation_id,
        "part_id": f"{generation_id}:{field}:{index}",
        "field": field,
        "part_number": index,
        "part_total": total,
        "payload": payload,
        "payload_hash": payload_hash,
    }


def _marker(*, generation_id: str, field: str, total: int) -> dict[str, Any]:
    return {
        EPICRISIS_PART_MARKER: {
            "generation_id": generation_id,
            "field": field,
            "part_total": total,
        }
    }


def _context_bson_bytes(context: dict[str, Any]) -> int:
    return len(BSON.encode({"contexto": context}))


def partition_epicrisis_context(
    context: dict[str, Any],
    *,
    username: str,
    case_key: str,
    max_bson_bytes: int,
    part_max_chars: int = EPICRISIS_CONTEXT_STORAGE_PART_MAX_CHARS,
    envelope_reserve_bytes: int = 0,
) -> PartitionedEpicrisisContext:
    """Externaliza campos completos como JSON para mantener pequeño el documento principal."""
    original = dict(context or {})
    context_budget = max(1, int(max_bson_bytes) - max(0, int(envelope_reserve_bytes)))
    generation_id = _generation_id(original, username=username, case_key=case_key)
    main = dict(original)
    main["storage_mode"] = "partitioned"
    main["context_parts_generation_id"] = generation_id
    parts: list[dict[str, Any]] = []
    partitioned_fields: list[str] = []

    candidates = sorted(
        (
            (key, value)
            for key, value in original.items()
            if key not in _PARTITION_METADATA_KEYS
        ),
        key=lambda item: len(_stable_json(item[1])),
        reverse=True,
    )

    def update_manifest() -> None:
        main["context_parts_manifest"] = [
            {
                "field": field,
                "part_total": sum(1 for part in parts if part["field"] == field),
            }
            for field in partitioned_fields
        ]

    for field, value in candidates:
        if _context_bson_bytes(main) <= context_budget:
            break
        serialized_parts = _split_serialized_value(value, max_chars=part_max_chars)
        total = len(serialized_parts)
        main[field] = _marker(generation_id=generation_id, field=field, total=total)
        parts.extend(
            _part_record(
                username=username,
                case_key=case_key,
                generation_id=generation_id,
                field=field,
                index=index,
                total=total,
                payload=payload,
            )
            for index, payload in enumerate(serialized_parts)
        )
        partitioned_fields.append(field)
        update_manifest()

    update_manifest()
    main["cache_bson_bytes"] = _context_bson_bytes(main)
    main_bson_bytes = _context_bson_bytes(main)
    if main_bson_bytes > context_budget:
        raise ValueError(
            "No fue posible compactar el read model de epicrisis dentro del presupuesto BSON."
        )
    return PartitionedEpicrisisContext(
        main_context=main,
        parts=tuple(parts),
        generation_id=generation_id,
        main_bson_bytes=main_bson_bytes,
        partitioned_fields=tuple(partitioned_fields),
    )


def hydrate_partitioned_context(
    context: dict[str, Any],
    parts: list[dict[str, Any]] | tuple[dict[str, Any], ...],
) -> dict[str, Any]:
    hydrated = dict(context or {})
    by_field: dict[str, list[dict[str, Any]]] = {}
    for part in parts:
        field = str(part.get("field") or "").strip()
        if field:
            by_field.setdefault(field, []).append(part)

    for field, marker in list(hydrated.items()):
        if not isinstance(marker, dict) or EPICRISIS_PART_MARKER not in marker:
            continue
        field_parts = sorted(by_field.get(field, []), key=lambda item: int(item.get("part_number") or 0))
        if not field_parts:
            continue
        serialized = "".join(str(part.get("payload") or "") for part in field_parts)
        try:
            hydrated[field] = json.loads(serialized)
        except (TypeError, ValueError, json.JSONDecodeError):
            continue

    hydrated.pop("storage_mode", None)
    hydrated.pop("context_parts_generation_id", None)
    hydrated.pop("context_parts_manifest", None)
    return hydrated
