Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions doris_mcp_server/auth/token_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
5 changes: 5 additions & 0 deletions doris_mcp_server/utils/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
16 changes: 16 additions & 0 deletions doris_mcp_server/utils/db.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,7 @@ class DatabasePoolConfig(TypedDict):
password: str
database: str
charset: str
init_command: NotRequired[str | None]


@dataclass
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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(),
}

Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
)
Expand Down
152 changes: 152 additions & 0 deletions test/utils/test_pool_init_command.py
Original file line number Diff line number Diff line change
@@ -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