diff --git a/backend/apps/datasource/embedding/table_embedding.py b/backend/apps/datasource/embedding/table_embedding.py index 7aea7272..b4d95ec9 100644 --- a/backend/apps/datasource/embedding/table_embedding.py +++ b/backend/apps/datasource/embedding/table_embedding.py @@ -12,6 +12,69 @@ # 预编译正则,避免每次调用都重新编译 _RE_UNDERSCORE_SPACE = re.compile(r'[_\s]+') +_IDENTIFIER_CHARS = r'A-Za-z0-9_$' +_RE_QUOTED_IDENTIFIER = re.compile(r'"(?:[^"]|"")*"|`(?:[^`]|``)*`|\[(?:[^\]]|\]\])*\]') + + +def _mention_spans(question: str, name: str): + """Match identifier boundaries while allowing adjacent Chinese text and SQL delimiters.""" + left = rf'(? furthest_end: + longest_spans.add((start, end)) + furthest_end = end + return {name for name, occurrences in spans_by_name.items() + if any(span in longest_spans for span in occurrences)} + + +def _select_tables(ranked_tables: list[dict], question: str, *, apply_limit: bool = True): + """Keep explicitly mentioned candidates, then fill remaining slots by existing rank. + + Inspect only candidates already filtered by the caller; retain duplicate comments. + Resolve overlapping names and comments separately. Error fallback can disable truncation. + """ + question = (question or '').lower() + quoted_names = [] + for match in _RE_QUOTED_IDENTIFIER.finditer(question): + token = match.group() + closing = token[-1] + quoted_names.append((match.span(), token[1:-1].replace(closing * 2, closing))) + name_spans, comment_spans = {}, {} + for table in ranked_tables: + name = (table.get('table_name') or '').strip().lower() + if name and name not in name_spans: + # Match quoted identifiers as a whole, even when the full name is not a candidate. + name_spans[name] = [ + (a, b) for a, b in _mention_spans(question, name) + if not any(a < end and start < b for (start, end), _ in quoted_names) + ] + name_spans[name].extend(span for span, quoted_name in quoted_names if quoted_name == name) + comment = (table.get('table_comment') or '').strip().lower() + if comment and comment not in comment_spans: + comment_spans[comment] = _mention_spans(question, comment) + + explicit_names = _longest_mentions(name_spans) + explicit_comments = _longest_mentions(comment_spans) + required, remaining = [], [] + for table in ranked_tables: + name = (table.get('table_name') or '').strip().lower() + comment = (table.get('table_comment') or '').strip().lower() + if name in explicit_names or comment in explicit_comments: + required.append(table) + else: + remaining.append(table) + if apply_limit: + remaining = remaining[:max(0, settings.TABLE_EMBEDDING_COUNT - len(required))] + return required + remaining def _parse_embedding(embedding): @@ -177,13 +240,13 @@ def calc_table_embedding(tables: list[dict], question: str, session=None, oid: i 1. 用 keywords 计算向量相似度(vec_score) 2. 计算关键词匹配度(keyword_score) 3. 融合评分:final_score = α * vec_score + (1-α) * keyword_score - 4. keyword_score=1.0 的表保证排在最前面 + 4. Keep explicitly mentioned table names or complete comments, then fill slots by rank. 当 keywords 为空或 TABLE_EMBEDDING_KEYWORD_ENABLED 为 False 时,回退到纯向量匹配。 Args: tables: 表列表,每个 dict 包含 id, table_name, schema_table, embedding, table_comment, fields - question: 原始用户问题(用于纯向量匹配回退) + question: Original user question for explicit matching and vector-only fallback. session: 数据库会话(预留) oid: 组织 ID(预留) keywords: 已提取并扩展的关键词(由调用方传入) @@ -272,7 +335,7 @@ def calc_table_embedding(tables: list[dict], question: str, session=None, oid: i other_tables.sort(key=lambda x: x['cosine_similarity'], reverse=True) _list = exact_matches + other_tables - _list = _list[:settings.TABLE_EMBEDDING_COUNT] + _list = _select_tables(_list, question) end_time = time.time() SQLBotLogUtil.info(f"[perf] 融合评分总耗时 {end_time - start_time:.3f}s") @@ -292,7 +355,7 @@ def calc_table_embedding(tables: list[dict], question: str, session=None, oid: i def _calc_vector_only(tables: list[dict], question: str): - """纯向量相似度评分(原始逻辑)。""" + """Rank by vector similarity while retaining explicitly mentioned tables.""" _list = [] for table in tables: _list.append({ @@ -329,7 +392,7 @@ def _calc_vector_only(tables: list[dict], question: str): _list[idx]['cosine_similarity'] = float(sim) _list.sort(key=lambda x: x['cosine_similarity'], reverse=True) - _list = _list[:settings.TABLE_EMBEDDING_COUNT] + _list = _select_tables(_list, question) end_time = time.time() SQLBotLogUtil.info(f"[perf] 纯向量匹配耗时 {end_time - start_time:.3f}s,共 {len(parsed_embeddings)} 张表") @@ -342,4 +405,5 @@ def _calc_vector_only(tables: list[dict], question: str): return _list except Exception: traceback.print_exc() - return _list + # Preserve all candidates when vector ranking fails; move explicit matches to the front. + return _select_tables(_list, question, apply_limit=False) diff --git a/backend/tests/test_table_embedding.py b/backend/tests/test_table_embedding.py index 92727bc8..b2b149d0 100644 --- a/backend/tests/test_table_embedding.py +++ b/backend/tests/test_table_embedding.py @@ -608,5 +608,157 @@ def test_fallback_on_unexpected_error(self, mock_embed_cache, mock_settings): assert len(result) == 2 +class TestExplicitTableSelection: + """Retain explicitly named tables and comments even with low vector rankings.""" + + TARGET = 'ads_jt_scm_purchase_process_chain_relation_detail_full_1d' + COMMENT = '采购全链路数据各节点表' + + @pytest.fixture(autouse=True) + def setup_models(self): + with patch(EMBEDDING_CACHE_PATCH) as cache, patch(SETTINGS_PATCH) as config: + config.TABLE_EMBEDDING_KEYWORD_ENABLED = True + config.TABLE_EMBEDDING_COUNT = 10 + config.TABLE_EMBEDDING_ALPHA = 0.4 + cache.get_model.return_value.embed_query.return_value = [1.0, 0.0] + self.cache, self.config = cache, config + yield + + def table(self, name, comment='', embedding=None): + return dict(id=name, table_name=name, table_comment=comment, + schema_table=f'# Table: {name}', fields=[], + embedding=json.dumps(embedding if embedding is not None else [0.95, 0.31225])) + + def candidates(self): + return [self.table(f'{self.TARGET}_{i}', self.COMMENT + f'(环节{i}明细)') + for i in range(19)] + [self.table(self.TARGET, self.COMMENT, [0.6, 0.8])] + + @pytest.mark.parametrize('keywords', [None, '采购,节点,金额', '查询采购金额']) + @pytest.mark.parametrize('mention', ['name', 'comment']) + def test_explicit_mention_survives_top_ten(self, keywords, mention): + question = f'查询{self.TARGET if mention == "name" else self.COMMENT}的采购金额' + result = calc_table_embedding(self.candidates(), question, keywords=keywords) + assert len(result) == 10 + assert result[0]['table_name'] == self.TARGET + + @pytest.mark.parametrize('mention', ['name', 'comment']) + def test_keyword_disabled(self, mention): + self.config.TABLE_EMBEDDING_KEYWORD_ENABLED = False + result = calc_table_embedding(self.candidates(), + self.TARGET if mention == 'name' else self.COMMENT, + keywords='采购') + assert len(result) == 10 + assert result[0]['table_name'] == self.TARGET + + @pytest.mark.parametrize('failure', ['load', 'encode', 'invalid_embedding', 'missing_embedding']) + def test_vector_failure_does_not_remove_explicit_table(self, failure): + tables = self.candidates() + if failure == 'load': + self.cache.get_model.side_effect = RuntimeError('model unavailable') + elif failure == 'encode': + self.cache.get_model.return_value.embed_query.side_effect = RuntimeError('encoding failed') + elif failure == 'invalid_embedding': + tables[0]['embedding'] = 'invalid json' + else: + tables[-1]['embedding'] = None + result = calc_table_embedding(tables, f'查询{self.COMMENT}', keywords='采购') + assert len(result) == (10 if failure == 'missing_embedding' else len(tables)) + assert result[0]['table_name'] == self.TARGET + + @pytest.mark.parametrize('keywords', [None, '销售额']) + @pytest.mark.parametrize('failure', ['load', 'encode', 'invalid_embedding']) + def test_vector_failure_keeps_all_candidates_without_explicit_name(self, keywords, failure): + tables = self.candidates() + if failure == 'load': + self.cache.get_model.side_effect = RuntimeError('model unavailable') + elif failure == 'encode': + self.cache.get_model.return_value.embed_query.side_effect = RuntimeError('encoding failed') + else: + tables[0]['embedding'] = 'invalid json' + result = calc_table_embedding(tables, '统计本月销售额', keywords=keywords) + assert [t['id'] for t in result] == [t['id'] for t in tables] + + @pytest.mark.parametrize('question,long_name,short_name', [ + ('查询订单明细', '订单明细', '订单'), + ('查询 "sales-order"', 'sales-order', 'sales'), + ('查询 `sales-order`', 'sales-order', 'sales'), + ('查询 [sales-order]', 'sales-order', 'sales'), + ('查询 public."sales-order"', 'sales-order', 'sales'), + ]) + def test_long_physical_name_does_not_require_short_name(self, question, long_name, short_name): + self.config.TABLE_EMBEDDING_COUNT = 1 + tables = [self.table(short_name), self.table(long_name, embedding=[0, 1])] + result = calc_table_embedding(tables, question) + assert [t['table_name'] for t in result] == [long_name] + result = calc_table_embedding(tables, question + ',以及 ' + short_name) + assert {t['table_name'] for t in result} == {long_name, short_name} + + @pytest.mark.parametrize('quoted', ['"sales-order"', '`sales-order`', '[sales-order]', '"订单明细"']) + def test_unknown_quoted_identifier_does_not_match_part(self, quoted): + self.config.TABLE_EMBEDDING_COUNT = 1 + tables = [self.table('other'), self.table('sales', embedding=[0, 1]), + self.table('订单', embedding=[0, 1])] + assert calc_table_embedding(tables, '查询 ' + quoted)[0]['table_name'] == 'other' + + @pytest.mark.parametrize('name,quoted', [('a"b', '"a""b"'), ('a`b', '`a``b`'), ('a]b', '[a]]b]')]) + def test_escaped_quoted_identifier(self, name, quoted): + self.config.TABLE_EMBEDDING_COUNT = 1 + tables = [self.table('other'), self.table(name, embedding=[0, 1])] + assert calc_table_embedding(tables, '查询 ' + quoted)[0]['table_name'] == name + + @pytest.mark.parametrize('question,expected', [ + ('查询 SALES_ORDER 的金额', 'sales_order'), + ('查询sales_order的金额', 'sales_order'), + ('查询 public.sales_order 的金额', 'sales_order'), + ('查询 "sales_order" 的金额', 'sales_order'), + ('查询 `sales_order` 的金额', 'sales_order'), + ('查询 [sales_order] 的金额', 'sales_order'), + ('查询 sales_order_detail 的金额', 'other'), + ('查询 old_sales_order 的金额', 'other'), + ('查询 sales_order2 的金额', 'other'), + ('查询 sales_order$backup 的金额', 'other'), + ]) + def test_physical_name_boundaries(self, question, expected): + self.config.TABLE_EMBEDDING_COUNT = 1 + tables = [self.table('other'), self.table('sales_order', embedding=[0, 1])] + result = calc_table_embedding(tables, question) + assert [t['table_name'] for t in result] == [expected] + + def test_longer_comment_wins_only_at_same_occurrence(self): + self.config.TABLE_EMBEDDING_COUNT = 1 + tables = [self.table('short', '采购订单'), self.table('long', '采购订单明细', [0, 1])] + result = calc_table_embedding(tables, '查询采购订单明细') + assert [t['table_name'] for t in result] == ['long'] + result = calc_table_embedding(tables, '比较采购订单和采购订单明细') + assert {t['table_name'] for t in result} == {'short', 'long'} + + def test_duplicate_comments_keep_all_candidates(self): + self.config.TABLE_EMBEDDING_COUNT = 1 + tables = [self.table('other'), self.table('a', '采购订单'), self.table('b', '采购订单')] + result = calc_table_embedding(tables, '查询采购订单') + assert {t['table_name'] for t in result} == {'a', 'b'} + + def test_explicit_names_can_exceed_limit_without_duplicates(self): + tables = [self.table(f'purchase_{i}') for i in range(20)] + names = [t['table_name'] for t in tables[-12:]] + result = calc_table_embedding(tables, '查询 ' + ','.join(names + names), keywords='采购') + assert len(result) == 12 + assert {t['table_name'] for t in result} == set(names) + + def test_explicit_table_not_in_candidates_is_not_added(self): + tables = self.candidates()[:-1] + result = calc_table_embedding(tables, self.TARGET) + assert len(result) == 10 + assert self.TARGET not in {t['table_name'] for t in result} + + def test_no_explicit_mention_keeps_existing_ranking(self): + tables = self.candidates() + result = calc_table_embedding(tables, '查询采购金额', keywords='采购,金额') + assert [t['id'] for t in result] == [t['id'] for t in tables[:10]] + + def test_empty_candidates(self): + assert calc_table_embedding([], self.TARGET, keywords='采购') == [] + + if __name__ == '__main__': pytest.main([__file__, '-v'])