Skip to content

Commit b0b1901

Browse files
authored
MAINT: Adding REQUIRED_VALUE as a placeholder for apply_defaults (microsoft#1182)
1 parent a99c07a commit b0b1901

28 files changed

Lines changed: 308 additions & 90 deletions

pyrit/common/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
reset_default_values,
1111
get_global_default_values,
1212
DefaultValueScope,
13+
REQUIRED_VALUE,
1314
)
1415
from pyrit.common.data_url_converter import convert_local_image_to_data_url
1516
from pyrit.common.default_values import get_non_required_value, get_required_value
@@ -53,6 +54,7 @@
5354
"is_in_ipython_session",
5455
"make_request_and_raise_if_error_async",
5556
"print_chat_messages_with_color",
57+
"REQUIRED_VALUE",
5658
"reset_default_values",
5759
"set_default_value",
5860
"Singleton",

pyrit/common/apply_defaults.py

Lines changed: 46 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,34 @@
2020
T = TypeVar("T")
2121

2222

23+
class _RequiredValueSentinel:
24+
"""
25+
Sentinel type to mark parameters as required but eligible for apply_defaults.
26+
27+
This allows parameters to have no default value in the function signature
28+
(making them appear required), while still allowing apply_defaults to provide
29+
a value if one is registered globally.
30+
31+
Usage:
32+
@apply_defaults
33+
def __init__(self, *, objective_target: PromptTarget = REQUIRED_VALUE):
34+
# If apply_defaults finds a default for objective_target, it will be used.
35+
# Otherwise, an error will be raised for missing required parameter.
36+
pass
37+
"""
38+
39+
def __repr__(self) -> str:
40+
return "REQUIRED_VALUE"
41+
42+
def __bool__(self) -> bool:
43+
# Ensure this evaluates to False in boolean context
44+
return False
45+
46+
47+
# Global sentinel instance
48+
REQUIRED_VALUE = _RequiredValueSentinel()
49+
50+
2351
@dataclass(frozen=True)
2452
class DefaultValueScope:
2553
"""
@@ -218,20 +246,35 @@ def wrapper(self, *args, **kwargs):
218246
bound_args = sig.bind(self, *args, **kwargs)
219247
bound_args.apply_defaults()
220248

221-
# Apply default values for parameters that are None
249+
# Apply default values for parameters that are None or REQUIRED_VALUE
222250
for param_name, param_value in bound_args.arguments.items():
223251
if param_name == "self":
224252
continue
225253

226-
# Only apply defaults if the parameter is None
227-
if param_value is None:
254+
# Apply defaults if parameter is None or REQUIRED_VALUE sentinel
255+
if param_value is None or isinstance(param_value, _RequiredValueSentinel):
228256
found, default_value = _global_default_values.get_default_value(
229257
class_type=cls,
230258
parameter_name=param_name,
231259
)
232260
if found:
233261
bound_args.arguments[param_name] = default_value
234262
logger.debug(f"Applied default value for {cls.__name__}.{param_name} = {default_value}")
263+
elif isinstance(param_value, _RequiredValueSentinel):
264+
# REQUIRED_VALUE was used but no default found - raise clear error
265+
raise ValueError(
266+
f"{param_name} is required for {cls.__name__}. "
267+
f"Either pass the parameter explicitly or register a default using set_default_value()."
268+
)
269+
# If None was explicitly passed and parameter has REQUIRED_VALUE as default, also raise
270+
elif param_value is None:
271+
# Check if the parameter's default in the signature is REQUIRED_VALUE
272+
param_obj = sig.parameters.get(param_name)
273+
if param_obj and isinstance(param_obj.default, _RequiredValueSentinel):
274+
raise ValueError(
275+
f"{param_name} is required for {cls.__name__}. "
276+
f"Either pass a valid value or register a default using set_default_value()."
277+
)
235278

236279
# Call the original method with updated arguments
237280
return method(*bound_args.args, **bound_args.kwargs)

pyrit/executor/attack/multi_turn/crescendo.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
from pathlib import Path
88
from typing import Optional, Union
99

10-
from pyrit.common.apply_defaults import apply_defaults
10+
from pyrit.common.apply_defaults import REQUIRED_VALUE, apply_defaults
1111
from pyrit.common.path import DATASETS_PATH
1212
from pyrit.common.utils import combine_dict
1313
from pyrit.exceptions import (
@@ -117,7 +117,7 @@ class CrescendoAttack(MultiTurnAttackStrategy[CrescendoAttackContext, CrescendoA
117117
def __init__(
118118
self,
119119
*,
120-
objective_target: PromptChatTarget,
120+
objective_target: PromptChatTarget = REQUIRED_VALUE, # type: ignore[assignment]
121121
attack_adversarial_config: AttackAdversarialConfig,
122122
attack_converter_config: Optional[AttackConverterConfig] = None,
123123
attack_scoring_config: Optional[AttackScoringConfig] = None,

pyrit/executor/attack/multi_turn/multi_prompt_sending.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
from dataclasses import dataclass, field
66
from typing import List, Optional
77

8-
from pyrit.common.apply_defaults import apply_defaults
8+
from pyrit.common.apply_defaults import REQUIRED_VALUE, apply_defaults
99
from pyrit.common.utils import combine_dict, get_kwarg_param
1010
from pyrit.executor.attack.component import ConversationManager
1111
from pyrit.executor.attack.core import (
@@ -66,7 +66,7 @@ class MultiPromptSendingAttack(MultiTurnAttackStrategy[MultiPromptSendingAttackC
6666
def __init__(
6767
self,
6868
*,
69-
objective_target: PromptTarget,
69+
objective_target: PromptTarget = REQUIRED_VALUE, # type: ignore[assignment]
7070
attack_converter_config: Optional[AttackConverterConfig] = None,
7171
attack_scoring_config: Optional[AttackScoringConfig] = None,
7272
prompt_normalizer: Optional[PromptNormalizer] = None,

pyrit/executor/attack/multi_turn/red_teaming.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
from pathlib import Path
99
from typing import Optional, Union
1010

11-
from pyrit.common.apply_defaults import apply_defaults
11+
from pyrit.common.apply_defaults import REQUIRED_VALUE, apply_defaults
1212
from pyrit.common.path import RED_TEAM_EXECUTOR_PATH
1313
from pyrit.common.utils import combine_dict, warn_if_set
1414
from pyrit.executor.attack.component import (
@@ -84,7 +84,7 @@ class RedTeamingAttack(MultiTurnAttackStrategy[MultiTurnAttackContext, AttackRes
8484
def __init__(
8585
self,
8686
*,
87-
objective_target: PromptTarget,
87+
objective_target: PromptTarget = REQUIRED_VALUE, # type: ignore[assignment]
8888
attack_adversarial_config: AttackAdversarialConfig,
8989
attack_converter_config: Optional[AttackConverterConfig] = None,
9090
attack_scoring_config: Optional[AttackScoringConfig] = None,

pyrit/executor/attack/multi_turn/tree_of_attacks.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111

1212
from treelib.tree import Tree
1313

14-
from pyrit.common.apply_defaults import apply_defaults
14+
from pyrit.common.apply_defaults import REQUIRED_VALUE, apply_defaults
1515
from pyrit.common.path import DATASETS_PATH
1616
from pyrit.common.utils import combine_dict, warn_if_set
1717
from pyrit.exceptions import (
@@ -934,7 +934,7 @@ class TreeOfAttacksWithPruningAttack(AttackStrategy[TAPAttackContext, TAPAttackR
934934
def __init__(
935935
self,
936936
*,
937-
objective_target: PromptChatTarget,
937+
objective_target: PromptChatTarget = REQUIRED_VALUE, # type: ignore[assignment]
938938
attack_adversarial_config: AttackAdversarialConfig,
939939
attack_converter_config: Optional[AttackConverterConfig] = None,
940940
attack_scoring_config: Optional[AttackScoringConfig] = None,

pyrit/executor/attack/single_turn/context_compliance.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
from pathlib import Path
66
from typing import Optional
77

8-
from pyrit.common.apply_defaults import apply_defaults
8+
from pyrit.common.apply_defaults import REQUIRED_VALUE, apply_defaults
99
from pyrit.common.path import DATASETS_PATH
1010
from pyrit.executor.attack.core import (
1111
AttackAdversarialConfig,
@@ -54,7 +54,7 @@ class ContextComplianceAttack(PromptSendingAttack):
5454
def __init__(
5555
self,
5656
*,
57-
objective_target: PromptChatTarget,
57+
objective_target: PromptChatTarget = REQUIRED_VALUE, # type: ignore[assignment]
5858
attack_adversarial_config: AttackAdversarialConfig,
5959
attack_converter_config: Optional[AttackConverterConfig] = None,
6060
attack_scoring_config: Optional[AttackScoringConfig] = None,

pyrit/executor/attack/single_turn/flip_attack.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66
import uuid
77
from typing import Optional
88

9-
from pyrit.common.apply_defaults import apply_defaults
9+
from pyrit.common.apply_defaults import REQUIRED_VALUE, apply_defaults
1010
from pyrit.common.path import DATASETS_PATH
1111
from pyrit.common.utils import combine_dict
1212
from pyrit.executor.attack.core import AttackConverterConfig, AttackScoringConfig
@@ -38,7 +38,7 @@ class FlipAttack(PromptSendingAttack):
3838
@apply_defaults
3939
def __init__(
4040
self,
41-
objective_target: PromptChatTarget,
41+
objective_target: PromptChatTarget = REQUIRED_VALUE, # type: ignore[assignment]
4242
attack_converter_config: Optional[AttackConverterConfig] = None,
4343
attack_scoring_config: Optional[AttackScoringConfig] = None,
4444
prompt_normalizer: Optional[PromptNormalizer] = None,

pyrit/executor/attack/single_turn/many_shot_jailbreak.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
import pathlib
66
from typing import Optional
77

8-
from pyrit.common.apply_defaults import apply_defaults
8+
from pyrit.common.apply_defaults import REQUIRED_VALUE, apply_defaults
99
from pyrit.common.path import DATASETS_PATH
1010
from pyrit.datasets import fetch_many_shot_jailbreaking_dataset
1111
from pyrit.executor.attack.core import AttackConverterConfig, AttackScoringConfig
@@ -31,7 +31,7 @@ class ManyShotJailbreakAttack(PromptSendingAttack):
3131
@apply_defaults
3232
def __init__(
3333
self,
34-
objective_target: PromptTarget,
34+
objective_target: PromptTarget = REQUIRED_VALUE, # type: ignore[assignment]
3535
attack_converter_config: Optional[AttackConverterConfig] = None,
3636
attack_scoring_config: Optional[AttackScoringConfig] = None,
3737
prompt_normalizer: Optional[PromptNormalizer] = None,

pyrit/executor/attack/single_turn/prompt_sending.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
import uuid
66
from typing import Optional
77

8-
from pyrit.common.apply_defaults import apply_defaults
8+
from pyrit.common.apply_defaults import REQUIRED_VALUE, apply_defaults
99
from pyrit.common.utils import combine_dict, warn_if_set
1010
from pyrit.executor.attack.component import ConversationManager
1111
from pyrit.executor.attack.core import AttackConverterConfig, AttackScoringConfig
@@ -53,7 +53,7 @@ class PromptSendingAttack(SingleTurnAttackStrategy):
5353
def __init__(
5454
self,
5555
*,
56-
objective_target: PromptTarget,
56+
objective_target: PromptTarget = REQUIRED_VALUE, # type: ignore[assignment]
5757
attack_converter_config: Optional[AttackConverterConfig] = None,
5858
attack_scoring_config: Optional[AttackScoringConfig] = None,
5959
prompt_normalizer: Optional[PromptNormalizer] = None,

0 commit comments

Comments
 (0)