Skip to content

Commit 7569e8d

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

4 files changed

Lines changed: 116 additions & 2 deletions

File tree

trpc_agent_sdk/sessions/_sql_session_service.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,8 @@
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
68+
from trpc_agent_sdk.types import GenerateContentResponseUsageMetadata
6769
from trpc_agent_sdk.utils import user_key
6870

6971
from ._base_session_service import BaseSessionService
@@ -165,6 +167,7 @@ class SessionStorageEvent(SessionStorageBase):
165167

166168
content: Mapped[dict[str, Any]] = mapped_column(DynamicJSON, nullable=True)
167169
grounding_metadata: Mapped[dict[str, Any]] = mapped_column(DynamicJSON, nullable=True)
170+
usage_metadata: Mapped[dict[str, Any]] = mapped_column(DynamicJSON, nullable=True)
168171
custom_metadata: Mapped[dict[str, Any]] = mapped_column(DynamicJSON, nullable=True)
169172

170173
partial: Mapped[bool] = mapped_column(Boolean, nullable=True)
@@ -218,6 +221,8 @@ def from_event(cls, session: Session, event: Event) -> SessionStorageEvent:
218221
storage_event.content = event.content.model_dump(exclude_none=True, mode="json")
219222
if event.grounding_metadata:
220223
storage_event.grounding_metadata = event.grounding_metadata.model_dump(exclude_none=True, mode="json")
224+
if event.usage_metadata:
225+
storage_event.usage_metadata = event.usage_metadata.model_dump(exclude_none=True, mode="json")
221226
if event.custom_metadata:
222227
storage_event.custom_metadata = event.custom_metadata
223228
return storage_event
@@ -238,6 +243,7 @@ def to_event(self) -> Event:
238243
error_message=self.error_message,
239244
interrupted=self.interrupted,
240245
grounding_metadata=decode_grounding_metadata(self.grounding_metadata),
246+
usage_metadata=decode_usage_metadata(self.usage_metadata),
241247
custom_metadata=self.custom_metadata,
242248
)
243249

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)