Files
OBLITERATUS/tests/test_prompt_sources.py
T

239 lines
8.2 KiB
Python

"""Offline tests for prompt registry parsing and external schema boundaries."""
from __future__ import annotations
import json
import sys
from types import SimpleNamespace
import pytest
from obliteratus import prompts
@pytest.fixture(autouse=True)
def clean_cache():
prompts.clear_dataset_cache()
yield
prompts.clear_dataset_cache()
def _datasets(monkeypatch, rows_or_error):
def load_dataset(*_args, **_kwargs):
if isinstance(rows_or_error, Exception):
raise rows_or_error
return rows_or_error
monkeypatch.setitem(sys.modules, "datasets", SimpleNamespace(load_dataset=load_dataset))
def test_cache_copies_results_and_builtin_loader_isolated():
calls = []
def loader():
calls.append(True)
return ["harm"], ["safe"]
first = prompts._cached_load("fixture", loader)
first[0].append("mutation")
assert prompts._cached_load("fixture", loader) == (["harm"], ["safe"])
assert len(calls) == 1
harmful, harmless = prompts._load_builtin()
assert len(harmful) == len(harmless) == 842
@pytest.mark.parametrize(
("loader", "rows", "expected"),
[
(prompts._load_harmbench, [{"Behavior": "A sufficiently long behavior"}], "A sufficiently long behavior"),
(prompts._load_advbench, [{"goal": "A sufficiently long goal prompt"}], "A sufficiently long goal prompt"),
],
)
def test_single_column_external_sources(monkeypatch, loader, rows, expected):
_datasets(monkeypatch, rows)
harmful, harmless = loader()
assert harmful == [expected]
assert len(harmless) == 1
def test_anthropic_parser_deduplicates_and_uses_fallback(monkeypatch):
calls = []
def load_dataset(*_args, **kwargs):
calls.append(kwargs)
if "data_dir" in kwargs:
raise RuntimeError("primary unavailable")
return [
{"chosen": "Human: A unique and sufficiently long prompt Assistant: response"},
{"rejected": "Human: A unique and sufficiently long prompt Assistant: other"},
{"chosen": "no conversation markers"},
]
monkeypatch.setitem(sys.modules, "datasets", SimpleNamespace(load_dataset=load_dataset))
harmful, harmless = prompts._load_anthropic_redteam()
assert harmful == ["A unique and sufficiently long prompt"]
assert len(harmless) == 1
assert len(calls) == 2
def test_wildjailbreak_requires_pairs_and_deduplicates(monkeypatch):
_datasets(monkeypatch, [
{"adversarial_query": "A sufficiently long adversarial prompt", "vanilla_query": "safe"},
{"adversarial": "A sufficiently long adversarial prompt", "vanilla": "duplicate"},
{"adversarial": "missing pair"},
])
harmful, harmless = prompts._load_wildjailbreak()
assert harmful == ["A sufficiently long adversarial prompt"]
assert harmless == ["safe"]
@pytest.mark.parametrize(
"loader",
[prompts._load_harmbench, prompts._load_advbench, prompts._load_wildjailbreak],
)
def test_external_loaders_explain_empty_schemas(monkeypatch, loader):
_datasets(monkeypatch, [{"unexpected": "value"}])
with pytest.raises(RuntimeError, match="0 prompts extracted"):
loader()
def test_anthropic_empty_parse_is_rejected(monkeypatch):
_datasets(monkeypatch, [{"chosen": "not a conversation"}])
with pytest.raises(RuntimeError, match="0 prompts extracted"):
prompts._load_anthropic_redteam()
def test_custom_prompt_validation_padding_and_registry_access():
harmful_text = "\n".join(f"harm {index}" for index in range(5))
with pytest.raises(ValueError, match="at least 5"):
prompts.load_custom_prompts("only one", "safe")
harmful, harmless = prompts.load_custom_prompts(harmful_text, "")
assert len(harmful) == len(harmless) == 5
harmful, harmless = prompts.load_custom_prompts(harmful_text, "safe one\nsafe two")
assert harmless[:2] == ["safe one", "safe two"]
assert len(harmless) == 5
with pytest.raises(ValueError, match="Unknown dataset source"):
prompts.load_dataset_source("missing")
assert prompts.get_source_key_from_label("missing label") == "builtin"
assert prompts.get_source_key_from_label(prompts.DATASET_SOURCES["advbench"].label) == "advbench"
assert len(prompts.get_source_choices()) == len(prompts.DATASET_SOURCES)
assert prompts.get_valid_volumes("missing") == ["all (use entire dataset)"]
assert prompts.get_valid_volumes("harmbench")[-1] == "all (use entire dataset)"
def test_prompt_pairs_file_loader_accepts_exact_schema(tmp_path):
path = tmp_path / "pairs.json"
path.write_text(
"""{
"harmful": ["harm 0", "harm 1", "harm 2", "harm 3", "harm 4"],
"harmless": ["safe 0", "safe 1", "safe 2", "safe 3", "safe 4"]
}""",
encoding="utf-8",
)
harmful, harmless = prompts.load_prompt_pairs_file(path)
assert harmful == [f"harm {index}" for index in range(5)]
assert harmless == [f"safe {index}" for index in range(5)]
@pytest.mark.parametrize(
("payload", "message"),
[
("not json", "malformed JSON"),
("[]", "JSON object"),
('{"harmful": ["h"], "harmless": ["s"], "extra": []}', "exactly"),
('{"harmful": "h", "harmless": ["s"]}', "must be an array"),
(
'{"harmful": ["h0", "h1", "h2", "h3", 4], '
'"harmless": ["s0", "s1", "s2", "s3", "s4"]}',
"must be a string",
),
('{"harmful": ["h"], "harmless": ["s"]}', "at least 5"),
(
'{"harmful": ["h0", "h1", "h2", "h3", "h4"], '
'"harmless": ["s0", "s1", "s2", "s3"]}',
"equal length",
),
(
'{"harmful": ["h0", "h1", "h2", "h3", "h4"], '
'"harmless": ["s0", "s1", "s2", "s3", ""]}',
"nonblank strings",
),
(
'{"harmful": ["h0", "h1", "h2", "h3", "h4"], '
'"harmless": ["s0", "s1", "s2", "s3", "s\\u0000"]}',
"NUL",
),
],
)
def test_prompt_pairs_file_loader_rejects_invalid_schema(tmp_path, payload, message):
path = tmp_path / "pairs.json"
path.write_text(payload, encoding="utf-8")
with pytest.raises(ValueError, match=message):
prompts.load_prompt_pairs_file(path)
def test_prompt_pairs_file_loader_rejects_missing_non_regular_oversized_and_utf8(tmp_path):
with pytest.raises(ValueError, match="does not exist"):
prompts.load_prompt_pairs_file(tmp_path / "missing.json")
with pytest.raises(ValueError, match="regular file"):
prompts.load_prompt_pairs_file(tmp_path)
oversized = tmp_path / "oversized.json"
oversized.write_bytes(b" " * (prompts.MAX_PROMPT_PAIRS_FILE_BYTES + 1))
with pytest.raises(ValueError, match="too large"):
prompts.load_prompt_pairs_file(oversized)
invalid_utf8 = tmp_path / "invalid.json"
invalid_utf8.write_bytes(b"\xff")
with pytest.raises(ValueError, match="UTF-8"):
prompts.load_prompt_pairs_file(invalid_utf8)
def test_prompt_pairs_file_loader_rejects_filesystem_errors_and_pair_limit(
tmp_path,
monkeypatch,
):
path = tmp_path / "pairs.json"
path.write_text('{"harmful": [], "harmless": []}', encoding="utf-8")
def stat_error(_self):
raise OSError("denied")
with monkeypatch.context() as m:
m.setattr(type(path), "stat", stat_error)
with pytest.raises(ValueError, match="Cannot access"):
prompts.load_prompt_pairs_file(path)
def read_error(*_args, **_kwargs):
raise OSError("denied")
with monkeypatch.context() as m:
m.setattr(type(path), "read_text", read_error)
with pytest.raises(ValueError, match="Cannot read"):
prompts.load_prompt_pairs_file(path)
monkeypatch.setattr(prompts, "MAX_PROMPT_PAIRS", 4)
path.write_text(
json.dumps(
{
"harmful": [f"harm {index}" for index in range(5)],
"harmless": [f"safe {index}" for index in range(5)],
}
),
encoding="utf-8",
)
with pytest.raises(ValueError, match="at most"):
prompts.load_prompt_pairs_file(path)
def test_harmless_generator_cycles_deterministically():
count = len(prompts._HARMLESS_POOL) + 1
generated = prompts._generate_harmless_counterparts(count)
assert generated[0] == generated[-1]