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

50 lines
1.7 KiB
Python

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