Skip to content

Commit d8d600c

Browse files
sunlishuo25XiaJunjie2020
authored andcommitted
fix(db): avoid duplicate Doris and StarRocks connection parameters
1 parent e02f717 commit d8d600c

2 files changed

Lines changed: 136 additions & 1 deletion

File tree

‎backend/apps/db/db.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -216,7 +216,7 @@ def get_driver_connection(ds: CoreDatasource | AssistantOutDsSchema, db_config:
216216
if not use_pool:
217217
conn = pymysql.connect(user=conf.username, passwd=conf.password, host=conf.host,
218218
port=conf.port, db=conf.database, connect_timeout=conf.timeout,
219-
read_timeout=conf.timeout, **conn_conf,
219+
read_timeout=conf.timeout,
220220
**args)
221221
else:
222222
conn = PooledDB(
Lines changed: 135 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,135 @@
1+
"""Regression coverage for Doris/StarRocks connection parameter forwarding."""
2+
3+
import ast
4+
import json
5+
from pathlib import Path
6+
from types import SimpleNamespace
7+
from unittest.mock import Mock
8+
9+
import pytest
10+
11+
from apps.datasource.models.datasource import DatasourceConf
12+
13+
SOURCE = Path(__file__).parents[1] / "apps" / "db" / "db.py"
14+
15+
16+
@pytest.fixture
17+
def driver_connection():
18+
"""Load production functions without importing unrelated database drivers."""
19+
nodes = ast.parse(SOURCE.read_text(encoding="utf-8")).body
20+
selected = [
21+
node
22+
for node in nodes
23+
if isinstance(node, ast.FunctionDef)
24+
and node.name in {"get_extra_config", "get_driver_connection"}
25+
]
26+
connect = Mock()
27+
pool = Mock()
28+
driver = SimpleNamespace(connect=connect)
29+
namespace = {
30+
"CoreDatasource": SimpleNamespace,
31+
"AssistantOutDsSchema": SimpleNamespace,
32+
"DatasourceConf": DatasourceConf,
33+
"json": json,
34+
"aes_decrypt": lambda value: value,
35+
"equals_ignore_case": lambda value, *options: (
36+
value.lower() in [option.lower() for option in options]
37+
),
38+
"pymysql": driver,
39+
"PooledDB": pool,
40+
}
41+
exec(
42+
compile(ast.Module(body=selected, type_ignores=[]), str(SOURCE), "exec"),
43+
namespace,
44+
)
45+
return namespace["get_driver_connection"], connect, pool, driver
46+
47+
48+
def _datasource(datasource_type, ssl, extra_jdbc):
49+
return SimpleNamespace(
50+
type=datasource_type,
51+
configuration=json.dumps(
52+
{
53+
"username": "test-user",
54+
"password": "test-password",
55+
"host": "test-host",
56+
"port": 9030,
57+
"database": "test-db",
58+
"timeout": 30,
59+
"ssl": ssl,
60+
"poolSize": 8,
61+
"extraJdbc": extra_jdbc,
62+
}
63+
),
64+
)
65+
66+
67+
def _connection_options(ssl):
68+
options = {
69+
"user": "test-user",
70+
"passwd": "test-password",
71+
"host": "test-host",
72+
"port": 9030,
73+
"db": "test-db",
74+
"connect_timeout": 30,
75+
"read_timeout": 30,
76+
}
77+
if ssl:
78+
options["ssl"] = {"ssl_mode": "REQUIRE"}
79+
return options
80+
81+
82+
@pytest.mark.parametrize("datasource_type", ["doris", "starrocks"])
83+
@pytest.mark.parametrize("ssl", [False, True])
84+
@pytest.mark.parametrize(
85+
("extra_jdbc", "db_config", "expected_extra"),
86+
[
87+
("", {}, {}),
88+
("charset=utf8mb4", {}, {"charset": "utf8mb4"}),
89+
(
90+
"charset=latin1",
91+
{"charset": "utf8mb4", "autocommit": True},
92+
{"charset": "utf8mb4", "autocommit": True},
93+
),
94+
],
95+
ids=["defaults", "extra-parameter", "config-override"],
96+
)
97+
def test_direct_connection_forwards_merged_parameters(
98+
driver_connection, datasource_type, ssl, extra_jdbc, db_config, expected_extra
99+
):
100+
get_connection, connect, pool, _ = driver_connection
101+
datasource = _datasource(datasource_type, ssl, extra_jdbc)
102+
103+
connection = get_connection(datasource, db_config)
104+
105+
assert connection is connect.return_value
106+
connect.assert_called_once_with(**(_connection_options(ssl) | expected_extra))
107+
pool.assert_not_called()
108+
109+
110+
@pytest.mark.parametrize("datasource_type", ["doris", "starrocks"])
111+
@pytest.mark.parametrize("ssl", [False, True])
112+
def test_pooled_connection_preserves_parameters(
113+
driver_connection, datasource_type, ssl
114+
):
115+
get_connection, connect, pool, driver = driver_connection
116+
datasource = _datasource(datasource_type, ssl, "charset=latin1")
117+
118+
connection = get_connection(
119+
datasource, {"charset": "utf8mb4", "autocommit": True}, use_pool=True
120+
)
121+
122+
assert connection is pool.return_value
123+
pool.assert_called_once_with(
124+
creator=driver,
125+
**_connection_options(ssl),
126+
charset="utf8mb4",
127+
autocommit=True,
128+
maxconnections=8,
129+
mincached=5,
130+
maxcached=10,
131+
blocking=True,
132+
maxusage=100,
133+
ping=1,
134+
)
135+
connect.assert_not_called()

0 commit comments

Comments
 (0)