From d7101f89c0e9b7c05cdcacb1169806107ea33e5c Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Mon, 21 Sep 2026 01:13:01 +0000 Subject: [PATCH] fix(providers): share the caller's own session with the model-provider port Catalog, models, organization pricing, organization routing, and the playground chat route take get_db for their session but resolved ModelProviderPort through get_model_provider_port, which takes its session from get_db_if_needed. FastAPI treats those as separate dependencies and opens a second, independent session for the port instead of sharing the caller's. Added get_model_provider_port_shared, which takes the session from get_db directly like the three sibling ports already do, and switched those five routes to it. Chat, messages, and responses keep the existing dependency: they take get_db_if_needed themselves, and hybrid mode has no local database at all. Fixes #1266 --- src/gateway/api/deps.py | 12 ++++ src/gateway/api/routes/catalog.py | 5 +- src/gateway/api/routes/models.py | 6 +- .../api/routes/organization_pricing.py | 4 +- .../api/routes/organization_routing.py | 6 +- src/gateway/api/routes/playground.py | 4 +- ..._provider_port_management_route_session.py | 62 +++++++++++++++++++ 7 files changed, 87 insertions(+), 12 deletions(-) create mode 100644 tests/unit/test_model_provider_port_management_route_session.py diff --git a/src/gateway/api/deps.py b/src/gateway/api/deps.py index f4394d6b28..efb1f03869 100644 --- a/src/gateway/api/deps.py +++ b/src/gateway/api/deps.py @@ -908,6 +908,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 @@ -1009,6 +1020,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 4a2504f233..8dd093b613 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, @@ -656,7 +657,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( @@ -719,7 +720,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 f2ae8a44e4..e299d08832 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 4f4dc314dd..f81fd74dfe 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 c6f43740c2..08e7b24f4b 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 cf100f33c7..809b955437 100644 --- a/src/gateway/api/routes/playground.py +++ b/src/gateway/api/routes/playground.py @@ -66,7 +66,7 @@ CurrentIdentity, FileServiceDep, McpServerPortDep, - ModelProviderPortDep, + ModelProviderPortSharedDep, get_config, get_db, get_log_writer, @@ -182,7 +182,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, key_format: ApiKeyFormatPortDep, 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 0000000000..78ec9118a3 --- /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}