Fix attack type normalization and safe classification leak in graph
Browse files- src/db/graph_writer.py +6 -1
- src/layers/pipeline.py +5 -2
src/db/graph_writer.py
CHANGED
|
@@ -61,9 +61,14 @@ async def write_threat_events_batch(events_data: list) -> None:
|
|
| 61 |
if hasattr(timestamp, "isoformat"):
|
| 62 |
timestamp = timestamp.isoformat()
|
| 63 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 64 |
formatted_events.append({
|
| 65 |
"key_id": str(log_entry.get("api_key_id", "unknown")),
|
| 66 |
-
"attack_type":
|
| 67 |
"flagged_layer": log_entry.get("flagged_layer") or "unknown_layer",
|
| 68 |
"flagged_pattern": str(log_entry.get("flagged_pattern") or "none"),
|
| 69 |
"prompt_hash": log_entry.get("prompt_hash", "unknown_hash"),
|
|
|
|
| 61 |
if hasattr(timestamp, "isoformat"):
|
| 62 |
timestamp = timestamp.isoformat()
|
| 63 |
|
| 64 |
+
raw_attack = log_entry.get("attack_type") or "unknown_attack"
|
| 65 |
+
if raw_attack.lower() == "safe":
|
| 66 |
+
raw_attack = "cumulative_risk_exceeded"
|
| 67 |
+
normalized_attack = raw_attack.lower().replace(" ", "_")
|
| 68 |
+
|
| 69 |
formatted_events.append({
|
| 70 |
"key_id": str(log_entry.get("api_key_id", "unknown")),
|
| 71 |
+
"attack_type": normalized_attack,
|
| 72 |
"flagged_layer": log_entry.get("flagged_layer") or "unknown_layer",
|
| 73 |
"flagged_pattern": str(log_entry.get("flagged_pattern") or "none"),
|
| 74 |
"prompt_hash": log_entry.get("prompt_hash", "unknown_hash"),
|
src/layers/pipeline.py
CHANGED
|
@@ -182,7 +182,7 @@ class ClassifierPipeline:
|
|
| 182 |
return build_short_circuit(
|
| 183 |
flagged_name="rule_based",
|
| 184 |
risk=current_risk,
|
| 185 |
-
attack=rule_res.attack_category,
|
| 186 |
pattern=rule_res.matched_pattern,
|
| 187 |
running_layers=layers_data
|
| 188 |
)
|
|
@@ -263,10 +263,13 @@ class ClassifierPipeline:
|
|
| 263 |
layers_data["ml_classifier"]["reason"] = ml_res.reason
|
| 264 |
|
| 265 |
if is_ml_triggered and ml_res.ran:
|
|
|
|
|
|
|
|
|
|
| 266 |
return build_short_circuit(
|
| 267 |
flagged_name="ml_classifier",
|
| 268 |
risk=current_risk,
|
| 269 |
-
attack=
|
| 270 |
pattern=f"ML classified threat: {ml_res.attack_class} (cumulative)",
|
| 271 |
running_layers=layers_data
|
| 272 |
)
|
|
|
|
| 182 |
return build_short_circuit(
|
| 183 |
flagged_name="rule_based",
|
| 184 |
risk=current_risk,
|
| 185 |
+
attack=rule_res.attack_category.lower().replace(" ", "_") if rule_res.attack_category else "unknown_rule",
|
| 186 |
pattern=rule_res.matched_pattern,
|
| 187 |
running_layers=layers_data
|
| 188 |
)
|
|
|
|
| 263 |
layers_data["ml_classifier"]["reason"] = ml_res.reason
|
| 264 |
|
| 265 |
if is_ml_triggered and ml_res.ran:
|
| 266 |
+
attack_class = ml_res.attack_class or "unknown"
|
| 267 |
+
if attack_class.lower() == "safe":
|
| 268 |
+
attack_class = "cumulative_risk_exceeded"
|
| 269 |
return build_short_circuit(
|
| 270 |
flagged_name="ml_classifier",
|
| 271 |
risk=current_risk,
|
| 272 |
+
attack=attack_class.lower().replace(" ", "_"),
|
| 273 |
pattern=f"ML classified threat: {ml_res.attack_class} (cumulative)",
|
| 274 |
running_layers=layers_data
|
| 275 |
)
|