from __future__ import annotations

from dataclasses import dataclass, field
from typing import Any


@dataclass
class tkPlotData:
    label: str
    plot_type: str
    axis: Any
    data: list[Any] = field(default_factory=list)
    xlist: Any = None
    xlabels: Any = None
    axis_scale: Any = None
    metadata: dict[str, Any] = field(default_factory=dict)

    def as_legacy_dict(self):
        result = dict(self.metadata)
        result.update(
            label=self.label,
            plot_type=self.plot_type,
            axis=self.axis,
            data=self.data,
            xlist=self.xlist,
            xlabels=self.xlabels,
            axis_scale=self.axis_scale,
        )
        return result


class tkDataSetRegistry:
    def __init__(self):
        self.items: list[dict] = []

    @staticmethod
    def _normalize_artists(axis_inf):
        plot_type = axis_inf.get("plot_type", "2D")
        data = axis_inf.get("data")
        if data is None:
            if plot_type in ("2D", "plot"):
                data = list(axis_inf["axis"].lines)
            elif plot_type == "scatter":
                data = list(axis_inf["axis"].collections)
            else:
                data = []
        elif not isinstance(data, (list, tuple)):
            data = [data]
        else:
            data = list(data)
        return data

    def add(self, axis_inf=None):
        if axis_inf is None:
            axis_inf = {}
        required = ("label", "plot_type", "axis")
        missing = [name for name in required if name not in axis_inf]
        if missing:
            raise ValueError(f"Missing axis information: {', '.join(missing)}")

        item = dict(axis_inf)
        axis = item["axis"]
        item["xlabel"] = item.get("xlabel", axis.get_xlabel())
        item["ylabel"] = item.get("ylabel", axis.get_ylabel())
        item["data"] = self._normalize_artists(item)

        for artist in item["data"]:
            if hasattr(artist, "get_xdata") and hasattr(artist, "get_ydata"):
                artist.xdata = artist.get_xdata()
                artist.ydata = artist.get_ydata()
            elif item["plot_type"] == "scatter" and hasattr(artist, "get_offsets"):
                offsets = artist.get_offsets()
                artist.xdata = offsets[:, 0]
                artist.ydata = offsets[:, 1]

        self.items.append(item)
        return item

    def remove(self, index_or_label):
        if isinstance(index_or_label, int):
            return self.items.pop(index_or_label)
        for index, item in enumerate(self.items):
            if item.get("label") == index_or_label:
                return self.items.pop(index)
        return None
