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,263 @@
1
+ """Curated exemplar tables per ``TableRole`` for embedding-based role voting.
2
+
3
+ The classifier in :mod:`pretensor.entities.classifier` uses several signals
4
+ (name patterns, column shape, structural degree). This module adds one
5
+ more: a nearest-centroid vote against a small per-role exemplar set.
6
+
7
+ For each role, we list a few real-ish ``(qualified_name, columns)`` pairs
8
+ that are *prototypical* of that role. At runtime, an
9
+ :class:`EmbeddingClient` embeds the exemplar text via
10
+ :func:`format_entity_text` and we mean-pool the vectors per role to get
11
+ one centroid per role. A new table's role vote is then
12
+ ``cosine(table_text_embedding, centroid)``, normalized so the best
13
+ matching role gets a score in roughly ``[0.0, 1.0]``.
14
+
15
+ The exemplar set ships in source so the wheel is self-contained — no
16
+ fixture files at runtime. Embeddings are computed lazily (first call)
17
+ and cached on the embedding client via ``_centroid_cache`` so repeated
18
+ classification runs reuse the same centroids without re-embedding.
19
+
20
+ Determinism contract:
21
+
22
+ * ``EmbeddingsConfig.role_weight=0.0`` (default) skips this signal entirely.
23
+ * When > 0, the vote is *added* to existing heuristic scores; it never
24
+ replaces them. ``role_weight`` should stay well below the typical
25
+ heuristic score magnitudes (~0.5–2.5) so the heuristic remains primary.
26
+ * If the embedding client is unavailable (extra not installed, or
27
+ ``embed`` raises), the function returns an all-zeros dict and the
28
+ classifier falls back to heuristic-only scoring.
29
+ """
30
+
31
+ from __future__ import annotations
32
+
33
+ import logging
34
+ import weakref
35
+
36
+ from pretensor.intelligence.embeddings import (
37
+ EmbeddingClient,
38
+ cosine_similarity,
39
+ format_entity_text,
40
+ )
41
+
42
+ __all__ = [
43
+ "ROLE_EXEMPLARS",
44
+ "compute_role_centroids",
45
+ "embedding_role_vote",
46
+ ]
47
+
48
+ logger = logging.getLogger(__name__)
49
+
50
+
51
+ # Per-role exemplars: (qualified_name, column_names). A handful per role is
52
+ # sufficient — mean-pooling smooths out idiosyncrasies of any single example.
53
+ # Names are illustrative; the embeddings encode the schema "shape" via
54
+ # format_entity_text, not the brand-name.
55
+ ROLE_EXEMPLARS: dict[str, list[tuple[str, list[str]]]] = {
56
+ "fact": [
57
+ ("public.orders", ["id", "customer_id", "order_date", "total_amount"]),
58
+ ("public.transactions", ["id", "account_id", "amount", "ts"]),
59
+ ("sales.invoice_lines", ["id", "invoice_id", "product_id", "qty", "price"]),
60
+ ("public.events", ["id", "user_id", "event_type", "occurred_at"]),
61
+ ("public.payments", ["id", "order_id", "amount", "paid_at", "method"]),
62
+ ],
63
+ "dimension": [
64
+ ("public.customers", ["id", "first_name", "last_name", "email"]),
65
+ ("public.products", ["id", "name", "category", "price"]),
66
+ ("public.users", ["id", "username", "email", "created_at"]),
67
+ ("public.locations", ["id", "country", "region", "city"]),
68
+ ("public.suppliers", ["id", "name", "address", "phone"]),
69
+ ],
70
+ "bridge": [
71
+ ("public.user_roles", ["user_id", "role_id"]),
72
+ ("public.product_categories", ["product_id", "category_id"]),
73
+ ("public.order_tags", ["order_id", "tag_id"]),
74
+ ("public.film_actor", ["film_id", "actor_id"]),
75
+ ],
76
+ "junction": [
77
+ ("public.user_groups", ["user_id", "group_id"]),
78
+ ("public.role_permissions", ["role_id", "permission_id"]),
79
+ ],
80
+ "staging": [
81
+ ("staging.raw_orders", ["id", "raw_payload", "ingested_at"]),
82
+ ("staging.tmp_users", ["id", "data", "loaded_at"]),
83
+ ("public.import_buffer", ["id", "source", "raw_json", "imported_at"]),
84
+ ],
85
+ "audit": [
86
+ ("public.audit_log", ["id", "actor", "action", "target", "ts"]),
87
+ ("public.access_log", ["id", "user_id", "endpoint", "ip", "ts"]),
88
+ ("public.history", ["id", "entity_id", "old_value", "new_value", "ts"]),
89
+ ],
90
+ "snapshot_scd": [
91
+ ("public.customer_history", ["id", "customer_id", "valid_from", "valid_to"]),
92
+ ("public.product_snapshot", ["id", "product_id", "snapshot_date", "price"]),
93
+ ("public.account_scd2", ["id", "account_id", "is_current", "effective_at"]),
94
+ ],
95
+ "aggregate": [
96
+ ("analytics.daily_revenue", ["day", "total", "order_count"]),
97
+ ("analytics.monthly_active_users", ["month", "mau", "growth_rate"]),
98
+ ("public.summary_stats", ["metric", "value", "computed_at"]),
99
+ ],
100
+ "system": [
101
+ ("public.schema_migrations", ["version", "applied_at"]),
102
+ ("public.ar_internal_metadata", ["key", "value", "updated_at"]),
103
+ ],
104
+ "entity_candidate": [
105
+ ("public.thing", ["id", "name"]),
106
+ ("public.unknown_table", ["id", "data"]),
107
+ ],
108
+ # 'unknown' has no exemplars by design — it's the fallback when nothing
109
+ # else fires; voting against it would defeat the purpose.
110
+ }
111
+
112
+
113
+ # Cache: maps client instance → {role → centroid}. Keyed on the actual
114
+ # object via WeakKeyDictionary so a freed client's entry vanishes
115
+ # automatically — guards against CPython recycling a deleted client's
116
+ # ``id()`` for a fresh, unrelated allocation, which would let a stale
117
+ # centroid leak into the new client's lookups in long-running processes.
118
+ _CENTROID_CACHE: "weakref.WeakKeyDictionary[EmbeddingClient, dict[str, list[float]]]" = weakref.WeakKeyDictionary()
119
+ # Bounded fallback for client types that aren't weakref-able (e.g. mocks
120
+ # built from raw ``object()``).
121
+ #
122
+ # CAVEAT: this fallback path keys on ``id(client)`` directly and therefore
123
+ # does NOT have the recycling protection of the WeakKeyDictionary above.
124
+ # If a non-weakref-able client A is freed and its id is recycled by a
125
+ # fresh, unrelated client B before A's fallback entry is evicted, B's
126
+ # lookup would return A's stale centroid. The hazard is bounded in
127
+ # practice because (a) only test mocks land here — production
128
+ # ``LocalEmbeddingClient`` is weakref-able, (b) the FIFO ceiling
129
+ # below caps the window in which a stale entry can survive, and (c) tests
130
+ # clear the cache between runs via ``_clear_centroid_cache_for_tests``.
131
+ # Treat the fallback as best-effort for tests, not a load-bearing
132
+ # correctness guarantee.
133
+ _CENTROID_CACHE_FALLBACK: dict[int, dict[str, list[float]]] = {}
134
+ # Maximum entries retained in the non-weakref-able fallback cache. The
135
+ # typical caller is one ``LocalEmbeddingClient`` per process; the only
136
+ # things landing here are short-lived test mocks. A small ceiling
137
+ # (insertion-order eviction, FIFO) is plenty to keep the dict bounded
138
+ # without losing the cache's value on the realistic call patterns.
139
+ _CENTROID_CACHE_FALLBACK_MAX = 32
140
+
141
+
142
+ def _cache_get(client: EmbeddingClient) -> dict[str, list[float]] | None:
143
+ try:
144
+ return _CENTROID_CACHE.get(client)
145
+ except TypeError:
146
+ return _CENTROID_CACHE_FALLBACK.get(id(client))
147
+
148
+
149
+ def _cache_set(client: EmbeddingClient, value: dict[str, list[float]]) -> None:
150
+ try:
151
+ _CENTROID_CACHE[client] = value
152
+ except TypeError:
153
+ # FIFO evict to keep the dict bounded.
154
+ while len(_CENTROID_CACHE_FALLBACK) >= _CENTROID_CACHE_FALLBACK_MAX:
155
+ oldest_key = next(iter(_CENTROID_CACHE_FALLBACK))
156
+ del _CENTROID_CACHE_FALLBACK[oldest_key]
157
+ _CENTROID_CACHE_FALLBACK[id(client)] = value
158
+
159
+
160
+ def compute_role_centroids(
161
+ client: EmbeddingClient,
162
+ ) -> dict[str, list[float]]:
163
+ """Return one centroid vector per role with at least one exemplar.
164
+
165
+ Embeds every exemplar via :func:`format_entity_text` and mean-pools
166
+ per role. Cached on the client (weak-ref keyed) so the second call
167
+ is free without leaking memory if the client is later garbage-collected.
168
+ Returns an empty dict (no centroids) when the client raises any
169
+ exception while embedding — caller falls back to heuristic-only.
170
+ """
171
+ cached = _cache_get(client)
172
+ if cached is not None:
173
+ return cached
174
+
175
+ centroids: dict[str, list[float]] = {}
176
+ try:
177
+ # Flatten the exemplar set into one batch for one embed() call.
178
+ flat_texts: list[str] = []
179
+ flat_owners: list[str] = []
180
+ for role, exemplars in ROLE_EXEMPLARS.items():
181
+ for qname, cols in exemplars:
182
+ flat_texts.append(format_entity_text(qname, cols))
183
+ flat_owners.append(role)
184
+ if not flat_texts:
185
+ _cache_set(client, centroids)
186
+ return centroids
187
+
188
+ vectors = client.embed(flat_texts)
189
+ if not vectors or len(vectors) != len(flat_texts):
190
+ # NullEmbeddingClient returns []; treat as no centroids.
191
+ _cache_set(client, centroids)
192
+ return centroids
193
+
194
+ per_role: dict[str, list[list[float]]] = {}
195
+ for role, vec in zip(flat_owners, vectors, strict=True):
196
+ per_role.setdefault(role, []).append(vec)
197
+
198
+ for role, vecs in per_role.items():
199
+ if not vecs:
200
+ continue
201
+ dim = len(vecs[0])
202
+ mean = [0.0] * dim
203
+ count = 0
204
+ for v in vecs:
205
+ if len(v) != dim:
206
+ continue
207
+ for i in range(dim):
208
+ mean[i] += float(v[i])
209
+ count += 1
210
+ if count == 0:
211
+ continue
212
+ centroids[role] = [x / count for x in mean]
213
+ except ImportError as exc:
214
+ logger.warning(
215
+ "role_exemplars: [embeddings] extra not installed (%s); "
216
+ "skipping role centroids",
217
+ exc,
218
+ )
219
+ centroids = {}
220
+ except Exception as exc: # noqa: BLE001 — additive signal, never raise
221
+ logger.warning(
222
+ "role_exemplars: failed to compute centroids (%s); "
223
+ "falling back to heuristic-only role vote",
224
+ exc,
225
+ )
226
+ centroids = {}
227
+
228
+ _cache_set(client, centroids)
229
+ return centroids
230
+
231
+
232
+ def embedding_role_vote(
233
+ table_embedding: list[float] | None,
234
+ centroids: dict[str, list[float]],
235
+ ) -> dict[str, float]:
236
+ """Cosine vote for the table's embedding against per-role centroids.
237
+
238
+ ``table_embedding`` is typically pre-computed at index time via
239
+ ``EmbeddingIndexStep`` and read back from the store; pass ``None``
240
+ when the table doesn't carry one and this function returns an
241
+ all-zeros vote.
242
+
243
+ Returns ``{role: cosine}`` for every role with a centroid; roles
244
+ without a centroid are not present in the output. Cosine values are
245
+ in ``[-1, 1]``; the classifier scales by ``role_weight`` before
246
+ blending.
247
+ """
248
+ if table_embedding is None or not centroids:
249
+ return {role: 0.0 for role in centroids}
250
+
251
+ out: dict[str, float] = {}
252
+ for role, centroid in centroids.items():
253
+ try:
254
+ out[role] = cosine_similarity(table_embedding, centroid)
255
+ except ValueError:
256
+ out[role] = 0.0
257
+ return out
258
+
259
+
260
+ def _clear_centroid_cache_for_tests() -> None:
261
+ """Test helper: drop the module-level centroid cache."""
262
+ _CENTROID_CACHE.clear()
263
+ _CENTROID_CACHE_FALLBACK.clear()
@@ -0,0 +1,360 @@
1
+ """Persist table roles and cluster-level schema patterns after clustering."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import logging
7
+ from collections import deque
8
+ from typing import Literal, Protocol, runtime_checkable
9
+
10
+ from pretensor.core.store import KuzuStore
11
+ from pretensor.entities.classifier import (
12
+ TableClassification,
13
+ TableClassifier,
14
+ TableClassifierInput,
15
+ )
16
+ from pretensor.intelligence.clustering import Cluster
17
+ from pretensor.intelligence.embeddings import EmbeddingClient
18
+ from pretensor.intelligence.role_exemplars import (
19
+ compute_role_centroids,
20
+ embedding_role_vote,
21
+ )
22
+
23
+ logger = logging.getLogger(__name__)
24
+
25
+ __all__ = [
26
+ "LlmTableClassificationClient",
27
+ "classify_database_tables",
28
+ "classify_database_tables_async",
29
+ "compute_cluster_schema_patterns",
30
+ "load_fk_reference_pairs",
31
+ ]
32
+
33
+ SchemaPattern = Literal["star", "snowflake", "constellation", "erd", "unknown"]
34
+
35
+ _CLASSIFIER = TableClassifier()
36
+
37
+ # Precedence for ``schema_pattern`` when multiple signals apply:
38
+ # 1. constellation — multiple facts share at least one dimension (direct FK or via bridges)
39
+ # 2. snowflake — star-like hub plus at least one dimension-to-dimension FK in the cluster
40
+ # 3. star — one (or primary) fact with 3+ distinct dimensions via outgoing FKs
41
+ # 4. erd — connected tables with FKs but not a clear star/snowflake/constellation
42
+ # 5. role-only fallback — used when the cluster has no FK edges (introspection gap)
43
+
44
+
45
+ @runtime_checkable
46
+ class LlmTableClassificationClient(Protocol):
47
+ """Async JSON batch classifier for low-confidence tables."""
48
+
49
+ async def classify_tables_json(self, user_prompt: str) -> str:
50
+ """Return JSON array aligned with batch order: {role, confidence, signals?}."""
51
+ ...
52
+
53
+
54
+ def load_fk_reference_pairs(
55
+ store: KuzuStore, database_key: str
56
+ ) -> list[tuple[str, str]]:
57
+ """Return all ``(src_table_id, dst_table_id)`` for ``FK_REFERENCES`` in one database."""
58
+ rows = store.query_all_rows(
59
+ """
60
+ MATCH (a:SchemaTable {database: $db})-[r:FK_REFERENCES]->(b:SchemaTable {database: $db})
61
+ RETURN a.node_id, b.node_id
62
+ """,
63
+ {"db": database_key},
64
+ )
65
+ return [(str(a), str(b)) for a, b in rows]
66
+
67
+
68
+ def _fk_degree_rows(store: KuzuStore, database_key: str) -> list[tuple[str, int, int]]:
69
+ return store.query_all_rows(
70
+ """
71
+ MATCH (t:SchemaTable {database: $db})
72
+ OPTIONAL MATCH (t)-[:FK_REFERENCES]->(o:SchemaTable)
73
+ OPTIONAL MATCH (i:SchemaTable)-[:FK_REFERENCES]->(t)
74
+ RETURN t.node_id, count(DISTINCT o.node_id), count(DISTINCT i.node_id)
75
+ """,
76
+ {"db": database_key},
77
+ )
78
+
79
+
80
+ def _table_context_rows(store: KuzuStore, database_key: str) -> list[tuple]:
81
+ return store.query_all_rows(
82
+ """
83
+ MATCH (t:SchemaTable {database: $db})
84
+ RETURN t.node_id, t.table_name, t.schema_name, t.row_count,
85
+ t.seq_scan_count, t.idx_scan_count, t.insert_count, t.update_count
86
+ """,
87
+ {"db": database_key},
88
+ )
89
+
90
+
91
+ def _columns_by_table(store: KuzuStore, database_key: str) -> dict[str, list[str]]:
92
+ rows = store.query_all_rows(
93
+ """
94
+ MATCH (t:SchemaTable {database: $db})-[:HAS_COLUMN]->(c:SchemaColumn)
95
+ RETURN t.node_id, c.column_name, c.ordinal_position
96
+ ORDER BY t.node_id, c.ordinal_position, c.column_name
97
+ """,
98
+ {"db": database_key},
99
+ )
100
+ out: dict[str, list[str]] = {}
101
+ for tid, cname, _ord in rows:
102
+ tid_s = str(tid)
103
+ out.setdefault(tid_s, []).append(str(cname))
104
+ return out
105
+
106
+
107
+ def _initial_classification(
108
+ inp: TableClassifierInput,
109
+ ) -> TableClassification:
110
+ return _CLASSIFIER.classify(inp)
111
+
112
+
113
+ def _fk_pairs_for_cluster(
114
+ table_ids: frozenset[str], all_pairs: list[tuple[str, str]]
115
+ ) -> list[tuple[str, str]]:
116
+ return [p for p in all_pairs if p[0] in table_ids and p[1] in table_ids]
117
+
118
+
119
+ def _dimensions_reachable_from(
120
+ start: str,
121
+ table_ids: frozenset[str],
122
+ adj_out: dict[str, list[str]],
123
+ dims: set[str],
124
+ ) -> set[str]:
125
+ """Directed reachability from ``start`` within ``table_ids``; return dimension nodes hit."""
126
+ found: set[str] = set()
127
+ q: deque[str] = deque([start])
128
+ seen = {start}
129
+ while q:
130
+ u = q.popleft()
131
+ if u in dims:
132
+ found.add(u)
133
+ for v in adj_out.get(u, ()):
134
+ if v not in table_ids or v in seen:
135
+ continue
136
+ seen.add(v)
137
+ q.append(v)
138
+ return found
139
+
140
+
141
+ def _facts_share_dimension_via_paths(
142
+ facts: set[str],
143
+ table_ids: frozenset[str],
144
+ fk_pairs: list[tuple[str, str]],
145
+ dims: set[str],
146
+ ) -> bool:
147
+ """True if some pair of facts reaches a common dimension (paths may cross bridges)."""
148
+ if len(facts) < 2:
149
+ return False
150
+ adj_out: dict[str, list[str]] = {}
151
+ for a, b in fk_pairs:
152
+ adj_out.setdefault(a, []).append(b)
153
+ dim_sets: dict[str, set[str]] = {}
154
+ for f in facts:
155
+ dim_sets[f] = _dimensions_reachable_from(f, table_ids, adj_out, dims)
156
+ fact_list = list(facts)
157
+ for i in range(len(fact_list)):
158
+ for j in range(i + 1, len(fact_list)):
159
+ if dim_sets[fact_list[i]] & dim_sets[fact_list[j]]:
160
+ return True
161
+ return False
162
+
163
+
164
+ def _shared_dimension_direct(
165
+ facts: set[str], dims: set[str], fk_pairs: list[tuple[str, str]]
166
+ ) -> dict[str, set[str]]:
167
+ """fact -> set of dimensions referenced by a direct FK from the fact."""
168
+ out: dict[str, set[str]] = {f: set() for f in facts}
169
+ for a, b in fk_pairs:
170
+ if a in facts and b in dims:
171
+ out.setdefault(a, set()).add(b)
172
+ return out
173
+
174
+
175
+ def _schema_pattern_role_only(
176
+ tset: frozenset[str],
177
+ roles: dict[str, str],
178
+ bridges: set[str],
179
+ ) -> SchemaPattern:
180
+ """Heuristic when the cluster has no FK edges (or none inside the cluster)."""
181
+ facts = {tid for tid in tset if roles.get(tid) == "fact"}
182
+ dims = {tid for tid in tset if roles.get(tid) == "dimension"}
183
+ if len(facts) >= 2 and len(dims) >= 2:
184
+ return "constellation"
185
+ if facts and len(dims) >= 3:
186
+ return "star"
187
+ if bridges and (facts or dims):
188
+ return "erd"
189
+ if len(tset) >= 2:
190
+ return "erd"
191
+ return "unknown"
192
+
193
+
194
+ def compute_cluster_schema_patterns(
195
+ clusters: list[Cluster],
196
+ role_by_table: dict[str, TableClassification],
197
+ fk_pairs: list[tuple[str, str]],
198
+ ) -> dict[int, SchemaPattern]:
199
+ """Map cluster index to ``schema_pattern`` (aligned with ``cluster::{idx}`` ids).
200
+
201
+ Args:
202
+ clusters: Leiden clusters.
203
+ role_by_table: Latest classification per table node id.
204
+ fk_pairs: All ``FK_REFERENCES`` (src, dst) for the logical database.
205
+ """
206
+ patterns: dict[int, SchemaPattern] = {}
207
+ for idx, cluster in enumerate(clusters):
208
+ tset = frozenset(cluster.table_ids)
209
+ if not tset:
210
+ patterns[idx] = "unknown"
211
+ continue
212
+ roles = {
213
+ tid: role_by_table.get(tid, TableClassification("unknown", 0.0, ())).role
214
+ for tid in tset
215
+ }
216
+ facts = {tid for tid, r in roles.items() if r == "fact"}
217
+ dims = {tid for tid, r in roles.items() if r == "dimension"}
218
+ bridges = {tid for tid, r in roles.items() if r == "bridge"}
219
+ cluster_pairs = _fk_pairs_for_cluster(tset, fk_pairs)
220
+
221
+ if not cluster_pairs:
222
+ patterns[idx] = _schema_pattern_role_only(tset, roles, bridges)
223
+ continue
224
+
225
+ dim_dim = sum(1 for a, b in cluster_pairs if a in dims and b in dims)
226
+ fact_to_dim = _shared_dimension_direct(facts, dims, cluster_pairs)
227
+ facts_per_dim: dict[str, set[str]] = {}
228
+ for f, dset in fact_to_dim.items():
229
+ for d in dset:
230
+ facts_per_dim.setdefault(d, set()).add(f)
231
+
232
+ shared_direct = any(len(fs) >= 2 for fs in facts_per_dim.values())
233
+ multi_hop_share = _facts_share_dimension_via_paths(
234
+ facts, tset, cluster_pairs, dims
235
+ )
236
+ constellation = len(facts) >= 2 and (shared_direct or multi_hop_share)
237
+
238
+ max_dim_from_one_fact = max((len(s) for s in fact_to_dim.values()), default=0)
239
+ star_like = bool(facts and dims and max_dim_from_one_fact >= 3)
240
+
241
+ pattern: SchemaPattern = "unknown"
242
+ if constellation:
243
+ pattern = "constellation"
244
+ elif star_like:
245
+ pattern = "snowflake" if dim_dim >= 1 else "star"
246
+ elif len(tset) >= 2 and cluster_pairs:
247
+ pattern = "erd"
248
+ elif bridges and cluster_pairs:
249
+ pattern = "erd"
250
+
251
+ patterns[idx] = pattern
252
+ return patterns
253
+
254
+
255
+ def classify_database_tables(
256
+ store: KuzuStore,
257
+ database_key: str,
258
+ *,
259
+ role_weight: float = 0.0,
260
+ embedding_client: EmbeddingClient | None = None,
261
+ ) -> dict[str, TableClassification]:
262
+ """Write classifier fields and return heuristic results.
263
+
264
+ When ``role_weight > 0`` and an ``embedding_client`` is provided, this
265
+ function computes per-role centroids (cached on the client identity)
266
+ and bulk-fetches every table's stored embedding. Each table's role
267
+ classification then receives an additive ``role_weight * cosine``
268
+ signal per role; heuristic scores remain primary.
269
+ Defaults keep behavior byte-identical to the pre-embedding pipeline.
270
+ """
271
+ store.ensure_schema()
272
+ fk_rows = _fk_degree_rows(store, database_key)
273
+ fk_out = {str(r[0]): int(r[1] or 0) for r in fk_rows}
274
+ fk_in = {str(r[0]): int(r[2] or 0) for r in fk_rows}
275
+ col_map = _columns_by_table(store, database_key)
276
+ ctx_rows = _table_context_rows(store, database_key)
277
+ results: dict[str, TableClassification] = {}
278
+
279
+ # Pre-pass: compute role centroids + bulk-fetch every table's embedding.
280
+ # Both are no-ops when role_weight is 0, when no embedding_client is
281
+ # supplied, or when the [embeddings] extra is missing (the client raises
282
+ # ImportError → role_exemplars logs once and returns {}).
283
+ centroids: dict[str, list[float]] = {}
284
+ embeddings_by_id: dict[str, list[float]] = {}
285
+ if role_weight > 0.0 and embedding_client is not None:
286
+ centroids = compute_role_centroids(embedding_client)
287
+ if centroids:
288
+ for emb_row in store.iter_table_embeddings(database=database_key):
289
+ if emb_row.embedding is None:
290
+ continue
291
+ embeddings_by_id[emb_row.node_id] = [
292
+ float(x) for x in emb_row.embedding
293
+ ]
294
+
295
+ for row in ctx_rows:
296
+ tid = str(row[0])
297
+ tname = str(row[1])
298
+ sname = str(row[2])
299
+ rcount = int(row[3]) if row[3] is not None else None
300
+ seq_sc = int(row[4]) if row[4] is not None else None
301
+ idx_sc = int(row[5]) if row[5] is not None else None
302
+ ins = int(row[6]) if row[6] is not None else None
303
+ upd = int(row[7]) if row[7] is not None else None
304
+ cols = col_map.get(tid, [])
305
+ inp = TableClassifierInput(
306
+ name=tname,
307
+ schema_name=sname,
308
+ columns=cols,
309
+ row_count=rcount,
310
+ fk_out_degree=fk_out.get(tid, 0),
311
+ fk_in_degree=fk_in.get(tid, 0),
312
+ seq_scan_count=seq_sc,
313
+ idx_scan_count=idx_sc,
314
+ insert_count=ins,
315
+ update_count=upd,
316
+ )
317
+
318
+ vote: dict[str, float] | None = None
319
+ emb = embeddings_by_id.get(tid)
320
+ if centroids and emb is not None:
321
+ vote = embedding_role_vote(emb, centroids)
322
+
323
+ cls = _CLASSIFIER.classify(
324
+ inp,
325
+ embedding_vote=vote,
326
+ role_weight=role_weight if vote is not None else 0.0,
327
+ )
328
+ results[tid] = cls
329
+ store.set_table_classification(
330
+ tid,
331
+ role=cls.role,
332
+ role_confidence=cls.confidence,
333
+ classification_signals_json=json.dumps(
334
+ list(cls.signals), ensure_ascii=False
335
+ ),
336
+ )
337
+
338
+ return results
339
+
340
+
341
+ async def classify_database_tables_async(
342
+ store: KuzuStore,
343
+ database_key: str,
344
+ *,
345
+ llm_client: LlmTableClassificationClient | None = None,
346
+ role_weight: float = 0.0,
347
+ embedding_client: EmbeddingClient | None = None,
348
+ ) -> dict[str, TableClassification]:
349
+ """Run heuristic classification with optional embedding role vote.
350
+
351
+ ``llm_client`` is accepted but ignored (legacy parameter). See the
352
+ sync :func:`classify_database_tables` for ``role_weight`` /
353
+ ``embedding_client`` semantics.
354
+ """
355
+ return classify_database_tables(
356
+ store,
357
+ database_key,
358
+ role_weight=role_weight,
359
+ embedding_client=embedding_client,
360
+ )
@@ -0,0 +1,76 @@
1
+ """Scoring contracts and registry for relationship discovery candidates."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from abc import ABC, abstractmethod
6
+ from collections.abc import Iterable, Iterator
7
+ from typing import TypeAlias
8
+
9
+ from pretensor.connectors.models import SchemaSnapshot
10
+ from pretensor.graph_models.relationship import RelationshipCandidate
11
+
12
+ __all__ = ["JoinKey", "RelationshipScorer", "ScorerRegistry"]
13
+
14
+ JoinKey: TypeAlias = tuple[str, str, str, str]
15
+
16
+
17
+ class RelationshipScorer(ABC):
18
+ """Contract for candidate generators used by relationship discovery."""
19
+
20
+ @abstractmethod
21
+ def name(self) -> str:
22
+ """Stable scorer name for observability and deduping."""
23
+
24
+ @abstractmethod
25
+ def score(
26
+ self,
27
+ snapshot: SchemaSnapshot,
28
+ explicit_fk_keys: set[JoinKey],
29
+ ) -> list[RelationshipCandidate]:
30
+ """Generate inferred relationship candidates for a snapshot."""
31
+
32
+
33
+ class ScorerRegistry:
34
+ """Ordered registry for relationship scorers.
35
+
36
+ The first scorer can emit broad candidates while later scorers can refine or
37
+ add alternatives. Discovery remains deterministic because the insertion order
38
+ is preserved.
39
+ """
40
+
41
+ def __init__(self, scorers: Iterable[RelationshipScorer] | None = None) -> None:
42
+ self._scorers_by_name: dict[str, RelationshipScorer] = {}
43
+ self._order: list[str] = []
44
+ for scorer in scorers or ():
45
+ self.register(scorer)
46
+
47
+ def register(self, scorer: RelationshipScorer) -> None:
48
+ """Register a scorer instance by unique name."""
49
+ name = scorer.name()
50
+ if name in self._scorers_by_name:
51
+ msg = f"Relationship scorer already registered: {name}"
52
+ raise ValueError(msg)
53
+ self._scorers_by_name[name] = scorer
54
+ self._order.append(name)
55
+
56
+ def score_all(
57
+ self,
58
+ snapshot: SchemaSnapshot,
59
+ explicit_fk_keys: set[JoinKey],
60
+ ) -> list[RelationshipCandidate]:
61
+ """Run all scorers in registration order and collect candidates."""
62
+ out: list[RelationshipCandidate] = []
63
+ for name in self._order:
64
+ scorer = self._scorers_by_name[name]
65
+ out.extend(scorer.score(snapshot, explicit_fk_keys))
66
+ return out
67
+
68
+ def __iter__(self) -> Iterator[RelationshipScorer]:
69
+ """Iterate scorers in registration order (read-only).
70
+
71
+ Used by ``extend_with_embedding_scorer`` (and any future helper
72
+ that builds a derived registry) to copy scorers without poking at
73
+ the private ``_scorers_by_name`` / ``_order`` state.
74
+ """
75
+ for name in self._order:
76
+ yield self._scorers_by_name[name]