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