from datetime import datetime, timedelta, timezone import unittest from unittest.mock import AsyncMock, patch from fastapi import HTTPException from jose import jwt from app import db from app.auth import ( create_access_token, get_current_user, require_admin, require_data_viewer, ) from app.config import settings from app.data_platform import interface_service from app.rate_limit import FixedWindowLimiter from app.security_baseline import enforce_security_baseline, security_readiness def _production_settings(**updates): values = { "app_environment": "production", "security_strict_mode": True, "auth_secret": "a" * 48, "interface_api_secret": "b" * 48, "auth_default_password": "S3cure-Admin-Credential-2026!", "database_url": "postgresql://app:" + "p" * 32 + "@127.0.0.1:5432/kg", "data_mysql_url": "mysql://app:" + "m" * 32 + "@127.0.0.1:3306/control", "mysql_host_bind": "127.0.0.1", "data_mysql_storage_backend": "dedicated-block", "data_mysql_storage_id": "prod-mysql-storage-01", "data_mysql_storage_mount": "/mnt/sdr", "data_mysql_data_dir": "/mnt/sdr/nianxx/mysql", "data_backup_enabled": True, "data_backup_encryption_required": True, "data_backup_retention_days": 30, "data_backup_root": "/mnt/backup/nianxx/mysql", "data_mysql_direct_access_enabled": True, "data_mysql_ssh_tunnel_required": True, "data_mysql_ssh_host": "db-admin.example.cn", "cors_allowed_origins": "https://data.example.cn", "trusted_hosts": "data.example.cn", "auth_algorithm": "HS256", "auth_token_expire_minutes": 60, "ingest_api_keys": "k" * 40, } values.update(updates) return settings.model_copy(update=values) class SecurityBaselineTests(unittest.TestCase): def test_safe_production_settings_have_no_critical_findings(self) -> None: config = _production_settings() report = security_readiness(config) self.assertEqual(report["status"], "ready") self.assertEqual(report["critical_count"], 0) enforce_security_baseline(config) def test_strict_mode_rejects_public_mysql_and_placeholder_secrets(self) -> None: config = _production_settings( mysql_host_bind="0.0.0.0", auth_secret="change-me", ) report = security_readiness(config) codes = {item["code"] for item in report["findings"]} self.assertIn("MYSQL_PUBLIC_BIND", codes) self.assertIn("AUTH_SECRET", codes) with self.assertRaisesRegex(RuntimeError, "生产安全基线检查失败"): enforce_security_baseline(config) def test_production_cors_requires_https(self) -> None: report = security_readiness( _production_settings(cors_allowed_origins="http://data.example.cn") ) self.assertIn( "CORS_WITHOUT_HTTPS", {item["code"] for item in report["findings"]}, ) def test_production_rejects_project_local_or_unidentified_mysql_storage(self) -> None: unsafe = _production_settings( data_mysql_storage_backend="docker-volume", data_mysql_storage_id="", data_mysql_data_dir="/app/mysql-data", ) codes = {item["code"] for item in security_readiness(unsafe)["findings"]} self.assertIn("MYSQL_STORAGE_BACKEND", codes) unsafe = _production_settings( data_mysql_storage_id="short", data_mysql_storage_mount="/app/mysql-storage", data_mysql_data_dir="/app/mysql-data", ) codes = {item["code"] for item in security_readiness(unsafe)["findings"]} self.assertIn("MYSQL_STORAGE_ID", codes) self.assertIn("MYSQL_STORAGE_MOUNT", codes) self.assertIn("MYSQL_DATA_DIR", codes) outside = _production_settings(data_mysql_data_dir="/srv/mysql") self.assertIn( "MYSQL_DATA_OUTSIDE_MOUNT", {item["code"] for item in security_readiness(outside)["findings"]}, ) def test_production_requires_encrypted_external_nonoverlapping_backups(self) -> None: unsafe = _production_settings( data_backup_enabled=False, data_backup_encryption_required=False, data_backup_root="/mnt/sdr/nianxx/mysql/backups", ) codes = {item["code"] for item in security_readiness(unsafe)["findings"]} self.assertIn("BACKUP_DISABLED", codes) unsafe = _production_settings( data_backup_encryption_required=False, data_backup_retention_days=1, data_backup_root="/mnt/sdr/nianxx/mysql/backups", ) codes = {item["code"] for item in security_readiness(unsafe)["findings"]} self.assertIn("BACKUP_ENCRYPTION_DISABLED", codes) self.assertIn("BACKUP_RETENTION", codes) self.assertIn("BACKUP_DATA_OVERLAP", codes) def test_managed_dbeaver_access_requires_an_isolated_provisioner(self) -> None: safe = _production_settings( data_mysql_managed_access_enabled=True, data_mysql_ssh_authorized_keys_file="/srv/dbeaver/authorized_keys", data_mysql_ssh_host_public_key_file="/srv/dbeaver/host.pub", data_mysql_provisioner_user="access_broker", data_mysql_provisioner_password="z" * 32, ) self.assertEqual(security_readiness(safe)["critical_count"], 0) unsafe = safe.model_copy(update={ "data_mysql_provisioner_user": "app", "data_mysql_provisioner_password": "change-me", }) codes = {item["code"] for item in security_readiness(unsafe)["findings"]} self.assertIn("DBEAVER_PROVISIONER_USER", codes) self.assertIn("DBEAVER_PROVISIONER_PASSWORD", codes) class AuthenticationSecurityTests(unittest.IsolatedAsyncioTestCase): async def test_token_uses_issuer_audience_and_live_database_roles(self) -> None: token = create_access_token({"sub": "admin@example.com", "roles": ["collector"]}) claims = jwt.decode( token, settings.auth_secret, algorithms=[settings.auth_algorithm], issuer=settings.auth_issuer, audience=settings.auth_audience, ) self.assertTrue(claims["jti"]) self.assertEqual(claims["iss"], settings.auth_issuer) with patch.object( db, "get_user_auth", AsyncMock( return_value={ "username": "admin@example.com", "full_name": "Admin", "status": "active", "roles": ["admin"], } ), ): user = await get_current_user(token) self.assertEqual(user["roles"], ["admin"]) async def test_disabled_account_invalidates_existing_token(self) -> None: token = create_access_token({"sub": "disabled@example.com", "roles": ["admin"]}) with patch.object( db, "get_user_auth", AsyncMock(return_value={"username": "disabled@example.com", "status": "disabled"}), ): with self.assertRaises(HTTPException) as raised: await get_current_user(token) self.assertEqual(raised.exception.status_code, 401) async def test_rate_limiter_blocks_at_configured_threshold(self) -> None: limiter = FixedWindowLimiter(limit=2, window_seconds=60, block_seconds=30) self.assertEqual(await limiter.record("user"), 0) self.assertGreater(await limiter.record("user"), 0) self.assertGreater(await limiter.check("user"), 0) await limiter.reset("user") self.assertEqual(await limiter.check("user"), 0) def test_admin_dependency_rejects_non_admin(self) -> None: with self.assertRaises(HTTPException) as raised: require_admin({"username": "collector", "roles": ["collector"]}) self.assertEqual(raised.exception.status_code, 403) def test_data_center_view_permission_rejects_unassigned_accounts(self) -> None: with self.assertRaises(HTTPException) as raised: require_data_viewer({"username": "unassigned", "roles": []}) self.assertEqual(raised.exception.status_code, 403) viewer = require_data_viewer( {"username": "viewer", "roles": ["operator"]} ) self.assertEqual(viewer["username"], "viewer") def test_retired_collector_accounts_are_read_only_compatible(self) -> None: legacy = require_data_viewer( {"username": "legacy", "roles": ["collector"]} ) self.assertEqual(legacy["username"], "legacy") with self.assertRaises(HTTPException): require_admin(legacy) class InterfaceCredentialSecurityTests(unittest.TestCase): def test_api_credentials_default_to_short_lived(self) -> None: before = datetime.now(timezone.utc).replace(tzinfo=None) with patch.object(interface_service.settings, "interface_api_default_expiry_days", 30): expiry = interface_service._parse_expiry(None) self.assertGreater(expiry, before + timedelta(days=29)) self.assertLess(expiry, before + timedelta(days=31)) def test_api_credentials_cannot_exceed_maximum_lifetime(self) -> None: too_far = datetime.now(timezone.utc) + timedelta(days=91) with patch.object(interface_service.settings, "interface_api_max_expiry_days", 90): with self.assertRaisesRegex(ValueError, "最长只能签发"): interface_service._parse_expiry(too_far) if __name__ == "__main__": unittest.main()