Skip to content

Commit fd6d7b7

Browse files
committed
fix(storage): 优化 SQL 存储兼容性
- 新增通用的 Content 清理逻辑,在 SQL Session 和 Memory 存储读写时过滤空的 Content Part。 - 修复 MySQL 下 DynamicPickleType 写入 LONGBLOB 时未序列化导致的类型错误。 - 支持通过 sessionmaker_kwargs 向 SqlStorage 传递 sessionmaker 配置,例如 expire_on_commit。 - 将 SQL Session 中显式更新 update_time 的逻辑改为使用 datetime.now(),避免 func.now() 带来的 ORM 状态问题。 - 统一通过当前事件循环创建 Session 和 Memory 服务的清理任务。 - 优化 SQL Session 示例输出,在助手回复前打印用户问题。 - 在 SQL Memory/Session 示例中关闭 thinking 输出,使演示结果更清晰。 - 增加 SQL Content 清理和 MySQL DynamicPickleType 序列化相关测试。
1 parent ca82991 commit fd6d7b7

13 files changed

Lines changed: 117 additions & 25 deletions

File tree

examples/memory_service_with_sql/agent/agent.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,8 @@
1010
from trpc_agent_sdk.models import OpenAIModel
1111
from trpc_agent_sdk.tools import FunctionTool
1212
from trpc_agent_sdk.tools import load_memory_tool
13+
from trpc_agent_sdk.types import GenerateContentConfig
14+
from trpc_agent_sdk.types import HttpOptions
1315

1416
from .config import get_model_config
1517
from .prompts import INSTRUCTION
@@ -25,12 +27,18 @@ def _create_model() -> LLMModel:
2527

2628
def create_agent() -> LlmAgent:
2729
""" Create an agent"""
30+
generate_content_config = GenerateContentConfig(
31+
http_options=HttpOptions(extra_body={"chat_template_kwargs": {
32+
"enable_thinking": False
33+
}}),
34+
)
2835
agent = LlmAgent(
2936
name="assistant",
3037
description="A helpful assistant for conversation",
3138
model=_create_model(), # You can change this to your preferred model
3239
instruction=INSTRUCTION,
3340
tools=[FunctionTool(get_weather_report), load_memory_tool],
41+
generate_content_config=generate_content_config,
3442
)
3543
return agent
3644

examples/session_service_with_sql/agent/agent.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,8 @@
99
from trpc_agent_sdk.models import LLMModel
1010
from trpc_agent_sdk.models import OpenAIModel
1111
from trpc_agent_sdk.tools import FunctionTool
12+
from trpc_agent_sdk.types import GenerateContentConfig
13+
from trpc_agent_sdk.types import HttpOptions
1214

1315
from .config import get_model_config
1416
from .prompts import INSTRUCTION
@@ -24,12 +26,18 @@ def _create_model() -> LLMModel:
2426

2527
def create_agent() -> LlmAgent:
2628
""" Create an agent"""
29+
generate_content_config = GenerateContentConfig(
30+
http_options=HttpOptions(extra_body={"chat_template_kwargs": {
31+
"enable_thinking": False
32+
}}),
33+
)
2734
agent = LlmAgent(
2835
name="assistant",
2936
description="A helpful assistant for conversation",
3037
model=_create_model(), # You can change this to your preferred model
3138
instruction=INSTRUCTION,
3239
tools=[FunctionTool(get_weather_report)],
40+
generate_content_config=generate_content_config,
3341
)
3442
return agent
3543

examples/session_service_with_sql/run_agent.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -76,10 +76,9 @@ async def run_weather_agent():
7676
]
7777

7878
for query in demo_queries:
79-
# Use a new session for each query
80-
8179
user_content = Content(parts=[Part.from_text(text=query)])
8280

81+
print(f"👤 User: {query}")
8382
print("🤖 Assistant: ", end="", flush=True)
8483
async for event in runner.run_async(user_id=user_id, session_id=current_session_id, new_message=user_content):
8584
# Check if event.content exists

tests/sessions/test_sql_session_service.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -90,12 +90,39 @@ def test_from_event_with_function_call(self):
9090
storage_event = SessionStorageEvent.from_event(session, event)
9191
assert storage_event.content is not None
9292

93+
def test_from_event_drops_empty_parts(self):
94+
session = Session(id="s1", app_name="app", user_id="user", save_key="k")
95+
event = Event(
96+
invocation_id="inv-1",
97+
author="agent",
98+
content=Content(parts=[Part()]),
99+
)
100+
storage_event = SessionStorageEvent.from_event(session, event)
101+
assert storage_event.content is None
102+
93103
def test_from_event_no_content(self):
94104
session = Session(id="s1", app_name="app", user_id="user", save_key="k")
95105
event = Event(invocation_id="inv-1", author="agent", actions=EventActions())
96106
storage_event = SessionStorageEvent.from_event(session, event)
97107
assert storage_event.content is None
98108

109+
def test_to_event_drops_legacy_empty_parts(self):
110+
storage_event = SessionStorageEvent(
111+
id="e1",
112+
app_name="app",
113+
user_id="user",
114+
session_id="s1",
115+
invocation_id="inv-1",
116+
author="agent",
117+
actions=EventActions(),
118+
long_running_tool_ids=set(),
119+
timestamp=datetime.now(),
120+
model_flags=1,
121+
content={"parts": [{}], "role": "model"},
122+
)
123+
event = storage_event.to_event()
124+
assert event.content is None
125+
99126
def test_long_running_tool_ids_property(self):
100127
session = Session(id="s1", app_name="app", user_id="user", save_key="k")
101128
event = _make_event()

tests/storage/test_sql_common.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -332,6 +332,14 @@ def test_process_bind_param_spanner(self):
332332
result = dpt.process_bind_param(value, dialect)
333333
assert pickle.loads(result) == value
334334

335+
def test_process_bind_param_mysql(self):
336+
dpt = DynamicPickleType()
337+
dialect = _make_dialect("mysql")
338+
value = {"key": "value", "nums": [1, 2, 3]}
339+
result = dpt.process_bind_param(value, dialect)
340+
assert isinstance(result, bytes)
341+
assert pickle.loads(result) == value
342+
335343
def test_process_bind_param_non_spanner(self):
336344
dpt = DynamicPickleType()
337345
dialect = _make_dialect("sqlite")
@@ -352,6 +360,14 @@ def test_process_result_value_spanner(self):
352360
result = dpt.process_result_value(pickled, dialect)
353361
assert result == original
354362

363+
def test_process_result_value_mysql(self):
364+
dpt = DynamicPickleType()
365+
dialect = _make_dialect("mysql")
366+
original = {"key": "value", "nums": [1, 2, 3]}
367+
pickled = pickle.dumps(original)
368+
result = dpt.process_result_value(pickled, dialect)
369+
assert result == original
370+
355371
def test_process_result_value_non_spanner(self):
356372
dpt = DynamicPickleType()
357373
dialect = _make_dialect("sqlite")

trpc_agent_sdk/memory/_in_memory_memory_service.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -200,7 +200,7 @@ def _start_cleanup_task(self) -> None:
200200
return
201201

202202
self.__cleanup_stop_event = asyncio.Event()
203-
self.__cleanup_task = asyncio.create_task(self._cleanup_loop())
203+
self.__cleanup_task = asyncio.get_event_loop().create_task(self._cleanup_loop())
204204
logger.debug("Cleanup task created")
205205

206206
def _stop_cleanup_task(self) -> None:

trpc_agent_sdk/memory/_sql_memory_service.py

Lines changed: 9 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -43,6 +43,7 @@
4343
from trpc_agent_sdk.storage import SqlStorage
4444
from trpc_agent_sdk.storage import decode_content
4545
from trpc_agent_sdk.storage import decode_grounding_metadata
46+
from trpc_agent_sdk.storage import sanitize_content_json
4647

4748
from ._utils import extract_words_lower
4849
from ._utils import format_timestamp
@@ -110,7 +111,7 @@ def update_event(self, session: Session, event: Event):
110111
self.error_message = event.error_message
111112
self.interrupted = event.interrupted
112113
if event.content:
113-
self.content = event.content.model_dump(exclude_none=True, mode="json")
114+
self.content = sanitize_content_json(event.content.model_dump(exclude_none=True, mode="json"))
114115
if event.grounding_metadata:
115116
self.grounding_metadata = event.grounding_metadata.model_dump(exclude_none=True, mode="json")
116117
if event.custom_metadata:
@@ -135,7 +136,7 @@ def from_event(cls, session: Session, event: Event) -> MemStorageEvent:
135136
interrupted=event.interrupted,
136137
)
137138
if event.content:
138-
storage_event.content = event.content.model_dump(exclude_none=True, mode="json")
139+
storage_event.content = sanitize_content_json(event.content.model_dump(exclude_none=True, mode="json"))
139140
if event.grounding_metadata:
140141
storage_event.grounding_metadata = event.grounding_metadata.model_dump(exclude_none=True, mode="json")
141142
if event.custom_metadata:
@@ -150,7 +151,7 @@ def to_event(self) -> Event:
150151
branch=self.branch,
151152
actions=self.actions, # type: ignore
152153
timestamp=self.timestamp.timestamp(),
153-
content=decode_content(self.content),
154+
content=decode_content(sanitize_content_json(self.content)),
154155
long_running_tool_ids=self.long_running_tool_ids,
155156
partial=self.partial,
156157
turn_complete=self.turn_complete,
@@ -194,15 +195,15 @@ async def store_session(self, session: Session, agent_context: Optional[AgentCon
194195
195196
Only stores events that are not expired based on event_ttl_seconds.
196197
"""
197-
if not isinstance(session, Session):
198-
raise TypeError(f"Content must be a Session, got {type(session)}")
199-
200198
async with self._sql_storage.create_db_session() as sql_session:
201199
is_exist = False
202200
for event in session.events:
203201
if not event.is_model_visible():
204202
continue
205-
if event.content and event.content.parts:
203+
if not event.content or not event.content.parts:
204+
continue
205+
content = sanitize_content_json(event.content.model_dump(exclude_none=True, mode="json"))
206+
if content:
206207
is_exist = True
207208
# Check if the event already exists
208209
event_key = SqlKey(key=(event.id, session.save_key, session.id), storage_cls=MemStorageEvent)
@@ -324,7 +325,7 @@ def _start_cleanup_task(self) -> None:
324325
return
325326

326327
self.__cleanup_stop_event = asyncio.Event()
327-
self.__cleanup_task = asyncio.create_task(self._cleanup_loop())
328+
self.__cleanup_task = asyncio.get_event_loop().create_task(self._cleanup_loop())
328329
logger.debug("Memory cleanup task created")
329330

330331
def _stop_cleanup_task(self) -> None:

trpc_agent_sdk/memory/mem0_memory_service.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -328,7 +328,7 @@ def _start_cleanup_task(self) -> None:
328328
return
329329

330330
self.__cleanup_stop_event = asyncio.Event()
331-
self.__cleanup_task = asyncio.create_task(self._cleanup_loop())
331+
self.__cleanup_task = asyncio.get_event_loop().create_task(self._cleanup_loop())
332332
logger.debug("Mem0 memory cleanup task created")
333333

334334
def _stop_cleanup_task(self) -> None:

trpc_agent_sdk/sessions/_in_memory_session_service.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -399,7 +399,7 @@ def _start_cleanup_task(self) -> None:
399399
return
400400

401401
self.__cleanup_stop_event = asyncio.Event()
402-
self.__cleanup_task = asyncio.create_task(self._cleanup_loop())
402+
self.__cleanup_task = asyncio.get_event_loop().create_task(self._cleanup_loop())
403403
logger.debug("Cleanup task created")
404404

405405
def _stop_cleanup_task(self) -> None:

trpc_agent_sdk/sessions/_sql_session_service.py

Lines changed: 9 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,7 @@
6565
from trpc_agent_sdk.storage import decode_content
6666
from trpc_agent_sdk.storage import decode_grounding_metadata
6767
from trpc_agent_sdk.storage import decode_usage_metadata
68+
from trpc_agent_sdk.storage import sanitize_content_json
6869
from trpc_agent_sdk.utils import user_key
6970

7071
from ._base_session_service import BaseSessionService
@@ -236,7 +237,7 @@ def from_event(cls, session: Session, event: Event) -> SessionStorageEvent:
236237
response_id=event.response_id,
237238
)
238239
if event.content:
239-
storage_event.content = event.content.model_dump(exclude_none=True, mode="json")
240+
storage_event.content = sanitize_content_json(event.content.model_dump(exclude_none=True, mode="json"))
240241
if event.grounding_metadata:
241242
storage_event.grounding_metadata = event.grounding_metadata.model_dump(exclude_none=True, mode="json")
242243
if event.custom_metadata:
@@ -269,7 +270,7 @@ def to_event(self) -> Event:
269270
error_message=self.error_message,
270271
interrupted=self.interrupted,
271272
response_id=self.response_id,
272-
content=decode_content(self.content),
273+
content=decode_content(sanitize_content_json(self.content)),
273274
grounding_metadata=decode_grounding_metadata(self.grounding_metadata),
274275
custom_metadata=self.custom_metadata,
275276
usage_metadata=decode_usage_metadata(self.usage_metadata),
@@ -574,7 +575,7 @@ async def _update_app_state(self, sql_session: SqlSession, app_name: str, state_
574575
await self._sql_storage.add(sql_session, storage_app_state)
575576
else:
576577
storage_app_state.state = app_state # type: ignore
577-
storage_app_state.update_time = func.now()
578+
storage_app_state.update_time = datetime.now()
578579

579580
return app_state
580581

@@ -604,7 +605,7 @@ async def _get_app_state(self, sql_session: SqlSession, app_name: str) -> dict[s
604605
if storage_app_state:
605606
if not self._session_config.is_expired_by_timestamp(storage_app_state.update_time.timestamp()):
606607
app_state = storage_app_state.state
607-
storage_app_state.update_time = func.now()
608+
storage_app_state.update_time = datetime.now()
608609
await self._sql_storage.commit(sql_session)
609610

610611
return app_state
@@ -617,7 +618,7 @@ async def _get_user_state(self, sql_session: SqlSession, app_name: str, user_id:
617618
if storage_user_state:
618619
if not self._session_config.is_expired_by_timestamp(storage_user_state.update_time.timestamp()):
619620
user_state = storage_user_state.state
620-
storage_user_state.update_time = func.now()
621+
storage_user_state.update_time = datetime.now()
621622
await self._sql_storage.commit(sql_session)
622623

623624
return user_state
@@ -633,7 +634,7 @@ async def _get_session(self, sql_session: SqlSession, app_name: str, user_id: st
633634
logger.debug("Session %s is expired", session_id)
634635
return None
635636

636-
storage_session.update_time = func.now()
637+
storage_session.update_time = datetime.now()
637638
await self._sql_storage.commit(sql_session)
638639

639640
return storage_session
@@ -645,7 +646,7 @@ async def _cleanup_expired_async(self) -> None:
645646
Deletes all expired data in three batch SQL DELETE statements.
646647
"""
647648
async with self._sql_storage.create_db_session() as sql_session:
648-
# Calculate expiration threshold once (using database local time)
649+
# Calculate expiration threshold once in application time for cross-database compatibility.
649650
expire_before = datetime.now() - timedelta(seconds=self._session_config.ttl.ttl_seconds)
650651
total_deleted = 0
651652

@@ -721,7 +722,7 @@ def _start_cleanup_task(self) -> None:
721722
return
722723

723724
self.__cleanup_stop_event = asyncio.Event()
724-
self.__cleanup_task = asyncio.create_task(self._cleanup_loop())
725+
self.__cleanup_task = asyncio.get_event_loop().create_task(self._cleanup_loop())
725726
logger.debug("Cleanup task created")
726727

727728
def _stop_cleanup_task(self) -> None:

0 commit comments

Comments
 (0)