from __future__ import annotations

from typing import Any

from pydantic import BaseModel, Field, model_validator


class SnapshotGraphNode(BaseModel):
    id: str
    kind: str
    name: str | None = None
    language: str | None = None
    path: str | None = None
    metadata: dict[str, Any] = Field(default_factory=dict)

    @model_validator(mode="before")
    @classmethod
    def _normalize_mongo_node(cls, data: Any) -> Any:
        if not isinstance(data, dict):
            return data
        normalized = dict(data)
        if normalized.get("metadata") is None:
            normalized["metadata"] = {}
        return normalized


class SnapshotGraphEdge(BaseModel):
    id: str
    kind: str
    source_id: str
    target_id: str
    metadata: dict[str, Any] = Field(default_factory=dict)

    @model_validator(mode="before")
    @classmethod
    def _normalize_mongo_edge(cls, data: Any) -> Any:
        if not isinstance(data, dict):
            return data
        normalized = dict(data)
        if "source_id" not in normalized and "sourceid" in normalized:
            normalized["source_id"] = normalized.pop("sourceid")
        if "target_id" not in normalized and "targetid" in normalized:
            normalized["target_id"] = normalized.pop("targetid")
        if normalized.get("metadata") is None:
            normalized["metadata"] = {}
        return normalized


class SnapshotGraphDocument(BaseModel):
    schema_version: str = "v1"
    repo_id: str
    snapshot_id: str
    commit_sha: str | None = None
    artifact_id: str | None = None
    chunked: bool = False
    part_count: int = 0
    node_count: int | None = None
    edge_count: int | None = None
    nodes: list[SnapshotGraphNode] = Field(default_factory=list)
    edges: list[SnapshotGraphEdge] = Field(default_factory=list)

    @classmethod
    def from_mongo(cls, doc: dict[str, Any]) -> SnapshotGraphDocument:
        payload = dict(doc)
        if "_id" in payload and "artifact_id" not in payload:
            payload["artifact_id"] = str(payload.pop("_id"))
        elif "_id" in payload:
            payload.pop("_id", None)
        return cls.model_validate(payload)
