File size: 7,462 Bytes
41ea373
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
258783b
 
 
 
41ea373
 
 
 
 
 
 
 
 
258783b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
41ea373
e6b6793
41ea373
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
98b952a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
import json
from pathlib import Path
from unittest.mock import MagicMock, patch

import pytest

from viral_script_engine.agents.critic import CritiqueOutput, CritiqueClaim
from viral_script_engine.agents.rewriter import RewriteResult
from viral_script_engine.environment.actions import ActionType, ArbitratorAction

FIXTURE_DIR = Path(__file__).parent.parent / "data" / "golden_fixtures"
SCRIPTS_PATH = str(Path(__file__).parent.parent / "data" / "test_scripts" / "scripts.json")


def load_fixture(script_id: str) -> dict:
    with open(FIXTURE_DIR / f"fixture_{script_id}.json") as f:
        return json.load(f)


def make_mock_critique() -> CritiqueOutput:
    fixture = load_fixture("S01")
    claims = [CritiqueClaim(**c) for c in fixture["critique"]["claims"]]
    return CritiqueOutput(
        claims=claims,
        overall_severity=fixture["critique"]["overall_severity"],
        raw_response=fixture["critique"]["raw_response"],
    )


def make_mock_rewrite(current_script: str, action: ArbitratorAction) -> RewriteResult:
    return RewriteResult(
        rewritten_script=current_script + " [REWRITTEN]",
        diff="@@ diff @@",
        word_count_delta=1,
    )


SAMPLE_ACTION = {
    "action_type": ActionType.HOOK_REWRITE.value,
    "target_section": "hook",
    "instruction": "Make the hook more attention-grabbing with a specific number.",
    "critique_claim_id": "C1",
    "reasoning": "Hook is weak per C1",
}


@pytest.fixture
def env():
    with (
        patch("viral_script_engine.environment.env.CriticAgent") as mock_critic_cls,
        patch("viral_script_engine.environment.env.RewriterAgent") as mock_rewriter_cls,
        patch("viral_script_engine.environment.env.DefenderAgent") as mock_defender_cls,
        patch("viral_script_engine.environment.env.CulturalAlignmentReward") as mock_r3_cls,
        patch("viral_script_engine.environment.env.DebateResolutionReward") as mock_r4_cls,
        patch("viral_script_engine.environment.env.DefenderPreservationReward") as mock_r5_cls,
    ):
        mock_critic = MagicMock()
        mock_critic.critique.return_value = make_mock_critique()
        mock_critic_cls.return_value = mock_critic

        mock_rewriter = MagicMock()
        mock_rewriter.rewrite.side_effect = make_mock_rewrite
        mock_rewriter_cls.return_value = mock_rewriter

        mock_defender = MagicMock()
        from viral_script_engine.agents.defender import DefenderOutput
        mock_defender.defend.return_value = DefenderOutput(
            core_strength="strong hook",
            core_strength_quote="test quote",
            defense_argument="preserve it",
            flagged_critic_claims=[],
            regional_voice_elements=[],
        )
        mock_defender_cls.return_value = mock_defender

        mock_r3 = MagicMock()
        mock_r3.score.return_value = MagicMock(score=0.6)
        mock_r3_cls.return_value = mock_r3

        mock_r4 = MagicMock()
        from viral_script_engine.rewards.r4_debate_resolution import DebateResolutionResult
        mock_r4.score.return_value = DebateResolutionResult(
            score=0.8,
            resolution_status="resolved",
            original_claim_id="C1",
            original_claim_class="hook_weakness",
            new_claims_count=2,
        )
        mock_r4_cls.return_value = mock_r4

        mock_r5 = MagicMock()
        from viral_script_engine.rewards.r5_defender_preservation import DefenderPreservationResult
        mock_r5.score.return_value = DefenderPreservationResult(
            score=0.9, max_similarity=0.9, best_matching_sentence="test quote"
        )
        mock_r5_cls.return_value = mock_r5

        from viral_script_engine.environment.env import ViralScriptEnv
        yield ViralScriptEnv(scripts_path=SCRIPTS_PATH, max_steps=5, difficulty="easy", use_escalation=False)


def test_reset_returns_valid_observation(env):
    obs, info = env.reset(seed=42)
    assert "current_script" in obs
    assert obs["step_num"] == 0
    assert obs["max_steps"] == 5
    assert obs["reward_components"]["r1_hook_strength"] is not None
    assert obs["reward_components"]["r2_coherence"] is not None


def test_step_completes_without_error(env):
    env.reset(seed=42)
    obs, reward, terminated, truncated, info = env.step(SAMPLE_ACTION)
    assert isinstance(reward, float)
    assert "reward_components" in info


def test_step_increments_step_num(env):
    env.reset(seed=42)
    obs, *_ = env.step(SAMPLE_ACTION)
    assert obs["step_num"] == 1
    obs, *_ = env.step(SAMPLE_ACTION)
    assert obs["step_num"] == 2


def test_anti_gaming_penalty_fires_on_repeated_action(env):
    env.reset(seed=42)
    for _ in range(3):
        obs, reward, _, _, info = env.step(SAMPLE_ACTION)
    assert info["anti_gaming_triggered"]


def test_episode_terminates_at_max_steps(env):
    env.reset(seed=42)
    terminated = False
    for _ in range(5):
        obs, reward, terminated, truncated, info = env.step(SAMPLE_ACTION)
    assert terminated


def test_reward_clipped_to_0_1(env):
    env.reset(seed=42)
    _, reward, _, _, _ = env.step(SAMPLE_ACTION)
    assert 0.0 <= reward <= 1.0


def test_timeout_truncates_episode(monkeypatch):
    """Verify that a hanging LLM call causes truncated=True, not an infinite hang."""
    import time

    def slow_generate(*args, **kwargs):
        time.sleep(200)

    with (
        patch("viral_script_engine.environment.env.CriticAgent") as mock_critic_cls,
        patch("viral_script_engine.environment.env.RewriterAgent") as mock_rewriter_cls,
        patch("viral_script_engine.environment.env.DefenderAgent") as mock_defender_cls,
        patch("viral_script_engine.environment.env.CulturalAlignmentReward") as mock_r3_cls,
        patch("viral_script_engine.environment.env.DebateResolutionReward") as mock_r4_cls,
        patch("viral_script_engine.environment.env.DefenderPreservationReward") as mock_r5_cls,
    ):
        from viral_script_engine.agents.llm_backend import LLMBackend

        mock_critic = MagicMock()
        mock_critic.critique.side_effect = TimeoutError("LLM call timed out after 30s")
        mock_critic_cls.return_value = mock_critic

        mock_rewriter_cls.return_value = MagicMock()
        mock_defender_cls.return_value = MagicMock()

        mock_r3 = MagicMock()
        mock_r3.score.return_value = MagicMock(score=0.6)
        mock_r3_cls.return_value = mock_r3

        mock_r4 = MagicMock()
        from viral_script_engine.rewards.r4_debate_resolution import DebateResolutionResult
        mock_r4.score.return_value = DebateResolutionResult(
            score=0.8, resolution_status="resolved",
            original_claim_id="C1", original_claim_class="hook_weakness", new_claims_count=2,
        )
        mock_r4_cls.return_value = mock_r4

        mock_r5 = MagicMock()
        from viral_script_engine.rewards.r5_defender_preservation import DefenderPreservationResult
        mock_r5.score.return_value = DefenderPreservationResult(
            score=0.9, max_similarity=0.9, best_matching_sentence="test quote"
        )
        mock_r5_cls.return_value = mock_r5

        from viral_script_engine.environment.env import ViralScriptEnv
        env = ViralScriptEnv(scripts_path=SCRIPTS_PATH, max_steps=5, difficulty="easy", use_escalation=False)
        env.reset(seed=42)
        _, _, terminated, truncated, info = env.step(SAMPLE_ACTION)

    assert truncated is True
    assert info.get("timeout") is True