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.
- pretensor/__init__.py +50 -0
- pretensor/benchmark/__init__.py +54 -0
- pretensor/benchmark/cli.py +294 -0
- pretensor/benchmark/fixtures.py +84 -0
- pretensor/benchmark/l1/__init__.py +23 -0
- pretensor/benchmark/l1/metrics.py +141 -0
- pretensor/benchmark/l1/pipeline.py +188 -0
- pretensor/benchmark/l1/runner.py +245 -0
- pretensor/benchmark/l2/__init__.py +27 -0
- pretensor/benchmark/l2/gold.py +236 -0
- pretensor/benchmark/l2/metrics.py +146 -0
- pretensor/benchmark/l2/pipeline.py +124 -0
- pretensor/benchmark/l2/runner.py +530 -0
- pretensor/benchmark/l3/__init__.py +73 -0
- pretensor/benchmark/l3/agent.py +316 -0
- pretensor/benchmark/l3/db.py +188 -0
- pretensor/benchmark/l3/gold.py +85 -0
- pretensor/benchmark/l3/llm_client.py +395 -0
- pretensor/benchmark/l3/mcp_client.py +357 -0
- pretensor/benchmark/l3/pretensor_runner.py +456 -0
- pretensor/benchmark/l3/prompt.py +132 -0
- pretensor/benchmark/l3/runner.py +358 -0
- pretensor/benchmark/l3/sql_equivalence.py +176 -0
- pretensor/benchmark/release_gate.py +448 -0
- pretensor/benchmark/results.py +298 -0
- pretensor/benchmark/runner.py +109 -0
- pretensor/cli/__init__.py +1 -0
- pretensor/cli/commands/_source_runner.py +147 -0
- pretensor/cli/commands/analyze.py +201 -0
- pretensor/cli/commands/connections/__init__.py +7 -0
- pretensor/cli/commands/connections/add_remove.py +126 -0
- pretensor/cli/commands/connections/register.py +12 -0
- pretensor/cli/commands/export.py +131 -0
- pretensor/cli/commands/index.py +559 -0
- pretensor/cli/commands/list.py +76 -0
- pretensor/cli/commands/quickstart.py +207 -0
- pretensor/cli/commands/reindex.py +646 -0
- pretensor/cli/commands/semantic.py +190 -0
- pretensor/cli/commands/serve.py +144 -0
- pretensor/cli/commands/sync_grants.py +149 -0
- pretensor/cli/commands/validate.py +176 -0
- pretensor/cli/config_file.py +442 -0
- pretensor/cli/constants.py +10 -0
- pretensor/cli/dbt_enrichment.py +96 -0
- pretensor/cli/main.py +109 -0
- pretensor/cli/paths.py +43 -0
- pretensor/cli/plugin.py +52 -0
- pretensor/config.py +226 -0
- pretensor/connectors/__init__.py +29 -0
- pretensor/connectors/base.py +165 -0
- pretensor/connectors/bigquery.py +468 -0
- pretensor/connectors/inspect.py +321 -0
- pretensor/connectors/lineage_sqlglot.py +97 -0
- pretensor/connectors/models.py +130 -0
- pretensor/connectors/mysql.py +402 -0
- pretensor/connectors/pg_array_parse.py +53 -0
- pretensor/connectors/postgres.py +938 -0
- pretensor/connectors/registry.py +93 -0
- pretensor/connectors/snapshot.py +244 -0
- pretensor/connectors/snowflake.py +908 -0
- pretensor/core/__init__.py +1 -0
- pretensor/core/builder.py +307 -0
- pretensor/core/dsn_crypto.py +51 -0
- pretensor/core/graph_schema_manager.py +246 -0
- pretensor/core/graph_store.py +1226 -0
- pretensor/core/ids.py +101 -0
- pretensor/core/portable_export.py +276 -0
- pretensor/core/query_runner.py +67 -0
- pretensor/core/registry.py +209 -0
- pretensor/core/schema.py +473 -0
- pretensor/core/secure_io.py +93 -0
- pretensor/core/store.py +469 -0
- pretensor/enrichment/__init__.py +1 -0
- pretensor/enrichment/analyze/__init__.py +0 -0
- pretensor/enrichment/analyze/classify.py +49 -0
- pretensor/enrichment/analyze/extract_python.py +196 -0
- pretensor/enrichment/analyze/parse.py +141 -0
- pretensor/enrichment/analyze/pipeline.py +195 -0
- pretensor/enrichment/analyze/summary.py +38 -0
- pretensor/enrichment/analyze/walker.py +98 -0
- pretensor/enrichment/analyze/writers.py +214 -0
- pretensor/enrichment/dbt/__init__.py +30 -0
- pretensor/enrichment/dbt/lineage.py +100 -0
- pretensor/enrichment/dbt/manifest.py +300 -0
- pretensor/enrichment/dbt/metadata.py +263 -0
- pretensor/enrichment/dbt/pipeline.py +77 -0
- pretensor/enrichment/dbt/resolution.py +101 -0
- pretensor/enrichment/dbt/signals.py +305 -0
- pretensor/entities/__init__.py +27 -0
- pretensor/entities/builder.py +63 -0
- pretensor/entities/classifier.py +383 -0
- pretensor/entities/llm_extract.py +66 -0
- pretensor/errors.py +35 -0
- pretensor/graph_models/__init__.py +17 -0
- pretensor/graph_models/base.py +11 -0
- pretensor/graph_models/consumer.py +71 -0
- pretensor/graph_models/edge.py +35 -0
- pretensor/graph_models/entity.py +21 -0
- pretensor/graph_models/node.py +79 -0
- pretensor/graph_models/relationship.py +33 -0
- pretensor/integrations/__init__.py +42 -0
- pretensor/integrations/_base.py +138 -0
- pretensor/integrations/google_adk.py +49 -0
- pretensor/integrations/langchain.py +55 -0
- pretensor/integrations/llamaindex.py +53 -0
- pretensor/intelligence/__init__.py +33 -0
- pretensor/intelligence/cluster_labeler.py +425 -0
- pretensor/intelligence/clustering.py +168 -0
- pretensor/intelligence/combining.py +32 -0
- pretensor/intelligence/discovery.py +114 -0
- pretensor/intelligence/embeddings.py +317 -0
- pretensor/intelligence/graph_export.py +200 -0
- pretensor/intelligence/heuristic.py +544 -0
- pretensor/intelligence/join_paths/__init__.py +130 -0
- pretensor/intelligence/join_paths/on_demand.py +516 -0
- pretensor/intelligence/join_paths/storage.py +70 -0
- pretensor/intelligence/llm_infer.py +78 -0
- pretensor/intelligence/llm_runtime.py +62 -0
- pretensor/intelligence/metric_templates.py +193 -0
- pretensor/intelligence/pipeline.py +364 -0
- pretensor/intelligence/role_exemplars.py +263 -0
- pretensor/intelligence/schema_classification.py +360 -0
- pretensor/intelligence/scoring.py +76 -0
- pretensor/intelligence/semantic.py +240 -0
- pretensor/intelligence/shadow_alias.py +101 -0
- pretensor/intelligence/statistical.py +50 -0
- pretensor/intelligence/steps.py +191 -0
- pretensor/intelligence/steps_embedding.py +168 -0
- pretensor/introspection/__init__.py +6 -0
- pretensor/introspection/inspector.py +5 -0
- pretensor/introspection/models/__init__.py +0 -0
- pretensor/introspection/models/base.py +5 -0
- pretensor/introspection/models/config.py +237 -0
- pretensor/introspection/models/dsn.py +550 -0
- pretensor/introspection/models/plan.py +116 -0
- pretensor/introspection/models/schema.py +10 -0
- pretensor/introspection/models/semantic.py +121 -0
- pretensor/introspection/models/validation.py +116 -0
- pretensor/introspection/snapshot.py +46 -0
- pretensor/mcp/__init__.py +16 -0
- pretensor/mcp/config_json.py +24 -0
- pretensor/mcp/payload_types.py +274 -0
- pretensor/mcp/resources/__init__.py +17 -0
- pretensor/mcp/resources/markdown.py +314 -0
- pretensor/mcp/server.py +285 -0
- pretensor/mcp/service.py +49 -0
- pretensor/mcp/service_context.py +142 -0
- pretensor/mcp/service_registry.py +294 -0
- pretensor/mcp/store_cache.py +43 -0
- pretensor/mcp/tool_registry.py +136 -0
- pretensor/mcp/tools/__init__.py +1 -0
- pretensor/mcp/tools/_rank.py +244 -0
- pretensor/mcp/tools/_timed.py +26 -0
- pretensor/mcp/tools/compile_metric.py +144 -0
- pretensor/mcp/tools/consumers.py +161 -0
- pretensor/mcp/tools/context.py +1121 -0
- pretensor/mcp/tools/cypher.py +509 -0
- pretensor/mcp/tools/detect_changes.py +254 -0
- pretensor/mcp/tools/impact.py +271 -0
- pretensor/mcp/tools/list.py +131 -0
- pretensor/mcp/tools/schema.py +170 -0
- pretensor/mcp/tools/search.py +316 -0
- pretensor/mcp/tools/semantic_search.py +282 -0
- pretensor/mcp/tools/traverse.py +1027 -0
- pretensor/mcp/tools/validate_sql.py +150 -0
- pretensor/observability.py +203 -0
- pretensor/py.typed +0 -0
- pretensor/quickstart/README.md +29 -0
- pretensor/quickstart/__init__.py +6 -0
- pretensor/quickstart/docker-compose.yml +18 -0
- pretensor/quickstart/pagila_data.sql +63 -0
- pretensor/quickstart/pagila_ddl.sql +92 -0
- pretensor/search/__init__.py +6 -0
- pretensor/search/base.py +80 -0
- pretensor/search/index.py +435 -0
- pretensor/semantic/__init__.py +24 -0
- pretensor/semantic/base.py +123 -0
- pretensor/semantic/compiler.py +487 -0
- pretensor/semantic/yaml_layer.py +180 -0
- pretensor/skills/__init__.py +5 -0
- pretensor/skills/generator.py +235 -0
- pretensor/staleness/__init__.py +15 -0
- pretensor/staleness/graph_patcher.py +355 -0
- pretensor/staleness/impact_analyzer.py +162 -0
- pretensor/staleness/snapshot_store.py +38 -0
- pretensor/validation/__init__.py +9 -0
- pretensor/validation/query_validator.py +436 -0
- pretensor/visibility/__init__.py +23 -0
- pretensor/visibility/config.py +126 -0
- pretensor/visibility/filter.py +143 -0
- pretensor/visibility/kuzu_helpers.py +32 -0
- pretensor/visibility/runtime.py +36 -0
- pretensor/visibility/sync_grants.py +188 -0
- pretensor-0.1.0.dist-info/METADATA +251 -0
- pretensor-0.1.0.dist-info/RECORD +198 -0
- pretensor-0.1.0.dist-info/WHEEL +4 -0
- pretensor-0.1.0.dist-info/entry_points.txt +2 -0
- 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]
|