diff --git a/doris_mcp_server/auth/token_manager.py b/doris_mcp_server/auth/token_manager.py index 45abd09..eb988c5 100644 --- a/doris_mcp_server/auth/token_manager.py +++ b/doris_mcp_server/auth/token_manager.py @@ -68,6 +68,9 @@ class DatabaseConfig: http_read_timeout_seconds: float = 15.0 http_total_timeout_seconds: float = 30.0 http_max_response_bytes: int = 4 * 1024 * 1024 + # Optional statement executed once on every NEW pooled connection of this + # token's pool (aiomysql `init_command`). None keeps pool creation unchanged. + init_command: str | None = None @dataclass @@ -350,6 +353,7 @@ def _add_token_from_config( 'http_max_response_bytes', 4 * 1024 * 1024, ), + init_command=db_config.get('init_command'), ) # Create token info diff --git a/doris_mcp_server/utils/config.py b/doris_mcp_server/utils/config.py index 7646365..b8c96df 100644 --- a/doris_mcp_server/utils/config.py +++ b/doris_mcp_server/utils/config.py @@ -409,6 +409,11 @@ class DatabaseConfig: be_hosts: list[str] = field(default_factory=list) be_webserver_port: int = 8040 + # Optional statement executed once on every NEW pooled connection + # (aiomysql `init_command`), e.g. a session-narrowing switch or a SET of + # session variables. None (default) keeps pool creation byte-identical. + init_command: str | None = None + # Shared FE/BE HTTP safety limits. Runtime enforcement also applies hard # caps so unsafe config values cannot remove the boundary. http_connect_timeout_seconds: float = 3.0 diff --git a/doris_mcp_server/utils/db.py b/doris_mcp_server/utils/db.py index e63d4be..0797c5a 100644 --- a/doris_mcp_server/utils/db.py +++ b/doris_mcp_server/utils/db.py @@ -127,6 +127,7 @@ class DatabasePoolConfig(TypedDict): password: str database: str charset: str + init_command: NotRequired[str | None] @dataclass @@ -639,6 +640,9 @@ def __init__( "database": config.database.database, "charset": config.database.charset, } + init_command = getattr(config.database, "init_command", None) + if init_command: + self.original_db_config["init_command"] = init_command # Current active database config (may be overridden by token-bound config) # NOTE: This is kept for backward compatibility with non-token requests @@ -758,6 +762,9 @@ def _get_current_token_db_config( "database": token_db_config.database, "charset": token_db_config.charset, } + init_command = getattr(token_db_config, "init_command", None) + if init_command: + db_config["init_command"] = init_command token_hosts = getattr(token_db_config, "hosts", []) if token_hosts: db_config["hosts"] = self._ordered_hosts( @@ -843,6 +850,11 @@ def _build_doris_user_db_config(self, user: str, password: str) -> dict: # still use fully qualified names or select a DB later. "database": "information_schema", "charset": self.original_db_config["charset"], + **( + {"init_command": self.original_db_config["init_command"]} + if "init_command" in self.original_db_config + else {} + ), "maxsize": self._get_doris_user_pool_maxsize(), } @@ -1260,6 +1272,9 @@ async def get_pool_for_token( "database": token_db_config.database, "charset": token_db_config.charset, } + init_command = getattr(token_db_config, "init_command", None) + if init_command: + db_config["init_command"] = init_command token_hosts = getattr(token_db_config, "hosts", []) if token_hosts: db_config["hosts"] = self._ordered_hosts( @@ -1384,6 +1399,7 @@ async def _create_pool_with_candidates( connect_timeout=self.connect_timeout, autocommit=True, pool_recycle=self.pool_recycle, + init_command=db_config.get("init_command"), ), timeout=pool_timeout, ) diff --git a/test/utils/test_pool_init_command.py b/test/utils/test_pool_init_command.py new file mode 100644 index 0000000..162430a --- /dev/null +++ b/test/utils/test_pool_init_command.py @@ -0,0 +1,152 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""An optional ``init_command`` rides every pool the manager creates.""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from doris_mcp_server.auth.token_manager import DatabaseConfig as TokenDatabaseConfig +from doris_mcp_server.utils import db as db_module +from doris_mcp_server.utils.config import DorisConfig +from doris_mcp_server.utils.db import DorisConnectionManager + +INIT = "SET SESSION time_zone = '+00:00'" + + +class _ProbeCursor: + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, traceback): + return False + + async def execute(self, sql: str) -> None: + assert sql == "SELECT 1" + + async def fetchone(self): + return (1,) + + +class _ProbeConnection: + closed = False + + def cursor(self): + return _ProbeCursor() + + +class _ProbePool: + def __init__(self): + self.closed = False + + async def acquire(self): + return _ProbeConnection() + + def release(self, connection): + return None + + def close(self): + self.closed = True + + async def wait_closed(self): + return None + + +def _config(init_command: str | None = None) -> DorisConfig: + config = DorisConfig() + config.database.host = "fe-1" + config.database.hosts = ["fe-1"] + config.database.user = "root" + config.database.password = "secret" + config.database.connection_timeout = 1 + config.database.max_connections = 2 + config.database.init_command = init_command + return config + + +def test_global_config_carries_init_command_into_the_pool_config(): + manager = DorisConnectionManager(_config(INIT)) + assert manager.original_db_config["init_command"] == INIT + + +def test_global_config_default_omits_the_key(): + # Unset = byte-identical pool config (no key), so equality-based callers + # and serialized configs are unaffected by this feature. + manager = DorisConnectionManager(_config()) + assert "init_command" not in manager.original_db_config + + +@pytest.mark.asyncio +async def test_pool_factory_passes_init_command_to_aiomysql(monkeypatch): + manager = DorisConnectionManager(_config(INIT)) + create_pool = AsyncMock(return_value=_ProbePool()) + monkeypatch.setattr(db_module.aiomysql, "create_pool", create_pool) + + await manager._create_global_pool() + + assert create_pool.await_args.kwargs["init_command"] == INIT + + +@pytest.mark.asyncio +async def test_pool_factory_passes_none_when_unset(monkeypatch): + manager = DorisConnectionManager(_config()) + create_pool = AsyncMock(return_value=_ProbePool()) + monkeypatch.setattr(db_module.aiomysql, "create_pool", create_pool) + + await manager._create_global_pool() + + assert create_pool.await_args.kwargs["init_command"] is None + + +@pytest.mark.asyncio +async def test_explicit_pool_config_wins(monkeypatch): + manager = DorisConnectionManager(_config()) + create_pool = AsyncMock(return_value=_ProbePool()) + monkeypatch.setattr(db_module.aiomysql, "create_pool", create_pool) + + db_config = dict(manager.original_db_config) + db_config["init_command"] = "SET SESSION sql_mode = ''" + await manager._create_pool_with_candidates(db_config, require_health=False) + + assert create_pool.await_args.kwargs["init_command"] == "SET SESSION sql_mode = ''" + + +def test_token_database_config_accepts_init_command(): + token_config = TokenDatabaseConfig(host="fe-1", init_command=INIT) + assert token_config.init_command == INIT + assert TokenDatabaseConfig(host="fe-1").init_command is None + + +@pytest.mark.asyncio +async def test_token_bound_pool_carries_the_token_init_command(monkeypatch): + manager = DorisConnectionManager(_config()) + token_config = TokenDatabaseConfig( + host="fe-1", user="svc", password="x", database="db", init_command=INIT + ) + manager.token_manager = SimpleNamespace( + get_database_config_by_token=lambda token: token_config + ) + create_pool = AsyncMock(return_value=_ProbePool()) + monkeypatch.setattr(db_module.aiomysql, "create_pool", create_pool) + + pool, _ = await manager.get_pool_for_token("token-1") + + assert create_pool.await_args.kwargs["init_command"] == INIT + assert create_pool.await_args.kwargs["user"] == "svc" + assert pool is not None