File size: 3,477 Bytes
2eef9ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4db2d34
 
2eef9ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4db2d34
2eef9ea
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
backend/agents/memory_agent.py

The Memory Agent — manages what the system remembers.

Two phases:
  1. RETRIEVE (before task): Pull relevant memories to inform planning
  2. STORE (after task): Extract and store learnings from this execution

Memory types:
  - Episodic: "Last time I searched for X, Y worked well"
  - Semantic: "The capital of France is Paris"
  - Procedural: "To analyze CSV data: use run_python with pandas"
"""
from __future__ import annotations
import asyncio
from datetime import datetime, timezone

from ..state.graph_state import WorkflowState, AgentRole, make_agent_event, make_memory_entry
from ..memory.memory_store import long_term, short_term
from ..core.config import get_settings
from ..core.logger import get_logger

log = get_logger(__name__)


def memory_retrieve_node(state: WorkflowState) -> WorkflowState:
    """Run at start of workflow — retrieve relevant memories."""
    settings = get_settings()
    if not settings.enable_memory:
        return state

    log.info("Memory retrieval", task_id=state["task_id"])

    # Check short-term first (Redis)
    cached_state = short_term.get_state(state["task_id"])
    if cached_state:
        log.info("Restored from short-term cache")
        return {**state, **cached_state, "task_id": state["task_id"]}

    # Use task text directly — no LLM call needed to generate queries
    queries = [state["task"][:80]]

    # Retrieve from long-term memory
    all_memories = []
    for q in queries[:3]:
        mems = asyncio.run(long_term.retrieve(q, limit=3))
        all_memories.extend(mems)

    # Deduplicate
    seen = set()
    unique_memories = []
    for m in all_memories:
        if m["memory_id"] not in seen:
            seen.add(m["memory_id"])
            unique_memories.append(m)

    log.info("Memories retrieved", count=len(unique_memories))

    return {
        **state,
        "memories": unique_memories[:5],
        "events": state["events"] + [
            make_agent_event(
                AgentRole.MEMORY, "memories_retrieved",
                f"Retrieved {len(unique_memories)} relevant memories",
                {"count": len(unique_memories), "queries": queries},
            )
        ],
    }


def memory_store_node(state: WorkflowState) -> WorkflowState:
    """Run at end of workflow — store learnings from this execution."""
    settings = get_settings()
    if not settings.enable_memory:
        return state

    log.info("Memory storage", task_id=state["task_id"])

    # Store task record to DB
    asyncio.run(long_term.store_task(state["task_id"], state["task"], dict(state)))

    # Cache state in Redis for fast resume
    short_term.store_state(state["task_id"], dict(state))

    # Extract and store new memories
    memories_stored = 0
    for mem in state.get("new_memories", []):
        try:
            asyncio.run(long_term.store(mem, task_id=state["task_id"]))
            memories_stored += 1
        except Exception as e:
            log.warning("Memory store failed", error=str(e))

    # No extra LLM call — critic already added memory entries to new_memories

    log.info("Memories stored", count=memories_stored)

    return {
        **state,
        "events": state["events"] + [
            make_agent_event(
                AgentRole.MEMORY, "memories_stored",
                f"Stored {memories_stored} memories from this execution",
                {"count": memories_stored},
            )
        ],
    }