344 lines
15 KiB
Python
344 lines
15 KiB
Python
"""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()
|