File size: 2,725 Bytes
79412de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e987a94
79412de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e987a94
 
 
79412de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
import pytest
import os
import io
import contextlib
from unittest.mock import patch, MagicMock
from utils.constants import DEFAULT_TASK_ID
from typing import Dict, List, Optional
import re

import inference

# A predictable sequence of LLM responses to test the robustness of the inference loop
MOCK_RESPONSES = [
    '{"U_target": 0.40, "F_target": 0.60}', # 1. Valid JSON
    '0.45 0.55',                            # 2. Valid Pair format
    '0.45, 0.55',                           # 3. Valid Pair format (comma)
    'Hello world',                          # 4. Invalid fallback
    '{"U_target": 1.5, "F_target": -0.5}',  # 5. Out of bounds (parser clamps)
]

def mock_get_model_response(
    client: MagicMock,
    task_description: str,
    step: int,
    observation: Dict[str, float],
    last_reward: float,
    history: List[str],
) -> str:
    # 0-indexed step from 1-indexed step argument
    idx = (step - 1) % len(MOCK_RESPONSES)
    return MOCK_RESPONSES[idx]

@pytest.fixture
def mock_inference_env(monkeypatch):
    monkeypatch.setenv("HF_TOKEN", "fake_token")
    monkeypatch.setenv("MODEL_NAME", "fake_model")
    monkeypatch.setenv("API_BASE_URL", "http://fake.api")
    monkeypatch.setenv("THERMAL_PLANT_EPISODE_ID", "1")
    # Tell the default task it's a short episode
    import tasks.task2
    monkeypatch.setattr(tasks.task2.Task2, "max_steps", 5)
    yield

def test_inference_main_prints_exact_stdout_format(mock_inference_env, monkeypatch):
    # Hijack the model call
    monkeypatch.setattr(inference, "get_model_response", mock_get_model_response)

    stdout_capture = io.StringIO()
    with contextlib.redirect_stdout(stdout_capture):
        inference.main()
        
    output = stdout_capture.getvalue()
    lines = output.strip().split("\n")
    
    # Needs at least START, END, and one STEP
    assert len(lines) >= 3
    
    assert lines[0].startswith("[START] task=")
    assert lines[-1].startswith("[END] success=")
    
    step_pattern = re.compile(
        r'^\[STEP\] step=\d+ action=\{"U_target":\d+\.\d{2},"F_target":\d+\.\d{2}\} '
        r'reward=-?\d+\.\d{2} done=(true|false) error=(null|.+)$'
    )
    
    step_count = 0
    for line in lines[1:-1]:
        # Filter debug lines if they slipped through to stdout (they shouldn't; inference prints to stderr for debug)
        if line.startswith("[STEP]"):
            assert step_pattern.match(line) is not None, f"Invalid STEP format: {line}"
            step_count += 1
            
            # Step 4 is 'Hello world', which triggers a parser fallback error
            if "step=4" in line:
                assert "error=parse" in line

    assert step_count == 5, f"Expected exactly 5 steps, got {step_count}"