refactor(task): 重构任务服务并迁移到postgres 数据库

将任务服务从server/services移动到src/services
添加TaskRepository实现数据库持久化
更新所有相关导入路径
This commit is contained in:
Wenjie Zhang 2026-01-22 02:09:50 +08:00
parent 63fed0dfd6
commit e585f796d2
10 changed files with 192 additions and 72 deletions

View File

@ -14,6 +14,7 @@
- 系统层面添加 apikey在智能体、知识库调用中支持 apikey 以支持外部调用
- 支持更多类型的文档源的导入功能
- 检查非 Agent 场景下的知识库的可见情况
- Tasker 新增删除任务的接口
### Bugs
- 部分异常状态下,智能体的模型名称出现重叠[#279](https://github.com/xerrors/Yuxi-Know/issues/279)

View File

@ -14,6 +14,7 @@ from src import config
from src.repositories.evaluation_repository import EvaluationRepository
from src.repositories.knowledge_base_repository import KnowledgeBaseRepository
from src.repositories.knowledge_file_repository import KnowledgeFileRepository
from src.repositories.task_repository import TaskRepository
from src.utils import logger
@ -64,7 +65,9 @@ async def rollback_all() -> None:
eval_repo = EvaluationRepository()
kb_repo = KnowledgeBaseRepository()
file_repo = KnowledgeFileRepository()
task_repo = TaskRepository()
await task_repo.delete_all()
await eval_repo.delete_all()
rows = await kb_repo.get_all()
@ -96,6 +99,7 @@ async def migrate(dry_run: bool, execute: bool, rollback: bool) -> None:
kb_repo = KnowledgeBaseRepository()
file_repo = KnowledgeFileRepository()
eval_repo = EvaluationRepository()
task_repo = TaskRepository()
kb_rows: list[dict[str, Any]] = []
file_rows: list[tuple[str, dict[str, Any]]] = []
@ -234,9 +238,13 @@ async def migrate(dry_run: bool, execute: bool, rollback: bool) -> None:
)
)
tasks_json_path = os.path.join(config.save_dir, "tasks", "tasks.json")
task_rows: list[dict[str, Any]] = _load_json(tasks_json_path).get("tasks", []) or []
logger.info(
f"Prepared: knowledge_bases={len(kb_rows)}, knowledge_files={len(file_rows)}, "
f"benchmarks={len(benchmark_rows)}, results={len(result_rows)}, result_details={len(result_detail_rows)}"
f"benchmarks={len(benchmark_rows)}, results={len(result_rows)}, result_details={len(result_detail_rows)}, "
f"tasks={len(task_rows)}"
)
if dry_run and not execute:
@ -286,6 +294,31 @@ async def migrate(dry_run: bool, execute: bool, rollback: bool) -> None:
for task_id, idx, data in result_detail_rows:
await eval_repo.upsert_result_detail(task_id=task_id, query_index=idx, data=data)
for item in task_rows:
task_id = item.get("id")
if not task_id:
continue
payload = item.get("payload") or {}
result = item.get("result")
await task_repo.upsert(
task_id,
{
"name": item.get("name") or "Unnamed Task",
"type": item.get("type") or "general",
"status": item.get("status") or "pending",
"progress": float(item.get("progress") or 0.0),
"message": item.get("message") or "",
"payload": payload,
"result": result,
"error": item.get("error"),
"cancel_requested": 1 if item.get("cancel_requested") else 0,
"created_at": _utc_dt(item.get("created_at")),
"updated_at": _utc_dt(item.get("updated_at")) or _utc_dt(item.get("created_at")),
"started_at": _utc_dt(item.get("started_at")),
"completed_at": _utc_dt(item.get("completed_at")),
},
)
logger.info("Migration completed")

View File

@ -10,7 +10,7 @@ from fastapi import APIRouter, Body, Depends, File, HTTPException, Query, Reques
from fastapi.responses import FileResponse
from starlette.responses import StreamingResponse
from server.services.tasker import TaskContext, tasker
from src.services.task_service import TaskContext, tasker
from server.utils.auth_middleware import get_admin_user, get_required_user
from src import config, knowledge_base
from src.knowledge.indexing import SUPPORTED_FILE_EXTENSIONS, is_supported_file_extension, process_file_to_markdown

View File

@ -1,7 +1,7 @@
from fastapi import APIRouter, Depends, HTTPException, Query
from src.storage.postgres.models_business import User
from server.services.tasker import tasker
from src.services.task_service import tasker
from server.utils.auth_middleware import get_admin_user
tasks = APIRouter(prefix="/tasks", tags=["tasks"])

View File

@ -1,3 +0,0 @@
from .tasker import TaskContext, Tasker, tasker
__all__ = ["TaskContext", "Tasker", "tasker"]

View File

@ -2,7 +2,7 @@ from contextlib import asynccontextmanager
from fastapi import FastAPI
from server.services import tasker
from src.services.task_service import tasker
from src.services.mcp_service import init_mcp_servers
from src.storage.postgres.manager import pg_manager
from src.knowledge import knowledge_base

View File

@ -0,0 +1,46 @@
from __future__ import annotations
from typing import Any
from sqlalchemy import delete, select
from src.storage.postgres.manager import pg_manager
from src.storage.postgres.models_business import TaskRecord
class TaskRepository:
async def get_by_id(self, task_id: str) -> TaskRecord | None:
async with pg_manager.get_async_session_context() as session:
result = await session.execute(select(TaskRecord).where(TaskRecord.id == task_id))
return result.scalar_one_or_none()
async def list(self, status: str | None = None, limit: int = 100) -> list[TaskRecord]:
async with pg_manager.get_async_session_context() as session:
stmt = select(TaskRecord)
if status:
stmt = stmt.where(TaskRecord.status == status)
stmt = stmt.order_by(TaskRecord.created_at.desc()).limit(max(limit, 0))
result = await session.execute(stmt)
return list(result.scalars().all())
async def list_all(self) -> list[TaskRecord]:
async with pg_manager.get_async_session_context() as session:
result = await session.execute(select(TaskRecord).order_by(TaskRecord.created_at.desc()))
return list(result.scalars().all())
async def upsert(self, task_id: str, data: dict[str, Any]) -> TaskRecord:
async with pg_manager.get_async_session_context() as session:
result = await session.execute(select(TaskRecord).where(TaskRecord.id == task_id))
record = result.scalar_one_or_none()
if record is None:
record = TaskRecord(id=task_id, **data)
session.add(record)
return record
for key, value in data.items():
setattr(record, key, value)
return record
async def delete_all(self) -> None:
async with pg_manager.get_async_session_context() as session:
await session.execute(delete(TaskRecord))

View File

@ -6,7 +6,7 @@ import uuid
from datetime import datetime
from typing import Any
from server.services.tasker import TaskContext, tasker
from src.services.task_service import TaskContext, tasker
from src.knowledge import knowledge_base
from src.models import select_model
from src.repositories.evaluation_repository import EvaluationRepository

View File

@ -1,23 +1,23 @@
import asyncio
import json
import os
import uuid
from dataclasses import asdict, dataclass, field
from pathlib import Path
from typing import Any
from collections.abc import Awaitable, Callable
from collections import Counter
from collections.abc import Awaitable, Callable
from dataclasses import asdict, dataclass, field
from datetime import datetime
from typing import Any
from src.config import config
from src.repositories.task_repository import TaskRepository
from src.utils.datetime_utils import coerce_any_to_utc_datetime, utc_isoformat
from src.utils.logging_config import logger
from src.utils.datetime_utils import utc_isoformat
TaskCoroutine = Callable[["TaskContext"], Awaitable[Any]]
TERMINAL_STATUSES = {"success", "failed", "cancelled"}
def _utc_timestamp() -> str:
return utc_isoformat()
def _iso_to_utc_naive(value: str | None) -> datetime | None:
if not value:
return None
return coerce_any_to_utc_datetime(value).replace(tzinfo=None)
@dataclass
@ -28,8 +28,8 @@ class Task:
status: str = "pending"
progress: float = 0.0
message: str = ""
created_at: str = field(default_factory=_utc_timestamp)
updated_at: str = field(default_factory=_utc_timestamp)
created_at: str = field(default_factory=utc_isoformat)
updated_at: str = field(default_factory=utc_isoformat)
started_at: str | None = None
completed_at: str | None = None
payload: dict[str, Any] = field(default_factory=dict)
@ -38,8 +38,7 @@ class Task:
cancel_requested: bool = False
def to_dict(self) -> dict[str, Any]:
data = asdict(self)
return data
return asdict(self)
def to_summary_dict(self) -> dict[str, Any]:
data = asdict(self)
@ -56,8 +55,8 @@ class Task:
status=data.get("status", "pending"),
progress=data.get("progress", 0.0),
message=data.get("message", ""),
created_at=data.get("created_at", _utc_timestamp()),
updated_at=data.get("updated_at", _utc_timestamp()),
created_at=data.get("created_at", utc_isoformat()),
updated_at=data.get("updated_at", utc_isoformat()),
started_at=data.get("started_at"),
completed_at=data.get("completed_at"),
payload=data.get("payload", {}),
@ -100,9 +99,8 @@ class Tasker:
self._tasks: dict[str, Task] = {}
self._lock = asyncio.Lock()
self._workers: list[asyncio.Task[Any]] = []
self._storage_path = Path(config.save_dir) / "tasks" / "tasks.json"
os.makedirs(self._storage_path.parent, exist_ok=True)
self._started = False
self._repo = TaskRepository()
async def start(self) -> None:
async with self._lock:
@ -123,7 +121,6 @@ class Tasker:
worker.cancel()
await asyncio.gather(*self._workers, return_exceptions=True)
self._workers.clear()
await self._persist_state()
self._started = False
logger.info("Tasker shutdown complete")
@ -139,7 +136,7 @@ class Tasker:
task = Task(id=task_id, name=name, type=task_type, payload=payload or {})
async with self._lock:
self._tasks[task_id] = task
await self._persist_state()
await self._persist_task(task)
await self._queue.put((task_id, coroutine))
logger.info("Enqueued task {} ({})", task_id, name)
return task
@ -180,11 +177,11 @@ class Tasker:
task = self._tasks.get(task_id)
if not task:
return False
if task.status in {"success", "failed", "cancelled"}:
if task.status in TERMINAL_STATUSES:
return False
task.cancel_requested = True
task.updated_at = _utc_timestamp()
await self._persist_state()
task.updated_at = utc_isoformat()
await self._persist_task(task)
logger.info("Cancellation requested for task {}", task_id)
return True
@ -200,7 +197,7 @@ class Tasker:
await self._mark_cancelled(task_id, "Task was cancelled before execution")
continue
await self._update_task(
task_id, status="running", progress=0.0, message="任务开始执行", started_at=_utc_timestamp()
task_id, status="running", progress=0.0, message="任务开始执行", started_at=utc_isoformat()
)
context = TaskContext(self, task_id)
try:
@ -214,7 +211,7 @@ class Tasker:
progress=100.0,
message="任务已完成",
result=result,
completed_at=_utc_timestamp(),
completed_at=utc_isoformat(),
)
except asyncio.CancelledError:
await self._mark_cancelled(task_id, "任务被取消")
@ -226,7 +223,7 @@ class Tasker:
progress=100.0,
message="任务执行失败",
error=str(exc),
completed_at=_utc_timestamp(),
completed_at=utc_isoformat(),
)
finally:
self._queue.task_done()
@ -245,7 +242,7 @@ class Tasker:
status="cancelled",
progress=100.0,
message=message,
completed_at=_utc_timestamp(),
completed_at=utc_isoformat(),
)
async def _update_task(
@ -278,52 +275,55 @@ class Tasker:
task.started_at = started_at
if completed_at is not None:
task.completed_at = completed_at
task.updated_at = _utc_timestamp()
await self._persist_state()
task.updated_at = utc_isoformat()
await self._persist_task(task)
def _is_cancel_requested(self, task_id: str) -> bool:
task = self._tasks.get(task_id)
return bool(task and task.cancel_requested)
async def _load_state(self) -> None:
if not self._storage_path.exists():
return
try:
content = await asyncio.to_thread(self._storage_path.read_text, encoding="utf-8")
if not content.strip():
return
data = json.loads(content)
tasks = data.get("tasks", [])
for item in tasks:
task = Task.from_dict(item)
if task.status == "running":
task.status = "failed"
task.message = "服务重启时任务中断"
task.updated_at = _utc_timestamp()
elif task.status not in TERMINAL_STATUSES:
task.status = "failed"
task.message = "服务重启时任务未继续执行"
task.updated_at = _utc_timestamp()
self._tasks[task.id] = task
logger.info("Loaded {} task records from storage", len(tasks))
except Exception as exc: # noqa: BLE001
logger.exception("Failed to load task state: {}", exc)
records = await self._repo.list_all()
updated: list[Task] = []
for record in records:
task = Task.from_dict(record.to_dict())
if task.status == "running":
task.status = "failed"
task.message = "服务重启时任务中断"
task.updated_at = utc_isoformat()
updated.append(task)
elif task.status not in TERMINAL_STATUSES:
task.status = "failed"
task.message = "服务重启时任务未继续执行"
task.updated_at = utc_isoformat()
updated.append(task)
self._tasks[task.id] = task
for task in updated:
await self._persist_task(task)
if records:
logger.info("Loaded {} task records from storage", len(records))
async def _persist_state(self) -> None:
tasks = [task.to_dict() for task in self._tasks.values()]
payload = {"tasks": tasks, "updated_at": _utc_timestamp()}
def _write() -> None:
self._storage_path.parent.mkdir(parents=True, exist_ok=True)
tmp_path = self._storage_path.with_suffix(".tmp")
with open(tmp_path, "w", encoding="utf-8") as fh:
json.dump(payload, fh, ensure_ascii=False, indent=2)
os.replace(tmp_path, self._storage_path)
await asyncio.to_thread(_write)
async def _persist_task(self, task: Task) -> None:
data: dict[str, Any] = {
"name": task.name,
"type": task.type,
"status": task.status,
"progress": task.progress,
"message": task.message,
"payload": task.payload,
"result": task.result,
"error": task.error,
"cancel_requested": 1 if task.cancel_requested else 0,
"created_at": _iso_to_utc_naive(task.created_at),
"updated_at": _iso_to_utc_naive(task.updated_at),
"started_at": _iso_to_utc_naive(task.started_at),
"completed_at": _iso_to_utc_naive(task.completed_at),
}
await self._repo.upsert(task.id, data)
tasker = Tasker()
__all__ = ["tasker", "TaskContext", "Tasker"]

View File

@ -2,7 +2,7 @@
from typing import Any
from sqlalchemy import JSON, Column, DateTime, ForeignKey, Integer, String, Text
from sqlalchemy import JSON, Column, DateTime, Float, ForeignKey, Integer, String, Text
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import relationship
@ -374,3 +374,46 @@ class MCPServer(Base):
if self.disabled_tools:
config["disabled_tools"] = self.disabled_tools
return config
class TaskRecord(Base):
__tablename__ = "tasks"
id = Column(String(32), primary_key=True)
name = Column(String(255), nullable=False)
type = Column(String(64), nullable=False, index=True)
status = Column(String(32), nullable=False, default="pending", index=True)
progress = Column(Float, nullable=False, default=0.0)
message = Column(Text, nullable=False, default="")
payload = Column(JSON, nullable=True)
result = Column(JSON, nullable=True)
error = Column(Text, nullable=True)
cancel_requested = Column(Integer, nullable=False, default=0)
created_at = Column(DateTime, default=utc_now_naive, index=True)
updated_at = Column(DateTime, default=utc_now_naive, onupdate=utc_now_naive)
started_at = Column(DateTime, nullable=True)
completed_at = Column(DateTime, nullable=True)
def to_dict(self) -> dict[str, Any]:
return {
"id": self.id,
"name": self.name,
"type": self.type,
"status": self.status,
"progress": self.progress,
"message": self.message,
"created_at": format_utc_datetime(self.created_at),
"updated_at": format_utc_datetime(self.updated_at),
"started_at": format_utc_datetime(self.started_at),
"completed_at": format_utc_datetime(self.completed_at),
"payload": self.payload or {},
"result": self.result,
"error": self.error,
"cancel_requested": bool(self.cancel_requested),
}
def to_summary_dict(self) -> dict[str, Any]:
data = self.to_dict()
data.pop("payload", None)
data.pop("result", None)
return data