Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 22 additions & 10 deletions conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,9 @@
from simple_module_core.discovery import discover_modules
from simple_module_db.base import all_module_bases
from simple_module_db.session import DatabaseState, init_db
from simple_module_hosting.csrf import SESSION_CSRF_TOKEN_KEY
from simple_module_hosting.settings import Settings
from simple_module_testing import forge_session_cookie
from sqlalchemy.ext.asyncio import (
AsyncEngine,
AsyncSession,
Expand Down Expand Up @@ -135,21 +137,30 @@ async def app(settings: Settings):
await ctx.__aexit__(None, None, None)


_TEST_CSRF_TOKEN = "test-csrf-token"


@pytest.fixture
async def client(app) -> AsyncGenerator[httpx.AsyncClient, None]:
"""Unauthenticated async HTTP client."""
"""Unauthenticated async HTTP client — still CSRF-equipped so POST/PATCH/DELETE
requests from anonymous test flows (login, accept-invite, register) pass
validation without first making a GET to mint a token."""
signed = forge_session_cookie(
app.state.settings.secret_key, {SESSION_CSRF_TOKEN_KEY: _TEST_CSRF_TOKEN}
)
transport = httpx.ASGITransport(app=app)
async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as c:
async with httpx.AsyncClient(
transport=transport,
base_url="http://testserver",
cookies={"session": signed},
headers={"X-CSRF-Token": _TEST_CSRF_TOKEN},
) as c:
yield c


@pytest.fixture
async def authenticated_client(app) -> AsyncGenerator[httpx.AsyncClient, None]:
"""HTTPX client with a signed session cookie carrying a seeded admin user's id."""
import json
from base64 import b64encode

from itsdangerous import TimestampSigner
from users.bootstrap import create_admin

async with app.state.db.session_factory() as session:
Expand All @@ -161,15 +172,16 @@ async def authenticated_client(app) -> AsyncGenerator[httpx.AsyncClient, None]:
)
user_id = str(result.user.id)

session_data = {"user_id": user_id}
data = b64encode(json.dumps(session_data).encode())
signer = TimestampSigner(str(app.state.settings.secret_key))
signed = signer.sign(data).decode("utf-8")
signed = forge_session_cookie(
app.state.settings.secret_key,
{"user_id": user_id, SESSION_CSRF_TOKEN_KEY: _TEST_CSRF_TOKEN},
)

transport = httpx.ASGITransport(app=app)
async with httpx.AsyncClient(
transport=transport,
base_url="http://testserver",
cookies={"session": signed},
headers={"X-CSRF-Token": _TEST_CSRF_TOKEN},
) as c:
yield c
15 changes: 15 additions & 0 deletions framework/core/simple_module_core/environments.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
"""Shared environment-classification constants.

Both the host and module settings validators need to know which ``SM_ENVIRONMENT``
values are "non-prod" — anything outside this set is treated as production
and subject to stricter defaults (e.g. placeholder-secret rejection).

Duplicating this constant in each settings module would mean an operator who
adds ``"staging"`` to one list but forgets the other gets inconsistent
validation. Lives in ``simple_module_core`` because both the host package and
module settings already depend on it.
"""

from __future__ import annotations

NON_PROD_ENVIRONMENTS: frozenset[str] = frozenset({"development", "testing"})
47 changes: 43 additions & 4 deletions framework/core/simple_module_core/i18n.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,10 @@

import json
import logging
from collections.abc import Mapping
from functools import lru_cache
from pathlib import Path
from types import MappingProxyType
from typing import Any

from babel import Locale
Expand Down Expand Up @@ -78,6 +80,13 @@ def __init__(self, default_locale: str, supported_locales: list[str]) -> None:
self.supported_locales = list(supported_locales)
self._sources: list[tuple[str, Path]] = []
self._messages: dict[str, dict[str, str]] = {}
# Immutable views into ``_messages`` — handed out by ``messages()`` to
# avoid a per-call dict copy. Rebuilt whenever ``load()`` runs.
self._message_views: dict[str, MappingProxyType[str, str]] = {}
self._available_locales: tuple[str, ...] = ()
self._available_locales_list: list[str] = []
self._empty_view: MappingProxyType[str, str] = MappingProxyType({})
self._loaded = False

def add_source(self, namespace: str, locale_dir: Path) -> None:
"""Queue a module's locale directory for loading under a namespace."""
Expand Down Expand Up @@ -110,13 +119,43 @@ def load(self) -> None:
flat = flatten_messages(raw, prefix=namespace)
self._messages[locale].update(flat)

# Cache the derived views now that loading is complete. Downstream
# (middleware, translator, switcher) reads these on every request.
self._message_views = {
locale: MappingProxyType(msgs) for locale, msgs in self._messages.items()
}
self._available_locales = tuple(locale for locale, msgs in self._messages.items() if msgs)
self._available_locales_list = list(self._available_locales)
self._loaded = True

def available_locales(self) -> list[str]:
"""Locales that have at least one loaded message."""
"""Locales that have at least one loaded message.

The list is cached at ``load()`` time; if ``load()`` hasn't run but
tests populated ``_messages`` directly, a one-off scan returns the
derived list without caching it (the test is outside the normal flow).
"""
if self._loaded:
return self._available_locales_list
return [locale for locale, msgs in self._messages.items() if msgs]

def messages(self, locale: str) -> dict[str, str]:
"""Flat dotted-key map for the given locale. Empty dict if unknown."""
return dict(self._messages.get(locale, {}))
def messages(self, locale: str) -> Mapping[str, str]:
"""Flat dotted-key map for the given locale. Empty mapping if unknown.

Returns an immutable view (``MappingProxyType``) into the cached
message dict — zero-copy. Callers that JSON-serialize the result
(e.g. Inertia shared props) should wrap with ``dict(...)`` at the
boundary.
"""
view = self._message_views.get(locale)
if view is not None:
return view
# Fallback: ``load()`` wasn't called (tests may populate _messages
# directly). Expose the raw dict as a proxy so Translator still works.
raw = self._messages.get(locale)
if raw is None:
return self._empty_view
return MappingProxyType(raw)


class _SafeFormatDict(dict):
Expand Down
17 changes: 16 additions & 1 deletion framework/core/simple_module_core/menu.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,9 @@

from dataclasses import dataclass, field
from enum import StrEnum
from typing import Literal

MenuItemMethod = Literal["get", "post"]


class MenuSection(StrEnum):
Expand All @@ -27,23 +30,34 @@ class MenuItem:
requires_auth: bool = True
roles: list[str] = field(default_factory=list)
"""Empty list = visible to all authenticated users."""
method: MenuItemMethod = "get"
"""HTTP method used when the item is activated. ``"post"`` renders as an
Inertia form submission so the target endpoint can be POST-only (e.g. logout)."""


class MenuRegistry:
"""Collects menu items from all modules and filters them per-request."""

def __init__(self) -> None:
self._items: list[MenuItem] = []
self._sorted: list[MenuItem] | None = None

def _invalidate(self) -> None:
self._sorted = None

def add(self, item: MenuItem) -> None:
self._items.append(item)
self._invalidate()

def add_many(self, items: list[MenuItem]) -> None:
self._items.extend(items)
self._invalidate()

@property
def all_items(self) -> list[MenuItem]:
return sorted(self._items, key=lambda i: i.order)
if self._sorted is None:
self._sorted = sorted(self._items, key=lambda i: i.order)
return self._sorted

def get_for_user(
self,
Expand All @@ -68,6 +82,7 @@ def get_for_user(
"label": item.label,
"url": item.url,
"icon": item.icon,
"method": item.method,
}
)

Expand Down
28 changes: 22 additions & 6 deletions framework/db/simple_module_db/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,22 +25,38 @@ class DatabaseState:
_listeners_registered: bool = field(default=False, repr=False)


def init_db(database_url: str, *, echo: bool = False) -> DatabaseState:
def init_db(
database_url: str,
*,
echo: bool = False,
pool_size: int = 10,
max_overflow: int = 20,
pool_pre_ping: bool = True,
pool_recycle: int = 1800,
) -> DatabaseState:
"""Create an async engine and session factory.

The pool options only take effect for server-side providers (Postgres).
SQLite uses SQLAlchemy's default pool (single-file, no network), so
passing ``pool_size``/etc. would raise ``TypeError`` — skipped below.

Returns a ``DatabaseState`` that should be stored on ``app.state.db``.
"""
provider = detect_provider(database_url)

connect_args: dict = {}
engine_kwargs: dict = {"echo": echo, "connect_args": connect_args}
if provider == DatabaseProvider.SQLITE:
connect_args["check_same_thread"] = False
else:
engine_kwargs.update(
pool_size=pool_size,
max_overflow=max_overflow,
pool_pre_ping=pool_pre_ping,
pool_recycle=pool_recycle,
)

engine = create_async_engine(
database_url,
echo=echo,
connect_args=connect_args,
)
engine = create_async_engine(database_url, **engine_kwargs)
# Scoped Session subclass so event listeners only fire for this engine's sessions
scoped_session_class = type("ScopedSession", (Session,), {})
session_factory = async_sessionmaker(
Expand Down
14 changes: 14 additions & 0 deletions framework/db/tests/test_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,20 @@ async def test_separate_init_db_calls_are_independent(self):
await db1.engine.dispose()
await db2.engine.dispose()

async def test_sqlite_ignores_pool_kwargs(self):
"""SQLite must not receive pool_size/max_overflow — they would raise TypeError."""
db_state = init_db(
"sqlite+aiosqlite:///:memory:",
pool_size=50,
max_overflow=100,
pool_pre_ping=True,
pool_recycle=60,
)
try:
assert db_state.engine is not None
finally:
await db_state.engine.dispose()


class TestGetDbDependency:
async def test_get_db_yields_session(self):
Expand Down
58 changes: 58 additions & 0 deletions framework/hosting/simple_module_hosting/_inertia_shared.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
"""Helpers for building Inertia shared-props payloads."""

from __future__ import annotations

import logging

from starlette.datastructures import Headers
from starlette.requests import Request
from starlette.types import Scope

logger = logging.getLogger(__name__)

_I18N_SESSION_LOCALE_KEY = "__i18n_locale"


def build_i18n_block(scope: Scope, request: Request) -> dict:
"""Assemble the ``i18n`` shared-props block for the current request.

Rules:

* No registry / no locale → serve an empty English block and log once.
* Inertia XHR partials (``X-Inertia: true``) reuse the client-side
cached messages; send ``messages: None`` unless the locale differs
from what was last served on this session.
* Full page loads and locale transitions ship the complete dict.
"""
registry = getattr(request.app.state, "i18n_registry", None)
locale = getattr(request.state, "locale", None)
if registry is None or locale is None:
logger.warning(
"InertiaLayoutDataMiddleware: i18n not fully wired "
"(registry_present=%s, locale_present=%s); serving empty messages",
registry is not None,
locale is not None,
)
return {"locale": "en", "supportedLocales": ["en"], "messages": {}}

is_inertia = Headers(scope=scope).get("x-inertia") == "true"
session_dict = scope.get("session")
# When the session is absent (pre-session-middleware routes, WebSocket
# upgrades), treat locale as "unchanged" so Inertia XHR requests still
# skip the messages payload. Non-Inertia requests will always ship them
# regardless of the session state.
if session_dict is not None:
last_locale = session_dict.get(_I18N_SESSION_LOCALE_KEY)
locale_changed = last_locale != locale
if locale_changed:
session_dict[_I18N_SESSION_LOCALE_KEY] = locale
else:
locale_changed = False
send_messages = (not is_inertia) or locale_changed
return {
"locale": locale,
"supportedLocales": registry.available_locales(),
# Inertia's JSON encoder doesn't accept the MappingProxyType view
# returned by messages(); copy to a plain dict at the boundary.
"messages": dict(registry.messages(locale)) if send_messages else None,
}
Loading