From 566a796d68ed01b1c9420a954067bb6474a34ce4 Mon Sep 17 00:00:00 2001 From: phernandez Date: Mon, 30 Dec 2024 14:08:21 -0600 Subject: [PATCH] fix sqlalchemy dep loading in deps.py --- src/basic_memory/api/app.py | 25 +++++++++++++------- src/basic_memory/db.py | 46 ++++++++++++++++++++++++++++++++++--- src/basic_memory/deps.py | 19 ++++++++------- 3 files changed, 69 insertions(+), 21 deletions(-) diff --git a/src/basic_memory/api/app.py b/src/basic_memory/api/app.py index b48a1939..6328d0ae 100644 --- a/src/basic_memory/api/app.py +++ b/src/basic_memory/api/app.py @@ -1,25 +1,34 @@ """FastAPI application for basic-memory knowledge graph API.""" +from contextlib import asynccontextmanager + from fastapi import FastAPI from loguru import logger +from basic_memory import db from .routers import documents from .routers import knowledge from .routers import discovery + +@asynccontextmanager +async def lifespan(app: FastAPI): + """Lifecycle manager for the FastAPI app.""" + logger.info("Starting Basic Memory API") + yield + logger.info("Shutting down Basic Memory API") + await db.shutdown_db() + + # Initialize FastAPI app app = FastAPI( - title="Basic Memory API", description="Knowledge graph API for basic-memory", version="0.1.0" + title="Basic Memory API", + description="Knowledge graph API for basic-memory", + version="0.1.0", + lifespan=lifespan, ) # Include routers app.include_router(knowledge.router) app.include_router(documents.router) app.include_router(discovery.router) - - -# Add startup event -@app.on_event("startup") -async def startup_event(): - """Log when the API starts.""" - logger.info("Starting Basic Memory API") diff --git a/src/basic_memory/db.py b/src/basic_memory/db.py index 21d51754..9ec02ce5 100644 --- a/src/basic_memory/db.py +++ b/src/basic_memory/db.py @@ -2,7 +2,7 @@ import asyncio from contextlib import asynccontextmanager from enum import Enum, auto from pathlib import Path -from typing import AsyncGenerator +from typing import AsyncGenerator, Optional from loguru import logger from sqlalchemy import text @@ -17,6 +17,11 @@ from sqlalchemy.ext.asyncio import ( from basic_memory.models import Base +# Module level state +_engine: Optional[AsyncEngine] = None +_session_maker: Optional[async_sessionmaker[AsyncSession]] = None + + class DatabaseType(Enum): """Types of supported databases.""" @@ -72,13 +77,48 @@ async def init_db(session: AsyncSession): await session.commit() +async def get_or_create_db( + db_path: Path, + db_type: DatabaseType = DatabaseType.FILESYSTEM, +) -> tuple[AsyncEngine, async_sessionmaker[AsyncSession]]: + """Get or create database engine and session maker.""" + global _engine, _session_maker + + if _engine is None: + db_url = DatabaseType.get_db_url(db_path, db_type) + logger.debug(f"Creating engine for db_url: {db_url}") + _engine = create_async_engine(db_url, connect_args={"check_same_thread": False}) + _session_maker = async_sessionmaker(_engine, expire_on_commit=False) + + # Initialize database + logger.debug("Initializing database...") + async with scoped_session(_session_maker) as db_session: + await init_db(db_session) + + return _engine, _session_maker + + +async def shutdown_db(): + """Clean up database connections.""" + global _engine, _session_maker + + if _engine: + await _engine.dispose() + _engine = None + _session_maker = None + + @asynccontextmanager async def engine_session_factory( db_path: Path, db_type: DatabaseType = DatabaseType.FILESYSTEM, init: bool = True, ) -> AsyncGenerator[tuple[AsyncEngine, async_sessionmaker[AsyncSession]], None]: - """Create engine and session factory.""" + """Create engine and session factory. + + Note: This is primarily used for testing where we want a fresh database + for each test. For production use, use get_or_create_db() instead. + """ db_url = DatabaseType.get_db_url(db_path, db_type) logger.debug(f"Creating engine for db_url: {db_url}") engine = create_async_engine(db_url, connect_args={"check_same_thread": False}) @@ -92,4 +132,4 @@ async def engine_session_factory( yield engine, factory finally: - await engine.dispose() + await engine.dispose() \ No newline at end of file diff --git a/src/basic_memory/deps.py b/src/basic_memory/deps.py index 191d206c..361f53d4 100644 --- a/src/basic_memory/deps.py +++ b/src/basic_memory/deps.py @@ -1,6 +1,6 @@ """Dependency injection functions for basic-memory services.""" -from typing import Annotated, AsyncGenerator +from typing import Annotated from fastapi import Depends from sqlalchemy.ext.asyncio import ( @@ -41,13 +41,10 @@ ProjectConfigDep = Annotated[ProjectConfig, Depends(get_project_config)] async def get_engine_factory( - project_config: ProjectConfigDep, db_type=DatabaseType.FILESYSTEM -) -> AsyncGenerator[tuple[AsyncEngine, async_sessionmaker[AsyncSession]], None]: - async with db.engine_session_factory(db_path=project_config.database_path, db_type=db_type) as ( - engine, - session_maker, - ): - yield engine, session_maker + project_config: ProjectConfigDep, +) -> tuple[AsyncEngine, async_sessionmaker[AsyncSession]]: + """Get engine and session maker.""" + return await db.get_or_create_db(project_config.database_path) EngineFactoryDep = Annotated[ @@ -56,13 +53,14 @@ EngineFactoryDep = Annotated[ async def get_session_maker(engine_factory: EngineFactoryDep) -> async_sessionmaker[AsyncSession]: - """Get session maker for tests.""" + """Get session maker.""" _, session_maker = engine_factory return session_maker SessionMakerDep = Annotated[async_sessionmaker, Depends(get_session_maker)] + ## repositories @@ -105,6 +103,7 @@ async def get_document_repository( DocumentRepositoryDep = Annotated[DocumentRepository, Depends(get_document_repository)] + ## services @@ -177,4 +176,4 @@ async def get_knowledge_service( ) -KnowledgeServiceDep = Annotated[KnowledgeService, Depends(get_knowledge_service)] +KnowledgeServiceDep = Annotated[KnowledgeService, Depends(get_knowledge_service)] \ No newline at end of file