Coverage for src/qdrant_loader_core/graph/falkor_store.py: 97%

211 statements  

« prev     ^ index     » next       coverage.py v7.15.0, created at 2026-07-20 10:12 +0000

1from __future__ import annotations 

2 

3import asyncio 

4import datetime 

5import json 

6from collections import defaultdict 

7from enum import Enum 

8from typing import Any 

9 

10from falkordb import FalkorDB 

11 

12from .base import GraphStore 

13from .models import ( 

14 CoreEdgeType, 

15 CoreNodeLabel, 

16 GraphEdge, 

17 GraphNode, 

18 SubGraph, 

19) 

20 

21MAX_ROWS = 5000 

22GLOBAL_PROJECT = "__global__" 

23 

24 

25class FalkorGraphStore(GraphStore): 

26 def __init__( 

27 self, 

28 host="localhost", 

29 port=6379, 

30 password=None, 

31 graph_name="default_graph", 

32 max_connections: int = 10, 

33 ): 

34 if max_connections < 1: 

35 raise ValueError("max_connections must be >= 1") 

36 self._semaphore = asyncio.Semaphore(max_connections) 

37 self._db = FalkorDB(host=host, port=port, password=password) 

38 self._graph = self._db.select_graph(graph_name) 

39 

40 def _validate_node(self, node: GraphNode): 

41 if node.label not in {e.value for e in CoreNodeLabel}: 

42 raise ValueError(f"Invalid node label: {node.label}") 

43 

44 def _validate_edge(self, edge: GraphEdge): 

45 if edge.edge_type not in {e.value for e in CoreEdgeType}: 

46 raise ValueError(f"Invalid edge type: {edge.edge_type}") 

47 

48 def _node_payload(self, node: GraphNode) -> dict[str, Any]: 

49 props = self._clean_props(node.properties or {}) 

50 # Pop (not get): leaving project in props would let `SET n += props` overwrite the MERGE key. 

51 props_project = props.pop("project", None) 

52 return { 

53 "id": node.id, 

54 "project": ( 

55 node.project 

56 if node.project is not None 

57 else props_project or GLOBAL_PROJECT 

58 ), 

59 "props": props, 

60 } 

61 

62 async def upsert_node(self, node: GraphNode) -> None: 

63 self._validate_node(node) 

64 

65 payload = self._node_payload(node) 

66 

67 # Match by (id, project) only, not label, so this MERGE can enrich a stub node an edge upsert already created. 

68 if payload["project"] is not None: 

69 query = f""" 

70 MERGE (n {{id: $id, project: $project}}) 

71 SET n:{node.label} 

72 SET n += $props 

73 """ 

74 else: 

75 query = f""" 

76 MERGE (n {{id: $id}}) 

77 SET n:{node.label} 

78 SET n += $props 

79 """ 

80 

81 await self._run_query(query, payload) 

82 

83 async def upsert_nodes_batch(self, nodes: list[GraphNode]) -> None: 

84 if not nodes: 

85 return 

86 for node in nodes: 

87 self._validate_node(node) 

88 grouped: dict[str, list[GraphNode]] = defaultdict(list) 

89 for node in nodes: 

90 grouped[node.label].append(node) 

91 

92 tasks = [] 

93 for label, group_nodes in grouped.items(): 

94 payload = [self._node_payload(node) for node in group_nodes] 

95 with_project = [n for n in payload if n["project"] is not None] 

96 without_project = [n for n in payload if n["project"] is None] 

97 

98 if with_project: 

99 tasks.append( 

100 self._run_query( 

101 f""" 

102 UNWIND $nodes AS node 

103 MERGE (n {{id: node.id, project: node.project}}) 

104 SET n:{label} 

105 SET n += node.props 

106 """, 

107 {"nodes": with_project}, 

108 ) 

109 ) 

110 

111 if without_project: 

112 tasks.append( 

113 self._run_query( 

114 f""" 

115 UNWIND $nodes AS node 

116 MERGE (n {{id: node.id}}) 

117 SET n:{label} 

118 SET n += node.props 

119 """, 

120 {"nodes": without_project}, 

121 ) 

122 ) 

123 

124 if tasks: 

125 await asyncio.gather(*tasks) 

126 

127 def _edge_payload(self, edge: GraphEdge) -> dict[str, Any]: 

128 props = self._clean_props(edge.properties or {}) 

129 # Pop (not get): leaving project in props would let `SET r += props` overwrite the MERGE key. 

130 props_project = props.pop("project", None) 

131 return { 

132 "source": edge.source, 

133 "target": edge.target, 

134 "project": ( 

135 edge.project 

136 if edge.project is not None 

137 else props_project or GLOBAL_PROJECT 

138 ), 

139 "props": props, 

140 } 

141 

142 async def upsert_edge(self, edge: GraphEdge) -> None: 

143 self._validate_edge(edge) 

144 payload = self._edge_payload(edge) 

145 

146 # MERGE (not MATCH): MATCH would silently drop an edge whose target isn't ingested yet. 

147 rel_pattern = ( 

148 f"r:{edge.edge_type} {{kind: $props.kind}}" 

149 if "kind" in payload["props"] 

150 else f"r:{edge.edge_type}" 

151 ) 

152 if payload["project"] is not None: 

153 query = f""" 

154 MERGE (a {{id: $source, project: $project}}) 

155 MERGE (b {{id: $target, project: $project}}) 

156 MERGE (a)-[{rel_pattern}]->(b) 

157 SET r += $props 

158 SET r.project = $project 

159 """ 

160 else: 

161 query = f""" 

162 MERGE (a {{id: $source}}) 

163 MERGE (b {{id: $target}}) 

164 MERGE (a)-[{rel_pattern}]->(b) 

165 SET r += $props 

166 """ 

167 await self._run_query(query, payload) 

168 

169 def _edge_batch_query(self, rel_pattern: str, with_project: bool) -> str: 

170 if with_project: 

171 endpoint_a = "a {id: e.source, project: e.project}" 

172 endpoint_b = "b {id: e.target, project: e.project}" 

173 set_project = "SET r.project = e.project" 

174 else: 

175 endpoint_a = "a {id: e.source}" 

176 endpoint_b = "b {id: e.target}" 

177 set_project = "" 

178 return f""" 

179 UNWIND $edges AS e 

180 MERGE ({endpoint_a}) 

181 MERGE ({endpoint_b}) 

182 MERGE (a)-[{rel_pattern}]->(b) 

183 SET r += e.props 

184 {set_project} 

185 """ 

186 

187 async def upsert_edges_batch( 

188 self, 

189 edges: list[GraphEdge], 

190 ) -> None: 

191 if not edges: 

192 return 

193 for edge in edges: 

194 self._validate_edge(edge) 

195 grouped: dict[str, list[GraphEdge]] = defaultdict(list) 

196 for edge in edges: 

197 grouped[edge.edge_type].append(edge) 

198 

199 tasks = [] 

200 for edge_type, group_edges in grouped.items(): 

201 payload = [self._edge_payload(edge) for edge in group_edges] 

202 

203 typed = [e for e in payload if "kind" in e["props"]] 

204 untyped = [e for e in payload if "kind" not in e["props"]] 

205 

206 for group, rel_pattern in ( 

207 (typed, f"r:{edge_type} {{kind: e.props.kind}}"), 

208 (untyped, f"r:{edge_type}"), 

209 ): 

210 if not group: 

211 continue 

212 with_project = [e for e in group if e["project"] is not None] 

213 without_project = [e for e in group if e["project"] is None] 

214 

215 if with_project: 

216 tasks.append( 

217 self._run_query( 

218 self._edge_batch_query(rel_pattern, with_project=True), 

219 {"edges": with_project}, 

220 ) 

221 ) 

222 if without_project: 

223 tasks.append( 

224 self._run_query( 

225 self._edge_batch_query(rel_pattern, with_project=False), 

226 {"edges": without_project}, 

227 ) 

228 ) 

229 if tasks: 

230 await asyncio.gather(*tasks) 

231 

232 async def neighbors( 

233 self, 

234 node_id: str, 

235 depth: int, 

236 edge_types: list[str] | None = None, 

237 project: str | None = None, 

238 ) -> SubGraph: 

239 edge_filter = "" 

240 if edge_types: 

241 invalid = [ 

242 e for e in edge_types if e not in {et.value for et in CoreEdgeType} 

243 ] 

244 if invalid: 

245 raise ValueError(f"Invalid edge types: {invalid}") 

246 edge_filter = ":" + "|".join(edge_types) 

247 params = {"id": node_id} 

248 if project: 

249 params["project"] = project 

250 query = f""" 

251 MATCH (n {{id: $id, project: $project}})-[r{edge_filter}*1..{depth}]-(m {{project: $project}}) 

252 RETURN n, r, m 

253 LIMIT {MAX_ROWS} 

254 """ 

255 else: 

256 query = f""" 

257 MATCH (n {{id: $id}})-[r{edge_filter}*1..{depth}]-(m) 

258 RETURN n, r, m 

259 LIMIT {MAX_ROWS} 

260 """ 

261 result = await self._run_query(query, params) 

262 nodes_map: dict[str, GraphNode] = {} 

263 internal_node_id_map: dict[int, str] = {} 

264 edges: list[GraphEdge] = [] 

265 

266 def _extract_node(node_value): 

267 if hasattr(node_value, "properties"): 

268 node_id = node_value.properties.get("id") 

269 label = ( 

270 node_value.labels[0] 

271 if getattr(node_value, "labels", None) 

272 else "Unknown" 

273 ) 

274 properties = node_value.properties or {} 

275 elif isinstance(node_value, dict): 

276 node_id = node_value.get("id") 

277 label = node_value.get("label", "Unknown") 

278 properties = node_value 

279 else: 

280 node_id = str(node_value) 

281 label = "Unknown" 

282 properties = {} 

283 return node_id, label, properties 

284 

285 def _normalize_relationships(rels_value): 

286 if rels_value is None: 

287 return [] 

288 if isinstance(rels_value, list): 

289 return rels_value 

290 return [rels_value] 

291 

292 def _resolve_internal_node_id(node_value): 

293 if isinstance(node_value, int): 

294 return internal_node_id_map.get(node_value, str(node_value)) 

295 return node_value 

296 

297 def _extract_edge(rel): 

298 if hasattr(rel, "src_node") and hasattr(rel, "dest_node"): 

299 source = _resolve_internal_node_id(rel.src_node) 

300 target = _resolve_internal_node_id(rel.dest_node) 

301 edge_type = getattr(rel, "relation", getattr(rel, "edge_type", None)) 

302 properties = rel.properties or {} 

303 return GraphEdge( 

304 source=source, 

305 target=target, 

306 edge_type=edge_type, 

307 properties=properties, 

308 ) 

309 if isinstance(rel, dict): 

310 return GraphEdge( 

311 source=rel.get("source"), 

312 target=rel.get("target"), 

313 edge_type=rel.get("edge_type"), 

314 properties=rel.get("properties", {}), 

315 ) 

316 raise ValueError("Unsupported relationship result type") 

317 

318 for row in result.result_set: 

319 if len(row) != 3: 

320 continue 

321 n, rels, m = row 

322 n_id, n_label, n_props = _extract_node(n) 

323 if hasattr(n, "id") and isinstance(n.id, int): 

324 internal_node_id_map[n.id] = n_id 

325 if n_id not in nodes_map: 

326 nodes_map[n_id] = GraphNode(id=n_id, label=n_label, properties=n_props) 

327 

328 m_id, m_label, m_props = _extract_node(m) 

329 if hasattr(m, "id") and isinstance(m.id, int): 

330 internal_node_id_map[m.id] = m_id 

331 if m_id not in nodes_map: 

332 nodes_map[m_id] = GraphNode(id=m_id, label=m_label, properties=m_props) 

333 

334 for rel in _normalize_relationships(rels): 

335 try: 

336 edge = _extract_edge(rel) 

337 except ValueError: 

338 continue 

339 edges.append(edge) 

340 return SubGraph( 

341 nodes=list(nodes_map.values()), 

342 edges=edges, 

343 ) 

344 

345 async def query_cypher( 

346 self, 

347 cypher: str, 

348 params: dict[str, Any], 

349 ) -> list[list[Any]]: 

350 result = await self._run_query(cypher, params or {}) 

351 return result.result_set 

352 

353 def _clean_props(self, props: dict) -> dict: 

354 def _normalize_value(v: Any): 

355 if v is None: 

356 return None 

357 

358 if isinstance(v, (str, int, float, bool)): 

359 return v 

360 

361 if isinstance(v, (datetime.datetime, datetime.date)): 

362 return v.isoformat() 

363 

364 if isinstance(v, Enum): 

365 return v.value 

366 

367 if isinstance(v, list): 

368 cleaned = [] 

369 for item in v: 

370 val = _normalize_value(item) 

371 

372 if isinstance(val, (str, int, float, bool)): 

373 cleaned.append(val) 

374 else: 

375 if val is not None: 

376 cleaned.append(str(val)) 

377 return cleaned 

378 

379 if isinstance(v, dict): 

380 try: 

381 return json.dumps(v, ensure_ascii=False) 

382 except Exception: 

383 return str(v) 

384 try: 

385 return str(v) 

386 except Exception: 

387 return None 

388 

389 clean = {} 

390 for k, v in props.items(): 

391 val = _normalize_value(v) 

392 if val is None: 

393 continue 

394 clean[k] = val 

395 return clean 

396 

397 async def _run_query(self, query, params): 

398 async with self._semaphore: 

399 return await asyncio.to_thread(self._graph.query, query, params)