Skip to content

Commit 0c9f4e1

Browse files
authored
feat: 补充数据库存储 usage_metadata 字段并新增自动迁移缺失列能力 (#43)
- SessionStorageEvent 新增 usage_metadata 列,支持存储和读取 token 用量统计 - 新增 decode_usage_metadata 工具函数用于反序列化 - 新增 _migrate_missing_columns 方法,自动检测并 ALTER TABLE 添加 ORM 模型中新增但数据库缺失的列
1 parent 70ec5a2 commit 0c9f4e1

6 files changed

Lines changed: 122 additions & 22 deletions

File tree

tests/sessions/test_sql_session_service.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -66,8 +66,7 @@ def _make_event_with_function_call():
6666
async def _create_service(config=None):
6767
config = config or _make_config()
6868
svc = SqlSessionService(db_url="sqlite:///:memory:", session_config=config, is_async=False)
69-
with patch('trpc_agent_sdk.storage._sql.event.listen'):
70-
await svc._sql_storage.create_sql_engine()
69+
await svc._sql_storage.create_sql_engine()
7170
return svc
7271

7372

tests/storage/test_sql.py

Lines changed: 6 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -51,19 +51,15 @@ def async_db_url(self):
5151
async def sync_storage(self, db_url):
5252
"""Synchronous SQL storage fixture with initialized engine."""
5353
storage = SqlStorage(is_async=False, db_url=db_url, metadata=StorageData.metadata)
54-
# Patch event.listen to work around SQLite pragma issue
55-
with patch('trpc_agent_sdk.storage._sql.event.listen'):
56-
await storage.create_sql_engine()
54+
await storage.create_sql_engine()
5755
yield storage
5856
await storage.close()
5957

6058
@pytest.fixture
6159
async def async_storage(self, async_db_url):
6260
"""Asynchronous SQL storage fixture with initialized engine."""
6361
storage = SqlStorage(is_async=True, db_url=async_db_url, metadata=StorageData.metadata)
64-
# Patch event.listen to work around async engine event limitation
65-
with patch('trpc_agent_sdk.storage._sql.event.listen'):
66-
await storage.create_sql_engine()
62+
await storage.create_sql_engine()
6763
yield storage
6864
await storage.close()
6965

@@ -93,9 +89,7 @@ async def test_create_sql_engine_async(self, async_db_url):
9389
"""Test creating async SQL engine."""
9490
storage = SqlStorage(is_async=True, db_url=async_db_url, metadata=StorageData.metadata)
9591

96-
# Patch event.listen to work around async engine event limitation
97-
with patch('trpc_agent_sdk.storage._sql.event.listen'):
98-
await storage.create_sql_engine()
92+
await storage.create_sql_engine()
9993

10094
assert storage._db_engine is not None
10195
assert storage._database_session_factory is not None
@@ -108,9 +102,7 @@ async def test_create_sql_engine_sync(self, db_url):
108102
"""Test creating sync SQL engine."""
109103
storage = SqlStorage(is_async=False, db_url=db_url, metadata=StorageData.metadata)
110104

111-
# Patch event.listen to work around SQLite pragma issue in tests
112-
with patch('trpc_agent_sdk.storage._sql.event.listen'):
113-
await storage.create_sql_engine()
105+
await storage.create_sql_engine()
114106

115107
assert storage._db_engine is not None
116108
assert storage._database_session_factory is not None
@@ -473,9 +465,7 @@ async def test_close_async(self, async_db_url):
473465
"""Test closing async SQL engine."""
474466
storage = SqlStorage(is_async=True, db_url=async_db_url, metadata=StorageData.metadata)
475467

476-
# Patch event.listen to work around async engine event limitation
477-
with patch('trpc_agent_sdk.storage._sql.event.listen'):
478-
await storage.create_sql_engine()
468+
await storage.create_sql_engine()
479469

480470
assert storage._db_engine is not None
481471

@@ -487,9 +477,7 @@ async def test_close_sync(self, db_url):
487477
"""Test closing sync SQL engine."""
488478
storage = SqlStorage(is_async=False, db_url=db_url, metadata=StorageData.metadata)
489479

490-
# Patch event.listen to work around SQLite pragma issue in tests
491-
with patch('trpc_agent_sdk.storage._sql.event.listen'):
492-
await storage.create_sql_engine()
480+
await storage.create_sql_engine()
493481

494482
assert storage._db_engine is not None
495483

trpc_agent_sdk/sessions/_sql_session_service.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,7 @@
6464
from trpc_agent_sdk.storage import UTF8MB4String
6565
from trpc_agent_sdk.storage import decode_content
6666
from trpc_agent_sdk.storage import decode_grounding_metadata
67+
from trpc_agent_sdk.storage import decode_usage_metadata
6768
from trpc_agent_sdk.utils import user_key
6869

6970
from ._base_session_service import BaseSessionService
@@ -165,6 +166,7 @@ class SessionStorageEvent(SessionStorageBase):
165166

166167
content: Mapped[dict[str, Any]] = mapped_column(DynamicJSON, nullable=True)
167168
grounding_metadata: Mapped[dict[str, Any]] = mapped_column(DynamicJSON, nullable=True)
169+
usage_metadata: Mapped[dict[str, Any]] = mapped_column(DynamicJSON, nullable=True)
168170
custom_metadata: Mapped[dict[str, Any]] = mapped_column(DynamicJSON, nullable=True)
169171

170172
partial: Mapped[bool] = mapped_column(Boolean, nullable=True)
@@ -218,6 +220,8 @@ def from_event(cls, session: Session, event: Event) -> SessionStorageEvent:
218220
storage_event.content = event.content.model_dump(exclude_none=True, mode="json")
219221
if event.grounding_metadata:
220222
storage_event.grounding_metadata = event.grounding_metadata.model_dump(exclude_none=True, mode="json")
223+
if event.usage_metadata:
224+
storage_event.usage_metadata = event.usage_metadata.model_dump(exclude_none=True, mode="json")
221225
if event.custom_metadata:
222226
storage_event.custom_metadata = event.custom_metadata
223227
return storage_event
@@ -238,6 +242,7 @@ def to_event(self) -> Event:
238242
error_message=self.error_message,
239243
interrupted=self.interrupted,
240244
grounding_metadata=decode_grounding_metadata(self.grounding_metadata),
245+
usage_metadata=decode_usage_metadata(self.usage_metadata),
241246
custom_metadata=self.custom_metadata,
242247
)
243248

trpc_agent_sdk/storage/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@
2929
from ._sql_common import UTF8MB4String
3030
from ._sql_common import decode_content
3131
from ._sql_common import decode_grounding_metadata
32+
from ._sql_common import decode_usage_metadata
3233

3334
__all__ = [
3435
"EXPIRE_METHOD",
@@ -55,4 +56,5 @@
5556
"UTF8MB4String",
5657
"decode_content",
5758
"decode_grounding_metadata",
59+
"decode_usage_metadata",
5860
]

trpc_agent_sdk/storage/_sql.py

Lines changed: 93 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -19,11 +19,16 @@
1919
from sqlalchemy import MetaData
2020
from sqlalchemy import and_
2121
from sqlalchemy import delete as sql_delete
22+
from sqlalchemy import Dialect
23+
from sqlalchemy.sql.compiler import IdentifierPreparer
2224
from sqlalchemy import event
2325
from sqlalchemy import select
26+
from sqlalchemy import text
2427
from sqlalchemy.engine import Engine
2528
from sqlalchemy.engine import create_engine
2629
from sqlalchemy.engine.interfaces import DBAPICursor
30+
from sqlalchemy.engine import Connection
31+
from sqlalchemy.engine.reflection import Inspector
2732
from sqlalchemy.exc import ArgumentError
2833
from sqlalchemy.ext.asyncio import AsyncEngine
2934
from sqlalchemy.ext.asyncio import AsyncSession
@@ -122,6 +127,89 @@ def __init__(self, is_async: bool, db_url: str, metadata: Optional[MetaData] = N
122127
self.__db_url = db_url
123128
self.__kwargs = kwargs
124129

130+
def _migrate_missing_columns(self, connection: Connection) -> None:
131+
"""Add columns that exist in the ORM model but are missing from the database,
132+
for forward compatibility across version changes.
133+
134+
SQLAlchemy's create_all only creates tables — it never ALTERs existing
135+
tables. This helper bridges the gap for lightweight forward-only migrations.
136+
Only handles adding new columns (forward-only).
137+
138+
All-or-nothing semantics: on databases that support transactional DDL
139+
(e.g. PostgreSQL) the caller's transaction handles rollback. On databases
140+
where DDL auto-commits (e.g. MySQL), a compensating DROP COLUMN is issued
141+
for every column that was already added before the failure.
142+
143+
Args:
144+
connection: A synchronous SQLAlchemy Connection object.
145+
"""
146+
insp: Inspector = inspect(connection)
147+
dialect: Dialect = connection.dialect
148+
preparer: IdentifierPreparer = dialect.identifier_preparer
149+
ddl_compiler = dialect.ddl_compiler(dialect, None)
150+
151+
pending_add_columns: list[tuple[str, str, str]] = []
152+
for table_name, table in self.__metadata.tables.items():
153+
if not insp.has_table(table_name):
154+
continue
155+
existing: set[str] = {col["name"] for col in insp.get_columns(table_name)}
156+
for column in table.columns:
157+
if column.name in existing:
158+
continue
159+
col_type: str = column.type.compile(dialect=dialect)
160+
# handle different types of default value
161+
nullable: str = "" if column.nullable else " NOT NULL"
162+
default: str = ""
163+
default_value = ddl_compiler.get_column_default_string(column)
164+
if default_value is not None:
165+
default = f" DEFAULT {default_value}"
166+
elif column.server_default is not None:
167+
# if the column has server_default, but it is not a DDL server_default, warning
168+
logger.warning(
169+
"Column '%s' on table '%s' has a non-DDL server_default "
170+
"(%s); skipping DEFAULT clause generation.",
171+
column.name,
172+
table_name,
173+
type(column.server_default).__name__,
174+
)
175+
elif not column.nullable:
176+
# if the column is NOT NULL and has no server_default, raise error
177+
logger.warning(
178+
"Column '%s' on table '%s' is NOT NULL without a server_default; "
179+
"migration may fail if the table already contains rows.",
180+
column.name,
181+
table_name,
182+
)
183+
quoted_table: str = preparer.quote_identifier(table_name)
184+
quoted_col: str = preparer.quote_identifier(column.name)
185+
stmt: str = f"ALTER TABLE {quoted_table} ADD COLUMN {quoted_col} {col_type}{default}{nullable}"
186+
pending_add_columns.append((stmt, column.name, table_name))
187+
188+
if not pending_add_columns:
189+
return
190+
191+
added_columns: list[tuple[str, str]] = []
192+
try:
193+
for stmt, col_name, table_name in pending_add_columns:
194+
connection.execute(text(stmt))
195+
added_columns.append((col_name, table_name))
196+
logger.info("Auto-migrated: added column '%s' to table '%s'", col_name, table_name)
197+
except Exception:
198+
logger.error("Migration failed, compensating %d already-added column(s).", len(added_columns))
199+
for col_name, tbl_name in reversed(added_columns):
200+
drop_stmt = (f"ALTER TABLE {preparer.quote_identifier(tbl_name)} "
201+
f"DROP COLUMN {preparer.quote_identifier(col_name)}")
202+
try:
203+
connection.execute(text(drop_stmt))
204+
logger.info("Compensated: dropped column '%s' from table '%s'", col_name, tbl_name)
205+
except Exception:
206+
logger.error(
207+
"Failed to compensate column '%s' on table '%s'; manual cleanup required.",
208+
col_name,
209+
tbl_name,
210+
)
211+
raise
212+
125213
async def create_sql_engine(self):
126214
"""Create the database engine."""
127215
if self._db_engine:
@@ -137,16 +225,19 @@ async def _async_inspect():
137225
self.inspector = await _async_inspect()
138226
async with db_engine.begin() as conn:
139227
await conn.run_sync(self.__metadata.create_all)
228+
await conn.run_sync(self._migrate_missing_columns)
140229
self._database_session_factory = async_sessionmaker(bind=db_engine)
141230
else:
142231
db_engine: SqlEngine = create_engine(self.__db_url, **self.__kwargs)
143232
self.inspector = inspect(db_engine)
144233
self.__metadata.create_all(db_engine)
234+
with db_engine.begin() as conn:
235+
self._migrate_missing_columns(conn)
145236
self._database_session_factory = sessionmaker(bind=db_engine)
146237

147238
if db_engine.dialect.name == "sqlite":
148-
# Set sqlite pragma to enable foreign keys constraints
149-
event.listen(db_engine, "connect", _set_sqlite_pragma)
239+
listen_target = db_engine.sync_engine if isinstance(db_engine, AsyncEngine) else db_engine
240+
event.listen(listen_target, "connect", _set_sqlite_pragma)
150241

151242
except Exception as ex: # pylint: disable=broad-except
152243
if isinstance(ex, ArgumentError):

trpc_agent_sdk/storage/_sql_common.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@
4242

4343
from trpc_agent_sdk.types import Content
4444
from trpc_agent_sdk.types import GroundingMetadata
45+
from trpc_agent_sdk.types import GenerateContentResponseUsageMetadata
4546

4647

4748
def decode_content(content: Optional[dict[str, Any]]) -> Optional[Content]:
@@ -58,6 +59,20 @@ def decode_content(content: Optional[dict[str, Any]]) -> Optional[Content]:
5859
return Content.model_validate(content)
5960

6061

62+
def decode_usage_metadata(usage_metadata: Optional[dict[str, Any]]) -> Optional[GenerateContentResponseUsageMetadata]:
63+
"""Decode a usage metadata object from a JSON dictionary.
64+
65+
Args:
66+
usage_metadata: JSON dictionary containing usage metadata
67+
68+
Returns:
69+
Decoded GenerateContentResponseUsageMetadata object or None if usage_metadata is None
70+
"""
71+
if not usage_metadata:
72+
return None
73+
return GenerateContentResponseUsageMetadata.model_validate(usage_metadata)
74+
75+
6176
def decode_grounding_metadata(grounding_metadata: Optional[dict[str, Any]]) -> Optional[GroundingMetadata]:
6277
"""Decode a grounding metadata object from a JSON dictionary.
6378

0 commit comments

Comments
 (0)