fix(dwh): align adapter sampling contract

This commit is contained in:
2026-07-11 20:14:23 +02:00
parent 717e5ecced
commit 216984aac8
11 changed files with 239 additions and 76 deletions
+56
View File
@@ -5,10 +5,66 @@ 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: