"""SQLAlchemy-backed embed session repository.""" from __future__ import annotations from datetime import UTC, datetime from typing import Any from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.embed_sessions.base import EmbedSessionStore from deerflow.persistence.embed_sessions.model import EmbedSessionRow def _row_to_dict(row: EmbedSessionRow) -> dict[str, Any]: data = row.to_dict() for key in ("created_at", "updated_at"): value = data.get(key) if isinstance(value, datetime): data[key] = value.isoformat() return data class EmbedSessionRepository(EmbedSessionStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory async def upsert(self, session_id: str, thread_id: str) -> dict[str, Any]: async with self._sf() as session: row = await session.get(EmbedSessionRow, session_id) now = datetime.now(UTC) if row is None: row = EmbedSessionRow( session_id=session_id, thread_id=thread_id, created_at=now, updated_at=now, ) session.add(row) else: row.thread_id = thread_id row.updated_at = now await session.commit() await session.refresh(row) return _row_to_dict(row) async def get(self, session_id: str) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(EmbedSessionRow, session_id) return _row_to_dict(row) if row is not None else None