|
| 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