81 lines
3.2 KiB
Python
81 lines
3.2 KiB
Python
"""SQLAlchemy-backed per-conversation share-link repository."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import datetime
|
|
from typing import Any
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|
|
|
from deerflow.persistence.thread_shares.base import ThreadShareStore, generate_share_code
|
|
from deerflow.persistence.thread_shares.model import ThreadShareRow
|
|
|
|
|
|
def _row_to_dict(row: ThreadShareRow) -> dict[str, Any]:
|
|
data = row.to_dict()
|
|
created = data.get("created_at")
|
|
if isinstance(created, datetime):
|
|
data["created_at"] = created.isoformat()
|
|
return data
|
|
|
|
|
|
class ThreadShareRepository(ThreadShareStore):
|
|
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
|
|
self._sf = session_factory
|
|
|
|
async def create(self, owner_user_id: str, thread_id: str, *, mode: str = "import") -> dict[str, Any]:
|
|
# Retry on the (astronomically unlikely) code collision.
|
|
async with self._sf() as session:
|
|
for _ in range(8):
|
|
code = generate_share_code()
|
|
if await session.get(ThreadShareRow, code) is not None:
|
|
continue
|
|
row = ThreadShareRow(
|
|
share_code=code,
|
|
owner_user_id=owner_user_id,
|
|
thread_id=thread_id,
|
|
revoked=False,
|
|
mode=mode,
|
|
)
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _row_to_dict(row)
|
|
raise RuntimeError("Failed to allocate a unique share code")
|
|
|
|
async def get(self, share_code: str) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
row = await session.get(ThreadShareRow, share_code)
|
|
return _row_to_dict(row) if row is not None else None
|
|
|
|
async def get_active_for_thread(self, owner_user_id: str, thread_id: str, *, mode: str = "import") -> dict[str, Any] | None:
|
|
stmt = (
|
|
select(ThreadShareRow)
|
|
.where(
|
|
ThreadShareRow.owner_user_id == owner_user_id,
|
|
ThreadShareRow.thread_id == thread_id,
|
|
ThreadShareRow.mode == mode,
|
|
ThreadShareRow.revoked.is_(False),
|
|
)
|
|
.order_by(ThreadShareRow.created_at.desc())
|
|
)
|
|
async with self._sf() as session:
|
|
row = (await session.execute(stmt)).scalars().first()
|
|
return _row_to_dict(row) if row is not None else None
|
|
|
|
async def list_for_owner(self, owner_user_id: str) -> list[dict[str, Any]]:
|
|
stmt = select(ThreadShareRow).where(ThreadShareRow.owner_user_id == owner_user_id).order_by(ThreadShareRow.created_at.desc())
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
return [_row_to_dict(row) for row in result.scalars()]
|
|
|
|
async def revoke(self, share_code: str, owner_user_id: str) -> bool:
|
|
async with self._sf() as session:
|
|
row = await session.get(ThreadShareRow, share_code)
|
|
if row is None or row.owner_user_id != owner_user_id:
|
|
return False
|
|
row.revoked = True
|
|
await session.commit()
|
|
return True
|