SIENTIAPDE-1241: Refactored method into smaller ones

This commit is contained in:
Kou-Kinoshita
2025-10-29 16:45:30 -03:00
parent 63fa264ec4
commit 26fca36d82
2 changed files with 401 additions and 62 deletions

View File

@@ -125,6 +125,99 @@ class ExperimentTracking(Postgres):
return await asyncio.to_thread(_run)
def _build_status_update_query(
self, status: str | None, experiment_run_id: int
) -> tuple[str, dict[str, Any]]:
"""Build SQL query for simple status update."""
if not isinstance(status, str) or not status:
raise ValueError('status is required for STATUS update type')
sql_query = """
UPDATE experiment_run
SET status = :status, updated_at = :updated_at
WHERE id = :experiment_run_id
"""
query_params = {
'status': status,
'updated_at': datetime.now(UTC),
'experiment_run_id': experiment_run_id,
}
return sql_query, query_params
def _build_status_with_error_query(
self, status: str | None, error_message: str | None, experiment_run_id: int
) -> tuple[str, dict[str, Any]]:
"""Build SQL query for status update with error message."""
if not isinstance(status, str) or not status:
raise ValueError('status is required for STATUS_WITH_ERROR update type')
if not isinstance(error_message, str) or not error_message:
raise ValueError('error_message is required for STATUS_WITH_ERROR update type')
# Truncate error message if too long
truncated_error = error_message[:1024] if len(error_message) > 1024 else error_message
sql_query = """
UPDATE experiment_run
SET status = :status, error_message = :error_message, updated_at = :updated_at
WHERE id = :experiment_run_id
"""
query_params = {
'status': status,
'error_message': truncated_error,
'updated_at': datetime.now(UTC),
'experiment_run_id': experiment_run_id,
}
return sql_query, query_params
def _build_model_saved_query(
self, run_name: str | None, status: str | None, experiment_run_id: int
) -> tuple[str, dict[str, Any]]:
"""Build SQL query for model saved update."""
if not isinstance(run_name, str) or not run_name:
raise ValueError('run_name is required for MODEL_SAVED update type')
if not isinstance(status, str) or not status:
raise ValueError('status is required for MODEL_SAVED update type')
sql_query = """
UPDATE experiment_run
SET run_name = :run_name, status = :status, updated_at = :updated_at
WHERE id = :experiment_run_id
"""
query_params = {
'run_name': run_name,
'status': status,
'updated_at': datetime.now(UTC),
'experiment_run_id': experiment_run_id,
}
return sql_query, query_params
def _get_update_query_and_params(
self, update_type: str, experiment_run_id: int, input_data: dict[str, Any]
) -> tuple[str, dict[str, Any]]:
"""Get SQL query and parameters based on update type."""
status = input_data.get('status')
error_message = input_data.get('error_message')
run_name = input_data.get('run_name')
if update_type == UpdateType.STATUS:
return self._build_status_update_query(status, experiment_run_id)
if update_type == UpdateType.STATUS_WITH_ERROR:
return self._build_status_with_error_query(status, error_message, experiment_run_id)
if update_type == UpdateType.MODEL_SAVED:
return self._build_model_saved_query(run_name, status, experiment_run_id)
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:
"""
@@ -153,70 +246,11 @@ class ExperimentTracking(Postgres):
experiment_run_id = input_data['experiment_run_id']
update_type = input_data['update_type']
status = input_data.get('status')
error_message = input_data.get('error_message')
run_name = input_data.get('run_name')
try:
query_params: dict[str, Any]
if update_type == UpdateType.STATUS:
if not isinstance(status, str) or not status:
raise ValueError('status is required for STATUS update type')
sql_query = """
UPDATE experiment_run
SET status = :status, updated_at = :updated_at
WHERE id = :experiment_run_id
"""
query_params = {
'status': status,
'updated_at': datetime.now(UTC),
'experiment_run_id': experiment_run_id,
}
elif update_type == UpdateType.STATUS_WITH_ERROR:
if not isinstance(status, str) or not status:
raise ValueError('status is required for STATUS_WITH_ERROR update type')
if not isinstance(error_message, str) or not error_message:
raise ValueError('error_message is required for STATUS_WITH_ERROR update type')
if len(error_message) > 1024:
error_message = error_message[:1024]
sql_query = """
UPDATE experiment_run
SET status = :status, error_message = :error_message, updated_at = :updated_at
WHERE id = :experiment_run_id
"""
query_params = {
'status': status,
'error_message': error_message,
'updated_at': datetime.now(UTC),
'experiment_run_id': experiment_run_id,
}
elif update_type == UpdateType.MODEL_SAVED:
if not isinstance(run_name, str) or not run_name:
raise ValueError('run_name is required for MODEL_SAVED update type')
if not isinstance(status, str) or not status:
raise ValueError('status is required for MODEL_SAVED update type')
sql_query = """
UPDATE experiment_run
SET run_name = :run_name, status = :status, updated_at = :updated_at
WHERE id = :experiment_run_id
"""
query_params = {
'run_name': run_name,
'status': status,
'updated_at': datetime.now(UTC),
'experiment_run_id': experiment_run_id,
}
else:
raise ValueError(f'Invalid update_type: {update_type}')
sql_query, query_params = self._get_update_query_and_params(
update_type, experiment_run_id, input_data
)
result = await self._execute_update(sql_query, query_params)