mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-19 02:56:17 +00:00
Group them by bounded context instead of by technology, one file per
port, and align the directory name with the AWS Prescriptive Guidance
layout (entrypoints / domain-with-ports / adapters).
app/infra/persistence/feedback.py
-> app/adapters/feedback/feedback_repository.py owned persistence
-> app/adapters/feedback/run_lookup.py anti-corruption layer
`persistence/` promised a technology-first classification that its own
contents contradicted: RunStoreRunLookup lived there while its docstring
said "no new SQL". Splitting per port makes that distinction structural.
SqlFeedbackRepository and _tz_aware move unchanged -- verified by
comparing their AST against the original rather than by eye. run_lookup.py
additionally gains a RunStore annotation behind TYPE_CHECKING (the module
is imported lazily by the composition root, so this keeps the runtime
import cost at zero), a docstring stating that this context owns no table
and writes no SQL against it, and a TODO recording the condition under
which the body is replaced: when the run context publishes a contract of
its own, the RunLookup port itself does not move.
Each module docstring opens with a fixed marker so the two kinds of
secondary adapter stay greppable:
grep -rl "anti-corruption layer" app/adapters/
Filenames deliberately carry no sql_ / acl_ prefix: a prefix encodes an
implementation property, so switching storage would force a rename even
though the port -- and therefore the import path -- has not changed. The
class name already carries it. A prefix earns its place once one port has
several production implementations, which is not yet the case here.
app/infra/ held nothing else and is removed.
144 lines
6.1 KiB
Python
144 lines
6.1 KiB
Python
"""Secondary adapter (owned persistence) -- the feedback table in SQL.
|
|
|
|
Implements ``FeedbackRepository`` from ``deerflow.domain.feedback.ports``.
|
|
This context owns the ``feedback`` table and writes its own queries, so
|
|
SQL/ORM vocabulary stops at this file: methods exchange domain objects,
|
|
translate ``IntegrityError`` into domain errors, and normalize SQLite's
|
|
tz-naive reads. Queries were migrated unchanged from the legacy repository
|
|
(now removed) to preserve behavior.
|
|
|
|
Sibling adapter: ``run_lookup.py`` serves the same context but owns no
|
|
table and writes no SQL -- see its docstring for why that distinction is
|
|
worth keeping visible.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import UTC, datetime
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.exc import IntegrityError
|
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
|
|
|
|
from deerflow.domain.feedback.model import DuplicateFeedbackError, Feedback
|
|
from deerflow.domain.feedback.ports import FeedbackRepository
|
|
|
|
# Transitional: the ORM row stays in the harness until PR-N moves engine,
|
|
# models, and migrations into app/adapters.
|
|
from deerflow.persistence.feedback.model import FeedbackRow
|
|
|
|
|
|
def _tz_aware(value: datetime) -> datetime:
|
|
"""SQLite drops tzinfo on read; stored values are always UTC."""
|
|
return value if value.tzinfo is not None else value.replace(tzinfo=UTC)
|
|
|
|
|
|
class SqlFeedbackRepository(FeedbackRepository):
|
|
"""SQL implementation of the ``FeedbackRepository`` port.
|
|
|
|
Explicit inheritance is a readability aid only: a missing method would
|
|
still instantiate fine (Protocol bodies are inherited), so the contract
|
|
test suite must cover every port method.
|
|
"""
|
|
|
|
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
|
|
self._session_factory = session_factory
|
|
|
|
@staticmethod
|
|
def _to_domain(row: FeedbackRow) -> Feedback:
|
|
"""ORM row -> aggregate. The only place reads are normalized."""
|
|
return Feedback(
|
|
feedback_id=row.feedback_id,
|
|
run_id=row.run_id,
|
|
thread_id=row.thread_id,
|
|
rating=row.rating,
|
|
user_id=row.user_id,
|
|
message_id=row.message_id,
|
|
comment=row.comment,
|
|
tags=tuple(row.tags or ()),
|
|
created_at=_tz_aware(row.created_at),
|
|
)
|
|
|
|
@staticmethod
|
|
def _to_row(feedback: Feedback) -> FeedbackRow:
|
|
"""Aggregate -> ORM row. Explicit field list: new columns stay
|
|
private until deliberately mapped here."""
|
|
return FeedbackRow(
|
|
feedback_id=feedback.feedback_id,
|
|
run_id=feedback.run_id,
|
|
thread_id=feedback.thread_id,
|
|
user_id=feedback.user_id,
|
|
message_id=feedback.message_id,
|
|
rating=feedback.rating,
|
|
comment=feedback.comment,
|
|
tags=list(feedback.tags) or None,
|
|
created_at=feedback.created_at,
|
|
)
|
|
|
|
async def save(self, feedback: Feedback) -> Feedback:
|
|
# Migrated from legacy ``upsert``: look up by aggregate identity,
|
|
# then update in place or insert.
|
|
async with self._session_factory() as session:
|
|
stmt = select(FeedbackRow).where(
|
|
FeedbackRow.thread_id == feedback.thread_id,
|
|
FeedbackRow.run_id == feedback.run_id,
|
|
FeedbackRow.user_id == feedback.user_id,
|
|
)
|
|
row = (await session.execute(stmt)).scalar_one_or_none()
|
|
if row is not None:
|
|
row.rating = feedback.rating
|
|
row.comment = feedback.comment
|
|
row.tags = list(feedback.tags) or None
|
|
row.created_at = feedback.created_at
|
|
else:
|
|
row = self._to_row(feedback)
|
|
session.add(row)
|
|
try:
|
|
await session.commit()
|
|
except IntegrityError as exc:
|
|
# Two concurrent upserts can both miss the lookup and insert;
|
|
# the loser hits the unique constraint. Translate instead of
|
|
# leaking the driver exception (legacy code returned a 500).
|
|
raise DuplicateFeedbackError(f"concurrent feedback upsert for run {feedback.run_id}") from exc
|
|
await session.refresh(row)
|
|
return self._to_domain(row)
|
|
|
|
async def latest_per_run_in_thread(self, thread_id: str, *, user_id: str | None) -> dict[str, Feedback]:
|
|
# Migrated from legacy ``list_by_thread_grouped``. With an ownership
|
|
# filter the unique constraint guarantees at most one row per run.
|
|
stmt = select(FeedbackRow).where(FeedbackRow.thread_id == thread_id)
|
|
if user_id is not None:
|
|
stmt = stmt.where(FeedbackRow.user_id == user_id)
|
|
async with self._session_factory() as session:
|
|
result = await session.execute(stmt)
|
|
return {row.run_id: self._to_domain(row) for row in result.scalars()}
|
|
|
|
async def latest_for_runs(self, thread_id: str, run_ids: set[str], *, user_id: str | None) -> dict[str, Feedback]:
|
|
# Migrated from legacy ``list_by_run_ids`` (paged message list).
|
|
if not run_ids:
|
|
return {}
|
|
stmt = select(FeedbackRow).where(
|
|
FeedbackRow.thread_id == thread_id,
|
|
FeedbackRow.run_id.in_(run_ids),
|
|
)
|
|
if user_id is not None:
|
|
stmt = stmt.where(FeedbackRow.user_id == user_id)
|
|
async with self._session_factory() as session:
|
|
result = await session.execute(stmt)
|
|
return {row.run_id: self._to_domain(row) for row in result.scalars()}
|
|
|
|
async def remove_for_run(self, thread_id: str, run_id: str, *, user_id: str | None) -> bool:
|
|
# Migrated from legacy ``delete_by_run``.
|
|
async with self._session_factory() as session:
|
|
stmt = select(FeedbackRow).where(
|
|
FeedbackRow.thread_id == thread_id,
|
|
FeedbackRow.run_id == run_id,
|
|
FeedbackRow.user_id == user_id,
|
|
)
|
|
row = (await session.execute(stmt)).scalar_one_or_none()
|
|
if row is None:
|
|
return False
|
|
await session.delete(row)
|
|
await session.commit()
|
|
return True
|