From f831721dcf3bfff2316bdf895aa4222c41b9f1a2 Mon Sep 17 00:00:00 2001 From: Raghvendra Singh Date: Mon, 7 Sep 2026 18:53:34 +0530 Subject: [PATCH] feat(db): honor an optional init_command on every connection pool Hosts that front Doris with a service account often need each pooled connection to start in a prepared session state: a session-narrowing switch, a SET of session variables, a workload-group pin. aiomysql's init_command is the right hook (it runs once per new connection), but _create_pool_with_candidates builds every pool without it, so a host has to subclass DorisConnectionManager and re-implement pool creation just to add one statement, and keep that subclass in lockstep with every refactor of the factory. On 0.6.1, where pool creation lived in six methods, that subclass silently missed paths. Add an optional init_command to DatabaseConfig (global and token-bound), carry it in DatabasePoolConfig as a NotRequired key, and pass it to aiomysql.create_pool in the single pool factory. When unset the key is omitted, so pool configs stay byte-identical for equality-based callers and serialized configs. Per-user pools inherit the global value. Tests cover the global and token-bound paths, the omitted default, and an explicit pool config taking precedence. Co-Authored-By: Claude Fable 5.1 Signed-off-by: Raghvendra Singh --- doris_mcp_server/auth/token_manager.py | 4 + doris_mcp_server/utils/config.py | 5 + doris_mcp_server/utils/db.py | 16 +++ test/utils/test_pool_init_command.py | 152 +++++++++++++++++++++++++ 4 files changed, 177 insertions(+) create mode 100644 test/utils/test_pool_init_command.py 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