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
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-20 10:12 +0000
1from __future__ import annotations
3import asyncio
4import datetime
5import json
6from collections import defaultdict
7from enum import Enum
8from typing import Any
10from falkordb import FalkorDB
12from .base import GraphStore
13from .models import (
14 CoreEdgeType,
15 CoreNodeLabel,
16 GraphEdge,
17 GraphNode,
18 SubGraph,
19)
21MAX_ROWS = 5000
22GLOBAL_PROJECT = "__global__"
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)
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}")
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}")
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 }
62 async def upsert_node(self, node: GraphNode) -> None:
63 self._validate_node(node)
65 payload = self._node_payload(node)
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 """
81 await self._run_query(query, payload)
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)
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]
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 )
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 )
124 if tasks:
125 await asyncio.gather(*tasks)
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 }
142 async def upsert_edge(self, edge: GraphEdge) -> None:
143 self._validate_edge(edge)
144 payload = self._edge_payload(edge)
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)
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 """
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)
199 tasks = []
200 for edge_type, group_edges in grouped.items():
201 payload = [self._edge_payload(edge) for edge in group_edges]
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"]]
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]
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)
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] = []
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
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]
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
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")
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)
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)
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 )
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
353 def _clean_props(self, props: dict) -> dict:
354 def _normalize_value(v: Any):
355 if v is None:
356 return None
358 if isinstance(v, (str, int, float, bool)):
359 return v
361 if isinstance(v, (datetime.datetime, datetime.date)):
362 return v.isoformat()
364 if isinstance(v, Enum):
365 return v.value
367 if isinstance(v, list):
368 cleaned = []
369 for item in v:
370 val = _normalize_value(item)
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
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
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
397 async def _run_query(self, query, params):
398 async with self._semaphore:
399 return await asyncio.to_thread(self._graph.query, query, params)