Skip to content

Commit 3705c20

Browse files
weimchraychen911
authored andcommitted
Bugfix: 修复SqlSessionService在sqlite下因时区设置不对导致频繁warn
- 问题:当前在SqlSessionService的实现中,默认创建DB的表中,update_time使用了sqlalchemy的func.now,但在更新时间时,使用了datetime.now,在sqlite实现里,func.now默认使用了utc的时间,而datetime.now不是utc的时间,导致append_event时,因时区不同,导致diff失败,warn告警 - 解决方案:总是使用func.now来更新时间
1 parent f49692a commit 3705c20

1 file changed

Lines changed: 38 additions & 14 deletions

File tree

trpc_agent_sdk/sessions/_sql_session_service.py

Lines changed: 38 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -103,6 +103,25 @@ class SessionStorageBase(DeclarativeBase):
103103
pass
104104

105105

106+
def _storage_dialect_name(storage: SessionStorageBase) -> Optional[str]:
107+
orm_session = inspect(storage).session
108+
if orm_session is None or orm_session.bind is None:
109+
return None
110+
return orm_session.bind.dialect.name
111+
112+
113+
def _timestamp_tz(value: datetime, dialect_name: Optional[str]) -> float:
114+
if dialect_name == "sqlite":
115+
return value.replace(tzinfo=timezone.utc).timestamp()
116+
return value.timestamp()
117+
118+
119+
def _expire_before(sql_session: SqlSession, ttl_seconds: int) -> datetime:
120+
if sql_session.bind is not None and sql_session.bind.dialect.name == "sqlite":
121+
return datetime.now(timezone.utc).replace(tzinfo=None) - timedelta(seconds=ttl_seconds)
122+
return datetime.now() - timedelta(seconds=ttl_seconds)
123+
124+
106125
class StorageSession(SessionStorageBase):
107126
"""Represents a session stored in the database with TTL support.
108127
@@ -135,14 +154,11 @@ def __repr__(self):
135154

136155
@property
137156
def _dialect_name(self) -> Optional[str]:
138-
session = inspect(self).session
139-
return session.bind.dialect.name if session else None # type: ignore
157+
return _storage_dialect_name(self)
140158

141159
@property
142160
def update_timestamp_tz(self) -> float:
143-
if self._dialect_name == "sqlite":
144-
return self.update_time.replace(tzinfo=timezone.utc).timestamp()
145-
return self.update_time.timestamp()
161+
return _timestamp_tz(self.update_time, self._dialect_name)
146162

147163
def to_session(
148164
self,
@@ -306,6 +322,10 @@ class StorageAppState(SessionStorageBase):
306322
state: Mapped[MutableDict[str, Any]] = mapped_column(MutableDict.as_mutable(DynamicJSON), default={})
307323
update_time: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now(), onupdate=func.now())
308324

325+
@property
326+
def update_timestamp_tz(self) -> float:
327+
return _timestamp_tz(self.update_time, _storage_dialect_name(self))
328+
309329

310330
class StorageUserState(SessionStorageBase):
311331
"""Represents a user state stored in the database with TTL support.
@@ -319,6 +339,10 @@ class StorageUserState(SessionStorageBase):
319339
state: Mapped[MutableDict[str, Any]] = mapped_column(MutableDict.as_mutable(DynamicJSON), default={})
320340
update_time: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now(), onupdate=func.now())
321341

342+
@property
343+
def update_timestamp_tz(self) -> float:
344+
return _timestamp_tz(self.update_time, _storage_dialect_name(self))
345+
322346

323347
class SqlSessionService(BaseSessionService):
324348
"""A SQL database implementation of the session service.
@@ -452,7 +476,7 @@ async def list_sessions(self, *, app_name: str, user_id: str) -> ListSessionsRes
452476

453477
sessions = []
454478
for storage_session in results:
455-
if self._session_config.is_expired_by_timestamp(storage_session.update_time.timestamp()):
479+
if self._session_config.is_expired_by_timestamp(storage_session.update_timestamp_tz):
456480
logger.debug("Cleaned up expired session: %s/%s/%s", storage_session.app_name,
457481
storage_session.user_id, storage_session.id)
458482
continue
@@ -593,7 +617,7 @@ async def _update_app_state(self, sql_session: SqlSession, app_name: str, state_
593617
await self._sql_storage.add(sql_session, storage_app_state)
594618
else:
595619
storage_app_state.state = app_state # type: ignore
596-
storage_app_state.update_time = datetime.now()
620+
storage_app_state.update_time = func.now()
597621

598622
return app_state
599623

@@ -621,9 +645,9 @@ async def _get_app_state(self, sql_session: SqlSession, app_name: str) -> dict[s
621645

622646
app_state = {}
623647
if storage_app_state:
624-
if not self._session_config.is_expired_by_timestamp(storage_app_state.update_time.timestamp()):
648+
if not self._session_config.is_expired_by_timestamp(storage_app_state.update_timestamp_tz):
625649
app_state = storage_app_state.state
626-
storage_app_state.update_time = datetime.now()
650+
storage_app_state.update_time = func.now()
627651
await self._sql_storage.commit(sql_session)
628652

629653
return app_state
@@ -634,9 +658,9 @@ async def _get_user_state(self, sql_session: SqlSession, app_name: str, user_id:
634658

635659
user_state = {}
636660
if storage_user_state:
637-
if not self._session_config.is_expired_by_timestamp(storage_user_state.update_time.timestamp()):
661+
if not self._session_config.is_expired_by_timestamp(storage_user_state.update_timestamp_tz):
638662
user_state = storage_user_state.state
639-
storage_user_state.update_time = datetime.now()
663+
storage_user_state.update_time = func.now()
640664
await self._sql_storage.commit(sql_session)
641665

642666
return user_state
@@ -648,11 +672,11 @@ async def _get_session(self, sql_session: SqlSession, app_name: str, user_id: st
648672
if storage_session is None:
649673
return None
650674

651-
if self._session_config.is_expired_by_timestamp(storage_session.update_time.timestamp()):
675+
if self._session_config.is_expired_by_timestamp(storage_session.update_timestamp_tz):
652676
logger.debug("Session %s is expired", session_id)
653677
return None
654678

655-
storage_session.update_time = datetime.now()
679+
storage_session.update_time = func.now()
656680
await self._sql_storage.commit(sql_session)
657681

658682
return storage_session
@@ -665,7 +689,7 @@ async def _cleanup_expired_async(self) -> None:
665689
"""
666690
async with self._sql_storage.create_db_session() as sql_session:
667691
# Calculate expiration threshold once in application time for cross-database compatibility.
668-
expire_before = datetime.now() - timedelta(seconds=self._session_config.ttl.ttl_seconds)
692+
expire_before = _expire_before(sql_session, self._session_config.ttl.ttl_seconds)
669693
total_deleted = 0
670694

671695
# Batch delete expired sessions

0 commit comments

Comments
 (0)