"""Reward system unit tests — REW-06 + REW-08. >= 33 tests total: 30 unit-test pairs + 3 @pytest.mark.reward_audit tests. """ from __future__ import annotations import random import pytest from unittest.mock import MagicMock from omegaconf import OmegaConf from rewards.format_gate import format_gate from rewards.correctness import correctness from rewards.token_budget import token_budget from rewards.recursion_efficiency import recursion_efficiency from rewards.compose import compose_reward_single, compose_reward_fn, make_reward_fn # --------------------------------------------------------------------------- # Fixture: minimal cfg_reward matching configs/reward/v1.yaml # --------------------------------------------------------------------------- @pytest.fixture def cfg_reward(): return OmegaConf.create({ "alpha": 0.2, "weights": {"correctness": 0.75, "token_budget": 0.2, "recursion_efficiency": 0.05}, "token_budget_variant": "capped_linear", "answer_regex": "(.*?)", "max_calls": 2, }) # --------------------------------------------------------------------------- # TestFormatGate — 7 cases # --------------------------------------------------------------------------- class TestFormatGate: def test_valid_tag(self): assert format_gate("Rome") == 1.0 def test_no_tag(self): assert format_gate("The answer is Rome.") == 0.0 def test_uppercase_tag(self): assert format_gate("Rome") == 1.0 def test_empty_answer_tag(self): assert format_gate("") == 1.0 def test_multiline_tag(self): assert format_gate("Some text\n\nRome\n") == 1.0 def test_no_closing_tag(self): assert format_gate("Rome") == 0.0 def test_nested_content(self): assert format_gate("The city of Rome") == 1.0 # --------------------------------------------------------------------------- # TestCorrectness — 8 cases (covers niah, extractive task types) # --------------------------------------------------------------------------- class TestCorrectness: def test_exact_match(self): assert correctness("Rome", "Rome") == 1.0 def test_case_insensitive(self): assert correctness("rome", "Rome") == 1.0 def test_trailing_period(self): assert correctness("Rome.", "Rome") == 1.0 def test_whitespace_padding(self): assert correctness(" Rome ", "Rome") == 1.0 def test_no_tag_returns_zero(self): assert correctness("The answer is Rome.", "Rome") == 0.0 def test_wrong_answer(self): assert correctness("Paris", "Rome") == 0.0 def test_multiword_gold(self): # niah task: "New York" vs "new york" assert correctness("new york", "New York") == 1.0 def test_empty_extraction(self): assert correctness("", "Rome") == 0.0 # --------------------------------------------------------------------------- # TestTokenBudget — 7 cases (covers REW-04 variants) # --------------------------------------------------------------------------- class TestTokenBudget: def test_short_completion_near_one(self): score = token_budget("hello world", 100, alpha=0.2) assert score > 0.95 def test_long_completion_reduced(self): long_text = " ".join(["word"] * 200) score = token_budget(long_text, 100, alpha=0.2) assert score < 0.8 def test_very_long_clamped_to_zero(self): long_text = " ".join(["word"] * 10000) score = token_budget(long_text, 100, alpha=0.2, variant="capped_linear") assert score == 0.0 def test_capped_quadratic_in_range(self): score = token_budget("short text here", 100, alpha=0.2, variant="capped_quadratic") assert 0.0 <= score <= 1.0 def test_uncapped_can_be_negative(self): long_text = " ".join(["word"] * 10000) score = token_budget(long_text, 10, alpha=2.0, variant="uncapped") assert score < 0.0 def test_unknown_variant_raises(self): with pytest.raises(ValueError, match="REW-04"): token_budget("text", 100, variant="unknown_variant") def test_zero_prompt_tokens_no_crash(self): score = token_budget("hello", 0) assert 0.0 <= score <= 1.0 # --------------------------------------------------------------------------- # TestRecursionEfficiency — 7 cases # --------------------------------------------------------------------------- class TestRecursionEfficiency: def test_zero_calls(self): assert recursion_efficiency(0) == 1.0 def test_one_call(self): assert recursion_efficiency(1) == 0.75 def test_two_calls(self): assert recursion_efficiency(2) == 0.50 def test_three_calls(self): assert recursion_efficiency(3) == 0.25 def test_four_or_more_calls_zero(self): assert recursion_efficiency(4) == 0.0 assert recursion_efficiency(10) == 0.0 def test_negative_treated_as_zero(self): assert recursion_efficiency(-1) == 1.0 def test_max_calls_ignored_in_v2(self): # max_calls kwarg is accepted via **_ but ignored by new linear decay assert recursion_efficiency(2, max_calls=1) == 0.50 # --------------------------------------------------------------------------- # TestCompose — 6 cases (covers counting + multi_needle task types implicitly) # --------------------------------------------------------------------------- class TestCompose: def test_format_fail_short_circuits(self, cfg_reward): score, _ = compose_reward_single("no tag here", "Rome", 100, cfg_reward, llm_call_count=0) assert 0.0 <= score <= 0.25 # A-02 cap applies, but > 0 because of token budget def test_perfect_score(self, cfg_reward): # Short completion, 0 llm calls, correct answer score, _ = compose_reward_single("Rome", "Rome", 1, cfg_reward, llm_call_count=0) assert score == pytest.approx(0.75 * 1.0 + 0.2 * 1.0 + 0.05 * 1.0 + 0.10, abs=0.05) def test_correctness_zero_capped_at_025(self, cfg_reward): # format passes but wrong answer → capped at 0.25 (since format bonus + cap) score, _ = compose_reward_single("Paris", "Rome", 100, cfg_reward, llm_call_count=0) assert score <= 0.25 def test_partial_mix(self, cfg_reward): # Correct answer but long completion (token_budget reduced) long_text = " ".join(["filler"] * 500) + " 42" score, _ = compose_reward_single(long_text, "42", 10, cfg_reward, llm_call_count=0) assert 0.5 < score <= 1.10 # correctness=1 dominates def test_make_reward_fn_factory(self, cfg_reward): bound_fn = make_reward_fn(cfg_reward) result = bound_fn( prompts=["q"], completions=["Rome"], gold_answer=["Rome"], prompt_token_count=[1], llm_call_count=[0], ) assert isinstance(result, list) and len(result) == 1 assert result[0] > 0.5 def test_trl_contract_batched(self, cfg_reward): results = compose_reward_fn( prompts=["q1", "q2"], completions=["42", "no answer here"], gold_answer=["42", "100"], prompt_token_count=[100, 100], llm_call_count=[0, 0], cfg_reward=cfg_reward, ) assert len(results) == 2 assert results[0] > 0.0 assert results[1] <= 0.25 # --------------------------------------------------------------------------- # REW-06: Automated anti-hacking audits (3 tests, @pytest.mark.reward_audit) # --------------------------------------------------------------------------- @pytest.mark.reward_audit def test_masked_context_reward_is_baseline(cfg_reward): """REW-06: gibberish answer with format tag must score <= 0.25 (A-02 cap).""" score, _ = compose_reward_single( "lorem ipsum dolor sit amet", "Rome", 100, cfg_reward, llm_call_count=0, ) assert score <= 0.25, f"Masked-context score too high: {score}" @pytest.mark.reward_audit def test_format_only_composite_is_capped(cfg_reward): """REW-06: 20 format-only completions with mismatched gold must all score <= 0.25.""" words = ["banana", "kiwi", "mango", "peach", "grape", "lemon", "plum", "cherry", "pear", "melon", "papaya", "guava", "fig", "date", "lime", "apricot", "coconut", "blueberry", "raspberry", "strawberry"] composites = [ compose_reward_single(f"{w}", "Rome", 100, cfg_reward, llm_call_count=0)[0] for w in words ] assert max(composites) <= 0.25, f"Format-only composites: {composites}" @pytest.mark.reward_audit def test_no_monotonic_length_exploit(cfg_reward): """REW-06: longer completions must NOT score higher (length-exploit guard).""" lengths = [10, 50, 100, 500, 1000, 2000, 5000, 8000, 12000, 20000] composites = [] for n in lengths: completion = " ".join(["filler"] * n) + " rome" score, _ = compose_reward_single(completion, "rome", 100, cfg_reward, llm_call_count=0) composites.append(score) # Longer completions must not score strictly higher assert composites[-1] < composites[0], ( f"Length exploit: shortest={composites[0]:.3f}, longest={composites[-1]:.3f}" ) # Sequence must be non-increasing (or flat) assert all(a >= b for a, b in zip(composites, composites[1:])), ( f"Non-monotonic sequence: {[round(x, 3) for x in composites]}" ) class TestComposeV3: """REW-02 v3: soft format + recursion-extraction + correctness-gated efficiency.""" @pytest.fixture def cfg_v3(self): return OmegaConf.create({ "alpha": 0.2, "weights": {"correctness": 0.70, "token_budget": 0.15, "recursion_efficiency": 0.15}, "token_budget_variant": "capped_linear", "answer_regex": "(.*?)", "max_calls": 4, }) def test_correct_no_recursion_scores_high(self, cfg_v3): c = "```python\nimport re\nm=re.search('silver', ctx)\nprint(m.group())\n```\nsilver" score, metrics = compose_reward_single(c, "silver", 100, cfg_v3) assert score >= 0.85, f"clean correct should score high, got {score}" assert metrics["llm_call_count"] == 0 def test_zero_calls_beats_one_call_when_both_correct(self, cfg_v3): c0 = "```python\nimport re\nm=re.search('silver', ctx)\nprint(m.group())\n```\nsilver" c1 = "```python\nans=llm('color', ctx[:5000])\nprint(ans)\n```\nsilver" s0, _ = compose_reward_single(c0, "silver", 100, cfg_v3) s1, _ = compose_reward_single(c1, "silver", 100, cfg_v3) assert s0 > s1, f"0-call ({s0:.3f}) should beat 1-call ({s1:.3f}) when both correct" def test_efficiency_gated_on_correctness(self, cfg_v3): # Wrong answer with 0 calls — must NOT earn efficiency bonus. c = "```python\nprint('done')\n```\ngold" score, metrics = compose_reward_single(c, "silver", 100, cfg_v3) assert metrics["recursion_eff_contribution"] == 0.0 assert score <= 0.25, f"wrong answer must be capped, got {score}" def test_recursion_spam_loses_to_minimal_recursion(self, cfg_v3): c2 = "```python\na=llm('q1',ctx[:1000])\nb=llm('q2',ctx[1000:2000])\n```\nsilver" c5 = "```python\n" + "\n".join(f"x{i}=llm('q{i}',ctx)" for i in range(5)) + "\n```\nsilver" s2, _ = compose_reward_single(c2, "silver", 100, cfg_v3) s5, _ = compose_reward_single(c5, "silver", 100, cfg_v3) assert s2 > s5, f"2-call ({s2:.3f}) should beat 5-call spam ({s5:.3f})" def test_format_only_capped(self, cfg_v3): c = "wrong" score, _ = compose_reward_single(c, "silver", 100, cfg_v3) assert 0.05 <= score <= 0.25, f"format-only wrong should be in [0.05, 0.25], got {score}" def test_no_format_gets_minimal_credit(self, cfg_v3): c = "silver" # right text but no tag score, _ = compose_reward_single(c, "silver", 100, cfg_v3) assert score <= 0.20 def test_group_variance_nonzero(self, cfg_v3): """Smoke check: a synthetic GRPO group of 8 must produce non-zero std. v1 had std=0.0 across all groups, which zeroed the GRPO advantage.""" gens = [ "```python\nimport re\nm=re.search('silver',ctx)\nprint(m.group())\n```\nsilver", "```python\nans=llm('color',ctx[:5000])\nprint(ans)\n```\nsilver", "silver", "gold", "the color is silver", "", "```python\n" + "\n".join(f"x{i}=llm('q{i}',ctx)" for i in range(5)) + "\n```\nsilver", "silver.", ] scores = [compose_reward_single(g, "silver", 200, cfg_v3)[0] for g in gens] import statistics as _st assert _st.stdev(scores) > 0.10, f"group std too low: {_st.stdev(scores)}"