diff --git a/routes/session_routes.py b/routes/session_routes.py index 12fd63b19..cd6f8c357 100644 --- a/routes/session_routes.py +++ b/routes/session_routes.py @@ -272,6 +272,58 @@ def _normalize_group_state(raw_state, parent_session_id: str) -> dict: } +def _group_participant_ids_from_state(state) -> set[str]: + if not isinstance(state, dict): + return set() + raw_participants = state.get("participantSessions") + if not isinstance(raw_participants, list): + return set() + return {str(session_id) for session_id in raw_participants if session_id} + + +def _group_state_query_for_user(db, user): + q = db.query(GroupChatState.parent_session_id, GroupChatState.state) + if user is not None: + q = q.filter(GroupChatState.owner == user) + return q + + +def _group_session_links_for_user(db, user) -> tuple[set[str], set[str]]: + parent_ids: set[str] = set() + participant_ids: set[str] = set() + for parent_id, state in _group_state_query_for_user(db, user).all(): + if parent_id: + parent_ids.add(parent_id) + participant_ids.update(_group_participant_ids_from_state(state)) + return parent_ids, participant_ids + + +def _group_parent_for_participant(db, session_id: str, user) -> str | None: + for parent_id, state in _group_state_query_for_user(db, user).all(): + if session_id in _group_participant_ids_from_state(state): + return parent_id + return None + + +def _set_group_participant_folders(db, participant_ids: set[str], folder: str | None, user) -> None: + if not participant_ids: + return + q = db.query(DbSession).filter(DbSession.id.in_(participant_ids)) + if user is not None: + q = q.filter(DbSession.owner == user) + now = datetime.utcnow() + for participant in q.all(): + participant.folder = folder + participant.updated_at = now + + +def _sync_group_participant_folder(db, parent_session_id: str, folder: str | None, user) -> None: + row = db.query(GroupChatState.state).filter(GroupChatState.parent_session_id == parent_session_id).first() + if row is None: + return + _set_group_participant_folders(db, _group_participant_ids_from_state(row.state), folder, user) + + _HIDDEN_SYSTEM_SESSION_NAMES = { "[Task] Chat Sessions Tidy", "[Task] Documents Tidy", @@ -368,6 +420,7 @@ def setup_session_routes( last_msg_map = {} mode_map = {} msg_count_map = {} + _, group_participant_ids = _group_session_links_for_user(db, user) q = db.query(DbSession.id, DbSession.folder, DbSession.total_input_tokens, DbSession.total_output_tokens, DbSession.is_important, DbSession.created_at, DbSession.updated_at, DbSession.last_message_at, DbSession.mode, DbSession.message_count).filter(DbSession.archived == False) q = owner_filter(q, DbSession, user) rows = q.all() @@ -421,6 +474,7 @@ def setup_session_routes( "message_count": msg_count_map.get(s.id, 0)} for s in user_sessions.values() if not s.archived + and s.id not in group_participant_ids and (s.name or "").strip() not in ("Nobody", "Incognito") and (s.name or "").strip() not in _HIDDEN_SYSTEM_SESSION_NAMES] @@ -577,12 +631,24 @@ def setup_session_routes( if folder is not None: db = SessionLocal() try: + user = effective_user(request) + parent_id = _group_parent_for_participant(db, sid, user) + if parent_id: + raise HTTPException(403, "Move the parent group chat instead") db_session = db.query(DbSession).filter(DbSession.id == sid).first() if db_session: - db_session.folder = folder if folder else None + folder_value = folder if folder else None + db_session.folder = folder_value db_session.updated_at = utcnow_naive() + _sync_group_participant_folder(db, sid, folder_value, user) db.commit() - result["folder"] = folder if folder else None + result["folder"] = folder_value + except HTTPException: + db.rollback() + raise + except Exception: + db.rollback() + raise finally: db.close() # Switch model/endpoint mid-session @@ -697,6 +763,13 @@ def setup_session_routes( group_state.mode = state["mode"] group_state.state = state group_state.updated_at = utcnow_naive() + parent_folder = db.query(DbSession.folder).filter(DbSession.id == sid).first() + _set_group_participant_folders( + db, + participant_ids, + parent_folder.folder if parent_folder else None, + user, + ) db.commit() return {"ok": True, "group_state": state} except HTTPException: @@ -927,6 +1000,9 @@ def setup_session_routes( if not user: raise HTTPException(403, "Authentication required") q = q.filter(DbSession.owner == user) + _, group_participant_ids = _group_session_links_for_user(db, user) + if group_participant_ids: + q = q.filter(~DbSession.id.in_(group_participant_ids)) if search: safe_search = search.replace('%', r'\%').replace('_', r'\_') q = q.filter(DbSession.name.ilike(f"%{safe_search}%", escape='\\')) @@ -1229,6 +1305,8 @@ def setup_session_routes( db = SessionLocal() deleted_empty = 0 deleted_throwaway = 0 + group_parent_ids: set[str] = set() + group_participant_ids: set[str] = set() # Names that indicate a throwaway/test session (case-insensitive exact or prefix match) _THROWAWAY_NAMES = { "test", "testing", "asdf", "asd", "hello", "hi", "hey", @@ -1246,6 +1324,8 @@ def setup_session_routes( elif not single_user_mode: rows_q = rows_q.filter(DbSession.owner == user) rows = rows_q.limit(2000).all() + group_parent_ids, group_participant_ids = _group_session_links_for_user(db, user) + protected_group_ids = group_parent_ids | group_participant_ids folder_map = {r.id: r.folder for r in rows} # Precompute per-session message counts in TWO aggregate queries # instead of 1–3 queries PER session — with many chats the per-row @@ -1258,6 +1338,8 @@ def setup_session_routes( ) cleanup_now = utcnow_naive() for row in rows: + if row.id in protected_group_ids: + continue # Never delete important sessions if getattr(row, 'is_important', False): continue @@ -1336,6 +1418,8 @@ def setup_session_routes( for s in user_sessions.values(): if s.archived or s.name == "Incognito": continue + if s.id in group_participant_ids: + continue if folder_map.get(s.id): # Already in a folder — skip on this pass. continue @@ -1474,6 +1558,7 @@ def setup_session_routes( if db_session: db_session.folder = folder_name db_session.updated_at = utcnow_naive() + _sync_group_participant_folder(db, sid, folder_name, user) updated += 1 db.commit() except Exception as e: diff --git a/src/session_actions.py b/src/session_actions.py index 072bb4c06..041e330e7 100644 --- a/src/session_actions.py +++ b/src/session_actions.py @@ -53,6 +53,15 @@ def is_session_recently_active(row, now=None, grace=_FRESH_SESSION_GRACE) -> boo return False +def _group_participant_ids_from_state(state) -> set[str]: + if not isinstance(state, dict): + return set() + participants = state.get("participantSessions") + if not isinstance(participants, list): + return set() + return {str(session_id) for session_id in participants if session_id} + + async def run_auto_sort(owner: str, skip_llm: bool = False, delete_throwaway: bool = True) -> str: """Run session cleanup + (optional) AI folder sort for the given owner. @@ -65,7 +74,7 @@ async def run_auto_sort(owner: str, skip_llm: bool = False, delete_throwaway: bo Returns a human-readable summary of what was done. """ - from core.database import SessionLocal, Session as DbSession, ChatMessage as DbMsg + from core.database import SessionLocal, Session as DbSession, ChatMessage as DbMsg, GroupChatState from src.llm_core import llm_call_async from src.task_endpoint import resolve_task_endpoint @@ -79,9 +88,21 @@ async def run_auto_sort(owner: str, skip_llm: bool = False, delete_throwaway: bo DbSession.archived == False, *([DbSession.owner == owner] if owner else []), ).all() + group_query = db.query(GroupChatState.parent_session_id, GroupChatState.state) + if owner: + group_query = group_query.filter(GroupChatState.owner == owner) + group_parent_to_participants = {} + group_participant_ids: set[str] = set() + for parent_id, state in group_query.all(): + participants = _group_participant_ids_from_state(state) + group_parent_to_participants[parent_id] = participants + group_participant_ids.update(participants) + protected_group_ids = set(group_parent_to_participants) | group_participant_ids cleanup_now = _utcnow_naive() for row in rows: + if row.id in protected_group_ids: + continue if getattr(row, 'is_important', False): continue created_at = _as_naive_utc(row.created_at or row.updated_at) or _utcnow_naive() @@ -150,6 +171,8 @@ async def run_auto_sort(owner: str, skip_llm: bool = False, delete_throwaway: bo for row in remaining: if row.name == "Incognito": continue + if row.id in group_participant_ids: + continue session_list.append({ "id": row.id, "name": row.name or "(unnamed)", @@ -240,6 +263,14 @@ async def run_auto_sort(owner: str, skip_llm: bool = False, delete_throwaway: bo if db_sess: db_sess.folder = folder_name db_sess.updated_at = _utcnow_naive() + participant_ids = group_parent_to_participants.get(full_id, set()) + if participant_ids: + child_query = db.query(DbSession).filter(DbSession.id.in_(participant_ids)) + if owner: + child_query = child_query.filter(DbSession.owner == owner) + for child in child_query.all(): + child.folder = folder_name + child.updated_at = _utcnow_naive() updated += 1 db.commit() diff --git a/src/session_search.py b/src/session_search.py index d8b994fa8..a17a1d840 100644 --- a/src/session_search.py +++ b/src/session_search.py @@ -11,6 +11,7 @@ from typing import Any, Iterable from sqlalchemy import text from core.database import ChatMessage as DBChatMessage +from core.database import GroupChatState from core.database import Session as DBSession from core.database import SessionLocal @@ -134,6 +135,30 @@ def _owner_filter(query, owner: str | None, include_legacy_owner: bool): return query.filter((DBSession.owner == owner) | (DBSession.owner.is_(None))) +def _group_participant_ids_from_state(state) -> set[str]: + if not isinstance(state, dict): + return set() + participants = state.get("participantSessions") + if not isinstance(participants, list): + return set() + return {str(session_id) for session_id in participants if session_id} + + +def _hidden_group_participant_ids(db, owner: str | None, restrict_owner: bool, include_legacy_owner: bool) -> set[str]: + q = db.query(GroupChatState.state) + if restrict_owner: + if owner is None: + q = q.filter(GroupChatState.owner.is_(None)) + elif include_legacy_owner: + q = q.filter((GroupChatState.owner == owner) | (GroupChatState.owner.is_(None))) + else: + q = q.filter(GroupChatState.owner == owner) + participant_ids: set[str] = set() + for (state,) in q.all(): + participant_ids.update(_group_participant_ids_from_state(state)) + return participant_ids + + def _context_for_message(db, msg: DBChatMessage, count: int) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: if count <= 0 or not msg.timestamp: return [], [] @@ -210,6 +235,9 @@ def _search_like( q = q.filter(~DBSession.name.like("SFT trace batch%")) if restrict_owner: q = _owner_filter(q, owner, include_legacy_owner) + hidden_participant_ids = _hidden_group_participant_ids(db, owner, restrict_owner, include_legacy_owner) + if hidden_participant_ids: + q = q.filter(~DBChatMessage.session_id.in_(hidden_participant_ids)) rows = q.order_by(DBChatMessage.timestamp.desc()).limit(limit).all() shaped = ((msg, session_name, _snippet(msg.content or "", query)) for msg, session_name in rows) return _rows_to_results(db, shaped, query, context_messages) @@ -278,6 +306,8 @@ def _search_fts( """ ) + hidden_participant_ids = _hidden_group_participant_ids(db, owner, restrict_owner, include_legacy_owner) + try: hits = db.execute(sql, params).fetchall() except Exception as e: @@ -293,6 +323,8 @@ def _search_fts( found = by_id.get(hit[0]) if found: msg, session_name = found + if msg.session_id in hidden_participant_ids: + continue rows.append((msg, session_name, hit[1] or "")) return _rows_to_results(db, rows, query, context_messages) diff --git a/static/js/sessions.js b/static/js/sessions.js index 4e6ba44d1..901e29d27 100644 --- a/static/js/sessions.js +++ b/static/js/sessions.js @@ -369,10 +369,24 @@ function getFolderNames() { async function moveToFolder(sessionId, folderName) { const fd = new FormData(); fd.append('folder', folderName || ''); - await fetch(`${API_BASE}/api/session/${sessionId}`, { method: 'PATCH', body: fd }); + const res = await fetch(`${API_BASE}/api/session/${sessionId}`, { + method: 'PATCH', + body: fd, + credentials: 'same-origin', + }); + if (!res.ok) { + let message = 'Failed to move session'; + try { + const data = await res.json(); + message = data.detail || data.message || message; + } catch (e) {} + uiModule.showError(message); + throw new Error(message); + } + const data = await res.json().catch(() => ({})); // Update local data const s = sessions.find(x => x.id === sessionId); - if (s) s.folder = folderName || null; + if (s) s.folder = data.folder !== undefined ? data.folder : (folderName || null); renderSessionList(); } diff --git a/tests/test_group_chat_state_routes.py b/tests/test_group_chat_state_routes.py index 4db5f677e..106ae5b1e 100644 --- a/tests/test_group_chat_state_routes.py +++ b/tests/test_group_chat_state_routes.py @@ -76,13 +76,24 @@ def _endpoint(router, path, method): ) -def _routes(monkeypatch): +def _routes(monkeypatch, session_manager=None): import routes.session_routes as sr _stub_multipart_if_missing(monkeypatch) monkeypatch.setattr(sr, "SessionLocal", _TS) monkeypatch.setattr(sr, "effective_user", lambda request: "alice") - return sr.setup_session_routes(MagicMock(), {}) + return sr.setup_session_routes(session_manager or MagicMock(), {}) + + +def _session_stub(session_id, name, archived=False): + return types.SimpleNamespace( + id=session_id, + name=name, + model="llama3", + endpoint_url="http://localhost:11434", + rag=False, + archived=archived, + ) def _group_payload(parent_id, participant_ids): @@ -115,6 +126,20 @@ def _group_payload(parent_id, participant_ids): } +def _add_group_state(parent_id, participant_ids, owner="alice"): + db = _TS() + try: + db.add(GroupChatState( + parent_session_id=parent_id, + owner=owner, + mode="round-robin", + state=_group_payload(parent_id, participant_ids) | {"parentSessionId": parent_id}, + )) + db.commit() + finally: + db.close() + + def test_group_chat_state_round_trips_personas_and_participant_sessions(monkeypatch): _reset_db() router = _routes(monkeypatch) @@ -140,6 +165,89 @@ def test_group_chat_state_round_trips_personas_and_participant_sessions(monkeypa assert state["models"][0]["character"]["characterPrompt"] == "Offer wise counsel." +def test_list_sessions_hides_group_participants(monkeypatch): + _reset_db() + parent_id = str(uuid.uuid4()) + child_ids = [str(uuid.uuid4()), str(uuid.uuid4())] + normal_id = str(uuid.uuid4()) + _add_session(parent_id, name="[GRP] Athena, Mistral") + for session_id in child_ids: + _add_session(session_id, name="[GRP] participant") + _add_session(normal_id, name="normal chat") + _add_group_state(parent_id, child_ids) + + sm = MagicMock() + sm.get_sessions_for_user.return_value = { + parent_id: _session_stub(parent_id, "[GRP] Athena, Mistral"), + child_ids[0]: _session_stub(child_ids[0], "[GRP] Athena"), + child_ids[1]: _session_stub(child_ids[1], "[GRP] Mistral"), + normal_id: _session_stub(normal_id, "normal chat"), + } + router = _routes(monkeypatch, sm) + list_sessions = _endpoint(router, "/api/sessions", "GET") + + returned_ids = {session["id"] for session in list_sessions(request=MagicMock())} + + assert parent_id in returned_ids + assert normal_id in returned_ids + assert not set(child_ids) & returned_ids + + +def test_group_parent_folder_move_cascades_and_child_move_is_blocked(monkeypatch): + _reset_db() + parent_id = str(uuid.uuid4()) + child_ids = [str(uuid.uuid4()), str(uuid.uuid4())] + _add_session(parent_id, name="[GRP] Athena, Mistral") + for session_id in child_ids: + _add_session(session_id, name="[GRP] participant") + _add_group_state(parent_id, child_ids) + + sm = MagicMock() + sm.get_session.return_value = _session_stub(parent_id, "[GRP] Athena, Mistral") + router = _routes(monkeypatch, sm) + patch_session = _endpoint(router, "/api/session/{sid}", "PATCH") + + result = patch_session( + request=MagicMock(), + sid=parent_id, + name=None, + folder="Research", + model=None, + endpoint_url=None, + endpoint_id=None, + ) + + db = _TS() + try: + folders = { + row.id: row.folder + for row in db.query(DbSession).filter(DbSession.id.in_([parent_id, *child_ids])).all() + } + finally: + db.close() + assert result["folder"] == "Research" + assert folders == {parent_id: "Research", child_ids[0]: "Research", child_ids[1]: "Research"} + + with pytest.raises(HTTPException) as exc: + patch_session( + request=MagicMock(), + sid=child_ids[0], + name=None, + folder="Solo", + model=None, + endpoint_url=None, + endpoint_id=None, + ) + assert exc.value.status_code == 403 + + db = _TS() + try: + child = db.query(DbSession).filter(DbSession.id == child_ids[0]).first() + assert child.folder == "Research" + finally: + db.close() + + def test_group_chat_state_rejects_participants_from_other_users(monkeypatch): _reset_db() router = _routes(monkeypatch) diff --git a/tests/test_session_actions_cleanup.py b/tests/test_session_actions_cleanup.py index 221713d33..898c92439 100644 --- a/tests/test_session_actions_cleanup.py +++ b/tests/test_session_actions_cleanup.py @@ -21,7 +21,7 @@ from sqlalchemy.orm import sessionmaker from sqlalchemy.pool import NullPool import core.database as cdb -from core.database import ChatMessage as DbMessage, Session as DbSession, utcnow_naive +from core.database import ChatMessage as DbMessage, GroupChatState, Session as DbSession, utcnow_naive import src.session_actions as session_actions @@ -164,3 +164,56 @@ def test_auto_sort_still_deletes_old_throwaway_sessions(monkeypatch): assert "Cleaned 1 sessions" in result finally: db.close() + + +def test_auto_sort_keeps_group_parent_and_participants(monkeypatch): + session_factory = _make_session_factory() + _install_session_factory(monkeypatch, session_factory) + + old_time = utcnow_naive() - timedelta(hours=2) + parent_id = "p-" + uuid.uuid4().hex + child_ids = ["c-" + uuid.uuid4().hex, "c-" + uuid.uuid4().hex] + db = session_factory() + try: + for sid, name in [ + (parent_id, "[GRP] Athena, Mistral"), + (child_ids[0], "[GRP] Athena"), + (child_ids[1], "[GRP] Mistral"), + ]: + db.add( + DbSession( + id=sid, + owner="alice", + name=name, + endpoint_url="", + model="", + archived=False, + message_count=0, + created_at=old_time, + updated_at=old_time, + last_accessed=old_time, + last_message_at=old_time, + ) + ) + db.commit() + db.add( + GroupChatState( + parent_session_id=parent_id, + owner="alice", + mode="parallel", + state={"participantSessions": child_ids}, + ) + ) + db.commit() + finally: + db.close() + + result = asyncio.run(session_actions.run_auto_sort("alice", skip_llm=True)) + + db = session_factory() + try: + remaining_ids = {row.id for row in db.query(DbSession).all()} + assert {parent_id, *child_ids} <= remaining_ids + assert "Cleaned 0 sessions" in result + finally: + db.close() diff --git a/tests/test_session_search.py b/tests/test_session_search.py index 467653635..7f83407dc 100644 --- a/tests/test_session_search.py +++ b/tests/test_session_search.py @@ -7,6 +7,7 @@ from sqlalchemy.orm import sessionmaker from core.database import Base from core.database import ChatMessage as DbChatMessage +from core.database import GroupChatState from core.database import Session as DbSession from src.session_search import SessionSearchResult, search_session_messages @@ -227,6 +228,32 @@ def test_session_search_excludes_archived_by_default(): db.close() +def test_session_search_hides_group_participant_sessions(): + db = _db(with_fts=True) + try: + base = datetime(2026, 1, 1, 12, 0, 0) + _add_session(db, "parent", owner="alice", name="[GRP] Athena, Mistral") + _add_session(db, "child", owner="alice", name="[GRP] Athena") + db.commit() + db.add( + GroupChatState( + parent_session_id="parent", + owner="alice", + mode="parallel", + state={"participantSessions": ["child"]}, + ) + ) + _add_message(db, "parent", "m-parent", "assistant", "group visible target", base) + _add_message(db, "child", "m-child", "assistant", "group hidden target", base + timedelta(minutes=1)) + db.commit() + + results = search_session_messages("group target", owner="alice", db=db) + + assert [r.message_id for r in results] == ["m-parent"] + finally: + db.close() + + def test_chat_messages_fts_migration_backfills_and_tracks_inserts(tmp_path, monkeypatch): from core import database as cdb