feat: persist imported DMARC reports
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user