Commit 815bfd1c authored by Sergio Gimenez's avatar Sergio Gimenez
Browse files

feat(fm): notify partners after federated deploy

Consume SRM operation.completed events from OOP_EVENTS with an explicit durable ack. Finalize the matching federation transaction and notify the partner once per instance using the fed-mgmt-notif token scope. Record callback delivery status and preserve the federation context in the transaction for delivery-time payload construction.
parent 1375ead9
Loading
Loading
Loading
Loading
+26 −9
Original line number Diff line number Diff line
from datetime import datetime
from uuid import UUID

from sqlalchemy import insert, select, update
from sqlalchemy import Select, insert, select, update
from sqlalchemy.ext.asyncio import AsyncSession

from federation_manager.adapters.database.tables import federation_transactions as transactions
@@ -45,15 +45,21 @@ class PostgresTransactionRepo:
    async def find_by_idempotency_key(
        self, partner_id: UUID, api_type: str, idempotency_key: str
    ) -> FederationTransaction | None:
        row = (
            await self._session.execute(
        return await self._one(
            select(transactions).where(
                transactions.c.partner_op_id == partner_id,
                transactions.c.api_type == api_type,
                transactions.c.idempotency_key == idempotency_key,
            )
        )
        ).one_or_none()

    async def find_by_operation_id(self, operation_id: UUID) -> FederationTransaction | None:
        return await self._one(
            select(transactions).where(transactions.c.operation_id == operation_id)
        )

    async def _one(self, stmt: Select[tuple[object, ...]]) -> FederationTransaction | None:
        row = (await self._session.execute(stmt)).one_or_none()
        if row is None:
            return None
        return FederationTransaction(
@@ -89,6 +95,17 @@ class PostgresTransactionRepo:
        )
        await self._session.commit()

    async def record_callback(self, transaction_id: UUID, *, status: str) -> None:
        await self._session.execute(
            update(transactions)
            .where(transactions.c.id == transaction_id)
            .values(
                callback_status=status,
                callback_attempts=transactions.c.callback_attempts + 1,
            )
        )
        await self._session.commit()

    async def record_outcome(
        self,
        transaction_id: UUID,
+82 −9
Original line number Diff line number Diff line
import json
from collections.abc import Awaitable, Callable
from typing import Any

import nats
from nats.aio.client import Client
from nats.aio.msg import Msg
from nats.js import JetStreamContext
from nats.js.api import RetentionPolicy, StreamConfig
from nats.js.api import AckPolicy, ConsumerConfig, RetentionPolicy, StreamConfig

from federation_manager.contracts.srm import TASK_STREAM
from federation_manager.contracts.srm import (
    EVENT_STREAM,
    TASK_STREAM,
    TASK_STREAM_MAX_AGE_SECONDS,
)


class NatsCommandPublisher:
@@ -26,16 +33,19 @@ class NatsCommandPublisher:

    async def ensure_task_stream(self) -> None:
        js = self._require_js()
        try:
            await js.stream_info(TASK_STREAM)
        except Exception:
            await js.add_stream(
                StreamConfig(
        config = StreamConfig(
            name=TASK_STREAM,
            subjects=["command.srm.>"],
            retention=RetentionPolicy.WORK_QUEUE,
            max_age=TASK_STREAM_MAX_AGE_SECONDS,
        )
            )
        try:
            await js.stream_info(TASK_STREAM)
        except Exception:
            await js.add_stream(config)
        else:
            # an older stream may predate the age limit, and unacked commands never expire
            await js.update_stream(config)

    async def publish(self, subject: str, payload: dict[str, object]) -> None:
        js = self._require_js()
@@ -45,3 +55,66 @@ class NatsCommandPublisher:
        if self._js is None:
            raise RuntimeError("NATS publisher is not connected")
        return self._js


class NatsEventConsumer:
    def __init__(self, url: str, durable: str = "fm-event-worker") -> None:
        self._url = url
        self._durable = durable
        self._nc: Client | None = None
        self._js: JetStreamContext | None = None

    async def connect(self) -> None:
        self._nc = await nats.connect(self._url)
        self._js = self._nc.jetstream()

    async def close(self) -> None:
        if self._nc is not None:
            await self._nc.drain()
            self._nc = None
            self._js = None

    async def ensure_event_stream(self) -> None:
        js = self._require_js()
        try:
            await js.stream_info(EVENT_STREAM)
        except Exception:
            await js.add_stream(
                StreamConfig(
                    name=EVENT_STREAM,
                    subjects=["event.srm.>"],
                    retention=RetentionPolicy.LIMITS,
                )
            )

    async def subscribe(
        self, subject: str, handler: Callable[[dict[str, Any]], Awaitable[None]]
    ) -> None:
        async def on_message(message: Msg) -> None:
            try:
                await handler(json.loads(message.data))
            except Exception:
                # Leave it unacked so JetStream redelivers up to max_deliver.
                return
            await message.ack()

        name = f"{self._durable}-{subject.replace('.', '-')}"
        # queue group so replicas of this FM share one durable instead of fighting over it
        await self._require_js().subscribe(
            subject,
            queue=name,
            durable=name,
            stream=EVENT_STREAM,
            cb=on_message,
            manual_ack=True,
            config=ConsumerConfig(max_deliver=3, ack_policy=AckPolicy.EXPLICIT),
        )

    async def delete_durable(self, subject: str) -> None:
        name = f"{self._durable}-{subject.replace('.', '-')}"
        await self._require_js().delete_consumer(EVENT_STREAM, name)

    def _require_js(self) -> JetStreamContext:
        if self._js is None:
            raise RuntimeError("NATS event consumer is not connected")
        return self._js
+41 −0
Original line number Diff line number Diff line
import httpx

from federation_manager.domain.errors import PartnerEndpointConfigurationError
from federation_manager.domain.models import PartnerOP
from federation_manager.domain.ports import PartnerTokenProviderPort

NOTIFICATION_SCOPE = "fed-mgmt-notif"


class HttpxCallbackClient:
    def __init__(
        self,
        client: httpx.AsyncClient,
        token_provider: PartnerTokenProviderPort,
        allow_insecure: bool = False,
    ) -> None:
        self._client = client
        self._token_provider = token_provider
        self._allow_insecure = allow_insecure

    async def deliver(self, partner: PartnerOP, url: str, payload: dict[str, object]) -> bool:
        self._check(partner, url)
        token = await self._token_provider.token_for(partner, scope=NOTIFICATION_SCOPE)
        try:
            response = await self._client.post(
                url,
                headers={"Accept": "application/json", "Authorization": f"Bearer {token}"},
                json=payload,
            )
        except httpx.HTTPError:
            return False
        return 200 <= response.status_code < 300

    def _check(self, partner: PartnerOP, url: str) -> None:
        try:
            parsed = httpx.URL(url)
        except httpx.InvalidURL:
            raise PartnerEndpointConfigurationError(partner.id) from None
        allowed = ("https", "http") if self._allow_insecure else ("https",)
        if not parsed.host or parsed.scheme not in allowed:
            raise PartnerEndpointConfigurationError(partner.id)
+1 −0
Original line number Diff line number Diff line
@@ -111,6 +111,7 @@ class InboundDeploymentService:
            callback_url=request.app_inst_callback_link,
            callback_status="pending",
            request_summary={
                "federation_context_id": federation_context_id,
                "app_id": request.app_id,
                "app_version": request.app_version,
                "flavour_id": request.zone_info.flavour_id,
+135 −0
Original line number Diff line number Diff line
from collections.abc import Callable
from datetime import datetime, timezone
from typing import Any

from federation_manager.contracts.ewbi import (
    AppInstanceInfo,
    InstanceState,
    InstanceStatusCallback,
)
from federation_manager.contracts.srm import CompletedInstanceV1, SrmOperationCompletedV1
from federation_manager.domain.models import FederationTransaction
from federation_manager.domain.ports import (
    CallbackClientPort,
    PartnerRepositoryPort,
    TransactionRepositoryPort,
)


def _utcnow() -> datetime:
    return datetime.now(timezone.utc)


class OperationCompletedConsumer:
    def __init__(
        self,
        transactions: TransactionRepositoryPort,
        partners: PartnerRepositoryPort,
        callbacks: CallbackClientPort,
        *,
        clock: Callable[[], datetime] = _utcnow,
    ) -> None:
        self._transactions = transactions
        self._partners = partners
        self._callbacks = callbacks
        self._clock = clock

    async def handle(self, payload: dict[str, Any]) -> None:
        event = SrmOperationCompletedV1.model_validate(payload)
        transaction = await self._transactions.find_by_operation_id(event.operation_id)
        if transaction is None:
            # Not ours: OEG-originated operations share the stream.
            return

        instances = [
            {
                "service_instance_id": str(instance.service_instance_id),
                "zone_id": str(instance.zone_id),
                "status": instance.status,
            }
            for instance in event.instances or []
        ]
        await self._transactions.record_outcome(
            transaction.id,
            status=event.status,
            completed_at=event.completed_at or self._clock(),
            response_summary={"instances": instances},
            error_detail=event.error,
        )
        await self._notify_partner(transaction, event)

    async def _notify_partner(
        self, transaction: FederationTransaction, event: SrmOperationCompletedV1
    ) -> None:
        if not transaction.callback_url:
            return
        partner = await self._partners.find_by_id(transaction.partner_op_id)
        if partner is None:
            return

        delivered = True
        for body in self._callback_bodies(transaction, event):
            payload = body.model_dump(mode="json", by_alias=True, exclude_none=True)
            delivered &= await self._callbacks.deliver(partner, transaction.callback_url, payload)
        await self._transactions.record_callback(
            transaction.id, status="delivered" if delivered else "failed"
        )

    def _callback_bodies(
        self, transaction: FederationTransaction, event: SrmOperationCompletedV1
    ) -> list[InstanceStatusCallback]:
        summary = transaction.request_summary
        context_id = str(summary.get("federation_context_id", ""))
        app_id = str(summary.get("app_id", ""))
        message = str(event.error.get("title")) if event.error else None

        if not event.instances:
            return [
                _callback(
                    context_id,
                    app_id,
                    str(transaction.external_resource_id),
                    str(summary.get("zone_id", "")),
                    "FAILED",
                    message,
                )
            ]
        return [
            _callback(
                context_id,
                app_id,
                instance.service_instance_id.hex,
                str(instance.zone_id),
                _instance_state(instance),
                _instance_message(instance) or message,
            )
            for instance in event.instances
        ]


def _instance_state(instance: CompletedInstanceV1) -> InstanceState:
    return "READY" if instance.status == "completed" else "FAILED"


def _instance_message(instance: CompletedInstanceV1) -> str | None:
    if not instance.error:
        return None
    title = instance.error.get("title")
    return str(title) if title is not None else None


def _callback(
    context_id: str,
    app_id: str,
    instance_id: str,
    zone_id: str,
    state: InstanceState,
    message: str | None,
) -> InstanceStatusCallback:
    return InstanceStatusCallback(
        federation_context_id=context_id,
        app_id=app_id,
        app_instance_id=instance_id,
        zone_id=zone_id,
        app_instance_info=AppInstanceInfo(app_instance_state=state, message=message),
    )
Loading