diff --git a/.gitignore b/.gitignore index fab542482..c53da76a2 100644 --- a/.gitignore +++ b/.gitignore @@ -189,3 +189,7 @@ test.py !/.venv/ + +sqlbot-xpack + +.claude diff --git a/Dockerfile b/Dockerfile index b1c1e9256..c8727b40a 100644 --- a/Dockerfile +++ b/Dockerfile @@ -95,6 +95,6 @@ EXPOSE 3000 8000 8001 5432 # Add health check HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \ - CMD curl -f http://localhost:8000 || exit 1 + CMD python3 -c "import os, urllib.request; urllib.request.urlopen(f'http://localhost:8000/{os.environ.get(\"CONTEXT_PATH\", \"\")}', timeout=3)" || exit 1 ENTRYPOINT ["sh", "start.sh"] diff --git a/README.md b/README.md index 9451714e1..c5a258844 100644 --- a/README.md +++ b/README.md @@ -83,7 +83,7 @@ docker run -d \ 如你有更多问题,可以加入我们的技术交流群与我们交流。 -contact_me_qr +contact_me_qr ## UI 展示 diff --git a/backend/alembic/env.py b/backend/alembic/env.py index 6b30c53e4..19e19fce6 100755 --- a/backend/alembic/env.py +++ b/backend/alembic/env.py @@ -25,8 +25,8 @@ # from apps.system.models.user import SQLModel # noqa # from apps.settings.models.setting_models import SQLModel #from apps.chat.models.chat_model import SQLModel -#from apps.terminology.models.terminology_model import SQLModel -#from apps.custom_prompt.models.custom_prompt_model import SQLModel +from apps.terminology.models.terminology_model import SQLModel +from sqlbot_xpack.custom_prompt.models.custom_prompt_model import SQLModel #from apps.data_training.models.data_training_model import SQLModel # from apps.dashboard.models.dashboard_model import SQLModel from common.core.config import settings # noqa diff --git a/backend/alembic/versions/067_ai_model_workspace_mapping.py b/backend/alembic/versions/067_ai_model_workspace_mapping.py new file mode 100644 index 000000000..aa558ba25 --- /dev/null +++ b/backend/alembic/versions/067_ai_model_workspace_mapping.py @@ -0,0 +1,34 @@ +"""067_ai_model_workspace_mapping + +Revision ID: e51127e9aa4a +Revises: 8adc3a4919be +Create Date: 2026-06-01 14:14:23.112843 + +""" +from alembic import op +import sqlalchemy as sa +import sqlmodel.sql.sqltypes +from sqlalchemy.dialects import postgresql + +# revision identifiers, used by Alembic. +revision = 'e51127e9aa4a' +down_revision = '8adc3a4919be' +branch_labels = None +depends_on = None + + +def upgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.create_table('ai_model_workspace_mapping', + sa.Column('id', sa.BigInteger(), nullable=False), + sa.Column('ai_model_id', sa.BigInteger(), nullable=True), + sa.Column('workspace_id', sa.BigInteger(), nullable=True), + sa.PrimaryKeyConstraint('id') + ) + # ### end Alembic commands ### + + +def downgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.drop_table('ai_model_workspace_mapping') + # ### end Alembic commands ### diff --git a/backend/alembic/versions/068_alter_sys_logs_user_agent_length.py b/backend/alembic/versions/068_alter_sys_logs_user_agent_length.py new file mode 100644 index 000000000..2ecbf966a --- /dev/null +++ b/backend/alembic/versions/068_alter_sys_logs_user_agent_length.py @@ -0,0 +1,29 @@ +"""068_alter_sys_logs_user_agent_length + +Revision ID: a1b2c3d4e5f6 +Revises: e51127e9aa4a +Create Date: 2026-06-09 00:00:00.000000 + +""" +from alembic import op +import sqlalchemy as sa + +# revision identifiers, used by Alembic. +revision = 'a1b2c3d4e5f6' +down_revision = 'e51127e9aa4a' +branch_labels = None +depends_on = None + + +def upgrade(): + op.alter_column('sys_logs', 'user_agent', + existing_type=sa.VARCHAR(length=255), + type_=sa.VARCHAR(length=500), + existing_nullable=True) + + +def downgrade(): + op.alter_column('sys_logs', 'user_agent', + existing_type=sa.VARCHAR(length=500), + type_=sa.VARCHAR(length=255), + existing_nullable=True) diff --git a/backend/alembic/versions/069_term_custom_prompt.py b/backend/alembic/versions/069_term_custom_prompt.py new file mode 100644 index 000000000..e02aaac37 --- /dev/null +++ b/backend/alembic/versions/069_term_custom_prompt.py @@ -0,0 +1,31 @@ +"""069_term_custom_prompt + +Revision ID: 1f82cad3546e +Revises: a1b2c3d4e5f6 +Create Date: 2026-06-15 14:51:12.280391 + +""" +from alembic import op +import sqlalchemy as sa +import sqlmodel.sql.sqltypes +from sqlalchemy.dialects import postgresql + +# revision identifiers, used by Alembic. +revision = '1f82cad3546e' +down_revision = 'a1b2c3d4e5f6' +branch_labels = None +depends_on = None + + +def upgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('custom_prompt', sa.Column('advanced_application', sa.BigInteger(), nullable=True)) + op.add_column('terminology', sa.Column('advanced_application', sa.BigInteger(), nullable=True)) + # ### end Alembic commands ### + + +def downgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.drop_column('terminology', 'advanced_application') + op.drop_column('custom_prompt', 'advanced_application') + # ### end Alembic commands ### diff --git a/backend/alembic/versions/070_add_table_column_comments.py b/backend/alembic/versions/070_add_table_column_comments.py new file mode 100644 index 000000000..2b0aa4e5a --- /dev/null +++ b/backend/alembic/versions/070_add_table_column_comments.py @@ -0,0 +1,440 @@ +"""070_add_table_column_comments + +为所有数据库表和字段添加中文备注(COMMENT ON)。 +Revision ID: 070a1b2c3d4e5 +Revises: 1f82cad3546e +Create Date: 2026-06-25 00:00:00.000000 + +""" +from alembic import op + +# revision identifiers, used by Alembic. +revision = '070a1b2c3d4e5' +down_revision = '1f82cad3546e' +branch_labels = None +depends_on = None + +# 表和列备注定义:key 为 "表名" 或 "表名.列名",value 为备注文本 +_COMMENTS = { + # ========== alembic_version 迁移版本记录表 ========== + 'alembic_version': 'Alembic迁移版本记录表', + 'alembic_version.version_num': '当前已应用的迁移版本号', + + # ========== sys_user 用户表 ========== + 'sys_user': '系统用户表', + 'sys_user.id': '用户ID', + 'sys_user.account': '用户账号', + 'sys_user.name': '用户名称', + 'sys_user.password': '用户密码', + 'sys_user.email': '用户邮箱', + 'sys_user.oid': '组织ID', + 'sys_user.status': '用户状态', + 'sys_user.origin': '用户来源', + 'sys_user.create_time': '创建时间', + 'sys_user.language': '用户语言偏好', + 'sys_user.system_variables': '用户自定义系统变量', + + # ========== sys_user_platform 用户平台关联表 ========== + 'sys_user_platform': '用户平台关联表', + 'sys_user_platform.id': '主键ID', + 'sys_user_platform.uid': '用户ID', + 'sys_user_platform.origin': '平台来源', + 'sys_user_platform.platform_uid': '平台用户ID', + + # ========== ai_model AI模型表 ========== + 'ai_model': 'AI模型配置表', + 'ai_model.id': '模型ID', + 'ai_model.supplier': '模型供应商', + 'ai_model.name': '模型名称', + 'ai_model.model_type': '模型类型', + 'ai_model.base_model': '基础模型标识', + 'ai_model.default_model': '是否默认模型', + 'ai_model.api_key': 'API密钥', + 'ai_model.api_domain': 'API域名地址', + 'ai_model.protocol': '通信协议', + 'ai_model.config': '模型配置信息', + 'ai_model.status': '模型状态', + 'ai_model.create_time': '创建时间', + + # ========== ai_model_workspace_mapping 模型工作空间映射表 ========== + 'ai_model_workspace_mapping': 'AI模型与工作空间映射表', + 'ai_model_workspace_mapping.id': '主键ID', + 'ai_model_workspace_mapping.ai_model_id': 'AI模型ID', + 'ai_model_workspace_mapping.workspace_id': '工作空间ID', + + # ========== sys_workspace 工作空间表 ========== + 'sys_workspace': '工作空间表', + 'sys_workspace.id': '工作空间ID', + 'sys_workspace.name': '工作空间名称', + 'sys_workspace.create_time': '创建时间', + + # ========== sys_user_ws 用户工作空间关联表 ========== + 'sys_user_ws': '用户工作空间权限表', + 'sys_user_ws.id': '主键ID', + 'sys_user_ws.uid': '用户ID', + 'sys_user_ws.oid': '组织ID', + 'sys_user_ws.weight': '权重值', + + # ========== sys_assistant AI助手表 ========== + 'sys_assistant': 'AI助手配置表', + 'sys_assistant.id': '助手ID', + 'sys_assistant.name': '助手名称', + 'sys_assistant.type': '助手类型', + 'sys_assistant.domain': '助手适用领域', + 'sys_assistant.description': '助手描述信息', + 'sys_assistant.configuration': '助手配置', + 'sys_assistant.create_time': '创建时间', + 'sys_assistant.app_id': '应用ID', + 'sys_assistant.app_secret': '应用密钥', + 'sys_assistant.oid': '组织ID', + 'sys_assistant.enable_custom_model': '是否启用自定义模型', + 'sys_assistant.custom_model': '自定义模型名称', + + # ========== sys_authentication 认证配置表 ========== + 'sys_authentication': '认证配置表', + 'sys_authentication.id': '认证配置ID', + 'sys_authentication.name': '认证配置名称', + 'sys_authentication.type': '认证类型', + 'sys_authentication.config': '认证配置详情', + 'sys_authentication.create_time': '创建时间', + 'sys_authentication.enable': '是否启用', + 'sys_authentication.valid': '配置是否有效', + + # ========== sys_apikey API密钥表 ========== + 'sys_apikey': 'API密钥表', + 'sys_apikey.id': '密钥ID', + 'sys_apikey.access_key': '访问密钥', + 'sys_apikey.secret_key': '密钥', + 'sys_apikey.create_time': '创建时间', + 'sys_apikey.uid': '绑定用户ID', + 'sys_apikey.status': '密钥状态', + + # ========== system_variable 系统变量表 ========== + 'system_variable': '系统变量表', + 'system_variable.id': '变量ID', + 'system_variable.name': '变量名称', + 'system_variable.var_type': '变量数据类型', + 'system_variable.type': '变量分类', + 'system_variable.value': '变量值', + 'system_variable.create_time': '创建时间', + 'system_variable.create_by': '创建人ID', + + # ========== terms 术语设置表 ========== + 'terms': '术语设置表', + 'terms.id': '术语ID', + 'terms.term': '术语名称', + 'terms.definition': '术语定义', + 'terms.domain': '所属领域', + 'terms.create_time': '创建时间', + + # ========== custom_prompt 自定义提示词表 ========== + 'custom_prompt': '自定义提示词表', + 'custom_prompt.id': '提示词ID', + 'custom_prompt.oid': '组织ID', + 'custom_prompt.type': '提示词类型', + 'custom_prompt.create_time': '创建时间', + 'custom_prompt.name': '提示词名称', + 'custom_prompt.prompt': '提示词内容', + 'custom_prompt.specific_ds': '是否关联特定数据源', + 'custom_prompt.datasource_ids': '关联数据源ID列表', + 'custom_prompt.advanced_application': '高级应用ID', + + # ========== core_datasource 数据源表 ========== + 'core_datasource': '数据源配置表', + 'core_datasource.id': '数据源ID', + 'core_datasource.name': '数据源名称', + 'core_datasource.description': '数据源描述', + 'core_datasource.type': '数据源类型', + 'core_datasource.type_name': '数据源类型名称', + 'core_datasource.configuration': '连接配置信息', + 'core_datasource.create_time': '创建时间', + 'core_datasource.create_by': '创建人ID', + 'core_datasource.status': '连接状态', + 'core_datasource.num': '数据源编号', + 'core_datasource.oid': '组织ID', + 'core_datasource.table_relation': '表关系配置', + 'core_datasource.embedding': '向量嵌入信息', + 'core_datasource.recommended_config': '推荐问题配置', + + # ========== core_table 数据源表信息 ========== + 'core_table': '数据源表信息', + 'core_table.id': '表记录ID', + 'core_table.ds_id': '所属数据源ID', + 'core_table.checked': '是否启用', + 'core_table.table_name': '表名称', + 'core_table.table_comment': '表注释', + 'core_table.custom_comment': '自定义表注释', + 'core_table.embedding': '表向量嵌入信息', + + # ========== ds_recommended_problem 推荐问题表 ========== + 'ds_recommended_problem': '数据源推荐问题表', + 'ds_recommended_problem.id': '问题ID', + 'ds_recommended_problem.datasource_id': '数据源ID', + 'ds_recommended_problem.question': '推荐问题内容', + 'ds_recommended_problem.remark': '问题备注', + 'ds_recommended_problem.sort': '排序序号', + 'ds_recommended_problem.create_time': '创建时间', + 'ds_recommended_problem.create_by': '创建人ID', + + # ========== core_field 数据源字段信息 ========== + 'core_field': '数据源字段信息表', + 'core_field.id': '字段记录ID', + 'core_field.ds_id': '所属数据源ID', + 'core_field.table_id': '所属表ID', + 'core_field.checked': '是否启用', + 'core_field.field_name': '字段名称', + 'core_field.field_type': '字段类型', + 'core_field.field_comment': '字段注释', + 'core_field.custom_comment': '自定义字段注释', + 'core_field.field_index': '字段序号', + + # ========== data_training 数据训练(示例库)表 ========== + 'data_training': 'SQL示例训练库表', + 'data_training.id': '训练记录ID', + 'data_training.oid': '组织ID', + 'data_training.datasource': '关联数据源ID', + 'data_training.create_time': '创建时间', + 'data_training.question': '训练问题', + 'data_training.description': '训练描述', + 'data_training.embedding': '向量嵌入数据', + 'data_training.enabled': '是否启用', + 'data_training.advanced_application': '高级应用ID', + + # ========== terminology 术语表 ========== + 'terminology': '术语管理表', + 'terminology.id': '术语ID', + 'terminology.oid': '组织ID', + 'terminology.pid': '父级术语ID', + 'terminology.create_time': '创建时间', + 'terminology.word': '术语词条', + 'terminology.description': '术语描述', + 'terminology.embedding': '向量嵌入数据', + 'terminology.specific_ds': '是否关联特定数据源', + 'terminology.datasource_ids': '关联数据源ID列表', + 'terminology.enabled': '是否启用', + 'terminology.advanced_application': '高级应用ID', + + # ========== core_dashboard 仪表板表 ========== + 'core_dashboard': '仪表板表', + 'core_dashboard.id': '仪表板ID', + 'core_dashboard.name': '仪表板名称', + 'core_dashboard.pid': '父级仪表板ID', + 'core_dashboard.workspace_id': '所属工作空间ID', + 'core_dashboard.org_id': '组织ID', + 'core_dashboard.level': '层级', + 'core_dashboard.node_type': '节点类型', + 'core_dashboard.type': '仪表板类型', + 'core_dashboard.canvas_style_data': '画布样式数据', + 'core_dashboard.component_data': '组件数据', + 'core_dashboard.canvas_view_info': '画布视图信息', + 'core_dashboard.mobile_layout': '是否移动端布局', + 'core_dashboard.status': '状态', + 'core_dashboard.self_watermark_status': '水印状态', + 'core_dashboard.sort': '排序序号', + 'core_dashboard.create_time': '创建时间', + 'core_dashboard.create_by': '创建人', + 'core_dashboard.update_time': '更新时间', + 'core_dashboard.update_by': '更新人', + 'core_dashboard.remark': '备注', + 'core_dashboard.source': '来源标识', + 'core_dashboard.delete_flag': '删除标记', + 'core_dashboard.delete_time': '删除时间', + 'core_dashboard.delete_by': '删除人', + 'core_dashboard.version': '版本号', + 'core_dashboard.content_id': '内容ID', + 'core_dashboard.check_version': '检查版本', + + # ========== chat_log 聊天日志表 ========== + 'chat_log': 'AI聊天执行日志表', + 'chat_log.id': '日志ID', + 'chat_log.type': '聊天类型', + 'chat_log.operate': '操作类型', + 'chat_log.pid': '父级日志ID', + 'chat_log.ai_modal_id': '使用的AI模型ID', + 'chat_log.base_modal': '基础模型标识', + 'chat_log.messages': '消息内容', + 'chat_log.reasoning_content': '推理内容', + 'chat_log.start_time': '开始时间', + 'chat_log.finish_time': '结束时间', + 'chat_log.token_usage': 'Token消耗统计', + 'chat_log.local_operation': '是否本地操作', + 'chat_log.error': '是否发生错误', + + # ========== chat 聊天主表 ========== + 'chat': 'AI聊天会话表', + 'chat.id': '聊天会话ID', + 'chat.oid': '组织ID', + 'chat.create_time': '创建时间', + 'chat.create_by': '创建人ID', + 'chat.brief': '会话摘要', + 'chat.chat_type': '聊天类型', + 'chat.datasource': '关联数据源ID', + 'chat.engine_type': '数据引擎类型', + 'chat.origin': '会话来源', + 'chat.brief_generate': '是否已生成摘要', + 'chat.recommended_question_answer': '推荐问题回答', + 'chat.recommended_question': '推荐问题内容', + 'chat.recommended_generate': '是否已生成推荐问题', + + # ========== chat_record 聊天记录表 ========== + 'chat_record': 'AI聊天记录表', + 'chat_record.id': '记录ID', + 'chat_record.chat_id': '关联聊天会话ID', + 'chat_record.ai_modal_id': '使用的AI模型ID', + 'chat_record.first_chat': '是否首轮对话', + 'chat_record.create_time': '创建时间', + 'chat_record.finish_time': '完成时间', + 'chat_record.create_by': '创建人ID', + 'chat_record.datasource': '关联数据源ID', + 'chat_record.engine_type': '数据引擎类型', + 'chat_record.question': '用户问题内容', + 'chat_record.sql_answer': 'SQL生成回答', + 'chat_record.sql': '生成的SQL语句', + 'chat_record.sql_exec_result': 'SQL执行结果', + 'chat_record.data': '查询返回数据', + 'chat_record.chart_answer': '图表生成回答', + 'chat_record.chart': '图表配置信息', + 'chat_record.analysis': '分析结果内容', + 'chat_record.predict': '预测结果内容', + 'chat_record.predict_data': '预测数据', + 'chat_record.recommended_question_answer': '推荐问题回答', + 'chat_record.recommended_question': '推荐问题列表', + 'chat_record.datasource_select_answer': '数据源选择回答', + 'chat_record.finish': '是否已完成', + 'chat_record.error': '错误信息', + 'chat_record.analysis_record_id': '关联分析记录ID', + 'chat_record.predict_record_id': '关联预测记录ID', + 'chat_record.regenerate_record_id': '关联重新生成记录ID', + + # ========== sys_logs 系统操作日志表 ========== + 'sys_logs': '系统操作日志表', + 'sys_logs.id': '日志ID', + 'sys_logs.operation_type': '操作类型', + 'sys_logs.operation_detail': '操作详情', + 'sys_logs.user_id': '操作用户ID', + 'sys_logs.operation_status': '操作状态', + 'sys_logs.ip_address': '操作IP地址', + 'sys_logs.user_agent': '用户代理信息', + 'sys_logs.execution_time': '执行耗时', + 'sys_logs.error_message': '错误信息', + 'sys_logs.create_time': '创建时间', + 'sys_logs.module': '操作模块', + 'sys_logs.oid': '组织ID', + 'sys_logs.resource_id': '操作资源ID', + 'sys_logs.request_method': '请求方法', + 'sys_logs.request_path': '请求路径', + 'sys_logs.remark': '备注', + 'sys_logs.user_name': '操作用户名称', + 'sys_logs.resource_name': '操作资源名称', + + # ========== sys_logs_resource 日志关联资源表 ========== + 'sys_logs_resource': '操作日志关联资源表', + 'sys_logs_resource.id': '主键ID', + 'sys_logs_resource.log_id': '关联日志ID', + 'sys_logs_resource.resource_id': '资源ID', + 'sys_logs_resource.resource_name': '资源名称', + 'sys_logs_resource.module': '资源模块', + + # ========== ds_permission 数据权限表 ========== + 'ds_permission': '数据权限配置表', + 'ds_permission.id': '权限ID', + 'ds_permission.enable': '是否启用', + 'ds_permission.name': '权限名称', + 'ds_permission.auth_target_type': '授权对象类型', + 'ds_permission.auth_target_id': '授权对象ID', + 'ds_permission.type': '权限类型', + 'ds_permission.ds_id': '数据源ID', + 'ds_permission.table_id': '表ID', + 'ds_permission.expression_tree': '权限表达式树', + 'ds_permission.permissions': '权限配置', + 'ds_permission.white_list_user': '白名单用户列表', + 'ds_permission.create_time': '创建时间', + + # ========== ds_rules 数据规则表 ========== + 'ds_rules': '数据规则组表', + 'ds_rules.id': '规则组ID', + 'ds_rules.enable': '是否启用', + 'ds_rules.name': '规则组名称', + 'ds_rules.description': '规则组描述', + 'ds_rules.permission_list': '权限列表', + 'ds_rules.user_list': '用户列表', + 'ds_rules.white_list_user': '白名单用户列表', + 'ds_rules.oid': '组织ID', + 'ds_rules.create_time': '创建时间', + + # ========== license 许可证表 ========== + 'license': '系统许可证表', + 'license.id': '许可证ID', + 'license.license_key': '许可证密钥', + 'license.f2c_license': 'F2C许可证信息', + 'license.create_time': '创建时间', + 'license.update_time': '更新时间', + + # ========== rsa 密钥表 ========== + 'rsa': 'RSA密钥存储表', + 'rsa.id': '密钥ID', + 'rsa.private_key': 'RSA私钥', + 'rsa.public_key': 'RSA公钥', + 'rsa.salt': '加密盐值', + 'rsa.create_time': '创建时间', + 'rsa.update_time': '更新时间', + + # ========== sys_arg 系统参数表 ========== + 'sys_arg': '系统参数配置表', + 'sys_arg.id': '参数ID', + 'sys_arg.pkey': '参数键', + 'sys_arg.pval': '参数值', + 'sys_arg.ptype': '参数类型', + 'sys_arg.sort_no': '排序序号', + + # ========== sys_platform_token 平台令牌表 ========== + 'sys_platform_token': '平台认证令牌表', + 'sys_platform_token.id': '令牌ID', + 'sys_platform_token.token': '令牌值', + 'sys_platform_token.create_time': '创建时间', + 'sys_platform_token.exp_time': '过期时间', +} + + +def _escape_comment(comment): + return comment.replace("'", "''") + + +def _apply_table_comment(table, comment): + if comment is None: + op.execute(f"COMMENT ON TABLE {table} IS NULL") + else: + op.execute(f"COMMENT ON TABLE {table} IS '{_escape_comment(comment)}'") + + +def _apply_column_comment(table, column, comment): + if comment is None: + op.execute(f"COMMENT ON COLUMN {table}.{column} IS NULL") + else: + op.execute(f"COMMENT ON COLUMN {table}.{column} IS '{_escape_comment(comment)}'") + + +def upgrade(): + """为所有表和字段添加中文备注""" + table_keys = sorted(k for k in _COMMENTS if '.' not in k) + column_keys = sorted(k for k in _COMMENTS if '.' in k) + + for table in table_keys: + _apply_table_comment(table, _COMMENTS[table]) + + for key in column_keys: + table, column = key.split('.', 1) + _apply_column_comment(table, column, _COMMENTS[key]) + + +def downgrade(): + """移除所有表和字段的备注""" + table_keys = sorted((k for k in _COMMENTS if '.' not in k), reverse=True) + column_keys = sorted((k for k in _COMMENTS if '.' in k), reverse=True) + + for key in column_keys: + table, column = key.split('.', 1) + _apply_column_comment(table, column, None) + + for table in table_keys: + _apply_table_comment(table, None) diff --git a/backend/alembic/versions/071_modify_permission_jsonb.py b/backend/alembic/versions/071_modify_permission_jsonb.py new file mode 100644 index 000000000..cac4c64b9 --- /dev/null +++ b/backend/alembic/versions/071_modify_permission_jsonb.py @@ -0,0 +1,45 @@ +"""071_modify_permission_jsonb + +Revision ID: a2e2ecfa5a9c +Revises: 070a1b2c3d4e5 +Create Date: 2026-07-16 16:56:47.093944 + +""" +from alembic import op +import sqlalchemy as sa +import sqlmodel.sql.sqltypes +from sqlalchemy.dialects import postgresql + +# revision identifiers, used by Alembic. +revision = 'a2e2ecfa5a9c' +down_revision = '070a1b2c3d4e5' +branch_labels = None +depends_on = None + + +def upgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.alter_column('ds_rules', 'permission_list', + existing_type=sa.Text(), + type_=postgresql.JSONB(astext_type=sa.Text()), + existing_nullable=True, + postgresql_using='permission_list::jsonb') + op.alter_column('ds_rules', 'user_list', + existing_type=sa.Text(), + type_=postgresql.JSONB(astext_type=sa.Text()), + existing_nullable=True, + postgresql_using='user_list::jsonb') + # ### end Alembic commands ### + + +def downgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.alter_column('ds_rules', 'permission_list', + existing_type=postgresql.JSONB(astext_type=sa.Text()), + type_=sa.Text(), + existing_nullable=True) + op.alter_column('ds_rules', 'user_list', + existing_type=postgresql.JSONB(astext_type=sa.Text()), + type_=sa.Text(), + existing_nullable=True) + # ### end Alembic commands ### diff --git a/backend/apps/ai_model/openai/llm.py b/backend/apps/ai_model/openai/llm.py index a03693c66..603fd50ea 100644 --- a/backend/apps/ai_model/openai/llm.py +++ b/backend/apps/ai_model/openai/llm.py @@ -91,19 +91,20 @@ def _default_params(self) -> dict[str, Any]: if max_tokens: params["max_tokens"] = max_tokens return params - + def _get_request_payload( - self, - input_: LanguageModelInput, - *, - stop: Optional[list[str]] = None, - **kwargs: Any, + self, + input_: LanguageModelInput, + *, + stop: Optional[list[str]] = None, + **kwargs: Any, ) -> dict: max_tokens = self.max_tokens payload = super()._get_request_payload(input_, stop=stop, **kwargs) if max_tokens: payload["max_tokens"] = max_tokens return payload + usage_metadata: dict = {} # custom_get_token_ids = custom_get_token_ids @@ -119,10 +120,10 @@ def _stream(self, *args: Any, **kwargs: Any) -> Iterator[ChatGenerationChunk]: yield chunk def _convert_chunk_to_generation_chunk( - self, - chunk: dict, - default_chunk_class: type, - base_generation_info: dict | None, + self, + chunk: dict, + default_chunk_class: type, + base_generation_info: dict | None, ) -> ChatGenerationChunk | None: if chunk.get("type") == "content.delta": # from beta.chat.completions.stream return None @@ -174,12 +175,12 @@ def _convert_chunk_to_generation_chunk( return generation_chunk def invoke( - self, - input: LanguageModelInput, - config: RunnableConfig | None = None, - *, - stop: list[str] | None = None, - **kwargs: Any, + self, + input: LanguageModelInput, + config: RunnableConfig | None = None, + *, + stop: list[str] | None = None, + **kwargs: Any, ) -> BaseMessage: config = ensure_config(config) chat_result = cast( diff --git a/backend/apps/chat/api/chat.py b/backend/apps/chat/api/chat.py index 0e6ff7ee7..8a2bed721 100644 --- a/backend/apps/chat/api/chat.py +++ b/backend/apps/chat/api/chat.py @@ -11,20 +11,20 @@ from starlette.responses import JSONResponse from apps.chat.curd.chat import delete_chat_with_user, get_chart_data_with_user, get_chat_predict_data_with_user, \ - list_chats, get_chat_with_records, create_chat, rename_chat, \ - delete_chat, get_chat_chart_data, get_chat_predict_data, get_chat_with_records_with_data, get_chat_record_by_id, \ - format_json_data, format_json_list_data, get_chart_config, list_recent_questions, get_chat as get_chat_exec, \ - rename_chat_with_user, get_chat_log_history, get_chart_data_with_user_live + list_chats, get_chat_with_records, create_chat, get_chat_chart_data, get_chat_predict_data, \ + get_chat_with_records_with_data, get_chat_record_by_id, \ + format_json_data, format_json_list_data, get_chart_config, list_recent_questions, rename_chat_with_user, \ + get_chat_log_history, get_chart_data_with_user_live from apps.chat.models.chat_model import CreateChat, ChatRecord, RenameChat, ChatQuestion, AxisObj, QuickCommand, \ - ChatInfo, Chat, ChatFinishStep + ChatInfo, Chat, ChatFinishStep, ChatQuestionBase, SimpleChat from apps.chat.task.llm import LLMService from apps.swagger.i18n import PLACEHOLDER_PREFIX from apps.system.schemas.permission import SqlbotPermission, require_permissions +from common.audit.models.log_model import OperationType, OperationModules +from common.audit.schemas.logger_decorator import LogConfig, system_log from common.core.deps import CurrentAssistant, SessionDep, CurrentUser, Trans from common.utils.command_utils import parse_quick_command from common.utils.data_format import DataFormat -from common.audit.models.log_model import OperationType, OperationModules -from common.audit.schemas.logger_decorator import LogConfig, system_log router = APIRouter(tags=["Data Q&A"], prefix="/chat") @@ -116,22 +116,6 @@ def inner(): return await asyncio.to_thread(inner) -""" @router.post("/rename", response_model=str, summary=f"{PLACEHOLDER_PREFIX}rename_chat") -@system_log(LogConfig( - operation_type=OperationType.UPDATE, - module=OperationModules.CHAT, - resource_id_expr="chat.id" -)) -async def rename(session: SessionDep, chat: RenameChat): - try: - return rename_chat(session=session, rename_object=chat) - except Exception as e: - raise HTTPException( - status_code=500, - detail=str(e) - ) """ - - @router.post("/rename", response_model=str, summary=f"{PLACEHOLDER_PREFIX}rename_chat") @system_log(LogConfig( operation_type=OperationType.UPDATE, @@ -148,31 +132,14 @@ async def rename(session: SessionDep, current_user: CurrentUser, chat: RenameCha ) -""" @router.delete("/{chart_id}/{brief}", response_model=str, summary=f"{PLACEHOLDER_PREFIX}delete_chat") -@system_log(LogConfig( - operation_type=OperationType.DELETE, - module=OperationModules.CHAT, - resource_id_expr="chart_id", - remark_expr="brief" -)) -async def delete(session: SessionDep, chart_id: int, brief: str): - try: - return delete_chat(session=session, chart_id=chart_id) - except Exception as e: - raise HTTPException( - status_code=500, - detail=str(e) - ) """ - - -@router.delete("/{chart_id}/{brief}", response_model=str, summary=f"{PLACEHOLDER_PREFIX}delete_chat") +@router.delete("/{chart_id}", response_model=str, summary=f"{PLACEHOLDER_PREFIX}delete_chat") @system_log(LogConfig( operation_type=OperationType.DELETE, module=OperationModules.CHAT, resource_id_expr="chart_id", - remark_expr="brief" + remark_expr="chat.brief" )) -async def delete(session: SessionDep, current_user: CurrentUser, chart_id: int, brief: str): +async def delete(session: SessionDep, current_user: CurrentUser, chart_id: int, chat: SimpleChat): try: return delete_chat_with_user(session=session, current_user=current_user, chart_id=chart_id) except Exception as e: @@ -269,9 +236,10 @@ def find_base_question(record_id: int, session: SessionDep): @router.post("/question", summary=f"{PLACEHOLDER_PREFIX}ask_question") @require_permissions(permission=SqlbotPermission(type='chat', keyExpression="request_question.chat_id")) -async def question_answer(session: SessionDep, current_user: CurrentUser, request_question: ChatQuestion, +async def question_answer(session: SessionDep, current_user: CurrentUser, request_question: ChatQuestionBase, current_assistant: CurrentAssistant): - return await question_answer_inner(session, current_user, request_question, current_assistant, embedding=True) + question = ChatQuestion(chat_id=request_question.chat_id, question=request_question.question) + return await question_answer_inner(session, current_user, question, current_assistant, embedding=True) async def question_answer_inner(session: SessionDep, current_user: CurrentUser, request_question: ChatQuestion, diff --git a/backend/apps/chat/curd/chat.py b/backend/apps/chat/curd/chat.py index eaddc671e..e261c2b6f 100644 --- a/backend/apps/chat/curd/chat.py +++ b/backend/apps/chat/curd/chat.py @@ -1,4 +1,5 @@ import datetime +from decimal import Decimal from typing import List, Optional, Union, Dict, Any import orjson @@ -186,7 +187,8 @@ def get_last_execute_sql_error(session: SessionDep, chart_id: int): def format_json_data(origin_data: dict): - result = {'fields': origin_data.get('fields') if origin_data.get('fields') else []} + result = {'fields': origin_data.get('fields') if origin_data.get('fields') else [], + 'fields_info': origin_data.get('fields_info') if origin_data.get('fields_info') else None} _list = origin_data.get('data') if origin_data.get('data') else [] data = format_json_list_data(_list) result['data'] = data @@ -207,7 +209,7 @@ def format_json_list_data(origin_data: list[dict]): value = str(value) # 小数且超过15位有效数字 → 转字符串并标记为文本列 elif isinstance(value, float): - decimal_str = format(value, '.16f').rstrip('0').rstrip('.') + decimal_str = str(Decimal(str(value))).rstrip('0').rstrip('.') if len(decimal_str) > 15: value = str(value) _row[key] = value @@ -237,21 +239,24 @@ def get_chart_data_with_user(session: SessionDep, current_user: CurrentUser, cha pass return {} + def get_chart_data_with_user_live(session: SessionDep, current_user: CurrentUser, chat_record_id: int): - stmt = select(ChatRecord.datasource,ChatRecord.sql).where(and_(ChatRecord.id == chat_record_id, ChatRecord.create_by == current_user.id)) + stmt = select(ChatRecord.datasource, ChatRecord.sql).where( + and_(ChatRecord.id == chat_record_id, ChatRecord.create_by == current_user.id)) row = session.execute(stmt).first() - return get_chart_data_ds(session,row.datasource, row.sql) + return get_chart_data_ds(session, row.datasource, row.sql) -def get_chart_data_ds(session: SessionDep,ds_id,sql): - json_result: Dict[str, Any] = {'status': 'success','data':[],'message':''} + +def get_chart_data_ds(session: SessionDep, ds_id, sql): + json_result: Dict[str, Any] = {'status': 'success', 'data': [], 'message': ''} try: - datasource = get_ds(session,ds_id) + datasource = get_ds(session, ds_id) if datasource is None: json_result['status'] = 'failed' json_result['message'] = 'Datasource not found' return json_result else: - result = exec_sql(ds=datasource,sql=sql, origin_column=False) + result = exec_sql(ds=datasource, sql=sql, origin_column=False) _data = DataFormat.convert_large_numbers_in_object_array(result.get('data')) _data = DataFormat.normalize_qualified_sql_column_keys_in_object_array(_data) json_result['data'] = _data @@ -263,6 +268,7 @@ def get_chart_data_ds(session: SessionDep,ds_id,sql): pass return json_result + def get_chat_chart_data(session: SessionDep, chat_record_id: int): stmt = select(ChatRecord.data).where(and_(ChatRecord.id == chat_record_id)) res = session.execute(stmt) @@ -335,7 +341,7 @@ def get_chat_with_records(session: SessionDep, chart_id: int, current_user: Curr predict_alias_log = aliased(ChatLog) stmt = (select(ChatRecord.id, ChatRecord.chat_id, ChatRecord.create_time, ChatRecord.finish_time, - ChatRecord.question, ChatRecord.sql_answer, ChatRecord.sql,ChatRecord.datasource, + ChatRecord.question, ChatRecord.sql_answer, ChatRecord.sql, ChatRecord.datasource, ChatRecord.chart_answer, ChatRecord.chart, ChatRecord.analysis, ChatRecord.predict, ChatRecord.datasource_select_answer, ChatRecord.analysis_record_id, ChatRecord.predict_record_id, ChatRecord.regenerate_record_id, @@ -362,7 +368,7 @@ def get_chat_with_records(session: SessionDep, chart_id: int, current_user: Curr ChatRecord.create_time)) if with_data: stmt = select(ChatRecord.id, ChatRecord.chat_id, ChatRecord.create_time, ChatRecord.finish_time, - ChatRecord.question, ChatRecord.sql_answer, ChatRecord.sql,ChatRecord.datasource, + ChatRecord.question, ChatRecord.sql_answer, ChatRecord.sql, ChatRecord.datasource, ChatRecord.chart_answer, ChatRecord.chart, ChatRecord.analysis, ChatRecord.predict, ChatRecord.datasource_select_answer, ChatRecord.analysis_record_id, ChatRecord.predict_record_id, ChatRecord.regenerate_record_id, @@ -428,7 +434,8 @@ def get_chat_with_records(session: SessionDep, chart_id: int, current_user: Curr finish_time=row.finish_time, duration=duration, total_tokens=total_tokens, - question=row.question, sql_answer=row.sql_answer, sql=row.sql, datasource=row.datasource, + question=row.question, sql_answer=row.sql_answer, sql=row.sql, + datasource=row.datasource, chart_answer=row.chart_answer, chart=row.chart, analysis=row.analysis, predict=row.predict, datasource_select_answer=row.datasource_select_answer, @@ -447,7 +454,8 @@ def get_chat_with_records(session: SessionDep, chart_id: int, current_user: Curr finish_time=row.finish_time, duration=duration, total_tokens=total_tokens, - question=row.question, sql_answer=row.sql_answer, sql=row.sql, datasource=row.datasource, + question=row.question, sql_answer=row.sql_answer, sql=row.sql, + datasource=row.datasource, chart_answer=row.chart_answer, chart=row.chart, analysis=row.analysis, predict=row.predict, datasource_select_answer=row.datasource_select_answer, diff --git a/backend/apps/chat/models/chat_model.py b/backend/apps/chat/models/chat_model.py index 6669e3f5f..6e550afb5 100644 --- a/backend/apps/chat/models/chat_model.py +++ b/backend/apps/chat/models/chat_model.py @@ -175,6 +175,9 @@ class RenameChat(BaseModel): brief: str = '' brief_generate: bool = True +class SimpleChat(BaseModel): + id: int = None + brief: str = '' class ChatInfo(BaseModel): id: Optional[int] = None @@ -250,16 +253,18 @@ def sql_sys_question(self, db_type: Union[str, DB], enable_query_limit: bool = T _example_answer_3 = _sql_template['example_answer_3_with_limit'] if enable_query_limit else _sql_template[ 'example_answer_3'] - templates['system'] = _base_template['system'].format(lang=self.lang, process_check=_process_check, sqlbot_name=self.sqlbot_name) + templates['system'] = _base_template['system'].format(lang=self.lang, process_check=_process_check, + sqlbot_name=self.sqlbot_name) templates['rules'] = _base_template['generate_rules'].format(lang=self.lang, - sqlbot_name = self.sqlbot_name, + sqlbot_name=self.sqlbot_name, base_sql_rules=_base_sql_rules, basic_sql_examples=_sql_examples, example_engine=_example_engine, example_answer_1=_example_answer_1, example_answer_2=_example_answer_2, example_answer_3=_example_answer_3) - templates['schema'] = _base_template['generate_basic_info'].format(engine=self.engine, schema=self.db_schema, sample_data=self.sample_data) + templates['schema'] = _base_template['generate_basic_info'].format(engine=self.engine, schema=self.db_schema, + sample_data=self.sample_data) if self.terminologies: templates['terminologies'] = _base_template['generate_terminologies_info'].format( @@ -303,7 +308,8 @@ def analysis_user_question(self): return get_analysis_template()['user'].format(fields=self.fields, data=self.data) def predict_sys_question(self): - return get_predict_template()['system'].format(lang=self.lang, custom_prompt=self.custom_prompt, sqlbot_name=self.sqlbot_name) + return get_predict_template()['system'].format(lang=self.lang, custom_prompt=self.custom_prompt, + sqlbot_name=self.sqlbot_name) def predict_user_question(self): return get_predict_template()['user'].format(fields=self.fields, data=self.data) @@ -315,14 +321,16 @@ def datasource_user_question(self, datasource_list: str = "[]"): return get_datasource_template()['user'].format(lang=self.lang, question=self.question, data=datasource_list) def guess_sys_question(self, articles_number: int = 4): - return get_guess_question_template()['system'].format(lang=self.lang, articles_number=articles_number, sqlbot_name=self.sqlbot_name) + return get_guess_question_template()['system'].format(lang=self.lang, articles_number=articles_number, + sqlbot_name=self.sqlbot_name) def guess_user_question(self, old_questions: str = "[]"): return get_guess_question_template()['user'].format(question=self.question, schema=self.db_schema, old_questions=old_questions) def filter_sys_question(self): - return get_permissions_template()['system'].format(lang=self.lang, engine=self.engine, sqlbot_name=self.sqlbot_name) + return get_permissions_template()['system'].format(lang=self.lang, engine=self.engine, + sqlbot_name=self.sqlbot_name) def filter_user_question(self): return get_permissions_template()['user'].format(sql=self.sql, filter=self.filter) @@ -348,20 +356,29 @@ class McpDs(BaseModel): oid: Optional[str] = Body(description='组织ID,如果不传则为最后一次登录SQLBot时所使用的组织ID', default=None) -class ChatStart(BaseModel): +class ChatToken(BaseModel): username: str = Body(description='用户名') password: str = Body(description='密码') -class McpQuestion(BaseModel): +class ChatStart(BaseModel): + username: str = Body(description='用户名', default=None) + password: str = Body(description='密码', default=None) + token: str = Body(description='token', default=None) + oid: Optional[str] = Body( + description='组织ID,仅当数据源ID为空时有效,如果不传则为最后一次登录SQLBot时所使用的组织ID', default=None) + + +class ChatQuestionBase(BaseModel): question: str = Body(description='用户提问') chat_id: int = Body(description='会话ID') + + +class McpQuestion(ChatQuestionBase): token: str = Body(description='token') stream: Optional[bool] = Body(description='是否流式输出,默认为true开启, 关闭false则返回JSON对象', default=True) lang: Optional[str] = Body(description='语言:zh-CN|zh-TW|en|ko-KR', default='zh-CN') datasource_id: Optional[int | str] = Body(description='数据源ID,仅当当前对话没有确定数据源时有效', default=None) - oid: Optional[str] = Body( - description='组织ID,仅当数据源ID为空时有效,如果不传则为最后一次登录SQLBot时所使用的组织ID', default=None) return_img: Optional[bool] = Body(description='是否返回图表,默认为true开启, 关闭false则仅返回数据', default=True) diff --git a/backend/apps/chat/task/llm.py b/backend/apps/chat/task/llm.py index 85d89dbaf..e9ef80420 100644 --- a/backend/apps/chat/task/llm.py +++ b/backend/apps/chat/task/llm.py @@ -6,12 +6,12 @@ import warnings from concurrent.futures import ThreadPoolExecutor, Future from datetime import datetime -from dis import specialized from typing import Any, List, Optional, Union, Dict, Iterator import orjson import pandas as pd import requests +import sqlglot import sqlparse from langchain.chat_models.base import BaseChatModel from langchain_community.utilities import SQLDatabase @@ -22,6 +22,7 @@ from sqlbot_xpack.custom_prompt.curd.custom_prompt import find_custom_prompts from sqlbot_xpack.custom_prompt.models.custom_prompt_model import CustomPromptTypeEnum from sqlbot_xpack.license.license_manage import SQLBotLicenseUtil +from sqlglot import exp from sqlmodel import Session from apps.ai_model.model_factory import LLMConfig, LLMFactory, get_default_config @@ -40,9 +41,11 @@ from apps.datasource.crud.permission import get_row_permission_filters, is_normal_user from apps.datasource.embedding.ds_embedding import get_ds_embedding from apps.datasource.models.datasource import CoreDatasource -from apps.db.db import exec_sql, get_version, check_connection +from apps.db.db import exec_sql, get_version, check_connection, get_sqlglot_dialect +from apps.system.crud.aimodel_manage import get_ai_model_list_by_workspace from apps.system.crud.assistant import AssistantOutDs, AssistantOutDsFactory, get_assistant_ds from apps.system.crud.parameter_manage import get_groups +from apps.system.crud.user import user_ws_list from apps.system.schemas.system_schema import AssistantOutDsSchema from apps.terminology.curd.terminology import get_terminology_template from common.core.config import settings @@ -65,9 +68,31 @@ i18n = I18n() +def extract_tables_from_sql(sql: str, ds_type: str = None) -> set: + """从 SQL 中提取真实表名(排除 CTE 别名)""" + tables = set() + dialect = get_sqlglot_dialect(ds_type) + try: + statements = sqlglot.parse(sql, dialect=dialect) + for stmt in statements: + if stmt: + # 收集 CTE 别名,排除嵌套 CTE + cte_names = set() + for cte in stmt.find_all(exp.CTE): + if cte.alias: + cte_names.add(cte.alias) + for table in stmt.find_all(exp.Table): + if table.name and table.name not in cte_names: + tables.add(table.name) + except Exception: + pass + return tables + + class LLMService: ds: CoreDatasource chat_question: ChatQuestion + oid: int record: ChatRecord config: LLMConfig llm: BaseChatModel @@ -105,15 +130,33 @@ def __init__(self, session: Session, current_user: CurrentUser, chat_question: C self.chunk_list = [] self.current_user = current_user self.current_assistant = current_assistant + + self.table_name_list = [] + + chat_question.lang = get_lang_name(current_user.language) + self.trans = i18n(lang=current_user.language) + chat_id = chat_question.chat_id chat: Chat | None = session.get(Chat, chat_id) if not chat: raise SingleMessageError(f"Chat with id {chat_id} not found") + self.oid = chat.oid + + if self.oid and not current_assistant: + w_list = user_ws_list(session, self.current_user.id) + oid_list = [item.id for item in w_list] + if int(self.oid) not in oid_list: + raise SingleMessageError("Current user cannot not access this chat") + if self.oid and current_assistant: + if self.oid != self.current_user.oid: + raise SingleMessageError("Current assistant user cannot not access this chat") + + self.current_user.oid = chat.oid ds: CoreDatasource | AssistantOutDsSchema | None = None if not chat.datasource and chat_question.datasource_id: _ds = session.get(CoreDatasource, chat_question.datasource_id) if _ds: - if _ds.oid != current_user.oid: + if _ds.oid != self.oid: raise SingleMessageError( f"Datasource with id {chat_question.datasource_id} does not belong to current workspace") chat.datasource = _ds.id @@ -143,9 +186,6 @@ def __init__(self, session: Session, current_user: CurrentUser, chat_question: C self.change_title = not get_chat_brief_generate(session=session, chat_id=chat_id) - chat_question.lang = get_lang_name(current_user.language) - self.trans = i18n(lang=current_user.language) - self.ds = ( ds if isinstance(ds, AssistantOutDsSchema) else CoreDatasource(**ds.model_dump())) if ds else None self.chat_question = chat_question @@ -176,11 +216,16 @@ def __init__(self, session: Session, current_user: CurrentUser, chat_question: C @classmethod async def create(cls, *args, **kwargs): specialized_model_id = None + _ai_model_list = [] if args[3]: + if args[1]: + ws_id = args[1].oid + _ai_model_list = get_ai_model_list_by_workspace(args[0], ws_id) if args[3].enable_custom_model: if args[3].custom_model: - specialized_model_id = args[3].custom_model - print("use custom model: id[" + args[3].custom_model + "]") + if any(str(model.id) == str(args[3].custom_model) for model in _ai_model_list): + specialized_model_id = args[3].custom_model + print("use custom model: id[" + specialized_model_id + "]") config: LLMConfig = await get_default_config(specialized_model_id) instance = cls(*args, **kwargs, config=config) @@ -216,7 +261,7 @@ def is_running(self, timeout=0.5): def init_messages(self, session: Session): - self.choose_table_schema(session) + self.table_name_list = self.choose_table_schema(session) last_sql_messages: List[dict[str, Any]] = self.generate_sql_logs[-1].messages if len( self.generate_sql_logs) > 0 else [] @@ -312,18 +357,26 @@ def get_fields_from_chart(self, _session: Session): return format_chart_fields(chart_info) def filter_terminology_template(self, _session: Session, oid: int = None, ds_id: int = None): + self.current_logs[OperationEnum.FILTER_TERMS] = start_log(session=_session, + operate=OperationEnum.FILTER_TERMS, + record_id=self.record.id, local_operation=True) calculate_oid = oid calculate_ds_id = ds_id if self.current_assistant: - calculate_oid = self.current_assistant.oid if self.current_assistant.type != 4 else self.current_user.oid + calculate_oid = self.current_assistant.oid if self.current_assistant.type != 4 else self.oid if self.current_assistant.type == 1: calculate_ds_id = None - self.current_logs[OperationEnum.FILTER_TERMS] = start_log(session=_session, - operate=OperationEnum.FILTER_TERMS, - record_id=self.record.id, local_operation=True) + if self.current_assistant and self.current_assistant.type == 1: + self.chat_question.terminologies, term_list = get_terminology_template(_session, + self.chat_question.question, + calculate_oid, + None, self.current_assistant.id) + else: + self.chat_question.terminologies, term_list = get_terminology_template(_session, + self.chat_question.question, + calculate_oid, + calculate_ds_id) - self.chat_question.terminologies, term_list = get_terminology_template(_session, self.chat_question.question, - calculate_oid, calculate_ds_id) self.current_logs[OperationEnum.FILTER_TERMS] = end_log(session=_session, log=self.current_logs[OperationEnum.FILTER_TERMS], full_message=term_list) @@ -331,19 +384,26 @@ def filter_terminology_template(self, _session: Session, oid: int = None, ds_id: def filter_custom_prompts(self, _session: Session, custom_prompt_type: CustomPromptTypeEnum, oid: int = None, ds_id: int = None): if SQLBotLicenseUtil.valid(): + self.current_logs[OperationEnum.FILTER_CUSTOM_PROMPT] = start_log(session=_session, + operate=OperationEnum.FILTER_CUSTOM_PROMPT, + record_id=self.record.id, + local_operation=True) calculate_oid = oid calculate_ds_id = ds_id if self.current_assistant: - calculate_oid = self.current_assistant.oid if self.current_assistant.type != 4 else self.current_user.oid + calculate_oid = self.current_assistant.oid if self.current_assistant.type != 4 else self.oid if self.current_assistant.type == 1: calculate_ds_id = None - self.current_logs[OperationEnum.FILTER_CUSTOM_PROMPT] = start_log(session=_session, - operate=OperationEnum.FILTER_CUSTOM_PROMPT, - record_id=self.record.id, - local_operation=True) - self.chat_question.custom_prompt, prompt_list = find_custom_prompts(_session, custom_prompt_type, - calculate_oid, - calculate_ds_id) + if self.current_assistant and self.current_assistant.type == 1: + self.chat_question.custom_prompt, prompt_list = find_custom_prompts(_session, + custom_prompt_type, + calculate_oid, + None, self.current_assistant.id) + else: + self.chat_question.custom_prompt, prompt_list = find_custom_prompts(_session, + custom_prompt_type, + calculate_oid, + calculate_ds_id) self.current_logs[OperationEnum.FILTER_CUSTOM_PROMPT] = end_log(session=_session, log=self.current_logs[ OperationEnum.FILTER_CUSTOM_PROMPT], @@ -357,7 +417,7 @@ def filter_training_template(self, _session: Session, oid: int = None, ds_id: in calculate_oid = oid calculate_ds_id = ds_id if self.current_assistant: - calculate_oid = self.current_assistant.oid if self.current_assistant.type != 4 else self.current_user.oid + calculate_oid = self.current_assistant.oid if self.current_assistant.type != 4 else self.oid if self.current_assistant.type == 1: calculate_ds_id = None if self.current_assistant and self.current_assistant.type == 1: @@ -398,6 +458,7 @@ def choose_table_schema(self, _session: Session): self.current_logs[OperationEnum.CHOOSE_TABLE] = end_log(session=_session, log=self.current_logs[OperationEnum.CHOOSE_TABLE], full_message=self.chat_question.db_schema) + return tables def generate_analysis(self, _session: Session): fields = self.get_fields_from_chart(_session) @@ -408,9 +469,9 @@ def generate_analysis(self, _session: Session): ds_id = self.ds.id if isinstance(self.ds, CoreDatasource) else None - self.filter_terminology_template(_session, self.current_user.oid, ds_id) + self.filter_terminology_template(_session, self.oid, ds_id) - self.filter_custom_prompts(_session, CustomPromptTypeEnum.ANALYSIS, self.current_user.oid, ds_id) + self.filter_custom_prompts(_session, CustomPromptTypeEnum.ANALYSIS, self.oid, ds_id) analysis_msg.append(SystemPromptMessage(content=self.chat_question.analysis_sys_question())) analysis_msg.append(HumanMessage(content=self.chat_question.analysis_user_question())) @@ -461,7 +522,7 @@ def generate_predict(self, _session: Session): self.chat_question.data = orjson.dumps(data.get('data')).decode() ds_id = self.ds.id if isinstance(self.ds, CoreDatasource) else None - self.filter_custom_prompts(_session, CustomPromptTypeEnum.PREDICT_DATA, self.current_user.oid, ds_id) + self.filter_custom_prompts(_session, CustomPromptTypeEnum.PREDICT_DATA, self.oid, ds_id) predict_msg: List[Union[BaseMessage, dict[str, Any]]] = [] predict_msg.append(SystemPromptMessage(content=self.chat_question.predict_sys_question())) @@ -581,7 +642,7 @@ def select_datasource(self, _session: Session): _ds_list = get_assistant_ds(session=_session, llm_service=self) else: stmt = select(CoreDatasource.id, CoreDatasource.name, CoreDatasource.description).where( - and_(CoreDatasource.oid == self.current_user.oid)) + and_(CoreDatasource.oid == self.oid)) _ds_list = [ { "id": ds.id, @@ -600,7 +661,7 @@ def select_datasource(self, _session: Session): if not ignore_auto_select: if settings.TABLE_EMBEDDING_ENABLED and ( not self.current_assistant or (self.current_assistant and self.current_assistant.type != 1)): - _ds_list = get_ds_embedding(_session, self.current_user, _ds_list, self.out_ds_instance, + _ds_list = get_ds_embedding(_session, _ds_list, self.out_ds_instance, self.chat_question.question, self.current_assistant) # yield {'content': '{"id":' + str(ds.get('id')) + '}'} @@ -1260,6 +1321,22 @@ def run_task(self, in_chat: bool = True, stream: bool = True, sql_operate = OperationEnum.GENERATE_SQL sql, tables = self.check_sql(session=_session, res=full_sql_text, operate=sql_operate) + + # 表名安全检查:用 sqlglot 解析真实 SQL,不信任 AI 返回的 tables + actual_tables = extract_tables_from_sql(sql, ds_type=self.ds.type) + if not actual_tables: + raise SingleMessageError( + "SQL parsing failed: unable to extract table names. " + "This may indicate an unsupported SQL syntax or a security issue." + ) + allowed_tables = set(self.table_name_list) + unauthorized_tables = actual_tables - allowed_tables + if unauthorized_tables: + raise SingleMessageError( + f"SQL contains unauthorized tables: {', '.join(unauthorized_tables)}. " + f"Allowed tables: {', '.join(allowed_tables)}" + ) + if ((not self.current_assistant or is_page_embedded) and is_normal_user( self.current_user)) or use_dynamic_ds: sql_result = None diff --git a/backend/apps/data_training/api/data_training.py b/backend/apps/data_training/api/data_training.py index e10c8164d..d3e662cfb 100644 --- a/backend/apps/data_training/api/data_training.py +++ b/backend/apps/data_training/api/data_training.py @@ -45,7 +45,8 @@ async def pager(session: SessionDep, current_user: CurrentUser, current_page: in @router.put("", response_model=int, summary=f"{PLACEHOLDER_PREFIX}create_or_update_dt") @require_permissions(permission=SqlbotPermission(role=['ws_admin'], type='ds', keyExpression="info.datasource")) -@system_log(LogConfig(operation_type=OperationType.CREATE_OR_UPDATE, module=OperationModules.DATA_TRAINING,resource_id_expr='info.id', result_id_expr="result_self")) +@system_log(LogConfig(operation_type=OperationType.CREATE_OR_UPDATE, module=OperationModules.DATA_TRAINING, + resource_id_expr='info.id', result_id_expr="result_self")) async def create_or_update(session: SessionDep, current_user: CurrentUser, trans: Trans, info: DataTrainingInfo): oid = current_user.oid if info.id: @@ -55,14 +56,16 @@ async def create_or_update(session: SessionDep, current_user: CurrentUser, trans @router.delete("", summary=f"{PLACEHOLDER_PREFIX}delete_dt") -@system_log(LogConfig(operation_type=OperationType.DELETE, module=OperationModules.DATA_TRAINING,resource_id_expr='id_list')) +@system_log( + LogConfig(operation_type=OperationType.DELETE, module=OperationModules.DATA_TRAINING, resource_id_expr='id_list')) @require_permissions(permission=SqlbotPermission(role=['ws_admin'])) async def delete(session: SessionDep, id_list: list[int]): delete_training(session, id_list) @router.get("/{id}/enable/{enabled}", summary=f"{PLACEHOLDER_PREFIX}enable_dt") -@system_log(LogConfig(operation_type=OperationType.UPDATE, module=OperationModules.DATA_TRAINING,resource_id_expr='id')) +@system_log( + LogConfig(operation_type=OperationType.UPDATE, module=OperationModules.DATA_TRAINING, resource_id_expr='id')) @require_permissions(permission=SqlbotPermission(role=['ws_admin'])) async def enable(session: SessionDep, id: int, enabled: bool, trans: Trans): enable_training(session, id, enabled, trans) @@ -89,9 +92,7 @@ def inner(): fields.append(AxisObj(name=trans('i18n_data_training.problem_description'), value='question')) fields.append(AxisObj(name=trans('i18n_data_training.sample_sql'), value='description')) fields.append(AxisObj(name=trans('i18n_data_training.effective_data_sources'), value='datasource_name')) - if current_user.oid == 1: - fields.append( - AxisObj(name=trans('i18n_data_training.advanced_application'), value='advanced_application_name')) + fields.append(AxisObj(name=trans('i18n_data_training.advanced_application'), value='advanced_application_name')) md_data, _fields_list = DataFormat.convert_object_array_for_pandas(fields, data_list) @@ -127,10 +128,9 @@ def inner(): fields.append(AxisObj(name=trans('i18n_data_training.sample_sql_template'), value='description')) fields.append( AxisObj(name=trans('i18n_data_training.effective_data_sources_template'), value='datasource_name')) - if current_user.oid == 1: - fields.append( - AxisObj(name=trans('i18n_data_training.advanced_application_template'), - value='advanced_application_name')) + fields.append( + AxisObj(name=trans('i18n_data_training.advanced_application_template'), + value='advanced_application_name')) md_data, _fields_list = DataFormat.convert_object_array_for_pandas(fields, data_list) @@ -175,10 +175,7 @@ async def upload_excel(trans: Trans, current_user: CurrentUser, file: UploadFile oid = current_user.oid - use_cols = [0, 1, 2] # 问题, 描述, 数据源名称 - # 根据oid确定要读取的列 - if oid == 1: - use_cols = [0, 1, 2, 3] # 问题, 描述, 数据源名称, 高级应用名称 + use_cols = [0, 1, 2, 3] # 问题, 描述, 数据源名称, 高级应用名称 def inner(): @@ -211,19 +208,14 @@ def inner(): description = row[1].strip() if pd.notna(row[1]) and row[1].strip() else '' datasource_name = row[2].strip() if pd.notna(row[2]) and row[2].strip() else '' - advanced_application_name = '' - if oid == 1 and len(row) > 3: + advanced_application_name = None + if len(row) > 3: advanced_application_name = row[3].strip() if pd.notna(row[3]) and row[3].strip() else '' - if oid == 1: - import_data.append( - DataTrainingInfo(oid=oid, question=question, description=description, - datasource_name=datasource_name, - advanced_application_name=advanced_application_name)) - else: - import_data.append( - DataTrainingInfo(oid=oid, question=question, description=description, - datasource_name=datasource_name)) + import_data.append( + DataTrainingInfo(oid=oid, question=question, description=description, + datasource_name=datasource_name, + advanced_application_name=advanced_application_name)) res = batch_create_training(session, import_data, oid, trans) @@ -247,9 +239,8 @@ def inner(): fields.append(AxisObj(name=trans('i18n_data_training.problem_description'), value='question')) fields.append(AxisObj(name=trans('i18n_data_training.sample_sql'), value='description')) fields.append(AxisObj(name=trans('i18n_data_training.effective_data_sources'), value='datasource_name')) - if current_user.oid == 1: - fields.append( - AxisObj(name=trans('i18n_data_training.advanced_application'), value='advanced_application_name')) + fields.append( + AxisObj(name=trans('i18n_data_training.advanced_application'), value='advanced_application_name')) fields.append(AxisObj(name=trans('i18n_data_training.error_info'), value='errors')) md_data, _fields_list = DataFormat.convert_object_array_for_pandas(fields, data_list) diff --git a/backend/apps/datasource/api/datasource.py b/backend/apps/datasource/api/datasource.py index d4e24b7a0..90801fdf2 100644 --- a/backend/apps/datasource/api/datasource.py +++ b/backend/apps/datasource/api/datasource.py @@ -4,6 +4,7 @@ import os import traceback import uuid +import re from io import StringIO from typing import List from urllib.parse import quote @@ -64,6 +65,7 @@ def inner(): @router.get("/check/{ds_id}", response_model=bool, summary=f"{PLACEHOLDER_PREFIX}ds_check") +@require_permissions(permission=SqlbotPermission(type='ds', keyExpression="ds_id")) async def check_by_id(session: SessionDep, trans: Trans, ds_id: int = Path(..., description=f"{PLACEHOLDER_PREFIX}ds_id")): def inner(): @@ -242,8 +244,9 @@ def inner(): # not used -@router.post("/fieldEnum/{id}", include_in_schema=False) -async def field_enum(session: SessionDep, id: int): +@router.post("/fieldEnum/{ds_id}/{id}", include_in_schema=False) +@require_permissions(permission=SqlbotPermission(type='ds', keyExpression="ds_id")) +async def field_enum(session: SessionDep, ds_id: int, id: int): def inner(): return fieldEnum(session, id) @@ -534,12 +537,12 @@ async def upload_ds_schema(session: SessionDep, id: int = Path(..., description= @router.post("/parseExcel", response_model=None, summary=f"{PLACEHOLDER_PREFIX}ds_parse_excel") @require_permissions(permission=SqlbotPermission(role=['ws_admin'])) async def parse_excel(file: UploadFile = File(..., description=f"{PLACEHOLDER_PREFIX}ds_excel")): - ALLOWED_EXTENSIONS = {"xlsx", "xls", "csv"} + ALLOWED_EXTENSIONS = {".xlsx", ".xls", ".csv"} if not file.filename.lower().endswith(tuple(ALLOWED_EXTENSIONS)): raise HTTPException(400, "Only support .xlsx/.xls/.csv") os.makedirs(path, exist_ok=True) - filename = f"{file.filename.split('.')[0]}_{hashlib.sha256(uuid.uuid4().bytes).hexdigest()[:10]}.{file.filename.split('.')[1]}" + filename = f"{file.filename.split('.')[0].split('/')[-1]}_{hashlib.sha256(uuid.uuid4().bytes).hexdigest()[:10]}.{file.filename.split('.')[-1]}" save_path = os.path.join(path, filename) with open(save_path, "wb") as f: f.write(await file.read()) @@ -567,7 +570,7 @@ def inner(): for sheet_info in import_req.sheets: sheet_name = sheet_info.sheetName - table_name = f"{sheet_name}_{hashlib.sha256(uuid.uuid4().bytes).hexdigest()[:10]}" + table_name = f"excel_{filter_string(sheet_name)}_{hashlib.sha256(uuid.uuid4().bytes).hexdigest()[:10]}" fields = sheet_info.fields field_mapping = {f.fieldName: f.fieldType for f in fields} @@ -617,3 +620,9 @@ def inner(): return {"filename": import_req.filePath, "sheets": results} return await asyncio.to_thread(inner) + + +# only allow chinese, a-z, A-Z, 0-9 +def filter_string(text): + pattern = r'[^\u4e00-\u9fa5a-zA-Z0-9]' + return re.sub(pattern, '', text) diff --git a/backend/apps/datasource/crud/datasource.py b/backend/apps/datasource/crud/datasource.py index 372e4ee37..11720a783 100644 --- a/backend/apps/datasource/crud/datasource.py +++ b/backend/apps/datasource/crud/datasource.py @@ -24,6 +24,7 @@ from ..crud.table import delete_table_by_ds_id, update_table from ..models.datasource import CoreDatasource, CreateDatasource, CoreTable, CoreField, ColumnSchema, TableObj, \ DatasourceConf, TableAndFields +from apps.db.db import pool_manager, driver_pool_manager def get_datasource_list(session: SessionDep, user: CurrentUser, oid: Optional[int] = None) -> List[CoreDatasource]: @@ -109,6 +110,10 @@ def update_ds(session: SessionDep, trans: Trans, user: CurrentUser, ds: CoreData session.add(record) session.commit() + # update pool + pool_manager.remove_pool(ds.id) + driver_pool_manager.remove_pool(ds.id) + run_save_ds_embeddings([ds.id]) return ds @@ -135,6 +140,11 @@ async def delete_ds(session: SessionDep, id: int): session.commit() delete_table_by_ds_id(session, id) delete_field_by_ds_id(session, id) + + # update pool + pool_manager.remove_pool(id) + driver_pool_manager.remove_pool(id) + if term: await clear_ws_ds_cache(term.oid) return { @@ -328,18 +338,19 @@ def preview(session: SessionDep, current_user: CurrentUser, id: int, data: Table if fields is None or len(fields) == 0: return {"fields": [], "data": [], "sql": ''} + table = session.query(CoreTable).filter(CoreTable.id == data.table.id).first() conf = DatasourceConf(**json.loads(aes_decrypt(ds.configuration))) if ds.type != "excel" else get_engine_config() sql: str = "" if ds.type == "mysql" or ds.type == "doris" or ds.type == "starrocks" or ds.type == "hive": - sql = f"""SELECT `{"`, `".join(fields)}` FROM `{data.table.table_name}` + sql = f"""SELECT `{"`, `".join(fields)}` FROM `{table.table_name}` {where} LIMIT 100""" elif ds.type == "sqlServer": - sql = f"""SELECT TOP 100 [{"], [".join(fields)}] FROM [{conf.dbSchema}].[{data.table.table_name}] + sql = f"""SELECT TOP 100 [{"], [".join(fields)}] FROM [{conf.dbSchema}].[{table.table_name}] {where} """ elif ds.type == "pg" or ds.type == "excel" or ds.type == "redshift" or ds.type == "kingbase": - sql = f"""SELECT "{'", "'.join(fields)}" FROM "{conf.dbSchema}"."{data.table.table_name}" + sql = f"""SELECT "{'", "'.join(fields)}" FROM "{conf.dbSchema}"."{table.table_name}" {where} LIMIT 100""" elif ds.type == "oracle": @@ -348,25 +359,21 @@ def preview(session: SessionDep, current_user: CurrentUser, id: int, data: Table # ORDER BY "{fields[0]}" # OFFSET 0 ROWS FETCH NEXT 100 ROWS ONLY""" sql = f"""SELECT * FROM - (SELECT "{'", "'.join(fields)}" FROM "{conf.dbSchema}"."{data.table.table_name}" + (SELECT "{'", "'.join(fields)}" FROM "{conf.dbSchema}"."{table.table_name}" {where} ORDER BY "{fields[0]}") WHERE ROWNUM <= 100 """ elif ds.type == "ck": - sql = f"""SELECT "{'", "'.join(fields)}" FROM "{data.table.table_name}" + sql = f"""SELECT "{'", "'.join(fields)}" FROM "{table.table_name}" {where} LIMIT 100""" elif ds.type == "dm": - sql = f"""SELECT "{'", "'.join(fields)}" FROM "{conf.dbSchema}"."{data.table.table_name}" + sql = f"""SELECT "{'", "'.join(fields)}" FROM "{conf.dbSchema}"."{table.table_name}" {where} LIMIT 100""" elif ds.type == "es": - sql = f"""SELECT "{'", "'.join(fields)}" FROM "{data.table.table_name}" - {where} - LIMIT 100""" - elif ds.type == "sqlite": - sql = f"""SELECT "{'", "'.join(fields)}" FROM "{data.table.table_name}" + sql = f"""SELECT "{'", "'.join(fields)}" FROM "{table.table_name}" {where} LIMIT 100""" return exec_sql(ds, sql, True) diff --git a/backend/apps/datasource/crud/permission.py b/backend/apps/datasource/crud/permission.py index 3e38d0447..dfa56c75f 100644 --- a/backend/apps/datasource/crud/permission.py +++ b/backend/apps/datasource/crud/permission.py @@ -1,7 +1,7 @@ import json from typing import List, Optional -from sqlalchemy import and_ +from sqlalchemy import and_, cast, or_ from sqlbot_xpack.permissions.api.permission import transRecord2DTO from sqlbot_xpack.permissions.models.ds_permission import DsPermission, PermissionDTO from sqlbot_xpack.permissions.models.ds_rules import DsRules @@ -9,6 +9,7 @@ from apps.datasource.crud.row_permission import transFilterTree from apps.datasource.models.datasource import CoreDatasource, CoreField, CoreTable from common.core.deps import CurrentUser, SessionDep +from sqlalchemy.dialects.postgresql import JSONB def get_row_permission_filters(session: SessionDep, current_user: CurrentUser, ds: CoreDatasource, @@ -22,7 +23,7 @@ def get_row_permission_filters(session: SessionDep, current_user: CurrentUser, d filters = [] if is_normal_user(current_user): - contain_rules = session.query(DsRules).all() + # contain_rules = session.query(DsRules).all() for table in table_list: row_permissions = session.query(DsPermission).filter( and_(DsPermission.table_id == table.id, DsPermission.type == 'row')).all() @@ -30,16 +31,24 @@ def get_row_permission_filters(session: SessionDep, current_user: CurrentUser, d if row_permissions is not None: for permission in row_permissions: # check permission and user in same rules - flag = False - for r in contain_rules: - p_list = json.loads(r.permission_list) - u_list = json.loads(r.user_list) - if p_list is not None and u_list is not None and permission.id in p_list and ( - current_user.id in u_list or f'{current_user.id}' in u_list): - flag = True - break - if flag: + obj = session.query(DsRules).filter( + and_(DsRules.permission_list.op('@>')(cast([permission.id], JSONB)), + or_(DsRules.user_list.op('@>')(cast([f'{current_user.id}'], JSONB)), + DsRules.user_list.op('@>')(cast([current_user.id], JSONB)))) + ).first() + if obj is not None: res.append(transRecord2DTO(session, permission)) + + # flag = False + # for r in contain_rules: + # p_list = json.loads(r.permission_list) + # u_list = json.loads(r.user_list) + # if p_list is not None and u_list is not None and permission.id in p_list and ( + # current_user.id in u_list or f'{current_user.id}' in u_list): + # flag = True + # break + # if flag: + # res.append(transRecord2DTO(session, permission)) where_str = transFilterTree(session, current_user, res, ds) if where_str: filters.append({"table": table.table_name, "filter": where_str}) @@ -54,22 +63,26 @@ def get_column_permission_fields(session: SessionDep, current_user: CurrentUser, if column_permissions is not None: for permission in column_permissions: # check permission and user in same rules - # obj = session.query(DsRules).filter( - # and_(DsRules.permission_list.op('@>')(cast([permission.id], JSONB)), - # or_(DsRules.user_list.op('@>')(cast([f'{current_user.id}'], JSONB)), - # DsRules.user_list.op('@>')(cast([current_user.id], JSONB)))) - # ).first() - flag = False - for r in contain_rules: - p_list = json.loads(r.permission_list) - u_list = json.loads(r.user_list) - if p_list is not None and u_list is not None and permission.id in p_list and ( - current_user.id in u_list or f'{current_user.id}' in u_list): - flag = True - break - if flag: + obj = session.query(DsRules).filter( + and_(DsRules.permission_list.op('@>')(cast([permission.id], JSONB)), + or_(DsRules.user_list.op('@>')(cast([f'{current_user.id}'], JSONB)), + DsRules.user_list.op('@>')(cast([current_user.id], JSONB)))) + ).first() + if obj is not None: permission_list = json.loads(permission.permissions) fields = filter_list(fields, permission_list) + + # flag = False + # for r in contain_rules: + # p_list = json.loads(r.permission_list) + # u_list = json.loads(r.user_list) + # if p_list is not None and u_list is not None and permission.id in p_list and ( + # current_user.id in u_list or f'{current_user.id}' in u_list): + # flag = True + # break + # if flag: + # permission_list = json.loads(permission.permissions) + # fields = filter_list(fields, permission_list) return fields diff --git a/backend/apps/datasource/crud/row_permission.py b/backend/apps/datasource/crud/row_permission.py index 86fbc9b76..b19063af5 100644 --- a/backend/apps/datasource/crud/row_permission.py +++ b/backend/apps/datasource/crud/row_permission.py @@ -9,6 +9,21 @@ from common.core.deps import SessionDep, CurrentUser +def _escape_sql_value(value: str) -> str: + """Escape a string value for safe inclusion in a SQL literal. + + Replaces single quotes with two single quotes (standard SQL escaping) + and strips characters that could break out of the string context. + """ + if value is None: + return value + # Standard SQL escaping: double any embedded single-quote characters + escaped = str(value).replace("'", "''") + # Remove backslashes that some drivers interpret as escape characters + escaped = escaped.replace("\\", "\\\\") + return escaped + + def transFilterTree(session: SessionDep, current_user: CurrentUser, tree_list: List[any], ds: CoreDatasource) -> str | None: if tree_list is None: @@ -24,10 +39,16 @@ def transFilterTree(session: SessionDep, current_user: CurrentUser, tree_list: L return " AND ".join(res) +_VALID_LOGIC_OPS = {"AND", "OR"} + + def transTreeToWhere(session: SessionDep, current_user: CurrentUser, tree: any, ds: CoreDatasource) -> str | None: if tree is None: return None logic = tree['logic'] + # Validate the logic operator to prevent injection via this field + if logic.upper() not in _VALID_LOGIC_OPS: + return None items = tree['items'] list: List[str] = [] @@ -56,11 +77,12 @@ def transTreeItem(session: SessionDep, current_user: CurrentUser, item: Dict, ds if item['filter_type'] == 'enum': if len(item['enum_value']) > 0: + escaped_values = [_escape_sql_value(v) for v in item['enum_value']] if ds['type'] == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - res = "(" + whereName + " IN (N'" + "',N'".join(item['enum_value']) + "'))" + res = "(" + whereName + " IN (N'" + "',N'".join(escaped_values) + "'))" else: - res = "(" + whereName + " IN ('" + "','".join(item['enum_value']) + "'))" + res = "(" + whereName + " IN ('" + "','".join(escaped_values) + "'))" else: # if system variable, do check and get value # new field: value_type(variable or normal), variable_id @@ -119,23 +141,26 @@ def transTreeItem(session: SessionDep, current_user: CurrentUser, item: Dict, ds elif item['term'] == 'not_empty': whereValue = "''" elif item['term'] == 'in' or item['term'] == 'not in': + escaped_values = [_escape_sql_value(v) for v in values] if ds.type == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - whereValue = "(N'" + "', N'".join(values) + "')" + whereValue = "(N'" + "', N'".join(escaped_values) + "')" else: - whereValue = "('" + "', '".join(values) + "')" + whereValue = "('" + "', '".join(escaped_values) + "')" elif item['term'] == 'like' or item['term'] == 'not like': + escaped_v = _escape_sql_value(values[0]) if ds.type == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - whereValue = f"N'%{values[0]}%'" + whereValue = f"N'%{escaped_v}%'" else: - whereValue = f"'%{values[0]}%'" + whereValue = f"'%{escaped_v}%'" else: + escaped_v = _escape_sql_value(values[0]) if ds.type == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - whereValue = f"N'{values[0]}'" + whereValue = f"N'{escaped_v}'" else: - whereValue = f"'{values[0]}'" + whereValue = f"'{escaped_v}'" res = whereName + whereTerm + whereValue else: @@ -153,23 +178,26 @@ def transTreeItem(session: SessionDep, current_user: CurrentUser, item: Dict, ds elif item['term'] == 'not_empty': whereValue = "''" elif item['term'] == 'in' or item['term'] == 'not in': + escaped_values = [_escape_sql_value(v) for v in value.split(",")] if ds.type == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - whereValue = "(N'" + "', N'".join(value.split(",")) + "')" + whereValue = "(N'" + "', N'".join(escaped_values) + "')" else: - whereValue = "('" + "', '".join(value.split(",")) + "')" + whereValue = "('" + "', '".join(escaped_values) + "')" elif item['term'] == 'like' or item['term'] == 'not like': + escaped_v = _escape_sql_value(value) if ds.type == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - whereValue = f"N'%{value}%'" + whereValue = f"N'%{escaped_v}%'" else: - whereValue = f"'%{value}%'" + whereValue = f"'%{escaped_v}%'" else: + escaped_v = _escape_sql_value(value) if ds.type == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - whereValue = f"N'{value}'" + whereValue = f"N'{escaped_v}'" else: - whereValue = f"'{value}'" + whereValue = f"'{escaped_v}'" res = whereName + whereTerm + whereValue return res @@ -226,6 +254,8 @@ def getSysVariableValue(sys_variable: SystemVariable, current_user: CurrentUser, if sys_variable.value[0] == 'email': v = current_user.email + escaped_v = _escape_sql_value(v) if v is not None else v + whereValue = '' if item['term'] == 'null': whereValue = '' @@ -238,20 +268,20 @@ def getSysVariableValue(sys_variable: SystemVariable, current_user: CurrentUser, elif item['term'] == 'in' or item['term'] == 'not in': if ds.type == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - whereValue = f"(N'{v}')" + whereValue = f"(N'{escaped_v}')" else: - whereValue = f"('{v}')" + whereValue = f"('{escaped_v}')" elif item['term'] == 'like' or item['term'] == 'not like': if ds.type == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - whereValue = f"N'%{v}%'" + whereValue = f"N'%{escaped_v}%'" else: - whereValue = f"'%{v}%'" + whereValue = f"'%{escaped_v}%'" else: if ds.type == 'sqlServer' and ( field.field_type == 'nchar' or field.field_type == 'NCHAR' or field.field_type == 'nvarchar' or field.field_type == 'NVARCHAR'): - whereValue = f"N'{v}'" + whereValue = f"N'{escaped_v}'" else: - whereValue = f"'{v}'" + whereValue = f"'{escaped_v}'" return whereValue diff --git a/backend/apps/datasource/embedding/ds_embedding.py b/backend/apps/datasource/embedding/ds_embedding.py index 9bfe4a48e..f75ee7df5 100644 --- a/backend/apps/datasource/embedding/ds_embedding.py +++ b/backend/apps/datasource/embedding/ds_embedding.py @@ -11,11 +11,11 @@ from apps.system.crud.assistant import AssistantOutDs from common.core.config import settings from common.core.deps import CurrentAssistant -from common.core.deps import SessionDep, CurrentUser +from common.core.deps import SessionDep from common.utils.utils import SQLBotLogUtil -def get_ds_embedding(session: SessionDep, current_user: CurrentUser, _ds_list, out_ds: AssistantOutDs, +def get_ds_embedding(session: SessionDep, _ds_list, out_ds: AssistantOutDs, question: str, current_assistant: Optional[CurrentAssistant] = None): _list = [] diff --git a/backend/apps/datasource/models/datasource.py b/backend/apps/datasource/models/datasource.py index 3971318cf..42f1dcd25 100644 --- a/backend/apps/datasource/models/datasource.py +++ b/backend/apps/datasource/models/datasource.py @@ -122,6 +122,7 @@ class DatasourceConf(BaseModel): timeout: int = 30 lowVersion: bool = False ssl: bool = False + poolSize: int = 5 def to_dict(self): return { @@ -138,7 +139,8 @@ def to_dict(self): "mode": self.mode, "timeout": self.timeout, "lowVersion": self.lowVersion, - "ssl": self.ssl + "ssl": self.ssl, + "poolSize": self.poolSize } diff --git a/backend/apps/db/constant.py b/backend/apps/db/constant.py index 6ee33f02f..dcec81901 100644 --- a/backend/apps/db/constant.py +++ b/backend/apps/db/constant.py @@ -28,7 +28,6 @@ class DB(Enum): oracle = ('oracle', 'Oracle', '"', '"', ConnectType.sqlalchemy, 'Oracle', []) pg = ('pg', 'PostgreSQL', '"', '"', ConnectType.sqlalchemy, 'PostgreSQL', []) starrocks = ('starrocks', 'StarRocks', '`', '`', ConnectType.py_driver, 'StarRocks', []) - sqlite = ('sqlite', 'SQLite', '"', '"', ConnectType.sqlalchemy, 'SQLite', []) hive = ('hive', 'Apache Hive', '`', '`', ConnectType.py_driver, 'Hive', []) def __init__(self, type, db_name, prefix, suffix, connect_type: ConnectType, template_name: str, diff --git a/backend/apps/db/db.py b/backend/apps/db/db.py index b87ff7891..8f27432ce 100644 --- a/backend/apps/db/db.py +++ b/backend/apps/db/db.py @@ -2,7 +2,6 @@ import json import os import platform -import re import urllib.parse from datetime import datetime, date, time, timedelta from decimal import Decimal @@ -35,8 +34,9 @@ from common.core.config import settings import sqlglot from sqlglot import expressions as exp -from sqlalchemy.pool import NullPool from pyhive import hive +from sqlalchemy.pool import NullPool +from dbutils.pooled_db import PooledDB try: if os.path.exists(settings.ORACLE_CLIENT_PATH): @@ -90,8 +90,6 @@ def get_uri_from_config(type: str, conf: DatasourceConf) -> str: db_url = f"clickhouse+http://{urllib.parse.quote(conf.username)}:{urllib.parse.quote(conf.password)}@{conf.host}:{conf.port}/{conf.database}?{conf.extraJdbc}" else: db_url = f"clickhouse+http://{urllib.parse.quote(conf.username)}:{urllib.parse.quote(conf.password)}@{conf.host}:{conf.port}/{conf.database}" - elif equals_ignore_case(type, "sqlite"): - db_url = f"sqlite:///{conf.filename}" else: raise 'The datasource type not support.' return db_url @@ -138,7 +136,7 @@ def get_origin_connect(type: str, conf: DatasourceConf): # use sqlalchemy -def get_engine(ds: CoreDatasource, timeout: int = 0) -> Engine: +def get_engine(ds: CoreDatasource, timeout: int = 0, use_pool: bool = False) -> Engine: conf = DatasourceConf(**json.loads(aes_decrypt(ds.configuration))) if not equals_ignore_case(ds.type, "excel") else get_engine_config() if conf.timeout is None: @@ -146,26 +144,32 @@ def get_engine(ds: CoreDatasource, timeout: int = 0) -> Engine: if timeout > 0: conf.timeout = timeout + db_config = { + 'pool_size': conf.poolSize if conf.poolSize else 5, + 'max_overflow': 20, + 'pool_recycle': 3600 + } if use_pool else { + 'poolclass': NullPool + } + if equals_ignore_case(ds.type, "pg"): if conf.dbSchema is not None and conf.dbSchema != "": engine = create_engine(get_uri(ds), connect_args={"options": f"-c search_path={urllib.parse.quote(conf.dbSchema)}", - "connect_timeout": conf.timeout}, poolclass=NullPool) + "connect_timeout": conf.timeout}, **db_config) else: - engine = create_engine(get_uri(ds), connect_args={"connect_timeout": conf.timeout}, poolclass=NullPool) + engine = create_engine(get_uri(ds), connect_args={"connect_timeout": conf.timeout}, **db_config) elif equals_ignore_case(ds.type, 'sqlServer'): engine = create_engine('mssql+pymssql://', creator=lambda: get_origin_connect(ds.type, conf), - poolclass=NullPool) + **db_config) elif equals_ignore_case(ds.type, 'oracle'): - engine = create_engine(get_uri(ds), poolclass=NullPool) + engine = create_engine(get_uri(ds), **db_config) elif equals_ignore_case(ds.type, 'mysql'): # mysql ssl_mode = {"require": True} if conf.ssl else None engine = create_engine(get_uri(ds), connect_args={"connect_timeout": conf.timeout, "ssl": ssl_mode}, - poolclass=NullPool) - elif equals_ignore_case(ds.type, 'sqlite'): - engine = create_engine(get_uri(ds), connect_args={"check_same_thread": False}, poolclass=NullPool) + **db_config) else: # ck - engine = create_engine(get_uri(ds), connect_args={"connect_timeout": conf.timeout}, poolclass=NullPool) + engine = create_engine(get_uri(ds), connect_args={"connect_timeout": conf.timeout}, **db_config) return engine @@ -175,10 +179,115 @@ def get_session(ds: CoreDatasource | AssistantOutDsSchema): out_conf = get_out_ds_conf(ds, 30) ds.configuration = out_conf - engine = get_engine(ds) - session_maker = sessionmaker(bind=engine) - session = session_maker() - return session + # engine = get_engine(ds) + # session_maker = sessionmaker(bind=engine) + + # get session from pool + session = pool_manager.get_pool(ds=ds, **{}) + return session() + + +def get_driver_connection(ds: CoreDatasource | AssistantOutDsSchema, db_config: dict = {}, use_pool: bool = False): + conf = DatasourceConf(**json.loads(aes_decrypt(ds.configuration))) + extra_config_dict = get_extra_config(conf) + + pool_config = { + 'maxconnections': conf.poolSize if conf.poolSize else 5, + 'mincached': 5, + 'maxcached': 10, + 'blocking': True, + 'maxusage': 100, + 'ping': 1, + } if use_pool else {} + conn_conf = extra_config_dict | db_config | pool_config + + conn = None + if equals_ignore_case(ds.type, 'dm'): + if not use_pool: + conn = dmPython.connect(user=conf.username, password=conf.password, server=conf.host, + port=conf.port, **conn_conf) + else: + conn = PooledDB( + creator=dmPython, + user=conf.username, + password=conf.password, + server=conf.host, + port=conf.port, + **conn_conf + ) + elif equals_ignore_case(ds.type, 'doris', 'starrocks'): + ssl_args = {'ssl': {'ssl_mode': 'REQUIRE'}} if conf.ssl else {} + args = conn_conf | ssl_args + if not use_pool: + conn = pymysql.connect(user=conf.username, passwd=conf.password, host=conf.host, + port=conf.port, db=conf.database, connect_timeout=conf.timeout, + read_timeout=conf.timeout, **conn_conf, + **args) + else: + conn = PooledDB( + creator=pymysql, + user=conf.username, + passwd=conf.password, + host=conf.host, + port=conf.port, + db=conf.database, + connect_timeout=conf.timeout, + read_timeout=conf.timeout, + **args + ) + elif equals_ignore_case(ds.type, 'redshift'): + if not use_pool: + conn = redshift_connector.connect(host=conf.host, port=conf.port, database=conf.database, + user=conf.username, + password=conf.password, + timeout=conf.timeout, **conn_conf) + else: + conn = PooledDB( + creator=redshift_connector, + host=conf.host, + port=conf.port, + database=conf.database, + user=conf.username, + password=conf.password, + timeout=conf.timeout, + **conn_conf + ) + elif equals_ignore_case(ds.type, 'kingbase'): + if not use_pool: + conn = psycopg2.connect(host=conf.host, port=conf.port, database=conf.database, user=conf.username, + password=conf.password, + options=f"-c statement_timeout={conf.timeout * 1000}", + **conn_conf) + else: + conn = PooledDB( + creator=psycopg2, + host=conf.host, + port=conf.port, + database=conf.database, + user=conf.username, + password=conf.password, + options=f"-c statement_timeout={conf.timeout * 1000}", + **conn_conf + ) + elif equals_ignore_case(ds.type, 'hive'): + if not use_pool: + conn = hive.connect(host=conf.host, port=conf.port, username=conf.username, + database=conf.database, **conn_conf) + else: + conn = PooledDB( + creator=hive, + host=conf.host, + port=conf.port, + username=conf.username, + database=conf.database, **conn_conf + ) + + return conn + + +def get_driver_pool(ds: CoreDatasource | AssistantOutDsSchema, db_config: dict = {}): + pool = driver_pool_manager.get_pool(ds, db_config) + return pool def check_connection(trans: Optional[Trans], ds: CoreDatasource | AssistantOutDsSchema, is_raise: bool = False): @@ -321,16 +430,13 @@ def get_version(ds: CoreDatasource | AssistantOutDsSchema): else: extra_config_dict = get_extra_config(conf) if equals_ignore_case(ds.type, 'dm'): - with dmPython.connect(user=conf.username, password=conf.password, server=conf.host, - port=conf.port) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: cursor.execute(sql, timeout=10, **extra_config_dict) res = cursor.fetchall() version = res[0][0] elif equals_ignore_case(ds.type, 'doris', 'starrocks'): - ssl_args = {'ssl': {'ssl_mode': 'REQUIRE'}} if conf.ssl else {} - with pymysql.connect(user=conf.username, passwd=conf.password, host=conf.host, - port=conf.port, db=conf.database, connect_timeout=10, - read_timeout=10, **extra_config_dict, **ssl_args) as conn, conn.cursor() as cursor: + t_conf = {'connect_timeout': 10, 'read_timeout': 10} + with get_driver_pool(ds, t_conf).connection() as conn, conn.cursor() as cursor: cursor.execute(sql) res = cursor.fetchall() version = res[0][0] @@ -346,7 +452,7 @@ def get_schema(ds: CoreDatasource): conf = DatasourceConf(**json.loads(aes_decrypt(ds.configuration))) if ds.type != "excel" else get_engine_config() db = DB.get_db(ds.type) if db.connect_type == ConnectType.sqlalchemy: - with get_session(ds) as session: + with sessionmaker(bind=get_engine(ds))() as session: sql: str = '' if equals_ignore_case(ds.type, "sqlServer"): sql = """select name @@ -357,37 +463,29 @@ def get_schema(ds: CoreDatasource): elif equals_ignore_case(ds.type, "oracle"): sql = """select * from all_users""" - elif equals_ignore_case(ds.type, "sqlite"): - return ['main'] with session.execute(text(sql)) as result: res = result.fetchall() res_list = [item[0] for item in res] return res_list else: - extra_config_dict = get_extra_config(conf) + # extra_config_dict = get_extra_config(conf) if equals_ignore_case(ds.type, 'dm'): - with dmPython.connect(user=conf.username, password=conf.password, server=conf.host, - port=conf.port, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_connection(ds) as conn, conn.cursor() as cursor: cursor.execute("""select OBJECT_NAME - from dba_objects + from all_objects where object_type = 'SCH'""", timeout=conf.timeout) res = cursor.fetchall() res_list = [item[0] for item in res] return res_list elif equals_ignore_case(ds.type, 'redshift'): - with redshift_connector.connect(host=conf.host, port=conf.port, database=conf.database, user=conf.username, - password=conf.password, - timeout=conf.timeout, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_connection(ds) as conn, conn.cursor() as cursor: cursor.execute("""SELECT nspname FROM pg_namespace""") res = cursor.fetchall() res_list = [item[0] for item in res] return res_list elif equals_ignore_case(ds.type, 'kingbase'): - with psycopg2.connect(host=conf.host, port=conf.port, database=conf.database, user=conf.username, - password=conf.password, - options=f"-c statement_timeout={conf.timeout * 1000}", - **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_connection(ds) as conn, conn.cursor() as cursor: cursor.execute("""SELECT nspname FROM pg_namespace""") res = cursor.fetchall() @@ -401,43 +499,34 @@ def get_tables(ds: CoreDatasource): db = DB.get_db(ds.type) sql, sql_param = get_table_sql(ds, conf, get_version(ds)) if db.connect_type == ConnectType.sqlalchemy: - with get_session(ds) as session: + with sessionmaker(bind=get_engine(ds))() as session: with session.execute(text(sql), {"param": sql_param}) as result: res = result.fetchall() res_list = [TableSchema(*item) for item in res] return res_list else: - extra_config_dict = get_extra_config(conf) + # extra_config_dict = get_extra_config(conf) if equals_ignore_case(ds.type, 'dm'): - with dmPython.connect(user=conf.username, password=conf.password, server=conf.host, - port=conf.port, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_connection(ds) as conn, conn.cursor() as cursor: cursor.execute(sql, {"param": sql_param}, timeout=conf.timeout) res = cursor.fetchall() res_list = [TableSchema(*item) for item in res] return res_list elif equals_ignore_case(ds.type, 'doris', 'starrocks'): - ssl_args = {'ssl': {'ssl_mode': 'REQUIRE'}} if conf.ssl else {} - with pymysql.connect(user=conf.username, passwd=conf.password, host=conf.host, - port=conf.port, db=conf.database, connect_timeout=conf.timeout, - read_timeout=conf.timeout, **extra_config_dict, - **ssl_args) as conn, conn.cursor() as cursor: + # ssl_args = {'ssl': {'ssl_mode': 'REQUIRE'}} if conf.ssl else {} + with get_driver_connection(ds) as conn, conn.cursor() as cursor: cursor.execute(sql, (sql_param,)) res = cursor.fetchall() res_list = [TableSchema(*item) for item in res] return res_list elif equals_ignore_case(ds.type, 'redshift'): - with redshift_connector.connect(host=conf.host, port=conf.port, database=conf.database, user=conf.username, - password=conf.password, - timeout=conf.timeout, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_connection(ds) as conn, conn.cursor() as cursor: cursor.execute(sql, (sql_param,)) res = cursor.fetchall() res_list = [TableSchema(*item) for item in res] return res_list elif equals_ignore_case(ds.type, 'kingbase'): - with psycopg2.connect(host=conf.host, port=conf.port, database=conf.database, user=conf.username, - password=conf.password, - options=f"-c statement_timeout={conf.timeout * 1000}", - **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_connection(ds) as conn, conn.cursor() as cursor: cursor.execute(sql.format(sql_param)) res = cursor.fetchall() res_list = [TableSchema(*item) for item in res] @@ -447,8 +536,7 @@ def get_tables(ds: CoreDatasource): res_list = [TableSchema(*item) for item in res] return res_list elif equals_ignore_case(ds.type, 'hive'): - with hive.connect(host=conf.host, port=conf.port, username=conf.username, - database=conf.database, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_connection(ds) as conn, conn.cursor() as cursor: cursor.execute(sql) res = cursor.fetchall() res_list = [TableSchema(*item) for item in res] @@ -464,43 +552,31 @@ def get_fields(ds: CoreDatasource, table_name: str = None): with get_session(ds) as session: with session.execute(text(sql), {"param1": p1, "param2": p2}) as result: res = result.fetchall() - if equals_ignore_case(ds.type, "sqlite"): - res_list = [ColumnSchema(item[1], item[2], '') for item in res] - else: - res_list = [ColumnSchema(*item) for item in res] + res_list = [ColumnSchema(*item) for item in res] return res_list else: - extra_config_dict = get_extra_config(conf) + # extra_config_dict = get_extra_config(conf) if equals_ignore_case(ds.type, 'dm'): - with dmPython.connect(user=conf.username, password=conf.password, server=conf.host, - port=conf.port, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"param1": p1, "param2": p2}, timeout=conf.timeout) res = cursor.fetchall() res_list = [ColumnSchema(*item) for item in res] return res_list elif equals_ignore_case(ds.type, 'doris', 'starrocks'): - ssl_args = {'ssl': {'ssl_mode': 'REQUIRE'}} if conf.ssl else {} - with pymysql.connect(user=conf.username, passwd=conf.password, host=conf.host, - port=conf.port, db=conf.database, connect_timeout=conf.timeout, - read_timeout=conf.timeout, **extra_config_dict, - **ssl_args) as conn, conn.cursor() as cursor: + # ssl_args = {'ssl': {'ssl_mode': 'REQUIRE'}} if conf.ssl else {} + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: cursor.execute(sql, (p1, p2)) res = cursor.fetchall() res_list = [ColumnSchema(*item) for item in res] return res_list elif equals_ignore_case(ds.type, 'redshift'): - with redshift_connector.connect(host=conf.host, port=conf.port, database=conf.database, user=conf.username, - password=conf.password, - timeout=conf.timeout, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: cursor.execute(sql, (p1, p2)) res = cursor.fetchall() res_list = [ColumnSchema(*item) for item in res] return res_list elif equals_ignore_case(ds.type, 'kingbase'): - with psycopg2.connect(host=conf.host, port=conf.port, database=conf.database, user=conf.username, - password=conf.password, - options=f"-c statement_timeout={conf.timeout * 1000}", - **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: cursor.execute(sql.format(p1, p2)) res = cursor.fetchall() res_list = [ColumnSchema(*item) for item in res] @@ -510,8 +586,7 @@ def get_fields(ds: CoreDatasource, table_name: str = None): res_list = [ColumnSchema(*item) for item in res] return res_list elif equals_ignore_case(ds.type, 'hive'): - with hive.connect(host=conf.host, port=conf.port, username=conf.username, - database=conf.database, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: cursor.execute(sql) res = cursor.fetchall() res_list = [ColumnSchema(*item) for item in res] @@ -591,122 +666,222 @@ def convert_value(value, datetime_format='space'): return value +def is_numeric_type_code(type_code, dialect_name: str) -> bool: + """ + 根据数据库方言和 type_code 判断是否为数值类型 + + Args: + type_code: cursor.description[col_idx][1] 的值 + dialect_name: SQLAlchemy dialect name (mysql/postgresql/mssql/oracle/sqlite 等) + + Returns: + bool: 是否为数值类型 + """ + dialect_name = dialect_name.lower() + + # ---------- MySQL (pymysql) ---------- + if dialect_name == 'mysql': + if isinstance(type_code, int): + return type_code in { + 1, # TINYINT + 2, # SMALLINT + 3, # INT + 4, # FLOAT + 5, # DOUBLE + 8, # BIGINT + 9, # MEDIUMINT + 16, # BIT + 246, # DECIMAL/NEWDECIMAL + } + return False + + # ---------- PostgreSQL (psycopg2) ---------- + if dialect_name == 'postgresql': + if isinstance(type_code, int): + return type_code in { + 20, # int8 + 21, # int2 + 23, # int4 + 700, # float4 + 701, # float8 + 1700, # numeric + 16, # boolean + } + return False + + # ---------- Oracle (cx_Oracle / oracledb) ---------- + if dialect_name == 'oracle': + type_str = str(type_code).upper() + return any(kw in type_str for kw in ['NUMBER', 'FLOAT', 'INTEGER', 'BINARY_FLOAT', 'BINARY_DOUBLE']) + + if dialect_name == 'clickhouse': + if isinstance(type_code, str): + upper_type = type_code.upper() + # 数值类型关键字 + numeric_prefixes = ( + 'INT', # Int8/16/32/64 + 'UINT', # UInt8/16/32/64 ✅ 加上 UINT + 'FLOAT', # Float32, Float64 + 'DECIMAL', # Decimal, Decimal32/64/128 + 'BOOL', # Bool + 'BIT', # 极少数场景 + ) + return any(upper_type.startswith(p) for p in numeric_prefixes) + + # ---------- SQL Server (pyodbc / pymssql) ---------- + if dialect_name == 'mssql': + if isinstance(type_code, int): + # SQL Server (pyodbc / ODBC) 数值类型码 + return type_code in { + 2, # smallint + 3, # int + 4, # tinyint + 5, # float / real / decimal / numeric / money / smallmoney + 6, # bit + 7, # bigint + } + + # ---------- SQLite ---------- + if dialect_name == 'sqlite': + if isinstance(type_code, int): + return type_code in {1, 2, 3, 4, 5} # INTEGER, FLOAT, NUMERIC, etc. + return False + + # ---------- 未知数据库,保守返回 False ---------- + return False + + def exec_sql(ds: CoreDatasource | AssistantOutDsSchema, sql: str, origin_column=False): while sql.endswith(';'): sql = sql[:-1] # check execute sql only contain read operations - if not check_sql_read(sql, ds): - raise ValueError(f"SQL can only contain read operations") + is_safe, error_reason = check_sql_read(sql, ds) + if not is_safe: + raise ValueError(f"SQL can only contain read operations: {error_reason}") db = DB.get_db(ds.type) if db.connect_type == ConnectType.sqlalchemy: with get_session(ds) as session: + # 获取当前数据库方言 + dialect_name = session.bind.dialect.name + with session.execute(text(sql)) as result: try: columns = result.keys()._keys if origin_column else [item.lower() for item in result.keys()._keys] + + fields_info = [] + + for col_idx, col_name in enumerate(columns): + is_numeric = False + try: + type_code = result.cursor.description[col_idx][1] + is_numeric = is_numeric_type_code(type_code, dialect_name) + except (IndexError, AttributeError): + pass + + fields_info.append({ + "name": col_name, + "is_numeric": is_numeric + }) + res = result.fetchall() result_list = [ {str(columns[i]): convert_value(value) for i, value in enumerate(tuple_item)} for tuple_item in res ] - return {"fields": columns, "data": result_list, + return {"fields": columns, "data": result_list, "fields_info": fields_info, "sql": bytes.decode(base64.b64encode(bytes(sql, 'utf-8')))} except Exception as ex: raise ParseSQLResultError(str(ex)) else: conf = DatasourceConf(**json.loads(aes_decrypt(ds.configuration))) - extra_config_dict = get_extra_config(conf) + # extra_config_dict = get_extra_config(conf) if equals_ignore_case(ds.type, 'dm'): - with dmPython.connect(user=conf.username, password=conf.password, server=conf.host, - port=conf.port, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: try: cursor.execute(sql, timeout=conf.timeout) res = cursor.fetchall() columns = [field[0] for field in cursor.description] if origin_column else [field[0].lower() for field in cursor.description] + fields_info = build_fields_info_from_cursor(cursor, origin_column, 'dm') result_list = [ {str(columns[i]): convert_value(value) for i, value in enumerate(tuple_item)} for tuple_item in res ] - return {"fields": columns, "data": result_list, + return {"fields": columns, "data": result_list, "fields_info": fields_info, "sql": bytes.decode(base64.b64encode(bytes(sql, 'utf-8')))} except Exception as ex: raise ParseSQLResultError(str(ex)) elif equals_ignore_case(ds.type, 'doris', 'starrocks'): - ssl_args = {'ssl': {'ssl_mode': 'REQUIRE'}} if conf.ssl else {} - with pymysql.connect(user=conf.username, passwd=conf.password, host=conf.host, - port=conf.port, db=conf.database, connect_timeout=conf.timeout, - read_timeout=conf.timeout, **extra_config_dict, - **ssl_args) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: try: cursor.execute(sql) res = cursor.fetchall() columns = [field[0] for field in cursor.description] if origin_column else [field[0].lower() for field in cursor.description] + fields_info = build_fields_info_from_cursor(cursor, origin_column, 'mysql') result_list = [ {str(columns[i]): convert_value(value) for i, value in enumerate(tuple_item)} for tuple_item in res ] - return {"fields": columns, "data": result_list, + return {"fields": columns, "data": result_list, "fields_info": fields_info, "sql": bytes.decode(base64.b64encode(bytes(sql, 'utf-8')))} except Exception as ex: raise ParseSQLResultError(str(ex)) elif equals_ignore_case(ds.type, 'redshift'): - with redshift_connector.connect(host=conf.host, port=conf.port, database=conf.database, user=conf.username, - password=conf.password, - timeout=conf.timeout, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: try: cursor.execute(sql) res = cursor.fetchall() columns = [field[0] for field in cursor.description] if origin_column else [field[0].lower() for field in cursor.description] + fields_info = build_fields_info_from_cursor(cursor, origin_column, 'postgresql') result_list = [ {str(columns[i]): convert_value(value) for i, value in enumerate(tuple_item)} for tuple_item in res ] - return {"fields": columns, "data": result_list, + return {"fields": columns, "data": result_list, "fields_info": fields_info, "sql": bytes.decode(base64.b64encode(bytes(sql, 'utf-8')))} except Exception as ex: raise ParseSQLResultError(str(ex)) elif equals_ignore_case(ds.type, 'kingbase'): - with psycopg2.connect(host=conf.host, port=conf.port, database=conf.database, user=conf.username, - password=conf.password, - options=f"-c statement_timeout={conf.timeout * 1000}", - **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: try: cursor.execute(sql) res = cursor.fetchall() columns = [field[0] for field in cursor.description] if origin_column else [field[0].lower() for field in cursor.description] + fields_info = build_fields_info_from_cursor(cursor, origin_column, 'postgresql') result_list = [ {str(columns[i]): convert_value(value) for i, value in enumerate(tuple_item)} for tuple_item in res ] - return {"fields": columns, "data": result_list, + return {"fields": columns, "data": result_list, "fields_info": fields_info, "sql": bytes.decode(base64.b64encode(bytes(sql, 'utf-8')))} except Exception as ex: raise ParseSQLResultError(str(ex)) elif equals_ignore_case(ds.type, 'es'): try: - res, columns = get_es_data_by_http(conf, sql) - columns = [field.get('name') for field in columns] if origin_column else [field.get('name').lower() for - field in - columns] + res, raw_columns = get_es_data_by_http(conf, sql) + columns = [field.get('name') for field in raw_columns] if origin_column else [field.get('name').lower() + for + field in + raw_columns] + fields_info = build_fields_info_from_es(raw_columns, origin_column) result_list = [ {str(columns[i]): convert_value(value) for i, value in enumerate(tuple_item)} for tuple_item in res ] - return {"fields": columns, "data": result_list, + return {"fields": columns, "data": result_list, "fields_info": fields_info, "sql": bytes.decode(base64.b64encode(bytes(sql, 'utf-8')))} except Exception as ex: raise Exception(str(ex)) elif equals_ignore_case(ds.type, 'hive'): - with hive.connect(host=conf.host, port=conf.port, username=conf.username, - database=conf.database, **extra_config_dict) as conn, conn.cursor() as cursor: + with get_driver_pool(ds).connection() as conn, conn.cursor() as cursor: try: # Hive uses backticks for identifiers; normalize quoted identifiers as a compatibility fallback. hive_sql = re.sub(r'"([A-Za-z_][A-Za-z0-9_]*)"', r'`\1`', sql) @@ -715,21 +890,191 @@ def exec_sql(ds: CoreDatasource | AssistantOutDsSchema, sql: str, origin_column= columns = [field[0] for field in cursor.description] if origin_column else [field[0].lower() for field in cursor.description] + fields_info = build_fields_info_from_cursor(cursor, origin_column, 'hive') result_list = [ {str(columns[i]): convert_value(value) for i, value in enumerate(tuple_item)} for tuple_item in res ] - return {"fields": columns, "data": result_list, + return {"fields": columns, "data": result_list, "fields_info": fields_info, "sql": bytes.decode(base64.b64encode(bytes(hive_sql, 'utf-8')))} except Exception as ex: raise ParseSQLResultError(str(ex)) -def check_sql_read(sql: str, ds: CoreDatasource | AssistantOutDsSchema): +def build_fields_info_from_cursor(cursor, origin_column, db_type='postgresql'): + """ + 根据数据库游标的 description 构建字段信息列表 + + Args: + cursor: 数据库游标对象 + origin_column: 是否保留原始列名大小写 + db_type: 数据库类型,支持 'mysql', 'postgresql', 'redshift', 'kingbase', 'dm', 'hive' + + Returns: + list: 包含字段名和是否数值类型的字典列表 + """ + fields_info = [] + + for col_info in cursor.description: + col_name = col_info[0] + + if db_type in ('mysql', 'mariadb', 'doris', 'starrocks'): + # MySQL/pymysql 类型码 + is_numeric = col_info[1] in ( + 1, # TINYINT + 2, # SMALLINT + 3, # INT + 4, # FLOAT + 5, # DOUBLE + 8, # BIGINT + 9, # MEDIUMINT + 16, # BIT + 246, # DECIMAL/NEWDECIMAL + ) + elif db_type in ('postgresql', 'redshift', 'kingbase'): + # PostgreSQL/psycopg2 类型 OID + is_numeric = col_info[1] in ( + 16, # bool + 20, # int8 + 21, # int2 + 23, # int4 + 700, # float4 + 701, # float8 + 790, # money + 1700, # numeric + ) + elif db_type == 'dm': + # 达梦数据库类型码 + # 获取类型类名 + type_name = col_info[1].__name__.upper() if hasattr(col_info[1], '__name__') else str(col_info[1]).upper() + + is_numeric = type_name in { + 'INT', 'INTEGER', + 'BIGINT', 'SMALLINT', 'TINYINT', + 'NUMBER', 'NUMERIC', 'DECIMAL', + 'FLOAT', 'DOUBLE', 'REAL', + 'BIT', 'BOOLEAN', + } + elif db_type == 'hive': + # Hive 类型对象转字符串判断 + type_str = str(col_info[1]).lower() + NUMERIC_PREFIXES = ('tinyint', 'smallint', 'int', 'bigint', 'float', 'double', 'decimal', 'numeric') + is_numeric = type_str == 'boolean' or any(type_str.startswith(p) for p in NUMERIC_PREFIXES) + else: + is_numeric = False + + fields_info.append({ + "name": col_name if origin_column else col_name.lower(), + "is_numeric": is_numeric + }) + + return fields_info + + +def build_fields_info_from_es(raw_columns, origin_column): + """ + 专门为 Elasticsearch 构建字段信息 + + Args: + raw_columns: ES 返回的列信息列表 + origin_column: 是否保留原始列名大小写 + + Returns: + list: 包含字段名和是否数值类型的字典列表 + """ + fields_info = [] + + for field in raw_columns: + field_name = field.get('name') if origin_column else field.get('name').lower() + field_type = field.get('type', '').lower() + + is_numeric = field_type in ( + 'long', 'integer', 'short', 'byte', + 'double', 'float', 'half_float', 'scaled_float', + 'unsigned_long', 'boolean' + ) + + fields_info.append({ + "name": field_name, + "is_numeric": is_numeric + }) + + return fields_info + + +def get_sqlglot_dialect(ds_type: str) -> str: + """根据数据源类型获取 sqlglot dialect""" + if equals_ignore_case(ds_type, 'mysql', 'doris', 'starrocks'): + return 'mysql' + elif equals_ignore_case(ds_type, 'sqlServer'): + return 'tsql' + elif equals_ignore_case(ds_type, 'hive'): + return 'hive' + return None + + +# 通用危险函数(适用于所有数据库) +COMMON_DANGEROUS_FUNCTIONS = {'version', 'current_user', 'user', 'database'} + +# 特定数据库的危险函数 +DS_SPECIFIC_DANGEROUS_FUNCTIONS = { + 'mysql': {'LOAD_FILE', 'INTO OUTFILE', 'INTO DUMPFILE'}, + 'doris': {'LOAD_FILE', 'INTO OUTFILE', 'INTO DUMPFILE'}, + 'starrocks': {'LOAD_FILE', 'INTO OUTFILE', 'INTO DUMPFILE'}, + 'postgresql': {'pg_read_file', 'pg_write_file', 'lo_import', 'lo_export'}, + 'sqlserver': {'EXEC', 'xp_cmdshell', 'sp_executesql'}, + 'oracle': {'UTL_FILE', 'DBMS_PIPE', 'DBMS_LOCK'}, + 'hive': {'ADD FILE', 'ADD JAR'}, +} + +# 危险模式正则表达式(用于检查特殊语法) +import re + +DANGEROUS_PATTERNS = [ + r'\bINTO\s+OUTFILE\b', + r'\bINTO\s+DUMPFILE\b', + r'\bEXEC\s*\(', + r'\bCOPY\s+.*\bTO\s+PROGRAM\b', +] + + +def get_dangerous_functions(ds_type: str) -> set: + """获取危险函数(通用 + 特定数据源)""" + functions = COMMON_DANGEROUS_FUNCTIONS.copy() + ds_key = ds_type.lower() if ds_type else '' + if ds_key in DS_SPECIFIC_DANGEROUS_FUNCTIONS: + functions.update(DS_SPECIFIC_DANGEROUS_FUNCTIONS[ds_key]) + return functions + + +def check_dangerous_functions(statements: list, ds_type: str) -> bool: + """检查是否使用了危险函数,返回 True 表示安全""" + dangerous_functions = get_dangerous_functions(ds_type) + dangerous_functions_upper = {f.upper() for f in dangerous_functions} + + for stmt in statements: + if stmt: + for func in stmt.find_all(exp.Anonymous): + if func.name.upper() in dangerous_functions_upper: + return False + return True + + +def check_sql_read(sql: str, ds: CoreDatasource | AssistantOutDsSchema) -> tuple[bool, str]: + """ + 检查 SQL 是否为安全的只读查询 + 返回: (是否安全, 错误原因) + """ try: normalized_sql = sql.strip().lstrip("(").strip() first_keyword = normalized_sql.split(None, 1)[0].upper() if normalized_sql else "" - allowed_read_commands = {"SELECT", "WITH", "SHOW", "DESCRIBE", "DESC", "EXPLAIN"} + + # 根据配置决定是否允许元数据查询 + if settings.SQLBOT_ALLOW_METADATA_QUERIES: + allowed_read_commands = {"SELECT", "WITH", "SHOW", "DESCRIBE", "DESC", "EXPLAIN"} + else: + allowed_read_commands = {"SELECT", "WITH"} + denied_write_commands = { "INSERT", "UPDATE", "DELETE", "CREATE", "DROP", "ALTER", "TRUNCATE", "MERGE", "COPY", "REPLACE", "GRANT", "REVOKE", @@ -739,21 +1084,29 @@ def check_sql_read(sql: str, ds: CoreDatasource | AssistantOutDsSchema): if not first_keyword: raise ValueError("Parse SQL Error") if first_keyword in denied_write_commands: - return False + return False, f"Write operation '{first_keyword}' is not allowed" - dialect = None - if equals_ignore_case(ds.type, 'mysql', 'doris', 'starrocks'): - dialect = 'mysql' - elif equals_ignore_case(ds.type, 'sqlServer'): - dialect = 'tsql' - elif equals_ignore_case(ds.type, 'hive'): - dialect = 'hive' + # 1. 使用正则检查特殊模式 + for pattern in DANGEROUS_PATTERNS: + if re.search(pattern, sql, re.IGNORECASE): + return False, f"SQL contains dangerous pattern: {pattern}" + dialect = get_sqlglot_dialect(ds.type) statements = sqlglot.parse(sql, dialect=dialect) if not statements: raise ValueError("Parse SQL Error") + # 2. 使用 sqlglot 检查函数调用 + dangerous_functions = get_dangerous_functions(ds.type) + dangerous_functions_upper = {f.upper() for f in dangerous_functions} + for stmt in statements: + if stmt: + for func in stmt.find_all(exp.Anonymous): + if func.name.upper() in dangerous_functions_upper: + return False, f"SQL contains dangerous function: {func.name}" + + # 3. 检查写操作类型 write_types = ( exp.Insert, exp.Update, exp.Delete, exp.Create, exp.Drop, exp.Alter, @@ -764,9 +1117,12 @@ def check_sql_read(sql: str, ds: CoreDatasource | AssistantOutDsSchema): if stmt is None: continue if isinstance(stmt, write_types): - return False + return False, f"SQL contains write operation: {type(stmt).__name__}" - return first_keyword in allowed_read_commands + if first_keyword not in allowed_read_commands: + return False, f"SQL command '{first_keyword}' is not allowed. Only SELECT and WITH are permitted" + + return True, "" except Exception as e: raise ValueError(f"Parse SQL Error: {e}") @@ -779,3 +1135,120 @@ def checkParams(extraParams: str, illegalParams: List[str]): k, v = kv.split('=') if k in illegalParams: raise HTTPException(status_code=500, detail=f'Illegal Parameter: {k}') + + +import threading +from collections import OrderedDict + + +class ConnectionPoolManager: + def __init__(self, max_pools=500): + """ + init + :param max_pools: max pool + """ + self.max_pools = max_pools + self._pools = OrderedDict() # 使用有序字典实现 LRU + self._lock = threading.Lock() # 保证多线程安全 + + def get_pool(self, ds: CoreDatasource | AssistantOutDsSchema, **db_config): + """ + get connection pool(lazy load + LRU update) + """ + with self._lock: + if ds.id: + # 1. 如果连接池已存在,将其移动到字典末尾(标记为最近使用) + if ds.id in self._pools: + self._pools.move_to_end(ds.id) + print(f"[LRU] return: {ds.id}") + return self._pools[ds.id] + + # 2. 如果连接池不存在,检查是否达到上限,若达到则淘汰最久未使用的(字典头部) + if len(self._pools) >= self.max_pools: + oldest_id, oldest_pool = self._pools.popitem(last=False) + oldest_pool.close() # 安全关闭被驱逐的连接池 + print(f"[LRU] remove oldest: {oldest_id}") + + # 3. 创建新连接池并放入字典末尾 + engine = get_engine(ds, use_pool=True) + new_pool = sessionmaker(bind=engine) + self._pools[ds.id] = new_pool + print(f"[LRU] create: {ds.id}") + return new_pool + + def remove_pool(self, datasource_id): + with self._lock: + if datasource_id in self._pools: + # 1. 从字典中移除并获取该连接池对象 + pool = self._pools.pop(datasource_id) + # 2. 安全关闭该连接池,释放底层所有数据库连接和内存 + pool.close() + print(f"[Manager] Closed pool and remove: {datasource_id}") + else: + print(f"[Manager] Warning: ds id {datasource_id} not exist in sqlalchemy") + + def close_all(self): + """stop""" + with self._lock: + for pool in self._pools.values(): + pool.close() + self._pools.clear() + + +pool_manager = ConnectionPoolManager(max_pools=500) + + +class DriverConnectionPoolManager: + def __init__(self, max_pools=500): + """ + init + :param max_pools: max pool + """ + self.max_pools = max_pools + self._pools = OrderedDict() # 使用有序字典实现 LRU + self._lock = threading.Lock() # 保证多线程安全 + + def get_pool(self, ds: CoreDatasource | AssistantOutDsSchema, db_config): + """ + get connection pool(lazy load + LRU update) + """ + with self._lock: + if ds.id: + # 1. 如果连接池已存在,将其移动到字典末尾(标记为最近使用) + if ds.id in self._pools: + self._pools.move_to_end(ds.id) + print(f"[LRU] return: {ds.id}") + return self._pools[ds.id] + + # 2. 如果连接池不存在,检查是否达到上限,若达到则淘汰最久未使用的(字典头部) + if len(self._pools) >= self.max_pools: + oldest_id, oldest_pool = self._pools.popitem(last=False) + oldest_pool.close() # 安全关闭被驱逐的连接池 + print(f"[LRU] remove oldest: {oldest_id}") + + # 3. 创建新连接池并放入字典末尾 + new_pool = get_driver_connection(ds, db_config, use_pool=True) + self._pools[ds.id] = new_pool + print(f"[LRU] create: {ds.id}") + return new_pool + + def remove_pool(self, datasource_id): + with self._lock: + if datasource_id in self._pools: + # 1. 从字典中移除并获取该连接池对象 + pool = self._pools.pop(datasource_id) + # 2. 安全关闭该连接池,释放底层所有数据库连接和内存 + pool.close() + print(f"[Manager] Closed pool and remove: {datasource_id}") + else: + print(f"[Manager] Warning: ds id {datasource_id} not exist in dbutils") + + def close_all(self): + """stop""" + with self._lock: + for pool in self._pools.values(): + pool.close() + self._pools.clear() + + +driver_pool_manager = DriverConnectionPoolManager(max_pools=500) diff --git a/backend/apps/db/db_sql.py b/backend/apps/db/db_sql.py index fa790c8a8..496075378 100644 --- a/backend/apps/db/db_sql.py +++ b/backend/apps/db/db_sql.py @@ -162,13 +162,6 @@ def get_table_sql(ds: CoreDatasource, conf: DatasourceConf, db_version: str = '' """, conf.dbSchema elif equals_ignore_case(ds.type, "es"): return "", None - elif equals_ignore_case(ds.type, "sqlite"): - return """ - SELECT name AS TABLE_NAME, '' - FROM sqlite_master - WHERE type='table' - ORDER BY name - """, None elif equals_ignore_case(ds.type, "hive"): return """ SHOW TABLES @@ -323,9 +316,6 @@ def get_field_sql(ds: CoreDatasource, conf: DatasourceConf, table_name: str = No return sql1 + sql2, conf.dbSchema, table_name elif equals_ignore_case(ds.type, "es"): return "", None, None - elif equals_ignore_case(ds.type, "sqlite"): - sql1 = f"PRAGMA table_info({table_name})" - return sql1, None, None elif equals_ignore_case(ds.type, "hive"): sql1 = f"DESCRIBE {table_name}" return sql1, None, None diff --git a/backend/apps/mcp/mcp.py b/backend/apps/mcp/mcp.py index 4afe1465e..2be45074c 100644 --- a/backend/apps/mcp/mcp.py +++ b/backend/apps/mcp/mcp.py @@ -13,7 +13,7 @@ from apps.chat.api.chat import create_chat, question_answer_inner from apps.chat.models.chat_model import ChatMcp, CreateChat, ChatStart, McpQuestion, McpAssistant, ChatQuestion, \ - ChatFinishStep, McpDs + ChatFinishStep, McpDs, ChatToken from apps.datasource.crud.datasource import get_datasource_list from apps.system.crud.user import authenticate, user_ws_options from apps.system.crud.user import get_db_user @@ -21,6 +21,9 @@ from apps.system.models.user import UserModel from apps.system.schemas.system_schema import BaseUserDTO, AssistantHeader from apps.system.schemas.system_schema import UserInfoDTO +from common.audit.models.log_model import OperationType, OperationModules +from common.audit.schemas.logger_decorator import LogConfig, system_log +from common.audit.schemas.request_context import RequestContext from common.core import security from common.core.config import settings from common.core.deps import SessionDep, Trans @@ -34,19 +37,21 @@ router = APIRouter(tags=["mcp"], prefix="/mcp") -# @router.post("/access_token", operation_id="access_token") -# def local_login( -# session: SessionDep, -# form_data: Annotated[OAuth2PasswordRequestForm, Depends()] -# ) -> Token: -# user = authenticate(session=session, account=form_data.username, password=form_data.password) -# if not user: -# raise HTTPException(status_code=400, detail="Incorrect account or password") -# access_token_expires = timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) -# user_dict = user.to_dict() -# return Token(access_token=create_access_token( -# user_dict, expires_delta=access_token_expires -# )) +@router.post("/access_token", operation_id="access_token") +async def access_token(session: SessionDep, chat: ChatToken): + user: BaseUserDTO = authenticate(session=session, account=chat.username, password=chat.password) + if not user: + raise HTTPException(status_code=400, detail="Incorrect account or password") + + if not user.oid or user.oid == 0: + raise HTTPException(status_code=400, detail="No associated workspace, Please contact the administrator") + access_token_expires = timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) + user_dict = user.to_dict() + t = Token(access_token=create_access_token( + user_dict, expires_delta=access_token_expires + )) + # c = create_chat(session, user, CreateChat(origin=1), False) + return {"access_token": t.access_token} def get_user(session: SessionDep, token: str): @@ -82,20 +87,45 @@ def get_user(session: SessionDep, token: str): @router.post("/mcp_start", operation_id="mcp_start") -async def mcp_start(session: SessionDep, chat: ChatStart): - user: BaseUserDTO = authenticate(session=session, account=chat.username, password=chat.password) - if not user: - raise HTTPException(status_code=400, detail="Incorrect account or password") +@system_log(LogConfig( + operation_type=OperationType.CREATE, + module=OperationModules.CHAT, + result_id_expr="id", + save_on_success_only=True +)) +async def mcp_start(session: SessionDep, trans: Trans, chat: ChatStart): + res_token = None + user = None + if chat.token: + res_token = chat.token + user = get_user(session, chat.token) + else: + user = authenticate(session=session, account=chat.username, password=chat.password) + if not user: + raise HTTPException(status_code=400, detail="Incorrect account or password") + + if not user.oid or user.oid == 0: + raise HTTPException(status_code=400, detail="No associated workspace, Please contact the administrator") + access_token_expires = timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) + user_dict = user.to_dict() + t = Token(access_token=create_access_token( + user_dict, expires_delta=access_token_expires + )) + res_token = t.access_token + + if chat.oid: + w_list = await user_ws_options(session, user.id, trans) + oid_list = [item.id for item in w_list] + if int(chat.oid) not in oid_list: + raise HTTPException(status_code=400, detail="The current user is not in the selected workspace") + + user.oid = int(chat.oid) + + request = RequestContext.get_request() + request.state.current_user = user - if not user.oid or user.oid == 0: - raise HTTPException(status_code=400, detail="No associated workspace, Please contact the administrator") - access_token_expires = timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) - user_dict = user.to_dict() - t = Token(access_token=create_access_token( - user_dict, expires_delta=access_token_expires - )) c = create_chat(session, user, CreateChat(origin=1), False) - return {"access_token": t.access_token, "chat_id": c.id} + return {"access_token": res_token, "chat_id": c.id} @router.post("/mcp_ws_list", operation_id="mcp_ws_list") @@ -105,9 +135,14 @@ async def ws_list(session: SessionDep, trans: Trans, token: str): @router.post("/mcp_ds_list", operation_id="mcp_datasource_list") -async def datasource_list(session: SessionDep, mcp_ds: McpDs): +async def datasource_list(session: SessionDep, trans: Trans, mcp_ds: McpDs): session_user = get_user(session, mcp_ds.token) if mcp_ds.oid: + w_list = await user_ws_options(session, session_user.id, trans) + oid_list = [item.id for item in w_list] + if int(mcp_ds.oid) not in oid_list: + raise HTTPException(status_code=400, detail="The current user is not in the selected workspace") + session_user.oid = int(mcp_ds.oid) ds_list = get_datasource_list(session=session, user=session_user) result = [] @@ -129,13 +164,18 @@ async def datasource_list(session: SessionDep, mcp_ds: McpDs): @router.post("/mcp_question", operation_id="mcp_question") -async def mcp_question(session: SessionDep, chat: McpQuestion): +async def mcp_question(session: SessionDep, trans: Trans, chat: McpQuestion): session_user = get_user(session, chat.token) lang = chat.lang if lang in ["zh-CN", "zh-TW", "en", "ko-KR"]: session_user.language = lang - if chat.oid: - session_user.oid = int(chat.oid) + # if chat.oid: + # w_list = await user_ws_options(session, session_user.id, trans) + # oid_list = [item.id for item in w_list] + # if int(chat.oid) not in oid_list: + # raise HTTPException(status_code=400, detail="The current user is not in the selected workspace") + # + # session_user.oid = int(chat.oid) ds_id: Optional[int] = None if chat.datasource_id: if isinstance(chat.datasource_id, str): diff --git a/backend/apps/swagger/locales/en.json b/backend/apps/swagger/locales/en.json index 512f7d2cc..5d84258f4 100644 --- a/backend/apps/swagger/locales/en.json +++ b/backend/apps/swagger/locales/en.json @@ -92,6 +92,13 @@ "system_model_create": "Save Model", "system_model_update": "Update Model", "system_model_del": "Delete Model", + "enable_custom_model": "Enable Custom Model", + "custom_model": "Custom Model", + "system_model_ws_mapping": "Query Model-Workspace Authorization Relationships", + "system_model_ws_mapping_update": "Update Model-Workspace Authorization Relationships", + "system_model_ws_mapping_add": "Add Model-Workspace Authorization Relationships", + "system_model_ws_mapping_delete": "Delete Model-Workspace Authorization Relationships", + "system_model_list_by_ws": "Get Model List by Workspace", "model_name": "Name", "model_type": "Type", "base_model": "Base Model", diff --git a/backend/apps/swagger/locales/zh.json b/backend/apps/swagger/locales/zh.json index b9552c4d7..397fda5e3 100644 --- a/backend/apps/swagger/locales/zh.json +++ b/backend/apps/swagger/locales/zh.json @@ -92,6 +92,13 @@ "system_model_create": "保存模型", "system_model_update": "更新模型", "system_model_del": "删除模型", + "enable_custom_model": "启用自定义模型", + "custom_model": "自定义模型", + "system_model_ws_mapping": "查询模型授权工作空间的关联关系", + "system_model_ws_mapping_update": "更新模型授权工作空间的关联关系", + "system_model_ws_mapping_add": "新增模型授权工作空间的关联关系", + "system_model_ws_mapping_delete": "删除模型授权工作空间的关联关系", + "system_model_list_by_ws": "根据工作空间获取模型列表", "model_name": "名称", "model_type": "类型", "base_model": "基础模型", diff --git a/backend/apps/system/api/aimodel.py b/backend/apps/system/api/aimodel.py index ce37229e2..620adac7b 100644 --- a/backend/apps/system/api/aimodel.py +++ b/backend/apps/system/api/aimodel.py @@ -1,16 +1,17 @@ import json from typing import List, Union +from fastapi import APIRouter, Path, Query, Body from fastapi.responses import StreamingResponse +from sqlmodel import func, select, update, delete + from apps.ai_model.model_factory import LLMConfig, LLMFactory from apps.swagger.i18n import PLACEHOLDER_PREFIX +from apps.system.crud.aimodel_manage import get_ai_model_list_by_workspace +from apps.system.models.system_model import AiModelDetail, AiModelWorkspaceMapping, AiModelBrief from apps.system.schemas.ai_model_schema import AiModelConfigItem, AiModelCreator, AiModelEditor, AiModelGridItem -from fastapi import APIRouter, Path, Query -from sqlmodel import func, select, update - -from apps.system.models.system_model import AiModelDetail from apps.system.schemas.permission import SqlbotPermission, require_permissions -from common.core.deps import SessionDep, Trans +from common.core.deps import SessionDep, Trans, CurrentUser from common.utils.crypto import sqlbot_decrypt from common.utils.time import get_timestamp from common.utils.utils import SQLBotLogUtil, prepare_model_arg @@ -19,12 +20,14 @@ from common.audit.models.log_model import OperationType, OperationModules from common.audit.schemas.logger_decorator import LogConfig, system_log + @router.post("/status", include_in_schema=False) -@require_permissions(permission=SqlbotPermission(role=['admin'])) +@require_permissions(permission=SqlbotPermission(role=['admin'])) async def check_llm(info: AiModelCreator, trans: Trans): async def generate(): try: - additional_params = {item.key: prepare_model_arg(item.val) for item in info.config_list if item.key and item.val} + additional_params = {item.key: prepare_model_arg(item.val) for item in info.config_list if + item.key and item.val} config = LLMConfig( model_type="openai" if info.protocol == 1 else "vllm", model_name=info.base_model, @@ -39,14 +42,15 @@ async def generate(): yield json.dumps({"content": chunk}) + "\n" if chunk and isinstance(chunk, dict) and chunk.content: yield json.dumps({"content": chunk.content}) + "\n" - + except Exception as e: SQLBotLogUtil.error(f"Error checking LLM: {e}") error_msg = trans('i18n_llm.validate_error', msg=str(e)) yield json.dumps({"error": error_msg}) + "\n" - + return StreamingResponse(generate(), media_type="application/x-ndjson") + @router.get("/default", include_in_schema=False) async def check_default(session: SessionDep, trans: Trans): db_model = session.exec( @@ -54,8 +58,10 @@ async def check_default(session: SessionDep, trans: Trans): ).first() if not db_model: raise Exception(trans('i18n_llm.miss_default')) - -@router.put("/default/{id}", summary=f"{PLACEHOLDER_PREFIX}system_model_default", description=f"{PLACEHOLDER_PREFIX}system_model_default") + + +@router.put("/default/{id}", summary=f"{PLACEHOLDER_PREFIX}system_model_default", + description=f"{PLACEHOLDER_PREFIX}system_model_default") @require_permissions(permission=SqlbotPermission(role=['admin'])) @system_log(LogConfig(operation_type=OperationType.UPDATE, module=OperationModules.AI_MODEL, resource_id_expr="id")) async def set_default(session: SessionDep, id: int = Path(description="ID")): @@ -76,27 +82,46 @@ async def set_default(session: SessionDep, id: int = Path(description="ID")): session.rollback() raise e -@router.get("", response_model=list[AiModelGridItem], summary=f"{PLACEHOLDER_PREFIX}system_model_grid", description=f"{PLACEHOLDER_PREFIX}system_model_grid") -@require_permissions(permission=SqlbotPermission(role=['admin'])) + +@router.get("", response_model=list[AiModelGridItem], summary=f"{PLACEHOLDER_PREFIX}system_model_grid", + description=f"{PLACEHOLDER_PREFIX}system_model_grid") +@require_permissions(permission=SqlbotPermission(role=['admin'])) async def query( session: SessionDep, keyword: Union[str, None] = Query(default=None, max_length=255, description=f"{PLACEHOLDER_PREFIX}keyword") ): - statement = select(AiModelDetail.id, - AiModelDetail.name, - AiModelDetail.model_type, - AiModelDetail.base_model, - AiModelDetail.supplier, - AiModelDetail.protocol, - AiModelDetail.default_model) + # 子查询:统计每个 model 绑定的 workspace 数量 + count_sub = ( + select( + AiModelWorkspaceMapping.ai_model_id, + func.count().label("ws_mapping_count") + ) + .group_by(AiModelWorkspaceMapping.ai_model_id) + .subquery() + ) + statement = ( + select( + AiModelDetail.id, + AiModelDetail.name, + AiModelDetail.model_type, + AiModelDetail.base_model, + AiModelDetail.supplier, + AiModelDetail.protocol, + AiModelDetail.default_model, + func.coalesce(count_sub.c.ws_mapping_count, 0).label("ws_mapping_count"), + ) + .outerjoin(count_sub, AiModelDetail.id == count_sub.c.ai_model_id) + ) if keyword is not None: statement = statement.where(AiModelDetail.name.like(f"%{keyword}%")) statement = statement.order_by(AiModelDetail.default_model.desc(), AiModelDetail.name, AiModelDetail.create_time) items = session.exec(statement).all() return items -@router.get("/{id}", response_model=AiModelEditor, summary=f"{PLACEHOLDER_PREFIX}system_model_query", description=f"{PLACEHOLDER_PREFIX}system_model_query") -@require_permissions(permission=SqlbotPermission(role=['admin'])) + +@router.get("/{id}", response_model=AiModelEditor, summary=f"{PLACEHOLDER_PREFIX}system_model_query", + description=f"{PLACEHOLDER_PREFIX}system_model_query") +@require_permissions(permission=SqlbotPermission(role=['admin'])) async def get_model_by_id( session: SessionDep, id: int = Path(description="ID") @@ -124,7 +149,9 @@ async def get_model_by_id( data["config_list"] = config_list return AiModelEditor(**data) -@router.post("", summary=f"{PLACEHOLDER_PREFIX}system_model_create", description=f"{PLACEHOLDER_PREFIX}system_model_create") + +@router.post("", summary=f"{PLACEHOLDER_PREFIX}system_model_create", + description=f"{PLACEHOLDER_PREFIX}system_model_create") @require_permissions(permission=SqlbotPermission(role=['admin'])) @system_log(LogConfig(operation_type=OperationType.CREATE, module=OperationModules.AI_MODEL, result_id_expr="id")) async def add_model( @@ -143,9 +170,12 @@ async def add_model( session.commit() return detail -@router.put("", summary=f"{PLACEHOLDER_PREFIX}system_model_update", description=f"{PLACEHOLDER_PREFIX}system_model_update") + +@router.put("", summary=f"{PLACEHOLDER_PREFIX}system_model_update", + description=f"{PLACEHOLDER_PREFIX}system_model_update") @require_permissions(permission=SqlbotPermission(role=['admin'])) -@system_log(LogConfig(operation_type=OperationType.UPDATE, module=OperationModules.AI_MODEL, resource_id_expr="editor.id")) +@system_log( + LogConfig(operation_type=OperationType.UPDATE, module=OperationModules.AI_MODEL, resource_id_expr="editor.id")) async def update_model( session: SessionDep, editor: AiModelEditor @@ -155,12 +185,14 @@ async def update_model( data["config"] = json.dumps([item.model_dump(exclude_unset=True) for item in editor.config_list]) data.pop("config_list", None) db_model = session.get(AiModelDetail, id) - #update_data = AiModelDetail.model_validate(data) + # update_data = AiModelDetail.model_validate(data) db_model.sqlmodel_update(data) session.add(db_model) session.commit() -@router.delete("/{id}", summary=f"{PLACEHOLDER_PREFIX}system_model_del", description=f"{PLACEHOLDER_PREFIX}system_model_del") + +@router.delete("/{id}", summary=f"{PLACEHOLDER_PREFIX}system_model_del", + description=f"{PLACEHOLDER_PREFIX}system_model_del") @require_permissions(permission=SqlbotPermission(role=['admin'])) @system_log(LogConfig(operation_type=OperationType.DELETE, module=OperationModules.AI_MODEL, resource_id_expr="id")) async def delete_model( @@ -170,9 +202,161 @@ async def delete_model( ): item = session.get(AiModelDetail, id) if item.default_model: - raise Exception(trans('i18n_llm.delete_default_error', key = item.name)) + raise Exception(trans('i18n_llm.delete_default_error', key=item.name)) session.delete(item) session.commit() - - \ No newline at end of file + +@router.get("/{id}/ws_mapping", response_model=List[str], summary=f"{PLACEHOLDER_PREFIX}system_model_ws_mapping", + description=f"{PLACEHOLDER_PREFIX}system_model_ws_mapping") +@require_permissions(permission=SqlbotPermission(role=['admin'])) +async def get_model_ws_mapping_by_id( + session: SessionDep, + id: int = Path(description="ID") +): + db_model = session.get(AiModelDetail, id) + if not db_model: + raise ValueError(f"AiModelDetail with id {id} not found") + + # 根据 ai_model_id 查询关联的 workspace_id 列表 + stmt = ( + select(AiModelWorkspaceMapping.workspace_id) + .where(AiModelWorkspaceMapping.ai_model_id == id) + .distinct() + ) + ws_ids: List[int] = session.exec(stmt).all() + + return [str(ws_id) for ws_id in ws_ids] + + +@router.put("/{id}/ws_mapping", response_model=List[str], summary=f"{PLACEHOLDER_PREFIX}system_model_ws_mapping_update", + description=f"{PLACEHOLDER_PREFIX}system_model_ws_mapping_update") +@require_permissions(permission=SqlbotPermission(role=['admin'])) +async def update_model_ws_mapping_by_id( + session: SessionDep, + id: int = Path(description="ID"), + ws_ids: List[str] = Body(description="workspace id list"), +): + if ws_ids is None: + ws_ids = [] + # 提前去重 + ws_ids = list({int(ws_id) for ws_id in ws_ids}) + + db_model = session.get(AiModelDetail, id) + if not db_model: + raise ValueError(f"AiModelDetail with id {id} not found") + + # 根据 ai_model_id 更新关联的 workspace_id 列表 + # 1. 批量删除旧映射 + session.execute( + delete(AiModelWorkspaceMapping) + .where(AiModelWorkspaceMapping.ai_model_id == id) + ) + + # 2. 插入去重后的映射关系 + for ws_id in ws_ids: + session.add( + AiModelWorkspaceMapping(ai_model_id=id, workspace_id=ws_id) + ) + + session.commit() + + return [str(ws_id) for ws_id in ws_ids] + + +# 新增映射(在已有基础上追加) +@router.post("/{id}/ws_mapping", response_model=List[str], summary=f"{PLACEHOLDER_PREFIX}system_model_ws_mapping_add", + description=f"{PLACEHOLDER_PREFIX}system_model_ws_mapping_add") +@require_permissions(permission=SqlbotPermission(role=['admin'])) +async def add_model_ws_mapping_by_id( + session: SessionDep, + id: int = Path(description="ID"), + ws_ids: List[str] = Body(description="workspace id list"), +): + if ws_ids is None: + ws_ids = [] + ws_ids = list({int(ws_id) for ws_id in ws_ids}) + + db_model = session.get(AiModelDetail, id) + if not db_model: + raise ValueError(f"AiModelDetail with id {id} not found") + + # 查询已存在的映射,过滤掉重复的 + existing_stmt = ( + select(AiModelWorkspaceMapping.workspace_id) + .where( + AiModelWorkspaceMapping.ai_model_id == id, + AiModelWorkspaceMapping.workspace_id.in_(ws_ids), + ) + ) + existing_ws_ids = set(session.exec(existing_stmt).all()) + + # 只插入不存在的映射 + new_ws_ids = [ws_id for ws_id in ws_ids if ws_id not in existing_ws_ids] + for ws_id in new_ws_ids: + session.add( + AiModelWorkspaceMapping(ai_model_id=id, workspace_id=ws_id) + ) + + session.commit() + + # 返回完整的映射列表 + all_stmt = ( + select(AiModelWorkspaceMapping.workspace_id) + .where(AiModelWorkspaceMapping.ai_model_id == id) + .distinct() + ) + all_ws_ids: List[int] = session.exec(all_stmt).all() + + return [str(ws_id) for ws_id in all_ws_ids] + + +# 删除指定映射 +@router.delete("/{id}/ws_mapping", response_model=List[str], + summary=f"{PLACEHOLDER_PREFIX}system_model_ws_mapping_delete", + description=f"{PLACEHOLDER_PREFIX}system_model_ws_mapping_delete") +@require_permissions(permission=SqlbotPermission(role=['admin'])) +async def delete_model_ws_mapping_by_id( + session: SessionDep, + id: int = Path(description="ID"), + ws_ids: List[str] = Body(description="workspace id list"), +): + if ws_ids is None: + ws_ids = [] + ws_ids = list({int(ws_id) for ws_id in ws_ids}) + + db_model = session.get(AiModelDetail, id) + if not db_model: + raise ValueError(f"AiModelDetail with id {id} not found") + + # 只删除指定的映射 + if ws_ids: + session.execute( + delete(AiModelWorkspaceMapping) + .where( + AiModelWorkspaceMapping.ai_model_id == id, + AiModelWorkspaceMapping.workspace_id.in_(ws_ids), + ) + ) + + session.commit() + + # 返回剩余的映射列表 + stmt = ( + select(AiModelWorkspaceMapping.workspace_id) + .where(AiModelWorkspaceMapping.ai_model_id == id) + .distinct() + ) + remaining_ws_ids: List[int] = session.exec(stmt).all() + + return [str(ws_id) for ws_id in remaining_ws_ids] + + +@router.get("/list/by_ws", response_model=List[AiModelBrief], summary=f"{PLACEHOLDER_PREFIX}system_model_list_by_ws", + description=f"{PLACEHOLDER_PREFIX}system_model_list_by_ws") +@require_permissions(permission=SqlbotPermission(role=['ws_admin'])) +async def get_model_by_ws( + session: SessionDep, + current_user: CurrentUser +): + return get_ai_model_list_by_workspace(session, current_user.oid, False) diff --git a/backend/apps/system/api/assistant.py b/backend/apps/system/api/assistant.py index 5db3f4032..138e534c2 100644 --- a/backend/apps/system/api/assistant.py +++ b/backend/apps/system/api/assistant.py @@ -1,6 +1,7 @@ import json import os -from datetime import timedelta +from datetime import datetime, timedelta, timezone +from zoneinfo import ZoneInfo from typing import List, Optional from fastapi import APIRouter, Form, HTTPException, Path, Query, Request, Response, UploadFile @@ -22,30 +23,83 @@ from common.core.security import create_access_token from common.core.sqlbot_cache import clear_cache from common.utils.utils import get_origin_from_referer, origin_match_domain - router = APIRouter(tags=["system_assistant"], prefix="/system/assistant") from common.audit.models.log_model import OperationType, OperationModules from common.audit.schemas.logger_decorator import LogConfig, system_log - +from sqlbot_xpack.core import decrypt_embedded_sign @router.get("/info/{id}", include_in_schema=False) -async def info(request: Request, response: Response, session: SessionDep, trans: Trans, id: int) -> AssistantModel: +async def info(request: Request, response: Response, session: SessionDep, trans: Trans, id: int, virtual: Optional[int] = Query(None)): if not id: raise Exception('miss assistant id') db_model = await get_assistant_info(session=session, assistant_id=id) if not db_model: raise RuntimeError(f"assistant application not exist") db_model = AssistantModel.model_validate(db_model) - - origin = request.headers.get("origin") or get_origin_from_referer(request) - if not origin: - raise RuntimeError(trans('i18n_embedded.invalid_origin', origin=origin or '')) - origin = origin.rstrip('/') - if not origin_match_domain(origin, db_model.domain): - raise RuntimeError(trans('i18n_embedded.invalid_origin', origin=origin or '')) - + + # 校验 SQLBOT-EMBEDDED-SIGN 请求头 + sign_header = request.headers.get("SQLBOT-EMBEDDED-SIGN") + if not sign_header: + raise RuntimeError(trans('i18n_embedded.invalid_origin', origin='')) + + sign_data = await decrypt_embedded_sign(sign_header) + + # 校验 assistant_id 与 id 参数一致 + if str(sign_data.get("assistant_id")) != str(id): + raise RuntimeError(trans('i18n_embedded.invalid_origin', origin='')) + + # 校验 target(来源域名)是否合法 + target = sign_data.get("target", "") + if not origin_match_domain(target, db_model.domain): + raise RuntimeError(trans('i18n_embedded.invalid_origin', origin=target or '')) + + # 校验 sign_time 是否在 10 秒内 + sign_time_str = sign_data.get("sign_time", "") + sign_time = datetime.fromisoformat(sign_time_str) + now_utc = datetime.now(timezone.utc) + sign_time_utc = sign_time.astimezone(timezone.utc) + if abs((now_utc - sign_time_utc).total_seconds()) > 10: + raise RuntimeError(trans('i18n_embedded.invalid_origin', origin=target or '')) + + # 校验是否为真实浏览器请求(非自动化工具) + if sign_data.get("webdriver", False): + raise RuntimeError(trans('i18n_embedded.invalid_origin', origin=target or '')) + + # 校验 User-Agent 一致性(签名中的 navigator.userAgent 与请求头一致) + sign_user_agent = sign_data.get("user_agent", "") + request_user_agent = request.headers.get("User-Agent", "") + if sign_user_agent != request_user_agent: + raise RuntimeError(trans('i18n_embedded.invalid_origin', origin=target or '')) + + # 校验 timezone 与 sign_time 偏移一致性(防时区伪造) + tz_name = sign_data.get("timezone", "") + tz = ZoneInfo(tz_name) + sign_time_naive = sign_time.replace(tzinfo=None) + if tz.utcoffset(sign_time_naive) != sign_time.utcoffset(): + raise RuntimeError(trans('i18n_embedded.invalid_origin', origin=target or '')) + + origin = target.rstrip('/') + response.headers["Access-Control-Allow-Origin"] = origin - return db_model + + + assistant_oid = 1 + if (db_model.type == 0): + configuration = db_model.configuration + config_obj = json.loads(configuration) if configuration else {} + assistant_oid = config_obj.get('oid', 1) + + access_token_expires = timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) + assistantDict = { + "id": virtual, "account": 'sqlbot-inner-assistant', "oid": assistant_oid, "assistant_id": id + } + access_token = create_access_token( + assistantDict, expires_delta=access_token_expires + ) + + result = db_model.model_dump() + result["token"] = access_token + return result @router.get("/app/{appId}", include_in_schema=False) @@ -67,7 +121,7 @@ async def getApp(request: Request, response: Response, session: SessionDep, tran return db_model -@router.get("/validator", response_model=AssistantValidator, include_in_schema=False) +""" @router.get("/validator", response_model=AssistantValidator, include_in_schema=False) async def validator(session: SessionDep, id: int, virtual: Optional[int] = Query(None)): if not id: raise Exception('miss assistant id') @@ -89,7 +143,7 @@ async def validator(session: SessionDep, id: int, virtual: Optional[int] = Query access_token = create_access_token( assistantDict, expires_delta=access_token_expires ) - return AssistantValidator(True, True, True, access_token) + return AssistantValidator(True, True, True, access_token) """ @router.get('/picture/{file_id}', summary=f"{PLACEHOLDER_PREFIX}assistant_picture_api", description=f"{PLACEHOLDER_PREFIX}assistant_picture_api") @@ -111,6 +165,7 @@ def iterfile(): @router.patch('/ui', summary=f"{PLACEHOLDER_PREFIX}assistant_ui_api", description=f"{PLACEHOLDER_PREFIX}assistant_ui_api") +@require_permissions(permission=SqlbotPermission(role=['ws_admin'])) @system_log(LogConfig(operation_type=OperationType.UPDATE, module=OperationModules.APPLICATION, result_id_expr="id")) async def ui(session: SessionDep, data: str = Form(), files: List[UploadFile] = []): json_data = json.loads(data) @@ -130,7 +185,7 @@ async def ui(session: SessionDep, data: str = Form(), files: List[UploadFile] = file.filename = file_name if flag_name == 'logo' or flag_name == 'float_icon': try: - SQLBotFileUtils.check_file(file=file, file_types=[".jpg", ".png", ".svg"], + SQLBotFileUtils.check_file(file=file, file_types=[".jpg", ".jpeg", ".png"], limit_file_size=(10 * 1024 * 1024)) except ValueError as e: error_msg = str(e) @@ -222,6 +277,8 @@ def get_db_type(type): async def query(session: SessionDep, current_user: CurrentUser): list_result = session.exec(select(AssistantModel).where(AssistantModel.oid == current_user.oid, AssistantModel.type != 4).order_by(AssistantModel.name, AssistantModel.create_time)).all() + for model in list_result: + model.enable_custom_model = model.enable_custom_model or False return list_result @@ -257,13 +314,13 @@ async def update(request: Request, session: SessionDep, editor: AssistantDTO): dynamic_upgrade_cors(request=request, session=session) -@router.get("/{id}", response_model=AssistantModel, summary=f"{PLACEHOLDER_PREFIX}assistant_query_api", description=f"{PLACEHOLDER_PREFIX}assistant_query_api") +""" @router.get("/{id}", response_model=AssistantModel, summary=f"{PLACEHOLDER_PREFIX}assistant_query_api", description=f"{PLACEHOLDER_PREFIX}assistant_query_api") async def get_one(session: SessionDep, id: int = Path(description="ID")): db_model = await get_assistant_info(session=session, assistant_id=id) if not db_model: raise ValueError(f"AssistantModel with id {id} not found") db_model = AssistantModel.model_validate(db_model) - return db_model + return db_model """ @router.delete("/{id}", summary=f"{PLACEHOLDER_PREFIX}assistant_del_api", description=f"{PLACEHOLDER_PREFIX}assistant_del_api") diff --git a/backend/apps/system/api/user.py b/backend/apps/system/api/user.py index 58033543a..20b379ef7 100644 --- a/backend/apps/system/api/user.py +++ b/backend/apps/system/api/user.py @@ -58,17 +58,34 @@ async def pager( status: Optional[int] = Query(None, description=f"{PLACEHOLDER_PREFIX}status"), origins: Optional[list[int]] = Query(None, description=f"{PLACEHOLDER_PREFIX}origin"), oidlist: Optional[list[int]] = Query(None, description=f"{PLACEHOLDER_PREFIX}oid"), + order_by: Optional[str] = Query(None, description="排序字段"), + desc: Optional[bool] = Query(False, description="是否降序"), ): pagination = PaginationParams(page=pageNum, size=pageSize) paginator = Paginator(session) - filters = {} - + + # 允许排序的字段白名单(防止 SQL 注入) + SORT_COLUMNS = { + 'account': UserModel.account, + 'create_time': UserModel.create_time, + 'name': UserModel.name, + 'email': UserModel.email, + 'status': UserModel.status, + } + sort_field = SORT_COLUMNS.get(order_by, UserModel.account) + sort_clause = sort_field.desc() if desc else sort_field.asc() + + # SELECT 列必须包含 ORDER BY 列(PostgreSQL DISTINCT 约束) + select_columns = [UserModel.id, UserModel.account] + if order_by and order_by != 'account': + select_columns.append(sort_field) + origin_stmt = ( - select(UserModel.id, UserModel.account) + select(*select_columns) .join(UserWsModel, UserModel.id == UserWsModel.uid, isouter=True) .where(UserModel.id != 1) .distinct() - .order_by(UserModel.account) + .order_by(sort_clause) ) if oidlist: @@ -89,8 +106,7 @@ async def pager( user_page = await paginator.get_paginated_response( stmt=origin_stmt, - pagination=pagination, - **filters) + pagination=pagination) uid_list = [item.get('id') for item in user_page.items] if not uid_list: return user_page @@ -98,7 +114,7 @@ async def pager( select(UserModel, UserWsModel.oid.label('ws_oid')) .join(UserWsModel, UserModel.id == UserWsModel.uid, isouter=True) .where(UserModel.id.in_(uid_list)) - .order_by(UserModel.account, UserModel.create_time) + .order_by(sort_clause) ) user_workspaces = session.exec(stmt).all() merged = defaultdict(list) diff --git a/backend/apps/system/api/variable_api.py b/backend/apps/system/api/variable_api.py index 751462391..ca5c952f5 100644 --- a/backend/apps/system/api/variable_api.py +++ b/backend/apps/system/api/variable_api.py @@ -8,27 +8,32 @@ from apps.system.models.system_variable_model import SystemVariable from common.core.config import settings from common.core.deps import SessionDep, CurrentUser, Trans +from apps.system.schemas.permission import SqlbotPermission, require_permissions router = APIRouter(tags=["System_variable"], prefix="/sys_variable") path = settings.EXCEL_PATH @router.post("/save", response_model=None, summary=f"{PLACEHOLDER_PREFIX}variable_save") +@require_permissions(permission=SqlbotPermission(role=['admin'])) async def save_variable(session: SessionDep, user: CurrentUser, trans: Trans, variable: SystemVariable): return save(session, user, trans, variable) @router.post("/delete",response_model=None, summary=f"{PLACEHOLDER_PREFIX}variable_delete") +@require_permissions(permission=SqlbotPermission(role=['admin'])) async def delete_variable(session: SessionDep, ids: List[int]): return delete(session, ids) @router.post("/listAll",response_model=None, summary=f"{PLACEHOLDER_PREFIX}variable_list") +@require_permissions(permission=SqlbotPermission(role=['ws_admin'])) async def list_all_data(session: SessionDep, trans: Trans, variable: SystemVariable = None): return list_all(session, trans, variable) @router.post("/listPage/{pageNum}/{pageSize}",response_model=None, summary=f"{PLACEHOLDER_PREFIX}variable_page") +@require_permissions(permission=SqlbotPermission(role=['admin'])) async def pager(session: SessionDep, trans: Trans, pageNum: int, pageSize: int, variable: SystemVariable = None): return await list_page(session, trans, pageNum, pageSize, variable) diff --git a/backend/apps/system/crud/aimodel_manage.py b/backend/apps/system/crud/aimodel_manage.py index 2d4f5b752..b0fabdbd1 100644 --- a/backend/apps/system/crud/aimodel_manage.py +++ b/backend/apps/system/crud/aimodel_manage.py @@ -1,10 +1,11 @@ +from sqlmodel import Session, select, or_ -from apps.system.models.system_model import AiModelDetail +from apps.system.models.system_model import AiModelDetail, AiModelBrief, AiModelWorkspaceMapping from common.core.db import engine -from sqlmodel import Session, select from common.utils.crypto import sqlbot_encrypt from common.utils.utils import SQLBotLogUtil + async def async_model_info(): with Session(engine) as session: model_list = session.exec(select(AiModelDetail)).all() @@ -27,7 +28,40 @@ async def async_model_info(): session.add(model) if any_model_change: session.commit() - SQLBotLogUtil.info("✅ 异步加密已有模型的密钥和地址完成") - - - \ No newline at end of file + SQLBotLogUtil.info("✅ 异步加密已有模型的密钥和地址完成") + + +def get_ai_model_list_by_workspace(session: Session, workspace_id: int, with_default: bool = True): + sub_stmt = ( + select(AiModelWorkspaceMapping.ai_model_id) + .where(AiModelWorkspaceMapping.workspace_id == workspace_id) + .distinct() + ) + + # 查询:关联的模型 + default_model 为 True 的模型,默认模型排第一 + base_condition = AiModelDetail.id.in_(sub_stmt) + if with_default: + where_condition = or_(base_condition, AiModelDetail.default_model == True) + else: + where_condition = base_condition + stmt = ( + select( + AiModelDetail.id, + AiModelDetail.name, + AiModelDetail.default_model, + AiModelDetail.supplier, + ) + .where(where_condition) + .order_by(AiModelDetail.default_model.desc()) + ) + rows = session.exec(stmt).all() + + return [ + AiModelBrief( + id=row[0], + name=row[1], + default_model=row[2], + supplier=row[3], + ) + for row in rows + ] diff --git a/backend/apps/system/crud/assistant.py b/backend/apps/system/crud/assistant.py index 1fa5eb27c..0c9d60083 100644 --- a/backend/apps/system/crud/assistant.py +++ b/backend/apps/system/crud/assistant.py @@ -22,6 +22,21 @@ from common.core.response_middleware import ResponseMiddleware +def _update_cors_middleware_instance(app: FastAPI, updated_origins: list[str]): + """遍历 middleware 栈,找到 CORSMiddleware 实例并更新其 allow_origins。 + + 仅修改 middleware.kwargs 不会影响已构建的中间件实例, + 需要直接更新实例的 allow_origins 属性。 + """ + stack = getattr(app, 'middleware_stack', None) + while stack is not None and hasattr(stack, 'app'): + if isinstance(stack, CORSMiddleware): + stack.allow_origins = updated_origins + return + stack = stack.app + + + @cache(namespace=CacheNamespace.EMBEDDED_INFO, cacheName=CacheName.ASSISTANT_INFO, keyExpression="assistant_id") async def get_assistant_info(*, session: Session, assistant_id: int) -> AssistantModel | None: db_model = session.get(AssistantModel, assistant_id) @@ -94,10 +109,11 @@ def init_dynamic_cors(app: FastAPI): response_middleware = middleware if cors_middleware and response_middleware: break - + updated_origins = list(set(settings.all_cors_origins + unique_domains)) if cors_middleware: cors_middleware.kwargs['allow_origins'] = updated_origins + _update_cors_middleware_instance(app, updated_origins) if response_middleware: for instance in ResponseMiddleware.instances: instance.update_allow_origins(updated_origins) @@ -187,6 +203,7 @@ def get_db_schema(self, ds_id: int, question: str = '', embedding: bool = True, db_name = ds.db_schema if ds.db_schema is not None and ds.db_schema != "" else ds.dataBase schema_str += f"【DB_ID】 {db_name}\n【Schema】\n" tables = [] + table_name_list = [] i = 0 for table in ds.tables: # 如果传入了 table_list,则只处理在列表中的表 @@ -213,6 +230,7 @@ def get_db_schema(self, ds_id: int, question: str = '', embedding: bool = True, schema_table += '\n]\n' t_obj = {"id": i, "schema_table": schema_table} tables.append(t_obj) + table_name_list.append(table.name) # do table embedding # if embedding and tables and settings.TABLE_EMBEDDING_ENABLED: @@ -222,7 +240,7 @@ def get_db_schema(self, ds_id: int, question: str = '', embedding: bool = True, for s in tables: schema_str += s.get('schema_table') - return schema_str, [] + return schema_str, table_name_list def get_ds(self, ds_id: int, trans: Trans = None): if self.ds_list: @@ -236,11 +254,11 @@ def get_ds(self, ds_id: int, trans: Trans = None): def convert2schema(self, ds_dict: dict, config: dict[any]) -> AssistantOutDsSchema: id_marker: str = '' - attr_list = ['name', 'type', 'host', 'port', 'user', 'dataBase', 'schema', 'mode'] + attr_list = ['name', 'type', 'host', 'port', 'user', 'dataBase', 'schema', 'mode', 'lowVersion'] if config.get('encrypt', False): key = config.get('aes_key', None) iv = config.get('aes_iv', None) - aes_attrs = ['host', 'user', 'password', 'dataBase', 'db_schema', 'schema', 'mode'] + aes_attrs = ['host', 'user', 'password', 'dataBase', 'db_schema', 'schema', 'mode', 'lowVersion'] for attr in aes_attrs: if attr in ds_dict and ds_dict[attr]: try: @@ -248,7 +266,7 @@ def convert2schema(self, ds_dict: dict, config: dict[any]) -> AssistantOutDsSche except Exception as e: raise Exception( f"Failed to encrypt {attr} for datasource {ds_dict.get('name')}, error: {str(e)}") - + id = ds_dict.get('id', None) if not id: for attr in attr_list: @@ -277,7 +295,8 @@ def get_out_ds_conf(ds: AssistantOutDsSchema, timeout: int = 30) -> str: "extraJdbc": ds.extraParams or '', "dbSchema": ds.db_schema or '', "timeout": timeout or 30, - "mode": ds.mode or '' + "mode": ds.mode or '', + "lowVersion": ds.lowVersion or False, } conf["extraJdbc"] = '' return aes_encrypt(json.dumps(conf)) diff --git a/backend/apps/system/crud/assistant_manage.py b/backend/apps/system/crud/assistant_manage.py index adec96793..f9dce3675 100644 --- a/backend/apps/system/crud/assistant_manage.py +++ b/backend/apps/system/crud/assistant_manage.py @@ -12,6 +12,17 @@ from common.core.response_middleware import ResponseMiddleware +def _update_cors_middleware_instance(app: FastAPI, updated_origins: list[str]): + """遍历 middleware 栈,找到 CORSMiddleware 实例并更新其 allow_origins。""" + stack = getattr(app, 'middleware_stack', None) + while stack is not None and hasattr(stack, 'app'): + if isinstance(stack, CORSMiddleware): + stack.allow_origins = updated_origins + return + stack = stack.app + + + def dynamic_upgrade_cors(request: Request, session: Session): list_result = session.exec(select(AssistantModel).order_by(AssistantModel.create_time)).all() seen = set() @@ -37,6 +48,7 @@ def dynamic_upgrade_cors(request: Request, session: Session): updated_origins = list(set(settings.all_cors_origins + unique_domains)) if cors_middleware: cors_middleware.kwargs['allow_origins'] = updated_origins + _update_cors_middleware_instance(app, updated_origins) if response_middleware: for instance in ResponseMiddleware.instances: instance.update_allow_origins(updated_origins) diff --git a/backend/apps/system/crud/parameter_manage.py b/backend/apps/system/crud/parameter_manage.py index f80edd29c..65aaf164a 100644 --- a/backend/apps/system/crud/parameter_manage.py +++ b/backend/apps/system/crud/parameter_manage.py @@ -15,7 +15,7 @@ async def get_groups(session: SessionDep, flag: str) -> list[SysArgModel]: async def save_parameter_args(session: SessionDep, request: Request): allow_file_mapping = { - """ "test_logo": { "types": [".jpg", ".jpeg", ".png", ".svg"], "size": 5 * 1024 * 1024 } """ + """ "test_logo": { "types": [".jpg", ".jpeg", ".png"], "size": 5 * 1024 * 1024 } """ } form_data = await request.form() files = form_data.getlist("files") diff --git a/backend/apps/system/crud/user.py b/backend/apps/system/crud/user.py index 1f5ccb59e..5f9928b1b 100644 --- a/backend/apps/system/crud/user.py +++ b/backend/apps/system/crud/user.py @@ -1,21 +1,23 @@ - from typing import Optional + from sqlmodel import Session, func, select, delete as sqlmodel_delete + from apps.system.models.system_model import UserWsModel, WorkspaceModel from apps.system.schemas.auth import CacheName, CacheNamespace from apps.system.schemas.system_schema import EMAIL_REGEX, PWD_REGEX, BaseUserDTO, UserInfoDTO, UserWs from common.core.deps import SessionDep +from common.core.security import verify_md5pwd from common.core.sqlbot_cache import cache, clear_cache -from common.utils.locale import I18n +from common.utils.locale import I18n, I18nHelper from common.utils.utils import SQLBotLogUtil from ..models.user import UserModel, UserPlatformModel -from common.core.security import verify_md5pwd -import re + def get_db_user(*, session: Session, user_id: int) -> UserModel: db_user = session.get(UserModel, user_id) return db_user + def get_user_by_account(*, session: Session, account: str) -> BaseUserDTO | None: statement = select(UserModel).where(UserModel.account == account) db_user = session.exec(statement).first() @@ -23,19 +25,22 @@ def get_user_by_account(*, session: Session, account: str) -> BaseUserDTO | None return None return BaseUserDTO.model_validate(db_user.model_dump()) + @cache(namespace=CacheNamespace.AUTH_INFO, cacheName=CacheName.USER_INFO, keyExpression="user_id") async def get_user_info(*, session: Session, user_id: int) -> UserInfoDTO | None: - db_user: UserModel = get_db_user(session = session, user_id = user_id) + db_user: UserModel = get_db_user(session=session, user_id=user_id) if not db_user: return None userInfo = UserInfoDTO.model_validate(db_user.model_dump()) userInfo.isAdmin = userInfo.id == 1 and userInfo.account == 'admin' if userInfo.isAdmin: return userInfo - ws_model: UserWsModel = session.exec(select(UserWsModel).where(UserWsModel.uid == userInfo.id, UserWsModel.oid == userInfo.oid)).first() + ws_model: UserWsModel = session.exec( + select(UserWsModel).where(UserWsModel.uid == userInfo.id, UserWsModel.oid == userInfo.oid)).first() userInfo.weight = ws_model.weight if ws_model else -1 return userInfo + def authenticate(*, session: Session, account: str, password: str) -> BaseUserDTO | None: db_user = get_user_by_account(session=session, account=account) if not db_user: @@ -44,7 +49,8 @@ def authenticate(*, session: Session, account: str, password: str) -> BaseUserDT return None return db_user -async def user_ws_options(session: Session, uid: int, trans: Optional[I18n] = None) -> list[UserWs]: + +def user_ws_list(session: Session, uid: int, trans: Optional[I18n | I18nHelper] = None) -> list[UserWs]: if uid == 1: stmt = select(WorkspaceModel.id, WorkspaceModel.name).order_by(WorkspaceModel.name, WorkspaceModel.create_time) else: @@ -57,16 +63,20 @@ async def user_ws_options(session: Session, uid: int, trans: Optional[I18n] = No if not trans: return result.all() list_result = [ - UserWs(id = id, name = trans(name) if name.startswith('i18n') else name) + UserWs(id=id, name=trans(name) if name.startswith('i18n') else name) for id, name in result.all() ] if list_result: list_result.sort(key=lambda x: x.name) return list_result - + +async def user_ws_options(session: Session, uid: int, trans: Optional[I18n | I18nHelper] = None) -> list[UserWs]: + return user_ws_list(session, uid, trans) + + @clear_cache(namespace=CacheNamespace.AUTH_INFO, cacheName=CacheName.USER_INFO, keyExpression="id") async def single_delete(session: SessionDep, id: int): - user_model: UserModel = get_db_user(session = session, user_id = id) + user_model: UserModel = get_db_user(session=session, user_id=id) del_stmt = sqlmodel_delete(UserWsModel).where(UserWsModel.uid == id) session.exec(del_stmt) if user_model and user_model.origin and user_model.origin != 0: @@ -75,20 +85,23 @@ async def single_delete(session: SessionDep, id: int): session.delete(user_model) session.commit() -@clear_cache(namespace=CacheNamespace.AUTH_INFO, cacheName=CacheName.USER_INFO, keyExpression="id") + +@clear_cache(namespace=CacheNamespace.AUTH_INFO, cacheName=CacheName.USER_INFO, keyExpression="id") async def clean_user_cache(id: int): SQLBotLogUtil.info(f"User cache for [{id}] has been cleaned") def check_account_exists(*, session: Session, account: str) -> bool: return session.exec(select(func.count()).select_from(UserModel).where(UserModel.account == account)).one() > 0 + + def check_email_exists(*, session: Session, email: str) -> bool: return session.exec(select(func.count()).select_from(UserModel).where(UserModel.email == email)).one() > 0 - def check_email_format(email: str) -> bool: return bool(EMAIL_REGEX.fullmatch(email)) + def check_pwd_format(pwd: str) -> bool: return bool(PWD_REGEX.fullmatch(pwd)) diff --git a/backend/apps/system/middleware/auth.py b/backend/apps/system/middleware/auth.py index 423f9ea78..feeb246aa 100644 --- a/backend/apps/system/middleware/auth.py +++ b/backend/apps/system/middleware/auth.py @@ -1,15 +1,17 @@ import base64 import json +import re from typing import Optional from fastapi import Request from fastapi.responses import JSONResponse +from starlette.responses import Response import jwt from sqlmodel import Session from starlette.middleware.base import BaseHTTPMiddleware from apps.system.crud.apikey_manage import get_api_key from apps.system.models.system_model import ApiKeyModel, AssistantModel -from common.core.db import engine +from common.core.db import engine from apps.system.crud.assistant import get_assistant_info, get_assistant_user from apps.system.crud.user import get_user_by_account, get_user_info from apps.system.schemas.system_schema import AssistantHeader, UserInfoDTO @@ -17,7 +19,7 @@ from common.core.config import settings from common.core.schemas import TokenPayload from common.utils.locale import I18n -from common.utils.utils import SQLBotLogUtil, get_origin_from_referer +from common.utils.utils import SQLBotLogUtil, get_origin_from_referer, origin_match_domain from common.utils.whitelist import whiteUtils from fastapi.security.utils import get_authorization_scheme_param from common.core.deps import get_i18n @@ -31,6 +33,26 @@ def __init__(self, app): async def dispatch(self, request, call_next): if self.is_options(request) or whiteUtils.is_whitelisted(request.url.path): + # 动态处理 /system/assistant/info/{id} 的 CORS 预检 + if request.method == "OPTIONS": + origin = request.headers.get("origin", "") + if origin: + match = re.search(r'/system/assistant/info/(\d+)', request.url.path) + if match: + assistant_id = int(match.group(1)) + with Session(engine) as session: + db_model = session.get(AssistantModel, assistant_id) + if db_model and origin_match_domain(origin, db_model.domain): + return Response( + status_code=200, + headers={ + "Access-Control-Allow-Origin": origin, + "Access-Control-Allow-Methods": "GET, OPTIONS", + "Access-Control-Allow-Headers": "*", + "Access-Control-Allow-Credentials": "true", + "Access-Control-Max-Age": "600", + }, + ) return await call_next(request) assistantTokenKey = settings.ASSISTANT_TOKEN_KEY assistantToken = request.headers.get(assistantTokenKey) diff --git a/backend/apps/system/models/system_model.py b/backend/apps/system/models/system_model.py index 73fa34435..6b3602329 100644 --- a/backend/apps/system/models/system_model.py +++ b/backend/apps/system/models/system_model.py @@ -1,6 +1,8 @@ - from typing import Optional + +from pydantic import field_serializer from sqlmodel import BigInteger, Field, Text, SQLModel + from common.core.models import SnowflakeBase from common.core.schemas import BaseCreatorDTO @@ -9,76 +11,98 @@ class AiModelBase: supplier: int = Field(nullable=False) name: str = Field(max_length=255, nullable=False) model_type: int = Field(nullable=False) - base_model: str = Field(max_length = 255, nullable=False) + base_model: str = Field(max_length=255, nullable=False) default_model: bool = Field(default=False, nullable=False) + class AiModelDetail(SnowflakeBase, AiModelBase, table=True): - __tablename__ = "ai_model" - api_key: str | None = Field(nullable=True) - api_domain: str = Field(nullable=False) - protocol: int = Field(nullable=False, default = 1) - config: str = Field(sa_type = Text()) - status: int = Field(nullable=False, default = 1) - create_time: int = Field(default=0, sa_type=BigInteger()) - + __tablename__ = "ai_model" + api_key: str | None = Field(default=None, nullable=True, sa_type=Text()) + api_domain: str = Field(nullable=False, sa_type=Text()) + protocol: int = Field(nullable=False, default=1) + config: str = Field(sa_type=Text()) + status: int = Field(nullable=False, default=1) + create_time: int = Field(default=0, sa_type=BigInteger()) + + +class AiModelWorkspaceMapping(SnowflakeBase, table=True): + __tablename__ = "ai_model_workspace_mapping" + ai_model_id: int = Field(default=None, nullable=True, sa_type=BigInteger()) + workspace_id: int = Field(default=None, nullable=True, sa_type=BigInteger()) +class AiModelBrief(SQLModel): + id: int + name: str + default_model: bool + supplier: int + + @field_serializer("id") + def id_to_str(self, v: int) -> str: + return str(v) + class WorkspaceBase(SQLModel): name: str = Field(max_length=255, nullable=False) + class WorkspaceEditor(WorkspaceBase, BaseCreatorDTO): pass - + + class WorkspaceModel(SnowflakeBase, WorkspaceBase, table=True): __tablename__ = "sys_workspace" create_time: int = Field(default=0, sa_type=BigInteger()) - + + class UserWsBaseModel(SQLModel): uid: int = Field(nullable=False, sa_type=BigInteger()) oid: int = Field(nullable=False, sa_type=BigInteger()) - weight: int = Field(default=0, nullable=False) - + weight: int = Field(default=0, nullable=False) + + class UserWsModel(SnowflakeBase, UserWsBaseModel, table=True): __tablename__ = "sys_user_ws" - + class AssistantBaseModel(SQLModel): name: str = Field(max_length=255, nullable=False) type: int = Field(nullable=False, default=0) domain: str = Field(max_length=255, nullable=False) - description: Optional[str] = Field(sa_type = Text(), nullable=True) - configuration: Optional[str] = Field(sa_type = Text(), nullable=True) + description: Optional[str] = Field(sa_type=Text(), nullable=True) + configuration: Optional[str] = Field(sa_type=Text(), nullable=True) create_time: int = Field(default=0, sa_type=BigInteger()) - app_id: Optional[str] = Field(default=None, max_length=255, nullable=True) + app_id: Optional[str] = Field(default=None, max_length=255, nullable=True) app_secret: Optional[str] = Field(default=None, max_length=255, nullable=True) oid: Optional[int] = Field(nullable=True, sa_type=BigInteger(), default=1) enable_custom_model: Optional[bool] = Field(default=False, nullable=True) custom_model: Optional[str] = Field(default=None, max_length=255, nullable=True) + class AssistantModel(SnowflakeBase, AssistantBaseModel, table=True): __tablename__ = "sys_assistant" - + class AuthenticationBaseModel(SQLModel): name: str = Field(max_length=255, nullable=False) type: int = Field(nullable=False, default=0) - config: Optional[str] = Field(sa_type = Text(), nullable=True) - - + config: Optional[str] = Field(sa_type=Text(), nullable=True) + + class AuthenticationModel(SnowflakeBase, AuthenticationBaseModel, table=True): __tablename__ = "sys_authentication" create_time: Optional[int] = Field(default=0, sa_type=BigInteger()) enable: bool = Field(default=False, nullable=False) valid: bool = Field(default=False, nullable=False) - + class ApiKeyBaseModel(SQLModel): access_key: str = Field(max_length=255, nullable=False) secret_key: str = Field(max_length=255, nullable=False) create_time: int = Field(default=0, sa_type=BigInteger()) - uid: int = Field(default=0,nullable=False, sa_type=BigInteger()) + uid: int = Field(default=0, nullable=False, sa_type=BigInteger()) status: bool = Field(default=True, nullable=False) - + + class ApiKeyModel(SnowflakeBase, ApiKeyBaseModel, table=True): - __tablename__ = "sys_apikey" \ No newline at end of file + __tablename__ = "sys_apikey" diff --git a/backend/apps/system/schemas/ai_model_schema.py b/backend/apps/system/schemas/ai_model_schema.py index 019aa358b..7f523eb93 100644 --- a/backend/apps/system/schemas/ai_model_schema.py +++ b/backend/apps/system/schemas/ai_model_schema.py @@ -14,7 +14,7 @@ class AiModelItem(BaseModel): default_model: bool = Field(default=False, description=f"{PLACEHOLDER_PREFIX}default_model") class AiModelGridItem(AiModelItem, BaseCreatorDTO): - pass + ws_mapping_count: int = Field(default=0, description="workspace mapping count") class AiModelConfigItem(BaseModel): key: str = Field(description=f"{PLACEHOLDER_PREFIX}arg_name") diff --git a/backend/apps/system/schemas/system_schema.py b/backend/apps/system/schemas/system_schema.py index 75eebc39c..db2f255b4 100644 --- a/backend/apps/system/schemas/system_schema.py +++ b/backend/apps/system/schemas/system_schema.py @@ -111,8 +111,8 @@ class AssistantBase(BaseModel): configuration: Optional[str] = Field(default=None, description=f"{PLACEHOLDER_PREFIX}assistant_configuration") description: Optional[str] = Field(default=None, description=f"{PLACEHOLDER_PREFIX}assistant_description") oid: Optional[int] = Field(default=1, description=f"{PLACEHOLDER_PREFIX}oid") - enable_custom_model: Optional[bool] = Field(default=False, description=f"{PLACEHOLDER_PREFIX}oid") - custom_model: Optional[str] = Field(description=f"{PLACEHOLDER_PREFIX}oid") + enable_custom_model: Optional[bool] = Field(default=False, description=f"{PLACEHOLDER_PREFIX}enable_custom_model") + custom_model: Optional[str] = Field(description=f"{PLACEHOLDER_PREFIX}custom_model") class AssistantDTO(AssistantBase, BaseCreatorDTO): @@ -197,6 +197,7 @@ class AssistantOutDsSchema(AssistantOutDsBase): db_schema: Optional[str] = None extraParams: Optional[str] = None mode: Optional[str] = None + lowVersion: Optional[bool] = False tables: Optional[list[AssistantTableSchema]] = None diff --git a/backend/apps/terminology/api/terminology.py b/backend/apps/terminology/api/terminology.py index b74cb19bd..979cdf41f 100644 --- a/backend/apps/terminology/api/terminology.py +++ b/backend/apps/terminology/api/terminology.py @@ -82,6 +82,7 @@ def inner(): "description": obj.description, "all_data_sources": 'N' if obj.specific_ds else 'Y', "datasource": ', '.join(obj.datasource_names) if obj.datasource_names and obj.specific_ds else '', + "advanced_application_name": obj.advanced_application_name or '', } data_list.append(_data) @@ -91,6 +92,7 @@ def inner(): fields.append(AxisObj(name=trans('i18n_terminology.term_description'), value='description')) fields.append(AxisObj(name=trans('i18n_terminology.effective_data_sources'), value='datasource')) fields.append(AxisObj(name=trans('i18n_terminology.all_data_sources'), value='all_data_sources')) + fields.append(AxisObj(name=trans('i18n_data_training.advanced_application'), value='advanced_application_name')) md_data, _fields_list = DataFormat.convert_object_array_for_pandas(fields, data_list) @@ -119,6 +121,7 @@ def inner(): "description": trans('i18n_terminology.term_description_template_example_1'), "all_data_sources": 'N', "datasource": trans('i18n_terminology.effective_data_sources_template_example_1'), + "advanced_application_name": '', } data_list.append(_data1) _data2 = { @@ -127,6 +130,7 @@ def inner(): "description": trans('i18n_terminology.term_description_template_example_2'), "all_data_sources": 'Y', "datasource": '', + "advanced_application_name": '', } data_list.append(_data2) @@ -136,6 +140,7 @@ def inner(): fields.append(AxisObj(name=trans('i18n_terminology.term_description_template'), value='description')) fields.append(AxisObj(name=trans('i18n_terminology.effective_data_sources_template'), value='datasource')) fields.append(AxisObj(name=trans('i18n_terminology.all_data_sources_template'), value='all_data_sources')) + fields.append(AxisObj(name=trans('i18n_data_training.advanced_application'), value='advanced_application_name')) md_data, _fields_list = DataFormat.convert_object_array_for_pandas(fields, data_list) @@ -180,7 +185,7 @@ async def upload_excel(trans: Trans, current_user: CurrentUser, file: UploadFile oid = current_user.oid - use_cols = [0, 1, 2, 3, 4] + use_cols = [0, 1, 2, 3, 4, 5] def inner(): @@ -217,9 +222,11 @@ def inner(): 3].strip() else [] all_datasource = True if pd.notna(row[4]) and row[4].lower().strip() in ['y', 'yes', 'true'] else False specific_ds = False if all_datasource else True + advanced_application_name = row[5].strip() if pd.notna(row[5]) and row[5].strip() else None import_data.append(TerminologyInfo(word=word, description=description, other_words=other_words, - datasource_names=datasource_names, specific_ds=specific_ds)) + datasource_names=datasource_names, specific_ds=specific_ds, + advanced_application_name=advanced_application_name)) res = batch_create_terminology(session, import_data, oid, trans) @@ -237,6 +244,7 @@ def inner(): "all_data_sources": 'N' if obj['data'].specific_ds else 'Y', "datasource": ', '.join(obj['data'].datasource_names) if obj['data'].datasource_names and obj[ 'data'].specific_ds else '', + "advanced_application_name": obj['data'].advanced_application_name or '', "errors": obj['errors'] } data_list.append(_data) @@ -247,6 +255,7 @@ def inner(): fields.append(AxisObj(name=trans('i18n_terminology.term_description'), value='description')) fields.append(AxisObj(name=trans('i18n_terminology.effective_data_sources'), value='datasource')) fields.append(AxisObj(name=trans('i18n_terminology.all_data_sources'), value='all_data_sources')) + fields.append(AxisObj(name=trans('i18n_data_training.advanced_application'), value='advanced_application_name')) fields.append(AxisObj(name=trans('i18n_data_training.error_info'), value='errors')) md_data, _fields_list = DataFormat.convert_object_array_for_pandas(fields, data_list) diff --git a/backend/apps/terminology/curd/terminology.py b/backend/apps/terminology/curd/terminology.py index 296fd1eba..7ca6128ab 100644 --- a/backend/apps/terminology/curd/terminology.py +++ b/backend/apps/terminology/curd/terminology.py @@ -10,8 +10,9 @@ from apps.ai_model.embedding import EmbeddingModelCache from apps.datasource.models.datasource import CoreDatasource +from apps.system.models.system_model import AssistantModel from apps.template.generate_chart.generator import get_base_terminology_template -from apps.terminology.models.terminology_model import Terminology, TerminologyInfo +from apps.terminology.models.terminology_model import Terminology, TerminologyInfo, TerminologyInfoResult from common.core.config import settings from common.core.deps import SessionDep, Trans from common.utils.embedding_threads import run_save_terminology_embeddings @@ -141,7 +142,9 @@ def build_terminology_query(session: SessionDep, oid: int, name: Optional[str] = Terminology.datasource_ids, children_subquery.c.other_words, func.jsonb_agg(CoreDatasource.name).filter(CoreDatasource.id.isnot(None)).label('datasource_names'), - Terminology.enabled + Terminology.enabled, + Terminology.advanced_application, + AssistantModel.name.label('advanced_application_name'), ) .outerjoin( children_subquery, @@ -155,6 +158,8 @@ def build_terminology_query(session: SessionDep, oid: int, name: Optional[str] = CoreDatasource, CoreDatasource.id == datasource_names_subquery.c.ds_id ) + .outerjoin(AssistantModel, + and_(Terminology.advanced_application == AssistantModel.id, AssistantModel.type == 1)) .where(and_(Terminology.id.in_(paginated_parent_ids), Terminology.oid == oid)) .group_by( Terminology.id, @@ -163,6 +168,8 @@ def build_terminology_query(session: SessionDep, oid: int, name: Optional[str] = Terminology.description, Terminology.specific_ds, Terminology.datasource_ids, + Terminology.advanced_application, + AssistantModel.name, children_subquery.c.other_words, Terminology.enabled ) @@ -172,7 +179,7 @@ def build_terminology_query(session: SessionDep, oid: int, name: Optional[str] = return stmt, total_count, total_pages, current_page, page_size -def execute_terminology_query(session: SessionDep, stmt) -> List[TerminologyInfo]: +def execute_terminology_query(session: SessionDep, stmt) -> List[TerminologyInfoResult]: """ 执行查询并返回术语信息列表 """ @@ -180,7 +187,7 @@ def execute_terminology_query(session: SessionDep, stmt) -> List[TerminologyInfo result = session.execute(stmt) for row in result: - _list.append(TerminologyInfo( + _list.append(TerminologyInfoResult( id=row.id, word=row.word, create_time=row.create_time, @@ -190,6 +197,8 @@ def execute_terminology_query(session: SessionDep, stmt) -> List[TerminologyInfo datasource_ids=row.datasource_ids if row.datasource_ids is not None else [], datasource_names=row.datasource_names if row.datasource_names is not None else [], enabled=row.enabled if row.enabled is not None else False, + advanced_application=str(row.advanced_application) if row.advanced_application else None, + advanced_application_name=row.advanced_application_name, )) return _list @@ -240,8 +249,8 @@ def create_terminology(session: SessionDep, info: TerminologyInfo, oid: int, tra datasource_ids = info.datasource_ids if info.datasource_ids is not None else [] if specific_ds: - if not datasource_ids: - raise Exception(trans("i18n_terminology.datasource_cannot_be_none")) + if not datasource_ids and info.advanced_application is None: + raise Exception(trans("i18n_data_training.datasource_assistant_cannot_be_none")) parent = Terminology( word=info.word.strip(), @@ -250,7 +259,8 @@ def create_terminology(session: SessionDep, info: TerminologyInfo, oid: int, tra oid=oid, specific_ds=specific_ds, enabled=info.enabled, - datasource_ids=datasource_ids + datasource_ids=datasource_ids, + advanced_application=info.advanced_application, ) words = [info.word.strip()] @@ -267,30 +277,63 @@ def create_terminology(session: SessionDep, info: TerminologyInfo, oid: int, tra # 基础查询条件(word 和 oid 必须满足) base_query = and_( Terminology.word.in_(words), - Terminology.oid == oid + Terminology.oid == oid, ) # 构建查询 query = session.query(Terminology).filter(base_query) - if specific_ds: - # 仅当 specific_ds=False 时,检查数据源条件 - query = query.where( - or_( - or_(Terminology.specific_ds == False, Terminology.specific_ds.is_(None)), - and_( - Terminology.specific_ds == True, - Terminology.datasource_ids.isnot(None), - text(""" - EXISTS ( - SELECT 1 FROM jsonb_array_elements(datasource_ids) AS elem - WHERE elem::text::int = ANY(:datasource_ids) - ) - """) + # 作用域重复检查 + scope_conditions = [] + + if not specific_ds: + # 全部数据源:与有数据源的记录冲突,也与同为全部数据源的记录冲突 + scope_conditions.append( + and_( + Terminology.specific_ds == True, + Terminology.datasource_ids.isnot(None), + func.jsonb_array_length(Terminology.datasource_ids) > 0, + ) + ) + scope_conditions.append( + Terminology.specific_ds == False, + ) + elif specific_ds and datasource_ids: + # 指定数据源:与全部数据源冲突,也与有数据源重叠的记录冲突 + scope_conditions.append( + Terminology.specific_ds == False, + ) + ds_overlap_conditions = [ + Terminology.datasource_ids.contains([ds_id]) + for ds_id in datasource_ids + ] + scope_conditions.append( + and_( + Terminology.specific_ds == True, + Terminology.datasource_ids.isnot(None), + or_(*ds_overlap_conditions) + ) + ) + else: + # 不选数据源:仅与同为不选数据源的记录冲突 + scope_conditions.append( + and_( + Terminology.specific_ds == True, + or_( + Terminology.datasource_ids.is_(None), + func.jsonb_array_length(Terminology.datasource_ids) == 0, ) ) ) - query = query.params(datasource_ids=datasource_ids) + + # 高级应用重复检查:advanced_application 相同时同名即重复 + if info.advanced_application is not None: + scope_conditions.append( + Terminology.advanced_application == info.advanced_application + ) + + if scope_conditions: + query = query.where(or_(*scope_conditions)) # 转换为 EXISTS 查询并获取结果 exists = session.query(query.exists()).scalar() @@ -396,6 +439,14 @@ def batch_create_terminology(session: SessionDep, info_list: List[TerminologyInf for ds in datasource_result: datasource_name_to_id[ds.name.strip()] = ds.id + # 预加载高级应用名称到ID的映射 + assistant_name_to_id = {} + assistant_stmt = select(AssistantModel.id, AssistantModel.name).where( + and_(AssistantModel.oid == oid, AssistantModel.type == 1)) + assistant_result = session.execute(assistant_stmt).all() + for a in assistant_result: + assistant_name_to_id[a.name.strip()] = a.id + # 验证和转换数据源名称 valid_records = [] for info in deduplicated_list: @@ -411,6 +462,7 @@ def batch_create_terminology(session: SessionDep, info_list: List[TerminologyInf # 根据specific_ds决定是否验证数据源 specific_ds = info.specific_ds if info.specific_ds is not None else False datasource_ids = [] + advanced_application = info.advanced_application if specific_ds: # specific_ds为True时需要验证数据源 @@ -424,12 +476,21 @@ def batch_create_terminology(session: SessionDep, info_list: List[TerminologyInf else: error_messages.append(trans("i18n_terminology.datasource_not_found").format(ds_name)) - # 检查specific_ds为True时必须有数据源 - if not datasource_ids: - error_messages.append(trans("i18n_terminology.datasource_cannot_be_none")) + # 解析高级应用名称到ID + if advanced_application is None and info.advanced_application_name: + if info.advanced_application_name.strip() in assistant_name_to_id: + advanced_application = assistant_name_to_id[info.advanced_application_name.strip()] + else: + error_messages.append(trans("i18n_data_training.advanced_application_not_found").format( + info.advanced_application_name)) + + # 检查specific_ds为True时datasource_ids和advanced_application不能同时为空 + if specific_ds and not datasource_ids and advanced_application is None: + error_messages.append(trans("i18n_data_training.datasource_assistant_cannot_be_none")) else: # specific_ds为False时忽略数据源名称 datasource_ids = [] + advanced_application = None # 检查主词和其他词是否重复(过滤空字符串) words = [info.word.strip().lower()] @@ -456,11 +517,13 @@ def batch_create_terminology(session: SessionDep, info_list: List[TerminologyInf processed_info = TerminologyInfo( word=info.word.strip(), description=info.description.strip(), - other_words=[w for w in info.other_words if w and w.strip()], # 过滤空字符串 + other_words=[w for w in info.other_words if w and w.strip()], datasource_ids=datasource_ids, datasource_names=info.datasource_names, specific_ds=specific_ds, - enabled=info.enabled if info.enabled is not None else True + enabled=info.enabled if info.enabled is not None else True, + advanced_application=advanced_application, + advanced_application_name=info.advanced_application_name, ) valid_records.append(processed_info) @@ -512,8 +575,8 @@ def update_terminology(session: SessionDep, info: TerminologyInfo, oid: int, tra datasource_ids = info.datasource_ids if info.datasource_ids is not None else [] if specific_ds: - if not datasource_ids: - raise Exception(trans("i18n_terminology.datasource_cannot_be_none")) + if not datasource_ids and info.advanced_application is None: + raise Exception(trans("i18n_data_training.datasource_assistant_cannot_be_none")) words = [info.word.strip()] for child in info.other_words: @@ -536,24 +599,57 @@ def update_terminology(session: SessionDep, info: TerminologyInfo, oid: int, tra # 构建查询 query = session.query(Terminology).filter(base_query) - if specific_ds: - # 仅当 specific_ds=False 时,检查数据源条件 - query = query.where( - or_( - or_(Terminology.specific_ds == False, Terminology.specific_ds.is_(None)), - and_( - Terminology.specific_ds == True, - Terminology.datasource_ids.isnot(None), - text(""" - EXISTS ( - SELECT 1 FROM jsonb_array_elements(datasource_ids) AS elem - WHERE elem::text::int = ANY(:datasource_ids) - ) - """) # 检查是否包含任意目标值 + # 作用域重复检查 + scope_conditions = [] + + if not specific_ds: + # 全部数据源:与有数据源的记录冲突,也与同为全部数据源的记录冲突 + scope_conditions.append( + and_( + Terminology.specific_ds == True, + Terminology.datasource_ids.isnot(None), + func.jsonb_array_length(Terminology.datasource_ids) > 0, + ) + ) + scope_conditions.append( + Terminology.specific_ds == False, + ) + elif specific_ds and datasource_ids: + # 指定数据源:与全部数据源冲突,也与有数据源重叠的记录冲突 + scope_conditions.append( + Terminology.specific_ds == False, + ) + ds_overlap_conditions = [ + Terminology.datasource_ids.contains([ds_id]) + for ds_id in datasource_ids + ] + scope_conditions.append( + and_( + Terminology.specific_ds == True, + Terminology.datasource_ids.isnot(None), + or_(*ds_overlap_conditions) + ) + ) + else: + # 不选数据源:仅与同为不选数据源的记录冲突 + scope_conditions.append( + and_( + Terminology.specific_ds == True, + or_( + Terminology.datasource_ids.is_(None), + func.jsonb_array_length(Terminology.datasource_ids) == 0, ) ) ) - query = query.params(datasource_ids=datasource_ids) + + # 高级应用重复检查:advanced_application 相同时同名即重复 + if info.advanced_application is not None: + scope_conditions.append( + Terminology.advanced_application == info.advanced_application + ) + + if scope_conditions: + query = query.where(or_(*scope_conditions)) # 转换为 EXISTS 查询并获取结果 exists = session.query(query.exists()).scalar() @@ -567,6 +663,7 @@ def update_terminology(session: SessionDep, info: TerminologyInfo, oid: int, tra specific_ds=specific_ds, datasource_ids=datasource_ids, enabled=info.enabled, + advanced_application=info.advanced_application, ) session.execute(stmt) session.commit() @@ -590,7 +687,8 @@ def update_terminology(session: SessionDep, info: TerminologyInfo, oid: int, tra oid=oid, enabled=info.enabled, specific_ds=specific_ds, - datasource_ids=datasource_ids + datasource_ids=datasource_ids, + advanced_application=info.advanced_application, ) ) @@ -711,8 +809,22 @@ def save_embeddings(session_maker, ids: List[int]): LIMIT {settings.EMBEDDING_TERMINOLOGY_TOP_COUNT} """ +embedding_sql_with_advanced_application = f""" +SELECT id, pid, word, similarity +FROM +(SELECT id, pid, word, oid, specific_ds, advanced_application, enabled, +( 1 - (embedding <=> :embedding_array) ) AS similarity +FROM terminology AS child +) TEMP +WHERE similarity > {settings.EMBEDDING_TERMINOLOGY_SIMILARITY} AND oid = :oid AND enabled = true +AND advanced_application = :advanced_application_id +ORDER BY similarity DESC +LIMIT {settings.EMBEDDING_TERMINOLOGY_TOP_COUNT} +""" -def select_terminology_by_word(session: SessionDep, word: str, oid: int, datasource: int = None): + +def select_terminology_by_word(session: SessionDep, word: str, oid: int, datasource: int = None, + advanced_application_id: Optional[int] = None): if word.strip() == "": return [] @@ -729,7 +841,9 @@ def select_terminology_by_word(session: SessionDep, word: str, oid: int, datasou ) ) - if datasource is not None: + if advanced_application_id is not None: + stmt = stmt.where(Terminology.advanced_application == advanced_application_id) + elif datasource is not None: stmt = stmt.where( or_( or_(Terminology.specific_ds == False, Terminology.specific_ds.is_(None)), @@ -760,7 +874,11 @@ def select_terminology_by_word(session: SessionDep, word: str, oid: int, datasou embedding = model.embed_query(word) - if datasource is not None: + if advanced_application_id is not None: + results = session.execute(text(embedding_sql_with_advanced_application), + {'embedding_array': str(embedding), 'oid': oid, + 'advanced_application_id': advanced_application_id}).fetchall() + elif datasource is not None: results = session.execute(text(embedding_sql_with_datasource), {'embedding_array': str(embedding), 'oid': oid, 'datasource': datasource}).fetchall() @@ -846,10 +964,11 @@ def to_xml_string(_dict: list[dict] | dict, root: str = 'terminologies') -> str: def get_terminology_template(session: SessionDep, question: str, oid: Optional[int] = 1, - datasource: Optional[int] = None) -> tuple[str, list[dict]]: + datasource: Optional[int] = None, + advanced_application_id: Optional[int] = None) -> tuple[str, list[dict]]: if not oid: oid = 1 - _results = select_terminology_by_word(session, question, oid, datasource) + _results = select_terminology_by_word(session, question, oid, datasource, advanced_application_id) if _results and len(_results) > 0: terminology = to_xml_string(_results) template = get_base_terminology_template().format(terminologies=terminology) diff --git a/backend/apps/terminology/models/terminology_model.py b/backend/apps/terminology/models/terminology_model.py index 850aadabe..a0887eb77 100644 --- a/backend/apps/terminology/models/terminology_model.py +++ b/backend/apps/terminology/models/terminology_model.py @@ -20,6 +20,7 @@ class Terminology(SQLModel, table=True): specific_ds: Optional[bool] = Field(sa_column=Column(Boolean, default=False)) datasource_ids: Optional[list[int]] = Field(sa_column=Column(JSONB), default=[]) enabled: Optional[bool] = Field(sa_column=Column(Boolean, default=True)) + advanced_application: Optional[int] = Field(sa_column=Column(BigInteger, nullable=True)) class TerminologyInfo(BaseModel): @@ -32,3 +33,18 @@ class TerminologyInfo(BaseModel): datasource_ids: Optional[list[int]] = [] datasource_names: Optional[list[str]] = [] enabled: Optional[bool] = True + advanced_application: Optional[int] = None + advanced_application_name: Optional[str] = None + +class TerminologyInfoResult(BaseModel): + id: Optional[int] = None + create_time: Optional[datetime] = None + word: Optional[str] = None + description: Optional[str] = None + other_words: Optional[List[str]] = [] + specific_ds: Optional[bool] = False + datasource_ids: Optional[list[int]] = [] + datasource_names: Optional[list[str]] = [] + enabled: Optional[bool] = True + advanced_application: Optional[str] = None + advanced_application_name: Optional[str] = None diff --git a/backend/common/audit/schemas/logger_decorator.py b/backend/common/audit/schemas/logger_decorator.py index bd7a777cb..957ff5659 100644 --- a/backend/common/audit/schemas/logger_decorator.py +++ b/backend/common/audit/schemas/logger_decorator.py @@ -127,12 +127,34 @@ def get_client_info(request: Request) -> Dict[str, Optional[str]]: user_agent = None if request: - # Obtain IP address - if request.client: + # Prefer real client IP headers, then fallback to socket peer IP. + header_candidates = [ + "x-forwarded-for", + "x-real-ip", + "cf-connecting-ip", + "x-client-ip", + "x-forwarded", + "forwarded", + ] + for key in header_candidates: + value = request.headers.get(key) + if value: + if key == "x-forwarded-for": + ip_address = value.split(",")[0].strip() + elif key == "forwarded": + # RFC 7239 format, e.g. for=203.0.113.43;proto=https;by=203.0.113.1 + parts = [part.strip() for part in value.split(";")] + for part in parts: + if part.lower().startswith("for="): + ip_address = part.split("=", 1)[1].strip().strip('"') + break + else: + ip_address = value.strip() + if ip_address: + break + + if not ip_address and request.client: ip_address = request.client.host - # Attempt to obtain the real IP from X-Forwarded-For - if "x-forwarded-for" in request.headers: - ip_address = request.headers["x-forwarded-for"].split(",")[0].strip() # Get User Agent user_agent = request.headers.get("user-agent") diff --git a/backend/common/core/config.py b/backend/common/core/config.py index 1b3cc24ef..4b9baeaec 100644 --- a/backend/common/core/config.py +++ b/backend/common/core/config.py @@ -75,6 +75,8 @@ def API_V1_STR(self) -> str: SCRIPT_DIR: str = f"{BASE_DIR}/scripts" UPLOAD_DIR: str = "/opt/sqlbot/data/file" SQLBOT_KEY_EXPIRED: int = 100 # License key expiration timestamp, 0 means no expiration + + SQLBOT_DOC_ENABLED: bool = True @computed_field # type: ignore[prop-decorator] @property @@ -111,6 +113,10 @@ def SQLALCHEMY_DATABASE_URI(self) -> PostgresDsn | str: GENERATE_SQL_QUERY_LIMIT_ENABLED: bool = True GENERATE_SQL_QUERY_HISTORY_ROUND_COUNT: int = 3 + # 安全配置:是否允许元数据查询(SHOW/DESCRIBE/DESC/EXPLAIN) + # 默认关闭,防止通过元数据查询泄露数据库结构 + SQLBOT_ALLOW_METADATA_QUERIES: bool = False + PARSE_REASONING_BLOCK_ENABLED: bool = True DEFAULT_REASONING_CONTENT_START: str = '' DEFAULT_REASONING_CONTENT_END: str = '' diff --git a/backend/common/utils/data_format.py b/backend/common/utils/data_format.py index bfa9e88b5..56a83f866 100644 --- a/backend/common/utils/data_format.py +++ b/backend/common/utils/data_format.py @@ -1,3 +1,5 @@ +from decimal import Decimal + import pandas as pd from apps.chat.models.chat_model import AxisObj @@ -53,7 +55,7 @@ def format_float_without_scientific(value): """格式化浮点数,避免科学记数法""" if value == 0: return "0" - formatted = f"{value:.15f}" + formatted = str(Decimal(str(value))) if '.' in formatted: formatted = formatted.rstrip('0').rstrip('.') return formatted diff --git a/backend/main.py b/backend/main.py index 3a72bb06e..a8bab9743 100644 --- a/backend/main.py +++ b/backend/main.py @@ -10,8 +10,10 @@ from fastapi.routing import APIRoute from fastapi.staticfiles import StaticFiles from fastapi_mcp import FastApiMCP +from starlette.datastructures import MutableHeaders from starlette.exceptions import HTTPException as StarletteHTTPException from starlette.middleware.cors import CORSMiddleware +from starlette.middleware.base import BaseHTTPMiddleware from alembic import command from apps.api import api_router @@ -70,13 +72,27 @@ def custom_generate_unique_id(route: APIRoute) -> str: app = FastAPI( title=settings.PROJECT_NAME, - openapi_url=f"{settings.CONTEXT_PATH}/openapi.json", + openapi_url=f"{settings.CONTEXT_PATH}/openapi.json" if settings.SQLBOT_DOC_ENABLED else None, generate_unique_id_function=custom_generate_unique_id, lifespan=lifespan, docs_url=None, redoc_url=None ) + +class McpClientIpForwardMiddleware(BaseHTTPMiddleware): + async def dispatch(self, request: Request, call_next): + client_host = request.client.host if request.client else None + if client_host: + headers = MutableHeaders(scope=request.scope) + if not headers.get("x-real-ip"): + headers["x-real-ip"] = client_host + if not headers.get("x-forwarded-for"): + headers["x-forwarded-for"] = client_host + if not headers.get("x-client-ip"): + headers["x-client-ip"] = client_host + return await call_next(request) + # cache docs for different text _openapi_cache: Dict[str, Dict[str, Any]] = {} @@ -152,27 +168,29 @@ def generate_openapi_for_lang(lang: str) -> Dict[str, Any]: # custom /openapi.json and /docs -@app.get(f"{settings.CONTEXT_PATH}/openapi.json", include_in_schema=False) -async def custom_openapi(request: Request): - lang = get_language_from_request(request) - schema = generate_openapi_for_lang(lang) - return JSONResponse(schema) - - -@app.get(f"{settings.CONTEXT_PATH}/docs", include_in_schema=False) -async def custom_swagger_ui(request: Request): - lang = get_language_from_request(request) - from fastapi.openapi.docs import get_swagger_ui_html - return get_swagger_ui_html( - openapi_url=f"./openapi.json?lang={lang}", - title="SQLBot API Docs", - swagger_favicon_url="https://fastapi.tiangolo.com/img/favicon.png", - swagger_js_url="./swagger-ui-bundle.js", - swagger_css_url="./swagger-ui.css", - ) +if settings.SQLBOT_DOC_ENABLED: + @app.get(f"{settings.CONTEXT_PATH}/openapi.json", include_in_schema=False) + async def custom_openapi(request: Request): + lang = get_language_from_request(request) + schema = generate_openapi_for_lang(lang) + return JSONResponse(schema) + + + @app.get(f"{settings.CONTEXT_PATH}/docs", include_in_schema=False) + async def custom_swagger_ui(request: Request): + lang = get_language_from_request(request) + from fastapi.openapi.docs import get_swagger_ui_html + return get_swagger_ui_html( + openapi_url=f"./openapi.json?lang={lang}", + title="SQLBot API Docs", + swagger_favicon_url="https://fastapi.tiangolo.com/img/favicon.png", + swagger_js_url="./swagger-ui-bundle.js", + swagger_css_url="./swagger-ui.css", + ) mcp_app = FastAPI() +mcp_app.add_middleware(McpClientIpForwardMiddleware) # mcp server, images path images_path = settings.MCP_IMAGE_PATH os.makedirs(images_path, exist_ok=True) @@ -184,7 +202,8 @@ async def custom_swagger_ui(request: Request): description="SQLBot MCP Server", describe_all_responses=True, describe_full_response_schema=True, - include_operations=["mcp_datasource_list", "get_model_list", "mcp_question", "mcp_start", "mcp_assistant", "mcp_ws_list"] + include_operations=["mcp_datasource_list", "get_model_list", "mcp_question", "mcp_start", "mcp_assistant", "mcp_ws_list", "access_token"], + headers=["Authorization", "X-Forwarded-For", "X-Real-IP", "CF-Connecting-IP", "X-Client-IP"] ) mcp.mount(mcp_app) diff --git a/backend/pyproject.toml b/backend/pyproject.toml index e0f345f19..1c38a592d 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "sqlbot" -version = "1.8.0" +version = "1.10.0" description = "" requires-python = "==3.11.*" dependencies = [ @@ -39,7 +39,7 @@ dependencies = [ "pyyaml (>=6.0.2,<7.0.0)", "fastapi-mcp (>=0.3.4,<0.4.0)", "tabulate>=0.9.0", - "sqlbot-xpack>=0.0.5.13,<0.0.6.0", + "sqlbot-xpack>=0.0.5.32,<0.0.6.0", "fastapi-cache2>=0.2.2", "sqlparse>=0.5.3", "redis>=6.2.0", @@ -55,7 +55,8 @@ dependencies = [ "sqlglot>=28.6.0", "numpy==2.3.5", "pyhive[hive_pure_sasl]>=0.7.0", - "thrift-sasl" + "thrift-sasl", + "dbutils>=3.1.2", ] [project.optional-dependencies] diff --git a/backend/templates/sql_examples/AWS_Redshift.yaml b/backend/templates/sql_examples/AWS_Redshift.yaml index 40036b975..768bdca6b 100644 --- a/backend/templates/sql_examples/AWS_Redshift.yaml +++ b/backend/templates/sql_examples/AWS_Redshift.yaml @@ -51,6 +51,7 @@ template: COUNT("t1"."订单ID") AS "total_orders", CONCAT(ROUND("t1"."折扣率" * 100, 2), '%') AS "discount_percent" FROM "TEST"."SALES" "t1" + GROUP BY "t1"."订单ID", "t1"."金额", "t1"."折扣率" LIMIT 100 diff --git a/backend/templates/sql_examples/DM.yaml b/backend/templates/sql_examples/DM.yaml index 29e65d82c..df050f6e8 100644 --- a/backend/templates/sql_examples/DM.yaml +++ b/backend/templates/sql_examples/DM.yaml @@ -52,6 +52,7 @@ template: COUNT("t1"."订单ID") AS "total_orders", TO_CHAR("t1"."折扣率" * 100, '990.99') || '%' AS "discount_percent" FROM "TEST"."ORDERS" "t1" + GROUP BY "t1"."订单ID", "t1"."金额", "t1"."折扣率" LIMIT 100 diff --git a/backend/templates/sql_examples/Doris.yaml b/backend/templates/sql_examples/Doris.yaml index b973c529d..9e16b7593 100644 --- a/backend/templates/sql_examples/Doris.yaml +++ b/backend/templates/sql_examples/Doris.yaml @@ -53,6 +53,7 @@ template: COUNT(`t1`.`订单ID`) AS `total_orders`, CONCAT(ROUND(`t1`.`折扣率` * 100, 2), '%') AS `discount_percent` FROM `test`.`orders` `t1` + GROUP BY `t1`.`订单ID`, `t1`.`金额`, `t1`.`折扣率` LIMIT 100 diff --git a/backend/templates/sql_examples/Hive.yaml b/backend/templates/sql_examples/Hive.yaml index 813f6ab50..f7998b930 100644 --- a/backend/templates/sql_examples/Hive.yaml +++ b/backend/templates/sql_examples/Hive.yaml @@ -52,6 +52,7 @@ template: COUNT(`t1`.`订单ID`) AS `total_orders`, CONCAT(CAST(ROUND(`t1`.`折扣率` * 100, 2) AS STRING), '%') AS `discount_percent` FROM `ods`.`orders` `t1` + GROUP BY `t1`.`订单ID`, `t1`.`金额`, `t1`.`折扣率` LIMIT 100 diff --git a/backend/templates/sql_examples/Microsoft_SQL_Server.yaml b/backend/templates/sql_examples/Microsoft_SQL_Server.yaml index 856ac2a6a..31bc77c5f 100644 --- a/backend/templates/sql_examples/Microsoft_SQL_Server.yaml +++ b/backend/templates/sql_examples/Microsoft_SQL_Server.yaml @@ -52,6 +52,7 @@ template: COUNT([o].[订单ID]) AS [total_orders], CONVERT(VARCHAR, ROUND([o].[折扣率] * 100, 2)) + '%' AS [discount_percent] FROM [Sales].[Orders] [o] + GROUP BY [o].[订单ID], [o].[金额], [o].[折扣率] diff --git a/backend/templates/sql_examples/MySQL.yaml b/backend/templates/sql_examples/MySQL.yaml index 1335ec9e4..36e27e818 100644 --- a/backend/templates/sql_examples/MySQL.yaml +++ b/backend/templates/sql_examples/MySQL.yaml @@ -51,6 +51,7 @@ template: COUNT(`t1`.`订单ID`) AS `total_orders`, CONCAT(ROUND(`t1`.`折扣率` * 100, 2), '%') AS `discount_percent` FROM `test`.`orders` `t1` + GROUP BY `t1`.`订单ID`, `t1`.`金额`, `t1`.`折扣率` LIMIT 100 diff --git a/backend/templates/sql_examples/Oracle.yaml b/backend/templates/sql_examples/Oracle.yaml index 26e75297b..e3c52288b 100644 --- a/backend/templates/sql_examples/Oracle.yaml +++ b/backend/templates/sql_examples/Oracle.yaml @@ -116,6 +116,7 @@ template: COUNT("t1"."订单ID") AS "total_orders", TO_CHAR("t1"."折扣率" * 100, '990.99') || '%' AS "discount_percent" FROM "TEST"."ORDERS" "t1" + GROUP BY "t1"."订单ID", "t1"."金额", "t1"."折扣率" WHERE ROWNUM <= 100 diff --git a/backend/templates/sql_examples/PostgreSQL.yaml b/backend/templates/sql_examples/PostgreSQL.yaml index 42a9cfe41..268e8b7b5 100644 --- a/backend/templates/sql_examples/PostgreSQL.yaml +++ b/backend/templates/sql_examples/PostgreSQL.yaml @@ -46,6 +46,7 @@ template: COUNT("t1"."订单ID") AS "total_orders", ROUND("t1"."折扣率" * 100, 2) || '%' AS "discount_percent" FROM "TEST"."ORDERS" "t1" + GROUP BY "t1"."订单ID", "t1"."金额", "t1"."折扣率" LIMIT 100 diff --git a/backend/templates/sql_examples/SQLite.yaml b/backend/templates/sql_examples/SQLite.yaml deleted file mode 100644 index bfcaacc94..000000000 --- a/backend/templates/sql_examples/SQLite.yaml +++ /dev/null @@ -1,81 +0,0 @@ -template: - quot_rule: | - - 必须对数据库名、表名、字段名、别名外层加双引号(")。 - - 1. 点号(.)不能包含在引号内,必须写成 "table" - 2. 即使标识符不含特殊字符或非关键字,也需强制加双引号 - - - - limit_rule: | - - 当需要限制行数时,必须使用标准的LIMIT语法 - - - other_rule: | - 必须为每个表生成别名(不加AS) - {multi_table_condition} - 禁止使用星号(*),必须明确字段名 - 中文/特殊字符字段需保留原名并添加英文别名 - 函数字段必须加别名 - 百分比字段保留两位小数并以%结尾 - 避免与数据库关键字冲突 - - basic_example: | - - - 📌 以下示例严格遵循中的 SQLite 规范,展示符合要求的 SQL 写法与典型错误案例。 - ⚠️ 注意:示例中的表名、字段名均为演示虚构,实际使用时需替换为用户提供的真实标识符。 - 🔍 重点观察: - 1. 双引号包裹所有数据库对象的规范用法 - 2. 中英别名/百分比/函数等特殊字段的处理 - 3. 关键字冲突的规避方式 - - - 查询 ORDERS 表的前100条订单(含中文字段和百分比) - - SELECT * FROM ORDERS LIMIT 100 -- 错误:未加引号、使用星号 - SELECT "订单ID", "金额" FROM "ORDERS" "t1" LIMIT 100 -- 错误:缺少英文别名 - SELECT COUNT("订单ID") FROM "ORDERS" "t1" -- 错误:函数未加别名 - - - SELECT - "t1"."订单ID" AS "order_id", - "t1"."金额" AS "amount", - COUNT("t1"."订单ID") AS "total_orders", - ROUND("t1"."折扣率" * 100, 2) || '%' AS "discount_percent" - FROM "ORDERS" "t1" - LIMIT 100 - - - - - 统计用户表 USERS(含关键字字段user)的活跃占比 - - SELECT user, status FROM USERS -- 错误:未处理关键字和引号 - SELECT "user", ROUND(active_ratio) FROM "USERS" -- 错误:百分比格式错误 - - - SELECT - "u"."user" AS "username", - ROUND("u"."active_ratio" * 100, 2) || '%' AS "active_percent" - FROM "USERS" "u" - WHERE "u"."status" = 1 - - - - - example_engine: SQLite 3.x - example_answer_1: | - {"success":true,"sql":"SELECT \"country_name\", \"continent_name\", \"year\", \"gdp\" FROM \"sample_country_gdp\" ORDER BY \"country_name\", \"year\"","tables":["sample_country_gdp"],"chart-type":"line"} - example_answer_1_with_limit: | - {"success":true,"sql":"SELECT \"country_name\", \"continent_name\", \"year\", \"gdp\" FROM \"sample_country_gdp\" ORDER BY \"country_name\", \"year\" LIMIT 1000","tables":["sample_country_gdp"],"chart-type":"line"} - example_answer_2: | - {"success":true,"sql":"SELECT \"country_name\", \"gdp\" FROM \"sample_country_gdp\" WHERE \"year\" = '2024' ORDER BY \"gdp\" DESC","tables":["sample_country_gdp"],"chart-type":"pie"} - example_answer_2_with_limit: | - {"success":true,"sql":"SELECT \"country_name\", \"gdp\" FROM \"sample_country_gdp\" WHERE \"year\" = '2024' ORDER BY \"gdp\" DESC LIMIT 1000","tables":["sample_country_gdp"],"chart-type":"pie"} - example_answer_3: | - {"success":true,"sql":"SELECT \"country_name\", \"gdp\" FROM \"sample_country_gdp\" WHERE \"year\" = '2025' AND \"country_name\" = '中国'","tables":["sample_country_gdp"],"chart-type":"table"} - example_answer_3_with_limit: | - {"success":true,"sql":"SELECT \"country_name\", \"gdp\" FROM \"sample_country_gdp\" WHERE \"year\" = '2025' AND \"country_name\" = '中国' LIMIT 1000","tables":["sample_country_gdp"],"chart-type":"table"} diff --git a/frontend/index.html b/frontend/index.html index 97c846b88..382730d02 100644 --- a/frontend/index.html +++ b/frontend/index.html @@ -4,7 +4,7 @@ - SQLBot +
diff --git a/frontend/package.json b/frontend/package.json index 14d623762..280710df8 100644 --- a/frontend/package.json +++ b/frontend/package.json @@ -11,7 +11,7 @@ }, "scripts": { "dev": "vue-tsc -b && vite", - "build": "vue-tsc -b && vite build", + "build": "vue-tsc -b &&NODE_OPTIONS=--max_old_space_size=4096 vite build", "preview": "vite preview", "lint": "eslint . --ext .vue,.js,.ts,.jsx,.tsx --fix" }, @@ -24,6 +24,7 @@ "@npkg/tinymce-plugins": "^0.0.7", "@tinymce/tinymce-vue": "^5.1.0", "@vueuse/core": "^14.1.0", + "core-js": "^3.49.0", "dayjs": "^1.11.13", "element-plus": "^2.10.1", "element-plus-secondary": "^1.0.0", @@ -55,17 +56,23 @@ "@types/node": "^22.14.1", "@typescript-eslint/eslint-plugin": "^8.34.0", "@typescript-eslint/parser": "^8.34.0", + "@vitejs/plugin-legacy": "^6.1.1", "@vitejs/plugin-vue": "^5.2.2", "@vue/tsconfig": "^0.7.0", + "autoprefixer": "^10.5.2", "axios": "^1.8.4", "crypto-js": "^4.2.0", + "css-has-pseudo": "^8.0.0", "eslint": "^9.28.0", "eslint-config-prettier": "^10.1.5", "eslint-plugin-prettier": "^5.4.1", "eslint-plugin-vue": "^10.2.0", + "flex-gap-polyfill": "^5.0.0", "globals": "^16.2.0", "less": "4.4.2", "pinia": "^3.0.2", + "postcss": "^8.5.16", + "postcss-preset-env": "^11.3.2", "prettier": "^3.5.3", "typescript": "~5.7.2", "typescript-eslint": "^8.34.0", @@ -76,4 +83,4 @@ "vite-svg-loader": "^5.1.0", "vue-tsc": "^2.2.8" } -} \ No newline at end of file +} diff --git a/frontend/postcss.config.js b/frontend/postcss.config.js new file mode 100644 index 000000000..bca102202 --- /dev/null +++ b/frontend/postcss.config.js @@ -0,0 +1,14 @@ +import autoprefixer from 'autoprefixer' +import postcssPresetEnv from 'postcss-preset-env' +import cssHasPseudo from 'css-has-pseudo' + +export default { + plugins: [ + cssHasPseudo({ preserve: true }), + autoprefixer(), + postcssPresetEnv({ + stage: 3, + browsers: 'Chrome 81', + }), + ], +} diff --git a/frontend/public/assistant.js b/frontend/public/assistant.js deleted file mode 100644 index c67613be4..000000000 --- a/frontend/public/assistant.js +++ /dev/null @@ -1,872 +0,0 @@ -; (function () { - window.sqlbot_assistant_handler = window.sqlbot_assistant_handler || {} - const defaultData = { - id: '1', - show_guide: false, - float_icon: '', - domain_url: 'http://localhost:5173', - header_font_color: 'rgb(100, 106, 115)', - x_type: 'right', - y_type: 'bottom', - x_val: '30', - y_val: '30', - float_icon_drag: false, - } - const script_id_prefix = 'sqlbot-assistant-float-script-' - const guideHtml = ` -
-
-
-
-
- - - -
- -
🌟 遇见问题,不再有障碍!
-

你好,我是你的智能小助手。
- 点我,开启高效解答模式,让问题变成过去式。

-
- -
- -
-` - - const chatButtonHtml = (data) => ` -
- - - - - - - - - - - - - - - - - - - - - - - - - - - - -
` - - const getChatContainerHtml = (data) => { - let srcUrl = `${data.domain_url}/#/assistant?id=${data.id}&online=${!!data.online}&name=${encodeURIComponent(data.name)}` - if (data.userFlag) { - srcUrl += `&userFlag=${data.userFlag || ''}` - } - if (data.history) { - srcUrl += `&history=${data.history}` - } - return ` -
- -
-
- - - -
-
- - - -
-
- - - -
-
-` - } - - function getHighestZIndexValue() { - try { - let maxZIndex = -Infinity - let foundAny = false - - const allElements = document.all || document.querySelectorAll('*') - - for (let i = 0; i < allElements.length; i++) { - const element = allElements[i] - - if (!element || element.nodeType !== 1) continue - - const styles = window.getComputedStyle(element) - - const position = styles.position - if (position === 'static') continue - - const zIndex = styles.zIndex - let zIndexValue - - if (zIndex === 'auto') { - zIndexValue = 0 - } else { - zIndexValue = parseInt(zIndex, 10) - if (isNaN(zIndexValue)) continue - } - - foundAny = true - - // 快速返回:如果找到很大的z-index,很可能就是最大值 - /* if (zIndexValue > 10000) { - return zIndexValue; - } */ - - if (zIndexValue > maxZIndex) { - maxZIndex = zIndexValue - } - } - return foundAny ? maxZIndex : 0 - } catch (error) { - console.warn('获取最高z-index时出错,返回默认值0:', error) - return 0 - } - } - - /** - * 初始化引导 - * @param {*} root - */ - const initGuide = (root) => { - root.insertAdjacentHTML('beforeend', guideHtml) - const button = root.querySelector('.sqlbot-assistant-button') - const close_icon = root.querySelector('.sqlbot-assistant-close') - const close_func = () => { - root.removeChild(root.querySelector('.sqlbot-assistant-tips')) - root.removeChild(root.querySelector('.sqlbot-assistant-mask')) - localStorage.setItem('sqlbot_assistant_mask_tip', true) - } - button.onclick = close_func - close_icon.onclick = close_func - } - const initChat = (root, data) => { - // 添加对话icon - root.insertAdjacentHTML('beforeend', chatButtonHtml(data)) - // 添加对话框 - root.insertAdjacentHTML('beforeend', getChatContainerHtml(data)) - // 按钮元素 - const chat_button = root.querySelector('.sqlbot-assistant-chat-button') - let chat_button_img = root.querySelector('.sqlbot-assistant-chat-button > svg') - if (data.float_icon) { - chat_button_img = root.querySelector('.sqlbot-assistant-chat-button > img') - } - chat_button_img.style.display = 'block' - function resizeImg() { - const rate = window.outerWidth / window.innerWidth; - chat_button_img.style.width = `${30 * (1 / rate)}px`; - chat_button_img.style.height = `${30 * (1 / rate)}px`; - } - resizeImg() - window.addEventListener('resize', resizeImg); - // 对话框元素 - const chat_container = root.querySelector('#sqlbot-assistant-chat-container') - // 引导层 - const mask_content = root.querySelector('.sqlbot-assistant-mask > .sqlbot-assistant-content') - const mask_tips = root.querySelector('.sqlbot-assistant-tips') - chat_button_img.onload = (event) => { - if (mask_content) { - mask_content.style.width = chat_button_img.width + 'px' - mask_content.style.height = chat_button_img.height + 'px' - if (data.x_type == 'left') { - mask_tips.style.marginLeft = - (chat_button_img.naturalWidth > 500 ? 500 : chat_button_img.naturalWidth) - 64 + 'px' - } else { - mask_tips.style.marginRight = - (chat_button_img.naturalWidth > 500 ? 500 : chat_button_img.naturalWidth) - 64 + 'px' - } - } - } - - const viewport = root.querySelector('.sqlbot-assistant-openviewport') - const closeviewport = root.querySelector('.sqlbot-assistant-closeviewport') - const close_func = () => { - chat_container.style['display'] = - chat_container.style['display'] == 'block' ? 'none' : 'block' - chat_button.style['display'] = chat_container.style['display'] == 'block' ? 'none' : 'block' - } - close_icon = chat_container.querySelector('.sqlbot-assistant-chat-close') - chat_button.onclick = close_func - close_icon.onclick = close_func - const viewport_func = () => { - if (chat_container.classList.contains('sqlbot-assistant-enlarge')) { - chat_container.classList.remove('sqlbot-assistant-enlarge') - viewport.classList.remove('sqlbot-assistant-viewportnone') - closeviewport.classList.add('sqlbot-assistant-viewportnone') - } else { - chat_container.classList.add('sqlbot-assistant-enlarge') - viewport.classList.add('sqlbot-assistant-viewportnone') - closeviewport.classList.remove('sqlbot-assistant-viewportnone') - } - } - if (data.float_icon_drag) { - chat_button.setAttribute('draggable', 'true') - - let startX = 0 - let startY = 0 - const img = new Image() - img.src = 'data:image/gif;base64,R0lGODlhAQABAIAAAAUEBAAAACwAAAAAAQABAAACAkQBADs=' - chat_button.addEventListener('dragstart', (e) => { - startX = e.clientX - chat_button.offsetLeft - startY = e.clientY - chat_button.offsetTop - e.dataTransfer.setDragImage(img, 0, 0) - }) - - chat_button.addEventListener('drag', (e) => { - if (e.clientX && e.clientY) { - const left = e.clientX - startX - const top = e.clientY - startY - - const maxX = window.innerWidth - chat_button.offsetWidth - const maxY = window.innerHeight - chat_button.offsetHeight - - chat_button.style.left = Math.min(Math.max(0, left), maxX) + 'px' - chat_button.style.top = Math.min(Math.max(0, top), maxY) + 'px' - } - }) - - let touchStartX = 0 - let touchStartY = 0 - - chat_button.addEventListener('touchstart', (e) => { - touchStartX = e.touches[0].clientX - chat_button.offsetLeft - touchStartY = e.touches[0].clientY - chat_button.offsetTop - e.preventDefault() - }) - - chat_button.addEventListener('touchmove', (e) => { - const left = e.touches[0].clientX - touchStartX - const top = e.touches[0].clientY - touchStartY - - const maxX = window.innerWidth - chat_button.offsetWidth - const maxY = window.innerHeight - chat_button.offsetHeight - - chat_button.style.left = Math.min(Math.max(0, left), maxX) + 'px' - chat_button.style.top = Math.min(Math.max(0, top), maxY) + 'px' - - e.preventDefault() - }) - } - /* const drag = (e) => { - if (['touchmove', 'touchstart'].includes(e.type)) { - chat_button.style.top = e.touches[0].clientY - chat_button_img.clientHeight / 2 + 'px' - chat_button.style.left = e.touches[0].clientX - chat_button_img.clientHeight / 2 + 'px' - } else { - chat_button.style.top = e.y - chat_button_img.clientHeight / 2 + 'px' - chat_button.style.left = e.x - chat_button_img.clientHeight / 2 + 'px' - } - chat_button.style.width = chat_button_img.clientHeight + 'px' - chat_button.style.height = chat_button_img.clientHeight + 'px' - } - if (data.float_icon_drag) { - chat_button.setAttribute('draggable', 'true') - chat_button.addEventListener('drag', drag) - chat_button.addEventListener('dragover', (e) => { - e.preventDefault() - }) - chat_button.addEventListener('dragend', drag) - chat_button.addEventListener('touchstart', drag) - chat_button.addEventListener('touchmove', drag) - } */ - viewport.onclick = viewport_func - closeviewport.onclick = viewport_func - } - /** - * 第一次进来的引导提示 - */ - function initsqlbot_assistant(data) { - const sqlbot_div = document.createElement('div') - const root = document.createElement('div') - const sqlbot_root_id = 'sqlbot-assistant-root-' + data.id - root.id = sqlbot_root_id - initsqlbot_assistantStyle(sqlbot_div, sqlbot_root_id, data) - sqlbot_div.appendChild(root) - document.body.appendChild(sqlbot_div) - const sqlbot_assistant_mask_tip = localStorage.getItem('sqlbot_assistant_mask_tip') - if (sqlbot_assistant_mask_tip == null && data.show_guide) { - initGuide(root) - } - initChat(root, data) - } - - // 初始化全局样式 - function initsqlbot_assistantStyle(root, sqlbot_assistantId, data) { - const maxZIndex = getHighestZIndexValue() - const zIndex = Math.max((maxZIndex || 0) + 1, 10000) - const maskZIndex = zIndex + 1 - style = document.createElement('style') - style.type = 'text/css' - style.innerText = ` - /* 放大 */ - #sqlbot-assistant .sqlbot-assistant-enlarge { - width: 50%!important; - height: 100%!important; - bottom: 0!important; - right: 0 !important; - } - @media only screen and (max-width: 768px){ - #sqlbot-assistant .sqlbot-assistant-enlarge { - width: 100%!important; - height: 100%!important; - right: 0 !important; - bottom: 0!important; - } - } - - /* 引导 */ - - #sqlbot-assistant .sqlbot-assistant-mask { - position: fixed; - z-index: ${maskZIndex}; - background-color: transparent; - height: 100%; - width: 100%; - top: 0; - left: 0; - } - #sqlbot-assistant .sqlbot-assistant-mask .sqlbot-assistant-content { - width: 64px; - height: 64px; - box-shadow: 1px 1px 1px 9999px rgba(0,0,0,.6); - position: absolute; - ${data.x_type}: ${data.x_val}px; - ${data.y_type}: ${data.y_val}px; - z-index: ${maskZIndex}; - } - #sqlbot-assistant .sqlbot-assistant-tips { - position: fixed; - ${data.x_type}:calc(${data.x_val}px + 75px); - ${data.y_type}: calc(${data.y_val}px + 0px); - padding: 22px 24px 24px; - border-radius: 6px; - color: #ffffff; - font-size: 14px; - background: #3370FF; - z-index: ${maskZIndex}; - } - #sqlbot-assistant .sqlbot-assistant-tips .sqlbot-assistant-arrow { - position: absolute; - background: #3370FF; - width: 10px; - height: 10px; - pointer-events: none; - transform: rotate(45deg); - box-sizing: border-box; - /* left */ - ${data.x_type}: -5px; - ${data.y_type}: 33px; - border-left-color: transparent; - border-bottom-color: transparent - } - #sqlbot-assistant .sqlbot-assistant-tips .sqlbot-assistant-title { - font-size: 20px; - font-weight: 500; - margin-bottom: 8px; - } - #sqlbot-assistant .sqlbot-assistant-tips .sqlbot-assistant-button { - text-align: right; - margin-top: 24px; - } - #sqlbot-assistant .sqlbot-assistant-tips .sqlbot-assistant-button button { - border-radius: 4px; - background: #FFF; - padding: 3px 12px; - color: #3370FF; - cursor: pointer; - outline: none; - border: none; - } - #sqlbot-assistant .sqlbot-assistant-tips .sqlbot-assistant-button button::after{ - border: none; - } - #sqlbot-assistant .sqlbot-assistant-tips .sqlbot-assistant-close { - position: absolute; - right: 20px; - top: 20px; - cursor: pointer; - - } - #sqlbot-assistant-chat-container { - width: 460px; - height: 640px; - display:none; - } - @media only screen and (max-width: 768px) { - #sqlbot-assistant-chat-container { - width: 100%; - height: 70%; - right: 0 !important; - } - } - - #sqlbot-assistant .sqlbot-assistant-chat-button{ - position: fixed; - ${data.x_type}: ${data.x_val}px; - ${data.y_type}: ${data.y_val}px; - cursor: pointer; - z-index: ${zIndex}; - } - #sqlbot-assistant #sqlbot-assistant-chat-container{ - z-index: ${zIndex}; - position: relative; - border-radius: 8px; - //border: 1px solid #ffffff; - background: linear-gradient(188deg, rgba(235, 241, 255, 0.20) 39.6%, rgba(231, 249, 255, 0.20) 94.3%), #EFF0F1; - box-shadow: 0px 4px 8px 0px rgba(31, 35, 41, 0.10); - position: fixed;bottom: 16px;right: 16px;overflow: hidden; - } - - .ed-overlay-dialog { - margin-top: 50px; - } - .ed-drawer { - margin-top: 50px; - } - - #sqlbot-assistant #sqlbot-assistant-chat-container .sqlbot-assistant-operate{ - top: 18px; - right: 15px; - position: absolute; - display: flex; - align-items: center; - line-height: 18px; - } - #sqlbot-assistant #sqlbot-assistant-chat-container .sqlbot-assistant-operate .sqlbot-assistant-chat-close{ - margin-left:15px; - cursor: pointer; - } - #sqlbot-assistant #sqlbot-assistant-chat-container .sqlbot-assistant-operate .sqlbot-assistant-openviewport{ - - cursor: pointer; - } - #sqlbot-assistant #sqlbot-assistant-chat-container .sqlbot-assistant-operate .sqlbot-assistant-closeviewport{ - - cursor: pointer; - } - #sqlbot-assistant #sqlbot-assistant-chat-container .sqlbot-assistant-viewportnone{ - display:none; - } - #sqlbot-assistant #sqlbot-assistant-chat-container #sqlbot-assistant-chat-iframe-${data.id} { - height:100%; - width:100%; - border: none; - } - #sqlbot-assistant #sqlbot-assistant-chat-container { - animation: appear .4s ease-in-out; - } - @keyframes appear { - from { - height: 0;; - } - - to { - height: 600px; - } - }`.replaceAll('#sqlbot-assistant ', `#${sqlbot_assistantId} `) - root.appendChild(style) - } - function getParam(src, key) { - const url = new URL(src) - return url.searchParams.get(key) - } - function parsrCertificate(config) { - const certificateList = config.certificate - if (!certificateList?.length) { - return null - } - const list = certificateList.map((item) => formatCertificate(item)).filter((item) => !!item) - return JSON.stringify(list) - } - function isEmpty(obj) { - return obj == null || typeof obj == 'undefined' - } - function formatCertificate(item) { - const { type, source, target, target_key, target_val } = item - let source_val = null - if (type.toLocaleLowerCase() == 'localstorage') { - source_val = localStorage.getItem(source) - } - if (type.toLocaleLowerCase() == 'sessionstorage') { - source_val = sessionStorage.getItem(source) - } - if (type.toLocaleLowerCase() == 'cookie') { - source_val = getCookie(source) - } - if (type.toLocaleLowerCase() == 'custom') { - source_val = source - } - if (isEmpty(source_val)) { - return null - } - return { - target, - key: target_key || source, - value: (target_val && eval(target_val)) || source_val, - } - } - function getCookie(key) { - if (!key || !document.cookie) { - return null - } - const cookies = document.cookie.split(';') - for (let i = 0; i < cookies.length; i++) { - const cookie = cookies[i].trim() - - if (cookie.startsWith(key + '=')) { - return decodeURIComponent(cookie.substring(key.length + 1)) - } - } - return null - } - function registerMessageEvent(id, data) { - const iframe = document.getElementById(`sqlbot-assistant-chat-iframe-${id}`) - const url = iframe.src - const eventName = 'sqlbot_assistant_event' - window.addEventListener('message', (event) => { - if (event.data?.eventName === eventName) { - if (event.data?.messageId !== id) { - return - } - if (event.data?.busi == 'ready' && event.data?.ready) { - params = { - eventName, - messageId: id, - hostOrigin: window.location.origin, - } - if (data.type === 1) { - const certificate = parsrCertificate(data) - params['busi'] = 'certificate' - params['certificate'] = certificate - } - const contentWindow = iframe.contentWindow - contentWindow.postMessage(params, url) - } - } - }) - } - function loadScript(src, id) { - const domain_url = getDomain(src) - const online = getParam(src, 'online') - const userFlag = getParam(src, 'userFlag') - const history = getParam(src, 'history') - let url = `${domain_url}/api/v1/system/assistant/info/${id}` - if (domain_url.includes('5173')) { - url = url.replace('5173', '8000') - } - fetch(url) - .then((response) => response.json()) - .then((res) => { - if (!res.data) { - throw new Error(res) - } - const data = res.data - const config_json = data.configuration - let tempData = Object.assign(defaultData, data) - if (tempData.configuration) { - delete tempData.configuration - } - if (config_json) { - const config = JSON.parse(config_json) - if (config) { - delete config.id - tempData = Object.assign(tempData, config) - } - } - tempData['id'] = id - tempData['domain_url'] = domain_url - - if (tempData['float_icon'] && !tempData['float_icon'].startsWith('http://')) { - tempData['float_icon'] = - `${domain_url}/api/v1/system/assistant/picture/${tempData['float_icon']}` - - if (domain_url.includes('5173')) { - tempData['float_icon'] = tempData['float_icon'].replace('5173', '8000') - } - } - - tempData['online'] = online && online.toString().toLowerCase() == 'true' - tempData['userFlag'] = userFlag - tempData['history'] = history - initsqlbot_assistant(tempData) - registerMessageEvent(id, tempData) - }) - .catch((e) => { - showMsg('嵌入失败', e.message) - }) - } - function getDomain(src) { - return src.substring(0, src.indexOf('/assistant.js')) - } - function init() { - const sqlbotScripts = document.querySelectorAll(`script[id^="${script_id_prefix}"]`) - const scriptsArray = Array.from(sqlbotScripts) - const src_list = scriptsArray.map((script) => script.src) - src_list.forEach((src) => { - const id = getParam(src, 'id') - window.sqlbot_assistant_handler[id] = window.sqlbot_assistant_handler[id] || {} - window.sqlbot_assistant_handler[id]['id'] = id - const propName = script_id_prefix + id + '-state' - if (window[propName]) { - return true - } - window[propName] = true - loadScript(src, id) - expposeGlobalMethods(id) - }) - } - - function showMsg(title, content) { - // 检查并创建容器(如果不存在) - let container = document.getElementById('messageContainer') - if (!container) { - container = document.createElement('div') - container.id = 'messageContainer' - container.style.position = 'fixed' - container.style.bottom = '20px' - container.style.right = '20px' - container.style.zIndex = '1000' - document.body.appendChild(container) - } else { - // 如果容器已存在,先移除旧弹窗 - const oldMessage = container.querySelector('div') - if (oldMessage) { - oldMessage.style.transform = 'translateX(120%)' - oldMessage.style.opacity = '0' - setTimeout(() => { - container.removeChild(oldMessage) - }, 300) - } - } - - // 创建弹窗元素 - const messageBox = document.createElement('div') - messageBox.style.width = '240px' - messageBox.style.minHeight = '100px' - messageBox.style.background = 'linear-gradient(135deg, #ff6b6b, #ff8e8e)' - messageBox.style.borderRadius = '8px' - messageBox.style.boxShadow = '0 4px 12px rgba(0, 0, 0, 0.15)' - messageBox.style.padding = '15px' - messageBox.style.color = 'white' - messageBox.style.fontFamily = 'Arial, sans-serif' - messageBox.style.display = 'flex' - messageBox.style.flexDirection = 'column' - messageBox.style.transform = 'translateX(120%)' - messageBox.style.transition = 'transform 0.3s ease-out' - messageBox.style.opacity = '0' - messageBox.style.transition = 'opacity 0.3s ease, transform 0.3s ease' - messageBox.style.overflow = 'hidden' - - // 创建标题元素 - const titleElement = document.createElement('div') - titleElement.style.fontSize = '18px' - titleElement.style.fontWeight = 'bold' - titleElement.style.marginBottom = '10px' - titleElement.style.borderBottom = '1px solid rgba(255, 255, 255, 0.3)' - titleElement.style.paddingBottom = '8px' - titleElement.textContent = title - - // 创建内容元素 - const contentElement = document.createElement('div') - contentElement.style.fontSize = '14px' - contentElement.style.flexGrow = '1' - contentElement.style.overflow = 'auto' - contentElement.textContent = content - - // 组装元素 - messageBox.appendChild(titleElement) - messageBox.appendChild(contentElement) - - // 添加到容器 - container.appendChild(messageBox) - - // 触发显示动画 - setTimeout(() => { - messageBox.style.transform = 'translateX(0)' - messageBox.style.opacity = '1' - }, 10) - - // 3秒后自动隐藏 - setTimeout(() => { - messageBox.style.transform = 'translateX(120%)' - messageBox.style.opacity = '0' - setTimeout(() => { - container.removeChild(messageBox) - // 如果容器是空的,也移除容器 - if (container.children.length === 0) { - document.body.removeChild(container) - } - }, 300) - }, 5000) - } - - /* function hideMsg() { - const container = document.getElementById('messageContainer'); - if (container) { - const messageBox = container.querySelector('div'); - if (messageBox) { - messageBox.style.transform = 'translateX(120%)'; - messageBox.style.opacity = '0'; - setTimeout(() => { - container.removeChild(messageBox); - // 如果容器是空的,也移除容器 - if (container.children.length === 0) { - document.body.removeChild(container); - } - }, 300); - } - } - } */ - - function updateParam(target_url, key, newValue) { - try { - const url = new URL(target_url) - const [hashPath, hashQuery] = url.hash.split('?') - let searchParams - if (hashQuery) { - searchParams = new URLSearchParams(hashQuery) - } else { - searchParams = url.searchParams - } - searchParams.set(key, newValue) - if (hashQuery) { - url.hash = `${hashPath}?${searchParams.toString()}` - } else { - url.search = searchParams.toString() - } - return url.toString() - } catch (e) { - console.error('Invalid URL:', target_url) - return target_url - } - } - function expposeGlobalMethods(id) { - window.sqlbot_assistant_handler[id]['setOnline'] = (online) => { - if (online != null && typeof online != 'boolean') { - throw new Error('The parameter can only be of type boolean') - } - const iframe = document.getElementById(`sqlbot-assistant-chat-iframe-${id}`) - if (iframe) { - const url = iframe.src - const eventName = 'sqlbot_assistant_event' - const params = { - busi: 'setOnline', - online, - eventName, - messageId: id, - } - const contentWindow = iframe.contentWindow - contentWindow.postMessage(params, url) - } - } - window.sqlbot_assistant_handler[id]['refresh'] = (online, userFlag) => { - if (online != null && typeof online != 'boolean') { - throw new Error('The parameter can only be of type boolean') - } - const iframe = document.getElementById(`sqlbot-assistant-chat-iframe-${id}`) - if (iframe) { - const url = iframe.src - let new_url = updateParam(url, 't', Date.now()) - if (online != null) { - new_url = updateParam(new_url, 'online', online) - } - if (userFlag != null) { - new_url = updateParam(new_url, 'userFlag', userFlag) - } - iframe.src = 'about:blank' - setTimeout(() => { - iframe.src = new_url - }, 500) - } - } - window.sqlbot_assistant_handler[id]['destroy'] = () => { - const sqlbot_root_id = 'sqlbot-assistant-root-' + id - const container_div = document.getElementById(sqlbot_root_id) - if (container_div) { - const root_div = container_div.parentNode - if (root_div?.parentNode) { - root_div.parentNode.removeChild(root_div) - } - } - - const scriptDom = document.getElementById(`sqlbot-assistant-float-script-${id}`) - if (scriptDom) { - scriptDom.parentNode.removeChild(scriptDom) - } - const propName = script_id_prefix + id + '-state' - if (window[propName]) { - delete window[propName] - } - delete window.sqlbot_assistant_handler[id] - } - window.sqlbot_assistant_handler[id]['setHistory'] = (show) => { - if (show != null && typeof show != 'boolean') { - throw new Error('The parameter can only be of type boolean') - } - const iframe = document.getElementById(`sqlbot-assistant-chat-iframe-${id}`) - if (iframe) { - const url = iframe.src - const eventName = 'sqlbot_assistant_event' - const params = { - busi: 'setHistory', - show, - eventName, - messageId: id, - } - const contentWindow = iframe.contentWindow - contentWindow.postMessage(params, url) - } - } - window.sqlbot_assistant_handler[id]['createConversation'] = (param) => { - const iframe = document.getElementById(`sqlbot-assistant-chat-iframe-${id}`) - if (iframe) { - const url = iframe.src - const eventName = 'sqlbot_assistant_event' - const params = { - busi: 'createConversation', - param, - eventName, - messageId: id, - } - const contentWindow = iframe.contentWindow - contentWindow.postMessage(params, url) - } - } - } - // window.addEventListener('load', init) - const executeWhenReady = (fn) => { - if ( - document.readyState === 'complete' || - (document.readyState !== 'loading' && !document.documentElement.doScroll) - ) { - setTimeout(fn, 0) - } else { - const onReady = () => { - document.removeEventListener('DOMContentLoaded', onReady) - window.removeEventListener('load', onReady) - fn() - } - document.addEventListener('DOMContentLoaded', onReady) - window.addEventListener('load', onReady) - } - } - - executeWhenReady(init) -})() diff --git a/frontend/src/api/assistant.ts b/frontend/src/api/assistant.ts index bb955e09f..b883976f0 100644 --- a/frontend/src/api/assistant.ts +++ b/frontend/src/api/assistant.ts @@ -6,6 +6,6 @@ export const assistantApi = { add: (data: any) => request.post('/system/assistant', data), edit: (data: any) => request.put('/system/assistant', data), delete: (id: number) => request.delete(`/system/assistant/${id}`), - query: (id: number) => request.get(`/system/assistant/${id}`), + // query: (id: number) => request.get(`/system/assistant/${id}`), validate: (data: any) => request.get('/system/assistant/validator', { params: data }), } diff --git a/frontend/src/api/chat.ts b/frontend/src/api/chat.ts index 2a5870632..4a52187ee 100644 --- a/frontend/src/api/chat.ts +++ b/frontend/src/api/chat.ts @@ -474,7 +474,7 @@ export const chatApi = { return request.post('/chat/rename', { id: chat_id, brief: brief }) }, deleteChat: (id: number | undefined, brief: any): Promise => { - return request.delete(`/chat/${id}/${brief}`) + return request.delete(`/chat/${id}`, { data: { id: id, brief: brief } }) }, analysis: (record_id: number | undefined, controller?: AbortController) => { return request.fetchStream(`/chat/record/${record_id}/analysis`, {}, controller) diff --git a/frontend/src/api/embedded.ts b/frontend/src/api/embedded.ts index 0733088ae..2e51008ae 100644 --- a/frontend/src/api/embedded.ts +++ b/frontend/src/api/embedded.ts @@ -5,7 +5,7 @@ export const getAdvancedApplicationList = () => request.get('/system/assistant/advanced_application') export const updateAssistant = (data: any) => request.put('/system/assistant', data) export const saveAssistant = (data: any) => request.post('/system/assistant', data) -export const getOne = (id: any) => request.get(`/system/assistant/${id}`) +// export const getOne = (id: any) => request.get(`/system/assistant/${id}`) export const delOne = (id: any) => request.delete(`/system/assistant/${id}`) export const dsApi = (id: any) => request.get(`/datasource/ws/${id}`) diff --git a/frontend/src/api/system.ts b/frontend/src/api/system.ts index f9ed3ec10..4d56fecb1 100644 --- a/frontend/src/api/system.ts +++ b/frontend/src/api/system.ts @@ -27,6 +27,8 @@ export const modelApi = { query: (id: number) => request.get(`/system/aimodel/${id}`), setDefault: (id: number) => request.put(`/system/aimodel/default/${id}`), check: (data: any) => request.fetchStream('/system/aimodel/status', data), - platform: (id: number) => request.get(`/system/platform/org/${id}`), + platform: (id: number, lazy?: number, pid?: string) => + request.post(`/system/platform/org/${id}`, { lazy, pid }), userSync: (data: any) => request.post(`/system/platform/user/sync`, data), + list_by_ws: () => request.get(`/system/aimodel/list/by_ws`), } diff --git a/frontend/src/api/workspace.ts b/frontend/src/api/workspace.ts index 5298045b5..46637678c 100644 --- a/frontend/src/api/workspace.ts +++ b/frontend/src/api/workspace.ts @@ -15,3 +15,12 @@ export const workspaceDelete = (id: any) => request.delete(`/system/workspace/${ export const workspaceList = () => request.get('/system/workspace') export const workspaceDetail = (id: any) => request.get(`/system/workspace/${id}`) export const uwsOption = (params: any) => request.get('system/workspace/uws/option', { params }) + +export const workspaceModelMapping = (aiModelId: any) => + request.get(`/system/aimodel/${aiModelId}/ws_mapping`) +export const workspaceModelMappingUpdate = (aiModelId: any, data: any) => + request.put(`/system/aimodel/${aiModelId}/ws_mapping`, data) +export const workspaceModelMappingAdd = (aiModelId: any, data: any) => + request.post(`/system/aimodel/${aiModelId}/ws_mapping`, data) +export const workspaceModelMappingDelete = (aiModelId: any, data: any) => + request.delete(`/system/aimodel/${aiModelId}/ws_mapping`, { data }) diff --git a/frontend/src/assets/svg/chart/icon-thousand-separator.svg b/frontend/src/assets/svg/chart/icon-thousand-separator.svg new file mode 100644 index 000000000..670e86901 --- /dev/null +++ b/frontend/src/assets/svg/chart/icon-thousand-separator.svg @@ -0,0 +1,6 @@ + + + diff --git a/frontend/src/assets/svg/icon_block_outlined.svg b/frontend/src/assets/svg/icon_block_outlined.svg new file mode 100644 index 000000000..ffefda2c0 --- /dev/null +++ b/frontend/src/assets/svg/icon_block_outlined.svg @@ -0,0 +1,11 @@ + + + + + + + + + + + diff --git a/frontend/src/assets/svg/icon_describe_outlined.svg b/frontend/src/assets/svg/icon_describe_outlined.svg new file mode 100644 index 000000000..7dc0951b4 --- /dev/null +++ b/frontend/src/assets/svg/icon_describe_outlined.svg @@ -0,0 +1,3 @@ + + + diff --git a/frontend/src/assets/svg/workspace-white.svg b/frontend/src/assets/svg/workspace-white.svg new file mode 100644 index 000000000..7e9cac0a5 --- /dev/null +++ b/frontend/src/assets/svg/workspace-white.svg @@ -0,0 +1,5 @@ + + + + + diff --git a/frontend/src/components/drawer-filter/src/DrawerEnumFilter.vue b/frontend/src/components/drawer-filter/src/DrawerEnumFilter.vue index a7cae5c5d..262b5c710 100644 --- a/frontend/src/components/drawer-filter/src/DrawerEnumFilter.vue +++ b/frontend/src/components/drawer-filter/src/DrawerEnumFilter.vue @@ -84,7 +84,7 @@ useEmitt({ padding: 1px 6px; background: var(--deTextPrimary5, #f5f6f7); color: var(--deTextPrimary, #1f2329); - border-radius: 4px; + border-radius: 6px; cursor: pointer; display: inline-block; margin-bottom: 12px; diff --git a/frontend/src/components/filter-text/src/FilterText.vue b/frontend/src/components/filter-text/src/FilterText.vue index 4eedaa269..d423543d4 100644 --- a/frontend/src/components/filter-text/src/FilterText.vue +++ b/frontend/src/components/filter-text/src/FilterText.vue @@ -181,7 +181,7 @@ watch( .arrow-filter:hover { background: rgba(31, 35, 41, 0.1); - border-radius: 4px; + border-radius: 6px; } .ed-icon-arrow-right.arrow-filter { diff --git a/frontend/src/components/layout/Person.vue b/frontend/src/components/layout/Person.vue index ae0448e26..a58a02afb 100644 --- a/frontend/src/components/layout/Person.vue +++ b/frontend/src/components/layout/Person.vue @@ -345,7 +345,7 @@ const logout = async () => { position: relative; cursor: pointer; margin: 0 4px; - border-radius: 4px; + border-radius: 6px; &:hover { background-color: #1f23291a; } @@ -382,7 +382,7 @@ const logout = async () => { padding-right: 8px; margin-bottom: 2px; position: relative; - border-radius: 4px; + border-radius: 6px; cursor: pointer; &:not(.empty):hover { background: #1f23291a; diff --git a/frontend/src/components/layout/Workspace.vue b/frontend/src/components/layout/Workspace.vue index 1f76e92ad..456c083cd 100644 --- a/frontend/src/components/layout/Workspace.vue +++ b/frontend/src/components/layout/Workspace.vue @@ -198,7 +198,7 @@ onMounted(async () => { padding-right: 8px; margin-bottom: 2px; position: relative; - border-radius: 4px; + border-radius: 6px; cursor: pointer; &:not(.empty):hover { background: #1f23291a; diff --git a/frontend/src/components/layout/index.vue b/frontend/src/components/layout/index.vue index f83ef983a..4a0097e41 100644 --- a/frontend/src/components/layout/index.vue +++ b/frontend/src/components/layout/index.vue @@ -362,7 +362,7 @@ onMounted(() => { column-gap: 12px; align-items: center; padding: 8px 16px; - border-radius: 4px; + border-radius: 6px; cursor: pointer; border: none; font-weight: 500; @@ -493,7 +493,7 @@ onMounted(() => { column-gap: 12px; align-items: center; padding: 8px 16px; - border-radius: 4px; + border-radius: 6px; cursor: pointer; border: none; font-weight: 500; diff --git a/frontend/src/entity/supplier.ts b/frontend/src/entity/supplier.ts index cd672f9d8..bd4ddc971 100644 --- a/frontend/src/entity/supplier.ts +++ b/frontend/src/entity/supplier.ts @@ -276,7 +276,7 @@ export const supplierList: Array<{ { name: 'doubao-1-5-pro-32k-character-250715' }, { name: 'kimi-k2-250711' }, { name: 'deepseek-v3-250324' }, - { name: 'deepseek-r1' }, + { name: 'deepseek-v4-pro-260425' }, ], }, }, @@ -290,11 +290,7 @@ export const supplierList: Array<{ 0: { api_domain: 'https://api.minimax.io/v1', common_args: [{ key: 'temperature', val: 0.7, type: 'number', range: '[0, 1]' }], - model_options: [ - { name: 'MiniMax-M2.7' }, - { name: 'MiniMax-M2.5' }, - { name: 'MiniMax-M2.5-highspeed' }, - ], + model_options: [{ name: 'MiniMax-M3' }, { name: 'MiniMax-M2.7' }], }, }, }, diff --git a/frontend/src/i18n/en.json b/frontend/src/i18n/en.json index f8a420250..4cdef7da8 100644 --- a/frontend/src/i18n/en.json +++ b/frontend/src/i18n/en.json @@ -32,12 +32,24 @@ "enter_variable_name": "Please Enter Variable Name", "enter_variable_value": "Please Enter Variable Value" }, + "authorized_space": { + "authorized_space": "Authorized Space", + "authorized_space_list": "Authorized Space List", + "select_space": "Select Space", + "modify_authorized_space": "Modify Authorized Space", + "workspaces_authorized": "Authorized {num} Workspaces", + "number_of_members": "Number of Members", + "delete_selected_workspaces": "Are you sure you want to delete the selected {msg} workspaces?", + "delete_workspace": "Are you sure you want to delete the workspace: {msg}?", + "no_workspace": "No Workspaces" + }, "sync": { "records": "Displaying {num} records out of {total}", "confirm_upload": "Confirm Upload", "field_details": "Field Details", "integration": "Platform integration needs to be enabled.", "the_existing_user": "If the user already exists, overwrite the existing user.", + "lazy_load": "Lazy Load", "sync_users": "Sync Users", "sync_wechat_users": "Sync WeChat Users", "sync_dingtalk_users": "Sync DingTalk Users", @@ -64,6 +76,7 @@ "context_record_count": "Context Record Count", "context_record_count_hint": "Number of user question rounds", "model_thinking_process": "Expand Model Thinking Process", + "hide_model_thinking_process": "Hide Model Thinking Process", "rows_of_data": "Limit 1000 Rows of Data", "third_party_platform_settings": "Authentication Settings", "by_third_party_platform": "Automatic User Creation", @@ -71,8 +84,8 @@ "platform_user_roles": "Third-Party Platform User Roles", "excessive_data_volume": "Disabling the 1000-row data limit may cause system lag due to excessive data volume.", "sqlbot_name": "Data Query Assistant Name", - "hide_sql": "Hide Show SQL Button", - "hide_log": "Hide Execution Log", + "show_sql": "Allow Viewing SQL Statements", + "show_log": "Show Execution Log", "prompt": "Prompt", "disabling_successfully": "Disabling Successfully", "closed_by_default": "In the Question Count window, control whether the model thinking process is expanded or closed by default.", @@ -178,6 +191,7 @@ "system_manage": "System Management", "update_success": "Update", "save_success": "Save Successful", + "operation_success": "Operation successful", "next": "Next", "save": "Save", "logout": "Logout", @@ -399,7 +413,8 @@ "address": "Address", "low_version": "Compatible with lower versions", "ssl": "Enable SSL", - "file_path": "File Path" + "file_path": "File Path", + "pool_size": "Connection Pool Size" }, "sync_fields": "Sync Fields", "sync_fields_success": "Sync fields successfully", @@ -571,6 +586,7 @@ "member_feng_yibudao": "Do you want to remove the member: {msg}?", "select_member": "Select member", "selected_2_people": "Selected: {msg} people", + "selected_number": "Selected: {msg} items", "clear": "Clear", "historical_dialogue": "No historical dialogue", "rename_a_workspace": "Rename a workspace", @@ -672,6 +688,9 @@ "application_description": "Application description", "cross_domain_settings": "Cross-domain settings", "enableCustomModel": "Use specified model", + "useModel": "Use model", + "defaultModel": "Default model", + "customModel": "Specified model", "third_party_address": "Please enter the embedded third party address,multiple items separated by semicolons", "set_to_private": "Set as private", "set_to_public": "Set as public", @@ -722,7 +741,7 @@ "display_settings": "Display Settings", "header_text_color": "Header Text Color", "app_logo": "App Logo", - "maximum_size_10mb": "Recommended size: 32 x 32, supports JPG, PNG, and SVG, maximum size: 10MB", + "maximum_size_10mb": "Recommended size: 32 x 32, supports JPG, PNG, maximum size: 10MB", "replace": "Replace", "default_icon_position": "Default Icon Position", "draggable_position": "Draggable Position", @@ -771,6 +790,8 @@ "no_data": "No Data", "loading_data": "Loading ...", "show_error_detail": "Show error info", + "thousands_separator_setting": "Thousands Separator Setting", + "thousands_separator_display": "Apply Thousands Separator Display", "log": { "GENERATE_SQL": "Generate SQL", "GENERATE_CHART": "Generate Chart Structure", @@ -835,11 +856,11 @@ "website_logo": "Website Logo", "tab": "Tab", "replace_image": "Replace Image", - "larger_than_200kb": "Logo displayed at the top of the website: Recommended size: 48 x 48 pixels, supports JPG, PNG, and SVG, and no larger than 200KB", + "larger_than_200kb": "Logo displayed at the top of the website: Recommended size: 48 x 48 pixels, supports JPG, PNG, and no larger than 200KB", "login_logo": "System Logo", - "larger_than_200kb_de": "Logo on the right side of the login page: Recommended size: 204 x 52 pixels, supports JPG, PNG, and SVG, and no larger than 200KB", + "larger_than_200kb_de": "Logo on the right side of the login page: Recommended size: 204 x 52 pixels, supports JPG, PNG, and no larger than 200KB", "login_background_image": "Login Background Image", - "larger_than_5mb": "Background image on the left: Recommended size: 576 x 900 for vector images, 1152 x 1800 for bitmap images, supports JPG, PNG, and SVG, and no larger than 5MB", + "larger_than_5mb": "Background image on the left: Recommended size: 576 x 900 for vector images, 1152 x 1800 for bitmap images, supports JPG, PNG, and no larger than 5MB", "website_name": "Website Name", "on_webpage_tabs": "Platform name displayed on webpage tabs", "welcome_message": "Display a welcome message", diff --git a/frontend/src/i18n/index.ts b/frontend/src/i18n/index.ts index 3b8078316..c06dfcd98 100644 --- a/frontend/src/i18n/index.ts +++ b/frontend/src/i18n/index.ts @@ -12,7 +12,31 @@ import { getBrowserLocale } from '@/utils/utils' const elementKoLocale = elementEnLocale const { wsCache } = useCache() +const isEmbeddedRoute = () => { + const hash = window.location.hash + if (!hash) return false + const hashPath = hash.substring(1).split('?')[0] + return ['/assistant', '/embeddedPage', '/embeddedCommon'].includes(hashPath) +} + +const getUrlLang = () => { + try { + const hash = window.location.hash + if (!hash) return null + const hashQuery = hash.substring(1).split('?')[1] + if (!hashQuery) return null + return new URLSearchParams(hashQuery).get('lang') + } catch { + return null + } +} + const getDefaultLocale = () => { + // 嵌入式页面的 URL lang 参数(第三种国际化渠道),优先级最高 + const urlLang = getUrlLang() + if (urlLang && isEmbeddedRoute()) { + return urlLang + } return wsCache.get('user.language') || getBrowserLocale() || 'zh-CN' } diff --git a/frontend/src/i18n/ko-KR.json b/frontend/src/i18n/ko-KR.json index 8c988192e..100efe262 100644 --- a/frontend/src/i18n/ko-KR.json +++ b/frontend/src/i18n/ko-KR.json @@ -32,12 +32,24 @@ "enter_variable_name": "변수 이름을 입력하세요", "enter_variable_value": "변수 값을 입력하세요" }, + "authorized_space": { + "authorized_space": "권한 있는 공간", + "authorized_space_list": "권한 있는 공간 목록", + "select_space": "공간 선택", + "modify_authorized_space": "권한 있는 공간 수정", + "workspaces_authorized": "{num}개의 권한 있는 작업 공간", + "number_of_members": "회원 수", + "delete_selected_workspaces": "선택한 {msg}개의 작업 공간을 삭제하시겠습니까?", + "delete_workspace": "작업 공간을 삭제하시겠습니까: {msg}?", + "no_workspace": "작업 공간 없음" + }, "sync": { "records": "{total}개 중 {num}개의 레코드 표시", "confirm_upload": "업로드 확인", "field_details": "필드 세부 정보", "integration": "플랫폼 통합을 활성화해야 합니다.", "the_existing_user": "해당 사용자가 이미 존재하는 경우 기존 사용자를 덮어씁니다.", + "lazy_load": "지연 로딩", "sync_users": "사용자 동기화", "sync_wechat_users": "위챗 사용자 동기화", "sync_dingtalk_users": "딩톡 사용자 동기화", @@ -64,6 +76,7 @@ "context_record_count": "컨텍스트 기록 수", "context_record_count_hint": "사용자 질문 라운드 수", "model_thinking_process": "모델 사고 프로세스 확장", + "hide_model_thinking_process": "모델 사고 과정 숨기기", "rows_of_data": "데이터 1,000행 제한", "third_party_platform_settings": "로그인 인증 설정", "by_third_party_platform": "자동 사용자 생성", @@ -71,8 +84,8 @@ "platform_user_roles": "타사 플랫폼 사용자 역할", "excessive_data_volume": "1,000행 데이터 제한을 비활성화하면 과도한 데이터 양으로 인해 시스템 지연이 발생할 수 있습니다.", "sqlbot_name": "데이터 질의 도우미 이름", - "hide_sql": "SQL 표시 버튼 숨기기", - "hide_log": "실행 로그 숨기기", + "show_sql": "SQL 문 보기 허용", + "show_log": "실행 로그 표시", "prompt": "프롬프트", "disabling_successfully": "비활성화 완료", "closed_by_default": "질문 수 창에서 모델 사고 프로세스를 기본적으로 확장할지 또는 닫을지 여부를 제어합니다.", @@ -178,6 +191,7 @@ "system_manage": "시스템 관리", "update_success": "업데이트 성공", "save_success": "저장 성공", + "operation_success": "작업 성공", "next": "다음 단계", "save": "저장", "logout": "로그아웃", @@ -399,7 +413,8 @@ "address": "주소", "low_version": "낮은 버전 호환", "ssl": "SSL 활성화", - "file_path": "파일 경로" + "file_path": "파일 경로", + "pool_size": "연결 풀 크기" }, "sync_fields": "동기화된 테이블 구조", "sync_fields_success": "테이블 구조 동기화 성공", @@ -571,6 +586,7 @@ "member_feng_yibudao": "멤버를 제거하시겠습니까: {msg}?", "select_member": "멤버 선택", "selected_2_people": "선택됨: {msg}명", + "selected_number": "선택됨: {msg}개", "clear": "지우기", "historical_dialogue": "과거 대화가 없습니다", "rename_a_workspace": "작업 공간 이름 바꾸기", @@ -672,6 +688,9 @@ "application_description": "애플리케이션 설명", "cross_domain_settings": "교차 도메인 설정", "enableCustomModel": "지정된 모델 사용", + "useModel": "사용 모델", + "defaultModel": "기본 모델", + "customModel": "지정 모델", "third_party_address": "임베디드할 제3자 주소를 입력하십시오, 여러 항목을 세미콜론으로 구분", "set_to_private": "비공개로 설정", "set_to_public": "공개로 설정", @@ -722,7 +741,7 @@ "display_settings": "표시 설정", "header_text_color": "헤더 텍스트 색상", "app_logo": "애플리케이션 로고", - "maximum_size_10mb": "권장 크기 32 x 32, JPG, PNG, SVG 지원, 크기 10MB 이하", + "maximum_size_10mb": "권장 크기 32 x 32, JPG, PNG 지원, 크기 10MB 이하", "replace": "교체", "default_icon_position": "아이콘 기본 위치", "draggable_position": "드래그 가능한 위치", @@ -771,6 +790,8 @@ "no_data": "데이터가 없습니다", "loading_data": "로딩 중 ...", "show_error_detail": "구체적인 정보 보기", + "thousands_separator_setting": "천 단위 구분 기호 설정", + "thousands_separator_display": "천 단위 구분 기호 표시 적용", "log": { "GENERATE_SQL": "SQL 생성", "GENERATE_CHART": "차트 구조 생성", @@ -835,11 +856,11 @@ "website_logo": "웹사이트 로고", "tab": "페이지 탭", "replace_image": "이미지 교체", - "larger_than_200kb": "상단 웹사이트에 표시되는 로고, 권장 크기 48 x 48, JPG, PNG, SVG 지원, 크기 200KB 이하", + "larger_than_200kb": "상단 웹사이트에 표시되는 로고, 권장 크기 48 x 48, JPG, PNG 지원, 크기 200KB 이하", "login_logo": "시스템 로고", - "larger_than_200kb_de": "로그인 페이지 오른쪽 로고, 권장 크기 204*52, JPG, PNG, SVG 지원, 크기 200KB 이하", + "larger_than_200kb_de": "로그인 페이지 오른쪽 로고, 권장 크기 204*52, JPG, PNG 지원, 크기 200KB 이하", "login_background_image": "로그인 배경 이미지", - "larger_than_5mb": "왼쪽 배경 이미지, 벡터 이미지 권장 크기 576*900, 비트맵 권장 크기 1152*1800; JPG, PNG, SVG 지원, 크기 5MB 이하", + "larger_than_5mb": "왼쪽 배경 이미지, 벡터 이미지 권장 크기 576*900, 비트맵 권장 크기 1152*1800; JPG, PNG 지원, 크기 5MB 이하", "website_name": "웹사이트 이름", "on_webpage_tabs": "웹페이지 탭에 표시되는 플랫폼 이름", "welcome_message": "환영 메시지 표시", diff --git a/frontend/src/i18n/zh-CN.json b/frontend/src/i18n/zh-CN.json index abbdd63c0..9f5cc9b24 100644 --- a/frontend/src/i18n/zh-CN.json +++ b/frontend/src/i18n/zh-CN.json @@ -32,12 +32,24 @@ "enter_variable_name": "请输入变量名称", "enter_variable_value": "请输入变量值" }, + "authorized_space": { + "authorized_space": "授权空间", + "authorized_space_list": "授权空间列表", + "select_space": "选择空间", + "modify_authorized_space": "修改授权空间", + "workspaces_authorized": "已授权 {num} 个工作空间", + "number_of_members": "成员数量", + "delete_selected_workspaces": "是否移除选中的 {msg} 个工作空间?", + "delete_workspace": "是否移除工作空间:{msg}?", + "no_workspace": "暂无工作空间" + }, "sync": { "records": "显示 {num} 条数据,共 {total} 条", "confirm_upload": "确定上传", "field_details": "字段详情", "integration": "需开启平台对接", "the_existing_user": "若用户已存在,覆盖旧用户", + "lazy_load": "懒加载", "sync_users": "同步用户", "sync_wechat_users": "同步企业微信用户", "sync_dingtalk_users": "同步钉钉用户", @@ -64,6 +76,7 @@ "context_record_count": "上下文记录数", "context_record_count_hint": "用户提问轮数", "model_thinking_process": "展开模型思考过程", + "hide_model_thinking_process": "隐藏模型思考过程", "rows_of_data": "限制 1000 行数据", "third_party_platform_settings": "登录认证设置", "by_third_party_platform": "自动创建用户", @@ -71,8 +84,8 @@ "platform_user_roles": "第三方平台用户角色", "excessive_data_volume": "关闭1000行的数据限制后,数据量过大,可能会造成系统卡顿", "sqlbot_name": "问数小助手名称", - "hide_sql": "隐藏展示SQL按钮", - "hide_log": "隐藏执行日志", + "show_sql": "允许查看SQL语句", + "show_log": "显示执行日志", "prompt": "提示", "disabling_successfully": "关闭成功", "closed_by_default": "在问数窗口中,控制模型思考过程默认展开或者关闭", @@ -178,6 +191,7 @@ "system_manage": "系统管理", "update_success": "更新成功", "save_success": "保存成功", + "operation_success": "操作成功", "next": "下一步", "save": "保存", "logout": "退出登录", @@ -399,7 +413,8 @@ "address": "地址", "low_version": "兼容低版本", "ssl": "启用 SSL", - "file_path": "文件路径" + "file_path": "文件路径", + "pool_size": "连接池大小" }, "sync_fields": "同步表结构", "sync_fields_success": "同步表结构成功", @@ -571,6 +586,7 @@ "member_feng_yibudao": "是否移除成员:{msg}?", "select_member": "选择成员", "selected_2_people": "已选:{msg} 人", + "selected_number": "已选:{msg} 个", "clear": "清空", "historical_dialogue": "暂无历史对话", "rename_a_workspace": "重命名工作空间", @@ -672,6 +688,9 @@ "application_description": "应用描述", "cross_domain_settings": "跨域设置", "enableCustomModel": "使用指定大模型", + "useModel": "使用模型", + "defaultModel": "默认模型", + "customModel": "指定模型", "third_party_address": "请输入嵌入的第三方地址,多个以分号分割", "set_to_private": "设为私有", "set_to_public": "设为公共", @@ -722,7 +741,7 @@ "display_settings": "显示设置", "header_text_color": "头部文本颜色", "app_logo": "应用 Logo", - "maximum_size_10mb": "建议尺寸 32 x 32,支持 JPG、PNG、SVG,大小不超过 10MB", + "maximum_size_10mb": "建议尺寸 32 x 32,支持 JPG、PNG,大小不超过 10MB", "replace": "替换", "default_icon_position": "图标默认位置", "draggable_position": "可拖拽位置", @@ -771,6 +790,8 @@ "no_data": "暂无数据", "loading_data": "加载中...", "show_error_detail": "查看具体信息", + "thousands_separator_setting": "千分位符设置", + "thousands_separator_display": "应用千分位符展示", "log": { "GENERATE_SQL": "生成 SQL", "GENERATE_CHART": "生成图表结构", @@ -835,11 +856,11 @@ "website_logo": "网站 Logo", "tab": "页签", "replace_image": "替换图片", - "larger_than_200kb": "顶部网站显示的 Logo,建议尺寸 48 x 48,支持 JPG、PNG、SVG,大小不超过 200KB", + "larger_than_200kb": "顶部网站显示的 Logo,建议尺寸 48 x 48,支持 JPG、PNG,大小不超过 200KB", "login_logo": "系统 Logo", - "larger_than_200kb_de": "登录页面右侧 Logo,建议尺寸 204*52,支持 JPG、PNG、SVG,大小不超过 200KB", + "larger_than_200kb_de": "登录页面右侧 Logo,建议尺寸 204*52,支持 JPG、PNG,大小不超过 200KB", "login_background_image": "登录背景图", - "larger_than_5mb": "左侧背景图,矢量图建议尺寸 576*900,位图建议尺寸 1152*1800;支持 JPG、PNG、SVG,大小不超过 5M", + "larger_than_5mb": "左侧背景图,矢量图建议尺寸 576*900,位图建议尺寸 1152*1800;支持 JPG、PNG,大小不超过 5M", "website_name": "网站名称", "on_webpage_tabs": "显示在网页 Tab 的平台名称", "welcome_message": "显示欢迎语", diff --git a/frontend/src/i18n/zh-TW.json b/frontend/src/i18n/zh-TW.json index 94668749f..b64649f9a 100644 --- a/frontend/src/i18n/zh-TW.json +++ b/frontend/src/i18n/zh-TW.json @@ -32,12 +32,24 @@ "enter_variable_name": "請輸入變數名稱", "enter_variable_value": "請輸入變數值" }, + "authorized_space": { + "authorized_space": "授權空間", + "authorized_space_list": "授權空間列表", + "select_space": "選擇空間", + "modify_authorized_space": "修改授權空間", + "workspaces_authorized": "已授權 {num} 個工作空間", + "number_of_members": "成員數量", + "delete_selected_workspaces": "是否移除選中的 {msg} 個工作空間?", + "delete_workspace": "是否移除工作空間:{msg}?", + "no_workspace": "暫無工作空間" + }, "sync": { "records": "顯示 {num} 筆資料,共 {total} 筆", "confirm_upload": "確定上傳", "field_details": "欄位詳情", "integration": "需開啟平台對接", "the_existing_user": "若使用者已存在,覆蓋舊使用者", + "lazy_load": "懶加載", "sync_users": "同步使用者", "sync_wechat_users": "同步企業微信使用者", "sync_dingtalk_users": "同步釘釘使用者", @@ -64,6 +76,7 @@ "context_record_count": "上下文記錄數", "context_record_count_hint": "使用者提問輪數", "model_thinking_process": "展開模型思考過程", + "hide_model_thinking_process": "隱藏模型思考過程", "rows_of_data": "限制 1000 列資料", "third_party_platform_settings": "登入認證設定", "by_third_party_platform": "自動建立使用者", @@ -71,8 +84,8 @@ "platform_user_roles": "第三方平台使用者角色", "excessive_data_volume": "關閉1000列的資料限制後,資料量過大,可能會造成系統卡頓", "sqlbot_name": "問數小助手名稱", - "hide_sql": "隱藏展示SQL按鈕", - "hide_log": "隱藏執行日誌", + "show_sql": "允許查看SQL語句", + "show_log": "顯示執行日誌", "prompt": "提示", "disabling_successfully": "關閉成功", "closed_by_default": "在問數視窗中,控制模型思考過程預設展開或者關閉", @@ -178,6 +191,7 @@ "system_manage": "系統管理", "update_success": "更新成功", "save_success": "儲存成功", + "operation_success": "操作成功", "next": "下一步", "save": "儲存", "logout": "登出", @@ -399,7 +413,8 @@ "address": "位址", "low_version": "相容低版本", "ssl": "啟用 SSL", - "file_path": "文件路徑" + "file_path": "文件路徑", + "pool_size": "連線池大小" }, "sync_fields": "同步表結構", "sync_fields_success": "同步表結構成功", @@ -571,6 +586,7 @@ "member_feng_yibudao": "是否移除成員:{msg}?", "select_member": "選擇成員", "selected_2_people": "已選:{msg} 人", + "selected_number": "已選:{msg} 個", "clear": "清空", "historical_dialogue": "暫無歷史對話", "rename_a_workspace": "重新命名工作區", @@ -672,6 +688,9 @@ "application_description": "應用描述", "cross_domain_settings": "跨網域設定", "enableCustomModel": "使用指定模型", + "useModel": "使用模型", + "defaultModel": "預設模型", + "customModel": "指定模型", "third_party_address": "請輸入嵌入的第三方位址,多個以分號分割", "set_to_private": "設為私有", "set_to_public": "設為公共", @@ -722,7 +741,7 @@ "display_settings": "顯示設定", "header_text_color": "頭部文字顏色", "app_logo": "應用 Logo", - "maximum_size_10mb": "建議尺寸 32 x 32,支援 JPG、PNG、SVG,大小不超過 10MB", + "maximum_size_10mb": "建議尺寸 32 x 32,支援 JPG、PNG,大小不超過 10MB", "replace": "取代", "default_icon_position": "圖示預設位置", "draggable_position": "可拖曳位置", @@ -771,6 +790,8 @@ "no_data": "暫無資料", "loading_data": "載入中...", "show_error_detail": "檢視具體資訊", + "thousands_separator_setting": "千分位符設定", + "thousands_separator_display": "套用千分位符顯示", "log": { "GENERATE_SQL": "產生 SQL", "GENERATE_CHART": "產生圖表結構", @@ -835,11 +856,11 @@ "website_logo": "網站 Logo", "tab": "頁籤", "replace_image": "取代圖片", - "larger_than_200kb": "頂部網站顯示的 Logo,建議尺寸 48 x 48,支援 JPG、PNG、SVG,大小不超過 200KB", + "larger_than_200kb": "頂部網站顯示的 Logo,建議尺寸 48 x 48,支援 JPG、PNG,大小不超過 200KB", "login_logo": "系統 Logo", - "larger_than_200kb_de": "登入頁面右側 Logo,建議尺寸 204*52,支援 JPG、PNG、SVG,大小不超過 200KB", + "larger_than_200kb_de": "登入頁面右側 Logo,建議尺寸 204*52,支援 JPG、PNG,大小不超過 200KB", "login_background_image": "登入背景圖", - "larger_than_5mb": "左側背景圖,向量圖建議尺寸 576*900,點陣圖建議尺寸 1152*1800;支援 JPG、PNG、SVG,大小不超過 5M", + "larger_than_5mb": "左側背景圖,向量圖建議尺寸 576*900,點陣圖建議尺寸 1152*1800;支援 JPG、PNG,大小不超過 5M", "website_name": "網站名稱", "on_webpage_tabs": "顯示在網頁 Tab 的平台名稱", "welcome_message": "顯示歡迎語", diff --git a/frontend/src/main.ts b/frontend/src/main.ts index f81e9fcbe..08c414c1a 100644 --- a/frontend/src/main.ts +++ b/frontend/src/main.ts @@ -1,3 +1,7 @@ +import 'core-js/features/object/has-own' +// @ts-ignore: css-has-pseudo/browser lacks official type definitions +import cssHasPseudo from 'css-has-pseudo/browser' + import { createApp } from 'vue' import { createPinia } from 'pinia' import './style.less' @@ -7,6 +11,26 @@ import { i18n } from './i18n' import VueDOMPurifyHTML from 'vue-dompurify-html' // import 'element-plus/dist/index.css' +cssHasPseudo(document) + +function supportsFlexGap() { + const flex = document.createElement('div') + flex.style.display = 'flex' + flex.style.flexDirection = 'column' + flex.style.rowGap = '1px' + + flex.appendChild(document.createElement('div')) + flex.appendChild(document.createElement('div')) + + document.body.appendChild(flex) + const isSupported = flex.scrollHeight === 1 + document.body.removeChild(flex) + + return isSupported +} + +document.documentElement.setAttribute('data-no-flex-gap', String(!supportsFlexGap())) + const app = createApp(App) const pinia = createPinia() diff --git a/frontend/src/stores/appearance.ts b/frontend/src/stores/appearance.ts index 2093f300a..ed368a7dc 100644 --- a/frontend/src/stores/appearance.ts +++ b/frontend/src/stores/appearance.ts @@ -315,7 +315,7 @@ const setLinkIcon = (linkWeb?: string) => { if (linkWeb) { link['href'] = baseUrl + linkWeb } else { - link['href'] = '/LOGO-fold.svg' + link['href'] = `${location.pathname}LOGO-fold.svg` } } } diff --git a/frontend/src/stores/assistant.ts b/frontend/src/stores/assistant.ts index 754d00a11..45474fcbd 100644 --- a/frontend/src/stores/assistant.ts +++ b/frontend/src/stores/assistant.ts @@ -1,10 +1,6 @@ import { defineStore } from 'pinia' import { store } from './index' import { chatApi, ChatInfo } from '@/api/chat' -import { useCache } from '@/utils/useCache' - -const { wsCache } = useCache() -const flagKey = 'sqlbit-assistant-flag' type Resolver = (value: T | PromiseLike) => void type Rejecter = (reason?: any) => void interface PendingRequest { @@ -16,7 +12,6 @@ interface AssistantState { id: string token: string assistant: boolean - flag: number type: number certificate: string online: boolean @@ -35,7 +30,6 @@ export const AssistantStore = defineStore('assistant', { id: '', token: '', assistant: false, - flag: 0, type: 0, certificate: '', online: false, @@ -61,9 +55,6 @@ export const AssistantStore = defineStore('assistant', { getAssistant(): boolean { return this.assistant }, - getFlag(): number { - return this.flag - }, getType(): number { return this.type }, @@ -176,14 +167,6 @@ export const AssistantStore = defineStore('assistant', { setAssistant(assistant: boolean) { this.assistant = assistant }, - setFlag(flag: number) { - if (wsCache.get(flagKey)) { - this.flag = wsCache.get(flagKey) - } else { - this.flag = flag - wsCache.set(flagKey, flag) - } - }, setPageEmbedded(embedded?: boolean) { this.pageEmbedded = !!embedded }, @@ -208,7 +191,6 @@ export const AssistantStore = defineStore('assistant', { return chat }, clear() { - wsCache.delete(flagKey) this.$reset() }, }, diff --git a/frontend/src/stores/chatConfig.ts b/frontend/src/stores/chatConfig.ts index 95ea95388..478547826 100644 --- a/frontend/src/stores/chatConfig.ts +++ b/frontend/src/stores/chatConfig.ts @@ -6,9 +6,10 @@ import { formatArg } from '@/utils/utils.ts' interface ChatConfig { sqlbot_name: string expand_thinking_block: boolean + hide_thinking_block: boolean limit_rows: boolean - hide_sql: boolean - hide_log: boolean + show_sql: boolean + show_log: boolean } export const chatConfigStore = defineStore('chatConfigStore', { @@ -16,9 +17,10 @@ export const chatConfigStore = defineStore('chatConfigStore', { return { sqlbot_name: 'SQLBot', expand_thinking_block: false, + hide_thinking_block: false, limit_rows: true, - hide_sql: false, - hide_log: false, + show_sql: true, + show_log: true, } }, getters: { @@ -28,11 +30,14 @@ export const chatConfigStore = defineStore('chatConfigStore', { getExpandThinkingBlock(): boolean { return this.expand_thinking_block }, - getHideSQL(): boolean { - return this.hide_sql + getHideThinkingBlock(): boolean { + return this.hide_thinking_block }, - getHideLog(): boolean { - return this.hide_log + getShowSQL(): boolean { + return this.show_sql + }, + getShowLog(): boolean { + return this.show_log }, getLimitRows(): boolean { return this.limit_rows @@ -43,14 +48,17 @@ export const chatConfigStore = defineStore('chatConfigStore', { request.get('/system/parameter/chat').then((res: any) => { if (res) { res.forEach((item: any) => { + if (item.pkey === 'chat.hide_thinking_block') { + this.hide_thinking_block = formatArg(item.pval) + } if (item.pkey === 'chat.expand_thinking_block') { this.expand_thinking_block = formatArg(item.pval) } - if (item.pkey === 'chat.hide_sql') { - this.hide_sql = formatArg(item.pval) + if (item.pkey === 'chat.show_sql') { + this.show_sql = formatArg(item.pval) } - if (item.pkey === 'chat.hide_log') { - this.hide_log = formatArg(item.pval) + if (item.pkey === 'chat.show_log') { + this.show_log = formatArg(item.pval) } if (item.pkey === 'chat.limit_rows') { this.limit_rows = formatArg(item.pval) diff --git a/frontend/src/style.less b/frontend/src/style.less index fe7799bbb..0d1d9263b 100644 --- a/frontend/src/style.less +++ b/frontend/src/style.less @@ -1,3 +1,23 @@ +/* ========================================== + 老旧浏览器 Flex Gap 完美物理隔离补丁 + ========================================== */ + +/* 1. 当浏览器【不支持 flex gap】时,激活此特定命名空间 */ +:root[data-no-flex-gap='true'] .flex-gap-fallback { + display: flex; +} + +/* 2. 仅对标记了该 class 的容器下的【相邻子元素】水平方向加间距 */ +:root[data-no-flex-gap='true'] .flex-gap-fallback > * + * { + margin-left: var(--gap-size, 16px); /* 默认 16px,可通过变量动态修改 */ +} + +/* 3. 如果需要处理垂直排列(flex-direction: column) */ +:root[data-no-flex-gap='true'] .flex-gap-fallback.flex-col > * + * { + margin-left: 0; + margin-top: var(--gap-size, 16px); +} + :root { font-family: system-ui, Avenir, Helvetica, Arial, sans-serif; line-height: 1.5; @@ -25,11 +45,15 @@ --shadow: 0 1px 3px rgba(0, 0, 0, 0.1); --ed-border-radius-base: 6px !important; --ed-border-color: #d9dcdf !important; +} + +:root:root { --ed-color-primary-light-7: #d2f1e9 !important; --ed-border-color: #d9dcdf !important; --ed-disabled-border-color: #d9dcdf !important; --ed-border-color-light: #dee0e3 !important; --ed-border-color-lighter: #dee0e3 !important; + --ed-border-radius-base: 6px !important; } a { @@ -84,7 +108,7 @@ body { .list-item_primary { height: 40px; - border-radius: 4px; + border-radius: 6px; padding: 8px 12px; cursor: pointer; display: flex; @@ -226,7 +250,7 @@ strong { .ed-select__popper { padding: 0 4px !important; .ed-select-dropdown__item { - border-radius: 4px; + border-radius: 6px; } .ed-select-dropdown__list { @@ -427,10 +451,9 @@ strong { } .ed-tree-node__content { - border-radius: 4px; + border-radius: 6px; } - .login-content { /* 针对 Webkit 浏览器的自动填充样式重置 */ input:-webkit-autofill, @@ -438,7 +461,7 @@ strong { input:-webkit-autofill:focus, input:-webkit-autofill:active { /* 1. 使用足够大的内阴影来覆盖背景色,把 #ffffff 替换成你输入框原本的背景色 */ - -webkit-box-shadow: 0 0 0px 1000px #f5f7fa inset !important; + -webkit-box-shadow: 0 0 0px 1000px #ffffff inset !important; /* 2. 由于常规的 color 属性也会失效,需要用这个属性修改文字颜色 */ -webkit-text-fill-color: #333333 !important; @@ -446,4 +469,12 @@ strong { /* 3. 保留光标的正常颜色 */ caret-color: #333333; } -} \ No newline at end of file +} + +.ed-popper.is-pure:has(.ed-select-dropdown) { + padding: 0 !important; + + .ed-select-dropdown { + padding: 0 4px !important; + } +} diff --git a/frontend/src/views/chat/ChatList.vue b/frontend/src/views/chat/ChatList.vue index c2a98aa3b..33799afc5 100644 --- a/frontend/src/views/chat/ChatList.vue +++ b/frontend/src/views/chat/ChatList.vue @@ -200,7 +200,7 @@ const handleConfirmPassword = () => {
+ + diff --git a/frontend/src/views/system/workspace/AuthorizedWorkspaceDialogForModelAdd.vue b/frontend/src/views/system/workspace/AuthorizedWorkspaceDialogForModelAdd.vue new file mode 100644 index 000000000..4f1fc4ba6 --- /dev/null +++ b/frontend/src/views/system/workspace/AuthorizedWorkspaceDialogForModelAdd.vue @@ -0,0 +1,349 @@ + + + + diff --git a/frontend/src/views/system/workspace/AuthorizedWorkspaceDraw.vue b/frontend/src/views/system/workspace/AuthorizedWorkspaceDraw.vue new file mode 100644 index 000000000..808eb16ac --- /dev/null +++ b/frontend/src/views/system/workspace/AuthorizedWorkspaceDraw.vue @@ -0,0 +1,521 @@ + + + + + diff --git a/frontend/src/views/system/workspace/index.vue b/frontend/src/views/system/workspace/index.vue index da3519cd4..303c888e6 100644 --- a/frontend/src/views/system/workspace/index.vue +++ b/frontend/src/views/system/workspace/index.vue @@ -701,7 +701,7 @@ const handleCurrentChange = (val: number) => { display: flex; align-items: center; padding-left: 8px; - border-radius: 4px; + border-radius: 6px; cursor: pointer; padding-right: 8px; margin-bottom: 2px; @@ -876,7 +876,7 @@ const handleCurrentChange = (val: number) => {