diff --git a/agent/conversation_loop.py b/agent/conversation_loop.py index c86d1b124256..3563a09d745d 100644 --- a/agent/conversation_loop.py +++ b/agent/conversation_loop.py @@ -623,7 +623,10 @@ def run_conversation( if agent._memory_manager: try: _query = original_user_message if isinstance(original_user_message, str) else "" - _ext_prefetch_cache = agent._memory_manager.prefetch_all(_query) or "" + _ext_prefetch_cache = agent._memory_manager.prefetch_all( + _query, + user_id=getattr(agent, "_user_id", "") or "", + ) or "" except Exception: pass diff --git a/agent/memory_manager.py b/agent/memory_manager.py index 6692d8f04e71..2140fd5d54bf 100644 --- a/agent/memory_manager.py +++ b/agent/memory_manager.py @@ -336,7 +336,7 @@ def build_system_prompt(self) -> str: # -- Prefetch / recall --------------------------------------------------- - def prefetch_all(self, query: str, *, session_id: str = "") -> str: + def prefetch_all(self, query: str, *, session_id: str = "", user_id: str = "") -> str: """Collect prefetch context from all providers. Returns merged context text labeled by provider. Empty providers @@ -345,7 +345,7 @@ def prefetch_all(self, query: str, *, session_id: str = "") -> str: parts = [] for provider in self._providers: try: - result = provider.prefetch(query, session_id=session_id) + result = provider.prefetch(query, session_id=session_id, user_id=user_id) if result and result.strip(): parts.append(result) except Exception as e: diff --git a/agent/memory_provider.py b/agent/memory_provider.py index 6678683d113f..a6c543a42cc4 100644 --- a/agent/memory_provider.py +++ b/agent/memory_provider.py @@ -89,7 +89,7 @@ def system_prompt_block(self) -> str: """ return "" - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: """Recall relevant context for the upcoming turn. Called before each API call. Return formatted text to inject as diff --git a/gateway/run.py b/gateway/run.py index 7e34d99138c6..67c50daddfe9 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -16363,6 +16363,14 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: except KeyError: pass self._init_cached_agent_for_turn(agent, _interrupt_depth) + # Refresh caller identity — source may differ from the prior turn + # in shared-thread sessions (thread_sessions_per_user=False). + agent._user_id = source.user_id or '' + agent._user_name = source.user_name or '' + agent._chat_id = source.chat_id or '' + agent._chat_name = source.chat_name or '' + agent._chat_type = source.chat_type or '' + agent._thread_id = source.thread_id or '' logger.debug("Reusing cached agent for session %s", session_key) if agent is None: diff --git a/plugins/memory/byterover/__init__.py b/plugins/memory/byterover/__init__.py index 469332c049c8..e9ff2112b8e7 100644 --- a/plugins/memory/byterover/__init__.py +++ b/plugins/memory/byterover/__init__.py @@ -212,7 +212,7 @@ def system_prompt_block(self) -> str: "important facts, brv_status to check state." ) - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: """Run brv query synchronously before the agent's first LLM call. Blocks until the query completes (up to _QUERY_TIMEOUT seconds), ensuring diff --git a/plugins/memory/hindsight/__init__.py b/plugins/memory/hindsight/__init__.py index 022cf7210639..8c50dc3acf5b 100644 --- a/plugins/memory/hindsight/__init__.py +++ b/plugins/memory/hindsight/__init__.py @@ -1278,7 +1278,7 @@ def system_prompt_block(self) -> str: f"hindsight_retain to store facts." ) - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: if self._prefetch_thread and self._prefetch_thread.is_alive(): logger.debug("Prefetch: waiting for background thread to complete") self._prefetch_thread.join(timeout=3.0) diff --git a/plugins/memory/holographic/__init__.py b/plugins/memory/holographic/__init__.py index 9969bcdec123..f6753023fb96 100644 --- a/plugins/memory/holographic/__init__.py +++ b/plugins/memory/holographic/__init__.py @@ -203,7 +203,7 @@ def system_prompt_block(self) -> str: f"Use fact_feedback to rate facts after using them (trains trust scores)." ) - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: if not self._retriever or not query: return "" try: diff --git a/plugins/memory/honcho/__init__.py b/plugins/memory/honcho/__init__.py index 6cf6c02ac4fa..8cc1d1d1a63a 100644 --- a/plugins/memory/honcho/__init__.py +++ b/plugins/memory/honcho/__init__.py @@ -544,7 +544,7 @@ def system_prompt_block(self) -> str: return header - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: """Return base context (representation + card) plus dialectic supplement. Assembles two layers: diff --git a/plugins/memory/mem0/__init__.py b/plugins/memory/mem0/__init__.py index 2b2729d00f44..47f1a6669554 100644 --- a/plugins/memory/mem0/__init__.py +++ b/plugins/memory/mem0/__init__.py @@ -234,7 +234,7 @@ def system_prompt_block(self) -> str: "mem0_profile for a full overview." ) - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: if self._prefetch_thread and self._prefetch_thread.is_alive(): self._prefetch_thread.join(timeout=3.0) with self._prefetch_lock: @@ -274,6 +274,8 @@ def sync_turn(self, user_content: str, assistant_content: str, *, session_id: st if self._is_breaker_open(): return + effective_user_id = user_id or self._user_id + def _sync(): try: client = self._get_client() @@ -281,7 +283,7 @@ def _sync(): {"role": "user", "content": user_content}, {"role": "assistant", "content": assistant_content}, ] - client.add(messages, **self._write_filters()) + client.add(messages, user_id=effective_user_id, agent_id=self._agent_id) self._record_success() except Exception as e: self._record_failure() diff --git a/plugins/memory/memgw/__init__.py b/plugins/memory/memgw/__init__.py index b1330a0e89d0..29612c84b24f 100644 --- a/plugins/memory/memgw/__init__.py +++ b/plugins/memory/memgw/__init__.py @@ -143,6 +143,7 @@ def __init__(self): self._prefetch_method = 'recall' self._user_id = '' self._prefetch_result = '' + self._prefetch_result_user: str = '' self._prefetch_lock = threading.Lock() self._prefetch_thread: threading.Thread | None = None # Monotonic generation: only the latest queued prefetch may store its @@ -284,12 +285,17 @@ def _format_recall(payload: dict) -> str: lines.append(f'- {snippet}') return '\n'.join(lines) - def prefetch(self, query: str, *, session_id: str = '') -> str: + def prefetch(self, query: str, *, session_id: str = '', user_id: str = '') -> str: if self._prefetch_thread and self._prefetch_thread.is_alive(): self._prefetch_thread.join(timeout=3.0) with self._prefetch_lock: + # Discard a result queued for a different user to prevent cross-user leak. + if user_id and self._prefetch_result_user and self._prefetch_result_user != user_id: + self._prefetch_result = '' + self._prefetch_result_user = '' result = self._prefetch_result self._prefetch_result = '' + self._prefetch_result_user = '' if not result: return '' return f'## Memory Gateway\n{result}' @@ -299,6 +305,7 @@ def on_session_switch(self, new_session_id: str, **kwargs) -> None: # its cached result, so the new session can't be fed stale context. with self._prefetch_lock: self._prefetch_result = '' + self._prefetch_result_user = '' self._prefetch_gen += 1 def queue_prefetch(self, query: str, *, session_id: str = '', user_id: str = '') -> None: @@ -329,6 +336,7 @@ def _run(): with self._prefetch_lock: if my_gen == self._prefetch_gen: self._prefetch_result = text + self._prefetch_result_user = user_id self._record_success() except Exception as e: self._record_failure() diff --git a/plugins/memory/openviking/__init__.py b/plugins/memory/openviking/__init__.py index 2a59c810737d..eae7831cbf56 100644 --- a/plugins/memory/openviking/__init__.py +++ b/plugins/memory/openviking/__init__.py @@ -518,7 +518,7 @@ def system_prompt_block(self) -> str: "viking_remember, viking_add_resource." ) - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: """Return prefetched results from the background thread.""" if self._prefetch_thread and self._prefetch_thread.is_alive(): self._prefetch_thread.join(timeout=3.0) diff --git a/plugins/memory/retaindb/__init__.py b/plugins/memory/retaindb/__init__.py index 79f0f994fbd3..b4773b2fd886 100644 --- a/plugins/memory/retaindb/__init__.py +++ b/plugins/memory/retaindb/__init__.py @@ -594,7 +594,7 @@ def _reasoning_level(query: str) -> str: return "medium" return "high" - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: """Consume prefetched results and return them as a context block.""" with self._lock: context = self._context_result diff --git a/plugins/memory/supermemory/__init__.py b/plugins/memory/supermemory/__init__.py index 7d9011349d00..2e76945dbf1b 100644 --- a/plugins/memory/supermemory/__init__.py +++ b/plugins/memory/supermemory/__init__.py @@ -543,7 +543,7 @@ def system_prompt_block(self) -> str: lines.append(f"\n{self._custom_container_instructions}") return "\n".join(lines) - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: if not self._active or not self._auto_recall or not self._client or not query.strip(): return "" try: diff --git a/tests/agent/test_memory_provider.py b/tests/agent/test_memory_provider.py index aeb5cac6c863..47fb31c491d0 100644 --- a/tests/agent/test_memory_provider.py +++ b/tests/agent/test_memory_provider.py @@ -45,7 +45,7 @@ def initialize(self, session_id, **kwargs): def system_prompt_block(self) -> str: return self._prompt_block - def prefetch(self, query, *, session_id=""): + def prefetch(self, query, *, session_id="", user_id=""): self.prefetch_queries.append(query) return self._prefetch_result @@ -1248,3 +1248,4 @@ def test_no_compressor_no_injection(self): """Gate is moot without a context_compressor.""" tools, names, engine_names = self._run_context_engine_injection(None, None) assert tools == [] + diff --git a/tests/agent/test_memory_user_id.py b/tests/agent/test_memory_user_id.py index 77edda431e25..0193caeadf30 100644 --- a/tests/agent/test_memory_user_id.py +++ b/tests/agent/test_memory_user_id.py @@ -40,7 +40,7 @@ def initialize(self, session_id: str, **kwargs) -> None: def system_prompt_block(self) -> str: return "" - def prefetch(self, query: str, *, session_id: str = "") -> str: + def prefetch(self, query: str, *, session_id: str = "", user_id: str = "") -> str: return "" def sync_turn(self, user_content, assistant_content, *, session_id="", user_id=""): @@ -357,3 +357,4 @@ def test_user_id_none_by_default(self): agent._user_id = None assert agent._user_id is None +