Enhance chat memory scoring and token handling by adding expanded query tokens and refining ranking weights.

This commit is contained in:
2026-06-15 22:54:10 +02:00
parent ebbc19c01d
commit 4b0ed88176

View File

@@ -316,6 +316,22 @@ def _query_tokens(prompt: str, limit: int = 16) -> List[str]:
return tokens
def _expanded_query_tokens(tokens: Sequence[str], limit: int = 24) -> List[str]:
expanded: List[str] = []
seen = set()
for token in tokens:
values = [token, *_QUERY_EXPANSIONS.get(token, [])]
for value in values:
cleaned = str(value or "").casefold().strip()
if not cleaned or cleaned in seen:
continue
seen.add(cleaned)
expanded.append(cleaned)
if len(expanded) >= limit:
return expanded
return expanded
def _fts_match_query(tokens: Sequence[str]) -> str:
safe_terms = []
for token in tokens:
@@ -325,6 +341,14 @@ def _fts_match_query(tokens: Sequence[str]) -> str:
return " OR ".join(safe_terms)
def _minimum_memory_score(token_count: int) -> float:
if token_count <= 1:
return 0.85
if token_count == 2:
return 0.9
return 0.65
def _row_mapping(row: Any) -> Dict[str, Any]:
if isinstance(row, dict):
return row
@@ -334,7 +358,13 @@ def _row_mapping(row: Any) -> Dict[str, Any]:
return dict(row)
def _score_hit(row: Dict[str, Any], tokens: Sequence[str], fts_rank: Optional[float] = None) -> float:
def _score_hit(
row: Dict[str, Any],
tokens: Sequence[str],
fts_rank: Optional[float] = None,
*,
query_token_count: Optional[int] = None,
) -> float:
title = str(row.get("session_name") or "").casefold()
user_text = str(row.get("user_content") or "").casefold()
assistant_text = str(row.get("assistant_content") or "").casefold()
@@ -343,19 +373,20 @@ def _score_hit(row: Dict[str, Any], tokens: Sequence[str], fts_rank: Optional[fl
hits = [token for token in tokens if token in combined]
if not hits:
return 0.0
if len(tokens) >= 5 and len(hits) < 2:
core_count = max(1, int(query_token_count or len(tokens) or 1))
if core_count >= 2 and len(hits) < 2:
return 0.0
score = len(hits) / max(1, min(len(tokens), 8))
score = len(hits) / max(1, min(core_count, 4))
score += 0.18 * sum(1 for token in hits if token in title)
score += 0.10 * sum(1 for token in hits if token in user_text)
if fts_rank is not None:
try:
score += min(0.25, 1.0 / (1.0 + abs(float(fts_rank))))
score += min(0.08, 1.0 / (1.0 + abs(float(fts_rank))))
except Exception:
pass
try:
assistant_row = int(row.get("assistant_message_row_id") or 0)
score += min(0.12, assistant_row / 1_000_000)
score += min(0.04, assistant_row / 2_000_000)
except Exception:
pass
return score
@@ -367,6 +398,7 @@ def _load_fts_candidates(
*,
exclude_session_id: Optional[str],
limit: int,
query_token_count: int,
) -> List[Dict[str, Any]]:
match = _fts_match_query(tokens)
if not match:
@@ -396,7 +428,7 @@ def _load_fts_candidates(
candidates: List[Dict[str, Any]] = []
for row in rows:
item = _row_mapping(row)
item["_score"] = _score_hit(item, tokens, item.get("fts_rank"))
item["_score"] = _score_hit(item, tokens, item.get("fts_rank"), query_token_count=query_token_count)
if item["_score"] > 0:
candidates.append(item)
return candidates
@@ -408,6 +440,7 @@ def _load_table_candidates(
*,
exclude_session_id: Optional[str],
limit: int,
query_token_count: int,
) -> List[Dict[str, Any]]:
try:
rows = db.execute(text(
@@ -430,7 +463,7 @@ def _load_table_candidates(
candidates: List[Dict[str, Any]] = []
for row in rows:
item = _row_mapping(row)
item["_score"] = _score_hit(item, tokens)
item["_score"] = _score_hit(item, tokens, query_token_count=query_token_count)
if item["_score"] > 0:
candidates.append(item)
return candidates
@@ -442,6 +475,7 @@ def _load_live_candidates(
*,
exclude_session_id: Optional[str],
limit: int,
query_token_count: int,
) -> List[Dict[str, Any]]:
try:
rows = db.execute(text(
@@ -492,7 +526,7 @@ def _load_live_candidates(
"attachment_summary": attachments,
"created_at": str(item.get("created_at") or ""),
}
turn["_score"] = _score_hit(turn, tokens)
turn["_score"] = _score_hit(turn, tokens, query_token_count=query_token_count)
if turn["_score"] > 0:
turns.append(turn)
return turns