Merge pull request #24 from christianlouis/copilot/update-pylint-warnings-in-main
Fix pylint/isort/flake8/black contradictions; consolidate linter config; pylint 10.00/10
This commit is contained in:
@@ -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
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
from app.api.api_v1.endpoints import domains, health, imap, reports, setup, stats
|
|
||||||
from fastapi import APIRouter
|
from fastapi import APIRouter
|
||||||
|
|
||||||
|
from app.api.api_v1.endpoints import domains, health, imap, reports, setup, stats
|
||||||
|
|
||||||
api_router = APIRouter()
|
api_router = APIRouter()
|
||||||
|
|
||||||
# Include all endpoint routers
|
# Include all endpoint routers
|
||||||
|
|||||||
@@ -2,10 +2,11 @@ import random # Used for mock data generation - TODO: Replace with actual histo
|
|||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
from app.services.report_store import ReportStore
|
|
||||||
from fastapi import APIRouter, HTTPException, Path, Query, status
|
from fastapi import APIRouter, HTTPException, Path, Query, status
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from app.services.report_store import ReportStore
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
|
|
||||||
@@ -248,9 +249,13 @@ async def get_domain_dns_records(domain_id: str = Path(..., title="The domain ID
|
|||||||
|
|
||||||
# For Milestone 1, return mock DNS record data
|
# For Milestone 1, return mock DNS record data
|
||||||
# In a future milestone, this will be replaced with actual DNS lookups
|
# In a future milestone, this will be replaced with actual DNS lookups
|
||||||
|
mock_dmarc_record = (
|
||||||
|
"v=DMARC1; p=none; rua=mailto:dmarc@example.com;"
|
||||||
|
" ruf=mailto:forensic@example.com; pct=100"
|
||||||
|
)
|
||||||
return DNSRecordResponse(
|
return DNSRecordResponse(
|
||||||
dmarc=True,
|
dmarc=True,
|
||||||
dmarcRecord="v=DMARC1; p=none; rua=mailto:dmarc@example.com; ruf=mailto:forensic@example.com; pct=100",
|
dmarcRecord=mock_dmarc_record,
|
||||||
spf=True,
|
spf=True,
|
||||||
spfRecord="v=spf1 include:_spf.google.com include:spf.protection.outlook.com -all",
|
spfRecord="v=spf1 include:_spf.google.com include:spf.protection.outlook.com -all",
|
||||||
dkim=True,
|
dkim=True,
|
||||||
@@ -299,7 +304,7 @@ async def get_domain_reports(
|
|||||||
date = datetime.now() - timedelta(days=i)
|
date = datetime.now() - timedelta(days=i)
|
||||||
date_str = date.strftime("%Y-%m-%d")
|
date_str = date.strftime("%Y-%m-%d")
|
||||||
|
|
||||||
# TODO: Replace with actual historical data in future milestone
|
# TODO: Replace with actual historical data in future milestone # pylint: disable=fixme
|
||||||
# For now, generate mock data with variation for demonstration purposes
|
# For now, generate mock data with variation for demonstration purposes
|
||||||
compliance_rate = random.uniform(80, 100) # nosec B311 - Mock data only
|
compliance_rate = random.uniform(80, 100) # nosec B311 - Mock data only
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
from app.api.api_v1.endpoints.setup import setup_status
|
|
||||||
from fastapi import APIRouter
|
from fastapi import APIRouter
|
||||||
|
|
||||||
|
from app.api.api_v1.endpoints.setup import setup_status
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2,11 +2,12 @@ import logging
|
|||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Any, Dict, Optional
|
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 fastapi import APIRouter, BackgroundTasks, Depends, HTTPException
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from app.core.security import require_admin_auth
|
||||||
|
from app.services.imap_client import IMAPClient
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|||||||
@@ -1,11 +1,12 @@
|
|||||||
import logging
|
import logging
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
|
from fastapi import APIRouter, File, HTTPException, UploadFile, status
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from app.services.dmarc_parser import DMARCParser
|
from app.services.dmarc_parser import DMARCParser
|
||||||
from app.services.report_store import ReportStore
|
from app.services.report_store import ReportStore
|
||||||
from app.utils.domain_validator import DomainValidationError, validate_domain
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -47,16 +48,16 @@ def _validate_mime_type(file_content: bytes) -> None:
|
|||||||
try:
|
try:
|
||||||
mime_type = magic.from_buffer(file_content, mime=True)
|
mime_type = magic.from_buffer(file_content, mime=True)
|
||||||
if mime_type not in ALLOWED_MIME_TYPES:
|
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(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
detail="Invalid file type. File must be XML, ZIP, or GZIP format.",
|
detail="Invalid file type. File must be XML, ZIP, or GZIP format.",
|
||||||
)
|
)
|
||||||
except HTTPException:
|
except HTTPException:
|
||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e: # pylint: disable=broad-exception-caught
|
||||||
# If magic fails, log but continue (fallback to extension check)
|
# 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:
|
def _validate_upload_file(file: UploadFile, file_content: bytes) -> None:
|
||||||
@@ -89,15 +90,14 @@ def _handle_upload_value_error(filename: str, error_message: str) -> None:
|
|||||||
|
|
||||||
Always raises — never returns.
|
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():
|
if "too large" in error_message.lower():
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, detail="File too large"
|
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, detail="File too large"
|
||||||
)
|
)
|
||||||
elif "zip bomb" in error_message.lower():
|
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 archive file")
|
||||||
else:
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid report format")
|
||||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid report format")
|
|
||||||
|
|
||||||
|
|
||||||
class UploadResponse(BaseModel):
|
class UploadResponse(BaseModel):
|
||||||
@@ -196,13 +196,13 @@ async def upload_report(file: UploadFile = File(...)):
|
|||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
# Security: Sanitize error messages from parser
|
# Security: Sanitize error messages from parser
|
||||||
_handle_upload_value_error(file.filename, str(e))
|
_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
|
# 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(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||||
detail="Error processing report. Please contact support if this persists.",
|
detail="Error processing report. Please contact support if this persists.",
|
||||||
)
|
) from e
|
||||||
|
|
||||||
|
|
||||||
@router.get("/domains", response_model=List[str])
|
@router.get("/domains", response_model=List[str])
|
||||||
@@ -300,10 +300,11 @@ async def get_domain_reports_paginated(
|
|||||||
|
|
||||||
if sort_field == "total_count":
|
if sort_field == "total_count":
|
||||||
all_reports.sort(
|
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:
|
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
|
# Apply pagination
|
||||||
total = len(all_reports)
|
total = len(all_reports)
|
||||||
|
|||||||
@@ -1,9 +1,10 @@
|
|||||||
from typing import Any, Dict
|
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.core.database import get_db
|
||||||
from app.utils.stats_summarizer import StatsSummarizer
|
from app.utils.stats_summarizer import StatsSummarizer
|
||||||
from fastapi import APIRouter, Depends, Path, Query
|
|
||||||
from sqlalchemy.orm import Session
|
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from typing import List, Optional, Union
|
|||||||
|
|
||||||
# Try to import from pydantic_settings first (newer versions)
|
# Try to import from pydantic_settings first (newer versions)
|
||||||
try:
|
try:
|
||||||
from pydantic import EmailStr, validator
|
from pydantic import EmailStr, validator # pylint: disable=ungrouped-imports
|
||||||
from pydantic_settings import BaseSettings
|
from pydantic_settings import BaseSettings
|
||||||
except ImportError:
|
except ImportError:
|
||||||
# Fall back to older pydantic version
|
# Fall back to older pydantic version
|
||||||
@@ -47,7 +47,7 @@ class Settings(BaseSettings):
|
|||||||
CLOUDFLARE_ZONE_ID: Optional[str] = None
|
CLOUDFLARE_ZONE_ID: Optional[str] = None
|
||||||
|
|
||||||
@validator("SECRET_KEY", pre=True, always=True)
|
@validator("SECRET_KEY", pre=True, always=True)
|
||||||
def validate_secret_key(cls, v: Optional[str]) -> str:
|
def validate_secret_key(cls, v: Optional[str]) -> str: # pylint: disable=no-self-argument
|
||||||
"""Validate and generate SECRET_KEY if not provided."""
|
"""Validate and generate SECRET_KEY if not provided."""
|
||||||
# Default insecure key that should never be used
|
# Default insecure key that should never be used
|
||||||
DEFAULT_INSECURE_KEY = "CHANGE_THIS_TO_A_RANDOM_SECRET_IN_PRODUCTION"
|
DEFAULT_INSECURE_KEY = "CHANGE_THIS_TO_A_RANDOM_SECRET_IN_PRODUCTION"
|
||||||
@@ -66,17 +66,20 @@ class Settings(BaseSettings):
|
|||||||
# Check if key is too short
|
# Check if key is too short
|
||||||
if len(v) < 32:
|
if len(v) < 32:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"SECRET_KEY is too short ({len(v)} characters). "
|
"SECRET_KEY is too short (%s characters). "
|
||||||
"Recommended minimum is 32 characters for security."
|
"Recommended minimum is 32 characters for security.",
|
||||||
|
len(v),
|
||||||
)
|
)
|
||||||
|
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@validator("BACKEND_CORS_ORIGINS", pre=True)
|
@validator("BACKEND_CORS_ORIGINS", pre=True)
|
||||||
def assemble_cors_origins(cls, v: Union[str, List[str]]) -> List[str]:
|
def assemble_cors_origins( # pylint: disable=no-self-argument
|
||||||
|
cls, v: Union[str, List[str]]
|
||||||
|
) -> List[str]:
|
||||||
if isinstance(v, str) and not v.startswith("["):
|
if isinstance(v, str) and not v.startswith("["):
|
||||||
return [i.strip() for i in v.split(",")]
|
return [i.strip() for i in v.split(",")]
|
||||||
elif isinstance(v, (list, str)):
|
if isinstance(v, (list, str)):
|
||||||
return v
|
return v
|
||||||
raise ValueError(v)
|
raise ValueError(v)
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +1,11 @@
|
|||||||
from typing import Generator
|
from typing import Generator
|
||||||
|
|
||||||
from app.core.config import get_settings
|
|
||||||
from sqlalchemy import create_engine
|
from sqlalchemy import create_engine
|
||||||
from sqlalchemy.ext.declarative import declarative_base
|
from sqlalchemy.ext.declarative import declarative_base
|
||||||
from sqlalchemy.orm import sessionmaker
|
from sqlalchemy.orm import sessionmaker
|
||||||
|
|
||||||
|
from app.core.config import get_settings
|
||||||
|
|
||||||
settings = get_settings()
|
settings = get_settings()
|
||||||
|
|
||||||
# Configure SQLAlchemy
|
# Configure SQLAlchemy
|
||||||
|
|||||||
@@ -4,12 +4,13 @@ import secrets
|
|||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
from typing import Any, Optional, Union
|
from typing import Any, Optional, Union
|
||||||
|
|
||||||
from app.core.config import get_settings
|
|
||||||
from fastapi import HTTPException, Security, status
|
from fastapi import HTTPException, Security, status
|
||||||
from fastapi.security import APIKeyHeader, HTTPAuthorizationCredentials, HTTPBearer
|
from fastapi.security import APIKeyHeader, HTTPAuthorizationCredentials, HTTPBearer
|
||||||
from jose import JWTError, jwt
|
from jose import JWTError, jwt
|
||||||
from passlib.context import CryptContext
|
from passlib.context import CryptContext
|
||||||
|
|
||||||
|
from app.core.config import get_settings
|
||||||
|
|
||||||
settings = get_settings()
|
settings = get_settings()
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -75,7 +76,7 @@ def add_api_key(api_key: str) -> bool:
|
|||||||
if api_key in _api_keys:
|
if api_key in _api_keys:
|
||||||
return False
|
return False
|
||||||
_api_keys.add(api_key)
|
_api_keys.add(api_key)
|
||||||
logger.info(f"API key added (ends with: ...{api_key[-8:]})")
|
logger.info("API key added (ends with: ...%s)", api_key[-8:])
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
@@ -92,7 +93,7 @@ def verify_api_key(api_key: str) -> bool:
|
|||||||
return api_key in _api_keys
|
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.
|
Dependency to verify API key authentication.
|
||||||
|
|
||||||
@@ -105,24 +106,23 @@ async def get_api_key(api_key_header: Optional[str] = Security(api_key_header))
|
|||||||
Raises:
|
Raises:
|
||||||
HTTPException: If API key is missing or invalid
|
HTTPException: If API key is missing or invalid
|
||||||
"""
|
"""
|
||||||
if not api_key_header:
|
if not api_key_value:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||||
detail="Missing API key",
|
detail="Missing API key",
|
||||||
headers={"WWW-Authenticate": "ApiKey"},
|
headers={"WWW-Authenticate": "ApiKey"},
|
||||||
)
|
)
|
||||||
|
|
||||||
if not verify_api_key(api_key_header):
|
if not verify_api_key(api_key_value):
|
||||||
logger.warning(
|
suffix = api_key_value[-8:] if len(api_key_value) >= 8 else "invalid"
|
||||||
f"Invalid API key attempt: ...{api_key_header[-8:] if len(api_key_header) >= 8 else 'invalid'}"
|
logger.warning("Invalid API key attempt: ...%s", suffix)
|
||||||
)
|
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||||
detail="Invalid API key",
|
detail="Invalid API key",
|
||||||
headers={"WWW-Authenticate": "ApiKey"},
|
headers={"WWW-Authenticate": "ApiKey"},
|
||||||
)
|
)
|
||||||
|
|
||||||
return api_key_header
|
return api_key_value
|
||||||
|
|
||||||
|
|
||||||
async def verify_token(
|
async def verify_token(
|
||||||
@@ -153,12 +153,12 @@ async def verify_token(
|
|||||||
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
|
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
|
||||||
return payload
|
return payload
|
||||||
except JWTError as e:
|
except JWTError as e:
|
||||||
logger.warning(f"Invalid JWT token: {str(e)}")
|
logger.warning("Invalid JWT token: %s", str(e))
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||||
detail="Invalid authentication token",
|
detail="Invalid authentication token",
|
||||||
headers={"WWW-Authenticate": "Bearer"},
|
headers={"WWW-Authenticate": "Bearer"},
|
||||||
)
|
) from e
|
||||||
|
|
||||||
|
|
||||||
async def require_admin_auth(
|
async def require_admin_auth(
|
||||||
@@ -192,7 +192,7 @@ async def require_admin_auth(
|
|||||||
)
|
)
|
||||||
return {"auth_type": "jwt", "payload": payload}
|
return {"auth_type": "jwt", "payload": payload}
|
||||||
except JWTError as e:
|
except JWTError as e:
|
||||||
logger.warning(f"Invalid JWT token: {str(e)}")
|
logger.warning("Invalid JWT token: %s", str(e))
|
||||||
|
|
||||||
# No valid authentication provided
|
# No valid authentication provided
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
|
|||||||
+36
-35
@@ -3,17 +3,18 @@ import logging
|
|||||||
import os
|
import os
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
|
||||||
|
from fastapi import Depends, FastAPI, Request
|
||||||
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
|
from fastapi.responses import HTMLResponse
|
||||||
|
from fastapi.staticfiles import StaticFiles
|
||||||
|
from fastapi.templating import Jinja2Templates
|
||||||
|
|
||||||
from app.api.api_v1.api import api_router
|
from app.api.api_v1.api import api_router
|
||||||
from app.core.config import get_settings
|
from app.core.config import get_settings
|
||||||
from app.core.security import add_api_key, generate_api_key, require_admin_auth
|
from app.core.security import add_api_key, generate_api_key, require_admin_auth
|
||||||
from app.middleware.security import SecurityHeadersMiddleware
|
from app.middleware.security import SecurityHeadersMiddleware
|
||||||
from app.services.imap_client import IMAPClient
|
from app.services.imap_client import IMAPClient
|
||||||
from app.services.report_store import ReportStore
|
from app.services.report_store import ReportStore
|
||||||
from fastapi import BackgroundTasks, Depends, FastAPI, Request
|
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
|
||||||
from fastapi.responses import HTMLResponse
|
|
||||||
from fastapi.staticfiles import StaticFiles
|
|
||||||
from fastapi.templating import Jinja2Templates
|
|
||||||
|
|
||||||
# Set up logging
|
# Set up logging
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -27,7 +28,7 @@ last_check_time = None
|
|||||||
|
|
||||||
async def scheduled_imap_polling():
|
async def scheduled_imap_polling():
|
||||||
"""Background task for periodically checking IMAP for new DMARC reports"""
|
"""Background task for periodically checking IMAP for new DMARC reports"""
|
||||||
global last_check_time
|
global last_check_time # pylint: disable=global-statement
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# How often to check for emails (in seconds)
|
# How often to check for emails (in seconds)
|
||||||
@@ -46,19 +47,20 @@ async def scheduled_imap_polling():
|
|||||||
|
|
||||||
if results["success"]:
|
if results["success"]:
|
||||||
logger.info(
|
logger.info(
|
||||||
f"IMAP polling completed: {results['processed']} emails processed, "
|
"IMAP polling completed: %s emails processed, %s reports found",
|
||||||
f"{results['reports_found']} reports found"
|
results["processed"],
|
||||||
|
results["reports_found"],
|
||||||
)
|
)
|
||||||
|
|
||||||
# If new domains were found, log them
|
# If new domains were found, log them
|
||||||
if results["new_domains"]:
|
if results["new_domains"]:
|
||||||
logger.info(f"New domains found: {', '.join(results['new_domains'])}")
|
logger.info("New domains found: %s", ", ".join(results["new_domains"]))
|
||||||
|
|
||||||
else:
|
else:
|
||||||
logger.error(f"IMAP polling failed: {results.get('error', 'Unknown error')}")
|
logger.error("IMAP polling failed: %s", results.get("error", "Unknown error"))
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e: # pylint: disable=broad-exception-caught
|
||||||
logger.error(f"Error in IMAP polling task: {str(e)}")
|
logger.error("Error in IMAP polling task: %s", str(e))
|
||||||
|
|
||||||
# Wait for the next check interval
|
# Wait for the next check interval
|
||||||
await asyncio.sleep(check_interval)
|
await asyncio.sleep(check_interval)
|
||||||
@@ -69,7 +71,7 @@ async def scheduled_imap_polling():
|
|||||||
|
|
||||||
def create_app() -> FastAPI:
|
def create_app() -> FastAPI:
|
||||||
"""Create and configure the FastAPI application"""
|
"""Create and configure the FastAPI application"""
|
||||||
app = FastAPI(
|
application = FastAPI(
|
||||||
title=settings.PROJECT_NAME,
|
title=settings.PROJECT_NAME,
|
||||||
openapi_url=f"{settings.API_V1_STR}/openapi.json",
|
openapi_url=f"{settings.API_V1_STR}/openapi.json",
|
||||||
version="0.1.0",
|
version="0.1.0",
|
||||||
@@ -78,11 +80,11 @@ def create_app() -> FastAPI:
|
|||||||
# Add security headers middleware
|
# Add security headers middleware
|
||||||
# Determine environment from settings or environment variable
|
# Determine environment from settings or environment variable
|
||||||
environment = os.getenv("ENVIRONMENT", "development")
|
environment = os.getenv("ENVIRONMENT", "development")
|
||||||
app.add_middleware(SecurityHeadersMiddleware, environment=environment)
|
application.add_middleware(SecurityHeadersMiddleware, environment=environment)
|
||||||
|
|
||||||
# Improved CORS configuration - restrict to specific methods and headers
|
# Improved CORS configuration - restrict to specific methods and headers
|
||||||
if settings.BACKEND_CORS_ORIGINS:
|
if settings.BACKEND_CORS_ORIGINS:
|
||||||
app.add_middleware(
|
application.add_middleware(
|
||||||
CORSMiddleware,
|
CORSMiddleware,
|
||||||
allow_origins=[str(origin) for origin in settings.BACKEND_CORS_ORIGINS],
|
allow_origins=[str(origin) for origin in settings.BACKEND_CORS_ORIGINS],
|
||||||
allow_credentials=True,
|
allow_credentials=True,
|
||||||
@@ -103,20 +105,20 @@ def create_app() -> FastAPI:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Include API router
|
# Include API router
|
||||||
app.include_router(api_router, prefix=settings.API_V1_STR)
|
application.include_router(api_router, prefix=settings.API_V1_STR)
|
||||||
|
|
||||||
# Mount static files directory
|
# Mount static files directory
|
||||||
app.mount(
|
application.mount(
|
||||||
"/static",
|
"/static",
|
||||||
StaticFiles(directory=os.path.join(os.path.dirname(__file__), "static")),
|
StaticFiles(directory=os.path.join(os.path.dirname(__file__), "static")),
|
||||||
name="static",
|
name="static",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Set up event handlers for startup and shutdown
|
# Set up event handlers for startup and shutdown
|
||||||
@app.on_event("startup")
|
@application.on_event("startup")
|
||||||
async def startup_event():
|
async def startup_event():
|
||||||
"""Initialize background tasks and security on application startup"""
|
"""Initialize background tasks and security on application startup"""
|
||||||
global background_task
|
global background_task # pylint: disable=global-statement
|
||||||
|
|
||||||
# Generate and provide admin API key
|
# Generate and provide admin API key
|
||||||
api_key = generate_api_key()
|
api_key = generate_api_key()
|
||||||
@@ -124,19 +126,20 @@ def create_app() -> FastAPI:
|
|||||||
|
|
||||||
# Security: Log only last 8 characters for reference
|
# Security: Log only last 8 characters for reference
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"=" * 80 + "\n"
|
"%s\nIMPORTANT: Admin API Key Generated\n"
|
||||||
"IMPORTANT: Admin API Key Generated\n"
|
"API Key (last 8 chars): ...%s\n"
|
||||||
f"API Key (last 8 chars): ...{api_key[-8:]}\n"
|
|
||||||
"Full key stored securely in memory.\n"
|
"Full key stored securely in memory.\n"
|
||||||
"For production, retrieve the key through secure configuration management.\n"
|
"For production, retrieve the key through secure configuration management.\n"
|
||||||
"Use this key in the X-API-Key header for admin endpoints.\n"
|
"Use this key in the X-API-Key header for admin endpoints.\n%s",
|
||||||
"=" * 80
|
"=" * 80,
|
||||||
|
api_key[-8:],
|
||||||
|
"=" * 80,
|
||||||
)
|
)
|
||||||
|
|
||||||
# In development, also log the full key for convenience
|
# In development, also log the full key for convenience
|
||||||
# This should be removed in production or controlled by environment variable
|
# This should be removed in production or controlled by environment variable
|
||||||
if os.getenv("ENVIRONMENT", "development") == "development":
|
if os.getenv("ENVIRONMENT", "development") == "development":
|
||||||
logger.info(f"Development Mode - Full API Key: {api_key}")
|
logger.info("Development Mode - Full API Key: %s", api_key)
|
||||||
|
|
||||||
# Check if IMAP credentials are configured
|
# Check if IMAP credentials are configured
|
||||||
if all([settings.IMAP_SERVER, settings.IMAP_USERNAME, settings.IMAP_PASSWORD]):
|
if all([settings.IMAP_SERVER, settings.IMAP_USERNAME, settings.IMAP_PASSWORD]):
|
||||||
@@ -145,7 +148,7 @@ def create_app() -> FastAPI:
|
|||||||
else:
|
else:
|
||||||
logger.warning("IMAP credentials not fully configured, polling disabled")
|
logger.warning("IMAP credentials not fully configured, polling disabled")
|
||||||
|
|
||||||
@app.on_event("shutdown")
|
@application.on_event("shutdown")
|
||||||
async def shutdown_event():
|
async def shutdown_event():
|
||||||
"""Clean up background tasks on application shutdown"""
|
"""Clean up background tasks on application shutdown"""
|
||||||
if background_task:
|
if background_task:
|
||||||
@@ -156,7 +159,7 @@ def create_app() -> FastAPI:
|
|||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
return app
|
return application
|
||||||
|
|
||||||
|
|
||||||
app = create_app()
|
app = create_app()
|
||||||
@@ -202,9 +205,9 @@ async def domains(request: Request):
|
|||||||
async def domain_details(request: Request, domain_id: str):
|
async def domain_details(request: Request, domain_id: str):
|
||||||
"""View detailed reports for a specific domain"""
|
"""View detailed reports for a specific domain"""
|
||||||
store = ReportStore.get_instance()
|
store = ReportStore.get_instance()
|
||||||
domains = store.get_domains()
|
known_domains = store.get_domains()
|
||||||
|
|
||||||
if domain_id not in domains:
|
if domain_id not in known_domains:
|
||||||
# Domain not found, redirect to domains list
|
# Domain not found, redirect to domains list
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
"domains.html", {"request": request, "error": f"Domain {domain_id} not found"}
|
"domains.html", {"request": request, "error": f"Domain {domain_id} not found"}
|
||||||
@@ -243,15 +246,13 @@ async def upload_page(request: Request):
|
|||||||
|
|
||||||
# API endpoint to manually trigger IMAP polling
|
# API endpoint to manually trigger IMAP polling
|
||||||
@app.post("/api/v1/admin/trigger-poll")
|
@app.post("/api/v1/admin/trigger-poll")
|
||||||
async def trigger_imap_poll(
|
async def trigger_imap_poll(auth: dict = Depends(require_admin_auth)):
|
||||||
background_tasks: BackgroundTasks, auth: dict = Depends(require_admin_auth)
|
|
||||||
):
|
|
||||||
"""
|
"""
|
||||||
Manually trigger IMAP polling (admin only - requires authentication)
|
Manually trigger IMAP polling (admin only - requires authentication)
|
||||||
|
|
||||||
Security: Requires either X-API-Key header or Bearer token
|
Security: Requires either X-API-Key header or Bearer token
|
||||||
"""
|
"""
|
||||||
global last_check_time
|
global last_check_time # pylint: disable=global-statement
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Create IMAP client and fetch reports
|
# Create IMAP client and fetch reports
|
||||||
@@ -269,8 +270,8 @@ async def trigger_imap_poll(
|
|||||||
"new_domains": results["new_domains"],
|
"new_domains": results["new_domains"],
|
||||||
"authenticated_by": auth.get("auth_type"),
|
"authenticated_by": auth.get("auth_type"),
|
||||||
}
|
}
|
||||||
except Exception as e:
|
except Exception as e: # pylint: disable=broad-exception-caught
|
||||||
logger.error(f"Error triggering IMAP poll: {str(e)}")
|
logger.error("Error triggering IMAP poll: %s", str(e))
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": False,
|
||||||
"error": "Failed to trigger IMAP poll. Check server logs for details.",
|
"error": "Failed to trigger IMAP poll. Check server logs for details.",
|
||||||
|
|||||||
@@ -78,11 +78,13 @@ class SecurityHeadersMiddleware(BaseHTTPMiddleware):
|
|||||||
# See: https://developer.mozilla.org/en-US/docs/Web/HTTP/CSP
|
# See: https://developer.mozilla.org/en-US/docs/Web/HTTP/CSP
|
||||||
csp_directives = [
|
csp_directives = [
|
||||||
"default-src 'self'",
|
"default-src 'self'",
|
||||||
# TODO: Remove 'unsafe-inline' - requires moving inline scripts to external files
|
# TODO: Remove 'unsafe-inline' - requires moving inline scripts to external files # pylint: disable=fixme
|
||||||
# TODO: Remove 'unsafe-eval' - no eval usage detected, safe to remove after testing
|
# TODO: Remove 'unsafe-eval' - no eval usage detected, safe to remove after testing # pylint: disable=fixme
|
||||||
"script-src 'self' 'unsafe-inline' 'unsafe-eval' https://cdn.tailwindcss.com https://cdn.jsdelivr.net",
|
"script-src 'self' 'unsafe-inline' 'unsafe-eval'"
|
||||||
# TODO: Remove 'unsafe-inline' - requires moving inline styles to CSS or using nonces
|
" https://cdn.tailwindcss.com https://cdn.jsdelivr.net",
|
||||||
"style-src 'self' 'unsafe-inline' https://fonts.googleapis.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",
|
"font-src 'self' https://fonts.gstatic.com",
|
||||||
"img-src 'self' data: https:",
|
"img-src 'self' data: https:",
|
||||||
"connect-src 'self'",
|
"connect-src 'self'",
|
||||||
|
|||||||
@@ -1,9 +1,10 @@
|
|||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
|
||||||
from app.core.database import Base
|
|
||||||
from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Index, Integer, String, Text
|
from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Index, Integer, String, Text
|
||||||
from sqlalchemy.orm import relationship
|
from sqlalchemy.orm import relationship
|
||||||
|
|
||||||
|
from app.core.database import Base
|
||||||
|
|
||||||
|
|
||||||
class Domain(Base):
|
class Domain(Base):
|
||||||
"""Domain model representing a monitored domain"""
|
"""Domain model representing a monitored domain"""
|
||||||
|
|||||||
@@ -1,9 +1,10 @@
|
|||||||
from datetime import datetime
|
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
|
||||||
from sqlalchemy.orm import relationship
|
from sqlalchemy.orm import relationship
|
||||||
|
|
||||||
|
from app.core.database import Base
|
||||||
|
|
||||||
|
|
||||||
class DMARCReport(Base):
|
class DMARCReport(Base):
|
||||||
"""DMARC Aggregate Report model"""
|
"""DMARC Aggregate Report model"""
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
from app.core.database import Base
|
|
||||||
from sqlalchemy import Boolean, Column, Integer, String
|
from sqlalchemy import Boolean, Column, Integer, String
|
||||||
from sqlalchemy.orm import relationship
|
from sqlalchemy.orm import relationship
|
||||||
|
|
||||||
|
from app.core.database import Base
|
||||||
|
|
||||||
|
|
||||||
class User(Base):
|
class User(Base):
|
||||||
"""User model"""
|
"""User model"""
|
||||||
|
|||||||
@@ -236,20 +236,22 @@ class DMARCParser:
|
|||||||
|
|
||||||
# Log parse results for debugging
|
# Log parse results for debugging
|
||||||
total_count = report["summary"]["total_count"]
|
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(f"Found {len(records)} record entries with {total_count} total messages")
|
logger.info("Found %s record entries with %s total messages", len(records), total_count)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Messages passed: {report['summary']['passed_count']}, "
|
"Messages passed: %s, failed: %s",
|
||||||
f"failed: {report['summary']['failed_count']}"
|
report["summary"]["passed_count"],
|
||||||
|
report["summary"]["failed_count"],
|
||||||
)
|
)
|
||||||
if records:
|
if records:
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Sample record - SPF: {records[0].get('spf_result')}, "
|
"Sample record - SPF: %s, DKIM: %s",
|
||||||
f"DKIM: {records[0].get('dkim_result')}"
|
records[0].get("spf_result"),
|
||||||
|
records[0].get("dkim_result"),
|
||||||
)
|
)
|
||||||
|
|
||||||
return report
|
return report
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(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)}")
|
raise ValueError(f"Error parsing DMARC XML: {str(e)}") from e
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ class IMAPClient:
|
|||||||
Client for retrieving DMARC reports from an IMAP mailbox
|
Client for retrieving DMARC reports from an IMAP mailbox
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__( # pylint: disable=too-many-positional-arguments,too-many-arguments
|
||||||
self,
|
self,
|
||||||
server: str = None,
|
server: str = None,
|
||||||
port: int = None,
|
port: int = None,
|
||||||
@@ -63,7 +63,7 @@ class IMAPClient:
|
|||||||
if mailbox_name.startswith(" "):
|
if mailbox_name.startswith(" "):
|
||||||
mailbox_name = mailbox_name[1:]
|
mailbox_name = mailbox_name[1:]
|
||||||
available_mailboxes.append(mailbox_name)
|
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
|
# Silently skip mailboxes that can't be parsed; they are simply
|
||||||
# omitted from the returned list so callers should expect it may
|
# omitted from the returned list so callers should expect it may
|
||||||
# be incomplete. Some IMAP servers return non-standard list
|
# be incomplete. Some IMAP servers return non-standard list
|
||||||
@@ -130,8 +130,8 @@ class IMAPClient:
|
|||||||
}
|
}
|
||||||
|
|
||||||
return True, "Connection successful", stats
|
return True, "Connection successful", stats
|
||||||
except Exception as e:
|
except Exception as e: # pylint: disable=broad-exception-caught
|
||||||
logger.error(f"IMAP connection test failed: {str(e)}")
|
logger.error("IMAP connection test failed: %s", str(e))
|
||||||
return False, f"Connection failed: {str(e)}", {}
|
return False, f"Connection failed: {str(e)}", {}
|
||||||
|
|
||||||
def _process_single_email(self, mail, email_id: bytes, stats: dict) -> None:
|
def _process_single_email(self, mail, email_id: bytes, stats: dict) -> None:
|
||||||
@@ -139,7 +139,7 @@ class IMAPClient:
|
|||||||
try:
|
try:
|
||||||
status, msg_data = mail.fetch(email_id, "(RFC822)")
|
status, msg_data = mail.fetch(email_id, "(RFC822)")
|
||||||
if status != "OK":
|
if status != "OK":
|
||||||
logger.error(f"Error fetching email ID {email_id}")
|
logger.error("Error fetching email ID %s", email_id)
|
||||||
return
|
return
|
||||||
|
|
||||||
raw_email = msg_data[0][1]
|
raw_email = msg_data[0][1]
|
||||||
@@ -155,7 +155,7 @@ class IMAPClient:
|
|||||||
mail.store(email_id, "+FLAGS", "\\Deleted")
|
mail.store(email_id, "+FLAGS", "\\Deleted")
|
||||||
|
|
||||||
stats["processed"] += 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)}"
|
error_msg = f"Error processing email ID {email_id}: {str(e)}"
|
||||||
logger.error(error_msg)
|
logger.error(error_msg)
|
||||||
stats["errors"].append(error_msg)
|
stats["errors"].append(error_msg)
|
||||||
@@ -225,8 +225,8 @@ class IMAPClient:
|
|||||||
|
|
||||||
return stats
|
return stats
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e: # pylint: disable=broad-exception-caught
|
||||||
logger.error(f"Error fetching DMARC reports: {str(e)}")
|
logger.error("Error fetching DMARC reports: %s", str(e))
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": False,
|
||||||
"error": f"Error connecting to mailbox: {str(e)}",
|
"error": f"Error connecting to mailbox: {str(e)}",
|
||||||
@@ -339,12 +339,12 @@ class IMAPClient:
|
|||||||
|
|
||||||
# Check content type
|
# Check content type
|
||||||
content_type = part.get_content_type()
|
content_type = part.get_content_type()
|
||||||
if (
|
if content_type in (
|
||||||
content_type == "application/zip"
|
"application/zip",
|
||||||
or content_type == "application/gzip"
|
"application/gzip",
|
||||||
or content_type == "application/x-gzip"
|
"application/x-gzip",
|
||||||
or content_type == "application/xml"
|
"application/xml",
|
||||||
or content_type == "text/xml"
|
"text/xml",
|
||||||
):
|
):
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -390,8 +390,8 @@ class IMAPClient:
|
|||||||
self.report_store.add_report(report)
|
self.report_store.add_report(report)
|
||||||
|
|
||||||
reports_found += 1
|
reports_found += 1
|
||||||
logger.info(f"Successfully processed DMARC report: {filename}")
|
logger.info("Successfully processed DMARC report: %s", filename)
|
||||||
except Exception as e:
|
except Exception as e: # pylint: disable=broad-exception-caught
|
||||||
logger.error(f"Error processing attachment {filename}: {str(e)}")
|
logger.error("Error processing attachment %s: %s", filename, str(e))
|
||||||
|
|
||||||
return reports_found
|
return reports_found
|
||||||
|
|||||||
@@ -155,7 +155,7 @@ class ReportStore:
|
|||||||
return sorted_reports[:limit]
|
return sorted_reports[:limit]
|
||||||
return sorted_reports
|
return sorted_reports
|
||||||
|
|
||||||
def get_domain_sources(self, domain: str, days: int = 30) -> List[Dict[str, Any]]:
|
def get_domain_sources(self, domain: str, _days: int = 30) -> List[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
Get sending sources for a domain
|
Get sending sources for a domain
|
||||||
|
|
||||||
@@ -206,6 +206,6 @@ class ReportStore:
|
|||||||
self.domain_summary.pop(domain, None)
|
self.domain_summary.pop(domain, None)
|
||||||
self.domain_sources.pop(domain, None)
|
self.domain_sources.pop(domain, None)
|
||||||
return True
|
return True
|
||||||
except Exception:
|
except Exception: # pylint: disable=broad-exception-caught
|
||||||
# If any exception occurs during deletion, return False
|
# If any exception occurs during deletion, return False
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -1,21 +1,21 @@
|
|||||||
# Import all models so Base.metadata knows every table
|
# Import all models so Base.metadata knows every table
|
||||||
import app.models.domain # noqa: F401
|
|
||||||
import app.models.report # noqa: F401
|
|
||||||
import app.models.user # noqa: F401
|
|
||||||
import pytest
|
import pytest
|
||||||
from app.core.database import Base, get_db
|
|
||||||
from app.services.report_store import ReportStore
|
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
from sqlalchemy import create_engine
|
from sqlalchemy import create_engine
|
||||||
from sqlalchemy.orm import sessionmaker
|
from sqlalchemy.orm import sessionmaker
|
||||||
|
|
||||||
|
import app.models.domain # noqa: F401 # pylint: disable=unused-import
|
||||||
|
import app.models.report # noqa: F401 # pylint: disable=unused-import
|
||||||
|
import app.models.user # noqa: F401 # pylint: disable=unused-import
|
||||||
|
from app.core.database import Base, get_db
|
||||||
|
from app.main import create_app
|
||||||
|
from app.services.report_store import ReportStore
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture()
|
@pytest.fixture()
|
||||||
def test_app() -> FastAPI:
|
def test_app() -> FastAPI:
|
||||||
"""Create a fresh FastAPI application instance for testing."""
|
"""Create a fresh FastAPI application instance for testing."""
|
||||||
from app.main import create_app
|
|
||||||
|
|
||||||
application = create_app()
|
application = create_app()
|
||||||
return application
|
return application
|
||||||
|
|
||||||
@@ -36,7 +36,7 @@ def db_session():
|
|||||||
|
|
||||||
|
|
||||||
@pytest.fixture()
|
@pytest.fixture()
|
||||||
def client(test_app: FastAPI, db_session):
|
def client(test_app: FastAPI, db_session): # pylint: disable=redefined-outer-name
|
||||||
"""Create a TestClient with a DB override for the test app."""
|
"""Create a TestClient with a DB override for the test app."""
|
||||||
|
|
||||||
def override_get_db():
|
def override_get_db():
|
||||||
|
|||||||
@@ -0,0 +1,49 @@
|
|||||||
|
"""Shared test data for DMARC report tests."""
|
||||||
|
|
||||||
|
SAMPLE_XML = """\
|
||||||
|
<?xml version="1.0" encoding="UTF-8" ?>
|
||||||
|
<feedback>
|
||||||
|
<report_metadata>
|
||||||
|
<org_name>google.com</org_name>
|
||||||
|
<email>noreply-dmarc-support@google.com</email>
|
||||||
|
<report_id>123456789</report_id>
|
||||||
|
<date_range>
|
||||||
|
<begin>1597449600</begin>
|
||||||
|
<end>1597535999</end>
|
||||||
|
</date_range>
|
||||||
|
</report_metadata>
|
||||||
|
<policy_published>
|
||||||
|
<domain>example.com</domain>
|
||||||
|
<adkim>r</adkim>
|
||||||
|
<aspf>r</aspf>
|
||||||
|
<p>none</p>
|
||||||
|
<sp>none</sp>
|
||||||
|
<pct>100</pct>
|
||||||
|
</policy_published>
|
||||||
|
<record>
|
||||||
|
<row>
|
||||||
|
<source_ip>203.0.113.1</source_ip>
|
||||||
|
<count>2</count>
|
||||||
|
<policy_evaluated>
|
||||||
|
<disposition>none</disposition>
|
||||||
|
<dkim>pass</dkim>
|
||||||
|
<spf>fail</spf>
|
||||||
|
</policy_evaluated>
|
||||||
|
</row>
|
||||||
|
<identifiers>
|
||||||
|
<header_from>example.com</header_from>
|
||||||
|
</identifiers>
|
||||||
|
<auth_results>
|
||||||
|
<dkim>
|
||||||
|
<domain>example.com</domain>
|
||||||
|
<result>pass</result>
|
||||||
|
<selector>default</selector>
|
||||||
|
</dkim>
|
||||||
|
<spf>
|
||||||
|
<domain>example.com</domain>
|
||||||
|
<result>fail</result>
|
||||||
|
</spf>
|
||||||
|
</auth_results>
|
||||||
|
</record>
|
||||||
|
</feedback>
|
||||||
|
"""
|
||||||
@@ -2,55 +2,9 @@ import io
|
|||||||
import zipfile
|
import zipfile
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from app.services.dmarc_parser import DMARCParser
|
|
||||||
|
|
||||||
SAMPLE_XML = """\
|
from app.services.dmarc_parser import DMARCParser
|
||||||
<?xml version="1.0" encoding="UTF-8" ?>
|
from app.tests.test_data import SAMPLE_XML
|
||||||
<feedback>
|
|
||||||
<report_metadata>
|
|
||||||
<org_name>google.com</org_name>
|
|
||||||
<email>noreply-dmarc-support@google.com</email>
|
|
||||||
<report_id>123456789</report_id>
|
|
||||||
<date_range>
|
|
||||||
<begin>1597449600</begin>
|
|
||||||
<end>1597535999</end>
|
|
||||||
</date_range>
|
|
||||||
</report_metadata>
|
|
||||||
<policy_published>
|
|
||||||
<domain>example.com</domain>
|
|
||||||
<adkim>r</adkim>
|
|
||||||
<aspf>r</aspf>
|
|
||||||
<p>none</p>
|
|
||||||
<sp>none</sp>
|
|
||||||
<pct>100</pct>
|
|
||||||
</policy_published>
|
|
||||||
<record>
|
|
||||||
<row>
|
|
||||||
<source_ip>203.0.113.1</source_ip>
|
|
||||||
<count>2</count>
|
|
||||||
<policy_evaluated>
|
|
||||||
<disposition>none</disposition>
|
|
||||||
<dkim>pass</dkim>
|
|
||||||
<spf>fail</spf>
|
|
||||||
</policy_evaluated>
|
|
||||||
</row>
|
|
||||||
<identifiers>
|
|
||||||
<header_from>example.com</header_from>
|
|
||||||
</identifiers>
|
|
||||||
<auth_results>
|
|
||||||
<dkim>
|
|
||||||
<domain>example.com</domain>
|
|
||||||
<result>pass</result>
|
|
||||||
<selector>default</selector>
|
|
||||||
</dkim>
|
|
||||||
<spf>
|
|
||||||
<domain>example.com</domain>
|
|
||||||
<result>fail</result>
|
|
||||||
</spf>
|
|
||||||
</auth_results>
|
|
||||||
</record>
|
|
||||||
</feedback>
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
class TestDMARCParser:
|
class TestDMARCParser:
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from app.models.domain import Domain
|
from app.models.domain import Domain
|
||||||
from app.models.report import DMARCReport, ReportRecord
|
from app.models.report import DMARCReport, ReportRecord
|
||||||
from sqlalchemy.orm import Session
|
|
||||||
|
|
||||||
|
|
||||||
class TestDomainModel:
|
class TestDomainModel:
|
||||||
|
|||||||
@@ -3,53 +3,7 @@ import zipfile
|
|||||||
|
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
SAMPLE_XML = """\
|
from app.tests.test_data import SAMPLE_XML
|
||||||
<?xml version="1.0" encoding="UTF-8" ?>
|
|
||||||
<feedback>
|
|
||||||
<report_metadata>
|
|
||||||
<org_name>google.com</org_name>
|
|
||||||
<email>noreply-dmarc-support@google.com</email>
|
|
||||||
<report_id>123456789</report_id>
|
|
||||||
<date_range>
|
|
||||||
<begin>1597449600</begin>
|
|
||||||
<end>1597535999</end>
|
|
||||||
</date_range>
|
|
||||||
</report_metadata>
|
|
||||||
<policy_published>
|
|
||||||
<domain>example.com</domain>
|
|
||||||
<adkim>r</adkim>
|
|
||||||
<aspf>r</aspf>
|
|
||||||
<p>none</p>
|
|
||||||
<sp>none</sp>
|
|
||||||
<pct>100</pct>
|
|
||||||
</policy_published>
|
|
||||||
<record>
|
|
||||||
<row>
|
|
||||||
<source_ip>203.0.113.1</source_ip>
|
|
||||||
<count>2</count>
|
|
||||||
<policy_evaluated>
|
|
||||||
<disposition>none</disposition>
|
|
||||||
<dkim>pass</dkim>
|
|
||||||
<spf>fail</spf>
|
|
||||||
</policy_evaluated>
|
|
||||||
</row>
|
|
||||||
<identifiers>
|
|
||||||
<header_from>example.com</header_from>
|
|
||||||
</identifiers>
|
|
||||||
<auth_results>
|
|
||||||
<dkim>
|
|
||||||
<domain>example.com</domain>
|
|
||||||
<result>pass</result>
|
|
||||||
<selector>default</selector>
|
|
||||||
</dkim>
|
|
||||||
<spf>
|
|
||||||
<domain>example.com</domain>
|
|
||||||
<result>fail</result>
|
|
||||||
</spf>
|
|
||||||
</auth_results>
|
|
||||||
</record>
|
|
||||||
</feedback>
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
def _make_zip(xml_content: str) -> bytes:
|
def _make_zip(xml_content: str) -> bytes:
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ Covers API key management, domain validation, file upload limits, and XML parsin
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
import app.services.dmarc_parser as parser_module
|
||||||
from app.core.security import add_api_key, generate_api_key, verify_api_key
|
from app.core.security import add_api_key, generate_api_key, verify_api_key
|
||||||
from app.services.dmarc_parser import DMARCParser
|
from app.services.dmarc_parser import DMARCParser
|
||||||
from app.utils.domain_validator import validate_domain, validate_domain_config
|
from app.utils.domain_validator import validate_domain, validate_domain_config
|
||||||
@@ -123,8 +125,6 @@ class TestXMLParsingSecurity:
|
|||||||
"""Test XML parsing security (defusedxml, XXE protection)."""
|
"""Test XML parsing security (defusedxml, XXE protection)."""
|
||||||
|
|
||||||
def test_defusedxml_is_used(self):
|
def test_defusedxml_is_used(self):
|
||||||
import app.services.dmarc_parser as parser_module
|
|
||||||
|
|
||||||
assert hasattr(parser_module, "ET")
|
assert hasattr(parser_module, "ET")
|
||||||
module_info = str(getattr(parser_module.ET, "__name__", "")) + str(
|
module_info = str(getattr(parser_module.ET, "__name__", "")) + str(
|
||||||
getattr(parser_module.ET, "__module__", "")
|
getattr(parser_module.ET, "__module__", "")
|
||||||
@@ -149,5 +149,5 @@ class TestXMLParsingSecurity:
|
|||||||
result = DMARCParser.parse_file(xxe_payload, "test.xml")
|
result = DMARCParser.parse_file(xxe_payload, "test.xml")
|
||||||
org_name = result.get("org_name", "")
|
org_name = result.get("org_name", "")
|
||||||
assert "root:" not in org_name and "/bin" not in org_name
|
assert "root:" not in org_name and "/bin" not in org_name
|
||||||
except Exception:
|
except Exception: # pylint: disable=broad-exception-caught
|
||||||
pass # Expected – defusedxml blocks DTD processing
|
pass # Expected – defusedxml blocks DTD processing
|
||||||
|
|||||||
@@ -56,7 +56,7 @@ def _validate_domain_labels(
|
|||||||
return True, None, None
|
return True, None, None
|
||||||
|
|
||||||
|
|
||||||
def validate_domain(
|
def validate_domain( # pylint: disable=too-many-return-statements
|
||||||
domain_name: str, check_dns: bool = True
|
domain_name: str, check_dns: bool = True
|
||||||
) -> Tuple[bool, Optional[str], Optional[str]]:
|
) -> Tuple[bool, Optional[str], Optional[str]]:
|
||||||
"""
|
"""
|
||||||
@@ -103,7 +103,6 @@ def validate_domain(
|
|||||||
if check_dns:
|
if check_dns:
|
||||||
try:
|
try:
|
||||||
socket.gethostbyname(domain_name)
|
socket.gethostbyname(domain_name)
|
||||||
return True, None, None
|
|
||||||
except socket.gaierror:
|
except socket.gaierror:
|
||||||
# We could consider this valid if we don't require DNS resolution,
|
# We could consider this valid if we don't require DNS resolution,
|
||||||
# but since DMARC requires valid DNS, we'll mark it as warning
|
# but since DMARC requires valid DNS, we'll mark it as warning
|
||||||
@@ -133,7 +132,7 @@ def validate_domain_config(domain_data: Dict) -> Dict[str, Union[bool, str]]:
|
|||||||
# Validate domain name
|
# Validate domain name
|
||||||
if "name" in domain_data:
|
if "name" in domain_data:
|
||||||
# Don't check DNS for domain config validation
|
# Don't check DNS for domain config validation
|
||||||
is_valid, error_msg, error_code = validate_domain(domain_data["name"], check_dns=False)
|
is_valid, error_msg, _ = validate_domain(domain_data["name"], check_dns=False)
|
||||||
if not is_valid:
|
if not is_valid:
|
||||||
errors["name"] = error_msg
|
errors["name"] = error_msg
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -63,10 +63,10 @@ class StatsSummarizer:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
# Read cache file
|
# Read cache file
|
||||||
with open(cache_file, "r") as f:
|
with open(cache_file, "r", encoding="utf-8") as f:
|
||||||
return json.load(f)
|
return json.load(f)
|
||||||
except Exception as e:
|
except Exception as e: # pylint: disable=broad-exception-caught
|
||||||
logger.warning(f"Error reading cache file {cache_file}: {str(e)}")
|
logger.warning("Error reading cache file %s: %s", cache_file, str(e))
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def save_summary(self, stats: Dict[str, Any], domain_id: Optional[str] = None) -> bool:
|
def save_summary(self, stats: Dict[str, Any], domain_id: Optional[str] = None) -> bool:
|
||||||
@@ -87,12 +87,12 @@ class StatsSummarizer:
|
|||||||
stats["cached_at"] = datetime.now().isoformat()
|
stats["cached_at"] = datetime.now().isoformat()
|
||||||
|
|
||||||
# Write to cache file
|
# Write to cache file
|
||||||
with open(cache_file, "w") as f:
|
with open(cache_file, "w", encoding="utf-8") as f:
|
||||||
json.dump(stats, f)
|
json.dump(stats, f)
|
||||||
|
|
||||||
return True
|
return True
|
||||||
except Exception as e:
|
except Exception as e: # pylint: disable=broad-exception-caught
|
||||||
logger.error(f"Error writing cache file {cache_file}: {str(e)}")
|
logger.error("Error writing cache file %s: %s", cache_file, str(e))
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def invalidate_cache(self, domain_id: Optional[str] = None) -> None:
|
def invalidate_cache(self, domain_id: Optional[str] = None) -> None:
|
||||||
@@ -126,12 +126,11 @@ class StatsSummarizer:
|
|||||||
"""
|
"""
|
||||||
if domain_id is None:
|
if domain_id is None:
|
||||||
return os.path.join(self.cache_dir, "global_summary.json")
|
return os.path.join(self.cache_dir, "global_summary.json")
|
||||||
else:
|
# Sanitize domain_id to use as filename
|
||||||
# Sanitize domain_id to use as filename
|
safe_domain = domain_id.replace(".", "_").replace("/", "_")
|
||||||
safe_domain = domain_id.replace(".", "_").replace("/", "_")
|
return os.path.join(self.cache_dir, f"domain_{safe_domain}.json")
|
||||||
return os.path.join(self.cache_dir, f"domain_{safe_domain}.json")
|
|
||||||
|
|
||||||
def calculate_summary_statistics(self, db, domain_id: Optional[str] = None) -> Dict[str, Any]:
|
def calculate_summary_statistics(self, _db, domain_id: Optional[str] = None) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
Calculate summary statistics from the database
|
Calculate summary statistics from the database
|
||||||
|
|
||||||
|
|||||||
@@ -49,6 +49,24 @@ include_trailing_comma = true
|
|||||||
force_grid_wrap = 0
|
force_grid_wrap = 0
|
||||||
use_parentheses = true
|
use_parentheses = true
|
||||||
ensure_newline_before_comments = true
|
ensure_newline_before_comments = true
|
||||||
|
known_first_party = ["app"]
|
||||||
|
skip = ["venv", ".venv", "migrations"]
|
||||||
|
|
||||||
|
[tool.pylint.main]
|
||||||
|
# Run from repo root: pylint backend/app
|
||||||
|
max-line-length = 100
|
||||||
|
|
||||||
|
[tool.pylint."messages control"]
|
||||||
|
disable = [
|
||||||
|
"C0111", # missing-docstring
|
||||||
|
"C0103", # invalid-name (e.g. SessionLocal, TestingSessionLocal)
|
||||||
|
"R0903", # too-few-public-methods
|
||||||
|
"R0913", # too-many-arguments
|
||||||
|
"W0212", # protected-access
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.pylint.basic]
|
||||||
|
good-names = ["i", "j", "k", "ex", "_", "id", "db"]
|
||||||
|
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
testpaths = ["backend/app/tests"]
|
testpaths = ["backend/app/tests"]
|
||||||
|
|||||||
@@ -31,59 +31,3 @@ skip_covered = False
|
|||||||
|
|
||||||
[coverage:html]
|
[coverage:html]
|
||||||
directory = htmlcov
|
directory = htmlcov
|
||||||
|
|
||||||
[flake8]
|
|
||||||
max-line-length = 100
|
|
||||||
exclude =
|
|
||||||
.git,
|
|
||||||
__pycache__,
|
|
||||||
.venv,
|
|
||||||
venv,
|
|
||||||
build,
|
|
||||||
dist,
|
|
||||||
*.egg-info,
|
|
||||||
migrations
|
|
||||||
extend-ignore = E203, W503, E501
|
|
||||||
per-file-ignores =
|
|
||||||
__init__.py:F401
|
|
||||||
max-complexity = 10
|
|
||||||
|
|
||||||
[mypy]
|
|
||||||
python_version = 3.10
|
|
||||||
warn_return_any = True
|
|
||||||
warn_unused_configs = True
|
|
||||||
disallow_untyped_defs = False
|
|
||||||
disallow_incomplete_defs = False
|
|
||||||
check_untyped_defs = True
|
|
||||||
disallow_untyped_calls = False
|
|
||||||
disallow_any_generics = False
|
|
||||||
ignore_missing_imports = True
|
|
||||||
no_implicit_optional = True
|
|
||||||
warn_redundant_casts = True
|
|
||||||
warn_unused_ignores = True
|
|
||||||
warn_no_return = True
|
|
||||||
strict_optional = True
|
|
||||||
|
|
||||||
[isort]
|
|
||||||
profile = black
|
|
||||||
line_length = 100
|
|
||||||
multi_line_output = 3
|
|
||||||
include_trailing_comma = True
|
|
||||||
force_grid_wrap = 0
|
|
||||||
use_parentheses = True
|
|
||||||
ensure_newline_before_comments = True
|
|
||||||
skip = venv,.venv,migrations
|
|
||||||
|
|
||||||
[pylint]
|
|
||||||
max-line-length = 100
|
|
||||||
disable =
|
|
||||||
C0111, # missing-docstring
|
|
||||||
C0103, # invalid-name
|
|
||||||
R0903, # too-few-public-methods
|
|
||||||
R0913, # too-many-arguments
|
|
||||||
W0212, # protected-access
|
|
||||||
good-names = i,j,k,ex,_,id,db
|
|
||||||
|
|
||||||
[bandit]
|
|
||||||
exclude_dirs = /tests/,/venv/,.venv/
|
|
||||||
skips = B101,B601
|
|
||||||
|
|||||||
Reference in New Issue
Block a user