ForcePilot/backend/package/yuxi/repositories/external_systems/base.py

234 lines
8.9 KiB
Python
Raw Normal View History

"""外部系统仓库基类,提供共享的分页、转义与软删除工具。
所有 ``external_systems`` 下的仓库均继承 ``BaseRepository``复用
- ``_escape_like_pattern`` LIKE 关键词转义
- ``_resolve_pagination`` 绗页参数解析page/page_size 优先其次 limit/offset
- ``_not_deleted`` 软删除过滤条件 ``is_deleted == 0``
- ``delete_by_id`` 按主键执行软删除``is_deleted=1, deleted_at=now``的公开接口
- ``_soft_delete_by_id`` 按主键执行软删除的内部实现
- ``get_deleted_by_id`` 按主键查询软删除记录恢复前精确查询避免 list_deleted 误匹配
- ``restore_by_id`` 按主键恢复软删除记录``is_deleted=0, deleted_at=None``
- ``list_deleted`` 查询软删除记录列表 ``deleted_at`` 降序支持时间范围与分页
- ``count_deleted`` 统计软删除记录数
- ``hard_delete_by_id`` 按主键物理删除记录仅删除 ``is_deleted=1`` 的记录
- ``hard_delete_before`` 物理删除 ``deleted_at`` 早于截止时间的软删除记录
- ``count_undeleted_by_fields`` 按字段条件统计未删除记录数恢复前唯一约束预检
"""
from __future__ import annotations
from datetime import datetime
from sqlalchemy import delete as sa_delete, func, select, update as sa_update
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi.utils.datetime_utils import utc_now_naive
class BaseRepository:
"""外部系统仓库基类。
子类需设置 ``model`` 类属性为对应的 SQLAlchemy 模型即可复用软删除
过滤分页解析与 LIKE 转义等通用逻辑
"""
model: type = None # 子类覆盖为具体模型
def __init__(self, db: AsyncSession) -> None:
self.db = db
# ------------------------------------------------------------------
# 通用工具
# ------------------------------------------------------------------
@staticmethod
def _escape_like_pattern(value: str) -> str:
"""转义 SQL LIKE 特殊字符,避免关键词注入错误匹配。"""
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
@staticmethod
def _resolve_pagination(
page: int | None,
page_size: int | None,
limit: int | None,
offset: int | None,
) -> tuple[int | None, int]:
"""解析分页参数,优先使用 page/page_size其次 limit/offset。"""
if page is not None and page_size is not None:
return page_size, max(page - 1, 0) * page_size
if limit is not None:
return limit, offset or 0
return None, 0
def _not_deleted(self):
"""返回当前模型的软删除过滤条件is_deleted == 0"""
return self.model.is_deleted == 0
async def delete_by_id(
self,
record_id: int,
*,
updated_by: str | None = None,
commit: bool = True,
) -> int:
"""按主键执行软删除的公开接口,返回受影响行数。
``commit=False`` 时仅执行 ``flush()`` 使变更进入会话但不提交
便于上层在事务中组合多个写操作
"""
return await self._soft_delete_by_id(
record_id, updated_by=updated_by, commit=commit
)
async def _soft_delete_by_id(
self,
record_id: int,
*,
updated_by: str | None = None,
commit: bool = True,
) -> int:
"""通用软删除:按主键标记 is_deleted=1, deleted_at=now。返回受影响行数。"""
now = utc_now_naive()
stmt = (
sa_update(self.model)
.where(self.model.id == record_id)
.where(self.model.is_deleted == 0)
.values(is_deleted=1, deleted_at=now, updated_at=now, updated_by=updated_by)
)
result = await self.db.execute(stmt)
if commit:
await self.db.commit()
else:
await self.db.flush()
return result.rowcount
# ------------------------------------------------------------------
# 回收站(软删除记录查询 / 恢复 / 物理删除)
# ------------------------------------------------------------------
async def get_deleted_by_id(self, record_id: int):
"""按主键查询软删除记录is_deleted=1。返回 ORM 实例或 None。
用于恢复前的记录查询与唯一约束预检避免 list_deleted(limit=1) 的误匹配
"""
stmt = (
select(self.model)
.where(self.model.id == record_id)
.where(self.model.is_deleted == 1)
)
result = await self.db.execute(stmt)
return result.scalars().first()
async def restore_by_id(
self,
record_id: int,
*,
updated_by: str | None = None,
commit: bool = True,
) -> int:
"""按主键恢复软删除记录is_deleted=0, deleted_at=None。返回受影响行数。
仅恢复 is_deleted=1 的记录避免误恢复未删除资源
"""
now = utc_now_naive()
stmt = (
sa_update(self.model)
.where(self.model.id == record_id)
.where(self.model.is_deleted == 1)
.values(is_deleted=0, deleted_at=None, updated_at=now, updated_by=updated_by)
)
result = await self.db.execute(stmt)
if commit:
await self.db.commit()
else:
await self.db.flush()
return result.rowcount
async def list_deleted(
self,
*,
limit: int | None = None,
offset: int = 0,
start_time: datetime | None = None,
end_time: datetime | None = None,
) -> list:
"""查询软删除记录is_deleted=1按 deleted_at 降序排序。"""
stmt = select(self.model).where(self.model.is_deleted == 1)
if start_time is not None:
stmt = stmt.where(self.model.deleted_at >= start_time)
if end_time is not None:
stmt = stmt.where(self.model.deleted_at <= end_time)
stmt = stmt.order_by(self.model.deleted_at.desc())
if limit is not None:
stmt = stmt.limit(limit).offset(offset)
result = await self.db.execute(stmt)
return list(result.scalars().all())
async def count_deleted(self) -> int:
"""统计软删除记录数。"""
stmt = select(func.count()).select_from(self.model).where(self.model.is_deleted == 1)
result = await self.db.execute(stmt)
return int(result.scalar() or 0)
async def hard_delete_by_id(
self,
record_id: int,
*,
commit: bool = True,
) -> int:
"""按主键物理删除记录(仅删除 is_deleted=1 的记录)。返回受影响行数。"""
stmt = (
sa_delete(self.model)
.where(self.model.id == record_id)
.where(self.model.is_deleted == 1)
)
result = await self.db.execute(stmt)
if commit:
await self.db.commit()
else:
await self.db.flush()
return result.rowcount
async def hard_delete_before(
self,
cutoff_at: datetime,
*,
commit: bool = True,
) -> int:
"""物理删除 deleted_at 早于 cutoff_at 的软删除记录。返回受影响行数。"""
stmt = (
sa_delete(self.model)
.where(self.model.is_deleted == 1)
.where(self.model.deleted_at < cutoff_at)
)
result = await self.db.execute(stmt)
if commit:
await self.db.commit()
else:
await self.db.flush()
return result.rowcount
async def count_undeleted_by_fields(self, **field_filters) -> int:
"""统计满足指定字段条件的未删除记录数is_deleted=0
用于恢复前唯一约束预检各仓储的 list 方法签名不一致 system_ids vs system_id
不支持 slug 查询等无法通用调用故在基类提供基于字段映射的通用预检方法
Args:
**field_filters: 字段名 值的映射将转为 ``where(model.field == value)`` 条件
字段名必须为 ORM 模型的实际列名 slug / system_id / env_key
Example:
await repo.count_undeleted_by_fields(slug="old_crm")
await repo.count_undeleted_by_fields(system_id=1, env_key="prod")
"""
stmt = select(func.count()).select_from(self.model).where(self.model.is_deleted == 0)
for field_name, value in field_filters.items():
column = getattr(self.model, field_name, None)
if column is None:
raise ValueError(f"模型 {self.model.__name__} 无字段 {field_name}")
stmt = stmt.where(column == value)
result = await self.db.execute(stmt)
return int(result.scalar() or 0)