@@ -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+
611763if __name__ == '__main__' :
612764 pytest .main ([__file__ , '-v' ])
0 commit comments