From c51cf5e25cf7911cff3b8ae2670a09c3ef85afd2 Mon Sep 17 00:00:00 2001 From: huaiyj <8699003+huaiyj@user.noreply.gitee.com> Date: Mon, 24 Aug 2026 20:53:07 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E4=BC=98=E5=8C=96=E4=BC=9A=E8=AF=9D?= =?UTF-8?q?=E6=90=9C=E7=B4=A2=E7=9B=B8=E5=85=B3=E6=80=A7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../ocai/repo/mysql/v2_session_repo.py | 96 +++++++++++++++---- ocai-service/ocai/services/chat_service.py | 6 +- ocai-service/tests/repo/test_mysql_repos.py | 38 ++++++++ 3 files changed, 118 insertions(+), 22 deletions(-) diff --git a/ocai-service/ocai/repo/mysql/v2_session_repo.py b/ocai-service/ocai/repo/mysql/v2_session_repo.py index 0887c6b..2aec619 100644 --- a/ocai-service/ocai/repo/mysql/v2_session_repo.py +++ b/ocai-service/ocai/repo/mysql/v2_session_repo.py @@ -296,35 +296,97 @@ class MysqlV2SessionRepository(AbstractV2SessionRepository): # ---- Search --------------------------------------------------- def search_sessions(self, user_id: str, keyword: str, limit: int = 30) -> list[dict]: - like = f"%{keyword}%" - # Match either the session summary or any of its messages' content. - message_session_ids = ( - select(V2SessionMessage.session_id).where(V2SessionMessage.content.like(like)).distinct().scalar_subquery() + keyword = keyword.strip() + if not keyword: + return [] + escaped = keyword.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + like = f"%{escaped}%" + content_match = ( + V2SessionMessage.content.match(keyword) + if self._engine.engine.dialect.name == "mysql" + else V2SessionMessage.content.like(like, escape="\\") ) + # Match either the session summary or any of its messages' content. + message_session_ids = select(V2SessionMessage.session_id).where(content_match).distinct().scalar_subquery() stmt = ( select(V2Session) .where( V2Session.user_id == user_id, V2Session.status == 1, or_( - V2Session.summary.like(like), + V2Session.summary.like(like, escape="\\"), V2Session.session_id.in_(message_session_ids), ), ) .order_by(desc(V2Session.updated_at)) - .limit(limit) + .limit(min(max(limit * 5, limit), 300)) ) with self._engine.session_scope() as s: rows = s.execute(stmt).scalars().all() - return [ - { - "session_id": r.session_id, - "summary": r.summary or "", - "created_at": r.created_at, - "updated_at": r.updated_at, - } - for r in rows - ] + if not rows: + return [] + + session_ids = [row.session_id for row in rows] + matched_messages = ( + s.execute( + select(V2SessionMessage) + .where( + V2SessionMessage.session_id.in_(session_ids), + content_match, + ) + .order_by(V2SessionMessage.session_id, V2SessionMessage.seq) + ) + .scalars() + .all() + ) + matches_by_session: dict[str, list[V2SessionMessage]] = {} + for message in matched_messages: + matches_by_session.setdefault(message.session_id, []).append(message) + + results = [] + lowered_keyword = keyword.casefold() + for row in rows: + matches = matches_by_session.get(row.session_id, []) + summary_hit = lowered_keyword in (row.summary or "").casefold() + first = matches[0] if matches else None + score = (4 if summary_hit else 0) + min(len(matches), 5) * 2 + if any(message.seq == 0 and message.role == "user" for message in matches): + score += 3 + results.append( + { + "session_id": row.session_id, + "summary": row.summary or "", + "created_at": row.created_at, + "updated_at": row.updated_at, + "relevance_score": score, + "match_count": len(matches) + int(summary_hit), + "message_id": first.message_id if first else None, + "snippet": ( + self._search_snippet(first.content, keyword) + if first + else self._search_snippet(row.summary, keyword) + ), + } + ) + + results.sort( + key=lambda item: (item["relevance_score"], item["updated_at"]), + reverse=True, + ) + return results[:limit] + + @staticmethod + def _search_snippet(content: str | None, keyword: str, radius: int = 80) -> str: + if not content: + return "" + index = content.casefold().find(keyword.casefold()) + if index < 0: + return content[: radius * 2] + start = max(0, index - radius) + end = min(len(content), index + len(keyword) + radius) + prefix = "…" if start else "" + suffix = "…" if end < len(content) else "" + return f"{prefix}{content[start:end]}{suffix}" def find_session_by_confirm_id(self, confirm_id: str) -> str | None: if not confirm_id: @@ -395,9 +457,7 @@ class MysqlV2SessionRepository(AbstractV2SessionRepository): @staticmethod def _lock_session(s, session_id: str) -> None: s.execute( - select(V2Session.session_id) - .where(V2Session.session_id == session_id) - .with_for_update() + select(V2Session.session_id).where(V2Session.session_id == session_id).with_for_update() ).scalar_one_or_none() @staticmethod diff --git a/ocai-service/ocai/services/chat_service.py b/ocai-service/ocai/services/chat_service.py index 0d9d45f..1f8d9a2 100644 --- a/ocai-service/ocai/services/chat_service.py +++ b/ocai-service/ocai/services/chat_service.py @@ -17,7 +17,6 @@ from __future__ import annotations import logging import queue as _queue -import re import threading import time import uuid @@ -289,10 +288,9 @@ class ChatService(MessagePersistenceMixin): def search_sessions(self, user_id: str, keyword: str, limit: int = 30) -> list[dict]: """根据关键字搜索用户的会话 - 对关键字进行正则转义防止注入。 + 仓储层使用参数化查询并转义 LIKE 通配符。 """ - safe_keyword = re.escape(keyword) - sessions = self._v2_repo.search_sessions(user_id, safe_keyword, limit=limit) + sessions = self._v2_repo.search_sessions(user_id, keyword, limit=limit) for s in sessions: for key in ("created_at", "updated_at"): if key in s and hasattr(s[key], "isoformat"): diff --git a/ocai-service/tests/repo/test_mysql_repos.py b/ocai-service/tests/repo/test_mysql_repos.py index 2f9b6f7..8d177a0 100644 --- a/ocai-service/tests/repo/test_mysql_repos.py +++ b/ocai-service/tests/repo/test_mysql_repos.py @@ -107,6 +107,44 @@ def test_v2_session_search_and_soft_delete(engine): assert repo.get_session("s1", "alice") is None +def test_v2_session_search_returns_ranked_snippets(engine): + repo = MysqlV2SessionRepository(engine) + repo.create_session("older-relevant", "alice", "Alice") + repo.append_message( + "older-relevant", + {"message_id": "first-hit", "role": "user", "content": "磁盘告警需要排查容量"}, + ) + repo.append_message( + "older-relevant", + {"message_id": "second-hit", "role": "assistant", "content": "磁盘告警来自根分区"}, + ) + repo.create_session("newer-single-hit", "alice", "Alice") + repo.append_message( + "newer-single-hit", + {"message_id": "only-hit", "role": "assistant", "content": "收到磁盘告警"}, + ) + + results = repo.search_sessions("alice", "磁盘告警") + + assert [item["session_id"] for item in results] == ["older-relevant", "newer-single-hit"] + assert results[0]["message_id"] == "first-hit" + assert results[0]["match_count"] == 2 + assert "磁盘告警" in results[0]["snippet"] + assert results[0]["relevance_score"] > results[1]["relevance_score"] + + +def test_v2_session_search_treats_like_wildcards_as_literals(engine): + repo = MysqlV2SessionRepository(engine) + repo.create_session("literal", "alice", "Alice") + repo.append_message("literal", {"role": "user", "content": "CPU 使用率达到 90%"}) + repo.create_session("other", "alice", "Alice") + repo.append_message("other", {"role": "user", "content": "CPU 使用率正常"}) + + results = repo.search_sessions("alice", "90%") + + assert [item["session_id"] for item in results] == ["literal"] + + # --------------------------------------------------------------------------- # File meta + many-to-many session links — Phase 3 起 FileMetaRepo 已下线 # Sub-agent rows / TTL cleanup — Phase 4 起 sub_agent_* 表已下线 -- Gitee