import asyncio
from datetime import datetime

import pytest
from fastapi.testclient import TestClient

from app.models.embedding_run import EmbeddingRun
from app.models.enums import ChunkType, EmbedStatus, SourceType
from app.models.stored_embedding_record import StoredEmbeddingRecord
from app.repositories import embedding_record_repo, embedding_run_repo


@pytest.fixture(autouse=True)
def _skip_startup_bootstrap(monkeypatch):
    monkeypatch.setattr("app.main.bootstrap_qdrant", lambda *_args, **_kwargs: "adpilot_embeddings")


@pytest.fixture
def client():
    from app.main import app

    return TestClient(app)


@pytest.fixture
def mock_repos(monkeypatch):
    embedding_record_repo._INDEXES_READY = False
    embedding_run_repo._INDEXES_READY = False

    run = EmbeddingRun(
        run_id="erun_test",
        dedupe_key="dedupe",
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
        event_id="evt_1",
        run_type="snapshot_batch",
        status=EmbedStatus.COMPLETED,
        total_records=2,
        embedded_records=2,
        celery_task_id="task-1",
        created_at=datetime.utcnow(),
        updated_at=datetime.utcnow(),
    )

    record = StoredEmbeddingRecord(
        record_id="emb_chk_1",
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
        commit_sha="a1b2c3d4e5f6",
        source_type=SourceType.DOCS,
        chunk_type=ChunkType.TEXT,
        text="Overview text for the architecture document.",
        doc_path="docs/architecture.md",
        section_title="Overview",
        graph_node_id="doc_chunk_chk_1",
        embed_status=EmbedStatus.EMBEDDED,
        qdrant_point_id="point-1",
        dedupe_key="dedupe-record",
        created_at=datetime.utcnow(),
        updated_at=datetime.utcnow(),
    )

    async def _get_run(run_id):
        return run if run_id == "erun_test" else None

    async def _list_runs(repo_id, snapshot_id, limit=50, statuses=None):
        return [run] if repo_id == "AD/example-repo" else []

    async def _get_record(record_id):
        return record if record_id == "emb_chk_1" else None

    async def _list_records(repo_id, snapshot_id, **kwargs):
        return [record] if repo_id == "AD/example-repo" else []

    async def _count_records(repo_id, snapshot_id, embed_status=None):
        if embed_status == EmbedStatus.EMBEDDED:
            return 1
        if embed_status == EmbedStatus.PENDING:
            return 0
        if embed_status == EmbedStatus.FAILED:
            return 0
        return 1

    monkeypatch.setattr(embedding_run_repo, "get_run", _get_run)
    monkeypatch.setattr(embedding_run_repo, "list_runs_by_snapshot", _list_runs)
    monkeypatch.setattr(embedding_record_repo, "get_record", _get_record)
    monkeypatch.setattr(embedding_record_repo, "list_records_by_snapshot", _list_records)
    monkeypatch.setattr(embedding_record_repo, "count_records_by_snapshot", _count_records)


def test_get_embedding_run(client, mock_repos):
    response = client.get("/api/v1/runs/erun_test")
    assert response.status_code == 200
    assert response.json()["run_id"] == "erun_test"
    assert response.json()["status"] == "completed"


def test_get_embedding_run_not_found(client, mock_repos):
    response = client.get("/api/v1/runs/missing")
    assert response.status_code == 404


def test_list_embedding_records(client, mock_repos):
    response = client.get(
        "/api/v1/records",
        params={"repo_id": "AD/example-repo", "snapshot_id": "snap_abc123"},
    )
    assert response.status_code == 200
    body = response.json()
    assert len(body) == 1
    assert body[0]["record_id"] == "emb_chk_1"
    assert "text_preview" in body[0]


def test_record_count_summary(client, mock_repos):
    response = client.get(
        "/api/v1/records/count/summary",
        params={"repo_id": "AD/example-repo", "snapshot_id": "snap_abc123"},
    )
    assert response.status_code == 200
    assert response.json()["total"] == 1
    assert response.json()["embedded"] == 1


def test_trigger_snapshot_embedding(client, monkeypatch):
    class _Task:
        id = "celery-task-99"

    async def _upsert_run(**kwargs):
        return EmbeddingRun(
            run_id="erun_new",
            dedupe_key="dedupe-new",
            repo_id=kwargs["repo_id"],
            snapshot_id=kwargs["snapshot_id"],
            event_id=kwargs["event_id"],
            run_type=kwargs.get("run_type", "snapshot_batch"),
            created_at=datetime.utcnow(),
            updated_at=datetime.utcnow(),
        )

    async def _update_run_state(*_args, **_kwargs):
        return True

    class _MockTask:
        def delay(self, *_args, **_kwargs):
            return _Task()

    monkeypatch.setattr(embedding_run_repo, "upsert_run", _upsert_run)
    monkeypatch.setattr(embedding_run_repo, "update_run_state", _update_run_state)
    monkeypatch.setattr(
        "app.api.routes.pipeline.process_snapshot_embedding_task",
        _MockTask(),
    )

    response = client.post(
        "/api/v1/pipeline/snapshots/embed",
        json={
            "repo_id": "AD/example-repo",
            "snapshot_id": "snap_abc123",
            "commit_sha": "a1b2c3d4e5f6",
            "artifact_uri": "mongo://snapshot_graphs/art_1",
        },
    )
    assert response.status_code == 202
    assert response.json()["run_id"] == "erun_new"
    assert response.json()["celery_task_id"] == "celery-task-99"


def test_version_endpoint(client):
    response = client.get("/api/v1/version")
    assert response.status_code == 200
    assert response.json()["service"] == "Embedding Engine"


def test_root_readyz_and_healthz(client, monkeypatch):
    async def _true():
        return True

    monkeypatch.setattr("app.api.readiness.is_db_ready", _true)
    monkeypatch.setattr("app.api.readiness.is_redis_ready", _true)
    monkeypatch.setattr("app.api.readiness.is_qdrant_ready", _true)

    healthz = client.get("/healthz")
    readyz = client.get("/readyz")
    assert healthz.status_code == 200
    assert readyz.status_code == 200
    assert readyz.json()["status"] == "ready"
