From 820e74c688264d9b23c9be43c873a8741bd3d840 Mon Sep 17 00:00:00 2001 From: Sergio Gimenez Date: Tue, 22 Sep 2026 09:32:16 +0200 Subject: [PATCH] feat(fm): partner registration admin API (#20) --- .../adapters/database/partner_repo.py | 51 ++++- .../adapters/database/tables.py | 1 + src/federation_manager/api/errors.py | 27 +++ .../api/internal/partners.py | 114 +++++++++++ .../application/partners.py | 76 ++++++++ src/federation_manager/dependencies.py | 7 + src/federation_manager/domain/errors.py | 13 ++ src/federation_manager/domain/models.py | 1 + src/federation_manager/domain/ports.py | 8 + src/federation_manager/main.py | 2 + .../migrations/versions/0003_partner_name.py | 31 +++ tests/fakes.py | 27 ++- .../test_federation_context_repo.py | 1 + tests/integration/test_fm_srm_loop.py | 1 + tests/integration/test_outbound_repos.py | 1 + .../integration/test_partner_registry_repo.py | 111 +++++++++++ tests/integration/test_partner_repo.py | 2 + tests/integration/test_row_timestamps.py | 1 + .../integration/test_two_stack_federation.py | 2 + tests/test_partner_registry.py | 183 ++++++++++++++++++ 20 files changed, 656 insertions(+), 4 deletions(-) create mode 100644 src/federation_manager/api/internal/partners.py create mode 100644 src/federation_manager/application/partners.py create mode 100644 src/federation_manager/migrations/versions/0003_partner_name.py create mode 100644 tests/integration/test_partner_registry_repo.py create mode 100644 tests/test_partner_registry.py diff --git a/src/federation_manager/adapters/database/partner_repo.py b/src/federation_manager/adapters/database/partner_repo.py index 550db28..8337c9f 100644 --- a/src/federation_manager/adapters/database/partner_repo.py +++ b/src/federation_manager/adapters/database/partner_repo.py @@ -1,12 +1,17 @@ from typing import Any from uuid import UUID -from sqlalchemy import Select, select +from sqlalchemy import Select, insert, select, update +from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession from federation_manager.adapters.database.tables import partner_ops +from federation_manager.domain.errors import PartnerRegistrationConflict from federation_manager.domain.models import PartnerOP +# columns carrying a UNIQUE constraint; their names appear in the violation message +_UNIQUE_FIELDS = ("mcc_mnc", "oauth2_client_id") + class PostgresPartnerRepo: def __init__(self, session: AsyncSession) -> None: @@ -25,6 +30,49 @@ class PostgresPartnerRepo: rows = (await self._session.execute(stmt)).all() return [self._to_partner(row) for row in rows] + async def find_by_mcc_mnc(self, mcc_mnc: str) -> PartnerOP | None: + return await self._one(select(partner_ops).where(partner_ops.c.mcc_mnc == mcc_mnc)) + + async def list_all(self) -> list[PartnerOP]: + stmt = select(partner_ops).order_by(partner_ops.c.created_at) + rows = (await self._session.execute(stmt)).all() + return [self._to_partner(row) for row in rows] + + async def add(self, partner: PartnerOP) -> None: + await self._write(insert(partner_ops).values(id=partner.id, **self._values(partner))) + + async def update(self, partner: PartnerOP) -> None: + await self._write( + update(partner_ops) + .where(partner_ops.c.id == partner.id) + .values(**self._values(partner)) + ) + + async def _write(self, stmt: Any) -> None: + try: + await self._session.execute(stmt) + await self._session.commit() + except IntegrityError as error: + await self._session.rollback() + # a concurrent registration won the unique index between our read and this write + field = next((f for f in _UNIQUE_FIELDS if f in str(error.orig)), None) + if field is None: + raise + raise PartnerRegistrationConflict(field) from None + + @staticmethod + def _values(partner: PartnerOP) -> dict[str, Any]: + return { + "name": partner.name, + "mcc_mnc": partner.mcc_mnc, + "oauth2_client_id": partner.oauth2_client_id, + "base_url": partner.base_url, + "our_client_id": partner.our_client_id, + "our_client_secret_ref": partner.our_client_secret_ref, + "token_endpoint": partner.token_endpoint, + "status": partner.status, + } + async def _one(self, stmt: Select[tuple[object, ...]]) -> PartnerOP | None: row = (await self._session.execute(stmt)).one_or_none() return None if row is None else self._to_partner(row) @@ -40,4 +88,5 @@ class PostgresPartnerRepo: our_client_secret_ref=row.our_client_secret_ref, token_endpoint=row.token_endpoint, base_url=row.base_url, + name=row.name, ) diff --git a/src/federation_manager/adapters/database/tables.py b/src/federation_manager/adapters/database/tables.py index 288594f..eda6e68 100644 --- a/src/federation_manager/adapters/database/tables.py +++ b/src/federation_manager/adapters/database/tables.py @@ -35,6 +35,7 @@ partner_ops = Table( "partner_ops", metadata, Column("id", PGUUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid()), + Column("name", String(255), nullable=False), Column("mcc_mnc", String(10), nullable=False, unique=True), Column("oauth2_client_id", String(255), nullable=False, unique=True), Column("base_url", Text, nullable=False), diff --git a/src/federation_manager/api/errors.py b/src/federation_manager/api/errors.py index 01648a9..1033db7 100644 --- a/src/federation_manager/api/errors.py +++ b/src/federation_manager/api/errors.py @@ -17,9 +17,11 @@ from federation_manager.domain.errors import ( IdempotencyKeyReused, NetworkQueryNotApplicable, NoRouteMatched, + PartnerCredentialsIncomplete, PartnerEndpointConfigurationError, PartnerNotActive, PartnerNotRegistered, + PartnerRegistrationConflict, PartnerRejectedRequest, PartnerRequestFailed, PartnerResponseInvalid, @@ -320,3 +322,28 @@ def register_exception_handlers(app: FastAPI) -> None: "Outbound federation is not correctly configured for the resolved partner.", request.url.path, ) + + # Internal admin surface: the caller is our own operator, so the detail can name the field. + @app.exception_handler(PartnerRegistrationConflict) + async def _registration_conflict( + request: Request, exc: PartnerRegistrationConflict + ) -> JSONResponse: + return problem( + 409, + "partner-registration-conflict", + "Partner Registration Conflict", + f"A partner is already registered with this {exc.field}.", + request.url.path, + ) + + @app.exception_handler(PartnerCredentialsIncomplete) + async def _credentials_incomplete( + request: Request, exc: PartnerCredentialsIncomplete + ) -> JSONResponse: + return problem( + 422, + "partner-credentials-incomplete", + "Partner Credentials Incomplete", + str(exc), + request.url.path, + ) diff --git a/src/federation_manager/api/internal/partners.py b/src/federation_manager/api/internal/partners.py new file mode 100644 index 0000000..cf3bdc0 --- /dev/null +++ b/src/federation_manager/api/internal/partners.py @@ -0,0 +1,114 @@ +from typing import Annotated, Literal, Self +from urllib.parse import urlsplit +from uuid import UUID + +from fastapi import APIRouter, Depends +from pydantic import AfterValidator, BaseModel, ConfigDict, Field, model_validator + +from federation_manager.application.partners import PartnerRegistration, PartnerRegistryService +from federation_manager.dependencies import get_partner_registry_service +from federation_manager.domain.models import PartnerOP + +router = APIRouter(prefix="/internal/partners", tags=["internal-partners"]) + +Service = Annotated[PartnerRegistryService, Depends(get_partner_registry_service)] + + +def _http_url(value: str) -> str: + # checked but stored as sent: HttpUrl would silently append a trailing slash + parts = urlsplit(value) + if parts.scheme not in ("http", "https") or not parts.netloc: + raise ValueError("must be an absolute http(s) URL") + return value + + +def _secret_ref(value: str) -> str: + # a path to the mounted secret, never the secret itself (persistence model: no inline secrets) + if not value.startswith("/"): + raise ValueError("must be an absolute path to the mounted secret file") + return value + + +Name = Annotated[str, Field(min_length=1, max_length=255)] +MccMnc = Annotated[str, Field(pattern=r"^[0-9]{3}-?[0-9]{2,3}$")] +ClientId = Annotated[str, Field(min_length=1, max_length=255)] +Url = Annotated[str, AfterValidator(_http_url)] +SecretRef = Annotated[str, AfterValidator(_secret_ref)] +Status = Literal["pending", "active", "suspended", "decommissioned"] + + +class PartnerCreate(BaseModel): + model_config = ConfigDict(extra="forbid") + + name: Name + mcc_mnc: MccMnc + oauth2_client_id: ClientId = Field( + description="Client id we issued to the partner in our Keycloak; matched to token azp" + ) + base_url: Url = Field(description="Partner's EWBI federation endpoint") + our_client_id: ClientId | None = Field( + default=None, description="Client id the partner issued to us for outbound calls" + ) + our_client_secret_ref: SecretRef | None = None + token_endpoint: Url | None = Field( + default=None, description="Partner's OAuth2 token endpoint for our outbound calls" + ) + status: Literal["pending", "active"] = "pending" + + +class PartnerUpdate(BaseModel): + """Partial update: fields left out are unchanged; null clears an optional field.""" + + model_config = ConfigDict(extra="forbid") + + name: Name | None = None + oauth2_client_id: ClientId | None = None + base_url: Url | None = None + our_client_id: ClientId | None = None + our_client_secret_ref: SecretRef | None = None + token_endpoint: Url | None = None + status: Status | None = None + + @model_validator(mode="after") + def _required_not_cleared(self) -> Self: + for field in ("name", "oauth2_client_id", "base_url", "status"): + if field in self.model_fields_set and getattr(self, field) is None: + raise ValueError(f"{field} cannot be null") + return self + + +class PartnerView(BaseModel): + id: UUID + name: str + mcc_mnc: str + oauth2_client_id: str + base_url: str | None + our_client_id: str | None + our_client_secret_ref: str | None + token_endpoint: str | None + status: Status + + @classmethod + def of(cls, partner: PartnerOP) -> "PartnerView": + return cls.model_validate(partner, from_attributes=True) + + +@router.post("", status_code=201, response_model=PartnerView) +async def register_partner(body: PartnerCreate, service: Service) -> PartnerView: + return PartnerView.of(await service.register(PartnerRegistration(**body.model_dump()))) + + +@router.get("", response_model=list[PartnerView]) +async def list_partners(service: Service) -> list[PartnerView]: + return [PartnerView.of(partner) for partner in await service.list()] + + +@router.get("/{partner_op_id}", response_model=PartnerView) +async def get_partner(partner_op_id: UUID, service: Service) -> PartnerView: + return PartnerView.of(await service.get(partner_op_id)) + + +@router.patch("/{partner_op_id}", response_model=PartnerView) +async def update_partner(partner_op_id: UUID, body: PartnerUpdate, service: Service) -> PartnerView: + changes = body.model_dump(exclude_unset=True) + return PartnerView.of(await service.update(partner_op_id, changes)) diff --git a/src/federation_manager/application/partners.py b/src/federation_manager/application/partners.py new file mode 100644 index 0000000..a1e35b8 --- /dev/null +++ b/src/federation_manager/application/partners.py @@ -0,0 +1,76 @@ +from collections.abc import Callable +from dataclasses import dataclass, replace +from typing import Any +from uuid import UUID, uuid4 + +from federation_manager.domain.errors import ( + PartnerCredentialsIncomplete, + PartnerNotRegistered, + PartnerRegistrationConflict, +) +from federation_manager.domain.models import PartnerOP +from federation_manager.domain.ports import PartnerRepositoryPort + + +@dataclass(frozen=True) +class PartnerRegistration: + """What operators exchange before federating (OPG.04 Table 371), plus FM's own keys.""" + + name: str + mcc_mnc: str + oauth2_client_id: str + base_url: str + our_client_id: str | None = None + our_client_secret_ref: str | None = None + token_endpoint: str | None = None + status: str = "pending" + + +class PartnerRegistryService: + def __init__( + self, partners: PartnerRepositoryPort, *, id_factory: Callable[[], UUID] = uuid4 + ) -> None: + self._partners = partners + self._new_id = id_factory + + async def register(self, registration: PartnerRegistration) -> PartnerOP: + _require_complete_credentials(registration) + if await self._partners.find_by_mcc_mnc(registration.mcc_mnc) is not None: + raise PartnerRegistrationConflict("mcc_mnc") + await self._require_free_client_id(registration.oauth2_client_id) + + partner = PartnerOP(id=self._new_id(), **vars(registration)) + # a concurrent registration can still win the unique index; the repo reports it as 409 + await self._partners.add(partner) + return partner + + async def get(self, partner_id: UUID) -> PartnerOP: + partner = await self._partners.find_by_id(partner_id) + if partner is None: + raise PartnerNotRegistered(partner_id) + return partner + + async def list(self) -> list[PartnerOP]: + return await self._partners.list_all() + + async def update(self, partner_id: UUID, changes: dict[str, Any]) -> PartnerOP: + partner = await self.get(partner_id) + client_id = changes.get("oauth2_client_id") + if client_id is not None and client_id != partner.oauth2_client_id: + await self._require_free_client_id(client_id) + + updated = replace(partner, **changes) + _require_complete_credentials(updated) + await self._partners.update(updated) + return updated + + async def _require_free_client_id(self, client_id: str) -> None: + if await self._partners.find_by_oauth2_client_id(client_id) is not None: + raise PartnerRegistrationConflict("oauth2_client_id") + + +def _require_complete_credentials(partner: PartnerRegistration | PartnerOP) -> None: + # the outbound token grant needs all three; a partial set fails only on first use + outbound = (partner.our_client_id, partner.our_client_secret_ref, partner.token_endpoint) + if any(v is not None for v in outbound) and any(v is None for v in outbound): + raise PartnerCredentialsIncomplete diff --git a/src/federation_manager/dependencies.py b/src/federation_manager/dependencies.py index c37dcb7..34edeaf 100644 --- a/src/federation_manager/dependencies.py +++ b/src/federation_manager/dependencies.py @@ -22,6 +22,7 @@ from federation_manager.application.federation import ( ) from federation_manager.application.outbound import OutboundFederationService from federation_manager.application.partner_federations import PartnerFederationService +from federation_manager.application.partners import PartnerRegistryService from federation_manager.application.queries import InboundQueryService from federation_manager.core.config import get_settings from federation_manager.domain.ports import ( @@ -145,6 +146,12 @@ def get_partner_federation_service( return PartnerFederationService(partner_repo, context_repo, ewbi_client, establishment, local) +def get_partner_registry_service( + partner_repo: Annotated[PartnerRepositoryPort, Depends(get_partner_repo)], +) -> PartnerRegistryService: + return PartnerRegistryService(partner_repo) + + def get_databus_health(request: Request) -> DataBusHealthPort: publisher: DataBusHealthPort = request.app.state.command_publisher return publisher diff --git a/src/federation_manager/domain/errors.py b/src/federation_manager/domain/errors.py index 4de1b04..2313a65 100644 --- a/src/federation_manager/domain/errors.py +++ b/src/federation_manager/domain/errors.py @@ -153,3 +153,16 @@ class UnsupportedApiType(FederationError): def __init__(self, api_type: str) -> None: super().__init__(f"api type {api_type!r} has no EWBI service path") self.api_type = api_type + + +class PartnerRegistrationConflict(FederationError): + def __init__(self, field: str) -> None: + super().__init__(f"a partner is already registered with this {field}") + self.field = field + + +class PartnerCredentialsIncomplete(FederationError): + def __init__(self) -> None: + super().__init__( + "our_client_id, our_client_secret_ref and token_endpoint must be set together" + ) diff --git a/src/federation_manager/domain/models.py b/src/federation_manager/domain/models.py index b4ae8c1..18fc537 100644 --- a/src/federation_manager/domain/models.py +++ b/src/federation_manager/domain/models.py @@ -13,6 +13,7 @@ class PartnerOP: our_client_secret_ref: str | None = None token_endpoint: str | None = None base_url: str | None = None + name: str = "" def is_active(self) -> bool: return self.status == "active" diff --git a/src/federation_manager/domain/ports.py b/src/federation_manager/domain/ports.py index 5f0d43b..771ef3b 100644 --- a/src/federation_manager/domain/ports.py +++ b/src/federation_manager/domain/ports.py @@ -21,6 +21,14 @@ class PartnerRepositoryPort(Protocol): async def list_active(self) -> list[PartnerOP]: ... + async def find_by_mcc_mnc(self, mcc_mnc: str) -> PartnerOP | None: ... + + async def list_all(self) -> list[PartnerOP]: ... + + async def add(self, partner: PartnerOP) -> None: ... + + async def update(self, partner: PartnerOP) -> None: ... + class JwtValidatorPort(Protocol): async def validate(self, token: str) -> ValidatedClaims: ... diff --git a/src/federation_manager/main.py b/src/federation_manager/main.py index f148a38..63ee325 100644 --- a/src/federation_manager/main.py +++ b/src/federation_manager/main.py @@ -36,6 +36,7 @@ from federation_manager.api.internal.federation import router as internal_federa from federation_manager.api.internal.partner_federations import ( router as internal_partner_federations_router, ) +from federation_manager.api.internal.partners import router as internal_partners_router from federation_manager.api.platform.health import router as health_router from federation_manager.application.events import OperationCompletedConsumer from federation_manager.application.federation import ( @@ -146,6 +147,7 @@ def create_app(lifespan: Lifespan[FastAPI] | None = None) -> FastAPI: app.include_router(ewbi_service_router) app.include_router(internal_federation_router) app.include_router(internal_partner_federations_router) + app.include_router(internal_partners_router) return app diff --git a/src/federation_manager/migrations/versions/0003_partner_name.py b/src/federation_manager/migrations/versions/0003_partner_name.py new file mode 100644 index 0000000..f38a0e3 --- /dev/null +++ b/src/federation_manager/migrations/versions/0003_partner_name.py @@ -0,0 +1,31 @@ +"""Add partner_ops.name for partner registration + +The persistence model declares name as NOT NULL, but the baseline carried only the +columns the inbound path read. Rows that predate this revision were inserted by hand, +so they are backfilled with their mcc_mnc, which is unique and non-null, before the +column becomes NOT NULL. + +Revision ID: 0003_partner_name +Revises: 0002_state_checks +Create Date: 2026-09-22 +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "0003_partner_name" +down_revision: str | None = "0002_state_checks" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.add_column("partner_ops", sa.Column("name", sa.String(255), nullable=True)) + op.execute("UPDATE partner_ops SET name = mcc_mnc WHERE name IS NULL") + op.alter_column("partner_ops", "name", nullable=False) + + +def downgrade() -> None: + op.drop_column("partner_ops", "name") diff --git a/tests/fakes.py b/tests/fakes.py index 1aa54a4..bd5f40e 100644 --- a/tests/fakes.py +++ b/tests/fakes.py @@ -3,7 +3,7 @@ from datetime import datetime, timezone from uuid import UUID from federation_manager.contracts.srm import LocationQueryRequestV1, LocationQueryResponseV1 -from federation_manager.domain.errors import AuthenticationFailed +from federation_manager.domain.errors import AuthenticationFailed, PartnerRegistrationConflict from federation_manager.domain.models import ( Agreement, EwbiResponse, @@ -17,11 +17,10 @@ from federation_manager.domain.models import ( class InMemoryPartnerRepo: def __init__(self, partners: list[PartnerOP]) -> None: - self._by_client_id = {p.oauth2_client_id: p for p in partners} self._by_id = {p.id: p for p in partners} async def find_by_oauth2_client_id(self, client_id: str) -> PartnerOP | None: - return self._by_client_id.get(client_id) + return next((p for p in self._by_id.values() if p.oauth2_client_id == client_id), None) async def find_by_id(self, partner_id: UUID) -> PartnerOP | None: return self._by_id.get(partner_id) @@ -29,6 +28,28 @@ class InMemoryPartnerRepo: async def list_active(self) -> list[PartnerOP]: return [p for p in self._by_id.values() if p.is_active()] + async def find_by_mcc_mnc(self, mcc_mnc: str) -> PartnerOP | None: + return next((p for p in self._by_id.values() if p.mcc_mnc == mcc_mnc), None) + + async def list_all(self) -> list[PartnerOP]: + return list(self._by_id.values()) + + async def add(self, partner: PartnerOP) -> None: + self._check_unique(partner) + self._by_id[partner.id] = partner + + async def update(self, partner: PartnerOP) -> None: + self._check_unique(partner) + self._by_id[partner.id] = partner + + def _check_unique(self, partner: PartnerOP) -> None: + for other in self._by_id.values(): + if other.id == partner.id: + continue + for field in ("mcc_mnc", "oauth2_client_id"): + if getattr(other, field) == getattr(partner, field): + raise PartnerRegistrationConflict(field) + class FakeJwtValidator: def __init__(self, tokens: dict[str, ValidatedClaims]) -> None: diff --git a/tests/integration/test_federation_context_repo.py b/tests/integration/test_federation_context_repo.py index e810c3f..29354f9 100644 --- a/tests/integration/test_federation_context_repo.py +++ b/tests/integration/test_federation_context_repo.py @@ -42,6 +42,7 @@ async def test_context_repo_round_trips_through_postgres() -> None: await session.execute( insert(partner_ops).values( id=pid, + name="Test partner", mcc_mnc=uuid4().hex[:10], oauth2_client_id=f"partner-{uuid4().hex[:8]}", base_url="https://partner.example", diff --git a/tests/integration/test_fm_srm_loop.py b/tests/integration/test_fm_srm_loop.py index 48c0d94..dfd7b00 100644 --- a/tests/integration/test_fm_srm_loop.py +++ b/tests/integration/test_fm_srm_loop.py @@ -114,6 +114,7 @@ async def _seed(specification_id: UUID, partner_id: UUID, zone_id: UUID) -> None await session.execute( insert(partner_ops).values( id=partner_id, + name="Test partner", mcc_mnc=uuid4().hex[:10], oauth2_client_id=CLIENT_ID, base_url="http://127.0.0.1:9", diff --git a/tests/integration/test_outbound_repos.py b/tests/integration/test_outbound_repos.py index 6c49391..87d11d7 100644 --- a/tests/integration/test_outbound_repos.py +++ b/tests/integration/test_outbound_repos.py @@ -41,6 +41,7 @@ async def _insert_partner(session: object, partner_id: object) -> str: await session.execute( # type: ignore[attr-defined] insert(partner_ops).values( id=partner_id, + name="Test partner", mcc_mnc=uuid4().hex[:10], oauth2_client_id=client_id, base_url="https://partner.example", diff --git a/tests/integration/test_partner_registry_repo.py b/tests/integration/test_partner_registry_repo.py new file mode 100644 index 0000000..caaa831 --- /dev/null +++ b/tests/integration/test_partner_registry_repo.py @@ -0,0 +1,111 @@ +import os +from dataclasses import replace +from uuid import uuid4 + +import pytest +from sqlalchemy import delete + +from federation_manager.adapters.database.core import ( + build_engine, + build_session_maker, + run_migrations, +) +from federation_manager.adapters.database.partner_repo import PostgresPartnerRepo +from federation_manager.adapters.database.tables import partner_ops +from federation_manager.application.partners import PartnerRegistration, PartnerRegistryService +from federation_manager.domain.errors import PartnerRegistrationConflict +from federation_manager.domain.models import PartnerOP + +pytestmark = pytest.mark.integration + +URL = os.getenv("FM_POSTGRES_URL", "postgresql+asyncpg://fm:fm@localhost:5433/fm_db") + + +def _partner() -> PartnerOP: + suffix = uuid4().hex[:6] + return PartnerOP( + id=uuid4(), + name=f"Partner {suffix}", + mcc_mnc=f"9{int(suffix, 16) % 100:02d}-{int(suffix, 16) % 1000:03d}", + oauth2_client_id=f"partner-{suffix}", + status="pending", + base_url="https://partner.example", + our_client_id="oop-i2cat", + our_client_secret_ref=f"/run/secrets/partners/{suffix}", + token_endpoint="https://partner.example/oauth2/token", + ) + + +async def test_partner_written_and_updated_in_postgres() -> None: + engine = build_engine(URL) + await run_migrations(engine) + session_maker = build_session_maker(engine) + partner = _partner() + + async with session_maker() as session: + repo = PostgresPartnerRepo(session) + await repo.add(partner) + assert await repo.find_by_id(partner.id) == partner + assert await repo.find_by_mcc_mnc(partner.mcc_mnc) == partner + + activated = replace(partner, status="active", name="Renamed") + await repo.update(activated) + assert await repo.find_by_id(partner.id) == activated + assert activated in await repo.list_active() + assert activated in await repo.list_all() + + await session.execute(delete(partner_ops).where(partner_ops.c.id == partner.id)) + await session.commit() + + await engine.dispose() + + +async def test_unique_index_race_surfaces_as_conflict() -> None: + engine = build_engine(URL) + await run_migrations(engine) + session_maker = build_session_maker(engine) + partner = _partner() + + async with session_maker() as session: + repo = PostgresPartnerRepo(session) + await repo.add(partner) + + # the service's pre-check would catch this; writing directly models losing the race + with pytest.raises(PartnerRegistrationConflict) as raised: + await repo.add(replace(partner, id=uuid4(), oauth2_client_id="someone-else")) + assert raised.value.field == "mcc_mnc" + + # the session is usable again after the rolled-back write + assert await repo.find_by_id(partner.id) == partner + + await session.execute(delete(partner_ops).where(partner_ops.c.id == partner.id)) + await session.commit() + + await engine.dispose() + + +async def test_registering_an_existing_partner_from_another_session_conflicts() -> None: + engine = build_engine(URL) + await run_migrations(engine) + session_maker = build_session_maker(engine) + partner = _partner() + registration = PartnerRegistration( + name=partner.name, + mcc_mnc=partner.mcc_mnc, + oauth2_client_id=partner.oauth2_client_id, + base_url="https://partner.example", + our_client_id=partner.our_client_id, + our_client_secret_ref=partner.our_client_secret_ref, + token_endpoint=partner.token_endpoint, + ) + + async with session_maker() as first, session_maker() as second: + created = await PartnerRegistryService(PostgresPartnerRepo(first)).register(registration) + + with pytest.raises(PartnerRegistrationConflict): + await PartnerRegistryService(PostgresPartnerRepo(second)).register(registration) + + await first.execute(delete(partner_ops).where(partner_ops.c.id == created.id)) + await first.commit() + + await engine.dispose() diff --git a/tests/integration/test_partner_repo.py b/tests/integration/test_partner_repo.py index e92a476..fa65581 100644 --- a/tests/integration/test_partner_repo.py +++ b/tests/integration/test_partner_repo.py @@ -38,6 +38,7 @@ async def test_repo_reads_partner_from_postgres() -> None: await session.execute( insert(partner_ops).values( id=uuid4(), + name="Test partner", mcc_mnc=uuid4().hex[:10], oauth2_client_id=client_id, base_url="https://partner.example", @@ -82,6 +83,7 @@ async def test_suspended_partner_rejected_against_postgres() -> None: await session.execute( insert(partner_ops).values( id=uuid4(), + name="Test partner", mcc_mnc=uuid4().hex[:10], oauth2_client_id=client_id, base_url="https://partner.example", diff --git a/tests/integration/test_row_timestamps.py b/tests/integration/test_row_timestamps.py index 06a68f1..2df23f5 100644 --- a/tests/integration/test_row_timestamps.py +++ b/tests/integration/test_row_timestamps.py @@ -29,6 +29,7 @@ async def test_updated_at_moves_when_a_row_changes() -> None: await session.execute( insert(partner_ops).values( id=partner_id, + name="Test partner", mcc_mnc=uuid4().hex[:10], oauth2_client_id=client_id, base_url="https://partner.example", diff --git a/tests/integration/test_two_stack_federation.py b/tests/integration/test_two_stack_federation.py index 47253e8..f52aaec 100644 --- a/tests/integration/test_two_stack_federation.py +++ b/tests/integration/test_two_stack_federation.py @@ -154,6 +154,7 @@ def stacks(tmp_path: Path) -> Iterator[Stacks]: DB_A, { "id": partner_b_id, + "name": "Partner B", "mcc_mnc": mcc_mnc_b, "oauth2_client_id": CLIENT_B, "base_url": f"http://127.0.0.1:{port_b}", @@ -167,6 +168,7 @@ def stacks(tmp_path: Path) -> Iterator[Stacks]: DB_B, { "id": partner_a_id, + "name": "Partner A", "mcc_mnc": mcc_mnc_a, "oauth2_client_id": CLIENT_A, "base_url": f"http://127.0.0.1:{port_a}", diff --git a/tests/test_partner_registry.py b/tests/test_partner_registry.py new file mode 100644 index 0000000..9c67a16 --- /dev/null +++ b/tests/test_partner_registry.py @@ -0,0 +1,183 @@ +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from typing import Any +from uuid import uuid4 + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from federation_manager.application.partners import PartnerRegistration, PartnerRegistryService +from federation_manager.dependencies import get_partner_repo +from federation_manager.domain.errors import PartnerRegistrationConflict +from federation_manager.domain.models import PartnerOP +from federation_manager.main import create_app +from tests.fakes import InMemoryPartnerRepo + +PARTNERS = "/internal/partners" + + +@asynccontextmanager +async def _no_infra(app: FastAPI) -> AsyncIterator[None]: + yield + + +def _registration(**overrides: Any) -> dict[str, Any]: + body: dict[str, Any] = { + "name": "Partner A", + "mcc_mnc": "208-01", + "oauth2_client_id": "partner-a", + "base_url": "https://partner-a.example", + "our_client_id": "oop-i2cat", + "our_client_secret_ref": "/run/secrets/partners/partner-a", + "token_endpoint": "https://partner-a.example/oauth2/token", + } + return {**body, **overrides} + + +@pytest.fixture +def client() -> TestClient: + repo = InMemoryPartnerRepo([]) + app = create_app(lifespan=_no_infra) + app.dependency_overrides[get_partner_repo] = lambda: repo + return TestClient(app) + + +def _register(client: TestClient, **overrides: Any) -> dict[str, Any]: + response = client.post(PARTNERS, json=_registration(**overrides)) + assert response.status_code == 201, response.text + body: dict[str, Any] = response.json() + return body + + +def test_registers_partner_as_pending_by_default(client: TestClient) -> None: + partner = _register(client) + + assert partner["status"] == "pending" + assert partner["name"] == "Partner A" + # stored exactly as sent: no trailing slash appended by URL normalisation + assert partner["base_url"] == "https://partner-a.example" + assert client.get(f"{PARTNERS}/{partner['id']}").json() == partner + + +@pytest.mark.parametrize( + ("overrides", "field"), + [ + ({}, "mcc_mnc"), # the very same registration again + ({"oauth2_client_id": "partner-b"}, "mcc_mnc"), + ({"mcc_mnc": "208-02"}, "oauth2_client_id"), + ], +) +def test_already_registered_partner_conflicts( + client: TestClient, overrides: dict[str, Any], field: str +) -> None: + _register(client) + + response = client.post(PARTNERS, json=_registration(**overrides)) + + assert response.status_code == 409 + assert response.json()["type"] == "urn:oop:ewbi:error:partner-registration-conflict" + assert field in response.json()["detail"] + assert len(client.get(PARTNERS).json()) == 1 + + +def test_partial_outbound_credentials_rejected(client: TestClient) -> None: + response = client.post(PARTNERS, json=_registration(token_endpoint=None)) + + assert response.status_code == 422 + assert response.json()["type"] == "urn:oop:ewbi:error:partner-credentials-incomplete" + + +def test_partner_without_outbound_credentials_is_valid(client: TestClient) -> None: + partner = _register(client, our_client_id=None, our_client_secret_ref=None, token_endpoint=None) + + assert partner["token_endpoint"] is None + + +@pytest.mark.parametrize( + "overrides", + [ + {"our_client_secret_ref": "s3cr3t-value"}, # an inline secret, not a path + {"base_url": "partner-a.example"}, + {"mcc_mnc": "20801x"}, + {"status": "suspended"}, # registration only starts a partner pending or active + {"name": ""}, + {"unexpected": "field"}, + ], +) +def test_invalid_registration_rejected(client: TestClient, overrides: dict[str, Any]) -> None: + assert client.post(PARTNERS, json=_registration(**overrides)).status_code == 422 + + +def test_unknown_partner_is_404(client: TestClient) -> None: + response = client.get(f"{PARTNERS}/{uuid4()}") + + assert response.status_code == 404 + assert client.patch(f"{PARTNERS}/{uuid4()}", json={"name": "x"}).status_code == 404 + + +def test_patch_updates_only_given_fields(client: TestClient) -> None: + partner = _register(client) + + response = client.patch( + f"{PARTNERS}/{partner['id']}", + json={"status": "active", "base_url": "https://new.partner-a.example"}, + ) + + assert response.status_code == 200 + assert response.json() == { + **partner, + "status": "active", + "base_url": "https://new.partner-a.example", + } + + +def test_patch_clears_outbound_credentials_only_as_a_set(client: TestClient) -> None: + partner = _register(client) + url = f"{PARTNERS}/{partner['id']}" + + assert client.patch(url, json={"token_endpoint": None}).status_code == 422 + cleared = client.patch( + url, json={"our_client_id": None, "our_client_secret_ref": None, "token_endpoint": None} + ) + assert cleared.status_code == 200 + assert cleared.json()["our_client_id"] is None + + +@pytest.mark.parametrize("body", [{"name": None}, {"status": None}, {"mcc_mnc": "208-02"}]) +def test_patch_rejects_clearing_required_or_changing_identity( + client: TestClient, body: dict[str, Any] +) -> None: + partner = _register(client) + + assert client.patch(f"{PARTNERS}/{partner['id']}", json=body).status_code == 422 + + +def test_patch_to_a_client_id_held_by_another_partner_conflicts(client: TestClient) -> None: + _register(client) + other = _register(client, mcc_mnc="208-02", oauth2_client_id="partner-b") + + response = client.patch(f"{PARTNERS}/{other['id']}", json={"oauth2_client_id": "partner-a"}) + + assert response.status_code == 409 + + +class _RacingRepo(InMemoryPartnerRepo): + """A concurrent registration lands between our pre-check read and our insert.""" + + def __init__(self, winner: PartnerOP) -> None: + super().__init__([winner]) + + async def find_by_mcc_mnc(self, mcc_mnc: str) -> PartnerOP | None: + return None + + async def find_by_oauth2_client_id(self, client_id: str) -> PartnerOP | None: + return None + + +async def test_losing_the_race_to_a_concurrent_registration_conflicts() -> None: + winner = PartnerOP(id=uuid4(), status="pending", **_registration()) + service = PartnerRegistryService(_RacingRepo(winner)) + + with pytest.raises(PartnerRegistrationConflict): + await service.register(PartnerRegistration(**_registration())) -- GitLab