mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-08-05 02:45:28 +00:00
fix(group): hide participant sessions
This commit is contained in:
parent
284e418c7c
commit
4ebcebfa3d
7 changed files with 358 additions and 8 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue