1919 import dmPython
2020import pymysql
2121import redshift_connector
22- from sqlalchemy import create_engine , text , Engine
22+ from sqlalchemy import create_engine , text , Engine , types
2323from sqlalchemy .orm import sessionmaker
2424
2525from apps .datasource .models .datasource import DatasourceConf , CoreDatasource , TableSchema , ColumnSchema
@@ -596,12 +596,35 @@ def exec_sql(ds: CoreDatasource | AssistantOutDsSchema, sql: str, origin_column=
596596 with session .execute (text (sql )) as result :
597597 try :
598598 columns = result .keys ()._keys if origin_column else [item .lower () for item in result .keys ()._keys ]
599+
600+ fields_info = []
601+ for col_info in result .cursor .description :
602+ # col_info 是 (name, type_code, display_size, internal_size, precision, scale, null_ok)
603+ col_name = col_info [0 ]
604+
605+ # 根据 type_code 判断是否为数值类型
606+ # psycopg2 的类型 OID 常量
607+ is_numeric = col_info [1 ] in (
608+ 20 , # int8
609+ 21 , # int2
610+ 23 , # int4
611+ 700 , # float4
612+ 701 , # float8
613+ 1700 , # numeric
614+ 16 , # boolean
615+ )
616+
617+ fields_info .append ({
618+ "name" : col_name if origin_column else col_name .lower (),
619+ "is_numeric" : is_numeric
620+ })
621+
599622 res = result .fetchall ()
600623 result_list = [
601624 {str (columns [i ]): convert_value (value ) for i , value in enumerate (tuple_item )} for tuple_item in
602625 res
603626 ]
604- return {"fields" : columns , "data" : result_list ,
627+ return {"fields" : columns , "data" : result_list , "fields_info" : fields_info ,
605628 "sql" : bytes .decode (base64 .b64encode (bytes (sql , 'utf-8' )))}
606629 except Exception as ex :
607630 raise ParseSQLResultError (str (ex ))
@@ -617,11 +640,12 @@ def exec_sql(ds: CoreDatasource | AssistantOutDsSchema, sql: str, origin_column=
617640 columns = [field [0 ] for field in cursor .description ] if origin_column else [field [0 ].lower () for
618641 field in
619642 cursor .description ]
643+ fields_info = build_fields_info_from_cursor (cursor , origin_column , 'dm' )
620644 result_list = [
621645 {str (columns [i ]): convert_value (value ) for i , value in enumerate (tuple_item )} for tuple_item in
622646 res
623647 ]
624- return {"fields" : columns , "data" : result_list ,
648+ return {"fields" : columns , "data" : result_list , "fields_info" : fields_info ,
625649 "sql" : bytes .decode (base64 .b64encode (bytes (sql , 'utf-8' )))}
626650 except Exception as ex :
627651 raise ParseSQLResultError (str (ex ))
@@ -637,11 +661,12 @@ def exec_sql(ds: CoreDatasource | AssistantOutDsSchema, sql: str, origin_column=
637661 columns = [field [0 ] for field in cursor .description ] if origin_column else [field [0 ].lower () for
638662 field in
639663 cursor .description ]
664+ fields_info = build_fields_info_from_cursor (cursor , origin_column , 'mysql' )
640665 result_list = [
641666 {str (columns [i ]): convert_value (value ) for i , value in enumerate (tuple_item )} for tuple_item in
642667 res
643668 ]
644- return {"fields" : columns , "data" : result_list ,
669+ return {"fields" : columns , "data" : result_list , "fields_info" : fields_info ,
645670 "sql" : bytes .decode (base64 .b64encode (bytes (sql , 'utf-8' )))}
646671 except Exception as ex :
647672 raise ParseSQLResultError (str (ex ))
@@ -655,11 +680,12 @@ def exec_sql(ds: CoreDatasource | AssistantOutDsSchema, sql: str, origin_column=
655680 columns = [field [0 ] for field in cursor .description ] if origin_column else [field [0 ].lower () for
656681 field in
657682 cursor .description ]
683+ fields_info = build_fields_info_from_cursor (cursor , origin_column , 'postgresql' )
658684 result_list = [
659685 {str (columns [i ]): convert_value (value ) for i , value in enumerate (tuple_item )} for tuple_item in
660686 res
661687 ]
662- return {"fields" : columns , "data" : result_list ,
688+ return {"fields" : columns , "data" : result_list , "fields_info" : fields_info ,
663689 "sql" : bytes .decode (base64 .b64encode (bytes (sql , 'utf-8' )))}
664690 except Exception as ex :
665691 raise ParseSQLResultError (str (ex ))
@@ -674,25 +700,28 @@ def exec_sql(ds: CoreDatasource | AssistantOutDsSchema, sql: str, origin_column=
674700 columns = [field [0 ] for field in cursor .description ] if origin_column else [field [0 ].lower () for
675701 field in
676702 cursor .description ]
703+ fields_info = build_fields_info_from_cursor (cursor , origin_column , 'postgresql' )
677704 result_list = [
678705 {str (columns [i ]): convert_value (value ) for i , value in enumerate (tuple_item )} for tuple_item in
679706 res
680707 ]
681- return {"fields" : columns , "data" : result_list ,
708+ return {"fields" : columns , "data" : result_list , "fields_info" : fields_info ,
682709 "sql" : bytes .decode (base64 .b64encode (bytes (sql , 'utf-8' )))}
683710 except Exception as ex :
684711 raise ParseSQLResultError (str (ex ))
685712 elif equals_ignore_case (ds .type , 'es' ):
686713 try :
687- res , columns = get_es_data_by_http (conf , sql )
688- columns = [field .get ('name' ) for field in columns ] if origin_column else [field .get ('name' ).lower () for
689- field in
690- columns ]
714+ res , raw_columns = get_es_data_by_http (conf , sql )
715+ columns = [field .get ('name' ) for field in raw_columns ] if origin_column else [field .get ('name' ).lower ()
716+ for
717+ field in
718+ raw_columns ]
719+ fields_info = build_fields_info_from_es (raw_columns , origin_column )
691720 result_list = [
692721 {str (columns [i ]): convert_value (value ) for i , value in enumerate (tuple_item )} for tuple_item in
693722 res
694723 ]
695- return {"fields" : columns , "data" : result_list ,
724+ return {"fields" : columns , "data" : result_list , "fields_info" : fields_info ,
696725 "sql" : bytes .decode (base64 .b64encode (bytes (sql , 'utf-8' )))}
697726 except Exception as ex :
698727 raise Exception (str (ex ))
@@ -707,16 +736,120 @@ def exec_sql(ds: CoreDatasource | AssistantOutDsSchema, sql: str, origin_column=
707736 columns = [field [0 ] for field in cursor .description ] if origin_column else [field [0 ].lower () for
708737 field in
709738 cursor .description ]
739+ fields_info = build_fields_info_from_cursor (cursor , origin_column , 'hive' )
710740 result_list = [
711741 {str (columns [i ]): convert_value (value ) for i , value in enumerate (tuple_item )} for tuple_item in
712742 res
713743 ]
714- return {"fields" : columns , "data" : result_list ,
744+ return {"fields" : columns , "data" : result_list , "fields_info" : fields_info ,
715745 "sql" : bytes .decode (base64 .b64encode (bytes (hive_sql , 'utf-8' )))}
716746 except Exception as ex :
717747 raise ParseSQLResultError (str (ex ))
718748
719749
750+ def build_fields_info_from_cursor (cursor , origin_column , db_type = 'postgresql' ):
751+ """
752+ 根据数据库游标的 description 构建字段信息列表
753+
754+ Args:
755+ cursor: 数据库游标对象
756+ origin_column: 是否保留原始列名大小写
757+ db_type: 数据库类型,支持 'mysql', 'postgresql', 'redshift', 'kingbase', 'dm', 'hive'
758+
759+ Returns:
760+ list: 包含字段名和是否数值类型的字典列表
761+ """
762+ fields_info = []
763+
764+ for col_info in cursor .description :
765+ col_name = col_info [0 ]
766+
767+ if db_type == 'mysql' :
768+ # MySQL/pymysql 类型码
769+ is_numeric = col_info [1 ] in (
770+ 1 , # TINYINT
771+ 2 , # SMALLINT
772+ 3 , # INT
773+ 4 , # FLOAT
774+ 5 , # DOUBLE
775+ 8 , # BIGINT
776+ 9 , # MEDIUMINT
777+ 16 , # BIT
778+ 246 , # DECIMAL
779+ )
780+ elif db_type in ('postgresql' , 'redshift' , 'kingbase' ):
781+ # PostgreSQL/psycopg2 类型 OID
782+ is_numeric = col_info [1 ] in (
783+ 20 , # int8
784+ 21 , # int2
785+ 23 , # int4
786+ 700 , # float4
787+ 701 , # float8
788+ 1700 , # numeric
789+ 16 , # boolean
790+ )
791+ elif db_type == 'dm' :
792+ # 达梦数据库类型码
793+ is_numeric = col_info [1 ] in (
794+ 3 , # DECIMAL/NUMERIC
795+ 2 , # NUMBER
796+ 4 , # INTEGER
797+ 5 , # INT
798+ 6 , # BIGINT
799+ 7 , # TINYINT
800+ 8 , # BYTE
801+ 9 , # FLOAT
802+ 10 , # DOUBLE
803+ 11 , # REAL
804+ 12 , # BOOLEAN
805+ )
806+ elif db_type == 'hive' :
807+ # Hive 类型对象转字符串判断
808+ type_str = str (col_info [1 ]).lower ()
809+ NUMERIC_PREFIXES = ('tinyint' , 'smallint' , 'int' , 'bigint' , 'float' , 'double' , 'decimal' , 'numeric' )
810+ is_numeric = type_str == 'boolean' or any (type_str .startswith (p ) for p in NUMERIC_PREFIXES )
811+ else :
812+ is_numeric = False
813+
814+ fields_info .append ({
815+ "name" : col_name if origin_column else col_name .lower (),
816+ "is_numeric" : is_numeric
817+ })
818+
819+ return fields_info
820+
821+
822+ def build_fields_info_from_es (raw_columns , origin_column ):
823+ """
824+ 专门为 Elasticsearch 构建字段信息
825+
826+ Args:
827+ raw_columns: ES 返回的列信息列表
828+ origin_column: 是否保留原始列名大小写
829+
830+ Returns:
831+ list: 包含字段名和是否数值类型的字典列表
832+ """
833+ fields_info = []
834+
835+ for field in raw_columns :
836+ field_name = field .get ('name' ) if origin_column else field .get ('name' ).lower ()
837+ field_type = field .get ('type' , '' ).lower ()
838+
839+ is_numeric = field_type in (
840+ 'long' , 'integer' , 'short' , 'byte' ,
841+ 'double' , 'float' , 'half_float' , 'scaled_float' ,
842+ 'unsigned_long' , 'boolean'
843+ )
844+
845+ fields_info .append ({
846+ "name" : field_name ,
847+ "is_numeric" : is_numeric
848+ })
849+
850+ return fields_info
851+
852+
720853def get_sqlglot_dialect (ds_type : str ) -> str :
721854 """根据数据源类型获取 sqlglot dialect"""
722855 if equals_ignore_case (ds_type , 'mysql' , 'doris' , 'starrocks' ):
@@ -744,6 +877,7 @@ def get_sqlglot_dialect(ds_type: str) -> str:
744877
745878# 危险模式正则表达式(用于检查特殊语法)
746879import re
880+
747881DANGEROUS_PATTERNS = [
748882 r'\bINTO\s+OUTFILE\b' ,
749883 r'\bINTO\s+DUMPFILE\b' ,
@@ -765,7 +899,7 @@ def check_dangerous_functions(statements: list, ds_type: str) -> bool:
765899 """检查是否使用了危险函数,返回 True 表示安全"""
766900 dangerous_functions = get_dangerous_functions (ds_type )
767901 dangerous_functions_upper = {f .upper () for f in dangerous_functions }
768-
902+
769903 for stmt in statements :
770904 if stmt :
771905 for func in stmt .find_all (exp .Anonymous ):
0 commit comments