Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions docs/changelog.rst
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,9 @@ Unreleased
* BigQuery supports native query resource controls,
explicit STRUCT parameters, typed empty arrays, and configurable Storage Write
stream modes while retaining the atomic PENDING default.
* SQLite and aiosqlite can register custom window functions on Python 3.11
and later when the SQLite runtime supports them. Choose a transaction lock
mode or set the batch size for Arrow imports. Defaults stay the same.

* Arrow ODBC runs ``execute_many()`` one row at a time. It reports an unknown
row count since the native driver does not return the number of changed rows.
Expand Down Expand Up @@ -45,6 +48,9 @@ Unreleased

**Fixed:**

* SQLite pools replace lost in-memory connections. Arrow imports roll back
writes on failure or cancellation when the adapter owns the transaction.

* Builder results keep CTE trees independent, and column pruning no longer
exposes its cached expression to mutation. SQL generation avoids redundant
copies of temporary trees while preserving caller and cache ownership.
Expand Down
2 changes: 2 additions & 0 deletions sqlspec/adapters/aiosqlite/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
AiosqliteDriverFeatures,
AiosqliteFunctionConfig,
AiosqlitePoolParams,
AiosqliteWindowFunctionConfig,
)
from sqlspec.adapters.aiosqlite.core import build_connection_config, default_statement_config
from sqlspec.adapters.aiosqlite.driver import AiosqliteDriver, AiosqliteExceptionHandler
Expand Down Expand Up @@ -34,6 +35,7 @@
"AiosqlitePoolConnection",
"AiosqlitePoolParams",
"AiosqliteRawCursor",
"AiosqliteWindowFunctionConfig",
"build_connection_config",
"default_statement_config",
)
91 changes: 47 additions & 44 deletions sqlspec/adapters/aiosqlite/adk/store.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
from typing_extensions import NotRequired

from sqlspec.adapters.aiosqlite._typing import aiosqlite_sqlite_module as sqlite3
from sqlspec.adapters.aiosqlite.config import _render_pragmas
from sqlspec.adapters.aiosqlite.core import end_transaction, render_pragmas
from sqlspec.config import ADKConfig
from sqlspec.exceptions import ImproperConfigurationError
from sqlspec.extensions.adk import BaseAsyncADKStore, StoredEvent, StoredSession, normalize_session_list_options
Expand Down Expand Up @@ -128,7 +128,7 @@ async def create_session(
async with self._config.provide_connection() as conn:
await self._apply_pragmas(conn)
await conn.execute(sql, params)
await conn.commit()
await end_transaction(conn, commit=True)

return StoredSession(
id=session_id, app_name=app_name, user_id=user_id, state=state, create_time=now, update_time=now
Expand Down Expand Up @@ -165,7 +165,7 @@ async def get_session(
WHERE app_name = ? AND user_id = ? AND id = ?
"""
await conn.execute(update_sql, (_datetime_to_julian(datetime.now(timezone.utc)), *params))
await conn.commit()
await end_transaction(conn, commit=True)
cursor = await conn.execute(sql, params)
row = await cursor.fetchone()

Expand Down Expand Up @@ -206,7 +206,7 @@ async def update_session_state(self, app_name: str, user_id: str, session_id: st
async with self._config.provide_connection() as conn:
await self._apply_pragmas(conn)
await conn.execute(sql, (state_json, now_julian, app_name, user_id, session_id))
await conn.commit()
await end_transaction(conn, commit=True)

async def list_sessions(
self,
Expand Down Expand Up @@ -274,7 +274,7 @@ async def delete_session(self, app_name: str, user_id: str, session_id: str) ->
async with self._config.provide_connection() as conn:
await self._apply_pragmas(conn)
await conn.execute(sql, (app_name, user_id, session_id))
await conn.commit()
await end_transaction(conn, commit=True)

async def append_event(self, event_record: StoredEvent) -> None:
"""Append an event to a session.
Expand Down Expand Up @@ -305,7 +305,7 @@ async def append_event(self, event_record: StoredEvent) -> None:
event_data_json,
),
)
await conn.commit()
await end_transaction(conn, commit=True)

async def append_event_and_update_state(
self,
Expand Down Expand Up @@ -391,13 +391,13 @@ async def append_event_and_update_state(
if user_state is not None:
await conn.execute(user_upsert_sql, (app_name, user_id, to_json(user_state), now_julian))
except Exception:
await conn.rollback()
await end_transaction(conn, commit=False)
raise
else:
if row is None:
await conn.rollback()
await end_transaction(conn, commit=False)
else:
await conn.commit()
await end_transaction(conn, commit=True)

if row is None:
msg = f"Session {session_id} not found during append_event_and_update_state."
Expand Down Expand Up @@ -488,7 +488,7 @@ async def delete_expired_events(self, before: datetime, app_name: "str | None" =
await self._apply_pragmas(conn)
cursor = await conn.execute(sql, tuple(params))
deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0
await conn.commit()
await end_transaction(conn, commit=True)
return deleted_count
except sqlite3.OperationalError as exc:
if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc):
Expand All @@ -508,7 +508,7 @@ async def delete_idle_sessions(self, updated_before: datetime, app_name: "str |
await self._apply_pragmas(conn)
cursor = await conn.execute(sql, tuple(params))
deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0
await conn.commit()
await end_transaction(conn, commit=True)
return deleted_count
except sqlite3.OperationalError as exc:
if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc):
Expand All @@ -528,7 +528,7 @@ async def delete_idle_user_states(self, updated_before: datetime, app_name: "str
await self._apply_pragmas(conn)
cursor = await conn.execute(sql, tuple(params))
deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0
await conn.commit()
await end_transaction(conn, commit=True)
return deleted_count
except sqlite3.OperationalError as exc:
if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc):
Expand Down Expand Up @@ -582,7 +582,7 @@ async def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None
async with self._config.provide_connection() as conn:
await self._apply_pragmas(conn)
await conn.execute(sql, (app_name, to_json(state), _datetime_to_julian(datetime.now(timezone.utc))))
await conn.commit()
await end_transaction(conn, commit=True)

async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None:
"""Insert or replace user-scoped state for an application user."""
Expand All @@ -599,7 +599,7 @@ async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str,
await conn.execute(
sql, (app_name, user_id, to_json(state), _datetime_to_julian(datetime.now(timezone.utc)))
)
await conn.commit()
await end_transaction(conn, commit=True)

async def get_metadata(self, key: str) -> "str | None":
"""Return a value from the ADK internal metadata table."""
Expand Down Expand Up @@ -627,7 +627,7 @@ async def set_metadata(self, key: str, value: str) -> None:
async with self._config.provide_connection() as conn:
await self._apply_pragmas(conn)
await conn.execute(sql, (key, value))
await conn.commit()
await end_transaction(conn, commit=True)

async def _apply_pragmas(self, connection: Any) -> None:
"""Apply PRAGMA optimization profile for this connection.
Expand Down Expand Up @@ -799,20 +799,19 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "
if not entries:
return 0

inserted_count = 0
async with self._config.provide_connection() as conn:
for entry in entries:
params: tuple[Any, ...]
scope = entry.get("scope", "user")
if self._owner_id_column_name:
sql = f"""
INSERT OR IGNORE INTO {self._memory_table}
(id, session_id, app_name, user_id, scope, event_id, author,
{self._owner_id_column_name}, timestamp, content_json,
content_text, metadata_json, inserted_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"""
params = (
params_list: list[tuple[Any, ...]] = []
if self._owner_id_column_name:
sql = f"""
INSERT OR IGNORE INTO {self._memory_table}
(id, session_id, app_name, user_id, scope, event_id, author,
{self._owner_id_column_name}, timestamp, content_json,
content_text, metadata_json, inserted_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"""
for entry in entries:
scope = entry.get("scope", "user")
params_list.append((
entry["id"],
entry["session_id"],
entry["app_name"],
Expand All @@ -826,15 +825,17 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "
entry["content_text"],
to_json(entry["metadata_json"]),
_datetime_to_julian(entry["inserted_at"]),
)
else:
sql = f"""
INSERT OR IGNORE INTO {self._memory_table}
(id, session_id, app_name, user_id, scope, event_id, author,
timestamp, content_json, content_text, metadata_json, inserted_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"""
params = (
))
else:
sql = f"""
INSERT OR IGNORE INTO {self._memory_table}
(id, session_id, app_name, user_id, scope, event_id, author,
timestamp, content_json, content_text, metadata_json, inserted_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"""
for entry in entries:
scope = entry.get("scope", "user")
params_list.append((
entry["id"],
entry["session_id"],
entry["app_name"],
Expand All @@ -847,11 +848,13 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "
entry["content_text"],
to_json(entry["metadata_json"]),
_datetime_to_julian(entry["inserted_at"]),
)
cursor = await conn.execute(sql, params)
inserted_count += cursor.rowcount
))
cursor = await conn.executemany(sql, params_list)
try:
inserted_count = cursor.rowcount if cursor.rowcount >= 0 else len(params_list)
finally:
await cursor.close()
await conn.commit()
await end_transaction(conn, commit=True)
return inserted_count

async def search_entries(
Expand Down Expand Up @@ -917,7 +920,7 @@ async def delete_entries_by_session(self, session_id: str) -> int:
sql = f"DELETE FROM {self._memory_table} WHERE session_id = ?"
async with self._config.provide_connection() as conn:
cursor = await conn.execute(sql, (session_id,))
await conn.commit()
await end_transaction(conn, commit=True)
return cursor.rowcount

async def delete_entries_older_than(
Expand All @@ -940,7 +943,7 @@ async def delete_entries_older_than(

async with self._config.provide_connection() as conn:
cursor = await conn.execute(sql, tuple(params))
await conn.commit()
await end_transaction(conn, commit=True)
return cursor.rowcount

async def _memory_table_ddl(self) -> str:
Expand Down Expand Up @@ -1036,7 +1039,7 @@ def _pragma_overrides(config: "AiosqliteConfig") -> "list[tuple[str, str]]":
msg = "extension_config['adk']['pragma_overrides'] must be a mapping of PRAGMA names to values"
raise ImproperConfigurationError(msg)
try:
return _render_pragmas(pragma_overrides)
return render_pragmas(pragma_overrides)
except ImproperConfigurationError as exc:
msg = str(exc).replace("driver_features['pragmas']", "extension_config['adk']['pragma_overrides']")
raise ImproperConfigurationError(msg) from exc
Expand Down
Loading
Loading