From 39ed61bc2d6f0785677926e503472591b94dd0c0 Mon Sep 17 00:00:00 2001 From: stentoumis Date: Wed, 29 Jul 2026 15:57:28 +0300 Subject: [PATCH 1/8] feat: nats adapter implementation. and some tests --- .env.example | 4 + pyproject.toml | 3 +- src/srm/adapters/databus/.gitkeep | 0 .../databus/nats_connection_manager.py | 60 +++++++++ src/srm/adapters/databus/nats_publisher.py | 20 +++ src/srm/api/databus/.gitkeep | 0 src/srm/api/databus/message_router.py | 22 ++++ src/srm/api/databus/nats_subscriber.py | 95 ++++++++++++++ src/srm/api/databus/schemas.py | 114 ++++++++++++++++ src/srm/api/health.py | 7 +- src/srm/app_state.py | 5 + src/srm/config.py | 7 + src/srm/domain/ports/databus/publisher.py | 10 ++ src/srm/main.py | 19 ++- tests/api/fakes.py | 27 ++++ tests/api/test_health.py | 34 ++++- tests/conftest.py | 3 + tests/test_config.py | 124 +++++++----------- uv.lock | 11 ++ 19 files changed, 487 insertions(+), 78 deletions(-) delete mode 100644 src/srm/adapters/databus/.gitkeep create mode 100644 src/srm/adapters/databus/nats_connection_manager.py create mode 100644 src/srm/adapters/databus/nats_publisher.py delete mode 100644 src/srm/api/databus/.gitkeep create mode 100644 src/srm/api/databus/message_router.py create mode 100644 src/srm/api/databus/nats_subscriber.py create mode 100644 src/srm/api/databus/schemas.py create mode 100644 src/srm/domain/ports/databus/publisher.py create mode 100644 tests/api/fakes.py diff --git a/.env.example b/.env.example index 06995e9..6722c82 100644 --- a/.env.example +++ b/.env.example @@ -5,3 +5,7 @@ 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_ATTEMPS = 3 diff --git a/pyproject.toml b/pyproject.toml index c99d698..a7a3783 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/.gitkeep b/src/srm/adapters/databus/.gitkeep deleted file mode 100644 index e69de29..0000000 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 0000000..0f45d79 --- /dev/null +++ b/src/srm/adapters/databus/nats_connection_manager.py @@ -0,0 +1,60 @@ +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 + + @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) + + if not self.is_connected: + try: + self._client = await nats.connect( + servers=[self._settings.url], + connect_timeout=self._settings.connect_timeout, + max_reconnect_attempts=self._settings.max_reconnect_attempts, + error_cb=_on_error, + disconnected_cb=_on_disconnect, + reconnected_cb=_on_reconnect, + ) + except Exception as e: + logger.error("nats_error", error=str(e)) + raise + else: + logger.info("nats_already_connected") + + async def close(self) -> None: + if self._client is not None: + await self._client.drain() + 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 0000000..c03ec67 --- /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 e69de29..0000000 diff --git a/src/srm/api/databus/message_router.py b/src/srm/api/databus/message_router.py new file mode 100644 index 0000000..9990f17 --- /dev/null +++ b/src/srm/api/databus/message_router.py @@ -0,0 +1,22 @@ +from collections.abc import Awaitable, Callable + +from nats.aio.msg import Msg + + +class MessageRouter: + def __init__(self) -> None: + self._handlers: dict[str, Callable[[Msg], Awaitable[None]]] = {} + + def register_handler( + self, + subject: str, + handler: Callable[[Msg], Awaitable[None]], + ) -> None: + self._handlers[subject] = handler + + async def route(self, msg: Msg) -> None: + handler = self._handlers.get(msg.subject) + if handler is None: + raise RuntimeError(f"No handler registered for subject: {msg.subject}") + + await handler(msg) diff --git a/src/srm/api/databus/nats_subscriber.py b/src/srm/api/databus/nats_subscriber.py new file mode 100644 index 0000000..2809dd8 --- /dev/null +++ b/src/srm/api/databus/nats_subscriber.py @@ -0,0 +1,95 @@ +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__) + + +async def _noop_router(message: InboundMessage) -> None: + logger.info("Message Received", subject=message.subject, data=message.payload) + return None + + +async def subscribe_to_subjects( + connection_manager: NatsConnectionManager, +) -> list["NatsSubscriber"]: + subscribers: list[NatsSubscriber] = [] + + service_deploy_sub = NatsSubscriber( + connection_manager=connection_manager, + subject="command.srm.service.deploy", + router=_noop_router, + ) + subscribers.append(service_deploy_sub) + + service_scale_sub = NatsSubscriber( + connection_manager=connection_manager, + subject="command.srm.service.scale", + router=_noop_router, + ) + subscribers.append(service_scale_sub) + + service_terminate_sub = NatsSubscriber( + connection_manager=connection_manager, + subject="command.srm.service.terminate", + router=_noop_router, + ) + subscribers.append(service_terminate_sub) + + service_network_activate_sub = NatsSubscriber( + connection_manager=connection_manager, + subject="command.srm.network.capability.activate", + router=_noop_router, + ) + subscribers.append(service_network_activate_sub) + + service_network_update_sub = NatsSubscriber( + connection_manager=connection_manager, + subject="command.srm.network.capability.update", + router=_noop_router, + ) + subscribers.append(service_network_update_sub) + + service_network_deactivate_sub = NatsSubscriber( + connection_manager=connection_manager, + subject="command.srm.network.capability.deactivate", + router=_noop_router, + ) + subscribers.append(service_network_deactivate_sub) + + 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 {}, + ) + await self._router(inbound_message) diff --git a/src/srm/api/databus/schemas.py b/src/srm/api/databus/schemas.py new file mode 100644 index 0000000..e108750 --- /dev/null +++ b/src/srm/api/databus/schemas.py @@ -0,0 +1,114 @@ +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: str = "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"] + + +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 + + +class SrmServiceDeployV1(CommandEnvelopeV1): + service_specification_public_id: UUID + targets: list[DeployTargetV1] + deploy: DeployPayloadV1 + + +class ScalePayloadV1(BaseModel): + replicas: int | None = Field(default=None, ge=0) + + +class SrmServiceScaleV1(CommandEnvelopeV1): + service_instance_public_id: UUID + service_specification_public_id: UUID | None = None + scale: ScalePayloadV1 + + +class TerminatePayloadV1(BaseModel): + grace_period_seconds: int = Field(default=0, ge=0) + + +class SrmServiceTerminateV1(CommandEnvelopeV1): + service_instance_public_id: UUID + service_specification_public_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_public_id: UUID + zone_public_id: UUID | None = None + network_capability: NetworkCapabilityPayloadV1 + + +class NetworkCapabilityRealizationRefV1(BaseModel): + external_ref: str | None = None + service_instance_public_id: UUID | None = None + capability_type: str | None = None + + @model_validator(mode="after") + def validate_reference_shape(self) -> "NetworkCapabilityRealizationRefV1": + if self.external_ref is not None: + return self + + if self.service_instance_public_id is not None and self.capability_type is not None: + return self + + raise ValueError( + "network capability realization must be identified by external_ref " + "or by service_instance_public_id and capability_type" + ) + + +class NetworkCapabilityUpdatePayloadV1(NetworkCapabilityRealizationRefV1): + parameters: CapabilityParameters + + +class SrmNetworkCapabilityUpdateV1(CommandEnvelopeV1): + network_capability: NetworkCapabilityUpdatePayloadV1 + + +class NetworkCapabilityDeactivatePayloadV1(NetworkCapabilityRealizationRefV1): + grace_period_seconds: int = Field(default=0, ge=0) + + +class SrmNetworkCapabilityDeactivateV1(CommandEnvelopeV1): + network_capability: NetworkCapabilityDeactivatePayloadV1 diff --git a/src/srm/api/health.py b/src/srm/api/health.py index d34ee2f..f3f4412 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 a21baa6..799d9cb 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 5a75a01..cb39787 100644 --- a/src/srm/config.py +++ b/src/srm/config.py @@ -29,6 +29,12 @@ class PostgreSQLSettings(BaseModel): create_schema_on_startup: bool +class NatsSettings(BaseModel): + url: str + connect_timeout: int + max_reconnect_attempts: int + + class Settings(BaseSettings): model_config = SettingsConfigDict(env_file=".env", env_nested_delimiter="__") @@ -37,6 +43,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 0000000..613f70d --- /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 6e56ba4..bfa7973 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,14 +54,23 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: except Exception as e: logger.error("Database engine init failed!", error=str(e)) raise + databus_manager: NatsConnectionManager = await init_databus_manager( + settings=settings.nats_settings + ) + + databus_subscribers: list[NatsSubscriber] = await subscribe_to_subjects(databus_manager) 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") await engine.dispose() + await app.state.databus_connection_manager.close() + app.state.databus_subscribers = [] def create_app() -> FastAPI: diff --git a/tests/api/fakes.py b/tests/api/fakes.py new file mode 100644 index 0000000..605bacd --- /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_health.py b/tests/api/test_health.py index e079c35..9476006 100644 --- a/tests/api/test_health.py +++ b/tests/api/test_health.py @@ -1,4 +1,16 @@ -from httpx import AsyncClient +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: @@ -11,3 +23,23 @@ 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 4206cf1..62f3feb 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -20,6 +20,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", } diff --git a/tests/test_config.py b/tests/test_config.py index 12e37e0..b1f8a58 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/uv.lock b/uv.lock index 80cae94..1911391 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" }, -- GitLab From b8a20797adc60b0b9fa2b7ffa95c9409c04e012f Mon Sep 17 00:00:00 2001 From: dgogos Date: Sun, 2 Aug 2026 10:19:09 +0300 Subject: [PATCH 2/8] feat: refactor NATS subscriber and publisher, add tests for connection manager and publisher --- .env.example | 2 +- src/srm/api/databus/message_router.py | 22 -- src/srm/api/databus/nats_subscriber.py | 62 ++-- src/srm/api/databus/schemas.py | 58 ++-- tests/api/databus/__init__.py | 0 tests/api/databus/test_nats_subscriber.py | 134 +++++++++ tests/api/databus/test_schemas.py | 321 +++++++++++++++++++++ tests/conftest.py | 18 +- tests/integration/test_databus.py | 122 ++++++++ tests/unit/test_nats_connection_manager.py | 107 +++++++ tests/unit/test_nats_publisher.py | 53 ++++ 11 files changed, 812 insertions(+), 87 deletions(-) delete mode 100644 src/srm/api/databus/message_router.py create mode 100644 tests/api/databus/__init__.py create mode 100644 tests/api/databus/test_nats_subscriber.py create mode 100644 tests/api/databus/test_schemas.py create mode 100644 tests/integration/test_databus.py create mode 100644 tests/unit/test_nats_connection_manager.py create mode 100644 tests/unit/test_nats_publisher.py diff --git a/.env.example b/.env.example index 6722c82..ec81351 100644 --- a/.env.example +++ b/.env.example @@ -8,4 +8,4 @@ POSTGRES_SETTINGS__CREATE_SCHEMA_ON_STARTUP = true NATS_SETTINGS__URL = "nats://localhost:4222" NATS_SETTINGS__CONNECT_TIMEOUT = 10 -NATS_SETTINGS__MAX_RECONNECT_ATTEMPS = 3 +NATS_SETTINGS__MAX_RECONNECT_ATTEMPTS = 3 diff --git a/src/srm/api/databus/message_router.py b/src/srm/api/databus/message_router.py deleted file mode 100644 index 9990f17..0000000 --- a/src/srm/api/databus/message_router.py +++ /dev/null @@ -1,22 +0,0 @@ -from collections.abc import Awaitable, Callable - -from nats.aio.msg import Msg - - -class MessageRouter: - def __init__(self) -> None: - self._handlers: dict[str, Callable[[Msg], Awaitable[None]]] = {} - - def register_handler( - self, - subject: str, - handler: Callable[[Msg], Awaitable[None]], - ) -> None: - self._handlers[subject] = handler - - async def route(self, msg: Msg) -> None: - handler = self._handlers.get(msg.subject) - if handler is None: - raise RuntimeError(f"No handler registered for subject: {msg.subject}") - - await handler(msg) diff --git a/src/srm/api/databus/nats_subscriber.py b/src/srm/api/databus/nats_subscriber.py index 2809dd8..f1d2d70 100644 --- a/src/srm/api/databus/nats_subscriber.py +++ b/src/srm/api/databus/nats_subscriber.py @@ -9,8 +9,18 @@ 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 @@ -18,49 +28,10 @@ async def _noop_router(message: InboundMessage) -> None: async def subscribe_to_subjects( connection_manager: NatsConnectionManager, ) -> list["NatsSubscriber"]: - subscribers: list[NatsSubscriber] = [] - - service_deploy_sub = NatsSubscriber( - connection_manager=connection_manager, - subject="command.srm.service.deploy", - router=_noop_router, - ) - subscribers.append(service_deploy_sub) - - service_scale_sub = NatsSubscriber( - connection_manager=connection_manager, - subject="command.srm.service.scale", - router=_noop_router, - ) - subscribers.append(service_scale_sub) - - service_terminate_sub = NatsSubscriber( - connection_manager=connection_manager, - subject="command.srm.service.terminate", - router=_noop_router, - ) - subscribers.append(service_terminate_sub) - - service_network_activate_sub = NatsSubscriber( - connection_manager=connection_manager, - subject="command.srm.network.capability.activate", - router=_noop_router, - ) - subscribers.append(service_network_activate_sub) - - service_network_update_sub = NatsSubscriber( - connection_manager=connection_manager, - subject="command.srm.network.capability.update", - router=_noop_router, - ) - subscribers.append(service_network_update_sub) - - service_network_deactivate_sub = NatsSubscriber( - connection_manager=connection_manager, - subject="command.srm.network.capability.deactivate", - router=_noop_router, - ) - subscribers.append(service_network_deactivate_sub) + subscribers = [ + NatsSubscriber(connection_manager=connection_manager, subject=subject, router=_noop_router) + for subject in COMMAND_SUBJECTS + ] for subscriber in subscribers: await subscriber.start() @@ -92,4 +63,7 @@ class NatsSubscriber: payload=msg.data, headers=dict(msg.headers) if msg.headers is not None else {}, ) - await self._router(inbound_message) + 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 index e108750..0da04d7 100644 --- a/src/srm/api/databus/schemas.py +++ b/src/srm/api/databus/schemas.py @@ -38,21 +38,29 @@ class DeployPayloadV1(BaseModel): 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_public_id: UUID - targets: list[DeployTargetV1] + service_specification_id: UUID + targets: list[DeployTargetV1] = Field(min_length=1) deploy: DeployPayloadV1 class ScalePayloadV1(BaseModel): - replicas: int | None = Field(default=None, ge=0) + replicas: int = Field(ge=0) class SrmServiceScaleV1(CommandEnvelopeV1): - service_instance_public_id: UUID - service_specification_public_id: UUID | None = None + service_instance_id: UUID + service_specification_id: UUID | None = None scale: ScalePayloadV1 @@ -61,8 +69,8 @@ class TerminatePayloadV1(BaseModel): class SrmServiceTerminateV1(CommandEnvelopeV1): - service_instance_public_id: UUID - service_specification_public_id: UUID | None = None + service_instance_id: UUID + service_specification_id: UUID | None = None terminate: TerminatePayloadV1 @@ -74,35 +82,46 @@ class NetworkCapabilityPayloadV1(BaseModel): class SrmNetworkCapabilityActivateV1(CommandEnvelopeV1): - service_specification_public_id: UUID - zone_public_id: UUID | None = None + 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_public_id: UUID | None = None - capability_type: str | None = None + service_instance_id: UUID | None = None @model_validator(mode="after") def validate_reference_shape(self) -> "NetworkCapabilityRealizationRefV1": - if self.external_ref is not None: - return self + has_external_ref = self.external_ref is not None + has_service_instance_id = self.service_instance_id is not None - if self.service_instance_public_id is not None and self.capability_type is not None: - return self + 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" + ) - raise ValueError( - "network capability realization must be identified by external_ref " - "or by service_instance_public_id and capability_type" - ) + 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 @@ -111,4 +130,5 @@ class NetworkCapabilityDeactivatePayloadV1(NetworkCapabilityRealizationRefV1): class SrmNetworkCapabilityDeactivateV1(CommandEnvelopeV1): + service_specification_id: UUID | None = None network_capability: NetworkCapabilityDeactivatePayloadV1 diff --git a/tests/api/databus/__init__.py b/tests/api/databus/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/api/databus/test_nats_subscriber.py b/tests/api/databus/test_nats_subscriber.py new file mode 100644 index 0000000..33e35f5 --- /dev/null +++ b/tests/api/databus/test_nats_subscriber.py @@ -0,0 +1,134 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from unittest.mock import AsyncMock, MagicMock + +import pytest + +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_swallows_router_exceptions(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"{}") + + await subscriber._handle_message(msg) # type: ignore[arg-type] + + router.assert_awaited_once() + + +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(COMMAND_SUBJECTS) + assert connection_manager.client.subscribe.await_count == len(COMMAND_SUBJECTS) + + +class TestCommandDeliveryDurability: + @pytest.mark.skip( + reason="TODO: JetStream durable consumer not implemented yet — " + "NatsSubscriber still uses core-NATS subscribe() (at-most-once)." + ) + async def test_subscriber_subscribes_durably_via_jetstream( + self, connection_manager: MagicMock + ) -> None: + """SRM must consume command.srm.* durably, at-least-once (the OOP_TASKS WorkQueue + stream, interface-contract.md §D.1-D.2). A core-NATS subscribe() is non-durable and + at-most-once: every command published while SRM is down or slow is lost with no + redelivery, and OEG/FM never see a terminal event for it.""" + jetstream = MagicMock() + jetstream.subscribe = AsyncMock() + connection_manager.client.jetstream.return_value = jetstream + + subscriber, _ = make_subscriber(connection_manager) + await subscriber.start() + + connection_manager.client.subscribe.assert_not_awaited() + jetstream.subscribe.assert_awaited_once() + assert jetstream.subscribe.await_args is not None + assert jetstream.subscribe.await_args.kwargs.get("durable"), ( + "the contract requires a named durable consumer" + ) diff --git a/tests/api/databus/test_schemas.py b/tests/api/databus/test_schemas.py new file mode 100644 index 0000000..5268dec --- /dev/null +++ b/tests/api/databus/test_schemas.py @@ -0,0 +1,321 @@ +from __future__ import annotations + +from uuid import uuid4 + +import pytest +from pydantic import ValidationError + +from srm.api.databus.schemas import ( + 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] = { + "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 + + +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/conftest.py b/tests/conftest.py index 62f3feb..dbc6a86 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 @@ -41,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") @@ -56,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 0000000..7665151 --- /dev/null +++ b/tests/integration/test_databus.py @@ -0,0 +1,122 @@ +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 ( + COMMAND_SUBJECTS, + NatsSubscriber, + subscribe_to_subjects, +) +from srm.api.databus.schemas import InboundMessage +from srm.config import NatsSettings + + +@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(COMMAND_SUBJECTS) + for sub in subscribers: + assert sub._subscription is not None diff --git a/tests/unit/test_nats_connection_manager.py b/tests/unit/test_nats_connection_manager.py new file mode 100644 index 0000000..6326a56 --- /dev/null +++ b/tests/unit/test_nats_connection_manager.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +from unittest.mock import AsyncMock, patch + +import pytest + +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) + + +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 callable(call_kwargs["error_cb"]) + assert callable(call_kwargs["disconnected_cb"]) + assert callable(call_kwargs["reconnected_cb"]) + assert manager.client is mock_client + + +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_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_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 0000000..e6b0662 --- /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() -- GitLab From b3407087c38403c072612c413f03c50783092e0d Mon Sep 17 00:00:00 2001 From: dgogos Date: Sun, 2 Aug 2026 11:06:30 +0300 Subject: [PATCH 3/8] feat: enforce schema_version validation in CommandEnvelopeV1 and add related tests --- src/srm/api/databus/schemas.py | 2 +- tests/api/databus/test_schemas.py | 31 +++++++++++++++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) diff --git a/src/srm/api/databus/schemas.py b/src/srm/api/databus/schemas.py index 0da04d7..d31c93c 100644 --- a/src/srm/api/databus/schemas.py +++ b/src/srm/api/databus/schemas.py @@ -17,7 +17,7 @@ class InboundMessage(BaseModel): class CommandEnvelopeV1(BaseModel): - schema_version: str = "1.0" + schema_version: Literal["1.0"] operation_id: UUID correlation_id: str requested_at: datetime diff --git a/tests/api/databus/test_schemas.py b/tests/api/databus/test_schemas.py index 5268dec..5e9a66f 100644 --- a/tests/api/databus/test_schemas.py +++ b/tests/api/databus/test_schemas.py @@ -6,6 +6,7 @@ import pytest from pydantic import ValidationError from srm.api.databus.schemas import ( + CommandEnvelopeV1, DeployPayloadV1, DeployTargetV1, NetworkCapabilityDeactivatePayloadV1, @@ -28,6 +29,7 @@ from srm.domain.models.canonical_parameters.parameters import ( 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", @@ -38,6 +40,35 @@ def _envelope(**overrides: object) -> dict[str, object]: 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) + + def test_deploy_v1_parses_minimal_valid_payload() -> None: payload = _envelope( service_specification_id=str(uuid4()), -- GitLab From 6061d7f93c1b79349952dfcab236fa29b904180c Mon Sep 17 00:00:00 2001 From: dgogos Date: Mon, 3 Aug 2026 11:46:40 +0300 Subject: [PATCH 4/8] feat: add validation for federation context in CommandEnvelopeV1 and enhance related tests --- src/srm/api/databus/schemas.py | 7 +++ tests/api/databus/test_nats_subscriber.py | 42 +++++++++++-- tests/api/databus/test_schemas.py | 46 +++++++++++++- tests/api/test_app.py | 70 ++++++++++++++++++++++ tests/integration/test_databus.py | 10 ++-- tests/unit/test_nats_connection_manager.py | 46 ++++++++++++++ 6 files changed, 209 insertions(+), 12 deletions(-) diff --git a/src/srm/api/databus/schemas.py b/src/srm/api/databus/schemas.py index d31c93c..2072cca 100644 --- a/src/srm/api/databus/schemas.py +++ b/src/srm/api/databus/schemas.py @@ -25,6 +25,13 @@ class CommandEnvelopeV1(BaseModel): 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"} diff --git a/tests/api/databus/test_nats_subscriber.py b/tests/api/databus/test_nats_subscriber.py index 33e35f5..07f9177 100644 --- a/tests/api/databus/test_nats_subscriber.py +++ b/tests/api/databus/test_nats_subscriber.py @@ -4,6 +4,7 @@ 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 ( @@ -88,14 +89,47 @@ async def test_handle_message_defaults_headers_to_empty_dict_when_none( assert inbound.headers == {} -async def test_handle_message_swallows_router_exceptions(connection_manager: MagicMock) -> None: +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"{}") - await subscriber._handle_message(msg) # type: ignore[arg-type] + 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_command_dropped", "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( @@ -103,8 +137,8 @@ async def test_subscribe_to_subjects_registers_all_command_subjects( ) -> None: subscribers = await subscribe_to_subjects(connection_manager) - assert sorted(sub._subject for sub in subscribers) == sorted(COMMAND_SUBJECTS) - assert connection_manager.client.subscribe.await_count == len(COMMAND_SUBJECTS) + assert sorted(sub._subject for sub in subscribers) == sorted(EXPECTED_COMMAND_SUBJECTS) + assert connection_manager.client.subscribe.await_count == len(EXPECTED_COMMAND_SUBJECTS) class TestCommandDeliveryDurability: diff --git a/tests/api/databus/test_schemas.py b/tests/api/databus/test_schemas.py index 5e9a66f..653ae0d 100644 --- a/tests/api/databus/test_schemas.py +++ b/tests/api/databus/test_schemas.py @@ -69,6 +69,50 @@ class TestSchemaVersion: 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()), @@ -109,7 +153,6 @@ def test_deploy_target_accepts_zone_and_domain_pin() -> None: 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()) @@ -255,7 +298,6 @@ def test_network_capability_payload_requires_capability_type() -> None: class TestNetworkCapabilityRealizationRef: - def test_accepts_external_ref_alone(self) -> None: ref = NetworkCapabilityUpdatePayloadV1( capability_type="qod_session", diff --git a/tests/api/test_app.py b/tests/api/test_app.py index d46ee3f..4096bab 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,68 @@ 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 + + async def test_startup_fails_when_subject_subscription_fails( + self, app: FastAPI, stub_database: AsyncMock + ) -> None: + with ( + patch("srm.main.init_databus_manager", new_callable=AsyncMock), + 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 + + 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 == [] diff --git a/tests/integration/test_databus.py b/tests/integration/test_databus.py index 7665151..fb581ed 100644 --- a/tests/integration/test_databus.py +++ b/tests/integration/test_databus.py @@ -10,14 +10,12 @@ 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 ( - COMMAND_SUBJECTS, - NatsSubscriber, - subscribe_to_subjects, -) +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]: @@ -117,6 +115,6 @@ async def test_subscribe_to_subjects_registers_all_command_subjects( ) -> None: subscribers = await subscribe_to_subjects(connection_manager) - assert sorted(sub._subject for sub in subscribers) == sorted(COMMAND_SUBJECTS) + 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/unit/test_nats_connection_manager.py b/tests/unit/test_nats_connection_manager.py index 6326a56..1f70a7f 100644 --- a/tests/unit/test_nats_connection_manager.py +++ b/tests/unit/test_nats_connection_manager.py @@ -1,8 +1,10 @@ from __future__ import annotations +from typing import Any from unittest.mock import AsyncMock, patch import pytest +import structlog.testing from srm.adapters.databus.nats_connection_manager import ( NatsConnectionManager, @@ -47,6 +49,50 @@ async def test_connect_uses_settings(settings: NatsSettings) -> None: 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 -- GitLab From e29de8ea8b94b99758f6dbfa8a7e3bdc884a0553 Mon Sep 17 00:00:00 2001 From: dgogos Date: Mon, 3 Aug 2026 13:05:28 +0300 Subject: [PATCH 5/8] feat: add drain_timeout to NATS settings and enhance connection management --- .env.example | 1 + .pre-commit-config.yaml | 1 + .../databus/nats_connection_manager.py | 26 +++++++-- src/srm/config.py | 1 + src/srm/main.py | 23 ++++++-- tests/api/databus/test_nats_subscriber.py | 3 +- tests/api/test_app.py | 27 +++++++++- tests/integration/test_databus.py | 1 - tests/unit/test_nats_connection_manager.py | 53 ++++++++++++++++++- 9 files changed, 121 insertions(+), 15 deletions(-) diff --git a/.env.example b/.env.example index ec81351..ca26bb5 100644 --- a/.env.example +++ b/.env.example @@ -9,3 +9,4 @@ 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/.pre-commit-config.yaml b/.pre-commit-config.yaml index df54d1b..8220cac 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/src/srm/adapters/databus/nats_connection_manager.py b/src/srm/adapters/databus/nats_connection_manager.py index 0f45d79..47693e0 100644 --- a/src/srm/adapters/databus/nats_connection_manager.py +++ b/src/srm/adapters/databus/nats_connection_manager.py @@ -1,3 +1,5 @@ +import asyncio + import nats import structlog from nats.aio.client import Client @@ -17,6 +19,7 @@ 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: @@ -38,12 +41,17 @@ class NatsConnectionManager: async def _on_reconnect() -> None: logger.info("nats_reconnected", url=self._settings.url) - if not self.is_connected: + 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, @@ -51,10 +59,18 @@ class NatsConnectionManager: except Exception as e: logger.error("nats_error", error=str(e)) raise - else: - logger.info("nats_already_connected") async def close(self) -> None: - if self._client is not None: - await self._client.drain() + 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/config.py b/src/srm/config.py index cb39787..20a4d87 100644 --- a/src/srm/config.py +++ b/src/srm/config.py @@ -33,6 +33,7 @@ class NatsSettings(BaseModel): url: str connect_timeout: int max_reconnect_attempts: int + drain_timeout: int = 30 class Settings(BaseSettings): diff --git a/src/srm/main.py b/src/srm/main.py index bfa7973..4e4e5f7 100644 --- a/src/srm/main.py +++ b/src/srm/main.py @@ -54,11 +54,22 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: except Exception as e: logger.error("Database engine init failed!", error=str(e)) raise - databus_manager: NatsConnectionManager = await init_databus_manager( - settings=settings.nats_settings - ) + 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 - databus_subscribers: list[NatsSubscriber] = await subscribe_to_subjects(databus_manager) + 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 @@ -68,9 +79,11 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: yield logger.info("Shutting down application") - await engine.dispose() + # 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() def create_app() -> FastAPI: diff --git a/tests/api/databus/test_nats_subscriber.py b/tests/api/databus/test_nats_subscriber.py index 07f9177..f10bc96 100644 --- a/tests/api/databus/test_nats_subscriber.py +++ b/tests/api/databus/test_nats_subscriber.py @@ -115,9 +115,10 @@ async def test_handle_message_logs_dropped_command_when_router_fails( router.assert_awaited_once() assert [(entry["event"], entry["subject"], entry["log_level"]) for entry in logs] == [ - ("databus_command_dropped", "command.srm.service.deploy", "error") + ("databus_handler_failed", "command.srm.service.deploy", "error") ] + EXPECTED_COMMAND_SUBJECTS = [ "command.srm.service.deploy", "command.srm.service.scale", diff --git a/tests/api/test_app.py b/tests/api/test_app.py index 4096bab..8d302e8 100644 --- a/tests/api/test_app.py +++ b/tests/api/test_app.py @@ -44,11 +44,15 @@ class TestLifespanDatabusInit: 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), + patch("srm.main.init_databus_manager", new_callable=AsyncMock, return_value=manager), patch( "srm.main.subscribe_to_subjects", new_callable=AsyncMock, @@ -59,6 +63,9 @@ class TestLifespanDatabusInit: 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: @@ -78,3 +85,21 @@ class TestLifespanDatabusInit: 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/integration/test_databus.py b/tests/integration/test_databus.py index fb581ed..82569c0 100644 --- a/tests/integration/test_databus.py +++ b/tests/integration/test_databus.py @@ -13,7 +13,6 @@ 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 diff --git a/tests/unit/test_nats_connection_manager.py b/tests/unit/test_nats_connection_manager.py index 1f70a7f..90d2e00 100644 --- a/tests/unit/test_nats_connection_manager.py +++ b/tests/unit/test_nats_connection_manager.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio from typing import Any from unittest.mock import AsyncMock, patch @@ -15,7 +16,17 @@ from srm.config import NatsSettings @pytest.fixture def settings() -> NatsSettings: - return NatsSettings(url="nats://broker:4222", connect_timeout=3, max_reconnect_attempts=2) + 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: @@ -43,6 +54,7 @@ async def test_connect_uses_settings(settings: NatsSettings) -> None: 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"]) @@ -50,7 +62,6 @@ async def test_connect_uses_settings(settings: NatsSettings) -> None: class TestConnectionCallbacks: - @staticmethod async def _connect_and_capture_callbacks(settings: NatsSettings) -> dict[str, Any]: with patch( @@ -109,6 +120,27 @@ async def test_connect_is_noop_when_already_connected(settings: NatsSettings) -> 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", @@ -133,6 +165,23 @@ async def test_close_drains_client_and_clears_state(settings: NatsSettings) -> N 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() -- GitLab From e461bb2679e5fa684ccf1431b51abf94747de664 Mon Sep 17 00:00:00 2001 From: dgogos Date: Mon, 3 Aug 2026 18:20:32 +0300 Subject: [PATCH 6/8] fix: remove skipped test for durable NATS subscriber and clean up code --- tests/api/databus/test_nats_subscriber.py | 27 ----------------------- tests/api/test_health.py | 2 ++ 2 files changed, 2 insertions(+), 27 deletions(-) diff --git a/tests/api/databus/test_nats_subscriber.py b/tests/api/databus/test_nats_subscriber.py index f10bc96..65a97ef 100644 --- a/tests/api/databus/test_nats_subscriber.py +++ b/tests/api/databus/test_nats_subscriber.py @@ -140,30 +140,3 @@ async def test_subscribe_to_subjects_registers_all_command_subjects( assert sorted(sub._subject for sub in subscribers) == sorted(EXPECTED_COMMAND_SUBJECTS) assert connection_manager.client.subscribe.await_count == len(EXPECTED_COMMAND_SUBJECTS) - - -class TestCommandDeliveryDurability: - @pytest.mark.skip( - reason="TODO: JetStream durable consumer not implemented yet — " - "NatsSubscriber still uses core-NATS subscribe() (at-most-once)." - ) - async def test_subscriber_subscribes_durably_via_jetstream( - self, connection_manager: MagicMock - ) -> None: - """SRM must consume command.srm.* durably, at-least-once (the OOP_TASKS WorkQueue - stream, interface-contract.md §D.1-D.2). A core-NATS subscribe() is non-durable and - at-most-once: every command published while SRM is down or slow is lost with no - redelivery, and OEG/FM never see a terminal event for it.""" - jetstream = MagicMock() - jetstream.subscribe = AsyncMock() - connection_manager.client.jetstream.return_value = jetstream - - subscriber, _ = make_subscriber(connection_manager) - await subscriber.start() - - connection_manager.client.subscribe.assert_not_awaited() - jetstream.subscribe.assert_awaited_once() - assert jetstream.subscribe.await_args is not None - assert jetstream.subscribe.await_args.kwargs.get("durable"), ( - "the contract requires a named durable consumer" - ) diff --git a/tests/api/test_health.py b/tests/api/test_health.py index 9476006..c25e238 100644 --- a/tests/api/test_health.py +++ b/tests/api/test_health.py @@ -1,3 +1,4 @@ +import pytest from fastapi import FastAPI from httpx import ASGITransport, AsyncClient @@ -19,6 +20,7 @@ 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 -- GitLab From 9e6aaab0ac72d10e1b4b207ba578af368436eb9c Mon Sep 17 00:00:00 2001 From: Dimitrios Gogos Date: Tue, 4 Aug 2026 07:29:01 +0000 Subject: [PATCH 7/8] fix: remove tags due to new gitlab runner --- .gitlab-ci.yml | 6 ------ 1 file changed, 6 deletions(-) diff --git a/.gitlab-ci.yml b/.gitlab-ci.yml index 4c98851..01c4d0b 100644 --- a/.gitlab-ci.yml +++ b/.gitlab-ci.yml @@ -1,6 +1,4 @@ default: - tags: - - docker image: python:3.12-slim cache: paths: @@ -57,8 +55,6 @@ test: build: stage: build - tags: - - shell before_script: - docker info script: @@ -69,8 +65,6 @@ build: push: stage: push - tags: - - shell needs: - build before_script: -- GitLab From 62768689a008566cf2030380b0b4f255305dc3bf Mon Sep 17 00:00:00 2001 From: dgogos Date: Tue, 4 Aug 2026 12:58:45 +0300 Subject: [PATCH 8/8] refactor: remove push stage and consolidate build script in GitLab CI configuration --- .gitlab-ci.yml | 22 ++++++++-------------- 1 file changed, 8 insertions(+), 14 deletions(-) diff --git a/.gitlab-ci.yml b/.gitlab-ci.yml index 01c4d0b..1e9ac23 100644 --- a/.gitlab-ci.yml +++ b/.gitlab-ci.yml @@ -17,7 +17,6 @@ stages: - format - test - build - - push variables: UV_CACHE_DIR: "$CI_PROJECT_DIR/.cache/uv" @@ -55,24 +54,19 @@ test: build: stage: build - 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 - 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"' -- GitLab