80 lines
3.0 KiB
Python
80 lines
3.0 KiB
Python
"""SQLAlchemy-backed per-user custom prompt repository."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import UTC, datetime
|
|
from typing import Any
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|
|
|
from deerflow.persistence.user_prompts.model import UserCustomPromptRow
|
|
|
|
|
|
def _row_to_dict(row: UserCustomPromptRow) -> 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 UserCustomPromptRepository:
|
|
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
|
|
self._sf = session_factory
|
|
|
|
async def list_for_user(self, user_id: str) -> list[dict[str, Any]]:
|
|
stmt = (
|
|
select(UserCustomPromptRow)
|
|
.where(UserCustomPromptRow.user_id == user_id)
|
|
.order_by(UserCustomPromptRow.sort_order.asc(), UserCustomPromptRow.created_at.asc())
|
|
)
|
|
async with self._sf() as session:
|
|
result = await session.execute(stmt)
|
|
return [_row_to_dict(row) for row in result.scalars()]
|
|
|
|
async def create(self, user_id: str, data: dict[str, Any]) -> dict[str, Any]:
|
|
now = datetime.now(UTC)
|
|
row = UserCustomPromptRow(
|
|
id=data["id"],
|
|
user_id=user_id,
|
|
title=data.get("title") or "",
|
|
content=data["content"],
|
|
sort_order=data.get("sort_order", 0),
|
|
agent_id=(data.get("agent_id") or None),
|
|
created_at=now,
|
|
updated_at=now,
|
|
)
|
|
async with self._sf() as session:
|
|
session.add(row)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _row_to_dict(row)
|
|
|
|
async def update(self, prompt_id: str, user_id: str, data: dict[str, Any]) -> dict[str, Any] | None:
|
|
async with self._sf() as session:
|
|
row = await session.get(UserCustomPromptRow, prompt_id)
|
|
if row is None or row.user_id != user_id:
|
|
return None
|
|
for key in ("title", "content", "sort_order"):
|
|
if key in data:
|
|
setattr(row, key, data[key])
|
|
# ``agent_id`` is explicitly nullable: an empty/None value clears the
|
|
# association (back to global), so accept None as a real update.
|
|
if "agent_id" in data:
|
|
row.agent_id = data["agent_id"] or None
|
|
row.updated_at = datetime.now(UTC)
|
|
await session.commit()
|
|
await session.refresh(row)
|
|
return _row_to_dict(row)
|
|
|
|
async def delete(self, prompt_id: str, user_id: str) -> bool:
|
|
async with self._sf() as session:
|
|
row = await session.get(UserCustomPromptRow, prompt_id)
|
|
if row is None or row.user_id != user_id:
|
|
return False
|
|
await session.delete(row)
|
|
await session.commit()
|
|
return True
|