bdd30f0a88
- deps.py: _create_redis_client 增加同步 ping 校验,Redis 不可达时正确降级为内存模式 - main.py: lifespan 复用 deps 的 UserStore 单例,避免 admin 与 API 请求实例不一致 - 前端: 强制改密弹窗(must_change_password 用户)、兼容新旧登录返回格式、markPasswordChanged - 新增 NAS SSH 部署脚本与批处理测试;gitignore 前端构建产物
308 lines
12 KiB
Python
308 lines
12 KiB
Python
"""入库异步任务管理器
|
||
|
||
将文档入库包装为后台异步任务: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)
|
||
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,
|
||
}
|
||
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,
|
||
}
|
||
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 _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
|
||
return {
|
||
"task_id": record.get("task_id"),
|
||
"status": record.get("status"),
|
||
"filename": record.get("filename"),
|
||
"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
|
||
)
|