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:
@@ -67,7 +67,6 @@ class Activities(ExperimentTracking, Training, Cleanup):
|
||||
Exception: If any parent class initialization fails
|
||||
"""
|
||||
|
||||
|
||||
ExperimentTracking.__init__(
|
||||
self,
|
||||
host=postgres_config['host'],
|
||||
|
||||
@@ -62,7 +62,7 @@ class Cleanup(SientiaMonitoring):
|
||||
) # name_YYYYMMDD_HHMMSS_microseconds
|
||||
|
||||
@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.
|
||||
|
||||
@@ -178,14 +178,14 @@ class Cleanup(SientiaMonitoring):
|
||||
|
||||
raise
|
||||
finally:
|
||||
await self._emit_metrics(
|
||||
self._emit_metrics(
|
||||
metadata=metadata,
|
||||
metrics_status=metrics_status,
|
||||
activity_name='cleanup_temp_directories',
|
||||
emit_workflow_metric=True,
|
||||
)
|
||||
|
||||
async def _emit_metrics(
|
||||
def _emit_metrics(
|
||||
self,
|
||||
metadata: dict[str, Any],
|
||||
metrics_status: str,
|
||||
@@ -201,7 +201,7 @@ class Cleanup(SientiaMonitoring):
|
||||
activity_name: Name of the activity being executed
|
||||
"""
|
||||
if emit_workflow_metric:
|
||||
await self.emit_metric(
|
||||
self.emit_metric_sync(
|
||||
metric_object=WORKFLOW_EXECUTION_TOTAL,
|
||||
tags={
|
||||
'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,
|
||||
tags={
|
||||
'pod_id': metadata.get('pod_id'),
|
||||
|
||||
@@ -11,7 +11,6 @@ import enum
|
||||
from temporalio import activity, workflow
|
||||
|
||||
with workflow.unsafe.imports_passed_through():
|
||||
import asyncio
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
from datetime import UTC, datetime
|
||||
@@ -110,9 +109,9 @@ class ExperimentTracking(Postgres):
|
||||
# Silently ignore errors during garbage collection
|
||||
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:
|
||||
query: Parameterized SQL string to execute.
|
||||
@@ -122,12 +121,9 @@ class ExperimentTracking(Postgres):
|
||||
dict: A dictionary containing the affected row count: {'rowcount': int}.
|
||||
"""
|
||||
|
||||
def _run() -> dict[str, Any]:
|
||||
with self.engine.begin() as connection:
|
||||
result = connection.execute(text(query), params)
|
||||
return {'rowcount': result.rowcount}
|
||||
|
||||
return await asyncio.to_thread(_run)
|
||||
with self.engine.begin() as connection:
|
||||
result = connection.execute(text(query), params)
|
||||
return {'rowcount': result.rowcount}
|
||||
|
||||
def _build_status_update_query(
|
||||
self, status: str | None, experiment_run_id: int
|
||||
@@ -223,7 +219,7 @@ class ExperimentTracking(Postgres):
|
||||
raise ValueError(f'Invalid update_type: {update_type}')
|
||||
|
||||
@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.
|
||||
|
||||
@@ -256,7 +252,7 @@ class ExperimentTracking(Postgres):
|
||||
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:
|
||||
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)}'
|
||||
trace = traceback.format_exc()
|
||||
|
||||
await self.send_notification_async(
|
||||
self.send_notification(
|
||||
metadata=metadata or {},
|
||||
notification_id='UPDATE_EXPERIMENT_RUN_ERROR',
|
||||
message=error_msg,
|
||||
|
||||
@@ -12,6 +12,7 @@ with workflow.unsafe.imports_passed_through():
|
||||
import traceback
|
||||
from typing import Any
|
||||
|
||||
import mlflow
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
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 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
|
||||
|
||||
import mlflow
|
||||
from model_manager.utils.repository.data_manager_repository import DataManagerRepository
|
||||
|
||||
|
||||
class Training(SientiaMonitoring):
|
||||
@@ -61,7 +60,7 @@ class Training(SientiaMonitoring):
|
||||
self.minio_repository = minio_repository
|
||||
|
||||
@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.
|
||||
|
||||
@@ -90,7 +89,7 @@ class Training(SientiaMonitoring):
|
||||
return train_params.to_dict()
|
||||
except Exception as exc:
|
||||
trace = traceback.format_exc()
|
||||
await self.send_notification_async(
|
||||
self.send_notification(
|
||||
metadata=metadata,
|
||||
notification_id='LOAD_MODEL_METADATA_ERROR',
|
||||
message=f'Error loading model metadata: {str(exc)}',
|
||||
@@ -99,9 +98,9 @@ class Training(SientiaMonitoring):
|
||||
attachment_content=trace,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
@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.
|
||||
|
||||
@@ -115,7 +114,7 @@ class Training(SientiaMonitoring):
|
||||
- All TrainModelParams fields (experiment_run_id, target_variable, etc.)
|
||||
|
||||
Returns:
|
||||
TrainModelParams: Validated and converted training parameters
|
||||
dict[str, Any]: Validated and converted training parameters as dictionary
|
||||
|
||||
Raises:
|
||||
Exception: If validation fails (after sending notification)
|
||||
@@ -123,7 +122,7 @@ class Training(SientiaMonitoring):
|
||||
metadata = input_data.get('metadata', {})
|
||||
try:
|
||||
train_params = TrainModelParams.from_dict(input_data)
|
||||
|
||||
|
||||
train_params.validate_business_rules()
|
||||
|
||||
self.info(
|
||||
@@ -133,12 +132,12 @@ class Training(SientiaMonitoring):
|
||||
metadata,
|
||||
)
|
||||
|
||||
return train_params
|
||||
return train_params.to_dict()
|
||||
except Exception as e:
|
||||
error_msg = f'Error validating training parameters: {str(e)}'
|
||||
trace = traceback.format_exc()
|
||||
|
||||
await self.send_notification_async(
|
||||
self.send_notification(
|
||||
metadata=metadata,
|
||||
notification_id='VALIDATE_TRAIN_PARAMS_ERROR',
|
||||
message=error_msg,
|
||||
@@ -149,7 +148,7 @@ class Training(SientiaMonitoring):
|
||||
raise
|
||||
|
||||
@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.
|
||||
|
||||
@@ -162,7 +161,7 @@ class Training(SientiaMonitoring):
|
||||
input_data: Training configuration containing:
|
||||
- metadata (dict): Workflow execution metadata.
|
||||
- uploaded_file (BytesIO): Training data already downloaded from MinIO.
|
||||
- train_params (TrainModelParams | dict): Training parameters.
|
||||
- train_params (dict): Training parameters.
|
||||
|
||||
Returns:
|
||||
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).
|
||||
"""
|
||||
metadata = input_data.get('metadata')
|
||||
train_params = input_data['train_params']
|
||||
|
||||
if isinstance(train_params, dict):
|
||||
train_params = TrainModelParams.from_dict(train_params)
|
||||
train_params = TrainModelParams.from_dict(input_data['train_params'])
|
||||
|
||||
try:
|
||||
# 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,
|
||||
bucket=train_params.bucket_name,
|
||||
metadata=metadata,
|
||||
@@ -189,7 +185,7 @@ class Training(SientiaMonitoring):
|
||||
val_bytes: bytes | None = None
|
||||
validation_name = train_params.val_file_name
|
||||
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,
|
||||
bucket=train_params.bucket_name,
|
||||
metadata=metadata,
|
||||
@@ -202,13 +198,13 @@ class Training(SientiaMonitoring):
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
wrapper = await self.plugin_store.get_model(
|
||||
wrapper = self.plugin_store.get_model(
|
||||
model_name=train_params.model_name,
|
||||
force_download=False,
|
||||
opt_params=train_params.opt_params or {},
|
||||
model_kwargs=train_params.model_kwargs or {},
|
||||
data_model_kwargs=train_params.data_model_kwargs or {},
|
||||
metadata=metadata
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
train_data = train_result.train_data
|
||||
@@ -238,7 +234,7 @@ class Training(SientiaMonitoring):
|
||||
wrapper,
|
||||
)
|
||||
|
||||
async with self.mlflow_repository.start_run(
|
||||
with self.mlflow_repository.start_run(
|
||||
model_name=train_params.model_name,
|
||||
run_name=None,
|
||||
experiment_name=f'{train_params.model_name}_experiment',
|
||||
@@ -247,33 +243,19 @@ class Training(SientiaMonitoring):
|
||||
) as run_info:
|
||||
train_result.run_name = run_info.run_name
|
||||
train_result.run_id = run_info.run_id
|
||||
|
||||
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)
|
||||
self._persist_training_artifacts(train_result, train_params, wrapper, metadata)
|
||||
|
||||
return {
|
||||
'run_name': train_result.run_name,
|
||||
'run_id': train_result.run_id,
|
||||
'run_dir': train_result.run_dir
|
||||
'run_dir': train_result.run_dir,
|
||||
}
|
||||
except Exception as e: # noqa: BLE001
|
||||
|
||||
error_msg = f'Error training model - error: {str(e)}'
|
||||
|
||||
trace = traceback.format_exc()
|
||||
|
||||
await self.send_notification_async(
|
||||
self.send_notification(
|
||||
metadata=metadata or {},
|
||||
notification_id='TRAIN_MODEL_ERROR',
|
||||
message=error_msg,
|
||||
@@ -284,8 +266,32 @@ class Training(SientiaMonitoring):
|
||||
|
||||
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')
|
||||
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.
|
||||
|
||||
@@ -317,4 +323,3 @@ class Training(SientiaMonitoring):
|
||||
)
|
||||
|
||||
raise
|
||||
|
||||
|
||||
@@ -62,8 +62,8 @@ class TrainModelParams:
|
||||
|
||||
# New Parameters
|
||||
val_file_name: str | None
|
||||
data_model_kwargs: dict | None # Removed params used in DataPreprocessor here
|
||||
model_kwargs: dict | None # Removed params used in Linear Regression Model here
|
||||
data_model_kwargs: dict | None # Removed params used in DataPreprocessor here
|
||||
model_kwargs: dict | None # Removed params used in Linear Regression Model here
|
||||
opt_params: dict | None
|
||||
model_type: str
|
||||
model_id: str | None
|
||||
@@ -102,12 +102,16 @@ class TrainModelParams:
|
||||
model_name = cls._check_none(data.get('model_name'), str, 'model_name')
|
||||
|
||||
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'),
|
||||
bucket_name=cls._check_none(data.get('bucket_name'), str, 'bucket_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'),
|
||||
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_format=data.get('date_format'),
|
||||
train_size=cls._check_none(data.get('train_size'), int, 'train_size'),
|
||||
@@ -117,12 +121,13 @@ class TrainModelParams:
|
||||
model_name=model_name,
|
||||
experiment_name=model_name + '_experiment',
|
||||
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'),
|
||||
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_id=data.get('model_id'),
|
||||
|
||||
model_metadata=cls._parse_optional_model_metadata(data.get('model_metadata')),
|
||||
)
|
||||
|
||||
@@ -233,9 +238,7 @@ class TrainModelParams:
|
||||
return None
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
raise TypeError(
|
||||
f'model_metadata must be a dict or None, but got {type(value).__name__}.'
|
||||
)
|
||||
raise TypeError(f'model_metadata must be a dict or None, but got {type(value).__name__}.')
|
||||
|
||||
def validate_business_rules(self) -> None:
|
||||
"""
|
||||
@@ -261,21 +264,20 @@ class TrainModelParams:
|
||||
|
||||
if not self.variable_columns:
|
||||
raise ValueError('variable_columns cannot be empty')
|
||||
|
||||
|
||||
def _validate_model_params(self) -> None:
|
||||
"""Validate model-related parameters."""
|
||||
if not self.model_metadata:
|
||||
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:
|
||||
return
|
||||
|
||||
data_model_schema = schemas.get("data_model")
|
||||
model_schema = schemas.get("model")
|
||||
opt_params_schema = schemas.get("opt_params")
|
||||
data_model_schema = schemas.get('data_model')
|
||||
model_schema = schemas.get('model')
|
||||
opt_params_schema = schemas.get('opt_params')
|
||||
|
||||
if data_model_schema:
|
||||
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)
|
||||
if opt_params_schema:
|
||||
self._validate_model_param(opt_params_schema, self.opt_params)
|
||||
|
||||
|
||||
|
||||
def _validate_model_param(self, schema: dict[str, Any], value: Any) -> None:
|
||||
"""Validate model parameter against schema."""
|
||||
@@ -292,9 +292,7 @@ class TrainModelParams:
|
||||
validator = Draft202012Validator(schema)
|
||||
validator.validate(value)
|
||||
except ValidationError as e:
|
||||
raise ValueError(f'Model parameters validation failed: {e.message}')
|
||||
except Exception as e:
|
||||
raise ValueError(f'Unexpected error: {e}')
|
||||
raise ValueError(f'Model parameters validation failed: {e.message}') from e
|
||||
|
||||
def _validate_required_strings(self) -> None:
|
||||
"""Validate required string fields are not empty."""
|
||||
@@ -313,4 +311,4 @@ class TrainModelParams:
|
||||
def _validate_date_format(self) -> None:
|
||||
"""Validate date_format is one of the allowed frontend formats when set."""
|
||||
if self.date_format:
|
||||
validate_frontend_date_format(self.date_format)
|
||||
validate_frontend_date_format(self.date_format)
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import pandas as pd
|
||||
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from io import BytesIO
|
||||
import json
|
||||
from os import makedirs, path
|
||||
from typing import Any
|
||||
from shutil import rmtree
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
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 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_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
|
||||
if random_state is not None:
|
||||
np.random.seed(random_state)
|
||||
|
||||
|
||||
# 2. Gerar índices e embaralhar se necessário
|
||||
indices = np.arange(len(data))
|
||||
|
||||
|
||||
if shuffle:
|
||||
np.random.shuffle(indices)
|
||||
|
||||
|
||||
# 3. Calcular o ponto de corte (split point)
|
||||
# Cálculo: N_treino = tamanho_total * proporcao_treino
|
||||
n_train = int(len(data) * train_size)
|
||||
|
||||
|
||||
# 4. Dividir os índices
|
||||
train_indices = indices[:n_train]
|
||||
test_indices = indices[n_train:]
|
||||
|
||||
|
||||
# 5. Retornar os dados fatiados (funciona para DataFrame ou Series)
|
||||
if isinstance(data, (pd.DataFrame, pd.Series)):
|
||||
return data.iloc[train_indices], data.iloc[test_indices]
|
||||
|
||||
|
||||
return data[train_indices], data[test_indices]
|
||||
|
||||
|
||||
def _ensure_date_column_parsed(data: pd.DataFrame, params: TrainModelParams) -> pd.DataFrame:
|
||||
"""
|
||||
If date_column is set, parse the column as timezone-aware
|
||||
@@ -178,7 +185,6 @@ class DataManagerRepository(SientiaMonitoring):
|
||||
if len(val_df) <= 0:
|
||||
raise ValueError('Validation data view is empty after transformation')
|
||||
|
||||
|
||||
val_data = pd.DataFrame(val_df[params.variable_columns + [params.target_variable]])
|
||||
else:
|
||||
# Fallback path: derive validation via train/test split from a single dataset.
|
||||
@@ -194,11 +200,7 @@ class DataManagerRepository(SientiaMonitoring):
|
||||
metadata,
|
||||
)
|
||||
|
||||
return TrainModelResult(
|
||||
params=params,
|
||||
train_data=train_data,
|
||||
val_data=val_data
|
||||
)
|
||||
return TrainModelResult(params=params, train_data=train_data, val_data=val_data)
|
||||
|
||||
def _as_series(self, pred: pd.DataFrame | pd.Series) -> pd.Series:
|
||||
if isinstance(pred, pd.Series):
|
||||
@@ -207,10 +209,8 @@ class DataManagerRepository(SientiaMonitoring):
|
||||
if pred.shape[1] == 1:
|
||||
return pred.iloc[:, 0]
|
||||
raise ValueError('y_pred/y_train_pred must be a Series or single-column DataFrame')
|
||||
|
||||
def _extract_model_equation(
|
||||
self, regr: Any, params: TrainModelParams
|
||||
) -> dict:
|
||||
|
||||
def _extract_model_equation(self, regr: Any, params: TrainModelParams) -> dict:
|
||||
"""
|
||||
Extract the linear regression equation coefficients and create equation metadata.
|
||||
|
||||
@@ -237,7 +237,7 @@ class DataManagerRepository(SientiaMonitoring):
|
||||
model_kwargs = params.model_kwargs or {}
|
||||
degree = model_kwargs.get('degree', 1)
|
||||
poly_feature_names = model_kwargs.get('poly_feature_names', None)
|
||||
|
||||
|
||||
if degree > 1 and poly_feature_names:
|
||||
feature_names = poly_feature_names
|
||||
else:
|
||||
@@ -271,7 +271,6 @@ class DataManagerRepository(SientiaMonitoring):
|
||||
'original_features': feature_names,
|
||||
}
|
||||
|
||||
|
||||
def compute_regression_metrics(
|
||||
self,
|
||||
tmr: TrainModelResult,
|
||||
@@ -301,7 +300,7 @@ class DataManagerRepository(SientiaMonitoring):
|
||||
|
||||
# True values are expected to come from val_data.
|
||||
y_true_val = tmr.val_data[target]
|
||||
|
||||
|
||||
y_pred_val = self._as_series(tmr.y_pred).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')
|
||||
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.
|
||||
|
||||
@@ -471,7 +472,6 @@ class DataManagerRepository(SientiaMonitoring):
|
||||
reference_data_float = reference_data.astype(np.float64)
|
||||
current_data_float = current_data.astype(np.float64)
|
||||
|
||||
|
||||
# Initialize report generator
|
||||
data.run_dir = self._create_run_directory(self._get_reports_directory(), data.run_name)
|
||||
report = Reports(
|
||||
|
||||
106
model_manager/worker/prepare_worker.py
Normal file
106
model_manager/worker/prepare_worker.py
Normal 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'],
|
||||
),
|
||||
)
|
||||
@@ -30,7 +30,6 @@ Environment Variables:
|
||||
|
||||
from temporalio import client, workflow
|
||||
from temporalio.runtime import PrometheusConfig, Runtime, TelemetryConfig
|
||||
from temporalio.worker import PollerBehaviorAutoscaling, Worker
|
||||
|
||||
with workflow.unsafe.imports_passed_through():
|
||||
import asyncio
|
||||
@@ -40,9 +39,8 @@ with workflow.unsafe.imports_passed_through():
|
||||
from prometheus_client import start_http_server
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
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.temporal.worker.prepare_worker import prepare_worker
|
||||
from sientia_model.model_repository.plugin_store import PluginStore
|
||||
|
||||
from model_manager import metrics
|
||||
from model_manager.activities.activities import Activities
|
||||
@@ -55,6 +53,7 @@ with workflow.unsafe.imports_passed_through():
|
||||
build_postgres_config,
|
||||
)
|
||||
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.train_model import TrainModel
|
||||
|
||||
@@ -96,7 +95,6 @@ async def main():
|
||||
'pod_id': POD_ID,
|
||||
'runtime': RUNTIME,
|
||||
}
|
||||
|
||||
|
||||
start_prometheus_server(logger, metadata)
|
||||
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'Initializing metrics controller', metadata)
|
||||
logger.custom_info('Initializing metrics controller', metadata)
|
||||
|
||||
metrics_controller = MetricsController(
|
||||
logger=logger
|
||||
)
|
||||
metrics_controller = MetricsController(logger=logger)
|
||||
|
||||
logger.custom_info(f'Installing runtime {RUNTIME}', metadata)
|
||||
|
||||
@@ -185,6 +181,7 @@ async def main():
|
||||
],
|
||||
temporal_client=temporal_client,
|
||||
logger=logger,
|
||||
runtime=RUNTIME,
|
||||
),
|
||||
prepare_worker(
|
||||
main_workflow=CleanupFiles,
|
||||
@@ -194,12 +191,11 @@ async def main():
|
||||
],
|
||||
temporal_client=temporal_client,
|
||||
logger=logger,
|
||||
runtime=RUNTIME,
|
||||
),
|
||||
]
|
||||
|
||||
handlers = [
|
||||
w.run() for w in workers
|
||||
]
|
||||
handlers = [w.run() for w in workers]
|
||||
|
||||
logger.custom_info('Model manager workers initialized', metadata)
|
||||
|
||||
|
||||
@@ -21,7 +21,6 @@ with workflow.unsafe.imports_passed_through():
|
||||
from model_manager.activities.activities import Activities
|
||||
from model_manager.activities.experiment_tracking import UpdateType
|
||||
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).
|
||||
# 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')
|
||||
class TrainModel:
|
||||
"""
|
||||
@@ -133,7 +131,7 @@ class TrainModel:
|
||||
)
|
||||
else:
|
||||
pass
|
||||
except Exception:
|
||||
except Exception: # noqa: BLE001
|
||||
# If cleanup fails after training failed, there is nothing extra to log (DB not committed).
|
||||
if training_succeeded: # pragma: no branch
|
||||
workflow.logger.warning(
|
||||
@@ -179,7 +177,7 @@ class TrainModel:
|
||||
input_data: dict[str, Any],
|
||||
experiment_run_id: int,
|
||||
metadata: dict[str, Any],
|
||||
) -> TrainModelParams:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Validate and convert training parameters from dict to TrainModelParams.
|
||||
|
||||
@@ -193,7 +191,7 @@ class TrainModel:
|
||||
metadata: Workflow execution metadata
|
||||
|
||||
Returns:
|
||||
TrainModelParams: Validated training parameters object
|
||||
dict[str, Any]: Validated training parameters
|
||||
|
||||
Raises:
|
||||
Exception: If validation fails (after updating DB status)
|
||||
@@ -220,7 +218,7 @@ class TrainModel:
|
||||
)
|
||||
|
||||
await self._update_experiment_run(
|
||||
metadata=metadata,
|
||||
metadata=metadata,
|
||||
experiment_run_id=experiment_run_id,
|
||||
update_type=UpdateType.STATUS,
|
||||
status=ExperimentStatus.ORCHESTRATOR_WAITING_PROC,
|
||||
@@ -236,7 +234,7 @@ class TrainModel:
|
||||
status=ExperimentStatus.ORCHESTRATOR_VALIDATION_ERROR,
|
||||
error_message=self._extract_error_message(e),
|
||||
)
|
||||
except Exception as secondary:
|
||||
except Exception as secondary: # noqa: BLE001
|
||||
workflow.logger.warning(
|
||||
'Failed to persist ORCHESTRATOR_VALIDATION_ERROR to experiment_run: %s',
|
||||
secondary,
|
||||
@@ -245,7 +243,7 @@ class TrainModel:
|
||||
|
||||
async def _train_model(
|
||||
self,
|
||||
train_params: TrainModelParams,
|
||||
train_params: dict[str, Any],
|
||||
experiment_run_id: int,
|
||||
metadata: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
@@ -288,7 +286,6 @@ class TrainModel:
|
||||
|
||||
return train_result
|
||||
except Exception as e:
|
||||
|
||||
try:
|
||||
await self._update_experiment_run(
|
||||
metadata=metadata,
|
||||
@@ -297,7 +294,7 @@ class TrainModel:
|
||||
status=ExperimentStatus.TRAINING_ERROR,
|
||||
error_message=self._extract_error_message(e),
|
||||
)
|
||||
except Exception as secondary:
|
||||
except Exception as secondary: # noqa: BLE001
|
||||
workflow.logger.warning(
|
||||
'Failed to persist TRAINING_ERROR status to experiment_run: %s',
|
||||
secondary,
|
||||
|
||||
@@ -56,6 +56,7 @@ ignore = [
|
||||
"S101", # assert allowed in tests
|
||||
"S105", # hardcoded passwords ok in tests
|
||||
"S106", # hardcoded passwords ok in tests
|
||||
"S108", # temp paths are expected in tests
|
||||
]
|
||||
|
||||
[tool.ruff.lint.mccabe]
|
||||
@@ -121,6 +122,13 @@ ignore_missing_imports = true
|
||||
module = "yaml"
|
||||
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]
|
||||
testpaths = ["tests"]
|
||||
python_files = ["test_*.py"]
|
||||
|
||||
@@ -4,7 +4,7 @@ sqlalchemy==2.0.44
|
||||
boto3==1.40.55
|
||||
botocore==1.40.55
|
||||
/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
|
||||
beautifulsoup4==4.12.3
|
||||
evidently
|
||||
@@ -44,7 +44,10 @@ def _minio(endpoint_url: str):
|
||||
)
|
||||
def test_activities_strips_minio_endpoint_scheme(endpoint, expected_endpoint):
|
||||
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.Cleanup.__init__', Mock(return_value=None)),
|
||||
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():
|
||||
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.Cleanup.__init__', Mock(return_value=None)),
|
||||
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():
|
||||
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.Cleanup.__init__', Mock(return_value=None)),
|
||||
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():
|
||||
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.Cleanup.__init__', Mock(return_value=None)),
|
||||
patch('model_manager.activities.activities.SientiaMLflowRepository'),
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
"""Unit tests for the Cleanup activity, ensuring 100% code coverage."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
@@ -126,12 +125,10 @@ def test_cleanup_temp_directories_nonexistent_path(
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup._emit_metrics = AsyncMock()
|
||||
cleanup._emit_metrics = 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._emit_metrics.assert_called_once()
|
||||
@@ -155,7 +152,7 @@ def test_cleanup_temp_directories_success_with_deletions(
|
||||
notification_handler=mock_notification_handler,
|
||||
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_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}')
|
||||
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 os.path.exists(recent_dir)
|
||||
@@ -190,13 +187,13 @@ def test_cleanup_temp_directories_dry_run(
|
||||
notification_handler=mock_notification_handler,
|
||||
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_dir = os.path.join(temp_dir, f'old_dir_{old_time}')
|
||||
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)
|
||||
cleanup._emit_metrics.assert_called_once()
|
||||
@@ -220,7 +217,7 @@ def test_cleanup_temp_directories_delete_error(
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup._emit_metrics = AsyncMock()
|
||||
cleanup._emit_metrics = MagicMock()
|
||||
cleanup.error = MagicMock()
|
||||
|
||||
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)
|
||||
|
||||
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._emit_metrics.assert_called_once()
|
||||
@@ -250,18 +247,16 @@ def test_emit_metrics(
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup.emit_metric = AsyncMock()
|
||||
cleanup.emit_metric_sync = MagicMock()
|
||||
|
||||
asyncio.run(
|
||||
cleanup._emit_metrics(
|
||||
metadata={'pod_id': 'p1', 'workflow_name': 'wf1'},
|
||||
metrics_status='success',
|
||||
activity_name='test_activity',
|
||||
emit_workflow_metric=True,
|
||||
)
|
||||
cleanup._emit_metrics(
|
||||
metadata={'pod_id': 'p1', 'workflow_name': 'wf1'},
|
||||
metrics_status='success',
|
||||
activity_name='test_activity',
|
||||
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(
|
||||
@@ -281,7 +276,7 @@ def test_cleanup_temp_directories_with_files_and_unmatched_dirs(
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup._emit_metrics = AsyncMock()
|
||||
cleanup._emit_metrics = MagicMock()
|
||||
cleanup.debug = MagicMock()
|
||||
|
||||
# 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')
|
||||
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
|
||||
cleanup.debug.assert_called_with(
|
||||
@@ -315,14 +310,14 @@ def test_cleanup_temp_directories_invalid_timestamp_format(
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup._emit_metrics = AsyncMock()
|
||||
cleanup._emit_metrics = MagicMock()
|
||||
cleanup.error = MagicMock()
|
||||
|
||||
# Create a directory with a malformed timestamp that matches the regex but fails parsing
|
||||
malformed_dir_name = 'dir_20239999_999999_999999'
|
||||
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._emit_metrics.assert_called_once()
|
||||
@@ -345,12 +340,12 @@ def test_cleanup_temp_directories_generic_exception(
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup._emit_metrics = AsyncMock()
|
||||
cleanup._emit_metrics = MagicMock()
|
||||
cleanup.send_notification = MagicMock()
|
||||
|
||||
with patch('os.listdir', side_effect=Exception('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._emit_metrics.assert_called_once()
|
||||
@@ -369,15 +364,13 @@ def test_emit_metrics_activity_only(
|
||||
notification_handler=mock_notification_handler,
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
cleanup.emit_metric = AsyncMock()
|
||||
cleanup.emit_metric_sync = MagicMock()
|
||||
|
||||
asyncio.run(
|
||||
cleanup._emit_metrics(
|
||||
metadata={'pod_id': 'p1', 'workflow_name': 'wf1'},
|
||||
metrics_status='success',
|
||||
activity_name='test_activity',
|
||||
emit_workflow_metric=False,
|
||||
)
|
||||
cleanup._emit_metrics(
|
||||
metadata={'pod_id': 'p1', 'workflow_name': 'wf1'},
|
||||
metrics_status='success',
|
||||
activity_name='test_activity',
|
||||
emit_workflow_metric=False,
|
||||
)
|
||||
|
||||
cleanup.emit_metric.assert_called_once()
|
||||
cleanup.emit_metric_sync.assert_called_once()
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
"""Unit tests for ExperimentTracking class with 100% coverage."""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -185,7 +184,7 @@ def test_execute_update_success(
|
||||
mock_engine.begin.return_value.__enter__.return_value = mock_connection
|
||||
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}
|
||||
mock_connection.execute.assert_called_once()
|
||||
@@ -212,7 +211,7 @@ def test_update_experiment_run_status_success(
|
||||
|
||||
mock_execute = MagicMock()
|
||||
|
||||
async def mock_execute_update(*args, **kwargs):
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
mock_execute(*args, **kwargs)
|
||||
return {'rowcount': 1}
|
||||
|
||||
@@ -226,7 +225,7 @@ def test_update_experiment_run_status_success(
|
||||
'status': 'running',
|
||||
}
|
||||
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
mock_execute.assert_called_once()
|
||||
call_args = mock_execute.call_args
|
||||
@@ -255,7 +254,7 @@ def test_update_experiment_run_status_missing_status(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification_async = AsyncMock()
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -264,9 +263,9 @@ def test_update_experiment_run_status_missing_status(
|
||||
}
|
||||
|
||||
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(
|
||||
@@ -290,7 +289,7 @@ def test_update_experiment_run_status_with_error_success(
|
||||
|
||||
mock_execute = MagicMock()
|
||||
|
||||
async def mock_execute_update(*args, **kwargs):
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
mock_execute(*args, **kwargs)
|
||||
return {'rowcount': 1}
|
||||
|
||||
@@ -305,7 +304,7 @@ def test_update_experiment_run_status_with_error_success(
|
||||
'error_message': 'Test error',
|
||||
}
|
||||
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
mock_execute.assert_called_once()
|
||||
call_args = mock_execute.call_args
|
||||
@@ -336,7 +335,7 @@ def test_update_experiment_run_status_with_error_truncate_message(
|
||||
|
||||
mock_execute = MagicMock()
|
||||
|
||||
async def mock_execute_update(*args, **kwargs):
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
mock_execute(*args, **kwargs)
|
||||
return {'rowcount': 1}
|
||||
|
||||
@@ -352,7 +351,7 @@ def test_update_experiment_run_status_with_error_truncate_message(
|
||||
'error_message': long_error,
|
||||
}
|
||||
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
call_args = mock_execute.call_args
|
||||
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,
|
||||
)
|
||||
|
||||
et.send_notification_async = AsyncMock()
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -387,9 +386,9 @@ def test_update_experiment_run_status_with_error_missing_error_message(
|
||||
}
|
||||
|
||||
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(
|
||||
@@ -413,7 +412,7 @@ def test_update_experiment_run_model_saved_success(
|
||||
|
||||
mock_execute = MagicMock()
|
||||
|
||||
async def mock_execute_update(*args, **kwargs):
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
mock_execute(*args, **kwargs)
|
||||
return {'rowcount': 1}
|
||||
|
||||
@@ -428,7 +427,7 @@ def test_update_experiment_run_model_saved_success(
|
||||
'run_name': 'run_001',
|
||||
}
|
||||
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
et.update_experiment_run(input_data)
|
||||
|
||||
mock_execute.assert_called_once()
|
||||
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,
|
||||
)
|
||||
|
||||
et.send_notification_async = AsyncMock()
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -467,9 +466,9 @@ def test_update_experiment_run_model_saved_missing_run_name(
|
||||
}
|
||||
|
||||
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(
|
||||
@@ -491,7 +490,7 @@ def test_update_experiment_run_invalid_update_type(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification_async = AsyncMock()
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -500,9 +499,9 @@ def test_update_experiment_run_invalid_update_type(
|
||||
}
|
||||
|
||||
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(
|
||||
@@ -524,11 +523,11 @@ def test_update_experiment_run_no_rows_updated(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
async def mock_execute_update(*args, **kwargs):
|
||||
def mock_execute_update(*args, **kwargs):
|
||||
return {'rowcount': 0}
|
||||
|
||||
et._execute_update = mock_execute_update
|
||||
et.send_notification_async = AsyncMock()
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -538,9 +537,9 @@ def test_update_experiment_run_no_rows_updated(
|
||||
}
|
||||
|
||||
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(
|
||||
@@ -562,7 +561,7 @@ def test_update_experiment_run_status_with_error_missing_status(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification_async = AsyncMock()
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -572,9 +571,9 @@ def test_update_experiment_run_status_with_error_missing_status(
|
||||
}
|
||||
|
||||
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(
|
||||
@@ -596,7 +595,7 @@ def test_update_experiment_run_model_saved_missing_status(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification_async = AsyncMock()
|
||||
et.send_notification = MagicMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -606,9 +605,9 @@ def test_update_experiment_run_model_saved_missing_status(
|
||||
}
|
||||
|
||||
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(
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Unit tests for Training activities."""
|
||||
|
||||
from contextlib import asynccontextmanager
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from contextlib import contextmanager
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
@@ -49,46 +49,43 @@ def training():
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_model_metadata_success(training):
|
||||
training.plugin_store.get_model_index = MagicMock(return_value={'schemas': {'components': {'schemas': {}}}})
|
||||
def test_load_model_metadata_success(training):
|
||||
training.plugin_store.get_model_index = MagicMock(
|
||||
return_value={'schemas': {'components': {'schemas': {}}}}
|
||||
)
|
||||
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 out['model_metadata']['schemas']
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_model_metadata_notifies_on_error(training):
|
||||
def test_load_model_metadata_notifies_on_error(training):
|
||||
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': {}}
|
||||
with pytest.raises(RuntimeError, match='idx'):
|
||||
await training.load_model_metadata(inp)
|
||||
training.send_notification_async.assert_awaited()
|
||||
training.load_model_metadata(inp)
|
||||
training.send_notification.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_train_params_success(training):
|
||||
def test_validate_train_params_success(training):
|
||||
pdict = _minimal_params_dict()
|
||||
pdict['model_metadata'] = {'schemas': {'components': {'schemas': {}}}}
|
||||
inp = {**pdict, 'metadata': {}}
|
||||
out = await training.validate_train_params(inp)
|
||||
assert isinstance(out, TrainModelParams)
|
||||
assert out.target_variable == 't'
|
||||
out = training.validate_train_params(inp)
|
||||
assert isinstance(out, dict)
|
||||
assert out['target_variable'] == 't'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_train_params_notifies(training):
|
||||
training.send_notification_async = AsyncMock()
|
||||
def test_validate_train_params_notifies(training):
|
||||
training.send_notification = MagicMock()
|
||||
inp = {'metadata': {}, 'experiment_run_id': 1}
|
||||
with pytest.raises(Exception):
|
||||
await training.validate_train_params(inp)
|
||||
training.send_notification_async.assert_awaited()
|
||||
with pytest.raises((KeyError, ValueError, TypeError)):
|
||||
training.validate_train_params(inp)
|
||||
training.send_notification.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_train_model_download_fails_notifies(training):
|
||||
def test_train_model_download_fails_notifies(training):
|
||||
"""train_model notifies and re-raises when MinIO download fails."""
|
||||
tp = TrainModelParams.from_dict(
|
||||
{
|
||||
@@ -96,32 +93,31 @@ async def test_train_model_download_fails_notifies(training):
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
)
|
||||
training.minio_repository.download_file = AsyncMock(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.minio_repository.download_file_sync = MagicMock(side_effect=OSError('minio'))
|
||||
training.send_notification = MagicMock()
|
||||
with pytest.raises(RuntimeError, match='rm'):
|
||||
await training.cleanup_resources({'metadata': {'pod': 'p'}, 'run_dir': '/tmp/x'})
|
||||
with pytest.raises(OSError, match='minio'):
|
||||
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()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@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)."""
|
||||
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]})
|
||||
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.compute_regression_metrics = MagicMock(
|
||||
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]})
|
||||
wrapper.predict = MagicMock(side_effect=[(pred_train, None), (pred_val, None)])
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = AsyncMock(return_value=wrapper)
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@asynccontextmanager
|
||||
async def _run_ctx(*_a, **_k):
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_name = 'run-n'
|
||||
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
|
||||
|
||||
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_id'] == 'run-i'
|
||||
assert out['run_dir'] == '/tmp/run'
|
||||
mock_mlflow.log_artifact.assert_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@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."""
|
||||
d = {
|
||||
**_minimal_params_dict(),
|
||||
@@ -190,9 +185,12 @@ async def test_train_model_train_params_as_dict(mock_mlflow, training):
|
||||
tp = TrainModelParams.from_dict(d)
|
||||
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.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):
|
||||
x.report_path = '/tmp/report.html'
|
||||
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)]
|
||||
)
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = AsyncMock(return_value=wrapper)
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@asynccontextmanager
|
||||
async def _run_ctx(*_a, **_k):
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_name = 'n'
|
||||
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
|
||||
|
||||
await training.train_model({'metadata': {}, 'train_params': d})
|
||||
training.train_model({'metadata': {}, 'train_params': d})
|
||||
mock_mlflow.log_artifact.assert_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@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)."""
|
||||
d = {
|
||||
**_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]})
|
||||
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:
|
||||
return b'train'
|
||||
if object_name == 'val.csv':
|
||||
return b'val'
|
||||
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.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):
|
||||
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)]
|
||||
)
|
||||
wrapper.store_model = MagicMock()
|
||||
training.plugin_store.get_model = AsyncMock(return_value=wrapper)
|
||||
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
||||
|
||||
@asynccontextmanager
|
||||
async def _run_ctx(*_a, **_k):
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_name = 'n'
|
||||
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
|
||||
|
||||
await training.train_model({'metadata': {}, 'train_params': tp})
|
||||
assert training.minio_repository.download_file.await_count == 2
|
||||
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
||||
assert training.minio_repository.download_file_sync.call_count == 2
|
||||
mock_mlflow.log_artifact.assert_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_train_model_value_error_when_paths_missing_after_report(training):
|
||||
def test_train_model_value_error_when_paths_missing_after_report(training):
|
||||
"""Raises ValueError when report paths are not populated after generate_report."""
|
||||
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]})
|
||||
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.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)
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
||||
wrapper.predict = MagicMock(
|
||||
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
|
||||
async def _run_ctx(*_a, **_k):
|
||||
@contextmanager
|
||||
def _run_ctx(*_a, **_k):
|
||||
info = MagicMock()
|
||||
info.run_name = 'n'
|
||||
info.run_id = 'i'
|
||||
yield info
|
||||
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
training.send_notification_async = AsyncMock()
|
||||
training.send_notification = MagicMock()
|
||||
|
||||
with pytest.raises(ValueError, match='Report path'):
|
||||
await training.train_model({'metadata': {}, 'train_params': tp})
|
||||
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
||||
|
||||
@@ -162,7 +162,11 @@ def test_validate_model_param_schema_validation_error(valid_train_params_dict):
|
||||
'schemas': {
|
||||
'components': {
|
||||
'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)
|
||||
with patch('model_manager.utils.models.train_model_params.Draft202012Validator') as m:
|
||||
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()
|
||||
|
||||
|
||||
@@ -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):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {'components': {'schemas': {'model': {'type': 'object'}}}}
|
||||
}
|
||||
d['model_metadata'] = {'schemas': {'components': {'schemas': {'model': {'type': 'object'}}}}}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.model_kwargs = {}
|
||||
p.validate_business_rules()
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import numpy as np
|
||||
@@ -33,7 +34,7 @@ def test_train_test_split_ndarray():
|
||||
|
||||
|
||||
def _params(**kwargs) -> TrainModelParams:
|
||||
base = {
|
||||
base: dict[str, Any] = {
|
||||
'variable_columns': ['v1'],
|
||||
'target_variable': 't',
|
||||
'bucket_name': 'b',
|
||||
@@ -276,7 +277,13 @@ def test_configure_datetime_index_already_datetime_index():
|
||||
def test_configure_datetime_index_from_common_column():
|
||||
repo = dmr.DataManagerRepository(MagicMock())
|
||||
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, {})
|
||||
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():
|
||||
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'):
|
||||
repo._create_run_directory('/tmp', 'run', {})
|
||||
|
||||
|
||||
def test_create_run_directory_os_error():
|
||||
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'):
|
||||
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('model_manager.utils.repository.data_manager_repository.Reports') as mrep,
|
||||
):
|
||||
inst = mrep.return_value
|
||||
instance = mrep.return_value
|
||||
instance.save_all_sections_html = Mock()
|
||||
out = repo.generate_report(tmr, {})
|
||||
|
||||
76
tests/worker/test_prepare_worker.py
Normal file
76
tests/worker/test_prepare_worker.py
Normal 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)
|
||||
@@ -125,7 +125,7 @@ def test_start_prometheus_server_success(
|
||||
mock_app_up = Mock()
|
||||
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)
|
||||
|
||||
@@ -148,7 +148,7 @@ def test_start_prometheus_server_custom_port(mock_metrics, mock_start_http_serve
|
||||
mock_app_up = Mock()
|
||||
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)
|
||||
|
||||
@@ -166,7 +166,7 @@ def test_start_prometheus_server_failure(
|
||||
|
||||
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)
|
||||
|
||||
@@ -252,9 +252,7 @@ async def test_main_successful_startup(
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(
|
||||
return_value={'plugins': [], 'interceptors': []}
|
||||
)
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
@@ -347,9 +345,7 @@ async def test_main_handles_exception(
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(
|
||||
return_value={'plugins': [], 'interceptors': []}
|
||||
)
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
@@ -486,9 +482,7 @@ async def test_main_temporal_client_configuration(
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(
|
||||
return_value={'plugins': [], 'interceptors': []}
|
||||
)
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
@@ -581,9 +575,7 @@ async def test_main_worker_configuration(
|
||||
mock_runtime_class.return_value = mock_runtime
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(
|
||||
return_value={'plugins': [], 'interceptors': []}
|
||||
)
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
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_client_instance = AsyncMock()
|
||||
mock_client_instance.config = Mock(
|
||||
return_value={'plugins': [], 'interceptors': []}
|
||||
)
|
||||
mock_client_instance.config = Mock(return_value={'plugins': [], 'interceptors': []})
|
||||
mock_client_class.connect = AsyncMock(return_value=mock_client_instance)
|
||||
|
||||
mock_worker_instance = Mock()
|
||||
@@ -844,7 +834,7 @@ def test_start_prometheus_server_prints_success(
|
||||
mock_app_up = Mock()
|
||||
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)
|
||||
|
||||
@@ -863,7 +853,7 @@ def test_start_prometheus_server_prints_failure(
|
||||
|
||||
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)
|
||||
|
||||
|
||||
@@ -243,7 +243,9 @@ async def test_run_validation_error(mock_wf, sample_input_data):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@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."""
|
||||
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
|
||||
@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."""
|
||||
from model_manager.workflows.train_model import TrainModel
|
||||
|
||||
@@ -297,7 +301,6 @@ def test_module_constants():
|
||||
from model_manager.workflows.train_model import (
|
||||
TIMEOUT_DELETE_FILE,
|
||||
TIMEOUT_TRAIN_MODEL,
|
||||
TIMEOUT_UPDATE_DATABASE,
|
||||
TIMEOUT_VALIDATE_PARAMS,
|
||||
database_retry_policy,
|
||||
network_retry_policy,
|
||||
|
||||
Reference in New Issue
Block a user