"""Focused tests for API server session-control endpoints."""

import asyncio
import json
import threading
from unittest.mock import AsyncMock, MagicMock, patch

import pytest
from aiohttp import web
from aiohttp.test_utils import TestClient, TestServer

from gateway.config import PlatformConfig
from gateway.platforms.api_server import APIServerAdapter
from hermes_state import SessionDB


@pytest.fixture
def session_db(tmp_path):
    db = SessionDB(tmp_path / "state.db")
    try:
        yield db
    finally:
        close = getattr(db, "close", None)
        if callable(close):
            close()


@pytest.fixture
def adapter(session_db):
    adapter = APIServerAdapter(PlatformConfig(enabled=True))
    adapter._session_db = session_db
    return adapter


@pytest.fixture
def auth_adapter(session_db):
    adapter = APIServerAdapter(PlatformConfig(enabled=True, extra={"key": "sk-test"}))
    adapter._session_db = session_db
    return adapter


def _create_session_app(adapter: APIServerAdapter) -> web.Application:
    app = web.Application()
    app.router.add_get("/v1/capabilities", adapter._handle_capabilities)
    app.router.add_get("/api/sessions", adapter._handle_list_sessions)
    app.router.add_post("/api/sessions", adapter._handle_create_session)
    app.router.add_get("/api/sessions/{session_id}", adapter._handle_get_session)
    app.router.add_patch("/api/sessions/{session_id}", adapter._handle_patch_session)
    app.router.add_delete("/api/sessions/{session_id}", adapter._handle_delete_session)
    app.router.add_get("/api/sessions/{session_id}/messages", adapter._handle_session_messages)
    app.router.add_post("/api/sessions/{session_id}/fork", adapter._handle_fork_session)
    app.router.add_post("/api/sessions/{session_id}/chat", adapter._handle_session_chat)
    app.router.add_post("/api/sessions/{session_id}/chat/stream", adapter._handle_session_chat_stream)
    return app


@pytest.mark.asyncio
async def test_capabilities_advertises_session_control_surface(adapter):
    app = _create_session_app(adapter)
    async with TestClient(TestServer(app)) as cli:
        resp = await cli.get("/v1/capabilities")
        assert resp.status == 200
        data = await resp.json()

    features = data["features"]
    assert features["session_resources"] is True
    assert features["session_chat"] is True
    assert features["session_chat_streaming"] is True
    assert features["session_fork"] is True
    assert features["run_steer"] is True
    assert data["endpoints"]["sessions"] == {"method": "GET", "path": "/api/sessions"}
    assert data["endpoints"]["session_chat_stream"] == {
        "method": "POST",
        "path": "/api/sessions/{session_id}/chat/stream",
    }
    assert data["endpoints"]["run_steer"] == {
        "method": "POST",
        "path": "/v1/runs/{run_id}/steer",
    }


@pytest.mark.asyncio
async def test_get_session_projects_delegate_provenance_without_model_config(
    adapter, session_db
):
    session_db.create_session("parent", "desktop")
    child_id = session_db.create_session(
        "delegate-child",
        "desktop",
        parent_session_id="parent",
        model_config={"_delegate_from": "parent", "api_key": "must-not-leak"},
    )
    branch_id = session_db.create_session(
        "visible-branch",
        "desktop",
        parent_session_id="parent",
        model_config={"_branched_from": "parent"},
    )

    app = _create_session_app(adapter)
    async with TestClient(TestServer(app)) as cli:
        child_resp = await cli.get(f"/api/sessions/{child_id}")
        branch_resp = await cli.get(f"/api/sessions/{branch_id}")
        assert child_resp.status == 200
        assert branch_resp.status == 200
        child = (await child_resp.json())["session"]
        branch = (await branch_resp.json())["session"]

    assert child["is_internal_child"] is True
    assert branch["is_internal_child"] is False
    assert "model_config" not in child
    assert "must-not-leak" not in str(child)


@pytest.mark.asyncio
async def test_session_messages_default_to_latest_bounded_page(adapter, session_db):
    session_id = session_db.create_session("bounded-messages", "api_server")
    session_db.replace_messages(
        session_id,
        [{"role": "user", "content": f"msg {i}"} for i in range(501)],
    )

    app = _create_session_app(adapter)
    async with TestClient(TestServer(app)) as cli:
        resp = await cli.get(f"/api/sessions/{session_id}/messages")
        assert resp.status == 200
        payload = await resp.json()

        explicit_resp = await cli.get(
            f"/api/sessions/{session_id}/messages?limit=2&offset=1"
        )
        assert explicit_resp.status == 200
        explicit = await explicit_resp.json()

    assert payload["pagination"] == {
        "limit": 500,
        "offset": 0,
        "order": "latest",
        "returned": 500,
    }
    assert payload["data"][0]["content"] == "msg 1"
    assert payload["data"][-1]["content"] == "msg 500"
    assert [message["content"] for message in explicit["data"]] == [
        "msg 1",
        "msg 2",
    ]


@pytest.mark.asyncio
async def test_forked_session_stays_listable_and_parent_survives_failed_fork(adapter, session_db):
    """A fork is created before the parent is ended (#11030) and carries the explicit branch
    marker, so it still shows in the default listing (the timestamp fallback no longer holds)."""
    session_db.create_session("parent", "api_server")
    app = _create_session_app(adapter)
    async with TestClient(TestServer(app)) as cli:
        resp = await cli.post("/api/sessions/parent/fork", json={"id": "child"})
        assert resp.status == 201
        listed = await (await cli.get("/api/sessions")).json()
        ids = {row["id"] for row in listed["data"]}
        assert {"parent", "child"} <= ids, ids

        session_db.create_session("solo", "api_server")
        with patch.object(session_db, "create_session", side_effect=RuntimeError("boom")):
            resp = await cli.post("/api/sessions/solo/fork", json={"id": "never"})
        assert resp.status >= 500
    assert session_db.get_session("solo")["end_reason"] is None


@pytest.mark.asyncio
async def test_list_sessions_resurrects_bot_chat_off_the_event_loop(adapter, session_db, monkeypatch):
    """Canonical Bot Chat recovery must not run SQLite work in the HTTP loop."""
    session_id = session_db.create_session("archived-bot-chat", "gateway_botmode")
    assert session_db.set_session_title(session_id, "Bot Chat")
    session_db.end_session(session_id, "ws_orphan_reap")
    assert session_db.set_session_archived(session_id, True)

    loop_thread = threading.get_ident()
    call_threads = {}
    get_by_title = session_db.get_session_by_title
    unarchive = session_db.unarchive_recoverable_session

    def record_get_by_title(title):
        call_threads["get_session_by_title"] = threading.get_ident()
        return get_by_title(title)

    def record_unarchive(stale_id):
        call_threads["unarchive_recoverable_session"] = threading.get_ident()
        return unarchive(stale_id)

    monkeypatch.setattr(session_db, "get_session_by_title", record_get_by_title)
    monkeypatch.setattr(session_db, "unarchive_recoverable_session", record_unarchive)

    app = _create_session_app(adapter)
    async with TestClient(TestServer(app)) as cli:
        response = await cli.get("/api/sessions?title=Bot%20Chat")
        payload = await response.json()

    assert response.status == 200
    assert [session["id"] for session in payload["data"]] == [session_id]
    assert set(call_threads) == {"get_session_by_title", "unarchive_recoverable_session"}
    assert all(thread_id != loop_thread for thread_id in call_threads.values())
    assert not session_db.get_session(session_id)["archived"]


@pytest.mark.asyncio
async def test_session_model_lock_persists_off_the_event_loop(adapter, session_db, monkeypatch):
    """POST /api/sessions/{id}/model writes the lock row through a worker thread: the same
    contended-write class as the Bot Chat resurrection, on a sibling handler."""
    session_id = session_db.create_session("lock-off-loop", "api_server", model="gpt-5.5")
    loop_thread = threading.get_ident()
    seen = []
    real = session_db.update_session_runtime_lock

    def record(*args, **kwargs):
        seen.append(threading.get_ident())
        return real(*args, **kwargs)

    monkeypatch.setattr(session_db, "update_session_runtime_lock", record)
    app = _create_session_app(adapter)
    _register_session_model_route(app, adapter)
    with patch.object(adapter, "_resolve_route", return_value=None):
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(
                f"/api/sessions/{session_id}/model",
                json={"provider": "nous", "model": "x-ai/grok-4.5", "require_model_lock": True})
            assert resp.status == 200, await resp.text()
    assert seen and all(tid != loop_thread for tid in seen)
    assert session_db.get_session(session_id)["model"] == "x-ai/grok-4.5"


@pytest.mark.asyncio
async def test_session_messages_returns_compression_ancestors(adapter, session_db):
    """GET /api/sessions/{id}/messages on a compression continuation returns the
    root→tip transcript, not just the tip's rows (#51058)."""
    source_id = session_db.create_session("compress-source", "api_server")
    session_db.replace_messages(
        source_id,
        [
            {"role": "user", "content": "before compression"},
            {"role": "assistant", "content": "before answer"},
        ],
    )
    session_db.end_session(source_id, "compression")
    child_id = session_db.create_session(
        "compress-tip", "api_server", parent_session_id=source_id
    )
    session_db.append_message(child_id, role="user", content="after compression")

    app = _create_session_app(adapter)
    async with TestClient(TestServer(app)) as cli:
        resp = await cli.get(f"/api/sessions/{child_id}/messages")
        assert resp.status == 200
        payload = await resp.json()

    assert payload["session_id"] == child_id
    assert [m["content"] for m in payload["data"]] == [
        "before compression",
        "before answer",
        "after compression",
    ]
    assert [m["role"] for m in payload["data"]] == ["user", "assistant", "user"]


@pytest.mark.asyncio
async def test_fork_session_writes_branched_from_marker(adapter, session_db):
    """The API fork must stamp _branched_from like the CLI/TUI branch paths, so the
    fork is never misclassified as a compression continuation."""
    source_id = session_db.create_session("fork-source", "api_server")
    session_db.replace_messages(source_id, [{"role": "user", "content": "hello"}])

    app = _create_session_app(adapter)
    async with TestClient(TestServer(app)) as cli:
        resp = await cli.post(f"/api/sessions/{source_id}/fork", json={"id": "fork-child"})
        assert resp.status == 201
        payload = await resp.json()

    assert payload["session"]["id"] == "fork-child"
    fork = session_db.get_session("fork-child")
    assert fork["parent_session_id"] == source_id
    assert session_db._is_explicit_branch_session("fork-child")
    cfg = fork["model_config"]
    if isinstance(cfg, str):
        cfg = json.loads(cfg)
    assert cfg["_branched_from"] == source_id

@pytest.mark.asyncio
async def test_run_agent_binds_api_session_context_for_tool_env(adapter, monkeypatch):
    """API-server request sessions should reach tools and terminal subprocess env."""
    monkeypatch.setenv("HERMES_SESSION_ID", "stale-session")
    observed = {}

    class FakeAgent:
        session_prompt_tokens = 0
        session_completion_tokens = 0
        session_total_tokens = 0

        def __init__(self, session_id: str):
            self.session_id = session_id

        def run_conversation(self, user_message, conversation_history, task_id):
            from gateway.session_context import get_session_env
            from tools.environments.local import _make_run_env

            observed["task_id"] = task_id
            observed["context_session_id"] = get_session_env("HERMES_SESSION_ID")
            observed["context_platform"] = get_session_env("HERMES_SESSION_PLATFORM")
            observed["context_session_key"] = get_session_env("HERMES_SESSION_KEY")
            observed["child_session_id"] = _make_run_env({}).get("HERMES_SESSION_ID")
            return {"final_response": "ok"}

    def fake_create_agent(**kwargs):
        return FakeAgent(kwargs["session_id"])

    monkeypatch.setattr(adapter, "_create_agent", fake_create_agent)

    result, usage = await adapter._run_agent(
        user_message="hello",
        conversation_history=[],
        session_id="request-session",
        gateway_session_key="request-key",
    )

    assert result["session_id"] == "request-session"
    assert usage["input_tokens"] == 0
    assert usage["output_tokens"] == 0
    assert usage["total_tokens"] == 0
    assert "runtime" not in usage
    assert observed == {
        "task_id": "request-session",
        "context_session_id": "request-session",
        "context_platform": "api_server",
        "context_session_key": "request-key",
        "child_session_id": "request-session",
    }


@pytest.mark.asyncio
async def test_run_agent_registers_active_run_id_for_steering(adapter, monkeypatch):
    observed = {}

    class FakeAgent:
        session_prompt_tokens = 0
        session_completion_tokens = 0
        session_total_tokens = 0

        def __init__(self, session_id: str):
            self.session_id = session_id

        def steer(self, text: str) -> bool:
            observed["steer_text"] = text
            return True

        def run_conversation(self, user_message, conversation_history, task_id):
            observed["registered"] = adapter._active_run_agents.get("run_steer_test") is self
            observed["task_id"] = task_id
            return {"final_response": "ok"}

    def fake_create_agent(**kwargs):
        return FakeAgent(kwargs["session_id"])

    monkeypatch.setattr(adapter, "_create_agent", fake_create_agent)

    result, usage = await adapter._run_agent(
        user_message="hello",
        conversation_history=[],
        session_id="request-session",
        active_run_id="run_steer_test",
    )

    assert result["session_id"] == "request-session"
    assert usage == {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
    assert observed == {"registered": True, "task_id": "request-session"}
    assert "run_steer_test" not in adapter._active_run_agents


@pytest.mark.asyncio
async def test_session_chat_stream_disconnect_keeps_control_refs_until_executor_finishes(
    adapter, session_db
):
    """Disconnects must interrupt the live run without dropping its control refs early."""
    session_id = session_db.create_session("disconnect-stream-session", "api_server")
    run_started = threading.Event()
    interrupt_called = threading.Event()
    allow_finish = threading.Event()
    write_calls = {"count": 0}

    class FakeAgent:
        session_prompt_tokens = 0
        session_completion_tokens = 0
        session_total_tokens = 0

        def __init__(self, stream_delta_callback):
            self._stream_delta_callback = stream_delta_callback
            self.session_id = session_id

        def interrupt(self, _message=None):
            interrupt_called.set()

        def run_conversation(self, user_message, conversation_history, task_id):
            del user_message, conversation_history, task_id
            run_started.set()
            self._stream_delta_callback("hello")
            allow_finish.wait(timeout=5)
            return {"final_response": "done", "session_id": session_id}

    class DisconnectingStreamResponse:
        async def prepare(self, request):
            del request

        async def write(self, payload):
            del payload
            write_calls["count"] += 1
            if write_calls["count"] >= 3:
                raise ConnectionResetError("simulated client disconnect")

    request = MagicMock()
    request.headers = {}
    request.match_info = {"session_id": session_id}

    def _create_agent(**kwargs):
        return FakeAgent(kwargs["stream_delta_callback"])

    with patch.object(
        adapter,
        "_get_existing_session_or_404",
        return_value=({"id": session_id}, None),
    ), patch.object(
        adapter,
        "_read_json_body",
        return_value=({"message": "stream please"}, None),
    ), patch.object(
        adapter,
        "_create_agent",
        side_effect=_create_agent,
    ), patch(
        "gateway.platforms.api_server.web.StreamResponse",
        return_value=DisconnectingStreamResponse(),
    ):
        handler_task = asyncio.create_task(adapter._handle_session_chat_stream(request))

        for _ in range(60):
            if run_started.is_set():
                break
            await asyncio.sleep(0.05)

        assert run_started.is_set()
        run_id = next(iter(adapter._run_statuses))

        for _ in range(40):
            if interrupt_called.is_set():
                break
            await asyncio.sleep(0.05)

        assert interrupt_called.is_set()
        assert run_id in adapter._active_run_agents
        # Not in _active_run_tasks: session-stream turns are counted via
        # _inflight_agent_runs; a task entry would double-count them in the
        # shutdown drain (active_agent_work_count).
        assert run_id not in adapter._active_run_tasks
        assert not handler_task.done()

        allow_finish.set()
        await handler_task

    assert run_id not in adapter._active_run_agents


@pytest.mark.asyncio
async def test_session_chat_stream_classifies_failed_tool_completions(adapter, session_db):
    session_id = session_db.create_session("tool-status-stream", "api_server")

    async def fake_run(**kwargs):
        progress = kwargs["tool_progress_callback"]
        progress("tool.completed", tool_name="read_file", is_error=False)
        progress("tool.completed", tool_name="terminal", is_error=True)
        progress("tool.failed", tool_name="web_search")
        return {"final_response": "done", "session_id": session_id}, {"total_tokens": 1}

    app = _create_session_app(adapter)
    with patch.object(adapter, "_run_agent", side_effect=fake_run):
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(
                f"/api/sessions/{session_id}/chat/stream",
                json={"message": "run tools"},
            )
            assert resp.status == 200
            body = await resp.text()

    blocks = body.split("\n\n")
    assert any("event: tool.completed" in b and '"tool_name": "read_file"' in b for b in blocks)
    assert any("event: tool.failed" in b and '"tool_name": "terminal"' in b for b in blocks)
    assert any("event: tool.failed" in b and '"tool_name": "web_search"' in b for b in blocks)
    assert body.count("event: tool.completed") == 1
    assert body.count("event: tool.failed") == 2


@pytest.mark.asyncio
async def test_session_chat_stream_run_completed_carries_turn_transcript(adapter, session_db):
    """run.completed must include the full interleaved turn transcript so a
    client that lost intermediate (pre-tool-call) assistant text from the live
    delta stream can reconcile without a separate /messages fetch. Refs #34703.
    """
    import json as _json

    session_id = session_db.create_session("transcript-session", "api_server")

    async def fake_run(**kwargs):
        # Stream the intermediate planning text the way a real turn would.
        kwargs["stream_delta_callback"]("Let me search for that:")
        kwargs["stream_delta_callback"]("Here is the summary.")
        result = {
            "final_response": "Here is the summary.",
            "session_id": session_id,
            "messages": [
                {"role": "user", "content": "search then summarize"},
                {
                    "role": "assistant",
                    "content": "Let me search for that:",
                    "tool_calls": [
                        {
                            "id": "call_1",
                            "type": "function",
                            "function": {"name": "web_search", "arguments": "{}"},
                        }
                    ],
                },
                {"role": "tool", "content": "results", "tool_call_id": "call_1", "tool_name": "web_search"},
                {"role": "assistant", "content": "Here is the summary."},
            ],
        }
        return result, {"total_tokens": 6}

    app = _create_session_app(adapter)
    with patch.object(adapter, "_run_agent", side_effect=fake_run):
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(
                f"/api/sessions/{session_id}/chat/stream",
                json={"message": "search then summarize"},
            )
            assert resp.status == 200
            body = await resp.text()

    # Pull the run.completed event payload out of the SSE body.
    run_completed_payload = None
    for block in body.split("\n\n"):
        if "event: run.completed" in block:
            for line in block.splitlines():
                if line.startswith("data: "):
                    run_completed_payload = _json.loads(line[len("data: "):])
            break
    assert run_completed_payload is not None, body
    messages = run_completed_payload.get("messages")
    assert isinstance(messages, list) and messages, run_completed_payload

    # The colon-ended intermediate text that preceded the tool call must be present.
    contents = [m.get("content") for m in messages]
    assert "Let me search for that:" in contents
    assert "Here is the summary." in contents
    # No prior-turn user message should leak into the per-turn slice.
    assert all(m.get("role") in ("assistant", "tool") for m in messages)
    # The tool call is preserved alongside the intermediate text.
    assert any(m.get("tool_calls") for m in messages)


@pytest.mark.asyncio
async def test_session_chat_stream_reports_interrupted_turn_as_not_completed(adapter, session_db):
    """The SSE terminal payload is derived from the result, never hard-coded: an interrupted
    turn streams ``completed: false`` / ``interrupted: true`` and ends with ``run.cancelled``,
    and the run status matches (#111770)."""
    import json as _json

    session_id = session_db.create_session("interrupted-session", "api_server")

    async def fake_run(**_kwargs):
        return {"final_response": "Operation interrupted.", "interrupted": True, "completed": False,
                "session_id": session_id}, {"total_tokens": 1}

    app = _create_session_app(adapter)
    with patch.object(adapter, "_run_agent", side_effect=fake_run):
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(f"/api/sessions/{session_id}/chat/stream", json={"message": "hello"})
            body = await resp.text()

    payloads = {}
    for block in body.split("\n\n"):
        lines = block.splitlines()
        event = next((ln[7:] for ln in lines if ln.startswith("event: ")), None)
        data = next((ln[6:] for ln in lines if ln.startswith("data: ")), None)
        if event and data:
            payloads[event] = _json.loads(data)

    assert payloads["assistant.completed"]["completed"] is False
    assert payloads["assistant.completed"]["interrupted"] is True
    assert "run.cancelled" in payloads and "run.completed" not in payloads
    assert payloads["run.cancelled"]["completed"] is False
    assert next(iter(adapter._run_statuses.values()))["status"] == "cancelled"


# ---------------------------------------------------------------------------
# Session-persisted model threading + provider-auth failure surfacing
# (salvaged from PR #57947 by @FvanW and PR #59941 by @kaishi00)
# ---------------------------------------------------------------------------


@pytest.mark.asyncio
async def test_session_chat_resolves_stored_model_route_alias(session_db, monkeypatch):
    """A session-persisted model that matches a model_routes alias must go
    through the route path (so route provider/credentials apply) and NOT be
    passed as a raw session_model (idea from PR #59941 by @kaishi00)."""
    adapter = APIServerAdapter(
        PlatformConfig(
            enabled=True,
            extra={"model_routes": {"alias": {"model": "route/model", "provider": "openrouter"}}},
        )
    )
    adapter._session_db = session_db
    session_id = session_db.create_session("route-pinned-session", "api_server", model="alias")

    mock_run = AsyncMock(return_value=({"final_response": "ok", "session_id": session_id}, {"total_tokens": 1}))
    app = _create_session_app(adapter)
    with patch.object(adapter, "_run_agent", mock_run):
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(
                f"/api/sessions/{session_id}/chat",
                json={"message": "hi"},
            )
            assert resp.status == 200

    _, kwargs = mock_run.call_args
    assert kwargs["route"] == {"model": "route/model", "provider": "openrouter"}
    assert kwargs["session_model"] is None


@pytest.mark.asyncio
async def test_session_chat_treats_pre_existing_poisoned_row_as_no_model(session_db):
    """A session row created before the alias-leak fix may still have the
    virtual model alias (e.g. "hermes-agent") persisted literally as its
    model. Reading that back must NOT thread it through as a raw
    session_model override — it must fall through to the global default,
    exactly like a row that never had a model at all (#session-model-
    alias-leak)."""
    adapter = APIServerAdapter(PlatformConfig(enabled=True))
    adapter._session_db = session_db
    session_id = session_db.create_session(
        "poisoned-session", "api_server", model=adapter._model_name
    )

    mock_run = AsyncMock(return_value=({"final_response": "ok", "session_id": session_id}, {"total_tokens": 1}))
    app = _create_session_app(adapter)
    with patch.object(adapter, "_run_agent", mock_run):
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(
                f"/api/sessions/{session_id}/chat",
                json={"message": "hi"},
            )
            assert resp.status == 200

    _, kwargs = mock_run.call_args
    assert kwargs["session_model"] is None


@pytest.mark.asyncio
async def test_session_chat_stream_treats_pre_existing_poisoned_row_as_no_model(session_db):
    """Streaming twin of the above: the SSE chat path must apply the same
    guard against a pre-existing poisoned session row."""
    adapter = APIServerAdapter(PlatformConfig(enabled=True))
    adapter._session_db = session_db
    session_id = session_db.create_session(
        "poisoned-stream-session", "api_server", model=adapter._model_name
    )

    async def fake_run(**kwargs):
        return {"final_response": "ok", "session_id": session_id}, {"total_tokens": 1}

    app = _create_session_app(adapter)
    with patch.object(adapter, "_run_agent", side_effect=fake_run) as mock_run:
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(
                f"/api/sessions/{session_id}/chat/stream",
                json={"message": "hi"},
            )
            assert resp.status == 200
            # Drain the SSE body: the 200 lands before the streaming task
            # invokes _run_agent, so asserting on call_args without reading
            # the body races the handler (flaked on loaded CI runners).
            await resp.text()

    _, kwargs = mock_run.call_args
    assert kwargs["session_model"] is None


def _register_session_model_route(app, adapter):
    app.router.add_post("/api/sessions/{session_id}/model", adapter._handle_session_model_lock)


def _patch_api_server_runtime(monkeypatch):
    monkeypatch.setattr(
        "gateway.run._resolve_runtime_agent_kwargs",
        lambda: {
            "provider": "openrouter",
            "api_key": "sk-global",
            "base_url": "https://openrouter.example/v1",
            "api_mode": "chat_completions",
        },
    )
    monkeypatch.setattr("gateway.run._resolve_gateway_model", lambda: "global/model")
    monkeypatch.setattr("gateway.run._load_gateway_config", lambda: {})
    monkeypatch.setattr(
        "gateway.run.GatewayRunner._load_reasoning_config",
        staticmethod(lambda model="": {}),
    )
    monkeypatch.setattr(
        "gateway.run.GatewayRunner._load_fallback_model",
        staticmethod(lambda: None),
    )
    monkeypatch.setattr("gateway.run._current_max_iterations", lambda: 90)
    monkeypatch.setattr("hermes_cli.tools_config._get_platform_tools", lambda *_: set())
    monkeypatch.setattr(
        "gateway.run._resolve_runtime_agent_kwargs_for_provider",
        lambda provider, target_model=None: {
            "provider": provider,
            "api_key": f"sk-{provider}",
            "base_url": f"https://{provider}.example/v1",
            "api_mode": "chat_completions",
        },
    )


@pytest.mark.asyncio
async def test_create_session_respects_browser_source_and_model_lock(adapter, session_db):
    app = _create_session_app(adapter)
    async with TestClient(TestServer(app)) as cli:
        resp = await cli.post(
            "/api/sessions",
            json={
                "id": "browser-lock-session",
                "source": "hermes_browser",
                "provider": "nous",
                "model": "x-ai/grok-4.5",
                "require_model_lock": True,
                "title": "Browser lock",
                "system_prompt": "browser prompt",
            },
        )
        assert resp.status == 201, await resp.text()
        payload = await resp.json()

    assert payload["session"]["source"] == "hermes_browser"
    assert payload["session"]["model"] == "x-ai/grok-4.5"
    row = session_db.get_session("browser-lock-session")
    assert row["source"] == "hermes_browser"
    assert row["model"] == "x-ai/grok-4.5"
    import json as _json
    model_config = row.get("model_config")
    if isinstance(model_config, str):
        model_config = _json.loads(model_config)
    assert model_config["browser_model_lock"]["provider"] == "nous"
    assert model_config["browser_model_lock"]["model"] == "x-ai/grok-4.5"
    assert model_config["browser_model_lock"]["confirmed"] is True


@pytest.mark.asyncio
async def test_session_model_lock_endpoint_then_chat_reuses_persisted_lock_and_provider_credentials(
    adapter,
    session_db,
    monkeypatch,
):
    session_id = session_db.create_session(
        "endpoint-lock-chat",
        "api_server",
        model="gpt-5.5",
        system_prompt="Conversation started:\nModel: gpt-5.5\nProvider: openai-codex\n",
    )
    captured = {}

    class FakeAgent:
        session_prompt_tokens = 0
        session_completion_tokens = 0
        session_total_tokens = 0

        def __init__(self, **kwargs):
            captured.update(kwargs)
            self.session_id = kwargs["session_id"]
            self.provider = kwargs.get("provider") or ""
            self.model = kwargs.get("model") or ""

        def run_conversation(self, user_message, conversation_history, task_id):
            return {"final_response": "locked", "session_id": self.session_id}

    _patch_api_server_runtime(monkeypatch)
    monkeypatch.setattr("run_agent.AIAgent", FakeAgent)
    monkeypatch.setattr(
        adapter,
        "_session_model_override_for",
        lambda *_: {
            "model": "session/override-model",
            "provider": "openai-codex",
            "api_key": "sk-session-override",
            "base_url": "https://override.example/v1",
            "api_mode": "codex_responses",
        },
    )

    app = _create_session_app(adapter)
    _register_session_model_route(app, adapter)
    with patch.object(adapter, "_resolve_route", return_value=None):
        async with TestClient(TestServer(app)) as cli:
            lock_resp = await cli.post(
                f"/api/sessions/{session_id}/model",
                json={
                    "provider": "nous",
                    "model": "x-ai/grok-4.5",
                    "require_model_lock": True,
                },
            )
            assert lock_resp.status == 200, await lock_resp.text()

            resp = await cli.post(
                f"/api/sessions/{session_id}/chat",
                json={"message": "use the stored lock"},
            )
            assert resp.status == 200, await resp.text()
            payload = await resp.json()

    assert captured["provider"] == "nous"
    assert captured["model"] == "x-ai/grok-4.5"
    assert captured["api_key"] == "sk-nous"
    assert captured["base_url"] == "https://nous.example/v1"
    assert payload["runtime"]["provider"] == "nous"
    assert payload["runtime"]["model"] == "x-ai/grok-4.5"
    assert payload["runtime"]["requested"] == {
        "provider": "nous",
        "model": "x-ai/grok-4.5",
    }
    assert payload["runtime"]["route_source"] == "session_model_lock"


@pytest.mark.asyncio
async def test_session_model_lock_endpoint_then_chat_stream_reuses_persisted_lock(
    adapter,
    session_db,
):
    session_id = session_db.create_session("endpoint-lock-stream", "api_server")
    captured = {}

    async def fake_run(**kwargs):
        captured.update(kwargs)
        kwargs["stream_delta_callback"]("hi")
        return (
            {
                "final_response": "hi",
                "session_id": session_id,
                "runtime": {
                    "provider": "nous",
                    "model": "x-ai/grok-4.5",
                    "requested": {"provider": "nous", "model": "x-ai/grok-4.5"},
                    "route_source": "session_model_lock",
                },
            },
            {
                "total_tokens": 1,
                "runtime": {
                    "provider": "nous",
                    "model": "x-ai/grok-4.5",
                    "requested": {"provider": "nous", "model": "x-ai/grok-4.5"},
                    "route_source": "session_model_lock",
                },
            },
        )

    app = _create_session_app(adapter)
    _register_session_model_route(app, adapter)
    with patch.object(adapter, "_resolve_route", return_value=None), patch.object(
        adapter,
        "_run_agent",
        side_effect=fake_run,
    ):
        async with TestClient(TestServer(app)) as cli:
            lock_resp = await cli.post(
                f"/api/sessions/{session_id}/model",
                json={
                    "provider": "nous",
                    "model": "x-ai/grok-4.5",
                    "require_model_lock": True,
                },
            )
            assert lock_resp.status == 200, await lock_resp.text()

            resp = await cli.post(
                f"/api/sessions/{session_id}/chat/stream",
                json={"message": "stream with stored lock"},
            )
            assert resp.status == 200, await resp.text()
            body = await resp.text()

    assert captured["route"] == {"provider": "nous", "model": "x-ai/grok-4.5"}
    assert captured["requested_runtime"]["provider"] == "nous"
    assert captured["requested_runtime"]["model"] == "x-ai/grok-4.5"
    assert captured["route_source"] == "session_model_lock"
    assert "x-ai/grok-4.5" in body


@pytest.mark.asyncio
async def test_run_agent_reports_actual_agent_runtime_not_requested_metadata(adapter, monkeypatch):
    class FakeAgent:
        session_prompt_tokens = 0
        session_completion_tokens = 0
        session_total_tokens = 0

        def __init__(self):
            self.session_id = "runtime-session"
            self.provider = "actual-provider"
            self.model = "actual-model"
            self._hermes_api_runtime = {
                "provider": "requested-provider",
                "model": "requested-model",
                "route_source": "raw_request",
            }

        def run_conversation(self, user_message, conversation_history, task_id):
            return {"final_response": "ok", "session_id": self.session_id}

    monkeypatch.setattr(adapter, "_create_agent", lambda **kwargs: FakeAgent())

    result, usage = await adapter._run_agent(
        user_message="hello",
        conversation_history=[],
        session_id="runtime-session",
        route={"provider": "requested-provider", "model": "requested-model"},
        requested_runtime={
            "provider": "requested-provider",
            "model": "requested-model",
        },
        route_source="session_model_lock",
    )

    assert result["runtime"]["provider"] == "actual-provider"
    assert result["runtime"]["model"] == "actual-model"
    assert result["runtime"]["requested"] == {
        "provider": "requested-provider",
        "model": "requested-model",
    }
    assert usage["runtime"]["provider"] == "actual-provider"
    assert usage["runtime"]["model"] == "actual-model"


@pytest.mark.asyncio
async def test_confirmed_runtime_lock_rejects_actual_runtime_mismatch(adapter, monkeypatch):
    class FakeAgent:
        session_prompt_tokens = 0
        session_completion_tokens = 0
        session_total_tokens = 0
        session_id = "mismatch-session"
        provider = "fallback-provider"
        model = "fallback-model"

        def run_conversation(self, user_message, conversation_history, task_id):
            return {"final_response": "wrong runtime", "session_id": self.session_id}

    monkeypatch.setattr(adapter, "_create_agent", lambda **kwargs: FakeAgent())

    with pytest.raises(RuntimeError, match="confirmed model lock runtime mismatch"):
        await adapter._run_agent(
            user_message="hello",
            conversation_history=[],
            session_id="mismatch-session",
            route={"provider": "nous", "model": "x-ai/grok-4.5"},
            requested_runtime={"provider": "nous", "model": "x-ai/grok-4.5"},
            route_source="session_model_lock",
            confirmed_runtime_lock=True,
        )


def test_confirmed_runtime_lock_rejects_provider_only_mismatch(adapter):
    class FakeAgent:
        provider = "fallback-provider"
        model = "some-model"
        _hermes_api_runtime = {
            "provider": "nous",
            "model": "some-model",
            "route_source": "session_model_lock",
        }

    with pytest.raises(RuntimeError, match="confirmed model lock runtime mismatch"):
        adapter._turn_runtime_metadata(
            FakeAgent(),
            route={"provider": "nous", "model": "some-model"},
            requested_runtime={"provider": "nous", "model": "some-model"},
            route_source="session_model_lock",
            confirmed_runtime_lock=True,
        )


@pytest.mark.parametrize(
    ("resolved_provider", "actual_provider"),
    [("custom", "custom"), ("my-endpoint", "my-endpoint")],
    ids=["custom-family", "plugin-canonical-name"],
)
def test_confirmed_runtime_lock_accepts_resolved_provider_identity(
    adapter, resolved_provider, actual_provider
):
    class FakeAgent:
        provider = actual_provider
        model = "some-model"
        _hermes_api_runtime = {
            "provider": resolved_provider,
            "model": "some-model",
            "route_source": "session_model_lock",
        }

    runtime = adapter._turn_runtime_metadata(
        FakeAgent(),
        route={"provider": "custom:my-endpoint", "model": "some-model"},
        requested_runtime={"provider": "custom:my-endpoint", "model": "some-model"},
        route_source="session_model_lock",
        confirmed_runtime_lock=True,
    )

    assert runtime["provider"] == actual_provider
    assert runtime["requested"]["provider"] == "custom:my-endpoint"


def test_confirmed_runtime_lock_disables_global_fallback_model(adapter, monkeypatch):
    _patch_api_server_runtime(monkeypatch)
    monkeypatch.setattr(
        "gateway.run.GatewayRunner._load_fallback_model",
        staticmethod(lambda: "openrouter/fallback-model"),
    )
    captured = {}

    class FakeAgent:
        provider = "nous"
        model = "x-ai/grok-4.5"

        def __init__(self, **kwargs):
            captured.update(kwargs)

    monkeypatch.setattr("run_agent.AIAgent", FakeAgent)

    adapter._create_agent(
        session_id="locked-session",
        route={"provider": "nous", "model": "x-ai/grok-4.5"},
        confirmed_runtime_lock=True,
    )

    assert captured["fallback_model"] is None


@pytest.mark.asyncio
async def test_unconfirmed_request_does_not_replace_confirmed_session_lock(adapter, session_db):
    session_id = session_db.create_session("one-off-override", "api_server")
    session_db.update_session_runtime_lock(
        session_id,
        provider="nous",
        model="x-ai/grok-4.5",
        route_source="raw_request",
        confirmed=True,
    )
    mock_run = AsyncMock(
        return_value=(
            {
                "final_response": "ok",
                "session_id": session_id,
                "runtime": {"provider": "openrouter", "model": "anthropic/claude-sonnet"},
            },
            {"total_tokens": 1},
        )
    )
    app = _create_session_app(adapter)
    with patch.object(adapter, "_resolve_route", return_value=None), patch.object(
        adapter,
        "_run_agent",
        mock_run,
    ):
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(
                f"/api/sessions/{session_id}/chat",
                json={
                    "message": "one turn only",
                    "provider": "openrouter",
                    "model": "anthropic/claude-sonnet",
                },
            )
            assert resp.status == 200, await resp.text()

    import json as _json

    row = session_db.get_session(session_id)
    config = row["model_config"]
    if isinstance(config, str):
        config = _json.loads(config)
    assert config["browser_model_lock"]["provider"] == "nous"
    assert config["browser_model_lock"]["model"] == "x-ai/grok-4.5"
    assert config["browser_model_lock"]["confirmed"] is True


@pytest.mark.asyncio
async def test_require_model_lock_hard_fails_when_global_default_would_be_used(adapter, session_db, monkeypatch):
    session_id = session_db.create_session("lock-fail-session", "api_server")
    monkeypatch.setattr(adapter, "_model_name", "gpt-5.5")
    app = _create_session_app(adapter)
    with patch.object(adapter, "_resolve_route", return_value=None), patch.object(adapter, "_run_agent", new_callable=AsyncMock) as mock_run:
        async with TestClient(TestServer(app)) as cli:
            # empty model + require_model_lock must not silently fall through
            resp = await cli.post(
                f"/api/sessions/{session_id}/chat",
                json={
                    "message": "hello",
                    "provider": "nous",
                    "model": "",
                    "require_model_lock": True,
                },
            )
            assert resp.status in (400, 409), await resp.text()
            body = await resp.json()
            assert body["error"]["code"] in {"model_lock_unavailable", "invalid_model_lock", "missing_model"}
    mock_run.assert_not_called()


_CHAT_REPLY = ({"final_response": "ok"}, {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0})


@pytest.mark.asyncio
@pytest.mark.parametrize("body, expected", [
    ({"message": "hello", "author": {"id": " bot:dixie ", "name": "dixie", "is_bot": 1, "role": "admin"}},
     {"id": "bot:dixie", "name": "dixie", "is_bot": True}),
    ({"message": "hello"}, None),
    ({"message": "hello", "author": None}, None),
], ids=["author", "no author", "null author"])
async def test_session_chat_passes_normalized_author_to_run_agent(adapter, session_db, body, expected):
    """A body ``author`` reaches ``_run_agent`` normalized with unknown keys dropped; absent or null is None."""
    session_id = session_db.create_session("author-session", "api_server")
    app = _create_session_app(adapter)
    with patch.object(adapter, "_run_agent", AsyncMock(return_value=_CHAT_REPLY)) as mock_run:
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(f"/api/sessions/{session_id}/chat", json=body)
            assert resp.status == 200, await resp.text()
    assert mock_run.call_args.kwargs["turn_author"] == expected


@pytest.mark.asyncio
@pytest.mark.parametrize("suffix", ["/chat", "/chat/stream"])
@pytest.mark.parametrize("author", ["dixie"])
async def test_session_chat_rejects_non_object_author(adapter, session_db, suffix, author):
    session_id = session_db.create_session("bad-author-session", "api_server")
    app = _create_session_app(adapter)
    with patch.object(adapter, "_run_agent", new_callable=AsyncMock) as mock_run:
        async with TestClient(TestServer(app)) as cli:
            resp = await cli.post(
                f"/api/sessions/{session_id}{suffix}", json={"message": "hello", "author": author})
            assert resp.status == 400, await resp.text()
            body = await resp.json()
    assert body["error"]["code"] == "invalid_author"
    mock_run.assert_not_called()


@pytest.mark.asyncio
async def test_run_agent_forwards_author_to_run_conversation_only_when_set(adapter, monkeypatch):
    """``turn_author`` reaches ``run_conversation`` when set; a human turn keeps today's call shape."""
    calls = []

    class FakeAgent:
        session_prompt_tokens = 0
        session_completion_tokens = 0
        session_total_tokens = 0
        session_id = "author-run"

        def run_conversation(self, user_message, conversation_history, task_id, **kwargs):
            calls.append(kwargs)
            return {"final_response": "ok", "session_id": self.session_id}

    monkeypatch.setattr(adapter, "_create_agent", lambda **kwargs: FakeAgent())
    author = {"id": "bot:dixie", "name": "dixie", "is_bot": True}
    await adapter._run_agent(
        user_message="hello", conversation_history=[], session_id="author-run", turn_author=author)
    await adapter._run_agent(user_message="hello", conversation_history=[], session_id="author-run")

    assert calls == [{"turn_author": author}, {}]


@pytest.mark.asyncio
async def test_patch_session_persists_pinned_and_archived(adapter, session_db):
    """PATCH must accept the durable pin/archive flags and round-trip them.

    These were rejected as unsupported fields, so every pin the desktop made
    400'd silently (the client swallows the error) and the pin only ever lived
    in that one app's localStorage. The auto-archive sweep reads
    `sessions.pinned` server-side, so an unpersisted pin does not protect the
    chat it was supposed to keep.
    """
    session_id = session_db.create_session("pin-session", "api_server")
    app = _create_session_app(adapter)

    async with TestClient(TestServer(app)) as cli:
        resp = await cli.patch(f"/api/sessions/{session_id}", json={"pinned": True})
        assert resp.status == 200, await resp.text()
        assert (await resp.json())["session"]["pinned"] is True

        # The flag is durable, not just echoed back from the request body.
        assert bool(session_db.get_session(session_id)["pinned"]) is True

        resp = await cli.get(f"/api/sessions/{session_id}")
        assert (await resp.json())["session"]["pinned"] is True

        resp = await cli.patch(f"/api/sessions/{session_id}", json={"pinned": False})
        assert (await resp.json())["session"]["pinned"] is False
        assert bool(session_db.get_session(session_id)["pinned"]) is False

        resp = await cli.patch(f"/api/sessions/{session_id}", json={"archived": True})
        assert (await resp.json())["session"]["archived"] is True
        assert bool(session_db.get_session(session_id)["archived"]) is True


@pytest.mark.asyncio
async def test_patch_session_rejects_non_boolean_pinned(adapter, session_db):
    session_id = session_db.create_session("pin-type-session", "api_server")
    app = _create_session_app(adapter)

    async with TestClient(TestServer(app)) as cli:
        resp = await cli.patch(f"/api/sessions/{session_id}", json={"pinned": "yes"})
        assert resp.status == 400, await resp.text()
        assert (await resp.json())["error"]["code"] == "invalid_session_field"


@pytest.mark.asyncio
async def test_patch_session_still_rejects_unknown_fields(adapter, session_db):
    session_id = session_db.create_session("unknown-field-session", "api_server")
    app = _create_session_app(adapter)

    async with TestClient(TestServer(app)) as cli:
        resp = await cli.patch(f"/api/sessions/{session_id}", json={"nonsense": 1})
        assert resp.status == 400, await resp.text()
        assert (await resp.json())["error"]["code"] == "unsupported_session_field"


@pytest.mark.asyncio
async def test_session_stream_records_reply_text_for_post_disconnect_recovery(
    adapter, session_db
):
    """A caller whose socket died must still be able to read what the agent said.

    The session-stream route put the reply only on the SSE queue, so a client
    that lost its connection saw the run reach "completed" with no way to learn
    the text — indistinguishable from, and as useless as, a run that produced
    nothing. POST /v1/runs has always recorded `output`; this pins the same for
    this route, which is what makes GET /v1/runs/{run_id} a recovery path.
    """
    session_id = session_db.create_session("recover-stream-session", "api_server")
    run_started = threading.Event()
    allow_finish = threading.Event()
    write_calls = {"count": 0}

    class FakeAgent:
        session_prompt_tokens = 0
        session_completion_tokens = 0
        session_total_tokens = 0

        def __init__(self, stream_delta_callback):
            self._stream_delta_callback = stream_delta_callback
            self.session_id = session_id

        def interrupt(self, _message=None):
            allow_finish.set()

        def run_conversation(self, user_message, conversation_history, task_id):
            del user_message, conversation_history, task_id
            run_started.set()
            self._stream_delta_callback("partial ")
            allow_finish.wait(timeout=5)
            return {"final_response": "the answer worth keeping", "session_id": session_id}

    class DisconnectingStreamResponse:
        async def prepare(self, request):
            del request

        async def write(self, payload):
            del payload
            write_calls["count"] += 1
            if write_calls["count"] >= 3:
                raise ConnectionResetError("simulated client disconnect")

    request = MagicMock()
    request.headers = {}
    request.match_info = {"session_id": session_id}

    with patch.object(
        adapter, "_get_existing_session_or_404", return_value=({"id": session_id}, None)
    ), patch.object(
        adapter, "_read_json_body", return_value=({"message": "stream please"}, None)
    ), patch.object(
        adapter, "_create_agent", side_effect=lambda **kw: FakeAgent(kw["stream_delta_callback"])
    ), patch(
        "gateway.platforms.api_server.web.StreamResponse",
        return_value=DisconnectingStreamResponse(),
    ):
        handler_task = asyncio.create_task(adapter._handle_session_chat_stream(request))

        for _ in range(60):
            if run_started.is_set():
                break
            await asyncio.sleep(0.05)
        assert run_started.is_set()
        run_id = next(iter(adapter._run_statuses))

        allow_finish.set()
        await handler_task

    record = adapter._run_statuses[run_id]
    assert record["status"] == "completed"
    # The whole point: the text survived the dead socket.
    assert record.get("output") == "the answer worth keeping"

    # And it is reachable through the documented read path, not just the dict.
    get_request = MagicMock()
    get_request.headers = {}
    get_request.match_info = {"run_id": run_id}
    response = await adapter._handle_get_run(get_request)
    assert response.status == 200
    assert "the answer worth keeping" in response.text


@pytest.mark.asyncio
async def test_interim_commentary_reaches_session_sse_and_responses_stream(adapter, session_db, monkeypatch):
    """Codex commentary / mid-turn assistant text is a typed ``assistant.commentary`` event on the
    session SSE endpoint and a ``phase: commentary`` message item on /v1/responses, never part of
    the final answer; ``display.interim_assistant_messages: false`` installs no callback (#67580)."""
    import json as _json

    session_id = session_db.create_session("commentary-session", "api_server")

    def fake_create_agent(**kwargs):
        interim = kwargs["interim_assistant_callback"]

        class FakeAgent:
            provider, model = "openai-codex", "gpt-5"
            session_prompt_tokens = session_completion_tokens = session_total_tokens = 0

            def run_conversation(self, **_kw):
                interim("Checking the docs first.", already_streamed=False)
                return {"final_response": "Done.", "messages": [], "api_calls": 1}

        return FakeAgent()

    app = _create_session_app(adapter)
    app.router.add_post("/v1/responses", adapter._handle_responses)
    with patch.object(adapter, "_create_agent", side_effect=fake_create_agent):
        async with TestClient(TestServer(app)) as cli:
            sse = await (await cli.post(f"/api/sessions/{session_id}/chat/stream", json={"message": "go"})).text()
            responses = await (await cli.post(
                "/v1/responses", json={"model": "hermes-agent", "input": "go", "stream": True})).text()

    def _events(body):
        out = []
        for block in body.split("\n\n"):
            lines = block.splitlines()
            name = next((ln[7:] for ln in lines if ln.startswith("event: ")), None)
            data = next((ln[6:] for ln in lines if ln.startswith("data: ")), None)
            if data:
                out.append((name, _json.loads(data)))
        return out

    sse_events = _events(sse)
    commentary = [d for n, d in sse_events if n == "assistant.commentary"]
    assert [(d["text"], d["already_streamed"]) for d in commentary] == [("Checking the docs first.", False)]
    assert next(d for n, d in sse_events if n == "assistant.completed")["content"] == "Done."

    done_items = [d["item"] for n, d in _events(responses) if n == "response.output_item.done"]
    assert [(i.get("phase"), i["content"][0]["text"]) for i in done_items if i["type"] == "message"] == [
        ("commentary", "Checking the docs first."), (None, "Done.")]

    # Display gate: the callback is dropped before it reaches AIAgent, like the gateway/TUI.
    _patch_api_server_runtime(monkeypatch)
    monkeypatch.setattr(
        "gateway.run._load_gateway_config", lambda: {"display": {"interim_assistant_messages": False}})
    captured = {}

    class CapturingAgent:
        provider, model = "openrouter", "global/model"

        def __init__(self, **kwargs):
            captured.update(kwargs)

    monkeypatch.setattr("run_agent.AIAgent", CapturingAgent)
    adapter._create_agent(session_id="gated", interim_assistant_callback=lambda *_a, **_k: None)
    assert captured["interim_assistant_callback"] is None
