From f5f3fca4ad51c296695d94790ec3a32f686ceda6 Mon Sep 17 00:00:00 2001 From: Chris Bartholomew Date: Tue, 13 Jan 2026 11:55:33 -0500 Subject: [PATCH] Fix: Load extensions in server.py for multi-worker deployments (#155) * Fix: Load extensions in server.py for multi-worker deployments When running with multiple workers (--workers 2), uvicorn uses `hindsight_api.server:app` import string instead of passing an app object. The server.py module was not loading tenant/operation validator extensions, causing authentication bypass in production. This fix: - Adds extension loading to server.py matching main.py behavior - Sets extension context on tenant extension for schema provisioning - Adds comprehensive unit tests for server.py extension loading The tests specifically verify: - TENANT extension is loaded when HINDSIGHT_API_TENANT_EXTENSION is set - OPERATION_VALIDATOR is loaded when configured - Extensions are passed to MemoryEngine constructor - Extension context is set on tenant extension - Server works correctly without extensions configured * Add unit tests for main.py extension loading (single-worker path) --- hindsight-api/hindsight_api/server.py | 32 +- hindsight-api/tests/test_main_module.py | 396 ++++++++++++++++++++++ hindsight-api/tests/test_server_module.py | 290 ++++++++++++++++ 3 files changed, 717 insertions(+), 1 deletion(-) create mode 100644 hindsight-api/tests/test_main_module.py create mode 100644 hindsight-api/tests/test_server_module.py diff --git a/hindsight-api/hindsight_api/server.py b/hindsight-api/hindsight_api/server.py index 38fe4238..10c6649e 100644 --- a/hindsight-api/hindsight_api/server.py +++ b/hindsight-api/hindsight_api/server.py @@ -7,6 +7,7 @@ This module provides the ASGI app for uvicorn import string usage: For CLI usage, use the hindsight-api command instead. """ +import logging import os import warnings @@ -17,6 +18,12 @@ warnings.filterwarnings("ignore", message="websockets.server.WebSocketServerProt from hindsight_api import MemoryEngine from hindsight_api.api import create_app from hindsight_api.config import get_config +from hindsight_api.extensions import ( + DefaultExtensionContext, + OperationValidatorExtension, + TenantExtension, + load_extension, +) # Disable tokenizers parallelism to avoid warnings os.environ["TOKENIZERS_PARALLELISM"] = "false" @@ -25,10 +32,33 @@ os.environ["TOKENIZERS_PARALLELISM"] = "false" config = get_config() config.configure_logging() +# Load operation validator extension if configured +operation_validator = load_extension("OPERATION_VALIDATOR", OperationValidatorExtension) +if operation_validator: + logging.info(f"Loaded operation validator: {operation_validator.__class__.__name__}") + +# Load tenant extension if configured +tenant_extension = load_extension("TENANT", TenantExtension) +if tenant_extension: + logging.info(f"Loaded tenant extension: {tenant_extension.__class__.__name__}") + # Create app at module level (required for uvicorn import string) # MemoryEngine reads configuration from environment variables automatically # Note: run_migrations=True by default, but migrations are idempotent so safe with workers -_memory = MemoryEngine(run_migrations=config.run_migrations_on_startup) +_memory = MemoryEngine( + operation_validator=operation_validator, + tenant_extension=tenant_extension, + run_migrations=config.run_migrations_on_startup, +) + +# Set extension context on tenant extension (needed for schema provisioning) +if tenant_extension: + extension_context = DefaultExtensionContext( + database_url=config.database_url, + memory_engine=_memory, + ) + tenant_extension.set_context(extension_context) + logging.info("Extension context set on tenant extension") # Create unified app with both HTTP and optionally MCP app = create_app( diff --git a/hindsight-api/tests/test_main_module.py b/hindsight-api/tests/test_main_module.py new file mode 100644 index 00000000..0923fad0 --- /dev/null +++ b/hindsight-api/tests/test_main_module.py @@ -0,0 +1,396 @@ +""" +Tests for hindsight_api.main module (single-worker code path). + +The main.py module is used when running with a single worker: + hindsight-api (or hindsight-api --workers 1) + +When workers=1, main.py creates the app directly and passes it to uvicorn. +These tests ensure that extensions are properly loaded in this code path. + +Compare with test_server_module.py which tests the multi-worker path (workers > 1). +""" + +import sys +from unittest.mock import MagicMock, patch + + +class TestMainModuleExtensionLoading: + """Tests that main.py correctly loads extensions when configured via environment.""" + + def test_main_loads_tenant_extension_when_configured(self, monkeypatch): + """ + Verify that main.py loads tenant extension from HINDSIGHT_API_TENANT_EXTENSION. + + This ensures extension loading works in the single-worker code path. + """ + # Set up environment to configure a tenant extension + monkeypatch.setenv( + "HINDSIGHT_API_TENANT_EXTENSION", + "tests.test_main_module:MockTenantExtension", + ) + # Ensure single worker mode + monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1") + + # Track what extensions were loaded via load_extension + loaded_extensions = {} + + # Get the real load_extension function + from hindsight_api.extensions.loader import load_extension as real_load_extension + + def tracking_load_extension(name, base_class): + """Track calls to load_extension and delegate to original.""" + result = real_load_extension(name, base_class) + loaded_extensions[name] = result + return result + + with patch("hindsight_api.main.MemoryEngine") as mock_engine, \ + patch("hindsight_api.main.create_app") as mock_create_app, \ + patch("hindsight_api.main.get_config") as mock_get_config, \ + patch("hindsight_api.main.load_extension", side_effect=tracking_load_extension), \ + patch("hindsight_api.main.DefaultExtensionContext"), \ + patch("hindsight_api.main.print_banner"), \ + patch("uvicorn.run"): # Don't actually start uvicorn + + mock_config = MagicMock() + mock_config.host = "0.0.0.0" + mock_config.port = 8888 + mock_config.log_level = "info" + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_engine.return_value = MagicMock() + mock_create_app.return_value = MagicMock() + + # Mock sys.argv to simulate CLI invocation + with patch.object(sys, 'argv', ['hindsight-api']): + from hindsight_api.main import main + main() + + # Verify TENANT extension was loaded + assert "TENANT" in loaded_extensions, \ + "main.py did not call load_extension('TENANT', ...) - extensions not loaded!" + assert loaded_extensions["TENANT"] is not None, \ + "load_extension('TENANT', ...) returned None despite env var being set" + assert isinstance(loaded_extensions["TENANT"], MockTenantExtension), \ + f"Expected MockTenantExtension, got {type(loaded_extensions['TENANT'])}" + + def test_main_loads_operation_validator_when_configured(self, monkeypatch): + """ + Verify that main.py loads operation validator from HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION. + """ + monkeypatch.setenv( + "HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION", + "tests.test_main_module:MockOperationValidator", + ) + monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1") + + loaded_extensions = {} + + from hindsight_api.extensions.loader import load_extension as real_load_extension + + def tracking_load_extension(name, base_class): + result = real_load_extension(name, base_class) + loaded_extensions[name] = result + return result + + with patch("hindsight_api.main.MemoryEngine") as mock_engine, \ + patch("hindsight_api.main.create_app") as mock_create_app, \ + patch("hindsight_api.main.get_config") as mock_get_config, \ + patch("hindsight_api.main.load_extension", side_effect=tracking_load_extension), \ + patch("hindsight_api.main.DefaultExtensionContext"), \ + patch("hindsight_api.main.print_banner"), \ + patch("uvicorn.run"): + + mock_config = MagicMock() + mock_config.host = "0.0.0.0" + mock_config.port = 8888 + mock_config.log_level = "info" + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_engine.return_value = MagicMock() + mock_create_app.return_value = MagicMock() + + with patch.object(sys, 'argv', ['hindsight-api']): + from hindsight_api.main import main + main() + + assert "OPERATION_VALIDATOR" in loaded_extensions, \ + "main.py did not call load_extension('OPERATION_VALIDATOR', ...)" + assert loaded_extensions["OPERATION_VALIDATOR"] is not None + assert isinstance(loaded_extensions["OPERATION_VALIDATOR"], MockOperationValidator) + + def test_main_passes_extensions_to_memory_engine(self, monkeypatch): + """ + Verify that main.py passes loaded extensions to MemoryEngine constructor. + + This is the critical test - even if extensions are loaded, they must be + passed to MemoryEngine for authentication to work. + """ + monkeypatch.setenv( + "HINDSIGHT_API_TENANT_EXTENSION", + "tests.test_main_module:MockTenantExtension", + ) + monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1") + + memory_engine_calls = [] + + def capture_memory_engine(*args, **kwargs): + memory_engine_calls.append({"args": args, "kwargs": kwargs}) + return MagicMock() + + with patch("hindsight_api.main.MemoryEngine", side_effect=capture_memory_engine), \ + patch("hindsight_api.main.create_app") as mock_create_app, \ + patch("hindsight_api.main.get_config") as mock_get_config, \ + patch("hindsight_api.main.DefaultExtensionContext"), \ + patch("hindsight_api.main.print_banner"), \ + patch("uvicorn.run"): + + mock_config = MagicMock() + mock_config.host = "0.0.0.0" + mock_config.port = 8888 + mock_config.log_level = "info" + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_create_app.return_value = MagicMock() + + with patch.object(sys, 'argv', ['hindsight-api']): + from hindsight_api.main import main + main() + + # Verify MemoryEngine was called + assert len(memory_engine_calls) == 1, "MemoryEngine should be called exactly once" + + call_kwargs = memory_engine_calls[0]["kwargs"] + + # THE CRITICAL ASSERTION: tenant_extension must be passed and not None + assert "tenant_extension" in call_kwargs, \ + "MemoryEngine was not called with tenant_extension parameter!" + assert call_kwargs["tenant_extension"] is not None, \ + "tenant_extension was None - main.py did not pass loaded extension to MemoryEngine!" + + def test_main_sets_extension_context_on_tenant_extension(self, monkeypatch): + """ + Verify that main.py sets the extension context on tenant extension. + + This is required for tenant extensions that need to provision schemas. + """ + monkeypatch.setenv( + "HINDSIGHT_API_TENANT_EXTENSION", + "tests.test_main_module:MockTenantExtension", + ) + monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1") + + captured_tenant_ext = [None] + + def capture_memory_engine(*args, **kwargs): + captured_tenant_ext[0] = kwargs.get("tenant_extension") + return MagicMock() + + context_created = [] + + def capture_context(*args, **kwargs): + ctx = MagicMock() + context_created.append(ctx) + return ctx + + with patch("hindsight_api.main.MemoryEngine", side_effect=capture_memory_engine), \ + patch("hindsight_api.main.create_app") as mock_create_app, \ + patch("hindsight_api.main.get_config") as mock_get_config, \ + patch("hindsight_api.main.DefaultExtensionContext", side_effect=capture_context), \ + patch("hindsight_api.main.print_banner"), \ + patch("uvicorn.run"): + + mock_config = MagicMock() + mock_config.host = "0.0.0.0" + mock_config.port = 8888 + mock_config.log_level = "info" + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_create_app.return_value = MagicMock() + + with patch.object(sys, 'argv', ['hindsight-api']): + from hindsight_api.main import main + main() + + # Verify context was created and set + assert len(context_created) == 1, "DefaultExtensionContext should be created" + assert captured_tenant_ext[0] is not None, "Tenant extension should be captured" + assert captured_tenant_ext[0]._context_set, \ + "set_context was not called on tenant extension" + + def test_main_works_without_extensions(self, monkeypatch): + """ + Verify that main.py works correctly when no extensions are configured. + """ + # Ensure no extension env vars are set + monkeypatch.delenv("HINDSIGHT_API_TENANT_EXTENSION", raising=False) + monkeypatch.delenv("HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION", raising=False) + monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1") + + memory_engine_calls = [] + + def capture_memory_engine(*args, **kwargs): + memory_engine_calls.append({"args": args, "kwargs": kwargs}) + return MagicMock() + + with patch("hindsight_api.main.MemoryEngine", side_effect=capture_memory_engine), \ + patch("hindsight_api.main.create_app") as mock_create_app, \ + patch("hindsight_api.main.get_config") as mock_get_config, \ + patch("hindsight_api.main.print_banner"), \ + patch("uvicorn.run"): + + mock_config = MagicMock() + mock_config.host = "0.0.0.0" + mock_config.port = 8888 + mock_config.log_level = "info" + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_create_app.return_value = MagicMock() + + with patch.object(sys, 'argv', ['hindsight-api']): + from hindsight_api.main import main + main() + + # Should work without extensions + assert len(memory_engine_calls) == 1 + call_kwargs = memory_engine_calls[0]["kwargs"] + + # Extensions should be None when not configured + assert call_kwargs.get("tenant_extension") is None + assert call_kwargs.get("operation_validator") is None + + def test_main_uses_app_object_for_single_worker(self, monkeypatch): + """ + Verify that main.py passes the app object (not import string) when workers=1. + + This is important because it means single-worker mode uses the app created + in main.py (with extensions loaded), not server.py. + """ + monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1") + monkeypatch.delenv("HINDSIGHT_API_TENANT_EXTENSION", raising=False) + + uvicorn_calls = [] + + def capture_uvicorn_run(**kwargs): + uvicorn_calls.append(kwargs) + + mock_app = MagicMock() + + with patch("hindsight_api.main.MemoryEngine") as mock_engine, \ + patch("hindsight_api.main.create_app", return_value=mock_app), \ + patch("hindsight_api.main.get_config") as mock_get_config, \ + patch("hindsight_api.main.print_banner"), \ + patch("uvicorn.run", side_effect=capture_uvicorn_run): + + mock_config = MagicMock() + mock_config.host = "0.0.0.0" + mock_config.port = 8888 + mock_config.log_level = "info" + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_engine.return_value = MagicMock() + + with patch.object(sys, 'argv', ['hindsight-api', '--workers', '1']): + from hindsight_api.main import main + main() + + assert len(uvicorn_calls) == 1 + # With workers=1, should pass app object, not import string + assert uvicorn_calls[0]["app"] is mock_app, \ + "main.py should pass app object (not import string) when workers=1" + + def test_main_uses_import_string_for_multiple_workers(self, monkeypatch): + """ + Verify that main.py uses import string when workers > 1. + + This is important because multi-worker mode requires server.py to be imported + by each worker process. + """ + monkeypatch.setenv("HINDSIGHT_API_WORKERS", "2") + monkeypatch.delenv("HINDSIGHT_API_TENANT_EXTENSION", raising=False) + + uvicorn_calls = [] + + def capture_uvicorn_run(**kwargs): + uvicorn_calls.append(kwargs) + + with patch("hindsight_api.main.MemoryEngine") as mock_engine, \ + patch("hindsight_api.main.create_app") as mock_create_app, \ + patch("hindsight_api.main.get_config") as mock_get_config, \ + patch("hindsight_api.main.print_banner"), \ + patch("uvicorn.run", side_effect=capture_uvicorn_run): + + mock_config = MagicMock() + mock_config.host = "0.0.0.0" + mock_config.port = 8888 + mock_config.log_level = "info" + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_engine.return_value = MagicMock() + mock_create_app.return_value = MagicMock() + + with patch.object(sys, 'argv', ['hindsight-api', '--workers', '2']): + from hindsight_api.main import main + main() + + assert len(uvicorn_calls) == 1 + # With workers > 1, should use import string + assert uvicorn_calls[0]["app"] == "hindsight_api.server:app", \ + "main.py should use import string when workers > 1" + assert uvicorn_calls[0]["workers"] == 2 + + +# Mock extensions for testing +from hindsight_api.extensions import ( + TenantExtension, + TenantContext, + RequestContext, + OperationValidatorExtension, + ValidationResult, + RetainContext, + RecallContext, + ReflectContext, +) + + +class MockTenantExtension(TenantExtension): + """Mock tenant extension for testing main.py extension loading.""" + + def __init__(self, config: dict): + super().__init__(config) + self._context_set = False + + async def authenticate(self, request_context: RequestContext) -> TenantContext: + return TenantContext(schema_name="public") + + def set_context(self, context) -> None: + self._context_set = True + + +class MockOperationValidator(OperationValidatorExtension): + """Mock operation validator for testing main.py extension loading.""" + + def __init__(self, config: dict): + super().__init__(config) + + async def validate_retain(self, ctx: RetainContext) -> ValidationResult: + return ValidationResult.accept() + + async def validate_recall(self, ctx: RecallContext) -> ValidationResult: + return ValidationResult.accept() + + async def validate_reflect(self, ctx: ReflectContext) -> ValidationResult: + return ValidationResult.accept() diff --git a/hindsight-api/tests/test_server_module.py b/hindsight-api/tests/test_server_module.py new file mode 100644 index 00000000..bc1d29bf --- /dev/null +++ b/hindsight-api/tests/test_server_module.py @@ -0,0 +1,290 @@ +""" +Tests for hindsight_api.server module (multi-worker code path). + +The server.py module is used when running with multiple workers: + uvicorn hindsight_api.server:app --workers 2 + +This module executes code at import time, creating the app at module level. +These tests ensure that extensions are properly loaded in this code path, +which was previously a regression that caused authentication bypass in production. +""" + +import importlib +import sys +from unittest.mock import MagicMock, patch + + +def _clean_server_module(): + """Remove hindsight_api.server from sys.modules for fresh import.""" + modules_to_remove = [k for k in sys.modules.keys() if k.startswith("hindsight_api.server")] + for mod in modules_to_remove: + del sys.modules[mod] + + +class TestServerModuleExtensionLoading: + """Tests that server.py correctly loads extensions when configured via environment.""" + + def test_server_loads_tenant_extension_when_configured(self, monkeypatch): + """ + Verify that server.py loads tenant extension from HINDSIGHT_API_TENANT_EXTENSION. + + This test catches the regression where server.py didn't call load_extension(), + causing authentication to be bypassed in multi-worker deployments. + """ + # Set up environment to configure a tenant extension + monkeypatch.setenv( + "HINDSIGHT_API_TENANT_EXTENSION", + "tests.test_server_module:MockTenantExtension", + ) + + _clean_server_module() + + # Track what extensions were loaded via load_extension + loaded_extensions = {} + + # Get the real load_extension function + from hindsight_api.extensions.loader import load_extension as real_load_extension + + def tracking_load_extension(name, base_class): + """Track calls to load_extension and delegate to original.""" + result = real_load_extension(name, base_class) + loaded_extensions[name] = result + return result + + # Patch at source level BEFORE importing server + # Note: We patch the entire hindsight_api module namespace + with patch("hindsight_api.MemoryEngine") as mock_engine, \ + patch("hindsight_api.api.create_app") as mock_create_app, \ + patch("hindsight_api.config.get_config") as mock_get_config, \ + patch("hindsight_api.extensions.load_extension", side_effect=tracking_load_extension), \ + patch("hindsight_api.extensions.DefaultExtensionContext"): + + mock_config = MagicMock() + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_engine.return_value = MagicMock() + mock_create_app.return_value = MagicMock() + + # Now import server - this triggers module-level code + import hindsight_api.server + + # Verify TENANT extension was loaded + assert "TENANT" in loaded_extensions, \ + "server.py did not call load_extension('TENANT', ...) - extensions not loaded!" + assert loaded_extensions["TENANT"] is not None, \ + "load_extension('TENANT', ...) returned None despite env var being set" + assert isinstance(loaded_extensions["TENANT"], MockTenantExtension), \ + f"Expected MockTenantExtension, got {type(loaded_extensions['TENANT'])}" + + def test_server_loads_operation_validator_when_configured(self, monkeypatch): + """ + Verify that server.py loads operation validator from HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION. + """ + monkeypatch.setenv( + "HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION", + "tests.test_server_module:MockOperationValidator", + ) + + _clean_server_module() + + loaded_extensions = {} + + from hindsight_api.extensions.loader import load_extension as real_load_extension + + def tracking_load_extension(name, base_class): + result = real_load_extension(name, base_class) + loaded_extensions[name] = result + return result + + with patch("hindsight_api.MemoryEngine") as mock_engine, \ + patch("hindsight_api.api.create_app") as mock_create_app, \ + patch("hindsight_api.config.get_config") as mock_get_config, \ + patch("hindsight_api.extensions.load_extension", side_effect=tracking_load_extension), \ + patch("hindsight_api.extensions.DefaultExtensionContext"): + + mock_config = MagicMock() + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_engine.return_value = MagicMock() + mock_create_app.return_value = MagicMock() + + import hindsight_api.server + + assert "OPERATION_VALIDATOR" in loaded_extensions, \ + "server.py did not call load_extension('OPERATION_VALIDATOR', ...)" + assert loaded_extensions["OPERATION_VALIDATOR"] is not None + assert isinstance(loaded_extensions["OPERATION_VALIDATOR"], MockOperationValidator) + + def test_server_passes_extensions_to_memory_engine(self, monkeypatch): + """ + Verify that server.py passes loaded extensions to MemoryEngine constructor. + + This is the critical test - even if extensions are loaded, they must be + passed to MemoryEngine for authentication to work. + """ + monkeypatch.setenv( + "HINDSIGHT_API_TENANT_EXTENSION", + "tests.test_server_module:MockTenantExtension", + ) + + _clean_server_module() + + memory_engine_calls = [] + + def capture_memory_engine(*args, **kwargs): + memory_engine_calls.append({"args": args, "kwargs": kwargs}) + return MagicMock() + + with patch("hindsight_api.MemoryEngine", side_effect=capture_memory_engine), \ + patch("hindsight_api.api.create_app") as mock_create_app, \ + patch("hindsight_api.config.get_config") as mock_get_config, \ + patch("hindsight_api.extensions.DefaultExtensionContext"): + + mock_config = MagicMock() + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_create_app.return_value = MagicMock() + + import hindsight_api.server + + # Verify MemoryEngine was called + assert len(memory_engine_calls) == 1, "MemoryEngine should be called exactly once" + + call_kwargs = memory_engine_calls[0]["kwargs"] + + # THE CRITICAL ASSERTION: tenant_extension must be passed and not None + assert "tenant_extension" in call_kwargs, \ + "MemoryEngine was not called with tenant_extension parameter!" + assert call_kwargs["tenant_extension"] is not None, \ + "tenant_extension was None - server.py did not pass loaded extension to MemoryEngine!" + + def test_server_sets_extension_context_on_tenant_extension(self, monkeypatch): + """ + Verify that server.py sets the extension context on tenant extension. + + This is required for tenant extensions that need to provision schemas. + """ + monkeypatch.setenv( + "HINDSIGHT_API_TENANT_EXTENSION", + "tests.test_server_module:MockTenantExtension", + ) + + _clean_server_module() + + context_set_calls = [] + captured_tenant_ext = [None] + + def capture_memory_engine(*args, **kwargs): + captured_tenant_ext[0] = kwargs.get("tenant_extension") + return MagicMock() + + def capture_context(*args, **kwargs): + ctx = MagicMock() + context_set_calls.append(ctx) + return ctx + + with patch("hindsight_api.MemoryEngine", side_effect=capture_memory_engine), \ + patch("hindsight_api.api.create_app") as mock_create_app, \ + patch("hindsight_api.config.get_config") as mock_get_config, \ + patch("hindsight_api.extensions.DefaultExtensionContext", side_effect=capture_context): + + mock_config = MagicMock() + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_create_app.return_value = MagicMock() + + import hindsight_api.server + + # Verify context was created and set + assert len(context_set_calls) == 1, "DefaultExtensionContext should be created" + assert captured_tenant_ext[0] is not None, "Tenant extension should be captured" + assert captured_tenant_ext[0]._context_set, \ + "set_context was not called on tenant extension" + + def test_server_works_without_extensions(self, monkeypatch): + """ + Verify that server.py works correctly when no extensions are configured. + """ + # Ensure no extension env vars are set + monkeypatch.delenv("HINDSIGHT_API_TENANT_EXTENSION", raising=False) + monkeypatch.delenv("HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION", raising=False) + + _clean_server_module() + + memory_engine_calls = [] + + def capture_memory_engine(*args, **kwargs): + memory_engine_calls.append({"args": args, "kwargs": kwargs}) + return MagicMock() + + with patch("hindsight_api.MemoryEngine", side_effect=capture_memory_engine), \ + patch("hindsight_api.api.create_app") as mock_create_app, \ + patch("hindsight_api.config.get_config") as mock_get_config: + + mock_config = MagicMock() + mock_config.mcp_enabled = False + mock_config.run_migrations_on_startup = False + mock_config.database_url = "postgresql://test:test@localhost/test" + mock_get_config.return_value = mock_config + mock_create_app.return_value = MagicMock() + + import hindsight_api.server + + # Should work without extensions + assert len(memory_engine_calls) == 1 + call_kwargs = memory_engine_calls[0]["kwargs"] + + # Extensions should be None when not configured + assert call_kwargs.get("tenant_extension") is None + assert call_kwargs.get("operation_validator") is None + + +# Mock extensions for testing +from hindsight_api.extensions import ( + TenantExtension, + TenantContext, + RequestContext, + OperationValidatorExtension, + ValidationResult, + RetainContext, + RecallContext, + ReflectContext, +) + + +class MockTenantExtension(TenantExtension): + """Mock tenant extension for testing server.py extension loading.""" + + def __init__(self, config: dict): + super().__init__(config) + self._context_set = False + + async def authenticate(self, request_context: RequestContext) -> TenantContext: + return TenantContext(schema_name="public") + + def set_context(self, context) -> None: + self._context_set = True + + +class MockOperationValidator(OperationValidatorExtension): + """Mock operation validator for testing server.py extension loading.""" + + def __init__(self, config: dict): + super().__init__(config) + + async def validate_retain(self, ctx: RetainContext) -> ValidationResult: + return ValidationResult.accept() + + async def validate_recall(self, ctx: RecallContext) -> ValidationResult: + return ValidationResult.accept() + + async def validate_reflect(self, ctx: ReflectContext) -> ValidationResult: + return ValidationResult.accept()