""" Tests for runpod.monitor module. Tests status polling, S3 log tailing, and completion detection. """ import pytest from unittest.mock import MagicMock, patch, call import time from runpod.monitor import PodMonitor, PodStatus from runpod.client import RunPodClient, RunPodAPIError from runpod.s3_client import S3Client class TestPodStatus: """Tests for PodStatus enum.""" def test_pod_status_values(self): """Test PodStatus enum values.""" assert PodStatus.PENDING.value == "PENDING" assert PodStatus.RUNNING.value == "RUNNING" assert PodStatus.COMPLETED.value == "COMPLETED" assert PodStatus.FAILED.value == "FAILED" assert PodStatus.STOPPED.value == "STOPPED" assert PodStatus.UNKNOWN.value == "UNKNOWN" class TestPodMonitor: """Tests for PodMonitor class.""" def test_monitor_initialization(self, sample_config): """Test monitor initializes correctly.""" client = RunPodClient(sample_config) monitor = PodMonitor(client) assert monitor.client == client assert monitor.s3_client is None def test_monitor_initialization_with_s3(self, sample_config, mock_s3_client): """Test monitor initializes with S3 client.""" client = RunPodClient(sample_config) s3_client = MagicMock(spec=S3Client) monitor = PodMonitor(client, s3_client) assert monitor.client == client assert monitor.s3_client == s3_client def test_parse_status_known_states(self, sample_config): """Test parsing known pod status states.""" client = RunPodClient(sample_config) monitor = PodMonitor(client) test_cases = [ ({"desiredStatus": "RUNNING"}, PodStatus.RUNNING), ({"desiredStatus": "PENDING"}, PodStatus.PENDING), ({"desiredStatus": "COMPLETED"}, PodStatus.COMPLETED), ({"desiredStatus": "FAILED"}, PodStatus.FAILED), ({"desiredStatus": "STOPPED"}, PodStatus.STOPPED), ] for pod_data, expected_status in test_cases: status = monitor._parse_status(pod_data) assert status == expected_status def test_parse_status_unknown_state(self, sample_config): """Test parsing unknown pod status.""" client = RunPodClient(sample_config) monitor = PodMonitor(client) pod_data = {"desiredStatus": "WEIRD_STATE"} status = monitor._parse_status(pod_data) assert status == PodStatus.UNKNOWN def test_parse_status_missing_field(self, sample_config): """Test parsing status with missing desiredStatus field.""" client = RunPodClient(sample_config) monitor = PodMonitor(client) pod_data = {} status = monitor._parse_status(pod_data) assert status == PodStatus.UNKNOWN def test_poll_status_completes_successfully(self, sample_config, mock_api_responses): """Test polling that completes successfully.""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = mock_api_responses["pod_status_completed"] monitor = PodMonitor(client) final_status = monitor.poll_status("test-pod-123456", interval_seconds=1, timeout_seconds=10) assert final_status == PodStatus.COMPLETED assert client.get_pod_status.call_count >= 1 def test_poll_status_fails(self, sample_config, mock_api_responses): """Test polling when pod fails.""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = mock_api_responses["pod_status_failed"] monitor = PodMonitor(client) final_status = monitor.poll_status("test-pod-123456", interval_seconds=1, timeout_seconds=10) assert final_status == PodStatus.FAILED def test_poll_status_timeout(self, sample_config, mock_api_responses): """Test polling timeout.""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = mock_api_responses["pod_status_running"] monitor = PodMonitor(client) with pytest.raises(TimeoutError, match="timeout after"): monitor.poll_status("test-pod-123456", interval_seconds=1, timeout_seconds=2) def test_poll_status_with_callback(self, sample_config, mock_api_responses): """Test polling with status change callback.""" client = MagicMock(spec=RunPodClient) # Simulate status progression: PENDING -> RUNNING -> COMPLETED status_sequence = [ {"desiredStatus": "PENDING"}, {"desiredStatus": "RUNNING"}, {"desiredStatus": "RUNNING"}, {"desiredStatus": "COMPLETED"}, ] client.get_pod_status.side_effect = status_sequence monitor = PodMonitor(client) callback = MagicMock() final_status = monitor.poll_status( "test-pod-123456", interval_seconds=0.1, timeout_seconds=10, on_status_change=callback ) assert final_status == PodStatus.COMPLETED # Callback should be called for each unique status assert callback.call_count == 3 # PENDING, RUNNING, COMPLETED def test_poll_status_handles_api_errors(self, sample_config, mock_api_responses): """Test polling continues despite API errors.""" client = MagicMock(spec=RunPodClient) # First call fails, second succeeds client.get_pod_status.side_effect = [ RunPodAPIError("Temporary error"), mock_api_responses["pod_status_completed"] ] monitor = PodMonitor(client) with patch('builtins.print'): # Suppress error output final_status = monitor.poll_status( "test-pod-123456", interval_seconds=0.1, timeout_seconds=10 ) assert final_status == PodStatus.COMPLETED def test_tail_logs_requires_s3_client(self, sample_config): """Test tail_logs raises error without S3 client.""" client = RunPodClient(sample_config) monitor = PodMonitor(client) with pytest.raises(ValueError, match="S3 client is required"): monitor.tail_logs("test-pod", "logs/test.log", follow=False) def test_tail_logs_without_follow(self, sample_config, mock_log_content): """Test tailing logs without following.""" client = RunPodClient(sample_config) s3_client = MagicMock(spec=S3Client) s3_client.stream_logs.return_value = mock_log_content monitor = PodMonitor(client, s3_client) with patch('builtins.print') as mock_print: monitor.tail_logs("test-pod", "logs/test.log", follow=False) # Should print log content assert mock_print.called s3_client.stream_logs.assert_called_once() def test_tail_logs_with_follow_until_completion(self, sample_config, mock_log_content, mock_api_responses): """Test tailing logs with follow mode until completion.""" client = MagicMock(spec=RunPodClient) s3_client = MagicMock(spec=S3Client) # Simulate log streaming with pod completion s3_client.stream_logs.side_effect = [ "Starting...\n", "Training...\n", "Completed!\n", "" # Final check after completion ] # Pod completes after a few checks client.get_pod_status.side_effect = [ {"desiredStatus": "RUNNING"}, {"desiredStatus": "RUNNING"}, {"desiredStatus": "COMPLETED"} ] monitor = PodMonitor(client, s3_client) with patch('builtins.print'): monitor.tail_logs("test-pod", "logs/test.log", follow=True, interval_seconds=0.1) # Should have checked multiple times assert s3_client.stream_logs.call_count >= 3 def test_tail_logs_handles_s3_errors(self, sample_config): """Test tail_logs handles S3 errors gracefully.""" client = MagicMock(spec=RunPodClient) s3_client = MagicMock(spec=S3Client) s3_client.stream_logs.side_effect = Exception("S3 error") monitor = PodMonitor(client, s3_client) with patch('builtins.print') as mock_print: monitor.tail_logs("test-pod", "logs/test.log", follow=False) # Should print error message assert any("Error" in str(call) for call in mock_print.call_args_list) def test_detect_completion_with_completed_status(self, sample_config, mock_api_responses): """Test completion detection via pod status.""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = mock_api_responses["pod_status_completed"] monitor = PodMonitor(client) result = monitor.detect_completion("test-pod", check_interval=0.1, max_wait_seconds=10) assert result is True def test_detect_completion_with_failed_status(self, sample_config, mock_api_responses): """Test completion detection with failed pod.""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = mock_api_responses["pod_status_failed"] monitor = PodMonitor(client) result = monitor.detect_completion("test-pod", check_interval=0.1, max_wait_seconds=10) assert result is False def test_detect_completion_with_log_markers(self, sample_config, mock_log_content): """Test completion detection via log markers.""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = {"desiredStatus": "RUNNING"} s3_client = MagicMock(spec=S3Client) s3_client.stream_logs.return_value = mock_log_content # Contains "Training complete" monitor = PodMonitor(client, s3_client) result = monitor.detect_completion( "test-pod", completion_markers=["Training complete"], check_interval=0.1, max_wait_seconds=10 ) assert result is True def test_detect_completion_custom_markers(self, sample_config): """Test completion detection with custom markers.""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = {"desiredStatus": "RUNNING"} s3_client = MagicMock(spec=S3Client) s3_client.stream_logs.return_value = "Custom completion marker found" monitor = PodMonitor(client, s3_client) result = monitor.detect_completion( "test-pod", completion_markers=["Custom completion marker"], check_interval=0.1, max_wait_seconds=10 ) assert result is True def test_detect_completion_timeout(self, sample_config): """Test completion detection timeout.""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = {"desiredStatus": "RUNNING"} monitor = PodMonitor(client) with pytest.raises(TimeoutError, match="timeout after"): monitor.detect_completion( "test-pod", check_interval=0.1, max_wait_seconds=1 ) def test_detect_completion_without_s3(self, sample_config, mock_api_responses): """Test completion detection without S3 client (status only).""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = mock_api_responses["pod_status_completed"] monitor = PodMonitor(client) # No S3 client result = monitor.detect_completion("test-pod", check_interval=0.1, max_wait_seconds=10) assert result is True def test_detect_completion_handles_s3_errors(self, sample_config, mock_api_responses): """Test completion detection handles S3 errors gracefully.""" client = MagicMock(spec=RunPodClient) s3_client = MagicMock(spec=S3Client) # S3 fails but pod eventually completes s3_client.stream_logs.side_effect = Exception("S3 error") client.get_pod_status.side_effect = [ {"desiredStatus": "RUNNING"}, mock_api_responses["pod_status_completed"] ] monitor = PodMonitor(client, s3_client) result = monitor.detect_completion("test-pod", check_interval=0.1, max_wait_seconds=10) assert result is True @pytest.mark.parametrize("marker", [ "Training complete", "Model saved successfully", "Checkpoint saved" ]) def test_detect_completion_default_markers(self, sample_config, marker): """Test default completion markers.""" client = MagicMock(spec=RunPodClient) client.get_pod_status.return_value = {"desiredStatus": "RUNNING"} s3_client = MagicMock(spec=S3Client) s3_client.stream_logs.return_value = f"Some log output\n{marker}\nMore output" monitor = PodMonitor(client, s3_client) result = monitor.detect_completion("test-pod", check_interval=0.1, max_wait_seconds=10) assert result is True