Files

447 lines
15 KiB
Python

"""Pytest configuration and fixtures for orchestrator E2E tests."""
import asyncio
import time
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from unittest.mock import MagicMock
import pytest
import pytest_asyncio
from google.protobuf.duration_pb2 import Duration
from pymongo import MongoClient
from redis import Redis
from sientia_do.notifications.handlers import CoreNotificationHandler
from sientia_do.observability.logger import Logger
from sientia_do.temporal.worker.prepare_worker import build_queue_name
from sqlalchemy import create_engine
from temporalio.api.enums.v1 import IndexedValueType
from temporalio.api.operatorservice.v1 import AddSearchAttributesRequest
from temporalio.api.workflowservice.v1 import (
DescribeNamespaceRequest,
RegisterNamespaceRequest,
)
from temporalio.client import Client
from temporalio.common import SearchAttributeKey
from temporalio.service import RPCError, RPCStatusCode
from temporalio.testing import WorkflowEnvironment
from temporalio.worker import Worker
from testcontainers.core.container import DockerContainer
from testcontainers.postgres import PostgresContainer
from e2e.helpers import MONGO_COLLECTIONS, ORCHESTRATOR_TASK_QUEUE
from e2e.smtp_test_server import SmtpTestServer
from e2e.stub_workflows import STUB_WORKFLOW_CLASSES
from orchestrator.activities.activities import Activities
from orchestrator.activities.formatters import schedule_types
from orchestrator.workflows.alerts import Alerts
from orchestrator.workflows.orchestrator import Orchestrator
from orchestrator.workflows.reports import Reports
from orchestrator.workflows.subworkflows.load_notification_package import (
LoadNotificationPackage,
)
from orchestrator.workflows.subworkflows.process_notifications import ProcessNotifications
DB_SCHEMA_SQL_PATH = Path(__file__).parent / 'db_schema.sql'
E2E_DATABASE = 'orchestrator_test'
E2E_RUNTIMES = ('legacy', 'gpu')
MANAGED_NAMESPACES = ('scouter', 'laborious')
E2E_SEARCH_ATTRIBUTES = [
SearchAttributeKey.for_keyword('model_id'),
SearchAttributeKey.for_keyword('model_name'),
SearchAttributeKey.for_keyword('orchestrated'),
]
E2E_NAMESPACE_SEARCH_ATTRIBUTES = {
'model_id': IndexedValueType.INDEXED_VALUE_TYPE_KEYWORD,
'model_name': IndexedValueType.INDEXED_VALUE_TYPE_KEYWORD,
'orchestrated': IndexedValueType.INDEXED_VALUE_TYPE_KEYWORD,
}
async def register_namespace_if_missing(env: WorkflowEnvironment, namespace: str) -> None:
"""
Register a Temporal namespace on the local dev server and wait until it is ready.
Args:
env: Session WorkflowEnvironment from start_local().
namespace: Namespace name to register.
"""
service = env.client.service_client
try:
await service.workflow_service.register_namespace(
RegisterNamespaceRequest(
namespace=namespace,
workflow_execution_retention_period=Duration(seconds=86400),
)
)
except RPCError as err:
if err.status != RPCStatusCode.ALREADY_EXISTS:
raise
deadline = time.monotonic() + 5.0
while time.monotonic() < deadline:
try:
await service.workflow_service.describe_namespace(
DescribeNamespaceRequest(namespace=namespace)
)
return
except RPCError:
await asyncio.sleep(0.1)
raise TimeoutError(f'Namespace {namespace} not ready within 5s')
async def ensure_namespace_search_attributes(
env: WorkflowEnvironment, namespace: str
) -> None:
"""
Register the orchestrator search attributes on a namespace, if missing.
The local Temporal dev server only registers search attributes on the default
namespace at start time. Schedules created in additional namespaces fail with
"no mapping defined for search attribute ..." unless we explicitly add the
same attribute mappings to those namespaces via the operator service.
Args:
env: Session WorkflowEnvironment from start_local().
namespace: Namespace where attributes must be available.
"""
service = env.client.service_client
try:
await service.operator_service.add_search_attributes(
AddSearchAttributesRequest(
namespace=namespace,
search_attributes=dict(E2E_NAMESPACE_SEARCH_ATTRIBUTES),
)
)
except RPCError as err:
if err.status != RPCStatusCode.ALREADY_EXISTS:
raise
def temporal_host_from_env(env: WorkflowEnvironment) -> str:
"""Return target host:port for the in-process Temporal dev server."""
return env.client.service_client.config.target_host
def mongo_uri_from_container(mongo_container) -> str:
"""Build a Mongo connection string for the testcontainer."""
port = mongo_container.get_exposed_port(27017)
return f'mongodb://localhost:{port}'
@pytest_asyncio.fixture(scope='session')
def postgres_container():
"""PostgreSQL testcontainer used by all E2E tests."""
postgres = PostgresContainer('postgres:15')
postgres.start()
yield postgres
postgres.stop()
@pytest_asyncio.fixture(scope='session')
def mongo_container():
"""MongoDB testcontainer used by real CoreNotificationHandler."""
mongo = DockerContainer('mongo:7').with_exposed_ports(27017)
mongo.start()
yield mongo
mongo.stop()
@pytest_asyncio.fixture(scope='session')
def redis_container():
"""Redis testcontainer for slot and notification timestamp paths."""
redis = DockerContainer('redis:7').with_exposed_ports(6379)
redis.start()
yield redis
redis.stop()
@pytest_asyncio.fixture
def postgres_engine(postgres_container):
"""SQLAlchemy engine bound to the PostgreSQL testcontainer."""
engine = create_engine(postgres_container.get_connection_url())
yield engine
engine.dispose()
def _create_schema_and_tables(engine):
sql_text = DB_SCHEMA_SQL_PATH.read_text(encoding='utf-8')
with engine.begin() as conn:
conn.exec_driver_sql(sql_text)
@pytest_asyncio.fixture(autouse=True)
def setup_postgres_schema_and_tables(postgres_engine):
"""Recreate Postgres schema from e2e/db_schema.sql before each test."""
_create_schema_and_tables(postgres_engine)
yield
@pytest.fixture
def mongo_uri(mongo_container):
return mongo_uri_from_container(mongo_container)
@pytest_asyncio.fixture(autouse=True)
def reset_mongo_collections(mongo_uri):
"""Drop orchestrator-managed Mongo collections between tests."""
client = MongoClient(mongo_uri)
try:
db = client[E2E_DATABASE]
for name in MONGO_COLLECTIONS:
db[name].drop()
finally:
client.close()
yield
@pytest.fixture
def redis_client(redis_container):
"""Redis client bound to the testcontainer."""
port = int(redis_container.get_exposed_port(6379))
client = Redis(host='localhost', port=port, decode_responses=True)
yield client
client.close()
@pytest_asyncio.fixture(autouse=True)
def reset_redis(redis_client):
"""Flush Redis between tests."""
redis_client.flushdb()
yield
@pytest.fixture(scope='session')
def smtp_server():
"""Session-scoped in-process SMTP server."""
server = SmtpTestServer()
server.start()
yield server
server.stop()
@pytest.fixture(autouse=True)
def smtp_messages_cleanup(smtp_server):
"""Clear captured SMTP messages between tests."""
smtp_server.clear()
yield
@pytest_asyncio.fixture(scope='session')
async def temporal_env():
"""Real Temporal dev server (schedule APIs supported)."""
env = await WorkflowEnvironment.start_local(search_attributes=E2E_SEARCH_ATTRIBUTES)
for namespace in MANAGED_NAMESPACES:
await register_namespace_if_missing(env, namespace)
await ensure_namespace_search_attributes(env, namespace)
yield env
await env.shutdown()
@pytest.fixture(scope='session')
def temporal_host(temporal_env):
return temporal_host_from_env(temporal_env)
@pytest_asyncio.fixture(autouse=True)
async def cleanup_temporal_schedules(temporal_env):
"""Delete orphan schedules in scouter/laborious before each test."""
host = temporal_host_from_env(temporal_env)
for namespace in MANAGED_NAMESPACES:
client = await Client.connect(host, namespace=namespace)
async for schedule in await client.list_schedules():
handle = client.get_schedule_handle(schedule.id)
await handle.delete()
yield
@pytest_asyncio.fixture
def mock_logger():
"""Logger double with readable console output for E2E runs."""
logger = MagicMock(spec=Logger)
logger.info = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}'))
logger.debug = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}'))
logger.error = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}'))
logger.warning = MagicMock(side_effect=lambda msg: print(f'[LOG] {msg}'))
logger.custom_info = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}'))
logger.custom_debug = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}'))
logger.custom_error = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}'))
logger.custom_warning = MagicMock(side_effect=lambda msg, _meta=None: print(f'[LOG] {msg}'))
return logger
@pytest_asyncio.fixture
async def notification_handler(mock_logger, mongo_container):
"""Real notification handler using MongoDB testcontainer."""
handler = CoreNotificationHandler(
connection_string=mongo_uri_from_container(mongo_container),
database=E2E_DATABASE,
logger=mock_logger,
project_name='orchestrator-e2e',
)
try:
yield handler
finally:
handler.shutdown()
@pytest.fixture
def notification_inserts(notification_handler):
"""Spy on real Mongo insert calls issued by notification handler."""
collection = notification_handler.mongo_collection
original_insert_one = collection.insert_one
spy = MagicMock(wraps=original_insert_one)
collection.insert_one = spy
try:
yield spy
finally:
collection.insert_one = original_insert_one
@pytest_asyncio.fixture
async def test_activities(
postgres_container,
mongo_container,
redis_container,
smtp_server,
mock_logger,
notification_handler,
temporal_host,
):
"""Real Activities wired to testcontainers and in-process SMTP."""
mongo_port = mongo_container.get_exposed_port(27017)
redis_port = int(redis_container.get_exposed_port(6379))
activities = Activities(
temporal_config={
'temporal_host': temporal_host,
'temporal_scouter_namespace': 'scouter',
'temporal_laborious_namespace': 'laborious',
},
redis_config={
'host': 'localhost',
'port': redis_port,
'username': '',
'password': '',
},
mongodb_config={
'connection_string': f'mongodb://localhost:{mongo_port}',
'database_name': E2E_DATABASE,
'ttl_index_seconds': 3600,
},
email_config={
'sender_email': 'e2e@example.com',
'sender_password': '',
'smtp_server': smtp_server.host,
'smtp_port': smtp_server.port,
},
postgres_config={
'host': 'localhost',
'port': int(postgres_container.get_exposed_port(5432)),
'user': postgres_container.username,
'password': postgres_container.password,
'dbname': postgres_container.dbname,
'min_connections': 1,
'max_connections': 5,
},
logger=mock_logger,
notification_handler=notification_handler,
)
await activities.connect_to_temporal()
try:
yield activities
finally:
activities.shutdown()
def _orchestrator_activity_list(activities: Activities) -> list:
return [
activities.load_active_ingestors,
activities.load_opc_slots,
activities.update_slots,
activities.delete_slots,
activities.aggregate_documents_in_mongodb,
activities.find_documents_in_mongodb,
activities.update_pipelines_timestamps,
activities.create_pipelines_timestamps,
activities.delete_pipelines_timestamps,
activities.create_collection_with_ttl_index,
activities.create_schedules,
activities.update_schedules,
activities.delete_schedules,
activities.normalize_schedules,
activities.process_schedules,
activities.process_slots,
activities.create_schedule_config,
activities.create_slot_config,
activities.report_schedule_orchestration,
activities.report_slot_orchestration,
activities.format_schedule_config,
activities.get_last_data_timestamp,
activities.load_latest_data,
activities.put_last_data_timestamp,
activities.filter_notification_alerts,
activities.filter_notification_reports,
activities.build_email_html,
activities.send_email,
activities.format_log_report,
activities.export_data_to_postgres,
activities.store_notification_cache,
]
@pytest_asyncio.fixture
async def orchestrator_worker(temporal_env, test_activities):
"""Worker for orchestrator workflows and all activities on the default namespace."""
with ThreadPoolExecutor(max_workers=32) as activity_executor:
async with Worker(
temporal_env.client,
task_queue=ORCHESTRATOR_TASK_QUEUE,
workflows=[
Orchestrator,
Alerts,
Reports,
LoadNotificationPackage,
ProcessNotifications,
],
activities=_orchestrator_activity_list(test_activities),
activity_executor=activity_executor,
) as worker:
yield worker
@pytest_asyncio.fixture
async def stub_workers(temporal_env):
"""No-op workers on scouter/laborious namespaces for every managed workflow type."""
host = temporal_host_from_env(temporal_env)
worker_contexts: list[Worker] = []
clients: list[Client] = []
stub_types = list(schedule_types.keys()) + [
'xgboost_predictions_batch',
'xgboost_minimal_retrain',
]
try:
for namespace in MANAGED_NAMESPACES:
client = await Client.connect(host, namespace=namespace)
clients.append(client)
queues = {
build_queue_name(workflow_type, runtime)
for workflow_type in stub_types
for runtime in E2E_RUNTIMES
}
for queue in queues:
worker = Worker(
client,
task_queue=queue,
workflows=STUB_WORKFLOW_CLASSES,
)
await worker.__aenter__()
worker_contexts.append(worker)
yield worker_contexts
finally:
for worker in reversed(worker_contexts):
await worker.__aexit__(None, None, None)