mirror of
https://github.com/supermemoryai/supermemory.git
synced 2026-10-01 02:01:40 +00:00
Merge 443874ab28 into cfa6c7cb17
This commit is contained in:
commit
2255bc3ab0
2 changed files with 65 additions and 32 deletions
|
|
@ -477,45 +477,46 @@ class SupermemoryCartesiaAgent:
|
|||
Yields:
|
||||
Output events from the wrapped agent.
|
||||
"""
|
||||
if type(event).__name__ != "UserTurnEnded":
|
||||
async for output in self.agent.process(env, event):
|
||||
yield output
|
||||
return
|
||||
|
||||
try:
|
||||
if type(event).__name__ == "UserTurnEnded":
|
||||
logger.info("[Supermemory] Processing UserTurnEnded event")
|
||||
event, memory_context = await self._enrich_event_with_memories(event)
|
||||
logger.info("[Supermemory] Processing UserTurnEnded event")
|
||||
event, memory_context = await self._enrich_event_with_memories(event)
|
||||
|
||||
# Store conversation in background
|
||||
if hasattr(event, 'history') and event.history:
|
||||
new_messages = self._new_history_messages(event.history)
|
||||
if new_messages:
|
||||
logger.info(
|
||||
f"[Supermemory] Queuing {len(new_messages)} messages for storage"
|
||||
)
|
||||
task = asyncio.create_task(self._store_messages(new_messages))
|
||||
self._background_tasks.add(task)
|
||||
task.add_done_callback(self._background_tasks.discard)
|
||||
else:
|
||||
# No history yet, store just the current user message
|
||||
current_messages = self._extract_conversation_from_history([event])
|
||||
if not current_messages:
|
||||
user_content = self._extract_user_message(event)
|
||||
if user_content:
|
||||
current_messages = [{"role": "user", "content": user_content}]
|
||||
new_messages = self._new_messages_from_sequence(current_messages)
|
||||
if new_messages:
|
||||
logger.info("[Supermemory] No history, storing current user message")
|
||||
task = asyncio.create_task(self._store_messages(new_messages))
|
||||
self._background_tasks.add(task)
|
||||
task.add_done_callback(self._background_tasks.discard)
|
||||
|
||||
async for output in self._process_agent(env, event, memory_context):
|
||||
yield output
|
||||
# Store conversation in background
|
||||
if hasattr(event, 'history') and event.history:
|
||||
new_messages = self._new_history_messages(event.history)
|
||||
if new_messages:
|
||||
logger.info(
|
||||
f"[Supermemory] Queuing {len(new_messages)} messages for storage"
|
||||
)
|
||||
task = asyncio.create_task(self._store_messages(new_messages))
|
||||
self._background_tasks.add(task)
|
||||
task.add_done_callback(self._background_tasks.discard)
|
||||
else:
|
||||
async for output in self.agent.process(env, event):
|
||||
yield output
|
||||
|
||||
# No history yet, store just the current user message
|
||||
current_messages = self._extract_conversation_from_history([event])
|
||||
if not current_messages:
|
||||
user_content = self._extract_user_message(event)
|
||||
if user_content:
|
||||
current_messages = [{"role": "user", "content": user_content}]
|
||||
new_messages = self._new_messages_from_sequence(current_messages)
|
||||
if new_messages:
|
||||
logger.info("[Supermemory] No history, storing current user message")
|
||||
task = asyncio.create_task(self._store_messages(new_messages))
|
||||
self._background_tasks.add(task)
|
||||
task.add_done_callback(self._background_tasks.discard)
|
||||
except Exception as e:
|
||||
logger.error(f"[Supermemory] Error in process: {e}")
|
||||
async for output in self.agent.process(env, event):
|
||||
yield output
|
||||
return
|
||||
|
||||
async for output in self._process_agent(env, event, memory_context):
|
||||
yield output
|
||||
|
||||
def reset_memory_tracking(self) -> None:
|
||||
"""Reset memory tracking for a new conversation."""
|
||||
|
|
|
|||
|
|
@ -92,6 +92,38 @@ class TestSupermemoryCartesiaNullProfile(unittest.IsolatedAsyncioTestCase):
|
|||
self.assertIsNotNone(context)
|
||||
self.assertIn(fact, context)
|
||||
|
||||
async def test_process_does_not_retry_agent_after_output(self) -> None:
|
||||
class UserTurnEnded:
|
||||
pass
|
||||
|
||||
class FailingAfterOutputAgent:
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def process(self, _env, _event, **_kwargs):
|
||||
self.calls += 1
|
||||
yield "partial"
|
||||
raise RuntimeError("agent failed after output")
|
||||
|
||||
wrapped_agent = FailingAfterOutputAgent()
|
||||
agent = SupermemoryCartesiaAgent(
|
||||
agent=wrapped_agent,
|
||||
api_key="mock_key",
|
||||
container_tag="user-123",
|
||||
custom_id="conversation-456",
|
||||
add_memory="never",
|
||||
)
|
||||
event = UserTurnEnded()
|
||||
agent._enrich_event_with_memories = AsyncMock(return_value=(event, None))
|
||||
|
||||
outputs = []
|
||||
with self.assertRaisesRegex(RuntimeError, "agent failed after output"):
|
||||
async for output in agent.process(None, event):
|
||||
outputs.append(output)
|
||||
|
||||
self.assertEqual(outputs, ["partial"])
|
||||
self.assertEqual(wrapped_agent.calls, 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue