diff --git a/backend/app/core/lifespan.py b/backend/app/core/lifespan.py new file mode 100644 index 0000000..7f3df4b --- /dev/null +++ b/backend/app/core/lifespan.py @@ -0,0 +1,50 @@ +from contextlib import asynccontextmanager + +import sqlalchemy +from fastapi import FastAPI +from llama_index.vector_stores.postgres import PGVectorStore +from openai import AsyncOpenAI + +from app.core.logging import logger + +from .config import settings +from .db import init_db + + +class AppState: + groq_client: AsyncOpenAI | None = None + vector_client: PGVectorStore | None = None + + +app_state = AppState() + + +@asynccontextmanager +async def lifespan(app: FastAPI): + + try: + init_db() + logger.info("Database initialized") + + app_state.groq_client = AsyncOpenAI( + base_url=settings.LLM_ENDPOINT, + api_key=settings.LLM_API_KEY or "missing-key", + ) + + url = sqlalchemy.make_url(settings.DATABASE_URL) + app_state.vector_client = PGVectorStore.from_params( + host=url.host, + port=str(url.port or 5432), + user=url.username, + password=url.password, + database=url.database, + table_name="regulations_vectors", + embed_dim=1024, + ) + yield + + finally: + if app_state.groq_client: + await app_state.groq_client.close() + if app_state.vector_client: + await app_state.vector_client.close() diff --git a/backend/app/regintel/rag.py b/backend/app/regintel/rag.py index 049e2e2..cf1f249 100644 --- a/backend/app/regintel/rag.py +++ b/backend/app/regintel/rag.py @@ -1,28 +1,20 @@ -from functools import lru_cache from pathlib import Path -import sqlalchemy from llama_index.core import SimpleDirectoryReader, VectorStoreIndex from llama_index.core.node_parser import SentenceSplitter from llama_index.core.schema import NodeWithScore from llama_index.embeddings.cohere import CohereEmbedding from llama_index.vector_stores.postgres import PGVectorStore -from openai import AsyncOpenAI from app.core.config import settings +from ..core.lifespan import app_state + # Detect project root (where .env lives) # rag.py is in backend/app/regintel/ BASE_DIR = Path(__file__).resolve().parent.parent.parent.parent -@lru_cache -def get_openai_client(): - return AsyncOpenAI( - base_url=settings.LLM_ENDPOINT, api_key=settings.LLM_API_KEY or "missing-key" - ) - - class VectorStoreConnection: def __init__(self): self.splitter = SentenceSplitter(chunk_size=512, chunk_overlap=60) @@ -30,17 +22,9 @@ def __init__(self): self.should_reset = False @property - def vector_store(self): - url = sqlalchemy.make_url(settings.DATABASE_URL) - return PGVectorStore.from_params( - host=url.host, - port=str(url.port or 5432), - user=url.username, - password=url.password, - database=url.database, - table_name="regulations_vectors", - embed_dim=1024, - ) + def vector_store(self) -> PGVectorStore: + assert app_state.vector_client is not None, "Vector client is not initialized" + return app_state.vector_client @property def embedding_model(self): @@ -89,8 +73,8 @@ async def vector_chat_async(self, query: str): {query} """ print(user_prompt) - - completion = await get_openai_client().chat.completions.create( + assert app_state.groq_client is not None, "Groq client is not initialized" + completion = await app_state.groq_client.chat.completions.create( model=settings.CHAT_MODEL, messages=[ {"role": "system", "content": system_prompt}, diff --git a/backend/main.py b/backend/main.py index ca1e6e6..b9b44d3 100644 --- a/backend/main.py +++ b/backend/main.py @@ -1,22 +1,11 @@ -from contextlib import asynccontextmanager - from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware from app.api import activity, ingest, manifest, validate from app.core.config import settings -from app.core.db import init_db -from app.core.logging import logger +from app.core.lifespan import lifespan from app.core.rate_limit import setup_rate_limiting - -@asynccontextmanager -async def lifespan(app: FastAPI): - init_db() - logger.info("Database initialized") - yield - - app = FastAPI(title=settings.PROJECT_NAME, lifespan=lifespan) # CORS Configuration