345 lines
13 KiB
Python
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
|