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,305 @@
1
+ """dbt exposures and ``sources.json`` freshness → ``SchemaTable`` signals.
2
+
3
+ Writes three independent signals onto matching ``SchemaTable`` rows:
4
+
5
+ * ``has_external_consumers`` — tables backing dbt exposures (and every
6
+ upstream model/source reachable through ``parent_map``) are marked ``true``.
7
+ * ``staleness_status`` / ``staleness_as_of`` — dbt source freshness results
8
+ from ``sources.json`` are persisted verbatim (``pass``/``warn``/``error``)
9
+ with the most recent ``max_loaded_at`` timestamp. This is staleness metadata,
10
+ NOT a rewrite of ``potentially_unused``.
11
+ * ``test_count`` — number of dbt tests attached to each model (via
12
+ ``manifest.tests[*].attached_node``).
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import json
18
+ import logging
19
+ from collections import Counter, deque
20
+ from dataclasses import dataclass
21
+ from pathlib import Path
22
+ from typing import Any, Mapping
23
+
24
+ from pretensor.core.store import KuzuStore
25
+ from pretensor.enrichment.dbt.manifest import DbtManifest
26
+ from pretensor.enrichment.dbt.resolution import resolve_dbt_parent_node_id
27
+
28
+ __all__ = ["DbtSignalsWriteStats", "write_dbt_signals"]
29
+
30
+ logger = logging.getLogger(__name__)
31
+
32
+ _VALID_STATUSES = frozenset({"pass", "warn", "error", "runtime error"})
33
+ _STALE_STATUSES = frozenset({"warn", "error", "runtime error"})
34
+
35
+
36
+ @dataclass(frozen=True, slots=True)
37
+ class DbtSignalsWriteStats:
38
+ """Counts from ``write_dbt_signals`` for CLI summaries."""
39
+
40
+ exposures_marked: int
41
+ freshness_rows_applied: int
42
+ tests_counted: int
43
+
44
+
45
+ def _as_str(value: Any) -> str | None:
46
+ if value is None:
47
+ return None
48
+ if isinstance(value, str):
49
+ return value
50
+ return str(value)
51
+
52
+
53
+ def _iter_sources_json_results(payload: Mapping[str, Any]) -> list[Mapping[str, Any]]:
54
+ """Normalize v3 ``results[]`` and legacy ``sources`` map shapes."""
55
+ results = payload.get("results")
56
+ if isinstance(results, list):
57
+ return [r for r in results if isinstance(r, Mapping)]
58
+ legacy = payload.get("sources")
59
+ if isinstance(legacy, dict):
60
+ out: list[Mapping[str, Any]] = []
61
+ for _key, row in legacy.items():
62
+ if isinstance(row, Mapping):
63
+ out.append(row)
64
+ return out
65
+ return []
66
+
67
+
68
+ def _freshness_status(row: Mapping[str, Any]) -> str | None:
69
+ raw = _as_str(row.get("status")) or _as_str(row.get("state"))
70
+ if raw is None:
71
+ return None
72
+ return raw.strip().lower()
73
+
74
+
75
+ def _load_sources_freshness_payload(path: Path) -> Mapping[str, Any] | None:
76
+ if not path.is_file():
77
+ logger.warning("dbt sources.json not found or not a file: %s", path)
78
+ return None
79
+ try:
80
+ text = path.read_text(encoding="utf-8")
81
+ except OSError as exc:
82
+ logger.warning("cannot read dbt sources.json %s: %s", path, exc)
83
+ return None
84
+ try:
85
+ data = json.loads(text)
86
+ except json.JSONDecodeError as exc:
87
+ logger.warning("dbt sources.json is not valid JSON %s: %s", path, exc)
88
+ return None
89
+ if not isinstance(data, dict):
90
+ logger.warning("dbt sources.json root must be a JSON object: %s", path)
91
+ return None
92
+ return data
93
+
94
+
95
+ def _apply_freshness_from_sources_json(
96
+ manifest: DbtManifest,
97
+ store: KuzuStore,
98
+ connection_name: str,
99
+ path: Path,
100
+ ) -> int:
101
+ payload = _load_sources_freshness_payload(path)
102
+ if payload is None:
103
+ return 0
104
+ applied = 0
105
+ for row in _iter_sources_json_results(payload):
106
+ uid = _as_str(row.get("unique_id"))
107
+ if uid is None or not uid.startswith("source."):
108
+ continue
109
+ status = _freshness_status(row)
110
+ if status is None or status not in _VALID_STATUSES:
111
+ logger.debug("dbt freshness: skipping unknown status %r for %s", status, uid)
112
+ continue
113
+ table_nid = resolve_dbt_parent_node_id(manifest, connection_name, uid)
114
+ if table_nid is None:
115
+ logger.debug(
116
+ "dbt freshness: could not resolve source %s to SchemaTable", uid
117
+ )
118
+ continue
119
+ max_loaded = _as_str(row.get("max_loaded_at")) or ""
120
+ updated = store.query_all_rows(
121
+ """
122
+ MATCH (t:SchemaTable {node_id: $nid, connection_name: $cn})
123
+ SET t.staleness_status = $status,
124
+ t.staleness_as_of = $as_of
125
+ RETURN t.node_id
126
+ """,
127
+ {
128
+ "nid": table_nid,
129
+ "cn": connection_name,
130
+ "status": status,
131
+ "as_of": max_loaded,
132
+ },
133
+ )
134
+ if not updated:
135
+ logger.debug(
136
+ "dbt freshness: SchemaTable missing for source %s (node_id=%s)",
137
+ uid,
138
+ table_nid,
139
+ )
140
+ continue
141
+ applied += 1
142
+ if status in _STALE_STATUSES:
143
+ ago = row.get("max_loaded_at_time_ago_in_s")
144
+ logger.warning(
145
+ "dbt source freshness: stale source %s status=%s max_loaded_at=%r "
146
+ "max_loaded_at_time_ago_in_s=%r",
147
+ uid,
148
+ status,
149
+ max_loaded,
150
+ ago,
151
+ )
152
+ return applied
153
+
154
+
155
+ def _collect_exposure_reachable_dbt_nodes(manifest: DbtManifest) -> set[str]:
156
+ """BFS upstream over ``parent_map`` from exposure ``depends_on`` model/source ids."""
157
+ seeds: list[str] = []
158
+ for exposure in manifest.exposures.values():
159
+ for dep in exposure.depends_on_nodes:
160
+ if dep.startswith("model.") or dep.startswith("source."):
161
+ seeds.append(dep)
162
+ if not seeds:
163
+ return set()
164
+ seen: set[str] = set()
165
+ frontier: deque[str] = deque(seeds)
166
+ while frontier:
167
+ dbid = frontier.popleft()
168
+ if dbid in seen:
169
+ continue
170
+ seen.add(dbid)
171
+ if not dbid.startswith("model."):
172
+ continue
173
+ for parent_id in manifest.parent_map.get(dbid, ()):
174
+ if parent_id.startswith("model.") or parent_id.startswith("source."):
175
+ frontier.append(parent_id)
176
+ return seen
177
+
178
+
179
+ def _mark_external_consumers(
180
+ manifest: DbtManifest,
181
+ store: KuzuStore,
182
+ connection_name: str,
183
+ dbt_node_ids: set[str],
184
+ ) -> int:
185
+ marked = 0
186
+ for dbid in dbt_node_ids:
187
+ if not (dbid.startswith("model.") or dbid.startswith("source.")):
188
+ continue
189
+ table_nid = resolve_dbt_parent_node_id(manifest, connection_name, dbid)
190
+ if table_nid is None:
191
+ logger.debug(
192
+ "dbt signals: could not resolve %s to SchemaTable for exposure propagation",
193
+ dbid,
194
+ )
195
+ continue
196
+ updated = store.query_all_rows(
197
+ """
198
+ MATCH (t:SchemaTable {node_id: $nid, connection_name: $cn})
199
+ SET t.has_external_consumers = true
200
+ RETURN t.node_id
201
+ """,
202
+ {"nid": table_nid, "cn": connection_name},
203
+ )
204
+ if not updated:
205
+ logger.debug(
206
+ "dbt signals: SchemaTable missing for exposure consumer %s (node_id=%s)",
207
+ dbid,
208
+ table_nid,
209
+ )
210
+ continue
211
+ marked += 1
212
+ return marked
213
+
214
+
215
+ def _apply_test_counts(
216
+ manifest: DbtManifest,
217
+ store: KuzuStore,
218
+ connection_name: str,
219
+ ) -> int:
220
+ """Count ``DbtTest.attached_node`` per model and write ``SchemaTable.test_count``."""
221
+ if not manifest.tests:
222
+ return 0
223
+ counter: Counter[str] = Counter()
224
+ for test in manifest.tests.values():
225
+ attached = test.attached_node
226
+ if attached is None or not attached.startswith("model."):
227
+ continue
228
+ counter[attached] += 1
229
+ if not counter:
230
+ return 0
231
+ applied = 0
232
+ for model_id, n in counter.items():
233
+ table_nid = resolve_dbt_parent_node_id(manifest, connection_name, model_id)
234
+ if table_nid is None:
235
+ continue
236
+ updated = store.query_all_rows(
237
+ """
238
+ MATCH (t:SchemaTable {node_id: $nid, connection_name: $cn})
239
+ SET t.test_count = $n
240
+ RETURN t.node_id
241
+ """,
242
+ {"nid": table_nid, "cn": connection_name, "n": int(n)},
243
+ )
244
+ if not updated:
245
+ logger.debug(
246
+ "dbt signals: SchemaTable missing for model %s (test_count=%d)",
247
+ model_id,
248
+ n,
249
+ )
250
+ continue
251
+ applied += 1
252
+ return applied
253
+
254
+
255
+ def write_dbt_signals(
256
+ manifest: DbtManifest,
257
+ store: KuzuStore,
258
+ connection_name: str,
259
+ sources_path: Path | None = None,
260
+ ) -> DbtSignalsWriteStats:
261
+ """Apply dbt exposure, freshness, and test-count signals to ``SchemaTable`` rows.
262
+
263
+ **Exposures:** Tables backing models or sources listed under any exposure, and every
264
+ upstream model/source reachable via ``parent_map``, get ``has_external_consumers=true``
265
+ when a matching ``SchemaTable`` exists.
266
+
267
+ **Freshness:** When ``sources_path`` points to dbt's ``sources.json``, every row with
268
+ a recognized status (``pass``, ``warn``, ``error``, ``runtime error``) is persisted as
269
+ ``staleness_status`` plus ``staleness_as_of`` (the ``max_loaded_at`` timestamp verbatim).
270
+ Stale statuses (``warn``/``error``) additionally emit a logger warning. This function
271
+ does not modify ``potentially_unused`` — that field remains governed by access stats.
272
+
273
+ **Test counts:** For each ``DbtTest`` with ``attached_node`` pointing at a model,
274
+ the number of tests is written to ``SchemaTable.test_count``.
275
+
276
+ Args:
277
+ manifest: Parsed dbt manifest (exposures, sources, models, tests, ``parent_map``).
278
+ store: Open Kuzu store.
279
+ connection_name: Pretensor connection name for ``SchemaTable`` nodes.
280
+ sources_path: Optional path to ``sources.json`` from ``dbt source freshness``.
281
+
282
+ Returns:
283
+ Counts of tables marked / freshness rows applied / tests counted for CLI summaries.
284
+ """
285
+ freshness_applied = 0
286
+ if sources_path is not None:
287
+ freshness_applied = _apply_freshness_from_sources_json(
288
+ manifest, store, connection_name, sources_path
289
+ )
290
+
291
+ exposures_marked = 0
292
+ if manifest.exposures:
293
+ reachable = _collect_exposure_reachable_dbt_nodes(manifest)
294
+ if reachable:
295
+ exposures_marked = _mark_external_consumers(
296
+ manifest, store, connection_name, reachable
297
+ )
298
+
299
+ tests_counted = _apply_test_counts(manifest, store, connection_name)
300
+
301
+ return DbtSignalsWriteStats(
302
+ exposures_marked=exposures_marked,
303
+ freshness_rows_applied=freshness_applied,
304
+ tests_counted=tests_counted,
305
+ )
@@ -0,0 +1,27 @@
1
+ """Business entity extraction: table classification, LLM grouping, Kuzu writes."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from pretensor.entities.builder import EntityBuilder
6
+ from pretensor.entities.classifier import (
7
+ TABLE_ROLES,
8
+ TableClassification,
9
+ TableClassifier,
10
+ TableClassifierInput,
11
+ is_entity_extraction_candidate,
12
+ )
13
+ from pretensor.entities.llm_extract import (
14
+ ExtractedEntity,
15
+ LLMEntityExtractor,
16
+ )
17
+
18
+ __all__ = [
19
+ "EntityBuilder",
20
+ "ExtractedEntity",
21
+ "LLMEntityExtractor",
22
+ "TABLE_ROLES",
23
+ "TableClassification",
24
+ "TableClassifier",
25
+ "TableClassifierInput",
26
+ "is_entity_extraction_candidate",
27
+ ]
@@ -0,0 +1,63 @@
1
+ """Write ``Entity`` nodes, ``REPRESENTS`` edges, and ``entity_type`` on tables."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from pretensor.connectors.models import SchemaSnapshot, Table
6
+ from pretensor.core.ids import entity_node_id, table_node_id
7
+ from pretensor.core.store import KuzuStore
8
+ from pretensor.entities.llm_extract import ExtractedEntity
9
+ from pretensor.graph_models.entity import EntityNode
10
+
11
+ __all__ = ["EntityBuilder"]
12
+
13
+
14
+ def _resolve_table(snapshot: SchemaSnapshot, ref: str) -> Table | None:
15
+ """Match ``schema.name`` or bare ``name`` (first match)."""
16
+ ref = ref.strip()
17
+ for t in snapshot.tables:
18
+ full = f"{t.schema_name}.{t.name}"
19
+ if ref == full or ref == t.name:
20
+ return t
21
+ return None
22
+
23
+
24
+ class EntityBuilder:
25
+ """Persists extracted business entities into Kuzu."""
26
+
27
+ def build(
28
+ self,
29
+ entities: list[ExtractedEntity],
30
+ store: KuzuStore,
31
+ snapshot: SchemaSnapshot,
32
+ ) -> None:
33
+ """Upsert entities and link them to existing ``SchemaTable`` nodes.
34
+
35
+ Args:
36
+ entities: Parsed LLM output (grouped tables per business entity).
37
+ store: Open Kuzu store (schema must exist; table nodes should exist).
38
+ snapshot: Used to resolve table references and stable node ids.
39
+ """
40
+ conn = snapshot.connection_name
41
+ db = snapshot.database
42
+ known = {table_node_id(conn, t.schema_name, t.name) for t in snapshot.tables}
43
+
44
+ for ext in entities:
45
+ eid = entity_node_id(conn, ext.name)
46
+ store.upsert_entity(
47
+ EntityNode(
48
+ node_id=eid,
49
+ connection_name=conn,
50
+ database=db,
51
+ name=ext.name,
52
+ description=ext.description,
53
+ )
54
+ )
55
+ for ref in ext.tables:
56
+ table = _resolve_table(snapshot, ref)
57
+ if table is None:
58
+ continue
59
+ tid = table_node_id(conn, table.schema_name, table.name)
60
+ if tid not in known:
61
+ continue
62
+ store.upsert_represents(eid, tid)
63
+ store.set_table_entity_type(tid, ext.name)