Newer
Older
navi-1 / tests / unit / tools / test_switch_profile.py
"""Tests for switch_profile.

The tool runs in the middle of a live agent turn that holds its own Session
object, so it must repoint the session row without saving a session of its own:
a whole-session load-and-save from here made it a second writer, and the turn's
next insert died on UNIQUE(session_id, sequence_number).
"""

import pytest

from navi.core.registry import ProfileRegistry, ToolRegistry
from navi.core.session import InMemorySessionStore
from navi.profiles.base import ToolConfig, ToolScopeConfig
from navi.tools._internal.base import ToolContext, current_event_sink
from navi.tools.switch_profile import SwitchProfileTool
from tests.conftest_factory import FakeTool, make_profile, make_profile_registry


class RecordingStore(InMemorySessionStore):
    """In-memory store that records whole-session saves and profile updates."""

    def __init__(self) -> None:
        super().__init__()
        self.saves = 0
        self.profile_updates: list[tuple[str, str]] = []

    async def save(self, session) -> None:
        self.saves += 1
        await super().save(session)

    async def set_profile(self, session_id: str, profile_id: str) -> bool:
        self.profile_updates.append((session_id, profile_id))
        return await super().set_profile(session_id, profile_id)


def _tool(store: RecordingStore) -> SwitchProfileTool:
    return SwitchProfileTool(session_store=store, profile_registry=make_profile_registry())


async def _make_session(store: RecordingStore, profile_id: str = "developer"):
    session = await store.create(profile_id=profile_id, user_id="u1")
    return session


@pytest.mark.asyncio
async def test_switch_repoints_the_row_without_saving_the_session():
    store = RecordingStore()
    session = await _make_session(store)

    result = await _tool(store).execute(
        {"profile_id": "secretary"}, ctx=ToolContext(user_id="u1", session_id=session.id)
    )

    assert result.success is True
    assert store.profile_updates == [(session.id, "secretary")]
    assert store.saves == 0
    assert (await store.get(session.id)).profile_id == "secretary"


@pytest.mark.asyncio
async def test_switch_reports_the_new_profile_to_the_client():
    """The UI learns about the switch from an event, not from a session save."""
    store = RecordingStore()
    session = await _make_session(store)
    sink: list = []

    class _Sink:
        async def put(self, event):
            sink.append(event)

    result = await _tool(store).execute(
        {"profile_id": "secretary"},
        ctx=ToolContext(user_id="u1", session_id=session.id, event_sink=_Sink()),
    )

    assert result.success is True
    assert [e.profile_id for e in sink] == ["secretary"]


@pytest.mark.asyncio
async def test_switching_to_the_current_profile_is_a_noop():
    store = RecordingStore()
    session = await _make_session(store, profile_id="secretary")

    result = await _tool(store).execute(
        {"profile_id": "secretary"}, ctx=ToolContext(user_id="u1", session_id=session.id)
    )

    assert result.success is True
    assert "Already on profile" in result.output
    assert store.profile_updates == []
    assert store.saves == 0


@pytest.mark.asyncio
async def test_unknown_profile_is_rejected():
    store = RecordingStore()
    session = await _make_session(store)

    result = await _tool(store).execute(
        {"profile_id": "nope"}, ctx=ToolContext(user_id="u1", session_id=session.id)
    )

    assert result.success is False
    assert "not found" in result.error
    assert store.profile_updates == []


@pytest.mark.asyncio
async def test_unknown_session_is_rejected():
    store = RecordingStore()

    result = await _tool(store).execute(
        {"profile_id": "secretary"}, ctx=ToolContext(user_id="u1", session_id="missing")
    )

    assert result.success is False
    assert store.profile_updates == []
    assert store.saves == 0


@pytest.mark.asyncio
async def test_event_reaches_the_contextvar_sink_when_ctx_has_none():
    """The agent loop builds tool_ctx with event_sink=None (the ContextVar is the
    real channel), so `if ctx else` silently dropped the event and the profile
    badge in the header never moved."""
    store = RecordingStore()
    session = await _make_session(store)
    sink: list = []

    class _Sink:
        async def put(self, event):
            sink.append(event)

    token = current_event_sink.set(_Sink())
    try:
        result = await _tool(store).execute(
            {"profile_id": "secretary"},
            ctx=ToolContext(user_id="u1", session_id=session.id, event_sink=None),
        )
    finally:
        current_event_sink.reset(token)

    assert result.success is True
    assert [e.profile_id for e in sink] == ["secretary"]


def _registry(*names: str) -> ToolRegistry:
    reg = ToolRegistry()
    for name in names:
        reg.register(FakeTool(name))
    return reg


def _profiles(**scopes: list[str]) -> ProfileRegistry:
    reg = ProfileRegistry()
    for profile_id, tools in scopes.items():
        reg.register(make_profile(
            profile_id,
            tools=ToolConfig(agent=ToolScopeConfig(native=list(tools))),
        ))
    return reg


@pytest.mark.asyncio
async def test_output_names_the_gained_and_lost_tools(monkeypatch):
    """The model called reload_tools after leaving tool_developer because nothing
    told it what the new profile has — report the delta instead."""
    monkeypatch.setattr("navi.core.tool_utils.load_user_enabled_tools", list)
    store = RecordingStore()
    session = await _make_session(store, profile_id="developer")
    tool = SwitchProfileTool(
        session_store=store,
        profile_registry=_profiles(
            developer=["shared", "reload_tools"],
            secretary=["shared", "send_email"],
        ),
        tool_registry=_registry("shared", "reload_tools", "send_email"),
    )

    result = await tool.execute(
        {"profile_id": "secretary"}, ctx=ToolContext(user_id="u1", session_id=session.id)
    )

    assert result.success is True
    assert "Tools available now: 2" in result.output
    assert "Newly available: send_email" in result.output
    assert "No longer available: reload_tools" in result.output
    assert "next step of this turn" in result.output


@pytest.mark.asyncio
async def test_output_survives_without_a_tool_registry():
    """No registry wired (legacy/test construction) — the switch still reports."""
    store = RecordingStore()
    session = await _make_session(store)

    result = await _tool(store).execute(
        {"profile_id": "secretary"}, ctx=ToolContext(user_id="u1", session_id=session.id)
    )

    assert result.success is True
    assert "Tools available now" not in result.output