"""Define base classes for serialization."""

import logging
import re
import sys
import warnings
from abc import abstractmethod
from collections.abc import Iterable
from functools import cached_property
from pathlib import Path
from typing import Annotated, Any, Optional, Union

from pydantic import (
    AnyUrl,
    BaseModel,
    ConfigDict,
    Field,
    NonNegativeInt,
    computed_field,
)
from typing_extensions import Self, override

from docling_core.transforms.serializer.base import (
    BaseAnnotationSerializer,
    BaseDocSerializer,
    BaseFallbackSerializer,
    BaseFormSerializer,
    BaseInlineSerializer,
    BaseKeyValueSerializer,
    BaseListSerializer,
    BaseMetaSerializer,
    BasePictureSerializer,
    BaseTableSerializer,
    BaseTextSerializer,
    SerializationResult,
    Span,
)
from docling_core.types.doc import (
    ContentLayer,
    DescriptionAnnotation,
    DocItem,
    DocItemLabel,
    DoclingDocument,
    FloatingItem,
    Formatting,
    FormItem,
    InlineGroup,
    KeyValueItem,
    ListGroup,
    NodeItem,
    PictureClassificationData,
    PictureDataType,
    PictureItem,
    PictureMoleculeData,
    Script,
    TableAnnotationType,
    TableItem,
    TextItem,
)
from docling_core.types.doc.document import DOCUMENT_TOKENS_EXPORT_LABELS

_DEFAULT_LABELS = DOCUMENT_TOKENS_EXPORT_LABELS
_DEFAULT_LAYERS = set(ContentLayer)


_logger = logging.getLogger(__name__)


class _PageBreakNode(NodeItem):
    """Page break node."""

    prev_page: int
    next_page: int


class _PageBreakSerResult(SerializationResult):
    """Page break serialization result."""

    node: _PageBreakNode


def _iterate_items(
    doc: DoclingDocument,
    layers: set[ContentLayer] | None,
    node: NodeItem | None = None,
    traverse_pictures: bool = False,
    add_page_breaks: bool = False,
    visited: set[str] | None = None,
) -> Iterable[tuple[NodeItem, int]]:
    my_visited: set[str] = visited if visited is not None else set()
    prev_page_nr: int | None = None
    page_break_i = 0
    for item, lvl in doc.iterate_items(
        root=node,
        with_groups=True,
        included_content_layers=layers,
        traverse_pictures=traverse_pictures,
    ):
        if add_page_breaks:
            if isinstance(item, ListGroup | InlineGroup) and item.self_ref not in my_visited:
                # if group starts with new page, yield page break before group node
                my_visited.add(item.self_ref)
                for it, _ in _iterate_items(
                    doc=doc,
                    layers=layers,
                    node=item,
                    traverse_pictures=traverse_pictures,
                    add_page_breaks=add_page_breaks,
                    visited=my_visited,
                ):
                    if isinstance(it, DocItem) and it.prov:
                        page_no = it.prov[0].page_no
                        if prev_page_nr is not None and page_no > prev_page_nr:
                            yield (
                                _PageBreakNode(
                                    self_ref=f"#/pb/{page_break_i}",
                                    prev_page=prev_page_nr,
                                    next_page=page_no,
                                ),
                                lvl,
                            )
                        break
            elif isinstance(item, DocItem) and item.prov:
                page_no = item.prov[0].page_no
                if prev_page_nr is None or page_no > prev_page_nr:
                    if prev_page_nr is not None:  # close previous range
                        yield (
                            _PageBreakNode(
                                self_ref=f"#/pb/{page_break_i}",
                                prev_page=prev_page_nr,
                                next_page=page_no,
                            ),
                            lvl,
                        )
                        page_break_i += 1
                    prev_page_nr = page_no
        yield item, lvl


def _get_annotation_text(
    annotation: PictureDataType | TableAnnotationType,
) -> str | None:
    result = None
    if isinstance(annotation, PictureClassificationData):
        predicted_class = annotation.predicted_classes[0].class_name if annotation.predicted_classes else None
        if predicted_class is not None:
            result = predicted_class.replace("_", " ")
    elif isinstance(annotation, DescriptionAnnotation):
        result = annotation.text
    elif isinstance(annotation, PictureMoleculeData):
        result = annotation.smi
    return result


def create_ser_result(
    *,
    text: str = "",
    span_source: DocItem | list[SerializationResult] = [],
) -> SerializationResult:
    """Function for creating `SerializationResult` instances.

    Args:
        text: the text the use. Defaults to "".
        span_source: the item or list of results to use as span source. Defaults to [].

    Returns:
        The created `SerializationResult`.
    """
    spans: list[Span]
    if isinstance(span_source, DocItem):
        spans = [Span(item=span_source)]
    else:
        results: list[SerializationResult] = span_source
        spans = []
        span_ids: set[str] = set()
        for ser_res in results:
            for span in ser_res.spans:
                if (span_id := span.item.self_ref) not in span_ids:
                    span_ids.add(span_id)
                    spans.append(span)
    return SerializationResult(
        text=text,
        spans=spans,
    )


class CommonParams(BaseModel):
    """Common serialization parameters."""

    # allowlists with non-recursive semantics, i.e. if a list group node is outside the
    # range and some of its children items are within, they will be serialized
    labels: set[DocItemLabel] = _DEFAULT_LABELS
    layers: set[ContentLayer] = _DEFAULT_LAYERS
    pages: set[int] | None = None  # None means all pages are allowed

    # slice-like semantics: start is included, stop is excluded
    start_idx: NonNegativeInt = 0
    stop_idx: NonNegativeInt = sys.maxsize

    include_non_meta: bool = True

    include_formatting: bool = True
    include_hyperlinks: bool = True
    caption_delim: str = " "
    use_legacy_annotations: bool = Field(
        default=False,
        description="Use legacy annotation serialization.",
        deprecated="Ignored field; legacy annotations considered only when meta not present.",
    )
    allowed_meta_names: set[str] | None = Field(
        default=None,
        description="Meta name to allow; None means all meta names are allowed.",
    )
    blocked_meta_names: set[str] = Field(
        default_factory=set,
        description="Meta name to block; takes precedence over allowed_meta_names.",
    )
    traverse_pictures: Annotated[
        bool,
        Field(
            description="Whether to traverse into picture objects to serialize their children (e.g., text objects).",
        ),
    ] = False

    def merge_with_patch(self, patch: dict[str, Any]) -> Self:
        """Create an instance by merging the provided patch dict on top of self."""
        res = self.model_copy(update=patch)
        return res


class DocSerializer(BaseModel, BaseDocSerializer):
    """Class for document serializers."""

    model_config = ConfigDict(arbitrary_types_allowed=True, extra="forbid")

    doc: DoclingDocument

    text_serializer: BaseTextSerializer
    table_serializer: BaseTableSerializer
    picture_serializer: BasePictureSerializer
    key_value_serializer: BaseKeyValueSerializer
    form_serializer: BaseFormSerializer
    fallback_serializer: BaseFallbackSerializer

    list_serializer: BaseListSerializer
    inline_serializer: BaseInlineSerializer

    meta_serializer: BaseMetaSerializer | None = None
    annotation_serializer: BaseAnnotationSerializer

    params: CommonParams = CommonParams()

    _excluded_refs_cache: dict[str, set[str]] = {}

    @computed_field  # type: ignore[misc]
    @cached_property
    def _captions_of_some_item(self) -> set[str]:
        layers = set(ContentLayer)  # TODO review
        refs = {
            cap.cref
            for (item, _) in self.doc.iterate_items(
                with_groups=True,
                traverse_pictures=True,
                included_content_layers=layers,
            )
            for cap in (item.captions if isinstance(item, FloatingItem) else [])
        }
        return refs

    @computed_field  # type: ignore[misc]
    @cached_property
    def _footnotes_of_some_item(self) -> set[str]:
        layers = set(ContentLayer)  # TODO review
        refs = {
            ftn.cref
            for (item, _) in self.doc.iterate_items(
                with_groups=True,
                traverse_pictures=True,
                included_content_layers=layers,
            )
            for ftn in (item.footnotes if isinstance(item, FloatingItem) else [])
        }
        return refs

    @override
    def get_excluded_refs(self, **kwargs: Any) -> set[str]:
        """References to excluded items."""
        params = self.params.merge_with_patch(patch=kwargs)
        params_json = params.model_dump_json()
        refs = self._excluded_refs_cache.get(params_json)
        if refs is None:
            refs = {
                item.self_ref
                for ix, (item, _) in enumerate(
                    _iterate_items(
                        doc=self.doc,
                        traverse_pictures=True,
                        layers=params.layers,
                    )
                )
                if (
                    (ix < params.start_idx or ix >= params.stop_idx)
                    or (
                        isinstance(item, DocItem)
                        and (
                            item.label not in params.labels
                            or item.content_layer not in params.layers
                            or (
                                params.pages is not None
                                and ((not item.prov) or item.prov[0].page_no not in params.pages)
                            )
                        )
                    )
                )
            }
            self._excluded_refs_cache[params_json] = refs
        return refs

    @abstractmethod
    def serialize_doc(
        self,
        *,
        parts: list[SerializationResult],
        **kwargs: Any,
    ) -> SerializationResult:
        """Serialize a document out of its parts."""
        ...

    def _serialize_body(self, **kwargs) -> SerializationResult:
        """Serialize the document body."""
        subparts = self.get_parts(**kwargs)
        res = self.serialize_doc(parts=subparts, **kwargs)
        return res

    def _meta_is_wrapped(self) -> bool:
        return False

    def _item_wraps_meta(self, item: NodeItem) -> bool:
        """Whether the item's serializer handles meta wrapping internally."""
        return False

    def _meta_position(self) -> str:
        """Where the meta block should be placed relative to the item content.

        Returns "after" (default) or "before".
        """
        return "after"

    @override
    def serialize(
        self,
        *,
        item: NodeItem | None = None,
        list_level: int = 0,
        is_inline_scope: bool = False,
        visited: set[str] | None = None,  # refs of visited items
        **kwargs: Any,
    ) -> SerializationResult:
        """Serialize a given node."""
        my_visited: set[str] = visited if visited is not None else set()
        parts: list[SerializationResult] = []
        delim: str = kwargs.get("delim", "\n")
        my_params = self.params.model_copy(update=kwargs)
        my_kwargs = {**self.params.model_dump(), **kwargs}
        empty_res = create_ser_result()

        my_item = item or self.doc.body

        meta_position = self._meta_position()

        if my_item == self.doc.body:
            body_meta_part: SerializationResult | None = None
            if my_item.meta and not self._meta_is_wrapped():
                candidate = self.serialize_meta(item=my_item, **my_kwargs)
                if candidate.text:
                    body_meta_part = candidate
                    if meta_position == "before":
                        parts.append(body_meta_part)

            if my_item.self_ref not in my_visited:
                my_visited.add(my_item.self_ref)
                part = self._serialize_body(**my_kwargs)
                if part.text:
                    parts.append(part)
                if body_meta_part is not None and meta_position == "after":
                    parts.append(body_meta_part)
                return create_ser_result(
                    text=delim.join([p.text for p in parts if p.text]),
                    span_source=parts,
                )
            else:
                return empty_res

        my_visited.add(my_item.self_ref)

        meta_part: SerializationResult | None = None
        if my_item.meta and not self._meta_is_wrapped() and not self._item_wraps_meta(my_item):
            candidate = self.serialize_meta(item=my_item, **my_kwargs)
            if candidate.text:
                meta_part = candidate
                if meta_position == "before":
                    parts.append(meta_part)

        if my_params.include_non_meta:
            ########
            # groups
            ########
            if isinstance(my_item, ListGroup):
                part = self.list_serializer.serialize(
                    item=my_item,
                    doc_serializer=self,
                    doc=self.doc,
                    list_level=list_level,
                    is_inline_scope=is_inline_scope,
                    visited=my_visited,
                    **my_kwargs,
                )
            elif isinstance(my_item, InlineGroup):
                part = self.inline_serializer.serialize(
                    item=my_item,
                    doc_serializer=self,
                    doc=self.doc,
                    list_level=list_level,
                    visited=my_visited,
                    **my_kwargs,
                )
            ###########
            # doc items
            ###########
            elif isinstance(my_item, TextItem):
                if my_item.self_ref in self._captions_of_some_item:
                    # those captions will be handled by the floating item holding them
                    return empty_res
                elif my_item.self_ref in self._footnotes_of_some_item:
                    # those footnotes will be handled by the floating item holding them
                    return empty_res
                else:
                    part = (
                        self.text_serializer.serialize(
                            item=my_item,
                            doc_serializer=self,
                            doc=self.doc,
                            is_inline_scope=is_inline_scope,
                            visited=my_visited,
                            **my_kwargs,
                        )
                        if my_item.self_ref not in self.get_excluded_refs(**kwargs)
                        else empty_res
                    )
            elif isinstance(my_item, TableItem):
                part = self.table_serializer.serialize(
                    item=my_item,
                    doc_serializer=self,
                    doc=self.doc,
                    visited=my_visited,
                    **my_kwargs,
                )
            elif isinstance(my_item, PictureItem):
                part = self.picture_serializer.serialize(
                    item=my_item,
                    doc_serializer=self,
                    doc=self.doc,
                    visited=my_visited,
                    **my_kwargs,
                )
            elif isinstance(my_item, KeyValueItem):
                part = self.key_value_serializer.serialize(
                    item=my_item,
                    doc_serializer=self,
                    doc=self.doc,
                    **my_kwargs,
                )
            elif isinstance(my_item, FormItem):
                part = self.form_serializer.serialize(
                    item=my_item,
                    doc_serializer=self,
                    doc=self.doc,
                    **my_kwargs,
                )
            elif isinstance(my_item, _PageBreakNode):
                # returned as such, so that the combining scope can tell a page
                # break apart from ordinary content
                return _PageBreakSerResult(
                    text=self._create_page_break(node=my_item),
                    node=my_item,
                )
            else:
                part = self.fallback_serializer.serialize(
                    item=my_item,
                    doc_serializer=self,
                    doc=self.doc,
                    visited=my_visited,
                    **my_kwargs,
                )
            parts.append(part)

        if meta_part is not None and meta_position == "after":
            parts.append(meta_part)

        return create_ser_result(text=delim.join([p.text for p in parts if p.text]), span_source=parts)

    # making some assumptions about the kwargs it can pass
    @override
    def get_parts(
        self,
        item: NodeItem | None = None,
        *,
        traverse_pictures: bool = False,
        list_level: int = 0,
        is_inline_scope: bool = False,
        visited: set[str] | None = None,  # refs of visited items
        **kwargs: Any,
    ) -> list[SerializationResult]:
        """Get the components to be combined for serializing this node."""
        parts: list[SerializationResult] = []
        my_visited: set[str] = visited if visited is not None else set()
        params = self.params.merge_with_patch(patch=kwargs)
        add_content = True

        if hasattr(params, "add_content"):
            add_content = getattr(params, "add_content")

        for node, lvl in _iterate_items(
            node=item,
            doc=self.doc,
            layers=params.layers,
            traverse_pictures=params.traverse_pictures,
            add_page_breaks=self.requires_page_break(),
        ):
            if node.self_ref in my_visited:
                continue
            else:
                my_visited.add(node.self_ref)

            part = self.serialize(
                item=node,
                list_level=list_level,
                is_inline_scope=is_inline_scope,
                visited=my_visited,
                **(dict(level=lvl) | kwargs),
            )
            if part.text or not add_content:
                parts.append(part)

        return parts

    @override
    def post_process(
        self,
        text: str,
        *,
        formatting: Formatting | None = None,
        hyperlink: AnyUrl | Path | None = None,
        **kwargs: Any,
    ) -> str:
        """Apply some text post-processing steps."""
        params = self.params.merge_with_patch(patch=kwargs)
        res = text
        if params.include_formatting and formatting:
            if formatting.bold:
                res = self.serialize_bold(text=res, **kwargs)
            if formatting.italic:
                res = self.serialize_italic(text=res, **kwargs)
            if formatting.underline:
                res = self.serialize_underline(text=res, **kwargs)
            if formatting.strikethrough:
                res = self.serialize_strikethrough(text=res, **kwargs)
            if formatting.script == Script.SUB:
                res = self.serialize_subscript(text=res, **kwargs)
            elif formatting.script == Script.SUPER:
                res = self.serialize_superscript(text=res, **kwargs)
        if params.include_hyperlinks and hyperlink:
            res = self.serialize_hyperlink(text=res, hyperlink=hyperlink, **kwargs)
        return res

    @override
    def serialize_bold(self, text: str, **kwargs: Any) -> str:
        """Hook for bold formatting serialization."""
        return text

    @override
    def serialize_italic(self, text: str, **kwargs: Any) -> str:
        """Hook for italic formatting serialization."""
        return text

    @override
    def serialize_underline(self, text: str, **kwargs: Any) -> str:
        """Hook for underline formatting serialization."""
        return text

    @override
    def serialize_strikethrough(self, text: str, **kwargs: Any) -> str:
        """Hook for strikethrough formatting serialization."""
        return text

    @override
    def serialize_subscript(self, text: str, **kwargs: Any) -> str:
        """Hook for subscript formatting serialization."""
        return text

    @override
    def serialize_superscript(self, text: str, **kwargs: Any) -> str:
        """Hook for superscript formatting serialization."""
        return text

    @override
    def serialize_hyperlink(
        self,
        text: str,
        hyperlink: AnyUrl | Path,
        **kwargs: Any,
    ) -> str:
        """Hook for hyperlink serialization."""
        return text

    @override
    def serialize_captions(
        self,
        item: FloatingItem,
        **kwargs: Any,
    ) -> SerializationResult:
        """Serialize the item's captions."""
        params = self.params.merge_with_patch(patch=kwargs)
        results: list[SerializationResult] = []
        if DocItemLabel.CAPTION in params.labels:
            results = [
                create_ser_result(text=it.text, span_source=it)
                for cap in item.captions
                if isinstance(it := cap.resolve(self.doc), TextItem)
                and it.self_ref not in self.get_excluded_refs(**kwargs)
            ]
            text_res = params.caption_delim.join([r.text for r in results])
            text_res = self.post_process(text=text_res)
        else:
            text_res = ""
        return create_ser_result(text=text_res, span_source=results)

    @override
    def serialize_footnotes(
        self,
        item: FloatingItem,
        **kwargs: Any,
    ) -> SerializationResult:
        """Serialize the item's footnotes."""
        params = self.params.merge_with_patch(patch=kwargs)
        results: list[SerializationResult] = []
        if DocItemLabel.FOOTNOTE in params.labels:
            results = [
                create_ser_result(text=it.text, span_source=it)
                for ftn in item.footnotes
                if isinstance(it := ftn.resolve(self.doc), TextItem)
                and it.self_ref not in self.get_excluded_refs(**kwargs)
            ]
            # FIXME: using the caption_delimiter for now ...
            text_res = params.caption_delim.join([r.text for r in results])
            text_res = self.post_process(text=text_res)
        else:
            text_res = ""
        return create_ser_result(text=text_res, span_source=results)

    @override
    def serialize_meta(
        self,
        item: NodeItem,
        **kwargs: Any,
    ) -> SerializationResult:
        """Serialize the item's meta."""
        if self.meta_serializer:
            if item.self_ref not in self.get_excluded_refs(**kwargs):
                return self.meta_serializer.serialize(
                    item=item,
                    doc=self.doc,
                    **(self.params.model_dump() | kwargs),
                )
            else:
                return create_ser_result(text="", span_source=item if isinstance(item, DocItem) else [])
        else:
            return create_ser_result(text="", span_source=item if isinstance(item, DocItem) else [])

    # TODO deprecate
    @override
    def serialize_annotations(
        self,
        item: DocItem,
        **kwargs: Any,
    ) -> SerializationResult:
        """Serialize the item's annotations."""
        return self.annotation_serializer.serialize(
            item=item,
            doc=self.doc,
            **kwargs,
        )

    def _get_applicable_pages(self) -> list[int] | None:
        pages = {
            item.prov[0].page_no: ...
            for ix, (item, _) in enumerate(
                self.doc.iterate_items(
                    with_groups=True,
                    included_content_layers=self.params.layers,
                    traverse_pictures=True,
                )
            )
            if (
                isinstance(item, DocItem)
                and item.prov
                and (self.params.pages is None or item.prov[0].page_no in self.params.pages)
                and ix >= self.params.start_idx
                and ix < self.params.stop_idx
            )
        }
        return list(pages) or None

    def _create_page_break(self, node: _PageBreakNode) -> str:
        return f"#_#_DOCLING_DOC_PAGE_BREAK_{node.prev_page}_{node.next_page}_#_#"

    def _get_page_breaks(self, text: str) -> Iterable[tuple[str, int, int]]:
        pattern = r"#_#_DOCLING_DOC_PAGE_BREAK_(\d+)_(\d+)_#_#"
        matches = re.finditer(pattern, text)
        for match in matches:
            full_match = match.group(0)
            prev_page_nr = int(match.group(1))
            next_page_nr = int(match.group(2))
            yield (full_match, prev_page_nr, next_page_nr)


def _should_use_legacy_annotations(
    *,
    params: CommonParams,
    item: PictureItem | TableItem,
    kind: str | None = None,
) -> bool:
    if item.meta:
        return False
    with warnings.catch_warnings(record=True) as caught_warnings:
        warnings.simplefilter("ignore", DeprecationWarning)
        if (incl_attr := getattr(params, "include_annotations", None)) is not None and not incl_attr:
            return False
        use_legacy = bool([ann for ann in item.annotations if ((ann.kind == kind) if kind is not None else True)])
        if use_legacy:
            for w in caught_warnings:
                warnings.warn(w.message, w.category)
        return use_legacy
