"""Tests for template auto-detection and stripping (no external deps).""" from __future__ import annotations from sdg.preprocessing.dedupe.templates import ( find_common_prefixes, strip_template, template_hit_stats, ) # ── find_common_prefixes ───────────────────────────────────────────────────── def test_find_common_prefixes_empty_input(): assert find_common_prefixes([], min_count=2) == [] def test_find_common_prefixes_returns_empty_when_no_common_prefix(): texts = ["alpha question one", "beta question two", "gamma question three"] assert find_common_prefixes(texts, min_count=2, prefix_lengths=(5, 10)) == [] def test_find_common_prefixes_detects_simple_repeated_prefix(): template = "Answer the question: " texts = [template + f"unique body number {i}" for i in range(20)] found = find_common_prefixes(texts, min_count=10, prefix_lengths=(len(template),)) assert template in found def test_find_common_prefixes_respects_min_count(): rare = "Rare prefix: " common = "Common prefix: " texts = ( [rare + f"a{i}" for i in range(5)] + [common + f"b{i}" for i in range(20)] ) found = find_common_prefixes(texts, min_count=10, prefix_lengths=(len(common),)) assert common in found assert rare not in found def test_find_common_prefixes_returns_sorted_longest_first(): long_tmpl = "A" * 100 + " specific" short_tmpl = "A" * 50 texts = [long_tmpl + " body " + str(i) for i in range(20)] found = find_common_prefixes(texts, min_count=10, prefix_lengths=(50, 100, 150)) assert found == sorted(found, key=len, reverse=True) def test_find_common_prefixes_handles_short_texts(): """Texts shorter than a probed prefix length must not crash.""" texts = ["xy"] * 20 # all shorter than every probed length found = find_common_prefixes(texts, min_count=2, prefix_lengths=(50, 100)) assert found == [] def test_find_common_prefixes_finds_multiple_distinct_templates(): t1 = "Template A: " t2 = "Different template B: " texts = ( [t1 + f"x{i}" for i in range(15)] + [t2 + f"y{i}" for i in range(15)] + [f"singleton {i}" for i in range(5)] ) found = find_common_prefixes(texts, min_count=10, prefix_lengths=(len(t1), len(t2))) assert t1 in found assert t2 in found # ── strip_template ─────────────────────────────────────────────────────────── def test_strip_template_returns_text_unchanged_with_no_match(): assert strip_template("hello world", ["FOO: "]) == "hello world" def test_strip_template_returns_text_unchanged_with_empty_templates(): assert strip_template("hello world", []) == "hello world" def test_strip_template_simple_match(): assert strip_template("PREFIX: actual content", ["PREFIX: "]) == "actual content" def test_strip_template_strips_leading_whitespace_after_removal(): assert strip_template("PREFIX: \n\ncontent", ["PREFIX:"]) == "content" def test_strip_template_longest_match_wins_when_sorted_correctly(): """Templates list sorted longest-first as find_common_prefixes returns.""" templates = ["PREFIX EXTRA: ", "PREFIX: "] assert strip_template("PREFIX EXTRA: content", templates) == "content" assert strip_template("PREFIX: content", templates) == "content" def test_strip_template_does_not_match_substring_in_middle(): """Template must occur as a *prefix*, not anywhere in the text.""" assert strip_template("hello PREFIX: world", ["PREFIX: "]) == "hello PREFIX: world" # ── template_hit_stats ─────────────────────────────────────────────────────── def test_template_hit_stats_partitions_texts_correctly(): texts = [ "MCQ: question 1", "MCQ: question 2", "MATH: question 3", "totally distinct prompt", ] templates = ["MCQ: ", "MATH: "] stats = template_hit_stats(texts, templates) assert stats["MCQ: "] == 2 assert stats["MATH: "] == 1 assert stats[""] == 1 assert sum(stats.values()) == len(texts) def test_template_hit_stats_uses_longest_match_when_templates_nest(): """If templates list has nested prefixes, longest-first prevents double counting.""" texts = ["PREFIX EXTRA: x"] * 10 + ["PREFIX: y"] * 5 templates = ["PREFIX EXTRA: ", "PREFIX: "] stats = template_hit_stats(texts, templates) assert stats["PREFIX EXTRA: "] == 10 assert stats["PREFIX: "] == 5 assert stats[""] == 0 def test_template_hit_stats_empty_templates_all_unmatched(): texts = ["a", "b", "c"] stats = template_hit_stats(texts, []) assert stats == {"": 3} # ── End-to-end: detect then strip ──────────────────────────────────────────── def test_detect_then_strip_round_trip(): """Realistic sequence: detect templates, then strip them from the corpus.""" template = "Answer the following multiple choice question: " questions = [ "What is 2+2?", "What color is the sky?", "Who wrote Hamlet?", "What is the capital of France?", "What is the speed of light?", ] * 4 # 20 questions texts = [template + q for q in questions] templates = find_common_prefixes(texts, min_count=10, prefix_lengths=(len(template),)) stripped = [strip_template(t, templates) for t in texts] # Every text should have been stripped down to just its question. assert all(s in [q for q in questions] for s in stripped)