"""SQLAlchemy-backed thread metadata repository.""" from __future__ import annotations from datetime import UTC, datetime from typing import Any from sqlalchemy import func, select, update from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.thread_meta.base import ThreadMetaStore, is_excluded_system_thread from deerflow.persistence.thread_meta.model import ThreadMetaRow from deerflow.runtime.user_context import AUTO, _AutoSentinel, resolve_user_id class ThreadMetaRepository(ThreadMetaStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory @staticmethod def _row_to_dict(row: ThreadMetaRow) -> dict[str, Any]: d = row.to_dict() d["metadata"] = d.pop("metadata_json", {}) for key in ("created_at", "updated_at"): val = d.get(key) if isinstance(val, datetime): d[key] = val.isoformat() return d async def create( self, thread_id: str, *, assistant_id: str | None = None, user_id: str | None | _AutoSentinel = AUTO, display_name: str | None = None, metadata: dict | None = None, ) -> dict: # Auto-resolve user_id from contextvar when AUTO; explicit None # creates an orphan row (used by migration scripts). resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.create") now = datetime.now(UTC) row = ThreadMetaRow( thread_id=thread_id, assistant_id=assistant_id, user_id=resolved_user_id, display_name=display_name, status="idle", metadata_json=metadata or {}, created_at=now, updated_at=now, ) async with self._sf() as session: session.add(row) await session.commit() # Deliberately no session.refresh(): every column above is assigned # explicitly and the session factory uses expire_on_commit=False, so # the in-memory row is already complete. The re-SELECT would run on a # different pooled connection than the INSERT, which a MySQL # read/write-splitting endpoint can route to a lagging replica -- # the row comes back empty and SQLAlchemy raises # "Could not refresh instance". Concurrency makes that far more # likely, so a single-user deployment never sees it. return self._row_to_dict(row) async def get( self, thread_id: str, *, user_id: str | None | _AutoSentinel = AUTO, ) -> dict | None: resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.get") async with self._sf() as session: row = await session.get(ThreadMetaRow, thread_id) if row is None: return None # Enforce owner filter unless explicitly bypassed (user_id=None). if resolved_user_id is not None and row.user_id != resolved_user_id: return None return self._row_to_dict(row) async def check_access(self, thread_id: str, user_id: str, *, require_existing: bool = False) -> bool: """Check if ``user_id`` has access to ``thread_id``. Two modes — one row, two distinct semantics depending on what the caller is about to do: - ``require_existing=False`` (default, permissive): Returns True for: row missing (untracked legacy thread), ``row.user_id`` is None (shared / pre-auth data), or ``row.user_id == user_id``. Use for **read-style** decorators where treating an untracked thread as accessible preserves backward-compat. - ``require_existing=True`` (strict): Returns True **only** when the row exists AND (``row.user_id == user_id`` OR ``row.user_id is None``). Use for **destructive / mutating** decorators (DELETE, PATCH, state-update) so a thread that has *already been deleted* cannot be re-targeted by any caller — closing the delete-idempotence cross-user gap where the row vanishing made every other user appear to "own" it. """ async with self._sf() as session: row = await session.get(ThreadMetaRow, thread_id) if row is None: return not require_existing md = row.metadata_json or {} # 全局共享、不区分用户(任何登录用户都可读写/续聊): # ① 任务深链对话(metadata.taskId); # ② 圆桌会商系统线程(metadata.thread_type=="roundtable")——其 Step2/Step3 # 沙箱产物(md/流程图/报告)需让历史草稿或其他账号打开同一会商时也能查看, # 与任务深链对话一致(线程 id 仅存在于草稿中,不会被无关账号发现)。 if md.get("taskId") or md.get("thread_type") == "roundtable": return True if row.user_id is None: return True return row.user_id == user_id async def search( self, *, metadata: dict | None = None, status: str | None = None, limit: int = 100, offset: int = 0, user_id: str | None | _AutoSentinel = AUTO, exclude_system: bool = False, query: str | None = None, ) -> list[dict]: """Search threads with optional metadata and status filters. Owner filter is enforced by default: caller must be in a user context. Pass ``user_id=None`` to bypass (migration/CLI). """ resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.search") # 任务深链对话(按 metadata.taskId 过滤)全局共享、不区分用户:跳过 owner 过滤。 if metadata and metadata.get("taskId"): resolved_user_id = None # thread_id tie-breaker keeps OFFSET pagination stable when many rows # share the same updated_at (bulk imports, scheduler bursts). stmt = select(ThreadMetaRow).order_by(ThreadMetaRow.updated_at.desc(), ThreadMetaRow.thread_id.desc()) if resolved_user_id is not None: stmt = stmt.where(ThreadMetaRow.user_id == resolved_user_id) if status: stmt = stmt.where(ThreadMetaRow.status == status) if query and query.strip(): # Case-insensitive title search. display_name is a real column, so # this filter runs in SQL even when the JSON filters below fall # back to Python-side scanning. stmt = stmt.where(func.lower(ThreadMetaRow.display_name).contains(query.strip().lower())) if not metadata and not exclude_system: stmt = stmt.limit(limit).offset(offset) async with self._sf() as session: result = await session.execute(stmt) return [self._row_to_dict(r) for r in result.scalars()] # Metadata filters live inside the JSON blob, so rows are scanned in # batches and filtered in Python until the requested window is full. # A fixed over-fetch window would under-fill pages when most rows are # filtered out (e.g. heavy scheduled-task usage), making callers that # paginate believe the list is exhausted. TODO(Phase 2): use JSON DB # operators (Postgres @>, SQLite json_extract) for server-side filtering. def _matches(row: dict) -> bool: md = row.get("metadata") or {} if metadata and not all(md.get(k) == v for k, v in metadata.items()): return False if exclude_system and is_excluded_system_thread(md): return False return True needed = offset + limit batch_size = max(500, needed) matched: list[dict] = [] scan_offset = 0 async with self._sf() as session: while len(matched) < needed: result = await session.execute(stmt.limit(batch_size).offset(scan_offset)) rows = result.scalars().all() if not rows: break matched.extend(d for d in (self._row_to_dict(r) for r in rows) if _matches(d)) if len(rows) < batch_size: break scan_offset += len(rows) return matched[offset:needed] async def count( self, *, metadata: dict | None = None, status: str | None = None, user_id: str | None | _AutoSentinel = AUTO, exclude_system: bool = False, query: str | None = None, ) -> int: """Count threads matching the same filters as :meth:`search`.""" resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.count") # 任务深链对话(按 metadata.taskId 过滤)全局共享、不区分用户:跳过 owner 过滤。 if metadata and metadata.get("taskId"): resolved_user_id = None conditions = [] if resolved_user_id is not None: conditions.append(ThreadMetaRow.user_id == resolved_user_id) if status: conditions.append(ThreadMetaRow.status == status) if query and query.strip(): conditions.append(func.lower(ThreadMetaRow.display_name).contains(query.strip().lower())) if not metadata and not exclude_system: stmt = select(func.count()).select_from(ThreadMetaRow) for condition in conditions: stmt = stmt.where(condition) async with self._sf() as session: return int((await session.execute(stmt)).scalar_one()) # JSON filters can't run in portable SQL — scan just the metadata # column in batches and count in Python (mirrors search()). stmt = select(ThreadMetaRow.metadata_json).order_by(ThreadMetaRow.thread_id) for condition in conditions: stmt = stmt.where(condition) total = 0 batch_size = 1000 scan_offset = 0 async with self._sf() as session: while True: result = await session.execute(stmt.limit(batch_size).offset(scan_offset)) rows = result.scalars().all() if not rows: break for metadata_json in rows: md = metadata_json or {} if metadata and not all(md.get(k) == v for k, v in metadata.items()): continue if exclude_system and is_excluded_system_thread(md): continue total += 1 if len(rows) < batch_size: break scan_offset += len(rows) return total async def _check_ownership(self, session: AsyncSession, thread_id: str, resolved_user_id: str | None) -> bool: """Return True if the row exists and is owned (or filter bypassed).""" if resolved_user_id is None: return True # explicit bypass row = await session.get(ThreadMetaRow, thread_id) if row is None: return False # 任务深链对话(metadata.taskId)全局共享:任何用户都可更新标题/状态/updated_at(续聊置顶)。 if (row.metadata_json or {}).get("taskId"): return True return row.user_id == resolved_user_id async def update_display_name( self, thread_id: str, display_name: str, *, user_id: str | None | _AutoSentinel = AUTO, ) -> None: """Update the display_name (title) for a thread.""" resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_display_name") async with self._sf() as session: if not await self._check_ownership(session, thread_id, resolved_user_id): return await session.execute(update(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).values(display_name=display_name, updated_at=datetime.now(UTC))) await session.commit() async def update_status( self, thread_id: str, status: str, *, user_id: str | None | _AutoSentinel = AUTO, ) -> None: resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_status") async with self._sf() as session: if not await self._check_ownership(session, thread_id, resolved_user_id): return await session.execute(update(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).values(status=status, updated_at=datetime.now(UTC))) await session.commit() async def update_assistant_id( self, thread_id: str, assistant_id: str, *, user_id: str | None | _AutoSentinel = AUTO, ) -> None: resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_assistant_id") async with self._sf() as session: if not await self._check_ownership(session, thread_id, resolved_user_id): return await session.execute(update(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id).values(assistant_id=assistant_id, updated_at=datetime.now(UTC))) await session.commit() async def update_metadata( self, thread_id: str, metadata: dict, *, user_id: str | None | _AutoSentinel = AUTO, ) -> None: """Merge ``metadata`` into ``metadata_json``. Read-modify-write inside a single session/transaction so concurrent callers see consistent state. No-op if the row does not exist or the user_id check fails. """ resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_metadata") async with self._sf() as session: row = await session.get(ThreadMetaRow, thread_id) if row is None: return # 任务深链对话(metadata.taskId)全局共享:跳过 owner 校验。 task_shared = bool((row.metadata_json or {}).get("taskId")) if resolved_user_id is not None and not task_shared and row.user_id != resolved_user_id: return merged = dict(row.metadata_json or {}) merged.update(metadata) row.metadata_json = merged row.updated_at = datetime.now(UTC) await session.commit() async def delete( self, thread_id: str, *, user_id: str | None | _AutoSentinel = AUTO, ) -> None: resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.delete") async with self._sf() as session: row = await session.get(ThreadMetaRow, thread_id) if row is None: return # 任务深链对话(metadata.taskId)全局共享:任何用户都可删除。 task_shared = bool((row.metadata_json or {}).get("taskId")) if resolved_user_id is not None and not task_shared and row.user_id != resolved_user_id: return await session.delete(row) await session.commit()