feat: add Cloudflare DNS integration

Closes #31
This commit is contained in:
Christian Krakau-Louis
2026-05-23 01:13:06 +02:00
parent a060bc59c2
commit 521dbc92d0
13 changed files with 1644 additions and 46 deletions
+427
View File
@@ -0,0 +1,427 @@
"""Cloudflare DNS discovery, analysis, and change tracking."""
from __future__ import annotations
import hashlib
import json
from dataclasses import dataclass
from datetime import UTC, datetime
from typing import Any, Dict, List, Optional
from sqlalchemy.orm import Session
from app.core.config import get_settings
from app.core.credential_encryption import decrypt_secret
from app.models.dns_cache import DNSRecordChange, DNSRecordSnapshot
from app.models.domain import Domain
from app.models.setting import Setting
from app.services.dns_resolver import CloudflareDNSProvider, extract_dmarc_policy
PROVIDER_NAME = "cloudflare"
@dataclass
class CloudflareCredentials:
"""Resolved Cloudflare credentials from persisted settings or environment."""
api_token: Optional[str] = None
zone_id: Optional[str] = None
@property
def configured(self) -> bool:
return bool(self.api_token)
def _utcnow_naive() -> datetime:
return datetime.now(UTC).replace(tzinfo=None)
def _plain_setting_value(db: Session, key: str) -> Optional[str]:
row = db.query(Setting).filter(Setting.key == key).first()
if row is None or not row.value:
return None
if key == "cloudflare.api_token":
return decrypt_secret(row.value)
return row.value
def get_cloudflare_credentials(db: Session) -> CloudflareCredentials:
"""Resolve Cloudflare credentials from app settings, falling back to env vars."""
settings = get_settings()
return CloudflareCredentials(
api_token=_plain_setting_value(db, "cloudflare.api_token") or settings.CLOUDFLARE_API_TOKEN,
zone_id=_plain_setting_value(db, "cloudflare.zone_id") or settings.CLOUDFLARE_ZONE_ID,
)
def build_cloudflare_provider(db: Session) -> CloudflareDNSProvider:
"""Return a Cloudflare provider configured from settings and environment."""
credentials = get_cloudflare_credentials(db)
if not credentials.configured:
raise LookupError("Cloudflare API token is not configured")
return CloudflareDNSProvider(
api_token=credentials.api_token,
zone_id=credentials.zone_id,
)
async def discover_cloudflare_zones(db: Session) -> List[Dict[str, Any]]:
"""Return zones visible to the configured Cloudflare token with import state."""
provider = build_cloudflare_provider(db)
known_domains = {name for (name,) in db.query(Domain.name).all()}
zones = await provider.list_zones()
return [
{
"id": zone.get("id"),
"name": zone.get("name"),
"status": zone.get("status"),
"account_name": (zone.get("account") or {}).get("name"),
"imported": zone.get("name") in known_domains,
}
for zone in zones
if zone.get("id") and zone.get("name")
]
async def import_cloudflare_domains(
db: Session,
*,
requested_domains: Optional[List[str]] = None,
) -> Dict[str, Any]:
"""Create Domain rows for Cloudflare zones, returning imported and existing names."""
zones = await discover_cloudflare_zones(db)
requested = {domain.strip().lower() for domain in requested_domains or [] if domain.strip()}
imported: List[str] = []
existing: List[str] = []
skipped: List[str] = []
for zone in zones:
name = str(zone["name"]).lower()
if requested and name not in requested:
skipped.append(name)
continue
domain = db.query(Domain).filter(Domain.name == name).first()
if domain is None:
db.add(Domain(name=name, active=True, verified=True))
imported.append(name)
else:
existing.append(name)
db.commit()
return {
"imported": imported,
"existing": existing,
"skipped": skipped,
"total_discovered": len(zones),
}
async def get_zone_for_domain(db: Session, domain: str) -> Dict[str, Any]:
"""Resolve the Cloudflare zone for a domain name."""
provider = build_cloudflare_provider(db)
credentials = get_cloudflare_credentials(db)
if credentials.zone_id:
records = await provider.list_dns_records(zone_id=credentials.zone_id)
try:
zone = await provider.find_zone_for_domain(domain)
except LookupError:
zone = None
return {
"id": credentials.zone_id,
"name": zone.get("name") if zone else domain,
"records": records,
}
zone = await provider.find_zone_for_domain(domain)
if not zone:
raise LookupError(f"No Cloudflare zone found for {domain}")
records = await provider.list_dns_records(zone_id=zone["id"])
return {"id": zone["id"], "name": zone["name"], "records": records}
def _record_key(record: Dict[str, Any]) -> str:
record_id = record.get("id")
if record_id:
return str(record_id)
payload = json.dumps(
{
"type": record.get("type"),
"name": record.get("name"),
"content": record.get("content"),
},
sort_keys=True,
separators=(",", ":"),
)
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
def _record_hash(record: Dict[str, Any]) -> str:
payload = json.dumps(
{
"type": record.get("type"),
"name": record.get("name"),
"content": record.get("content"),
"proxied": record.get("proxied"),
"ttl": record.get("ttl"),
},
sort_keys=True,
separators=(",", ":"),
)
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
def _change_to_dict(change: DNSRecordChange) -> Dict[str, Any]:
return {
"id": change.id,
"domain": change.domain,
"provider": change.provider,
"zone_id": change.zone_id,
"record_type": change.record_type,
"record_name": change.record_name,
"change_type": change.change_type,
"previous_content": change.previous_content,
"current_content": change.current_content,
"observed_at": change.observed_at.isoformat() if change.observed_at else None,
}
def sync_dns_record_changes(
db: Session,
*,
domain: str,
zone_id: str,
records: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
"""Track additions, modifications, and removals for a Cloudflare DNS snapshot."""
now = _utcnow_naive()
existing = {
snapshot.record_key: snapshot
for snapshot in db.query(DNSRecordSnapshot)
.filter(
DNSRecordSnapshot.domain == domain,
DNSRecordSnapshot.provider == PROVIDER_NAME,
DNSRecordSnapshot.zone_id == zone_id,
DNSRecordSnapshot.active == True, # noqa: E712
)
.all()
}
seen: set[str] = set()
changes: List[DNSRecordChange] = []
for record in records:
record_type = str(record.get("type") or "").upper()
record_name = str(record.get("name") or "")
if not record_type or not record_name:
continue
key = _record_key(record)
seen.add(key)
content = record.get("content")
current_hash = _record_hash(record)
snapshot = existing.get(key)
if snapshot is None:
snapshot = DNSRecordSnapshot(
domain=domain,
provider=PROVIDER_NAME,
zone_id=zone_id,
record_key=key,
record_id=record.get("id"),
record_type=record_type,
record_name=record_name,
content=content,
proxied=record.get("proxied"),
ttl=record.get("ttl"),
record_hash=current_hash,
active=True,
first_seen_at=now,
last_seen_at=now,
)
db.add(snapshot)
changes.append(
DNSRecordChange(
domain=domain,
provider=PROVIDER_NAME,
zone_id=zone_id,
record_key=key,
record_id=record.get("id"),
record_type=record_type,
record_name=record_name,
change_type="added",
current_content=content,
observed_at=now,
)
)
continue
if snapshot.record_hash != current_hash:
changes.append(
DNSRecordChange(
domain=domain,
provider=PROVIDER_NAME,
zone_id=zone_id,
record_key=key,
record_id=record.get("id"),
record_type=record_type,
record_name=record_name,
change_type="modified",
previous_content=snapshot.content,
current_content=content,
observed_at=now,
)
)
snapshot.content = content
snapshot.proxied = record.get("proxied")
snapshot.ttl = record.get("ttl")
snapshot.record_hash = current_hash
snapshot.record_id = record.get("id")
snapshot.record_type = record_type
snapshot.record_name = record_name
snapshot.active = True
snapshot.last_seen_at = now
for key, snapshot in existing.items():
if key in seen:
continue
snapshot.active = False
snapshot.last_seen_at = now
changes.append(
DNSRecordChange(
domain=domain,
provider=PROVIDER_NAME,
zone_id=zone_id,
record_key=key,
record_id=snapshot.record_id,
record_type=snapshot.record_type,
record_name=snapshot.record_name,
change_type="removed",
previous_content=snapshot.content,
observed_at=now,
)
)
for change in changes:
db.add(change)
db.commit()
for change in changes:
db.refresh(change)
return [_change_to_dict(change) for change in changes]
def list_dns_record_changes(db: Session, domain: str, *, limit: int = 50) -> List[Dict[str, Any]]:
"""Return recent DNS record change events for a domain."""
rows = (
db.query(DNSRecordChange)
.filter(DNSRecordChange.domain == domain)
.order_by(DNSRecordChange.observed_at.desc(), DNSRecordChange.id.desc())
.limit(max(1, min(limit, 200)))
.all()
)
return [_change_to_dict(row) for row in rows]
def _txt_contents(records: List[Dict[str, Any]], name: str) -> List[str]:
target = name.rstrip(".").lower()
return [
str(record.get("content") or "")
for record in records
if str(record.get("type") or "").upper() == "TXT"
and str(record.get("name") or "").rstrip(".").lower() == target
]
def _cloudflare_record_to_dict(record: Dict[str, Any]) -> Dict[str, Any]:
return {
"id": record.get("id"),
"type": record.get("type"),
"name": record.get("name"),
"content": record.get("content"),
"ttl": record.get("ttl"),
"proxied": record.get("proxied"),
"modified_on": record.get("modified_on"),
}
def analyze_dns_records(domain: str, records: List[Dict[str, Any]]) -> Dict[str, Any]:
"""Analyze Cloudflare DNS records and return checks plus actionable suggestions."""
root_txt = _txt_contents(records, domain)
dmarc_records = _txt_contents(records, f"_dmarc.{domain}")
spf_records = [record for record in root_txt if record.lower().startswith("v=spf1")]
dmarc_auth_records = [
record for record in dmarc_records if record.lower().startswith("v=dmarc1")
]
dkim_records = [
record
for record in records
if str(record.get("type") or "").upper() == "TXT"
and "._domainkey." in str(record.get("name") or "").lower()
and ("v=dkim1" in str(record.get("content") or "").lower())
]
suggestions: List[Dict[str, str]] = []
if not dmarc_auth_records:
suggestions.append(
{
"type": "missing_dmarc",
"severity": "error",
"message": "Add a TXT record at _dmarc with a v=DMARC1 policy.",
}
)
elif len(dmarc_auth_records) > 1:
suggestions.append(
{
"type": "duplicate_dmarc",
"severity": "error",
"message": "Keep exactly one DMARC TXT record at _dmarc.",
}
)
elif extract_dmarc_policy(dmarc_auth_records[0]) is None:
suggestions.append(
{
"type": "malformed_dmarc",
"severity": "error",
"message": "Add a p=none, p=quarantine, or p=reject tag to the DMARC record.",
}
)
if not spf_records:
suggestions.append(
{
"type": "missing_spf",
"severity": "warning",
"message": "Add an SPF TXT record at the root domain for authorized senders.",
}
)
elif len(spf_records) > 1:
suggestions.append(
{
"type": "duplicate_spf",
"severity": "error",
"message": "Merge multiple SPF records into a single v=spf1 TXT record.",
}
)
if not dkim_records:
suggestions.append(
{
"type": "missing_dkim",
"severity": "warning",
"message": "No DKIM TXT records were found; configure DKIM for active mail providers.",
}
)
return {
"records": [_cloudflare_record_to_dict(record) for record in records],
"checks": {
"dmarc": bool(dmarc_auth_records),
"dmarc_record": dmarc_auth_records[0] if dmarc_auth_records else None,
"dmarc_policy": (
extract_dmarc_policy(dmarc_auth_records[0]) if dmarc_auth_records else None
),
"spf": len(spf_records) == 1,
"spf_record": spf_records[0] if spf_records else None,
"dkim": bool(dkim_records),
"dkim_records": [
_cloudflare_record_to_dict(record)
for record in sorted(dkim_records, key=lambda item: str(item.get("name") or ""))
],
},
"suggestions": suggestions,
}
+151 -24
View File
@@ -11,7 +11,7 @@ import ipaddress
import logging
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
from typing import Any, Dict, List, Optional, Tuple
logger = logging.getLogger(__name__)
@@ -239,24 +239,16 @@ class SystemDNSProvider(BaseDNSProvider):
class CloudflareDNSProvider(BaseDNSProvider):
"""DNS provider using Cloudflare's DNS-over-HTTPS (DoH) endpoint.
"""DNS provider using Cloudflare DoH and, when configured, the REST API.
This provider resolves DNS queries via Cloudflare's public DoH API
(``1.1.1.1`` / ``cloudflare-dns.com``). When *api_token* and *zone_id*
are supplied, future versions will also support reading and writing DNS
records directly through the Cloudflare REST API, enabling automated DNS
synchronisation.
Current status
--------------
* DoH-based lookups are fully functional.
* Direct Cloudflare API integration (zone management, record sync) is
reserved for a future release.
Public DNS lookups continue to use Cloudflare's DNS-over-HTTPS endpoint.
If an API token is supplied, the provider can also discover account zones
and read managed DNS records directly from the Cloudflare REST API.
"""
#: Cloudflare DNS-over-HTTPS endpoint (JSON wire format)
CLOUDFLARE_DOH_URL: str = "https://cloudflare-dns.com/dns-query"
#: Cloudflare REST API base URL (for future zone-management support)
#: Cloudflare REST API base URL
CLOUDFLARE_API_BASE: str = "https://api.cloudflare.com/client/v4"
def __init__(
@@ -268,15 +260,121 @@ class CloudflareDNSProvider(BaseDNSProvider):
Parameters
----------
api_token:
Cloudflare API token. Required for future DNS record management;
not needed for read-only DoH lookups.
Cloudflare API token. Required for zone discovery and managed
DNS record reads; not needed for read-only DoH lookups.
zone_id:
Cloudflare zone identifier. Required for future DNS record
management.
Optional Cloudflare zone identifier used as a preferred zone.
"""
self.api_token = api_token
self.zone_id = zone_id
def _auth_headers(self) -> Dict[str, str]:
if not self.api_token:
raise LookupError("Cloudflare API token is not configured")
return {
"Authorization": f"Bearer {self.api_token}",
"Accept": "application/json",
}
async def _api_get(
self,
path: str,
*,
params: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""Call Cloudflare's REST API and return the decoded response."""
import httpx # type: ignore[import]
url = f"{self.CLOUDFLARE_API_BASE}{path}"
try:
async with httpx.AsyncClient() as client:
response = await client.get(
url,
params=params,
headers=self._auth_headers(),
timeout=DNS_TIMEOUT,
)
response.raise_for_status()
data = response.json()
except (httpx.RequestError, httpx.HTTPStatusError, httpx.TimeoutException) as exc:
raise LookupError(f"Cloudflare API request failed for {path}: {exc}") from exc
if not data.get("success", False):
errors = data.get("errors") or []
message = "; ".join(str(error.get("message", error)) for error in errors[:3])
raise LookupError(message or f"Cloudflare API request failed for {path}")
return data
async def list_zones(self) -> List[Dict[str, Any]]:
"""Return all zones visible to the configured Cloudflare API token."""
zones: List[Dict[str, Any]] = []
page = 1
while True:
data = await self._api_get(
"/zones",
params={"page": page, "per_page": 50, "status": "active"},
)
result = data.get("result") or []
if not isinstance(result, list):
return zones
zones.extend(result)
info = data.get("result_info") or {}
total_pages = int(info.get("total_pages") or 1)
if page >= total_pages:
return zones
page += 1
async def find_zone_for_domain(self, domain: str) -> Optional[Dict[str, Any]]:
"""Return the best matching Cloudflare zone for *domain*."""
zones = await self.list_zones()
domain_lc = domain.rstrip(".").lower()
matches = [
zone
for zone in zones
if isinstance(zone.get("name"), str)
and (
domain_lc == zone["name"].lower() or domain_lc.endswith(f".{zone['name'].lower()}")
)
]
if not matches:
return None
return sorted(matches, key=lambda zone: len(zone.get("name", "")), reverse=True)[0]
async def list_dns_records(
self,
*,
zone_id: Optional[str] = None,
name: Optional[str] = None,
record_type: Optional[str] = None,
) -> List[Dict[str, Any]]:
"""Return DNS records for a Cloudflare zone."""
resolved_zone_id = zone_id or self.zone_id
if not resolved_zone_id:
raise LookupError("Cloudflare zone ID is not configured")
records: List[Dict[str, Any]] = []
page = 1
while True:
params: Dict[str, Any] = {"page": page, "per_page": 100}
if name:
params["name"] = name
if record_type:
params["type"] = record_type
data = await self._api_get(
f"/zones/{resolved_zone_id}/dns_records",
params=params,
)
result = data.get("result") or []
if not isinstance(result, list):
return records
records.extend(result)
info = data.get("result_info") or {}
total_pages = int(info.get("total_pages") or 1)
if page >= total_pages:
return records
page += 1
async def lookup_txt(self, name: str) -> List[str]:
"""Resolve TXT records via Cloudflare's DoH endpoint (JSON format)."""
import httpx # type: ignore[import]
@@ -332,13 +430,42 @@ class CloudflareDNSProvider(BaseDNSProvider):
return None
def get_default_provider() -> BaseDNSProvider:
"""Return the default DNS provider (system resolver).
def _decrypt_setting_value(value: Optional[str]) -> Optional[str]:
if not value:
return value
try:
from app.core.credential_encryption import decrypt_secret
In a future release this function will inspect application settings and
return a ``CloudflareDNSProvider`` when Cloudflare credentials are
configured.
"""
return decrypt_secret(value)
except Exception:
return value
def _setting_value(db: Any, key: str) -> Optional[str]:
if db is None:
return None
try:
from app.models.setting import Setting
row = db.query(Setting).filter(Setting.key == key).first()
return row.value if row is not None else None
except Exception:
return None
def get_default_provider(db: Any = None) -> BaseDNSProvider:
"""Return the configured default DNS provider."""
resolver = (_setting_value(db, "dns.resolver") or "").strip().lower()
if resolver == "cloudflare":
from app.core.config import get_settings
settings = get_settings()
api_token = _decrypt_setting_value(_setting_value(db, "cloudflare.api_token"))
zone_id = _setting_value(db, "cloudflare.zone_id")
return CloudflareDNSProvider(
api_token=api_token or settings.CLOUDFLARE_API_TOKEN,
zone_id=zone_id or settings.CLOUDFLARE_ZONE_ID,
)
return SystemDNSProvider()