pretensor 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (198) hide show
  1. pretensor/__init__.py +50 -0
  2. pretensor/benchmark/__init__.py +54 -0
  3. pretensor/benchmark/cli.py +294 -0
  4. pretensor/benchmark/fixtures.py +84 -0
  5. pretensor/benchmark/l1/__init__.py +23 -0
  6. pretensor/benchmark/l1/metrics.py +141 -0
  7. pretensor/benchmark/l1/pipeline.py +188 -0
  8. pretensor/benchmark/l1/runner.py +245 -0
  9. pretensor/benchmark/l2/__init__.py +27 -0
  10. pretensor/benchmark/l2/gold.py +236 -0
  11. pretensor/benchmark/l2/metrics.py +146 -0
  12. pretensor/benchmark/l2/pipeline.py +124 -0
  13. pretensor/benchmark/l2/runner.py +530 -0
  14. pretensor/benchmark/l3/__init__.py +73 -0
  15. pretensor/benchmark/l3/agent.py +316 -0
  16. pretensor/benchmark/l3/db.py +188 -0
  17. pretensor/benchmark/l3/gold.py +85 -0
  18. pretensor/benchmark/l3/llm_client.py +395 -0
  19. pretensor/benchmark/l3/mcp_client.py +357 -0
  20. pretensor/benchmark/l3/pretensor_runner.py +456 -0
  21. pretensor/benchmark/l3/prompt.py +132 -0
  22. pretensor/benchmark/l3/runner.py +358 -0
  23. pretensor/benchmark/l3/sql_equivalence.py +176 -0
  24. pretensor/benchmark/release_gate.py +448 -0
  25. pretensor/benchmark/results.py +298 -0
  26. pretensor/benchmark/runner.py +109 -0
  27. pretensor/cli/__init__.py +1 -0
  28. pretensor/cli/commands/_source_runner.py +147 -0
  29. pretensor/cli/commands/analyze.py +201 -0
  30. pretensor/cli/commands/connections/__init__.py +7 -0
  31. pretensor/cli/commands/connections/add_remove.py +126 -0
  32. pretensor/cli/commands/connections/register.py +12 -0
  33. pretensor/cli/commands/export.py +131 -0
  34. pretensor/cli/commands/index.py +559 -0
  35. pretensor/cli/commands/list.py +76 -0
  36. pretensor/cli/commands/quickstart.py +207 -0
  37. pretensor/cli/commands/reindex.py +646 -0
  38. pretensor/cli/commands/semantic.py +190 -0
  39. pretensor/cli/commands/serve.py +144 -0
  40. pretensor/cli/commands/sync_grants.py +149 -0
  41. pretensor/cli/commands/validate.py +176 -0
  42. pretensor/cli/config_file.py +442 -0
  43. pretensor/cli/constants.py +10 -0
  44. pretensor/cli/dbt_enrichment.py +96 -0
  45. pretensor/cli/main.py +109 -0
  46. pretensor/cli/paths.py +43 -0
  47. pretensor/cli/plugin.py +52 -0
  48. pretensor/config.py +226 -0
  49. pretensor/connectors/__init__.py +29 -0
  50. pretensor/connectors/base.py +165 -0
  51. pretensor/connectors/bigquery.py +468 -0
  52. pretensor/connectors/inspect.py +321 -0
  53. pretensor/connectors/lineage_sqlglot.py +97 -0
  54. pretensor/connectors/models.py +130 -0
  55. pretensor/connectors/mysql.py +402 -0
  56. pretensor/connectors/pg_array_parse.py +53 -0
  57. pretensor/connectors/postgres.py +938 -0
  58. pretensor/connectors/registry.py +93 -0
  59. pretensor/connectors/snapshot.py +244 -0
  60. pretensor/connectors/snowflake.py +908 -0
  61. pretensor/core/__init__.py +1 -0
  62. pretensor/core/builder.py +307 -0
  63. pretensor/core/dsn_crypto.py +51 -0
  64. pretensor/core/graph_schema_manager.py +246 -0
  65. pretensor/core/graph_store.py +1226 -0
  66. pretensor/core/ids.py +101 -0
  67. pretensor/core/portable_export.py +276 -0
  68. pretensor/core/query_runner.py +67 -0
  69. pretensor/core/registry.py +209 -0
  70. pretensor/core/schema.py +473 -0
  71. pretensor/core/secure_io.py +93 -0
  72. pretensor/core/store.py +469 -0
  73. pretensor/enrichment/__init__.py +1 -0
  74. pretensor/enrichment/analyze/__init__.py +0 -0
  75. pretensor/enrichment/analyze/classify.py +49 -0
  76. pretensor/enrichment/analyze/extract_python.py +196 -0
  77. pretensor/enrichment/analyze/parse.py +141 -0
  78. pretensor/enrichment/analyze/pipeline.py +195 -0
  79. pretensor/enrichment/analyze/summary.py +38 -0
  80. pretensor/enrichment/analyze/walker.py +98 -0
  81. pretensor/enrichment/analyze/writers.py +214 -0
  82. pretensor/enrichment/dbt/__init__.py +30 -0
  83. pretensor/enrichment/dbt/lineage.py +100 -0
  84. pretensor/enrichment/dbt/manifest.py +300 -0
  85. pretensor/enrichment/dbt/metadata.py +263 -0
  86. pretensor/enrichment/dbt/pipeline.py +77 -0
  87. pretensor/enrichment/dbt/resolution.py +101 -0
  88. pretensor/enrichment/dbt/signals.py +305 -0
  89. pretensor/entities/__init__.py +27 -0
  90. pretensor/entities/builder.py +63 -0
  91. pretensor/entities/classifier.py +383 -0
  92. pretensor/entities/llm_extract.py +66 -0
  93. pretensor/errors.py +35 -0
  94. pretensor/graph_models/__init__.py +17 -0
  95. pretensor/graph_models/base.py +11 -0
  96. pretensor/graph_models/consumer.py +71 -0
  97. pretensor/graph_models/edge.py +35 -0
  98. pretensor/graph_models/entity.py +21 -0
  99. pretensor/graph_models/node.py +79 -0
  100. pretensor/graph_models/relationship.py +33 -0
  101. pretensor/integrations/__init__.py +42 -0
  102. pretensor/integrations/_base.py +138 -0
  103. pretensor/integrations/google_adk.py +49 -0
  104. pretensor/integrations/langchain.py +55 -0
  105. pretensor/integrations/llamaindex.py +53 -0
  106. pretensor/intelligence/__init__.py +33 -0
  107. pretensor/intelligence/cluster_labeler.py +425 -0
  108. pretensor/intelligence/clustering.py +168 -0
  109. pretensor/intelligence/combining.py +32 -0
  110. pretensor/intelligence/discovery.py +114 -0
  111. pretensor/intelligence/embeddings.py +317 -0
  112. pretensor/intelligence/graph_export.py +200 -0
  113. pretensor/intelligence/heuristic.py +544 -0
  114. pretensor/intelligence/join_paths/__init__.py +130 -0
  115. pretensor/intelligence/join_paths/on_demand.py +516 -0
  116. pretensor/intelligence/join_paths/storage.py +70 -0
  117. pretensor/intelligence/llm_infer.py +78 -0
  118. pretensor/intelligence/llm_runtime.py +62 -0
  119. pretensor/intelligence/metric_templates.py +193 -0
  120. pretensor/intelligence/pipeline.py +364 -0
  121. pretensor/intelligence/role_exemplars.py +263 -0
  122. pretensor/intelligence/schema_classification.py +360 -0
  123. pretensor/intelligence/scoring.py +76 -0
  124. pretensor/intelligence/semantic.py +240 -0
  125. pretensor/intelligence/shadow_alias.py +101 -0
  126. pretensor/intelligence/statistical.py +50 -0
  127. pretensor/intelligence/steps.py +191 -0
  128. pretensor/intelligence/steps_embedding.py +168 -0
  129. pretensor/introspection/__init__.py +6 -0
  130. pretensor/introspection/inspector.py +5 -0
  131. pretensor/introspection/models/__init__.py +0 -0
  132. pretensor/introspection/models/base.py +5 -0
  133. pretensor/introspection/models/config.py +237 -0
  134. pretensor/introspection/models/dsn.py +550 -0
  135. pretensor/introspection/models/plan.py +116 -0
  136. pretensor/introspection/models/schema.py +10 -0
  137. pretensor/introspection/models/semantic.py +121 -0
  138. pretensor/introspection/models/validation.py +116 -0
  139. pretensor/introspection/snapshot.py +46 -0
  140. pretensor/mcp/__init__.py +16 -0
  141. pretensor/mcp/config_json.py +24 -0
  142. pretensor/mcp/payload_types.py +274 -0
  143. pretensor/mcp/resources/__init__.py +17 -0
  144. pretensor/mcp/resources/markdown.py +314 -0
  145. pretensor/mcp/server.py +285 -0
  146. pretensor/mcp/service.py +49 -0
  147. pretensor/mcp/service_context.py +142 -0
  148. pretensor/mcp/service_registry.py +294 -0
  149. pretensor/mcp/store_cache.py +43 -0
  150. pretensor/mcp/tool_registry.py +136 -0
  151. pretensor/mcp/tools/__init__.py +1 -0
  152. pretensor/mcp/tools/_rank.py +244 -0
  153. pretensor/mcp/tools/_timed.py +26 -0
  154. pretensor/mcp/tools/compile_metric.py +144 -0
  155. pretensor/mcp/tools/consumers.py +161 -0
  156. pretensor/mcp/tools/context.py +1121 -0
  157. pretensor/mcp/tools/cypher.py +509 -0
  158. pretensor/mcp/tools/detect_changes.py +254 -0
  159. pretensor/mcp/tools/impact.py +271 -0
  160. pretensor/mcp/tools/list.py +131 -0
  161. pretensor/mcp/tools/schema.py +170 -0
  162. pretensor/mcp/tools/search.py +316 -0
  163. pretensor/mcp/tools/semantic_search.py +282 -0
  164. pretensor/mcp/tools/traverse.py +1027 -0
  165. pretensor/mcp/tools/validate_sql.py +150 -0
  166. pretensor/observability.py +203 -0
  167. pretensor/py.typed +0 -0
  168. pretensor/quickstart/README.md +29 -0
  169. pretensor/quickstart/__init__.py +6 -0
  170. pretensor/quickstart/docker-compose.yml +18 -0
  171. pretensor/quickstart/pagila_data.sql +63 -0
  172. pretensor/quickstart/pagila_ddl.sql +92 -0
  173. pretensor/search/__init__.py +6 -0
  174. pretensor/search/base.py +80 -0
  175. pretensor/search/index.py +435 -0
  176. pretensor/semantic/__init__.py +24 -0
  177. pretensor/semantic/base.py +123 -0
  178. pretensor/semantic/compiler.py +487 -0
  179. pretensor/semantic/yaml_layer.py +180 -0
  180. pretensor/skills/__init__.py +5 -0
  181. pretensor/skills/generator.py +235 -0
  182. pretensor/staleness/__init__.py +15 -0
  183. pretensor/staleness/graph_patcher.py +355 -0
  184. pretensor/staleness/impact_analyzer.py +162 -0
  185. pretensor/staleness/snapshot_store.py +38 -0
  186. pretensor/validation/__init__.py +9 -0
  187. pretensor/validation/query_validator.py +436 -0
  188. pretensor/visibility/__init__.py +23 -0
  189. pretensor/visibility/config.py +126 -0
  190. pretensor/visibility/filter.py +143 -0
  191. pretensor/visibility/kuzu_helpers.py +32 -0
  192. pretensor/visibility/runtime.py +36 -0
  193. pretensor/visibility/sync_grants.py +188 -0
  194. pretensor-0.1.0.dist-info/METADATA +251 -0
  195. pretensor-0.1.0.dist-info/RECORD +198 -0
  196. pretensor-0.1.0.dist-info/WHEEL +4 -0
  197. pretensor-0.1.0.dist-info/entry_points.txt +2 -0
  198. pretensor-0.1.0.dist-info/licenses/LICENSE +21 -0
@@ -0,0 +1,516 @@
1
+ """Graph adjacency, Dijkstra + Yen's K-shortest path scoring."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hashlib
6
+ import heapq
7
+ import json
8
+ from collections.abc import Iterable
9
+ from dataclasses import dataclass
10
+ from typing import Any, Literal
11
+
12
+ from pretensor.core.store import KuzuStore
13
+
14
+ EdgeKind = Literal["fk", "inferred"]
15
+
16
+ _FK_CONFIDENCE = 1.0
17
+ _MAX_PATHS_PER_PAIR = 12
18
+ # Default Yen's K. The MCP traverse tool exposes this as ``top_k``.
19
+ _DEFAULT_TOP_K = 3
20
+ # Default cap on inferred-join hops per path; enforced inside Dijkstra so the
21
+ # search itself prunes weak chains rather than over-enumerating then filtering.
22
+ _DEFAULT_MAX_INFERRED_HOPS = 2
23
+
24
+
25
+ def edge_cost(edge: AdjEdge) -> float:
26
+ """Lower is better. Hop baseline + inferred penalty + low-confidence penalty.
27
+
28
+ A single inferred hop must cost more than an extra FK hop, else short
29
+ inferred paths hide longer authoritative chains (e.g. pagila
30
+ ``film → inventory → customer`` via an inferred ``store_id`` shortcut
31
+ winning over the 3-hop all-FK ``film → inventory → rental → customer``).
32
+ With ``inferred_penalty=1.5``, a 2-hop-with-1-inferred path costs
33
+ ``1.0 + 2.5 = 3.5``, losing to the 3-hop FK chain at ``3.0`` by a
34
+ comfortable margin. Inferred paths only win when the FK alternative
35
+ is ≥4 hops or absent.
36
+ """
37
+ base = 1.0
38
+ inferred_penalty = 1.5 if edge.kind == "inferred" else 0.0
39
+ low_conf_penalty = max(0.0, 0.7 - edge.confidence)
40
+ return base + inferred_penalty + low_conf_penalty
41
+
42
+
43
+ @dataclass(frozen=True, slots=True)
44
+ class JoinStep:
45
+ """One hop along a join path (table → table via columns)."""
46
+
47
+ from_schema: str
48
+ from_table: str
49
+ to_schema: str
50
+ to_table: str
51
+ from_column: str
52
+ to_column: str
53
+ edge_type: EdgeKind
54
+ confidence: float
55
+
56
+
57
+ @dataclass(frozen=True, slots=True)
58
+ class StoredJoinPath:
59
+ """Materialized path ready for MCP JSON."""
60
+
61
+ path_id: str
62
+ from_table_id: str
63
+ to_table_id: str
64
+ depth: int
65
+ confidence: float
66
+ ambiguous: bool
67
+ steps: tuple[JoinStep, ...]
68
+ semantic_label: str
69
+ stale: bool = False
70
+ # Dijkstra cost (lower = better). Persisted on-disk paths predating Yen
71
+ # leave this at 0.0; in-memory paths from ``best_paths`` carry the real cost.
72
+ cost: float = 0.0
73
+
74
+
75
+ @dataclass(frozen=True, slots=True)
76
+ class AdjEdge:
77
+ to_id: str
78
+ source_column: str
79
+ target_column: str
80
+ kind: EdgeKind
81
+ confidence: float
82
+
83
+
84
+ def table_meta(store: KuzuStore, database_key: str) -> dict[str, tuple[str, str]]:
85
+ rows = store.query_all_rows(
86
+ """
87
+ MATCH (t:SchemaTable)
88
+ WHERE t.database = $db
89
+ RETURN t.node_id, t.schema_name, t.table_name
90
+ """,
91
+ {"db": database_key},
92
+ )
93
+ return {str(r[0]): (str(r[1]), str(r[2])) for r in rows if r[0] is not None}
94
+
95
+
96
+ def table_to_cluster(store: KuzuStore, database_key: str) -> dict[str, str]:
97
+ rows = store.query_all_rows(
98
+ """
99
+ MATCH (t:SchemaTable)-[:IN_CLUSTER]->(c:Cluster)
100
+ WHERE c.database_key = $db
101
+ RETURN t.node_id, c.node_id
102
+ """,
103
+ {"db": database_key},
104
+ )
105
+ return {str(r[0]): str(r[1]) for r in rows if r[0] is not None and r[1] is not None}
106
+
107
+
108
+ def build_adjacency(store: KuzuStore, database_key: str) -> dict[str, list[AdjEdge]]:
109
+ adj: dict[str, list[AdjEdge]] = {}
110
+ fk_rows = store.query_all_rows(
111
+ """
112
+ MATCH (a:SchemaTable)-[r:FK_REFERENCES]->(b:SchemaTable)
113
+ WHERE a.database = $db AND b.database = $db
114
+ RETURN a.node_id, b.node_id, r.source_column, r.target_column
115
+ """,
116
+ {"db": database_key},
117
+ )
118
+ for a, b, sc, tc in fk_rows:
119
+ if a is None or b is None:
120
+ continue
121
+ sa, sb = str(a), str(b)
122
+ s_col = str(sc) if sc is not None else ""
123
+ t_col = str(tc) if tc is not None else ""
124
+ adj.setdefault(sa, []).append(AdjEdge(sb, s_col, t_col, "fk", _FK_CONFIDENCE))
125
+ adj.setdefault(sb, []).append(AdjEdge(sa, t_col, s_col, "fk", _FK_CONFIDENCE))
126
+
127
+ inf_rows = store.query_all_rows(
128
+ """
129
+ MATCH (a:SchemaTable)-[r:INFERRED_JOIN]->(b:SchemaTable)
130
+ WHERE a.database = $db AND b.database = $db
131
+ RETURN a.node_id, b.node_id, r.source_column, r.target_column, r.confidence
132
+ """,
133
+ {"db": database_key},
134
+ )
135
+ for row in inf_rows:
136
+ a, b, sc, tc, conf = row
137
+ if a is None or b is None:
138
+ continue
139
+ sa, sb = str(a), str(b)
140
+ w = float(conf) if conf is not None else 0.5
141
+ s_col = str(sc) if sc is not None else ""
142
+ t_col = str(tc) if tc is not None else ""
143
+ adj.setdefault(sa, []).append(AdjEdge(sb, s_col, t_col, "inferred", w))
144
+ adj.setdefault(sb, []).append(AdjEdge(sa, t_col, s_col, "inferred", w))
145
+ return adj
146
+
147
+
148
+ def best_path(
149
+ adj: dict[str, list[AdjEdge]],
150
+ table_meta: dict[str, tuple[str, str]],
151
+ start: str,
152
+ goal: str,
153
+ max_depth: int,
154
+ ) -> StoredJoinPath | None:
155
+ """Single best path. Asks Yen for top-2 to detect cost ties for ``ambiguous``."""
156
+ paths = best_paths(adj, table_meta, start, goal, max_depth, top_k=2)
157
+ if not paths:
158
+ return None
159
+ winner = paths[0]
160
+ if len(paths) > 1 and abs(paths[0].cost - paths[1].cost) < 1e-9:
161
+ # Tied-cost peer exists — surface ambiguity so callers re-enumerate.
162
+ from dataclasses import replace
163
+ winner = replace(winner, ambiguous=True)
164
+ else:
165
+ from dataclasses import replace
166
+ winner = replace(winner, ambiguous=False)
167
+ return winner
168
+
169
+
170
+ def best_paths(
171
+ adj: dict[str, list[AdjEdge]],
172
+ table_meta: dict[str, tuple[str, str]],
173
+ start: str,
174
+ goal: str,
175
+ max_depth: int,
176
+ *,
177
+ top_k: int = _DEFAULT_TOP_K,
178
+ edge_kinds: tuple[EdgeKind, ...] | None = None,
179
+ max_inferred_hops: int = _DEFAULT_MAX_INFERRED_HOPS,
180
+ ) -> list[StoredJoinPath]:
181
+ """Return up to ``top_k`` shortest paths under the cost model in :func:`edge_cost`.
182
+
183
+ Uses Yen's K-shortest paths over Dijkstra. ``edge_kinds`` restricts the set
184
+ of allowed edge kinds (``("fk",)`` excludes inferred joins entirely).
185
+ ``max_inferred_hops`` caps inferred-join hops per path.
186
+ """
187
+ if start == goal or top_k < 1:
188
+ return []
189
+ ranked = yen_k_shortest(
190
+ adj,
191
+ start,
192
+ goal,
193
+ max_depth,
194
+ k=top_k,
195
+ edge_kinds=edge_kinds,
196
+ max_inferred_hops=max_inferred_hops,
197
+ )
198
+ if not ranked:
199
+ return []
200
+ out: list[StoredJoinPath] = []
201
+ # Belt-and-suspenders dedup: yen_k's signature dedup is keyed on the edge
202
+ # tuple, but two distinct edge sequences can still project to an identical
203
+ # ``path_id`` (the sorted-JSON hash of the projected steps). Drop collisions
204
+ # at the output boundary so MCP clients never see duplicate ``steps``.
205
+ seen_path_ids: set[str] = set()
206
+ for cost, edges in ranked:
207
+ steps: list[JoinStep] = []
208
+ cur = start
209
+ conf = 1.0
210
+ for e in edges:
211
+ fs, ft = table_meta[cur]
212
+ ts, tt = table_meta[e.to_id]
213
+ steps.append(
214
+ JoinStep(
215
+ from_schema=fs,
216
+ from_table=ft,
217
+ to_schema=ts,
218
+ to_table=tt,
219
+ from_column=e.source_column,
220
+ to_column=e.target_column,
221
+ edge_type=e.kind,
222
+ confidence=e.confidence,
223
+ )
224
+ )
225
+ conf *= e.confidence
226
+ cur = e.to_id
227
+ steps_t = tuple(steps)
228
+ length = len(edges)
229
+ payload = steps_to_json_payload(steps_t)
230
+ raw = json.dumps(payload, sort_keys=True)
231
+ digest = hashlib.sha256(raw.encode()).hexdigest()[:16]
232
+ path_id = f"{start}::{goal}::{digest}"
233
+ if path_id in seen_path_ids:
234
+ continue
235
+ seen_path_ids.add(path_id)
236
+ label = f"{steps[0].from_table} → {steps[-1].to_table} ({length} hop)"
237
+ out.append(
238
+ StoredJoinPath(
239
+ path_id=path_id,
240
+ from_table_id=start,
241
+ to_table_id=goal,
242
+ depth=length,
243
+ confidence=conf,
244
+ cost=cost,
245
+ ambiguous=len(ranked) > 1,
246
+ steps=steps_t,
247
+ semantic_label=label,
248
+ stale=False,
249
+ )
250
+ )
251
+ return out
252
+
253
+
254
+ def dijkstra_path(
255
+ adj: dict[str, list[AdjEdge]],
256
+ start: str,
257
+ goal: str,
258
+ max_depth: int,
259
+ *,
260
+ edge_kinds: tuple[EdgeKind, ...] | None = None,
261
+ max_inferred_hops: int = _DEFAULT_MAX_INFERRED_HOPS,
262
+ blocked_nodes: frozenset[str] = frozenset(),
263
+ blocked_edges: frozenset[tuple[str, str, str, str]] = frozenset(),
264
+ ) -> tuple[float, list[AdjEdge]] | None:
265
+ """Single-source shortest path using :func:`edge_cost`.
266
+
267
+ ``blocked_edges`` entries are ``(from_id, to_id, from_col, to_col)`` tuples
268
+ so Yen's spur step can prohibit specific edges without disabling the
269
+ underlying node entirely.
270
+ """
271
+ if start == goal:
272
+ return (0.0, [])
273
+ if start in blocked_nodes:
274
+ return None
275
+ counter = 0
276
+ # heap entry: (cum_cost, tie, node, edges, inferred_hops)
277
+ heap: list[tuple[float, int, str, list[AdjEdge], int]] = [
278
+ (0.0, counter, start, [], 0)
279
+ ]
280
+ best_cost: dict[str, float] = {start: 0.0}
281
+ while heap:
282
+ cost, _, cur, edges, inf_hops = heapq.heappop(heap)
283
+ if cur == goal:
284
+ return (cost, edges)
285
+ if cost > best_cost.get(cur, cost):
286
+ continue
287
+ if len(edges) >= max_depth:
288
+ continue
289
+ for e in adj.get(cur, []):
290
+ if edge_kinds is not None and e.kind not in edge_kinds:
291
+ continue
292
+ if e.to_id in blocked_nodes:
293
+ continue
294
+ edge_key = (cur, e.to_id, e.source_column, e.target_column)
295
+ if edge_key in blocked_edges:
296
+ continue
297
+ new_inf = inf_hops + (1 if e.kind == "inferred" else 0)
298
+ if new_inf > max_inferred_hops:
299
+ continue
300
+ # Avoid loops: skip if we'd revisit a node already on this path.
301
+ if any(prev.to_id == e.to_id for prev in edges) or e.to_id == start:
302
+ continue
303
+ new_cost = cost + edge_cost(e)
304
+ if new_cost >= best_cost.get(e.to_id, float("inf")):
305
+ # Yen needs to allow equal-cost spurs through the same node, so
306
+ # only prune *strictly* worse routes; track the cheapest seen.
307
+ continue
308
+ best_cost[e.to_id] = new_cost
309
+ counter += 1
310
+ heapq.heappush(
311
+ heap,
312
+ (new_cost, counter, e.to_id, edges + [e], new_inf),
313
+ )
314
+ return None
315
+
316
+
317
+ def yen_k_shortest(
318
+ adj: dict[str, list[AdjEdge]],
319
+ start: str,
320
+ goal: str,
321
+ max_depth: int,
322
+ k: int,
323
+ *,
324
+ edge_kinds: tuple[EdgeKind, ...] | None = None,
325
+ max_inferred_hops: int = _DEFAULT_MAX_INFERRED_HOPS,
326
+ ) -> list[tuple[float, list[AdjEdge]]]:
327
+ """Yen's K-shortest-paths over :func:`dijkstra_path`."""
328
+ if k < 1:
329
+ return []
330
+ first = dijkstra_path(
331
+ adj,
332
+ start,
333
+ goal,
334
+ max_depth,
335
+ edge_kinds=edge_kinds,
336
+ max_inferred_hops=max_inferred_hops,
337
+ )
338
+ if first is None:
339
+ return []
340
+ accepted: list[tuple[float, list[AdjEdge]]] = [first]
341
+ # Min-heap of candidate spur paths: (cost, tie, edges)
342
+ candidates: list[tuple[float, int, list[AdjEdge]]] = []
343
+ seen_signatures: set[tuple[tuple[str, str, str, str, str], ...]] = {
344
+ _path_signature(start, first[1])
345
+ }
346
+ counter = 0
347
+ while len(accepted) < k:
348
+ prev_cost, prev_edges = accepted[-1]
349
+ # Walk each prefix of the previous path, treating the next edge as a
350
+ # spur point. We block (a) the edge taken at the spur from any prior
351
+ # accepted path that shares this prefix, and (b) the prefix nodes
352
+ # themselves (so the spur stays simple).
353
+ for i in range(len(prev_edges)):
354
+ spur_node = start if i == 0 else prev_edges[i - 1].to_id
355
+ root_edges = prev_edges[:i]
356
+ # Block every prefix node *before* the spur — keeping the spur node
357
+ # itself reachable so the new Dijkstra can leave from it.
358
+ root_nodes: set[str] = set()
359
+ if i > 0:
360
+ root_nodes.add(start)
361
+ for re in root_edges[:-1]:
362
+ root_nodes.add(re.to_id)
363
+ blocked_edges: set[tuple[str, str, str, str]] = set()
364
+ for _, edges in accepted:
365
+ if len(edges) > i and edges[:i] == root_edges:
366
+ e = edges[i]
367
+ blocked_edges.add(
368
+ (spur_node, e.to_id, e.source_column, e.target_column)
369
+ )
370
+ spur_path = dijkstra_path(
371
+ adj,
372
+ spur_node,
373
+ goal,
374
+ max_depth - i,
375
+ edge_kinds=edge_kinds,
376
+ max_inferred_hops=max_inferred_hops - sum(
377
+ 1 for e in root_edges if e.kind == "inferred"
378
+ ),
379
+ blocked_nodes=frozenset(root_nodes),
380
+ blocked_edges=frozenset(blocked_edges),
381
+ )
382
+ if spur_path is None:
383
+ continue
384
+ spur_cost, spur_edges = spur_path
385
+ total_edges = root_edges + spur_edges
386
+ sig = _path_signature(start, total_edges)
387
+ if sig in seen_signatures:
388
+ continue
389
+ seen_signatures.add(sig)
390
+ root_cost = sum(edge_cost(e) for e in root_edges)
391
+ counter += 1
392
+ heapq.heappush(
393
+ candidates,
394
+ (root_cost + spur_cost, counter, total_edges),
395
+ )
396
+ if not candidates:
397
+ break
398
+ cost, _, edges = heapq.heappop(candidates)
399
+ accepted.append((cost, edges))
400
+ return accepted
401
+
402
+
403
+ def _path_signature(
404
+ start: str, edges: Iterable[AdjEdge]
405
+ ) -> tuple[tuple[str, str, str, str, str], ...]:
406
+ """Identity tuple for Yen dedup. Includes ``kind`` so parallel FK + inferred
407
+ edges on the same columns aren't conflated (otherwise yen_k would shadow the
408
+ inferred alternative and callers would never see it)."""
409
+ sig: list[tuple[str, str, str, str, str]] = []
410
+ cur = start
411
+ for e in edges:
412
+ sig.append((cur, e.to_id, e.source_column, e.target_column, e.kind))
413
+ cur = e.to_id
414
+ return tuple(sig)
415
+
416
+
417
+ def steps_to_json_payload(steps: tuple[JoinStep, ...]) -> list[dict[str, Any]]:
418
+ return [
419
+ {
420
+ "from_schema": s.from_schema,
421
+ "from_table": s.from_table,
422
+ "to_schema": s.to_schema,
423
+ "to_table": s.to_table,
424
+ "from_column": s.from_column,
425
+ "to_column": s.to_column,
426
+ "edge_type": s.edge_type,
427
+ "confidence": s.confidence,
428
+ }
429
+ for s in steps
430
+ ]
431
+
432
+
433
+ def reverse_stored_path(path: StoredJoinPath) -> StoredJoinPath:
434
+ rev_steps: list[JoinStep] = []
435
+ for s in reversed(path.steps):
436
+ rev_steps.append(
437
+ JoinStep(
438
+ from_schema=s.to_schema,
439
+ from_table=s.to_table,
440
+ to_schema=s.from_schema,
441
+ to_table=s.from_table,
442
+ from_column=s.to_column,
443
+ to_column=s.from_column,
444
+ edge_type=s.edge_type,
445
+ confidence=s.confidence,
446
+ )
447
+ )
448
+ rt = tuple(rev_steps)
449
+ payload = steps_to_json_payload(rt)
450
+ raw = json.dumps(payload, sort_keys=True)
451
+ digest = hashlib.sha256(raw.encode()).hexdigest()[:16]
452
+ path_id = f"{path.to_table_id}::{path.from_table_id}::{digest}"
453
+ label = f"{rev_steps[0].from_table} → {rev_steps[-1].to_table} ({path.depth} hop)"
454
+ return StoredJoinPath(
455
+ path_id=path_id,
456
+ from_table_id=path.to_table_id,
457
+ to_table_id=path.from_table_id,
458
+ depth=path.depth,
459
+ confidence=path.confidence,
460
+ ambiguous=path.ambiguous,
461
+ steps=rt,
462
+ semantic_label=label,
463
+ stale=path.stale,
464
+ cost=path.cost,
465
+ )
466
+
467
+
468
+ def parse_steps_json(raw: str) -> tuple[JoinStep, ...]:
469
+ try:
470
+ data = json.loads(raw)
471
+ except json.JSONDecodeError:
472
+ return ()
473
+ if not isinstance(data, list):
474
+ return ()
475
+ steps: list[JoinStep] = []
476
+ for item in data:
477
+ if not isinstance(item, dict):
478
+ continue
479
+ steps.append(
480
+ JoinStep(
481
+ from_schema=str(item.get("from_schema", "")),
482
+ from_table=str(item.get("from_table", "")),
483
+ to_schema=str(item.get("to_schema", "")),
484
+ to_table=str(item.get("to_table", "")),
485
+ from_column=str(item.get("from_column", "")),
486
+ to_column=str(item.get("to_column", "")),
487
+ edge_type=edge_kind(item.get("edge_type")),
488
+ confidence=float(item.get("confidence", 1.0)),
489
+ )
490
+ )
491
+ return tuple(steps)
492
+
493
+
494
+ def edge_kind(raw: Any) -> EdgeKind:
495
+ if raw == "inferred":
496
+ return "inferred"
497
+ return "fk"
498
+
499
+
500
+ __all__ = [
501
+ "AdjEdge",
502
+ "EdgeKind",
503
+ "JoinStep",
504
+ "StoredJoinPath",
505
+ "best_path",
506
+ "best_paths",
507
+ "build_adjacency",
508
+ "dijkstra_path",
509
+ "edge_cost",
510
+ "parse_steps_json",
511
+ "reverse_stored_path",
512
+ "steps_to_json_payload",
513
+ "table_meta",
514
+ "table_to_cluster",
515
+ "yen_k_shortest",
516
+ ]
@@ -0,0 +1,70 @@
1
+ """Load and persist ``JoinPath`` rows in Kuzu."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+
7
+ from pretensor.core.store import KuzuStore
8
+ from pretensor.intelligence.join_paths.on_demand import (
9
+ StoredJoinPath,
10
+ parse_steps_json,
11
+ steps_to_json_payload,
12
+ )
13
+
14
+
15
+ def load_stored_paths(
16
+ store: KuzuStore,
17
+ database_key: str,
18
+ from_table_id: str,
19
+ to_table_id: str,
20
+ ) -> list[StoredJoinPath]:
21
+ """Return precomputed paths for an ordered table pair."""
22
+ rows = store.query_all_rows(
23
+ """
24
+ MATCH (p:JoinPath)
25
+ WHERE p.database_key = $db
26
+ AND p.from_table_id = $from_id
27
+ AND p.to_table_id = $to_id
28
+ RETURN p.node_id, p.depth, p.confidence, p.ambiguous, p.steps_json,
29
+ p.semantic_label, p.stale
30
+ ORDER BY p.confidence DESC
31
+ """,
32
+ {"db": database_key, "from_id": from_table_id, "to_id": to_table_id},
33
+ )
34
+ out: list[StoredJoinPath] = []
35
+ for row in rows:
36
+ nid, depth, conf, amb, steps_json, sem, stale = row
37
+ steps = parse_steps_json(str(steps_json))
38
+ stale_b = bool(stale) if stale is not None else False
39
+ out.append(
40
+ StoredJoinPath(
41
+ path_id=str(nid),
42
+ from_table_id=from_table_id,
43
+ to_table_id=to_table_id,
44
+ depth=int(depth) if depth is not None else len(steps),
45
+ confidence=float(conf) if conf is not None else 0.0,
46
+ ambiguous=bool(amb),
47
+ steps=steps,
48
+ semantic_label=str(sem or ""),
49
+ stale=stale_b,
50
+ )
51
+ )
52
+ return out
53
+
54
+
55
+ def persist_path(store: KuzuStore, database_key: str, path: StoredJoinPath) -> None:
56
+ payload = steps_to_json_payload(path.steps)
57
+ store.upsert_join_path(
58
+ node_id=path.path_id,
59
+ database_key=database_key,
60
+ from_table_id=path.from_table_id,
61
+ to_table_id=path.to_table_id,
62
+ depth=path.depth,
63
+ confidence=path.confidence,
64
+ ambiguous=path.ambiguous,
65
+ steps_json=json.dumps(payload),
66
+ semantic_label=path.semantic_label,
67
+ )
68
+
69
+
70
+ __all__ = ["load_stored_paths", "persist_path"]
@@ -0,0 +1,78 @@
1
+ """LLM relationship inference Protocol and null client (extension point)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import logging
7
+ from typing import Any, Protocol
8
+
9
+ from pretensor.connectors.models import SchemaSnapshot
10
+ from pretensor.observability import run_timed_async
11
+
12
+ __all__ = [
13
+ "LlmRelationshipClient",
14
+ "NullLlmRelationshipClient",
15
+ "parse_llm_join_response_json",
16
+ ]
17
+
18
+ logger = logging.getLogger(__name__)
19
+
20
+
21
+ class LlmRelationshipClient(Protocol):
22
+ """Async client that returns structured join hypotheses (implement per provider)."""
23
+
24
+ async def suggest_joins(
25
+ self,
26
+ *,
27
+ snapshot_summary: dict[str, Any],
28
+ batch_table_names: list[str],
29
+ ) -> list[dict[str, Any]]:
30
+ """Return dicts with keys: source_table, source_column, target_table, target_column, confidence, reasoning."""
31
+ ...
32
+
33
+
34
+ class NullLlmRelationshipClient:
35
+ """No-op client: returns no candidates (default until a provider is wired)."""
36
+
37
+ async def suggest_joins(
38
+ self,
39
+ *,
40
+ snapshot_summary: dict[str, Any],
41
+ batch_table_names: list[str],
42
+ ) -> list[dict[str, Any]]:
43
+ async def _no_results() -> list[dict[str, Any]]:
44
+ return []
45
+
46
+ return await run_timed_async(
47
+ logger,
48
+ event="llm.suggest_joins",
49
+ callback=_no_results,
50
+ provider="null",
51
+ batch_size=len(batch_table_names),
52
+ snapshot_table_count=len(snapshot_summary.get("tables", []))
53
+ if isinstance(snapshot_summary.get("tables"), list)
54
+ else None,
55
+ )
56
+
57
+
58
+ def _snapshot_summary(snapshot: SchemaSnapshot) -> dict[str, Any]:
59
+ return {
60
+ "connection_name": snapshot.connection_name,
61
+ "database": snapshot.database,
62
+ "tables": [
63
+ {
64
+ "schema": t.schema_name,
65
+ "name": t.name,
66
+ "columns": [c.name for c in t.columns],
67
+ }
68
+ for t in snapshot.tables
69
+ ],
70
+ }
71
+
72
+
73
+ def parse_llm_join_response_json(raw: str) -> list[dict[str, Any]]:
74
+ """Parse a JSON array of join objects from an LLM response body."""
75
+ data = json.loads(raw)
76
+ if not isinstance(data, list):
77
+ raise ValueError("LLM join response must be a JSON array")
78
+ return [x for x in data if isinstance(x, dict)]