feat: persist imported DMARC reports

This commit is contained in:
Christian Krakau-Louis
2026-05-22 19:38:03 +02:00
parent 1fdfa016ed
commit b9041012de
14 changed files with 492 additions and 82 deletions
+34 -13
View File
@@ -15,6 +15,10 @@ from app.services.dns_resolver import (
extract_dmarc_policy,
get_default_provider,
)
from app.services.report_persistence import (
delete_persisted_domain,
hydrate_report_store_from_db,
)
from app.services.report_store import ReportStore
logger = logging.getLogger(__name__)
@@ -146,9 +150,7 @@ def _get_domain_selectors_from_db(db: Session, domain_name: str) -> List[str]:
return []
def _get_domain_selectors_map_from_db(
db: Session, domain_names: List[str]
) -> Dict[str, List[str]]:
def _get_domain_selectors_map_from_db(db: Session, domain_names: List[str]) -> Dict[str, List[str]]:
"""Return manually configured DKIM selectors for all requested domains."""
if not domain_names:
return {}
@@ -157,11 +159,7 @@ def _get_domain_selectors_map_from_db(
selectors_by_domain: Dict[str, List[str]] = {}
for index in range(0, len(unique_names), DOMAIN_SELECTOR_LOOKUP_CHUNK_SIZE):
chunk = unique_names[index : index + DOMAIN_SELECTOR_LOOKUP_CHUNK_SIZE]
rows = (
db.query(Domain.name, Domain.dkim_selectors)
.filter(Domain.name.in_(chunk))
.all()
)
rows = db.query(Domain.name, Domain.dkim_selectors).filter(Domain.name.in_(chunk)).all()
for name, selectors in rows:
selectors_by_domain[name] = [
selector.strip() for selector in (selectors or "").split(",") if selector.strip()
@@ -180,6 +178,7 @@ async def get_domains_summary(db: Session = Depends(get_db)):
blocking the page load.
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
summaries = store.get_all_domain_summaries()
@@ -256,12 +255,13 @@ async def get_domains_summary(db: Session = Depends(get_db)):
@router.get("/domains", response_model=List[DomainResponse])
async def read_domains():
async def read_domains(db: Session = Depends(get_db)):
"""
Retrieve domains with their statistics.
For Milestone 1, this simply returns domains from the in-memory store.
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
summaries = store.get_all_domain_summaries()
@@ -281,11 +281,12 @@ async def read_domains():
@router.get("/domains/{domain_name}", response_model=DomainResponse)
async def read_domain(domain_name: str):
async def read_domain(domain_name: str, db: Session = Depends(get_db)):
"""
Get statistics for a specific domain.
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
if domain_name not in domains:
@@ -309,11 +310,15 @@ async def read_domain(domain_name: str):
@router.get("/{domain_id}/stats", response_model=DomainStatsResponse)
async def get_domain_stats(domain_id: str = Path(..., title="The domain ID or name")):
async def get_domain_stats(
domain_id: str = Path(..., title="The domain ID or name"),
db: Session = Depends(get_db),
):
"""
Get detailed statistics for a specific domain
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
# For Milestone 1, domain_id is simply the domain name
@@ -351,6 +356,7 @@ async def get_domain_dns_records(
selectors used as a final fallback.
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
if domain_id not in domains:
@@ -380,11 +386,13 @@ async def get_domain_dns_records(
async def get_domain_reports(
domain_id: str = Path(..., title="The domain ID or name"),
limit: int = Query(10, title="Maximum number of reports to return"),
db: Session = Depends(get_db),
):
"""
Get recent DMARC reports for a specific domain, along with compliance timeline
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
if domain_id not in domains:
@@ -493,12 +501,14 @@ async def _safe_ptr_lookup(provider: Any, ip: str, timeout: float = 3.0) -> Opti
async def get_domain_sources(
domain_id: str = Path(..., title="The domain ID or name"),
days: int = Query(30, title="Number of days to look back"),
db: Session = Depends(get_db),
):
"""
Get sending sources for a specific domain, including reverse-DNS hostnames
and SPF fix hints for sources that fail authentication.
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
if domain_id not in domains:
@@ -546,6 +556,7 @@ async def get_domain_selectors(
DMARC reports, read-only).
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
if domain_id not in store.get_domains():
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@@ -571,6 +582,7 @@ async def add_domain_selector(
any received DMARC report.
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
if domain_id not in store.get_domains():
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@@ -627,12 +639,16 @@ async def delete_domain_selector(
@router.delete("/{domain_id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_domain(domain_id: str = Path(..., title="The domain ID or name")):
async def delete_domain(
domain_id: str = Path(..., title="The domain ID or name"),
db: Session = Depends(get_db),
):
"""
Delete a domain and all associated data.
This performs a full cleanup of all reports and records related to this domain.
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
if domain_id not in domains:
@@ -642,7 +658,10 @@ async def delete_domain(domain_id: str = Path(..., title="The domain ID or name"
)
# Perform deletion with cleanup
deleted = store.delete_domain_with_cleanup(domain_id)
deleted_from_db = delete_persisted_domain(db, domain_id)
if deleted_from_db:
db.commit()
deleted = store.delete_domain_with_cleanup(domain_id) or deleted_from_db
if not deleted:
raise HTTPException(
@@ -660,6 +679,7 @@ async def search_domains(
policy: Optional[str] = Query(None, title="Filter by DMARC policy"),
page: int = Query(1, title="Page number", ge=1),
limit: int = Query(10, title="Number of domains per page", ge=1, le=100),
db: Session = Depends(get_db),
):
"""
Search domains with filtering and pagination.
@@ -672,6 +692,7 @@ async def search_domains(
limit: Number of domains per page (max 100)
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
summaries = store.get_all_domain_summaries()
+20 -2
View File
@@ -4,7 +4,9 @@ from typing import Any, Dict, Optional
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException
from pydantic import BaseModel
from sqlalchemy.orm import Session
from app.core.database import SessionLocal, get_db
from app.core.security import require_admin_auth
from app.services.imap_client import IMAPClient
@@ -22,6 +24,20 @@ class IMAPTestRequest(BaseModel):
ssl: bool = True
def _fetch_imap_reports_background(days: int, delete_emails: bool) -> None:
"""Fetch IMAP reports with a standalone DB session for background imports."""
db = SessionLocal()
try:
imap_client = IMAPClient(delete_emails=delete_emails, db=db)
imap_client.fetch_reports(days=days)
db.commit()
except Exception:
db.rollback()
raise
finally:
db.close()
@router.post("/test-connection")
async def test_imap_connection(
request: IMAPTestRequest,
@@ -57,6 +73,7 @@ async def test_imap_connection(
async def fetch_imap_reports(
background_tasks: BackgroundTasks,
_auth: dict = Depends(require_admin_auth),
db: Session = Depends(get_db),
days: int = 7,
delete_emails: bool = False,
) -> Dict[str, Any]:
@@ -69,11 +86,11 @@ async def fetch_imap_reports(
if days < 1 or days > 365:
raise HTTPException(status_code=400, detail="Days parameter must be between 1 and 365")
imap_client = IMAPClient(delete_emails=delete_emails)
imap_client = IMAPClient(delete_emails=delete_emails, db=db)
# Run in background if it might take a while
if days > 14:
background_tasks.add_task(imap_client.fetch_reports, days)
background_tasks.add_task(_fetch_imap_reports_background, days, delete_emails)
return {
"success": True,
"message": f"Background task started to fetch {days} days of reports",
@@ -83,6 +100,7 @@ async def fetch_imap_reports(
# Otherwise run immediately
try:
results = imap_client.fetch_reports(days=days)
db.commit()
return {
"success": results["success"],
@@ -679,6 +679,7 @@ async def gmail_fetch_reports(
access_token=source.gmail_access_token,
refresh_token=source.gmail_refresh_token or "",
already_ingested_ids=already,
db=db,
)
started_at = datetime.utcnow()
+34 -11
View File
@@ -1,10 +1,18 @@
import logging
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, File, HTTPException, UploadFile, status
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile, status
from pydantic import BaseModel
from sqlalchemy.orm import Session
from app.core.database import get_db
from app.services.dmarc_parser import DMARCParser
from app.services.report_persistence import (
delete_persisted_report,
hydrate_report_store_from_db,
report_exists,
save_parsed_report,
)
from app.services.report_store import ReportStore
from app.utils.domain_validator import DomainValidationError, validate_domain
@@ -157,7 +165,7 @@ class PaginatedReportResponse(BaseModel):
@router.post("/upload", response_model=UploadResponse)
async def upload_report(file: UploadFile = File(...)):
async def upload_report(file: UploadFile = File(...), db: Session = Depends(get_db)):
"""
Upload and process a DMARC aggregate report file (XML, ZIP, or GZIP)
@@ -195,7 +203,9 @@ async def upload_report(file: UploadFile = File(...)):
# Check for duplicate report before storing
store = ReportStore.get_instance()
report_id = report.get("report_id", "")
if report_id and store.has_report(domain, report_id):
if report_id and (
store.has_report(domain, report_id) or report_exists(db, domain, report_id)
):
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=(
@@ -205,6 +215,8 @@ async def upload_report(file: UploadFile = File(...)):
)
# Store the report
save_parsed_report(db, report)
db.commit()
store.add_report(report)
processed_records = report.get("summary", {}).get("total_count", 0)
@@ -231,11 +243,12 @@ async def upload_report(file: UploadFile = File(...)):
@router.get("", response_model=List[AllReportsItem])
async def get_all_reports():
async def get_all_reports(db: Session = Depends(get_db)):
"""
Get all DMARC reports across all domains, sorted by end_date descending.
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
domains = store.get_domains()
all_reports: List[AllReportsItem] = []
@@ -267,20 +280,22 @@ async def get_all_reports():
@router.get("/domains", response_model=List[str])
async def get_domains():
async def get_domains(db: Session = Depends(get_db)):
"""
Get list of all domains with reports
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
return store.get_domains()
@router.get("/domain/{domain}/summary", response_model=DomainSummary)
async def get_domain_summary(domain: str):
async def get_domain_summary(domain: str, db: Session = Depends(get_db)):
"""
Get summary statistics for a specific domain
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
summary = store.get_domain_summary(domain)
if not summary:
@@ -292,22 +307,24 @@ async def get_domain_summary(domain: str):
@router.get("/summary", response_model=List[DomainSummary])
async def get_all_summaries():
async def get_all_summaries(db: Session = Depends(get_db)):
"""
Get summary statistics for all domains
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
all_summaries = store.get_all_domain_summaries()
return [DomainSummary(domain=domain, **summary) for domain, summary in all_summaries.items()]
@router.get("/domain/{domain}/reports", response_model=List[ReportSummary])
async def get_domain_reports(domain: str):
async def get_domain_reports(domain: str, db: Session = Depends(get_db)):
"""
Get all reports for a specific domain
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
reports = store.get_domain_reports(domain)
if not reports:
@@ -336,6 +353,7 @@ async def get_domain_reports_paginated(
page_size: int = 10,
sort_by: str = "end_date",
sort_order: str = "desc",
db: Session = Depends(get_db),
):
"""
Get paginated reports for a specific domain with sorting options
@@ -348,6 +366,7 @@ async def get_domain_reports_paginated(
sort_order: Sort order (asc or desc)
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
all_reports = store.get_domain_reports(domain)
if not all_reports:
@@ -404,7 +423,7 @@ class DeleteReportResponse(BaseModel):
"/domain/{domain}/reports/{report_id}",
response_model=DeleteReportResponse,
)
async def delete_report(domain: str, report_id: str):
async def delete_report(domain: str, report_id: str, db: Session = Depends(get_db)):
"""
Delete a single DMARC report for a domain.
@@ -412,7 +431,10 @@ async def delete_report(domain: str, report_id: str):
that aggregated numbers remain accurate after deletion.
"""
store = ReportStore.get_instance()
deleted = store.delete_report(domain, report_id)
deleted_from_db = delete_persisted_report(db, domain, report_id)
if deleted_from_db:
db.commit()
deleted = store.delete_report(domain, report_id) or deleted_from_db
if not deleted:
raise HTTPException(
@@ -473,11 +495,12 @@ class ReportDetail(BaseModel):
@router.get("/{report_id}", response_model=ReportDetail)
async def get_report_by_id(report_id: str):
async def get_report_by_id(report_id: str, db: Session = Depends(get_db)):
"""
Get full details for a single DMARC report by its report ID.
"""
store = ReportStore.get_instance()
hydrate_report_store_from_db(db, store)
report = store.get_report_by_id(report_id)
if report is None: