feat: enhance training and experiment tracking functionality

- Updated `Activities` class to improve garbage collection handling.
- Enhanced error messaging in `ExperimentTracking` for better clarity on update failures.
- Refactored `Training` class to streamline exception handling and improve type hints.
- Introduced new methods in `TrainModelParams` for better handling of experiment run IDs and model metadata.
- Added functionality to extract model equations in `DataManagerRepository` for linear regression models.
This commit is contained in:
vitor-aignosi
2026-04-06 15:05:57 -03:00
parent 1352d1ac8f
commit 6b1df7c3a7
22 changed files with 1751 additions and 2085 deletions

View File

@@ -1,7 +1,7 @@
"""Unit tests for ExperimentTracking class with 100% coverage."""
import asyncio
from unittest.mock import MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@@ -22,12 +22,8 @@ def mock_notification_handler():
def mock_metrics_controller():
"""Create a mock metrics controller."""
controller = MagicMock()
# Make shutdown an async coroutine
async def mock_shutdown():
pass
controller.shutdown = mock_shutdown
controller.shutdown = AsyncMock()
controller.emit = AsyncMock()
return controller
@@ -259,7 +255,7 @@ def test_update_experiment_run_status_missing_status(
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
et.send_notification_async = AsyncMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
@@ -270,7 +266,7 @@ def test_update_experiment_run_status_missing_status(
with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data))
et.send_notification.assert_called_once()
et.send_notification_async.assert_awaited_once()
def test_update_experiment_run_status_with_error_success(
@@ -381,7 +377,7 @@ def test_update_experiment_run_status_with_error_missing_error_message(
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
et.send_notification_async = AsyncMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
@@ -393,7 +389,7 @@ def test_update_experiment_run_status_with_error_missing_error_message(
with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data))
et.send_notification.assert_called_once()
et.send_notification_async.assert_awaited_once()
def test_update_experiment_run_model_saved_success(
@@ -461,7 +457,7 @@ def test_update_experiment_run_model_saved_missing_run_name(
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
et.send_notification_async = AsyncMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
@@ -473,7 +469,7 @@ def test_update_experiment_run_model_saved_missing_run_name(
with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data))
et.send_notification.assert_called_once()
et.send_notification_async.assert_awaited_once()
def test_update_experiment_run_invalid_update_type(
@@ -495,7 +491,7 @@ def test_update_experiment_run_invalid_update_type(
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
et.send_notification_async = AsyncMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
@@ -506,7 +502,7 @@ def test_update_experiment_run_invalid_update_type(
with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data))
et.send_notification.assert_called_once()
et.send_notification_async.assert_awaited_once()
def test_update_experiment_run_no_rows_updated(
@@ -532,7 +528,7 @@ def test_update_experiment_run_no_rows_updated(
return {'rowcount': 0}
et._execute_update = mock_execute_update
et.send_notification = MagicMock()
et.send_notification_async = AsyncMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
@@ -544,7 +540,7 @@ def test_update_experiment_run_no_rows_updated(
with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data))
et.send_notification.assert_called_once()
et.send_notification_async.assert_awaited_once()
def test_update_experiment_run_status_with_error_missing_status(
@@ -566,7 +562,7 @@ def test_update_experiment_run_status_with_error_missing_status(
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
et.send_notification_async = AsyncMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
@@ -578,7 +574,7 @@ def test_update_experiment_run_status_with_error_missing_status(
with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data))
et.send_notification.assert_called_once()
et.send_notification_async.assert_awaited_once()
def test_update_experiment_run_model_saved_missing_status(
@@ -600,7 +596,7 @@ def test_update_experiment_run_model_saved_missing_status(
metrics_controller=mock_metrics_controller,
)
et.send_notification = MagicMock()
et.send_notification_async = AsyncMock()
input_data = {
'metadata': {'workflow_id': 'test-123'},
@@ -612,7 +608,7 @@ def test_update_experiment_run_model_saved_missing_status(
with pytest.raises(RuntimeError):
asyncio.run(et.update_experiment_run(input_data))
et.send_notification.assert_called_once()
et.send_notification_async.assert_awaited_once()
def test_experiment_tracking_del_with_engine_no_super_del(