diff --git a/backend/.gitignore b/backend/.gitignore index 3967bcb3a..23b010539 100644 --- a/backend/.gitignore +++ b/backend/.gitignore @@ -31,3 +31,7 @@ config.yaml # Claude Code settings .claude/settings.local.json + +# pytest --basetemp workaround dirs (sandbox tmp_path permission) +.pytest_tmp/ +.pytest_tmp_run/ diff --git a/backend/AGENTS.md b/backend/AGENTS.md index 59226f6b9..c31e058bd 100644 --- a/backend/AGENTS.md +++ b/backend/AGENTS.md @@ -810,6 +810,7 @@ Focused regression coverage for the updater lives in `backend/tests/test_memory_ - `consolidation_min_facts` - Minimum facts in a category to trigger consolidation review (default: 8; range: 3–30) - `consolidation_max_groups_per_cycle` - Maximum categories the LLM can merge in one cycle (default: 3; range: 1–10; also controls the LLM's prompt instruction) - `consolidation_max_sources` - Maximum source facts per merge group; prevents over-merging (default: 8; range: 2–20) +- `watermark_max_keys` - Soft cap on the in-memory conversation-watermark cache (one entry per distinct thread/user/agent). A bounded LRU: when over capacity the least-recently-used entry is dropped, and a dropped key re-extracts one batch on that thread's next turn (same as a restart). Bounds memory in long-lived gateways handling many threads (default: 4096; 0 = unbounded) ### Reflection System (`packages/harness/deerflow/reflection/`) diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deer_mem.py b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deer_mem.py index a78d217d9..f9d7cf563 100644 --- a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deer_mem.py +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deer_mem.py @@ -33,14 +33,15 @@ from deerflow.agents.memory.manager import MemoryConflictError, MemoryCorruption from .deermem.config import DeerMemConfig from .deermem.core.llm import build_llm from .deermem.core.message_processing import ( - detect_correction, - detect_reinforcement, + SIGNAL_NAMES, + detect_signals, filter_messages_for_memory, + filter_trivial, load_patterns, ) from .deermem.core.paths import DEFAULT_AGENT_BUCKET from .deermem.core.prompt import format_memory_for_injection, load_prompt, load_prompt_messages, warm_tiktoken_cache -from .deermem.core.queue import MemoryUpdateQueue +from .deermem.core.queue import MemoryUpdateQueue, QueueFull from .deermem.core.storage import MemoryRevisionConflict, MemoryStorageCorruption, create_storage from .deermem.core.updater import MemoryUpdater, _coerce_source_confidence @@ -100,8 +101,7 @@ class DeerMem(MemoryManager): _llm: Any = PrivateAttr(default=None) _updater: Any = PrivateAttr(default=None) _queue: Any = PrivateAttr(default=None) - _correction_patterns: Any = PrivateAttr(default=None) - _reinforcement_patterns: Any = PrivateAttr(default=None) + _trivial_patterns: Any = PrivateAttr(default=None) # DeerMem implements search() (case-insensitive substring over stored facts), # so it is valid for mode="tool" (the base invariant validator requires this @@ -121,8 +121,12 @@ class DeerMem(MemoryManager): # Signal-detection patterns (externalized YAML; ``patterns_dir`` override # or bundled defaults = pre-externalization behavior). Loaded once at # construction and reused by ``_prepare_update``'s detect_* calls. - self._correction_patterns = load_patterns("correction", patterns_dir=self._config.patterns_dir) - self._reinforcement_patterns = load_patterns("reinforcement", patterns_dir=self._config.patterns_dir) + # Pre-load trivial + signal patterns at construction so a misconfigured + # patterns_dir (missing / invalid yaml) surfaces at startup, not on the + # first update. Compiled patterns are cached by load_patterns. + self._trivial_patterns = load_patterns("trivial", patterns_dir=self._config.patterns_dir) + for _signal_name in SIGNAL_NAMES: + load_patterns(_signal_name, patterns_dir=self._config.patterns_dir) # host_llm (host-injected default model) takes precedence over build_llm(model) # so zero-config DeerMem (empty `model`) still extracts via the app default, # mirroring pre-abstraction `model_name: null`. Standalone (no factory) -> None. @@ -168,7 +172,7 @@ class DeerMem(MemoryManager): ``model_post_init`` (shared with direct construction). """ config_dict = dict(backend_config or {}) - for key in ("should_keep_hidden_message", "trace_context_manager"): + for key in ("should_keep_hidden_message", "trace_context_manager", "extraction_callback"): if key not in config_dict and key in host_hooks: config_dict[key] = host_hooks[key] if "host_llm" not in config_dict: @@ -207,16 +211,25 @@ class DeerMem(MemoryManager): prepared = self._prepare_update(messages) if prepared is None: return - filtered, correction_detected, reinforcement_detected = prepared - self._queue.add( - thread_id=thread_id, - messages=filtered, - agent_name=_resolve_agent_name(agent_name), - user_id=user_id, - trace_id=trace_id, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, - ) + filtered, signals = prepared + # DeerMem owns the queue, so it owns the backpressure degradation: a + # QueueFull here is logged + dropped so memory backpressure degrades to + # "update skipped" rather than propagating into + # MemoryMiddleware.after_agent and breaking the agent run (peer + # middlewares self-guard the same way). The dropped update is re-fed + # next turn (the middleware passes the full conversation each cycle, and + # the watermark does not advance on a non-enqueued turn). + try: + self._queue.add( + thread_id=thread_id, + messages=filtered, + agent_name=_resolve_agent_name(agent_name), + user_id=user_id, + trace_id=trace_id, + signals=signals, + ) + except QueueFull as e: + logger.warning("Memory update rejected under backpressure (thread=%s): %s", thread_id, e) def add_nowait( self, @@ -234,37 +247,44 @@ class DeerMem(MemoryManager): prepared = self._prepare_update(messages) if prepared is None: return - filtered, correction_detected, reinforcement_detected = prepared - self._queue.add_nowait( - thread_id=thread_id, - messages=filtered, - agent_name=_resolve_agent_name(agent_name), - user_id=user_id, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, - ) + filtered, signals = prepared + # Defense-in-depth: the emergency path always admits under backpressure + # (see _enqueue_locked), so QueueFull is not expected here -- but the + # emergency flush is invoked from summarization_hook, so a propagated + # exception would break summarization. Catch + log to be safe. + try: + self._queue.add_nowait( + thread_id=thread_id, + messages=filtered, + agent_name=_resolve_agent_name(agent_name), + user_id=user_id, + signals=signals, + ) + except QueueFull as e: + logger.warning("Memory emergency flush rejected under backpressure (thread=%s): %s", thread_id, e) def _prepare_update( self, messages: list[Any], - ) -> tuple[list[Any], bool, bool] | None: + ) -> tuple[list[Any], frozenset[str]] | None: """Filter to user+final-AI messages, require both, detect signals. - Returns ``(filtered, correction_detected, reinforcement_detected)`` - or ``None`` when there is no meaningful conversation (missing a user - or an assistant turn). + Returns ``(filtered, signals)`` where ``signals`` is the set of signal + classes detected in the recent turns, or ``None`` when there is no + meaningful conversation (missing a user or an assistant turn, or every + turn dropped as a trivial pure-acknowledgment). """ filtered = filter_messages_for_memory( messages, should_keep_hidden_message=self._config.should_keep_hidden_message, ) + filtered = filter_trivial(filtered, patterns=self._trivial_patterns) user_messages = [m for m in filtered if getattr(m, "type", None) == "human"] assistant_messages = [m for m in filtered if getattr(m, "type", None) == "ai"] if not user_messages or not assistant_messages: return None - correction_detected = detect_correction(filtered, patterns=self._correction_patterns) - reinforcement_detected = not correction_detected and detect_reinforcement(filtered, patterns=self._reinforcement_patterns) - return filtered, correction_detected, reinforcement_detected + signals = detect_signals(filtered, patterns_dir=self._config.patterns_dir) + return filtered, frozenset(signals) # ── Read ───────────────────────────────────────────────────────────── def get_context( diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/config.py b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/config.py index 5af788615..651c24832 100644 --- a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/config.py +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/config.py @@ -79,6 +79,11 @@ class DeerMemConfig(BaseModel): le=300, description="Seconds to wait before processing queued updates (debounce).", ) + queue_max_depth: int = Field( + default=1000, + ge=0, + description=("Backpressure cap on pending items. 0 = unlimited. When the cap is reached, new non-signal updates are rejected (QueueFull); signal updates are always admitted so important memories are never shed."), + ) # ── Facts ──────────────────────────────────────────────────────────── max_facts: int = Field(default=100, ge=10, le=500, description="Maximum number of facts to store.") fact_confidence_threshold: float = Field( @@ -199,6 +204,29 @@ class DeerMemConfig(BaseModel): le=20, description=("Maximum number of source facts per consolidation group. Prevents the LLM from merging too many facts into one and losing important details."), ) + # ── Extraction quality callback (post-invoke observability) ───────── + extraction_callback: Any = Field( + default=None, + description=( + "Optional ``callback(metrics)`` invoked AFTER the extraction LLM " + "call (token usage, facts passing/rejected by the confidence " + "filter, rejection rate, prompt version). The host injects a " + "Langfuse-based callback to emit an extraction span; None = no " + "post-invoke observability. Set programmatically (not from YAML)." + ), + ) + # ── Watermark cache (in-memory, bounded LRU) ───────────────────────── + watermark_max_keys: int = Field( + default=4096, + ge=0, + description=( + "Soft cap on the in-memory conversation-watermark cache (one entry " + "per distinct thread/user/agent). The cache is a bounded LRU: when " + "over capacity the least-recently-used entry is dropped, and a " + "dropped key re-extracts one batch on that thread's next turn (the " + "same as a restart). 0 = unbounded." + ), + ) # ── Message processing (externalized patterns / prompts) ── patterns_dir: str | None = Field( default=None, diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/decision.yaml b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/decision.yaml new file mode 100644 index 000000000..7c503f6ba --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/decision.yaml @@ -0,0 +1,27 @@ +# Decision signal patterns for detect_signals (message_processing). +# +# Detects that the user made a decision / chose an option -- a high-value "what +# to remember" signal. Matched via ``search`` over the last 6 human turns. Keep +# these narrow to avoid false positives. +# +# Each list entry is either: +# - a string -> compiled with no flags +# - a mapping {pattern, flags} -> flags is a list of names: "ignorecase" +# +# In single-quoted strings backslashes are literal (regex \b stays \b); double +# '' for a literal single quote (apostrophe). + +- pattern: '\blet.s (?:go with|use|pick|choose)\b' + flags: [ignorecase] +- pattern: '\bI.ll (?:go with|use|pick|choose)\b' + flags: [ignorecase] +- pattern: '\bwe (?:should|will) (?:go with|use|pick|choose)\b' + flags: [ignorecase] +- pattern: '\bI (?:decide|chose|selected) to\b' + flags: [ignorecase] +- '就用' +- '决定用' +- '决定采用' +- '我们选' +- '我选' +- '就采用' diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/goal.yaml b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/goal.yaml new file mode 100644 index 000000000..4f65b3acd --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/goal.yaml @@ -0,0 +1,27 @@ +# Goal signal patterns for detect_signals (message_processing). +# +# Detects that the user stated an objective / intent / plan -- a high-value +# "what to remember" signal. Matched via ``search`` over the last 6 human +# turns. Keep these narrow to avoid false positives. +# +# Each list entry is either: +# - a string -> compiled with no flags +# - a mapping {pattern, flags} -> flags is a list of names: "ignorecase" +# +# In single-quoted strings backslashes are literal (regex \b stays \b); double +# '' for a literal single quote (apostrophe). + +- pattern: '\bI (?:plan to|want to|am going to|aim to|intend to)\b' + flags: [ignorecase] +- pattern: '\bmy goal is\b' + flags: [ignorecase] +- pattern: '\bI.m going to\b' + flags: [ignorecase] +- pattern: '\bnext I.ll\b' + flags: [ignorecase] +- '我的目标是' +- '我打算' +- '我准备' +- '接下来我要' +- '接下来我打算' +- '我想做' diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/identity.yaml b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/identity.yaml new file mode 100644 index 000000000..41e43b23d --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/identity.yaml @@ -0,0 +1,26 @@ +# Identity signal patterns for detect_signals (message_processing). +# +# Detects that the user stated something about who they are (role, profession, +# background) -- a high-value "what to remember" signal. Matched via ``search`` +# over the last 6 human turns. Keep these narrow to avoid false positives. +# +# Each list entry is either: +# - a string -> compiled with no flags +# - a mapping {pattern, flags} -> flags is a list of names: "ignorecase" +# +# In single-quoted strings backslashes are literal (regex \b stays \b); double +# '' for a literal single quote (apostrophe). + +- pattern: '\bI am (?:a |an )?[a-z][a-z -]{2,}\b' + flags: [ignorecase] +- pattern: '\bI work as\b' + flags: [ignorecase] +- pattern: '\bI.m (?:a |an )?[a-z][a-z -]{2,}\b' + flags: [ignorecase] +- pattern: '\bmy job is\b' + flags: [ignorecase] +- '我是' +- '我的职业是' +- '我的工作是' +- '我担任' +- '我是一名' diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/preference.yaml b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/preference.yaml new file mode 100644 index 000000000..c35a05c39 --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/preference.yaml @@ -0,0 +1,29 @@ +# Preference signal patterns for detect_signals (message_processing). +# +# Detects that the user stated a preference / dislike -- a high-value "what to +# remember" signal. Matched via ``search`` over the last 6 human turns. Keep +# these narrow to avoid false positives; extend here without touching code. +# +# Each list entry is either: +# - a string -> compiled with no flags +# - a mapping {pattern, flags} -> flags is a list of names: "ignorecase" +# +# In single-quoted strings backslashes are literal (regex \b stays \b); double +# '' for a literal single quote (apostrophe). + +- pattern: '\bI (?:really )?(?:prefer|like|love|enjoy|favor)\b' + flags: [ignorecase] +- pattern: '\bI (?:really )?(?:hate|dislike|don.t like|cannot stand)\b' + flags: [ignorecase] +- pattern: '\bmy favorite\b' + flags: [ignorecase] +- pattern: '\bI.d rather\b' + flags: [ignorecase] +- '我喜欢' +- '我偏好' +- '我更喜欢' +- '我偏爱' +- '我讨厌' +- '我不喜欢' +- '我钟爱' +- '我最喜欢' diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/trivial.yaml b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/trivial.yaml new file mode 100644 index 000000000..a7e44534f --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_patterns/trivial.yaml @@ -0,0 +1,36 @@ +# Trivial pure-acknowledgment patterns for filter_trivial (message_processing). +# +# Each entry is a regex matched against the WHOLE (stripped) human message via +# ``fullmatch`` -- a message is trivial only if it is nothing but an ack, so a +# substantive turn that happens to contain "ok" is never dropped. Trailing +# punctuation/whitespace is stripped before matching, so "ok." / "好的!" count. +# +# Each list entry is either: +# - a string -> compiled with no flags +# - a mapping {pattern, flags} -> flags is a list of names: "ignorecase" +# +# Keep these narrow: a false positive silently drops a real user turn from +# memory. Extend here without touching code. + +- pattern: '嗯+' +- pattern: 'ok(?:ay)?' + flags: [ignorecase] +- pattern: '好的[呢]?' +- pattern: '好[的呀]?' +- pattern: '谢谢' +- pattern: '感谢' +- pattern: '多谢' +- pattern: '对' +- pattern: '是的' +- pattern: '收到' +- pattern: '明白' +- pattern: '了解' +- pattern: '知道了' +- pattern: 'got it' + flags: [ignorecase] +- pattern: 'thanks' + flags: [ignorecase] +- pattern: 'thank you' + flags: [ignorecase] +- pattern: 'cheers' + flags: [ignorecase] diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_processing.py b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_processing.py index 0407b42d3..5d58e7849 100644 --- a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_processing.py +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/message_processing.py @@ -246,3 +246,95 @@ def detect_reinforcement(messages: list[Any], *, patterns: list[re.Pattern[str]] return True return False + + +# Signal classes detected by :func:`detect_signals`. Names align with the fact +# ``category`` enum (CORE_CATEGORIES) so a signal can drive an extraction +# category hint directly, except ``reinforcement`` (no same-named category; it +# maps to preference/behavior in the extraction hint). Keep new signal names in +# sync with CORE_CATEGORIES before adding them here. +SIGNAL_NAMES: tuple[str, ...] = ( + "correction", + "reinforcement", + "preference", + "identity", + "goal", + "decision", +) + + +def detect_signals( + messages: list[Any], + *, + patterns_dir: str | None = None, +) -> set[str]: + """Detect signal classes in the recent conversation turns. + + Returns the set of signal names whose patterns match any of the last 6 + human turns. This generalizes :func:`detect_correction` / + :func:`detect_reinforcement` (which remain for backward compatibility) to + the full signal set. The window stays ``messages[-6:]``. + """ + recent_user_msgs = [msg for msg in messages[-6:] if getattr(msg, "type", None) == "human"] + if not recent_user_msgs: + return set() + + hits: set[str] = set() + for name in SIGNAL_NAMES: + patterns = load_patterns(name, patterns_dir=patterns_dir) + if not patterns: + continue + for msg in recent_user_msgs: + content = extract_message_text(msg).strip() + if content and any(pattern.search(content) for pattern in patterns): + hits.add(name) + break + return hits + + +# Trailing characters stripped before a whole-message trivial match: a pure +# acknowledgment with trailing punctuation ("ok.", "好的!") is still trivial. +_TRIVIAL_TRAIL = " \t\n\r.。,,!!??;;" + + +def filter_trivial( + messages: list[Any], + *, + patterns: list[re.Pattern[str]] | None = None, + patterns_dir: str | None = None, +) -> list[Any]: + """Drop pure-acknowledgment human turns and their AI replies. + + A human turn is "trivial" when its whole (stripped) text matches a trivial + pattern (e.g. "嗯", "ok", "好的", "谢谢") -- matched via ``fullmatch`` so a + substantive turn containing "ok" is never dropped. The matched human turn + and its following assistant reply are both removed (reusing the + ``skip_next_ai`` discipline from :func:`filter_messages_for_memory`). When + every turn is trivial, the result is empty, which the caller treats as "do + not enqueue" (saving an extraction LLM call). + """ + if patterns is None: + patterns = load_patterns("trivial", patterns_dir=patterns_dir) + if not patterns: + return list(messages) + + result: list[Any] = [] + skip_next_ai = False + for msg in messages: + msg_type = getattr(msg, "type", None) + if msg_type == "human": + content = extract_message_text(msg).strip().rstrip(_TRIVIAL_TRAIL) + is_trivial = bool(content) and any(pattern.fullmatch(content) for pattern in patterns) + if is_trivial: + skip_next_ai = True + continue + result.append(msg) + skip_next_ai = False + elif msg_type == "ai": + tool_calls = getattr(msg, "tool_calls", None) + if not tool_calls: + if skip_next_ai: + skip_next_ai = False + continue + result.append(msg) + return result diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/prompt.py b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/prompt.py index 950d329b0..36c49eedb 100644 --- a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/prompt.py +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/prompt.py @@ -762,9 +762,16 @@ def format_conversation_for_update(messages: list[Any]) -> str: if not content: continue - # Truncate very long messages + # Truncate very long messages: keep the head (topic / opening) and the + # tail (conclusion / "remember X" instruction), dropping the middle. + # A head-only chop loses the tail's directives; a head+tail split + # preserves both. The separator is plain ASCII (no < > &) so the + # html.escape below leaves it intact and tells the LLM where text was + # cut. Escape happens after truncation, so the boundary never splits an + # entity (entities only exist after escaping). if len(str(content)) > 1000: - content = str(content)[:1000] + "..." + s = str(content) + content = s[:500] + "\n...[truncated]...\n" + s[-500:] # Escape < > & before embedding into the block of # the memory_update prompt. This raw user turn is the most attacker-influenced diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/queue.py b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/queue.py index 14a8f0f1c..b422262fb 100644 --- a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/queue.py +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/queue.py @@ -1,4 +1,18 @@ -"""Memory update queue with debounce mechanism.""" +"""Memory update queue with debounce mechanism. + +The queue collects conversation contexts and processes them after a +configurable debounce period; multiple contexts for the same +``(thread_id, user_id, agent_name)`` key are coalesced into one update. + +The queue is a process-local in-memory list plus a debounce +:class:`~threading.Timer`. Items still pending at process exit are lost +(best-effort :meth:`MemoryUpdateQueue.flush_sync` drain softens this for +graceful shutdown). Memory updates are best-effort: a failed or lost update is +re-fed on the next conversation turn (the middleware passes the full +conversation each cycle, and the updater's watermark does not advance on +failure), so an in-memory queue covers the realistic graceful-deploy case +without a persistence layer. +""" from __future__ import annotations @@ -17,6 +31,25 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +class QueueFull(Exception): + """Raised when a non-signal update is rejected under backpressure. + + Signal-bearing updates (any detected signal) are always admitted so that + important memories are never shed; only non-signal updates are rejected + once ``queue_max_depth`` is reached. Callers may catch this to degrade + (e.g. fall back to a synchronous write on the emergency path). + """ + + +def queue_key( + thread_id: str, + user_id: str | None, + agent_name: str | None, +) -> tuple[str, str | None, str | None]: + """Return the debounce identity for a memory update target.""" + return (thread_id, user_id, agent_name) + + @dataclass class ConversationContext: """Context for a conversation to be processed for memory update.""" @@ -27,8 +60,14 @@ class ConversationContext: agent_name: str | None = None user_id: str | None = None trace_id: str | None = None - correction_detected: bool = False - reinforcement_detected: bool = False + signals: frozenset[str] = field(default_factory=frozenset) + # Emergency (summarization) flushes bypass the updater's index watermark: + # the subset they carry is a one-shot "extract before removal" snapshot whose + # own length would otherwise regress the conversation watermark. Such contexts + # also coexist with (do not replace) a pending normal update for the same key + # so a flush cannot drop a pending normal update's un-extracted tail. See + # ``_enqueue_locked``'s match-key + backpressure handling. + bypass_watermark: bool = False class MemoryUpdateQueue: @@ -43,7 +82,7 @@ class MemoryUpdateQueue: """Initialize the memory update queue with injected config + updater.""" self._config = config self._updater = updater - self._queue: list[ConversationContext] = [] + self._items: list[ConversationContext] = [] self._lock = threading.Lock() self._timer: threading.Timer | None = None self._processing = False @@ -54,15 +93,6 @@ class MemoryUpdateQueue: self._processing_thread: threading.Thread | None = None self._reprocess_pending = False - @staticmethod - def _queue_key( - thread_id: str, - user_id: str | None, - agent_name: str | None, - ) -> tuple[str, str | None, str | None]: - """Return the debounce identity for a memory update target.""" - return (thread_id, user_id, agent_name) - def add( self, thread_id: str, @@ -70,8 +100,7 @@ class MemoryUpdateQueue: agent_name: str | None = None, user_id: str | None = None, trace_id: str | None = None, - correction_detected: bool = False, - reinforcement_detected: bool = False, + signals: frozenset[str] | None = None, ) -> None: """Add a conversation to the update queue. @@ -84,8 +113,9 @@ class MemoryUpdateQueue: raw threads). trace_id: Request trace id captured at enqueue time so the later Timer thread can attach it to memory LLM tracing metadata. - correction_detected: Whether recent turns include an explicit correction signal. - reinforcement_detected: Whether recent turns include a positive reinforcement signal. + signals: Signal classes detected in the conversation (correction / + reinforcement / preference / ...), used as extraction hints. Any + signal is admitted under backpressure. """ with self._lock: self._enqueue_locked( @@ -94,12 +124,12 @@ class MemoryUpdateQueue: agent_name=agent_name, user_id=user_id, trace_id=trace_id, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, + signals=frozenset(signals) if signals else frozenset(), + bypass_watermark=False, ) self._reset_timer() - logger.info("Memory update queued for thread %s, queue size: %d", thread_id, len(self._queue)) + logger.info("Memory update queued for thread %s, queue size: %d", thread_id, len(self._items)) def add_nowait( self, @@ -108,8 +138,7 @@ class MemoryUpdateQueue: agent_name: str | None = None, user_id: str | None = None, trace_id: str | None = None, - correction_detected: bool = False, - reinforcement_detected: bool = False, + signals: frozenset[str] | None = None, ) -> None: """Add a conversation and start processing immediately in the background.""" with self._lock: @@ -119,12 +148,12 @@ class MemoryUpdateQueue: agent_name=agent_name, user_id=user_id, trace_id=trace_id, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, + signals=frozenset(signals) if signals else frozenset(), + bypass_watermark=True, ) self._schedule_timer(0) - logger.info("Memory update queued for immediate processing on thread %s, queue size: %d", thread_id, len(self._queue)) + logger.info("Memory update queued for immediate processing on thread %s, queue size: %d", thread_id, len(self._items)) def _enqueue_locked( self, @@ -134,28 +163,45 @@ class MemoryUpdateQueue: agent_name: str | None, user_id: str | None, trace_id: str | None, - correction_detected: bool, - reinforcement_detected: bool, - ) -> None: - queue_key = self._queue_key(thread_id, user_id, agent_name) - existing_context = next( - (context for context in self._queue if self._queue_key(context.thread_id, context.user_id, context.agent_name) == queue_key), + signals: frozenset[str], + bypass_watermark: bool = False, + ) -> ConversationContext: + key = queue_key(thread_id, user_id, agent_name) + # Emergency (bypass) and normal updates coexist: the match key includes + # ``bypass_watermark`` so a summarization flush (bypass=True) never + # replaces a pending normal update for the same (thread, user, agent) -- + # replacing it would drop the normal update's un-extracted tail, which + # the next turn may not re-feed if the user stops. Both are processed + # independently instead. + existing = next( + (c for c in self._items if queue_key(c.thread_id, c.user_id, c.agent_name) == key and c.bypass_watermark == bypass_watermark), None, ) - merged_correction_detected = correction_detected or (existing_context.correction_detected if existing_context is not None else False) - merged_reinforcement_detected = reinforcement_detected or (existing_context.reinforcement_detected if existing_context is not None else False) + # Backpressure: once depth reaches the cap, reject NEW non-signal normal + # items. Same-key updates merge (do not grow depth); signal-bearing items + # and emergency (bypass) flushes are always admitted. Signals capture + # important memories, and the emergency path captures messages about to + # be removed by summarization -- neither can be re-fed next turn, so + # shedding them under load would lose data rather than merely defer it. + max_depth = self._config.queue_max_depth + if max_depth > 0 and not bypass_watermark and not signals and existing is None and len(self._items) >= max_depth: + raise QueueFull(f"memory update queue is full (depth {len(self._items)} >= {max_depth}); non-signal update for thread {thread_id} rejected") + + # Merge by signal union: a signal seen on any update for this key stays. + merged_signals = signals | (existing.signals if existing is not None else frozenset()) context = ConversationContext( thread_id=thread_id, messages=messages, agent_name=agent_name, user_id=user_id, trace_id=trace_id, - correction_detected=merged_correction_detected, - reinforcement_detected=merged_reinforcement_detected, + signals=merged_signals, + bypass_watermark=bypass_watermark, ) - - self._queue = [context for context in self._queue if self._queue_key(context.thread_id, context.user_id, context.agent_name) != queue_key] - self._queue.append(context) + if existing is not None: + self._items = [c for c in self._items if not (queue_key(c.thread_id, c.user_id, c.agent_name) == key and c.bypass_watermark == bypass_watermark)] + self._items.append(context) + return context def _reset_timer(self) -> None: """Reset the debounce timer.""" @@ -196,13 +242,13 @@ class MemoryUpdateQueue: self._reprocess_pending = True return - if not self._queue: + if not self._items: return self._processing = True self._processing_thread = threading.current_thread() - contexts_to_process = self._queue.copy() - self._queue.clear() + contexts_to_process = self._items + self._items = [] self._timer = None logger.info("Processing %d queued memory updates", len(contexts_to_process)) @@ -217,10 +263,10 @@ class MemoryUpdateQueue: messages=context.messages, thread_id=context.thread_id, agent_name=context.agent_name, - correction_detected=context.correction_detected, - reinforcement_detected=context.reinforcement_detected, + signals=context.signals, user_id=context.user_id, trace_id=context.trace_id, + bypass_watermark=context.bypass_watermark, ) if success: succeeded += 1 @@ -248,9 +294,16 @@ class MemoryUpdateQueue: with self._lock: self._processing = False self._processing_thread = None + # Reschedule inside the lock: ``_schedule_timer`` read-cancels- + # reassigns ``self._timer`` non-atomically, and a concurrent + # ``add``'s ``_reset_timer`` (also under the lock) touches the + # same field. Holding the lock makes the reschedule atomic w.r.t. + # ``add``. ``_schedule_timer`` only calls ``Timer.start()`` (no + # synchronous lock acquisition), so this cannot deadlock. if self._reprocess_pending: self._reprocess_pending = False - if self._queue: + if self._items: + # New work arrived mid-processing: re-run immediately. self._schedule_timer(0) def flush(self, *, skip_inter_item_delay: bool = False) -> None: @@ -306,7 +359,7 @@ class MemoryUpdateQueue: # (1) Wait for an in-flight _process_queue first (bounded). Otherwise # flush() would see _processing=True, no-op, and we would report # success while that worker is still mid-LLM-call on a daemon thread - # that exit will kill — losing the contexts it already pulled out. + # that exit will kill - losing the contexts it already pulled out. with self._lock: in_flight = self._processing_thread if in_flight is not None: @@ -356,7 +409,7 @@ class MemoryUpdateQueue: if self._timer is not None: self._timer.cancel() self._timer = None - self._queue.clear() + self._items = [] self._processing = False self._processing_thread = None self._reprocess_pending = False @@ -365,7 +418,7 @@ class MemoryUpdateQueue: def pending_count(self) -> int: """Get the number of pending updates.""" with self._lock: - return len(self._queue) + return len(self._items) @property def is_processing(self) -> bool: diff --git a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/updater.py b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/updater.py index 608edbab3..7f81c0e3e 100644 --- a/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/updater.py +++ b/backend/packages/harness/deerflow/agents/memory/backends/deermem/deermem/core/updater.py @@ -10,10 +10,12 @@ import logging import math import re import uuid +from collections import OrderedDict from datetime import UTC, datetime, timedelta from typing import Any from ..config import DeerMemConfig +from .message_processing import detect_signals, extract_message_text from .prompt import ( format_conversation_for_update, load_prompt, @@ -610,6 +612,52 @@ def _escape_memory_for_prompt(memory: Any) -> Any: return memory +def _memory_with_manual_markers(memory: Any) -> Any: + """Return a deep copy of ``memory`` with ``[MANUAL]`` prefixed onto the + content of manually-authored facts (``source.type == "manual"``). + + The marker is a prompt-only signal that tells the extraction LLM a fact is a + high-trust user edit; the persisted memory is untouched (this copy is only + fed to the prompt). Idempotent: a fact already carrying the prefix is not + double-marked. + """ + display = copy.deepcopy(memory) + if not isinstance(display, dict): + return display + for fact in display.get("facts", []): + if not isinstance(fact, dict): + continue + src = fact.get("source") + src_type = src.get("type") if isinstance(src, dict) else None + if src_type == "manual": + content = fact.get("content") + if isinstance(content, str) and not content.startswith("[MANUAL]"): + fact["content"] = "[MANUAL] " + content + return display + + +def _message_identity(msg: Any) -> tuple[str, ...] | None: + """Return a hashable identity for ``msg`` for watermark tracking. + + The watermark is content/identity based rather than index based so it stays + valid when summarization removes the conversation front (an index watermark + would point at the wrong message after a front removal, silently skipping + un-extracted turns). Prefers the langgraph message ``id`` (unique, robust to + duplicate content); falls back to ``(type, content)`` when no id is set + (e.g. plain ``HumanMessage(content=...)`` in tests). Returns ``None`` for a + message with neither id nor extractable text -- the caller then feeds the + full list, which is safe over-extraction and never loss. + """ + mid = getattr(msg, "id", None) + if isinstance(mid, str) and mid: + return ("id", mid) + text = extract_message_text(msg) + if not text: + return None + msg_type = getattr(msg, "type", "") or "" + return ("content", msg_type, text) + + class MemoryUpdater: """Updates memory using LLM based on conversation context.""" @@ -632,6 +680,12 @@ class MemoryUpdater: self._llm = llm self._prompts_dir = prompts_dir self._callbacks = callbacks + # Watermark: last-extracted message identity per (thread_id, user_id, + # agent_name), held in memory so a restart re-extracts one batch. The + # cache is a bounded LRU (config.watermark_max_keys) so a long-lived + # gateway handling many threads cannot grow it without limit; a dropped + # key re-extracts one batch on that thread's next turn. + self._watermarks: OrderedDict[tuple[str | None, str | None, str | None], tuple[str, ...] | None] = OrderedDict() # ── Data access + fact CRUD (formerly module-level functions; use self._storage) ── @@ -917,37 +971,45 @@ class MemoryUpdater: raise OSError(f"Failed to save memory data after updating fact '{fact_id}'") return updated_memory - def _build_correction_hint( - self, - correction_detected: bool, - reinforcement_detected: bool, - ) -> str: - """Build optional prompt hints for correction and reinforcement signals.""" - correction_hint = "" - if correction_detected: - correction_hint = ( + def _build_signal_hints(self, signals: frozenset[str]) -> str: + """Build optional prompt hints for the detected signal classes. + + Each present signal contributes one instruction nudging the extraction + LLM toward the right category and confidence. The variable is still + rendered into the template's ``{correction_hint}`` slot (the name is + historical -- it now carries the full signal-hint set, plus the manual + fact note appended by :meth:`_prepare_update_prompt`). + """ + hints: list[str] = [] + if "correction" in signals: + hints.append( "IMPORTANT: Explicit correction signals were detected in this conversation. " "Pay special attention to what the agent got wrong, what the user corrected, " "and record the correct approach as a fact with category " '"correction" and confidence >= 0.95 when appropriate.' ) - if reinforcement_detected: - reinforcement_hint = ( + if "reinforcement" in signals: + hints.append( "IMPORTANT: Positive reinforcement signals were detected in this conversation. " "The user explicitly confirmed the agent's approach was correct or helpful. " "Record the confirmed approach, style, or preference as a fact with category " '"preference" or "behavior" and confidence >= 0.9 when appropriate.' ) - correction_hint = (correction_hint + "\n" + reinforcement_hint).strip() if correction_hint else reinforcement_hint - - return correction_hint + if "preference" in signals: + hints.append('IMPORTANT: A preference signal was detected. Record the user\'s stated preference or dislike as a fact with category "preference" and high confidence.') + if "identity" in signals: + hints.append('IMPORTANT: An identity signal was detected. Record the user\'s stated role, profession, or background as a fact with category "identity" and high confidence.') + if "goal" in signals: + hints.append('IMPORTANT: A goal signal was detected. Record the user\'s stated objective or intent as a fact with category "goal" and high confidence.') + if "decision" in signals: + hints.append('IMPORTANT: A decision signal was detected. Record the user\'s decision or chosen option as a fact with category "decision" and high confidence.') + return "\n".join(hints) def _prepare_update_prompt( self, messages: list[Any], agent_name: str | None, - correction_detected: bool, - reinforcement_detected: bool, + signals: frozenset[str], user_id: str | None = None, ) -> tuple[dict[str, Any], list[Any]] | None: """Load memory and build the update prompt for a conversation.""" @@ -960,10 +1022,16 @@ class MemoryUpdater: if not conversation_text.strip(): return None - correction_hint = self._build_correction_hint( - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, - ) + correction_hint = self._build_signal_hints(signals) + + # Manual-fact signal: tag high-trust user-authored facts with a [MANUAL] + # prefix in the prompt's current_memory and instruct the model to preserve + # them unless the new conversation is an explicit, unambiguous correction. + display_memory = current_memory + if self._has_manual_facts(current_memory): + display_memory = _memory_with_manual_markers(current_memory) + manual_hint = "NOTE: Facts marked [MANUAL] are high-trust user-authored edits. Update them only when the new conversation is an explicit, unambiguous correction; otherwise preserve them as-is." + correction_hint = (correction_hint + "\n" + manual_hint).strip() if correction_hint else manual_hint # ── Build staleness review section ── staleness_section = "" @@ -986,7 +1054,7 @@ class MemoryUpdater: ) variables = { - "current_memory": json.dumps(_escape_memory_for_prompt(current_memory), indent=2, ensure_ascii=False), + "current_memory": json.dumps(_escape_memory_for_prompt(display_memory), indent=2, ensure_ascii=False), "conversation": conversation_text, "correction_hint": correction_hint, "staleness_review_section": staleness_section, @@ -995,6 +1063,45 @@ class MemoryUpdater: prompt = load_prompt_messages("memory_update", variables, agent_name=agent_name, prompts_dir=self._prompts_dir) return current_memory, prompt + def _has_manual_facts(self, memory: dict[str, Any]) -> bool: + """Return whether ``memory`` contains any user-authored (manual) fact.""" + return any(isinstance(f, dict) and isinstance(f.get("source"), dict) and f.get("source", {}).get("type") == "manual" for f in memory.get("facts", [])) + + def _emit_extraction_metrics( + self, + metrics: dict[str, Any], + *, + thread_id: str | None, + user_id: str | None, + trace_id: str | None, + model_name: str | None, + response: Any, + success: bool, + ) -> None: + """Invoke the post-extraction observability callback (Langfuse span etc.). + + No-op when ``extraction_callback`` is unset (default). Exceptions from + the callback are logged and swallowed so observability never breaks the + update path. + """ + callback = self._config.extraction_callback + if callback is None: + return + usage = getattr(response, "usage_metadata", None) + payload: dict[str, Any] = { + "thread_id": thread_id, + "user_id": user_id, + "trace_id": trace_id, + "model_name": model_name, + "success": success, + "token_usage": usage if isinstance(usage, dict) else None, + } + payload.update(metrics) + try: + callback(payload) + except Exception: + logger.warning("extraction_callback raised; ignoring", exc_info=True) + def _finalize_update( self, current_memory: dict[str, Any], @@ -1002,9 +1109,18 @@ class MemoryUpdater: thread_id: str | None, agent_name: str | None, user_id: str | None = None, + *, + metrics: dict[str, Any] | None = None, ) -> bool: """Parse the model response, apply updates, and persist memory.""" update_data = _parse_memory_update_response(response_content) + if metrics is not None: + extracted = update_data.get("newFacts", []) + extracted_list = extracted if isinstance(extracted, list) else [] + metrics["facts_extracted"] = len(extracted_list) + # facts_passed_confidence / rejected_low_confidence are populated + # inside _apply_updates at the real confidence-filter site, so the + # metric tracks the actual filter rather than a re-derived copy here. if getattr(type(self._storage), "apply_changes", None) is not MemoryStorage.apply_changes: for attempt in range(3): # Deep-copy before in-place mutation so a failed commit cannot @@ -1012,7 +1128,7 @@ class MemoryUpdater: # complete extraction result is reapplied to a fresh document; # its trim/consolidation/delete decisions are snapshot-wide and # must never be replayed as disjoint point writes. - updated_memory = self._apply_updates(copy.deepcopy(current_memory), update_data, thread_id) + updated_memory = self._apply_updates(copy.deepcopy(current_memory), update_data, thread_id, metrics=metrics) updated_memory = _strip_upload_mentions_from_memory(updated_memory) current_by_id = {str(fact.get("id")): fact for fact in current_memory.get("facts", [])} updated_by_id = {str(fact.get("id")): fact for fact in updated_memory.get("facts", [])} @@ -1044,7 +1160,7 @@ class MemoryUpdater: raise AssertionError("bounded extracted-update retry did not return or raise") # Deep-copy before in-place mutation so a subsequent save() failure # cannot corrupt the still-cached original object reference. - updated_memory = self._apply_updates(copy.deepcopy(current_memory), update_data, thread_id) + updated_memory = self._apply_updates(copy.deepcopy(current_memory), update_data, thread_id, metrics=metrics) updated_memory = _strip_upload_mentions_from_memory(updated_memory) return self._storage.save( updated_memory, @@ -1058,10 +1174,11 @@ class MemoryUpdater: messages: list[Any], thread_id: str | None = None, agent_name: str | None = None, - correction_detected: bool = False, - reinforcement_detected: bool = False, + signals: frozenset[str] = frozenset(), user_id: str | None = None, trace_id: str | None = None, + *, + bypass_watermark: bool = False, ) -> bool: """Update memory asynchronously by delegating to the sync path. @@ -1076,10 +1193,10 @@ class MemoryUpdater: messages=messages, thread_id=thread_id, agent_name=agent_name, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, + signals=signals, user_id=user_id, trace_id=trace_id, + bypass_watermark=bypass_watermark, ) def _do_update_memory_sync( @@ -1087,10 +1204,11 @@ class MemoryUpdater: messages: list[Any], thread_id: str | None = None, agent_name: str | None = None, - correction_detected: bool = False, - reinforcement_detected: bool = False, + signals: frozenset[str] = frozenset(), user_id: str | None = None, trace_id: str | None = None, + *, + bypass_watermark: bool = False, ) -> bool: """Pure-sync memory update; bind ``trace_id`` into the request-trace ContextVar for the worker thread, then delegate to the impl. @@ -1110,45 +1228,128 @@ class MemoryUpdater: messages=messages, thread_id=thread_id, agent_name=agent_name, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, + signals=signals, user_id=user_id, trace_id=trace_id, + bypass_watermark=bypass_watermark, ) return self._do_update_memory_sync_impl( messages=messages, thread_id=thread_id, agent_name=agent_name, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, + signals=signals, user_id=user_id, trace_id=trace_id, + bypass_watermark=bypass_watermark, ) + def _watermark_get(self, key: tuple[str | None, str | None, str | None]) -> tuple[str, ...] | None: + """Return the watermark for ``key``, marking it most-recently-used. + + Uses key presence (not value truthiness) so a stored ``None`` identity + still counts as a live entry for LRU ordering. + """ + if key not in self._watermarks: + return None + self._watermarks.move_to_end(key) + return self._watermarks[key] + + def _watermark_set( + self, + key: tuple[str | None, str | None, str | None], + value: tuple[str, ...] | None, + ) -> None: + """Store ``value`` for ``key``, evicting the least-recently-used entry + when the bounded LRU cache exceeds ``config.watermark_max_keys``. + + A dropped key is safe: the next turn for that thread finds no watermark + and re-extracts one batch (the documented restart behavior). ``0`` = + unbounded (no eviction). + """ + self._watermarks[key] = value + self._watermarks.move_to_end(key) + cap = self._config.watermark_max_keys + if cap > 0 and len(self._watermarks) > cap: + self._watermarks.popitem(last=False) + + def _feed_after_watermark( + self, + watermark_key: tuple[str | None, str | None, str | None], + messages: list[Any], + ) -> list[Any]: + """Return the slice of ``messages`` not yet extracted. + + The watermark stores the identity of the last-extracted message (see + :func:`_message_identity`). If that message is still present, everything + *after* it is fed; if it is absent (front removed by summarization, or + the first-ever extraction for this key) the full list is fed. Re-feeding + is safe over-extraction -- it never skips a turn, which is the only + failure direction that would lose facts. + """ + last_id = self._watermark_get(watermark_key) + if last_id is None: + return messages + for i, msg in enumerate(messages): + if _message_identity(msg) == last_id: + return messages[i + 1 :] + return messages + def _do_update_memory_sync_impl( self, messages: list[Any], thread_id: str | None = None, agent_name: str | None = None, - correction_detected: bool = False, - reinforcement_detected: bool = False, + signals: frozenset[str] = frozenset(), user_id: str | None = None, trace_id: str | None = None, + *, + bypass_watermark: bool = False, ) -> bool: """Pure-sync memory update using ``model.invoke()``. Uses the *sync* LLM call path so no event loop is created. This guarantees that the langchain provider's globally cached async httpx ``AsyncClient`` / connection pool (the one shared with the - lead agent) is never touched — no cross-loop connection reuse is + lead agent) is never touched - no cross-loop connection reuse is possible. + + Watermark: the middleware passes the full conversation each turn, so + without skipping already-extracted turns every update re-feeds old + messages. The watermark stores the identity of the last-extracted + message (content/id based, in-memory only) so it stays correct when + summarization removes the conversation front; a restart loses it and + re-extracts one batch. ``bypass_watermark`` is set by the emergency + (summarization) flush path: the subset it carries is a one-shot + "extract before removal" snapshot, so it is fed in full and does not + read or advance the conversation watermark (advancing it from the + subset's own length would regress the watermark and skip un-extracted + tail turns on the next normal feed). """ + metrics: dict[str, Any] = {} + response: Any = None + model_name: str | None = None + success = False + attempted = False try: + watermark_key = (thread_id, user_id, agent_name) + if bypass_watermark: + # Emergency flush: extract the carried subset in full. + feed_messages = messages + else: + feed_messages = self._feed_after_watermark(watermark_key, messages) + if not feed_messages: + logger.debug("Memory update skipped: no new messages since watermark (thread=%s)", thread_id) + return True + # Re-detect signals on the post-watermark feed so extraction hints + # reference only turns the LLM will actually see. The admission-time + # ``signals`` (detected on the full conversation in DeerMem) already + # served their purpose (backpressure admission at enqueue); the hint + # is a soft nudge and must not point at turns the watermark excluded. + feed_signals = detect_signals(feed_messages, patterns_dir=self._config.patterns_dir) prepared = self._prepare_update_prompt( - messages=messages, + messages=feed_messages, agent_name=agent_name, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, + signals=feed_signals, user_id=user_id, ) if prepared is None: @@ -1173,30 +1374,56 @@ class MemoryUpdater: model_name=model_name, ) logger.info("Invoking memory-update LLM (thread=%s trace_id=%s)", thread_id, trace_id) + attempted = True response = model.invoke(prompt, config=invoke_config) - return self._finalize_update( + success = self._finalize_update( current_memory=current_memory, response_content=response.content, thread_id=thread_id, agent_name=agent_name, user_id=user_id, + metrics=metrics, ) + if success and not bypass_watermark: + # Advance the watermark to the last message fed (the feed is a + # suffix, so this is messages[-1]). Skipped on the emergency + # path -- the subset's last message is older than the + # conversation's latest, so advancing from it would regress. + self._watermark_set(watermark_key, _message_identity(messages[-1])) + return success except json.JSONDecodeError as e: logger.warning("Failed to parse LLM response for memory update: %s", e) return False except Exception as e: logger.exception("Memory update failed: %s", e) return False + finally: + # Emit metrics even when _finalize_update (or invoke) raises, so the + # observability callback sees exception failures (parse errors, + # storage errors after retry) rather than only the happy path. The + # pre-attempt early returns (no new messages, empty conversation, no + # model) do not emit, matching the prior behavior. + if attempted: + self._emit_extraction_metrics( + metrics, + thread_id=thread_id, + user_id=user_id, + trace_id=trace_id, + model_name=model_name, + response=response, + success=success, + ) def update_memory( self, messages: list[Any], thread_id: str | None = None, agent_name: str | None = None, - correction_detected: bool = False, - reinforcement_detected: bool = False, + signals: frozenset[str] = frozenset(), user_id: str | None = None, trace_id: str | None = None, + *, + bypass_watermark: bool = False, ) -> bool: """Synchronously update memory using the sync LLM path. @@ -1213,12 +1440,15 @@ class MemoryUpdater: messages: List of conversation messages. thread_id: Optional thread ID for tracking source. agent_name: If provided, updates per-agent memory. If None, updates global memory. - correction_detected: Whether recent turns include an explicit correction signal. - reinforcement_detected: Whether recent turns include a positive reinforcement signal. + signals: Signal classes detected in the conversation (correction / + reinforcement / preference / ...), used as extraction hints. user_id: If provided, scopes memory to a specific user. Returns: - True if update was successful, False otherwise. + True if the update persisted. False on any failure (no content, + unparseable response, LLM error); failures are swallowed (best-effort) + -- a failed update is re-fed on the next conversation turn because the + watermark does not advance on failure. """ try: loop = asyncio.get_running_loop() @@ -1232,10 +1462,10 @@ class MemoryUpdater: messages=messages, thread_id=thread_id, agent_name=agent_name, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, + signals=signals, user_id=user_id, trace_id=trace_id, + bypass_watermark=bypass_watermark, ) return future.result() except Exception: @@ -1246,10 +1476,10 @@ class MemoryUpdater: messages=messages, thread_id=thread_id, agent_name=agent_name, - correction_detected=correction_detected, - reinforcement_detected=reinforcement_detected, + signals=signals, user_id=user_id, trace_id=trace_id, + bypass_watermark=bypass_watermark, ) def _apply_updates( @@ -1257,6 +1487,8 @@ class MemoryUpdater: current_memory: dict[str, Any], update_data: dict[str, Any], thread_id: str | None = None, + *, + metrics: dict[str, Any] | None = None, ) -> dict[str, Any]: """Apply LLM-generated updates to memory. @@ -1264,6 +1496,11 @@ class MemoryUpdater: current_memory: Current memory data. update_data: Updates from LLM. thread_id: Optional thread ID for tracking. + metrics: Optional observability dict. When provided, populated with + ``facts_passed_confidence`` / ``rejected_low_confidence`` counted + at the real confidence-filter site below (the only acceptance + gate for new facts), so the metric cannot drift from the actual + filter the way a re-derived count in the caller could. Returns: Updated memory data. @@ -1398,9 +1635,18 @@ class MemoryUpdater: # Creation-time lifetime cap shared with the consolidation path below, so # both fact-creation sites apply the identical bound in one place. creation_cap = int(config.staleness_age_days * config.staleness_max_lifetime_multiplier) + # Counted at the confidence-gate site (the only real accept filter for new + # facts) so the ``facts_passed_confidence`` metric mirrors the actual + # filter and cannot drift from it. Facts below the threshold are the + # reject count; duplicate / empty / over-cap facts that pass the + # threshold are still counted here -- the metric is a confidence-gate + # signal (the host's rejection-rate warning monitors confidence + # filtering, not dedup / over-cap), not a persisted-fact count. + passed_threshold = 0 for fact in new_facts: confidence = fact.get("confidence", 0.5) if confidence >= config.fact_confidence_threshold: + passed_threshold += 1 raw_content = fact.get("content", "") if not isinstance(raw_content, str): continue @@ -1439,6 +1685,10 @@ class MemoryUpdater: if fact_key is not None: existing_fact_keys.add(fact_key) + if metrics is not None: + metrics["facts_passed_confidence"] = passed_threshold + metrics["rejected_low_confidence"] = len(new_facts) - passed_threshold + # Enforce max facts limit (coerced confidence -- see _trim_facts_to_max). current_memory["facts"] = _trim_facts_to_max(current_memory["facts"], config.max_facts) diff --git a/backend/packages/harness/deerflow/agents/memory/manager.py b/backend/packages/harness/deerflow/agents/memory/manager.py index fd88495e7..dd6d7c7f8 100644 --- a/backend/packages/harness/deerflow/agents/memory/manager.py +++ b/backend/packages/harness/deerflow/agents/memory/manager.py @@ -656,6 +656,51 @@ def _host_default_llm() -> Any: return None +def _host_default_extraction_callback(payload: Any) -> None: + """deer-flow default for DeerMem's ``extraction_callback`` slot. + + Logs post-extraction metrics (token usage, facts passing/rejected by the + confidence filter, gate rejection rate) for ops observability, and flags a + high rejection rate + (>60%) so a prompt/threshold regression is visible without inspecting every + trace. A Langfuse-aware callback can replace this to emit a dedicated + extraction span; the metrics keys are stable for that handoff. Exceptions + are never raised (the DeerMem side already wraps the call). + """ + if not isinstance(payload, dict): + return + extracted = payload.get("facts_extracted") + passed_confidence = payload.get("facts_passed_confidence") + rejected = payload.get("rejected_low_confidence", 0) + thread_id = payload.get("thread_id") + model_name = payload.get("model_name") + if isinstance(extracted, int) and isinstance(passed_confidence, int) and extracted > 0: + rejection_rate = (extracted - passed_confidence) / extracted + logger.info( + "Memory extraction metrics: thread=%s model=%s extracted=%d passed_confidence=%d rejected=%d rejection_rate=%.2f", + thread_id, + model_name, + extracted, + passed_confidence, + rejected, + rejection_rate, + ) + if rejection_rate > 0.6: + logger.warning( + "Memory extraction rejection rate %.0f%% exceeds 60%% - review extraction prompt / confidence threshold (thread=%s)", + rejection_rate * 100, + thread_id, + ) + else: + logger.info( + "Memory extraction metrics: thread=%s model=%s success=%s token_usage=%s", + thread_id, + model_name, + payload.get("success"), + payload.get("token_usage"), + ) + + def _collect_host_hooks() -> dict[str, Any]: """Provide host hook callables for backends to consume in ``from_config``. @@ -674,6 +719,7 @@ def _collect_host_hooks() -> dict[str, Any]: "should_keep_hidden_message": _host_default_should_keep_hidden_message, "trace_context_manager": request_trace_context, "host_llm_factory": _host_default_llm, + "extraction_callback": _host_default_extraction_callback, } diff --git a/backend/tests/test_deermem_self_contained.py b/backend/tests/test_deermem_self_contained.py index c34a71143..54e2d6eb3 100644 --- a/backend/tests/test_deermem_self_contained.py +++ b/backend/tests/test_deermem_self_contained.py @@ -55,6 +55,33 @@ def _deermem_with_fake_llm(backend_config=None, payload=None, callbacks=None) -> return dm +def test_add_swallows_queue_full_so_backpressure_does_not_break_caller(deermem_data_dir, caplog) -> None: + """Regression: QueueFull raised under backpressure is caught in + DeerMem.add (the backend owns the queue, so it owns the degradation) so + memory backpressure degrades to "update skipped" instead of propagating into + MemoryMiddleware.after_agent and breaking the agent run -- peer middlewares + self-guard the same way.""" + import logging + + dm = _deermem_with_fake_llm(backend_config={"storage_path": str(deermem_data_dir), "queue_max_depth": 1}) + # Stop the debounce timer so enqueued items stay pending (the cap persists + # across the second add instead of being drained by a timer fire). + dm._queue._schedule_timer = lambda *a, **k: None + + conv = [HumanMessage("Please explain quantum computing in detail"), AIMessage("Quantum computing uses qubits and superposition.")] + # First add fills the queue to its depth cap (non-signal, new key). + dm.add("thread-A", conv, agent_name="lead_agent", user_id="u") + assert dm._queue.pending_count == 1 + + # Second add for a different key hits the cap -> QueueFull internally. It + # must be caught: no exception escapes DeerMem.add. + with caplog.at_level(logging.WARNING, logger="deerflow.agents.memory.backends.deermem.deer_mem"): + dm.add("thread-B", conv, agent_name="lead_agent", user_id="u") + assert "rejected under backpressure" in caplog.text + # thread-B was rejected (not enqueued); only thread-A remains. + assert dm._queue.pending_count == 1 + + def test_di_construction_owns_dependencies(): dm = DeerMem(backend_config={"max_facts": 50, "storage_path": "/tmp/x"}) assert dm._config.max_facts == 50 diff --git a/backend/tests/test_memory_consolidation.py b/backend/tests/test_memory_consolidation.py index fd6bf470b..a2c92f988 100644 --- a/backend/tests/test_memory_consolidation.py +++ b/backend/tests/test_memory_consolidation.py @@ -1472,8 +1472,7 @@ class TestPrepareUpdatePromptConsolidation: result = updater._prepare_update_prompt( messages=[msg], agent_name=None, - correction_detected=False, - reinforcement_detected=False, + signals=frozenset(), ) assert result is not None @@ -1497,8 +1496,7 @@ class TestPrepareUpdatePromptConsolidation: result = updater._prepare_update_prompt( messages=[msg], agent_name=None, - correction_detected=False, - reinforcement_detected=False, + signals=frozenset(), ) assert result is not None @@ -1521,8 +1519,7 @@ class TestPrepareUpdatePromptConsolidation: result = updater._prepare_update_prompt( messages=[msg], agent_name=None, - correction_detected=False, - reinforcement_detected=False, + signals=frozenset(), ) assert result is not None diff --git a/backend/tests/test_memory_manager_pluggable.py b/backend/tests/test_memory_manager_pluggable.py index e7b4ce33a..7de42b9c2 100644 --- a/backend/tests/test_memory_manager_pluggable.py +++ b/backend/tests/test_memory_manager_pluggable.py @@ -193,7 +193,7 @@ def test_deermem_shutdown_flush_drains_a_pending_update() -> None: mock_updater = MagicMock() mock_updater.update_memory.return_value = True deermem._queue._updater = mock_updater - deermem._queue._queue = [ConversationContext(thread_id=f"t{i}", messages=["m"], agent_name="lead_agent") for i in range(3)] + deermem._queue._items = [ConversationContext(thread_id=f"t{i}", messages=["m"], agent_name="lead_agent") for i in range(3)] assert deermem.shutdown_flush(5.0) is True assert deermem._queue.pending_count == 0 assert mock_updater.update_memory.call_count == 3 diff --git a/backend/tests/test_memory_queue.py b/backend/tests/test_memory_queue.py index 13d3c5e34..5c5b9de28 100644 --- a/backend/tests/test_memory_queue.py +++ b/backend/tests/test_memory_queue.py @@ -13,20 +13,20 @@ def _queue(updater: MagicMock | None = None) -> MemoryUpdateQueue: def test_queue_add_preserves_existing_correction_flag_for_same_thread() -> None: queue = _queue() - with patch.object(queue, "_reset_timer"): - queue.add(thread_id="thread-1", messages=["first"], correction_detected=True) - queue.add(thread_id="thread-1", messages=["second"], correction_detected=False) + with patch.object(queue, "_schedule_timer"): + queue.add(thread_id="thread-1", messages=["first"], signals=frozenset({"correction"})) + queue.add(thread_id="thread-1", messages=["second"], signals=frozenset()) - assert len(queue._queue) == 1 - assert queue._queue[0].messages == ["second"] - assert queue._queue[0].correction_detected is True + assert len(queue._items) == 1 + assert queue._items[0].messages == ["second"] + assert "correction" in queue._items[0].signals def test_process_queue_forwards_correction_flag_to_updater() -> None: mock_updater = MagicMock() mock_updater.update_memory.return_value = True queue = _queue(mock_updater) - queue._queue = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent", correction_detected=True)] + queue._items = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent", signals=frozenset({"correction"}))] queue._process_queue() @@ -34,29 +34,29 @@ def test_process_queue_forwards_correction_flag_to_updater() -> None: messages=["conversation"], thread_id="thread-1", agent_name="lead_agent", - correction_detected=True, - reinforcement_detected=False, + signals=frozenset({"correction"}), user_id=None, trace_id=None, + bypass_watermark=False, ) def test_queue_add_preserves_existing_reinforcement_flag_for_same_thread() -> None: queue = _queue() - with patch.object(queue, "_reset_timer"): - queue.add(thread_id="thread-1", messages=["first"], reinforcement_detected=True) - queue.add(thread_id="thread-1", messages=["second"], reinforcement_detected=False) + with patch.object(queue, "_schedule_timer"): + queue.add(thread_id="thread-1", messages=["first"], signals=frozenset({"reinforcement"})) + queue.add(thread_id="thread-1", messages=["second"], signals=frozenset()) - assert len(queue._queue) == 1 - assert queue._queue[0].messages == ["second"] - assert queue._queue[0].reinforcement_detected is True + assert len(queue._items) == 1 + assert queue._items[0].messages == ["second"] + assert "reinforcement" in queue._items[0].signals def test_process_queue_forwards_reinforcement_flag_to_updater() -> None: mock_updater = MagicMock() mock_updater.update_memory.return_value = True queue = _queue(mock_updater) - queue._queue = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent", reinforcement_detected=True)] + queue._items = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent", signals=frozenset({"reinforcement"}))] queue._process_queue() @@ -64,10 +64,10 @@ def test_process_queue_forwards_reinforcement_flag_to_updater() -> None: messages=["conversation"], thread_id="thread-1", agent_name="lead_agent", - correction_detected=False, - reinforcement_detected=True, + signals=frozenset({"reinforcement"}), user_id=None, trace_id=None, + bypass_watermark=False, ) @@ -99,7 +99,7 @@ def test_add_nowait_cancels_existing_timer_and_starts_immediate_timer() -> None: existing_timer.cancel.assert_called_once_with() timer_cls.assert_called_once_with(0, queue._process_queue) assert queue.pending_count == 1 - assert queue._queue[0].agent_name == "lead-agent" + assert queue._items[0].agent_name == "lead-agent" assert created_timer.daemon is True created_timer.start.assert_called_once_with() @@ -127,14 +127,14 @@ def test_finishing_worker_reschedules_once_when_reprocess_pending() -> None: schedules exactly one follow-up run (not a per-arrival timer spin).""" mock_updater = MagicMock() queue = _queue(mock_updater) - queue._queue = [ConversationContext(thread_id="thread-1", messages=["first"], agent_name="lead_agent")] + queue._items = [ConversationContext(thread_id="thread-1", messages=["first"], agent_name="lead_agent")] queue._reprocess_pending = True created_timer = MagicMock() def _enqueue_more_while_processing(**_kwargs) -> bool: # Simulate a new update arriving mid-processing so the finally block sees # remaining work and reschedules exactly once. - queue._queue.append(ConversationContext(thread_id="thread-2", messages=["second"], agent_name="lead_agent")) + queue._items.append(ConversationContext(thread_id="thread-2", messages=["second"], agent_name="lead_agent")) return True mock_updater.update_memory.side_effect = _enqueue_more_while_processing @@ -154,7 +154,7 @@ def test_finishing_worker_does_not_reschedule_when_no_work_remains() -> None: mock_updater = MagicMock() mock_updater.update_memory.return_value = True queue = _queue(mock_updater) - queue._queue = [ConversationContext(thread_id="thread-1", messages=["only"], agent_name="lead_agent")] + queue._items = [ConversationContext(thread_id="thread-1", messages=["only"], agent_name="lead_agent")] queue._reprocess_pending = True with patch("deerflow.agents.memory.backends.deermem.deermem.core.queue.threading.Timer") as timer_cls: @@ -193,19 +193,19 @@ def test_queue_keeps_updates_for_different_agents_in_same_thread() -> None: queue.add(thread_id="thread-1", messages=["agent-b"], agent_name="agent-b") assert queue.pending_count == 2 - assert [context.agent_name for context in queue._queue] == ["agent-a", "agent-b"] + assert [context.agent_name for context in queue._items] == ["agent-a", "agent-b"] def test_queue_still_coalesces_updates_for_same_agent_in_same_thread() -> None: queue = _queue() - with patch.object(queue, "_reset_timer"): - queue.add(thread_id="thread-1", messages=["first"], agent_name="agent-a", correction_detected=True) - queue.add(thread_id="thread-1", messages=["second"], agent_name="agent-a", correction_detected=False) + with patch.object(queue, "_schedule_timer"): + queue.add(thread_id="thread-1", messages=["first"], agent_name="agent-a", signals=frozenset({"correction"})) + queue.add(thread_id="thread-1", messages=["second"], agent_name="agent-a", signals=frozenset()) assert queue.pending_count == 1 - assert queue._queue[0].agent_name == "agent-a" - assert queue._queue[0].messages == ["second"] - assert queue._queue[0].correction_detected is True + assert queue._items[0].agent_name == "agent-a" + assert queue._items[0].messages == ["second"] + assert "correction" in queue._items[0].signals def test_process_queue_updates_different_agents_in_same_thread_separately() -> None: @@ -224,8 +224,8 @@ def test_process_queue_updates_different_agents_in_same_thread_separately() -> N assert mock_updater.update_memory.call_count == 2 mock_updater.update_memory.assert_has_calls( [ - call(messages=["agent-a"], thread_id="thread-1", agent_name="agent-a", correction_detected=False, reinforcement_detected=False, user_id=None, trace_id=None), - call(messages=["agent-b"], thread_id="thread-1", agent_name="agent-b", correction_detected=False, reinforcement_detected=False, user_id=None, trace_id=None), + call(messages=["agent-a"], thread_id="thread-1", agent_name="agent-a", signals=frozenset(), user_id=None, trace_id=None, bypass_watermark=False), + call(messages=["agent-b"], thread_id="thread-1", agent_name="agent-b", signals=frozenset(), user_id=None, trace_id=None, bypass_watermark=False), ] ) @@ -234,7 +234,7 @@ def test_process_queue_forwards_trace_id_to_updater() -> None: mock_updater = MagicMock() mock_updater.update_memory.return_value = True queue = _queue(mock_updater) - queue._queue = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent", trace_id="trace-memory-1")] + queue._items = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent", trace_id="trace-memory-1")] queue._process_queue() @@ -242,10 +242,10 @@ def test_process_queue_forwards_trace_id_to_updater() -> None: messages=["conversation"], thread_id="thread-1", agent_name="lead_agent", - correction_detected=False, - reinforcement_detected=False, + signals=frozenset(), user_id=None, trace_id="trace-memory-1", + bypass_watermark=False, ) @@ -272,7 +272,7 @@ def test_flush_sync_drains_pending_queue_and_returns_true() -> None: mock_updater = MagicMock() mock_updater.update_memory.return_value = True queue = _queue(mock_updater) - queue._queue = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent")] + queue._items = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent")] with ( patch(_QUEUE_MODULE + ".MemoryUpdater", create=True), @@ -286,17 +286,17 @@ def test_flush_sync_drains_pending_queue_and_returns_true() -> None: messages=["conversation"], thread_id="thread-1", agent_name="lead_agent", - correction_detected=False, - reinforcement_detected=False, + signals=frozenset(), user_id=None, trace_id=None, + bypass_watermark=False, ) def test_flush_sync_returns_false_when_flush_exceeds_timeout() -> None: """flush_sync does not block past ``timeout``; a slow flush returns False.""" queue = _queue() - queue._queue = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent")] + queue._items = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent")] release = threading.Event() def _slow_flush() -> None: @@ -379,7 +379,7 @@ def test_flush_sync_returns_false_when_flush_raises() -> None: caller never logs a contradictory 'completed' next to the exception (review comment #2).""" queue = _queue() - queue._queue = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent")] + queue._items = [ConversationContext(thread_id="thread-1", messages=["conversation"], agent_name="lead_agent")] with patch.object(queue, "flush", side_effect=RuntimeError("boom")): completed = queue.flush_sync(timeout=5.0) @@ -393,7 +393,7 @@ def test_flush_sync_skips_inter_item_delay_on_drain_path() -> None: mock_updater = MagicMock() mock_updater.update_memory.return_value = True queue = _queue(mock_updater) - queue._queue = [ConversationContext(thread_id=f"thread-{i}", messages=["conversation"], agent_name="lead_agent") for i in range(3)] + queue._items = [ConversationContext(thread_id=f"thread-{i}", messages=["conversation"], agent_name="lead_agent") for i in range(3)] with patch(_QUEUE_MODULE + ".time.sleep") as mock_sleep: completed = queue.flush_sync(timeout=5.0) diff --git a/backend/tests/test_memory_queue_user_isolation.py b/backend/tests/test_memory_queue_user_isolation.py index 2de38bfcd..033f6f19c 100644 --- a/backend/tests/test_memory_queue_user_isolation.py +++ b/backend/tests/test_memory_queue_user_isolation.py @@ -24,8 +24,8 @@ def test_queue_add_stores_user_id(): q = _queue() with patch.object(q, "_reset_timer"): q.add(thread_id="t1", messages=["msg"], user_id="alice") - assert len(q._queue) == 1 - assert q._queue[0].user_id == "alice" + assert len(q._items) == 1 + assert q._items[0].user_id == "alice" q.clear() @@ -49,8 +49,8 @@ def test_queue_keeps_updates_for_different_users_in_same_thread_and_agent(): q.add(thread_id="main", messages=["bob update"], agent_name="researcher", user_id="bob") assert q.pending_count == 2 - assert [context.user_id for context in q._queue] == ["alice", "bob"] - assert [context.messages for context in q._queue] == [["alice update"], ["bob update"]] + assert [context.user_id for context in q._items] == ["alice", "bob"] + assert [context.messages for context in q._items] == [["alice update"], ["bob update"]] def test_queue_still_coalesces_updates_for_same_user_thread_and_agent(): @@ -60,9 +60,9 @@ def test_queue_still_coalesces_updates_for_same_user_thread_and_agent(): q.add(thread_id="main", messages=["second"], agent_name="researcher", user_id="alice") assert q.pending_count == 1 - assert q._queue[0].messages == ["second"] - assert q._queue[0].user_id == "alice" - assert q._queue[0].agent_name == "researcher" + assert q._items[0].messages == ["second"] + assert q._items[0].user_id == "alice" + assert q._items[0].agent_name == "researcher" def test_add_nowait_keeps_different_users_separate(): @@ -72,4 +72,4 @@ def test_add_nowait_keeps_different_users_separate(): q.add_nowait(thread_id="main", messages=["bob update"], agent_name="researcher", user_id="bob") assert q.pending_count == 2 - assert [context.user_id for context in q._queue] == ["alice", "bob"] + assert [context.user_id for context in q._items] == ["alice", "bob"] diff --git a/backend/tests/test_memory_staleness_review.py b/backend/tests/test_memory_staleness_review.py index 3db4a15dc..1fcabc59e 100644 --- a/backend/tests/test_memory_staleness_review.py +++ b/backend/tests/test_memory_staleness_review.py @@ -1156,8 +1156,7 @@ class TestPrepareUpdatePromptStaleness: result = updater._prepare_update_prompt( messages=[msg], agent_name=None, - correction_detected=False, - reinforcement_detected=False, + signals=frozenset(), ) assert result is not None @@ -1180,8 +1179,7 @@ class TestPrepareUpdatePromptStaleness: result = updater._prepare_update_prompt( messages=[msg], agent_name=None, - correction_detected=False, - reinforcement_detected=False, + signals=frozenset(), ) assert result is not None @@ -1205,8 +1203,7 @@ class TestPrepareUpdatePromptStaleness: result = updater._prepare_update_prompt( messages=[msg], agent_name=None, - correction_detected=False, - reinforcement_detected=False, + signals=frozenset(), ) assert result is not None @@ -1229,8 +1226,7 @@ class TestPrepareUpdatePromptStaleness: result = updater._prepare_update_prompt( messages=[msg], agent_name=None, - correction_detected=False, - reinforcement_detected=False, + signals=frozenset(), ) assert result is not None diff --git a/backend/tests/test_message_processing.py b/backend/tests/test_message_processing.py index d40ad95ca..d36654eb4 100644 --- a/backend/tests/test_message_processing.py +++ b/backend/tests/test_message_processing.py @@ -107,13 +107,12 @@ def test_prepare_update_missing_role_returns_none(tmp_path): assert m._prepare_update([]) is None -def test_prepare_update_returns_3tuple_with_correction_true(tmp_path): +def test_prepare_update_returns_signals_with_correction(tmp_path): m = _make_deermem(tmp_path) r = m._prepare_update([_human("That's wrong, use uv"), _ai("ok")]) - assert r is not None and len(r) == 3 - filtered, corr, rein = r - assert corr is True - assert rein is False + assert r is not None and len(r) == 2 + filtered, signals = r + assert "correction" in signals assert len(filtered) == 2 diff --git a/backend/tests/test_message_processing_signals.py b/backend/tests/test_message_processing_signals.py new file mode 100644 index 000000000..5d32f1203 --- /dev/null +++ b/backend/tests/test_message_processing_signals.py @@ -0,0 +1,139 @@ +"""Tests for message-processing signal detection and trivial filtering. + +Pins three behaviors: ``detect_signals`` recognizes all six signal classes +(correction, reinforcement, preference, identity, goal, decision); +``filter_trivial`` drops pure-ack turns and their replies while keeping +substantive turns; and ``_prepare_update`` returns the full signal set as a +``frozenset`` (not just correction/reinforcement), so every detected class +flows through to the extraction prompt. +""" + +from __future__ import annotations + +from langchain_core.messages import AIMessage, HumanMessage + +from deerflow.agents.memory.backends.deermem.deer_mem import DeerMem +from deerflow.agents.memory.backends.deermem.deermem.core.message_processing import ( + detect_signals, + extract_message_text, + filter_trivial, +) + + +def _human(text: str) -> HumanMessage: + return HumanMessage(content=text) + + +def _ai(text: str) -> AIMessage: + return AIMessage(content=text) + + +# ── detect_signals: the 6 signal classes ─────────────────────────────────── + + +def test_detect_signals_correction() -> None: + assert "correction" in detect_signals([_human("That's wrong, use uv"), _ai("ok")]) + + +def test_detect_signals_reinforcement() -> None: + assert "reinforcement" in detect_signals([_human("perfect, exactly right"), _ai("ok")]) + + +def test_detect_signals_preference() -> None: + assert "preference" in detect_signals([_human("I prefer uv over pip"), _ai("ok")]) + + +def test_detect_signals_identity() -> None: + assert "identity" in detect_signals([_human("I am an engineer"), _ai("ok")]) + + +def test_detect_signals_goal() -> None: + assert "goal" in detect_signals([_human("I plan to migrate to uv"), _ai("ok")]) + + +def test_detect_signals_decision() -> None: + assert "decision" in detect_signals([_human("let's go with uv"), _ai("ok")]) + + +def test_detect_signals_none_for_substantive_turn() -> None: + assert detect_signals([_human("what is the weather"), _ai("sunny")]) == set() + + +def test_detect_signals_multiple_classes_in_one_turn() -> None: + # A turn that states both a preference and an identity surfaces both. + signals = detect_signals([_human("I am an engineer and I prefer uv"), _ai("ok")]) + assert "identity" in signals + assert "preference" in signals + + +# ── filter_trivial ───────────────────────────────────────────────────────── + + +def test_filter_trivial_drops_pure_ack_and_its_reply() -> None: + msgs = [_human("嗯"), _ai("thanks"), _human("what next"), _ai("let's see")] + result = filter_trivial(msgs) + # "嗯" + its AI "thanks" dropped; the substantive pair is kept. + assert len(result) == 2 + assert extract_message_text(result[0]) == "what next" + + +def test_filter_trivial_keeps_substantive_message_containing_ok() -> None: + msgs = [_human("use uv to install, ok?"), _ai("done")] + result = filter_trivial(msgs) + assert len(result) == 2 # not dropped: not a whole-message ack + + +def test_filter_trivial_all_trivial_returns_empty() -> None: + msgs = [_human("好的"), _ai("嗯")] + assert filter_trivial(msgs) == [] + + +def test_filter_trivial_tolerates_trailing_punctuation() -> None: + msgs = [_human("ok."), _ai("ok!")] + assert filter_trivial(msgs) == [] + + +def test_filter_trivial_no_patterns_keeps_all() -> None: + msgs = [_human("ok"), _ai("ok")] + assert filter_trivial(msgs, patterns=[]) == msgs + + +# ── _prepare_update: seam-stable 3-tuple projection ──────────────────────── + + +def _make_deermem(tmp_path, **overrides) -> DeerMem: + cfg = {"storage_path": str(tmp_path)} + cfg.update(overrides) + return DeerMem(backend_config=cfg) + + +def test_prepare_update_all_trivial_returns_none(tmp_path) -> None: + m = _make_deermem(tmp_path) + assert m._prepare_update([_human("好的"), _ai("嗯")]) is None + + +def test_prepare_update_returns_correction_signal(tmp_path) -> None: + m = _make_deermem(tmp_path) + r = m._prepare_update([_human("That's wrong, use uv"), _ai("ok")]) + assert r is not None and len(r) == 2 + _filtered, signals = r + assert "correction" in signals + + +def test_prepare_update_returns_reinforcement_signal(tmp_path) -> None: + m = _make_deermem(tmp_path) + r = m._prepare_update([_human("perfect, exactly right"), _ai("ok")]) + assert r is not None and len(r) == 2 + _filtered, signals = r + assert "reinforcement" in signals + + +def test_prepare_update_returns_new_signals_after_swap(tmp_path) -> None: + # After the signals-seam swap, the full signal set flows through (not just + # correction/reinforcement): a preference turn surfaces "preference". + m = _make_deermem(tmp_path) + r = m._prepare_update([_human("I prefer uv over pip"), _ai("ok")]) + assert r is not None and len(r) == 2 + _filtered, signals = r + assert "preference" in signals + assert len(_filtered) == 2 # not trivial -> kept diff --git a/backend/tests/test_updater_truncation.py b/backend/tests/test_updater_truncation.py new file mode 100644 index 000000000..23b48bbf4 --- /dev/null +++ b/backend/tests/test_updater_truncation.py @@ -0,0 +1,54 @@ +"""Tests for the head500 + tail500 message truncation in format_conversation_for_update.""" + +from __future__ import annotations + +from langchain_core.messages import AIMessage, HumanMessage + +from deerflow.agents.memory.backends.deermem.deermem.core.prompt import format_conversation_for_update + + +def test_long_message_keeps_head_and_tail_drops_middle() -> None: + # 600 head chars + 400 middle + 600 tail = 1600 (> 1000 -> truncated). + long_content = "H" * 600 + "M" * 400 + "T" * 600 + result = format_conversation_for_update([HumanMessage(content=long_content)]) + + assert "[truncated]" in result + # The first 500 and last 500 characters survive. + assert "H" * 500 in result + assert "T" * 500 in result + # The middle block is dropped. + assert "M" * 400 not in result + + +def test_message_under_threshold_is_not_truncated() -> None: + result = format_conversation_for_update([HumanMessage(content="a short message")]) + assert "[truncated]" not in result + assert "a short message" in result + + +def test_message_exactly_1000_chars_is_not_truncated() -> None: + # The guard is strictly greater-than 1000, so 1000 chars pass through whole. + result = format_conversation_for_update([HumanMessage(content="x" * 1000)]) + assert "[truncated]" not in result + + +def test_message_1001_chars_is_truncated() -> None: + result = format_conversation_for_update([HumanMessage(content="x" * 1001)]) + assert "[truncated]" in result + + +def test_truncation_then_html_escape_preserves_head_marker() -> None: + # A leading "" must be HTML-escaped after truncation (block-breakout + # defense), and the head is preserved up to the 500-char boundary. + long_content = "" + "y" * 1500 + result = format_conversation_for_update([HumanMessage(content=long_content)]) + assert "<b>" in result + assert "[truncated]" in result + + +def test_truncation_applies_to_ai_messages_too() -> None: + long_content = "A" * 700 + "B" * 700 + result = format_conversation_for_update([AIMessage(content=long_content)]) + assert "[truncated]" in result + assert "A" * 500 in result + assert "B" * 500 in result diff --git a/backend/tests/test_updater_watermark.py b/backend/tests/test_updater_watermark.py new file mode 100644 index 000000000..989394ea4 --- /dev/null +++ b/backend/tests/test_updater_watermark.py @@ -0,0 +1,183 @@ +"""Tests for the in-memory watermark (skip already-extracted messages).""" + +from __future__ import annotations + +import json +from typing import Any + +from langchain_core.messages import AIMessage, HumanMessage + +from deerflow.agents.memory.backends.deermem.deermem.config import DeerMemConfig +from deerflow.agents.memory.backends.deermem.deermem.core.storage import MemoryStorage +from deerflow.agents.memory.backends.deermem.deermem.core.updater import MemoryUpdater, _message_identity + + +class _FakeLLM: + """Returns a canned empty-update response; counts invocations.""" + + def __init__(self) -> None: + self.invoke_count = 0 + self._response = _EmptyResponse() + + def invoke(self, prompt: Any, config: Any = None) -> Any: + self.invoke_count += 1 + return self._response + + +class _EmptyResponse: + content = json.dumps( + { + "user": {}, + "history": {}, + "newFacts": [], + "factsToRemove": [], + "staleFactsToRemove": [], + "staleFactsToExtend": [], + "factsToConsolidate": [], + } + ) + usage_metadata: dict[str, int] | None = None + + +class _FakeStorage(MemoryStorage): + """Minimal in-memory storage stub (load/save) for the save() path.""" + + def __init__(self) -> None: + self.memory: dict[str, Any] = {"version": "2.0", "revision": 0, "user": {}, "history": {}, "facts": []} + + def load(self, agent_name: str | None = None, *, user_id: str | None = None) -> dict[str, Any]: + return json.loads(json.dumps(self.memory)) + + def reload(self, agent_name: str | None = None, *, user_id: str | None = None) -> dict[str, Any]: + return self.load(agent_name, user_id=user_id) + + def save(self, memory_data: dict[str, Any], agent_name: str | None = None, *, user_id: str | None = None, expected_revision: int | None = None) -> bool: + self.memory = json.loads(json.dumps(memory_data)) + return True + + +def _config(**overrides: Any) -> DeerMemConfig: + base: dict[str, Any] = {} + base.update(overrides) + return DeerMemConfig(**base) + + +def _msgs(*texts: str) -> list[Any]: + out: list[Any] = [] + for t in texts: + out.append(HumanMessage(content=t)) + out.append(AIMessage(content=f"reply-{t}")) + return out + + +def test_watermark_skips_already_extracted_messages() -> None: + llm = _FakeLLM() + updater = MemoryUpdater(_config(), _FakeStorage(), llm) + messages = _msgs("first") + + updater.update_memory(messages, thread_id="t1", agent_name="a", user_id="u") + assert llm.invoke_count == 1 + + # Same messages again -> nothing new since the watermark -> skipped, no LLM call. + result = updater.update_memory(messages, thread_id="t1", agent_name="a", user_id="u") + assert result is True + assert llm.invoke_count == 1 + + +def test_watermark_feeds_only_new_messages_on_growth() -> None: + llm = _FakeLLM() + updater = MemoryUpdater(_config(), _FakeStorage(), llm) + messages = _msgs("first") + updater.update_memory(messages, thread_id="t1", agent_name="a", user_id="u") + assert llm.invoke_count == 1 + + # Append a new turn; only the new turn is fed (watermark = prior length). + messages += _msgs("second") + updater.update_memory(messages, thread_id="t1", agent_name="a", user_id="u") + assert llm.invoke_count == 2 + + +def test_watermark_is_per_thread() -> None: + llm = _FakeLLM() + updater = MemoryUpdater(_config(), _FakeStorage(), llm) + messages = _msgs("first") + # Thread t1 extracts; thread t2 has its own watermark (starts at 0). + updater.update_memory(messages, thread_id="t1", agent_name="a", user_id="u") + updater.update_memory(messages, thread_id="t2", agent_name="a", user_id="u") + assert llm.invoke_count == 2 + + +def test_watermark_resets_when_conversation_shrinks() -> None: + llm = _FakeLLM() + updater = MemoryUpdater(_config(), _FakeStorage(), llm) + long_msgs = _msgs("first", "second", "third") + updater.update_memory(long_msgs, thread_id="t1", agent_name="a", user_id="u") + assert llm.invoke_count == 1 + # A shorter message list (e.g. after summarization) must not get stuck at a + # watermark past the end; it re-extracts from the start. + short_msgs = _msgs("only") + updater.update_memory(short_msgs, thread_id="t1", agent_name="a", user_id="u") + assert llm.invoke_count == 2 + + +def test_watermark_front_removal_does_not_skip_pending_tail() -> None: + """Regression: an index watermark skips un-extracted turns when + summarization removes the conversation front. The watermark is + content/identity based, so after a front removal it finds the last-extracted + message at its NEW index and feeds the real pending tail instead of slicing + past the (now shorter) list. + + Setup: 6 messages; pre-set the watermark to msg[3] so msgs[0..3] are + "already extracted" and msgs[4..5] are pending. Summarization then removes + the front pair (msgs[0..1]). The surviving 4-message list must still feed + the pending pair (msgs[4..5]) -- an index watermark (=4) would slice [4:] + on a 4-element list and feed nothing, losing them. + """ + llm = _FakeLLM() + updater = MemoryUpdater(_config(), _FakeStorage(), llm) + msgs = _msgs("a", "b", "c") # [H a, A ra, H b, A rb, H c, A rc] + updater._watermarks[("t1", "u", "a")] = _message_identity(msgs[3]) # ...A rb extracted + surviving = msgs[2:] # summarization removed the front pair + updater.update_memory(surviving, thread_id="t1", agent_name="a", user_id="u") + # The pending tail (H c, A rc) was fed -> exactly one extraction. + assert llm.invoke_count == 1 + + +def test_emergency_flush_bypasses_watermark_and_does_not_regress() -> None: + """Regression: the emergency (summarization) flush path + bypasses the watermark -- it extracts its subset in full and does NOT + advance the conversation watermark (advancing from the subset's own last + message, which is older than the conversation's latest, would regress it + and skip the real tail on the next normal feed).""" + llm = _FakeLLM() + updater = MemoryUpdater(_config(), _FakeStorage(), llm) + msgs = _msgs("a", "b") # [H a, A ra, H b, A rb] + key = ("t1", "u", "a") + updater._watermarks[key] = _message_identity(msgs[1]) # ...A ra extracted; H b, A rb pending + # Emergency flush of the front subset about to be removed. + updater.update_memory(msgs[:2], thread_id="t1", agent_name="a", user_id="u", bypass_watermark=True) + assert llm.invoke_count == 1 # subset extracted in full + # Watermark did not regress to the subset's last message. + assert updater._watermarks[key] == _message_identity(msgs[1]) + # A subsequent normal feed of the full conversation still extracts the tail. + updater.update_memory(msgs, thread_id="t1", agent_name="a", user_id="u") + assert llm.invoke_count == 2 + + +def test_watermark_cache_is_bounded_lru() -> None: + """The watermark cache is a bounded LRU: over capacity it drops the + least-recently-used key, and a dropped key re-extracts one batch on the + next turn for that thread (no loss).""" + llm = _FakeLLM() + updater = MemoryUpdater(_config(watermark_max_keys=2), _FakeStorage(), llm) + msgs = _msgs("first") + # Three distinct threads fill the cache (cap=2); the least-recently-used + # key (t1) is evicted. + updater.update_memory(msgs, thread_id="t1", agent_name="a", user_id="u") + updater.update_memory(msgs, thread_id="t2", agent_name="a", user_id="u") + updater.update_memory(msgs, thread_id="t3", agent_name="a", user_id="u") + assert len(updater._watermarks) == 2 + assert ("t1", "u", "a") not in updater._watermarks # evicted (LRU) + # t1's next turn finds no watermark -> re-extracts one batch (not skipped). + updater.update_memory(msgs, thread_id="t1", agent_name="a", user_id="u") + assert llm.invoke_count == 4 # t1, t2, t3, then t1 re-extract diff --git a/config.example.yaml b/config.example.yaml index a21d8ab3b..476e82123 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -1649,6 +1649,10 @@ memory: file_lock_timeout_seconds: 10 # per-scope cross-process advisory lock timeout (single-machine local filesystem) retrieval_adapter: "" # optional dotted factory(config) supplied by the retrieval module debounce_seconds: 30 # Wait time before processing queued updates + # Backpressure cap on pending items. 0 = unlimited. At the cap, new + # non-signal updates are rejected (QueueFull); signal updates are always + # admitted so important memories are never shed. + queue_max_depth: 1000 model: # LLM for memory extraction; omit all fields = no extraction (non-LLM ops still work; an update raises) # provider: openai # model: gpt-4o-mini @@ -1696,6 +1700,9 @@ memory: consolidation_min_facts: 8 # min facts in one category to trigger review (3-30) consolidation_max_groups_per_cycle: 3 # max groups merged per update cycle (1-10) consolidation_max_sources: 8 # max source facts per consolidation group (2-20) + # extraction_callback is a host-injected post-extraction observability hook + # (token usage, facts accepted/rejected, rejection rate). The factory + # injects a logging default; set programmatically to emit a Langfuse span. # Message processing (externalized patterns / prompt templates): # patterns_dir - dir with correction.yaml / reinforcement.yaml overriding # the bundled signal-detection patterns. None (default) =