diff --git a/.env.example b/.env.example index 06995e9d47fda8e74b7ef3ed8a118934ddeef97a..ca26bb544337076f370e6bcbee370ab337cf163b 100644 --- a/.env.example +++ b/.env.example @@ -5,3 +5,8 @@ APP_VERSION="1.5.0" POSTGRES_SETTINGS__URL = "postgresql+asyncpg://postgres:postgres@localhost:5432/srm" POSTGRES_SETTINGS__ECHO = true POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP = true + +NATS_SETTINGS__URL = "nats://localhost:4222" +NATS_SETTINGS__CONNECT_TIMEOUT = 10 +NATS_SETTINGS__MAX_RECONNECT_ATTEMPTS = 3 +NATS_SETTINGS__DRAIN_TIMEOUT = 30 diff --git a/.gitlab-ci.yml b/.gitlab-ci.yml index 4c98851a768bf403cf50a16d623a5ff31e492d07..1e9ac238ba6a47edb238b62ef2d6a8ac8ae3172d 100644 --- a/.gitlab-ci.yml +++ b/.gitlab-ci.yml @@ -1,6 +1,4 @@ default: - tags: - - docker image: python:3.12-slim cache: paths: @@ -19,7 +17,6 @@ stages: - format - test - build - - push variables: UV_CACHE_DIR: "$CI_PROJECT_DIR/.cache/uv" @@ -57,28 +54,19 @@ test: build: stage: build - tags: - - shell - before_script: - - docker info - script: - - export TEST_IMAGE_TAG="ci-${CI_COMMIT_REF_SLUG}-${CI_COMMIT_SHORT_SHA}" - - docker build --network=host -t "$CI_REGISTRY_IMAGE:$TEST_IMAGE_TAG" . - rules: - - if: '$CI_COMMIT_BRANCH' - -push: - stage: push - tags: - - shell - needs: - - build + image: docker:cli + variables: + DOCKER_HOST: tcp://docker:2375 + DOCKER_TLS_CERTDIR: "" + services: + - docker:dind before_script: - docker info script: - export TEST_IMAGE_TAG="ci-${CI_COMMIT_REF_SLUG}-${CI_COMMIT_SHORT_SHA}" - echo "$CI_REGISTRY_PASSWORD" | docker login -u "$CI_REGISTRY_USER" "$CI_REGISTRY" --password-stdin + - docker build --network=host -t "$CI_REGISTRY_IMAGE:$TEST_IMAGE_TAG" . - docker push "$CI_REGISTRY_IMAGE:$TEST_IMAGE_TAG" - docker logout "$CI_REGISTRY" rules: - - if: '$CI_COMMIT_BRANCH' + - if: '$CI_COMMIT_BRANCH == "main" || $CI_COMMIT_BRANCH == "develop"' diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index df54d1b50493cc434531c676bc5d6379c784f5c2..8220cacc1d0be84b7a9665b6ae9b79ff7bd5893f 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -27,6 +27,7 @@ repos: - "docker>=7.0.0" - "fastapi[standard]>=0.135.1" - "httpx>=0.27.0" + - "nats-py>=2.10.0" - "pydantic-settings>=2.13.1" - "pytest>=9.0.2" - "pytest-asyncio>=0.24" diff --git a/pyproject.toml b/pyproject.toml index c99d698f3411a44f9e4335955c8c6788d6d9d5fe..a7a3783e663b2993575a4ac1b733f9ef65956564 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,7 +20,8 @@ dependencies = [ "structlog>=25.5.0", "sunrise6g-opensdk==2.0.0", "sqlalchemy>=2.0.48", - "asyncpg>=0.31.0" + "asyncpg>=0.31.0", + "nats-py>=2.10.0", ] [project.optional-dependencies] diff --git a/src/srm/adapters/databus/nats_connection_manager.py b/src/srm/adapters/databus/nats_connection_manager.py new file mode 100644 index 0000000000000000000000000000000000000000..47693e0af5135150b1f3b45caeae2bfeac53d5fd --- /dev/null +++ b/src/srm/adapters/databus/nats_connection_manager.py @@ -0,0 +1,76 @@ +import asyncio + +import nats +import structlog +from nats.aio.client import Client + +from srm.config import NatsSettings + +logger: structlog.BoundLogger = structlog.get_logger(__name__) + + +async def init_databus_manager(settings: NatsSettings) -> "NatsConnectionManager": + connection_manager: NatsConnectionManager = NatsConnectionManager(settings=settings) + await connection_manager.connect() + return connection_manager + + +class NatsConnectionManager: + def __init__(self, settings: NatsSettings) -> None: + self._settings = settings + self._client: Client | None = None + self._connect_lock = asyncio.Lock() + + @property + def is_connected(self) -> bool: + return self._client is not None and self._client.is_connected + + @property + def client(self) -> Client: + if self._client is None: + raise RuntimeError("NATS client is not connected") + return self._client + + async def connect(self) -> None: + async def _on_error(e: Exception) -> None: + logger.error("nats_error", error=str(e)) + + async def _on_disconnect() -> None: + logger.warning("nats_disconnected", url=self._settings.url) + + async def _on_reconnect() -> None: + logger.info("nats_reconnected", url=self._settings.url) + + async with self._connect_lock: + if self.is_connected: + logger.info("nats_already_connected") + return + + try: + self._client = await nats.connect( + servers=[self._settings.url], + connect_timeout=self._settings.connect_timeout, + max_reconnect_attempts=self._settings.max_reconnect_attempts, + drain_timeout=self._settings.drain_timeout, + error_cb=_on_error, + disconnected_cb=_on_disconnect, + reconnected_cb=_on_reconnect, + ) + except Exception as e: + logger.error("nats_error", error=str(e)) + raise + + async def close(self) -> None: + if self._client is None: + return + + client = self._client + try: + # drain() is bounded by the drain_timeout passed to connect(); on timeout it + # reports through error_cb and closes anyway, so shutdown cannot hang here. + await client.drain() + except Exception as e: + logger.warning("nats_drain_failed", error=str(e)) + await client.close() + finally: + self._client = None diff --git a/src/srm/adapters/databus/nats_publisher.py b/src/srm/adapters/databus/nats_publisher.py new file mode 100644 index 0000000000000000000000000000000000000000..c03ec677e45aa6e385acb956f259d9ecd6ddcaaf --- /dev/null +++ b/src/srm/adapters/databus/nats_publisher.py @@ -0,0 +1,20 @@ +import json +from typing import Any + +from srm.adapters.databus.nats_connection_manager import NatsConnectionManager +from srm.domain.ports.databus.publisher import DataBusPublisher + + +class NatsPublisher(DataBusPublisher): + def __init__(self, connection_manager: NatsConnectionManager) -> None: + self._connection_manager = connection_manager + + async def publish( + self, + subject: str, + payload: dict[str, Any], + headers: dict[str, str] | None = None, + ) -> None: + + body = json.dumps(payload).encode("utf-8") + await self._connection_manager.client.publish(subject, body, headers=headers) diff --git a/src/srm/api/databus/.gitkeep b/src/srm/api/databus/.gitkeep deleted file mode 100644 index e69de29bb2d1d6434b8b29ae775ad8c2e48c5391..0000000000000000000000000000000000000000 diff --git a/src/srm/api/databus/nats_subscriber.py b/src/srm/api/databus/nats_subscriber.py new file mode 100644 index 0000000000000000000000000000000000000000..f1d2d705251ecd320b05265cb419df20d06f5966 --- /dev/null +++ b/src/srm/api/databus/nats_subscriber.py @@ -0,0 +1,69 @@ +from collections.abc import Awaitable, Callable + +import structlog +from nats.aio.msg import Msg +from nats.aio.subscription import Subscription + +from srm.adapters.databus.nats_connection_manager import NatsConnectionManager +from srm.api.databus.schemas import InboundMessage + +logger: structlog.BoundLogger = structlog.get_logger(__name__) + +COMMAND_SUBJECTS = [ + "command.srm.service.deploy", + "command.srm.service.scale", + "command.srm.service.terminate", + "command.srm.network.capability.activate", + "command.srm.network.capability.update", + "command.srm.network.capability.deactivate", +] + + +async def _noop_router(message: InboundMessage) -> None: + # TODO: the real router must validate in two stages (interface-contract.md §A, ADR-0032). + logger.info("Message Received", subject=message.subject, data=message.payload) + return None + + +async def subscribe_to_subjects( + connection_manager: NatsConnectionManager, +) -> list["NatsSubscriber"]: + subscribers = [ + NatsSubscriber(connection_manager=connection_manager, subject=subject, router=_noop_router) + for subject in COMMAND_SUBJECTS + ] + + for subscriber in subscribers: + await subscriber.start() + + return subscribers + + +class NatsSubscriber: + def __init__( + self, + connection_manager: NatsConnectionManager, + subject: str, + router: Callable[[InboundMessage], Awaitable[None]], + ) -> None: + self._connection_manager = connection_manager + self._subject = subject + self._router = router + self._subscription: Subscription | None = None + + async def start(self) -> None: + self._subscription = await self._connection_manager.client.subscribe( + self._subject, + cb=self._handle_message, + ) + + async def _handle_message(self, msg: Msg) -> None: + inbound_message = InboundMessage( + subject=msg.subject, + payload=msg.data, + headers=dict(msg.headers) if msg.headers is not None else {}, + ) + try: + await self._router(inbound_message) + except Exception: + logger.exception("databus_handler_failed", subject=msg.subject) diff --git a/src/srm/api/databus/schemas.py b/src/srm/api/databus/schemas.py new file mode 100644 index 0000000000000000000000000000000000000000..2072cca213fd137462f154701c915c7a0667440f --- /dev/null +++ b/src/srm/api/databus/schemas.py @@ -0,0 +1,141 @@ +from datetime import datetime +from typing import Literal +from uuid import UUID + +from pydantic import BaseModel, Field, model_validator + +from srm.domain.models.canonical_parameters.parameters import ( + CapabilityParameters, + CapabilityTarget, +) + + +class InboundMessage(BaseModel): + subject: str + payload: bytes + headers: dict[str, str] + + +class CommandEnvelopeV1(BaseModel): + schema_version: Literal["1.0"] + operation_id: UUID + correlation_id: str + requested_at: datetime + app_provider_id: str + federation_partner_ref: str | None = None + source: Literal["nbi_camara", "nbi_tmf", "operator_portal", "federation"] + + @model_validator(mode="after") + def validate_federation_context(self) -> "CommandEnvelopeV1": + if self.source == "federation" and self.federation_partner_ref is None: + raise ValueError("federation_partner_ref is required when source=federation") + + return self + + +class PlacementConstraintsV1(BaseModel): + model_config = {"extra": "allow"} + + +class DeployPayloadV1(BaseModel): + instance_name: str | None = None + placement_constraints: PlacementConstraintsV1 | None = None + + +class DeployTargetV1(BaseModel): + app_instance_id: UUID + zone_id: UUID | None = None + domain_id: UUID | None = None + + @model_validator(mode="after") + def validate_pin_shape(self) -> "DeployTargetV1": + if self.domain_id is not None and self.zone_id is None: + raise ValueError("domain_id requires zone_id: a domain pin must name its zone") + + return self + + +class SrmServiceDeployV1(CommandEnvelopeV1): + service_specification_id: UUID + targets: list[DeployTargetV1] = Field(min_length=1) + deploy: DeployPayloadV1 + + +class ScalePayloadV1(BaseModel): + replicas: int = Field(ge=0) + + +class SrmServiceScaleV1(CommandEnvelopeV1): + service_instance_id: UUID + service_specification_id: UUID | None = None + scale: ScalePayloadV1 + + +class TerminatePayloadV1(BaseModel): + grace_period_seconds: int = Field(default=0, ge=0) + + +class SrmServiceTerminateV1(CommandEnvelopeV1): + service_instance_id: UUID + service_specification_id: UUID | None = None + terminate: TerminatePayloadV1 + + +class NetworkCapabilityPayloadV1(BaseModel): + capability_type: str + target: CapabilityTarget + profile_ref: str | None = None + parameters: CapabilityParameters + + +class SrmNetworkCapabilityActivateV1(CommandEnvelopeV1): + service_specification_id: UUID + zone_id: UUID | None = None + domain_id: UUID | None = None + network_capability: NetworkCapabilityPayloadV1 + + @model_validator(mode="after") + def validate_pin_shape(self) -> "SrmNetworkCapabilityActivateV1": + if self.domain_id is not None and self.zone_id is None: + raise ValueError("domain_id requires zone_id: a domain pin must name its zone") + + return self + + +class NetworkCapabilityRealizationRefV1(BaseModel): + capability_type: str + external_ref: str | None = None + service_instance_id: UUID | None = None + + @model_validator(mode="after") + def validate_reference_shape(self) -> "NetworkCapabilityRealizationRefV1": + has_external_ref = self.external_ref is not None + has_service_instance_id = self.service_instance_id is not None + + if has_external_ref == has_service_instance_id: + raise ValueError( + "network capability realization must be identified by exactly one of " + "external_ref or service_instance_id" + ) + + return self + + +class NetworkCapabilityUpdatePayloadV1(NetworkCapabilityRealizationRefV1): + target: CapabilityTarget | None = None + profile_ref: str | None = None + parameters: CapabilityParameters + + +class SrmNetworkCapabilityUpdateV1(CommandEnvelopeV1): + service_specification_id: UUID | None = None + network_capability: NetworkCapabilityUpdatePayloadV1 + + +class NetworkCapabilityDeactivatePayloadV1(NetworkCapabilityRealizationRefV1): + grace_period_seconds: int = Field(default=0, ge=0) + + +class SrmNetworkCapabilityDeactivateV1(CommandEnvelopeV1): + service_specification_id: UUID | None = None + network_capability: NetworkCapabilityDeactivatePayloadV1 diff --git a/src/srm/api/health.py b/src/srm/api/health.py index d34ee2f0204cb0d88732d2eba8b84cfca26bee1e..f3f441295c566ada3c876a2a9ee3b9fe6f72c8b9 100644 --- a/src/srm/api/health.py +++ b/src/srm/api/health.py @@ -13,10 +13,15 @@ async def liveness() -> bool: @health_router.get("/health/readyz") async def readiness(request: Request) -> bool: + app_state = get_app_state(request) + try: - async with get_app_state(request).db_engine.connect() as connection: + async with app_state.db_engine.connect() as connection: await connection.execute(text("SELECT 1")) except Exception as exc: raise HTTPException(status_code=503, detail="Postgres not ready") from exc + if not app_state.databus_connection_manager.is_connected: + raise HTTPException(status_code=503, detail="NATS not ready") + return True diff --git a/src/srm/app_state.py b/src/srm/app_state.py index a21baa6340beb330781d024a994b46c7f5f22b4c..799d9cbc14eca4aa990e15d8a86c3b6c04cf4445 100644 --- a/src/srm/app_state.py +++ b/src/srm/app_state.py @@ -2,7 +2,12 @@ from typing import Protocol from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker +from srm.adapters.databus.nats_connection_manager import NatsConnectionManager +from srm.api.databus.nats_subscriber import NatsSubscriber + class AppState(Protocol): db_engine: AsyncEngine session_maker: async_sessionmaker[AsyncSession] + databus_connection_manager: NatsConnectionManager + databus_subscribers: list[NatsSubscriber] diff --git a/src/srm/config.py b/src/srm/config.py index 5a75a01938a49c0fb19272f94b26971332127d35..20a4d87a63a9f7703d0fd02b406358492bfaa588 100644 --- a/src/srm/config.py +++ b/src/srm/config.py @@ -29,6 +29,13 @@ class PostgreSQLSettings(BaseModel): create_schema_on_startup: bool +class NatsSettings(BaseModel): + url: str + connect_timeout: int + max_reconnect_attempts: int + drain_timeout: int = 30 + + class Settings(BaseSettings): model_config = SettingsConfigDict(env_file=".env", env_nested_delimiter="__") @@ -37,6 +44,7 @@ class Settings(BaseSettings): app_description: str postgres_settings: PostgreSQLSettings + nats_settings: NatsSettings @lru_cache() diff --git a/src/srm/domain/ports/databus/publisher.py b/src/srm/domain/ports/databus/publisher.py new file mode 100644 index 0000000000000000000000000000000000000000..613f70db017bf0ad413a2c9d325b7504348c42e3 --- /dev/null +++ b/src/srm/domain/ports/databus/publisher.py @@ -0,0 +1,10 @@ +from typing import Any, Protocol + + +class DataBusPublisher(Protocol): + async def publish( + self, + subject: str, + payload: dict[str, Any], + headers: dict[str, str] | None = None, + ) -> None: ... diff --git a/src/srm/main.py b/src/srm/main.py index 6e56ba471d3b0da5015e751da45f4cb5d8a8f27c..4e4e5f7100619460e03de948b996619b41108411 100644 --- a/src/srm/main.py +++ b/src/srm/main.py @@ -17,7 +17,15 @@ from typing import AsyncIterator import structlog from fastapi import FastAPI -from srm.adapters.database.core import build_engine_and_session_maker, schema_initialization +from srm.adapters.database.core import ( + build_engine_and_session_maker, + schema_initialization, +) +from srm.adapters.databus.nats_connection_manager import ( + NatsConnectionManager, + init_databus_manager, +) +from srm.api.databus.nats_subscriber import NatsSubscriber, subscribe_to_subjects from srm.api.health import health_router from srm.api.middlewares.middlewares import register_middlewares from srm.config import Settings, get_settings @@ -46,13 +54,35 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: except Exception as e: logger.error("Database engine init failed!", error=str(e)) raise + try: + databus_manager: NatsConnectionManager = await init_databus_manager( + settings=settings.nats_settings + ) + except Exception as e: + logger.error("Databus connection init failed!", error=str(e)) + await engine.dispose() + raise + + try: + databus_subscribers: list[NatsSubscriber] = await subscribe_to_subjects(databus_manager) + except Exception as e: + logger.error("Databus subscription failed!", error=str(e)) + await databus_manager.close() + await engine.dispose() + raise app.state.db_engine = engine app.state.session_maker = session_maker + app.state.databus_connection_manager = databus_manager + app.state.databus_subscribers = databus_subscribers yield logger.info("Shutting down application") + # Drain the bus first: draining delivers in-flight messages to their handlers, and + # those handlers still need a live DB engine. + await app.state.databus_connection_manager.close() + app.state.databus_subscribers = [] await engine.dispose() diff --git a/src/srm/adapters/databus/.gitkeep b/tests/api/databus/__init__.py similarity index 100% rename from src/srm/adapters/databus/.gitkeep rename to tests/api/databus/__init__.py diff --git a/tests/api/databus/test_nats_subscriber.py b/tests/api/databus/test_nats_subscriber.py new file mode 100644 index 0000000000000000000000000000000000000000..65a97ef29f97a32bb19003aa89764b083cddaede --- /dev/null +++ b/tests/api/databus/test_nats_subscriber.py @@ -0,0 +1,142 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from unittest.mock import AsyncMock, MagicMock + +import pytest +import structlog.testing + +from srm.adapters.databus.nats_connection_manager import NatsConnectionManager +from srm.api.databus.nats_subscriber import ( + COMMAND_SUBJECTS, + NatsSubscriber, + subscribe_to_subjects, +) +from srm.api.databus.schemas import InboundMessage + + +@dataclass +class FakeMsg: + subject: str + data: bytes + headers: dict[str, str] | None = field(default=None) + + +@pytest.fixture +def connection_manager() -> MagicMock: + manager = MagicMock(spec=NatsConnectionManager) + manager.client = AsyncMock() + return manager + + +def make_subscriber( + connection_manager: MagicMock, subject: str = "command.srm.service.deploy" +) -> tuple[NatsSubscriber, AsyncMock]: + router = AsyncMock() + subscriber = NatsSubscriber( + connection_manager=connection_manager, + subject=subject, + router=router, + ) + return subscriber, router + + +async def test_start_subscribes_via_connection_manager_client( + connection_manager: MagicMock, +) -> None: + subscriber, _ = make_subscriber(connection_manager) + + await subscriber.start() + + connection_manager.client.subscribe.assert_awaited_once() + call_kwargs = connection_manager.client.subscribe.call_args + assert call_kwargs.args == ("command.srm.service.deploy",) + assert call_kwargs.kwargs["cb"] == subscriber._handle_message + assert subscriber._subscription is not None + + +async def test_handle_message_builds_inbound_message_and_invokes_router( + connection_manager: MagicMock, +) -> None: + subscriber, router = make_subscriber(connection_manager) + msg = FakeMsg( + subject="command.srm.service.deploy", + data=b'{"operation_id": "abc-123"}', + headers={"x-correlation-id": "corr-1"}, + ) + + await subscriber._handle_message(msg) # type: ignore[arg-type] + + router.assert_awaited_once() + assert router.await_args is not None + (inbound,) = router.await_args.args + assert isinstance(inbound, InboundMessage) + assert inbound.subject == "command.srm.service.deploy" + assert inbound.payload == b'{"operation_id": "abc-123"}' + assert inbound.headers == {"x-correlation-id": "corr-1"} + + +async def test_handle_message_defaults_headers_to_empty_dict_when_none( + connection_manager: MagicMock, +) -> None: + subscriber, router = make_subscriber(connection_manager) + msg = FakeMsg(subject="command.srm.service.deploy", data=b"{}", headers=None) + + await subscriber._handle_message(msg) # type: ignore[arg-type] + + assert router.await_args is not None + (inbound,) = router.await_args.args + assert inbound.headers == {} + + +async def test_handle_message_passes_malformed_payload_through_unparsed( + connection_manager: MagicMock, +) -> None: + subscriber, router = make_subscriber(connection_manager) + msg = FakeMsg(subject="command.srm.service.deploy", data=b'{"operation_id": ') + + await subscriber._handle_message(msg) # type: ignore[arg-type] + + router.assert_awaited_once() + assert router.await_args is not None + (inbound,) = router.await_args.args + assert inbound.payload == b'{"operation_id": ' + + +async def test_handle_message_logs_dropped_command_when_router_fails( + connection_manager: MagicMock, +) -> None: + subscriber, router = make_subscriber(connection_manager) + router.side_effect = RuntimeError("handler blew up") + msg = FakeMsg(subject="command.srm.service.deploy", data=b"{}") + + with structlog.testing.capture_logs() as logs: + await subscriber._handle_message(msg) # type: ignore[arg-type] + + router.assert_awaited_once() + assert [(entry["event"], entry["subject"], entry["log_level"]) for entry in logs] == [ + ("databus_handler_failed", "command.srm.service.deploy", "error") + ] + + +EXPECTED_COMMAND_SUBJECTS = [ + "command.srm.service.deploy", + "command.srm.service.scale", + "command.srm.service.terminate", + "command.srm.network.capability.activate", + "command.srm.network.capability.update", + "command.srm.network.capability.deactivate", +] + + +def test_command_subjects_matches_the_contract() -> None: + assert sorted(COMMAND_SUBJECTS) == sorted(EXPECTED_COMMAND_SUBJECTS) + + +async def test_subscribe_to_subjects_registers_all_command_subjects( + connection_manager: MagicMock, +) -> None: + subscribers = await subscribe_to_subjects(connection_manager) + + assert sorted(sub._subject for sub in subscribers) == sorted(EXPECTED_COMMAND_SUBJECTS) + assert connection_manager.client.subscribe.await_count == len(EXPECTED_COMMAND_SUBJECTS) diff --git a/tests/api/databus/test_schemas.py b/tests/api/databus/test_schemas.py new file mode 100644 index 0000000000000000000000000000000000000000..653ae0da440a2293f0206e91ca302b215aebb823 --- /dev/null +++ b/tests/api/databus/test_schemas.py @@ -0,0 +1,394 @@ +from __future__ import annotations + +from uuid import uuid4 + +import pytest +from pydantic import ValidationError + +from srm.api.databus.schemas import ( + CommandEnvelopeV1, + DeployPayloadV1, + DeployTargetV1, + NetworkCapabilityDeactivatePayloadV1, + NetworkCapabilityPayloadV1, + NetworkCapabilityUpdatePayloadV1, + ScalePayloadV1, + SrmNetworkCapabilityActivateV1, + SrmNetworkCapabilityDeactivateV1, + SrmNetworkCapabilityUpdateV1, + SrmServiceDeployV1, + SrmServiceScaleV1, + SrmServiceTerminateV1, + TerminatePayloadV1, +) +from srm.domain.models.canonical_parameters.parameters import ( + CapabilityParameters, + CapabilityTarget, +) + + +def _envelope(**overrides: object) -> dict[str, object]: + envelope: dict[str, object] = { + "schema_version": "1.0", + "operation_id": str(uuid4()), + "correlation_id": "corr-1", + "requested_at": "2026-07-03T12:00:00+00:00", + "app_provider_id": "provider-1", + "source": "nbi_camara", + } + envelope.update(overrides) + return envelope + + +class TestSchemaVersion: + """§B.1: schema_version is required. §A: an unsupported version is unanswerable + and must be dead-lettered, so it must never parse as if it were 1.0.""" + + def test_envelope_requires_schema_version(self) -> None: + payload = _envelope() + del payload["schema_version"] + + with pytest.raises(ValidationError, match="schema_version"): + CommandEnvelopeV1.model_validate(payload) + + @pytest.mark.parametrize("version", ["2.0", "banana"]) + def test_envelope_rejects_unsupported_schema_version(self, version: str) -> None: + with pytest.raises(ValidationError, match="schema_version"): + CommandEnvelopeV1.model_validate(_envelope(schema_version=version)) + + def test_every_command_inherits_the_rule(self) -> None: + """The rule lives on the envelope, so a v2 producer cannot have its message + silently processed as v1 on any command.srm.* subject.""" + payload = _envelope( + schema_version="2.0", + service_instance_id=str(uuid4()), + scale={"replicas": 3}, + ) + + with pytest.raises(ValidationError, match="schema_version"): + SrmServiceScaleV1.model_validate(payload) + + +class TestSource: + @pytest.mark.parametrize("source", ["nbi_camara", "nbi_tmf", "operator_portal"]) + def test_envelope_accepts_each_non_federation_source(self, source: str) -> None: + assert CommandEnvelopeV1.model_validate(_envelope(source=source)).source == source + + @pytest.mark.parametrize("source", ["nbi_rest", "FEDERATION", ""]) + def test_envelope_rejects_unknown_source(self, source: str) -> None: + with pytest.raises(ValidationError, match="source"): + CommandEnvelopeV1.model_validate(_envelope(source=source)) + + +class TestFederationContext: + def test_federation_source_requires_partner_ref(self) -> None: + with pytest.raises(ValidationError, match="federation_partner_ref is required"): + CommandEnvelopeV1.model_validate(_envelope(source="federation")) + + def test_federation_source_rejects_explicit_null_partner_ref(self) -> None: + with pytest.raises(ValidationError, match="federation_partner_ref is required"): + CommandEnvelopeV1.model_validate( + _envelope(source="federation", federation_partner_ref=None) + ) + + def test_federation_source_accepts_partner_ref(self) -> None: + command = CommandEnvelopeV1.model_validate( + _envelope(source="federation", federation_partner_ref="ptr-OperatorB") + ) + + assert command.federation_partner_ref == "ptr-OperatorB" + + def test_non_federation_source_may_omit_partner_ref(self) -> None: + assert CommandEnvelopeV1.model_validate(_envelope()).federation_partner_ref is None + + def test_every_command_inherits_the_rule(self) -> None: + payload = _envelope( + source="federation", + service_specification_id=str(uuid4()), + targets=[{"app_instance_id": str(uuid4())}], + deploy={}, + ) + + with pytest.raises(ValidationError, match="federation_partner_ref is required"): + SrmServiceDeployV1.model_validate(payload) + + +def test_deploy_v1_parses_minimal_valid_payload() -> None: + payload = _envelope( + service_specification_id=str(uuid4()), + targets=[{"app_instance_id": str(uuid4())}], + deploy={}, + ) + + command = SrmServiceDeployV1.model_validate(payload) + + assert command.schema_version == "1.0" + assert command.deploy == DeployPayloadV1() + assert command.targets[0].zone_id is None + assert command.targets[0].domain_id is None + + +def test_deploy_v1_rejects_empty_targets() -> None: + payload = _envelope( + service_specification_id=str(uuid4()), + targets=[], + deploy={}, + ) + + with pytest.raises(ValidationError): + SrmServiceDeployV1.model_validate(payload) + + +def test_deploy_target_pins_default_to_none() -> None: + target = DeployTargetV1(app_instance_id=uuid4()) + assert target.zone_id is None + assert target.domain_id is None + + +def test_deploy_target_accepts_zone_and_domain_pin() -> None: + zone_id, domain_id = uuid4(), uuid4() + target = DeployTargetV1(app_instance_id=uuid4(), zone_id=zone_id, domain_id=domain_id) + assert target.zone_id == zone_id + assert target.domain_id == domain_id + + +class TestPinShape: + def test_deploy_target_rejects_domain_pin_without_zone(self) -> None: + with pytest.raises(ValidationError, match="domain_id requires zone_id"): + DeployTargetV1(app_instance_id=uuid4(), domain_id=uuid4()) + + def test_deploy_v1_rejects_domain_pin_without_zone_on_any_target(self) -> None: + payload = _envelope( + service_specification_id=str(uuid4()), + targets=[ + {"app_instance_id": str(uuid4()), "zone_id": str(uuid4())}, + {"app_instance_id": str(uuid4()), "domain_id": str(uuid4())}, + ], + deploy={}, + ) + + with pytest.raises(ValidationError, match="domain_id requires zone_id"): + SrmServiceDeployV1.model_validate(payload) + + def test_deploy_target_accepts_explicit_null_pins(self) -> None: + """§B.2: null and omitted optional pin fields are equivalent.""" + target = DeployTargetV1(app_instance_id=uuid4(), zone_id=None, domain_id=None) + assert target.zone_id is None + assert target.domain_id is None + + def test_network_capability_activate_v1_rejects_domain_pin_without_zone(self) -> None: + payload = _envelope( + service_specification_id=str(uuid4()), + domain_id=str(uuid4()), + network_capability={ + "capability_type": "qod_session", + "target": _capability_target().model_dump(), + "parameters": _capability_parameters().model_dump(), + }, + ) + + with pytest.raises(ValidationError, match="domain_id requires zone_id"): + SrmNetworkCapabilityActivateV1.model_validate(payload) + + def test_update_ignores_pins_and_does_not_apply_the_rule(self) -> None: + """§B.6: zone_id/domain_id are ignored when identifying an existing realization, + so a bare domain_id must not be rejected on update.""" + payload = _envelope( + domain_id=str(uuid4()), + network_capability={ + "capability_type": "qod_session", + "external_ref": "sess-123", + "parameters": _capability_parameters().model_dump(), + }, + ) + + command = SrmNetworkCapabilityUpdateV1.model_validate(payload) + + assert command.network_capability.external_ref == "sess-123" + + +def test_scale_payload_requires_replicas() -> None: + with pytest.raises(ValidationError): + ScalePayloadV1() # type: ignore[call-arg] + + +def test_scale_payload_requires_non_negative_replicas() -> None: + with pytest.raises(ValidationError): + ScalePayloadV1(replicas=-1) + + +def test_scale_v1_allows_omitted_service_specification_id() -> None: + payload = _envelope( + service_instance_id=str(uuid4()), + scale={"replicas": 3}, + ) + + command = SrmServiceScaleV1.model_validate(payload) + + assert command.service_specification_id is None + assert command.scale.replicas == 3 + + +def test_terminate_payload_defaults_grace_period_to_zero() -> None: + payload = _envelope( + service_instance_id=str(uuid4()), + terminate={}, + ) + + command = SrmServiceTerminateV1.model_validate(payload) + + assert command.terminate == TerminatePayloadV1(grace_period_seconds=0) + + +def test_terminate_payload_rejects_negative_grace_period() -> None: + with pytest.raises(ValidationError): + TerminatePayloadV1(grace_period_seconds=-1) + + +def _capability_target() -> CapabilityTarget: + return CapabilityTarget() + + +def _capability_parameters() -> CapabilityParameters: + return CapabilityParameters() + + +def test_network_capability_activate_v1_parses_minimal_valid_payload() -> None: + payload = _envelope( + service_specification_id=str(uuid4()), + network_capability={ + "capability_type": "qod_session", + "target": _capability_target().model_dump(), + "parameters": _capability_parameters().model_dump(), + }, + ) + + command = SrmNetworkCapabilityActivateV1.model_validate(payload) + + assert command.zone_id is None + assert command.domain_id is None + assert command.network_capability.capability_type == "qod_session" + + +def test_network_capability_activate_v1_accepts_zone_and_domain_pin() -> None: + zone_id, domain_id = uuid4(), uuid4() + payload = _envelope( + service_specification_id=str(uuid4()), + zone_id=str(zone_id), + domain_id=str(domain_id), + network_capability={ + "capability_type": "qod_session", + "target": _capability_target().model_dump(), + "parameters": _capability_parameters().model_dump(), + }, + ) + + command = SrmNetworkCapabilityActivateV1.model_validate(payload) + + assert command.zone_id == zone_id + assert command.domain_id == domain_id + + +def test_network_capability_payload_requires_capability_type() -> None: + with pytest.raises(ValidationError): + NetworkCapabilityPayloadV1( # type: ignore[call-arg] + target=_capability_target(), + parameters=_capability_parameters(), + ) + + +class TestNetworkCapabilityRealizationRef: + def test_accepts_external_ref_alone(self) -> None: + ref = NetworkCapabilityUpdatePayloadV1( + capability_type="qod_session", + external_ref="sess-123", + parameters=_capability_parameters(), + ) + assert ref.external_ref == "sess-123" + assert ref.service_instance_id is None + + def test_accepts_service_instance_id_alone(self) -> None: + instance_id = uuid4() + ref = NetworkCapabilityUpdatePayloadV1( + capability_type="qod_session", + service_instance_id=instance_id, + parameters=_capability_parameters(), + ) + assert ref.service_instance_id == instance_id + assert ref.external_ref is None + + def test_rejects_neither_present(self) -> None: + with pytest.raises(ValidationError, match="exactly one"): + NetworkCapabilityUpdatePayloadV1( + capability_type="qod_session", + parameters=_capability_parameters(), + ) + + def test_rejects_both_present(self) -> None: + with pytest.raises(ValidationError, match="exactly one"): + NetworkCapabilityUpdatePayloadV1( + capability_type="qod_session", + external_ref="sess-123", + service_instance_id=uuid4(), + parameters=_capability_parameters(), + ) + + def test_requires_capability_type(self) -> None: + with pytest.raises(ValidationError): + NetworkCapabilityUpdatePayloadV1( # type: ignore[call-arg] + external_ref="sess-123", + parameters=_capability_parameters(), + ) + + def test_deactivate_payload_shares_the_same_rule(self) -> None: + with pytest.raises(ValidationError, match="exactly one"): + NetworkCapabilityDeactivatePayloadV1(capability_type="qod_session") + + +def test_network_capability_update_v1_parses_valid_payload() -> None: + payload = _envelope( + network_capability={ + "capability_type": "qod_session", + "external_ref": "sess-123", + "parameters": _capability_parameters().model_dump(), + }, + ) + + command = SrmNetworkCapabilityUpdateV1.model_validate(payload) + + assert command.service_specification_id is None + assert command.network_capability.external_ref == "sess-123" + assert command.network_capability.target is None + assert command.network_capability.profile_ref is None + + +def test_network_capability_update_v1_accepts_retarget_and_profile_ref() -> None: + payload = _envelope( + service_specification_id=str(uuid4()), + network_capability={ + "capability_type": "qod_session", + "external_ref": "sess-123", + "target": _capability_target().model_dump(), + "profile_ref": "profile-1", + "parameters": _capability_parameters().model_dump(), + }, + ) + + command = SrmNetworkCapabilityUpdateV1.model_validate(payload) + + assert command.network_capability.profile_ref == "profile-1" + assert command.network_capability.target is not None + + +def test_network_capability_deactivate_v1_defaults_grace_period_to_zero() -> None: + payload = _envelope( + network_capability={ + "capability_type": "qod_session", + "external_ref": "sess-123", + }, + ) + + command = SrmNetworkCapabilityDeactivateV1.model_validate(payload) + + assert command.service_specification_id is None + assert command.network_capability.grace_period_seconds == 0 diff --git a/tests/api/fakes.py b/tests/api/fakes.py new file mode 100644 index 0000000000000000000000000000000000000000..605bacd54cf246258765fb53368b063b458f914d --- /dev/null +++ b/tests/api/fakes.py @@ -0,0 +1,27 @@ +from types import TracebackType + + +class FakeConnection: + async def __aenter__(self) -> "FakeConnection": + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + traceback: TracebackType | None, + ) -> None: + return None + + async def execute(self, statement: object) -> None: + return None + + +class FakeEngine: + def connect(self) -> FakeConnection: + return FakeConnection() + + +class FakeDatabusConnectionManager: + def __init__(self, *, is_connected: bool) -> None: + self.is_connected = is_connected diff --git a/tests/api/test_app.py b/tests/api/test_app.py index d46ee3f968e08393967a97cb927d50f30d59074c..8d302e8fda368ab9bf1c078cf753711fb5f74226 100644 --- a/tests/api/test_app.py +++ b/tests/api/test_app.py @@ -1,6 +1,11 @@ +from collections.abc import Iterator +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest from fastapi import FastAPI from srm.config import get_settings +from srm.main import lifespan def test_create_app_uses_settings(app: FastAPI) -> None: @@ -8,3 +13,93 @@ def test_create_app_uses_settings(app: FastAPI) -> None: assert app.title == settings.app_name assert app.version == settings.app_version assert app.description == settings.app_description + + +@pytest.fixture +def stub_database() -> Iterator[AsyncMock]: + """Lifespan builds the DB engine before touching NATS; stubbing it keeps these tests + focused on the DataBus stage and off Docker.""" + engine = AsyncMock() + with ( + patch( + "srm.main.build_engine_and_session_maker", + new_callable=AsyncMock, + return_value=(engine, MagicMock()), + ), + patch("srm.main.schema_initialization", new_callable=AsyncMock), + ): + yield engine + + +class TestLifespanDatabusInit: + async def test_startup_fails_when_databus_connection_fails( + self, app: FastAPI, stub_database: AsyncMock + ) -> None: + with patch( + "srm.main.init_databus_manager", + new_callable=AsyncMock, + side_effect=OSError("nats unavailable"), + ): + with pytest.raises(OSError, match="nats unavailable"): + async with lifespan(app): + pass + + stub_database.dispose.assert_awaited_once() + + async def test_startup_fails_when_subject_subscription_fails( + self, app: FastAPI, stub_database: AsyncMock + ) -> None: + manager = AsyncMock() + + with ( + patch("srm.main.init_databus_manager", new_callable=AsyncMock, return_value=manager), + patch( + "srm.main.subscribe_to_subjects", + new_callable=AsyncMock, + side_effect=OSError("subscription refused"), + ), + ): + with pytest.raises(OSError, match="subscription refused"): + async with lifespan(app): + pass + + manager.close.assert_awaited_once() + stub_database.dispose.assert_awaited_once() + + async def test_successful_startup_exposes_databus_state( + self, app: FastAPI, stub_database: AsyncMock + ) -> None: + manager, subscribers = AsyncMock(), [MagicMock(), MagicMock()] + + with ( + patch("srm.main.init_databus_manager", new_callable=AsyncMock, return_value=manager), + patch( + "srm.main.subscribe_to_subjects", + new_callable=AsyncMock, + return_value=subscribers, + ), + ): + async with lifespan(app): + assert app.state.databus_connection_manager is manager + assert app.state.databus_subscribers == subscribers + + manager.close.assert_awaited_once() + assert app.state.databus_subscribers == [] + + async def test_shutdown_drains_databus_before_disposing_the_engine( + self, app: FastAPI, stub_database: AsyncMock + ) -> None: + """Draining delivers in-flight messages to handlers, which still need the engine.""" + order: list[str] = [] + manager = AsyncMock() + manager.close.side_effect = lambda: order.append("databus_closed") + stub_database.dispose.side_effect = lambda: order.append("engine_disposed") + + with ( + patch("srm.main.init_databus_manager", new_callable=AsyncMock, return_value=manager), + patch("srm.main.subscribe_to_subjects", new_callable=AsyncMock, return_value=[]), + ): + async with lifespan(app): + pass + + assert order == ["databus_closed", "engine_disposed"] diff --git a/tests/api/test_health.py b/tests/api/test_health.py index e079c355a8b39b4b6b1772988fb61ec570ee2481..c25e23832c090fbd8a694056d3d3fad9261c9b4a 100644 --- a/tests/api/test_health.py +++ b/tests/api/test_health.py @@ -1,4 +1,17 @@ -from httpx import AsyncClient +import pytest +from fastapi import FastAPI +from httpx import ASGITransport, AsyncClient + +from srm.api.health import health_router +from tests.api.fakes import FakeDatabusConnectionManager, FakeEngine + + +def make_health_app(*, nats_connected: bool) -> FastAPI: + app = FastAPI() + app.include_router(health_router) + app.state.db_engine = FakeEngine() + app.state.databus_connection_manager = FakeDatabusConnectionManager(is_connected=nats_connected) + return app async def test_livez_returns_200(client: AsyncClient) -> None: @@ -7,7 +20,28 @@ async def test_livez_returns_200(client: AsyncClient) -> None: assert response.json() is True +@pytest.mark.integration async def test_readyz_returns_200(client_with_db: AsyncClient) -> None: response = await client_with_db.get("/health/readyz") assert response.status_code == 200 assert response.json() is True + + +async def test_readyz_returns_200_when_nats_connected() -> None: + app = make_health_app(nats_connected=True) + + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as c: + response = await c.get("/health/readyz") + + assert response.status_code == 200 + assert response.json() is True + + +async def test_readyz_returns_503_when_nats_disconnected() -> None: + app = make_health_app(nats_connected=False) + + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as c: + response = await c.get("/health/readyz") + + assert response.status_code == 503 + assert response.json() == {"detail": "NATS not ready"} diff --git a/tests/conftest.py b/tests/conftest.py index 4206cf1265dbd45f46643ee9fa825c934f441911..dbc6a8676ff8e683e90c69c249b819c59a2779f7 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -8,6 +8,7 @@ from fastapi import FastAPI from httpx import ASGITransport, AsyncClient from sqlalchemy import text from sqlalchemy.ext.asyncio import create_async_engine +from testcontainers.nats import NatsContainer from testcontainers.postgres import PostgresContainer from srm.adapters.database.sql import get_metadata @@ -20,6 +21,9 @@ TEST_SETTINGS_ENV = { "APP_DESCRIPTION": "Test SRM", "POSTGRES_SETTINGS__ECHO": "true", "POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP": "true", + "NATS_SETTINGS__URL": "nats://localhost:4222", + "NATS_SETTINGS__CONNECT_TIMEOUT": "10", + "NATS_SETTINGS__MAX_RECONNECT_ATTEMPTS": "3", } @@ -38,10 +42,14 @@ def _as_asyncpg_url(url: str) -> str: raise ValueError(f"Unsupported postgres URL: {url}") -def _set_test_settings_env(monkeypatch: pytest.MonkeyPatch, *, postgres_url: str) -> None: +def _set_test_settings_env( + monkeypatch: pytest.MonkeyPatch, *, postgres_url: str, nats_url: str | None = None +) -> None: for key, value in TEST_SETTINGS_ENV.items(): monkeypatch.setenv(key, value) monkeypatch.setenv("POSTGRES_SETTINGS__URL", postgres_url) + if nats_url is not None: + monkeypatch.setenv("NATS_SETTINGS__URL", nats_url) @pytest.fixture(scope="session") @@ -53,15 +61,26 @@ def postgres_container() -> Iterator[PostgresContainer]: pytest.skip(f"Docker is not available for testcontainers: {exc}") +@pytest.fixture(scope="session") +def nats_container() -> Iterator[NatsContainer]: + try: + with NatsContainer() as container: + yield container + except DockerException as exc: + pytest.skip(f"Docker is not available for testcontainers: {exc}") + + @pytest.fixture def app_with_db( monkeypatch: pytest.MonkeyPatch, postgres_container: PostgresContainer, + nats_container: NatsContainer, clean_db: None, ) -> FastAPI: _set_test_settings_env( monkeypatch, postgres_url=_as_asyncpg_url(postgres_container.get_connection_url()), + nats_url=nats_container.nats_uri(), ) return create_app() diff --git a/tests/integration/test_databus.py b/tests/integration/test_databus.py new file mode 100644 index 0000000000000000000000000000000000000000..82569c00d51f5f41e64513649a1aaee89b92d628 --- /dev/null +++ b/tests/integration/test_databus.py @@ -0,0 +1,119 @@ +import asyncio +import json +from collections.abc import AsyncIterator, Generator + +import nats +import pytest +import pytest_asyncio +from nats.aio.client import Client +from testcontainers.nats import NatsContainer + +from srm.adapters.databus.nats_connection_manager import NatsConnectionManager +from srm.adapters.databus.nats_publisher import NatsPublisher +from srm.api.databus.nats_subscriber import NatsSubscriber, subscribe_to_subjects +from srm.api.databus.schemas import InboundMessage +from srm.config import NatsSettings +from tests.api.databus.test_nats_subscriber import EXPECTED_COMMAND_SUBJECTS + + +@pytest.fixture(scope="session") +def nats_url() -> Generator[str, None, None]: + with NatsContainer(image="nats:2.10-alpine") as container: + yield container.nats_uri() + + +@pytest_asyncio.fixture +async def connection_manager(nats_url: str) -> AsyncIterator[NatsConnectionManager]: + manager = NatsConnectionManager( + settings=NatsSettings(url=nats_url, connect_timeout=5, max_reconnect_attempts=3) + ) + await manager.connect() + try: + yield manager + finally: + await manager.close() + + +@pytest_asyncio.fixture +async def raw_client(nats_url: str) -> AsyncIterator[Client]: + client = await nats.connect(nats_url) + try: + yield client + finally: + await client.drain() + + +async def test_connection_manager_connects(connection_manager: NatsConnectionManager) -> None: + assert connection_manager.is_connected is True + + +async def test_connection_manager_close_disconnects(nats_url: str) -> None: + manager = NatsConnectionManager( + settings=NatsSettings(url=nats_url, connect_timeout=5, max_reconnect_attempts=3) + ) + await manager.connect() + + await manager.close() + + assert manager.is_connected is False + + +async def test_publisher_publishes_json_message_to_subject( + connection_manager: NatsConnectionManager, + raw_client: Client, +) -> None: + publisher = NatsPublisher(connection_manager) + received: list[dict[str, object]] = [] + ready = asyncio.Event() + + async def handler(msg: object) -> None: + received.append(json.loads(msg.data)) # type: ignore[attr-defined] + ready.set() + + sub = await raw_client.subscribe("command.srm.service.deploy", cb=handler) + await raw_client.flush() + await publisher.publish("command.srm.service.deploy", {"operation_id": "abc-123"}) + await asyncio.wait_for(ready.wait(), timeout=2.0) + await sub.unsubscribe() + + assert received == [{"operation_id": "abc-123"}] + + +async def test_subscriber_invokes_router_when_message_arrives( + connection_manager: NatsConnectionManager, + raw_client: Client, +) -> None: + received: list[InboundMessage] = [] + ready = asyncio.Event() + + async def router(message: InboundMessage) -> None: + received.append(message) + ready.set() + + subscriber = NatsSubscriber( + connection_manager=connection_manager, + subject="command.srm.service.deploy", + router=router, + ) + await subscriber.start() + await connection_manager.client.flush() + + await raw_client.publish( + "command.srm.service.deploy", + json.dumps({"operation_id": "abc-123"}).encode("utf-8"), + ) + await asyncio.wait_for(ready.wait(), timeout=2.0) + + assert len(received) == 1 + assert received[0].subject == "command.srm.service.deploy" + assert json.loads(received[0].payload) == {"operation_id": "abc-123"} + + +async def test_subscribe_to_subjects_registers_all_command_subjects( + connection_manager: NatsConnectionManager, +) -> None: + subscribers = await subscribe_to_subjects(connection_manager) + + assert sorted(sub._subject for sub in subscribers) == sorted(EXPECTED_COMMAND_SUBJECTS) + for sub in subscribers: + assert sub._subscription is not None diff --git a/tests/test_config.py b/tests/test_config.py index 12e37e050e5e7c001cd9158f7e2369b6755b7ed6..b1f8a58c74f532c13b0d81e5aaf02cb6c0b6eedf 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -3,14 +3,26 @@ from pydantic import ValidationError from srm.config import Settings, get_settings +VALID_ENV = { + "APP_NAME": "my-srm", + "APP_VERSION": "1.2.3", + "APP_DESCRIPTION": "a description", + "POSTGRES_SETTINGS__URL": "postgresql://localhost:5432/srm", + "POSTGRES_SETTINGS__ECHO": "true", + "POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP": "true", + "NATS_SETTINGS__URL": "nats://localhost:4222", + "NATS_SETTINGS__CONNECT_TIMEOUT": "10", + "NATS_SETTINGS__MAX_RECONNECT_ATTEMPTS": "3", +} + + +def set_valid_env(monkeypatch: pytest.MonkeyPatch) -> None: + for key, value in VALID_ENV.items(): + monkeypatch.setenv(key, value) + def test_loads_all_fields_from_env(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("APP_NAME", "my-srm") - monkeypatch.setenv("APP_VERSION", "1.2.3") - monkeypatch.setenv("APP_DESCRIPTION", "a description") - monkeypatch.setenv("POSTGRES_SETTINGS__URL", "postgresql://localhost:5432/srm") - monkeypatch.setenv("POSTGRES_SETTINGS__ECHO", "true") - monkeypatch.setenv("POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP", "true") + set_valid_env(monkeypatch) s = Settings(_env_file=None) # type: ignore[call-arg] assert s.app_name == "my-srm" assert s.app_version == "1.2.3" @@ -18,79 +30,43 @@ def test_loads_all_fields_from_env(monkeypatch: pytest.MonkeyPatch) -> None: assert s.postgres_settings.url == "postgresql://localhost:5432/srm" assert s.postgres_settings.echo assert s.postgres_settings.create_schema_on_startup - - -def test_raises_if_app_name_missing(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.delenv("APP_NAME", raising=False) - monkeypatch.setenv("APP_VERSION", "1.0.0") - monkeypatch.setenv("APP_DESCRIPTION", "test") - monkeypatch.setenv("POSTGRES_SETTINGS__URL", "postgresql://localhost:5432/srm") - monkeypatch.setenv("POSTGRES_SETTINGS__ECHO", "true") - monkeypatch.setenv("POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP", "true") - with pytest.raises(ValidationError): - Settings(_env_file=None) # type: ignore[call-arg] - - -def test_raises_if_app_version_missing(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("APP_NAME", "test") - monkeypatch.delenv("APP_VERSION", raising=False) - monkeypatch.setenv("APP_DESCRIPTION", "test") - monkeypatch.setenv("POSTGRES_SETTINGS__URL", "postgresql://localhost:5432/srm") - monkeypatch.setenv("POSTGRES_SETTINGS__ECHO", "true") - monkeypatch.setenv("POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP", "true") - with pytest.raises(ValidationError): - Settings(_env_file=None) # type: ignore[call-arg] - - -def test_raises_if_app_description_missing(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("APP_NAME", "test") - monkeypatch.setenv("APP_VERSION", "1.0.0") - monkeypatch.delenv("APP_DESCRIPTION", raising=False) - monkeypatch.setenv("POSTGRES_SETTINGS__URL", "postgresql://localhost:5432/srm") - monkeypatch.setenv("POSTGRES_SETTINGS__ECHO", "true") - monkeypatch.setenv("POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP", "true") - with pytest.raises(ValidationError): - Settings(_env_file=None) # type: ignore[call-arg] - - -def test_raises_if_postgres_url_missing(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("APP_NAME", "test") - monkeypatch.setenv("APP_VERSION", "1.0.0") - monkeypatch.setenv("APP_DESCRIPTION", "a description") - monkeypatch.delenv("POSTGRES_SETTINGS__URL", raising=False) - monkeypatch.setenv("POSTGRES_SETTINGS__ECHO", "true") - monkeypatch.setenv("POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP", "true") - with pytest.raises(ValidationError): - Settings(_env_file=None) # type: ignore[call-arg] - - -def test_raises_if_postgres_echo_missing(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("APP_NAME", "test") - monkeypatch.setenv("APP_VERSION", "1.0.0") - monkeypatch.setenv("APP_DESCRIPTION", "a description") - monkeypatch.setenv("POSTGRES_SETTINGS__URL", "postgresql://localhost:5432/srm") - monkeypatch.delenv("POSTGRES_SETTINGS__ECHO", raising=False) - monkeypatch.setenv("POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP", "true") + assert s.nats_settings.url == "nats://localhost:4222" + assert s.nats_settings.connect_timeout == 10 + assert s.nats_settings.max_reconnect_attempts == 3 + + +@pytest.mark.parametrize( + "missing_key", + [ + "APP_NAME", + "APP_VERSION", + "APP_DESCRIPTION", + "POSTGRES_SETTINGS__URL", + "POSTGRES_SETTINGS__ECHO", + "POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP", + "NATS_SETTINGS__URL", + "NATS_SETTINGS__CONNECT_TIMEOUT", + "NATS_SETTINGS__MAX_RECONNECT_ATTEMPTS", + ], +) +def test_raises_if_required_setting_missing( + monkeypatch: pytest.MonkeyPatch, + missing_key: str, +) -> None: + set_valid_env(monkeypatch) + monkeypatch.delenv(missing_key, raising=False) with pytest.raises(ValidationError): Settings(_env_file=None) # type: ignore[call-arg] -def test_raises_if_postgres_create_schema_missing(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("APP_NAME", "test") - monkeypatch.setenv("APP_VERSION", "1.0.0") - monkeypatch.setenv("APP_DESCRIPTION", "a description") - monkeypatch.setenv("POSTGRES_SETTINGS__URL", "postgresql://localhost:5432/srm") - monkeypatch.setenv("POSTGRES_SETTINGS__ECHO", "true") - monkeypatch.delenv("POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP", raising=False) - with pytest.raises(ValidationError): - Settings(_env_file=None) # type: ignore[call-arg] +def test_loads_nested_nats_fields(monkeypatch: pytest.MonkeyPatch) -> None: + set_valid_env(monkeypatch) + s = Settings(_env_file=None) # type: ignore[call-arg] + assert s.nats_settings.url == "nats://localhost:4222" + assert s.nats_settings.connect_timeout == 10 + assert s.nats_settings.max_reconnect_attempts == 3 def test_get_settings_returns_same_instance(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("APP_NAME", "test") - monkeypatch.setenv("APP_VERSION", "1.0.0") - monkeypatch.setenv("APP_DESCRIPTION", "test") - monkeypatch.setenv("POSTGRES_SETTINGS__URL", "postgresql://localhost:5432/srm") - monkeypatch.setenv("POSTGRES_SETTINGS__ECHO", "true") - monkeypatch.setenv("POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP", "true") + set_valid_env(monkeypatch) assert get_settings() is get_settings() diff --git a/tests/unit/test_nats_connection_manager.py b/tests/unit/test_nats_connection_manager.py new file mode 100644 index 0000000000000000000000000000000000000000..90d2e007ef408569002933db1e9a9951fae1ad53 --- /dev/null +++ b/tests/unit/test_nats_connection_manager.py @@ -0,0 +1,202 @@ +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock, patch + +import pytest +import structlog.testing + +from srm.adapters.databus.nats_connection_manager import ( + NatsConnectionManager, + init_databus_manager, +) +from srm.config import NatsSettings + + +@pytest.fixture +def settings() -> NatsSettings: + return NatsSettings( + url="nats://broker:4222", + connect_timeout=3, + max_reconnect_attempts=2, + drain_timeout=7, + ) + + +def test_drain_timeout_defaults_when_unset() -> None: + settings = NatsSettings(url="nats://broker:4222", connect_timeout=3, max_reconnect_attempts=2) + assert settings.drain_timeout == 30 + + +def test_is_connected_false_when_no_client(settings: NatsSettings) -> None: + manager = NatsConnectionManager(settings=settings) + assert manager.is_connected is False + + +def test_client_raises_when_not_connected(settings: NatsSettings) -> None: + manager = NatsConnectionManager(settings=settings) + with pytest.raises(RuntimeError, match="not connected"): + _ = manager.client + + +async def test_connect_uses_settings(settings: NatsSettings) -> None: + mock_client = AsyncMock() + with patch( + "srm.adapters.databus.nats_connection_manager.nats.connect", + new_callable=AsyncMock, + return_value=mock_client, + ) as connect: + manager = NatsConnectionManager(settings=settings) + await manager.connect() + + call_kwargs = connect.call_args.kwargs + assert call_kwargs["servers"] == ["nats://broker:4222"] + assert call_kwargs["connect_timeout"] == 3 + assert call_kwargs["max_reconnect_attempts"] == 2 + assert call_kwargs["drain_timeout"] == 7 + assert callable(call_kwargs["error_cb"]) + assert callable(call_kwargs["disconnected_cb"]) + assert callable(call_kwargs["reconnected_cb"]) + assert manager.client is mock_client + + +class TestConnectionCallbacks: + @staticmethod + async def _connect_and_capture_callbacks(settings: NatsSettings) -> dict[str, Any]: + with patch( + "srm.adapters.databus.nats_connection_manager.nats.connect", + new_callable=AsyncMock, + return_value=AsyncMock(), + ) as connect: + await NatsConnectionManager(settings=settings).connect() + + return dict(connect.call_args.kwargs) + + async def test_error_cb_logs_the_error(self, settings: NatsSettings) -> None: + callbacks = await self._connect_and_capture_callbacks(settings) + + with structlog.testing.capture_logs() as logs: + await callbacks["error_cb"](OSError("stale connection")) + + assert [(e["event"], e["error"], e["log_level"]) for e in logs] == [ + ("nats_error", "stale connection", "error") + ] + + async def test_disconnected_cb_logs_the_url(self, settings: NatsSettings) -> None: + callbacks = await self._connect_and_capture_callbacks(settings) + + with structlog.testing.capture_logs() as logs: + await callbacks["disconnected_cb"]() + + assert [(e["event"], e["url"], e["log_level"]) for e in logs] == [ + ("nats_disconnected", "nats://broker:4222", "warning") + ] + + async def test_reconnected_cb_logs_the_url(self, settings: NatsSettings) -> None: + callbacks = await self._connect_and_capture_callbacks(settings) + + with structlog.testing.capture_logs() as logs: + await callbacks["reconnected_cb"]() + + assert [(e["event"], e["url"], e["log_level"]) for e in logs] == [ + ("nats_reconnected", "nats://broker:4222", "info") + ] + + +async def test_connect_is_noop_when_already_connected(settings: NatsSettings) -> None: + mock_client = AsyncMock() + mock_client.is_connected = True + + with patch( + "srm.adapters.databus.nats_connection_manager.nats.connect", + new_callable=AsyncMock, + return_value=mock_client, + ) as connect: + manager = NatsConnectionManager(settings=settings) + await manager.connect() + await manager.connect() + + connect.assert_awaited_once() + + +async def test_concurrent_connect_dials_only_once(settings: NatsSettings) -> None: + """Without the lock both callers see is_connected == False and dial, leaking a client.""" + mock_client = AsyncMock() + mock_client.is_connected = True + + async def slow_connect(**kwargs: Any) -> AsyncMock: + await asyncio.sleep(0) # hand control back so the second caller can interleave + return mock_client + + with patch( + "srm.adapters.databus.nats_connection_manager.nats.connect", + new_callable=AsyncMock, + side_effect=slow_connect, + ) as connect: + manager = NatsConnectionManager(settings=settings) + await asyncio.gather(manager.connect(), manager.connect()) + + connect.assert_awaited_once() + assert manager.client is mock_client + + +async def test_connect_propagates_and_logs_connection_errors(settings: NatsSettings) -> None: + with patch( + "srm.adapters.databus.nats_connection_manager.nats.connect", + new_callable=AsyncMock, + side_effect=OSError("connection refused"), + ): + manager = NatsConnectionManager(settings=settings) + with pytest.raises(OSError, match="connection refused"): + await manager.connect() + + assert manager.is_connected is False + + +async def test_close_drains_client_and_clears_state(settings: NatsSettings) -> None: + mock_client = AsyncMock() + manager = NatsConnectionManager(settings=settings) + manager._client = mock_client + + await manager.close() + + mock_client.drain.assert_awaited_once() + assert manager.is_connected is False + + +async def test_close_force_closes_when_drain_fails(settings: NatsSettings) -> None: + """A failed drain must not leave the connection open or block shutdown.""" + mock_client = AsyncMock() + mock_client.drain.side_effect = RuntimeError("drain exploded") + manager = NatsConnectionManager(settings=settings) + manager._client = mock_client + + with structlog.testing.capture_logs() as logs: + await manager.close() + + mock_client.close.assert_awaited_once() + assert manager.is_connected is False + assert [(e["event"], e["error"], e["log_level"]) for e in logs] == [ + ("nats_drain_failed", "drain exploded", "warning") + ] + + +async def test_close_when_not_connected_is_noop(settings: NatsSettings) -> None: + manager = NatsConnectionManager(settings=settings) + await manager.close() + assert manager.is_connected is False + + +async def test_init_databus_manager_connects_and_returns_manager(settings: NatsSettings) -> None: + mock_client = AsyncMock() + mock_client.is_connected = True + with patch( + "srm.adapters.databus.nats_connection_manager.nats.connect", + new_callable=AsyncMock, + return_value=mock_client, + ): + manager = await init_databus_manager(settings) + + assert isinstance(manager, NatsConnectionManager) + assert manager.is_connected is True diff --git a/tests/unit/test_nats_publisher.py b/tests/unit/test_nats_publisher.py new file mode 100644 index 0000000000000000000000000000000000000000..e6b06620ecc40e8a28d2ef3c7ced2be9e14d5c14 --- /dev/null +++ b/tests/unit/test_nats_publisher.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +import json +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from srm.adapters.databus.nats_connection_manager import NatsConnectionManager +from srm.adapters.databus.nats_publisher import NatsPublisher + + +@pytest.fixture +def connection_manager() -> MagicMock: + manager = MagicMock(spec=NatsConnectionManager) + manager.client = AsyncMock() + return manager + + +async def test_publish_serializes_json_payload(connection_manager: MagicMock) -> None: + publisher = NatsPublisher(connection_manager) + + await publisher.publish("command.srm.service.deploy", {"operation_id": "abc-123"}) + + connection_manager.client.publish.assert_awaited_once_with( + "command.srm.service.deploy", + json.dumps({"operation_id": "abc-123"}).encode("utf-8"), + headers=None, + ) + + +async def test_publish_passes_headers(connection_manager: MagicMock) -> None: + publisher = NatsPublisher(connection_manager) + + await publisher.publish( + "command.srm.service.deploy", + {"operation_id": "abc-123"}, + headers={"x-correlation-id": "corr-1"}, + ) + + connection_manager.client.publish.assert_awaited_once_with( + "command.srm.service.deploy", + json.dumps({"operation_id": "abc-123"}).encode("utf-8"), + headers={"x-correlation-id": "corr-1"}, + ) + + +async def test_publish_rejects_non_json_serializable_payload(connection_manager: MagicMock) -> None: + publisher = NatsPublisher(connection_manager) + + with pytest.raises(TypeError): + await publisher.publish("command.srm.service.deploy", {"bad": object()}) + + connection_manager.client.publish.assert_not_awaited() diff --git a/uv.lock b/uv.lock index 80cae94fdf83b7aac0f1173c47e88a26f4d42f66..19113912835f94ab68290d99d08c02e8ea80392b 100644 --- a/uv.lock +++ b/uv.lock @@ -689,6 +689,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/79/7b/2c79738432f5c924bef5071f933bcc9efd0473bac3b4aa584a6f7c1c8df8/mypy_extensions-1.1.0-py3-none-any.whl", hash = "sha256:1be4cccdb0f2482337c4743e60421de3a356cd97508abadd57d47403e94f5505", size = 4963, upload-time = "2025-04-22T14:54:22.983Z" }, ] +[[package]] +name = "nats-py" +version = "2.15.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/02/f0/fc5e93f2b0dd14a202590ad9d30eda1955ea872039b5204357348d0f4b1e/nats_py-2.15.0.tar.gz", hash = "sha256:6622c547d9a7d2313d9c147d46c386188f4ec2c7b5c9f9a0438a4d1b55f54a93", size = 75995, upload-time = "2026-06-05T07:34:03.904Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/db/a8/b55606c7c621fb813c8ec78baf201d2c78bf6051091ec0c7ada572999e95/nats_py-2.15.0-py3-none-any.whl", hash = "sha256:9f8d36aa52a9926a88b8f1d70cf1fdce0ad387941479b500ee9ab3e51073cefd", size = 90334, upload-time = "2026-06-05T07:34:02.81Z" }, +] + [[package]] name = "nodeenv" version = "1.10.0" @@ -1135,6 +1144,7 @@ source = { editable = "." } dependencies = [ { name = "asyncpg" }, { name = "fastapi", extra = ["standard"] }, + { name = "nats-py" }, { name = "pydantic-settings" }, { name = "python-dotenv" }, { name = "sqlalchemy" }, @@ -1162,6 +1172,7 @@ requires-dist = [ { name = "fastapi", extras = ["standard"], specifier = ">=0.135.1" }, { name = "import-linter", marker = "extra == 'dev'", specifier = ">=2.11" }, { name = "mypy", marker = "extra == 'dev'", specifier = ">=1.19.1" }, + { name = "nats-py", specifier = ">=2.10.0" }, { name = "pre-commit", marker = "extra == 'dev'", specifier = ">=4.5.1" }, { name = "pydantic-settings", specifier = ">=2.13.1" }, { name = "pytest", marker = "extra == 'dev'", specifier = ">=9.0.2" },