mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-04 23:17:20 +08:00
refactor: isolate media server and site userdata queries
This commit is contained in:
@@ -7,7 +7,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from app.db.base import Base, execute_dml, get_id_column
|
||||
from app.db.decorators import async_db_query, db_query
|
||||
from app.db.decorators import legacy_async_db_query, legacy_db_query
|
||||
from app.db.models._constraints import media_identity_constraint
|
||||
from app.schemas.types import MediaSource
|
||||
|
||||
@@ -53,12 +53,12 @@ class MediaServerItem(Base):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def get_by_itemid(cls, db: Session, item_id: str):
|
||||
return db.execute(select(cls).where(cls.item_id == item_id)).scalars().first()
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def get_by_server_itemid(cls, db: Session, server: str, item_id: str):
|
||||
return db.execute(
|
||||
select(cls).where(cls.server == server, cls.item_id == item_id)
|
||||
@@ -97,7 +97,7 @@ class MediaServerItem(Base):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def exist_by_media_identity(
|
||||
cls, db: Session, media_source: MediaSource, media_id: str, mtype: str,
|
||||
):
|
||||
@@ -109,7 +109,7 @@ class MediaServerItem(Base):
|
||||
)).scalars().first()
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def exists_by_title(cls, db: Session, title: str, mtype: str, year: str):
|
||||
statement = select(cls).where(cls.title == title)
|
||||
if mtype:
|
||||
@@ -119,13 +119,13 @@ class MediaServerItem(Base):
|
||||
return db.execute(statement).scalars().first()
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
@legacy_async_db_query
|
||||
async def async_get_by_itemid(cls, db: AsyncSession, item_id: str):
|
||||
result = await db.execute(select(cls).filter(cls.item_id == item_id))
|
||||
return result.scalars().first()
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
@legacy_async_db_query
|
||||
async def async_exist_by_media_identity(
|
||||
cls, db: AsyncSession, media_source: MediaSource, media_id: str, mtype: str,
|
||||
):
|
||||
@@ -138,7 +138,7 @@ class MediaServerItem(Base):
|
||||
return result.scalars().first()
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
@legacy_async_db_query
|
||||
async def async_exists_by_title(cls, db: AsyncSession, title: str, mtype: str, year: str):
|
||||
if not mtype and not year:
|
||||
result = await db.execute(select(cls).filter(cls.title == title))
|
||||
|
||||
@@ -6,7 +6,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from app.db.base import Base, execute_dml, get_id_column
|
||||
from app.db.decorators import async_db_query, db_query
|
||||
from app.db.decorators import legacy_async_db_query, legacy_db_query
|
||||
|
||||
|
||||
class SiteUserData(Base):
|
||||
@@ -61,7 +61,7 @@ class SiteUserData(Base):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def get_by_domain(cls, db: Session, domain: str, workdate: Optional[str] = None, worktime: Optional[str] = None):
|
||||
statement = select(cls).where(cls.domain == domain)
|
||||
if workdate and worktime:
|
||||
@@ -72,7 +72,7 @@ class SiteUserData(Base):
|
||||
return list(db.execute(statement).scalars().all())
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
@legacy_async_db_query
|
||||
async def async_get_by_domain(cls, db: AsyncSession, domain: str, workdate: Optional[str] = None, worktime: Optional[str] = None):
|
||||
query = select(cls).filter(cls.domain == domain)
|
||||
if workdate and worktime:
|
||||
@@ -83,12 +83,12 @@ class SiteUserData(Base):
|
||||
return list(result.scalars().all())
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def get_by_date(cls, db: Session, date: str):
|
||||
return list(db.execute(select(cls).where(cls.updated_day == date)).scalars().all())
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def get_latest(cls, db: Session):
|
||||
"""
|
||||
获取各站点最新一天的数据
|
||||
@@ -113,7 +113,7 @@ class SiteUserData(Base):
|
||||
).scalars().all())
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
@legacy_async_db_query
|
||||
async def async_get_latest(cls, db: AsyncSession):
|
||||
"""
|
||||
异步获取各站点最新一天的数据
|
||||
|
||||
+48
-18
@@ -1,5 +1,6 @@
|
||||
from typing import Optional
|
||||
from typing import Optional, Union
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.db.base import DbOper
|
||||
@@ -11,7 +12,11 @@ class MediaServerOper(DbOper):
|
||||
媒体服务器数据管理
|
||||
"""
|
||||
|
||||
def __init__(self, db: Optional[Session] = None):
|
||||
def __init__(
|
||||
self,
|
||||
db: Optional[Union[Session, AsyncSession]] = None,
|
||||
) -> None:
|
||||
"""保存调用方提供的同步或异步查询会话。"""
|
||||
super().__init__(db)
|
||||
|
||||
@staticmethod
|
||||
@@ -34,7 +39,12 @@ class MediaServerOper(DbOper):
|
||||
if not server or not item_id:
|
||||
return False
|
||||
item = MediaServerItem(**kwargs)
|
||||
if not item.get_by_server_itemid(self._db, server, item_id):
|
||||
existing = self._execute_sync_query(
|
||||
lambda session: MediaServerItem.get_by_server_itemid(
|
||||
session, server, item_id
|
||||
)
|
||||
)
|
||||
if not existing:
|
||||
self._stage_create(item)
|
||||
return True
|
||||
return False
|
||||
@@ -49,7 +59,11 @@ class MediaServerOper(DbOper):
|
||||
if not server or not item_id:
|
||||
return False
|
||||
|
||||
item = MediaServerItem.get_by_server_itemid(self._db, server, item_id)
|
||||
item = self._execute_sync_query(
|
||||
lambda session: MediaServerItem.get_by_server_itemid(
|
||||
session, server, item_id
|
||||
)
|
||||
)
|
||||
if item:
|
||||
self._stage_update(item, kwargs)
|
||||
return False
|
||||
@@ -93,16 +107,24 @@ class MediaServerOper(DbOper):
|
||||
判断媒体服务器数据是否存在
|
||||
"""
|
||||
if kwargs.get("media_source") and kwargs.get("media_id"):
|
||||
item = MediaServerItem.exist_by_media_identity(
|
||||
self._db,
|
||||
media_source=kwargs.get("media_source"),
|
||||
media_id=kwargs.get("media_id"),
|
||||
mtype=kwargs.get("mtype"),
|
||||
item = self._execute_sync_query(
|
||||
lambda session: MediaServerItem.exist_by_media_identity(
|
||||
session,
|
||||
media_source=kwargs.get("media_source"),
|
||||
media_id=kwargs.get("media_id"),
|
||||
mtype=kwargs.get("mtype"),
|
||||
)
|
||||
)
|
||||
elif kwargs.get("title"):
|
||||
# 按标题、类型、年份查
|
||||
item = MediaServerItem.exists_by_title(self._db, title=kwargs.get("title"),
|
||||
mtype=kwargs.get("mtype"), year=kwargs.get("year"))
|
||||
item = self._execute_sync_query(
|
||||
lambda session: MediaServerItem.exists_by_title(
|
||||
session,
|
||||
title=kwargs.get("title"),
|
||||
mtype=kwargs.get("mtype"),
|
||||
year=kwargs.get("year"),
|
||||
)
|
||||
)
|
||||
else:
|
||||
return None
|
||||
if not item:
|
||||
@@ -122,16 +144,24 @@ class MediaServerOper(DbOper):
|
||||
异步判断媒体服务器数据是否存在
|
||||
"""
|
||||
if kwargs.get("media_source") and kwargs.get("media_id"):
|
||||
item = await MediaServerItem.async_exist_by_media_identity(
|
||||
self._db,
|
||||
media_source=kwargs.get("media_source"),
|
||||
media_id=kwargs.get("media_id"),
|
||||
mtype=kwargs.get("mtype"),
|
||||
item = await self._execute_async_query(
|
||||
lambda session: MediaServerItem.async_exist_by_media_identity(
|
||||
session,
|
||||
media_source=kwargs.get("media_source"),
|
||||
media_id=kwargs.get("media_id"),
|
||||
mtype=kwargs.get("mtype"),
|
||||
)
|
||||
)
|
||||
elif kwargs.get("title"):
|
||||
# 按标题、类型、年份查
|
||||
item = await MediaServerItem.async_exists_by_title(self._db, title=kwargs.get("title"),
|
||||
mtype=kwargs.get("mtype"), year=kwargs.get("year"))
|
||||
item = await self._execute_async_query(
|
||||
lambda session: MediaServerItem.async_exists_by_title(
|
||||
session,
|
||||
title=kwargs.get("title"),
|
||||
mtype=kwargs.get("mtype"),
|
||||
year=kwargs.get("year"),
|
||||
)
|
||||
)
|
||||
else:
|
||||
return None
|
||||
if not item:
|
||||
|
||||
+32
-8
@@ -280,7 +280,13 @@ class SiteOper(DbOper):
|
||||
"err_msg": payload.get("err_msg") or ""
|
||||
})
|
||||
# 按站点+天判断是否存在数据
|
||||
siteuserdatas = SiteUserData.get_by_domain(self._db, domain=domain, workdate=current_day)
|
||||
siteuserdatas = self._execute_sync_query(
|
||||
lambda session: SiteUserData.get_by_domain(
|
||||
session,
|
||||
domain=domain,
|
||||
workdate=current_day,
|
||||
)
|
||||
)
|
||||
if siteuserdatas:
|
||||
# 存在则更新
|
||||
if not payload.get("err_msg"):
|
||||
@@ -294,13 +300,21 @@ class SiteOper(DbOper):
|
||||
"""
|
||||
获取站点用户数据
|
||||
"""
|
||||
return SiteUserData.list(self._db)
|
||||
return self._execute_sync_query(
|
||||
lambda session: SiteUserData.list(session)
|
||||
)
|
||||
|
||||
def get_userdata_by_domain(self, domain: str, workdate: Optional[str] = None) -> List[SiteUserData]:
|
||||
"""
|
||||
获取站点用户数据
|
||||
"""
|
||||
return SiteUserData.get_by_domain(self._db, domain=domain, workdate=workdate)
|
||||
return self._execute_sync_query(
|
||||
lambda session: SiteUserData.get_by_domain(
|
||||
session,
|
||||
domain=domain,
|
||||
workdate=workdate,
|
||||
)
|
||||
)
|
||||
|
||||
async def async_get_userdata_by_domain(
|
||||
self, domain: str, workdate: Optional[str] = None
|
||||
@@ -308,13 +322,19 @@ class SiteOper(DbOper):
|
||||
"""
|
||||
异步获取站点用户数据。
|
||||
"""
|
||||
return await SiteUserData.async_get_by_domain(
|
||||
self._db, domain=domain, workdate=workdate
|
||||
return await self._execute_async_query(
|
||||
lambda session: SiteUserData.async_get_by_domain(
|
||||
session,
|
||||
domain=domain,
|
||||
workdate=workdate,
|
||||
)
|
||||
)
|
||||
|
||||
async def async_get_userdata_latest(self) -> List[SiteUserData]:
|
||||
"""异步获取各站点最新用户数据。"""
|
||||
return await SiteUserData.async_get_latest(self._db)
|
||||
return await self._execute_async_query(
|
||||
lambda session: SiteUserData.async_get_latest(session)
|
||||
)
|
||||
|
||||
async def async_get_icon_by_domain(self, domain: str) -> Optional[SiteIcon]:
|
||||
"""异步按域名获取站点图标。"""
|
||||
@@ -335,13 +355,17 @@ class SiteOper(DbOper):
|
||||
"""
|
||||
获取站点用户数据
|
||||
"""
|
||||
return SiteUserData.get_by_date(self._db, date)
|
||||
return self._execute_sync_query(
|
||||
lambda session: SiteUserData.get_by_date(session, date)
|
||||
)
|
||||
|
||||
def get_userdata_latest(self) -> List[SiteUserData]:
|
||||
"""
|
||||
获取站点最新数据
|
||||
"""
|
||||
return SiteUserData.get_latest(self._db)
|
||||
return self._execute_sync_query(
|
||||
lambda session: SiteUserData.get_latest(session)
|
||||
)
|
||||
|
||||
def get_icon_by_domain(self, domain: str) -> Optional[SiteIcon]:
|
||||
"""
|
||||
|
||||
Reference in New Issue
Block a user