Compare commits
403 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 2ff47950fa | |||
| e9e0013ab8 | |||
| 8d8f3ace21 | |||
| 549d35e5a8 | |||
| 0d15c16e6a | |||
| e566a2d1e9 | |||
| 3b6b64afbe | |||
| 26bda88166 | |||
| 539c0b6283 | |||
| 5eea21c6f2 | |||
| cec72faf34 | |||
| 07069fd6f3 | |||
| d4f3e1155f | |||
| fec19fa6ed | |||
| d78a6db5a6 | |||
| d180f984a6 | |||
| 6ddd42bb3e | |||
| 26e9463f22 | |||
| b503d9c9fe | |||
| 60193c990b | |||
| 531a06c923 | |||
| 453a9b7e2a | |||
| 7f648e01a6 | |||
| 98116a4c6d | |||
| c1b53cf81f | |||
| 946925d370 | |||
| 63a3d300ff | |||
| d61a05bad0 | |||
| aedd868a9c | |||
| fcb6f57818 | |||
| e2b95f183d | |||
| ee663afffb | |||
| 312808a002 | |||
| ee3fff3936 | |||
| 2ff0c6d144 | |||
| 6883b40080 | |||
| bd0499f9ef | |||
| 6681aa4da0 | |||
| 805c36344f | |||
| f2021aece1 | |||
| 00ea392bee | |||
| 1f9f2b00be | |||
| 4a8d6a5067 | |||
| 57e164d282 | |||
| b30e28838b | |||
| f82dfaffe7 | |||
| c029574717 | |||
| 4a99a16160 | |||
| 39cd54eabf | |||
| 8556ea30bf | |||
| 78df802a98 | |||
| 5b4c8e8426 | |||
| 10b0467424 | |||
| 8450b52d43 | |||
| bfe442cecd | |||
| 5caefb06db | |||
| 08af80fa0e | |||
| f12461f5cf | |||
| 950c1c9a5a | |||
| 6818d52a8b | |||
| 9d0b12efea | |||
| d9e68e721a | |||
| 0200a13069 | |||
| cd6fa25f5d | |||
| 85c9435a84 | |||
| 04785ced18 | |||
| e781b32a01 | |||
| 271ce5d231 | |||
| 5445508a71 | |||
| db49e7e481 | |||
| 5060467c7e | |||
| 27628a8bce | |||
| 8bce8a8595 | |||
| 8407a63a92 | |||
| c2057acc71 | |||
| 8a4795187e | |||
| 817270b1fb | |||
| 908d4cd2cd | |||
| 1f954c9e08 | |||
| ca9e0bfd4a | |||
| 60fa437d02 | |||
| 77feb9d92e | |||
| 24f36a8252 | |||
| 65b154ae7f | |||
| 26c5c5b425 | |||
| e01095bca2 | |||
| a6c795ee16 | |||
| 0512f8cfa9 | |||
| da4489d36f | |||
| 127dc8e486 | |||
| f513f97947 | |||
| beae6e7469 | |||
| ef506ca654 | |||
| 521dbc92d0 | |||
| 0c6798e355 | |||
| e737e65e7a | |||
| a060bc59c2 | |||
| 236fa41df5 | |||
| 947428afb6 | |||
| 1a7a13b0d2 | |||
| e3a8ccdb45 | |||
| fbe97938e6 | |||
| 88a8325a9d | |||
| d5693c6d1d | |||
| e219048b5c | |||
| 5e6e408778 | |||
| 32129a41c9 | |||
| e70df29213 | |||
| b6db3bb014 | |||
| 8d6acc25ea | |||
| fe2c24de74 | |||
| 653cbf257b | |||
| 88af60ed2e | |||
| 302159aeb6 | |||
| 9f5cbdb851 | |||
| 55a5faa7e6 | |||
| 77fd0f9552 | |||
| a39b273419 | |||
| 20c245c9a2 | |||
| af66c93829 | |||
| c384200098 | |||
| d3f170c3a8 | |||
| f4ff1f151d | |||
| 2fc7c40189 | |||
| a20f00d204 | |||
| 05b687b525 | |||
| aacce812a5 | |||
| 9d0375569b | |||
| dd4b90c9de | |||
| 701a3603b5 | |||
| ef80619726 | |||
| 87ecdc1d9c | |||
| fd18aa9685 | |||
| b6852dcf0f | |||
| 7973620e42 | |||
| 51f9f303ca | |||
| 6c4bf72f31 | |||
| b5508e428c | |||
| be983e7110 | |||
| 016e63f202 | |||
| ed38cef7d0 | |||
| 2c528137db | |||
| f626675609 | |||
| 40debda0b2 | |||
| fbbfecf634 | |||
| b17c4c309b | |||
| bc0a64c108 | |||
| eeff24fa5c | |||
| 4d1ed329ec | |||
| 77de03d8b5 | |||
| b9bc0d3e4f | |||
| 6ece4e02f0 | |||
| 67e5fedce5 | |||
| 40b7bc459a | |||
| 3125fe3c3f | |||
| 364d517e07 | |||
| c11ba15611 | |||
| 8e8d90b8d0 | |||
| 34d8a655a2 | |||
| dad6f16f82 | |||
| a15b9fe5d3 | |||
| ccc634a10b | |||
| b1ff08590f | |||
| 672d426231 | |||
| 692ca913e0 | |||
| 81d302c867 | |||
| 884d83ef83 | |||
| 8485514445 | |||
| 1c8eaa5733 | |||
| c60a99f844 | |||
| 14f9f6defd | |||
| 2165fb1054 | |||
| d2323e5d1f | |||
| b5d0ecaf85 | |||
| b46087a1a1 | |||
| 9e7ba94338 | |||
| c1b81d3e06 | |||
| b21ebd3249 | |||
| 5f416d042f | |||
| 1b7f00e7a0 | |||
| cf21890064 | |||
| b104f26974 | |||
| 69aa1459eb | |||
| 648e3742f8 | |||
| 40f7d52df2 | |||
| b9041012de | |||
| 9630944460 | |||
| 1fdfa016ed | |||
| 4fc69602b3 | |||
| d72c27d881 | |||
| 07b2ea400e | |||
| 9e9df108c8 | |||
| bf76e2d450 | |||
| 00e1ee2a72 | |||
| 181b43bff9 | |||
| ee7c15ce4b | |||
| b4191e956c | |||
| 8f7c41193f | |||
| 20dddd0824 | |||
| 523b0529f4 | |||
| b0e643bdf3 | |||
| f8d4f936a7 | |||
| 32351a2134 | |||
| 2c9e5f4ca9 | |||
| c81173e417 | |||
| 9f1494a090 | |||
| 6805ae301a | |||
| c4a5d7da53 | |||
| 6581cf05b2 | |||
| 3a5b75e964 | |||
| 8d75f2e045 | |||
| 74ce9a7e37 | |||
| 92d6a06a33 | |||
| 8d98d96afc | |||
| 9e579972ca | |||
| 3fb90dbcf6 | |||
| 3c6d737139 | |||
| a463599d05 | |||
| 420975c1d7 | |||
| 9f5e2c1258 | |||
| bf7cb918bf | |||
| a808c62d4f | |||
| 274b0b09c5 | |||
| f31877b964 | |||
| 11440f4a02 | |||
| a23c69011a | |||
| 975f14c3f7 | |||
| 797567d441 | |||
| 3e35eab5ed | |||
| 63ce49ab66 | |||
| 04931172dd | |||
| 531dc968a8 | |||
| 308e6f8d91 | |||
| 6dd5c1e8f6 | |||
| 13a82ee148 | |||
| 9a331f18c5 | |||
| bd26f3c607 | |||
| 42e3a1f148 | |||
| 241713083b | |||
| ecf0976301 | |||
| 564824af6b | |||
| 24c134f1db | |||
| 09dbda26b1 | |||
| da34af977d | |||
| dc4e8377ee | |||
| 61cb428f54 | |||
| 3524344a4a | |||
| f998804766 | |||
| e75fde4311 | |||
| 7b932bb41d | |||
| 7955dac8e1 | |||
| 172d6465ab | |||
| 1cecde6120 | |||
| 227dfe1819 | |||
| 12305fbaef | |||
| 06a007d9db | |||
| 596a9b882e | |||
| 14ac9572b2 | |||
| bf9d1f6755 | |||
| 04e6dcfb8c | |||
| dfef73040c | |||
| 4f9e3b4b4b | |||
| 1e09d34736 | |||
| 4d77af791c | |||
| ba4dcd9f81 | |||
| 985f49d61f | |||
| de225bc1cd | |||
| aa8c55d27c | |||
| cb9eb39579 | |||
| ccb3f3fb7e | |||
| a61bf4df0a | |||
| dda4224129 | |||
| 12567f77c0 | |||
| f2dd2333b6 | |||
| f5c93a1e89 | |||
| 55fad8fcef | |||
| 58b5755f72 | |||
| 30007d21d6 | |||
| daa38b8337 | |||
| b6c15bda9c | |||
| 1815747fba | |||
| c144fc7a86 | |||
| 8665e39ff4 | |||
| 506346537c | |||
| 1455db9298 | |||
| af122881a5 | |||
| cd1f82aa79 | |||
| 56fe1db857 | |||
| 584a73d44d | |||
| 70c6a70af3 | |||
| a918e002be | |||
| 69c3bac480 | |||
| 2d364509bd | |||
| dbbf539182 | |||
| 962aec5bc3 | |||
| 204131ee2a | |||
| 11daa15335 | |||
| 8e051ac74d | |||
| 0c05ded74f | |||
| f3ad1fadaf | |||
| dadfd4ae3c | |||
| c48050ea04 | |||
| e033e11e34 | |||
| 27b74d26a6 | |||
| 7f3c908f71 | |||
| 5a753ef3cf | |||
| 551d604b8f | |||
| 874bc9110a | |||
| 6bae2f8146 | |||
| 25b3567414 | |||
| 4b2b904421 | |||
| f0c62ea95d | |||
| 2d32830270 | |||
| fb82643747 | |||
| 5edfac0397 | |||
| 6eeda1c5bd | |||
| 298ca938ec | |||
| 0bdd789944 | |||
| af457e4e48 | |||
| 2247115a7e | |||
| 087c2b375b | |||
| 8480b18b95 | |||
| a77a89f475 | |||
| 8d5b713bf1 | |||
| 23d5377a50 | |||
| 76e5e4d8f7 | |||
| 6294e4c935 | |||
| a04006f75b | |||
| 1c94844602 | |||
| 0d1a4fdac3 | |||
| 423c596c28 | |||
| 7c471b2800 | |||
| dd9aeb0aac | |||
| 15bbe5f137 | |||
| 4e4db14d36 | |||
| 53d543f49c | |||
| d2a557354f | |||
| e67375770f | |||
| c251a58340 | |||
| 24d7a7fd19 | |||
| a679459f5c | |||
| a0cbdb7b43 | |||
| 22946d619b | |||
| a64ed4acb3 | |||
| 0844981bc0 | |||
| 413bdab116 | |||
| 17116d8603 | |||
| f7200d973c | |||
| 0617cd6d2d | |||
| e9c1cd48bd | |||
| c631c5abbc | |||
| 9bd64c0429 | |||
| 4210450d2e | |||
| 33918ebc21 | |||
| af8620fecc | |||
| d91f424f3f | |||
| ae55ef46ea | |||
| a67900902b | |||
| 2a2fc9bfd5 | |||
| 84117abd13 | |||
| 078e33d273 | |||
| 4350bdcc16 | |||
| fca3eb3061 | |||
| 42d93e7f84 | |||
| 834de600df | |||
| 883d54a428 | |||
| 6d9392e5e7 | |||
| 7a0444a619 | |||
| 276cea7732 | |||
| 70dc77011a | |||
| c3ee0d8f5e | |||
| 1b83b2cce1 | |||
| 6329d3aa35 | |||
| e51fbe755e | |||
| f8887b50f4 | |||
| e09f39954f | |||
| 91cd221ef1 | |||
| 8e6ab244ad | |||
| de4d1b2b34 | |||
| 9588ce4e0d | |||
| e057a9c013 | |||
| dd73af5d37 | |||
| 9f99b6c6d7 | |||
| 4713e0bec7 | |||
| ea317bac32 | |||
| e471404538 | |||
| 2192944acf | |||
| 3ae290823e | |||
| c549e5dfbb | |||
| a90c9b393e | |||
| bad7a9ddc2 | |||
| 01d0331136 | |||
| f2a015971d | |||
| f7774dc3b2 | |||
| 50aa5bd5da | |||
| c07649a6bc | |||
| 286012fdd4 | |||
| 20affacd0f | |||
| 6eb99da749 | |||
| 9b21f5df80 | |||
| 714dd0a241 | |||
| 821d9dc863 | |||
| d580da6fee |
@@ -11,6 +11,12 @@ PROJECT_NAME="DMARQ"
|
||||
# NEVER use the default value in production!
|
||||
SECRET_KEY="CHANGE_THIS_TO_A_RANDOM_SECRET_IN_PRODUCTION"
|
||||
|
||||
# Admin API Key for the X-API-Key header (optional)
|
||||
# If set, this key is used instead of generating a random one at startup.
|
||||
# Generate with: openssl rand -hex 32
|
||||
# If not set, a random key is generated each restart (key length logged only).
|
||||
# ADMIN_API_KEY="your_admin_api_key_here"
|
||||
|
||||
# Environment (development/production)
|
||||
# Affects HSTS and other security settings
|
||||
ENVIRONMENT="development"
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
[flake8]
|
||||
# Keep in sync with black's line-length in pyproject.toml [tool.black]
|
||||
max-line-length = 100
|
||||
max-complexity = 10
|
||||
exclude =
|
||||
.git,
|
||||
__pycache__,
|
||||
.venv,
|
||||
venv,
|
||||
build,
|
||||
dist,
|
||||
*.egg-info,
|
||||
migrations
|
||||
# Ignored rules – must not conflict with black:
|
||||
# E203 – whitespace before ':' (black formats slices this way)
|
||||
# W503 – line break before binary operator (black prefers this style)
|
||||
# E501 – line too long (black already enforces max-line-length; avoid double-reporting)
|
||||
extend-ignore = E203, W503, E501
|
||||
per-file-ignores =
|
||||
# Allow unused imports in __init__.py (re-exports)
|
||||
__init__.py: F401
|
||||
+155
-16
@@ -3,12 +3,31 @@ name: CI
|
||||
on:
|
||||
push:
|
||||
branches: [main, develop]
|
||||
tags: ['v*']
|
||||
pull_request:
|
||||
branches: [main, develop]
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
release_ref:
|
||||
description: Git ref to build, for example v1.8.4
|
||||
required: false
|
||||
type: string
|
||||
release_tag:
|
||||
description: Release image tag to publish, for example v1.8.4
|
||||
required: false
|
||||
type: string
|
||||
promote_stable:
|
||||
description: Also publish the stable image tag
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
schedule:
|
||||
# Weekly security scan on Mondays at 00:00 UTC
|
||||
- cron: '0 0 * * 1'
|
||||
|
||||
env:
|
||||
K8S_STATE_REPO: christianlouis/k8s-cluster-state
|
||||
|
||||
jobs:
|
||||
# ── Stage 1: Lint (gates everything else) ────────────────────────────────
|
||||
lint:
|
||||
@@ -21,10 +40,10 @@ jobs:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python 3.10
|
||||
- name: Set up Python 3.13
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
python-version: '3.13'
|
||||
|
||||
- name: Cache pip packages
|
||||
uses: actions/cache@v4
|
||||
@@ -39,6 +58,11 @@ jobs:
|
||||
pip install black isort flake8 pylint
|
||||
cd backend && pip install -r requirements.txt
|
||||
|
||||
- name: Auto-format with Black and isort
|
||||
run: |
|
||||
black backend/app
|
||||
isort backend/app
|
||||
|
||||
- name: Black – format check
|
||||
run: black --check backend/app
|
||||
|
||||
@@ -64,10 +88,10 @@ jobs:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python 3.10
|
||||
- name: Set up Python 3.13
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
python-version: '3.13'
|
||||
|
||||
- name: Cache pip packages
|
||||
uses: actions/cache@v4
|
||||
@@ -84,7 +108,7 @@ jobs:
|
||||
- name: Run tests with coverage
|
||||
run: |
|
||||
cd backend
|
||||
pytest --cov=app --cov-report=xml --cov-report=term-missing
|
||||
pytest --cov=app --cov-report=xml --cov-report=term-missing --cov-fail-under=80
|
||||
|
||||
- name: Upload coverage to Codecov
|
||||
uses: codecov/codecov-action@v4
|
||||
@@ -105,10 +129,10 @@ jobs:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python 3.10
|
||||
- name: Set up Python 3.13
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
python-version: '3.13'
|
||||
|
||||
- name: Cache pip packages
|
||||
uses: actions/cache@v4
|
||||
@@ -189,7 +213,10 @@ jobs:
|
||||
name: Docker Build & Publish
|
||||
runs-on: ubuntu-latest
|
||||
needs: [test, security]
|
||||
if: github.event_name == 'push' && github.ref == 'refs/heads/main'
|
||||
if: >-
|
||||
(github.event_name == 'push' && github.ref == 'refs/heads/main') ||
|
||||
(github.event_name == 'push' && startsWith(github.ref, 'refs/tags/v')) ||
|
||||
github.event_name == 'workflow_dispatch'
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
@@ -197,6 +224,8 @@ jobs:
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ github.event.inputs.release_ref || github.ref }}
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
@@ -208,15 +237,53 @@ jobs:
|
||||
username: ${{ github.actor }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Extract metadata for Docker
|
||||
- name: Compute Docker metadata
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ghcr.io/${{ github.repository }}
|
||||
tags: |
|
||||
type=ref,event=branch
|
||||
type=sha,prefix=
|
||||
type=raw,value=latest,enable={{is_default_branch}}
|
||||
env:
|
||||
DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
|
||||
EVENT_NAME: ${{ github.event_name }}
|
||||
REF_NAME: ${{ github.ref_name }}
|
||||
REF_TYPE: ${{ github.ref_type }}
|
||||
RELEASE_TAG: ${{ github.event.inputs.release_tag }}
|
||||
PROMOTE_STABLE: ${{ github.event.inputs.promote_stable }}
|
||||
run: |
|
||||
IMAGE="ghcr.io/$(echo "${GITHUB_REPOSITORY}" | tr '[:upper:]' '[:lower:]')"
|
||||
SHORT_SHA="$(echo "${GITHUB_SHA}" | cut -c1-7)"
|
||||
TAGS=("${IMAGE}:${SHORT_SHA}")
|
||||
|
||||
if [ "${REF_TYPE}" = "branch" ]; then
|
||||
SAFE_BRANCH="$(echo "${REF_NAME}" | tr '/:@' '---')"
|
||||
TAGS+=("${IMAGE}:${SAFE_BRANCH}")
|
||||
if [ "${REF_NAME}" = "${DEFAULT_BRANCH}" ]; then
|
||||
TAGS+=("${IMAGE}:latest")
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "${REF_TYPE}" = "tag" ]; then
|
||||
TAGS+=("${IMAGE}:${REF_NAME}")
|
||||
if [[ "${REF_NAME}" =~ ^v[0-9] ]]; then
|
||||
TAGS+=("${IMAGE}:${REF_NAME#v}")
|
||||
TAGS+=("${IMAGE}:stable")
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "${EVENT_NAME}" = "workflow_dispatch" ] && [ -n "${RELEASE_TAG}" ]; then
|
||||
TAGS+=("${IMAGE}:${RELEASE_TAG}")
|
||||
fi
|
||||
|
||||
if [ "${PROMOTE_STABLE}" = "true" ]; then
|
||||
TAGS+=("${IMAGE}:stable")
|
||||
fi
|
||||
|
||||
{
|
||||
echo "tags<<EOF"
|
||||
printf '%s\n' "${TAGS[@]}" | sort -u
|
||||
echo "EOF"
|
||||
echo "labels<<EOF"
|
||||
echo "org.opencontainers.image.source=https://github.com/${GITHUB_REPOSITORY}"
|
||||
echo "org.opencontainers.image.revision=${GITHUB_SHA}"
|
||||
echo "EOF"
|
||||
} >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Build and push Docker image
|
||||
uses: docker/build-push-action@v6
|
||||
@@ -228,3 +295,75 @@ jobs:
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
|
||||
# ── Stage 4: GitOps – update preprod k8s manifest ────────────────────────
|
||||
update-k8s-manifest:
|
||||
name: Update Preprod K8s Manifest
|
||||
runs-on: ubuntu-latest
|
||||
needs: [docker]
|
||||
if: github.event_name == 'push' && github.ref == 'refs/heads/main'
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
steps:
|
||||
- name: Compute image tag
|
||||
id: tag
|
||||
run: |
|
||||
SHORT_SHA=$(echo "${{ github.sha }}" | cut -c1-7)
|
||||
echo "image=ghcr.io/${{ github.repository }}:${SHORT_SHA}" >> "$GITHUB_OUTPUT"
|
||||
echo "short_sha=${SHORT_SHA}" >> "$GITHUB_OUTPUT"
|
||||
echo "image_pattern=^ghcr\\.io/${{ github.repository }}:" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Check if GH_PAT is configured and has repo access
|
||||
id: pat-check
|
||||
env:
|
||||
GH_PAT: ${{ secrets.GH_PAT }}
|
||||
run: |
|
||||
if [ -z "$GH_PAT" ]; then
|
||||
echo "::warning::GH_PAT secret is not configured. Skipping k8s manifest update."
|
||||
echo "available=false" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
HTTP_CODE=$(curl -s -o /dev/null -w "%{http_code}" \
|
||||
-H "Authorization: Bearer $GH_PAT" \
|
||||
"https://api.github.com/repos/${{ env.K8S_STATE_REPO }}")
|
||||
if [ "$HTTP_CODE" = "200" ]; then
|
||||
echo "available=true" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
echo "::warning::GH_PAT does not have access to ${{ env.K8S_STATE_REPO }} (HTTP $HTTP_CODE). Skipping k8s manifest update."
|
||||
echo "available=false" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
fi
|
||||
|
||||
- name: Checkout k8s-cluster-state
|
||||
if: steps.pat-check.outputs.available == 'true'
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
repository: ${{ env.K8S_STATE_REPO }}
|
||||
token: ${{ secrets.GH_PAT }}
|
||||
path: k8s-cluster-state
|
||||
ref: main
|
||||
|
||||
- name: Update image tag in preprod manifest
|
||||
if: steps.pat-check.outputs.available == 'true'
|
||||
uses: mikefarah/yq@v4.44.6
|
||||
env:
|
||||
IMAGE: ${{ steps.tag.outputs.image }}
|
||||
IMAGE_PATTERN: ${{ steps.tag.outputs.image_pattern }}
|
||||
with:
|
||||
cmd: |
|
||||
yq -i '(.. | select(tag == "!!str") | select(test(strenv(IMAGE_PATTERN)))) = strenv(IMAGE)' \
|
||||
k8s-cluster-state/apps/dmarq/preprod/dmarq-stack.yaml
|
||||
|
||||
- name: Commit and push manifest update
|
||||
if: steps.pat-check.outputs.available == 'true'
|
||||
run: |
|
||||
cd k8s-cluster-state
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
git add apps/dmarq/preprod/dmarq-stack.yaml
|
||||
if git diff --staged --quiet; then
|
||||
echo "No changes to commit"
|
||||
else
|
||||
git commit -m "chore(preprod): update dmarq image to ${{ steps.tag.outputs.short_sha }}"
|
||||
git push
|
||||
fi
|
||||
|
||||
@@ -22,7 +22,7 @@ jobs:
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
python-version: "3.13"
|
||||
|
||||
- name: Semantic Release
|
||||
uses: python-semantic-release/python-semantic-release@v10
|
||||
|
||||
+2
-1
@@ -130,6 +130,7 @@ celerybeat.pid
|
||||
# Environments
|
||||
.env
|
||||
.venv
|
||||
.pipcache/
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
@@ -216,4 +217,4 @@ temp/
|
||||
*.backup
|
||||
|
||||
# Docker override
|
||||
docker-compose.override.yml
|
||||
docker-compose.override.yml
|
||||
|
||||
+1
-1
@@ -3,7 +3,7 @@ version: 2
|
||||
build:
|
||||
os: ubuntu-22.04
|
||||
tools:
|
||||
python: "3.9"
|
||||
python: "3.13"
|
||||
|
||||
mkdocs:
|
||||
configuration: mkdocs.yml
|
||||
|
||||
+1
-1
@@ -61,7 +61,7 @@ Unsure where to begin? You can start by looking through issues tagged with:
|
||||
|
||||
### Prerequisites
|
||||
|
||||
- Python 3.10 or higher
|
||||
- Python 3.13 or higher
|
||||
- Docker and Docker Compose (for full stack testing)
|
||||
- Git
|
||||
|
||||
|
||||
@@ -4,34 +4,38 @@
|
||||
|
||||
🌐 [Live Demo (soon)](https://app.dmarq.org)
|
||||
🔒 Self-hosted. Secure. Beautifully visual.
|
||||
🛠️ Docker-deployable. Cloudflare-integrated.
|
||||
📬 Aggregate & forensic report support.
|
||||
🛠️ Docker-deployable. DNS posture checks with optional Cloudflare inspection.
|
||||
📬 Aggregate report support (failure/forensic reports planned).
|
||||
|
||||
---
|
||||
|
||||
## 💡 What is DMARQ?
|
||||
|
||||
DMARQ ingests and visualizes DMARC (Domain-based Message Authentication, Reporting & Conformance) reports — both aggregate and forensic — to help domain owners understand who is sending emails on their behalf and whether those messages are properly authenticated using SPF and DKIM.
|
||||
DMARQ ingests and visualizes DMARC (Domain-based Message Authentication, Reporting & Conformance) reports — primarily aggregate (RUA) reports today — to help domain owners understand who is sending emails on their behalf and whether those messages are properly authenticated using SPF and DKIM.
|
||||
|
||||
No more guessing. See which services are passing DMARC, which are failing, and how to fix them — all in one clear dashboard.
|
||||
|
||||
---
|
||||
|
||||
## 🚀 Current Status - Milestone 1 Completed
|
||||
## 🚀 Current Status (Milestones 1–8 Complete)
|
||||
|
||||
We have achieved **Milestone 1: Basic DMARC Monitoring**. This milestone includes:
|
||||
DMARQ currently supports end-to-end aggregate DMARC monitoring with mailbox ingestion, persistence, reporting, DNS checks, and notifications.
|
||||
|
||||
- ✅ DMARC XML report parsing (supports XML, ZIP, and GZIP formats)
|
||||
- ✅ In-memory storage of report data for up to 5 domains
|
||||
- ✅ Simple dashboard UI showing DMARC compliance statistics
|
||||
- ✅ Support for uploading and processing DMARC aggregate reports
|
||||
- ✅ Domain overview with compliance rates and email statistics
|
||||
Included:
|
||||
|
||||
You can now:
|
||||
1. Upload DMARC aggregate reports via the web interface
|
||||
2. View summary statistics across all monitored domains
|
||||
3. Drill down into domain-specific details and reports
|
||||
4. Track compliance rates and authentication failures
|
||||
- ✅ DMARC aggregate XML report parsing (XML, ZIP, GZIP)
|
||||
- ✅ Upload ingestion + mailbox ingestion (IMAP + Gmail OAuth)
|
||||
- ✅ Database persistence (SQLite/PostgreSQL) + migrations
|
||||
- ✅ Dashboard trends, domain timelines, and sender/source analytics
|
||||
- ✅ Import history + backfills for mail sources
|
||||
- ✅ Alerts & notifications via Apprise (test send, alert rules, daily/weekly summaries)
|
||||
- ✅ DNS checks (DMARC/SPF/DKIM) with DKIM selector discovery from report data
|
||||
- ✅ Cloudflare read-only domain discovery, DNS inspection, recommendations, and change tracking
|
||||
|
||||
Up next:
|
||||
|
||||
- 🔜 Setup and operations polish (Milestone 9)
|
||||
- 🧊 Failure/forensic report support (RUF) (Milestone 10)
|
||||
|
||||
---
|
||||
|
||||
@@ -39,24 +43,24 @@ You can now:
|
||||
|
||||
### 📊 Dashboard & Reports
|
||||
- **DMARC Compliance Rate**: Track pass/fail rates over time
|
||||
- **Enforcement Rate**: Visualize policy strength and adoption
|
||||
- **Volume & Trends**: Identify traffic spikes and anomalies
|
||||
- **Top Sending Sources**: Detect unknown or unauthorized senders
|
||||
- **Forensic Reports**: Analyze failure samples (RFC 6591 support)
|
||||
- **Actionable Recommendations**: Prioritize what to fix next
|
||||
- **Export**: CSV export for selected domains and date ranges
|
||||
|
||||
### 🛡 DNS Record Health
|
||||
- Inspect **SPF**, **DKIM**, **DMARC**, **MX**, and **BIMI** records
|
||||
- Inspect **SPF**, **DKIM**, and **DMARC** records
|
||||
- Discover likely **DKIM selectors** from report data
|
||||
- Show which records are missing, broken, or invalid
|
||||
- Get **fix suggestions** tailored to your provider (e.g., Google, Microsoft)
|
||||
- 🔒 No automatic changes — all DNS updates require explicit confirmation
|
||||
- 🔒 No automatic changes — all DNS updates require explicit confirmation (when remediation workflows are added)
|
||||
|
||||
### 🌐 Cloudflare Integration
|
||||
- Automatically discover domains in your Cloudflare account
|
||||
- Fetch and analyze relevant DNS records
|
||||
- Suggest missing or malformed entries
|
||||
- Track configuration changes over time (coming soon)
|
||||
- Optional read-only domain discovery and DNS inspection
|
||||
- Import Cloudflare zones as monitored domains from Settings
|
||||
- Suggestions for missing or malformed entries
|
||||
- Track configuration changes over time
|
||||
|
||||
### ⚙️ Web-Based Setup Wizard
|
||||
### ⚙️ Web-Based Setup Wizard (Planned / In Progress)
|
||||
- Guided onboarding experience (no CLI setup required)
|
||||
- Store all configuration in a secure internal database
|
||||
- Seed config with environment variables for headless deployment
|
||||
@@ -64,12 +68,13 @@ You can now:
|
||||
### 🚨 Alerts & Notifications
|
||||
- Integration with [Apprise](https://github.com/caronc/apprise)
|
||||
- Email, Slack, webhook, and more
|
||||
- Alert on new failures, compliance drops, or unknown senders
|
||||
- Alert on new senders, compliance drops, elevated failures, or missing reports
|
||||
- Daily and weekly DMARC summaries
|
||||
- Alert history for active and resolved alerts
|
||||
|
||||
### 🔐 User Management
|
||||
- Built-in authentication via **FastAPI Users**
|
||||
- JWT-secured API endpoints
|
||||
- Admin dashboard access control
|
||||
### 🔐 Authentication
|
||||
- Logto-based authentication integration
|
||||
- Explicit auth-disabled mode for local development
|
||||
|
||||
---
|
||||
|
||||
@@ -100,24 +105,55 @@ Then visit [http://localhost:8080](http://localhost:8080)
|
||||
|
||||
---
|
||||
|
||||
## 🚨 Integration with Apprise
|
||||
|
||||
DMARQ sends notifications through Apprise target URLs configured in
|
||||
**Settings** > **Notifications**. Add one target URL per line, enable
|
||||
notifications, and use **Send Test** to verify delivery. Target URLs are
|
||||
encrypted in the database and redacted in API responses. Apprise supports email,
|
||||
Slack, Teams, Discord, generic webhooks, and many other targets.
|
||||
|
||||
Notification settings include alert-rule toggles and thresholds for:
|
||||
|
||||
- New sending sources
|
||||
- Compliance-rate drops
|
||||
- DMARC failures above a daily threshold
|
||||
- Missing reports for monitored domains
|
||||
|
||||
DMARQ can also send daily and weekly summaries. Use **Preview Summary** to see
|
||||
the current summary payload, **Send Summary Now** for an immediate message, and
|
||||
the daily/weekly toggles to enable scheduled delivery. Outbound messages are
|
||||
rate-limited by the configured cooldown and email addresses are redacted by
|
||||
default before delivery.
|
||||
|
||||
Alert history is available in **Settings** > **Notifications** after alerts have
|
||||
been evaluated or sent. History rows track active/resolved status, first seen,
|
||||
last seen, observed count, and alert metadata. Notification and alert-rule
|
||||
configuration changes are recorded in the configuration audit trail without
|
||||
storing raw notification secrets.
|
||||
|
||||
See [Settings](docs/user_guide/settings.md) and
|
||||
[Configuration](docs/deployment/configuration.md) for examples and available
|
||||
settings.
|
||||
|
||||
---
|
||||
|
||||
## 📦 Requirements
|
||||
|
||||
- DMARC aggregate reports (XML, ZIP, or GZIP format)
|
||||
- Docker + Docker Compose (for production deployment)
|
||||
- Python 3.10+ (for development)
|
||||
- Python 3.13+ (for development)
|
||||
|
||||
---
|
||||
|
||||
## 🧪 Development Roadmap
|
||||
## 🧭 Development Roadmap
|
||||
|
||||
- ✅ **Milestone 1**: Basic DMARC Monitoring (up to 5 domains)
|
||||
- ✅ **Milestone 2**: IMAP Integration
|
||||
- ✅ **Milestone 3**: Database Persistence
|
||||
- 🔜 **Milestone 4**: Enhanced Dashboard & Visualization
|
||||
- 🔜 **Milestone 5**: User Authentication & Multi-User Support
|
||||
- ✅ **Milestones 1–8**: Parsing, ingestion (upload/IMAP/Gmail), persistence, reporting, notifications, production hardening, DNS health, Cloudflare read-only inspection
|
||||
- 🔜 **Milestone 9**: Setup and operations polish
|
||||
- 🧊 **Milestone 10**: Failure/forensic report support (RUF)
|
||||
- 🧠 **Milestones 11–16**: DMARC format compatibility, Microsoft 365 ingestion, broader email posture, APIs/webhooks, workspaces/MSP, AI/MCP (see docs)
|
||||
|
||||
See the full [Roadmap](docs/development/roadmap.md) and [TODO](TODO.md) for details
|
||||
on what is planned vs. what is currently implemented.
|
||||
See the full [Roadmap](docs/development/roadmap.md) and [Milestones](docs/milestones.md).
|
||||
|
||||
---
|
||||
|
||||
@@ -142,7 +178,7 @@ automated versioning and changelog generation.
|
||||
Unlike most commercial DMARC tools, DMARQ gives you:
|
||||
- 🔍 Full visibility without third-party access to your reports
|
||||
- 🧠 Intelligence-driven suggestions, not just raw data
|
||||
- 🎨 A beautiful, intuitive dashboard with real-time insights
|
||||
- 🎨 A beautiful, intuitive dashboard with actionable insights
|
||||
- 💻 Self-hosted flexibility with modern developer practices
|
||||
|
||||
Let's build better email security — together.
|
||||
|
||||
@@ -30,6 +30,10 @@ These features are documented and confirmed working in the codebase:
|
||||
(`docker-compose.yml`, `backend/Dockerfile`)
|
||||
- [x] **Setup Wizard** — Basic guided onboarding endpoints, though in-memory only
|
||||
(`backend/app/api/api_v1/endpoints/setup.py`)
|
||||
- [x] **DNS Record Health Checks** — Real DNS lookups for DMARC, SPF, DKIM, and
|
||||
reverse-DNS (PTR) via `SystemDNSProvider` (dnspython async) and
|
||||
`CloudflareDNSProvider` (DNS-over-HTTPS); used by `/dns`, `/summary`, and
|
||||
`/sources` API endpoints (`backend/app/services/dns_resolver.py`)
|
||||
|
||||
---
|
||||
|
||||
@@ -59,24 +63,14 @@ have no working implementation in the codebase yet.
|
||||
|
||||
### Forensic Reports (RFC 6591)
|
||||
- **Documented in**: README.md ("Forensic Reports: Analyze failure samples (RFC 6591 support)")
|
||||
- **Current state**: The DMARC parser (`backend/app/services/dmarc_parser.py`) only
|
||||
handles aggregate reports. There is no forensic report parsing, UI, or storage.
|
||||
- [ ] Forensic report parsing
|
||||
- [ ] Failure sample analysis
|
||||
- [ ] PII redaction options
|
||||
- [ ] Detailed authentication failure views
|
||||
|
||||
### DNS Record Health Checks
|
||||
- **Documented in**: README.md ("Inspect SPF, DKIM, DMARC, MX, and BIMI records"),
|
||||
docs/development/roadmap.md (Milestone 8)
|
||||
- **Current state**: The `/api/v1/domains/{domain_id}/dns` endpoint
|
||||
(`backend/app/api/api_v1/endpoints/domains.py`) returns hardcoded mock data.
|
||||
`dnspython>=2.3.0` is in `requirements.txt` but is never imported or used.
|
||||
- [ ] Real DNS lookups for SPF, DKIM, DMARC, and MX records
|
||||
- [ ] BIMI record support (zero code exists)
|
||||
- [ ] Identify missing, broken, or invalid records
|
||||
- [ ] Provider-specific fix suggestions (Google, Microsoft, etc.)
|
||||
- [ ] DNSSEC validation
|
||||
- **Current state**: Aggregate and forensic reports are now parsed separately. Forensic
|
||||
reports are stored in dedicated database rows and surfaced through authenticated APIs
|
||||
without affecting aggregate compliance statistics. Operators can configure forensic
|
||||
email-address and token redaction under Settings.
|
||||
- [x] Forensic report parsing
|
||||
- [x] Failure sample analysis
|
||||
- [x] PII redaction options
|
||||
- [x] Detailed authentication failure views
|
||||
|
||||
### User Authentication & Multi-User Support
|
||||
- **Documented in**: README.md ("Built-in authentication via FastAPI Users"),
|
||||
@@ -94,15 +88,15 @@ have no working implementation in the codebase yet.
|
||||
|
||||
### Dashboard Visualizations (Real Data)
|
||||
- **Documented in**: README.md ("Track pass/fail rates over time", "Volume & Trends")
|
||||
- **Current state**: The stats endpoints (`backend/app/utils/stats_summarizer.py`,
|
||||
`backend/app/api/api_v1/endpoints/domains.py`) return mock/random data with TODO
|
||||
comments like `# For now, mock statistics` and `# TODO: Replace with actual
|
||||
historical data`. Chart.js is integrated in templates but fed with mock data.
|
||||
- [ ] Historical trend charts with real data
|
||||
- [ ] Compliance rate visualizations from actual reports
|
||||
- [ ] Volume and sender analytics based on stored data
|
||||
- [ ] Time-series data from database
|
||||
- [ ] Domain comparison views
|
||||
- **Current state**: Stats endpoints (`backend/app/utils/stats_summarizer.py`,
|
||||
`backend/app/api/api_v1/endpoints/domains.py`) now query real data from the
|
||||
database and in-memory ReportStore. Chart.js visualizations display actual
|
||||
compliance trends derived from uploaded DMARC reports.
|
||||
- [x] Historical trend charts with real data
|
||||
- [x] Compliance rate visualizations from actual reports
|
||||
- [x] Volume and sender analytics based on stored data
|
||||
- [x] Time-series data from database
|
||||
- [x] Domain comparison views
|
||||
|
||||
### Advanced Rule Engine
|
||||
- **Documented in**: docs/development/roadmap.md (Milestone 7)
|
||||
@@ -151,14 +145,34 @@ have no working implementation in the codebase yet.
|
||||
- [ ] Add vault integration for secure credential storage
|
||||
- [ ] Audit logging for IMAP operations
|
||||
|
||||
### DNS Record Health Checks
|
||||
- **Documented in**: README.md ("Inspect SPF, DKIM, DMARC, MX, and BIMI records"),
|
||||
docs/development/roadmap.md (Milestone 8)
|
||||
- **Current state**: The `/api/v1/domains/{domain_id}/dns` and `/api/v1/domains/summary`
|
||||
endpoints perform real DNS lookups via `SystemDNSProvider` (dnspython async) and
|
||||
`CloudflareDNSProvider` (DNS-over-HTTPS), implemented in
|
||||
`backend/app/services/dns_resolver.py`. DKIM selectors from DMARC reports and manually
|
||||
configured selectors are both checked. PTR reverse-DNS lookups power the `/sources`
|
||||
endpoint. Full Cloudflare zone-management integration is planned for a future milestone.
|
||||
- [x] Real DNS lookups for SPF, DKIM, and DMARC records
|
||||
- [x] Proper error handling for DNS failures and timeouts (per-domain timeouts, graceful fallback)
|
||||
- [x] Unit and integration tests for DNS resolution (55 tests passing)
|
||||
- [ ] MX record lookups
|
||||
- [ ] BIMI record support
|
||||
- [ ] Cloudflare API zone-management integration (read/write DNS records)
|
||||
- [ ] Provider-specific fix suggestions (Google, Microsoft, etc.)
|
||||
- [ ] DNSSEC validation
|
||||
|
||||
---
|
||||
|
||||
## Housekeeping
|
||||
|
||||
- [ ] Remove unused `apprise` from `requirements.txt` or implement alerts
|
||||
- [ ] Remove unused `dnspython` from `requirements.txt` or implement DNS checks
|
||||
- [x] Remove unused `dnspython` from `requirements.txt` or implement DNS checks (done — `dnspython` is used by `SystemDNSProvider`)
|
||||
- [ ] Remove or wire up `fastapi-users` (currently installed but unused)
|
||||
- [ ] Replace mock data in stats endpoints with real database queries
|
||||
- [ ] Replace mock DNS data with actual DNS lookups
|
||||
- [ ] Add CI/CD pipeline
|
||||
- [x] Replace mock data in stats endpoints with real database queries
|
||||
- [x] Replace mock DNS data with actual DNS lookups
|
||||
- [x] Add CI/CD pipeline — GitHub Actions workflows in `.github/workflows/ci.yml`
|
||||
(lint → test/security/CodeQL/dependency-review → Docker build/push → GitOps)
|
||||
and `.github/workflows/release.yml` (semantic versioning)
|
||||
- [ ] Reach >80% test coverage
|
||||
|
||||
+6
-2
@@ -1,4 +1,4 @@
|
||||
FROM python:3.10-slim
|
||||
FROM python:3.13-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
@@ -19,7 +19,11 @@ RUN pip install --no-cache-dir -r requirements.txt
|
||||
# Copy application code including templates and static assets
|
||||
COPY . .
|
||||
|
||||
# Make the entrypoint executable and create the default data directory so
|
||||
# SQLite has a place to write its file even without an explicit volume mount.
|
||||
RUN chmod +x /app/entrypoint.sh && mkdir -p /app/data
|
||||
|
||||
# Expose application port
|
||||
EXPOSE 8080
|
||||
|
||||
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8080"]
|
||||
ENTRYPOINT ["/app/entrypoint.sh"]
|
||||
@@ -0,0 +1,149 @@
|
||||
# A generic, single database configuration.
|
||||
|
||||
[alembic]
|
||||
# path to migration scripts.
|
||||
# this is typically a path given in POSIX (e.g. forward slashes)
|
||||
# format, relative to the token %(here)s which refers to the location of this
|
||||
# ini file
|
||||
script_location = %(here)s/alembic
|
||||
|
||||
# template used to generate migration file names; The default value is %%(rev)s_%%(slug)s
|
||||
# Uncomment the line below if you want the files to be prepended with date and time
|
||||
# see https://alembic.sqlalchemy.org/en/latest/tutorial.html#editing-the-ini-file
|
||||
# for all available tokens
|
||||
# file_template = %%(year)d_%%(month).2d_%%(day).2d_%%(hour).2d%%(minute).2d-%%(rev)s_%%(slug)s
|
||||
# Or organize into date-based subdirectories (requires recursive_version_locations = true)
|
||||
# file_template = %%(year)d/%%(month).2d/%%(day).2d_%%(hour).2d%%(minute).2d_%%(second).2d_%%(rev)s_%%(slug)s
|
||||
|
||||
# sys.path path, will be prepended to sys.path if present.
|
||||
# defaults to the current working directory. for multiple paths, the path separator
|
||||
# is defined by "path_separator" below.
|
||||
prepend_sys_path = .
|
||||
|
||||
|
||||
# timezone to use when rendering the date within the migration file
|
||||
# as well as the filename.
|
||||
# If specified, requires the tzdata library which can be installed by adding
|
||||
# `alembic[tz]` to the pip requirements.
|
||||
# string value is passed to ZoneInfo()
|
||||
# leave blank for localtime
|
||||
# timezone =
|
||||
|
||||
# max length of characters to apply to the "slug" field
|
||||
# truncate_slug_length = 40
|
||||
|
||||
# set to 'true' to run the environment during
|
||||
# the 'revision' command, regardless of autogenerate
|
||||
# revision_environment = false
|
||||
|
||||
# set to 'true' to allow .pyc and .pyo files without
|
||||
# a source .py file to be detected as revisions in the
|
||||
# versions/ directory
|
||||
# sourceless = false
|
||||
|
||||
# version location specification; This defaults
|
||||
# to <script_location>/versions. When using multiple version
|
||||
# directories, initial revisions must be specified with --version-path.
|
||||
# The path separator used here should be the separator specified by "path_separator"
|
||||
# below.
|
||||
# version_locations = %(here)s/bar:%(here)s/bat:%(here)s/alembic/versions
|
||||
|
||||
# path_separator; This indicates what character is used to split lists of file
|
||||
# paths, including version_locations and prepend_sys_path within configparser
|
||||
# files such as alembic.ini.
|
||||
# The default rendered in new alembic.ini files is "os", which uses os.pathsep
|
||||
# to provide os-dependent path splitting.
|
||||
#
|
||||
# Note that in order to support legacy alembic.ini files, this default does NOT
|
||||
# take place if path_separator is not present in alembic.ini. If this
|
||||
# option is omitted entirely, fallback logic is as follows:
|
||||
#
|
||||
# 1. Parsing of the version_locations option falls back to using the legacy
|
||||
# "version_path_separator" key, which if absent then falls back to the legacy
|
||||
# behavior of splitting on spaces and/or commas.
|
||||
# 2. Parsing of the prepend_sys_path option falls back to the legacy
|
||||
# behavior of splitting on spaces, commas, or colons.
|
||||
#
|
||||
# Valid values for path_separator are:
|
||||
#
|
||||
# path_separator = :
|
||||
# path_separator = ;
|
||||
# path_separator = space
|
||||
# path_separator = newline
|
||||
#
|
||||
# Use os.pathsep. Default configuration used for new projects.
|
||||
path_separator = os
|
||||
|
||||
# set to 'true' to search source files recursively
|
||||
# in each "version_locations" directory
|
||||
# new in Alembic version 1.10
|
||||
# recursive_version_locations = false
|
||||
|
||||
# the output encoding used when revision files
|
||||
# are written from script.py.mako
|
||||
# output_encoding = utf-8
|
||||
|
||||
# database URL. This is consumed by the user-maintained env.py script only.
|
||||
# other means of configuring database URLs may be customized within the env.py
|
||||
# file.
|
||||
sqlalchemy.url = sqlite:///./dmarq.db
|
||||
|
||||
|
||||
[post_write_hooks]
|
||||
# post_write_hooks defines scripts or Python functions that are run
|
||||
# on newly generated revision scripts. See the documentation for further
|
||||
# detail and examples
|
||||
|
||||
# format using "black" - use the console_scripts runner, against the "black" entrypoint
|
||||
# hooks = black
|
||||
# black.type = console_scripts
|
||||
# black.entrypoint = black
|
||||
# black.options = -l 79 REVISION_SCRIPT_FILENAME
|
||||
|
||||
# lint with attempts to fix using "ruff" - use the module runner, against the "ruff" module
|
||||
# hooks = ruff
|
||||
# ruff.type = module
|
||||
# ruff.module = ruff
|
||||
# ruff.options = check --fix REVISION_SCRIPT_FILENAME
|
||||
|
||||
# Alternatively, use the exec runner to execute a binary found on your PATH
|
||||
# hooks = ruff
|
||||
# ruff.type = exec
|
||||
# ruff.executable = ruff
|
||||
# ruff.options = check --fix REVISION_SCRIPT_FILENAME
|
||||
|
||||
# Logging configuration. This is also consumed by the user-maintained
|
||||
# env.py script only.
|
||||
[loggers]
|
||||
keys = root,sqlalchemy,alembic
|
||||
|
||||
[handlers]
|
||||
keys = console
|
||||
|
||||
[formatters]
|
||||
keys = generic
|
||||
|
||||
[logger_root]
|
||||
level = WARNING
|
||||
handlers = console
|
||||
qualname =
|
||||
|
||||
[logger_sqlalchemy]
|
||||
level = WARNING
|
||||
handlers =
|
||||
qualname = sqlalchemy.engine
|
||||
|
||||
[logger_alembic]
|
||||
level = INFO
|
||||
handlers =
|
||||
qualname = alembic
|
||||
|
||||
[handler_console]
|
||||
class = StreamHandler
|
||||
args = (sys.stderr,)
|
||||
level = NOTSET
|
||||
formatter = generic
|
||||
|
||||
[formatter_generic]
|
||||
format = %(levelname)-5.5s [%(name)s] %(message)s
|
||||
datefmt = %H:%M:%S
|
||||
@@ -0,0 +1 @@
|
||||
Generic single-database configuration.
|
||||
@@ -0,0 +1,92 @@
|
||||
import os
|
||||
from logging.config import fileConfig
|
||||
|
||||
from alembic import context
|
||||
from sqlalchemy import engine_from_config, pool
|
||||
|
||||
# this is the Alembic Config object, which provides
|
||||
# access to the values within the .ini file in use.
|
||||
config = context.config
|
||||
|
||||
# Interpret the config file for Python logging.
|
||||
# This line sets up loggers basically.
|
||||
if config.config_file_name is not None:
|
||||
fileConfig(config.config_file_name)
|
||||
|
||||
# Override sqlalchemy.url with DATABASE_URL environment variable when present
|
||||
database_url = os.environ.get("DATABASE_URL")
|
||||
if database_url:
|
||||
from app.core.database import _make_sync_db_url # noqa: E402
|
||||
|
||||
config.set_main_option("sqlalchemy.url", _make_sync_db_url(database_url))
|
||||
|
||||
import app.models.alert # noqa: E402, F401
|
||||
import app.models.api_token # noqa: E402, F401
|
||||
import app.models.dns_cache # noqa: E402, F401
|
||||
import app.models.domain # noqa: E402, F401
|
||||
import app.models.mail_source # noqa: E402, F401
|
||||
import app.models.mail_source_import # noqa: E402, F401
|
||||
import app.models.report # noqa: E402, F401
|
||||
import app.models.setting # noqa: E402, F401
|
||||
import app.models.user # noqa: E402, F401
|
||||
import app.models.webhook # noqa: E402, F401
|
||||
import app.models.workspace # noqa: E402, F401
|
||||
import app.models.workspace_access # noqa: E402, F401
|
||||
|
||||
# Import all models so that autogenerate can detect them
|
||||
from app.core.database import Base # noqa: E402
|
||||
|
||||
target_metadata = Base.metadata
|
||||
|
||||
|
||||
def run_migrations_offline() -> None:
|
||||
"""Run migrations in 'offline' mode.
|
||||
|
||||
This configures the context with just a URL
|
||||
and not an Engine, though an Engine is acceptable
|
||||
here as well. By skipping the Engine creation
|
||||
we don't even need a DBAPI to be available.
|
||||
|
||||
Calls to context.execute() here emit the given string to the
|
||||
script output.
|
||||
|
||||
"""
|
||||
url = config.get_main_option("sqlalchemy.url")
|
||||
context.configure(
|
||||
url=url,
|
||||
target_metadata=target_metadata,
|
||||
literal_binds=True,
|
||||
dialect_opts={"paramstyle": "named"},
|
||||
)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
def run_migrations_online() -> None:
|
||||
"""Run migrations in 'online' mode.
|
||||
|
||||
In this scenario we need to create an Engine
|
||||
and associate a connection with the context.
|
||||
|
||||
"""
|
||||
connectable = engine_from_config(
|
||||
config.get_section(config.config_ini_section, {}),
|
||||
prefix="sqlalchemy.",
|
||||
poolclass=pool.NullPool,
|
||||
)
|
||||
|
||||
with connectable.connect() as connection:
|
||||
context.configure(
|
||||
connection=connection,
|
||||
target_metadata=target_metadata,
|
||||
)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
if context.is_offline_mode():
|
||||
run_migrations_offline()
|
||||
else:
|
||||
run_migrations_online()
|
||||
@@ -0,0 +1,28 @@
|
||||
"""${message}
|
||||
|
||||
Revision ID: ${up_revision}
|
||||
Revises: ${down_revision | comma,n}
|
||||
Create Date: ${create_date}
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
${imports if imports else ""}
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = ${repr(up_revision)}
|
||||
down_revision: Union[str, Sequence[str], None] = ${repr(down_revision)}
|
||||
branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)}
|
||||
depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)}
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Upgrade schema."""
|
||||
${upgrades if upgrades else "pass"}
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Downgrade schema."""
|
||||
${downgrades if downgrades else "pass"}
|
||||
@@ -0,0 +1,62 @@
|
||||
"""add scoped api tokens
|
||||
|
||||
Revision ID: 0a1b2c3d4e5f
|
||||
Revises: f7a8b9c0d1e2
|
||||
Create Date: 2026-05-23 00:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "0a1b2c3d4e5f"
|
||||
down_revision: Union[str, Sequence[str], None] = "f7a8b9c0d1e2"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create scoped API token storage."""
|
||||
op.create_table(
|
||||
"api_tokens",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("name", sa.String(length=120), nullable=False),
|
||||
sa.Column("key_hash", sa.String(length=255), nullable=False),
|
||||
sa.Column("key_prefix", sa.String(length=16), nullable=False),
|
||||
sa.Column("scopes", sa.Text(), nullable=False),
|
||||
sa.Column("active", sa.Boolean(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("revoked_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("last_used_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("last_used_ip", sa.String(length=64), nullable=True),
|
||||
sa.Column("usage_count", sa.Integer(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("key_hash"),
|
||||
)
|
||||
op.create_index(op.f("ix_api_tokens_id"), "api_tokens", ["id"])
|
||||
op.create_index(op.f("ix_api_tokens_key_hash"), "api_tokens", ["key_hash"])
|
||||
op.create_index(op.f("ix_api_tokens_key_prefix"), "api_tokens", ["key_prefix"])
|
||||
op.create_index(op.f("ix_api_tokens_active"), "api_tokens", ["active"])
|
||||
op.create_index(op.f("ix_api_tokens_created_at"), "api_tokens", ["created_at"])
|
||||
op.create_index(op.f("ix_api_tokens_revoked_at"), "api_tokens", ["revoked_at"])
|
||||
op.create_index(op.f("ix_api_tokens_last_used_at"), "api_tokens", ["last_used_at"])
|
||||
op.create_index("ix_api_tokens_active_scope", "api_tokens", ["active", "scopes"])
|
||||
op.create_index("ix_api_tokens_last_used", "api_tokens", ["last_used_at"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop scoped API token storage."""
|
||||
op.drop_index("ix_api_tokens_last_used", table_name="api_tokens")
|
||||
op.drop_index("ix_api_tokens_active_scope", table_name="api_tokens")
|
||||
op.drop_index(op.f("ix_api_tokens_last_used_at"), table_name="api_tokens")
|
||||
op.drop_index(op.f("ix_api_tokens_revoked_at"), table_name="api_tokens")
|
||||
op.drop_index(op.f("ix_api_tokens_created_at"), table_name="api_tokens")
|
||||
op.drop_index(op.f("ix_api_tokens_active"), table_name="api_tokens")
|
||||
op.drop_index(op.f("ix_api_tokens_key_prefix"), table_name="api_tokens")
|
||||
op.drop_index(op.f("ix_api_tokens_key_hash"), table_name="api_tokens")
|
||||
op.drop_index(op.f("ix_api_tokens_id"), table_name="api_tokens")
|
||||
op.drop_table("api_tokens")
|
||||
@@ -0,0 +1,137 @@
|
||||
"""add webhook event framework
|
||||
|
||||
Revision ID: 1b2c3d4e5f6a
|
||||
Revises: 0a1b2c3d4e5f
|
||||
Create Date: 2026-05-23 00:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "1b2c3d4e5f6a"
|
||||
down_revision: Union[str, Sequence[str], None] = "0a1b2c3d4e5f"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create outbound webhook endpoint and delivery tables."""
|
||||
op.create_table(
|
||||
"webhook_endpoints",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("name", sa.String(length=120), nullable=False),
|
||||
sa.Column("url", sa.Text(), nullable=False),
|
||||
sa.Column("secret", sa.Text(), nullable=False),
|
||||
sa.Column("event_types", sa.Text(), nullable=False),
|
||||
sa.Column("enabled", sa.Boolean(), nullable=False),
|
||||
sa.Column("max_attempts", sa.Integer(), nullable=False),
|
||||
sa.Column("timeout_seconds", sa.Integer(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("last_success_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("last_failure_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("failure_count", sa.Integer(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(op.f("ix_webhook_endpoints_id"), "webhook_endpoints", ["id"])
|
||||
op.create_index(op.f("ix_webhook_endpoints_enabled"), "webhook_endpoints", ["enabled"])
|
||||
op.create_index(op.f("ix_webhook_endpoints_created_at"), "webhook_endpoints", ["created_at"])
|
||||
op.create_index(
|
||||
op.f("ix_webhook_endpoints_last_success_at"),
|
||||
"webhook_endpoints",
|
||||
["last_success_at"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_webhook_endpoints_last_failure_at"),
|
||||
"webhook_endpoints",
|
||||
["last_failure_at"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_webhook_endpoints_enabled_events",
|
||||
"webhook_endpoints",
|
||||
["enabled", "event_types"],
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"webhook_deliveries",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("endpoint_id", sa.Integer(), nullable=False),
|
||||
sa.Column("event_type", sa.String(length=80), nullable=False),
|
||||
sa.Column("payload", sa.Text(), nullable=False),
|
||||
sa.Column("idempotency_key", sa.String(length=160), nullable=False),
|
||||
sa.Column("status", sa.String(length=24), nullable=False),
|
||||
sa.Column("attempt_count", sa.Integer(), nullable=False),
|
||||
sa.Column("max_attempts", sa.Integer(), nullable=False),
|
||||
sa.Column("next_attempt_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("last_attempt_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("delivered_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("last_status_code", sa.Integer(), nullable=True),
|
||||
sa.Column("last_error", sa.Text(), nullable=True),
|
||||
sa.Column("response_excerpt", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=True),
|
||||
sa.ForeignKeyConstraint(["endpoint_id"], ["webhook_endpoints.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(op.f("ix_webhook_deliveries_id"), "webhook_deliveries", ["id"])
|
||||
op.create_index(
|
||||
op.f("ix_webhook_deliveries_endpoint_id"), "webhook_deliveries", ["endpoint_id"]
|
||||
)
|
||||
op.create_index(op.f("ix_webhook_deliveries_event_type"), "webhook_deliveries", ["event_type"])
|
||||
op.create_index(
|
||||
op.f("ix_webhook_deliveries_idempotency_key"), "webhook_deliveries", ["idempotency_key"]
|
||||
)
|
||||
op.create_index(op.f("ix_webhook_deliveries_status"), "webhook_deliveries", ["status"])
|
||||
op.create_index(
|
||||
op.f("ix_webhook_deliveries_next_attempt_at"),
|
||||
"webhook_deliveries",
|
||||
["next_attempt_at"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_webhook_deliveries_last_attempt_at"),
|
||||
"webhook_deliveries",
|
||||
["last_attempt_at"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_webhook_deliveries_delivered_at"), "webhook_deliveries", ["delivered_at"]
|
||||
)
|
||||
op.create_index(op.f("ix_webhook_deliveries_created_at"), "webhook_deliveries", ["created_at"])
|
||||
op.create_index(
|
||||
"ix_webhook_delivery_endpoint_idempotency",
|
||||
"webhook_deliveries",
|
||||
["endpoint_id", "idempotency_key"],
|
||||
unique=True,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_webhook_delivery_due",
|
||||
"webhook_deliveries",
|
||||
["status", "next_attempt_at"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop outbound webhook endpoint and delivery tables."""
|
||||
op.drop_index("ix_webhook_delivery_due", table_name="webhook_deliveries")
|
||||
op.drop_index("ix_webhook_delivery_endpoint_idempotency", table_name="webhook_deliveries")
|
||||
op.drop_index(op.f("ix_webhook_deliveries_created_at"), table_name="webhook_deliveries")
|
||||
op.drop_index(op.f("ix_webhook_deliveries_delivered_at"), table_name="webhook_deliveries")
|
||||
op.drop_index(op.f("ix_webhook_deliveries_last_attempt_at"), table_name="webhook_deliveries")
|
||||
op.drop_index(op.f("ix_webhook_deliveries_next_attempt_at"), table_name="webhook_deliveries")
|
||||
op.drop_index(op.f("ix_webhook_deliveries_status"), table_name="webhook_deliveries")
|
||||
op.drop_index(op.f("ix_webhook_deliveries_idempotency_key"), table_name="webhook_deliveries")
|
||||
op.drop_index(op.f("ix_webhook_deliveries_event_type"), table_name="webhook_deliveries")
|
||||
op.drop_index(op.f("ix_webhook_deliveries_endpoint_id"), table_name="webhook_deliveries")
|
||||
op.drop_index(op.f("ix_webhook_deliveries_id"), table_name="webhook_deliveries")
|
||||
op.drop_table("webhook_deliveries")
|
||||
|
||||
op.drop_index("ix_webhook_endpoints_enabled_events", table_name="webhook_endpoints")
|
||||
op.drop_index(op.f("ix_webhook_endpoints_last_failure_at"), table_name="webhook_endpoints")
|
||||
op.drop_index(op.f("ix_webhook_endpoints_last_success_at"), table_name="webhook_endpoints")
|
||||
op.drop_index(op.f("ix_webhook_endpoints_created_at"), table_name="webhook_endpoints")
|
||||
op.drop_index(op.f("ix_webhook_endpoints_enabled"), table_name="webhook_endpoints")
|
||||
op.drop_index(op.f("ix_webhook_endpoints_id"), table_name="webhook_endpoints")
|
||||
op.drop_table("webhook_endpoints")
|
||||
@@ -0,0 +1,129 @@
|
||||
"""add workspace foundations
|
||||
|
||||
Revision ID: 2c3d4e5f6a7b
|
||||
Revises: 1b2c3d4e5f6a
|
||||
Create Date: 2026-05-23 18:58:00.000000
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "2c3d4e5f6a7b"
|
||||
down_revision = "1b2c3d4e5f6a"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _default_workspace_id_sql() -> str:
|
||||
return "(SELECT id FROM workspaces WHERE slug = 'default')"
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create workspaces and attach existing single-tenant rows to default."""
|
||||
op.create_table(
|
||||
"workspaces",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("slug", sa.String(), nullable=False),
|
||||
sa.Column("name", sa.String(), nullable=False),
|
||||
sa.Column("description", sa.Text(), nullable=True),
|
||||
sa.Column("active", sa.Boolean(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("slug"),
|
||||
)
|
||||
op.create_index(op.f("ix_workspaces_id"), "workspaces", ["id"])
|
||||
op.create_index(op.f("ix_workspaces_slug"), "workspaces", ["slug"])
|
||||
op.create_index(op.f("ix_workspaces_active"), "workspaces", ["active"])
|
||||
op.create_index(op.f("ix_workspaces_created_at"), "workspaces", ["created_at"])
|
||||
op.create_index("ix_workspaces_active_slug", "workspaces", ["active", "slug"])
|
||||
|
||||
now = datetime.utcnow()
|
||||
workspaces = sa.table(
|
||||
"workspaces",
|
||||
sa.column("slug", sa.String()),
|
||||
sa.column("name", sa.String()),
|
||||
sa.column("description", sa.Text()),
|
||||
sa.column("active", sa.Boolean()),
|
||||
sa.column("created_at", sa.DateTime()),
|
||||
sa.column("updated_at", sa.DateTime()),
|
||||
)
|
||||
op.bulk_insert(
|
||||
workspaces,
|
||||
[
|
||||
{
|
||||
"slug": "default",
|
||||
"name": "Default Workspace",
|
||||
"description": "Automatically created for existing single-tenant installs.",
|
||||
"active": True,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
with op.batch_alter_table("domains") as batch_op:
|
||||
batch_op.add_column(sa.Column("workspace_id", sa.Integer(), nullable=True))
|
||||
batch_op.create_foreign_key(
|
||||
"fk_domains_workspace_id_workspaces",
|
||||
"workspaces",
|
||||
["workspace_id"],
|
||||
["id"],
|
||||
)
|
||||
batch_op.create_index(op.f("ix_domains_workspace_id"), ["workspace_id"])
|
||||
batch_op.create_index("ix_domains_workspace_name", ["workspace_id", "name"])
|
||||
|
||||
with op.batch_alter_table("mail_sources") as batch_op:
|
||||
batch_op.add_column(sa.Column("workspace_id", sa.Integer(), nullable=True))
|
||||
batch_op.create_foreign_key(
|
||||
"fk_mail_sources_workspace_id_workspaces",
|
||||
"workspaces",
|
||||
["workspace_id"],
|
||||
["id"],
|
||||
)
|
||||
batch_op.create_index(op.f("ix_mail_sources_workspace_id"), ["workspace_id"])
|
||||
batch_op.create_index("ix_mail_sources_workspace_enabled", ["workspace_id", "enabled"])
|
||||
|
||||
with op.batch_alter_table("users") as batch_op:
|
||||
batch_op.add_column(sa.Column("workspace_id", sa.Integer(), nullable=True))
|
||||
batch_op.create_foreign_key(
|
||||
"fk_users_workspace_id_workspaces",
|
||||
"workspaces",
|
||||
["workspace_id"],
|
||||
["id"],
|
||||
)
|
||||
batch_op.create_index(op.f("ix_users_workspace_id"), ["workspace_id"])
|
||||
|
||||
op.execute(f"UPDATE domains SET workspace_id = {_default_workspace_id_sql()}")
|
||||
op.execute(f"UPDATE mail_sources SET workspace_id = {_default_workspace_id_sql()}")
|
||||
op.execute(f"UPDATE users SET workspace_id = {_default_workspace_id_sql()}")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove workspace ownership columns and table."""
|
||||
with op.batch_alter_table("users") as batch_op:
|
||||
batch_op.drop_index(op.f("ix_users_workspace_id"))
|
||||
batch_op.drop_constraint("fk_users_workspace_id_workspaces", type_="foreignkey")
|
||||
batch_op.drop_column("workspace_id")
|
||||
|
||||
with op.batch_alter_table("mail_sources") as batch_op:
|
||||
batch_op.drop_index("ix_mail_sources_workspace_enabled")
|
||||
batch_op.drop_index(op.f("ix_mail_sources_workspace_id"))
|
||||
batch_op.drop_constraint("fk_mail_sources_workspace_id_workspaces", type_="foreignkey")
|
||||
batch_op.drop_column("workspace_id")
|
||||
|
||||
with op.batch_alter_table("domains") as batch_op:
|
||||
batch_op.drop_index("ix_domains_workspace_name")
|
||||
batch_op.drop_index(op.f("ix_domains_workspace_id"))
|
||||
batch_op.drop_constraint("fk_domains_workspace_id_workspaces", type_="foreignkey")
|
||||
batch_op.drop_column("workspace_id")
|
||||
|
||||
op.drop_index("ix_workspaces_active_slug", table_name="workspaces")
|
||||
op.drop_index(op.f("ix_workspaces_created_at"), table_name="workspaces")
|
||||
op.drop_index(op.f("ix_workspaces_active"), table_name="workspaces")
|
||||
op.drop_index(op.f("ix_workspaces_slug"), table_name="workspaces")
|
||||
op.drop_index(op.f("ix_workspaces_id"), table_name="workspaces")
|
||||
op.drop_table("workspaces")
|
||||
@@ -0,0 +1,147 @@
|
||||
"""add workspace rbac audit foundations
|
||||
|
||||
Revision ID: 3d4e5f6a7b8c
|
||||
Revises: 2c3d4e5f6a7b
|
||||
Create Date: 2026-05-23 19:20:00.000000
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "3d4e5f6a7b8c"
|
||||
down_revision = "2c3d4e5f6a7b"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create workspace membership and sanitized audit log tables."""
|
||||
op.create_table(
|
||||
"workspace_memberships",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("workspace_id", sa.Integer(), nullable=False),
|
||||
sa.Column("user_id", sa.Integer(), nullable=False),
|
||||
sa.Column("role", sa.String(length=50), nullable=False),
|
||||
sa.Column("active", sa.Boolean(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=True),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"]),
|
||||
sa.ForeignKeyConstraint(["workspace_id"], ["workspaces.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(op.f("ix_workspace_memberships_id"), "workspace_memberships", ["id"])
|
||||
op.create_index(
|
||||
op.f("ix_workspace_memberships_workspace_id"),
|
||||
"workspace_memberships",
|
||||
["workspace_id"],
|
||||
)
|
||||
op.create_index(op.f("ix_workspace_memberships_user_id"), "workspace_memberships", ["user_id"])
|
||||
op.create_index(op.f("ix_workspace_memberships_role"), "workspace_memberships", ["role"])
|
||||
op.create_index(op.f("ix_workspace_memberships_active"), "workspace_memberships", ["active"])
|
||||
op.create_index(
|
||||
op.f("ix_workspace_memberships_created_at"),
|
||||
"workspace_memberships",
|
||||
["created_at"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_workspace_memberships_workspace_user",
|
||||
"workspace_memberships",
|
||||
["workspace_id", "user_id"],
|
||||
unique=True,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_workspace_memberships_workspace_role",
|
||||
"workspace_memberships",
|
||||
["workspace_id", "role"],
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"workspace_audit_logs",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("workspace_id", sa.Integer(), nullable=False),
|
||||
sa.Column("actor_type", sa.String(length=50), nullable=False),
|
||||
sa.Column("actor_id", sa.String(length=120), nullable=True),
|
||||
sa.Column("action", sa.String(length=100), nullable=False),
|
||||
sa.Column("entity_type", sa.String(length=80), nullable=False),
|
||||
sa.Column("entity_id", sa.String(length=120), nullable=True),
|
||||
sa.Column("entity_name", sa.String(length=255), nullable=True),
|
||||
sa.Column("details", sa.Text(), nullable=True),
|
||||
sa.Column("ip_address", sa.String(length=64), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.ForeignKeyConstraint(["workspace_id"], ["workspaces.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(op.f("ix_workspace_audit_logs_id"), "workspace_audit_logs", ["id"])
|
||||
op.create_index(
|
||||
op.f("ix_workspace_audit_logs_workspace_id"),
|
||||
"workspace_audit_logs",
|
||||
["workspace_id"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_workspace_audit_logs_actor_type"),
|
||||
"workspace_audit_logs",
|
||||
["actor_type"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_workspace_audit_logs_actor_id"),
|
||||
"workspace_audit_logs",
|
||||
["actor_id"],
|
||||
)
|
||||
op.create_index(op.f("ix_workspace_audit_logs_action"), "workspace_audit_logs", ["action"])
|
||||
op.create_index(
|
||||
op.f("ix_workspace_audit_logs_entity_type"),
|
||||
"workspace_audit_logs",
|
||||
["entity_type"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_workspace_audit_logs_entity_id"),
|
||||
"workspace_audit_logs",
|
||||
["entity_id"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_workspace_audit_logs_created_at"),
|
||||
"workspace_audit_logs",
|
||||
["created_at"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_workspace_audit_workspace_created",
|
||||
"workspace_audit_logs",
|
||||
["workspace_id", "created_at"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_workspace_audit_workspace_action",
|
||||
"workspace_audit_logs",
|
||||
["workspace_id", "action"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_workspace_audit_entity",
|
||||
"workspace_audit_logs",
|
||||
["entity_type", "entity_id"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove workspace RBAC and audit foundations."""
|
||||
op.drop_index("ix_workspace_audit_entity", table_name="workspace_audit_logs")
|
||||
op.drop_index("ix_workspace_audit_workspace_action", table_name="workspace_audit_logs")
|
||||
op.drop_index("ix_workspace_audit_workspace_created", table_name="workspace_audit_logs")
|
||||
op.drop_index(op.f("ix_workspace_audit_logs_created_at"), table_name="workspace_audit_logs")
|
||||
op.drop_index(op.f("ix_workspace_audit_logs_entity_id"), table_name="workspace_audit_logs")
|
||||
op.drop_index(op.f("ix_workspace_audit_logs_entity_type"), table_name="workspace_audit_logs")
|
||||
op.drop_index(op.f("ix_workspace_audit_logs_action"), table_name="workspace_audit_logs")
|
||||
op.drop_index(op.f("ix_workspace_audit_logs_actor_id"), table_name="workspace_audit_logs")
|
||||
op.drop_index(op.f("ix_workspace_audit_logs_actor_type"), table_name="workspace_audit_logs")
|
||||
op.drop_index(op.f("ix_workspace_audit_logs_workspace_id"), table_name="workspace_audit_logs")
|
||||
op.drop_index(op.f("ix_workspace_audit_logs_id"), table_name="workspace_audit_logs")
|
||||
op.drop_table("workspace_audit_logs")
|
||||
|
||||
op.drop_index("ix_workspace_memberships_workspace_role", table_name="workspace_memberships")
|
||||
op.drop_index("ix_workspace_memberships_workspace_user", table_name="workspace_memberships")
|
||||
op.drop_index(op.f("ix_workspace_memberships_created_at"), table_name="workspace_memberships")
|
||||
op.drop_index(op.f("ix_workspace_memberships_active"), table_name="workspace_memberships")
|
||||
op.drop_index(op.f("ix_workspace_memberships_role"), table_name="workspace_memberships")
|
||||
op.drop_index(op.f("ix_workspace_memberships_user_id"), table_name="workspace_memberships")
|
||||
op.drop_index(op.f("ix_workspace_memberships_workspace_id"), table_name="workspace_memberships")
|
||||
op.drop_index(op.f("ix_workspace_memberships_id"), table_name="workspace_memberships")
|
||||
op.drop_table("workspace_memberships")
|
||||
@@ -0,0 +1,37 @@
|
||||
"""add workspace operator controls
|
||||
|
||||
Revision ID: 4e5f6a7b8c9d
|
||||
Revises: 3d4e5f6a7b8c
|
||||
Create Date: 2026-05-23 20:05:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "4e5f6a7b8c9d"
|
||||
down_revision: Union[str, None] = "3d4e5f6a7b8c"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"workspaces",
|
||||
sa.Column("report_retention_days", sa.Integer(), nullable=False, server_default="400"),
|
||||
)
|
||||
op.add_column(
|
||||
"workspaces",
|
||||
sa.Column("forensic_retention_days", sa.Integer(), nullable=False, server_default="90"),
|
||||
)
|
||||
op.add_column(
|
||||
"workspaces",
|
||||
sa.Column("tls_report_retention_days", sa.Integer(), nullable=False, server_default="400"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("workspaces", "tls_report_retention_days")
|
||||
op.drop_column("workspaces", "forensic_retention_days")
|
||||
op.drop_column("workspaces", "report_retention_days")
|
||||
@@ -0,0 +1,159 @@
|
||||
"""initial schema
|
||||
|
||||
Revision ID: 88b549786e2d
|
||||
Revises:
|
||||
Create Date: 2026-03-29 11:36:11.612752
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = '88b549786e2d'
|
||||
down_revision: Union[str, Sequence[str], None] = None
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Upgrade schema."""
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_table('domains',
|
||||
sa.Column('id', sa.Integer(), nullable=False),
|
||||
sa.Column('name', sa.String(), nullable=False),
|
||||
sa.Column('description', sa.Text(), nullable=True),
|
||||
sa.Column('active', sa.Boolean(), nullable=True),
|
||||
sa.Column('dmarc_policy', sa.String(), nullable=True),
|
||||
sa.Column('spf_record', sa.String(), nullable=True),
|
||||
sa.Column('dkim_selectors', sa.String(), nullable=True),
|
||||
sa.Column('verified', sa.Boolean(), nullable=True),
|
||||
sa.Column('verification_token', sa.String(), nullable=True),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=True),
|
||||
sa.Column('updated_at', sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_domains_active'), 'domains', ['active'], unique=False)
|
||||
op.create_index('ix_domains_active_verified', 'domains', ['active', 'verified'], unique=False)
|
||||
op.create_index(op.f('ix_domains_created_at'), 'domains', ['created_at'], unique=False)
|
||||
op.create_index(op.f('ix_domains_id'), 'domains', ['id'], unique=False)
|
||||
op.create_index(op.f('ix_domains_name'), 'domains', ['name'], unique=True)
|
||||
op.create_index('ix_domains_policy', 'domains', ['dmarc_policy'], unique=False)
|
||||
op.create_index('ix_domains_updated', 'domains', ['updated_at'], unique=False)
|
||||
op.create_index(op.f('ix_domains_verified'), 'domains', ['verified'], unique=False)
|
||||
op.create_table('users',
|
||||
sa.Column('id', sa.Integer(), nullable=False),
|
||||
sa.Column('email', sa.String(), nullable=False),
|
||||
sa.Column('hashed_password', sa.String(), nullable=False),
|
||||
sa.Column('is_active', sa.Boolean(), nullable=True),
|
||||
sa.Column('is_superuser', sa.Boolean(), nullable=True),
|
||||
sa.Column('is_verified', sa.Boolean(), nullable=True),
|
||||
sa.Column('full_name', sa.String(), nullable=True),
|
||||
sa.Column('organization', sa.String(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_users_email'), 'users', ['email'], unique=True)
|
||||
op.create_index(op.f('ix_users_id'), 'users', ['id'], unique=False)
|
||||
op.create_table('dmarc_reports',
|
||||
sa.Column('id', sa.Integer(), nullable=False),
|
||||
sa.Column('domain_id', sa.Integer(), nullable=False),
|
||||
sa.Column('report_id', sa.String(), nullable=False),
|
||||
sa.Column('org_name', sa.String(), nullable=False),
|
||||
sa.Column('begin_date', sa.Integer(), nullable=False),
|
||||
sa.Column('end_date', sa.Integer(), nullable=False),
|
||||
sa.Column('source_email', sa.String(), nullable=True),
|
||||
sa.Column('policy', sa.String(), nullable=True),
|
||||
sa.Column('subdomain_policy', sa.String(), nullable=True),
|
||||
sa.Column('adkim', sa.String(length=1), nullable=True),
|
||||
sa.Column('aspf', sa.String(length=1), nullable=True),
|
||||
sa.Column('percentage', sa.Integer(), nullable=True),
|
||||
sa.Column('processed_at', sa.DateTime(), nullable=True),
|
||||
sa.Column('raw_data', sa.Text(), nullable=True),
|
||||
sa.ForeignKeyConstraint(['domain_id'], ['domains.id'], ),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_dmarc_reports_begin_date'), 'dmarc_reports', ['begin_date'], unique=False)
|
||||
op.create_index('ix_dmarc_reports_domain_dates', 'dmarc_reports', ['domain_id', 'begin_date', 'end_date'], unique=False)
|
||||
op.create_index(op.f('ix_dmarc_reports_domain_id'), 'dmarc_reports', ['domain_id'], unique=False)
|
||||
op.create_index(op.f('ix_dmarc_reports_end_date'), 'dmarc_reports', ['end_date'], unique=False)
|
||||
op.create_index(op.f('ix_dmarc_reports_id'), 'dmarc_reports', ['id'], unique=False)
|
||||
op.create_index(op.f('ix_dmarc_reports_org_name'), 'dmarc_reports', ['org_name'], unique=False)
|
||||
op.create_index('ix_dmarc_reports_policy', 'dmarc_reports', ['policy'], unique=False)
|
||||
op.create_index('ix_dmarc_reports_processed', 'dmarc_reports', ['processed_at'], unique=False)
|
||||
op.create_index(op.f('ix_dmarc_reports_report_id'), 'dmarc_reports', ['report_id'], unique=False)
|
||||
op.create_table('user_domains',
|
||||
sa.Column('id', sa.Integer(), nullable=False),
|
||||
sa.Column('user_id', sa.Integer(), nullable=False),
|
||||
sa.Column('domain_id', sa.Integer(), nullable=False),
|
||||
sa.Column('role', sa.String(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=True),
|
||||
sa.ForeignKeyConstraint(['domain_id'], ['domains.id'], ),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_user_domains_id'), 'user_domains', ['id'], unique=False)
|
||||
op.create_table('report_records',
|
||||
sa.Column('id', sa.Integer(), nullable=False),
|
||||
sa.Column('report_id', sa.Integer(), nullable=False),
|
||||
sa.Column('source_ip', sa.String(), nullable=False),
|
||||
sa.Column('count', sa.Integer(), nullable=False),
|
||||
sa.Column('disposition', sa.String(), nullable=False),
|
||||
sa.Column('dkim', sa.String(), nullable=True),
|
||||
sa.Column('spf', sa.String(), nullable=True),
|
||||
sa.Column('header_from', sa.String(), nullable=True),
|
||||
sa.Column('envelope_from', sa.String(), nullable=True),
|
||||
sa.Column('dkim_auth_details', sa.Text(), nullable=True),
|
||||
sa.Column('spf_auth_details', sa.Text(), nullable=True),
|
||||
sa.ForeignKeyConstraint(['report_id'], ['dmarc_reports.id'], ),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index('ix_report_records_disposition', 'report_records', ['disposition', 'count'], unique=False)
|
||||
op.create_index(op.f('ix_report_records_dkim'), 'report_records', ['dkim'], unique=False)
|
||||
op.create_index(op.f('ix_report_records_header_from'), 'report_records', ['header_from'], unique=False)
|
||||
op.create_index(op.f('ix_report_records_id'), 'report_records', ['id'], unique=False)
|
||||
op.create_index(op.f('ix_report_records_report_id'), 'report_records', ['report_id'], unique=False)
|
||||
op.create_index('ix_report_records_source_auth', 'report_records', ['source_ip', 'dkim', 'spf'], unique=False)
|
||||
op.create_index(op.f('ix_report_records_source_ip'), 'report_records', ['source_ip'], unique=False)
|
||||
op.create_index(op.f('ix_report_records_spf'), 'report_records', ['spf'], unique=False)
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Downgrade schema."""
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_index(op.f('ix_report_records_spf'), table_name='report_records')
|
||||
op.drop_index(op.f('ix_report_records_source_ip'), table_name='report_records')
|
||||
op.drop_index('ix_report_records_source_auth', table_name='report_records')
|
||||
op.drop_index(op.f('ix_report_records_report_id'), table_name='report_records')
|
||||
op.drop_index(op.f('ix_report_records_id'), table_name='report_records')
|
||||
op.drop_index(op.f('ix_report_records_header_from'), table_name='report_records')
|
||||
op.drop_index(op.f('ix_report_records_dkim'), table_name='report_records')
|
||||
op.drop_index('ix_report_records_disposition', table_name='report_records')
|
||||
op.drop_table('report_records')
|
||||
op.drop_index(op.f('ix_user_domains_id'), table_name='user_domains')
|
||||
op.drop_table('user_domains')
|
||||
op.drop_index(op.f('ix_dmarc_reports_report_id'), table_name='dmarc_reports')
|
||||
op.drop_index('ix_dmarc_reports_processed', table_name='dmarc_reports')
|
||||
op.drop_index('ix_dmarc_reports_policy', table_name='dmarc_reports')
|
||||
op.drop_index(op.f('ix_dmarc_reports_org_name'), table_name='dmarc_reports')
|
||||
op.drop_index(op.f('ix_dmarc_reports_id'), table_name='dmarc_reports')
|
||||
op.drop_index(op.f('ix_dmarc_reports_end_date'), table_name='dmarc_reports')
|
||||
op.drop_index(op.f('ix_dmarc_reports_domain_id'), table_name='dmarc_reports')
|
||||
op.drop_index('ix_dmarc_reports_domain_dates', table_name='dmarc_reports')
|
||||
op.drop_index(op.f('ix_dmarc_reports_begin_date'), table_name='dmarc_reports')
|
||||
op.drop_table('dmarc_reports')
|
||||
op.drop_index(op.f('ix_users_id'), table_name='users')
|
||||
op.drop_index(op.f('ix_users_email'), table_name='users')
|
||||
op.drop_table('users')
|
||||
op.drop_index(op.f('ix_domains_verified'), table_name='domains')
|
||||
op.drop_index('ix_domains_updated', table_name='domains')
|
||||
op.drop_index('ix_domains_policy', table_name='domains')
|
||||
op.drop_index(op.f('ix_domains_name'), table_name='domains')
|
||||
op.drop_index(op.f('ix_domains_id'), table_name='domains')
|
||||
op.drop_index(op.f('ix_domains_created_at'), table_name='domains')
|
||||
op.drop_index('ix_domains_active_verified', table_name='domains')
|
||||
op.drop_index(op.f('ix_domains_active'), table_name='domains')
|
||||
op.drop_table('domains')
|
||||
# ### end Alembic commands ###
|
||||
@@ -0,0 +1,107 @@
|
||||
"""add alert history
|
||||
|
||||
Revision ID: a7b8c9d0e1f2
|
||||
Revises: f6a7b8c9d0e1
|
||||
Create Date: 2026-05-22 23:12:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "a7b8c9d0e1f2"
|
||||
down_revision: Union[str, Sequence[str], None] = "f6a7b8c9d0e1"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create persisted alert history rows."""
|
||||
op.create_table(
|
||||
"alert_history",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("fingerprint", sa.String(length=64), nullable=False),
|
||||
sa.Column("rule", sa.String(), nullable=False),
|
||||
sa.Column("severity", sa.String(), nullable=False),
|
||||
sa.Column("domain", sa.String(), nullable=True),
|
||||
sa.Column("title", sa.String(), nullable=False),
|
||||
sa.Column("detail", sa.Text(), nullable=False),
|
||||
sa.Column("payload", sa.Text(), nullable=True),
|
||||
sa.Column("observed_count", sa.Integer(), nullable=False),
|
||||
sa.Column("is_active", sa.Boolean(), nullable=False),
|
||||
sa.Column("first_seen_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("last_seen_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("resolved_at", sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("fingerprint"),
|
||||
)
|
||||
op.create_index(op.f("ix_alert_history_id"), "alert_history", ["id"], unique=False)
|
||||
op.create_index(
|
||||
op.f("ix_alert_history_fingerprint"),
|
||||
"alert_history",
|
||||
["fingerprint"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(op.f("ix_alert_history_rule"), "alert_history", ["rule"], unique=False)
|
||||
op.create_index(
|
||||
op.f("ix_alert_history_severity"),
|
||||
"alert_history",
|
||||
["severity"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(op.f("ix_alert_history_domain"), "alert_history", ["domain"], unique=False)
|
||||
op.create_index(
|
||||
op.f("ix_alert_history_is_active"),
|
||||
"alert_history",
|
||||
["is_active"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_alert_history_first_seen_at"),
|
||||
"alert_history",
|
||||
["first_seen_at"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_alert_history_last_seen_at"),
|
||||
"alert_history",
|
||||
["last_seen_at"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_alert_history_resolved_at"),
|
||||
"alert_history",
|
||||
["resolved_at"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_alert_history_active_last_seen",
|
||||
"alert_history",
|
||||
["is_active", "last_seen_at"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_alert_history_rule_domain",
|
||||
"alert_history",
|
||||
["rule", "domain"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop persisted alert history rows."""
|
||||
op.drop_index("ix_alert_history_rule_domain", table_name="alert_history")
|
||||
op.drop_index("ix_alert_history_active_last_seen", table_name="alert_history")
|
||||
op.drop_index(op.f("ix_alert_history_resolved_at"), table_name="alert_history")
|
||||
op.drop_index(op.f("ix_alert_history_last_seen_at"), table_name="alert_history")
|
||||
op.drop_index(op.f("ix_alert_history_first_seen_at"), table_name="alert_history")
|
||||
op.drop_index(op.f("ix_alert_history_is_active"), table_name="alert_history")
|
||||
op.drop_index(op.f("ix_alert_history_domain"), table_name="alert_history")
|
||||
op.drop_index(op.f("ix_alert_history_severity"), table_name="alert_history")
|
||||
op.drop_index(op.f("ix_alert_history_rule"), table_name="alert_history")
|
||||
op.drop_index(op.f("ix_alert_history_fingerprint"), table_name="alert_history")
|
||||
op.drop_index(op.f("ix_alert_history_id"), table_name="alert_history")
|
||||
op.drop_table("alert_history")
|
||||
@@ -0,0 +1,40 @@
|
||||
"""add Gmail API fields to mail_sources
|
||||
|
||||
Revision ID: b2c3d4e5f6a7
|
||||
Revises: a1b2c3d4e5f6
|
||||
Create Date: 2026-03-29 19:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "b2c3d4e5f6a7"
|
||||
down_revision: Union[str, Sequence[str], None] = "a1b2c3d4e5f6"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add Gmail OAuth2 credential columns to mail_sources."""
|
||||
op.add_column("mail_sources", sa.Column("gmail_client_id", sa.String(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("gmail_client_secret", sa.Text(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("gmail_access_token", sa.Text(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("gmail_refresh_token", sa.Text(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("gmail_email", sa.String(), nullable=True))
|
||||
op.add_column(
|
||||
"mail_sources",
|
||||
sa.Column("gmail_ingested_ids", sa.Text(), nullable=True, server_default="[]"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove Gmail OAuth2 credential columns from mail_sources."""
|
||||
op.drop_column("mail_sources", "gmail_ingested_ids")
|
||||
op.drop_column("mail_sources", "gmail_email")
|
||||
op.drop_column("mail_sources", "gmail_refresh_token")
|
||||
op.drop_column("mail_sources", "gmail_access_token")
|
||||
op.drop_column("mail_sources", "gmail_client_secret")
|
||||
op.drop_column("mail_sources", "gmail_client_id")
|
||||
@@ -0,0 +1,50 @@
|
||||
"""add mail_sources table
|
||||
|
||||
Revision ID: a1b2c3d4e5f6
|
||||
Revises: 88b549786e2d
|
||||
Create Date: 2026-03-29 17:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "a1b2c3d4e5f6"
|
||||
down_revision: Union[str, Sequence[str], None] = "88b549786e2d"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create the mail_sources table."""
|
||||
op.create_table(
|
||||
"mail_sources",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("name", sa.String(), nullable=False),
|
||||
sa.Column("method", sa.String(), nullable=False),
|
||||
sa.Column("server", sa.String(), nullable=True),
|
||||
sa.Column("port", sa.Integer(), nullable=True),
|
||||
sa.Column("username", sa.String(), nullable=True),
|
||||
sa.Column("password", sa.Text(), nullable=True),
|
||||
sa.Column("use_ssl", sa.Boolean(), nullable=True),
|
||||
sa.Column("folder", sa.String(), nullable=True),
|
||||
sa.Column("polling_interval", sa.Integer(), nullable=True),
|
||||
sa.Column("enabled", sa.Boolean(), nullable=True),
|
||||
sa.Column("last_checked", sa.DateTime(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(op.f("ix_mail_sources_id"), "mail_sources", ["id"], unique=False)
|
||||
op.create_index(
|
||||
op.f("ix_mail_sources_enabled"), "mail_sources", ["enabled"], unique=False
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop the mail_sources table."""
|
||||
op.drop_index(op.f("ix_mail_sources_enabled"), table_name="mail_sources")
|
||||
op.drop_index(op.f("ix_mail_sources_id"), table_name="mail_sources")
|
||||
op.drop_table("mail_sources")
|
||||
@@ -0,0 +1,38 @@
|
||||
"""add settings table
|
||||
|
||||
Revision ID: c3d4e5f6a7b8
|
||||
Revises: b2c3d4e5f6a7
|
||||
Create Date: 2026-03-30 07:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "c3d4e5f6a7b8"
|
||||
down_revision: Union[str, Sequence[str], None] = "b2c3d4e5f6a7"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create the settings table."""
|
||||
op.create_table(
|
||||
"settings",
|
||||
sa.Column("key", sa.String(100), nullable=False),
|
||||
sa.Column("value", sa.Text(), nullable=True),
|
||||
sa.Column("description", sa.String(255), nullable=True),
|
||||
sa.Column("value_type", sa.String(20), nullable=False, server_default="string"),
|
||||
sa.Column("category", sa.String(50), nullable=False, server_default="general"),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("updated_by", sa.Integer(), nullable=True),
|
||||
sa.ForeignKeyConstraint(["updated_by"], ["users.id"]),
|
||||
sa.PrimaryKeyConstraint("key"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop the settings table."""
|
||||
op.drop_table("settings")
|
||||
@@ -0,0 +1,60 @@
|
||||
"""add dns cache
|
||||
|
||||
Revision ID: b8c9d0e1f2a3
|
||||
Revises: a7b8c9d0e1f2
|
||||
Create Date: 2026-05-22 23:28:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "b8c9d0e1f2a3"
|
||||
down_revision: Union[str, Sequence[str], None] = "a7b8c9d0e1f2"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create a DNS result cache table."""
|
||||
op.create_table(
|
||||
"dns_cache",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("domain", sa.String(), nullable=False),
|
||||
sa.Column("provider", sa.String(), nullable=False),
|
||||
sa.Column("selectors_key", sa.String(length=64), nullable=False),
|
||||
sa.Column("result_json", sa.Text(), nullable=False),
|
||||
sa.Column("checked_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("domain", "provider", "selectors_key", name="uq_dns_cache_lookup"),
|
||||
)
|
||||
op.create_index(op.f("ix_dns_cache_id"), "dns_cache", ["id"], unique=False)
|
||||
op.create_index(op.f("ix_dns_cache_domain"), "dns_cache", ["domain"], unique=False)
|
||||
op.create_index(op.f("ix_dns_cache_provider"), "dns_cache", ["provider"], unique=False)
|
||||
op.create_index(
|
||||
op.f("ix_dns_cache_selectors_key"),
|
||||
"dns_cache",
|
||||
["selectors_key"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(op.f("ix_dns_cache_checked_at"), "dns_cache", ["checked_at"], unique=False)
|
||||
op.create_index(
|
||||
"ix_dns_cache_domain_checked",
|
||||
"dns_cache",
|
||||
["domain", "checked_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop the DNS result cache table."""
|
||||
op.drop_index("ix_dns_cache_domain_checked", table_name="dns_cache")
|
||||
op.drop_index(op.f("ix_dns_cache_checked_at"), table_name="dns_cache")
|
||||
op.drop_index(op.f("ix_dns_cache_selectors_key"), table_name="dns_cache")
|
||||
op.drop_index(op.f("ix_dns_cache_provider"), table_name="dns_cache")
|
||||
op.drop_index(op.f("ix_dns_cache_domain"), table_name="dns_cache")
|
||||
op.drop_index(op.f("ix_dns_cache_id"), table_name="dns_cache")
|
||||
op.drop_table("dns_cache")
|
||||
@@ -0,0 +1,82 @@
|
||||
"""add alert configuration audit
|
||||
|
||||
Revision ID: b9c0d1e2f3a4
|
||||
Revises: b8c9d0e1f2a3
|
||||
Create Date: 2026-05-23 00:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "b9c0d1e2f3a4"
|
||||
down_revision: Union[str, Sequence[str], None] = "b8c9d0e1f2a3"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create alert configuration audit trail rows."""
|
||||
op.create_table(
|
||||
"alert_configuration_audit",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("key", sa.String(length=100), nullable=False),
|
||||
sa.Column("old_value", sa.Text(), nullable=True),
|
||||
sa.Column("new_value", sa.Text(), nullable=True),
|
||||
sa.Column("changed_by", sa.String(length=100), nullable=True),
|
||||
sa.Column("auth_type", sa.String(length=50), nullable=True),
|
||||
sa.Column("changed_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_alert_configuration_audit_id"),
|
||||
"alert_configuration_audit",
|
||||
["id"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_alert_configuration_audit_key"),
|
||||
"alert_configuration_audit",
|
||||
["key"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_alert_configuration_audit_changed_by"),
|
||||
"alert_configuration_audit",
|
||||
["changed_by"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_alert_configuration_audit_changed_at"),
|
||||
"alert_configuration_audit",
|
||||
["changed_at"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_alert_configuration_audit_key_changed_at",
|
||||
"alert_configuration_audit",
|
||||
["key", "changed_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop alert configuration audit trail rows."""
|
||||
op.drop_index(
|
||||
"ix_alert_configuration_audit_key_changed_at",
|
||||
table_name="alert_configuration_audit",
|
||||
)
|
||||
op.drop_index(
|
||||
op.f("ix_alert_configuration_audit_changed_at"),
|
||||
table_name="alert_configuration_audit",
|
||||
)
|
||||
op.drop_index(
|
||||
op.f("ix_alert_configuration_audit_changed_by"),
|
||||
table_name="alert_configuration_audit",
|
||||
)
|
||||
op.drop_index(op.f("ix_alert_configuration_audit_key"), table_name="alert_configuration_audit")
|
||||
op.drop_index(op.f("ix_alert_configuration_audit_id"), table_name="alert_configuration_audit")
|
||||
op.drop_table("alert_configuration_audit")
|
||||
@@ -0,0 +1,172 @@
|
||||
"""add dns record change tracking
|
||||
|
||||
Revision ID: c0d1e2f3a4b5
|
||||
Revises: b9c0d1e2f3a4
|
||||
Create Date: 2026-05-23 00:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "c0d1e2f3a4b5"
|
||||
down_revision: Union[str, Sequence[str], None] = "b9c0d1e2f3a4"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create DNS record snapshot and change history tables."""
|
||||
op.create_table(
|
||||
"dns_record_snapshots",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("domain", sa.String(), nullable=False),
|
||||
sa.Column("provider", sa.String(), nullable=False),
|
||||
sa.Column("zone_id", sa.String(), nullable=True),
|
||||
sa.Column("record_key", sa.String(length=128), nullable=False),
|
||||
sa.Column("record_id", sa.String(), nullable=True),
|
||||
sa.Column("record_type", sa.String(length=20), nullable=False),
|
||||
sa.Column("record_name", sa.String(), nullable=False),
|
||||
sa.Column("content", sa.Text(), nullable=True),
|
||||
sa.Column("proxied", sa.Boolean(), nullable=True),
|
||||
sa.Column("ttl", sa.Integer(), nullable=True),
|
||||
sa.Column("record_hash", sa.String(length=64), nullable=False),
|
||||
sa.Column("active", sa.Boolean(), nullable=False),
|
||||
sa.Column("first_seen_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("last_seen_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint(
|
||||
"domain", "provider", "record_key", name="uq_dns_record_snapshot_lookup"
|
||||
),
|
||||
)
|
||||
op.create_index(op.f("ix_dns_record_snapshots_id"), "dns_record_snapshots", ["id"])
|
||||
op.create_index(op.f("ix_dns_record_snapshots_domain"), "dns_record_snapshots", ["domain"])
|
||||
op.create_index(op.f("ix_dns_record_snapshots_provider"), "dns_record_snapshots", ["provider"])
|
||||
op.create_index(op.f("ix_dns_record_snapshots_zone_id"), "dns_record_snapshots", ["zone_id"])
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_snapshots_record_key"),
|
||||
"dns_record_snapshots",
|
||||
["record_key"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_snapshots_record_id"),
|
||||
"dns_record_snapshots",
|
||||
["record_id"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_snapshots_record_type"),
|
||||
"dns_record_snapshots",
|
||||
["record_type"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_snapshots_record_name"),
|
||||
"dns_record_snapshots",
|
||||
["record_name"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_snapshots_record_hash"),
|
||||
"dns_record_snapshots",
|
||||
["record_hash"],
|
||||
)
|
||||
op.create_index(op.f("ix_dns_record_snapshots_active"), "dns_record_snapshots", ["active"])
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_snapshots_first_seen_at"),
|
||||
"dns_record_snapshots",
|
||||
["first_seen_at"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_snapshots_last_seen_at"),
|
||||
"dns_record_snapshots",
|
||||
["last_seen_at"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_dns_record_snapshots_domain_active",
|
||||
"dns_record_snapshots",
|
||||
["domain", "active"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_dns_record_snapshots_domain_seen",
|
||||
"dns_record_snapshots",
|
||||
["domain", "last_seen_at"],
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"dns_record_changes",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("domain", sa.String(), nullable=False),
|
||||
sa.Column("provider", sa.String(), nullable=False),
|
||||
sa.Column("zone_id", sa.String(), nullable=True),
|
||||
sa.Column("record_key", sa.String(length=128), nullable=False),
|
||||
sa.Column("record_id", sa.String(), nullable=True),
|
||||
sa.Column("record_type", sa.String(length=20), nullable=False),
|
||||
sa.Column("record_name", sa.String(), nullable=False),
|
||||
sa.Column("change_type", sa.String(length=20), nullable=False),
|
||||
sa.Column("previous_content", sa.Text(), nullable=True),
|
||||
sa.Column("current_content", sa.Text(), nullable=True),
|
||||
sa.Column("observed_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(op.f("ix_dns_record_changes_id"), "dns_record_changes", ["id"])
|
||||
op.create_index(op.f("ix_dns_record_changes_domain"), "dns_record_changes", ["domain"])
|
||||
op.create_index(op.f("ix_dns_record_changes_provider"), "dns_record_changes", ["provider"])
|
||||
op.create_index(op.f("ix_dns_record_changes_zone_id"), "dns_record_changes", ["zone_id"])
|
||||
op.create_index(op.f("ix_dns_record_changes_record_key"), "dns_record_changes", ["record_key"])
|
||||
op.create_index(op.f("ix_dns_record_changes_record_id"), "dns_record_changes", ["record_id"])
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_changes_record_type"), "dns_record_changes", ["record_type"]
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_changes_record_name"), "dns_record_changes", ["record_name"]
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_changes_change_type"), "dns_record_changes", ["change_type"]
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_dns_record_changes_observed_at"), "dns_record_changes", ["observed_at"]
|
||||
)
|
||||
op.create_index(
|
||||
"ix_dns_record_changes_domain_observed",
|
||||
"dns_record_changes",
|
||||
["domain", "observed_at"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_dns_record_changes_record_observed",
|
||||
"dns_record_changes",
|
||||
["record_key", "observed_at"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop DNS record snapshot and change history tables."""
|
||||
op.drop_index("ix_dns_record_changes_record_observed", table_name="dns_record_changes")
|
||||
op.drop_index("ix_dns_record_changes_domain_observed", table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_observed_at"), table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_change_type"), table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_record_name"), table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_record_type"), table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_record_id"), table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_record_key"), table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_zone_id"), table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_provider"), table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_domain"), table_name="dns_record_changes")
|
||||
op.drop_index(op.f("ix_dns_record_changes_id"), table_name="dns_record_changes")
|
||||
op.drop_table("dns_record_changes")
|
||||
|
||||
op.drop_index("ix_dns_record_snapshots_domain_seen", table_name="dns_record_snapshots")
|
||||
op.drop_index("ix_dns_record_snapshots_domain_active", table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_last_seen_at"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_first_seen_at"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_active"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_record_hash"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_record_name"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_record_type"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_record_id"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_record_key"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_zone_id"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_provider"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_domain"), table_name="dns_record_snapshots")
|
||||
op.drop_index(op.f("ix_dns_record_snapshots_id"), table_name="dns_record_snapshots")
|
||||
op.drop_table("dns_record_snapshots")
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Add Microsoft 365 folder selection.
|
||||
|
||||
Revision ID: c4d5e6f7a8b9
|
||||
Revises: f3a4b5c6d7e8
|
||||
Create Date: 2026-05-23 00:00:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "c4d5e6f7a8b9"
|
||||
down_revision: Union[str, Sequence[str], None] = "f3a4b5c6d7e8"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("mail_sources", sa.Column("m365_folder_id", sa.String(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("mail_sources", "m365_folder_id")
|
||||
@@ -0,0 +1,88 @@
|
||||
"""add forensic reports
|
||||
|
||||
Revision ID: d1e2f3a4b5c6
|
||||
Revises: c0d1e2f3a4b5
|
||||
Create Date: 2026-05-23 00:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "d1e2f3a4b5c6"
|
||||
down_revision: Union[str, Sequence[str], None] = "c0d1e2f3a4b5"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create DMARC forensic/failure report storage."""
|
||||
op.create_table(
|
||||
"forensic_reports",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("domain_id", sa.Integer(), nullable=True),
|
||||
sa.Column("report_id", sa.String(), nullable=False),
|
||||
sa.Column("source_email", sa.String(), nullable=True),
|
||||
sa.Column("feedback_type", sa.String(), nullable=True),
|
||||
sa.Column("user_agent", sa.String(), nullable=True),
|
||||
sa.Column("version", sa.String(), nullable=True),
|
||||
sa.Column("reported_domain", sa.String(), nullable=True),
|
||||
sa.Column("source_ip", sa.String(), nullable=True),
|
||||
sa.Column("auth_failure", sa.String(), nullable=True),
|
||||
sa.Column("delivery_result", sa.String(), nullable=True),
|
||||
sa.Column("arrival_date", sa.DateTime(), nullable=True),
|
||||
sa.Column("authentication_results", sa.Text(), nullable=True),
|
||||
sa.Column("original_mail_from", sa.String(), nullable=True),
|
||||
sa.Column("original_from", sa.String(), nullable=True),
|
||||
sa.Column("original_to", sa.String(), nullable=True),
|
||||
sa.Column("original_subject", sa.String(), nullable=True),
|
||||
sa.Column("original_message_id", sa.String(), nullable=True),
|
||||
sa.Column("original_date", sa.String(), nullable=True),
|
||||
sa.Column("feedback_headers", sa.Text(), nullable=True),
|
||||
sa.Column("processed_at", sa.DateTime(), nullable=True),
|
||||
sa.ForeignKeyConstraint(["domain_id"], ["domains.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("report_id", name="uq_forensic_reports_report_id"),
|
||||
)
|
||||
op.create_index(op.f("ix_forensic_reports_id"), "forensic_reports", ["id"])
|
||||
op.create_index(op.f("ix_forensic_reports_domain_id"), "forensic_reports", ["domain_id"])
|
||||
op.create_index(op.f("ix_forensic_reports_report_id"), "forensic_reports", ["report_id"])
|
||||
op.create_index(
|
||||
op.f("ix_forensic_reports_feedback_type"), "forensic_reports", ["feedback_type"]
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_forensic_reports_reported_domain"), "forensic_reports", ["reported_domain"]
|
||||
)
|
||||
op.create_index(op.f("ix_forensic_reports_source_ip"), "forensic_reports", ["source_ip"])
|
||||
op.create_index(op.f("ix_forensic_reports_auth_failure"), "forensic_reports", ["auth_failure"])
|
||||
op.create_index(op.f("ix_forensic_reports_arrival_date"), "forensic_reports", ["arrival_date"])
|
||||
op.create_index(op.f("ix_forensic_reports_processed_at"), "forensic_reports", ["processed_at"])
|
||||
op.create_index(
|
||||
"ix_forensic_reports_domain_arrival",
|
||||
"forensic_reports",
|
||||
["domain_id", "arrival_date"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_forensic_reports_failure_source",
|
||||
"forensic_reports",
|
||||
["auth_failure", "source_ip"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop DMARC forensic/failure report storage."""
|
||||
op.drop_index("ix_forensic_reports_failure_source", table_name="forensic_reports")
|
||||
op.drop_index("ix_forensic_reports_domain_arrival", table_name="forensic_reports")
|
||||
op.drop_index(op.f("ix_forensic_reports_processed_at"), table_name="forensic_reports")
|
||||
op.drop_index(op.f("ix_forensic_reports_arrival_date"), table_name="forensic_reports")
|
||||
op.drop_index(op.f("ix_forensic_reports_auth_failure"), table_name="forensic_reports")
|
||||
op.drop_index(op.f("ix_forensic_reports_source_ip"), table_name="forensic_reports")
|
||||
op.drop_index(op.f("ix_forensic_reports_reported_domain"), table_name="forensic_reports")
|
||||
op.drop_index(op.f("ix_forensic_reports_feedback_type"), table_name="forensic_reports")
|
||||
op.drop_index(op.f("ix_forensic_reports_report_id"), table_name="forensic_reports")
|
||||
op.drop_index(op.f("ix_forensic_reports_domain_id"), table_name="forensic_reports")
|
||||
op.drop_index(op.f("ix_forensic_reports_id"), table_name="forensic_reports")
|
||||
op.drop_table("forensic_reports")
|
||||
@@ -0,0 +1,87 @@
|
||||
"""Add Logto fields to users table.
|
||||
|
||||
Adds the columns required for Logto OIDC integration and general user-profile
|
||||
enhancements:
|
||||
|
||||
- ``logto_id`` – the Logto subject claim (``sub``); acts as the stable
|
||||
external identity reference.
|
||||
- ``username`` – optional display username synced from Logto.
|
||||
- ``picture`` – profile-picture URL synced from Logto.
|
||||
- ``created_at`` – row-creation timestamp.
|
||||
- ``updated_at`` – last-update timestamp.
|
||||
|
||||
``hashed_password`` is made nullable because Logto-authenticated users
|
||||
authenticate externally and have no local password.
|
||||
|
||||
Revision ID: d4e5f6a7b8c9
|
||||
Revises: c3d4e5f6a7b8
|
||||
Create Date: 2026-03-30 10:00:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "d4e5f6a7b8c9"
|
||||
down_revision: Union[str, Sequence[str], None] = "c3d4e5f6a7b8"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Apply schema changes."""
|
||||
with op.batch_alter_table("users") as batch_op:
|
||||
# New columns
|
||||
batch_op.add_column(
|
||||
sa.Column("logto_id", sa.String(), nullable=True)
|
||||
)
|
||||
batch_op.add_column(
|
||||
sa.Column("username", sa.String(), nullable=True)
|
||||
)
|
||||
batch_op.add_column(
|
||||
sa.Column("picture", sa.String(), nullable=True)
|
||||
)
|
||||
batch_op.add_column(
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(),
|
||||
nullable=True,
|
||||
server_default=sa.func.now(),
|
||||
)
|
||||
)
|
||||
batch_op.add_column(
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(),
|
||||
nullable=True,
|
||||
server_default=sa.func.now(),
|
||||
)
|
||||
)
|
||||
# Make hashed_password nullable (Logto users have no local password)
|
||||
batch_op.alter_column("hashed_password", nullable=True)
|
||||
# Set is_superuser default to True (all users are admins for now)
|
||||
batch_op.alter_column("is_superuser", server_default=sa.true())
|
||||
|
||||
# Add unique index on logto_id
|
||||
op.create_index(
|
||||
op.f("ix_users_logto_id"),
|
||||
"users",
|
||||
["logto_id"],
|
||||
unique=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Revert schema changes."""
|
||||
op.drop_index(op.f("ix_users_logto_id"), table_name="users")
|
||||
|
||||
with op.batch_alter_table("users") as batch_op:
|
||||
batch_op.drop_column("updated_at")
|
||||
batch_op.drop_column("created_at")
|
||||
batch_op.drop_column("picture")
|
||||
batch_op.drop_column("username")
|
||||
batch_op.drop_column("logto_id")
|
||||
batch_op.alter_column("hashed_password", nullable=False)
|
||||
batch_op.alter_column("is_superuser", server_default=sa.false())
|
||||
@@ -0,0 +1,54 @@
|
||||
"""add dmarcbis aggregate fields
|
||||
|
||||
Revision ID: e2f3a4b5c6d7
|
||||
Revises: d1e2f3a4b5c6
|
||||
Create Date: 2026-05-23 00:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "e2f3a4b5c6d7"
|
||||
down_revision: Union[str, Sequence[str], None] = "d1e2f3a4b5c6"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Persist optional RFC 9990 / DMARCbis aggregate metadata."""
|
||||
op.add_column("dmarc_reports", sa.Column("extra_contact_info", sa.String(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("generator", sa.String(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("report_errors", sa.Text(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("non_subdomain_policy", sa.String(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("failure_options", sa.String(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("testing", sa.String(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("discovery_method", sa.String(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("schema_version", sa.String(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("report_variant", sa.String(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("xml_namespace", sa.String(), nullable=True))
|
||||
op.add_column("dmarc_reports", sa.Column("report_extensions", sa.Text(), nullable=True))
|
||||
op.add_column("report_records", sa.Column("envelope_to", sa.String(), nullable=True))
|
||||
op.add_column("report_records", sa.Column("policy_override_reasons", sa.Text(), nullable=True))
|
||||
op.add_column("report_records", sa.Column("record_extensions", sa.Text(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove optional RFC 9990 / DMARCbis aggregate metadata."""
|
||||
op.drop_column("report_records", "record_extensions")
|
||||
op.drop_column("report_records", "policy_override_reasons")
|
||||
op.drop_column("report_records", "envelope_to")
|
||||
op.drop_column("dmarc_reports", "report_extensions")
|
||||
op.drop_column("dmarc_reports", "xml_namespace")
|
||||
op.drop_column("dmarc_reports", "report_variant")
|
||||
op.drop_column("dmarc_reports", "schema_version")
|
||||
op.drop_column("dmarc_reports", "discovery_method")
|
||||
op.drop_column("dmarc_reports", "testing")
|
||||
op.drop_column("dmarc_reports", "failure_options")
|
||||
op.drop_column("dmarc_reports", "non_subdomain_policy")
|
||||
op.drop_column("dmarc_reports", "report_errors")
|
||||
op.drop_column("dmarc_reports", "generator")
|
||||
op.drop_column("dmarc_reports", "extra_contact_info")
|
||||
@@ -0,0 +1,73 @@
|
||||
"""add mail source import history
|
||||
|
||||
Revision ID: e5f6a7b8c9d0
|
||||
Revises: d4e5f6a7b8c9
|
||||
Create Date: 2026-05-22 19:40:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "e5f6a7b8c9d0"
|
||||
down_revision: Union[str, Sequence[str], None] = "d4e5f6a7b8c9"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create the mail_source_imports table."""
|
||||
op.create_table(
|
||||
"mail_source_imports",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("mail_source_id", sa.Integer(), nullable=False),
|
||||
sa.Column("trigger", sa.String(), nullable=False, server_default="manual"),
|
||||
sa.Column("status", sa.String(), nullable=False),
|
||||
sa.Column("processed", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("reports_found", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("duplicate_reports", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("error_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("new_domains", sa.Text(), nullable=True),
|
||||
sa.Column("errors", sa.Text(), nullable=True),
|
||||
sa.Column("started_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("finished_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.ForeignKeyConstraint(["mail_source_id"], ["mail_sources.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(op.f("ix_mail_source_imports_id"), "mail_source_imports", ["id"])
|
||||
op.create_index(
|
||||
op.f("ix_mail_source_imports_mail_source_id"),
|
||||
"mail_source_imports",
|
||||
["mail_source_id"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_mail_source_imports_status"),
|
||||
"mail_source_imports",
|
||||
["status"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_mail_source_imports_started_at"),
|
||||
"mail_source_imports",
|
||||
["started_at"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_mail_source_imports_finished_at"),
|
||||
"mail_source_imports",
|
||||
["finished_at"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop the mail_source_imports table."""
|
||||
op.drop_index(op.f("ix_mail_source_imports_finished_at"), table_name="mail_source_imports")
|
||||
op.drop_index(op.f("ix_mail_source_imports_started_at"), table_name="mail_source_imports")
|
||||
op.drop_index(op.f("ix_mail_source_imports_status"), table_name="mail_source_imports")
|
||||
op.drop_index(
|
||||
op.f("ix_mail_source_imports_mail_source_id"),
|
||||
table_name="mail_source_imports",
|
||||
)
|
||||
op.drop_index(op.f("ix_mail_source_imports_id"), table_name="mail_source_imports")
|
||||
op.drop_table("mail_source_imports")
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Add Microsoft 365 Graph mail source fields.
|
||||
|
||||
Revision ID: f3a4b5c6d7e8
|
||||
Revises: e2f3a4b5c6d7
|
||||
Create Date: 2026-05-23 00:00:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "f3a4b5c6d7e8"
|
||||
down_revision: Union[str, Sequence[str], None] = "e2f3a4b5c6d7"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("mail_sources", sa.Column("m365_tenant_id", sa.String(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("m365_client_id", sa.String(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("m365_client_secret", sa.Text(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("m365_access_token", sa.Text(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("m365_refresh_token", sa.Text(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("m365_mailbox", sa.String(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("m365_email", sa.String(), nullable=True))
|
||||
op.add_column("mail_sources", sa.Column("m365_ingested_ids", sa.Text(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("mail_sources", "m365_ingested_ids")
|
||||
op.drop_column("mail_sources", "m365_email")
|
||||
op.drop_column("mail_sources", "m365_mailbox")
|
||||
op.drop_column("mail_sources", "m365_refresh_token")
|
||||
op.drop_column("mail_sources", "m365_access_token")
|
||||
op.drop_column("mail_sources", "m365_client_secret")
|
||||
op.drop_column("mail_sources", "m365_client_id")
|
||||
op.drop_column("mail_sources", "m365_tenant_id")
|
||||
@@ -0,0 +1,27 @@
|
||||
"""add mail source import details
|
||||
|
||||
Revision ID: f6a7b8c9d0e1
|
||||
Revises: e5f6a7b8c9d0
|
||||
Create Date: 2026-05-22 20:05:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "f6a7b8c9d0e1"
|
||||
down_revision: Union[str, Sequence[str], None] = "e5f6a7b8c9d0"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add a sanitized details payload to import history rows."""
|
||||
op.add_column("mail_source_imports", sa.Column("details", sa.Text(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove import details payloads."""
|
||||
op.drop_column("mail_source_imports", "details")
|
||||
@@ -0,0 +1,128 @@
|
||||
"""add tls reports
|
||||
|
||||
Revision ID: f7a8b9c0d1e2
|
||||
Revises: c4d5e6f7a8b9
|
||||
Create Date: 2026-05-23 00:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "f7a8b9c0d1e2"
|
||||
down_revision: Union[str, Sequence[str], None] = "c4d5e6f7a8b9"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create privacy-conscious SMTP TLS report storage."""
|
||||
op.create_table(
|
||||
"tls_reports",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("domain_id", sa.Integer(), nullable=True),
|
||||
sa.Column("report_id", sa.String(), nullable=False),
|
||||
sa.Column("org_name", sa.String(), nullable=True),
|
||||
sa.Column("contact_info", sa.String(), nullable=True),
|
||||
sa.Column("policy_domain", sa.String(), nullable=False),
|
||||
sa.Column("policy_type", sa.String(), nullable=True),
|
||||
sa.Column("begin_date", sa.DateTime(), nullable=True),
|
||||
sa.Column("end_date", sa.DateTime(), nullable=True),
|
||||
sa.Column("total_successful_sessions", sa.Integer(), nullable=False),
|
||||
sa.Column("total_failure_sessions", sa.Integer(), nullable=False),
|
||||
sa.Column("raw_policy", sa.Text(), nullable=True),
|
||||
sa.Column("processed_at", sa.DateTime(), nullable=True),
|
||||
sa.ForeignKeyConstraint(["domain_id"], ["domains.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("report_id", "policy_domain", name="uq_tls_reports_report_domain"),
|
||||
)
|
||||
op.create_index(op.f("ix_tls_reports_id"), "tls_reports", ["id"])
|
||||
op.create_index(op.f("ix_tls_reports_domain_id"), "tls_reports", ["domain_id"])
|
||||
op.create_index(op.f("ix_tls_reports_report_id"), "tls_reports", ["report_id"])
|
||||
op.create_index(op.f("ix_tls_reports_org_name"), "tls_reports", ["org_name"])
|
||||
op.create_index(op.f("ix_tls_reports_policy_domain"), "tls_reports", ["policy_domain"])
|
||||
op.create_index(op.f("ix_tls_reports_policy_type"), "tls_reports", ["policy_type"])
|
||||
op.create_index(op.f("ix_tls_reports_begin_date"), "tls_reports", ["begin_date"])
|
||||
op.create_index(op.f("ix_tls_reports_end_date"), "tls_reports", ["end_date"])
|
||||
op.create_index(op.f("ix_tls_reports_processed_at"), "tls_reports", ["processed_at"])
|
||||
op.create_index(
|
||||
"ix_tls_reports_domain_dates",
|
||||
"tls_reports",
|
||||
["domain_id", "begin_date", "end_date"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_tls_reports_policy_domain_dates",
|
||||
"tls_reports",
|
||||
["policy_domain", "begin_date", "end_date"],
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"tls_report_failures",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("report_id", sa.Integer(), nullable=False),
|
||||
sa.Column("result_type", sa.String(), nullable=False),
|
||||
sa.Column("failed_session_count", sa.Integer(), nullable=False),
|
||||
sa.Column("sending_mta_ip", sa.String(), nullable=True),
|
||||
sa.Column("receiving_mx_hostname", sa.String(), nullable=True),
|
||||
sa.Column("receiving_mx_helo", sa.String(), nullable=True),
|
||||
sa.Column("receiving_ip", sa.String(), nullable=True),
|
||||
sa.Column("failure_reason_code", sa.String(), nullable=True),
|
||||
sa.Column("additional_information", sa.Text(), nullable=True),
|
||||
sa.ForeignKeyConstraint(["report_id"], ["tls_reports.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(op.f("ix_tls_report_failures_id"), "tls_report_failures", ["id"])
|
||||
op.create_index(op.f("ix_tls_report_failures_report_id"), "tls_report_failures", ["report_id"])
|
||||
op.create_index(
|
||||
op.f("ix_tls_report_failures_result_type"), "tls_report_failures", ["result_type"]
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_tls_report_failures_sending_mta_ip"),
|
||||
"tls_report_failures",
|
||||
["sending_mta_ip"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_tls_report_failures_receiving_mx_hostname"),
|
||||
"tls_report_failures",
|
||||
["receiving_mx_hostname"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_tls_report_failures_result_count",
|
||||
"tls_report_failures",
|
||||
["result_type", "failed_session_count"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_tls_report_failures_report_result",
|
||||
"tls_report_failures",
|
||||
["report_id", "result_type"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop SMTP TLS report storage."""
|
||||
op.drop_index("ix_tls_report_failures_report_result", table_name="tls_report_failures")
|
||||
op.drop_index("ix_tls_report_failures_result_count", table_name="tls_report_failures")
|
||||
op.drop_index(
|
||||
op.f("ix_tls_report_failures_receiving_mx_hostname"),
|
||||
table_name="tls_report_failures",
|
||||
)
|
||||
op.drop_index(op.f("ix_tls_report_failures_sending_mta_ip"), table_name="tls_report_failures")
|
||||
op.drop_index(op.f("ix_tls_report_failures_result_type"), table_name="tls_report_failures")
|
||||
op.drop_index(op.f("ix_tls_report_failures_report_id"), table_name="tls_report_failures")
|
||||
op.drop_index(op.f("ix_tls_report_failures_id"), table_name="tls_report_failures")
|
||||
op.drop_table("tls_report_failures")
|
||||
op.drop_index("ix_tls_reports_policy_domain_dates", table_name="tls_reports")
|
||||
op.drop_index("ix_tls_reports_domain_dates", table_name="tls_reports")
|
||||
op.drop_index(op.f("ix_tls_reports_processed_at"), table_name="tls_reports")
|
||||
op.drop_index(op.f("ix_tls_reports_end_date"), table_name="tls_reports")
|
||||
op.drop_index(op.f("ix_tls_reports_begin_date"), table_name="tls_reports")
|
||||
op.drop_index(op.f("ix_tls_reports_policy_type"), table_name="tls_reports")
|
||||
op.drop_index(op.f("ix_tls_reports_policy_domain"), table_name="tls_reports")
|
||||
op.drop_index(op.f("ix_tls_reports_org_name"), table_name="tls_reports")
|
||||
op.drop_index(op.f("ix_tls_reports_report_id"), table_name="tls_reports")
|
||||
op.drop_index(op.f("ix_tls_reports_domain_id"), table_name="tls_reports")
|
||||
op.drop_index(op.f("ix_tls_reports_id"), table_name="tls_reports")
|
||||
op.drop_table("tls_reports")
|
||||
@@ -1,3 +1,3 @@
|
||||
"""DMARQ - DMARC monitoring and analysis platform."""
|
||||
|
||||
__version__ = "1.0.0"
|
||||
__version__ = "1.56.0"
|
||||
|
||||
@@ -1,12 +1,50 @@
|
||||
from app.api.api_v1.endpoints import domains, health, imap, reports, setup, stats
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.api_v1.endpoints import (
|
||||
ai,
|
||||
api_tokens,
|
||||
audit,
|
||||
auth,
|
||||
domains,
|
||||
forensics,
|
||||
health,
|
||||
imap,
|
||||
integrations,
|
||||
mail_sources,
|
||||
mcp,
|
||||
onboarding,
|
||||
operator,
|
||||
public,
|
||||
reports,
|
||||
settings,
|
||||
setup,
|
||||
stats,
|
||||
tls_reports,
|
||||
webhook,
|
||||
webhooks,
|
||||
)
|
||||
|
||||
api_router = APIRouter()
|
||||
|
||||
# Include all endpoint routers
|
||||
api_router.include_router(auth.router, prefix="/auth", tags=["auth"])
|
||||
api_router.include_router(ai.router, prefix="/ai", tags=["ai"])
|
||||
api_router.include_router(api_tokens.router, prefix="/api-tokens", tags=["api-tokens"])
|
||||
api_router.include_router(audit.router, prefix="/audit", tags=["audit"])
|
||||
api_router.include_router(health.router, tags=["health"])
|
||||
api_router.include_router(public.router, prefix="/public", tags=["public-api"])
|
||||
api_router.include_router(domains.router, prefix="/domains", tags=["domains"])
|
||||
api_router.include_router(reports.router, prefix="/reports", tags=["reports"])
|
||||
api_router.include_router(forensics.router, prefix="/forensics", tags=["forensics"])
|
||||
api_router.include_router(setup.router, prefix="/setup", tags=["setup"])
|
||||
api_router.include_router(imap.router, prefix="/imap", tags=["imap"])
|
||||
api_router.include_router(integrations.router, prefix="/integrations", tags=["integrations"])
|
||||
api_router.include_router(stats.router, prefix="/stats", tags=["stats"])
|
||||
api_router.include_router(mail_sources.router, prefix="/mail-sources", tags=["mail-sources"])
|
||||
api_router.include_router(mcp.router, prefix="/mcp", tags=["mcp"])
|
||||
api_router.include_router(onboarding.router, prefix="/onboarding", tags=["onboarding"])
|
||||
api_router.include_router(operator.router, prefix="/operator", tags=["operator"])
|
||||
api_router.include_router(settings.router, prefix="/settings", tags=["settings"])
|
||||
api_router.include_router(tls_reports.router, prefix="/tls-reports", tags=["tls-reports"])
|
||||
api_router.include_router(webhook.router, prefix="/webhook", tags=["webhook"])
|
||||
api_router.include_router(webhooks.router, prefix="/webhooks", tags=["webhooks"])
|
||||
|
||||
@@ -0,0 +1,194 @@
|
||||
"""Optional AI and automation endpoints."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.services.ai_assistance import (
|
||||
build_action_proposals,
|
||||
build_evidence_summary,
|
||||
build_safe_context,
|
||||
get_assistance_config,
|
||||
)
|
||||
from app.services.workspace_audit import record_workspace_audit_log
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class SafeContextResponse(BaseModel):
|
||||
"""Redacted model/agent context response."""
|
||||
|
||||
context: Dict[str, Any]
|
||||
|
||||
|
||||
class EvidenceSummaryResponse(BaseModel):
|
||||
"""Evidence-first summary response."""
|
||||
|
||||
summary: Dict[str, Any]
|
||||
|
||||
|
||||
class ActionProposalResponse(BaseModel):
|
||||
"""Reviewable action proposals response."""
|
||||
|
||||
domain: str
|
||||
action_tools_enabled: bool
|
||||
proposals: list[Dict[str, Any]]
|
||||
|
||||
|
||||
class ProposalConfirmation(BaseModel):
|
||||
"""Human confirmation payload for a proposal."""
|
||||
|
||||
proposal_id: str
|
||||
confirmation_text: str
|
||||
note: Optional[str] = None
|
||||
|
||||
|
||||
def _require_ai_enabled(db: Session) -> None:
|
||||
if not get_assistance_config(db).ai_enabled:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="AI assistance is disabled. Enable ai.enabled before using this endpoint.",
|
||||
)
|
||||
|
||||
|
||||
@router.get("/config")
|
||||
async def get_ai_config(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Return safe AI/MCP configuration without secrets."""
|
||||
return {"config": get_assistance_config(db).to_dict()}
|
||||
|
||||
|
||||
@router.get("/domains/{domain}/context", response_model=SafeContextResponse)
|
||||
async def get_domain_safe_context(
|
||||
domain: str,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> SafeContextResponse:
|
||||
"""Return a redacted, evidence-linked context payload for one domain."""
|
||||
_require_ai_enabled(db)
|
||||
try:
|
||||
return {"context": build_safe_context(db, domain)}
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.get("/domains/{domain}/summary", response_model=EvidenceSummaryResponse)
|
||||
async def get_domain_evidence_summary(
|
||||
domain: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> EvidenceSummaryResponse:
|
||||
"""Return deterministic evidence-first assistance for one domain."""
|
||||
_require_ai_enabled(db)
|
||||
try:
|
||||
summary = build_evidence_summary(db, domain)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db, commit=False)
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="ai.summary_generated",
|
||||
entity_type="domain",
|
||||
entity_id=domain,
|
||||
entity_name=domain,
|
||||
details={
|
||||
"provider": summary["provider"]["provider"],
|
||||
"recommendations": len(summary["recommendations"]),
|
||||
},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
)
|
||||
db.commit()
|
||||
return {"summary": summary}
|
||||
|
||||
|
||||
@router.get("/domains/{domain}/action-proposals", response_model=ActionProposalResponse)
|
||||
async def get_domain_action_proposals(
|
||||
domain: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> ActionProposalResponse:
|
||||
"""Return reviewable proposals; this endpoint never applies changes."""
|
||||
_require_ai_enabled(db)
|
||||
try:
|
||||
payload = build_action_proposals(db, domain)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db, commit=False)
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="ai.action_proposals_generated",
|
||||
entity_type="domain",
|
||||
entity_id=domain,
|
||||
entity_name=domain,
|
||||
details={"proposal_count": len(payload["proposals"]), "mutates_state": False},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
)
|
||||
db.commit()
|
||||
return payload
|
||||
|
||||
|
||||
@router.post("/domains/{domain}/action-proposals/confirm")
|
||||
async def confirm_action_proposal(
|
||||
domain: str,
|
||||
payload: ProposalConfirmation,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Audit human confirmation for a proposal without applying external changes."""
|
||||
_require_ai_enabled(db)
|
||||
config = get_assistance_config(db)
|
||||
if not config.action_tools_enabled:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=(
|
||||
"Action tools are disabled. Enable ai.action_tools_enabled to confirm "
|
||||
"proposals."
|
||||
),
|
||||
)
|
||||
proposals = build_action_proposals(db, domain)["proposals"]
|
||||
proposal = next(
|
||||
(item for item in proposals if item["proposal_id"] == payload.proposal_id),
|
||||
None,
|
||||
)
|
||||
if proposal is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Proposal not found")
|
||||
if payload.confirmation_text != payload.proposal_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="confirmation_text must match proposal_id",
|
||||
)
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db, commit=False)
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="ai.action_proposal_confirmed",
|
||||
entity_type="action_proposal",
|
||||
entity_id=payload.proposal_id,
|
||||
entity_name=proposal["title"],
|
||||
details={
|
||||
"domain": domain,
|
||||
"proposal_id": payload.proposal_id,
|
||||
"mutates_state": False,
|
||||
"note": payload.note,
|
||||
},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
)
|
||||
db.commit()
|
||||
return {"status": "confirmed", "applied": False, "proposal": proposal}
|
||||
@@ -0,0 +1,137 @@
|
||||
"""Admin API token management endpoints."""
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.models.api_token import APIToken
|
||||
from app.services.api_tokens import (
|
||||
PUBLIC_READ_SCOPES,
|
||||
create_api_token,
|
||||
revoke_api_token,
|
||||
token_to_dict,
|
||||
)
|
||||
from app.services.workspace_audit import record_workspace_audit_log
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class APITokenCreateRequest(BaseModel):
|
||||
"""Request body for creating a scoped API token."""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=120)
|
||||
scopes: List[str] = Field(default_factory=lambda: sorted(PUBLIC_READ_SCOPES))
|
||||
|
||||
|
||||
class APITokenResponse(BaseModel):
|
||||
"""API-safe token metadata."""
|
||||
|
||||
id: int
|
||||
name: str
|
||||
key_prefix: str
|
||||
scopes: List[str]
|
||||
active: bool
|
||||
created_at: str
|
||||
last_used_at: Optional[str] = None
|
||||
last_used_ip: Optional[str] = None
|
||||
usage_count: int
|
||||
revoked_at: Optional[str] = None
|
||||
|
||||
|
||||
class APITokenCreateResponse(BaseModel):
|
||||
"""New token response. The secret is returned once."""
|
||||
|
||||
token: str
|
||||
metadata: APITokenResponse
|
||||
|
||||
|
||||
class APITokenListResponse(BaseModel):
|
||||
"""List of API token metadata rows."""
|
||||
|
||||
tokens: List[APITokenResponse]
|
||||
available_scopes: List[str]
|
||||
|
||||
|
||||
@router.get("", response_model=APITokenListResponse)
|
||||
async def list_api_tokens(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""List API token metadata without exposing raw secrets or hashes."""
|
||||
rows = db.query(APIToken).order_by(APIToken.created_at.desc(), APIToken.id.desc()).all()
|
||||
return APITokenListResponse(
|
||||
tokens=[APITokenResponse(**token_to_dict(row)) for row in rows],
|
||||
available_scopes=sorted(PUBLIC_READ_SCOPES),
|
||||
)
|
||||
|
||||
|
||||
@router.post("", response_model=APITokenCreateResponse, status_code=status.HTTP_201_CREATED)
|
||||
async def create_public_api_token(
|
||||
payload: APITokenCreateRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""Create a scoped API token for read-only automation."""
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db)
|
||||
try:
|
||||
created = create_api_token(db, name=payload.name, scopes=payload.scopes)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="api_token.created",
|
||||
entity_type="api_token",
|
||||
entity_id=created.token.id,
|
||||
entity_name=created.token.name,
|
||||
details={
|
||||
"scopes": sorted(created.token.scopes.split(",")),
|
||||
"key_prefix": created.token.key_prefix,
|
||||
},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
commit=True,
|
||||
)
|
||||
return APITokenCreateResponse(
|
||||
token=created.secret,
|
||||
metadata=APITokenResponse(**token_to_dict(created.token)),
|
||||
)
|
||||
|
||||
|
||||
@router.delete("/{token_id}", status_code=status.HTTP_200_OK)
|
||||
async def revoke_public_api_token(
|
||||
token_id: int,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""Revoke a scoped API token."""
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db)
|
||||
token = db.query(APIToken).filter(APIToken.id == token_id).first()
|
||||
if not revoke_api_token(db, token_id):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="API token not found",
|
||||
)
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="api_token.revoked",
|
||||
entity_type="api_token",
|
||||
entity_id=token_id,
|
||||
entity_name=token.name if token else None,
|
||||
details={"key_prefix": token.key_prefix if token else None},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
commit=True,
|
||||
)
|
||||
return {"revoked": True}
|
||||
@@ -0,0 +1,61 @@
|
||||
"""Workspace RBAC and audit endpoints."""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.services.workspace_access import (
|
||||
PERMISSION_AUDIT_READ,
|
||||
list_workspace_roles,
|
||||
require_workspace_permission,
|
||||
)
|
||||
from app.services.workspace_audit import list_workspace_audit_logs
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class WorkspaceRoleResponse(BaseModel):
|
||||
"""Workspace role and permission definitions."""
|
||||
|
||||
roles: List[Dict[str, Any]]
|
||||
|
||||
|
||||
class WorkspaceAuditLogResponse(BaseModel):
|
||||
"""Workspace audit log list response."""
|
||||
|
||||
audit: List[Dict[str, Any]]
|
||||
|
||||
|
||||
@router.get("/roles", response_model=WorkspaceRoleResponse)
|
||||
async def get_workspace_roles(
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> WorkspaceRoleResponse:
|
||||
"""Return the supported workspace role definitions."""
|
||||
return {"roles": list_workspace_roles()}
|
||||
|
||||
|
||||
@router.get("/logs", response_model=WorkspaceAuditLogResponse)
|
||||
async def get_workspace_audit_logs(
|
||||
limit: int = Query(50, ge=1, le=200),
|
||||
action: Optional[str] = None,
|
||||
entity_type: Optional[str] = None,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> WorkspaceAuditLogResponse:
|
||||
"""Return recent sanitized audit events for the default workspace."""
|
||||
require_workspace_permission(_auth, PERMISSION_AUDIT_READ)
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db)
|
||||
return {
|
||||
"audit": list_workspace_audit_logs(
|
||||
db,
|
||||
workspace=workspace,
|
||||
limit=limit,
|
||||
action=action,
|
||||
entity_type=entity_type,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,354 @@
|
||||
"""
|
||||
Authentication endpoints (Logto OIDC).
|
||||
|
||||
Routes
|
||||
------
|
||||
GET /sign-in – Initiate the Logto sign-in flow.
|
||||
GET /callback – Handle the Logto authorization-code callback.
|
||||
GET /sign-out – Sign the user out (clears session + redirects to Logto).
|
||||
GET /me – Return the currently authenticated user's profile.
|
||||
GET /forgot-password – Redirect to Logto's forgot-password screen (unauthenticated).
|
||||
GET /change-password – Redirect to Logto Account Center password page (authenticated).
|
||||
GET /manage-mfa – Redirect to Logto Account Center MFA page (authenticated).
|
||||
GET /account-portal – Redirect to the Logto Account Center root.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.responses import RedirectResponse
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_db
|
||||
from app.core.logto import (
|
||||
SESSION_COOKIE,
|
||||
CookieStorage,
|
||||
create_session_token,
|
||||
decode_session_token,
|
||||
make_logto_client,
|
||||
sync_logto_user,
|
||||
)
|
||||
from app.models.user import User
|
||||
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
# ── Helpers ───────────────────────────────────────────────────────────────────
|
||||
|
||||
_SAFE_NEXT_PREFIXES = ("/",) # only allow relative redirects after login
|
||||
|
||||
|
||||
def _safe_next(next_url: Optional[str]) -> str:
|
||||
"""Validate and return a safe post-login redirect path."""
|
||||
if next_url and next_url.startswith("/") and not next_url.startswith("//"):
|
||||
return next_url
|
||||
return "/"
|
||||
|
||||
|
||||
def _logto_not_configured() -> HTTPException:
|
||||
return HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=(
|
||||
"Logto is not configured. "
|
||||
"Set LOGTO_ENDPOINT, LOGTO_APP_ID, and LOGTO_APP_SECRET "
|
||||
"in your environment."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _get_redirect_uri(request: Request) -> str:
|
||||
"""Build the callback redirect URI, preferring the configured override."""
|
||||
if settings.LOGTO_REDIRECT_URI:
|
||||
return settings.LOGTO_REDIRECT_URI
|
||||
base = str(request.base_url).rstrip("/")
|
||||
return f"{base}/api/v1/auth/callback"
|
||||
|
||||
|
||||
# ── Endpoints ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/sign-in")
|
||||
async def sign_in(
|
||||
request: Request,
|
||||
next: Optional[str] = None,
|
||||
) -> RedirectResponse:
|
||||
"""
|
||||
Initiate the Logto OIDC sign-in flow.
|
||||
|
||||
Stores the PKCE sign-in session in a short-lived cookie and redirects the
|
||||
browser to Logto's authorization endpoint. The optional ``next`` query
|
||||
parameter is persisted in a separate cookie and used to redirect the user
|
||||
to their original page after a successful login.
|
||||
"""
|
||||
if not settings.logto_configured:
|
||||
raise _logto_not_configured()
|
||||
|
||||
storage = CookieStorage(request)
|
||||
client = make_logto_client(storage)
|
||||
|
||||
sign_in_url: str = await client.signIn(redirectUri=_get_redirect_uri(request))
|
||||
|
||||
response = RedirectResponse(url=sign_in_url, status_code=302)
|
||||
storage.apply_to_response(response)
|
||||
|
||||
# Persist the post-login destination so the callback can redirect there.
|
||||
safe = _safe_next(next)
|
||||
if safe != "/":
|
||||
response.set_cookie(
|
||||
key="logto_next",
|
||||
value=safe,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
max_age=600, # 10 minutes – must survive the Logto redirect round-trip
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
@router.get("/callback")
|
||||
async def callback(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> RedirectResponse:
|
||||
"""
|
||||
Handle the Logto authorization-code callback.
|
||||
|
||||
Exchanges the code for tokens, validates the ID token, upserts the local
|
||||
user shadow record, issues the app-level session cookie, and clears the
|
||||
temporary Logto cookies.
|
||||
"""
|
||||
if not settings.logto_configured:
|
||||
raise _logto_not_configured()
|
||||
|
||||
storage = CookieStorage(request)
|
||||
client = make_logto_client(storage)
|
||||
|
||||
try:
|
||||
await client.handleSignInCallback(str(request.url))
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.warning("Logto callback error: %s", exc)
|
||||
return RedirectResponse(url="/login?error=callback_failed", status_code=302)
|
||||
|
||||
try:
|
||||
claims = client.getIdTokenClaims()
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.warning("Failed to extract ID-token claims: %s", exc)
|
||||
return RedirectResponse(url="/login?error=token_error", status_code=302)
|
||||
|
||||
user = sync_logto_user(claims, db)
|
||||
|
||||
# Where to go after login
|
||||
next_url = _safe_next(request.cookies.get("logto_next"))
|
||||
|
||||
response = RedirectResponse(url=next_url, status_code=302)
|
||||
|
||||
# Issue our own session cookie (independent of Logto from here on)
|
||||
session_token = create_session_token(user.id)
|
||||
response.set_cookie(
|
||||
key=SESSION_COOKIE,
|
||||
value=session_token,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
max_age=86_400, # 24 hours
|
||||
)
|
||||
|
||||
# Clean up all temporary Logto & next cookies
|
||||
storage.clear_all_logto_cookies(response)
|
||||
response.delete_cookie(key="logto_next", httponly=True, samesite="lax")
|
||||
|
||||
logger.info("User id=%d logged in via Logto.", user.id)
|
||||
return response
|
||||
|
||||
|
||||
@router.get("/sign-out")
|
||||
async def sign_out(request: Request) -> RedirectResponse:
|
||||
"""
|
||||
Sign the user out.
|
||||
|
||||
When ``AUTH_DISABLED=true`` there is nothing to sign out of; redirects to ``/``.
|
||||
|
||||
Otherwise clears the app session cookie and redirects to Logto's end-session
|
||||
endpoint (if available) so that the Logto session is terminated too.
|
||||
"""
|
||||
if settings.AUTH_DISABLED:
|
||||
return RedirectResponse(url="/", status_code=302)
|
||||
|
||||
post_logout_url = str(request.base_url).rstrip("/")
|
||||
|
||||
# Best-effort: obtain Logto's end-session URL from OIDC metadata.
|
||||
end_session_url: Optional[str] = None
|
||||
if settings.logto_configured:
|
||||
try:
|
||||
storage = CookieStorage(request)
|
||||
client = make_logto_client(storage)
|
||||
core = await client.getOidcCore()
|
||||
end_session_url = getattr(core.metadata, "end_session_endpoint", None)
|
||||
except Exception: # pylint: disable=broad-exception-caught
|
||||
pass
|
||||
|
||||
if end_session_url:
|
||||
redirect_to = f"{end_session_url}?post_logout_redirect_uri={post_logout_url}"
|
||||
else:
|
||||
redirect_to = "/login"
|
||||
|
||||
response = RedirectResponse(url=redirect_to, status_code=302)
|
||||
response.delete_cookie(key=SESSION_COOKIE, httponly=True, samesite="lax")
|
||||
return response
|
||||
|
||||
|
||||
@router.get("/change-password")
|
||||
async def change_password(request: Request) -> RedirectResponse:
|
||||
"""
|
||||
Redirect an authenticated user to the Logto Account Center password page.
|
||||
|
||||
Uses Logto's prebuilt Account Center flow at ``{LOGTO_ENDPOINT}/account/password``
|
||||
so the user can change their existing password directly. A ``redirect``
|
||||
query parameter is appended so that Logto returns the user to the Profile &
|
||||
Security page after a successful update.
|
||||
"""
|
||||
if not settings.logto_configured:
|
||||
raise _logto_not_configured()
|
||||
|
||||
base = str(request.base_url).rstrip("/")
|
||||
password_url = f"{settings.LOGTO_ENDPOINT.rstrip('/')}/account/password?redirect={base}/profile"
|
||||
return RedirectResponse(url=password_url, status_code=302)
|
||||
|
||||
|
||||
@router.get("/forgot-password")
|
||||
async def forgot_password(request: Request) -> RedirectResponse:
|
||||
"""
|
||||
Redirect the user to Logto's forgot-password screen.
|
||||
|
||||
Builds a standard Logto authorization URL and appends the
|
||||
``first_screen=forgot_password`` parameter so that Logto shows the
|
||||
password-reset form immediately instead of the normal sign-in form.
|
||||
After the user resets their password they are returned via the normal
|
||||
callback flow and land on the app dashboard.
|
||||
|
||||
This endpoint is kept for unauthenticated / "I forgot my password" use
|
||||
cases. Authenticated users should use ``/change-password`` instead.
|
||||
"""
|
||||
if not settings.logto_configured:
|
||||
raise _logto_not_configured()
|
||||
|
||||
storage = CookieStorage(request)
|
||||
client = make_logto_client(storage)
|
||||
|
||||
sign_in_url: str = await client.signIn(redirectUri=_get_redirect_uri(request))
|
||||
|
||||
# Append the Logto-specific first_screen parameter so the password-reset
|
||||
# form is shown directly. The sign-in URL normally already contains a "?"
|
||||
# but we defensively detect the right separator in case the structure varies.
|
||||
separator = "&" if "?" in sign_in_url else "?"
|
||||
forgot_url = f"{sign_in_url}{separator}first_screen=forgot_password"
|
||||
|
||||
response = RedirectResponse(url=forgot_url, status_code=302)
|
||||
storage.apply_to_response(response)
|
||||
return response
|
||||
|
||||
|
||||
@router.get("/manage-mfa")
|
||||
async def manage_mfa(request: Request) -> RedirectResponse:
|
||||
"""
|
||||
Redirect an authenticated user to the Logto Account Center MFA page.
|
||||
|
||||
Uses Logto's prebuilt Account Center flow at
|
||||
``{LOGTO_ENDPOINT}/account/authenticator-app`` so the user can enable,
|
||||
configure, or remove TOTP authenticator-app MFA directly. A ``redirect``
|
||||
query parameter is appended so that Logto returns the user to the Profile &
|
||||
Security page after a successful update.
|
||||
"""
|
||||
if not settings.logto_configured:
|
||||
raise _logto_not_configured()
|
||||
|
||||
base = str(request.base_url).rstrip("/")
|
||||
mfa_url = (
|
||||
f"{settings.LOGTO_ENDPOINT.rstrip('/')}/account/authenticator-app"
|
||||
f"?redirect={base}/profile"
|
||||
)
|
||||
return RedirectResponse(url=mfa_url, status_code=302)
|
||||
|
||||
|
||||
@router.get("/account-portal")
|
||||
async def account_portal(request: Request) -> RedirectResponse:
|
||||
"""
|
||||
Redirect an authenticated user to the Logto Account Center.
|
||||
|
||||
The Logto account portal (``{LOGTO_ENDPOINT}/account``) lets users manage
|
||||
their profile, linked identities, and multi-factor authentication settings
|
||||
without leaving the Logto-hosted UI. A ``redirect`` query parameter is
|
||||
appended so that Logto returns the user to the Profile & Security page
|
||||
after a successful update.
|
||||
"""
|
||||
if not settings.logto_configured:
|
||||
raise _logto_not_configured()
|
||||
|
||||
base = str(request.base_url).rstrip("/")
|
||||
portal_url = f"{settings.LOGTO_ENDPOINT.rstrip('/')}/account?redirect={base}/profile"
|
||||
return RedirectResponse(url=portal_url, status_code=302)
|
||||
|
||||
|
||||
@router.get("/me", response_model=None)
|
||||
async def get_current_user(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Return the profile of the currently authenticated user.
|
||||
|
||||
When ``AUTH_DISABLED=true`` a synthetic anonymous-admin profile is returned
|
||||
so that UI components (e.g. the navbar user menu) work without a real session.
|
||||
|
||||
Otherwise reads the ``dmarq_session`` cookie (issued at callback time) and
|
||||
looks up the corresponding local ``User`` record.
|
||||
"""
|
||||
# Auth-disabled: return a synthetic profile so the UI renders correctly.
|
||||
if settings.AUTH_DISABLED:
|
||||
return {
|
||||
"id": 0,
|
||||
"email": "admin@localhost",
|
||||
"full_name": "Local Admin",
|
||||
"username": "admin",
|
||||
"picture": None,
|
||||
"is_superuser": True,
|
||||
"logto_id": None,
|
||||
"auth_disabled": True,
|
||||
}
|
||||
|
||||
token = request.cookies.get(SESSION_COOKIE)
|
||||
if not token:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Not authenticated",
|
||||
)
|
||||
|
||||
user_id = decode_session_token(token)
|
||||
if user_id is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Invalid or expired session",
|
||||
)
|
||||
|
||||
user: Optional[User] = (
|
||||
db.query(User).filter(User.id == user_id, User.is_active == True).first() # noqa: E712
|
||||
)
|
||||
if user is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="User not found or inactive",
|
||||
)
|
||||
|
||||
return {
|
||||
"id": user.id,
|
||||
"email": user.email,
|
||||
"full_name": user.full_name,
|
||||
"username": user.username,
|
||||
"picture": user.picture,
|
||||
"is_superuser": user.is_superuser,
|
||||
"logto_id": user.logto_id,
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,280 @@
|
||||
import logging
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.models.domain import Domain
|
||||
from app.models.report import ForensicReport
|
||||
from app.services.forensic_analysis import analyze_forensic_report, summarize_forensic_samples
|
||||
from app.services.forensic_parser import ForensicParser, MAX_FORENSIC_REPORT_SIZE
|
||||
from app.services.forensic_persistence import (
|
||||
forensic_report_exists,
|
||||
forensic_report_to_dict,
|
||||
save_forensic_report,
|
||||
)
|
||||
from app.services.forensic_redaction import get_forensic_redaction_policy
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class ForensicSampleAnalysisResponse(BaseModel):
|
||||
id: int
|
||||
report_id: str
|
||||
domain: Optional[str] = None
|
||||
source_ip: Optional[str] = None
|
||||
auth_failure: str
|
||||
delivery_result: Optional[str] = None
|
||||
priority: str
|
||||
diagnosis: str
|
||||
recommendations: List[str] = Field(default_factory=list)
|
||||
signals: List[str] = Field(default_factory=list)
|
||||
authentication_results: Dict[str, str] = Field(default_factory=dict)
|
||||
dkim_domain: Optional[str] = None
|
||||
mail_from_domain: Optional[str] = None
|
||||
privacy_note: str
|
||||
|
||||
|
||||
class ForensicAnalysisGroupResponse(BaseModel):
|
||||
key: str
|
||||
domain: str
|
||||
source_ip: str
|
||||
auth_failure: str
|
||||
delivery_result: str
|
||||
count: int
|
||||
priority: str
|
||||
latest_arrival: Optional[str] = None
|
||||
diagnosis: str
|
||||
recommendations: List[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class ForensicAnalysisResponse(BaseModel):
|
||||
total: int
|
||||
priority_counts: Dict[str, int] = Field(default_factory=dict)
|
||||
failure_counts: Dict[str, int] = Field(default_factory=dict)
|
||||
result_counts: Dict[str, int] = Field(default_factory=dict)
|
||||
groups: List[ForensicAnalysisGroupResponse] = Field(default_factory=list)
|
||||
samples: List[ForensicSampleAnalysisResponse] = Field(default_factory=list)
|
||||
|
||||
|
||||
class ForensicReportResponse(BaseModel):
|
||||
id: int
|
||||
report_id: str
|
||||
domain: Optional[str] = None
|
||||
reported_domain: Optional[str] = None
|
||||
source_email: Optional[str] = None
|
||||
feedback_type: Optional[str] = None
|
||||
user_agent: Optional[str] = None
|
||||
version: Optional[str] = None
|
||||
source_ip: Optional[str] = None
|
||||
auth_failure: Optional[str] = None
|
||||
delivery_result: Optional[str] = None
|
||||
arrival_date: Optional[str] = None
|
||||
authentication_results: Optional[str] = None
|
||||
original_mail_from: Optional[str] = None
|
||||
original_from: Optional[str] = None
|
||||
original_to: Optional[str] = None
|
||||
original_subject: Optional[str] = None
|
||||
original_message_id: Optional[str] = None
|
||||
original_date: Optional[str] = None
|
||||
feedback_headers: Dict[str, Any] = Field(default_factory=dict)
|
||||
processed_at: Optional[str] = None
|
||||
analysis: Optional[ForensicSampleAnalysisResponse] = None
|
||||
|
||||
|
||||
class ForensicListResponse(BaseModel):
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
total_pages: int
|
||||
reports: List[ForensicReportResponse]
|
||||
|
||||
|
||||
class ForensicUploadResponse(BaseModel):
|
||||
success: bool
|
||||
report_id: str
|
||||
domain: Optional[str] = None
|
||||
message: str
|
||||
duplicate: bool = False
|
||||
|
||||
|
||||
def _validate_upload(file: UploadFile, content: bytes) -> None:
|
||||
filename = file.filename or ""
|
||||
if not filename:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Filename is required")
|
||||
if len(content) == 0:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="File is empty")
|
||||
if len(content) > MAX_FORENSIC_REPORT_SIZE:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, detail="File too large"
|
||||
)
|
||||
if not filename.lower().endswith((".eml", ".txt")):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Invalid file type. Upload a forensic report email as .eml or .txt.",
|
||||
)
|
||||
|
||||
|
||||
def _filtered_forensic_query(
|
||||
db: Session,
|
||||
*,
|
||||
domain: Optional[str] = None,
|
||||
source_ip: Optional[str] = None,
|
||||
auth_failure: Optional[str] = None,
|
||||
delivery_result: Optional[str] = None,
|
||||
):
|
||||
query = db.query(ForensicReport).options(selectinload(ForensicReport.domain))
|
||||
if domain:
|
||||
normalized = domain.lower()
|
||||
query = query.outerjoin(Domain).filter(
|
||||
(Domain.name == normalized) | (ForensicReport.reported_domain == normalized)
|
||||
)
|
||||
if source_ip:
|
||||
query = query.filter(ForensicReport.source_ip == source_ip.strip())
|
||||
if auth_failure:
|
||||
query = query.filter(ForensicReport.auth_failure == auth_failure.strip().lower())
|
||||
if delivery_result:
|
||||
query = query.filter(ForensicReport.delivery_result == delivery_result.strip().lower())
|
||||
return query
|
||||
|
||||
|
||||
def _response_for_row(row: ForensicReport, redaction_policy) -> ForensicReportResponse:
|
||||
data = forensic_report_to_dict(row, redaction_policy=redaction_policy)
|
||||
data["analysis"] = analyze_forensic_report(row)
|
||||
return ForensicReportResponse(**data)
|
||||
|
||||
|
||||
@router.post("/upload", response_model=ForensicUploadResponse)
|
||||
async def upload_forensic_report(
|
||||
file: UploadFile = File(...),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""Upload and store a DMARC forensic/failure report email."""
|
||||
try:
|
||||
content = await file.read()
|
||||
_validate_upload(file, content)
|
||||
redaction_policy = get_forensic_redaction_policy(db)
|
||||
parsed = ForensicParser.parse_bytes(content, redaction_policy=redaction_policy)
|
||||
if forensic_report_exists(db, parsed["report_id"]):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail="Forensic report has already been uploaded.",
|
||||
)
|
||||
|
||||
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.refresh(row)
|
||||
return ForensicUploadResponse(
|
||||
success=True,
|
||||
report_id=row.report_id,
|
||||
domain=row.reported_domain,
|
||||
message="Forensic report processed successfully.",
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Invalid forensic report format.",
|
||||
) from exc
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Unexpected forensic upload failure for %s: %s", file.filename, exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Error processing forensic report.",
|
||||
) from exc
|
||||
|
||||
|
||||
@router.get("", response_model=ForensicListResponse)
|
||||
async def list_forensic_reports(
|
||||
domain: Optional[str] = Query(default=None),
|
||||
source_ip: Optional[str] = Query(default=None),
|
||||
auth_failure: Optional[str] = Query(default=None),
|
||||
delivery_result: Optional[str] = Query(default=None),
|
||||
page: int = Query(default=1, ge=1),
|
||||
page_size: int = Query(default=50, ge=1, le=200),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""List stored forensic reports, newest first."""
|
||||
query = _filtered_forensic_query(
|
||||
db,
|
||||
domain=domain,
|
||||
source_ip=source_ip,
|
||||
auth_failure=auth_failure,
|
||||
delivery_result=delivery_result,
|
||||
)
|
||||
total = query.count()
|
||||
rows = (
|
||||
query.order_by(ForensicReport.arrival_date.desc().nullslast(), ForensicReport.id.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
.all()
|
||||
)
|
||||
total_pages = (total + page_size - 1) // page_size if total else 0
|
||||
redaction_policy = get_forensic_redaction_policy(db)
|
||||
return ForensicListResponse(
|
||||
total=total,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
total_pages=total_pages,
|
||||
reports=[_response_for_row(row, redaction_policy) for row in rows],
|
||||
)
|
||||
|
||||
|
||||
@router.get("/analysis", response_model=ForensicAnalysisResponse)
|
||||
async def analyze_forensic_reports(
|
||||
domain: Optional[str] = Query(default=None),
|
||||
source_ip: Optional[str] = Query(default=None),
|
||||
auth_failure: Optional[str] = Query(default=None),
|
||||
delivery_result: Optional[str] = Query(default=None),
|
||||
page_size: int = Query(default=200, ge=1, le=500),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""Summarize stored forensic samples into operator investigation groups."""
|
||||
query = _filtered_forensic_query(
|
||||
db,
|
||||
domain=domain,
|
||||
source_ip=source_ip,
|
||||
auth_failure=auth_failure,
|
||||
delivery_result=delivery_result,
|
||||
)
|
||||
rows = (
|
||||
query.order_by(ForensicReport.arrival_date.desc().nullslast(), ForensicReport.id.desc())
|
||||
.limit(page_size)
|
||||
.all()
|
||||
)
|
||||
return ForensicAnalysisResponse(**summarize_forensic_samples(rows))
|
||||
|
||||
|
||||
@router.get("/{report_id}", response_model=ForensicReportResponse)
|
||||
async def get_forensic_report(
|
||||
report_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""Return one stored forensic report by numeric ID."""
|
||||
row = (
|
||||
db.query(ForensicReport)
|
||||
.options(selectinload(ForensicReport.domain))
|
||||
.filter(ForensicReport.id == report_id)
|
||||
.first()
|
||||
)
|
||||
if row is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail="Forensic report not found"
|
||||
)
|
||||
redaction_policy = get_forensic_redaction_policy(db)
|
||||
return _response_for_row(row, redaction_policy)
|
||||
@@ -1,5 +1,14 @@
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy import func, text
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.api_v1.endpoints.setup import setup_status
|
||||
from fastapi import APIRouter
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.models.mail_source import MailSource
|
||||
from app.models.mail_source_import import MailSourceImport
|
||||
from app.models.report import DMARCReport
|
||||
from app.services.runtime_status import get_scheduler_status
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -16,3 +25,94 @@ async def health_check():
|
||||
"service": "dmarq",
|
||||
"is_setup_complete": setup_status["is_setup_complete"],
|
||||
}
|
||||
|
||||
|
||||
def _iso(value):
|
||||
return value.isoformat() if value else None
|
||||
|
||||
|
||||
@router.get("/health/operations", status_code=200)
|
||||
async def operations_health(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""Return operational health details for the web health page."""
|
||||
database = {"ok": True, "detail": "Connected"}
|
||||
try:
|
||||
db.execute(text("SELECT 1"))
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
database = {"ok": False, "detail": str(exc)}
|
||||
|
||||
enabled_sources = 0
|
||||
total_sources = 0
|
||||
report_count = 0
|
||||
latest_report = None
|
||||
latest_import = None
|
||||
latest_successful_import = None
|
||||
if database["ok"]:
|
||||
enabled_sources = (
|
||||
db.query(func.count(MailSource.id))
|
||||
.filter(MailSource.enabled == True) # noqa: E712
|
||||
.scalar()
|
||||
)
|
||||
total_sources = db.query(func.count(MailSource.id)).scalar()
|
||||
report_count = db.query(func.count(DMARCReport.id)).scalar()
|
||||
latest_report = db.query(func.max(DMARCReport.processed_at)).scalar()
|
||||
latest_import = (
|
||||
db.query(MailSourceImport)
|
||||
.order_by(MailSourceImport.finished_at.desc(), MailSourceImport.id.desc())
|
||||
.first()
|
||||
)
|
||||
latest_successful_import = (
|
||||
db.query(MailSourceImport)
|
||||
.filter(MailSourceImport.status.in_(["success", "warning"]))
|
||||
.order_by(MailSourceImport.finished_at.desc(), MailSourceImport.id.desc())
|
||||
.first()
|
||||
)
|
||||
|
||||
scheduler = get_scheduler_status()
|
||||
status = "ok"
|
||||
checks = []
|
||||
if not database["ok"]:
|
||||
status = "degraded"
|
||||
checks.append("Database connectivity failed.")
|
||||
if total_sources and enabled_sources == 0:
|
||||
status = "degraded"
|
||||
checks.append("All mail sources are disabled.")
|
||||
if scheduler.get("last_error"):
|
||||
status = "degraded"
|
||||
checks.append("The scheduler reported a recent error.")
|
||||
|
||||
return {
|
||||
"status": status,
|
||||
"service": "dmarq",
|
||||
"database": database,
|
||||
"scheduler": {
|
||||
**scheduler,
|
||||
"enabled_sources": int(enabled_sources or 0),
|
||||
"total_sources": int(total_sources or 0),
|
||||
},
|
||||
"imports": {
|
||||
"latest": {
|
||||
"status": latest_import.status,
|
||||
"trigger": latest_import.trigger,
|
||||
"reports_found": latest_import.reports_found,
|
||||
"finished_at": _iso(latest_import.finished_at),
|
||||
}
|
||||
if latest_import
|
||||
else None,
|
||||
"latest_successful": {
|
||||
"status": latest_successful_import.status,
|
||||
"trigger": latest_successful_import.trigger,
|
||||
"reports_found": latest_successful_import.reports_found,
|
||||
"finished_at": _iso(latest_successful_import.finished_at),
|
||||
}
|
||||
if latest_successful_import
|
||||
else None,
|
||||
},
|
||||
"reports": {
|
||||
"count": int(report_count or 0),
|
||||
"latest_processed_at": _iso(latest_report),
|
||||
},
|
||||
"checks": checks,
|
||||
}
|
||||
|
||||
@@ -2,10 +2,13 @@ import logging
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from app.core.security import require_admin_auth
|
||||
from app.services.imap_client import IMAPClient
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException
|
||||
from pydantic import BaseModel
|
||||
from starlette.concurrency import run_in_threadpool
|
||||
|
||||
from app.core.database import SessionLocal
|
||||
from app.core.security import require_admin_auth
|
||||
from app.services.imap_client import IMAPClient
|
||||
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -21,6 +24,26 @@ class IMAPTestRequest(BaseModel):
|
||||
ssl: bool = True
|
||||
|
||||
|
||||
def _fetch_imap_reports_sync(days: int, delete_emails: Optional[bool]) -> Dict[str, Any]:
|
||||
"""Fetch IMAP reports with a standalone DB session."""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
imap_client = IMAPClient(delete_emails=delete_emails, db=db)
|
||||
results = imap_client.fetch_reports(days=days)
|
||||
db.commit()
|
||||
return results
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _fetch_imap_reports_background(days: int, delete_emails: Optional[bool]) -> None:
|
||||
"""Fetch IMAP reports from a FastAPI background task."""
|
||||
_fetch_imap_reports_sync(days, delete_emails)
|
||||
|
||||
|
||||
@router.post("/test-connection")
|
||||
async def test_imap_connection(
|
||||
request: IMAPTestRequest,
|
||||
@@ -57,7 +80,7 @@ async def fetch_imap_reports(
|
||||
background_tasks: BackgroundTasks,
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
days: int = 7,
|
||||
delete_emails: bool = False,
|
||||
delete_emails: Optional[bool] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Fetch DMARC reports from the configured IMAP mailbox
|
||||
@@ -68,11 +91,9 @@ 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)
|
||||
|
||||
# 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",
|
||||
@@ -81,12 +102,14 @@ async def fetch_imap_reports(
|
||||
|
||||
# Otherwise run immediately
|
||||
try:
|
||||
results = imap_client.fetch_reports(days=days)
|
||||
results = await run_in_threadpool(_fetch_imap_reports_sync, days, delete_emails)
|
||||
|
||||
return {
|
||||
"success": results["success"],
|
||||
"processed_emails": results["processed"],
|
||||
"reports_found": results["reports_found"],
|
||||
"forensic_reports_found": int(results.get("forensic_reports_found", 0)),
|
||||
"duplicate_forensic_reports": int(results.get("duplicate_forensic_reports", 0)),
|
||||
"new_domains": results["new_domains"],
|
||||
"errors": results["errors"] if "errors" in results and results["errors"] else None,
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
"""Integration template endpoints."""
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from app.core.security import require_admin_auth
|
||||
from app.services.siem_templates import get_siem_templates
|
||||
from app.services.ticketing_chatops_templates import get_ticketing_chatops_templates
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("/siem/templates")
|
||||
async def siem_templates(_auth: dict = Depends(require_admin_auth)):
|
||||
"""Return versioned schemas and examples for SIEM ingestion."""
|
||||
return get_siem_templates()
|
||||
|
||||
|
||||
@router.get("/ticketing-chatops/templates")
|
||||
async def ticketing_chatops_templates(_auth: dict = Depends(require_admin_auth)):
|
||||
"""Return ticketing and chatops workflow templates."""
|
||||
return get_ticketing_chatops_templates()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,163 @@
|
||||
"""Read-only MCP-style JSON-RPC endpoint for agent integrations."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_api_token_scope
|
||||
from app.services.ai_assistance import (
|
||||
build_action_proposals,
|
||||
build_evidence_summary,
|
||||
get_assistance_config,
|
||||
)
|
||||
from app.services.api_tokens import MCP_READ_SCOPE
|
||||
from app.services.report_persistence import hydrate_report_store_from_db
|
||||
from app.services.report_store import ReportStore
|
||||
from app.services.workspace_audit import record_workspace_audit_log
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class MCPRequest(BaseModel):
|
||||
"""Minimal JSON-RPC request for MCP HTTP integrations."""
|
||||
|
||||
jsonrpc: str = "2.0"
|
||||
id: Optional[Any] = None
|
||||
method: str
|
||||
params: Dict[str, Any] = {}
|
||||
|
||||
|
||||
READ_ONLY_TOOLS = [
|
||||
{
|
||||
"name": "list_domains",
|
||||
"description": "List monitored domains and aggregate counts.",
|
||||
"inputSchema": {"type": "object", "properties": {}},
|
||||
"readOnlyHint": True,
|
||||
},
|
||||
{
|
||||
"name": "domain_summary",
|
||||
"description": "Return an evidence-first summary for one domain.",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {"domain": {"type": "string"}},
|
||||
"required": ["domain"],
|
||||
},
|
||||
"readOnlyHint": True,
|
||||
},
|
||||
{
|
||||
"name": "action_proposals",
|
||||
"description": "Return reviewable remediation proposals without applying changes.",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {"domain": {"type": "string"}},
|
||||
"required": ["domain"],
|
||||
},
|
||||
"readOnlyHint": True,
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _jsonrpc_response(request_id: Any, result: Any = None, error: Optional[Dict[str, Any]] = None):
|
||||
payload = {"jsonrpc": "2.0", "id": request_id}
|
||||
if error is not None:
|
||||
payload["error"] = error
|
||||
else:
|
||||
payload["result"] = result
|
||||
return payload
|
||||
|
||||
|
||||
def _require_mcp_enabled(db: Session) -> None:
|
||||
if not get_assistance_config(db).mcp_enabled:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="MCP access is disabled. Enable mcp.enabled before using this endpoint.",
|
||||
)
|
||||
|
||||
|
||||
def _list_domains(db: Session) -> Dict[str, Any]:
|
||||
store = ReportStore.get_instance()
|
||||
hydrate_report_store_from_db(db, store)
|
||||
summaries = store.get_all_domain_summaries()
|
||||
return {
|
||||
"domains": [
|
||||
{
|
||||
"domain": domain,
|
||||
"total_messages": int(summary.get("total_count", 0) or 0),
|
||||
"failed_messages": int(summary.get("failed_count", 0) or 0),
|
||||
"compliance_rate": float(summary.get("compliance_rate", 0.0) or 0.0),
|
||||
"reports_processed": int(summary.get("reports_processed", 0) or 0),
|
||||
}
|
||||
for domain, summary in sorted(summaries.items())
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@router.post("")
|
||||
async def mcp_jsonrpc(
|
||||
payload: MCPRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_api_token_scope(MCP_READ_SCOPE)),
|
||||
) -> Dict[str, Any]:
|
||||
"""Handle a small read-only MCP tool surface over JSON-RPC."""
|
||||
_require_mcp_enabled(db)
|
||||
if payload.method == "initialize":
|
||||
return _jsonrpc_response(
|
||||
payload.id,
|
||||
{
|
||||
"protocolVersion": "2024-11-05",
|
||||
"serverInfo": {"name": "dmarq", "version": "1"},
|
||||
"capabilities": {"tools": {}},
|
||||
},
|
||||
)
|
||||
if payload.method == "tools/list":
|
||||
return _jsonrpc_response(payload.id, {"tools": READ_ONLY_TOOLS})
|
||||
if payload.method != "tools/call":
|
||||
return _jsonrpc_response(
|
||||
payload.id,
|
||||
error={"code": -32601, "message": f"Unsupported method: {payload.method}"},
|
||||
)
|
||||
|
||||
name = payload.params.get("name")
|
||||
arguments = payload.params.get("arguments") or {}
|
||||
try:
|
||||
if name == "list_domains":
|
||||
result = _list_domains(db)
|
||||
elif name == "domain_summary":
|
||||
result = build_evidence_summary(db, str(arguments.get("domain", "")))
|
||||
elif name == "action_proposals":
|
||||
result = build_action_proposals(db, str(arguments.get("domain", "")))
|
||||
else:
|
||||
return _jsonrpc_response(
|
||||
payload.id,
|
||||
error={"code": -32602, "message": f"Unsupported tool: {name}"},
|
||||
)
|
||||
except ValueError as exc:
|
||||
return _jsonrpc_response(payload.id, error={"code": -32004, "message": str(exc)})
|
||||
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db, commit=False)
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="mcp.tool_called",
|
||||
entity_type="mcp_tool",
|
||||
entity_id=name,
|
||||
entity_name=name,
|
||||
details={"tool": name, "read_only": True},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
)
|
||||
db.commit()
|
||||
return _jsonrpc_response(
|
||||
payload.id,
|
||||
{
|
||||
"content": [{"type": "json", "json": result}],
|
||||
"isError": False,
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,126 @@
|
||||
"""Workspace onboarding template endpoints."""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.services.workspace_access import (
|
||||
PERMISSION_WORKSPACE_ADMIN,
|
||||
require_workspace_permission,
|
||||
)
|
||||
from app.services.workspace_onboarding import (
|
||||
apply_onboarding_plan,
|
||||
build_onboarding_plan,
|
||||
list_onboarding_templates,
|
||||
public_onboarding_plan,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class OnboardingWorkspace(BaseModel):
|
||||
"""Workspace target for an onboarding plan."""
|
||||
|
||||
slug: Optional[str] = None
|
||||
name: str
|
||||
description: Optional[str] = None
|
||||
|
||||
|
||||
class OnboardingPlanRequest(BaseModel):
|
||||
"""Request body for rendering or applying an onboarding template."""
|
||||
|
||||
template_id: str
|
||||
workspace: OnboardingWorkspace
|
||||
variables: Dict[str, Any] = Field(default_factory=dict)
|
||||
domains: Optional[List[Dict[str, Any]]] = None
|
||||
mail_sources: Optional[List[Dict[str, Any]]] = None
|
||||
notification_defaults: Optional[Dict[str, Any]] = None
|
||||
overwrite_existing: bool = False
|
||||
|
||||
|
||||
class OnboardingTemplatesResponse(BaseModel):
|
||||
"""Available workspace onboarding templates."""
|
||||
|
||||
templates: List[Dict[str, Any]]
|
||||
|
||||
|
||||
class OnboardingPlanResponse(BaseModel):
|
||||
"""Rendered onboarding plan response."""
|
||||
|
||||
plan: Dict[str, Any]
|
||||
|
||||
|
||||
class OnboardingApplyResponse(BaseModel):
|
||||
"""Applied onboarding plan response."""
|
||||
|
||||
result: Dict[str, Any]
|
||||
|
||||
|
||||
def _build_plan_or_422(payload: OnboardingPlanRequest) -> Dict[str, Any]:
|
||||
try:
|
||||
plan = build_onboarding_plan(
|
||||
template_id=payload.template_id,
|
||||
workspace=payload.workspace.model_dump(),
|
||||
variables=payload.variables,
|
||||
domains=payload.domains,
|
||||
mail_sources=payload.mail_sources,
|
||||
notification_defaults=payload.notification_defaults,
|
||||
overwrite_existing=payload.overwrite_existing,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
|
||||
if plan["errors"]:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=plan["errors"],
|
||||
)
|
||||
return plan
|
||||
|
||||
|
||||
@router.get("/templates", response_model=OnboardingTemplatesResponse)
|
||||
async def get_onboarding_templates(
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> OnboardingTemplatesResponse:
|
||||
"""Return versioned workspace onboarding templates."""
|
||||
require_workspace_permission(_auth, PERMISSION_WORKSPACE_ADMIN)
|
||||
return {"templates": list_onboarding_templates()}
|
||||
|
||||
|
||||
@router.post("/preview", response_model=OnboardingPlanResponse)
|
||||
async def preview_onboarding_plan(
|
||||
payload: OnboardingPlanRequest,
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> OnboardingPlanResponse:
|
||||
"""Render an onboarding template without changing the database."""
|
||||
require_workspace_permission(_auth, PERMISSION_WORKSPACE_ADMIN)
|
||||
return {"plan": public_onboarding_plan(_build_plan_or_422(payload))}
|
||||
|
||||
|
||||
@router.post("/apply", response_model=OnboardingApplyResponse)
|
||||
async def apply_workspace_onboarding(
|
||||
payload: OnboardingPlanRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> OnboardingApplyResponse:
|
||||
"""Apply a workspace onboarding template."""
|
||||
require_workspace_permission(_auth, PERMISSION_WORKSPACE_ADMIN)
|
||||
plan = _build_plan_or_422(payload)
|
||||
try:
|
||||
result = apply_onboarding_plan(db, plan=plan, auth_context=_auth, request=request)
|
||||
except IntegrityError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail="Onboarding plan conflicts with existing data",
|
||||
) from exc
|
||||
if not result.get("applied"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=result,
|
||||
)
|
||||
return {"result": result}
|
||||
@@ -0,0 +1,111 @@
|
||||
"""MSP operator endpoints."""
|
||||
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.models.workspace import Workspace
|
||||
from app.services.workspace_access import (
|
||||
PERMISSION_AUDIT_READ,
|
||||
PERMISSION_WORKSPACE_ADMIN,
|
||||
require_workspace_permission,
|
||||
)
|
||||
from app.services.workspace_audit import record_workspace_audit_log
|
||||
from app.services.workspace_operator import (
|
||||
list_workspace_operator_summaries,
|
||||
retention_to_dict,
|
||||
workspace_operator_summary,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class OperatorWorkspacesResponse(BaseModel):
|
||||
"""Cross-workspace operator summaries."""
|
||||
|
||||
workspaces: List[Dict[str, Any]]
|
||||
|
||||
|
||||
class WorkspaceRetentionUpdate(BaseModel):
|
||||
"""Workspace retention controls."""
|
||||
|
||||
aggregate_reports_days: int = Field(..., ge=1, le=3650)
|
||||
forensic_reports_days: int = Field(..., ge=1, le=3650)
|
||||
tls_reports_days: int = Field(..., ge=1, le=3650)
|
||||
|
||||
|
||||
class WorkspaceRetentionResponse(BaseModel):
|
||||
"""Updated workspace retention response."""
|
||||
|
||||
workspace: Dict[str, Any]
|
||||
retention: Dict[str, int]
|
||||
|
||||
|
||||
def _workspace_or_404(db: Session, workspace_id: int) -> Workspace:
|
||||
workspace = db.query(Workspace).filter(Workspace.id == workspace_id).first()
|
||||
if workspace is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Workspace {workspace_id} not found",
|
||||
)
|
||||
return workspace
|
||||
|
||||
|
||||
@router.get("/workspaces", response_model=OperatorWorkspacesResponse)
|
||||
async def list_operator_workspaces(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> OperatorWorkspacesResponse:
|
||||
"""Return safe cross-workspace health, drift, import, and retention summaries."""
|
||||
require_workspace_permission(_auth, PERMISSION_AUDIT_READ)
|
||||
return {"workspaces": list_workspace_operator_summaries(db)}
|
||||
|
||||
|
||||
@router.get("/workspaces/{workspace_id}", response_model=Dict[str, Any])
|
||||
async def get_operator_workspace(
|
||||
workspace_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Return one workspace operator summary."""
|
||||
require_workspace_permission(_auth, PERMISSION_AUDIT_READ)
|
||||
return workspace_operator_summary(db, _workspace_or_404(db, workspace_id))
|
||||
|
||||
|
||||
@router.put("/workspaces/{workspace_id}/retention", response_model=WorkspaceRetentionResponse)
|
||||
async def update_workspace_retention(
|
||||
workspace_id: int,
|
||||
payload: WorkspaceRetentionUpdate,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> WorkspaceRetentionResponse:
|
||||
"""Update workspace retention controls and audit the change."""
|
||||
require_workspace_permission(_auth, PERMISSION_WORKSPACE_ADMIN)
|
||||
workspace = _workspace_or_404(db, workspace_id)
|
||||
old_retention = retention_to_dict(workspace)
|
||||
workspace.report_retention_days = payload.aggregate_reports_days
|
||||
workspace.forensic_retention_days = payload.forensic_reports_days
|
||||
workspace.tls_report_retention_days = payload.tls_reports_days
|
||||
new_retention = retention_to_dict(workspace)
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="workspace.retention_updated",
|
||||
entity_type="workspace",
|
||||
entity_id=workspace.id,
|
||||
entity_name=workspace.slug,
|
||||
details={"old": old_retention, "new": new_retention},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
)
|
||||
db.commit()
|
||||
db.refresh(workspace)
|
||||
return {
|
||||
"workspace": {"id": workspace.id, "slug": workspace.slug, "name": workspace.name},
|
||||
"retention": retention_to_dict(workspace),
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
"""Stable read-only public API endpoints."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Path, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.api_v1.endpoints import domains, tls_reports
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_api_token_scope
|
||||
from app.services.api_tokens import READ_POSTURE_SCOPE, READ_REPORTS_SCOPE, READ_TLS_SCOPE
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("/domains", response_model=domains.DomainSummaryResponse)
|
||||
async def public_domain_summary(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_api_token_scope(READ_REPORTS_SCOPE)),
|
||||
):
|
||||
"""List monitored domains with report and DNS posture summary fields."""
|
||||
return await domains.get_domains_summary(db=db)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/domains/{domain_id}/posture",
|
||||
response_model=domains.PostureDashboardResponse,
|
||||
)
|
||||
async def public_domain_posture(
|
||||
domain_id: str = Path(..., title="The domain ID or name"),
|
||||
refresh: bool = Query(False, title="Refresh cached DNS posture"),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_api_token_scope(READ_POSTURE_SCOPE)),
|
||||
):
|
||||
"""Return the stable evidence-first posture payload for one domain."""
|
||||
return await domains.get_domain_posture_dashboard(
|
||||
domain_id=domain_id,
|
||||
refresh=refresh,
|
||||
db=db,
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/domains/{domain_id}/reports",
|
||||
response_model=domains.DomainReportsResponse,
|
||||
)
|
||||
async def public_domain_reports(
|
||||
domain_id: str = Path(..., title="The domain ID or name"),
|
||||
limit: int = Query(10, ge=1, le=200),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_api_token_scope(READ_REPORTS_SCOPE)),
|
||||
):
|
||||
"""Return recent DMARC aggregate report summaries for one domain."""
|
||||
return await domains.get_domain_reports(domain_id=domain_id, limit=limit, db=db)
|
||||
|
||||
|
||||
@router.get("/tls-reports/summary", response_model=tls_reports.TLSSummaryResponse)
|
||||
async def public_tls_report_summary(
|
||||
domain: Optional[str] = Query(default=None),
|
||||
days: int = Query(default=30, ge=1, le=365),
|
||||
limit: int = Query(default=10, ge=1, le=50),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_api_token_scope(READ_TLS_SCOPE)),
|
||||
):
|
||||
"""Return aggregate SMTP TLS reporting posture trends."""
|
||||
return await tls_reports.tls_report_summary(
|
||||
domain=domain,
|
||||
days=days,
|
||||
limit=limit,
|
||||
db=db,
|
||||
_auth=_auth,
|
||||
)
|
||||
@@ -1,11 +1,20 @@
|
||||
import logging
|
||||
from typing import List
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
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
|
||||
from fastapi import APIRouter, File, HTTPException, UploadFile, status
|
||||
from pydantic import BaseModel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -47,16 +56,16 @@ def _validate_mime_type(file_content: bytes) -> None:
|
||||
try:
|
||||
mime_type = magic.from_buffer(file_content, mime=True)
|
||||
if mime_type not in ALLOWED_MIME_TYPES:
|
||||
logger.warning(f"Rejected file with MIME type: {mime_type}")
|
||||
logger.warning("Rejected file with MIME type: %s", mime_type)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Invalid file type. File must be XML, ZIP, or GZIP format.",
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
except Exception as e: # pylint: disable=broad-exception-caught
|
||||
# If magic fails, log but continue (fallback to extension check)
|
||||
logger.warning(f"MIME type detection failed: {str(e)}")
|
||||
logger.warning("MIME type detection failed: %s", str(e))
|
||||
|
||||
|
||||
def _validate_upload_file(file: UploadFile, file_content: bytes) -> None:
|
||||
@@ -66,9 +75,7 @@ def _validate_upload_file(file: UploadFile, file_content: bytes) -> None:
|
||||
"""
|
||||
# Security: Validate filename is provided
|
||||
if not file.filename:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail="Filename is required"
|
||||
)
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Filename is required")
|
||||
|
||||
# Security: Validate file extension
|
||||
file_ext = "." + file.filename.rsplit(".", 1)[-1].lower() if "." in file.filename else ""
|
||||
@@ -91,19 +98,14 @@ def _handle_upload_value_error(filename: str, error_message: str) -> None:
|
||||
|
||||
Always raises — never returns.
|
||||
"""
|
||||
logger.error(f"ValueError processing report {filename}: {error_message}")
|
||||
logger.error("ValueError processing report %s: %s", filename, error_message)
|
||||
if "too large" in error_message.lower():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, detail="File too large"
|
||||
)
|
||||
elif "zip bomb" in error_message.lower():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid archive file"
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid report format"
|
||||
)
|
||||
if "zip bomb" in error_message.lower():
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid archive file")
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid report format")
|
||||
|
||||
|
||||
class UploadResponse(BaseModel):
|
||||
@@ -138,6 +140,20 @@ class ReportSummary(BaseModel):
|
||||
failed_count: int
|
||||
|
||||
|
||||
class AllReportsItem(BaseModel):
|
||||
"""Single report item for the cross-domain reports list"""
|
||||
|
||||
report_id: str
|
||||
domain: str
|
||||
org_name: str
|
||||
begin_date: str
|
||||
end_date: str
|
||||
total_count: int
|
||||
passed_count: int
|
||||
failed_count: int
|
||||
pass_rate: float
|
||||
|
||||
|
||||
class PaginatedReportResponse(BaseModel):
|
||||
"""Paginated reports response model"""
|
||||
|
||||
@@ -149,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)
|
||||
|
||||
@@ -184,8 +200,23 @@ async def upload_report(file: UploadFile = File(...)):
|
||||
detail=f"Invalid domain in report: {error_msg}",
|
||||
)
|
||||
|
||||
# Store the report
|
||||
# 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) or report_exists(db, domain, report_id)
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=(
|
||||
f"Report '{report_id}' for domain '{domain}' has already been uploaded. "
|
||||
"Duplicate reports are not stored to keep statistics accurate."
|
||||
),
|
||||
)
|
||||
|
||||
# Store the report
|
||||
save_parsed_report(db, report)
|
||||
db.commit()
|
||||
store.add_report(report)
|
||||
|
||||
processed_records = report.get("summary", {}).get("total_count", 0)
|
||||
@@ -202,30 +233,69 @@ async def upload_report(file: UploadFile = File(...)):
|
||||
except ValueError as e:
|
||||
# Security: Sanitize error messages from parser
|
||||
_handle_upload_value_error(file.filename, str(e))
|
||||
except Exception as e:
|
||||
except Exception as e: # pylint: disable=broad-exception-caught
|
||||
# Security: Don't expose internal errors to client
|
||||
logger.error(f"Unexpected error processing report {file.filename}: {str(e)}")
|
||||
logger.error("Unexpected error processing report %s: %s", file.filename, str(e))
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Error processing report. Please contact support if this persists.",
|
||||
)
|
||||
) from e
|
||||
|
||||
|
||||
@router.get("", response_model=List[AllReportsItem])
|
||||
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] = []
|
||||
for domain in domains:
|
||||
domain_reports = store.get_domain_reports(domain)
|
||||
for report in domain_reports:
|
||||
summary = report.get("summary", {})
|
||||
total = summary.get("total_count", 0)
|
||||
passed = summary.get("passed_count", 0)
|
||||
pass_rate = round(passed / total * 100, 1) if total > 0 else 0.0
|
||||
all_reports.append(
|
||||
AllReportsItem(
|
||||
report_id=report.get("report_id", ""),
|
||||
domain=domain,
|
||||
org_name=report.get("org_name", ""),
|
||||
begin_date=str(report.get("begin_date", "")),
|
||||
end_date=str(report.get("end_date", "")),
|
||||
total_count=total,
|
||||
passed_count=passed,
|
||||
failed_count=summary.get("failed_count", 0),
|
||||
pass_rate=pass_rate,
|
||||
)
|
||||
)
|
||||
|
||||
# end_date is stored in ISO 8601 format (YYYY-MM-DDTHH:MM:SS), so lexicographic
|
||||
# sorting produces correct chronological order.
|
||||
all_reports.sort(key=lambda r: r.end_date, reverse=True)
|
||||
return 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:
|
||||
@@ -237,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:
|
||||
@@ -281,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
|
||||
@@ -293,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:
|
||||
@@ -306,10 +380,11 @@ async def get_domain_reports_paginated(
|
||||
|
||||
if sort_field == "total_count":
|
||||
all_reports.sort(
|
||||
key=lambda r: r.get("summary", {}).get("total_count", 0), reverse=(sort_order == "desc")
|
||||
key=lambda r: r.get("summary", {}).get("total_count", 0),
|
||||
reverse=sort_order == "desc",
|
||||
)
|
||||
else:
|
||||
all_reports.sort(key=lambda r: r.get(sort_field, ""), reverse=(sort_order == "desc"))
|
||||
all_reports.sort(key=lambda r: r.get(sort_field, ""), reverse=sort_order == "desc")
|
||||
|
||||
# Apply pagination
|
||||
total = len(all_reports)
|
||||
@@ -335,3 +410,149 @@ async def get_domain_reports_paginated(
|
||||
return PaginatedReportResponse(
|
||||
total=total, page=page, page_size=page_size, total_pages=total_pages, reports=report_entries
|
||||
)
|
||||
|
||||
|
||||
class DeleteReportResponse(BaseModel):
|
||||
"""Response model for report deletion"""
|
||||
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/domain/{domain}/reports/{report_id}",
|
||||
response_model=DeleteReportResponse,
|
||||
)
|
||||
async def delete_report(domain: str, report_id: str, db: Session = Depends(get_db)):
|
||||
"""
|
||||
Delete a single DMARC report for a domain.
|
||||
|
||||
Removes the report from the store and recomputes all domain statistics so
|
||||
that aggregated numbers remain accurate after deletion.
|
||||
"""
|
||||
store = ReportStore.get_instance()
|
||||
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(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Report '{report_id}' not found for domain '{domain}'.",
|
||||
)
|
||||
|
||||
return DeleteReportResponse(
|
||||
success=True,
|
||||
message=f"Report '{report_id}' for domain '{domain}' deleted successfully.",
|
||||
)
|
||||
|
||||
|
||||
class ReportRecordDetail(BaseModel):
|
||||
"""Detailed record from a DMARC report"""
|
||||
|
||||
source_ip: str
|
||||
count: int
|
||||
disposition: str
|
||||
dkim_result: str
|
||||
spf_result: str
|
||||
header_from: str
|
||||
spf: Optional[List[Dict[str, Any]]] = None
|
||||
dkim: Optional[List[Dict[str, Any]]] = None
|
||||
|
||||
|
||||
class ReportPolicyDetail(BaseModel):
|
||||
"""Published policy from a DMARC report"""
|
||||
|
||||
p: str
|
||||
sp: str = ""
|
||||
pct: str = "100"
|
||||
|
||||
|
||||
class ReportSummaryDetail(BaseModel):
|
||||
"""Summary statistics for a DMARC report"""
|
||||
|
||||
total_count: int
|
||||
passed_count: int
|
||||
failed_count: int
|
||||
pass_rate: float
|
||||
|
||||
|
||||
class ReportDetail(BaseModel):
|
||||
"""Full detail of a single DMARC report"""
|
||||
|
||||
report_id: str
|
||||
org_name: str
|
||||
email: str
|
||||
domain: str
|
||||
begin_date: str
|
||||
end_date: str
|
||||
begin_timestamp: int
|
||||
end_timestamp: int
|
||||
policy: ReportPolicyDetail
|
||||
records: List[ReportRecordDetail]
|
||||
summary: ReportSummaryDetail
|
||||
|
||||
|
||||
@router.get("/{report_id}", response_model=ReportDetail)
|
||||
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:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Report '{report_id}' not found.",
|
||||
)
|
||||
|
||||
# Normalize the policy field
|
||||
policy_val = report.get("policy", {})
|
||||
if isinstance(policy_val, str):
|
||||
policy_val = {"p": policy_val, "sp": "", "pct": "100"}
|
||||
policy_detail = ReportPolicyDetail(
|
||||
p=policy_val.get("p", "none"),
|
||||
sp=policy_val.get("sp", ""),
|
||||
pct=str(policy_val.get("pct", "100")),
|
||||
)
|
||||
|
||||
# Normalize records
|
||||
record_details = []
|
||||
for rec in report.get("records", []):
|
||||
record_details.append(
|
||||
ReportRecordDetail(
|
||||
source_ip=rec.get("source_ip", ""),
|
||||
count=rec.get("count", 0),
|
||||
disposition=rec.get("disposition", "none"),
|
||||
dkim_result=rec.get("dkim_result", ""),
|
||||
spf_result=rec.get("spf_result", ""),
|
||||
header_from=rec.get("header_from", ""),
|
||||
spf=rec.get("spf") if isinstance(rec.get("spf"), list) else None,
|
||||
dkim=rec.get("dkim") if isinstance(rec.get("dkim"), list) else None,
|
||||
)
|
||||
)
|
||||
|
||||
raw_summary = report.get("summary", {})
|
||||
summary_detail = ReportSummaryDetail(
|
||||
total_count=raw_summary.get("total_count", 0),
|
||||
passed_count=raw_summary.get("passed_count", 0),
|
||||
failed_count=raw_summary.get("failed_count", 0),
|
||||
pass_rate=raw_summary.get("pass_rate", 0.0),
|
||||
)
|
||||
|
||||
return ReportDetail(
|
||||
report_id=report.get("report_id", ""),
|
||||
org_name=report.get("org_name", ""),
|
||||
email=report.get("email", ""),
|
||||
domain=report.get("domain", ""),
|
||||
begin_date=str(report.get("begin_date", "")),
|
||||
end_date=str(report.get("end_date", "")),
|
||||
begin_timestamp=report.get("begin_timestamp", 0),
|
||||
end_timestamp=report.get("end_timestamp", 0),
|
||||
policy=policy_detail,
|
||||
records=record_details,
|
||||
summary=summary_detail,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,783 @@
|
||||
"""
|
||||
Settings API endpoints.
|
||||
|
||||
Provides endpoints to read and write application-level settings persisted
|
||||
in the ``settings`` database table. Settings are organised into categories:
|
||||
|
||||
- ``general`` – App name, base URL, reports-per-page, etc.
|
||||
- ``dmarc`` – Default DMARC policy, percentage, etc.
|
||||
- ``dns`` – Default DNS resolver, Cloudflare DoH toggle.
|
||||
- ``cloudflare`` – Cloudflare API token and Zone ID.
|
||||
- ``forensics`` – Forensic report privacy and retention controls.
|
||||
- ``notifications`` – Future alerting/notification settings.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.credential_encryption import decrypt_secret, encrypt_secret, is_encrypted_secret
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.models.setting import Setting
|
||||
from app.services.alert_history import (
|
||||
list_alert_config_audit,
|
||||
list_alert_history,
|
||||
record_alert_config_change,
|
||||
record_alert_evaluation,
|
||||
)
|
||||
from app.services.alert_rules import (
|
||||
enqueue_alert_webhook_events,
|
||||
evaluate_alert_rules,
|
||||
send_current_alerts,
|
||||
)
|
||||
from app.services.notifications import send_notification
|
||||
from app.services.summary_notifications import build_summary, send_summary_notification
|
||||
from app.services.workspace_audit import record_workspace_audit_log
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Defaults – used to seed missing keys on first read
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
SETTING_DEFAULTS: List[Dict[str, Any]] = [
|
||||
# ── General ─────────────────────────────────────────────────────────────
|
||||
{
|
||||
"key": "general.app_name",
|
||||
"value": "DMARQ",
|
||||
"description": "Application display name shown in the UI",
|
||||
"value_type": "string",
|
||||
"category": "general",
|
||||
},
|
||||
{
|
||||
"key": "general.base_url",
|
||||
"value": "",
|
||||
"description": "Public base URL (e.g. https://dmarc.example.com)",
|
||||
"value_type": "string",
|
||||
"category": "general",
|
||||
},
|
||||
{
|
||||
"key": "general.reports_per_page",
|
||||
"value": "25",
|
||||
"description": "Number of reports shown per page in the reports list",
|
||||
"value_type": "integer",
|
||||
"category": "general",
|
||||
},
|
||||
{
|
||||
"key": "general.session_lifetime_minutes",
|
||||
"value": "1440",
|
||||
"description": "How long a login session stays valid (minutes)",
|
||||
"value_type": "integer",
|
||||
"category": "general",
|
||||
},
|
||||
# ── DMARC ────────────────────────────────────────────────────────────────
|
||||
{
|
||||
"key": "dmarc.default_policy",
|
||||
"value": "none",
|
||||
"description": "Default DMARC policy applied when adding a new domain",
|
||||
"value_type": "string",
|
||||
"category": "dmarc",
|
||||
},
|
||||
{
|
||||
"key": "dmarc.default_percentage",
|
||||
"value": "100",
|
||||
"description": "Default DMARC percentage (pct) tag for new domains",
|
||||
"value_type": "integer",
|
||||
"category": "dmarc",
|
||||
},
|
||||
{
|
||||
"key": "dmarc.default_adkim",
|
||||
"value": "r",
|
||||
"description": "Default DKIM alignment mode: r (relaxed) or s (strict)",
|
||||
"value_type": "string",
|
||||
"category": "dmarc",
|
||||
},
|
||||
{
|
||||
"key": "dmarc.default_aspf",
|
||||
"value": "r",
|
||||
"description": "Default SPF alignment mode: r (relaxed) or s (strict)",
|
||||
"value_type": "string",
|
||||
"category": "dmarc",
|
||||
},
|
||||
# ── DNS ──────────────────────────────────────────────────────────────────
|
||||
{
|
||||
"key": "dns.resolver",
|
||||
"value": "system",
|
||||
"description": "DNS resolver to use: system or cloudflare",
|
||||
"value_type": "string",
|
||||
"category": "dns",
|
||||
},
|
||||
# ── Cloudflare ───────────────────────────────────────────────────────────
|
||||
{
|
||||
"key": "cloudflare.api_token",
|
||||
"value": "",
|
||||
"description": "Cloudflare API token for DNS record management",
|
||||
"value_type": "string",
|
||||
"category": "cloudflare",
|
||||
},
|
||||
{
|
||||
"key": "cloudflare.zone_id",
|
||||
"value": "",
|
||||
"description": "Cloudflare Zone ID for DNS record management",
|
||||
"value_type": "string",
|
||||
"category": "cloudflare",
|
||||
},
|
||||
# ── Forensics ────────────────────────────────────────────────────────────
|
||||
{
|
||||
"key": "forensics.redaction_mode",
|
||||
"value": "balanced",
|
||||
"description": "Forensic report email-address redaction mode: balanced, domain_only, or strict",
|
||||
"value_type": "string",
|
||||
"category": "forensics",
|
||||
},
|
||||
{
|
||||
"key": "forensics.redact_long_tokens_enabled",
|
||||
"value": "true",
|
||||
"description": "Redact long opaque tokens in forensic report metadata",
|
||||
"value_type": "boolean",
|
||||
"category": "forensics",
|
||||
},
|
||||
# ── Notifications ─────────────────────────────────────────────────────────
|
||||
{
|
||||
"key": "notifications.apprise_enabled",
|
||||
"value": "false",
|
||||
"description": "Send notifications through configured Apprise target URLs",
|
||||
"value_type": "boolean",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.apprise_urls",
|
||||
"value": "",
|
||||
"description": "Newline-separated Apprise notification target URLs",
|
||||
"value_type": "string",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.min_send_interval_minutes",
|
||||
"value": "15",
|
||||
"description": "Minimum minutes between outbound notification deliveries",
|
||||
"value_type": "integer",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.redact_pii_enabled",
|
||||
"value": "true",
|
||||
"description": "Redact email addresses from outbound notification titles and bodies",
|
||||
"value_type": "boolean",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.last_sent_at",
|
||||
"value": "",
|
||||
"description": "Internal timestamp for outbound notification rate limiting",
|
||||
"value_type": "string",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.alert_new_sources_enabled",
|
||||
"value": "true",
|
||||
"description": "Alert when a new sending source appears in recent DMARC reports",
|
||||
"value_type": "boolean",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.alert_compliance_drop_enabled",
|
||||
"value": "true",
|
||||
"description": "Alert when DMARC compliance drops by the configured percentage points",
|
||||
"value_type": "boolean",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.alert_compliance_drop_points",
|
||||
"value": "10",
|
||||
"description": "Minimum compliance-rate drop, in percentage points, before alerting",
|
||||
"value_type": "integer",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.alert_failure_threshold_enabled",
|
||||
"value": "true",
|
||||
"description": "Alert when DMARC failures exceed the configured daily threshold",
|
||||
"value_type": "boolean",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.alert_failure_threshold_count",
|
||||
"value": "100",
|
||||
"description": "Minimum failed message count in the last day before alerting",
|
||||
"value_type": "integer",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.alert_missing_reports_enabled",
|
||||
"value": "true",
|
||||
"description": "Alert when a monitored domain has not received recent DMARC reports",
|
||||
"value_type": "boolean",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.alert_missing_reports_days",
|
||||
"value": "2",
|
||||
"description": "Number of days without DMARC reports before alerting",
|
||||
"value_type": "integer",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.summary_daily_enabled",
|
||||
"value": "false",
|
||||
"description": "Send one daily DMARC activity summary notification",
|
||||
"value_type": "boolean",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.summary_weekly_enabled",
|
||||
"value": "false",
|
||||
"description": "Send one weekly DMARC activity summary notification",
|
||||
"value_type": "boolean",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.summary_send_hour_utc",
|
||||
"value": "8",
|
||||
"description": "UTC hour when scheduled summary notifications can be sent",
|
||||
"value_type": "integer",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.summary_weekday_utc",
|
||||
"value": "0",
|
||||
"description": "UTC weekday for weekly summaries, where 0 is Monday",
|
||||
"value_type": "integer",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.summary_daily_last_sent_date",
|
||||
"value": "",
|
||||
"description": "Internal date marker for the last sent daily summary",
|
||||
"value_type": "string",
|
||||
"category": "notifications",
|
||||
},
|
||||
{
|
||||
"key": "notifications.summary_weekly_last_sent_week",
|
||||
"value": "",
|
||||
"description": "Internal ISO week marker for the last sent weekly summary",
|
||||
"value_type": "string",
|
||||
"category": "notifications",
|
||||
},
|
||||
# ── Optional AI / MCP ───────────────────────────────────────────────────
|
||||
{
|
||||
"key": "ai.enabled",
|
||||
"value": "false",
|
||||
"description": "Enable optional AI assistance endpoints",
|
||||
"value_type": "boolean",
|
||||
"category": "ai",
|
||||
},
|
||||
{
|
||||
"key": "ai.provider",
|
||||
"value": "template",
|
||||
"description": "AI provider: template, local, or remote",
|
||||
"value_type": "string",
|
||||
"category": "ai",
|
||||
},
|
||||
{
|
||||
"key": "ai.model",
|
||||
"value": "",
|
||||
"description": "Optional model name for local or remote providers",
|
||||
"value_type": "string",
|
||||
"category": "ai",
|
||||
},
|
||||
{
|
||||
"key": "ai.remote_base_url",
|
||||
"value": "",
|
||||
"description": "Optional remote provider base URL; credentials should be injected by environment",
|
||||
"value_type": "string",
|
||||
"category": "ai",
|
||||
},
|
||||
{
|
||||
"key": "ai.redaction_mode",
|
||||
"value": "strict",
|
||||
"description": "Redaction mode for AI-safe context: strict or balanced",
|
||||
"value_type": "string",
|
||||
"category": "ai",
|
||||
},
|
||||
{
|
||||
"key": "ai.action_tools_enabled",
|
||||
"value": "false",
|
||||
"description": "Allow human-confirmed action proposals to be recorded",
|
||||
"value_type": "boolean",
|
||||
"category": "ai",
|
||||
},
|
||||
{
|
||||
"key": "mcp.enabled",
|
||||
"value": "false",
|
||||
"description": "Enable the scoped read-only MCP endpoint",
|
||||
"value_type": "boolean",
|
||||
"category": "mcp",
|
||||
},
|
||||
]
|
||||
|
||||
# Keys whose values should be redacted in GET responses (treated as secrets)
|
||||
_SECRET_KEYS = {
|
||||
"cloudflare.api_token",
|
||||
"notifications.apprise_urls",
|
||||
"notifications.smtp_password",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _seed_defaults(db: Session) -> None:
|
||||
"""Insert any missing default settings rows (idempotent)."""
|
||||
for defaults in SETTING_DEFAULTS:
|
||||
key = defaults["key"]
|
||||
if db.query(Setting).filter(Setting.key == key).first() is None:
|
||||
db.add(
|
||||
Setting(
|
||||
key=key,
|
||||
value=defaults["value"],
|
||||
description=defaults["description"],
|
||||
value_type=defaults["value_type"],
|
||||
category=defaults["category"],
|
||||
)
|
||||
)
|
||||
_migrate_plaintext_secret_settings(db)
|
||||
db.commit()
|
||||
|
||||
|
||||
def _get_setting(key: str, db: Session) -> Optional[Setting]:
|
||||
return db.query(Setting).filter(Setting.key == key).first()
|
||||
|
||||
|
||||
def _migrate_plaintext_secret_settings(db: Session) -> None:
|
||||
"""Encrypt legacy plaintext secret settings opportunistically."""
|
||||
rows = db.query(Setting).filter(Setting.key.in_(_SECRET_KEYS)).all()
|
||||
for row in rows:
|
||||
if row.value and not is_encrypted_secret(row.value):
|
||||
row.value = encrypt_secret(row.value)
|
||||
|
||||
|
||||
def _stored_value_for_setting(key: str, value: Optional[str]) -> Optional[str]:
|
||||
if key in _SECRET_KEYS:
|
||||
return encrypt_secret(value)
|
||||
return value
|
||||
|
||||
|
||||
def _plain_value_for_setting(key: str, value: Optional[str]) -> Optional[str]:
|
||||
if key not in _SECRET_KEYS:
|
||||
return value
|
||||
return decrypt_secret(value)
|
||||
|
||||
|
||||
def _audit_value_for_setting(key: str, value: Optional[str]) -> Optional[str]:
|
||||
if key in _SECRET_KEYS:
|
||||
return "[redacted]" if value else ""
|
||||
return value
|
||||
|
||||
|
||||
def _should_audit_setting(key: str) -> bool:
|
||||
return key.startswith(("notifications.", "forensics.", "ai.", "mcp."))
|
||||
|
||||
|
||||
def _audit_setting_change(
|
||||
db: Session,
|
||||
*,
|
||||
key: str,
|
||||
old_plain: Optional[str],
|
||||
new_plain: Optional[str],
|
||||
auth_context: Optional[Dict[str, Any]],
|
||||
request: Optional[Request] = None,
|
||||
) -> None:
|
||||
if not _should_audit_setting(key) or old_plain == new_plain:
|
||||
return
|
||||
record_alert_config_change(
|
||||
db,
|
||||
key=key,
|
||||
old_value=_audit_value_for_setting(key, old_plain),
|
||||
new_value=_audit_value_for_setting(key, new_plain),
|
||||
auth_context=auth_context,
|
||||
)
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db, commit=False)
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="setting.changed",
|
||||
entity_type="setting",
|
||||
entity_id=key,
|
||||
entity_name=key,
|
||||
details={
|
||||
"key": key,
|
||||
"old_value": _audit_value_for_setting(key, old_plain),
|
||||
"new_value": _audit_value_for_setting(key, new_plain),
|
||||
},
|
||||
auth_context=auth_context,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
def _row_to_dict(row: Setting, redact_secrets: bool = True) -> Dict[str, Any]:
|
||||
value = row.value
|
||||
if redact_secrets and row.key in _SECRET_KEYS and value:
|
||||
value = "**redacted**"
|
||||
return {
|
||||
"key": row.key,
|
||||
"value": value,
|
||||
"description": row.description,
|
||||
"value_type": row.value_type,
|
||||
"category": row.category,
|
||||
"updated_at": row.updated_at.isoformat() if row.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SettingUpdate(BaseModel):
|
||||
"""Payload for updating a single setting."""
|
||||
|
||||
value: Optional[str] = None
|
||||
|
||||
|
||||
class BulkSettingsUpdate(BaseModel):
|
||||
"""Payload for updating multiple settings at once."""
|
||||
|
||||
settings: Dict[str, Optional[str]]
|
||||
|
||||
|
||||
class SettingResponse(BaseModel):
|
||||
"""Response for a single setting."""
|
||||
|
||||
key: str
|
||||
value: Optional[str]
|
||||
description: Optional[str]
|
||||
value_type: str
|
||||
category: str
|
||||
updated_at: Optional[str]
|
||||
|
||||
|
||||
class NotificationTestResponse(BaseModel):
|
||||
"""Sanitized response from a test notification send."""
|
||||
|
||||
success: bool
|
||||
message: str
|
||||
configured_targets: int = 0
|
||||
invalid_targets: int = 0
|
||||
skipped: bool = False
|
||||
rate_limited: bool = False
|
||||
error: Optional[str] = None
|
||||
|
||||
|
||||
class AlertRulesResponse(BaseModel):
|
||||
"""Current alert-rule evaluation response."""
|
||||
|
||||
alerts: List[Dict[str, Any]]
|
||||
|
||||
|
||||
class AlertNotificationResponse(BaseModel):
|
||||
"""Alert-rule evaluation plus notification delivery status."""
|
||||
|
||||
alerts: List[Dict[str, Any]]
|
||||
notification: Dict[str, Any]
|
||||
|
||||
|
||||
class AlertHistoryResponse(BaseModel):
|
||||
"""Persisted alert history response."""
|
||||
|
||||
history: List[Dict[str, Any]]
|
||||
|
||||
|
||||
class AlertConfigurationAuditResponse(BaseModel):
|
||||
"""Persisted alert configuration audit response."""
|
||||
|
||||
audit: List[Dict[str, Any]]
|
||||
|
||||
|
||||
class SummaryResponse(BaseModel):
|
||||
"""Current DMARC summary notification preview."""
|
||||
|
||||
summary: Dict[str, Any]
|
||||
|
||||
|
||||
class SummaryNotificationResponse(BaseModel):
|
||||
"""DMARC summary plus notification delivery status."""
|
||||
|
||||
summary: Dict[str, Any]
|
||||
notification: Dict[str, Any]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("", response_model=List[SettingResponse])
|
||||
async def list_settings(
|
||||
category: Optional[str] = None,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> List[SettingResponse]:
|
||||
"""
|
||||
Return all persisted settings, optionally filtered by category.
|
||||
|
||||
Missing rows are seeded from defaults before returning.
|
||||
"""
|
||||
_seed_defaults(db)
|
||||
query = db.query(Setting)
|
||||
if category:
|
||||
query = query.filter(Setting.category == category)
|
||||
rows = query.order_by(Setting.category, Setting.key).all()
|
||||
return [_row_to_dict(row) for row in rows]
|
||||
|
||||
|
||||
@router.post("/notifications/test", response_model=NotificationTestResponse)
|
||||
async def test_notification_settings(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> NotificationTestResponse:
|
||||
"""Send a test notification using the configured Apprise targets."""
|
||||
_seed_defaults(db)
|
||||
result = send_notification(
|
||||
db,
|
||||
title="DMARQ test notification",
|
||||
body="This confirms that DMARQ can reach the configured notification target.",
|
||||
force=True,
|
||||
)
|
||||
if not result.success:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=result.to_dict(),
|
||||
)
|
||||
return result.to_dict()
|
||||
|
||||
|
||||
@router.get("/notifications/alerts", response_model=AlertRulesResponse)
|
||||
async def evaluate_notification_alerts(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> AlertRulesResponse:
|
||||
"""Evaluate enabled notification alert rules against current DMARC data."""
|
||||
_seed_defaults(db)
|
||||
alerts = evaluate_alert_rules(db)
|
||||
record_alert_evaluation(db, alerts)
|
||||
enqueue_alert_webhook_events(db, alerts)
|
||||
return {"alerts": alerts}
|
||||
|
||||
|
||||
@router.post("/notifications/alerts/send", response_model=AlertNotificationResponse)
|
||||
async def send_notification_alerts(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> AlertNotificationResponse:
|
||||
"""Evaluate current alert rules and send a notification summary when needed."""
|
||||
_seed_defaults(db)
|
||||
result = send_current_alerts(db)
|
||||
notification = result["notification"]
|
||||
if result["alerts"] and not notification.get("success"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=result,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/notifications/alerts/history", response_model=AlertHistoryResponse)
|
||||
async def get_notification_alert_history(
|
||||
active: Optional[bool] = None,
|
||||
limit: int = 50,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> AlertHistoryResponse:
|
||||
"""Return persisted alert history rows."""
|
||||
return {"history": list_alert_history(db, active=active, limit=max(1, min(limit, 200)))}
|
||||
|
||||
|
||||
@router.get("/notifications/config-audit", response_model=AlertConfigurationAuditResponse)
|
||||
async def get_notification_config_audit(
|
||||
limit: int = 50,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> AlertConfigurationAuditResponse:
|
||||
"""Return recent notification and alert-rule configuration changes."""
|
||||
return {"audit": list_alert_config_audit(db, limit=max(1, min(limit, 200)))}
|
||||
|
||||
|
||||
@router.get("/notifications/summary", response_model=SummaryResponse)
|
||||
async def preview_notification_summary(
|
||||
period: str = "daily",
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> SummaryResponse:
|
||||
"""Preview a daily or weekly DMARC summary notification."""
|
||||
_seed_defaults(db)
|
||||
try:
|
||||
summary = build_summary(db, period)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
return {"summary": summary}
|
||||
|
||||
|
||||
@router.post("/notifications/summary/send", response_model=SummaryNotificationResponse)
|
||||
async def send_notification_summary(
|
||||
period: str = "daily",
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> SummaryNotificationResponse:
|
||||
"""Send a daily or weekly DMARC summary notification immediately."""
|
||||
_seed_defaults(db)
|
||||
try:
|
||||
result = send_summary_notification(db, period)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
notification = result["notification"]
|
||||
if not notification.get("success"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=result,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/{key:path}", response_model=SettingResponse)
|
||||
async def get_setting(
|
||||
key: str,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> SettingResponse:
|
||||
"""Return a single setting by key."""
|
||||
_seed_defaults(db)
|
||||
row = _get_setting(key, db)
|
||||
if row is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Setting '{key}' not found",
|
||||
)
|
||||
return _row_to_dict(row)
|
||||
|
||||
|
||||
@router.put("/{key:path}", response_model=SettingResponse)
|
||||
async def update_setting(
|
||||
key: str,
|
||||
payload: SettingUpdate,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> SettingResponse:
|
||||
"""Update or create a single setting."""
|
||||
row = _get_setting(key, db)
|
||||
new_value = payload.value
|
||||
if row is None:
|
||||
# Find matching default metadata
|
||||
default_meta = next((d for d in SETTING_DEFAULTS if d["key"] == key), None)
|
||||
new_plain = _plain_value_for_setting(key, new_value)
|
||||
row = Setting(
|
||||
key=key,
|
||||
value=_stored_value_for_setting(key, new_value),
|
||||
description=default_meta["description"] if default_meta else None,
|
||||
value_type=default_meta["value_type"] if default_meta else "string",
|
||||
category=default_meta["category"] if default_meta else "general",
|
||||
)
|
||||
db.add(row)
|
||||
_audit_setting_change(
|
||||
db,
|
||||
key=key,
|
||||
old_plain=None,
|
||||
new_plain=new_plain,
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
)
|
||||
else:
|
||||
# For secret keys, only update if not the redacted placeholder
|
||||
if key in _SECRET_KEYS and payload.value == "**redacted**":
|
||||
db.refresh(row)
|
||||
return _row_to_dict(row)
|
||||
old_plain = _plain_value_for_setting(key, row.value)
|
||||
new_plain = _plain_value_for_setting(key, new_value)
|
||||
row.value = _stored_value_for_setting(key, new_value)
|
||||
_audit_setting_change(
|
||||
db,
|
||||
key=key,
|
||||
old_plain=old_plain,
|
||||
new_plain=new_plain,
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
)
|
||||
db.commit()
|
||||
db.refresh(row)
|
||||
return _row_to_dict(row)
|
||||
|
||||
|
||||
@router.post("/bulk", response_model=List[SettingResponse])
|
||||
async def bulk_update_settings(
|
||||
payload: BulkSettingsUpdate,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> List[SettingResponse]:
|
||||
"""
|
||||
Update multiple settings in a single request.
|
||||
|
||||
Accepts ``{"settings": {"key1": "value1", "key2": "value2", ...}}``.
|
||||
"""
|
||||
results = []
|
||||
for key, value in payload.settings.items():
|
||||
row = _get_setting(key, db)
|
||||
if row is None:
|
||||
default_meta = next((d for d in SETTING_DEFAULTS if d["key"] == key), None)
|
||||
new_plain = _plain_value_for_setting(key, value)
|
||||
row = Setting(
|
||||
key=key,
|
||||
value=_stored_value_for_setting(key, value),
|
||||
description=default_meta["description"] if default_meta else None,
|
||||
value_type=default_meta["value_type"] if default_meta else "string",
|
||||
category=default_meta["category"] if default_meta else "general",
|
||||
)
|
||||
db.add(row)
|
||||
_audit_setting_change(
|
||||
db,
|
||||
key=key,
|
||||
old_plain=None,
|
||||
new_plain=new_plain,
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
)
|
||||
else:
|
||||
# Skip secret placeholder updates
|
||||
if key in _SECRET_KEYS and value == "**redacted**":
|
||||
results.append(_row_to_dict(row))
|
||||
continue
|
||||
old_plain = _plain_value_for_setting(key, row.value)
|
||||
new_plain = _plain_value_for_setting(key, value)
|
||||
row.value = _stored_value_for_setting(key, value)
|
||||
_audit_setting_change(
|
||||
db,
|
||||
key=key,
|
||||
old_plain=old_plain,
|
||||
new_plain=new_plain,
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
)
|
||||
results.append(_row_to_dict(row))
|
||||
db.commit()
|
||||
# Re-read rows to get updated_at timestamps
|
||||
refreshed = []
|
||||
for item in results:
|
||||
row = _get_setting(item["key"], db)
|
||||
if row:
|
||||
refreshed.append(_row_to_dict(row))
|
||||
return refreshed
|
||||
@@ -1,5 +1,15 @@
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Security, status
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
from pydantic import BaseModel, EmailStr
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import api_key_header, require_admin_auth, security_bearer
|
||||
from app.models.domain import Domain
|
||||
from app.models.mail_source import MailSource
|
||||
from app.models.setting import Setting
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -10,12 +20,80 @@ setup_status = {
|
||||
"app_name": "DMARQ",
|
||||
}
|
||||
|
||||
SETUP_COMPLETE_KEY = "setup.is_complete"
|
||||
SETUP_ADMIN_EMAIL_KEY = "setup.admin_email"
|
||||
GENERAL_APP_NAME_KEY = "general.app_name"
|
||||
GENERAL_BASE_URL_KEY = "general.base_url"
|
||||
|
||||
|
||||
def _setting_value(db: Session, key: str) -> Optional[str]:
|
||||
row = db.query(Setting).filter(Setting.key == key).first()
|
||||
return row.value if row else None
|
||||
|
||||
|
||||
def _is_true(value: Optional[str]) -> bool:
|
||||
return str(value).strip().lower() in {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def _upsert_setting(
|
||||
db: Session,
|
||||
key: str,
|
||||
value: Optional[str],
|
||||
*,
|
||||
description: str,
|
||||
value_type: str = "string",
|
||||
category: str = "setup",
|
||||
) -> None:
|
||||
row = db.query(Setting).filter(Setting.key == key).first()
|
||||
if row is None:
|
||||
db.add(
|
||||
Setting(
|
||||
key=key,
|
||||
value=value,
|
||||
description=description,
|
||||
value_type=value_type,
|
||||
category=category,
|
||||
)
|
||||
)
|
||||
return
|
||||
|
||||
row.value = value
|
||||
row.description = row.description or description
|
||||
row.value_type = row.value_type or value_type
|
||||
row.category = row.category or category
|
||||
|
||||
|
||||
def _refresh_setup_status_from_db(db: Session) -> dict:
|
||||
"""Merge persisted setup state into the legacy in-memory setup status."""
|
||||
persisted_complete = _is_true(_setting_value(db, SETUP_COMPLETE_KEY))
|
||||
if persisted_complete:
|
||||
setup_status["is_setup_complete"] = True
|
||||
|
||||
persisted_admin_email = _setting_value(db, SETUP_ADMIN_EMAIL_KEY)
|
||||
if persisted_admin_email:
|
||||
setup_status["admin_email"] = persisted_admin_email
|
||||
|
||||
persisted_app_name = _setting_value(db, GENERAL_APP_NAME_KEY)
|
||||
if persisted_app_name:
|
||||
setup_status["app_name"] = persisted_app_name
|
||||
|
||||
return setup_status
|
||||
|
||||
|
||||
def _setup_is_complete(db: Session) -> bool:
|
||||
return bool(setup_status["is_setup_complete"]) or _is_true(
|
||||
_setting_value(db, SETUP_COMPLETE_KEY)
|
||||
)
|
||||
|
||||
|
||||
class SetupStatusResponse(BaseModel):
|
||||
"""Setup status response"""
|
||||
|
||||
is_setup_complete: bool
|
||||
app_name: str
|
||||
total_domains: int = 0
|
||||
total_mail_sources: int = 0
|
||||
enabled_mail_sources: int = 0
|
||||
|
||||
|
||||
class AdminSetupRequest(BaseModel):
|
||||
@@ -33,34 +111,67 @@ class SystemConfigRequest(BaseModel):
|
||||
base_url: str
|
||||
|
||||
|
||||
async def require_setup_write_auth(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
api_key: Optional[str] = Security(api_key_header),
|
||||
bearer: Optional[HTTPAuthorizationCredentials] = Security(security_bearer),
|
||||
) -> dict:
|
||||
"""Allow unauthenticated first-time setup writes, then require admin auth."""
|
||||
if not _setup_is_complete(db):
|
||||
return {"auth_type": "initial_setup"}
|
||||
return await require_admin_auth(request=request, api_key=api_key, bearer=bearer)
|
||||
|
||||
|
||||
@router.get("/status", response_model=SetupStatusResponse)
|
||||
async def get_setup_status():
|
||||
async def get_setup_status(db: Session = Depends(get_db)):
|
||||
"""Get the current setup status"""
|
||||
current_status = _refresh_setup_status_from_db(db)
|
||||
return SetupStatusResponse(
|
||||
is_setup_complete=setup_status["is_setup_complete"],
|
||||
app_name=setup_status["app_name"],
|
||||
is_setup_complete=current_status["is_setup_complete"],
|
||||
app_name=current_status["app_name"],
|
||||
total_domains=db.query(Domain.id).count(),
|
||||
total_mail_sources=db.query(MailSource.id).count(),
|
||||
enabled_mail_sources=db.query(MailSource.id)
|
||||
.filter(MailSource.enabled == True) # noqa: E712
|
||||
.count(),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/admin", status_code=201)
|
||||
async def setup_admin(request: AdminSetupRequest):
|
||||
async def setup_admin(
|
||||
request: AdminSetupRequest,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_setup_write_auth),
|
||||
):
|
||||
"""
|
||||
Setup admin user during initial system configuration.
|
||||
For Milestone 1, this simply stores the admin email in memory.
|
||||
"""
|
||||
if setup_status["is_setup_complete"]:
|
||||
if _setup_is_complete(db):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail="Setup already completed"
|
||||
)
|
||||
|
||||
# Store admin email
|
||||
setup_status["admin_email"] = request.email
|
||||
_upsert_setting(
|
||||
db,
|
||||
SETUP_ADMIN_EMAIL_KEY,
|
||||
request.email,
|
||||
description="Email address configured during initial setup",
|
||||
)
|
||||
db.commit()
|
||||
|
||||
return {"message": "Admin user setup completed"}
|
||||
|
||||
|
||||
@router.post("/system", status_code=200)
|
||||
async def setup_system(request: SystemConfigRequest):
|
||||
async def setup_system(
|
||||
request: SystemConfigRequest,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_setup_write_auth),
|
||||
):
|
||||
"""
|
||||
Setup system configuration.
|
||||
For Milestone 1, this simply stores the app name in memory.
|
||||
@@ -68,5 +179,27 @@ async def setup_system(request: SystemConfigRequest):
|
||||
# Store app name
|
||||
setup_status["app_name"] = request.app_name
|
||||
setup_status["is_setup_complete"] = True
|
||||
_upsert_setting(
|
||||
db,
|
||||
GENERAL_APP_NAME_KEY,
|
||||
request.app_name,
|
||||
description="Application display name shown in the UI",
|
||||
category="general",
|
||||
)
|
||||
_upsert_setting(
|
||||
db,
|
||||
GENERAL_BASE_URL_KEY,
|
||||
request.base_url,
|
||||
description="Public base URL for this DMARQ instance",
|
||||
category="general",
|
||||
)
|
||||
_upsert_setting(
|
||||
db,
|
||||
SETUP_COMPLETE_KEY,
|
||||
"true",
|
||||
description="Whether initial setup has been completed",
|
||||
value_type="boolean",
|
||||
)
|
||||
db.commit()
|
||||
|
||||
return {"message": "System settings saved successfully"}
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
from typing import Any, Dict
|
||||
|
||||
from fastapi import APIRouter, Depends, Path, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.utils.stats_summarizer import StatsSummarizer
|
||||
from fastapi import APIRouter, Depends, Path, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -12,7 +13,7 @@ router = APIRouter()
|
||||
async def get_dashboard_statistics(
|
||||
db: Session = Depends(get_db),
|
||||
force_refresh: bool = Query(False, title="Force refresh of statistics"),
|
||||
period_days: int = Query(30, title="Period in days for time-based statistics"),
|
||||
period_days: int = Query(30, ge=1, le=365, title="Period in days for time-based statistics"),
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Get optimized statistics for the dashboard using cached data when possible.
|
||||
@@ -33,7 +34,7 @@ async def get_dashboard_statistics(
|
||||
stats_summarizer.invalidate_cache()
|
||||
|
||||
# Get statistics (from cache or calculate if needed)
|
||||
stats = stats_summarizer.calculate_summary_statistics(db)
|
||||
stats = stats_summarizer.calculate_summary_statistics(db, period_days=period_days)
|
||||
|
||||
# Add version and timestamp
|
||||
stats["api_version"] = "1.0"
|
||||
@@ -47,7 +48,7 @@ async def get_domain_statistics(
|
||||
domain_id: str = Path(..., title="The domain ID or name"),
|
||||
db: Session = Depends(get_db),
|
||||
force_refresh: bool = Query(False, title="Force refresh of statistics"),
|
||||
period_days: int = Query(30, title="Period in days for time-based statistics"),
|
||||
period_days: int = Query(30, ge=1, le=365, title="Period in days for time-based statistics"),
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Get optimized statistics for a specific domain using cached data when possible.
|
||||
@@ -68,7 +69,7 @@ async def get_domain_statistics(
|
||||
stats_summarizer.invalidate_cache(domain_id)
|
||||
|
||||
# Get domain statistics (from cache or calculate if needed)
|
||||
stats = stats_summarizer.calculate_summary_statistics(db, domain_id)
|
||||
stats = stats_summarizer.calculate_summary_statistics(db, domain_id, period_days=period_days)
|
||||
|
||||
# Add version and timestamp
|
||||
stats["api_version"] = "1.0"
|
||||
|
||||
@@ -0,0 +1,189 @@
|
||||
import logging
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.models.domain import Domain
|
||||
from app.models.report import TLSReport
|
||||
from app.services.tls_report_parser import MAX_TLS_REPORT_SIZE, TLSReportParser
|
||||
from app.services.tls_report_persistence import (
|
||||
TLS_REPORT_PRIVACY_CONTROLS,
|
||||
save_tls_report,
|
||||
summarize_tls_reports,
|
||||
tls_report_to_dict,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class TLSFailureResponse(BaseModel):
|
||||
result_type: str
|
||||
failed_session_count: int
|
||||
sending_mta_ip: Optional[str] = None
|
||||
receiving_mx_hostname: Optional[str] = None
|
||||
receiving_mx_helo: Optional[str] = None
|
||||
receiving_ip: Optional[str] = None
|
||||
failure_reason_code: Optional[str] = None
|
||||
additional_information: Optional[str] = None
|
||||
|
||||
|
||||
class TLSReportResponse(BaseModel):
|
||||
id: int
|
||||
report_id: str
|
||||
domain: Optional[str] = None
|
||||
org_name: Optional[str] = None
|
||||
contact_info: Optional[str] = None
|
||||
policy_domain: str
|
||||
policy_type: Optional[str] = None
|
||||
begin_date: Optional[str] = None
|
||||
end_date: Optional[str] = None
|
||||
total_successful_sessions: int
|
||||
total_failure_sessions: int
|
||||
processed_at: Optional[str] = None
|
||||
failures: List[TLSFailureResponse] = Field(default_factory=list)
|
||||
|
||||
|
||||
class TLSReportListResponse(BaseModel):
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
total_pages: int
|
||||
reports: List[TLSReportResponse]
|
||||
privacy: Dict[str, Any]
|
||||
|
||||
|
||||
class TLSReportUploadResponse(BaseModel):
|
||||
success: bool
|
||||
report_id: str
|
||||
policies_created: int
|
||||
policies_skipped: int
|
||||
duplicate: bool = False
|
||||
message: str
|
||||
privacy: Dict[str, Any]
|
||||
|
||||
|
||||
class TLSSummaryResponse(BaseModel):
|
||||
domain: Optional[str] = None
|
||||
days: int
|
||||
totals: Dict[str, Any]
|
||||
trends: List[Dict[str, Any]] = Field(default_factory=list)
|
||||
top_failures: List[Dict[str, Any]] = Field(default_factory=list)
|
||||
affected_domains: List[Dict[str, Any]] = Field(default_factory=list)
|
||||
privacy: Dict[str, Any]
|
||||
|
||||
|
||||
def _validate_upload(file: UploadFile, content: bytes) -> None:
|
||||
filename = file.filename or ""
|
||||
if not filename:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Filename is required")
|
||||
if len(content) == 0:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="File is empty")
|
||||
if len(content) > MAX_TLS_REPORT_SIZE:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail="File too large",
|
||||
)
|
||||
if not filename.lower().endswith((".json", ".json.gz", ".gzip", ".zip")):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Invalid file type. Upload a TLS report as .json, .json.gz, or .zip.",
|
||||
)
|
||||
|
||||
|
||||
def _filtered_tls_query(db: Session, *, domain: Optional[str] = None):
|
||||
query = db.query(TLSReport).options(
|
||||
selectinload(TLSReport.domain), selectinload(TLSReport.failures)
|
||||
)
|
||||
if domain:
|
||||
normalized = domain.lower().strip(".")
|
||||
query = query.outerjoin(Domain).filter(
|
||||
(Domain.name == normalized) | (TLSReport.policy_domain == normalized)
|
||||
)
|
||||
return query
|
||||
|
||||
|
||||
@router.post("/upload", response_model=TLSReportUploadResponse)
|
||||
async def upload_tls_report(
|
||||
file: UploadFile = File(...),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""Upload and store an SMTP TLS Reporting aggregate."""
|
||||
try:
|
||||
content = await file.read()
|
||||
_validate_upload(file, content)
|
||||
parsed = TLSReportParser.parse_file(content, file.filename or "")
|
||||
result = save_tls_report(db, parsed)
|
||||
db.commit()
|
||||
return TLSReportUploadResponse(
|
||||
success=True,
|
||||
report_id=parsed["report_id"],
|
||||
policies_created=result["created"],
|
||||
policies_skipped=result["skipped"],
|
||||
duplicate=result["created"] == 0 and result["skipped"] > 0,
|
||||
message=(
|
||||
"TLS report imported."
|
||||
if result["created"]
|
||||
else "TLS report had already been imported."
|
||||
),
|
||||
privacy=TLS_REPORT_PRIVACY_CONTROLS,
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=str(exc) or "Invalid TLS report format.",
|
||||
) from exc
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Unexpected TLS report upload failure for %s: %s", file.filename, exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Error processing TLS report.",
|
||||
) from exc
|
||||
|
||||
|
||||
@router.get("", response_model=TLSReportListResponse)
|
||||
async def list_tls_reports(
|
||||
domain: Optional[str] = Query(default=None),
|
||||
page: int = Query(default=1, ge=1),
|
||||
page_size: int = Query(default=50, ge=1, le=200),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""List stored SMTP TLS reports, newest first."""
|
||||
query = _filtered_tls_query(db, domain=domain)
|
||||
total = query.count()
|
||||
rows = (
|
||||
query.order_by(TLSReport.begin_date.desc().nullslast(), TLSReport.id.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
.all()
|
||||
)
|
||||
total_pages = (total + page_size - 1) // page_size if total else 0
|
||||
return TLSReportListResponse(
|
||||
total=total,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
total_pages=total_pages,
|
||||
reports=[TLSReportResponse(**tls_report_to_dict(row)) for row in rows],
|
||||
privacy=TLS_REPORT_PRIVACY_CONTROLS,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/summary", response_model=TLSSummaryResponse)
|
||||
async def tls_report_summary(
|
||||
domain: Optional[str] = Query(default=None),
|
||||
days: int = Query(default=30, ge=1, le=365),
|
||||
limit: int = Query(default=10, ge=1, le=50),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
):
|
||||
"""Summarize TLS reports into trends and top failure causes."""
|
||||
return TLSSummaryResponse(**summarize_tls_reports(db, domain=domain, days=days, limit=limit))
|
||||
@@ -0,0 +1,190 @@
|
||||
"""Webhook ingestion endpoints for inbound DMARC report emails."""
|
||||
|
||||
import base64
|
||||
import email
|
||||
import hmac
|
||||
import logging
|
||||
from email.header import decode_header
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_db
|
||||
from app.core.redaction import sanitize_for_log
|
||||
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
|
||||
from app.services.webhook_events import EVENT_REPORT_IMPORTED, enqueue_webhook_event
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class EmailWebhookPayload(BaseModel):
|
||||
"""Payload for JSON webhook delivery from an email worker."""
|
||||
|
||||
raw_email: str
|
||||
from_address: Optional[str] = None
|
||||
to_address: Optional[str] = None
|
||||
subject: Optional[str] = None
|
||||
|
||||
|
||||
def _decode_email_header(header: Optional[str]) -> str:
|
||||
"""Decode an RFC 2047 email header to display text."""
|
||||
if not header:
|
||||
return ""
|
||||
decoded_parts = []
|
||||
for text, encoding in decode_header(header):
|
||||
if isinstance(text, bytes):
|
||||
decoded_parts.append(text.decode(encoding or "utf-8", errors="replace"))
|
||||
else:
|
||||
decoded_parts.append(text)
|
||||
return " ".join(decoded_parts)
|
||||
|
||||
|
||||
def _is_dmarc_filename(filename: str) -> bool:
|
||||
lower = filename.lower()
|
||||
return lower.endswith((".xml", ".zip", ".gz", ".gzip"))
|
||||
|
||||
|
||||
def _require_webhook_secret(x_webhook_secret: Optional[str]) -> None:
|
||||
settings = get_settings()
|
||||
if not settings.WEBHOOK_SECRET:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Webhook ingestion is not configured.",
|
||||
)
|
||||
if not x_webhook_secret or not hmac.compare_digest(x_webhook_secret, settings.WEBHOOK_SECRET):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Invalid webhook secret.",
|
||||
)
|
||||
|
||||
|
||||
def _store_report(db: Session, store: ReportStore, report: Dict[str, Any]) -> str:
|
||||
domain = report.get("domain") or "unknown"
|
||||
report_id = report.get("report_id") or ""
|
||||
if report_id and (store.has_report(domain, report_id) or report_exists(db, domain, report_id)):
|
||||
return "duplicate"
|
||||
save_parsed_report(db, report)
|
||||
try:
|
||||
enqueue_webhook_event(
|
||||
db,
|
||||
event_type=EVENT_REPORT_IMPORTED,
|
||||
payload={
|
||||
"domain": domain,
|
||||
"report_id": report_id,
|
||||
"org_name": report.get("org_name"),
|
||||
"begin_date": report.get("begin_date"),
|
||||
"end_date": report.get("end_date"),
|
||||
"records": len(report.get("records") or []),
|
||||
},
|
||||
idempotency_key=f"{EVENT_REPORT_IMPORTED}:{domain}:{report_id or 'unknown'}",
|
||||
)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.warning("Failed to queue report-import webhook event: %s", sanitize_for_log(exc))
|
||||
store.add_report(report)
|
||||
return "imported"
|
||||
|
||||
|
||||
def _process_email_attachments(msg: email.message.Message, db: Session) -> Dict[str, Any]:
|
||||
store = ReportStore.get_instance()
|
||||
results: Dict[str, Any] = {
|
||||
"reports_found": 0,
|
||||
"imported": 0,
|
||||
"duplicates": 0,
|
||||
"errors": [],
|
||||
}
|
||||
|
||||
for part in msg.walk():
|
||||
if part.get_content_disposition() != "attachment":
|
||||
continue
|
||||
|
||||
filename = _decode_email_header(part.get_filename())
|
||||
if not filename or not _is_dmarc_filename(filename):
|
||||
continue
|
||||
|
||||
try:
|
||||
content = part.get_payload(decode=True)
|
||||
if not content:
|
||||
continue
|
||||
report = DMARCParser.parse_file(content, filename)
|
||||
outcome = _store_report(db, store, report)
|
||||
results["reports_found"] += 1
|
||||
if outcome == "duplicate":
|
||||
results["duplicates"] += 1
|
||||
else:
|
||||
results["imported"] += 1
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.warning(
|
||||
"Webhook failed to process DMARC attachment %s: %s",
|
||||
sanitize_for_log(filename),
|
||||
sanitize_for_log(exc),
|
||||
)
|
||||
results["errors"].append(filename)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def _subject_from_message(msg: email.message.Message, fallback: Optional[str] = None) -> str:
|
||||
return fallback or _decode_email_header(msg.get("Subject"))
|
||||
|
||||
|
||||
@router.post("/email")
|
||||
async def receive_email(
|
||||
payload: EmailWebhookPayload,
|
||||
x_webhook_secret: Optional[str] = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
) -> Dict[str, Any]:
|
||||
"""Receive a base64 encoded raw email from an email worker webhook."""
|
||||
_require_webhook_secret(x_webhook_secret)
|
||||
try:
|
||||
raw_email = base64.b64decode(payload.raw_email, validate=True)
|
||||
except (ValueError, TypeError) as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="raw_email must be valid base64.",
|
||||
) from exc
|
||||
|
||||
return _handle_raw_email(raw_email, db, subject=payload.subject)
|
||||
|
||||
|
||||
@router.post("/email/raw")
|
||||
async def receive_raw_email(
|
||||
request: Request,
|
||||
x_webhook_secret: Optional[str] = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
) -> Dict[str, Any]:
|
||||
"""Receive raw RFC 822 email bytes from an email worker webhook."""
|
||||
_require_webhook_secret(x_webhook_secret)
|
||||
return _handle_raw_email(await request.body(), db)
|
||||
|
||||
|
||||
def _handle_raw_email(
|
||||
raw_email: bytes,
|
||||
db: Session,
|
||||
*,
|
||||
subject: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
try:
|
||||
msg = email.message_from_bytes(raw_email)
|
||||
attachment_results = _process_email_attachments(msg, db)
|
||||
db.commit()
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
db.rollback()
|
||||
logger.warning("Webhook failed to process email: %s", sanitize_for_log(exc))
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Error processing email.",
|
||||
) from exc
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"subject": _subject_from_message(msg, subject),
|
||||
**attachment_results,
|
||||
}
|
||||
@@ -0,0 +1,284 @@
|
||||
"""Admin endpoints for outbound webhook event delivery."""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.security import require_admin_auth
|
||||
from app.models.webhook import WebhookDelivery, WebhookEndpoint
|
||||
from app.services.webhook_events import (
|
||||
SUPPORTED_EVENT_TYPES,
|
||||
create_webhook_endpoint,
|
||||
deliver_due_webhooks,
|
||||
delivery_to_dict,
|
||||
endpoint_to_dict,
|
||||
queue_test_webhook,
|
||||
update_webhook_endpoint,
|
||||
)
|
||||
from app.services.workspace_audit import changed_fields, record_workspace_audit_log
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class WebhookEndpointCreate(BaseModel):
|
||||
"""Create payload for outbound webhook endpoints."""
|
||||
|
||||
name: str
|
||||
url: str
|
||||
secret: Optional[str] = None
|
||||
event_types: List[str] = ["*"]
|
||||
enabled: bool = True
|
||||
max_attempts: int = 5
|
||||
timeout_seconds: int = 10
|
||||
|
||||
|
||||
class WebhookEndpointUpdate(BaseModel):
|
||||
"""Update payload for outbound webhook endpoints."""
|
||||
|
||||
name: Optional[str] = None
|
||||
url: Optional[str] = None
|
||||
secret: Optional[str] = None
|
||||
event_types: Optional[List[str]] = None
|
||||
enabled: Optional[bool] = None
|
||||
max_attempts: Optional[int] = None
|
||||
timeout_seconds: Optional[int] = None
|
||||
|
||||
|
||||
class WebhookEndpointResponse(BaseModel):
|
||||
"""API-safe webhook endpoint metadata."""
|
||||
|
||||
id: int
|
||||
name: str
|
||||
url: str
|
||||
event_types: List[str]
|
||||
enabled: bool
|
||||
max_attempts: int
|
||||
timeout_seconds: int
|
||||
created_at: Optional[str]
|
||||
updated_at: Optional[str]
|
||||
last_success_at: Optional[str]
|
||||
last_failure_at: Optional[str]
|
||||
failure_count: int
|
||||
secret_configured: bool
|
||||
url_encrypted: bool
|
||||
secret: Optional[str] = None
|
||||
|
||||
|
||||
class WebhookEndpointListResponse(BaseModel):
|
||||
"""List response for outbound webhook endpoints."""
|
||||
|
||||
endpoints: List[WebhookEndpointResponse]
|
||||
supported_event_types: List[str]
|
||||
|
||||
|
||||
class WebhookDeliveryResponse(BaseModel):
|
||||
"""API-safe webhook delivery metadata."""
|
||||
|
||||
id: int
|
||||
endpoint_id: int
|
||||
event_type: str
|
||||
idempotency_key: str
|
||||
status: str
|
||||
attempt_count: int
|
||||
max_attempts: int
|
||||
next_attempt_at: Optional[str]
|
||||
last_attempt_at: Optional[str]
|
||||
delivered_at: Optional[str]
|
||||
last_status_code: Optional[int]
|
||||
last_error: Optional[str]
|
||||
response_excerpt: Optional[str]
|
||||
created_at: Optional[str]
|
||||
updated_at: Optional[str]
|
||||
|
||||
|
||||
class WebhookDeliveryListResponse(BaseModel):
|
||||
"""List response for outbound webhook deliveries."""
|
||||
|
||||
deliveries: List[WebhookDeliveryResponse]
|
||||
|
||||
|
||||
class WebhookTestResponse(BaseModel):
|
||||
"""Response for a webhook test delivery."""
|
||||
|
||||
delivery: WebhookDeliveryResponse
|
||||
|
||||
|
||||
@router.get("", response_model=WebhookEndpointListResponse)
|
||||
async def list_webhook_endpoints(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Return configured outbound webhook endpoints."""
|
||||
endpoints = db.query(WebhookEndpoint).order_by(WebhookEndpoint.created_at.desc()).all()
|
||||
return {
|
||||
"endpoints": [endpoint_to_dict(endpoint) for endpoint in endpoints],
|
||||
"supported_event_types": SUPPORTED_EVENT_TYPES,
|
||||
}
|
||||
|
||||
|
||||
@router.post("", response_model=WebhookEndpointResponse)
|
||||
async def create_webhook(
|
||||
payload: WebhookEndpointCreate,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Create an outbound webhook endpoint."""
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db)
|
||||
try:
|
||||
endpoint, raw_secret = create_webhook_endpoint(db, **payload.model_dump())
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="webhook.created",
|
||||
entity_type="webhook_endpoint",
|
||||
entity_id=endpoint.id,
|
||||
entity_name=endpoint.name,
|
||||
details={"event_types": payload.event_types, "enabled": endpoint.enabled},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
commit=True,
|
||||
)
|
||||
body = endpoint_to_dict(endpoint)
|
||||
body["secret"] = raw_secret
|
||||
return body
|
||||
|
||||
|
||||
@router.put("/{endpoint_id}", response_model=WebhookEndpointResponse)
|
||||
async def update_webhook(
|
||||
endpoint_id: int,
|
||||
payload: WebhookEndpointUpdate,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Update an outbound webhook endpoint."""
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db)
|
||||
endpoint = db.query(WebhookEndpoint).filter(WebhookEndpoint.id == endpoint_id).first()
|
||||
if endpoint is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail="Webhook endpoint not found"
|
||||
)
|
||||
try:
|
||||
endpoint, raw_secret = update_webhook_endpoint(
|
||||
db,
|
||||
endpoint,
|
||||
**payload.model_dump(exclude_unset=True),
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="webhook.updated",
|
||||
entity_type="webhook_endpoint",
|
||||
entity_id=endpoint.id,
|
||||
entity_name=endpoint.name,
|
||||
details={"changed_fields": changed_fields(payload.model_dump(exclude_unset=True))},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
commit=True,
|
||||
)
|
||||
body = endpoint_to_dict(endpoint)
|
||||
body["secret"] = raw_secret
|
||||
return body
|
||||
|
||||
|
||||
@router.delete("/{endpoint_id}", response_model=WebhookEndpointResponse)
|
||||
async def disable_webhook(
|
||||
endpoint_id: int,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Disable a webhook endpoint without deleting delivery history."""
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db)
|
||||
endpoint = db.query(WebhookEndpoint).filter(WebhookEndpoint.id == endpoint_id).first()
|
||||
if endpoint is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail="Webhook endpoint not found"
|
||||
)
|
||||
endpoint.enabled = False
|
||||
db.commit()
|
||||
db.refresh(endpoint)
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="webhook.disabled",
|
||||
entity_type="webhook_endpoint",
|
||||
entity_id=endpoint.id,
|
||||
entity_name=endpoint.name,
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
commit=True,
|
||||
)
|
||||
body = endpoint_to_dict(endpoint)
|
||||
body["secret"] = None
|
||||
return body
|
||||
|
||||
|
||||
@router.get("/deliveries", response_model=WebhookDeliveryListResponse)
|
||||
async def list_webhook_deliveries(
|
||||
endpoint_id: Optional[int] = None,
|
||||
delivery_status: Optional[str] = Query(None, alias="status"),
|
||||
limit: int = 50,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Return recent outbound webhook deliveries."""
|
||||
query = db.query(WebhookDelivery)
|
||||
if endpoint_id is not None:
|
||||
query = query.filter(WebhookDelivery.endpoint_id == endpoint_id)
|
||||
if delivery_status:
|
||||
query = query.filter(WebhookDelivery.status == delivery_status)
|
||||
deliveries = (
|
||||
query.order_by(WebhookDelivery.created_at.desc(), WebhookDelivery.id.desc())
|
||||
.limit(max(1, min(limit, 200)))
|
||||
.all()
|
||||
)
|
||||
return {"deliveries": [delivery_to_dict(delivery) for delivery in deliveries]}
|
||||
|
||||
|
||||
@router.post("/{endpoint_id}/test", response_model=WebhookTestResponse)
|
||||
async def test_webhook(
|
||||
endpoint_id: int,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Queue and immediately attempt a test delivery for an endpoint."""
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db)
|
||||
try:
|
||||
delivery = queue_test_webhook(db, endpoint_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
|
||||
delivered = deliver_due_webhooks(db, endpoint_id=endpoint_id, limit=1)
|
||||
record_workspace_audit_log(
|
||||
db,
|
||||
workspace=workspace,
|
||||
action="webhook.tested",
|
||||
entity_type="webhook_endpoint",
|
||||
entity_id=endpoint_id,
|
||||
details={"delivery_id": delivery.id},
|
||||
auth_context=_auth,
|
||||
request=request,
|
||||
commit=True,
|
||||
)
|
||||
return {"delivery": delivery_to_dict(delivered[0] if delivered else delivery)}
|
||||
|
||||
|
||||
@router.post("/deliveries/process", response_model=WebhookDeliveryListResponse)
|
||||
async def process_due_webhooks(
|
||||
limit: int = 25,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: dict = Depends(require_admin_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""Attempt due pending webhook deliveries."""
|
||||
deliveries = deliver_due_webhooks(db, limit=max(1, min(limit, 100)))
|
||||
return {"deliveries": [delivery_to_dict(delivery) for delivery in deliveries]}
|
||||
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import logging
|
||||
import secrets
|
||||
from functools import lru_cache
|
||||
@@ -5,7 +6,7 @@ from typing import List, Optional, Union
|
||||
|
||||
# Try to import from pydantic_settings first (newer versions)
|
||||
try:
|
||||
from pydantic import EmailStr, validator
|
||||
from pydantic import EmailStr, validator # pylint: disable=ungrouped-imports
|
||||
from pydantic_settings import BaseSettings
|
||||
except ImportError:
|
||||
# Fall back to older pydantic version
|
||||
@@ -20,9 +21,12 @@ class Settings(BaseSettings):
|
||||
# Base
|
||||
PROJECT_NAME: str = "DMARQ"
|
||||
API_V1_STR: str = "/api/v1"
|
||||
ENVIRONMENT: str = "development"
|
||||
|
||||
# Database
|
||||
DATABASE_URL: str = "sqlite:///./dmarq.db"
|
||||
# Default to a sub-directory so the SQLite file lives in a location that
|
||||
# can be persisted via a Docker volume mount (e.g. /app/data).
|
||||
DATABASE_URL: str = "sqlite:///./data/dmarq.db"
|
||||
|
||||
# JWT Authentication
|
||||
SECRET_KEY: Optional[str] = None
|
||||
@@ -37,6 +41,8 @@ class Settings(BaseSettings):
|
||||
IMAP_PORT: int = 993
|
||||
IMAP_USERNAME: Optional[str] = None
|
||||
IMAP_PASSWORD: Optional[str] = None
|
||||
IMAP_FOLDER: str = "INBOX"
|
||||
DELETE_IMPORTED_EMAILS: bool = False
|
||||
|
||||
# Admin User
|
||||
FIRST_SUPERUSER: Optional[EmailStr] = None
|
||||
@@ -45,14 +51,81 @@ class Settings(BaseSettings):
|
||||
# Optional Cloudflare Integration
|
||||
CLOUDFLARE_API_TOKEN: Optional[str] = None
|
||||
CLOUDFLARE_ZONE_ID: Optional[str] = None
|
||||
WEBHOOK_SECRET: Optional[str] = None
|
||||
|
||||
# Admin API Key (optional)
|
||||
# If set, this key is used directly instead of generating a random one at startup.
|
||||
# Use: openssl rand -hex 32
|
||||
ADMIN_API_KEY: Optional[str] = None
|
||||
|
||||
# ── Authentication mode ───────────────────────────────────────────────────
|
||||
# Set AUTH_DISABLED=true to run without any authentication.
|
||||
# Every request is treated as an anonymous admin.
|
||||
#
|
||||
# ⚠️ Only use this for local development or deployments that are protected
|
||||
# by an external auth proxy (e.g. Authelia, OAuth2 Proxy, Traefik Forward Auth).
|
||||
# Never expose an AUTH_DISABLED instance directly to the internet.
|
||||
AUTH_DISABLED: bool = False
|
||||
ALLOW_AUTH_DISABLED_IN_PRODUCTION: bool = False
|
||||
|
||||
# ── Logto OIDC ────────────────────────────────────────────────────────────
|
||||
# Set these to enable Logto-based authentication.
|
||||
# LOGTO_ENDPOINT: the base URL of your Logto instance,
|
||||
# e.g. "https://your-tenant.logto.app" or a self-hosted URL.
|
||||
# LOGTO_APP_ID: the Client ID of the "Traditional Web" application in Logto.
|
||||
# LOGTO_APP_SECRET: the Client Secret of the same application.
|
||||
# LOGTO_REDIRECT_URI (optional): override the default callback URL.
|
||||
# Defaults to <base_url>/api/v1/auth/callback.
|
||||
# LOGTO_SKIP_SSL_VERIFY (optional): set to true only when connecting to a
|
||||
# self-hosted Logto endpoint with a self-signed certificate.
|
||||
# Defaults to false so TLS certificates are verified.
|
||||
LOGTO_ENDPOINT: Optional[str] = None
|
||||
LOGTO_APP_ID: Optional[str] = None
|
||||
LOGTO_APP_SECRET: Optional[str] = None
|
||||
LOGTO_REDIRECT_URI: Optional[str] = None
|
||||
LOGTO_SKIP_SSL_VERIFY: bool = False
|
||||
ALLOW_LOGTO_SKIP_SSL_VERIFY_IN_PRODUCTION: bool = False
|
||||
|
||||
@property
|
||||
def logto_configured(self) -> bool:
|
||||
"""Return True when the minimum Logto settings are present."""
|
||||
return bool(self.LOGTO_ENDPOINT and self.LOGTO_APP_ID and self.LOGTO_APP_SECRET)
|
||||
|
||||
@property
|
||||
def is_production(self) -> bool:
|
||||
"""Return True when the app is explicitly running in production mode."""
|
||||
return self.ENVIRONMENT.strip().lower() in {"prod", "production"}
|
||||
|
||||
@validator("ADMIN_API_KEY", pre=True, always=True)
|
||||
@classmethod
|
||||
def validate_admin_api_key(
|
||||
cls, v: Optional[str]
|
||||
) -> Optional[str]: # pylint: disable=no-self-argument
|
||||
"""Warn if ADMIN_API_KEY is set but too short."""
|
||||
if v is not None and len(v) < 32:
|
||||
logger.warning(
|
||||
"ADMIN_API_KEY is too short (%s characters). "
|
||||
"Recommended minimum is 32 characters for security. "
|
||||
"Generate a strong key with: openssl rand -hex 32",
|
||||
len(v),
|
||||
)
|
||||
return v or None
|
||||
|
||||
@validator("SECRET_KEY", pre=True, always=True)
|
||||
def validate_secret_key(cls, v: Optional[str]) -> str:
|
||||
def validate_secret_key( # pylint: disable=no-self-argument
|
||||
cls, v: Optional[str], values
|
||||
) -> str:
|
||||
"""Validate and generate SECRET_KEY if not provided."""
|
||||
# Default insecure key that should never be used
|
||||
DEFAULT_INSECURE_KEY = "CHANGE_THIS_TO_A_RANDOM_SECRET_IN_PRODUCTION"
|
||||
environment = str(values.get("ENVIRONMENT", "development")).strip().lower()
|
||||
is_production = environment in {"prod", "production"}
|
||||
|
||||
if v is None or v == "" or v == DEFAULT_INSECURE_KEY:
|
||||
if is_production:
|
||||
raise ValueError(
|
||||
"SECRET_KEY must be set to a stable random value when ENVIRONMENT=production."
|
||||
)
|
||||
# Generate a secure random key
|
||||
generated_key = secrets.token_hex(32)
|
||||
logger.warning(
|
||||
@@ -65,24 +138,37 @@ class Settings(BaseSettings):
|
||||
|
||||
# Check if key is too short
|
||||
if len(v) < 32:
|
||||
if is_production:
|
||||
raise ValueError(
|
||||
"SECRET_KEY must be at least 32 characters when ENVIRONMENT=production."
|
||||
)
|
||||
logger.warning(
|
||||
f"SECRET_KEY is too short ({len(v)} characters). "
|
||||
"Recommended minimum is 32 characters for security."
|
||||
"SECRET_KEY is too short (%s characters). "
|
||||
"Recommended minimum is 32 characters for security.",
|
||||
len(v),
|
||||
)
|
||||
|
||||
return v
|
||||
|
||||
@validator("BACKEND_CORS_ORIGINS", pre=True)
|
||||
def assemble_cors_origins(cls, v: Union[str, List[str]]) -> List[str]:
|
||||
if isinstance(v, str) and not v.startswith("["):
|
||||
return [i.strip() for i in v.split(",")]
|
||||
elif isinstance(v, (list, str)):
|
||||
def assemble_cors_origins( # pylint: disable=no-self-argument
|
||||
cls, v: Union[str, List[str]]
|
||||
) -> List[str]:
|
||||
if isinstance(v, str):
|
||||
v = v.strip()
|
||||
if not v:
|
||||
return []
|
||||
if v.startswith("["):
|
||||
return json.loads(v)
|
||||
return [i.strip() for i in v.split(",") if i.strip()]
|
||||
if isinstance(v, list):
|
||||
return v
|
||||
raise ValueError(v)
|
||||
|
||||
class Config:
|
||||
env_file = ".env"
|
||||
case_sensitive = True
|
||||
env_ignore_empty = True
|
||||
|
||||
|
||||
@lru_cache()
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
import base64
|
||||
import hashlib
|
||||
from functools import lru_cache
|
||||
from typing import Optional
|
||||
|
||||
from cryptography.fernet import Fernet, InvalidToken
|
||||
|
||||
from app.core.config import get_settings
|
||||
|
||||
ENCRYPTED_SECRET_PREFIX = "enc:v1:"
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _get_fernet() -> Fernet:
|
||||
"""Build a Fernet instance from the stable application secret key."""
|
||||
secret_key = get_settings().SECRET_KEY
|
||||
digest = hashlib.sha256(secret_key.encode("utf-8")).digest()
|
||||
return Fernet(base64.urlsafe_b64encode(digest))
|
||||
|
||||
|
||||
def is_encrypted_secret(value: Optional[str]) -> bool:
|
||||
return bool(value and value.startswith(ENCRYPTED_SECRET_PREFIX))
|
||||
|
||||
|
||||
def encrypt_secret(value: Optional[str]) -> Optional[str]:
|
||||
"""Encrypt a secret for database storage, preserving empty and encrypted values."""
|
||||
if value is None or value == "":
|
||||
return value
|
||||
if is_encrypted_secret(value):
|
||||
return value
|
||||
|
||||
token = _get_fernet().encrypt(value.encode("utf-8")).decode("ascii")
|
||||
return f"{ENCRYPTED_SECRET_PREFIX}{token}"
|
||||
|
||||
|
||||
def decrypt_secret(value: Optional[str]) -> Optional[str]:
|
||||
"""Return plaintext for encrypted values and legacy plaintext unchanged."""
|
||||
if value is None or value == "":
|
||||
return value
|
||||
if not is_encrypted_secret(value):
|
||||
return value
|
||||
|
||||
token = value[len(ENCRYPTED_SECRET_PREFIX) :]
|
||||
try:
|
||||
return _get_fernet().decrypt(token.encode("ascii")).decode("utf-8")
|
||||
except InvalidToken as exc:
|
||||
raise ValueError(
|
||||
"Stored credential could not be decrypted with the configured SECRET_KEY"
|
||||
) from exc
|
||||
@@ -1,14 +1,65 @@
|
||||
import os
|
||||
from typing import Generator
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
from app.core.config import get_settings
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.engine import make_url
|
||||
from sqlalchemy.ext.declarative import declarative_base
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from app.core.config import get_settings
|
||||
|
||||
_ASYNC_TO_SYNC_SCHEMES = {
|
||||
"postgresql+asyncpg": "postgresql+psycopg2",
|
||||
}
|
||||
|
||||
|
||||
def _make_sync_db_url(url: str) -> str:
|
||||
"""Return the synchronous-driver equivalent of *url*.
|
||||
|
||||
Kubernetes and docker-compose deployments sometimes configure DATABASE_URL
|
||||
with an async driver scheme (e.g. ``postgresql+asyncpg://``). Alembic and
|
||||
the synchronous SQLAlchemy engine used here require a sync driver, so we
|
||||
map known async schemes to their psycopg2 equivalents.
|
||||
|
||||
Only the scheme component of the URL is rewritten; all other parts
|
||||
(credentials, host, path, query) are left untouched.
|
||||
"""
|
||||
parsed = urlparse(url)
|
||||
sync_scheme = _ASYNC_TO_SYNC_SCHEMES.get(parsed.scheme)
|
||||
if sync_scheme is None:
|
||||
return url
|
||||
return urlunparse(parsed._replace(scheme=sync_scheme))
|
||||
|
||||
|
||||
def _ensure_sqlite_dir(url: str) -> None:
|
||||
"""Create the parent directory for a SQLite database file if needed.
|
||||
|
||||
For SQLite URLs (``sqlite:///relative/path`` or ``sqlite:////absolute/path``),
|
||||
the parent directory must exist before SQLAlchemy tries to open (or create)
|
||||
the file. This is a no-op for in-memory databases (``sqlite://``) and for
|
||||
non-SQLite URLs.
|
||||
"""
|
||||
sa_url = make_url(url)
|
||||
if not sa_url.drivername.startswith("sqlite"):
|
||||
return
|
||||
db_path = sa_url.database
|
||||
if not db_path or db_path == ":memory:":
|
||||
return # in-memory – nothing to create
|
||||
parent = os.path.dirname(db_path)
|
||||
if parent:
|
||||
os.makedirs(parent, exist_ok=True)
|
||||
|
||||
|
||||
settings = get_settings()
|
||||
|
||||
# Configure SQLAlchemy
|
||||
engine = create_engine(settings.DATABASE_URL, pool_pre_ping=True)
|
||||
_sync_url = _make_sync_db_url(settings.DATABASE_URL)
|
||||
|
||||
# Ensure the parent directory exists before SQLAlchemy tries to open the file
|
||||
_ensure_sqlite_dir(_sync_url)
|
||||
|
||||
# Configure SQLAlchemy (normalise async driver schemes to their sync equivalents)
|
||||
engine = create_engine(_sync_url, pool_pre_ping=True)
|
||||
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
||||
|
||||
# Create base class for SQLAlchemy models
|
||||
|
||||
@@ -0,0 +1,314 @@
|
||||
"""
|
||||
Logto OIDC integration helpers.
|
||||
|
||||
Provides:
|
||||
- ``CookieStorage`` – Logto SDK Storage adapter backed by HTTP cookies.
|
||||
- ``make_logto_client`` – Factory that builds a per-request LogtoClient.
|
||||
- ``create_session_token``/``decode_session_token`` – thin JWT helpers for the
|
||||
app-level session cookie (independent of Logto after the initial callback).
|
||||
- ``sync_logto_user`` – Upserts the local User shadow record from Logto claims.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import ssl
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Optional
|
||||
|
||||
import aiohttp
|
||||
from fastapi import Request, Response
|
||||
from jose import JWTError, jwt
|
||||
from logto import IdTokenClaims, LogtoClient, LogtoConfig, PersistKey, Storage, UserInfoScope
|
||||
from logto.models.oidc import OAuthScope
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.models.user import User
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
# ── Constants ────────────────────────────────────────────────────────────────
|
||||
|
||||
SESSION_COOKIE = "dmarq_session"
|
||||
|
||||
# Short-lived: only needed while the browser is being redirected to Logto and back.
|
||||
_SIGN_IN_SESSION_MAX_AGE = 600 # 10 minutes
|
||||
# The app-level session lasts 24 hours by default; the Logto ID-token has its own
|
||||
# expiry but we don't keep it in the browser beyond the callback request.
|
||||
_SESSION_MAX_AGE = 86_400 # 24 hours
|
||||
|
||||
|
||||
# ── SSL configuration for Logto SDK ──────────────────────────────────────────
|
||||
|
||||
|
||||
def _apply_logto_ssl_patch() -> None:
|
||||
"""
|
||||
If ``LOGTO_SKIP_SSL_VERIFY`` is ``True``, monkey-patch both
|
||||
``aiohttp.ClientSession`` and the ``PyJWKClient`` used by the Logto SDK so
|
||||
that every connection to the Logto OIDC endpoint skips SSL certificate
|
||||
verification.
|
||||
|
||||
Two patches are applied:
|
||||
|
||||
1. **aiohttp.ClientSession** – The Logto SDK creates its own
|
||||
``aiohttp.ClientSession`` objects internally (for the OIDC discovery
|
||||
document and token-endpoint requests) and provides no mechanism to
|
||||
inject an SSL context. Replacing the class at module level is the
|
||||
only way to propagate the setting without forking the SDK.
|
||||
|
||||
2. **PyJWKClient** inside ``logto.OidcCore`` – The Logto SDK uses
|
||||
``PyJWKClient`` (from PyJWT) to fetch and verify the JWKS for
|
||||
ID-token signature validation. ``PyJWKClient`` uses ``urllib``
|
||||
internally, *not* ``aiohttp``, so the first patch does not cover it.
|
||||
We replace the ``PyJWKClient`` reference in the ``logto.OidcCore``
|
||||
module so that every ``OidcCore`` instance gets a client that passes
|
||||
the non-verifying SSL context to ``urllib``.
|
||||
|
||||
**Scope note:** ``aiohttp`` is not used anywhere else in this application
|
||||
– only the Logto SDK pulls it in. If additional code in this repository
|
||||
starts using ``aiohttp`` directly, review whether those connections should
|
||||
also skip verification before enabling this setting.
|
||||
|
||||
.. warning::
|
||||
Disabling SSL verification removes protection against man-in-the-middle
|
||||
attacks. Only enable this when connecting to a Logto instance that uses
|
||||
a self-signed certificate that you control.
|
||||
"""
|
||||
if not settings.LOGTO_SKIP_SSL_VERIFY:
|
||||
return
|
||||
|
||||
logger.warning(
|
||||
"LOGTO_SKIP_SSL_VERIFY is enabled – SSL certificate verification for "
|
||||
"Logto OIDC connections is DISABLED. Use this only when your Logto "
|
||||
"instance uses a self-signed certificate. Never enable this in a "
|
||||
"production environment that faces the public internet."
|
||||
)
|
||||
|
||||
ssl_ctx = ssl.create_default_context()
|
||||
ssl_ctx.check_hostname = False
|
||||
ssl_ctx.verify_mode = ssl.CERT_NONE
|
||||
|
||||
# ── Patch 1: aiohttp.ClientSession ───────────────────────────────────────
|
||||
# Covers OIDC discovery-document and token-endpoint requests.
|
||||
|
||||
_OriginalClientSession = aiohttp.ClientSession
|
||||
|
||||
class _NoVerifyClientSession(_OriginalClientSession): # type: ignore[misc]
|
||||
"""``aiohttp.ClientSession`` subclass that disables SSL verification."""
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None: # type: ignore[override]
|
||||
if "connector" not in kwargs:
|
||||
kwargs["connector"] = aiohttp.TCPConnector(ssl=ssl_ctx)
|
||||
kwargs.setdefault("connector_owner", True)
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
aiohttp.ClientSession = _NoVerifyClientSession # type: ignore[assignment]
|
||||
|
||||
# ── Patch 2: PyJWKClient inside logto.OidcCore ────────────────────────────
|
||||
# Covers JWKS fetching for ID-token signature verification.
|
||||
# PyJWKClient uses urllib internally, so Patch 1 does not cover it.
|
||||
try:
|
||||
import logto.OidcCore as _oidc_module # noqa: PLC0415
|
||||
from jwt import PyJWKClient as _OrigPyJWKClient # noqa: PLC0415
|
||||
|
||||
class _NoVerifyPyJWKClient(_OrigPyJWKClient): # type: ignore[misc]
|
||||
"""``PyJWKClient`` subclass that injects a non-verifying SSL context."""
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None: # type: ignore[override]
|
||||
kwargs.setdefault("ssl_context", ssl_ctx)
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
_oidc_module.PyJWKClient = _NoVerifyPyJWKClient # type: ignore[attr-defined]
|
||||
except Exception as _exc: # pylint: disable=broad-exception-caught
|
||||
logger.warning(
|
||||
"Failed to patch PyJWKClient for LOGTO_SKIP_SSL_VERIFY: %s. "
|
||||
"JWKS fetching will still verify SSL certificates, which may cause "
|
||||
"ID-token verification to fail when using a self-signed certificate.",
|
||||
_exc,
|
||||
)
|
||||
|
||||
|
||||
_apply_logto_ssl_patch()
|
||||
|
||||
|
||||
# ── Cookie-backed Logto Storage ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class CookieStorage(Storage):
|
||||
"""
|
||||
Storage adapter for the Logto SDK that persists the OIDC session data
|
||||
(sign-in session, tokens) in HTTP-only cookies.
|
||||
|
||||
Usage::
|
||||
|
||||
storage = CookieStorage(request)
|
||||
client = make_logto_client(storage)
|
||||
url = await client.signIn(redirect_uri=…)
|
||||
# build a response, then:
|
||||
storage.apply_to_response(response)
|
||||
return response
|
||||
"""
|
||||
|
||||
_COOKIE_PREFIX = "logto_"
|
||||
|
||||
def __init__(self, request: Request) -> None:
|
||||
self._request = request
|
||||
# Pending writes/deletes – applied to the Response via apply_to_response().
|
||||
self._writes: dict[str, Optional[str]] = {}
|
||||
self._deletes: set[str] = set()
|
||||
|
||||
# ── Storage protocol ──────────────────────────────────────────────────────
|
||||
|
||||
def get(self, key: PersistKey) -> Optional[str]: # type: ignore[override]
|
||||
if key in self._writes:
|
||||
return self._writes[key]
|
||||
if key in self._deletes:
|
||||
return None
|
||||
return self._request.cookies.get(self._COOKIE_PREFIX + key)
|
||||
|
||||
def set(self, key: PersistKey, value: Optional[str]) -> None: # type: ignore[override]
|
||||
self._writes[key] = value
|
||||
self._deletes.discard(key)
|
||||
|
||||
def delete(self, key: PersistKey) -> None: # type: ignore[override]
|
||||
self._deletes.add(key)
|
||||
self._writes.pop(key, None)
|
||||
|
||||
# ── Response helper ───────────────────────────────────────────────────────
|
||||
|
||||
def apply_to_response(self, response: Response) -> None:
|
||||
"""Flush pending cookie mutations onto *response*."""
|
||||
for key, value in self._writes.items():
|
||||
if value is None:
|
||||
continue
|
||||
max_age = _SIGN_IN_SESSION_MAX_AGE if key == "signInSession" else _SESSION_MAX_AGE
|
||||
response.set_cookie(
|
||||
key=self._COOKIE_PREFIX + key,
|
||||
value=value,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
max_age=max_age,
|
||||
)
|
||||
for key in self._deletes:
|
||||
response.delete_cookie(
|
||||
key=self._COOKIE_PREFIX + key,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
)
|
||||
|
||||
def clear_all_logto_cookies(self, response: Response) -> None:
|
||||
"""Remove every Logto cookie (called after we've issued our own session)."""
|
||||
for key in ("signInSession", "idToken", "accessTokenMap", "refreshToken"):
|
||||
response.delete_cookie(
|
||||
key=self._COOKIE_PREFIX + key,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
)
|
||||
|
||||
|
||||
# ── LogtoClient factory ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def make_logto_client(storage: CookieStorage) -> LogtoClient:
|
||||
"""Return a per-request ``LogtoClient`` bound to *storage*."""
|
||||
return LogtoClient(
|
||||
LogtoConfig(
|
||||
endpoint=settings.LOGTO_ENDPOINT or "",
|
||||
appId=settings.LOGTO_APP_ID or "",
|
||||
appSecret=settings.LOGTO_APP_SECRET,
|
||||
scopes=[
|
||||
UserInfoScope.email,
|
||||
UserInfoScope.profile,
|
||||
OAuthScope.offlineAccess,
|
||||
],
|
||||
),
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
|
||||
# ── App-level session JWT (independent of Logto after first login) ────────────
|
||||
|
||||
|
||||
def create_session_token(user_id: int) -> str:
|
||||
"""Mint a signed HS256 JWT for *user_id* with a 24-hour lifetime."""
|
||||
payload = {
|
||||
"sub": str(user_id),
|
||||
"type": "dmarq_session",
|
||||
"exp": datetime.utcnow() + timedelta(seconds=_SESSION_MAX_AGE),
|
||||
}
|
||||
return jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM)
|
||||
|
||||
|
||||
def decode_session_token(token: str) -> Optional[int]:
|
||||
"""
|
||||
Validate *token* and return the user's local DB id.
|
||||
|
||||
Returns ``None`` on any error (expired, wrong type, bad signature, …).
|
||||
"""
|
||||
try:
|
||||
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
|
||||
if payload.get("type") != "dmarq_session":
|
||||
return None
|
||||
return int(payload["sub"])
|
||||
except (JWTError, ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
# ── Local user sync ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def sync_logto_user(claims: IdTokenClaims, db: Session) -> User:
|
||||
"""
|
||||
Upsert the local ``User`` shadow record from Logto ID-token claims.
|
||||
|
||||
Lookup order:
|
||||
1. Match on ``logto_id`` (``sub`` claim) – fastest, stable.
|
||||
2. Fall back to matching on email if the user was created before Logto
|
||||
integration and doesn't have a ``logto_id`` yet.
|
||||
3. Create a brand-new record if neither match.
|
||||
|
||||
All users are treated as admins (``is_superuser=True``) until RBAC is
|
||||
added in a future milestone.
|
||||
"""
|
||||
logto_id: str = claims.sub
|
||||
email: str = claims.email or f"{logto_id}@logto.local"
|
||||
|
||||
# 1. Try existing Logto-linked user
|
||||
user: Optional[User] = db.query(User).filter(User.logto_id == logto_id).first()
|
||||
|
||||
if user is None:
|
||||
# 2. Try to link a legacy user by email
|
||||
user = db.query(User).filter(User.email == email).first()
|
||||
if user is not None:
|
||||
user.logto_id = logto_id
|
||||
logger.info(
|
||||
"Linked existing user id=%d (%s) to Logto sub=%s",
|
||||
user.id,
|
||||
email,
|
||||
logto_id,
|
||||
)
|
||||
|
||||
if user is None:
|
||||
# 3. Create new user
|
||||
user = User(
|
||||
logto_id=logto_id,
|
||||
email=email,
|
||||
is_active=True,
|
||||
is_superuser=True,
|
||||
is_verified=bool(getattr(claims, "email_verified", False)),
|
||||
)
|
||||
db.add(user)
|
||||
db.flush() # populate user.id before commit
|
||||
logger.info("Created new user id=%d from Logto sub=%s (%s)", user.id, logto_id, email)
|
||||
|
||||
# Always refresh profile from latest claims
|
||||
user.full_name = getattr(claims, "name", None) or user.full_name
|
||||
user.username = getattr(claims, "username", None) or user.username
|
||||
user.picture = getattr(claims, "picture", None) or user.picture
|
||||
user.updated_at = datetime.utcnow()
|
||||
|
||||
db.commit()
|
||||
db.refresh(user)
|
||||
return user
|
||||
@@ -0,0 +1,30 @@
|
||||
"""
|
||||
Utilities for making diagnostic text safe to store or log.
|
||||
"""
|
||||
|
||||
import re
|
||||
|
||||
SENSITIVE_VALUE = "**redacted**"
|
||||
|
||||
_SENSITIVE_KEY_PATTERN = re.compile(
|
||||
r"(?i)([\"']?\b(?:access_token|api_key|apikey|bearer|client_secret|"
|
||||
r"gmail_client_secret|id_token|passwd|password|refresh_token|secret|token)\b[\"']?"
|
||||
r"\s*[:=]\s*[\"']?)([^\"'\s,;&}]+)([\"']?)"
|
||||
)
|
||||
_BEARER_TOKEN_PATTERN = re.compile(r"(?i)\b(bearer)\s+([A-Za-z0-9._~+/=-]{8,})")
|
||||
_AUTHORIZATION_BEARER_PATTERN = re.compile(
|
||||
r"(?i)\b(authorization\s*[:=]\s*bearer)\s+([A-Za-z0-9._~+/=-]{8,})"
|
||||
)
|
||||
|
||||
|
||||
def sanitize_for_log(value: object) -> str:
|
||||
"""Remove CR/LF characters from a value to prevent log injection attacks."""
|
||||
return str(value).replace("\r", "").replace("\n", " ")
|
||||
|
||||
|
||||
def redact_sensitive_text(value: object) -> str:
|
||||
"""Sanitize text and redact common secret-bearing key/value fragments."""
|
||||
text = sanitize_for_log(value)
|
||||
text = _AUTHORIZATION_BEARER_PATTERN.sub(r"\1 " + SENSITIVE_VALUE, text)
|
||||
text = _SENSITIVE_KEY_PATTERN.sub(r"\1" + SENSITIVE_VALUE + r"\3", text)
|
||||
return _BEARER_TOKEN_PATTERN.sub(r"\1 " + SENSITIVE_VALUE, text)
|
||||
@@ -2,13 +2,17 @@ import logging
|
||||
import os
|
||||
import secrets
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Optional, Union
|
||||
from typing import Any, Callable, Optional, Union
|
||||
|
||||
from app.core.config import get_settings
|
||||
from fastapi import HTTPException, Security, status
|
||||
from fastapi import Depends, HTTPException, Request, Security, status
|
||||
from fastapi.security import APIKeyHeader, HTTPAuthorizationCredentials, HTTPBearer
|
||||
from jose import JWTError, jwt
|
||||
from passlib.context import CryptContext
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.database import get_db
|
||||
from app.services.api_tokens import find_api_token, parse_scopes, record_api_token_use
|
||||
|
||||
settings = get_settings()
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -75,7 +79,7 @@ def add_api_key(api_key: str) -> bool:
|
||||
if api_key in _api_keys:
|
||||
return False
|
||||
_api_keys.add(api_key)
|
||||
logger.info(f"API key added (ends with: ...{api_key[-8:]})")
|
||||
logger.info("API key added (length: %d chars)", len(api_key))
|
||||
return True
|
||||
|
||||
|
||||
@@ -92,7 +96,7 @@ def verify_api_key(api_key: str) -> bool:
|
||||
return api_key in _api_keys
|
||||
|
||||
|
||||
async def get_api_key(api_key_header: Optional[str] = Security(api_key_header)) -> str:
|
||||
async def get_api_key(api_key_value: Optional[str] = Security(api_key_header)) -> str:
|
||||
"""
|
||||
Dependency to verify API key authentication.
|
||||
|
||||
@@ -105,24 +109,23 @@ async def get_api_key(api_key_header: Optional[str] = Security(api_key_header))
|
||||
Raises:
|
||||
HTTPException: If API key is missing or invalid
|
||||
"""
|
||||
if not api_key_header:
|
||||
if not api_key_value:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Missing API key",
|
||||
headers={"WWW-Authenticate": "ApiKey"},
|
||||
)
|
||||
|
||||
if not verify_api_key(api_key_header):
|
||||
logger.warning(
|
||||
f"Invalid API key attempt: ...{api_key_header[-8:] if len(api_key_header) >= 8 else 'invalid'}"
|
||||
)
|
||||
if not verify_api_key(api_key_value):
|
||||
suffix = api_key_value[-8:] if len(api_key_value) >= 8 else "invalid"
|
||||
logger.warning("Invalid API key attempt: ...%s", suffix)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Invalid API key",
|
||||
headers={"WWW-Authenticate": "ApiKey"},
|
||||
)
|
||||
|
||||
return api_key_header
|
||||
return api_key_value
|
||||
|
||||
|
||||
async def verify_token(
|
||||
@@ -153,55 +156,116 @@ async def verify_token(
|
||||
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
|
||||
return payload
|
||||
except JWTError as e:
|
||||
logger.warning(f"Invalid JWT token: {str(e)}")
|
||||
logger.warning("Invalid JWT token: %s", str(e))
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Invalid authentication token",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
) from e
|
||||
|
||||
|
||||
async def require_admin_auth(
|
||||
request: Request,
|
||||
api_key: Optional[str] = Security(api_key_header),
|
||||
bearer: Optional[HTTPAuthorizationCredentials] = Security(security_bearer),
|
||||
) -> dict:
|
||||
"""
|
||||
Dependency to require either API key or JWT token authentication for admin endpoints.
|
||||
Dependency to require authentication for admin/API endpoints.
|
||||
|
||||
Checks API key first, then falls back to JWT token.
|
||||
Accepts (in priority order):
|
||||
1. ``AUTH_DISABLED=true`` env var – passes through with a synthetic context.
|
||||
2. ``dmarq_session`` cookie – set after a successful Logto login.
|
||||
3. ``X-API-Key`` header – static admin key for programmatic access.
|
||||
4. ``Authorization: Bearer <token>`` header – app-issued JWT.
|
||||
|
||||
Args:
|
||||
api_key: Optional API key from X-API-Key header
|
||||
bearer: Optional JWT token from Authorization header
|
||||
|
||||
Returns:
|
||||
Authentication context (api_key or token payload)
|
||||
|
||||
Raises:
|
||||
HTTPException: If no valid authentication is provided
|
||||
Returns an authentication context dict describing how the request was
|
||||
authenticated. Raises ``HTTP 401`` when no valid credential is present.
|
||||
"""
|
||||
# Try API key first
|
||||
if api_key and verify_api_key(api_key):
|
||||
return {"auth_type": "api_key", "api_key": api_key}
|
||||
# 0. Auth globally disabled
|
||||
if settings.AUTH_DISABLED:
|
||||
return {"auth_type": "disabled"}
|
||||
|
||||
# Try JWT token
|
||||
# 1. Session cookie (Logto-backed app session)
|
||||
from app.core.logto import SESSION_COOKIE, decode_session_token # local import
|
||||
|
||||
session_token = request.cookies.get(SESSION_COOKIE)
|
||||
if session_token:
|
||||
user_id = decode_session_token(session_token)
|
||||
if user_id is not None:
|
||||
return {"auth_type": "session", "user_id": user_id}
|
||||
|
||||
# 2. Static admin API key
|
||||
if api_key and verify_api_key(api_key):
|
||||
return {"auth_type": "api_key"}
|
||||
|
||||
# 3. Bearer JWT (app-issued; also covers Bearer tokens set by older clients)
|
||||
if bearer:
|
||||
from app.core.logto import decode_session_token as _dec # local import
|
||||
|
||||
user_id = _dec(bearer.credentials)
|
||||
if user_id is not None:
|
||||
return {"auth_type": "bearer", "user_id": user_id}
|
||||
|
||||
# Fallback: legacy python-jose JWT (pre-Logto API keys / CI tokens)
|
||||
try:
|
||||
payload = jwt.decode(
|
||||
bearer.credentials, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]
|
||||
)
|
||||
return {"auth_type": "jwt", "payload": payload}
|
||||
except JWTError as e:
|
||||
logger.warning(f"Invalid JWT token: {str(e)}")
|
||||
logger.warning("Invalid Bearer JWT: %s", str(e))
|
||||
|
||||
# No valid authentication provided
|
||||
# No valid authentication
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Authentication required. Provide either X-API-Key header or Bearer token.",
|
||||
detail="Authentication required. Provide a session cookie, X-API-Key header, or Bearer token.",
|
||||
headers={"WWW-Authenticate": "ApiKey, Bearer"},
|
||||
)
|
||||
|
||||
|
||||
def require_api_token_scope(required_scope: str) -> Callable:
|
||||
"""Build a dependency that requires a scoped persistent API token."""
|
||||
|
||||
async def _require_api_token_scope(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
api_key: Optional[str] = Security(api_key_header),
|
||||
) -> dict:
|
||||
if not api_key:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Missing API token",
|
||||
headers={"WWW-Authenticate": "ApiKey"},
|
||||
)
|
||||
|
||||
token = find_api_token(db, api_key)
|
||||
if token is None:
|
||||
logger.warning("Invalid public API token attempt")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Invalid API token",
|
||||
headers={"WWW-Authenticate": "ApiKey"},
|
||||
)
|
||||
|
||||
scopes = parse_scopes(token.scopes)
|
||||
if required_scope not in scopes:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=f"API token requires scope: {required_scope}",
|
||||
)
|
||||
|
||||
client_host = request.client.host if request.client else None
|
||||
record_api_token_use(db, token, ip_address=client_host)
|
||||
return {
|
||||
"auth_type": "api_token",
|
||||
"token_id": token.id,
|
||||
"token_name": token.name,
|
||||
"scopes": sorted(scopes),
|
||||
}
|
||||
|
||||
return _require_api_token_scope
|
||||
|
||||
|
||||
def create_access_token(subject: Union[str, Any], expires_delta: timedelta = None) -> str:
|
||||
"""
|
||||
Create a JWT access token for authentication
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
"""
|
||||
Startup configuration checks for production deployments.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy.engine import make_url
|
||||
|
||||
from app.core.config import Settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class StartupConfigurationError(RuntimeError):
|
||||
"""Raised when production configuration is unsafe enough to block startup."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StartupCheckResult:
|
||||
"""Result of validating startup configuration."""
|
||||
|
||||
errors: tuple[str, ...]
|
||||
warnings: tuple[str, ...]
|
||||
|
||||
@property
|
||||
def ok(self) -> bool:
|
||||
return not self.errors
|
||||
|
||||
|
||||
def _uses_sqlite(database_url: str) -> bool:
|
||||
try:
|
||||
return make_url(database_url).drivername.startswith("sqlite")
|
||||
except Exception: # pylint: disable=broad-exception-caught
|
||||
return False
|
||||
|
||||
|
||||
def validate_startup_configuration(settings: Settings) -> StartupCheckResult:
|
||||
"""Return production startup errors and warnings for the provided settings."""
|
||||
errors: list[str] = []
|
||||
warnings: list[str] = []
|
||||
|
||||
if not settings.is_production:
|
||||
return StartupCheckResult(errors=(), warnings=())
|
||||
|
||||
if settings.AUTH_DISABLED and not settings.ALLOW_AUTH_DISABLED_IN_PRODUCTION:
|
||||
errors.append(
|
||||
"AUTH_DISABLED=true is not allowed in production unless "
|
||||
"ALLOW_AUTH_DISABLED_IN_PRODUCTION=true is also set."
|
||||
)
|
||||
|
||||
if (
|
||||
settings.LOGTO_SKIP_SSL_VERIFY
|
||||
and not settings.ALLOW_LOGTO_SKIP_SSL_VERIFY_IN_PRODUCTION
|
||||
):
|
||||
errors.append(
|
||||
"LOGTO_SKIP_SSL_VERIFY=true is not allowed in production unless "
|
||||
"ALLOW_LOGTO_SKIP_SSL_VERIFY_IN_PRODUCTION=true is also set."
|
||||
)
|
||||
|
||||
if not settings.AUTH_DISABLED and not settings.ADMIN_API_KEY and not settings.logto_configured:
|
||||
errors.append(
|
||||
"Production startup requires Logto settings or ADMIN_API_KEY. "
|
||||
"Set LOGTO_ENDPOINT, LOGTO_APP_ID, and LOGTO_APP_SECRET, or set ADMIN_API_KEY."
|
||||
)
|
||||
|
||||
if settings.ADMIN_API_KEY and len(settings.ADMIN_API_KEY) < 32:
|
||||
errors.append("ADMIN_API_KEY must be at least 32 characters in production.")
|
||||
|
||||
if settings.SECRET_KEY is None or len(settings.SECRET_KEY) < 32:
|
||||
errors.append("SECRET_KEY must be at least 32 characters in production.")
|
||||
|
||||
if _uses_sqlite(settings.DATABASE_URL):
|
||||
warnings.append(
|
||||
"DATABASE_URL uses SQLite in production. This is supported for small "
|
||||
"single-node deployments, but PostgreSQL is recommended for durable production use."
|
||||
)
|
||||
|
||||
return StartupCheckResult(errors=tuple(errors), warnings=tuple(warnings))
|
||||
|
||||
|
||||
def run_startup_checks(settings: Settings) -> StartupCheckResult:
|
||||
"""Log startup validation results and raise when production config is unsafe."""
|
||||
result = validate_startup_configuration(settings)
|
||||
for warning in result.warnings:
|
||||
logger.warning("Startup configuration warning: %s", warning)
|
||||
|
||||
if result.errors:
|
||||
for error in result.errors:
|
||||
logger.error("Startup configuration error: %s", error)
|
||||
raise StartupConfigurationError(
|
||||
"Unsafe production configuration: " + " ".join(result.errors)
|
||||
)
|
||||
|
||||
return result
|
||||
+771
-118
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,105 @@
|
||||
"""
|
||||
Authentication redirect middleware.
|
||||
|
||||
Intercepts browser requests for protected HTML pages and redirects
|
||||
unauthenticated visitors to ``/login`` (or ``/setup`` if Logto is not yet
|
||||
configured).
|
||||
|
||||
API routes (``/api/…``) are intentionally left to handle their own 401
|
||||
responses so that programmatic clients are not broken.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import RedirectResponse, Response
|
||||
from starlette.types import ASGIApp
|
||||
|
||||
from app.core.logto import SESSION_COOKIE, decode_session_token
|
||||
|
||||
# Paths that are always publicly accessible
|
||||
_PUBLIC_PATHS: frozenset[str] = frozenset(
|
||||
{
|
||||
"/login",
|
||||
"/setup",
|
||||
"/health",
|
||||
"/healthz",
|
||||
}
|
||||
)
|
||||
|
||||
# Request path prefixes that bypass auth checks
|
||||
_PUBLIC_PREFIXES: tuple[str, ...] = (
|
||||
"/api/",
|
||||
"/static/",
|
||||
"/docs",
|
||||
"/redoc",
|
||||
"/openapi",
|
||||
)
|
||||
|
||||
# File extensions for static assets that are always publicly accessible
|
||||
_STATIC_EXTENSIONS: tuple[str, ...] = (
|
||||
".ico",
|
||||
".png",
|
||||
".jpg",
|
||||
".jpeg",
|
||||
".gif",
|
||||
".svg",
|
||||
".webp",
|
||||
".css",
|
||||
".js",
|
||||
".woff",
|
||||
".woff2",
|
||||
".ttf",
|
||||
".eot",
|
||||
".map",
|
||||
)
|
||||
|
||||
|
||||
class AuthRedirectMiddleware(BaseHTTPMiddleware):
|
||||
"""
|
||||
Redirect unauthenticated browser requests to the appropriate page.
|
||||
|
||||
Decision tree
|
||||
-------------
|
||||
1. Path is public → pass through.
|
||||
2. Session cookie present and valid → pass through.
|
||||
3. Logto not configured → redirect to ``/setup``.
|
||||
4. Otherwise → redirect to ``/login?next=<original_path>``.
|
||||
"""
|
||||
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
super().__init__(app)
|
||||
|
||||
async def dispatch(self, request: Request, call_next) -> Response: # type: ignore[override]
|
||||
path = request.url.path
|
||||
|
||||
# ── 0. Auth disabled globally ─────────────────────────────────────────
|
||||
from app.core.config import get_settings # local import avoids circular dep
|
||||
|
||||
cfg = get_settings()
|
||||
if cfg.AUTH_DISABLED:
|
||||
return await call_next(request)
|
||||
|
||||
# ── 1. Public paths & prefixes ────────────────────────────────────────
|
||||
if path in _PUBLIC_PATHS:
|
||||
return await call_next(request)
|
||||
if any(path.startswith(p) for p in _PUBLIC_PREFIXES):
|
||||
return await call_next(request)
|
||||
if any(path.endswith(ext) for ext in _STATIC_EXTENSIONS):
|
||||
return await call_next(request)
|
||||
|
||||
# ── 2. Valid session cookie ───────────────────────────────────────────
|
||||
token = request.cookies.get(SESSION_COOKIE)
|
||||
if token and decode_session_token(token) is not None:
|
||||
return await call_next(request)
|
||||
|
||||
# ── 3. Logto not configured ───────────────────────────────────────────
|
||||
if not cfg.logto_configured:
|
||||
return RedirectResponse(url="/setup", status_code=302)
|
||||
|
||||
# ── 4. Redirect to login ──────────────────────────────────────────────
|
||||
next_path = request.url.path
|
||||
if request.url.query:
|
||||
next_path = f"{next_path}?{request.url.query}"
|
||||
return RedirectResponse(url=f"/login?next={next_path}", status_code=302)
|
||||
@@ -53,19 +53,14 @@ class SecurityHeadersMiddleware(BaseHTTPMiddleware):
|
||||
# Content Security Policy (CSP)
|
||||
# Restricts sources of content that can be loaded
|
||||
#
|
||||
# SECURITY TODO: Current CSP includes 'unsafe-inline' and 'unsafe-eval' which
|
||||
# weaken XSS protection. To remove these:
|
||||
# SECURITY TODO: Current CSP includes 'unsafe-inline' which weakens
|
||||
# XSS protection. To remove it:
|
||||
#
|
||||
# For script-src 'unsafe-inline':
|
||||
# 1. Move all inline <script> tags from templates to external .js files
|
||||
# 2. OR implement CSP nonces for inline scripts (requires template changes)
|
||||
# 3. Convert any inline event handlers (onclick, etc.) to addEventListener
|
||||
#
|
||||
# For script-src 'unsafe-eval':
|
||||
# 1. Verify no code uses eval(), Function(), setTimeout/setInterval with strings
|
||||
# 2. If using libraries that require eval, consider alternatives
|
||||
# 3. Current scan shows no eval usage - can likely remove this directive
|
||||
#
|
||||
# For style-src 'unsafe-inline':
|
||||
# 1. Move inline styles to CSS files or use style tags with nonces
|
||||
# 2. Replace style="" attributes with CSS classes
|
||||
@@ -78,11 +73,14 @@ class SecurityHeadersMiddleware(BaseHTTPMiddleware):
|
||||
# See: https://developer.mozilla.org/en-US/docs/Web/HTTP/CSP
|
||||
csp_directives = [
|
||||
"default-src 'self'",
|
||||
# TODO: Remove 'unsafe-inline' - requires moving inline scripts to external files
|
||||
# TODO: Remove 'unsafe-eval' - no eval usage detected, safe to remove after testing
|
||||
"script-src 'self' 'unsafe-inline' 'unsafe-eval' https://cdn.tailwindcss.com https://cdn.jsdelivr.net",
|
||||
# TODO: Remove 'unsafe-inline' - requires moving inline styles to CSS or using nonces
|
||||
"style-src 'self' 'unsafe-inline' https://fonts.googleapis.com https://cdn.jsdelivr.net",
|
||||
# TODO: Remove 'unsafe-inline' and 'unsafe-eval' - requires moving inline
|
||||
# scripts to external files and replacing the standard Alpine CDN build
|
||||
# with the CSP-compatible build. # pylint: disable=fixme
|
||||
"script-src 'self' 'unsafe-inline' 'unsafe-eval'"
|
||||
" https://cdn.tailwindcss.com https://cdn.jsdelivr.net",
|
||||
# TODO: Remove 'unsafe-inline' - requires moving inline styles to CSS or using nonces # pylint: disable=fixme
|
||||
"style-src 'self' 'unsafe-inline' https://fonts.googleapis.com"
|
||||
" https://cdn.jsdelivr.net",
|
||||
"font-src 'self' https://fonts.gstatic.com",
|
||||
"img-src 'self' data: https:",
|
||||
"connect-src 'self'",
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, Index, Integer, String, Text
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class AlertHistory(Base):
|
||||
"""Persisted alert lifecycle record."""
|
||||
|
||||
__tablename__ = "alert_history"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
fingerprint = Column(String(64), unique=True, nullable=False, index=True)
|
||||
rule = Column(String, nullable=False, index=True)
|
||||
severity = Column(String, nullable=False, index=True)
|
||||
domain = Column(String, nullable=True, index=True)
|
||||
title = Column(String, nullable=False)
|
||||
detail = Column(Text, nullable=False)
|
||||
payload = Column(Text, nullable=True)
|
||||
observed_count = Column(Integer, nullable=False, default=1)
|
||||
is_active = Column(Boolean, nullable=False, default=True, index=True)
|
||||
first_seen_at = Column(DateTime, default=datetime.utcnow, nullable=False, index=True)
|
||||
last_seen_at = Column(DateTime, default=datetime.utcnow, nullable=False, index=True)
|
||||
resolved_at = Column(DateTime, nullable=True, index=True)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_alert_history_active_last_seen", "is_active", "last_seen_at"),
|
||||
Index("ix_alert_history_rule_domain", "rule", "domain"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<AlertHistory {self.rule} active={self.is_active}>"
|
||||
|
||||
|
||||
class AlertConfigurationAudit(Base):
|
||||
"""Audit trail for notification and alert-rule configuration changes."""
|
||||
|
||||
__tablename__ = "alert_configuration_audit"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
key = Column(String(100), nullable=False, index=True)
|
||||
old_value = Column(Text, nullable=True)
|
||||
new_value = Column(Text, nullable=True)
|
||||
changed_by = Column(String(100), nullable=True, index=True)
|
||||
auth_type = Column(String(50), nullable=True)
|
||||
changed_at = Column(DateTime, default=datetime.utcnow, nullable=False, index=True)
|
||||
|
||||
__table_args__ = (Index("ix_alert_configuration_audit_key_changed_at", "key", "changed_at"),)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<AlertConfigurationAudit {self.key}>"
|
||||
@@ -0,0 +1,32 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, Index, Integer, String, Text
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class APIToken(Base):
|
||||
"""Scoped API token for stable automation access."""
|
||||
|
||||
__tablename__ = "api_tokens"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
name = Column(String(120), nullable=False)
|
||||
key_hash = Column(String(255), unique=True, nullable=False, index=True)
|
||||
key_prefix = Column(String(16), nullable=False, index=True)
|
||||
scopes = Column(Text, nullable=False)
|
||||
active = Column(Boolean, default=True, nullable=False, index=True)
|
||||
created_at = Column(DateTime, default=datetime.utcnow, nullable=False, index=True)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
revoked_at = Column(DateTime, nullable=True, index=True)
|
||||
last_used_at = Column(DateTime, nullable=True, index=True)
|
||||
last_used_ip = Column(String(64), nullable=True)
|
||||
usage_count = Column(Integer, default=0, nullable=False)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_api_tokens_active_scope", "active", "scopes"),
|
||||
Index("ix_api_tokens_last_used", "last_used_at"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<APIToken {self.name} active={self.active}>"
|
||||
@@ -0,0 +1,93 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, Index, Integer, String, Text, UniqueConstraint
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
def _utcnow_naive() -> datetime:
|
||||
return datetime.now(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
|
||||
class DNSCache(Base):
|
||||
"""Cached DNS authentication result for a domain and selector set."""
|
||||
|
||||
__tablename__ = "dns_cache"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
domain = Column(String, nullable=False, index=True)
|
||||
provider = Column(String, nullable=False, index=True)
|
||||
selectors_key = Column(String(64), nullable=False, index=True)
|
||||
result_json = Column(Text, nullable=False)
|
||||
checked_at = Column(DateTime, default=_utcnow_naive, nullable=False, index=True)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint("domain", "provider", "selectors_key", name="uq_dns_cache_lookup"),
|
||||
Index("ix_dns_cache_domain_checked", "domain", "checked_at"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<DNSCache {self.domain} provider={self.provider}>"
|
||||
|
||||
|
||||
class DNSRecordSnapshot(Base):
|
||||
"""Last observed DNS record state for provider-backed DNS integrations."""
|
||||
|
||||
__tablename__ = "dns_record_snapshots"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
domain = Column(String, nullable=False, index=True)
|
||||
provider = Column(String, nullable=False, index=True)
|
||||
zone_id = Column(String, nullable=True, index=True)
|
||||
record_key = Column(String(128), nullable=False, index=True)
|
||||
record_id = Column(String, nullable=True, index=True)
|
||||
record_type = Column(String(20), nullable=False, index=True)
|
||||
record_name = Column(String, nullable=False, index=True)
|
||||
content = Column(Text, nullable=True)
|
||||
proxied = Column(Boolean, nullable=True)
|
||||
ttl = Column(Integer, nullable=True)
|
||||
record_hash = Column(String(64), nullable=False, index=True)
|
||||
active = Column(Boolean, default=True, nullable=False, index=True)
|
||||
first_seen_at = Column(DateTime, default=_utcnow_naive, nullable=False, index=True)
|
||||
last_seen_at = Column(DateTime, default=_utcnow_naive, nullable=False, index=True)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
"domain",
|
||||
"provider",
|
||||
"record_key",
|
||||
name="uq_dns_record_snapshot_lookup",
|
||||
),
|
||||
Index("ix_dns_record_snapshots_domain_active", "domain", "active"),
|
||||
Index("ix_dns_record_snapshots_domain_seen", "domain", "last_seen_at"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<DNSRecordSnapshot {self.domain} {self.record_type} {self.record_name}>"
|
||||
|
||||
|
||||
class DNSRecordChange(Base):
|
||||
"""Append-only DNS record change event detected during provider sync."""
|
||||
|
||||
__tablename__ = "dns_record_changes"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
domain = Column(String, nullable=False, index=True)
|
||||
provider = Column(String, nullable=False, index=True)
|
||||
zone_id = Column(String, nullable=True, index=True)
|
||||
record_key = Column(String(128), nullable=False, index=True)
|
||||
record_id = Column(String, nullable=True, index=True)
|
||||
record_type = Column(String(20), nullable=False, index=True)
|
||||
record_name = Column(String, nullable=False, index=True)
|
||||
change_type = Column(String(20), nullable=False, index=True)
|
||||
previous_content = Column(Text, nullable=True)
|
||||
current_content = Column(Text, nullable=True)
|
||||
observed_at = Column(DateTime, default=_utcnow_naive, nullable=False, index=True)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_dns_record_changes_domain_observed", "domain", "observed_at"),
|
||||
Index("ix_dns_record_changes_record_observed", "record_key", "observed_at"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<DNSRecordChange {self.domain} {self.change_type} {self.record_name}>"
|
||||
@@ -1,9 +1,10 @@
|
||||
from datetime import datetime
|
||||
|
||||
from app.core.database import Base
|
||||
from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Index, Integer, String, Text
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class Domain(Base):
|
||||
"""Domain model representing a monitored domain"""
|
||||
@@ -11,12 +12,13 @@ class Domain(Base):
|
||||
__tablename__ = "domains"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
workspace_id = Column(Integer, ForeignKey("workspaces.id"), nullable=True, index=True)
|
||||
name = Column(String, unique=True, index=True, nullable=False)
|
||||
description = Column(Text, nullable=True)
|
||||
active = Column(Boolean, default=True, index=True)
|
||||
|
||||
# DMARC policy information
|
||||
dmarc_policy = Column(String, nullable=True, index=True)
|
||||
dmarc_policy = Column(String, nullable=True)
|
||||
spf_record = Column(String, nullable=True)
|
||||
dkim_selectors = Column(String, nullable=True) # Comma-separated list of DKIM selectors
|
||||
|
||||
@@ -29,13 +31,20 @@ class Domain(Base):
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
# Relationships
|
||||
workspace = relationship("Workspace", back_populates="domains")
|
||||
reports = relationship("DMARCReport", back_populates="domain", cascade="all, delete-orphan")
|
||||
forensic_reports = relationship(
|
||||
"ForensicReport", back_populates="domain", cascade="all, delete-orphan"
|
||||
)
|
||||
tls_reports = relationship("TLSReport", back_populates="domain", cascade="all, delete-orphan")
|
||||
user_domains = relationship("UserDomain", back_populates="domain", cascade="all, delete-orphan")
|
||||
|
||||
# Indexes for common queries
|
||||
__table_args__ = (
|
||||
# Index for finding active and verified domains
|
||||
Index("ix_domains_active_verified", "active", "verified"),
|
||||
# Workspace-scoped domain lookups for MSP mode.
|
||||
Index("ix_domains_workspace_name", "workspace_id", "name"),
|
||||
# Index for finding domains by policy
|
||||
Index("ix_domains_policy", "dmarc_policy"),
|
||||
# Index for finding recently updated domains
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Index, Integer, String, Text
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.core.credential_encryption import decrypt_secret, encrypt_secret, is_encrypted_secret
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class MailSource(Base):
|
||||
"""
|
||||
Mail source configuration model.
|
||||
|
||||
Stores credentials and settings for a mail inbox used to retrieve DMARC
|
||||
aggregate reports. The ``method`` field determines how the connection is
|
||||
made and which additional fields are relevant:
|
||||
|
||||
- ``IMAP`` – standard IMAP4 (over SSL/TLS or STARTTLS)
|
||||
- ``POP3`` – POP3 inbox (stub for future implementation)
|
||||
- ``GMAIL_API`` – Gmail API with OAuth 2.0
|
||||
- ``M365_GRAPH`` – Microsoft 365 / Exchange Online via Microsoft Graph
|
||||
"""
|
||||
|
||||
__tablename__ = "mail_sources"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
workspace_id = Column(Integer, ForeignKey("workspaces.id"), nullable=True, index=True)
|
||||
|
||||
# Human-readable label for the source
|
||||
name = Column(String, nullable=False)
|
||||
|
||||
# Connection method – determines which fields are used at runtime
|
||||
method = Column(String, nullable=False, default="IMAP") # IMAP | POP3 | GMAIL_API | M365_GRAPH
|
||||
|
||||
# Connection details (used by IMAP and POP3)
|
||||
server = Column(String, nullable=True)
|
||||
port = Column(Integer, nullable=True, default=993)
|
||||
username = Column(String, nullable=True)
|
||||
_password = Column("password", Text, nullable=True)
|
||||
use_ssl = Column(Boolean, default=True)
|
||||
folder = Column(String, default="INBOX")
|
||||
|
||||
# Gmail API OAuth2 credentials (used by GMAIL_API method)
|
||||
gmail_client_id = Column(String, nullable=True)
|
||||
_gmail_client_secret = Column("gmail_client_secret", Text, nullable=True)
|
||||
_gmail_access_token = Column("gmail_access_token", Text, nullable=True)
|
||||
_gmail_refresh_token = Column("gmail_refresh_token", Text, nullable=True)
|
||||
# Email address of the authorised Gmail account
|
||||
gmail_email = Column(String, nullable=True)
|
||||
# JSON-encoded list of Gmail message IDs that have already been ingested
|
||||
gmail_ingested_ids = Column(Text, nullable=True, default="[]")
|
||||
|
||||
# Microsoft 365 / Graph OAuth2 credentials (used by M365_GRAPH method)
|
||||
m365_tenant_id = Column(String, nullable=True, default="common")
|
||||
m365_client_id = Column(String, nullable=True)
|
||||
_m365_client_secret = Column("m365_client_secret", Text, nullable=True)
|
||||
_m365_access_token = Column("m365_access_token", Text, nullable=True)
|
||||
_m365_refresh_token = Column("m365_refresh_token", Text, nullable=True)
|
||||
# Optional user/shared mailbox to poll. Empty means the authorised account (/me).
|
||||
m365_mailbox = Column(String, nullable=True)
|
||||
# Optional Microsoft Graph mailFolder id. Empty means use ``folder`` as a well-known name.
|
||||
m365_folder_id = Column(String, nullable=True)
|
||||
# Email address reported by Microsoft Graph for the authorised account.
|
||||
m365_email = Column(String, nullable=True)
|
||||
# JSON-encoded list of Graph message IDs that have already been ingested
|
||||
m365_ingested_ids = Column(Text, nullable=True, default="[]")
|
||||
|
||||
# Polling behaviour
|
||||
polling_interval = Column(Integer, default=60) # minutes
|
||||
|
||||
# Source lifecycle
|
||||
enabled = Column(Boolean, default=True, index=True)
|
||||
last_checked = Column(DateTime, nullable=True)
|
||||
|
||||
# Timestamps
|
||||
created_at = Column(DateTime, default=datetime.utcnow)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
imports = relationship(
|
||||
"MailSourceImport",
|
||||
back_populates="mail_source",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
workspace = relationship("Workspace", back_populates="mail_sources")
|
||||
|
||||
__table_args__ = (Index("ix_mail_sources_workspace_enabled", "workspace_id", "enabled"),)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<MailSource id={self.id} name={self.name!r} method={self.method!r}>"
|
||||
|
||||
def encrypt_legacy_secrets(self) -> bool:
|
||||
"""Encrypt any legacy plaintext secrets already stored on this row."""
|
||||
changed = False
|
||||
secret_fields = {
|
||||
"password": self._password,
|
||||
"gmail_client_secret": self._gmail_client_secret,
|
||||
"gmail_access_token": self._gmail_access_token,
|
||||
"gmail_refresh_token": self._gmail_refresh_token,
|
||||
"m365_client_secret": self._m365_client_secret,
|
||||
"m365_access_token": self._m365_access_token,
|
||||
"m365_refresh_token": self._m365_refresh_token,
|
||||
}
|
||||
|
||||
for public_name, stored_value in secret_fields.items():
|
||||
if stored_value and not is_encrypted_secret(stored_value):
|
||||
setattr(self, public_name, stored_value)
|
||||
changed = True
|
||||
|
||||
return changed
|
||||
|
||||
@property
|
||||
def password(self):
|
||||
"""Return the decrypted IMAP password, if present."""
|
||||
return decrypt_secret(self._password)
|
||||
|
||||
@password.setter
|
||||
def password(self, value):
|
||||
self._password = encrypt_secret(value)
|
||||
|
||||
@property
|
||||
def gmail_client_secret(self):
|
||||
"""Return the decrypted Gmail OAuth client secret, if present."""
|
||||
return decrypt_secret(self._gmail_client_secret)
|
||||
|
||||
@gmail_client_secret.setter
|
||||
def gmail_client_secret(self, value):
|
||||
self._gmail_client_secret = encrypt_secret(value)
|
||||
|
||||
@property
|
||||
def gmail_access_token(self):
|
||||
"""Return the decrypted Gmail OAuth access token, if present."""
|
||||
return decrypt_secret(self._gmail_access_token)
|
||||
|
||||
@gmail_access_token.setter
|
||||
def gmail_access_token(self, value):
|
||||
self._gmail_access_token = encrypt_secret(value)
|
||||
|
||||
@property
|
||||
def gmail_refresh_token(self):
|
||||
"""Return the decrypted Gmail OAuth refresh token, if present."""
|
||||
return decrypt_secret(self._gmail_refresh_token)
|
||||
|
||||
@gmail_refresh_token.setter
|
||||
def gmail_refresh_token(self, value):
|
||||
self._gmail_refresh_token = encrypt_secret(value)
|
||||
|
||||
@property
|
||||
def m365_client_secret(self):
|
||||
"""Return the decrypted Microsoft 365 OAuth client secret, if present."""
|
||||
return decrypt_secret(self._m365_client_secret)
|
||||
|
||||
@m365_client_secret.setter
|
||||
def m365_client_secret(self, value):
|
||||
self._m365_client_secret = encrypt_secret(value)
|
||||
|
||||
@property
|
||||
def m365_access_token(self):
|
||||
"""Return the decrypted Microsoft Graph access token, if present."""
|
||||
return decrypt_secret(self._m365_access_token)
|
||||
|
||||
@m365_access_token.setter
|
||||
def m365_access_token(self, value):
|
||||
self._m365_access_token = encrypt_secret(value)
|
||||
|
||||
@property
|
||||
def m365_refresh_token(self):
|
||||
"""Return the decrypted Microsoft Graph refresh token, if present."""
|
||||
return decrypt_secret(self._m365_refresh_token)
|
||||
|
||||
@m365_refresh_token.setter
|
||||
def m365_refresh_token(self, value):
|
||||
self._m365_refresh_token = encrypt_secret(value)
|
||||
@@ -0,0 +1,39 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Column, DateTime, ForeignKey, Integer, String, Text
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class MailSourceImport(Base):
|
||||
"""Sanitized audit record for one mail source import attempt."""
|
||||
|
||||
__tablename__ = "mail_source_imports"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
mail_source_id = Column(Integer, ForeignKey("mail_sources.id"), nullable=False, index=True)
|
||||
|
||||
trigger = Column(String, nullable=False, default="manual")
|
||||
status = Column(String, nullable=False, index=True)
|
||||
|
||||
processed = Column(Integer, nullable=False, default=0)
|
||||
reports_found = Column(Integer, nullable=False, default=0)
|
||||
duplicate_reports = Column(Integer, nullable=False, default=0)
|
||||
error_count = Column(Integer, nullable=False, default=0)
|
||||
|
||||
new_domains = Column(Text, nullable=True)
|
||||
errors = Column(Text, nullable=True)
|
||||
details = Column(Text, nullable=True)
|
||||
|
||||
started_at = Column(DateTime, nullable=False, default=datetime.utcnow, index=True)
|
||||
finished_at = Column(DateTime, nullable=False, default=datetime.utcnow, index=True)
|
||||
created_at = Column(DateTime, nullable=False, default=datetime.utcnow)
|
||||
|
||||
mail_source = relationship("MailSource", back_populates="imports")
|
||||
|
||||
def __repr__(self):
|
||||
return (
|
||||
f"<MailSourceImport id={self.id} source={self.mail_source_id} "
|
||||
f"status={self.status!r}>"
|
||||
)
|
||||
@@ -1,9 +1,10 @@
|
||||
from datetime import datetime
|
||||
|
||||
from app.core.database import Base
|
||||
from sqlalchemy import Column, DateTime, ForeignKey, Index, Integer, String, Text
|
||||
from sqlalchemy import Column, DateTime, ForeignKey, Index, Integer, String, Text, UniqueConstraint
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class DMARCReport(Base):
|
||||
"""DMARC Aggregate Report model"""
|
||||
@@ -19,16 +20,27 @@ class DMARCReport(Base):
|
||||
begin_date = Column(Integer, nullable=False, index=True) # Unix timestamp
|
||||
end_date = Column(Integer, nullable=False, index=True) # Unix timestamp
|
||||
source_email = Column(String, nullable=True)
|
||||
extra_contact_info = Column(String, nullable=True)
|
||||
generator = Column(String, nullable=True)
|
||||
report_errors = Column(Text, nullable=True)
|
||||
|
||||
# Policy information
|
||||
policy = Column(String, nullable=True) # none, quarantine, reject (indexed via __table_args__)
|
||||
subdomain_policy = Column(String, nullable=True)
|
||||
non_subdomain_policy = Column(String, nullable=True)
|
||||
adkim = Column(String(1), nullable=True) # r (relaxed) or s (strict)
|
||||
aspf = Column(String(1), nullable=True) # r (relaxed) or s (strict)
|
||||
percentage = Column(Integer, nullable=True)
|
||||
failure_options = Column(String, nullable=True)
|
||||
testing = Column(String, nullable=True)
|
||||
discovery_method = Column(String, nullable=True)
|
||||
|
||||
# Processing metadata
|
||||
processed_at = Column(DateTime, default=datetime.utcnow, index=True)
|
||||
schema_version = Column(String, nullable=True)
|
||||
report_variant = Column(String, nullable=True)
|
||||
xml_namespace = Column(String, nullable=True)
|
||||
report_extensions = Column(Text, nullable=True)
|
||||
processed_at = Column(DateTime, default=datetime.utcnow)
|
||||
raw_data = Column(Text, nullable=True) # Original XML content (optional)
|
||||
|
||||
# Relationships
|
||||
@@ -71,10 +83,13 @@ class ReportRecord(Base):
|
||||
# Identifiers
|
||||
header_from = Column(String, nullable=True, index=True)
|
||||
envelope_from = Column(String, nullable=True)
|
||||
envelope_to = Column(String, nullable=True)
|
||||
|
||||
# Authentication details (optional JSON fields)
|
||||
dkim_auth_details = Column(Text, nullable=True) # JSON array of DKIM results
|
||||
spf_auth_details = Column(Text, nullable=True) # JSON array of SPF results
|
||||
policy_override_reasons = Column(Text, nullable=True) # JSON array of policy reasons
|
||||
record_extensions = Column(Text, nullable=True) # JSON object of extension values
|
||||
|
||||
# Relationships
|
||||
report = relationship("DMARCReport", back_populates="records")
|
||||
@@ -89,3 +104,114 @@ class ReportRecord(Base):
|
||||
|
||||
def __repr__(self):
|
||||
return f"<ReportRecord {self.id} ({self.source_ip})>"
|
||||
|
||||
|
||||
class ForensicReport(Base):
|
||||
"""DMARC forensic/failure report model (RFC 6591 / ARF)."""
|
||||
|
||||
__tablename__ = "forensic_reports"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
domain_id = Column(Integer, ForeignKey("domains.id"), nullable=True, index=True)
|
||||
|
||||
# Report metadata
|
||||
report_id = Column(String, nullable=False, index=True)
|
||||
source_email = Column(String, nullable=True)
|
||||
feedback_type = Column(String, nullable=True, index=True)
|
||||
user_agent = Column(String, nullable=True)
|
||||
version = Column(String, nullable=True)
|
||||
|
||||
# DMARC failure fields
|
||||
reported_domain = Column(String, nullable=True, index=True)
|
||||
source_ip = Column(String, nullable=True, index=True)
|
||||
auth_failure = Column(String, nullable=True, index=True)
|
||||
delivery_result = Column(String, nullable=True)
|
||||
arrival_date = Column(DateTime, nullable=True, index=True)
|
||||
authentication_results = Column(Text, nullable=True)
|
||||
|
||||
# Redacted original-message metadata. Never store original body content here.
|
||||
original_mail_from = Column(String, nullable=True)
|
||||
original_from = Column(String, nullable=True)
|
||||
original_to = Column(String, nullable=True)
|
||||
original_subject = Column(String, nullable=True)
|
||||
original_message_id = Column(String, nullable=True)
|
||||
original_date = Column(String, nullable=True)
|
||||
|
||||
# Sanitized parser details for operators/debugging.
|
||||
feedback_headers = Column(Text, nullable=True)
|
||||
processed_at = Column(DateTime, default=datetime.utcnow, index=True)
|
||||
|
||||
domain = relationship("Domain", back_populates="forensic_reports")
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint("report_id", name="uq_forensic_reports_report_id"),
|
||||
Index("ix_forensic_reports_domain_arrival", "domain_id", "arrival_date"),
|
||||
Index("ix_forensic_reports_failure_source", "auth_failure", "source_ip"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<ForensicReport {self.report_id}>"
|
||||
|
||||
|
||||
class TLSReport(Base):
|
||||
"""SMTP TLS reporting (TLS-RPT) aggregate report."""
|
||||
|
||||
__tablename__ = "tls_reports"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
domain_id = Column(Integer, ForeignKey("domains.id"), nullable=True, index=True)
|
||||
|
||||
report_id = Column(String, nullable=False, index=True)
|
||||
org_name = Column(String, nullable=True, index=True)
|
||||
contact_info = Column(String, nullable=True)
|
||||
policy_domain = Column(String, nullable=False, index=True)
|
||||
policy_type = Column(String, nullable=True, index=True)
|
||||
begin_date = Column(DateTime, nullable=True, index=True)
|
||||
end_date = Column(DateTime, nullable=True, index=True)
|
||||
total_successful_sessions = Column(Integer, nullable=False, default=0)
|
||||
total_failure_sessions = Column(Integer, nullable=False, default=0)
|
||||
raw_policy = Column(Text, nullable=True)
|
||||
processed_at = Column(DateTime, default=datetime.utcnow, index=True)
|
||||
|
||||
domain = relationship("Domain", back_populates="tls_reports")
|
||||
failures = relationship(
|
||||
"TLSReportFailure",
|
||||
back_populates="report",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint("report_id", "policy_domain", name="uq_tls_reports_report_domain"),
|
||||
Index("ix_tls_reports_domain_dates", "domain_id", "begin_date", "end_date"),
|
||||
Index("ix_tls_reports_policy_domain_dates", "policy_domain", "begin_date", "end_date"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<TLSReport {self.report_id} {self.policy_domain}>"
|
||||
|
||||
|
||||
class TLSReportFailure(Base):
|
||||
"""Grouped TLS-RPT failure detail without message-level identifiers."""
|
||||
|
||||
__tablename__ = "tls_report_failures"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
report_id = Column(Integer, ForeignKey("tls_reports.id"), nullable=False, index=True)
|
||||
result_type = Column(String, nullable=False, index=True)
|
||||
failed_session_count = Column(Integer, nullable=False, default=0)
|
||||
sending_mta_ip = Column(String, nullable=True, index=True)
|
||||
receiving_mx_hostname = Column(String, nullable=True, index=True)
|
||||
receiving_mx_helo = Column(String, nullable=True)
|
||||
receiving_ip = Column(String, nullable=True)
|
||||
failure_reason_code = Column(String, nullable=True)
|
||||
additional_information = Column(Text, nullable=True)
|
||||
|
||||
report = relationship("TLSReport", back_populates="failures")
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_tls_report_failures_result_count", "result_type", "failed_session_count"),
|
||||
Index("ix_tls_report_failures_report_result", "report_id", "result_type"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<TLSReportFailure {self.result_type} count={self.failed_session_count}>"
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Column, DateTime, ForeignKey, Integer, String, Text
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class Setting(Base):
|
||||
"""
|
||||
Key-value store for system-wide application settings.
|
||||
|
||||
Settings are grouped by a ``category`` prefix (e.g. ``general``,
|
||||
``dmarc``, ``cloudflare``) to make bulk retrieval and UI grouping easy.
|
||||
The ``value`` is always stored as text; callers are responsible for
|
||||
serialising/deserialising typed values (int, bool, JSON) via the
|
||||
``value_type`` hint.
|
||||
"""
|
||||
|
||||
__tablename__ = "settings"
|
||||
|
||||
key = Column(String(100), primary_key=True)
|
||||
value = Column(Text, nullable=True)
|
||||
# Human-readable description shown in the admin UI
|
||||
description = Column(String(255), nullable=True)
|
||||
# Hint for the UI / API on how to interpret the value: string | integer | boolean | json
|
||||
value_type = Column(String(20), nullable=False, default="string")
|
||||
# Category / section grouping (e.g. "general", "dmarc", "cloudflare", "dns")
|
||||
category = Column(String(50), nullable=False, default="general")
|
||||
# Audit fields
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
updated_by = Column(Integer, ForeignKey("users.id"), nullable=True)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<Setting key={self.key!r} category={self.category!r}>"
|
||||
@@ -1,23 +1,46 @@
|
||||
from app.core.database import Base
|
||||
from sqlalchemy import Boolean, Column, Integer, String
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Integer, String
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class User(Base):
|
||||
"""User model"""
|
||||
"""User model – local shadow of the identity managed by Logto."""
|
||||
|
||||
__tablename__ = "users"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
workspace_id = Column(Integer, ForeignKey("workspaces.id"), nullable=True, index=True)
|
||||
email = Column(String, unique=True, index=True, nullable=False)
|
||||
hashed_password = Column(String, nullable=False)
|
||||
# Logto subject claim (the user's stable ID inside Logto).
|
||||
# Null for users that pre-date Logto integration or for
|
||||
# programmatic/service accounts created directly in the DB.
|
||||
logto_id = Column(String, unique=True, index=True, nullable=True)
|
||||
# hashed_password kept for possible future local-auth fallback; nullable
|
||||
# because Logto users authenticate externally and have no local password.
|
||||
hashed_password = Column(String, nullable=True)
|
||||
is_active = Column(Boolean, default=True)
|
||||
is_superuser = Column(Boolean, default=False)
|
||||
# For now all users are treated as admin. RBAC tiers are planned.
|
||||
is_superuser = Column(Boolean, default=True)
|
||||
is_verified = Column(Boolean, default=False)
|
||||
|
||||
# Additional fields
|
||||
# Profile – synced from Logto claims on every login
|
||||
full_name = Column(String, nullable=True)
|
||||
username = Column(String, nullable=True)
|
||||
organization = Column(String, nullable=True)
|
||||
picture = Column(String, nullable=True)
|
||||
|
||||
# Timestamps
|
||||
created_at = Column(DateTime, default=datetime.utcnow, nullable=True)
|
||||
updated_at = Column(
|
||||
DateTime,
|
||||
default=datetime.utcnow,
|
||||
onupdate=datetime.utcnow,
|
||||
nullable=True,
|
||||
)
|
||||
|
||||
# Relationships
|
||||
workspace = relationship("Workspace", back_populates="users")
|
||||
user_domains = relationship("UserDomain", back_populates="user", cascade="all, delete-orphan")
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Index, Integer, String, Text
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class WebhookEndpoint(Base):
|
||||
"""Outbound webhook endpoint configured by an operator."""
|
||||
|
||||
__tablename__ = "webhook_endpoints"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
name = Column(String(120), nullable=False)
|
||||
url = Column(Text, nullable=False)
|
||||
secret = Column(Text, nullable=False)
|
||||
event_types = Column(Text, nullable=False, default="*")
|
||||
enabled = Column(Boolean, default=True, nullable=False, index=True)
|
||||
max_attempts = Column(Integer, default=5, nullable=False)
|
||||
timeout_seconds = Column(Integer, default=10, nullable=False)
|
||||
created_at = Column(DateTime, default=datetime.utcnow, nullable=False, index=True)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
last_success_at = Column(DateTime, nullable=True, index=True)
|
||||
last_failure_at = Column(DateTime, nullable=True, index=True)
|
||||
failure_count = Column(Integer, default=0, nullable=False)
|
||||
|
||||
__table_args__ = (Index("ix_webhook_endpoints_enabled_events", "enabled", "event_types"),)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<WebhookEndpoint {self.name} enabled={self.enabled}>"
|
||||
|
||||
|
||||
class WebhookDelivery(Base):
|
||||
"""Single outbound webhook delivery attempt state."""
|
||||
|
||||
__tablename__ = "webhook_deliveries"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
endpoint_id = Column(Integer, ForeignKey("webhook_endpoints.id"), nullable=False, index=True)
|
||||
event_type = Column(String(80), nullable=False, index=True)
|
||||
payload = Column(Text, nullable=False)
|
||||
idempotency_key = Column(String(160), nullable=False, index=True)
|
||||
status = Column(String(24), nullable=False, default="pending", index=True)
|
||||
attempt_count = Column(Integer, default=0, nullable=False)
|
||||
max_attempts = Column(Integer, default=5, nullable=False)
|
||||
next_attempt_at = Column(DateTime, default=datetime.utcnow, nullable=False, index=True)
|
||||
last_attempt_at = Column(DateTime, nullable=True, index=True)
|
||||
delivered_at = Column(DateTime, nullable=True, index=True)
|
||||
last_status_code = Column(Integer, nullable=True)
|
||||
last_error = Column(Text, nullable=True)
|
||||
response_excerpt = Column(Text, nullable=True)
|
||||
created_at = Column(DateTime, default=datetime.utcnow, nullable=False, index=True)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
__table_args__ = (
|
||||
Index(
|
||||
"ix_webhook_delivery_endpoint_idempotency",
|
||||
"endpoint_id",
|
||||
"idempotency_key",
|
||||
unique=True,
|
||||
),
|
||||
Index("ix_webhook_delivery_due", "status", "next_attempt_at"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<WebhookDelivery endpoint={self.endpoint_id} event={self.event_type} status={self.status}>"
|
||||
@@ -0,0 +1,32 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, Index, Integer, String, Text
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class Workspace(Base):
|
||||
"""Tenant/workspace boundary for monitored DMARC assets."""
|
||||
|
||||
__tablename__ = "workspaces"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
slug = Column(String, unique=True, nullable=False, index=True)
|
||||
name = Column(String, nullable=False)
|
||||
description = Column(Text, nullable=True)
|
||||
active = Column(Boolean, default=True, nullable=False, index=True)
|
||||
report_retention_days = Column(Integer, default=400, nullable=False)
|
||||
forensic_retention_days = Column(Integer, default=90, nullable=False)
|
||||
tls_report_retention_days = Column(Integer, default=400, nullable=False)
|
||||
created_at = Column(DateTime, default=datetime.utcnow, index=True)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
domains = relationship("Domain", back_populates="workspace")
|
||||
mail_sources = relationship("MailSource", back_populates="workspace")
|
||||
users = relationship("User", back_populates="workspace")
|
||||
|
||||
__table_args__ = (Index("ix_workspaces_active_slug", "active", "slug"),)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<Workspace {self.slug}>"
|
||||
@@ -0,0 +1,60 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Index, Integer, String, Text
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.core.database import Base
|
||||
|
||||
|
||||
class WorkspaceMembership(Base):
|
||||
"""User role assignment inside one workspace."""
|
||||
|
||||
__tablename__ = "workspace_memberships"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
workspace_id = Column(Integer, ForeignKey("workspaces.id"), nullable=False, index=True)
|
||||
user_id = Column(Integer, ForeignKey("users.id"), nullable=False, index=True)
|
||||
role = Column(String(50), nullable=False, index=True)
|
||||
active = Column(Boolean, default=True, nullable=False, index=True)
|
||||
created_at = Column(DateTime, default=datetime.utcnow, nullable=False, index=True)
|
||||
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
|
||||
|
||||
workspace = relationship("Workspace")
|
||||
user = relationship("User")
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_workspace_memberships_workspace_user", "workspace_id", "user_id", unique=True),
|
||||
Index("ix_workspace_memberships_workspace_role", "workspace_id", "role"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<WorkspaceMembership workspace={self.workspace_id} user={self.user_id}>"
|
||||
|
||||
|
||||
class WorkspaceAuditLog(Base):
|
||||
"""Sanitized audit trail for sensitive workspace-scoped changes."""
|
||||
|
||||
__tablename__ = "workspace_audit_logs"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
workspace_id = Column(Integer, ForeignKey("workspaces.id"), nullable=False, index=True)
|
||||
actor_type = Column(String(50), nullable=False, index=True)
|
||||
actor_id = Column(String(120), nullable=True, index=True)
|
||||
action = Column(String(100), nullable=False, index=True)
|
||||
entity_type = Column(String(80), nullable=False, index=True)
|
||||
entity_id = Column(String(120), nullable=True, index=True)
|
||||
entity_name = Column(String(255), nullable=True)
|
||||
details = Column(Text, nullable=True)
|
||||
ip_address = Column(String(64), nullable=True)
|
||||
created_at = Column(DateTime, default=datetime.utcnow, nullable=False, index=True)
|
||||
|
||||
workspace = relationship("Workspace")
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_workspace_audit_workspace_created", "workspace_id", "created_at"),
|
||||
Index("ix_workspace_audit_workspace_action", "workspace_id", "action"),
|
||||
Index("ix_workspace_audit_entity", "entity_type", "entity_id"),
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<WorkspaceAuditLog {self.action} {self.entity_type}>"
|
||||
@@ -0,0 +1,317 @@
|
||||
"""Privacy-preserving AI and agent assistance helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.redaction import SENSITIVE_VALUE, redact_sensitive_text
|
||||
from app.models.domain import Domain
|
||||
from app.models.setting import Setting
|
||||
from app.services.report_persistence import hydrate_report_store_from_db
|
||||
from app.services.report_store import ReportStore
|
||||
|
||||
AI_DEFAULTS = {
|
||||
"ai.enabled": "false",
|
||||
"ai.provider": "template",
|
||||
"ai.model": "",
|
||||
"ai.remote_base_url": "",
|
||||
"ai.redaction_mode": "strict",
|
||||
"ai.action_tools_enabled": "false",
|
||||
"mcp.enabled": "false",
|
||||
}
|
||||
|
||||
EMAIL_PATTERN = re.compile(r"\b[A-Z0-9._%+-]+@([A-Z0-9.-]+\.[A-Z]{2,})\b", re.IGNORECASE)
|
||||
LONG_TOKEN_PATTERN = re.compile(r"\b[A-Za-z0-9._~+/=-]{24,}\b")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AssistanceConfig:
|
||||
"""Operator-controlled automation settings."""
|
||||
|
||||
ai_enabled: bool
|
||||
provider: str
|
||||
model: str
|
||||
remote_base_url: str
|
||||
redaction_mode: str
|
||||
action_tools_enabled: bool
|
||||
mcp_enabled: bool
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Return a UI/API-safe provider configuration."""
|
||||
return {
|
||||
"ai_enabled": self.ai_enabled,
|
||||
"provider": self.provider,
|
||||
"model": self.model,
|
||||
"remote_base_url_configured": bool(self.remote_base_url),
|
||||
"redaction_mode": self.redaction_mode,
|
||||
"action_tools_enabled": self.action_tools_enabled,
|
||||
"mcp_enabled": self.mcp_enabled,
|
||||
"data_handling": {
|
||||
"default_provider": "template",
|
||||
"secrets_in_prompts": "never",
|
||||
"raw_message_content": "never",
|
||||
"remote_provider_requires_opt_in": True,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _setting_value(db: Session, key: str) -> str:
|
||||
row = db.query(Setting).filter(Setting.key == key).first()
|
||||
if row is None:
|
||||
return AI_DEFAULTS.get(key, "")
|
||||
return row.value or ""
|
||||
|
||||
|
||||
def _setting_bool(db: Session, key: str) -> bool:
|
||||
return _setting_value(db, key).strip().lower() in {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def get_assistance_config(db: Session) -> AssistanceConfig:
|
||||
"""Load AI/MCP settings with safe defaults."""
|
||||
provider = (_setting_value(db, "ai.provider") or "template").strip().lower()
|
||||
if provider not in {"template", "local", "remote"}:
|
||||
provider = "template"
|
||||
redaction_mode = (_setting_value(db, "ai.redaction_mode") or "strict").strip().lower()
|
||||
if redaction_mode not in {"strict", "balanced"}:
|
||||
redaction_mode = "strict"
|
||||
return AssistanceConfig(
|
||||
ai_enabled=_setting_bool(db, "ai.enabled"),
|
||||
provider=provider,
|
||||
model=_setting_value(db, "ai.model").strip(),
|
||||
remote_base_url=_setting_value(db, "ai.remote_base_url").strip(),
|
||||
redaction_mode=redaction_mode,
|
||||
action_tools_enabled=_setting_bool(db, "ai.action_tools_enabled"),
|
||||
mcp_enabled=_setting_bool(db, "mcp.enabled"),
|
||||
)
|
||||
|
||||
|
||||
def redact_safe_value(value: Any, *, mode: str = "strict") -> Any:
|
||||
"""Redact values before they can be shared with model or agent surfaces."""
|
||||
if isinstance(value, dict):
|
||||
return {str(key): redact_safe_value(item, mode=mode) for key, item in value.items()}
|
||||
if isinstance(value, list):
|
||||
return [redact_safe_value(item, mode=mode) for item in value]
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
redacted = redact_sensitive_text(value)
|
||||
redacted = LONG_TOKEN_PATTERN.sub(SENSITIVE_VALUE, redacted)
|
||||
if mode == "strict":
|
||||
redacted = EMAIL_PATTERN.sub(r"*@\1", redacted)
|
||||
return redacted
|
||||
|
||||
|
||||
def _domain_exists(db: Session, store: ReportStore, domain: str) -> bool:
|
||||
if domain in store.get_domains():
|
||||
return True
|
||||
return (
|
||||
db.query(Domain.id).filter(Domain.name == domain, Domain.active.is_(True)).first()
|
||||
is not None
|
||||
)
|
||||
|
||||
|
||||
def _evidence(label: str, value: Any, href: str) -> Dict[str, str]:
|
||||
return {
|
||||
"label": label,
|
||||
"value": str(value),
|
||||
"href": href,
|
||||
}
|
||||
|
||||
|
||||
def build_safe_context(db: Session, domain: str) -> Dict[str, Any]:
|
||||
"""Build a redacted, evidence-linked context bundle for one domain."""
|
||||
store = ReportStore.get_instance()
|
||||
hydrate_report_store_from_db(db, store)
|
||||
if not _domain_exists(db, store, domain):
|
||||
raise ValueError("Domain not found")
|
||||
|
||||
config = get_assistance_config(db)
|
||||
summary = store.get_domain_summary(domain)
|
||||
total = int(summary.get("total_count", 0) or 0)
|
||||
passed = int(summary.get("passed_count", 0) or 0)
|
||||
failed = int(summary.get("failed_count", max(0, total - passed)) or 0)
|
||||
compliance = float(summary.get("compliance_rate", 0.0) or 0.0)
|
||||
reports_processed = int(summary.get("reports_processed", 0) or 0)
|
||||
sources = store.get_domain_sources(domain, days=30)[:5]
|
||||
reports = store.get_domain_reports(domain, limit=5)
|
||||
|
||||
context = {
|
||||
"domain": domain,
|
||||
"generated_at": datetime.utcnow().isoformat() + "Z",
|
||||
"config": config.to_dict(),
|
||||
"summary": {
|
||||
"total_messages": total,
|
||||
"passed_messages": passed,
|
||||
"failed_messages": failed,
|
||||
"compliance_rate": compliance,
|
||||
"reports_processed": reports_processed,
|
||||
"policy": summary.get("policy", "unknown"),
|
||||
},
|
||||
"top_sources": [
|
||||
{
|
||||
"source_ip": source.get("source_ip", "unknown"),
|
||||
"count": int(source.get("count", 0) or 0),
|
||||
"spf": source.get("spf_result", "unknown"),
|
||||
"dkim": source.get("dkim_result", "unknown"),
|
||||
"dmarc": source.get("dmarc_result", "unknown"),
|
||||
"disposition": source.get("disposition", "none"),
|
||||
}
|
||||
for source in sources
|
||||
],
|
||||
"recent_reports": [
|
||||
{
|
||||
"report_id": report.get("report_id", "unknown"),
|
||||
"org_name": report.get("org_name", "Unknown Organization"),
|
||||
"total_messages": int(report.get("summary", {}).get("total_count", 0) or 0),
|
||||
"pass_rate": report.get("pass_rate", 0.0),
|
||||
}
|
||||
for report in reports
|
||||
],
|
||||
"evidence": [
|
||||
_evidence("Domain summary", domain, f"/domain/{domain}"),
|
||||
_evidence("Total messages", total, f"/domain/{domain}#compliance-chart"),
|
||||
_evidence("Compliance rate", f"{compliance}%", f"/domain/{domain}#compliance-chart"),
|
||||
_evidence("Failed messages", failed, f"/domain/{domain}#sending-sources"),
|
||||
],
|
||||
"redaction": {
|
||||
"mode": config.redaction_mode,
|
||||
"applied": True,
|
||||
"rules": [
|
||||
"secret-like key/value fragments",
|
||||
"bearer tokens",
|
||||
"long opaque tokens",
|
||||
"email local-parts in strict mode",
|
||||
],
|
||||
},
|
||||
}
|
||||
return redact_safe_value(context, mode=config.redaction_mode)
|
||||
|
||||
|
||||
def _headline_for_context(context: Dict[str, Any]) -> str:
|
||||
summary = context["summary"]
|
||||
total = int(summary["total_messages"])
|
||||
failed = int(summary["failed_messages"])
|
||||
compliance = float(summary["compliance_rate"])
|
||||
if total == 0:
|
||||
return "No DMARC aggregate volume has been observed yet."
|
||||
if failed == 0 and compliance >= 99:
|
||||
return "Observed DMARC traffic is passing cleanly."
|
||||
if compliance >= 90:
|
||||
return "DMARC posture is mostly healthy, with a small failure set to review."
|
||||
return "DMARC posture needs remediation before policy enforcement."
|
||||
|
||||
|
||||
def build_evidence_summary(db: Session, domain: str) -> Dict[str, Any]:
|
||||
"""Return an evidence-first operator summary and remediation plan."""
|
||||
context = build_safe_context(db, domain)
|
||||
summary = context["summary"]
|
||||
failed = int(summary["failed_messages"])
|
||||
total = int(summary["total_messages"])
|
||||
compliance = float(summary["compliance_rate"])
|
||||
recommendations: List[Dict[str, Any]] = []
|
||||
|
||||
if total == 0:
|
||||
recommendations.append(
|
||||
{
|
||||
"priority": "medium",
|
||||
"title": "Confirm report ingestion",
|
||||
"detail": "No aggregate reports are available for this domain.",
|
||||
"action": "Check mailbox sources and senders for rua delivery.",
|
||||
"evidence": [context["evidence"][0]],
|
||||
}
|
||||
)
|
||||
elif failed > 0:
|
||||
recommendations.append(
|
||||
{
|
||||
"priority": "high" if compliance < 90 else "medium",
|
||||
"title": "Review failing sending sources",
|
||||
"detail": f"{failed} of {total} observed messages failed DMARC alignment.",
|
||||
"action": (
|
||||
"Open the sending-source evidence and confirm whether each failing "
|
||||
"source is legitimate."
|
||||
),
|
||||
"evidence": [context["evidence"][2], context["evidence"][3]],
|
||||
}
|
||||
)
|
||||
|
||||
unknown_or_failing_sources = [
|
||||
source
|
||||
for source in context["top_sources"]
|
||||
if source.get("dmarc") in {"fail", "mixed", "unknown", "none"}
|
||||
]
|
||||
if unknown_or_failing_sources:
|
||||
recommendations.append(
|
||||
{
|
||||
"priority": "medium",
|
||||
"title": "Triage top unauthenticated sources",
|
||||
"detail": "At least one high-volume source is not consistently passing DMARC.",
|
||||
"action": "Verify ownership before adding SPF mechanisms or enabling DKIM signing.",
|
||||
"evidence": [
|
||||
_evidence(
|
||||
"Top source",
|
||||
f"{source['source_ip']} ({source['count']} messages)",
|
||||
f"/domain/{context['domain']}#sending-sources",
|
||||
)
|
||||
for source in unknown_or_failing_sources[:3]
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"enabled": get_assistance_config(db).ai_enabled,
|
||||
"provider": context["config"],
|
||||
"summary": {
|
||||
"domain": context["domain"],
|
||||
"headline": _headline_for_context(context),
|
||||
"total_messages": total,
|
||||
"failed_messages": failed,
|
||||
"compliance_rate": compliance,
|
||||
},
|
||||
"recommendations": recommendations,
|
||||
"safe_context": context,
|
||||
}
|
||||
|
||||
|
||||
def _proposal_id(payload: Dict[str, Any]) -> str:
|
||||
serialized = json.dumps(payload, sort_keys=True, separators=(",", ":"))
|
||||
return hashlib.sha256(serialized.encode("utf-8")).hexdigest()[:16]
|
||||
|
||||
|
||||
def build_action_proposals(db: Session, domain: str) -> Dict[str, Any]:
|
||||
"""Generate reviewable, reproducible action proposals without mutating state."""
|
||||
summary = build_evidence_summary(db, domain)
|
||||
proposals = []
|
||||
for index, recommendation in enumerate(summary["recommendations"], start=1):
|
||||
payload = {
|
||||
"domain": domain,
|
||||
"index": index,
|
||||
"title": recommendation["title"],
|
||||
"action": recommendation["action"],
|
||||
"evidence": recommendation.get("evidence", []),
|
||||
}
|
||||
proposal_id = _proposal_id(payload)
|
||||
proposals.append(
|
||||
{
|
||||
"proposal_id": proposal_id,
|
||||
"domain": domain,
|
||||
"status": "proposed",
|
||||
"title": recommendation["title"],
|
||||
"rationale": recommendation["detail"],
|
||||
"proposed_action": recommendation["action"],
|
||||
"requires_human_confirmation": True,
|
||||
"confirmation_text": proposal_id,
|
||||
"mutates_state": False,
|
||||
"evidence": recommendation.get("evidence", []),
|
||||
}
|
||||
)
|
||||
return {
|
||||
"domain": domain,
|
||||
"action_tools_enabled": get_assistance_config(db).action_tools_enabled,
|
||||
"proposals": proposals,
|
||||
}
|
||||
@@ -0,0 +1,188 @@
|
||||
"""Persistence helpers for alert history."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.alert import AlertConfigurationAudit, AlertHistory
|
||||
|
||||
|
||||
def _json_dumps(value: Dict[str, Any]) -> str:
|
||||
return json.dumps(value, sort_keys=True, separators=(",", ":"), default=str)
|
||||
|
||||
|
||||
def alert_fingerprint(alert: Dict[str, Any]) -> str:
|
||||
"""Return a stable fingerprint for one alert signal."""
|
||||
identity = {
|
||||
"rule": alert.get("rule"),
|
||||
"domain": alert.get("domain"),
|
||||
"source_ip": alert.get("source_ip"),
|
||||
"threshold": alert.get("threshold"),
|
||||
}
|
||||
return hashlib.sha256(_json_dumps(identity).encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _row_to_dict(row: AlertHistory) -> Dict[str, Any]:
|
||||
payload = {}
|
||||
if row.payload:
|
||||
try:
|
||||
payload = json.loads(row.payload)
|
||||
except json.JSONDecodeError:
|
||||
payload = {}
|
||||
return {
|
||||
"id": row.id,
|
||||
"fingerprint": row.fingerprint,
|
||||
"rule": row.rule,
|
||||
"severity": row.severity,
|
||||
"domain": row.domain,
|
||||
"title": row.title,
|
||||
"detail": row.detail,
|
||||
"payload": payload,
|
||||
"observed_count": row.observed_count,
|
||||
"is_active": row.is_active,
|
||||
"first_seen_at": row.first_seen_at.isoformat() if row.first_seen_at else None,
|
||||
"last_seen_at": row.last_seen_at.isoformat() if row.last_seen_at else None,
|
||||
"resolved_at": row.resolved_at.isoformat() if row.resolved_at else None,
|
||||
}
|
||||
|
||||
|
||||
def record_alert_evaluation(
|
||||
db: Session,
|
||||
alerts: List[Dict[str, Any]],
|
||||
*,
|
||||
observed_at: Optional[datetime] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Upsert current alerts and resolve active alerts that are no longer present."""
|
||||
timestamp = observed_at or datetime.utcnow()
|
||||
fingerprints = {alert_fingerprint(alert): alert for alert in alerts}
|
||||
|
||||
if fingerprints:
|
||||
existing_rows = (
|
||||
db.query(AlertHistory)
|
||||
.filter(AlertHistory.fingerprint.in_(list(fingerprints.keys())))
|
||||
.all()
|
||||
)
|
||||
else:
|
||||
existing_rows = []
|
||||
existing_by_fingerprint = {row.fingerprint: row for row in existing_rows}
|
||||
|
||||
for fingerprint, alert in fingerprints.items():
|
||||
row = existing_by_fingerprint.get(fingerprint)
|
||||
if row is None:
|
||||
row = AlertHistory(
|
||||
fingerprint=fingerprint,
|
||||
rule=str(alert.get("rule") or "unknown"),
|
||||
severity=str(alert.get("severity") or "warning"),
|
||||
domain=alert.get("domain"),
|
||||
title=str(alert.get("title") or "DMARC alert"),
|
||||
detail=str(alert.get("detail") or ""),
|
||||
payload=_json_dumps(alert),
|
||||
observed_count=1,
|
||||
is_active=True,
|
||||
first_seen_at=timestamp,
|
||||
last_seen_at=timestamp,
|
||||
)
|
||||
db.add(row)
|
||||
continue
|
||||
|
||||
row.rule = str(alert.get("rule") or row.rule)
|
||||
row.severity = str(alert.get("severity") or row.severity)
|
||||
row.domain = alert.get("domain")
|
||||
row.title = str(alert.get("title") or row.title)
|
||||
row.detail = str(alert.get("detail") or row.detail)
|
||||
row.payload = _json_dumps(alert)
|
||||
row.observed_count = int(row.observed_count or 0) + 1
|
||||
row.is_active = True
|
||||
row.last_seen_at = timestamp
|
||||
row.resolved_at = None
|
||||
|
||||
active_rows = db.query(AlertHistory).filter(AlertHistory.is_active == True).all() # noqa: E712
|
||||
for row in active_rows:
|
||||
if row.fingerprint not in fingerprints:
|
||||
row.is_active = False
|
||||
row.resolved_at = timestamp
|
||||
|
||||
db.commit()
|
||||
return list_alert_history(db, limit=max(50, len(alerts)))
|
||||
|
||||
|
||||
def list_alert_history(
|
||||
db: Session,
|
||||
*,
|
||||
active: Optional[bool] = None,
|
||||
limit: int = 50,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Return alert history rows ordered by most recent observation."""
|
||||
query = db.query(AlertHistory)
|
||||
if active is not None:
|
||||
query = query.filter(AlertHistory.is_active == active)
|
||||
rows = (
|
||||
query.order_by(AlertHistory.last_seen_at.desc(), AlertHistory.id.desc()).limit(limit).all()
|
||||
)
|
||||
return [_row_to_dict(row) for row in rows]
|
||||
|
||||
|
||||
def _actor_from_auth(auth_context: Optional[Dict[str, Any]]) -> Dict[str, Optional[str]]:
|
||||
auth_context = auth_context or {}
|
||||
user_id = auth_context.get("user_id")
|
||||
if user_id is not None:
|
||||
changed_by = str(user_id)
|
||||
elif auth_context.get("payload", {}).get("sub"):
|
||||
changed_by = str(auth_context["payload"]["sub"])
|
||||
else:
|
||||
changed_by = str(auth_context.get("auth_type") or "unknown")
|
||||
return {
|
||||
"changed_by": changed_by,
|
||||
"auth_type": str(auth_context.get("auth_type") or "unknown"),
|
||||
}
|
||||
|
||||
|
||||
def record_alert_config_change(
|
||||
db: Session,
|
||||
*,
|
||||
key: str,
|
||||
old_value: Optional[str],
|
||||
new_value: Optional[str],
|
||||
auth_context: Optional[Dict[str, Any]] = None,
|
||||
changed_at: Optional[datetime] = None,
|
||||
) -> None:
|
||||
"""Record one sanitized alert/notification setting change."""
|
||||
actor = _actor_from_auth(auth_context)
|
||||
db.add(
|
||||
AlertConfigurationAudit(
|
||||
key=key,
|
||||
old_value=old_value,
|
||||
new_value=new_value,
|
||||
changed_by=actor["changed_by"],
|
||||
auth_type=actor["auth_type"],
|
||||
changed_at=changed_at or datetime.utcnow(),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _config_audit_row_to_dict(row: AlertConfigurationAudit) -> Dict[str, Any]:
|
||||
return {
|
||||
"id": row.id,
|
||||
"key": row.key,
|
||||
"old_value": row.old_value,
|
||||
"new_value": row.new_value,
|
||||
"changed_by": row.changed_by,
|
||||
"auth_type": row.auth_type,
|
||||
"changed_at": row.changed_at.isoformat() if row.changed_at else None,
|
||||
}
|
||||
|
||||
|
||||
def list_alert_config_audit(db: Session, *, limit: int = 50) -> List[Dict[str, Any]]:
|
||||
"""Return recent alert/notification configuration changes."""
|
||||
rows = (
|
||||
db.query(AlertConfigurationAudit)
|
||||
.order_by(AlertConfigurationAudit.changed_at.desc(), AlertConfigurationAudit.id.desc())
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
return [_config_audit_row_to_dict(row) for row in rows]
|
||||
@@ -0,0 +1,292 @@
|
||||
"""Alert rule evaluation for DMARC report data."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from sqlalchemy import case, func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.domain import Domain
|
||||
from app.models.report import DMARCReport, ReportRecord
|
||||
from app.models.setting import Setting
|
||||
from app.services.alert_history import record_alert_evaluation
|
||||
from app.services.notifications import NotificationResult, send_notification
|
||||
from app.services.webhook_events import (
|
||||
EVENT_ALERT_CREATED,
|
||||
EVENT_COMPLIANCE_DROP,
|
||||
EVENT_REPORTS_MISSING,
|
||||
EVENT_SENDER_NEW,
|
||||
enqueue_webhook_event,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _truthy(value: Optional[str], default: bool = True) -> bool:
|
||||
if value in (None, ""):
|
||||
return default
|
||||
return str(value).strip().lower() in {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def _int_setting(value: Optional[str], default: int) -> int:
|
||||
try:
|
||||
return int(value or default)
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
|
||||
def _settings(db: Session) -> Dict[str, Optional[str]]:
|
||||
rows = db.query(Setting).filter(Setting.category == "notifications").all()
|
||||
return {row.key: row.value for row in rows}
|
||||
|
||||
|
||||
def _days_ago_ts(days: int) -> int:
|
||||
return int((datetime.now(timezone.utc) - timedelta(days=days)).timestamp())
|
||||
|
||||
|
||||
def _new_source_alerts(db: Session, window_days: int = 7) -> List[Dict[str, Any]]:
|
||||
cutoff_ts = _days_ago_ts(window_days)
|
||||
previous_sources = {
|
||||
(row.domain, row.source_ip)
|
||||
for row in (
|
||||
db.query(Domain.name.label("domain"), ReportRecord.source_ip.label("source_ip"))
|
||||
.join(DMARCReport, DMARCReport.domain_id == Domain.id)
|
||||
.join(ReportRecord, ReportRecord.report_id == DMARCReport.id)
|
||||
.filter(DMARCReport.begin_date < cutoff_ts)
|
||||
.distinct()
|
||||
.all()
|
||||
)
|
||||
}
|
||||
|
||||
current_sources = (
|
||||
db.query(
|
||||
Domain.name.label("domain"),
|
||||
ReportRecord.source_ip.label("source_ip"),
|
||||
func.sum(ReportRecord.count).label("message_count"),
|
||||
)
|
||||
.join(DMARCReport, DMARCReport.domain_id == Domain.id)
|
||||
.join(ReportRecord, ReportRecord.report_id == DMARCReport.id)
|
||||
.filter(DMARCReport.begin_date >= cutoff_ts)
|
||||
.group_by(Domain.name, ReportRecord.source_ip)
|
||||
.order_by(func.sum(ReportRecord.count).desc())
|
||||
.all()
|
||||
)
|
||||
|
||||
alerts = []
|
||||
for row in current_sources:
|
||||
if (row.domain, row.source_ip) in previous_sources:
|
||||
continue
|
||||
count = int(row.message_count or 0)
|
||||
alerts.append(
|
||||
{
|
||||
"rule": "new_sender_source",
|
||||
"severity": "warning",
|
||||
"domain": row.domain,
|
||||
"source_ip": row.source_ip,
|
||||
"message_count": count,
|
||||
"title": "New sending source",
|
||||
"detail": f"{row.source_ip} first appeared for {row.domain} with {count} messages.",
|
||||
}
|
||||
)
|
||||
return alerts
|
||||
|
||||
|
||||
def _failure_threshold_alerts(
|
||||
db: Session, threshold: int, window_days: int = 1
|
||||
) -> List[Dict[str, Any]]:
|
||||
cutoff_ts = _days_ago_ts(window_days)
|
||||
rows = (
|
||||
db.query(
|
||||
Domain.name.label("domain"),
|
||||
func.sum(ReportRecord.count).label("failed_messages"),
|
||||
)
|
||||
.join(DMARCReport, DMARCReport.domain_id == Domain.id)
|
||||
.join(ReportRecord, ReportRecord.report_id == DMARCReport.id)
|
||||
.filter(DMARCReport.begin_date >= cutoff_ts)
|
||||
.filter(func.coalesce(ReportRecord.dkim, "fail") != "pass")
|
||||
.filter(func.coalesce(ReportRecord.spf, "fail") != "pass")
|
||||
.group_by(Domain.name)
|
||||
.having(func.sum(ReportRecord.count) >= threshold)
|
||||
.all()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"rule": "dmarc_failures_above_threshold",
|
||||
"severity": "error",
|
||||
"domain": row.domain,
|
||||
"failed_messages": int(row.failed_messages or 0),
|
||||
"threshold": threshold,
|
||||
"title": "DMARC failures above threshold",
|
||||
"detail": (
|
||||
f"{row.domain} had {int(row.failed_messages or 0)} DMARC failures "
|
||||
f"in the last {window_days} day(s)."
|
||||
),
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
|
||||
|
||||
def _missing_report_alerts(db: Session, missing_days: int) -> List[Dict[str, Any]]:
|
||||
cutoff_ts = _days_ago_ts(missing_days)
|
||||
latest_rows = (
|
||||
db.query(
|
||||
Domain.name.label("domain"),
|
||||
func.max(DMARCReport.end_date).label("last_report_ts"),
|
||||
)
|
||||
.outerjoin(DMARCReport, DMARCReport.domain_id == Domain.id)
|
||||
.group_by(Domain.id, Domain.name)
|
||||
.all()
|
||||
)
|
||||
alerts = []
|
||||
for row in latest_rows:
|
||||
last_report_ts = int(row.last_report_ts or 0)
|
||||
if last_report_ts >= cutoff_ts:
|
||||
continue
|
||||
alerts.append(
|
||||
{
|
||||
"rule": "missing_reports",
|
||||
"severity": "warning",
|
||||
"domain": row.domain,
|
||||
"missing_days": missing_days,
|
||||
"last_report_at": (
|
||||
datetime.fromtimestamp(last_report_ts, tz=timezone.utc).isoformat()
|
||||
if last_report_ts
|
||||
else None
|
||||
),
|
||||
"title": "Missing DMARC reports",
|
||||
"detail": f"{row.domain} has no DMARC report in the last {missing_days} day(s).",
|
||||
}
|
||||
)
|
||||
return alerts
|
||||
|
||||
|
||||
def _compliance_drop_alerts(
|
||||
db: Session, drop_points: int, window_days: int = 2
|
||||
) -> List[Dict[str, Any]]:
|
||||
cutoff_ts = _days_ago_ts(window_days)
|
||||
rows = (
|
||||
db.query(
|
||||
Domain.name.label("domain"),
|
||||
DMARCReport.begin_date.label("begin_date"),
|
||||
func.sum(ReportRecord.count).label("total"),
|
||||
func.sum(
|
||||
case(
|
||||
(
|
||||
(ReportRecord.dkim == "pass") | (ReportRecord.spf == "pass"),
|
||||
ReportRecord.count,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
).label("passed"),
|
||||
)
|
||||
.join(DMARCReport, DMARCReport.domain_id == Domain.id)
|
||||
.join(ReportRecord, ReportRecord.report_id == DMARCReport.id)
|
||||
.filter(DMARCReport.begin_date >= cutoff_ts)
|
||||
.group_by(Domain.name, DMARCReport.begin_date)
|
||||
.order_by(Domain.name, DMARCReport.begin_date)
|
||||
.all()
|
||||
)
|
||||
|
||||
by_domain: Dict[str, List[Dict[str, float]]] = {}
|
||||
for row in rows:
|
||||
total = int(row.total or 0)
|
||||
passed = int(row.passed or 0)
|
||||
rate = round((passed / total) * 100, 1) if total else 0.0
|
||||
by_domain.setdefault(row.domain, []).append({"date": row.begin_date, "rate": rate})
|
||||
|
||||
alerts = []
|
||||
for domain, points in by_domain.items():
|
||||
if len(points) < 2:
|
||||
continue
|
||||
previous = points[-2]
|
||||
current = points[-1]
|
||||
drop = round(previous["rate"] - current["rate"], 1)
|
||||
if drop < drop_points:
|
||||
continue
|
||||
alerts.append(
|
||||
{
|
||||
"rule": "compliance_drop",
|
||||
"severity": "error" if drop >= 25 else "warning",
|
||||
"domain": domain,
|
||||
"previous_rate": previous["rate"],
|
||||
"current_rate": current["rate"],
|
||||
"drop": drop,
|
||||
"threshold": drop_points,
|
||||
"title": "Compliance dropped",
|
||||
"detail": f"{domain} compliance fell by {drop} percentage points.",
|
||||
}
|
||||
)
|
||||
return alerts
|
||||
|
||||
|
||||
def evaluate_alert_rules(db: Session) -> List[Dict[str, Any]]:
|
||||
"""Evaluate all enabled alert rules against persisted DMARC data."""
|
||||
settings = _settings(db)
|
||||
alerts: List[Dict[str, Any]] = []
|
||||
|
||||
if _truthy(settings.get("notifications.alert_new_sources_enabled")):
|
||||
alerts.extend(_new_source_alerts(db))
|
||||
|
||||
if _truthy(settings.get("notifications.alert_failure_threshold_enabled")):
|
||||
threshold = _int_setting(settings.get("notifications.alert_failure_threshold_count"), 100)
|
||||
alerts.extend(_failure_threshold_alerts(db, threshold))
|
||||
|
||||
if _truthy(settings.get("notifications.alert_missing_reports_enabled")):
|
||||
missing_days = _int_setting(settings.get("notifications.alert_missing_reports_days"), 2)
|
||||
alerts.extend(_missing_report_alerts(db, missing_days))
|
||||
|
||||
if _truthy(settings.get("notifications.alert_compliance_drop_enabled")):
|
||||
drop_points = _int_setting(settings.get("notifications.alert_compliance_drop_points"), 10)
|
||||
alerts.extend(_compliance_drop_alerts(db, drop_points))
|
||||
|
||||
return alerts
|
||||
|
||||
|
||||
def enqueue_alert_webhook_events(db: Session, alerts: List[Dict[str, Any]]) -> None:
|
||||
"""Queue webhook events for alert-rule results without failing alert evaluation."""
|
||||
event_by_rule = {
|
||||
"new_sender_source": EVENT_SENDER_NEW,
|
||||
"missing_reports": EVENT_REPORTS_MISSING,
|
||||
"compliance_drop": EVENT_COMPLIANCE_DROP,
|
||||
}
|
||||
for alert in alerts:
|
||||
rule = alert.get("rule", "alert")
|
||||
event_type = event_by_rule.get(rule, EVENT_ALERT_CREATED)
|
||||
domain = alert.get("domain", "global")
|
||||
idempotency_key = f"{event_type}:{domain}:{rule}:{alert.get('detail', '')}"
|
||||
try:
|
||||
enqueue_webhook_event(
|
||||
db,
|
||||
event_type=event_type,
|
||||
payload=alert,
|
||||
idempotency_key=idempotency_key,
|
||||
)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.warning("Failed to queue alert webhook event: %s", exc)
|
||||
|
||||
|
||||
def send_current_alerts(db: Session) -> Dict[str, Any]:
|
||||
"""Evaluate current alert rules and send one summary notification when needed."""
|
||||
alerts = evaluate_alert_rules(db)
|
||||
record_alert_evaluation(db, alerts)
|
||||
enqueue_alert_webhook_events(db, alerts)
|
||||
if not alerts:
|
||||
return {
|
||||
"alerts": [],
|
||||
"notification": NotificationResult(
|
||||
success=True, skipped=True, message="No alerts."
|
||||
).to_dict(),
|
||||
}
|
||||
|
||||
lines = [f"{alert['title']}: {alert['detail']}" for alert in alerts[:10]]
|
||||
if len(alerts) > 10:
|
||||
lines.append(f"...and {len(alerts) - 10} more alert(s).")
|
||||
result = send_notification(
|
||||
db,
|
||||
title=f"DMARQ alert summary: {len(alerts)} alert(s)",
|
||||
body="\n".join(lines),
|
||||
)
|
||||
return {"alerts": alerts, "notification": result.to_dict()}
|
||||
@@ -0,0 +1,141 @@
|
||||
"""Persistent scoped API token helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Iterable, List, Optional, Set
|
||||
|
||||
import bcrypt
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.api_token import APIToken
|
||||
|
||||
READ_REPORTS_SCOPE = "reports:read"
|
||||
READ_POSTURE_SCOPE = "posture:read"
|
||||
READ_TLS_SCOPE = "tls-reports:read"
|
||||
MCP_READ_SCOPE = "mcp:read"
|
||||
|
||||
PUBLIC_READ_SCOPES = {
|
||||
READ_REPORTS_SCOPE,
|
||||
READ_POSTURE_SCOPE,
|
||||
READ_TLS_SCOPE,
|
||||
MCP_READ_SCOPE,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class CreatedAPIToken:
|
||||
"""Return value for newly created API tokens."""
|
||||
|
||||
token: APIToken
|
||||
secret: str
|
||||
|
||||
|
||||
def normalize_scopes(scopes: Iterable[str]) -> List[str]:
|
||||
"""Normalize and validate requested API token scopes."""
|
||||
normalized = sorted({scope.strip().lower() for scope in scopes if scope and scope.strip()})
|
||||
invalid = [scope for scope in normalized if scope not in PUBLIC_READ_SCOPES]
|
||||
if invalid:
|
||||
raise ValueError(f"Unsupported API token scope: {', '.join(invalid)}")
|
||||
if not normalized:
|
||||
raise ValueError("At least one API token scope is required")
|
||||
return normalized
|
||||
|
||||
|
||||
def scopes_to_string(scopes: Iterable[str]) -> str:
|
||||
"""Serialize scopes for storage."""
|
||||
return ",".join(normalize_scopes(scopes))
|
||||
|
||||
|
||||
def parse_scopes(value: str) -> Set[str]:
|
||||
"""Parse stored scope text into a set."""
|
||||
return {scope.strip().lower() for scope in (value or "").split(",") if scope.strip()}
|
||||
|
||||
|
||||
def generate_public_api_key() -> str:
|
||||
"""Generate an operator-facing API token secret."""
|
||||
return f"dmarq_{secrets.token_urlsafe(32)}"
|
||||
|
||||
|
||||
def hash_api_key(secret: str) -> str:
|
||||
"""Hash an API token for database storage."""
|
||||
return bcrypt.hashpw(secret.encode("utf-8"), bcrypt.gensalt()).decode("utf-8")
|
||||
|
||||
|
||||
def verify_api_key_secret(secret: str, hashed_secret: str) -> bool:
|
||||
"""Return True when a raw API token matches the stored hash."""
|
||||
try:
|
||||
return bcrypt.checkpw(secret.encode("utf-8"), hashed_secret.encode("utf-8"))
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def create_api_token(db: Session, *, name: str, scopes: Iterable[str]) -> CreatedAPIToken:
|
||||
"""Create a persistent API token and return the raw secret once."""
|
||||
clean_name = name.strip()
|
||||
if not clean_name:
|
||||
raise ValueError("Token name is required")
|
||||
secret = generate_public_api_key()
|
||||
token = APIToken(
|
||||
name=clean_name,
|
||||
key_hash=hash_api_key(secret),
|
||||
key_prefix=secret[:12],
|
||||
scopes=scopes_to_string(scopes),
|
||||
active=True,
|
||||
)
|
||||
db.add(token)
|
||||
db.commit()
|
||||
db.refresh(token)
|
||||
return CreatedAPIToken(token=token, secret=secret)
|
||||
|
||||
|
||||
def find_api_token(db: Session, secret: str) -> Optional[APIToken]:
|
||||
"""Return the active token row matching *secret*, if any."""
|
||||
if not secret:
|
||||
return None
|
||||
candidates = (
|
||||
db.query(APIToken)
|
||||
.filter(APIToken.key_prefix == secret[:12], APIToken.active == True) # noqa: E712
|
||||
.all()
|
||||
)
|
||||
for token in candidates:
|
||||
if verify_api_key_secret(secret, token.key_hash):
|
||||
return token
|
||||
return None
|
||||
|
||||
|
||||
def record_api_token_use(db: Session, token: APIToken, *, ip_address: Optional[str]) -> None:
|
||||
"""Persist minimal audit data for a successful API token use."""
|
||||
token.last_used_at = datetime.utcnow()
|
||||
token.last_used_ip = ip_address
|
||||
token.usage_count = int(token.usage_count or 0) + 1
|
||||
db.commit()
|
||||
|
||||
|
||||
def revoke_api_token(db: Session, token_id: int) -> bool:
|
||||
"""Deactivate an API token by id."""
|
||||
token = db.query(APIToken).filter(APIToken.id == token_id).first()
|
||||
if token is None or not token.active:
|
||||
return False
|
||||
token.active = False
|
||||
token.revoked_at = datetime.utcnow()
|
||||
db.commit()
|
||||
return True
|
||||
|
||||
|
||||
def token_to_dict(token: APIToken) -> dict:
|
||||
"""Return an API-safe token representation without the secret hash."""
|
||||
return {
|
||||
"id": token.id,
|
||||
"name": token.name,
|
||||
"key_prefix": token.key_prefix,
|
||||
"scopes": sorted(parse_scopes(token.scopes)),
|
||||
"active": token.active,
|
||||
"created_at": token.created_at.isoformat() if token.created_at else None,
|
||||
"last_used_at": token.last_used_at.isoformat() if token.last_used_at else None,
|
||||
"last_used_ip": token.last_used_ip,
|
||||
"usage_count": token.usage_count,
|
||||
"revoked_at": token.revoked_at.isoformat() if token.revoked_at else None,
|
||||
}
|
||||
@@ -0,0 +1,203 @@
|
||||
"""BIMI DNS posture checks for monitored domains."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List, Optional, Tuple
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.dns_cache import DNSCache
|
||||
from app.services.dns_cache import DEFAULT_DNS_CACHE_TTL_SECONDS
|
||||
from app.services.dns_resolver import BaseDNSProvider
|
||||
|
||||
_CACHE_KEY_PREFIX = "bimi-v1"
|
||||
|
||||
|
||||
@dataclass
|
||||
class BIMIResult:
|
||||
"""Operator-facing BIMI posture evidence."""
|
||||
|
||||
status: str = "fail"
|
||||
selector: str = "default"
|
||||
query_name: str = ""
|
||||
dns_record: Optional[str] = None
|
||||
logo_url: Optional[str] = None
|
||||
certificate_url: Optional[str] = None
|
||||
evidence_url: Optional[str] = None
|
||||
errors: List[str] = field(default_factory=list)
|
||||
warnings: List[str] = field(default_factory=list)
|
||||
|
||||
|
||||
def _utcnow_naive() -> datetime:
|
||||
return datetime.now(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
|
||||
def _is_fresh(row: DNSCache, ttl_seconds: int, now: datetime) -> bool:
|
||||
return row.checked_at >= now - timedelta(seconds=ttl_seconds)
|
||||
|
||||
|
||||
def _result_from_json(value: str) -> BIMIResult:
|
||||
data = json.loads(value)
|
||||
return BIMIResult(
|
||||
status=str(data.get("status") or "fail"),
|
||||
selector=str(data.get("selector") or "default"),
|
||||
query_name=str(data.get("query_name") or ""),
|
||||
dns_record=data.get("dns_record"),
|
||||
logo_url=data.get("logo_url"),
|
||||
certificate_url=data.get("certificate_url"),
|
||||
evidence_url=data.get("evidence_url"),
|
||||
errors=list(data.get("errors") or []),
|
||||
warnings=list(data.get("warnings") or []),
|
||||
)
|
||||
|
||||
|
||||
def _https_url(value: Optional[str]) -> bool:
|
||||
if not value:
|
||||
return False
|
||||
parsed = urlparse(value)
|
||||
return parsed.scheme == "https" and bool(parsed.netloc)
|
||||
|
||||
|
||||
def _tags(record: str) -> dict[str, str]:
|
||||
return {
|
||||
part.split("=", 1)[0].strip().lower(): part.split("=", 1)[1].strip()
|
||||
for part in record.split(";")
|
||||
if "=" in part
|
||||
}
|
||||
|
||||
|
||||
def parse_bimi_record(records: List[str]) -> Tuple[Optional[BIMIResult], List[str], List[str]]:
|
||||
"""Parse BIMI TXT records into a normalized result, warnings, and errors."""
|
||||
bimi_records = [record for record in records if record.lower().startswith("v=bimi1")]
|
||||
if not bimi_records:
|
||||
return None, [], ["No BIMI TXT record was found at the selector."]
|
||||
|
||||
warnings: List[str] = []
|
||||
errors: List[str] = []
|
||||
if len(bimi_records) > 1:
|
||||
warnings.append("Multiple BIMI TXT records were found; publish exactly one.")
|
||||
|
||||
record = bimi_records[0]
|
||||
tags = _tags(record)
|
||||
logo_url = tags.get("l")
|
||||
certificate_url = tags.get("a")
|
||||
if tags.get("v", "").lower() != "bimi1":
|
||||
errors.append("The BIMI TXT record must start with v=BIMI1.")
|
||||
if not logo_url:
|
||||
errors.append("The BIMI TXT record must include an l= HTTPS SVG logo URL.")
|
||||
elif not _https_url(logo_url):
|
||||
errors.append("The BIMI logo URL must use HTTPS.")
|
||||
elif not urlparse(logo_url).path.lower().endswith(".svg"):
|
||||
warnings.append("The BIMI logo URL should point to an SVG file.")
|
||||
|
||||
if certificate_url and not _https_url(certificate_url):
|
||||
errors.append("The BIMI certificate URL must use HTTPS when present.")
|
||||
elif not certificate_url:
|
||||
warnings.append("No BIMI certificate URL is published; some mailbox providers require one.")
|
||||
|
||||
result = BIMIResult(
|
||||
status="pass" if not errors else "fail",
|
||||
dns_record=record,
|
||||
logo_url=logo_url,
|
||||
certificate_url=certificate_url,
|
||||
evidence_url=logo_url,
|
||||
errors=errors,
|
||||
warnings=warnings,
|
||||
)
|
||||
return result, warnings, errors
|
||||
|
||||
|
||||
async def check_bimi(
|
||||
domain: str,
|
||||
provider: BaseDNSProvider,
|
||||
*,
|
||||
selector: str = "default",
|
||||
) -> BIMIResult:
|
||||
"""Resolve and validate the BIMI TXT record for a domain selector."""
|
||||
normalized_selector = (selector or "default").strip().lower()
|
||||
query_name = f"{normalized_selector}._bimi.{domain}"
|
||||
result = BIMIResult(selector=normalized_selector, query_name=query_name)
|
||||
try:
|
||||
records = await provider.lookup_txt(query_name)
|
||||
except LookupError as exc:
|
||||
result.errors.append(f"BIMI DNS lookup failed: {exc}")
|
||||
return result
|
||||
|
||||
parsed, warnings, errors = parse_bimi_record(records)
|
||||
if parsed is None:
|
||||
result.errors.extend(errors)
|
||||
return result
|
||||
parsed.selector = normalized_selector
|
||||
parsed.query_name = query_name
|
||||
parsed.warnings = warnings
|
||||
parsed.errors = errors
|
||||
return parsed
|
||||
|
||||
|
||||
async def check_bimi_cached(
|
||||
db: Session,
|
||||
provider: BaseDNSProvider,
|
||||
domain: str,
|
||||
*,
|
||||
selector: str = "default",
|
||||
ttl_seconds: int = DEFAULT_DNS_CACHE_TTL_SECONDS,
|
||||
refresh: bool = False,
|
||||
) -> Tuple[BIMIResult, bool, datetime]:
|
||||
"""Resolve BIMI posture, reusing the shared DNS cache semantics."""
|
||||
normalized_selector = (selector or "default").strip().lower()
|
||||
cache_key = f"{_CACHE_KEY_PREFIX}:{normalized_selector}"
|
||||
now = _utcnow_naive()
|
||||
provider_name = f"{provider.__class__.__name__}:bimi"
|
||||
row = (
|
||||
db.query(DNSCache)
|
||||
.filter(
|
||||
DNSCache.domain == domain,
|
||||
DNSCache.provider == provider_name,
|
||||
DNSCache.selectors_key == cache_key,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if row and not refresh and _is_fresh(row, ttl_seconds, now):
|
||||
return _result_from_json(row.result_json), True, row.checked_at
|
||||
|
||||
result = await check_bimi(domain, provider, selector=normalized_selector)
|
||||
payload = json.dumps(asdict(result), sort_keys=True, separators=(",", ":"))
|
||||
if row is None:
|
||||
row = DNSCache(
|
||||
domain=domain,
|
||||
provider=provider_name,
|
||||
selectors_key=cache_key,
|
||||
result_json=payload,
|
||||
checked_at=now,
|
||||
)
|
||||
db.add(row)
|
||||
else:
|
||||
row.result_json = payload
|
||||
row.checked_at = now
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except IntegrityError:
|
||||
db.rollback()
|
||||
row = (
|
||||
db.query(DNSCache)
|
||||
.filter(
|
||||
DNSCache.domain == domain,
|
||||
DNSCache.provider == provider_name,
|
||||
DNSCache.selectors_key == cache_key,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if row is None:
|
||||
raise
|
||||
row.result_json = payload
|
||||
row.checked_at = now
|
||||
db.commit()
|
||||
|
||||
db.refresh(row)
|
||||
return result, False, row.checked_at
|
||||
@@ -0,0 +1,433 @@
|
||||
"""Cloudflare DNS discovery, analysis, and change tracking."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
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
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
|
||||
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(timezone.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)
|
||||
workspace = assign_default_workspace_to_unscoped_rows(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, Domain.workspace_id == workspace.id)
|
||||
.first()
|
||||
)
|
||||
if domain is None:
|
||||
db.add(Domain(name=name, active=True, verified=True, workspace_id=workspace.id))
|
||||
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,
|
||||
}
|
||||
@@ -3,7 +3,7 @@ import io
|
||||
import logging
|
||||
import zipfile
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import defusedxml.ElementTree as ET
|
||||
|
||||
@@ -128,76 +128,239 @@ class DMARCParser:
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _strip_namespace(el) -> None:
|
||||
"""Recursively remove XML namespace prefixes from element tags in-place."""
|
||||
if "}" in el.tag:
|
||||
el.tag = el.tag.split("}", 1)[1]
|
||||
for child in el:
|
||||
DMARCParser._strip_namespace(child)
|
||||
|
||||
@staticmethod
|
||||
def _namespace(tag: str) -> str:
|
||||
"""Return the XML namespace from an ElementTree tag."""
|
||||
return tag[1:].split("}", 1)[0] if tag.startswith("{") and "}" in tag else ""
|
||||
|
||||
@staticmethod
|
||||
def _safe_int(value: Any, default: int = 0) -> int:
|
||||
"""Parse integer fields without failing the entire report on bad optional data."""
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
@staticmethod
|
||||
def _text(parent, name: str, default: str = "") -> str:
|
||||
"""Return stripped child text for a parsed XML element."""
|
||||
return (parent.findtext(name, default) or default).strip()
|
||||
|
||||
@staticmethod
|
||||
def _parse_text_list(parent, name: str) -> List[str]:
|
||||
"""Return all non-empty child text values for repeated simple elements."""
|
||||
return [
|
||||
text
|
||||
for text in (DMARCParser._text(child, ".") for child in parent.findall(name))
|
||||
if text
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _collect_extension_values(parent) -> Dict[str, Any]:
|
||||
"""Capture namespaced extension values without coupling to vendor-specific schemas."""
|
||||
values: Dict[str, Any] = {}
|
||||
for child in list(parent):
|
||||
key = child.tag
|
||||
if len(child):
|
||||
values[key] = DMARCParser._collect_extension_values(child)
|
||||
else:
|
||||
values[key] = (child.text or "").strip()
|
||||
return values
|
||||
|
||||
@staticmethod
|
||||
def _extension_value(element) -> Any:
|
||||
"""Return a scalar or nested mapping for a vendor extension element."""
|
||||
if len(element):
|
||||
return DMARCParser._collect_extension_values(element)
|
||||
return (element.text or "").strip()
|
||||
|
||||
@staticmethod
|
||||
def _detect_variant(root, xml_namespace: str) -> dict:
|
||||
"""Identify the aggregate report format variant for debugging/import history."""
|
||||
version = DMARCParser._text(root, "version", "1.0")
|
||||
has_rfc9990_fields = any(
|
||||
root.find(path) is not None
|
||||
for path in (
|
||||
"report_metadata/generator",
|
||||
"policy_published/discovery_method",
|
||||
"policy_published/np",
|
||||
"policy_published/testing",
|
||||
"record/identifiers/envelope_to",
|
||||
)
|
||||
)
|
||||
if xml_namespace or has_rfc9990_fields:
|
||||
variant = "rfc9990"
|
||||
else:
|
||||
variant = "rfc7489-compatible"
|
||||
return {
|
||||
"variant": variant,
|
||||
"schema_version": version,
|
||||
"xml_namespace": xml_namespace,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _parse_metadata(root) -> dict:
|
||||
"""Parse the report_metadata section of a DMARC XML report."""
|
||||
report: dict = {}
|
||||
metadata = root.find("report_metadata")
|
||||
if metadata is not None:
|
||||
report["report_id"] = metadata.findtext("report_id", "")
|
||||
report["org_name"] = metadata.findtext("org_name", "")
|
||||
report["email"] = metadata.findtext("email", "")
|
||||
report["report_id"] = DMARCParser._text(metadata, "report_id")
|
||||
report["org_name"] = DMARCParser._text(metadata, "org_name")
|
||||
report["email"] = DMARCParser._text(metadata, "email")
|
||||
report["extra_contact_info"] = DMARCParser._text(metadata, "extra_contact_info")
|
||||
report["generator"] = DMARCParser._text(metadata, "generator")
|
||||
errors = DMARCParser._parse_text_list(metadata, "error")
|
||||
if errors:
|
||||
report["errors"] = errors
|
||||
|
||||
date_range = metadata.find("date_range")
|
||||
if date_range is not None:
|
||||
begin_ts = int(date_range.findtext("begin", 0))
|
||||
end_ts = int(date_range.findtext("end", 0))
|
||||
begin_ts = DMARCParser._safe_int(date_range.findtext("begin", 0))
|
||||
end_ts = DMARCParser._safe_int(date_range.findtext("end", 0))
|
||||
report["begin_date"] = datetime.fromtimestamp(begin_ts).isoformat()
|
||||
report["end_date"] = datetime.fromtimestamp(end_ts).isoformat()
|
||||
report["begin_timestamp"] = begin_ts
|
||||
report["end_timestamp"] = end_ts
|
||||
return report
|
||||
|
||||
@staticmethod
|
||||
def _parse_policy(root) -> dict:
|
||||
"""Parse policy_published, including RFC 9990 optional fields."""
|
||||
policy = root.find("policy_published")
|
||||
if policy is None:
|
||||
return {}
|
||||
parsed = {
|
||||
"domain": DMARCParser._text(policy, "domain"),
|
||||
"policy": {
|
||||
"p": DMARCParser._text(policy, "p", "none"),
|
||||
"sp": DMARCParser._text(policy, "sp"),
|
||||
"pct": DMARCParser._text(policy, "pct", "100"),
|
||||
"np": DMARCParser._text(policy, "np"),
|
||||
"fo": DMARCParser._text(policy, "fo"),
|
||||
"adkim": DMARCParser._text(policy, "adkim"),
|
||||
"aspf": DMARCParser._text(policy, "aspf"),
|
||||
"testing": DMARCParser._text(policy, "testing"),
|
||||
"discovery_method": DMARCParser._text(policy, "discovery_method"),
|
||||
},
|
||||
}
|
||||
parsed["policy"] = {key: value for key, value in parsed["policy"].items() if value}
|
||||
parsed["policy"].setdefault("pct", "100")
|
||||
return parsed
|
||||
|
||||
@staticmethod
|
||||
def _parse_policy_reasons(policy_evaluated) -> List[dict]:
|
||||
"""Parse policy_evaluated/reason override data."""
|
||||
reasons = []
|
||||
for reason in policy_evaluated.findall("reason"):
|
||||
parsed = {
|
||||
"type": DMARCParser._text(reason, "type"),
|
||||
"comment": DMARCParser._text(reason, "comment"),
|
||||
}
|
||||
if parsed["type"] or parsed["comment"]:
|
||||
reasons.append(parsed)
|
||||
return reasons
|
||||
|
||||
@staticmethod
|
||||
def _parse_row(record_elem) -> dict:
|
||||
"""Parse the record row and policy_evaluated section."""
|
||||
parsed: dict = {}
|
||||
row = record_elem.find("row")
|
||||
if row is None:
|
||||
return parsed
|
||||
|
||||
parsed["source_ip"] = DMARCParser._text(row, "source_ip")
|
||||
parsed["count"] = DMARCParser._safe_int(row.findtext("count", 0))
|
||||
policy_evaluated = row.find("policy_evaluated")
|
||||
if policy_evaluated is None:
|
||||
return parsed
|
||||
|
||||
parsed["disposition"] = DMARCParser._text(policy_evaluated, "disposition", "none")
|
||||
parsed["dkim_result"] = DMARCParser._text(policy_evaluated, "dkim").lower()
|
||||
parsed["spf_result"] = DMARCParser._text(policy_evaluated, "spf").lower()
|
||||
reasons = DMARCParser._parse_policy_reasons(policy_evaluated)
|
||||
if reasons:
|
||||
parsed["policy_override_reasons"] = reasons
|
||||
return parsed
|
||||
|
||||
@staticmethod
|
||||
def _parse_identifiers(record_elem) -> dict:
|
||||
"""Parse identifier fields used for aggregate policy evaluation."""
|
||||
parsed: dict = {}
|
||||
identifiers = record_elem.find("identifiers")
|
||||
if identifiers is None:
|
||||
return parsed
|
||||
parsed["header_from"] = DMARCParser._text(identifiers, "header_from")
|
||||
parsed["envelope_from"] = DMARCParser._text(identifiers, "envelope_from")
|
||||
parsed["envelope_to"] = DMARCParser._text(identifiers, "envelope_to")
|
||||
return parsed
|
||||
|
||||
@staticmethod
|
||||
def _parse_auth_results(record_elem) -> dict:
|
||||
"""Parse uninterpreted DKIM/SPF authentication results."""
|
||||
parsed: dict = {}
|
||||
auth_results = record_elem.find("auth_results")
|
||||
if auth_results is None:
|
||||
return parsed
|
||||
|
||||
spf_entries = [
|
||||
{
|
||||
"domain": DMARCParser._text(spf, "domain"),
|
||||
"scope": DMARCParser._text(spf, "scope"),
|
||||
"result": DMARCParser._text(spf, "result").lower(),
|
||||
"human_result": DMARCParser._text(spf, "human_result"),
|
||||
}
|
||||
for spf in auth_results.findall("spf")
|
||||
]
|
||||
if spf_entries:
|
||||
parsed["spf"] = spf_entries
|
||||
|
||||
dkim_entries = [
|
||||
{
|
||||
"domain": DMARCParser._text(dkim, "domain"),
|
||||
"result": DMARCParser._text(dkim, "result").lower(),
|
||||
"selector": DMARCParser._text(dkim, "selector"),
|
||||
"human_result": DMARCParser._text(dkim, "human_result"),
|
||||
}
|
||||
for dkim in auth_results.findall("dkim")
|
||||
]
|
||||
if dkim_entries:
|
||||
parsed["dkim"] = dkim_entries
|
||||
return parsed
|
||||
|
||||
@staticmethod
|
||||
def _parse_record_extensions(record_elem) -> dict:
|
||||
"""Parse record-level extension elements."""
|
||||
extension_values = {}
|
||||
for child in record_elem:
|
||||
if child.tag not in {"row", "identifiers", "auth_results"}:
|
||||
extension_values[child.tag] = DMARCParser._extension_value(child)
|
||||
return {"extensions": extension_values} if extension_values else {}
|
||||
|
||||
@staticmethod
|
||||
def _parse_record(record_elem) -> dict:
|
||||
"""Parse a single <record> element into a dictionary."""
|
||||
record: dict = {}
|
||||
|
||||
row = record_elem.find("row")
|
||||
if row is not None:
|
||||
record["source_ip"] = row.findtext("source_ip", "")
|
||||
record["count"] = int(row.findtext("count", 0))
|
||||
policy_evaluated = row.find("policy_evaluated")
|
||||
if policy_evaluated is not None:
|
||||
record["disposition"] = policy_evaluated.findtext("disposition", "none")
|
||||
record["dkim_result"] = policy_evaluated.findtext("dkim", "").lower()
|
||||
record["spf_result"] = policy_evaluated.findtext("spf", "").lower()
|
||||
|
||||
identifiers = record_elem.find("identifiers")
|
||||
if identifiers is not None:
|
||||
record["header_from"] = identifiers.findtext("header_from", "")
|
||||
|
||||
auth_results = record_elem.find("auth_results")
|
||||
if auth_results is not None:
|
||||
spf_entries = [
|
||||
{
|
||||
"domain": spf.findtext("domain", ""),
|
||||
"result": spf.findtext("result", "").lower(),
|
||||
}
|
||||
for spf in auth_results.findall("spf")
|
||||
]
|
||||
if spf_entries:
|
||||
record["spf"] = spf_entries
|
||||
|
||||
dkim_entries = [
|
||||
{
|
||||
"domain": dkim.findtext("domain", ""),
|
||||
"result": dkim.findtext("result", "").lower(),
|
||||
"selector": dkim.findtext("selector", ""),
|
||||
}
|
||||
for dkim in auth_results.findall("dkim")
|
||||
]
|
||||
if dkim_entries:
|
||||
record["dkim"] = dkim_entries
|
||||
record.update(DMARCParser._parse_row(record_elem))
|
||||
record.update(DMARCParser._parse_identifiers(record_elem))
|
||||
record.update(DMARCParser._parse_auth_results(record_elem))
|
||||
record.update(DMARCParser._parse_record_extensions(record_elem))
|
||||
|
||||
return record
|
||||
|
||||
@staticmethod
|
||||
def _compute_summary(records: list) -> dict:
|
||||
"""Compute aggregate pass/fail statistics for a list of records."""
|
||||
total_count = sum(r["count"] for r in records)
|
||||
total_count = sum(r.get("count", 0) for r in records)
|
||||
passed_count = sum(
|
||||
r["count"]
|
||||
r.get("count", 0)
|
||||
for r in records
|
||||
if r.get("spf_result") == "pass" or r.get("dkim_result") == "pass"
|
||||
)
|
||||
@@ -216,18 +379,18 @@ class DMARCParser:
|
||||
"""
|
||||
try:
|
||||
root = ET.fromstring(xml_content)
|
||||
xml_namespace = DMARCParser._namespace(root.tag)
|
||||
DMARCParser._strip_namespace(root)
|
||||
|
||||
report = DMARCParser._parse_metadata(root)
|
||||
report.update(DMARCParser._detect_variant(root, xml_namespace))
|
||||
|
||||
# Parse policy published
|
||||
policy = root.find("policy_published")
|
||||
if policy is not None:
|
||||
report["domain"] = policy.findtext("domain", "")
|
||||
report["policy"] = {
|
||||
"p": policy.findtext("p", "none"),
|
||||
"sp": policy.findtext("sp", ""),
|
||||
"pct": policy.findtext("pct", "100"),
|
||||
}
|
||||
report.update(DMARCParser._parse_policy(root))
|
||||
|
||||
extension = root.find("extension")
|
||||
if extension is not None:
|
||||
report["extensions"] = DMARCParser._collect_extension_values(extension)
|
||||
|
||||
# Parse records
|
||||
records = [DMARCParser._parse_record(elem) for elem in root.findall("record")]
|
||||
@@ -236,22 +399,22 @@ class DMARCParser:
|
||||
|
||||
# Log parse results for debugging
|
||||
total_count = report["summary"]["total_count"]
|
||||
logger.info(f"Parsed DMARC report for domain: {report.get('domain')}")
|
||||
logger.info("Parsed DMARC report for domain: %s", report.get("domain"))
|
||||
logger.info("Found %s record entries with %s total messages", len(records), total_count)
|
||||
logger.info(
|
||||
f"Found {len(records)} record entries with {total_count} total messages"
|
||||
)
|
||||
logger.info(
|
||||
f"Messages passed: {report['summary']['passed_count']}, "
|
||||
f"failed: {report['summary']['failed_count']}"
|
||||
"Messages passed: %s, failed: %s",
|
||||
report["summary"]["passed_count"],
|
||||
report["summary"]["failed_count"],
|
||||
)
|
||||
if records:
|
||||
logger.info(
|
||||
f"Sample record - SPF: {records[0].get('spf_result')}, "
|
||||
f"DKIM: {records[0].get('dkim_result')}"
|
||||
"Sample record - SPF: %s, DKIM: %s",
|
||||
records[0].get("spf_result"),
|
||||
records[0].get("dkim_result"),
|
||||
)
|
||||
|
||||
return report
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error parsing DMARC XML: {str(e)}")
|
||||
raise ValueError(f"Error parsing DMARC XML: {str(e)}")
|
||||
logger.error("Error parsing DMARC XML: %s", str(e))
|
||||
raise ValueError(f"Error parsing DMARC XML: {str(e)}") from e
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
"""Database-backed DNS result cache."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from dataclasses import asdict
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List, Tuple
|
||||
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.dns_cache import DNSCache
|
||||
from app.services.dns_resolver import BaseDNSProvider, DomainDNSResult
|
||||
|
||||
DEFAULT_DNS_CACHE_TTL_SECONDS = 900
|
||||
|
||||
|
||||
def _utcnow_naive() -> datetime:
|
||||
return datetime.now(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
|
||||
def _selectors_key(selectors: List[str]) -> str:
|
||||
payload = json.dumps(list(dict.fromkeys(selectors or [])), separators=(",", ":"))
|
||||
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _result_to_json(result: DomainDNSResult) -> str:
|
||||
return json.dumps(asdict(result), sort_keys=True, separators=(",", ":"))
|
||||
|
||||
|
||||
def _result_from_json(value: str) -> DomainDNSResult:
|
||||
data = json.loads(value)
|
||||
return DomainDNSResult(
|
||||
dmarc=bool(data.get("dmarc")),
|
||||
dmarc_record=data.get("dmarc_record"),
|
||||
spf=bool(data.get("spf")),
|
||||
spf_record=data.get("spf_record"),
|
||||
dkim=bool(data.get("dkim")),
|
||||
dkim_selectors=list(data.get("dkim_selectors") or []),
|
||||
dkim_record=data.get("dkim_record"),
|
||||
selectors_checked=list(data.get("selectors_checked") or []),
|
||||
)
|
||||
|
||||
|
||||
def _is_fresh(row: DNSCache, ttl_seconds: int, now: datetime) -> bool:
|
||||
return row.checked_at >= now - timedelta(seconds=ttl_seconds)
|
||||
|
||||
|
||||
async def resolve_domain_dns_cached(
|
||||
db: Session,
|
||||
provider: BaseDNSProvider,
|
||||
domain: str,
|
||||
*,
|
||||
selectors: List[str],
|
||||
ttl_seconds: int = DEFAULT_DNS_CACHE_TTL_SECONDS,
|
||||
refresh: bool = False,
|
||||
) -> Tuple[DomainDNSResult, bool, datetime]:
|
||||
"""Resolve DNS for a domain, reusing a fresh cached result when available."""
|
||||
now = _utcnow_naive()
|
||||
provider_name = provider.__class__.__name__
|
||||
selectors_key = _selectors_key(selectors)
|
||||
row = (
|
||||
db.query(DNSCache)
|
||||
.filter(
|
||||
DNSCache.domain == domain,
|
||||
DNSCache.provider == provider_name,
|
||||
DNSCache.selectors_key == selectors_key,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if row and not refresh and _is_fresh(row, ttl_seconds, now):
|
||||
return _result_from_json(row.result_json), True, row.checked_at
|
||||
|
||||
result = await provider.check_domain(domain, selectors=selectors)
|
||||
payload = _result_to_json(result)
|
||||
if row is None:
|
||||
row = DNSCache(
|
||||
domain=domain,
|
||||
provider=provider_name,
|
||||
selectors_key=selectors_key,
|
||||
result_json=payload,
|
||||
checked_at=now,
|
||||
)
|
||||
db.add(row)
|
||||
else:
|
||||
row.result_json = payload
|
||||
row.checked_at = now
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except IntegrityError:
|
||||
db.rollback()
|
||||
row = (
|
||||
db.query(DNSCache)
|
||||
.filter(
|
||||
DNSCache.domain == domain,
|
||||
DNSCache.provider == provider_name,
|
||||
DNSCache.selectors_key == selectors_key,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if row is None:
|
||||
raise
|
||||
row.result_json = payload
|
||||
row.checked_at = now
|
||||
db.commit()
|
||||
|
||||
db.refresh(row)
|
||||
return result, False, row.checked_at
|
||||
@@ -0,0 +1,484 @@
|
||||
"""
|
||||
DNS resolver service for DMARC, SPF, DKIM, and PTR record lookups.
|
||||
|
||||
Provides an extensible provider architecture so that DNS data can be fetched
|
||||
either via the system resolver (dnspython) or via the Cloudflare DNS API for
|
||||
future Cloudflare integration.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import ipaddress
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _sanitize_for_log(value: str) -> str:
|
||||
"""Remove newline and carriage-return characters to prevent log injection."""
|
||||
return value.replace("\r", "").replace("\n", "")
|
||||
|
||||
|
||||
def _ip_to_arpa_name(ip: str) -> str:
|
||||
"""Convert an IP address string to its reverse-DNS ARPA lookup name.
|
||||
|
||||
E.g. ``"1.2.3.4"`` → ``"4.3.2.1.in-addr.arpa"``
|
||||
``"2001:db8::1"`` → ``"...ip6.arpa"``
|
||||
|
||||
Raises ``ValueError`` for invalid IP address strings.
|
||||
"""
|
||||
addr = ipaddress.ip_address(ip)
|
||||
if isinstance(addr, ipaddress.IPv4Address):
|
||||
parts = ip.split(".")
|
||||
return ".".join(reversed(parts)) + ".in-addr.arpa"
|
||||
# IPv6: expand, strip colons, reverse nibbles
|
||||
expanded = addr.exploded.replace(":", "")
|
||||
return ".".join(reversed(expanded)) + ".ip6.arpa"
|
||||
|
||||
|
||||
# Well-known DKIM selectors tried when no selectors are configured
|
||||
COMMON_DKIM_SELECTORS: List[str] = [
|
||||
"default",
|
||||
"google",
|
||||
"mail",
|
||||
"selector1",
|
||||
"selector2",
|
||||
"dkim",
|
||||
"k1",
|
||||
"key1",
|
||||
"mta",
|
||||
"email",
|
||||
"smtp",
|
||||
"s1",
|
||||
"s2",
|
||||
"pm",
|
||||
"mandrill",
|
||||
"sendgrid",
|
||||
]
|
||||
|
||||
# Seconds to wait for a single DNS query before giving up
|
||||
DNS_TIMEOUT: float = 5.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class DomainDNSResult:
|
||||
"""Aggregated DNS authentication record results for one domain."""
|
||||
|
||||
dmarc: bool = False
|
||||
dmarc_record: Optional[str] = None
|
||||
spf: bool = False
|
||||
spf_record: Optional[str] = None
|
||||
dkim: bool = False
|
||||
# All selectors that resolved to a valid DKIM record (may be multiple)
|
||||
dkim_selectors: List[str] = field(default_factory=list)
|
||||
dkim_record: Optional[str] = None
|
||||
# Track which selectors were tried so callers can surface this information
|
||||
selectors_checked: List[str] = field(default_factory=list)
|
||||
|
||||
|
||||
class BaseDNSProvider(ABC):
|
||||
"""
|
||||
Abstract base class for DNS providers.
|
||||
|
||||
Subclasses implement ``lookup_txt`` and inherit the higher-level helper
|
||||
methods for DMARC, SPF, and DKIM checks so that provider-specific
|
||||
differences stay confined to a single method.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def lookup_txt(self, name: str) -> List[str]:
|
||||
"""Return TXT record strings for *name*.
|
||||
|
||||
Raises ``LookupError`` on failure (NXDOMAIN, timeout, network error
|
||||
etc.). Returns an empty list when the name exists but has no TXT
|
||||
records.
|
||||
"""
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# High-level record checks built on top of lookup_txt
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def check_dmarc(self, domain: str) -> Tuple[bool, Optional[str]]:
|
||||
"""Return *(found, record_string)* for the domain's DMARC TXT record."""
|
||||
try:
|
||||
records = await self.lookup_txt(f"_dmarc.{domain}")
|
||||
for record in records:
|
||||
if record.lower().startswith("v=dmarc1"):
|
||||
return True, record
|
||||
except LookupError as exc:
|
||||
logger.debug("DMARC lookup failed for %s: %s", _sanitize_for_log(domain), exc)
|
||||
return False, None
|
||||
|
||||
async def check_spf(self, domain: str) -> Tuple[bool, Optional[str]]:
|
||||
"""Return *(found, record_string)* for the domain's SPF TXT record."""
|
||||
try:
|
||||
records = await self.lookup_txt(domain)
|
||||
for record in records:
|
||||
if record.lower().startswith("v=spf1"):
|
||||
return True, record
|
||||
except LookupError as exc:
|
||||
logger.debug("SPF lookup failed for %s: %s", _sanitize_for_log(domain), exc)
|
||||
return False, None
|
||||
|
||||
async def lookup_ptr(self, ip: str) -> Optional[str]:
|
||||
"""Return the PTR (reverse DNS) hostname for *ip*, or ``None`` if unavailable.
|
||||
|
||||
The base implementation always returns ``None``. Concrete providers
|
||||
override this to perform an actual DNS PTR lookup so that existing
|
||||
test doubles (which only implement ``lookup_txt``) keep working without
|
||||
modification.
|
||||
"""
|
||||
return None
|
||||
|
||||
async def check_dkim(
|
||||
self, domain: str, selectors: List[str]
|
||||
) -> Tuple[bool, List[str], Optional[str]]:
|
||||
"""Return *(found, matching_selectors, first_record_string)* for all working DKIM selectors.
|
||||
|
||||
All selectors in *selectors* are checked and every one that resolves to
|
||||
a valid DKIM TXT record is collected. The boolean is ``True`` when at
|
||||
least one selector resolved. *first_record_string* is the record text
|
||||
for the first matching selector (useful for display purposes).
|
||||
"""
|
||||
matching_selectors: List[str] = []
|
||||
first_record: Optional[str] = None
|
||||
for selector in selectors:
|
||||
try:
|
||||
records = await self.lookup_txt(f"{selector}._domainkey.{domain}")
|
||||
for record in records:
|
||||
if "v=dkim1" in record.lower() or "p=" in record.lower():
|
||||
matching_selectors.append(selector)
|
||||
if first_record is None:
|
||||
first_record = record
|
||||
break
|
||||
except LookupError as exc:
|
||||
logger.debug(
|
||||
"DKIM lookup failed for selector=%s domain=%s: %s",
|
||||
selector,
|
||||
_sanitize_for_log(domain),
|
||||
exc,
|
||||
)
|
||||
return bool(matching_selectors), matching_selectors, first_record
|
||||
|
||||
async def check_domain(
|
||||
self, domain: str, selectors: Optional[List[str]] = None
|
||||
) -> DomainDNSResult:
|
||||
"""Run DMARC, SPF, and DKIM checks concurrently for *domain*.
|
||||
|
||||
*selectors* are tried first; common well-known selectors are appended
|
||||
as a fallback so that a domain with no explicitly configured selectors
|
||||
can still be verified.
|
||||
"""
|
||||
# Deduplicate while preserving priority order (manual selectors first)
|
||||
all_selectors: List[str] = list(selectors or [])
|
||||
for s in COMMON_DKIM_SELECTORS:
|
||||
if s not in all_selectors:
|
||||
all_selectors.append(s)
|
||||
|
||||
dmarc_coro = self.check_dmarc(domain)
|
||||
spf_coro = self.check_spf(domain)
|
||||
dkim_coro = self.check_dkim(domain, all_selectors)
|
||||
|
||||
(dmarc_ok, dmarc_record), (spf_ok, spf_record), (dkim_ok, dkim_sels, dkim_record) = (
|
||||
await asyncio.gather(dmarc_coro, spf_coro, dkim_coro)
|
||||
)
|
||||
|
||||
return DomainDNSResult(
|
||||
dmarc=dmarc_ok,
|
||||
dmarc_record=dmarc_record,
|
||||
spf=spf_ok,
|
||||
spf_record=spf_record,
|
||||
dkim=dkim_ok,
|
||||
dkim_selectors=dkim_sels,
|
||||
dkim_record=dkim_record,
|
||||
selectors_checked=all_selectors,
|
||||
)
|
||||
|
||||
|
||||
class SystemDNSProvider(BaseDNSProvider):
|
||||
"""DNS provider that resolves records via the system resolver using dnspython."""
|
||||
|
||||
async def lookup_txt(self, name: str) -> List[str]:
|
||||
"""Resolve TXT records using dnspython's async resolver."""
|
||||
# Import here so the module can be imported even if dnspython is absent
|
||||
# (tests can mock this method directly without needing the library).
|
||||
import dns.asyncresolver # type: ignore[import]
|
||||
import dns.exception # type: ignore[import]
|
||||
|
||||
try:
|
||||
answers = await dns.asyncresolver.resolve(
|
||||
name, "TXT", lifetime=DNS_TIMEOUT, raise_on_no_answer=False
|
||||
)
|
||||
result: List[str] = []
|
||||
if answers:
|
||||
for rdata in answers:
|
||||
for string in rdata.strings:
|
||||
result.append(string.decode("utf-8", errors="replace"))
|
||||
return result
|
||||
except dns.exception.DNSException as exc:
|
||||
raise LookupError(f"TXT lookup failed for {name}: {exc}") from exc
|
||||
|
||||
async def lookup_ptr(self, ip: str) -> Optional[str]:
|
||||
"""Resolve a PTR record for *ip* via the system resolver."""
|
||||
import dns.asyncresolver # type: ignore[import]
|
||||
import dns.exception # type: ignore[import]
|
||||
|
||||
try:
|
||||
ptr_name = _ip_to_arpa_name(ip)
|
||||
answers = await dns.asyncresolver.resolve(
|
||||
ptr_name, "PTR", lifetime=DNS_TIMEOUT, raise_on_no_answer=False
|
||||
)
|
||||
if answers:
|
||||
for rdata in answers:
|
||||
return str(rdata).rstrip(".")
|
||||
except (dns.exception.DNSException, ValueError):
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
class CloudflareDNSProvider(BaseDNSProvider):
|
||||
"""DNS provider using Cloudflare DoH and, when configured, the REST API.
|
||||
|
||||
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
|
||||
CLOUDFLARE_API_BASE: str = "https://api.cloudflare.com/client/v4"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_token: Optional[str] = None,
|
||||
zone_id: Optional[str] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
api_token:
|
||||
Cloudflare API token. Required for zone discovery and managed
|
||||
DNS record reads; not needed for read-only DoH lookups.
|
||||
zone_id:
|
||||
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]
|
||||
|
||||
params = {"name": name, "type": "TXT"}
|
||||
headers = {"Accept": "application/dns-json"}
|
||||
try:
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.get(
|
||||
self.CLOUDFLARE_DOH_URL,
|
||||
params=params,
|
||||
headers=headers,
|
||||
timeout=DNS_TIMEOUT,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
records: List[str] = []
|
||||
for answer in data.get("Answer", []):
|
||||
if answer.get("type") == 16: # TXT record type
|
||||
# Cloudflare wraps TXT values in double-quotes
|
||||
txt = answer.get("data", "").strip('"')
|
||||
records.append(txt)
|
||||
return records
|
||||
except (httpx.RequestError, httpx.HTTPStatusError, httpx.TimeoutException) as exc:
|
||||
raise LookupError(f"Cloudflare DoH lookup failed for {name}: {exc}") from exc
|
||||
|
||||
async def lookup_ptr(self, ip: str) -> Optional[str]:
|
||||
"""Resolve a PTR record for *ip* via Cloudflare's DoH endpoint."""
|
||||
import httpx # type: ignore[import]
|
||||
|
||||
try:
|
||||
ptr_name = _ip_to_arpa_name(ip)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
params = {"name": ptr_name, "type": "PTR"}
|
||||
headers = {"Accept": "application/dns-json"}
|
||||
try:
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.get(
|
||||
self.CLOUDFLARE_DOH_URL,
|
||||
params=params,
|
||||
headers=headers,
|
||||
timeout=DNS_TIMEOUT,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
for answer in data.get("Answer", []):
|
||||
if answer.get("type") == 12: # PTR record type
|
||||
return answer.get("data", "").rstrip(".")
|
||||
except (httpx.RequestError, httpx.HTTPStatusError, httpx.TimeoutException):
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def _decrypt_setting_value(value: Optional[str]) -> Optional[str]:
|
||||
if not value:
|
||||
return value
|
||||
try:
|
||||
from app.core.credential_encryption import decrypt_secret
|
||||
|
||||
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()
|
||||
|
||||
|
||||
def extract_dmarc_policy(dmarc_record: Optional[str]) -> Optional[str]:
|
||||
"""Parse the *p=* tag from a DMARC TXT record string.
|
||||
|
||||
Returns the policy value (e.g. ``"none"``, ``"quarantine"``,
|
||||
``"reject"``) or ``None`` if the record is absent or unparsable.
|
||||
"""
|
||||
if not dmarc_record:
|
||||
return None
|
||||
for part in dmarc_record.split(";"):
|
||||
part = part.strip()
|
||||
if part.lower().startswith("p="):
|
||||
return part[2:].strip().lower()
|
||||
return None
|
||||
@@ -0,0 +1,271 @@
|
||||
import json
|
||||
import re
|
||||
from collections import Counter
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple
|
||||
|
||||
from app.models.report import ForensicReport
|
||||
|
||||
|
||||
AUTH_RESULT_PATTERN = re.compile(r"\b(dkim|spf|dmarc)=([a-zA-Z0-9_-]+)", re.IGNORECASE)
|
||||
HEADER_DOMAIN_PATTERN = re.compile(r"\bheader\.d=([^;\s]+)", re.IGNORECASE)
|
||||
MAILFROM_DOMAIN_PATTERN = re.compile(r"\bsmtp\.mailfrom=([^;\s]+)", re.IGNORECASE)
|
||||
PRIORITY_ORDER = {"high": 3, "medium": 2, "low": 1}
|
||||
|
||||
|
||||
def _clean_value(value: Any) -> str:
|
||||
return str(value or "").strip()
|
||||
|
||||
|
||||
def _normalize(value: Any) -> str:
|
||||
return _clean_value(value).lower()
|
||||
|
||||
|
||||
def _feedback_headers(row: ForensicReport) -> Dict[str, Any]:
|
||||
if not row.feedback_headers:
|
||||
return {}
|
||||
try:
|
||||
parsed = json.loads(row.feedback_headers)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return {}
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
|
||||
|
||||
def _parse_authentication_results(value: str) -> Dict[str, str]:
|
||||
results: Dict[str, str] = {}
|
||||
for mechanism, result in AUTH_RESULT_PATTERN.findall(value or ""):
|
||||
results[mechanism.lower()] = result.lower()
|
||||
return results
|
||||
|
||||
|
||||
def _first_match(pattern: re.Pattern[str], value: str) -> str:
|
||||
match = pattern.search(value or "")
|
||||
return match.group(1).lower().strip(".,") if match else ""
|
||||
|
||||
|
||||
def _failure_kind(row: ForensicReport, auth_results: Dict[str, str]) -> str:
|
||||
reported = _normalize(row.auth_failure)
|
||||
if reported in {"dkim", "spf", "dmarc", "both"}:
|
||||
return reported
|
||||
failed = {name for name, result in auth_results.items() if result in {"fail", "softfail"}}
|
||||
if {"dkim", "spf"}.issubset(failed):
|
||||
return "both"
|
||||
for mechanism in ("dmarc", "dkim", "spf"):
|
||||
if mechanism in failed:
|
||||
return mechanism
|
||||
return reported or "unknown"
|
||||
|
||||
|
||||
def _priority(row: ForensicReport, failure_kind: str) -> str:
|
||||
delivery = _normalize(row.delivery_result)
|
||||
if delivery in {"reject", "quarantine"}:
|
||||
return "high"
|
||||
if failure_kind in {"both", "dmarc"}:
|
||||
return "high"
|
||||
if failure_kind in {"dkim", "spf"}:
|
||||
return "medium"
|
||||
return "low"
|
||||
|
||||
|
||||
def _diagnosis(failure_kind: str, auth_results: Dict[str, str], delivery_result: str) -> str:
|
||||
delivery = _normalize(delivery_result)
|
||||
rejected = delivery in {"reject", "quarantine"}
|
||||
suffix = " The receiver enforced the failure." if rejected else ""
|
||||
if failure_kind == "both":
|
||||
return "Both DKIM and SPF failed, so DMARC could not find an aligned pass." + suffix
|
||||
if failure_kind == "dmarc":
|
||||
return "DMARC failed after the receiver evaluated DKIM and SPF alignment." + suffix
|
||||
if failure_kind == "dkim":
|
||||
if auth_results.get("spf") == "pass":
|
||||
return "DKIM failed while SPF passed; focus on DKIM signing and alignment." + suffix
|
||||
return "DKIM failed for the reported message sample." + suffix
|
||||
if failure_kind == "spf":
|
||||
if auth_results.get("dkim") == "pass":
|
||||
return (
|
||||
"SPF failed while DKIM passed; focus on SPF authorization and alignment." + suffix
|
||||
)
|
||||
return "SPF failed for the reported message sample." + suffix
|
||||
return "The receiver reported an authentication failure, but did not include a clear mechanism."
|
||||
|
||||
|
||||
def _recommendations(
|
||||
failure_kind: str,
|
||||
auth_results: Dict[str, str],
|
||||
source_ip: str,
|
||||
reported_domain: str,
|
||||
) -> List[str]:
|
||||
actions: List[str] = []
|
||||
if failure_kind in {"dkim", "both", "dmarc"}:
|
||||
actions.append(
|
||||
"Confirm the sending system signs mail with a DKIM domain aligned to the visible From domain."
|
||||
)
|
||||
actions.append(
|
||||
"Check recent DKIM key, selector, and canonicalization changes for this sender."
|
||||
)
|
||||
if failure_kind in {"spf", "both", "dmarc"}:
|
||||
actions.append(
|
||||
"Verify the source IP or provider include is authorized in the domain SPF record."
|
||||
)
|
||||
actions.append(
|
||||
"Review forwarding paths, because forwarding commonly breaks SPF while preserving DKIM."
|
||||
)
|
||||
if auth_results.get("spf") == "pass" and failure_kind == "dkim":
|
||||
actions.append(
|
||||
"If SPF is aligned and passing, this may be a DKIM-only repair rather than a sender authorization issue."
|
||||
)
|
||||
if auth_results.get("dkim") == "pass" and failure_kind == "spf":
|
||||
actions.append(
|
||||
"If DKIM is aligned and passing, treat SPF repair as lower risk before changing DMARC policy."
|
||||
)
|
||||
if source_ip:
|
||||
actions.append(
|
||||
f"Compare {source_ip} with known mail sources for {reported_domain or 'this domain'}."
|
||||
)
|
||||
actions.append(
|
||||
"Keep using redacted forensic metadata; do not import or retain message bodies for this investigation."
|
||||
)
|
||||
return actions
|
||||
|
||||
|
||||
def _signals(
|
||||
row: ForensicReport,
|
||||
feedback_headers: Dict[str, Any],
|
||||
auth_results: Dict[str, str],
|
||||
header_domain: str,
|
||||
mailfrom_domain: str,
|
||||
) -> List[str]:
|
||||
signals = []
|
||||
if row.source_ip:
|
||||
signals.append(f"Source IP: {row.source_ip}")
|
||||
if row.reported_domain:
|
||||
signals.append(f"Reported domain: {row.reported_domain}")
|
||||
if row.auth_failure:
|
||||
signals.append(f"Failure: {row.auth_failure}")
|
||||
if row.delivery_result:
|
||||
signals.append(f"Delivery result: {row.delivery_result}")
|
||||
if header_domain:
|
||||
signals.append(f"DKIM header domain: {header_domain}")
|
||||
if mailfrom_domain:
|
||||
signals.append(f"SPF mail-from domain: {mailfrom_domain}")
|
||||
identity_alignment = _clean_value(feedback_headers.get("identity_alignment"))
|
||||
if identity_alignment:
|
||||
signals.append(f"Identity alignment: {identity_alignment}")
|
||||
for mechanism, result in sorted(auth_results.items()):
|
||||
signals.append(f"{mechanism.upper()} result: {result}")
|
||||
return signals
|
||||
|
||||
|
||||
def analyze_forensic_report(row: ForensicReport) -> Dict[str, Any]:
|
||||
"""Build a privacy-preserving operator analysis for one forensic sample."""
|
||||
feedback_headers = _feedback_headers(row)
|
||||
auth_results = _parse_authentication_results(row.authentication_results or "")
|
||||
header_domain = _first_match(
|
||||
HEADER_DOMAIN_PATTERN, row.authentication_results or ""
|
||||
) or _normalize(feedback_headers.get("dkim_domain"))
|
||||
mailfrom_domain = _first_match(MAILFROM_DOMAIN_PATTERN, row.authentication_results or "")
|
||||
failure_kind = _failure_kind(row, auth_results)
|
||||
priority = _priority(row, failure_kind)
|
||||
reported_domain = _clean_value(row.reported_domain or (row.domain.name if row.domain else ""))
|
||||
source_ip = _clean_value(row.source_ip)
|
||||
|
||||
return {
|
||||
"id": row.id,
|
||||
"report_id": row.report_id,
|
||||
"domain": reported_domain,
|
||||
"source_ip": source_ip,
|
||||
"auth_failure": failure_kind,
|
||||
"delivery_result": _clean_value(row.delivery_result),
|
||||
"priority": priority,
|
||||
"diagnosis": _diagnosis(failure_kind, auth_results, row.delivery_result or ""),
|
||||
"recommendations": _recommendations(
|
||||
failure_kind,
|
||||
auth_results,
|
||||
source_ip,
|
||||
reported_domain,
|
||||
),
|
||||
"signals": _signals(row, feedback_headers, auth_results, header_domain, mailfrom_domain),
|
||||
"authentication_results": auth_results,
|
||||
"dkim_domain": header_domain,
|
||||
"mail_from_domain": mailfrom_domain,
|
||||
"privacy_note": "Analysis uses redacted headers and metadata only; message bodies are not stored.",
|
||||
}
|
||||
|
||||
|
||||
def _group_key(row: ForensicReport) -> Tuple[str, str, str, str]:
|
||||
return (
|
||||
_clean_value(row.reported_domain or (row.domain.name if row.domain else "")) or "unknown",
|
||||
_clean_value(row.source_ip) or "unknown",
|
||||
_normalize(row.auth_failure) or "unknown",
|
||||
_normalize(row.delivery_result) or "unknown",
|
||||
)
|
||||
|
||||
|
||||
def _latest(left: Optional[datetime], right: Optional[datetime]) -> Optional[datetime]:
|
||||
if left is None:
|
||||
return right
|
||||
if right is None:
|
||||
return left
|
||||
return max(left, right)
|
||||
|
||||
|
||||
def summarize_forensic_samples(rows: Iterable[ForensicReport]) -> Dict[str, Any]:
|
||||
"""Summarize forensic samples into investigation groups and top examples."""
|
||||
reports = list(rows)
|
||||
analyses = [analyze_forensic_report(row) for row in reports]
|
||||
priority_counts = Counter(item["priority"] for item in analyses)
|
||||
failure_counts = Counter(item["auth_failure"] for item in analyses)
|
||||
result_counts = Counter(_normalize(row.delivery_result) or "unknown" for row in reports)
|
||||
grouped: Dict[Tuple[str, str, str, str], Dict[str, Any]] = {}
|
||||
|
||||
for row, analysis in zip(reports, analyses):
|
||||
key = _group_key(row)
|
||||
group = grouped.setdefault(
|
||||
key,
|
||||
{
|
||||
"key": "|".join(key),
|
||||
"domain": key[0],
|
||||
"source_ip": key[1],
|
||||
"auth_failure": analysis["auth_failure"],
|
||||
"delivery_result": key[3],
|
||||
"count": 0,
|
||||
"priority": analysis["priority"],
|
||||
"latest_arrival": None,
|
||||
"diagnosis": analysis["diagnosis"],
|
||||
"recommendations": analysis["recommendations"][:3],
|
||||
},
|
||||
)
|
||||
group["count"] += 1
|
||||
group["latest_arrival"] = _latest(
|
||||
group["latest_arrival"], row.arrival_date or row.processed_at
|
||||
)
|
||||
if PRIORITY_ORDER[analysis["priority"]] > PRIORITY_ORDER[group["priority"]]:
|
||||
group["priority"] = analysis["priority"]
|
||||
group["diagnosis"] = analysis["diagnosis"]
|
||||
group["recommendations"] = analysis["recommendations"][:3]
|
||||
|
||||
groups = sorted(
|
||||
grouped.values(),
|
||||
key=lambda item: (
|
||||
PRIORITY_ORDER[item["priority"]],
|
||||
item["count"],
|
||||
item["latest_arrival"] or datetime.min,
|
||||
),
|
||||
reverse=True,
|
||||
)
|
||||
for group in groups:
|
||||
if group["latest_arrival"] is not None:
|
||||
group["latest_arrival"] = group["latest_arrival"].isoformat()
|
||||
|
||||
samples = sorted(
|
||||
analyses,
|
||||
key=lambda item: (PRIORITY_ORDER[item["priority"]], item["id"] or 0),
|
||||
reverse=True,
|
||||
)
|
||||
return {
|
||||
"total": len(reports),
|
||||
"priority_counts": dict(priority_counts),
|
||||
"failure_counts": dict(failure_counts),
|
||||
"result_counts": dict(result_counts),
|
||||
"groups": groups,
|
||||
"samples": samples,
|
||||
}
|
||||
@@ -0,0 +1,237 @@
|
||||
import hashlib
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
from email import message_from_bytes
|
||||
from email.message import Message
|
||||
from email.parser import Parser
|
||||
from email.utils import getaddresses, parsedate_to_datetime
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from app.services.forensic_redaction import (
|
||||
ForensicRedactionPolicy,
|
||||
normalize_forensic_redaction_policy,
|
||||
redact_forensic_text,
|
||||
)
|
||||
|
||||
|
||||
MAX_FORENSIC_REPORT_SIZE = 10 * 1024 * 1024
|
||||
|
||||
|
||||
def _coerce_text(value: Any) -> str:
|
||||
if value is None:
|
||||
return ""
|
||||
if isinstance(value, bytes):
|
||||
return value.decode("utf-8", errors="replace")
|
||||
return str(value)
|
||||
|
||||
|
||||
def _clean(
|
||||
value: Any,
|
||||
*,
|
||||
redact: bool = True,
|
||||
redaction_policy: Optional[ForensicRedactionPolicy] = None,
|
||||
) -> str:
|
||||
text = " ".join(_coerce_text(value).replace("\r", " ").replace("\n", " ").split())
|
||||
return redact_text(text, redaction_policy=redaction_policy) if redact else text
|
||||
|
||||
|
||||
def redact_text(
|
||||
value: str,
|
||||
*,
|
||||
redaction_policy: Optional[ForensicRedactionPolicy] = None,
|
||||
) -> str:
|
||||
"""Redact email local-parts and long opaque tokens from forensic metadata."""
|
||||
return redact_forensic_text(value, redaction_policy)
|
||||
|
||||
|
||||
def _header(
|
||||
msg: Optional[Message],
|
||||
name: str,
|
||||
*,
|
||||
redact: bool = True,
|
||||
redaction_policy: Optional[ForensicRedactionPolicy] = None,
|
||||
) -> str:
|
||||
return _clean(
|
||||
msg.get(name, "") if msg is not None else "",
|
||||
redact=redact,
|
||||
redaction_policy=redaction_policy,
|
||||
)
|
||||
|
||||
|
||||
def _payload_text(part: Message) -> str:
|
||||
payload = part.get_payload(decode=True)
|
||||
if payload is not None:
|
||||
charset = part.get_content_charset() or "utf-8"
|
||||
return payload.decode(charset, errors="replace")
|
||||
payload_value = part.get_payload()
|
||||
if isinstance(payload_value, list):
|
||||
return ""
|
||||
return _coerce_text(payload_value)
|
||||
|
||||
|
||||
def _message_part_payload(part: Message) -> Optional[Message]:
|
||||
payload = part.get_payload()
|
||||
if isinstance(payload, list) and payload:
|
||||
return payload[0]
|
||||
return None
|
||||
|
||||
|
||||
def _parse_feedback_headers(text: str) -> Message:
|
||||
return Parser().parsestr(text or "")
|
||||
|
||||
|
||||
def _domain_from_address(value: str) -> str:
|
||||
addresses = getaddresses([value])
|
||||
for _, addr in addresses:
|
||||
if "@" in addr:
|
||||
return addr.rsplit("@", 1)[-1].lower()
|
||||
return ""
|
||||
|
||||
|
||||
def _parse_datetime(value: str) -> Optional[datetime]:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
parsed = parsedate_to_datetime(value)
|
||||
except (TypeError, ValueError, IndexError):
|
||||
return None
|
||||
if parsed is None:
|
||||
return None
|
||||
if parsed.tzinfo is not None:
|
||||
parsed = parsed.astimezone(timezone.utc)
|
||||
return parsed.replace(tzinfo=None)
|
||||
|
||||
|
||||
def _message_id_hash(value: str) -> str:
|
||||
cleaned = _clean(value, redact=False)
|
||||
if not cleaned:
|
||||
return ""
|
||||
return hashlib.sha256(cleaned.encode("utf-8")).hexdigest()[:24]
|
||||
|
||||
|
||||
class ForensicParser:
|
||||
"""Parse DMARC forensic/failure report emails without retaining message bodies."""
|
||||
|
||||
@staticmethod
|
||||
def is_forensic_report(msg: Message) -> bool:
|
||||
if msg.get_content_type() == "multipart/report":
|
||||
report_type = (msg.get_param("report-type") or "").lower()
|
||||
if report_type == "feedback-report":
|
||||
return True
|
||||
|
||||
for part in msg.walk():
|
||||
content_type = part.get_content_type().lower()
|
||||
if content_type == "message/feedback-report":
|
||||
return True
|
||||
if content_type == "text/rfc822-headers" and "dmarc" in _payload_text(part).lower():
|
||||
return True
|
||||
|
||||
subject = _header(msg, "Subject", redact=False).lower()
|
||||
return "dmarc" in subject and any(
|
||||
term in subject for term in ("failure", "forensic", "ruf")
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def parse_bytes(
|
||||
cls,
|
||||
content: bytes,
|
||||
*,
|
||||
message_id_hint: Optional[str] = None,
|
||||
redaction_policy: Optional[ForensicRedactionPolicy] = None,
|
||||
) -> Dict[str, Any]:
|
||||
if len(content) > MAX_FORENSIC_REPORT_SIZE:
|
||||
raise ValueError("Forensic report is too large")
|
||||
if not content:
|
||||
raise ValueError("Forensic report is empty")
|
||||
|
||||
msg = message_from_bytes(content)
|
||||
if not cls.is_forensic_report(msg):
|
||||
raise ValueError("Email is not a DMARC forensic report")
|
||||
|
||||
redaction_policy = normalize_forensic_redaction_policy(redaction_policy)
|
||||
feedback = None
|
||||
original_headers = None
|
||||
|
||||
for part in msg.walk():
|
||||
content_type = part.get_content_type().lower()
|
||||
if content_type == "message/feedback-report":
|
||||
feedback = _message_part_payload(part) or _parse_feedback_headers(
|
||||
_payload_text(part)
|
||||
)
|
||||
elif content_type == "text/rfc822-headers":
|
||||
original_headers = _parse_feedback_headers(_payload_text(part))
|
||||
elif content_type == "message/rfc822" and original_headers is None:
|
||||
original_headers = _message_part_payload(part)
|
||||
|
||||
feedback = feedback or msg
|
||||
reported_domain = (
|
||||
_header(feedback, "Reported-Domain", redact=False)
|
||||
or _header(feedback, "DKIM-Domain", redact=False)
|
||||
or _domain_from_address(_header(feedback, "Original-Mail-From", redact=False))
|
||||
or _domain_from_address(_header(original_headers, "From", redact=False))
|
||||
).lower()
|
||||
source_ip = _header(feedback, "Source-IP", redact=False)
|
||||
auth_failure = _header(feedback, "Auth-Failure", redact=False)
|
||||
original_message_id = _header(original_headers, "Message-ID", redact=False)
|
||||
top_message_id = _header(msg, "Message-ID", redact=False)
|
||||
|
||||
report_id = (
|
||||
_clean(message_id_hint, redact=False)
|
||||
or _message_id_hash(top_message_id)
|
||||
or _message_id_hash(original_message_id)
|
||||
or hashlib.sha256(content).hexdigest()[:24]
|
||||
)
|
||||
if not report_id.startswith("ruf-"):
|
||||
report_id = f"ruf-{report_id}"
|
||||
|
||||
source_email = _header(msg, "From", redaction_policy=redaction_policy)
|
||||
arrival_date = _parse_datetime(_header(feedback, "Arrival-Date", redact=False))
|
||||
|
||||
details = {
|
||||
"identity_alignment": _header(feedback, "Identity-Alignment", redact=False),
|
||||
"dkim_domain": _header(feedback, "DKIM-Domain", redact=False),
|
||||
"spf_dns": _header(feedback, "SPF-DNS", redact=False),
|
||||
"reported_uri": _header(
|
||||
feedback,
|
||||
"Reported-URI",
|
||||
redaction_policy=redaction_policy,
|
||||
),
|
||||
}
|
||||
details = {key: value for key, value in details.items() if value}
|
||||
|
||||
return {
|
||||
"report_id": report_id,
|
||||
"source_email": source_email,
|
||||
"feedback_type": _header(feedback, "Feedback-Type", redact=False) or "auth-failure",
|
||||
"user_agent": _header(feedback, "User-Agent", redaction_policy=redaction_policy),
|
||||
"version": _header(feedback, "Version", redact=False),
|
||||
"reported_domain": reported_domain,
|
||||
"source_ip": source_ip,
|
||||
"auth_failure": auth_failure,
|
||||
"delivery_result": _header(feedback, "Delivery-Result", redact=False),
|
||||
"arrival_date": arrival_date,
|
||||
"authentication_results": _header(
|
||||
feedback,
|
||||
"Authentication-Results",
|
||||
redaction_policy=redaction_policy,
|
||||
),
|
||||
"original_mail_from": _header(
|
||||
feedback,
|
||||
"Original-Mail-From",
|
||||
redaction_policy=redaction_policy,
|
||||
),
|
||||
"original_from": _header(
|
||||
original_headers,
|
||||
"From",
|
||||
redaction_policy=redaction_policy,
|
||||
),
|
||||
"original_to": _header(original_headers, "To", redaction_policy=redaction_policy),
|
||||
"original_subject": _header(
|
||||
original_headers,
|
||||
"Subject",
|
||||
redaction_policy=redaction_policy,
|
||||
),
|
||||
"original_message_id": _message_id_hash(original_message_id),
|
||||
"original_date": _header(original_headers, "Date", redact=False),
|
||||
"feedback_headers": json.dumps(details, sort_keys=True) if details else None,
|
||||
}
|
||||
@@ -0,0 +1,148 @@
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.domain import Domain
|
||||
from app.models.report import ForensicReport
|
||||
from app.services.forensic_redaction import ForensicRedactionPolicy, redact_forensic_value
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
from app.utils.domain_validator import DomainValidationError, validate_domain
|
||||
|
||||
|
||||
def forensic_report_exists(db: Session, report_id: str) -> bool:
|
||||
"""Return True when a forensic report ID is already persisted."""
|
||||
if not str(report_id or "").strip():
|
||||
return False
|
||||
return (
|
||||
db.query(ForensicReport.id).filter(ForensicReport.report_id == report_id).first()
|
||||
is not None
|
||||
)
|
||||
|
||||
|
||||
def _domain_for_report(db: Session, domain_name: Optional[str]) -> Optional[Domain]:
|
||||
if not domain_name:
|
||||
return None
|
||||
normalized = domain_name.lower().strip(".")
|
||||
is_valid, _, error_code = validate_domain(normalized, check_dns=False)
|
||||
if not is_valid and error_code != DomainValidationError.DNS_RESOLUTION_FAILED:
|
||||
return None
|
||||
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db, commit=False)
|
||||
domain = (
|
||||
db.query(Domain)
|
||||
.filter(Domain.name == normalized, Domain.workspace_id == workspace.id)
|
||||
.first()
|
||||
)
|
||||
if domain is None:
|
||||
domain = Domain(name=normalized, workspace_id=workspace.id)
|
||||
db.add(domain)
|
||||
db.flush()
|
||||
return domain
|
||||
|
||||
|
||||
def save_forensic_report(db: Session, report: Dict[str, Any]) -> tuple[ForensicReport, bool]:
|
||||
"""Persist a parsed forensic report.
|
||||
|
||||
Returns ``(row, created)``. The caller owns the transaction and should
|
||||
commit after related work has completed.
|
||||
"""
|
||||
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()
|
||||
if existing is not None:
|
||||
return existing, False
|
||||
|
||||
domain = _domain_for_report(db, report.get("reported_domain"))
|
||||
feedback_headers = report.get("feedback_headers")
|
||||
if isinstance(feedback_headers, dict):
|
||||
feedback_headers = json.dumps(feedback_headers, sort_keys=True)
|
||||
|
||||
row = ForensicReport(
|
||||
domain_id=domain.id if domain else None,
|
||||
report_id=report_id,
|
||||
source_email=report.get("source_email"),
|
||||
feedback_type=report.get("feedback_type"),
|
||||
user_agent=report.get("user_agent"),
|
||||
version=report.get("version"),
|
||||
reported_domain=report.get("reported_domain"),
|
||||
source_ip=report.get("source_ip"),
|
||||
auth_failure=report.get("auth_failure"),
|
||||
delivery_result=report.get("delivery_result"),
|
||||
arrival_date=report.get("arrival_date"),
|
||||
authentication_results=report.get("authentication_results"),
|
||||
original_mail_from=report.get("original_mail_from"),
|
||||
original_from=report.get("original_from"),
|
||||
original_to=report.get("original_to"),
|
||||
original_subject=report.get("original_subject"),
|
||||
original_message_id=report.get("original_message_id"),
|
||||
original_date=report.get("original_date"),
|
||||
feedback_headers=feedback_headers,
|
||||
)
|
||||
db.add(row)
|
||||
try:
|
||||
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
|
||||
|
||||
|
||||
_REDACTABLE_RESPONSE_FIELDS = {
|
||||
"source_email",
|
||||
"user_agent",
|
||||
"authentication_results",
|
||||
"original_mail_from",
|
||||
"original_from",
|
||||
"original_to",
|
||||
"original_subject",
|
||||
"feedback_headers",
|
||||
}
|
||||
|
||||
|
||||
def forensic_report_to_dict(
|
||||
row: ForensicReport,
|
||||
*,
|
||||
redaction_policy: Optional[ForensicRedactionPolicy] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""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 = {}
|
||||
|
||||
result = {
|
||||
"id": row.id,
|
||||
"report_id": row.report_id,
|
||||
"domain": row.domain.name if row.domain else row.reported_domain,
|
||||
"reported_domain": row.reported_domain,
|
||||
"source_email": row.source_email,
|
||||
"feedback_type": row.feedback_type,
|
||||
"user_agent": row.user_agent,
|
||||
"version": row.version,
|
||||
"source_ip": row.source_ip,
|
||||
"auth_failure": row.auth_failure,
|
||||
"delivery_result": row.delivery_result,
|
||||
"arrival_date": row.arrival_date.isoformat() if row.arrival_date else None,
|
||||
"authentication_results": row.authentication_results,
|
||||
"original_mail_from": row.original_mail_from,
|
||||
"original_from": row.original_from,
|
||||
"original_to": row.original_to,
|
||||
"original_subject": row.original_subject,
|
||||
"original_message_id": row.original_message_id,
|
||||
"original_date": row.original_date,
|
||||
"feedback_headers": feedback_headers,
|
||||
"processed_at": row.processed_at.isoformat() if row.processed_at else None,
|
||||
}
|
||||
if redaction_policy is not None:
|
||||
for field in _REDACTABLE_RESPONSE_FIELDS:
|
||||
result[field] = redact_forensic_value(result.get(field), redaction_policy)
|
||||
return result
|
||||
@@ -0,0 +1,106 @@
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.setting import Setting
|
||||
|
||||
|
||||
FORENSIC_REDACTION_MODE_KEY = "forensics.redaction_mode"
|
||||
FORENSIC_REDACT_LONG_TOKENS_KEY = "forensics.redact_long_tokens_enabled"
|
||||
DEFAULT_FORENSIC_REDACTION_MODE = "balanced"
|
||||
FORENSIC_REDACTION_MODES = {"balanced", "domain_only", "strict"}
|
||||
|
||||
_EMAIL_RE = re.compile(
|
||||
r"\b([A-Z0-9._%+\-*]{1,64})@([A-Z0-9.-]+\.[A-Z]{2,})\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
_LONG_TOKEN_RE = re.compile(r"\b[A-Za-z0-9_./+=-]{28,}\b")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ForensicRedactionPolicy:
|
||||
"""Privacy policy for forensic report metadata."""
|
||||
|
||||
mode: str = DEFAULT_FORENSIC_REDACTION_MODE
|
||||
redact_long_tokens: bool = True
|
||||
|
||||
|
||||
def _truthy(value: Optional[str], *, default: bool = True) -> bool:
|
||||
if value is None:
|
||||
return default
|
||||
return str(value).strip().lower() in {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def normalize_forensic_redaction_policy(
|
||||
policy: Optional[ForensicRedactionPolicy] = None,
|
||||
*,
|
||||
mode: Optional[str] = None,
|
||||
redact_long_tokens: Optional[bool] = None,
|
||||
) -> ForensicRedactionPolicy:
|
||||
requested_mode = (mode if mode is not None else policy.mode if policy else "").strip().lower()
|
||||
if requested_mode not in FORENSIC_REDACTION_MODES:
|
||||
requested_mode = DEFAULT_FORENSIC_REDACTION_MODE
|
||||
return ForensicRedactionPolicy(
|
||||
mode=requested_mode,
|
||||
redact_long_tokens=(
|
||||
policy.redact_long_tokens
|
||||
if redact_long_tokens is None and policy is not None
|
||||
else bool(True if redact_long_tokens is None else redact_long_tokens)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def get_forensic_redaction_policy(db: Optional[Session]) -> ForensicRedactionPolicy:
|
||||
"""Load the current forensic redaction policy from persisted settings."""
|
||||
if db is None:
|
||||
return ForensicRedactionPolicy()
|
||||
|
||||
rows = (
|
||||
db.query(Setting.key, Setting.value)
|
||||
.filter(Setting.key.in_([FORENSIC_REDACTION_MODE_KEY, FORENSIC_REDACT_LONG_TOKENS_KEY]))
|
||||
.all()
|
||||
)
|
||||
values = {key: value for key, value in rows}
|
||||
return normalize_forensic_redaction_policy(
|
||||
mode=values.get(FORENSIC_REDACTION_MODE_KEY),
|
||||
redact_long_tokens=_truthy(values.get(FORENSIC_REDACT_LONG_TOKENS_KEY), default=True),
|
||||
)
|
||||
|
||||
|
||||
def redact_forensic_text(
|
||||
value: str,
|
||||
policy: Optional[ForensicRedactionPolicy] = None,
|
||||
) -> str:
|
||||
"""Redact forensic metadata according to the selected privacy policy."""
|
||||
policy = normalize_forensic_redaction_policy(policy)
|
||||
|
||||
def _redact_email(match: re.Match[str]) -> str:
|
||||
local = match.group(1)
|
||||
domain = match.group(2).lower()
|
||||
if policy.mode == "strict":
|
||||
return "[redacted-email]"
|
||||
if policy.mode == "domain_only":
|
||||
return f"***@{domain}"
|
||||
prefix = local[:2] if len(local) > 2 else local[:1]
|
||||
return f"{prefix}***@{domain}"
|
||||
|
||||
redacted = _EMAIL_RE.sub(_redact_email, value)
|
||||
if policy.redact_long_tokens:
|
||||
redacted = _LONG_TOKEN_RE.sub("[redacted-token]", redacted)
|
||||
return redacted
|
||||
|
||||
|
||||
def redact_forensic_value(
|
||||
value: Any,
|
||||
policy: Optional[ForensicRedactionPolicy] = None,
|
||||
) -> Any:
|
||||
"""Redact strings inside a response value while preserving container shape."""
|
||||
if isinstance(value, str):
|
||||
return redact_forensic_text(value, policy)
|
||||
if isinstance(value, dict):
|
||||
return {key: redact_forensic_value(item, policy) for key, item in value.items()}
|
||||
if isinstance(value, list):
|
||||
return [redact_forensic_value(item, policy) for item in value]
|
||||
return value
|
||||
@@ -0,0 +1,543 @@
|
||||
"""
|
||||
Gmail API client for retrieving DMARC reports.
|
||||
|
||||
Connects to Gmail via OAuth 2.0, searches for emails that are likely to
|
||||
contain DMARC aggregate-report attachments, and processes any new ones.
|
||||
Already-ingested message IDs are tracked so the same email is never
|
||||
processed twice (no messages are modified or deleted).
|
||||
"""
|
||||
|
||||
import base64
|
||||
import email
|
||||
import logging
|
||||
from typing import Any, Dict, List, Optional
|
||||
from urllib.parse import urlencode
|
||||
|
||||
import httpx
|
||||
from google.auth.transport.requests import Request
|
||||
from google.oauth2.credentials import Credentials
|
||||
from googleapiclient.discovery import build
|
||||
from googleapiclient.errors import HttpError
|
||||
|
||||
from app.services.dmarc_parser import DMARCParser
|
||||
from app.services.forensic_parser import ForensicParser
|
||||
from app.services.forensic_persistence import forensic_report_exists, save_forensic_report
|
||||
from app.services.forensic_redaction import get_forensic_redaction_policy
|
||||
from app.services.mail_connector import (
|
||||
append_import_detail,
|
||||
connector_failure_stats,
|
||||
dump_ingested_ids,
|
||||
initial_import_stats,
|
||||
load_ingested_ids,
|
||||
sanitize_connector_error,
|
||||
)
|
||||
from app.services.report_persistence import report_exists, save_parsed_report
|
||||
from app.services.report_store import ReportStore
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OAuth2 scopes – read-only access to Gmail messages is all we need
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
GMAIL_SCOPES = [
|
||||
"https://www.googleapis.com/auth/gmail.readonly",
|
||||
]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Gmail search query used to find emails likely containing DMARC reports.
|
||||
#
|
||||
# Strategy:
|
||||
# • Require at least one attachment whose name ends in .zip, .gz, or .xml
|
||||
# (the three formats used by virtually every DMARC sender).
|
||||
# • Additionally require *either* a keyword in the subject that DMARC senders
|
||||
# use, or an envelope-from that belongs to a well-known DMARC reporting
|
||||
# address. This keeps false-positive rates low while catching reports
|
||||
# from providers that don't follow naming conventions perfectly.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
DMARC_GMAIL_QUERY = (
|
||||
"((has:attachment (filename:zip OR filename:gz OR filename:xml)) "
|
||||
'OR subject:"DMARC failure" OR subject:"failure report" OR subject:forensic OR subject:ruf) '
|
||||
"(subject:dmarc OR subject:report OR subject:rua OR subject:submitter "
|
||||
'OR subject:"aggregate report" OR subject:"domain report" '
|
||||
'OR subject:"report domain" OR from:dmarc OR from:dmarc-noreply '
|
||||
"OR from:noreply-dmarc-support OR from:reports OR from:postmaster)"
|
||||
)
|
||||
|
||||
# How many message results to fetch per API page
|
||||
_PAGE_SIZE = 100
|
||||
RETRYABLE_MESSAGE_FAILURE = -1
|
||||
|
||||
|
||||
class GmailClient:
|
||||
"""
|
||||
Client for retrieving DMARC reports from a Gmail account via the Gmail API.
|
||||
|
||||
OAuth2 tokens are accepted at construction time and auto-refreshed when
|
||||
expired. The caller is responsible for persisting any refreshed tokens
|
||||
returned by :meth:`get_refreshed_tokens`.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client_id: str,
|
||||
client_secret: str,
|
||||
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,
|
||||
refresh_token=refresh_token,
|
||||
token_uri="https://oauth2.googleapis.com/token",
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
scopes=GMAIL_SCOPES,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Public helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def get_refreshed_tokens(self) -> Optional[Dict[str, str]]:
|
||||
"""
|
||||
Return updated tokens if the google-auth library has refreshed them.
|
||||
|
||||
Call this after :meth:`fetch_reports` and persist any non-None result
|
||||
so the next run doesn't need an extra refresh round-trip.
|
||||
"""
|
||||
current = self.credentials.token
|
||||
if current and current != self._initial_access_token:
|
||||
result: Dict[str, str] = {"access_token": current}
|
||||
if self.credentials.refresh_token:
|
||||
result["refresh_token"] = self.credentials.refresh_token
|
||||
return result
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# OAuth2 helpers (static / class methods used by the endpoint layer)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def build_authorization_url(
|
||||
client_id: str,
|
||||
redirect_uri: str,
|
||||
state: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Construct the Google OAuth2 authorization URL.
|
||||
|
||||
Requests offline access so a refresh token is issued, and forces
|
||||
the consent screen so the refresh token is always returned even if
|
||||
the user has authorised this app before.
|
||||
"""
|
||||
params: Dict[str, str] = {
|
||||
"client_id": client_id,
|
||||
"response_type": "code",
|
||||
"scope": " ".join(GMAIL_SCOPES),
|
||||
"redirect_uri": redirect_uri,
|
||||
"access_type": "offline",
|
||||
"prompt": "consent",
|
||||
}
|
||||
if state:
|
||||
params["state"] = state
|
||||
return "https://accounts.google.com/o/oauth2/v2/auth?" + urlencode(params)
|
||||
|
||||
@staticmethod
|
||||
def exchange_code_for_tokens(
|
||||
client_id: str,
|
||||
client_secret: str,
|
||||
code: str,
|
||||
redirect_uri: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Synchronously exchange an authorization code for access+refresh tokens.
|
||||
|
||||
Returns the raw JSON from Google's token endpoint. The caller
|
||||
should check for ``access_token`` in the result before using it.
|
||||
|
||||
Raises:
|
||||
ValueError: if Google returns a non-200 response.
|
||||
"""
|
||||
resp = httpx.post(
|
||||
"https://oauth2.googleapis.com/token",
|
||||
data={
|
||||
"code": code,
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
"redirect_uri": redirect_uri,
|
||||
"grant_type": "authorization_code",
|
||||
},
|
||||
)
|
||||
if resp.status_code != 200:
|
||||
raise ValueError(f"Token exchange failed ({resp.status_code}): {resp.text}")
|
||||
return resp.json()
|
||||
|
||||
@staticmethod
|
||||
def get_gmail_email(access_token: str) -> Optional[str]:
|
||||
"""
|
||||
Return the email address associated with an access token.
|
||||
|
||||
Uses the OAuth2 userinfo endpoint. Returns None on failure.
|
||||
"""
|
||||
try:
|
||||
resp = httpx.get(
|
||||
"https://www.googleapis.com/oauth2/v2/userinfo",
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
)
|
||||
if resp.status_code == 200:
|
||||
return resp.json().get("email")
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Failed to fetch Gmail email address: %s", exc)
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Core fetching logic
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def fetch_reports(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Search Gmail for DMARC report emails and ingest any new ones.
|
||||
|
||||
Emails that have already been ingested (tracked via
|
||||
``already_ingested_ids``) are silently skipped. No messages are
|
||||
modified or deleted.
|
||||
|
||||
Returns:
|
||||
A dict with keys ``success``, ``processed``, ``reports_found``,
|
||||
``new_domains``, ``errors``, and ``new_ingested_ids`` (the IDs
|
||||
added in this run so the caller can persist them).
|
||||
"""
|
||||
stats = initial_import_stats()
|
||||
|
||||
try:
|
||||
service = self._build_service()
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Gmail API: failed to build service: %s", exc)
|
||||
return connector_failure_stats(stats, "Failed to initialize Gmail API.", error=exc)
|
||||
|
||||
try:
|
||||
message_ids = self._list_dmarc_message_ids(service)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Gmail API: failed to list messages: %s", exc)
|
||||
return connector_failure_stats(stats, "Failed to list Gmail messages.", error=exc)
|
||||
|
||||
domains_before = set(self.report_store.get_domains())
|
||||
|
||||
for msg_id in message_ids:
|
||||
if msg_id in self.already_ingested_ids:
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="already_ingested_message",
|
||||
message_id=msg_id,
|
||||
)
|
||||
continue
|
||||
|
||||
stats["processed"] += 1
|
||||
found = self._process_message(service, msg_id, stats)
|
||||
if found >= 0:
|
||||
# Track it even when no report is found so we don't re-examine
|
||||
# unrelated messages on every poll. Retryable failures return -1.
|
||||
stats["new_ingested_ids"].append(msg_id)
|
||||
self.already_ingested_ids.append(msg_id)
|
||||
|
||||
domains_after = set(self.report_store.get_domains())
|
||||
stats["new_domains"] = list(domains_after - domains_before)
|
||||
return stats
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Private helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _build_service(self):
|
||||
"""Build (and auto-refresh if needed) the Gmail API service object."""
|
||||
if self.credentials.expired and self.credentials.refresh_token:
|
||||
try:
|
||||
self.credentials.refresh(Request())
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Gmail token refresh failed: %s", exc)
|
||||
raise
|
||||
|
||||
return build("gmail", "v1", credentials=self.credentials, cache_discovery=False)
|
||||
|
||||
def _list_dmarc_message_ids(self, service) -> List[str]:
|
||||
"""Return all Gmail message IDs matching the DMARC search query."""
|
||||
ids: List[str] = []
|
||||
page_token: Optional[str] = None
|
||||
|
||||
while True:
|
||||
kwargs: Dict[str, Any] = {
|
||||
"userId": "me",
|
||||
"q": DMARC_GMAIL_QUERY,
|
||||
"maxResults": _PAGE_SIZE,
|
||||
}
|
||||
if page_token:
|
||||
kwargs["pageToken"] = page_token
|
||||
|
||||
try:
|
||||
result = service.users().messages().list(**kwargs).execute()
|
||||
except HttpError as exc:
|
||||
logger.error("Gmail API list error: %s", exc)
|
||||
raise
|
||||
|
||||
for msg in result.get("messages", []):
|
||||
ids.append(msg["id"])
|
||||
|
||||
page_token = result.get("nextPageToken")
|
||||
if not page_token:
|
||||
break
|
||||
|
||||
return ids
|
||||
|
||||
@staticmethod
|
||||
def _append_detail(stats: dict, **detail: str) -> None:
|
||||
"""Append a compact attachment/message outcome to the import stats."""
|
||||
append_import_detail(stats, **detail)
|
||||
|
||||
def _process_message(self, service, msg_id: str, stats: dict) -> int:
|
||||
"""
|
||||
Download a Gmail message and process any DMARC-report attachments.
|
||||
|
||||
Returns the number of DMARC reports found in this message.
|
||||
"""
|
||||
try:
|
||||
msg_data = (
|
||||
service.users().messages().get(userId="me", id=msg_id, format="raw").execute()
|
||||
)
|
||||
except HttpError as exc:
|
||||
logger.error("Gmail API: failed to fetch message %s: %s", msg_id, exc)
|
||||
stats["errors"].append(sanitize_connector_error(f"Failed to fetch message {msg_id}"))
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="error",
|
||||
reason="message_fetch_failed",
|
||||
message_id=msg_id,
|
||||
)
|
||||
return 0
|
||||
|
||||
raw_bytes = base64.urlsafe_b64decode(msg_data.get("raw", ""))
|
||||
msg = email.message_from_bytes(raw_bytes)
|
||||
if ForensicParser.is_forensic_report(msg):
|
||||
return self._process_forensic_message(raw_bytes, stats, message_id=msg_id)
|
||||
return self._process_attachments(msg, stats, message_id=msg_id)
|
||||
|
||||
@staticmethod
|
||||
def _decode_part_filename(part: email.message.Message) -> str:
|
||||
"""Return the decoded filename for a MIME part (handles RFC 2047 encoding)."""
|
||||
from email.header import decode_header
|
||||
|
||||
raw_name = part.get_filename() or ""
|
||||
decoded_parts = []
|
||||
for fragment, charset in decode_header(raw_name):
|
||||
if isinstance(fragment, bytes):
|
||||
decoded_parts.append(fragment.decode(charset or "utf-8", errors="replace"))
|
||||
else:
|
||||
decoded_parts.append(fragment)
|
||||
return "".join(decoded_parts)
|
||||
|
||||
@staticmethod
|
||||
def _is_dmarc_attachment(filename: str) -> bool:
|
||||
"""Return True if *filename* looks like a DMARC aggregate-report file."""
|
||||
lower = filename.lower()
|
||||
return (
|
||||
lower.endswith(".xml")
|
||||
or lower.endswith(".zip")
|
||||
or lower.endswith(".gz")
|
||||
or lower.endswith(".gzip")
|
||||
)
|
||||
|
||||
def _store_report_if_new(self, report: Dict[str, Any]) -> bool:
|
||||
"""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)
|
||||
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
|
||||
|
||||
def _process_forensic_message(
|
||||
self,
|
||||
raw_bytes: bytes,
|
||||
stats: dict,
|
||||
message_id: Optional[str] = None,
|
||||
) -> int:
|
||||
"""Parse and persist one DMARC forensic report message."""
|
||||
try:
|
||||
report = ForensicParser.parse_bytes(
|
||||
raw_bytes,
|
||||
message_id_hint=message_id,
|
||||
redaction_policy=get_forensic_redaction_policy(self.db),
|
||||
)
|
||||
report_id = str(report.get("report_id", ""))
|
||||
domain = str(report.get("reported_domain") or "unknown")
|
||||
|
||||
if self.db is None:
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="forensic_report_requires_database",
|
||||
message_id=message_id,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
return RETRYABLE_MESSAGE_FAILURE
|
||||
|
||||
if forensic_report_exists(self.db, report_id):
|
||||
stats["duplicate_forensic_reports"] = stats.get("duplicate_forensic_reports", 0) + 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="duplicate",
|
||||
reason="duplicate_forensic_report",
|
||||
message_id=message_id,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
return 0
|
||||
|
||||
_row, created = save_forensic_report(self.db, report)
|
||||
if not created:
|
||||
stats["duplicate_forensic_reports"] = stats.get("duplicate_forensic_reports", 0) + 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="duplicate",
|
||||
reason="duplicate_forensic_report",
|
||||
message_id=message_id,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
return 0
|
||||
|
||||
stats["forensic_reports_found"] = stats.get("forensic_reports_found", 0) + 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="imported",
|
||||
reason="forensic_report",
|
||||
message_id=message_id,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
return 1
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Failed to parse Gmail forensic report %s: %s", message_id, exc)
|
||||
stats.setdefault("errors", []).append(
|
||||
sanitize_connector_error(f"Failed to parse forensic report {message_id}: {exc}")
|
||||
)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="error",
|
||||
reason="forensic_parse_failed",
|
||||
message_id=message_id,
|
||||
error=str(exc),
|
||||
)
|
||||
return RETRYABLE_MESSAGE_FAILURE
|
||||
|
||||
def _process_attachments(
|
||||
self,
|
||||
msg: email.message.Message,
|
||||
stats: dict,
|
||||
message_id: Optional[str] = None,
|
||||
) -> int:
|
||||
"""Walk a parsed email message and extract DMARC report attachments."""
|
||||
reports_found = 0
|
||||
|
||||
for part in msg.walk():
|
||||
filename = self._decode_part_filename(part)
|
||||
if not filename:
|
||||
continue
|
||||
|
||||
disposition = part.get_content_disposition()
|
||||
if disposition not in ("attachment", None):
|
||||
continue
|
||||
|
||||
if not self._is_dmarc_attachment(filename):
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="unsupported_attachment",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
)
|
||||
continue
|
||||
|
||||
content = part.get_payload(decode=True)
|
||||
if not content:
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="empty_attachment",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
report = DMARCParser.parse_file(content, filename)
|
||||
domain = str(report.get("domain", "unknown"))
|
||||
report_id = str(report.get("report_id", ""))
|
||||
if self._store_report_if_new(report):
|
||||
stats["reports_found"] += 1
|
||||
reports_found += 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="imported",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
else:
|
||||
stats["duplicate_reports"] = stats.get("duplicate_reports", 0) + 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="duplicate",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Failed to parse DMARC attachment %s: %s", filename, exc)
|
||||
stats["errors"].append(
|
||||
sanitize_connector_error(f"Failed to parse {filename}: {exc}")
|
||||
)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="error",
|
||||
reason="parse_failed",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
error=str(exc),
|
||||
)
|
||||
|
||||
return reports_found
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Convenience: load / save ingested IDs from/to the JSON text column
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def load_ingested_ids(json_text: Optional[str]) -> List[str]:
|
||||
"""Deserialise the gmail_ingested_ids text column into a list."""
|
||||
return load_ingested_ids(json_text)
|
||||
|
||||
@staticmethod
|
||||
def dump_ingested_ids(ids: List[str]) -> str:
|
||||
"""Serialise the list of ingested IDs back to a JSON string."""
|
||||
return dump_ingested_ids(ids)
|
||||
@@ -3,10 +3,19 @@ import imaplib
|
||||
import logging
|
||||
from datetime import datetime, timedelta
|
||||
from email.header import decode_header
|
||||
from typing import Any, Dict, Tuple
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.services.dmarc_parser import DMARCParser
|
||||
from app.services.forensic_parser import ForensicParser
|
||||
from app.services.forensic_persistence import forensic_report_exists, save_forensic_report
|
||||
from app.services.forensic_redaction import get_forensic_redaction_policy
|
||||
from app.services.mail_connector import (
|
||||
append_import_detail,
|
||||
initial_import_stats,
|
||||
sanitize_connector_error,
|
||||
)
|
||||
from app.services.report_persistence import report_exists, save_parsed_report
|
||||
from app.services.report_store import ReportStore
|
||||
|
||||
# Setup logger
|
||||
@@ -18,13 +27,15 @@ class IMAPClient:
|
||||
Client for retrieving DMARC reports from an IMAP mailbox
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
def __init__( # pylint: disable=too-many-positional-arguments,too-many-arguments
|
||||
self,
|
||||
server: str = None,
|
||||
port: int = None,
|
||||
username: str = None,
|
||||
password: str = None,
|
||||
delete_emails: bool = False,
|
||||
delete_emails: Optional[bool] = None,
|
||||
folder: str = None,
|
||||
db: Any = None,
|
||||
):
|
||||
"""
|
||||
Initialize the IMAP client with credentials
|
||||
@@ -34,21 +45,36 @@ class IMAPClient:
|
||||
port: IMAP server port (if None, uses settings)
|
||||
username: IMAP username (if None, uses settings)
|
||||
password: IMAP password (if None, uses settings)
|
||||
delete_emails: Whether to delete emails after processing (default: False)
|
||||
delete_emails: Whether to delete emails after successful report imports.
|
||||
If omitted, uses DELETE_IMPORTED_EMAILS from settings.
|
||||
folder: IMAP mailbox folder to read (if None, uses settings or INBOX)
|
||||
db: Optional SQLAlchemy session used to persist imported reports
|
||||
"""
|
||||
settings = get_settings()
|
||||
settings_folder = getattr(settings, "IMAP_FOLDER", None)
|
||||
if not isinstance(settings_folder, str):
|
||||
settings_folder = None
|
||||
|
||||
self.server = server or settings.IMAP_SERVER
|
||||
self.port = port or settings.IMAP_PORT
|
||||
self.username = username or settings.IMAP_USERNAME
|
||||
self.password = password or settings.IMAP_PASSWORD
|
||||
self.delete_emails = delete_emails
|
||||
configured_delete = getattr(settings, "DELETE_IMPORTED_EMAILS", False)
|
||||
if not isinstance(configured_delete, bool):
|
||||
configured_delete = False
|
||||
self.delete_emails = configured_delete if delete_emails is None else delete_emails
|
||||
self.folder = folder or settings_folder or "INBOX"
|
||||
self.db = db
|
||||
|
||||
self.report_store = ReportStore.get_instance()
|
||||
|
||||
if not all([self.server, self.username, self.password]):
|
||||
logger.warning("IMAP credentials not fully configured")
|
||||
|
||||
def _quoted_folder(self) -> str:
|
||||
escaped = self.folder.replace("\\", "\\\\").replace('"', '\\"')
|
||||
return f'"{escaped}"'
|
||||
|
||||
def _list_mailboxes(self, mailbox_data: list) -> list:
|
||||
"""Parse the raw IMAP LIST response into a list of mailbox name strings."""
|
||||
available_mailboxes = []
|
||||
@@ -63,7 +89,7 @@ class IMAPClient:
|
||||
if mailbox_name.startswith(" "):
|
||||
mailbox_name = mailbox_name[1:]
|
||||
available_mailboxes.append(mailbox_name)
|
||||
except Exception:
|
||||
except Exception: # pylint: disable=broad-exception-caught
|
||||
# Silently skip mailboxes that can't be parsed; they are simply
|
||||
# omitted from the returned list so callers should expect it may
|
||||
# be incomplete. Some IMAP servers return non-standard list
|
||||
@@ -84,7 +110,11 @@ class IMAPClient:
|
||||
- stats: Dictionary with mailbox statistics (if successful)
|
||||
"""
|
||||
if not all([self.server, self.username, self.password]):
|
||||
return False, "IMAP credentials not fully configured", {}
|
||||
return (
|
||||
False,
|
||||
"IMAP credentials not fully configured.",
|
||||
{"diagnostic_detail": "missing server, username, or password"},
|
||||
)
|
||||
|
||||
try:
|
||||
# Create IMAP4 connection
|
||||
@@ -96,18 +126,28 @@ class IMAPClient:
|
||||
status, mailbox_list = mail.list()
|
||||
available_mailboxes = self._list_mailboxes(mailbox_list) if status == "OK" else []
|
||||
|
||||
# Select inbox and get message count
|
||||
status, data = mail.select("INBOX")
|
||||
# Select configured mailbox and get message count
|
||||
status, data = mail.select(self._quoted_folder())
|
||||
message_count = 0
|
||||
unread_count = 0
|
||||
|
||||
if status == "OK":
|
||||
message_count = int(data[0])
|
||||
if status != "OK":
|
||||
mail.logout()
|
||||
return (
|
||||
False,
|
||||
"Configured mailbox folder could not be opened.",
|
||||
{
|
||||
"available_mailboxes": available_mailboxes,
|
||||
"diagnostic_detail": f"select failed for folder {self.folder}",
|
||||
},
|
||||
)
|
||||
|
||||
# Count unread messages
|
||||
status, data = mail.search(None, "UNSEEN")
|
||||
if status == "OK":
|
||||
unread_count = len(data[0].split())
|
||||
message_count = int(data[0])
|
||||
|
||||
# Count unread messages
|
||||
status, data = mail.search(None, "UNSEEN")
|
||||
if status == "OK":
|
||||
unread_count = len(data[0].split())
|
||||
|
||||
# Gather some stats about potential DMARC reports
|
||||
dmarc_count = 0
|
||||
@@ -130,35 +170,81 @@ class IMAPClient:
|
||||
}
|
||||
|
||||
return True, "Connection successful", stats
|
||||
except Exception as e:
|
||||
logger.error(f"IMAP connection test failed: {str(e)}")
|
||||
return False, f"Connection failed: {str(e)}", {}
|
||||
except imaplib.IMAP4.error as e:
|
||||
logger.error("IMAP connection test failed: %s", str(e))
|
||||
return (
|
||||
False,
|
||||
"IMAP authentication failed or the mailbox server rejected the request.",
|
||||
{"diagnostic_detail": str(e)},
|
||||
)
|
||||
except (TimeoutError, OSError) as e:
|
||||
logger.error("IMAP connection test failed: %s", str(e))
|
||||
return (
|
||||
False,
|
||||
"Could not reach the IMAP server.",
|
||||
{"diagnostic_detail": str(e)},
|
||||
)
|
||||
except Exception as e: # pylint: disable=broad-exception-caught
|
||||
logger.error("IMAP connection test failed: %s", str(e))
|
||||
return (
|
||||
False,
|
||||
"Connection failed. Check mailbox settings and try again.",
|
||||
{"diagnostic_detail": str(e)},
|
||||
)
|
||||
|
||||
def _process_single_email(self, mail, email_id: bytes, stats: dict) -> None:
|
||||
"""Fetch, parse, and store DMARC attachments from one email message."""
|
||||
message_id = email_id.decode("utf-8", errors="replace")
|
||||
try:
|
||||
status, msg_data = mail.fetch(email_id, "(RFC822)")
|
||||
if status != "OK":
|
||||
logger.error(f"Error fetching email ID {email_id}")
|
||||
logger.error("Error fetching email ID %s", email_id)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="error",
|
||||
reason="message_fetch_failed",
|
||||
message_id=message_id,
|
||||
)
|
||||
return
|
||||
|
||||
raw_email = msg_data[0][1]
|
||||
msg = email.message_from_bytes(raw_email)
|
||||
|
||||
if ForensicParser.is_forensic_report(msg):
|
||||
imported = self._process_forensic_email(
|
||||
raw_email,
|
||||
stats=stats,
|
||||
message_id=message_id,
|
||||
)
|
||||
mail.store(email_id, "+FLAGS", "\\Seen")
|
||||
if self.delete_emails and imported:
|
||||
mail.store(email_id, "+FLAGS", "\\Deleted")
|
||||
stats["deleted"] = stats.get("deleted", 0) + 1
|
||||
stats["processed"] += 1
|
||||
return
|
||||
|
||||
if self._is_dmarc_report_email(msg):
|
||||
reports_found = self._process_attachments(msg)
|
||||
reports_found = self._process_attachments(msg, stats, message_id=message_id)
|
||||
stats["reports_found"] += reports_found
|
||||
|
||||
# Mark email as read (and optionally delete)
|
||||
# Mark DMARC-looking email as read, and delete only after a successful import.
|
||||
mail.store(email_id, "+FLAGS", "\\Seen")
|
||||
if self.delete_emails:
|
||||
if self.delete_emails and reports_found > 0:
|
||||
mail.store(email_id, "+FLAGS", "\\Deleted")
|
||||
stats["deleted"] = stats.get("deleted", 0) + 1
|
||||
|
||||
stats["processed"] += 1
|
||||
except Exception as e:
|
||||
except Exception as e: # pylint: disable=broad-exception-caught
|
||||
error_msg = f"Error processing email ID {email_id}: {str(e)}"
|
||||
logger.error(error_msg)
|
||||
stats["errors"].append(error_msg)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="error",
|
||||
reason="message_processing_failed",
|
||||
message_id=message_id,
|
||||
error=str(e),
|
||||
)
|
||||
|
||||
def fetch_reports(self, days: int = 7) -> Dict[str, Any]:
|
||||
"""
|
||||
@@ -174,19 +260,13 @@ class IMAPClient:
|
||||
logger.error("IMAP credentials not fully configured")
|
||||
return {"success": False, "error": "IMAP credentials not configured", "processed": 0}
|
||||
|
||||
stats = {
|
||||
"success": True,
|
||||
"processed": 0,
|
||||
"reports_found": 0,
|
||||
"new_domains": [],
|
||||
"errors": [],
|
||||
}
|
||||
stats = initial_import_stats(deleted=True)
|
||||
|
||||
try:
|
||||
# Connect to the mail server
|
||||
mail = imaplib.IMAP4_SSL(self.server, self.port)
|
||||
mail.login(self.username, self.password)
|
||||
mail.select("INBOX")
|
||||
mail.select(self._quoted_folder())
|
||||
|
||||
# Calculate the date range for search
|
||||
date_since = (datetime.now() - timedelta(days=days)).strftime("%d-%b-%Y")
|
||||
@@ -213,7 +293,7 @@ class IMAPClient:
|
||||
self._process_single_email(mail, email_id, stats)
|
||||
|
||||
# Actually remove emails marked for deletion
|
||||
if self.delete_emails:
|
||||
if self.delete_emails and stats["deleted"] > 0:
|
||||
mail.expunge()
|
||||
|
||||
# Logout
|
||||
@@ -225,12 +305,13 @@ class IMAPClient:
|
||||
|
||||
return stats
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error fetching DMARC reports: {str(e)}")
|
||||
except Exception as e: # pylint: disable=broad-exception-caught
|
||||
logger.error("Error fetching DMARC reports: %s", str(e))
|
||||
return {
|
||||
"success": False,
|
||||
"error": f"Error connecting to mailbox: {str(e)}",
|
||||
"error": "Error connecting to mailbox. Check server logs for details.",
|
||||
"processed": 0,
|
||||
"errors": [sanitize_connector_error(e)],
|
||||
}
|
||||
|
||||
def _is_dmarc_report_email(self, msg: email.message.Message) -> bool:
|
||||
@@ -329,28 +410,212 @@ class IMAPClient:
|
||||
filename = self._decode_email_header(filename)
|
||||
|
||||
# Check file extension
|
||||
if (
|
||||
filename.lower().endswith(".xml")
|
||||
or filename.lower().endswith(".zip")
|
||||
or filename.lower().endswith(".gz")
|
||||
or filename.lower().endswith(".gzip")
|
||||
):
|
||||
if self._is_dmarc_filename(filename):
|
||||
return True
|
||||
|
||||
# Check content type
|
||||
content_type = part.get_content_type()
|
||||
if (
|
||||
content_type == "application/zip"
|
||||
or content_type == "application/gzip"
|
||||
or content_type == "application/x-gzip"
|
||||
or content_type == "application/xml"
|
||||
or content_type == "text/xml"
|
||||
if content_type in (
|
||||
"application/zip",
|
||||
"application/gzip",
|
||||
"application/x-gzip",
|
||||
"application/xml",
|
||||
"text/xml",
|
||||
):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def _process_attachments(self, msg: email.message.Message) -> int:
|
||||
@staticmethod
|
||||
def _is_dmarc_filename(filename: str) -> bool:
|
||||
lower = filename.lower()
|
||||
return (
|
||||
lower.endswith(".xml")
|
||||
or lower.endswith(".zip")
|
||||
or lower.endswith(".gz")
|
||||
or lower.endswith(".gzip")
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _append_detail(stats: Optional[Dict[str, Any]], **detail: str) -> None:
|
||||
"""Append a compact attachment/message outcome to the import stats."""
|
||||
if stats is None:
|
||||
return
|
||||
append_import_detail(stats, **detail)
|
||||
|
||||
def _store_report_if_new(
|
||||
self,
|
||||
report: Dict[str, Any],
|
||||
*,
|
||||
filename: str,
|
||||
stats: Optional[Dict[str, Any]],
|
||||
message_id: Optional[str],
|
||||
) -> bool:
|
||||
domain = report.get("domain", "unknown")
|
||||
report_id = report.get("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)
|
||||
if stats is not None:
|
||||
stats["duplicate_reports"] = stats.get("duplicate_reports", 0) + 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="duplicate",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
domain=str(domain),
|
||||
report_id=str(report_id),
|
||||
)
|
||||
return False
|
||||
|
||||
if self.db is not None:
|
||||
save_parsed_report(self.db, report)
|
||||
self.report_store.add_report(report)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="imported",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
domain=str(domain),
|
||||
report_id=str(report_id),
|
||||
)
|
||||
return True
|
||||
|
||||
def _process_forensic_email(
|
||||
self,
|
||||
raw_email: bytes,
|
||||
*,
|
||||
stats: Optional[Dict[str, Any]],
|
||||
message_id: Optional[str],
|
||||
) -> bool:
|
||||
try:
|
||||
report = ForensicParser.parse_bytes(
|
||||
raw_email,
|
||||
redaction_policy=get_forensic_redaction_policy(self.db),
|
||||
)
|
||||
report_id = str(report.get("report_id", ""))
|
||||
domain = str(report.get("reported_domain") or "unknown")
|
||||
|
||||
if self.db is not None and forensic_report_exists(self.db, report_id):
|
||||
if stats is not None:
|
||||
stats["duplicate_forensic_reports"] = (
|
||||
stats.get("duplicate_forensic_reports", 0) + 1
|
||||
)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="duplicate",
|
||||
reason="duplicate_forensic_report",
|
||||
message_id=message_id,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
return False
|
||||
|
||||
if self.db is None:
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="forensic_report_requires_database",
|
||||
message_id=message_id,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
return False
|
||||
|
||||
_row, created = save_forensic_report(self.db, report)
|
||||
if created:
|
||||
if stats is not None:
|
||||
stats["forensic_reports_found"] = stats.get("forensic_reports_found", 0) + 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="imported",
|
||||
reason="forensic_report",
|
||||
message_id=message_id,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
return True
|
||||
|
||||
if stats is not None:
|
||||
stats["duplicate_forensic_reports"] = stats.get("duplicate_forensic_reports", 0) + 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="duplicate",
|
||||
reason="duplicate_forensic_report",
|
||||
message_id=message_id,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
return False
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Error processing forensic report email %s: %s", message_id, exc)
|
||||
if stats is not None:
|
||||
stats.setdefault("errors", []).append(
|
||||
sanitize_connector_error(f"Failed to parse forensic report {message_id}: {exc}")
|
||||
)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="error",
|
||||
reason="forensic_parse_failed",
|
||||
message_id=message_id,
|
||||
error=str(exc),
|
||||
)
|
||||
return False
|
||||
|
||||
def _process_dmarc_attachment(
|
||||
self,
|
||||
part: email.message.Message,
|
||||
*,
|
||||
filename: str,
|
||||
stats: Optional[Dict[str, Any]],
|
||||
message_id: Optional[str],
|
||||
) -> bool:
|
||||
try:
|
||||
content = part.get_payload(decode=True)
|
||||
if not content:
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="empty_attachment",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
)
|
||||
return False
|
||||
|
||||
report = DMARCParser.parse_file(content, filename)
|
||||
stored = self._store_report_if_new(
|
||||
report,
|
||||
filename=filename,
|
||||
stats=stats,
|
||||
message_id=message_id,
|
||||
)
|
||||
if stored:
|
||||
logger.info("Successfully processed DMARC report: %s", filename)
|
||||
return stored
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Error processing attachment %s: %s", filename, str(exc))
|
||||
if stats is not None:
|
||||
stats.setdefault("errors", []).append(
|
||||
sanitize_connector_error(f"Failed to parse {filename}: {exc}")
|
||||
)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="error",
|
||||
reason="parse_failed",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
error=str(exc),
|
||||
)
|
||||
return False
|
||||
|
||||
def _process_attachments(
|
||||
self,
|
||||
msg: email.message.Message,
|
||||
stats: Optional[Dict[str, Any]] = None,
|
||||
message_id: Optional[str] = None,
|
||||
) -> int:
|
||||
"""
|
||||
Process email attachments that might be DMARC reports
|
||||
|
||||
@@ -363,35 +628,30 @@ class IMAPClient:
|
||||
reports_found = 0
|
||||
|
||||
for part in msg.walk():
|
||||
content_disposition = part.get_content_disposition()
|
||||
if part.get_content_disposition() != "attachment":
|
||||
continue
|
||||
|
||||
if content_disposition == "attachment":
|
||||
filename = part.get_filename()
|
||||
if filename:
|
||||
# Decode filename if needed
|
||||
filename = self._decode_email_header(filename)
|
||||
filename = part.get_filename()
|
||||
if not filename:
|
||||
continue
|
||||
|
||||
# Check if it's a likely DMARC report file
|
||||
if (
|
||||
filename.lower().endswith(".xml")
|
||||
or filename.lower().endswith(".zip")
|
||||
or filename.lower().endswith(".gz")
|
||||
or filename.lower().endswith(".gzip")
|
||||
):
|
||||
filename = self._decode_email_header(filename)
|
||||
if not self._is_dmarc_filename(filename):
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="unsupported_attachment",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
# Get attachment content
|
||||
content = part.get_payload(decode=True)
|
||||
|
||||
# Parse the DMARC report
|
||||
report = DMARCParser.parse_file(content, filename)
|
||||
|
||||
# Add the report to the store
|
||||
self.report_store.add_report(report)
|
||||
|
||||
reports_found += 1
|
||||
logger.info(f"Successfully processed DMARC report: {filename}")
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing attachment {filename}: {str(e)}")
|
||||
if self._process_dmarc_attachment(
|
||||
part,
|
||||
filename=filename,
|
||||
stats=stats,
|
||||
message_id=message_id,
|
||||
):
|
||||
reports_found += 1
|
||||
|
||||
return reports_found
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Iterable, Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.redaction import redact_sensitive_text
|
||||
from app.models.mail_source import MailSource
|
||||
from app.models.mail_source_import import MailSourceImport
|
||||
|
||||
MAX_STORED_ERRORS = 10
|
||||
MAX_ERROR_LENGTH = 500
|
||||
MAX_STORED_DETAILS = 50
|
||||
MAX_DETAIL_VALUE_LENGTH = 300
|
||||
DETAIL_FIELDS = {
|
||||
"status",
|
||||
"reason",
|
||||
"message_id",
|
||||
"filename",
|
||||
"domain",
|
||||
"report_id",
|
||||
"mailbox",
|
||||
"folder",
|
||||
"error",
|
||||
}
|
||||
|
||||
|
||||
def _sanitize_error(value: object) -> str:
|
||||
"""Return a compact, log-safe error string for storage and UI display."""
|
||||
text = redact_sensitive_text(value).strip()
|
||||
if len(text) > MAX_ERROR_LENGTH:
|
||||
return text[: MAX_ERROR_LENGTH - 3] + "..."
|
||||
return text
|
||||
|
||||
|
||||
def _json_list(values: Optional[Iterable[Any]]) -> str:
|
||||
return json.dumps([str(value) for value in values or []])
|
||||
|
||||
|
||||
def _sanitize_detail_value(value: object) -> str:
|
||||
text = _sanitize_error(value)
|
||||
if len(text) > MAX_DETAIL_VALUE_LENGTH:
|
||||
return text[: MAX_DETAIL_VALUE_LENGTH - 3] + "..."
|
||||
return text
|
||||
|
||||
|
||||
def _json_details(values: Optional[Iterable[Any]]) -> str:
|
||||
details = []
|
||||
for value in list(values or [])[:MAX_STORED_DETAILS]:
|
||||
if not isinstance(value, dict):
|
||||
continue
|
||||
entry = {}
|
||||
for key in DETAIL_FIELDS:
|
||||
if key in value and value[key] not in (None, ""):
|
||||
entry[key] = _sanitize_detail_value(value[key])
|
||||
if entry:
|
||||
details.append(entry)
|
||||
return json.dumps(details)
|
||||
|
||||
|
||||
def record_import_attempt(
|
||||
db: Session,
|
||||
source: MailSource,
|
||||
results: Dict[str, Any],
|
||||
*,
|
||||
started_at: datetime,
|
||||
trigger: str,
|
||||
) -> MailSourceImport:
|
||||
"""Persist a sanitized summary of a mail source import attempt."""
|
||||
result_errors = list(results.get("errors") or [])
|
||||
errors = [_sanitize_error(error) for error in result_errors[:MAX_STORED_ERRORS]]
|
||||
success = bool(results.get("success", False))
|
||||
status = "success" if success and not errors else "warning" if success else "failed"
|
||||
|
||||
attempt = MailSourceImport(
|
||||
mail_source_id=source.id,
|
||||
trigger=trigger,
|
||||
status=status,
|
||||
processed=int(results.get("processed", 0) or 0),
|
||||
reports_found=int(results.get("reports_found", 0) or 0),
|
||||
duplicate_reports=int(results.get("duplicate_reports", 0) or 0),
|
||||
error_count=len(result_errors),
|
||||
new_domains=_json_list(results.get("new_domains", [])),
|
||||
errors=json.dumps(errors),
|
||||
details=_json_details(results.get("details", [])),
|
||||
started_at=started_at,
|
||||
finished_at=datetime.utcnow(),
|
||||
)
|
||||
db.add(attempt)
|
||||
return attempt
|
||||
@@ -0,0 +1,164 @@
|
||||
"""Shared contracts and helpers for mail-source connectors."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Iterable, List, Optional, Protocol
|
||||
|
||||
from app.core.redaction import redact_sensitive_text
|
||||
|
||||
MAX_CONNECTOR_ERROR_LENGTH = 500
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ConnectorImportContext:
|
||||
"""Safe import context that can be returned to operators and import history."""
|
||||
|
||||
source_type: str
|
||||
mailbox: Optional[str] = None
|
||||
folder: Optional[str] = None
|
||||
search_window_days: Optional[int] = None
|
||||
|
||||
def as_stats(self) -> Dict[str, Any]:
|
||||
stats: Dict[str, Any] = {"source_type": self.source_type}
|
||||
if self.mailbox:
|
||||
stats["target_mailbox"] = self.mailbox
|
||||
if self.folder:
|
||||
stats["target_folder"] = self.folder
|
||||
if self.search_window_days is not None:
|
||||
stats["search_window_days"] = self.search_window_days
|
||||
return stats
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ConnectorMessage:
|
||||
"""Provider-neutral message metadata used by connector implementations."""
|
||||
|
||||
message_id: str
|
||||
subject: str = ""
|
||||
sender: str = ""
|
||||
received_at: Optional[str] = None
|
||||
has_attachments: bool = False
|
||||
raw: Any = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ConnectorAttachment:
|
||||
"""Provider-neutral attachment payload used by connector implementations."""
|
||||
|
||||
filename: str
|
||||
content: bytes
|
||||
content_type: str = ""
|
||||
raw: Any = None
|
||||
|
||||
|
||||
class MailSourceConnector(Protocol):
|
||||
"""Interface new mailbox connectors should satisfy before endpoint wiring."""
|
||||
|
||||
def import_context(self, days: Optional[int] = None) -> ConnectorImportContext:
|
||||
"""Return safe, non-secret context for history and API responses."""
|
||||
|
||||
def search_messages(self, days: int) -> Iterable[Any]:
|
||||
"""Return provider messages in the requested search window."""
|
||||
|
||||
def iter_attachments(self, message: Any) -> Iterable[Any]:
|
||||
"""Yield provider attachments for one message."""
|
||||
|
||||
def fetch_reports(self, days: int = 7) -> Dict[str, Any]:
|
||||
"""Fetch, parse, and persist DMARC reports."""
|
||||
|
||||
|
||||
def clamp_search_window(days: Optional[int], *, default: int = 7, maximum: int = 365) -> int:
|
||||
"""Normalize user-supplied backfill windows for connector fetches."""
|
||||
try:
|
||||
value = int(days or default)
|
||||
except (TypeError, ValueError):
|
||||
value = default
|
||||
return max(1, min(value, maximum))
|
||||
|
||||
|
||||
def initial_import_stats(
|
||||
context: Optional[ConnectorImportContext] = None,
|
||||
*,
|
||||
deleted: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
"""Return the shared import-result shape used by mailbox connectors."""
|
||||
stats: Dict[str, Any] = {
|
||||
"success": True,
|
||||
"processed": 0,
|
||||
"reports_found": 0,
|
||||
"forensic_reports_found": 0,
|
||||
"duplicate_reports": 0,
|
||||
"duplicate_forensic_reports": 0,
|
||||
"new_domains": [],
|
||||
"errors": [],
|
||||
"new_ingested_ids": [],
|
||||
"details": [],
|
||||
}
|
||||
if deleted:
|
||||
stats["deleted"] = 0
|
||||
if context:
|
||||
stats.update(context.as_stats())
|
||||
return stats
|
||||
|
||||
|
||||
def append_import_detail(
|
||||
stats: Optional[Dict[str, Any]],
|
||||
*,
|
||||
context: Optional[ConnectorImportContext] = None,
|
||||
**detail: Any,
|
||||
) -> None:
|
||||
"""Append one compact, sanitized message or attachment outcome."""
|
||||
if stats is None:
|
||||
return
|
||||
if context:
|
||||
detail.setdefault("mailbox", context.mailbox)
|
||||
detail.setdefault("folder", context.folder)
|
||||
clean_detail = {
|
||||
str(key): sanitize_connector_error(value)
|
||||
for key, value in detail.items()
|
||||
if value not in (None, "")
|
||||
}
|
||||
if clean_detail:
|
||||
stats.setdefault("details", []).append(clean_detail)
|
||||
|
||||
|
||||
def sanitize_connector_error(value: object) -> str:
|
||||
"""Return a compact, log-safe connector diagnostic with secrets redacted."""
|
||||
text = redact_sensitive_text(value).strip()
|
||||
if len(text) > MAX_CONNECTOR_ERROR_LENGTH:
|
||||
return text[: MAX_CONNECTOR_ERROR_LENGTH - 3] + "..."
|
||||
return text
|
||||
|
||||
|
||||
def connector_failure_stats(
|
||||
stats: Dict[str, Any],
|
||||
message: str,
|
||||
*,
|
||||
error: Optional[object] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Return a standardized failed import payload for provider errors."""
|
||||
safe_message = sanitize_connector_error(error if error is not None else message)
|
||||
return {
|
||||
**stats,
|
||||
"success": False,
|
||||
"error": safe_message,
|
||||
"errors": [safe_message],
|
||||
}
|
||||
|
||||
|
||||
def load_ingested_ids(json_text: Optional[str]) -> List[str]:
|
||||
"""Deserialize a connector ingested-message-id JSON column."""
|
||||
if not json_text:
|
||||
return []
|
||||
try:
|
||||
decoded = json.loads(json_text)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return []
|
||||
return [str(item) for item in decoded] if isinstance(decoded, list) else []
|
||||
|
||||
|
||||
def dump_ingested_ids(ids: Iterable[Any]) -> str:
|
||||
"""Serialize connector ingested-message IDs for database storage."""
|
||||
return json.dumps([str(item) for item in ids])
|
||||
@@ -0,0 +1,586 @@
|
||||
"""Microsoft Graph client for retrieving DMARC aggregate reports."""
|
||||
|
||||
import base64
|
||||
import logging
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Callable, Dict, Iterable, List, Optional
|
||||
from urllib.parse import quote, urlencode
|
||||
|
||||
import httpx
|
||||
|
||||
from app.services.dmarc_parser import DMARCParser
|
||||
from app.services.mail_connector import (
|
||||
ConnectorImportContext,
|
||||
MailSourceConnector,
|
||||
append_import_detail,
|
||||
clamp_search_window,
|
||||
connector_failure_stats,
|
||||
dump_ingested_ids,
|
||||
initial_import_stats,
|
||||
load_ingested_ids,
|
||||
sanitize_connector_error,
|
||||
)
|
||||
from app.services.report_persistence import report_exists, save_parsed_report
|
||||
from app.services.report_store import ReportStore
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
GRAPH_BASE_URL = "https://graph.microsoft.com/v1.0"
|
||||
LOGIN_BASE_URL = "https://login.microsoftonline.com"
|
||||
|
||||
M365_SCOPES = [
|
||||
"offline_access",
|
||||
"https://graph.microsoft.com/User.Read",
|
||||
"https://graph.microsoft.com/Mail.Read",
|
||||
]
|
||||
|
||||
_PAGE_SIZE = 100
|
||||
_MAX_FOLDER_DEPTH = 5
|
||||
_MAX_GRAPH_RETRIES = 3
|
||||
_MAX_RETRY_DELAY_SECONDS = 30
|
||||
_RETRYABLE_STATUS_CODES = {429, 503, 504}
|
||||
_DMARC_SUBJECT_TERMS = (
|
||||
"dmarc",
|
||||
"aggregate report",
|
||||
"domain report",
|
||||
"report domain",
|
||||
"rua",
|
||||
"submitter",
|
||||
)
|
||||
_DMARC_SENDER_TERMS = (
|
||||
"dmarc",
|
||||
"reports",
|
||||
"postmaster",
|
||||
)
|
||||
|
||||
|
||||
class MicrosoftGraphError(RuntimeError):
|
||||
"""Raised when Microsoft Graph or the token endpoint returns a failure."""
|
||||
|
||||
|
||||
class MicrosoftGraphClient(MailSourceConnector):
|
||||
"""
|
||||
Retrieve DMARC aggregate reports from Microsoft 365 through Microsoft Graph.
|
||||
|
||||
The client uses delegated OAuth tokens and read-only Graph scopes. Messages
|
||||
are never modified or deleted; already-ingested Graph message IDs are stored
|
||||
by the caller to avoid reprocessing the same email.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tenant_id: str,
|
||||
client_id: str,
|
||||
client_secret: str,
|
||||
access_token: str,
|
||||
refresh_token: str,
|
||||
mailbox: Optional[str] = None,
|
||||
folder: str = "inbox",
|
||||
folder_id: Optional[str] = None,
|
||||
already_ingested_ids: Optional[List[str]] = None,
|
||||
db: Any = None,
|
||||
sleep: Optional[Callable[[float], None]] = None,
|
||||
):
|
||||
self.tenant_id = tenant_id or "common"
|
||||
self.client_id = client_id
|
||||
self.client_secret = client_secret
|
||||
self.access_token = access_token
|
||||
self.refresh_token = refresh_token
|
||||
self.mailbox = (mailbox or "").strip()
|
||||
self.folder = folder or "inbox"
|
||||
self.folder_id = (folder_id or "").strip()
|
||||
self.already_ingested_ids: List[str] = list(already_ingested_ids or [])
|
||||
self.report_store = ReportStore.get_instance()
|
||||
self.db = db
|
||||
self._sleep = sleep or time.sleep
|
||||
self._refreshed_tokens: Optional[Dict[str, str]] = None
|
||||
|
||||
def get_refreshed_tokens(self) -> Optional[Dict[str, str]]:
|
||||
"""Return refreshed OAuth tokens, if a request had to refresh them."""
|
||||
return self._refreshed_tokens
|
||||
|
||||
@staticmethod
|
||||
def build_authorization_url(
|
||||
tenant_id: str,
|
||||
client_id: str,
|
||||
redirect_uri: str,
|
||||
state: Optional[str] = None,
|
||||
) -> str:
|
||||
"""Build a Microsoft identity platform authorization-code URL."""
|
||||
tenant = quote(tenant_id or "common", safe="")
|
||||
params: Dict[str, str] = {
|
||||
"client_id": client_id,
|
||||
"response_type": "code",
|
||||
"redirect_uri": redirect_uri,
|
||||
"response_mode": "query",
|
||||
"scope": " ".join(M365_SCOPES),
|
||||
"prompt": "select_account",
|
||||
}
|
||||
if state:
|
||||
params["state"] = state
|
||||
return f"{LOGIN_BASE_URL}/{tenant}/oauth2/v2.0/authorize?" + urlencode(params)
|
||||
|
||||
@staticmethod
|
||||
def exchange_code_for_tokens(
|
||||
tenant_id: str,
|
||||
client_id: str,
|
||||
client_secret: str,
|
||||
code: str,
|
||||
redirect_uri: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""Exchange an authorization code for Microsoft Graph tokens."""
|
||||
data = {
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
"code": code,
|
||||
"redirect_uri": redirect_uri,
|
||||
"grant_type": "authorization_code",
|
||||
"scope": " ".join(M365_SCOPES),
|
||||
}
|
||||
resp = httpx.post(MicrosoftGraphClient._token_url(tenant_id), data=data, timeout=30)
|
||||
if resp.status_code != 200:
|
||||
raise MicrosoftGraphError(
|
||||
f"Microsoft token exchange failed ({resp.status_code}): {resp.text}"
|
||||
)
|
||||
return resp.json()
|
||||
|
||||
@staticmethod
|
||||
def get_account_email(access_token: str) -> Optional[str]:
|
||||
"""Return the mailbox identity exposed by Graph /me for an access token."""
|
||||
try:
|
||||
resp = httpx.get(
|
||||
f"{GRAPH_BASE_URL}/me",
|
||||
headers={"Authorization": f"Bearer {access_token}"},
|
||||
params={"$select": "mail,userPrincipalName"},
|
||||
timeout=30,
|
||||
)
|
||||
if resp.status_code == 200:
|
||||
profile = resp.json()
|
||||
return profile.get("mail") or profile.get("userPrincipalName")
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Failed to fetch Microsoft 365 account email: %s", exc)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def load_ingested_ids(json_text: Optional[str]) -> List[str]:
|
||||
"""Deserialize the m365_ingested_ids text column into a list."""
|
||||
return load_ingested_ids(json_text)
|
||||
|
||||
@staticmethod
|
||||
def dump_ingested_ids(ids: List[str]) -> str:
|
||||
"""Serialize Graph message IDs for database storage."""
|
||||
return dump_ingested_ids(ids)
|
||||
|
||||
def test_connection(self) -> Dict[str, Any]:
|
||||
"""Verify that the saved delegated token can read the target mailbox."""
|
||||
data = self._request(
|
||||
"GET",
|
||||
self._messages_path(),
|
||||
params={"$top": 1, "$select": "id"},
|
||||
)
|
||||
return {
|
||||
"success": True,
|
||||
"message_count": len(data.get("value", [])),
|
||||
"target_mailbox": self._target_mailbox_label(),
|
||||
"target_folder": self._target_folder_label(),
|
||||
"diagnostic_detail": "Microsoft Graph mailbox read succeeded.",
|
||||
}
|
||||
|
||||
def list_mail_folders(self) -> List[Dict[str, str]]:
|
||||
"""Return selectable mail folders for the configured mailbox."""
|
||||
folders: List[Dict[str, str]] = []
|
||||
self._collect_mail_folders(
|
||||
f"{self._mailbox_path()}/mailFolders",
|
||||
folders,
|
||||
parent_path="",
|
||||
depth=0,
|
||||
)
|
||||
return folders
|
||||
|
||||
def _collect_mail_folders(
|
||||
self,
|
||||
start_url: str,
|
||||
folders: List[Dict[str, str]],
|
||||
*,
|
||||
parent_path: str,
|
||||
depth: int,
|
||||
) -> None:
|
||||
params: Optional[Dict[str, Any]] = {
|
||||
"$top": _PAGE_SIZE,
|
||||
"$select": "id,displayName,parentFolderId,childFolderCount",
|
||||
}
|
||||
url: Optional[str] = start_url
|
||||
|
||||
while url:
|
||||
data = self._request("GET", url, params=params)
|
||||
for folder in data.get("value", []):
|
||||
folder_id = str(folder.get("id") or "")
|
||||
display_name = str(folder.get("displayName") or folder_id)
|
||||
if not folder_id:
|
||||
continue
|
||||
folder_path = f"{parent_path} / {display_name}" if parent_path else display_name
|
||||
folders.append(
|
||||
{
|
||||
"id": folder_id,
|
||||
"display_name": display_name,
|
||||
"path": folder_path,
|
||||
"parent_folder_id": str(folder.get("parentFolderId") or ""),
|
||||
}
|
||||
)
|
||||
if int(folder.get("childFolderCount") or 0) > 0 and depth < _MAX_FOLDER_DEPTH:
|
||||
self._collect_mail_folders(
|
||||
f"{self._mailbox_path()}/mailFolders/{quote(folder_id, safe='')}/childFolders",
|
||||
folders,
|
||||
parent_path=folder_path,
|
||||
depth=depth + 1,
|
||||
)
|
||||
url = data.get("@odata.nextLink")
|
||||
params = None
|
||||
|
||||
def import_context(self, days: Optional[int] = None) -> ConnectorImportContext:
|
||||
"""Return safe Microsoft 365 import context for API responses/history."""
|
||||
return ConnectorImportContext(
|
||||
source_type="M365_GRAPH",
|
||||
mailbox=self._target_mailbox_label(),
|
||||
folder=self._target_folder_label(),
|
||||
search_window_days=days,
|
||||
)
|
||||
|
||||
def search_messages(self, days: int) -> Iterable[Dict[str, Any]]:
|
||||
"""Return Microsoft Graph messages that look like DMARC reports."""
|
||||
return self._list_dmarc_messages(days=days)
|
||||
|
||||
def iter_attachments(self, message: Dict[str, Any]) -> Iterable[Dict[str, Any]]:
|
||||
"""Yield Microsoft Graph attachments for one message."""
|
||||
message_id = str(message.get("id") or "")
|
||||
return self._list_attachments(message_id)
|
||||
|
||||
def fetch_reports(self, days: int = 7) -> Dict[str, Any]:
|
||||
"""Fetch and ingest DMARC report attachments from Microsoft Graph."""
|
||||
safe_days = clamp_search_window(days)
|
||||
stats = initial_import_stats(self.import_context(days=safe_days))
|
||||
|
||||
try:
|
||||
messages = self.search_messages(days=safe_days)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Microsoft Graph: failed to list messages: %s", exc)
|
||||
return connector_failure_stats(
|
||||
stats, "Failed to list Microsoft Graph messages.", error=exc
|
||||
)
|
||||
|
||||
domains_before = set(self.report_store.get_domains())
|
||||
|
||||
for message in messages:
|
||||
message_id = str(message.get("id") or "")
|
||||
if not message_id:
|
||||
continue
|
||||
if message_id in self.already_ingested_ids:
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="already_ingested_message",
|
||||
message_id=message_id,
|
||||
)
|
||||
continue
|
||||
|
||||
stats["processed"] += 1
|
||||
found = self._process_message(message, stats)
|
||||
if found >= 0:
|
||||
stats["new_ingested_ids"].append(message_id)
|
||||
self.already_ingested_ids.append(message_id)
|
||||
|
||||
domains_after = set(self.report_store.get_domains())
|
||||
stats["new_domains"] = list(domains_after - domains_before)
|
||||
return stats
|
||||
|
||||
@staticmethod
|
||||
def _token_url(tenant_id: str) -> str:
|
||||
tenant = quote(tenant_id or "common", safe="")
|
||||
return f"{LOGIN_BASE_URL}/{tenant}/oauth2/v2.0/token"
|
||||
|
||||
def _append_detail(self, stats: dict, **detail: str) -> None:
|
||||
append_import_detail(stats, context=self.import_context(), **detail)
|
||||
|
||||
def _target_mailbox_label(self) -> str:
|
||||
return self.mailbox or "authorized account"
|
||||
|
||||
def _target_folder_label(self) -> str:
|
||||
if self.folder:
|
||||
return self.folder
|
||||
if self.folder_id:
|
||||
return self.folder_id
|
||||
return "All messages"
|
||||
|
||||
def _mailbox_path(self) -> str:
|
||||
if not self.mailbox or self.mailbox.lower() == "me":
|
||||
return "/me"
|
||||
return f"/users/{quote(self.mailbox, safe='')}"
|
||||
|
||||
def _messages_path(self) -> str:
|
||||
mailbox_path = self._mailbox_path()
|
||||
if self.folder_id:
|
||||
return f"{mailbox_path}/mailFolders/{quote(self.folder_id, safe='')}/messages"
|
||||
folder = (self.folder or "").strip()
|
||||
if not folder:
|
||||
return f"{mailbox_path}/messages"
|
||||
if folder.upper() == "INBOX":
|
||||
folder = "inbox"
|
||||
return f"{mailbox_path}/mailFolders/{quote(folder, safe='')}/messages"
|
||||
|
||||
def _headers(self) -> Dict[str, str]:
|
||||
return {"Authorization": f"Bearer {self.access_token}"}
|
||||
|
||||
def _request(
|
||||
self,
|
||||
method: str,
|
||||
path_or_url: str,
|
||||
*,
|
||||
params: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
url = path_or_url if path_or_url.startswith("http") else f"{GRAPH_BASE_URL}{path_or_url}"
|
||||
resp: Optional[httpx.Response] = None
|
||||
|
||||
for attempt in range(_MAX_GRAPH_RETRIES + 1):
|
||||
resp = httpx.request(method, url, headers=self._headers(), params=params, timeout=30)
|
||||
if resp.status_code == 401 and self.refresh_token:
|
||||
self._refresh_access_token()
|
||||
resp = httpx.request(
|
||||
method, url, headers=self._headers(), params=params, timeout=30
|
||||
)
|
||||
if resp.status_code in _RETRYABLE_STATUS_CODES and attempt < _MAX_GRAPH_RETRIES:
|
||||
delay = self._retry_delay_seconds(resp, attempt)
|
||||
logger.warning(
|
||||
"Microsoft Graph request throttled/unavailable; retrying in %.1fs",
|
||||
delay,
|
||||
)
|
||||
self._sleep(delay)
|
||||
continue
|
||||
break
|
||||
|
||||
if resp is None:
|
||||
raise MicrosoftGraphError("Microsoft Graph request failed before receiving a response.")
|
||||
if resp.status_code < 200 or resp.status_code >= 300:
|
||||
raise MicrosoftGraphError(self._format_error(resp))
|
||||
return resp.json() if resp.content else {}
|
||||
|
||||
@staticmethod
|
||||
def _retry_delay_seconds(resp: httpx.Response, attempt: int) -> float:
|
||||
retry_after = resp.headers.get("Retry-After")
|
||||
if retry_after:
|
||||
try:
|
||||
return min(float(retry_after), _MAX_RETRY_DELAY_SECONDS)
|
||||
except ValueError:
|
||||
pass
|
||||
return min(float(2**attempt), _MAX_RETRY_DELAY_SECONDS)
|
||||
|
||||
def _refresh_access_token(self) -> None:
|
||||
data = {
|
||||
"client_id": self.client_id,
|
||||
"client_secret": self.client_secret,
|
||||
"refresh_token": self.refresh_token,
|
||||
"grant_type": "refresh_token",
|
||||
"scope": " ".join(M365_SCOPES),
|
||||
}
|
||||
resp = httpx.post(self._token_url(self.tenant_id), data=data, timeout=30)
|
||||
if resp.status_code != 200:
|
||||
raise MicrosoftGraphError(
|
||||
f"Microsoft token refresh failed ({resp.status_code}): {resp.text}"
|
||||
)
|
||||
token_data = resp.json()
|
||||
access_token = token_data.get("access_token")
|
||||
if not access_token:
|
||||
raise MicrosoftGraphError("Microsoft token refresh did not return an access token.")
|
||||
self.access_token = access_token
|
||||
refreshed = {"access_token": access_token}
|
||||
if token_data.get("refresh_token"):
|
||||
self.refresh_token = token_data["refresh_token"]
|
||||
refreshed["refresh_token"] = token_data["refresh_token"]
|
||||
self._refreshed_tokens = refreshed
|
||||
|
||||
@staticmethod
|
||||
def _format_error(resp: httpx.Response) -> str:
|
||||
try:
|
||||
payload = resp.json()
|
||||
except ValueError:
|
||||
payload = {}
|
||||
message = payload.get("error_description")
|
||||
if not message and isinstance(payload.get("error"), dict):
|
||||
message = payload["error"].get("message")
|
||||
code = payload["error"].get("code")
|
||||
if code:
|
||||
message = f"{code}: {message}" if message else code
|
||||
return message or f"Microsoft Graph request failed ({resp.status_code}): {resp.text}"
|
||||
|
||||
@staticmethod
|
||||
def _looks_like_dmarc_message(message: Dict[str, Any]) -> bool:
|
||||
if not message.get("hasAttachments"):
|
||||
return False
|
||||
subject = str(message.get("subject") or "").lower()
|
||||
sender = (
|
||||
((message.get("from") or {}).get("emailAddress") or {}).get("address") or ""
|
||||
).lower()
|
||||
return any(term in subject for term in _DMARC_SUBJECT_TERMS) or any(
|
||||
term in sender for term in _DMARC_SENDER_TERMS
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_dmarc_attachment(filename: str) -> bool:
|
||||
lower = filename.lower()
|
||||
return (
|
||||
lower.endswith(".xml")
|
||||
or lower.endswith(".zip")
|
||||
or lower.endswith(".gz")
|
||||
or lower.endswith(".gzip")
|
||||
)
|
||||
|
||||
def _list_dmarc_messages(self, days: int) -> List[Dict[str, Any]]:
|
||||
messages: List[Dict[str, Any]] = []
|
||||
url = self._messages_path()
|
||||
cutoff = datetime.utcnow() - timedelta(days=days)
|
||||
params: Optional[Dict[str, Any]] = {
|
||||
"$top": _PAGE_SIZE,
|
||||
"$select": "id,subject,from,hasAttachments,receivedDateTime",
|
||||
"$orderby": "receivedDateTime desc",
|
||||
"$filter": f"receivedDateTime ge {cutoff.strftime('%Y-%m-%dT%H:%M:%SZ')}",
|
||||
}
|
||||
|
||||
while url:
|
||||
data = self._request("GET", url, params=params)
|
||||
for message in data.get("value", []):
|
||||
if self._looks_like_dmarc_message(message):
|
||||
messages.append(message)
|
||||
url = data.get("@odata.nextLink")
|
||||
params = None
|
||||
|
||||
return messages
|
||||
|
||||
def _process_message(self, message: Dict[str, Any], stats: Dict[str, Any]) -> int:
|
||||
message_id = str(message.get("id") or "")
|
||||
try:
|
||||
attachments = self._list_attachments(message_id)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Microsoft Graph: failed to fetch attachments for %s: %s", message_id, exc)
|
||||
stats["errors"].append(
|
||||
sanitize_connector_error(
|
||||
f"Failed to fetch attachments for message {message_id}: {exc}"
|
||||
)
|
||||
)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="error",
|
||||
reason="attachment_fetch_failed",
|
||||
message_id=message_id,
|
||||
error=str(exc),
|
||||
)
|
||||
return -1
|
||||
return self._process_attachments(message_id, attachments, stats)
|
||||
|
||||
def _list_attachments(self, message_id: str) -> List[Dict[str, Any]]:
|
||||
mailbox_path = self._mailbox_path()
|
||||
data = self._request(
|
||||
"GET",
|
||||
f"{mailbox_path}/messages/{quote(message_id, safe='')}/attachments",
|
||||
)
|
||||
return list(data.get("value", []))
|
||||
|
||||
def _store_report_if_new(self, report: Dict[str, Any]) -> bool:
|
||||
domain = report.get("domain", "unknown")
|
||||
report_id = report.get("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
|
||||
|
||||
def _process_attachments(
|
||||
self,
|
||||
message_id: str,
|
||||
attachments: List[Dict[str, Any]],
|
||||
stats: Dict[str, Any],
|
||||
) -> int:
|
||||
reports_found = 0
|
||||
|
||||
for attachment in attachments:
|
||||
filename = str(attachment.get("name") or "")
|
||||
if not filename:
|
||||
continue
|
||||
if not self._is_dmarc_attachment(filename):
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="unsupported_attachment",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
)
|
||||
continue
|
||||
|
||||
attachment_type = str(attachment.get("@odata.type") or "").lower()
|
||||
if "fileattachment" not in attachment_type:
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="unsupported_attachment_type",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
)
|
||||
continue
|
||||
|
||||
content_b64 = attachment.get("contentBytes")
|
||||
if not content_b64:
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="skipped",
|
||||
reason="empty_attachment",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
content = base64.b64decode(content_b64)
|
||||
report = DMARCParser.parse_file(content, filename)
|
||||
domain = str(report.get("domain", "unknown"))
|
||||
report_id = str(report.get("report_id", ""))
|
||||
if self._store_report_if_new(report):
|
||||
stats["reports_found"] += 1
|
||||
reports_found += 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="imported",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
else:
|
||||
stats["duplicate_reports"] = stats.get("duplicate_reports", 0) + 1
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="duplicate",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
domain=domain,
|
||||
report_id=report_id,
|
||||
)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
logger.error("Failed to parse Graph DMARC attachment %s: %s", filename, exc)
|
||||
stats["errors"].append(
|
||||
sanitize_connector_error(f"Failed to parse {filename}: {exc}")
|
||||
)
|
||||
self._append_detail(
|
||||
stats,
|
||||
status="error",
|
||||
reason="parse_failed",
|
||||
message_id=message_id,
|
||||
filename=filename,
|
||||
error=str(exc),
|
||||
)
|
||||
|
||||
return reports_found
|
||||
@@ -0,0 +1,216 @@
|
||||
"""MTA-STS posture checks for monitored domains."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.dns_cache import DNSCache
|
||||
from app.services.dns_cache import DEFAULT_DNS_CACHE_TTL_SECONDS
|
||||
from app.services.dns_resolver import BaseDNSProvider
|
||||
|
||||
_CACHE_KEY = "mta-sts-v1"
|
||||
_POLICY_TIMEOUT_SECONDS = 5.0
|
||||
_VALID_MODES = {"enforce", "testing", "none"}
|
||||
|
||||
|
||||
@dataclass
|
||||
class MTAStsResult:
|
||||
"""Operator-facing MTA-STS posture evidence."""
|
||||
|
||||
status: str = "fail"
|
||||
dns_record: Optional[str] = None
|
||||
policy_url: Optional[str] = None
|
||||
policy_text: Optional[str] = None
|
||||
mode: Optional[str] = None
|
||||
max_age: Optional[int] = None
|
||||
mx: List[str] = field(default_factory=list)
|
||||
errors: List[str] = field(default_factory=list)
|
||||
warnings: List[str] = field(default_factory=list)
|
||||
|
||||
|
||||
def _utcnow_naive() -> datetime:
|
||||
return datetime.now(timezone.utc).replace(tzinfo=None)
|
||||
|
||||
|
||||
def _is_fresh(row: DNSCache, ttl_seconds: int, now: datetime) -> bool:
|
||||
return row.checked_at >= now - timedelta(seconds=ttl_seconds)
|
||||
|
||||
|
||||
def _result_from_json(value: str) -> MTAStsResult:
|
||||
data = json.loads(value)
|
||||
return MTAStsResult(
|
||||
status=str(data.get("status") or "fail"),
|
||||
dns_record=data.get("dns_record"),
|
||||
policy_url=data.get("policy_url"),
|
||||
policy_text=data.get("policy_text"),
|
||||
mode=data.get("mode"),
|
||||
max_age=data.get("max_age"),
|
||||
mx=list(data.get("mx") or []),
|
||||
errors=list(data.get("errors") or []),
|
||||
warnings=list(data.get("warnings") or []),
|
||||
)
|
||||
|
||||
|
||||
def parse_mta_sts_record(records: List[str]) -> Tuple[Optional[str], List[str], List[str]]:
|
||||
"""Return the selected MTA-STS TXT record, warnings, and errors."""
|
||||
sts_records = [record for record in records if record.lower().startswith("v=stsv1")]
|
||||
if not sts_records:
|
||||
return None, [], ["No _mta-sts TXT record was found."]
|
||||
warnings = []
|
||||
if len(sts_records) > 1:
|
||||
warnings.append("Multiple _mta-sts TXT records were found; publish exactly one.")
|
||||
record = sts_records[0]
|
||||
tags = {
|
||||
part.split("=", 1)[0].strip().lower(): part.split("=", 1)[1].strip()
|
||||
for part in record.split(";")
|
||||
if "=" in part
|
||||
}
|
||||
errors = []
|
||||
if tags.get("v", "").lower() != "stsv1":
|
||||
errors.append("The _mta-sts TXT record must start with v=STSv1.")
|
||||
if not tags.get("id"):
|
||||
errors.append("The _mta-sts TXT record must include a non-empty id tag.")
|
||||
return record, warnings, errors
|
||||
|
||||
|
||||
def parse_mta_sts_policy( # noqa: C901
|
||||
policy_text: str,
|
||||
) -> Tuple[Dict[str, Any], List[str], List[str]]:
|
||||
"""Parse and validate an MTA-STS policy file."""
|
||||
data: Dict[str, Any] = {"mx": []}
|
||||
for raw_line in policy_text.splitlines():
|
||||
line = raw_line.strip()
|
||||
if not line or line.startswith("#") or ":" not in line:
|
||||
continue
|
||||
key, value = line.split(":", 1)
|
||||
key = key.strip().lower()
|
||||
value = value.strip()
|
||||
if key == "mx":
|
||||
data.setdefault("mx", []).append(value)
|
||||
else:
|
||||
data[key] = value
|
||||
|
||||
errors = []
|
||||
warnings = []
|
||||
if str(data.get("version", "")).upper() != "STSV1":
|
||||
errors.append("The policy file must contain version: STSv1.")
|
||||
mode = str(data.get("mode", "")).lower()
|
||||
if mode not in _VALID_MODES:
|
||||
errors.append("The policy file must contain mode: enforce, testing, or none.")
|
||||
elif mode in {"testing", "none"}:
|
||||
warnings.append(f"MTA-STS policy is valid but not enforcing mail delivery ({mode}).")
|
||||
try:
|
||||
max_age = int(str(data.get("max_age", "")))
|
||||
if max_age <= 0:
|
||||
errors.append("The policy max_age must be greater than zero.")
|
||||
data["max_age"] = max_age
|
||||
except ValueError:
|
||||
errors.append("The policy file must contain an integer max_age value.")
|
||||
if not data.get("mx"):
|
||||
errors.append("The policy file must contain at least one mx entry.")
|
||||
return data, warnings, errors
|
||||
|
||||
|
||||
async def check_mta_sts(domain: str, provider: BaseDNSProvider) -> MTAStsResult:
|
||||
"""Resolve the MTA-STS TXT record and validate the HTTPS policy file."""
|
||||
result = MTAStsResult(policy_url=f"https://mta-sts.{domain}/.well-known/mta-sts.txt")
|
||||
try:
|
||||
records = await provider.lookup_txt(f"_mta-sts.{domain}")
|
||||
except LookupError as exc:
|
||||
result.errors.append(f"MTA-STS DNS lookup failed: {exc}")
|
||||
return result
|
||||
|
||||
record, warnings, errors = parse_mta_sts_record(records)
|
||||
result.dns_record = record
|
||||
result.warnings.extend(warnings)
|
||||
result.errors.extend(errors)
|
||||
if record is None:
|
||||
return result
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=_POLICY_TIMEOUT_SECONDS, follow_redirects=False
|
||||
) as client:
|
||||
response = await client.get(result.policy_url)
|
||||
response.raise_for_status()
|
||||
result.policy_text = response.text
|
||||
except (httpx.RequestError, httpx.HTTPStatusError, httpx.TimeoutException) as exc:
|
||||
result.errors.append(f"MTA-STS policy fetch failed: {exc}")
|
||||
return result
|
||||
|
||||
policy, policy_warnings, policy_errors = parse_mta_sts_policy(result.policy_text or "")
|
||||
result.warnings.extend(policy_warnings)
|
||||
result.errors.extend(policy_errors)
|
||||
result.mode = policy.get("mode")
|
||||
result.max_age = policy.get("max_age")
|
||||
result.mx = list(policy.get("mx") or [])
|
||||
result.status = "pass" if not result.errors else "fail"
|
||||
return result
|
||||
|
||||
|
||||
async def check_mta_sts_cached(
|
||||
db: Session,
|
||||
provider: BaseDNSProvider,
|
||||
domain: str,
|
||||
*,
|
||||
ttl_seconds: int = DEFAULT_DNS_CACHE_TTL_SECONDS,
|
||||
refresh: bool = False,
|
||||
) -> Tuple[MTAStsResult, bool, datetime]:
|
||||
"""Resolve MTA-STS posture, reusing the shared DNS cache semantics."""
|
||||
now = _utcnow_naive()
|
||||
provider_name = f"{provider.__class__.__name__}:mta-sts"
|
||||
row = (
|
||||
db.query(DNSCache)
|
||||
.filter(
|
||||
DNSCache.domain == domain,
|
||||
DNSCache.provider == provider_name,
|
||||
DNSCache.selectors_key == _CACHE_KEY,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if row and not refresh and _is_fresh(row, ttl_seconds, now):
|
||||
return _result_from_json(row.result_json), True, row.checked_at
|
||||
|
||||
result = await check_mta_sts(domain, provider)
|
||||
payload = json.dumps(asdict(result), sort_keys=True, separators=(",", ":"))
|
||||
if row is None:
|
||||
row = DNSCache(
|
||||
domain=domain,
|
||||
provider=provider_name,
|
||||
selectors_key=_CACHE_KEY,
|
||||
result_json=payload,
|
||||
checked_at=now,
|
||||
)
|
||||
db.add(row)
|
||||
else:
|
||||
row.result_json = payload
|
||||
row.checked_at = now
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except IntegrityError:
|
||||
db.rollback()
|
||||
row = (
|
||||
db.query(DNSCache)
|
||||
.filter(
|
||||
DNSCache.domain == domain,
|
||||
DNSCache.provider == provider_name,
|
||||
DNSCache.selectors_key == _CACHE_KEY,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if row is None:
|
||||
raise
|
||||
row.result_json = payload
|
||||
row.checked_at = now
|
||||
db.commit()
|
||||
|
||||
db.refresh(row)
|
||||
return result, False, row.checked_at
|
||||
@@ -0,0 +1,270 @@
|
||||
"""Notification delivery helpers backed by Apprise."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import asdict, dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
import apprise
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.credential_encryption import decrypt_secret
|
||||
from app.models.setting import Setting
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class NotificationResult:
|
||||
"""Sanitized result for a notification send attempt."""
|
||||
|
||||
success: bool
|
||||
message: str
|
||||
configured_targets: int = 0
|
||||
invalid_targets: int = 0
|
||||
skipped: bool = False
|
||||
rate_limited: bool = False
|
||||
error: Optional[str] = None
|
||||
|
||||
def to_dict(self) -> Dict[str, object]:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
def _truthy(value: Optional[str], default: bool = False) -> bool:
|
||||
if value in (None, ""):
|
||||
return default
|
||||
return str(value).strip().lower() in {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def _split_apprise_urls(value: Optional[str]) -> List[str]:
|
||||
if not value:
|
||||
return []
|
||||
return [
|
||||
line.strip()
|
||||
for line in value.splitlines()
|
||||
if line.strip() and not line.strip().startswith("#")
|
||||
]
|
||||
|
||||
|
||||
def _notification_settings(db: Session) -> Dict[str, Optional[str]]:
|
||||
rows = db.query(Setting).filter(Setting.category == "notifications").all()
|
||||
return {row.key: row.value for row in rows}
|
||||
|
||||
|
||||
def _decrypted_setting(settings: Dict[str, Optional[str]], key: str) -> Optional[str]:
|
||||
try:
|
||||
return decrypt_secret(settings.get(key))
|
||||
except ValueError:
|
||||
logger.exception("Encrypted notification setting could not be decrypted: %s", key)
|
||||
return None
|
||||
|
||||
|
||||
def _int_setting(value: Optional[str], default: int) -> int:
|
||||
try:
|
||||
return int(value or default)
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
|
||||
def _parse_timestamp(value: Optional[str]) -> Optional[datetime]:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError:
|
||||
return None
|
||||
if parsed.tzinfo is None:
|
||||
parsed = parsed.replace(tzinfo=timezone.utc)
|
||||
return parsed.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _set_notification_setting(db: Session, key: str, value: str) -> None:
|
||||
row = db.query(Setting).filter(Setting.key == key).first()
|
||||
if row is None:
|
||||
row = Setting(
|
||||
key=key,
|
||||
value=value,
|
||||
value_type="string",
|
||||
category="notifications",
|
||||
)
|
||||
db.add(row)
|
||||
else:
|
||||
row.value = value
|
||||
|
||||
|
||||
def _rate_limit_result(settings: Dict[str, Optional[str]]) -> Optional[NotificationResult]:
|
||||
interval_minutes = max(
|
||||
0, _int_setting(settings.get("notifications.min_send_interval_minutes"), 15)
|
||||
)
|
||||
if interval_minutes <= 0:
|
||||
return None
|
||||
|
||||
last_sent_at = _parse_timestamp(settings.get("notifications.last_sent_at"))
|
||||
if last_sent_at is None:
|
||||
return None
|
||||
|
||||
next_allowed = last_sent_at + timedelta(minutes=interval_minutes)
|
||||
now = datetime.now(timezone.utc)
|
||||
if now >= next_allowed:
|
||||
return None
|
||||
|
||||
retry_after = max(1, int((next_allowed - now).total_seconds() // 60) + 1)
|
||||
return NotificationResult(
|
||||
success=False,
|
||||
skipped=True,
|
||||
rate_limited=True,
|
||||
message=f"Notification rate limit active. Try again in about {retry_after} minute(s).",
|
||||
error="rate_limited",
|
||||
)
|
||||
|
||||
|
||||
_EMAIL_LOCAL_CHARS = frozenset(
|
||||
"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789._%+-"
|
||||
)
|
||||
_EMAIL_DOMAIN_CHARS = frozenset("ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789.-")
|
||||
_EMAIL_LEADING_PUNCTUATION = frozenset("\"'(<[{")
|
||||
_EMAIL_TRAILING_PUNCTUATION = frozenset("\"'.,;:!?)]}>")
|
||||
|
||||
|
||||
def _redact_email_token(token: str) -> str:
|
||||
leading = ""
|
||||
trailing = ""
|
||||
core = token
|
||||
|
||||
while core and core[0] in _EMAIL_LEADING_PUNCTUATION:
|
||||
leading += core[0]
|
||||
core = core[1:]
|
||||
while core and core[-1] in _EMAIL_TRAILING_PUNCTUATION:
|
||||
trailing = core[-1] + trailing
|
||||
core = core[:-1]
|
||||
|
||||
if core.count("@") != 1:
|
||||
return token
|
||||
|
||||
local_part, domain_part = core.split("@", 1)
|
||||
domain_labels = domain_part.split(".")
|
||||
if (
|
||||
not local_part
|
||||
or not domain_part
|
||||
or len(domain_labels) < 2
|
||||
or len(domain_labels[-1]) < 2
|
||||
or not domain_labels[-1].isalpha()
|
||||
or any(not label for label in domain_labels)
|
||||
or any(char not in _EMAIL_LOCAL_CHARS for char in local_part)
|
||||
or any(char not in _EMAIL_DOMAIN_CHARS for char in domain_part)
|
||||
):
|
||||
return token
|
||||
|
||||
return f"{leading}[redacted-email]@{domain_part}{trailing}"
|
||||
|
||||
|
||||
def redact_notification_text(value: str) -> str:
|
||||
"""Remove common PII from outbound notification text."""
|
||||
redacted = []
|
||||
token = []
|
||||
for char in value:
|
||||
if char.isspace():
|
||||
if token:
|
||||
redacted.append(_redact_email_token("".join(token)))
|
||||
token = []
|
||||
redacted.append(char)
|
||||
else:
|
||||
token.append(char)
|
||||
if token:
|
||||
redacted.append(_redact_email_token("".join(token)))
|
||||
return "".join(redacted)
|
||||
|
||||
|
||||
def _add_apprise_targets(notifier: apprise.Apprise, urls: List[str]) -> Tuple[int, int]:
|
||||
configured_targets = 0
|
||||
invalid_targets = 0
|
||||
for url in urls:
|
||||
try:
|
||||
if notifier.add(url):
|
||||
configured_targets += 1
|
||||
else:
|
||||
invalid_targets += 1
|
||||
except Exception: # pylint: disable=broad-exception-caught
|
||||
invalid_targets += 1
|
||||
logger.warning("Invalid Apprise notification target was ignored.")
|
||||
return configured_targets, invalid_targets
|
||||
|
||||
|
||||
def send_notification(
|
||||
db: Session,
|
||||
*,
|
||||
title: str,
|
||||
body: str,
|
||||
force: bool = False,
|
||||
) -> NotificationResult:
|
||||
"""Send a notification through configured Apprise target URLs."""
|
||||
settings = _notification_settings(db)
|
||||
enabled = _truthy(settings.get("notifications.apprise_enabled"))
|
||||
|
||||
if not enabled and not force:
|
||||
return NotificationResult(
|
||||
success=False,
|
||||
skipped=True,
|
||||
message="Notifications are disabled.",
|
||||
)
|
||||
|
||||
rate_limit = None if force else _rate_limit_result(settings)
|
||||
if rate_limit:
|
||||
return rate_limit
|
||||
|
||||
urls = _split_apprise_urls(_decrypted_setting(settings, "notifications.apprise_urls"))
|
||||
if not urls:
|
||||
return NotificationResult(
|
||||
success=False,
|
||||
message="No notification targets are configured.",
|
||||
)
|
||||
|
||||
notifier = apprise.Apprise()
|
||||
configured_targets, invalid_targets = _add_apprise_targets(notifier, urls)
|
||||
|
||||
if configured_targets == 0:
|
||||
return NotificationResult(
|
||||
success=False,
|
||||
message="No valid notification targets are configured.",
|
||||
invalid_targets=invalid_targets,
|
||||
)
|
||||
|
||||
try:
|
||||
if _truthy(settings.get("notifications.redact_pii_enabled"), default=True):
|
||||
title = redact_notification_text(title)
|
||||
body = redact_notification_text(body)
|
||||
success = bool(notifier.notify(title=title, body=body))
|
||||
except Exception: # pylint: disable=broad-exception-caught
|
||||
logger.exception("Apprise notification delivery failed.")
|
||||
return NotificationResult(
|
||||
success=False,
|
||||
message="Notification delivery failed.",
|
||||
configured_targets=configured_targets,
|
||||
invalid_targets=invalid_targets,
|
||||
error="delivery_failed",
|
||||
)
|
||||
|
||||
if not success:
|
||||
return NotificationResult(
|
||||
success=False,
|
||||
message="Notification delivery was not accepted by any configured target.",
|
||||
configured_targets=configured_targets,
|
||||
invalid_targets=invalid_targets,
|
||||
error="not_delivered",
|
||||
)
|
||||
|
||||
_set_notification_setting(
|
||||
db,
|
||||
"notifications.last_sent_at",
|
||||
datetime.now(timezone.utc).isoformat(),
|
||||
)
|
||||
db.commit()
|
||||
|
||||
return NotificationResult(
|
||||
success=True,
|
||||
message="Notification sent.",
|
||||
configured_targets=configured_targets,
|
||||
invalid_targets=invalid_targets,
|
||||
)
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user