Skip to content
Merged
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
76 changes: 70 additions & 6 deletions backend/apps/datasource/embedding/table_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'(?<![{_IDENTIFIER_CHARS}])' if re.match(rf'[{_IDENTIFIER_CHARS}]', name[0]) else ''
right = rf'(?![{_IDENTIFIER_CHARS}])' if re.match(rf'[{_IDENTIFIER_CHARS}]', name[-1]) else ''
return [match.span() for match in re.finditer(left + re.escape(name) + right, question)]


def _longest_mentions(spans_by_name: dict):
"""Keep the longest overlapping name and any independently mentioned shorter names."""
spans = {span for occurrences in spans_by_name.values() for span in occurrences}
longest_spans = set()
furthest_end = -1
for start, end in sorted(spans, key=lambda span: (span[0], -span[1])):
if end > 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):
Expand Down Expand Up @@ -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: 已提取并扩展的关键词(由调用方传入)
Expand Down Expand Up @@ -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")
Expand All @@ -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({
Expand Down Expand Up @@ -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)} 张表")
Expand All @@ -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)
152 changes: 152 additions & 0 deletions backend/tests/test_table_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'])
Loading