fix: refactor db schema migrate handling · basicmachines-co/basic-memory@ca632be · GitHub
Skip to content

Commit ca632be

Browse files
author
phernandez
committed
fix: refactor db schema migrate handling
1 parent a491c2b commit ca632be

8 files changed

Lines changed: 85 additions & 103 deletions

File tree

.gitignore

Lines changed: 1 addition & 0 deletions

src/basic_memory/api/app.py

Lines changed: 1 addition & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -10,35 +10,13 @@
1010
from basic_memory import db
1111
from basic_memory.config import config as app_config
1212
from basic_memory.api.routers import knowledge, search, memory, resource
13-
from alembic import command
14-
from alembic.config import Config
15-
16-
from basic_memory.db import DatabaseType
17-
from basic_memory.repository.search_repository import SearchRepository
18-
19-
20-
async def run_migrations(): # pragma: no cover
21-
"""Run any pending alembic migrations."""
22-
logger.info("Running database migrations...")
23-
try:
24-
config = Config("alembic.ini")
25-
command.upgrade(config, "head")
26-
logger.info("Migrations completed successfully")
27-
28-
_, session_maker = await db.get_or_create_db(
29-
app_config.database_path, DatabaseType.FILESYSTEM
30-
)
31-
await SearchRepository(session_maker).init_search_index()
32-
except Exception as e:
33-
logger.error(f"Error running migrations: {e}")
34-
raise
3513

3614

3715
@asynccontextmanager
3816
async def lifespan(app: FastAPI): # pragma: no cover
3917
"""Lifecycle manager for the FastAPI app."""
4018
logger.info(f"Starting Basic Memory API {basic_memory.__version__}")
41-
await run_migrations()
19+
await db.run_migrations(app_config)
4220
yield
4321
logger.info("Shutting down Basic Memory API")
4422
await db.shutdown_db()

src/basic_memory/cli/app.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,13 @@
1+
import asyncio
2+
13
import typer
24

5+
from basic_memory import db
6+
from basic_memory.config import config
7+
from basic_memory.utils import setup_logging
8+
9+
setup_logging(log_file=".basic-memory/basic-memory-cli.log") # pragma: no cover
10+
11+
asyncio.run(db.run_migrations(config))
12+
313
app = typer.Typer()

src/basic_memory/cli/commands/db.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,4 +22,4 @@ def reset(
2222
from basic_memory.cli.commands.sync import sync
2323

2424
logger.info("Rebuilding search index from filesystem...")
25-
asyncio.run(sync()) # pyright: ignore
25+
sync(watch=False) # pyright: ignore

src/basic_memory/cli/commands/status.py

Lines changed: 5 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -25,13 +25,11 @@ async def get_file_change_scanner(
2525
db_type=DatabaseType.FILESYSTEM,
2626
) -> FileChangeScanner: # pragma: no cover
2727
"""Get sync service instance."""
28-
async with db.engine_session_factory(db_path=config.database_path, db_type=db_type) as (
29-
engine,
30-
session_maker,
31-
):
32-
entity_repository = EntityRepository(session_maker)
33-
file_change_scanner = FileChangeScanner(entity_repository)
34-
return file_change_scanner
28+
_, session_maker = await db.get_or_create_db(db_path=config.database_path, db_type=db_type)
29+
30+
entity_repository = EntityRepository(session_maker)
31+
file_change_scanner = FileChangeScanner(entity_repository)
32+
return file_change_scanner
3533

3634

3735
def add_files_to_tree(

src/basic_memory/cli/commands/sync.py

Lines changed: 43 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -39,50 +39,48 @@ class ValidationIssue:
3939
error: str
4040

4141

42-
async def get_sync_service(db_type=DatabaseType.FILESYSTEM): # pragma: no cover
42+
async def get_sync_service(): # pragma: no cover
4343
"""Get sync service instance with all dependencies."""
44-
async with db.engine_session_factory(db_path=config.database_path, db_type=db_type) as (
45-
engine,
46-
session_maker,
47-
):
48-
entity_parser = EntityParser(config.home)
49-
markdown_processor = MarkdownProcessor(entity_parser)
50-
file_service = FileService(config.home, markdown_processor)
51-
52-
# Initialize repositories
53-
entity_repository = EntityRepository(session_maker)
54-
observation_repository = ObservationRepository(session_maker)
55-
relation_repository = RelationRepository(session_maker)
56-
search_repository = SearchRepository(session_maker)
57-
58-
# Initialize services
59-
search_service = SearchService(search_repository, entity_repository, file_service)
60-
link_resolver = LinkResolver(entity_repository, search_service)
61-
62-
# Initialize scanner
63-
file_change_scanner = FileChangeScanner(entity_repository)
64-
65-
# Initialize services
66-
entity_service = EntityService(
67-
entity_parser,
68-
entity_repository,
69-
observation_repository,
70-
relation_repository,
71-
file_service,
72-
link_resolver,
73-
)
74-
75-
# Create sync service
76-
sync_service = SyncService(
77-
scanner=file_change_scanner,
78-
entity_service=entity_service,
79-
entity_parser=entity_parser,
80-
entity_repository=entity_repository,
81-
relation_repository=relation_repository,
82-
search_service=search_service,
83-
)
84-
85-
return sync_service
44+
_, session_maker = await db.get_or_create_db(db_path=config.database_path, db_type=db.DatabaseType.FILESYSTEM)
45+
46+
entity_parser = EntityParser(config.home)
47+
markdown_processor = MarkdownProcessor(entity_parser)
48+
file_service = FileService(config.home, markdown_processor)
49+
50+
# Initialize repositories
51+
entity_repository = EntityRepository(session_maker)
52+
observation_repository = ObservationRepository(session_maker)
53+
relation_repository = RelationRepository(session_maker)
54+
search_repository = SearchRepository(session_maker)
55+
56+
# Initialize services
57+
search_service = SearchService(search_repository, entity_repository, file_service)
58+
link_resolver = LinkResolver(entity_repository, search_service)
59+
60+
# Initialize scanner
61+
file_change_scanner = FileChangeScanner(entity_repository)
62+
63+
# Initialize services
64+
entity_service = EntityService(
65+
entity_parser,
66+
entity_repository,
67+
observation_repository,
68+
relation_repository,
69+
file_service,
70+
link_resolver,
71+
)
72+
73+
# Create sync service
74+
sync_service = SyncService(
75+
scanner=file_change_scanner,
76+
entity_service=entity_service,
77+
entity_parser=entity_parser,
78+
entity_repository=entity_repository,
79+
relation_repository=relation_repository,
80+
search_service=search_service,
81+
)
82+
83+
return sync_service
8684

8785

8886
def group_issues_by_directory(issues: List[ValidationIssue]) -> Dict[str, List[ValidationIssue]]:
@@ -154,6 +152,8 @@ def display_detailed_sync_results(knowledge: SyncReport):
154152

155153
async def run_sync(verbose: bool = False, watch: bool = False):
156154
"""Run sync operation."""
155+
156+
157157
sync_service = await get_sync_service()
158158

159159
# Start watching if requested

src/basic_memory/db.py

Lines changed: 21 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,10 @@
44
from pathlib import Path
55
from typing import AsyncGenerator, Optional
66

7+
from basic_memory.config import ProjectConfig
8+
from alembic import command
9+
from alembic.config import Config
10+
711
from loguru import logger
812
from sqlalchemy import text
913
from sqlalchemy.ext.asyncio import (
@@ -14,8 +18,8 @@
1418
async_scoped_session,
1519
)
1620

17-
from basic_memory.models import Base
1821
from basic_memory.models.search import CREATE_SEARCH_INDEX
22+
from basic_memory.repository.search_repository import SearchRepository
1923

2024
# Module level state
2125
_engine: Optional[AsyncEngine] = None
@@ -35,7 +39,7 @@ def get_db_url(cls, db_path: Path, db_type: "DatabaseType") -> str:
3539
logger.info("Using in-memory SQLite database")
3640
return "sqlite+aiosqlite://"
3741

38-
return f"sqlite+aiosqlite:///{db_path}"
42+
return f"sqlite+aiosqlite:///{db_path}" # pragma: no cover
3943

4044

4145
def get_scoped_session_factory(
@@ -69,21 +73,6 @@ async def scoped_session(
6973
await factory.remove()
7074

7175

72-
async def init_db() -> None:
73-
"""Initialize database with required tables."""
74-
if _session_maker is None: # pragma: no cover
75-
raise RuntimeError("Database session maker not initialized")
76-
77-
logger.info("Initializing database...")
78-
79-
async with scoped_session(_session_maker) as session:
80-
await session.execute(text("PRAGMA foreign_keys=ON"))
81-
82-
# recreate search index
83-
await session.execute(CREATE_SEARCH_INDEX)
84-
85-
await session.commit()
86-
8776

8877
async def get_or_create_db(
8978
db_path: Path,
@@ -98,9 +87,6 @@ async def get_or_create_db(
9887
_engine = create_async_engine(db_url, connect_args={"check_same_thread": False})
9988
_session_maker = async_sessionmaker(_engine, expire_on_commit=False)
10089

101-
# Initialize database
102-
await init_db()
103-
10490
assert _engine is not None # for type checker
10591
assert _session_maker is not None # for type checker
10692
return _engine, _session_maker
@@ -120,7 +106,6 @@ async def shutdown_db() -> None: # pragma: no cover
120106
async def engine_session_factory(
121107
db_path: Path,
122108
db_type: DatabaseType = DatabaseType.MEMORY,
123-
init: bool = True,
124109
) -> AsyncGenerator[tuple[AsyncEngine, async_sessionmaker[AsyncSession]], None]:
125110
"""Create engine and session factory.
126111
@@ -137,9 +122,6 @@ async def engine_session_factory(
137122
try:
138123
_session_maker = async_sessionmaker(_engine, expire_on_commit=False)
139124

140-
if init:
141-
await init_db()
142-
143125
assert _engine is not None # for type checker
144126
assert _session_maker is not None # for type checker
145127
yield _engine, _session_maker
@@ -148,3 +130,18 @@ async def engine_session_factory(
148130
await _engine.dispose()
149131
_engine = None
150132
_session_maker = None
133+
134+
135+
async def run_migrations(app_config: ProjectConfig, database_type=DatabaseType.FILESYSTEM):
136+
"""Run any pending alembic migrations."""
137+
logger.info("Running database migrations...")
138+
try:
139+
config = Config("alembic.ini")
140+
command.upgrade(config, "head")
141+
logger.info("Migrations completed successfully")
142+
143+
_, session_maker = await get_or_create_db(app_config.database_path, database_type)
144+
await SearchRepository(session_maker).init_search_index()
145+
except Exception as e: # pragma: no cover
146+
logger.error(f"Error running migrations: {e}")
147+
raise

tests/conftest.py

Lines changed: 3 additions & 5 deletions

0 commit comments

Comments
 (0)