Files
tianxuan/analysis/entity_detector.py
T

178 lines
6.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Auto-detect entity column from CSV schema using heuristic scoring.
Provides :func:`detect_entity_column` which scores all string/categorical
columns by keyword-name match, unique-value ratio, and dtype, returning
a ranked list of candidates and a recommended column.
Usage::
from analysis.entity_detector import detect_entity_column
result = detect_entity_column(dataset_id)
result['recommended'] # e.g. 'src_ip'
result['candidates'] # sorted list of scored columns
"""
from __future__ import annotations
from typing import Optional
import polars as pl
from .session_store import SessionStore
# ---------------------------------------------------------------------------
# Entity-related keywords (lowercased)
# ---------------------------------------------------------------------------
ENTITY_KEYWORDS: tuple[str, ...] = (
# User exact column names only — no generic/fuzzy keywords
':ips', ':ipd', ':prs', ':prd',
'server_ip', 'client_ip', 'scnt', 'dcnt',
'cnam', 'snam', '0ver', '4dur', '8ses', '2tmo', '4ksz',
'cnrs', 'isrs', '8ack', '8ppk', '8dbd',
'1ipp', '4dbn', 'tabl', 'name', 'source_node',
'cipher_suite', 'ecdhe_named_curve',
'0cph', '0crv', '0rnd', '0rnt', 'row', 'time', 'timestamp',
'latd', 'lond', 'ispn', 'orgn', 'city',
)
def _score_name(column: str) -> tuple[float, str]:
"""Score a column name against *ENTITY_KEYWORDS*.
Returns ``(score, matched_keyword_or_empty)`` where *score* is 3.0
for every matched keyword. Partial matches (keyword is a substring
of the lowercased column name) also count.
"""
col_lower = column.lower().replace('-', '_').replace('.', '_').replace(' ', '_')
for kw in ENTITY_KEYWORDS:
if kw == col_lower or col_lower.endswith(f'_{kw}') or col_lower.startswith(f'{kw}_'):
return 3.0, kw
return 0.0, ''
def detect_entity_column(dataset_id: str, entity_columns: Optional[list[str]] = None) -> dict:
"""Score all string/categorical columns and return candidates.
The scoring algorithm:
* **Name match** (``+3``): column name (normalised) is a keyword, or
ends/starts with a keyword via ``_`` separator.
* **Unique ratio** (010): columns whose ratio of unique values to
non-null values falls in the range ``[0.01, 0.50]`` receive up to
10 points (``min(ratio * 100, 10)``).
* **String dtype** (``+2``): ``Utf8`` / ``String`` / ``Categorical``
columns get a bonus.
* **Null penalty** (``*0.1``): if null ratio > 50%, score severely reduced.
* **Penalty** (``-5``): unique ratio > 0.50 (likely random IDs or
transaction-level identifiers).
Parameters
----------
dataset_id:
Dataset identifier in the session store.
entity_columns:
Optional list of column names to restrict detection to. When given,
only these columns are scored (must exist in the schema).
Returns
-------
dict
``candidates``: list of ``{column, score, unique_ratio, matched_keyword}``
sorted descending by score.
``recommended``: highest-scoring column name (or ``''`` when none found).
``recommended_multi``: list of the top 3 candidate column names (may be
shorter than 3 when fewer candidates exist).
``dataset_id``: the input identifier (echoed).
"""
store = SessionStore()
entry = store.get_dataset(dataset_id)
if entry is None:
return {
'dataset_id': dataset_id,
'candidates': [],
'recommended': '',
'error': f'Dataset not found: {dataset_id}',
}
lf: pl.LazyFrame = entry['lazyframe']
schema: dict[str, str] = entry.get('schema', {})
# Identify string / categorical columns
string_types = frozenset({'Utf8', 'String', 'Categorical', 'str', 'cat'})
candidate_scores: list[dict] = []
try:
# Collect a sample to compute unique ratios efficiently
df_sample = lf.collect(streaming=True)
except Exception as exc:
return {
'dataset_id': dataset_id,
'candidates': [],
'recommended': '',
'error': f'Failed to collect dataset: {exc}',
}
total_rows = len(df_sample)
for col_name, dtype_str in schema.items():
# If entity_columns specified, skip columns not in the list
if entity_columns is not None and col_name not in entity_columns:
continue
# Only score string-like columns
dtype_base = dtype_str.split('(')[0].strip()
if dtype_base not in string_types:
continue
series_all = df_sample[col_name]
non_null_series = series_all.drop_nulls()
non_null_count = len(non_null_series)
null_count = total_rows - non_null_count
if non_null_count == 0:
continue
unique_count = int(non_null_series.n_unique())
unique_ratio = unique_count / non_null_count
# --- scoring ---
name_score, matched_kw = _score_name(col_name)
dtype_score = 2.0 # String/categorical bonus
ratio_score = min(unique_ratio * 100.0, 10.0)
penalty = -5.0 if unique_ratio > 0.50 else 0.0
total_score = name_score + dtype_score + ratio_score + penalty
# Null ratio penalty: if more than 50% null, severely reduce score
if null_count / max(1, non_null_count) > 0.5:
total_score *= 0.1
candidate_scores.append({
'column': col_name,
'score': round(total_score, 2),
'unique_ratio': round(unique_ratio, 4),
'matched_keyword': matched_kw,
'non_null_count': non_null_count,
'unique_count': unique_count,
'null_count': null_count,
})
# Sort descending by score, then by unique_count desc (prefer richer columns)
candidate_scores.sort(key=lambda c: (c['score'], c['unique_count']), reverse=True)
recommended = candidate_scores[0]['column'] if candidate_scores else ''
recommended_multi = [c['column'] for c in candidate_scores[:3]]
return {
'dataset_id': dataset_id,
'candidates': candidate_scores,
'recommended': recommended,
'recommended_multi': recommended_multi, # NEW
'total_rows': total_rows,
'scored_columns': len(candidate_scores),
}