-
Notifications
You must be signed in to change notification settings - Fork 78
Expand file tree
/
Copy path_redis_memory_service.py
More file actions
124 lines (110 loc) · 5.48 KB
/
Copy path_redis_memory_service.py
File metadata and controls
124 lines (110 loc) · 5.48 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
# Tencent is pleased to support the open source community by making tRPC-Agent-Python available.
#
# Copyright (C) 2026 Tencent. All rights reserved.
#
# tRPC-Agent-Python is licensed under Apache-2.0.
"""A Redis-based memory service for prototyping and multi-node sharing."""
from typing import Any
from typing import Optional
from typing_extensions import override
from trpc_agent_sdk.abc import MemoryServiceABC as BaseMemoryService
from trpc_agent_sdk.abc import MemoryServiceConfig
from trpc_agent_sdk.context import AgentContext
from trpc_agent_sdk.events import Event as EventCls
from trpc_agent_sdk.log import logger
from trpc_agent_sdk.sessions import Session
from trpc_agent_sdk.storage import RedisCommand
from trpc_agent_sdk.storage import RedisCondition
from trpc_agent_sdk.storage import RedisExpire
from trpc_agent_sdk.storage import RedisStorage
from trpc_agent_sdk.types import MemoryEntry
from trpc_agent_sdk.types import SearchMemoryResponse
from ._utils import extract_words_lower
from ._utils import format_timestamp
class RedisMemoryService(BaseMemoryService):
"""A Redis-based memory service for prototyping and multi-node sharing.
Uses keyword matching instead of semantic search.
Stores events in Redis as JSON.
"""
def __init__(
self,
db_url: str,
enabled: bool = False,
is_async: bool = False,
memory_service_config: Optional[MemoryServiceConfig] = None,
**kwargs: Any,
):
super().__init__(memory_service_config=memory_service_config, enabled=enabled)
# Redis needs default TTL configuration
self._redis_storage = self._create_storage(db_url=db_url, is_async=is_async, **kwargs)
def _create_storage(self, db_url: str, is_async: bool, **kwargs: Any) -> RedisStorage:
"""Create the backing storage.
Subclasses override this factory to preserve memory behavior while
selecting a deployment-specific Redis client.
"""
return RedisStorage(is_async=is_async, redis_url=db_url, **kwargs)
@override
async def store_session(self, session: Session, agent_context: Optional[AgentContext] = None) -> None:
# Store all events for the session in a Redis list as JSON
async with self._redis_storage.create_db_session() as redis_session:
key = f"memory:{session.save_key}:{session.id}"
events_json = [event.model_dump_json() for event in session.events if event.content and event.content.parts]
if events_json:
args = [key]
args.extend(events_json)
await self._redis_storage.delete(redis_session, key) # Remove old events
expire = RedisExpire(key=key, ttl=self._memory_service_config.ttl)
command = RedisCommand(method='rpush', args=tuple(args), expire=expire)
await self._redis_storage.execute_command(redis_session, command)
@override
async def search_memory(self,
key: str,
query: str,
limit: int = 10,
agent_context: Optional[AgentContext] = None) -> SearchMemoryResponse:
response = SearchMemoryResponse()
async with self._redis_storage.create_db_session() as redis_session:
pattern = f"memory:{key}:*"
events_json_list = await self._redis_storage.query(redis_session, pattern, RedisCondition(limit=-1))
# Extract words from query (handles both English and Chinese)
words_in_query = extract_words_lower(query)
count = 0
for redis_key, event_json in events_json_list:
has_valid_event = False
event = None
if not isinstance(event_json, list):
event_json = [event_json]
for data in event_json:
try:
event = EventCls.model_validate_json(data)
except Exception as ex: # pylint: disable=broad-except
logger.error("Error parsing event JSON: %s", ex)
continue
if not event or not event.content or not event.content.parts:
continue
words_in_event = extract_words_lower(' '.join(
[part.text for part in event.content.parts if part.text]))
if not words_in_event:
continue
if any(query_word in words_in_event for query_word in words_in_query):
response.memories.append(
MemoryEntry(
content=event.content,
author=event.author,
timestamp=format_timestamp(event.timestamp),
))
count += 1
has_valid_event = True
if limit > 0 and count >= limit:
break
# Refresh TTL on accessed keys that contain valid (non-expired) events
if has_valid_event:
expire = RedisExpire(key=redis_key, ttl=self._memory_service_config.ttl)
await self._redis_storage.expire(redis_session, expire)
if limit > 0 and count >= limit:
break
return response
@override
async def close(self) -> None:
await self._redis_storage.close()
await super().close()