refactor(repos): 统一仓库软删除过滤,优化多个查询逻辑
1. 为渠道绑定统计查询添加软删除过滤 2. 为会话列表查询添加limit参数和文档注释 3. 修复内容审核记录删除的批量删除逻辑,改用IN子句避免全表扫描风险 4. 替换datetime.utc导入为UTC常量 5. 优化用户身份仓库的更新逻辑,添加乐观锁、pending_review字段支持和合并来源查询适配
This commit is contained in:
parent
fd885e5323
commit
63b981f130
@ -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())
|
||||
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user