Skip to content

Commit 0cbbe37

Browse files
committed
fix(alembic): auto-discover models for fresh databases
1 parent bf05265 commit 0cbbe37

3 files changed

Lines changed: 141 additions & 0 deletions

File tree

src/backend/bisheng/core/database/alembic/env.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -393,8 +393,10 @@ def run_migrations_online() -> None:
393393
"""
394394

395395
from bisheng.core.database.manager import sync_get_database_connection
396+
from bisheng.core.database.model_discovery import import_all_sqlmodel_models
396397

397398
database_conn_manager = sync_get_database_connection()
399+
import_all_sqlmodel_models()
398400

399401
with database_conn_manager.engine.connect() as connection:
400402
ensure_alembic_version_table(connection)
Lines changed: 63 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,63 @@
1+
"""Convention-based SQLModel discovery for schema bootstrap tooling."""
2+
3+
import ast
4+
import importlib
5+
from pathlib import Path
6+
7+
_BISHENG_PACKAGE_ROOT = Path(__file__).resolve().parents[2]
8+
_SQLMODEL_DIRECTORY_PATTERNS = (
9+
"database/models",
10+
"common/models",
11+
"*/domain/models",
12+
)
13+
14+
15+
def _declares_sqlmodel_table(module_path: Path) -> bool:
16+
"""Return whether a module declares a literal ``table=True`` class."""
17+
tree = ast.parse(module_path.read_text(encoding="utf-8"), filename=str(module_path))
18+
return any(
19+
isinstance(node, ast.ClassDef)
20+
and any(
21+
keyword.arg == "table" and isinstance(keyword.value, ast.Constant) and keyword.value.value is True
22+
for keyword in node.keywords
23+
)
24+
for node in ast.walk(tree)
25+
)
26+
27+
28+
def _module_name(module_path: Path) -> str:
29+
relative_path = module_path.relative_to(_BISHENG_PACKAGE_ROOT.parent).with_suffix("")
30+
module_parts = list(relative_path.parts)
31+
if module_parts[-1] == "__init__":
32+
module_parts.pop()
33+
return ".".join(module_parts)
34+
35+
36+
def discover_sqlmodel_module_names() -> tuple[str, ...]:
37+
"""Find table modules under the repository's model-directory conventions."""
38+
model_directories = {
39+
model_directory
40+
for pattern in _SQLMODEL_DIRECTORY_PATTERNS
41+
for model_directory in _BISHENG_PACKAGE_ROOT.glob(pattern)
42+
if model_directory.is_dir()
43+
}
44+
module_names = {
45+
_module_name(module_path)
46+
for model_directory in model_directories
47+
for module_path in model_directory.rglob("*.py")
48+
if _declares_sqlmodel_table(module_path)
49+
}
50+
return tuple(sorted(module_names))
51+
52+
53+
def import_all_sqlmodel_models() -> None:
54+
"""Strictly import every discovered model so metadata is complete."""
55+
module_names = discover_sqlmodel_module_names()
56+
if not module_names:
57+
raise RuntimeError("No SQLModel table modules found under the model-directory conventions")
58+
59+
for module_name in module_names:
60+
try:
61+
importlib.import_module(module_name)
62+
except Exception as exc:
63+
raise RuntimeError(f"Failed to import SQLModel module {module_name}") from exc
Lines changed: 76 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,76 @@
1+
"""Guards for convention-based SQLModel discovery used by DB bootstrap."""
2+
3+
import ast
4+
import subprocess
5+
import sys
6+
from pathlib import Path
7+
8+
import pytest
9+
10+
from bisheng.core.database import model_discovery
11+
from bisheng.core.database.model_discovery import discover_sqlmodel_module_names, import_all_sqlmodel_models
12+
13+
_BISHENG_ROOT = Path(__file__).resolve().parents[2] / "bisheng"
14+
15+
16+
def _declares_sqlmodel_table(path: Path) -> bool:
17+
tree = ast.parse(path.read_text(encoding="utf-8"))
18+
return any(
19+
isinstance(node, ast.ClassDef)
20+
and any(
21+
keyword.arg == "table" and isinstance(keyword.value, ast.Constant) and keyword.value.value is True
22+
for keyword in node.keywords
23+
)
24+
for node in ast.walk(tree)
25+
)
26+
27+
28+
def _module_name(path: Path) -> str:
29+
module_parts = list(path.relative_to(_BISHENG_ROOT.parent).with_suffix("").parts)
30+
if module_parts[-1] == "__init__":
31+
module_parts.pop()
32+
return ".".join(module_parts)
33+
34+
35+
def test_discovery_covers_every_literal_sqlmodel_table_module():
36+
discovered_modules = {_module_name(path) for path in _BISHENG_ROOT.rglob("*.py") if _declares_sqlmodel_table(path)}
37+
convention_modules = set(discover_sqlmodel_module_names())
38+
39+
missing_modules = discovered_modules - convention_modules
40+
41+
assert not missing_modules, (
42+
f"SQLModel table modules outside the model-directory conventions: {sorted(missing_modules)}"
43+
)
44+
45+
46+
def test_strict_model_loading_registers_space_channel_member():
47+
result = subprocess.run(
48+
[
49+
sys.executable,
50+
"-c",
51+
"from sqlmodel import SQLModel; "
52+
"from bisheng.core.database.model_discovery import import_all_sqlmodel_models; "
53+
"import_all_sqlmodel_models(); "
54+
"assert 'space_channel_member' in SQLModel.metadata.tables",
55+
],
56+
check=False,
57+
capture_output=True,
58+
text=True,
59+
)
60+
61+
assert result.returncode == 0, result.stderr
62+
63+
64+
def test_strict_model_loading_raises_when_a_discovered_module_cannot_import(monkeypatch):
65+
missing_module = "bisheng.common.models.does_not_exist"
66+
monkeypatch.setattr(model_discovery, "discover_sqlmodel_module_names", lambda: (missing_module,))
67+
68+
with pytest.raises(RuntimeError, match=missing_module):
69+
import_all_sqlmodel_models()
70+
71+
72+
def test_strict_model_loading_rejects_empty_discovery(monkeypatch):
73+
monkeypatch.setattr(model_discovery, "discover_sqlmodel_module_names", tuple)
74+
75+
with pytest.raises(RuntimeError, match="No SQLModel table modules"):
76+
import_all_sqlmodel_models()

0 commit comments

Comments
 (0)