Skip to content
Open
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
12 changes: 12 additions & 0 deletions src/gateway/api/deps.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
5 changes: 3 additions & 2 deletions src/gateway/api/routes/catalog.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@

from gateway.api.deps import (
ModelProviderPortDep,
ModelProviderPortSharedDep,
get_config,
get_db,
get_session_identity,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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.

Expand Down
6 changes: 3 additions & 3 deletions src/gateway/api/routes/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
from sqlalchemy.ext.asyncio import AsyncSession

from gateway.api.deps import (
ModelProviderPortDep,
ModelProviderPortSharedDep,
get_config,
get_db,
get_session_identity,
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down
4 changes: 2 additions & 2 deletions src/gateway/api/routes/organization_pricing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
6 changes: 3 additions & 3 deletions src/gateway/api/routes/organization_routing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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.

Expand Down
4 changes: 2 additions & 2 deletions src/gateway/api/routes/playground.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@
CurrentIdentity,
FileServiceDep,
McpServerPortDep,
ModelProviderPortDep,
ModelProviderPortSharedDep,
get_config,
get_db,
get_log_writer,
Expand Down Expand Up @@ -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,
Expand Down
62 changes: 62 additions & 0 deletions tests/unit/test_model_provider_port_management_route_session.py
Original file line number Diff line number Diff line change
@@ -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}
Loading