Skip to content

Commit 6fa74aa

Browse files
fix(datasource): retain explicitly referenced tables during schema selection
1 parent 13f0076 commit 6fa74aa

2 files changed

Lines changed: 222 additions & 6 deletions

File tree

‎backend/apps/datasource/embedding/table_embedding.py‎

Lines changed: 70 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,69 @@
1212

1313
# 预编译正则,避免每次调用都重新编译
1414
_RE_UNDERSCORE_SPACE = re.compile(r'[_\s]+')
15+
_IDENTIFIER_CHARS = r'A-Za-z0-9_$'
16+
_RE_QUOTED_IDENTIFIER = re.compile(r'"(?:[^"]|"")*"|`(?:[^`]|``)*`|\[(?:[^\]]|\]\])*\]')
17+
18+
19+
def _mention_spans(question: str, name: str):
20+
"""Match identifier boundaries while allowing adjacent Chinese text and SQL delimiters."""
21+
left = rf'(?<![{_IDENTIFIER_CHARS}])' if re.match(rf'[{_IDENTIFIER_CHARS}]', name[0]) else ''
22+
right = rf'(?![{_IDENTIFIER_CHARS}])' if re.match(rf'[{_IDENTIFIER_CHARS}]', name[-1]) else ''
23+
return [match.span() for match in re.finditer(left + re.escape(name) + right, question)]
24+
25+
26+
def _longest_mentions(spans_by_name: dict):
27+
"""Keep the longest overlapping name and any independently mentioned shorter names."""
28+
spans = {span for occurrences in spans_by_name.values() for span in occurrences}
29+
longest_spans = set()
30+
furthest_end = -1
31+
for start, end in sorted(spans, key=lambda span: (span[0], -span[1])):
32+
if end > furthest_end:
33+
longest_spans.add((start, end))
34+
furthest_end = end
35+
return {name for name, occurrences in spans_by_name.items()
36+
if any(span in longest_spans for span in occurrences)}
37+
38+
39+
def _select_tables(ranked_tables: list[dict], question: str, *, apply_limit: bool = True):
40+
"""Keep explicitly mentioned candidates, then fill remaining slots by existing rank.
41+
42+
Inspect only candidates already filtered by the caller; retain duplicate comments.
43+
Resolve overlapping names and comments separately. Error fallback can disable truncation.
44+
"""
45+
question = (question or '').lower()
46+
quoted_names = []
47+
for match in _RE_QUOTED_IDENTIFIER.finditer(question):
48+
token = match.group()
49+
closing = token[-1]
50+
quoted_names.append((match.span(), token[1:-1].replace(closing * 2, closing)))
51+
name_spans, comment_spans = {}, {}
52+
for table in ranked_tables:
53+
name = (table.get('table_name') or '').strip().lower()
54+
if name and name not in name_spans:
55+
# Match quoted identifiers as a whole, even when the full name is not a candidate.
56+
name_spans[name] = [
57+
(a, b) for a, b in _mention_spans(question, name)
58+
if not any(a < end and start < b for (start, end), _ in quoted_names)
59+
]
60+
name_spans[name].extend(span for span, quoted_name in quoted_names if quoted_name == name)
61+
comment = (table.get('table_comment') or '').strip().lower()
62+
if comment and comment not in comment_spans:
63+
comment_spans[comment] = _mention_spans(question, comment)
64+
65+
explicit_names = _longest_mentions(name_spans)
66+
explicit_comments = _longest_mentions(comment_spans)
67+
required, remaining = [], []
68+
for table in ranked_tables:
69+
name = (table.get('table_name') or '').strip().lower()
70+
comment = (table.get('table_comment') or '').strip().lower()
71+
if name in explicit_names or comment in explicit_comments:
72+
required.append(table)
73+
else:
74+
remaining.append(table)
75+
if apply_limit:
76+
remaining = remaining[:max(0, settings.TABLE_EMBEDDING_COUNT - len(required))]
77+
return required + remaining
1578

1679

1780
def _parse_embedding(embedding):
@@ -177,13 +240,13 @@ def calc_table_embedding(tables: list[dict], question: str, session=None, oid: i
177240
1. 用 keywords 计算向量相似度(vec_score)
178241
2. 计算关键词匹配度(keyword_score)
179242
3. 融合评分:final_score = α * vec_score + (1-α) * keyword_score
180-
4. keyword_score=1.0 的表保证排在最前面
243+
4. Keep explicitly mentioned table names or complete comments, then fill slots by rank.
181244
182245
当 keywords 为空或 TABLE_EMBEDDING_KEYWORD_ENABLED 为 False 时,回退到纯向量匹配。
183246
184247
Args:
185248
tables: 表列表,每个 dict 包含 id, table_name, schema_table, embedding, table_comment, fields
186-
question: 原始用户问题(用于纯向量匹配回退)
249+
question: Original user question for explicit matching and vector-only fallback.
187250
session: 数据库会话(预留)
188251
oid: 组织 ID(预留)
189252
keywords: 已提取并扩展的关键词(由调用方传入)
@@ -272,7 +335,7 @@ def calc_table_embedding(tables: list[dict], question: str, session=None, oid: i
272335
other_tables.sort(key=lambda x: x['cosine_similarity'], reverse=True)
273336

274337
_list = exact_matches + other_tables
275-
_list = _list[:settings.TABLE_EMBEDDING_COUNT]
338+
_list = _select_tables(_list, question)
276339

277340
end_time = time.time()
278341
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
292355

293356

294357
def _calc_vector_only(tables: list[dict], question: str):
295-
"""纯向量相似度评分(原始逻辑)。"""
358+
"""Rank by vector similarity while retaining explicitly mentioned tables."""
296359
_list = []
297360
for table in tables:
298361
_list.append({
@@ -329,7 +392,7 @@ def _calc_vector_only(tables: list[dict], question: str):
329392
_list[idx]['cosine_similarity'] = float(sim)
330393

331394
_list.sort(key=lambda x: x['cosine_similarity'], reverse=True)
332-
_list = _list[:settings.TABLE_EMBEDDING_COUNT]
395+
_list = _select_tables(_list, question)
333396

334397
end_time = time.time()
335398
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):
342405
return _list
343406
except Exception:
344407
traceback.print_exc()
345-
return _list
408+
# Preserve all candidates when vector ranking fails; move explicit matches to the front.
409+
return _select_tables(_list, question, apply_limit=False)

‎backend/tests/test_table_embedding.py‎

Lines changed: 152 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -608,5 +608,157 @@ def test_fallback_on_unexpected_error(self, mock_embed_cache, mock_settings):
608608
assert len(result) == 2
609609

610610

611+
class TestExplicitTableSelection:
612+
"""Retain explicitly named tables and comments even with low vector rankings."""
613+
614+
TARGET = 'ads_jt_scm_purchase_process_chain_relation_detail_full_1d'
615+
COMMENT = '采购全链路数据各节点表'
616+
617+
@pytest.fixture(autouse=True)
618+
def setup_models(self):
619+
with patch(EMBEDDING_CACHE_PATCH) as cache, patch(SETTINGS_PATCH) as config:
620+
config.TABLE_EMBEDDING_KEYWORD_ENABLED = True
621+
config.TABLE_EMBEDDING_COUNT = 10
622+
config.TABLE_EMBEDDING_ALPHA = 0.4
623+
cache.get_model.return_value.embed_query.return_value = [1.0, 0.0]
624+
self.cache, self.config = cache, config
625+
yield
626+
627+
def table(self, name, comment='', embedding=None):
628+
return dict(id=name, table_name=name, table_comment=comment,
629+
schema_table=f'# Table: {name}', fields=[],
630+
embedding=json.dumps(embedding if embedding is not None else [0.95, 0.31225]))
631+
632+
def candidates(self):
633+
return [self.table(f'{self.TARGET}_{i}', self.COMMENT + f'(环节{i}明细)')
634+
for i in range(19)] + [self.table(self.TARGET, self.COMMENT, [0.6, 0.8])]
635+
636+
@pytest.mark.parametrize('keywords', [None, '采购,节点,金额', '查询采购金额'])
637+
@pytest.mark.parametrize('mention', ['name', 'comment'])
638+
def test_explicit_mention_survives_top_ten(self, keywords, mention):
639+
question = f'查询{self.TARGET if mention == "name" else self.COMMENT}的采购金额'
640+
result = calc_table_embedding(self.candidates(), question, keywords=keywords)
641+
assert len(result) == 10
642+
assert result[0]['table_name'] == self.TARGET
643+
644+
@pytest.mark.parametrize('mention', ['name', 'comment'])
645+
def test_keyword_disabled(self, mention):
646+
self.config.TABLE_EMBEDDING_KEYWORD_ENABLED = False
647+
result = calc_table_embedding(self.candidates(),
648+
self.TARGET if mention == 'name' else self.COMMENT,
649+
keywords='采购')
650+
assert len(result) == 10
651+
assert result[0]['table_name'] == self.TARGET
652+
653+
@pytest.mark.parametrize('failure', ['load', 'encode', 'invalid_embedding', 'missing_embedding'])
654+
def test_vector_failure_does_not_remove_explicit_table(self, failure):
655+
tables = self.candidates()
656+
if failure == 'load':
657+
self.cache.get_model.side_effect = RuntimeError('model unavailable')
658+
elif failure == 'encode':
659+
self.cache.get_model.return_value.embed_query.side_effect = RuntimeError('encoding failed')
660+
elif failure == 'invalid_embedding':
661+
tables[0]['embedding'] = 'invalid json'
662+
else:
663+
tables[-1]['embedding'] = None
664+
result = calc_table_embedding(tables, f'查询{self.COMMENT}', keywords='采购')
665+
assert len(result) == (10 if failure == 'missing_embedding' else len(tables))
666+
assert result[0]['table_name'] == self.TARGET
667+
668+
@pytest.mark.parametrize('keywords', [None, '销售额'])
669+
@pytest.mark.parametrize('failure', ['load', 'encode', 'invalid_embedding'])
670+
def test_vector_failure_keeps_all_candidates_without_explicit_name(self, keywords, failure):
671+
tables = self.candidates()
672+
if failure == 'load':
673+
self.cache.get_model.side_effect = RuntimeError('model unavailable')
674+
elif failure == 'encode':
675+
self.cache.get_model.return_value.embed_query.side_effect = RuntimeError('encoding failed')
676+
else:
677+
tables[0]['embedding'] = 'invalid json'
678+
result = calc_table_embedding(tables, '统计本月销售额', keywords=keywords)
679+
assert [t['id'] for t in result] == [t['id'] for t in tables]
680+
681+
@pytest.mark.parametrize('question,long_name,short_name', [
682+
('查询订单明细', '订单明细', '订单'),
683+
('查询 "sales-order"', 'sales-order', 'sales'),
684+
('查询 `sales-order`', 'sales-order', 'sales'),
685+
('查询 [sales-order]', 'sales-order', 'sales'),
686+
('查询 public."sales-order"', 'sales-order', 'sales'),
687+
])
688+
def test_long_physical_name_does_not_require_short_name(self, question, long_name, short_name):
689+
self.config.TABLE_EMBEDDING_COUNT = 1
690+
tables = [self.table(short_name), self.table(long_name, embedding=[0, 1])]
691+
result = calc_table_embedding(tables, question)
692+
assert [t['table_name'] for t in result] == [long_name]
693+
result = calc_table_embedding(tables, question + ',以及 ' + short_name)
694+
assert {t['table_name'] for t in result} == {long_name, short_name}
695+
696+
@pytest.mark.parametrize('quoted', ['"sales-order"', '`sales-order`', '[sales-order]', '"订单明细"'])
697+
def test_unknown_quoted_identifier_does_not_match_part(self, quoted):
698+
self.config.TABLE_EMBEDDING_COUNT = 1
699+
tables = [self.table('other'), self.table('sales', embedding=[0, 1]),
700+
self.table('订单', embedding=[0, 1])]
701+
assert calc_table_embedding(tables, '查询 ' + quoted)[0]['table_name'] == 'other'
702+
703+
@pytest.mark.parametrize('name,quoted', [('a"b', '"a""b"'), ('a`b', '`a``b`'), ('a]b', '[a]]b]')])
704+
def test_escaped_quoted_identifier(self, name, quoted):
705+
self.config.TABLE_EMBEDDING_COUNT = 1
706+
tables = [self.table('other'), self.table(name, embedding=[0, 1])]
707+
assert calc_table_embedding(tables, '查询 ' + quoted)[0]['table_name'] == name
708+
709+
@pytest.mark.parametrize('question,expected', [
710+
('查询 SALES_ORDER 的金额', 'sales_order'),
711+
('查询sales_order的金额', 'sales_order'),
712+
('查询 public.sales_order 的金额', 'sales_order'),
713+
('查询 "sales_order" 的金额', 'sales_order'),
714+
('查询 `sales_order` 的金额', 'sales_order'),
715+
('查询 [sales_order] 的金额', 'sales_order'),
716+
('查询 sales_order_detail 的金额', 'other'),
717+
('查询 old_sales_order 的金额', 'other'),
718+
('查询 sales_order2 的金额', 'other'),
719+
('查询 sales_order$backup 的金额', 'other'),
720+
])
721+
def test_physical_name_boundaries(self, question, expected):
722+
self.config.TABLE_EMBEDDING_COUNT = 1
723+
tables = [self.table('other'), self.table('sales_order', embedding=[0, 1])]
724+
result = calc_table_embedding(tables, question)
725+
assert [t['table_name'] for t in result] == [expected]
726+
727+
def test_longer_comment_wins_only_at_same_occurrence(self):
728+
self.config.TABLE_EMBEDDING_COUNT = 1
729+
tables = [self.table('short', '采购订单'), self.table('long', '采购订单明细', [0, 1])]
730+
result = calc_table_embedding(tables, '查询采购订单明细')
731+
assert [t['table_name'] for t in result] == ['long']
732+
result = calc_table_embedding(tables, '比较采购订单和采购订单明细')
733+
assert {t['table_name'] for t in result} == {'short', 'long'}
734+
735+
def test_duplicate_comments_keep_all_candidates(self):
736+
self.config.TABLE_EMBEDDING_COUNT = 1
737+
tables = [self.table('other'), self.table('a', '采购订单'), self.table('b', '采购订单')]
738+
result = calc_table_embedding(tables, '查询采购订单')
739+
assert {t['table_name'] for t in result} == {'a', 'b'}
740+
741+
def test_explicit_names_can_exceed_limit_without_duplicates(self):
742+
tables = [self.table(f'purchase_{i}') for i in range(20)]
743+
names = [t['table_name'] for t in tables[-12:]]
744+
result = calc_table_embedding(tables, '查询 ' + ','.join(names + names), keywords='采购')
745+
assert len(result) == 12
746+
assert {t['table_name'] for t in result} == set(names)
747+
748+
def test_explicit_table_not_in_candidates_is_not_added(self):
749+
tables = self.candidates()[:-1]
750+
result = calc_table_embedding(tables, self.TARGET)
751+
assert len(result) == 10
752+
assert self.TARGET not in {t['table_name'] for t in result}
753+
754+
def test_no_explicit_mention_keeps_existing_ranking(self):
755+
tables = self.candidates()
756+
result = calc_table_embedding(tables, '查询采购金额', keywords='采购,金额')
757+
assert [t['id'] for t in result] == [t['id'] for t in tables[:10]]
758+
759+
def test_empty_candidates(self):
760+
assert calc_table_embedding([], self.TARGET, keywords='采购') == []
761+
762+
611763
if __name__ == '__main__':
612764
pytest.main([__file__, '-v'])

0 commit comments

Comments
 (0)