Files
ThothII/harness/tht/db/sampling.py
T

217 lines
8.4 KiB
Python

import logging
from dataclasses import dataclass
from sqlalchemy import Engine, text
from tht.config import ExamplesConfig, LshConfig
from tht.mschema.models import Annotations, PhysicalSchema
from tht.ports.dwh import DistinctValues
logger = logging.getLogger(__name__)
TEXT_TYPE_PREFIXES = ("text", "varchar", "character", "char")
DEFAULT_DISTINCT_VALUES_LIMIT = 1000
def _quoted_top_values_query(engine: Engine, schema: str, table: str, column: str):
quote = engine.dialect.identifier_preparer.quote
identifier = quote(column)
return text(
f"SELECT {identifier} AS value FROM {quote(schema)}.{quote(table)} "
f"WHERE {identifier} IS NOT NULL GROUP BY {identifier} "
f"ORDER BY count(*) DESC, {identifier} LIMIT :lim"
)
def sample_column(
engine: Engine, schema: str, table: str, column: str, *, limit: int
) -> list[object]:
if limit <= 0:
raise ValueError("limit must be a positive integer")
query = _quoted_top_values_query(engine, schema, table, column)
with engine.connect() as conn:
rows = conn.execute(query, {"lim": limit}).fetchall()
return [row[0] for row in rows]
def sample_column_rest(
client, schema: str, table: str, column: str, *, limit: int
) -> list[object]:
if limit <= 0:
raise ValueError("limit must be a positive integer")
rows = client.top_values(schema, table, column, limit)
return [row["value"] for row in rows if row.get("value") is not None]
def distinct_values(
engine: Engine,
schema: str,
table: str,
column: str,
*,
max_values: int = DEFAULT_DISTINCT_VALUES_LIMIT,
) -> DistinctValues:
values = sample_column(engine, schema, table, column, limit=max_values + 1)
return DistinctValues(values=values[:max_values], truncated=len(values) > max_values)
def distinct_values_rest(
client,
schema: str,
table: str,
column: str,
*,
max_values: int = DEFAULT_DISTINCT_VALUES_LIMIT,
) -> DistinctValues:
values = sample_column_rest(client, schema, table, column, limit=max_values + 1)
return DistinctValues(values=values[:max_values], truncated=len(values) > max_values)
def is_text_type(pg_type: str) -> bool:
return pg_type.lower().startswith(TEXT_TYPE_PREFIXES)
def add_examples(engine: Engine, physical: PhysicalSchema, cfg: ExamplesConfig) -> None:
"""Campiona i valori distinti piu' frequenti delle colonne testuali (in-place)."""
schema = physical.db_schema
with engine.connect() as conn:
for table_name, table in physical.tables.items():
for column_name, column in table.columns.items():
if not is_text_type(column.type):
continue
q = text(f'''
SELECT "{column_name}" FROM (
SELECT "{column_name}", count(*) AS _freq
FROM "{schema}"."{table_name}"
WHERE "{column_name}" IS NOT NULL AND length("{column_name}") > 0
GROUP BY "{column_name}"
ORDER BY _freq DESC
LIMIT :lim
) AS sub
''')
try:
rows = conn.execute(q, {"lim": cfg.max_per_column}).fetchall()
except Exception as e: # colonna non leggibile: si salta, non si interrompe
logger.warning("Campionamento saltato per %s.%s: %s", table_name, column_name, e)
continue
column.examples = [str(r[0]) for r in rows]
def add_examples_rest(client, physical: PhysicalSchema, cfg: ExamplesConfig) -> None:
"""Variante REST di add_examples: valori più frequenti via rpc `top_values`."""
schema = physical.db_schema
for table_name, table in physical.tables.items():
for column_name, column in table.columns.items():
if not is_text_type(column.type):
continue
rows = client.top_values(schema, table_name, column_name, cfg.max_per_column)
column.examples = [str(r["value"]) for r in rows if r["value"] not in (None, "")]
@dataclass
class SkippedColumn:
table: str
column: str
reason: str
@dataclass
class TruncatedColumn:
table: str
column: str
indexed: int # quanti valori (i più frequenti) sono stati indicizzati
def unique_values_for_lsh(
engine: Engine,
physical: PhysicalSchema,
cfg: LshConfig,
annotations: Annotations | None = None,
) -> tuple[dict[str, dict[str, list[str]]], list[SkippedColumn], list[TruncatedColumn]]:
"""Valori delle colonne testuali *eligible* per l'indice LSH.
Indicizza solo colonne con eligibilità effettiva True (le `wide_text` sono escluse:
vedi principio di column eligibility). Estrae i valori distinti *più frequenti*
(ORDER BY frequenza); se superano `max_values_per_column` la colonna è troncata e
segnalata (mai tagliata in silenzio).
"""
from tht.mschema.eligibility import effective_eligibility
annotations = annotations or Annotations()
schema = physical.db_schema
values: dict[str, dict[str, list[str]]] = {}
skipped: list[SkippedColumn] = []
truncated: list[TruncatedColumn] = []
with engine.connect() as conn:
for table_name, table in physical.tables.items():
table_ann = annotations.tables.get(table_name)
for column_name, column in table.columns.items():
if not is_text_type(column.type):
continue
ann_col = table_ann.columns.get(column_name) if table_ann else None
if not effective_eligibility(column, ann_col)[0]:
continue
q = text(f'''
SELECT "{column_name}" FROM (
SELECT "{column_name}", count(*) AS _freq
FROM "{schema}"."{table_name}"
WHERE "{column_name}" IS NOT NULL AND length("{column_name}") > 0
GROUP BY "{column_name}"
ORDER BY _freq DESC, "{column_name}"
LIMIT :lim
) AS sub
''')
try:
rows = conn.execute(q, {"lim": cfg.max_values_per_column}).fetchall()
except Exception as e:
skipped.append(SkippedColumn(table_name, column_name, f"errore: {e}"))
continue
vals = [str(r[0]) for r in rows]
if not vals:
continue
values.setdefault(table_name, {})[column_name] = vals
if len(vals) >= cfg.max_values_per_column:
truncated.append(TruncatedColumn(table_name, column_name, len(vals)))
return values, skipped, truncated
def unique_values_for_lsh_rest(
client,
physical: PhysicalSchema,
cfg: LshConfig,
annotations: Annotations | None = None,
) -> tuple[dict[str, dict[str, list[str]]], list[SkippedColumn], list[TruncatedColumn]]:
"""Variante REST di unique_values_for_lsh: valori più frequenti via rpc `top_values`.
Stessa logica di eligibility e di segnalazione del troncamento del transport diretto.
"""
from tht.mschema.eligibility import effective_eligibility
annotations = annotations or Annotations()
schema = physical.db_schema
values: dict[str, dict[str, list[str]]] = {}
skipped: list[SkippedColumn] = []
truncated: list[TruncatedColumn] = []
for table_name, table in physical.tables.items():
table_ann = annotations.tables.get(table_name)
for column_name, column in table.columns.items():
if not is_text_type(column.type):
continue
ann_col = table_ann.columns.get(column_name) if table_ann else None
if not effective_eligibility(column, ann_col)[0]:
continue
try:
rows = client.top_values(
schema, table_name, column_name, cfg.max_values_per_column
)
except Exception as e:
skipped.append(SkippedColumn(table_name, column_name, f"errore: {e}"))
continue
vals = [str(r["value"]) for r in rows if r["value"] not in (None, "")]
if not vals:
continue
values.setdefault(table_name, {})[column_name] = vals
if len(vals) >= cfg.max_values_per_column:
truncated.append(TruncatedColumn(table_name, column_name, len(vals)))
return values, skipped, truncated