SIENTIAPDE-1241: Refactored method into smaller ones
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user