diff --git a/src/basic_memory/deps.py b/src/basic_memory/deps.py index 1ad7465c..9dccf51e 100644 --- a/src/basic_memory/deps.py +++ b/src/basic_memory/deps.py @@ -23,13 +23,15 @@ def get_project_path(project_config: ProjectConfigDep) -> Path: ProjectPathDep = Annotated[Path, Depends(get_project_path)] -async def get_engine(project_path: ProjectPathDep, db_type=db.DatabaseType.FILESYSTEM): - yield db.engine(project_path, db_type) +async def get_engine(project_path: ProjectPathDep, db_type=db.DatabaseType.FILESYSTEM) -> AsyncGenerator[AsyncEngine, None]: + async with db.engine(project_path, db_type) as engine: + yield engine EngineDep = Annotated[AsyncEngine, Depends(get_engine)] -async def get_session(engine: EngineDep) : - yield db.session(engine) +async def get_session(engine: EngineDep) -> AsyncGenerator[AsyncSession, None]: + async with db.session(engine) as session: + yield session AsyncSessionDep = Annotated[AsyncSession, Depends(get_session)] @@ -80,7 +82,7 @@ async def get_relation_service( RelationServiceDep = Annotated[RelationService, Depends(get_relation_service)] @asynccontextmanager -async def get_memory_service( +async def memory_service( project_path: ProjectPathDep, entity_service: EntityServiceDep, relation_service: RelationServiceDep, @@ -94,6 +96,15 @@ async def get_memory_service( observation_service=observation_service ) +async def get_memory_service( + project_path: ProjectPathDep, + entity_service: EntityServiceDep, + relation_service: RelationServiceDep, + observation_service: ObservationServiceDep +) -> AsyncGenerator[MemoryService, None]: + async with memory_service(project_path, entity_service, relation_service, observation_service) as service: + yield service + MemoryServiceDep = Annotated[MemoryService, Depends(get_memory_service)] diff --git a/tests/api/test_knowledge.py b/tests/api/test_knowledge.py index 139274cd..276200ad 100644 --- a/tests/api/test_knowledge.py +++ b/tests/api/test_knowledge.py @@ -10,12 +10,12 @@ from unittest.mock import AsyncMock from icecream import ic from loguru import logger +from basic_memory.deps import get_project_config, get_engine from basic_memory.models import Entity -from basic_memory.deps import get_project_services @pytest_asyncio.fixture -def app(memory_service_mock: AsyncMock) -> FastAPI: +def app(test_config, engine) -> FastAPI: """Create FastAPI test application.""" # Lazy import router to avoid app startup issues from basic_memory.api.routers.knowledge import router @@ -23,11 +23,8 @@ def app(memory_service_mock: AsyncMock) -> FastAPI: app = FastAPI() app.include_router(router) - # 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 + app.dependency_overrides[get_project_config] = lambda: test_config + app.dependency_overrides[get_engine] = lambda: engine return app @@ -41,18 +38,9 @@ async def client(app: FastAPI) -> AsyncGenerator[AsyncClient, None]: yield client -@pytest_asyncio.fixture -def memory_service_mock() -> AsyncMock: - """Create service mock.""" - return AsyncMock() - - @pytest.mark.asyncio -async def test_create_entities(client: AsyncClient, memory_service_mock): +async def test_create_entities(client: AsyncClient): """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 like a real client would response = await client.post("/knowledge/entities", json={ @@ -69,7 +57,4 @@ async def test_create_entities(client: AsyncClient, memory_service_mock): assert response.status_code == 200 data = response.json() assert len(data["entities"]) == 1 - assert data["entities"][0]["id"] == "test-1" - - # Verify service called - memory_service_mock.create_entities.assert_called_once() \ No newline at end of file + assert data["entities"][0]["id"] == "test/test_entity" diff --git a/tests/conftest.py b/tests/conftest.py index 0a3d3f3f..e441fc4b 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -7,7 +7,7 @@ from loguru import logger from sqlalchemy import text from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker, AsyncSession -from basic_memory import deps +from basic_memory import deps, db from basic_memory.db import DatabaseType from basic_memory.models import Base from basic_memory.repository.entity_repository import EntityRepository @@ -20,6 +20,8 @@ from basic_memory.deps import ( from basic_memory.schemas import EntityIn from basic_memory.debug_utils import dump_db_state from basic_memory.config import ProjectConfig +from basic_memory.services import MemoryService + @pytest_asyncio.fixture def anyio_backend(): @@ -38,7 +40,7 @@ def test_config(tmp_path): @pytest_asyncio.fixture(scope="function") async def engine(test_config): """Create an async engine using in-memory SQLite database""" - async with get_engine(project_path=test_config.path, db_type=DatabaseType.MEMORY) as engine: + async with db.engine(project_path=test_config.path, db_type=DatabaseType.MEMORY) as engine: yield engine @@ -96,7 +98,7 @@ async def memory_service( observation_service ): """Fixture providing initialized MemoryService.""" - return await get_memory_service( + return MemoryService( test_project_path, entity_service, relation_service, @@ -117,9 +119,9 @@ async def sample_entity(entity_repository: EntityRepository): @pytest_asyncio.fixture async def test_entity(entity_service): """Create a test entity for reuse in tests.""" - entity_data = EntityIn( + entity_data = EntityIn( # pyright: ignore [reportCallIssue] name="Test Entity", - entity_type="test", + entity_type="test", # pyright: ignore [reportCallIssue] ) return await entity_service.create_entity(entity_data)