From 4e218118470f6963e4f03f36dd8acf7e8685dbea Mon Sep 17 00:00:00 2001 From: phernandez Date: Sat, 14 Dec 2024 13:47:51 -0600 Subject: [PATCH] fix deps for fastapi --- src/basic_memory/api/app.py | 2 - src/basic_memory/api/deps.py | 11 -- src/basic_memory/api/routers/knowledge.py | 9 +- src/basic_memory/config.py | 1 - src/basic_memory/db.py | 21 ++- src/basic_memory/deps.py | 106 ++++++----- tests/api/test_knowledge.py | 212 +++------------------- 7 files changed, 105 insertions(+), 257 deletions(-) delete mode 100644 src/basic_memory/api/deps.py diff --git a/src/basic_memory/api/app.py b/src/basic_memory/api/app.py index 4339927f..bd8b25cb 100644 --- a/src/basic_memory/api/app.py +++ b/src/basic_memory/api/app.py @@ -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 diff --git a/src/basic_memory/api/deps.py b/src/basic_memory/api/deps.py deleted file mode 100644 index 0e871602..00000000 --- a/src/basic_memory/api/deps.py +++ /dev/null @@ -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))] - diff --git a/src/basic_memory/api/routers/knowledge.py b/src/basic_memory/api/routers/knowledge.py index 09cdcade..9df1e553 100644 --- a/src/basic_memory/api/routers/knowledge.py +++ b/src/basic_memory/api/routers/knowledge.py @@ -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, diff --git a/src/basic_memory/config.py b/src/basic_memory/config.py index 68adc03d..45ce35d6 100644 --- a/src/basic_memory/config.py +++ b/src/basic_memory/config.py @@ -38,4 +38,3 @@ class ProjectConfig(BaseSettings): # Load project config config = ProjectConfig() -project_path = Path(config.path) diff --git a/src/basic_memory/db.py b/src/basic_memory/db.py index 18853d6e..d0daabd6 100644 --- a/src/basic_memory/db.py +++ b/src/basic_memory/db.py @@ -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. diff --git a/src/basic_memory/deps.py b/src/basic_memory/deps.py index 6825c99b..1ad7465c 100644 --- a/src/basic_memory/deps.py +++ b/src/basic_memory/deps.py @@ -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 \ No newline at end of file diff --git a/tests/api/test_knowledge.py b/tests/api/test_knowledge.py index 370b7550..139274cd 100644 --- a/tests/api/test_knowledge.py +++ b/tests/api/test_knowledge.py @@ -1,4 +1,5 @@ """Tests for knowledge graph API endpoints.""" +from pathlib import Path from typing import AsyncGenerator import pytest import pytest_asyncio @@ -9,49 +10,51 @@ from unittest.mock import AsyncMock from icecream import ic from loguru import logger -from basic_memory.api.routers.knowledge import router -from basic_memory.api.deps import MemoryServiceDep -from basic_memory.models import Entity, Relation, Observation +from basic_memory.models import Entity +from basic_memory.deps import get_project_services @pytest_asyncio.fixture def app(memory_service_mock: AsyncMock) -> FastAPI: """Create FastAPI test application.""" + # Lazy import router to avoid app startup issues + from basic_memory.api.routers.knowledge import router + app = FastAPI() app.include_router(router) - - async def override_get_memory_service() -> MemoryServiceDep: - async def get_service(): - yield memory_service_mock - return get_service() - - app.dependency_overrides[MemoryServiceDep] = override_get_memory_service - return app + + # Override service dependency with mock + async def memory_service_override(project_path: Path): + yield memory_service_mock + + app.dependency_overrides[get_project_services] = memory_service_override return app + @pytest_asyncio.fixture async def client(app: FastAPI) -> AsyncGenerator[AsyncClient, None]: - """Create async test client.""" - transport = ASGITransport(app=app) - base_url = "http://test" - - async with AsyncClient(transport=transport, base_url=base_url) as client: + """Create client using ASGI transport - same as CLI will use.""" + async with AsyncClient( + transport=ASGITransport(app=app), + base_url="http://test" + ) as client: yield client + @pytest_asyncio.fixture def memory_service_mock() -> AsyncMock: - """Create mock memory service.""" + """Create service mock.""" return AsyncMock() @pytest.mark.asyncio -async def test_create_entities(client: AsyncClient, memory_service_mock: AsyncMock): - """Should create new entities successfully.""" - # Setup mock response +async def test_create_entities(client: AsyncClient, memory_service_mock): + """Should create entities successfully.""" + # Setup mock entity = Entity(id="test-1", name="Test Entity", entity_type="test") memory_service_mock.create_entities.return_value = [entity] - - # Make request + + # Make request like a real client would response = await client.post("/knowledge/entities", json={ "entities": [{ "name": "Test Entity", @@ -60,168 +63,13 @@ async def test_create_entities(client: AsyncClient, memory_service_mock: AsyncMo }) logger.debug(ic(response.content)) - # Assert response + + + # Verify response assert response.status_code == 200 data = response.json() assert len(data["entities"]) == 1 assert data["entities"][0]["id"] == "test-1" - assert data["entities"][0]["name"] == "Test Entity" - + # Verify service called - memory_service_mock.create_entities.assert_called_once() - -@pytest.mark.asyncio -async def test_get_entity(client: AsyncClient, memory_service_mock: AsyncMock): - """Should retrieve a single entity by ID.""" - # Setup mock - entity = Entity( - id="test-1", - name="Test Entity", - entity_type="test", - observations=[ - Observation(id=1, content="Test observation") - ] - ) - memory_service_mock.get_entity.return_value = entity - - # Make request - response = await client.get("/knowledge/entities/test-1") - - # Assert response - assert response.status_code == 200 - data = response.json() - assert data["id"] == "test-1" - assert data["name"] == "Test Entity" - assert len(data["observations"]) == 1 - - # Verify service called - memory_service_mock.get_entity.assert_called_once_with("test-1") - - -@pytest.mark.asyncio -async def test_create_relations(client: AsyncClient, memory_service_mock: AsyncMock): - """Should create relations between entities.""" - # Setup mock - relation = Relation( - id=1, - from_entity_id="test-1", - to_entity_id="test-2", - relation_type="related_to" - ) - memory_service_mock.create_relations.return_value = [relation] - - # Make request - response = await client.post("/knowledge/relations", json={ - "relations": [{ - "from_entity_id": "test-1", - "to_entity_id": "test-2", - "relation_type": "related_to" - }] - }) - - # Assert response - assert response.status_code == 200 - data = response.json() - assert len(data["relations"]) == 1 - assert data["relations"][0]["from_entity_id"] == "test-1" - assert data["relations"][0]["to_entity_id"] == "test-2" - - # Verify service called - memory_service_mock.create_relations.assert_called_once() - - -@pytest.mark.asyncio -async def test_add_observations(client: AsyncClient, memory_service_mock: AsyncMock): - """Should add observations to an entity.""" - # Setup mock - observations = [ - Observation(id=1, content="Test observation 1"), - Observation(id=2, content="Test observation 2") - ] - memory_service_mock.add_observations.return_value = observations - - # Make request - response = await client.post("/knowledge/observations", json={ - "entity_id": "test-1", - "observations": [ - "Test observation 1", - "Test observation 2" - ] - }) - - # Assert response - assert response.status_code == 200 - data = response.json() - assert data["entity_id"] == "test-1" - assert len(data["observations"]) == 2 - assert data["observations"][0]["content"] == "Test observation 1" - - # Verify service called - memory_service_mock.add_observations.assert_called_once() - - -@pytest.mark.asyncio -async def test_search_nodes(client: AsyncClient, memory_service_mock: AsyncMock): - """Should search for entities in the knowledge graph.""" - # Setup mock - entity = Entity(id="test-1", name="Test Entity", entity_type="test") - memory_service_mock.search_nodes.return_value = [entity] - - # Make request - response = await client.post("/knowledge/search", json={ - "query": "test" - }) - - # Assert response - assert response.status_code == 200 - data = response.json() - assert data["query"] == "test" - assert len(data["matches"]) == 1 - assert data["matches"][0]["id"] == "test-1" - - # Verify service called - memory_service_mock.search_nodes.assert_called_once_with("test") - - -@pytest.mark.asyncio -async def test_create_entities_validation(client: AsyncClient): - """Should validate entity creation input.""" - # Make request with invalid data - response = await client.post("/knowledge/entities", json={ - "entities": [{ - "name": "", # Empty name should fail validation - "entity_type": "test" - }] - }) - - # Assert validation error - assert response.status_code == 422 - - -@pytest.mark.asyncio -async def test_create_relations_validation(client: AsyncClient): - """Should validate relation creation input.""" - # Make request with invalid data - response = await client.post("/knowledge/relations", json={ - "relations": [{ - "from_entity_id": "", # Empty ID should fail validation - "to_entity_id": "test-2", - "relation_type": "related_to" - }] - }) - - # Assert validation error - assert response.status_code == 422 - - -@pytest.mark.asyncio -async def test_add_observations_validation(client: AsyncClient): - """Should validate observation input.""" - # Make request with invalid data - response = await client.post("/knowledge/observations", json={ - "entity_id": "", # Empty ID should fail validation - "observations": [] # Empty observations should fail validation - }) - - # Assert validation error - assert response.status_code == 422 + memory_service_mock.create_entities.assert_called_once() \ No newline at end of file