refactor: make model sessions explicit

This commit is contained in:
jxxghp
2026-08-23 23:33:07 +08:00
parent 820582ab12
commit 6e69258e3c
65 changed files with 1299 additions and 2010 deletions
+12 -31
View File
@@ -4,7 +4,6 @@ from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import Mapped, Session, mapped_column
from app.db.base import get_id_column, Base
from app.db.decorators import legacy_async_db_query, legacy_db_query
class PluginData(Base):
@@ -21,44 +20,32 @@ class PluginData(Base):
)
@classmethod
@legacy_db_query
def get_plugin_data(cls, db: Session | None = None, plugin_id: str | None = None):
"""在调用方 Session 中读取插件全部数据,并兼容旧无会话入口。"""
if plugin_id is None:
raise TypeError("plugin_id is required")
def get_plugin_data(cls, db: Session, plugin_id: str):
"""在调用方 Session 中读取插件全部数据。"""
return list(db.execute(select(cls).where(cls.plugin_id == plugin_id)).scalars().all())
@classmethod
@legacy_async_db_query
async def async_get_plugin_data(
cls, db: AsyncSession | None = None, plugin_id: str | None = None
cls, db: AsyncSession, plugin_id: str
):
"""在调用方 AsyncSession 中读取插件全部数据,并兼容旧无会话入口"""
if plugin_id is None:
raise TypeError("plugin_id is required")
"""在调用方 AsyncSession 中读取插件全部数据。"""
result = await db.execute(select(cls).where(cls.plugin_id == plugin_id))
return list(result.scalars().all())
@classmethod
@legacy_db_query
def get_plugin_data_by_key(
cls, db: Session | None = None, plugin_id: str | None = None, key: str | None = None
cls, db: Session, plugin_id: str, key: str
):
"""在调用方 Session 中按键读取插件数据,并兼容旧无会话入口"""
if plugin_id is None or key is None:
raise TypeError("plugin_id and key are required")
"""在调用方 Session 中按键读取插件数据。"""
return db.execute(
select(cls).where(cls.plugin_id == plugin_id, cls.key == key)
).scalars().first()
@classmethod
@legacy_async_db_query
async def async_get_plugin_data_by_key(
cls, db: AsyncSession | None = None, plugin_id: str | None = None, key: str | None = None
cls, db: AsyncSession, plugin_id: str, key: str
):
"""在调用方 AsyncSession 中按键读取插件数据,并兼容旧无会话入口"""
if plugin_id is None or key is None:
raise TypeError("plugin_id and key are required")
"""在调用方 AsyncSession 中按键读取插件数据。"""
result = await db.execute(
select(cls).where(cls.plugin_id == plugin_id, cls.key == key)
)
@@ -75,22 +62,16 @@ class PluginData(Base):
db.execute(delete(cls).where(cls.plugin_id == plugin_id))
@classmethod
@legacy_db_query
def get_plugin_data_by_plugin_id(
cls, db: Session | None = None, plugin_id: str | None = None
cls, db: Session, plugin_id: str
):
"""在调用方 Session 中按插件 ID 读取数据,并兼容旧无会话入口"""
if plugin_id is None:
raise TypeError("plugin_id is required")
"""在调用方 Session 中按插件 ID 读取数据。"""
return list(db.execute(select(cls).where(cls.plugin_id == plugin_id)).scalars().all())
@classmethod
@legacy_async_db_query
async def async_get_plugin_data_by_plugin_id(
cls, db: AsyncSession | None = None, plugin_id: str | None = None
cls, db: AsyncSession, plugin_id: str
):
"""在调用方 AsyncSession 中按插件 ID 读取数据,并兼容旧无会话入口"""
if plugin_id is None:
raise TypeError("plugin_id is required")
"""在调用方 AsyncSession 中按插件 ID 读取数据。"""
result = await db.execute(select(cls).where(cls.plugin_id == plugin_id))
return list(result.scalars().all())