feat: add durable document indexing pipeline

This commit is contained in:
gqt
2026-07-13 11:32:14 +08:00
parent eebcd5d6a7
commit 303116f950
12 changed files with 598 additions and 1 deletions
+130
View File
@@ -0,0 +1,130 @@
import asyncio
import hashlib
import os
from pathlib import Path
from uuid import uuid4
from fastapi import UploadFile
from sqlalchemy.ext.asyncio import async_sessionmaker
from kbqa.api.errors import AppError
from kbqa.config import Settings
from kbqa.documents.models import Document
from kbqa.documents.repository import DocumentRepository, StoredUpload
from kbqa.indexing.worker import IndexWorker
from kbqa.rag.vector_store import MilvusVectorStore
ALLOWED_EXTENSIONS = {
".pdf": ("pdf", "application/pdf"),
".md": ("md", "text/markdown"),
".markdown": ("md", "text/markdown"),
".txt": ("txt", "text/plain"),
}
class DocumentService:
def __init__(
self,
settings: Settings,
session_factory: async_sessionmaker,
worker: IndexWorker,
vector_store: MilvusVectorStore,
) -> None:
self.settings = settings
self.session_factory = session_factory
self.worker = worker
self.vector_store = vector_store
async def upload(self, file: UploadFile) -> Document:
original_name = Path(file.filename or "").name
suffix = Path(original_name).suffix.lower()
if suffix not in ALLOWED_EXTENSIONS:
raise AppError("UNSUPPORTED_FILE_TYPE", "仅支持 PDF、Markdown 和 TXT 文件", 415)
file_type, expected_mime = ALLOWED_EXTENSIONS[suffix]
document_id = str(uuid4())
target = self.settings.raw_data_dir / f"{document_id}{suffix}"
temporary = target.with_suffix(f"{suffix}.part")
maximum = self.settings.max_upload_mib * 1024 * 1024
digest = hashlib.sha256()
size = 0
prefix = b""
try:
with temporary.open("wb") as handle:
while chunk := await file.read(1024 * 1024):
size += len(chunk)
if size > maximum:
raise AppError("FILE_TOO_LARGE", "上传文件超过大小限制", 413)
if len(prefix) < 5:
prefix = (prefix + chunk)[:5]
digest.update(chunk)
handle.write(chunk)
if size == 0:
raise AppError("EMPTY_DOCUMENT", "上传文件为空", 415)
self._validate_content(file_type, prefix, temporary)
os.replace(temporary, target)
upload = StoredUpload(
document_id=document_id,
filename=original_name,
stored_path=str(target),
file_type=file_type,
mime_type=expected_mime,
file_size=size,
sha256=digest.hexdigest(),
)
async with self.session_factory() as session:
document, job = await DocumentRepository(session).create_pending(upload)
self.worker.enqueue(job.id)
return document
except Exception:
temporary.unlink(missing_ok=True)
if target.exists():
target.unlink(missing_ok=True)
raise
finally:
await file.close()
@staticmethod
def _validate_content(file_type: str, prefix: bytes, path: Path) -> None:
if file_type == "pdf":
if prefix != b"%PDF-":
raise AppError("INVALID_PDF", "PDF 文件格式无效", 415)
return
try:
path.read_text(encoding="utf-8")
except UnicodeDecodeError as exc:
raise AppError("TEXT_ENCODING_ERROR", "文本文件必须使用 UTF-8 编码", 415) from exc
async def retry(self, document_id: str) -> Document:
async with self.session_factory() as session:
repository = DocumentRepository(session)
job = await repository.requeue_failed(document_id)
document = await repository.get(document_id)
if document is None:
raise AppError("DOCUMENT_NOT_FOUND", "文档不存在", 404)
if job is None:
raise AppError("DOCUMENT_NOT_RETRYABLE", "当前文档状态不能重试", 409)
self.worker.enqueue(job.id)
return document
async def delete(self, document_id: str) -> None:
async with self.session_factory() as session:
repository = DocumentRepository(session)
document = await repository.mark_deleting(document_id)
if document is None:
raise AppError("DOCUMENT_NOT_FOUND", "文档不存在", 404)
await self._finish_delete(document)
async def recover_deletions(self) -> None:
async with self.session_factory() as session:
documents = await DocumentRepository(session).deleting_documents()
for document in documents:
try:
await self._finish_delete(document)
except Exception:
continue
async def _finish_delete(self, document: Document) -> None:
await asyncio.to_thread(self.vector_store.delete_document, document.id)
await asyncio.to_thread(Path(document.stored_path).unlink, True)
async with self.session_factory() as session:
await DocumentRepository(session).delete_document(document.id)
+1 -1
View File
@@ -21,7 +21,7 @@ async def health(request: Request) -> dict[str, object]:
components["sqlite"] = {"status": "error"} components["sqlite"] = {"status": "error"}
try: try:
await asyncio.to_thread(request.app.state.milvus_client.list_collections) await asyncio.to_thread(request.app.state.vector_store.ping)
except Exception: except Exception:
components["milvus"] = {"status": "error"} components["milvus"] = {"status": "error"}
+1
View File
@@ -0,0 +1 @@
"""Document indexing worker."""
+39
View File
@@ -0,0 +1,39 @@
from dataclasses import dataclass
from pathlib import Path
from pypdf import PdfReader
from kbqa.api.errors import AppError
@dataclass(frozen=True, slots=True)
class LoadedPage:
content: str
page_number: int | None
def load_document(path: Path, file_type: str) -> list[LoadedPage]:
if file_type == "pdf":
with path.open("rb") as handle:
if handle.read(5) != b"%PDF-":
raise AppError("INVALID_PDF", "PDF 文件格式无效", 415)
reader = PdfReader(str(path))
pages = [
LoadedPage(content=(page.extract_text() or "").strip(), page_number=index + 1)
for index, page in enumerate(reader.pages)
]
pages = [page for page in pages if page.content]
if not pages:
raise AppError(
"PDF_TEXT_NOT_EXTRACTABLE",
"PDF 中没有可提取文本,暂不支持扫描件 OCR",
415,
)
return pages
try:
content = path.read_text(encoding="utf-8").strip()
except UnicodeDecodeError as exc:
raise AppError("TEXT_ENCODING_ERROR", "文本文件必须使用 UTF-8 编码", 415) from exc
if not content:
raise AppError("EMPTY_DOCUMENT", "文档内容为空", 415)
return [LoadedPage(content=content, page_number=None)]
+34
View File
@@ -0,0 +1,34 @@
from uuid import UUID, uuid5
from langchain_text_splitters import RecursiveCharacterTextSplitter
from kbqa.indexing.loaders import LoadedPage
from kbqa.rag.types import PreparedChunk
def split_pages(
document_id: str,
pages: list[LoadedPage],
chunk_size: int,
chunk_overlap: int,
) -> list[PreparedChunk]:
splitter = RecursiveCharacterTextSplitter(
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
length_function=len,
)
chunks: list[PreparedChunk] = []
namespace = UUID(document_id)
for page in pages:
for content in splitter.split_text(page.content):
order_index = len(chunks)
chunks.append(
PreparedChunk(
id=str(uuid5(namespace, str(order_index))),
document_id=document_id,
order_index=order_index,
page_number=page.page_number,
content=content,
)
)
return chunks
+125
View File
@@ -0,0 +1,125 @@
import asyncio
import logging
from pathlib import Path
from sqlalchemy.ext.asyncio import async_sessionmaker
from kbqa.api.errors import AppError
from kbqa.config import Settings
from kbqa.documents.models import Chunk
from kbqa.documents.repository import DocumentRepository
from kbqa.indexing.loaders import load_document
from kbqa.indexing.splitter import split_pages
from kbqa.rag.embeddings import EmbeddingProvider
from kbqa.rag.types import PreparedChunk
from kbqa.rag.vector_store import MilvusVectorStore
logger = logging.getLogger(__name__)
class IndexWorker:
def __init__(
self,
settings: Settings,
session_factory: async_sessionmaker,
embeddings: EmbeddingProvider,
vector_store: MilvusVectorStore,
) -> None:
self.settings = settings
self.session_factory = session_factory
self.embeddings = embeddings
self.vector_store = vector_store
self.queue: asyncio.Queue[str] = asyncio.Queue()
self._task: asyncio.Task[None] | None = None
async def start(self) -> None:
async with self.session_factory() as session:
pending_ids = await DocumentRepository(session).requeue_unfinished()
for job_id in pending_ids:
self.queue.put_nowait(job_id)
self._task = asyncio.create_task(self._run(), name="kbqa-index-worker")
async def stop(self) -> None:
if self._task is None:
return
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
self._task = None
def enqueue(self, job_id: str) -> None:
self.queue.put_nowait(job_id)
async def _run(self) -> None:
while True:
job_id = await self.queue.get()
try:
await self._process(job_id)
except asyncio.CancelledError:
raise
except Exception:
logger.exception("Unexpected indexing worker failure job_id=%s", job_id)
finally:
self.queue.task_done()
async def _process(self, job_id: str) -> None:
async with self.session_factory() as session:
repository = DocumentRepository(session)
job = await repository.claim_job(job_id)
if job is None:
return
document = await repository.get(job.document_id)
if document is None:
return
document_id = document.id
stored_path = document.stored_path
file_type = document.file_type
try:
await asyncio.to_thread(self.vector_store.delete_document, document_id)
async with self.session_factory() as session:
await DocumentRepository(session).replace_chunks(document_id, [])
pages = await asyncio.to_thread(load_document, Path(stored_path), file_type)
prepared = split_pages(
document_id,
pages,
self.settings.chunk_size,
self.settings.chunk_overlap,
)
if not prepared:
raise AppError("EMPTY_DOCUMENT", "文档分块结果为空", 415)
orm_chunks = [
Chunk(
id=chunk.id,
document_id=chunk.document_id,
order_index=chunk.order_index,
content=chunk.content,
page_number=chunk.page_number,
char_count=len(chunk.content),
)
for chunk in prepared
]
async with self.session_factory() as session:
await DocumentRepository(session).replace_chunks(document_id, orm_chunks)
await self._embed_and_upsert(prepared)
async with self.session_factory() as session:
await DocumentRepository(session).mark_indexed(document_id, len(prepared))
except asyncio.CancelledError:
raise
except AppError as exc:
await self._mark_failed(document_id, exc.code, exc.message)
except Exception:
logger.exception("Document indexing failed document_id=%s", document_id)
await self._mark_failed(document_id, "INDEXING_FAILED", "文档索引失败,请重试")
async def _embed_and_upsert(self, chunks: list[PreparedChunk]) -> None:
batch_size = 20
for start in range(0, len(chunks), batch_size):
batch = chunks[start : start + batch_size]
vectors = await self.embeddings.embed_documents([chunk.content for chunk in batch])
await asyncio.to_thread(self.vector_store.upsert, batch, vectors)
async def _mark_failed(self, document_id: str, code: str, message: str) -> None:
async with self.session_factory() as session:
await DocumentRepository(session).mark_failed(document_id, code, message)
+25
View File
@@ -1,12 +1,18 @@
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from pathlib import Path from pathlib import Path
import asyncio
from fastapi import FastAPI from fastapi import FastAPI
from pymilvus import MilvusClient from pymilvus import MilvusClient
from kbqa.config import get_settings from kbqa.config import get_settings
from kbqa.database import Database from kbqa.database import Database
from kbqa.documents.service import DocumentService
from kbqa.indexing.worker import IndexWorker
from kbqa.rag.embeddings import EmbeddingProvider
from kbqa.rag.retriever import KnowledgeRetriever
from kbqa.rag.vector_store import MilvusVectorStore
def _ensure_runtime_directories(database_url: str, milvus_uri: str, raw_data_dir: Path) -> None: def _ensure_runtime_directories(database_url: str, milvus_uri: str, raw_data_dir: Path) -> None:
@@ -28,11 +34,30 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
) )
database = Database(settings.database_url) database = Database(settings.database_url)
milvus_client = MilvusClient(uri=settings.milvus_uri) milvus_client = MilvusClient(uri=settings.milvus_uri)
vector_store = MilvusVectorStore(milvus_client, settings.embedding_dim)
await asyncio.to_thread(vector_store.ensure_collection)
embeddings = EmbeddingProvider(settings)
retriever = KnowledgeRetriever(embeddings, vector_store, database.session_factory)
worker = IndexWorker(settings, database.session_factory, embeddings, vector_store)
document_service = DocumentService(
settings,
database.session_factory,
worker,
vector_store,
)
app.state.settings = settings app.state.settings = settings
app.state.database = database app.state.database = database
app.state.milvus_client = milvus_client app.state.milvus_client = milvus_client
app.state.vector_store = vector_store
app.state.embeddings = embeddings
app.state.retriever = retriever
app.state.index_worker = worker
app.state.document_service = document_service
await worker.start()
await document_service.recover_deletions()
try: try:
yield yield
finally: finally:
await worker.stop()
milvus_client.close() milvus_client.close()
await database.dispose() await database.dispose()
+1
View File
@@ -0,0 +1 @@
"""Retrieval-augmented generation infrastructure."""
+63
View File
@@ -0,0 +1,63 @@
from langchain_openai import OpenAIEmbeddings
from kbqa.api.errors import AppError
from kbqa.config import Settings, validate_live_settings
class EmbeddingProvider:
def __init__(self, settings: Settings) -> None:
self.settings = settings
self._client: OpenAIEmbeddings | None = None
def _get_client(self) -> OpenAIEmbeddings:
if self._client is None:
validate_live_settings(self.settings)
self._client = OpenAIEmbeddings(
api_key=self.settings.dashscope_api_key,
base_url=self.settings.dashscope_base_url,
model=self.settings.embedding_model,
dimensions=self.settings.embedding_dim,
chunk_size=20,
max_retries=2,
timeout=60.0,
check_embedding_ctx_length=False,
)
return self._client
def _validate_vectors(self, vectors: list[list[float]]) -> list[list[float]]:
if any(len(vector) != self.settings.embedding_dim for vector in vectors):
raise AppError(
"EMBEDDING_DIMENSION_MISMATCH",
"Embedding 返回的向量维度不正确",
502,
retryable=True,
)
return vectors
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
try:
vectors = await self._get_client().aembed_documents(texts)
except AppError:
raise
except Exception as exc:
raise AppError(
"EMBEDDING_UPSTREAM_ERROR",
"文档向量化失败,请稍后重试",
502,
retryable=True,
) from exc
return self._validate_vectors(vectors)
async def embed_query(self, text: str) -> list[float]:
try:
vector = await self._get_client().aembed_query(text)
except AppError:
raise
except Exception as exc:
raise AppError(
"EMBEDDING_UPSTREAM_ERROR",
"查询向量化失败,请稍后重试",
502,
retryable=True,
) from exc
return self._validate_vectors([vector])[0]
+56
View File
@@ -0,0 +1,56 @@
import asyncio
from sqlalchemy import select
from sqlalchemy.ext.asyncio import async_sessionmaker
from kbqa.documents.models import Chunk, Document
from kbqa.rag.embeddings import EmbeddingProvider
from kbqa.rag.types import SourceCandidate
from kbqa.rag.vector_store import MilvusVectorStore
class KnowledgeRetriever:
def __init__(
self,
embeddings: EmbeddingProvider,
vector_store: MilvusVectorStore,
session_factory: async_sessionmaker,
) -> None:
self.embeddings = embeddings
self.vector_store = vector_store
self.session_factory = session_factory
async def retrieve(self, query: str, top_k: int) -> list[SourceCandidate]:
vector = await self.embeddings.embed_query(query)
candidate_limit = max(top_k * 3, top_k)
hits = await asyncio.to_thread(self.vector_store.search, vector, candidate_limit)
if not hits:
return []
ids = [hit.chunk_id for hit in hits]
async with self.session_factory() as session:
rows = await session.execute(
select(Chunk, Document)
.join(Document, Document.id == Chunk.document_id)
.where(Chunk.id.in_(ids), Document.status == "indexed")
)
lookup = {chunk.id: (chunk, document) for chunk, document in rows}
sources: list[SourceCandidate] = []
for hit in hits:
record = lookup.get(hit.chunk_id)
if record is None:
continue
chunk, document = record
sources.append(
SourceCandidate(
chunk_id=chunk.id,
document_id=document.id,
filename=document.filename,
order_index=chunk.order_index,
page_number=chunk.page_number,
content=chunk.content,
score=hit.score,
)
)
if len(sources) == top_k:
break
return sources
+29
View File
@@ -0,0 +1,29 @@
from dataclasses import dataclass
@dataclass(frozen=True, slots=True)
class SourceCandidate:
chunk_id: str
document_id: str
filename: str
order_index: int
page_number: int | None
content: str
score: float
@dataclass(frozen=True, slots=True)
class PreparedChunk:
id: str
document_id: str
order_index: int
page_number: int | None
content: str
@dataclass(frozen=True, slots=True)
class VectorHit:
chunk_id: str
document_id: str
order_index: int
score: float
+94
View File
@@ -0,0 +1,94 @@
import json
from pymilvus import DataType, MilvusClient
from kbqa.rag.types import PreparedChunk, VectorHit
class MilvusVectorStore:
collection_name = "kbqa_chunks"
def __init__(self, client: MilvusClient, embedding_dim: int) -> None:
self.client = client
self.embedding_dim = embedding_dim
def ensure_collection(self) -> None:
if self.client.has_collection(self.collection_name):
return
schema = self.client.create_schema(auto_id=False, enable_dynamic_field=False)
schema.add_field(
field_name="chunk_id",
datatype=DataType.VARCHAR,
is_primary=True,
max_length=36,
)
schema.add_field(
field_name="document_id",
datatype=DataType.VARCHAR,
max_length=36,
)
schema.add_field(field_name="order_index", datatype=DataType.INT64)
schema.add_field(
field_name="vector",
datatype=DataType.FLOAT_VECTOR,
dim=self.embedding_dim,
)
index_params = self.client.prepare_index_params()
index_params.add_index(
field_name="vector",
index_type="AUTOINDEX",
metric_type="COSINE",
)
self.client.create_collection(
collection_name=self.collection_name,
schema=schema,
index_params=index_params,
)
def upsert(self, chunks: list[PreparedChunk], vectors: list[list[float]]) -> None:
if len(chunks) != len(vectors):
raise ValueError("chunks and vectors must have equal lengths")
if not chunks:
return
data = [
{
"chunk_id": chunk.id,
"document_id": chunk.document_id,
"order_index": chunk.order_index,
"vector": vector,
}
for chunk, vector in zip(chunks, vectors, strict=True)
]
self.client.upsert(collection_name=self.collection_name, data=data)
def search(self, vector: list[float], limit: int) -> list[VectorHit]:
results = self.client.search(
collection_name=self.collection_name,
data=[vector],
anns_field="vector",
limit=limit,
output_fields=["document_id", "order_index"],
search_params={"metric_type": "COSINE"},
)
hits: list[VectorHit] = []
for item in results[0] if results else []:
entity = item.get("entity", {})
hits.append(
VectorHit(
chunk_id=str(item["id"]),
document_id=str(entity["document_id"]),
order_index=int(entity["order_index"]),
score=float(item["distance"]),
)
)
return hits
def delete_document(self, document_id: str) -> None:
encoded = json.dumps(document_id)
self.client.delete(
collection_name=self.collection_name,
filter=f"document_id == {encoded}",
)
def ping(self) -> None:
self.client.list_collections()