From bb0e0316a7e0063be258965cfbaf5b36d5bd47ee Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicol=C3=B2=20Boschi?= Date: Wed, 28 Jan 2026 14:51:25 +0100 Subject: [PATCH] fix: graph endpoint not showing links for observations (#214) --- .../hindsight_api/engine/memory_engine.py | 93 +++++++++++++++++-- .../hindsight_api/engine/reflect/agent.py | 6 ++ .../engine/search/link_expansion_retrieval.py | 1 - hindsight-api/tests/test_consolidation.py | 90 ++++++++++++++++++ hindsight-api/tests/test_reflect_agent.py | 11 +++ 5 files changed, 193 insertions(+), 8 deletions(-) diff --git a/hindsight-api/hindsight_api/engine/memory_engine.py b/hindsight-api/hindsight_api/engine/memory_engine.py index d211960f..bea556b8 100644 --- a/hindsight-api/hindsight_api/engine/memory_engine.py +++ b/hindsight-api/hindsight_api/engine/memory_engine.py @@ -2764,7 +2764,7 @@ class MemoryEngine(MemoryEngineInterface): param_count += 1 units = await conn.fetch( f""" - SELECT id, text, event_date, context, occurred_start, occurred_end, mentioned_at, document_id, chunk_id, fact_type, tags, created_at, proof_count + SELECT id, text, event_date, context, occurred_start, occurred_end, mentioned_at, document_id, chunk_id, fact_type, tags, created_at, proof_count, source_memory_ids FROM {fq_table("memory_units")} {where_clause} ORDER BY mentioned_at DESC NULLS LAST, event_date DESC @@ -2777,7 +2777,18 @@ class MemoryEngine(MemoryEngineInterface): # Get links, filtering to only include links between units of the selected agent # Use DISTINCT ON with LEAST/GREATEST to deduplicate bidirectional links unit_ids = [row["id"] for row in units] - if unit_ids: + unit_id_set = set(unit_ids) + + # Collect source memory IDs from observations + source_memory_ids = [] + for unit in units: + if unit["source_memory_ids"]: + source_memory_ids.extend(unit["source_memory_ids"]) + source_memory_ids = list(set(source_memory_ids)) # Deduplicate + + # Fetch links involving both visible units AND source memories + all_relevant_ids = unit_ids + source_memory_ids + if all_relevant_ids: links = await conn.fetch( f""" SELECT DISTINCT ON (LEAST(ml.from_unit_id, ml.to_unit_id), GREATEST(ml.from_unit_id, ml.to_unit_id), ml.link_type, COALESCE(ml.entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) @@ -2788,14 +2799,69 @@ class MemoryEngine(MemoryEngineInterface): e.canonical_name as entity_name FROM {fq_table("memory_links")} ml LEFT JOIN {fq_table("entities")} e ON ml.entity_id = e.id - WHERE ml.from_unit_id = ANY($1::uuid[]) AND ml.to_unit_id = ANY($1::uuid[]) + WHERE ml.from_unit_id = ANY($1::uuid[]) OR ml.to_unit_id = ANY($1::uuid[]) ORDER BY LEAST(ml.from_unit_id, ml.to_unit_id), GREATEST(ml.from_unit_id, ml.to_unit_id), ml.link_type, COALESCE(ml.entity_id, '00000000-0000-0000-0000-000000000000'::uuid), ml.weight DESC """, - unit_ids, + all_relevant_ids, ) else: links = [] + # Copy links from source memories to observations + # Observations inherit links from their source memories via source_memory_ids + # Build a map from source_id to observation_ids + source_to_observations = {} + for unit in units: + if unit["source_memory_ids"]: + for source_id in unit["source_memory_ids"]: + if source_id not in source_to_observations: + source_to_observations[source_id] = [] + source_to_observations[source_id].append(unit["id"]) + + copied_links = [] + for link in links: + from_id = link["from_unit_id"] + to_id = link["to_unit_id"] + + # Get observations that should inherit this link + from_observations = source_to_observations.get(from_id, []) + to_observations = source_to_observations.get(to_id, []) + + # If from_id is a source memory, copy links to its observations + if from_observations: + for obs_id in from_observations: + # Only include if the target is visible + if to_id in unit_id_set or to_observations: + target = to_observations[0] if to_observations and to_id not in unit_id_set else to_id + if target in unit_id_set: + copied_links.append( + { + "from_unit_id": obs_id, + "to_unit_id": target, + "link_type": link["link_type"], + "weight": link["weight"], + "entity_name": link["entity_name"], + } + ) + + # If to_id is a source memory, copy links to its observations + if to_observations and from_id in unit_id_set: + for obs_id in to_observations: + copied_links.append( + { + "from_unit_id": from_id, + "to_unit_id": obs_id, + "link_type": link["link_type"], + "weight": link["weight"], + "entity_name": link["entity_name"], + } + ) + + # Keep only direct links between visible nodes + direct_links = [ + link for link in links if link["from_unit_id"] in unit_id_set and link["to_unit_id"] in unit_id_set + ] + # Get entity information unit_entities = await conn.fetch(f""" SELECT ue.unit_id, e.canonical_name @@ -2813,6 +2879,18 @@ class MemoryEngine(MemoryEngineInterface): entity_map[unit_id] = [] entity_map[unit_id].append(entity_name) + # For observations, inherit entities from source memories + for unit in units: + if unit["source_memory_ids"] and unit["id"] not in entity_map: + # Collect entities from all source memories + source_entities = [] + for source_id in unit["source_memory_ids"]: + if source_id in entity_map: + source_entities.extend(entity_map[source_id]) + if source_entities: + # Deduplicate while preserving order + entity_map[unit["id"]] = list(dict.fromkeys(source_entities)) + # Build nodes nodes = [] for row in units: @@ -2846,14 +2924,15 @@ class MemoryEngine(MemoryEngineInterface): } ) - # Build edges + # Build edges (combine direct links and copied links from sources) edges = [] - for row in links: + all_links = direct_links + copied_links + for row in all_links: from_id = str(row["from_unit_id"]) to_id = str(row["to_unit_id"]) link_type = row["link_type"] weight = row["weight"] - entity_name = row["entity_name"] + entity_name = row.get("entity_name") # Color by link type if link_type == "temporal": diff --git a/hindsight-api/hindsight_api/engine/reflect/agent.py b/hindsight-api/hindsight_api/engine/reflect/agent.py index ba923751..289f0b6b 100644 --- a/hindsight-api/hindsight_api/engine/reflect/agent.py +++ b/hindsight-api/hindsight_api/engine/reflect/agent.py @@ -58,6 +58,7 @@ def _normalize_tool_name(name: str) -> str: - 'functions.done' (OpenAI-style prefix) - 'call=functions.done' (some models) - 'call=done' (some models) + - 'done<|channel|>commentary' (malformed special tokens appended) Returns the normalized tool name (e.g., 'done', 'recall', etc.) """ @@ -69,6 +70,11 @@ def _normalize_tool_name(name: str) -> str: if name.startswith("functions."): name = name[len("functions.") :] + # Handle malformed special tokens appended to tool name + # e.g., 'done<|channel|>commentary' -> 'done' + if "<|" in name: + name = name.split("<|")[0] + return name diff --git a/hindsight-api/hindsight_api/engine/search/link_expansion_retrieval.py b/hindsight-api/hindsight_api/engine/search/link_expansion_retrieval.py index 0c43d343..8d6fe99f 100644 --- a/hindsight-api/hindsight_api/engine/search/link_expansion_retrieval.py +++ b/hindsight-api/hindsight_api/engine/search/link_expansion_retrieval.py @@ -155,7 +155,6 @@ class LinkExpansionRetriever(GraphRetriever): all_seeds.extend(temporal_seeds) if not all_seeds: - logger.info("[LinkExpansion] No seeds found, returning empty results") return [], timings seed_ids = list({s.id for s in all_seeds}) diff --git a/hindsight-api/tests/test_consolidation.py b/hindsight-api/tests/test_consolidation.py index 0112fd2f..3fe745c5 100644 --- a/hindsight-api/tests/test_consolidation.py +++ b/hindsight-api/tests/test_consolidation.py @@ -1897,3 +1897,93 @@ class TestMentalModelRefreshAfterConsolidation: # Cleanup await memory.delete_bank(bank_id, request_context=request_context) + + @pytest.mark.asyncio + async def test_graph_endpoint_observations_inherit_links_and_entities( + self, memory: MemoryEngine, request_context + ): + """Test that graph endpoint shows links and entities for observations filtered by type. + + When filtering graph by type=observation: + - Observations should inherit links from their source memories + - Observations should show entities inherited from source memories + - Even when source memories are not visible, their links should be copied to observations + """ + bank_id = f"test-graph-obs-{uuid.uuid4().hex[:8]}" + + # Create the bank + await memory.get_bank_profile(bank_id=bank_id, request_context=request_context) + + # Retain content that will create world facts with shared entities + # This should create facts that are linked by shared entities + await memory.retain_async( + bank_id=bank_id, + content="Alice works at Google as a software engineer.", + request_context=request_context, + ) + + await memory.retain_async( + bank_id=bank_id, + content="Bob also works at Google in the sales department.", + request_context=request_context, + ) + + # Wait for consolidation to create observations + import asyncio + + await asyncio.sleep(2) + + # Get graph data filtered by observation type only + graph_data = await memory.get_graph_data( + bank_id=bank_id, + fact_type="observation", + limit=1000, + request_context=request_context, + ) + + # Should have observations + assert graph_data["total_units"] > 0, "Should have observations" + assert len(graph_data["nodes"]) > 0, "Should have observation nodes" + + # Verify all nodes are observations + for row in graph_data["table_rows"]: + assert row["fact_type"] == "observation", f"All nodes should be observations, got {row['fact_type']}" + + # Should have edges (inherited from source memories) + # Even though we're only showing observations, they should inherit links from their sources + assert len(graph_data["edges"]) > 0, ( + "Observations should have edges inherited from source memories. " + f"Found {len(graph_data['edges'])} edges" + ) + + # Should have entities (inherited from source memories) + observations_with_entities = [ + row for row in graph_data["table_rows"] if row["entities"] and row["entities"] != "None" + ] + assert len(observations_with_entities) > 0, ( + "Observations should inherit entities from source memories. " + f"Found {len(observations_with_entities)} observations with entities" + ) + + # Verify entities contain expected values + all_entities = " ".join([row["entities"] for row in graph_data["table_rows"]]) + assert "Alice" in all_entities or "Bob" in all_entities or "Google" in all_entities, ( + f"Expected to find Alice, Bob, or Google in entities, got: {all_entities}" + ) + + # Verify edge types are valid + valid_link_types = {"semantic", "temporal", "entity"} + for edge in graph_data["edges"]: + link_type = edge["data"]["linkType"] + assert link_type in valid_link_types, f"Invalid link type: {link_type}" + + # Verify all edges connect visible observation nodes + visible_node_ids = {row["id"] for row in graph_data["table_rows"]} + for edge in graph_data["edges"]: + source_id = edge["data"]["source"] + target_id = edge["data"]["target"] + assert source_id in visible_node_ids, f"Edge source {source_id[:8]} not in visible nodes" + assert target_id in visible_node_ids, f"Edge target {target_id[:8]} not in visible nodes" + + # Cleanup + await memory.delete_bank(bank_id, request_context=request_context) diff --git a/hindsight-api/tests/test_reflect_agent.py b/hindsight-api/tests/test_reflect_agent.py index 8d44b836..d867b44d 100644 --- a/hindsight-api/tests/test_reflect_agent.py +++ b/hindsight-api/tests/test_reflect_agent.py @@ -163,6 +163,12 @@ class TestToolNameNormalization: assert _normalize_tool_name("call=functions.recall") == "recall" assert _normalize_tool_name("call=functions.search_observations") == "search_observations" + def test_normalize_special_token_suffix(self): + """Tool names with malformed special tokens should be normalized.""" + assert _normalize_tool_name("done<|channel|>commentary") == "done" + assert _normalize_tool_name("recall<|endoftext|>") == "recall" + assert _normalize_tool_name("search_observations<|im_end|>extra") == "search_observations" + def test_is_done_tool(self): """Test _is_done_tool helper.""" # Standard @@ -174,9 +180,14 @@ class TestToolNameNormalization: assert _is_done_tool("call=done") is True assert _is_done_tool("call=functions.done") is True + # With malformed special tokens + assert _is_done_tool("done<|channel|>commentary") is True + assert _is_done_tool("done<|endoftext|>") is True + # Not done assert _is_done_tool("functions.recall") is False assert _is_done_tool("call=functions.recall") is False + assert _is_done_tool("recall<|channel|>done") is False class TestReflectAgentMocked: