import time

import pytest
import pytest_asyncio

from app.services.infra import jobs

TENANT = "org_job_test"

pytestmark = pytest.mark.asyncio


@pytest_asyncio.fixture(autouse=True)
async def clean():
    yield
    await jobs.clear_tenant(TENANT)


async def test_a_new_job_starts_queued():
    job_id = await jobs.create(TENANT, "enrich")
    state = await jobs.get(job_id)
    assert state["status"] == "queued"
    assert state["percent"] == 0
    assert state["kind"] == "enrich"
    await jobs.finish(job_id, result={"ok": True})


async def test_progress_moves_the_job_to_running():
    job_id = await jobs.create(TENANT, "enrich")
    jobs.tick(job_id, 45, "batch 5 of 11")
    state = await jobs.get(job_id)
    assert state["status"] == "running"
    assert state["percent"] == 45
    assert state["step"] == "batch 5 of 11"
    await jobs.finish(job_id, result={})


async def test_a_finished_job_carries_its_report():
    # The report must be the same object the endpoint used to return, so
    # nothing downstream has to learn a new shape.
    job_id = await jobs.create(TENANT, "enrich")
    await jobs.finish(job_id, result={"products": 218, "extracted": 218})
    state = await jobs.get(job_id)
    assert state["status"] == "done"
    assert state["percent"] == 100
    assert state["result"]["extracted"] == 218


async def test_a_failed_job_carries_only_the_exception_class():
    job_id = await jobs.create(TENANT, "pair")
    await jobs.finish(job_id, error="RuntimeError")
    state = await jobs.get(job_id)
    assert state["status"] == "failed"
    assert state["error"] == "RuntimeError"


async def test_a_second_job_of_the_same_kind_is_refused():
    # Two concurrent pair runs both rewrite the neighbour table.
    first = await jobs.create(TENANT, "pair")
    with pytest.raises(jobs.JobAlreadyRunning) as caught:
        await jobs.create(TENANT, "pair")
    assert caught.value.job_id == first
    await jobs.finish(first, result={})


async def test_a_different_kind_may_run_concurrently():
    pair = await jobs.create(TENANT, "pair")
    enrich = await jobs.create(TENANT, "enrich")
    assert pair != enrich
    await jobs.finish(pair, result={})
    await jobs.finish(enrich, result={})


async def test_finishing_releases_the_lock():
    first = await jobs.create(TENANT, "pair")
    await jobs.finish(first, result={})
    second = await jobs.create(TENANT, "pair")
    assert second != first
    await jobs.finish(second, result={})


async def test_a_silent_job_is_reported_lost(monkeypatch):
    # An in-process job dies with the process. Without this a client polls
    # "running" forever and never learns.
    monkeypatch.setattr(jobs, "STALE_AFTER_SECONDS", 0)
    job_id = await jobs.create(TENANT, "pair")
    jobs.tick(job_id, 10, "embedding")
    time.sleep(1)
    assert (await jobs.get(job_id))["status"] == "lost"
    await jobs.finish(job_id, error="Lost")


async def test_a_finished_job_is_never_lost(monkeypatch):
    monkeypatch.setattr(jobs, "STALE_AFTER_SECONDS", 0)
    job_id = await jobs.create(TENANT, "enrich")
    await jobs.finish(job_id, result={})
    time.sleep(1)
    assert (await jobs.get(job_id))["status"] == "done"


async def test_an_unknown_job_is_none():
    assert await jobs.get("job_does_not_exist") is None


async def test_recent_jobs_are_listed_newest_first():
    first = await jobs.create(TENANT, "enrich")
    await jobs.finish(first, result={})
    second = await jobs.create(TENANT, "pair")
    await jobs.finish(second, result={})
    listed = [j["job_id"] for j in await jobs.recent(TENANT)]
    assert listed[:2] == [second, first]


async def test_a_tick_on_an_unknown_job_does_not_raise():
    # Progress reporting is best-effort: it must never take down work that is
    # already spending money on model calls.
    jobs.tick("job_gone", 50, "step")


async def test_a_failed_job_keeps_the_percentage_it_reached():
    # Claiming 100 for a job that died at 30 tells a merchant the work finished.
    job_id = await jobs.create(TENANT, "pair")
    jobs.tick(job_id, 30, "embedding products")
    await jobs.finish(job_id, error="RuntimeError")

    state = await jobs.get(job_id)
    assert state["status"] == "failed"
    assert state["percent"] == 30
