fix: serialize vector indexing and deletion
This commit is contained in:
@@ -60,6 +60,12 @@ class DocumentRepository:
|
|||||||
async def get(self, document_id: str) -> Document | None:
|
async def get(self, document_id: str) -> Document | None:
|
||||||
return await self.session.get(Document, document_id)
|
return await self.session.get(Document, document_id)
|
||||||
|
|
||||||
|
async def is_indexing(self, document_id: str) -> bool:
|
||||||
|
status = await self.session.scalar(
|
||||||
|
select(Document.status).where(Document.id == document_id)
|
||||||
|
)
|
||||||
|
return status == "indexing"
|
||||||
|
|
||||||
async def get_job_for_document(self, document_id: str) -> IndexJob | None:
|
async def get_job_for_document(self, document_id: str) -> IndexJob | None:
|
||||||
return await self.session.scalar(
|
return await self.session.scalar(
|
||||||
select(IndexJob).where(IndexJob.document_id == document_id)
|
select(IndexJob).where(IndexJob.document_id == document_id)
|
||||||
|
|||||||
@@ -35,11 +35,13 @@ class DocumentService:
|
|||||||
session_factory: async_sessionmaker,
|
session_factory: async_sessionmaker,
|
||||||
worker: IndexWorker,
|
worker: IndexWorker,
|
||||||
vector_store: MilvusVectorStore,
|
vector_store: MilvusVectorStore,
|
||||||
|
vector_mutation_lock: asyncio.Lock,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.settings = settings
|
self.settings = settings
|
||||||
self.session_factory = session_factory
|
self.session_factory = session_factory
|
||||||
self.worker = worker
|
self.worker = worker
|
||||||
self.vector_store = vector_store
|
self.vector_store = vector_store
|
||||||
|
self.vector_mutation_lock = vector_mutation_lock
|
||||||
|
|
||||||
async def upload(self, file: UploadFile) -> Document:
|
async def upload(self, file: UploadFile) -> Document:
|
||||||
original_name = Path(file.filename or "").name
|
original_name = Path(file.filename or "").name
|
||||||
@@ -132,6 +134,7 @@ class DocumentService:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
async def _finish_delete(self, document: Document) -> None:
|
async def _finish_delete(self, document: Document) -> None:
|
||||||
|
async with self.vector_mutation_lock:
|
||||||
await asyncio.to_thread(self.vector_store.delete_document, document.id)
|
await asyncio.to_thread(self.vector_store.delete_document, document.id)
|
||||||
await asyncio.to_thread(Path(document.stored_path).unlink, True)
|
await asyncio.to_thread(Path(document.stored_path).unlink, True)
|
||||||
async with self.session_factory() as session:
|
async with self.session_factory() as session:
|
||||||
|
|||||||
@@ -24,11 +24,13 @@ class IndexWorker:
|
|||||||
session_factory: async_sessionmaker,
|
session_factory: async_sessionmaker,
|
||||||
embeddings: EmbeddingProvider,
|
embeddings: EmbeddingProvider,
|
||||||
vector_store: MilvusVectorStore,
|
vector_store: MilvusVectorStore,
|
||||||
|
vector_mutation_lock: asyncio.Lock,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.settings = settings
|
self.settings = settings
|
||||||
self.session_factory = session_factory
|
self.session_factory = session_factory
|
||||||
self.embeddings = embeddings
|
self.embeddings = embeddings
|
||||||
self.vector_store = vector_store
|
self.vector_store = vector_store
|
||||||
|
self.vector_mutation_lock = vector_mutation_lock
|
||||||
self.queue: asyncio.Queue[str] = asyncio.Queue()
|
self.queue: asyncio.Queue[str] = asyncio.Queue()
|
||||||
self._task: asyncio.Task[None] | None = None
|
self._task: asyncio.Task[None] | None = None
|
||||||
|
|
||||||
@@ -77,6 +79,7 @@ class IndexWorker:
|
|||||||
stored_path = document.stored_path
|
stored_path = document.stored_path
|
||||||
file_type = document.file_type
|
file_type = document.file_type
|
||||||
try:
|
try:
|
||||||
|
async with self.vector_mutation_lock:
|
||||||
await asyncio.to_thread(self.vector_store.delete_document, document_id)
|
await asyncio.to_thread(self.vector_store.delete_document, document_id)
|
||||||
async with self.session_factory() as session:
|
async with self.session_factory() as session:
|
||||||
await DocumentRepository(session).replace_chunks(document_id, [])
|
await DocumentRepository(session).replace_chunks(document_id, [])
|
||||||
@@ -102,7 +105,9 @@ class IndexWorker:
|
|||||||
]
|
]
|
||||||
async with self.session_factory() as session:
|
async with self.session_factory() as session:
|
||||||
await DocumentRepository(session).replace_chunks(document_id, orm_chunks)
|
await DocumentRepository(session).replace_chunks(document_id, orm_chunks)
|
||||||
await self._embed_and_upsert(prepared)
|
completed = await self._embed_and_upsert(document_id, prepared)
|
||||||
|
if not completed:
|
||||||
|
return
|
||||||
async with self.session_factory() as session:
|
async with self.session_factory() as session:
|
||||||
await DocumentRepository(session).mark_indexed(document_id, len(prepared))
|
await DocumentRepository(session).mark_indexed(document_id, len(prepared))
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
@@ -113,12 +118,17 @@ class IndexWorker:
|
|||||||
logger.exception("Document indexing failed document_id=%s", document_id)
|
logger.exception("Document indexing failed document_id=%s", document_id)
|
||||||
await self._mark_failed(document_id, "INDEXING_FAILED", "文档索引失败,请重试")
|
await self._mark_failed(document_id, "INDEXING_FAILED", "文档索引失败,请重试")
|
||||||
|
|
||||||
async def _embed_and_upsert(self, chunks: list[PreparedChunk]) -> None:
|
async def _embed_and_upsert(self, document_id: str, chunks: list[PreparedChunk]) -> bool:
|
||||||
batch_size = 20
|
batch_size = 20
|
||||||
for start in range(0, len(chunks), batch_size):
|
for start in range(0, len(chunks), batch_size):
|
||||||
batch = chunks[start : start + batch_size]
|
batch = chunks[start : start + batch_size]
|
||||||
vectors = await self.embeddings.embed_documents([chunk.content for chunk in batch])
|
vectors = await self.embeddings.embed_documents([chunk.content for chunk in batch])
|
||||||
|
async with self.vector_mutation_lock:
|
||||||
|
async with self.session_factory() as session:
|
||||||
|
if not await DocumentRepository(session).is_indexing(document_id):
|
||||||
|
return False
|
||||||
await asyncio.to_thread(self.vector_store.upsert, batch, vectors)
|
await asyncio.to_thread(self.vector_store.upsert, batch, vectors)
|
||||||
|
return True
|
||||||
|
|
||||||
async def _mark_failed(self, document_id: str, code: str, message: str) -> None:
|
async def _mark_failed(self, document_id: str, code: str, message: str) -> None:
|
||||||
async with self.session_factory() as session:
|
async with self.session_factory() as session:
|
||||||
|
|||||||
@@ -38,12 +38,20 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
|||||||
await asyncio.to_thread(vector_store.ensure_collection)
|
await asyncio.to_thread(vector_store.ensure_collection)
|
||||||
embeddings = EmbeddingProvider(settings)
|
embeddings = EmbeddingProvider(settings)
|
||||||
retriever = KnowledgeRetriever(embeddings, vector_store, database.session_factory)
|
retriever = KnowledgeRetriever(embeddings, vector_store, database.session_factory)
|
||||||
worker = IndexWorker(settings, database.session_factory, embeddings, vector_store)
|
vector_mutation_lock = asyncio.Lock()
|
||||||
|
worker = IndexWorker(
|
||||||
|
settings,
|
||||||
|
database.session_factory,
|
||||||
|
embeddings,
|
||||||
|
vector_store,
|
||||||
|
vector_mutation_lock,
|
||||||
|
)
|
||||||
document_service = DocumentService(
|
document_service = DocumentService(
|
||||||
settings,
|
settings,
|
||||||
database.session_factory,
|
database.session_factory,
|
||||||
worker,
|
worker,
|
||||||
vector_store,
|
vector_store,
|
||||||
|
vector_mutation_lock,
|
||||||
)
|
)
|
||||||
app.state.settings = settings
|
app.state.settings = settings
|
||||||
app.state.database = database
|
app.state.database = database
|
||||||
|
|||||||
Reference in New Issue
Block a user