refactor(repos): 统一仓库软删除过滤,优化多个查询逻辑

1. 为渠道绑定统计查询添加软删除过滤
2. 为会话列表查询添加limit参数和文档注释
3. 修复内容审核记录删除的批量删除逻辑,改用IN子句避免全表扫描风险
4. 替换datetime.utc导入为UTC常量
5. 优化用户身份仓库的更新逻辑,添加乐观锁、pending_review字段支持和合并来源查询适配
This commit is contained in:
Kris 2026-07-08 23:02:15 +08:00
parent fd885e5323
commit 63b981f130
4 changed files with 41 additions and 10 deletions

View File

@ -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())

View File

@ -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()

View File

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

View File

@ -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)