from __future__ import annotations

from datetime import datetime
from typing import Any

from pymongo import ReturnDocument, UpdateOne


class _FakeUpdateResult:
    def __init__(self, *, matched_count: int = 0, upserted_count: int = 0, modified_count: int = 0):
        self.matched_count = matched_count
        self.upserted_count = upserted_count
        self.modified_count = modified_count


class _FakeDeleteResult:
    def __init__(self, deleted_count: int):
        self.deleted_count = deleted_count


class _FakeBulkResult:
    def __init__(self, *, upserted_count: int = 0, modified_count: int = 0):
        self.upserted_count = upserted_count
        self.modified_count = modified_count


class _FakeCursor:
    def __init__(self, docs: list[dict]):
        self._docs = list(docs)

    def skip(self, _):
        return self

    def limit(self, n: int):
        self._docs = self._docs[:n]
        return self

    def sort(self, *_args, **_kwargs):
        return self

    def __aiter__(self):
        self._iter = iter(self._docs)
        return self

    async def __anext__(self):
        try:
            return next(self._iter)
        except StopIteration as exc:
            raise StopAsyncIteration from exc


class _FakeCollection:
    def __init__(self):
        self.docs: dict[str, dict] = {}
        self.indexes: list = []

    async def create_index(self, keys, **kwargs):
        self.indexes.append((keys, kwargs))

    async def find_one(self, query: dict):
        if "record_id" in query:
            return self.docs.get(query["record_id"])
        if "run_id" in query:
            return next((doc for doc in self.docs.values() if doc.get("run_id") == query["run_id"]), None)
        if "dedupe_key" in query:
            return self.docs.get(query["dedupe_key"])
        if "analysis_id" in query:
            return self.docs.get(query["analysis_id"])
        if "_id" in query:
            return self.docs.get(str(query["_id"]))
        if len(query) > 0:
            for doc in self.docs.values():
                if _matches(doc, query):
                    return doc
        return None

    def find(self, query: dict):
        results = []
        for doc in self.docs.values():
            if not _matches(doc, query):
                continue
            results.append(doc)
        return _FakeCursor(results)

    async def count_documents(self, query: dict):
        return len([doc for doc in self.docs.values() if _matches(doc, query)])

    async def find_one_and_update(self, query, update, upsert, return_document):
        key = _primary_key(query, self.docs)
        existing = self.docs.get(key, {})
        created = existing.get("created_at") or update.get("$setOnInsert", {}).get("created_at") or datetime.utcnow()
        merged = {
            **update.get("$setOnInsert", {}),
            **existing,
            **update.get("$set", {}),
            "created_at": created,
        }
        self.docs[key] = merged
        return merged

    async def update_one(self, query, update):
        if "record_id" in query:
            key = query["record_id"]
        elif "run_id" in query:
            key = next((k for k, doc in self.docs.items() if doc.get("run_id") == query["run_id"]), None)
        else:
            key = None
        if key is None or key not in self.docs:
            return _FakeUpdateResult(matched_count=0)
        self.docs[key].update(update.get("$set", {}))
        if "$inc" in update:
            self.docs[key]["retry_count"] = self.docs[key].get("retry_count", 0) + update["$inc"]["retry_count"]
        return _FakeUpdateResult(matched_count=1)

    async def bulk_write(self, operations, ordered=False):
        upserted = 0
        modified = 0
        for op in operations:
            if not isinstance(op, UpdateOne):
                continue
            record_id = op._filter["record_id"]
            existing = self.docs.get(record_id, {})
            created = existing.get("created_at") or op._doc["$setOnInsert"]["created_at"]
            merged = {**existing, **op._doc["$set"], "created_at": created}
            if record_id not in self.docs:
                upserted += 1
            else:
                modified += 1
            self.docs[record_id] = merged
        return _FakeBulkResult(upserted_count=upserted, modified_count=modified)

    async def delete_many(self, query):
        if "record_id" in query and "$in" in query["record_id"]:
            ids = set(query["record_id"]["$in"])
            deleted = [key for key in list(self.docs) if key in ids]
            for key in deleted:
                del self.docs[key]
            return _FakeDeleteResult(deleted_count=len(deleted))
        return _FakeDeleteResult(deleted_count=0)

    async def insert_many(self, docs):
        for doc in docs:
            key = doc.get("chunk_id") or doc.get("analysis_id") or str(len(self.docs))
            self.docs[key] = dict(doc)
        return _FakeUpdateResult(upserted_count=len(docs))


def _primary_key(query: dict, docs: dict[str, dict]) -> str:
    if "dedupe_key" in query:
        return query["dedupe_key"]
    if "record_id" in query:
        return query["record_id"]
    return str(len(docs))


def _matches(doc: dict, query: dict) -> bool:
    for key, expected in query.items():
        if key == "$or":
            if not any(_matches(doc, clause) for clause in expected):
                return False
            continue
        if key == "metadata.snapshot_id":
            if doc.get("metadata", {}).get("snapshot_id") != expected:
                return False
            continue
        if isinstance(expected, dict) and "$in" in expected:
            if doc.get(key) not in expected["$in"]:
                return False
            continue
        if doc.get(key) != expected:
            return False
    return True


class FakeIntegrationDB:
    def __init__(self):
        self.doc_chunks = _FakeCollection()
        self.code_chunks = _FakeCollection()
        self.commit_analyses = _FakeCollection()
        self.embedding_records = _FakeCollection()
        self.embedding_runs = _FakeCollection()
        self.snapshot_graphs = _FakeCollection()
        self.snapshot_graph_parts = _FakeCollection()

    def get_collection(self, name: str, codec_options=None):
        return getattr(self, name, _FakeCollection())


class FakeQdrantStore:
    def __init__(self):
        self.points: dict[str, dict] = {}

    def upsert(self, collection_name, points, wait=True):
        for point in points:
            self.points[str(point.id)] = {
                "vector": point.vector,
                "payload": point.payload,
                "collection": collection_name,
            }

    def delete(self, collection_name, points_selector, wait=True):
        from qdrant_client.http import models

        if isinstance(points_selector, models.PointIdsList):
            for point_id in points_selector.points:
                self.points.pop(str(point_id), None)
        elif isinstance(points_selector, models.FilterSelector):
            to_delete = []
            for point_id, data in self.points.items():
                if data.get("collection") != collection_name:
                    continue
                if _qdrant_payload_matches_filter(data["payload"], points_selector.filter):
                    to_delete.append(point_id)
            for point_id in to_delete:
                del self.points[point_id]

    def _search_points(
        self,
        collection_name,
        query_vector,
        query_filter=None,
        limit=10,
        score_threshold=None,
    ):
        from qdrant_client.http import models

        matches: list[models.ScoredPoint] = []
        for point_id, data in self.points.items():
            if data.get("collection") != collection_name:
                continue
            if not _qdrant_payload_matches_filter(data["payload"], query_filter):
                continue
            score = 0.9
            if score_threshold is not None and score < score_threshold:
                continue
            matches.append(
                models.ScoredPoint(
                    id=point_id,
                    version=0,
                    score=score,
                    payload=data["payload"],
                )
            )
        matches.sort(key=lambda item: item.score, reverse=True)
        return matches[:limit]

    def query_points(
        self,
        collection_name,
        query,
        query_filter=None,
        limit=10,
        score_threshold=None,
        with_payload=True,
        **kwargs,
    ):
        points = self._search_points(
            collection_name,
            query,
            query_filter=query_filter,
            limit=limit,
            score_threshold=score_threshold,
        )

        class _FakeQueryResponse:
            def __init__(self, scored_points):
                self.points = scored_points

        return _FakeQueryResponse(points)


def _qdrant_payload_matches_filter(payload: dict, query_filter) -> bool:
    from qdrant_client.http import models

    if query_filter is None:
        return True
    for condition in query_filter.must:
        if not isinstance(condition, models.FieldCondition):
            continue
        key = condition.key
        if isinstance(condition.match, models.MatchValue):
            if payload.get(key) != condition.match.value:
                return False
        elif isinstance(condition.match, models.MatchAny):
            if payload.get(key) not in condition.match.any:
                return False
    return True


class FakeQdrantClient:
    """Minimal Qdrant client surface for integration tests."""

    def __init__(self, store: FakeQdrantStore):
        self.store = store

    def upsert(self, collection_name, points, wait=True):
        self.store.upsert(collection_name, points, wait=wait)

    def delete(self, collection_name, points_selector, wait=True):
        self.store.delete(collection_name, points_selector, wait=wait)

    def query_points(
        self,
        collection_name,
        query,
        query_filter=None,
        limit=10,
        score_threshold=None,
        with_payload=True,
        **kwargs,
    ):
        return self.store.query_points(
            collection_name,
            query,
            query_filter=query_filter,
            limit=limit,
            score_threshold=score_threshold,
            with_payload=with_payload,
            **kwargs,
        )

    def close(self):
        return None


class FakeEmbeddingProvider:
    def __init__(self) -> None:
        self.call_count = 0

    async def embed_batch(self, texts: list[str]) -> list[list[float]]:
        self.call_count += 1
        return [[0.01] * 1536 for _ in texts]

    async def embed_single(self, text: str) -> list[float]:
        self.call_count += 1
        return [0.01] * 1536


def seed_snapshot_upstream(db: FakeIntegrationDB) -> None:
    db.doc_chunks.docs["chk_doc_1"] = {
        "chunk_id": "chk_doc_1",
        "repo_id": "AD/example-repo",
        "doc_path": "docs/architecture.md",
        "section_title": "Overview",
        "chunk_text": "The service exposes a REST API on port 6001.",
        "chunk_hash": "sha256:doc1",
        "metadata": {
            "snapshot_id": "snap_abc123",
            "commit_hash": "a1b2c3d4e5f6",
            "section_level": 2,
        },
    }
    db.code_chunks.docs["chk_sym_Load"] = {
        "chunk_id": "chk_sym_config_go_Load",
        "repo_id": "AD/example-repo",
        "snapshot_id": "snap_abc123",
        "commit_sha": "a1b2c3d4e5f6",
        "file_path": "internal/config/config.go",
        "symbol_name": "Load",
        "symbol_type": "function",
        "text": "func Load(path string) (*Config, error) { ... }",
        "start_line": 24,
        "end_line": 58,
        "chunk_hash": "sha256:code1",
        "language": "go",
    }


def _snapshot_graph_nodes_edges() -> tuple[list[dict], list[dict]]:
    return (
        [
            {
                "id": "sym_config_go_Load",
                "kind": "symbol",
                "name": "Load",
                "path": "internal/config/config.go",
            },
            {
                "id": "chk_doc_1",
                "kind": "doc_chunk",
                "name": "Overview",
            },
        ],
        [
            {
                "id": "e_graph_doc_to_symbol",
                "kind": "documents",
                "source_id": "chk_doc_1",
                "target_id": "sym_config_go_Load",
            }
        ],
    )


def seed_snapshot_graph(
    db: FakeIntegrationDB,
    *,
    repo_id: str = "AD/example-repo",
    snapshot_id: str = "snap_abc123",
) -> None:
    artifact_id = f"art_{snapshot_id}_merged"
    nodes, edges = _snapshot_graph_nodes_edges()
    db.snapshot_graphs.docs[artifact_id] = {
        "_id": artifact_id,
        "schema_version": "v1",
        "repo_id": repo_id,
        "snapshot_id": snapshot_id,
        "commit_sha": "a1b2c3d4e5f6",
        "artifact_id": artifact_id,
        "chunked": False,
        "node_count": len(nodes),
        "edge_count": len(edges),
        "nodes": nodes,
        "edges": edges,
    }


def seed_chunked_snapshot_graph(
    db: FakeIntegrationDB,
    *,
    repo_id: str = "AD/example-repo",
    snapshot_id: str = "snap_abc123",
) -> None:
    artifact_id = f"art_{snapshot_id}_merged"
    nodes, edges = _snapshot_graph_nodes_edges()
    db.snapshot_graphs.docs[artifact_id] = {
        "_id": artifact_id,
        "schema_version": "v1",
        "repo_id": repo_id,
        "snapshot_id": snapshot_id,
        "commit_sha": "a1b2c3d4e5f6",
        "artifact_id": artifact_id,
        "chunked": True,
        "part_count": 2,
        "node_count": len(nodes),
        "edge_count": len(edges),
    }
    db.snapshot_graph_parts.docs[f"{artifact_id}:0"] = {
        "artifact_id": artifact_id,
        "part_index": 0,
        "nodes": nodes,
        "edges": [],
    }
    db.snapshot_graph_parts.docs[f"{artifact_id}:1"] = {
        "artifact_id": artifact_id,
        "part_index": 1,
        "nodes": [],
        "edges": edges,
    }


def seed_commit_upstream(db: FakeIntegrationDB) -> None:
    db.commit_analyses.docs["analysis_b2c3d4"] = {
        "analysis_id": "analysis_b2c3d4",
        "repo_id": "AD/example-repo",
        "commit_sha": "b2c3d4e5f6a7",
        "summary": "Introduced Reload and refactored configuration loading.",
        "impacted_symbols": ["Load", "Reload"],
        "changed_files": ["internal/config/config.go"],
    }
