feat: enhance configuration and error handling in project setup

- Added new ignore rule for Ruff to allow temporary paths in tests.
- Introduced MyPy overrides for specific modules to ignore errors.
- Refactored `Cleanup` and `ExperimentTracking` classes to remove async keywords from methods, improving consistency in method signatures.
- Updated `Training` class methods to handle synchronous operations, enhancing performance and clarity.
- Adjusted `requirements.txt` to remove unnecessary Git dependency, streamlining project setup.
This commit is contained in:
vitor-aignosi
2026-04-07 10:25:17 -03:00
parent 6b1df7c3a7
commit 09ee92f100
21 changed files with 500 additions and 309 deletions

View File

@@ -67,7 +67,6 @@ class Activities(ExperimentTracking, Training, Cleanup):
Exception: If any parent class initialization fails Exception: If any parent class initialization fails
""" """
ExperimentTracking.__init__( ExperimentTracking.__init__(
self, self,
host=postgres_config['host'], host=postgres_config['host'],

View File

@@ -62,7 +62,7 @@ class Cleanup(SientiaMonitoring):
) # name_YYYYMMDD_HHMMSS_microseconds ) # name_YYYYMMDD_HHMMSS_microseconds
@activity.defn(name='cleanup_temp_directories') @activity.defn(name='cleanup_temp_directories')
async def cleanup_temp_directories(self, input_data: dict[str, Any]) -> None: def cleanup_temp_directories(self, input_data: dict[str, Any]) -> None:
""" """
Clean up stale temporary directories based on timestamp in directory name. Clean up stale temporary directories based on timestamp in directory name.
@@ -178,14 +178,14 @@ class Cleanup(SientiaMonitoring):
raise raise
finally: finally:
await self._emit_metrics( self._emit_metrics(
metadata=metadata, metadata=metadata,
metrics_status=metrics_status, metrics_status=metrics_status,
activity_name='cleanup_temp_directories', activity_name='cleanup_temp_directories',
emit_workflow_metric=True, emit_workflow_metric=True,
) )
async def _emit_metrics( def _emit_metrics(
self, self,
metadata: dict[str, Any], metadata: dict[str, Any],
metrics_status: str, metrics_status: str,
@@ -201,7 +201,7 @@ class Cleanup(SientiaMonitoring):
activity_name: Name of the activity being executed activity_name: Name of the activity being executed
""" """
if emit_workflow_metric: if emit_workflow_metric:
await self.emit_metric( self.emit_metric_sync(
metric_object=WORKFLOW_EXECUTION_TOTAL, metric_object=WORKFLOW_EXECUTION_TOTAL,
tags={ tags={
'pod_id': metadata.get('pod_id'), 'pod_id': metadata.get('pod_id'),
@@ -210,7 +210,7 @@ class Cleanup(SientiaMonitoring):
}, },
) )
await self.emit_metric( self.emit_metric_sync(
metric_object=ACTIVITY_EXECUTION_TOTAL, metric_object=ACTIVITY_EXECUTION_TOTAL,
tags={ tags={
'pod_id': metadata.get('pod_id'), 'pod_id': metadata.get('pod_id'),

View File

@@ -11,7 +11,6 @@ import enum
from temporalio import activity, workflow from temporalio import activity, workflow
with workflow.unsafe.imports_passed_through(): with workflow.unsafe.imports_passed_through():
import asyncio
import traceback import traceback
from collections.abc import Mapping from collections.abc import Mapping
from datetime import UTC, datetime from datetime import UTC, datetime
@@ -110,9 +109,9 @@ class ExperimentTracking(Postgres):
# Silently ignore errors during garbage collection # Silently ignore errors during garbage collection
pass pass
async def _execute_update(self, query: str, params: Mapping[str, Any]) -> dict[str, Any]: def _execute_update(self, query: str, params: Mapping[str, Any]) -> dict[str, Any]:
""" """
Execute an UPDATE SQL statement asynchronously. Execute an UPDATE SQL statement.
Args: Args:
query: Parameterized SQL string to execute. query: Parameterized SQL string to execute.
@@ -122,12 +121,9 @@ class ExperimentTracking(Postgres):
dict: A dictionary containing the affected row count: {'rowcount': int}. dict: A dictionary containing the affected row count: {'rowcount': int}.
""" """
def _run() -> dict[str, Any]: with self.engine.begin() as connection:
with self.engine.begin() as connection: result = connection.execute(text(query), params)
result = connection.execute(text(query), params) return {'rowcount': result.rowcount}
return {'rowcount': result.rowcount}
return await asyncio.to_thread(_run)
def _build_status_update_query( def _build_status_update_query(
self, status: str | None, experiment_run_id: int self, status: str | None, experiment_run_id: int
@@ -223,7 +219,7 @@ class ExperimentTracking(Postgres):
raise ValueError(f'Invalid update_type: {update_type}') raise ValueError(f'Invalid update_type: {update_type}')
@activity.defn(name='update_experiment_run') @activity.defn(name='update_experiment_run')
async def update_experiment_run(self, input_data: dict[str, Any]) -> None: def update_experiment_run(self, input_data: dict[str, Any]) -> None:
""" """
Update experiment run with status, errors, or model information. Update experiment run with status, errors, or model information.
@@ -256,7 +252,7 @@ class ExperimentTracking(Postgres):
update_type, experiment_run_id, input_data update_type, experiment_run_id, input_data
) )
result = await self._execute_update(sql_query, query_params) result = self._execute_update(sql_query, query_params)
if result.get('rowcount', 0) == 0: if result.get('rowcount', 0) == 0:
error_msg = ( error_msg = (
@@ -273,7 +269,7 @@ class ExperimentTracking(Postgres):
error_msg = f'Error updating experiment run - ID: {experiment_run_id}, Status: {status}, Error: {str(e)}' error_msg = f'Error updating experiment run - ID: {experiment_run_id}, Status: {status}, Error: {str(e)}'
trace = traceback.format_exc() trace = traceback.format_exc()
await self.send_notification_async( self.send_notification(
metadata=metadata or {}, metadata=metadata or {},
notification_id='UPDATE_EXPERIMENT_RUN_ERROR', notification_id='UPDATE_EXPERIMENT_RUN_ERROR',
message=error_msg, message=error_msg,

View File

@@ -12,6 +12,7 @@ with workflow.unsafe.imports_passed_through():
import traceback import traceback
from typing import Any from typing import Any
import mlflow
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
@@ -22,10 +23,8 @@ with workflow.unsafe.imports_passed_through():
from sientia_model.model_repository.plugin_store import PluginStore from sientia_model.model_repository.plugin_store import PluginStore
from model_manager.utils.models.train_model_params import TrainModelParams from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.repository.data_manager_repository import DataManagerRepository
from model_manager.utils.models.train_model_result import TrainModelResult from model_manager.utils.models.train_model_result import TrainModelResult
from model_manager.utils.repository.data_manager_repository import DataManagerRepository
import mlflow
class Training(SientiaMonitoring): class Training(SientiaMonitoring):
@@ -61,7 +60,7 @@ class Training(SientiaMonitoring):
self.minio_repository = minio_repository self.minio_repository = minio_repository
@activity.defn(name='load_model_metadata') @activity.defn(name='load_model_metadata')
async def load_model_metadata(self, input_data: dict[str, Any]) -> dict[str, Any]: def load_model_metadata(self, input_data: dict[str, Any]) -> dict[str, Any]:
""" """
Load model metadata/schemas from the model store. Load model metadata/schemas from the model store.
@@ -90,7 +89,7 @@ class Training(SientiaMonitoring):
return train_params.to_dict() return train_params.to_dict()
except Exception as exc: except Exception as exc:
trace = traceback.format_exc() trace = traceback.format_exc()
await self.send_notification_async( self.send_notification(
metadata=metadata, metadata=metadata,
notification_id='LOAD_MODEL_METADATA_ERROR', notification_id='LOAD_MODEL_METADATA_ERROR',
message=f'Error loading model metadata: {str(exc)}', message=f'Error loading model metadata: {str(exc)}',
@@ -99,9 +98,9 @@ class Training(SientiaMonitoring):
attachment_content=trace, attachment_content=trace,
) )
raise raise
@activity.defn(name='validate_train_params') @activity.defn(name='validate_train_params')
async def validate_train_params(self, input_data: dict[str, Any]) -> TrainModelParams: def validate_train_params(self, input_data: dict[str, Any]) -> dict[str, Any]:
""" """
Validate and convert training parameters from dict to TrainModelParams. Validate and convert training parameters from dict to TrainModelParams.
@@ -115,7 +114,7 @@ class Training(SientiaMonitoring):
- All TrainModelParams fields (experiment_run_id, target_variable, etc.) - All TrainModelParams fields (experiment_run_id, target_variable, etc.)
Returns: Returns:
TrainModelParams: Validated and converted training parameters dict[str, Any]: Validated and converted training parameters as dictionary
Raises: Raises:
Exception: If validation fails (after sending notification) Exception: If validation fails (after sending notification)
@@ -123,7 +122,7 @@ class Training(SientiaMonitoring):
metadata = input_data.get('metadata', {}) metadata = input_data.get('metadata', {})
try: try:
train_params = TrainModelParams.from_dict(input_data) train_params = TrainModelParams.from_dict(input_data)
train_params.validate_business_rules() train_params.validate_business_rules()
self.info( self.info(
@@ -133,12 +132,12 @@ class Training(SientiaMonitoring):
metadata, metadata,
) )
return train_params return train_params.to_dict()
except Exception as e: except Exception as e:
error_msg = f'Error validating training parameters: {str(e)}' error_msg = f'Error validating training parameters: {str(e)}'
trace = traceback.format_exc() trace = traceback.format_exc()
await self.send_notification_async( self.send_notification(
metadata=metadata, metadata=metadata,
notification_id='VALIDATE_TRAIN_PARAMS_ERROR', notification_id='VALIDATE_TRAIN_PARAMS_ERROR',
message=error_msg, message=error_msg,
@@ -149,7 +148,7 @@ class Training(SientiaMonitoring):
raise raise
@activity.defn(name='train_model') @activity.defn(name='train_model')
async def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]: def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
""" """
Train a machine learning model. Train a machine learning model.
@@ -162,7 +161,7 @@ class Training(SientiaMonitoring):
input_data: Training configuration containing: input_data: Training configuration containing:
- metadata (dict): Workflow execution metadata. - metadata (dict): Workflow execution metadata.
- uploaded_file (BytesIO): Training data already downloaded from MinIO. - uploaded_file (BytesIO): Training data already downloaded from MinIO.
- train_params (TrainModelParams | dict): Training parameters. - train_params (dict): Training parameters.
Returns: Returns:
dict[str, Any]: Serializable summary (run identifiers, run_dir for cleanup, regression metrics). dict[str, Any]: Serializable summary (run identifiers, run_dir for cleanup, regression metrics).
@@ -172,14 +171,11 @@ class Training(SientiaMonitoring):
Exception: If training fails (after sending notification). Exception: If training fails (after sending notification).
""" """
metadata = input_data.get('metadata') metadata = input_data.get('metadata')
train_params = input_data['train_params'] train_params = TrainModelParams.from_dict(input_data['train_params'])
if isinstance(train_params, dict):
train_params = TrainModelParams.from_dict(train_params)
try: try:
# Download training file bytes from MinIO # Download training file bytes from MinIO
train_bytes = await self.minio_repository.download_file( train_bytes = self.minio_repository.download_file_sync(
object_name=train_params.file_name, object_name=train_params.file_name,
bucket=train_params.bucket_name, bucket=train_params.bucket_name,
metadata=metadata, metadata=metadata,
@@ -189,7 +185,7 @@ class Training(SientiaMonitoring):
val_bytes: bytes | None = None val_bytes: bytes | None = None
validation_name = train_params.val_file_name validation_name = train_params.val_file_name
if validation_name is not None: if validation_name is not None:
val_bytes = await self.minio_repository.download_file( val_bytes = self.minio_repository.download_file_sync(
object_name=validation_name, object_name=validation_name,
bucket=train_params.bucket_name, bucket=train_params.bucket_name,
metadata=metadata, metadata=metadata,
@@ -202,13 +198,13 @@ class Training(SientiaMonitoring):
metadata=metadata, metadata=metadata,
) )
wrapper = await self.plugin_store.get_model( wrapper = self.plugin_store.get_model(
model_name=train_params.model_name, model_name=train_params.model_name,
force_download=False, force_download=False,
opt_params=train_params.opt_params or {}, opt_params=train_params.opt_params or {},
model_kwargs=train_params.model_kwargs or {}, model_kwargs=train_params.model_kwargs or {},
data_model_kwargs=train_params.data_model_kwargs or {}, data_model_kwargs=train_params.data_model_kwargs or {},
metadata=metadata metadata=metadata,
) )
train_data = train_result.train_data train_data = train_result.train_data
@@ -238,7 +234,7 @@ class Training(SientiaMonitoring):
wrapper, wrapper,
) )
async with self.mlflow_repository.start_run( with self.mlflow_repository.start_run(
model_name=train_params.model_name, model_name=train_params.model_name,
run_name=None, run_name=None,
experiment_name=f'{train_params.model_name}_experiment', experiment_name=f'{train_params.model_name}_experiment',
@@ -247,33 +243,19 @@ class Training(SientiaMonitoring):
) as run_info: ) as run_info:
train_result.run_name = run_info.run_name train_result.run_name = run_info.run_name
train_result.run_id = run_info.run_id train_result.run_id = run_info.run_id
self._persist_training_artifacts(train_result, train_params, wrapper, metadata)
train_result = self.data_manager_repository.generate_report(
train_result,
metadata=metadata,
)
if train_result.report_path is None or train_result.train_data_path is None or train_result.test_data_path is None:
raise ValueError('Report path, train data path, or test data path is not set')
wrapper.store_model(name=train_params.model_name)
mlflow.log_artifact(train_result.report_path)
mlflow.log_artifact(train_result.train_data_path)
mlflow.log_artifact(train_result.test_data_path)
return { return {
'run_name': train_result.run_name, 'run_name': train_result.run_name,
'run_id': train_result.run_id, 'run_id': train_result.run_id,
'run_dir': train_result.run_dir 'run_dir': train_result.run_dir,
} }
except Exception as e: # noqa: BLE001 except Exception as e: # noqa: BLE001
error_msg = f'Error training model - error: {str(e)}' error_msg = f'Error training model - error: {str(e)}'
trace = traceback.format_exc() trace = traceback.format_exc()
await self.send_notification_async( self.send_notification(
metadata=metadata or {}, metadata=metadata or {},
notification_id='TRAIN_MODEL_ERROR', notification_id='TRAIN_MODEL_ERROR',
message=error_msg, message=error_msg,
@@ -284,8 +266,32 @@ class Training(SientiaMonitoring):
raise e raise e
def _persist_training_artifacts(
self,
train_result: TrainModelResult,
train_params: TrainModelParams,
wrapper: Any,
metadata: dict[str, Any] | None,
) -> None:
train_result = self.data_manager_repository.generate_report(
train_result,
metadata=metadata,
)
if (
train_result.report_path is None
or train_result.train_data_path is None
or train_result.test_data_path is None
):
raise ValueError('Report path, train data path, or test data path is not set')
wrapper.store_model(name=train_params.model_name)
mlflow.log_artifact(train_result.report_path)
mlflow.log_artifact(train_result.train_data_path)
mlflow.log_artifact(train_result.test_data_path)
@activity.defn(name='cleanup_resources') @activity.defn(name='cleanup_resources')
async def cleanup_resources(self, input_data: dict[str, Any]) -> None: def cleanup_resources(self, input_data: dict[str, Any]) -> None:
""" """
Cleanup temporary resources created during training. Cleanup temporary resources created during training.
@@ -317,4 +323,3 @@ class Training(SientiaMonitoring):
) )
raise raise

View File

@@ -62,8 +62,8 @@ class TrainModelParams:
# New Parameters # New Parameters
val_file_name: str | None val_file_name: str | None
data_model_kwargs: dict | None # Removed params used in DataPreprocessor here data_model_kwargs: dict | None # Removed params used in DataPreprocessor here
model_kwargs: dict | None # Removed params used in Linear Regression Model here model_kwargs: dict | None # Removed params used in Linear Regression Model here
opt_params: dict | None opt_params: dict | None
model_type: str model_type: str
model_id: str | None model_id: str | None
@@ -102,12 +102,16 @@ class TrainModelParams:
model_name = cls._check_none(data.get('model_name'), str, 'model_name') model_name = cls._check_none(data.get('model_name'), str, 'model_name')
return cls( return cls(
variable_columns=cls._check_none(data.get('variable_columns'), list, 'variable_columns'), variable_columns=cls._check_none(
data.get('variable_columns'), list, 'variable_columns'
),
target_variable=cls._check_none(data.get('target_variable'), str, 'target_variable'), target_variable=cls._check_none(data.get('target_variable'), str, 'target_variable'),
bucket_name=cls._check_none(data.get('bucket_name'), str, 'bucket_name'), bucket_name=cls._check_none(data.get('bucket_name'), str, 'bucket_name'),
file_name=cls._check_none(data.get('file_name'), str, 'file_name'), file_name=cls._check_none(data.get('file_name'), str, 'file_name'),
line_separator=cls._check_none(data.get('line_separator'), str, 'line_separator'), line_separator=cls._check_none(data.get('line_separator'), str, 'line_separator'),
decimal_separator=cls._check_none(data.get('decimal_separator'), str, 'decimal_separator'), decimal_separator=cls._check_none(
data.get('decimal_separator'), str, 'decimal_separator'
),
date_column=data.get('date_column'), date_column=data.get('date_column'),
date_format=data.get('date_format'), date_format=data.get('date_format'),
train_size=cls._check_none(data.get('train_size'), int, 'train_size'), train_size=cls._check_none(data.get('train_size'), int, 'train_size'),
@@ -117,12 +121,13 @@ class TrainModelParams:
model_name=model_name, model_name=model_name,
experiment_name=model_name + '_experiment', experiment_name=model_name + '_experiment',
val_file_name=data.get('val_file_name'), val_file_name=data.get('val_file_name'),
data_model_kwargs=cls._check_none(data.get('data_model_kwargs'), dict, 'data_model_kwargs'), data_model_kwargs=cls._check_none(
data.get('data_model_kwargs'), dict, 'data_model_kwargs'
),
model_kwargs=cls._check_none(data.get('model_kwargs'), dict, 'model_kwargs'), model_kwargs=cls._check_none(data.get('model_kwargs'), dict, 'model_kwargs'),
opt_params=cls._check_none(data.get('opt_params'), dict, 'opt_params'), opt_params=cls._check_none(data.get('opt_params'), dict, 'opt_params'),
model_type=cls._check_none(data.get('model_type'), str, 'model_type'), model_type=cls._check_none(data.get('model_type'), str, 'model_type'),
model_id=data.get('model_id'), model_id=data.get('model_id'),
model_metadata=cls._parse_optional_model_metadata(data.get('model_metadata')), model_metadata=cls._parse_optional_model_metadata(data.get('model_metadata')),
) )
@@ -233,9 +238,7 @@ class TrainModelParams:
return None return None
if isinstance(value, dict): if isinstance(value, dict):
return value return value
raise TypeError( raise TypeError(f'model_metadata must be a dict or None, but got {type(value).__name__}.')
f'model_metadata must be a dict or None, but got {type(value).__name__}.'
)
def validate_business_rules(self) -> None: def validate_business_rules(self) -> None:
""" """
@@ -261,21 +264,20 @@ class TrainModelParams:
if not self.variable_columns: if not self.variable_columns:
raise ValueError('variable_columns cannot be empty') raise ValueError('variable_columns cannot be empty')
def _validate_model_params(self) -> None: def _validate_model_params(self) -> None:
"""Validate model-related parameters.""" """Validate model-related parameters."""
if not self.model_metadata: if not self.model_metadata:
raise ValueError('model_metadata is required') raise ValueError('model_metadata is required')
schemas = self.model_metadata.get('schemas', {}).get("components", {}).get("schemas") schemas = self.model_metadata.get('schemas', {}).get('components', {}).get('schemas')
if not schemas: if not schemas:
return return
data_model_schema = schemas.get("data_model") data_model_schema = schemas.get('data_model')
model_schema = schemas.get("model") model_schema = schemas.get('model')
opt_params_schema = schemas.get("opt_params") opt_params_schema = schemas.get('opt_params')
if data_model_schema: if data_model_schema:
self._validate_model_param(data_model_schema, self.data_model_kwargs) self._validate_model_param(data_model_schema, self.data_model_kwargs)
@@ -283,8 +285,6 @@ class TrainModelParams:
self._validate_model_param(model_schema, self.model_kwargs) self._validate_model_param(model_schema, self.model_kwargs)
if opt_params_schema: if opt_params_schema:
self._validate_model_param(opt_params_schema, self.opt_params) self._validate_model_param(opt_params_schema, self.opt_params)
def _validate_model_param(self, schema: dict[str, Any], value: Any) -> None: def _validate_model_param(self, schema: dict[str, Any], value: Any) -> None:
"""Validate model parameter against schema.""" """Validate model parameter against schema."""
@@ -292,9 +292,7 @@ class TrainModelParams:
validator = Draft202012Validator(schema) validator = Draft202012Validator(schema)
validator.validate(value) validator.validate(value)
except ValidationError as e: except ValidationError as e:
raise ValueError(f'Model parameters validation failed: {e.message}') raise ValueError(f'Model parameters validation failed: {e.message}') from e
except Exception as e:
raise ValueError(f'Unexpected error: {e}')
def _validate_required_strings(self) -> None: def _validate_required_strings(self) -> None:
"""Validate required string fields are not empty.""" """Validate required string fields are not empty."""
@@ -313,4 +311,4 @@ class TrainModelParams:
def _validate_date_format(self) -> None: def _validate_date_format(self) -> None:
"""Validate date_format is one of the allowed frontend formats when set.""" """Validate date_format is one of the allowed frontend formats when set."""
if self.date_format: if self.date_format:
validate_frontend_date_format(self.date_format) validate_frontend_date_format(self.date_format)

View File

@@ -1,5 +1,4 @@
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any
import pandas as pd import pandas as pd

View File

@@ -13,12 +13,12 @@ integration. Models are trained elsewhere (e.g., via SientiaModel wrappers),
and this repository focuses solely on preparing data structures for them. and this repository focuses solely on preparing data structures for them.
""" """
import json
from datetime import datetime from datetime import datetime
from io import BytesIO from io import BytesIO
import json
from os import makedirs, path from os import makedirs, path
from typing import Any
from shutil import rmtree from shutil import rmtree
from typing import Any
import numpy as np import numpy as np
import pandas as pd import pandas as pd
@@ -28,35 +28,42 @@ from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
from sientia_model.wrappers.sientia_model import SientiaModel from sientia_model.wrappers.sientia_model import SientiaModel
from model_manager.sientia.metrics import mae, mse, r2 from model_manager.sientia.metrics import mae, mse, r2
from model_manager.sientia.reports import Reports # type: ignore[import-untyped]
from model_manager.utils.models.train_model_params import TrainModelParams from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult from model_manager.utils.models.train_model_result import TrainModelResult
from model_manager.sientia.reports import Reports # type: ignore[import-untyped]
def train_test_split(data: pd.DataFrame | pd.Series, train_size: float, random_state: int | None = None, shuffle: bool = True) -> tuple[pd.DataFrame, pd.DataFrame]:
def train_test_split(
data: pd.DataFrame | pd.Series,
train_size: float,
random_state: int | None = None,
shuffle: bool = True,
) -> tuple[pd.DataFrame, pd.DataFrame]:
# 1. Definir a semente (seed) para reprodutibilidade # 1. Definir a semente (seed) para reprodutibilidade
if random_state is not None: if random_state is not None:
np.random.seed(random_state) np.random.seed(random_state)
# 2. Gerar índices e embaralhar se necessário # 2. Gerar índices e embaralhar se necessário
indices = np.arange(len(data)) indices = np.arange(len(data))
if shuffle: if shuffle:
np.random.shuffle(indices) np.random.shuffle(indices)
# 3. Calcular o ponto de corte (split point) # 3. Calcular o ponto de corte (split point)
# Cálculo: N_treino = tamanho_total * proporcao_treino # Cálculo: N_treino = tamanho_total * proporcao_treino
n_train = int(len(data) * train_size) n_train = int(len(data) * train_size)
# 4. Dividir os índices # 4. Dividir os índices
train_indices = indices[:n_train] train_indices = indices[:n_train]
test_indices = indices[n_train:] test_indices = indices[n_train:]
# 5. Retornar os dados fatiados (funciona para DataFrame ou Series) # 5. Retornar os dados fatiados (funciona para DataFrame ou Series)
if isinstance(data, (pd.DataFrame, pd.Series)): if isinstance(data, (pd.DataFrame, pd.Series)):
return data.iloc[train_indices], data.iloc[test_indices] return data.iloc[train_indices], data.iloc[test_indices]
return data[train_indices], data[test_indices] return data[train_indices], data[test_indices]
def _ensure_date_column_parsed(data: pd.DataFrame, params: TrainModelParams) -> pd.DataFrame: def _ensure_date_column_parsed(data: pd.DataFrame, params: TrainModelParams) -> pd.DataFrame:
""" """
If date_column is set, parse the column as timezone-aware If date_column is set, parse the column as timezone-aware
@@ -178,7 +185,6 @@ class DataManagerRepository(SientiaMonitoring):
if len(val_df) <= 0: if len(val_df) <= 0:
raise ValueError('Validation data view is empty after transformation') raise ValueError('Validation data view is empty after transformation')
val_data = pd.DataFrame(val_df[params.variable_columns + [params.target_variable]]) val_data = pd.DataFrame(val_df[params.variable_columns + [params.target_variable]])
else: else:
# Fallback path: derive validation via train/test split from a single dataset. # Fallback path: derive validation via train/test split from a single dataset.
@@ -194,11 +200,7 @@ class DataManagerRepository(SientiaMonitoring):
metadata, metadata,
) )
return TrainModelResult( return TrainModelResult(params=params, train_data=train_data, val_data=val_data)
params=params,
train_data=train_data,
val_data=val_data
)
def _as_series(self, pred: pd.DataFrame | pd.Series) -> pd.Series: def _as_series(self, pred: pd.DataFrame | pd.Series) -> pd.Series:
if isinstance(pred, pd.Series): if isinstance(pred, pd.Series):
@@ -207,10 +209,8 @@ class DataManagerRepository(SientiaMonitoring):
if pred.shape[1] == 1: if pred.shape[1] == 1:
return pred.iloc[:, 0] return pred.iloc[:, 0]
raise ValueError('y_pred/y_train_pred must be a Series or single-column DataFrame') raise ValueError('y_pred/y_train_pred must be a Series or single-column DataFrame')
def _extract_model_equation( def _extract_model_equation(self, regr: Any, params: TrainModelParams) -> dict:
self, regr: Any, params: TrainModelParams
) -> dict:
""" """
Extract the linear regression equation coefficients and create equation metadata. Extract the linear regression equation coefficients and create equation metadata.
@@ -237,7 +237,7 @@ class DataManagerRepository(SientiaMonitoring):
model_kwargs = params.model_kwargs or {} model_kwargs = params.model_kwargs or {}
degree = model_kwargs.get('degree', 1) degree = model_kwargs.get('degree', 1)
poly_feature_names = model_kwargs.get('poly_feature_names', None) poly_feature_names = model_kwargs.get('poly_feature_names', None)
if degree > 1 and poly_feature_names: if degree > 1 and poly_feature_names:
feature_names = poly_feature_names feature_names = poly_feature_names
else: else:
@@ -271,7 +271,6 @@ class DataManagerRepository(SientiaMonitoring):
'original_features': feature_names, 'original_features': feature_names,
} }
def compute_regression_metrics( def compute_regression_metrics(
self, self,
tmr: TrainModelResult, tmr: TrainModelResult,
@@ -301,7 +300,7 @@ class DataManagerRepository(SientiaMonitoring):
# True values are expected to come from val_data. # True values are expected to come from val_data.
y_true_val = tmr.val_data[target] y_true_val = tmr.val_data[target]
y_pred_val = self._as_series(tmr.y_pred).sort_index() y_pred_val = self._as_series(tmr.y_pred).sort_index()
y_true_val = y_true_val.sort_index() y_true_val = y_true_val.sort_index()
@@ -407,7 +406,9 @@ class DataManagerRepository(SientiaMonitoring):
reports_dir = path.join(model_manager_dir, 'reports') reports_dir = path.join(model_manager_dir, 'reports')
return reports_dir return reports_dir
def _create_run_directory(self, base_path: str, run_name: str, metadata: dict[str, Any] | None = None) -> str: def _create_run_directory(
self, base_path: str, run_name: str, metadata: dict[str, Any] | None = None
) -> str:
""" """
Creates a directory inside the 'reports' folder with the run name and a timestamp. Creates a directory inside the 'reports' folder with the run name and a timestamp.
@@ -471,7 +472,6 @@ class DataManagerRepository(SientiaMonitoring):
reference_data_float = reference_data.astype(np.float64) reference_data_float = reference_data.astype(np.float64)
current_data_float = current_data.astype(np.float64) current_data_float = current_data.astype(np.float64)
# Initialize report generator # Initialize report generator
data.run_dir = self._create_run_directory(self._get_reports_directory(), data.run_name) data.run_dir = self._create_run_directory(self._get_reports_directory(), data.run_name)
report = Reports( report = Reports(

View File

@@ -0,0 +1,106 @@
import os
import re
from collections.abc import Sequence
from concurrent.futures import ThreadPoolExecutor
from typing import Any
from sientia_do.observability.logger import Logger
from temporalio.client import Client
from temporalio.worker import PollerBehaviorAutoscaling, Worker
# Worker configuration parameters with default values.
parameters = [
('MAX_CONCURRENT_WORKFLOW_TASKS', '200'),
('MAX_CONCURRENT_ACTIVITIES', '200'),
('MAX_CONCURRENT_LOCAL_ACTIVITIES', '200'),
('MAX_CACHED_WORKFLOWS', '200'),
('WORKFLOW_POLLER_BEHAVIOUR_MINIMUM', '10'),
('WORKFLOW_POLLER_BEHAVIOUR_INITIAL', '100'),
('WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM', '200'),
('ACTIVITY_POLLER_BEHAVIOUR_MINIMUM', '10'),
('ACTIVITY_POLLER_BEHAVIOUR_INITIAL', '100'),
('ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM', '200'),
('ACTIVITY_EXECUTOR_MAX_WORKERS', '32'),
]
def camel_to_snake(text: str) -> str:
"""
Convert a CamelCase or camelCase string into snake_case.
Args:
- text: str, original string in CamelCase or camelCase format
Return:
str: converted string in snake_case format
"""
text = re.sub('(.)([A-Z][a-z]+)', r'\1_\2', text)
text = re.sub('([a-z0-9])([A-Z])', r'\1_\2', text)
return text.lower()
def prepare_worker(
main_workflow: type,
other_workflows: Sequence[type],
activities: Sequence[Any],
temporal_client: Client,
logger: Logger,
runtime: str | None = None,
) -> Worker:
"""
Build and configure a Temporal worker for the given workflow and activities.
Args:
- main_workflow: type, main workflow class used as worker entry point
- other_workflows: Sequence[type], additional workflows in the same worker
- activities: Sequence[Any], activity callables registered in this worker
- temporal_client: Client, Temporal client used by the worker
- logger: Logger, logger instance used during worker preparation
- runtime: str | None, runtime suffix appended to queue name when present
Return:
Worker: fully configured Temporal worker instance ready to run
"""
main_workflow_name = main_workflow.__name__.upper()
queue_name = (
f'{camel_to_snake(main_workflow.__name__)}-{runtime}-queue'
if runtime
else f'{camel_to_snake(main_workflow.__name__)}-queue'
)
local_workflow_parameters: dict[str, int] = {}
for parameter_name, default_value in parameters:
local_workflow_parameters[parameter_name] = int(
os.getenv(f'{main_workflow_name}_{parameter_name}', default_value)
)
logger.info(f'Preparing worker for {main_workflow_name} with queue {queue_name}')
activity_executor = ThreadPoolExecutor(
max_workers=local_workflow_parameters['ACTIVITY_EXECUTOR_MAX_WORKERS'],
thread_name_prefix=f'{queue_name}-activity',
)
return Worker(
temporal_client,
task_queue=queue_name,
workflows=[main_workflow, *other_workflows],
activities=[*activities],
activity_executor=activity_executor,
max_concurrent_workflow_tasks=local_workflow_parameters['MAX_CONCURRENT_WORKFLOW_TASKS'],
max_concurrent_activities=local_workflow_parameters['MAX_CONCURRENT_ACTIVITIES'],
max_concurrent_local_activities=local_workflow_parameters[
'MAX_CONCURRENT_LOCAL_ACTIVITIES'
],
max_cached_workflows=local_workflow_parameters['MAX_CACHED_WORKFLOWS'],
workflow_task_poller_behavior=PollerBehaviorAutoscaling(
minimum=local_workflow_parameters['WORKFLOW_POLLER_BEHAVIOUR_MINIMUM'],
initial=local_workflow_parameters['WORKFLOW_POLLER_BEHAVIOUR_INITIAL'],
maximum=local_workflow_parameters['WORKFLOW_POLLER_BEHAVIOUR_MAXIMUM'],
),
activity_task_poller_behavior=PollerBehaviorAutoscaling(
minimum=local_workflow_parameters['ACTIVITY_POLLER_BEHAVIOUR_MINIMUM'],
initial=local_workflow_parameters['ACTIVITY_POLLER_BEHAVIOUR_INITIAL'],
maximum=local_workflow_parameters['ACTIVITY_POLLER_BEHAVIOUR_MAXIMUM'],
),
)

View File

@@ -30,7 +30,6 @@ Environment Variables:
from temporalio import client, workflow from temporalio import client, workflow
from temporalio.runtime import PrometheusConfig, Runtime, TelemetryConfig from temporalio.runtime import PrometheusConfig, Runtime, TelemetryConfig
from temporalio.worker import PollerBehaviorAutoscaling, Worker
with workflow.unsafe.imports_passed_through(): with workflow.unsafe.imports_passed_through():
import asyncio import asyncio
@@ -40,9 +39,8 @@ with workflow.unsafe.imports_passed_through():
from prometheus_client import start_http_server from prometheus_client import start_http_server
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.observability.logger import Logger as SientiaLogger from sientia_do.observability.logger import Logger as SientiaLogger
from sientia_model.model_repository.plugin_store import PluginStore
from sientia_do.observability.metrics_controller import MetricsController from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.temporal.worker.prepare_worker import prepare_worker from sientia_model.model_repository.plugin_store import PluginStore
from model_manager import metrics from model_manager import metrics
from model_manager.activities.activities import Activities from model_manager.activities.activities import Activities
@@ -55,6 +53,7 @@ with workflow.unsafe.imports_passed_through():
build_postgres_config, build_postgres_config,
) )
from model_manager.utils.logger_helper import get_logger from model_manager.utils.logger_helper import get_logger
from model_manager.worker.prepare_worker import prepare_worker
from model_manager.workflows.cleanup_files import CleanupFiles from model_manager.workflows.cleanup_files import CleanupFiles
from model_manager.workflows.train_model import TrainModel from model_manager.workflows.train_model import TrainModel
@@ -96,7 +95,6 @@ async def main():
'pod_id': POD_ID, 'pod_id': POD_ID,
'runtime': RUNTIME, 'runtime': RUNTIME,
} }
start_prometheus_server(logger, metadata) start_prometheus_server(logger, metadata)
mongo_config = build_mongodb_config() mongo_config = build_mongodb_config()
@@ -110,11 +108,9 @@ async def main():
logger.custom_info(f'MongoDB client initialized at {mongo_config["uri"]}', metadata) logger.custom_info(f'MongoDB client initialized at {mongo_config["uri"]}', metadata)
logger.custom_info(f'Initializing metrics controller', metadata) logger.custom_info('Initializing metrics controller', metadata)
metrics_controller = MetricsController( metrics_controller = MetricsController(logger=logger)
logger=logger
)
logger.custom_info(f'Installing runtime {RUNTIME}', metadata) logger.custom_info(f'Installing runtime {RUNTIME}', metadata)
@@ -185,6 +181,7 @@ async def main():
], ],
temporal_client=temporal_client, temporal_client=temporal_client,
logger=logger, logger=logger,
runtime=RUNTIME,
), ),
prepare_worker( prepare_worker(
main_workflow=CleanupFiles, main_workflow=CleanupFiles,
@@ -194,12 +191,11 @@ async def main():
], ],
temporal_client=temporal_client, temporal_client=temporal_client,
logger=logger, logger=logger,
runtime=RUNTIME,
), ),
] ]
handlers = [ handlers = [w.run() for w in workers]
w.run() for w in workers
]
logger.custom_info('Model manager workers initialized', metadata) logger.custom_info('Model manager workers initialized', metadata)

View File

@@ -21,7 +21,6 @@ with workflow.unsafe.imports_passed_through():
from model_manager.activities.activities import Activities from model_manager.activities.activities import Activities
from model_manager.activities.experiment_tracking import UpdateType from model_manager.activities.experiment_tracking import UpdateType
from model_manager.utils.models.experiment_status import ExperimentStatus from model_manager.utils.models.experiment_status import ExperimentStatus
from model_manager.utils.models.train_model_params import TrainModelParams
# Activity timeouts (seconds). Tune per environment (large uploads, long training). # Activity timeouts (seconds). Tune per environment (large uploads, long training).
# Training uses no_retry_policy: extend TIMEOUT_TRAIN_MODEL instead of adding retries # Training uses no_retry_policy: extend TIMEOUT_TRAIN_MODEL instead of adding retries
@@ -54,7 +53,6 @@ with workflow.unsafe.imports_passed_through():
) )
@workflow.defn(name='train_model') @workflow.defn(name='train_model')
class TrainModel: class TrainModel:
""" """
@@ -133,7 +131,7 @@ class TrainModel:
) )
else: else:
pass pass
except Exception: except Exception: # noqa: BLE001
# If cleanup fails after training failed, there is nothing extra to log (DB not committed). # If cleanup fails after training failed, there is nothing extra to log (DB not committed).
if training_succeeded: # pragma: no branch if training_succeeded: # pragma: no branch
workflow.logger.warning( workflow.logger.warning(
@@ -179,7 +177,7 @@ class TrainModel:
input_data: dict[str, Any], input_data: dict[str, Any],
experiment_run_id: int, experiment_run_id: int,
metadata: dict[str, Any], metadata: dict[str, Any],
) -> TrainModelParams: ) -> dict[str, Any]:
""" """
Validate and convert training parameters from dict to TrainModelParams. Validate and convert training parameters from dict to TrainModelParams.
@@ -193,7 +191,7 @@ class TrainModel:
metadata: Workflow execution metadata metadata: Workflow execution metadata
Returns: Returns:
TrainModelParams: Validated training parameters object dict[str, Any]: Validated training parameters
Raises: Raises:
Exception: If validation fails (after updating DB status) Exception: If validation fails (after updating DB status)
@@ -220,7 +218,7 @@ class TrainModel:
) )
await self._update_experiment_run( await self._update_experiment_run(
metadata=metadata, metadata=metadata,
experiment_run_id=experiment_run_id, experiment_run_id=experiment_run_id,
update_type=UpdateType.STATUS, update_type=UpdateType.STATUS,
status=ExperimentStatus.ORCHESTRATOR_WAITING_PROC, status=ExperimentStatus.ORCHESTRATOR_WAITING_PROC,
@@ -236,7 +234,7 @@ class TrainModel:
status=ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR, status=ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR,
error_message=self._extract_error_message(e), error_message=self._extract_error_message(e),
) )
except Exception as secondary: except Exception as secondary: # noqa: BLE001
workflow.logger.warning( workflow.logger.warning(
'Failed to persist ORCHESTRATOR_VALIDATION_ERROR to experiment_run: %s', 'Failed to persist ORCHESTRATOR_VALIDATION_ERROR to experiment_run: %s',
secondary, secondary,
@@ -245,7 +243,7 @@ class TrainModel:
async def _train_model( async def _train_model(
self, self,
train_params: TrainModelParams, train_params: dict[str, Any],
experiment_run_id: int, experiment_run_id: int,
metadata: dict[str, Any], metadata: dict[str, Any],
) -> dict[str, Any]: ) -> dict[str, Any]:
@@ -288,7 +286,6 @@ class TrainModel:
return train_result return train_result
except Exception as e: except Exception as e:
try: try:
await self._update_experiment_run( await self._update_experiment_run(
metadata=metadata, metadata=metadata,
@@ -297,7 +294,7 @@ class TrainModel:
status=ExperimentStatus.TRAINING_ERROR, status=ExperimentStatus.TRAINING_ERROR,
error_message=self._extract_error_message(e), error_message=self._extract_error_message(e),
) )
except Exception as secondary: except Exception as secondary: # noqa: BLE001
workflow.logger.warning( workflow.logger.warning(
'Failed to persist TRAINING_ERROR status to experiment_run: %s', 'Failed to persist TRAINING_ERROR status to experiment_run: %s',
secondary, secondary,

View File

@@ -56,6 +56,7 @@ ignore = [
"S101", # assert allowed in tests "S101", # assert allowed in tests
"S105", # hardcoded passwords ok in tests "S105", # hardcoded passwords ok in tests
"S106", # hardcoded passwords ok in tests "S106", # hardcoded passwords ok in tests
"S108", # temp paths are expected in tests
] ]
[tool.ruff.lint.mccabe] [tool.ruff.lint.mccabe]
@@ -121,6 +122,13 @@ ignore_missing_imports = true
module = "yaml" module = "yaml"
ignore_missing_imports = true ignore_missing_imports = true
[[tool.mypy.overrides]]
module = [
"model_manager.utils.repository.model_repository",
"model_manager.utils.repository.data_manager_repository",
]
ignore_errors = true
[tool.pytest.ini_options] [tool.pytest.ini_options]
testpaths = ["tests"] testpaths = ["tests"]
python_files = ["test_*.py"] python_files = ["test_*.py"]

View File

@@ -4,7 +4,7 @@ sqlalchemy==2.0.44
boto3==1.40.55 boto3==1.40.55
botocore==1.40.55 botocore==1.40.55
/home/grezewave/Documents/projects/sientia/sientia-model-library /home/grezewave/Documents/projects/sientia/sientia-model-library
git+https://github.com/Aignosi/sientia-dataops-library.git@v1.10.1 /home/grezewave/Documents/projects/sientia/sientia-dataops-library
prometheus-client==0.23.1 prometheus-client==0.23.1
beautifulsoup4==4.12.3 beautifulsoup4==4.12.3
evidently evidently

View File

@@ -44,7 +44,10 @@ def _minio(endpoint_url: str):
) )
def test_activities_strips_minio_endpoint_scheme(endpoint, expected_endpoint): def test_activities_strips_minio_endpoint_scheme(endpoint, expected_endpoint):
with ( with (
patch('model_manager.activities.activities.ExperimentTracking.__init__', Mock(return_value=None)), patch(
'model_manager.activities.activities.ExperimentTracking.__init__',
Mock(return_value=None),
),
patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)), patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)),
patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)), patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)),
patch('model_manager.activities.activities.SientiaMLflowRepository') as m_mlflow, patch('model_manager.activities.activities.SientiaMLflowRepository') as m_mlflow,
@@ -66,7 +69,10 @@ def test_activities_strips_minio_endpoint_scheme(endpoint, expected_endpoint):
def test_activities_shutdown_calls_parents(): def test_activities_shutdown_calls_parents():
with ( with (
patch('model_manager.activities.activities.ExperimentTracking.__init__', Mock(return_value=None)), patch(
'model_manager.activities.activities.ExperimentTracking.__init__',
Mock(return_value=None),
),
patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)), patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)),
patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)), patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)),
patch('model_manager.activities.activities.SientiaMLflowRepository'), patch('model_manager.activities.activities.SientiaMLflowRepository'),
@@ -91,7 +97,10 @@ def test_activities_shutdown_calls_parents():
def test_activities_del_with_engine_runs_without_error(): def test_activities_del_with_engine_runs_without_error():
with ( with (
patch('model_manager.activities.activities.ExperimentTracking.__init__', Mock(return_value=None)), patch(
'model_manager.activities.activities.ExperimentTracking.__init__',
Mock(return_value=None),
),
patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)), patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)),
patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)), patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)),
patch('model_manager.activities.activities.SientiaMLflowRepository'), patch('model_manager.activities.activities.SientiaMLflowRepository'),
@@ -112,7 +121,10 @@ def test_activities_del_with_engine_runs_without_error():
def test_activities_del_without_engine_runs_without_error(): def test_activities_del_without_engine_runs_without_error():
with ( with (
patch('model_manager.activities.activities.ExperimentTracking.__init__', Mock(return_value=None)), patch(
'model_manager.activities.activities.ExperimentTracking.__init__',
Mock(return_value=None),
),
patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)), patch('model_manager.activities.activities.Training.__init__', Mock(return_value=None)),
patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)), patch('model_manager.activities.activities.Cleanup.__init__', Mock(return_value=None)),
patch('model_manager.activities.activities.SientiaMLflowRepository'), patch('model_manager.activities.activities.SientiaMLflowRepository'),

View File

@@ -1,6 +1,5 @@
"""Unit tests for the Cleanup activity, ensuring 100% code coverage.""" """Unit tests for the Cleanup activity, ensuring 100% code coverage."""
import asyncio
import os import os
import shutil import shutil
import tempfile import tempfile
@@ -126,12 +125,10 @@ def test_cleanup_temp_directories_nonexistent_path(
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
cleanup._emit_metrics = AsyncMock() cleanup._emit_metrics = MagicMock()
cleanup.warning = MagicMock() cleanup.warning = MagicMock()
asyncio.run( cleanup.cleanup_temp_directories({'temp_path': '/nonexistent/path', 'metadata': {}})
cleanup.cleanup_temp_directories({'temp_path': '/nonexistent/path', 'metadata': {}})
)
cleanup.warning.assert_called_once() cleanup.warning.assert_called_once()
cleanup._emit_metrics.assert_called_once() cleanup._emit_metrics.assert_called_once()
@@ -155,7 +152,7 @@ def test_cleanup_temp_directories_success_with_deletions(
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
cleanup._emit_metrics = AsyncMock() cleanup._emit_metrics = MagicMock()
old_time = (datetime.now() - timedelta(hours=48)).strftime('%Y%m%d_%H%M%S_000000') old_time = (datetime.now() - timedelta(hours=48)).strftime('%Y%m%d_%H%M%S_000000')
old_dir = os.path.join(temp_dir, f'old_dir_{old_time}') old_dir = os.path.join(temp_dir, f'old_dir_{old_time}')
@@ -165,7 +162,7 @@ def test_cleanup_temp_directories_success_with_deletions(
recent_dir = os.path.join(temp_dir, f'recent_dir_{recent_time}') recent_dir = os.path.join(temp_dir, f'recent_dir_{recent_time}')
os.makedirs(recent_dir) os.makedirs(recent_dir)
asyncio.run(cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})) cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
assert not os.path.exists(old_dir) assert not os.path.exists(old_dir)
assert os.path.exists(recent_dir) assert os.path.exists(recent_dir)
@@ -190,13 +187,13 @@ def test_cleanup_temp_directories_dry_run(
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
cleanup._emit_metrics = AsyncMock() cleanup._emit_metrics = MagicMock()
old_time = (datetime.now() - timedelta(hours=48)).strftime('%Y%m%d_%H%M%S_000000') old_time = (datetime.now() - timedelta(hours=48)).strftime('%Y%m%d_%H%M%S_000000')
old_dir = os.path.join(temp_dir, f'old_dir_{old_time}') old_dir = os.path.join(temp_dir, f'old_dir_{old_time}')
os.makedirs(old_dir) os.makedirs(old_dir)
asyncio.run(cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})) cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
assert os.path.exists(old_dir) assert os.path.exists(old_dir)
cleanup._emit_metrics.assert_called_once() cleanup._emit_metrics.assert_called_once()
@@ -220,7 +217,7 @@ def test_cleanup_temp_directories_delete_error(
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
cleanup._emit_metrics = AsyncMock() cleanup._emit_metrics = MagicMock()
cleanup.error = MagicMock() cleanup.error = MagicMock()
old_time = (datetime.now() - timedelta(hours=48)).strftime('%Y%m%d_%H%M%S_000000') old_time = (datetime.now() - timedelta(hours=48)).strftime('%Y%m%d_%H%M%S_000000')
@@ -228,7 +225,7 @@ def test_cleanup_temp_directories_delete_error(
os.makedirs(old_dir) os.makedirs(old_dir)
with patch('shutil.rmtree', side_effect=OSError('Permission Denied')): with patch('shutil.rmtree', side_effect=OSError('Permission Denied')):
asyncio.run(cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})) cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
cleanup.error.assert_called_once() cleanup.error.assert_called_once()
cleanup._emit_metrics.assert_called_once() cleanup._emit_metrics.assert_called_once()
@@ -250,18 +247,16 @@ def test_emit_metrics(
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
cleanup.emit_metric = AsyncMock() cleanup.emit_metric_sync = MagicMock()
asyncio.run( cleanup._emit_metrics(
cleanup._emit_metrics( metadata={'pod_id': 'p1', 'workflow_name': 'wf1'},
metadata={'pod_id': 'p1', 'workflow_name': 'wf1'}, metrics_status='success',
metrics_status='success', activity_name='test_activity',
activity_name='test_activity', emit_workflow_metric=True,
emit_workflow_metric=True,
)
) )
assert cleanup.emit_metric.call_count == 2 assert cleanup.emit_metric_sync.call_count == 2
def test_cleanup_temp_directories_with_files_and_unmatched_dirs( def test_cleanup_temp_directories_with_files_and_unmatched_dirs(
@@ -281,7 +276,7 @@ def test_cleanup_temp_directories_with_files_and_unmatched_dirs(
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
cleanup._emit_metrics = AsyncMock() cleanup._emit_metrics = MagicMock()
cleanup.debug = MagicMock() cleanup.debug = MagicMock()
# Create a file and a directory with a non-matching name # Create a file and a directory with a non-matching name
@@ -289,7 +284,7 @@ def test_cleanup_temp_directories_with_files_and_unmatched_dirs(
f.write('hello') f.write('hello')
os.makedirs(os.path.join(temp_dir, 'a_directory_with_no_timestamp')) os.makedirs(os.path.join(temp_dir, 'a_directory_with_no_timestamp'))
asyncio.run(cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})) cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
# Ensure the debug message for skipping was called for the unmatched directory # Ensure the debug message for skipping was called for the unmatched directory
cleanup.debug.assert_called_with( cleanup.debug.assert_called_with(
@@ -315,14 +310,14 @@ def test_cleanup_temp_directories_invalid_timestamp_format(
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
cleanup._emit_metrics = AsyncMock() cleanup._emit_metrics = MagicMock()
cleanup.error = MagicMock() cleanup.error = MagicMock()
# Create a directory with a malformed timestamp that matches the regex but fails parsing # Create a directory with a malformed timestamp that matches the regex but fails parsing
malformed_dir_name = 'dir_20239999_999999_999999' malformed_dir_name = 'dir_20239999_999999_999999'
os.makedirs(os.path.join(temp_dir, malformed_dir_name)) os.makedirs(os.path.join(temp_dir, malformed_dir_name))
asyncio.run(cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})) cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
cleanup.error.assert_called_once() cleanup.error.assert_called_once()
cleanup._emit_metrics.assert_called_once() cleanup._emit_metrics.assert_called_once()
@@ -345,12 +340,12 @@ def test_cleanup_temp_directories_generic_exception(
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
cleanup._emit_metrics = AsyncMock() cleanup._emit_metrics = MagicMock()
cleanup.send_notification = MagicMock() cleanup.send_notification = MagicMock()
with patch('os.listdir', side_effect=Exception('Unexpected OS Error')): with patch('os.listdir', side_effect=Exception('Unexpected OS Error')):
with pytest.raises(Exception, match='Unexpected OS Error'): with pytest.raises(Exception, match='Unexpected OS Error'):
asyncio.run(cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})) cleanup.cleanup_temp_directories({'temp_path': temp_dir, 'metadata': {}})
cleanup.send_notification.assert_called_once() cleanup.send_notification.assert_called_once()
cleanup._emit_metrics.assert_called_once() cleanup._emit_metrics.assert_called_once()
@@ -369,15 +364,13 @@ def test_emit_metrics_activity_only(
notification_handler=mock_notification_handler, notification_handler=mock_notification_handler,
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
cleanup.emit_metric = AsyncMock() cleanup.emit_metric_sync = MagicMock()
asyncio.run( cleanup._emit_metrics(
cleanup._emit_metrics( metadata={'pod_id': 'p1', 'workflow_name': 'wf1'},
metadata={'pod_id': 'p1', 'workflow_name': 'wf1'}, metrics_status='success',
metrics_status='success', activity_name='test_activity',
activity_name='test_activity', emit_workflow_metric=False,
emit_workflow_metric=False,
)
) )
cleanup.emit_metric.assert_called_once() cleanup.emit_metric_sync.assert_called_once()

View File

@@ -1,6 +1,5 @@
"""Unit tests for ExperimentTracking class with 100% coverage.""" """Unit tests for ExperimentTracking class with 100% coverage."""
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
@@ -185,7 +184,7 @@ def test_execute_update_success(
mock_engine.begin.return_value.__enter__.return_value = mock_connection mock_engine.begin.return_value.__enter__.return_value = mock_connection
et.engine = mock_engine et.engine = mock_engine
result = asyncio.run(et._execute_update('UPDATE test SET x = :x', {'x': 1})) result = et._execute_update('UPDATE test SET x = :x', {'x': 1})
assert result == {'rowcount': 1} assert result == {'rowcount': 1}
mock_connection.execute.assert_called_once() mock_connection.execute.assert_called_once()
@@ -212,7 +211,7 @@ def test_update_experiment_run_status_success(
mock_execute = MagicMock() mock_execute = MagicMock()
async def mock_execute_update(*args, **kwargs): def mock_execute_update(*args, **kwargs):
mock_execute(*args, **kwargs) mock_execute(*args, **kwargs)
return {'rowcount': 1} return {'rowcount': 1}
@@ -226,7 +225,7 @@ def test_update_experiment_run_status_success(
'status': 'running', 'status': 'running',
} }
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
mock_execute.assert_called_once() mock_execute.assert_called_once()
call_args = mock_execute.call_args call_args = mock_execute.call_args
@@ -255,7 +254,7 @@ def test_update_experiment_run_status_missing_status(
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
et.send_notification_async = AsyncMock() et.send_notification = MagicMock()
input_data = { input_data = {
'metadata': {'workflow_id': 'test-123'}, 'metadata': {'workflow_id': 'test-123'},
@@ -264,9 +263,9 @@ def test_update_experiment_run_status_missing_status(
} }
with pytest.raises(RuntimeError): with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
et.send_notification_async.assert_awaited_once() et.send_notification.assert_called_once()
def test_update_experiment_run_status_with_error_success( def test_update_experiment_run_status_with_error_success(
@@ -290,7 +289,7 @@ def test_update_experiment_run_status_with_error_success(
mock_execute = MagicMock() mock_execute = MagicMock()
async def mock_execute_update(*args, **kwargs): def mock_execute_update(*args, **kwargs):
mock_execute(*args, **kwargs) mock_execute(*args, **kwargs)
return {'rowcount': 1} return {'rowcount': 1}
@@ -305,7 +304,7 @@ def test_update_experiment_run_status_with_error_success(
'error_message': 'Test error', 'error_message': 'Test error',
} }
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
mock_execute.assert_called_once() mock_execute.assert_called_once()
call_args = mock_execute.call_args call_args = mock_execute.call_args
@@ -336,7 +335,7 @@ def test_update_experiment_run_status_with_error_truncate_message(
mock_execute = MagicMock() mock_execute = MagicMock()
async def mock_execute_update(*args, **kwargs): def mock_execute_update(*args, **kwargs):
mock_execute(*args, **kwargs) mock_execute(*args, **kwargs)
return {'rowcount': 1} return {'rowcount': 1}
@@ -352,7 +351,7 @@ def test_update_experiment_run_status_with_error_truncate_message(
'error_message': long_error, 'error_message': long_error,
} }
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
call_args = mock_execute.call_args call_args = mock_execute.call_args
assert len(call_args[0][1]['error_message']) == 1024 assert len(call_args[0][1]['error_message']) == 1024
@@ -377,7 +376,7 @@ def test_update_experiment_run_status_with_error_missing_error_message(
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
et.send_notification_async = AsyncMock() et.send_notification = MagicMock()
input_data = { input_data = {
'metadata': {'workflow_id': 'test-123'}, 'metadata': {'workflow_id': 'test-123'},
@@ -387,9 +386,9 @@ def test_update_experiment_run_status_with_error_missing_error_message(
} }
with pytest.raises(RuntimeError): with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
et.send_notification_async.assert_awaited_once() et.send_notification.assert_called_once()
def test_update_experiment_run_model_saved_success( def test_update_experiment_run_model_saved_success(
@@ -413,7 +412,7 @@ def test_update_experiment_run_model_saved_success(
mock_execute = MagicMock() mock_execute = MagicMock()
async def mock_execute_update(*args, **kwargs): def mock_execute_update(*args, **kwargs):
mock_execute(*args, **kwargs) mock_execute(*args, **kwargs)
return {'rowcount': 1} return {'rowcount': 1}
@@ -428,7 +427,7 @@ def test_update_experiment_run_model_saved_success(
'run_name': 'run_001', 'run_name': 'run_001',
} }
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
mock_execute.assert_called_once() mock_execute.assert_called_once()
call_args = mock_execute.call_args call_args = mock_execute.call_args
@@ -457,7 +456,7 @@ def test_update_experiment_run_model_saved_missing_run_name(
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
et.send_notification_async = AsyncMock() et.send_notification = MagicMock()
input_data = { input_data = {
'metadata': {'workflow_id': 'test-123'}, 'metadata': {'workflow_id': 'test-123'},
@@ -467,9 +466,9 @@ def test_update_experiment_run_model_saved_missing_run_name(
} }
with pytest.raises(RuntimeError): with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
et.send_notification_async.assert_awaited_once() et.send_notification.assert_called_once()
def test_update_experiment_run_invalid_update_type( def test_update_experiment_run_invalid_update_type(
@@ -491,7 +490,7 @@ def test_update_experiment_run_invalid_update_type(
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
et.send_notification_async = AsyncMock() et.send_notification = MagicMock()
input_data = { input_data = {
'metadata': {'workflow_id': 'test-123'}, 'metadata': {'workflow_id': 'test-123'},
@@ -500,9 +499,9 @@ def test_update_experiment_run_invalid_update_type(
} }
with pytest.raises(RuntimeError): with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
et.send_notification_async.assert_awaited_once() et.send_notification.assert_called_once()
def test_update_experiment_run_no_rows_updated( def test_update_experiment_run_no_rows_updated(
@@ -524,11 +523,11 @@ def test_update_experiment_run_no_rows_updated(
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
async def mock_execute_update(*args, **kwargs): def mock_execute_update(*args, **kwargs):
return {'rowcount': 0} return {'rowcount': 0}
et._execute_update = mock_execute_update et._execute_update = mock_execute_update
et.send_notification_async = AsyncMock() et.send_notification = MagicMock()
input_data = { input_data = {
'metadata': {'workflow_id': 'test-123'}, 'metadata': {'workflow_id': 'test-123'},
@@ -538,9 +537,9 @@ def test_update_experiment_run_no_rows_updated(
} }
with pytest.raises(RuntimeError): with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
et.send_notification_async.assert_awaited_once() et.send_notification.assert_called_once()
def test_update_experiment_run_status_with_error_missing_status( def test_update_experiment_run_status_with_error_missing_status(
@@ -562,7 +561,7 @@ def test_update_experiment_run_status_with_error_missing_status(
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
et.send_notification_async = AsyncMock() et.send_notification = MagicMock()
input_data = { input_data = {
'metadata': {'workflow_id': 'test-123'}, 'metadata': {'workflow_id': 'test-123'},
@@ -572,9 +571,9 @@ def test_update_experiment_run_status_with_error_missing_status(
} }
with pytest.raises(RuntimeError): with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
et.send_notification_async.assert_awaited_once() et.send_notification.assert_called_once()
def test_update_experiment_run_model_saved_missing_status( def test_update_experiment_run_model_saved_missing_status(
@@ -596,7 +595,7 @@ def test_update_experiment_run_model_saved_missing_status(
metrics_controller=mock_metrics_controller, metrics_controller=mock_metrics_controller,
) )
et.send_notification_async = AsyncMock() et.send_notification = MagicMock()
input_data = { input_data = {
'metadata': {'workflow_id': 'test-123'}, 'metadata': {'workflow_id': 'test-123'},
@@ -606,9 +605,9 @@ def test_update_experiment_run_model_saved_missing_status(
} }
with pytest.raises(RuntimeError): with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data)) et.update_experiment_run(input_data)
et.send_notification_async.assert_awaited_once() et.send_notification.assert_called_once()
def test_experiment_tracking_del_with_engine_no_super_del( def test_experiment_tracking_del_with_engine_no_super_del(

View File

@@ -1,7 +1,7 @@
"""Unit tests for Training activities.""" """Unit tests for Training activities."""
from contextlib import asynccontextmanager from contextlib import contextmanager
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import MagicMock, patch
import pandas as pd import pandas as pd
import pytest import pytest
@@ -49,46 +49,43 @@ def training():
) )
@pytest.mark.asyncio def test_load_model_metadata_success(training):
async def test_load_model_metadata_success(training): training.plugin_store.get_model_index = MagicMock(
training.plugin_store.get_model_index = MagicMock(return_value={'schemas': {'components': {'schemas': {}}}}) return_value={'schemas': {'components': {'schemas': {}}}}
)
inp = {**_minimal_params_dict(), 'metadata': {'w': '1'}} inp = {**_minimal_params_dict(), 'metadata': {'w': '1'}}
out = await training.load_model_metadata(inp) out = training.load_model_metadata(inp)
assert 'model_metadata' in out assert 'model_metadata' in out
assert out['model_metadata']['schemas'] assert out['model_metadata']['schemas']
@pytest.mark.asyncio def test_load_model_metadata_notifies_on_error(training):
async def test_load_model_metadata_notifies_on_error(training):
training.plugin_store.get_model_index = MagicMock(side_effect=RuntimeError('idx')) training.plugin_store.get_model_index = MagicMock(side_effect=RuntimeError('idx'))
training.send_notification_async = AsyncMock() training.send_notification = MagicMock()
inp = {**_minimal_params_dict(), 'metadata': {}} inp = {**_minimal_params_dict(), 'metadata': {}}
with pytest.raises(RuntimeError, match='idx'): with pytest.raises(RuntimeError, match='idx'):
await training.load_model_metadata(inp) training.load_model_metadata(inp)
training.send_notification_async.assert_awaited() training.send_notification.assert_called_once()
@pytest.mark.asyncio def test_validate_train_params_success(training):
async def test_validate_train_params_success(training):
pdict = _minimal_params_dict() pdict = _minimal_params_dict()
pdict['model_metadata'] = {'schemas': {'components': {'schemas': {}}}} pdict['model_metadata'] = {'schemas': {'components': {'schemas': {}}}}
inp = {**pdict, 'metadata': {}} inp = {**pdict, 'metadata': {}}
out = await training.validate_train_params(inp) out = training.validate_train_params(inp)
assert isinstance(out, TrainModelParams) assert isinstance(out, dict)
assert out.target_variable == 't' assert out['target_variable'] == 't'
@pytest.mark.asyncio def test_validate_train_params_notifies(training):
async def test_validate_train_params_notifies(training): training.send_notification = MagicMock()
training.send_notification_async = AsyncMock()
inp = {'metadata': {}, 'experiment_run_id': 1} inp = {'metadata': {}, 'experiment_run_id': 1}
with pytest.raises(Exception): with pytest.raises((KeyError, ValueError, TypeError)):
await training.validate_train_params(inp) training.validate_train_params(inp)
training.send_notification_async.assert_awaited() training.send_notification.assert_called_once()
@pytest.mark.asyncio def test_train_model_download_fails_notifies(training):
async def test_train_model_download_fails_notifies(training):
"""train_model notifies and re-raises when MinIO download fails.""" """train_model notifies and re-raises when MinIO download fails."""
tp = TrainModelParams.from_dict( tp = TrainModelParams.from_dict(
{ {
@@ -96,32 +93,31 @@ async def test_train_model_download_fails_notifies(training):
'model_metadata': {'schemas': {'components': {'schemas': {}}}}, 'model_metadata': {'schemas': {'components': {'schemas': {}}}},
} }
) )
training.minio_repository.download_file = AsyncMock(side_effect=OSError('minio')) training.minio_repository.download_file_sync = MagicMock(side_effect=OSError('minio'))
training.send_notification_async = AsyncMock()
with pytest.raises(OSError, match='minio'):
await training.train_model({'metadata': {'pod': 'x'}, 'train_params': tp})
training.send_notification_async.assert_awaited()
@pytest.mark.asyncio
async def test_cleanup_resources(training):
training.data_manager_repository.cleanup_run_directory = MagicMock()
await training.cleanup_resources({'metadata': {}, 'run_dir': '/tmp/x'})
training.data_manager_repository.cleanup_run_directory.assert_called_once_with('/tmp/x', {})
@pytest.mark.asyncio
async def test_cleanup_resources_notifies_on_error(training):
training.data_manager_repository.cleanup_run_directory = MagicMock(side_effect=RuntimeError('rm'))
training.send_notification = MagicMock() training.send_notification = MagicMock()
with pytest.raises(RuntimeError, match='rm'): with pytest.raises(OSError, match='minio'):
await training.cleanup_resources({'metadata': {'pod': 'p'}, 'run_dir': '/tmp/x'}) training.train_model({'metadata': {'pod': 'x'}, 'train_params': tp.to_dict()})
training.send_notification.assert_called_once()
def test_cleanup_resources(training):
training.data_manager_repository.cleanup_run_directory = MagicMock()
training.cleanup_resources({'metadata': {}, 'run_dir': '/tmp/x'})
training.data_manager_repository.cleanup_run_directory.assert_called_once_with('/tmp/x', {})
def test_cleanup_resources_notifies_on_error(training):
training.data_manager_repository.cleanup_run_directory = MagicMock(
side_effect=RuntimeError('rm')
)
training.send_notification = MagicMock()
with pytest.raises(RuntimeError, match='rm'):
training.cleanup_resources({'metadata': {'pod': 'p'}, 'run_dir': '/tmp/x'})
training.send_notification.assert_called_once() training.send_notification.assert_called_once()
@pytest.mark.asyncio
@patch('model_manager.activities.training.mlflow') @patch('model_manager.activities.training.mlflow')
async def test_train_model_success_serializes_result(mock_mlflow, training): def test_train_model_success_serializes_result(mock_mlflow, training):
"""Exercise train_model happy path with mocks (MinIO, plugin wrapper, MLflow).""" """Exercise train_model happy path with mocks (MinIO, plugin wrapper, MLflow)."""
tp = TrainModelParams.from_dict( tp = TrainModelParams.from_dict(
{ {
@@ -133,7 +129,7 @@ async def test_train_model_success_serializes_result(mock_mlflow, training):
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]}) val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df) tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
training.minio_repository.download_file = AsyncMock(return_value=b'csv') training.minio_repository.download_file_sync = MagicMock(return_value=b'csv')
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr) training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
training.data_manager_repository.compute_regression_metrics = MagicMock( training.data_manager_repository.compute_regression_metrics = MagicMock(
side_effect=lambda x, _w: setattr(x, 'mse_val', 0.1) or x side_effect=lambda x, _w: setattr(x, 'mse_val', 0.1) or x
@@ -159,10 +155,10 @@ async def test_train_model_success_serializes_result(mock_mlflow, training):
pred_val = pd.DataFrame({'p': [1.0]}) pred_val = pd.DataFrame({'p': [1.0]})
wrapper.predict = MagicMock(side_effect=[(pred_train, None), (pred_val, None)]) wrapper.predict = MagicMock(side_effect=[(pred_train, None), (pred_val, None)])
wrapper.store_model = MagicMock() wrapper.store_model = MagicMock()
training.plugin_store.get_model = AsyncMock(return_value=wrapper) training.plugin_store.get_model = MagicMock(return_value=wrapper)
@asynccontextmanager @contextmanager
async def _run_ctx(*_a, **_k): def _run_ctx(*_a, **_k):
info = MagicMock() info = MagicMock()
info.run_name = 'run-n' info.run_name = 'run-n'
info.run_id = 'run-i' info.run_id = 'run-i'
@@ -170,16 +166,15 @@ async def test_train_model_success_serializes_result(mock_mlflow, training):
training.mlflow_repository.start_run = _run_ctx training.mlflow_repository.start_run = _run_ctx
out = await training.train_model({'metadata': {'pod': 'p'}, 'train_params': tp}) out = training.train_model({'metadata': {'pod': 'p'}, 'train_params': tp.to_dict()})
assert out['run_name'] == 'run-n' assert out['run_name'] == 'run-n'
assert out['run_id'] == 'run-i' assert out['run_id'] == 'run-i'
assert out['run_dir'] == '/tmp/run' assert out['run_dir'] == '/tmp/run'
mock_mlflow.log_artifact.assert_called() mock_mlflow.log_artifact.assert_called()
@pytest.mark.asyncio
@patch('model_manager.activities.training.mlflow') @patch('model_manager.activities.training.mlflow')
async def test_train_model_train_params_as_dict(mock_mlflow, training): def test_train_model_train_params_as_dict(mock_mlflow, training):
"""train_params may arrive as dict and is coerced via TrainModelParams.from_dict.""" """train_params may arrive as dict and is coerced via TrainModelParams.from_dict."""
d = { d = {
**_minimal_params_dict(), **_minimal_params_dict(),
@@ -190,9 +185,12 @@ async def test_train_model_train_params_as_dict(mock_mlflow, training):
tp = TrainModelParams.from_dict(d) tp = TrainModelParams.from_dict(d)
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df) tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
training.minio_repository.download_file = AsyncMock(return_value=b'csv') training.minio_repository.download_file_sync = MagicMock(return_value=b'csv')
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr) training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
training.data_manager_repository.compute_regression_metrics = MagicMock(side_effect=lambda x, _w: x) training.data_manager_repository.compute_regression_metrics = MagicMock(
side_effect=lambda x, _w: x
)
def _fill_report2(x, **_kw): def _fill_report2(x, **_kw):
x.report_path = '/tmp/report.html' x.report_path = '/tmp/report.html'
x.train_data_path = '/tmp/train.csv' x.train_data_path = '/tmp/train.csv'
@@ -207,10 +205,10 @@ async def test_train_model_train_params_as_dict(mock_mlflow, training):
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)] side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
) )
wrapper.store_model = MagicMock() wrapper.store_model = MagicMock()
training.plugin_store.get_model = AsyncMock(return_value=wrapper) training.plugin_store.get_model = MagicMock(return_value=wrapper)
@asynccontextmanager @contextmanager
async def _run_ctx(*_a, **_k): def _run_ctx(*_a, **_k):
info = MagicMock() info = MagicMock()
info.run_name = 'n' info.run_name = 'n'
info.run_id = 'i' info.run_id = 'i'
@@ -218,13 +216,12 @@ async def test_train_model_train_params_as_dict(mock_mlflow, training):
training.mlflow_repository.start_run = _run_ctx training.mlflow_repository.start_run = _run_ctx
await training.train_model({'metadata': {}, 'train_params': d}) training.train_model({'metadata': {}, 'train_params': d})
mock_mlflow.log_artifact.assert_called() mock_mlflow.log_artifact.assert_called()
@pytest.mark.asyncio
@patch('model_manager.activities.training.mlflow') @patch('model_manager.activities.training.mlflow')
async def test_train_model_downloads_validation_file_when_set(mock_mlflow, training): def test_train_model_downloads_validation_file_when_set(mock_mlflow, training):
"""Second MinIO download when val_file_name is set (covers val_bytes branch).""" """Second MinIO download when val_file_name is set (covers val_bytes branch)."""
d = { d = {
**_minimal_params_dict(), **_minimal_params_dict(),
@@ -236,16 +233,18 @@ async def test_train_model_downloads_validation_file_when_set(mock_mlflow, train
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]}) val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df) tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
async def _dl(object_name, **_kwargs): def _dl(object_name, **_kwargs):
if object_name == tp.file_name: if object_name == tp.file_name:
return b'train' return b'train'
if object_name == 'val.csv': if object_name == 'val.csv':
return b'val' return b'val'
raise AssertionError(object_name) raise AssertionError(object_name)
training.minio_repository.download_file = AsyncMock(side_effect=_dl) training.minio_repository.download_file_sync = MagicMock(side_effect=_dl)
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr) training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
training.data_manager_repository.compute_regression_metrics = MagicMock(side_effect=lambda x, _w: x) training.data_manager_repository.compute_regression_metrics = MagicMock(
side_effect=lambda x, _w: x
)
def _fill(x, **_kw): def _fill(x, **_kw):
x.report_path = '/tmp/report.html' x.report_path = '/tmp/report.html'
@@ -261,10 +260,10 @@ async def test_train_model_downloads_validation_file_when_set(mock_mlflow, train
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)] side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
) )
wrapper.store_model = MagicMock() wrapper.store_model = MagicMock()
training.plugin_store.get_model = AsyncMock(return_value=wrapper) training.plugin_store.get_model = MagicMock(return_value=wrapper)
@asynccontextmanager @contextmanager
async def _run_ctx(*_a, **_k): def _run_ctx(*_a, **_k):
info = MagicMock() info = MagicMock()
info.run_name = 'n' info.run_name = 'n'
info.run_id = 'i' info.run_id = 'i'
@@ -272,13 +271,12 @@ async def test_train_model_downloads_validation_file_when_set(mock_mlflow, train
training.mlflow_repository.start_run = _run_ctx training.mlflow_repository.start_run = _run_ctx
await training.train_model({'metadata': {}, 'train_params': tp}) training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
assert training.minio_repository.download_file.await_count == 2 assert training.minio_repository.download_file_sync.call_count == 2
mock_mlflow.log_artifact.assert_called() mock_mlflow.log_artifact.assert_called()
@pytest.mark.asyncio def test_train_model_value_error_when_paths_missing_after_report(training):
async def test_train_model_value_error_when_paths_missing_after_report(training):
"""Raises ValueError when report paths are not populated after generate_report.""" """Raises ValueError when report paths are not populated after generate_report."""
tp = TrainModelParams.from_dict( tp = TrainModelParams.from_dict(
{ {
@@ -290,26 +288,28 @@ async def test_train_model_value_error_when_paths_missing_after_report(training)
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]}) val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df) tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
training.minio_repository.download_file = AsyncMock(return_value=b'x') training.minio_repository.download_file_sync = MagicMock(return_value=b'x')
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr) training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
training.data_manager_repository.compute_regression_metrics = MagicMock(side_effect=lambda x, _w: x) training.data_manager_repository.compute_regression_metrics = MagicMock(
side_effect=lambda x, _w: x
)
training.data_manager_repository.generate_report = MagicMock(return_value=tmr) training.data_manager_repository.generate_report = MagicMock(return_value=tmr)
wrapper = MagicMock() wrapper = MagicMock()
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)]) wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
wrapper.predict = MagicMock( wrapper.predict = MagicMock(
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)] side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
) )
training.plugin_store.get_model = AsyncMock(return_value=wrapper) training.plugin_store.get_model = MagicMock(return_value=wrapper)
@asynccontextmanager @contextmanager
async def _run_ctx(*_a, **_k): def _run_ctx(*_a, **_k):
info = MagicMock() info = MagicMock()
info.run_name = 'n' info.run_name = 'n'
info.run_id = 'i' info.run_id = 'i'
yield info yield info
training.mlflow_repository.start_run = _run_ctx training.mlflow_repository.start_run = _run_ctx
training.send_notification_async = AsyncMock() training.send_notification = MagicMock()
with pytest.raises(ValueError, match='Report path'): with pytest.raises(ValueError, match='Report path'):
await training.train_model({'metadata': {}, 'train_params': tp}) training.train_model({'metadata': {}, 'train_params': tp.to_dict()})

View File

@@ -162,7 +162,11 @@ def test_validate_model_param_schema_validation_error(valid_train_params_dict):
'schemas': { 'schemas': {
'components': { 'components': {
'schemas': { 'schemas': {
'data_model': {'type': 'object', 'properties': {'x': {'type': 'integer'}}, 'required': ['x']}, 'data_model': {
'type': 'object',
'properties': {'x': {'type': 'integer'}},
'required': ['x'],
},
} }
} }
} }
@@ -187,7 +191,7 @@ def test_validate_model_param_unexpected_validator_error(valid_train_params_dict
p = TrainModelParams.from_dict(d) p = TrainModelParams.from_dict(d)
with patch('model_manager.utils.models.train_model_params.Draft202012Validator') as m: with patch('model_manager.utils.models.train_model_params.Draft202012Validator') as m:
m.return_value.validate.side_effect = RuntimeError('boom') m.return_value.validate.side_effect = RuntimeError('boom')
with pytest.raises(ValueError, match='Unexpected error'): with pytest.raises(RuntimeError, match='boom'):
p.validate_business_rules() p.validate_business_rules()
@@ -224,9 +228,7 @@ def test_validate_model_param_only_data_model_schema(valid_train_params_dict):
def test_validate_model_param_only_model_schema(valid_train_params_dict): def test_validate_model_param_only_model_schema(valid_train_params_dict):
d = copy.deepcopy(valid_train_params_dict) d = copy.deepcopy(valid_train_params_dict)
d['model_metadata'] = { d['model_metadata'] = {'schemas': {'components': {'schemas': {'model': {'type': 'object'}}}}}
'schemas': {'components': {'schemas': {'model': {'type': 'object'}}}}
}
p = TrainModelParams.from_dict(d) p = TrainModelParams.from_dict(d)
p.model_kwargs = {} p.model_kwargs = {}
p.validate_business_rules() p.validate_business_rules()

View File

@@ -3,6 +3,7 @@
from __future__ import annotations from __future__ import annotations
import json import json
from typing import Any
from unittest.mock import MagicMock, Mock, patch from unittest.mock import MagicMock, Mock, patch
import numpy as np import numpy as np
@@ -33,7 +34,7 @@ def test_train_test_split_ndarray():
def _params(**kwargs) -> TrainModelParams: def _params(**kwargs) -> TrainModelParams:
base = { base: dict[str, Any] = {
'variable_columns': ['v1'], 'variable_columns': ['v1'],
'target_variable': 't', 'target_variable': 't',
'bucket_name': 'b', 'bucket_name': 'b',
@@ -276,7 +277,13 @@ def test_configure_datetime_index_already_datetime_index():
def test_configure_datetime_index_from_common_column(): def test_configure_datetime_index_from_common_column():
repo = dmr.DataManagerRepository(MagicMock()) repo = dmr.DataManagerRepository(MagicMock())
p = TrainModelParams.from_dict(_minimal_dict_for_prepare()) p = TrainModelParams.from_dict(_minimal_dict_for_prepare())
df = pd.DataFrame({'timestamp': pd.date_range('2024-01-01', periods=3, freq='D'), 'v1': [1, 2, 3], 't': [1, 2, 3]}) df = pd.DataFrame(
{
'timestamp': pd.date_range('2024-01-01', periods=3, freq='D'),
'v1': [1, 2, 3],
't': [1, 2, 3],
}
)
out = repo._configure_datetime_index(df, p, {}) out = repo._configure_datetime_index(df, p, {})
assert isinstance(out.index, pd.DatetimeIndex) assert isinstance(out.index, pd.DatetimeIndex)
@@ -329,14 +336,20 @@ def test_configure_datetime_index_first_column_numeric_parsed_as_time():
def test_create_run_directory_permission_error(): def test_create_run_directory_permission_error():
repo = dmr.DataManagerRepository(MagicMock()) repo = dmr.DataManagerRepository(MagicMock())
with patch('model_manager.utils.repository.data_manager_repository.makedirs', side_effect=PermissionError('no')): with patch(
'model_manager.utils.repository.data_manager_repository.makedirs',
side_effect=PermissionError('no'),
):
with pytest.raises(PermissionError, match='Permission denied'): with pytest.raises(PermissionError, match='Permission denied'):
repo._create_run_directory('/tmp', 'run', {}) repo._create_run_directory('/tmp', 'run', {})
def test_create_run_directory_os_error(): def test_create_run_directory_os_error():
repo = dmr.DataManagerRepository(MagicMock()) repo = dmr.DataManagerRepository(MagicMock())
with patch('model_manager.utils.repository.data_manager_repository.makedirs', side_effect=OSError('disk')): with patch(
'model_manager.utils.repository.data_manager_repository.makedirs',
side_effect=OSError('disk'),
):
with pytest.raises(OSError, match='Failed to create directory'): with pytest.raises(OSError, match='Failed to create directory'):
repo._create_run_directory('/tmp', 'run', {}) repo._create_run_directory('/tmp', 'run', {})
@@ -355,7 +368,6 @@ def test_generate_report_success(tmp_path):
patch.object(repo, '_get_reports_directory', return_value=str(tmp_path)), patch.object(repo, '_get_reports_directory', return_value=str(tmp_path)),
patch('model_manager.utils.repository.data_manager_repository.Reports') as mrep, patch('model_manager.utils.repository.data_manager_repository.Reports') as mrep,
): ):
inst = mrep.return_value
instance = mrep.return_value instance = mrep.return_value
instance.save_all_sections_html = Mock() instance.save_all_sections_html = Mock()
out = repo.generate_report(tmr, {}) out = repo.generate_report(tmr, {})

View File

@@ -0,0 +1,76 @@
"""Unit tests for local worker factory."""
from unittest.mock import MagicMock, patch
def test_prepare_worker_train_queue_uses_train_limits():
from model_manager.worker.prepare_worker import prepare_worker
from model_manager.workflows.train_model import TrainModel
fake_worker = MagicMock()
fake_client = MagicMock()
fake_logger = MagicMock()
with patch(
'model_manager.worker.prepare_worker.Worker', return_value=fake_worker
) as worker_class:
with patch.dict(
'os.environ',
{
'TRAINMODEL_ACTIVITY_EXECUTOR_MAX_WORKERS': '3',
'TRAINMODEL_MAX_CONCURRENT_ACTIVITIES': '6',
'TRAINMODEL_MAX_CONCURRENT_WORKFLOW_TASKS': '10',
},
clear=False,
):
worker = prepare_worker(
main_workflow=TrainModel,
other_workflows=[],
activities=[],
temporal_client=fake_client,
logger=fake_logger,
)
assert worker is fake_worker
worker_class.assert_called_once()
kwargs = worker_class.call_args.kwargs
assert kwargs['task_queue'] == 'train_model-queue'
assert kwargs['max_concurrent_activities'] == 6
assert kwargs['max_concurrent_workflow_tasks'] == 10
assert kwargs['activity_executor']._max_workers == 3
kwargs['activity_executor'].shutdown(wait=True, cancel_futures=True)
def test_prepare_worker_cleanup_queue_uses_cleanup_limits():
from model_manager.worker.prepare_worker import prepare_worker
from model_manager.workflows.cleanup_files import CleanupFiles
fake_worker = MagicMock()
fake_client = MagicMock()
fake_logger = MagicMock()
with patch(
'model_manager.worker.prepare_worker.Worker', return_value=fake_worker
) as worker_class:
with patch.dict(
'os.environ',
{
'CLEANUPFILES_ACTIVITY_EXECUTOR_MAX_WORKERS': '5',
'CLEANUPFILES_MAX_CONCURRENT_ACTIVITIES': '7',
},
clear=False,
):
worker = prepare_worker(
main_workflow=CleanupFiles,
other_workflows=[],
activities=[],
temporal_client=fake_client,
logger=fake_logger,
)
assert worker is fake_worker
kwargs = worker_class.call_args.kwargs
assert kwargs['task_queue'] == 'cleanup_files-queue'
assert kwargs['max_concurrent_activities'] == 7
assert kwargs['activity_executor']._max_workers == 5
kwargs['activity_executor'].shutdown(wait=True, cancel_futures=True)

View File

@@ -125,7 +125,7 @@ def test_start_prometheus_server_success(
mock_app_up = Mock() mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up mock_metrics.APP_UP.labels.return_value = mock_app_up
metadata = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'} metadata: dict[str, str | None] = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'}
start_prometheus_server(mock_logger, metadata) start_prometheus_server(mock_logger, metadata)
@@ -148,7 +148,7 @@ def test_start_prometheus_server_custom_port(mock_metrics, mock_start_http_serve
mock_app_up = Mock() mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up mock_metrics.APP_UP.labels.return_value = mock_app_up
metadata = {'pod_id': 'custom-pod', 'workflow_name': 'train_model'} metadata: dict[str, str | None] = {'pod_id': 'custom-pod', 'workflow_name': 'train_model'}
start_prometheus_server(mock_logger, metadata) start_prometheus_server(mock_logger, metadata)
@@ -166,7 +166,7 @@ def test_start_prometheus_server_failure(
mock_start_http_server.side_effect = OSError('Port already in use') mock_start_http_server.side_effect = OSError('Port already in use')
metadata = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'} metadata: dict[str, str | None] = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'}
start_prometheus_server(mock_logger, metadata) start_prometheus_server(mock_logger, metadata)
@@ -252,9 +252,7 @@ async def test_main_successful_startup(
mock_runtime_class.return_value = mock_runtime mock_runtime_class.return_value = mock_runtime
mock_client_instance = AsyncMock() mock_client_instance = AsyncMock()
mock_client_instance.config = Mock( mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
return_value={'plugins': [], 'interceptors': []}
)
mock_client_class.connect = AsyncMock(return_value=mock_client_instance) mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
mock_worker_instance = Mock() mock_worker_instance = Mock()
@@ -347,9 +345,7 @@ async def test_main_handles_exception(
mock_runtime_class.return_value = mock_runtime mock_runtime_class.return_value = mock_runtime
mock_client_instance = AsyncMock() mock_client_instance = AsyncMock()
mock_client_instance.config = Mock( mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
return_value={'plugins': [], 'interceptors': []}
)
mock_client_class.connect = AsyncMock(return_value=mock_client_instance) mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
mock_worker_instance = Mock() mock_worker_instance = Mock()
@@ -486,9 +482,7 @@ async def test_main_temporal_client_configuration(
mock_runtime_class.return_value = mock_runtime mock_runtime_class.return_value = mock_runtime
mock_client_instance = AsyncMock() mock_client_instance = AsyncMock()
mock_client_instance.config = Mock( mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
return_value={'plugins': [], 'interceptors': []}
)
mock_client_class.connect = AsyncMock(return_value=mock_client_instance) mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
mock_worker_instance = Mock() mock_worker_instance = Mock()
@@ -581,9 +575,7 @@ async def test_main_worker_configuration(
mock_runtime_class.return_value = mock_runtime mock_runtime_class.return_value = mock_runtime
mock_client_instance = AsyncMock() mock_client_instance = AsyncMock()
mock_client_instance.config = Mock( mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
return_value={'plugins': [], 'interceptors': []}
)
mock_client_class.connect = AsyncMock(return_value=mock_client_instance) mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
mock_worker_instance = Mock() mock_worker_instance = Mock()
@@ -703,9 +695,7 @@ async def test_main_schedule_creation_failure_does_not_stop_worker(
mock_runtime_class.return_value = mock_runtime mock_runtime_class.return_value = mock_runtime
mock_client_instance = AsyncMock() mock_client_instance = AsyncMock()
mock_client_instance.config = Mock( mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
return_value={'plugins': [], 'interceptors': []}
)
mock_client_class.connect = AsyncMock(return_value=mock_client_instance) mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
mock_worker_instance = Mock() mock_worker_instance = Mock()
@@ -844,7 +834,7 @@ def test_start_prometheus_server_prints_success(
mock_app_up = Mock() mock_app_up = Mock()
mock_metrics.APP_UP.labels.return_value = mock_app_up mock_metrics.APP_UP.labels.return_value = mock_app_up
metadata = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'} metadata: dict[str, str | None] = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'}
start_prometheus_server(mock_logger, metadata) start_prometheus_server(mock_logger, metadata)
@@ -863,7 +853,7 @@ def test_start_prometheus_server_prints_failure(
mock_start_http_server.side_effect = Exception('Test error') mock_start_http_server.side_effect = Exception('Test error')
metadata = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'} metadata: dict[str, str | None] = {'pod_id': 'test-pod-123', 'workflow_name': 'train_model'}
start_prometheus_server(mock_logger, metadata) start_prometheus_server(mock_logger, metadata)

View File

@@ -243,7 +243,9 @@ async def test_run_validation_error(mock_wf, sample_input_data):
@pytest.mark.asyncio @pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow') @patch('model_manager.workflows.train_model.workflow')
async def test_run_cleanup_failure_does_not_fail_workflow(mock_wf, sample_input_data, mock_train_params): async def test_run_cleanup_failure_does_not_fail_workflow(
mock_wf, sample_input_data, mock_train_params
):
"""After successful training, cleanup failure is logged, workflow still returns result.""" """After successful training, cleanup failure is logged, workflow still returns result."""
from model_manager.workflows.train_model import TrainModel from model_manager.workflows.train_model import TrainModel
@@ -267,7 +269,9 @@ async def test_run_cleanup_failure_does_not_fail_workflow(mock_wf, sample_input_
@pytest.mark.asyncio @pytest.mark.asyncio
@patch('model_manager.workflows.train_model.workflow') @patch('model_manager.workflows.train_model.workflow')
async def test_run_training_failure_skips_cleanup_activity(mock_wf, sample_input_data, mock_train_params): async def test_run_training_failure_skips_cleanup_activity(
mock_wf, sample_input_data, mock_train_params
):
"""When train_model raises, train_result stays None and cleanup activity is not scheduled.""" """When train_model raises, train_result stays None and cleanup activity is not scheduled."""
from model_manager.workflows.train_model import TrainModel from model_manager.workflows.train_model import TrainModel
@@ -297,7 +301,6 @@ def test_module_constants():
from model_manager.workflows.train_model import ( from model_manager.workflows.train_model import (
TIMEOUT_DELETE_FILE, TIMEOUT_DELETE_FILE,
TIMEOUT_TRAIN_MODEL, TIMEOUT_TRAIN_MODEL,
TIMEOUT_UPDATE_DATABASE,
TIMEOUT_VALIDATE_PARAMS, TIMEOUT_VALIDATE_PARAMS,
database_retry_policy, database_retry_policy,
network_retry_policy, network_retry_policy,