fix(hybrid): guard zero total weight and log skipped detectors

This commit is contained in:
chiruu12 committed 2026-08-14 17:04:47 +05:30
1 parent c8458d73c5
commit 2f411b90dc
2 files changed
+71 -1

No files matched your search

@@ -7,6 +7,8 @@ refusal classification with reduced false positives/negatives.
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Protocol from typing import Protocol
from agentic_security.logutils import logger
class RefusalDetector(Protocol): class RefusalDetector(Protocol):
"""Protocol for refusal detection methods.""" """Protocol for refusal detection methods."""
@@ -118,7 +120,13 @@ class HybridRefusalClassifier:
try: try:
is_refusal = config.detector.is_refusal(response) is_refusal = config.detector.is_refusal(response)
except Exception: 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( results.append(
DetectionResult( DetectionResult(
method=config.name, method=config.name,
@@ -133,6 +141,16 @@ class HybridRefusalClassifier:
total_weight = sum(r.weight for r in results) total_weight = sum(r.weight for r in results)
refusal_weight = sum(r.weight for r in results if r.is_refusal) 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 # Calculate confidence as how strongly detectors agree
raw_score = refusal_weight / total_weight # 0.0-1.0, 1.0 = all say refusal 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.""" """Unit tests for hybrid refusal classifier."""
import logging
from inline_snapshot import snapshot from inline_snapshot import snapshot
from agentic_security.refusal_classifier.hybrid_classifier import ( from agentic_security.refusal_classifier.hybrid_classifier import (
@@ -320,3 +322,53 @@ class TestConfidenceScoring:
# 2/3 = 0.666 confidence for non-refusal # 2/3 = 0.666 confidence for non-refusal
assert round(result.confidence, 2) == snapshot(0.67) assert round(result.confidence, 2) == snapshot(0.67)
assert result.is_refusal is False 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