fix: harden forensic report ingestion

This commit is contained in:
Christian Krakau-Louis
2026-05-23 13:02:01 +02:00
parent 908d4cd2cd
commit 817270b1fb
8 changed files with 89 additions and 23 deletions
@@ -95,7 +95,12 @@ async def upload_forensic_report(
detail="Forensic report has already been uploaded.", detail="Forensic report has already been uploaded.",
) )
row, _created = save_forensic_report(db, parsed) row, created = save_forensic_report(db, parsed)
if not created:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail="Forensic report has already been uploaded.",
)
db.commit() db.commit()
db.refresh(row) db.refresh(row)
return ForensicUploadResponse( return ForensicUploadResponse(
+3 -1
View File
@@ -1,7 +1,7 @@
import hashlib import hashlib
import json import json
import re import re
from datetime import datetime from datetime import datetime, timezone
from email import message_from_bytes from email import message_from_bytes
from email.message import Message from email.message import Message
from email.parser import Parser from email.parser import Parser
@@ -84,6 +84,8 @@ def _parse_datetime(value: str) -> Optional[datetime]:
return None return None
if parsed is None: if parsed is None:
return None return None
if parsed.tzinfo is not None:
parsed = parsed.astimezone(timezone.utc)
return parsed.replace(tzinfo=None) return parsed.replace(tzinfo=None)
+21 -3
View File
@@ -1,6 +1,7 @@
import json import json
from typing import Any, Dict, Optional from typing import Any, Dict, Optional
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.models.domain import Domain from app.models.domain import Domain
@@ -10,7 +11,7 @@ from app.utils.domain_validator import DomainValidationError, validate_domain
def forensic_report_exists(db: Session, report_id: str) -> bool: def forensic_report_exists(db: Session, report_id: str) -> bool:
"""Return True when a forensic report ID is already persisted.""" """Return True when a forensic report ID is already persisted."""
if not report_id: if not str(report_id or "").strip():
return False return False
return ( return (
db.query(ForensicReport.id).filter(ForensicReport.report_id == report_id).first() db.query(ForensicReport.id).filter(ForensicReport.report_id == report_id).first()
@@ -40,7 +41,10 @@ def save_forensic_report(db: Session, report: Dict[str, Any]) -> tuple[ForensicR
Returns ``(row, created)``. The caller owns the transaction and should Returns ``(row, created)``. The caller owns the transaction and should
commit after related work has completed. commit after related work has completed.
""" """
report_id = str(report.get("report_id") or "") report_id = str(report.get("report_id") or "").strip()
if not report_id:
raise ValueError("Forensic report_id is required")
existing = db.query(ForensicReport).filter(ForensicReport.report_id == report_id).first() existing = db.query(ForensicReport).filter(ForensicReport.report_id == report_id).first()
if existing is not None: if existing is not None:
return existing, False return existing, False
@@ -72,12 +76,26 @@ def save_forensic_report(db: Session, report: Dict[str, Any]) -> tuple[ForensicR
feedback_headers=feedback_headers, feedback_headers=feedback_headers,
) )
db.add(row) db.add(row)
try:
db.flush() db.flush()
except IntegrityError:
db.rollback()
existing = db.query(ForensicReport).filter(ForensicReport.report_id == report_id).first()
if existing is not None:
return existing, False
raise
return row, True return row, True
def forensic_report_to_dict(row: ForensicReport) -> Dict[str, Any]: def forensic_report_to_dict(row: ForensicReport) -> Dict[str, Any]:
"""Convert a forensic report row to an API-safe dictionary.""" """Convert a forensic report row to an API-safe dictionary."""
feedback_headers = {}
if row.feedback_headers:
try:
feedback_headers = json.loads(row.feedback_headers)
except (json.JSONDecodeError, TypeError):
feedback_headers = {}
return { return {
"id": row.id, "id": row.id,
"report_id": row.report_id, "report_id": row.report_id,
@@ -98,6 +116,6 @@ def forensic_report_to_dict(row: ForensicReport) -> Dict[str, Any]:
"original_subject": row.original_subject, "original_subject": row.original_subject,
"original_message_id": row.original_message_id, "original_message_id": row.original_message_id,
"original_date": row.original_date, "original_date": row.original_date,
"feedback_headers": json.loads(row.feedback_headers) if row.feedback_headers else {}, "feedback_headers": feedback_headers,
"processed_at": row.processed_at.isoformat() if row.processed_at else None, "processed_at": row.processed_at.isoformat() if row.processed_at else None,
} }
+14 -13
View File
@@ -59,6 +59,7 @@ DMARC_GMAIL_QUERY = (
# How many message results to fetch per API page # How many message results to fetch per API page
_PAGE_SIZE = 100 _PAGE_SIZE = 100
RETRYABLE_MESSAGE_FAILURE = -1
class GmailClient: class GmailClient:
@@ -247,11 +248,9 @@ class GmailClient:
stats["processed"] += 1 stats["processed"] += 1
found = self._process_message(service, msg_id, stats) found = self._process_message(service, msg_id, stats)
if found > 0: if found >= 0:
stats["new_ingested_ids"].append(msg_id) # Track it even when no report is found so we don't re-examine
self.already_ingested_ids.append(msg_id) # unrelated messages on every poll. Retryable failures return -1.
else:
# Track it anyway so we don't re-examine it next run
stats["new_ingested_ids"].append(msg_id) stats["new_ingested_ids"].append(msg_id)
self.already_ingested_ids.append(msg_id) self.already_ingested_ids.append(msg_id)
@@ -334,7 +333,7 @@ class GmailClient:
raw_bytes = base64.urlsafe_b64decode(msg_data.get("raw", "")) raw_bytes = base64.urlsafe_b64decode(msg_data.get("raw", ""))
msg = email.message_from_bytes(raw_bytes) msg = email.message_from_bytes(raw_bytes)
if ForensicParser.is_forensic_report(msg): if ForensicParser.is_forensic_report(msg):
return 1 if self._process_forensic_message(raw_bytes, stats, message_id=msg_id) else 0 return self._process_forensic_message(raw_bytes, stats, message_id=msg_id)
return self._process_attachments(msg, stats, message_id=msg_id) return self._process_attachments(msg, stats, message_id=msg_id)
@staticmethod @staticmethod
@@ -383,7 +382,7 @@ class GmailClient:
raw_bytes: bytes, raw_bytes: bytes,
stats: dict, stats: dict,
message_id: Optional[str] = None, message_id: Optional[str] = None,
) -> bool: ) -> int:
"""Parse and persist one DMARC forensic report message.""" """Parse and persist one DMARC forensic report message."""
try: try:
report = ForensicParser.parse_bytes(raw_bytes, message_id_hint=message_id) report = ForensicParser.parse_bytes(raw_bytes, message_id_hint=message_id)
@@ -399,7 +398,7 @@ class GmailClient:
domain=domain, domain=domain,
report_id=report_id, report_id=report_id,
) )
return False return RETRYABLE_MESSAGE_FAILURE
if forensic_report_exists(self.db, report_id): if forensic_report_exists(self.db, report_id):
stats["duplicate_forensic_reports"] = stats.get("duplicate_forensic_reports", 0) + 1 stats["duplicate_forensic_reports"] = stats.get("duplicate_forensic_reports", 0) + 1
@@ -411,7 +410,7 @@ class GmailClient:
domain=domain, domain=domain,
report_id=report_id, report_id=report_id,
) )
return False return 0
_row, created = save_forensic_report(self.db, report) _row, created = save_forensic_report(self.db, report)
if not created: if not created:
@@ -424,7 +423,7 @@ class GmailClient:
domain=domain, domain=domain,
report_id=report_id, report_id=report_id,
) )
return False return 0
stats["forensic_reports_found"] = stats.get("forensic_reports_found", 0) + 1 stats["forensic_reports_found"] = stats.get("forensic_reports_found", 0) + 1
self._append_detail( self._append_detail(
@@ -435,10 +434,12 @@ class GmailClient:
domain=domain, domain=domain,
report_id=report_id, report_id=report_id,
) )
return True return 1
except Exception as exc: # pylint: disable=broad-exception-caught except Exception as exc: # pylint: disable=broad-exception-caught
logger.error("Failed to parse Gmail forensic report %s: %s", message_id, exc) logger.error("Failed to parse Gmail forensic report %s: %s", message_id, exc)
stats["errors"].append(f"Failed to parse forensic report {message_id}: {exc}") stats.setdefault("errors", []).append(
f"Failed to parse forensic report {message_id}: {exc}"
)
self._append_detail( self._append_detail(
stats, stats,
status="error", status="error",
@@ -446,7 +447,7 @@ class GmailClient:
message_id=message_id, message_id=message_id,
error=str(exc), error=str(exc),
) )
return False return RETRYABLE_MESSAGE_FAILURE
def _process_attachments( def _process_attachments(
self, self,
+12
View File
@@ -92,6 +92,18 @@ def test_parse_forensic_email_handles_invalid_dates():
assert parsed["arrival_date"] is None assert parsed["arrival_date"] is None
def test_parse_forensic_email_normalizes_offset_dates_to_utc():
content = SAMPLE_FORENSIC_EMAIL.replace(
b"Arrival-Date: Fri, 22 May 2026 10:15:00 +0000",
b"Arrival-Date: Fri, 22 May 2026 12:15:00 +0200",
)
parsed = ForensicParser.parse_bytes(content)
assert parsed["arrival_date"].hour == 10
assert parsed["arrival_date"].tzinfo is None
def test_parse_forensic_email_falls_back_to_dkim_domain_and_content_hash(): def test_parse_forensic_email_falls_back_to_dkim_domain_and_content_hash():
content = SAMPLE_FORENSIC_EMAIL.replace(b"Message-ID: <report-1@example.net>\n", b"") content = SAMPLE_FORENSIC_EMAIL.replace(b"Message-ID: <report-1@example.net>\n", b"")
content = content.replace(b"Reported-Domain: example.com\n", b"DKIM-Domain: fallback.test\n") content = content.replace(b"Reported-Domain: example.com\n", b"DKIM-Domain: fallback.test\n")
+14
View File
@@ -1,3 +1,5 @@
import pytest
from app.models.report import ForensicReport from app.models.report import ForensicReport
from app.services.forensic_parser import ForensicParser from app.services.forensic_parser import ForensicParser
from app.services.forensic_persistence import ( from app.services.forensic_persistence import (
@@ -19,6 +21,10 @@ def test_upload_forensic_report_persists_redacted_metadata(authed_client, db_ses
assert data["success"] is True assert data["success"] is True
assert data["domain"] == "example.com" assert data["domain"] == "example.com"
assert db_session.query(ForensicReport).count() == 1 assert db_session.query(ForensicReport).count() == 1
report = db_session.query(ForensicReport).one()
assert report.original_mail_from == "al***@example.com"
assert "original-message@example.com" not in report.original_message_id
assert not report.original_message_id.startswith("<")
def test_upload_forensic_report_rejects_duplicates(authed_client): def test_upload_forensic_report_rejects_duplicates(authed_client):
@@ -99,6 +105,14 @@ def test_save_forensic_report_duplicate_and_invalid_domain_paths(db_session):
assert forensic_report_exists(db_session, "") is False assert forensic_report_exists(db_session, "") is False
assert forensic_report_to_dict(first)["feedback_headers"] == {"identity_alignment": "dkim"} assert forensic_report_to_dict(first)["feedback_headers"] == {"identity_alignment": "dkim"}
first.feedback_headers = "{not-json"
assert forensic_report_to_dict(first)["feedback_headers"] == {}
missing_id = dict(parsed)
missing_id["report_id"] = " "
with pytest.raises(ValueError, match="report_id"):
save_forensic_report(db_session, missing_id)
invalid = dict(parsed) invalid = dict(parsed)
invalid["report_id"] = "ruf-invalid-domain" invalid["report_id"] = "ruf-invalid-domain"
invalid["reported_domain"] = "bad domain" invalid["reported_domain"] = "bad domain"
+17 -4
View File
@@ -20,7 +20,7 @@ from unittest.mock import MagicMock, patch
import pytest import pytest
from app.models.report import DMARCReport, ForensicReport from app.models.report import DMARCReport, ForensicReport
from app.services.gmail_client import GmailClient from app.services.gmail_client import GmailClient, RETRYABLE_MESSAGE_FAILURE
from app.services.report_store import ReportStore from app.services.report_store import ReportStore
from app.tests.test_data import SAMPLE_XML from app.tests.test_data import SAMPLE_XML
from app.tests.test_forensic_parser import SAMPLE_FORENSIC_EMAIL from app.tests.test_forensic_parser import SAMPLE_FORENSIC_EMAIL
@@ -475,7 +475,7 @@ class TestProcessMessage:
message_id="msg-forensic", message_id="msg-forensic",
) )
assert imported is False assert imported == RETRYABLE_MESSAGE_FAILURE
assert stats["details"][0]["reason"] == "forensic_report_requires_database" assert stats["details"][0]["reason"] == "forensic_report_requires_database"
def test_duplicate_forensic_report_is_skipped(self, db_session): def test_duplicate_forensic_report_is_skipped(self, db_session):
@@ -489,7 +489,7 @@ class TestProcessMessage:
message_id="msg-1", message_id="msg-1",
) )
assert imported is False assert imported == 0
assert stats["duplicate_forensic_reports"] == 1 assert stats["duplicate_forensic_reports"] == 1
assert stats["details"][1]["status"] == "duplicate" assert stats["details"][1]["status"] == "duplicate"
@@ -503,7 +503,7 @@ class TestProcessMessage:
message_id="bad", message_id="bad",
) )
assert imported is False assert imported == RETRYABLE_MESSAGE_FAILURE
assert stats["errors"] assert stats["errors"]
assert stats["details"][0]["reason"] == "forensic_parse_failed" assert stats["details"][0]["reason"] == "forensic_parse_failed"
@@ -743,6 +743,19 @@ class TestFetchReports:
assert "id1" in result["new_ingested_ids"] assert "id1" in result["new_ingested_ids"]
assert "id2" in result["new_ingested_ids"] assert "id2" in result["new_ingested_ids"]
def test_retryable_processing_failure_is_not_marked_ingested(self):
client = _make_client()
mock_service = MagicMock()
with (
patch.object(client, "_build_service", return_value=mock_service),
patch.object(client, "_list_dmarc_message_ids", return_value=["id1"]),
patch.object(client, "_process_message", return_value=RETRYABLE_MESSAGE_FAILURE),
):
result = client.fetch_reports()
assert result["new_ingested_ids"] == []
assert client.already_ingested_ids == []
def test_reports_new_domains(self): def test_reports_new_domains(self):
"""fetch_reports should report domains that appear after ingestion.""" """fetch_reports should report domains that appear after ingestion."""
from app.services.report_store import ReportStore from app.services.report_store import ReportStore
+1
View File
@@ -618,6 +618,7 @@ class TestProcessSingleEmail:
assert stats["details"][0]["status"] == "imported" assert stats["details"][0]["status"] == "imported"
def test_processes_forensic_report_without_aggregate_count(self, db_session): def test_processes_forensic_report_without_aggregate_count(self, db_session):
ReportStore.get_instance().clear()
client = self._make_client(db=db_session) client = self._make_client(db=db_session)
mock_mail = MagicMock() mock_mail = MagicMock()
mock_mail.fetch.return_value = ("OK", [(b"1", SAMPLE_FORENSIC_EMAIL)]) mock_mail.fetch.return_value = ("OK", [(b"1", SAMPLE_FORENSIC_EMAIL)])