diff --git a/backend/alembic/env.py b/backend/alembic/env.py index 19e19fce6..224e56540 100755 --- a/backend/alembic/env.py +++ b/backend/alembic/env.py @@ -24,14 +24,14 @@ # from apps.system.models.user import SQLModel # noqa # from apps.settings.models.setting_models import SQLModel -#from apps.chat.models.chat_model import SQLModel -from apps.terminology.models.terminology_model import SQLModel -from sqlbot_xpack.custom_prompt.models.custom_prompt_model import SQLModel +from apps.chat.models.chat_model import SQLModel +#from apps.terminology.models.terminology_model import SQLModel +#from sqlbot_xpack.custom_prompt.models.custom_prompt_model import SQLModel #from apps.data_training.models.data_training_model import SQLModel # from apps.dashboard.models.dashboard_model import SQLModel from common.core.config import settings # noqa #from apps.datasource.models.datasource import SQLModel -from apps.system.models.system_model import SQLModel +#from apps.system.models.system_model import SQLModel target_metadata = SQLModel.metadata diff --git a/backend/alembic/versions/072_chat_record.py b/backend/alembic/versions/072_chat_record.py new file mode 100644 index 000000000..f1895bb69 --- /dev/null +++ b/backend/alembic/versions/072_chat_record.py @@ -0,0 +1,31 @@ +"""072_chat_record + +Revision ID: f898322341a5 +Revises: a2e2ecfa5a9c +Create Date: 2026-09-14 15:54:47.921806 + +""" +from alembic import op +import sqlalchemy as sa +import sqlmodel.sql.sqltypes +from sqlalchemy.dialects import postgresql + +# revision identifiers, used by Alembic. +revision = 'f898322341a5' +down_revision = 'a2e2ecfa5a9c' +branch_labels = None +depends_on = None + + +def upgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('chat_record', sa.Column('extracted_keywords', sa.Text(), nullable=True, comment='关键词')) + op.add_column('chat_record', sa.Column('expanded_keywords', sa.Text(), nullable=True, comment='扩展后的关键词')) + # ### end Alembic commands ### + + +def downgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.drop_column('chat_record', 'extracted_keywords') + op.drop_column('chat_record', 'expanded_keywords') + # ### end Alembic commands ### diff --git a/backend/apps/chat/curd/chat.py b/backend/apps/chat/curd/chat.py index 292a3d5cf..f00a296e2 100644 --- a/backend/apps/chat/curd/chat.py +++ b/backend/apps/chat/curd/chat.py @@ -933,6 +933,19 @@ def save_sql_answer(session: SessionDep, record_id: int, answer: str) -> ChatRec return record +def save_extracted_keywords(session: SessionDep, record_id: int, expanded_keywords: str, + extracted_keywords: str = None) -> None: + """保存提取的关键词到 record,供后续独立接口复用。""" + if not record_id: + return + values = {'expanded_keywords': expanded_keywords} + if extracted_keywords: + values['extracted_keywords'] = extracted_keywords + stmt = update(ChatRecord).where(and_(ChatRecord.id == record_id)).values(**values) + session.execute(stmt) + session.commit() + + def save_analysis_answer(session: SessionDep, record_id: int, answer: str = '') -> ChatRecord: if not record_id: raise Exception("Record id cannot be None") diff --git a/backend/apps/chat/models/chat_model.py b/backend/apps/chat/models/chat_model.py index fe6b4d0a4..d8005e266 100644 --- a/backend/apps/chat/models/chat_model.py +++ b/backend/apps/chat/models/chat_model.py @@ -47,6 +47,7 @@ class OperationEnum(Enum): FILTER_CUSTOM_PROMPT = '11' EXECUTE_SQL = '12' GENERATE_PICTURE = '13' + EXTRACT_KEYWORDS = '14' class ChatFinishStep(Enum): @@ -129,6 +130,8 @@ class ChatRecord(SQLModel, table=True): analysis_record_id: int = Field(sa_column=Column(BigInteger, nullable=True)) predict_record_id: int = Field(sa_column=Column(BigInteger, nullable=True)) regenerate_record_id: int = Field(sa_column=Column(BigInteger, nullable=True)) + extracted_keywords: str = Field(sa_column=Column(Text, nullable=True)) + expanded_keywords: str = Field(sa_column=Column(Text, nullable=True)) class ChatRecordResult(BaseModel): @@ -238,6 +241,8 @@ class AiModelQuestion(BaseModel): question: str = None ai_modal_id: int = None ai_modal_name: str = None # Specific model name + extracted_keywords: str = "" # LLM 提取的核心业务关键词(逗号分隔,未经扩展) + expanded_keywords: str = "" # 经术语同义词扩展后的关键词(用于表匹配) engine: str = "" db_schema: str = "" sql: str = "" @@ -255,6 +260,10 @@ class AiModelQuestion(BaseModel): sample_data: str = "" sqlbot_name: str = "SQLBot" + def extract_keywords_sys_prompt(self) -> str: + return get_sql_template().get("extract_keywords").format(lang=self.lang, + sqlbot_name=self.sqlbot_name) + def sql_sys_question(self, db_type: Union[str, DB], enable_query_limit: bool = True): templates: dict[str, str] = {} _sql_template = get_sql_example_template(db_type) diff --git a/backend/apps/chat/task/llm.py b/backend/apps/chat/task/llm.py index cd6fc2036..c39666248 100644 --- a/backend/apps/chat/task/llm.py +++ b/backend/apps/chat/task/llm.py @@ -2,6 +2,7 @@ import concurrent import json import os +import re import traceback import urllib.parse import warnings @@ -30,7 +31,7 @@ from apps.chat.curd.chat import save_question, save_sql_answer, save_sql, \ save_error_message, save_sql_exec_data, save_chart_answer, save_chart, \ finish_record, save_analysis_answer, save_predict_answer, save_predict_data, \ - save_select_datasource_answer, save_recommend_question_answer, \ + save_select_datasource_answer, save_recommend_question_answer, save_extracted_keywords, \ get_old_questions, save_analysis_predict_record, rename_chat, get_chart_config, \ get_chat_chart_data, list_generate_sql_logs, list_generate_chart_logs, start_log, end_log, \ get_last_execute_sql_error, format_json_data, format_chart_fields, get_chat_brief_generate, get_chat_predict_data, \ @@ -48,6 +49,7 @@ from apps.system.crud.parameter_manage import get_groups from apps.system.crud.user import user_ws_list from apps.system.schemas.system_schema import AssistantOutDsSchema +from apps.terminology.curd.terminology import expand_with_terminology from apps.terminology.curd.terminology import get_terminology_template from common.core.config import settings from common.core.db import engine @@ -378,14 +380,16 @@ def filter_terminology_template(self, _session: Session, oid: int = None, ds_id: calculate_oid = self.current_assistant.oid if self.current_assistant.type != 4 else self.oid if self.current_assistant.type == 1: calculate_ds_id = None + # 使用原始提取的关键词(未扩展)进行术语匹配,避免同义词重复查询 + match_text = self.chat_question.extracted_keywords or self.chat_question.question if self.current_assistant and self.current_assistant.type == 1: self.chat_question.terminologies, term_list = get_terminology_template(_session, - self.chat_question.question, + match_text, calculate_oid, None, self.current_assistant.id) else: self.chat_question.terminologies, term_list = get_terminology_template(_session, - self.chat_question.question, + match_text, calculate_oid, calculate_ds_id) @@ -447,17 +451,108 @@ def filter_training_template(self, _session: Session, oid: int = None, ds_id: in OperationEnum.FILTER_SQL_EXAMPLE], full_message=example_list) + def extract_keywords(self, _session: Session) -> tuple[bool, str]: + """从用户问题中提取核心业务实体词。 + + 使用 LLM 提取名词性业务实体,忽略时间词、动作词、修饰词。 + 支持上下文历史,可根据对话上下文理解用户意图。 + 提取结果可复用于表匹配、术语扩展等环节。 + + Returns: + (is_error, result) 元组: + - (False, keywords) 成功,关键词逗号分隔 + - (True, error_message) 错误(LLM 失败或非查数据意图) + """ + system_prompt = self.chat_question.extract_keywords_sys_prompt() + + keywords_msg: List[Union[BaseMessage, dict[str, Any]]] = [] + keywords_msg.append(SystemPromptMessage(content=system_prompt)) + + # 加载上下文历史(仅取用户的历史提问,不含 AI 生成的 SQL 等回复) + last_sql_messages: List[dict[str, Any]] = self.generate_sql_logs[-1].messages if len( + self.generate_sql_logs) > 0 else [] + if self.chat_question.regenerate_record_id: + _temp_log = next( + filter(lambda obj: obj.pid == self.chat_question.regenerate_record_id, self.generate_sql_logs), None) + last_sql_messages: List[dict[str, Any]] = _temp_log.messages if _temp_log else [] + + # 排除系统提示词和 AI 回复,只保留用户的历史提问 + last_user_messages = [ + obj for obj in last_sql_messages + if obj.get("sqlbot_system") != True and obj.get('type') == 'human' + ] + + if last_user_messages: + last_rounds = get_last_conversation_rounds(last_user_messages, rounds=self.base_message_round_count_limit) + for _msg_dict in last_rounds: + content = _msg_dict.get('content', '') + # 提取 标签内的用户原始提问,排除 error-msg 等干扰信息 + match = re.search(r'(.*?)', content, re.DOTALL) + if match: + question_text = match.group(1).strip() + if question_text: + keywords_msg.append(HumanMessage(content=question_text)) + + # 当前问题 + keywords_msg.append(HumanMessage(content=self.chat_question.question)) + + self.current_logs[OperationEnum.EXTRACT_KEYWORDS] = start_log( + session=_session, + ai_modal_id=self.chat_question.ai_modal_id, + ai_modal_name=self.chat_question.ai_modal_name, + operate=OperationEnum.EXTRACT_KEYWORDS, + record_id=self.record.id, + full_message=[{'type': msg.type, 'content': msg.content} for msg in keywords_msg] + ) + + full_thinking_text = '' + full_text = '' + token_usage = {} + res = process_stream(self.llm.stream(keywords_msg), token_usage) + for chunk in res: + if chunk.get('content'): + full_text += chunk.get('content') + if chunk.get('reasoning_content'): + full_thinking_text += chunk.get('reasoning_content') + + result_text = full_text.strip() + + keywords_msg.append(AIMessage(result_text)) + self.current_logs[OperationEnum.EXTRACT_KEYWORDS] = end_log( + session=_session, + log=self.current_logs[OperationEnum.EXTRACT_KEYWORDS], + full_message=[{'type': msg.type, 'content': msg.content} for msg in keywords_msg], + reasoning_content=full_thinking_text if full_thinking_text else None, + token_usage=token_usage + ) + + if result_text.startswith('ERROR:NOT_DATA_QUERY'): + # return True, '您的问题似乎不是查询数据的问题,请提出与数据查询相关的问题。' + return False, '' + + if result_text == 'EMPTY' or not result_text: + self.chat_question.extracted_keywords = self.chat_question.question + return False, '' + + self.chat_question.extracted_keywords = result_text + return False, result_text + def choose_table_schema(self, _session: Session): self.current_logs[OperationEnum.CHOOSE_TABLE] = start_log(session=_session, operate=OperationEnum.CHOOSE_TABLE, record_id=self.record.id, local_operation=True) + + # 使用扩展后的关键词进行表匹配(包含同义词,匹配更全面) + keywords = self.chat_question.expanded_keywords or self.chat_question.question + self.chat_question.db_schema, tables = self.out_ds_instance.get_db_schema( self.ds.id, self.chat_question.question) if self.out_ds_instance else get_table_schema( session=_session, current_user=self.current_user, ds=self.ds, - question=self.chat_question.question) + question=self.chat_question.question, + keywords=keywords) # Get sample data for all tables if not self.out_ds_instance: @@ -479,6 +574,12 @@ def generate_analysis(self, _session: Session): self.chat_question.data = orjson.dumps(data.get('data')).decode() analysis_msg: List[Union[BaseMessage, dict[str, Any]]] = [] + # 从 record 加载已提取的关键词(SQL 生成阶段已持久化) + if self.record.extracted_keywords: + self.chat_question.extracted_keywords = self.record.extracted_keywords + if self.record.expanded_keywords: + self.chat_question.expanded_keywords = self.record.expanded_keywords + ds_id = self.ds.id if isinstance(self.ds, CoreDatasource) else None self.filter_terminology_template(_session, self.oid, ds_id) @@ -781,6 +882,22 @@ def select_datasource(self, _session: Session): oid = self.ds.oid if isinstance(self.ds, CoreDatasource) else 1 ds_id = self.ds.id if isinstance(self.ds, CoreDatasource) else None + # 提取关键词(前置步骤,结果复用于术语匹配和表匹配) + if settings.TABLE_EMBEDDING_KEYWORD_ENABLED: + is_error, kw_result = self.extract_keywords(_session) + if is_error: + raise SingleMessageError(kw_result) + if kw_result: + self.chat_question.extracted_keywords = kw_result + SQLBotLogUtil.info(f"提取的关键词: {kw_result}") + # 术语同义词扩展(用于表匹配,不用于术语模板查询) + expanded = expand_with_terminology(kw_result, _session, oid) + self.chat_question.expanded_keywords = expanded + SQLBotLogUtil.info(f"扩展后的关键词: {expanded}") + # 持久化到 record,供后续独立接口(如 generate_analysis)使用 + save_extracted_keywords(_session, self.record.id, + expanded_keywords=expanded, extracted_keywords=kw_result) + self.filter_terminology_template(_session, oid, ds_id) self.filter_training_template(_session, oid, ds_id) @@ -1247,6 +1364,22 @@ def run_task(self, in_chat: bool = True, stream: bool = True, oid = self.ds.oid if isinstance(self.ds, CoreDatasource) else 1 ds_id = self.ds.id if isinstance(self.ds, CoreDatasource) else None + # 提取关键词(前置步骤,结果复用于术语匹配和表匹配) + if settings.TABLE_EMBEDDING_KEYWORD_ENABLED: + is_error, kw_result = self.extract_keywords(_session) + if is_error: + raise SingleMessageError(kw_result) + if kw_result: + self.chat_question.extracted_keywords = kw_result + SQLBotLogUtil.info(f"提取的关键词: {kw_result}") + # 术语同义词扩展(用于表匹配,不用于术语模板查询) + expanded = expand_with_terminology(kw_result, _session, oid) + self.chat_question.expanded_keywords = expanded + SQLBotLogUtil.info(f"扩展后的关键词: {expanded}") + # 持久化到 record,供后续独立接口(如 generate_analysis)使用 + save_extracted_keywords(_session, self.record.id, + expanded_keywords=expanded, extracted_keywords=kw_result) + self.filter_terminology_template(_session, oid, ds_id) self.filter_training_template(_session, oid, ds_id) diff --git a/backend/apps/datasource/crud/datasource.py b/backend/apps/datasource/crud/datasource.py index 855881e6e..9e727a3bb 100644 --- a/backend/apps/datasource/crud/datasource.py +++ b/backend/apps/datasource/crud/datasource.py @@ -517,7 +517,8 @@ def get_tables_sample_data(session: SessionDep, current_user: CurrentUser, ds: C def get_table_schema(session: SessionDep, current_user: CurrentUser, ds: CoreDatasource, question: str, - embedding: bool = True, table_list: list[str] = None) -> tuple[str, list]: + embedding: bool = True, table_list: list[str] = None, + keywords: str = None) -> tuple[str, list]: schema_str = "" table_objs = get_table_obj_by_ds(session=session, current_user=current_user, ds=ds) if len(table_objs) == 0: @@ -538,17 +539,25 @@ def get_table_schema(session: SessionDep, current_user: CurrentUser, ds: CoreDat table_comment = '' if obj.table.custom_comment: table_comment = obj.table.custom_comment.strip() + if not table_comment and obj.table.table_comment: + table_comment = obj.table.table_comment.strip() if table_comment == '': schema_table += '\n[\n' else: schema_table += f", {table_comment}\n[\n" + fields_info = [] if obj.fields: field_list = [] for field in obj.fields: field_comment = '' if field.custom_comment: field_comment = field.custom_comment.strip() + fields_info.append({ + "field_name": field.field_name, + "field_comment": (field.field_comment or '').strip(), + "custom_comment": (field.custom_comment or '').strip() + }) if field_comment == '': field_list.append(f"({field.field_name}:{field.field_type})") else: @@ -557,7 +566,7 @@ def get_table_schema(session: SessionDep, current_user: CurrentUser, ds: CoreDat schema_table += '\n]\n' t_obj = {"id": obj.table.id, "table_name": obj.table.table_name, "schema_table": schema_table, - "embedding": obj.table.embedding} + "embedding": obj.table.embedding, "table_comment": table_comment, "fields": fields_info} tables.append(t_obj) all_tables.append(t_obj) @@ -565,9 +574,14 @@ def get_table_schema(session: SessionDep, current_user: CurrentUser, ds: CoreDat if not tables: return schema_str, [] - # do table embedding + # 执行表 embedding 匹配 if embedding and tables and settings.TABLE_EMBEDDING_ENABLED: - tables = calc_table_embedding(tables, question) + try: + tables = calc_table_embedding(tables, question, session=session, oid=current_user.oid, + keywords=keywords) + except ValueError as e: + # LLM 判断意图不是查数据或关键词提取失败 + return str(e), [] # splice schema if tables: for s in tables: diff --git a/backend/apps/datasource/embedding/table_embedding.py b/backend/apps/datasource/embedding/table_embedding.py index 186debec4..e1fbfa286 100644 --- a/backend/apps/datasource/embedding/table_embedding.py +++ b/backend/apps/datasource/embedding/table_embedding.py @@ -1,6 +1,7 @@ # Author: Junjun # Date: 2025/9/23 import json +import re import time import traceback @@ -10,6 +11,15 @@ from common.utils.utils import SQLBotLogUtil +def _parse_embedding(embedding): + """解析 embedding,兼容 JSON 字符串和原生列表两种格式。""" + if isinstance(embedding, list): + return embedding + if isinstance(embedding, str): + return json.loads(embedding) + return None + + def get_table_embedding(tables: list[dict], question: str): _list = [] for table in tables: @@ -40,38 +50,244 @@ def get_table_embedding(tables: list[dict], question: str): return _list -def calc_table_embedding(tables: list[dict], question: str): +def calc_keyword_score( + question: str, + table_name: str, + table_comment: str, + fields: list[dict] +) -> float: + """计算关键词与表的匹配分数。 + + 综合考虑所有关键词的匹配情况: + - 表名完整匹配:1.0 + - 表注释完整匹配:0.9 + - 部分匹配时,按匹配质量 * 关键词覆盖率计算 + + Args: + question: 提取后的关键词(逗号分隔)或原始问题 + table_name: 表名 + table_comment: 表注释(已合并 custom_comment 优先) + fields: 字段列表,每个 dict 包含 field_name, field_comment, custom_comment + + Returns: + 0.0~1.0 的匹配分数 + """ + if not question: + return 0.0 + + keywords = [kw.strip().lower() for kw in question.split(',') if kw.strip()] + if not keywords: + return 0.0 + + tn = table_name.lower() + tc = (table_comment or '').lower() + total_kw = len(keywords) + + # 收集每个关键词的最佳匹配分数 + kw_scores = {} # keyword -> best_score + + for kw in keywords: + # 1. 表名完整匹配:1.0 + if kw == tn: + kw_scores[kw] = 1.0 + continue + + # 2. 表注释完整匹配:0.9 + if tc and kw == tc: + kw_scores[kw] = max(kw_scores.get(kw, 0), 0.9) + continue + + # 3. 表注释子串匹配:0.7(关键词是注释的子串,如 "用户" 匹配 "用户表") + if tc and kw in tc: + kw_scores[kw] = max(kw_scores.get(kw, 0), 0.7) + continue + + # 4. 表名部分匹配:收集匹配的关键词 + tn_parts = set(re.split(r'[_\s]+', tn)) + tn_parts.discard('') + if kw in tn_parts: + kw_scores[kw] = max(kw_scores.get(kw, 0), 0) # 标记匹配,分数稍后计算 + + # 4. 字段注释匹配:0.5(精确)/ 0.4(子串) + for field in fields: + field_comment = '' + if field.get('custom_comment'): + field_comment = field['custom_comment'].strip().lower() + if not field_comment and field.get('field_comment'): + field_comment = field['field_comment'].strip().lower() + if field_comment: + if kw == field_comment: + kw_scores[kw] = max(kw_scores.get(kw, 0), 0.5) + break + elif kw in field_comment: + kw_scores[kw] = max(kw_scores.get(kw, 0), 0.4) + + # 5. 字段名匹配:0.3 + if kw not in kw_scores or kw_scores[kw] < 0.3: + for field in fields: + fname = field.get('field_name', '').lower() + if fname: + fname_parts = set(re.split(r'[_\s]+', fname)) + fname_parts.discard('') + if kw in fname_parts: + kw_scores[kw] = max(kw_scores.get(kw, 0), 0.3) + break + + if not kw_scores: + return 0.0 + + # 表名/注释完整匹配 → 直接返回 1.0 或 0.9 + if kw_scores.get(tn) == 1.0: + return 1.0 + if tc and kw_scores.get(tc) == 0.9: + return 0.9 + + # 表名部分匹配:按覆盖率计算(多个关键词共同覆盖表名) + tn_parts = set(re.split(r'[_\s]+', tn)) + tn_parts.discard('') + if tn_parts: + matched_parts = {kw for kw in keywords if kw in tn_parts} + if matched_parts: + coverage = len(matched_parts) / len(tn_parts) + tn_score = round(0.25 + 0.25 * coverage, 2) + # 更新这些关键词的分数为表名匹配分数 + for kw in matched_parts: + kw_scores[kw] = max(kw_scores.get(kw, 0), tn_score) + + # 综合评分:最佳匹配分数 * 关键词覆盖率 + best_score = max(kw_scores.values()) + match_rate = len(kw_scores) / total_kw + return round(best_score * match_rate, 2) + + +def calc_table_embedding(tables: list[dict], question: str, session=None, oid: int = None, + keywords: str = None): + """使用向量相似度和关键词匹配的融合评分计算表的相关性分数。 + + 当 keywords 不为空时(已由调用方提取并扩展): + 1. 用 keywords 计算向量相似度(vec_score) + 2. 计算关键词匹配度(keyword_score) + 3. 融合评分:final_score = α * vec_score + (1-α) * keyword_score + 4. keyword_score=1.0 的表保证排在最前面 + + 当 keywords 为空或 TABLE_EMBEDDING_KEYWORD_ENABLED 为 False 时,回退到纯向量匹配。 + + Args: + tables: 表列表,每个 dict 包含 id, table_name, schema_table, embedding, table_comment, fields + question: 原始用户问题(用于纯向量匹配回退) + session: 数据库会话(预留) + oid: 组织 ID(预留) + keywords: 已提取并扩展的关键词(由调用方传入) + """ + if not keywords or not settings.TABLE_EMBEDDING_KEYWORD_ENABLED: + # 回退到纯向量匹配 + return _calc_vector_only(tables, question) + + _list = [] + for table in tables: + _list.append({ + "id": table.get('id'), + "schema_table": table.get('schema_table'), + "embedding": table.get('embedding'), + "cosine_similarity": 0.0, + "table_name": table.get('table_name'), + "table_comment": table.get('table_comment', ''), + "fields": table.get('fields', []) + }) + + if not _list: + return _list + + try: + start_time = time.time() + + # 步骤 1:用关键词计算向量相似度 + model = EmbeddingModelCache.get_model() + results = [item.get('embedding') for item in _list] + + q_embedding = model.embed_query(keywords) + for index in range(len(results)): + item = results[index] + if item: + _list[index]['cosine_similarity'] = cosine_similarity(q_embedding, _parse_embedding(item)) + + # 步骤 2 & 3:计算关键词分数并融合 + alpha = settings.TABLE_EMBEDDING_ALPHA + for table in _list: + vec_score = table['cosine_similarity'] + kw_score = calc_keyword_score( + keywords, + table.get('table_name', ''), + table.get('table_comment', ''), + table.get('fields', []) + ) + table['keyword_score'] = kw_score + table['cosine_similarity'] = alpha * vec_score + (1 - alpha) * kw_score + + # 步骤 4:排序 - 精确匹配(keyword_score=1.0)始终排最前 + exact_matches = [t for t in _list if t.get('keyword_score') == 1.0] + other_tables = [t for t in _list if t.get('keyword_score') != 1.0] + + exact_matches.sort(key=lambda x: x['cosine_similarity'], reverse=True) + other_tables.sort(key=lambda x: x['cosine_similarity'], reverse=True) + + _list = exact_matches + other_tables + _list = _list[:settings.TABLE_EMBEDDING_COUNT] + + end_time = time.time() + SQLBotLogUtil.info(f"融合评分耗时 {end_time - start_time:.3f}s") + SQLBotLogUtil.info(json.dumps([{ + "id": ele.get('id'), + "schema_table": ele.get('schema_table'), + "cosine_similarity": ele.get('cosine_similarity'), + "keyword_score": ele.get('keyword_score'), + "table_name": ele.get('table_name') + } for ele in _list])) + + return _list + except Exception: + traceback.print_exc() + # 异常回退到纯向量匹配 + return _calc_vector_only(tables, question) + + +def _calc_vector_only(tables: list[dict], question: str): + """纯向量相似度评分(原始逻辑)。""" _list = [] for table in tables: - _list.append( - {"id": table.get('id'), "schema_table": table.get('schema_table'), "embedding": table.get('embedding'), - "cosine_similarity": 0.0, "table_name": table.get('table_name')}) + _list.append({ + "id": table.get('id'), + "schema_table": table.get('schema_table'), + "embedding": table.get('embedding'), + "cosine_similarity": 0.0, + "table_name": table.get('table_name'), + "table_comment": table.get('table_comment', ''), + "fields": table.get('fields', []) + }) if _list: try: - # text = [s.get('schema_table') for s in _list] - # model = EmbeddingModelCache.get_model() start_time = time.time() - # results = model.embed_documents(text) - # end_time = time.time() - # SQLBotLogUtil.info(str(end_time - start_time)) results = [item.get('embedding') for item in _list] q_embedding = model.embed_query(question) for index in range(len(results)): item = results[index] if item: - _list[index]['cosine_similarity'] = cosine_similarity(q_embedding, json.loads(item)) + _list[index]['cosine_similarity'] = cosine_similarity(q_embedding, _parse_embedding(item)) _list.sort(key=lambda x: x['cosine_similarity'], reverse=True) _list = _list[:settings.TABLE_EMBEDDING_COUNT] - # print(len(_list)) + end_time = time.time() SQLBotLogUtil.info(str(end_time - start_time)) - SQLBotLogUtil.info(json.dumps([{"id": ele.get('id'), "schema_table": ele.get('schema_table'), - "cosine_similarity": ele.get('cosine_similarity'), "table_name": ele.get('table_name')} - for ele in _list])) + SQLBotLogUtil.info(json.dumps([{ + "id": ele.get('id'), + "schema_table": ele.get('schema_table'), + "cosine_similarity": ele.get('cosine_similarity'), + "table_name": ele.get('table_name') + } for ele in _list])) return _list except Exception: traceback.print_exc() diff --git a/backend/apps/terminology/curd/terminology.py b/backend/apps/terminology/curd/terminology.py index c69aa77d8..a3864bd72 100644 --- a/backend/apps/terminology/curd/terminology.py +++ b/backend/apps/terminology/curd/terminology.py @@ -977,10 +977,122 @@ def get_terminology_template(session: SessionDep, question: str, oid: Optional[i advanced_application_id: Optional[int] = None) -> tuple[str, list[dict]]: if not oid: oid = 1 - _results = select_terminology_by_word(session, question, oid, datasource, advanced_application_id) - if _results and len(_results) > 0: + + # 有逗号时视为关键词列表,逐个匹配(精度更高,避免整体 ILIKE 误匹配) + # 无逗号时视为自然语言问题,整体匹配 + if ',' in question: + parts = [kw.strip() for kw in question.split(',') if kw.strip()] + else: + parts = [question] + + # select_terminology_by_word 返回 list[dict],每个 dict 结构: + # {'words': ['销售', 'sales'], 'description': '...'} + seen_words = set() + _results = [] + for part in parts: + for item in select_terminology_by_word(session, part, oid, datasource, advanced_application_id): + item_words = set(item.get('words', [])) if isinstance(item, dict) else set() + if item_words and not item_words.issubset(seen_words): + seen_words.update(item_words) + _results.append(item) + + if _results: terminology = to_xml_string(_results) template = get_base_terminology_template().format(terminologies=terminology) return template, _results else: return '', [] + + +def expand_with_terminology(keywords: str, session: SessionDep, oid: int) -> str: + """使用术语同义词扩展关键词。 + + 查询当前租户启用的术语表,将匹配的关键词扩展为同义词集合。 + pid=None 表示父术语(主词),pid= 表示子术语(同义词)。 + + Args: + keywords: 逗号分隔的关键词字符串 + session: 数据库会话 + oid: 组织 ID + + Returns: + 扩展后的关键词字符串(逗号分隔),最多 10 个关键词 + """ + if not keywords or not session or oid is None: + return keywords + + keyword_list = [kw.strip() for kw in keywords.split(',') if kw.strip()] + if not keyword_list: + return keywords + + try: + # 查询当前租户启用的术语(限制 500 条避免内存问题) + terms = session.query(Terminology).filter( + Terminology.oid == oid, + Terminology.enabled == True + ).limit(500).all() + + if not terms: + return keywords + + # 构建父子结构 + parent_map = {} # id -> Terminology(仅父术语) + children_map = {} # parent_id -> [Terminology](子术语列表) + word_to_parent_ids = {} # 小写词 -> 父术语 id 集合 + + for term in terms: + if term.pid is None or term.pid == 0: + parent_map[term.id] = term + if term.word: + w = term.word.strip().lower() + if w not in word_to_parent_ids: + word_to_parent_ids[w] = set() + word_to_parent_ids[w].add(term.id) + else: + if term.pid not in children_map: + children_map[term.pid] = [] + children_map[term.pid].append(term) + if term.word: + w = term.word.strip().lower() + if w not in word_to_parent_ids: + word_to_parent_ids[w] = set() + word_to_parent_ids[w].add(term.pid) + + # 扩展关键词 + expanded = list(keyword_list) + seen = set(kw.lower() for kw in keyword_list) + max_keywords = 10 + + for kw in keyword_list: + if len(expanded) >= max_keywords: + break + + kw_lower = kw.lower() + matched_parent_ids = word_to_parent_ids.get(kw_lower, set()) + + for pid in matched_parent_ids: + if len(expanded) >= max_keywords: + break + + # 添加父术语 + parent = parent_map.get(pid) + if parent and parent.word: + parent_word = parent.word.strip() + if parent_word.lower() not in seen: + expanded.append(parent_word) + seen.add(parent_word.lower()) + + # 添加子术语(同义词) + for child in children_map.get(pid, []): + if len(expanded) >= max_keywords: + break + if child.word: + child_word = child.word.strip() + if child_word.lower() not in seen: + expanded.append(child_word) + seen.add(child_word.lower()) + + return ','.join(expanded) + except Exception: + traceback.print_exc() + return keywords diff --git a/backend/common/core/config.py b/backend/common/core/config.py index 4b9baeaec..cb9f9b2dd 100644 --- a/backend/common/core/config.py +++ b/backend/common/core/config.py @@ -128,6 +128,8 @@ def SQLALCHEMY_DATABASE_URI(self) -> PostgresDsn | str: TABLE_EMBEDDING_ENABLED: bool = True TABLE_EMBEDDING_COUNT: int = 10 + TABLE_EMBEDDING_ALPHA: float = 0.4 # weight for vector score; (1-alpha) is keyword weight + TABLE_EMBEDDING_KEYWORD_ENABLED: bool = True DS_EMBEDDING_COUNT: int = 10 ORACLE_CLIENT_PATH: str = '/opt/sqlbot/db_client/oracle_instant_client' @@ -138,6 +140,7 @@ def SQLALCHEMY_DATABASE_URI(self) -> PostgresDsn | str: 'PARSE_REASONING_BLOCK_ENABLED', 'PG_POOL_PRE_PING', 'TABLE_EMBEDDING_ENABLED', + 'TABLE_EMBEDDING_KEYWORD_ENABLED', mode='before') @classmethod def lowercase_bool(cls, v: Any) -> Any: @@ -150,5 +153,15 @@ def lowercase_bool(cls, v: Any) -> Any: return False return v + @field_validator('TABLE_EMBEDDING_ALPHA', mode='before') + @classmethod + def clamp_alpha(cls, v: Any) -> float: + """将 TABLE_EMBEDDING_ALPHA 限制在 [0.0, 1.0] 范围内""" + try: + v = float(v) + except (TypeError, ValueError): + return 0.4 + return max(0.0, min(1.0, v)) + settings = Settings() # type: ignore diff --git a/backend/templates/template.yaml b/backend/templates/template.yaml index 2bad9fd55..79f36a828 100644 --- a/backend/templates/template.yaml +++ b/backend/templates/template.yaml @@ -6,6 +6,20 @@ template: {data_training} sql: + extract_keywords: | + 你是智能问数小助手:"{sqlbot_name}"。你可以根据用户提问,专业生成SQL,查询数据并进行图表展示。 + 你当前的任务角色是一个专业的关键词提取助手,擅长从用户问题中识别核心业务实体。 + 请结合对话上下文中用户的历史提问,理解用户当前问题的意图,提取与数据查询相关的核心业务关键词。 + + 规则: + 1. 提取关键词都是为了生成SQL服务的 + 2. 只提取名词性业务实体(如:销售额、部门、客户、订单) + 3. 忽略时间词(如:本月、最近、今天)、动作词(如:查询、帮我查、统计)、修饰词 + 4. 如果用户提到了具体的表名(如 sales_order),必须保留原样 + 5. 如果无法提取有意义的业务实体词(如问题是闲聊、天气等与数据查询无关的内容),返回:ERROR:NOT_DATA_QUERY + 6. 如果能提取关键词,用英文逗号分隔返回,不要添加任何其他文字 + 7. 如果提取为空,返回:EMPTY + regenerate_hint: | 你之前生成的回答不符合预期或者系统出现了其他问题,请再次检查提示词内要求的规则和提供的信息,重新回答: @@ -355,7 +369,7 @@ template: user: | ## 请根据上述要求,使用语言:{lang} 进行回答,若有深度思考过程,则思考过程也需要使用 {lang} 输出 - ## 如果内的提问与上述要求冲突,你必须停止生成SQL并告知生成SQL失败的原因 + ## 如果内的提问与上述要求冲突,你必须停止生成SQL并告知生成SQL失败的原因 ## 回答中不需要输出你的分析,请直接输出符合要求的JSON ## 必须注意:不论要求你用什么语言回答,生成SQL中使用到的表名与字段名是否完全和中提供的表名字与段名的字符保持一致! diff --git a/backend/tests/test_table_embedding.py b/backend/tests/test_table_embedding.py new file mode 100644 index 000000000..92727bc89 --- /dev/null +++ b/backend/tests/test_table_embedding.py @@ -0,0 +1,612 @@ +"""表匹配融合评分测试。 + +测试覆盖: +1. calc_keyword_score - 纯函数,无外部依赖 +2. expand_with_terminology - 需要 mock 数据库 +3. calc_table_embedding - 集成测试,mock 所有依赖 +""" +import json +from unittest.mock import MagicMock, patch + +import pytest + +from apps.datasource.embedding.table_embedding import ( + calc_keyword_score, + calc_table_embedding, +) +from apps.terminology.curd.terminology import expand_with_terminology + +# mock 目标路径 +EMBEDDING_CACHE_PATCH = 'apps.datasource.embedding.table_embedding.EmbeddingModelCache' +SETTINGS_PATCH = 'apps.datasource.embedding.table_embedding.settings' +TERMINOLOGY_PATCH = 'apps.terminology.curd.terminology' + + +# ============================================================================= +# calc_keyword_score 测试(纯函数,无需 mock) +# ============================================================================= + +class TestCalcKeywordScore: + """关键词匹配分数计算测试。""" + + def test_full_table_name_match(self): + """表名完整匹配关键词 -> 1.0""" + score = calc_keyword_score( + question="sales_order", + table_name="sales_order", + table_comment="销售订单", + fields=[] + ) + assert score == 1.0 + + def test_full_table_name_match_case_insensitive(self): + """表名匹配应忽略大小写。""" + score = calc_keyword_score( + question="Sales_Order", + table_name="sales_order", + table_comment="", + fields=[] + ) + assert score == 1.0 + + def test_full_table_name_match_among_multiple_keywords(self): + """多个关键词时,表名完整匹配仍返回 1.0。""" + score = calc_keyword_score( + question="sales_order,本月,销售额", + table_name="sales_order", + table_comment="", + fields=[] + ) + assert score == 1.0 + + def test_table_comment_full_match(self): + """表注释完整匹配关键词 -> 0.9""" + score = calc_keyword_score( + question="销售订单", + table_name="so_header", + table_comment="销售订单", + fields=[] + ) + assert score == 0.9 + + def test_table_comment_full_match_custom_comment_priority(self): + """自定义备注应作为 table_comment 使用(调用方负责合并)。""" + # 调用方传入已合并的注释,这里模拟该场景 + score = calc_keyword_score( + question="销售订单", + table_name="so_header", + table_comment="自定义备注", + fields=[] + ) + # "销售订单" != "自定义备注",不匹配注释 + assert score == 0.0 + + def test_partial_table_name_match_single_keyword(self): + """关键词匹配表名部分 -> 0.25~0.5(按覆盖比例)。""" + # "sales" 匹配 "sales_order" 的一部分(2 个部分:sales, order) + # 覆盖率 = 1/2 = 0.5 -> 0.25 + 0.25*0.5 = 0.375 + score = calc_keyword_score( + question="sales", + table_name="sales_order", + table_comment="", + fields=[] + ) + assert score == 0.38 # 0.25 + 0.25 * 0.5 = 0.375,四舍五入为 0.38 + + def test_partial_table_name_match_full_coverage(self): + """表名所有部分都被关键词覆盖 -> 0.5。""" + # "sales" 和 "order" 覆盖了 "sales_order" 的所有部分 + score = calc_keyword_score( + question="sales,order", + table_name="sales_order", + table_comment="", + fields=[] + ) + assert score == 0.5 + + def test_field_comment_match(self): + """关键词匹配字段注释 -> 0.5""" + fields = [ + {"field_name": "amount", "field_comment": "销售额", "custom_comment": ""}, + {"field_name": "qty", "field_comment": "数量", "custom_comment": ""}, + ] + score = calc_keyword_score( + question="销售额", + table_name="order_detail", + table_comment="", + fields=fields + ) + assert score == 0.5 + + def test_field_custom_comment_priority(self): + """自定义注释应优先于字段注释。""" + fields = [ + {"field_name": "amount", "field_comment": "原始注释", "custom_comment": "销售额"}, + ] + score = calc_keyword_score( + question="销售额", + table_name="order_detail", + table_comment="", + fields=fields + ) + assert score == 0.5 + + def test_field_name_match(self): + """关键词匹配字段名部分 -> 0.3""" + fields = [ + {"field_name": "order_amount", "field_comment": "", "custom_comment": ""}, + ] + score = calc_keyword_score( + question="amount", + table_name="transactions", + table_comment="", + fields=fields + ) + assert score == 0.3 + + def test_field_name_match_no_duplicate(self): + """同一词根匹配多个字段时只计算一次。""" + fields = [ + {"field_name": "order_amount", "field_comment": "", "custom_comment": ""}, + {"field_name": "total_amount", "field_comment": "", "custom_comment": ""}, + ] + # "amount" 匹配了两个字段,但仍应返回 0.3(不会更高) + score = calc_keyword_score( + question="amount", + table_name="transactions", + table_comment="", + fields=fields + ) + assert score == 0.3 + + def test_no_match(self): + """无匹配 -> 0.0""" + fields = [ + {"field_name": "id", "field_comment": "主键", "custom_comment": ""}, + ] + score = calc_keyword_score( + question="销售额", + table_name="users", + table_comment="用户表", + fields=fields + ) + assert score == 0.0 + + def test_empty_question(self): + """空问题 -> 0.0""" + score = calc_keyword_score( + question="", + table_name="sales_order", + table_comment="", + fields=[] + ) + assert score == 0.0 + + def test_empty_keywords_after_split(self): + """问题只有逗号/空格 -> 0.0""" + score = calc_keyword_score( + question=",, ,", + table_name="sales_order", + table_comment="", + fields=[] + ) + assert score == 0.0 + + def test_priority_table_name_over_comment(self): + """表名匹配应优先于注释匹配。""" + score = calc_keyword_score( + question="sales_order,销售订单", + table_name="sales_order", + table_comment="销售订单", + fields=[] + ) + assert score == 1.0 + + def test_priority_comment_over_field_comment(self): + """注释匹配(0.9)应优先于字段注释匹配(0.5)。""" + fields = [ + {"field_name": "amt", "field_comment": "金额", "custom_comment": ""}, + ] + score = calc_keyword_score( + question="销售订单", + table_name="so_detail", + table_comment="销售订单", + fields=fields + ) + assert score == 0.9 + + def test_priority_field_comment_over_field_name(self): + """字段注释匹配(0.5)应优先于字段名匹配(0.3)。""" + fields = [ + {"field_name": "销售额_column", "field_comment": "", "custom_comment": ""}, + {"field_name": "amt", "field_comment": "销售额", "custom_comment": ""}, + ] + # 字段注释匹配 "销售额" -> 0.5 + score = calc_keyword_score( + question="销售额", + table_name="detail", + table_comment="", + fields=fields + ) + assert score == 0.5 + + def test_multiple_keywords_best_score_wins(self): + """多个关键词时,返回最佳匹配分数。""" + fields = [ + {"field_name": "qty", "field_comment": "数量", "custom_comment": ""}, + ] + # "sales_order" 匹配表名 -> 1.0 + # "数量" 匹配字段注释 -> 0.5 + # 最佳为 1.0 + score = calc_keyword_score( + question="sales_order,数量", + table_name="sales_order", + table_comment="", + fields=fields + ) + assert score == 1.0 + + def test_underscore_table_name_parts(self): + """带下划线的表名应正确分割。""" + # "sales_order_detail" -> 部分: {"sales", "order", "detail"} + # "sales" 覆盖 1/3 -> 0.25 + 0.25*(1/3) ≈ 0.33 + score = calc_keyword_score( + question="sales", + table_name="sales_order_detail", + table_comment="", + fields=[] + ) + assert 0.3 <= score <= 0.35 + + def test_field_name_with_underscores(self): + """带下划线的字段名应正确分割。""" + fields = [ + {"field_name": "customer_order_count", "field_comment": "", "custom_comment": ""}, + ] + score = calc_keyword_score( + question="order", + table_name="stats", + table_comment="", + fields=fields + ) + assert score == 0.3 + + def test_real_scenario_sales_query(self): + """真实场景:'查询销售额' 应匹配包含 '销售额' 字段的表。""" + fields = [ + {"field_name": "id", "field_comment": "主键", "custom_comment": ""}, + {"field_name": "order_amount", "field_comment": "订单金额", "custom_comment": ""}, + {"field_name": "sales_amount", "field_comment": "销售额", "custom_comment": ""}, + ] + score = calc_keyword_score( + question="销售额", + table_name="daily_report", + table_comment="日报表", + fields=fields + ) + assert score == 0.5 # 字段注释匹配 + + def test_real_scenario_table_name_query(self): + """真实场景:'查 sales_order 表最近的订单' -> 提取 'sales_order'。""" + score = calc_keyword_score( + question="sales_order,订单", + table_name="sales_order", + table_comment="销售订单", + fields=[] + ) + assert score == 1.0 # 表名完整匹配 + + +# ============================================================================= +# expand_with_terminology 测试(mock 数据库) +# ============================================================================= + +class TestExpandWithTerminology: + """术语同义词扩展测试。""" + + def _make_term(self, id, word, pid=None, enabled=True, oid=1): + """创建 mock 术语对象的辅助方法。""" + term = MagicMock() + term.id = id + term.word = word + term.pid = pid + term.enabled = enabled + term.oid = oid + return term + + def test_expand_hit(self): + """匹配术语的关键词应扩展同义词。""" + session = MagicMock() + session.query.return_value.filter.return_value.limit.return_value.all.return_value = [ + self._make_term(1, "销售", pid=None), + self._make_term(2, "sales", pid=1), + self._make_term(3, "selling", pid=1), + ] + + result = expand_with_terminology("销售", session, oid=1) + keywords = [kw.strip() for kw in result.split(',')] + assert "销售" in keywords + assert "sales" in keywords + assert "selling" in keywords + + def test_no_hit(self): + """未匹配任何术语的关键词应保持不变。""" + session = MagicMock() + session.query.return_value.filter.return_value.limit.return_value.all.return_value = [ + self._make_term(1, "客户", pid=None), + ] + + result = expand_with_terminology("销售额", session, oid=1) + assert result == "销售额" + + def test_max_10_keywords(self): + """扩展后的关键词总数不应超过 10 个。""" + session = MagicMock() + # 创建一个有多个同义词的术语 + terms = [self._make_term(1, "销售", pid=None)] + for i in range(15): + terms.append(self._make_term(100 + i, f"synonym_{i}", pid=1)) + session.query.return_value.filter.return_value.limit.return_value.all.return_value = terms + + result = expand_with_terminology("销售", session, oid=1) + keywords = [kw.strip() for kw in result.split(',') if kw.strip()] + assert len(keywords) <= 10 + + def test_empty_terminology(self): + """空术语表应返回原始关键词。""" + session = MagicMock() + session.query.return_value.filter.return_value.limit.return_value.all.return_value = [] + + result = expand_with_terminology("销售额,部门", session, oid=1) + assert result == "销售额,部门" + + def test_none_session(self): + """session 为 None 时应返回原始关键词。""" + result = expand_with_terminology("销售额", None, oid=1) + assert result == "销售额" + + def test_none_oid(self): + """oid 为 None 时应返回原始关键词。""" + session = MagicMock() + result = expand_with_terminology("销售额", session, oid=None) + assert result == "销售额" + + def test_oid_zero_should_still_query(self): + """oid=0 不应被视为假值 — 它是有效的组织 ID。""" + session = MagicMock() + session.query.return_value.filter.return_value.limit.return_value.all.return_value = [ + self._make_term(1, "销售", pid=None, oid=0), + self._make_term(2, "revenue", pid=1, oid=0), + ] + result = expand_with_terminology("销售", session, oid=0) + keywords = [kw.strip() for kw in result.split(',')] + assert "销售" in keywords + assert "revenue" in keywords + + def test_empty_keywords(self): + """空关键词字符串应原样返回。""" + session = MagicMock() + result = expand_with_terminology("", session, oid=1) + assert result == "" + + def test_case_insensitive_match(self): + """术语匹配应忽略大小写。""" + session = MagicMock() + session.query.return_value.filter.return_value.limit.return_value.all.return_value = [ + self._make_term(1, "Sales", pid=None), + self._make_term(2, "revenue", pid=1), + ] + + result = expand_with_terminology("sales", session, oid=1) + keywords = [kw.strip() for kw in result.split(',')] + # "sales"(输入)匹配 "Sales"(术语),忽略大小写 + # "Sales" 不会被添加,因为 "sales" 已在列表中(避免重复) + # "revenue"(子术语同义词)被添加 + assert "sales" in keywords + assert "revenue" in keywords + assert len(keywords) == 2 + + def test_multiple_keywords_expansion(self): + """多个关键词应各自独立扩展。""" + session = MagicMock() + session.query.return_value.filter.return_value.limit.return_value.all.return_value = [ + self._make_term(1, "销售", pid=None), + self._make_term(2, "sales", pid=1), + self._make_term(3, "部门", pid=None), + self._make_term(4, "department", pid=3), + ] + + result = expand_with_terminology("销售,部门", session, oid=1) + keywords = [kw.strip() for kw in result.split(',')] + assert "sales" in keywords + assert "department" in keywords + + def test_child_term_word_also_matches(self): + """如果关键词匹配子术语的词,应扩展父术语和兄弟术语。""" + session = MagicMock() + session.query.return_value.filter.return_value.limit.return_value.all.return_value = [ + self._make_term(1, "销售", pid=None), + self._make_term(2, "sales", pid=1), + self._make_term(3, "selling", pid=1), + ] + + # "sales" 匹配子术语,应扩展包含父术语 "销售" 和兄弟术语 "selling" + result = expand_with_terminology("sales", session, oid=1) + keywords = [kw.strip() for kw in result.split(',')] + assert "sales" in keywords + assert "销售" in keywords + assert "selling" in keywords + + +# ============================================================================= +# calc_table_embedding 集成测试(mock embedding) +# ============================================================================= + +class TestCalcTableEmbedding: + """融合评分管线集成测试。""" + + def _make_table(self, id, table_name, schema_table, embedding, table_comment="", fields=None): + return { + "id": id, + "table_name": table_name, + "schema_table": schema_table, + "embedding": json.dumps(embedding), + "table_comment": table_comment, + "fields": fields or [] + } + + @patch('apps.datasource.embedding.table_embedding.settings') + @patch('apps.datasource.embedding.table_embedding.EmbeddingModelCache') + def test_exact_table_name_match_ranks_first(self, mock_embed_cache, mock_settings): + """表名完整匹配关键词时,应排第一。""" + mock_settings.TABLE_EMBEDDING_KEYWORD_ENABLED = True + mock_settings.TABLE_EMBEDDING_ALPHA = 0.4 + mock_settings.TABLE_EMBEDDING_COUNT = 10 + + mock_model = MagicMock() + mock_model.embed_query.return_value = [1.0, 0.0, 0.0] + mock_embed_cache.get_model.return_value = mock_model + + tables = [ + self._make_table(1, "customer", "# Table: customer", [0.0, 1.0, 0.0], "客户表", + [{"field_name": "name", "field_comment": "姓名", "custom_comment": ""}]), + self._make_table(2, "product", "# Table: product", [0.0, 0.0, 1.0], "产品表", + [{"field_name": "name", "field_comment": "产品名", "custom_comment": ""}]), + self._make_table(3, "sales_order", "# Table: sales_order", [0.9, 0.1, 0.0], "销售订单", + [{"field_name": "amount", "field_comment": "金额", "custom_comment": ""}]), + ] + + # 直接传入已提取的关键词 + result = calc_table_embedding(tables, "查 sales_order 表最近的订单", + keywords="sales_order,订单") + + # sales_order 应排第一,因为 keyword_score=1.0 + assert result[0]['table_name'] == "sales_order" + assert result[0].get('keyword_score') == 1.0 + + @patch('apps.datasource.embedding.table_embedding.settings') + @patch('apps.datasource.embedding.table_embedding.EmbeddingModelCache') + def test_keyword_disabled_uses_vector_only(self, mock_embed_cache, mock_settings): + """关键词匹配关闭时,应使用纯向量匹配。""" + mock_settings.TABLE_EMBEDDING_KEYWORD_ENABLED = False + mock_settings.TABLE_EMBEDDING_COUNT = 10 + + mock_model = MagicMock() + mock_model.embed_query.return_value = [1.0, 0.0, 0.0] + mock_embed_cache.get_model.return_value = mock_model + + tables = [ + self._make_table(1, "customer", "# Table: customer", [0.0, 1.0, 0.0]), + self._make_table(2, "sales_order", "# Table: sales_order", [0.9, 0.1, 0.0]), + ] + + result = calc_table_embedding(tables, "查 sales_order 表", keywords="sales_order,订单") + + # 关键词匹配关闭,即使传入了 keywords 也应使用纯向量匹配 + assert result[0]['table_name'] == "sales_order" + + @patch('apps.datasource.embedding.table_embedding.settings') + @patch('apps.datasource.embedding.table_embedding.EmbeddingModelCache') + def test_no_keywords_falls_back_to_vector(self, mock_embed_cache, mock_settings): + """未传入 keywords 时,应回退到纯向量匹配。""" + mock_settings.TABLE_EMBEDDING_KEYWORD_ENABLED = True + mock_settings.TABLE_EMBEDDING_COUNT = 10 + + mock_model = MagicMock() + mock_model.embed_query.return_value = [1.0, 0.0, 0.0] + mock_embed_cache.get_model.return_value = mock_model + + tables = [ + self._make_table(1, "customer", "# Table: customer", [0.0, 1.0, 0.0]), + self._make_table(2, "sales_order", "# Table: sales_order", [0.9, 0.1, 0.0]), + ] + + # 不传 keywords + result = calc_table_embedding(tables, "查 sales_order 表") + + # 应回退到纯向量匹配 + assert result[0]['table_name'] == "sales_order" + + @patch('apps.datasource.embedding.table_embedding.settings') + @patch('apps.datasource.embedding.table_embedding.EmbeddingModelCache') + def test_multiple_exact_matches_all_rank_first(self, mock_embed_cache, mock_settings): + """多个 keyword_score=1.0 的表都应排在其他表之前。""" + mock_settings.TABLE_EMBEDDING_KEYWORD_ENABLED = True + mock_settings.TABLE_EMBEDDING_ALPHA = 0.4 + mock_settings.TABLE_EMBEDDING_COUNT = 10 + + mock_model = MagicMock() + mock_model.embed_query.return_value = [1.0, 0.0, 0.0] + mock_embed_cache.get_model.return_value = mock_model + + tables = [ + self._make_table(1, "product", "# Table: product", [0.0, 0.0, 1.0], "产品", + [{"field_name": "name", "field_comment": "名称", "custom_comment": ""}]), + self._make_table(2, "sales_order", "# Table: sales_order", [0.9, 0.1, 0.0], "", + [{"field_name": "amount", "field_comment": "", "custom_comment": ""}]), + self._make_table(3, "customer", "# Table: customer", [0.8, 0.2, 0.0], "", + [{"field_name": "name", "field_comment": "", "custom_comment": ""}]), + ] + + result = calc_table_embedding(tables, "查 sales_order 和 customer", + keywords="sales_order,customer") + + # 前两个应是精确匹配(它们之间的顺序无所谓) + exact_match_names = {result[0]['table_name'], result[1]['table_name']} + assert exact_match_names == {"sales_order", "customer"} + # product 应排最后 + assert result[2]['table_name'] == "product" + + @patch('apps.datasource.embedding.table_embedding.settings') + @patch('apps.datasource.embedding.table_embedding.EmbeddingModelCache') + def test_fusion_score_formula(self, mock_embed_cache, mock_settings): + """验证融合公式:final = α * vec + (1-α) * keyword。""" + mock_settings.TABLE_EMBEDDING_KEYWORD_ENABLED = True + mock_settings.TABLE_EMBEDDING_ALPHA = 0.4 + mock_settings.TABLE_EMBEDDING_COUNT = 10 + + mock_model = MagicMock() + mock_model.embed_query.return_value = [1.0, 0.0] + mock_embed_cache.get_model.return_value = mock_model + + # 表的 vec_score=0.8(通过余弦相似度),keyword_score 来自部分匹配 + tables = [ + self._make_table(1, "sales_order", "# Table: sales_order", [0.8, 0.6], "", + [{"field_name": "id", "field_comment": "", "custom_comment": ""}]), + ] + + result = calc_table_embedding(tables, "sales", keywords="sales") + + # "sales" 部分匹配 "sales_order"(2 个部分中的 1 个) + # keyword_score = 0.25 + 0.25 * 0.5 = 0.375 + # vec_score = cosine([1,0], [0.8, 0.6]) = 0.8 / (1 * 1) = 0.8 + # final = 0.4 * 0.8 + 0.6 * 0.375 = 0.32 + 0.225 = 0.545 + assert len(result) == 1 + score = result[0]['cosine_similarity'] + assert 0.5 <= score <= 0.6 # 允许小的浮点误差 + + @patch('apps.datasource.embedding.table_embedding.settings') + @patch('apps.datasource.embedding.table_embedding.EmbeddingModelCache') + def test_fallback_on_unexpected_error(self, mock_embed_cache, mock_settings): + """遇到意外错误时,应回退到纯向量评分。""" + mock_settings.TABLE_EMBEDDING_KEYWORD_ENABLED = True + mock_settings.TABLE_EMBEDDING_ALPHA = 0.4 + mock_settings.TABLE_EMBEDDING_COUNT = 10 + + # 让 embed_query 抛出意外错误 + mock_model = MagicMock() + mock_model.embed_query.side_effect = RuntimeError("模型加载失败") + mock_embed_cache.get_model.return_value = mock_model + + tables = [ + self._make_table(1, "t1", "# Table: t1", [0.8, 0.6]), + self._make_table(2, "t2", "# Table: t2", [0.3, 0.7]), + ] + + # 不应抛出异常,应回退到纯向量匹配 + result = calc_table_embedding(tables, "test", keywords="test") + assert len(result) == 2 + + +if __name__ == '__main__': + pytest.main([__file__, '-v']) diff --git a/frontend/src/i18n/en.json b/frontend/src/i18n/en.json index 430fa369c..1816cf4c5 100644 --- a/frontend/src/i18n/en.json +++ b/frontend/src/i18n/en.json @@ -809,6 +809,7 @@ "GENERATE_DYNAMIC_SQL": "Generate Dynamic SQL", "CHOOSE_TABLE": "Match Data Table (Schema)", "FILTER_TERMS": "Match Terms", + "EXTRACT_KEYWORDS": "Extract Keywords", "FILTER_SQL_EXAMPLE": "Match SQL Examples", "FILTER_CUSTOM_PROMPT": "Match Custom Prompts", "EXECUTE_SQL": "Execute SQL", diff --git a/frontend/src/i18n/ko-KR.json b/frontend/src/i18n/ko-KR.json index 850102f9b..58686d303 100644 --- a/frontend/src/i18n/ko-KR.json +++ b/frontend/src/i18n/ko-KR.json @@ -809,6 +809,7 @@ "GENERATE_DYNAMIC_SQL": "동적 SQL 생성", "CHOOSE_TABLE": "데이터 테이블 매칭 (스키마)", "FILTER_TERMS": "용어 매칭", + "EXTRACT_KEYWORDS": "키워드 추출", "FILTER_SQL_EXAMPLE": "SQL 예시 매칭", "FILTER_CUSTOM_PROMPT": "사용자 정의 프롬프트 매칭", "EXECUTE_SQL": "SQL 실행", diff --git a/frontend/src/i18n/zh-CN.json b/frontend/src/i18n/zh-CN.json index 30cea8f2d..16e05da63 100644 --- a/frontend/src/i18n/zh-CN.json +++ b/frontend/src/i18n/zh-CN.json @@ -809,6 +809,7 @@ "GENERATE_DYNAMIC_SQL": "生成动态 SQL", "CHOOSE_TABLE": "匹配数据表 (Schema)", "FILTER_TERMS": "匹配术语", + "EXTRACT_KEYWORDS": "提取关键词", "FILTER_SQL_EXAMPLE": "匹配 SQL 示例", "FILTER_CUSTOM_PROMPT": "匹配自定义提示词", "EXECUTE_SQL": "执行 SQL", diff --git a/frontend/src/i18n/zh-TW.json b/frontend/src/i18n/zh-TW.json index 295d2d016..8acb5198f 100644 --- a/frontend/src/i18n/zh-TW.json +++ b/frontend/src/i18n/zh-TW.json @@ -809,6 +809,7 @@ "GENERATE_DYNAMIC_SQL": "產生動態 SQL", "CHOOSE_TABLE": "匹配資料表 (Schema)", "FILTER_TERMS": "匹配術語", + "EXTRACT_KEYWORDS": "提取關鍵詞", "FILTER_SQL_EXAMPLE": "匹配 SQL 範例", "FILTER_CUSTOM_PROMPT": "匹配自訂提示詞", "EXECUTE_SQL": "執行 SQL",