如何在类中通过SQLAlchemy事件移除pyodbc的Trusted_Connection=Yes并使用Azure AD令牌
问题:SQLAlchemy类实例化引擎后如何注册do_connect事件
我参考SQLAlchemy文档的示例,在全局创建引擎时可以用@event.listens_for装饰器移除连接字符串中的Trusted_Connection=Yes并结合Azure AD令牌认证:
import struct from sqlalchemy import create_engine, event from sqlalchemy.engine.url import URL from azure import identity SQL_COPT_SS_ACCESS_TOKEN = 1256 # msodbcsql.h中定义的访问令牌连接选项 TOKEN_URL = "https://database.windows.net/" # Azure SQL数据库的令牌URL connection_string = "mssql+pyodbc://@my-server.database.windows.net/myDb?driver=ODBC+Driver+17+for+SQL+Server" engine = create_engine(connection_string) azure_credentials = identity.DefaultAzureCredential() @event.listens_for(engine, "do_connect") def provide_token(dialect, conn_rec, cargs, cparams): # 移除SQLAlchemy自动添加的"Trusted_Connection"参数 cargs[0] = cargs[0].replace(";Trusted_Connection=Yes", "") # 创建令牌凭证 raw_token = azure_credentials.get_token(TOKEN_URL).token.encode("utf-16-le") token_struct = struct.pack(f"<I{len(raw_token)}s", len(raw_token), raw_token) # 将令牌应用到连接参数 cparams["attrs_before"] = {SQL_COPT_SS_ACCESS_TOKEN: token_struct}
但我现在通过Db类管理数据库连接,engine是类实例化后才创建的属性(依赖环境配置),尝试用装饰器注册事件时出现错误:AttributeError: type object 'Db' has no attribute 'engine',我的代码如下:
import os import struct from azure.identity import DefaultAzureCredential from sqlalchemy.engine.url import URL from sqlalchemy import create_engine, event from sqlalchemy import inspect class Db: def __init__(self, config: object) -> None: url = URL.create( drivername="mssql+pyodbc", port=1433, query=dict(driver='ODBC Driver 18 for SQL Server'), host=f"tcp:{os.environ.get(f'SERVER_{config.environment}')}", database=os.environ.get(f'DATABASE_{config.environment}') ) self.engine = create_engine(url=url, connect_args={"autocommit": True}) self.connection = self.engine.connect() self.inspector = inspect(subject=self.engine) def close(self) -> None: self.connection.close() @event.listens_for(target=Db.engine, identifier="do_connect") def provide_token(dialect, conn_rec, cargs, cparams): # 移除SQLAlchemy自动添加的"Trusted_Connection"参数 cargs[0] = cargs[0].replace(";Trusted_Connection=Yes", "") azure_credentials = DefaultAzureCredential() raw_token = azure_credentials.get_token(TOKEN_URL).token.encode("utf-16-le") token_struct = struct.pack(f"<I{len(raw_token)}s", len(raw_token), raw_token) # 将令牌应用到连接参数 cparams["attrs_before"] = {1256: token_struct}
错误原因
@event.listens_for的target参数需要接收已存在的Engine实例或者对应的类,但Db.engine是类的实例属性,只有当类被实例化后才会存在,直接通过Db.engine访问会报错,因为类本身没有这个属性。
解决方案
方法1:在类的__init__中注册事件(推荐)
在创建self.engine之后,使用event.listen()函数手动注册事件,而不是装饰器。这样可以确保事件绑定到已经创建的实例引擎上。
修正后的代码:
import os import struct from azure.identity import DefaultAzureCredential from sqlalchemy.engine.url import URL from sqlalchemy import create_engine, event from sqlalchemy import inspect SQL_COPT_SS_ACCESS_TOKEN = 1256 TOKEN_URL = "https://database.windows.net/" # 提前创建凭证实例,避免每次连接都初始化 azure_credentials = DefaultAzureCredential() class Db: def __init__(self, config: object) -> None: url = URL.create( drivername="mssql+pyodbc", port=1433, query=dict(driver='ODBC Driver 18 for SQL Server'), host=f"tcp:{os.environ.get(f'SERVER_{config.environment}')}", database=os.environ.get(f'DATABASE_{config.environment}') ) self.engine = create_engine(url=url, connect_args={"autocommit": True}) # 注册do_connect事件到当前实例的engine event.listen(self.engine, "do_connect", provide_token) self.connection = self.engine.connect() self.inspector = inspect(subject=self.engine) def close(self) -> None: self.connection.close() def provide_token(dialect, conn_rec, cargs, cparams): # 移除Trusted_Connection参数 cargs[0] = cargs[0].replace(";Trusted_Connection=Yes", "") # 获取并处理Azure AD令牌 raw_token = azure_credentials.get_token(TOKEN_URL).token.encode("utf-16-le") token_struct = struct.pack(f"<I{len(raw_token)}s", len(raw_token), raw_token) cparams["attrs_before"] = {SQL_COPT_SS_ACCESS_TOKEN: token_struct}
方法2:监听Engine类(全局生效)
如果希望所有Engine实例都应用这个事件逻辑,可以直接监听Engine类,这样任何新创建的Engine都会触发该事件:
from sqlalchemy.engine import Engine @event.listens_for(Engine, "do_connect") def provide_token(dialect, conn_rec, cargs, cparams): # 仅对Azure SQL的pyodbc连接生效,避免影响其他数据库 if dialect.name == "mssql" and dialect.driver == "pyodbc": cargs[0] = cargs[0].replace(";Trusted_Connection=Yes", "") raw_token = azure_credentials.get_token(TOKEN_URL).token.encode("utf-16-le") token_struct = struct.pack(f"<I{len(raw_token)}s", len(raw_token), raw_token) cparams["attrs_before"] = {SQL_COPT_SS_ACCESS_TOKEN: token_struct}
这种方法适合所有Engine实例都需要相同处理逻辑的场景,但要注意添加判断条件,避免影响其他数据库连接。
内容的提问来源于stack exchange,提问作者Adventure-Knorrig
相关产品推荐
相关产品推荐

