Skip to content

Commit 190ee4a

Browse files
author
Victor Valbuena
committed
FEAT: Register ADVERSARIAL_CHAT singleturn/multiturn/reasoning variants
Adds three TargetConfig entries (singleturn, multiturn, reasoning), each tagged [DEFAULT, ADVERSARIAL], for the env-driven variants already declared in .env_example. Tightens _register_target to skip with a warning when a TargetConfig declares model_var but the env var is unset; without this guard the target silently falls back to the global OPENAI_CHAT_MODEL default and sends requests to the wrong model. New tests covering naming-related failure modes flagged for the eventual PR review: - test_register_instance_with_duplicate_name_silently_overwrites pins the current "second write wins" behavior so future hardening (warn / raise / idempotent skip) is intentional. - test_target_configs_have_unique_registry_names guards against typos in ENV_TARGET_CONFIGS that would otherwise silently drop a target. - test_double_initialize_async_is_idempotent regression-guards the re-init path that depends on the silent-overwrite semantics above. - test_variant_skips_when_model_env_var_missing parameterizes the missing-_MODEL skip+warning for all three new variants. Failure modes surfaced during this change but not addressed here (tracked for the PR description batch): - Duplicate registry_name silently overwrites in BaseInstanceRegistry. - registry_name has no format validation; risk grows with per-user TargetConfig support in P1. - No-adversarial-models-found error message UX is owned by the upcoming BenchmarkInitializer commit and needs a clear, actionable message.
1 parent b7af0de commit 190ee4a

3 files changed

Lines changed: 220 additions & 0 deletions

File tree

pyrit/setup/initializers/components/targets.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -189,6 +189,33 @@ class TargetConfig:
189189
temperature=1.2,
190190
tags=[TargetInitializerTags.DEFAULT, TargetInitializerTags.ADVERSARIAL],
191191
),
192+
TargetConfig(
193+
registry_name="adversarial_chat_singleturn",
194+
target_class=OpenAIChatTarget,
195+
endpoint_var="ADVERSARIAL_CHAT_SINGLETURN_ENDPOINT",
196+
key_var="ADVERSARIAL_CHAT_SINGLETURN_KEY",
197+
model_var="ADVERSARIAL_CHAT_SINGLETURN_MODEL",
198+
temperature=1.2,
199+
tags=[TargetInitializerTags.DEFAULT, TargetInitializerTags.ADVERSARIAL],
200+
),
201+
TargetConfig(
202+
registry_name="adversarial_chat_multiturn",
203+
target_class=OpenAIChatTarget,
204+
endpoint_var="ADVERSARIAL_CHAT_MULTITURN_ENDPOINT",
205+
key_var="ADVERSARIAL_CHAT_MULTITURN_KEY",
206+
model_var="ADVERSARIAL_CHAT_MULTITURN_MODEL",
207+
temperature=1.2,
208+
tags=[TargetInitializerTags.DEFAULT, TargetInitializerTags.ADVERSARIAL],
209+
),
210+
TargetConfig(
211+
registry_name="adversarial_chat_reasoning",
212+
target_class=OpenAIChatTarget,
213+
endpoint_var="ADVERSARIAL_CHAT_REASONING_ENDPOINT",
214+
key_var="ADVERSARIAL_CHAT_REASONING_KEY",
215+
model_var="ADVERSARIAL_CHAT_REASONING_MODEL",
216+
temperature=1.2,
217+
tags=[TargetInitializerTags.DEFAULT, TargetInitializerTags.ADVERSARIAL],
218+
),
192219
TargetConfig(
193220
registry_name="objective_scorer_chat",
194221
target_class=OpenAIChatTarget,
@@ -573,6 +600,18 @@ def _register_target(self, config: TargetConfig) -> None:
573600
model_name = os.getenv(config.model_var) if config.model_var else None
574601
underlying_model = os.getenv(config.underlying_model_var) if config.underlying_model_var else None
575602

603+
# Guard against silent fallback to a global OPENAI_CHAT_MODEL default when the
604+
# declared per-config model env var is unset. Without this skip, the target
605+
# registers cleanly but sends requests to the wrong model at runtime.
606+
if config.model_var and not model_name:
607+
logger.warning(
608+
"Skipping target '%s': %s is not set. "
609+
"All declared env vars (endpoint, key, model) must be present for this target to register.",
610+
config.registry_name,
611+
config.model_var,
612+
)
613+
return
614+
576615
# Build kwargs for the target constructor
577616
kwargs: dict[str, Any] = {
578617
"endpoint": endpoint,

tests/unit/registry/test_target_registry.py

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -138,6 +138,25 @@ def test_register_instance_same_target_type_different_config(self):
138138

139139
assert len(self.registry) == 2
140140

141+
def test_register_instance_with_duplicate_name_silently_overwrites(self):
142+
"""Characterization: re-registering an existing name silently replaces the prior entry.
143+
144+
BaseInstanceRegistry.register is plain dict assignment; there is no
145+
collision check, warning, or error. This test pins the current behavior
146+
so any future tightening (warn, raise, idempotent skip) is an
147+
intentional decision rather than a silent regression. Tracked as
148+
``duplicate-registry-name`` in failure_mode_followups for the PR
149+
review batch.
150+
"""
151+
first = MockPromptTarget(model_name="first")
152+
second = MockPromptTarget(model_name="second")
153+
154+
self.registry.register_instance(first, name="same_name")
155+
self.registry.register_instance(second, name="same_name")
156+
157+
assert len(self.registry) == 1
158+
assert self.registry.get("same_name") is second
159+
141160

142161
@pytest.mark.usefixtures("patch_central_database")
143162
class TestTargetRegistryGetInstanceByName:

tests/unit/setup/test_targets_initializer.py

Lines changed: 162 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -211,6 +211,23 @@ def test_expected_targets_in_configs(self):
211211
assert "groq" in registry_names
212212
assert "google_gemini" in registry_names
213213

214+
def test_target_configs_have_unique_registry_names(self):
215+
"""Guard against typos: every ``registry_name`` in ``ENV_TARGET_CONFIGS`` must be unique.
216+
217+
Duplicate names would silently overwrite each other when
218+
``TargetInitializer`` registers them (per ``BaseInstanceRegistry.register``
219+
semantics, characterized in ``test_target_registry.py``). Only the
220+
second entry would survive in the registry, which breaks downstream
221+
fan-out (``BenchmarkInitializer``) and is hard to diagnose. Tracked
222+
as ``duplicate-registry-name`` in failure_mode_followups.
223+
"""
224+
registry_names = [config.registry_name for config in TARGET_CONFIGS]
225+
seen: dict[str, int] = {}
226+
for name in registry_names:
227+
seen[name] = seen.get(name, 0) + 1
228+
duplicates = {name: count for name, count in seen.items() if count > 1}
229+
assert not duplicates, f"Duplicate registry_name(s) in TARGET_CONFIGS: {duplicates}"
230+
214231

215232
class TestTargetInitializerGetInfo:
216233
"""Tests for TargetInitializer.get_info_async method."""
@@ -500,3 +517,148 @@ async def test_register_target_default_objective_tag_still_applied(self) -> None
500517
assert any(entry.name == "openai_chat" for entry in default_entries), (
501518
"openai_chat's config.tags=[DEFAULT] must propagate even when default_objective_target=True"
502519
)
520+
521+
522+
ADVERSARIAL_CHAT_VARIANTS: list[tuple[str, str]] = [
523+
("adversarial_chat_singleturn", "ADVERSARIAL_CHAT_SINGLETURN"),
524+
("adversarial_chat_multiturn", "ADVERSARIAL_CHAT_MULTITURN"),
525+
("adversarial_chat_reasoning", "ADVERSARIAL_CHAT_REASONING"),
526+
]
527+
528+
529+
@pytest.mark.usefixtures("patch_central_database")
530+
class TestTargetInitializerAdversarialChatVariants:
531+
"""Tests for the ``ADVERSARIAL_CHAT_{SINGLETURN,MULTITURN,REASONING}_*`` env-driven variants."""
532+
533+
def setup_method(self) -> None:
534+
"""Reset registry and clear variant env vars."""
535+
TargetRegistry.reset_instance()
536+
self._clear_variant_env_vars()
537+
538+
def teardown_method(self) -> None:
539+
"""Reset registry and clear variant env vars."""
540+
TargetRegistry.reset_instance()
541+
self._clear_variant_env_vars()
542+
543+
@staticmethod
544+
def _clear_variant_env_vars() -> None:
545+
for _, prefix in ADVERSARIAL_CHAT_VARIANTS:
546+
for suffix in ("ENDPOINT", "KEY", "MODEL"):
547+
os.environ.pop(f"{prefix}_{suffix}", None)
548+
549+
@staticmethod
550+
def _set_variant_env_vars(prefix: str) -> None:
551+
os.environ[f"{prefix}_ENDPOINT"] = "https://variant.openai.azure.com/openai/v1"
552+
os.environ[f"{prefix}_KEY"] = "test_key"
553+
os.environ[f"{prefix}_MODEL"] = "deployment-name"
554+
555+
@pytest.mark.parametrize(("registry_name", "env_prefix"), ADVERSARIAL_CHAT_VARIANTS)
556+
async def test_variant_registers_with_default_and_adversarial_tags(
557+
self, registry_name: str, env_prefix: str
558+
) -> None:
559+
"""Each variant registers with ``[DEFAULT, ADVERSARIAL]`` tags when its env vars are set."""
560+
from pyrit.setup.initializers.components.targets import TargetInitializerTags
561+
562+
self._set_variant_env_vars(env_prefix)
563+
564+
init = TargetInitializer()
565+
await init.initialize_async()
566+
567+
registry = TargetRegistry.get_registry_singleton()
568+
assert registry_name in registry
569+
570+
adversarial_entries = registry.get_by_tag(tag=TargetInitializerTags.ADVERSARIAL)
571+
assert any(entry.name == registry_name for entry in adversarial_entries)
572+
573+
default_entries = registry.get_by_tag(tag=TargetInitializerTags.DEFAULT)
574+
assert any(entry.name == registry_name for entry in default_entries)
575+
576+
@pytest.mark.parametrize(("registry_name", "env_prefix"), ADVERSARIAL_CHAT_VARIANTS)
577+
async def test_variant_skips_when_env_vars_missing(self, registry_name: str, env_prefix: str) -> None:
578+
"""Variants skip gracefully when their env vars are missing (matches existing adversarial_chat behavior)."""
579+
init = TargetInitializer()
580+
await init.initialize_async()
581+
582+
registry = TargetRegistry.get_registry_singleton()
583+
assert registry_name not in registry
584+
585+
@pytest.mark.parametrize(("registry_name", "env_prefix"), ADVERSARIAL_CHAT_VARIANTS)
586+
async def test_variant_skips_when_model_env_var_missing(
587+
self, registry_name: str, env_prefix: str, caplog: pytest.LogCaptureFixture
588+
) -> None:
589+
"""Endpoint+key set but _MODEL unset must skip with a warning, not silently fall back to OPENAI_CHAT_MODEL."""
590+
import logging
591+
592+
os.environ[f"{env_prefix}_ENDPOINT"] = "https://variant.openai.azure.com/openai/v1"
593+
os.environ[f"{env_prefix}_KEY"] = "test_key"
594+
595+
try:
596+
with caplog.at_level(logging.WARNING, logger="pyrit.setup.initializers.components.targets"):
597+
init = TargetInitializer()
598+
await init.initialize_async()
599+
600+
registry = TargetRegistry.get_registry_singleton()
601+
assert registry_name not in registry
602+
603+
captured_messages = [r.message for r in caplog.records]
604+
assert any(f"{env_prefix}_MODEL" in m for m in captured_messages), (
605+
f"Expected a warning naming the missing {env_prefix}_MODEL env var; got: {captured_messages}"
606+
)
607+
finally:
608+
os.environ.pop(f"{env_prefix}_ENDPOINT", None)
609+
os.environ.pop(f"{env_prefix}_KEY", None)
610+
611+
async def test_all_variants_discoverable_via_adversarial_tag_query(self) -> None:
612+
"""End-to-end: variants + ``adversarial_chat`` are returned by adversarial-tag ``get_by_tag_query``."""
613+
from pyrit.registry.tag_query import TagQuery
614+
615+
os.environ["ADVERSARIAL_CHAT_ENDPOINT"] = "https://parent.openai.azure.com/openai/v1"
616+
os.environ["ADVERSARIAL_CHAT_KEY"] = "test_key"
617+
os.environ["ADVERSARIAL_CHAT_MODEL"] = "deployment-name"
618+
619+
for _, prefix in ADVERSARIAL_CHAT_VARIANTS:
620+
self._set_variant_env_vars(prefix)
621+
622+
try:
623+
init = TargetInitializer()
624+
await init.initialize_async()
625+
626+
registry = TargetRegistry.get_registry_singleton()
627+
matches = registry.get_by_tag_query(query=TagQuery.all("adversarial"))
628+
match_names = {entry.name for entry in matches}
629+
630+
expected = {"adversarial_chat"} | {name for name, _ in ADVERSARIAL_CHAT_VARIANTS}
631+
assert expected <= match_names, (
632+
f"Missing variants from tag query result. Expected superset: {expected}, got: {match_names}"
633+
)
634+
finally:
635+
for var in ("ADVERSARIAL_CHAT_ENDPOINT", "ADVERSARIAL_CHAT_KEY", "ADVERSARIAL_CHAT_MODEL"):
636+
os.environ.pop(var, None)
637+
638+
async def test_double_initialize_async_is_idempotent(self) -> None:
639+
"""Re-running ``initialize_async`` with the same env state produces the same registry contents.
640+
641+
Regression guard for the duplicate-registration silent-overwrite path:
642+
because env vars haven't changed between calls, the rebuilt entries
643+
carry identical configuration. If anyone introduces non-idempotent
644+
side-effects (e.g. tag accumulation, instance leaks) into
645+
``_register_target``, this test will catch it. Tracked as
646+
``duplicate-registry-name`` in failure_mode_followups.
647+
"""
648+
from pyrit.setup.initializers.components.targets import TargetInitializerTags
649+
650+
for _, prefix in ADVERSARIAL_CHAT_VARIANTS:
651+
self._set_variant_env_vars(prefix)
652+
653+
init = TargetInitializer()
654+
await init.initialize_async()
655+
registry = TargetRegistry.get_registry_singleton()
656+
first_names = sorted(registry.get_names())
657+
first_adversarial_count = len(registry.get_by_tag(tag=TargetInitializerTags.ADVERSARIAL))
658+
659+
await init.initialize_async()
660+
second_names = sorted(registry.get_names())
661+
second_adversarial_count = len(registry.get_by_tag(tag=TargetInitializerTags.ADVERSARIAL))
662+
663+
assert first_names == second_names
664+
assert first_adversarial_count == second_adversarial_count

0 commit comments

Comments
 (0)