Files
2026-09-08 19:17:35 +08:00

48 lines
1.6 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""MySQL 通用仓储基类:封装 CRUD(AsyncSession 会话注入,函数内自提交)。"""
from __future__ import annotations
from typing import Any, Sequence
from sqlalchemy import delete as sa_delete
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
class BaseRepository:
"""子类只需设 model;主键列固定名为 id。"""
model: type | None = None
def __init__(self, db: AsyncSession):
self.db = db
async def get(self, pk: int) -> Any | None:
return await self.db.get(self.model, pk)
async def list(self, *, where: Sequence | None = None,
order_by: Any | None = None, limit: int = 100,
offset: int = 0) -> list:
stmt = select(self.model)
if where:
stmt = stmt.where(*where)
if order_by is not None:
stmt = stmt.order_by(order_by)
return list((await self.db.scalars(stmt.limit(limit).offset(offset))).all())
async def add(self, obj: Any) -> Any:
"""新增并入参刷新出 id 等 DB 生成字段。"""
self.db.add(obj)
await self.db.commit()
await self.db.refresh(obj)
return obj
async def delete(self, pk: int) -> bool:
result = await self.db.execute(sa_delete(self.model).where(self.model.id == pk))
await self.db.commit()
return result.rowcount > 0
async def count(self, *, where: Sequence | None = None) -> int:
stmt = select(func.count()).select_from(self.model)
if where:
stmt = stmt.where(*where)
return (await self.db.scalar(stmt)) or 0