"""入库异步任务管理器 将文档入库包装为后台异步任务:submit 登记任务并立即返回 task_id, 后台受并发上限控制执行 Ingester.ingest,并按阶段推进任务状态。 内存注册表为主(记录 status/created_at/updated_at/result/error), Redis 为持久镜像(key: ingest_task:{task_id}),每次状态迁移同步写入; Redis 不可用或写入失败仅记录 warning,不影响任务执行。 文本去重委托给 app.core.dedup.get_dedup_strategy(),按 runtime_settings.dedup 选择策略(none/sha256/simhash)。 """ 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.dedup import get_dedup_strategy 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 文本去重:按 runtime_settings.dedup 选择策略(none/sha256/simhash); 命中则直接复用旧 IngestionResult(仅置 deduplicated=True),不重跑流水线; 未命中走原异步入库流程,完成后写入去重记录供后续命中复用。 Redis 不可用或策略=none 时跳过,按原流程执行。 """ task_id = uuid.uuid4().hex now = _utc_now_iso() filename = self._extract_filename(doc) source = self._extract_source(doc) dedup = get_dedup_strategy(self._redis) # 1. 去重命中:直接置 done,复用旧结果,不调 _run dedup_record = await dedup.lookup(doc.text) if dedup_record is not None: result_dict = dict(dedup_record) result_dict["deduplicated"] = True self._tasks[task_id] = { "task_id": task_id, "status": IngestTaskStatus.DONE, "created_at": now, "updated_at": now, "result": result_dict, "error": None, "filename": filename, "source": source, } self._schedule_mirror(task_id) logger.info( "入库任务命中去重,复用既有文档", task_id=task_id, document_id=result_dict.get("document_id"), ) return task_id # 2. 未命中:登记 pending 并后台跑流水线 self._tasks[task_id] = { "task_id": task_id, "status": IngestTaskStatus.PENDING, "created_at": now, "updated_at": now, "result": None, "error": None, "filename": filename, "source": source, } self._schedule_mirror(task_id) background = asyncio.create_task(self._run(task_id, doc, dedup)) 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 list_tasks(self, limit: int = 20) -> list[dict[str, Any]]: """合并内存注册表与 Redis 镜像,去重,按 updated_at 降序,截断 limit 每项提取 {task_id, status, filename, created_at, updated_at, doc_id}: doc_id 从 done 任务的 result.document_id 提取,其余状态为 None。 """ records: dict[str, dict[str, Any]] = {} # 1. 内存注册表(主) for task_id, record in self._tasks.items(): records[task_id] = record # 2. Redis 镜像补充内存中没有的(内存未命中的任务,如重启后只存在 Redis 的历史记录) for task_id, record in await self._scan_redis_tasks(): records.setdefault(task_id, record) # 3. 按 updated_at 降序,截断 limit sorted_records = sorted( records.values(), key=lambda r: r.get("updated_at") or "", reverse=True, )[:limit] # 4. 提取展示字段 return [self._to_list_item(r) for r in sorted_records] async def _scan_redis_tasks(self) -> list[tuple[str, dict[str, Any]]]: """扫描 Redis 中 ingest_task:* 键,返回 (task_id, record) 列表 RedisCache 未暴露公开扫描接口,借道底层 _get_client().scan_iter; 任何异常降级为空列表,不影响 list_tasks 主流程。 """ if self._redis is None: return [] get_client = getattr(self._redis, "_get_client", None) if not callable(get_client): return [] try: client = get_client() tasks: list[tuple[str, dict[str, Any]]] = [] async for key in client.scan_iter(match=f"{REDIS_KEY_PREFIX}*"): key_str = key if isinstance(key, str) else key.decode("utf-8", "replace") task_id = key_str[len(REDIS_KEY_PREFIX) :] record = await self._redis.get_json(key_str) if isinstance(record, dict): tasks.append((task_id, record)) return tasks except Exception: logger.warning("扫描 Redis 任务键失败", exc_info=True) return [] @staticmethod def _extract_filename(doc: DocumentInput) -> str | None: """从文档输入提取展示用文件名:优先 metadata.original_filename,其次 title""" return doc.metadata.get("original_filename") or doc.title or None @staticmethod def _extract_source(doc: DocumentInput) -> dict[str, Any]: """从文档输入提取重试所需的源信息(供 retry 复用,含文件路径/大小)""" return { "text": doc.text, "title": doc.title, "source": doc.source, "metadata": dict(doc.metadata), } async def retry(self, task_id: str) -> str | None: """重试失败/已完成的入库任务:用原 source 重新提交一个新任务 返回新 task_id;任务不存在或缺少 source 时返回 None。 """ record = await self.get(task_id) if record is None: return None source = record.get("source") if not isinstance(source, dict) or not source.get("text"): return None doc = DocumentInput( text=source["text"], title=source.get("title") or "", source=source.get("source") or "", metadata=source.get("metadata") or {}, ) return await self.submit(doc) async def delete(self, task_id: str) -> bool: """删除入库任务(内存注册表 + Redis 镜像),不存在返回 False""" existed = self._tasks.pop(task_id, None) is not None if self._redis is not None: try: get_client = getattr(self._redis, "_get_client", None) if callable(get_client): await get_client().delete(f"{REDIS_KEY_PREFIX}{task_id}") existed = True except Exception: logger.warning("删除 Redis 任务镜像失败", task_id=task_id, exc_info=True) return existed @staticmethod def _to_list_item(record: dict[str, Any]) -> dict[str, Any]: """从完整任务记录提取列表展示字段""" result = record.get("result") doc_id = result.get("document_id") if isinstance(result, dict) else None source = record.get("source") or {} metadata = source.get("metadata") or {} return { "task_id": record.get("task_id"), "status": record.get("status"), "filename": record.get("filename"), "title": source.get("title"), "source": source.get("source"), "size_bytes": metadata.get("original_size_bytes"), "created_at": record.get("created_at"), "updated_at": record.get("updated_at"), "doc_id": doc_id, } 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, dedup: Any) -> 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: result_dict = result.model_dump(mode="json") self._tasks[task_id].update( status=IngestTaskStatus.DONE, updated_at=_utc_now_iso(), result=result_dict, ) self._schedule_mirror(task_id) await dedup.record(doc.text, result_dict) 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 )