230 lines
9.5 KiB
Python
230 lines
9.5 KiB
Python
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()
|