diff --git a/api/apps/restful_apis/agent_api.py b/api/apps/restful_apis/agent_api.py index 0d545e7649..4d07842592 100644 --- a/api/apps/restful_apis/agent_api.py +++ b/api/apps/restful_apis/agent_api.py @@ -791,6 +791,8 @@ async def test_db_connection(): port=req["port"], password=req["password"], ) + with db.connection_context(): + db.execute_sql("SELECT 1") elif req["db_type"] == "oceanbase": db = MySQLDatabase( req["database"], @@ -800,6 +802,8 @@ async def test_db_connection(): password=req["password"], charset="utf8mb4", ) + with db.connection_context(): + db.execute_sql("SELECT 1") elif req["db_type"] == "postgres": db = PostgresqlDatabase( req["database"], @@ -808,6 +812,8 @@ async def test_db_connection(): port=req["port"], password=req["password"], ) + with db.connection_context(): + db.execute_sql("SELECT 1") elif req["db_type"] == "mssql": import pyodbc @@ -819,9 +825,14 @@ async def test_db_connection(): f"PWD={req['password']};" ) db = pyodbc.connect(connection_string) - cursor = db.cursor() - cursor.execute("SELECT 1") - cursor.close() + try: + cursor = db.cursor() + try: + cursor.execute("SELECT 1") + finally: + cursor.close() + finally: + db.close() elif req["db_type"] == "IBM DB2": import ibm_db @@ -844,7 +855,6 @@ async def test_db_connection(): stmt = ibm_db.exec_immediate(conn, "SELECT 1 FROM sysibm.sysdummy1") ibm_db.fetch_assoc(stmt) ibm_db.close(conn) - return get_json_result(data="Database Connection Successful!") elif req["db_type"] == "trino": import os import trino @@ -871,18 +881,18 @@ async def test_db_connection(): http_scheme=http_scheme, auth=auth, ) - cur = conn.cursor() - cur.execute("SELECT 1") - cur.fetchall() - cur.close() - conn.close() - return get_json_result(data="Database Connection Successful!") + try: + cur = conn.cursor() + try: + cur.execute("SELECT 1") + cur.fetchall() + finally: + cur.close() + finally: + conn.close() else: return server_error_response("Unsupported database type.") - if req["db_type"] != "mssql": - db.connect() - db.close() return get_json_result(data="Database Connection Successful!") except Exception as exc: return server_error_response(exc)