Files
foxhunt/tests/runpod/test_monitor.py

345 lines
13 KiB
Python

"""
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