1515# ===============================================================================
1616
1717import asyncio
18+ import copy
1819import getpass
1920import os
2021from contextlib import contextmanager
2930)
3031from sqlalchemy .util import await_only
3132
33+ from services .util import get_bool_env
34+
3235load_dotenv ()
3336driver = os .environ .get ("DB_DRIVER" , "" )
3437
3538
39+ def get_iam_login_token () -> str :
40+ """
41+ Return a short-lived IAM DB auth token for Cloud SQL Postgres.
42+ """
43+ from google .auth import default
44+ from google .auth .transport .requests import Request
45+
46+ scopes = ["https://www.googleapis.com/auth/sqlservice.login" ]
47+ creds , _ = default ()
48+ if hasattr (creds , "with_scopes" ):
49+ creds = creds .with_scopes (scopes = scopes )
50+ else :
51+ creds = copy .copy (creds )
52+ creds ._scopes = scopes # type: ignore[attr-defined]
53+ creds .refresh (Request ())
54+ if not getattr (creds , "token" , None ):
55+ raise RuntimeError ("Unable to acquire IAM DB auth token." )
56+ return creds .token
57+
58+
3659async def get_async_engine ():
3760 """
3861 Asynchronous database session generator.
@@ -48,14 +71,21 @@ def asyncify_connection():
4871 user = os .environ .get ("CLOUD_SQL_USER" )
4972 password = os .environ .get ("CLOUD_SQL_PASSWORD" )
5073 database = os .environ .get ("CLOUD_SQL_DATABASE" )
51-
52- connection = connector .connect_async (
53- instance_name ,
54- "asyncpg" ,
55- db = database ,
56- password = password ,
57- user = user ,
58- )
74+ use_iam_auth = get_bool_env ("CLOUD_SQL_IAM_AUTH" , False )
75+ ip_type = os .environ .get ("CLOUD_SQL_IP_TYPE" , "public" )
76+
77+ connect_kwargs = {
78+ "db" : database ,
79+ "user" : user ,
80+ "enable_iam_auth" : use_iam_auth ,
81+ "ip_type" : ip_type ,
82+ }
83+ if use_iam_auth :
84+ connect_kwargs ["password" ] = get_iam_login_token ()
85+ else :
86+ connect_kwargs ["password" ] = password
87+
88+ connection = connector .connect_async (instance_name , "asyncpg" , ** connect_kwargs )
5989
6090 return AsyncAdapt_asyncpg_connection (
6191 engine .dialect .dbapi ,
@@ -78,15 +108,25 @@ def init_connection_pool(connector):
78108 user = os .environ .get ("CLOUD_SQL_USER" )
79109 password = os .environ .get ("CLOUD_SQL_PASSWORD" )
80110 database = os .environ .get ("CLOUD_SQL_DATABASE" )
111+ use_iam_auth = get_bool_env ("CLOUD_SQL_IAM_AUTH" , False )
112+ ip_type = os .environ .get ("CLOUD_SQL_IP_TYPE" , "public" )
81113
82114 def getconn ():
115+ connect_kwargs = {
116+ "user" : user ,
117+ "db" : database ,
118+ "ip_type" : ip_type ,
119+ "enable_iam_auth" : use_iam_auth ,
120+ }
121+ if use_iam_auth :
122+ connect_kwargs ["password" ] = get_iam_login_token ()
123+ else :
124+ connect_kwargs ["password" ] = password
125+
83126 conn = connector .connect (
84127 instance_name , # The Cloud SQL instance name
85128 "pg8000" ,
86- user = user ,
87- password = password ,
88- db = database ,
89- ip_type = "public" ,
129+ ** connect_kwargs ,
90130 )
91131 return conn
92132
@@ -107,7 +147,7 @@ def getconn():
107147 connector = Connector ()
108148 engine = init_connection_pool (connector )
109149
110- async_engine = asyncio .run (get_async_engine ())
150+ # async_engine = asyncio.run(get_async_engine())
111151
112152else :
113153 # if driver == "sqlite":
@@ -161,7 +201,7 @@ def getconn():
161201 # listen(engine, "connect", on_connect)
162202
163203
164- async_database_sessionmaker = async_sessionmaker (async_engine )
204+ # async_database_sessionmaker = async_sessionmaker(async_engine)
165205database_sessionmaker = sessionmaker (engine , expire_on_commit = False )
166206
167207
0 commit comments