diff --git a/navi/core/pg_session_store.py b/navi/core/pg_session_store.py index 11c1631..8c5fb6d 100644 --- a/navi/core/pg_session_store.py +++ b/navi/core/pg_session_store.py @@ -465,7 +465,10 @@ async def save(self, session: Session) -> None: session.last_active = datetime.now(timezone.utc) pool = await self._get_pool() - async with pool.acquire() as conn: + # One transaction: the sequence range handed out below and the rows that + # consume it must commit together, or a crash in between burns numbers + # and leaves the counter ahead of the messages. + async with pool.acquire() as conn, conn.transaction(): # Upsert the session row. On the first save() after a lazy create() # the row does not exist yet — INSERT ... ON CONFLICT creates it so # the session_messages FK is satisfied before messages are inserted. @@ -495,7 +498,6 @@ async with self._pending_lock: self._pending.pop(session.id, None) - db_next = session.db_next_sequence messages = session.messages # 1. Update mutable flags for already-persisted rows (sequence_number >= 0) @@ -540,9 +542,31 @@ # 2. Insert new messages (sequence_number < 0 means "not yet persisted") new_msgs = [m for m in messages if m.sequence_number < 0] if new_msgs: + # Claim the range in the DB, never from session.db_next_sequence: + # another writer of the same session may have consumed those + # numbers since this object was loaded. switch_profile used to do + # exactly that — it loaded the session mid-run and saved it back — + # and the running turn then inserted onto numbers already taken, + # dying on UNIQUE(session_id, sequence_number). GREATEST() also + # heals rows whose counter is still 0 (pre-counter sessions). + base = await conn.fetchval( + """ + UPDATE sessions s + SET next_sequence = GREATEST( + s.next_sequence, + COALESCE((SELECT MAX(m.sequence_number) + 1 + FROM session_messages m + WHERE m.session_id = s.id), 0) + ) + $2 + WHERE s.id = $1 + RETURNING s.next_sequence - $2 + """, + session.id, + len(new_msgs), + ) insert_rows = [] for i, m in enumerate(new_msgs): - seq = db_next + i + seq = base + i m.sequence_number = seq insert_rows.append( ( @@ -579,12 +603,7 @@ """, insert_rows, ) - new_next = db_next + len(new_msgs) - await conn.execute( - "UPDATE sessions SET next_sequence = $1 WHERE id = $2", - new_next, session.id, - ) - session.db_next_sequence = new_next + session.db_next_sequence = base + len(new_msgs) session.db_message_count = len(messages) @@ -786,6 +805,31 @@ ) return result == "UPDATE 1" + async def set_profile(self, session_id: str, profile_id: str) -> bool: + """Repoint a session at another profile — one narrow UPDATE, no save(). + + switch_profile runs *inside* a live agent run that holds its own Session + object. Loading the session and saving it back (as it once did) makes the + tool a second writer of the same session, and the run's next insert then + lands on a sequence number this tool already took. Touching only the + profile column keeps the two out of each other's way; the run picks the + change up through its post-switch profile reload. + """ + # A never-persisted session lives only in _pending (lazy persistence): + # the bare UPDATE would touch 0 rows and silently drop the switch. + async with self._pending_lock: + pending = self._pending.get(session_id) + if pending is not None: + pending.profile_id = profile_id + return True + pool = await self._get_pool() + async with pool.acquire() as conn: + result = await conn.execute( + "UPDATE sessions SET profile_id = $1 WHERE id = $2", + profile_id, session_id, + ) + return result == "UPDATE 1" + async def list_all( self, user_id: str | None = None, is_admin: bool = False, special: bool | None = None ) -> list[Session]: diff --git a/navi/core/session.py b/navi/core/session.py index 4ad2074..c17894f 100644 --- a/navi/core/session.py +++ b/navi/core/session.py @@ -93,6 +93,14 @@ async def set_name(self, session_id: str, name: str) -> bool: ... @abstractmethod + async def set_profile(self, session_id: str, profile_id: str) -> bool: + """Repoint a session at another profile without loading or saving it. + + A whole-session load-and-save from inside a running turn (switch_profile) + would make the caller a second writer of the session's messages. + """ + + @abstractmethod async def archive_old_messages(self, session_id: str, keep_seq_threshold: int) -> int: ... @abstractmethod @@ -231,6 +239,13 @@ s.name = name return True + async def set_profile(self, session_id: str, profile_id: str) -> bool: + s = self._sessions.get(session_id) + if s is None: + return False + s.profile_id = profile_id + return True + async def archive_old_messages(self, session_id: str, keep_seq_threshold: int) -> int: # In-memory store: no-op, everything stays in RAM return 0 diff --git a/navi/tools/switch_profile.py b/navi/tools/switch_profile.py index 65666aa..397543b 100644 --- a/navi/tools/switch_profile.py +++ b/navi/tools/switch_profile.py @@ -70,8 +70,12 @@ output=f"Already on profile '{profile.name}' — no change.", ) - session.profile_id = profile_id - await self._sessions.save(session) + # Narrow UPDATE — deliberately not save(). This tool runs in the middle of + # an agent turn that holds its own Session object; saving the copy we just + # loaded would write to the session behind that turn's back and hand out + # sequence numbers it is still claiming. The run picks the new profile up + # after the turn (agent.py: profile_reloaded). + await self._sessions.set_profile(sid, profile_id) # Notify the client immediately so it can update the UI. sink = ctx.event_sink if ctx else current_event_sink.get() diff --git a/tests/conftest_factory.py b/tests/conftest_factory.py index 397a01e..b47b5af 100644 --- a/tests/conftest_factory.py +++ b/tests/conftest_factory.py @@ -229,6 +229,17 @@ async def executemany(self, query: str, args_list: list) -> None: self.calls.append(("executemany", query, args_list)) + def transaction(self): + """Stand-in for asyncpg's ``conn.transaction()`` async context manager.""" + class _Tx: + async def __aenter__(_self): + return _self + + async def __aexit__(_self, *exc): + return False + + return _Tx() + async def __aenter__(self): return self diff --git a/tests/unit/core/test_pg_session_store.py b/tests/unit/core/test_pg_session_store.py index a71ae61..e2025b4 100644 --- a/tests/unit/core/test_pg_session_store.py +++ b/tests/unit/core/test_pg_session_store.py @@ -12,9 +12,7 @@ async def test_save_assigns_sequence_numbers_to_new_messages(): conn = FakeConnection() conn.enqueue("OK") # upsert sessions row (INSERT ... ON CONFLICT) - conn.enqueue(None) # executemany UPDATE (empty) - conn.enqueue("OK") # INSERT executemany - conn.enqueue("OK") # UPDATE next_sequence + conn.enqueue(0) # sequence range claimed in the DB: RETURNING base pool = FakePool(conn) store = PgSessionStore(pool) @@ -63,6 +61,73 @@ @pytest.mark.asyncio +async def test_save_numbers_from_the_range_the_db_returns_not_the_local_snapshot(): + """Regression: a second writer of the same session (switch_profile mid-run) + already consumed the numbers this object still believes are free. + + Numbering from session.db_next_sequence put the message back on 94, which was + taken, and the turn died on UNIQUE(session_id, sequence_number). + """ + conn = FakeConnection() + conn.enqueue("OK") # upsert sessions row + conn.enqueue(95) # another writer took 94 → the DB hands back 95 + store = PgSessionStore(FakePool(conn)) + store._initialized = True + + session = Session(profile_id="test") + session.db_next_sequence = 94 # stale snapshot — 94 is no longer free + tool_result = Message(role="tool", content="the real result") + session.messages = [tool_result] + + await store.save(session) + + assert tool_result.sequence_number == 95 + assert session.db_next_sequence == 96 + # The range is claimed before the rows that consume it, in one transaction. + kinds = [c[0] for c in conn.calls] + assert kinds.index("fetchval") < kinds.index("executemany") + allocation = next(c for c in conn.calls if c[0] == "fetchval") + assert "RETURNING" in allocation[1] + # GREATEST() heals the pre-counter sessions whose next_sequence is still 0. + assert "GREATEST" in allocation[1] + # ...and the old read-modify-write of the counter is gone. + assert not any("SET next_sequence = $1" in c[1] for c in conn.calls) + + +@pytest.mark.asyncio +async def test_set_profile_is_one_narrow_update_without_a_session_load(): + store, conn = _make_store() + conn.enqueue("UPDATE 1") + + assert await store.set_profile("sid-1", "secretary") is True + + assert [c[0] for c in conn.calls] == ["execute"] + _, query, args = conn.calls[0] + assert query == "UPDATE sessions SET profile_id = $1 WHERE id = $2" + assert args == ("secretary", "sid-1") + + +@pytest.mark.asyncio +async def test_set_profile_missing_row_returns_false(): + store, conn = _make_store() + conn.enqueue("UPDATE 0") + + assert await store.set_profile("nope", "secretary") is False + + +@pytest.mark.asyncio +async def test_set_profile_moves_a_pending_session_in_memory(): + """A never-persisted session has no row to update — the switch is kept in memory.""" + store, conn = _make_store() + session = await store.create(profile_id="developer", user_id="u1") + + assert await store.set_profile(session.id, "secretary") is True + + assert session.profile_id == "secretary" + assert conn.calls == [] # no round-trip for a session with no row yet + + +@pytest.mark.asyncio async def test_archive_old_messages_sends_correct_sql(): conn = FakeConnection() conn.enqueue("INSERT 0 3") # copy to archive @@ -133,8 +198,7 @@ """First save() upserts the sessions row and removes it from _pending.""" conn = FakeConnection() conn.enqueue("INSERT 0 1") # upsert sessions row - conn.enqueue("OK") # INSERT executemany (one new message) - conn.enqueue("OK") # UPDATE next_sequence + conn.enqueue(0) # sequence range claimed in the DB (one new message) pool = FakePool(conn) store = PgSessionStore(pool) store._initialized = True diff --git a/tests/unit/tools/test_switch_profile.py b/tests/unit/tools/test_switch_profile.py new file mode 100644 index 0000000..4c740fe --- /dev/null +++ b/tests/unit/tools/test_switch_profile.py @@ -0,0 +1,117 @@ +"""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.session import InMemorySessionStore +from navi.tools._internal.base import ToolContext +from navi.tools.switch_profile import SwitchProfileTool +from tests.conftest_factory import 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