Repository navigation
Expand file tree
/
Copy pathgraph_read_cache.py
More file actions
464 lines (406 loc) · 18.1 KB
/
Copy pathgraph_read_cache.py
File metadata and controls
464 lines (406 loc) · 18.1 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
"""Per-search read-through cache for graph traversal.
Graph expansion and PPR are intentionally read-heavy: they walk from the same
seed nodes through several structural paths. This helper keeps those reads
inside one search call so each edge/node lookup is paid for once without
changing backend semantics or persisting any cache state.
"""
from __future__ import annotations
from collections.abc import Sequence
from typing import TYPE_CHECKING, Literal
from synaptic.models import Edge, EdgeKind, Node
if TYPE_CHECKING:
from synaptic.protocols import StorageBackend
Direction = Literal["both", "incoming", "outgoing"]
class GraphReadCache:
"""Small async read-through cache scoped to one graph traversal/search."""
__slots__ = (
"_backend",
"_edge_cache",
"_edge_kind_cache",
"_edge_kind_light_cache",
"_edge_light_cache",
"_neighbor_cache",
"_node_cache",
)
def __init__(self, backend: StorageBackend) -> None:
self._backend = backend
self._edge_cache: dict[tuple[str, Direction], list[Edge]] = {}
self._edge_kind_cache: dict[tuple[str, Direction, tuple[str, ...]], list[Edge]] = {}
self._edge_kind_light_cache: dict[tuple[str, Direction, tuple[str, ...]], list[Edge]] = {}
self._edge_light_cache: dict[tuple[str, Direction], list[Edge]] = {}
self._neighbor_cache: dict[tuple[str, int], list[tuple[Node, Edge]]] = {}
self._node_cache: dict[str, Node | None] = {}
async def get_node(self, node_id: str) -> Node | None:
if node_id not in self._node_cache:
self._node_cache[node_id] = await self._backend.get_node(node_id)
return self._node_cache[node_id]
async def get_nodes(self, node_ids: list[str]) -> dict[str, Node]:
"""Return found nodes keyed by id, batch-loading misses when possible."""
if not node_ids:
return {}
missing = [nid for nid in dict.fromkeys(node_ids) if nid not in self._node_cache]
if missing:
get_batch = getattr(self._backend, "get_nodes_batch", None)
if callable(get_batch):
nodes = await get_batch(missing)
found = {node.id: node for node in nodes}
for nid in missing:
self._node_cache[nid] = found.get(nid)
else:
for nid in missing:
self._node_cache[nid] = await self._backend.get_node(nid)
return {
nid: node
for nid in dict.fromkeys(node_ids)
if (node := self._node_cache.get(nid)) is not None
}
async def get_edges(self, node_id: str, *, direction: Direction = "both") -> list[Edge]:
return (await self.get_edges_many([node_id], direction=direction)).get(node_id, [])
async def get_edges_many(
self, node_ids: list[str], *, direction: Direction = "both"
) -> dict[str, list[Edge]]:
"""Return edge lists for multiple nodes, using backend batch reads if available."""
unique_ids = list(dict.fromkeys(node_ids))
if not unique_ids:
return {}
result: dict[str, list[Edge]] = {}
missing: list[str] = []
for node_id in unique_ids:
cached = self._cached_edges(node_id, direction)
if cached is None:
missing.append(node_id)
else:
result[node_id] = cached
if missing:
get_batch = getattr(self._backend, "get_edges_batch", None)
if callable(get_batch):
fetched = await get_batch(missing, direction=direction)
for node_id in missing:
edges = list(fetched.get(node_id, []))
self._edge_cache[(node_id, direction)] = edges
result[node_id] = edges
else:
for node_id in missing:
edges = await self._backend.get_edges(node_id, direction=direction)
self._edge_cache[(node_id, direction)] = edges
result[node_id] = edges
return result
async def get_edges_many_light(
self, node_ids: list[str], *, direction: Direction = "both"
) -> dict[str, list[Edge]]:
"""Return traversal-only edges, using lightweight backend reads when available.
PPR only needs source/target/kind/weight. SQLite can skip loading and
parsing provenance JSON for that path, while callers that need full
edge metadata continue to use ``get_edges_many``.
"""
unique_ids = list(dict.fromkeys(node_ids))
if not unique_ids:
return {}
result: dict[str, list[Edge]] = {}
missing: list[str] = []
for node_id in unique_ids:
cached = self._cached_edges_light(node_id, direction)
if cached is None:
missing.append(node_id)
else:
result[node_id] = cached
if missing:
get_light = getattr(self._backend, "get_edges_batch_light", None)
if callable(get_light):
fetched = await get_light(missing, direction=direction)
for node_id in missing:
edges = list(fetched.get(node_id, []))
self._edge_light_cache[(node_id, direction)] = edges
result[node_id] = edges
else:
fetched = await self.get_edges_many(missing, direction=direction)
for node_id in missing:
edges = list(fetched.get(node_id, []))
self._edge_light_cache[(node_id, direction)] = edges
result[node_id] = edges
return result
async def get_edges_many_by_kind(
self,
node_ids: list[str],
*,
direction: Direction = "both",
kinds: Sequence[EdgeKind | str],
) -> dict[str, list[Edge]]:
"""Return edges for multiple nodes, limited to the requested kinds.
Backends that expose a filtered batch read avoid materializing noisy
edge kinds just so GraphExpander can discard them. Backends without the
optional method degrade to the normal full edge read plus in-memory
filtering, preserving semantics.
"""
unique_ids = list(dict.fromkeys(node_ids))
if not unique_ids:
return {}
kind_key = _kind_key(kinds)
if not kind_key:
return {node_id: [] for node_id in unique_ids}
result: dict[str, list[Edge]] = {}
missing: list[str] = []
for node_id in unique_ids:
cached = self._cached_edges_by_kind(node_id, direction, kind_key)
if cached is None:
missing.append(node_id)
else:
result[node_id] = cached
if missing:
get_filtered = getattr(self._backend, "get_edges_batch_filtered", None)
if callable(get_filtered):
fetched = await get_filtered(missing, direction=direction, kinds=list(kind_key))
for node_id in missing:
edges = _filter_edges_by_kind(fetched.get(node_id, []), kind_key)
self._edge_kind_cache[(node_id, direction, kind_key)] = edges
result[node_id] = edges
else:
fetched = await self.get_edges_many(missing, direction=direction)
for node_id in missing:
edges = _filter_edges_by_kind(fetched.get(node_id, []), kind_key)
self._edge_kind_cache[(node_id, direction, kind_key)] = edges
result[node_id] = edges
return result
async def get_edges_many_by_kind_light(
self,
node_ids: list[str],
*,
direction: Direction = "both",
kinds: Sequence[EdgeKind | str],
) -> dict[str, list[Edge]]:
"""Return traversal-only edges limited to the requested kinds.
This mirrors ``get_edges_many_by_kind`` for expansion paths that only
need source/target/kind/weight. SQLite can skip loading and parsing
``properties_json`` for those paths; callers that need provenance
metadata should keep using ``get_edges_many_by_kind``.
"""
unique_ids = list(dict.fromkeys(node_ids))
if not unique_ids:
return {}
kind_key = _kind_key(kinds)
if not kind_key:
return {node_id: [] for node_id in unique_ids}
result: dict[str, list[Edge]] = {}
missing: list[str] = []
for node_id in unique_ids:
cached = self._cached_edges_by_kind_light(node_id, direction, kind_key)
if cached is None:
missing.append(node_id)
else:
result[node_id] = cached
if missing:
get_filtered_light = getattr(self._backend, "get_edges_batch_filtered_light", None)
if callable(get_filtered_light):
fetched = await get_filtered_light(
missing,
direction=direction,
kinds=list(kind_key),
)
for node_id in missing:
edges = _filter_edges_by_kind(fetched.get(node_id, []), kind_key)
self._edge_kind_light_cache[(node_id, direction, kind_key)] = edges
result[node_id] = edges
else:
fetched = await self.get_edges_many_by_kind(
missing,
direction=direction,
kinds=kind_key,
)
for node_id in missing:
edges = list(fetched.get(node_id, []))
self._edge_kind_light_cache[(node_id, direction, kind_key)] = edges
result[node_id] = edges
return result
async def get_edges_many_by_kind_selective_light(
self,
node_ids: list[str],
*,
direction: Direction = "both",
light_kinds: Sequence[EdgeKind | str],
full_kinds: Sequence[EdgeKind | str],
) -> dict[str, list[Edge]]:
"""Return mixed metadata edges in one filtered batch read.
``light_kinds`` are materialized without properties, while
``full_kinds`` keep provenance metadata. This is useful for relation
expansion where generic RELATED edges only need traversal fields, but
typed OpenIE edges need ``is_openie`` and ``confidence``.
"""
unique_ids = list(dict.fromkeys(node_ids))
if not unique_ids:
return {}
light_key = _kind_key(light_kinds)
full_key = _kind_key(full_kinds)
if not light_key and not full_key:
return {node_id: [] for node_id in unique_ids}
result: dict[str, list[Edge]] = {}
missing: list[str] = []
for node_id in unique_ids:
light_edges = (
self._cached_edges_by_kind_light(node_id, direction, light_key) if light_key else []
)
full_edges = (
self._cached_edges_by_kind(node_id, direction, full_key) if full_key else []
)
if light_edges is None or full_edges is None:
missing.append(node_id)
else:
result[node_id] = _merge_edges(list(light_edges), list(full_edges))
if missing:
get_selective = getattr(self._backend, "get_edges_batch_filtered_selective_light", None)
if callable(get_selective):
fetched = await get_selective(
missing,
direction=direction,
light_kinds=list(light_key),
full_kinds=list(full_key),
)
else:
all_fetched = await self.get_edges_many_by_kind(
missing,
direction=direction,
kinds=[*light_key, *full_key],
)
fetched = {
node_id: _strip_light_kind_properties(
list(all_fetched.get(node_id, [])),
light_key,
)
for node_id in missing
}
for node_id in missing:
edges = list(fetched.get(node_id, []))
if light_key:
self._edge_kind_light_cache[(node_id, direction, light_key)] = (
_filter_edges_by_kind(edges, light_key)
)
if full_key:
self._edge_kind_cache[(node_id, direction, full_key)] = _filter_edges_by_kind(
edges, full_key
)
result[node_id] = edges
return result
def _cached_edges(self, node_id: str, direction: Direction) -> list[Edge] | None:
key = (node_id, direction)
cached = self._edge_cache.get(key)
if cached is None:
if direction != "both":
both = self._edge_cache.get((node_id, "both"))
if both is not None:
cached = _filter_edges(node_id, both, direction)
self._edge_cache[key] = cached
return cached
else:
outgoing = self._edge_cache.get((node_id, "outgoing"))
incoming = self._edge_cache.get((node_id, "incoming"))
if outgoing is not None and incoming is not None:
cached = _merge_edges(outgoing, incoming)
self._edge_cache[key] = cached
return cached
return cached
def _cached_edges_light(self, node_id: str, direction: Direction) -> list[Edge] | None:
full = self._cached_edges(node_id, direction)
if full is not None:
return full
cached = self._edge_light_cache.get((node_id, direction))
if cached is None:
if direction != "both":
both = self._edge_light_cache.get((node_id, "both"))
if both is not None:
cached = _filter_edges(node_id, both, direction)
self._edge_light_cache[(node_id, direction)] = cached
return cached
else:
outgoing = self._edge_light_cache.get((node_id, "outgoing"))
incoming = self._edge_light_cache.get((node_id, "incoming"))
if outgoing is not None and incoming is not None:
cached = _merge_edges(outgoing, incoming)
self._edge_light_cache[(node_id, direction)] = cached
return cached
return cached
def _cached_edges_by_kind(
self, node_id: str, direction: Direction, kind_key: tuple[str, ...]
) -> list[Edge] | None:
full = self._cached_edges(node_id, direction)
if full is not None:
return _filter_edges_by_kind(full, kind_key)
cached = self._edge_kind_cache.get((node_id, direction, kind_key))
if cached is not None:
return cached
if direction != "both":
both = self._edge_kind_cache.get((node_id, "both", kind_key))
if both is not None:
cached = _filter_edges(node_id, both, direction)
self._edge_kind_cache[(node_id, direction, kind_key)] = cached
return cached
return None
def _cached_edges_by_kind_light(
self, node_id: str, direction: Direction, kind_key: tuple[str, ...]
) -> list[Edge] | None:
light = self._cached_edges_light(node_id, direction)
if light is not None:
return _filter_edges_by_kind(light, kind_key)
full = self._cached_edges_by_kind(node_id, direction, kind_key)
if full is not None:
return full
cached = self._edge_kind_light_cache.get((node_id, direction, kind_key))
if cached is not None:
return cached
if direction != "both":
both = self._edge_kind_light_cache.get((node_id, "both", kind_key))
if both is not None:
cached = _filter_edges(node_id, both, direction)
self._edge_kind_light_cache[(node_id, direction, kind_key)] = cached
return cached
return None
async def get_neighbors(self, node_id: str, *, depth: int = 1) -> list[tuple[Node, Edge]]:
key = (node_id, depth)
cached = self._neighbor_cache.get(key)
if cached is None:
cached = await self._backend.get_neighbors(node_id, depth=depth)
self._neighbor_cache[key] = cached
for node, _edge in cached:
self._node_cache.setdefault(node.id, node)
return cached
def _filter_edges(node_id: str, edges: list[Edge], direction: Direction) -> list[Edge]:
if direction == "outgoing":
return [edge for edge in edges if edge.source_id == node_id]
if direction == "incoming":
return [edge for edge in edges if edge.target_id == node_id]
return edges
def _merge_edges(first: list[Edge], second: list[Edge]) -> list[Edge]:
merged: list[Edge] = []
seen: set[str] = set()
for edge in first + second:
if edge.id in seen:
continue
seen.add(edge.id)
merged.append(edge)
return merged
def _kind_key(kinds: Sequence[EdgeKind | str]) -> tuple[str, ...]:
return tuple(
sorted({kind.value if isinstance(kind, EdgeKind) else str(kind) for kind in kinds})
)
def _filter_edges_by_kind(edges: Sequence[Edge], kind_key: tuple[str, ...]) -> list[Edge]:
kind_set = set(kind_key)
return [edge for edge in edges if edge.kind.value in kind_set]
def _strip_light_kind_properties(edges: list[Edge], light_key: tuple[str, ...]) -> list[Edge]:
if not light_key:
return edges
light_set = set(light_key)
stripped: list[Edge] = []
for edge in edges:
if edge.kind.value not in light_set:
stripped.append(edge)
continue
stripped.append(
Edge(
id=edge.id,
source_id=edge.source_id,
target_id=edge.target_id,
kind=edge.kind,
weight=edge.weight,
properties={},
created_at=edge.created_at,
)
)
return stripped