Improve context explainability outputs

This commit is contained in:
KaifAhmad1
2026-03-13 06:31:38 +05:30
parent 2d90bdaad5
commit c2a6e944fe
18 changed files with 344 additions and 86 deletions
+2 -2
View File
@@ -2073,7 +2073,7 @@ class AgentContext:
def find_similar_entities( def find_similar_entities(
self, entity_id: str, similarity_type: str = "content", top_k: int = 10 self, entity_id: str, similarity_type: str = "content", top_k: int = 10
) -> List[Tuple[str, float]]: ) -> List[Dict[str, Any]]:
""" """
Find similar entities using advanced similarity measures. Find similar entities using advanced similarity measures.
@@ -2083,7 +2083,7 @@ class AgentContext:
top_k: Number of similar entities to return top_k: Number of similar entities to return
Returns: Returns:
List of (entity_id, similarity_score) tuples List of dicts with entity ID, content, type, and similarity score
""" """
if not self._graph_builder: if not self._graph_builder:
return [] return []
+2 -1
View File
@@ -258,7 +258,8 @@ class CausalChainAnalyzer:
WHERE ALL(i IN range(0, length(path)-2) | WHERE ALL(i IN range(0, length(path)-2) |
path[i].decision_id <> path[i+1].decision_id) path[i].decision_id <> path[i+1].decision_id)
RETURN d1.decision_id as decision_id, RETURN d1.decision_id as decision_id,
[node in nodes(path) | node.decision_id] as loop_path, d1.scenario as decision_scenario,
[node in nodes(path) | {{decision_id: node.decision_id, scenario: node.scenario, category: node.category}}] as loop_path,
length(path) as loop_length length(path) as loop_length
ORDER BY loop_length ORDER BY loop_length
""" """
+79 -22
View File
@@ -403,8 +403,8 @@ class ContextGraph:
def has_node(self, node_id: str) -> bool: def has_node(self, node_id: str) -> bool:
return node_id in self.nodes return node_id in self.nodes
def neighbors(self, node_id: str) -> List[str]: def neighbors(self, node_id: str) -> List[Dict[str, Any]]:
return self.get_neighbor_ids(node_id) return self.get_neighbors(node_id, hops=1)
def get_neighbor_ids( def get_neighbor_ids(
self, self,
@@ -421,8 +421,18 @@ class ContextGraph:
neighbor_ids.append(edge.target_id) neighbor_ids.append(edge.target_id)
return neighbor_ids return neighbor_ids
def get_nodes_by_label(self, label: str) -> List[str]: def get_nodes_by_label(self, label: str) -> List[Dict[str, Any]]:
return list(self.node_type_index.get(label, set())) result = []
for nid in self.node_type_index.get(label, set()):
node = self.nodes.get(nid)
if node:
result.append({
"id": node.node_id,
"content": node.content,
"type": node.node_type,
"metadata": node.properties,
})
return result
def get_node_property(self, node_id: str, property_name: str) -> Any: def get_node_property(self, node_id: str, property_name: str) -> Any:
node = self.nodes.get(node_id) node = self.nodes.get(node_id)
@@ -1424,7 +1434,7 @@ class ContextGraph:
decision = Decision( decision = Decision(
decision_id=current_id, decision_id=current_id,
category=decision_data.get("category", ""), category=decision_data.get("category", ""),
scenario=node.content, scenario=decision_data.get("scenario", node.content),
reasoning=decision_data.get("reasoning", ""), reasoning=decision_data.get("reasoning", ""),
outcome=decision_data.get("outcome", ""), outcome=decision_data.get("outcome", ""),
confidence=decision_data.get("confidence", 0.0), confidence=decision_data.get("confidence", 0.0),
@@ -1433,7 +1443,7 @@ class ContextGraph:
reasoning_embedding=decision_data.get("reasoning_embedding"), reasoning_embedding=decision_data.get("reasoning_embedding"),
node2vec_embedding=decision_data.get("node2vec_embedding"), node2vec_embedding=decision_data.get("node2vec_embedding"),
metadata={k: v for k, v in decision_data.items() if k not in [ metadata={k: v for k, v in decision_data.items() if k not in [
"category", "reasoning", "outcome", "confidence", "category", "scenario", "reasoning", "outcome", "confidence",
"timestamp", "decision_maker", "reasoning_embedding", "node2vec_embedding" "timestamp", "decision_maker", "reasoning_embedding", "node2vec_embedding"
]} ]}
) )
@@ -1489,7 +1499,7 @@ class ContextGraph:
decision = Decision( decision = Decision(
decision_id=pid, decision_id=pid,
category=decision_data.get("category", ""), category=decision_data.get("category", ""),
scenario=node.content, scenario=decision_data.get("scenario", node.content),
reasoning=decision_data.get("reasoning", ""), reasoning=decision_data.get("reasoning", ""),
outcome=decision_data.get("outcome", ""), outcome=decision_data.get("outcome", ""),
confidence=decision_data.get("confidence", 0.0), confidence=decision_data.get("confidence", 0.0),
@@ -1498,7 +1508,7 @@ class ContextGraph:
reasoning_embedding=decision_data.get("reasoning_embedding"), reasoning_embedding=decision_data.get("reasoning_embedding"),
node2vec_embedding=decision_data.get("node2vec_embedding"), node2vec_embedding=decision_data.get("node2vec_embedding"),
metadata={k: v for k, v in decision_data.items() if k not in [ metadata={k: v for k, v in decision_data.items() if k not in [
"category", "reasoning", "outcome", "confidence", "category", "scenario", "reasoning", "outcome", "confidence",
"timestamp", "decision_maker", "reasoning_embedding", "node2vec_embedding" "timestamp", "decision_maker", "reasoning_embedding", "node2vec_embedding"
]} ]}
) )
@@ -1609,7 +1619,7 @@ class ContextGraph:
def find_similar_nodes( def find_similar_nodes(
self, node_id: str, similarity_type: str = "content", top_k: int = 10 self, node_id: str, similarity_type: str = "content", top_k: int = 10
) -> List[Tuple[str, float]]: ) -> List[Dict[str, Any]]:
""" """
Find similar nodes using various similarity measures. Find similar nodes using various similarity measures.
@@ -1619,7 +1629,7 @@ class ContextGraph:
top_k: Number of similar nodes to return top_k: Number of similar nodes to return
Returns: Returns:
List of (node_id, similarity_score) tuples List of dicts with node ID, type, content, and similarity score
""" """
if node_id not in self.nodes: if node_id not in self.nodes:
return [] return []
@@ -1637,10 +1647,15 @@ class ContextGraph:
else: else:
similarity = self._calculate_content_similarity(reference_node, other_node) similarity = self._calculate_content_similarity(reference_node, other_node)
similar_nodes.append((other_id, similarity)) similar_nodes.append({
"id": other_id,
"content": other_node.content,
"type": other_node.node_type,
"score": similarity,
})
# Sort by similarity and return top_k # Sort by similarity and return top_k
similar_nodes.sort(key=lambda x: x[1], reverse=True) similar_nodes.sort(key=lambda x: x["score"], reverse=True)
return similar_nodes[:top_k] return similar_nodes[:top_k]
except Exception as e: except Exception as e:
@@ -2013,11 +2028,23 @@ class ContextGraph:
reverse=True reverse=True
) )
def _enrich(did: str) -> Dict[str, Any]:
dec = self._decisions.get(did, {})
return {
"decision_id": did,
"scenario": dec.get("scenario", ""),
"outcome": dec.get("outcome", ""),
"category": dec.get("category", ""),
}
return { return {
"decision_id": decision_id, "decision_id": decision_id,
"direct_influence": list(direct_influence), "direct_influence": [_enrich(did) for did in direct_influence],
"indirect_influence": list(indirect_influence), "indirect_influence": [_enrich(did) for did in indirect_influence],
"influence_scores": sorted_influence, "influence_scores": [
{**_enrich(did), "score": score}
for did, score in sorted_influence
],
"total_influenced": len(influence_scores), "total_influenced": len(influence_scores),
"max_influence_score": max(influence_scores.values()) if influence_scores else 0.0 "max_influence_score": max(influence_scores.values()) if influence_scores else 0.0
} }
@@ -2115,7 +2142,15 @@ class ContextGraph:
potential_causes.append(other_decision_id) potential_causes.append(other_decision_id)
for cause_id in potential_causes: for cause_id in potential_causes:
cause_path = path + [{"from": cause_id, "to": current_id, "type": "influences"}] cause_dec = self._decisions.get(cause_id, {})
hop = {
"from": cause_id,
"from_scenario": cause_dec.get("scenario", ""),
"to": current_id,
"to_scenario": current_decision.get("scenario", ""),
"type": "influences",
}
cause_path = path + [hop]
causal_chain.append(cause_path) causal_chain.append(cause_path)
trace_recursive(cause_id, depth + 1, cause_path) trace_recursive(cause_id, depth + 1, cause_path)
@@ -2185,17 +2220,36 @@ class ContextGraph:
def _add_decision_to_graph(self, decision: Dict[str, Any]) -> None: def _add_decision_to_graph(self, decision: Dict[str, Any]) -> None:
"""Add decision to context graph.""" """Add decision to context graph."""
try: try:
extra_properties = {
key: value
for key, value in decision.items()
if key not in {
"id",
"category",
"scenario",
"reasoning",
"outcome",
"confidence",
"entities",
"decision_maker",
"timestamp",
"metadata",
}
}
# Add decision node # Add decision node
self.add_node( self.add_node(
decision["id"], decision["id"],
"decision", "decision",
content=decision["scenario"],
category=decision["category"], category=decision["category"],
outcome=decision["outcome"], outcome=decision["outcome"],
confidence=decision["confidence"], confidence=decision["confidence"],
timestamp=decision["timestamp"], timestamp=decision["timestamp"],
scenario=decision["scenario"][:100] + "..." if len(decision["scenario"]) > 100 else decision["scenario"], scenario=decision["scenario"],
decision_maker=decision.get("decision_maker", ""), decision_maker=decision.get("decision_maker", ""),
reasoning=decision["reasoning"][:200] + "..." if len(decision["reasoning"]) > 200 else decision["reasoning"] reasoning=decision["reasoning"],
**(decision.get("metadata") or {}),
**extra_properties,
) )
# Add entity nodes and relationships # Add entity nodes and relationships
@@ -2281,8 +2335,11 @@ class ContextGraph:
) )
if similar_nodes: if similar_nodes:
# similar_nodes is List[Tuple[str, float]], extract similarity scores return max(
return max(similarity for node_id, similarity in similar_nodes) item.get("score", 0.0)
for item in similar_nodes
if isinstance(item, dict)
)
except Exception as e: except Exception as e:
self.logger.exception("Structural similarity calculation failed") self.logger.exception("Structural similarity calculation failed")
@@ -2538,7 +2595,7 @@ class ContextGraph:
node_id: str, node_id: str,
how_many: int = 10, how_many: int = 10,
similarity_type: str = "content" similarity_type: str = "content"
) -> List[Tuple[str, float]]: ) -> List[Dict[str, Any]]:
""" """
Easy way to find nodes similar to a given node. Easy way to find nodes similar to a given node.
@@ -2548,7 +2605,7 @@ class ContextGraph:
similarity_type: Type of similarity ("content", "structural") similarity_type: Type of similarity ("content", "structural")
Returns: Returns:
List of (node_id, similarity_score) tuples List of dicts with node ID, type, content, and similarity score
""" """
return self.find_similar_nodes( return self.find_similar_nodes(
node_id=node_id, node_id=node_id,
+27 -11
View File
@@ -2023,7 +2023,11 @@ Answer:"""
elif hasattr(self.knowledge_graph, "get_neighbor_ids"): elif hasattr(self.knowledge_graph, "get_neighbor_ids"):
return list(self.knowledge_graph.get_neighbor_ids(node)) return list(self.knowledge_graph.get_neighbor_ids(node))
elif hasattr(self.knowledge_graph, "neighbors"): elif hasattr(self.knowledge_graph, "neighbors"):
return list(self.knowledge_graph.neighbors(node)) return [
n.get("id") if isinstance(n, dict) else n
for n in self.knowledge_graph.neighbors(node)
if n
]
return [] return []
visited: set = {entity_name} visited: set = {entity_name}
@@ -2111,7 +2115,11 @@ Answer:"""
# Simplified centrality calculation # Simplified centrality calculation
if hasattr(self.knowledge_graph, 'get_neighbors'): if hasattr(self.knowledge_graph, 'get_neighbors'):
if hasattr(self.knowledge_graph, "neighbors"): if hasattr(self.knowledge_graph, "neighbors"):
neighbor_ids = list(self.knowledge_graph.neighbors(entity_name)) neighbor_ids = [
n.get("id") if isinstance(n, dict) else n
for n in self.knowledge_graph.neighbors(entity_name)
if n
]
elif hasattr(self.knowledge_graph, "get_neighbor_ids"): elif hasattr(self.knowledge_graph, "get_neighbor_ids"):
neighbor_ids = self.knowledge_graph.get_neighbor_ids(entity_name) neighbor_ids = self.knowledge_graph.get_neighbor_ids(entity_name)
else: else:
@@ -2162,8 +2170,13 @@ Answer:"""
if hasattr(self.knowledge_graph, 'get_nodes_by_label'): if hasattr(self.knowledge_graph, 'get_nodes_by_label'):
policy_nodes = self.knowledge_graph.get_nodes_by_label("Policy") policy_nodes = self.knowledge_graph.get_nodes_by_label("Policy")
for policy in policy_nodes[:5]: # Limit results for policy in policy_nodes[:5]: # Limit results
policy_name = (
policy.get("content")
or policy.get("metadata", {}).get("name", "")
or policy.get("id", "")
) if isinstance(policy, dict) else policy
policies.append({ policies.append({
"name": policy, "name": policy_name,
"type": "policy", "type": "policy",
"source": "policy_search", "source": "policy_search",
"related_category": category "related_category": category
@@ -2482,17 +2495,20 @@ Answer:"""
decision_nodes = self.knowledge_graph.get_nodes_by_label("Decision") decision_nodes = self.knowledge_graph.get_nodes_by_label("Decision")
for node_data in decision_nodes[:limit]: for node_data in decision_nodes[:limit]:
metadata = {}
if isinstance(node_data, dict):
metadata = node_data.get("metadata") or node_data.get("properties") or {}
# Convert to Decision object # Convert to Decision object
decision = Decision( decision = Decision(
decision_id=node_data.get("id", ""), decision_id=node_data.get("id", "") if isinstance(node_data, dict) else "",
category=node_data.get("properties", {}).get("category", ""), category=metadata.get("category", ""),
scenario=node_data.get("content", ""), scenario=node_data.get("content", "") if isinstance(node_data, dict) else "",
reasoning=node_data.get("properties", {}).get("reasoning", ""), reasoning=metadata.get("reasoning", ""),
outcome=node_data.get("properties", {}).get("outcome", ""), outcome=metadata.get("outcome", ""),
confidence=node_data.get("properties", {}).get("confidence", 0.0), confidence=metadata.get("confidence", 0.0),
timestamp=datetime.now(), timestamp=datetime.now(),
decision_maker=node_data.get("properties", {}).get("decision_maker", ""), decision_maker=metadata.get("decision_maker", ""),
metadata=node_data.get("properties", {}) metadata=metadata
) )
# Filter by category if specified # Filter by category if specified
+13 -4
View File
@@ -623,13 +623,22 @@ def analyze_decision_impact(
# Get root causes # Get root causes
root_causes = analyzer.find_root_causes(decision_id, max_depth=5) root_causes = analyzer.find_root_causes(decision_id, max_depth=5)
def _decision_dict(d) -> Dict[str, Any]:
return {
"decision_id": d.decision_id,
"scenario": d.scenario,
"category": d.category,
"outcome": d.outcome,
"confidence": d.confidence,
}
return { return {
"decision_id": decision_id, "decision_id": decision_id,
"impact_score": impact_score, "impact_score": impact_score,
"influenced_decisions": len(influenced), "influenced_decisions": [_decision_dict(d) for d in influenced],
"root_causes": len(root_causes), "root_causes": [_decision_dict(d) for d in root_causes],
"influenced_decision_ids": [d.decision_id for d in influenced], "total_influenced": len(influenced),
"root_cause_ids": [d.decision_id for d in root_causes], "total_root_causes": len(root_causes),
"analysis_timestamp": datetime.now().isoformat() "analysis_timestamp": datetime.now().isoformat()
} }
+6 -3
View File
@@ -573,7 +573,9 @@ class DecisionQuery:
query = f""" query = f"""
MATCH (d:Decision {{decision_id: $decision_id}}) MATCH (d:Decision {{decision_id: $decision_id}})
MATCH path = (d)-[:{rel_filter}*]-(related) MATCH path = (d)-[:{rel_filter}*]-(related)
RETURN path, length(path) as path_length RETURN [node in nodes(path) | {{decision_id: node.decision_id, scenario: node.scenario, category: node.category}}] as path_nodes,
[r in relationships(path) | {{from: startNode(r).decision_id, to: endNode(r).decision_id, type: type(r)}}] as path_rels,
length(path) as path_length
ORDER BY path_length ORDER BY path_length
""" """
@@ -585,8 +587,9 @@ class DecisionQuery:
paths = [] paths = []
for record in results: for record in results:
path_info = { path_info = {
"path": record.get("path"), "path_length": record.get("path_length", 0),
"path_length": record.get("path_length", 0) "nodes": record.get("path_nodes", []),
"relationships": record.get("path_rels", []),
} }
paths.append(path_info) paths.append(path_info)
+22 -8
View File
@@ -366,7 +366,7 @@ class EntityLinker:
entity_text: str, entity_text: str,
entity_type: Optional[str] = None, entity_type: Optional[str] = None,
threshold: Optional[float] = None, threshold: Optional[float] = None,
) -> List[Tuple[str, float]]: ) -> List[Dict[str, Any]]:
""" """
Find similar entities in knowledge graph. Find similar entities in knowledge graph.
@@ -376,7 +376,7 @@ class EntityLinker:
threshold: Similarity threshold (uses default if None) threshold: Similarity threshold (uses default if None)
Returns: Returns:
List of (entity_id, similarity_score) tuples List of dicts with entity_id, text, type, uri, and similarity
""" """
threshold = threshold or self.similarity_threshold threshold = threshold or self.similarity_threshold
@@ -404,10 +404,16 @@ class EntityLinker:
if similarity >= threshold: if similarity >= threshold:
entity_id = entity.get("id") or entity.get("entity_id") entity_id = entity.get("id") or entity.get("entity_id")
if entity_id: if entity_id:
similar_entities.append((entity_id, similarity)) similar_entities.append({
"entity_id": entity_id,
"text": entity_text2,
"type": entity.get("type", ""),
"uri": self.entity_registry.get(entity_id, ""),
"similarity": similarity,
})
# Sort by similarity # Sort by similarity
similar_entities.sort(key=lambda x: x[1], reverse=True) similar_entities.sort(key=lambda x: x["similarity"], reverse=True)
return similar_entities return similar_entities
@@ -425,7 +431,11 @@ class EntityLinker:
# Find similar entities in knowledge graph # Find similar entities in knowledge graph
if self.knowledge_graph: if self.knowledge_graph:
similar = self.find_similar_entities(entity_text, entity_type) similar = self.find_similar_entities(entity_text, entity_type)
for similar_id, similarity in similar: for similar_entity in similar:
similar_id = similar_entity.get("entity_id")
similarity = similar_entity.get("similarity", 0.0)
if not similar_id:
continue
if similar_id != entity_id: if similar_id != entity_id:
links.append( links.append(
EntityLink( EntityLink(
@@ -578,7 +588,11 @@ class EntityLinker:
) )
linked_entities = [] linked_entities = []
for similar_id, similarity in similar: for similar_entity in similar:
similar_id = similar_entity.get("entity_id")
similarity = similar_entity.get("similarity", 0.0)
if not similar_id:
continue
linked_entities.append( linked_entities.append(
EntityLink( EntityLink(
source_entity_id=entity.get("id", ""), source_entity_id=entity.get("id", ""),
@@ -604,7 +618,7 @@ class EntityLinker:
# Search Methods # Search Methods
def find_similar( def find_similar(
self, entity: Union[str, EntityDict], threshold: float = 0.8 self, entity: Union[str, EntityDict], threshold: float = 0.8
) -> List[Tuple[str, float]]: ) -> List[Dict[str, Any]]:
""" """
Find similar entities. Find similar entities.
@@ -613,7 +627,7 @@ class EntityLinker:
threshold: Similarity threshold (default: 0.8) threshold: Similarity threshold (default: 0.8)
Returns: Returns:
List of (entity_id, similarity) tuples List of dicts with entity_id, text, type, uri, and similarity
Example: Example:
>>> similar = linker.find_similar("Python", threshold=0.8) >>> similar = linker.find_similar("Python", threshold=0.8)
+81 -10
View File
@@ -601,7 +601,7 @@ class PolicyEngine:
to_version: New version to_version: New version
Returns: Returns:
List of affected decision IDs List of affected decisions with readable metadata
""" """
try: try:
if self._supports_cypher: if self._supports_cypher:
@@ -610,16 +610,23 @@ class PolicyEngine:
policy_id: $policy_id, policy_id: $policy_id,
version: $from_version version: $from_version
}) })
RETURN d.decision_id as decision_id RETURN d.decision_id as decision_id,
d.scenario as scenario,
d.category as category,
d.outcome as outcome,
d.confidence as confidence
""" """
results = self.graph_store.execute_query(query, { results = self._extract_records(self.graph_store.execute_query(query, {
"policy_id": policy_id, "policy_id": policy_id,
"from_version": from_version "from_version": from_version
}) }))
decisions = [] decisions = [
for record in results: self._enrich_affected_decision(
decisions.append(record if isinstance(record, dict) else {"decision_id": record}) record if isinstance(record, dict) else {"decision_id": record}
)
for record in results
]
self.logger.info(f"Found {len(decisions)} decisions affected by policy change") self.logger.info(f"Found {len(decisions)} decisions affected by policy change")
return decisions return decisions
@@ -627,16 +634,80 @@ class PolicyEngine:
if not hasattr(self.graph_store, "find_edges"): if not hasattr(self.graph_store, "find_edges"):
return [] return []
policy_node_id = f"{policy_id}:{from_version}" policy_node_id = f"{policy_id}:{from_version}"
decision_ids: List[str] = [] decisions: List[Dict[str, Any]] = []
for edge in self.graph_store.find_edges(edge_type="APPLIED_POLICY"): for edge in self.graph_store.find_edges(edge_type="APPLIED_POLICY"):
if edge.get("target") == policy_node_id: if edge.get("target") == policy_node_id:
decision_ids.append(edge.get("source")) decisions.append(
return decision_ids self._enrich_affected_decision(
{"decision_id": edge.get("source")}
)
)
return decisions
except Exception as e: except Exception as e:
self.logger.exception("Failed to get affected decisions") self.logger.exception("Failed to get affected decisions")
raise raise
def _enrich_affected_decision(self, record: Dict[str, Any]) -> Dict[str, Any]:
"""Normalize an affected decision record with readable fields."""
decision_id = record.get("decision_id") or record.get("source") or ""
node_record = self._get_decision_node_record(decision_id)
enriched = {
"decision_id": decision_id,
"scenario": record.get("scenario", node_record.get("scenario", "")),
"category": record.get("category", node_record.get("category", "")),
"outcome": record.get("outcome", node_record.get("outcome", "")),
"confidence": record.get("confidence", node_record.get("confidence", 0.0)),
}
for key, value in record.items():
if key not in enriched:
enriched[key] = value
return enriched
def _get_decision_node_record(self, decision_id: str) -> Dict[str, Any]:
"""Best-effort lookup of a decision node from the backing graph store."""
if not decision_id:
return {}
if hasattr(self.graph_store, "nodes"):
nodes = getattr(self.graph_store, "nodes", {})
if isinstance(nodes, dict):
node = nodes.get(decision_id)
if node:
properties = getattr(node, "properties", {}) or {}
content = getattr(node, "content", "") or ""
return {
"scenario": properties.get("scenario", content),
"category": properties.get("category", ""),
"outcome": properties.get("outcome", ""),
"confidence": properties.get("confidence", 0.0),
}
get_node = getattr(self.graph_store, "get_node", None)
if callable(get_node):
try:
node = get_node(decision_id)
except Exception:
return {}
if isinstance(node, dict):
properties = node.get("properties", {}) or {}
return {
"scenario": (
properties.get("scenario")
or node.get("content")
or node.get("scenario", "")
),
"category": properties.get("category", node.get("category", "")),
"outcome": properties.get("outcome", node.get("outcome", "")),
"confidence": properties.get("confidence", node.get("confidence", 0.0)),
}
return {}
def analyze_policy_impact( def analyze_policy_impact(
self, self,
policy_id: str, policy_id: str,
+2 -1
View File
@@ -757,7 +757,8 @@ class CentralityCalculator:
) -> List[str]: ) -> List[str]:
"""Get neighbors filtered by relationship types.""" """Get neighbors filtered by relationship types."""
if hasattr(graph, 'neighbors'): if hasattr(graph, 'neighbors'):
neighbors = list(graph.neighbors(node)) _raw = list(graph.neighbors(node))
neighbors = [n.get("id") if isinstance(n, dict) else n for n in _raw]
elif hasattr(graph, 'get_neighbors'): elif hasattr(graph, 'get_neighbors'):
neighbors = graph.get_neighbors(node) neighbors = graph.get_neighbors(node)
if neighbors and isinstance(neighbors[0], dict): if neighbors and isinstance(neighbors[0], dict):
+6 -2
View File
@@ -398,7 +398,11 @@ class LinkPredictor:
if hasattr(graph_store, 'get_nodes_by_label') and callable(graph_store.get_nodes_by_label): if hasattr(graph_store, 'get_nodes_by_label') and callable(graph_store.get_nodes_by_label):
result = graph_store.get_nodes_by_label(label) result = graph_store.get_nodes_by_label(label)
if isinstance(result, list): if isinstance(result, list):
nodes.extend(result) nodes.extend(
item.get("id") if isinstance(item, dict) else item
for item in result
if item and (not isinstance(item, dict) or item.get("id"))
)
else: else:
# Fallback - get all nodes and filter by label if possible # Fallback - get all nodes and filter by label if possible
all_nodes = self._get_all_nodes(graph_store) all_nodes = self._get_all_nodes(graph_store)
@@ -500,7 +504,7 @@ class LinkPredictor:
return neighbors return neighbors
if hasattr(graph_store, 'neighbors') and callable(graph_store.neighbors): if hasattr(graph_store, 'neighbors') and callable(graph_store.neighbors):
try: try:
raw = list(graph_store.neighbors(node_id)) raw = [n.get("id") if isinstance(n, dict) else n for n in graph_store.neighbors(node_id)]
if not isinstance(raw, list): if not isinstance(raw, list):
return [] return []
if relationship_types and hasattr(graph_store, 'get_edge_data') and callable(graph_store.get_edge_data): if relationship_types and hasattr(graph_store, 'get_edge_data') and callable(graph_store.get_edge_data):
+19 -2
View File
@@ -356,7 +356,13 @@ class NodeEmbedder:
nodes = [] nodes = []
if hasattr(graph_store, 'get_nodes_by_label'): if hasattr(graph_store, 'get_nodes_by_label'):
for label in node_labels: for label in node_labels:
nodes.extend(graph_store.get_nodes_by_label(label)) for node in graph_store.get_nodes_by_label(label):
if isinstance(node, dict):
node_id = node.get("id")
if node_id:
nodes.append(node_id)
elif node:
nodes.append(node)
else: else:
# Fallback for different graph store implementations # Fallback for different graph store implementations
nodes = list(graph_store.nodes()) nodes = list(graph_store.nodes())
@@ -365,7 +371,14 @@ class NodeEmbedder:
for node in nodes: for node in nodes:
if hasattr(graph_store, 'get_neighbors'): if hasattr(graph_store, 'get_neighbors'):
try: try:
try:
neighbor_details = graph_store.get_neighbors(
node,
relationship_types=relationship_types,
)
except TypeError:
neighbor_details = graph_store.get_neighbors(node, relationship_types) neighbor_details = graph_store.get_neighbors(node, relationship_types)
if isinstance(neighbor_details, list): if isinstance(neighbor_details, list):
adjacency[node] = [ adjacency[node] = [
n.get("id") if isinstance(n, dict) else n n.get("id") if isinstance(n, dict) else n
@@ -380,7 +393,11 @@ class NodeEmbedder:
adjacency[node] = list(graph_store.get_neighbor_ids(node, relationship_types)) adjacency[node] = list(graph_store.get_neighbor_ids(node, relationship_types))
elif hasattr(graph_store, 'neighbors'): elif hasattr(graph_store, 'neighbors'):
try: try:
adjacency[node] = list(graph_store.neighbors(node)) adjacency[node] = [
n.get("id") if isinstance(n, dict) else n
for n in graph_store.neighbors(node)
if n
]
except TypeError: except TypeError:
adjacency[node] = [] adjacency[node] = []
else: else:
+2 -1
View File
@@ -569,7 +569,8 @@ class PathFinder:
neighbors = [] neighbors = []
if hasattr(graph, 'neighbors'): if hasattr(graph, 'neighbors'):
for neighbor in graph.neighbors(node): for _raw in graph.neighbors(node):
neighbor = _raw.get("id") if isinstance(_raw, dict) else _raw
edge_data = self._get_edge_data(graph, node, neighbor) edge_data = self._get_edge_data(graph, node, neighbor)
neighbors.append((neighbor, edge_data)) neighbors.append((neighbor, edge_data))
elif hasattr(graph, 'get_neighbors'): elif hasattr(graph, 'get_neighbors'):
+9 -1
View File
@@ -217,7 +217,13 @@ class TestCausalChainAnalyzer:
mock_graph_store.execute_query.return_value = [ mock_graph_store.execute_query.return_value = [
{ {
"decision_id": "decision_001", "decision_id": "decision_001",
"loop_path": ["decision_001", "decision_002", "decision_003", "decision_001"], "decision_scenario": "Approve LATAM expansion",
"loop_path": [
{"decision_id": "decision_001", "scenario": "Approve LATAM expansion", "category": "strategy"},
{"decision_id": "decision_002", "scenario": "Fund regional hiring", "category": "finance"},
{"decision_id": "decision_003", "scenario": "Open Sao Paulo office", "category": "operations"},
{"decision_id": "decision_001", "scenario": "Approve LATAM expansion", "category": "strategy"},
],
"loop_length": 3, "loop_length": 3,
"cycle_strength": 0.7 "cycle_strength": 0.7
} }
@@ -227,6 +233,8 @@ class TestCausalChainAnalyzer:
assert len(loops) == 1 assert len(loops) == 1
assert loops[0]["decision_id"] == "decision_001" assert loops[0]["decision_id"] == "decision_001"
assert loops[0]["decision_scenario"] == "Approve LATAM expansion"
assert loops[0]["loop_path"][0]["scenario"] == "Approve LATAM expansion"
assert len(loops[0]["loop_path"]) == 4 # Including return to start assert len(loops[0]["loop_path"]) == 4 # Including return to start
assert loops[0]["loop_length"] == 3 assert loops[0]["loop_length"] == 3
+42
View File
@@ -64,6 +64,23 @@ class TestContextModule(unittest.TestCase):
# However, checking it runs without error is a good start. # However, checking it runs without error is a good start.
self.assertIsInstance(linked, list) self.assertIsInstance(linked, list)
def test_entity_linker_find_similar_entities_returns_dicts(self):
linker = EntityLinker(
knowledge_graph={
"entities": [
{"id": "lang_python", "text": "Python programming language", "type": "Technology"}
]
}
)
linker.assign_uri("lang_python", "Python programming language", "Technology")
similar = linker.find_similar_entities("Python programming language", threshold=0.5)
self.assertEqual(len(similar), 1)
self.assertEqual(similar[0]["entity_id"], "lang_python")
self.assertEqual(similar[0]["text"], "Python programming language")
self.assertIn("similarity", similar[0])
# --- ContextGraph Tests --- # --- ContextGraph Tests ---
def test_context_graph_operations(self): def test_context_graph_operations(self):
graph = ContextGraph() graph = ContextGraph()
@@ -92,6 +109,31 @@ class TestContextModule(unittest.TestCase):
self.assertEqual(neighbors[0]["id"], "n2") self.assertEqual(neighbors[0]["id"], "n2")
self.assertEqual(neighbors[0]["relationship"], "knows") self.assertEqual(neighbors[0]["relationship"], "knows")
def test_context_graph_preserves_full_decision_text(self):
graph = ContextGraph()
scenario = "Launch regional expansion plan " + ("X" * 140)
root_id = graph.record_decision(
category="strategy",
scenario=scenario,
reasoning="Growth opportunity with strong local demand",
outcome="approved",
confidence=0.91,
)
child_id = graph.record_decision(
category="operations",
scenario="Open Sao Paulo office",
reasoning="Needed to support expansion",
outcome="pending",
confidence=0.74,
)
graph.add_causal_relationship(root_id, child_id, "CAUSED")
chain = graph.get_causal_chain(child_id)
self.assertEqual(graph.nodes[root_id].content, scenario)
self.assertEqual(graph.nodes[root_id].properties["scenario"], scenario)
self.assertEqual(chain[0].scenario, scenario)
# --- AgentMemory Tests --- # --- AgentMemory Tests ---
def test_agent_memory_store(self): def test_agent_memory_store(self):
memory = AgentMemory(vector_store=self.mock_vector_store) memory = AgentMemory(vector_store=self.mock_vector_store)
+8 -3
View File
@@ -302,11 +302,13 @@ class TestDecisionQuery:
mock_graph_store.execute_query.return_value = [ mock_graph_store.execute_query.return_value = [
{ {
"path": "mock_path_1", "path_nodes": [{"decision_id": "d1", "scenario": "S1", "category": "C1"}],
"path_rels": [{"from": "d1", "to": "d2", "type": "CAUSED"}],
"path_length": 2 "path_length": 2
}, },
{ {
"path": "mock_path_2", "path_nodes": [{"decision_id": "d2", "scenario": "S2", "category": "C2"}],
"path_rels": [],
"path_length": 3 "path_length": 3
} }
] ]
@@ -314,8 +316,11 @@ class TestDecisionQuery:
paths = decision_query.trace_decision_path(decision_id, relationship_types) paths = decision_query.trace_decision_path(decision_id, relationship_types)
assert len(paths) == 2 assert len(paths) == 2
assert paths[0]["path"] == "mock_path_1"
assert paths[0]["path_length"] == 2 assert paths[0]["path_length"] == 2
assert isinstance(paths[0]["nodes"], list)
assert isinstance(paths[0]["relationships"], list)
assert paths[0]["nodes"][0]["scenario"] == "S1"
assert paths[0]["relationships"][0]["type"] == "CAUSED"
# Verify query was called with relationship types # Verify query was called with relationship types
call_args = mock_graph_store.execute_query.call_args call_args = mock_graph_store.execute_query.call_args
+4
View File
@@ -460,6 +460,10 @@ class TestPolicyEngine:
assert len(affected_decisions) == 2 assert len(affected_decisions) == 2
assert affected_decisions[0]["decision_id"] == "decision_001" assert affected_decisions[0]["decision_id"] == "decision_001"
assert affected_decisions[0]["applied_policy_version"] == from_version assert affected_decisions[0]["applied_policy_version"] == from_version
assert "scenario" in affected_decisions[0]
assert "category" in affected_decisions[0]
assert "outcome" in affected_decisions[0]
assert "confidence" in affected_decisions[0]
def test_analyze_policy_impact_success(self, policy_engine, mock_graph_store): def test_analyze_policy_impact_success(self, policy_engine, mock_graph_store):
"""Test policy impact analysis.""" """Test policy impact analysis."""
@@ -346,9 +346,9 @@ class TestContextGraphAdvancedDecisionMethods:
g, ids = self._build_loan_graph() g, ids = self._build_loan_graph()
result = g.analyze_decision_influence(ids["alice"]) result = g.analyze_decision_influence(ids["alice"])
# alice, bob, and carol all share underwriter_ai_v5 — they should appear in influence # alice, bob, and carol all share underwriter_ai_v5 — they should appear in influence
direct = set(result["direct_influence"]) direct = result["direct_influence"]
# At minimum the category-shared decisions should appear # At minimum the category-shared decisions should appear
assert isinstance(direct, set) assert isinstance(direct, list)
assert result["total_influenced"] >= 0 # May be 0 if no category overlap assert result["total_influenced"] >= 0 # May be 0 if no category overlap
def test_analyze_decision_influence_category_cross(self): def test_analyze_decision_influence_category_cross(self):
@@ -356,7 +356,8 @@ class TestContextGraphAdvancedDecisionMethods:
g, ids = self._build_loan_graph() g, ids = self._build_loan_graph()
result = g.analyze_decision_influence(ids["bob"]) result = g.analyze_decision_influence(ids["bob"])
# alice and bob are both "mortgage" category — alice should appear in influence # alice and bob are both "mortgage" category — alice should appear in influence
assert ids["alice"] in result["direct_influence"] or result["total_influenced"] >= 0 direct_ids = [d["decision_id"] for d in result["direct_influence"]]
assert ids["alice"] in direct_ids or result["total_influenced"] >= 0
def test_analyze_decision_influence_nonexistent_raises(self): def test_analyze_decision_influence_nonexistent_raises(self):
g, _ = self._build_loan_graph() g, _ = self._build_loan_graph()
@@ -2019,8 +2020,8 @@ class TestContextGraphFindSimilarNodes:
similar = g.find_similar_nodes("repo_pytorch", similarity_type="structural", top_k=5) similar = g.find_similar_nodes("repo_pytorch", similarity_type="structural", top_k=5)
assert isinstance(similar, list) assert isinstance(similar, list)
for item in similar: for item in similar:
node_id, score = item assert isinstance(item, dict)
assert 0.0 <= score <= 1.0 assert 0.0 <= item["score"] <= 1.0
def test_similar_nodes_nonexistent_returns_empty(self): def test_similar_nodes_nonexistent_returns_empty(self):
g = _build_research_graph() g = _build_research_graph()
+7 -3
View File
@@ -421,9 +421,13 @@ class TestContextGraphKGAnalytics:
g = _build_tech_graph() g = _build_tech_graph()
similar = g.find_similar_nodes("apple", similarity_type="structural", top_k=5) similar = g.find_similar_nodes("apple", similarity_type="structural", top_k=5)
assert isinstance(similar, list) assert isinstance(similar, list)
for node_id, score in similar: for item in similar:
assert isinstance(score, float) assert isinstance(item, dict)
assert 0.0 <= score <= 1.0 assert isinstance(item.get("id"), str)
assert isinstance(item.get("content"), str)
assert isinstance(item.get("type"), str)
assert isinstance(item.get("score"), float)
assert 0.0 <= item["score"] <= 1.0
def test_find_similar_nodes_missing_node_returns_empty(self): def test_find_similar_nodes_missing_node_returns_empty(self):
g = _build_tech_graph() g = _build_tech_graph()