05be7f9ad0
`ThreadMetaStore.check_access(thread_id, user_id, workspace_id, *, require_existing)` is now a three-positional method. Cross-workspace is denied unconditionally — even when the row exists and `user_id` matches — so the decorator layer can convert the False into a **404** and never leak the existence of a thread across tenants. Inside the workspace, the existing legacy semantics still hold (NULL `row.user_id` stays "shared in workspace", `require_existing` still gates the missing-row path against ghost-row re-targeting). `@require_permission(owner_check=True)` in `app/gateway/authz.py` now reads the active workspace from `get_effective_workspace_id()` (set by PR4 AuthMiddleware; falls back to "default" in no-auth dev) and passes it through. The existing 404-not-403 mapping is unchanged. Existing positional callers in `test_thread_meta_repo.py` and the permissive mock in `test_threads_router.py` were updated for the new arity. 58 thread_meta / router / memory tests stay green; 90 auth / uploads / suggestions tests stay green.
265 lines
11 KiB
Python
265 lines
11 KiB
Python
"""SQLAlchemy-backed thread metadata repository."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import UTC, datetime
|
|
from typing import Any
|
|
|
|
from sqlalchemy import select, update
|
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|
|
|
from deerflow.persistence.thread_meta.base import ThreadMetaStore
|
|
from deerflow.persistence.thread_meta.model import ThreadMetaRow
|
|
from deerflow.runtime.user_context import AUTO, _AutoSentinel, resolve_user_id
|
|
from deerflow.runtime.workspace_context import AUTO as WORKSPACE_AUTO
|
|
from deerflow.runtime.workspace_context import (
|
|
_AutoSentinel as _WorkspaceAutoSentinel,
|
|
)
|
|
from deerflow.runtime.workspace_context import (
|
|
resolve_workspace_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,
|
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
|
display_name: str | None = None,
|
|
metadata: dict | None = None,
|
|
) -> dict:
|
|
# Auto-resolve both user_id and workspace_id from contextvars when
|
|
# AUTO; explicit None creates an orphan row (used by migration
|
|
# scripts that intentionally bypass scope).
|
|
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.create")
|
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.create")
|
|
now = datetime.now(UTC)
|
|
row = ThreadMetaRow(
|
|
thread_id=thread_id,
|
|
assistant_id=assistant_id,
|
|
user_id=resolved_user_id,
|
|
workspace_id=resolved_workspace_id,
|
|
display_name=display_name,
|
|
metadata_json=metadata or {},
|
|
created_at=now,
|
|
updated_at=now,
|
|
)
|
|
async with self._sf() as session:
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return self._row_to_dict(row)
|
|
|
|
async def get(
|
|
self,
|
|
thread_id: str,
|
|
*,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
|
) -> dict | None:
|
|
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.get")
|
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.get")
|
|
stmt = select(ThreadMetaRow).where(ThreadMetaRow.thread_id == thread_id)
|
|
if resolved_workspace_id is not None:
|
|
stmt = stmt.where(ThreadMetaRow.workspace_id == resolved_workspace_id)
|
|
async with self._sf() as session:
|
|
row = (await session.execute(stmt)).scalar_one_or_none()
|
|
if row is None:
|
|
return None
|
|
# Owner filter still applies inside the workspace scope.
|
|
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,
|
|
workspace_id: str,
|
|
*,
|
|
require_existing: bool = False,
|
|
) -> bool:
|
|
"""Check if ``user_id`` in ``workspace_id`` has access to ``thread_id``.
|
|
|
|
Three filters layered, from outside in:
|
|
|
|
- Cross-workspace is **always** denied (returns False), even when
|
|
the row exists and ``user_id`` matches. The decorator layer
|
|
converts a False into a 404 so cross-tenant access never leaks
|
|
the existence of a thread.
|
|
- Missing row honours ``require_existing``: False by default
|
|
(permissive — untracked legacy threads still readable), True
|
|
for destructive routes (DELETE / PATCH) so a re-targeted ghost
|
|
row cannot be claimed.
|
|
- Within the workspace, ``row.user_id IS NULL`` keeps the legacy
|
|
"shared / pre-auth" semantics — readable by anyone in the
|
|
workspace. ``row.user_id == user_id`` is the normal case.
|
|
"""
|
|
async with self._sf() as session:
|
|
row = await session.get(ThreadMetaRow, thread_id)
|
|
if row is None:
|
|
return not require_existing
|
|
if row.workspace_id is not None and row.workspace_id != workspace_id:
|
|
return False
|
|
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,
|
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
|
) -> list[dict]:
|
|
"""Search threads with optional metadata and status filters.
|
|
|
|
Both workspace and owner filters are enforced by default. Pass
|
|
``workspace_id=None`` and / or ``user_id=None`` to bypass for
|
|
migration / CLI paths.
|
|
"""
|
|
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.search")
|
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.search")
|
|
stmt = select(ThreadMetaRow).order_by(ThreadMetaRow.updated_at.desc())
|
|
if resolved_workspace_id is not None:
|
|
stmt = stmt.where(ThreadMetaRow.workspace_id == resolved_workspace_id)
|
|
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 metadata:
|
|
# When metadata filter is active, fetch a larger window and filter
|
|
# in Python. TODO(Phase 2): use JSON DB operators (Postgres @>,
|
|
# SQLite json_extract) for server-side filtering.
|
|
stmt = stmt.limit(limit * 5 + offset)
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
rows = [self._row_to_dict(r) for r in result.scalars()]
|
|
rows = [r for r in rows if all(r.get("metadata", {}).get(k) == v for k, v in metadata.items())]
|
|
return rows[offset : offset + limit]
|
|
else:
|
|
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()]
|
|
|
|
async def _check_ownership(
|
|
self,
|
|
session: AsyncSession,
|
|
thread_id: str,
|
|
resolved_user_id: str | None,
|
|
resolved_workspace_id: str | None,
|
|
) -> bool:
|
|
"""Return True if the row exists, is in scope, and is owned (or filter bypassed)."""
|
|
row = await session.get(ThreadMetaRow, thread_id)
|
|
if row is None:
|
|
return False
|
|
if resolved_workspace_id is not None and row.workspace_id != resolved_workspace_id:
|
|
return False
|
|
if resolved_user_id is not None and row.user_id != resolved_user_id:
|
|
return False
|
|
return True
|
|
|
|
async def update_display_name(
|
|
self,
|
|
thread_id: str,
|
|
display_name: str,
|
|
*,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
|
) -> None:
|
|
"""Update the display_name (title) for a thread."""
|
|
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_display_name")
|
|
resolved_workspace_id = resolve_workspace_id(workspace_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, resolved_workspace_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,
|
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
|
) -> None:
|
|
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_status")
|
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.update_status")
|
|
async with self._sf() as session:
|
|
if not await self._check_ownership(session, thread_id, resolved_user_id, resolved_workspace_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_metadata(
|
|
self,
|
|
thread_id: str,
|
|
metadata: dict,
|
|
*,
|
|
user_id: str | None | _AutoSentinel = AUTO,
|
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_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 workspace / user check fails.
|
|
"""
|
|
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.update_metadata")
|
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.update_metadata")
|
|
async with self._sf() as session:
|
|
row = await session.get(ThreadMetaRow, thread_id)
|
|
if row is None:
|
|
return
|
|
if resolved_workspace_id is not None and row.workspace_id != resolved_workspace_id:
|
|
return
|
|
if resolved_user_id is not None 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,
|
|
workspace_id: str | None | _WorkspaceAutoSentinel = WORKSPACE_AUTO,
|
|
) -> None:
|
|
resolved_user_id = resolve_user_id(user_id, method_name="ThreadMetaRepository.delete")
|
|
resolved_workspace_id = resolve_workspace_id(workspace_id, method_name="ThreadMetaRepository.delete")
|
|
async with self._sf() as session:
|
|
row = await session.get(ThreadMetaRow, thread_id)
|
|
if row is None:
|
|
return
|
|
if resolved_workspace_id is not None and row.workspace_id != resolved_workspace_id:
|
|
return
|
|
if resolved_user_id is not None and row.user_id != resolved_user_id:
|
|
return
|
|
await session.delete(row)
|
|
await session.commit()
|