deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/persistence/thread_meta/sql.py
2026-09-07 18:24:55 +08:00

344 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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