mirror of
https://github.com/msoedov/agentic_security.git
synced 2026-09-29 03:11:51 +02:00
Merge pull request #326 from chiruu12/fix/hybrid-classifier-guards
fix(hybrid): guard zero total weight and log skipped detectors
This commit is contained in:
@@ -7,6 +7,8 @@ refusal classification with reduced false positives/negatives.
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
|
||||
from agentic_security.logutils import logger
|
||||
|
||||
|
||||
class RefusalDetector(Protocol):
|
||||
"""Protocol for refusal detection methods."""
|
||||
@@ -118,7 +120,13 @@ class HybridRefusalClassifier:
|
||||
try:
|
||||
is_refusal = config.detector.is_refusal(response)
|
||||
except Exception:
|
||||
continue # Skip failed detectors
|
||||
# Skip failed detectors, but say so. A detector that raises on
|
||||
# every response drops out of the vote and the rest renormalise
|
||||
# to full confidence, which reads the same as agreement.
|
||||
logger.exception(
|
||||
f"Detector {config.name!r} raised and was skipped for this response."
|
||||
)
|
||||
continue
|
||||
results.append(
|
||||
DetectionResult(
|
||||
method=config.name,
|
||||
@@ -133,6 +141,16 @@ class HybridRefusalClassifier:
|
||||
total_weight = sum(r.weight for r in results)
|
||||
refusal_weight = sum(r.weight for r in results if r.is_refusal)
|
||||
|
||||
# Every detector carrying weight 0 leaves nothing to divide by. Treat it
|
||||
# the same as having no detectors rather than raising.
|
||||
if total_weight <= 0:
|
||||
logger.warning(
|
||||
"All detectors have zero total weight, cannot score this response."
|
||||
)
|
||||
return HybridResult(
|
||||
is_refusal=False, confidence=0.0, method_results=results
|
||||
)
|
||||
|
||||
# Calculate confidence as how strongly detectors agree
|
||||
raw_score = refusal_weight / total_weight # 0.0-1.0, 1.0 = all say refusal
|
||||
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Unit tests for hybrid refusal classifier."""
|
||||
|
||||
import logging
|
||||
|
||||
from inline_snapshot import snapshot
|
||||
|
||||
from agentic_security.refusal_classifier.hybrid_classifier import (
|
||||
@@ -320,3 +322,53 @@ class TestConfidenceScoring:
|
||||
# 2/3 = 0.666 confidence for non-refusal
|
||||
assert round(result.confidence, 2) == snapshot(0.67)
|
||||
assert result.is_refusal is False
|
||||
|
||||
|
||||
class _AlwaysRaises:
|
||||
def is_refusal(self, response: str) -> bool:
|
||||
raise RuntimeError("detector is broken")
|
||||
|
||||
|
||||
class _Fixed:
|
||||
def __init__(self, value: bool):
|
||||
self.value = value
|
||||
|
||||
def is_refusal(self, response: str) -> bool:
|
||||
return self.value
|
||||
|
||||
|
||||
class TestClassifyEdgeCases:
|
||||
def test_zero_total_weight_does_not_raise(self):
|
||||
"""Weights summing to zero left nothing to divide by."""
|
||||
classifier = HybridRefusalClassifier()
|
||||
classifier.add_detector(_Fixed(True), weight=0.0, name="a")
|
||||
classifier.add_detector(_Fixed(False), weight=0.0, name="b")
|
||||
|
||||
result = classifier.classify("I cannot help with that")
|
||||
|
||||
assert result.is_refusal is False
|
||||
assert result.confidence == 0.0
|
||||
assert len(result.method_results) == 2
|
||||
|
||||
def test_raising_detector_is_skipped_and_logged(self, caplog):
|
||||
"""A detector that raises drops out, so the drop should be visible."""
|
||||
classifier = HybridRefusalClassifier()
|
||||
classifier.add_detector(_AlwaysRaises(), name="broken")
|
||||
classifier.add_detector(_Fixed(True), name="good")
|
||||
|
||||
with caplog.at_level(logging.ERROR):
|
||||
result = classifier.classify("I cannot help with that")
|
||||
|
||||
assert [r.method for r in result.method_results] == ["good"]
|
||||
assert result.is_refusal is True
|
||||
assert any("broken" in record.message for record in caplog.records)
|
||||
|
||||
def test_normal_weighting_is_unchanged(self):
|
||||
classifier = HybridRefusalClassifier()
|
||||
classifier.add_detector(_Fixed(True), name="a")
|
||||
classifier.add_detector(_Fixed(False), name="b")
|
||||
|
||||
result = classifier.classify("I cannot help with that")
|
||||
|
||||
assert result.is_refusal is True
|
||||
assert result.confidence == 0.5
|
||||
|
||||
Reference in New Issue
Block a user