75 lines
2.5 KiB
Python
75 lines
2.5 KiB
Python
from collections.abc import AsyncIterator
|
|
from contextlib import asynccontextmanager
|
|
from pathlib import Path
|
|
import asyncio
|
|
|
|
from fastapi import FastAPI
|
|
from pymilvus import MilvusClient
|
|
|
|
from kbqa.config import get_settings
|
|
from kbqa.chat.service import ChatService
|
|
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:
|
|
raw_data_dir.mkdir(parents=True, exist_ok=True)
|
|
if database_url.startswith("sqlite+aiosqlite:///./"):
|
|
Path(database_url.removeprefix("sqlite+aiosqlite:///./")).parent.mkdir(
|
|
parents=True, exist_ok=True
|
|
)
|
|
Path(milvus_uri).parent.mkdir(parents=True, exist_ok=True)
|
|
|
|
|
|
@asynccontextmanager
|
|
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
|
settings = get_settings()
|
|
_ensure_runtime_directories(
|
|
settings.database_url,
|
|
settings.milvus_uri,
|
|
settings.raw_data_dir,
|
|
)
|
|
database = Database(settings.database_url)
|
|
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)
|
|
vector_mutation_lock = asyncio.Lock()
|
|
worker = IndexWorker(
|
|
settings,
|
|
database.session_factory,
|
|
embeddings,
|
|
vector_store,
|
|
vector_mutation_lock,
|
|
)
|
|
document_service = DocumentService(
|
|
settings,
|
|
database.session_factory,
|
|
worker,
|
|
vector_store,
|
|
vector_mutation_lock,
|
|
)
|
|
chat_service = ChatService(settings, database.session_factory, retriever)
|
|
app.state.settings = settings
|
|
app.state.database = database
|
|
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
|
|
app.state.chat_service = chat_service
|
|
await worker.start()
|
|
await document_service.recover_deletions()
|
|
try:
|
|
yield
|
|
finally:
|
|
await worker.stop()
|
|
milvus_client.close()
|
|
await database.dispose()
|