fix deps for fastapi

This commit is contained in:
phernandez
2024-12-14 13:47:51 -06:00
parent 551a36e92b
commit 4e21811847
7 changed files with 105 additions and 257 deletions
-2
View File
@@ -1,11 +1,9 @@
"""FastAPI application for basic-memory knowledge graph API."""
from pathlib import Path
from fastapi import FastAPI
from loguru import logger
from .routers import knowledge
from ..config import ProjectConfig
# Initialize FastAPI app
-11
View File
@@ -1,11 +0,0 @@
"""FastAPI dependency functions."""
from typing import Annotated
from fastapi import Depends
from basic_memory.config import project_path
from basic_memory.deps import get_project_services
from basic_memory.services import MemoryService
MemoryServiceDep = Annotated[MemoryService, Depends(get_project_services(project_path))]
+6 -3
View File
@@ -1,8 +1,9 @@
"""Router for knowledge graph operations."""
from fastapi import APIRouter
from typing import Annotated
from fastapi import APIRouter, Depends
from basic_memory.api.deps import MemoryServiceDep
from basic_memory.deps import MemoryServiceDep
from basic_memory.schemas import (
CreateEntitiesInput, CreateEntitiesResponse,
SearchNodesInput, SearchNodesResponse,
@@ -12,7 +13,6 @@ from basic_memory.schemas import (
router = APIRouter(prefix="/knowledge", tags=["knowledge"])
@router.post("/entities", response_model=CreateEntitiesResponse)
async def create_entities(
data: CreateEntitiesInput,
@@ -22,6 +22,7 @@ async def create_entities(
entities = await memory_service.create_entities(data.entities)
return CreateEntitiesResponse(entities=[EntityOut.model_validate(entity) for entity in entities])
@router.get("/entities/{entity_id}", response_model=EntityOut)
async def get_entity(
entity_id: str,
@@ -31,6 +32,7 @@ async def get_entity(
entity = await memory_service.get_entity(entity_id)
return EntityOut.model_validate(entity)
@router.post("/relations", response_model=CreateRelationsResponse)
async def create_relations(
data: CreateRelationsInput,
@@ -40,6 +42,7 @@ async def create_relations(
relations = await memory_service.create_relations(data.relations)
return CreateRelationsResponse(relations=[RelationOut.model_validate(relation) for relation in relations])
@router.post("/observations", response_model=ObservationsOut)
async def add_observations(
data: ObservationsIn,
-1
View File
@@ -38,4 +38,3 @@ class ProjectConfig(BaseSettings):
# Load project config
config = ProjectConfig()
project_path = Path(config.path)
+17 -4
View File
@@ -1,10 +1,11 @@
"""Database configuration and initialization for basic-memory."""
from enum import Enum
from pathlib import Path
from typing import Optional
from typing import Optional, AsyncGenerator
from contextlib import asynccontextmanager
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker, AsyncEngine
from loguru import logger
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker, AsyncEngine, AsyncSession
from sqlalchemy.pool import StaticPool
from basic_memory.models import Base
@@ -14,7 +15,7 @@ class DatabaseType(Enum):
MEMORY = "memory" # In-memory SQLite for testing
FILESYSTEM = "file" # File-based SQLite for projects
def get_database_url(db_type: DatabaseType, project_path: Optional[Path] = None) -> str:
def get_database_url(project_path: Path, db_type: DatabaseType, ) -> str:
"""
Get database URL based on type and optional project path.
@@ -77,7 +78,19 @@ async def init_database(url: str, echo: bool = False) -> AsyncEngine:
return engine
@asynccontextmanager
async def get_session(engine: AsyncEngine):
async def engine(project_path: Path, db_type=DatabaseType.FILESYSTEM) -> AsyncGenerator[AsyncEngine, None]:
"""Get database engine for project with proper lifecycle management."""
url = get_database_url(project_path, db_type=db_type)
engine = await init_database(url, echo=True)
engine = await init_database(url)
logger.debug(f"engine url: {engine.url}")
try:
yield engine
finally:
await engine.dispose()
@asynccontextmanager
async def session(engine: AsyncEngine) -> AsyncGenerator[AsyncSession, None]:
"""
Get database session with proper lifecycle management.
+52 -54
View File
@@ -1,104 +1,102 @@
"""Dependency injection functions for basic-memory services."""
from contextlib import asynccontextmanager
from pathlib import Path
from typing import AsyncGenerator
from typing import AsyncGenerator, Annotated
from fastapi import Depends
from loguru import logger
from sqlalchemy.ext.asyncio import AsyncSession, AsyncEngine
from basic_memory.config import ProjectConfig
from basic_memory.config import ProjectConfig, config
from basic_memory.repository.entity_repository import EntityRepository
from basic_memory.repository.observation_repository import ObservationRepository
from basic_memory.repository.relation_repository import RelationRepository
from basic_memory.services import EntityService, ObservationService, RelationService, MemoryService
from basic_memory.db import DatabaseType, get_database_url, init_database, get_session
from basic_memory import db
def get_project_config() -> ProjectConfig:
return config
ProjectConfigDep = Annotated[ProjectConfig, Depends(get_project_config)]
def get_project_path(project_config: ProjectConfigDep) -> Path:
return Path(project_config.path)
ProjectPathDep = Annotated[Path, Depends(get_project_path)]
async def get_entity_repo(session: AsyncSession) -> EntityRepository:
async def get_engine(project_path: ProjectPathDep, db_type=db.DatabaseType.FILESYSTEM):
yield db.engine(project_path, db_type)
EngineDep = Annotated[AsyncEngine, Depends(get_engine)]
async def get_session(engine: EngineDep) :
yield db.session(engine)
AsyncSessionDep = Annotated[AsyncSession, Depends(get_session)]
async def get_entity_repo(session: AsyncSessionDep) -> EntityRepository:
"""Get an EntityRepository instance."""
return EntityRepository(session) # Entity type is handled in EntityRepository.__init__
async def get_observation_repo(session: AsyncSession) -> ObservationRepository:
EntityRepositoryDep = Annotated[EntityRepository, Depends(get_entity_repo)]
async def get_observation_repo(session: AsyncSessionDep) -> ObservationRepository:
"""Get an ObservationRepository instance."""
return ObservationRepository(session)
async def get_relation_repo(session: AsyncSession) -> RelationRepository:
ObservationRepositoryDep = Annotated[ObservationRepository, Depends(get_observation_repo)]
async def get_relation_repo(session: AsyncSessionDep) -> RelationRepository:
"""Get a RelationRepository instance."""
return RelationRepository(session)
RelationRepositoryDep = Annotated[RelationRepository, Depends(get_relation_repo)]
async def get_entity_service(
project_path: Path,
entity_repo: EntityRepository
project_path: ProjectPathDep,
entity_repo: EntityRepositoryDep
) -> EntityService:
"""Get an EntityService instance."""
return EntityService(project_path, entity_repo)
EntityServiceDep = Annotated[EntityService, Depends(get_entity_service)]
async def get_observation_service(
project_path: Path,
observation_repo: ObservationRepository
project_path: ProjectPathDep,
observation_repo: ObservationRepositoryDep
) -> ObservationService:
"""Get an ObservationService instance."""
return ObservationService(project_path, observation_repo)
ObservationServiceDep = Annotated[ObservationService, Depends(get_observation_service)]
async def get_relation_service(
project_path: Path,
relation_repo: RelationRepository
project_path: ProjectPathDep,
relation_repo: RelationRepositoryDep
) -> RelationService:
"""Get a RelationService instance."""
return RelationService(project_path, relation_repo)
RelationServiceDep = Annotated[RelationService, Depends(get_relation_service)]
@asynccontextmanager
async def get_memory_service(
project_path: Path,
entity_service: EntityService,
relation_service: RelationService,
observation_service: ObservationService
) -> MemoryService:
project_path: ProjectPathDep,
entity_service: EntityServiceDep,
relation_service: RelationServiceDep,
observation_service: ObservationServiceDep
) -> AsyncGenerator[MemoryService, None]:
"""Get a fully configured MemoryService instance."""
return MemoryService(
yield MemoryService(
project_path=project_path,
entity_service=entity_service,
relation_service=relation_service,
observation_service=observation_service
)
@asynccontextmanager
async def get_engine(project_path: Path, db_type=DatabaseType.FILESYSTEM ):
"""Get database engine for project with proper lifecycle management."""
url = get_database_url(db_type, project_path)
engine = await init_database(url)
logger.debug(f"engine url: {engine.url}")
try:
yield engine
finally:
await engine.dispose()
MemoryServiceDep = Annotated[MemoryService, Depends(get_memory_service)]
@asynccontextmanager
async def get_memory_service_session(engine: AsyncEngine, project_path: Path):
"""Get all services with proper session and lifecycle management."""
async with get_session(engine) as session:
# Create repos
entity_repo = await get_entity_repo(session)
observation_repo = await get_observation_repo(session)
relation_repo = await get_relation_repo(session)
# Create services
entity_service = await get_entity_service(project_path, entity_repo)
observation_service = await get_observation_service(project_path, observation_repo)
relation_service = await get_relation_service(project_path, relation_repo)
# Create memory service
memory_service = await get_memory_service(
project_path=project_path,
entity_service=entity_service,
relation_service=relation_service,
observation_service=observation_service
)
yield memory_service
@asynccontextmanager
async def get_project_services(project_path: Path) -> AsyncGenerator[MemoryService, None]:
"""Get all services for a project with full lifecycle management."""
async with get_engine(project_path=project_path) as engine:
async with get_memory_service_session(engine, project_path) as service_session:
yield service_session