Files
QMDSearch/app/core/ingest_tasks.py
T
kplam dce9e31bde feat: 新增多格式文件上传入库与认证体系
- 新增 JWT 认证模块,支持登录/注册/用户管理
- 新增文件上传接口,支持 .txt/.md/.html/.pdf/.docx 等格式解析入库
- 新增检索结果 AI 总结功能
- 新增文本去重缓存机制
- 新增全局认证夹具简化测试
- 新增配置项与环境变量支持
- 完善文档与测试覆盖
2026-07-30 10:30:15 +08:00

231 lines
9.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""入库异步任务管理器
将文档入库包装为后台异步任务:submit 登记任务并立即返回 task_id,
后台受并发上限控制执行 Ingester.ingest,并按阶段推进任务状态。
内存注册表为主(记录 status/created_at/updated_at/result/error),
Redis 为持久镜像(key: ingest_task:{task_id}),每次状态迁移同步写入;
Redis 不可用或写入失败仅记录 warning,不影响任务执行。
"""
import asyncio
import hashlib
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:"
# Redis 文本去重 key 前缀(value 为已入库文档的 IngestionResult JSON
DEDUP_KEY_PREFIX = "dedup:sha256:"
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
文本去重:基于 doc.text 的 sha256 在 Redis 中查重;命中则直接复用旧
IngestionResult(仅置 deduplicated=True),不重跑流水线;未命中走原
异步入库流程,完成后写入去重记录供后续命中复用。Redis 不可用时跳过
去重,按原流程执行,不影响主流程。
"""
task_id = uuid.uuid4().hex
now = _utc_now_iso()
text_hash = hashlib.sha256(doc.text.encode("utf-8")).hexdigest()
# 1. 去重命中:直接置 done,复用旧结果,不调 _run
dedup_record = await self._lookup_dedup(text_hash)
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,
}
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,
}
self._schedule_mirror(task_id)
background = asyncio.create_task(self._run(task_id, doc, text_hash))
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, text_hash: str) -> 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 self._record_dedup(text_hash, result_dict)
logger.info("入库任务完成", task_id=task_id)
async def _lookup_dedup(self, text_hash: str) -> dict[str, Any] | None:
"""查询文本去重记录;Redis 不可用或异常时降级为未命中"""
if self._redis is None:
return None
try:
return await self._redis.get_json(f"{DEDUP_KEY_PREFIX}{text_hash}")
except Exception:
logger.warning("去重记录查询失败,降级为未命中", text_hash=text_hash, exc_info=True)
return None
async def _record_dedup(self, text_hash: str, result_dict: dict[str, Any]) -> None:
"""写入文本去重记录(含完整 IngestionResult),供后续命中复用;失败仅告警"""
if self._redis is None:
return
try:
await self._redis.set_json(
f"{DEDUP_KEY_PREFIX}{text_hash}",
result_dict,
ttl=self._settings.ingest_task_ttl_done,
)
except Exception:
logger.warning("去重记录写入失败", text_hash=text_hash, exc_info=True)
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_donefailed 用 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)