diff --git a/.github/workflows/pr.yml b/.github/workflows/pr.yml index 5bebe1ef..5b389c45 100644 --- a/.github/workflows/pr.yml +++ b/.github/workflows/pr.yml @@ -93,6 +93,18 @@ jobs: - run: make gen-pages - run: make ci-js-typecheck + file-size-check: + name: File size (300-line cap) + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v6 + - uses: astral-sh/setup-uv@v8.0.0 + with: + enable-cache: true + cache-dependency-glob: ${{ env.UV_CACHE_GLOB }} + - run: make install-py + - run: make ci-check-file-size + # Single required status check for branch protection. # Protect `main` with this one check and every leaf job is required transitively. pr-checks: @@ -104,6 +116,7 @@ jobs: - python-tests - js-lint - js-typecheck + - file-size-check if: always() steps: - if: contains(needs.*.result, 'failure') || contains(needs.*.result, 'cancelled') diff --git a/Makefile b/Makefile index f85bf191..3fb1862d 100644 --- a/Makefile +++ b/Makefile @@ -1,4 +1,4 @@ -.PHONY: install install-py install-js dev dev-api dev-ui build test lint doctor migrate migration downgrade migration-history docker-up docker-down kill new-module gen-pages ci-python-lint ci-python-typecheck ci-js-lint ci-js-typecheck +.PHONY: install install-py install-js dev dev-api dev-ui build test lint doctor migrate migration downgrade migration-history docker-up docker-down kill new-module gen-pages ci-python-lint ci-python-typecheck ci-js-lint ci-js-typecheck ci-check-file-size # Install install: @@ -35,7 +35,7 @@ build: test: uv run pytest -lint: ci-python-lint ci-python-typecheck ci-js-lint ci-js-typecheck +lint: ci-python-lint ci-python-typecheck ci-js-lint ci-js-typecheck ci-check-file-size # Kept granular so pr.yml can run them in parallel. ci-python-lint: @@ -51,6 +51,11 @@ ci-js-lint: ci-js-typecheck: npx tsc --noEmit -p host/client_app/tsconfig.json +# Enforce a max of 300 lines per .py/.ts/.tsx file. +# Exempts vendored shadcn components under packages/ui/src/components/ui/**. +ci-check-file-size: + uv run python scripts/check_file_size.py + # Diagnostics doctor: uv run python -m simple_module_core diff --git a/framework/core/simple_module_core/diagnostics/__init__.py b/framework/core/simple_module_core/diagnostics/__init__.py new file mode 100644 index 00000000..ca9367e8 --- /dev/null +++ b/framework/core/simple_module_core/diagnostics/__init__.py @@ -0,0 +1,22 @@ +"""Module diagnostics — validates structure and patterns at startup or via CLI. + +This is the public surface re-exported from the submodules below. +Callers import from ``simple_module_core.diagnostics`` and should not +need to reach into ``._module`` or ``._migration`` directly. +""" + +from __future__ import annotations + +from simple_module_core.diagnostics._migration import MigrationDiagnostics +from simple_module_core.diagnostics._module import ModuleDiagnostics +from simple_module_core.diagnostics._runner import print_diagnostics, run_diagnostics +from simple_module_core.diagnostics._types import Diagnostic, DiagnosticLevel + +__all__ = [ + "Diagnostic", + "DiagnosticLevel", + "MigrationDiagnostics", + "ModuleDiagnostics", + "print_diagnostics", + "run_diagnostics", +] diff --git a/framework/core/simple_module_core/diagnostics/_migration.py b/framework/core/simple_module_core/diagnostics/_migration.py new file mode 100644 index 00000000..5c135b0d --- /dev/null +++ b/framework/core/simple_module_core/diagnostics/_migration.py @@ -0,0 +1,45 @@ +"""Alembic migration-state diagnostics (SM010, SM011).""" + +from __future__ import annotations + +from simple_module_core.diagnostics._types import Diagnostic, DiagnosticLevel + + +class MigrationDiagnostics: + """Validates database migration state.""" + + def check_revision_mismatch( + self, + current_revision: str | None, + head_revision: str | None, + ) -> list[Diagnostic]: + """SM010: Error if database is not at the migration head.""" + if current_revision == head_revision: + return [] + return [ + Diagnostic( + level=DiagnosticLevel.ERROR, + code="SM010", + message=(f"Database at revision {current_revision!r}, expected {head_revision!r}"), + module_name="migrations", + suggestion="Run: make migrate", + ) + ] + + def check_table_coverage( + self, + module_tables: set[str], + migrated_tables: set[str], + ) -> list[Diagnostic]: + """SM011: Warning if module tables are missing from migration history.""" + missing = module_tables - migrated_tables + return [ + Diagnostic( + level=DiagnosticLevel.WARNING, + code="SM011", + message=f"Table '{table}' declared in models but not found in migration history", + module_name="migrations", + suggestion=f'Run: make migration msg="add {table}"', + ) + for table in sorted(missing) + ] diff --git a/framework/core/simple_module_core/diagnostics.py b/framework/core/simple_module_core/diagnostics/_module.py similarity index 73% rename from framework/core/simple_module_core/diagnostics.py rename to framework/core/simple_module_core/diagnostics/_module.py index 228d130b..69735400 100644 --- a/framework/core/simple_module_core/diagnostics.py +++ b/framework/core/simple_module_core/diagnostics/_module.py @@ -1,48 +1,17 @@ -"""Module diagnostics — validates structure and patterns at startup or via CLI.""" +"""Structural diagnostics that validate discovered modules against invariants.""" from __future__ import annotations import ast import importlib.util -import logging -import sys -from dataclasses import dataclass -from enum import StrEnum from pathlib import Path from typing import TYPE_CHECKING +from simple_module_core.diagnostics._types import Diagnostic, DiagnosticLevel + if TYPE_CHECKING: from simple_module_core.module import ModuleBase -logger = logging.getLogger(__name__) - - -class DiagnosticLevel(StrEnum): - ERROR = "error" - WARNING = "warning" - INFO = "info" - - -@dataclass -class Diagnostic: - """A single diagnostic finding.""" - - level: DiagnosticLevel - code: str - message: str - module_name: str - file: str | None = None - suggestion: str | None = None - - def __str__(self) -> str: - prefix = {"error": "\u2717", "warning": "\u26a0", "info": "\u2139"}[self.level] - parts = [f"{prefix} {self.code} [{self.level.upper()}] {self.module_name}: {self.message}"] - if self.file: - parts.append(f" \u21b3 {self.file}") - if self.suggestion: - parts.append(f" \u21b3 Suggestion: {self.suggestion}") - return "\n".join(parts) - class ModuleDiagnostics: """Validates module structure and configuration.""" @@ -162,7 +131,6 @@ def _check_framework_module_coupling(self, modules: list[ModuleBase]) -> list[Di discovered module's package (e.g. ``auth``, ``products``). All interaction should go through the ``ModuleBase`` lifecycle hooks. """ - # Collect top-level package names for every discovered module. module_packages: dict[str, str] = {} # package -> module name for mod in modules: top_pkg = type(mod).__module__.split(".")[0] @@ -171,7 +139,6 @@ def _check_framework_module_coupling(self, modules: list[ModuleBase]) -> list[Di if not module_packages: return [] - # Locate framework package source directories. framework_dirs: list[tuple[str, Path]] = [] for fw_pkg in self.FRAMEWORK_PACKAGES: fw_dir = self._find_package_dir(fw_pkg) @@ -309,90 +276,3 @@ def _find_source_dir(self, mod: ModuleBase) -> Path | None: """Locate the source directory for a module's package.""" pkg_name = type(mod).__module__.rsplit(".", 1)[0] return self._find_package_dir(pkg_name) - - -class MigrationDiagnostics: - """Validates database migration state.""" - - def check_revision_mismatch( - self, - current_revision: str | None, - head_revision: str | None, - ) -> list[Diagnostic]: - """SM010: Error if database is not at the migration head.""" - if current_revision == head_revision: - return [] - return [ - Diagnostic( - level=DiagnosticLevel.ERROR, - code="SM010", - message=(f"Database at revision {current_revision!r}, expected {head_revision!r}"), - module_name="migrations", - suggestion="Run: make migrate", - ) - ] - - def check_table_coverage( - self, - module_tables: set[str], - migrated_tables: set[str], - ) -> list[Diagnostic]: - """SM011: Warning if module tables are missing from migration history.""" - missing = module_tables - migrated_tables - return [ - Diagnostic( - level=DiagnosticLevel.WARNING, - code="SM011", - message=f"Table '{table}' declared in models but not found in migration history", - module_name="migrations", - suggestion=f'Run: make migration msg="add {table}"', - ) - for table in sorted(missing) - ] - - -def run_diagnostics( - modules: list[ModuleBase], - *, - migration_state: dict | None = None, - module_tables: set[str] | None = None, - migrated_tables: set[str] | None = None, -) -> list[Diagnostic]: - """Convenience function to run all diagnostics. - - When ``migration_state`` is provided, also runs migration diagnostics. - """ - diagnostics = ModuleDiagnostics().run(modules) - - if migration_state is not None: - migration_diag = MigrationDiagnostics() - diagnostics.extend( - migration_diag.check_revision_mismatch( - current_revision=migration_state.get("current_revision"), - head_revision=migration_state.get("head_revision"), - ) - ) - if module_tables is not None and migrated_tables is not None: - diagnostics.extend(migration_diag.check_table_coverage(module_tables, migrated_tables)) - - return diagnostics - - -def print_diagnostics(diagnostics: list[Diagnostic]) -> None: - """Pretty-print diagnostics to stderr.""" - if not diagnostics: - logger.info("\u2713 No module diagnostics issues found") - return - - errors = [d for d in diagnostics if d.level == DiagnosticLevel.ERROR] - warnings = [d for d in diagnostics if d.level == DiagnosticLevel.WARNING] - infos = [d for d in diagnostics if d.level == DiagnosticLevel.INFO] - - for d in diagnostics: - print(str(d), file=sys.stderr) - print(file=sys.stderr) - - print( - f"Results: {len(errors)} error(s), {len(warnings)} warning(s), {len(infos)} info", - file=sys.stderr, - ) diff --git a/framework/core/simple_module_core/diagnostics/_runner.py b/framework/core/simple_module_core/diagnostics/_runner.py new file mode 100644 index 00000000..5f659a84 --- /dev/null +++ b/framework/core/simple_module_core/diagnostics/_runner.py @@ -0,0 +1,63 @@ +"""Entry points that assemble and print diagnostic output.""" + +from __future__ import annotations + +import logging +import sys +from typing import TYPE_CHECKING + +from simple_module_core.diagnostics._migration import MigrationDiagnostics +from simple_module_core.diagnostics._module import ModuleDiagnostics +from simple_module_core.diagnostics._types import Diagnostic, DiagnosticLevel + +if TYPE_CHECKING: + from simple_module_core.module import ModuleBase + +logger = logging.getLogger(__name__) + + +def run_diagnostics( + modules: list[ModuleBase], + *, + migration_state: dict | None = None, + module_tables: set[str] | None = None, + migrated_tables: set[str] | None = None, +) -> list[Diagnostic]: + """Convenience function to run all diagnostics. + + When ``migration_state`` is provided, also runs migration diagnostics. + """ + diagnostics = ModuleDiagnostics().run(modules) + + if migration_state is not None: + migration_diag = MigrationDiagnostics() + diagnostics.extend( + migration_diag.check_revision_mismatch( + current_revision=migration_state.get("current_revision"), + head_revision=migration_state.get("head_revision"), + ) + ) + if module_tables is not None and migrated_tables is not None: + diagnostics.extend(migration_diag.check_table_coverage(module_tables, migrated_tables)) + + return diagnostics + + +def print_diagnostics(diagnostics: list[Diagnostic]) -> None: + """Pretty-print diagnostics to stderr.""" + if not diagnostics: + logger.info("\u2713 No module diagnostics issues found") + return + + errors = [d for d in diagnostics if d.level == DiagnosticLevel.ERROR] + warnings = [d for d in diagnostics if d.level == DiagnosticLevel.WARNING] + infos = [d for d in diagnostics if d.level == DiagnosticLevel.INFO] + + for d in diagnostics: + print(str(d), file=sys.stderr) + print(file=sys.stderr) + + print( + f"Results: {len(errors)} error(s), {len(warnings)} warning(s), {len(infos)} info", + file=sys.stderr, + ) diff --git a/framework/core/simple_module_core/diagnostics/_types.py b/framework/core/simple_module_core/diagnostics/_types.py new file mode 100644 index 00000000..8745d794 --- /dev/null +++ b/framework/core/simple_module_core/diagnostics/_types.py @@ -0,0 +1,33 @@ +"""Core diagnostic types: level enum and finding dataclass.""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import StrEnum + + +class DiagnosticLevel(StrEnum): + ERROR = "error" + WARNING = "warning" + INFO = "info" + + +@dataclass +class Diagnostic: + """A single diagnostic finding.""" + + level: DiagnosticLevel + code: str + message: str + module_name: str + file: str | None = None + suggestion: str | None = None + + def __str__(self) -> str: + prefix = {"error": "\u2717", "warning": "\u26a0", "info": "\u2139"}[self.level] + parts = [f"{prefix} {self.code} [{self.level.upper()}] {self.module_name}: {self.message}"] + if self.file: + parts.append(f" \u21b3 {self.file}") + if self.suggestion: + parts.append(f" \u21b3 Suggestion: {self.suggestion}") + return "\n".join(parts) diff --git a/framework/core/tests/test_core.py b/framework/core/tests/test_core.py deleted file mode 100644 index 10194407..00000000 --- a/framework/core/tests/test_core.py +++ /dev/null @@ -1,1061 +0,0 @@ -"""Tests for the framework core: module system, menu, permissions, feature flags, events.""" - -from __future__ import annotations - -from dataclasses import dataclass - -import pytest -from simple_module_core.diagnostics import ( - Diagnostic, - DiagnosticLevel, - MigrationDiagnostics, - print_diagnostics, -) -from simple_module_core.discovery import discover_modules, topological_sort -from simple_module_core.events import Event, EventBus -from simple_module_core.exceptions import ( - CircularDependencyError, - FrameworkVersionError, - InvalidModuleError, -) -from simple_module_core.feature_flags import FeatureFlagDefinition, FeatureFlagRegistry -from simple_module_core.health import HealthCheck, HealthCheckResult, HealthRegistry, HealthStatus -from simple_module_core.menu import MenuItem, MenuRegistry, MenuSection -from simple_module_core.module import ModuleBase, ModuleMeta -from simple_module_core.permissions import PermissionRegistry - -# ── ModuleMeta ─────────────────────────────────────────────────────── - - -class TestModuleMeta: - async def test_defaults(self): - meta = ModuleMeta(name="TestModule") - assert meta.name == "TestModule" - assert meta.route_prefix == "" - assert meta.view_prefix == "" - assert meta.depends_on == [] - assert meta.version == "1.0.0" - - async def test_custom_fields(self): - meta = ModuleMeta( - name="Products", - route_prefix="/api/products", - view_prefix="/products", - depends_on=["Auth"], - version="2.0.0", - ) - assert meta.route_prefix == "/api/products" - assert meta.depends_on == ["Auth"] - assert meta.version == "2.0.0" - - async def test_frozen(self): - meta = ModuleMeta(name="Frozen") - with pytest.raises(AttributeError): - meta.name = "Changed" # type: ignore[misc] # ty: ignore[invalid-assignment] - - -# ── ModuleBase ─────────────────────────────────────────────────────── - - -class DummyModule(ModuleBase): - meta = ModuleMeta(name="Dummy", route_prefix="/api/dummy") - - def __init__(self): - self.routes_registered = False - - def register_routes(self, api_router, view_router): - self.routes_registered = True - - -class TestModuleBase: - async def test_subclass_has_meta(self): - mod = DummyModule() - assert mod.meta.name == "Dummy" - - async def test_register_routes_override(self): - mod = DummyModule() - mod.register_routes(None, None) # type: ignore[arg-type] - assert mod.routes_registered is True - - async def test_default_noop_methods(self): - """Default implementations should not raise.""" - mod = DummyModule() - mod.register_menu_items(MenuRegistry()) - mod.register_permissions(PermissionRegistry()) - - -# ── MenuRegistry ───────────────────────────────────────────────────── - - -class TestMenuRegistry: - async def test_add_and_all_items(self): - reg = MenuRegistry() - reg.add(MenuItem(label="Dashboard", url="/dashboard", order=1)) - reg.add(MenuItem(label="Products", url="/products", order=2)) - assert len(reg.all_items) == 2 - assert reg.all_items[0].label == "Dashboard" - - async def test_add_many(self): - reg = MenuRegistry() - reg.add_many( - [ - MenuItem(label="A", url="/a", order=1), - MenuItem(label="B", url="/b", order=2), - ] - ) - assert len(reg.all_items) == 2 - - async def test_sorted_by_order(self): - reg = MenuRegistry() - reg.add(MenuItem(label="Z", url="/z", order=99)) - reg.add(MenuItem(label="A", url="/a", order=1)) - assert reg.all_items[0].label == "A" - assert reg.all_items[1].label == "Z" - - async def test_filter_unauthenticated(self): - reg = MenuRegistry() - reg.add(MenuItem(label="Public", url="/pub", requires_auth=False)) - reg.add(MenuItem(label="Private", url="/priv", requires_auth=True)) - - result = reg.get_for_user(is_authenticated=False) - sidebar = result["sidebar"] - assert len(sidebar) == 1 - assert sidebar[0]["label"] == "Public" - - async def test_filter_authenticated_sees_all(self): - reg = MenuRegistry() - reg.add(MenuItem(label="Public", url="/pub", requires_auth=False)) - reg.add(MenuItem(label="Private", url="/priv", requires_auth=True)) - - result = reg.get_for_user(is_authenticated=True) - sidebar = result["sidebar"] - assert len(sidebar) == 2 - - async def test_filter_by_roles(self): - reg = MenuRegistry() - reg.add(MenuItem(label="Admin Panel", url="/admin", roles=["admin"])) - reg.add(MenuItem(label="Dashboard", url="/dash")) - - # User without admin role - result = reg.get_for_user(is_authenticated=True, roles=["user"]) - sidebar = result["sidebar"] - labels = [i["label"] for i in sidebar] - assert "Dashboard" in labels - assert "Admin Panel" not in labels - - # User with admin role - result = reg.get_for_user(is_authenticated=True, roles=["admin"]) - sidebar = result["sidebar"] - labels = [i["label"] for i in sidebar] - assert "Admin Panel" in labels - - async def test_sections(self): - reg = MenuRegistry() - reg.add(MenuItem(label="Side", url="/s", section=MenuSection.SIDEBAR)) - reg.add(MenuItem(label="Nav", url="/n", section=MenuSection.NAVBAR)) - reg.add(MenuItem(label="Drop", url="/d", section=MenuSection.USER_DROPDOWN)) - - result = reg.get_for_user(is_authenticated=True) - assert len(result["sidebar"]) == 1 - assert len(result["navbar"]) == 1 - assert len(result["userDropdown"]) == 1 - - -# ── PermissionRegistry ─────────────────────────────────────────────── - - -class TestPermissionRegistry: - async def test_add_group(self): - reg = PermissionRegistry() - reg.add_group("Products", ["products.view", "products.create"]) - assert "products.view" in reg.all_permissions - assert "products.create" in reg.all_permissions - - async def test_add_single(self): - reg = PermissionRegistry() - reg.add("orders.view") - assert reg.has("orders.view") - - async def test_auto_grouping(self): - reg = PermissionRegistry() - reg.add("orders.view") - reg.add("orders.create") - groups = reg.groups - assert any(g.name == "orders" for g in groups) - - async def test_has(self): - reg = PermissionRegistry() - reg.add("test.perm") - assert reg.has("test.perm") is True - assert reg.has("nonexistent") is False - - async def test_admin_role_gets_all(self): - reg = PermissionRegistry() - reg.add_group("Products", ["products.view", "products.edit"]) - perms = reg.get_permissions_for_roles(["admin"]) - assert "products.view" in perms - assert "products.edit" in perms - - async def test_non_admin_gets_none_by_default(self): - reg = PermissionRegistry() - reg.add_group("Products", ["products.view"]) - perms = reg.get_permissions_for_roles(["user"]) - assert len(perms) == 0 - - async def test_custom_role_map(self): - reg = PermissionRegistry() - reg.add_group("Products", ["products.view", "products.edit"]) - role_map = {"editor": ["products.edit"]} - perms = reg.get_permissions_for_roles(["editor"], role_permission_map=role_map) - assert "products.edit" in perms - assert "products.view" not in perms - - async def test_extend_existing_group(self): - reg = PermissionRegistry() - reg.add_group("Products", ["products.view"]) - reg.add_group("Products", ["products.delete"]) - perms = reg.all_permissions - assert "products.view" in perms - assert "products.delete" in perms - - -# ── FeatureFlagRegistry ────────────────────────────────────────────── - - -class TestFeatureFlagRegistry: - async def test_add_and_check_default(self): - reg = FeatureFlagRegistry() - reg.add(FeatureFlagDefinition(name="beta_ui", default_enabled=False)) - assert reg.is_enabled("beta_ui") is False - - async def test_default_enabled(self): - reg = FeatureFlagRegistry() - reg.add(FeatureFlagDefinition(name="stable_feature", default_enabled=True)) - assert reg.is_enabled("stable_feature") is True - - async def test_override(self): - reg = FeatureFlagRegistry() - reg.add(FeatureFlagDefinition(name="beta_ui", default_enabled=False)) - reg.set_override("beta_ui", True) - assert reg.is_enabled("beta_ui") is True - - async def test_clear_override(self): - reg = FeatureFlagRegistry() - reg.add(FeatureFlagDefinition(name="beta_ui", default_enabled=False)) - reg.set_override("beta_ui", True) - reg.clear_override("beta_ui") - assert reg.is_enabled("beta_ui") is False - - async def test_unknown_flag_is_disabled(self): - reg = FeatureFlagRegistry() - assert reg.is_enabled("nonexistent") is False - - async def test_all_flags(self): - reg = FeatureFlagRegistry() - reg.add(FeatureFlagDefinition(name="a")) - reg.add(FeatureFlagDefinition(name="b")) - assert len(reg.all_flags) == 2 - - -# ── EventBus ───────────────────────────────────────────────────────── - - -@dataclass -class OrderCreated(Event): - order_id: int = 0 - - -class TestEventBus: - async def test_subscribe_and_publish(self): - bus = EventBus() - received: list[Event] = [] - - async def handler(event: OrderCreated): - received.append(event) - - bus.subscribe(OrderCreated, handler) - await bus.publish(OrderCreated(order_id=42)) - - assert len(received) == 1 - assert received[0].order_id == 42 # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] - - async def test_multiple_handlers(self): - bus = EventBus() - calls: list[str] = [] - - async def handler_a(event: OrderCreated): - calls.append("a") - - async def handler_b(event: OrderCreated): - calls.append("b") - - bus.subscribe(OrderCreated, handler_a) - bus.subscribe(OrderCreated, handler_b) - await bus.publish(OrderCreated()) - - assert "a" in calls - assert "b" in calls - - async def test_no_handlers_no_error(self): - bus = EventBus() - await bus.publish(OrderCreated()) # Should not raise - - async def test_handler_error_does_not_propagate(self): - bus = EventBus() - calls: list[str] = [] - - async def bad_handler(event: OrderCreated): - raise ValueError("boom") - - async def good_handler(event: OrderCreated): - calls.append("ok") - - bus.subscribe(OrderCreated, bad_handler) - bus.subscribe(OrderCreated, good_handler) - await bus.publish(OrderCreated()) - - # The good handler should still have been called - assert "ok" in calls - - -# ── Discovery / topological_sort ───────────────────────────────────── - - -class ModA(ModuleBase): - meta = ModuleMeta(name="A") - - -class ModB(ModuleBase): - meta = ModuleMeta(name="B", depends_on=["A"]) - - -class ModC(ModuleBase): - meta = ModuleMeta(name="C", depends_on=["B"]) - - -class CycleX(ModuleBase): - meta = ModuleMeta(name="X", depends_on=["Y"]) - - -class CycleY(ModuleBase): - meta = ModuleMeta(name="Y", depends_on=["X"]) - - -class TestTopologicalSort: - async def test_valid_dag(self): - modules = [ModC(), ModA(), ModB()] - sorted_mods = topological_sort(modules) - names = [m.meta.name for m in sorted_mods] - assert names.index("A") < names.index("B") - assert names.index("B") < names.index("C") - - async def test_no_dependencies(self): - modules = [ModA()] - sorted_mods = topological_sort(modules) - assert len(sorted_mods) == 1 - assert sorted_mods[0].meta.name == "A" - - async def test_circular_dependency_raises(self): - modules = [CycleX(), CycleY()] - with pytest.raises(CircularDependencyError): - topological_sort(modules) - - -class TestDiscoverModules: - async def test_discover_finds_installed_modules(self): - """discover_modules() should find modules registered via entry_points.""" - from simple_module_core.discovery import discover_modules - - modules = discover_modules() - names = [m.meta.name for m in modules] - # The workspace has Auth, Products, Dashboard registered - assert "Products" in names - assert "Auth" in names - assert "Dashboard" in names - - -# ── Topological Sort Edge Cases ───────────────────────────────────── - - -class TestTopologicalSortEdgeCases: - async def test_diamond_dependency(self): - """A -> B, A -> C, B -> D, C -> D (diamond, not cycle).""" - - class ModD(ModuleBase): - meta = ModuleMeta(name="D") - - class ModB2(ModuleBase): - meta = ModuleMeta(name="B2", depends_on=["D"]) - - class ModC2(ModuleBase): - meta = ModuleMeta(name="C2", depends_on=["D"]) - - class ModA2(ModuleBase): - meta = ModuleMeta(name="A2", depends_on=["B2", "C2"]) - - modules = [ModA2(), ModC2(), ModB2(), ModD()] - sorted_mods = topological_sort(modules) - names = [m.meta.name for m in sorted_mods] - assert names.index("D") < names.index("B2") - assert names.index("D") < names.index("C2") - assert names.index("B2") < names.index("A2") - - async def test_missing_dependency_ignored(self): - """A module depending on a non-installed module should not crash.""" - - class ModWithMissing(ModuleBase): - meta = ModuleMeta(name="Lonely", depends_on=["NonExistent"]) - - sorted_mods = topological_sort([ModWithMissing()]) - assert len(sorted_mods) == 1 - - async def test_self_dependency_raises(self): - """A module that depends on itself is a cycle.""" - - class SelfDep(ModuleBase): - meta = ModuleMeta(name="Self", depends_on=["Self"]) - - with pytest.raises(CircularDependencyError): - topological_sort([SelfDep()]) - - async def test_three_node_cycle(self): - """A -> B -> C -> A should raise.""" - - class CA(ModuleBase): - meta = ModuleMeta(name="CA", depends_on=["CC"]) - - class CB(ModuleBase): - meta = ModuleMeta(name="CB", depends_on=["CA"]) - - class CC(ModuleBase): - meta = ModuleMeta(name="CC", depends_on=["CB"]) - - with pytest.raises(CircularDependencyError): - topological_sort([CA(), CB(), CC()]) - - async def test_empty_list(self): - assert topological_sort([]) == [] - - -# ── EventBus Advanced ─────────────────────────────────────────────── - - -class TestEventBusAdvanced: - async def test_different_event_types_isolated(self): - """Handlers only receive events of their subscribed type.""" - bus = EventBus() - - @dataclass - class EventA(Event): - pass - - @dataclass - class EventB(Event): - pass - - a_calls: list = [] - b_calls: list = [] - - async def handle_a(e): - a_calls.append(e) - - async def handle_b(e): - b_calls.append(e) - - bus.subscribe(EventA, handle_a) - bus.subscribe(EventB, handle_b) - - await bus.publish(EventA()) - assert len(a_calls) == 1 - assert len(b_calls) == 0 - - await bus.publish(EventB()) - assert len(b_calls) == 1 - - async def test_publish_nowait(self): - """publish_nowait should schedule without blocking.""" - import asyncio - - bus = EventBus() - received: list = [] - - async def handler(e): - received.append(e) - - bus.subscribe(OrderCreated, handler) - bus.publish_nowait(OrderCreated(order_id=99)) - await asyncio.sleep(0.05) - assert len(received) == 1 - assert received[0].order_id == 99 - - async def test_subclass_events_do_not_match_parent_subscription(self): - """Subscribing to a base Event class should not receive subclass events.""" - bus = EventBus() - - @dataclass - class Parent(Event): - pass - - @dataclass - class Child(Parent): - pass - - calls: list = [] - - async def parent_handler(e): - calls.append(("parent", e)) - - bus.subscribe(Parent, parent_handler) - await bus.publish(Child()) - - # Child events should not trigger Parent handlers — strict type match. - assert calls == [] - - async def test_publish_with_no_subscribers_returns_none(self): - """publish() should resolve to None when nothing is listening.""" - bus = EventBus() - - @dataclass - class Orphan(Event): - pass - - result = await bus.publish(Orphan()) - assert result is None - - async def test_publish_nowait_with_no_subscribers_is_noop(self): - """publish_nowait() on an unheard event should not raise.""" - bus = EventBus() - - @dataclass - class Orphan(Event): - pass - - bus.publish_nowait(Orphan()) # must not raise - - async def test_handlers_dispatched_concurrently(self): - """All handlers for an event should run concurrently via gather.""" - import asyncio - - bus = EventBus() - order: list[str] = [] - - async def slow(e): - await asyncio.sleep(0.02) - order.append("slow") - - async def fast(e): - order.append("fast") - - bus.subscribe(OrderCreated, slow) - bus.subscribe(OrderCreated, fast) - await bus.publish(OrderCreated(order_id=1)) - - # "fast" should complete before "slow" because they run concurrently. - assert order == ["fast", "slow"] - - -# ── MenuRegistry Advanced ─────────────────────────────────────────── - - -class TestMenuRegistryAdvanced: - async def test_multiple_roles_any_match(self): - reg = MenuRegistry() - reg.add(MenuItem(label="Editor", url="/edit", roles=["editor", "admin"])) - result = reg.get_for_user(is_authenticated=True, roles=["editor"]) - assert len(result["sidebar"]) == 1 - - async def test_empty_registry(self): - reg = MenuRegistry() - result = reg.get_for_user(is_authenticated=True) - assert all(len(v) == 0 for v in result.values()) - - async def test_admin_sidebar_section(self): - reg = MenuRegistry() - reg.add(MenuItem(label="Users", url="/admin/users", section=MenuSection.ADMIN_SIDEBAR)) - result = reg.get_for_user(is_authenticated=True) - assert len(result["adminSidebar"]) == 1 - - async def test_icon_preserved(self): - reg = MenuRegistry() - reg.add(MenuItem(label="Home", url="/", icon="home")) - result = reg.get_for_user(is_authenticated=True) - assert result["sidebar"][0]["icon"] == "home" - - -# ── PermissionRegistry Advanced ───────────────────────────────────── - - -class TestPermissionRegistryAdvanced: - async def test_no_duplicates(self): - reg = PermissionRegistry() - reg.add("products.view") - reg.add("products.view") - assert reg.all_permissions.count("products.view") == 1 - - async def test_multiple_roles_union(self): - reg = PermissionRegistry() - reg.add_group("Products", ["products.view", "products.edit"]) - role_map = {"viewer": ["products.view"], "editor": ["products.edit"]} - perms = reg.get_permissions_for_roles(["viewer", "editor"], role_permission_map=role_map) - assert "products.view" in perms - assert "products.edit" in perms - - async def test_groups_list(self): - reg = PermissionRegistry() - reg.add_group("Auth", ["auth.login"]) - reg.add_group("Products", ["products.view"]) - assert len(reg.groups) == 2 - - async def test_permissions_sorted(self): - reg = PermissionRegistry() - reg.add("z.last") - reg.add("a.first") - assert reg.all_permissions == ["a.first", "z.last"] - - -# ── DiscoverModules Advanced ──────────────────────────────────────── - - -class TestDiscoverModulesAdvanced: - async def test_discover_returns_module_instances(self): - from simple_module_core.discovery import discover_modules - - modules = discover_modules() - for mod in modules: - assert isinstance(mod, ModuleBase) - assert hasattr(mod, "meta") - - async def test_discover_modules_have_valid_meta(self): - from simple_module_core.discovery import discover_modules - - modules = discover_modules() - for mod in modules: - assert isinstance(mod.meta, ModuleMeta) - assert isinstance(mod.meta.depends_on, list) - assert mod.meta.name != "" - - -# ── discover_modules validation & strict mode ────────────────────── - - -class _FakeEntryPoint: - """Minimal EntryPoint shim for testing the validation path. - - Pass a class to return on ``load()``, or a zero-arg callable to - raise/return something custom (for load-failure cases). - """ - - def __init__(self, name: str, target): - self.name = name - self._target = target - - def load(self): - return ( - self._target() - if callable(self._target) and not isinstance(self._target, type) - else self._target - ) - - -def _patch_entry_points(monkeypatch, eps): - import simple_module_core.discovery as discovery_mod - - monkeypatch.setattr(discovery_mod, "entry_points", lambda group: eps) - - -def _boom_loader(): - raise ImportError("boom") - - -class TestDiscoverModulesValidation: - async def test_missing_meta_strict_raises(self, monkeypatch): - class NoMeta(ModuleBase): # intentionally no meta - pass - - _patch_entry_points(monkeypatch, [_FakeEntryPoint("nometa", NoMeta)]) - - with pytest.raises(InvalidModuleError, match="missing 'meta"): - discover_modules(strict=True) - - async def test_missing_meta_non_strict_skips(self, monkeypatch): - class NoMeta(ModuleBase): - pass - - _patch_entry_points(monkeypatch, [_FakeEntryPoint("nometa", NoMeta)]) - - assert discover_modules(strict=False) == [] - - async def test_non_modulebase_strict_raises(self, monkeypatch): - class NotAModule: - pass - - _patch_entry_points(monkeypatch, [_FakeEntryPoint("notmod", NotAModule)]) - - with pytest.raises(InvalidModuleError, match="not a ModuleBase"): - discover_modules(strict=True) - - async def test_load_failure_strict_raises(self, monkeypatch): - _patch_entry_points(monkeypatch, [_FakeEntryPoint("broken", _boom_loader)]) - - with pytest.raises(InvalidModuleError, match="Failed to load"): - discover_modules(strict=True) - - async def test_load_failure_non_strict_logs_and_skips(self, monkeypatch, caplog): - import logging - - _patch_entry_points(monkeypatch, [_FakeEntryPoint("broken", _boom_loader)]) - - with caplog.at_level(logging.ERROR, logger="simple_module_core.discovery"): - modules = discover_modules(strict=False) - - assert modules == [] - assert any("Failed to load" in r.message for r in caplog.records) - - async def test_meta_must_be_modulemeta_instance(self, monkeypatch): - class BadMeta(ModuleBase): - meta = "not a ModuleMeta" # type: ignore[assignment] - - _patch_entry_points(monkeypatch, [_FakeEntryPoint("bad", BadMeta)]) - - with pytest.raises(InvalidModuleError, match="missing 'meta"): - discover_modules(strict=True) - - -# ── ModuleBase Lifecycle ──────────────────────────────────────────── - - -class TestModuleLifecycle: - async def test_on_startup_default_noop(self): - mod = DummyModule() - await mod.on_startup(None) # type: ignore - - async def test_on_shutdown_default_noop(self): - mod = DummyModule() - await mod.on_shutdown(None) # type: ignore - - async def test_register_event_handlers_default_noop(self): - mod = DummyModule() - bus = EventBus() - mod.register_event_handlers(bus) - - async def test_register_feature_flags_default_noop(self): - mod = DummyModule() - reg = FeatureFlagRegistry() - mod.register_feature_flags(reg) - assert len(reg.all_flags) == 0 - - -# ── HealthRegistry ───────────────────────────────────────────────── - - -class TestHealthRegistry: - async def test_add_and_list(self): - reg = HealthRegistry() - - async def check_db() -> HealthCheckResult: - return HealthCheckResult(status=HealthStatus.HEALTHY) - - reg.add(HealthCheck(name="db", check=check_db)) - assert len(reg.all_checks) == 1 - assert reg.all_checks[0].name == "db" - - async def test_empty_registry(self): - reg = HealthRegistry() - assert reg.all_checks == [] - - async def test_multiple_checks(self): - reg = HealthRegistry() - - async def check_a() -> HealthCheckResult: - return HealthCheckResult(status=HealthStatus.HEALTHY) - - async def check_b() -> HealthCheckResult: - return HealthCheckResult(status=HealthStatus.DEGRADED, detail="slow") - - reg.add(HealthCheck(name="a", check=check_a)) - reg.add(HealthCheck(name="b", check=check_b)) - assert len(reg.all_checks) == 2 - - async def test_check_result_defaults(self): - result = HealthCheckResult(status=HealthStatus.HEALTHY) - assert result.detail is None - - async def test_check_result_with_detail(self): - result = HealthCheckResult(status=HealthStatus.DEGRADED, detail="reindexing") - assert result.detail == "reindexing" - - async def test_health_status_ordering(self): - """Verify enum values exist for aggregation logic.""" - assert HealthStatus.HEALTHY == "healthy" - assert HealthStatus.DEGRADED == "degraded" - assert HealthStatus.UNHEALTHY == "unhealthy" - - -class TestModuleNewHooks: - async def test_register_exception_handlers_default_noop(self): - mod = DummyModule() - mod.register_exception_handlers(None) # type: ignore - - async def test_register_health_checks_default_noop(self): - mod = DummyModule() - reg = HealthRegistry() - mod.register_health_checks(reg) - assert len(reg.all_checks) == 0 - - async def test_register_settings_default_noop(self): - mod = DummyModule() - mod.register_settings(None) # type: ignore - - -# ── MigrationDiagnostics ────────────────────────────────────────── - - -class TestMigrationDiagnostics: - async def test_sm010_migration_mismatch(self): - """SM010 should fire when current revision != head.""" - diag = MigrationDiagnostics() - results = diag.check_revision_mismatch( - current_revision="abc123", - head_revision="def456", - ) - assert len(results) == 1 - assert results[0].code == "SM010" - assert results[0].level == DiagnosticLevel.ERROR - - async def test_sm010_no_error_when_current(self): - """SM010 should not fire when DB is at head.""" - diag = MigrationDiagnostics() - results = diag.check_revision_mismatch( - current_revision="abc123", - head_revision="abc123", - ) - assert len(results) == 0 - - async def test_sm011_missing_tables(self): - """SM011 should fire when module tables aren't in migration tables.""" - diag = MigrationDiagnostics() - results = diag.check_table_coverage( - module_tables={"products_product", "products_category"}, - migrated_tables={"products_product"}, - ) - assert len(results) == 1 - assert results[0].code == "SM011" - assert results[0].level == DiagnosticLevel.WARNING - assert "products_category" in results[0].message - - async def test_sm011_no_warning_when_covered(self): - """SM011 should not fire when all tables are covered.""" - diag = MigrationDiagnostics() - results = diag.check_table_coverage( - module_tables={"products_product"}, - migrated_tables={"products_product"}, - ) - assert len(results) == 0 - - -# ── Framework API Version Compatibility (Gap 3) ───────────────────── - - -class TestFrameworkVersion: - async def test_framework_exposes_api_version(self): - """`simple_module_core.FRAMEWORK_API_VERSION` must be importable and semver-shaped.""" - # Must be a non-empty string that parses as PEP 440 version. - from packaging.version import Version - from simple_module_core import FRAMEWORK_API_VERSION - - assert isinstance(FRAMEWORK_API_VERSION, str) - assert FRAMEWORK_API_VERSION != "" - Version(FRAMEWORK_API_VERSION) # raises if malformed - - async def test_module_meta_accepts_requires_framework(self): - """ModuleMeta should accept an optional requires_framework field.""" - meta = ModuleMeta(name="X", requires_framework=">=1.0,<2.0") - assert meta.requires_framework == ">=1.0,<2.0" - - async def test_module_meta_requires_framework_defaults_to_none(self): - """When not set, requires_framework is None (no compat check applied).""" - meta = ModuleMeta(name="X") - assert meta.requires_framework is None - - async def test_check_compat_passes_when_version_matches(self): - """A module declaring a spec that matches the framework version passes.""" - from simple_module_core import FRAMEWORK_API_VERSION - from simple_module_core.versioning import check_framework_compatibility - - class ModGood(ModuleBase): - meta = ModuleMeta( - name="Good", - requires_framework=f"=={FRAMEWORK_API_VERSION}", - ) - - # Must not raise. - check_framework_compatibility([ModGood()]) - - async def test_check_compat_raises_on_mismatch(self): - """A module with an unsatisfiable spec raises FrameworkVersionError at boot.""" - from simple_module_core.versioning import check_framework_compatibility - - class ModStale(ModuleBase): - # Requires a version far in the future — guaranteed to fail. - meta = ModuleMeta(name="Stale", requires_framework=">=999.0") - - with pytest.raises(FrameworkVersionError) as exc_info: - check_framework_compatibility([ModStale()]) - - # Error should name the offending module and the incompatible spec. - msg = str(exc_info.value) - assert "Stale" in msg - assert ">=999.0" in msg - - async def test_check_compat_skips_modules_without_spec(self): - """Modules that don't declare requires_framework are not checked.""" - from simple_module_core.versioning import check_framework_compatibility - - class ModLegacy(ModuleBase): - meta = ModuleMeta(name="Legacy") # no requires_framework - - check_framework_compatibility([ModLegacy()]) # must not raise - - async def test_check_compat_rejects_malformed_spec(self): - """A malformed version specifier raises FrameworkVersionError (not something cryptic).""" - from simple_module_core.versioning import check_framework_compatibility - - class ModBadSpec(ModuleBase): - meta = ModuleMeta(name="BadSpec", requires_framework="not-a-spec") - - with pytest.raises(FrameworkVersionError) as exc_info: - check_framework_compatibility([ModBadSpec()]) - assert "BadSpec" in str(exc_info.value) - - async def test_check_compat_reports_all_failures(self): - """When multiple modules are incompatible, the error mentions all of them.""" - from simple_module_core.versioning import check_framework_compatibility - - class ModBadA(ModuleBase): - meta = ModuleMeta(name="BadA", requires_framework=">=999.0") - - class ModBadB(ModuleBase): - meta = ModuleMeta(name="BadB", requires_framework=">=999.0") - - with pytest.raises(FrameworkVersionError) as exc_info: - check_framework_compatibility([ModBadA(), ModBadB()]) - - msg = str(exc_info.value) - assert "BadA" in msg - assert "BadB" in msg - - -# ── Selective Module Loading (Gap 4) ──────────────────────────────── - - -class TestSelectiveModuleLoading: - async def test_discover_with_none_loads_all(self): - """Passing enabled=None keeps existing behaviour (load all installed modules).""" - from simple_module_core.discovery import discover_modules - - all_mods = discover_modules(enabled=None) - names = {m.meta.name for m in all_mods} - assert {"Auth", "Products", "Dashboard"}.issubset(names) - - async def test_discover_with_allowlist_filters(self): - """Passing enabled=['Auth'] loads only Auth, even if other modules are installed.""" - from simple_module_core.discovery import discover_modules - - filtered = discover_modules(enabled=["Auth"]) - names = [m.meta.name for m in filtered] - assert names == ["Auth"] - - async def test_discover_with_empty_list_loads_none(self): - """Passing enabled=[] loads no modules (explicit opt-out of everything).""" - from simple_module_core.discovery import discover_modules - - assert discover_modules(enabled=[]) == [] - - async def test_discover_allowlist_case_insensitive(self): - """Allowlist matching ignores case so 'products' and 'Products' both work.""" - from simple_module_core.discovery import discover_modules - - names = [m.meta.name for m in discover_modules(enabled=["products"])] - assert names == ["Products"] - - async def test_discover_unknown_name_logged_and_ignored(self, caplog): - """Names in enabled that don't match any installed module log a warning but don't raise.""" - import logging - - from simple_module_core.discovery import discover_modules - - with caplog.at_level(logging.WARNING, logger="simple_module_core.discovery"): - result = discover_modules(enabled=["Auth", "Nonexistent"]) - - names = [m.meta.name for m in result] - assert names == ["Auth"] - assert any("nonexistent" in rec.message.lower() for rec in caplog.records) - - -# ── Module Template & Static Contribution Hooks (Gap 5) ───────────── - - -class TestModuleAssetHooks: - async def test_template_dirs_default_empty(self): - """ModuleBase.template_dirs() returns an empty list by default.""" - mod = DummyModule() - assert mod.template_dirs() == [] - - async def test_static_mounts_default_empty(self): - """ModuleBase.static_mounts() returns an empty dict by default.""" - mod = DummyModule() - assert mod.static_mounts() == {} - - async def test_template_dirs_override(self, tmp_path): - """A module can return its own template directory.""" - tpl_dir = tmp_path / "my_templates" - tpl_dir.mkdir() - - class ModWithTpl(ModuleBase): - meta = ModuleMeta(name="WithTpl") - - def template_dirs(self): - return [tpl_dir] - - mod = ModWithTpl() - result = mod.template_dirs() - assert result == [tpl_dir] - - async def test_static_mounts_override(self, tmp_path): - """A module can map URL prefixes to filesystem directories.""" - assets = tmp_path / "assets" - assets.mkdir() - - class ModWithStatic(ModuleBase): - meta = ModuleMeta(name="WithStatic") - - def static_mounts(self): - return {"/modules/with-static": assets} - - mod = ModWithStatic() - mounts = mod.static_mounts() - assert mounts == {"/modules/with-static": assets} - - -# ── print_diagnostics ───────────────────────────────────────────── - - -class TestPrintDiagnostics: - async def test_writes_to_stderr(self, capsys): - diag = Diagnostic( - level=DiagnosticLevel.ERROR, - code="SM001", - message="test error", - module_name="TestMod", - ) - print_diagnostics([diag]) - captured = capsys.readouterr() - assert captured.out == "" - assert "SM001" in captured.err - assert "Results: 1 error(s)" in captured.err - - async def test_empty_is_quiet(self, capsys): - print_diagnostics([]) - captured = capsys.readouterr() - assert captured.out == "" - assert captured.err == "" diff --git a/framework/core/tests/test_diagnostics.py b/framework/core/tests/test_diagnostics.py new file mode 100644 index 00000000..8333b04c --- /dev/null +++ b/framework/core/tests/test_diagnostics.py @@ -0,0 +1,74 @@ +"""Tests for MigrationDiagnostics and print_diagnostics output.""" + +from __future__ import annotations + +from simple_module_core.diagnostics import ( + Diagnostic, + DiagnosticLevel, + MigrationDiagnostics, + print_diagnostics, +) + + +class TestMigrationDiagnostics: + async def test_sm010_migration_mismatch(self): + """SM010 should fire when current revision != head.""" + diag = MigrationDiagnostics() + results = diag.check_revision_mismatch( + current_revision="abc123", + head_revision="def456", + ) + assert len(results) == 1 + assert results[0].code == "SM010" + assert results[0].level == DiagnosticLevel.ERROR + + async def test_sm010_no_error_when_current(self): + """SM010 should not fire when DB is at head.""" + diag = MigrationDiagnostics() + results = diag.check_revision_mismatch( + current_revision="abc123", + head_revision="abc123", + ) + assert len(results) == 0 + + async def test_sm011_missing_tables(self): + """SM011 should fire when module tables aren't in migration tables.""" + diag = MigrationDiagnostics() + results = diag.check_table_coverage( + module_tables={"products_product", "products_category"}, + migrated_tables={"products_product"}, + ) + assert len(results) == 1 + assert results[0].code == "SM011" + assert results[0].level == DiagnosticLevel.WARNING + assert "products_category" in results[0].message + + async def test_sm011_no_warning_when_covered(self): + """SM011 should not fire when all tables are covered.""" + diag = MigrationDiagnostics() + results = diag.check_table_coverage( + module_tables={"products_product"}, + migrated_tables={"products_product"}, + ) + assert len(results) == 0 + + +class TestPrintDiagnostics: + async def test_writes_to_stderr(self, capsys): + diag = Diagnostic( + level=DiagnosticLevel.ERROR, + code="SM001", + message="test error", + module_name="TestMod", + ) + print_diagnostics([diag]) + captured = capsys.readouterr() + assert captured.out == "" + assert "SM001" in captured.err + assert "Results: 1 error(s)" in captured.err + + async def test_empty_is_quiet(self, capsys): + print_diagnostics([]) + captured = capsys.readouterr() + assert captured.out == "" + assert captured.err == "" diff --git a/framework/core/tests/test_discovery.py b/framework/core/tests/test_discovery.py new file mode 100644 index 00000000..30403d1d --- /dev/null +++ b/framework/core/tests/test_discovery.py @@ -0,0 +1,248 @@ +"""Tests for module discovery and topological_sort.""" + +from __future__ import annotations + +import logging + +import pytest +from simple_module_core.discovery import discover_modules, topological_sort +from simple_module_core.exceptions import CircularDependencyError, InvalidModuleError +from simple_module_core.module import ModuleBase, ModuleMeta + + +class ModA(ModuleBase): + meta = ModuleMeta(name="A") + + +class ModB(ModuleBase): + meta = ModuleMeta(name="B", depends_on=["A"]) + + +class ModC(ModuleBase): + meta = ModuleMeta(name="C", depends_on=["B"]) + + +class CycleX(ModuleBase): + meta = ModuleMeta(name="X", depends_on=["Y"]) + + +class CycleY(ModuleBase): + meta = ModuleMeta(name="Y", depends_on=["X"]) + + +class TestTopologicalSort: + async def test_valid_dag(self): + modules = [ModC(), ModA(), ModB()] + sorted_mods = topological_sort(modules) + names = [m.meta.name for m in sorted_mods] + assert names.index("A") < names.index("B") + assert names.index("B") < names.index("C") + + async def test_no_dependencies(self): + modules = [ModA()] + sorted_mods = topological_sort(modules) + assert len(sorted_mods) == 1 + assert sorted_mods[0].meta.name == "A" + + async def test_circular_dependency_raises(self): + modules = [CycleX(), CycleY()] + with pytest.raises(CircularDependencyError): + topological_sort(modules) + + +class TestTopologicalSortEdgeCases: + async def test_diamond_dependency(self): + """A -> B, A -> C, B -> D, C -> D (diamond, not cycle).""" + + class ModD(ModuleBase): + meta = ModuleMeta(name="D") + + class ModB2(ModuleBase): + meta = ModuleMeta(name="B2", depends_on=["D"]) + + class ModC2(ModuleBase): + meta = ModuleMeta(name="C2", depends_on=["D"]) + + class ModA2(ModuleBase): + meta = ModuleMeta(name="A2", depends_on=["B2", "C2"]) + + modules = [ModA2(), ModC2(), ModB2(), ModD()] + sorted_mods = topological_sort(modules) + names = [m.meta.name for m in sorted_mods] + assert names.index("D") < names.index("B2") + assert names.index("D") < names.index("C2") + assert names.index("B2") < names.index("A2") + + async def test_missing_dependency_ignored(self): + """A module depending on a non-installed module should not crash.""" + + class ModWithMissing(ModuleBase): + meta = ModuleMeta(name="Lonely", depends_on=["NonExistent"]) + + sorted_mods = topological_sort([ModWithMissing()]) + assert len(sorted_mods) == 1 + + async def test_self_dependency_raises(self): + """A module that depends on itself is a cycle.""" + + class SelfDep(ModuleBase): + meta = ModuleMeta(name="Self", depends_on=["Self"]) + + with pytest.raises(CircularDependencyError): + topological_sort([SelfDep()]) + + async def test_three_node_cycle(self): + """A -> B -> C -> A should raise.""" + + class CA(ModuleBase): + meta = ModuleMeta(name="CA", depends_on=["CC"]) + + class CB(ModuleBase): + meta = ModuleMeta(name="CB", depends_on=["CA"]) + + class CC(ModuleBase): + meta = ModuleMeta(name="CC", depends_on=["CB"]) + + with pytest.raises(CircularDependencyError): + topological_sort([CA(), CB(), CC()]) + + async def test_empty_list(self): + assert topological_sort([]) == [] + + +class TestDiscoverModules: + async def test_discover_finds_installed_modules(self): + """discover_modules() should find modules registered via entry_points.""" + modules = discover_modules() + names = [m.meta.name for m in modules] + assert "Products" in names + assert "Auth" in names + assert "Dashboard" in names + + +class TestDiscoverModulesAdvanced: + async def test_discover_returns_module_instances(self): + modules = discover_modules() + for mod in modules: + assert isinstance(mod, ModuleBase) + assert hasattr(mod, "meta") + + async def test_discover_modules_have_valid_meta(self): + modules = discover_modules() + for mod in modules: + assert isinstance(mod.meta, ModuleMeta) + assert isinstance(mod.meta.depends_on, list) + assert mod.meta.name != "" + + +class _FakeEntryPoint: + """Minimal EntryPoint shim for testing the validation path. + + Pass a class to return on ``load()``, or a zero-arg callable to + raise/return something custom (for load-failure cases). + """ + + def __init__(self, name: str, target): + self.name = name + self._target = target + + def load(self): + return ( + self._target() + if callable(self._target) and not isinstance(self._target, type) + else self._target + ) + + +def _patch_entry_points(monkeypatch, eps): + import simple_module_core.discovery as discovery_mod + + monkeypatch.setattr(discovery_mod, "entry_points", lambda group: eps) + + +def _boom_loader(): + raise ImportError("boom") + + +class TestDiscoverModulesValidation: + async def test_missing_meta_strict_raises(self, monkeypatch): + class NoMeta(ModuleBase): # intentionally no meta + pass + + _patch_entry_points(monkeypatch, [_FakeEntryPoint("nometa", NoMeta)]) + + with pytest.raises(InvalidModuleError, match="missing 'meta"): + discover_modules(strict=True) + + async def test_missing_meta_non_strict_skips(self, monkeypatch): + class NoMeta(ModuleBase): + pass + + _patch_entry_points(monkeypatch, [_FakeEntryPoint("nometa", NoMeta)]) + + assert discover_modules(strict=False) == [] + + async def test_non_modulebase_strict_raises(self, monkeypatch): + class NotAModule: + pass + + _patch_entry_points(monkeypatch, [_FakeEntryPoint("notmod", NotAModule)]) + + with pytest.raises(InvalidModuleError, match="not a ModuleBase"): + discover_modules(strict=True) + + async def test_load_failure_strict_raises(self, monkeypatch): + _patch_entry_points(monkeypatch, [_FakeEntryPoint("broken", _boom_loader)]) + + with pytest.raises(InvalidModuleError, match="Failed to load"): + discover_modules(strict=True) + + async def test_load_failure_non_strict_logs_and_skips(self, monkeypatch, caplog): + _patch_entry_points(monkeypatch, [_FakeEntryPoint("broken", _boom_loader)]) + + with caplog.at_level(logging.ERROR, logger="simple_module_core.discovery"): + modules = discover_modules(strict=False) + + assert modules == [] + assert any("Failed to load" in r.message for r in caplog.records) + + async def test_meta_must_be_modulemeta_instance(self, monkeypatch): + class BadMeta(ModuleBase): + meta = "not a ModuleMeta" # type: ignore[assignment] + + _patch_entry_points(monkeypatch, [_FakeEntryPoint("bad", BadMeta)]) + + with pytest.raises(InvalidModuleError, match="missing 'meta"): + discover_modules(strict=True) + + +class TestSelectiveModuleLoading: + async def test_discover_with_none_loads_all(self): + """Passing enabled=None keeps existing behaviour (load all installed modules).""" + all_mods = discover_modules(enabled=None) + names = {m.meta.name for m in all_mods} + assert {"Auth", "Products", "Dashboard"}.issubset(names) + + async def test_discover_with_allowlist_filters(self): + """Passing enabled=['Auth'] loads only Auth, even if other modules are installed.""" + filtered = discover_modules(enabled=["Auth"]) + names = [m.meta.name for m in filtered] + assert names == ["Auth"] + + async def test_discover_with_empty_list_loads_none(self): + """Passing enabled=[] loads no modules (explicit opt-out of everything).""" + assert discover_modules(enabled=[]) == [] + + async def test_discover_allowlist_case_insensitive(self): + """Allowlist matching ignores case so 'products' and 'Products' both work.""" + names = [m.meta.name for m in discover_modules(enabled=["products"])] + assert names == ["Products"] + + async def test_discover_unknown_name_logged_and_ignored(self, caplog): + """Names in enabled that don't match any installed module log a warning but don't raise.""" + with caplog.at_level(logging.WARNING, logger="simple_module_core.discovery"): + result = discover_modules(enabled=["Auth", "Nonexistent"]) + + names = [m.meta.name for m in result] + assert names == ["Auth"] + assert any("nonexistent" in rec.message.lower() for rec in caplog.records) diff --git a/framework/core/tests/test_events.py b/framework/core/tests/test_events.py new file mode 100644 index 00000000..a2d5ae6a --- /dev/null +++ b/framework/core/tests/test_events.py @@ -0,0 +1,175 @@ +"""Tests for EventBus: basic subscribe/publish, isolation, async behavior.""" + +from __future__ import annotations + +import asyncio +from dataclasses import dataclass + +from simple_module_core.events import Event, EventBus + + +@dataclass +class OrderCreated(Event): + order_id: int = 0 + + +class TestEventBus: + async def test_subscribe_and_publish(self): + bus = EventBus() + received: list[Event] = [] + + async def handler(event: OrderCreated): + received.append(event) + + bus.subscribe(OrderCreated, handler) + await bus.publish(OrderCreated(order_id=42)) + + assert len(received) == 1 + assert received[0].order_id == 42 # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + + async def test_multiple_handlers(self): + bus = EventBus() + calls: list[str] = [] + + async def handler_a(event: OrderCreated): + calls.append("a") + + async def handler_b(event: OrderCreated): + calls.append("b") + + bus.subscribe(OrderCreated, handler_a) + bus.subscribe(OrderCreated, handler_b) + await bus.publish(OrderCreated()) + + assert "a" in calls + assert "b" in calls + + async def test_no_handlers_no_error(self): + bus = EventBus() + await bus.publish(OrderCreated()) + + async def test_handler_error_does_not_propagate(self): + bus = EventBus() + calls: list[str] = [] + + async def bad_handler(event: OrderCreated): + raise ValueError("boom") + + async def good_handler(event: OrderCreated): + calls.append("ok") + + bus.subscribe(OrderCreated, bad_handler) + bus.subscribe(OrderCreated, good_handler) + await bus.publish(OrderCreated()) + + assert "ok" in calls + + +class TestEventBusAdvanced: + async def test_different_event_types_isolated(self): + """Handlers only receive events of their subscribed type.""" + bus = EventBus() + + @dataclass + class EventA(Event): + pass + + @dataclass + class EventB(Event): + pass + + a_calls: list = [] + b_calls: list = [] + + async def handle_a(e): + a_calls.append(e) + + async def handle_b(e): + b_calls.append(e) + + bus.subscribe(EventA, handle_a) + bus.subscribe(EventB, handle_b) + + await bus.publish(EventA()) + assert len(a_calls) == 1 + assert len(b_calls) == 0 + + await bus.publish(EventB()) + assert len(b_calls) == 1 + + async def test_publish_nowait(self): + """publish_nowait should schedule without blocking.""" + bus = EventBus() + received: list = [] + + async def handler(e): + received.append(e) + + bus.subscribe(OrderCreated, handler) + bus.publish_nowait(OrderCreated(order_id=99)) + await asyncio.sleep(0.05) + assert len(received) == 1 + assert received[0].order_id == 99 + + async def test_subclass_events_do_not_match_parent_subscription(self): + """Subscribing to a base Event class should not receive subclass events.""" + bus = EventBus() + + @dataclass + class Parent(Event): + pass + + @dataclass + class Child(Parent): + pass + + calls: list = [] + + async def parent_handler(e): + calls.append(("parent", e)) + + bus.subscribe(Parent, parent_handler) + await bus.publish(Child()) + + # Child events should not trigger Parent handlers — strict type match. + assert calls == [] + + async def test_publish_with_no_subscribers_returns_none(self): + """publish() should resolve to None when nothing is listening.""" + bus = EventBus() + + @dataclass + class Orphan(Event): + pass + + result = await bus.publish(Orphan()) + assert result is None + + async def test_publish_nowait_with_no_subscribers_is_noop(self): + """publish_nowait() on an unheard event should not raise.""" + bus = EventBus() + + @dataclass + class Orphan(Event): + pass + + bus.publish_nowait(Orphan()) + + async def test_handlers_dispatched_concurrently(self): + """All handlers for an event should run concurrently via gather.""" + bus = EventBus() + order: list[str] = [] + + async def slow(e): + await asyncio.sleep(0.02) + order.append("slow") + + async def fast(e): + order.append("fast") + + bus.subscribe(OrderCreated, slow) + bus.subscribe(OrderCreated, fast) + await bus.publish(OrderCreated(order_id=1)) + + # "fast" should complete before "slow" because they run concurrently. + assert order == ["fast", "slow"] diff --git a/framework/core/tests/test_feature_flags.py b/framework/core/tests/test_feature_flags.py new file mode 100644 index 00000000..75ca4669 --- /dev/null +++ b/framework/core/tests/test_feature_flags.py @@ -0,0 +1,40 @@ +"""Tests for FeatureFlagRegistry: defaults, overrides, listing.""" + +from __future__ import annotations + +from simple_module_core.feature_flags import FeatureFlagDefinition, FeatureFlagRegistry + + +class TestFeatureFlagRegistry: + async def test_add_and_check_default(self): + reg = FeatureFlagRegistry() + reg.add(FeatureFlagDefinition(name="beta_ui", default_enabled=False)) + assert reg.is_enabled("beta_ui") is False + + async def test_default_enabled(self): + reg = FeatureFlagRegistry() + reg.add(FeatureFlagDefinition(name="stable_feature", default_enabled=True)) + assert reg.is_enabled("stable_feature") is True + + async def test_override(self): + reg = FeatureFlagRegistry() + reg.add(FeatureFlagDefinition(name="beta_ui", default_enabled=False)) + reg.set_override("beta_ui", True) + assert reg.is_enabled("beta_ui") is True + + async def test_clear_override(self): + reg = FeatureFlagRegistry() + reg.add(FeatureFlagDefinition(name="beta_ui", default_enabled=False)) + reg.set_override("beta_ui", True) + reg.clear_override("beta_ui") + assert reg.is_enabled("beta_ui") is False + + async def test_unknown_flag_is_disabled(self): + reg = FeatureFlagRegistry() + assert reg.is_enabled("nonexistent") is False + + async def test_all_flags(self): + reg = FeatureFlagRegistry() + reg.add(FeatureFlagDefinition(name="a")) + reg.add(FeatureFlagDefinition(name="b")) + assert len(reg.all_flags) == 2 diff --git a/framework/core/tests/test_health_registry.py b/framework/core/tests/test_health_registry.py new file mode 100644 index 00000000..1c0fb390 --- /dev/null +++ b/framework/core/tests/test_health_registry.py @@ -0,0 +1,48 @@ +"""Tests for HealthRegistry, HealthCheck, HealthCheckResult, HealthStatus.""" + +from __future__ import annotations + +from simple_module_core.health import HealthCheck, HealthCheckResult, HealthRegistry, HealthStatus + + +class TestHealthRegistry: + async def test_add_and_list(self): + reg = HealthRegistry() + + async def check_db() -> HealthCheckResult: + return HealthCheckResult(status=HealthStatus.HEALTHY) + + reg.add(HealthCheck(name="db", check=check_db)) + assert len(reg.all_checks) == 1 + assert reg.all_checks[0].name == "db" + + async def test_empty_registry(self): + reg = HealthRegistry() + assert reg.all_checks == [] + + async def test_multiple_checks(self): + reg = HealthRegistry() + + async def check_a() -> HealthCheckResult: + return HealthCheckResult(status=HealthStatus.HEALTHY) + + async def check_b() -> HealthCheckResult: + return HealthCheckResult(status=HealthStatus.DEGRADED, detail="slow") + + reg.add(HealthCheck(name="a", check=check_a)) + reg.add(HealthCheck(name="b", check=check_b)) + assert len(reg.all_checks) == 2 + + async def test_check_result_defaults(self): + result = HealthCheckResult(status=HealthStatus.HEALTHY) + assert result.detail is None + + async def test_check_result_with_detail(self): + result = HealthCheckResult(status=HealthStatus.DEGRADED, detail="reindexing") + assert result.detail == "reindexing" + + async def test_health_status_ordering(self): + """Verify enum values exist for aggregation logic.""" + assert HealthStatus.HEALTHY == "healthy" + assert HealthStatus.DEGRADED == "degraded" + assert HealthStatus.UNHEALTHY == "unhealthy" diff --git a/framework/core/tests/test_menu.py b/framework/core/tests/test_menu.py new file mode 100644 index 00000000..7ec00bb1 --- /dev/null +++ b/framework/core/tests/test_menu.py @@ -0,0 +1,102 @@ +"""Tests for MenuRegistry: adding items, sorting, filtering, sections.""" + +from __future__ import annotations + +from simple_module_core.menu import MenuItem, MenuRegistry, MenuSection + + +class TestMenuRegistry: + async def test_add_and_all_items(self): + reg = MenuRegistry() + reg.add(MenuItem(label="Dashboard", url="/dashboard", order=1)) + reg.add(MenuItem(label="Products", url="/products", order=2)) + assert len(reg.all_items) == 2 + assert reg.all_items[0].label == "Dashboard" + + async def test_add_many(self): + reg = MenuRegistry() + reg.add_many( + [ + MenuItem(label="A", url="/a", order=1), + MenuItem(label="B", url="/b", order=2), + ] + ) + assert len(reg.all_items) == 2 + + async def test_sorted_by_order(self): + reg = MenuRegistry() + reg.add(MenuItem(label="Z", url="/z", order=99)) + reg.add(MenuItem(label="A", url="/a", order=1)) + assert reg.all_items[0].label == "A" + assert reg.all_items[1].label == "Z" + + async def test_filter_unauthenticated(self): + reg = MenuRegistry() + reg.add(MenuItem(label="Public", url="/pub", requires_auth=False)) + reg.add(MenuItem(label="Private", url="/priv", requires_auth=True)) + + result = reg.get_for_user(is_authenticated=False) + sidebar = result["sidebar"] + assert len(sidebar) == 1 + assert sidebar[0]["label"] == "Public" + + async def test_filter_authenticated_sees_all(self): + reg = MenuRegistry() + reg.add(MenuItem(label="Public", url="/pub", requires_auth=False)) + reg.add(MenuItem(label="Private", url="/priv", requires_auth=True)) + + result = reg.get_for_user(is_authenticated=True) + sidebar = result["sidebar"] + assert len(sidebar) == 2 + + async def test_filter_by_roles(self): + reg = MenuRegistry() + reg.add(MenuItem(label="Admin Panel", url="/admin", roles=["admin"])) + reg.add(MenuItem(label="Dashboard", url="/dash")) + + result = reg.get_for_user(is_authenticated=True, roles=["user"]) + sidebar = result["sidebar"] + labels = [i["label"] for i in sidebar] + assert "Dashboard" in labels + assert "Admin Panel" not in labels + + result = reg.get_for_user(is_authenticated=True, roles=["admin"]) + sidebar = result["sidebar"] + labels = [i["label"] for i in sidebar] + assert "Admin Panel" in labels + + async def test_sections(self): + reg = MenuRegistry() + reg.add(MenuItem(label="Side", url="/s", section=MenuSection.SIDEBAR)) + reg.add(MenuItem(label="Nav", url="/n", section=MenuSection.NAVBAR)) + reg.add(MenuItem(label="Drop", url="/d", section=MenuSection.USER_DROPDOWN)) + + result = reg.get_for_user(is_authenticated=True) + assert len(result["sidebar"]) == 1 + assert len(result["navbar"]) == 1 + assert len(result["userDropdown"]) == 1 + + +class TestMenuRegistryAdvanced: + async def test_multiple_roles_any_match(self): + reg = MenuRegistry() + reg.add(MenuItem(label="Editor", url="/edit", roles=["editor", "admin"])) + result = reg.get_for_user(is_authenticated=True, roles=["editor"]) + assert len(result["sidebar"]) == 1 + + async def test_empty_registry(self): + reg = MenuRegistry() + result = reg.get_for_user(is_authenticated=True) + assert all(len(v) == 0 for v in result.values()) + + async def test_admin_sidebar_section(self): + reg = MenuRegistry() + reg.add(MenuItem(label="Users", url="/admin/users", section=MenuSection.ADMIN_SIDEBAR)) + result = reg.get_for_user(is_authenticated=True) + assert len(result["adminSidebar"]) == 1 + + async def test_icon_preserved(self): + reg = MenuRegistry() + reg.add(MenuItem(label="Home", url="/", icon="home")) + result = reg.get_for_user(is_authenticated=True) + assert result["sidebar"][0]["icon"] == "home" diff --git a/framework/core/tests/test_module_base.py b/framework/core/tests/test_module_base.py new file mode 100644 index 00000000..ef57d107 --- /dev/null +++ b/framework/core/tests/test_module_base.py @@ -0,0 +1,144 @@ +"""Tests for ModuleMeta and ModuleBase lifecycle/hooks.""" + +from __future__ import annotations + +import pytest +from simple_module_core.events import EventBus +from simple_module_core.feature_flags import FeatureFlagRegistry +from simple_module_core.health import HealthRegistry +from simple_module_core.menu import MenuRegistry +from simple_module_core.module import ModuleBase, ModuleMeta +from simple_module_core.permissions import PermissionRegistry + + +class TestModuleMeta: + async def test_defaults(self): + meta = ModuleMeta(name="TestModule") + assert meta.name == "TestModule" + assert meta.route_prefix == "" + assert meta.view_prefix == "" + assert meta.depends_on == [] + assert meta.version == "1.0.0" + + async def test_custom_fields(self): + meta = ModuleMeta( + name="Products", + route_prefix="/api/products", + view_prefix="/products", + depends_on=["Auth"], + version="2.0.0", + ) + assert meta.route_prefix == "/api/products" + assert meta.depends_on == ["Auth"] + assert meta.version == "2.0.0" + + async def test_frozen(self): + meta = ModuleMeta(name="Frozen") + with pytest.raises(AttributeError): + meta.name = "Changed" # type: ignore[misc] # ty: ignore[invalid-assignment] + + +class DummyModule(ModuleBase): + meta = ModuleMeta(name="Dummy", route_prefix="/api/dummy") + + def __init__(self): + self.routes_registered = False + + def register_routes(self, api_router, view_router): + self.routes_registered = True + + +class TestModuleBase: + async def test_subclass_has_meta(self): + mod = DummyModule() + assert mod.meta.name == "Dummy" + + async def test_register_routes_override(self): + mod = DummyModule() + mod.register_routes(None, None) # type: ignore[arg-type] + assert mod.routes_registered is True + + async def test_default_noop_methods(self): + """Default implementations should not raise.""" + mod = DummyModule() + mod.register_menu_items(MenuRegistry()) + mod.register_permissions(PermissionRegistry()) + + +class TestModuleLifecycle: + async def test_on_startup_default_noop(self): + mod = DummyModule() + await mod.on_startup(None) # type: ignore + + async def test_on_shutdown_default_noop(self): + mod = DummyModule() + await mod.on_shutdown(None) # type: ignore + + async def test_register_event_handlers_default_noop(self): + mod = DummyModule() + bus = EventBus() + mod.register_event_handlers(bus) + + async def test_register_feature_flags_default_noop(self): + mod = DummyModule() + reg = FeatureFlagRegistry() + mod.register_feature_flags(reg) + assert len(reg.all_flags) == 0 + + +class TestModuleNewHooks: + async def test_register_exception_handlers_default_noop(self): + mod = DummyModule() + mod.register_exception_handlers(None) # type: ignore + + async def test_register_health_checks_default_noop(self): + mod = DummyModule() + reg = HealthRegistry() + mod.register_health_checks(reg) + assert len(reg.all_checks) == 0 + + async def test_register_settings_default_noop(self): + mod = DummyModule() + mod.register_settings(None) # type: ignore + + +class TestModuleAssetHooks: + async def test_template_dirs_default_empty(self): + """ModuleBase.template_dirs() returns an empty list by default.""" + mod = DummyModule() + assert mod.template_dirs() == [] + + async def test_static_mounts_default_empty(self): + """ModuleBase.static_mounts() returns an empty dict by default.""" + mod = DummyModule() + assert mod.static_mounts() == {} + + async def test_template_dirs_override(self, tmp_path): + """A module can return its own template directory.""" + tpl_dir = tmp_path / "my_templates" + tpl_dir.mkdir() + + class ModWithTpl(ModuleBase): + meta = ModuleMeta(name="WithTpl") + + def template_dirs(self): + return [tpl_dir] + + mod = ModWithTpl() + result = mod.template_dirs() + assert result == [tpl_dir] + + async def test_static_mounts_override(self, tmp_path): + """A module can map URL prefixes to filesystem directories.""" + assets = tmp_path / "assets" + assets.mkdir() + + class ModWithStatic(ModuleBase): + meta = ModuleMeta(name="WithStatic") + + def static_mounts(self): + return {"/modules/with-static": assets} + + mod = ModWithStatic() + mounts = mod.static_mounts() + assert mounts == {"/modules/with-static": assets} diff --git a/framework/core/tests/test_permissions.py b/framework/core/tests/test_permissions.py new file mode 100644 index 00000000..a53602b1 --- /dev/null +++ b/framework/core/tests/test_permissions.py @@ -0,0 +1,88 @@ +"""Tests for PermissionRegistry: groups, roles, admin bypass.""" + +from __future__ import annotations + +from simple_module_core.permissions import PermissionRegistry + + +class TestPermissionRegistry: + async def test_add_group(self): + reg = PermissionRegistry() + reg.add_group("Products", ["products.view", "products.create"]) + assert "products.view" in reg.all_permissions + assert "products.create" in reg.all_permissions + + async def test_add_single(self): + reg = PermissionRegistry() + reg.add("orders.view") + assert reg.has("orders.view") + + async def test_auto_grouping(self): + reg = PermissionRegistry() + reg.add("orders.view") + reg.add("orders.create") + groups = reg.groups + assert any(g.name == "orders" for g in groups) + + async def test_has(self): + reg = PermissionRegistry() + reg.add("test.perm") + assert reg.has("test.perm") is True + assert reg.has("nonexistent") is False + + async def test_admin_role_gets_all(self): + reg = PermissionRegistry() + reg.add_group("Products", ["products.view", "products.edit"]) + perms = reg.get_permissions_for_roles(["admin"]) + assert "products.view" in perms + assert "products.edit" in perms + + async def test_non_admin_gets_none_by_default(self): + reg = PermissionRegistry() + reg.add_group("Products", ["products.view"]) + perms = reg.get_permissions_for_roles(["user"]) + assert len(perms) == 0 + + async def test_custom_role_map(self): + reg = PermissionRegistry() + reg.add_group("Products", ["products.view", "products.edit"]) + role_map = {"editor": ["products.edit"]} + perms = reg.get_permissions_for_roles(["editor"], role_permission_map=role_map) + assert "products.edit" in perms + assert "products.view" not in perms + + async def test_extend_existing_group(self): + reg = PermissionRegistry() + reg.add_group("Products", ["products.view"]) + reg.add_group("Products", ["products.delete"]) + perms = reg.all_permissions + assert "products.view" in perms + assert "products.delete" in perms + + +class TestPermissionRegistryAdvanced: + async def test_no_duplicates(self): + reg = PermissionRegistry() + reg.add("products.view") + reg.add("products.view") + assert reg.all_permissions.count("products.view") == 1 + + async def test_multiple_roles_union(self): + reg = PermissionRegistry() + reg.add_group("Products", ["products.view", "products.edit"]) + role_map = {"viewer": ["products.view"], "editor": ["products.edit"]} + perms = reg.get_permissions_for_roles(["viewer", "editor"], role_permission_map=role_map) + assert "products.view" in perms + assert "products.edit" in perms + + async def test_groups_list(self): + reg = PermissionRegistry() + reg.add_group("Auth", ["auth.login"]) + reg.add_group("Products", ["products.view"]) + assert len(reg.groups) == 2 + + async def test_permissions_sorted(self): + reg = PermissionRegistry() + reg.add("z.last") + reg.add("a.first") + assert reg.all_permissions == ["a.first", "z.last"] diff --git a/framework/core/tests/test_versioning.py b/framework/core/tests/test_versioning.py new file mode 100644 index 00000000..651819c6 --- /dev/null +++ b/framework/core/tests/test_versioning.py @@ -0,0 +1,92 @@ +"""Tests for framework API version compatibility checks (Gap 3).""" + +from __future__ import annotations + +import pytest +from simple_module_core.exceptions import FrameworkVersionError +from simple_module_core.module import ModuleBase, ModuleMeta + + +class TestFrameworkVersion: + async def test_framework_exposes_api_version(self): + """`simple_module_core.FRAMEWORK_API_VERSION` must be importable and semver-shaped.""" + from packaging.version import Version + from simple_module_core import FRAMEWORK_API_VERSION + + assert isinstance(FRAMEWORK_API_VERSION, str) + assert FRAMEWORK_API_VERSION != "" + Version(FRAMEWORK_API_VERSION) # raises if malformed + + async def test_module_meta_accepts_requires_framework(self): + """ModuleMeta should accept an optional requires_framework field.""" + meta = ModuleMeta(name="X", requires_framework=">=1.0,<2.0") + assert meta.requires_framework == ">=1.0,<2.0" + + async def test_module_meta_requires_framework_defaults_to_none(self): + """When not set, requires_framework is None (no compat check applied).""" + meta = ModuleMeta(name="X") + assert meta.requires_framework is None + + async def test_check_compat_passes_when_version_matches(self): + """A module declaring a spec that matches the framework version passes.""" + from simple_module_core import FRAMEWORK_API_VERSION + from simple_module_core.versioning import check_framework_compatibility + + class ModGood(ModuleBase): + meta = ModuleMeta( + name="Good", + requires_framework=f"=={FRAMEWORK_API_VERSION}", + ) + + check_framework_compatibility([ModGood()]) + + async def test_check_compat_raises_on_mismatch(self): + """A module with an unsatisfiable spec raises FrameworkVersionError at boot.""" + from simple_module_core.versioning import check_framework_compatibility + + class ModStale(ModuleBase): + meta = ModuleMeta(name="Stale", requires_framework=">=999.0") + + with pytest.raises(FrameworkVersionError) as exc_info: + check_framework_compatibility([ModStale()]) + + msg = str(exc_info.value) + assert "Stale" in msg + assert ">=999.0" in msg + + async def test_check_compat_skips_modules_without_spec(self): + """Modules that don't declare requires_framework are not checked.""" + from simple_module_core.versioning import check_framework_compatibility + + class ModLegacy(ModuleBase): + meta = ModuleMeta(name="Legacy") + + check_framework_compatibility([ModLegacy()]) + + async def test_check_compat_rejects_malformed_spec(self): + """A malformed version specifier raises FrameworkVersionError (not something cryptic).""" + from simple_module_core.versioning import check_framework_compatibility + + class ModBadSpec(ModuleBase): + meta = ModuleMeta(name="BadSpec", requires_framework="not-a-spec") + + with pytest.raises(FrameworkVersionError) as exc_info: + check_framework_compatibility([ModBadSpec()]) + assert "BadSpec" in str(exc_info.value) + + async def test_check_compat_reports_all_failures(self): + """When multiple modules are incompatible, the error mentions all of them.""" + from simple_module_core.versioning import check_framework_compatibility + + class ModBadA(ModuleBase): + meta = ModuleMeta(name="BadA", requires_framework=">=999.0") + + class ModBadB(ModuleBase): + meta = ModuleMeta(name="BadB", requires_framework=">=999.0") + + with pytest.raises(FrameworkVersionError) as exc_info: + check_framework_compatibility([ModBadA(), ModBadB()]) + + msg = str(exc_info.value) + assert "BadA" in msg + assert "BadB" in msg diff --git a/framework/db/tests/conftest.py b/framework/db/tests/conftest.py new file mode 100644 index 00000000..7feb5cde --- /dev/null +++ b/framework/db/tests/conftest.py @@ -0,0 +1,45 @@ +"""Shared fixtures and test models for the database test suite.""" + +from __future__ import annotations + +from collections.abc import AsyncGenerator + +import pytest +from simple_module_db.base import create_module_base +from simple_module_db.listeners import register_listeners +from simple_module_db.mixins import MultiTenantMixin, SoftDeleteMixin +from simple_module_db.provider import DatabaseProvider +from simple_module_db.session import init_db +from sqlalchemy import String +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import Mapped, mapped_column + +_TenantBase = create_module_base("mt_test", provider=DatabaseProvider.SQLITE) + + +class _TenantItem(_TenantBase, MultiTenantMixin): # ty: ignore[unsupported-base] + __tablename__ = "mt_test_item" + id: Mapped[int] = mapped_column(primary_key=True, autoincrement=True) + name: Mapped[str] = mapped_column(String(100)) + + +class _TenantSoftItem(_TenantBase, MultiTenantMixin, SoftDeleteMixin): # ty: ignore[unsupported-base] + """Combines multi-tenant and soft-delete mixins to test filter composition.""" + + __tablename__ = "mt_test_soft_item" + id: Mapped[int] = mapped_column(primary_key=True, autoincrement=True) + name: Mapped[str] = mapped_column(String(100)) + + +@pytest.fixture +async def tenant_session() -> AsyncGenerator[AsyncSession, None]: + """Session backed by in-memory SQLite with tenant listeners registered.""" + db_state = init_db("sqlite+aiosqlite:///:memory:") + try: + register_listeners(db_state) + async with db_state.engine.begin() as conn: + await conn.run_sync(_TenantBase.metadata.create_all) + async with db_state.session_factory() as session: + yield session + finally: + await db_state.engine.dispose() diff --git a/framework/db/tests/test_base.py b/framework/db/tests/test_base.py new file mode 100644 index 00000000..e66a4ef9 --- /dev/null +++ b/framework/db/tests/test_base.py @@ -0,0 +1,86 @@ +"""Tests for create_module_base, detect_provider, and mixin field surfaces.""" + +from __future__ import annotations + +from simple_module_db.base import create_module_base +from simple_module_db.mixins import AuditMixin, MultiTenantMixin, SoftDeleteMixin, VersionedMixin +from simple_module_db.provider import DatabaseProvider, detect_provider + + +class TestCreateModuleBase: + async def test_returns_declarative_base(self): + base = create_module_base("test_mod_base", provider=DatabaseProvider.SQLITE) + assert hasattr(base, "metadata") + assert base.__abstract__ is True # ty: ignore[unresolved-attribute] + + async def test_caching_same_args(self): + base1 = create_module_base("cache_test", provider=DatabaseProvider.SQLITE) + base2 = create_module_base("cache_test", provider=DatabaseProvider.SQLITE) + assert base1 is base2 + + async def test_different_providers_different_bases(self): + base_sqlite = create_module_base("multi_prov", provider=DatabaseProvider.SQLITE) + base_pg = create_module_base("multi_prov", provider=DatabaseProvider.POSTGRESQL) + assert base_sqlite is not base_pg + + async def test_postgresql_uses_schema(self): + base = create_module_base("schemamod", provider=DatabaseProvider.POSTGRESQL) + assert base.metadata.schema == "schemamod" + + async def test_sqlite_no_schema(self): + base = create_module_base("noschemod", provider=DatabaseProvider.SQLITE) + assert base.metadata.schema is None + + async def test_module_name_stored(self): + base = create_module_base("named_mod", provider=DatabaseProvider.SQLITE) + assert base.__module_name__ == "named_mod" # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + + async def test_all_module_bases_is_deduped(self): + """Re-creating the same module must not grow ``all_module_bases``.""" + from simple_module_db import base as base_mod + + create_module_base("dedupe_test", provider=DatabaseProvider.SQLITE) + before = len(base_mod.all_module_bases) + + create_module_base("dedupe_test", provider=DatabaseProvider.SQLITE) + after = len(base_mod.all_module_bases) + + assert after == before + assert create_module_base("dedupe_test", provider=DatabaseProvider.SQLITE) in ( + base_mod.all_module_bases + ) + + +class TestDetectProvider: + async def test_sqlite(self): + assert detect_provider("sqlite+aiosqlite:///:memory:") == DatabaseProvider.SQLITE + + async def test_postgresql(self): + assert ( + detect_provider("postgresql+asyncpg://user:pass@localhost/db") + == DatabaseProvider.POSTGRESQL + ) + + async def test_postgres_prefix(self): + assert detect_provider("postgres://user:pass@localhost/db") == DatabaseProvider.POSTGRESQL + + +class TestMixins: + async def test_audit_mixin_fields(self): + """AuditMixin should define created_at, updated_at, created_by, updated_by.""" + fields = ["created_at", "updated_at", "created_by", "updated_by"] + for field_name in fields: + assert hasattr(AuditMixin, field_name), f"AuditMixin missing {field_name}" + + async def test_soft_delete_mixin_fields(self): + """SoftDeleteMixin should define is_deleted, deleted_at, deleted_by.""" + fields = ["is_deleted", "deleted_at", "deleted_by"] + for field_name in fields: + assert hasattr(SoftDeleteMixin, field_name), f"SoftDeleteMixin missing {field_name}" + + async def test_versioned_mixin_fields(self): + assert hasattr(VersionedMixin, "version") + + async def test_multi_tenant_mixin_fields(self): + """MultiTenantMixin should define tenant_id.""" + assert hasattr(MultiTenantMixin, "tenant_id") diff --git a/framework/db/tests/test_db.py b/framework/db/tests/test_db.py deleted file mode 100644 index 7292177e..00000000 --- a/framework/db/tests/test_db.py +++ /dev/null @@ -1,676 +0,0 @@ -"""Tests for the database layer: base creation, mixins, session, deps.""" - -from __future__ import annotations - -from collections.abc import AsyncGenerator - -import pytest -from simple_module_db.base import create_module_base -from simple_module_db.listeners import TenantIsolationError, current_tenant_id, register_listeners -from simple_module_db.mixins import AuditMixin, MultiTenantMixin, SoftDeleteMixin, VersionedMixin -from simple_module_db.provider import DatabaseProvider, detect_provider -from simple_module_db.session import DatabaseState, init_db -from sqlalchemy import String, select -from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Mapped, mapped_column - -# ── Test model for multi-tenancy ──────────────────────────────────── -_TenantBase = create_module_base("mt_test", provider=DatabaseProvider.SQLITE) - - -class _TenantItem(_TenantBase, MultiTenantMixin): # ty: ignore[unsupported-base] - __tablename__ = "mt_test_item" - id: Mapped[int] = mapped_column(primary_key=True, autoincrement=True) - name: Mapped[str] = mapped_column(String(100)) - - -class _TenantSoftItem(_TenantBase, MultiTenantMixin, SoftDeleteMixin): # ty: ignore[unsupported-base] - """Combines multi-tenant and soft-delete mixins to test filter composition.""" - - __tablename__ = "mt_test_soft_item" - id: Mapped[int] = mapped_column(primary_key=True, autoincrement=True) - name: Mapped[str] = mapped_column(String(100)) - - -@pytest.fixture -async def tenant_session() -> AsyncGenerator[AsyncSession, None]: - """Session backed by in-memory SQLite with tenant listeners registered.""" - db_state = init_db("sqlite+aiosqlite:///:memory:") - try: - register_listeners(db_state) - async with db_state.engine.begin() as conn: - await conn.run_sync(_TenantBase.metadata.create_all) - async with db_state.session_factory() as session: - yield session - finally: - await db_state.engine.dispose() - - -# ── create_module_base ─────────────────────────────────────────────── - - -class TestCreateModuleBase: - async def test_returns_declarative_base(self): - base = create_module_base("test_mod_base", provider=DatabaseProvider.SQLITE) - # It should be a class that can be used as a SQLAlchemy base - assert hasattr(base, "metadata") - assert base.__abstract__ is True # ty: ignore[unresolved-attribute] - - async def test_caching_same_args(self): - base1 = create_module_base("cache_test", provider=DatabaseProvider.SQLITE) - base2 = create_module_base("cache_test", provider=DatabaseProvider.SQLITE) - assert base1 is base2 - - async def test_different_providers_different_bases(self): - base_sqlite = create_module_base("multi_prov", provider=DatabaseProvider.SQLITE) - base_pg = create_module_base("multi_prov", provider=DatabaseProvider.POSTGRESQL) - assert base_sqlite is not base_pg - - async def test_postgresql_uses_schema(self): - base = create_module_base("schemamod", provider=DatabaseProvider.POSTGRESQL) - assert base.metadata.schema == "schemamod" - - async def test_sqlite_no_schema(self): - base = create_module_base("noschemod", provider=DatabaseProvider.SQLITE) - assert base.metadata.schema is None - - async def test_module_name_stored(self): - base = create_module_base("named_mod", provider=DatabaseProvider.SQLITE) - assert base.__module_name__ == "named_mod" # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] - - async def test_all_module_bases_is_deduped(self): - """Re-creating the same module must not grow ``all_module_bases``.""" - from simple_module_db import base as base_mod - - # Prime both caches, then snapshot. - create_module_base("dedupe_test", provider=DatabaseProvider.SQLITE) - before = len(base_mod.all_module_bases) - - # Second call returns cached base; list length must not change. - create_module_base("dedupe_test", provider=DatabaseProvider.SQLITE) - after = len(base_mod.all_module_bases) - - assert after == before - assert create_module_base("dedupe_test", provider=DatabaseProvider.SQLITE) in ( - base_mod.all_module_bases - ) - - -# ── detect_provider ────────────────────────────────────────────────── - - -class TestDetectProvider: - async def test_sqlite(self): - assert detect_provider("sqlite+aiosqlite:///:memory:") == DatabaseProvider.SQLITE - - async def test_postgresql(self): - assert ( - detect_provider("postgresql+asyncpg://user:pass@localhost/db") - == DatabaseProvider.POSTGRESQL - ) - - async def test_postgres_prefix(self): - assert detect_provider("postgres://user:pass@localhost/db") == DatabaseProvider.POSTGRESQL - - -# ── Mixins ─────────────────────────────────────────────────────────── - - -class TestMixins: - async def test_audit_mixin_fields(self): - """AuditMixin should define created_at, updated_at, created_by, updated_by.""" - fields = ["created_at", "updated_at", "created_by", "updated_by"] - for field_name in fields: - assert hasattr(AuditMixin, field_name), f"AuditMixin missing {field_name}" - - async def test_soft_delete_mixin_fields(self): - """SoftDeleteMixin should define is_deleted, deleted_at, deleted_by.""" - fields = ["is_deleted", "deleted_at", "deleted_by"] - for field_name in fields: - assert hasattr(SoftDeleteMixin, field_name), f"SoftDeleteMixin missing {field_name}" - - async def test_versioned_mixin_fields(self): - assert hasattr(VersionedMixin, "version") - - -# ── init_db / DatabaseState ────────────────────────────────────────── - - -class TestSessionManagement: - async def test_init_db_returns_database_state(self): - """init_db should return a DatabaseState with engine and session factory.""" - db_state = init_db("sqlite+aiosqlite:///:memory:") - try: - assert isinstance(db_state, DatabaseState) - assert db_state.engine is not None - assert db_state.session_factory is not None - finally: - await db_state.engine.dispose() - - async def test_separate_init_db_calls_are_independent(self): - """Two init_db calls should produce independent state.""" - db1 = init_db("sqlite+aiosqlite:///:memory:") - db2 = init_db("sqlite+aiosqlite:///:memory:") - try: - assert db1.engine is not db2.engine - assert db1.session_factory is not db2.session_factory - finally: - await db1.engine.dispose() - await db2.engine.dispose() - - -# ── get_db dependency ──────────────────────────────────────────────── - - -class TestGetDbDependency: - async def test_get_db_yields_session(self): - """get_db should yield an AsyncSession from app.state.db.""" - import contextlib - from unittest.mock import MagicMock - - from simple_module_db.deps import get_db - - db_state = init_db("sqlite+aiosqlite:///:memory:") - try: - mock_request = MagicMock() - mock_request.app.state.db = db_state - - gen = get_db(mock_request) - session = await gen.__anext__() - assert isinstance(session, AsyncSession) - - with contextlib.suppress(StopAsyncIteration): - await gen.__anext__() - finally: - await db_state.engine.dispose() - - -# ── Multi-tenancy ─────────────────────────────────────────────────── - - -class TestMultiTenancy: - """Automatic tenant isolation: auto-populate, query filtering, enforcement.""" - - # ── auto-populate ─────────────────────────────────────── - - async def test_auto_populate_tenant_id(self, tenant_session: AsyncSession): - """New objects should get tenant_id from the current context.""" - token = current_tenant_id.set("tenant-a") - try: - item = _TenantItem(name="Widget") - tenant_session.add(item) - await tenant_session.flush() - assert item.tenant_id == "tenant-a" - finally: - current_tenant_id.reset(token) - - async def test_explicit_tenant_id_preserved(self, tenant_session: AsyncSession): - """Explicitly set tenant_id matching the context should be kept.""" - token = current_tenant_id.set("tenant-a") - try: - item = _TenantItem(name="Explicit", tenant_id="tenant-a") - tenant_session.add(item) - await tenant_session.flush() - assert item.tenant_id == "tenant-a" - finally: - current_tenant_id.reset(token) - - # ── query filtering ───────────────────────────────────── - - async def test_query_returns_only_current_tenant(self, tenant_session: AsyncSession): - """SELECT should be automatically filtered to the current tenant.""" - # Seed two tenants - token_a = current_tenant_id.set("tenant-a") - try: - tenant_session.add_all([_TenantItem(name="A1"), _TenantItem(name="A2")]) - await tenant_session.flush() - finally: - current_tenant_id.reset(token_a) - - token_b = current_tenant_id.set("tenant-b") - try: - tenant_session.add(_TenantItem(name="B1")) - await tenant_session.flush() - finally: - current_tenant_id.reset(token_b) - - # Query as tenant-a - token = current_tenant_id.set("tenant-a") - try: - result = await tenant_session.execute(select(_TenantItem)) - items = result.scalars().all() - assert len(items) == 2 - assert all(i.tenant_id == "tenant-a" for i in items) - finally: - current_tenant_id.reset(token) - - # Query as tenant-b - token = current_tenant_id.set("tenant-b") - try: - result = await tenant_session.execute(select(_TenantItem)) - items = result.scalars().all() - assert len(items) == 1 - assert items[0].name == "B1" - finally: - current_tenant_id.reset(token) - - async def test_no_filtering_without_tenant_context(self, tenant_session: AsyncSession): - """Without a tenant context, all rows should be visible (system query).""" - token_a = current_tenant_id.set("tenant-a") - try: - tenant_session.add(_TenantItem(name="A1")) - await tenant_session.flush() - finally: - current_tenant_id.reset(token_a) - - token_b = current_tenant_id.set("tenant-b") - try: - tenant_session.add(_TenantItem(name="B1")) - await tenant_session.flush() - finally: - current_tenant_id.reset(token_b) - - # No tenant context → see everything - result = await tenant_session.execute(select(_TenantItem)) - items = result.scalars().all() - assert len(items) == 2 - - # ── isolation enforcement ─────────────────────────────── - - async def test_cross_tenant_creation_rejected(self, tenant_session: AsyncSession): - """Creating an object with a tenant_id that doesn't match the context should raise.""" - token = current_tenant_id.set("tenant-a") - try: - item = _TenantItem(name="Imposter", tenant_id="tenant-b") - tenant_session.add(item) - with pytest.raises(TenantIsolationError, match="Cannot create object"): - await tenant_session.flush() - finally: - await tenant_session.rollback() - current_tenant_id.reset(token) - - async def test_tenant_id_change_rejected(self, tenant_session: AsyncSession): - """Changing tenant_id on an existing object should raise.""" - token = current_tenant_id.set("tenant-a") - try: - item = _TenantItem(name="Stable") - tenant_session.add(item) - await tenant_session.flush() - - # Attempt to move to another tenant - item.tenant_id = "tenant-b" - with pytest.raises(TenantIsolationError, match="Cannot change tenant_id"): - await tenant_session.flush() - finally: - await tenant_session.rollback() - current_tenant_id.reset(token) - - async def test_multi_tenant_mixin_fields(self): - """MultiTenantMixin should define tenant_id.""" - assert hasattr(MultiTenantMixin, "tenant_id") - - -# ── Multi-tenancy edge cases ──────────────────────────────────────── - - -class TestMultiTenancyEdgeCases: - """Edge cases: updates, session.get, filter composition, and error recovery.""" - - async def test_update_within_same_tenant(self, tenant_session: AsyncSession): - """Updating a non-tenant field on a same-tenant object should succeed.""" - token = current_tenant_id.set("tenant-a") - try: - item = _TenantItem(name="Original") - tenant_session.add(item) - await tenant_session.flush() - - item.name = "Renamed" - await tenant_session.flush() - - assert item.name == "Renamed" - assert item.tenant_id == "tenant-a" - finally: - current_tenant_id.reset(token) - - async def test_reassigning_same_tenant_id_is_allowed(self, tenant_session: AsyncSession): - """Setting tenant_id to its current value is a no-op, not a violation.""" - token = current_tenant_id.set("tenant-a") - try: - item = _TenantItem(name="A") - tenant_session.add(item) - await tenant_session.flush() - - item.tenant_id = "tenant-a" - await tenant_session.flush() - - assert item.tenant_id == "tenant-a" - finally: - current_tenant_id.reset(token) - - async def test_session_get_respects_tenant_filter(self, tenant_session: AsyncSession): - """session.get() for an item owned by another tenant should return None.""" - # Seed an item for tenant-a - token_a = current_tenant_id.set("tenant-a") - try: - item = _TenantItem(name="Belongs to A") - tenant_session.add(item) - await tenant_session.flush() - tenant_a_item_id = item.id - tenant_session.expunge(item) # Drop from identity map so get() hits DB - finally: - current_tenant_id.reset(token_a) - - # Try to fetch from tenant-b — filter should hide it - token_b = current_tenant_id.set("tenant-b") - try: - found = await tenant_session.get(_TenantItem, tenant_a_item_id) - assert found is None - finally: - current_tenant_id.reset(token_b) - - # Same id, correct tenant → visible - token_a = current_tenant_id.set("tenant-a") - try: - found = await tenant_session.get(_TenantItem, tenant_a_item_id) - assert found is not None - assert found.tenant_id == "tenant-a" - finally: - current_tenant_id.reset(token_a) - - async def test_tenant_filter_composes_with_soft_delete_filter( - self, tenant_session: AsyncSession - ): - """A model using both mixins should be filtered by tenant AND is_deleted=False.""" - # Seed: two items for tenant-a, one alive and one soft-deleted; one item for tenant-b. - token_a = current_tenant_id.set("tenant-a") - try: - alive = _TenantSoftItem(name="alive-a") - deleted = _TenantSoftItem(name="deleted-a") - tenant_session.add_all([alive, deleted]) - await tenant_session.flush() - - await tenant_session.delete(deleted) - await tenant_session.flush() - finally: - current_tenant_id.reset(token_a) - - token_b = current_tenant_id.set("tenant-b") - try: - tenant_session.add(_TenantSoftItem(name="alive-b")) - await tenant_session.flush() - finally: - current_tenant_id.reset(token_b) - - # tenant-a should only see the alive item (soft-deleted one filtered out) - token = current_tenant_id.set("tenant-a") - try: - result = await tenant_session.execute(select(_TenantSoftItem)) - rows = result.scalars().all() - assert len(rows) == 1 - assert rows[0].name == "alive-a" - assert rows[0].is_deleted is False - finally: - current_tenant_id.reset(token) - - async def test_cross_tenant_violation_does_not_leak_state(self, tenant_session: AsyncSession): - """After a raised TenantIsolationError and rollback, the session is usable.""" - token = current_tenant_id.set("tenant-a") - try: - bad = _TenantItem(name="bad", tenant_id="tenant-b") - tenant_session.add(bad) - with pytest.raises(TenantIsolationError): - await tenant_session.flush() - - await tenant_session.rollback() - - # Session should be usable again - good = _TenantItem(name="good") - tenant_session.add(good) - await tenant_session.flush() - assert good.tenant_id == "tenant-a" - finally: - current_tenant_id.reset(token) - - async def test_creation_without_tenant_or_context_fails_at_db( - self, tenant_session: AsyncSession - ): - """No tenant context and no explicit tenant_id → NOT NULL constraint fires.""" - from sqlalchemy.exc import IntegrityError - - item = _TenantItem(name="Orphan") - tenant_session.add(item) - with pytest.raises(IntegrityError): - await tenant_session.flush() - await tenant_session.rollback() - - async def test_system_operation_sees_all_tenants(self, tenant_session: AsyncSession): - """Without a tenant context, a system query can read across all tenants.""" - for tenant in ("tenant-a", "tenant-b", "tenant-c"): - token = current_tenant_id.set(tenant) - try: - tenant_session.add(_TenantItem(name=f"item-{tenant}")) - await tenant_session.flush() - finally: - current_tenant_id.reset(token) - - result = await tenant_session.execute(select(_TenantItem)) - rows = result.scalars().all() - tenants = {r.tenant_id for r in rows} - assert tenants == {"tenant-a", "tenant-b", "tenant-c"} - - async def test_tenant_context_scoped_to_async_task(self, tenant_session: AsyncSession): - """ContextVar changes in one task don't leak into a concurrent task.""" - import asyncio - - seen: list[str | None] = [] - - async def read_tenant() -> None: - # New task inherits the ContextVar snapshot at task creation time - seen.append(current_tenant_id.get()) - - token = current_tenant_id.set("tenant-a") - try: - await asyncio.create_task(read_tenant()) - finally: - current_tenant_id.reset(token) - - assert current_tenant_id.get() is None - # The task saw the value that was current when it was created - assert seen == ["tenant-a"] - - -# ── DB logging ────────────────────────────────────────────────────── - - -class TestGetDbLogging: - """The ``db_state`` fixture handles engine setup/teardown; these tests - just need to create tables and drive ``get_db`` against a mock request. - """ - - @staticmethod - async def _drive_get_db(db_state, populate=None): - """Yield the session, let ``populate`` touch it, then let the - dependency close — this mirrors FastAPI's request lifecycle. - """ - import contextlib - from unittest.mock import MagicMock - - from simple_module_db.deps import get_db - - mock_request = MagicMock() - mock_request.app.state.db = db_state - gen = get_db(mock_request) - session = await gen.__anext__() - if populate is not None: - await populate(session) - with contextlib.suppress(StopAsyncIteration): - await gen.__anext__() - - async def test_commit_logs_on_write(self, db_state, caplog): - import logging - - async with db_state.engine.begin() as conn: - await conn.run_sync(_TenantBase.metadata.create_all) - - async def add_one(session): - session.add(_TenantItem(name="w", tenant_id="t1")) - - with caplog.at_level(logging.INFO, logger="simple_module.db"): - await self._drive_get_db(db_state, populate=add_one) - - commits = [ - r - for r in caplog.records - if r.name == "simple_module.db" and r.message == "db.session.commit" - ] - assert len(commits) == 1 - assert commits[0].operation == "commit" # type: ignore[attr-defined] - assert hasattr(commits[0], "db_duration_ms") - - async def test_commit_fires_even_after_explicit_flush(self, db_state, caplog): - """After flush clears ``session.new`` the ``has_writes`` tag from - the after_flush listener must still drive the commit path. This - guards the real-world pattern used by service.create(). - """ - import logging - - async with db_state.engine.begin() as conn: - await conn.run_sync(_TenantBase.metadata.create_all) - - async def add_and_flush(session): - session.add(_TenantItem(name="w", tenant_id="t1")) - await session.flush() - assert not session.new # flush emptied the live collection - - with caplog.at_level(logging.INFO, logger="simple_module.db"): - await self._drive_get_db(db_state, populate=add_and_flush) - - commits = [ - r - for r in caplog.records - if r.name == "simple_module.db" and r.message == "db.session.commit" - ] - assert len(commits) == 1 - - async def test_read_only_skips_commit(self, db_state, caplog): - import logging - - with caplog.at_level(logging.DEBUG, logger="simple_module.db"): - await self._drive_get_db(db_state) - - records = [r for r in caplog.records if r.name == "simple_module.db"] - assert [r for r in records if r.message == "db.session.commit"] == [] - read_only = [r for r in records if r.message == "db.session.read_only"] - assert len(read_only) == 1 - assert read_only[0].operation == "read_only_rollback" # type: ignore[attr-defined] - - -class TestEntityListenerLogging: - async def test_create_logs_entity_created(self, db_session: AsyncSession, caplog): - """Inserting a new entity should log db.entity.created.""" - import logging - - from products.models import Product - - with caplog.at_level(logging.INFO, logger="simple_module.db"): - product = Product(name="Widget", price=9.99) - db_session.add(product) - await db_session.flush() - - created_msgs = [ - r - for r in caplog.records - if r.name == "simple_module.db" and r.message == "db.entity.created" - ] - assert len(created_msgs) == 1 - assert created_msgs[0].entity == "Product" # type: ignore[attr-defined] - assert created_msgs[0].operation == "create" # type: ignore[attr-defined] - - async def test_update_logs_entity_updated(self, db_session: AsyncSession, caplog): - """Modifying an entity should log db.entity.updated.""" - import logging - - from products.models import Product - - product = Product(name="Widget", price=9.99) - db_session.add(product) - await db_session.flush() - - caplog.clear() - - product.name = "Updated Widget" - with caplog.at_level(logging.INFO, logger="simple_module.db"): - await db_session.flush() - - updated_msgs = [ - r - for r in caplog.records - if r.name == "simple_module.db" and r.message == "db.entity.updated" - ] - assert len(updated_msgs) == 1 - assert updated_msgs[0].entity == "Product" # type: ignore[attr-defined] - assert updated_msgs[0].operation == "update" # type: ignore[attr-defined] - assert updated_msgs[0].entity_id is not None # type: ignore[attr-defined] - - -# ── Migration metadata helper (Gap 1) ─────────────────────────────── - - -class TestMigrationsHelper: - async def test_combined_metadata_includes_installed_module_tables(self): - """ - build_module_metadata() discovers every installed module's models and - aggregates their tables into a single SQLAlchemy MetaData. This is the - helper an Alembic env.py calls to get its `target_metadata`, and it is - the crux of the pip-installed module story: the logic does not care - whether modules are editable installs or wheel installs — it uses - importlib to locate them. - """ - from simple_module_db.migrations import build_module_metadata - - metadata = build_module_metadata() - table_names = set(metadata.tables.keys()) - - # Products ships models and must contribute at least one table. - # (Dashboard is event-driven with no models; Auth's tables are - # currently not part of this workspace's ORM surface.) - assert any("product" in name.lower() for name in table_names) - assert len(table_names) >= 1 - - async def test_combined_metadata_only_returns_module_tables(self): - """ - The helper's allowlist must exclude host-defined or framework-internal - tables so autogenerate doesn't try to drop them on the first run. - """ - from simple_module_db.migrations import build_module_metadata - - metadata = build_module_metadata() - allowlist = set(metadata.tables.keys()) - - # Sanity — table names are non-empty and typed. - assert allowlist - assert all(isinstance(name, str) and name for name in allowlist) - - async def test_include_object_allowlist_filters_unknown_tables(self): - """ - `make_include_object` returns an Alembic `include_object` filter that - passes only tables the framework manages — protecting user-added - tables in the host DB from being destroyed by autogenerate. - """ - from simple_module_db.migrations import build_module_metadata, make_include_object - - metadata = build_module_metadata() - include = make_include_object(metadata) - - # A known module table is included. - some_known_table = next(iter(metadata.tables.values())) - assert include(some_known_table, some_known_table.name, "table", False, None) is True - - # A stranger table (e.g., a host-owned user-auth table) is excluded. - # Alembic passes the candidate Table as the first argument for type_=="table". - from sqlalchemy import Column, Integer, MetaData, Table - - stranger = Table( - "unrelated_host_table", MetaData(), Column("id", Integer, primary_key=True) - ) - assert include(stranger, "unrelated_host_table", "table", False, None) is False diff --git a/framework/db/tests/test_db_logging.py b/framework/db/tests/test_db_logging.py new file mode 100644 index 00000000..64d01f82 --- /dev/null +++ b/framework/db/tests/test_db_logging.py @@ -0,0 +1,128 @@ +"""Tests for session- and entity-level database logging.""" + +from __future__ import annotations + +import contextlib +import logging +from unittest.mock import MagicMock + +from simple_module_db.deps import get_db +from sqlalchemy.ext.asyncio import AsyncSession + +from conftest import _TenantBase, _TenantItem # noqa: E402 # ty: ignore[unresolved-import] + + +async def _drive_get_db(db_state, populate=None): + """Yield the session, let ``populate`` touch it, then let the dependency + close — this mirrors FastAPI's request lifecycle. + """ + mock_request = MagicMock() + mock_request.app.state.db = db_state + gen = get_db(mock_request) + session = await gen.__anext__() + if populate is not None: + await populate(session) + with contextlib.suppress(StopAsyncIteration): + await gen.__anext__() + + +class TestGetDbLogging: + """The ``db_state`` fixture handles engine setup/teardown; these tests + just need to create tables and drive ``get_db`` against a mock request. + """ + + async def test_commit_logs_on_write(self, db_state, caplog): + async with db_state.engine.begin() as conn: + await conn.run_sync(_TenantBase.metadata.create_all) + + async def add_one(session): + session.add(_TenantItem(name="w", tenant_id="t1")) + + with caplog.at_level(logging.INFO, logger="simple_module.db"): + await _drive_get_db(db_state, populate=add_one) + + commits = [ + r + for r in caplog.records + if r.name == "simple_module.db" and r.message == "db.session.commit" + ] + assert len(commits) == 1 + assert commits[0].operation == "commit" # type: ignore[attr-defined] + assert hasattr(commits[0], "db_duration_ms") + + async def test_commit_fires_even_after_explicit_flush(self, db_state, caplog): + """After flush clears ``session.new`` the ``has_writes`` tag from + the after_flush listener must still drive the commit path. This + guards the real-world pattern used by service.create(). + """ + async with db_state.engine.begin() as conn: + await conn.run_sync(_TenantBase.metadata.create_all) + + async def add_and_flush(session): + session.add(_TenantItem(name="w", tenant_id="t1")) + await session.flush() + assert not session.new + + with caplog.at_level(logging.INFO, logger="simple_module.db"): + await _drive_get_db(db_state, populate=add_and_flush) + + commits = [ + r + for r in caplog.records + if r.name == "simple_module.db" and r.message == "db.session.commit" + ] + assert len(commits) == 1 + + async def test_read_only_skips_commit(self, db_state, caplog): + with caplog.at_level(logging.DEBUG, logger="simple_module.db"): + await _drive_get_db(db_state) + + records = [r for r in caplog.records if r.name == "simple_module.db"] + assert [r for r in records if r.message == "db.session.commit"] == [] + read_only = [r for r in records if r.message == "db.session.read_only"] + assert len(read_only) == 1 + assert read_only[0].operation == "read_only_rollback" # type: ignore[attr-defined] + + +class TestEntityListenerLogging: + async def test_create_logs_entity_created(self, db_session: AsyncSession, caplog): + """Inserting a new entity should log db.entity.created.""" + from products.models import Product + + with caplog.at_level(logging.INFO, logger="simple_module.db"): + product = Product(name="Widget", price=9.99) + db_session.add(product) + await db_session.flush() + + created_msgs = [ + r + for r in caplog.records + if r.name == "simple_module.db" and r.message == "db.entity.created" + ] + assert len(created_msgs) == 1 + assert created_msgs[0].entity == "Product" # type: ignore[attr-defined] + assert created_msgs[0].operation == "create" # type: ignore[attr-defined] + + async def test_update_logs_entity_updated(self, db_session: AsyncSession, caplog): + """Modifying an entity should log db.entity.updated.""" + from products.models import Product + + product = Product(name="Widget", price=9.99) + db_session.add(product) + await db_session.flush() + + caplog.clear() + + product.name = "Updated Widget" + with caplog.at_level(logging.INFO, logger="simple_module.db"): + await db_session.flush() + + updated_msgs = [ + r + for r in caplog.records + if r.name == "simple_module.db" and r.message == "db.entity.updated" + ] + assert len(updated_msgs) == 1 + assert updated_msgs[0].entity == "Product" # type: ignore[attr-defined] + assert updated_msgs[0].operation == "update" # type: ignore[attr-defined] + assert updated_msgs[0].entity_id is not None # type: ignore[attr-defined] diff --git a/framework/db/tests/test_migrations.py b/framework/db/tests/test_migrations.py new file mode 100644 index 00000000..cfd4e176 --- /dev/null +++ b/framework/db/tests/test_migrations.py @@ -0,0 +1,56 @@ +"""Tests for build_module_metadata and make_include_object (Gap 1).""" + +from __future__ import annotations + + +class TestMigrationsHelper: + async def test_combined_metadata_includes_installed_module_tables(self): + """build_module_metadata() discovers every installed module's models and + aggregates their tables into a single SQLAlchemy MetaData. This is the + helper an Alembic env.py calls to get its `target_metadata`, and it is + the crux of the pip-installed module story: the logic does not care + whether modules are editable installs or wheel installs — it uses + importlib to locate them. + """ + from simple_module_db.migrations import build_module_metadata + + metadata = build_module_metadata() + table_names = set(metadata.tables.keys()) + + # Products ships models and must contribute at least one table. + # (Dashboard is event-driven with no models; Auth's tables are + # currently not part of this workspace's ORM surface.) + assert any("product" in name.lower() for name in table_names) + assert len(table_names) >= 1 + + async def test_combined_metadata_only_returns_module_tables(self): + """The helper's allowlist must exclude host-defined or framework-internal + tables so autogenerate doesn't try to drop them on the first run. + """ + from simple_module_db.migrations import build_module_metadata + + metadata = build_module_metadata() + allowlist = set(metadata.tables.keys()) + + assert allowlist + assert all(isinstance(name, str) and name for name in allowlist) + + async def test_include_object_allowlist_filters_unknown_tables(self): + """`make_include_object` returns an Alembic `include_object` filter that + passes only tables the framework manages — protecting user-added + tables in the host DB from being destroyed by autogenerate. + """ + from simple_module_db.migrations import build_module_metadata, make_include_object + + metadata = build_module_metadata() + include = make_include_object(metadata) + + some_known_table = next(iter(metadata.tables.values())) + assert include(some_known_table, some_known_table.name, "table", False, None) is True + + from sqlalchemy import Column, Integer, MetaData, Table + + stranger = Table( + "unrelated_host_table", MetaData(), Column("id", Integer, primary_key=True) + ) + assert include(stranger, "unrelated_host_table", "table", False, None) is False diff --git a/framework/db/tests/test_multi_tenancy.py b/framework/db/tests/test_multi_tenancy.py new file mode 100644 index 00000000..b9516a41 --- /dev/null +++ b/framework/db/tests/test_multi_tenancy.py @@ -0,0 +1,274 @@ +"""Tests for automatic tenant isolation: auto-populate, filtering, enforcement.""" + +from __future__ import annotations + +import asyncio + +import pytest +from simple_module_db.listeners import TenantIsolationError, current_tenant_id +from sqlalchemy import select +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession + +from conftest import _TenantItem, _TenantSoftItem # noqa: E402 # ty: ignore[unresolved-import] + + +class TestMultiTenancy: + """Automatic tenant isolation: auto-populate, query filtering, enforcement.""" + + async def test_auto_populate_tenant_id(self, tenant_session: AsyncSession): + """New objects should get tenant_id from the current context.""" + token = current_tenant_id.set("tenant-a") + try: + item = _TenantItem(name="Widget") + tenant_session.add(item) + await tenant_session.flush() + assert item.tenant_id == "tenant-a" + finally: + current_tenant_id.reset(token) + + async def test_explicit_tenant_id_preserved(self, tenant_session: AsyncSession): + """Explicitly set tenant_id matching the context should be kept.""" + token = current_tenant_id.set("tenant-a") + try: + item = _TenantItem(name="Explicit", tenant_id="tenant-a") + tenant_session.add(item) + await tenant_session.flush() + assert item.tenant_id == "tenant-a" + finally: + current_tenant_id.reset(token) + + async def test_query_returns_only_current_tenant(self, tenant_session: AsyncSession): + """SELECT should be automatically filtered to the current tenant.""" + token_a = current_tenant_id.set("tenant-a") + try: + tenant_session.add_all([_TenantItem(name="A1"), _TenantItem(name="A2")]) + await tenant_session.flush() + finally: + current_tenant_id.reset(token_a) + + token_b = current_tenant_id.set("tenant-b") + try: + tenant_session.add(_TenantItem(name="B1")) + await tenant_session.flush() + finally: + current_tenant_id.reset(token_b) + + token = current_tenant_id.set("tenant-a") + try: + result = await tenant_session.execute(select(_TenantItem)) + items = result.scalars().all() + assert len(items) == 2 + assert all(i.tenant_id == "tenant-a" for i in items) + finally: + current_tenant_id.reset(token) + + token = current_tenant_id.set("tenant-b") + try: + result = await tenant_session.execute(select(_TenantItem)) + items = result.scalars().all() + assert len(items) == 1 + assert items[0].name == "B1" + finally: + current_tenant_id.reset(token) + + async def test_no_filtering_without_tenant_context(self, tenant_session: AsyncSession): + """Without a tenant context, all rows should be visible (system query).""" + token_a = current_tenant_id.set("tenant-a") + try: + tenant_session.add(_TenantItem(name="A1")) + await tenant_session.flush() + finally: + current_tenant_id.reset(token_a) + + token_b = current_tenant_id.set("tenant-b") + try: + tenant_session.add(_TenantItem(name="B1")) + await tenant_session.flush() + finally: + current_tenant_id.reset(token_b) + + result = await tenant_session.execute(select(_TenantItem)) + items = result.scalars().all() + assert len(items) == 2 + + async def test_cross_tenant_creation_rejected(self, tenant_session: AsyncSession): + """Creating an object with a tenant_id that doesn't match the context should raise.""" + token = current_tenant_id.set("tenant-a") + try: + item = _TenantItem(name="Imposter", tenant_id="tenant-b") + tenant_session.add(item) + with pytest.raises(TenantIsolationError, match="Cannot create object"): + await tenant_session.flush() + finally: + await tenant_session.rollback() + current_tenant_id.reset(token) + + async def test_tenant_id_change_rejected(self, tenant_session: AsyncSession): + """Changing tenant_id on an existing object should raise.""" + token = current_tenant_id.set("tenant-a") + try: + item = _TenantItem(name="Stable") + tenant_session.add(item) + await tenant_session.flush() + + item.tenant_id = "tenant-b" + with pytest.raises(TenantIsolationError, match="Cannot change tenant_id"): + await tenant_session.flush() + finally: + await tenant_session.rollback() + current_tenant_id.reset(token) + + +class TestMultiTenancyEdgeCases: + """Edge cases: updates, session.get, filter composition, and error recovery.""" + + async def test_update_within_same_tenant(self, tenant_session: AsyncSession): + """Updating a non-tenant field on a same-tenant object should succeed.""" + token = current_tenant_id.set("tenant-a") + try: + item = _TenantItem(name="Original") + tenant_session.add(item) + await tenant_session.flush() + + item.name = "Renamed" + await tenant_session.flush() + + assert item.name == "Renamed" + assert item.tenant_id == "tenant-a" + finally: + current_tenant_id.reset(token) + + async def test_reassigning_same_tenant_id_is_allowed(self, tenant_session: AsyncSession): + """Setting tenant_id to its current value is a no-op, not a violation.""" + token = current_tenant_id.set("tenant-a") + try: + item = _TenantItem(name="A") + tenant_session.add(item) + await tenant_session.flush() + + item.tenant_id = "tenant-a" + await tenant_session.flush() + + assert item.tenant_id == "tenant-a" + finally: + current_tenant_id.reset(token) + + async def test_session_get_respects_tenant_filter(self, tenant_session: AsyncSession): + """session.get() for an item owned by another tenant should return None.""" + token_a = current_tenant_id.set("tenant-a") + try: + item = _TenantItem(name="Belongs to A") + tenant_session.add(item) + await tenant_session.flush() + tenant_a_item_id = item.id + tenant_session.expunge(item) + finally: + current_tenant_id.reset(token_a) + + token_b = current_tenant_id.set("tenant-b") + try: + found = await tenant_session.get(_TenantItem, tenant_a_item_id) + assert found is None + finally: + current_tenant_id.reset(token_b) + + token_a = current_tenant_id.set("tenant-a") + try: + found = await tenant_session.get(_TenantItem, tenant_a_item_id) + assert found is not None + assert found.tenant_id == "tenant-a" + finally: + current_tenant_id.reset(token_a) + + async def test_tenant_filter_composes_with_soft_delete_filter( + self, tenant_session: AsyncSession + ): + """A model using both mixins should be filtered by tenant AND is_deleted=False.""" + token_a = current_tenant_id.set("tenant-a") + try: + alive = _TenantSoftItem(name="alive-a") + deleted = _TenantSoftItem(name="deleted-a") + tenant_session.add_all([alive, deleted]) + await tenant_session.flush() + + await tenant_session.delete(deleted) + await tenant_session.flush() + finally: + current_tenant_id.reset(token_a) + + token_b = current_tenant_id.set("tenant-b") + try: + tenant_session.add(_TenantSoftItem(name="alive-b")) + await tenant_session.flush() + finally: + current_tenant_id.reset(token_b) + + token = current_tenant_id.set("tenant-a") + try: + result = await tenant_session.execute(select(_TenantSoftItem)) + rows = result.scalars().all() + assert len(rows) == 1 + assert rows[0].name == "alive-a" + assert rows[0].is_deleted is False + finally: + current_tenant_id.reset(token) + + async def test_cross_tenant_violation_does_not_leak_state(self, tenant_session: AsyncSession): + """After a raised TenantIsolationError and rollback, the session is usable.""" + token = current_tenant_id.set("tenant-a") + try: + bad = _TenantItem(name="bad", tenant_id="tenant-b") + tenant_session.add(bad) + with pytest.raises(TenantIsolationError): + await tenant_session.flush() + + await tenant_session.rollback() + + good = _TenantItem(name="good") + tenant_session.add(good) + await tenant_session.flush() + assert good.tenant_id == "tenant-a" + finally: + current_tenant_id.reset(token) + + async def test_creation_without_tenant_or_context_fails_at_db( + self, tenant_session: AsyncSession + ): + """No tenant context and no explicit tenant_id → NOT NULL constraint fires.""" + item = _TenantItem(name="Orphan") + tenant_session.add(item) + with pytest.raises(IntegrityError): + await tenant_session.flush() + await tenant_session.rollback() + + async def test_system_operation_sees_all_tenants(self, tenant_session: AsyncSession): + """Without a tenant context, a system query can read across all tenants.""" + for tenant in ("tenant-a", "tenant-b", "tenant-c"): + token = current_tenant_id.set(tenant) + try: + tenant_session.add(_TenantItem(name=f"item-{tenant}")) + await tenant_session.flush() + finally: + current_tenant_id.reset(token) + + result = await tenant_session.execute(select(_TenantItem)) + rows = result.scalars().all() + tenants = {r.tenant_id for r in rows} + assert tenants == {"tenant-a", "tenant-b", "tenant-c"} + + async def test_tenant_context_scoped_to_async_task(self, tenant_session: AsyncSession): + """ContextVar changes in one task don't leak into a concurrent task.""" + seen: list[str | None] = [] + + async def read_tenant() -> None: + seen.append(current_tenant_id.get()) + + token = current_tenant_id.set("tenant-a") + try: + await asyncio.create_task(read_tenant()) + finally: + current_tenant_id.reset(token) + + assert current_tenant_id.get() is None + assert seen == ["tenant-a"] diff --git a/framework/db/tests/test_session.py b/framework/db/tests/test_session.py new file mode 100644 index 00000000..c1b01930 --- /dev/null +++ b/framework/db/tests/test_session.py @@ -0,0 +1,51 @@ +"""Tests for init_db / DatabaseState and the get_db FastAPI dependency.""" + +from __future__ import annotations + +import contextlib +from unittest.mock import MagicMock + +from simple_module_db.deps import get_db +from simple_module_db.session import DatabaseState, init_db +from sqlalchemy.ext.asyncio import AsyncSession + + +class TestSessionManagement: + async def test_init_db_returns_database_state(self): + """init_db should return a DatabaseState with engine and session factory.""" + db_state = init_db("sqlite+aiosqlite:///:memory:") + try: + assert isinstance(db_state, DatabaseState) + assert db_state.engine is not None + assert db_state.session_factory is not None + finally: + await db_state.engine.dispose() + + async def test_separate_init_db_calls_are_independent(self): + """Two init_db calls should produce independent state.""" + db1 = init_db("sqlite+aiosqlite:///:memory:") + db2 = init_db("sqlite+aiosqlite:///:memory:") + try: + assert db1.engine is not db2.engine + assert db1.session_factory is not db2.session_factory + finally: + await db1.engine.dispose() + await db2.engine.dispose() + + +class TestGetDbDependency: + async def test_get_db_yields_session(self): + """get_db should yield an AsyncSession from app.state.db.""" + db_state = init_db("sqlite+aiosqlite:///:memory:") + try: + mock_request = MagicMock() + mock_request.app.state.db = db_state + + gen = get_db(mock_request) + session = await gen.__anext__() + assert isinstance(session, AsyncSession) + + with contextlib.suppress(StopAsyncIteration): + await gen.__anext__() + finally: + await db_state.engine.dispose() diff --git a/framework/hosting/simple_module_hosting/_error_handlers.py b/framework/hosting/simple_module_hosting/_error_handlers.py new file mode 100644 index 00000000..4f9701e5 --- /dev/null +++ b/framework/hosting/simple_module_hosting/_error_handlers.py @@ -0,0 +1,54 @@ +"""Framework-wide exception handlers that render Inertia error pages.""" + +from __future__ import annotations + +import logging + +from fastapi.responses import JSONResponse +from inertia import ( + Inertia, + InertiaConfig, + InertiaVersionConflictException, + inertia_version_conflict_exception_handler, +) +from simple_module_core.exceptions import NotFoundError +from starlette.exceptions import HTTPException +from starlette.requests import Request +from starlette.responses import Response + +logger = logging.getLogger(__name__) + +_INERTIA_ERROR_STATUSES = frozenset({403, 404, 500}) + + +async def render_error_page(request: Request, status_code: int, message: str) -> Response: + config: InertiaConfig = request.app.state.inertia_config + try: + inertia = Inertia(request, config) + response = await inertia.render("Error", {"status": status_code, "message": message}) + response.status_code = status_code + return response + except InertiaVersionConflictException as exc: + return await inertia_version_conflict_exception_handler(request, exc) + except Exception: + # Fallback if Inertia rendering itself fails (e.g. missing session) + logger.exception("Error page rendering failed, falling back to JSON") + return JSONResponse( + status_code=status_code, content={"detail": message or "Internal Server Error"} + ) + + +async def http_exception_handler(request: Request, exc: HTTPException) -> Response: + if exc.status_code in _INERTIA_ERROR_STATUSES: + detail = str(exc.detail) if exc.detail else "" + return await render_error_page(request, exc.status_code, detail) + return JSONResponse(status_code=exc.status_code, content={"detail": exc.detail}) + + +async def not_found_error_handler(request: Request, exc: NotFoundError) -> Response: + return await render_error_page(request, 404, str(exc)) + + +async def unhandled_exception_handler(request: Request, exc: Exception) -> Response: + logger.exception("Unhandled exception: %s", exc) + return await render_error_page(request, 500, "") diff --git a/framework/hosting/simple_module_hosting/_inertia_setup.py b/framework/hosting/simple_module_hosting/_inertia_setup.py new file mode 100644 index 00000000..579787b8 --- /dev/null +++ b/framework/hosting/simple_module_hosting/_inertia_setup.py @@ -0,0 +1,68 @@ +"""Configure fastapi-inertia with the Jinja2 template.""" + +from __future__ import annotations + +import logging +from pathlib import Path + +from fastapi import FastAPI +from inertia import InertiaConfig, inertia_dependency_factory + +from simple_module_hosting.settings import Settings + +logger = logging.getLogger(__name__) + + +def setup_inertia( + app: FastAPI, + settings: Settings, + modules: list, + project_root: Path, +) -> None: + """Configure fastapi-inertia and attach the dependency factory to app.state. + + The host's own ``host/templates`` directory is first in the search path so + it can override module-contributed templates. Each installed module + contributes additional directories via ``ModuleBase.template_dirs()``. + """ + from fastapi.templating import Jinja2Templates + + host_templates = project_root / "host" / "templates" + directories: list[Path] = [] + + if host_templates.is_dir(): + directories.append(host_templates) + else: + logger.warning("Host templates directory not found at %s", host_templates) + + for mod in modules: + for path in mod.template_dirs(): + if Path(path).is_dir(): + directories.append(Path(path)) + else: + logger.warning( + "Module '%s' declared template dir %s but it does not exist", + mod.meta.name, + path, + ) + + if not directories: + logger.warning("No usable template directories — Inertia will fail to render views") + return + + templates = Jinja2Templates(directory=directories) + + inertia_config = InertiaConfig( + environment=settings.environment, # ty: ignore[invalid-argument-type] + version="1.0", + dev_url=settings.vite_dev_url if settings.is_development else "", + templates=templates, + root_template_filename="index.html", + entrypoint_filename="main.tsx", + root_directory=".", + use_flash_errors=True, + ) + + inertia_dep = inertia_dependency_factory(inertia_config) + app.state.inertia_config = inertia_config + app.state.inertia_dependency = inertia_dep diff --git a/framework/hosting/simple_module_hosting/_migrations.py b/framework/hosting/simple_module_hosting/_migrations.py new file mode 100644 index 00000000..692d4d77 --- /dev/null +++ b/framework/hosting/simple_module_hosting/_migrations.py @@ -0,0 +1,58 @@ +"""Alembic migration check performed during app startup.""" + +from __future__ import annotations + +import logging + +logger = logging.getLogger(__name__) + + +async def check_migrations(engine, alembic_ini_path: str = "host/alembic.ini") -> dict: + """Check database migration state. Raises RuntimeError if not at head. + + Returns a dict with migration status for storage on app.state. + """ + from alembic.config import Config as AlembicConfig + from alembic.runtime.migration import MigrationContext + from alembic.script import ScriptDirectory + from alembic.util.exc import CommandError + + _no_migrations = { + "current_revision": None, + "head_revision": None, + "is_current": True, + "pending_count": 0, + } + + try: + alembic_cfg = AlembicConfig(alembic_ini_path) + script = ScriptDirectory.from_config(alembic_cfg) + head = script.get_current_head() + except (CommandError, FileNotFoundError) as exc: + logger.debug("Alembic not available: %s — skipping migration check", exc) + return _no_migrations + + if head is None: + return _no_migrations + + async with engine.connect() as conn: + + def _get_current(sync_conn): + ctx = MigrationContext.configure(sync_conn) + return ctx.get_current_revision() + + current = await conn.run_sync(_get_current) + + if current != head: + pending = list(script.iterate_revisions(head, current)) + raise RuntimeError( + f"Database is {len(pending)} revision(s) behind " + f"(at {current!r}, head is {head!r}). Run: make migrate" + ) + + return { + "current_revision": current, + "head_revision": head, + "is_current": True, + "pending_count": 0, + } diff --git a/framework/hosting/simple_module_hosting/app_builder.py b/framework/hosting/simple_module_hosting/app_builder.py index f29da641..3ef38f6f 100644 --- a/framework/hosting/simple_module_hosting/app_builder.py +++ b/framework/hosting/simple_module_hosting/app_builder.py @@ -9,11 +9,8 @@ from pathlib import Path from fastapi import APIRouter, FastAPI -from fastapi.responses import JSONResponse from fastapi.staticfiles import StaticFiles from inertia import ( - Inertia, - InertiaConfig, InertiaVersionConflictException, inertia_version_conflict_exception_handler, ) @@ -34,9 +31,14 @@ from simple_module_db.session import init_db from starlette.exceptions import HTTPException from starlette.middleware.sessions import SessionMiddleware -from starlette.requests import Request -from starlette.responses import Response +from simple_module_hosting._error_handlers import ( + http_exception_handler, + not_found_error_handler, + unhandled_exception_handler, +) +from simple_module_hosting._inertia_setup import setup_inertia +from simple_module_hosting._migrations import check_migrations from simple_module_hosting.health import router as health_router from simple_module_hosting.middleware import ( CorrelationIdMiddleware, @@ -85,57 +87,6 @@ def wire_module_routes(app: FastAPI, module) -> None: app.include_router(view_router) -async def _check_migrations(engine, alembic_ini_path: str = "host/alembic.ini") -> dict: - """Check database migration state. Raises RuntimeError if not at head. - - Returns a dict with migration status for storage on app.state. - """ - from alembic.config import Config as AlembicConfig - from alembic.runtime.migration import MigrationContext - from alembic.script import ScriptDirectory - from alembic.util.exc import CommandError - - _no_migrations = { - "current_revision": None, - "head_revision": None, - "is_current": True, - "pending_count": 0, - } - - try: - alembic_cfg = AlembicConfig(alembic_ini_path) - script = ScriptDirectory.from_config(alembic_cfg) - head = script.get_current_head() - except (CommandError, FileNotFoundError) as exc: - logger.debug("Alembic not available: %s — skipping migration check", exc) - return _no_migrations - - if head is None: - return _no_migrations - - async with engine.connect() as conn: - - def _get_current(sync_conn): - ctx = MigrationContext.configure(sync_conn) - return ctx.get_current_revision() - - current = await conn.run_sync(_get_current) - - if current != head: - pending = list(script.iterate_revisions(head, current)) - raise RuntimeError( - f"Database is {len(pending)} revision(s) behind " - f"(at {current!r}, head is {head!r}). Run: make migrate" - ) - - return { - "current_revision": current, - "head_revision": head, - "is_current": True, - "pending_count": 0, - } - - def create_app(settings: Settings | None = None) -> FastAPI: """Build and configure the full FastAPI application. @@ -196,7 +147,7 @@ def create_app(settings: Settings | None = None) -> FastAPI: @asynccontextmanager async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: - app.state.migration = await _check_migrations(app.state.db.engine) + app.state.migration = await check_migrations(app.state.db.engine) for mod in modules: await mod.on_startup(app) @@ -252,15 +203,15 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: app.state.db = db_state # ── Phase 7: Inertia + exception handlers ────────────── - _setup_inertia(app, settings, modules) + setup_inertia(app, settings, modules, _PROJECT_ROOT) app.add_exception_handler( InertiaVersionConflictException, inertia_version_conflict_exception_handler, # ty: ignore[invalid-argument-type] ) - app.add_exception_handler(HTTPException, _http_exception_handler) # ty: ignore[invalid-argument-type] - app.add_exception_handler(NotFoundError, _not_found_error_handler) # ty: ignore[invalid-argument-type] - app.add_exception_handler(Exception, _unhandled_exception_handler) + app.add_exception_handler(HTTPException, http_exception_handler) # ty: ignore[invalid-argument-type] + app.add_exception_handler(NotFoundError, not_found_error_handler) # ty: ignore[invalid-argument-type] + app.add_exception_handler(Exception, unhandled_exception_handler) for mod in modules: mod.register_exception_handlers(app) @@ -315,95 +266,6 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: return app -_INERTIA_ERROR_STATUSES = frozenset({403, 404, 500}) - - -async def _render_error_page(request: Request, status_code: int, message: str) -> Response: - config: InertiaConfig = request.app.state.inertia_config - try: - inertia = Inertia(request, config) - response = await inertia.render("Error", {"status": status_code, "message": message}) - response.status_code = status_code - return response - except InertiaVersionConflictException as exc: - return await inertia_version_conflict_exception_handler(request, exc) - except Exception: - # Fallback if Inertia rendering itself fails (e.g. missing session) - logger.exception("Error page rendering failed, falling back to JSON") - return JSONResponse( - status_code=status_code, content={"detail": message or "Internal Server Error"} - ) - - -async def _http_exception_handler(request: Request, exc: HTTPException) -> Response: - if exc.status_code in _INERTIA_ERROR_STATUSES: - detail = str(exc.detail) if exc.detail else "" - return await _render_error_page(request, exc.status_code, detail) - return JSONResponse(status_code=exc.status_code, content={"detail": exc.detail}) - - -async def _not_found_error_handler(request: Request, exc: NotFoundError) -> Response: - return await _render_error_page(request, 404, str(exc)) - - -async def _unhandled_exception_handler(request: Request, exc: Exception) -> Response: - logger.exception("Unhandled exception: %s", exc) - return await _render_error_page(request, 500, "") - - -def _setup_inertia(app: FastAPI, settings: Settings, modules: list) -> None: - """Configure fastapi-inertia with the Jinja2 template. - - The host's own ``host/templates`` directory is first in the search path so - it can override module-contributed templates. Each installed module - contributes additional directories via ``ModuleBase.template_dirs()``. - """ - from fastapi.templating import Jinja2Templates - - host_templates = _PROJECT_ROOT / "host" / "templates" - directories: list[Path] = [] - - if host_templates.is_dir(): - directories.append(host_templates) - else: - logger.warning("Host templates directory not found at %s", host_templates) - - for mod in modules: - for path in mod.template_dirs(): - if Path(path).is_dir(): - directories.append(Path(path)) - else: - logger.warning( - "Module '%s' declared template dir %s but it does not exist", - mod.meta.name, - path, - ) - - if not directories: - logger.warning("No usable template directories — Inertia will fail to render views") - return - - templates = Jinja2Templates(directory=directories) - - inertia_config = InertiaConfig( - environment=settings.environment, # ty: ignore[invalid-argument-type] - version="1.0", - dev_url=settings.vite_dev_url if settings.is_development else "", - templates=templates, - root_template_filename="index.html", - entrypoint_filename="main.tsx", - root_directory=".", - use_flash_errors=True, - ) - - # Register the Inertia dependency globally - from inertia import inertia_dependency_factory - - inertia_dep = inertia_dependency_factory(inertia_config) - app.state.inertia_config = inertia_config - app.state.inertia_dependency = inertia_dep - - def _check_settings_registration(modules: list, added_keys: set[str]) -> None: """SM012: warn if a module overrides register_settings but added nothing to app.state. diff --git a/framework/hosting/tests/test_app.py b/framework/hosting/tests/test_app.py index f17a6659..9d19c070 100644 --- a/framework/hosting/tests/test_app.py +++ b/framework/hosting/tests/test_app.py @@ -1,19 +1,14 @@ -"""Tests for app creation, routing, and overall integration.""" +"""Tests for app creation, routing, protected pages, security, and migration.""" from __future__ import annotations -from types import SimpleNamespace +from collections import defaultdict import httpx -import pytest from fastapi import FastAPI -from simple_module_db import current_tenant_id from simple_module_hosting.app_builder import _resolve_project_root, create_app -from simple_module_hosting.middleware import TenantMiddleware from simple_module_hosting.settings import Settings -# ── App creation ───────────────────────────────────────────────────── - class TestCreateApp: async def test_returns_fastapi_instance(self, settings: Settings): @@ -59,7 +54,6 @@ class FakeStaticMod(ModuleBase): def static_mounts(self): return {"/modules/fakestatic/static": asset_dir} - # Monkey-patch discovery to return our fake module alongside the real ones. real_discover = app_builder.discover_modules def fake_discover(enabled=None, *, strict=False): @@ -72,352 +66,6 @@ def fake_discover(enabled=None, *, strict=False): assert "/modules/fakestatic/static" in paths -# ── Frontend module manifest (Gap 2a) ──────────────────────────────── - - -class TestModulePagesManifest: - async def test_compute_returns_existing_page_dirs(self): - """Returns {ModuleName: Path} for installed modules that ship a pages/ dir.""" - from simple_module_core import discover_modules - from simple_module_hosting.scaffolding import compute_module_pages - - modules = discover_modules() - result = compute_module_pages(modules) - - # Products + Dashboard ship pages/; Auth is API-only (no frontend pages). - assert {"Products", "Dashboard"}.issubset(result.keys()) - assert "Auth" not in result - for name, path in result.items(): - assert path.is_dir(), f"{name} -> {path} should exist" - assert path.name == "pages" - - async def test_compute_skips_modules_without_pages_dir(self, tmp_path, monkeypatch): - """A module whose package has no pages/ dir is omitted (not an error).""" - from simple_module_core import ModuleBase, ModuleMeta - from simple_module_hosting.scaffolding import compute_module_pages - - class HeadlessMod(ModuleBase): - # Its __module__ is tests' package, which has no pages/ dir. - meta = ModuleMeta(name="Headless") - - result = compute_module_pages([HeadlessMod()]) - assert "Headless" not in result - - async def test_write_manifest_emits_json_and_ts(self, tmp_path): - """write_module_pages_manifest emits both the JSON manifest and the TS glob file.""" - import json - - from simple_module_core import discover_modules - from simple_module_hosting.scaffolding import write_module_pages_manifest - - modules = discover_modules() - written = write_module_pages_manifest(modules, tmp_path) - - manifest = tmp_path / "modules.manifest.json" - generated = tmp_path / "modules.generated.ts" - assert manifest.is_file() - assert generated.is_file() - assert written == {"manifest": manifest, "generated": generated} - - data = json.loads(manifest.read_text(encoding="utf-8")) - assert "Products" in data - assert data["Products"].endswith("pages") or data["Products"].endswith("pages/") - - ts = generated.read_text(encoding="utf-8") - # Should contain an import.meta.glob call per discovered module with pages. - assert "import.meta.glob" in ts - assert "Products" in ts - # And a header marking it auto-generated so devs don't hand-edit. - assert "AUTO-GENERATED" in ts or "auto-generated" in ts.lower() - - -# ── Host scaffold (Gap 6) ──────────────────────────────────────────── - - -class TestCreateHost: - async def test_creates_expected_backend_files(self, tmp_path): - """create_host writes the full backend + frontend scaffold.""" - from simple_module_hosting.scaffolding import create_host - - dest = tmp_path / "demo" - create_host(dest, name="demo-host", modules=["Products", "Auth"]) - - for relpath in [ - # Backend - "pyproject.toml", - "main.py", - "alembic.ini", - "migrations/env.py", - "migrations/script.py.mako", - "migrations/versions/.gitkeep", - ".env.example", - ".gitignore", - "README.md", - "Makefile", - # Frontend - "client_app/package.json", - "client_app/tsconfig.json", - "client_app/vite.config.ts", - "client_app/main.tsx", - "client_app/app.tsx", - "client_app/pages.ts", - "client_app/styles.css", - "client_app/pages/Error.tsx", - "templates/index.html", - ]: - assert (dest / relpath).exists(), f"missing: {relpath}" - - async def test_package_json_carries_host_name(self, tmp_path): - """client_app/package.json has its `name` prefixed with the host name.""" - from simple_module_hosting.scaffolding import create_host - - dest = tmp_path / "demo" - create_host(dest, name="my-host", modules=[]) - pkg = (dest / "client_app" / "package.json").read_text(encoding="utf-8") - assert '"name": "my-host-client-app"' in pkg - - async def test_substitutes_host_name_into_pyproject(self, tmp_path): - """The host name lands in pyproject.toml's [project].name field.""" - from simple_module_hosting.scaffolding import create_host - - dest = tmp_path / "demo" - create_host(dest, name="my-acme-app", modules=[]) - pyproject = (dest / "pyproject.toml").read_text(encoding="utf-8") - assert 'name = "my-acme-app"' in pyproject - - async def test_declares_selected_module_deps(self, tmp_path): - """Each module from --with appears as a PyPI dep in pyproject.toml.""" - from simple_module_hosting.scaffolding import create_host - - dest = tmp_path / "demo" - create_host(dest, name="demo", modules=["Products", "Auth"]) - pyproject = (dest / "pyproject.toml").read_text(encoding="utf-8") - # Module names get converted to PyPI names: simple-module-. - assert "simple-module-products" in pyproject - assert "simple-module-auth" in pyproject - - async def test_refuses_existing_non_empty_dir(self, tmp_path): - """create_host aborts if the destination exists and is non-empty — no clobbering.""" - from simple_module_hosting.scaffolding import create_host - - dest = tmp_path / "existing" - dest.mkdir() - (dest / "unrelated.txt").write_text("do not delete me", encoding="utf-8") - - with pytest.raises(FileExistsError): - create_host(dest, name="demo", modules=[]) - - async def test_env_py_uses_shared_helper(self, tmp_path): - """Scaffolded migrations/env.py delegates to the shared helper, not inline logic.""" - from simple_module_hosting.scaffolding import create_host - - dest = tmp_path / "demo" - create_host(dest, name="demo", modules=[]) - env_py = (dest / "migrations" / "env.py").read_text(encoding="utf-8") - assert "build_module_metadata" in env_py - assert "make_include_object" in env_py - # Must NOT embed the old inline loop — that's the refactor we locked in at Gap 1. - assert "for mod in modules:" not in env_py - - async def test_cli_create_host_runs_end_to_end(self, tmp_path): - """The Click `sm create-host` command produces a working scaffold.""" - from click.testing import CliRunner - from simple_module_hosting.cli import main - - runner = CliRunner() - result = runner.invoke( - main, - ["create-host", "smoke-host", "--dest", str(tmp_path / "out"), "--with", "Products"], - ) - assert result.exit_code == 0, result.output - assert (tmp_path / "out" / "main.py").is_file() - assert (tmp_path / "out" / "pyproject.toml").is_file() - assert "simple-module-products" in (tmp_path / "out" / "pyproject.toml").read_text( - encoding="utf-8" - ) - - -class TestCreateModule: - async def test_creates_expected_module_files(self, tmp_path): - """create_module writes a PyPI-ready module package.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-my-feature" - create_module(dest, name="MyFeature") - - for relpath in [ - "pyproject.toml", - "my_feature/__init__.py", - "my_feature/module.py", - "my_feature/endpoints/__init__.py", - "my_feature/endpoints/api.py", - "tests/__init__.py", - "tests/test_module.py", - ".gitignore", - "README.md", - ]: - assert (dest / relpath).is_file(), f"missing: {relpath}" - - async def test_pyproject_declares_entry_point_and_deps(self, tmp_path): - """pyproject.toml sets the entry_point and pins the framework API range.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-my-feature" - create_module(dest, name="MyFeature") - pyproject = (dest / "pyproject.toml").read_text(encoding="utf-8") - - assert 'name = "simple-module-my-feature"' in pyproject - assert "[project.entry-points.simple_module]" in pyproject - assert "my_feature = " in pyproject # entry-point key - assert "simple-module-core" in pyproject - - async def test_module_py_subclasses_module_base(self, tmp_path): - """The generated module.py has a ModuleBase subclass with the right Meta.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-my-feature" - create_module(dest, name="MyFeature") - module_py = (dest / "my_feature" / "module.py").read_text(encoding="utf-8") - - assert "class MyFeatureModule(ModuleBase)" in module_py - assert 'name="MyFeature"' in module_py - assert "requires_framework=" in module_py - - async def test_snake_case_derivation(self, tmp_path): - """Module names with dashes, spaces, or camel case convert to snake_case packages.""" - from simple_module_hosting.scaffolding import create_module - - # Caller supplies a PascalCase-ish name; package dir is snake_case. - dest = tmp_path / "simple-module-order-tracker" - create_module(dest, name="OrderTracker") - assert (dest / "order_tracker" / "module.py").is_file() - - async def test_refuses_existing_non_empty_dir(self, tmp_path): - """create_module aborts rather than clobber an existing directory.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "existing" - dest.mkdir() - (dest / "sentinel").write_text("keep me", encoding="utf-8") - with pytest.raises(FileExistsError): - create_module(dest, name="MyFeature") - - async def test_cli_create_module_runs_end_to_end(self, tmp_path): - """The Click `sm create-module` command produces a working scaffold.""" - from click.testing import CliRunner - from simple_module_hosting.cli import main - - runner = CliRunner() - dest = tmp_path / "simple-module-smoke" - result = runner.invoke( - main, - ["create-module", "Smoke", "--dest", str(dest)], - ) - assert result.exit_code == 0, result.output - assert (dest / "smoke" / "module.py").is_file() - assert "class SmokeModule(ModuleBase)" in (dest / "smoke" / "module.py").read_text( - encoding="utf-8" - ) - - async def test_scaffold_ships_github_workflows(self, tmp_path): - """Gap 8: scaffolded modules include publish.yml + ci.yml.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-widget" - create_module(dest, name="Widget") - - publish = dest / ".github" / "workflows" / "publish.yml" - ci = dest / ".github" / "workflows" / "ci.yml" - assert publish.is_file(), "publish.yml missing" - assert ci.is_file(), "ci.yml missing" - - async def test_publish_workflow_uses_trusted_publishing(self, tmp_path): - """publish.yml must request OIDC token and use pypa/gh-action-pypi-publish.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-widget" - create_module(dest, name="Widget") - publish = (dest / ".github" / "workflows" / "publish.yml").read_text(encoding="utf-8") - - # Trusted publishing requires these two knobs — without them the - # workflow falls back to API-token auth, which defeats the point. - assert "id-token: write" in publish - assert "pypa/gh-action-pypi-publish" in publish - # Should NOT pin a PyPI API token env var — that's the old way. - assert "PYPI_API_TOKEN" not in publish - - async def test_publish_workflow_triggers_on_version_tag(self, tmp_path): - """publish.yml fires only on tag push, not every commit to main.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-widget" - create_module(dest, name="Widget") - publish = (dest / ".github" / "workflows" / "publish.yml").read_text(encoding="utf-8") - assert "tags:" in publish - - async def test_workflows_parse_as_valid_yaml(self, tmp_path): - """Both workflow files must be parseable YAML — catches template substitution bugs.""" - import yaml # PyYAML ships transitively via uvicorn[standard] - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-widget" - create_module(dest, name="Widget") - - for wf in ("publish.yml", "ci.yml"): - path = dest / ".github" / "workflows" / wf - parsed = yaml.safe_load(path.read_text(encoding="utf-8")) - assert isinstance(parsed, dict), f"{wf} did not parse to a mapping" - assert "jobs" in parsed, f"{wf} has no jobs: key" - - async def test_scaffold_has_pages_dir(self, tmp_path): - """Gap 2b: modules intended to ship TSX pages get a pages/ dir from day one.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-widget" - create_module(dest, name="Widget") - pages_dir = dest / "widget" / "pages" - assert pages_dir.is_dir() - # A .gitkeep avoids an empty dir getting lost during git operations. - assert (pages_dir / ".gitkeep").is_file() - - async def test_pyproject_force_includes_static_dist(self, tmp_path): - """Gap 2b: pyproject.toml must ship /static/dist/ inside the wheel.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-widget" - create_module(dest, name="Widget") - pyproject = (dest / "pyproject.toml").read_text(encoding="utf-8") - - # The built JS is normally gitignored, but hatch needs an explicit - # directive to copy it into the wheel at build time. - assert "force-include" in pyproject - assert "widget/static/dist" in pyproject - - async def test_module_py_mounts_static_dist_conditionally(self, tmp_path): - """Generated module.py exposes static_mounts() that tolerates a missing dist/.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-widget" - create_module(dest, name="Widget") - module_py = (dest / "widget" / "module.py").read_text(encoding="utf-8") - - assert "static_mounts" in module_py - # The URL prefix must match the host's convention (Gap 5's docstring). - assert "/modules/widget/static" in module_py - - async def test_gitignore_excludes_built_assets(self, tmp_path): - """Built JS lives in source control's blind spot; only wheels carry it.""" - from simple_module_hosting.scaffolding import create_module - - dest = tmp_path / "simple-module-widget" - create_module(dest, name="Widget") - gitignore = (dest / ".gitignore").read_text(encoding="utf-8") - assert "static/dist" in gitignore - - -# ── Project root resolution ───────────────────────────────────────── - - class TestResolveProjectRoot: async def test_honours_env_override(self, monkeypatch, tmp_path): monkeypatch.setenv("SM_PROJECT_ROOT", str(tmp_path)) @@ -425,55 +73,23 @@ async def test_honours_env_override(self, monkeypatch, tmp_path): async def test_empty_env_var_uses_fallback(self, monkeypatch, tmp_path): # Empty string is falsy — must fall through to the path walk. - # Use the override path to assert fallback runs, since an empty - # value falls through to the workspace-relative fallback. monkeypatch.setenv("SM_PROJECT_ROOT", str(tmp_path)) monkeypatch.setenv("SM_PROJECT_ROOT", "") - # With empty override, _resolve_project_root should not return tmp_path; - # it returns the parents[3] fallback instead. assert _resolve_project_root() != tmp_path -# ── Health endpoints ───────────────────────────────────────────────── - - -class TestHealthEndpoints: - async def test_health(self, client: httpx.AsyncClient): - resp = await client.get("/health") - assert resp.status_code == 200 - assert resp.json()["status"] == "healthy" - - async def test_health_live(self, client: httpx.AsyncClient): - resp = await client.get("/health/live") - assert resp.status_code == 200 - assert resp.json()["status"] == "alive" - - async def test_health_ready(self, client: httpx.AsyncClient): - resp = await client.get("/health/ready") - assert resp.status_code == 200 - data = resp.json() - assert data["status"] == "healthy" - assert "checks" in data - - -# ── Route registration ─────────────────────────────────────────────── - - class TestRouteRegistration: async def test_expected_routes_registered(self, app: FastAPI): """All modules should have their routes registered in the app.""" route_paths = [r.path for r in app.routes if hasattr(r, "path")] - # Health assert "/health" in route_paths assert "/health/live" in route_paths assert "/health/ready" in route_paths - # Products API assert "/api/products/" in route_paths assert "/api/products/{product_id}" in route_paths - # Auth assert "/auth/login" in route_paths assert "/auth/callback" in route_paths assert "/auth/logout" in route_paths @@ -486,8 +102,6 @@ async def test_expected_routes_registered(self, app: FastAPI): async def test_products_api_methods(self, app: FastAPI): """Products endpoints should support the correct HTTP methods.""" - from collections import defaultdict - routes_by_path: dict[str, set[str]] = defaultdict(set) for route in app.routes: if hasattr(route, "path") and hasattr(route, "methods"): @@ -500,9 +114,6 @@ async def test_products_api_methods(self, app: FastAPI): assert "DELETE" in routes_by_path.get("/api/products/{product_id}", set()) -# ── Unauthenticated access to protected pages ─────────────────────── - - class TestProtectedPages: async def test_dashboard_redirects_unauthenticated(self, client: httpx.AsyncClient): resp = await client.get("/dashboard", follow_redirects=False) @@ -515,9 +126,6 @@ async def test_products_page_redirects_unauthenticated(self, client: httpx.Async assert "/auth/login" in resp.headers["location"] -# ── Security headers ───────────────────────────────────────────────── - - class TestSecurityHeaders: async def test_security_headers_present(self, client: httpx.AsyncClient): resp = await client.get("/health") @@ -527,76 +135,6 @@ async def test_security_headers_present(self, client: httpx.AsyncClient): assert resp.headers["referrer-policy"] == "strict-origin-when-cross-origin" -class TestHealthReady: - async def test_ready_includes_module_checks(self, app: FastAPI, client: httpx.AsyncClient): - """If modules registered health checks, /health/ready should include them.""" - from simple_module_core.health import ( - HealthCheck, - HealthCheckResult, - HealthRegistry, - HealthStatus, - ) - - registry: HealthRegistry = app.state.health_registry - - async def check_test_service() -> HealthCheckResult: - return HealthCheckResult(status=HealthStatus.HEALTHY) - - registry.add(HealthCheck(name="test_service", check=check_test_service)) - - resp = await client.get("/health/ready") - assert resp.status_code == 200 - data = resp.json() - assert data["status"] == "healthy" - assert "checks" in data - assert data["checks"]["test_service"]["status"] == "healthy" - - async def test_ready_degraded_status(self, app: FastAPI, client: httpx.AsyncClient): - from simple_module_core.health import ( - HealthCheck, - HealthCheckResult, - HealthRegistry, - HealthStatus, - ) - - registry: HealthRegistry = app.state.health_registry - - async def check_degraded() -> HealthCheckResult: - return HealthCheckResult(status=HealthStatus.DEGRADED, detail="slow") - - registry.add(HealthCheck(name="slow_service", check=check_degraded)) - - resp = await client.get("/health/ready") - data = resp.json() - assert data["status"] == "degraded" - assert data["checks"]["slow_service"]["detail"] == "slow" - - async def test_ready_unhealthy_on_exception(self, app: FastAPI, client: httpx.AsyncClient): - from simple_module_core.health import HealthCheck, HealthRegistry - - registry: HealthRegistry = app.state.health_registry - - async def check_broken(): - raise ConnectionError("connection refused") - - registry.add(HealthCheck(name="broken_service", check=check_broken)) - - resp = await client.get("/health/ready") - data = resp.json() - assert data["status"] == "unhealthy" - assert data["checks"]["broken_service"]["status"] == "unhealthy" - assert "connection refused" in data["checks"]["broken_service"]["detail"] - - async def test_ready_no_checks_is_healthy(self, client: httpx.AsyncClient): - resp = await client.get("/health/ready") - data = resp.json() - assert data["status"] == "healthy" - assert data["checks"] == {} - - -# ── Migration check ───────────────────────────────────────────────── - - class TestHealthMigrationStatus: async def test_health_includes_migration(self, client: httpx.AsyncClient): resp = await client.get("/health") @@ -614,154 +152,3 @@ async def test_app_state_has_migration_info(self, app: FastAPI): assert migration["pending_count"] == 0 assert "current_revision" in migration assert "head_revision" in migration - - -# ── TenantMiddleware ───────────────────────────────────────────────── - - -def _http_scope(headers: list[tuple[bytes, bytes]] | None = None) -> dict: - return { - "type": "http", - "method": "GET", - "path": "/", - "headers": headers or [], - "state": {}, - } - - -async def _noop_receive(): # pragma: no cover - receive is unused in these tests - return {"type": "http.request", "body": b"", "more_body": False} - - -async def _noop_send(message): # pragma: no cover - nothing inspects responses - return None - - -class TestTenantMiddleware: - """Unit tests exercising the raw-ASGI TenantMiddleware directly.""" - - async def test_skips_non_http_scopes(self): - """Lifespan / websocket scopes should pass through unchanged.""" - calls = {"count": 0} - - async def inner_app(scope, receive, send): - calls["count"] += 1 - assert current_tenant_id.get() is None - - mw = TenantMiddleware(inner_app) - await mw({"type": "lifespan"}, _noop_receive, _noop_send) - assert calls["count"] == 1 - - async def test_tenant_from_user_state_sets_context(self): - """If request.state.user.tenant_id is set, it becomes the current tenant.""" - captured: dict = {} - - async def inner_app(scope, receive, send): - captured["tenant_id"] = current_tenant_id.get() - captured["state_tenant_id"] = scope["state"].get("tenant_id") - - scope = _http_scope() - scope["state"]["user"] = SimpleNamespace(tenant_id="acme-corp") - - await TenantMiddleware(inner_app)(scope, _noop_receive, _noop_send) - - assert captured["tenant_id"] == "acme-corp" - assert captured["state_tenant_id"] == "acme-corp" - - async def test_tenant_from_header_fallback(self): - """With no authenticated user, the configured header should be used.""" - captured: dict = {} - - async def inner_app(scope, receive, send): - captured["tenant_id"] = current_tenant_id.get() - - scope = _http_scope(headers=[(b"x-tenant-id", b"header-tenant")]) - await TenantMiddleware(inner_app, header="X-Tenant-ID")(scope, _noop_receive, _noop_send) - - assert captured["tenant_id"] == "header-tenant" - - async def test_header_ignored_when_header_is_none(self): - """``header=None`` disables the header source entirely.""" - captured: dict = {} - - async def inner_app(scope, receive, send): - captured["tenant_id"] = current_tenant_id.get() - - scope = _http_scope(headers=[(b"x-tenant-id", b"header-tenant")]) - await TenantMiddleware(inner_app)(scope, _noop_receive, _noop_send) - - assert captured["tenant_id"] is None - - async def test_user_tenant_id_takes_precedence_over_header(self): - """Authenticated user's tenant_id must win over the X-Tenant-ID header.""" - captured: dict = {} - - async def inner_app(scope, receive, send): - captured["tenant_id"] = current_tenant_id.get() - - scope = _http_scope(headers=[(b"x-tenant-id", b"header-tenant")]) - scope["state"]["user"] = SimpleNamespace(tenant_id="user-tenant") - - await TenantMiddleware(inner_app, header="X-Tenant-ID")(scope, _noop_receive, _noop_send) - - assert captured["tenant_id"] == "user-tenant" - - async def test_no_tenant_leaves_context_unset(self): - """No user tenant + no header means context stays None and state is None.""" - captured: dict = {} - - async def inner_app(scope, receive, send): - captured["tenant_id"] = current_tenant_id.get() - captured["state_tenant_id"] = scope["state"].get("tenant_id") - - await TenantMiddleware(inner_app)(_http_scope(), _noop_receive, _noop_send) - - assert captured["tenant_id"] is None - assert captured["state_tenant_id"] is None - - async def test_context_reset_after_request(self): - """ContextVar must be reset after the inner app returns, even on error.""" - - async def failing_app(scope, receive, send): - raise RuntimeError("boom") - - scope = _http_scope() - scope["state"]["user"] = SimpleNamespace(tenant_id="leaked") - - with pytest.raises(RuntimeError, match="boom"): - await TenantMiddleware(failing_app)(scope, _noop_receive, _noop_send) - - assert current_tenant_id.get() is None - - async def test_user_without_tenant_id_falls_back_to_header(self): - """An authenticated user whose tenant_id is None shouldn't block header fallback.""" - captured: dict = {} - - async def inner_app(scope, receive, send): - captured["tenant_id"] = current_tenant_id.get() - - scope = _http_scope(headers=[(b"x-tenant-id", b"from-header")]) - scope["state"]["user"] = SimpleNamespace(tenant_id=None) - - await TenantMiddleware(inner_app, header="X-Tenant-ID")(scope, _noop_receive, _noop_send) - - assert captured["tenant_id"] == "from-header" - - -class TestTenantMiddlewareIntegration: - async def test_app_pipeline_includes_tenant_middleware(self, app: FastAPI): - """TenantMiddleware should be registered when multi_tenant=True (fixture default).""" - middleware_classes = [m.cls for m in app.user_middleware] - assert TenantMiddleware in middleware_classes - - async def test_tenant_middleware_absent_when_opted_out(self): - """With ``multi_tenant=False`` the middleware must not be installed.""" - single_tenant_settings = Settings( - database_url="sqlite+aiosqlite:///:memory:", - environment="testing", - secret_key="test-secret-key", - multi_tenant=False, - ) - app = create_app(single_tenant_settings) - middleware_classes = [m.cls for m in app.user_middleware] - assert TenantMiddleware not in middleware_classes diff --git a/framework/hosting/tests/test_health.py b/framework/hosting/tests/test_health.py new file mode 100644 index 00000000..90b488a1 --- /dev/null +++ b/framework/hosting/tests/test_health.py @@ -0,0 +1,82 @@ +"""Tests for /health and /health/ready endpoints, including module health checks.""" + +from __future__ import annotations + +import httpx +from fastapi import FastAPI +from simple_module_core.health import ( + HealthCheck, + HealthCheckResult, + HealthRegistry, + HealthStatus, +) + + +class TestHealthEndpoints: + async def test_health(self, client: httpx.AsyncClient): + resp = await client.get("/health") + assert resp.status_code == 200 + assert resp.json()["status"] == "healthy" + + async def test_health_live(self, client: httpx.AsyncClient): + resp = await client.get("/health/live") + assert resp.status_code == 200 + assert resp.json()["status"] == "alive" + + async def test_health_ready(self, client: httpx.AsyncClient): + resp = await client.get("/health/ready") + assert resp.status_code == 200 + data = resp.json() + assert data["status"] == "healthy" + assert "checks" in data + + +class TestHealthReady: + async def test_ready_includes_module_checks(self, app: FastAPI, client: httpx.AsyncClient): + """If modules registered health checks, /health/ready should include them.""" + registry: HealthRegistry = app.state.health_registry + + async def check_test_service() -> HealthCheckResult: + return HealthCheckResult(status=HealthStatus.HEALTHY) + + registry.add(HealthCheck(name="test_service", check=check_test_service)) + + resp = await client.get("/health/ready") + assert resp.status_code == 200 + data = resp.json() + assert data["status"] == "healthy" + assert "checks" in data + assert data["checks"]["test_service"]["status"] == "healthy" + + async def test_ready_degraded_status(self, app: FastAPI, client: httpx.AsyncClient): + registry: HealthRegistry = app.state.health_registry + + async def check_degraded() -> HealthCheckResult: + return HealthCheckResult(status=HealthStatus.DEGRADED, detail="slow") + + registry.add(HealthCheck(name="slow_service", check=check_degraded)) + + resp = await client.get("/health/ready") + data = resp.json() + assert data["status"] == "degraded" + assert data["checks"]["slow_service"]["detail"] == "slow" + + async def test_ready_unhealthy_on_exception(self, app: FastAPI, client: httpx.AsyncClient): + registry: HealthRegistry = app.state.health_registry + + async def check_broken(): + raise ConnectionError("connection refused") + + registry.add(HealthCheck(name="broken_service", check=check_broken)) + + resp = await client.get("/health/ready") + data = resp.json() + assert data["status"] == "unhealthy" + assert data["checks"]["broken_service"]["status"] == "unhealthy" + assert "connection refused" in data["checks"]["broken_service"]["detail"] + + async def test_ready_no_checks_is_healthy(self, client: httpx.AsyncClient): + resp = await client.get("/health/ready") + data = resp.json() + assert data["status"] == "healthy" + assert data["checks"] == {} diff --git a/framework/hosting/tests/test_scaffolding_host.py b/framework/hosting/tests/test_scaffolding_host.py new file mode 100644 index 00000000..20769451 --- /dev/null +++ b/framework/hosting/tests/test_scaffolding_host.py @@ -0,0 +1,157 @@ +"""Tests for the module-pages manifest and `sm create-host` scaffolding.""" + +from __future__ import annotations + +import pytest + + +class TestModulePagesManifest: + async def test_compute_returns_existing_page_dirs(self): + """Returns {ModuleName: Path} for installed modules that ship a pages/ dir.""" + from simple_module_core import discover_modules + from simple_module_hosting.scaffolding import compute_module_pages + + modules = discover_modules() + result = compute_module_pages(modules) + + # Products + Dashboard ship pages/; Auth is API-only (no frontend pages). + assert {"Products", "Dashboard"}.issubset(result.keys()) + assert "Auth" not in result + for name, path in result.items(): + assert path.is_dir(), f"{name} -> {path} should exist" + assert path.name == "pages" + + async def test_compute_skips_modules_without_pages_dir(self, tmp_path, monkeypatch): + """A module whose package has no pages/ dir is omitted (not an error).""" + from simple_module_core import ModuleBase, ModuleMeta + from simple_module_hosting.scaffolding import compute_module_pages + + class HeadlessMod(ModuleBase): + meta = ModuleMeta(name="Headless") + + result = compute_module_pages([HeadlessMod()]) + assert "Headless" not in result + + async def test_write_manifest_emits_json_and_ts(self, tmp_path): + """write_module_pages_manifest emits both the JSON manifest and the TS glob file.""" + import json + + from simple_module_core import discover_modules + from simple_module_hosting.scaffolding import write_module_pages_manifest + + modules = discover_modules() + written = write_module_pages_manifest(modules, tmp_path) + + manifest = tmp_path / "modules.manifest.json" + generated = tmp_path / "modules.generated.ts" + assert manifest.is_file() + assert generated.is_file() + assert written == {"manifest": manifest, "generated": generated} + + data = json.loads(manifest.read_text(encoding="utf-8")) + assert "Products" in data + assert data["Products"].endswith("pages") or data["Products"].endswith("pages/") + + ts = generated.read_text(encoding="utf-8") + assert "import.meta.glob" in ts + assert "Products" in ts + assert "AUTO-GENERATED" in ts or "auto-generated" in ts.lower() + + +class TestCreateHost: + async def test_creates_expected_backend_files(self, tmp_path): + """create_host writes the full backend + frontend scaffold.""" + from simple_module_hosting.scaffolding import create_host + + dest = tmp_path / "demo" + create_host(dest, name="demo-host", modules=["Products", "Auth"]) + + for relpath in [ + "pyproject.toml", + "main.py", + "alembic.ini", + "migrations/env.py", + "migrations/script.py.mako", + "migrations/versions/.gitkeep", + ".env.example", + ".gitignore", + "README.md", + "Makefile", + "client_app/package.json", + "client_app/tsconfig.json", + "client_app/vite.config.ts", + "client_app/main.tsx", + "client_app/app.tsx", + "client_app/pages.ts", + "client_app/styles.css", + "client_app/pages/Error.tsx", + "templates/index.html", + ]: + assert (dest / relpath).exists(), f"missing: {relpath}" + + async def test_package_json_carries_host_name(self, tmp_path): + """client_app/package.json has its `name` prefixed with the host name.""" + from simple_module_hosting.scaffolding import create_host + + dest = tmp_path / "demo" + create_host(dest, name="my-host", modules=[]) + pkg = (dest / "client_app" / "package.json").read_text(encoding="utf-8") + assert '"name": "my-host-client-app"' in pkg + + async def test_substitutes_host_name_into_pyproject(self, tmp_path): + """The host name lands in pyproject.toml's [project].name field.""" + from simple_module_hosting.scaffolding import create_host + + dest = tmp_path / "demo" + create_host(dest, name="my-acme-app", modules=[]) + pyproject = (dest / "pyproject.toml").read_text(encoding="utf-8") + assert 'name = "my-acme-app"' in pyproject + + async def test_declares_selected_module_deps(self, tmp_path): + """Each module from --with appears as a PyPI dep in pyproject.toml.""" + from simple_module_hosting.scaffolding import create_host + + dest = tmp_path / "demo" + create_host(dest, name="demo", modules=["Products", "Auth"]) + pyproject = (dest / "pyproject.toml").read_text(encoding="utf-8") + assert "simple-module-products" in pyproject + assert "simple-module-auth" in pyproject + + async def test_refuses_existing_non_empty_dir(self, tmp_path): + """create_host aborts if the destination exists and is non-empty — no clobbering.""" + from simple_module_hosting.scaffolding import create_host + + dest = tmp_path / "existing" + dest.mkdir() + (dest / "unrelated.txt").write_text("do not delete me", encoding="utf-8") + + with pytest.raises(FileExistsError): + create_host(dest, name="demo", modules=[]) + + async def test_env_py_uses_shared_helper(self, tmp_path): + """Scaffolded migrations/env.py delegates to the shared helper, not inline logic.""" + from simple_module_hosting.scaffolding import create_host + + dest = tmp_path / "demo" + create_host(dest, name="demo", modules=[]) + env_py = (dest / "migrations" / "env.py").read_text(encoding="utf-8") + assert "build_module_metadata" in env_py + assert "make_include_object" in env_py + assert "for mod in modules:" not in env_py + + async def test_cli_create_host_runs_end_to_end(self, tmp_path): + """The Click `sm create-host` command produces a working scaffold.""" + from click.testing import CliRunner + from simple_module_hosting.cli import main + + runner = CliRunner() + result = runner.invoke( + main, + ["create-host", "smoke-host", "--dest", str(tmp_path / "out"), "--with", "Products"], + ) + assert result.exit_code == 0, result.output + assert (tmp_path / "out" / "main.py").is_file() + assert (tmp_path / "out" / "pyproject.toml").is_file() + assert "simple-module-products" in (tmp_path / "out" / "pyproject.toml").read_text( + encoding="utf-8" + ) diff --git a/framework/hosting/tests/test_scaffolding_module.py b/framework/hosting/tests/test_scaffolding_module.py new file mode 100644 index 00000000..b8cdeed8 --- /dev/null +++ b/framework/hosting/tests/test_scaffolding_module.py @@ -0,0 +1,179 @@ +"""Tests for `sm create-module` scaffolding: module package, CI, static bundling.""" + +from __future__ import annotations + +import pytest + + +class TestCreateModule: + async def test_creates_expected_module_files(self, tmp_path): + """create_module writes a PyPI-ready module package.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-my-feature" + create_module(dest, name="MyFeature") + + for relpath in [ + "pyproject.toml", + "my_feature/__init__.py", + "my_feature/module.py", + "my_feature/endpoints/__init__.py", + "my_feature/endpoints/api.py", + "tests/__init__.py", + "tests/test_module.py", + ".gitignore", + "README.md", + ]: + assert (dest / relpath).is_file(), f"missing: {relpath}" + + async def test_pyproject_declares_entry_point_and_deps(self, tmp_path): + """pyproject.toml sets the entry_point and pins the framework API range.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-my-feature" + create_module(dest, name="MyFeature") + pyproject = (dest / "pyproject.toml").read_text(encoding="utf-8") + + assert 'name = "simple-module-my-feature"' in pyproject + assert "[project.entry-points.simple_module]" in pyproject + assert "my_feature = " in pyproject + assert "simple-module-core" in pyproject + + async def test_module_py_subclasses_module_base(self, tmp_path): + """The generated module.py has a ModuleBase subclass with the right Meta.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-my-feature" + create_module(dest, name="MyFeature") + module_py = (dest / "my_feature" / "module.py").read_text(encoding="utf-8") + + assert "class MyFeatureModule(ModuleBase)" in module_py + assert 'name="MyFeature"' in module_py + assert "requires_framework=" in module_py + + async def test_snake_case_derivation(self, tmp_path): + """Module names with dashes, spaces, or camel case convert to snake_case packages.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-order-tracker" + create_module(dest, name="OrderTracker") + assert (dest / "order_tracker" / "module.py").is_file() + + async def test_refuses_existing_non_empty_dir(self, tmp_path): + """create_module aborts rather than clobber an existing directory.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "existing" + dest.mkdir() + (dest / "sentinel").write_text("keep me", encoding="utf-8") + with pytest.raises(FileExistsError): + create_module(dest, name="MyFeature") + + async def test_cli_create_module_runs_end_to_end(self, tmp_path): + """The Click `sm create-module` command produces a working scaffold.""" + from click.testing import CliRunner + from simple_module_hosting.cli import main + + runner = CliRunner() + dest = tmp_path / "simple-module-smoke" + result = runner.invoke( + main, + ["create-module", "Smoke", "--dest", str(dest)], + ) + assert result.exit_code == 0, result.output + assert (dest / "smoke" / "module.py").is_file() + assert "class SmokeModule(ModuleBase)" in (dest / "smoke" / "module.py").read_text( + encoding="utf-8" + ) + + async def test_scaffold_ships_github_workflows(self, tmp_path): + """Gap 8: scaffolded modules include publish.yml + ci.yml.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-widget" + create_module(dest, name="Widget") + + publish = dest / ".github" / "workflows" / "publish.yml" + ci = dest / ".github" / "workflows" / "ci.yml" + assert publish.is_file(), "publish.yml missing" + assert ci.is_file(), "ci.yml missing" + + async def test_publish_workflow_uses_trusted_publishing(self, tmp_path): + """publish.yml must request OIDC token and use pypa/gh-action-pypi-publish.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-widget" + create_module(dest, name="Widget") + publish = (dest / ".github" / "workflows" / "publish.yml").read_text(encoding="utf-8") + + # Trusted publishing requires these two knobs — without them the + # workflow falls back to API-token auth, which defeats the point. + assert "id-token: write" in publish + assert "pypa/gh-action-pypi-publish" in publish + assert "PYPI_API_TOKEN" not in publish + + async def test_publish_workflow_triggers_on_version_tag(self, tmp_path): + """publish.yml fires only on tag push, not every commit to main.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-widget" + create_module(dest, name="Widget") + publish = (dest / ".github" / "workflows" / "publish.yml").read_text(encoding="utf-8") + assert "tags:" in publish + + async def test_workflows_parse_as_valid_yaml(self, tmp_path): + """Both workflow files must be parseable YAML — catches template substitution bugs.""" + import yaml + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-widget" + create_module(dest, name="Widget") + + for wf in ("publish.yml", "ci.yml"): + path = dest / ".github" / "workflows" / wf + parsed = yaml.safe_load(path.read_text(encoding="utf-8")) + assert isinstance(parsed, dict), f"{wf} did not parse to a mapping" + assert "jobs" in parsed, f"{wf} has no jobs: key" + + async def test_scaffold_has_pages_dir(self, tmp_path): + """Gap 2b: modules intended to ship TSX pages get a pages/ dir from day one.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-widget" + create_module(dest, name="Widget") + pages_dir = dest / "widget" / "pages" + assert pages_dir.is_dir() + assert (pages_dir / ".gitkeep").is_file() + + async def test_pyproject_force_includes_static_dist(self, tmp_path): + """Gap 2b: pyproject.toml must ship /static/dist/ inside the wheel.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-widget" + create_module(dest, name="Widget") + pyproject = (dest / "pyproject.toml").read_text(encoding="utf-8") + + # The built JS is normally gitignored, but hatch needs an explicit + # directive to copy it into the wheel at build time. + assert "force-include" in pyproject + assert "widget/static/dist" in pyproject + + async def test_module_py_mounts_static_dist_conditionally(self, tmp_path): + """Generated module.py exposes static_mounts() that tolerates a missing dist/.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-widget" + create_module(dest, name="Widget") + module_py = (dest / "widget" / "module.py").read_text(encoding="utf-8") + + assert "static_mounts" in module_py + assert "/modules/widget/static" in module_py + + async def test_gitignore_excludes_built_assets(self, tmp_path): + """Built JS lives in source control's blind spot; only wheels carry it.""" + from simple_module_hosting.scaffolding import create_module + + dest = tmp_path / "simple-module-widget" + create_module(dest, name="Widget") + gitignore = (dest / ".gitignore").read_text(encoding="utf-8") + assert "static/dist" in gitignore diff --git a/framework/hosting/tests/test_tenant_middleware.py b/framework/hosting/tests/test_tenant_middleware.py new file mode 100644 index 00000000..5acb429e --- /dev/null +++ b/framework/hosting/tests/test_tenant_middleware.py @@ -0,0 +1,160 @@ +"""Tests for TenantMiddleware: user/header sources, context lifecycle, opt-in.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +from fastapi import FastAPI +from simple_module_db import current_tenant_id +from simple_module_hosting.app_builder import create_app +from simple_module_hosting.middleware import TenantMiddleware +from simple_module_hosting.settings import Settings + + +def _http_scope(headers: list[tuple[bytes, bytes]] | None = None) -> dict: + return { + "type": "http", + "method": "GET", + "path": "/", + "headers": headers or [], + "state": {}, + } + + +async def _noop_receive(): # pragma: no cover - receive is unused in these tests + return {"type": "http.request", "body": b"", "more_body": False} + + +async def _noop_send(message): # pragma: no cover - nothing inspects responses + return None + + +class TestTenantMiddleware: + """Unit tests exercising the raw-ASGI TenantMiddleware directly.""" + + async def test_skips_non_http_scopes(self): + """Lifespan / websocket scopes should pass through unchanged.""" + calls = {"count": 0} + + async def inner_app(scope, receive, send): + calls["count"] += 1 + assert current_tenant_id.get() is None + + mw = TenantMiddleware(inner_app) + await mw({"type": "lifespan"}, _noop_receive, _noop_send) + assert calls["count"] == 1 + + async def test_tenant_from_user_state_sets_context(self): + """If request.state.user.tenant_id is set, it becomes the current tenant.""" + captured: dict = {} + + async def inner_app(scope, receive, send): + captured["tenant_id"] = current_tenant_id.get() + captured["state_tenant_id"] = scope["state"].get("tenant_id") + + scope = _http_scope() + scope["state"]["user"] = SimpleNamespace(tenant_id="acme-corp") + + await TenantMiddleware(inner_app)(scope, _noop_receive, _noop_send) + + assert captured["tenant_id"] == "acme-corp" + assert captured["state_tenant_id"] == "acme-corp" + + async def test_tenant_from_header_fallback(self): + """With no authenticated user, the configured header should be used.""" + captured: dict = {} + + async def inner_app(scope, receive, send): + captured["tenant_id"] = current_tenant_id.get() + + scope = _http_scope(headers=[(b"x-tenant-id", b"header-tenant")]) + await TenantMiddleware(inner_app, header="X-Tenant-ID")(scope, _noop_receive, _noop_send) + + assert captured["tenant_id"] == "header-tenant" + + async def test_header_ignored_when_header_is_none(self): + """``header=None`` disables the header source entirely.""" + captured: dict = {} + + async def inner_app(scope, receive, send): + captured["tenant_id"] = current_tenant_id.get() + + scope = _http_scope(headers=[(b"x-tenant-id", b"header-tenant")]) + await TenantMiddleware(inner_app)(scope, _noop_receive, _noop_send) + + assert captured["tenant_id"] is None + + async def test_user_tenant_id_takes_precedence_over_header(self): + """Authenticated user's tenant_id must win over the X-Tenant-ID header.""" + captured: dict = {} + + async def inner_app(scope, receive, send): + captured["tenant_id"] = current_tenant_id.get() + + scope = _http_scope(headers=[(b"x-tenant-id", b"header-tenant")]) + scope["state"]["user"] = SimpleNamespace(tenant_id="user-tenant") + + await TenantMiddleware(inner_app, header="X-Tenant-ID")(scope, _noop_receive, _noop_send) + + assert captured["tenant_id"] == "user-tenant" + + async def test_no_tenant_leaves_context_unset(self): + """No user tenant + no header means context stays None and state is None.""" + captured: dict = {} + + async def inner_app(scope, receive, send): + captured["tenant_id"] = current_tenant_id.get() + captured["state_tenant_id"] = scope["state"].get("tenant_id") + + await TenantMiddleware(inner_app)(_http_scope(), _noop_receive, _noop_send) + + assert captured["tenant_id"] is None + assert captured["state_tenant_id"] is None + + async def test_context_reset_after_request(self): + """ContextVar must be reset after the inner app returns, even on error.""" + + async def failing_app(scope, receive, send): + raise RuntimeError("boom") + + scope = _http_scope() + scope["state"]["user"] = SimpleNamespace(tenant_id="leaked") + + with pytest.raises(RuntimeError, match="boom"): + await TenantMiddleware(failing_app)(scope, _noop_receive, _noop_send) + + assert current_tenant_id.get() is None + + async def test_user_without_tenant_id_falls_back_to_header(self): + """An authenticated user whose tenant_id is None shouldn't block header fallback.""" + captured: dict = {} + + async def inner_app(scope, receive, send): + captured["tenant_id"] = current_tenant_id.get() + + scope = _http_scope(headers=[(b"x-tenant-id", b"from-header")]) + scope["state"]["user"] = SimpleNamespace(tenant_id=None) + + await TenantMiddleware(inner_app, header="X-Tenant-ID")(scope, _noop_receive, _noop_send) + + assert captured["tenant_id"] == "from-header" + + +class TestTenantMiddlewareIntegration: + async def test_app_pipeline_includes_tenant_middleware(self, app: FastAPI): + """TenantMiddleware should be registered when multi_tenant=True (fixture default).""" + middleware_classes = [m.cls for m in app.user_middleware] + assert TenantMiddleware in middleware_classes + + async def test_tenant_middleware_absent_when_opted_out(self): + """With ``multi_tenant=False`` the middleware must not be installed.""" + single_tenant_settings = Settings( + database_url="sqlite+aiosqlite:///:memory:", + environment="testing", + secret_key="test-secret-key", + multi_tenant=False, + ) + app = create_app(single_tenant_settings) + middleware_classes = [m.cls for m in app.user_middleware] + assert TenantMiddleware not in middleware_classes diff --git a/modules/auth/tests/test_auth.py b/modules/auth/tests/test_auth.py deleted file mode 100644 index 15850c61..00000000 --- a/modules/auth/tests/test_auth.py +++ /dev/null @@ -1,375 +0,0 @@ -"""Tests for the Auth module: UserContext, dependencies, middleware, endpoints.""" - -from __future__ import annotations - -import httpx -import pytest -from auth.contracts.schemas import UserContext - -# ── UserContext ─────────────────────────────────────────────────────── - - -class TestUserContext: - async def test_from_keycloak_userinfo_basic(self): - userinfo = { - "sub": "user-123", - "email": "alice@example.com", - "name": "Alice Smith", - "realm_access": {"roles": ["user", "editor"]}, - } - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert ctx.id == "user-123" - assert ctx.email == "alice@example.com" - assert ctx.name == "Alice Smith" - assert ctx.roles == ["user", "editor"] - - async def test_from_keycloak_userinfo_fallback_username(self): - userinfo = { - "sub": "user-456", - "email": "bob@example.com", - "preferred_username": "bob", - } - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert ctx.name == "bob" - - async def test_from_keycloak_userinfo_missing_fields(self): - userinfo = {} - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert ctx.id == "" - assert ctx.email == "" - assert ctx.name == "" - assert ctx.roles == [] - - async def test_has_role(self): - ctx = UserContext(id="1", email="a@b.com", name="A", roles=["admin", "user"]) - assert ctx.has_role("admin") is True - assert ctx.has_role("superadmin") is False - - async def test_has_any_role(self): - ctx = UserContext(id="1", email="a@b.com", name="A", roles=["editor"]) - assert ctx.has_any_role(["admin", "editor"]) is True - assert ctx.has_any_role(["admin", "superadmin"]) is False - - -# ── UserContext tenant_id ───────────────────────────────────────────── - - -class TestUserContextTenantId: - async def test_tenant_id_from_custom_claim(self): - """UserContext should read a custom ``tenant_id`` claim from userinfo.""" - userinfo = { - "sub": "user-123", - "tenant_id": "acme-corp", - } - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert ctx.tenant_id == "acme-corp" - - async def test_tenant_id_from_organization_claim(self): - """UserContext should fall back to Keycloak's organization.id claim.""" - userinfo = { - "sub": "user-123", - "organization": {"id": "org-42", "name": "Acme"}, - } - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert ctx.tenant_id == "org-42" - - async def test_tenant_id_custom_claim_takes_precedence(self): - """Custom tenant_id claim should take precedence over organization.id.""" - userinfo = { - "sub": "user-123", - "tenant_id": "custom-tenant", - "organization": {"id": "org-42"}, - } - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert ctx.tenant_id == "custom-tenant" - - async def test_tenant_id_missing_is_none(self): - """When no tenant claim is present, tenant_id should be None.""" - userinfo = {"sub": "user-123"} - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert ctx.tenant_id is None - - async def test_tenant_id_organization_without_id_is_none(self): - """An organization claim without an id should leave tenant_id as None.""" - userinfo = { - "sub": "user-123", - "organization": {"name": "no-id"}, - } - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert ctx.tenant_id is None - - async def test_tenant_id_organization_as_non_dict_is_ignored(self): - """A non-dict organization claim should not crash; tenant_id stays None.""" - userinfo = { - "sub": "user-123", - "organization": "not-a-dict", - } - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert ctx.tenant_id is None - - async def test_tenant_id_default_is_none(self): - """Direct construction without tenant_id should default to None.""" - ctx = UserContext(id="1", email="a@b.com", name="A") - assert ctx.tenant_id is None - - -# ── Auth dependencies (unit tests) ────────────────────────────────── - - -class TestGetCurrentUser: - async def test_raises_401_when_no_user(self): - """get_current_user raises 401 when request.state has no user.""" - from unittest.mock import MagicMock - - from auth.deps import get_current_user - from fastapi import HTTPException - - request = MagicMock() - # Simulate no user on request.state - del request.state.user - - with pytest.raises(HTTPException) as exc_info: - await get_current_user(request) - assert exc_info.value.status_code == 401 - - async def test_returns_user_when_present(self): - """get_current_user returns the user from request.state.""" - from unittest.mock import MagicMock - - from auth.deps import get_current_user - - user = UserContext(id="u1", email="u@test.com", name="User", roles=["user"]) - request = MagicMock() - request.state.user = user - - result = await get_current_user(request) - assert result.id == "u1" - - -class TestRequirePermission: - async def test_raises_403_when_missing_permission(self, app): - """The require_permission check function raises 403 when user lacks permissions.""" - from unittest.mock import MagicMock - - from auth.deps import require_permission - from fastapi import HTTPException - - # Get the inner check function from the Depends wrapper - dep = require_permission("products.delete") - check_fn = dep.dependency - - # Create a mock request with the app's permission registry - request = MagicMock() - request.app.state.perm_registry = app.state.perm_registry - - # User without the required permission - user = UserContext(id="u1", email="u@test.com", name="User", roles=["viewer"]) - - with pytest.raises(HTTPException) as exc_info: - await check_fn(request, user) - assert exc_info.value.status_code == 403 - - async def test_admin_bypasses_permission_check(self, app): - """The require_permission check allows admin users through.""" - from unittest.mock import MagicMock - - from auth.deps import require_permission - - dep = require_permission("products.delete") - check_fn = dep.dependency - - request = MagicMock() - request.app.state.perm_registry = app.state.perm_registry - - # Admin user should pass without raising - admin_user = UserContext(id="a1", email="admin@test.com", name="Admin", roles=["admin"]) - await check_fn(request, admin_user) # Should not raise - - -# ── AuthMiddleware ─────────────────────────────────────────────────── - - -class TestAuthMiddleware: - async def test_unauthenticated_request_redirects(self, client: httpx.AsyncClient): - """Accessing a protected page without a session should redirect to /auth/login.""" - resp = await client.get("/dashboard/", follow_redirects=False) - assert resp.status_code == 302 - assert "/auth/login" in resp.headers["location"] - - async def test_public_paths_not_redirected(self, client: httpx.AsyncClient): - """Health and auth paths should be accessible without authentication.""" - resp = await client.get("/health") - # In testing mode docs are disabled (404), but should NOT redirect to login - assert resp.status_code != 302 - - async def test_auth_me_unauthenticated(self, client: httpx.AsyncClient): - """/auth/me should return authenticated:false when no session.""" - resp = await client.get("/auth/me") - # In testing mode docs are disabled (404), but should NOT redirect to login - assert resp.status_code != 302 - data = resp.json() - assert data["authenticated"] is False - - async def test_authenticated_user_not_redirected(self, authenticated_client: httpx.AsyncClient): - """An authenticated user should not be redirected from protected API endpoints.""" - resp = await authenticated_client.get("/api/products/") - # In testing mode docs are disabled (404), but should NOT redirect to login - assert resp.status_code != 302 - - async def test_auth_me_authenticated(self, authenticated_client: httpx.AsyncClient): - """/auth/me should return user info when authenticated.""" - resp = await authenticated_client.get("/auth/me") - # In testing mode docs are disabled (404), but should NOT redirect to login - assert resp.status_code != 302 - data = resp.json() - assert data["authenticated"] is True - assert data["user"]["email"] == "test@example.com" - - -# ── UserContext Advanced ───────────────────────────────────────────── - - -class TestUserContextAdvanced: - async def test_from_keycloak_with_realm_access_roles(self): - userinfo = { - "sub": "u1", - "name": "Admin", - "email": "admin@test.com", - "realm_access": {"roles": ["admin", "user"]}, - } - ctx = UserContext.from_keycloak_userinfo(userinfo) - assert "admin" in ctx.roles - assert "user" in ctx.roles - - async def test_has_any_role_empty_user_roles(self): - ctx = UserContext(id="1", email="a@b.com", name="A", roles=[]) - assert ctx.has_any_role(["admin"]) is False - - async def test_has_any_role_empty_check_list(self): - ctx = UserContext(id="1", email="a@b.com", name="A", roles=["admin"]) - assert ctx.has_any_role([]) is False - - async def test_has_role_case_sensitive(self): - ctx = UserContext(id="1", email="a@b.com", name="A", roles=["Admin"]) - assert ctx.has_role("admin") is False - assert ctx.has_role("Admin") is True - - -# ── Auth Middleware Advanced ───────────────────────────────────────── - - -class TestAuthMiddlewareAdvanced: - async def test_landing_page_is_public(self, client: httpx.AsyncClient): - """The root / page should be accessible without auth.""" - resp = await client.get("/", follow_redirects=False) - # In testing mode docs are disabled (404), but should NOT redirect to login - assert resp.status_code != 302 - - async def test_health_endpoints_public(self, client: httpx.AsyncClient): - for path in ["/health", "/health/live", "/health/ready"]: - resp = await client.get(path) - # In testing mode docs are disabled (404), but should NOT redirect to login - assert resp.status_code != 302 - - async def test_api_docs_path_not_redirected(self, client: httpx.AsyncClient): - resp = await client.get("/api/docs", follow_redirects=False) - # In testing mode docs are disabled (404), but should NOT redirect to login - assert resp.status_code != 302 - - async def test_static_paths_public(self, client: httpx.AsyncClient): - # Static path won't have a file, but shouldn't redirect to login - resp = await client.get("/static/nonexistent.js", follow_redirects=False) - # 404 is fine, just not 302 to login - assert resp.status_code != 302 - - async def test_products_api_requires_auth(self, client: httpx.AsyncClient): - resp = await client.get("/api/products/", follow_redirects=False) - assert resp.status_code == 302 - assert "/auth/login" in resp.headers["location"] - - async def test_products_page_requires_auth(self, client: httpx.AsyncClient): - resp = await client.get("/products/", follow_redirects=False) - assert resp.status_code == 302 - - async def test_authenticated_can_access_products_api( - self, authenticated_client: httpx.AsyncClient - ): - resp = await authenticated_client.get("/api/products/") - # In testing mode docs are disabled (404), but should NOT redirect to login - assert resp.status_code != 302 - assert resp.json() == [] - - async def test_authenticated_can_access_dashboard( - self, authenticated_client: httpx.AsyncClient - ): - resp = await authenticated_client.get("/dashboard", follow_redirects=False) - # In testing mode docs are disabled (404), but should NOT redirect to login - assert resp.status_code != 302 - - -# ── Require Permission Advanced ────────────────────────────────────── - - -class TestRequirePermissionAdvanced: - async def test_multiple_permissions_any_match(self, app): - """User with any of the required permissions should pass.""" - from unittest.mock import MagicMock - - from auth.deps import require_permission - - dep = require_permission("products.view", "products.edit") - check_fn = dep.dependency - - request = MagicMock() - request.app.state.perm_registry = app.state.perm_registry - - # Admin passes regardless - admin = UserContext(id="a1", email="a@t.com", name="Admin", roles=["admin"]) - await check_fn(request, admin) # Should not raise - - async def test_non_admin_without_permission_fails(self, app): - from unittest.mock import MagicMock - - from auth.deps import require_permission - from fastapi import HTTPException - - dep = require_permission("products.delete") - check_fn = dep.dependency - - request = MagicMock() - request.app.state.perm_registry = app.state.perm_registry - - user = UserContext(id="u1", email="u@t.com", name="User", roles=["user"]) - with pytest.raises(HTTPException) as exc_info: - await check_fn(request, user) - assert exc_info.value.status_code == 403 - assert "products.delete" in str(exc_info.value.detail) - - -# ── Auth Module Registration ───────────────────────────────────────── - - -class TestAuthModuleRegistration: - async def test_auth_module_has_correct_meta(self): - from auth.module import AuthModule - - mod = AuthModule() - assert mod.meta.name == "Auth" - assert mod.meta.route_prefix == "/auth" - - async def test_auth_module_registers_menu_items(self): - from auth.module import AuthModule - - mod = AuthModule() - from simple_module_core.menu import MenuRegistry - - reg = MenuRegistry() - mod.register_menu_items(reg) - assert len(reg.all_items) == 1 - assert reg.all_items[0].label == "Logout" - assert reg.all_items[0].url == "/auth/logout" - - async def test_auth_logout_endpoint_exists(self, client: httpx.AsyncClient): - """The /auth/logout endpoint should exist (even if it redirects).""" - resp = await client.get("/auth/logout", follow_redirects=False) - assert resp.status_code == 302 diff --git a/modules/auth/tests/test_deps.py b/modules/auth/tests/test_deps.py new file mode 100644 index 00000000..729036df --- /dev/null +++ b/modules/auth/tests/test_deps.py @@ -0,0 +1,83 @@ +"""Tests for auth FastAPI dependencies (get_current_user, require_permission).""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest +from auth.contracts.schemas import UserContext +from auth.deps import get_current_user, require_permission +from fastapi import HTTPException + + +class TestGetCurrentUser: + async def test_raises_401_when_no_user(self): + """get_current_user raises 401 when request.state has no user.""" + request = MagicMock() + del request.state.user + + with pytest.raises(HTTPException) as exc_info: + await get_current_user(request) + assert exc_info.value.status_code == 401 + + async def test_returns_user_when_present(self): + """get_current_user returns the user from request.state.""" + user = UserContext(id="u1", email="u@test.com", name="User", roles=["user"]) + request = MagicMock() + request.state.user = user + + result = await get_current_user(request) + assert result.id == "u1" + + +class TestRequirePermission: + async def test_raises_403_when_missing_permission(self, app): + """The require_permission check raises 403 when user lacks permissions.""" + dep = require_permission("products.delete") + check_fn = dep.dependency + + request = MagicMock() + request.app.state.perm_registry = app.state.perm_registry + + user = UserContext(id="u1", email="u@test.com", name="User", roles=["viewer"]) + + with pytest.raises(HTTPException) as exc_info: + await check_fn(request, user) + assert exc_info.value.status_code == 403 + + async def test_admin_bypasses_permission_check(self, app): + """The require_permission check allows admin users through.""" + dep = require_permission("products.delete") + check_fn = dep.dependency + + request = MagicMock() + request.app.state.perm_registry = app.state.perm_registry + + admin_user = UserContext(id="a1", email="admin@test.com", name="Admin", roles=["admin"]) + await check_fn(request, admin_user) + + +class TestRequirePermissionAdvanced: + async def test_multiple_permissions_any_match(self, app): + """User with any of the required permissions should pass.""" + dep = require_permission("products.view", "products.edit") + check_fn = dep.dependency + + request = MagicMock() + request.app.state.perm_registry = app.state.perm_registry + + admin = UserContext(id="a1", email="a@t.com", name="Admin", roles=["admin"]) + await check_fn(request, admin) + + async def test_non_admin_without_permission_fails(self, app): + dep = require_permission("products.delete") + check_fn = dep.dependency + + request = MagicMock() + request.app.state.perm_registry = app.state.perm_registry + + user = UserContext(id="u1", email="u@t.com", name="User", roles=["user"]) + with pytest.raises(HTTPException) as exc_info: + await check_fn(request, user) + assert exc_info.value.status_code == 403 + assert "products.delete" in str(exc_info.value.detail) diff --git a/modules/auth/tests/test_middleware.py b/modules/auth/tests/test_middleware.py new file mode 100644 index 00000000..0e64a37f --- /dev/null +++ b/modules/auth/tests/test_middleware.py @@ -0,0 +1,80 @@ +"""Tests for AuthMiddleware: redirects, public paths, authenticated access.""" + +from __future__ import annotations + +import httpx + + +class TestAuthMiddleware: + async def test_unauthenticated_request_redirects(self, client: httpx.AsyncClient): + """Accessing a protected page without a session should redirect to /auth/login.""" + resp = await client.get("/dashboard/", follow_redirects=False) + assert resp.status_code == 302 + assert "/auth/login" in resp.headers["location"] + + async def test_public_paths_not_redirected(self, client: httpx.AsyncClient): + """Health and auth paths should be accessible without authentication.""" + resp = await client.get("/health") + assert resp.status_code != 302 + + async def test_auth_me_unauthenticated(self, client: httpx.AsyncClient): + """/auth/me should return authenticated:false when no session.""" + resp = await client.get("/auth/me") + assert resp.status_code != 302 + data = resp.json() + assert data["authenticated"] is False + + async def test_authenticated_user_not_redirected(self, authenticated_client: httpx.AsyncClient): + """An authenticated user should not be redirected from protected API endpoints.""" + resp = await authenticated_client.get("/api/products/") + assert resp.status_code != 302 + + async def test_auth_me_authenticated(self, authenticated_client: httpx.AsyncClient): + """/auth/me should return user info when authenticated.""" + resp = await authenticated_client.get("/auth/me") + assert resp.status_code != 302 + data = resp.json() + assert data["authenticated"] is True + assert data["user"]["email"] == "test@example.com" + + +class TestAuthMiddlewareAdvanced: + async def test_landing_page_is_public(self, client: httpx.AsyncClient): + """The root / page should be accessible without auth.""" + resp = await client.get("/", follow_redirects=False) + assert resp.status_code != 302 + + async def test_health_endpoints_public(self, client: httpx.AsyncClient): + for path in ["/health", "/health/live", "/health/ready"]: + resp = await client.get(path) + assert resp.status_code != 302 + + async def test_api_docs_path_not_redirected(self, client: httpx.AsyncClient): + resp = await client.get("/api/docs", follow_redirects=False) + assert resp.status_code != 302 + + async def test_static_paths_public(self, client: httpx.AsyncClient): + resp = await client.get("/static/nonexistent.js", follow_redirects=False) + assert resp.status_code != 302 + + async def test_products_api_requires_auth(self, client: httpx.AsyncClient): + resp = await client.get("/api/products/", follow_redirects=False) + assert resp.status_code == 302 + assert "/auth/login" in resp.headers["location"] + + async def test_products_page_requires_auth(self, client: httpx.AsyncClient): + resp = await client.get("/products/", follow_redirects=False) + assert resp.status_code == 302 + + async def test_authenticated_can_access_products_api( + self, authenticated_client: httpx.AsyncClient + ): + resp = await authenticated_client.get("/api/products/") + assert resp.status_code != 302 + assert resp.json() == [] + + async def test_authenticated_can_access_dashboard( + self, authenticated_client: httpx.AsyncClient + ): + resp = await authenticated_client.get("/dashboard", follow_redirects=False) + assert resp.status_code != 302 diff --git a/modules/auth/tests/test_module.py b/modules/auth/tests/test_module.py new file mode 100644 index 00000000..d5cbdd14 --- /dev/null +++ b/modules/auth/tests/test_module.py @@ -0,0 +1,27 @@ +"""Tests for AuthModule registration (metadata, menu, logout endpoint).""" + +from __future__ import annotations + +import httpx +from auth.module import AuthModule +from simple_module_core.menu import MenuRegistry + + +class TestAuthModuleRegistration: + async def test_auth_module_has_correct_meta(self): + mod = AuthModule() + assert mod.meta.name == "Auth" + assert mod.meta.route_prefix == "/auth" + + async def test_auth_module_registers_menu_items(self): + mod = AuthModule() + reg = MenuRegistry() + mod.register_menu_items(reg) + assert len(reg.all_items) == 1 + assert reg.all_items[0].label == "Logout" + assert reg.all_items[0].url == "/auth/logout" + + async def test_auth_logout_endpoint_exists(self, client: httpx.AsyncClient): + """The /auth/logout endpoint should exist (even if it redirects).""" + resp = await client.get("/auth/logout", follow_redirects=False) + assert resp.status_code == 302 diff --git a/modules/auth/tests/test_user_context.py b/modules/auth/tests/test_user_context.py new file mode 100644 index 00000000..67292e73 --- /dev/null +++ b/modules/auth/tests/test_user_context.py @@ -0,0 +1,132 @@ +"""Tests for the UserContext value object (construction, roles, tenant).""" + +from __future__ import annotations + +from auth.contracts.schemas import UserContext + + +class TestUserContext: + async def test_from_keycloak_userinfo_basic(self): + userinfo = { + "sub": "user-123", + "email": "alice@example.com", + "name": "Alice Smith", + "realm_access": {"roles": ["user", "editor"]}, + } + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert ctx.id == "user-123" + assert ctx.email == "alice@example.com" + assert ctx.name == "Alice Smith" + assert ctx.roles == ["user", "editor"] + + async def test_from_keycloak_userinfo_fallback_username(self): + userinfo = { + "sub": "user-456", + "email": "bob@example.com", + "preferred_username": "bob", + } + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert ctx.name == "bob" + + async def test_from_keycloak_userinfo_missing_fields(self): + userinfo = {} + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert ctx.id == "" + assert ctx.email == "" + assert ctx.name == "" + assert ctx.roles == [] + + async def test_has_role(self): + ctx = UserContext(id="1", email="a@b.com", name="A", roles=["admin", "user"]) + assert ctx.has_role("admin") is True + assert ctx.has_role("superadmin") is False + + async def test_has_any_role(self): + ctx = UserContext(id="1", email="a@b.com", name="A", roles=["editor"]) + assert ctx.has_any_role(["admin", "editor"]) is True + assert ctx.has_any_role(["admin", "superadmin"]) is False + + +class TestUserContextTenantId: + async def test_tenant_id_from_custom_claim(self): + """UserContext should read a custom ``tenant_id`` claim from userinfo.""" + userinfo = { + "sub": "user-123", + "tenant_id": "acme-corp", + } + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert ctx.tenant_id == "acme-corp" + + async def test_tenant_id_from_organization_claim(self): + """UserContext should fall back to Keycloak's organization.id claim.""" + userinfo = { + "sub": "user-123", + "organization": {"id": "org-42", "name": "Acme"}, + } + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert ctx.tenant_id == "org-42" + + async def test_tenant_id_custom_claim_takes_precedence(self): + """Custom tenant_id claim should take precedence over organization.id.""" + userinfo = { + "sub": "user-123", + "tenant_id": "custom-tenant", + "organization": {"id": "org-42"}, + } + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert ctx.tenant_id == "custom-tenant" + + async def test_tenant_id_missing_is_none(self): + """When no tenant claim is present, tenant_id should be None.""" + userinfo = {"sub": "user-123"} + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert ctx.tenant_id is None + + async def test_tenant_id_organization_without_id_is_none(self): + """An organization claim without an id should leave tenant_id as None.""" + userinfo = { + "sub": "user-123", + "organization": {"name": "no-id"}, + } + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert ctx.tenant_id is None + + async def test_tenant_id_organization_as_non_dict_is_ignored(self): + """A non-dict organization claim should not crash; tenant_id stays None.""" + userinfo = { + "sub": "user-123", + "organization": "not-a-dict", + } + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert ctx.tenant_id is None + + async def test_tenant_id_default_is_none(self): + """Direct construction without tenant_id should default to None.""" + ctx = UserContext(id="1", email="a@b.com", name="A") + assert ctx.tenant_id is None + + +class TestUserContextAdvanced: + async def test_from_keycloak_with_realm_access_roles(self): + userinfo = { + "sub": "u1", + "name": "Admin", + "email": "admin@test.com", + "realm_access": {"roles": ["admin", "user"]}, + } + ctx = UserContext.from_keycloak_userinfo(userinfo) + assert "admin" in ctx.roles + assert "user" in ctx.roles + + async def test_has_any_role_empty_user_roles(self): + ctx = UserContext(id="1", email="a@b.com", name="A", roles=[]) + assert ctx.has_any_role(["admin"]) is False + + async def test_has_any_role_empty_check_list(self): + ctx = UserContext(id="1", email="a@b.com", name="A", roles=["admin"]) + assert ctx.has_any_role([]) is False + + async def test_has_role_case_sensitive(self): + ctx = UserContext(id="1", email="a@b.com", name="A", roles=["Admin"]) + assert ctx.has_role("admin") is False + assert ctx.has_role("Admin") is True diff --git a/modules/products/products/pages/Browse.tsx b/modules/products/products/pages/Browse.tsx index 19b15a71..10c85427 100644 --- a/modules/products/products/pages/Browse.tsx +++ b/modules/products/products/pages/Browse.tsx @@ -16,15 +16,6 @@ import { Button } from '@ui/components/ui/button'; import { Card } from '@ui/components/ui/card'; import { Empty, EmptyDescription, EmptyMedia, EmptyTitle } from '@ui/components/ui/empty'; import { Input } from '@ui/components/ui/input'; -import { - Pagination, - PaginationContent, - PaginationEllipsis, - PaginationItem, - PaginationLink, - PaginationNext, - PaginationPrevious, -} from '@ui/components/ui/pagination'; import { Table, TableBody, @@ -38,6 +29,7 @@ import { AuthenticatedLayout } from '@ui/layouts/AuthenticatedLayout'; import { Package, Pencil, Plus, Search, Trash2 } from 'lucide-react'; import { useEffect, useMemo, useState } from 'react'; import { toast } from 'sonner'; +import { ProductsPagination } from './components/ProductsPagination'; interface Product { id: number; @@ -248,61 +240,11 @@ function Browse() { - {/* Pagination */} - {totalPages > 1 && ( -
- - - - navigate(pagination.page - 1)} - className={ - pagination.page <= 1 ? 'pointer-events-none opacity-50' : 'cursor-pointer' - } - /> - - {Array.from({ length: totalPages }, (_, i) => i + 1) - .filter((p) => { - if (p === 1 || p === totalPages) return true; - if (Math.abs(p - pagination.page) <= 1) return true; - return false; - }) - .reduce<(number | 'ellipsis')[]>((acc, p, i, arr) => { - if (i > 0 && p - (arr[i - 1] as number) > 1) acc.push('ellipsis'); - acc.push(p); - return acc; - }, []) - .map((item, i) => - item === 'ellipsis' ? ( - - - - ) : ( - - navigate(item)} - className="cursor-pointer" - > - {item} - - - ), - )} - - navigate(pagination.page + 1)} - className={ - pagination.page >= totalPages - ? 'pointer-events-none opacity-50' - : 'cursor-pointer' - } - /> - - - -
- )} + navigate(p)} + /> ); } diff --git a/modules/products/products/pages/components/ProductsPagination.tsx b/modules/products/products/pages/components/ProductsPagination.tsx new file mode 100644 index 00000000..622fb960 --- /dev/null +++ b/modules/products/products/pages/components/ProductsPagination.tsx @@ -0,0 +1,70 @@ +import { + Pagination, + PaginationContent, + PaginationEllipsis, + PaginationItem, + PaginationLink, + PaginationNext, + PaginationPrevious, +} from '@ui/components/ui/pagination'; +import { useMemo } from 'react'; + +interface Props { + page: number; + totalPages: number; + onNavigate: (page: number) => void; +} + +export function ProductsPagination({ page, totalPages, onNavigate }: Props) { + const items = useMemo( + () => + Array.from({ length: totalPages }, (_, i) => i + 1) + .filter((p) => p === 1 || p === totalPages || Math.abs(p - page) <= 1) + .reduce<(number | 'ellipsis')[]>((acc, p, i, arr) => { + if (i > 0 && p - (arr[i - 1] as number) > 1) acc.push('ellipsis'); + acc.push(p); + return acc; + }, []), + [page, totalPages], + ); + + if (totalPages <= 1) return null; + + return ( +
+ + + + onNavigate(page - 1)} + className={page <= 1 ? 'pointer-events-none opacity-50' : 'cursor-pointer'} + /> + + {items.map((item, i) => + item === 'ellipsis' ? ( + + + + ) : ( + + onNavigate(item)} + className="cursor-pointer" + > + {item} + + + ), + )} + + onNavigate(page + 1)} + className={page >= totalPages ? 'pointer-events-none opacity-50' : 'cursor-pointer'} + /> + + + +
+ ); +} diff --git a/scripts/_templates_contracts.py b/scripts/_templates_contracts.py new file mode 100644 index 00000000..6bc3753c --- /dev/null +++ b/scripts/_templates_contracts.py @@ -0,0 +1,93 @@ +"""Template generators for the contracts/ sub-package (schemas + service protocol).""" + +from __future__ import annotations + +from _templates_py import ScaffoldContext + + +def contracts_init(ctx: ScaffoldContext) -> str: + return f'''\ + """{ctx.class_name} contracts — public interface for other modules.""" + + from {ctx.pkg}.contracts.schemas import ( + {ctx.singular_class}Create, + {ctx.singular_class}Out, + {ctx.singular_class}Update, + ) + from {ctx.pkg}.contracts.service import I{ctx.singular_class}Service + + __all__ = [ + "{ctx.singular_class}Create", + "{ctx.singular_class}Out", + "{ctx.singular_class}Update", + "I{ctx.singular_class}Service", + ] + ''' + + +def schemas_py(ctx: ScaffoldContext) -> str: + return f'''\ + """Pydantic DTOs for the {ctx.class_name} module.""" + + from __future__ import annotations + + from datetime import datetime + + from pydantic import BaseModel, ConfigDict, Field + + + class {ctx.singular_class}Out(BaseModel): + """{ctx.singular_class} data returned by the API.""" + + model_config = ConfigDict(from_attributes=True) + + id: int + name: str + description: str | None = None + is_active: bool + created_at: datetime | None = None + updated_at: datetime | None = None + + + class {ctx.singular_class}Create(BaseModel): + """Data required to create a new {ctx.singular}.""" + + name: str = Field(min_length=1, max_length=200) + description: str | None = None + + + class {ctx.singular_class}Update(BaseModel): + """Data to update an existing {ctx.singular}. All fields optional.""" + + name: str | None = Field(default=None, min_length=1, max_length=200) + description: str | None = None + is_active: bool | None = None + ''' + + +def contracts_service(ctx: ScaffoldContext) -> str: + return f'''\ + """{ctx.singular_class} service protocol — the public contract other modules depend on.""" + + from __future__ import annotations + + from typing import Protocol + + from {ctx.pkg}.contracts.schemas import ( + {ctx.singular_class}Create, + {ctx.singular_class}Out, + {ctx.singular_class}Update, + ) + + + class I{ctx.singular_class}Service(Protocol): + """Interface for {ctx.singular} operations.""" + + async def get_all(self) -> list[{ctx.singular_class}Out]: ... + async def get_by_id(self, {ctx.singular}_id: int) -> {ctx.singular_class}Out | None: ... + async def create(self, data: {ctx.singular_class}Create) -> {ctx.singular_class}Out: ... + async def update( + self, {ctx.singular}_id: int, data: {ctx.singular_class}Update + ) -> {ctx.singular_class}Out | None: ... + async def delete(self, {ctx.singular}_id: int) -> bool: ... + ''' diff --git a/scripts/_templates_endpoints.py b/scripts/_templates_endpoints.py new file mode 100644 index 00000000..0ee69105 --- /dev/null +++ b/scripts/_templates_endpoints.py @@ -0,0 +1,125 @@ +"""Template generators for FastAPI endpoint files (api.py, views.py).""" + +from __future__ import annotations + +from _templates_py import ScaffoldContext + + +def api_py(ctx: ScaffoldContext) -> str: + return f'''\ + """REST API endpoints for {ctx.class_name}.""" + + from __future__ import annotations + + from fastapi import APIRouter, Depends, HTTPException + + from {ctx.pkg}.contracts.schemas import ( + {ctx.singular_class}Create, + {ctx.singular_class}Out, + {ctx.singular_class}Update, + ) + from {ctx.pkg}.deps import get_{ctx.singular}_service + from {ctx.pkg}.service import {ctx.singular_class}Service + + router = APIRouter() + + + @router.get("/", response_model=list[{ctx.singular_class}Out]) + async def list_{ctx.name}( + service: {ctx.singular_class}Service = Depends(get_{ctx.singular}_service), + ) -> list[{ctx.singular_class}Out]: + return await service.get_all() + + + @router.get("/{{{ctx.singular}_id}}", response_model={ctx.singular_class}Out) + async def get_{ctx.singular}( + {ctx.singular}_id: int, + service: {ctx.singular_class}Service = Depends(get_{ctx.singular}_service), + ) -> {ctx.singular_class}Out: + result = await service.get_by_id({ctx.singular}_id) + if result is None: + raise HTTPException(status_code=404, detail="{ctx.singular_class} not found") + return result + + + @router.post("/", response_model={ctx.singular_class}Out, status_code=201) + async def create_{ctx.singular}( + data: {ctx.singular_class}Create, + service: {ctx.singular_class}Service = Depends(get_{ctx.singular}_service), + ) -> {ctx.singular_class}Out: + return await service.create(data) + + + @router.put("/{{{ctx.singular}_id}}", response_model={ctx.singular_class}Out) + async def update_{ctx.singular}( + {ctx.singular}_id: int, + data: {ctx.singular_class}Update, + service: {ctx.singular_class}Service = Depends(get_{ctx.singular}_service), + ) -> {ctx.singular_class}Out: + result = await service.update({ctx.singular}_id, data) + if result is None: + raise HTTPException(status_code=404, detail="{ctx.singular_class} not found") + return result + + + @router.delete("/{{{ctx.singular}_id}}", status_code=204) + async def delete_{ctx.singular}( + {ctx.singular}_id: int, + service: {ctx.singular_class}Service = Depends(get_{ctx.singular}_service), + ) -> None: + deleted = await service.delete({ctx.singular}_id) + if not deleted: + raise HTTPException(status_code=404, detail="{ctx.singular_class} not found") + ''' + + +def views_py(ctx: ScaffoldContext) -> str: + return f'''\ + """Inertia view endpoints for {ctx.class_name}.""" + + from __future__ import annotations + + from fastapi import APIRouter, Depends + from inertia import InertiaResponse + from simple_module_hosting.inertia_deps import InertiaDep + + from {ctx.pkg}.deps import get_{ctx.singular}_service + from {ctx.pkg}.service import {ctx.singular_class}Service + + router = APIRouter() + + + @router.get("/", response_model=None) + async def browse( + inertia: InertiaDep, + service: {ctx.singular_class}Service = Depends(get_{ctx.singular}_service), + ) -> InertiaResponse: + items = await service.get_all() + return await inertia.render( + "{ctx.class_name}/Browse", + {{"{ctx.name}": [item.model_dump(mode="json") for item in items]}}, + ) + + + @router.get("/create", response_model=None) + async def create_view(inertia: InertiaDep) -> InertiaResponse: + return await inertia.render("{ctx.class_name}/Create") + + + @router.get("/{{{ctx.singular}_id}}/edit", response_model=None) + async def edit_view( + {ctx.singular}_id: int, + inertia: InertiaDep, + service: {ctx.singular_class}Service = Depends(get_{ctx.singular}_service), + ) -> InertiaResponse: + item = await service.get_by_id({ctx.singular}_id) + if item is None: + return await inertia.render( + "{ctx.class_name}/Browse", + {{"error": "{ctx.singular_class} not found"}}, + ) + return await inertia.render( + "{ctx.class_name}/Edit", + {{"{ctx.singular}": item.model_dump(mode="json")}}, + ) + ''' diff --git a/scripts/_templates_py.py b/scripts/_templates_py.py new file mode 100644 index 00000000..eaf5cca6 --- /dev/null +++ b/scripts/_templates_py.py @@ -0,0 +1,219 @@ +"""Template generators for Python files scaffolded by new_module. + +Each function returns the file content for a single generated Python file. +Context is passed via the ``ScaffoldContext`` dataclass so the signatures +stay compact. +""" + +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class ScaffoldContext: + """Naming variables for all generated templates.""" + + name: str # snake_case plural, e.g. "orders" + class_name: str # PascalCase plural, e.g. "Orders" + singular: str # snake_case singular, e.g. "order" + singular_class: str # PascalCase singular, e.g. "Order" + pkg: str # Python package name (same as name today) + + +def pyproject_toml(ctx: ScaffoldContext) -> str: + return f"""\ + [project] + name = "{ctx.pkg.replace("_", "-")}" + version = "0.1.0" + description = "The {ctx.class_name} module" + authors = [] + requires-python = ">=3.12" + dependencies = [ + "simple-module-core", + "simple-module-db", + "simple-module-hosting", + ] + + [project.entry-points.simple_module] + {ctx.name} = "{ctx.pkg}.module:{ctx.class_name}Module" + + [build-system] + requires = ["hatchling"] + build-backend = "hatchling.build" + + [tool.uv.sources] + simple-module-core = {{ workspace = true }} + simple-module-db = {{ workspace = true }} + simple-module-hosting = {{ workspace = true }} + """ + + +def package_init(ctx: ScaffoldContext) -> str: + return f'''\ + """{ctx.class_name} module.""" + ''' + + +def module_py(ctx: ScaffoldContext) -> str: + return f'''\ + """{ctx.class_name} module definition.""" + + from __future__ import annotations + + from fastapi import APIRouter + from simple_module_core.menu import MenuItem, MenuRegistry, MenuSection + from simple_module_core.module import ModuleBase, ModuleMeta + from simple_module_core.permissions import PermissionRegistry + + + class {ctx.class_name}Module(ModuleBase): + meta = ModuleMeta( + name="{ctx.class_name}", + route_prefix="/api/{ctx.name}", + view_prefix="/{ctx.name}", + ) + + def register_routes(self, api_router: APIRouter, view_router: APIRouter) -> None: + from {ctx.pkg}.endpoints.api import router as api + from {ctx.pkg}.endpoints.views import router as views + + api_router.include_router(api) + view_router.include_router(views) + + def register_menu_items(self, registry: MenuRegistry) -> None: + registry.add( + MenuItem( + label="{ctx.class_name}", + url="/{ctx.name}", + icon="box", + order=30, + section=MenuSection.SIDEBAR, + ) + ) + + def register_permissions(self, registry: PermissionRegistry) -> None: + registry.add_group( + "{ctx.class_name}", + [ + "{ctx.name}.view", + "{ctx.name}.create", + "{ctx.name}.edit", + "{ctx.name}.delete", + ], + ) + ''' + + +def models_py(ctx: ScaffoldContext) -> str: + return f'''\ + """SQLAlchemy models for the {ctx.class_name} module.""" + + from __future__ import annotations + + from simple_module_db.base import create_module_base + from simple_module_db.mixins import AuditMixin + from sqlalchemy import String + from sqlalchemy.orm import Mapped, mapped_column + + # Provider is auto-detected from SM_DATABASE_URL (falls back to SQLite). + # On PostgreSQL this gives the module its own `{ctx.name}` schema; on SQLite + # all modules share one schema, so __tablename__ is prefixed for isolation. + Base = create_module_base("{ctx.name}") + + + class {ctx.singular_class}(Base, AuditMixin): # ty: ignore[unsupported-base] + """A {ctx.singular} entity.""" + + __tablename__ = "{ctx.name}_{ctx.singular}" + + id: Mapped[int] = mapped_column(primary_key=True, autoincrement=True) + name: Mapped[str] = mapped_column(String(200)) + description: Mapped[str | None] = mapped_column(String(2000), default=None) + is_active: Mapped[bool] = mapped_column(default=True) + ''' + + +def service_py(ctx: ScaffoldContext) -> str: + return f'''\ + """{ctx.singular_class} service implementation.""" + + from __future__ import annotations + + from sqlalchemy import select + from sqlalchemy.ext.asyncio import AsyncSession + + from {ctx.pkg}.contracts.schemas import ( + {ctx.singular_class}Create, + {ctx.singular_class}Out, + {ctx.singular_class}Update, + ) + from {ctx.pkg}.models import {ctx.singular_class} + + + class {ctx.singular_class}Service: + """CRUD operations for {ctx.name}.""" + + def __init__(self, db: AsyncSession) -> None: + self.db = db + + async def get_all(self) -> list[{ctx.singular_class}Out]: + result = await self.db.execute( + select({ctx.singular_class}) + .where({ctx.singular_class}.is_active.is_(True)) + .order_by({ctx.singular_class}.id) + ) + return [{ctx.singular_class}Out.model_validate(row) for row in result.scalars()] + + async def get_by_id(self, {ctx.singular}_id: int) -> {ctx.singular_class}Out | None: + entity = await self.db.get({ctx.singular_class}, {ctx.singular}_id) + if entity is None: + return None + return {ctx.singular_class}Out.model_validate(entity) + + async def create(self, data: {ctx.singular_class}Create) -> {ctx.singular_class}Out: + entity = {ctx.singular_class}(**data.model_dump()) + self.db.add(entity) + await self.db.flush() + await self.db.refresh(entity) + return {ctx.singular_class}Out.model_validate(entity) + + async def update( + self, {ctx.singular}_id: int, data: {ctx.singular_class}Update + ) -> {ctx.singular_class}Out | None: + entity = await self.db.get({ctx.singular_class}, {ctx.singular}_id) + if entity is None: + return None + for field, value in data.model_dump(exclude_unset=True).items(): + setattr(entity, field, value) + await self.db.flush() + await self.db.refresh(entity) + return {ctx.singular_class}Out.model_validate(entity) + + async def delete(self, {ctx.singular}_id: int) -> bool: + entity = await self.db.get({ctx.singular_class}, {ctx.singular}_id) + if entity is None: + return False + await self.db.delete(entity) + return True + ''' + + +def deps_py(ctx: ScaffoldContext) -> str: + return f'''\ + """FastAPI dependencies for the {ctx.class_name} module.""" + + from __future__ import annotations + + from fastapi import Depends + from simple_module_db.deps import get_db + from sqlalchemy.ext.asyncio import AsyncSession + + from {ctx.pkg}.service import {ctx.singular_class}Service + + + async def get_{ctx.singular}_service( + db: AsyncSession = Depends(get_db), + ) -> {ctx.singular_class}Service: + return {ctx.singular_class}Service(db) + ''' diff --git a/scripts/_templates_tests.py b/scripts/_templates_tests.py new file mode 100644 index 00000000..8f00f9b1 --- /dev/null +++ b/scripts/_templates_tests.py @@ -0,0 +1,181 @@ +"""Template generator for the scaffolded module's test file.""" + +from __future__ import annotations + +from _templates_py import ScaffoldContext + + +def test_module_py(ctx: ScaffoldContext) -> str: + return f'''\ + """Tests for the {ctx.class_name} module: service CRUD, API endpoints, schema validation.""" + + from __future__ import annotations + + import httpx + import pytest + from pydantic import ValidationError + from {ctx.pkg}.contracts.schemas import ( + {ctx.singular_class}Create, + {ctx.singular_class}Update, + ) + from {ctx.pkg}.service import {ctx.singular_class}Service + from sqlalchemy.ext.asyncio import AsyncSession + + # ── Schema validation ──────────────────────────────────────────────── + + + class Test{ctx.singular_class}Schemas: + async def test_create_valid(self): + data = {ctx.singular_class}Create(name="Test {ctx.singular_class}") + assert data.name == "Test {ctx.singular_class}" + assert data.description is None + + async def test_create_empty_name_rejected(self): + with pytest.raises(ValidationError): + {ctx.singular_class}Create(name="") + + async def test_update_all_optional(self): + data = {ctx.singular_class}Update() + assert data.name is None + assert data.is_active is None + + + # ── {ctx.singular_class}Service CRUD ────────────────────────────────────────────── + + + class Test{ctx.singular_class}Service: + async def test_create(self, db_session: AsyncSession): + svc = {ctx.singular_class}Service(db_session) + item = await svc.create({ctx.singular_class}Create(name="Test")) + assert item.id is not None + assert item.name == "Test" + assert item.is_active is True + + async def test_get_all(self, db_session: AsyncSession): + svc = {ctx.singular_class}Service(db_session) + await svc.create({ctx.singular_class}Create(name="A")) + await svc.create({ctx.singular_class}Create(name="B")) + items = await svc.get_all() + assert len(items) == 2 + + async def test_get_by_id(self, db_session: AsyncSession): + svc = {ctx.singular_class}Service(db_session) + created = await svc.create({ctx.singular_class}Create(name="X")) + found = await svc.get_by_id(created.id) + assert found is not None + assert found.name == "X" + + async def test_get_by_id_not_found(self, db_session: AsyncSession): + svc = {ctx.singular_class}Service(db_session) + found = await svc.get_by_id(999) + assert found is None + + async def test_update(self, db_session: AsyncSession): + svc = {ctx.singular_class}Service(db_session) + created = await svc.create({ctx.singular_class}Create(name="Old")) + updated = await svc.update(created.id, {ctx.singular_class}Update(name="New")) + assert updated is not None + assert updated.name == "New" + + async def test_update_not_found(self, db_session: AsyncSession): + svc = {ctx.singular_class}Service(db_session) + result = await svc.update(999, {ctx.singular_class}Update(name="Ghost")) + assert result is None + + async def test_delete(self, db_session: AsyncSession): + svc = {ctx.singular_class}Service(db_session) + created = await svc.create({ctx.singular_class}Create(name="Doomed")) + deleted = await svc.delete(created.id) + assert deleted is True + + async def test_delete_not_found(self, db_session: AsyncSession): + svc = {ctx.singular_class}Service(db_session) + deleted = await svc.delete(999) + assert deleted is False + + + # ── API endpoints ─────────────────────────────────────────────── + + + class Test{ctx.class_name}API: + async def test_list_empty(self, authenticated_client: httpx.AsyncClient): + resp = await authenticated_client.get("/api/{ctx.name}/") + assert resp.status_code == 200 + assert resp.json() == [] + + async def test_create(self, authenticated_client: httpx.AsyncClient): + resp = await authenticated_client.post( + "/api/{ctx.name}/", + json={{"name": "Test {ctx.singular_class}"}}, + ) + assert resp.status_code == 201 + data = resp.json() + assert data["name"] == "Test {ctx.singular_class}" + assert data["id"] is not None + + async def test_get_by_id(self, authenticated_client: httpx.AsyncClient): + create_resp = await authenticated_client.post( + "/api/{ctx.name}/", + json={{"name": "Findable"}}, + ) + item_id = create_resp.json()["id"] + resp = await authenticated_client.get(f"/api/{ctx.name}/{{item_id}}") + assert resp.status_code == 200 + assert resp.json()["name"] == "Findable" + + async def test_get_not_found(self, authenticated_client: httpx.AsyncClient): + resp = await authenticated_client.get("/api/{ctx.name}/99999") + assert resp.status_code == 404 + + async def test_update(self, authenticated_client: httpx.AsyncClient): + create_resp = await authenticated_client.post( + "/api/{ctx.name}/", + json={{"name": "Original"}}, + ) + item_id = create_resp.json()["id"] + resp = await authenticated_client.put( + f"/api/{ctx.name}/{{item_id}}", + json={{"name": "Updated"}}, + ) + assert resp.status_code == 200 + assert resp.json()["name"] == "Updated" + + async def test_delete(self, authenticated_client: httpx.AsyncClient): + create_resp = await authenticated_client.post( + "/api/{ctx.name}/", + json={{"name": "Deletable"}}, + ) + item_id = create_resp.json()["id"] + resp = await authenticated_client.delete(f"/api/{ctx.name}/{{item_id}}") + assert resp.status_code == 204 + + async def test_delete_not_found(self, authenticated_client: httpx.AsyncClient): + resp = await authenticated_client.delete("/api/{ctx.name}/99999") + assert resp.status_code == 404 + + async def test_create_invalid_data(self, authenticated_client: httpx.AsyncClient): + resp = await authenticated_client.post( + "/api/{ctx.name}/", + json={{"name": ""}}, + ) + assert resp.status_code == 422 + + + # ── Module lifecycle ──────────────────────────────────────────────── + + + class Test{ctx.class_name}ModuleLifecycle: + async def test_on_startup_does_not_call_create_all(self): + """on_startup should not create tables — Alembic manages schema.""" + from unittest.mock import AsyncMock, MagicMock + + from {ctx.pkg}.module import {ctx.class_name}Module + + mod = {ctx.class_name}Module() + mock_app = MagicMock() + mock_app.state.db.engine = AsyncMock() + + await mod.on_startup(mock_app) + + mock_app.state.db.engine.begin.assert_not_called() + ''' diff --git a/scripts/_templates_tsx.py b/scripts/_templates_tsx.py new file mode 100644 index 00000000..9877b76e --- /dev/null +++ b/scripts/_templates_tsx.py @@ -0,0 +1,121 @@ +"""Template generators for the three React/Inertia page TSX files.""" + +from __future__ import annotations + +from _templates_py import ScaffoldContext + + +def browse_tsx(ctx: ScaffoldContext) -> str: + return f"""\ + type {ctx.singular_class} = {{ + id: number; + name: string; + description: string | null; + is_active: boolean; + }}; + + type Props = {{ {ctx.name}: {ctx.singular_class}[] }}; + + export default function Browse({{ {ctx.name} }}: Props) {{ + return ( +
+
+

{ctx.class_name}

+ + New {ctx.singular_class} + +
+
    + {{{ctx.name}.map(({ctx.singular}) => ( +
  • + {{{ctx.singular}.name}} + Edit +
  • + ))}} +
+
+ ); + }} + """ + + +def create_tsx(ctx: ScaffoldContext) -> str: + return f"""\ + export default function Create() {{ + return ( +
+

New {ctx.singular_class}

+
+ +