From fd75501280e83a25db7063f66b098b5c74573e2a Mon Sep 17 00:00:00 2001 From: xuwei-fit2cloud Date: Sun, 20 Sep 2026 11:13:35 +0800 Subject: [PATCH] fix(db): avoid detached datasource in SQL Server pool creator (#1330) --- backend/apps/db/db.py | 4 +- .../tests/test_sqlserver_pool_lifecycle.py | 84 +++++++++++++++++++ 2 files changed, 87 insertions(+), 1 deletion(-) create mode 100644 backend/tests/test_sqlserver_pool_lifecycle.py diff --git a/backend/apps/db/db.py b/backend/apps/db/db.py index b8ca4dba..ca58024e 100644 --- a/backend/apps/db/db.py +++ b/backend/apps/db/db.py @@ -162,7 +162,9 @@ def get_engine(ds: CoreDatasource, timeout: int = 0, use_pool: bool = False) -> else: engine = create_engine(get_uri(ds), connect_args={"connect_timeout": conf.timeout}, **db_config) elif equals_ignore_case(ds.type, 'sqlServer'): - engine = create_engine('mssql+pymssql://', creator=lambda: get_origin_connect(ds.type, conf), + # A pooled connection may be recreated after the datasource's ORM session closes. + ds_type = ds.type + engine = create_engine('mssql+pymssql://', creator=lambda: get_origin_connect(ds_type, conf), **db_config) elif equals_ignore_case(ds.type, 'oracle'): engine = create_engine(get_uri(ds), connect_args={"tcp_connect_timeout": conf.timeout}, **db_config) diff --git a/backend/tests/test_sqlserver_pool_lifecycle.py b/backend/tests/test_sqlserver_pool_lifecycle.py new file mode 100644 index 00000000..1048504a --- /dev/null +++ b/backend/tests/test_sqlserver_pool_lifecycle.py @@ -0,0 +1,84 @@ +"""Regression coverage for issue 1330's SQL Server pool reconnect path.""" + +import ast +import json +import sqlite3 +import threading +from collections import OrderedDict +from pathlib import Path + +import pytest +from sqlalchemy import create_engine, text +from sqlalchemy.orm import make_transient_to_detached, sessionmaker +from sqlalchemy.pool import NullPool, QueuePool +from sqlmodel import Session + +from apps.datasource.models.datasource import CoreDatasource, DatasourceConf + + +SOURCE = Path(__file__).parents[1] / "apps" / "db" / "db.py" + + +def _load_pool_code(connect): + """Load the production functions without importing unrelated DB drivers.""" + nodes = ast.parse(SOURCE.read_text()).body + selected = [node for node in nodes if + (isinstance(node, ast.FunctionDef) and node.name == "get_engine") or + (isinstance(node, ast.ClassDef) and node.name == "ConnectionPoolManager")] + + def build_engine(url, creator, **options): + assert url == "mssql+pymssql://" + assert options["pool_recycle"] == 3600 + # Force a new DBAPI connection at the next checkout rather than waiting an hour. + return create_engine("sqlite://", creator=creator, poolclass=QueuePool, + pool_recycle=0) + + namespace = { + "CoreDatasource": CoreDatasource, + "AssistantOutDsSchema": type("AssistantOutDsSchema", (), {}), + "DatasourceConf": DatasourceConf, + "Engine": object, + "json": json, + "aes_decrypt": lambda value: value, + "equals_ignore_case": lambda left, right: left.lower() == right.lower(), + "create_engine": build_engine, + "get_origin_connect": connect, + "NullPool": NullPool, + "threading": threading, + "OrderedDict": OrderedDict, + "sessionmaker": sessionmaker, + } + exec(compile(ast.Module(body=selected, type_ignores=[]), str(SOURCE), "exec"), namespace) + return namespace["ConnectionPoolManager"]() + + +@pytest.mark.parametrize("finish", ["commit", "rollback"]) +def test_preview_created_pool_reconnects_after_request_session_closes(finish): + connections = [] + + def connect(ds_type, conf): + connections.append((ds_type, conf.host, conf.database)) + return sqlite3.connect(":memory:") + + manager = _load_pool_code(connect) + metadata_engine = create_engine("sqlite://") + configuration = json.dumps({"host": "test-host", "database": "test-db"}) + datasource = CoreDatasource(id=1330, type="sqlServer", configuration=configuration) + make_transient_to_detached(datasource) + try: + with Session(metadata_engine) as request_session: + request_session.add(datasource) + with manager.get_pool(datasource)() as sql_session: + assert sql_session.execute(text("SELECT 1")).scalar_one() == 1 + getattr(request_session, finish)() + + # A later chat/MCP request reuses the pool created by the preview request. + later_datasource = CoreDatasource(id=1330, type="sqlServer", configuration=configuration) + with manager.get_pool(later_datasource)() as sql_session: + assert sql_session.execute(text("SELECT 1")).scalar_one() == 1 + assert len(connections) >= 2 + assert all(connection == ("sqlServer", "test-host", "test-db") + for connection in connections) + finally: + manager.close_all() + metadata_engine.dispose()