imDrizzle commited on
Commit
08e7bbd
·
1 Parent(s): d468fc4

Fix attack type normalization and safe classification leak in graph

Browse files
Files changed (2) hide show
  1. src/db/graph_writer.py +6 -1
  2. 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": log_entry.get("attack_type") or "unknown_attack",
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=ml_res.attack_class,
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
  )