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

1"""Core knowledge graph implementation (migrated from legacy single-file module). 

2 

3This module contains the canonical, package-local implementations and should be 

4imported from `lexigram.ai.rag.knowledge_graph`. 

5""" 

6 

7from __future__ import annotations 

8 

9from collections import defaultdict, deque 

10from datetime import UTC, datetime 

11from typing import TYPE_CHECKING, Any 

12 

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) 

24 

25if TYPE_CHECKING: 

26 from lexigram.contracts.ai import ( 

27 DocumentVectorStoreProtocol, 

28 LLMClientProtocol, 

29 ) 

30 

31 

32class KnowledgeGraph: 

33 """In-memory knowledge graph for storing and querying entities and relationships.""" 

34 

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) 

42 

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) 

48 

49 async def add_entities(self, entities: list[Entity]) -> None: 

50 for entity in entities: 

51 await self.add_entity(entity) 

52 

53 async def add_relationship(self, relationship: Relationship) -> None: 

54 source_key = relationship.source.lower() 

55 target_key = relationship.target.lower() 

56 

57 self._relationships[source_key].append(relationship) 

58 self._reverse_relationships[target_key].append(relationship) 

59 self._updated_at = datetime.now(UTC) 

60 

61 async def add_relationships(self, relationships: list[Relationship]) -> None: 

62 for relationship in relationships: 

63 await self.add_relationship(relationship) 

64 

65 async def get_entity(self, name: str) -> Entity | None: 

66 return self._entities.get(name.lower()) 

67 

68 async def get_all_entities(self) -> list[Entity]: 

69 """Get all entities in the knowledge graph.""" 

70 return list(self._entities.values()) 

71 

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 

78 

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] 

83 

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() 

91 

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) 

97 

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) 

103 

104 return list(neighbors) 

105 

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 = [] 

113 

114 if direction in ("outgoing", "both"): 

115 rels.extend(self._relationships.get(key, [])) 

116 

117 if direction in ("incoming", "both"): 

118 rels.extend(self._reverse_relationships.get(key, [])) 

119 

120 return rels 

121 

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() 

131 

132 if source_key not in self._entities or target_key not in self._entities: 

133 return None 

134 

135 queue: deque[tuple[str, list[Relationship]]] = deque([(source_key, [])]) 

136 visited = {source_key} 

137 

138 while queue: 

139 current, path_rels = queue.popleft() 

140 

141 if len(path_rels) >= max_depth: 

142 continue 

143 

144 if current == target_key: 

145 entities = [source] 

146 for rel in path_rels: 

147 entities.append(rel.target) 

148 

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 ) 

159 

160 for rel in self._relationships.get(current, []): 

161 if relationship_types and rel.type not in relationship_types: 

162 continue 

163 

164 neighbor = rel.target.lower() 

165 if neighbor not in visited: 

166 visited.add(neighbor) 

167 queue.append((neighbor, [*path_rels, rel])) 

168 

169 return None 

170 

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() 

180 

181 if source_key not in self._entities or target_key not in self._entities: 

182 return [] 

183 

184 paths: list[GraphPath] = [] 

185 

186 def dfs(current: str, path_rels: list[Relationship], visited: set[str]) -> Any: 

187 if len(paths) >= max_paths: 

188 return 

189 

190 if len(path_rels) >= max_depth: 

191 return 

192 

193 if current == target_key and path_rels: 

194 entities = [source] 

195 for rel in path_rels: 

196 entities.append(rel.target) 

197 

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 

206 

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) 

213 

214 dfs(source_key, [], {source_key}) 

215 

216 paths.sort(key=lambda p: p.score, reverse=True) 

217 return paths 

218 

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 [], [] 

227 

228 entities_found = {key} 

229 relationships_found = [] 

230 

231 current_level = {key} 

232 for _ in range(depth): 

233 next_level = set() 

234 

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) 

241 

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) 

247 

248 current_level = next_level 

249 

250 entities = [ 

251 self._entities[name] for name in entities_found if name in self._entities 

252 ] 

253 

254 return entities, relationships_found 

255 

256 def get_stats(self) -> dict[str, Any]: 

257 total_rels = sum(len(rels) for rels in self._relationships.values()) 

258 

259 entity_counts = { 

260 str(etype): len(names) for etype, names in self._entity_index.items() 

261 } 

262 

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 

267 

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 } 

276 

277 def __len__(self) -> int: 

278 return len(self._entities) 

279 

280 def __repr__(self) -> str: 

281 return f"KnowledgeGraph(entities={len(self._entities)}, relationships={sum(len(r) for r in self._relationships.values())})" 

282 

283 

284class KnowledgeGraphBuilder: 

285 """Builder for constructing knowledge graphs from text/documents.""" 

286 

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 ) 

304 

305 async def build_from_text(self, text: str) -> KnowledgeGraph: 

306 try: 

307 import importlib 

308 

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 

316 

317 kg = public_kg() if public_kg is not None else KnowledgeGraph() 

318 

319 entities = await self.entity_extractor.extract(text) 

320 await kg.add_entities(entities) 

321 

322 relationships = await self.relationship_extractor.extract(text, entities) 

323 await kg.add_relationships(relationships) 

324 

325 return kg 

326 

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) 

336 

337 for entity in await doc_kg.get_all_entities(): 

338 await kg.add_entity(entity) 

339 

340 for rel in await doc_kg.get_all_relationships(): 

341 await kg.add_relationship(rel) 

342 

343 try: 

344 import importlib 

345 

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 

357 

358 return kg 

359 graphs = [] 

360 for doc in documents: 

361 kg = await self.build_from_text(doc) 

362 graphs.append(kg) 

363 

364 try: 

365 import importlib 

366 

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 

384 

385 return graphs 

386 

387 

388class GraphRAGIntegration: 

389 """Integrates knowledge graph with RAG pipeline for enhanced retrieval.""" 

390 

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 

400 

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) 

407 

408 queries = [query] 

409 

410 for entity in entities[:3]: 

411 neighbors = await self.kg.get_neighbors(entity.name, direction="both") 

412 

413 for neighbor in neighbors[: max_expansions - len(queries)]: 

414 expanded = f"{query} {neighbor.name}" 

415 queries.append(expanded) 

416 

417 if len(queries) >= max_expansions: 

418 break 

419 

420 if len(queries) >= max_expansions: 

421 break 

422 

423 return queries 

424 

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] 

432 

433 if expand: 

434 queries = await self.expand_query_with_graph(query, max_expansions=3) 

435 

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([])) 

440 

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) 

448 

449 return unique_results[:top_k]