Skip to content

Commit 76601c6

Browse files
GWealecopybara-github
authored andcommitted
fix: name the async driver when a database URL uses a sync one
Co-authored-by: George Weale <gweale@google.com> PiperOrigin-RevId: 974018508
1 parent b499fbe commit 76601c6

2 files changed

Lines changed: 32 additions & 0 deletions

File tree

src/google/adk/sessions/database_session_service.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,7 @@
4040
from sqlalchemy.engine import make_url
4141
from sqlalchemy.exc import ArgumentError
4242
from sqlalchemy.exc import IntegrityError
43+
from sqlalchemy.exc import InvalidRequestError
4344
from sqlalchemy.ext.asyncio import async_sessionmaker
4445
from sqlalchemy.ext.asyncio import AsyncEngine
4546
from sqlalchemy.ext.asyncio import AsyncSession as DatabaseSessionFactory
@@ -95,6 +96,15 @@
9596
_MYSQL_DIALECT,
9697
_MARIADB_DIALECT,
9798
)
99+
# The driver a URL falls back to when it names none is synchronous for each of
100+
# these backends, and the asyncio extension of SQLAlchemy refuses a synchronous
101+
# driver, so such a URL cannot be used here.
102+
_ASYNC_DRIVER_BY_BACKEND = {
103+
_SQLITE_DIALECT: "aiosqlite",
104+
_POSTGRESQL_DIALECT: "asyncpg",
105+
_MYSQL_DIALECT: "aiomysql",
106+
_MARIADB_DIALECT: "asyncmy",
107+
}
98108
# Tuple key order for in-process per-session lock maps:
99109
# (app_name, user_id, session_id).
100110
_SessionLockKey: TypeAlias = tuple[str, str, str]
@@ -358,6 +368,16 @@ def __init__(
358368
raise ValueError(
359369
f"Invalid database URL format or argument '{redacted_url}'."
360370
) from e
371+
if isinstance(e, InvalidRequestError):
372+
backend = make_url(db_url).get_backend_name()
373+
async_driver = _ASYNC_DRIVER_BY_BACKEND.get(backend)
374+
message = (
375+
f"Database URL '{redacted_url}' resolves to a synchronous"
376+
" driver, but this service requires an asynchronous one."
377+
)
378+
if async_driver:
379+
message += f" Use a '{backend}+{async_driver}://' URL instead."
380+
raise ValueError(message) from e
361381
if isinstance(e, ImportError):
362382
raise ValueError(
363383
f"Database related module not found for URL '{redacted_url}'."

tests/unittests/sessions/test_session_service.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -46,6 +46,7 @@
4646
from sqlalchemy import text
4747
from sqlalchemy import update
4848
from sqlalchemy.exc import ArgumentError
49+
from sqlalchemy.exc import InvalidRequestError
4950
from sqlalchemy.ext.asyncio import create_async_engine
5051
from sqlalchemy.pool import StaticPool
5152

@@ -2799,6 +2800,7 @@ async def test_database_session_service_requires_one_argument():
27992800
RuntimeError('boom'),
28002801
ArgumentError('bad argument'),
28012802
ImportError('no driver'),
2803+
InvalidRequestError('not an async driver'),
28022804
],
28032805
)
28042806
def test_database_session_service_engine_error_hides_password(raised_error):
@@ -2835,6 +2837,16 @@ def test_database_session_service_malformed_url_reports_usable_error():
28352837
assert isinstance(exc_info.value.__cause__, ArgumentError)
28362838

28372839

2840+
def test_database_session_service_sync_driver_url_names_async_driver():
2841+
"""A synchronous URL is the common mistake, so name the driver that works."""
2842+
with pytest.raises(ValueError) as exc_info:
2843+
DatabaseSessionService('sqlite:///sessions.db')
2844+
2845+
message = str(exc_info.value)
2846+
assert 'synchronous' in message
2847+
assert 'sqlite+aiosqlite' in message
2848+
2849+
28382850
@pytest.mark.asyncio
28392851
async def test_database_session_service_sqlite_file_timestamp_read_after_reopen(
28402852
tmp_path,

0 commit comments

Comments
 (0)