"""Unit tests for the extracted turn prologue (``agent/turn_context.py``).

These exercise ``build_turn_context`` against a lightweight fake agent to
confirm the prologue produces the right ``TurnContext`` and applies the
``agent`` side effects the loop relies on — without spinning up a real
``AIAgent`` or hitting any provider.
"""

from __future__ import annotations

import threading
import types
from unittest.mock import MagicMock, patch

import pytest

from agent.turn_context import (
    PreflightCompressionTimedOut,
    TurnContext,
    build_turn_context,
)


class _FakeTodoStore:
    def has_items(self):
        return True

    def _hydrate(self, *_a, **_k):
        pass


class _FakeGuardrails:
    def __init__(self):
        self.reset_called = False

    def reset_for_turn(self):
        self.reset_called = True


class _FakeAgent:
    """Minimal stand-in covering only what the prologue touches."""

    def __init__(self):
        self.session_id = "sess-1"
        self.model = "test/model"
        self.provider = "openrouter"
        self.requested_provider = "openrouter"
        self.base_url = "https://openrouter.ai/api/v1"
        self.api_key = "sk-x"
        self.api_mode = "chat_completions"
        self.platform = "cli"
        self.quiet_mode = True
        self.max_iterations = 90
        self.tools = []
        self.valid_tool_names = set()
        self.enabled_toolsets = None
        self.disabled_toolsets = None
        self._skip_mcp_refresh = False
        self.compression_enabled = False
        self.context_compressor = types.SimpleNamespace(
            protect_first_n=2, protect_last_n=2
        )
        # Make the fake compressor honour the ContextEngine contract that the
        # real code now relies on (should_compress_info returns a (bool, reason)
        # tuple). Without it build_turn_context raises AttributeError.
        def _fake_should_compress(tokens=None):
            return False

        def _fake_should_compress_info(tokens=None):
            return (False, None)

        self.context_compressor.should_compress = _fake_should_compress
        self.context_compressor.should_compress_info = _fake_should_compress_info
        self._cached_system_prompt = "SYSTEM"
        self._memory_store = None
        self._memory_manager = None
        self._memory_nudge_interval = 0
        self._turns_since_memory = 0
        self._user_turn_count = 0
        self._todo_store = _FakeTodoStore()
        self._tool_guardrails = _FakeGuardrails()
        self._compression_warning = None
        self._emit_warning = MagicMock()
        self._last_ctx_overflow_warn = None
        self._interrupt_requested = False
        self._memory_write_origin = "assistant_tool"
        self._stream_context_scrubber = None
        self._stream_think_scrubber = None
        # Attributes the prologue assigns; recorded for assertions.
        self._invalid_tool_retries = -1
        self._persist_calls = 0
        self._session_messages = []
        self._pending_cli_user_message = None
        self._session_persist_lock = threading.RLock()
        # Records _cached_system_prompt at the moment _ensure_db_session()
        # is called (regression guard for #45499 turn-setup ordering).
        self._ensure_db_prompt_at_call = "<unset>"

    def _warn_context_overflow_blocked(self, reason, preflight_tokens, threshold_tokens):
        # Mirror the real AIAgent helper so tests can assert the warning fired.
        _warn_kind = (reason or "unknown").split(":", 1)[0]
        _warn_key = ("ctx_overflow_blocked", _warn_kind)
        if self._last_ctx_overflow_warn != _warn_key:
            self._last_ctx_overflow_warn = _warn_key
            self._emit_warning(
                f"⚠ Context is over the compression threshold "
                f"(~{preflight_tokens:,} tokens >= {threshold_tokens:,}) "
                f"but compression is currently blocked ({reason})."
            )

    def _clear_context_overflow_warn(self):
        self._last_ctx_overflow_warn = None

    # --- methods the prologue calls ---
    def _ensure_db_session(self):
        self._ensure_db_prompt_at_call = self._cached_system_prompt

    def _restore_primary_runtime(self):
        pass

    def _cleanup_dead_connections(self):
        return False

    def _emit_status(self, _msg):
        pass

    def _replay_compression_warning(self):
        pass

    def _hydrate_todo_store(self, *_a, **_k):
        pass

    def _safe_print(self, *_a, **_k):
        pass

    def _persist_session(self, *_a, **_k):
        self._persist_calls += 1


@pytest.fixture(autouse=True)
def _stub_runtime_main():
    """``build_turn_context`` calls ``auxiliary_client.set_runtime_main`` as a
    production side effect (telling aux tools the live main provider/model).
    That writes a module-level global these unit tests don't care about and
    which would otherwise leak into sibling tests (e.g. provider-parity
    resolution) when the per-test process isolation plugin is disabled. Stub
    it out so the prologue tests stay hermetic.
    """
    with patch("agent.auxiliary_client.set_runtime_main", lambda *a, **k: None):
        yield


def _build(agent, **overrides):
    kwargs = dict(
        agent=agent,
        user_message="hello",
        system_message=None,
        conversation_history=None,
        task_id=None,
        stream_callback=None,
        persist_user_message=None,
        restore_or_build_system_prompt=lambda *a, **k: None,
        install_safe_stdio=lambda: None,
        sanitize_surrogates=lambda s: s,
        summarize_user_message_for_log=lambda s: s,
        set_session_context=lambda _sid: None,
        set_current_write_origin=lambda _o: None,
        ra=lambda: types.SimpleNamespace(_set_interrupt=lambda *a, **k: None),
    )
    kwargs.update(overrides)
    return build_turn_context(**kwargs)


def test_returns_turn_context_with_user_message_appended():
    agent = _FakeAgent()
    ctx = _build(agent)
    assert isinstance(ctx, TurnContext)
    assert ctx.user_message == "hello"
    # The user turn was appended and indexed.
    assert ctx.messages[-1]["role"] == "user"
    assert ctx.messages[-1]["content"] == "hello"
    assert isinstance(ctx.messages[-1]["timestamp"], float)
    assert ctx.current_turn_user_idx == len(ctx.messages) - 1
    assert ctx.active_system_prompt == "SYSTEM"


def test_preflight_timeout_stops_turn_before_provider_boundary():
    """An unchanged payload above the model window must not escape turn construction (a request that
    still fits its window is sent uncompressed instead — see test_preflight_compression_timeout_fail_closed)."""
    agent = _FakeAgent()
    agent.compression_enabled = True
    agent.max_compression_attempts = 3
    agent.context_compressor = types.SimpleNamespace(
        protect_first_n=2,
        protect_last_n=2,
        threshold_tokens=1_000,
        context_length=1_500,
        summary_target_ratio=0.3,
        last_prompt_tokens=0,
        should_compress=lambda tokens=None: True,
        should_compress_info=lambda tokens=None: (True, None),
        get_active_compression_failure_cooldown=lambda: None,
    )

    def stalled_compression(messages, *_args, **_kwargs):
        agent._last_compression_timed_out = True
        return messages, "SYSTEM"

    agent._compress_context = MagicMock(side_effect=stalled_compression)
    oversized_history = [
        {"role": "assistant", "content": "x" * 8_000},
    ]
    provider_call = MagicMock()

    def run_turn():
        context = _build(agent, conversation_history=oversized_history)
        provider_call(context.messages)

    with pytest.raises(PreflightCompressionTimedOut, match="provider call was not sent"):
        run_turn()

    agent._compress_context.assert_called_once()
    provider_call.assert_not_called()


def test_user_message_preserves_platform_event_timestamp():
    agent = _FakeAgent()

    ctx = _build(agent, persist_user_timestamp=123.5)

    assert ctx.messages[-1]["timestamp"] == 123.5


# ── Trivial-prompt prefetch gate (PR #25350 salvage) ─────────────────────────
#
# The prologue is the ONLY place the per-turn synchronous
# memory_manager.prefetch_all() fires; a bare greeting must not block the
# turn on provider network round-trips, while a substantive question must
# still prefetch. These assert the gate at the call site (the classifier
# itself is covered in tests/agent/test_memory_provider.py).


def _agent_with_memory_manager():
    agent = _FakeAgent()
    mm = MagicMock()
    mm.prefetch_all.return_value = "REMEMBERED CONTEXT"
    agent._memory_manager = mm
    return agent, mm


def test_prefetch_skipped_for_trivial_user_message():
    agent, mm = _agent_with_memory_manager()
    ctx = _build(agent, user_message="hi!")
    mm.prefetch_all.assert_not_called()
    assert ctx.ext_prefetch_cache == ""


def test_prefetch_runs_for_substantive_user_message():
    agent, mm = _agent_with_memory_manager()
    query = "what did we decide about the deploy pipeline?"
    ctx = _build(agent, user_message=query)
    mm.prefetch_all.assert_called_once_with(query, session_id=agent.session_id)
    assert ctx.ext_prefetch_cache == "REMEMBERED CONTEXT"


# ── Per-turn author ──────────────────────────────────────────────────────────


def test_turn_author_is_normalized_then_reaches_on_turn_start_and_the_agent_stash():
    agent, mm = _agent_with_memory_manager()

    _build(agent, user_message="what did we decide about the deploy pipeline?",
           turn_author={"id": " bot:al\x00pha ", "name": "Alpha", "is_bot": 1, "x": 1})

    kwargs = mm.on_turn_start.call_args.kwargs
    assert (kwargs["author_id"], kwargs["author_name"], kwargs["author_is_bot"]) == ("bot:alpha", "Alpha", True)
    # The end-of-turn sync reads this back off the agent.
    assert agent._turn_author == {"id": "bot:alpha", "name": "Alpha", "is_bot": True}


def test_turn_without_author_clears_previous_bot_author():
    agent, mm = _agent_with_memory_manager()
    _build(agent, user_message="first turn", turn_author={"id": "bot:alpha", "name": "Alpha", "is_bot": True})
    assert agent._turn_author["id"] == "bot:alpha"

    _build(agent, user_message="second turn")

    assert agent._turn_author is None
    kwargs = mm.on_turn_start.call_args.kwargs
    assert kwargs["author_id"] is None
    assert kwargs["author_is_bot"] is False


def test_turn_author_reaches_the_agent_through_the_real_facade(monkeypatch):
    """``AIAgent.run_conversation(turn_author=...)`` crosses the facade and the loop entry point, not only
    ``build_turn_context``; a kwarg dropped at either hop raised TypeError on every real turn."""
    from types import SimpleNamespace
    from run_agent import AIAgent

    class _Completions:
        def create(self, **_kw):
            return SimpleNamespace(choices=[SimpleNamespace(
                message=SimpleNamespace(content="ok", tool_calls=None, reasoning=None, reasoning_content=None),
                finish_reason="stop")], usage=None, model="test-model")

    monkeypatch.setattr("agent.process_bootstrap.OpenAI",
                        lambda **_kw: SimpleNamespace(chat=SimpleNamespace(completions=_Completions())))
    monkeypatch.setattr("model_tools.get_tool_definitions", lambda *a, **k: [])
    agent = AIAgent(model="test-model", api_key="k", base_url="http://localhost:1/v1", platform="cli",
                    max_iterations=2, quiet_mode=True, skip_memory=True)
    agent._disable_streaming = True

    agent.run_conversation("hi", turn_author={"id": "bot:alpha", "name": "Alpha", "is_bot": True})
    assert agent._turn_author == {"id": "bot:alpha", "name": "Alpha", "is_bot": True}
    agent.run_conversation("hi again")
    assert agent._turn_author is None


def test_turn_start_replaces_stale_parent_history_with_compression_child():
    agent = _FakeAgent()
    stale_history = [{"role": "user", "content": "stale parent"}]
    compacted_history = [
        {"role": "user", "content": "[CONTEXT COMPACTION] summary"},
        {"role": "assistant", "content": "child tail"},
    ]

    def _recover(_agent):
        _agent.session_id = "compression-child"
        return compacted_history

    log_context = MagicMock()
    with patch(
        "agent.turn_context.recover_rotated_compression_session",
        side_effect=_recover,
    ):
        ctx = _build(
            agent,
            conversation_history=stale_history,
            set_session_context=log_context,
        )

    assert agent.session_id == "compression-child"
    assert agent._current_turn_id.startswith("compression-child:")
    log_context.assert_called_once_with("compression-child")
    assert ctx.conversation_history == compacted_history
    assert ctx.messages[:-1] == compacted_history
    assert ctx.messages[-1]["role"] == "user"
    assert ctx.messages[-1]["content"] == "hello"
    assert isinstance(ctx.messages[-1]["timestamp"], float)
    assert all(message.get("content") != "stale parent" for message in ctx.messages)


def test_applies_agent_side_effects():
    agent = _FakeAgent()
    _build(agent)
    # Retry counters reset, guardrails reset, turn counted.
    assert agent._invalid_tool_retries == 0
    assert agent._tool_guardrails.reset_called is True
    assert agent._user_turn_count == 1
    # Crash-resilience persistence fired once.
    assert agent._persist_calls == 1
    # task/turn ids assigned on the agent.
    assert agent._current_task_id
    assert agent._current_turn_id


def test_pending_cli_message_uses_clean_override_for_api_local_note():
    """A noted API message reuses the clean staged dict and its DB marker."""
    agent = _FakeAgent()
    staged = {"role": "user", "content": "clean prompt", "_db_persisted": True}
    agent._pending_cli_user_message = staged

    ctx = _build(
        agent,
        user_message="[MODEL NOTE]\n\nclean prompt",
        persist_user_message="clean prompt",
    )

    assert ctx.messages[-1] is staged
    assert ctx.messages[-1]["content"] == "[MODEL NOTE]\n\nclean prompt"
    assert ctx.messages[-1]["_db_persisted"] is True
    assert isinstance(ctx.messages[-1]["timestamp"], float)
    assert agent._pending_cli_user_message is None


def test_recall_indicator_emitted_when_memory_injected():
    """When prefetch injects memory, the deterministic indicator is emitted."""
    agent = _FakeAgent()
    agent._emit_status = MagicMock()
    mm = MagicMock()
    mm.prefetch_all.return_value = "- recalled fact"
    mm.describe_recall.return_value = "👁️ Hindsight — recalled 2 memories"
    agent._memory_manager = mm

    # A substantive query — a trivial prompt ("hi", "hello") skips prefetch_all
    # entirely, so there'd be nothing to indicate. See is_trivial_prompt.
    _build(agent, user_message="what did we decide about the deploy pipeline?")

    agent._emit_status.assert_any_call("👁️ Hindsight — recalled 2 memories")


def test_recall_indicator_skipped_when_nothing_injected():
    """No memory injected → describe_recall isn't consulted, nothing emitted."""
    agent = _FakeAgent()
    agent._emit_status = MagicMock()
    mm = MagicMock()
    mm.prefetch_all.return_value = ""
    agent._memory_manager = mm

    # Substantive query so prefetch_all actually runs; it returns nothing, so the
    # indicator path must stay silent (as opposed to being skipped as trivial).
    _build(agent, user_message="what did we decide about the deploy pipeline?")

    mm.describe_recall.assert_not_called()
    for call in agent._emit_status.call_args_list:
        assert "👁️" not in str(call)


def test_ensure_db_session_runs_after_system_prompt_restore():
    """Regression for #45499.

    On a fresh API/gateway agent (``_cached_system_prompt is None``) the DB
    session row must be created AFTER the system prompt is restored/built, so
    the persisted snapshot is written non-NULL. If ``_ensure_db_session()``
    ran first it would insert ``system_prompt=NULL`` and trip the misleading
    "stored system prompt is null; rebuilding" warning plus a first-turn
    prefix cache miss.
    """
    agent = _FakeAgent()
    agent._cached_system_prompt = None  # fresh agent, no cached prompt yet

    def _restore(_agent, _system_message, _history):
        _agent._cached_system_prompt = "REBUILT-SYSTEM"

    _build(agent, restore_or_build_system_prompt=_restore)

    # The prompt was populated before the DB row was created.
    assert agent._ensure_db_prompt_at_call == "REBUILT-SYSTEM"
    assert agent._cached_system_prompt == "REBUILT-SYSTEM"


# ── Between-turns MCP refresh (cache-safe late-binding) ──────────────────────
#
# A slow MCP server that connects after the agent's build-time tool snapshot
# must become callable by the user's NEXT turn — without mutating an in-flight
# turn's cached request prefix. The prologue is exactly that boundary, so the
# refresh hook lives here. These assert the contract (R1/R2/R6 in the spec),
# not timing permutations.


def test_between_turns_refresh_adds_late_tool_when_servers_registered():
    """R1: a tool that registered since build lands in this turn's snapshot."""
    agent = _FakeAgent()

    new_def = {"type": "function", "function": {"name": "mcp_x_tool", "description": "", "parameters": {}}}

    import model_tools
    import tools.mcp_tool  # noqa: F401 — the prologue's import-cost gate requires it in sys.modules
    with patch("tools.mcp_tool_discovery.has_registered_mcp_tools", return_value=True), \
         patch.object(model_tools, "get_tool_definitions", return_value=[new_def]):
        _build(agent)

    assert "mcp_x_tool" in agent.valid_tool_names
    assert any(t["function"]["name"] == "mcp_x_tool" for t in agent.tools)


class _TitlingAgent:
    """Only what ``_maybe_title_session_at_turn_start`` reads off an agent."""

    def __init__(self, platform):
        self.platform = platform
        self.session_id = "sess-1"
        self.model = "test/model"
        self.provider = "openrouter"
        self.base_url = "https://openrouter.ai/api/v1"
        self.api_key = "sk-x"
        self.api_mode = "chat_completions"
        self._session_db = MagicMock()
        self._session_db_created = True


def _title_turn(platform, message="Fix the login button"):
    """Run the prologue's titling step and return the maybe_auto_title mock."""
    from agent import turn_context

    with patch("agent.title_generator.maybe_auto_title") as titler:
        turn_context._maybe_title_session_at_turn_start(
            _TitlingAgent(platform),
            [{"role": "user", "content": message}],
        )
    return titler


@pytest.mark.parametrize("platform", ["cli", "telegram", "desktop", "acp", None])
def test_prologue_titles_the_surfaces_a_person_reads(platform):
    assert _title_turn(platform).called


@pytest.mark.parametrize("platform", ["cron", "CRON", "subagent"])
def test_prologue_does_not_title_machine_driven_runs(platform):
    """Cron names its own session after the job, and nobody opens a subagent's.

    Both would otherwise pay a side-LLM call per run for a name that is either
    overwritten or never read.
    """
    assert not _title_turn(platform).called


def test_prologue_names_a_subagent_run_after_its_goal_without_a_model_call():
    """A delegate run gets ``Subagent: <goal>`` at derived authority so it reads as machinery
    wherever ``sessions.show_subagents`` lists it, instead of staying untitled (#97202)."""
    from agent import turn_context

    agent = _TitlingAgent("subagent")
    agent._session_db.set_auto_title.return_value = True
    with patch("agent.title_generator.maybe_auto_title") as titler:
        turn_context._maybe_title_session_at_turn_start(
            agent, [{"role": "user", "content": "Audit the billing module\n\nContext: ..."}])
    assert not titler.called
    agent._session_db.set_auto_title.assert_called_once_with(
        "sess-1", "Subagent: Audit the billing module", source="derived")


def test_prologue_leaves_cron_runs_untitled():
    agent = _TitlingAgent("cron")
    from agent import turn_context

    turn_context._maybe_title_session_at_turn_start(agent, [{"role": "user", "content": "Run the job"}])
    assert not agent._session_db.set_auto_title.called


def test_prologue_forwards_the_submit_title_preview_to_the_titler():
    """A paste-shrunk ``display_metadata.title_preview`` from prompt.submit is the text the
    titler should read, not the full pasted body."""
    from agent import turn_context

    with patch("agent.title_generator.maybe_auto_title") as titler:
        turn_context._maybe_title_session_at_turn_start(
            _TitlingAgent("desktop"),
            [{"role": "user", "content": "x" * 5000,
              "display_metadata": {"title_preview": "Pasted 5000 chars"}}],
        )
    assert titler.call_args.kwargs["title_preview"] == "Pasted 5000 chars"
