refactor: consolidate duplicate code and fix deadlock bug

- Extract duplicate attack/defense config merging into _merge_component_configs()
- Extract duplicate lazy loading logic into _get_component()
- Move content policy detection to BaseModel base class
- Fix BatchSaveManager deadlock by splitting flush logic
- Add TypeError to ValueError conversion for consistent config errors
- Move _determine_load_model() to BaseComponent (explicit field only)
This commit is contained in:
Liao, Jie
2026-01-23 12:55:34 +08:00
parent 84a4d1708e
commit 04a1cbe8d1
12 changed files with 384 additions and 400 deletions
+23 -39
View File
@@ -47,7 +47,11 @@ class BaseComponent(ABC):
cfg_dict = config
allowed_fields = {f.name for f in fields(self.CONFIG_CLASS)}
filtered = {k: v for k, v in cfg_dict.items() if k in allowed_fields}
self.cfg = self.CONFIG_CLASS(**filtered)
try:
self.cfg = self.CONFIG_CLASS(**filtered)
except TypeError as e:
# Convert TypeError to ValueError for consistency
raise ValueError(f"Invalid configuration: {e}") from e
else:
# No configuration class or not a dataclass, use dictionary directly
self.cfg = config
@@ -77,6 +81,24 @@ class BaseComponent(ABC):
pass
# Subclasses can add more validation logic
def _determine_load_model(self) -> bool:
"""Determine if local model needs to be loaded.
Checks the load_model field in configuration (dict or dataclass).
Returns False if load_model is not explicitly set.
"""
if hasattr(self, "cfg"):
config_obj = self.cfg
else:
config_obj = self.config
if isinstance(config_obj, dict):
return config_obj.get("load_model", False)
elif hasattr(config_obj, "load_model"):
return getattr(config_obj, "load_model", False)
return False
class BaseAttack(BaseComponent, ABC):
"""Attack method base class (enhanced version)"""
@@ -109,25 +131,6 @@ class BaseAttack(BaseComponent, ABC):
# Determine if local model needs to be loaded
self.load_model = self._determine_load_model()
def _determine_load_model(self) -> bool:
"""Determine if local model needs to be loaded
Only decide based on the load_model field in configuration
"""
# Check if load_model field exists in configuration
if hasattr(self, "cfg"):
config_obj = self.cfg
else:
config_obj = self.config
# Only check load_model configuration item
if isinstance(config_obj, dict):
return config_obj.get("load_model", False)
elif hasattr(config_obj, "load_model"):
return getattr(config_obj, "load_model", False)
return False
@abstractmethod
def generate_test_case(
self, original_prompt: str, image_path: str, case_id: str, **kwargs
@@ -259,25 +262,6 @@ class BaseDefense(BaseComponent, ABC):
# Determine if local model needs to be loaded
self.load_model = self._determine_load_model()
def _determine_load_model(self) -> bool:
"""Determine if local model needs to be loaded
Only based on the load_model field in configuration
"""
# Check if load_model field exists in configuration
if hasattr(self, "cfg"):
config_obj = self.cfg
else:
config_obj = self.config
# Only check load_model configuration item
if isinstance(config_obj, dict):
return config_obj.get("load_model", False)
elif hasattr(config_obj, "load_model"):
return getattr(config_obj, "load_model", False)
return False
@abstractmethod
def apply_defense(self, test_case: TestCase, **kwargs) -> TestCase:
"""