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()