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
+30 -182
View File
@@ -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()