Newer
Older
navi-1 / tests / unit / api / test_websocket.py
@Eugene Sukhodolskiy Eugene Sukhodolskiy 4 hours ago 26 KB background tasks: review fixes B1-B11 batch
"""Unit tests for WebSocket handler internals and reconnect logic."""

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

import pytest
from fastapi import WebSocketDisconnect

from navi.api import websocket as ws_mod
from navi.core.orchestrator import AgentSessionOrchestrator, SessionRun


@pytest.fixture(autouse=True)
def _clear_state(monkeypatch):
    """Clear global state before every WS test."""
    yield


@pytest.fixture
def mock_websocket():
    ws = AsyncMock()
    ws.accept = AsyncMock()
    ws.close = AsyncMock()
    ws.send_json = AsyncMock()
    return ws


@pytest.fixture
def mock_session():
    session = MagicMock()
    # user_id=None = legacy/local-owned session (deps.settings may carry
    # env-derived NAVI_AUTH_ENABLED=false in unit runs; NULL owner passes).
    session.user_id = None
    return session


@pytest.fixture
def mock_user():
    user = MagicMock()
    user.id = "test-user-id"
    user.role = "admin"
    return user


@pytest.fixture
def fake_orchestrator():
    container = MagicMock()
    container.profile_registry = None
    container.tool_registry = None
    container.backend_registry = None
    container.cp_registry = None
    container.workers = []
    container.memory_store = None
    container.mcp_manager = None
    return AgentSessionOrchestrator(container)


# ── SessionRun buffer tests ─────────────────────────────────────────────────

@pytest.mark.anyio
async def test_event_buffer_appended_and_replayed():
    """Broadcast stores serialised events; oldest evicted when cap exceeded."""
    run = SessionRun()

    class FakeEvent:
        def __init__(self, idx: int) -> None:
            self.idx = idx

        def to_wire(self) -> dict:
            return {"type": "stream_delta", "delta": str(self.idx)}

    for i in range(3):
        await run.broadcast(("event", FakeEvent(i)))

    assert len(run.events) == 3
    assert run.events[0] == {"type": "stream_delta", "delta": "0"}
    assert run.events[2] == {"type": "stream_delta", "delta": "2"}

    # Fill buffer past limit
    run.events.clear()
    for i in range(500 + 5):
        await run.broadcast(("event", FakeEvent(i)))

    assert len(run.events) == 500
    assert run.events[0] == {"type": "stream_delta", "delta": "5"}
    assert run.events[-1] == {"type": "stream_delta", "delta": str(500 + 4)}


# ── Reconnect / replay tests ─────────────────────────────────────────────────

@pytest.mark.anyio
async def test_reconnect_replays_buffered_events(mock_websocket, mock_session, mock_user, monkeypatch):
    """Re-attach to active run yields replay_start, buffered events, replay_end, then session_sync."""
    monkeypatch.setattr(ws_mod, "get_current_user_ws", AsyncMock(return_value=mock_user))
    mock_store = MagicMock()
    mock_store.get = AsyncMock(return_value=mock_session)
    monkeypatch.setattr(ws_mod, "get_session_store", lambda: mock_store)
    monkeypatch.setattr(ws_mod, "_stream_to_client", AsyncMock(return_value=True))

    fake_container = MagicMock()
    fake_container.profile_registry = None
    fake_container.tool_registry = None
    fake_container.backend_registry = None
    fake_container.cp_registry = None
    fake_container.orchestrator = AgentSessionOrchestrator(fake_container)
    monkeypatch.setattr("navi.api.deps._resolve_container", lambda: fake_container)

    orchestrator = fake_container.orchestrator
    run = orchestrator.create_run("s1")
    run.events = [
        {"type": "stream_delta", "delta": "hello"},
        {"type": "thinking_delta", "delta": "hmm"},
    ]

    mock_websocket.receive_text = AsyncMock(side_effect=WebSocketDisconnect())

    await ws_mod.websocket_session("s1", mock_websocket)

    calls = [c.args[0] for c in mock_websocket.send_json.call_args_list]
    types = [c["type"] for c in calls]

    assert types == [
        "stream_start",
        "replay_start",
        "stream_delta",
        "thinking_delta",
        "replay_end",
        "session_sync",
    ]
    assert calls[1]["count"] == 2

    orchestrator._sessions.pop("s1", None)


@pytest.mark.anyio
async def test_session_sync_after_reconnect_when_done(mock_websocket, mock_session, mock_user, monkeypatch):
    """Reconnect when no run is active → only session_sync, no replay."""
    monkeypatch.setattr(ws_mod, "get_current_user_ws", AsyncMock(return_value=mock_user))
    mock_store = MagicMock()
    mock_store.get = AsyncMock(return_value=mock_session)
    monkeypatch.setattr(ws_mod, "get_session_store", lambda: mock_store)
    monkeypatch.setattr(ws_mod, "_stream_to_client", AsyncMock(return_value=True))

    fake_container = MagicMock()
    fake_container.profile_registry = None
    fake_container.tool_registry = None
    fake_container.backend_registry = None
    fake_container.cp_registry = None
    fake_container.orchestrator = AgentSessionOrchestrator(fake_container)
    monkeypatch.setattr("navi.api.deps._resolve_container", lambda: fake_container)

    mock_websocket.receive_text = AsyncMock(side_effect=WebSocketDisconnect())

    await ws_mod.websocket_session("s1", mock_websocket)

    calls = [c.args[0] for c in mock_websocket.send_json.call_args_list]
    assert len(calls) == 1
    assert calls[0]["type"] == "session_sync"


@pytest.mark.anyio
async def test_session_sync_after_recall_run(mock_websocket, mock_session, mock_user, monkeypatch):
    """Reconnect while a headless recall run is active → session_sync (no replay, no error)."""
    monkeypatch.setattr(ws_mod, "get_current_user_ws", AsyncMock(return_value=mock_user))
    mock_store = MagicMock()
    mock_store.get = AsyncMock(return_value=mock_session)
    monkeypatch.setattr(ws_mod, "get_session_store", lambda: mock_store)
    monkeypatch.setattr(ws_mod, "_stream_to_client", AsyncMock(return_value=True))

    fake_container = MagicMock()
    fake_container.profile_registry = None
    fake_container.tool_registry = None
    fake_container.backend_registry = None
    fake_container.cp_registry = None
    fake_container.orchestrator = AgentSessionOrchestrator(fake_container)
    monkeypatch.setattr("navi.api.deps._resolve_container", lambda: fake_container)

    orchestrator = fake_container.orchestrator
    orchestrator.mark_busy("s1")

    mock_websocket.receive_text = AsyncMock(side_effect=WebSocketDisconnect())

    await ws_mod.websocket_session("s1", mock_websocket)

    calls = [c.args[0] for c in mock_websocket.send_json.call_args_list]
    types = [c["type"] for c in calls]
    assert types == ["session_sync"]

    await orchestrator.clear_busy("s1")


# ── Concurrent run guard ─────────────────────────────────────────────────────

@pytest.mark.anyio
async def test_concurrent_run_guard_queues_second_message(mock_websocket, mock_session, mock_user, monkeypatch):
    """Sending a second message while a run is active queues it (message_queued)."""
    monkeypatch.setattr(ws_mod, "get_current_user_ws", AsyncMock(return_value=mock_user))
    mock_store = MagicMock()
    mock_store.get = AsyncMock(return_value=mock_session)
    monkeypatch.setattr(ws_mod, "get_session_store", lambda: mock_store)
    monkeypatch.setattr(ws_mod, "_stream_to_client", AsyncMock(return_value=True))

    fake_container = MagicMock()
    fake_container.profile_registry = None
    fake_container.tool_registry = None
    fake_container.backend_registry = None
    fake_container.cp_registry = None
    fake_container.orchestrator = AgentSessionOrchestrator(fake_container)
    monkeypatch.setattr("navi.api.deps._resolve_container", lambda: fake_container)

    orchestrator = fake_container.orchestrator

    # _run_agent sleeps so the run stays registered
    async def fake_run_agent(*a, **kw):
        await asyncio.sleep(3600)

    monkeypatch.setattr(orchestrator, "run_agent", fake_run_agent)

    message_count = 0

    async def fake_receive_text():
        nonlocal message_count
        message_count += 1
        if message_count == 1:
            return json.dumps({"type": "message", "content": "first"})
        if message_count == 2:
            return json.dumps({"type": "message", "content": "second"})
        raise WebSocketDisconnect()

    mock_websocket.receive_text = fake_receive_text

    await ws_mod.websocket_session("s1", mock_websocket)

    calls = [c.args[0] for c in mock_websocket.send_json.call_args_list]
    queued_calls = [c for c in calls if c["type"] == "message_queued"]
    assert len(queued_calls) == 1
    assert queued_calls[0]["position"] == 1
    assert queued_calls[0]["queue_len"] == 1
    # No error was sent for the queued message
    assert not [c for c in calls if c["type"] == "error"]
    # The message is retained in the session queue
    assert orchestrator.has_pending("s1")

    # Cleanup background task
    state = orchestrator._sessions.get("s1")
    if state and state.run and state.run.task:
        state.run.task.cancel()
        try:
            await state.run.task
        except asyncio.CancelledError:
            pass
    orchestrator._sessions.pop("s1", None)


# ── /compact (forced compression) ───────────────────────────────────────────


@pytest.mark.anyio
async def test_compact_message_runs_forced_compression(mock_websocket, mock_session, mock_user, monkeypatch):
    """A {"type": "compact"} message dispatches to run_compact (not run_agent),
    and sends NO stream_start — the client distinguishes a forced compact from an
    in-turn auto-compress by the absence of a streaming turn."""
    monkeypatch.setattr(ws_mod, "get_current_user_ws", AsyncMock(return_value=mock_user))
    mock_store = MagicMock()
    mock_store.get = AsyncMock(return_value=mock_session)
    monkeypatch.setattr(ws_mod, "get_session_store", lambda: mock_store)
    monkeypatch.setattr(ws_mod, "_stream_to_client", AsyncMock(return_value=True))

    fake_container = MagicMock()
    fake_container.profile_registry = None
    fake_container.tool_registry = None
    fake_container.backend_registry = None
    fake_container.cp_registry = None
    fake_container.orchestrator = AgentSessionOrchestrator(fake_container)
    monkeypatch.setattr("navi.api.deps._resolve_container", lambda: fake_container)

    orchestrator = fake_container.orchestrator
    compact_calls: list = []

    async def fake_run_compact(sid, store):
        compact_calls.append((sid, store))

    monkeypatch.setattr(orchestrator, "run_compact", fake_run_compact)

    received: list[str] = []

    async def fake_receive_text():
        if not received:
            received.append(json.dumps({"type": "compact"}))
            return received[-1]
        raise WebSocketDisconnect()

    mock_websocket.receive_text = fake_receive_text

    await ws_mod.websocket_session("s1", mock_websocket)

    # The handler schedules run_compact via create_task but the mocked
    # _stream_to_client / receive_text never yield to it, so drive it here.
    state = orchestrator._sessions.get("s1")
    if state and state.run and state.run.task:
        try:
            await state.run.task
        except asyncio.CancelledError:
            pass

    assert len(compact_calls) == 1
    assert compact_calls[0][0] == "s1"
    # No stream_start is sent for a compact.
    calls = [c.args[0] for c in mock_websocket.send_json.call_args_list]
    assert all(c.get("type") != "stream_start" for c in calls)

    orchestrator._sessions.pop("s1", None)


@pytest.mark.anyio
async def test_compact_while_agent_running_is_rejected(mock_websocket, mock_session, mock_user, monkeypatch):
    """A /compact while an agent turn is active is rejected — compressing
    session.context mid-turn would race the running agent."""
    monkeypatch.setattr(ws_mod, "get_current_user_ws", AsyncMock(return_value=mock_user))
    mock_store = MagicMock()
    mock_store.get = AsyncMock(return_value=mock_session)
    monkeypatch.setattr(ws_mod, "get_session_store", lambda: mock_store)
    monkeypatch.setattr(ws_mod, "_stream_to_client", AsyncMock(return_value=True))

    fake_container = MagicMock()
    fake_container.profile_registry = None
    fake_container.tool_registry = None
    fake_container.backend_registry = None
    fake_container.cp_registry = None
    fake_container.orchestrator = AgentSessionOrchestrator(fake_container)
    monkeypatch.setattr("navi.api.deps._resolve_container", lambda: fake_container)

    orchestrator = fake_container.orchestrator

    async def fake_run_agent(*a, **kw):
        await asyncio.sleep(3600)

    monkeypatch.setattr(orchestrator, "run_agent", fake_run_agent)

    message_count = 0

    async def fake_receive_text():
        nonlocal message_count
        message_count += 1
        if message_count == 1:
            return json.dumps({"type": "message", "content": "first"})
        if message_count == 2:
            return json.dumps({"type": "compact"})
        raise WebSocketDisconnect()

    mock_websocket.receive_text = fake_receive_text

    await ws_mod.websocket_session("s1", mock_websocket)

    calls = [c.args[0] for c in mock_websocket.send_json.call_args_list]
    error_calls = [c for c in calls if c["type"] == "error"]
    assert len(error_calls) == 1
    assert "wait for it to finish" in error_calls[0]["message"]

    state = orchestrator._sessions.get("s1")
    if state and state.run and state.run.task:
        state.run.task.cancel()
        try:
            await state.run.task
        except asyncio.CancelledError:
            pass
    orchestrator._sessions.pop("s1", None)


# ── run_compact event/error broadcasting ────────────────────────────────────


@pytest.mark.anyio
async def test_run_compact_broadcasts_events_and_done(fake_orchestrator, monkeypatch):
    """run_compact forwards agent.compact_stream events to subscribers and ends
    with a done marker so _stream_to_client returns."""
    from navi.core.events import CompressionStarted, ContextCompressed

    async def fake_compact_stream(session_id):
        yield CompressionStarted(context_tokens=100, max_context_tokens=4096)
        yield ContextCompressed(messages_before=20, messages_after=5, context_tokens=80, max_context_tokens=4096)

    class _FakeAgent:
        async def compact_stream(self, session_id):
            async for ev in fake_compact_stream(session_id):
                yield ev

    def fake_build_agent(self, _store):
        return _FakeAgent()

    monkeypatch.setattr(AgentSessionOrchestrator, "_build_agent", fake_build_agent)

    run = fake_orchestrator.create_run("s1")
    queue = run.subscribe()
    await fake_orchestrator.run_compact("s1", session_store=MagicMock())

    items = []
    while not queue.empty():
        items.append(queue.get_nowait())
    kinds = [item[0] for item in items]
    assert kinds == ["event", "event", "done"]
    assert isinstance(items[0][1], CompressionStarted)
    assert isinstance(items[1][1], ContextCompressed)
    # run_compact's finally clears state.run, and with no websockets attached
    # _cleanup then removes the session entry entirely.
    assert "s1" not in fake_orchestrator._sessions


@pytest.mark.anyio
async def test_run_compact_surfaces_nothing_to_compact_as_error(fake_orchestrator, monkeypatch):
    """NothingToCompactError is broadcast as an error event (user feedback),
    followed by done — not raised into the WS handler."""
    from navi.exceptions import NothingToCompactError

    class _FakeAgent:
        async def compact_stream(self, session_id):
            raise NothingToCompactError("nothing")
            yield  # pragma: no cover - makes this an async generator

    def fake_build_agent(self, _store):
        return _FakeAgent()

    monkeypatch.setattr(AgentSessionOrchestrator, "_build_agent", fake_build_agent)

    run = fake_orchestrator.create_run("s1")
    queue = run.subscribe()
    await fake_orchestrator.run_compact("s1", session_store=MagicMock())

    items = []
    while not queue.empty():
        items.append(queue.get_nowait())
    kinds = [item[0] for item in items]
    assert kinds == ["error", "done"]
    assert "nothing" in items[0][1]
    assert "s1" not in fake_orchestrator._sessions


# ── Anonymous / legacy-session access control ────────────────────────────────


def _auth_settings(enabled: bool):
    from navi.config import Settings

    return Settings(_env_file=None, navi_persona_file="", navi_auth_enabled=enabled)


def _wire(monkeypatch, user, session, orchestrator):
    """Wire the WS handler to the given user/session/orchestrator."""
    monkeypatch.setattr(ws_mod, "get_current_user_ws", AsyncMock(return_value=user))
    mock_store = MagicMock()
    mock_store.get = AsyncMock(return_value=session)
    monkeypatch.setattr(ws_mod, "get_session_store", lambda: mock_store)
    monkeypatch.setattr(ws_mod, "_stream_to_client", AsyncMock(return_value=True))

    fake_container = MagicMock()
    fake_container.profile_registry = None
    fake_container.tool_registry = None
    fake_container.backend_registry = None
    fake_container.cp_registry = None
    fake_container.orchestrator = orchestrator
    monkeypatch.setattr("navi.api.deps._resolve_container", lambda: fake_container)


@pytest.mark.anyio
async def test_anonymous_rejected_when_auth_enabled(mock_websocket, mock_session, monkeypatch):
    """With auth on, anonymous clients are rejected — even for legacy (ownerless)
    sessions, which would otherwise expose the full unrestricted tool surface."""
    monkeypatch.setattr(ws_mod, "settings", _auth_settings(True))
    mock_session.user_id = None  # legacy session
    orchestrator = AgentSessionOrchestrator(MagicMock())
    _wire(monkeypatch, None, mock_session, orchestrator)

    mock_websocket.receive_text = AsyncMock(side_effect=WebSocketDisconnect())

    await ws_mod.websocket_session("s1", mock_websocket)

    mock_websocket.close.assert_awaited_with(code=4003, reason="Authentication required")
    # The socket must NOT be registered with the orchestrator after a denial.
    assert "s1" not in orchestrator._sessions


@pytest.mark.anyio
async def test_anonymous_legacy_session_allowed_when_auth_disabled(mock_websocket, monkeypatch):
    """Without auth (trusted local mode), legacy ownerless sessions stay reachable."""
    monkeypatch.setattr(ws_mod, "settings", _auth_settings(False))
    session = MagicMock()
    session.user_id = None
    orchestrator = AgentSessionOrchestrator(MagicMock())
    _wire(monkeypatch, None, session, orchestrator)

    mock_websocket.receive_text = AsyncMock(side_effect=WebSocketDisconnect())

    await ws_mod.websocket_session("s1", mock_websocket)

    mock_websocket.close.assert_not_awaited()
    calls = [c.args[0] for c in mock_websocket.send_json.call_args_list]
    assert calls[0]["type"] == "session_sync"


@pytest.mark.anyio
async def test_anonymous_owned_session_denied_when_auth_disabled(mock_websocket, mock_session, monkeypatch):
    """Even without auth, an anonymous client may not attach to someone else's session."""
    monkeypatch.setattr(ws_mod, "settings", _auth_settings(False))
    mock_session.user_id = "someone-else"
    orchestrator = AgentSessionOrchestrator(MagicMock())
    _wire(monkeypatch, None, mock_session, orchestrator)

    mock_websocket.receive_text = AsyncMock(side_effect=WebSocketDisconnect())

    await ws_mod.websocket_session("s1", mock_websocket)

    mock_websocket.close.assert_awaited_with(code=4003, reason="Authentication required")
    assert "s1" not in orchestrator._sessions


@pytest.mark.anyio
async def test_non_owner_denied_leaves_no_socket(mock_websocket, monkeypatch):
    """A non-owner authenticated user is denied and no socket lingers behind."""
    from navi.auth import User

    monkeypatch.setattr(ws_mod, "settings", _auth_settings(True))
    # check_session_access reads navi.auth.deps.settings directly
    import navi.auth.deps as auth_deps
    monkeypatch.setattr(auth_deps, "settings", _auth_settings(True))
    session = MagicMock()
    session.user_id = "owner-1"
    orchestrator = AgentSessionOrchestrator(MagicMock())
    _wire(monkeypatch, User(id="intruder", email="x@test.com", role="user"), session, orchestrator)

    mock_websocket.receive_text = AsyncMock(side_effect=WebSocketDisconnect())

    await ws_mod.websocket_session("s1", mock_websocket)

    mock_websocket.close.assert_awaited_with(code=4003, reason="Access denied")
    assert "s1" not in orchestrator._sessions


# ── Message queue (busy → message_queued → drain) ───────────────────────────


@pytest.mark.anyio
async def test_drain_runs_queued_messages_after_run(mock_websocket, mock_user, monkeypatch):
    """After a run finishes on the socket, queued messages execute back-to-back."""
    orchestrator = AgentSessionOrchestrator(MagicMock())
    ran = []

    async def fake_single(**kwargs):
        ran.append(kwargs["user_content"])
        return True, True

    monkeypatch.setattr(ws_mod, "_run_single_message", AsyncMock(side_effect=fake_single))
    orchestrator.queue_message("s1", {
        "user_content": "queued one",
        "raw_images": ["img"],
        "display_content": "queued one",
        "uploaded_files": [],
        "hidden": False,
        "user": mock_user,
        "cwd": "/tmp",
    })

    connected = await ws_mod._start_agent_run(
        session_id="s1", user_content="original", raw_images=None,
        display_content="original", uploaded_files=[], hidden=False,
        websocket=mock_websocket, orchestrator=orchestrator,
        session_store=MagicMock(), user=mock_user,
    )

    assert connected is True
    assert ran == ["original", "queued one"]
    # drained entry preserved identity and fields
    second = ws_mod._run_single_message.call_args_list[1].kwargs
    assert second["raw_images"] == ["img"]
    assert second["user"] is mock_user
    assert second["cwd"] == "/tmp"
    assert not orchestrator.has_pending("s1")


@pytest.mark.anyio
async def test_drain_stops_when_socket_dies(mock_websocket, mock_user, monkeypatch):
    """A disconnect mid-drain leaves the remaining queue intact."""
    orchestrator = AgentSessionOrchestrator(MagicMock())
    ran = []

    async def fake_single(**kwargs):
        ran.append(kwargs["user_content"])
        # the second drained message "disconnects" the socket
        return (False, True) if kwargs["user_content"] == "second" else (True, True)

    monkeypatch.setattr(ws_mod, "_run_single_message", AsyncMock(side_effect=fake_single))
    orchestrator.queue_message("s1", {"user_content": "second", "uploaded_files": []})
    orchestrator.queue_message("s1", {"user_content": "third", "uploaded_files": []})

    connected = await ws_mod._start_agent_run(
        session_id="s1", user_content="first", raw_images=None,
        display_content="first", uploaded_files=[], hidden=False,
        websocket=mock_websocket, orchestrator=orchestrator,
        session_store=MagicMock(), user=mock_user,
    )

    assert connected is False
    assert ran == ["first", "second"]
    assert orchestrator.pop_pending("s1")["user_content"] == "third"


@pytest.mark.anyio
async def test_queued_entry_keeps_user_identity(mock_websocket, monkeypatch):
    """Messages queued from another socket run with the queueing user."""
    from navi.auth import User

    orchestrator = AgentSessionOrchestrator(MagicMock())
    other = User(id="user-b", email="b@test.com", role="user")
    ws_a = AsyncMock()
    ws_a.send_json = AsyncMock()

    # Simulate socket B queueing while a run is "active"
    orchestrator.create_run("s1")
    connected, ran = await ws_mod._run_single_message(
        session_id="s1", user_content="from b", raw_images=None,
        display_content="from b", uploaded_files=[], hidden=False,
        websocket=ws_a, orchestrator=orchestrator,
        session_store=MagicMock(), user=other,
    )
    assert connected is True
    sent = ws_a.send_json.call_args_list[0].args[0]
    assert sent["type"] == "message_queued"

    entry = orchestrator.pop_pending("s1")
    assert entry["user"] is other
    assert entry["user_content"] == "from b"
    state = orchestrator._sessions.get("s1")
    if state and state.run:
        state.run = None
    orchestrator._sessions.pop("s1", None)


# ── Two-tab visibility (B5) and stranded queue (B4) ──────────────────────────


@pytest.mark.anyio
async def test_run_marks_owner_socket(mock_websocket, mock_user, monkeypatch):
    """The run-owning socket is recorded so teardown excludes it from rebroadcasts."""
    orchestrator = AgentSessionOrchestrator(MagicMock())
    monkeypatch.setattr(orchestrator, "run_agent", AsyncMock())
    monkeypatch.setattr(ws_mod, "_stream_to_client", AsyncMock(return_value=True))

    connected, ran = await ws_mod._run_single_message(
        session_id="s1", user_content="hi", raw_images=None,
        display_content=None, uploaded_files=[], hidden=False,
        websocket=mock_websocket, orchestrator=orchestrator,
        session_store=MagicMock(), user=mock_user,
    )
    assert connected is True and ran is True
    run = orchestrator.get_run("s1")
    assert run is not None
    assert run.owner_socket is mock_websocket


@pytest.mark.anyio
async def test_broadcast_session_sync_reaches_other_tabs_only(mock_websocket, mock_user, monkeypatch):
    """Watchers reload via session_sync; the excluded (owner) socket does not."""
    orchestrator = AgentSessionOrchestrator(MagicMock())
    other = AsyncMock()
    orchestrator.add_websocket("s1", mock_websocket)
    orchestrator.add_websocket("s1", other)

    orchestrator.broadcast_session_sync("s1", session=None, exclude=mock_websocket)
    await asyncio.sleep(0.02)

    other.send_json.assert_awaited_once()
    payload = other.send_json.call_args.args[0]
    assert payload["type"] == "session_sync"
    assert payload["session_id"] == "s1"
    mock_websocket.send_json.assert_not_awaited()


@pytest.mark.anyio
async def test_headless_drain_broadcasts_session_sync(mock_websocket, monkeypatch):
    """After the headless drain, remaining watchers reload the outcome (B4 tail)."""
    orchestrator = AgentSessionOrchestrator(MagicMock())

    async def fake_run_agent(session_id, content, *args, **kwargs):
        state = orchestrator._sessions.get(session_id)
        if state is not None:
            state.run = None

    monkeypatch.setattr(orchestrator, "run_agent", fake_run_agent)
    orchestrator.queue_message("s1", {"user_content": "queued", "uploaded_files": []})
    orchestrator.add_websocket("s1", mock_websocket)

    await orchestrator._drain_queued("s1", None)
    await asyncio.sleep(0.02)  # let the fire-and-forget broadcast run

    assert not orchestrator.has_pending("s1")
    mock_websocket.send_json.assert_awaited()
    payload = mock_websocket.send_json.call_args_list[-1].args[0]
    assert payload["type"] == "session_sync"