mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-05 23:47:41 +08:00
refactor: isolate agent history queries
This commit is contained in:
@@ -5,7 +5,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from app.db.base import Base, 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 AgentChat(Base):
|
||||
@@ -50,7 +50,7 @@ class AgentChat(Base):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def get_by_session(
|
||||
cls, db: Session, session_id: str, user_id: Optional[str] = None
|
||||
) -> Optional["AgentChat"]:
|
||||
@@ -63,7 +63,7 @@ class AgentChat(Base):
|
||||
return db.execute(statement.order_by(cls.id.desc())).scalars().first()
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
@legacy_async_db_query
|
||||
async def async_get_by_session(
|
||||
cls, db: AsyncSession, session_id: str, user_id: Optional[str] = None
|
||||
) -> Optional["AgentChat"]:
|
||||
@@ -77,7 +77,7 @@ class AgentChat(Base):
|
||||
return result.scalars().first()
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def list_by_page(
|
||||
cls,
|
||||
db: Session,
|
||||
@@ -103,7 +103,7 @@ class AgentChat(Base):
|
||||
).scalars().all())
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
@legacy_async_db_query
|
||||
async def async_list_by_page(
|
||||
cls,
|
||||
db: AsyncSession,
|
||||
|
||||
@@ -4,7 +4,7 @@ from sqlalchemy import Index, Integer, String, Text, delete, select, update
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from app.db.base import Base, execute_dml, get_id_column
|
||||
from app.db.decorators import db_query
|
||||
from app.db.decorators import legacy_db_query
|
||||
from app.db.models.agenttask import AgentTask
|
||||
|
||||
|
||||
@@ -249,7 +249,7 @@ class AgentTaskRun(Base):
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def get_by_run_id(
|
||||
cls,
|
||||
db: Session,
|
||||
@@ -261,7 +261,7 @@ class AgentTaskRun(Base):
|
||||
).scalars().first()
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
@legacy_db_query
|
||||
def list_for_task(
|
||||
cls,
|
||||
db: Session,
|
||||
|
||||
@@ -77,7 +77,9 @@ class AgentChatOper(DbOper):
|
||||
"""
|
||||
获取 Agent 会话。
|
||||
"""
|
||||
return AgentChat.get_by_session(self._db, session_id, user_id)
|
||||
return self._execute_sync_query(
|
||||
lambda session: AgentChat.get_by_session(session, session_id, user_id)
|
||||
)
|
||||
|
||||
async def async_get(
|
||||
self, session_id: str, user_id: Optional[str] = None
|
||||
@@ -85,7 +87,9 @@ class AgentChatOper(DbOper):
|
||||
"""
|
||||
异步获取 Agent 会话。
|
||||
"""
|
||||
return await AgentChat.async_get_by_session(self._db, session_id, user_id)
|
||||
return await self._execute_async_query(
|
||||
lambda session: AgentChat.async_get_by_session(session, session_id, user_id)
|
||||
)
|
||||
|
||||
def ensure_session(
|
||||
self,
|
||||
@@ -295,12 +299,14 @@ class AgentChatOper(DbOper):
|
||||
"""
|
||||
异步分页获取 Agent 会话历史。
|
||||
"""
|
||||
return await AgentChat.async_list_by_page(
|
||||
self._db,
|
||||
page=page,
|
||||
count=count,
|
||||
user_id=user_id,
|
||||
username=username,
|
||||
return await self._execute_async_query(
|
||||
lambda session: AgentChat.async_list_by_page(
|
||||
session,
|
||||
page=page,
|
||||
count=count,
|
||||
user_id=user_id,
|
||||
username=username,
|
||||
)
|
||||
)
|
||||
|
||||
async def async_delete(
|
||||
|
||||
@@ -182,7 +182,9 @@ class AgentTaskOper(DbOper):
|
||||
|
||||
def get_run(self, run_id: str) -> Optional[AgentTaskRun]:
|
||||
"""查询一次 Agent 任务运行。"""
|
||||
return AgentTaskRun.get_by_run_id(self._db, run_id=run_id)
|
||||
return self._execute_sync_query(
|
||||
lambda session: AgentTaskRun.get_by_run_id(session, run_id=run_id)
|
||||
)
|
||||
|
||||
def list_runs(
|
||||
self,
|
||||
@@ -191,11 +193,13 @@ class AgentTaskOper(DbOper):
|
||||
limit: int = 10,
|
||||
) -> list[AgentTaskRun]:
|
||||
"""查询任务最近的有界运行历史。"""
|
||||
return AgentTaskRun.list_for_task(
|
||||
self._db,
|
||||
task_id=task_id,
|
||||
user_id=user_id,
|
||||
limit=limit,
|
||||
return self._execute_sync_query(
|
||||
lambda session: AgentTaskRun.list_for_task(
|
||||
session,
|
||||
task_id=task_id,
|
||||
user_id=user_id,
|
||||
limit=limit,
|
||||
)
|
||||
)
|
||||
|
||||
def finish_run(
|
||||
|
||||
Reference in New Issue
Block a user