Skip to content

Commit 36a2d02

Browse files
committed
feat: add thousands separator settings
1 parent a2196d9 commit 36a2d02

3 files changed

Lines changed: 174 additions & 26 deletions

File tree

backend/apps/chat/curd/chat.py

Lines changed: 18 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -187,7 +187,8 @@ def get_last_execute_sql_error(session: SessionDep, chart_id: int):
187187

188188

189189
def format_json_data(origin_data: dict):
190-
result = {'fields': origin_data.get('fields') if origin_data.get('fields') else []}
190+
result = {'fields': origin_data.get('fields') if origin_data.get('fields') else [],
191+
'fields_info': origin_data.get('fields_info') if origin_data.get('fields_info') else None}
191192
_list = origin_data.get('data') if origin_data.get('data') else []
192193
data = format_json_list_data(_list)
193194
result['data'] = data
@@ -238,21 +239,24 @@ def get_chart_data_with_user(session: SessionDep, current_user: CurrentUser, cha
238239
pass
239240
return {}
240241

242+
241243
def get_chart_data_with_user_live(session: SessionDep, current_user: CurrentUser, chat_record_id: int):
242-
stmt = select(ChatRecord.datasource,ChatRecord.sql).where(and_(ChatRecord.id == chat_record_id, ChatRecord.create_by == current_user.id))
244+
stmt = select(ChatRecord.datasource, ChatRecord.sql).where(
245+
and_(ChatRecord.id == chat_record_id, ChatRecord.create_by == current_user.id))
243246
row = session.execute(stmt).first()
244-
return get_chart_data_ds(session,row.datasource, row.sql)
247+
return get_chart_data_ds(session, row.datasource, row.sql)
245248

246-
def get_chart_data_ds(session: SessionDep,ds_id,sql):
247-
json_result: Dict[str, Any] = {'status': 'success','data':[],'message':''}
249+
250+
def get_chart_data_ds(session: SessionDep, ds_id, sql):
251+
json_result: Dict[str, Any] = {'status': 'success', 'data': [], 'message': ''}
248252
try:
249-
datasource = get_ds(session,ds_id)
253+
datasource = get_ds(session, ds_id)
250254
if datasource is None:
251255
json_result['status'] = 'failed'
252256
json_result['message'] = 'Datasource not found'
253257
return json_result
254258
else:
255-
result = exec_sql(ds=datasource,sql=sql, origin_column=False)
259+
result = exec_sql(ds=datasource, sql=sql, origin_column=False)
256260
_data = DataFormat.convert_large_numbers_in_object_array(result.get('data'))
257261
_data = DataFormat.normalize_qualified_sql_column_keys_in_object_array(_data)
258262
json_result['data'] = _data
@@ -264,6 +268,7 @@ def get_chart_data_ds(session: SessionDep,ds_id,sql):
264268
pass
265269
return json_result
266270

271+
267272
def get_chat_chart_data(session: SessionDep, chat_record_id: int):
268273
stmt = select(ChatRecord.data).where(and_(ChatRecord.id == chat_record_id))
269274
res = session.execute(stmt)
@@ -336,7 +341,7 @@ def get_chat_with_records(session: SessionDep, chart_id: int, current_user: Curr
336341
predict_alias_log = aliased(ChatLog)
337342

338343
stmt = (select(ChatRecord.id, ChatRecord.chat_id, ChatRecord.create_time, ChatRecord.finish_time,
339-
ChatRecord.question, ChatRecord.sql_answer, ChatRecord.sql,ChatRecord.datasource,
344+
ChatRecord.question, ChatRecord.sql_answer, ChatRecord.sql, ChatRecord.datasource,
340345
ChatRecord.chart_answer, ChatRecord.chart, ChatRecord.analysis, ChatRecord.predict,
341346
ChatRecord.datasource_select_answer, ChatRecord.analysis_record_id, ChatRecord.predict_record_id,
342347
ChatRecord.regenerate_record_id,
@@ -363,7 +368,7 @@ def get_chat_with_records(session: SessionDep, chart_id: int, current_user: Curr
363368
ChatRecord.create_time))
364369
if with_data:
365370
stmt = select(ChatRecord.id, ChatRecord.chat_id, ChatRecord.create_time, ChatRecord.finish_time,
366-
ChatRecord.question, ChatRecord.sql_answer, ChatRecord.sql,ChatRecord.datasource,
371+
ChatRecord.question, ChatRecord.sql_answer, ChatRecord.sql, ChatRecord.datasource,
367372
ChatRecord.chart_answer, ChatRecord.chart, ChatRecord.analysis, ChatRecord.predict,
368373
ChatRecord.datasource_select_answer, ChatRecord.analysis_record_id, ChatRecord.predict_record_id,
369374
ChatRecord.regenerate_record_id,
@@ -429,7 +434,8 @@ def get_chat_with_records(session: SessionDep, chart_id: int, current_user: Curr
429434
finish_time=row.finish_time,
430435
duration=duration,
431436
total_tokens=total_tokens,
432-
question=row.question, sql_answer=row.sql_answer, sql=row.sql, datasource=row.datasource,
437+
question=row.question, sql_answer=row.sql_answer, sql=row.sql,
438+
datasource=row.datasource,
433439
chart_answer=row.chart_answer, chart=row.chart,
434440
analysis=row.analysis, predict=row.predict,
435441
datasource_select_answer=row.datasource_select_answer,
@@ -448,7 +454,8 @@ def get_chat_with_records(session: SessionDep, chart_id: int, current_user: Curr
448454
finish_time=row.finish_time,
449455
duration=duration,
450456
total_tokens=total_tokens,
451-
question=row.question, sql_answer=row.sql_answer, sql=row.sql, datasource=row.datasource,
457+
question=row.question, sql_answer=row.sql_answer, sql=row.sql,
458+
datasource=row.datasource,
452459
chart_answer=row.chart_answer, chart=row.chart,
453460
analysis=row.analysis, predict=row.predict,
454461
datasource_select_answer=row.datasource_select_answer,

backend/apps/db/db.py

Lines changed: 147 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@
1919
import dmPython
2020
import pymysql
2121
import redshift_connector
22-
from sqlalchemy import create_engine, text, Engine
22+
from sqlalchemy import create_engine, text, Engine, types
2323
from sqlalchemy.orm import sessionmaker
2424

2525
from 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+
720853
def 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
# 危险模式正则表达式(用于检查特殊语法)
746879
import re
880+
747881
DANGEROUS_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):

frontend/src/views/chat/chat-block/ChartBlock.vue

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ import DisplayChartBlock from '@/views/chat/component/DisplayChartBlock.vue'
44
import ChartPopover from '@/views/chat/chat-block/ChartPopover.vue'
55
import { computed, ref, watch } from 'vue'
66
import { useClipboard } from '@vueuse/core'
7-
import { concat } from 'lodash-es'
7+
import { concat, filter, includes, map } from 'lodash-es'
88
import type { ChartTypes } from '@/views/chat/component/BaseChart.ts'
99
import ICON_BAR from '@/assets/svg/chart/icon_bar_outlined.svg'
1010
import ICON_COLUMN from '@/assets/svg/chart/icon_dashboard_outlined.svg'
@@ -61,6 +61,7 @@ const emits = defineEmits(['exitFullScreen', 'update:thousandsSeparatorList', 'u
6161
6262
const dataObject = computed<{
6363
fields: Array<string>
64+
fields_info: Array<{ name: string; is_numeric: boolean }>
6465
data: Array<{ [key: string]: any }>
6566
limit: number | undefined
6667
datasource: number | undefined
@@ -388,7 +389,13 @@ const enableThousandsSeparatorList = computed({
388389
389390
const optionList = ref<Array<{ name: string; value: string }>>([])
390391
function getBaseAxis() {
391-
optionList.value = chartRef.value?.getBaseAxis()
392+
const _list = chartRef.value?.getBaseAxis()
393+
if (dataObject.value.fields_info) {
394+
const numberList = map(filter(dataObject.value.fields_info, { is_numeric: true }), 'name')
395+
optionList.value = filter(_list, (obj) => includes(numberList, obj.value))
396+
} else {
397+
optionList.value = _list
398+
}
392399
}
393400
</script>
394401

0 commit comments

Comments
 (0)