diff --git a/scripts/batch_upload.py b/scripts/batch_upload.py
index b8ce49ac..33c23b4b 100644
--- a/scripts/batch_upload.py
+++ b/scripts/batch_upload.py
@@ -88,8 +88,18 @@ async def process_document(
response.raise_for_status()
result = response.json()
- # Check if the overall request was successful
- if result.get("status") != "success":
+ # Handle asynchronous ingest response
+ overall_status = result.get("status")
+ if overall_status == "queued":
+ task_id = result.get("task_id")
+ extra = f" (task id: {task_id})" if task_id else ""
+ console.print(
+ f"[bold cyan]Ingestion queued for {server_file_path}{extra}. Track progress in the task center.[/bold cyan]"
+ )
+ return True
+
+ # Check if the overall request was successful for synchronous responses
+ if overall_status != "success":
console.print(
f"[bold yellow]Processing warning for {server_file_path}: {result.get('message')}[/bold yellow]"
)
diff --git a/server/main.py b/server/main.py
index a26ead72..701fd773 100644
--- a/server/main.py
+++ b/server/main.py
@@ -4,6 +4,7 @@ from fastapi.middleware.cors import CORSMiddleware
from starlette.middleware.base import BaseHTTPMiddleware
from server.routers import router
+from server.services.tasker import tasker
from server.utils.auth_middleware import is_public_path
from server.utils.common_utils import setup_logging
@@ -59,5 +60,16 @@ class AuthMiddleware(BaseHTTPMiddleware):
# 添加鉴权中间件
app.add_middleware(AuthMiddleware)
+
+@app.on_event("startup")
+async def start_tasker() -> None:
+ await tasker.start()
+
+
+@app.on_event("shutdown")
+async def stop_tasker() -> None:
+ await tasker.shutdown()
+
+
if __name__ == "__main__":
uvicorn.run(app, host="0.0.0.0", port=5050, threads=10, workers=10, reload=True)
diff --git a/server/routers/__init__.py b/server/routers/__init__.py
index b2afa84f..987853ae 100644
--- a/server/routers/__init__.py
+++ b/server/routers/__init__.py
@@ -6,6 +6,7 @@ from server.routers.dashboard_router import dashboard
from server.routers.graph_router import graph
from server.routers.knowledge_router import knowledge
from server.routers.system_router import system
+from server.routers.task_router import tasks
router = APIRouter()
@@ -16,3 +17,4 @@ router.include_router(chat) # /api/chat/*
router.include_router(dashboard) # /api/dashboard/*
router.include_router(knowledge) # /api/knowledge/*
router.include_router(graph) # /api/graph/*
+router.include_router(tasks) # /api/tasks/*
diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py
index e417be6e..8ecff08c 100644
--- a/server/routers/knowledge_router.py
+++ b/server/routers/knowledge_router.py
@@ -1,3 +1,4 @@
+import asyncio
import os
import traceback
from urllib.parse import quote, unquote
@@ -8,6 +9,7 @@ from starlette.responses import FileResponse as StarletteFileResponse
from src.storage.db.models import User
from server.utils.auth_middleware import get_admin_user
+from server.services.tasker import TaskContext, tasker
from src import config, knowledge_base
from src.knowledge.indexing import SUPPORTED_FILE_EXTENSIONS, is_supported_file_extension, process_file_to_markdown
from src.models.embed import test_embedding_model_status, test_all_embedding_models_status
@@ -160,15 +162,63 @@ async def add_documents(
except ValueError as e:
raise HTTPException(status_code=403, detail=str(e))
+ async def run_ingest(context: TaskContext):
+ await context.set_message("任务初始化")
+ await context.set_progress(5.0, "准备处理文档")
+
+ total = len(items)
+ processed_items = []
+
+ try:
+ # 逐个处理文档并更新进度
+ for idx, item in enumerate(items, 1):
+ await context.raise_if_cancelled()
+
+ # 更新进度
+ progress = 5.0 + (idx / total) * 90.0 # 5% ~ 95%
+ await context.set_progress(progress, f"正在处理第 {idx}/{total} 个文档")
+
+ # 处理单个文档
+ result = await knowledge_base.add_content(db_id, [item], params=params)
+ processed_items.extend(result)
+
+ except asyncio.CancelledError:
+ await context.set_progress(100.0, "任务已取消")
+ raise
+
+ item_type = "URL" if content_type == "url" else "文件"
+ failed_count = len([_p for _p in processed_items if _p.get("status") == "failed"])
+ summary = {
+ "db_id": db_id,
+ "item_type": item_type,
+ "submitted": len(processed_items),
+ "failed": failed_count,
+ }
+ message = f"{item_type}处理完成,失败 {failed_count} 个" if failed_count else f"{item_type}处理完成"
+ await context.set_result(summary | {"items": processed_items})
+ await context.set_progress(100.0, message)
+ return summary | {"items": processed_items}
+
try:
- processed_items = await knowledge_base.add_content(db_id, items, params=params)
- item_type = "URLs" if content_type == "url" else "files"
- processed_failed_count = len([_p for _p in processed_items if _p["status"] == "failed"])
- processed_info = f"Processed {len(processed_items)} {item_type}, {processed_failed_count} {item_type} failed"
- return {"message": processed_info, "items": processed_items, "status": "success"}
- except Exception as e:
- logger.error(f"Failed to process {content_type}s: {e}, {traceback.format_exc()}")
- return {"message": f"Failed to process {content_type}s: {e}", "status": "failed"}
+ task = await tasker.enqueue(
+ name=f"知识库文档处理({db_id})",
+ task_type="knowledge_ingest",
+ payload={
+ "db_id": db_id,
+ "items": items,
+ "params": params,
+ "content_type": content_type,
+ },
+ coroutine=run_ingest,
+ )
+ return {
+ "message": "任务已提交,请在任务中心查看进度",
+ "status": "queued",
+ "task_id": task.id,
+ }
+ except Exception as e: # noqa: BLE001
+ logger.error(f"Failed to enqueue {content_type}s: {e}, {traceback.format_exc()}")
+ return {"message": f"Failed to enqueue task: {e}", "status": "failed"}
@knowledge.get("/databases/{db_id}/documents/{doc_id}")
diff --git a/server/routers/task_router.py b/server/routers/task_router.py
new file mode 100644
index 00000000..c4c4246d
--- /dev/null
+++ b/server/routers/task_router.py
@@ -0,0 +1,35 @@
+from fastapi import APIRouter, Depends, HTTPException, Query
+
+from src.storage.db.models import User
+from server.services.tasker import tasker
+from server.utils.auth_middleware import get_admin_user
+
+tasks = APIRouter(prefix="/tasks", tags=["tasks"])
+
+
+@tasks.get("")
+async def list_tasks(
+ status: str | None = Query(default=None),
+ current_user: User = Depends(get_admin_user),
+):
+ """List tasks, optionally filtered by status."""
+ task_list = await tasker.list_tasks(status=status)
+ return {"tasks": task_list}
+
+
+@tasks.get("/{task_id}")
+async def get_task(task_id: str, current_user: User = Depends(get_admin_user)):
+ """Retrieve a single task by id."""
+ task = await tasker.get_task(task_id)
+ if not task:
+ raise HTTPException(status_code=404, detail="Task not found")
+ return {"task": task}
+
+
+@tasks.post("/{task_id}/cancel")
+async def cancel_task(task_id: str, current_user: User = Depends(get_admin_user)):
+ """Request cancellation of a task."""
+ success = await tasker.cancel_task(task_id)
+ if not success:
+ raise HTTPException(status_code=400, detail="Task cannot be cancelled")
+ return {"task_id": task_id, "status": "cancelled"}
diff --git a/server/services/__init__.py b/server/services/__init__.py
new file mode 100644
index 00000000..5bbcfe71
--- /dev/null
+++ b/server/services/__init__.py
@@ -0,0 +1,3 @@
+from .tasker import TaskContext, Tasker, tasker
+
+__all__ = ["TaskContext", "Tasker", "tasker"]
diff --git a/server/services/tasker.py b/server/services/tasker.py
new file mode 100644
index 00000000..ffe8a0aa
--- /dev/null
+++ b/server/services/tasker.py
@@ -0,0 +1,301 @@
+import asyncio
+import json
+import os
+import uuid
+from dataclasses import asdict, dataclass, field
+from datetime import datetime
+from pathlib import Path
+from typing import Any, Awaitable, Callable, Dict, List, Optional
+
+from src.config import config
+from src.utils.logging_config import logger
+
+TaskCoroutine = Callable[["TaskContext"], Awaitable[Any]]
+TERMINAL_STATUSES = {"success", "failed", "cancelled"}
+
+
+def _utc_timestamp() -> str:
+ return datetime.utcnow().isoformat() + "Z"
+
+
+@dataclass
+class Task:
+ id: str
+ name: str
+ type: str
+ 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)
+ started_at: Optional[str] = None
+ completed_at: Optional[str] = None
+ payload: Dict[str, Any] = field(default_factory=dict)
+ result: Optional[Any] = None
+ error: Optional[str] = None
+ cancel_requested: bool = False
+
+ def to_dict(self) -> Dict[str, Any]:
+ data = asdict(self)
+ return data
+
+ @classmethod
+ def from_dict(cls, data: Dict[str, Any]) -> "Task":
+ return cls(
+ id=data["id"],
+ name=data.get("name", "Unnamed Task"),
+ type=data.get("type", "general"),
+ 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()),
+ started_at=data.get("started_at"),
+ completed_at=data.get("completed_at"),
+ payload=data.get("payload", {}),
+ result=data.get("result"),
+ error=data.get("error"),
+ cancel_requested=data.get("cancel_requested", False),
+ )
+
+
+class TaskContext:
+ def __init__(self, tasker: "Tasker", task_id: str):
+ self._tasker = tasker
+ self.task_id = task_id
+
+ async def set_progress(self, progress: float, message: Optional[str] = None) -> None:
+ await self._tasker._update_task(
+ self.task_id,
+ progress=max(0.0, min(progress, 100.0)),
+ message=message,
+ )
+
+ async def set_message(self, message: str) -> None:
+ await self._tasker._update_task(self.task_id, message=message)
+
+ async def set_result(self, result: Any) -> None:
+ await self._tasker._update_task(self.task_id, result=result)
+
+ def is_cancel_requested(self) -> bool:
+ return self._tasker._is_cancel_requested(self.task_id)
+
+ async def raise_if_cancelled(self) -> None:
+ if self.is_cancel_requested():
+ raise asyncio.CancelledError("Task was cancelled")
+
+
+class Tasker:
+ def __init__(self, worker_count: int = 2):
+ self.worker_count = max(1, worker_count)
+ self._queue: "asyncio.Queue[tuple[str, TaskCoroutine]]" = asyncio.Queue()
+ 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
+
+ async def start(self) -> None:
+ async with self._lock:
+ if self._started:
+ return
+ await self._load_state()
+ for _ in range(self.worker_count):
+ worker = asyncio.create_task(self._worker_loop(), name="tasker-worker")
+ self._workers.append(worker)
+ self._started = True
+ logger.info("Tasker started with %s workers", self.worker_count)
+
+ async def shutdown(self) -> None:
+ async with self._lock:
+ if not self._started:
+ return
+ for worker in self._workers:
+ 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")
+
+ async def enqueue(
+ self,
+ *,
+ name: str,
+ task_type: str,
+ payload: Optional[Dict[str, Any]] = None,
+ coroutine: TaskCoroutine,
+ ) -> Task:
+ task_id = uuid.uuid4().hex
+ 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._queue.put((task_id, coroutine))
+ logger.info("Enqueued task %s (%s)", task_id, name)
+ return task
+
+ async def list_tasks(self, status: Optional[str] = None) -> List[Dict[str, Any]]:
+ async with self._lock:
+ tasks = list(self._tasks.values())
+ if status:
+ tasks = [task for task in tasks if task.status == status]
+ tasks.sort(key=lambda item: item.created_at, reverse=True)
+ return [task.to_dict() for task in tasks]
+
+ async def get_task(self, task_id: str) -> Optional[Dict[str, Any]]:
+ async with self._lock:
+ task = self._tasks.get(task_id)
+ return task.to_dict() if task else None
+
+ async def cancel_task(self, task_id: str) -> bool:
+ async with self._lock:
+ task = self._tasks.get(task_id)
+ if not task:
+ return False
+ if task.status in {"success", "failed", "cancelled"}:
+ return False
+ task.cancel_requested = True
+ task.updated_at = _utc_timestamp()
+ await self._persist_state()
+ logger.info("Cancellation requested for task %s", task_id)
+ return True
+
+ async def _worker_loop(self) -> None:
+ while True:
+ try:
+ task_id, coroutine = await self._queue.get()
+ try:
+ task = await self._get_task_instance(task_id)
+ if not task:
+ continue
+ if task.cancel_requested:
+ 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())
+ context = TaskContext(self, task_id)
+ try:
+ result = await coroutine(context)
+ if task.cancel_requested:
+ await self._mark_cancelled(task_id, "Task cancelled during execution")
+ continue
+ await self._update_task(
+ task_id,
+ status="success",
+ progress=100.0,
+ message="任务已完成",
+ result=result,
+ completed_at=_utc_timestamp(),
+ )
+ except asyncio.CancelledError:
+ await self._mark_cancelled(task_id, "任务被取消")
+ except Exception as exc: # noqa: BLE001
+ logger.error("Task %s failed: %s", task_id, exc, exc_info=True)
+ await self._update_task(
+ task_id,
+ status="failed",
+ progress=100.0,
+ message="任务执行失败",
+ error=str(exc),
+ completed_at=_utc_timestamp(),
+ )
+ finally:
+ self._queue.task_done()
+ except asyncio.CancelledError:
+ break
+ except Exception as exc: # noqa: BLE001
+ logger.error("Tasker worker error: %s", exc, exc_info=True)
+
+ async def _get_task_instance(self, task_id: str) -> Optional[Task]:
+ async with self._lock:
+ return self._tasks.get(task_id)
+
+ async def _mark_cancelled(self, task_id: str, message: str) -> None:
+ await self._update_task(
+ task_id,
+ status="cancelled",
+ progress=100.0,
+ message=message,
+ completed_at=_utc_timestamp(),
+ )
+
+ async def _update_task(
+ self,
+ task_id: str,
+ *,
+ status: Optional[str] = None,
+ progress: Optional[float] = None,
+ message: Optional[str] = None,
+ result: Any = None,
+ error: Optional[str] = None,
+ started_at: Optional[str] = None,
+ completed_at: Optional[str] = None,
+ ) -> None:
+ async with self._lock:
+ task = self._tasks.get(task_id)
+ if not task:
+ return
+ if status:
+ task.status = status
+ if progress is not None:
+ task.progress = max(0.0, min(progress, 100.0))
+ if message is not None:
+ task.message = message
+ if result is not None:
+ task.result = result
+ if error is not None:
+ task.error = error
+ if started_at is not None:
+ 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()
+
+ 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 %s task records from storage", len(tasks))
+ except Exception as exc: # noqa: BLE001
+ logger.error("Failed to load task state: %s", exc, exc_info=True)
+
+ async def _persist_state(self) -> None:
+ tasks = [task.to_dict() for task in self._tasks.values()]
+ payload = {"tasks": tasks, "updated_at": _utc_timestamp()}
+
+ async 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)
+
+
+tasker = Tasker()
+
+
+__all__ = ["tasker", "TaskContext", "Tasker"]
diff --git a/test/api/test_task_router.py b/test/api/test_task_router.py
new file mode 100644
index 00000000..8080f25d
--- /dev/null
+++ b/test/api/test_task_router.py
@@ -0,0 +1,93 @@
+"""
+Integration tests for the task management router.
+"""
+
+from __future__ import annotations
+
+import asyncio
+
+import pytest
+
+pytestmark = [pytest.mark.asyncio, pytest.mark.integration]
+
+
+async def test_task_routes_require_admin(test_client, standard_user):
+ """Non-admin users should be blocked from accessing task APIs."""
+ headers = standard_user["headers"]
+
+ list_response = await test_client.get("/api/tasks", headers=headers)
+ assert list_response.status_code == 403
+
+ detail_response = await test_client.get("/api/tasks/some-task", headers=headers)
+ assert detail_response.status_code == 403
+
+ cancel_response = await test_client.post("/api/tasks/some-task/cancel", headers=headers)
+ assert cancel_response.status_code == 403
+
+
+async def test_admin_can_list_tasks(test_client, admin_headers):
+ """Admin should receive a well-formed task list payload."""
+ response = await test_client.get("/api/tasks", headers=admin_headers)
+ assert response.status_code == 200, response.text
+
+ payload = response.json()
+ assert "tasks" in payload
+ assert isinstance(payload["tasks"], list)
+
+
+async def test_cancel_unknown_task_returns_client_error(test_client, admin_headers):
+ """Cancelling a non-existent task should surface a 400 response."""
+ response = await test_client.post("/api/tasks/not-real/cancel", headers=admin_headers)
+ assert response.status_code == 400, response.text
+
+
+async def test_enqueue_document_creates_task(
+ test_client,
+ admin_headers,
+ knowledge_database,
+):
+ """Trigger knowledge ingestion to ensure a task record is materialised."""
+ db_id = knowledge_database["db_id"]
+
+ enqueue_response = await test_client.post(
+ f"/api/knowledge/databases/{db_id}/documents",
+ json={
+ "items": [],
+ "params": {"content_type": "file"},
+ },
+ headers=admin_headers,
+ )
+ assert enqueue_response.status_code == 200, enqueue_response.text
+
+ enqueue_payload = enqueue_response.json()
+ assert enqueue_payload.get("status") == "queued"
+ task_id = enqueue_payload.get("task_id")
+ assert task_id, "Knowledge ingestion did not return a task_id"
+
+ # The task should be queryable immediately after enqueueing.
+ detail_response = await test_client.get(f"/api/tasks/{task_id}", headers=admin_headers)
+ assert detail_response.status_code == 200, detail_response.text
+ detail_payload = detail_response.json().get("task", {})
+ assert detail_payload.get("id") == task_id
+ assert detail_payload.get("status") in {"queued", "pending", "running", "failed", "success", "cancelled"}
+
+ # Ensure the task surfaces in the list endpoint within a short window.
+ for _ in range(10):
+ list_response = await test_client.get("/api/tasks", headers=admin_headers)
+ assert list_response.status_code == 200, list_response.text
+ all_tasks = list_response.json().get("tasks", [])
+ if any(entry.get("id") == task_id for entry in all_tasks):
+ break
+ await asyncio.sleep(0.2)
+ else:
+ pytest.fail("Task did not appear in list endpoint within timeout window")
+
+ # Poll for terminal state to validate worker bookkeeping.
+ for _ in range(20):
+ detail_response = await test_client.get(f"/api/tasks/{task_id}", headers=admin_headers)
+ task_status = detail_response.json().get("task", {}).get("status")
+ if task_status in {"success", "failed", "cancelled"}:
+ break
+ await asyncio.sleep(0.5)
+ else:
+ pytest.fail("Task did not reach a terminal status within timeout window")
diff --git a/web/src/apis/index.js b/web/src/apis/index.js
index 2131228f..259ec55d 100644
--- a/web/src/apis/index.js
+++ b/web/src/apis/index.js
@@ -8,6 +8,7 @@ export * from './system_api' // 系统管理API
export * from './knowledge_api' // 知识库管理API
export * from './graph_api' // 图谱API
export * from './agent_api' // 智能体API
+export * from './tasker' // 任务管理API
// 导出基础工具函数
export { apiGet, apiPost, apiPut, apiDelete,
@@ -36,4 +37,4 @@ export { apiGet, apiPost, apiPut, apiDelete,
* - 智能体管理、聊天、配置等功能
*
* 注意:API模块已处理权限验证和请求头,使用时无需再手动添加认证头
- */
\ No newline at end of file
+ */
diff --git a/web/src/apis/tasker.js b/web/src/apis/tasker.js
new file mode 100644
index 00000000..612d510b
--- /dev/null
+++ b/web/src/apis/tasker.js
@@ -0,0 +1,19 @@
+import { apiAdminGet, apiAdminPost } from './base'
+
+const BASE_URL = '/api/tasks'
+
+export const taskerApi = {
+ fetchTasks: async (params = {}) => {
+ const query = new URLSearchParams(params).toString()
+ const url = query ? `${BASE_URL}?${query}` : BASE_URL
+ return apiAdminGet(url)
+ },
+
+ fetchTaskDetail: async (taskId) => {
+ return apiAdminGet(`${BASE_URL}/${taskId}`)
+ },
+
+ cancelTask: async (taskId) => {
+ return apiAdminPost(`${BASE_URL}/${taskId}/cancel`, {})
+ }
+}
diff --git a/web/src/components/StatusBar.vue b/web/src/components/StatusBar.vue
index 23a37bf5..3f3f4546 100644
--- a/web/src/components/StatusBar.vue
+++ b/web/src/components/StatusBar.vue
@@ -21,6 +21,18 @@