Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 36 additions & 8 deletions hermes_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -1705,6 +1705,24 @@ def _count_cjk(cls, text: str) -> int:
"""Count CJK characters in text."""
return sum(1 for ch in text if cls._is_cjk_codepoint(ord(ch)))

@classmethod
def _cjk_search_terms(cls, query: str) -> List[str]:
"""Return non-operator CJK search terms from a simple FTS query."""
terms: List[str] = []
for tok in query.split():
if tok.upper() in ("AND", "OR", "NOT"):
continue
term = tok.strip('"').strip()
if term:
terms.append(term)
return terms

@classmethod
def _cjk_trigram_can_match(cls, query: str) -> bool:
"""Whether every CJK term is long enough for SQLite trigram FTS."""
terms = cls._cjk_search_terms(query)
return bool(terms) and all(cls._count_cjk(term) >= 3 for term in terms)

def search_messages(
self,
query: str,
Expand Down Expand Up @@ -1787,9 +1805,8 @@ def search_messages(
is_cjk = self._contains_cjk(query)
if is_cjk:
raw_query = query.strip('"').strip()
cjk_count = self._count_cjk(raw_query)

if cjk_count >= 3:
if self._cjk_trigram_can_match(raw_query):
# Trigram FTS5 path — quote each non-operator token to handle
# FTS5 special chars (%, *, etc.) while preserving boolean
# operators (AND, OR, NOT) for multi-term queries.
Expand Down Expand Up @@ -1841,10 +1858,22 @@ def search_messages(
matches = [dict(row) for row in tri_cursor.fetchall()]
else:
# Short CJK query (1-2 chars) — trigram needs ≥3 CJK chars.
# Fall back to LIKE substring search.
escaped = raw_query.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
like_where = ["(m.content LIKE ? ESCAPE '\\' OR m.tool_name LIKE ? ESCAPE '\\' OR m.tool_calls LIKE ? ESCAPE '\\')"]
like_params: list = [f"%{escaped}%", f"%{escaped}%", f"%{escaped}%"]
# Fall back to LIKE substring search. For simple OR/AND CJK
# queries, apply the boolean operator across each short term so
# natural searches such as "广西 OR 桂林" can still match.
terms = self._cjk_search_terms(raw_query) or [raw_query]
operator = " OR " if re.search(r"(?i)\bOR\b", raw_query) else " AND "
like_clauses = []
like_params: list = []
for term in terms:
escaped = term.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
like_clauses.append(
"(m.content LIKE ? ESCAPE '\\' "
"OR m.tool_name LIKE ? ESCAPE '\\' "
"OR m.tool_calls LIKE ? ESCAPE '\\')"
)
like_params.extend([f"%{escaped}%", f"%{escaped}%", f"%{escaped}%"])
like_where = [f"({operator.join(like_clauses)})"]
if source_filter is not None:
like_where.append(f"s.source IN ({','.join('?' for _ in source_filter)})")
like_params.extend(source_filter)
Expand All @@ -1869,7 +1898,7 @@ def search_messages(
"""
like_params.extend([limit, offset])
# instr() parameter goes first in the bound list
like_params = [raw_query] + like_params
like_params = [terms[0]] + like_params
with self._lock:
like_cursor = self._conn.execute(like_sql, like_params)
matches = [dict(row) for row in like_cursor.fetchall()]
Expand Down Expand Up @@ -2666,4 +2695,3 @@ def maybe_auto_prune_and_vacuum(
result["error"] = str(exc)

return result

29 changes: 28 additions & 1 deletion tests/test_hermes_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -957,6 +957,34 @@ def test_cjk_trigram_preserves_boolean_operators(self, db):
session_ids = {r["session_id"] for r in results}
assert session_ids == {"s1", "s2"}

def test_cjk_or_short_terms_uses_like_fallback(self, db):
"""OR-combined 2-char CJK terms must not be routed to trigram FTS."""
db.create_session(session_id="s1", source="telegram")
db.create_session(session_id="s2", source="telegram")
db.create_session(session_id="s3", source="cli")
db.append_message("s1", role="user", content="这次广西旅游安排得很好")
db.append_message("s2", role="user", content="桂林和漓江行程需要补充")
db.append_message("s3", role="user", content="无关内容")

results = db.search_messages("广西 OR 桂林 OR 漓江 OR 旅游")

session_ids = {r["session_id"] for r in results}
assert session_ids == {"s1", "s2"}

def test_cjk_or_short_terms_preserves_source_filter(self, db):
db.create_session(session_id="s1", source="telegram")
db.create_session(session_id="s2", source="cli")
db.append_message("s1", role="user", content="广西旅游")
db.append_message("s2", role="user", content="桂林旅游")

results = db.search_messages(
"广西 OR 桂林",
source_filter=["telegram"],
)

assert len(results) == 1
assert results[0]["session_id"] == "s1"


# =========================================================================
# Session search and listing
Expand Down Expand Up @@ -2909,4 +2937,3 @@ def test_v10_to_v11_upgrade_backfills_tool_fields(self, tmp_path):
assert version == 11
finally:
session_db.close()

Loading