diff --git a/src/gateway/api/deps.py b/src/gateway/api/deps.py index 264e12aa2..729c14e94 100644 --- a/src/gateway/api/deps.py +++ b/src/gateway/api/deps.py @@ -916,6 +916,17 @@ def get_model_provider_port(db: PortSessionDep, container: ContainerDep) -> Mode return container.resolve(ModelProviderPort, db) +# ``get_db`` rather than ``PortSessionDep``: every management-plane caller here already +# holds a session from ``get_db``, and naming the same dependency shares it instead of +# opening a second one. Data-plane routes (chat/messages/responses) keep the port above. +def get_model_provider_port_shared( + db: Annotated[AsyncSession, Depends(get_db)], + container: ContainerDep, +) -> ModelProviderPort: + """Resolve the model-provider adapter, sharing the caller's own database session.""" + return container.resolve(ModelProviderPort, db) + + # Deliberately ``get_db`` and not ``PortSessionDep``: every surface that # resolves this port (the OTLP receiver, the telemetry read and purge # endpoints, user deletion) is standalone-only and already holds a session from @@ -1022,6 +1033,7 @@ def get_organization_guardrail_definition_service( IdentityProviderPortDep = Annotated[IdentityProviderPort, Depends(get_identity_provider_port)] McpServerPortDep = Annotated[McpServerPort, Depends(get_mcp_server_port)] ModelProviderPortDep = Annotated[ModelProviderPort, Depends(get_model_provider_port)] +ModelProviderPortSharedDep = Annotated[ModelProviderPort, Depends(get_model_provider_port_shared)] def get_org_provider_model_service( diff --git a/src/gateway/api/routes/catalog.py b/src/gateway/api/routes/catalog.py index d0fa2786f..041bed127 100644 --- a/src/gateway/api/routes/catalog.py +++ b/src/gateway/api/routes/catalog.py @@ -38,6 +38,7 @@ from gateway.api.deps import ( ModelProviderPortDep, + ModelProviderPortSharedDep, get_config, get_db, get_session_identity, @@ -661,7 +662,7 @@ async def list_catalog( config: Annotated[GatewayConfig, Depends(get_config)], caller: Annotated[CatalogCaller, Depends(verify_catalog_reader_or_public)], session_identity: Annotated[TenancyUser | None, Depends(get_session_identity)], - model_provider: ModelProviderPortDep, + model_provider: ModelProviderPortSharedDep, at_context: Annotated[ int | None, Query( @@ -724,7 +725,7 @@ async def get_catalog_model( config: Annotated[GatewayConfig, Depends(get_config)], caller: Annotated[CatalogCaller, Depends(verify_catalog_reader_or_public)], session_identity: Annotated[TenancyUser | None, Depends(get_session_identity)], - model_provider: ModelProviderPortDep, + model_provider: ModelProviderPortSharedDep, ) -> CatalogModelDetail: """One model and every offering of it this caller may use. diff --git a/src/gateway/api/routes/models.py b/src/gateway/api/routes/models.py index f2ae8a44e..e299d0883 100644 --- a/src/gateway/api/routes/models.py +++ b/src/gateway/api/routes/models.py @@ -8,7 +8,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from gateway.api.deps import ( - ModelProviderPortDep, + ModelProviderPortSharedDep, get_config, get_db, get_session_identity, @@ -183,7 +183,7 @@ async def list_models( config: Annotated[GatewayConfig, Depends(get_config)], auth: Annotated[tuple[APIKey | None, bool], Depends(verify_catalog_reader)], session_identity: Annotated[TenancyUser | None, Depends(get_session_identity)], - model_provider: ModelProviderPortDep, + model_provider: ModelProviderPortSharedDep, provider: Annotated[str | None, Query(description="Filter models by provider name")] = None, ) -> ModelListResponse: """List all available models. @@ -276,7 +276,7 @@ async def get_model( config: Annotated[GatewayConfig, Depends(get_config)], auth: Annotated[tuple[APIKey | None, bool], Depends(verify_catalog_reader)], session_identity: Annotated[TenancyUser | None, Depends(get_session_identity)], - model_provider: ModelProviderPortDep, + model_provider: ModelProviderPortSharedDep, ) -> ModelObject: """Get details for a specific model.""" api_key, _is_master_key = auth diff --git a/src/gateway/api/routes/organization_pricing.py b/src/gateway/api/routes/organization_pricing.py index 4f4dc314d..f81fd74df 100644 --- a/src/gateway/api/routes/organization_pricing.py +++ b/src/gateway/api/routes/organization_pricing.py @@ -27,7 +27,7 @@ from sqlalchemy.exc import IntegrityError, SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession -from gateway.api.deps import CurrentIdentity, ModelProviderPortDep, get_config, get_db, verify_master_key +from gateway.api.deps import CurrentIdentity, ModelProviderPortSharedDep, get_config, get_db, verify_master_key from gateway.core.config import GatewayConfig from gateway.models.money import as_float from gateway.models.pricing import OrganizationModelPricing @@ -212,7 +212,7 @@ class OrganizationModelPricingsPublic(BaseModel): def get_organization_pricing_service( db: Annotated[AsyncSession, Depends(get_db)], config: Annotated[GatewayConfig, Depends(get_config)], - model_provider: ModelProviderPortDep, + model_provider: ModelProviderPortSharedDep, ) -> OrganizationPricingService: """Build the pricing service on the request's session, provider map, and hosted-credential port.""" return OrganizationPricingService(db, config, model_provider=model_provider) diff --git a/src/gateway/api/routes/organization_routing.py b/src/gateway/api/routes/organization_routing.py index c6f43740c..08e7b24f4 100644 --- a/src/gateway/api/routes/organization_routing.py +++ b/src/gateway/api/routes/organization_routing.py @@ -61,7 +61,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlmodel import col -from gateway.api.deps import CurrentIdentity, ModelProviderPortDep, get_config, get_db, verify_master_key +from gateway.api.deps import CurrentIdentity, ModelProviderPortSharedDep, get_config, get_db, verify_master_key from gateway.api.routes.aliases import ( AliasRequest, AliasResponse, @@ -298,7 +298,7 @@ async def set_organization_routing_policy( db: Annotated[AsyncSession, Depends(get_db)], config: Annotated[GatewayConfig, Depends(get_config)], current_identity: CurrentIdentity, - model_provider: ModelProviderPortDep, + model_provider: ModelProviderPortSharedDep, ) -> PolicyResponse: """Create or update a stored policy in one of the organization's workspaces. @@ -380,7 +380,7 @@ async def set_organization_alias( db: Annotated[AsyncSession, Depends(get_db)], config: Annotated[GatewayConfig, Depends(get_config)], current_identity: CurrentIdentity, - model_provider: ModelProviderPortDep, + model_provider: ModelProviderPortSharedDep, ) -> AliasResponse: """Create or update a stored alias in one of the organization's workspaces. diff --git a/src/gateway/api/routes/playground.py b/src/gateway/api/routes/playground.py index 936476487..a8fd6e2df 100644 --- a/src/gateway/api/routes/playground.py +++ b/src/gateway/api/routes/playground.py @@ -66,7 +66,7 @@ CurrentIdentity, FileServiceDep, McpServerPortDep, - ModelProviderPortDep, + ModelProviderPortSharedDep, WebSearchPolicyPortDep, get_config, get_db, @@ -183,7 +183,7 @@ async def playground_chat_completions( files: FileServiceDep, config: Annotated[GatewayConfig, Depends(get_config)], log_writer: Annotated[LogWriter, Depends(get_log_writer)], - model_provider: ModelProviderPortDep, + model_provider: ModelProviderPortSharedDep, code_execution_port: CodeExecutionPortDep, mcp_server_port: McpServerPortDep, web_search_policy_port: WebSearchPolicyPortDep, diff --git a/tests/unit/test_model_provider_port_management_route_session.py b/tests/unit/test_model_provider_port_management_route_session.py new file mode 100644 index 000000000..78ec9118a --- /dev/null +++ b/tests/unit/test_model_provider_port_management_route_session.py @@ -0,0 +1,62 @@ +"""Regression for otari#1266: a management-plane route's model-provider port must +share the route's own database session, not open a second, independent one. +""" + +from typing import Annotated, cast + +import pytest +from fastapi import Depends, FastAPI +from fastapi.testclient import TestClient +from sqlalchemy.ext.asyncio import AsyncSession + +from gateway.api.routes.organization_pricing import get_organization_pricing_service +from gateway.core import database +from gateway.core.config import GatewayConfig +from gateway.ports.model_provider_port import ModelProviderPort +from gateway.services.organization_pricing_service import OrganizationPricingService + + +class _FakeSession: + """A distinguishable stand-in for ``AsyncSession``; only its identity matters here.""" + + async def __aenter__(self) -> "_FakeSession": + return self + + async def __aexit__(self, *exc: object) -> None: + return None + + +class _NullModelProviderPort: + """A do-nothing port; the stub container below only records which session built it.""" + + +class _RecordingContainer: + """Stands in for the composition-root container, recording ``resolve``'s session arg.""" + + def __init__(self) -> None: + self.resolved_with: list[AsyncSession | None] = [] + + def resolve(self, port: type[ModelProviderPort], session: AsyncSession | None) -> ModelProviderPort: + self.resolved_with.append(session) + return cast(ModelProviderPort, _NullModelProviderPort()) + + +def test_organization_pricing_service_shares_its_own_request_session(monkeypatch: pytest.MonkeyPatch) -> None: + # A fresh fake per call, so a route that opens two sessions is caught by identity. + monkeypatch.setattr(database, "_SessionLocal", lambda: _FakeSession()) + + container = _RecordingContainer() + app = FastAPI() + app.state.container = container + app.state.config = GatewayConfig(database_url="sqlite:///./test.db") + + @app.get("/probe") + async def probe( + service: Annotated[OrganizationPricingService, Depends(get_organization_pricing_service)], + ) -> dict[str, bool]: + return {"shared": len(container.resolved_with) == 1 and container.resolved_with[0] is service.db} + + response = TestClient(app).get("/probe") + + assert response.status_code == 200 + assert response.json() == {"shared": True}