diff --git a/docs/latest/changelog/roadmap.md b/docs/latest/changelog/roadmap.md index a31c5c69..a7d0889c 100644 --- a/docs/latest/changelog/roadmap.md +++ b/docs/latest/changelog/roadmap.md @@ -14,6 +14,7 @@ - 系统层面添加 apikey,在智能体、知识库调用中支持 apikey 以支持外部调用 - 支持更多类型的文档源的导入功能 - 检查非 Agent 场景下的知识库的可见情况 +- Tasker 新增删除任务的接口 ### Bugs - 部分异常状态下,智能体的模型名称出现重叠[#279](https://github.com/xerrors/Yuxi-Know/issues/279) diff --git a/scripts/migrate_kb_metadata_to_db.py b/scripts/migrate_kb_metadata_to_db.py index 0498fd71..7229d70b 100644 --- a/scripts/migrate_kb_metadata_to_db.py +++ b/scripts/migrate_kb_metadata_to_db.py @@ -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") diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index 30921ea7..ae133301 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -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 diff --git a/server/routers/task_router.py b/server/routers/task_router.py index e6d87fee..848af3f9 100644 --- a/server/routers/task_router.py +++ b/server/routers/task_router.py @@ -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"]) diff --git a/server/services/__init__.py b/server/services/__init__.py deleted file mode 100644 index 5bbcfe71..00000000 --- a/server/services/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .tasker import TaskContext, Tasker, tasker - -__all__ = ["TaskContext", "Tasker", "tasker"] diff --git a/server/utils/lifespan.py b/server/utils/lifespan.py index eb38bf13..175f9dca 100644 --- a/server/utils/lifespan.py +++ b/server/utils/lifespan.py @@ -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 diff --git a/src/repositories/task_repository.py b/src/repositories/task_repository.py new file mode 100644 index 00000000..8e8af1dd --- /dev/null +++ b/src/repositories/task_repository.py @@ -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)) + diff --git a/src/services/evaluation_service.py b/src/services/evaluation_service.py index b1dace53..6720facb 100644 --- a/src/services/evaluation_service.py +++ b/src/services/evaluation_service.py @@ -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 diff --git a/server/services/tasker.py b/src/services/task_service.py similarity index 78% rename from server/services/tasker.py rename to src/services/task_service.py index 1bdfdb35..a23b3dbd 100644 --- a/server/services/tasker.py +++ b/src/services/task_service.py @@ -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"] + diff --git a/src/storage/postgres/models_business.py b/src/storage/postgres/models_business.py index 79fc733b..1aaa0419 100644 --- a/src/storage/postgres/models_business.py +++ b/src/storage/postgres/models_business.py @@ -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