refactor(task): 重构任务服务并迁移到postgres 数据库
将任务服务从server/services移动到src/services 添加TaskRepository实现数据库持久化 更新所有相关导入路径
This commit is contained in:
parent
63fed0dfd6
commit
e585f796d2
@ -14,6 +14,7 @@
|
||||
- 系统层面添加 apikey,在智能体、知识库调用中支持 apikey 以支持外部调用
|
||||
- 支持更多类型的文档源的导入功能
|
||||
- 检查非 Agent 场景下的知识库的可见情况
|
||||
- Tasker 新增删除任务的接口
|
||||
|
||||
### Bugs
|
||||
- 部分异常状态下,智能体的模型名称出现重叠[#279](https://github.com/xerrors/Yuxi-Know/issues/279)
|
||||
|
||||
@ -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")
|
||||
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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"])
|
||||
|
||||
@ -1,3 +0,0 @@
|
||||
from .tasker import TaskContext, Tasker, tasker
|
||||
|
||||
__all__ = ["TaskContext", "Tasker", "tasker"]
|
||||
@ -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
|
||||
|
||||
46
src/repositories/task_repository.py
Normal file
46
src/repositories/task_repository.py
Normal 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))
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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"]
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user