"""Fail-closed production security checks without exposing secret values.""" from __future__ import annotations from typing import Any from urllib.parse import unquote, urlsplit from app.config import Settings, settings _PLACEHOLDER_MARKERS = ( "change-me", "change-this", "password", "dev-key", "example", ) _LOOPBACK_HOSTS = {"127.0.0.1", "localhost", "::1"} def _finding(code: str, severity: str, message: str) -> dict[str, str]: return {"code": code, "severity": severity, "message": message} def _looks_insecure_secret(value: str, *, minimum: int = 32) -> bool: normalized = value.strip().lower() return ( len(value.strip()) < minimum or any(marker in normalized for marker in _PLACEHOLDER_MARKERS) ) def _url_password(value: str) -> str: try: parsed = urlsplit(value) return unquote(parsed.password or "") except ValueError: return "" def _csv_values(value: str) -> list[str]: return [item.strip() for item in value.split(",") if item.strip()] def security_readiness(config: Settings = settings) -> dict[str, Any]: """Return a sanitized security report suitable for an admin endpoint.""" production = config.app_environment.strip().lower() in { "production", "prod", "server", } findings: list[dict[str, str]] = [] if _looks_insecure_secret(config.auth_secret): findings.append( _finding("AUTH_SECRET", "critical", "JWT 签名密钥未替换或强度不足") ) if _looks_insecure_secret(config.interface_api_secret): findings.append( _finding( "INTERFACE_API_SECRET", "critical", "接口密钥散列密钥未独立配置或强度不足", ) ) if ( config.interface_api_secret.strip() and config.interface_api_secret == config.auth_secret ): findings.append( _finding( "SECRET_REUSE", "critical", "JWT 与接口密钥不能共用同一签名密钥", ) ) if _looks_insecure_secret(config.auth_default_password, minimum=12): findings.append( _finding( "DEFAULT_ADMIN_PASSWORD", "critical", "默认管理员密码未替换或少于 12 位", ) ) for code, url in ( ("POSTGRES_PASSWORD", config.database_url), ("MYSQL_PASSWORD", config.data_mysql_url), ): if _looks_insecure_secret(_url_password(url), minimum=16): findings.append( _finding(code, "critical", f"{code} 未替换或强度不足") ) if (config.mysql_host_bind.strip() or "127.0.0.1") not in _LOOPBACK_HOSTS: findings.append( _finding("MYSQL_PUBLIC_BIND", "critical", "MySQL 不能绑定公网地址") ) if config.data_mysql_direct_access_enabled and not config.data_mysql_ssh_tunnel_required: findings.append( _finding( "DBEAVER_WITHOUT_TUNNEL", "critical", "启用 DBeaver 管理接入时必须强制 SSH 隧道", ) ) origins = _csv_values(config.cors_allowed_origins) if not origins or "*" in origins: findings.append( _finding( "CORS_WILDCARD", "critical" if production else "warning", "生产环境必须明确配置允许访问后台的 HTTPS 来源", ) ) if production and any( not origin.startswith("https://") and not origin.startswith("http://localhost") and not origin.startswith("http://127.0.0.1") for origin in origins ): findings.append( _finding( "CORS_WITHOUT_HTTPS", "critical", "生产后台来源必须使用 HTTPS", ) ) trusted_hosts = _csv_values(config.trusted_hosts) if not trusted_hosts or "*" in trusted_hosts: findings.append( _finding( "TRUSTED_HOSTS_WILDCARD", "critical" if production else "warning", "生产环境必须明确配置可信域名", ) ) if config.auth_algorithm not in {"HS256", "HS384", "HS512"}: findings.append( _finding("JWT_ALGORITHM", "critical", "JWT 算法不在允许列表中") ) if not 5 <= config.auth_token_expire_minutes <= 120: findings.append( _finding( "JWT_LIFETIME", "critical" if production else "warning", "后台登录令牌有效期必须为 5–120 分钟", ) ) if any( _looks_insecure_secret(item, minimum=24) for item in _csv_values(config.ingest_api_keys) ): findings.append( _finding( "INGEST_API_KEYS", "critical" if production else "warning", "外部问答接口仍包含开发密钥或弱密钥", ) ) if not config.data_backup_enabled: findings.append( _finding("BACKUP_DISABLED", "warning", "尚未声明已启用自动加密备份") ) if not config.data_mysql_audit_enabled: findings.append( _finding("DB_AUDIT_DISABLED", "warning", "DBeaver 数据库审计尚未启用") ) critical_count = sum(item["severity"] == "critical" for item in findings) warning_count = sum(item["severity"] == "warning" for item in findings) return { "environment": config.app_environment, "strict_mode": config.security_strict_mode, "status": "blocked" if critical_count else "ready", "critical_count": critical_count, "warning_count": warning_count, "findings": findings, } def enforce_security_baseline(config: Settings = settings) -> None: """Abort a strict deployment before it starts with unsafe settings.""" if not config.security_strict_mode: return report = security_readiness(config) critical = [item for item in report["findings"] if item["severity"] == "critical"] if critical: codes = ", ".join(item["code"] for item in critical) raise RuntimeError(f"生产安全基线检查失败:{codes}")