diff --git a/pyrit/common/lazy_imports.py b/pyrit/common/lazy_imports.py new file mode 100644 index 0000000000..07fe09f963 --- /dev/null +++ b/pyrit/common/lazy_imports.py @@ -0,0 +1,63 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Helpers for resolving public package exports on first access.""" + +from collections.abc import Mapping +from importlib import import_module +from typing import Any + +LazyExport = str | tuple[str, str | None] + + +def resolve_lazy_export( + *, + name: str, + module_name: str, + module_globals: dict[str, Any], + exports: Mapping[str, LazyExport], +) -> Any: + """ + Resolve and cache one package export. + + Args: + name (str): Public attribute requested from the package. + module_name (str): Package name used in an ``AttributeError``. + module_globals (dict[str, Any]): Package globals where the resolved value is cached. + exports (Mapping[str, LazyExport]): Public names and their implementation locations. + + Returns: + Any: The resolved public value. + + Raises: + AttributeError: If ``name`` is not a declared lazy export. + """ + try: + export = exports[name] + except KeyError: + raise AttributeError(f"module {module_name!r} has no attribute {name!r}") from None + + if isinstance(export, str): + target_module = export + target_attribute: str | None = name + else: + target_module, target_attribute = export + + imported_module = import_module(target_module) + value = imported_module if target_attribute is None else getattr(imported_module, target_attribute) + module_globals[name] = value + return value + + +def get_lazy_dir(*, module_globals: dict[str, Any], exports: Mapping[str, LazyExport]) -> list[str]: + """ + Return package attributes and unresolved public exports. + + Args: + module_globals (dict[str, Any]): The package globals. + exports (Mapping[str, LazyExport]): Public names and their implementation locations. + + Returns: + list[str]: Sorted package attribute names. + """ + return sorted(module_globals.keys() | exports.keys()) diff --git a/pyrit/models/__init__.py b/pyrit/models/__init__.py index 667a2ca033..4a6dfc36ca 100644 --- a/pyrit/models/__init__.py +++ b/pyrit/models/__init__.py @@ -1,5 +1,6 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +# ruff: noqa: F401 """ Public model exports for PyRIT core data structures and helpers. @@ -16,224 +17,261 @@ ComponentIdentifier``). """ -from pyrit.models.additional_initializer import AdditionalInitializer -from pyrit.models.conversation_stats import ConversationStats -from pyrit.models.embeddings import EmbeddingData, EmbeddingResponse, EmbeddingSupport, EmbeddingUsageInformation -from pyrit.models.harm_definition import HarmDefinition, ScaleDescription, get_all_harm_definitions -from pyrit.models.identifiers import ( - REGISTRY_NAME_PATTERN, - TARGET_EVAL_PARAM_FALLBACKS, - TARGET_EVAL_PARAMS, - AtomicAttackEvaluationIdentifier, - AtomicAttackIdentifier, - AttackIdentifier, - AttackTechniqueIdentifier, - ChildEvalRule, - ComponentIdentifier, - ConverterIdentifier, - Evaluate, - EvaluationIdentifier, - Identifiable, - IdentifierFilter, - IdentifierType, - JSONValue, - ObjectiveTargetEvaluationIdentifier, - ScenarioEvaluationIdentifier, - ScenarioIdentifier, - ScorerEvaluationIdentifier, - ScorerIdentifier, - SeedIdentifier, - TargetIdentifier, - class_name_to_snake_case, - compute_eval_hash, - config_hash, - snake_case_to_class_name, - validate_registry_name, -) -from pyrit.models.literals import ( - MEDIA_PATH_DATA_TYPES, - ChatMessageRole, - Modality, - PromptDataType, - PromptResponseError, - SeedType, -) -from pyrit.models.messages import ( - Conversation, - Message, - MessagePiece, - construct_response_from_request, - flatten_to_message_pieces, - get_all_values, - group_conversation_message_pieces_by_sequence, - group_message_pieces_into_conversations, - sort_message_pieces, -) -from pyrit.models.messages.chat_message import ALLOWED_CHAT_MESSAGE_ROLES, ChatMessage, ChatMessagesDataset, ToolCall -from pyrit.models.messages.conversation_reference import ConversationReference, ConversationType -from pyrit.models.messages.conversation_retry import ConversationRetry, ConversationRetryReason -from pyrit.models.parameter import ComponentType, Parameter, ParameterDestination, RegistryReference, display_choices -from pyrit.models.question_answering import QuestionAnsweringDataset, QuestionAnsweringEntry, QuestionChoice -from pyrit.models.results.attack_result import AttackOutcome, AttackResult, AttackResultT -from pyrit.models.results.scenario_result import ScenarioResult, ScenarioRunState -from pyrit.models.results.strategy_result import StrategyResult, StrategyResultT -from pyrit.models.retry_event import RetryEvent -from pyrit.models.score import ( - Condition, - ContentScorable, - MatchesObjective, - MessageScorable, - Scorable, - Score, - ScoreType, - ScoringExpectation, - UnvalidatedScore, -) +from typing import TYPE_CHECKING -# Seeds - import from new seeds submodule for forward compatibility -# Also keep imports from old locations for backward compatibility -from pyrit.models.seeds import ( - AttackSeedGroup, - AttackTechniqueSeedGroup, - NextMessageSystemPromptPaths, - Seed, - SeedDataset, - SeedGroup, - SeedObjective, - SeedPrompt, - SeedSimulatedConversation, - SeedUnion, - SimulatedTargetSystemPromptPaths, - group_seeds_into_attack_groups, -) -from pyrit.models.target import ( - COMMON_JSON_SCHEMAS, - JSON_SCHEMA_METADATA_KEY, - SEED_RESPONSE_JSON_SCHEMA_METADATA_KEY, - TOKEN_USAGE_METADATA_PREFIX, - CapabilityName, - JsonResponseConfig, - JsonSchemaDefinition, - TargetCapabilities, - TokenUsage, - get_common_json_schema, - read_usage_int, - read_usage_value, - register_common_json_schema, - unregister_common_json_schema, -) +from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export -__all__ = [ - "ALLOWED_CHAT_MESSAGE_ROLES", - "AdditionalInitializer", - "AtomicAttackEvaluationIdentifier", - "AtomicAttackIdentifier", - "AttackIdentifier", - "AttackTechniqueIdentifier", - "AttackResult", - "AttackResultT", - "AttackOutcome", - "ChatMessage", - "ChatMessagesDataset", - "ChatMessageRole", - "ChildEvalRule", - "class_name_to_snake_case", - "CapabilityName", - "ComponentIdentifier", - "ComponentType", - "compute_eval_hash", - "Condition", - "config_hash", - "ConverterIdentifier", - "Conversation", - "ConversationReference", - "ConversationRetry", - "ConversationRetryReason", - "ConversationStats", - "ConversationType", - "ContentScorable", - "construct_response_from_request", - "display_choices", - "EmbeddingData", - "EmbeddingResponse", - "EmbeddingSupport", - "EmbeddingUsageInformation", - "Evaluate", - "EvaluationIdentifier", - "flatten_to_message_pieces", - "get_all_harm_definitions", - "get_all_values", - "group_conversation_message_pieces_by_sequence", - "group_message_pieces_into_conversations", - "group_seeds_into_attack_groups", - "HarmDefinition", - "Identifiable", - "IdentifierFilter", - "IdentifierType", - "JSONValue", - "COMMON_JSON_SCHEMAS", - "JsonResponseConfig", - "get_common_json_schema", - "register_common_json_schema", - "unregister_common_json_schema", - "JSON_SCHEMA_METADATA_KEY", - "SEED_RESPONSE_JSON_SCHEMA_METADATA_KEY", - "JsonSchemaDefinition", - "MatchesObjective", - "MEDIA_PATH_DATA_TYPES", - "Message", - "MessagePiece", - "MessageScorable", - "Modality", - "NextMessageSystemPromptPaths", - "ObjectiveTargetEvaluationIdentifier", - "Parameter", - "ParameterDestination", - "PromptDataType", - "PromptResponseError", - "QuestionAnsweringDataset", - "QuestionAnsweringEntry", - "RegistryReference", - "QuestionChoice", - "REGISTRY_NAME_PATTERN", - "ScaleDescription", - "Scorable", - "Score", - "ScoreType", - "ScoringExpectation", - "ScenarioEvaluationIdentifier", - "ScorerEvaluationIdentifier", - "ScorerIdentifier", - "ScenarioIdentifier", - "ScenarioResult", - "ScenarioRunState", - "Seed", - "AttackSeedGroup", - "AttackTechniqueSeedGroup", - "SeedObjective", - "SeedPrompt", - "SeedDataset", - "SeedGroup", - "SeedIdentifier", - "SeedSimulatedConversation", - "SeedType", - "SeedUnion", - "SimulatedTargetSystemPromptPaths", - "snake_case_to_class_name", - "sort_message_pieces", - "StrategyResult", - "StrategyResultT", - "TARGET_EVAL_PARAM_FALLBACKS", - "TARGET_EVAL_PARAMS", - "TargetCapabilities", - "TargetIdentifier", - "TextDataTypeSerializer", - "TOKEN_USAGE_METADATA_PREFIX", - "TokenUsage", - "ToolCall", - "UnvalidatedScore", - "read_usage_int", - "read_usage_value", - "validate_registry_name", - "RetryEvent", -] +if TYPE_CHECKING: + from pyrit.models.additional_initializer import AdditionalInitializer + from pyrit.models.conversation_stats import ConversationStats + from pyrit.models.embeddings import EmbeddingData, EmbeddingResponse, EmbeddingSupport, EmbeddingUsageInformation + from pyrit.models.harm_definition import HarmDefinition, ScaleDescription, get_all_harm_definitions + from pyrit.models.identifiers import ( + REGISTRY_NAME_PATTERN, + TARGET_EVAL_PARAM_FALLBACKS, + TARGET_EVAL_PARAMS, + AtomicAttackEvaluationIdentifier, + AtomicAttackIdentifier, + AttackIdentifier, + AttackTechniqueIdentifier, + ChildEvalRule, + ComponentIdentifier, + ConverterIdentifier, + Evaluate, + EvaluationIdentifier, + Identifiable, + IdentifierFilter, + IdentifierType, + JSONValue, + ObjectiveTargetEvaluationIdentifier, + ScenarioEvaluationIdentifier, + ScenarioIdentifier, + ScorerEvaluationIdentifier, + ScorerIdentifier, + SeedIdentifier, + TargetIdentifier, + class_name_to_snake_case, + compute_eval_hash, + config_hash, + snake_case_to_class_name, + validate_registry_name, + ) + from pyrit.models.literals import ( + MEDIA_PATH_DATA_TYPES, + ChatMessageRole, + Modality, + PromptDataType, + PromptResponseError, + SeedType, + ) + from pyrit.models.messages import ( + Conversation, + Message, + MessagePiece, + construct_response_from_request, + flatten_to_message_pieces, + get_all_values, + group_conversation_message_pieces_by_sequence, + group_message_pieces_into_conversations, + sort_message_pieces, + ) + from pyrit.models.messages.chat_message import ( + ALLOWED_CHAT_MESSAGE_ROLES, + ChatMessage, + ChatMessagesDataset, + ToolCall, + ) + from pyrit.models.messages.conversation_reference import ConversationReference, ConversationType + from pyrit.models.messages.conversation_retry import ConversationRetry, ConversationRetryReason + from pyrit.models.parameter import ( + ComponentType, + Parameter, + ParameterDestination, + RegistryReference, + display_choices, + ) + from pyrit.models.question_answering import QuestionAnsweringDataset, QuestionAnsweringEntry, QuestionChoice + from pyrit.models.results.attack_result import AttackOutcome, AttackResult, AttackResultT + from pyrit.models.results.scenario_result import ScenarioResult, ScenarioRunState + from pyrit.models.results.strategy_result import StrategyResult, StrategyResultT + from pyrit.models.retry_event import RetryEvent + from pyrit.models.score import ( + Condition, + ContentScorable, + MatchesObjective, + MessageScorable, + Scorable, + Score, + ScoreType, + ScoringExpectation, + UnvalidatedScore, + ) + from pyrit.models.seeds import ( + AttackSeedGroup, + AttackTechniqueSeedGroup, + NextMessageSystemPromptPaths, + Seed, + SeedDataset, + SeedGroup, + SeedObjective, + SeedPrompt, + SeedSimulatedConversation, + SeedUnion, + SimulatedTargetSystemPromptPaths, + group_seeds_into_attack_groups, + ) + from pyrit.models.target import ( + COMMON_JSON_SCHEMAS, + JSON_SCHEMA_METADATA_KEY, + SEED_RESPONSE_JSON_SCHEMA_METADATA_KEY, + TOKEN_USAGE_METADATA_PREFIX, + CapabilityName, + JsonResponseConfig, + JsonSchemaDefinition, + TargetCapabilities, + TokenUsage, + get_common_json_schema, + read_usage_int, + read_usage_value, + register_common_json_schema, + unregister_common_json_schema, + ) + +_LAZY_EXPORTS: dict[str, str] = { + "ALLOWED_CHAT_MESSAGE_ROLES": "pyrit.models.messages.chat_message", + "AdditionalInitializer": "pyrit.models.additional_initializer", + "AtomicAttackEvaluationIdentifier": "pyrit.models.identifiers", + "AtomicAttackIdentifier": "pyrit.models.identifiers", + "AttackIdentifier": "pyrit.models.identifiers", + "AttackTechniqueIdentifier": "pyrit.models.identifiers", + "AttackResult": "pyrit.models.results.attack_result", + "AttackResultT": "pyrit.models.results.attack_result", + "AttackOutcome": "pyrit.models.results.attack_result", + "ChatMessage": "pyrit.models.messages.chat_message", + "ChatMessagesDataset": "pyrit.models.messages.chat_message", + "ChatMessageRole": "pyrit.models.literals", + "ChildEvalRule": "pyrit.models.identifiers", + "class_name_to_snake_case": "pyrit.models.identifiers", + "CapabilityName": "pyrit.models.target", + "ComponentIdentifier": "pyrit.models.identifiers", + "ComponentType": "pyrit.models.parameter", + "compute_eval_hash": "pyrit.models.identifiers", + "Condition": "pyrit.models.score", + "config_hash": "pyrit.models.identifiers", + "ConverterIdentifier": "pyrit.models.identifiers", + "Conversation": "pyrit.models.messages.conversations", + "ConversationReference": "pyrit.models.messages.conversation_reference", + "ConversationRetry": "pyrit.models.messages.conversation_retry", + "ConversationRetryReason": "pyrit.models.messages.conversation_retry", + "ConversationStats": "pyrit.models.conversation_stats", + "ConversationType": "pyrit.models.messages.conversation_reference", + "ContentScorable": "pyrit.models.score", + "construct_response_from_request": "pyrit.models.messages.conversations", + "display_choices": "pyrit.models.parameter", + "EmbeddingData": "pyrit.models.embeddings", + "EmbeddingResponse": "pyrit.models.embeddings", + "EmbeddingSupport": "pyrit.models.embeddings", + "EmbeddingUsageInformation": "pyrit.models.embeddings", + "Evaluate": "pyrit.models.identifiers", + "EvaluationIdentifier": "pyrit.models.identifiers", + "flatten_to_message_pieces": "pyrit.models.messages.conversations", + "get_all_harm_definitions": "pyrit.models.harm_definition", + "get_all_values": "pyrit.models.messages.conversations", + "group_conversation_message_pieces_by_sequence": "pyrit.models.messages.conversations", + "group_message_pieces_into_conversations": "pyrit.models.messages.conversations", + "group_seeds_into_attack_groups": "pyrit.models.seeds", + "HarmDefinition": "pyrit.models.harm_definition", + "Identifiable": "pyrit.models.identifiers", + "IdentifierFilter": "pyrit.models.identifiers", + "IdentifierType": "pyrit.models.identifiers", + "JSONValue": "pyrit.models.identifiers", + "COMMON_JSON_SCHEMAS": "pyrit.models.target", + "JsonResponseConfig": "pyrit.models.target", + "get_common_json_schema": "pyrit.models.target", + "register_common_json_schema": "pyrit.models.target", + "unregister_common_json_schema": "pyrit.models.target", + "JSON_SCHEMA_METADATA_KEY": "pyrit.models.target", + "SEED_RESPONSE_JSON_SCHEMA_METADATA_KEY": "pyrit.models.target", + "JsonSchemaDefinition": "pyrit.models.target", + "MatchesObjective": "pyrit.models.score", + "MEDIA_PATH_DATA_TYPES": "pyrit.models.literals", + "Message": "pyrit.models.messages.message", + "MessagePiece": "pyrit.models.messages.message_piece", + "MessageScorable": "pyrit.models.score", + "Modality": "pyrit.models.literals", + "NextMessageSystemPromptPaths": "pyrit.models.seeds", + "ObjectiveTargetEvaluationIdentifier": "pyrit.models.identifiers", + "Parameter": "pyrit.models.parameter", + "ParameterDestination": "pyrit.models.parameter", + "PromptDataType": "pyrit.models.literals", + "PromptResponseError": "pyrit.models.literals", + "QuestionAnsweringDataset": "pyrit.models.question_answering", + "QuestionAnsweringEntry": "pyrit.models.question_answering", + "RegistryReference": "pyrit.models.parameter", + "QuestionChoice": "pyrit.models.question_answering", + "REGISTRY_NAME_PATTERN": "pyrit.models.identifiers", + "ScaleDescription": "pyrit.models.harm_definition", + "Scorable": "pyrit.models.score", + "Score": "pyrit.models.score", + "ScoreType": "pyrit.models.score", + "ScoringExpectation": "pyrit.models.score", + "ScenarioEvaluationIdentifier": "pyrit.models.identifiers", + "ScorerEvaluationIdentifier": "pyrit.models.identifiers", + "ScorerIdentifier": "pyrit.models.identifiers", + "ScenarioIdentifier": "pyrit.models.identifiers", + "ScenarioResult": "pyrit.models.results.scenario_result", + "ScenarioRunState": "pyrit.models.results.scenario_result", + "Seed": "pyrit.models.seeds", + "AttackSeedGroup": "pyrit.models.seeds", + "AttackTechniqueSeedGroup": "pyrit.models.seeds", + "SeedObjective": "pyrit.models.seeds", + "SeedPrompt": "pyrit.models.seeds", + "SeedDataset": "pyrit.models.seeds", + "SeedGroup": "pyrit.models.seeds", + "SeedIdentifier": "pyrit.models.identifiers", + "SeedSimulatedConversation": "pyrit.models.seeds", + "SeedType": "pyrit.models.literals", + "SeedUnion": "pyrit.models.seeds", + "SimulatedTargetSystemPromptPaths": "pyrit.models.seeds", + "snake_case_to_class_name": "pyrit.models.identifiers", + "sort_message_pieces": "pyrit.models.messages.message_piece", + "StrategyResult": "pyrit.models.results.strategy_result", + "StrategyResultT": "pyrit.models.results.strategy_result", + "TARGET_EVAL_PARAM_FALLBACKS": "pyrit.models.identifiers", + "TARGET_EVAL_PARAMS": "pyrit.models.identifiers", + "TargetCapabilities": "pyrit.models.target", + "TargetIdentifier": "pyrit.models.identifiers", + "TOKEN_USAGE_METADATA_PREFIX": "pyrit.models.target", + "TokenUsage": "pyrit.models.target", + "ToolCall": "pyrit.models.messages.chat_message", + "UnvalidatedScore": "pyrit.models.score", + "read_usage_int": "pyrit.models.target", + "read_usage_value": "pyrit.models.target", + "validate_registry_name": "pyrit.models.identifiers", + "RetryEvent": "pyrit.models.retry_event", +} + +__all__ = list(_LAZY_EXPORTS) + + +def __getattr__(name: str) -> object: + """ + Resolve a public model export on first access. + + Args: + name (str): The requested public name. + + Returns: + object: The resolved export. + """ + return resolve_lazy_export( + name=name, + module_name=__name__, + module_globals=globals(), + exports=_LAZY_EXPORTS, + ) + + +def __dir__() -> list[str]: + """Return package attributes, including unresolved exports.""" + return get_lazy_dir(module_globals=globals(), exports=_LAZY_EXPORTS) diff --git a/pyrit/models/catalog/__init__.py b/pyrit/models/catalog/__init__.py index 6d8e2e15d3..2624c812a2 100644 --- a/pyrit/models/catalog/__init__.py +++ b/pyrit/models/catalog/__init__.py @@ -1,5 +1,6 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +# ruff: noqa: F401 """ Catalog sub-package - registry/wire-format types for scenarios, initializers, @@ -13,26 +14,52 @@ ``pyrit.backend.models`` """ -from pyrit.models.catalog.initializer import ( - RegisteredInitializer, -) -from pyrit.models.catalog.scenario import ( - AttackErrorSummary, - AttackRetrySummary, - RegisteredScenario, - RunScenarioRequest, - ScenarioRunSummary, -) -from pyrit.models.catalog.target import ( - TargetInstance, -) - -__all__ = [ - "AttackErrorSummary", - "AttackRetrySummary", - "RegisteredInitializer", - "RegisteredScenario", - "RunScenarioRequest", - "ScenarioRunSummary", - "TargetInstance", -] +from typing import TYPE_CHECKING + +from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export + +if TYPE_CHECKING: + from pyrit.models.catalog.initializer import RegisteredInitializer + from pyrit.models.catalog.scenario import ( + AttackErrorSummary, + AttackRetrySummary, + RegisteredScenario, + RunScenarioRequest, + ScenarioRunSummary, + ) + from pyrit.models.catalog.target import TargetInstance + +_LAZY_EXPORTS: dict[str, str] = { + "AttackErrorSummary": "pyrit.models.catalog.scenario", + "AttackRetrySummary": "pyrit.models.catalog.scenario", + "RegisteredInitializer": "pyrit.models.catalog.initializer", + "RegisteredScenario": "pyrit.models.catalog.scenario", + "RunScenarioRequest": "pyrit.models.catalog.scenario", + "ScenarioRunSummary": "pyrit.models.catalog.scenario", + "TargetInstance": "pyrit.models.catalog.target", +} + +__all__ = list(_LAZY_EXPORTS) + + +def __getattr__(name: str) -> object: + """ + Resolve a public catalog export on first access. + + Args: + name (str): The requested public name. + + Returns: + object: The resolved export. + """ + return resolve_lazy_export( + name=name, + module_name=__name__, + module_globals=globals(), + exports=_LAZY_EXPORTS, + ) + + +def __dir__() -> list[str]: + """Return package attributes, including unresolved exports.""" + return get_lazy_dir(module_globals=globals(), exports=_LAZY_EXPORTS) diff --git a/pyrit/models/identifiers/__init__.py b/pyrit/models/identifiers/__init__.py index aea1532581..bfff6d3ca9 100644 --- a/pyrit/models/identifiers/__init__.py +++ b/pyrit/models/identifiers/__init__.py @@ -1,82 +1,111 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +# ruff: noqa: F401 """Identifiers module for PyRIT components.""" -from pyrit.models.identifiers.atomic_attack_identifier import ( - AtomicAttackIdentifier, -) -from pyrit.models.identifiers.attack_identifier import AttackIdentifier -from pyrit.models.identifiers.attack_technique_identifier import AttackTechniqueIdentifier -from pyrit.models.identifiers.class_name_utils import ( - REGISTRY_NAME_PATTERN, - class_name_to_snake_case, - snake_case_to_class_name, - validate_registry_name, -) -from pyrit.models.identifiers.component_identifier import ( - ComponentIdentifier, - Identifiable, - JSONValue, - config_hash, -) -from pyrit.models.identifiers.converter_identifier import ConverterIdentifier -from pyrit.models.identifiers.evaluation_identifier import ( - TARGET_EVAL_PARAM_FALLBACKS, - TARGET_EVAL_PARAMS, - AtomicAttackEvaluationIdentifier, - ChildEvalRule, - EvaluationIdentifier, - ObjectiveTargetEvaluationIdentifier, - ScenarioEvaluationIdentifier, - ScorerEvaluationIdentifier, - compute_eval_hash, - compute_inner_attack_eval_hash, - derive_eval_config, -) -from pyrit.models.identifiers.evaluation_markers import EvalMarker, Evaluate, Exclude, Include, Unwrap -from pyrit.models.identifiers.identifier_filters import IdentifierFilter, IdentifierType -from pyrit.models.identifiers.param_markers import Param, ParamMarker -from pyrit.models.identifiers.scenario_identifier import ScenarioIdentifier -from pyrit.models.identifiers.scorer_identifier import ScorerIdentifier -from pyrit.models.identifiers.seed_identifier import SeedIdentifier -from pyrit.models.identifiers.target_identifier import TargetIdentifier +from typing import TYPE_CHECKING -__all__ = [ - "AtomicAttackEvaluationIdentifier", - "AtomicAttackIdentifier", - "AttackIdentifier", - "AttackTechniqueIdentifier", - "ChildEvalRule", - "class_name_to_snake_case", - "ComponentIdentifier", - "compute_eval_hash", - "compute_inner_attack_eval_hash", - "ConverterIdentifier", - "derive_eval_config", - "EvalMarker", - "Evaluate", - "EvaluationIdentifier", - "Exclude", - "Identifiable", - "Include", - "JSONValue", - "ObjectiveTargetEvaluationIdentifier", - "REGISTRY_NAME_PATTERN", - "Param", - "ParamMarker", - "ScenarioEvaluationIdentifier", - "ScorerEvaluationIdentifier", - "ScorerIdentifier", - "ScenarioIdentifier", - "SeedIdentifier", - "snake_case_to_class_name", - "TARGET_EVAL_PARAM_FALLBACKS", - "TARGET_EVAL_PARAMS", - "TargetIdentifier", - "Unwrap", - "validate_registry_name", - "config_hash", - "IdentifierFilter", - "IdentifierType", -] +from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export + +if TYPE_CHECKING: + from pyrit.models.identifiers.atomic_attack_identifier import AtomicAttackIdentifier + from pyrit.models.identifiers.attack_identifier import AttackIdentifier + from pyrit.models.identifiers.attack_technique_identifier import AttackTechniqueIdentifier + from pyrit.models.identifiers.class_name_utils import ( + REGISTRY_NAME_PATTERN, + class_name_to_snake_case, + snake_case_to_class_name, + validate_registry_name, + ) + from pyrit.models.identifiers.component_identifier import ( + ComponentIdentifier, + Identifiable, + JSONValue, + config_hash, + ) + from pyrit.models.identifiers.converter_identifier import ConverterIdentifier + from pyrit.models.identifiers.evaluation_identifier import ( + TARGET_EVAL_PARAM_FALLBACKS, + TARGET_EVAL_PARAMS, + AtomicAttackEvaluationIdentifier, + ChildEvalRule, + EvaluationIdentifier, + ObjectiveTargetEvaluationIdentifier, + ScenarioEvaluationIdentifier, + ScorerEvaluationIdentifier, + compute_eval_hash, + compute_inner_attack_eval_hash, + derive_eval_config, + ) + from pyrit.models.identifiers.evaluation_markers import EvalMarker, Evaluate, Exclude, Include, Unwrap + from pyrit.models.identifiers.identifier_filters import IdentifierFilter, IdentifierType + from pyrit.models.identifiers.param_markers import Param, ParamMarker + from pyrit.models.identifiers.scenario_identifier import ScenarioIdentifier + from pyrit.models.identifiers.scorer_identifier import ScorerIdentifier + from pyrit.models.identifiers.seed_identifier import SeedIdentifier + from pyrit.models.identifiers.target_identifier import TargetIdentifier + +_LAZY_EXPORTS: dict[str, str] = { + "AtomicAttackEvaluationIdentifier": "pyrit.models.identifiers.evaluation_identifier", + "AtomicAttackIdentifier": "pyrit.models.identifiers.atomic_attack_identifier", + "AttackIdentifier": "pyrit.models.identifiers.attack_identifier", + "AttackTechniqueIdentifier": "pyrit.models.identifiers.attack_technique_identifier", + "ChildEvalRule": "pyrit.models.identifiers.evaluation_identifier", + "class_name_to_snake_case": "pyrit.models.identifiers.class_name_utils", + "ComponentIdentifier": "pyrit.models.identifiers.component_identifier", + "compute_eval_hash": "pyrit.models.identifiers.evaluation_identifier", + "compute_inner_attack_eval_hash": "pyrit.models.identifiers.evaluation_identifier", + "ConverterIdentifier": "pyrit.models.identifiers.converter_identifier", + "derive_eval_config": "pyrit.models.identifiers.evaluation_identifier", + "EvalMarker": "pyrit.models.identifiers.evaluation_markers", + "Evaluate": "pyrit.models.identifiers.evaluation_markers", + "EvaluationIdentifier": "pyrit.models.identifiers.evaluation_identifier", + "Exclude": "pyrit.models.identifiers.evaluation_markers", + "Identifiable": "pyrit.models.identifiers.component_identifier", + "Include": "pyrit.models.identifiers.evaluation_markers", + "JSONValue": "pyrit.models.identifiers.component_identifier", + "ObjectiveTargetEvaluationIdentifier": "pyrit.models.identifiers.evaluation_identifier", + "REGISTRY_NAME_PATTERN": "pyrit.models.identifiers.class_name_utils", + "Param": "pyrit.models.identifiers.param_markers", + "ParamMarker": "pyrit.models.identifiers.param_markers", + "ScenarioEvaluationIdentifier": "pyrit.models.identifiers.evaluation_identifier", + "ScorerEvaluationIdentifier": "pyrit.models.identifiers.evaluation_identifier", + "ScorerIdentifier": "pyrit.models.identifiers.scorer_identifier", + "ScenarioIdentifier": "pyrit.models.identifiers.scenario_identifier", + "SeedIdentifier": "pyrit.models.identifiers.seed_identifier", + "snake_case_to_class_name": "pyrit.models.identifiers.class_name_utils", + "TARGET_EVAL_PARAM_FALLBACKS": "pyrit.models.identifiers.evaluation_identifier", + "TARGET_EVAL_PARAMS": "pyrit.models.identifiers.evaluation_identifier", + "TargetIdentifier": "pyrit.models.identifiers.target_identifier", + "Unwrap": "pyrit.models.identifiers.evaluation_markers", + "validate_registry_name": "pyrit.models.identifiers.class_name_utils", + "config_hash": "pyrit.models.identifiers.component_identifier", + "IdentifierFilter": "pyrit.models.identifiers.identifier_filters", + "IdentifierType": "pyrit.models.identifiers.identifier_filters", +} + +__all__ = list(_LAZY_EXPORTS) + + +def __getattr__(name: str) -> object: + """ + Resolve a public identifier export on first access. + + Args: + name (str): The requested public name. + + Returns: + object: The resolved export. + """ + return resolve_lazy_export( + name=name, + module_name=__name__, + module_globals=globals(), + exports=_LAZY_EXPORTS, + ) + + +def __dir__() -> list[str]: + """Return package attributes, including unresolved exports.""" + return get_lazy_dir(module_globals=globals(), exports=_LAZY_EXPORTS) diff --git a/pyrit/models/messages/__init__.py b/pyrit/models/messages/__init__.py index c18b2a0964..75639bbbba 100644 --- a/pyrit/models/messages/__init__.py +++ b/pyrit/models/messages/__init__.py @@ -1,5 +1,6 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +# ruff: noqa: F401 """ Messages module - message types and helpers for PyRIT. @@ -12,38 +13,68 @@ - conversations: Free functions that operate on collections of messages/pieces. """ -from pyrit.models.messages.chat_message import ( - ALLOWED_CHAT_MESSAGE_ROLES, - ChatMessage, - ChatMessagesDataset, - ToolCall, -) -from pyrit.models.messages.conversation_reference import ConversationReference, ConversationType -from pyrit.models.messages.conversations import ( - Conversation, - construct_response_from_request, - flatten_to_message_pieces, - get_all_values, - group_conversation_message_pieces_by_sequence, - group_message_pieces_into_conversations, -) -from pyrit.models.messages.message import Message -from pyrit.models.messages.message_piece import MessagePiece, sort_message_pieces - -__all__ = [ - "ALLOWED_CHAT_MESSAGE_ROLES", - "ChatMessage", - "ChatMessagesDataset", - "Conversation", - "ConversationReference", - "ConversationType", - "Message", - "MessagePiece", - "ToolCall", - "construct_response_from_request", - "flatten_to_message_pieces", - "get_all_values", - "group_conversation_message_pieces_by_sequence", - "group_message_pieces_into_conversations", - "sort_message_pieces", -] +from typing import TYPE_CHECKING + +from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export + +if TYPE_CHECKING: + from pyrit.models.messages.chat_message import ( + ALLOWED_CHAT_MESSAGE_ROLES, + ChatMessage, + ChatMessagesDataset, + ToolCall, + ) + from pyrit.models.messages.conversation_reference import ConversationReference, ConversationType + from pyrit.models.messages.conversations import ( + Conversation, + construct_response_from_request, + flatten_to_message_pieces, + get_all_values, + group_conversation_message_pieces_by_sequence, + group_message_pieces_into_conversations, + ) + from pyrit.models.messages.message import Message + from pyrit.models.messages.message_piece import MessagePiece, sort_message_pieces + +_LAZY_EXPORTS: dict[str, str] = { + "ALLOWED_CHAT_MESSAGE_ROLES": "pyrit.models.messages.chat_message", + "ChatMessage": "pyrit.models.messages.chat_message", + "ChatMessagesDataset": "pyrit.models.messages.chat_message", + "Conversation": "pyrit.models.messages.conversations", + "ConversationReference": "pyrit.models.messages.conversation_reference", + "ConversationType": "pyrit.models.messages.conversation_reference", + "Message": "pyrit.models.messages.message", + "MessagePiece": "pyrit.models.messages.message_piece", + "ToolCall": "pyrit.models.messages.chat_message", + "construct_response_from_request": "pyrit.models.messages.conversations", + "flatten_to_message_pieces": "pyrit.models.messages.conversations", + "get_all_values": "pyrit.models.messages.conversations", + "group_conversation_message_pieces_by_sequence": "pyrit.models.messages.conversations", + "group_message_pieces_into_conversations": "pyrit.models.messages.conversations", + "sort_message_pieces": "pyrit.models.messages.message_piece", +} + +__all__ = list(_LAZY_EXPORTS) + + +def __getattr__(name: str) -> object: + """ + Resolve a public message export on first access. + + Args: + name (str): The requested public name. + + Returns: + object: The resolved export. + """ + return resolve_lazy_export( + name=name, + module_name=__name__, + module_globals=globals(), + exports=_LAZY_EXPORTS, + ) + + +def __dir__() -> list[str]: + """Return package attributes, including unresolved exports.""" + return get_lazy_dir(module_globals=globals(), exports=_LAZY_EXPORTS) diff --git a/pyrit/models/results/__init__.py b/pyrit/models/results/__init__.py index 4332080b4f..fbb01363f1 100644 --- a/pyrit/models/results/__init__.py +++ b/pyrit/models/results/__init__.py @@ -1,5 +1,6 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +# ruff: noqa: F401 """ Results module - strategy, attack, and scenario result types for PyRIT. @@ -12,21 +13,48 @@ - ScenarioRunState: Lifecycle state of a scenario run. """ -from pyrit.models.identifiers.scenario_identifier import ScenarioIdentifier -from pyrit.models.results.attack_result import AttackOutcome, AttackResult, AttackResultT -from pyrit.models.results.scenario_result import ( - ScenarioResult, - ScenarioRunState, -) -from pyrit.models.results.strategy_result import StrategyResult, StrategyResultT - -__all__ = [ - "AttackOutcome", - "AttackResult", - "AttackResultT", - "ScenarioIdentifier", - "ScenarioResult", - "ScenarioRunState", - "StrategyResult", - "StrategyResultT", -] +from typing import TYPE_CHECKING + +from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export + +if TYPE_CHECKING: + from pyrit.models.identifiers.scenario_identifier import ScenarioIdentifier + from pyrit.models.results.attack_result import AttackOutcome, AttackResult, AttackResultT + from pyrit.models.results.scenario_result import ScenarioResult, ScenarioRunState + from pyrit.models.results.strategy_result import StrategyResult, StrategyResultT + +_LAZY_EXPORTS: dict[str, str] = { + "AttackOutcome": "pyrit.models.results.attack_result", + "AttackResult": "pyrit.models.results.attack_result", + "AttackResultT": "pyrit.models.results.attack_result", + "ScenarioIdentifier": "pyrit.models.identifiers.scenario_identifier", + "ScenarioResult": "pyrit.models.results.scenario_result", + "ScenarioRunState": "pyrit.models.results.scenario_result", + "StrategyResult": "pyrit.models.results.strategy_result", + "StrategyResultT": "pyrit.models.results.strategy_result", +} + +__all__ = list(_LAZY_EXPORTS) + + +def __getattr__(name: str) -> object: + """ + Resolve a public result export on first access. + + Args: + name (str): The requested public name. + + Returns: + object: The resolved export. + """ + return resolve_lazy_export( + name=name, + module_name=__name__, + module_globals=globals(), + exports=_LAZY_EXPORTS, + ) + + +def __dir__() -> list[str]: + """Return package attributes, including unresolved exports.""" + return get_lazy_dir(module_globals=globals(), exports=_LAZY_EXPORTS) diff --git a/pyrit/models/score/__init__.py b/pyrit/models/score/__init__.py index b870e813b9..63b583fa2f 100644 --- a/pyrit/models/score/__init__.py +++ b/pyrit/models/score/__init__.py @@ -1,5 +1,6 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +# ruff: noqa: F401 """ Score types: what a scorer looks at, what it scores against, and the result. @@ -9,20 +10,50 @@ are inert canonical data; scoring-layer resolvers acquire the evidence they name. """ -from pyrit.models.score.condition import Condition, MatchesObjective -from pyrit.models.score.expectation import ScoringExpectation -from pyrit.models.score.scorable import ContentScorable, MessageScorable, Scorable -from pyrit.models.score.score import ComponentIdentifierField, Score, ScoreType, UnvalidatedScore - -__all__ = [ - "ComponentIdentifierField", - "Condition", - "ContentScorable", - "MatchesObjective", - "MessageScorable", - "Scorable", - "Score", - "ScoreType", - "ScoringExpectation", - "UnvalidatedScore", -] +from typing import TYPE_CHECKING + +from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export + +if TYPE_CHECKING: + from pyrit.models.score.condition import Condition, MatchesObjective + from pyrit.models.score.expectation import ScoringExpectation + from pyrit.models.score.scorable import ContentScorable, MessageScorable, Scorable + from pyrit.models.score.score import ComponentIdentifierField, Score, ScoreType, UnvalidatedScore + +_LAZY_EXPORTS: dict[str, str] = { + "ComponentIdentifierField": "pyrit.models.score.score", + "Condition": "pyrit.models.score.condition", + "ContentScorable": "pyrit.models.score.scorable", + "MatchesObjective": "pyrit.models.score.condition", + "MessageScorable": "pyrit.models.score.scorable", + "Scorable": "pyrit.models.score.scorable", + "Score": "pyrit.models.score.score", + "ScoreType": "pyrit.models.score.score", + "ScoringExpectation": "pyrit.models.score.expectation", + "UnvalidatedScore": "pyrit.models.score.score", +} + +__all__ = list(_LAZY_EXPORTS) + + +def __getattr__(name: str) -> object: + """ + Resolve a public score export on first access. + + Args: + name (str): The requested public name. + + Returns: + object: The resolved export. + """ + return resolve_lazy_export( + name=name, + module_name=__name__, + module_globals=globals(), + exports=_LAZY_EXPORTS, + ) + + +def __dir__() -> list[str]: + """Return package attributes, including unresolved exports.""" + return get_lazy_dir(module_globals=globals(), exports=_LAZY_EXPORTS) diff --git a/pyrit/models/seeds/__init__.py b/pyrit/models/seeds/__init__.py index a4e01145ff..dbb015bf05 100644 --- a/pyrit/models/seeds/__init__.py +++ b/pyrit/models/seeds/__init__.py @@ -1,5 +1,6 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +# ruff: noqa: F401 """ Seeds module - Contains all seed-related classes for PyRIT. @@ -15,39 +16,69 @@ - SeedDataset: Container for managing collections of seeds """ -from pyrit.models.seeds.attack_seed_group import AttackSeedGroup -from pyrit.models.seeds.attack_technique_seed_group import AttackTechniqueSeedGroup -from pyrit.models.seeds.seed import Seed -from pyrit.models.seeds.seed_dataset import SeedDataset -from pyrit.models.seeds.seed_group import SeedGroup, SeedUnion -from pyrit.models.seeds.seed_grouping import group_seeds_into_attack_groups -from pyrit.models.seeds.seed_objective import SeedObjective -from pyrit.models.seeds.seed_prompt import SeedPrompt -from pyrit.models.seeds.seed_simulated_conversation import ( - NextMessageSystemPromptPaths, - SeedSimulatedConversation, - SimulatedTargetSystemPromptPaths, -) -from pyrit.models.seeds.yaml_seed_loader import ( - load_seed_dataset_from_yaml, - load_seed_from_yaml, - load_seed_prompt_from_yaml_with_required_parameters, -) - -__all__ = [ - "load_seed_dataset_from_yaml", - "load_seed_from_yaml", - "load_seed_prompt_from_yaml_with_required_parameters", - "group_seeds_into_attack_groups", - "NextMessageSystemPromptPaths", - "Seed", - "AttackSeedGroup", - "AttackTechniqueSeedGroup", - "SeedDataset", - "SeedGroup", - "SeedObjective", - "SeedPrompt", - "SeedSimulatedConversation", - "SeedUnion", - "SimulatedTargetSystemPromptPaths", -] +from typing import TYPE_CHECKING + +from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export + +if TYPE_CHECKING: + from pyrit.models.seeds.attack_seed_group import AttackSeedGroup + from pyrit.models.seeds.attack_technique_seed_group import AttackTechniqueSeedGroup + from pyrit.models.seeds.seed import Seed + from pyrit.models.seeds.seed_dataset import SeedDataset + from pyrit.models.seeds.seed_group import SeedGroup, SeedUnion + from pyrit.models.seeds.seed_grouping import group_seeds_into_attack_groups + from pyrit.models.seeds.seed_objective import SeedObjective + from pyrit.models.seeds.seed_prompt import SeedPrompt + from pyrit.models.seeds.seed_simulated_conversation import ( + NextMessageSystemPromptPaths, + SeedSimulatedConversation, + SimulatedTargetSystemPromptPaths, + ) + from pyrit.models.seeds.yaml_seed_loader import ( + load_seed_dataset_from_yaml, + load_seed_from_yaml, + load_seed_prompt_from_yaml_with_required_parameters, + ) + +_LAZY_EXPORTS: dict[str, str] = { + "load_seed_dataset_from_yaml": "pyrit.models.seeds.yaml_seed_loader", + "load_seed_from_yaml": "pyrit.models.seeds.yaml_seed_loader", + "load_seed_prompt_from_yaml_with_required_parameters": "pyrit.models.seeds.yaml_seed_loader", + "group_seeds_into_attack_groups": "pyrit.models.seeds.seed_grouping", + "NextMessageSystemPromptPaths": "pyrit.models.seeds.seed_simulated_conversation", + "Seed": "pyrit.models.seeds.seed", + "AttackSeedGroup": "pyrit.models.seeds.attack_seed_group", + "AttackTechniqueSeedGroup": "pyrit.models.seeds.attack_technique_seed_group", + "SeedDataset": "pyrit.models.seeds.seed_dataset", + "SeedGroup": "pyrit.models.seeds.seed_group", + "SeedObjective": "pyrit.models.seeds.seed_objective", + "SeedPrompt": "pyrit.models.seeds.seed_prompt", + "SeedSimulatedConversation": "pyrit.models.seeds.seed_simulated_conversation", + "SeedUnion": "pyrit.models.seeds.seed_group", + "SimulatedTargetSystemPromptPaths": "pyrit.models.seeds.seed_simulated_conversation", +} + +__all__ = list(_LAZY_EXPORTS) + + +def __getattr__(name: str) -> object: + """ + Resolve a public seed export on first access. + + Args: + name (str): The requested public name. + + Returns: + object: The resolved export. + """ + return resolve_lazy_export( + name=name, + module_name=__name__, + module_globals=globals(), + exports=_LAZY_EXPORTS, + ) + + +def __dir__() -> list[str]: + """Return package attributes, including unresolved exports.""" + return get_lazy_dir(module_globals=globals(), exports=_LAZY_EXPORTS) diff --git a/pyrit/models/target/__init__.py b/pyrit/models/target/__init__.py index 554154bdd8..92254b659f 100644 --- a/pyrit/models/target/__init__.py +++ b/pyrit/models/target/__init__.py @@ -1,5 +1,6 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +# ruff: noqa: F401 """ Canonical data models for how PyRIT interacts with targets. @@ -17,37 +18,67 @@ ``from pyrit.models import TokenUsage``). """ -from pyrit.models.target.json_response_config import JsonResponseConfig -from pyrit.models.target.json_schema_definition import ( - COMMON_JSON_SCHEMAS, - JSON_SCHEMA_METADATA_KEY, - SEED_RESPONSE_JSON_SCHEMA_METADATA_KEY, - JsonSchemaDefinition, - get_common_json_schema, - register_common_json_schema, - unregister_common_json_schema, -) -from pyrit.models.target.target_capabilities import CapabilityName, TargetCapabilities -from pyrit.models.target.token_usage import ( - TOKEN_USAGE_METADATA_PREFIX, - TokenUsage, - read_usage_int, - read_usage_value, -) - -__all__ = [ - "COMMON_JSON_SCHEMAS", - "CapabilityName", - "JSON_SCHEMA_METADATA_KEY", - "JsonResponseConfig", - "JsonSchemaDefinition", - "SEED_RESPONSE_JSON_SCHEMA_METADATA_KEY", - "TOKEN_USAGE_METADATA_PREFIX", - "TargetCapabilities", - "TokenUsage", - "get_common_json_schema", - "read_usage_int", - "read_usage_value", - "register_common_json_schema", - "unregister_common_json_schema", -] +from typing import TYPE_CHECKING + +from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export + +if TYPE_CHECKING: + from pyrit.models.target.json_response_config import JsonResponseConfig + from pyrit.models.target.json_schema_definition import ( + COMMON_JSON_SCHEMAS, + JSON_SCHEMA_METADATA_KEY, + SEED_RESPONSE_JSON_SCHEMA_METADATA_KEY, + JsonSchemaDefinition, + get_common_json_schema, + register_common_json_schema, + unregister_common_json_schema, + ) + from pyrit.models.target.target_capabilities import CapabilityName, TargetCapabilities + from pyrit.models.target.token_usage import ( + TOKEN_USAGE_METADATA_PREFIX, + TokenUsage, + read_usage_int, + read_usage_value, + ) + +_LAZY_EXPORTS: dict[str, str] = { + "COMMON_JSON_SCHEMAS": "pyrit.models.target.json_schema_definition", + "CapabilityName": "pyrit.models.target.target_capabilities", + "JSON_SCHEMA_METADATA_KEY": "pyrit.models.target.json_schema_definition", + "JsonResponseConfig": "pyrit.models.target.json_response_config", + "JsonSchemaDefinition": "pyrit.models.target.json_schema_definition", + "SEED_RESPONSE_JSON_SCHEMA_METADATA_KEY": "pyrit.models.target.json_schema_definition", + "TOKEN_USAGE_METADATA_PREFIX": "pyrit.models.target.token_usage", + "TargetCapabilities": "pyrit.models.target.target_capabilities", + "TokenUsage": "pyrit.models.target.token_usage", + "get_common_json_schema": "pyrit.models.target.json_schema_definition", + "read_usage_int": "pyrit.models.target.token_usage", + "read_usage_value": "pyrit.models.target.token_usage", + "register_common_json_schema": "pyrit.models.target.json_schema_definition", + "unregister_common_json_schema": "pyrit.models.target.json_schema_definition", +} + +__all__ = list(_LAZY_EXPORTS) + + +def __getattr__(name: str) -> object: + """ + Resolve a public target export on first access. + + Args: + name (str): The requested public name. + + Returns: + object: The resolved export. + """ + return resolve_lazy_export( + name=name, + module_name=__name__, + module_globals=globals(), + exports=_LAZY_EXPORTS, + ) + + +def __dir__() -> list[str]: + """Return package attributes, including unresolved exports.""" + return get_lazy_dir(module_globals=globals(), exports=_LAZY_EXPORTS) diff --git a/tests/unit/common/test_lazy_imports.py b/tests/unit/common/test_lazy_imports.py new file mode 100644 index 0000000000..172c3f2efe --- /dev/null +++ b/tests/unit/common/test_lazy_imports.py @@ -0,0 +1,76 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from types import ModuleType +from unittest.mock import patch + +import pytest + +from pyrit.common.lazy_imports import get_lazy_dir, resolve_lazy_export + + +def test_resolve_lazy_export_resolves_and_caches_attribute() -> None: + implementation = ModuleType("example.implementation") + implementation.ExportedClass = object + module_globals: dict[str, object] = {} + + with patch("pyrit.common.lazy_imports.import_module", return_value=implementation) as import_module_mock: + result = resolve_lazy_export( + name="ExportedClass", + module_name="example", + module_globals=module_globals, + exports={"ExportedClass": "example.implementation"}, + ) + + assert result is object + assert module_globals["ExportedClass"] is object + import_module_mock.assert_called_once_with("example.implementation") + + +def test_resolve_lazy_export_supports_renamed_attribute() -> None: + implementation = ModuleType("example.implementation") + implementation.OriginalName = object + + with patch("pyrit.common.lazy_imports.import_module", return_value=implementation): + result = resolve_lazy_export( + name="PublicName", + module_name="example", + module_globals={}, + exports={"PublicName": ("example.implementation", "OriginalName")}, + ) + + assert result is object + + +def test_resolve_lazy_export_supports_module_export() -> None: + implementation = ModuleType("example.implementation") + + with patch("pyrit.common.lazy_imports.import_module", return_value=implementation): + result = resolve_lazy_export( + name="implementation", + module_name="example", + module_globals={}, + exports={"implementation": ("example.implementation", None)}, + ) + + assert result is implementation + + +def test_resolve_lazy_export_rejects_unknown_name() -> None: + with pytest.raises(AttributeError, match="module 'example' has no attribute 'missing'"): + resolve_lazy_export( + name="missing", + module_name="example", + module_globals={}, + exports={}, + ) + + +def test_get_lazy_dir_includes_unresolved_exports() -> None: + result = get_lazy_dir( + module_globals={"existing": object()}, + exports={"LazyName": "example.implementation"}, + ) + + assert "existing" in result + assert "LazyName" in result diff --git a/tests/unit/common/test_lazy_package_imports.py b/tests/unit/common/test_lazy_package_imports.py new file mode 100644 index 0000000000..71e65e35c7 --- /dev/null +++ b/tests/unit/common/test_lazy_package_imports.py @@ -0,0 +1,420 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import ast +import importlib +import subprocess +import sys +from pathlib import Path +from textwrap import dedent + +import pytest + +_REPOSITORY_ROOT = Path(__file__).parents[3] +_PACKAGE_ROOT = _REPOSITORY_ROOT / "pyrit" +_MODEL_PACKAGE_NAMES = ( + "pyrit.models", + "pyrit.models.catalog", + "pyrit.models.identifiers", + "pyrit.models.messages", + "pyrit.models.results", + "pyrit.models.score", + "pyrit.models.seeds", + "pyrit.models.target", +) + +# Remove an entry when its package adopts the standard lazy export contract. +# The inventory test rejects both unlisted eager packages and stale exceptions. +_EAGER_PACKAGE_EXCEPTIONS = frozenset( + { + "pyrit", + "pyrit.analytics", + "pyrit.auth", + "pyrit.backend.mappers", + "pyrit.backend.middleware", + "pyrit.backend.models", + "pyrit.backend.routes", + "pyrit.backend.services", + "pyrit.common", + "pyrit.converter", + "pyrit.converter.token_smuggling", + "pyrit.datasets", + "pyrit.datasets.seed_datasets.local", + "pyrit.datasets.seed_datasets.remote", + "pyrit.embedding", + "pyrit.exceptions", + "pyrit.executor.attack", + "pyrit.executor.attack.component", + "pyrit.executor.attack.compound", + "pyrit.executor.attack.core", + "pyrit.executor.attack.multi_turn", + "pyrit.executor.attack.single_turn", + "pyrit.executor.attack.streaming", + "pyrit.executor.benchmark", + "pyrit.executor.core", + "pyrit.executor.promptgen", + "pyrit.executor.promptgen.core", + "pyrit.executor.promptgen.fuzzer", + "pyrit.executor.promptgen.gcg", + "pyrit.executor.workflow", + "pyrit.executor.workflow.core", + "pyrit.memory", + "pyrit.memory.storage", + "pyrit.message_normalizer", + "pyrit.output", + "pyrit.prompt_normalizer", + "pyrit.prompt_target", + "pyrit.registry", + "pyrit.registry.components", + "pyrit.scenario", + "pyrit.scenario.core", + "pyrit.scenario.scenarios.adaptive", + "pyrit.scenario.scenarios.adaptive.selectors", + "pyrit.scenario.scenarios.airt", + "pyrit.scenario.scenarios.benchmark", + "pyrit.scenario.scenarios.foundry", + "pyrit.scenario.scenarios.garak", + "pyrit.score", + "pyrit.score.true_false.regex", + "pyrit.setup", + "pyrit.setup.initializers", + "pyrit.setup.initializers.techniques", + } +) + +_LAZY_IMPORT_SPOT_CHECKS = [ + ( + "pyrit.models", + "Message", + "pyrit.models.messages.message", + "pyrit.models.question_answering", + ), + ( + "pyrit.models.catalog", + "RegisteredInitializer", + "pyrit.models.catalog.initializer", + "pyrit.models.catalog.scenario", + ), + ( + "pyrit.models.identifiers", + "validate_registry_name", + "pyrit.models.identifiers.class_name_utils", + "pyrit.models.identifiers.evaluation_identifier", + ), + ( + "pyrit.models.messages", + "Message", + "pyrit.models.messages.message", + "pyrit.models.messages.chat_message", + ), + ( + "pyrit.models.results", + "StrategyResult", + "pyrit.models.results.strategy_result", + "pyrit.models.results.scenario_result", + ), + ( + "pyrit.models.score", + "Condition", + "pyrit.models.score.condition", + "pyrit.models.score.score", + ), + ( + "pyrit.models.seeds", + "Seed", + "pyrit.models.seeds.seed", + "pyrit.models.seeds.yaml_seed_loader", + ), + ( + "pyrit.models.target", + "TokenUsage", + "pyrit.models.target.token_usage", + "pyrit.models.target.json_schema_definition", + ), +] + + +def _is_type_checking_guard(test: ast.expr) -> bool: + """Return whether an expression is a ``TYPE_CHECKING`` guard.""" + return (isinstance(test, ast.Name) and test.id == "TYPE_CHECKING") or ( + isinstance(test, ast.Attribute) + and isinstance(test.value, ast.Name) + and test.value.id == "typing" + and test.attr == "TYPE_CHECKING" + ) + + +class _RuntimeImportCollector(ast.NodeVisitor): + """Collect non-standard imports that execute when a package is imported.""" + + def __init__(self) -> None: + self.imports: list[str] = [] + + def visit_FunctionDef(self, node: ast.FunctionDef) -> None: + return + + def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> None: + return + + def visit_If(self, node: ast.If) -> None: + if not _is_type_checking_guard(node.test): + for statement in node.body: + self.visit(statement) + for statement in node.orelse: + self.visit(statement) + + def visit_Import(self, node: ast.Import) -> None: + for alias in node.names: + if alias.name == "pyrit.common.lazy_imports": + continue + if alias.name.split(".", maxsplit=1)[0] not in sys.stdlib_module_names: + self.imports.append(alias.name) + + def visit_ImportFrom(self, node: ast.ImportFrom) -> None: + module_name = node.module or "" + if node.level == 0 and module_name == "pyrit.common.lazy_imports": + return + if node.level or module_name.split(".", maxsplit=1)[0] not in sys.stdlib_module_names: + self.imports.append(f"{'.' * node.level}{module_name}") + + +def _initializer_paths() -> tuple[Path, ...]: + """Return every PyRIT package initializer.""" + return tuple(sorted(_PACKAGE_ROOT.rglob("__init__.py"))) + + +def _package_name(init_path: Path) -> str: + """Return the dotted package name for an initializer.""" + relative_parts = init_path.parent.relative_to(_PACKAGE_ROOT).parts + return ".".join(("pyrit", *relative_parts)) + + +def _assigned_value(*, tree: ast.Module, name: str) -> ast.expr | None: + """Return the module-level value assigned to a name.""" + for statement in tree.body: + if ( + isinstance(statement, ast.AnnAssign) + and isinstance(statement.target, ast.Name) + and statement.target.id == name + ): + return statement.value + if isinstance(statement, ast.Assign) and any( + isinstance(target, ast.Name) and target.id == name for target in statement.targets + ): + return statement.value + return None + + +def _is_empty_all(value: ast.expr | None) -> bool: + """Return whether ``__all__`` is an empty literal.""" + return isinstance(value, (ast.List, ast.Tuple)) and not value.elts + + +def _lazy_export_names(tree: ast.Module) -> tuple[str, ...] | None: + """Return names from a literal ``_LAZY_EXPORTS`` map.""" + value = _assigned_value(tree=tree, name="_LAZY_EXPORTS") + if not isinstance(value, ast.Dict): + return None + + names: list[str] = [] + for key in value.keys: + if not isinstance(key, ast.Constant) or not isinstance(key.value, str): + return None + names.append(key.value) + return tuple(names) + + +def _all_derives_from_lazy_exports(tree: ast.Module) -> bool: + """Return whether ``__all__`` is exactly ``list(_LAZY_EXPORTS)``.""" + value = _assigned_value(tree=tree, name="__all__") + return ( + isinstance(value, ast.Call) + and isinstance(value.func, ast.Name) + and value.func.id == "list" + and len(value.args) == 1 + and isinstance(value.args[0], ast.Name) + and value.args[0].id == "_LAZY_EXPORTS" + and not value.keywords + ) + + +def _type_checking_import_names(tree: ast.Module) -> set[str]: + """Return names imported below module-level ``TYPE_CHECKING`` guards.""" + names: set[str] = set() + for statement in tree.body: + if not isinstance(statement, ast.If) or not _is_type_checking_guard(statement.test): + continue + for node in ast.walk(ast.Module(body=statement.body, type_ignores=[])): + if isinstance(node, ast.Import): + names.update(alias.asname or alias.name.split(".")[0] for alias in node.names) + elif isinstance(node, ast.ImportFrom): + names.update(alias.asname or alias.name for alias in node.names) + return names + + +def _module_function_names(tree: ast.Module) -> set[str]: + """Return function names defined directly in a module.""" + return {statement.name for statement in tree.body if isinstance(statement, (ast.FunctionDef, ast.AsyncFunctionDef))} + + +def _lazy_helper_import_names(tree: ast.Module) -> set[str]: + """Return names imported from the shared lazy import helper.""" + names: set[str] = set() + for statement in tree.body: + if isinstance(statement, ast.ImportFrom) and statement.module == "pyrit.common.lazy_imports": + names.update(alias.asname or alias.name for alias in statement.names) + return names + + +def _runtime_non_standard_imports(tree: ast.Module) -> list[str]: + """Return non-standard imports that execute while importing a package.""" + collector = _RuntimeImportCollector() + collector.visit(tree) + return collector.imports + + +def _non_exempt_public_initializers() -> tuple[Path, ...]: + """Return public package initializers that must use the lazy contract.""" + paths: list[Path] = [] + for init_path in _initializer_paths(): + if _package_name(init_path) in _EAGER_PACKAGE_EXCEPTIONS: + continue + tree = ast.parse(init_path.read_text(encoding="utf-8")) + all_value = _assigned_value(tree=tree, name="__all__") + lazy_exports_value = _assigned_value(tree=tree, name="_LAZY_EXPORTS") + if lazy_exports_value is not None or (all_value is not None and not _is_empty_all(all_value)): + paths.append(init_path) + return tuple(paths) + + +def _assert_subprocess_succeeds(code: str) -> None: + """Run Python code in a clean interpreter and require success.""" + result = subprocess.run( + [sys.executable, "-c", dedent(code)], + capture_output=True, + text=True, + timeout=30, + ) + + assert result.returncode == 0, result.stderr + + +def test_all_package_initializers_are_lazy_or_exempt() -> None: + eager_imports: dict[str, list[str]] = {} + for init_path in _initializer_paths(): + tree = ast.parse(init_path.read_text(encoding="utf-8")) + imports = _runtime_non_standard_imports(tree) + if imports: + eager_imports[_package_name(init_path)] = imports + + eager_packages = set(eager_imports) + unexpected = sorted(eager_packages - _EAGER_PACKAGE_EXCEPTIONS) + stale = sorted(_EAGER_PACKAGE_EXCEPTIONS - eager_packages) + unexpected_details = [f"{package}: {eager_imports[package]}" for package in unexpected] + + assert not unexpected and not stale, ( + f"Unlisted eager packages: {unexpected_details}\n" + f"Stale eager-package exceptions: {stale}\n" + "Convert unlisted packages to the standard lazy export contract. " + "Remove an exception when its package no longer imports PyRIT modules eagerly." + ) + + +@pytest.mark.parametrize( + "init_path", + _non_exempt_public_initializers(), + ids=lambda path: _package_name(path), +) +def test_non_exempt_public_package_uses_standard_lazy_contract(init_path: Path) -> None: + tree = ast.parse(init_path.read_text(encoding="utf-8")) + lazy_export_names = _lazy_export_names(tree) + function_names = _module_function_names(tree) + helper_import_names = _lazy_helper_import_names(tree) + + assert lazy_export_names is not None, f"{_package_name(init_path)} must define a literal _LAZY_EXPORTS map" + assert _all_derives_from_lazy_exports(tree), f"{_package_name(init_path)} must set __all__ = list(_LAZY_EXPORTS)" + assert {"__getattr__", "__dir__"} <= function_names + assert {"get_lazy_dir", "resolve_lazy_export"} <= helper_import_names + + missing_type_imports = set(lazy_export_names) - _type_checking_import_names(tree) + assert not missing_type_imports, ( + f"{_package_name(init_path)} lacks TYPE_CHECKING imports for: {sorted(missing_type_imports)}" + ) + + +def test_lazy_packages_do_not_load_child_modules() -> None: + package_names = [_package_name(path) for path in _non_exempt_public_initializers()] + _assert_subprocess_succeeds( + f""" + import importlib + import sys + + package_names = {package_names!r} + for package_name in sorted(package_names, key=lambda name: name.count(".")): + package = importlib.import_module(package_name) + descendants = [name for name in sys.modules if name.startswith(f"{{package_name}}.")] + + assert not descendants, (package_name, descendants) + assert package.__all__ == list(package._LAZY_EXPORTS) + assert set(package._LAZY_EXPORTS) <= set(dir(package)) + """ + ) + + +@pytest.mark.parametrize("package_name", _MODEL_PACKAGE_NAMES) +def test_model_package_dir_includes_lazy_exports_in_process(package_name: str) -> None: + package = importlib.import_module(package_name) + + assert set(package.__all__) <= set(dir(package)) + + +@pytest.mark.parametrize( + ("package_name", "export_name", "declaring_module"), + [(package, export, declaring) for package, export, declaring, _ in _LAZY_IMPORT_SPOT_CHECKS], + ids=[case[0] for case in _LAZY_IMPORT_SPOT_CHECKS], +) +def test_lazy_import_resolves_and_caches_export_in_process( + package_name: str, + export_name: str, + declaring_module: str, +) -> None: + package = importlib.import_module(package_name) + package.__dict__.pop(export_name, None) + + exported_value = package.__getattr__(export_name) + direct_value = getattr(importlib.import_module(declaring_module), export_name) + + assert exported_value is direct_value + assert getattr(package, export_name) is direct_value + + +@pytest.mark.parametrize( + ("package_name", "export_name", "declaring_module", "unrelated_module"), + _LAZY_IMPORT_SPOT_CHECKS, + ids=[case[0] for case in _LAZY_IMPORT_SPOT_CHECKS], +) +def test_lazy_import_spot_check( + package_name: str, + export_name: str, + declaring_module: str, + unrelated_module: str, +) -> None: + _assert_subprocess_succeeds( + f""" + import importlib + import sys + + package = importlib.import_module({package_name!r}) + + assert {declaring_module!r} not in sys.modules + assert {unrelated_module!r} not in sys.modules + assert {export_name!r} in dir(package) + + exported_value = getattr(package, {export_name!r}) + direct_value = getattr(importlib.import_module({declaring_module!r}), {export_name!r}) + + assert exported_value is direct_value + assert getattr(package, {export_name!r}) is exported_value + assert {unrelated_module!r} not in sys.modules + """ + )