141 lines
5.8 KiB
Python
141 lines
5.8 KiB
Python
from pydantic import BaseModel
|
|
|
|
from tht.vectorstore.store import VectorStore
|
|
|
|
|
|
class SearchResult(BaseModel):
|
|
key: str # column:<t>.<c> | table:<t> | evidence:<id>
|
|
label: str
|
|
kind: str # schema_column | schema_table | evidence | values
|
|
signals: dict # {"lsh": {"rank","score","value"?}, "vector": {"rank","score"}}
|
|
rrf: float
|
|
status: str = "" # status evidence, se applicabile
|
|
content: str = "" # testo matchato (per --explain)
|
|
|
|
|
|
def rrf_fuse(rankings: dict[str, list[tuple[str, float]]], k: int) -> dict[str, dict]:
|
|
"""rankings: nome_segnale -> lista (key, raw_score) gia' ordinata per rilevanza.
|
|
Ritorna key -> {"rrf": float, "signals": {segnale: {"rank", "score"}}}."""
|
|
fused: dict[str, dict] = {}
|
|
for signal, ranked in rankings.items():
|
|
for rank, (key, score) in enumerate(ranked, start=1):
|
|
entry = fused.setdefault(key, {"rrf": 0.0, "signals": {}})
|
|
entry["rrf"] += 1.0 / (k + rank)
|
|
entry["signals"][signal] = {"rank": rank, "score": round(score, 4)}
|
|
return fused
|
|
|
|
|
|
def _aggregate_lsh(lsh_hits: list[tuple[str, str, str, float]]) -> list[tuple[str, float, str]]:
|
|
"""Aggrega i match LSH per tabella.colonna tenendo il migliore: (key, score, value)."""
|
|
best: dict[str, tuple[float, str]] = {}
|
|
for table, column, value, score in lsh_hits:
|
|
key = f"column:{table}.{column}"
|
|
if key not in best or score > best[key][0]:
|
|
best[key] = (score, value)
|
|
ordered = sorted(best.items(), key=lambda kv: kv[1][0], reverse=True)
|
|
return [(key, score, value) for key, (score, value) in ordered]
|
|
|
|
|
|
def aggregate_lsh_multi(hits: list[dict]) -> dict[str, list[dict]]:
|
|
"""Value grounding (spec D14a): group LSH hits by table, keeping EVERY column
|
|
where the value appears -- NOT collapsed to a single best column.
|
|
|
|
The old _aggregate_lsh collapsed matches to one column per table.column key,
|
|
hiding alternative groundings (e.g. 'ablazione' matching both a boolean flag
|
|
and a free-text patologia field). This function exposes all of them so the
|
|
value-grounding widget can let the reviewer choose which column(s) anchor a
|
|
cited value.
|
|
|
|
hits: list of {table, column, value, score}.
|
|
Returns: {table -> [{column, value, score}, ...]}, each table's columns ordered
|
|
by score desc; within one (table, column) the best-scored value is kept.
|
|
"""
|
|
best: dict[tuple[str, str], dict] = {}
|
|
for h in hits:
|
|
key = (h["table"], h["column"])
|
|
cur = best.get(key)
|
|
if cur is None or h["score"] > cur["score"]:
|
|
best[key] = {"column": h["column"], "value": h["value"], "score": h["score"]}
|
|
grouped: dict[str, list[dict]] = {}
|
|
for (table, _), row in best.items():
|
|
grouped.setdefault(table, []).append(row)
|
|
for rows in grouped.values():
|
|
rows.sort(key=lambda r: r["score"], reverse=True)
|
|
return grouped
|
|
|
|
|
|
def _vector_key(hit) -> str:
|
|
if hit.kind == "schema_column":
|
|
return f"column:{hit.ref}"
|
|
if hit.kind == "schema_table":
|
|
return f"table:{hit.ref}"
|
|
return f"evidence:{hit.id}"
|
|
|
|
|
|
def schema_tables(results: list["SearchResult"], top_tables: int) -> list[tuple[str, float]]:
|
|
"""Aggrega i risultati schema a livello di tabella per lo schema-linking: ogni chunk
|
|
(tabella o colonna) contribuisce alla sua tabella tenendo il miglior RRF. Ritorna le
|
|
prime top_tables tabelle, ordinate per RRF desc (poi nome). I chunk non-schema sono
|
|
ignorati."""
|
|
best: dict[str, float] = {}
|
|
for r in results:
|
|
if r.kind not in ("schema_table", "schema_column"):
|
|
continue
|
|
table = r.key.split(":", 1)[1].split(".", 1)[0]
|
|
if table not in best or r.rrf > best[table]:
|
|
best[table] = r.rrf
|
|
ordered = sorted(best.items(), key=lambda kv: (-kv[1], kv[0]))
|
|
return ordered[:top_tables]
|
|
|
|
|
|
def combined_search(
|
|
keyword: str,
|
|
*,
|
|
lsh_hits: list[tuple[str, str, str, float]] | None,
|
|
store: VectorStore,
|
|
embedder,
|
|
top: int,
|
|
rrf_k: int,
|
|
kinds: list[str] | None,
|
|
query_vec: list[float] | None = None,
|
|
) -> list[SearchResult]:
|
|
"""Fonde LSH (valori di campo) e ricerca semantica con Reciprocal Rank Fusion.
|
|
|
|
`query_vec` permette di riusare un embedding gia' calcolato della stessa
|
|
keyword (es. `tht search pack`, che fa piu' ricerche sulla stessa domanda)."""
|
|
rankings: dict[str, list[tuple[str, float]]] = {}
|
|
lsh_values: dict[str, str] = {}
|
|
if lsh_hits:
|
|
aggregated = _aggregate_lsh(lsh_hits)
|
|
rankings["lsh"] = [(key, score) for key, score, _ in aggregated]
|
|
lsh_values = {key: value for key, _, value in aggregated}
|
|
|
|
if query_vec is None:
|
|
query_vec = embedder.embed_query(keyword)
|
|
search_kwargs = {"top_n": top * 2, "kinds": kinds}
|
|
if kinds is not None and "evidence" in kinds:
|
|
search_kwargs["query_text"] = keyword
|
|
vector_hits = store.search(query_vec, **search_kwargs)
|
|
rankings["vector"] = [(_vector_key(h), h.similarity) for h in vector_hits]
|
|
by_key = {_vector_key(h): h for h in vector_hits}
|
|
|
|
fused = rrf_fuse(rankings, k=rrf_k)
|
|
results: list[SearchResult] = []
|
|
for key, data in fused.items():
|
|
hit = by_key.get(key)
|
|
if "lsh" in data["signals"] and key in lsh_values:
|
|
data["signals"]["lsh"]["value"] = lsh_values[key]
|
|
results.append(
|
|
SearchResult(
|
|
key=key,
|
|
label=hit.title if hit else key.removeprefix("column:"),
|
|
kind=hit.kind if hit else "values",
|
|
signals=data["signals"],
|
|
rrf=data["rrf"],
|
|
status=(hit.metadata.get("status", "") if hit else ""),
|
|
content=(hit.content if hit else lsh_values.get(key, "")),
|
|
)
|
|
)
|
|
results.sort(key=lambda r: r.rrf, reverse=True)
|
|
return results[:top]
|