diff --git a/src/splunk_ao/experiment.py b/src/splunk_ao/experiment.py index 3f579818..5e10aa7a 100644 --- a/src/splunk_ao/experiment.py +++ b/src/splunk_ao/experiment.py @@ -3,9 +3,13 @@ import builtins import datetime import re +import warnings from collections.abc import Iterator +from time import sleep from typing import TYPE_CHECKING, Any +from tqdm.auto import tqdm + from splunk_ao.config import SplunkAOConfig from splunk_ao.datasets import Dataset as LegacyDataset from splunk_ao.exceptions import NotFoundError @@ -13,7 +17,6 @@ from splunk_ao.experiments import Experiments as ExperimentsService from splunk_ao.experiments import _default_prompt_settings from splunk_ao.export import ExportClient -from splunk_ao.job_progress import get_run_scorer_jobs, job_progress from splunk_ao.prompts import PromptTemplate, get_prompt from splunk_ao.resources.api.experiment import ( delete_experiment_projects_project_id_experiments_experiment_id_delete, @@ -185,7 +188,6 @@ class Experiment(StateManagementMixin): _prompt_template: PromptTemplate | None _model_obj: Model | None _experiment_response: ExperimentResponse | None - _job_id: str | None def __str__(self) -> str: """String representation of the experiment.""" @@ -366,7 +368,6 @@ def __init__( # Private runtime state self._experiment_response: ExperimentResponse | None = None - self._job_id: str | None = None self._run_result: ExperimentRunResult | None = None self._run_result_consumed: bool = False @@ -535,7 +536,6 @@ def _create_empty(cls) -> Experiment: instance._prompt_template = None instance._model_obj = None instance._experiment_response: ExperimentResponse | None = None - instance._job_id: str | None = None instance._run_result: ExperimentRunResult | None = None instance._run_result_consumed: bool = False return instance @@ -627,7 +627,6 @@ def _from_api_response(cls, retrieved_experiment: ExperimentResponse) -> Experim instance._prompt_template = None instance._model_obj = None instance._experiment_response = retrieved_experiment - instance._job_id = None # Set state to synced since we just retrieved from API instance._set_state(SyncState.SYNCED) return instance @@ -1122,54 +1121,99 @@ def get_status(self) -> ExperimentStatusInfo: return ExperimentStatusInfo(self._experiment_response) - def monitor_progress(self, job_id: str | None = None) -> str: + def monitor_progress( + self, + poll_interval_seconds: float = 2.0, + *, + timeout_seconds: float | None = 3600.0, + job_id: str | None = None, + ) -> None: """ - Monitor the progress of the experiment job with a progress bar. + Monitor the progress of the experiment with a progress bar. - Args: - job_id: Optional job ID to monitor. If not provided, will attempt to find - the primary job for this experiment. + Polls the experiment status via the API until the experiment completes, + displaying a tqdm progress bar reflecting `log_generation` progress. + + Parameters + ---------- + poll_interval_seconds : float, optional + Seconds to wait between status polls. Defaults to 2.0. + Note: in a prior version, ``job_id`` was the first positional + parameter. That parameter has been removed; callers that passed a + job ID string positionally will now receive a ``TypeError`` from + ``sleep()``. Use ``job_id=`` as a keyword argument instead. + timeout_seconds : float or None, optional + Maximum seconds to wait before raising TimeoutError. Defaults to 3600.0 + (one hour). Pass None to wait indefinitely (not recommended). + job_id : str or None, optional + Deprecated. This parameter is ignored; it existed in a prior version + that polled the jobs table, which has been retired. Returns ------- - str: The unique identifier of the completed job. + None Raises ------ - ValueError: If the experiment lacks required id or project_id attributes, - or if no job_id is provided and no job can be found. + ValueError + If the experiment lacks required id or project_id attributes. + RuntimeError + If the experiment enters a failed state. + TimeoutError + If the experiment does not complete within timeout_seconds. Examples -------- - experiment = Experiment.get(name="ml-evaluation", project_name="My AI Project") - result = experiment.run() + experiment = Experiment( + name="ml-evaluation", + dataset_name="ml-dataset", + project_name="My AI Project" + ).create() - # Monitor the job progress - completed_job_id = experiment.monitor_progress() + experiment.monitor_progress() """ + if job_id is not None: + warnings.warn( + "The 'job_id' parameter of monitor_progress() is deprecated and will be removed in a future release. " + "Progress is now tracked directly via experiment status; the job_id value is ignored.", + DeprecationWarning, + stacklevel=2, + ) + if self.id is None: raise ValueError("Experiment ID is not set. Cannot monitor progress for a local-only experiment.") if self.project_id is None: raise ValueError("Project ID is not set. Cannot monitor progress without project_id.") - if job_id is None: - # Try to get job from stored state or query for it - if self._job_id: - job_id = self._job_id - else: - # Get the first scorer job - scorer_jobs = get_run_scorer_jobs(project_id=self.project_id, run_id=self.id) - if not scorer_jobs: - raise ValueError("No job found for this experiment. Run the experiment first.") - job_id = str(scorer_jobs[0].id) + _logger.info(f"Experiment.monitor_progress: experiment_id='{self.id}' - started") + + import time - _logger.info(f"Experiment.monitor_progress: experiment_id='{self.id}' job_id='{job_id}' - started") + deadline = time.monotonic() + timeout_seconds if timeout_seconds is not None else None - # Monitor job progress with progress bar - completed_job_id = job_progress(job_id=job_id, project_id=self.project_id, run_id=self.id) + status = self.get_status() + progress_bar = tqdm(total=100, unit="%", desc="Experiment progress") + try: + while not status.is_complete: + if status.is_failed: + raise RuntimeError( + f"Experiment '{self.id}' entered a failed state. " + "Check the experiment results for details." + ) + if deadline is not None and time.monotonic() >= deadline: + raise TimeoutError( + f"Experiment '{self.id}' did not complete within {timeout_seconds}s. " + "Increase timeout_seconds or pass None to wait indefinitely." + ) + new_progress = status.overall_progress + progress_bar.update(new_progress - progress_bar.n) + sleep(poll_interval_seconds) + status = self.get_status() + progress_bar.update(100 - progress_bar.n) + finally: + progress_bar.close() _logger.info(f"Experiment.monitor_progress: experiment_id='{self.id}' - completed") - return str(completed_job_id) # Query and export methods - similar to LogStream diff --git a/src/splunk_ao/job_progress.py b/src/splunk_ao/job_progress.py deleted file mode 100644 index dc60279c..00000000 --- a/src/splunk_ao/job_progress.py +++ /dev/null @@ -1,119 +0,0 @@ -import contextlib -import random -from time import sleep - -from pydantic import UUID4 -from tqdm.auto import tqdm - -from galileo_core.constants.job import JobName, JobStatus -from galileo_core.constants.scorers import Scorers -from splunk_ao.config import SplunkAOConfig -from splunk_ao.resources.api.jobs import ( - get_job_jobs_job_id_get, - get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get, -) -from splunk_ao.resources.models import HTTPValidationError, JobDB -from splunk_ao.utils.log_config import get_logger - -_logger = get_logger(__name__) - - -def get_job(job_id: str) -> JobDB: - config = SplunkAOConfig.get() - - response = get_job_jobs_job_id_get.sync(client=config.api_client, job_id=str(job_id)) - - if isinstance(response, HTTPValidationError): - raise ValueError(response.detail) - if not response: - raise ValueError(f"Failed to get job status for job {job_id}") - return response - - -def get_run_scorer_jobs(project_id: str, run_id: str) -> list[JobDB]: - config = SplunkAOConfig.get() - - response = get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get.sync( - client=config.api_client, project_id=str(project_id), run_id=str(run_id) - ) - if isinstance(response, HTTPValidationError): - raise ValueError(response.detail) - if response is None: - raise ValueError(f"Failed to get scorer jobs for project {project_id}, run {run_id}") - - _logger.debug(f"Scorer jobs: {response}") - - return [job for job in response if job.job_name == JobName.log_stream_scorer] - - -def scorer_jobs_status(project_id: str, run_id: str) -> None: - """Gets the status of all scorer jobs for a given project and run. - - Parameters - ---------- - project_id - The unique identifier of the project. - run_id - The unique identifier of the run. - """ - scorer_jobs = get_run_scorer_jobs(project_id, run_id) - for job in scorer_jobs: - scorer_name = None - if "prompt_scorer_settings" in job.request_data: - scorer_name = job.request_data["prompt_scorer_settings"]["scorer_name"] - elif "scorer_config" in job.request_data: - scorer_name = job.request_data["scorer_config"]["name"] - - if not scorer_name: - _logger.debug(f"Scorer job {job.id} has no scorer name.") - continue - - with contextlib.suppress(ValueError): - scorer_name = Scorers(scorer_name).name - - _logger.debug(f"Scorer job {job.id} has scorer {scorer_name}.") - - if JobStatus.is_incomplete(job.status): - _logger.info(f"{scorer_name.lstrip('_')}: Computing 🚧") - elif JobStatus.is_failed(job.status): - _logger.info(f"{scorer_name.lstrip('_')}: Failed ❌, error was: {job.error_message}") - else: - _logger.info(f"{scorer_name.lstrip('_')}: Done ✅") - - -def job_progress(job_id: str, project_id: str, run_id: str) -> UUID4: - """Monitors the progress of a job and displays a progress bar. - - Parameters - ---------- - job_id - The unique identifier of the job to monitor. - project_id - The unique identifier of the project. - run_id - The unique identifier of the run. - - Returns - ------- - The unique identifier of the completed job. - """ - job_status = get_job(job_id) - backoff = random.random() - - if JobStatus.is_incomplete(job_status.status): - job_progress_bar = tqdm(total=job_status.steps_total, position=0, leave=True, desc=job_status.progress_message) - while JobStatus.is_incomplete(job_status.status): - sleep(backoff) - job_status = get_job(job_id) - job_progress_bar.set_description(job_status.progress_message) - job_progress_bar.update(job_status.steps_completed - job_progress_bar.n) - backoff = random.random() - job_progress_bar.close() - - _logger.debug(f"Job {job_id} status: {job_status.status}.") - if JobStatus.is_failed(job_status.status): - raise ValueError(f"Job failed with error message {job_status.error_message}.") from None - - _logger.info("Initial job complete, executing scorers asynchronously. Current status:") - scorer_jobs_status(project_id=project_id, run_id=run_id) - return job_status.id diff --git a/src/splunk_ao/jobs.py b/src/splunk_ao/jobs.py deleted file mode 100644 index f001c237..00000000 --- a/src/splunk_ao/jobs.py +++ /dev/null @@ -1,55 +0,0 @@ -import logging - -from splunk_ao.config import SplunkAOConfig -from splunk_ao.resources.api.jobs import create_job_jobs_post -from splunk_ao.resources.models import ( - CreateJobRequest, - CreateJobResponse, - HTTPValidationError, - PromptRunSettings, - ScorerConfig, - TaskType, -) -from splunk_ao.utils.exceptions import _format_http_validation_error - -_logger = logging.getLogger(__name__) - - -class Jobs: - config: SplunkAOConfig - - def __init__(self) -> None: - self.config = SplunkAOConfig.get() - - def create( - self, - project_id: str, - name: str, - run_id: str, - dataset_id: str, - prompt_template_id: str | None, - task_type: TaskType, - scorers: list[ScorerConfig] | None, - prompt_settings: PromptRunSettings | None, - ) -> CreateJobResponse: - create_params: dict = { - "project_id": project_id, - "dataset_id": dataset_id, - "job_name": name, - "run_id": run_id, - "task_type": task_type, - "scorers": scorers, - } - if prompt_template_id is not None: - create_params["prompt_template_version_id"] = prompt_template_id - if prompt_settings is not None: - create_params["prompt_settings"] = prompt_settings - _logger.info(f"create job: {create_params}") - result = create_job_jobs_post.sync_detailed( - client=self.config.api_client, body=CreateJobRequest(**create_params) - ) - if not result.parsed or not isinstance(result.parsed, CreateJobResponse): - if isinstance(result.parsed, HTTPValidationError): - raise ValueError(_format_http_validation_error(result.parsed)) - raise ValueError(f"Create job failed (HTTP {result.status_code}): {result.content.decode(errors='ignore')}") - return result.parsed diff --git a/tests/test_experiment_progress.py b/tests/test_experiment_progress.py new file mode 100644 index 00000000..b956f786 --- /dev/null +++ b/tests/test_experiment_progress.py @@ -0,0 +1,125 @@ +from unittest.mock import MagicMock, patch +from uuid import uuid4 + +import pytest + +from splunk_ao.experiment import Experiment +from splunk_ao.shared.base import SyncState + +FIXED_PROJECT_ID = str(uuid4()) +FIXED_EXPERIMENT_ID = str(uuid4()) + + +def _make_status(progress_percent: float, *, failed: bool = False) -> MagicMock: + """Build a status mock with the given progress, is_complete, and is_failed.""" + status = MagicMock() + status.overall_progress = progress_percent + status.is_complete = progress_percent >= 100.0 + status.is_failed = failed + return status + + +def _make_experiment() -> Experiment: + exp = Experiment._create_empty() + exp.id = FIXED_EXPERIMENT_ID + exp.project_id = FIXED_PROJECT_ID + exp.name = "test-experiment" + exp._set_state(SyncState.SYNCED) + return exp + + +class TestMonitorProgress: + @patch("splunk_ao.experiment.Experiment.get_status") + @patch("splunk_ao.experiment.sleep", return_value=None) + def test_completes_when_status_reaches_100(self, mock_sleep, mock_get_status): + # Given: an experiment that progresses through 0%, 50%, then 100% + mock_get_status.side_effect = [_make_status(0.0), _make_status(50.0), _make_status(100.0)] + exp = _make_experiment() + + # When: monitoring progress until completion + exp.monitor_progress(poll_interval_seconds=0.0) + + # Then: get_status is polled until 100% is reached + assert mock_get_status.call_count == 3 + + @patch("splunk_ao.experiment.Experiment.get_status") + @patch("splunk_ao.experiment.sleep", return_value=None) + def test_already_complete_on_first_poll(self, mock_sleep, mock_get_status): + # Given: an experiment that is already at 100% on the first poll + mock_get_status.return_value = _make_status(100.0) + exp = _make_experiment() + + # When: monitoring progress + exp.monitor_progress(poll_interval_seconds=0.0) + + # Then: get_status is called once and sleep is never called + assert mock_get_status.call_count == 1 + mock_sleep.assert_not_called() + + @patch("splunk_ao.experiment.Experiment.get_status") + @patch("splunk_ao.experiment.sleep", return_value=None) + def test_uses_poll_interval_seconds(self, mock_sleep, mock_get_status): + # Given: an experiment that completes on the second poll + mock_get_status.side_effect = [_make_status(0.0), _make_status(100.0)] + exp = _make_experiment() + + # When: monitoring with a custom poll interval + exp.monitor_progress(poll_interval_seconds=5.0) + + # Then: sleep is called once with the specified interval + mock_sleep.assert_called_once_with(5.0) + + def test_raises_without_experiment_id(self): + # Given: an experiment without an id + exp = Experiment._create_empty() + exp.id = None + exp.project_id = FIXED_PROJECT_ID + + # When/Then: monitoring raises ValueError about the missing experiment id + with pytest.raises(ValueError, match="Experiment ID is not set"): + exp.monitor_progress() + + def test_raises_without_project_id(self): + # Given: an experiment without a project_id + exp = Experiment._create_empty() + exp.id = FIXED_EXPERIMENT_ID + exp.project_id = None + + # When/Then: monitoring raises ValueError about the missing project id + with pytest.raises(ValueError, match="Project ID is not set"): + exp.monitor_progress() + + @patch("splunk_ao.experiment.Experiment.get_status") + @patch("splunk_ao.experiment.sleep", return_value=None) + def test_deprecated_job_id_warns(self, mock_sleep, mock_get_status): + # Given: an experiment that is already complete, and a caller passing the deprecated job_id + mock_get_status.return_value = _make_status(100.0) + exp = _make_experiment() + + # When/Then: monitor_progress emits a DeprecationWarning when job_id is supplied + with pytest.warns(DeprecationWarning, match="job_id"): + exp.monitor_progress(job_id="some-old-job-id") + + @patch("splunk_ao.experiment.Experiment.get_status") + @patch("splunk_ao.experiment.sleep", return_value=None) + def test_raises_runtime_error_on_failed_experiment(self, mock_sleep, mock_get_status): + # Given: an experiment that enters a failed state on the second poll + mock_get_status.side_effect = [_make_status(30.0), _make_status(30.0, failed=True)] + exp = _make_experiment() + + # When/Then: monitor_progress raises RuntimeError instead of polling forever + with pytest.raises(RuntimeError, match="failed state"): + exp.monitor_progress(poll_interval_seconds=0.0) + + @patch("splunk_ao.experiment.Experiment.get_status") + @patch("splunk_ao.experiment.sleep", return_value=None) + def test_raises_timeout_error_when_deadline_exceeded(self, mock_sleep, mock_get_status): + # Given: an experiment stuck at 50% forever with a very short timeout + mock_get_status.return_value = _make_status(50.0) + exp = _make_experiment() + + # When/Then: monitor_progress raises TimeoutError + # Use a tiny timeout and a non-zero interval so the real clock eventually trips it + with pytest.raises(TimeoutError, match="did not complete within"): + exp.monitor_progress(poll_interval_seconds=0.001, timeout_seconds=0.0) + diff --git a/tests/test_experiments.py b/tests/test_experiments.py index 05729f81..2e248d0d 100644 --- a/tests/test_experiments.py +++ b/tests/test_experiments.py @@ -12,7 +12,6 @@ from time_machine import travel import splunk_ao.experiments -import splunk_ao.jobs import splunk_ao.utils.datasets from galileo_core.schemas.logging.span import Span, StepWithChildSpans from galileo_core.schemas.shared.metric import MetricValueType @@ -543,20 +542,12 @@ def test_load_dataset_and_records_error(self) -> None: assert str(exc_info.value) == "To load dataset records, dataset, dataset_name, or dataset_id must be provided" @patch.object(splunk_ao.datasets.Datasets, "get") - @patch.object(splunk_ao.jobs.Jobs, "create") @patch.object(splunk_ao.experiments.Experiments, "create", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Experiments, "get", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Projects, "get_with_env_fallbacks", return_value=project()) def test_run_experiment_with_project_name_loads_project( - self, - mock_get_project: Mock, - mock_get_experiment: Mock, - mock_create_job: Mock, - mock_get_dataset: Mock, - dataset_content: DatasetContent, + self, mock_get_project: Mock, mock_get_experiment: Mock, mock_get_dataset: Mock, dataset_content: DatasetContent ) -> None: - mock_create_job.return_value = MagicMock() - dataset_id = str(UUID(int=0)) run_experiment( "test_experiment", project="awesome-new-project", dataset_id=dataset_id, prompt_template=prompt_template() @@ -565,20 +556,12 @@ def test_run_experiment_with_project_name_loads_project( mock_get_project.assert_called_once_with(id=None, name="awesome-new-project") @patch.object(splunk_ao.datasets.Datasets, "get") - @patch.object(splunk_ao.jobs.Jobs, "create") @patch.object(splunk_ao.experiments.Experiments, "create", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Experiments, "get", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Projects, "get_with_env_fallbacks", return_value=project()) def test_run_experiment_with_project_id_loads_project( - self, - mock_get_project: Mock, - mock_get_experiment: Mock, - mock_create_job: Mock, - mock_get_dataset: Mock, - dataset_content: DatasetContent, + self, mock_get_project: Mock, mock_get_experiment: Mock, mock_get_dataset: Mock, dataset_content: DatasetContent ) -> None: - mock_create_job.return_value = MagicMock() - dataset_id = str(UUID(int=0)) run_experiment( "test_experiment", @@ -590,20 +573,12 @@ def test_run_experiment_with_project_id_loads_project( mock_get_project.assert_called_once_with(id="awesome-new-project", name=None) @patch.object(splunk_ao.datasets.Datasets, "get") - @patch.object(splunk_ao.jobs.Jobs, "create") @patch.object(splunk_ao.experiments.Experiments, "create", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Experiments, "get", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Projects, "get_with_env_fallbacks", return_value=None) def test_run_experiment_with_invalid_project_id_gives_error( - self, - mock_get_project: Mock, - mock_get_experiment: Mock, - mock_create_job: Mock, - mock_get_dataset: Mock, - dataset_content: DatasetContent, + self, mock_get_project: Mock, mock_get_experiment: Mock, mock_get_dataset: Mock, dataset_content: DatasetContent ) -> None: - mock_create_job.return_value = MagicMock() - dataset_id = str(UUID(int=0)) with pytest.raises(ValueError) as exc_info: run_experiment( @@ -616,20 +591,12 @@ def test_run_experiment_with_invalid_project_id_gives_error( assert str(exc_info.value) == "Project with Id awesome-new-project does not exist" @patch.object(splunk_ao.datasets.Datasets, "get") - @patch.object(splunk_ao.jobs.Jobs, "create") @patch.object(splunk_ao.experiments.Experiments, "create", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Experiments, "get", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Projects, "get_with_env_fallbacks", return_value=None) def test_run_experiment_with_invalid_project_name_gives_error( - self, - mock_get_project: Mock, - mock_get_experiment: Mock, - mock_create_job: Mock, - mock_get_dataset: Mock, - dataset_content: DatasetContent, + self, mock_get_project: Mock, mock_get_experiment: Mock, mock_get_dataset: Mock, dataset_content: DatasetContent ) -> None: - mock_create_job.return_value = MagicMock() - dataset_id = str(UUID(int=0)) with pytest.raises(ValueError) as exc_info: run_experiment( @@ -677,7 +644,6 @@ def test_run_experiment_without_metrics( @pytest.mark.parametrize("console_url", ["http://fake.test:8088", "http://fake.test:8088/"]) @travel(datetime(2012, 1, 1), tick=False) @patch.object(splunk_ao.datasets.Datasets, "get") - @patch.object(splunk_ao.jobs.Jobs, "create") @patch.object(splunk_ao.experiments.Experiments, "create", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Experiments, "get", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Projects, "get_with_env_fallbacks", return_value=project()) @@ -686,13 +652,11 @@ def test_run_experiment_link_no_double_slash( mock_get_project: Mock, mock_get_experiment: Mock, mock_create_experiment: Mock, - mock_create_job: Mock, mock_get_dataset: Mock, console_url: str, dataset_content: DatasetContent, ) -> None: # Given: a console_url with or without a trailing slash - mock_create_job.return_value = MagicMock() mock_config = MagicMock() mock_config.console_url = console_url @@ -776,17 +740,11 @@ def test_run_experiment_prompt_takes_precedence_over_generated_output( ) @patch.object(splunk_ao.datasets.Datasets, "get", return_value=None) - @patch.object(splunk_ao.jobs.Jobs, "create") @patch.object(splunk_ao.experiments.Experiments, "create", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Experiments, "get", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Projects, "get_with_env_fallbacks", return_value=project()) def test_run_experiment_no_prompt_no_dataset_raises( - self, - mock_get_project: Mock, - mock_get_experiment: Mock, - mock_create_experiment: Mock, - mock_create_job: Mock, - mock_get_dataset: Mock, + self, mock_get_project: Mock, mock_get_experiment: Mock, mock_create_experiment: Mock, mock_get_dataset: Mock ) -> None: # Given: no prompt_template and no dataset # When/Then: ValueError is raised requiring a dataset @@ -1166,7 +1124,6 @@ def test_experiments_run_with_prompt_settings_as_dict(self, mock_create: Mock) - @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") @patch.object(splunk_ao.datasets.Datasets, "get") - @patch.object(splunk_ao.jobs.Jobs, "create") @patch.object(splunk_ao.experiments.Experiments, "create", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Experiments, "get", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Projects, "get_with_env_fallbacks", return_value=project()) @@ -1175,7 +1132,6 @@ def test_run_experiment_with_runner_and_dataset( mock_get_project: Mock, mock_get_experiment: Mock, mock_create_experiment: Mock, - mock_create_job: Mock, mock_get_dataset: Mock, mock_traces_client: Mock, mock_projects_client: Mock, @@ -1187,8 +1143,6 @@ def test_run_experiment_with_runner_and_dataset( setup_mock_projects_client(mock_projects_client) setup_mock_logstreams_client(mock_logstreams_client) - mock_create_job.return_value = MagicMock() - # mock dataset.get_content # Return dataset_content on first call (starting_token=0), then None to signal end of pagination mock_get_dataset_instance = mock_get_dataset.return_value @@ -1523,7 +1477,6 @@ def test_run_experiment_job_creation_failure( @patch("splunk_ao.experiments.upsert_experiment_tag") @patch.object(splunk_ao.datasets.Datasets, "get") - @patch.object(splunk_ao.jobs.Jobs, "create") @patch.object(splunk_ao.experiments.Experiments, "create", return_value=experiment_response()) @patch.object(splunk_ao.experiments.Experiments, "get", return_value=None) @patch.object(splunk_ao.experiments.Projects, "get_with_env_fallbacks", return_value=project()) @@ -1532,13 +1485,11 @@ def test_run_experiment_with_experiment_tags_basic( mock_get_project: Mock, mock_get_experiment: Mock, mock_create_experiment: Mock, - mock_create_job: Mock, mock_get_dataset: Mock, mock_upsert_tag: Mock, dataset_content: DatasetContent, ) -> None: """Test that experiment_tags are applied when running experiments.""" - mock_create_job.return_value = MagicMock() mock_get_dataset_instance = mock_get_dataset.return_value mock_get_dataset_instance.get_content = MagicMock(return_value=dataset_content) diff --git a/tests/test_job_progress.py b/tests/test_job_progress.py deleted file mode 100644 index 0c77203d..00000000 --- a/tests/test_job_progress.py +++ /dev/null @@ -1,180 +0,0 @@ -import logging -import re -from unittest.mock import ANY, Mock, patch -from uuid import uuid4 - -import pytest -from pytest import CaptureFixture, LogCaptureFixture - -from galileo_core.constants.job import JobStatus -from splunk_ao.job_progress import job_progress, scorer_jobs_status -from splunk_ao.resources.models import HTTPValidationError, JobDB, ValidationError - -FIXED_PROJECT_ID = str(uuid4()) -FIXED_RUN_ID = str(uuid4()) -FIXED_JOB_ID = str(uuid4()) - - -def _job_db_factory( - *, - project_id: str = FIXED_PROJECT_ID, - run_id: str = FIXED_RUN_ID, - job_id: str | None = None, - job_name: str = "test-job", - status: JobStatus = JobStatus.completed, - request_data: dict | None = None, - error_message: str | None = None, - steps_total: int = 100, - steps_completed: int = 100, - progress_message: str = "Done", -) -> JobDB: - data = { - "id": str(job_id or uuid4()), - "project_id": str(project_id), - "run_id": str(run_id), - "job_name": job_name, - "status": status.value, - "request_data": request_data or {}, - "error_message": error_message, - "steps_total": steps_total, - "steps_completed": steps_completed, - "progress_message": progress_message, - "created_at": "2023-01-01T00:00:00", - "updated_at": "2023-01-01T00:00:00", - "retries": 0, - } - return JobDB.from_dict(data) - - -class TestJobProgress: - @patch("splunk_ao.job_progress.get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get.sync") - @patch("splunk_ao.job_progress.get_job_jobs_job_id_get.sync") - def test_completed(self, mock_get_job: Mock, mock_get_scorer_jobs: Mock): - mock_get_job.return_value = _job_db_factory(status=JobStatus.completed) - mock_get_scorer_jobs.return_value = [] - - job_progress(job_id=FIXED_JOB_ID, project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - mock_get_job.assert_called_with(client=ANY, job_id=FIXED_JOB_ID) - - @patch("splunk_ao.job_progress.get_job_jobs_job_id_get.sync") - def test_failed(self, mock_get_job: Mock): - mock_get_job.return_value = _job_db_factory(status=JobStatus.failed, error_message="Test error") - - with pytest.raises(ValueError, match="Job failed with error message Test error."): - job_progress(job_id=FIXED_JOB_ID, project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - mock_get_job.assert_called_with(client=ANY, job_id=FIXED_JOB_ID) - - @patch("splunk_ao.job_progress.get_job_jobs_job_id_get.sync") - def test_get_job_fails(self, mock_get_job: Mock): - mock_get_job.return_value = None - - with pytest.raises(ValueError, match=f"Failed to get job status for job {FIXED_JOB_ID}"): - job_progress(job_id=FIXED_JOB_ID, project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - mock_get_job.assert_called_with(client=ANY, job_id=FIXED_JOB_ID) - - @patch("splunk_ao.job_progress.get_job_jobs_job_id_get.sync") - def test_get_job_http_validation_error(self, mock_get_job: Mock): - detail = [ValidationError(loc=["path", "job_id"], msg="value is not a valid uuid", type_="type_error.uuid")] - mock_get_job.return_value = HTTPValidationError(detail=detail) - - with pytest.raises(ValueError, match=re.escape(str(detail))): - job_progress(job_id=FIXED_JOB_ID, project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - - -class TestScorerJobsStatus: - @patch("splunk_ao.job_progress.get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get.sync") - def test_simple(self, mock_get_jobs: Mock, caplog: LogCaptureFixture, enable_galileo_logging): - mock_get_jobs.return_value = [ - _job_db_factory( - job_name="log_stream_scorer", - status=JobStatus.in_progress, - request_data={"prompt_scorer_settings": {"scorer_name": "pii"}}, - ) - ] - - with caplog.at_level(logging.INFO, logger="splunk_ao.job_progress"): - scorer_jobs_status(project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - assert "pii: Computing 🚧" in caplog.text - - @patch("splunk_ao.job_progress.get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get.sync") - def test_skips_prompt_run(self, mock_get_jobs: Mock, caplog: LogCaptureFixture, enable_galileo_logging): - mock_get_jobs.return_value = [ - _job_db_factory(job_name="log_stream_run"), - _job_db_factory( - job_name="log_stream_scorer", - status=JobStatus.in_progress, - request_data={"prompt_scorer_settings": {"scorer_name": "pii"}}, - ), - ] - - with caplog.at_level(logging.INFO, logger="splunk_ao.job_progress"): - scorer_jobs_status(project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - assert "pii: Computing 🚧" in caplog.text - - @patch("splunk_ao.job_progress.get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get.sync") - def test_no_scorer_jobs(self, mock_get_jobs: Mock, capsys: CaptureFixture[str]): - mock_get_jobs.return_value = [_job_db_factory(job_name="log_stream_run")] - - scorer_jobs_status(project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - captured = capsys.readouterr() - assert captured.out == "" - - @patch("splunk_ao.job_progress.get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get.sync") - def test_one_of_each(self, mock_get_jobs: Mock, caplog: LogCaptureFixture, enable_galileo_logging): - mock_get_jobs.return_value = [ - _job_db_factory(job_name="log_stream_run"), - _job_db_factory( - job_name="log_stream_scorer", - status=JobStatus.in_progress, - request_data={"prompt_scorer_settings": {"scorer_name": "pii"}}, - ), - _job_db_factory( - job_name="log_stream_scorer", - status=JobStatus.error, - error_message="An error occurred.", - request_data={"prompt_scorer_settings": {"scorer_name": "toxicity"}}, - ), - _job_db_factory( - job_name="log_stream_scorer", - status=JobStatus.completed, - request_data={"prompt_scorer_settings": {"scorer_name": "chunk_attribution_utilization_plus"}}, - ), - ] - - with caplog.at_level(logging.INFO, logger="splunk_ao.job_progress"): - scorer_jobs_status(project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - assert "pii: Computing 🚧" in caplog.text - assert "toxicity: Failed ❌, error was: An error occurred." in caplog.text - assert "chunk_attribution_utilization_plus: Done ✅" in caplog.text - - @patch("splunk_ao.job_progress.get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get.sync") - def test_unknown_scorer_name(self, mock_get_jobs: Mock, caplog: LogCaptureFixture, enable_galileo_logging): - mock_get_jobs.return_value = [ - _job_db_factory(job_name="log_stream_run"), - _job_db_factory( - job_name="log_stream_scorer", - status=JobStatus.in_progress, - request_data={"prompt_scorer_settings": {"scorer_name": "abc"}}, - ), - ] - - with caplog.at_level(logging.INFO, logger="splunk_ao.job_progress"): - scorer_jobs_status(project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - assert "abc: Computing 🚧" in caplog.text - - @patch("splunk_ao.job_progress.get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get.sync") - def test_get_run_scorer_jobs_fails(self, mock_get_jobs: Mock): - mock_get_jobs.return_value = None - - with pytest.raises( - ValueError, match=f"Failed to get scorer jobs for project {FIXED_PROJECT_ID}, run {FIXED_RUN_ID}" - ): - scorer_jobs_status(project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) - - @patch("splunk_ao.job_progress.get_jobs_for_project_run_projects_project_id_runs_run_id_jobs_get.sync") - def test_get_run_scorer_jobs_http_validation_error(self, mock_get_jobs: Mock): - detail = [ValidationError(loc=["path", "project_id"], msg="value is not a valid uuid", type_="type_error.uuid")] - mock_get_jobs.return_value = HTTPValidationError(detail=detail) - - with pytest.raises(ValueError, match=re.escape(str(detail))): - scorer_jobs_status(project_id=FIXED_PROJECT_ID, run_id=FIXED_RUN_ID) diff --git a/tests/test_jobs.py b/tests/test_jobs.py deleted file mode 100644 index aa41c870..00000000 --- a/tests/test_jobs.py +++ /dev/null @@ -1,106 +0,0 @@ -from http import HTTPStatus -from unittest.mock import MagicMock, patch - -import pytest - -from splunk_ao.jobs import Jobs -from splunk_ao.resources.models import HTTPValidationError, PromptRunSettings, TaskType, ValidationError -from splunk_ao.resources.types import Response - - -def _make_422_response(msg: str = "Invalid model alias: 'gpt-4o-mini'") -> Response: - return Response( - status_code=HTTPStatus(422), - content=b"{}", - headers={}, - parsed=HTTPValidationError( - detail=[ValidationError(loc=["body", "prompt_settings", "model_alias"], msg=msg, type_="value_error")] - ), - ) - - -def _make_job_kwargs(**overrides): - defaults = dict( - project_id="proj-id", - name="test-job", - run_id="run-id", - dataset_id="ds-id", - prompt_template_id=None, - task_type=TaskType.VALUE_16, - scorers=None, - prompt_settings=PromptRunSettings(model_alias="gpt-4o-mini"), - ) - return {**defaults, **overrides} - - -class TestJobsCreate: - @patch("splunk_ao.jobs.create_job_jobs_post") - def test_raises_value_error_with_clear_message_on_invalid_model_alias(self, mock_post: MagicMock) -> None: - """Jobs.create() with invalid model_alias (HTTP 422) raises ValueError with readable message.""" - # Given: the API returns a 422 with validation error for model_alias - mock_post.sync_detailed = MagicMock(return_value=_make_422_response()) - - # When/Then: ValueError is raised with a human-readable message - with pytest.raises(ValueError, match="Request validation failed"): - Jobs().create(**_make_job_kwargs()) - - @patch("splunk_ao.jobs.create_job_jobs_post") - def test_error_message_contains_field_path_and_backend_message(self, mock_post: MagicMock) -> None: - """The ValueError includes the field path and backend error message.""" - # Given: the API returns a 422 with specific validation detail - mock_post.sync_detailed = MagicMock(return_value=_make_422_response("Invalid model alias: 'gpt-4o-mini'")) - - # When/Then: the error includes the field path and backend message - with pytest.raises(ValueError) as exc_info: - Jobs().create(**_make_job_kwargs()) - msg = str(exc_info.value) - assert "model_alias" in msg - assert "gpt-4o-mini" in msg - - @patch("splunk_ao.jobs.create_job_jobs_post") - def test_does_not_raise_galileo_http_exception(self, mock_post: MagicMock) -> None: - """Jobs.create() no longer raises the empty GalileoHTTPException on 422.""" - # Given: the API returns a 422 response - mock_post.sync_detailed = MagicMock(return_value=_make_422_response()) - - # When/Then: the raised exception is ValueError, not GalileoHTTPException - from galileo_core.exceptions.http import GalileoHTTPException - - with pytest.raises(ValueError): - Jobs().create(**_make_job_kwargs()) - # Ensure GalileoHTTPException is NOT raised - try: - Jobs().create(**_make_job_kwargs()) - except GalileoHTTPException: - pytest.fail("GalileoHTTPException should no longer be raised") - except ValueError: - pass # expected - - @patch("splunk_ao.jobs.create_job_jobs_post") - def test_raises_value_error_on_unexpected_non_200(self, mock_post: MagicMock) -> None: - """Jobs.create() raises ValueError with status code for non-422 unexpected responses.""" - # Given: the API returns an unexpected non-200 response with no parsed body - mock_post.sync_detailed = MagicMock( - return_value=Response(status_code=HTTPStatus(503), content=b"Service Unavailable", headers={}, parsed=None) - ) - - # When/Then: ValueError is raised with the status code in the message - with pytest.raises(ValueError, match="503"): - Jobs().create(**_make_job_kwargs()) - - @patch("splunk_ao.jobs.create_job_jobs_post") - def test_raises_value_error_when_api_returns_string_detail(self, mock_post: MagicMock) -> None: - """Jobs.create() surfaces the API message when the 422 detail is a plain string.""" - # Given: the API returns a 422 whose 'detail' is a plain string (not the standard list shape) - mock_post.sync_detailed = MagicMock( - return_value=Response( - status_code=HTTPStatus(422), - content=b'{"detail": "Model alias not found"}', - headers={}, - parsed=HTTPValidationError.from_dict({"detail": "Model alias not found"}), - ) - ) - - # When/Then: ValueError contains the original API message - with pytest.raises(ValueError, match="Model alias not found"): - Jobs().create(**_make_job_kwargs())