Commit 1065c335 authored by George Papathanail's avatar George Papathanail
Browse files

test(traffic-influence): add fakes and fixtures for flow tests

FakeSRMClient gains get_traffic_influence(s); FakeTrafficInfluenceRepository
mirrors FakeQodSessionRepository; wire_traffic_influence_operation_consumer
mirrors the QoD equivalent. conftest.py wires matching fixtures and
registers get_traffic_influence_service in api_client's overrides.
parent a579aa78
Loading
Loading
Loading
Loading
+42 −0
Original line number Diff line number Diff line
@@ -14,11 +14,15 @@ from open_exposure_gateway.application.services.edge_application_management_serv
from open_exposure_gateway.application.services.quality_on_demand_service import (
    QualityOnDemandService,
)
from open_exposure_gateway.application.services.traffic_influence_service import (
    TrafficInfluenceService,
)
from open_exposure_gateway.dependencies import (
    get_database_health,
    get_edge_app_service,
    get_publisher,
    get_qod_service,
    get_traffic_influence_service,
)
from open_exposure_gateway.main import app
from tests.unit.fakes import (
@@ -32,9 +36,11 @@ from tests.unit.fakes import (
    FakeQodCallbackDeliveryPort,
    FakeQodSessionRepository,
    FakeSRMClient,
    FakeTrafficInfluenceRepository,
    wire_operation_consumer,
    wire_qod_operation_consumer,
    wire_srm_worker,
    wire_traffic_influence_operation_consumer,
)


@@ -84,6 +90,11 @@ def app_instance_repo() -> FakeAppInstanceRepository:
    return FakeAppInstanceRepository()


@pytest.fixture()
def traffic_influence_repo() -> FakeTrafficInfluenceRepository:
    return FakeTrafficInfluenceRepository()


@pytest.fixture()
def callback_registration_repo() -> FakeCallbackRegistrationRepository:
    return FakeCallbackRegistrationRepository()
@@ -237,14 +248,45 @@ def register_app(api_client: TestClient) -> Callable[[Any], None]:
    return _register


@pytest.fixture()
def live_ti(
    fake_bus: FakeDataBus,
    operation_repo: FakeOperationRepository,
    traffic_influence_repo: FakeTrafficInfluenceRepository,
) -> None:
    """Wires Traffic Influence's real completion handler behind the fake bus, sharing
    the same operation_repo/traffic_influence_repo instances the api_client's
    traffic_influence_service writes to."""
    wire_traffic_influence_operation_consumer(fake_bus, operation_repo, traffic_influence_repo)


@pytest.fixture()
def traffic_influence_service(
    fake_srm: FakeSRMClient,
    fake_bus: FakeDataBus,
    operation_repo: FakeOperationRepository,
    traffic_influence_repo: FakeTrafficInfluenceRepository,
    callback_registration_repo: FakeCallbackRegistrationRepository,
) -> TrafficInfluenceService:
    return TrafficInfluenceService(
        srm_client=fake_srm,
        publisher=fake_bus,
        operation_repo=operation_repo,
        traffic_influence_repo=traffic_influence_repo,
        callback_registration_repo=callback_registration_repo,
    )


@pytest.fixture()
def api_client(
    eam_service: EdgeApplicationManagementService,
    qod_service: QualityOnDemandService,
    traffic_influence_service: TrafficInfluenceService,
    fake_bus: FakeDataBus,
) -> Generator[TestClient, None, None]:
    app.dependency_overrides[get_edge_app_service] = lambda: eam_service
    app.dependency_overrides[get_qod_service] = lambda: qod_service
    app.dependency_overrides[get_traffic_influence_service] = lambda: traffic_influence_service
    app.dependency_overrides[get_publisher] = lambda: fake_bus
    app.dependency_overrides[get_database_health] = lambda: True
    # raise_server_exceptions=False: unhandled errors surface as the 500 envelope
+79 −0
Original line number Diff line number Diff line
@@ -33,12 +33,18 @@ from open_exposure_gateway.adapters.errors import (
from open_exposure_gateway.api.camara.quality_on_demand.v0_10_1.schemas import (
    SessionInfo,
)
from open_exposure_gateway.api.camara.traffic_influence.vwip.schemas import (
    TrafficInfluence,
)
from open_exposure_gateway.application.services.edge_application_management_service import (
    EdgeApplicationManagementService,
)
from open_exposure_gateway.application.services.quality_on_demand_service import (
    QualityOnDemandService,
)
from open_exposure_gateway.application.services.traffic_influence_service import (
    TrafficInfluenceService,
)
from open_exposure_gateway.core.exceptions import NotFoundException
from open_exposure_gateway.domain.edge_application_management import (
    AppInstanceStatusChangeCloudEvent,
@@ -58,8 +64,12 @@ from open_exposure_gateway.domain.models import (
    Operation,
    QodSession,
)
from open_exposure_gateway.domain.models import (
    TrafficInfluence as TrafficInfluenceRecord,
)
from open_exposure_gateway.domain.quality_on_demand import QosStatusChangedCloudEvent
from open_exposure_gateway.domain.quality_on_demand import Subject as QodSubject
from open_exposure_gateway.domain.traffic_influence import Subject as TrafficInfluenceSubject
from open_exposure_gateway.ports.callback_delivery_port import CallbackDeliveryPort
from open_exposure_gateway.ports.database.callbacks import (
    CallbackDeliveryRepository,
@@ -69,6 +79,7 @@ from open_exposure_gateway.ports.database.instances import AppInstanceRepository
from open_exposure_gateway.ports.database.operations import OperationRepository
from open_exposure_gateway.ports.database.qod_sessions import QodSessionRepository
from open_exposure_gateway.ports.database.registration import AppRegistrationRepository
from open_exposure_gateway.ports.database.traffic_influences import TrafficInfluenceRepository
from open_exposure_gateway.ports.qod_callback_port import QodCallbackDeliveryPort

Handler = Callable[[dict[str, Any]], Awaitable[None]]
@@ -110,6 +121,7 @@ class FakeSRMClient:
        self.catalog: dict[str, dict[str, Any]] = {}
        self.instances: dict[str, SRMServiceInstance] = {}
        self.qod_sessions: dict[str, SessionInfo] = {}
        self.traffic_influences: dict[str, TrafficInfluence] = {}

    async def get_zones(
        self,
@@ -160,6 +172,22 @@ class FakeSRMClient:
            raise NotFoundException(message=f"Session {session_id} not found")
        return session

    async def get_traffic_influence(
        self, traffic_influence_id: str, x_correlator: str | None = None
    ) -> TrafficInfluence:
        influence = self.traffic_influences.get(traffic_influence_id)
        if influence is None:
            raise NotFoundException(message=f"Traffic influence {traffic_influence_id} not found")
        return influence

    async def get_traffic_influences(
        self, app_id: UUID | None = None, x_correlator: str | None = None
    ) -> list[TrafficInfluence]:
        result = list(self.traffic_influences.values())
        if app_id is not None:
            result = [t for t in result if t.appId == app_id]
        return result


def completion_payload(operation_id: str, **overrides: Any) -> dict[str, Any]:
    payload: dict[str, Any] = {
@@ -329,6 +357,37 @@ def wire_qod_operation_consumer(
    return completed_consumer


def wire_traffic_influence_operation_consumer(
    bus: FakeDataBus,
    operation_repo: OperationRepository,
    traffic_influence_repo: TrafficInfluenceRepository,
) -> NatsOperationConsumer:
    """Wires TrafficInfluenceService.handle_completed behind the fake bus -- pass the
    same repo instances used to build the service under test so a completion event
    updates the rows the test can see."""
    service = TrafficInfluenceService(
        srm_client=AsyncMock(),
        operation_repo=operation_repo,
        traffic_influence_repo=traffic_influence_repo,
    )
    consumer = NatsOperationConsumer(
        client=AsyncMock(),
        subject=TrafficInfluenceSubject.OPERATION_COMPLETED,
        handler=service.handle_completed,
    )

    async def deliver(payload: dict[str, Any]) -> None:
        await consumer._handle_message(
            FakeMsg(
                data=json.dumps(payload).encode(),
                subject=str(TrafficInfluenceSubject.OPERATION_COMPLETED),
            )
        )

    bus.subscribe(TrafficInfluenceSubject.OPERATION_COMPLETED, deliver)
    return consumer


class FakeAppRegistrationRepository(AppRegistrationRepository):
    def __init__(self) -> None:
        self.rows: dict[UUID, AppRegistration] = {}
@@ -451,6 +510,26 @@ class FakeQodSessionRepository(QodSessionRepository):
        return stored.model_copy(deep=True)


class FakeTrafficInfluenceRepository(TrafficInfluenceRepository):
    def __init__(self) -> None:
        self.rows: dict[UUID, TrafficInfluenceRecord] = {}

    async def get_by_id(self, traffic_influence_id: UUID) -> TrafficInfluenceRecord | None:
        found = self.rows.get(traffic_influence_id)
        return found.model_copy(deep=True) if found is not None else None

    async def get_by_operation_id(self, operation_id: UUID) -> TrafficInfluenceRecord | None:
        for row in self.rows.values():
            if row.operation_id == operation_id:
                return row.model_copy(deep=True)
        return None

    async def save(self, traffic_influence: TrafficInfluenceRecord) -> TrafficInfluenceRecord:
        stored = traffic_influence.model_copy(deep=True)
        self.rows[stored.traffic_influence_id] = stored
        return stored.model_copy(deep=True)


class FakeCallbackRegistrationRepository(CallbackRegistrationRepository):
    def __init__(self) -> None:
        self.rows: dict[UUID, CallbackRegistration] = {}