Commit ac5e755f authored by George Papathanail's avatar George Papathanail
Browse files

feat: add TrafficInfluence persistence layer

parent d8e6fff7
Loading
Loading
Loading
Loading
+36 −0
Original line number Diff line number Diff line
@@ -5,6 +5,7 @@ from open_exposure_gateway.adapters.database.sql import (
    CallbackRegistrationRow,
    OperationRow,
    QodSessionRow,
    TrafficInfluenceRow,
)
from open_exposure_gateway.domain.models import (
    AppInstance,
@@ -13,6 +14,7 @@ from open_exposure_gateway.domain.models import (
    CallbackRegistration,
    Operation,
    QodSession,
    TrafficInfluence,
)


@@ -138,6 +140,40 @@ class QodSessionMapper:
        )


class TrafficInfluenceMapper:
    @staticmethod
    def to_domain(row: TrafficInfluenceRow) -> TrafficInfluence:
        return TrafficInfluence(
            traffic_influence_id=row.traffic_influence_id,
            operation_id=row.operation_id,
            app_id=row.app_id,
            edge_cloud_zone_id=row.edge_cloud_zone_id,
            edge_cloud_region=row.edge_cloud_region,
            source_port=row.source_port,
            destination_port=row.destination_port,
            destination_protocol=row.destination_protocol,
            state=row.state,
            external_ref=row.external_ref,
            created_at=row.created_at,
            updated_at=row.updated_at,
        )

    @staticmethod
    def to_row(domain: TrafficInfluence) -> TrafficInfluenceRow:
        return TrafficInfluenceRow(
            traffic_influence_id=domain.traffic_influence_id,
            operation_id=domain.operation_id,
            app_id=domain.app_id,
            edge_cloud_zone_id=domain.edge_cloud_zone_id,
            edge_cloud_region=domain.edge_cloud_region,
            source_port=domain.source_port,
            destination_port=domain.destination_port,
            destination_protocol=domain.destination_protocol,
            state=domain.state,
            external_ref=domain.external_ref,
        )


class CallbackRegistrationMapper:
    @staticmethod
    def to_domain(row: CallbackRegistrationRow) -> CallbackRegistration:
+34 −0
Original line number Diff line number Diff line
from uuid import UUID

from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession

from open_exposure_gateway.adapters.database.mappers import TrafficInfluenceMapper
from open_exposure_gateway.adapters.database.sql import TrafficInfluenceRow
from open_exposure_gateway.domain.models import TrafficInfluence
from open_exposure_gateway.ports.database.traffic_influences import TrafficInfluenceRepository


class SqlTrafficInfluenceRepository(TrafficInfluenceRepository):
    def __init__(self, session: AsyncSession) -> None:
        self._session = session

    async def get_by_id(self, traffic_influence_id: UUID) -> TrafficInfluence | None:
        stmt = select(TrafficInfluenceRow).where(
            TrafficInfluenceRow.traffic_influence_id == traffic_influence_id
        )
        row = await self._session.scalar(stmt)
        return TrafficInfluenceMapper.to_domain(row) if row is not None else None

    async def get_by_operation_id(self, operation_id: UUID) -> TrafficInfluence | None:
        stmt = select(TrafficInfluenceRow).where(TrafficInfluenceRow.operation_id == operation_id)
        row = await self._session.scalar(stmt)
        return TrafficInfluenceMapper.to_domain(row) if row is not None else None

    async def save(self, traffic_influence: TrafficInfluence) -> TrafficInfluence:
        merged = await self._session.merge(TrafficInfluenceMapper.to_row(traffic_influence))
        await self._session.flush()
        saved = await self.get_by_id(merged.traffic_influence_id)
        if saved is None:
            raise RuntimeError("Saved traffic influence could not be reloaded")
        return saved
+21 −0
Original line number Diff line number Diff line
@@ -31,6 +31,7 @@ from open_exposure_gateway.domain.models.registration.enums import (
    AppRegistrationStatus,
    PackageType,
)
from open_exposure_gateway.domain.models.traffic_influences.enums import TrafficInfluenceState


def get_metadata() -> MetaData:
@@ -200,3 +201,23 @@ class QodSessionRow(AuditedMixin, Base):
    duration_seconds: Mapped[int] = mapped_column(Integer, nullable=False)
    state: Mapped[QodSessionState] = mapped_column(_enum_type(QodSessionState), nullable=False)
    external_ref: Mapped[str | None] = mapped_column(String(255))


class TrafficInfluenceRow(AuditedMixin, Base):
    __tablename__ = "traffic_influences"
    __table_args__ = (Index("idx_traffic_influences_operation", "operation_id"),)

    traffic_influence_id: Mapped[UUID] = mapped_column(PG_UUID(as_uuid=True), primary_key=True)
    operation_id: Mapped[UUID] = mapped_column(
        ForeignKey("operations.operation_id"), nullable=False
    )
    app_id: Mapped[UUID] = mapped_column(PG_UUID(as_uuid=True), nullable=False)
    edge_cloud_zone_id: Mapped[UUID | None] = mapped_column(PG_UUID(as_uuid=True))
    edge_cloud_region: Mapped[str | None] = mapped_column(String(256))
    source_port: Mapped[int | None] = mapped_column(Integer)
    destination_port: Mapped[int | None] = mapped_column(Integer)
    destination_protocol: Mapped[str | None] = mapped_column(String(32))
    state: Mapped[TrafficInfluenceState] = mapped_column(
        _enum_type(TrafficInfluenceState), nullable=False
    )
    external_ref: Mapped[str | None] = mapped_column(String(255))
+6 −0
Original line number Diff line number Diff line
@@ -14,6 +14,10 @@ from open_exposure_gateway.domain.models.registration import (
    AppRegistrationStatus,
    PackageType,
)
from open_exposure_gateway.domain.models.traffic_influences import (
    TrafficInfluence,
    TrafficInfluenceState,
)

__all__ = [
    "AppInstance",
@@ -28,4 +32,6 @@ __all__ = [
    "PackageType",
    "QodSession",
    "QodSessionState",
    "TrafficInfluence",
    "TrafficInfluenceState",
]
+4 −0
Original line number Diff line number Diff line
from open_exposure_gateway.domain.models.traffic_influences.enums import TrafficInfluenceState
from open_exposure_gateway.domain.models.traffic_influences.models import TrafficInfluence

__all__ = ["TrafficInfluence", "TrafficInfluenceState"]
Loading