feat: persist imported DMARC reports
This commit is contained in:
@@ -21,6 +21,7 @@ from googleapiclient.discovery import build
|
||||
from googleapiclient.errors import HttpError
|
||||
|
||||
from app.services.dmarc_parser import DMARCParser
|
||||
from app.services.report_persistence import report_exists, save_parsed_report
|
||||
from app.services.report_store import ReportStore
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -74,12 +75,14 @@ class GmailClient:
|
||||
access_token: str,
|
||||
refresh_token: str,
|
||||
already_ingested_ids: Optional[List[str]] = None,
|
||||
db: Any = None,
|
||||
):
|
||||
self.client_id = client_id
|
||||
self.client_secret = client_secret
|
||||
self._initial_access_token = access_token
|
||||
self.already_ingested_ids: List[str] = list(already_ingested_ids or [])
|
||||
self.report_store = ReportStore.get_instance()
|
||||
self.db = db
|
||||
|
||||
self.credentials = Credentials(
|
||||
token=access_token,
|
||||
@@ -337,10 +340,15 @@ class GmailClient:
|
||||
"""Store a parsed report unless that domain/report ID is already present."""
|
||||
domain = report.get("domain", "unknown")
|
||||
report_id = report.get("report_id", "")
|
||||
if report_id and self.report_store.has_report(domain, report_id):
|
||||
if report_id and (
|
||||
self.report_store.has_report(domain, report_id)
|
||||
or (self.db is not None and report_exists(self.db, domain, report_id))
|
||||
):
|
||||
logger.info("Skipping duplicate DMARC report %s for %s", report_id, domain)
|
||||
return False
|
||||
|
||||
if self.db is not None:
|
||||
save_parsed_report(self.db, report)
|
||||
self.report_store.add_report(report)
|
||||
return True
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ from typing import Any, Dict, Tuple
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.services.dmarc_parser import DMARCParser
|
||||
from app.services.report_persistence import report_exists, save_parsed_report
|
||||
from app.services.report_store import ReportStore
|
||||
|
||||
# Setup logger
|
||||
@@ -25,6 +26,7 @@ class IMAPClient:
|
||||
username: str = None,
|
||||
password: str = None,
|
||||
delete_emails: bool = False,
|
||||
db: Any = None,
|
||||
):
|
||||
"""
|
||||
Initialize the IMAP client with credentials
|
||||
@@ -35,6 +37,7 @@ class IMAPClient:
|
||||
username: IMAP username (if None, uses settings)
|
||||
password: IMAP password (if None, uses settings)
|
||||
delete_emails: Whether to delete emails after processing (default: False)
|
||||
db: Optional SQLAlchemy session used to persist imported reports
|
||||
"""
|
||||
settings = get_settings()
|
||||
|
||||
@@ -43,6 +46,7 @@ class IMAPClient:
|
||||
self.username = username or settings.IMAP_USERNAME
|
||||
self.password = password or settings.IMAP_PASSWORD
|
||||
self.delete_emails = delete_emails
|
||||
self.db = db
|
||||
|
||||
self.report_store = ReportStore.get_instance()
|
||||
|
||||
@@ -389,7 +393,13 @@ class IMAPClient:
|
||||
|
||||
domain = report.get("domain", "unknown")
|
||||
report_id = report.get("report_id", "")
|
||||
if report_id and self.report_store.has_report(domain, report_id):
|
||||
if report_id and (
|
||||
self.report_store.has_report(domain, report_id)
|
||||
or (
|
||||
self.db is not None
|
||||
and report_exists(self.db, domain, report_id)
|
||||
)
|
||||
):
|
||||
logger.info(
|
||||
"Skipping duplicate DMARC report %s for %s",
|
||||
report_id,
|
||||
@@ -402,6 +412,8 @@ class IMAPClient:
|
||||
continue
|
||||
|
||||
# Add the report to the store
|
||||
if self.db is not None:
|
||||
save_parsed_report(self.db, report)
|
||||
self.report_store.add_report(report)
|
||||
|
||||
reports_found += 1
|
||||
|
||||
@@ -0,0 +1,236 @@
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from app.models.domain import Domain
|
||||
from app.models.report import DMARCReport, ReportRecord
|
||||
from app.services.report_store import ReportStore
|
||||
|
||||
|
||||
def _parse_timestamp(value: Any) -> int:
|
||||
"""Return a Unix timestamp from an int-like or ISO date value."""
|
||||
if value in (None, ""):
|
||||
return 0
|
||||
if isinstance(value, (int, float)):
|
||||
return int(value)
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
try:
|
||||
return int(datetime.fromisoformat(str(value)).timestamp())
|
||||
except (TypeError, ValueError):
|
||||
return 0
|
||||
|
||||
|
||||
def _iso_from_timestamp(value: int) -> str:
|
||||
if not value:
|
||||
return ""
|
||||
return datetime.fromtimestamp(value).isoformat()
|
||||
|
||||
|
||||
def _loads_json_list(value: Optional[str]) -> Optional[List[Dict[str, Any]]]:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
decoded = json.loads(value)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return None
|
||||
return decoded if isinstance(decoded, list) else None
|
||||
|
||||
|
||||
def _policy_parts(report: Dict[str, Any]) -> Dict[str, Any]:
|
||||
policy = report.get("policy") or {}
|
||||
if isinstance(policy, str):
|
||||
return {"p": policy, "sp": "", "pct": "100"}
|
||||
if not isinstance(policy, dict):
|
||||
return {"p": "none", "sp": "", "pct": "100"}
|
||||
return {
|
||||
"p": policy.get("p", "none"),
|
||||
"sp": policy.get("sp", ""),
|
||||
"pct": str(policy.get("pct", "100")),
|
||||
"adkim": policy.get("adkim") or report.get("adkim"),
|
||||
"aspf": policy.get("aspf") or report.get("aspf"),
|
||||
}
|
||||
|
||||
|
||||
def report_exists(db: Session, domain_name: str, report_id: str) -> bool:
|
||||
"""Return True when the domain/report ID pair is already persisted."""
|
||||
if not report_id:
|
||||
return False
|
||||
return (
|
||||
db.query(DMARCReport.id)
|
||||
.join(Domain, DMARCReport.domain_id == Domain.id)
|
||||
.filter(Domain.name == domain_name, DMARCReport.report_id == report_id)
|
||||
.first()
|
||||
is not None
|
||||
)
|
||||
|
||||
|
||||
def save_parsed_report(db: Session, report: Dict[str, Any]) -> tuple[DMARCReport, bool]:
|
||||
"""Persist a parsed DMARC report and its records.
|
||||
|
||||
Returns ``(row, created)``. The caller owns the transaction and should
|
||||
commit after all related work has completed.
|
||||
"""
|
||||
domain_name = report.get("domain") or "unknown"
|
||||
report_id = report.get("report_id") or ""
|
||||
policy = _policy_parts(report)
|
||||
|
||||
domain = db.query(Domain).filter(Domain.name == domain_name).first()
|
||||
if domain is None:
|
||||
domain = Domain(name=domain_name, dmarc_policy=policy["p"])
|
||||
db.add(domain)
|
||||
db.flush()
|
||||
elif policy.get("p"):
|
||||
domain.dmarc_policy = policy["p"]
|
||||
|
||||
existing = (
|
||||
db.query(DMARCReport)
|
||||
.filter(DMARCReport.domain_id == domain.id, DMARCReport.report_id == report_id)
|
||||
.first()
|
||||
)
|
||||
if existing is not None:
|
||||
return existing, False
|
||||
|
||||
begin_ts = _parse_timestamp(report.get("begin_timestamp") or report.get("begin_date"))
|
||||
end_ts = _parse_timestamp(report.get("end_timestamp") or report.get("end_date"))
|
||||
pct = _parse_timestamp(policy.get("pct")) or 100
|
||||
|
||||
db_report = DMARCReport(
|
||||
domain_id=domain.id,
|
||||
report_id=report_id,
|
||||
org_name=report.get("org_name") or "",
|
||||
begin_date=begin_ts,
|
||||
end_date=end_ts,
|
||||
source_email=report.get("email") or report.get("source_email"),
|
||||
policy=policy["p"],
|
||||
subdomain_policy=policy.get("sp") or None,
|
||||
adkim=policy.get("adkim") or None,
|
||||
aspf=policy.get("aspf") or None,
|
||||
percentage=pct,
|
||||
)
|
||||
db.add(db_report)
|
||||
db.flush()
|
||||
|
||||
for record in report.get("records", []):
|
||||
db.add(
|
||||
ReportRecord(
|
||||
report_id=db_report.id,
|
||||
source_ip=record.get("source_ip") or "unknown",
|
||||
count=int(record.get("count") or 0),
|
||||
disposition=record.get("disposition") or "none",
|
||||
dkim=record.get("dkim_result") or record.get("dkim") or "unknown",
|
||||
spf=record.get("spf_result") or record.get("spf") or "unknown",
|
||||
header_from=record.get("header_from"),
|
||||
envelope_from=record.get("envelope_from"),
|
||||
dkim_auth_details=(
|
||||
json.dumps(record.get("dkim")) if isinstance(record.get("dkim"), list) else None
|
||||
),
|
||||
spf_auth_details=(
|
||||
json.dumps(record.get("spf")) if isinstance(record.get("spf"), list) else None
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
return db_report, True
|
||||
|
||||
|
||||
def persisted_report_to_dict(report: DMARCReport) -> Dict[str, Any]:
|
||||
"""Convert persisted report rows into the parsed-report shape used by the UI."""
|
||||
records: List[Dict[str, Any]] = []
|
||||
total_count = 0
|
||||
passed_count = 0
|
||||
|
||||
for record in report.records:
|
||||
count = int(record.count or 0)
|
||||
dkim_result = record.dkim or "unknown"
|
||||
spf_result = record.spf or "unknown"
|
||||
total_count += count
|
||||
if dkim_result == "pass" or spf_result == "pass":
|
||||
passed_count += count
|
||||
|
||||
records.append(
|
||||
{
|
||||
"source_ip": record.source_ip,
|
||||
"count": count,
|
||||
"disposition": record.disposition or "none",
|
||||
"dkim_result": dkim_result,
|
||||
"spf_result": spf_result,
|
||||
"header_from": record.header_from or "",
|
||||
"dkim": _loads_json_list(record.dkim_auth_details),
|
||||
"spf": _loads_json_list(record.spf_auth_details),
|
||||
}
|
||||
)
|
||||
|
||||
failed_count = total_count - passed_count
|
||||
pass_rate = round(passed_count / total_count * 100, 1) if total_count > 0 else 0.0
|
||||
return {
|
||||
"domain": report.domain.name if report.domain else "unknown",
|
||||
"report_id": report.report_id,
|
||||
"org_name": report.org_name,
|
||||
"email": report.source_email or "",
|
||||
"begin_date": _iso_from_timestamp(report.begin_date),
|
||||
"end_date": _iso_from_timestamp(report.end_date),
|
||||
"begin_timestamp": report.begin_date,
|
||||
"end_timestamp": report.end_date,
|
||||
"policy": {
|
||||
"p": report.policy or "none",
|
||||
"sp": report.subdomain_policy or "",
|
||||
"pct": str(report.percentage or 100),
|
||||
},
|
||||
"records": records,
|
||||
"summary": {
|
||||
"total_count": total_count,
|
||||
"passed_count": passed_count,
|
||||
"failed_count": failed_count,
|
||||
"pass_rate": pass_rate,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def hydrate_report_store_from_db(db: Session, store: ReportStore | None = None) -> int:
|
||||
"""Load persisted reports into ReportStore when the database has report rows."""
|
||||
report_count = db.query(DMARCReport.id).count()
|
||||
if report_count == 0:
|
||||
return 0
|
||||
|
||||
store = store or ReportStore.get_instance()
|
||||
store.clear()
|
||||
reports = (
|
||||
db.query(DMARCReport)
|
||||
.options(
|
||||
selectinload(DMARCReport.domain),
|
||||
selectinload(DMARCReport.records),
|
||||
)
|
||||
.order_by(DMARCReport.end_date.desc())
|
||||
.all()
|
||||
)
|
||||
for report in reports:
|
||||
store.add_report(persisted_report_to_dict(report))
|
||||
return len(reports)
|
||||
|
||||
|
||||
def delete_persisted_report(db: Session, domain_name: str, report_id: str) -> bool:
|
||||
"""Delete a persisted report by domain/report ID."""
|
||||
report = (
|
||||
db.query(DMARCReport)
|
||||
.join(Domain, DMARCReport.domain_id == Domain.id)
|
||||
.filter(Domain.name == domain_name, DMARCReport.report_id == report_id)
|
||||
.first()
|
||||
)
|
||||
if report is None:
|
||||
return False
|
||||
db.delete(report)
|
||||
return True
|
||||
|
||||
|
||||
def delete_persisted_domain(db: Session, domain_name: str) -> bool:
|
||||
"""Delete a domain row and cascaded report data."""
|
||||
domain = db.query(Domain).filter(Domain.name == domain_name).first()
|
||||
if domain is None:
|
||||
return False
|
||||
db.delete(domain)
|
||||
return True
|
||||
Reference in New Issue
Block a user