Files
Cloud-Tour-to-Libo/tests/test_mysql_data_center.py

312 lines
12 KiB
Python

import unittest
from unittest.mock import AsyncMock, Mock, patch
import base64
from pathlib import Path
import tempfile
from app.data_platform import interface_service, mysql_db, mysql_service
from app.data_platform.mysql_service import (
_managed_index_name,
_mysql_default,
_mysql_type,
project_database_name,
validate_console_sql,
)
from app.main import app
class _FakeCursor:
def __init__(self) -> None:
self.executed: list[tuple[str, tuple[object, ...]]] = []
self.rowcount = 1
async def __aenter__(self):
return self
async def __aexit__(self, _exc_type, _exc, _traceback) -> None:
return None
async def execute(self, statement: str, params: tuple[object, ...] = ()) -> None:
self.executed.append((statement, params))
async def fetchone(self):
return None
class _FakeConnection:
def __init__(self) -> None:
self.cursor_instance = _FakeCursor()
self.commit_count = 0
def cursor(self) -> _FakeCursor:
return self.cursor_instance
async def commit(self) -> None:
self.commit_count += 1
class _FakeConnectionContext:
def __init__(self, connection: _FakeConnection) -> None:
self.connection = connection
async def __aenter__(self) -> _FakeConnection:
return self.connection
async def __aexit__(self, _exc_type, _exc, _traceback) -> None:
return None
class MySQLDataCenterContractTests(unittest.TestCase):
def test_connection_options_use_database_from_mysql_url(self) -> None:
with patch.object(
mysql_db.settings,
"data_mysql_url",
"mysql://data_user:secret@mysql:3306/platform_control",
):
options = mysql_db._connection_options(include_database=True)
self.assertEqual(options["host"], "mysql")
self.assertEqual(options["port"], 3306)
self.assertEqual(options["db"], "platform_control")
self.assertEqual(mysql_db.control_database_name(), "platform_control")
def test_physical_database_name_is_stable_and_mysql_safe(self) -> None:
self.assertEqual(project_database_name("yunyou_libo"), "yunyou_libo")
long_name = "a" * 63
physical = project_database_name(long_name)
self.assertEqual(physical, long_name)
self.assertLessEqual(len(_managed_index_name("uq", long_name)), 64)
def test_postgresql_registry_types_translate_to_mysql(self) -> None:
self.assertEqual(_mysql_type("UUID"), "CHAR(36)")
self.assertEqual(_mysql_type("JSONB"), "JSON")
self.assertEqual(_mysql_type("TIMESTAMPTZ"), "DATETIME(6)")
self.assertEqual(_mysql_type("DOUBLE PRECISION"), "DOUBLE")
self.assertEqual(_mysql_type("NUMERIC(12,2)"), "DECIMAL(12,2)")
self.assertEqual(_mysql_default("'{}'::jsonb", "JSON"), "(JSON_OBJECT())")
def test_console_accepts_mysql_discovery_queries_and_frontend_sql(self) -> None:
statement, keyword = validate_console_sql("SHOW TABLES;", "demo")
self.assertEqual((statement, keyword), ("SHOW TABLES", "show"))
statement, keyword = validate_console_sql(
'SELECT * FROM "hotel_profiles" WHERE deleted_at IS NULL LIMIT 100;',
"demo",
)
self.assertEqual(keyword, "select")
self.assertIn('"hotel_profiles"', statement)
def test_console_rejects_cross_database_or_privileged_sql(self) -> None:
for sql in (
"SHOW DATABASES",
"SELECT * FROM platform_control.api_clients",
"SELECT * FROM other_database.hotel_profiles",
"SELECT * FROM unrelated_database.hotel_profiles",
"SHOW TABLES FROM unrelated_database",
"DROP TABLE hotel_profiles",
):
with self.subTest(sql=sql), self.assertRaises(ValueError):
validate_console_sql(sql, "demo")
def test_console_is_read_only_and_blocks_high_risk_queries(self) -> None:
for sql in (
"DELETE FROM hotel_profiles WHERE id='1'",
"UPDATE hotel_profiles SET name='x' WHERE id='1'",
"SELECT SLEEP(10)",
"SELECT * FROM hotel_profiles FOR UPDATE",
"SELECT id INTO @captured FROM hotel_profiles LIMIT 1",
"WITH changed AS (SELECT 1) DELETE FROM hotel_profiles WHERE id='1'",
):
with self.subTest(sql=sql), self.assertRaises(ValueError):
validate_console_sql(sql, "demo")
def test_enabled_console_writes_still_require_scoped_update(self) -> None:
with patch.object(mysql_service.settings, "data_sql_console_write_enabled", True):
with self.assertRaisesRegex(ValueError, "WHERE"):
validate_console_sql("UPDATE hotel_profiles SET name='x'", "demo")
statement, keyword = validate_console_sql(
"UPDATE hotel_profiles SET name='x' WHERE id='1'",
"demo",
)
self.assertEqual(keyword, "update")
self.assertIn("WHERE id='1'", statement)
with self.assertRaisesRegex(ValueError, "物理删除"):
validate_console_sql(
"DELETE FROM hotel_profiles WHERE id='1'",
"demo",
)
def test_interface_center_routes_are_registered(self) -> None:
paths = {route.path for route in app.routes}
for path in (
"/v1/admin/interface-center/summary",
"/v1/admin/interface-center/clients",
"/v1/admin/interface-center/policies",
"/v1/admin/interface-center/logs",
"/v1/admin/interface-center/security-readiness",
"/v1/admin/interface-center/dbeaver-access",
"/v1/admin/interface-center/dbeaver-access/{grant_id}/revoke",
"/v1/admin/data-platform/security/audit-logs",
"/v1/openapi/data/catalog",
"/v1/openapi/data/databases/{database_id}/tables/{table_code}/records",
):
self.assertIn(path, paths)
class MySQLDatabaseProvisioningTests(unittest.IsolatedAsyncioTestCase):
async def test_system_database_name_is_rejected(self) -> None:
with patch.object(mysql_service, "_ready", AsyncMock()):
with self.assertRaisesRegex(ValueError, "系统数据库"):
await mysql_service.ensure_project_database("mysql", "mysql", "MySQL")
async def test_new_database_is_created_without_template_tables(self) -> None:
connection = _FakeConnection()
database = {
"project_id": "empty_demo",
"tenant_id": "empty_demo",
"display_name": "Empty Demo",
"database_name": "empty_demo",
"schema_name": "empty_demo",
"engine": "mysql",
"status": "ready",
}
get_connection = Mock(
side_effect=lambda *_args, **_kwargs: _FakeConnectionContext(connection),
)
with (
patch.object(mysql_service, "_ready", AsyncMock()),
patch.object(mysql_service, "get_data_conn", get_connection),
patch.object(
mysql_service,
"list_project_table_definitions",
AsyncMock(return_value=()),
),
patch.object(mysql_service, "_create_physical_table", AsyncMock()) as create_table,
patch.object(
mysql_service,
"get_project_database",
AsyncMock(return_value=database),
),
):
result = await mysql_service.ensure_project_database(
"empty_demo",
"empty_demo",
"Empty Demo",
)
statements = "\n".join(statement for statement, _params in connection.cursor_instance.executed)
self.assertEqual(result, database)
self.assertIn("CREATE DATABASE IF NOT EXISTS `empty_demo`", statements)
self.assertIn("INSERT INTO project_databases", statements)
self.assertNotIn("project_table_definitions", statements)
create_table.assert_not_awaited()
class InterfaceCenterSecuritySummaryTests(unittest.IsolatedAsyncioTestCase):
def test_parameterized_mysql_account_escapes_host_wildcard(self) -> None:
username = "dba_123456abcdef"
self.assertEqual(
interface_service._mysql_account_sql(username),
"'dba_123456abcdef'@'%'",
)
self.assertEqual(
interface_service._mysql_account_sql(username, parameterized=True),
"'dba_123456abcdef'@'%%'",
)
async def test_dbeaver_summary_separates_ssh_and_database_endpoints(self) -> None:
class SummaryCursor:
async def __aenter__(self):
return self
async def __aexit__(self, _exc_type, _exc, _traceback) -> None:
return None
async def execute(self, _statement: str) -> None:
return None
async def fetchone(self):
return {
"active_clients": 0,
"active_credentials": 0,
"active_policies": 0,
"calls_24h": 0,
"errors_24h": 0,
}
class SummaryConnection:
def cursor(self):
return SummaryCursor()
with (
patch.object(interface_service, "_ready", AsyncMock()),
patch.object(
interface_service,
"get_data_conn",
Mock(return_value=_FakeConnectionContext(SummaryConnection())),
),
patch.object(interface_service.settings, "data_mysql_ssh_host", "db.example.com"),
patch.object(interface_service.settings, "data_mysql_ssh_port", 22),
patch.object(interface_service.settings, "data_mysql_admin_host", "127.0.0.1"),
patch.object(interface_service.settings, "data_mysql_admin_port", 3307),
patch.object(interface_service.settings, "mysql_host_bind", "127.0.0.1"),
):
summary = await interface_service.interface_summary()
self.assertEqual(summary["ssh_host"], "db.example.com")
self.assertEqual(summary["ssh_port"], 22)
self.assertEqual(summary["database_host"], "127.0.0.1")
self.assertEqual(summary["database_port"], 3307)
self.assertFalse(summary["mysql_publicly_bound"])
def test_managed_ssh_key_is_restricted_to_mysql_tunnel(self) -> None:
def field(value: bytes) -> bytes:
return len(value).to_bytes(4, "big") + value
rsa_modulus = b"\x00\x80" + (b"\x00" * 383)
blob = field(b"ssh-rsa") + field(b"\x01\x00\x01") + field(rsa_modulus)
encoded = base64.b64encode(blob).decode()
public_key = f"ssh-rsa {encoded} user-supplied-comment"
with tempfile.TemporaryDirectory() as directory:
target = Path(directory) / "authorized_keys"
host_key = Path(directory) / "ssh_host_ed25519_key.pub"
host_key.write_text(public_key)
with (
patch.object(
interface_service.settings,
"data_mysql_managed_access_enabled",
True,
),
patch.object(
interface_service.settings,
"data_mysql_ssh_authorized_keys_file",
str(target),
),
patch.object(
interface_service.settings,
"data_mysql_ssh_host_public_key_file",
str(host_key),
),
):
interface_service._write_authorized_keys(
[
{
"id": "grant-1",
"mysql_username": "dba_123456abcdef",
"ssh_public_key": public_key,
}
]
)
host_fingerprint = interface_service._ssh_host_key_fingerprint()
content = target.read_text()
self.assertIn("restrict,port-forwarding", content)
self.assertIn('permitopen="127.0.0.1:3307"', content)
self.assertNotIn("user-supplied-comment", content)
self.assertIn("nianxx:grant-1:dba_123456abcdef", content)
self.assertTrue(host_fingerprint.startswith("SHA256:"))
if __name__ == "__main__":
unittest.main()