fix(group): hide participant sessions

This commit is contained in:
Matyas Fenyves 2026-06-12 11:09:03 +02:00
parent 284e418c7c
commit 4ebcebfa3d
7 changed files with 358 additions and 8 deletions

View file

@ -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 13 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:

View file

@ -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()

View file

@ -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)

View file

@ -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();
}

View file

@ -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)

View file

@ -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()

View file

@ -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