mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-08-15 11:04:12 +08:00
151 lines
5.2 KiB
Python
151 lines
5.2 KiB
Python
"""
|
||
ORM 基类与数据访问基类。
|
||
|
||
Base 提供声明式基类与通用的行为(字典转换、增删改查便利方法);
|
||
DbOper 是各业务 Oper 的基类,持有一个可注入的会话。
|
||
"""
|
||
from typing import Any, List, Optional, Self, Union, cast
|
||
|
||
from sqlalchemy import (CursorResult, Executable, Identity, Integer, Sequence,
|
||
and_, delete, inspect, select)
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
from sqlalchemy.orm import DeclarativeBase, Mapped, Session, declared_attr, mapped_column
|
||
|
||
from app.runtime.config import settings
|
||
from app.db.decorators import async_db_query, async_db_update, db_query, db_update
|
||
|
||
|
||
def execute_dml(db: Session, statement: Executable,
|
||
execution_options: Optional[dict] = None) -> int:
|
||
"""
|
||
执行 DML 语句并返回影响行数。
|
||
|
||
``Session.execute`` 的类型标注一律是 ``Result``,只有运行期真正拿到的
|
||
``CursorResult`` 才带 ``rowcount``——2.0 只为 ``Connection.execute`` 加了
|
||
``CursorResult`` 重载。这里把转换收口一次,免得每个模型各写一遍 cast。
|
||
:param db: 数据库会话
|
||
:param statement: delete()/update() 等 DML 语句
|
||
:param execution_options: 执行选项;不传即沿用 SQLAlchemy 默认的会话同步策略
|
||
:return: 影响行数
|
||
"""
|
||
if execution_options is None:
|
||
result = db.execute(statement)
|
||
else:
|
||
result = db.execute(statement, execution_options=execution_options)
|
||
return cast(CursorResult[Any], result).rowcount
|
||
|
||
|
||
def get_id_column() -> Mapped[int]:
|
||
"""
|
||
根据数据库类型返回合适的ID列定义
|
||
"""
|
||
if settings.DB_TYPE.lower() == "postgresql":
|
||
# PostgreSQL使用SERIAL类型,让数据库自动处理序列
|
||
return mapped_column(Integer, Identity(start=1, cycle=True), primary_key=True)
|
||
else:
|
||
# SQLite使用Sequence
|
||
return mapped_column(Integer, Sequence('id'), primary_key=True)
|
||
|
||
|
||
class Base(DeclarativeBase):
|
||
"""
|
||
声明式基类。
|
||
|
||
2.0 的声明式系统会解释类级 PEP 484 注解,未包裹在 Mapped[] 中的注解会直接报错。
|
||
仓内模型已全部迁移到 mapped_column() + Mapped[] 注解,因此不设 __allow_unmapped__:
|
||
该标志此前只为「仓外插件可能继承本 Base 自定义 legacy 注解模型」保留,插件生态
|
||
确定迭代后这条理由不再成立。留着它反而会让回流的 1.x 写法在 import 期悄悄通过,
|
||
等到运行期才以「列不存在」的形式暴露。
|
||
|
||
继承本类的模型一律使用 mapped_column() + Mapped[] 注解;确需非映射的类级属性时
|
||
用 ClassVar 显式声明,而不是把这个标志加回来。
|
||
"""
|
||
|
||
# 由 get_id_column() 在各模型中提供实际的列定义,这里只声明类型供 IDE 使用
|
||
id: Mapped[int]
|
||
|
||
@db_update
|
||
def create(self, db: Session):
|
||
db.add(self)
|
||
|
||
@async_db_update
|
||
async def async_create(self, db: AsyncSession):
|
||
db.add(self)
|
||
await db.flush()
|
||
return self
|
||
|
||
@classmethod
|
||
@db_query
|
||
def get(cls, db: Session, rid: int) -> Optional[Self]:
|
||
return db.execute(select(cls).where(and_(cls.id == rid))).scalars().first()
|
||
|
||
@classmethod
|
||
@async_db_query
|
||
async def async_get(cls, db: AsyncSession, rid: int) -> Optional[Self]:
|
||
result = await db.execute(select(cls).where(and_(cls.id == rid)))
|
||
return result.scalars().first()
|
||
|
||
@db_update
|
||
def update(self, db: Session, payload: dict):
|
||
for key, value in payload.items():
|
||
setattr(self, key, value)
|
||
if inspect(self).detached:
|
||
db.add(self)
|
||
|
||
@async_db_update
|
||
async def async_update(self, db: AsyncSession, payload: dict):
|
||
for key, value in payload.items():
|
||
setattr(self, key, value)
|
||
if inspect(self).detached:
|
||
db.add(self)
|
||
|
||
@classmethod
|
||
@db_update
|
||
def delete(cls, db: Session, rid):
|
||
db.execute(delete(cls).where(and_(cls.id == rid)))
|
||
|
||
@classmethod
|
||
@async_db_update
|
||
async def async_delete(cls, db: AsyncSession, rid):
|
||
result = await db.execute(select(cls).where(and_(cls.id == rid)))
|
||
user = result.scalars().first()
|
||
if user:
|
||
await db.delete(user)
|
||
|
||
@classmethod
|
||
@db_update
|
||
def truncate(cls, db: Session):
|
||
db.execute(delete(cls))
|
||
|
||
@classmethod
|
||
@async_db_update
|
||
async def async_truncate(cls, db: AsyncSession):
|
||
await db.execute(delete(cls))
|
||
|
||
@classmethod
|
||
@db_query
|
||
def list(cls, db: Session) -> List[Self]:
|
||
return list(db.execute(select(cls)).scalars().all())
|
||
|
||
@classmethod
|
||
@async_db_query
|
||
async def async_list(cls, db: AsyncSession) -> List[Self]:
|
||
result = await db.execute(select(cls))
|
||
return list(result.scalars().all())
|
||
|
||
def to_dict(self):
|
||
return {c.name: getattr(self, c.name, None) for c in self.__table__.columns} # noqa
|
||
|
||
@declared_attr.directive
|
||
def __tablename__(cls) -> str: # noqa: N805 declared_attr 的第一个参数即类本身
|
||
return cls.__name__.lower()
|
||
|
||
|
||
class DbOper:
|
||
"""
|
||
数据库操作基类
|
||
"""
|
||
|
||
def __init__(self, db: Optional[Union[Session, AsyncSession]] = None):
|
||
self._db = db
|