"""入库异步任务管理器 将文档入库包装为后台异步任务:submit 登记任务并立即返回 task_id, 后台受并发上限控制执行 Ingester.ingest,并按阶段推进任务状态。 内存注册表为主(记录 status/created_at/updated_at/result/error), Redis 为持久镜像(key: ingest_task:{task_id}),每次状态迁移同步写入; Redis 不可用或写入失败仅记录 warning,不影响任务执行。 """ import asyncio import uuid from datetime import UTC, datetime from enum import StrEnum from typing import Any import structlog from app.config import Settings from app.core.ingestion import Ingester, IngestionError from app.models.document import DocumentInput from app.services.redis import RedisCache logger = structlog.get_logger() # Redis 任务状态 key 前缀 REDIS_KEY_PREFIX = "ingest_task:" class IngestTaskStatus(StrEnum): """入库任务状态""" PENDING = "pending" # 已登记,排队等待执行 SUMMARIZING = "summarizing" # 三级总结中 CLASSIFYING = "classifying" # 分类判定中 EMBEDDING = "embedding" # 向量化中 WRITING = "writing" # 写入 Qdrant 中 DONE = "done" # 入库完成 FAILED = "failed" # 入库失败 # 终态集合 TERMINAL_STATUSES: frozenset[str] = frozenset({IngestTaskStatus.DONE, IngestTaskStatus.FAILED}) def _utc_now_iso() -> str: """当前 UTC 时间的 ISO8601 字符串""" return datetime.now(UTC).isoformat() class IngestTaskManager: """入库异步任务管理器:登记、后台执行、状态查询与 Redis 持久镜像""" def __init__(self, ingester: Ingester, redis: RedisCache | None, settings: Settings) -> None: self._ingester = ingester self._redis = redis self._settings = settings self._tasks: dict[str, dict[str, Any]] = {} self._semaphore = asyncio.Semaphore(settings.ingest_max_concurrency) # 持有后台任务与镜像任务引用,避免被 GC 提前回收 self._background_tasks: set[asyncio.Task[None]] = set() self._mirror_tasks: set[asyncio.Task[None]] = set() async def submit(self, doc: DocumentInput) -> str: """登记入库任务并后台执行,立即返回 task_id""" task_id = uuid.uuid4().hex now = _utc_now_iso() self._tasks[task_id] = { "task_id": task_id, "status": IngestTaskStatus.PENDING, "created_at": now, "updated_at": now, "result": None, "error": None, } self._schedule_mirror(task_id) background = asyncio.create_task(self._run(task_id, doc)) self._background_tasks.add(background) background.add_done_callback(self._background_tasks.discard) logger.info("入库任务已登记", task_id=task_id, title=doc.title) return task_id async def get(self, task_id: str) -> dict[str, Any] | None: """查询任务状态:先查内存注册表,miss 再查 Redis 镜像,都没有返回 None""" record = self._tasks.get(task_id) if record is not None: return record if self._redis is None: return None try: return await self._redis.get_json(f"{REDIS_KEY_PREFIX}{task_id}") except Exception: logger.warning("入库任务状态读取 Redis 失败,降级为未命中", task_id=task_id, exc_info=True) return None async def wait_done(self, task_id: str, timeout: float = 30.0) -> dict[str, Any]: """轮询内存注册表直到任务进入终态(done/failed)或超时 进入终态后会等待已调度的 Redis 镜像写完再返回;超时抛 TimeoutError。 """ loop = asyncio.get_running_loop() deadline = loop.time() + timeout while True: record = self._tasks.get(task_id) if record is not None and record["status"] in TERMINAL_STATUSES: if self._mirror_tasks: await asyncio.gather(*self._mirror_tasks, return_exceptions=True) return record if loop.time() >= deadline: raise TimeoutError(f"入库任务 {task_id} 在 {timeout}s 内未进入终态") await asyncio.sleep(0.01) async def _run(self, task_id: str, doc: DocumentInput) -> None: """后台执行入库:并发限流 + 阶段状态推进 + 结果/错误落账""" async with self._semaphore: try: result = await self._ingester.ingest(doc, progress_cb=lambda stage: self._on_progress(task_id, stage)) except IngestionError as exc: # 入库已知失败:透传阶段与已产出的部分总结 self._finish_failed( task_id, { "stage": exc.stage, "message": str(exc), "partial_summary": exc.summary.model_dump(mode="json") if exc.summary is not None else None, }, ) except Exception as exc: self._finish_failed(task_id, {"stage": "unknown", "message": str(exc), "partial_summary": None}) else: self._tasks[task_id].update( status=IngestTaskStatus.DONE, updated_at=_utc_now_iso(), result=result.model_dump(mode="json"), ) self._schedule_mirror(task_id) logger.info("入库任务完成", task_id=task_id) def _finish_failed(self, task_id: str, error: dict[str, Any]) -> None: """将任务置为 failed 并记录错误信息""" self._tasks[task_id].update(status=IngestTaskStatus.FAILED, updated_at=_utc_now_iso(), error=error) self._schedule_mirror(task_id) logger.error("入库任务失败", task_id=task_id, stage=error["stage"], error=error["message"]) def _on_progress(self, task_id: str, stage: str) -> None: """Ingester 阶段回调:推进任务状态并同步镜像(同步函数,供 progress_cb 使用)""" record = self._tasks.get(task_id) if record is None: return record["status"] = stage record["updated_at"] = _utc_now_iso() self._schedule_mirror(task_id) def _schedule_mirror(self, task_id: str) -> None: """将当前任务状态快照异步镜像到 Redis(同步上下文也可调用)""" if self._redis is None: return mirror = asyncio.create_task(self._mirror_to_redis(task_id, dict(self._tasks[task_id]))) self._mirror_tasks.add(mirror) mirror.add_done_callback(self._mirror_tasks.discard) async def _mirror_to_redis(self, task_id: str, snapshot: dict[str, Any]) -> None: """写入 Redis 镜像:进行中与 done 用 ttl_done,failed 用 ttl_failed;写失败仅告警""" ttl = ( self._settings.ingest_task_ttl_failed if snapshot["status"] == IngestTaskStatus.FAILED else self._settings.ingest_task_ttl_done ) try: await self._redis.set_json(f"{REDIS_KEY_PREFIX}{task_id}", snapshot, ttl=ttl) except Exception: logger.warning("入库任务状态镜像 Redis 失败", task_id=task_id, exc_info=True)