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,62 @@
|
|
|
1
|
+
"""LLM interface types and text utilities (LLM infrastructure removed)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import json
|
|
6
|
+
import re
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
from typing import Any, Literal
|
|
9
|
+
|
|
10
|
+
from pretensor.errors import PretensorError
|
|
11
|
+
|
|
12
|
+
__all__ = [
|
|
13
|
+
"ChatMessage",
|
|
14
|
+
"LlmUsage",
|
|
15
|
+
"LlmBudgetExceededError",
|
|
16
|
+
"strip_json_fence",
|
|
17
|
+
"strip_markdown_fence",
|
|
18
|
+
"parse_json_array",
|
|
19
|
+
]
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class LlmBudgetExceededError(PretensorError):
|
|
23
|
+
"""Raised when estimated spend for an index run reaches the configured ceiling."""
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@dataclass(frozen=True, slots=True)
|
|
27
|
+
class ChatMessage:
|
|
28
|
+
"""Single chat turn for the shared completion API."""
|
|
29
|
+
|
|
30
|
+
role: Literal["system", "user", "assistant"]
|
|
31
|
+
content: str
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
@dataclass(frozen=True, slots=True)
|
|
35
|
+
class LlmUsage:
|
|
36
|
+
"""Token usage reported by a provider (zeros if absent)."""
|
|
37
|
+
|
|
38
|
+
input_tokens: int
|
|
39
|
+
output_tokens: int
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def strip_json_fence(raw: str) -> str:
|
|
43
|
+
"""Remove optional ```json ... ``` wrapping from model output."""
|
|
44
|
+
return strip_markdown_fence(raw)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def strip_markdown_fence(raw: str) -> str:
|
|
48
|
+
"""Remove optional ```lang ... ``` wrapping (json, yaml, or bare)."""
|
|
49
|
+
s = raw.strip()
|
|
50
|
+
m = re.match(r"^```\w*\s*([\s\S]*?)\s*```$", s)
|
|
51
|
+
if m:
|
|
52
|
+
return m.group(1).strip()
|
|
53
|
+
return s
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def parse_json_array(raw: str) -> list[Any]:
|
|
57
|
+
"""Parse a JSON array, tolerating markdown fences."""
|
|
58
|
+
s = strip_json_fence(raw)
|
|
59
|
+
data = json.loads(s)
|
|
60
|
+
if not isinstance(data, list):
|
|
61
|
+
raise ValueError("expected JSON array")
|
|
62
|
+
return data
|
|
@@ -0,0 +1,193 @@
|
|
|
1
|
+
"""Heuristic metric templates from classified fact tables.
|
|
2
|
+
|
|
3
|
+
Templates are generated deterministically as ``SELECT SUM(col) AS total FROM schema.table``
|
|
4
|
+
using ANSI double-quote identifier quoting. LLM-based refinement (COUNT, AVG, multi-column
|
|
5
|
+
aggregations) is deferred to a follow-up task.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import logging
|
|
11
|
+
import re
|
|
12
|
+
from datetime import datetime, timezone
|
|
13
|
+
|
|
14
|
+
from pretensor.core.ids import metric_template_node_id
|
|
15
|
+
from pretensor.core.store import KuzuStore
|
|
16
|
+
from pretensor.validation.query_validator import QueryValidator
|
|
17
|
+
|
|
18
|
+
__all__ = ["MetricTemplateBuilder"]
|
|
19
|
+
|
|
20
|
+
logger = logging.getLogger(__name__)
|
|
21
|
+
|
|
22
|
+
_METRIC_NAME_PARTS = frozenset(
|
|
23
|
+
{
|
|
24
|
+
"amount",
|
|
25
|
+
"total",
|
|
26
|
+
"price",
|
|
27
|
+
"cost",
|
|
28
|
+
"revenue",
|
|
29
|
+
"qty",
|
|
30
|
+
"quantity",
|
|
31
|
+
"count",
|
|
32
|
+
"fee",
|
|
33
|
+
"fees",
|
|
34
|
+
"tax",
|
|
35
|
+
"payment",
|
|
36
|
+
"payments",
|
|
37
|
+
"subtotal",
|
|
38
|
+
"discount",
|
|
39
|
+
}
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
_DIALECT_POSTGRESQL = "postgresql"
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def _pg_ident(part: str) -> str:
|
|
47
|
+
"""ANSI / PostgreSQL double-quote identifier escaping."""
|
|
48
|
+
return '"' + part.replace('"', '""') + '"'
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _is_numeric_type(data_type: str) -> bool:
|
|
52
|
+
t = data_type.lower()
|
|
53
|
+
return any(
|
|
54
|
+
x in t
|
|
55
|
+
for x in (
|
|
56
|
+
"int",
|
|
57
|
+
"numeric",
|
|
58
|
+
"decimal",
|
|
59
|
+
"float",
|
|
60
|
+
"double",
|
|
61
|
+
"real",
|
|
62
|
+
"money",
|
|
63
|
+
"smallserial",
|
|
64
|
+
"serial",
|
|
65
|
+
"bigserial",
|
|
66
|
+
)
|
|
67
|
+
)
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def _column_signals_metric(column_name: str) -> bool:
|
|
71
|
+
base = column_name.lower().strip()
|
|
72
|
+
if not base:
|
|
73
|
+
return False
|
|
74
|
+
for part in _METRIC_NAME_PARTS:
|
|
75
|
+
if part in base:
|
|
76
|
+
return True
|
|
77
|
+
return False
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def _slugify_metric_name(schema: str, table: str, column: str) -> str:
|
|
81
|
+
raw = f"sum_{schema}_{table}_{column}"
|
|
82
|
+
s = re.sub(r"[^a-zA-Z0-9]+", "_", raw).strip("_").lower()
|
|
83
|
+
return s or "metric"
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
class MetricTemplateBuilder:
|
|
87
|
+
"""Create validated ``MetricTemplate`` nodes from fact tables and metric-like columns."""
|
|
88
|
+
|
|
89
|
+
def __init__(self, store: KuzuStore) -> None:
|
|
90
|
+
self._store = store
|
|
91
|
+
|
|
92
|
+
def build(self, database_key: str) -> int:
|
|
93
|
+
"""Clear prior templates for this DB, then write new validated templates. Returns count."""
|
|
94
|
+
rows_cn = self._store.query_all_rows(
|
|
95
|
+
"""
|
|
96
|
+
MATCH (t:SchemaTable {database: $db})
|
|
97
|
+
RETURN t.connection_name
|
|
98
|
+
LIMIT 1
|
|
99
|
+
""",
|
|
100
|
+
{"db": database_key},
|
|
101
|
+
)
|
|
102
|
+
if not rows_cn:
|
|
103
|
+
return 0
|
|
104
|
+
connection_name = str(rows_cn[0][0])
|
|
105
|
+
|
|
106
|
+
self._store.execute_write(
|
|
107
|
+
"""
|
|
108
|
+
MATCH (m:MetricTemplate {connection_name: $cn, database: $db})
|
|
109
|
+
DETACH DELETE m
|
|
110
|
+
""",
|
|
111
|
+
{"cn": connection_name, "db": database_key},
|
|
112
|
+
)
|
|
113
|
+
|
|
114
|
+
candidates = self._store.query_all_rows(
|
|
115
|
+
"""
|
|
116
|
+
MATCH (t:SchemaTable {connection_name: $cn, database: $db, role: 'fact'})
|
|
117
|
+
-[:HAS_COLUMN]->(c:SchemaColumn)
|
|
118
|
+
RETURN t.node_id, t.schema_name, t.table_name, c.column_name, c.data_type
|
|
119
|
+
ORDER BY t.schema_name, t.table_name, c.ordinal_position, c.column_name
|
|
120
|
+
""",
|
|
121
|
+
{"cn": connection_name, "db": database_key},
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
written = 0
|
|
125
|
+
now = datetime.now(timezone.utc).isoformat()
|
|
126
|
+
validator = QueryValidator(
|
|
127
|
+
self._store,
|
|
128
|
+
connection_name=connection_name,
|
|
129
|
+
database_key=database_key,
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
for tid, sn, tn, cname, dtype in candidates:
|
|
133
|
+
if not _is_numeric_type(str(dtype or "")):
|
|
134
|
+
continue
|
|
135
|
+
if not _column_signals_metric(str(cname)):
|
|
136
|
+
continue
|
|
137
|
+
schema_s = str(sn)
|
|
138
|
+
table_s = str(tn)
|
|
139
|
+
col_s = str(cname)
|
|
140
|
+
qschema = _pg_ident(schema_s)
|
|
141
|
+
qtable = _pg_ident(table_s)
|
|
142
|
+
qcol = _pg_ident(col_s)
|
|
143
|
+
sql = f"SELECT SUM({qcol}) AS total FROM {qschema}.{qtable}"
|
|
144
|
+
result = validator.validate(sql)
|
|
145
|
+
name = _slugify_metric_name(schema_s, table_s, col_s)
|
|
146
|
+
node_id = metric_template_node_id(connection_name, database_key, name)
|
|
147
|
+
display = f"Sum of {col_s} ({schema_s}.{table_s})"
|
|
148
|
+
desc = (
|
|
149
|
+
f"Aggregate total of column {col_s} on fact table {schema_s}.{table_s}."
|
|
150
|
+
)
|
|
151
|
+
tables_used = [f"{schema_s}.{table_s}"]
|
|
152
|
+
err_list: list[str] = []
|
|
153
|
+
if not result.valid:
|
|
154
|
+
err_list = (
|
|
155
|
+
result.syntax_errors
|
|
156
|
+
+ [f"missing_table:{t}" for t in result.missing_tables]
|
|
157
|
+
+ [f"missing_column:{c}" for c in result.missing_columns]
|
|
158
|
+
+ [f"join:{j.message}" for j in result.invalid_joins]
|
|
159
|
+
)
|
|
160
|
+
self._store.upsert_metric_template(
|
|
161
|
+
node_id=node_id,
|
|
162
|
+
connection_name=connection_name,
|
|
163
|
+
database=database_key,
|
|
164
|
+
dialect=_DIALECT_POSTGRESQL,
|
|
165
|
+
name=name,
|
|
166
|
+
display_name=display,
|
|
167
|
+
description=desc,
|
|
168
|
+
sql_template=sql,
|
|
169
|
+
tables_used=tables_used,
|
|
170
|
+
validated=result.valid,
|
|
171
|
+
validation_errors=err_list,
|
|
172
|
+
generated_at_iso=now,
|
|
173
|
+
stale=False,
|
|
174
|
+
depends_on_table_node_ids=[str(tid)],
|
|
175
|
+
)
|
|
176
|
+
written += 1
|
|
177
|
+
if not result.valid:
|
|
178
|
+
logger.debug(
|
|
179
|
+
"Metric template %s not validated: %s", name, err_list[:3]
|
|
180
|
+
)
|
|
181
|
+
|
|
182
|
+
return written
|
|
183
|
+
|
|
184
|
+
@staticmethod
|
|
185
|
+
def mark_stale_for_database(store: KuzuStore, database_key: str) -> None:
|
|
186
|
+
"""Set ``stale`` on templates when the indexed graph may have drifted (e.g. reindex)."""
|
|
187
|
+
store.execute_write(
|
|
188
|
+
"""
|
|
189
|
+
MATCH (m:MetricTemplate {database: $db})
|
|
190
|
+
SET m.stale = true
|
|
191
|
+
""",
|
|
192
|
+
{"db": database_key},
|
|
193
|
+
)
|
|
@@ -0,0 +1,364 @@
|
|
|
1
|
+
"""Run clustering, labeling, and join-path precomputation after a graph build.
|
|
2
|
+
|
|
3
|
+
The pipeline is expressed as a sequence of named :class:`PipelineStep` objects
|
|
4
|
+
executed by :class:`PipelineRunner`. Plugins can inject additional steps (e.g.
|
|
5
|
+
``llm_refine``, ``feedback_score``, ``semantic_propose``) between the built-in
|
|
6
|
+
OSS steps by calling :func:`build_oss_pipeline` and registering extra steps
|
|
7
|
+
before calling ``runner.run(ctx)``.
|
|
8
|
+
|
|
9
|
+
OSS step order (resolved from dependencies):
|
|
10
|
+
embedding_index → classify → cluster → label → join_paths
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import asyncio
|
|
16
|
+
import logging
|
|
17
|
+
import os
|
|
18
|
+
import time
|
|
19
|
+
from collections.abc import Callable
|
|
20
|
+
from typing import TYPE_CHECKING
|
|
21
|
+
|
|
22
|
+
from pretensor.config import GraphConfig
|
|
23
|
+
from pretensor.core.store import KuzuStore
|
|
24
|
+
from pretensor.intelligence.cluster_labeler import ClusterLabeler
|
|
25
|
+
from pretensor.intelligence.clustering import Cluster, ClusteringEngine
|
|
26
|
+
from pretensor.intelligence.embeddings import embeddings_disabled_via_env
|
|
27
|
+
from pretensor.intelligence.graph_export import GraphExporter
|
|
28
|
+
from pretensor.intelligence.join_paths import JoinPathEngine
|
|
29
|
+
from pretensor.intelligence.schema_classification import (
|
|
30
|
+
classify_database_tables_async,
|
|
31
|
+
compute_cluster_schema_patterns,
|
|
32
|
+
load_fk_reference_pairs,
|
|
33
|
+
)
|
|
34
|
+
from pretensor.intelligence.steps import PipelineContext, PipelineRunner, PipelineStep
|
|
35
|
+
from pretensor.intelligence.steps_embedding import EmbeddingIndexStep
|
|
36
|
+
|
|
37
|
+
if TYPE_CHECKING:
|
|
38
|
+
from pretensor.config import PretensorConfig
|
|
39
|
+
|
|
40
|
+
__all__ = [
|
|
41
|
+
"run_intelligence_layer",
|
|
42
|
+
"run_intelligence_layer_sync",
|
|
43
|
+
"build_oss_pipeline",
|
|
44
|
+
]
|
|
45
|
+
|
|
46
|
+
logger = logging.getLogger(__name__)
|
|
47
|
+
|
|
48
|
+
_PROFILE_INDEX = os.environ.get("PRETENSOR_PROFILE_INDEX", "").lower() not in (
|
|
49
|
+
"",
|
|
50
|
+
"0",
|
|
51
|
+
"false",
|
|
52
|
+
"no",
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
# ---------------------------------------------------------------------------
|
|
56
|
+
# Context keys — stable names shared between steps
|
|
57
|
+
# ---------------------------------------------------------------------------
|
|
58
|
+
# Canonical context keys live in ``steps.py`` so step modules can import
|
|
59
|
+
# the same constants instead of defining sibling string literals (which
|
|
60
|
+
# would silently no-op on drift). Re-exported here for backward compat
|
|
61
|
+
# with any tests / cloud extensions that imported them from this module.
|
|
62
|
+
from pretensor.intelligence.steps import ( # noqa: E402 (post-imports to avoid cycle re-order churn)
|
|
63
|
+
_CTX_CLUSTERS,
|
|
64
|
+
_CTX_CONFIG,
|
|
65
|
+
_CTX_DATABASE_KEY,
|
|
66
|
+
_CTX_EMBEDDINGS_CONFIG,
|
|
67
|
+
_CTX_EMBEDDINGS_PRECOMPUTED,
|
|
68
|
+
_CTX_GRAPH,
|
|
69
|
+
_CTX_PATTERNS,
|
|
70
|
+
_CTX_ROLE_BY_TABLE,
|
|
71
|
+
_CTX_STORE,
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
# ---------------------------------------------------------------------------
|
|
75
|
+
# Built-in OSS steps
|
|
76
|
+
# ---------------------------------------------------------------------------
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
class _ClassifyStep:
|
|
80
|
+
"""Classify each table in the database by role (fact/dimension/bridge/lookup)."""
|
|
81
|
+
|
|
82
|
+
name = "classify"
|
|
83
|
+
# Depends on embedding_index so the optional role vote (role_weight > 0)
|
|
84
|
+
# reads vectors computed in THIS run; without the edge the vote would
|
|
85
|
+
# silently see an empty store on the first index and degrade to the
|
|
86
|
+
# heuristic-only path until the next reindex.
|
|
87
|
+
dependencies: list[str] = ["embedding_index"]
|
|
88
|
+
|
|
89
|
+
async def execute(self, ctx: PipelineContext) -> None:
|
|
90
|
+
store: KuzuStore = ctx.get(_CTX_STORE)
|
|
91
|
+
database_key: str = ctx.get(_CTX_DATABASE_KEY)
|
|
92
|
+
emb_cfg = ctx.get(_CTX_EMBEDDINGS_CONFIG)
|
|
93
|
+
|
|
94
|
+
# When the user opts into the role-classification embedding vote,
|
|
95
|
+
# instantiate a real client so the classifier can compute per-role
|
|
96
|
+
# centroids + per-table votes. Default keeps role_weight=0 and
|
|
97
|
+
# embedding_client=None → heuristic-only path.
|
|
98
|
+
embedding_client = None
|
|
99
|
+
role_weight = 0.0
|
|
100
|
+
if emb_cfg is not None and emb_cfg.role_weight > 0.0:
|
|
101
|
+
from pretensor.intelligence.embeddings import (
|
|
102
|
+
get_default_embedding_client,
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
embedding_client = get_default_embedding_client()
|
|
106
|
+
role_weight = emb_cfg.role_weight
|
|
107
|
+
|
|
108
|
+
_t = time.perf_counter() if _PROFILE_INDEX else 0.0
|
|
109
|
+
role_by_table = await classify_database_tables_async(
|
|
110
|
+
store,
|
|
111
|
+
database_key,
|
|
112
|
+
role_weight=role_weight,
|
|
113
|
+
embedding_client=embedding_client,
|
|
114
|
+
)
|
|
115
|
+
if _PROFILE_INDEX:
|
|
116
|
+
print(
|
|
117
|
+
f"[profile] intelligence.classify_database_tables: {(time.perf_counter() - _t) * 1000:.0f}ms",
|
|
118
|
+
flush=True,
|
|
119
|
+
)
|
|
120
|
+
ctx.set(_CTX_ROLE_BY_TABLE, role_by_table)
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
class _ClusterStep:
|
|
124
|
+
"""Run community detection on the FK graph to group tables into clusters."""
|
|
125
|
+
|
|
126
|
+
name = "cluster"
|
|
127
|
+
dependencies: list[str] = ["classify"]
|
|
128
|
+
|
|
129
|
+
async def execute(self, ctx: PipelineContext) -> None:
|
|
130
|
+
store: KuzuStore = ctx.get(_CTX_STORE)
|
|
131
|
+
database_key: str = ctx.get(_CTX_DATABASE_KEY)
|
|
132
|
+
cfg: GraphConfig = ctx.get(_CTX_CONFIG)
|
|
133
|
+
graph = ctx.get(_CTX_GRAPH)
|
|
134
|
+
|
|
135
|
+
_t = time.perf_counter() if _PROFILE_INDEX else 0.0
|
|
136
|
+
clusters: list[Cluster] = ClusteringEngine(cfg).cluster(graph)
|
|
137
|
+
if _PROFILE_INDEX:
|
|
138
|
+
print(
|
|
139
|
+
f"[profile] intelligence.clustering: {(time.perf_counter() - _t) * 1000:.0f}ms ({len(clusters)} clusters)",
|
|
140
|
+
flush=True,
|
|
141
|
+
)
|
|
142
|
+
|
|
143
|
+
_t = time.perf_counter() if _PROFILE_INDEX else 0.0
|
|
144
|
+
fk_pairs = load_fk_reference_pairs(store, database_key)
|
|
145
|
+
role_by_table = ctx.get(_CTX_ROLE_BY_TABLE)
|
|
146
|
+
patterns = compute_cluster_schema_patterns(clusters, role_by_table, fk_pairs)
|
|
147
|
+
if _PROFILE_INDEX:
|
|
148
|
+
print(
|
|
149
|
+
f"[profile] intelligence.cluster_schema_patterns: {(time.perf_counter() - _t) * 1000:.0f}ms",
|
|
150
|
+
flush=True,
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
ctx.set(_CTX_CLUSTERS, clusters)
|
|
154
|
+
ctx.set(_CTX_PATTERNS, patterns)
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
class _LabelStep:
|
|
158
|
+
"""Assign domain labels to clusters and persist them to Kuzu.
|
|
159
|
+
|
|
160
|
+
Declares ``embedding_index`` as a dependency so the topological
|
|
161
|
+
runner orders the embedding writes before the labeler reads them.
|
|
162
|
+
Without this, correctness would rely on insertion-order tie-breaking
|
|
163
|
+
inside the runner — fragile to a future scheduling change.
|
|
164
|
+
"""
|
|
165
|
+
|
|
166
|
+
name = "label"
|
|
167
|
+
dependencies: list[str] = ["cluster", "embedding_index"]
|
|
168
|
+
|
|
169
|
+
async def execute(self, ctx: PipelineContext) -> None:
|
|
170
|
+
store: KuzuStore = ctx.get(_CTX_STORE)
|
|
171
|
+
database_key: str = ctx.get(_CTX_DATABASE_KEY)
|
|
172
|
+
clusters: list[Cluster] = ctx.get(_CTX_CLUSTERS)
|
|
173
|
+
patterns = ctx.get(_CTX_PATTERNS)
|
|
174
|
+
emb_cfg = ctx.get(_CTX_EMBEDDINGS_CONFIG)
|
|
175
|
+
|
|
176
|
+
# When the user opts into clustering blend, thread an embedding
|
|
177
|
+
# resolver into the labeler so per-cluster centroids can act as a
|
|
178
|
+
# tiebreaker after role-weighted degree + row count. Default is
|
|
179
|
+
# ``None`` → labeler behavior is byte-identical to the pre-embedding
|
|
180
|
+
# tiebreaker chain.
|
|
181
|
+
embedding_resolver = None
|
|
182
|
+
if emb_cfg is not None and emb_cfg.cluster_blend > 0.0:
|
|
183
|
+
embedding_resolver = _build_embedding_resolver(store, database_key)
|
|
184
|
+
|
|
185
|
+
_t = time.perf_counter() if _PROFILE_INDEX else 0.0
|
|
186
|
+
labeler = ClusterLabeler(store, embedding_resolver=embedding_resolver)
|
|
187
|
+
await labeler.label_and_persist(
|
|
188
|
+
clusters, database_key, cluster_schema_patterns=patterns
|
|
189
|
+
)
|
|
190
|
+
if _PROFILE_INDEX:
|
|
191
|
+
print(
|
|
192
|
+
f"[profile] intelligence.cluster_labeling: {(time.perf_counter() - _t) * 1000:.0f}ms",
|
|
193
|
+
flush=True,
|
|
194
|
+
)
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
def _build_embedding_resolver(
|
|
198
|
+
store: KuzuStore, database_key: str
|
|
199
|
+
) -> Callable[[str], list[float] | None]:
|
|
200
|
+
"""One-shot fetch of every embedded table for the database, return dict lookup.
|
|
201
|
+
|
|
202
|
+
Issuing one bulk query and resolving via ``dict.get`` is materially cheaper
|
|
203
|
+
than per-table Kuzu round-trips inside ``label_and_persist`` (which the
|
|
204
|
+
labeler can call many times per cluster).
|
|
205
|
+
|
|
206
|
+
Reuses :meth:`KuzuStore.iter_table_embeddings` so the embedding-fetch
|
|
207
|
+
Cypher lives in exactly one place — if the ``SchemaTable.embedding``
|
|
208
|
+
storage shape changes (column type, optional-match shape, etc.) the
|
|
209
|
+
iterator is the single source of truth that needs updating.
|
|
210
|
+
"""
|
|
211
|
+
cache: dict[str, list[float]] = {}
|
|
212
|
+
for row in store.iter_table_embeddings(database=database_key):
|
|
213
|
+
if row.embedding is None:
|
|
214
|
+
continue
|
|
215
|
+
# ``iter_table_embeddings`` may yield the same node_id multiple
|
|
216
|
+
# times (one row per cluster the table belongs to via OPTIONAL
|
|
217
|
+
# MATCH); the embedding is identical across those duplicates so
|
|
218
|
+
# last-write-wins is safe and the dict naturally dedups.
|
|
219
|
+
cache[row.node_id] = row.embedding
|
|
220
|
+
|
|
221
|
+
def resolver(nid: str) -> list[float] | None:
|
|
222
|
+
return cache.get(nid)
|
|
223
|
+
|
|
224
|
+
return resolver
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
class _JoinPathsStep:
|
|
228
|
+
"""Precompute join paths between table pairs within and across clusters."""
|
|
229
|
+
|
|
230
|
+
name = "join_paths"
|
|
231
|
+
dependencies: list[str] = ["label"]
|
|
232
|
+
|
|
233
|
+
async def execute(self, ctx: PipelineContext) -> None:
|
|
234
|
+
store: KuzuStore = ctx.get(_CTX_STORE)
|
|
235
|
+
database_key: str = ctx.get(_CTX_DATABASE_KEY)
|
|
236
|
+
cfg: GraphConfig = ctx.get(_CTX_CONFIG)
|
|
237
|
+
|
|
238
|
+
_t = time.perf_counter() if _PROFILE_INDEX else 0.0
|
|
239
|
+
JoinPathEngine(store).precompute(database_key, cfg)
|
|
240
|
+
if _PROFILE_INDEX:
|
|
241
|
+
print(
|
|
242
|
+
f"[profile] intelligence.join_paths_precompute: {(time.perf_counter() - _t) * 1000:.0f}ms",
|
|
243
|
+
flush=True,
|
|
244
|
+
)
|
|
245
|
+
|
|
246
|
+
|
|
247
|
+
# ---------------------------------------------------------------------------
|
|
248
|
+
# Public factory and orchestration helpers
|
|
249
|
+
# ---------------------------------------------------------------------------
|
|
250
|
+
|
|
251
|
+
_OSS_STEPS: list[PipelineStep] = [
|
|
252
|
+
_ClassifyStep(),
|
|
253
|
+
_ClusterStep(),
|
|
254
|
+
EmbeddingIndexStep(),
|
|
255
|
+
_LabelStep(),
|
|
256
|
+
_JoinPathsStep(),
|
|
257
|
+
]
|
|
258
|
+
|
|
259
|
+
|
|
260
|
+
def build_oss_pipeline() -> PipelineRunner:
|
|
261
|
+
"""Return a :class:`PipelineRunner` pre-loaded with the five OSS steps.
|
|
262
|
+
|
|
263
|
+
Plugins can append additional steps via :meth:`PipelineRunner.register`
|
|
264
|
+
before calling ``runner.run(ctx)``.
|
|
265
|
+
|
|
266
|
+
Returns:
|
|
267
|
+
A :class:`PipelineRunner` with steps: embedding_index → classify →
|
|
268
|
+
cluster → label → join_paths.
|
|
269
|
+
"""
|
|
270
|
+
return PipelineRunner(list(_OSS_STEPS))
|
|
271
|
+
|
|
272
|
+
|
|
273
|
+
async def run_intelligence_layer(
|
|
274
|
+
store: KuzuStore,
|
|
275
|
+
database_key: str,
|
|
276
|
+
*,
|
|
277
|
+
config: GraphConfig | PretensorConfig | None = None,
|
|
278
|
+
embeddings_precomputed: bool = False,
|
|
279
|
+
) -> None:
|
|
280
|
+
"""Clear prior intelligence rows, cluster tables, label, precompute join paths.
|
|
281
|
+
|
|
282
|
+
Delegates to :func:`build_oss_pipeline` so external callers get identical
|
|
283
|
+
semantics while plugins can extend the pipeline via :func:`build_oss_pipeline`
|
|
284
|
+
directly.
|
|
285
|
+
|
|
286
|
+
Args:
|
|
287
|
+
store: The Kuzu graph store.
|
|
288
|
+
database_key: Logical database name (``SchemaTable.database``).
|
|
289
|
+
config: Optional tuning overrides; accepts :class:`GraphConfig` or
|
|
290
|
+
:class:`PretensorConfig`. When a :class:`PretensorConfig` is given,
|
|
291
|
+
its ``graph`` sub-field is used for clustering/join-path tuning.
|
|
292
|
+
Defaults to :class:`GraphConfig` with OSS defaults.
|
|
293
|
+
embeddings_precomputed: Set True by callers (index/reindex) that
|
|
294
|
+
already ran :func:`compute_table_embeddings` before relationship
|
|
295
|
+
discovery; the ``embedding_index`` pipeline step then skips its
|
|
296
|
+
redundant recompute.
|
|
297
|
+
"""
|
|
298
|
+
from pretensor.config import EmbeddingsConfig, PretensorConfig
|
|
299
|
+
|
|
300
|
+
if isinstance(config, PretensorConfig):
|
|
301
|
+
cfg = config.graph
|
|
302
|
+
emb_cfg = config.embeddings
|
|
303
|
+
else:
|
|
304
|
+
cfg = config or GraphConfig()
|
|
305
|
+
emb_cfg = EmbeddingsConfig()
|
|
306
|
+
# Single choke point for the env kill switch: every embedding consumer
|
|
307
|
+
# below (classify role vote, cluster blend, label resolver, index step)
|
|
308
|
+
# reads this config, so swapping in the all-off default forces the null
|
|
309
|
+
# path even when the caller's toggles are on.
|
|
310
|
+
if embeddings_disabled_via_env() and emb_cfg != EmbeddingsConfig():
|
|
311
|
+
logger.info(
|
|
312
|
+
"intelligence: PRETENSOR_EMBEDDINGS_DISABLED set; "
|
|
313
|
+
"forcing the null embeddings config"
|
|
314
|
+
)
|
|
315
|
+
emb_cfg = EmbeddingsConfig()
|
|
316
|
+
store.ensure_schema()
|
|
317
|
+
store.clear_intelligence_artifacts()
|
|
318
|
+
|
|
319
|
+
_t = time.perf_counter() if _PROFILE_INDEX else 0.0
|
|
320
|
+
exporter = GraphExporter(store)
|
|
321
|
+
graph = exporter.to_igraph(
|
|
322
|
+
database_key, config=cfg, cluster_blend=emb_cfg.cluster_blend
|
|
323
|
+
)
|
|
324
|
+
if _PROFILE_INDEX:
|
|
325
|
+
print(
|
|
326
|
+
f"[profile] intelligence.to_igraph: {(time.perf_counter() - _t) * 1000:.0f}ms (vcount={graph.vcount()}, ecount={graph.ecount()})",
|
|
327
|
+
flush=True,
|
|
328
|
+
)
|
|
329
|
+
if graph.vcount() == 0:
|
|
330
|
+
logger.info(
|
|
331
|
+
"Intelligence layer skipped: no tables for database %s", database_key
|
|
332
|
+
)
|
|
333
|
+
return
|
|
334
|
+
|
|
335
|
+
ctx = PipelineContext(
|
|
336
|
+
**{
|
|
337
|
+
_CTX_STORE: store,
|
|
338
|
+
_CTX_DATABASE_KEY: database_key,
|
|
339
|
+
_CTX_CONFIG: cfg,
|
|
340
|
+
_CTX_GRAPH: graph,
|
|
341
|
+
_CTX_EMBEDDINGS_CONFIG: emb_cfg,
|
|
342
|
+
_CTX_EMBEDDINGS_PRECOMPUTED: embeddings_precomputed,
|
|
343
|
+
}
|
|
344
|
+
)
|
|
345
|
+
runner = build_oss_pipeline()
|
|
346
|
+
await runner.run(ctx)
|
|
347
|
+
|
|
348
|
+
|
|
349
|
+
def run_intelligence_layer_sync(
|
|
350
|
+
store: KuzuStore,
|
|
351
|
+
database_key: str,
|
|
352
|
+
*,
|
|
353
|
+
config: GraphConfig | PretensorConfig | None = None,
|
|
354
|
+
embeddings_precomputed: bool = False,
|
|
355
|
+
) -> None:
|
|
356
|
+
"""Sync wrapper for :func:`run_intelligence_layer` (CLI / builder)."""
|
|
357
|
+
asyncio.run(
|
|
358
|
+
run_intelligence_layer(
|
|
359
|
+
store,
|
|
360
|
+
database_key,
|
|
361
|
+
config=config,
|
|
362
|
+
embeddings_precomputed=embeddings_precomputed,
|
|
363
|
+
)
|
|
364
|
+
)
|