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_operator, ) 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_direct_access_enabled": True, "data_mysql_ssh_tunnel_required": True, "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"]}, ) 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_rejects_collector_role(self) -> None: with self.assertRaises(HTTPException) as raised: require_data_operator({"username": "collector", "roles": ["collector"]}) self.assertEqual(raised.exception.status_code, 403) operator = require_data_operator( {"username": "operator", "roles": ["operator"]} ) self.assertEqual(operator["username"], "operator") 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()