Skip to content

Commit 518403e

Browse files
fix(db): avoid detached datasource in SQL Server pool creator (#1330)
1 parent 6d273f1 commit 518403e

2 files changed

Lines changed: 87 additions & 1 deletion

File tree

‎backend/apps/db/db.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -162,7 +162,9 @@ def get_engine(ds: CoreDatasource, timeout: int = 0, use_pool: bool = False) ->
162162
else:
163163
engine = create_engine(get_uri(ds), connect_args={"connect_timeout": conf.timeout}, **db_config)
164164
elif equals_ignore_case(ds.type, 'sqlServer'):
165-
engine = create_engine('mssql+pymssql://', creator=lambda: get_origin_connect(ds.type, conf),
165+
# A pooled connection may be recreated after the datasource's ORM session closes.
166+
ds_type = ds.type
167+
engine = create_engine('mssql+pymssql://', creator=lambda: get_origin_connect(ds_type, conf),
166168
**db_config)
167169
elif equals_ignore_case(ds.type, 'oracle'):
168170
engine = create_engine(get_uri(ds), connect_args={"tcp_connect_timeout": conf.timeout}, **db_config)
Lines changed: 84 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,84 @@
1+
"""Regression coverage for issue 1330's SQL Server pool reconnect path."""
2+
3+
import ast
4+
import json
5+
import sqlite3
6+
import threading
7+
from collections import OrderedDict
8+
from pathlib import Path
9+
10+
import pytest
11+
from sqlalchemy import create_engine, text
12+
from sqlalchemy.orm import make_transient_to_detached, sessionmaker
13+
from sqlalchemy.pool import NullPool, QueuePool
14+
from sqlmodel import Session
15+
16+
from apps.datasource.models.datasource import CoreDatasource, DatasourceConf
17+
18+
19+
SOURCE = Path(__file__).parents[1] / "apps" / "db" / "db.py"
20+
21+
22+
def _load_pool_code(connect):
23+
"""Load the production functions without importing unrelated DB drivers."""
24+
nodes = ast.parse(SOURCE.read_text()).body
25+
selected = [node for node in nodes if
26+
(isinstance(node, ast.FunctionDef) and node.name == "get_engine") or
27+
(isinstance(node, ast.ClassDef) and node.name == "ConnectionPoolManager")]
28+
29+
def build_engine(url, creator, **options):
30+
assert url == "mssql+pymssql://"
31+
assert options["pool_recycle"] == 3600
32+
# Force a new DBAPI connection at the next checkout rather than waiting an hour.
33+
return create_engine("sqlite://", creator=creator, poolclass=QueuePool,
34+
pool_recycle=0)
35+
36+
namespace = {
37+
"CoreDatasource": CoreDatasource,
38+
"AssistantOutDsSchema": type("AssistantOutDsSchema", (), {}),
39+
"DatasourceConf": DatasourceConf,
40+
"Engine": object,
41+
"json": json,
42+
"aes_decrypt": lambda value: value,
43+
"equals_ignore_case": lambda left, right: left.lower() == right.lower(),
44+
"create_engine": build_engine,
45+
"get_origin_connect": connect,
46+
"NullPool": NullPool,
47+
"threading": threading,
48+
"OrderedDict": OrderedDict,
49+
"sessionmaker": sessionmaker,
50+
}
51+
exec(compile(ast.Module(body=selected, type_ignores=[]), str(SOURCE), "exec"), namespace)
52+
return namespace["ConnectionPoolManager"]()
53+
54+
55+
@pytest.mark.parametrize("finish", ["commit", "rollback"])
56+
def test_preview_created_pool_reconnects_after_request_session_closes(finish):
57+
connections = []
58+
59+
def connect(ds_type, conf):
60+
connections.append((ds_type, conf.host, conf.database))
61+
return sqlite3.connect(":memory:")
62+
63+
manager = _load_pool_code(connect)
64+
metadata_engine = create_engine("sqlite://")
65+
configuration = json.dumps({"host": "test-host", "database": "test-db"})
66+
datasource = CoreDatasource(id=1330, type="sqlServer", configuration=configuration)
67+
make_transient_to_detached(datasource)
68+
try:
69+
with Session(metadata_engine) as request_session:
70+
request_session.add(datasource)
71+
with manager.get_pool(datasource)() as sql_session:
72+
assert sql_session.execute(text("SELECT 1")).scalar_one() == 1
73+
getattr(request_session, finish)()
74+
75+
# A later chat/MCP request reuses the pool created by the preview request.
76+
later_datasource = CoreDatasource(id=1330, type="sqlServer", configuration=configuration)
77+
with manager.get_pool(later_datasource)() as sql_session:
78+
assert sql_session.execute(text("SELECT 1")).scalar_one() == 1
79+
assert len(connections) >= 2
80+
assert all(connection == ("sqlServer", "test-host", "test-db")
81+
for connection in connections)
82+
finally:
83+
manager.close_all()
84+
metadata_engine.dispose()

0 commit comments

Comments
 (0)