312 lines
12 KiB
Python
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()
|