diff --git a/src/crud/session.py b/src/crud/session.py index efa3f664..1cb0a06a 100644 --- a/src/crud/session.py +++ b/src/crud/session.py @@ -1084,20 +1084,21 @@ async def set_peers_for_session( f"Session {session_name} not found in workspace {workspace_name}" ) - # Soft delete every *ordinary* active membership. Scope memberships are - # deliberately preserved: this route replaces the peers the caller names, and a - # caller detaches a scope by simply *omitting* it from an otherwise valid - # replacement map — never naming it, so no request-level guard can see it. - # Without the exclusion a plain replacement silently bypasses the facade that - # owns scope membership and its removal reconciliation. Being part of the - # UPDATE, this holds regardless of the request body or concurrent scope - # creation. + # Soft delete every *ordinary* active membership not in the incoming map. + # Scope memberships are deliberately preserved: this route replaces the peers + # the caller names, and a caller detaches a scope by simply *omitting* it from + # an otherwise valid replacement map — never naming it, so no request-level + # guard can see it. Without the exclusion a plain replacement silently + # bypasses the facade that owns scope membership and its removal + # reconciliation. Being part of the UPDATE, this holds regardless of the + # request body or concurrent scope creation. update_stmt = ( update(models.SessionPeer) .where( models.SessionPeer.session_name == session_name, models.SessionPeer.workspace_name == workspace_name, models.SessionPeer.left_at.is_(None), # Only update active peers + models.SessionPeer.peer_name.notin_(peer_names.keys()), ~exists( select(models.Peer.id) .where(models.Peer.workspace_name == workspace_name) @@ -1254,13 +1255,14 @@ async def _get_or_add_peers_to_session( ] ) - # On conflict, update joined_at and clear left_at (rejoin scenario) - # If left_at is not None (peer has left the session): Use the new configuration (stmt.excluded.configuration) - # If left_at is None (peer is still active): Keep the existing configuration (models.SessionPeer.configuration) + # On conflict, rejoin departed peers and leave active memberships unchanged. stmt = stmt.on_conflict_do_update( index_elements=["session_name", "peer_name", "workspace_name"], set_={ - "joined_at": func.now(), + "joined_at": case( + (models.SessionPeer.left_at.is_not(None), func.now()), + else_=models.SessionPeer.joined_at, + ), "left_at": None, "configuration": case( (models.SessionPeer.left_at.is_not(None), stmt.excluded.configuration), diff --git a/tests/crud/test_session.py b/tests/crud/test_session.py index 8b4b5795..8427622a 100644 --- a/tests/crud/test_session.py +++ b/tests/crud/test_session.py @@ -1,5 +1,8 @@ +from datetime import datetime, timezone + import pytest from nanoid import generate as generate_nanoid +from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from src import crud, models, schemas @@ -9,6 +12,179 @@ from src.exceptions import ResourceNotFoundException class TestSessionCRUD: """Test suite for session CRUD operations""" + @pytest.mark.asyncio + async def test_get_or_create_session_preserves_active_joined_at( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Active re-adds keep joined_at and config; a genuine rejoin starts a new window.""" + test_workspace, test_peer = sample_data + session_name = str(generate_nanoid()) + original_config = schemas.SessionPeerConfig( + observe_others=True, observe_me=False + ) + updated_config = schemas.SessionPeerConfig( + observe_others=False, observe_me=True + ) + session_peer_stmt = select( + models.SessionPeer.joined_at, + models.SessionPeer.left_at, + models.SessionPeer.configuration, + ).where( + models.SessionPeer.session_name == session_name, + models.SessionPeer.peer_name == test_peer.name, + models.SessionPeer.workspace_name == test_workspace.name, + ) + + await crud.get_or_create_session( + db_session, + schemas.SessionCreate( + name=session_name, peers={test_peer.name: original_config} + ), + test_workspace.name, + ) + first_joined_at, first_left_at, first_config = ( + await db_session.execute(session_peer_stmt) + ).one() + assert first_left_at is None + assert first_config == original_config.model_dump() + + await crud.get_or_create_session( + db_session, + schemas.SessionCreate( + name=session_name, peers={test_peer.name: updated_config} + ), + test_workspace.name, + ) + second_joined_at, second_left_at, second_config = ( + await db_session.execute(session_peer_stmt) + ).one() + assert second_joined_at == first_joined_at + assert second_left_at is None + assert second_config == original_config.model_dump() + + session_peer = ( + await db_session.execute( + select(models.SessionPeer).where( + models.SessionPeer.session_name == session_name, + models.SessionPeer.peer_name == test_peer.name, + models.SessionPeer.workspace_name == test_workspace.name, + ) + ) + ).scalar_one() + session_peer.left_at = datetime.now(timezone.utc) + await db_session.commit() + + await crud.get_or_create_session( + db_session, + schemas.SessionCreate( + name=session_name, peers={test_peer.name: updated_config} + ), + test_workspace.name, + ) + rejoined_joined_at, rejoined_left_at, rejoined_config = ( + await db_session.execute(session_peer_stmt) + ).one() + assert rejoined_joined_at > second_joined_at + assert rejoined_left_at is None + assert rejoined_config == updated_config.model_dump() + + @pytest.mark.asyncio + async def test_set_peers_preserves_active_joined_at( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """PUT-peers keeps active membership windows and refreshes real rejoins.""" + test_workspace, test_peer = sample_data + session_name = str(generate_nanoid()) + original_config = schemas.SessionPeerConfig( + observe_others=True, observe_me=False + ) + updated_config = schemas.SessionPeerConfig( + observe_others=False, observe_me=True + ) + db_session.add( + models.Session(name=session_name, workspace_name=test_workspace.name) + ) + await db_session.flush() + + session_peer_stmt = select( + models.SessionPeer.joined_at, + models.SessionPeer.left_at, + models.SessionPeer.configuration, + ).where( + models.SessionPeer.session_name == session_name, + models.SessionPeer.peer_name == test_peer.name, + models.SessionPeer.workspace_name == test_workspace.name, + ) + + await crud.set_peers_for_session( + db_session, + workspace_name=test_workspace.name, + session_name=session_name, + peer_names={test_peer.name: original_config}, + ) + first_left_at, first_config = ( + await db_session.execute( + select( + models.SessionPeer.left_at, + models.SessionPeer.configuration, + ).where( + models.SessionPeer.session_name == session_name, + models.SessionPeer.peer_name == test_peer.name, + models.SessionPeer.workspace_name == test_workspace.name, + ) + ) + ).one() + assert first_left_at is None + assert first_config == original_config.model_dump() + + session_peer = ( + await db_session.execute( + select(models.SessionPeer).where( + models.SessionPeer.session_name == session_name, + models.SessionPeer.peer_name == test_peer.name, + models.SessionPeer.workspace_name == test_workspace.name, + ) + ) + ).scalar_one() + session_peer.joined_at = datetime(2020, 1, 1, tzinfo=timezone.utc) + await db_session.commit() + + await crud.set_peers_for_session( + db_session, + workspace_name=test_workspace.name, + session_name=session_name, + peer_names={test_peer.name: updated_config}, + ) + active_joined_at, active_left_at, active_config = ( + await db_session.execute(session_peer_stmt) + ).one() + assert active_joined_at == datetime(2020, 1, 1, tzinfo=timezone.utc) + assert active_left_at is None + assert active_config == original_config.model_dump() + + await crud.set_peers_for_session( + db_session, + workspace_name=test_workspace.name, + session_name=session_name, + peer_names={}, + ) + await crud.set_peers_for_session( + db_session, + workspace_name=test_workspace.name, + session_name=session_name, + peer_names={test_peer.name: updated_config}, + ) + rejoined_joined_at, rejoined_left_at, rejoined_config = ( + await db_session.execute(session_peer_stmt) + ).one() + assert rejoined_joined_at > active_joined_at + assert rejoined_left_at is None + assert rejoined_config == updated_config.model_dump() + @pytest.mark.asyncio async def test_get_session_peer_configuration( self, diff --git a/tests/test_search.py b/tests/test_search.py index 84f3ffa3..4ad0559f 100644 --- a/tests/test_search.py +++ b/tests/test_search.py @@ -4,9 +4,10 @@ import datetime import pytest from nanoid import generate as generate_nanoid +from sqlalchemy import update from sqlalchemy.ext.asyncio import AsyncSession -from src import crud, models +from src import crud, models, schemas from src.utils.search import search @@ -704,3 +705,70 @@ async def test_grep_messages_observer_scoping_left_session_still_visible( matched_ids = [m.public_id for matches, _ in results for m in matches] assert msg_during.public_id in matched_ids assert msg_after.public_id in matched_ids + + +@pytest.mark.asyncio +async def test_peer_perspective_search_after_active_readd( + db_session: AsyncSession, +): + """Active re-add keeps existing messages visible; a genuine rejoin starts a new window.""" + workspace = models.Workspace(name=generate_nanoid()) + peer1 = models.Peer(name="peer1", workspace_name=workspace.name) + peer2 = models.Peer(name="peer2", workspace_name=workspace.name) + session = models.Session(name="session1", workspace_name=workspace.name) + db_session.add_all([workspace, peer1, peer2, session]) + await db_session.flush() + + past_time = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta( + hours=1 + ) + await db_session.execute( + models.session_peers_table.insert().values( + workspace_name=workspace.name, + session_name=session.name, + peer_name=peer1.name, + joined_at=past_time, + left_at=None, + ) + ) + msg_old = models.Message( + content="old persistent message", + session_name=session.name, + peer_name=peer2.name, + workspace_name=workspace.name, + seq_in_session=1, + created_at=past_time + datetime.timedelta(minutes=1), + ) + db_session.add(msg_old) + await db_session.commit() + + session_create = schemas.SessionCreate( + name=session.name, + peers={peer1.name: schemas.SessionPeerConfig()}, + ) + await crud.get_or_create_session(db_session, session_create, workspace.name) + results = await search( + "persistent", + filters={"peer_perspective": peer1.name, "workspace_id": workspace.name}, + limit=10, + ) + assert msg_old.public_id in [m.public_id for m in results] + + await db_session.execute( + update(models.SessionPeer) + .where( + models.SessionPeer.session_name == session.name, + models.SessionPeer.peer_name == peer1.name, + models.SessionPeer.workspace_name == workspace.name, + ) + .values(left_at=datetime.datetime.now(datetime.timezone.utc)) + ) + await db_session.commit() + + await crud.get_or_create_session(db_session, session_create, workspace.name) + results = await search( + "persistent", + filters={"peer_perspective": peer1.name, "workspace_id": workspace.name}, + limit=10, + ) + assert msg_old.public_id not in [m.public_id for m in results]