diff --git a/backend/package/yuxi/repositories/channels/channel_session_repository.py b/backend/package/yuxi/repositories/channels/channel_session_repository.py index cbccd430..2cb4fb60 100644 --- a/backend/package/yuxi/repositories/channels/channel_session_repository.py +++ b/backend/package/yuxi/repositories/channels/channel_session_repository.py @@ -299,14 +299,27 @@ class ChannelSessionRepository(BaseRepository): result = await self.db.execute(stmt) return list(result.scalars().all()) - async def list_by_conversation(self, conversation_id: int) -> list[ChannelSession]: - """按内部会话 ID 反查渠道会话(FR-06 跨渠道关联)。""" + async def list_by_conversation( + self, + conversation_id: int, + *, + limit: int | None = None, + ) -> list[ChannelSession]: + """按内部会话 ID 反查渠道会话(FR-06 跨渠道关联)。 + + ``limit`` 提供时附加 ``LIMIT`` 子句,供仅需首条会话的调用方 + (如 ``getSessionOwner`` / ``transferSessionOwner`` / + ``isTemporarySession``)避免加载全部会话;为 ``None`` 时返回全部 + (供 ``mergeConversations`` 等需要枚举全部会话的场景)。 + """ stmt = ( select(ChannelSession) .where(ChannelSession.conversation_id == conversation_id) .where(self._not_deleted()) .order_by(ChannelSession.created_at.desc()) ) + if limit is not None: + stmt = stmt.limit(limit) result = await self.db.execute(stmt) return list(result.scalars().all()) diff --git a/backend/package/yuxi/repositories/channels/content_review_record_repository.py b/backend/package/yuxi/repositories/channels/content_review_record_repository.py index c494eda5..7bbcbba9 100644 --- a/backend/package/yuxi/repositories/channels/content_review_record_repository.py +++ b/backend/package/yuxi/repositories/channels/content_review_record_repository.py @@ -3,7 +3,7 @@ from __future__ import annotations -from datetime import datetime, timezone +from datetime import UTC, datetime from typing import Any @@ -92,7 +92,7 @@ class ChannelContentReviewRecordRepository(BaseRepository): orm.verdict = verdict orm.reviewer = reviewer orm.updated_by = reviewer - orm.decision_completed_at = datetime.now(timezone.utc).replace(tzinfo=None) + orm.decision_completed_at = datetime.now(UTC).replace(tzinfo=None) if commit: await self.db.commit() else: @@ -120,7 +120,10 @@ class ChannelContentReviewRecordRepository(BaseRepository): Returns: 被删除的记录数。 """ - stmt = sa_delete(ChannelContentReviewRecord).where(ChannelContentReviewRecord.reviewed_at < before).limit(limit) + id_select = ( + select(ChannelContentReviewRecord.id).where(ChannelContentReviewRecord.reviewed_at < before).limit(limit) + ) + stmt = sa_delete(ChannelContentReviewRecord).where(ChannelContentReviewRecord.id.in_(id_select)) result = await self.db.execute(stmt) if commit: await self.db.commit() diff --git a/backend/package/yuxi/repositories/channels/route_binding_repository.py b/backend/package/yuxi/repositories/channels/route_binding_repository.py index df4ff827..c4fc448a 100644 --- a/backend/package/yuxi/repositories/channels/route_binding_repository.py +++ b/backend/package/yuxi/repositories/channels/route_binding_repository.py @@ -112,7 +112,7 @@ class ChannelRouteBindingRepository(BaseRepository): filter: RouteBindingFilter, ) -> int: """按过滤条件统计总数(始终排除已软删除)。""" - stmt = select(func.count(ChannelRouteBinding.id)) + stmt = select(func.count(ChannelRouteBinding.id)).where(self._not_deleted()) stmt = self._apply_filter_to_stmt(stmt, filter) result = await self.db.execute(stmt) return result.scalar() or 0 diff --git a/backend/package/yuxi/repositories/channels/user_identity_repository.py b/backend/package/yuxi/repositories/channels/user_identity_repository.py index c7cf9d6b..0727320d 100644 --- a/backend/package/yuxi/repositories/channels/user_identity_repository.py +++ b/backend/package/yuxi/repositories/channels/user_identity_repository.py @@ -132,6 +132,7 @@ class UserIdentityRepository(BaseRepository): "channel_bindings", "merged_from", "confidence", + "pending_review", "updated_by", ) for field in updatable_fields: @@ -149,20 +150,28 @@ class UserIdentityRepository(BaseRepository): self, record_id: int, channel_bindings: dict[str, Any], - merged_from: list[str], + merged_from: list[dict], + expected_version: int, *, updated_by: str | None = None, commit: bool = True, ) -> int: - """按主键更新渠道绑定与合并来源(FR-07 合并场景,不加载 ORM 实例)。""" + """按主键更新渠道绑定与合并来源(FR-07 合并场景,不加载 ORM 实例)。 + + 通过 ``WHERE version = expected_version`` 实现乐观锁,版本不匹配 + (并发修改或记录不存在)时返回 ``rowcount=0``,由调用方翻译为 + ``ConflictError``。 + """ now = utc_now_naive() stmt = ( update(ChannelUserIdentity) .where(ChannelUserIdentity.id == record_id) + .where(ChannelUserIdentity.version == expected_version) .where(self._not_deleted()) .values( channel_bindings=channel_bindings, merged_from=merged_from, + version=expected_version + 1, updated_at=now, updated_by=updated_by, ) @@ -209,6 +218,8 @@ class UserIdentityRepository(BaseRepository): values["user_id"] = updates["user_id"] if "identity_type" in updates: values["identity_type"] = updates["identity_type"] + if "pending_review" in updates: + values["pending_review"] = updates["pending_review"] if updates.get("updated_by") is not None: values["updated_by"] = updates["updated_by"] stmt = ( @@ -264,10 +275,14 @@ class UserIdentityRepository(BaseRepository): return list(result.scalars().all()) async def list_by_merged_from(self, identity_id: str) -> list[ChannelUserIdentity]: - """查询合并来源包含指定 identity_id 的记录(FR-07 合并回滚)。""" + """查询合并来源包含指定 identity_id 的记录(FR-07 合并回滚)。 + + merged_from 为结构化 dict 列表(H-21),每项含 ``identity_id`` 键, + 通过 JSONB containment ``@>`` 查询匹配 ``[{"identity_id": ...}]``。 + """ stmt = ( select(ChannelUserIdentity) - .where(ChannelUserIdentity.merged_from.contains([identity_id])) + .where(ChannelUserIdentity.merged_from.contains([{"identity_id": identity_id}])) .where(self._not_deleted()) ) result = await self.db.execute(stmt)