"""Tests for recommend_papers tool."""

from __future__ import annotations

import asyncio
import json
from contextlib import asynccontextmanager

import httpx
import pytest
import respx
from fastmcp import FastMCP
from fastmcp.client import Client

from scholar_mcp._server_deps import ServiceBundle
from scholar_mcp._tools_recommendations import register_recommendation_tools

S2_REC = "https://api.semanticscholar.org/recommendations/v1"


@pytest.fixture
def mcp(bundle: ServiceBundle) -> FastMCP:
    @asynccontextmanager
    async def lifespan(app: FastMCP):  # type: ignore[type-arg]
        yield {"bundle": bundle}

    app = FastMCP("test", lifespan=lifespan)
    register_recommendation_tools(app)
    return app


async def test_recommend_papers(mcp: FastMCP) -> None:
    with respx.mock:
        respx.post(f"{S2_REC}/papers").mock(
            return_value=httpx.Response(
                200,
                json={
                    "recommendedPapers": [
                        {
                            "paperId": "r1",
                            "title": "Recommended 1",
                            "year": 2023,
                            "citationCount": 10,
                        }
                    ]
                },
            )
        )
        async with Client(mcp) as client:
            result = await client.call_tool(
                "recommend_papers", {"positive_ids": ["p1", "p2"]}
            )
    data = json.loads(result.content[0].text)
    assert len(data) == 1
    assert data[0]["paperId"] == "r1"


async def test_recommend_papers_with_negatives(mcp: FastMCP) -> None:
    with respx.mock:
        respx.post(f"{S2_REC}/papers").mock(
            return_value=httpx.Response(200, json={"recommendedPapers": []})
        )
        async with Client(mcp) as client:
            result = await client.call_tool(
                "recommend_papers",
                {"positive_ids": ["p1"], "negative_ids": ["n1"], "limit": 5},
            )
    data = json.loads(result.content[0].text)
    assert isinstance(data, list)


async def test_recommend_papers_caps_positive_ids(mcp: FastMCP) -> None:
    """Only first 5 positive IDs are sent (spec: 1-5)."""
    captured_body: dict = {}

    async def capture(request: httpx.Request) -> httpx.Response:
        import json as _json

        captured_body.update(_json.loads(request.content))
        return httpx.Response(200, json={"recommendedPapers": []})

    with respx.mock:
        respx.post(f"{S2_REC}/papers").mock(side_effect=capture)
        async with Client(mcp) as client:
            await client.call_tool(
                "recommend_papers",
                {"positive_ids": ["p1", "p2", "p3", "p4", "p5", "p6"]},
            )
    assert len(captured_body.get("positivePaperIds", [])) <= 5


async def test_recommend_papers_upstream_error(mcp: FastMCP) -> None:
    with respx.mock:
        respx.post(f"{S2_REC}/papers").mock(
            return_value=httpx.Response(500, text="Internal Server Error")
        )
        async with Client(mcp) as client:
            result = await client.call_tool(
                "recommend_papers", {"positive_ids": ["p1"]}
            )
    data = json.loads(result.content[0].text)
    assert data["error"] == "upstream_error"


async def test_recommend_papers_queued_on_429(
    bundle: ServiceBundle,
) -> None:
    """recommend_papers returns queued on 429, background completes."""
    call_count = 0

    def _side_effect(request: httpx.Request) -> httpx.Response:
        nonlocal call_count
        call_count += 1
        if call_count == 1:
            return httpx.Response(429)
        return httpx.Response(
            200,
            json={
                "recommendedPapers": [{"paperId": "r1", "title": "Rec 1", "year": 2024}]
            },
        )

    with respx.mock:
        respx.post(f"{S2_REC}/papers").mock(side_effect=_side_effect)

        @asynccontextmanager
        async def lifespan(app: FastMCP):  # type: ignore[type-arg]
            yield {"bundle": bundle}

        app = FastMCP("test", lifespan=lifespan)
        register_recommendation_tools(app)
        from scholar_mcp._tools_tasks import register_task_tools

        register_task_tools(app)

        async with Client(app) as client:
            result = await client.call_tool(
                "recommend_papers", {"positive_ids": ["p1"]}
            )
            data = json.loads(result.content[0].text)
            assert data["queued"] is True
            assert data["tool"] == "recommend_papers"

            for _ in range(40):
                poll = await client.call_tool(
                    "get_task_result", {"task_id": data["task_id"]}
                )
                poll_data = json.loads(poll.content[0].text)
                if poll_data["status"] in ("completed", "failed"):
                    break
                await asyncio.sleep(0.05)
            assert poll_data["status"] == "completed"
            inner = json.loads(poll_data["result"])
            assert inner[0]["paperId"] == "r1"
