Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-rag/src/lexigram/ai/rag/knowledge_graph/core.py: 15%
248 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-25 07:19 +0800
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-25 07:19 +0800
1"""Core knowledge graph implementation (migrated from legacy single-file module).
3This module contains the canonical, package-local implementations and should be
4imported from `lexigram.ai.rag.knowledge_graph`.
5"""
7from __future__ import annotations
9from collections import defaultdict, deque
10from datetime import UTC, datetime
11from typing import TYPE_CHECKING, Any
13from lexigram.ai.rag.knowledge_graph.extractors import (
14 EntityExtractor,
15 RelationshipExtractor,
16)
17from lexigram.ai.rag.knowledge_graph.types import (
18 Entity,
19 EntityType,
20 GraphPath,
21 Relationship,
22 RelationshipType,
23)
25if TYPE_CHECKING:
26 from lexigram.contracts.ai import (
27 DocumentVectorStoreProtocol,
28 LLMClientProtocol,
29 )
32class KnowledgeGraph:
33 """In-memory knowledge graph for storing and querying entities and relationships."""
35 def __init__(self) -> None:
36 self._entities: dict[str, Entity] = {}
37 self._relationships: dict[str, list[Relationship]] = defaultdict(list)
38 self._reverse_relationships: dict[str, list[Relationship]] = defaultdict(list)
39 self._entity_index: dict[EntityType, set[str]] = defaultdict(set)
40 self._created_at = datetime.now(UTC)
41 self._updated_at = datetime.now(UTC)
43 async def add_entity(self, entity: Entity) -> None:
44 key = entity.name.lower()
45 self._entities[key] = entity
46 self._entity_index[EntityType(entity.type)].add(key)
47 self._updated_at = datetime.now(UTC)
49 async def add_entities(self, entities: list[Entity]) -> None:
50 for entity in entities:
51 await self.add_entity(entity)
53 async def add_relationship(self, relationship: Relationship) -> None:
54 source_key = relationship.source.lower()
55 target_key = relationship.target.lower()
57 self._relationships[source_key].append(relationship)
58 self._reverse_relationships[target_key].append(relationship)
59 self._updated_at = datetime.now(UTC)
61 async def add_relationships(self, relationships: list[Relationship]) -> None:
62 for relationship in relationships:
63 await self.add_relationship(relationship)
65 async def get_entity(self, name: str) -> Entity | None:
66 return self._entities.get(name.lower())
68 async def get_all_entities(self) -> list[Entity]:
69 """Get all entities in the knowledge graph."""
70 return list(self._entities.values())
72 async def get_all_relationships(self) -> list[Relationship]:
73 """Get all relationships in the knowledge graph."""
74 all_relationships = []
75 for rels in self._relationships.values():
76 all_relationships.extend(rels)
77 return all_relationships
79 async def get_entities_by_type(self, entity_type: EntityType | str) -> list[Entity]:
80 entity_type = EntityType(entity_type)
81 names = self._entity_index.get(entity_type, set())
82 return [self._entities[name] for name in names if name in self._entities]
84 async def get_neighbors(
85 self,
86 entity_name: str,
87 direction: str = "outgoing",
88 ) -> list[Entity]:
89 key = entity_name.lower()
90 neighbors = set()
92 if direction in ("outgoing", "both"):
93 for rel in self._relationships.get(key, []):
94 target = await self.get_entity(rel.target)
95 if target:
96 neighbors.add(target)
98 if direction in ("incoming", "both"):
99 for rel in self._reverse_relationships.get(key, []):
100 source = await self.get_entity(rel.source)
101 if source:
102 neighbors.add(source)
104 return list(neighbors)
106 async def get_relationships(
107 self,
108 entity_name: str,
109 direction: str = "outgoing",
110 ) -> list[Relationship]:
111 key = entity_name.lower()
112 rels = []
114 if direction in ("outgoing", "both"):
115 rels.extend(self._relationships.get(key, []))
117 if direction in ("incoming", "both"):
118 rels.extend(self._reverse_relationships.get(key, []))
120 return rels
122 async def find_path(
123 self,
124 source: str,
125 target: str,
126 max_depth: int = 5,
127 relationship_types: list[RelationshipType | str] | None = None,
128 ) -> GraphPath | None:
129 source_key = source.lower()
130 target_key = target.lower()
132 if source_key not in self._entities or target_key not in self._entities:
133 return None
135 queue: deque[tuple[str, list[Relationship]]] = deque([(source_key, [])])
136 visited = {source_key}
138 while queue:
139 current, path_rels = queue.popleft()
141 if len(path_rels) >= max_depth:
142 continue
144 if current == target_key:
145 entities = [source]
146 for rel in path_rels:
147 entities.append(rel.target)
149 return GraphPath(
150 entities=entities,
151 relationships=path_rels,
152 length=len(path_rels),
153 score=(
154 sum(r.confidence for r in path_rels) / len(path_rels)
155 if path_rels
156 else 1.0
157 ),
158 )
160 for rel in self._relationships.get(current, []):
161 if relationship_types and rel.type not in relationship_types:
162 continue
164 neighbor = rel.target.lower()
165 if neighbor not in visited:
166 visited.add(neighbor)
167 queue.append((neighbor, [*path_rels, rel]))
169 return None
171 async def find_all_paths(
172 self,
173 source: str,
174 target: str,
175 max_depth: int = 5,
176 max_paths: int = 10,
177 ) -> list[GraphPath]:
178 source_key = source.lower()
179 target_key = target.lower()
181 if source_key not in self._entities or target_key not in self._entities:
182 return []
184 paths: list[GraphPath] = []
186 def dfs(current: str, path_rels: list[Relationship], visited: set[str]) -> Any:
187 if len(paths) >= max_paths:
188 return
190 if len(path_rels) >= max_depth:
191 return
193 if current == target_key and path_rels:
194 entities = [source]
195 for rel in path_rels:
196 entities.append(rel.target)
198 path = GraphPath(
199 entities=entities,
200 relationships=path_rels,
201 length=len(path_rels),
202 score=sum(r.confidence for r in path_rels) / len(path_rels),
203 )
204 paths.append(path)
205 return
207 for rel in self._relationships.get(current, []):
208 neighbor = rel.target.lower()
209 if neighbor not in visited:
210 visited.add(neighbor)
211 dfs(neighbor, [*path_rels, rel], visited)
212 visited.remove(neighbor)
214 dfs(source_key, [], {source_key})
216 paths.sort(key=lambda p: p.score, reverse=True)
217 return paths
219 async def query_subgraph(
220 self,
221 entity_name: str,
222 depth: int = 2,
223 ) -> tuple[list[Entity], list[Relationship]]:
224 key = entity_name.lower()
225 if key not in self._entities:
226 return [], []
228 entities_found = {key}
229 relationships_found = []
231 current_level = {key}
232 for _ in range(depth):
233 next_level = set()
235 for entity_key in current_level:
236 for rel in self._relationships.get(entity_key, []):
237 relationships_found.append(rel)
238 target = rel.target.lower()
239 entities_found.add(target)
240 next_level.add(target)
242 for rel in self._reverse_relationships.get(entity_key, []):
243 relationships_found.append(rel)
244 source = rel.source.lower()
245 entities_found.add(source)
246 next_level.add(source)
248 current_level = next_level
250 entities = [
251 self._entities[name] for name in entities_found if name in self._entities
252 ]
254 return entities, relationships_found
256 def get_stats(self) -> dict[str, Any]:
257 total_rels = sum(len(rels) for rels in self._relationships.values())
259 entity_counts = {
260 str(etype): len(names) for etype, names in self._entity_index.items()
261 }
263 rel_counts: dict[str, int] = defaultdict(int)
264 for rels in self._relationships.values():
265 for rel in rels:
266 rel_counts[str(rel.type)] += 1
268 return {
269 "total_entities": len(self._entities),
270 "total_relationships": total_rels,
271 "entity_counts": entity_counts,
272 "relationship_counts": dict(rel_counts),
273 "created_at": self._created_at.isoformat(),
274 "updated_at": self._updated_at.isoformat(),
275 }
277 def __len__(self) -> int:
278 return len(self._entities)
280 def __repr__(self) -> str:
281 return f"KnowledgeGraph(entities={len(self._entities)}, relationships={sum(len(r) for r in self._relationships.values())})"
284class KnowledgeGraphBuilder:
285 """Builder for constructing knowledge graphs from text/documents."""
287 def __init__(
288 self,
289 llm_client: LLMClientProtocol,
290 entity_types: list[EntityType | str] | None = None,
291 relationship_types: list[RelationshipType | str] | None = None,
292 min_confidence: float = 0.5,
293 ):
294 self.entity_extractor = EntityExtractor(
295 llm_client=llm_client,
296 entity_types=entity_types,
297 min_confidence=min_confidence,
298 )
299 self.relationship_extractor = RelationshipExtractor(
300 llm_client=llm_client,
301 relationship_types=relationship_types,
302 min_confidence=min_confidence,
303 )
305 async def build_from_text(self, text: str) -> KnowledgeGraph:
306 try:
307 import importlib
309 public_kg = getattr(
310 importlib.import_module("lexigram.ai.rag.knowledge_graph"),
311 "KnowledgeGraph",
312 None,
313 )
314 except (ImportError, ModuleNotFoundError, AttributeError):
315 public_kg = None
317 kg = public_kg() if public_kg is not None else KnowledgeGraph()
319 entities = await self.entity_extractor.extract(text)
320 await kg.add_entities(entities)
322 relationships = await self.relationship_extractor.extract(text, entities)
323 await kg.add_relationships(relationships)
325 return kg
327 async def build_from_documents(
328 self,
329 documents: list[str],
330 merge: bool = True,
331 ) -> KnowledgeGraph | list[KnowledgeGraph]:
332 if merge:
333 kg = KnowledgeGraph()
334 for doc in documents:
335 doc_kg = await self.build_from_text(doc)
337 for entity in await doc_kg.get_all_entities():
338 await kg.add_entity(entity)
340 for rel in await doc_kg.get_all_relationships():
341 await kg.add_relationship(rel)
343 try:
344 import importlib
346 _pkg_mod = importlib.import_module("lexigram.ai.rag.knowledge_graph")
347 public_kg = getattr(_pkg_mod, "KnowledgeGraph", None)
348 if public_kg and not isinstance(kg, public_kg):
349 new_kg = public_kg()
350 for entity in await kg.get_all_entities():
351 await new_kg.add_entity(entity)
352 for rel in await kg.get_all_relationships():
353 await new_kg.add_relationship(rel)
354 kg = new_kg
355 except (ImportError, ModuleNotFoundError, AttributeError):
356 pass
358 return kg
359 graphs = []
360 for doc in documents:
361 kg = await self.build_from_text(doc)
362 graphs.append(kg)
364 try:
365 import importlib
367 _pkg_mod = importlib.import_module("lexigram.ai.rag.knowledge_graph")
368 public_kg = getattr(_pkg_mod, "KnowledgeGraph", None)
369 if public_kg:
370 normalized: list[KnowledgeGraph] = []
371 for g in graphs:
372 if isinstance(g, public_kg):
373 normalized.append(g)
374 else:
375 new_kg = public_kg()
376 for entity in await g.get_all_entities():
377 await new_kg.add_entity(entity)
378 for rel in await g.get_all_relationships():
379 await new_kg.add_relationship(rel)
380 normalized.append(new_kg)
381 return normalized
382 except (ImportError, ModuleNotFoundError, AttributeError):
383 pass
385 return graphs
388class GraphRAGIntegration:
389 """Integrates knowledge graph with RAG pipeline for enhanced retrieval."""
391 def __init__(
392 self,
393 knowledge_graph: KnowledgeGraph,
394 vector_store: DocumentVectorStoreProtocol,
395 llm_client: LLMClientProtocol,
396 ):
397 self.kg = knowledge_graph
398 self.vector_store = vector_store
399 self.llm_client = llm_client
401 async def expand_query_with_graph(
402 self,
403 query: str,
404 max_expansions: int = 5,
405 ) -> list[str]:
406 entities = await EntityExtractor(self.llm_client).extract(query)
408 queries = [query]
410 for entity in entities[:3]:
411 neighbors = await self.kg.get_neighbors(entity.name, direction="both")
413 for neighbor in neighbors[: max_expansions - len(queries)]:
414 expanded = f"{query} {neighbor.name}"
415 queries.append(expanded)
417 if len(queries) >= max_expansions:
418 break
420 if len(queries) >= max_expansions:
421 break
423 return queries
425 async def retrieve_with_graph(
426 self,
427 query: str,
428 top_k: int = 5,
429 expand: bool = True,
430 ) -> list[Any]:
431 queries = [query]
433 if expand:
434 queries = await self.expand_query_with_graph(query, max_expansions=3)
436 all_results = []
437 for q in queries:
438 search_result = await self.vector_store.search(q, top_k=top_k) # type: ignore[arg-type]
439 all_results.extend(search_result.unwrap_or([]))
441 seen = set()
442 unique_results = []
443 for result in all_results:
444 doc_id = getattr(result, "id", hash(getattr(result, "text", "")))
445 if doc_id not in seen:
446 seen.add(doc_id)
447 unique_results.append(result)
449 return unique_results[:top_k]