refactor: migrate agent task query ownership

This commit is contained in:
jxxghp
2026-08-23 00:37:03 +08:00
parent ab51b5116a
commit 2031f3c420
6 changed files with 92 additions and 32 deletions
+30 -13
View File
@@ -4,7 +4,32 @@ from sqlalchemy import Boolean, Index, Integer, String, Text, 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
def _get_for_user_statement(
model: type["AgentTask"],
task_id: int,
user_id: Optional[str] = None,
):
"""构造按任务 ID 与可选用户归属收窄的查询语句。"""
statement = select(model).where(model.id == task_id)
if user_id is not None:
statement = statement.where(model.user_id == user_id)
return statement
def _list_for_user_statement(
model: type["AgentTask"],
user_id: Optional[str] = None,
enabled: Optional[bool] = None,
):
"""构造按用户、启用状态和创建时间排序的任务列表语句。"""
statement = select(model)
if user_id is not None:
statement = statement.where(model.user_id == user_id)
if enabled is not None:
statement = statement.where(model.enabled.is_(enabled))
return statement.order_by(model.created_at.desc(), model.id.desc())
class AgentTask(Base):
@@ -59,7 +84,6 @@ class AgentTask(Base):
return task.id
@classmethod
@db_query
def get_for_user(
cls,
db: Session,
@@ -69,13 +93,11 @@ class AgentTask(Base):
"""
按任务 ID 和可选用户 ID 查询 Agent 定时任务。
"""
statement = select(cls).where(cls.id == task_id)
if user_id is not None:
statement = statement.where(cls.user_id == user_id)
return db.execute(statement).scalars().first()
return db.execute(
_get_for_user_statement(cls, task_id=task_id, user_id=user_id)
).scalars().first()
@classmethod
@db_query
def list_for_user(
cls,
db: Session,
@@ -85,13 +107,8 @@ class AgentTask(Base):
"""
按用户和启用状态查询 Agent 定时任务。
"""
statement = select(cls)
if user_id is not None:
statement = statement.where(cls.user_id == user_id)
if enabled is not None:
statement = statement.where(cls.enabled.is_(enabled))
return list(db.execute(
statement.order_by(cls.created_at.desc(), cls.id.desc())
_list_for_user_statement(cls, user_id=user_id, enabled=enabled)
).scalars().all())
@classmethod