diff --git a/.drone.yml b/.drone.yml index 23bc7f938..ce53612f2 100644 --- a/.drone.yml +++ b/.drone.yml @@ -4,6 +4,7 @@ name: cicd # 定义流水线名称 clone: disable: true + steps: # 定义流水线执行步骤,这些步骤将顺序执行 - name: clone image: alpine/git @@ -125,22 +126,22 @@ steps: # 定义流水线执行步骤,这些步骤将顺序执行 - cd ./src/frontend/ - docker login -u $docker_user -p $docker_password $docker_registry - docker build -t $docker_repo:$version . - # - docker push $docker_repo:$version + - docker push $docker_repo:$version - # - name: ssh deploy - # image: appleboy/drone-ssh - # pull: if-not-exists - # settings: - # host: 192.168.106.116 - # username: root - # password: - # from_secret: sshpwd - # script: - # - echo =======找到目录======= - # - cd /opt/server/bisheng-test - # - echo =======直接启动======= - # - docker compose pull - # - docker compose up -d + - name: ssh deploy + image: appleboy/drone-ssh + pull: if-not-exists + settings: + host: 192.168.106.116 + username: root + password: + from_secret: sshpwd + script: + - echo =======找到目录======= + - cd /opt/server/bisheng-test + - echo =======直接启动======= + - docker compose pull + - docker compose up -d - name: notify-start # notify pull: if-not-exists @@ -176,7 +177,196 @@ steps: # 定义流水线执行步骤,这些步骤将顺序执行 trigger: branch: - release - - feat/workstation + event: + - push + +volumes: +- name: bisheng-cache + host: + path: /opt/drone/data/bisheng/ +- name: pro-cache + host: + path: /opt/drone/data/pro/ +- name: apt-cache + host: + path: /opt/drone/data/bisheng/apt/ +- name: socket + host: + path: /var/run/docker.sock + + +--- + +kind: pipeline # 定义对象类型,还有secret和signature两种类型 +type: docker # 定义流水线类型,还有kubernetes、exec、ssh等类型 +name: feat_cicd # 定义流水线名称 + +clone: + disable: true + +steps: # 定义流水线执行步骤,这些步骤将顺序执行 + - name: clone + image: alpine/git + pull: if-not-exists + environment: + http_proxy: + from_secret: PROXY + https_proxy: + from_secret: PROXY + commands: + - git config --global core.compression 0 + - git clone https://github.com/dataelement/bisheng.git . + - git checkout $DRONE_COMMIT + + - name: package # 流水线名称 + pull: if-not-exists + image: python:3.10-slim # 定义创建容器的Docker镜像 + volumes: # 将容器内目录挂载到宿主机,仓库需要开启Trusted设置 + - name: bisheng-cache + path: /app/build # 将应用打包好的Jar和执行脚本挂载出来 + environment: + RELEASE_VERSION: 99.99.90 + NEXUS_USER: + from_secret: NEXUS_USER + NEXUS_PASSWORD: + from_secret: NEXUS_PASSWORD + REPO: + from_secret: PY_NEXUS + commands: # 定义在Docker容器中执行的shell命令 + - pip config set global.index-url https://pypi.tuna.tsinghua.edu.cn/simple + - pip install Cython + - pip install wheel + - pip install twine + - cd ./src/bisheng-langchain + - python setup.py bdist_wheel + - twine upload --verbose -u $NEXUS_USER -p $NEXUS_PASSWORD --repository-url $REPO dist/*.whl + + - name: set poetry + pull: if-not-exists + image: golang + environment: + RELEASE_VERSION: 99.99.90 + NEXUS_PUBLIC: + from_secret: NEXUS_PUBLIC + NEXUS_PUBLIC_PASSWORD: + from_secret: NEXUS_PUBLIC_PASSWORD + REPO: + from_secret: PY_NEXUS + PROXY: + from_secret: APT-GET + volumes: # 将容器内目录挂载到宿主机,仓库需要开启Trusted设置 + - name: bisheng-cache + path: /app/build/ + commands: + - cd ./src/backend + - cp -r /app/build/nltk_data ./ + - echo $REPO + - REPO2=$(echo $REPO | sed 's/http:\\/\\///g') + - sed '/apt-get/ s|$| '"$PROXY"'|' Dockerfile + - sed -i 's/^bisheng_langchain.*/bisheng_langchain = "'$RELEASE_VERSION'"/g' pyproject.toml + - sed -i '6i\RUN pip config set global.index-url https://pypi.tuna.tsinghua.edu.cn/simple' Dockerfile + - sed -i '7i\RUN poetry source add --priority=supplemental foo http://'$NEXUS_PUBLIC':'$NEXUS_PUBLIC_PASSWORD'@'$REPO2'simple' Dockerfile + - sed -i '8i\RUN poetry source add --priority=primary qh https://pypi.tuna.tsinghua.edu.cn/simple' Dockerfile + - cat Dockerfile + + - name: build_docker + pull: if-not-exists + image: docker:24.0.6 + privileged: true + volumes: # 将容器内目录挂载到宿主机,仓库需要开启Trusted设置 + - name: apt-cache + path: /var/cache/apt/archives # 将应用打包好的Jar和执行脚本挂载出来 + - name: socket + path: /var/run/docker.sock + - name: pro-cache + path: /root/.local/share/pypoetry + environment: + http_proxy: + from_secret: PROXY + https_proxy: + from_secret: PROXY + no_proxy: 192.168.106.8 + version: ${DRONE_BRANCH} + docker_registry: http://192.168.106.8:6082 + docker_repo: 192.168.106.8:6082/dataelement/bisheng-backend + docker_user: + from_secret: NEXUS_USER + docker_password: + from_secret: NEXUS_PASSWORD + commands: + - echo "old tag is $version" + - version=$(echo $version | sed 's/\\//_/g') + - echo "build image tag is $version" + - cd ./src/backend/ + - docker login -u $docker_user -p $docker_password $docker_registry + - docker build -t $docker_repo:$version . + - docker push $docker_repo:$version + + - name: build_docker_frontend + pull: if-not-exists + image: docker:24.0.6 + privileged: true + volumes: # 将容器内目录挂载到宿主机,仓库需要开启Trusted设置 + - name: apt-cache + path: /var/cache/apt/archives # 将应用打包好的Jar和执行脚本挂载出来 + - name: socket + path: /var/run/docker.sock + environment: + http_proxy: + from_secret: PROXY + https_proxy: + from_secret: PROXY + no_proxy: 192.168.106.8 + version: ${DRONE_BRANCH} + docker_registry: http://192.168.106.8:6082 + docker_repo: 192.168.106.8:6082/dataelement/bisheng-frontend + docker_user: + from_secret: NEXUS_USER + docker_password: + from_secret: NEXUS_PASSWORD + commands: + - echo "old tag is $version" + - version=$(echo $version | sed 's/\\//_/g') + - echo "build image tag is $version" + - cd ./src/frontend/ + - docker login -u $docker_user -p $docker_password $docker_registry + - docker build -t $docker_repo:$version . + - docker push $docker_repo:$version + + - name: notify-start # notify + pull: if-not-exists + image: plugins/webhook + settings: + debug: true + urls: + from_secret: FEISHU_URL + content_type: application/json + template: | + { + "msg_type": "interactive", + "card": { + "type": "template", + "data": { + "template_id": "AAqkI9bnY5FUs", + "template_variable": { + "repo_name": "{{ repo.name }}", + "build_branch": "{{build.branch}}", + "build_author": "{{ DRONE_COMMIT_AUTHOR }}", + "link": "{{build.link}}", + "commit_msg": "{{ trim build.message }}", + "build_tag":"{{build.tag}}", + "build_start":"{{build.started}}", + "status": "{{ build.status }}" + } + } + } + } + when: # 成功 + status: + - success +trigger: + branch: + - add_some_branch_you_need event: - push diff --git a/src/backend/Dockerfile b/src/backend/Dockerfile index 6c2f03ba8..e244f8738 100644 --- a/src/backend/Dockerfile +++ b/src/backend/Dockerfile @@ -1,4 +1,4 @@ -FROM dataelement/bisheng-backend:base.v1 +FROM dataelement/bisheng-backend:base.v2 WORKDIR /app @@ -10,4 +10,8 @@ RUN poetry update --without dev # patch langchain-openai lib. remove this when upgrade langchain-openai RUN patch -p1 < /app/bisheng/patches/langchain_openai.patch /usr/local/lib/python3.10/site-packages/langchain_openai/chat_models/base.py +# patch fastapi-jwt-auth lib. remove this when remove fastapi-jwt-auth +# fix fastapi-jwt-auth not support pydantic:v2 +RUN patch -p1 < /app/bisheng/patches/fastapi_jwt_auth.patch /usr/local/lib/python3.10/site-packages/fastapi_jwt_auth/config.py + CMD ["sh entrypoint.sh"] diff --git a/src/backend/README.md b/src/backend/README.md index af4779ac9..8abb4ce91 100644 --- a/src/backend/README.md +++ b/src/backend/README.md @@ -1,3 +1,4 @@ # 毕昇后端代码 * Dockerfile 使用 poetry 进行 Python 依赖管理 + diff --git a/src/backend/bisheng/api/errcode/base.py b/src/backend/bisheng/api/errcode/base.py index e0ceb9d17..b55c0beb7 100644 --- a/src/backend/bisheng/api/errcode/base.py +++ b/src/backend/bisheng/api/errcode/base.py @@ -25,3 +25,8 @@ class UnAuthorizedError(BaseErrorCode): class NotFoundError(BaseErrorCode): Code: int = 404 Msg: str = '资源不存在' + + +class ServerError(BaseErrorCode): + Code: int = 500 + Msg: str = '服务器错误' diff --git a/src/backend/bisheng/api/errcode/knowledge.py b/src/backend/bisheng/api/errcode/knowledge.py index 3e4a8c37a..e3d026622 100644 --- a/src/backend/bisheng/api/errcode/knowledge.py +++ b/src/backend/bisheng/api/errcode/knowledge.py @@ -24,7 +24,7 @@ class KnowledgeSimilarError(BaseErrorCode): class KnowledgeQAError(BaseErrorCode): Code: int = 10930 - Msg: str = '该问题已被标注过' + Msg: str = '该问题已存在' class KnowledgeCPError(BaseErrorCode): diff --git a/src/backend/bisheng/api/services/assistant.py b/src/backend/bisheng/api/services/assistant.py index 58fb8cb1c..c2151fcf4 100644 --- a/src/backend/bisheng/api/services/assistant.py +++ b/src/backend/bisheng/api/services/assistant.py @@ -7,22 +7,24 @@ from loguru import logger from bisheng.api.errcode.assistant import (AssistantInitError, AssistantNameRepeatError, AssistantNotEditError, AssistantNotExistsError, ToolTypeRepeatError, - ToolTypeEmptyError, ToolTypeNotExistsError, ToolTypeIsPresetError) + ToolTypeNotExistsError, ToolTypeIsPresetError) from bisheng.api.errcode.base import UnAuthorizedError, NotFoundError from bisheng.api.services.assistant_agent import AssistantAgent from bisheng.api.services.assistant_base import AssistantUtils from bisheng.api.services.audit_log import AuditLogService from bisheng.api.services.base import BaseService from bisheng.api.services.llm import LLMService +from bisheng.api.services.tool import ToolServices from bisheng.api.services.user_service import UserPayload from bisheng.api.utils import get_request_ip from bisheng.api.v1.schemas import (AssistantInfo, AssistantSimpleInfo, AssistantUpdateReq, StreamData, UnifiedResponseModel, resp_200, resp_500) from bisheng.cache import InMemoryCache +from bisheng.database.constants import ToolPresetType from bisheng.database.models.assistant import (Assistant, AssistantDao, AssistantLinkDao, AssistantStatus) from bisheng.database.models.flow import Flow, FlowDao -from bisheng.database.models.gpts_tools import GptsToolsDao, GptsToolsRead, GptsToolsTypeRead, GptsTools +from bisheng.database.models.gpts_tools import GptsToolsDao, GptsToolsTypeRead, GptsTools from bisheng.database.models.group_resource import GroupResourceDao, GroupResource, ResourceTypeEnum from bisheng.database.models.knowledge import KnowledgeDao from bisheng.database.models.role_access import AccessType, RoleAccessDao @@ -107,7 +109,7 @@ class AssistantService(BaseService, AssistantUtils): @classmethod def get_assistant_info(cls, assistant_id: str, login_user: UserPayload): assistant = AssistantDao.get_one_assistant(assistant_id) - if not assistant: + if not assistant or assistant.is_delete: return AssistantNotExistsError.return_resp() # 检查是否有权限获取信息 if not login_user.access_check(assistant.user_id, assistant.id, AccessType.ASSISTANT_READ): @@ -368,11 +370,11 @@ class AssistantService(BaseService, AssistantUtils): return resp_200() @classmethod - def get_gpts_tools(cls, user: UserPayload, is_preset: Optional[bool] = None) -> List[GptsToolsTypeRead]: + def get_gpts_tools(cls, user: UserPayload, is_preset: Optional[int] = None) -> List[GptsToolsTypeRead]: """ 获取用户可见的工具列表 """ # 获取用户可见的工具类别 tool_type_ids_extra = [] - if not is_preset: + if is_preset != ToolPresetType.PRESET.value: # 获取自定义工具列表时,需要包含用户可用的工具列表 user_role = UserRoleDao.get_user_roles(user.user_id) if user_role: @@ -383,12 +385,13 @@ class AssistantService(BaseService, AssistantUtils): # 获取用户可见的所有工具列表 if is_preset is None: all_tool_type = GptsToolsDao.get_user_tool_type(user.user_id, tool_type_ids_extra) - elif is_preset: + elif is_preset == ToolPresetType.PRESET.value: # 获取预置工具列表 all_tool_type = GptsToolsDao.get_preset_tool_type() else: # 获取用户可见的自定义工具列表 - all_tool_type = GptsToolsDao.get_user_tool_type(user.user_id, tool_type_ids_extra, False) + all_tool_type = GptsToolsDao.get_user_tool_type(user.user_id, tool_type_ids_extra, False, + ToolPresetType(is_preset)) tool_type_id = [one.id for one in all_tool_type] res = [] tool_type_children = {} @@ -426,24 +429,29 @@ class AssistantService(BaseService, AssistantUtils): return tool_type @classmethod - def add_gpts_tools(cls, user: UserPayload, req: GptsToolsTypeRead) -> UnifiedResponseModel: + async def add_gpts_tools(cls, user: UserPayload, req: GptsToolsTypeRead) -> UnifiedResponseModel: """ 添加自定义工具 """ + # 尝试解析下openapi schema看下是否可以正常解析, 不能的话保存不允许保存 + tool_service = ToolServices() + if req.is_preset == ToolPresetType.API.value: + await tool_service.parse_openapi_schema('', req.openapi_schema) + elif req.is_preset == ToolPresetType.MCP.value: + await tool_service.parse_mcp_schema(req.openapi_schema) + req.id = None - if req.name.__len__() > 30 or req.name.__len__() == 0: - return resp_500(message="名字不符合规范:至少1个字符,不能超过30个字符") + if req.name.__len__() > 1000 or req.name.__len__() == 0: + return resp_500(message="名字不符合规范:至少1个字符,不能超过1000个字符") # 判断类别是否已存在 tool_type = GptsToolsDao.get_one_tool_type_by_name(user.user_id, req.name) if tool_type: return ToolTypeRepeatError.return_resp() - if len(req.children) == 0: - return ToolTypeEmptyError.return_resp() req.user_id = user.user_id for one in req.children: one.id = None one.user_id = user.user_id one.is_delete = 0 - one.is_preset = False + one.is_preset = req.is_preset # 添加工具类别和对应的 工具列表 res = GptsToolsDao.insert_tool_type(req) @@ -468,17 +476,22 @@ class AssistantService(BaseService, AssistantUtils): return True @classmethod - def update_gpts_tools(cls, user: UserPayload, req: GptsToolsTypeRead) -> UnifiedResponseModel: + async def update_gpts_tools(cls, user: UserPayload, req: GptsToolsTypeRead) -> UnifiedResponseModel: """ 更新工具类别,包括更新工具类别的名称和删除、新增工具类别的API """ + # 尝试解析下openapi schema看下是否可以正常解析, 不能的话保存不允许保存 + tool_service = ToolServices() + if req.is_preset == ToolPresetType.API.value: + await tool_service.parse_openapi_schema('', req.openapi_schema) + elif req.is_preset == ToolPresetType.MCP.value: + await tool_service.parse_mcp_schema(req.openapi_schema) + exist_tool_type = GptsToolsDao.get_one_tool_type(req.id) if not exist_tool_type: return ToolTypeNotExistsError.return_resp() - if len(req.children) == 0: - return ToolTypeEmptyError.return_resp() - if req.name.__len__() > 30 or req.name.__len__() == 0: - return resp_500(message="名字不符合规范:最少一个字符,不能超过30个字符") + if req.name.__len__() > 1000 or req.name.__len__() == 0: + return resp_500(message="名字不符合规范:至少1个字符,不能超过1000个字符") # 判断工具类别名称是否重复 tool_type = GptsToolsDao.get_one_tool_type_by_name(user.user_id, req.name) @@ -496,8 +509,8 @@ class AssistantService(BaseService, AssistantUtils): exist_tool_type.api_key = req.api_key exist_tool_type.auth_type = req.auth_type exist_tool_type.openapi_schema = req.openapi_schema - tool_extra = {"api_location":req.api_location,"parameter_name":req.parameter_name} - exist_tool_type.extra= json.dumps(tool_extra, ensure_ascii=False) + tool_extra = {"api_location": req.api_location, "parameter_name": req.parameter_name} + exist_tool_type.extra = json.dumps(tool_extra, ensure_ascii=False) children_map = {} for one in req.children: @@ -532,7 +545,7 @@ class AssistantService(BaseService, AssistantUtils): for one in children_map.values(): one.id = None one.user_id = user.user_id - one.is_preset = False + one.is_preset = exist_tool_type.is_preset one.is_delete = 0 add_children.append(one) @@ -549,7 +562,7 @@ class AssistantService(BaseService, AssistantUtils): exist_tool_type = GptsToolsDao.get_one_tool_type(tool_type_id) if not exist_tool_type: return resp_200() - if exist_tool_type.is_preset: + if exist_tool_type.is_preset == ToolPresetType.PRESET.value: return ToolTypeIsPresetError.return_resp() # 判断是否有更新权限 if not user.access_check(exist_tool_type.user_id, str(exist_tool_type.id), AccessType.GPTS_TOOL_WRITE): diff --git a/src/backend/bisheng/api/services/assistant_agent.py b/src/backend/bisheng/api/services/assistant_agent.py index 9bdb2d517..122353a51 100644 --- a/src/backend/bisheng/api/services/assistant_agent.py +++ b/src/backend/bisheng/api/services/assistant_agent.py @@ -6,18 +6,6 @@ import uuid from pathlib import Path from typing import Any, Dict, List -from bisheng.api.services.assistant_base import AssistantUtils -from bisheng.api.services.knowledge_imp import decide_vectorstores -from bisheng.api.services.llm import LLMService -from bisheng.api.services.openapi import OpenApiSchema -from bisheng.api.utils import build_flow_no_yield -from bisheng.api.v1.schemas import InputRequest -from bisheng.database.models.assistant import Assistant, AssistantLink, AssistantLinkDao -from bisheng.database.models.flow import FlowDao, FlowStatus -from bisheng.database.models.gpts_tools import GptsTools, GptsToolsDao, GptsToolsType -from bisheng.database.models.knowledge import Knowledge, KnowledgeDao -from bisheng.settings import settings -from bisheng.utils.embedding import decide_embeddings from bisheng_langchain.gpts.assistant import ConfigurableAssistant from bisheng_langchain.gpts.auto_optimization import (generate_breif_description, generate_opening_dialog, @@ -35,6 +23,23 @@ from langchain_core.utils.function_calling import format_tool_to_openai_tool from langchain_core.vectorstores import VectorStoreRetriever from loguru import logger +from bisheng.api.services.assistant_base import AssistantUtils +from bisheng.api.services.knowledge_imp import decide_vectorstores +from bisheng.api.services.llm import LLMService +from bisheng.api.services.openapi import OpenApiSchema +from bisheng.api.utils import build_flow_no_yield +from bisheng.api.v1.schemas import InputRequest +from bisheng.database.constants import ToolPresetType +from bisheng.database.models.assistant import Assistant, AssistantLink, AssistantLinkDao +from bisheng.database.models.flow import FlowDao, FlowStatus +from bisheng.database.models.gpts_tools import GptsTools, GptsToolsDao, GptsToolsType +from bisheng.database.models.knowledge import Knowledge, KnowledgeDao +from bisheng.mcp_manage.constant import McpClientType +from bisheng.mcp_manage.langchain.tool import McpTool +from bisheng.mcp_manage.manager import ClientManager +from bisheng.settings import settings +from bisheng.utils.embedding import decide_embeddings + class AssistantAgent(AssistantUtils): # cohere的模型需要的特殊prompt @@ -185,12 +190,43 @@ class AssistantAgent(AssistantUtils): tool_langchain.append(openapi_tool) return tool_langchain - async def init_personal_tools(self, tool_list: List[GptsTools], callbacks: Callbacks = None): + @staticmethod + def sync_init_mcp_tools(tool_list: List[GptsTools], callbacks: Callbacks = None): """ - 初始化自定义工具列表 + 初始化mcp工具列表 """ - return asyncio.get_running_loop().run_in_executor(None, self.sync_init_personal_tools, - tool_list, callbacks) + tool_type_ids = [one.type for one in tool_list] + all_tool_type = GptsToolsDao.get_all_tool_type(tool_type_ids) + all_tool_type = {one.id: one for one in all_tool_type} + tool_langchain = [] + for one in tool_list: + tool_type = all_tool_type.get(one.type) + input_schema = json.loads(one.extra) + mcp_client = ClientManager.sync_connect_mcp_from_json(tool_type.openapi_schema) + mcp_tool = McpTool.get_mcp_tool(name=one.tool_key, description=one.desc, mcp_client=mcp_client, + mcp_tool_name=one.name, arg_schema=input_schema['inputSchema'], + callbacks=callbacks) + tool_langchain.append(mcp_tool) + return tool_langchain + + @staticmethod + async def async_init_mcp_tools(tool_list: List[GptsTools], callbacks: Callbacks = None): + """ + 初始化mcp工具列表 + """ + tool_type_ids = [one.type for one in tool_list] + all_tool_type = GptsToolsDao.get_all_tool_type(tool_type_ids) + all_tool_type = {one.id: one for one in all_tool_type} + tool_langchain = [] + for one in tool_list: + tool_type = all_tool_type.get(one.type) + input_schema = json.loads(one.extra) + mcp_client = await ClientManager.connect_mcp_from_json(tool_type.openapi_schema) + mcp_tool = McpTool.get_mcp_tool(name=one.tool_key, description=one.desc, mcp_client=mcp_client, + mcp_tool_name=one.name, arg_schema=input_schema['inputSchema'], + callbacks=callbacks) + tool_langchain.append(mcp_tool) + return tool_langchain @staticmethod def sync_init_knowledge_tool(knowledge: Knowledge, @@ -231,22 +267,33 @@ class AssistantAgent(AssistantUtils): callbacks, self.knowledge_retriever) + @staticmethod + def parse_tools_type(tool_ids: List[int]) -> (list, list, list): + """ + 解析工具类型 + """ + tools_model: List[GptsTools] = GptsToolsDao.get_list_by_ids(tool_ids) + preset_tools = [] + personal_tools = [] + mcp_tools = [] + for one in tools_model: + if one.is_preset == ToolPresetType.PRESET.value: + preset_tools.append(one) + elif one.is_preset == ToolPresetType.API.value: + personal_tools.append(one) + else: + mcp_tools.append(one) + return preset_tools, personal_tools, mcp_tools + @staticmethod def init_tools_by_toolid( tool_ids: List[int], llm: BaseLanguageModel, callbacks: Callbacks = None, ): - """通过id初始化tool""" - tools_model: List[GptsTools] = GptsToolsDao.get_list_by_ids(tool_ids) - preset_tools = [] - personal_tools = [] - tools: List[BaseTool] = [] - for one in tools_model: - if one.is_preset: - preset_tools.append(one) - else: - personal_tools.append(one) + """ 通过id初始化tool !!! 只能在没有事件循环的线程中调用 """ + tools = [] + preset_tools, personal_tools, mcp_tools = AssistantAgent.parse_tools_type(tool_ids) if preset_tools: tool_langchain = AssistantAgent.sync_init_preset_tools(preset_tools, llm, callbacks) logger.info('act=build_preset_tools size={} return_tools={}', len(preset_tools), @@ -257,6 +304,34 @@ class AssistantAgent(AssistantUtils): logger.info('act=build_personal_tools size={} return_tools={}', len(personal_tools), len(tool_langchain)) tools += tool_langchain + if mcp_tools: + tool_langchain = AssistantAgent.sync_init_mcp_tools(mcp_tools, callbacks) + logger.info('act=build_mcp_tools size={} return_tools={}', len(mcp_tools), + len(tool_langchain)) + tools += tool_langchain + return tools + + @staticmethod + async def init_tools_by_tool_ids(tool_ids: List[int], + llm: BaseLanguageModel, + callbacks: Callbacks = None, ): + tools = [] + preset_tools, personal_tools, mcp_tools = AssistantAgent.parse_tools_type(tool_ids) + if preset_tools: + tool_langchain = AssistantAgent.sync_init_preset_tools(preset_tools, llm, callbacks) + logger.info('act=build_preset_tools size={} return_tools={}', len(preset_tools), + len(tool_langchain)) + tools += tool_langchain + if personal_tools: + tool_langchain = AssistantAgent.sync_init_personal_tools(personal_tools, callbacks) + logger.info('act=build_personal_tools size={} return_tools={}', len(personal_tools), + len(tool_langchain)) + tools += tool_langchain + if mcp_tools: + tools_langchain = await AssistantAgent.async_init_mcp_tools(mcp_tools, callbacks) + logger.info('act=build_mcp_tools size={} return_tools={}', len(mcp_tools), + len(tools_langchain)) + tools += tools_langchain return tools async def init_tools(self, callbacks: Callbacks = None): @@ -275,7 +350,7 @@ class AssistantAgent(AssistantUtils): else: flow_links.append(link) if tool_ids: - tools = self.init_tools_by_toolid(tool_ids, self.llm, callbacks) + tools = await self.init_tools_by_tool_ids(tool_ids, self.llm, callbacks) # flow + knowledge flow_data = FlowDao.get_flow_by_ids([link.flow_id for link in flow_links if link.flow_id]) diff --git a/src/backend/bisheng/api/services/audit_log.py b/src/backend/bisheng/api/services/audit_log.py index aa897636e..70677876e 100644 --- a/src/backend/bisheng/api/services/audit_log.py +++ b/src/backend/bisheng/api/services/audit_log.py @@ -102,12 +102,13 @@ class AuditLogService: flow_id, flow_info.name, ResourceTypeEnum.FLOW) @classmethod - def create_chat_workflow(cls, user: UserPayload, ip_address: str, flow_id: str): + def create_chat_workflow(cls, user: UserPayload, ip_address: str, flow_id: str, flow_info=None): """ 新建工作流会话的审计日志 """ logger.info(f"act=create_chat_workflow user={user.user_name} ip={ip_address} flow={flow_id}") - flow_info = FlowDao.get_flow_by_id(flow_id) + if not flow_info: + flow_info = FlowDao.get_flow_by_id(flow_id) cls._chat_log(user, ip_address, EventType.CREATE_CHAT, ObjectType.WORK_FLOW, flow_id, flow_info.name, ResourceTypeEnum.WORK_FLOW) diff --git a/src/backend/bisheng/api/services/chat_imp.py b/src/backend/bisheng/api/services/chat_imp.py index 11cb990d1..74ea3dd3e 100644 --- a/src/backend/bisheng/api/services/chat_imp.py +++ b/src/backend/bisheng/api/services/chat_imp.py @@ -53,7 +53,7 @@ async def clean_inactive_queues(queue: defaultdict, timeout_threshold: timedelta # 维护一个连接池 connection_pool = defaultdict(TimedQueue) -clean_inactive_queues(connection_pool, timedelta(minutes=5)) +# clean_inactive_queues(connection_pool, timedelta(minutes=5)) async def get_connection(uri, identifier): diff --git a/src/backend/bisheng/api/services/knowledge.py b/src/backend/bisheng/api/services/knowledge.py index 1ddd62b26..cfeba49ef 100644 --- a/src/backend/bisheng/api/services/knowledge.py +++ b/src/backend/bisheng/api/services/knowledge.py @@ -5,7 +5,7 @@ import os import time from typing import Any, Dict, List -from bisheng.api.errcode.base import NotFoundError, UnAuthorizedError +from bisheng.api.errcode.base import NotFoundError, UnAuthorizedError, ServerError from bisheng.api.errcode.knowledge import (KnowledgeChunkError, KnowledgeExistError, KnowledgeNoEmbeddingError) from bisheng.api.services.audit_log import AuditLogService @@ -840,7 +840,7 @@ class KnowledgeService(KnowledgeUtils): knowldge_dict = knowledge.model_dump() knowldge_dict.pop('id') knowldge_dict.pop('create_time') - knowldge_dict['update_time'] = '' + knowldge_dict.pop('update_time', None) knowldge_dict['user_id'] = login_user.user_id knowldge_dict['index_name'] = f'col_{int(time.time())}_{generate_uuid()[:8]}' knowldge_dict['name'] = f'{knowledge.name} 副本' @@ -857,3 +857,17 @@ class KnowledgeService(KnowledgeUtils): cls.create_knowledge_hook(request, login_user, target_knowlege) background_tasks.add_task(file_worker.file_copy_celery, params) return target_knowlege + + @classmethod + def judge_qa_knowledge_write(cls, login_user: UserPayload, qa_knowledge_id: int) -> Knowledge: + db_knowledge = KnowledgeDao.query_by_id(qa_knowledge_id) + # 查询当前知识库,是否有写入权限 + if not db_knowledge: + raise ServerError.http_exception(msg='当前知识库不可用,返回上级目录') + if not login_user.access_check(db_knowledge.user_id, str(qa_knowledge_id), + AccessType.KNOWLEDGE): + raise UnAuthorizedError.http_exception() + + if db_knowledge.type == KnowledgeTypeEnum.NORMAL.value: + raise ServerError.http_exception(msg='知识库为普通知识库') + return db_knowledge diff --git a/src/backend/bisheng/api/services/knowledge_imp.py b/src/backend/bisheng/api/services/knowledge_imp.py index d78801d78..bea816636 100644 --- a/src/backend/bisheng/api/services/knowledge_imp.py +++ b/src/backend/bisheng/api/services/knowledge_imp.py @@ -16,7 +16,7 @@ from bisheng.database.base import session_getter from bisheng.database.models.knowledge import Knowledge, KnowledgeDao from bisheng.database.models.knowledge_file import (KnowledgeFile, KnowledgeFileDao, KnowledgeFileStatus, ParseType, QAKnoweldgeDao, - QAKnowledge, QAKnowledgeUpsert) + QAKnowledge, QAKnowledgeUpsert, QAStatus) from bisheng.interface.embeddings.custom import FakeEmbedding from bisheng.interface.importing.utils import import_vectorstore from bisheng.interface.initialize.loading import instantiate_vectorstore @@ -86,8 +86,9 @@ class KnowledgeUtils: if not chunk.startswith('{'): return chunk.split(cls.chunk_split)[-1] - return chunk.split('\n')[-1].rstrip( - '}') + chunk = chunk.split('')[-1] + chunk = chunk.split('')[0] + return chunk @classmethod def save_preview_cache(cls, @@ -258,7 +259,7 @@ def decide_knowledge_llm() -> Any: return None # 获取llm对象 - return LLMService.get_bisheng_llm(model_id=knowledge_llm.extract_title_model_id) + return LLMService.get_bisheng_llm(model_id=knowledge_llm.extract_title_model_id, cache=False) def addEmbedding(collection_name: str, @@ -370,7 +371,7 @@ def add_file_embedding(vector_client, metadatas.append(val['metadata']) for index, one in enumerate(texts): if len(one) > 10000: - raise ValueError('分段结果超长,请尝试在自定义策略中使用更多切分符(例如 \n)进行切分') + raise ValueError('分段结果超长,请尝试在自定义策略中使用更多切分符(例如 \\n、。、\\.)进行切分') # 入库时 拼接文件名和文档摘要 texts[index] = KnowledgeUtils.aggregate_chunk_metadata(one, metadatas[index]) @@ -483,6 +484,8 @@ def read_chunk_text(input_file, file_name, separator: List[str], separator_rule: for one in documents: # 配置了相关llm的话,就对文档做总结 title = extract_title(llm, one.page_content) + # remove .* tag content + title = re.sub('.*', '', title, flags=re.S).strip() one.metadata['title'] = title logger.info('file_extract_title=success timecost={}', time.time() - t) @@ -669,14 +672,14 @@ def QA_save_knowledge(db_knowledge: Knowledge, QA: QAKnowledge): for vectore_client in vectore_client_list: vectore_client.add_texts(texts=[t.page_content for t in docs], metadatas=metadata) - QA.status = 1 + QA.status = QAStatus.ENABLED.value with session_getter() as session: session.add(QA) session.commit() session.refresh(QA) except Exception as e: logger.error(e) - setattr(QA, 'status', 0) + setattr(QA, 'status', QAStatus.FAILED.value) setattr(QA, 'remark', str(e)[:500]) with session_getter() as session: session.add(QA) @@ -721,12 +724,12 @@ def qa_status_change(qa_id: int, target_status: int): return db_knowledge = KnowledgeDao.query_by_id(qa_db.knowledge_id) - if target_status == 0: + if target_status == QAStatus.DISABLED.value: delete_vector_data(db_knowledge, [qa_id]) qa_db.status = target_status QAKnoweldgeDao.update(qa_db) else: - qa_db.status = target_status + qa_db.status = QAStatus.PROCESSING.value QAKnoweldgeDao.update(qa_db) QA_save_knowledge(db_knowledge, qa_db) return qa_db diff --git a/src/backend/bisheng/api/services/llm.py b/src/backend/bisheng/api/services/llm.py index 32825977f..fc3abea8a 100644 --- a/src/backend/bisheng/api/services/llm.py +++ b/src/backend/bisheng/api/services/llm.py @@ -1,6 +1,11 @@ import json from typing import List, Optional +from fastapi import Request +from langchain_core.embeddings import Embeddings +from langchain_core.language_models import BaseChatModel +from loguru import logger + from bisheng.api.errcode.base import NotFoundError from bisheng.api.errcode.llm import ServerExistError, ModelNameRepeatError, ServerAddError, ServerAddAllError from bisheng.api.services.user_service import UserPayload @@ -10,10 +15,6 @@ from bisheng.database.models.config import ConfigDao, ConfigKeyEnum, Config from bisheng.database.models.llm_server import LLMDao, LLMServer, LLMModel, LLMModelType from bisheng.interface.importing import import_by_type from bisheng.interface.initialize.loading import instantiate_llm, instantiate_embedding -from fastapi import Request -from langchain_core.embeddings import Embeddings -from langchain_core.language_models import BaseChatModel -from loguru import logger class LLMService: @@ -33,7 +34,7 @@ class LLMService: for one in llm_models: if one.server_id not in server_dicts: server_dicts[one.server_id] = [] - server_dicts[one.server_id].append(one.model_dump(exclude={'config'})) + server_dicts[one.server_id].append(LLMModelInfo(**one.model_dump(exclude={'config'}))) for one in ret: one.models = server_dicts.get(one.id, []) @@ -115,6 +116,8 @@ class LLMService: handle_types = [] for one in server.models: + # test model status + cls.test_model_status(one) if one.model_type in handle_types: continue handle_types.append(one.model_type) @@ -124,6 +127,19 @@ class LLMService: cls.set_default_model(request, login_user, model_info) return True + @classmethod + def test_model_status(cls, model: LLMModel | LLMModelInfo): + try: + if model.model_type == LLMModelType.LLM.value: + bisheng_model = cls.get_bisheng_llm(model_id=model.id, ignore_online=True, cache=False) + bisheng_model.invoke('hello') + elif model.model_type == LLMModelType.EMBEDDING.value: + bisheng_embed = cls.get_bisheng_embedding(model_id=model.id, ignore_online=True, cache=False) + bisheng_embed.embed_query('hello') + except Exception as e: + LLMDao.update_model_status(model.id, 1, str(e)) + logger.exception(f'test model status: {model.id} {model.model_name}') + @classmethod def set_default_model(cls, request: Request, login_user: UserPayload, model: LLMModel): """ 设置默认的模型配置 """ @@ -176,6 +192,11 @@ class LLMService: exist_server = LLMDao.get_server_by_id(server.id) if not exist_server: raise NotFoundError.http_exception() + + old_models = LLMDao.get_model_by_server_ids([exist_server.id]) + old_model_dict = { + one.id: one for one in old_models + } if exist_server.name != server.name: # 改名的话判断下是否已经存在 name_server = LLMDao.get_server_by_name(server.name) @@ -185,7 +206,7 @@ class LLMService: model_dict = {} for one in server.models: if one.model_name not in model_dict: - model_dict[one.model_name] = LLMModel(**one.dict()) + model_dict[one.model_name] = LLMModel(**one.model_dump()) # 说明是新增模型 if not one.id: model_dict[one.model_name].user_id = login_user.user_id @@ -201,8 +222,15 @@ class LLMService: exist_server.config = server.config db_server = LLMDao.update_server_with_models(exist_server, list(model_dict.values())) + new_server_info = cls.get_one_llm(request, login_user, db_server.id) - return cls.get_one_llm(request, login_user, db_server.id) + # 判断是否需要重新判断模型状态 + for one in new_server_info.models: + # 新增的模型,或者模型名字或者类型发生了变化 + if (one.id not in old_model_dict or old_model_dict[one.id].model_name != one.model_name + or old_model_dict[one.id].model_type != one.model_type): + cls.test_model_status(one) + return new_server_info @classmethod def update_model_online(cls, request: Request, login_user: UserPayload, model_id: int, diff --git a/src/backend/bisheng/api/services/role_group_service.py b/src/backend/bisheng/api/services/role_group_service.py index 332202916..ae895e82d 100644 --- a/src/backend/bisheng/api/services/role_group_service.py +++ b/src/backend/bisheng/api/services/role_group_service.py @@ -200,7 +200,7 @@ class RoleGroupService(): for one in group_ids: note += f'{group_dict.get(one, one)}、' note = note.rstrip('、') - AuditLogService.update_user(login_user, get_request_ip(request), user_id, group_dict.keys(), note) + AuditLogService.update_user(login_user, get_request_ip(request), user_id, list(group_dict.keys()), note) return None def get_user_groups_list(self, user_id: int) -> List[GroupRead]: diff --git a/src/backend/bisheng/api/services/tool.py b/src/backend/bisheng/api/services/tool.py new file mode 100644 index 000000000..50d5544fe --- /dev/null +++ b/src/backend/bisheng/api/services/tool.py @@ -0,0 +1,198 @@ +import json +from typing import Optional + +import yaml +from fastapi import Request +from loguru import logger +from pydantic import BaseModel, ConfigDict + +from bisheng.api.errcode.base import ServerError +from bisheng.api.services.openapi import OpenApiSchema +from bisheng.api.services.user_service import UserPayload +from bisheng.api.utils import get_url_content +from bisheng.database.constants import ToolPresetType +from bisheng.database.models.gpts_tools import GptsToolsDao, GptsTools, GptsToolsType, GptsToolsTypeRead +from bisheng.mcp_manage.manager import ClientManager +from bisheng.utils import md5_hash + + +class ToolServices(BaseModel): + """ 工具服务类 """ + model_config = ConfigDict(arbitrary_types_allowed=True) + + request: Optional[Request] = None + login_user: Optional[UserPayload] = None + + async def parse_openapi_schema(self, download_url: str, file_content: str) -> GptsToolsTypeRead: + if download_url: + try: + file_content = await get_url_content(download_url) + except Exception as e: + logger.exception(f'file {download_url} download error') + raise ServerError.http_exception(msg='url文件下载失败:' + str(e)) + if not file_content: + raise ServerError.http_exception(msg='schema内容不能为空') + # 根据文件内容是否以`{`开头判断用什么解析方式 + try: + if file_content.startswith('{'): + res = json.loads(file_content) + else: + res = yaml.safe_load(file_content) + except Exception as e: + logger.exception(f'openapi schema parse error {e}') + raise ServerError.http_exception(msg=f'openapi schema解析报错,请检查内容是否符合json或者yaml格式: {str(e)}') + + # 解析openapi schema转为助手工具的格式 + try: + schema = OpenApiSchema(res) + schema.parse_server() + if not schema.default_server.startswith(('http', 'https')): + raise ServerError.http_exception(msg=f'server中的url必须以http或者https开头: {schema.default_server}') + tool_type = GptsToolsTypeRead(name=schema.title, + description=schema.description, + is_preset=ToolPresetType.API.value, + server_host=schema.default_server, + openapi_schema=file_content, + api_location=schema.api_location, + parameter_name=schema.parameter_name, + auth_type=schema.auth_type, + auth_method=schema.auth_method, + children=[]) + # 解析获取所有的api + schema.parse_paths() + for one in schema.apis: + tool_type.children.append( + GptsTools( + name=one['operationId'], + desc=one['description'], + tool_key=md5_hash(one['operationId']), + is_preset=0, + is_delete=0, + api_params=one['parameters'], + extra=json.dumps(one, ensure_ascii=False), + )) + return tool_type + except Exception as e: + logger.exception(f'openapi schema parse error {e}') + raise ServerError.http_exception(msg='openapi schema解析失败:' + str(e)) + + async def parse_mcp_schema(self, file_content: str) -> GptsToolsTypeRead: + try: + result = json.loads(file_content) + mcp_servers = result['mcpServers'] + except Exception as e: + logger.exception(f'mcp tool schema parse error {e}') + raise ServerError.http_exception(msg=f'mcp工具配置解析失败,请检查内容是否符合mcp配置格式: {str(e)}') + tool_type = None + for key, value in mcp_servers.items(): + # 解析mcp服务配置 + tool_type = GptsToolsTypeRead(name=value.get('name', ''), + server_host=value.get('url', ''), + description=value.get('description', ''), + is_preset=ToolPresetType.MCP.value, + openapi_schema=file_content, + children=[]) + # 实例化mcp服务对象,获取工具列表 + client = await ClientManager.connect_mcp_from_json(result) + + tools = await client.list_tools() + + for one in tools: + tool_type.children.append(GptsTools( + name=one.name, + desc=one.description, + tool_key=md5_hash(one.name), + is_preset=ToolPresetType.MCP.value, + api_params=ToolServices.convert_input_schema(one.inputSchema), + extra=one.model_dump_json(), + )) + break + if tool_type is None: + raise ServerError.http_exception(msg='mcp服务配置解析失败,请检查配置里是否配置了mcpServers') + return tool_type + + async def refresh_all_mcp(self) -> str: + """ return mcp server error msg """ + # get user all mcp tool + tool_types = GptsToolsDao.get_user_tool_type(self.login_user.user_id, is_preset=ToolPresetType.MCP) + if not tool_types: + return '' + + tools = GptsToolsDao.get_list_by_type(tool_type_ids=[one.id for one in tool_types]) + tools_map = {} + for one in tools: + if one.type not in tools_map: + tools_map[one.type] = [] + tools_map[one.type].append(one) + error_msg = '' + for one in tool_types: + try: + await self.refresh_mcp_tools(one, tools_map.get(one.id, [])) + except Exception as e: + logger.exception(f'{one.name}刷新工具失败:') + error_msg += f'{one.name}工具获取失败,请重试\n' + return error_msg + + async def refresh_mcp_tools(self, tool_type: GptsToolsType, old_tools: list[GptsTools]): + """ refresh mcp tools """ + # 1. get all new tools + # 实例化mcp服务对象,获取工具列表 + client = await ClientManager.connect_mcp_from_json(tool_type.openapi_schema) + tools = await client.list_tools() + new_tools = {} + for one in tools: + tool_key = GptsToolsDao.get_tool_key(tool_type.id, md5_hash(one.name)) + new_tools[tool_key] = GptsTools( + name=one.name, + desc=one.description, + tool_key=tool_key, + is_preset=ToolPresetType.MCP.value, + api_params=self.convert_input_schema(one.inputSchema), + extra=one.model_dump_json(), + type=tool_type.id, + ) + + # 2. get need add or update or delete tool + need_delete_tool = [] # list[int] + need_update_tool = [] # list[GptsTools] + for one in old_tools: + if one.tool_key not in new_tools: + # 需要删除的工具 + logger.info(f'delete mcp tool: {one.name}') + need_delete_tool.append(one.id) + else: + logger.info(f'update mcp tool: {one.name}') + one.name = new_tools[one.tool_key].name + one.desc = new_tools[one.tool_key].desc + one.tool_key = new_tools[one.tool_key].tool_key + one.api_params = new_tools[one.tool_key].api_params + one.extra = new_tools[one.tool_key].extra + need_update_tool.append(one) + del new_tools[one.tool_key] + need_add_tool = list(new_tools.values()) + + # 3. update db + if need_delete_tool: + GptsToolsDao.delete_tool_by_ids(need_delete_tool) + if need_update_tool: + GptsToolsDao.update_tool_list(need_update_tool) + if need_add_tool: + GptsToolsDao.update_tool_list(need_add_tool) + + @classmethod + def convert_input_schema(cls, input_schema: dict): + """ 转换mcp工具的输入参数 为自定义工具的格式""" + required = input_schema.get('required', []) + properties = input_schema.get('properties', {}) + res = [] + for filed, field_info in properties.items(): + res.append({ + 'in': "query", + 'name': filed, + 'description': field_info.get('description'), + 'required': filed in required, + 'schema': { + 'type': field_info.get('type'), + } + }) + return res diff --git a/src/backend/bisheng/api/services/utils.py b/src/backend/bisheng/api/services/utils.py index e66bd8832..6ff69747c 100644 --- a/src/backend/bisheng/api/services/utils.py +++ b/src/backend/bisheng/api/services/utils.py @@ -1,6 +1,6 @@ from bisheng.template.field.base import TemplateField from bisheng.template.template.base import Template -from langchain.pydantic_v1 import BaseModel +from pydantic import BaseModel from langchain_core.language_models import BaseLanguageModel diff --git a/src/backend/bisheng/api/services/workflow.py b/src/backend/bisheng/api/services/workflow.py index 89c9e3452..c68828c16 100644 --- a/src/backend/bisheng/api/services/workflow.py +++ b/src/backend/bisheng/api/services/workflow.py @@ -258,8 +258,17 @@ class WorkFlowService(BaseService): key=event_input_schema.get('key'), type='text', required=True, + value='' ) ] + for one in event_input_schema.get('value', []): + tmp = WorkflowInputItem(**one) + if tmp.key == 'dialog_files_content': + tmp.type = 'dialog_file' + tmp.value = [] + elif tmp.key == 'dialog_file_accept': + tmp.type = 'dialog_file_accept' + input_schema.value.append(tmp) workflow_event.input_schema = input_schema return workflow_event diff --git a/src/backend/bisheng/api/services/workstation.py b/src/backend/bisheng/api/services/workstation.py index 5649b4e66..b254023fa 100644 --- a/src/backend/bisheng/api/services/workstation.py +++ b/src/backend/bisheng/api/services/workstation.py @@ -1,24 +1,27 @@ import asyncio import json from datetime import datetime -from typing import Optional +from typing import Optional, Any + +from pydantic import field_validator from bisheng.api.services import knowledge_imp, llm +from bisheng.api.services.base import BaseService from bisheng.api.services.knowledge import KnowledgeService from bisheng.api.services.user_service import UserPayload from bisheng.api.v1.schemas import KnowledgeFileOne, KnowledgeFileProcess, WorkstationConfig +from bisheng.database.constants import MessageCategory from bisheng.database.models.config import Config, ConfigDao, ConfigKeyEnum from bisheng.database.models.knowledge import KnowledgeCreate, KnowledgeDao, KnowledgeTypeEnum from bisheng.database.models.message import ChatMessage, ChatMessageDao from bisheng.database.models.session import MessageSession -from bisheng.restructure.assistants.services import MsgCategory from fastapi import BackgroundTasks, Request from langchain_core.messages import AIMessage, HumanMessage from loguru import logger from openai import BaseModel -class WorkStationService: +class WorkStationService(BaseService): @classmethod def update_config(cls, request: Request, login_user: UserPayload, data: WorkstationConfig) \ @@ -35,11 +38,15 @@ class WorkStationService: @classmethod def get_config(cls) -> WorkstationConfig | None: """ 获取评测功能的默认模型配置 """ - ret = {} config = ConfigDao.get_config(ConfigKeyEnum.WORKSTATION) if config: ret = json.loads(config.value) - return WorkstationConfig(**ret) + ret = WorkstationConfig(**ret) + if ret.assistantIcon and ret.assistantIcon.relative_path: + ret.assistantIcon.image = cls.get_logo_share_link(ret.assistantIcon.relative_path) + if ret.sidebarIcon and ret.sidebarIcon.relative_path: + ret.sidebarIcon.image = cls.get_logo_share_link(ret.sidebarIcon.relative_path) + return ret return None @classmethod @@ -124,9 +131,9 @@ class WorkStationService: # the user and the assistant were reversed, leading to incorrect question-and-answer sequences. extra = json.loads(one.extra) or {} content = extra['prompt'] if 'prompt' in extra else one.message - if one.category == MsgCategory.Question: + if one.category == MessageCategory.QUESTION.value: chat_history.append(HumanMessage(content=content)) - elif one.category == MsgCategory.Answer: + elif one.category == MessageCategory.ANSWER.value: chat_history.append(AIMessage(content=content)) logger.info(f'loaded {len(chat_history)} chat history for chat_id {chat_id}') return chat_history @@ -146,6 +153,20 @@ class WorkstationMessage(BaseModel): error: Optional[bool] = False unfinished: Optional[bool] = False + @field_validator('messageId', mode='before') + @classmethod + def convert_message_id(cls, value: Any) -> str: + if isinstance(value, str): + return value + return str(value) + + @field_validator('parentMessageId', mode='before') + @classmethod + def convert_parent_message_id(cls, value: Any) -> str: + if isinstance(value, str): + return value + return str(value) + @classmethod def from_chat_message(cls, message: ChatMessage): files = json.loads(message.files) if message.files else [] @@ -184,6 +205,13 @@ class WorkstationConversation(BaseModel): title=session.flow_name, ) + @field_validator('user', mode='before') + @classmethod + def convert_user(cls, v: Any) -> str: + if isinstance(v, str): + return v + return str(v) + class SSECallbackClient: diff --git a/src/backend/bisheng/api/v1/assistant.py b/src/backend/bisheng/api/v1/assistant.py index 24e7fd388..3aa3448f2 100644 --- a/src/backend/bisheng/api/v1/assistant.py +++ b/src/backend/bisheng/api/v1/assistant.py @@ -3,7 +3,6 @@ import json from typing import Dict, List, Optional import yaml -from bisheng.utils import generate_uuid from bisheng_langchain.gpts.tools.api_tools.openapi import OpenApiTools from fastapi import (APIRouter, Body, Depends, HTTPException, Query, Request, WebSocket, WebSocketException) @@ -13,23 +12,28 @@ from fastapi_jwt_auth import AuthJWT from bisheng.api.services.assistant import AssistantService from bisheng.api.services.openapi import OpenApiSchema +from bisheng.api.services.tool import ToolServices from bisheng.api.services.user_service import UserPayload, get_admin_user, get_login_user -from bisheng.api.utils import get_url_content -from bisheng.api.v1.schemas import (AssistantCreateReq, AssistantInfo, AssistantUpdateReq, +from bisheng.api.utils import get_url_content, md5_hash +from bisheng.api.v1.schemas import (AssistantCreateReq, AssistantUpdateReq, DeleteToolTypeReq, StreamData, TestToolReq, - UnifiedResponseModel, resp_200, resp_500) + resp_200, resp_500) from bisheng.cache.redis import redis_client from bisheng.chat.manager import ChatManager from bisheng.chat.types import WorkType +from bisheng.database.constants import ToolPresetType from bisheng.database.models.assistant import Assistant from bisheng.database.models.gpts_tools import GptsTools, GptsToolsTypeRead +from bisheng.mcp_manage.constant import McpClientType +from bisheng.mcp_manage.manager import ClientManager +from bisheng.utils import generate_uuid from bisheng.utils.logger import logger router = APIRouter(prefix='/assistant', tags=['Assistant']) chat_manager = ChatManager() -@router.get('', response_model=UnifiedResponseModel[List[AssistantInfo]]) +@router.get('') def get_assistant(*, name: str = Query(default=None, description='助手名称,模糊匹配, 包含描述的模糊匹配'), tag_id: int = Query(default=None, description='标签ID'), @@ -41,13 +45,13 @@ def get_assistant(*, # 获取某个助手的详细信息 -@router.get('/info/{assistant_id}', response_model=UnifiedResponseModel[AssistantInfo]) +@router.get('/info/{assistant_id}') def get_assistant_info(*, assistant_id: str, login_user: UserPayload = Depends(get_login_user)): """获取助手信息""" return AssistantService.get_assistant_info(assistant_id, login_user) -@router.post('/delete', response_model=UnifiedResponseModel) +@router.post('/delete') def delete_assistant(*, request: Request, assistant_id: str, @@ -56,7 +60,7 @@ def delete_assistant(*, return AssistantService.delete_assistant(request, login_user, assistant_id) -@router.post('', response_model=UnifiedResponseModel[AssistantInfo]) +@router.post('') async def create_assistant(*, request: Request, req: AssistantCreateReq, @@ -70,7 +74,7 @@ async def create_assistant(*, return resp_500(message=f'创建助手出错:{str(e)}') -@router.put('', response_model=UnifiedResponseModel[AssistantInfo]) +@router.put('') async def update_assistant(*, request: Request, req: AssistantUpdateReq, @@ -79,7 +83,7 @@ async def update_assistant(*, return await AssistantService.update_assistant(request, login_user, req) -@router.post('/status', response_model=UnifiedResponseModel) +@router.post('/status') async def update_status(*, request: Request, assistant_id: str = Body(description='助手唯一ID', alias='id'), @@ -129,7 +133,7 @@ async def auto_update_assistant(*, task_id: str = Query(description='优化任 # 更新助手的提示词 -@router.post('/prompt', response_model=UnifiedResponseModel) +@router.post('/prompt') async def update_prompt(*, assistant_id: str = Body(description='助手唯一ID', alias='id'), prompt: str = Body(description='用户使用的prompt'), @@ -137,7 +141,7 @@ async def update_prompt(*, return AssistantService.update_prompt(assistant_id, prompt, login_user) -@router.post('/flow', response_model=UnifiedResponseModel) +@router.post('/flow') async def update_flow_list(*, assistant_id: str = Body(description='助手唯一ID', alias='id'), flow_list: List[str] = Body(description='用户选择的技能列表'), @@ -145,7 +149,7 @@ async def update_flow_list(*, return AssistantService.update_flow_list(assistant_id, flow_list, login_user) -@router.post('/tool', response_model=UnifiedResponseModel) +@router.post('/tool') async def update_tool_list(*, assistant_id: str = Body(description='助手唯一ID', alias='id'), tool_list: List[int] = Body(description='用户选择的工具列表'), @@ -186,15 +190,17 @@ async def chat(*, await websocket.close(code=http_status.WS_1011_INTERNAL_ERROR, reason=message) -@router.get('/tool_list', response_model=UnifiedResponseModel) +@router.get('/tool_list') def get_tool_list(*, - is_preset: Optional[bool] = None, + is_preset: Optional[int | bool] = None, login_user: UserPayload = Depends(get_login_user)): """查询所有可见的tool 列表""" + if is_preset is not None and type(is_preset) == bool: + is_preset = ToolPresetType.PRESET.value if is_preset else ToolPresetType.API.value return resp_200(AssistantService.get_gpts_tools(login_user, is_preset)) -@router.post('/tool/config', response_model=UnifiedResponseModel) +@router.post('/tool/config') async def update_tool_config(*, login_user: UserPayload = Depends(get_admin_user), tool_id: int = Body(description='工具类别唯一ID'), @@ -204,94 +210,79 @@ async def update_tool_config(*, return resp_200(data=data) -@router.post('/tool_schema', response_model=UnifiedResponseModel) -async def get_tool_schema(*, +@router.post('/tool_schema') +async def get_tool_schema(request: Request, login_user: UserPayload = Depends(get_login_user), download_url: Optional[str] = Body(default=None, description='下载url不为空的话优先用下载url'), - file_content: Optional[str] = Body(default=None, description='上传的文件'), - login_user: UserPayload = Depends(get_login_user)): + file_content: Optional[str] = Body(default=None, description='上传的文件')): """ 下载或者解析openapi schema的内容 转为助手自定义工具的格式 """ - if download_url: - try: - file_content = await get_url_content(download_url) - except Exception as e: - logger.exception(f'file {download_url} download error') - return resp_500(message='url文件下载失败:' + str(e)) + services = ToolServices(request=request, login_user=login_user) + tool_type = await services.parse_openapi_schema(download_url, file_content) + return resp_200(data=tool_type) - if not file_content: - return resp_500(message='schema内容不能为空') - # 根据文件内容是否以`{`开头判断用什么解析方式 + +@router.post('/mcp/tool_schema') +async def get_mcp_tool_schema(request: Request, login_user: UserPayload = Depends(get_login_user), + file_content: Optional[str] = Body(default=None, embed=True, + description='mcp服务配置内容')): + """ 解析mcp的工具配置文件 """ + services = ToolServices(request=request, login_user=login_user) + tool_type = await services.parse_mcp_schema(file_content) + return resp_200(data=tool_type) + + +@router.post('/mcp/tool_test') +async def mcp_tool_run(login_user: UserPayload = Depends(get_login_user), + req: TestToolReq = None): + """ 测试mcp服务的工具 """ try: - if file_content.startswith('{'): - res = json.loads(file_content) - else: - res = yaml.safe_load(file_content) + # 实例化mcp服务对象,获取工具列表 + client = await ClientManager.connect_mcp_from_json(req.openapi_schema) + extra = json.loads(req.extra) + tool_name = extra.get('name') + resp = await client.call_tool(tool_name, req.request_params) + return resp_200(data=resp) except Exception as e: - logger.exception(f'openapi schema parse error {e}') - return resp_500(message=f'openapi schema解析报错,请检查内容是否符合json或者yaml格式: {str(e)}') - - # 解析openapi schema转为助手工具的格式 - try: - schema = OpenApiSchema(res) - schema.parse_server() - if not schema.default_server.startswith(('http', 'https')): - return resp_500(message=f'server中的url必须以http或者https开头: {schema.default_server}') - tool_type = GptsToolsTypeRead(name=schema.title, - description=schema.description, - is_preset=0, - is_delete=0, - server_host=schema.default_server, - openapi_schema=file_content, - api_location=schema.api_location, - parameter_name=schema.parameter_name, - auth_type=schema.auth_type, - auth_method=schema.auth_method, - children=[]) - # 解析获取所有的api - schema.parse_paths() - for one in schema.apis: - tool_type.children.append( - GptsTools( - name=one['operationId'], - desc=one['description'], - tool_key=hashlib.md5(one['operationId'].encode('utf-8')).hexdigest(), - is_preset=0, - is_delete=0, - api_params=one['parameters'], - extra=json.dumps(one, ensure_ascii=False), - )) - return resp_200(data=tool_type) - except Exception as e: - logger.exception(f'openapi schema parse error {e}') - return resp_500(message='openapi schema解析失败:' + str(e)) + logger.exception('mcp_tool_run error') + return resp_500(message=f'测试请求出错:{str(e)}') -@router.post('/tool_list', response_model=UnifiedResponseModel[GptsToolsTypeRead]) -def add_tool_type(*, +@router.post('/mcp/refresh') +async def refresh_all_mcp_tools(request: Request, login_user: UserPayload = Depends(get_login_user)): + """ 刷新用户当前所有的mcp工具列表 """ + services = ToolServices(request=request, login_user=login_user) + error_msg = await services.refresh_all_mcp() + if error_msg: + return resp_500(message=error_msg) + return resp_200(message='刷新成功') + + +@router.post('/tool_list') +async def add_tool_type(*, req: Dict = Body(default={}, description='openapi解析后的工具对象'), login_user: UserPayload = Depends(get_login_user)): """ 新增自定义tool """ req = GptsToolsTypeRead(**req) - return AssistantService.add_gpts_tools(login_user, req) + return await AssistantService.add_gpts_tools(login_user, req) -@router.put('/tool_list', response_model=UnifiedResponseModel[GptsToolsTypeRead]) -def update_tool_type(*, +@router.put('/tool_list') +async def update_tool_type(*, login_user: UserPayload = Depends(get_login_user), req: Dict = Body(default={}, description='通过openapi 解析后的内容,包含类别的唯一ID')): """ 更新自定义tool """ req = GptsToolsTypeRead(**req) - return AssistantService.update_gpts_tools(login_user, req) + return await AssistantService.update_gpts_tools(login_user, req) -@router.delete('/tool_list', response_model=UnifiedResponseModel) +@router.delete('/tool_list') def delete_tool_type(*, login_user: UserPayload = Depends(get_login_user), req: DeleteToolTypeReq): """ 删除自定义工具 """ return AssistantService.delete_gpts_tools(login_user, req.tool_type_id) -@router.post('/tool_test', response_model=UnifiedResponseModel) -async def test_tool_type(*, login_user: UserPayload = Depends(get_login_user), req: TestToolReq): +@router.post('/tool_test') +async def tool_run(*, login_user: UserPayload = Depends(get_login_user), req: TestToolReq): """ 测试自定义工具 """ extra = json.loads(req.extra) extra.update({'api_location': req.api_location, 'parameter_name': req.parameter_name}) diff --git a/src/backend/bisheng/api/v1/audit.py b/src/backend/bisheng/api/v1/audit.py index ab5be9a70..a52b15d47 100644 --- a/src/backend/bisheng/api/v1/audit.py +++ b/src/backend/bisheng/api/v1/audit.py @@ -10,7 +10,7 @@ from bisheng.api.v1.schemas import UnifiedResponseModel, resp_200 router = APIRouter(prefix='/audit', tags=['AuditLog']) -@router.get('', response_model=UnifiedResponseModel) +@router.get('') def get_audit_logs(*, group_ids: Optional[List[str]] = Query(default=[], description='分组id列表'), operator_ids: Optional[List[int]] = Query(default=[], description='操作人id列表'), @@ -27,7 +27,7 @@ def get_audit_logs(*, start_time, end_time, system_id, event_type, page, limit) -@router.get('/operators', response_model=UnifiedResponseModel) +@router.get('/operators') def get_all_operators(*, login_user: UserPayload = Depends(get_login_user)): """ 获取操作过组下资源的所有用户 @@ -35,7 +35,7 @@ def get_all_operators(*, login_user: UserPayload = Depends(get_login_user)): return AuditLogService.get_all_operators(login_user) -@router.get('/session', response_model=UnifiedResponseModel) +@router.get('/session') def get_session_list(login_user: UserPayload = Depends(get_login_user), flow_ids: Optional[List[str]] = Query(default=[], description='应用id列表'), user_ids: Optional[List[int]] = Query(default=[], description='用户id列表'), @@ -55,7 +55,7 @@ def get_session_list(login_user: UserPayload = Depends(get_login_user), }) -@router.get('/session/export', response_model=UnifiedResponseModel) +@router.get('/session/export') def export_session_messages(login_user: UserPayload = Depends(get_login_user), flow_ids: Optional[List[str]] = Query(default=[], description='应用id列表'), user_ids: Optional[List[int]] = Query(default=[], description='用户id列表'), @@ -73,7 +73,7 @@ def export_session_messages(login_user: UserPayload = Depends(get_login_user), }) -@router.get('/session/export/data', response_model=UnifiedResponseModel) +@router.get('/session/export/data') def get_session_messages(login_user: UserPayload = Depends(get_login_user), flow_ids: Optional[List[str]] = Query(default=[], description='应用id列表'), user_ids: Optional[List[int]] = Query(default=[], description='用户id列表'), diff --git a/src/backend/bisheng/api/v1/base.py b/src/backend/bisheng/api/v1/base.py index 6d75892e3..7d8e07343 100644 --- a/src/backend/bisheng/api/v1/base.py +++ b/src/backend/bisheng/api/v1/base.py @@ -1,7 +1,7 @@ from bisheng.interface.utils import extract_input_variables_from_prompt from bisheng.template.frontend_node.base import FrontendNode from langchain.prompts import PromptTemplate -from pydantic import BaseModel, validator +from pydantic import field_validator, BaseModel class CacheResponse(BaseModel): @@ -27,11 +27,13 @@ class CodeValidationResponse(BaseModel): imports: dict function: dict - @validator('imports') + @field_validator('imports') + @classmethod def validate_imports(cls, v): return v or {'errors': []} - @validator('function') + @field_validator('function') + @classmethod def validate_function(cls, v): return v or {'errors': []} diff --git a/src/backend/bisheng/api/v1/callback.py b/src/backend/bisheng/api/v1/callback.py index 196265932..f98fcdb2c 100644 --- a/src/backend/bisheng/api/v1/callback.py +++ b/src/backend/bisheng/api/v1/callback.py @@ -14,6 +14,7 @@ from langchain.schema import AgentFinish, LLMResult from langchain.schema.agent import AgentAction from langchain.schema.document import Document from langchain.schema.messages import BaseMessage +from langchain_core.messages import ToolMessage # https://github.com/hwchase17/chat-langchain/blob/master/callback.py @@ -120,7 +121,7 @@ class AsyncStreamingLLMCallbackHandler(AsyncCallbackHandler): observation_prefix = kwargs.get('observation_prefix', 'Tool output: ') # from langchain.docstore.document import Document # noqa # result = eval(output).get('result') - result = output + result = output if isinstance(output, str) else getattr(output, 'content', output) # Create a formatted message. intermediate_steps = f'{observation_prefix}{result[:100]}' @@ -332,7 +333,7 @@ class StreamingLLMCallbackHandler(BaseCallbackHandler): # from langchain.docstore.document import Document # noqa # result = eval(output).get('result') - result = output + result = output if isinstance(output, str) else getattr(output, 'content', output) # Create a formatted message. intermediate_steps = f'{observation_prefix}{result}' @@ -495,12 +496,13 @@ class AsyncGptsDebugCallbackHandler(AsyncGptsLLMCallbackHandler): extra=json.dumps({'run_id': kwargs.get('run_id').hex})) await self.websocket.send_json(resp.dict()) - async def on_tool_end(self, output: str, **kwargs: Any) -> Any: + async def on_tool_end(self, output: ToolMessage, **kwargs: Any) -> Any: """Run when tool ends running.""" + output = output.content logger.debug(f'on_tool_end output={output} kwargs={kwargs}') observation_prefix = kwargs.get('observation_prefix', 'Tool output: ') - result = output + result = output if isinstance(output, str) else getattr(output, 'content', output) # Create a formatted message. intermediate_steps = f'{observation_prefix}\n\n{result}' tool_name, tool_category = self.parse_tool_category(kwargs.get('name')) diff --git a/src/backend/bisheng/api/v1/chat.py b/src/backend/bisheng/api/v1/chat.py index 8cddd542f..866232fb0 100644 --- a/src/backend/bisheng/api/v1/chat.py +++ b/src/backend/bisheng/api/v1/chat.py @@ -38,6 +38,7 @@ from fastapi import (APIRouter, Body, HTTPException, Query, Request, WebSocket, from fastapi.params import Depends from fastapi.responses import StreamingResponse from fastapi_jwt_auth import AuthJWT +from sqlalchemy import func from sqlmodel import select router = APIRouter(tags=['Chat']) @@ -70,9 +71,7 @@ async def chat_completions(request: APIChatCompletion, Authorize: AuthJWT = Depe media_type='text/event-stream') -@router.get('/chat/app/list', - response_model=UnifiedResponseModel[PageList[AppChatList]], - status_code=200) +@router.get('/chat/app/list') def get_app_chat_list(*, keyword: Optional[str] = None, mark_user: Optional[str] = None, @@ -132,10 +131,7 @@ def get_app_chat_list(*, flow_ids = group_flow_ids # 获取会话列表 - res = MessageSessionDao.filter_session( - flow_ids=flow_ids, - user_ids=user_ids, - ) + res = MessageSessionDao.filter_session(flow_ids=flow_ids, user_ids=user_ids) total = len(res) # 查询会话的状态 @@ -172,11 +168,29 @@ def get_app_chat_list(*, continue result.append(tmp) - result = result[(page_num - 1) * page_size:page_num * page_size] + result = result[(page_num - 1) * page_size: page_num * page_size] return resp_200(PageList(list=result, total=total)) +@router.get('/chat/history') +def get_chatmessage(*, + chat_id: str, + flow_id: str, + id: Optional[str] = None, + page_size: Optional[int] = 20, + login_user: UserPayload = Depends(get_login_user)): + if not chat_id or not flow_id: + return {'code': 500, 'message': 'chat_id 和 flow_id 必传参数'} + where = select(ChatMessage).where(ChatMessage.flow_id == flow_id, + ChatMessage.chat_id == chat_id) + if id: + where = where.where(ChatMessage.id < int(id)) + with session_getter() as session: + db_message = session.exec(where.order_by(ChatMessage.id.desc()).limit(page_size)).all() + return resp_200(db_message) + + @router.post('/chat/conversation/rename') def rename(conversationId: str = Body(..., description='会话id', embed=True), name: str = Body(..., description='会话名称', embed=True), @@ -238,15 +252,14 @@ def del_chat_id(*, if session_chat.flow_type == FlowType.ASSISTANT.value: assistant_info = AssistantDao.get_one_assistant(session_chat.flow_id) if assistant_info: - AuditLogService.delete_chat_assistant(login_user, get_request_ip(request), - assistant_info) + AuditLogService.delete_chat_assistant(login_user, get_request_ip(request), assistant_info) else: # 判断下是助手还是技能, 写审计日志 flow_info = FlowDao.get_flow_by_id(session_chat.flow_id) if flow_info and flow_info.flow_type == FlowType.FLOW.value: - AuditLogService.delete_chat_flow(login_user, get_request_ip(request), flow_info) + AuditLogService.delete_chat_flow(login_user, get_request_ip(request), flow_info) elif flow_info: - AuditLogService.delete_chat_workflow(login_user, get_request_ip(request), flow_info) + AuditLogService.delete_chat_workflow(login_user, get_request_ip(request), flow_info) # 设置会话的删除状态 MessageSessionDao.delete_session(chat_id) @@ -262,15 +275,27 @@ def add_chat_messages(*, """ 添加一条完整问答记录, 安全检查写入使用 """ + logger.debug(f'gateway add_chat_messages {data}') flow_id = data.flow_id chat_id = data.chat_id if not chat_id or not flow_id: raise HTTPException(status_code=500, detail='chat_id 和 flow_id 必传参数') + save_human_message = data.human_message + flow_info = FlowDao.get_flow_by_id(flow_id) + if flow_info and flow_info.flow_type == FlowType.WORKFLOW.value: + # 工作流的输入,需要从输入里解析出来实际的输入内容 + try: + tmp_human_message = json.loads(data.human_message) + for node_id, node_input in tmp_human_message.items(): + save_human_message = node_input.get('message') + except: + save_human_message = data.human_message + human_message = ChatMessage(flow_id=flow_id, chat_id=chat_id, user_id=login_user.user_id, is_bot=False, - message=data.human_message, + message=save_human_message, sensitive_status=SensitiveStatus.VIOLATIONS.value, type='human', category='question') @@ -282,49 +307,50 @@ def add_chat_messages(*, sensitive_status=SensitiveStatus.PASS.value, type='bot', category='answer') - ChatMessageDao.insert_batch([human_message, bot_message]) + message_dbs = ChatMessageDao.insert_batch([human_message, bot_message]) # 更新会话的状态 MessageSessionDao.update_sensitive_status(chat_id, SensitiveStatus.VIOLATIONS) # 写审计日志, 判断是否是新建会话 - res = ChatMessageDao.get_messages_by_chat_id(chat_id=chat_id) - if len(res) <= 2: + session_info = MessageSessionDao.get_one(chat_id=chat_id) + if not session_info: # 新建会话 # 判断下是助手还是技能, 写审计日志 - flow_info = FlowDao.get_flow_by_id(flow_id) if flow_info: - MessageSessionDao.insert_one( - MessageSession( - chat_id=chat_id, - flow_id=flow_id, - flow_type=FlowType.FLOW.value, - flow_name=flow_info.name, - user_id=login_user.user_id, - sensitive_status=SensitiveStatus.VIOLATIONS.value, - )) - AuditLogService.create_chat_flow(login_user, get_request_ip(request), flow_id, - flow_info) + MessageSessionDao.insert_one(MessageSession( + chat_id=chat_id, + flow_id=flow_id, + flow_type=flow_info.flow_type, + flow_name=flow_info.name, + user_id=login_user.user_id, + sensitive_status=SensitiveStatus.VIOLATIONS.value, + )) + if flow_info.flow_type == FlowType.FLOW.value: + AuditLogService.create_chat_flow(login_user, get_request_ip(request), flow_id, flow_info) + elif flow_info.flow_type == FlowType.WORKFLOW.value: + AuditLogService.create_chat_workflow(login_user, get_request_ip(request), flow_id, flow_info) else: assistant_info = AssistantDao.get_one_assistant(flow_id) if assistant_info: - MessageSessionDao.insert_one( - MessageSession( - chat_id=chat_id, - flow_id=flow_id, - flow_type=FlowType.ASSISTANT.value, - flow_name=assistant_info.name, - user_id=login_user.user_id, - sensitive_status=SensitiveStatus.VIOLATIONS.value, - )) - AuditLogService.create_chat_assistant(login_user, get_request_ip(request), flow_id) + MessageSessionDao.insert_one(MessageSession( + chat_id=chat_id, + flow_id=flow_id, + flow_type=FlowType.ASSISTANT.value, + flow_name=assistant_info.name, + user_id=login_user.user_id, + sensitive_status=SensitiveStatus.VIOLATIONS.value, + )) + AuditLogService.create_chat_assistant(login_user, get_request_ip(request), + flow_id) - return resp_200(message='添加成功') + return resp_200(data=message_dbs, message='添加成功') @router.put('/chat/message/{message_id}', status_code=200) def update_chat_message(*, message_id: int, message: str = Body(embed=True), + category: str = Body(default=None, embed=True), login_user: UserPayload = Depends(get_login_user)): """ 更新一条消息的内容 安全检查使用""" logger.info( @@ -337,6 +363,8 @@ def update_chat_message(*, return resp_200(message='用户不一致') chat_message.message = message + if category: + chat_message.category = category chat_message.source = False chat_message.sensitive_status = SensitiveStatus.VIOLATIONS.value @@ -349,7 +377,6 @@ def update_chat_message(*, @router.delete('/chat/message/{message_id}', status_code=200) def del_message_id(*, message_id: str, login_user: UserPayload = Depends(get_login_user)): - # 删除一条消息,安全检查使用 ChatMessageDao.delete_by_message_id(login_user.user_id, message_id) return resp_200(message='删除成功') @@ -410,7 +437,7 @@ def comment_resp(*, data: ChatInput): return resp_200(message='操作成功') -@router.get('/chat/list', response_model=UnifiedResponseModel[List[ChatList]], status_code=200) +@router.get('/chat/list') def get_session_list(*, page: Optional[int] = 1, limit: Optional[int] = 10, @@ -431,33 +458,31 @@ def get_session_list(*, assistant_list = AssistantDao.get_assistants_by_ids(flow_ids) logo_map = {one.id: BaseService.get_logo_share_link(one.logo) for one in flow_list} logo_map.update({one.id: BaseService.get_logo_share_link(one.logo) for one in assistant_list}) - latest_messages = ChatMessageDao.get_latest_message_by_chat_ids( - chat_ids, exclude_category=WorkflowEventType.UserInput.value) + latest_messages = ChatMessageDao.get_latest_message_by_chat_ids(chat_ids, + exclude_category=WorkflowEventType.UserInput.value) latest_messages = {one.chat_id: one for one in latest_messages} return resp_200([ - ChatList(chat_id=one.chat_id, - flow_id=one.flow_id, - flow_name=one.flow_name, - flow_type=one.flow_type, - logo=logo_map.get(one.flow_id, ''), - latest_message=latest_messages.get(one.chat_id, None), - create_time=one.create_time, - update_time=one.update_time) for one in res + ChatList( + chat_id=one.chat_id, + flow_id=one.flow_id, + flow_name=one.flow_name, + flow_type=one.flow_type, + logo=logo_map.get(one.flow_id, ''), + latest_message=latest_messages.get(one.chat_id, None), + create_time=one.create_time, + update_time=one.update_time) for one in res ]) # 获取所有已上线的技能和助手 -@router.get('/chat/online', - response_model=UnifiedResponseModel[List[FlowGptsOnlineList]], - status_code=200) +@router.get('/chat/online') def get_online_chat(*, keyword: Optional[str] = None, tag_id: Optional[int] = None, page: Optional[int] = 1, limit: Optional[int] = 10, user: UserPayload = Depends(get_login_user)): - data, _ = WorkFlowService.get_all_flows(user, keyword, FlowStatus.ONLINE.value, tag_id, None, - page, limit) + data, _ = WorkFlowService.get_all_flows(user, keyword, FlowStatus.ONLINE.value, tag_id, None, page, limit) return resp_200(data=data) @@ -529,9 +554,7 @@ async def chat( await websocket.close(code=status.WS_1011_INTERNAL_ERROR, reason=messsage) -@router.post('/build/init/{flow_id}', - response_model=UnifiedResponseModel[InitResponse], - status_code=201) +@router.post('/build/init/{flow_id}') async def init_build(*, graph_data: dict, flow_id: str, diff --git a/src/backend/bisheng/api/v1/component.py b/src/backend/bisheng/api/v1/component.py index 6764e8b5c..a57db87b5 100644 --- a/src/backend/bisheng/api/v1/component.py +++ b/src/backend/bisheng/api/v1/component.py @@ -18,7 +18,7 @@ from bisheng.interface.custom.utils import build_custom_component_template router = APIRouter(prefix='/component', tags=['Component'], dependencies=[Depends(get_login_user)]) -@router.get('', response_model=UnifiedResponseModel[List[Component]]) +@router.get('') def get_all_components(*, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -26,7 +26,7 @@ def get_all_components(*, Authorize: AuthJWT = Depends()): return ComponentService.get_all_component(current_user) -@router.post('', response_model=UnifiedResponseModel[Component]) +@router.post('') def save_components(*, data: CreateComponentReq, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -38,7 +38,7 @@ def save_components(*, data: CreateComponentReq, Authorize: AuthJWT = Depends()) return ComponentService.save_component(component) -@router.patch('', response_model=UnifiedResponseModel[Component]) +@router.patch('') def update_component(*, data: CreateComponentReq, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -50,7 +50,7 @@ def update_component(*, data: CreateComponentReq, Authorize: AuthJWT = Depends() return ComponentService.update_component(component) -@router.delete('', response_model=UnifiedResponseModel[Component]) +@router.delete('') def delete_component(*, name: str = Body(..., embed=True, description='组件名'), Authorize: AuthJWT = Depends()): @@ -60,7 +60,7 @@ def delete_component(*, return ComponentService.delete_component(current_user.get('user_id'), name) -@router.post('/custom_component', response_model=UnifiedResponseModel[Component]) +@router.post('/custom_component') async def custom_component( raw_code: CustomComponentCode, Authorize: AuthJWT = Depends(), @@ -78,7 +78,7 @@ async def custom_component( return resp_200(data=built_frontend_node) -@router.post('/custom_component/reload', response_model=UnifiedResponseModel[Component]) +@router.post('/custom_component/reload') async def reload_custom_component(path: str, Authorize: AuthJWT = Depends()): from bisheng.interface.custom.utils import build_custom_component_template @@ -98,7 +98,7 @@ async def reload_custom_component(path: str, Authorize: AuthJWT = Depends()): return resp_500(message=str(exc)) -@router.post('/custom_component/update', response_model=UnifiedResponseModel[Component]) +@router.post('/custom_component/update') async def custom_component_update( raw_code: CustomComponentCode, Authorize: AuthJWT = Depends(), diff --git a/src/backend/bisheng/api/v1/endpoints.py b/src/backend/bisheng/api/v1/endpoints.py index c129fa856..fa2a94bbf 100644 --- a/src/backend/bisheng/api/v1/endpoints.py +++ b/src/backend/bisheng/api/v1/endpoints.py @@ -41,7 +41,8 @@ router = APIRouter(tags=['Base']) @router.get('/all') def get_all(): """获取所有参数""" - return resp_200(get_all_types_dict()) + all_types = get_all_types_dict() + return resp_200(all_types) @router.get('/env') @@ -137,8 +138,8 @@ async def process_flow_old( # For backwards compatibility we will keep the old endpoint -# @router.post('/predict/{flow_id}', response_model=UnifiedResponseModel[ProcessResponse]) -@router.post('/process', response_model=UnifiedResponseModel[ProcessResponse]) +# @router.post('/predict/{flow_id}') +@router.post('/process') async def process_flow( flow_id: Annotated[UUID, Body(embed=True)], inputs: Optional[dict] = None, @@ -320,9 +321,7 @@ async def upload_icon_workflow(request: Request, return resp_200(data=resp) -@router.post('/upload/{flow_id}', - response_model=UnifiedResponseModel[UploadFileResponse], - status_code=201) +@router.post('/upload/{flow_id}') async def create_upload_file(file: UploadFile, flow_id: str): # Cache file try: diff --git a/src/backend/bisheng/api/v1/evaluation.py b/src/backend/bisheng/api/v1/evaluation.py index 0544e6686..d19e48e10 100644 --- a/src/backend/bisheng/api/v1/evaluation.py +++ b/src/backend/bisheng/api/v1/evaluation.py @@ -15,7 +15,7 @@ from bisheng.cache.utils import convert_encoding_cchardet router = APIRouter(prefix='/evaluation', tags=['Skills'], dependencies=[Depends(get_login_user)]) -@router.get('', response_model=UnifiedResponseModel[List[Evaluation]]) +@router.get('') def get_evaluation(*, page: Optional[int] = Query(default=1, gt=0, description='页码'), limit: Optional[int] = Query(default=10, gt=0, description='每页条数'), @@ -27,7 +27,7 @@ def get_evaluation(*, return EvaluationService.get_evaluation(user, page, limit) -@router.post('', response_model=UnifiedResponseModel[EvaluationRead], status_code=201) +@router.post('') def create_evaluation(*, file: UploadFile, prompt: str = Form(), @@ -78,7 +78,7 @@ def delete_evaluation(*, evaluation_id: int, Authorize: AuthJWT = Depends()): return EvaluationService.delete_evaluation(evaluation_id, user_payload=user) -@router.get('/result/file/download', response_model=UnifiedResponseModel) +@router.get('/result/file/download') async def get_download_url(*, file_url: str, Authorize: AuthJWT = Depends()): diff --git a/src/backend/bisheng/api/v1/finetune.py b/src/backend/bisheng/api/v1/finetune.py index b2200078f..3fe4211df 100644 --- a/src/backend/bisheng/api/v1/finetune.py +++ b/src/backend/bisheng/api/v1/finetune.py @@ -23,7 +23,7 @@ router = APIRouter(prefix='/finetune', tags=['Finetune'], dependencies=[Depends( # create finetune job -@router.post('/job', response_model=UnifiedResponseModel[Finetune]) +@router.post('/job') async def create_job(*, finetune: FinetuneCreateReq, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -36,7 +36,7 @@ async def create_job(*, finetune: FinetuneCreateReq, Authorize: AuthJWT = Depend # 删除训练任务 -@router.delete('/job', response_model=UnifiedResponseModel) +@router.delete('/job') async def delete_job(*, job_id: str, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -45,7 +45,7 @@ async def delete_job(*, job_id: str, Authorize: AuthJWT = Depends()): # 中止训练任务 -@router.post('/job/cancel', response_model=UnifiedResponseModel) +@router.post('/job/cancel') async def cancel_job(*, job_id: str, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -54,7 +54,7 @@ async def cancel_job(*, job_id: str, Authorize: AuthJWT = Depends()): # 发布训练任务 -@router.post('/job/publish', response_model=UnifiedResponseModel) +@router.post('/job/publish') async def publish_job(*, job_id: str, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -62,7 +62,7 @@ async def publish_job(*, job_id: str, Authorize: AuthJWT = Depends()): return FinetuneService.publish_job(job_id, current_user) -@router.post('/job/publish/cancel', response_model=UnifiedResponseModel) +@router.post('/job/publish/cancel') async def cancel_publish_job(*, job_id: str, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -71,7 +71,7 @@ async def cancel_publish_job(*, job_id: str, Authorize: AuthJWT = Depends()): # 获取训练任务列表,支持分页 -@router.get('/job', response_model=UnifiedResponseModel[List[Finetune]]) +@router.get('/job') async def get_job(*, server: str = Query(default=None, description='关联的RT服务名字'), status: str = Query( @@ -96,14 +96,14 @@ async def get_job(*, # 获取任务最新详细信息,此接口会同步查询SFT-backend侧将任务状态更新到最新 -@router.get('/job/info', response_model=UnifiedResponseModel[Finetune]) +@router.get('/job/info') async def get_job_info(*, job_id: str, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() return FinetuneService.get_job_info(job_id) -@router.patch('/job/model', response_model=UnifiedResponseModel) +@router.patch('/job/model') async def update_job(*, req_data: FinetuneChangeModelName, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -111,7 +111,7 @@ async def update_job(*, req_data: FinetuneChangeModelName, Authorize: AuthJWT = return FinetuneService.change_job_model_name(req_data, current_user) -@router.post('/job/file', response_model=UnifiedResponseModel[List[PresetTrain]]) +@router.post('/job/file') async def upload_file(*, files: list[UploadFile] = File(description='训练文件列表'), Authorize: AuthJWT = Depends()): @@ -122,7 +122,7 @@ async def upload_file(*, return FinetuneFileService.upload_file(files, False, current_user) -@router.post('/job/file/preset', response_model=UnifiedResponseModel[List[PresetTrain]]) +@router.post('/job/file/preset') async def upload_preset_file(*, files: Optional[str] = Body(default=None, description='预置训练文件列表'), name: Optional[str] = Body(description='数据集名字'), @@ -154,7 +154,7 @@ async def upload_preset_file(*, # 获取预置训练文件列表 -@router.get('/job/file/preset', response_model=UnifiedResponseModel[List[PresetTrain]]) +@router.get('/job/file/preset') async def get_preset_file(*, page_size: Optional[int] = None, page_num: Optional[int] = None, @@ -166,7 +166,7 @@ async def get_preset_file(*, return resp_200(ret) -@router.delete('/job/file/preset', response_model=UnifiedResponseModel) +@router.delete('/job/file/preset') async def delete_preset_file(*, file_id: str, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() @@ -174,7 +174,7 @@ async def delete_preset_file(*, file_id: str, Authorize: AuthJWT = Depends()): return FinetuneFileService.delete_preset_file(file_id, current_user) -@router.get('/job/file/download', response_model=UnifiedResponseModel) +@router.get('/job/file/download') async def get_download_url(*, file_url: str, Authorize: AuthJWT = Depends()): Authorize.jwt_required() minio_client = MinioClient() @@ -182,14 +182,14 @@ async def get_download_url(*, file_url: str, Authorize: AuthJWT = Depends()): return resp_200(data={'url': download_url}) -@router.get('/server/filters', response_model=UnifiedResponseModel) +@router.get('/server/filters') async def get_server_filters(*, Authorize: AuthJWT = Depends()): Authorize.jwt_required() return FinetuneService.get_server_filters() -@router.get('/model/list', response_model=UnifiedResponseModel[List[ModelDeploy]]) +@router.get('/model/list') async def get_model_list(request: Request, login_user: UserPayload = Depends(get_login_user), server_id: int = Query(..., description='ft服务唯一ID')): @@ -198,7 +198,7 @@ async def get_model_list(request: Request, return resp_200(data=ret) -@router.get('/gpu', response_model=UnifiedResponseModel) +@router.get('/gpu') async def get_gpu_info(*, Authorize: AuthJWT = Depends()): # get login user Authorize.jwt_required() diff --git a/src/backend/bisheng/api/v1/flows.py b/src/backend/bisheng/api/v1/flows.py index c86e38529..c5a9d6987 100644 --- a/src/backend/bisheng/api/v1/flows.py +++ b/src/backend/bisheng/api/v1/flows.py @@ -137,13 +137,13 @@ def read_flows(*, raise HTTPException(status_code=500, detail=str(e)) from e -@router.get('/{flow_id}', response_model=UnifiedResponseModel[FlowReadWithStyle], status_code=200) +@router.get('/{flow_id}') def read_flow(*, flow_id: str, login_user: UserPayload = Depends(get_login_user)): """Read a flow.""" return FlowService.get_one_flow(login_user, flow_id) -@router.patch('/{flow_id}', response_model=UnifiedResponseModel[FlowRead], status_code=200) +@router.patch('/{flow_id}') async def update_flow(*, request: Request, flow_id: str, @@ -205,14 +205,14 @@ def delete_flow(*, return resp_200(message='删除成功') -@router.get('/download/', response_model=UnifiedResponseModel[FlowListRead], status_code=200) +@router.get('/download/') async def download_file(): """Download all flows as a file.""" flows = read_flows() return resp_200(FlowListRead(flows=flows)) -@router.post('/compare', response_model=UnifiedResponseModel, status_code=200) +@router.post('/compare') async def compare_flow_node(*, item: FlowCompareReq, Authorize: AuthJWT = Depends()): """ 技能多版本对比 """ Authorize.jwt_required() diff --git a/src/backend/bisheng/api/v1/knowledge.py b/src/backend/bisheng/api/v1/knowledge.py index 4e29c5bff..93a1ecc53 100644 --- a/src/backend/bisheng/api/v1/knowledge.py +++ b/src/backend/bisheng/api/v1/knowledge.py @@ -1,33 +1,36 @@ import json import urllib.parse -from typing import List, Optional +from datetime import datetime +from io import BytesIO +from typing import List, Optional, Any -from bisheng.api.errcode.base import UnAuthorizedError -from bisheng.api.errcode.knowledge import KnowledgeCPError, KnowledgeQAError -from bisheng.api.services import knowledge_imp -from bisheng.api.services.knowledge import KnowledgeService -from bisheng.api.services.knowledge_imp import add_qa -from bisheng.api.services.user_service import UserPayload, get_login_user -from bisheng.api.v1.schemas import (KnowledgeFileProcess, PreviewFileChunk, UnifiedResponseModel, - UpdatePreviewFileChunk, UploadFileResponse, resp_200, resp_500) -from bisheng.cache.utils import save_uploaded_file -from bisheng.database.base import session_getter -from bisheng.database.models.knowledge import (Knowledge, KnowledgeCreate, KnowledgeDao, - KnowledgeRead, KnowledgeTypeEnum, KnowledgeUpdate) -from bisheng.database.models.knowledge_file import (KnowledgeFileDao, KnowledgeFileStatus, - QAKnoweldgeDao, QAKnowledgeUpsert) -from bisheng.database.models.role_access import AccessType -from bisheng.database.models.user import UserDao -from bisheng.utils.logger import logger +import numpy as np +import pandas as pd from fastapi import (APIRouter, BackgroundTasks, Body, Depends, File, HTTPException, Query, Request, UploadFile) from fastapi.encoders import jsonable_encoder +from bisheng.api.errcode.base import UnAuthorizedError, ServerError +from bisheng.api.errcode.knowledge import KnowledgeCPError, KnowledgeQAError +from bisheng.api.services import knowledge_imp +from bisheng.api.services.knowledge import KnowledgeService +from bisheng.api.services.knowledge_imp import add_qa, QA_save_knowledge +from bisheng.api.services.user_service import UserPayload, get_login_user +from bisheng.api.v1.schemas import (KnowledgeFileProcess, PreviewFileChunk, UpdatePreviewFileChunk, UploadFileResponse, + resp_200, resp_500) +from bisheng.cache.utils import save_uploaded_file +from bisheng.database.models.knowledge import (KnowledgeCreate, KnowledgeDao, KnowledgeTypeEnum, KnowledgeUpdate) +from bisheng.database.models.knowledge_file import (KnowledgeFileDao, KnowledgeFileStatus, + QAKnoweldgeDao, QAKnowledgeUpsert, QAStatus) +from bisheng.database.models.role_access import AccessType +from bisheng.database.models.user import UserDao +from bisheng.utils.logger import logger + # build router router = APIRouter(prefix='/knowledge', tags=['Knowledge']) -@router.post('/upload', response_model=UnifiedResponseModel[UploadFileResponse], status_code=201) +@router.post('/upload') async def upload_file(*, file: UploadFile = File(...)): try: file_name = file.filename @@ -95,7 +98,7 @@ async def process_knowledge_file(*, return resp_200(res) -@router.post('/create', response_model=UnifiedResponseModel[KnowledgeRead], status_code=201) +@router.post('/create') def create_knowledge(*, request: Request, login_user: UserPayload = Depends(get_login_user), @@ -105,7 +108,7 @@ def create_knowledge(*, return resp_200(db_knowledge) -@router.post('/copy', response_model=UnifiedResponseModel[KnowledgeRead], status_code=201) +@router.post('/copy') async def copy_knowledge(*, request: Request, background_tasks: BackgroundTasks, @@ -123,7 +126,7 @@ async def copy_knowledge(*, ) if knowledge.state != 1 or knowledge_count > 0: return KnowledgeCPError.return_resp() - knowledge = KnowledgeService.copy_knowledge(request,background_tasks, login_user, knowledge) + knowledge = KnowledgeService.copy_knowledge(request, background_tasks, login_user, knowledge) return resp_200(knowledge) @@ -142,6 +145,7 @@ def get_knowledge(*, page_num, page_size) return resp_200(data={'data': res, 'total': total}) + @router.get('/info', status_code=200) def get_knowledge_info(*, request: Request, @@ -151,6 +155,7 @@ def get_knowledge_info(*, res = KnowledgeService.get_knowledge_info(request, login_user, knowledge_id) return resp_200(data=res) + @router.put('/', status_code=200) async def update_knowledge(*, request: Request, @@ -203,25 +208,11 @@ def get_QA_list(*, status: Optional[int] = None, login_user: UserPayload = Depends(get_login_user)): """ 获取知识库文件信息. """ - - # 查询当前知识库,是否有写入权限 - with session_getter() as session: - db_knowledge: Knowledge = session.get(Knowledge, qa_knowledge_id) - if not db_knowledge: - raise HTTPException(status_code=500, detail='当前知识库不可用,返回上级目录') - if not login_user.access_check(db_knowledge.user_id, str(qa_knowledge_id), - AccessType.KNOWLEDGE): - return UnAuthorizedError.return_resp() - - if db_knowledge.type == KnowledgeTypeEnum.NORMAL.value: - return HTTPException(status_code=500, detail='知识库为普通知识库') - - if keyword: - question = keyword + db_knowledge = KnowledgeService.judge_qa_knowledge_write(login_user, qa_knowledge_id) qa_list, total_count = knowledge_imp.list_qa_by_knowledge_id(qa_knowledge_id, page_size, page_num, question, answer, - status) + keyword, status) user_list = UserDao.get_user_by_ids([qa.user_id for qa in qa_list]) user_map = {user.user_id: user.user_name for user in user_list} data = [jsonable_encoder(qa) for qa in qa_list] @@ -232,12 +223,12 @@ def get_QA_list(*, return resp_200({ 'data': - data, + data, 'total': - total_count, + total_count, 'writeable': - login_user.access_check(db_knowledge.user_id, str(qa_knowledge_id), - AccessType.KNOWLEDGE_WRITE) + login_user.access_check(db_knowledge.user_id, str(qa_knowledge_id), + AccessType.KNOWLEDGE_WRITE) }) @@ -330,8 +321,9 @@ async def qa_add(*, QACreate: QAKnowledgeUpsert, if db_knowledge.type != KnowledgeTypeEnum.QA.value: raise HTTPException(status_code=404, detail='知识库类型错误') - db_q = QAKnoweldgeDao.get_qa_knowledge_by_name(QACreate.questions, QACreate.knowledge_id) - if db_q and not QACreate.id: + db_q = QAKnoweldgeDao.get_qa_knowledge_by_name(QACreate.questions, QACreate.knowledge_id, exclude_id=QACreate.id) + # create repeat question or update + if (db_q and not QACreate.id) or (db_q and QACreate.id and db_q.id != QACreate.id): raise KnowledgeQAError.http_exception() add_qa(db_knowledge=db_knowledge, data=QACreate) @@ -409,3 +401,199 @@ def qa_auto_question( """通过大模型自动生成问题""" questions = knowledge_imp.recommend_question(ori_question, number=number, answer=answer) return resp_200(data={'questions': questions}) + + +@router.get('/qa/export/template', status_code=200) +def get_export_url(): + data = [{"问题": "", "答案": "", "相似问题1": "", "相似问题2": ""}] + df = pd.DataFrame(data) + bio = BytesIO() + with pd.ExcelWriter(bio, engine="openpyxl") as writer: + df.to_excel(writer, sheet_name="Sheet1", index=False) + file_name = f"QA知识库导入模板.xlsx" + file_path = save_uploaded_file(bio, 'bisheng', file_name) + return resp_200({"url": file_path}) + + +@router.get('/qa/export/{qa_knowledge_id}', status_code=200) +def get_export_url(*, + qa_knowledge_id: int, + question: Optional[str] = None, + answer: Optional[str] = None, + keyword: Optional[str] = None, + status: Optional[int] = None, + max_lines: Optional[int] = 10000, + login_user: UserPayload = Depends(get_login_user)): + # 查询当前知识库,是否有写入权限 + db_knowledge = KnowledgeService.judge_qa_knowledge_write(login_user, qa_knowledge_id) + + if keyword: + question = keyword + + def get_qa_source(source): + '0: 未知 1: 手动;2: 审计, 3: api' + if int(source) == 1: + return "手动创建" + elif int(source) == 2: + return "审计创建" + elif int(source) == 3: + return "api创建" + return "未知" + + def get_status(statu): + if int(statu) == 1: + return "开启" + return "关闭" + + page_num = 1 + total_num = 0 + page_size = max_lines + user_list = UserDao.get_all_users() + user_map = {user.user_id: user.user_name for user in user_list} + file_list = [] + file_pr = datetime.now().strftime('%Y%m%d%H%M%S') + file_index = 1 + while True: + qa_list, total_count = knowledge_imp.list_qa_by_knowledge_id(qa_knowledge_id, page_size, + page_num, question, answer, + status) + + data = [jsonable_encoder(qa) for qa in qa_list] + qa_dict_list = [] + all_title = ["问题", "答案"] + for qa in data: + qa_dict_list.append({ + "问题": qa['questions'][0], + "答案": json.loads(qa['answers'])[0], + # "类型":get_qa_source(qa['source']), + # "创建时间":qa['create_time'], + # "更新时间":qa['update_time'], + # "创建者":user_map.get(qa['user_id'], qa['user_id']), + # "状态":get_status(qa['status']), + }) + for index, question in enumerate(qa['questions']): + if index == 0: + continue + key = f"相似问题{index}" + if key not in all_title: + all_title.append(key) + qa_dict_list[-1][key] = question + if len(qa_dict_list) != 0: + df = pd.DataFrame(qa_dict_list) + else: + df = pd.DataFrame([{"问题": "", "答案": "", "相似问题1": "", "相似问题2": ""}]) + df = df[all_title] + bio = BytesIO() + with pd.ExcelWriter(bio, engine="openpyxl") as writer: + df.to_excel(writer, sheet_name="Sheet1", index=False) + file_name = f"{file_pr}_{file_index}.xlsx" + file_index = file_index + 1 + file_path = save_uploaded_file(bio, 'bisheng', file_name) + file_list.append(file_path) + total_num += len(qa_list) + if len(qa_list) < page_size or total_num >= total_count: + break + + return resp_200({"file_list": file_list}) + + +def convert_excel_value(value: Any): + if value is None or value == "": + return '' + if str(value) == 'nan' or str(value) == 'null': + return '' + return str(value) + +@router.post('/qa/preview/{qa_knowledge_id}', status_code=200) +def post_import_file(*, + qa_knowledge_id: int, + file_url: str = Body(..., embed=True), + size: Optional[int] = Body(default=0, embed=True), + offset: Optional[int] = Body(default=0, embed=True), + login_user: UserPayload = Depends(get_login_user)): + df = pd.read_excel(file_url) + columns = df.columns.to_list() + if '答案' not in columns or '问题' not in columns: + raise HTTPException(status_code=500, detail='文件格式错误,没有 ‘问题’ 或 ‘答案’ 列') + data = df.T.to_dict().values() + insert_data = [] + for dd in data: + d = QAKnowledgeUpsert( + user_id=login_user.user_id, + knowledge_id=qa_knowledge_id, + answers=[convert_excel_value(dd['答案'])], + questions=[convert_excel_value(dd['问题'])], + source=4, + create_time=datetime.now(), + update_time=datetime.now()) + for key, value in dd.items(): + if key.startswith('相似问题') and convert_excel_value(value): + d.questions.append(convert_excel_value(value)) + insert_data.append(d) + try: + if size > 0 and offset >= 0: + if offset >= len(insert_data): + insert_data = [] + else: + insert_data = insert_data[offset:size] + except Exception as e: + raise HTTPException(status_code=500, detail=e) + return resp_200({"result": insert_data}) + + +@router.post('/qa/import/{qa_knowledge_id}', status_code=200) +def post_import_file(*, + qa_knowledge_id: int, + file_list: list[str] = Body(..., embed=True), + background_tasks: BackgroundTasks, + login_user: UserPayload = Depends(get_login_user)): + # 查询当前知识库,是否有写入权限 + db_knowledge = KnowledgeService.judge_qa_knowledge_write(login_user, qa_knowledge_id) + + insert_result = [] + error_result = [] + have_question = [] + for file_url in file_list: + df = pd.read_excel(file_url) + columns = df.columns.to_list() + if '答案' not in columns or '问题' not in columns: + insert_result.append(0) + continue + data = df.T.to_dict().values() + insert_data = [] + have_data = [] + all_questions = set() + for index, dd in enumerate(data): + tmp_questions = set() + dd_question = convert_excel_value(dd['问题']) + dd_answer = convert_excel_value(dd['答案']) + QACreate = QAKnowledgeUpsert( + user_id=login_user.user_id, + knowledge_id=qa_knowledge_id, + answers=[dd_answer], + questions=[dd_question], + source=4, + status=QAStatus.PROCESSING.value) + tmp_questions.add(QACreate.questions[0]) + for key, value in dd.items(): + if key.startswith('相似问题'): + if tmp_value := convert_excel_value(value): + if tmp_value not in tmp_questions: + QACreate.questions.append(tmp_value) + tmp_questions.add(tmp_value) + + db_q = QAKnoweldgeDao.get_qa_knowledge_by_name(QACreate.questions, QACreate.knowledge_id) + if (db_q and not QACreate.id) or len(tmp_questions & all_questions) > 0 or not dd_question or not dd_answer: + have_data.append(index) + else: + insert_data.append(QACreate) + all_questions = all_questions | tmp_questions + result = QAKnoweldgeDao.batch_insert_qa(insert_data) + + # async task add qa into milvus and es + for one in result: + background_tasks.add_task(QA_save_knowledge, db_knowledge, one) + + error_result.append(have_data) + + return resp_200({"errors": error_result}) diff --git a/src/backend/bisheng/api/v1/llm.py b/src/backend/bisheng/api/v1/llm.py index cf6546768..7272399c2 100644 --- a/src/backend/bisheng/api/v1/llm.py +++ b/src/backend/bisheng/api/v1/llm.py @@ -1,163 +1,116 @@ -from typing import List - from fastapi import APIRouter, Request, Depends, Body, Query from bisheng.api.services.llm import LLMService from bisheng.api.services.user_service import UserPayload, get_login_user, get_admin_user -from bisheng.api.v1.schemas import UnifiedResponseModel, LLMServerInfo, resp_200, KnowledgeLLMConfig, \ - AssistantLLMConfig, EvaluationLLMConfig, LLMServerCreateReq, LLMModelInfo +from bisheng.api.v1.schemas import resp_200, KnowledgeLLMConfig, \ + AssistantLLMConfig, EvaluationLLMConfig, LLMServerCreateReq router = APIRouter(prefix='/llm', tags=['LLM']) -@router.get('', response_model=UnifiedResponseModel[List[LLMServerInfo]]) -def get_all_llm( - request: Request, - login_user: UserPayload = Depends(get_login_user), -) -> UnifiedResponseModel[List[LLMServerInfo]]: +@router.get('') +def get_all_llm(request: Request, login_user: UserPayload = Depends(get_login_user)): ret = LLMService.get_all_llm(request, login_user) return resp_200(data=ret) -@router.post('', response_model=UnifiedResponseModel[LLMServerInfo]) -def add_llm_server( - request: Request, - login_user: UserPayload = Depends(get_admin_user), - server: LLMServerCreateReq = Body(..., description="服务提供方所有数据"), -) -> UnifiedResponseModel[LLMServerInfo]: +@router.post('') +def add_llm_server(request: Request, login_user: UserPayload = Depends(get_admin_user), + server: LLMServerCreateReq = Body(..., description="服务提供方所有数据")): ret = LLMService.add_llm_server(request, login_user, server) return resp_200(data=ret) -@router.delete('', response_model=UnifiedResponseModel) -def delete_llm_server( - request: Request, - login_user: UserPayload = Depends(get_admin_user), - server_id: int = Body(..., embed=True, description="服务提供方唯一ID"), -) -> UnifiedResponseModel: +@router.delete('') +def delete_llm_server(request: Request, login_user: UserPayload = Depends(get_admin_user), + server_id: int = Body(..., embed=True, description="服务提供方唯一ID")): LLMService.delete_llm_server(request, login_user, server_id) return resp_200() -@router.put('', response_model=UnifiedResponseModel[LLMServerInfo]) -def update_llm_server( - request: Request, - login_user: UserPayload = Depends(get_admin_user), - server: LLMServerCreateReq = Body(..., description="服务提供方所有数据"), -) -> UnifiedResponseModel[LLMServerInfo]: +@router.put('') +def update_llm_server(request: Request, login_user: UserPayload = Depends(get_admin_user), + server: LLMServerCreateReq = Body(..., description="服务提供方所有数据")): ret = LLMService.update_llm_server(request, login_user, server) return resp_200(data=ret) -@router.get('/info', response_model=UnifiedResponseModel[LLMServerInfo]) -def get_one_llm( - request: Request, - login_user: UserPayload = Depends(get_admin_user), - server_id: int = Query(..., description="服务提供方唯一ID"), -) -> UnifiedResponseModel[LLMServerInfo]: +@router.get('/info') +def get_one_llm(request: Request, login_user: UserPayload = Depends(get_admin_user), + server_id: int = Query(..., description="服务提供方唯一ID")): ret = LLMService.get_one_llm(request, login_user, server_id) return resp_200(data=ret) -@router.post('/online', response_model=UnifiedResponseModel[LLMModelInfo]) -def update_model_online( - request: Request, - login_user: UserPayload = Depends(get_admin_user), - model_id: int = Body(..., embed=True, description="模型的唯一ID"), - online: bool = Body(..., embed=True, description="是否上线"), -) -> UnifiedResponseModel[LLMModelInfo]: +@router.post('/online') +def update_model_online(request: Request, login_user: UserPayload = Depends(get_admin_user), + model_id: int = Body(..., embed=True, description="模型的唯一ID"), + online: bool = Body(..., embed=True, description="是否上线")): ret = LLMService.update_model_online(request, login_user, model_id, online) return resp_200(data=ret) -@router.get('/knowledge', response_model=UnifiedResponseModel[KnowledgeLLMConfig]) -def get_knowledge_llm( - request: Request, - login_user: UserPayload = Depends(get_login_user), -) -> UnifiedResponseModel[KnowledgeLLMConfig]: +@router.get('/knowledge') +def get_knowledge_llm(request: Request, login_user: UserPayload = Depends(get_login_user)): ret = LLMService.get_knowledge_llm() return resp_200(data=ret) -@router.post('/knowledge', response_model=UnifiedResponseModel[KnowledgeLLMConfig]) -def update_knowledge_llm( - request: Request, - login_user: UserPayload = Depends(get_admin_user), - data: KnowledgeLLMConfig = Body(..., description="知识库默认模型配置"), -) -> UnifiedResponseModel[KnowledgeLLMConfig]: +@router.post('/knowledge') +def update_knowledge_llm(request: Request, login_user: UserPayload = Depends(get_admin_user), + data: KnowledgeLLMConfig = Body(..., description="知识库默认模型配置")): """ 更新知识库相关的默认模型配置 """ ret = LLMService.update_knowledge_llm(request, login_user, data) return resp_200(data=ret) -@router.get('/assistant', response_model=UnifiedResponseModel[AssistantLLMConfig]) -def get_assistant_llm( - request: Request, - login_user: UserPayload = Depends(get_login_user), -) -> UnifiedResponseModel[AssistantLLMConfig]: +@router.get('/assistant') +def get_assistant_llm(request: Request, login_user: UserPayload = Depends(get_login_user)): """ 获取助手相关的模型配置 """ ret = LLMService.get_assistant_llm() return resp_200(data=ret) -@router.post('/assistant', response_model=UnifiedResponseModel[AssistantLLMConfig]) -def update_assistant_llm( - request: Request, - login_user: UserPayload = Depends(get_admin_user), - data: AssistantLLMConfig = Body(..., description="助手默认模型配置"), -) -> UnifiedResponseModel[AssistantLLMConfig]: +@router.post('/assistant') +def update_assistant_llm(request: Request, login_user: UserPayload = Depends(get_admin_user), + data: AssistantLLMConfig = Body(..., description="助手默认模型配置")): """ 更新助手相关的模型配置 """ ret = LLMService.update_assistant_llm(request, login_user, data) return resp_200(data=ret) -@router.get('/evaluation', response_model=UnifiedResponseModel[EvaluationLLMConfig]) -def get_evaluation_llm( - request: Request, - login_user: UserPayload = Depends(get_login_user), -) -> UnifiedResponseModel[EvaluationLLMConfig]: +@router.get('/evaluation') +def get_evaluation_llm(request: Request, login_user: UserPayload = Depends(get_login_user)): """ 获取评价相关的模型配置 """ ret = LLMService.get_evaluation_llm() return resp_200(data=ret) -@router.post('/evaluation', response_model=UnifiedResponseModel[EvaluationLLMConfig]) -def update_evaluation_llm( - request: Request, - login_user: UserPayload = Depends(get_admin_user), - data: EvaluationLLMConfig = Body(..., description="评价默认模型配置"), -) -> UnifiedResponseModel[EvaluationLLMConfig]: +@router.post('/evaluation') +def update_evaluation_llm(request: Request, login_user: UserPayload = Depends(get_admin_user), + data: EvaluationLLMConfig = Body(..., description="评价默认模型配置")): """ 更新评价相关的模型配置 """ ret = LLMService.update_evaluation_llm(request, login_user, data) return resp_200(data=ret) -@router.get('/workflow', response_model=UnifiedResponseModel[EvaluationLLMConfig]) -def get_workflow_llm( - request: Request, - login_user: UserPayload = Depends(get_login_user), -) -> UnifiedResponseModel[EvaluationLLMConfig]: +@router.get('/workflow') +def get_workflow_llm(request: Request, login_user: UserPayload = Depends(get_login_user)): """ 获取评价相关的模型配置 """ ret = LLMService.get_workflow_llm() return resp_200(data=ret) -@router.post('/workflow', response_model=UnifiedResponseModel[EvaluationLLMConfig]) -def update_workflow_llm( - request: Request, - login_user: UserPayload = Depends(get_admin_user), - data: EvaluationLLMConfig = Body(..., description="工作流默认模型配置"), -) -> UnifiedResponseModel[EvaluationLLMConfig]: +@router.post('/workflow') +def update_workflow_llm(request: Request, login_user: UserPayload = Depends(get_admin_user), + data: EvaluationLLMConfig = Body(..., description="工作流默认模型配置")): """ 更新评价相关的模型配置 """ ret = LLMService.update_workflow_llm(request, login_user, data) return resp_200(data=ret) -@router.get('/assistant/llm_list', response_model=UnifiedResponseModel[List[LLMServerInfo]]) -async def get_assistant_llm_list( - request: Request, - login_user: UserPayload = Depends(get_login_user), -) -> UnifiedResponseModel[List[LLMServerInfo]]: +@router.get('/assistant/llm_list') +async def get_assistant_llm_list(request: Request, login_user: UserPayload = Depends(get_login_user)): """ 获取助手可选的模型列表 """ ret = LLMService.get_assistant_llm_list(request, login_user) return resp_200(data=ret) diff --git a/src/backend/bisheng/api/v1/mark_task.py b/src/backend/bisheng/api/v1/mark_task.py index 4b9f3bd78..43aaf45ae 100644 --- a/src/backend/bisheng/api/v1/mark_task.py +++ b/src/backend/bisheng/api/v1/mark_task.py @@ -190,17 +190,20 @@ async def pre_or_next(chat_id: str, action: str, task_id: int, login_user: UserP if action == "prev": record = MarkRecordDao.get_prev_task(login_user.user_id, task_id) + top_queue = deque() + bottom_queue = deque() if record: - queue = deque() + queue = top_queue for r in record: if r.session_id == chat_id: + queue = bottom_queue continue queue.append(r) - if len(queue) == 0: + logger.info("top_queue={} bottom_queue={}", top_queue, bottom_queue) + if len(top_queue) == 0 and len(bottom_queue) == 0: return resp_200() - record = queue.pop() - logger.info("queue={} record={}", queue, record) + record = bottom_queue.popleft() if len(bottom_queue) else top_queue.popleft() chat = MessageSessionDao.get_one(record.session_id) result["chat_id"] = chat.chat_id result["flow_type"] = chat.flow_type @@ -219,6 +222,8 @@ async def pre_or_next(chat_id: str, action: str, task_id: int, login_user: UserP linked.append(m.chat_id) cur = linked.find(chat_id) + if not k_list: + return resp_200() logger.info("k_list={} cur={}", k_list, cur) diff --git a/src/backend/bisheng/api/v1/qa.py b/src/backend/bisheng/api/v1/qa.py index a6906d373..d3d2c8f93 100644 --- a/src/backend/bisheng/api/v1/qa.py +++ b/src/backend/bisheng/api/v1/qa.py @@ -16,7 +16,7 @@ from bisheng.utils.minio_client import MinioClient router = APIRouter(prefix='/qa', tags=['QA']) -@router.get('/keyword', response_model=UnifiedResponseModel[List[str]], status_code=200) +@router.get('/keyword') async def get_answer_keyword(message_id: int): # 获取命中的key conter = 3 diff --git a/src/backend/bisheng/api/v1/schema/base_schema.py b/src/backend/bisheng/api/v1/schema/base_schema.py index cb20f602c..573d31cc1 100644 --- a/src/backend/bisheng/api/v1/schema/base_schema.py +++ b/src/backend/bisheng/api/v1/schema/base_schema.py @@ -6,6 +6,6 @@ from pydantic import BaseModel DataT = TypeVar('DataT') -class PageList(Generic[DataT], BaseModel): +class PageList(BaseModel, Generic[DataT]): list: List[DataT] total: int diff --git a/src/backend/bisheng/api/v1/schema/chat_schema.py b/src/backend/bisheng/api/v1/schema/chat_schema.py index d5de1bc7a..241235ea5 100644 --- a/src/backend/bisheng/api/v1/schema/chat_schema.py +++ b/src/backend/bisheng/api/v1/schema/chat_schema.py @@ -2,7 +2,7 @@ import json from datetime import datetime from typing import Any, Dict, List, Optional -from pydantic import BaseModel +from pydantic import BaseModel, field_validator class AppChatList(BaseModel): @@ -13,15 +13,22 @@ class AppChatList(BaseModel): flow_id: str flow_type: int create_time: datetime - like_count: Optional[int] - dislike_count: Optional[int] - copied_count: Optional[int] - sensitive_status: Optional[int] # 敏感词审查状态 - user_groups: Optional[List[Any]] # 用户所属的分组 - mark_user: Optional[str] - mark_status: Optional[int] - mark_id: Optional[int] - messages: Optional[List[dict]] # 会话的所有消息列表数据 + like_count: Optional[int] = None + dislike_count: Optional[int] = None + copied_count: Optional[int] = None + sensitive_status: Optional[int] = None # 敏感词审查状态 + user_groups: Optional[List[Any]] = None # 用户所属的分组 + mark_user: Optional[str] = None + mark_status: Optional[int] = None + mark_id: Optional[int] = None + messages: Optional[List[dict]] = None # 会话的所有消息列表数据 + + @field_validator('user_name', mode='before') + @classmethod + def convert_user_name(cls, v: Any): + if not isinstance(v, str): + return str(v) + return v class APIAddQAParam(BaseModel): diff --git a/src/backend/bisheng/api/v1/schema/mark_schema.py b/src/backend/bisheng/api/v1/schema/mark_schema.py index 76394db99..589dee782 100644 --- a/src/backend/bisheng/api/v1/schema/mark_schema.py +++ b/src/backend/bisheng/api/v1/schema/mark_schema.py @@ -1,15 +1,26 @@ -from typing import List, Optional +from typing import List, Optional, Any -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, field_validator class MarkTaskCreate(BaseModel): - app_list: List[str] = Field(max_items=30) + app_list: List[str] = Field(max_length=30) user_list: List[str] + @field_validator('user_list', mode='before') + @classmethod + def convert_user_list(cls, v: Any): + ret = [] + for one in v: + if isinstance(one, str): + ret.append(one) + else: + ret.append(str(one)) + return ret + + class MarkData(BaseModel): session_id: str task_id: int status: int - flow_type:Optional[int] - + flow_type: Optional[int] = None diff --git a/src/backend/bisheng/api/v1/schema/workflow.py b/src/backend/bisheng/api/v1/schema/workflow.py index 94fdb2716..5d705b5db 100644 --- a/src/backend/bisheng/api/v1/schema/workflow.py +++ b/src/backend/bisheng/api/v1/schema/workflow.py @@ -1,7 +1,7 @@ from enum import Enum from typing import Optional, Any, List -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, field_validator class WorkflowEventType(Enum): @@ -32,11 +32,12 @@ class WorkflowOutputSchema(BaseModel): class WorkflowInputItem(BaseModel): key: str = Field(default=None, description='Unique key corresponding to user input') type: str = Field(default=None, description='The input type, select or dialog or file') - value: str = Field(default=None, description='The input default value') + value: Any = Field(default=None, description='The input default value') label: str = Field(default=None, description='The key label') multiple: bool = Field(default=False, description='The input is multi select') required: bool = Field(default=False, description='The input is required') options: Optional[Any] = Field(default=None, description='The select type options') + file_type: Optional[str] = Field(default=None, description='The allow upload file type') class WorkflowInputSchema(BaseModel): @@ -53,6 +54,13 @@ class WorkflowEvent(BaseModel): output_schema: Optional[WorkflowOutputSchema] = Field(default=None, description='The output schema') input_schema: Optional[WorkflowInputSchema] = Field(default=None, description='The input schema') + @field_validator('message_id', mode='before') + @classmethod + def validate_message_id(cls, v: Any) -> Optional[str]: + if isinstance(v, str) or v is None: + return v + return str(v) + class WorkflowStream(BaseModel): session_id: str = Field(default=None, description='The session id') diff --git a/src/backend/bisheng/api/v1/schemas.py b/src/backend/bisheng/api/v1/schemas.py index 685cd5b15..05800ddbe 100644 --- a/src/backend/bisheng/api/v1/schemas.py +++ b/src/backend/bisheng/api/v1/schemas.py @@ -12,7 +12,7 @@ from bisheng.database.models.message import ChatMessageRead from bisheng.database.models.tag import Tag from langchain.docstore.document import Document from orjson import orjson -from pydantic import BaseModel, Field, root_validator, validator +from pydantic import BaseModel, Field, model_validator, field_validator class CaptchaInput(BaseModel): @@ -55,7 +55,7 @@ class InputRequest(BaseModel): class TweaksRequest(BaseModel): - tweaks: Optional[Dict[str, Dict[str, str]]] = Field(default_factory=dict) + tweaks: Optional[Dict[str, Dict[str, str]]] = Field(default_factory=dict, description='List of dictionaries') class UpdateTemplateRequest(BaseModel): @@ -66,7 +66,7 @@ class UpdateTemplateRequest(BaseModel): DataT = TypeVar('DataT') -class UnifiedResponseModel(Generic[DataT], BaseModel): +class UnifiedResponseModel(BaseModel, Generic[DataT]): """统一响应模型""" status_code: int status_message: str @@ -90,7 +90,7 @@ def resp_500(code: int = 500, class ProcessResponse(BaseModel): """Process response schema.""" - result: Any + result: Any = None # task: Optional[TaskResponse] = None session_id: Optional[str] = None backend: Optional[str] = None @@ -121,7 +121,7 @@ class ChatList(BaseModel): create_time: datetime = None update_time: datetime = None flow_type: int = None - latest_message: ChatMessageRead = None + latest_message: Optional[ChatMessageRead] = None logo: Optional[str] = None @@ -143,34 +143,35 @@ class ChatMessage(BaseModel): message: Union[str, None, dict, list] = '' type: str = 'human' category: str = 'processing' # system processing answer tool - intermediate_steps: str = None + intermediate_steps: Optional[str] = None files: Optional[list] = [] - user_id: int = None - message_id: int | str = None - source: int = 0 - sender: str = None - receiver: dict = None + user_id: Optional[int] = None + message_id: Optional[int | str] = None + source: Optional[int] = 0 + sender: Optional[str] = None + receiver: Optional[dict] = None liked: int = 0 extra: Optional[str | dict] = '{}' - flow_id: str = None - chat_id: str = None + flow_id: Optional[str] = None + chat_id: Optional[str] = None class ChatResponse(ChatMessage): """Chat response schema.""" intermediate_steps: str = '' - is_bot: bool = True + is_bot: bool | int = True category: str = 'processing' - @validator('type') + @field_validator('type') + @classmethod def validate_message_type(cls, v): """ end_cover: 结束并覆盖上一条message """ if v not in [ - 'start', 'stream', 'end', 'error', 'info', 'file', 'begin', 'close', 'end_cover', - 'over' + 'start', 'stream', 'end', 'error', 'info', 'file', 'begin', 'close', 'end_cover', + 'over' ]: raise ValueError('type must be start, stream, end, error, info, or file') return v @@ -179,12 +180,13 @@ class ChatResponse(ChatMessage): class FileResponse(ChatMessage): """File response schema.""" - data: Any + data: Any = None data_type: str type: str = 'file' is_bot: bool = True - @validator('data_type') + @field_validator('data_type') + @classmethod def validate_data_type(cls, v): if v not in ['image', 'csv']: raise ValueError('data_type must be image or csv') @@ -210,9 +212,9 @@ class BuiltResponse(BaseModel): class UploadFileResponse(BaseModel): """Upload file response schema.""" - flowId: Optional[str] + flowId: Optional[str] = None file_path: str - relative_path: Optional[str] # minio的相对路径,即object_name + relative_path: Optional[str] = None # minio的相对路径,即object_name class StreamData(BaseModel): @@ -270,6 +272,11 @@ class AssistantUpdateReq(BaseModel): flow_list: List[str] | None = Field(default=None, description='助手的技能ID列表,为None则不更新') knowledge_list: List[int] | None = Field(default=None, description='知识库ID列表,为None则不更新') + @field_validator('model_name', mode='before') + @classmethod + def convert_model_name(cls, v): + return str(v) + class AssistantSimpleInfo(BaseModel): id: str @@ -279,22 +286,22 @@ class AssistantSimpleInfo(BaseModel): user_id: int user_name: str status: int - flow_type: Optional[int] + flow_type: Optional[int] = None write: Optional[bool] = Field(default=False) - group_ids: Optional[List[int]] - tags: Optional[List[Tag]] + group_ids: Optional[List[int]] = None + tags: Optional[List[Tag]] = None create_time: datetime update_time: datetime class AssistantInfo(AssistantBase): - tool_list: List[GptsToolsRead] = Field(default=[], description='助手的工具ID列表') - flow_list: List[FlowRead] = Field(default=[], description='助手的技能ID列表') - knowledge_list: List[KnowledgeRead] = Field(default=[], description='知识库ID列表') + tool_list: List[GptsToolsRead] = Field(default_factory=list, description='助手的工具ID列表') + flow_list: List[FlowRead] = Field(default_factory=list, description='助手的技能ID列表') + knowledge_list: List[KnowledgeRead] = Field(default_factory=list, description='知识库ID列表') class FlowVersionCreate(BaseModel): - name: Optional[str] = Field(default=..., description='版本的名字') + name: Optional[str] = Field(default=None, description='版本的名字') description: Optional[str] = Field(default=None, description='版本的描述') data: Optional[Dict] = Field(default=None, description='技能版本的节点数据数据') original_version_id: Optional[int] = Field(default=None, description='版本的来源版本ID') @@ -303,8 +310,8 @@ class FlowVersionCreate(BaseModel): class FlowCompareReq(BaseModel): inputs: Any = Field(default=None, description='技能运行所需要的输入') - question_list: List[str] = Field(default=[], description='测试case列表') - version_list: List[int] = Field(default=[], description='对比版本ID列表') + question_list: List[str] = Field(default_factory=list, description='测试case列表') + version_list: List[int] = Field(default_factory=list, description='对比版本ID列表') node_id: str = Field(default=None, description='需要对比的节点唯一ID') thread_num: Optional[int] = Field(default=1, description='对比线程数') @@ -315,6 +322,7 @@ class DeleteToolTypeReq(BaseModel): class TestToolReq(BaseModel): server_host: str = Field(default='', description='服务的根地址') + openapi_schema: Optional[str] = Field(default='', description='openapi schema') extra: str = Field(default='', description='Api 对象解析后的extra字段') auth_method: int = Field(default=AuthMethod.NO.value, description='认证类型') auth_type: Optional[str] = Field(default=AuthType.BASIC.value, description='Auth Type') @@ -342,7 +350,7 @@ class OpenAIChatCompletionReq(BaseModel): n: int = Field(default=1, description='返回的答案个数, 助手侧默认为1,暂不支持多个回答') stream: bool = Field(default=False, description='是否开启流式回复') temperature: float = Field(default=0.0, description='模型温度, 传入0或者不传表示不覆盖') - tools: List[dict] = Field(default=[], description='工具列表, 助手暂不支持,使用助手的配置') + tools: List[dict] = Field(default_factory=list, description='工具列表, 助手暂不支持,使用助手的配置') class OpenAIChoice(BaseModel): @@ -380,27 +388,27 @@ class LLMServerCreateReq(BaseModel): limit_flag: Optional[bool] = Field(default=False, description='是否开启每日调用次数限制') limit: Optional[int] = Field(default=0, description='每日调用次数限制') config: Optional[dict] = Field(default=None, description='服务提供方配置') - models: Optional[List[LLMModelCreateReq]] = Field(default=[], description='服务提供方下的模型列表') + models: Optional[List[LLMModelCreateReq]] = Field(default_factory=list, description='服务提供方下的模型列表') class LLMModelInfo(LLMModelBase): - id: Optional[int] + id: Optional[int] = None class LLMServerInfo(LLMServerBase): - id: Optional[int] - models: List[LLMModelInfo] = Field(default=[], description='模型列表') + id: Optional[int] = None + models: List[LLMModelInfo] = Field(default_factory=list, description='模型列表') class KnowledgeLLMConfig(BaseModel): - embedding_model_id: Optional[int] = Field(description='知识库默认embedding模型的ID') - source_model_id: Optional[int] = Field(description='知识库溯源模型的ID') - extract_title_model_id: Optional[int] = Field(description='文档知识库提取标题模型的ID') - qa_similar_model_id: Optional[int] = Field(description='QA知识库相似问模型的ID') + embedding_model_id: Optional[int] = Field(None, description='知识库默认embedding模型的ID') + source_model_id: Optional[int] = Field(None, description='知识库溯源模型的ID') + extract_title_model_id: Optional[int] = Field(None, description='文档知识库提取标题模型的ID') + qa_similar_model_id: Optional[int] = Field(None, description='QA知识库相似问模型的ID') class AssistantLLMItem(BaseModel): - model_id: Optional[int] = Field(description='模型的ID') + model_id: Optional[int] = Field(None, description='模型的ID') agent_executor_type: Optional[str] = Field(default='ReAct', description='执行模式。function call 或者 ReAct') knowledge_max_content: Optional[int] = Field(default=15000, description='知识库检索最大字符串数') @@ -410,69 +418,71 @@ class AssistantLLMItem(BaseModel): class AssistantLLMConfig(BaseModel): - llm_list: Optional[List[AssistantLLMItem]] = Field(default=[], description='助手可选的LLM列表') - auto_llm: Optional[AssistantLLMItem] = Field(description='助手画像自动优化模型的配置') + llm_list: Optional[List[AssistantLLMItem]] = Field(default_factory=list, description='助手可选的LLM列表') + auto_llm: Optional[AssistantLLMItem] = Field(None, description='助手画像自动优化模型的配置') class EvaluationLLMConfig(BaseModel): - model_id: Optional[int] = Field(description='评测功能默认模型的ID') + model_id: Optional[int] = Field(None, description='评测功能默认模型的ID') class Icon(BaseModel): enabled: bool - image: Optional[str] + image: Optional[str] = None + relative_path: Optional[str] = None class WSModel(BaseModel): - key: Optional[str] + key: Optional[str] = None id: str - name: Optional[str] - displayName: Optional[str] + name: Optional[str] = None + displayName: Optional[str] = None class WSPrompt(BaseModel): enabled: bool - prompt: Optional[str] - model: Optional[str] - tool: Optional[str] - bingKey: Optional[str] - bingUrl: Optional[str] + prompt: Optional[str] = None + model: Optional[str] = None + tool: Optional[str] = None + bingKey: Optional[str] = None + bingUrl: Optional[str] = None class WorkstationConfig(BaseModel): menuShow: bool = Field(default=True, description='是否显示左侧菜单栏') maxTokens: Optional[int] = Field(default=1500, description='最大token数') - sidebarIcon: Optional[Icon] - assistantIcon: Optional[Icon] - sidebarSlogan: Optional[str] = Field(description='侧边栏slogan') - welcomeMessage: Optional[str] = Field() - functionDescription: Optional[str] = Field() - inputPlaceholder: Optional[str] - models: Optional[Union[List[WSModel], str]] - voiceInput: Optional[WSPrompt] - webSearch: Optional[WSPrompt] - knowledgeBase: Optional[WSPrompt] - fileUpload: Optional[WSPrompt] + sidebarIcon: Optional[Icon] = None + assistantIcon: Optional[Icon] = None + sidebarSlogan: Optional[str] = Field(default='', description='侧边栏slogan') + welcomeMessage: Optional[str] = Field(default='') + functionDescription: Optional[str] = Field(default='') + inputPlaceholder: Optional[str] = '' + models: Optional[Union[List[WSModel], str]] = None + voiceInput: Optional[WSPrompt] = None + webSearch: Optional[WSPrompt] = None + knowledgeBase: Optional[WSPrompt] = None + fileUpload: Optional[WSPrompt] = None # 文件切分请求基础参数 class FileProcessBase(BaseModel): knowledge_id: int = Field(..., description='知识库ID') - separator: Optional[List[str]] = Field(default=['\n\n', '\n'], description='切分文本规则, 不传则为默认') - separator_rule: Optional[List[str]] = Field(default=['after', 'after'], + separator: Optional[List[str]] = Field(default=None, description='切分文本规则, 不传则为默认') + separator_rule: Optional[List[str]] = Field(default=None, description='切分规则前还是后进行切分;before/after') chunk_size: Optional[int] = Field(default=1000, description='切分文本长度,不传则为默认') chunk_overlap: Optional[int] = Field(default=100, description='切分文本重叠长度,不传则为默认') - @root_validator - def check_separator_rule(cls, values): - if values['separator'] is None: + @model_validator(mode='before') + @classmethod + def check_separator_rule(cls, values: Any): + if values.get('separator', None) is None: values['separator'] = ['\n\n', '\n'] - if values['separator_rule'] is None: + if values.get('separator_rule', None) is None: values['separator_rule'] = ['after' for _ in values['separator']] - if values['chunk_size'] is None: + if values.get('chunk_size', None) is None: values['chunk_size'] = 1000 - if values['chunk_overlap'] is None: + if values.get('chunk_overlap') is None: values['chunk_overlap'] = 100 return values diff --git a/src/backend/bisheng/api/v1/server.py b/src/backend/bisheng/api/v1/server.py index cb14e0fad..dd2c20d14 100644 --- a/src/backend/bisheng/api/v1/server.py +++ b/src/backend/bisheng/api/v1/server.py @@ -26,7 +26,7 @@ thread_pool = ThreadPoolExecutor(3) required_param = ['type', 'pymodel_type', 'gpu_memory', 'instance_groups'] -@router.post('/add', response_model=UnifiedResponseModel[ServerRead], status_code=201) +@router.post('/add') async def add_server(*, server: ServerCreate): try: db_server = Server.from_orm(server) @@ -42,7 +42,7 @@ async def add_server(*, server: ServerCreate): raise HTTPException(status_code=500, detail=str(exc)) from exc -@router.get('/list_server', response_model=UnifiedResponseModel[List[ServerRead]], status_code=200) +@router.get('/list_server') async def list_server(): try: with session_getter() as session: @@ -73,7 +73,7 @@ async def delete_server(*, server_id: int): raise HTTPException(status_code=500, detail=str(exc)) from exc -@router.get('/list', response_model=UnifiedResponseModel[List[ModelDeployRead]], status_code=201) +@router.get('/list') async def list(*, query: ModelDeployQuery = None): try: # 更新模型 @@ -105,7 +105,7 @@ async def list(*, query: ModelDeployQuery = None): raise HTTPException(status_code=500, detail=str(exc)) from exc -@router.get('/model/{deploy_id}', response_model=UnifiedResponseModel[ModelDeployRead], status_code=201) +@router.get('/model/{deploy_id}') async def get_model_deploy(*, deploy_id: int): try: model_deploy = ModelDeployDao.find_model(deploy_id) @@ -117,7 +117,7 @@ async def get_model_deploy(*, deploy_id: int): raise HTTPException(status_code=500, detail=str(exc)) from exc -@router.post('/update', response_model=UnifiedResponseModel[ModelDeployRead], status_code=201) +@router.post('/update') async def update_deploy(*, deploy: ModelDeployUpdate): try: with session_getter() as session: diff --git a/src/backend/bisheng/api/v1/skillcenter.py b/src/backend/bisheng/api/v1/skillcenter.py index 7a2c19b44..d68411694 100644 --- a/src/backend/bisheng/api/v1/skillcenter.py +++ b/src/backend/bisheng/api/v1/skillcenter.py @@ -17,9 +17,7 @@ router = APIRouter(prefix='/skill', tags=['Skills'], dependencies=[Depends(get_l ORDER_GAP = 65535 -@router.post('/template/create', - response_model=UnifiedResponseModel[TemplateRead], - status_code=201) +@router.post('/template/create') def create_template(*, template: TemplateCreate): """Create a new flow.""" db_template = Template.model_validate(template) @@ -46,7 +44,7 @@ def create_template(*, template: TemplateCreate): return resp_200(db_template) -@router.get('/template', response_model=UnifiedResponseModel[list[Template]], status_code=200) +@router.get('/template') def read_template(page_size: Optional[int] = None, page_name: Optional[int] = None, flow_type: Optional[int] = None, @@ -70,12 +68,16 @@ def read_template(page_size: Optional[int] = None, with session_getter() as session: template_session = session.exec(sql) templates = template_session.mappings().all() + res = [] + for one in templates: + res.append(Template.model_validate(one)) + return resp_200(res) + except Exception as e: raise HTTPException(status_code=500, detail=str(e)) from e - return resp_200(templates) -@router.post('/template/{id}', response_model=UnifiedResponseModel[TemplateRead], status_code=200) +@router.post('/template/{id}') def update_template(*, id: int, template: TemplateUpdate): """Update a flow.""" with session_getter() as session: diff --git a/src/backend/bisheng/api/v1/tag.py b/src/backend/bisheng/api/v1/tag.py index 5fcfa1891..76e8b809a 100644 --- a/src/backend/bisheng/api/v1/tag.py +++ b/src/backend/bisheng/api/v1/tag.py @@ -11,7 +11,7 @@ from bisheng.database.models.tag import Tag, TagLink router = APIRouter(prefix='/tag', tags=['Tag']) -@router.get('', response_model=UnifiedResponseModel[List[Tag]]) +@router.get('') def get_all_tag(request: Request, login_user: UserPayload = Depends(get_login_user), keyword: str = Query(default=None, description='搜索关键字'), @@ -24,7 +24,7 @@ def get_all_tag(request: Request, }) -@router.post('', response_model=UnifiedResponseModel[Tag]) +@router.post('') def create_tag(request: Request, login_user: UserPayload = Depends(get_admin_user), name: str = Body(..., embed=True, description='标签名称')): @@ -32,7 +32,7 @@ def create_tag(request: Request, return resp_200(result) -@router.put('', response_model=UnifiedResponseModel[Tag]) +@router.put('') def update_tag(request: Request, login_user: UserPayload = Depends(get_admin_user), tag_id: int = Body(..., embed=True, description='标签ID'), @@ -41,7 +41,7 @@ def update_tag(request: Request, return resp_200(result) -@router.delete('', response_model=UnifiedResponseModel) +@router.delete('') def delete_tag(request: Request, login_user: UserPayload = Depends(get_admin_user), tag_id: int = Body(..., embed=True, description='标签ID')): @@ -49,7 +49,7 @@ def delete_tag(request: Request, return resp_200() -@router.post('/link', response_model=UnifiedResponseModel[TagLink]) +@router.post('/link') def create_tag_link(request: Request, login_user: UserPayload = Depends(get_login_user), tag_id: int = Body(..., embed=True, description='标签ID'), @@ -59,7 +59,7 @@ def create_tag_link(request: Request, return resp_200(result) -@router.delete('/link', response_model=UnifiedResponseModel) +@router.delete('/link') def delete_tag_link( request: Request, login_user: UserPayload = Depends(get_login_user), @@ -70,7 +70,7 @@ def delete_tag_link( return resp_200() -@router.get('/home', response_model=UnifiedResponseModel[List[Tag]]) +@router.get('/home') def get_home_tag(request: Request, login_user: UserPayload = Depends(get_login_user)): """ @@ -81,7 +81,7 @@ def get_home_tag(request: Request, return resp_200(result) -@router.post('/home', response_model=UnifiedResponseModel[List[Tag]]) +@router.post('/home') def update_home_tag(request: Request, login_user: UserPayload = Depends(get_admin_user), tag_ids: List[int] = Body(..., embed=True, description='标签ID列表')): diff --git a/src/backend/bisheng/api/v1/user.py b/src/backend/bisheng/api/v1/user.py index b87edecee..e60a429ce 100644 --- a/src/backend/bisheng/api/v1/user.py +++ b/src/backend/bisheng/api/v1/user.py @@ -47,7 +47,7 @@ router = APIRouter(prefix='', tags=['User']) oauth2_scheme = OAuth2PasswordBearer(tokenUrl='token') -@router.post('/user/regist', response_model=UnifiedResponseModel[UserRead], status_code=201) +@router.post('/user/regist') async def regist(*, user: UserCreate): # 验证码校验 if settings.get_from_db('use_captcha'): @@ -78,7 +78,7 @@ async def regist(*, user: UserCreate): return resp_200(db_user) -@router.post('/user/sso', response_model=UnifiedResponseModel[UserRead], status_code=201) +@router.post('/user/sso') async def sso(*, request: Request, user: UserCreate): """ 给闭源网关提供的登录接口 """ if settings.get_system_login_method().bisheng_pro: # 判断sso 是否打开 @@ -125,7 +125,7 @@ def clear_error_password_key(username: str): redis_client.delete(error_key) -@router.post('/user/login', response_model=UnifiedResponseModel[UserRead], status_code=201) +@router.post('/user/login') async def login(*, request: Request, user: UserLogin, Authorize: AuthJWT = Depends()): # 验证码校验 if settings.get_from_db('use_captcha'): @@ -184,7 +184,7 @@ async def login(*, request: Request, user: UserLogin, Authorize: AuthJWT = Depen return resp_200(UserRead(role=str(role), web_menu=web_menu, access_token=access_token, **db_user.__dict__)) -@router.get('/user/admin', response_model=UnifiedResponseModel[UserRead], status_code=200) +@router.get('/user/admin') async def get_admins(login_user: UserPayload = Depends(get_login_user)): """ 获取所有的超级管理员账号 @@ -203,7 +203,7 @@ async def get_admins(login_user: UserPayload = Depends(get_login_user)): raise HTTPException(status_code=500, detail='用户信息失败') -@router.get('/user/info', response_model=UnifiedResponseModel[UserRead], status_code=201) +@router.get('/user/info') async def get_info(login_user: UserPayload = Depends(get_login_user)): # check if user already exist try: diff --git a/src/backend/bisheng/api/v1/usergroup.py b/src/backend/bisheng/api/v1/usergroup.py index ff20517de..ed5fc410f 100644 --- a/src/backend/bisheng/api/v1/usergroup.py +++ b/src/backend/bisheng/api/v1/usergroup.py @@ -23,7 +23,7 @@ from bisheng.database.models.user_group import UserGroupDao, UserGroupRead router = APIRouter(prefix='/group', tags=['User'], dependencies=[Depends(get_login_user)]) -@router.get('/list', response_model=UnifiedResponseModel[List[GroupRead]]) +@router.get('/list') async def get_all_group(Authorize: AuthJWT = Depends()): """ 获取所有分组 @@ -48,7 +48,7 @@ async def get_all_group(Authorize: AuthJWT = Depends()): return resp_200({'records': groups_res}) -@router.post('/create', response_model=UnifiedResponseModel[GroupRead], status_code=200) +@router.post('/create') async def create_group(request: Request, group: GroupCreate, Authorize: AuthJWT = Depends()): """ 新建用户组 @@ -59,7 +59,7 @@ async def create_group(request: Request, group: GroupCreate, Authorize: AuthJWT return resp_200(RoleGroupService().create_group(request, login_user, group)) -@router.put('/create', response_model=UnifiedResponseModel[GroupRead], status_code=200) +@router.put('/create') async def update_group(request: Request, group: Group, login_user: UserPayload = Depends(get_login_user)): @@ -84,9 +84,7 @@ async def delete_group(request: Request, return RoleGroupService().delete_group(request, login_user, group_id) -@router.post('/set_user_group', - response_model=UnifiedResponseModel[UserGroupRead], - status_code=200) +@router.post('/set_user_group') async def set_user_group(request: Request, user_id: Annotated[int, Body(embed=True)], group_id: Annotated[List[int], Body(embed=True)], @@ -100,9 +98,7 @@ async def set_user_group(request: Request, return resp_200(RoleGroupService().replace_user_groups(request, login_user, user_id, group_id)) -@router.get('/get_user_group', - response_model=UnifiedResponseModel[List[GroupRead]], - status_code=200) +@router.get('/get_user_group') async def get_user_group(user_id: int, Authorize: AuthJWT = Depends()): """ 获取用户所属分组 @@ -111,7 +107,7 @@ async def get_user_group(user_id: int, Authorize: AuthJWT = Depends()): return resp_200(RoleGroupService().get_user_groups_list(user_id)) -@router.get('/get_group_user', response_model=UnifiedResponseModel[List[User]], status_code=200) +@router.get('/get_group_user') async def get_group_user(group_id: int, page_size: int = None, page_num: int = None, @@ -123,9 +119,7 @@ async def get_group_user(group_id: int, return RoleGroupService().get_group_user_list(group_id, page_size, page_num) -@router.post('/set_group_admin', - response_model=UnifiedResponseModel[List[UserGroupRead]], - status_code=200) +@router.post('/set_group_admin') async def set_group_admin( request: Request, user_ids: Annotated[List[int], Body(embed=True)], @@ -152,9 +146,7 @@ async def set_update_user(group_id: Annotated[int, Body(embed=True)], return resp_200(RoleGroupService().set_group_update_user(login_user, group_id)) -@router.get('/get_group_resources', - response_model=UnifiedResponseModel[List[UserGroupRead]], - status_code=200) +@router.get('/get_group_resources') async def get_group_resources(*, group_id: int, resource_type: int, @@ -180,7 +172,7 @@ async def get_group_resources(*, }) -@router.get("/roles", response_model=UnifiedResponseModel) +@router.get("/roles") async def get_group_roles(*, group_id: List[int] = Query(..., description="用户组ID列表"), keyword: str = Query(None, description="搜索关键字"), @@ -203,7 +195,7 @@ async def get_group_roles(*, }) -@router.get("/manage/resources", response_model=UnifiedResponseModel) +@router.get("/manage/resources") async def get_manage_resources(login_user: UserPayload = Depends(get_login_user), keyword: str = Query(None, description="搜索关键字"), page: int = 1, diff --git a/src/backend/bisheng/api/v1/validate.py b/src/backend/bisheng/api/v1/validate.py index dfb4535ce..12730c55e 100644 --- a/src/backend/bisheng/api/v1/validate.py +++ b/src/backend/bisheng/api/v1/validate.py @@ -10,7 +10,7 @@ from fastapi import APIRouter, HTTPException router = APIRouter(prefix='/validate', tags=['Validate']) -@router.post('/code', status_code=200, response_model=UnifiedResponseModel[CodeValidationResponse]) +@router.post('/code', status_code=200) def post_validate_code(code: Code): try: errors = validate_code(code.code) @@ -23,9 +23,7 @@ def post_validate_code(code: Code): return HTTPException(status_code=500, detail=str(e)) -@router.post('/prompt', - status_code=200, - response_model=UnifiedResponseModel[PromptValidationResponse]) +@router.post('/prompt') def post_validate_prompt(prompt_request: ValidatePromptRequest): try: input_variables = validate_prompt(prompt_request.template) diff --git a/src/backend/bisheng/api/v1/variable.py b/src/backend/bisheng/api/v1/variable.py index 554f67884..6da441415 100644 --- a/src/backend/bisheng/api/v1/variable.py +++ b/src/backend/bisheng/api/v1/variable.py @@ -13,7 +13,7 @@ from bisheng.database.models.variable_value import Variable, VariableCreate, Var router = APIRouter(prefix='/variable', tags=['variable']) -@router.post('/', status_code=200, response_model=UnifiedResponseModel[VariableRead]) +@router.post('/', status_code=200) def post_variable(variable: Variable): try: if not variable.version_id: @@ -47,7 +47,7 @@ def post_variable(variable: Variable): return HTTPException(status_code=500, detail=str(e)) -@router.get('/list', response_model=UnifiedResponseModel[List[VariableRead]], status_code=200) +@router.get('/list') def get_variables(*, flow_id: str, node_id: Optional[str] = None, diff --git a/src/backend/bisheng/api/v1/workflow.py b/src/backend/bisheng/api/v1/workflow.py index 7838c5088..ef58fc705 100644 --- a/src/backend/bisheng/api/v1/workflow.py +++ b/src/backend/bisheng/api/v1/workflow.py @@ -28,7 +28,7 @@ from bisheng.utils import generate_uuid router = APIRouter(prefix='/workflow', tags=['Workflow']) -@router.get("/report/file", response_model=UnifiedResponseModel, status_code=200) +@router.get("/report/file") async def get_report_file( request: Request, login_user: UserPayload = Depends(get_login_user), @@ -216,13 +216,13 @@ def change_version(*, return FlowService.change_current_version(request, login_user, flow_id, version_id) -@router.get('/get_one_flow/{flow_id}', response_model=UnifiedResponseModel[FlowReadWithStyle], status_code=200) +@router.get('/get_one_flow/{flow_id}') def read_flow(*, flow_id: str, login_user: UserPayload = Depends(get_login_user)): """Read a flow.""" return FlowService.get_one_flow(login_user, flow_id) -@router.patch('/update/{flow_id}', response_model=UnifiedResponseModel[FlowRead], status_code=200) +@router.patch('/update/{flow_id}') async def update_flow(*, request: Request, flow_id: str, @@ -253,7 +253,7 @@ async def update_flow(*, return resp_200(db_flow) -@router.patch('/status', response_model=UnifiedResponseModel[FlowRead], status_code=200) +@router.patch('/status') async def update_flow_status(request: Request, login_user: UserPayload = Depends(get_login_user), flow_id: str = Body(..., description='技能ID'), version_id: int = Body(..., description='版本ID'), diff --git a/src/backend/bisheng/api/v1/workstation.py b/src/backend/bisheng/api/v1/workstation.py index 5afdaaa2d..70135457f 100644 --- a/src/backend/bisheng/api/v1/workstation.py +++ b/src/backend/bisheng/api/v1/workstation.py @@ -103,22 +103,22 @@ def final_message(conversation: MessageSession, title: str, requestMessage: Chat return f'event: message\ndata: {msg}\n\n' -@router.get('/config', response_model=UnifiedResponseModel[WorkstationConfig]) +@router.get('/config') def get_config( request: Request, login_user: UserPayload = Depends(get_login_user), -) -> UnifiedResponseModel[WorkstationConfig]: +): """ 获取评价相关的模型配置 """ ret = WorkStationService.get_config() return resp_200(data=ret) -@router.post('/config', response_model=UnifiedResponseModel[WorkstationConfig]) +@router.post('/config') def update_config( request: Request, login_user: UserPayload = Depends(get_admin_user), data: WorkstationConfig = Body(..., description='默认模型配置'), -) -> UnifiedResponseModel[WorkstationConfig]: +): """ 更新评价相关的模型配置 """ ret = WorkStationService.update_config(request, login_user, data) return resp_200(data=ret) diff --git a/src/backend/bisheng/api/v2/assistant.py b/src/backend/bisheng/api/v2/assistant.py index 4621482b1..7ed7ba2e0 100644 --- a/src/backend/bisheng/api/v2/assistant.py +++ b/src/backend/bisheng/api/v2/assistant.py @@ -24,7 +24,7 @@ from bisheng.utils import generate_uuid router = APIRouter(prefix='/assistant', tags=['OpenAPI', 'Assistant']) -@router.post('/chat/completions', response_model=OpenAIChatCompletionResp) +@router.post('/chat/completions') async def assistant_chat_completions(request: Request, req_data: OpenAIChatCompletionReq): """ 兼容openai接口格式,所有的错误必须返回非http200的状态码 @@ -105,7 +105,7 @@ async def assistant_chat_completions(request: Request, req_data: OpenAIChatCompl return ORJSONResponse(status_code=500, content=str(exc)) -@router.get('/info/{assistant_id}', response_model=UnifiedResponseModel[AssistantInfo]) +@router.get('/info/{assistant_id}') async def get_assistant_info(request: Request, assistant_id: UUID): """ 获取助手信息, 用系统配置里的default_operator.user的用户信息来做权限校验 diff --git a/src/backend/bisheng/api/v2/filelib.py b/src/backend/bisheng/api/v2/filelib.py index fcabbae45..2b4d49054 100644 --- a/src/backend/bisheng/api/v2/filelib.py +++ b/src/backend/bisheng/api/v2/filelib.py @@ -76,9 +76,7 @@ def clear_knowledge_files(*, request: Request, knowledge_id: int): return resp_200(message='knowledge clear successfully') -@router.post('/file/{knowledge_id}', - response_model=UnifiedResponseModel[KnowledgeFileRead], - status_code=200) +@router.post('/file/{knowledge_id}') async def upload_file(request: Request, knowledge_id: int, separator: Optional[List[str]] = Form(default=None, @@ -147,7 +145,7 @@ def get_filelist(request: Request, return resp_200(data={'data': data, 'total': total, 'writeable': flag}) -@router.post('/chunks', response_model=UnifiedResponseModel[KnowledgeFileRead], status_code=200) +@router.post('/chunks') async def post_chunks(request: Request, knowledge_id: int = Form(...), metadata: str = Form(...), @@ -177,9 +175,7 @@ async def post_chunks(request: Request, return resp_200(data=res[0]) -@router.post('/chunks_string', - response_model=UnifiedResponseModel[KnowledgeFileRead], - status_code=200) +@router.post('/chunks_string') async def post_string_chunks(request: Request, document: ChunkInput): """ 获取知识库文件信息. """ @@ -276,7 +272,7 @@ def add_qa(*, return resp_200(res) -@router.post('/add_relative_qa', response_model=QAKnowledge) +@router.post('/add_relative_qa') def append_qa(*, knowledge_id: int = Body(embed=True), data: APIAppendQAParam = Body(embed=True), diff --git a/src/backend/bisheng/api/v2/workflow.py b/src/backend/bisheng/api/v2/workflow.py index 05b2a84a2..86fabbdd0 100644 --- a/src/backend/bisheng/api/v2/workflow.py +++ b/src/backend/bisheng/api/v2/workflow.py @@ -27,7 +27,7 @@ async def invoke_workflow(request: Request, workflow_id: UUID = Body(..., description='工作流唯一ID'), stream: Optional[bool] = Body(default=True, description='是否流式调用'), user_input: Optional[dict] = Body(default=None, description='用户输入', alias='input'), - message_id: Optional[str] = Body(default=None, description='消息ID'), + message_id: Optional[int] = Body(default=None, description='消息ID'), session_id: Optional[str] = Body(default=None, description='会话ID,一次workflow调用的唯一标识')): login_user = get_default_operator() workflow_id = workflow_id.hex @@ -69,10 +69,13 @@ async def invoke_workflow(request: Request, async for event in workflow.get_response_until_break(): if event.category == WorkflowEventType.NodeRun.value: continue + # 非流式请求,过滤掉节点产生的流式输出事件 + if not stream and event.category == WorkflowEventType.StreamMsg.value and event.type == 'stream': + continue workflow_stream = WorkflowStream(session_id=session_id, data=WorkFlowService.convert_chat_response_to_workflow_event(event)) event_list.append(workflow_stream.data) - yield f'data: {workflow_stream.json()}\n\n' + yield f'data: {workflow_stream.model_dump_json()}\n\n' tmp_status_info = workflow.get_workflow_status() if tmp_status_info['status'] in [WorkflowStatus.SUCCESS.value, WorkflowStatus.FAILED.value]: workflow.clear_workflow_status() @@ -80,7 +83,7 @@ async def invoke_workflow(request: Request, workflow_stream = WorkflowStream(session_id=session_id, data=WorkflowEvent(event=WorkflowEventType.Close.value)) event_list.append(workflow_stream.data) - yield f'data: {workflow_stream.json()}\n\n' + yield f'data: {workflow_stream.model_dump_json()}\n\n' res = [] # 非流式返回累计的事件列表 diff --git a/src/backend/bisheng/cache/redis.py b/src/backend/bisheng/cache/redis.py index 6a0dde30a..f7db529e8 100644 --- a/src/backend/bisheng/cache/redis.py +++ b/src/backend/bisheng/cache/redis.py @@ -35,7 +35,7 @@ class RedisClient: hosts = [eval(x) for x in redis_conf.pop('sentinel_hosts')] password = redis_conf.pop('sentinel_password') master = redis_conf.pop('sentinel_master') - sentinel = Sentinel(sentinels=hosts, socket_timeout=0.1, password=password) + sentinel = Sentinel(sentinels=hosts, socket_timeout=0.1, sentinel_kwargs={'password': password}) # 获取主节点的连接 self.connection = sentinel.master_for(master, socket_timeout=0.1, **redis_conf) diff --git a/src/backend/bisheng/chat/client.py b/src/backend/bisheng/chat/client.py index 7fe5dbd7b..2333dfe16 100644 --- a/src/backend/bisheng/chat/client.py +++ b/src/backend/bisheng/chat/client.py @@ -233,12 +233,18 @@ class ChatClient: # 将流式输出的内容写到数据库内 answer = '' + reasoning_answer = '' while not self.stream_queue.empty(): msg = self.stream_queue.get() if msg.get('type') == 'answer': answer += msg.get('content', '') + elif msg.get('type') == 'reasoning': + reasoning_answer += msg.get('content', '') # 有流式输出内容的话,记录流式输出内容到数据库 + if reasoning_answer.split(): + res = await self.add_message('bot', answer, 'reasoning_answer', 'break_answer') + await self.send_response('reasoning_answer', 'end', '', message_id=res.id if res else None) if answer.strip(): res = await self.add_message('bot', answer, 'answer', 'break_answer') await self.send_response('answer', 'end', '', message_id=res.id if res else None) diff --git a/src/backend/bisheng/chat/clients/workflow_client.py b/src/backend/bisheng/chat/clients/workflow_client.py index 55d3ea255..729f9be09 100644 --- a/src/backend/bisheng/chat/clients/workflow_client.py +++ b/src/backend/bisheng/chat/clients/workflow_client.py @@ -160,7 +160,7 @@ class WorkflowClient(BaseClient): return True status_info = self.workflow.get_workflow_status() - if status_info['status'] in [WorkflowStatus.FAILED.value, WorkflowStatus.SUCCESS.value]: + if not status_info or status_info['status'] in [WorkflowStatus.FAILED.value, WorkflowStatus.SUCCESS.value]: await self.send_response('processing', 'close', '') self.workflow.clear_workflow_status() self.workflow = None diff --git a/src/backend/bisheng/chat/handlers.py b/src/backend/bisheng/chat/handlers.py index 1cfcd5a66..a3984dfa3 100644 --- a/src/backend/bisheng/chat/handlers.py +++ b/src/backend/bisheng/chat/handlers.py @@ -10,6 +10,7 @@ from bisheng.chat.manager import ChatManager from bisheng.chat.utils import judge_source, process_graph, process_source_document from bisheng.database.base import session_getter from bisheng.database.models.report import Report +from bisheng.database.models.message import ChatMessage as ChatMessageDB, ChatMessageDao from bisheng.interface.importing.utils import import_by_type from bisheng.interface.initialize.loading import instantiate_llm from bisheng.settings import settings @@ -67,6 +68,38 @@ class Handler: else: logger.error(f'act=auto_gen act={action}') else: + # 将流式输出的内容写到数据库内 + answer = '' + reasoning_answer = '' + while not self.stream_queue.empty(): + msg = self.stream_queue.get() + if msg.get('type') == 'answer': + answer += msg.get('content', '') + elif msg.get('type') == 'reasoning': + reasoning_answer += msg.get('content', '') + if reasoning_answer.strip(): + chat_message = ChatMessageDB(flow_id=client_id, chat_id=chat_id, + message=reasoning_answer, + category='answer', + type='end', + user_id=user_id, + remark='break_answer', + is_bot=True) + if chat_id: + db_message = ChatMessageDao.insert_one(chat_message) + await session.send_json(client_id, chat_id, ChatMessage(**db_message.model_dump(), message_id=db_message.id), add=False) + + if answer.strip(): + chat_message = ChatMessageDB(flow_id=client_id, chat_id=chat_id, + message=answer, + category='answer', + type='end', + user_id=user_id, + remark='break_answer', + is_bot=True) + if chat_id: + db_message = ChatMessageDao.insert_one(chat_message) + await session.send_json(client_id, chat_id, ChatMessage(**db_message.model_dump(), message_id=db_message.id), add=False) # 普通技能的stop res = thread_pool.cancel_task([key]) # 将进行中的任务进行cancel if res[0]: @@ -75,18 +108,7 @@ class Handler: close = ChatResponse(type='close') await session.send_json(client_id, chat_id, res, add=False) await session.send_json(client_id, chat_id, close, add=False) - answer = '' - # 记录中止后产生的流式输出内容 - while not self.stream_queue.empty(): - answer += self.stream_queue.get() - if answer.strip(): - chat_message = ChatMessage(message=answer, - category='answer', - type='end', - user_id=user_id, - remark='break_answer', - is_bot=True) - session.chat_history.add_message(client_id, chat_id, chat_message) + logger.info('process_stop done') async def process_report(self, diff --git a/src/backend/bisheng/chat/manager.py b/src/backend/bisheng/chat/manager.py index a3a9ef3e7..62b8e62a0 100644 --- a/src/backend/bisheng/chat/manager.py +++ b/src/backend/bisheng/chat/manager.py @@ -51,10 +51,11 @@ class ChatHistory(Subject): from bisheng.database.models.message import ChatMessage message.flow_id = client_id message.chat_id = chat_id + db_message = None if chat_id and (message.message or message.intermediate_steps or message.files) and message.type != 'stream': msg = message.copy() - msg.message = json.dumps(msg.message) if isinstance(msg.message, dict) else msg.message + msg.message = json.dumps(msg.message, ensure_ascii=False) if isinstance(msg.message, dict) else msg.message files = json.dumps(msg.files) if msg.files else '' msg.__dict__.pop('files') db_message = ChatMessage(files=files, **msg.__dict__) @@ -67,6 +68,7 @@ class ChatHistory(Subject): if not isinstance(message, FileResponse): self.notify() + return db_message def empty_history(self, client_id: str, chat_id: str): """Empty the chat history for a client.""" diff --git a/src/backend/bisheng/chat/utils.py b/src/backend/bisheng/chat/utils.py index cfdb0b2a3..0ae1e297d 100644 --- a/src/backend/bisheng/chat/utils.py +++ b/src/backend/bisheng/chat/utils.py @@ -1,4 +1,6 @@ import json +import re +import ast from enum import Enum from typing import Dict, List from urllib.parse import unquote, urlparse @@ -81,7 +83,8 @@ def extract_answer_keys(answer, llm): llm_chain = LLMChain(llm=llm, prompt=PromptTemplate.from_template(prompt_template)) try: keywords_str = llm_chain.run(answer) - keywords = eval(keywords_str[9:]) + keywords_str = re.sub('.*', '', keywords_str, flags=re.S).strip() + keywords = ast.literal_eval(keywords_str[9:]) except Exception: import jieba.analyse logger.warning(f'llm {llm} extract_not_support, change to jieba') diff --git a/src/backend/bisheng/database/constants.py b/src/backend/bisheng/database/constants.py index 84708fc2c..67f1bbd0e 100644 --- a/src/backend/bisheng/database/constants.py +++ b/src/backend/bisheng/database/constants.py @@ -1,5 +1,18 @@ +from enum import Enum # 默认普通用户角色的ID DefaultRole = 2 # 超级管理员角色ID AdminRole = 1 + + +class ToolPresetType(Enum): + PRESET = 1 # 预置工具 + API = 0 # 自定义API工具 + MCP = 2 # mcp类型的工具 + + +# 消息表里一些基础的category类型 +class MessageCategory(Enum): + QUESTION = 'question' # 用户问题 + ANSWER = 'answer' # 答案 diff --git a/src/backend/bisheng/database/data/template.json b/src/backend/bisheng/database/data/template.json index ba9fa05fe..7d154e842 100644 --- a/src/backend/bisheng/database/data/template.json +++ b/src/backend/bisheng/database/data/template.json @@ -27838,7 +27838,7 @@ { "update_time": "2025-03-03 20:06:19", "parameters": null, - "name": "「多助手并行+穿行报告生成」", + "name": "「多助手并行+串行报告生成」", "description": "", "flow_id": "497051c652b94cc3a2ffb482c520b1bd", "api_parameters": null, diff --git a/src/backend/bisheng/database/init_data.py b/src/backend/bisheng/database/init_data.py index 68a8473ed..a35c1c2e5 100644 --- a/src/backend/bisheng/database/init_data.py +++ b/src/backend/bisheng/database/init_data.py @@ -173,7 +173,7 @@ def read_from_conf(file_path: str) -> str: def upload_preset_minio_file(): """ 上传预置文件到minio, 为了和工作流模板配合 """ minio_client = MinioClient() - # 上传 「多助手并行+穿行报告生成」 工作流模板需要的docx文件 + # 上传 「多助手并行+串行报告生成」 工作流模板需要的docx文件 template_data = read_from_conf('data/0254d1808a5247d2a3ee0d0011819acb.docx') minio_client.upload_minio_data('workflow/report/0254d1808a5247d2a3ee0d0011819acb.docx', template_data, len(template_data), 'application/vnd.openxmlformats-officedocument.wordprocessingml.document') diff --git a/src/backend/bisheng/database/models/assistant.py b/src/backend/bisheng/database/models/assistant.py index 712cc0ec1..02b7fb8f0 100644 --- a/src/backend/bisheng/database/models/assistant.py +++ b/src/backend/bisheng/database/models/assistant.py @@ -2,12 +2,13 @@ from datetime import datetime from enum import Enum from typing import List, Optional, Tuple +from sqlalchemy import JSON, Column, DateTime, Text, and_, func, or_, text +from sqlmodel import Field, select + from bisheng.database.base import session_getter from bisheng.database.models.base import SQLModelSerializable from bisheng.database.models.role_access import AccessType, RoleAccess from bisheng.utils import generate_uuid -from sqlalchemy import JSON, Column, DateTime, Text, and_, func, or_, text -from sqlmodel import Field, select class AssistantStatus(Enum): @@ -23,33 +24,32 @@ class AssistantBase(SQLModelSerializable): system_prompt: str = Field(default='', sa_column=Column(Text), description='系统提示词') prompt: str = Field(default='', sa_column=Column(Text), description='用户可见描述词') guide_word: Optional[str] = Field(default='', sa_column=Column(Text), description='开场白') - guide_question: Optional[List] = Field(sa_column=Column(JSON), description='引导问题') + guide_question: Optional[List] = Field(default_factory=list, sa_column=Column(JSON), description='引导问题') model_name: str = Field(default='', description='对应模型管理里模型的唯一ID') temperature: float = Field(default=0.5, description='模型温度') max_token: int = Field(default=32000, description='最大token数') status: int = Field(default=AssistantStatus.OFFLINE.value, description='助手是否上线') user_id: int = Field(default=0, description='创建用户ID') is_delete: int = Field(default=0, description='删除标志') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, + sa_column=Column(DateTime, + nullable=False, + server_default=text( + 'CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) class AssistantLinkBase(SQLModelSerializable): - id: Optional[int] = Field(nullable=False, primary_key=True, description='唯一ID') - assistant_id: Optional[str] = Field(index=True, description='助手ID') + id: Optional[int] = Field(default=None, nullable=False, primary_key=True, description='唯一ID') + assistant_id: Optional[str] = Field(default=0, index=True, description='助手ID') tool_id: Optional[int] = Field(default=0, index=True, description='工具ID') flow_id: Optional[str] = Field(default='', index=True, description='技能ID') knowledge_id: Optional[int] = Field(default=0, index=True, description='知识库ID') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) class Assistant(AssistantBase, table=True): diff --git a/src/backend/bisheng/database/models/audit_log.py b/src/backend/bisheng/database/models/audit_log.py index cf505f1c4..637a47358 100644 --- a/src/backend/bisheng/database/models/audit_log.py +++ b/src/backend/bisheng/database/models/audit_log.py @@ -67,7 +67,7 @@ class AuditLogBase(SQLModelSerializable): system_id: Optional[str] = Field(index=True, description="系统模块") event_type: Optional[str] = Field(index=True, description="操作行为") object_type: Optional[str] = Field(index=True, description="操作对象类型") - object_id: Optional[int] = Field(index=True, description="操作对象ID") + object_id: Optional[str] = Field(index=True, description="操作对象ID") object_name: Optional[str] = Field(sa_column=Column(Text), description="操作对象名称") note: Optional[str] = Field(sa_column=Column(Text), description="操作备注") ip_address: Optional[str] = Field(index=True, description="操作时客户端的IP地址") diff --git a/src/backend/bisheng/database/models/base.py b/src/backend/bisheng/database/models/base.py index 2aca81001..bee308e98 100644 --- a/src/backend/bisheng/database/models/base.py +++ b/src/backend/bisheng/database/models/base.py @@ -1,7 +1,12 @@ from datetime import datetime +from typing import Union, Dict, Any + +from sqlmodel.main import IncEx +from typing_extensions import Literal import orjson from sqlmodel import SQLModel +from pydantic import ConfigDict def orjson_dumps(v, *, default=None, sort_keys=False, indent_2=True): @@ -21,11 +26,7 @@ def orjson_dumps(v, *, default=None, sort_keys=False, indent_2=True): class SQLModelSerializable(SQLModel): - - class Config: - orm_mode = True - json_loads = orjson.loads - json_dumps = orjson_dumps + model_config = ConfigDict(from_attributes=True) def to_dict(self): result = self.model_dump() @@ -36,3 +37,27 @@ class SQLModelSerializable(SQLModel): value = value.isoformat() result[column] = value return result + + def model_dump( + self, + *, + mode: Union[Literal["json", "python"], str] = "json", + include: IncEx = None, + exclude: IncEx = None, + by_alias: bool = False, + exclude_unset: bool = False, + exclude_defaults: bool = False, + exclude_none: bool = False, + round_trip: bool = False, + warnings: bool = True, + ) -> Dict[str, Any]: + return super().model_dump( + mode=mode, + include=include, + exclude=exclude, + by_alias=by_alias, + exclude_unset=exclude_unset, + exclude_defaults=exclude_defaults, + exclude_none=exclude_none, + round_trip=round_trip, + warnings=warnings) diff --git a/src/backend/bisheng/database/models/component.py b/src/backend/bisheng/database/models/component.py index 3f9a08625..ee2737995 100644 --- a/src/backend/bisheng/database/models/component.py +++ b/src/backend/bisheng/database/models/component.py @@ -15,9 +15,9 @@ class ComponentBase(SQLModelSerializable): version: str = Field(default='', index=True, description='组件版本') user_id: int = Field(default=None, index=True, description='创建人ID') user_name: str = Field(default=None, description='创建人姓名') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field(sa_column=Column( + update_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) diff --git a/src/backend/bisheng/database/models/config.py b/src/backend/bisheng/database/models/config.py index b6618bbc6..1218e67b6 100644 --- a/src/backend/bisheng/database/models/config.py +++ b/src/backend/bisheng/database/models/config.py @@ -2,11 +2,12 @@ from datetime import datetime from enum import Enum from typing import Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable -from sqlalchemy import Column, DateTime, Text, text +from sqlalchemy import Column, DateTime, text, Text from sqlmodel import Field, select +from bisheng.database.models.base import SQLModelSerializable +from bisheng.database.base import session_getter + class ConfigKeyEnum(Enum): INIT_DB = 'initdb_config' # 默认系统配置 @@ -22,14 +23,11 @@ class ConfigKeyEnum(Enum): class ConfigBase(SQLModelSerializable): key: str = Field(index=True, unique=True) value: str = Field(sa_column=Column(Text)) - comment: Optional[str] = Field(index=False) - create_time: Optional[datetime] = Field(sa_column=Column( + comment: Optional[str] = Field(default=None, index=False) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) class Config(ConfigBase, table=True): @@ -46,8 +44,8 @@ class ConfigCreate(ConfigBase): class ConfigUpdate(SQLModelSerializable): key: str - value: Optional[str] - comment: Optional[str] + value: Optional[str] = None + comment: Optional[str] = None class ConfigDao(ConfigBase): diff --git a/src/backend/bisheng/database/models/dataset.py b/src/backend/bisheng/database/models/dataset.py index cf85663a8..c318ed082 100644 --- a/src/backend/bisheng/database/models/dataset.py +++ b/src/backend/bisheng/database/models/dataset.py @@ -1,25 +1,23 @@ from datetime import datetime from typing import Any, List, Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import Column, DateTime, delete, text from sqlmodel import Field, select +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable + class DatasetBase(SQLModelSerializable): user_id: Optional[int] = Field(index=True, description='创建用户id') name: str = Field(index=True, description='数据集名称') type: str = Field(index=False, default=0, description='预留字段') - description: Optional[str] = Field(index=False, description='数据集描述') - object_name: Optional[str] = Field(index=False, description='数据集S3名称') - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + description: Optional[str] = Field(default=None, index=False, description='数据集描述') + object_name: Optional[str] = Field(default=None, index=False, description='数据集S3名称') + create_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=True, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) class Dataset(DatasetBase, table=True): diff --git a/src/backend/bisheng/database/models/evaluation.py b/src/backend/bisheng/database/models/evaluation.py index 7e07756d0..df2466e3b 100644 --- a/src/backend/bisheng/database/models/evaluation.py +++ b/src/backend/bisheng/database/models/evaluation.py @@ -2,11 +2,12 @@ from datetime import datetime from enum import Enum from typing import List, Optional, Dict -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import Column, DateTime, Text, text, func, and_, JSON from sqlmodel import Field, select +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable + class ExecType(Enum): FLOW = 'flow' @@ -29,15 +30,16 @@ class EvaluationBase(SQLModelSerializable): status: int = Field(index=True, default=1, description='任务执行状态。1:执行中 2: 执行失败 3:执行成功') prompt: str = Field(default='', sa_column=Column(Text), description='评测指令文本') result_file_path: str = Field(default='', description='评测结果的 minio 地址') - result_score: Optional[Dict] = Field(sa_column=Column(JSON), default=None, description='最终评测分数') + result_score: Optional[Dict] = Field(default=None, sa_column=Column(JSON), description='最终评测分数') is_delete: int = Field(default=0, description='是否删除') - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, + sa_column=Column(DateTime, nullable=False, + server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, + sa_column=Column(DateTime, + nullable=True, + server_default=text('CURRENT_TIMESTAMP'), + onupdate=text('CURRENT_TIMESTAMP'))) class Evaluation(EvaluationBase, table=True): @@ -46,7 +48,7 @@ class Evaluation(EvaluationBase, table=True): class EvaluationRead(EvaluationBase): id: int - user_name: Optional[str] + user_name: Optional[str] = None class EvaluationCreate(EvaluationBase): diff --git a/src/backend/bisheng/database/models/finetune.py b/src/backend/bisheng/database/models/finetune.py index 970ee0db2..f26b14f45 100644 --- a/src/backend/bisheng/database/models/finetune.py +++ b/src/backend/bisheng/database/models/finetune.py @@ -2,7 +2,7 @@ from datetime import datetime from enum import Enum from typing import Any, Dict, List, Optional -from pydantic import BaseModel, validator +from pydantic import field_validator, BaseModel from sqlalchemy.dialects.mysql import LONGTEXT from sqlmodel import JSON, Column, DateTime, Field, func, select, text, update @@ -11,7 +11,6 @@ from bisheng.database.models.base import SQLModelSerializable from bisheng.utils import generate_uuid - class TrainMethod(Enum): FULL = 'full' FREEZE = 'freeze' @@ -44,17 +43,17 @@ class FinetuneBase(SQLModelSerializable): model_name: str = Field(index=True, max_length=50, description='训练模型的名称') method: str = Field(default=TrainMethod.FULL.value, nullable=False, max_length=20, description='训练方法') extra_params: Dict = Field(sa_column=Column(JSON), description='训练任务所需的额外参数') - train_data: Optional[List[Dict]] = Field(sa_column=Column(JSON), description='个人训练数据集信息') - preset_data: Optional[List[Dict]] = Field(sa_column=Column(JSON), description='预置训练数据集信息') + train_data: Optional[List[Dict]] = Field(default=None, sa_column=Column(JSON), description='个人训练数据集信息') + preset_data: Optional[List[Dict]] = Field(default=None, sa_column=Column(JSON), description='预置训练数据集信息') status: int = Field(default=FinetuneStatus.TRAINING.value, index=True, description='训练任务的状态') reason: Optional[str] = Field(default='', sa_column=Column(LONGTEXT), description='任务失败原因') log_path: Optional[str] = Field(default='', max_length=512, description='训练日志在minio上的路径') - report: Optional[Dict] = Field(sa_column=Column(JSON), description='训练任务的评估报告数据') + report: Optional[Dict] = Field(default=None, sa_column=Column(JSON), description='训练任务的评估报告数据') user_id: int = Field(default=None, index=True, description='创建人ID') user_name: str = Field(default=None, description='创建人姓名') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field(sa_column=Column( + update_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) # 检查训练集数据格式 @@ -69,17 +68,20 @@ class FinetuneBase(SQLModelSerializable): raise ValueError('Finetune.train_data each item must be {name:"",url:"",num:0}') return v - @validator('extra_params') + @field_validator('extra_params') + @classmethod def validate_params(cls, v: Optional[Dict]): if v is None or not isinstance(v, dict): raise ValueError('Finetune.extra_params must be a valid json') return v - @validator('train_data') + @field_validator('train_data') + @classmethod def validate_train_data(cls, v: Optional[Dict]): return cls.validate_train(v) - @validator('preset_data') + @field_validator('preset_data') + @classmethod def validate_preset_data(cls, v: Optional[Dict]): return cls.validate_train(v) @@ -89,10 +91,10 @@ class Finetune(FinetuneBase, table=True): class FinetuneList(BaseModel): - server: Optional[int] = Field(description='关联的RT服务ID') - server_name: Optional[str] = Field(description='关联的RT服务名称') - status: Optional[List[int]] = Field(description='训练任务的状态') - model_name: Optional[str] = Field(description='模型名称, 模糊搜索') + server: Optional[int] = Field(None, description='关联的RT服务ID') + server_name: Optional[str] = Field(None, description='关联的RT服务名称') + status: Optional[List[int]] = Field(None, description='训练任务的状态') + model_name: Optional[str] = Field(None, description='模型名称, 模糊搜索') page: Optional[int] = Field(default=1, description='页码') limit: Optional[int] = Field(default=10, description='每页条数') @@ -111,7 +113,8 @@ class FinetuneExtraParams(BaseModel): max_seq_len: int = Field(8192, gt=0, description='最大序列长度') cpu_load: str = Field('false', description='是否cpu载入') - @validator('per_device_train_batch_size') + @field_validator('per_device_train_batch_size') + @classmethod def validate_batch_size(cls, v: str): try: batch_size = int(v) @@ -122,7 +125,8 @@ class FinetuneExtraParams(BaseModel): except Exception as e: raise ValueError(f'per_device_train_batch_size must be an integer {e}') - @validator('gpus') + @field_validator('gpus') + @classmethod def validate_gpus(cls, v: str): try: gpu_list = v.split(',') diff --git a/src/backend/bisheng/database/models/flow.py b/src/backend/bisheng/database/models/flow.py index d7a5ead56..ffbf69737 100644 --- a/src/backend/bisheng/database/models/flow.py +++ b/src/backend/bisheng/database/models/flow.py @@ -4,19 +4,22 @@ from datetime import datetime from enum import Enum from typing import Dict, List, Optional, Tuple, Union +from pydantic import field_validator +from sqlalchemy import Column, DateTime, String, and_, func, or_, text +from sqlmodel import JSON, Field, select, update + from bisheng.database.base import session_getter from bisheng.database.models.assistant import Assistant from bisheng.database.models.base import SQLModelSerializable from bisheng.database.models.role_access import AccessType, RoleAccess, RoleAccessDao from bisheng.database.models.user_role import UserRoleDao from bisheng.utils import generate_uuid -from pydantic import validator -from sqlalchemy import Column, DateTime, String, and_, func, or_, text -from sqlmodel import JSON, Field, select, update + # if TYPE_CHECKING: + class FlowStatus(Enum): OFFLINE = 1 ONLINE = 2 @@ -31,23 +34,21 @@ class FlowType(Enum): class FlowBase(SQLModelSerializable): name: str = Field(index=True) - user_id: Optional[int] = Field(index=True) - description: Optional[str] = Field(index=False) + user_id: Optional[int] = Field(default=None, index=True) + description: Optional[str] = Field(default=None, index=False) data: Optional[Dict] = Field(default=None) - logo: Optional[str] = Field(index=False) + logo: Optional[str] = Field(default=None, index=False) status: Optional[int] = Field(index=False, default=1) flow_type: Optional[int] = Field(index=False, default=1) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) - create_time: Optional[datetime] = Field(sa_column=Column( + guide_word: Optional[str] = Field(default=None, sa_column=Column(String(length=1000))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=True, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - guide_word: Optional[str] = Field(sa_column=Column(String(length=1000))) - @validator('data') - def validate_json(v): + @field_validator('data', mode='before') + @classmethod + def validate_json(cls, v): if not v: return v if not isinstance(v, dict): @@ -68,13 +69,13 @@ class Flow(FlowBase, table=True): class FlowCreate(FlowBase): - flow_id: Optional[str] + flow_id: Optional[str] = None class FlowRead(FlowBase): id: str - user_name: Optional[str] - version_id: Optional[int] + user_name: Optional[str] = None + version_id: Optional[int] = None class FlowReadWithStyle(FlowRead): @@ -210,7 +211,7 @@ class FlowDao(FlowBase): if status is not None: statement = statement.where(Flow.status == status) if flow_type is not None: - statement = statement.where(Flow.flow_type == flow_type) + statement = statement.where(Flow.flow_type== flow_type) if flow_ids: statement = statement.where(Flow.id.in_(flow_ids)) statement = statement.order_by(Flow.update_time.desc()) diff --git a/src/backend/bisheng/database/models/flow_version.py b/src/backend/bisheng/database/models/flow_version.py index 80f44bb89..25ac43036 100644 --- a/src/backend/bisheng/database/models/flow_version.py +++ b/src/backend/bisheng/database/models/flow_version.py @@ -8,7 +8,7 @@ from sqlalchemy import func from bisheng.database.base import session_getter from bisheng.database.models.base import SQLModelSerializable # if TYPE_CHECKING: -from pydantic import validator +from pydantic import field_validator from sqlmodel import JSON, Field, select, update, text, Column, DateTime from bisheng.database.models.flow import Flow @@ -19,19 +19,20 @@ class FlowVersionBase(SQLModelSerializable): flow_id: str = Field(index=True, max_length=32, description="所属的技能ID") name: str = Field(index=True, description="版本的名字") data: Optional[Dict] = Field(default=None, description="版本的数据") - description: Optional[str] = Field(index=False, description="版本的描述") - user_id: Optional[int] = Field(index=True, description="创建者") + description: Optional[str] = Field(default=None, index=False, description="版本的描述") + user_id: Optional[int] = Field(default=None, index=True, description="创建者") flow_type: Optional[int] = Field(default=1, description="版本的类型") is_current: Optional[int] = Field(default=0, description="是否为正在使用版本") is_delete: Optional[int] = Field(default=0, description="是否删除") original_version_id: Optional[int] = Field(default=None, description="来源版本的ID") - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field(sa_column=Column( + update_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) - @validator('data') - def validate_json(v): + @field_validator('data') + @classmethod + def validate_json(cls, v): # dict_keys(['description', 'name', 'id', 'data']) if not v: return v @@ -54,7 +55,7 @@ class FlowVersionRead(FlowVersionBase): pass -class FlowVersionDao(FlowVersion): +class FlowVersionDao(FlowVersionBase): @classmethod def create_version(cls, version: FlowVersion) -> FlowVersion: diff --git a/src/backend/bisheng/database/models/gpts_tools.py b/src/backend/bisheng/database/models/gpts_tools.py index 0e38037e1..794991ac0 100644 --- a/src/backend/bisheng/database/models/gpts_tools.py +++ b/src/backend/bisheng/database/models/gpts_tools.py @@ -2,11 +2,13 @@ from datetime import datetime from enum import Enum from typing import Dict, List, Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import JSON, Column, DateTime, String, text, func from sqlmodel import Field, or_, select, Text, update +from bisheng.database.base import session_getter +from bisheng.database.constants import ToolPresetType +from bisheng.database.models.base import SQLModelSerializable + class AuthMethod(Enum): NO = 0 @@ -16,54 +18,57 @@ class AuthMethod(Enum): class AuthType(Enum): BASIC = "basic" BEARER = "bearer" - CUSTOM= "custom" + CUSTOM = "custom" class GptsToolsBase(SQLModelSerializable): name: str = Field(sa_column=Column(String(length=125), index=True)) - logo: Optional[str] = Field(sa_column=Column(String(length=512), index=False)) - desc: Optional[str] = Field(sa_column=Column(String(length=2048), index=False)) + logo: Optional[str] = Field(default=None, sa_column=Column(String(length=512), index=False)) + desc: Optional[str] = Field(default=None, sa_column=Column(String(length=2048), index=False)) tool_key: str = Field(sa_column=Column(String(length=125), index=False)) type: int = Field(default=0, description='所属类别的ID') - is_preset: bool = Field(default=True) + is_preset: int = Field(default=ToolPresetType.API.value, description="工具的类别,历史原因字段就不改名了") is_delete: int = Field(default=0, description='1 表示逻辑删除') - api_params: Optional[List[Dict]] = Field(sa_column=Column(JSON), description='用来存储api参数等信息') - user_id: Optional[int] = Field(index=True, description='创建用户ID, null表示系统创建') - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + api_params: Optional[List[Dict]] = Field(default=None, sa_column=Column(JSON), description='用来存储api参数等信息') + user_id: Optional[int] = Field(default=None, index=True, description='创建用户ID, null表示系统创建') + create_time: Optional[datetime] = Field(default=None, + sa_column=Column(DateTime, nullable=False, + server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, + sa_column=Column(DateTime, + nullable=False, + server_default=text('CURRENT_TIMESTAMP'), + onupdate=text('CURRENT_TIMESTAMP'))) class GptsToolsTypeBase(SQLModelSerializable): - id: Optional[int] = Field(index=True, primary_key=True) - name: str = Field(default='', index=True, description="工具类别名字") + id: Optional[int] = Field(default=None, index=True, primary_key=True) + name: str = Field(default='', sa_column=Column(String(length=1024), index=True), description="工具类别名字") logo: Optional[str] = Field(default='', description="工具类别的logo文件地址") extra: Optional[str] = Field(default='', sa_column=Column(String(length=2048)), description="工具类别的配置信息,用来存储工具类别所需的配置信息") description: str = Field(default='', description="工具类别的描述") server_host: Optional[str] = Field(default='', description="自定义工具的访问根地址,必须以http或者https开头") auth_method: Optional[int] = Field(default=0, description="工具类别的鉴权方式") - api_key: Optional[str] = Field(default='', description="工具鉴权的api_key",sa_column=Column(String(length=2048)),max_length=1000) + api_key: Optional[str] = Field(default='', description="工具鉴权的api_key", sa_column=Column(String(length=2048)), + max_length=1000) auth_type: Optional[str] = Field(default=AuthType.BASIC.value, description="工具鉴权的鉴权方式") - is_preset: Optional[int] = Field(default=0, description="是否是预置工具类别") - user_id: Optional[int] = Field(index=True, description='创建用户ID, null表示系统创建') + is_preset: Optional[int] = Field(default=ToolPresetType.API.value, description="工具的类别,历史原因字段就不改名了") + user_id: Optional[int] = Field(default=None, index=True, description='创建用户ID, null表示系统创建') is_delete: int = Field(default=0, description='1 表示逻辑删除') - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, + sa_column=Column(DateTime, nullable=False, + server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, + sa_column=Column(DateTime, + nullable=False, + server_default=text('CURRENT_TIMESTAMP'), + onupdate=text('CURRENT_TIMESTAMP'))) class GptsTools(GptsToolsBase, table=True): __tablename__ = 't_gpts_tools' - extra: Optional[str] = Field(sa_column=Column(String(length=2048), index=False), + extra: Optional[str] = Field(default=None, sa_column=Column(String(length=2048), index=False), description='用来存储额外信息,比如参数需求等,包含 &initdb_conf_key 字段' '表示配置信息从系统配置里获取,多层级用.隔开') id: Optional[int] = Field(default=None, primary_key=True) @@ -77,7 +82,7 @@ class GptsToolsType(GptsToolsTypeBase, table=True): class GptsToolsTypeRead(GptsToolsTypeBase): openapi_schema: Optional[str] = Field(default="", description="工具类别的schema内容,符合openapi规范的数据") - children: Optional[List[GptsTools]] = Field(default=[], description="工具类别下的工具列表") + children: Optional[List[GptsTools]] = Field(default_factory=list, description="工具类别下的工具列表") parameter_name: Optional[str] = Field(default="", description="自定义请求头参数名") api_location: Optional[str] = Field(default="", description="自定义请求头参数位置 header or query") @@ -110,11 +115,26 @@ class GptsToolsDao(GptsToolsBase): session.refresh(data) return data + @classmethod + def update_tool_list(cls, data: List[GptsTools]) -> List[GptsTools]: + with session_getter() as session: + for one in data: + session.add(one) + session.commit() + return data + @classmethod def delete_tool(cls, data: GptsTools) -> GptsTools: data.is_delete = 1 return cls.update_tools(data) + @classmethod + def delete_tool_by_ids(cls, tool_ids: List[int]) -> None: + with session_getter() as session: + statement = update(GptsTools).where(GptsTools.id.in_(tool_ids)).values(is_delete=1) + session.exec(statement) + session.commit() + @classmethod def get_one_tool(cls, tool_id: int) -> GptsTools: with session_getter() as session: @@ -135,7 +155,7 @@ class GptsToolsDao(GptsToolsBase): with session_getter() as session: statement = select(GptsTools).where( or_(GptsTools.user_id == user_id, - GptsTools.is_preset == 1)).where(GptsTools.is_delete == 0) + GptsTools.is_preset == ToolPresetType.PRESET.value)).where(GptsTools.is_delete == 0) if page and page_size: statement = statement.offset((page - 1) * page_size).limit(page_size) statement = statement.order_by(GptsTools.create_time.desc()) @@ -171,12 +191,14 @@ class GptsToolsDao(GptsToolsBase): 获得所有的预置工具类别 """ with session_getter() as session: - statement = select(GptsToolsType).where(GptsToolsType.is_preset == 1, GptsToolsType.is_delete == 0) + statement = select(GptsToolsType).where(GptsToolsType.is_preset == ToolPresetType.PRESET.value, + GptsToolsType.is_delete == 0) + statement = statement.order_by(GptsToolsType.update_time.desc()) return session.exec(statement).all() @classmethod - def get_user_tool_type(cls, user_id: int, extra_tool_type_ids: List[int], include_preset: bool = True) \ - -> List[GptsToolsType]: + def get_user_tool_type(cls, user_id: int, extra_tool_type_ids: List[int] = None, include_preset: bool = True, + is_preset: ToolPresetType = None) -> List[GptsToolsType]: """ 获取用户可见的所有工具类别 """ @@ -190,8 +212,12 @@ class GptsToolsDao(GptsToolsBase): else: filters.append(GptsToolsType.user_id == user_id) if include_preset: - filters.append(GptsToolsType.is_preset == 1) + filters.append(GptsToolsType.is_preset == ToolPresetType.PRESET.value) + if is_preset is not None: + statement = statement.where(GptsToolsType.is_preset == is_preset.value) statement = statement.where(or_(*filters)) + statement = statement.order_by(func.field(GptsToolsType.is_preset, + ToolPresetType.PRESET.value).desc() ,GptsToolsType.update_time.desc()) with session_getter() as session: return session.exec(statement).all() @@ -204,8 +230,8 @@ class GptsToolsDao(GptsToolsBase): statement = select(GptsToolsType).where(GptsToolsType.is_delete == 0) count_statement = select(func.count(GptsToolsType.id)).where(GptsToolsType.is_delete == 0) if not include_preset: - statement = statement.where(GptsToolsType.is_preset == 0) - count_statement = count_statement.where(GptsToolsType.is_preset == 0) + statement = statement.where(GptsToolsType.is_preset != ToolPresetType.PRESET.value) + count_statement = count_statement.where(GptsToolsType.is_preset != ToolPresetType.PRESET.value) if tool_type_ids: statement = statement.where(GptsToolsType.id.in_(tool_type_ids)) @@ -261,13 +287,14 @@ class GptsToolsDao(GptsToolsBase): session.add(gpts_tool_type) session.commit() session.refresh(gpts_tool_type) - with session_getter() as session: - # 插入工具列表 - for one in children: - one.type = gpts_tool_type.id - one.tool_key = cls.get_tool_key(gpts_tool_type.id, one.tool_key) - session.add_all(children) - session.commit() + if children: + with session_getter() as session: + # 插入工具列表 + for one in children: + one.type = gpts_tool_type.id + one.tool_key = cls.get_tool_key(gpts_tool_type.id, one.tool_key) + session.add_all(children) + session.commit() res = GptsToolsTypeRead(**gpts_tool_type.model_dump(), children=children) return res @@ -313,13 +340,13 @@ class GptsToolsDao(GptsToolsBase): session.exec( update(GptsToolsType).filter( GptsToolsType.id == tool_type_id, - GptsToolsType.is_preset == 0, + GptsToolsType.is_preset != ToolPresetType.PRESET.value, ).values(is_delete=1) ) session.exec( update(GptsTools).filter( GptsTools.type == tool_type_id, - GptsToolsType.is_preset == False + GptsToolsType.is_preset != ToolPresetType.PRESET.value ).values(is_delete=1) ) session.commit() diff --git a/src/backend/bisheng/database/models/group.py b/src/backend/bisheng/database/models/group.py index 926628ef2..bb7e3f79f 100644 --- a/src/backend/bisheng/database/models/group.py +++ b/src/backend/bisheng/database/models/group.py @@ -12,12 +12,12 @@ DefaultGroup = 2 class GroupBase(SQLModelSerializable): group_name: str = Field(index=False, description='前端展示名称', unique=True) - remark: Optional[str] = Field(index=False) - create_user: Optional[int] = Field(index=True, description="创建用户的ID") - update_user: Optional[int] = Field(description="更新用户的ID") - create_time: Optional[datetime] = Field(sa_column=Column( + remark: Optional[str] = Field(default=None, index=False) + create_user: Optional[int] = Field(default=None, index=True, description="创建用户的ID") + update_user: Optional[int] = Field(default=None, description="更新用户的ID") + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( + update_time: Optional[datetime] = Field(default=None, sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), @@ -30,17 +30,17 @@ class Group(GroupBase, table=True): class GroupRead(GroupBase): - id: Optional[int] - group_admins: Optional[List[Dict]] + id: Optional[int] = None + group_admins: Optional[List[Dict]] = None class GroupUpdate(GroupBase): - role_name: Optional[str] - remark: Optional[str] + role_name: Optional[str] = None + remark: Optional[str] = None class GroupCreate(GroupBase): - group_admins: Optional[List[int]] + group_admins: Optional[List[int]] = None class GroupDao(GroupBase): diff --git a/src/backend/bisheng/database/models/group_resource.py b/src/backend/bisheng/database/models/group_resource.py index 5a2902690..4c014ed93 100644 --- a/src/backend/bisheng/database/models/group_resource.py +++ b/src/backend/bisheng/database/models/group_resource.py @@ -2,31 +2,29 @@ from datetime import datetime from enum import Enum from typing import Dict, List, Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import Column, DateTime, text, delete from sqlmodel import Field, select +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable + class ResourceTypeEnum(Enum): KNOWLEDGE = 1 FLOW = 2 ASSISTANT = 3 GPTS_TOOL = 4 - WORK_FLOW= 5 + WORK_FLOW = 5 class GroupResourceBase(SQLModelSerializable): group_id: str = Field(index=True) third_id: str = Field(index=False) type: int = Field(index=False, description='资源类别 1:知识库 2:技能 3:助手 4:工具 5:工作流') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) class GroupResource(GroupResourceBase, table=True): @@ -34,13 +32,13 @@ class GroupResource(GroupResourceBase, table=True): class GroupResourceRead(GroupResourceBase): - id: Optional[int] - group_admins: Optional[List[Dict]] + id: Optional[int] = None + group_admins: Optional[List[Dict]] = None class GroupResourceUpdate(GroupResourceBase): - role_name: Optional[str] - remark: Optional[str] + role_name: Optional[str] = None + remark: Optional[str] = None class GroupResourceCreate(GroupResourceBase): diff --git a/src/backend/bisheng/database/models/knowledge.py b/src/backend/bisheng/database/models/knowledge.py index 8913b0072..1efe07550 100644 --- a/src/backend/bisheng/database/models/knowledge.py +++ b/src/backend/bisheng/database/models/knowledge.py @@ -8,8 +8,8 @@ from bisheng.database.models.knowledge_file import KnowledgeFile, KnowledgeFileD from bisheng.database.models.role_access import AccessType, RoleAccessDao from bisheng.database.models.user import UserDao from bisheng.database.models.user_role import UserRoleDao -from pydantic import BaseModel -from sqlmodel import Column, DateTime, Field, delete, func, or_, select, text, update +from pydantic import BaseModel, field_validator +from sqlmodel import Column, DateTime, Field, delete, func, or_, select, text, update, CHAR from sqlmodel.sql.expression import Select, SelectOfScalar @@ -20,21 +20,25 @@ class KnowledgeTypeEnum(Enum): class KnowledgeBase(SQLModelSerializable): - user_id: Optional[int] = Field(index=True) + user_id: Optional[int] = Field(default=None, index=True) name: str = Field(index=True, min_length=1, max_length=30, description='知识库名, 最少一个字符,最多30个字符') type: int = Field(index=False, default=0, description='0 为普通知识库,1 为QA知识库') - description: Optional[str] = Field(index=True) - model: Optional[str] = Field(index=False) - collection_name: Optional[str] = Field(index=False) - index_name: Optional[str] = Field(index=False) + description: Optional[str] = Field(default=None, index=True) + model: Optional[str] = Field(default=None, index=False) + collection_name: Optional[str] = Field(default=None, index=False) + index_name: Optional[str] = Field(default=None, index=False) state: Optional[int] = Field(index=False, default=1, description='0 为未发布,1 为已发布, 2 为复制中') - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=True, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) + + @field_validator('model', mode='before') + @classmethod + def convert_model(cls, v: Any) -> str: + if isinstance(v, int): + v = str(v) + return v class Knowledge(KnowledgeBase, table=True): @@ -43,14 +47,14 @@ class Knowledge(KnowledgeBase, table=True): class KnowledgeRead(KnowledgeBase): id: int - user_name: Optional[str] - copiable: Optional[bool] + user_name: Optional[str] = None + copiable: Optional[bool] = None class KnowledgeUpdate(BaseModel): knowledge_id: int - name: Optional[str] - description: Optional[str] + name: Optional[str] = None + description: Optional[str] = None class KnowledgeCreate(KnowledgeBase): diff --git a/src/backend/bisheng/database/models/knowledge_file.py b/src/backend/bisheng/database/models/knowledge_file.py index 9455e2cee..afa3ee406 100644 --- a/src/backend/bisheng/database/models/knowledge_file.py +++ b/src/backend/bisheng/database/models/knowledge_file.py @@ -3,18 +3,26 @@ from datetime import datetime from enum import Enum from typing import List, Optional +# if TYPE_CHECKING: +from pydantic import field_validator +from sqlalchemy import JSON, Column, DateTime, String, or_, text, Text +from sqlmodel import Field, delete, func, select + from bisheng.database.base import session_getter from bisheng.database.models.base import SQLModelSerializable -# if TYPE_CHECKING: -from pydantic import validator -from sqlalchemy import JSON, Column, DateTime, String, or_, text -from sqlmodel import Field, delete, func, select class KnowledgeFileStatus(Enum): - PROCESSING = 1 - SUCCESS = 2 - FAILED = 3 + PROCESSING = 1 # 处理中 + SUCCESS = 2 # 成功 + FAILED = 3 # 解析失败 + + +class QAStatus(Enum): + DISABLED = 0 # 用户手动关闭QA + ENABLED = 1 # 启用成功 + PROCESSING = 2 # 处理中 + FAILED = 3 # QA插入向量库失败 class ParseType(Enum): @@ -23,10 +31,10 @@ class ParseType(Enum): class KnowledgeFileBase(SQLModelSerializable): - user_id: Optional[int] = Field(index=True) + user_id: Optional[int] = Field(default=None, index=True) knowledge_id: int = Field(index=True) file_name: str = Field(index=True) - md5: Optional[str] = Field(index=False) + md5: Optional[str] = Field(default=None, index=False) parse_type: Optional[str] = Field(default=ParseType.LOCAL.value, index=False, description='采用什么模式解析的文件') @@ -35,37 +43,33 @@ class KnowledgeFileBase(SQLModelSerializable): status: Optional[int] = Field(default=KnowledgeFileStatus.PROCESSING.value, index=False, description='1: 解析中;2: 解析成功;3: 解析失败') - object_name: Optional[str] = Field(index=False, description='文件在minio存储的对象名称') - extra_meta: Optional[str] = Field(index=False) + object_name: Optional[str] = Field(default=None, index=False, description='文件在minio存储的对象名称') + extra_meta: Optional[str] = Field(default=None, index=False) remark: Optional[str] = Field(default='', sa_column=Column(String(length=512))) - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) class QAKnowledgeBase(SQLModelSerializable): - user_id: Optional[int] = Field(index=True) + user_id: Optional[int] = Field(default=None, index=True) knowledge_id: int = Field(index=True) questions: List[str] = Field(index=False) answers: str = Field(index=False) - source: Optional[int] = Field(index=False, description='0: 未知 1: 手动;2: 审计, 3: api') - status: Optional[int] = Field(index=False, description='1: 解析中;2: 解析成功;3: 解析失败') - extra_meta: Optional[str] = Field(index=False) - remark: Optional[str] = Field(sa_column=Column(String(length=512))) - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + source: Optional[int] = Field(default=0, index=False, description='0: 未知 1: 手动;2: 审计, 3: api, 4: 批量导入') + status: Optional[int] = Field(default=1, index=False, + description='1: 开启;0: 关闭,用户手动关闭;2: 处理中;3:插入失败') + extra_meta: Optional[str] = Field(default=None, index=False) + remark: Optional[str] = Field(default=None, sa_column=Column(String(length=512))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) - @validator('questions') - def validate_json(v): + @field_validator('questions') + @classmethod + def validate_json(cls, v): # dict_keys(['description', 'name', 'id', 'data']) if not v: return v @@ -74,8 +78,9 @@ class QAKnowledgeBase(SQLModelSerializable): return v - @validator('answers') - def validate_answer(v): + @field_validator('answers') + @classmethod + def validate_answer(cls, v): # dict_keys(['description', 'name', 'id', 'data']) if not v: return v @@ -92,7 +97,7 @@ class KnowledgeFile(KnowledgeFileBase, table=True): class QAKnowledge(QAKnowledgeBase, table=True): id: Optional[int] = Field(default=None, primary_key=True) questions: Optional[List[str]] = Field(default=None, sa_column=Column(JSON)) - answers: Optional[str] = Field(default=None, sa_column=Column(String(length=2048))) + answers: Optional[str] = Field(default=None, sa_column=Column(Text)) class KnowledgeFileRead(KnowledgeFileBase): @@ -105,8 +110,8 @@ class KnowledgeFileCreate(KnowledgeFileBase): class QAKnowledgeUpsert(QAKnowledgeBase): """支持修改""" - id: Optional[int] - answers: Optional[List[str]] + id: Optional[int] = None + answers: Optional[List[str] | str] = None class KnowledgeFileDao(KnowledgeFileBase): @@ -238,13 +243,15 @@ class QAKnoweldgeDao(QAKnowledgeBase): return session.exec(select(QAKnowledge).where(QAKnowledge.id == qa_id)).first() @classmethod - def get_qa_knowledge_by_name(cls, question: List[str], knowledge_id: int) -> QAKnowledge: + def get_qa_knowledge_by_name(cls, question: List[str], knowledge_id: int, exclude_id: int = None) -> QAKnowledge: with session_getter() as session: group_filters = [] for one in question: group_filters.append(func.json_contains(QAKnowledge.questions, json.dumps(one))) statement = select(QAKnowledge).where( or_(*group_filters)).where(QAKnowledge.knowledge_id == knowledge_id) + if exclude_id: + statement = statement.where(QAKnowledge.id != exclude_id) return session.exec(statement).first() @classmethod @@ -252,7 +259,7 @@ class QAKnoweldgeDao(QAKnowledgeBase): if qa_knowledge.id is None: raise ValueError('id不能为空') with session_getter() as session: - session.add(QAKnowledge.validate(qa_knowledge)) + session.add(qa_knowledge) session.commit() session.refresh(qa_knowledge) return qa_knowledge @@ -275,12 +282,25 @@ class QAKnoweldgeDao(QAKnowledgeBase): @classmethod def insert_qa(cls, qa_knowledge: QAKnowledgeUpsert): with session_getter() as session: - qa = QAKnowledge.validate(qa_knowledge) + qa = QAKnowledge.model_validate(qa_knowledge) session.add(qa) session.commit() session.refresh(qa) return qa + @classmethod + def batch_insert_qa(cls, qa_knowledges: List[QAKnowledgeUpsert]) -> List[QAKnowledge]: + with session_getter() as session: + qas = [] + for qa_knowledge in qa_knowledges: + qa = QAKnowledge.model_validate(qa_knowledge) + qas.append(qa) + session.add_all(qas) + session.commit() + for qa in qas: + session.refresh(qa) + return qas + @classmethod def total_count(cls, sql): with session_getter() as session: diff --git a/src/backend/bisheng/database/models/llm_server.py b/src/backend/bisheng/database/models/llm_server.py index b14b4f46a..e07a239e5 100644 --- a/src/backend/bisheng/database/models/llm_server.py +++ b/src/backend/bisheng/database/models/llm_server.py @@ -2,11 +2,12 @@ from datetime import datetime from enum import Enum from typing import Dict, List, Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import CHAR, JSON, Column, DateTime, Text, UniqueConstraint, delete, text, update from sqlmodel import Field, select +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable + # 服务提供方枚举 class LLMServerType(Enum): @@ -24,6 +25,10 @@ class LLMServerType(Enum): DEEPSEEK = 'deepseek' SPARK = 'spark' # 讯飞星火大模型 BISHENG_RT = 'bisheng_rt' + TENCENT = 'tencent' # 腾讯云 + MOONSHOT = 'moonshot' # 月之暗面的kimi + VOLCENGINE = 'volcengine' # 火山引擎的大模型 + SILICON = 'silicon' # 硅基流动 # 模型类型枚举 @@ -39,46 +44,42 @@ class LLMServerBase(SQLModelSerializable): type: str = Field(sa_column=Column(CHAR(20)), description='服务提供方类型') limit_flag: bool = Field(default=False, description='是否开启每日调用次数限制') limit: int = Field(default=0, description='每日调用次数限制') - config: Optional[Dict] = Field(sa_column=Column(JSON), description='服务提供方公共配置') + config: Optional[Dict] = Field(default=None, sa_column=Column(JSON), description='服务提供方公共配置') user_id: int = Field(default=0, description='创建人ID') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) class LLMModelBase(SQLModelSerializable): - server_id: Optional[int] = Field(nullable=False, index=True, description='服务ID') + server_id: Optional[int] = Field(default=None, nullable=False, index=True, description='服务ID') name: str = Field(default='', description='模型展示名') description: Optional[str] = Field(default='', sa_column=Column(Text), description='模型描述') model_name: str = Field(default='', description='模型名称,实例化组件时用的参数') model_type: str = Field(sa_column=Column(CHAR(20)), description='模型类型') - config: Optional[Dict] = Field(sa_column=Column(JSON), description='服务提供方公共配置') + config: Optional[Dict] = Field(default=None, sa_column=Column(JSON), description='服务提供方公共配置') status: int = Field(default=2, description='模型状态。0:正常,1:异常, 2: 未知') remark: Optional[str] = Field(default='', sa_column=Column(Text), description='异常原因') online: bool = Field(default=True, description='是否在线') user_id: int = Field(default=0, description='创建人ID') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) class LLMServer(LLMServerBase, table=True): __tablename__ = 'llm_server' - id: Optional[int] = Field(nullable=False, primary_key=True, description='服务唯一ID') + id: Optional[int] = Field(default=None, nullable=False, primary_key=True, description='服务唯一ID') class LLMModel(LLMModelBase, table=True): __tablename__ = 'llm_model' - __table_args__ = (UniqueConstraint('server_id', 'model_name', name='server_model_uniq'), ) + __table_args__ = (UniqueConstraint('server_id', 'model_name', name='server_model_uniq'),) - id: Optional[int] = Field(nullable=False, primary_key=True, description='模型唯一ID') + id: Optional[int] = Field(default=None, nullable=False, primary_key=True, description='模型唯一ID') class LLMDao: diff --git a/src/backend/bisheng/database/models/mark_app_user.py b/src/backend/bisheng/database/models/mark_app_user.py index 6c5e52323..3d1295d32 100644 --- a/src/backend/bisheng/database/models/mark_app_user.py +++ b/src/backend/bisheng/database/models/mark_app_user.py @@ -1,12 +1,12 @@ from datetime import datetime from typing import List, Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable # if TYPE_CHECKING: from sqlalchemy import Column, DateTime, text from sqlmodel import Field +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable class MarkAppUserBase(SQLModelSerializable): @@ -15,20 +15,16 @@ class MarkAppUserBase(SQLModelSerializable): task_id: int = Field(index=True) create_id: int = Field(index=True) status: Optional[int] = Field(index=False, default=1) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) - create_time: Optional[datetime] = Field(sa_column=Column( + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=True, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) -class MarkAppUser(MarkAppUserBase,table=True): +class MarkAppUser(MarkAppUserBase, table=True): id: Optional[int] = Field(default=None, primary_key=True) - class MarkAppUserDao(MarkAppUserBase): @classmethod @@ -36,4 +32,4 @@ class MarkAppUserDao(MarkAppUserBase): with session_getter() as session: session.add_all(task_info) session.commit() - return task_info + return task_info diff --git a/src/backend/bisheng/database/models/mark_record.py b/src/backend/bisheng/database/models/mark_record.py index dedf79430..d8738f0c6 100644 --- a/src/backend/bisheng/database/models/mark_record.py +++ b/src/backend/bisheng/database/models/mark_record.py @@ -1,16 +1,13 @@ - from datetime import datetime from enum import Enum -from typing import Dict, List, Optional, Tuple, Union +from typing import List, Optional + +# if TYPE_CHECKING: +from sqlalchemy import Column, DateTime, delete, text +from sqlmodel import Field, select from bisheng.database.base import session_getter from bisheng.database.models.base import SQLModelSerializable -from bisheng.database.models.role_access import AccessType, RoleAccess, RoleAccessDao -from bisheng.database.models.user_role import UserRoleDao -# if TYPE_CHECKING: -from pydantic import validator -from sqlalchemy import Column, DateTime, String, and_, delete, func, or_, text -from sqlmodel import JSON, Field, select, update class MarkRecordStatus(Enum): @@ -27,23 +24,18 @@ class MarkRecordBase(SQLModelSerializable): task_id: int = Field(index=True) session_id: str = Field(index=True) status: int = Field(index=False, default=1) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) - create_time: Optional[datetime] = Field(sa_column=Column( + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=True, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) -class MarkRecord(MarkRecordBase,table=True): +class MarkRecord(MarkRecordBase, table=True): id: Optional[int] = Field(default=None, primary_key=True) - class MarkRecordDao(MarkRecordBase): - @classmethod def update_record(cls, record_info: MarkRecord) -> MarkRecord: with session_getter() as session: @@ -53,26 +45,24 @@ class MarkRecordDao(MarkRecordBase): return record_info @classmethod - def get_prev_task(cls,user_id:int,task_id:int): + def get_prev_task(cls, user_id: int, task_id: int): with session_getter() as session: - statement = select(MarkRecord).where(MarkRecord.create_id==user_id).where(MarkRecord.task_id==task_id).order_by(MarkRecord.id) + statement = select(MarkRecord).where(MarkRecord.create_id == user_id).where( + MarkRecord.task_id == task_id).order_by(MarkRecord.id) return session.exec(statement).all() - - @classmethod def create_record(cls, record_info: MarkRecord) -> MarkRecord: with session_getter() as session: session.add(record_info) session.commit() session.refresh(record_info) - return record_info - + return record_info @classmethod - def del_record(cls,task_id:int): + def del_record(cls, task_id: int): with session_getter() as session: - st = delete(MarkRecord).where(MarkRecord.task_id==task_id) + st = delete(MarkRecord).where(MarkRecord.task_id == task_id) session.exec(st) session.commit() return @@ -86,29 +76,30 @@ class MarkRecordDao(MarkRecordBase): return @classmethod - def get_list_by_taskid(cls,task_id:int): + def get_list_by_taskid(cls, task_id: int): with session_getter() as session: - statement = select(MarkRecord).where(MarkRecord.task_id==task_id) + statement = select(MarkRecord).where(MarkRecord.task_id == task_id) return session.exec(statement).all() @classmethod - def get_count(cls,task_id:int): + def get_count(cls, task_id: int): with session_getter() as session: - sql = text("select create_user,count(*) as user_count,create_id from markrecord where task_id=:task_id group by create_id") - query = session.execute(sql,{"task_id":task_id}).fetchall() + sql = text( + "select create_user,count(*) as user_count,create_id from markrecord where task_id=:task_id group by create_id") + query = session.execute(sql, {"task_id": task_id}).fetchall() return query - @classmethod - def get_record(cls,task_id:int,session_id:str) -> MarkRecord: + def get_record(cls, task_id: int, session_id: str) -> MarkRecord: with session_getter() as session: - statement = select(MarkRecord).where(MarkRecord.task_id==task_id).where(MarkRecord.session_id==session_id) + statement = select(MarkRecord).where(MarkRecord.task_id == task_id).where( + MarkRecord.session_id == session_id) return session.exec(statement).first() - @classmethod - def filter_records(cls, task_id: int, chat_ids: list[str] = None, status: int = None, mark_user: int = None) -> List[MarkRecord]: + def filter_records(cls, task_id: int, chat_ids: list[str] = None, status: int = None, mark_user: int = None) -> \ + List[MarkRecord]: statement = select(MarkRecord).where(MarkRecord.task_id == task_id) if chat_ids: statement = statement.where(MarkRecord.session_id.in_(chat_ids)) diff --git a/src/backend/bisheng/database/models/mark_task.py b/src/backend/bisheng/database/models/mark_task.py index 4e09a8804..3263e5f42 100644 --- a/src/backend/bisheng/database/models/mark_task.py +++ b/src/backend/bisheng/database/models/mark_task.py @@ -2,12 +2,13 @@ from datetime import datetime from enum import Enum from typing import List, Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable # if TYPE_CHECKING: from sqlalchemy import Column, DateTime, and_, delete, func, or_, text from sqlmodel import Field, select, update +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable + class MarkTaskStatus(Enum): DEFAULT = 1 @@ -20,14 +21,11 @@ class MarkTaskBase(SQLModelSerializable): create_id: int = Field(index=True) app_id: str = Field(index=False, max_length=2048) process_users: str = Field(index=False) # 23,2323 - mark_user: Optional[str] = Field(index=True, nullable=True) + mark_user: Optional[str] = Field(default=None, index=True, nullable=True) status: Optional[int] = Field(index=False, default=1) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) - create_time: Optional[datetime] = Field(sa_column=Column( + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=True, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) @@ -36,8 +34,8 @@ class MarkTask(MarkTaskBase, table=True): class MarkTaskRead(MarkTaskBase): - id: Optional[int] - mark_process: Optional[List[str]] + id: Optional[int] = None + mark_process: Optional[List[str]] = None class MarkTaskDao(MarkTaskBase): @@ -84,9 +82,9 @@ class MarkTaskDao(MarkTaskBase): @classmethod def get_all_task( - cls, - page_size: int = 10, - page_num: int = 1, + cls, + page_size: int = 10, + page_num: int = 1, ): with session_getter() as session: statement = select(MarkTask) @@ -98,12 +96,12 @@ class MarkTaskDao(MarkTaskBase): @classmethod def get_task_list( - cls, - status: int, - create_id: Optional[int], - user_id: Optional[int], - page_size: int = 10, - page_num: int = 1, + cls, + status: int, + create_id: Optional[int], + user_id: Optional[int], + page_size: int = 10, + page_num: int = 1, ): with session_getter() as session: statement = select(MarkTask) diff --git a/src/backend/bisheng/database/models/message.py b/src/backend/bisheng/database/models/message.py index 7f5b3f4d7..65c37aed5 100644 --- a/src/backend/bisheng/database/models/message.py +++ b/src/backend/bisheng/database/models/message.py @@ -18,35 +18,32 @@ class LikedType(Enum): class MessageBase(SQLModelSerializable): is_bot: bool = Field(index=False, description='聊天角色') - source: Optional[int] = Field(index=False, description='是否支持溯源') + source: Optional[int] = Field(default=None, index=False, description='是否支持溯源') mark_status: Optional[int] = Field(index=False, default=1, description='标记状态') - mark_user: Optional[int] = Field(index=False, description='标记用户') - mark_user_name: Optional[str] = Field(index=False, description='标记用户') - message: Optional[str] = Field(sa_column=Column(Text), description='聊天消息') - extra: Optional[str] = Field(sa_column=Column(Text), description='连接信息等') + mark_user: Optional[int] = Field(default=None, index=False, description='标记用户') + mark_user_name: Optional[str] = Field(default=None, index=False, description='标记用户') + message: Optional[str] = Field(default=None, sa_column=Column(Text), description='聊天消息') + extra: Optional[str] = Field(default=None, sa_column=Column(Text), description='连接信息等') type: str = Field(index=False, description='消息类型') category: str = Field(index=False, max_length=32, description='消息类别, question等') flow_id: str = Field(index=True, description='对应的技能id') - chat_id: Optional[str] = Field(index=True, description='chat_id, 前端生成') - user_id: Optional[str] = Field(index=True, description='用户id') + chat_id: Optional[str] = Field(default=None, index=True, description='chat_id, 前端生成') + user_id: Optional[int] = Field(default=None, index=True, description='用户id') liked: Optional[int] = Field(index=False, default=0, description='用户是否喜欢 0未评价/1 喜欢/2 不喜欢') solved: Optional[int] = Field(index=False, default=0, description='用户是否喜欢 0未评价/1 解决/2 未解决') copied: Optional[int] = Field(index=False, default=0, description='用户是否复制 0:未复制 1:已复制') sensitive_status: Optional[int] = Field(index=False, default=1, description='敏感词状态 1:通过 2:违规') sender: Optional[str] = Field(index=False, default='', description='autogen 的发送方') receiver: Optional[Dict] = Field(index=False, default=None, description='autogen 的发送方') - intermediate_steps: Optional[str] = Field(sa_column=Column(Text), description='过程日志') - files: Optional[str] = Field(sa_column=Column(String(length=4096)), description='上传的文件等') - remark: Optional[str] = Field(sa_column=Column(String(length=4096)), + intermediate_steps: Optional[str] = Field(default=None, sa_column=Column(Text), description='过程日志') + files: Optional[str] = Field(default=None, sa_column=Column(String(length=4096)), description='上传的文件等') + remark: Optional[str] = Field(default=None, sa_column=Column(String(length=4096)), description='备注。break_answer: 中断的回复不作为history传给模型') - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - index=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'), + onupdate=text('CURRENT_TIMESTAMP'))) class ChatMessage(MessageBase, table=True): @@ -55,11 +52,11 @@ class ChatMessage(MessageBase, table=True): class ChatMessageRead(MessageBase): - id: Optional[int] + id: Optional[int] = None class ChatMessageQuery(BaseModel): - id: Optional[int] + id: Optional[int] = None flow_id: str chat_id: str @@ -257,6 +254,11 @@ class ChatMessageDao(MessageBase): with session_getter() as session: session.add_all(messages) session.commit() + ret = [] + for one in messages: + session.refresh(one) + ret.append(one) + return ret @classmethod def get_message_by_id(cls, message_id: int) -> Optional[ChatMessage]: diff --git a/src/backend/bisheng/database/models/model_deploy.py b/src/backend/bisheng/database/models/model_deploy.py index fac3364f0..40025ad68 100644 --- a/src/backend/bisheng/database/models/model_deploy.py +++ b/src/backend/bisheng/database/models/model_deploy.py @@ -1,28 +1,26 @@ from datetime import datetime from typing import Optional, List -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import Column, DateTime, String, UniqueConstraint, delete, text from sqlmodel import Field, select +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable + class ModelDeployBase(SQLModelSerializable): endpoint: str = Field(index=False, unique=False) server: str = Field(index=True) model: str = Field(index=False) - config: Optional[str] = Field(sa_column=Column(String(length=512))) - status: Optional[str] = Field(index=False) - remark: Optional[str] = Field(sa_column=Column(String(length=4096))) + config: Optional[str] = Field(default=None, sa_column=Column(String(length=512))) + status: Optional[str] = Field(default=None, index=False) + remark: Optional[str] = Field(default=None, sa_column=Column(String(length=4096))) - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - index=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'), + onupdate=text('CURRENT_TIMESTAMP'))) class ModelDeploy(ModelDeployBase, table=True): diff --git a/src/backend/bisheng/database/models/preset_train.py b/src/backend/bisheng/database/models/preset_train.py index 6f96b74f2..1f9ee242d 100644 --- a/src/backend/bisheng/database/models/preset_train.py +++ b/src/backend/bisheng/database/models/preset_train.py @@ -1,11 +1,12 @@ from datetime import datetime from typing import List, Optional, Tuple +from sqlalchemy import func +from sqlmodel import Column, DateTime, Field, select, text + from bisheng.database.base import session_getter from bisheng.database.models.base import SQLModelSerializable from bisheng.utils import generate_uuid -from sqlalchemy import func -from sqlmodel import Column, DateTime, Field, select, text # Finetune任务的预置训练集 @@ -16,12 +17,10 @@ class PresetTrainBase(SQLModelSerializable): user_id: str = Field(default='', index=True, description='创建人ID') user_name: str = Field(default='', index=True, description='创建人姓名') type: int = Field(default=0, index=True, description='0 文件 1 QA') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) class PresetTrain(PresetTrainBase, table=True): diff --git a/src/backend/bisheng/database/models/recall_chunk.py b/src/backend/bisheng/database/models/recall_chunk.py index 2e2d7efae..6046ab98e 100644 --- a/src/backend/bisheng/database/models/recall_chunk.py +++ b/src/backend/bisheng/database/models/recall_chunk.py @@ -1,27 +1,24 @@ from datetime import datetime from typing import Optional -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import Column, DateTime, Text, text from sqlmodel import Field +from bisheng.database.models.base import SQLModelSerializable + class RecallBase(SQLModelSerializable): - message_id: Optional[int] = Field(index=True, unique=False) + message_id: Optional[int] = Field(default=None, index=True, unique=False) chat_id: str = Field(index=False) keywords: str = Field(sa_column=Column(Text)) - chunk: Optional[str] = Field(sa_column=Column(Text)) - meta_data: Optional[str] = Field(sa_column=Column(Text)) - file_id: Optional[int] = Field(index=False) - - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - index=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + chunk: Optional[str] = Field(default=None, sa_column=Column(Text)) + meta_data: Optional[str] = Field(default=None, sa_column=Column(Text)) + file_id: Optional[int] = Field(default=None, index=False) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'), + onupdate=text('CURRENT_TIMESTAMP'))) class RecallChunk(RecallBase, table=True): @@ -29,13 +26,13 @@ class RecallChunk(RecallBase, table=True): class RecallChunkRead(RecallBase): - id: Optional[int] - score: Optional[int] + id: Optional[int] = None + score: Optional[int] = None class RecallChunkQuery(SQLModelSerializable): - id: Optional[int] - server: Optional[str] + id: Optional[int] = None + server: Optional[str] = None class RecallChunkCreate(RecallBase): diff --git a/src/backend/bisheng/database/models/report.py b/src/backend/bisheng/database/models/report.py index 42af7c2e0..455a5c0ee 100644 --- a/src/backend/bisheng/database/models/report.py +++ b/src/backend/bisheng/database/models/report.py @@ -1,29 +1,27 @@ from datetime import datetime from typing import Optional -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import Column, DateTime, text from sqlmodel import Field +from bisheng.database.models.base import SQLModelSerializable + class ReportBase(SQLModelSerializable): # ``` # 使用flow_id 按更新时间倒排获取最新模板的路径 # 会存储模板生成的最终报告的记录` flow_id: str = Field(index=False, description='技能名字') - file_name: Optional[str] = Field(index=False, description='生成报告名字') - template_name: Optional[str] = Field(index=False, description='报告模板数据存储路径') - version_key: Optional[str] = Field(index=True, unique=True, description='前端模板唯一key') + file_name: Optional[str] = Field(default=None, index=False, description='生成报告名字') + template_name: Optional[str] = Field(default=None, index=False, description='报告模板数据存储路径') + version_key: Optional[str] = Field(default=None, index=True, unique=True, description='前端模板唯一key') newversion_key: Optional[str] = Field(index=False, default=None, description='前端模板下一个key') - object_name: Optional[str] = Field(index=False, description='报告模板数据存储路径') + object_name: Optional[str] = Field(default=None, index=False, description='报告模板数据存储路径') del_yn: Optional[int] = Field(index=False, default=0, description='删除状态, 1表示删除') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) class Report(ReportBase, table=True): @@ -32,12 +30,12 @@ class Report(ReportBase, table=True): class ReportRead(ReportBase): - id: Optional[int] + id: Optional[int] = None class RoleUpdate(ReportBase): - role_name: Optional[str] - remark: Optional[str] + role_name: Optional[str] = None + remark: Optional[str] = None class RoleCreate(ReportBase): diff --git a/src/backend/bisheng/database/models/role.py b/src/backend/bisheng/database/models/role.py index 41327e7b2..a622cf680 100644 --- a/src/backend/bisheng/database/models/role.py +++ b/src/backend/bisheng/database/models/role.py @@ -1,26 +1,24 @@ from datetime import datetime from typing import List, Optional +from sqlalchemy import Column, DateTime, text, func, delete, and_, UniqueConstraint +from sqlmodel import Field, select + from bisheng.database.base import session_getter -from bisheng.database.constants import DefaultRole, AdminRole +from bisheng.database.constants import AdminRole from bisheng.database.models.base import SQLModelSerializable from bisheng.database.models.role_access import RoleAccess from bisheng.database.models.user_role import UserRole -from sqlalchemy import Column, DateTime, text, func, delete, and_, UniqueConstraint -from sqlmodel import Field, select class RoleBase(SQLModelSerializable): role_name: str = Field(index=False, description='前端展示名称') - group_id: Optional[int] = Field(index=True) - remark: Optional[str] = Field(index=False) - create_time: Optional[datetime] = Field(sa_column=Column( + group_id: Optional[int] = Field(default=None, index=True) + remark: Optional[str] = Field(default=None, index=False) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) class Role(RoleBase, table=True): @@ -29,12 +27,12 @@ class Role(RoleBase, table=True): class RoleRead(RoleBase): - id: Optional[int] + id: Optional[int] = None class RoleUpdate(RoleBase): - role_name: Optional[str] - remark: Optional[str] + role_name: Optional[str] = None + remark: Optional[str] = None class RoleCreate(RoleBase): @@ -115,6 +113,6 @@ class RoleDao(RoleBase): Role, and_(UserRole.role_id == Role.id, Role.group_id == group_id)).group_by(UserRole.id) all_user = session.exec(all_user).all() - session.exec(delete(UserRole).where(UserRole.id.in_([one.id for one in all_user]))) + session.exec(delete(UserRole).where(UserRole.id.in_([one.UserRole.id for one in all_user]))) session.exec(delete(Role).where(Role.group_id == group_id)) session.commit() diff --git a/src/backend/bisheng/database/models/role_access.py b/src/backend/bisheng/database/models/role_access.py index f44cad507..b576241a8 100644 --- a/src/backend/bisheng/database/models/role_access.py +++ b/src/backend/bisheng/database/models/role_access.py @@ -2,25 +2,23 @@ from datetime import datetime from enum import Enum from typing import List, Optional, Union -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from pydantic import BaseModel from sqlalchemy import Column, DateTime, text from sqlmodel import Field, select +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable + class RoleAccessBase(SQLModelSerializable): role_id: str = Field(index=True) third_id: str = Field(index=False) type: int = Field(index=False) - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - index=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'), + onupdate=text('CURRENT_TIMESTAMP'))) class RoleAccess(RoleAccessBase, table=True): @@ -28,7 +26,7 @@ class RoleAccess(RoleAccessBase, table=True): class RoleAccessRead(RoleAccessBase): - id: Optional[int] + id: Optional[int] = None class RoleAccessCreate(RoleAccessBase): @@ -46,8 +44,8 @@ class AccessType(Enum): ASSISTANT_WRITE = 6 GPTS_TOOL_READ = 7 GPTS_TOOL_WRITE = 8 - WORK_FLOW= 9 - WORK_FLOW_WRITE= 10 + WORK_FLOW = 9 + WORK_FLOW_WRITE = 10 WEB_MENU = 99 # 前端菜单栏权限限制 @@ -76,6 +74,7 @@ class RoleAccessDao(RoleAccessBase): return session.exec( select(RoleAccess).where(RoleAccess.role_id.in_(role_ids), RoleAccess.type.in_([x.value for x in access_type]))).all() + @classmethod def judge_role_access(cls, role_ids: List[int], third_id: str, access_type: AccessType) -> Optional[RoleAccess]: with session_getter() as session: diff --git a/src/backend/bisheng/database/models/server.py b/src/backend/bisheng/database/models/server.py index 41623277a..9dfb92932 100644 --- a/src/backend/bisheng/database/models/server.py +++ b/src/backend/bisheng/database/models/server.py @@ -1,24 +1,22 @@ from datetime import datetime from typing import Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import Column, DateTime, text from sqlmodel import Field, select +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable + class ServerBase(SQLModelSerializable): endpoint: str = Field(index=False, unique=True) sft_endpoint: str = Field(default='', index=False, description='Finetune服务地址') server: str = Field(index=True) - remark: Optional[str] = Field(index=False) - create_time: Optional[datetime] = Field(sa_column=Column( + remark: Optional[str] = Field(default=None, index=False) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) class Server(ServerBase, table=True): @@ -41,12 +39,12 @@ class ServerDao(ServerBase): class ServerRead(ServerBase): - id: Optional[int] + id: Optional[int] = None class ServerQuery(ServerBase): - id: Optional[int] - server: Optional[str] + id: Optional[int] = None + server: Optional[str] = None class ServerCreate(ServerBase): diff --git a/src/backend/bisheng/database/models/session.py b/src/backend/bisheng/database/models/session.py index 37bd77b5d..97529d5f2 100644 --- a/src/backend/bisheng/database/models/session.py +++ b/src/backend/bisheng/database/models/session.py @@ -1,12 +1,12 @@ from datetime import datetime from enum import Enum -from typing import List, Optional +from typing import Optional, List + +from sqlmodel import Field, Column, DateTime, text, select, func, update from bisheng.database.base import session_getter from bisheng.database.models.base import SQLModelSerializable from bisheng.database.models.flow import FlowType -from sqlmodel import Column, DateTime, Field, func, select, text, update - class SensitiveStatus(Enum): PASS = 1 # 通过 @@ -25,14 +25,11 @@ class MessageSessionBase(SQLModelSerializable): dislike: Optional[int] = Field(default=0, description='点踩的消息数量') copied: Optional[int] = Field(default=0, description='已复制的消息数量') sensitive_status: int = Field(default=SensitiveStatus.PASS.value, description='审查状态') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - index=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'), + onupdate=text('CURRENT_TIMESTAMP'))) class MessageSession(MessageSessionBase, table=True): @@ -51,8 +48,7 @@ class MessageSessionDao(MessageSessionBase): @classmethod def delete_session(cls, chat_id: str): - statement = update(MessageSession).where(MessageSession.chat_id == chat_id).values( - is_delete=True) + statement = update(MessageSession).where(MessageSession.chat_id == chat_id).values(is_delete=True) with session_getter() as session: session.exec(statement) session.commit() diff --git a/src/backend/bisheng/database/models/sft_model.py b/src/backend/bisheng/database/models/sft_model.py index 9da8deaeb..9191078e6 100644 --- a/src/backend/bisheng/database/models/sft_model.py +++ b/src/backend/bisheng/database/models/sft_model.py @@ -1,18 +1,19 @@ from datetime import datetime from typing import Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import Column, DateTime, delete, text, update from sqlmodel import Field, select +from bisheng.database.base import session_getter +from bisheng.database.models.base import SQLModelSerializable + # 可用于训练的model列表 class SftModelBase(SQLModelSerializable): id: int = Field(default=None, nullable=False, primary_key=True, description='唯一ID') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field(sa_column=Column( + update_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP'))) diff --git a/src/backend/bisheng/database/models/tag.py b/src/backend/bisheng/database/models/tag.py index bd219b2ac..907a8ca38 100644 --- a/src/backend/bisheng/database/models/tag.py +++ b/src/backend/bisheng/database/models/tag.py @@ -13,19 +13,17 @@ class TagBase(SQLModelSerializable): """ 标签表 """ - name: Optional[str] = Field(index=True, unique=True, description="标签名称") + name: Optional[str] = Field(default=None, index=True, unique=True, description="标签名称") user_id: int = Field(default=0, description='创建用户ID') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP')), description="创建时间") - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP')), description="更新时间") + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP')), + description="更新时间") class Tag(TagBase, table=True): - id: Optional[int] = Field(index=True, primary_key=True, description="标签唯一ID") + id: Optional[int] = Field(default=None, index=True, primary_key=True, description="标签唯一ID") class TagLinkBase(SQLModelSerializable): @@ -36,18 +34,16 @@ class TagLinkBase(SQLModelSerializable): resource_id: str = Field(description="资源唯一ID") resource_type: int = Field(description="资源类型") # 使用group_resource.ResourceTypeEnum枚举值 user_id: int = Field(default=0, description='创建用户ID') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP')), description="创建时间") - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP')), description="更新时间") + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP')), + description="更新时间") class TagLink(TagLinkBase, table=True): __table_args__ = (UniqueConstraint('resource_id', 'resource_type', 'tag_id', name='resource_tag_uniq'),) - id: Optional[int] = Field(index=True, primary_key=True, description="标签关联唯一ID") + id: Optional[int] = Field(default=None, index=True, primary_key=True, description="标签关联唯一ID") class TagDao(Tag): diff --git a/src/backend/bisheng/database/models/template.py b/src/backend/bisheng/database/models/template.py index f45415c35..cae93f817 100644 --- a/src/backend/bisheng/database/models/template.py +++ b/src/backend/bisheng/database/models/template.py @@ -1,10 +1,11 @@ from datetime import datetime from typing import Dict, Optional -from bisheng.database.models.base import SQLModelSerializable from sqlalchemy import JSON, Column, DateTime, text, String from sqlmodel import Field +from bisheng.database.models.base import SQLModelSerializable + class TemplateSkillBase(SQLModelSerializable): name: str = Field(index=True) @@ -13,15 +14,12 @@ class TemplateSkillBase(SQLModelSerializable): order_num: Optional[int] = Field(default=True) # 1 flow 5 assistant 10 workflow flow_type: Optional[int] = Field(default=1) - flow_id: Optional[str] = Field(index=False) - create_time: Optional[datetime] = Field(sa_column=Column( + flow_id: Optional[str] = Field(default=None, index=False) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) - guide_word: Optional[str] = Field(sa_column=Column(String(length=1000))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) + guide_word: Optional[str] = Field(default=None, sa_column=Column(String(length=1000))) class Template(TemplateSkillBase, table=True): diff --git a/src/backend/bisheng/database/models/user.py b/src/backend/bisheng/database/models/user.py index ae11ef235..41744c34c 100644 --- a/src/backend/bisheng/database/models/user.py +++ b/src/backend/bisheng/database/models/user.py @@ -1,33 +1,32 @@ from datetime import datetime from typing import List, Optional +from pydantic import field_validator +from sqlalchemy import Column, DateTime, func, text +from sqlmodel import Field, select + from bisheng.database.base import session_getter from bisheng.database.constants import AdminRole, DefaultRole from bisheng.database.models.base import SQLModelSerializable from bisheng.database.models.user_group import UserGroup from bisheng.database.models.user_role import UserRole -from pydantic import validator -from sqlalchemy import Column, DateTime, func, text -from sqlmodel import Field, select class UserBase(SQLModelSerializable): user_name: str = Field(index=True, unique=True) - email: Optional[str] = Field(index=True) - phone_number: Optional[str] = Field(index=True) - dept_id: Optional[str] = Field(index=True) - remark: Optional[str] = Field(index=False) - delete: int = Field(index=False, default=0) - create_time: Optional[datetime] = Field(sa_column=Column( + email: Optional[str] = Field(default=None, index=True) + phone_number: Optional[str] = Field(default=None, index=True) + dept_id: Optional[str] = Field(default=None, index=True) + remark: Optional[str] = Field(default=None, index=False) + delete: int = Field(default=0, index=False) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) - @validator('user_name') - def validate_str(v): + @field_validator('user_name') + @classmethod + def validate_str(cls, v): # dict_keys(['description', 'name', 'id', 'data']) if not v: raise ValueError('user_name 不能为空') @@ -37,35 +36,34 @@ class UserBase(SQLModelSerializable): class User(UserBase, table=True): user_id: Optional[int] = Field(default=None, primary_key=True) password: str = Field(index=False) - password_update_time: Optional[datetime] = Field(sa_column=Column( - DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP')), - description='密码最近的修改时间') + password_update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP')), description='密码最近的修改时间') class UserRead(UserBase): - user_id: Optional[int] - role: Optional[str] - access_token: Optional[str] - web_menu: Optional[List[str]] - admin_groups: Optional[List[int]] # 所管理的用户组ID列表 + user_id: Optional[int] = None + role: Optional[str] = None + access_token: Optional[str] = None + web_menu: Optional[List[str]] = None + admin_groups: Optional[List[int]] = None # 所管理的用户组ID列表 class UserQuery(UserBase): - user_id: Optional[int] - user_name: Optional[str] + user_id: Optional[int] = None + user_name: Optional[str] = None class UserLogin(UserBase): password: str user_name: str - captcha_key: Optional[str] - captcha: Optional[str] + captcha_key: Optional[str] = None + captcha: Optional[str] = None class UserCreate(UserBase): password: Optional[str] = Field(default='') - captcha_key: Optional[str] - captcha: Optional[str] + captcha_key: Optional[str] = None + captcha: Optional[str] = None class UserUpdate(SQLModelSerializable): diff --git a/src/backend/bisheng/database/models/user_group.py b/src/backend/bisheng/database/models/user_group.py index 9043ae38f..1a4ce871b 100644 --- a/src/backend/bisheng/database/models/user_group.py +++ b/src/backend/bisheng/database/models/user_group.py @@ -1,11 +1,11 @@ from datetime import datetime from typing import List, Optional +from sqlalchemy import Column, DateTime, delete, text +from sqlmodel import Field, select + from bisheng.database.base import session_getter from bisheng.database.models.base import SQLModelSerializable -from sqlalchemy import Column, DateTime, delete, text -from sqlmodel import Field, select, update - from bisheng.database.models.group import DefaultGroup @@ -13,14 +13,11 @@ class UserGroupBase(SQLModelSerializable): user_id: int = Field(index=True, description='用户id') group_id: int = Field(index=True, description='组id') is_group_admin: bool = Field(default=False, index=False, description='是否是组管理员') # 管理员不属于此用户组 - remark: Optional[str] = Field(index=False) - create_time: Optional[datetime] = Field(sa_column=Column( + remark: Optional[str] = Field(default=None, index=False) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) class UserGroup(UserGroupBase, table=True): @@ -28,14 +25,14 @@ class UserGroup(UserGroupBase, table=True): class UserGroupRead(UserGroupBase): - id: Optional[int] + id: Optional[int] = None class UserGroupUpdate(UserGroupBase): - user_id: Optional[int] - group_id: Optional[int] - is_group_admin: Optional[bool] - remark: Optional[str] + user_id: Optional[int] = None + group_id: Optional[int] = None + is_group_admin: Optional[bool] = None + remark: Optional[str] = None class UserGroupCreate(UserGroupBase): diff --git a/src/backend/bisheng/database/models/user_role.py b/src/backend/bisheng/database/models/user_role.py index efd910508..aa0208f2e 100644 --- a/src/backend/bisheng/database/models/user_role.py +++ b/src/backend/bisheng/database/models/user_role.py @@ -1,26 +1,23 @@ from datetime import datetime from typing import List, Optional -from bisheng.database.base import session_getter -from bisheng.database.models.base import SQLModelSerializable from pydantic import BaseModel from sqlalchemy import Column, DateTime, text, delete from sqlmodel import Field, select +from bisheng.database.base import session_getter from bisheng.database.constants import AdminRole +from bisheng.database.models.base import SQLModelSerializable class UserRoleBase(SQLModelSerializable): user_id: int = Field(index=True) role_id: int = Field(index=True) - create_time: Optional[datetime] = Field( - sa_column=Column(DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - index=True, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + create_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'), + onupdate=text('CURRENT_TIMESTAMP'))) class UserRole(UserRoleBase, table=True): @@ -28,7 +25,7 @@ class UserRole(UserRoleBase, table=True): class UserRoleRead(UserRoleBase): - id: Optional[int] + id: Optional[int] = None class UserRoleCreate(BaseModel): diff --git a/src/backend/bisheng/database/models/variable_value.py b/src/backend/bisheng/database/models/variable_value.py index d8fd3eadd..322f8a67a 100644 --- a/src/backend/bisheng/database/models/variable_value.py +++ b/src/backend/bisheng/database/models/variable_value.py @@ -1,12 +1,13 @@ from datetime import datetime from typing import Optional, List +# if TYPE_CHECKING: +from pydantic import field_validator +from sqlalchemy import Column, DateTime, text +from sqlmodel import Field, select + from bisheng.database.base import session_getter from bisheng.database.models.base import SQLModelSerializable -# if TYPE_CHECKING: -from pydantic import validator -from sqlalchemy import Column, DateTime, text -from sqlmodel import Field, select, or_ class VariableBase(SQLModelSerializable): @@ -17,16 +18,14 @@ class VariableBase(SQLModelSerializable): value_type: int = Field(index=False, description='变量类型,1=文本 2=list 3=file') is_option: int = Field(index=False, default=1, description='是否必填 1=必填 0=非必填') value: str = Field(index=False, default=0, description='变量值,当文本的时候,传入文本长度') - create_time: Optional[datetime] = Field(sa_column=Column( + create_time: Optional[datetime] = Field(default=None, sa_column=Column( DateTime, nullable=False, index=True, server_default=text('CURRENT_TIMESTAMP'))) - update_time: Optional[datetime] = Field( - sa_column=Column(DateTime, - nullable=False, - server_default=text('CURRENT_TIMESTAMP'), - onupdate=text('CURRENT_TIMESTAMP'))) + update_time: Optional[datetime] = Field(default=None, sa_column=Column( + DateTime, nullable=False, server_default=text('CURRENT_TIMESTAMP'), onupdate=text('CURRENT_TIMESTAMP'))) - @validator('variable_name') - def validate_length(v): + @field_validator('variable_name') + @classmethod + def validate_length(cls, v): # dict_keys(['description', 'name', 'id', 'data']) if not v: return v @@ -35,8 +34,9 @@ class VariableBase(SQLModelSerializable): return v - @validator('value') - def validate_value(v): + @field_validator('value') + @classmethod + def validate_value(cls, v): # dict_keys(['description', 'name', 'id', 'data']) if not v: return v @@ -103,4 +103,3 @@ class VariableDao(Variable): session.add(new_version) session.commit() return old_version - diff --git a/src/backend/bisheng/database/service.py b/src/backend/bisheng/database/service.py index 1b80751d3..585b6ad65 100644 --- a/src/backend/bisheng/database/service.py +++ b/src/backend/bisheng/database/service.py @@ -27,7 +27,7 @@ class DatabaseService(Service): connect_args = {'check_same_thread': False} else: connect_args = {} - return create_engine(self.database_url, connect_args=connect_args, pool_size=100, max_overflow=20, pool_pre_ping=True) + return create_engine(self.database_url, connect_args=connect_args, pool_size=100, max_overflow=20, pool_timeout=3, pool_pre_ping=True) def __enter__(self): self._session = Session(self.engine) diff --git a/src/backend/bisheng/field_typing/range_spec.py b/src/backend/bisheng/field_typing/range_spec.py index 6563147cd..e55f4aa0a 100644 --- a/src/backend/bisheng/field_typing/range_spec.py +++ b/src/backend/bisheng/field_typing/range_spec.py @@ -1,4 +1,4 @@ -from pydantic import BaseModel, validator +from pydantic import field_validator, BaseModel, model_validator class RangeSpec(BaseModel): @@ -6,14 +6,14 @@ class RangeSpec(BaseModel): max: float = 1.0 step: float = 0.1 - @validator('max') + @model_validator(mode='before') @classmethod - def max_must_be_greater_than_min(cls, v, values, **kwargs): - if 'min' in values.data and v <= values.data['min']: + def max_must_be_greater_than_min(cls, values): + if 'min' in values and values['max'] <= values['min']: raise ValueError('max must be greater than min') - return v + return values - @validator('step') + @field_validator('step') @classmethod def step_must_be_positive(cls, v): if v <= 0: diff --git a/src/backend/bisheng/interface/chains/custom.py b/src/backend/bisheng/interface/chains/custom.py index f308e926a..2907fb9d0 100644 --- a/src/backend/bisheng/interface/chains/custom.py +++ b/src/backend/bisheng/interface/chains/custom.py @@ -12,8 +12,7 @@ from langchain.prompts import PromptTemplate from langchain.schema import BaseMemory from langchain.schema.prompt_template import BasePromptTemplate from langchain_community.utilities.dalle_image_generator import DallEAPIWrapper -from langchain_core.pydantic_v1 import Field, root_validator -from pydantic import BaseModel +from pydantic import Field, model_validator DEFAULT_SUFFIX = """" Current conversation: @@ -30,7 +29,8 @@ class BaseCustomConversationChain(ConversationChain): ai_prefix_value: Optional[str] """Field to use as the ai_prefix. It needs to be set and has to be in the template""" - @root_validator(pre=False) + @model_validator(mode='before') + @classmethod def build_template(cls, values): format_dict = {} input_variables = extract_input_variables_from_prompt(values['template']) @@ -168,7 +168,7 @@ prompt_default = PromptTemplate( {image_desc}""") -class DalleGeneratorChain(CustomChain, BaseModel): +class DalleGeneratorChain(CustomChain): """Implementation of dall-e generate images""" dalle: DallEAPIWrapper llm: Optional[BaseLanguageModel] diff --git a/src/backend/bisheng/interface/custom/schema.py b/src/backend/bisheng/interface/custom/schema.py index 89d75006f..7c5975150 100644 --- a/src/backend/bisheng/interface/custom/schema.py +++ b/src/backend/bisheng/interface/custom/schema.py @@ -15,9 +15,6 @@ class ClassCodeDetails(BaseModel): methods: list init: Optional[dict] = Field(default_factory=dict) - def model_dump(self): - return self.dict() - class CallableCodeDetails(BaseModel): """ @@ -30,6 +27,3 @@ class CallableCodeDetails(BaseModel): body: list return_type: Optional[Any] = None has_return: bool = False - - def model_dump(self): - return self.dict() diff --git a/src/backend/bisheng/interface/custom_lists.py b/src/backend/bisheng/interface/custom_lists.py index cc64bb904..4c9bd865a 100644 --- a/src/backend/bisheng/interface/custom_lists.py +++ b/src/backend/bisheng/interface/custom_lists.py @@ -1,20 +1,23 @@ import inspect from typing import Any +from langchain import llms, memory, text_splitter +from langchain_anthropic import ChatAnthropic +from langchain_community import agent_toolkits, document_loaders, embeddings +from langchain_community.chat_models import ChatVertexAI, MiniMaxChat, ChatTongyi, QianfanChatEndpoint, ChatZhipuAI, \ + ChatHunyuan, MoonshotChat +from langchain_community.utilities import requests +from langchain_deepseek import ChatDeepSeek +from langchain_ollama import ChatOllama +from langchain_openai import AzureChatOpenAI, ChatOpenAI, OpenAIEmbeddings, AzureOpenAIEmbeddings, OpenAI + +from bisheng_langchain import chat_models +from bisheng_langchain import document_loaders as contribute_loader +from bisheng_langchain import embeddings as contribute_embeddings from bisheng.interface.agents.custom import CUSTOM_AGENTS from bisheng.interface.chains.custom import CUSTOM_CHAINS from bisheng.interface.embeddings.custom import CUSTOM_EMBEDDING from bisheng.interface.importing.utils import import_class -from bisheng_langchain import chat_models -from bisheng_langchain import document_loaders as contribute_loader -from bisheng_langchain import embeddings as contribute_embeddings -from langchain import llms, memory, text_splitter -from langchain_community.utilities import requests -from langchain_anthropic import ChatAnthropic -from langchain_community import agent_toolkits, document_loaders, embeddings -from langchain_community.chat_models import ChatVertexAI, MiniMaxChat -from langchain_openai import AzureChatOpenAI, ChatOpenAI, OpenAIEmbeddings, AzureOpenAIEmbeddings, OpenAI -from langchain_ollama.chat_models import ChatOllama # LLMs llm_type_to_cls_dict = {} @@ -29,7 +32,13 @@ llm_type_to_cls_dict['ChatOpenAI'] = ChatOpenAI # type: ignore llm_type_to_cls_dict['ChatVertexAI'] = ChatVertexAI # type: ignore llm_type_to_cls_dict['MiniMaxChat'] = MiniMaxChat llm_type_to_cls_dict['ChatOllama'] = ChatOllama +llm_type_to_cls_dict['ChatTongyi'] = ChatTongyi +llm_type_to_cls_dict['QianfanChatEndpoint'] = QianfanChatEndpoint llm_type_to_cls_dict["OpenAI"] = OpenAI +llm_type_to_cls_dict['ChatZhipuAI'] = ChatZhipuAI +llm_type_to_cls_dict['ChatDeepSeek'] = ChatDeepSeek +llm_type_to_cls_dict['ChatHunyuan'] = ChatHunyuan +llm_type_to_cls_dict['MoonshotChat'] = MoonshotChat # llm contribute llm_type_to_cls_dict.update({ @@ -58,13 +67,13 @@ for memory_name in memory.__all__: elif memory_name == "ConversationKGMemory": memory_type_to_cls_dict[memory_name] = import_class(f"langchain_community.memory.kg.{memory_name}") elif memory_name == "MotorheadMemory": - memory_type_to_cls_dict[memory_name] = import_class(f"langchain_community.memory.motorhead_memory.{memory_name}") + memory_type_to_cls_dict[memory_name] = import_class( + f"langchain_community.memory.motorhead_memory.{memory_name}") elif memory_name == "ZepMemory": memory_type_to_cls_dict[memory_name] = import_class(f"langchain_community.memory.zep_memory.{memory_name}") else: memory_type_to_cls_dict[memory_name] = import_class(f'langchain.memory.{memory_name}') - # Wrappers wrapper_type_to_cls_dict: dict[str, Any] = { wrapper.__name__: wrapper diff --git a/src/backend/bisheng/interface/embeddings/custom.py b/src/backend/bisheng/interface/embeddings/custom.py index 0393c3b4b..77dba4a58 100644 --- a/src/backend/bisheng/interface/embeddings/custom.py +++ b/src/backend/bisheng/interface/embeddings/custom.py @@ -1,4 +1,4 @@ -from typing import List, Optional +from typing import List, Optional, Dict import numpy as np from bisheng.database.models.llm_server import (LLMDao, LLMModel, LLMModelType, LLMServer, @@ -6,9 +6,8 @@ from bisheng.database.models.llm_server import (LLMDao, LLMModel, LLMModelType, from bisheng.interface.importing import import_by_type from bisheng.interface.utils import wrapper_bisheng_model_limit_check from langchain.embeddings.base import Embeddings -from langchain_core.pydantic_v1 import BaseModel from loguru import logger -from pydantic import Field +from pydantic import ConfigDict, Field, BaseModel class OpenAIProxyEmbedding(Embeddings): @@ -57,7 +56,7 @@ class BishengEmbedding(BaseModel, Embeddings): model_kwargs: dict = Field(default={}, description='embedding模型调用参数') embeddings: Optional[Embeddings] = Field(default=None) - llm_node_type = { + llm_node_type: Dict = { # 开源推理框架 LLMServerType.OLLAMA.value: 'OllamaEmbeddings', LLMServerType.XINFERENCE.value: 'OpenAIEmbeddings', @@ -72,16 +71,15 @@ class BishengEmbedding(BaseModel, Embeddings): LLMServerType.QIAN_FAN.value: 'QianfanEmbeddingsEndpoint', LLMServerType.MINIMAX.value: 'OpenAIEmbeddings', LLMServerType.ZHIPU.value: 'OpenAIEmbeddings', + LLMServerType.TENCENT.value: 'OpenAIEmbeddings', + LLMServerType.VOLCENGINE.value: 'OpenAIEmbeddings', + LLMServerType.SILICON.value: 'OpenAIEmbeddings', } # bisheng强相关的业务参数 model_info: Optional[LLMModel] = Field(default=None) server_info: Optional[LLMServer] = Field(default=None) - - class Config: - """Configuration for this pydantic object.""" - allow_population_by_field_name = True - arbitrary_types_allowed = True + model_config = ConfigDict(validate_by_name=True, arbitrary_types_allowed=True) def __init__(self, **kwargs): from bisheng.interface.initialize.loading import instantiate_embedding @@ -187,7 +185,10 @@ class BishengEmbedding(BaseModel, Embeddings): def _update_model_status(self, status: int, remark: str = ''): """更新模型状态""" - LLMDao.update_model_status(self.model_id, status, remark) + # todo 接入到异步任务模块 累计5分钟更新一次 + if self.model_info.status != status: + self.model_info.status = status + LLMDao.update_model_status(self.model_id, status, remark) CUSTOM_EMBEDDING = { diff --git a/src/backend/bisheng/interface/llms/base.py b/src/backend/bisheng/interface/llms/base.py index cb38df9d9..8a2cabb79 100644 --- a/src/backend/bisheng/interface/llms/base.py +++ b/src/backend/bisheng/interface/llms/base.py @@ -7,7 +7,7 @@ from bisheng.settings import settings from bisheng.template.frontend_node.llms import LLMFrontendNode from bisheng.utils.logger import logger from bisheng.utils.util import build_template_from_class - +from bisheng.interface.llms.chat_spark import ChatSparkOpenAI class LLMCreator(LangChainTypeCreator): type_name: str = 'llms' @@ -21,7 +21,8 @@ class LLMCreator(LangChainTypeCreator): if self.type_dict is None: self.type_dict = llm_type_to_cls_dict self.type_dict.update({ - 'BishengLLM': BishengLLM + 'BishengLLM': BishengLLM, + 'ChatSparkOpenAI': ChatSparkOpenAI }) return self.type_dict diff --git a/src/backend/bisheng/interface/llms/chat_spark.py b/src/backend/bisheng/interface/llms/chat_spark.py new file mode 100644 index 000000000..fb6ea37cd --- /dev/null +++ b/src/backend/bisheng/interface/llms/chat_spark.py @@ -0,0 +1,21 @@ +from typing import Optional, Any + +from langchain_core.language_models import LanguageModelInput +from langchain_openai import ChatOpenAI + + +class ChatSparkOpenAI(ChatOpenAI): + + def _get_request_payload( + self, + input_: LanguageModelInput, + *, + stop: Optional[list[str]] = None, + **kwargs: Any, + ) -> dict: + payload = super()._get_request_payload(input_, stop=stop, **kwargs) + # max_tokens was deprecated in favor of max_completion_tokens + # in September 2024 release + if "max_completion_tokens" in payload: + payload["max_tokens"] = payload.pop("max_completion_tokens") + return payload \ No newline at end of file diff --git a/src/backend/bisheng/interface/llms/custom.py b/src/backend/bisheng/interface/llms/custom.py index cf7b93aab..e2e8f666b 100644 --- a/src/backend/bisheng/interface/llms/custom.py +++ b/src/backend/bisheng/interface/llms/custom.py @@ -1,18 +1,139 @@ -from typing import List, Optional, Any, Sequence, Union, Dict, Type, Callable +import json +from typing import List, Optional, Any, Sequence, Union, Dict, Type, Callable, Iterator, AsyncIterator -from langchain_core.messages import BaseMessage -from langchain_core.outputs import ChatResult +from langchain_core.callbacks import AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun +from langchain_core.language_models import BaseLanguageModel, BaseChatModel, LanguageModelInput +from langchain_core.messages import BaseMessage, ToolMessage, HumanMessage, BaseMessageChunk +from langchain_core.outputs import ChatResult, ChatGenerationChunk from langchain_core.runnables import Runnable from langchain_core.tools import BaseTool from loguru import logger from pydantic import Field -from langchain_core.callbacks import AsyncCallbackManagerForLLMRun -from langchain_core.language_models import BaseLanguageModel, BaseChatModel, LanguageModelInput from bisheng.database.models.llm_server import LLMDao, LLMModelType, LLMServerType, LLMModel, LLMServer from bisheng.interface.importing import import_by_type from bisheng.interface.initialize.loading import instantiate_llm -from bisheng.interface.utils import wrapper_bisheng_model_limit_check, wrapper_bisheng_model_limit_check_async +from bisheng.interface.utils import wrapper_bisheng_model_limit_check, wrapper_bisheng_model_limit_check_async, \ + wrapper_bisheng_model_generator, wrapper_bisheng_model_generator_async + + +def _get_ollama_params(params: dict, server_config: dict, model_config: dict) -> dict: + params['base_url'] = server_config.get('base_url', '') + # some bugs + params['extract_reasoning'] = False + params['stream'] = params.pop('streaming', True) + if params.get('max_tokens'): + params['num_ctx'] = params.pop('max_tokens', None) + return params + + +def _get_xinference_params(params: dict, server_config: dict, model_config: dict) -> dict: + params = _get_openai_params(params, server_config, model_config) + if not params.get('api_key', None): + params['api_key'] = 'Empty' + return params + + +def _get_bisheng_rt_params(params: dict, server_config: dict, model_config: dict) -> dict: + params.update(server_config) + return params + + +def _get_openai_params(params: dict, server_config: dict, model_config: dict) -> dict: + if server_config: + params.update({ + 'api_key': server_config.get('openai_api_key') or server_config.get('api_key'), + 'base_url': server_config.get('openai_api_base') or server_config.get('base_url'), + }) + if server_config.get('openai_proxy'): + params['openai_proxy'] = server_config.get('openai_proxy') + return params + + +def _get_azure_openai_params(params: dict, server_config: dict, model_config: dict) -> dict: + params.update({ + 'azure_endpoint': server_config.get('azure_endpoint'), + 'openai_api_key': server_config.get('openai_api_key'), + 'openai_api_version': server_config.get('openai_api_version'), + 'azure_deployment': params.pop('model'), + }) + return params + + +def _get_qwen_params(params: dict, server_config: dict, model_config: dict) -> dict: + params['dashscope_api_key'] = server_config.get('openai_api_key', '') + params['model_kwargs'] = { + 'enable_search': model_config.get('enable_web_search', False), + 'temperature': params.pop('temperature', 0.3), + } + if params.get('max_tokens'): + params['model_kwargs']['max_tokens'] = params.get('max_tokens') + return params + + +def _get_qianfan_params(params: dict, server_config: dict, model_config: dict) -> dict: + params['qianfan_ak'] = server_config.get('wenxin_api_key') + params['qianfan_sk'] = server_config.get('wenxin_secret_key') + if params.get('max_tokens'): + params['model_kwargs'] = {"max_output_tokens": params.pop('max_tokens')} + return params + + +def _get_minimax_params(params: dict, server_config: dict, model_config: dict) -> dict: + params['minimax_api_key'] = server_config.get('openai_api_key') + params['base_url'] = server_config.get('openai_api_base') + if 'max_tokens' not in params: + params['max_tokens'] = 2048 + if '/chat/completions' not in params['base_url']: + params['base_url'] = f"{params['base_url']}/chat/completions" + return params + + +def _get_anthropic_params(params: dict, server_config: dict, model_config: dict) -> dict: + params.update(server_config) + return params + + +def _get_zhipu_params(params: dict, server_config: dict, model_config: dict) -> dict: + params['zhipuai_api_key'] = server_config.get('openai_api_key') + params['zhipuai_api_base'] = server_config.get('openai_api_base') + if 'chat/completions' not in params['zhipuai_api_base']: + params['zhipuai_api_base'] = f"{params['zhipuai_api_base'].rstrip('/')}/chat/completions" + return params + + +def _get_spark_params(params: dict, server_config: dict, model_config: dict) -> dict: + params.update({ + 'api_key': f'{server_config.get("api_key")}:{server_config.get("api_secret")}', + 'base_url': server_config.get('openai_api_base'), + }) + return params + + +_llm_node_type: Dict = { + # 开源推理框架 + LLMServerType.OLLAMA.value: {'client': 'ChatOllama', 'params_handler': _get_ollama_params}, + LLMServerType.XINFERENCE.value: {'client': 'ChatOpenAI', 'params_handler': _get_xinference_params}, + LLMServerType.LLAMACPP.value: {'client': 'ChatOpenAI', 'params_handler': _get_openai_params}, + # 此组件是加载本地的模型文件,待确认是否有api服务提供 + LLMServerType.VLLM.value: {'client': 'ChatOpenAI', 'params_handler': _get_openai_params}, + LLMServerType.BISHENG_RT.value: {'client': 'HostChatGLM', 'params_handler': _get_bisheng_rt_params}, + + # 官方api服务 + LLMServerType.OPENAI.value: {'client': 'ChatOpenAI', 'params_handler': _get_openai_params}, + LLMServerType.AZURE_OPENAI.value: {'client': 'AzureChatOpenAI', 'params_handler': _get_azure_openai_params}, + LLMServerType.QWEN.value: {'client': 'ChatTongyi', 'params_handler': _get_qwen_params}, + LLMServerType.QIAN_FAN.value: {'client': 'QianfanChatEndpoint', 'params_handler': _get_qianfan_params}, + LLMServerType.ZHIPU.value: {'client': 'ChatZhipuAI', 'params_handler': _get_zhipu_params}, + LLMServerType.MINIMAX.value: {'client': 'MiniMaxChat', 'params_handler': _get_minimax_params}, + LLMServerType.ANTHROPIC.value: {'client': 'ChatAnthropic', 'params_handler': _get_anthropic_params}, + LLMServerType.DEEPSEEK.value: {'client': 'ChatDeepSeek', 'params_handler': _get_openai_params}, + LLMServerType.SPARK.value: {'client': 'ChatSparkOpenAI', 'params_handler': _get_spark_params}, + LLMServerType.TENCENT.value: {'client': 'ChatSparkOpenAI', 'params_handler': _get_openai_params}, + LLMServerType.MOONSHOT.value: {'client': 'MoonshotChat', 'params_handler': _get_openai_params}, + LLMServerType.VOLCENGINE.value: {'client': 'ChatSparkOpenAI', 'params_handler': _get_openai_params}, + LLMServerType.SILICON.value: {'client': 'ChatSparkOpenAI', 'params_handler': _get_openai_params}, +} class BishengLLM(BaseChatModel): @@ -28,25 +149,6 @@ class BishengLLM(BaseChatModel): cache: bool = Field(default=False, description="是否使用缓存") llm: Optional[BaseChatModel] = Field(default=None) - llm_node_type = { - # 开源推理框架 - LLMServerType.OLLAMA.value: 'ChatOllama', - LLMServerType.XINFERENCE.value: 'ChatOpenAI', - LLMServerType.LLAMACPP.value: 'ChatOpenAI', # 此组件是加载本地的模型文件,待确认是否有api服务提供 - LLMServerType.VLLM.value: 'ChatOpenAI', - LLMServerType.BISHENG_RT.value: "HostChatGLM", - - # 官方api服务 - LLMServerType.OPENAI.value: 'ChatOpenAI', - LLMServerType.AZURE_OPENAI.value: 'AzureChatOpenAI', - LLMServerType.QWEN.value: 'ChatOpenAI', - LLMServerType.QIAN_FAN.value: 'ChatWenxin', - LLMServerType.ZHIPU.value: 'ChatOpenAI', - LLMServerType.MINIMAX.value: 'ChatOpenAI', - LLMServerType.ANTHROPIC.value: 'ChatAnthropic', - LLMServerType.DEEPSEEK.value: 'ChatOpenAI', - LLMServerType.SPARK.value: 'ChatOpenAI', - } # bisheng强相关的业务参数 model_info: Optional[LLMModel] = Field(default=None) @@ -80,25 +182,38 @@ class BishengLLM(BaseChatModel): self.model_info = model_info self.server_info = server_info - class_object = self._get_llm_class(server_info.type) + class_object, class_name = self._get_llm_class(server_info.type) params = self._get_llm_params(server_info, model_info) try: - self.llm = instantiate_llm(self.llm_node_type.get(server_info.type), class_object, params) + self.llm = instantiate_llm(class_name, class_object, params) except Exception as e: logger.exception('init bisheng llm error') raise Exception(f'初始化llm失败,请检查配置或联系管理员。错误信息:{e}') - def _get_llm_class(self, server_type: str) -> BaseLanguageModel: - node_type = self.llm_node_type[server_type] + def _get_llm_class(self, server_type: str) -> (BaseLanguageModel, str): + if server_type not in _llm_node_type: + raise Exception(f'not support llm type: {server_type}') + node_type = _llm_node_type[server_type]['client'] class_object = import_by_type(_type='llms', name=node_type) - return class_object + return class_object, node_type def _get_llm_params(self, server_info: LLMServer, model_info: LLMModel) -> dict: + server_config = self.get_server_info_config() + model_config = self.get_model_info_config() + default_params = self._get_default_params(server_config, model_config) + + params_handler = _llm_node_type[server_info.type]['params_handler'] + params = params_handler(default_params, server_config, model_config) + return params + params = {} if server_info.config: params.update(server_info.config) + enable_web_search = False if model_info.config: - params.update(model_info.config) + enable_web_search = model_info.config.get('enable_web_search', False) + if model_info.config.get('max_tokens'): + params['max_tokens'] = model_info.config.get('max_tokens') params.update({ 'model_name': model_info.model_name, @@ -109,21 +224,116 @@ class BishengLLM(BaseChatModel): }) if server_info.type == LLMServerType.OLLAMA.value: params['model'] = params.pop('model_name') + if params.get('max_tokens'): + params['num_ctx'] = params.pop('max_tokens') elif server_info.type == LLMServerType.AZURE_OPENAI.value: params['azure_deployment'] = params.pop('model_name') elif server_info.type == LLMServerType.QIAN_FAN.value: params['model'] = params.pop('model_name') + params['qianfan_ak'] = params.pop('wenxin_api_key') + params['qianfan_sk'] = params.pop('wenxin_secret_key') + if params.get('max_tokens'): + params['model_kwargs'] = {"max_output_tokens": params.pop('max_tokens')} elif server_info.type == LLMServerType.SPARK.value: params['openai_api_key'] = f'{params.pop("api_key")}:{params.pop("api_secret")}' elif server_info.type in [LLMServerType.XINFERENCE.value, LLMServerType.LLAMACPP.value, LLMServerType.VLLM.value]: params['openai_api_key'] = params.pop('openai_api_key', None) or "EMPTY" + elif server_info.type == LLMServerType.QWEN.value: + params['dashscope_api_key'] = params.pop('openai_api_key') + params.pop('openai_api_base', None) + params['model_kwargs'] = {'enable_search': enable_web_search} + if params.get('max_tokens'): + params['model_kwargs']['max_tokens'] = params.pop('max_tokens') + elif server_info.type == LLMServerType.TENCENT.value: + params['extra_body'] = {'enable_enhancement': enable_web_search} + elif server_info.type == LLMServerType.MINIMAX.value: + params['api_key'] = params.pop('openai_api_key', None) + params.pop('openai_api_base', None) return params + def _get_default_params(self, server_config: dict, model_config: dict) -> dict: + default_params = { + 'model': self.model_info.model_name, + 'streaming': self.streaming, + 'temperature': self.temperature, + 'top_p': self.top_p, + 'cache': self.cache + } + if model_config.get('max_tokens'): + default_params['max_tokens'] = model_config.get('max_tokens') + + return default_params + @property def _llm_type(self): return self.llm._llm_type + def get_server_info_config(self): + if self.server_info.config: + return self.server_info.config + return {} + + def get_model_info_config(self): + if self.model_info.config: + return self.model_info.config + return {} + + def parse_kwargs(self, messages: List[BaseMessage], kwargs: Dict[str, Any]) -> (List[BaseMessage], Dict[str, Any]): + if self.server_info.type == LLMServerType.MINIMAX.value: + if self.get_model_info_config().get('enable_web_search'): + if 'tools' not in kwargs: + kwargs.update({ + 'tools': [{'type': 'web_search'}], + }) + else: + tool_exists = False + for tool in kwargs['tools']: + if tool.get('type') == 'web_search': + tool_exists = True + break + if not tool_exists: + kwargs['tools'].append({ + 'type': 'web_search', + }) + elif self.server_info.type == LLMServerType.MOONSHOT.value: + if self.get_model_info_config().get('enable_web_search'): + if 'tools' not in kwargs: + kwargs.update({ + 'tools': [{ + "type": "builtin_function", + "function": { + "name": "$web_search", + }, + }], + }) + else: + tool_exists = False + for tool in kwargs['tools']: + if tool.get('type') == 'builtin_function': + tool_exists = True + break + if not tool_exists: + kwargs['tools'].append({ + "type": "builtin_function", + "function": { + "name": "$web_search", + }, + }) + elif self.server_info.type == LLMServerType.QWEN.value: + # ChatTongYi 对多模态的入参比较特殊,需要转换以支持 + user_message = messages[-1] + if isinstance(user_message, HumanMessage): + if isinstance(user_message.content, list): + for one in user_message.content: + if one.get('type') == 'image' and one.get('data'): + one['type'] = 'image' + one['image'] = f"data:{one.get('mime_type')};{one.get('source_type')},{one.get('data')}" + elif one.get('type') == 'image_url' and one.get('image_url'): + one['type'] = 'image' + one['image'] = one.pop('image_url', {}).get('url') + return messages, kwargs + @wrapper_bisheng_model_limit_check def _generate( self, @@ -134,13 +344,50 @@ class BishengLLM(BaseChatModel): **kwargs: Any, ) -> ChatResult: try: - ret = self.llm._generate(messages, stop, run_manager, **kwargs) + messages, kwargs = self.parse_kwargs(messages, kwargs) + if self.server_info.type == LLMServerType.MOONSHOT.value: + ret = self.moonshot_generate(messages, stop, run_manager, **kwargs) + else: + ret = self.llm._generate(messages, stop, run_manager, **kwargs) + if self.server_info.type == LLMServerType.QWEN.value: + ret.generations[0].message = self.convert_qwen_result(ret.generations[0].message) self._update_model_status(0) except Exception as e: self._update_model_status(1, str(e)) raise e return ret + def moonshot_generate( + self, + messages: List[BaseMessage], + stop: Optional[List[str]] = None, + run_manager: Optional[AsyncCallbackManagerForLLMRun] = None, + **kwargs: Any, + ) -> ChatResult: + try: + result = None + finish_reason = None + while finish_reason is None or finish_reason == 'tool_calls': + result = self.llm._generate(messages, stop, run_manager, **kwargs) + result_message = result.generations[0].message + finish_reason = result.generations[0].generation_info.get('finish_reason') + for tool_call in result_message.tool_calls: + tool_call_name = tool_call['name'] + if tool_call_name == "$web_search": + messages.append(result_message) + messages.append(ToolMessage( + tool_call_id=tool_call['id'], + name=tool_call_name, + content=json.dumps(tool_call['args'], ensure_ascii=False), + )) + else: + break + self._update_model_status(0) + except Exception as e: + self._update_model_status(1, str(e)) + raise e + return result + @wrapper_bisheng_model_limit_check_async async def _agenerate( self, @@ -151,7 +398,13 @@ class BishengLLM(BaseChatModel): **kwargs: Any, ) -> ChatResult: try: - ret = await self.llm._agenerate(messages, stop, run_manager, **kwargs) + messages, kwargs = self.parse_kwargs(messages, kwargs) + if self.server_info.type == LLMServerType.MOONSHOT.value: + ret = await self.moonshot_agenerate(messages, stop, run_manager, **kwargs) + else: + ret = await self.llm._agenerate(messages, stop, run_manager, **kwargs) + if self.server_info.type == LLMServerType.QWEN.value: + ret.generations[0].message = self.convert_qwen_result(ret.generations[0].message) self._update_model_status(0) except Exception as e: self._update_model_status(1, str(e)) @@ -159,14 +412,89 @@ class BishengLLM(BaseChatModel): raise e return ret + async def moonshot_agenerate( + self, + messages: List[BaseMessage], + stop: Optional[List[str]] = None, + run_manager: Optional[AsyncCallbackManagerForLLMRun] = None, + **kwargs: Any, + ) -> ChatResult: + try: + result = None + finish_reason = None + while finish_reason is None or finish_reason == 'tool_calls': + result = await self.llm._agenerate(messages, stop, run_manager, **kwargs) + result_message = result.generations[0].message + finish_reason = result.generations[0].generation_info.get('finish_reason') + for tool_call in result_message.tool_calls: + tool_call_name = tool_call['name'] + if tool_call_name == "$web_search": + messages.append(result_message) + messages.append(ToolMessage( + tool_call_id=tool_call['id'], + name=tool_call_name, + content=json.dumps(tool_call['args'], ensure_ascii=False), + )) + else: + break + self._update_model_status(0) + except Exception as e: + self._update_model_status(1, str(e)) + raise e + return result + def _update_model_status(self, status: int, remark: str = ''): """更新模型状态""" - # todo 接入到异步任务模块 - LLMDao.update_model_status(self.model_id, status, remark) + # todo 接入到异步任务模块 累计5分钟更新一次 + if self.model_info.status != status: + self.model_info.status = status + LLMDao.update_model_status(self.model_id, status, remark) def bind_tools( - self, - tools: Sequence[Union[Dict[str, Any], Type, Callable, BaseTool]], - **kwargs: Any, + self, + tools: Sequence[Union[Dict[str, Any], Type, Callable, BaseTool]], + **kwargs: Any, ) -> Runnable[LanguageModelInput, BaseMessage]: return self.llm.bind_tools(tools, **kwargs) + + def convert_qwen_result(self, message: BaseMessageChunk | BaseMessage) -> BaseMessageChunk | BaseMessage: + # ChatTongYi model vl model message.content is list + if isinstance(message.content, list): + message.content = ''.join([one.get('text', '') for one in message.content]) + return message + + @wrapper_bisheng_model_generator + def _stream( + self, + messages: list[BaseMessage], + stop: Optional[list[str]] = None, + run_manager: Optional[CallbackManagerForLLMRun] = None, + **kwargs: Any, + ) -> Iterator[ChatGenerationChunk]: + try: + for one in self.llm._stream(messages, stop=stop, run_manager=run_manager, **kwargs): + if self.server_info.type == LLMServerType.QWEN.value: + one.message = self.convert_qwen_result(one.message) + yield one + self._update_model_status(0) + except Exception as e: + self._update_model_status(1, str(e)) + raise e + + @wrapper_bisheng_model_generator_async + async def _astream( + self, + messages: list[BaseMessage], + stop: Optional[list[str]] = None, + run_manager: Optional[AsyncCallbackManagerForLLMRun] = None, + **kwargs: Any, + ) -> AsyncIterator[ChatGenerationChunk]: + try: + async for one in self.llm._astream(messages, stop=stop, run_manager=run_manager, **kwargs): + if self.server_info.type == LLMServerType.QWEN.value: + one.message = self.convert_qwen_result(one.message) + yield one + self._update_model_status(0) + except Exception as e: + self._update_model_status(1, str(e)) + raise e diff --git a/src/backend/bisheng/interface/prompts/custom.py b/src/backend/bisheng/interface/prompts/custom.py index 35d0ce443..81a18aa51 100644 --- a/src/backend/bisheng/interface/prompts/custom.py +++ b/src/backend/bisheng/interface/prompts/custom.py @@ -2,7 +2,7 @@ from typing import Dict, List, Optional, Type from bisheng.interface.utils import extract_input_variables_from_prompt from langchain.prompts import PromptTemplate -from langchain_core.pydantic_v1 import root_validator +from pydantic import model_validator # Steps to create a BaseCustomPrompt: # 1. Create a prompt template that endes with: @@ -27,7 +27,8 @@ class BaseCustomPrompt(PromptTemplate): description: Optional[str] ai_prefix: Optional[str] - @root_validator(pre=False) + @model_validator(mode='before') + @classmethod def build_template(cls, values): format_dict = {} ai_prefix_format_dict = {} diff --git a/src/backend/bisheng/interface/tools/base.py b/src/backend/bisheng/interface/tools/base.py index a8df56a88..a2d3bbb89 100644 --- a/src/backend/bisheng/interface/tools/base.py +++ b/src/backend/bisheng/interface/tools/base.py @@ -121,8 +121,8 @@ class ToolCreator(LangChainTypeCreator): # Pop unnecessary fields and add name fields.pop('_type') # type: ignore - fields.pop('return_direct') # type: ignore - fields.pop('verbose') # type: ignore + fields.pop('return_direct', None) # type: ignore + fields.pop('verbose', None) # type: ignore tool_params = { 'name': fields.pop('name')['value'], # type: ignore diff --git a/src/backend/bisheng/interface/tools/custom.py b/src/backend/bisheng/interface/tools/custom.py index 0079cf634..64c80f3e5 100644 --- a/src/backend/bisheng/interface/tools/custom.py +++ b/src/backend/bisheng/interface/tools/custom.py @@ -3,7 +3,7 @@ from typing import Callable, Optional from bisheng.interface.importing.utils import get_function from bisheng.utils import validate from langchain_community.tools import Tool -from pydantic import BaseModel, validator +from pydantic import field_validator, BaseModel class Function(BaseModel): @@ -16,7 +16,8 @@ class Function(BaseModel): super().__init__(**data) # Validate the function - @validator('code') + @field_validator('code') + @classmethod def validate_func(cls, v): try: validate.eval_function(v) @@ -32,7 +33,7 @@ class Function(BaseModel): return validate.create_function(self.code, function_name) -class PythonFunctionTool(Function, Tool): +class PythonFunctionTool(Tool): """Python function""" name: str = 'Custom Tool' diff --git a/src/backend/bisheng/interface/utils.py b/src/backend/bisheng/interface/utils.py index 93e4bd4c0..e3040d76e 100644 --- a/src/backend/bisheng/interface/utils.py +++ b/src/backend/bisheng/interface/utils.py @@ -1,20 +1,19 @@ import base64 import datetime import functools -import inspect import json import os import re from io import BytesIO import yaml +from PIL.Image import Image +from langchain.base_language import BaseLanguageModel from bisheng.cache.redis import redis_client from bisheng.chat.config import ChatConfig from bisheng.settings import settings from bisheng.utils.logger import logger -from langchain.base_language import BaseLanguageModel -from PIL.Image import Image def load_file_into_dict(file_path: str) -> dict: @@ -100,7 +99,7 @@ def bisheng_model_limit_check(self: 'BishengLLM | BishengEmbedding'): cache_key = f"model_limit:{now}:{self.server_info.id}" use_num = redis_client.incr(cache_key) if use_num > self.server_info.limit: - raise Exception(f'额度已用完') + raise Exception(f'{self.server_info.name}/{self.model_info.model_name} 额度已用完') def wrapper_bisheng_model_limit_check_async(func): @@ -127,3 +126,31 @@ def wrapper_bisheng_model_limit_check(func): return func(*args, **kwargs) return wrapper + + +def wrapper_bisheng_model_generator(func): + """ + 调用次数检查的装饰器 装饰同步生成器函数 + """ + + @functools.wraps(func) + def wrapper(*args, **kwargs): + bisheng_model_limit_check(args[0]) + for item in func(*args, **kwargs): + yield item + + return wrapper + + +def wrapper_bisheng_model_generator_async(func): + """ + 调用次数检查的装饰器 装饰异步生成器函数 + """ + + @functools.wraps(func) + async def wrapper(*args, **kwargs): + bisheng_model_limit_check(args[0]) + async for item in func(*args, **kwargs): + yield item + + return wrapper diff --git a/src/backend/bisheng/main.py b/src/backend/bisheng/main.py index 4143a79a3..4485f81e1 100644 --- a/src/backend/bisheng/main.py +++ b/src/backend/bisheng/main.py @@ -5,7 +5,6 @@ from typing import Optional from bisheng.api import router, router_rpc from bisheng.database.init_data import init_default_data from bisheng.interface.utils import setup_llm_caching -from bisheng.restructure.register import register_restructure from bisheng.services.utils import initialize_services, teardown_services from bisheng.settings import settings from bisheng.utils.http_middleware import CustomMiddleware @@ -95,7 +94,6 @@ def create_app(): app.include_router(router) app.include_router(router_rpc) - register_restructure(app) return app diff --git a/src/backend/bisheng/restructure/__init__.py b/src/backend/bisheng/mcp_manage/__init__.py similarity index 100% rename from src/backend/bisheng/restructure/__init__.py rename to src/backend/bisheng/mcp_manage/__init__.py diff --git a/src/backend/bisheng/restructure/assistants/__init__.py b/src/backend/bisheng/mcp_manage/clients/__init__.py similarity index 100% rename from src/backend/bisheng/restructure/assistants/__init__.py rename to src/backend/bisheng/mcp_manage/clients/__init__.py diff --git a/src/backend/bisheng/mcp_manage/clients/base.py b/src/backend/bisheng/mcp_manage/clients/base.py new file mode 100644 index 000000000..14fa0144e --- /dev/null +++ b/src/backend/bisheng/mcp_manage/clients/base.py @@ -0,0 +1,46 @@ +from abc import abstractmethod +from contextlib import AsyncExitStack, asynccontextmanager +from typing import Any + +from mcp import ClientSession + + +class BaseMcpClient(object): + """ + Base class for MCP clients. + """ + + def __init__(self, **kwargs): + self.exit_stack = AsyncExitStack() + + self.client_session: ClientSession | None = None + + @abstractmethod + async def get_transport(self): + raise NotImplementedError("get_mcp_client_transport() must be implemented in subclasses.") + + @asynccontextmanager + async def initialize(self): + """ + Initialize the client. + """ + async with self.get_transport() as (read, write): + async with ClientSession(read, write) as session: + await session.initialize() + yield session + + async def list_tools(self): + async with self.initialize() as client_session: + tools = await client_session.list_tools() + return tools.tools + + async def call_tool(self, name: str, arguments: dict[str, Any] | None = None) -> str: + """ + Call a tool. + """ + async with self.initialize() as client_session: + try: + resp = await client_session.call_tool(name, arguments) + except Exception as e: + return f"Tool call failed: {str(e)}" + return resp.model_dump_json() diff --git a/src/backend/bisheng/mcp_manage/clients/sse.py b/src/backend/bisheng/mcp_manage/clients/sse.py new file mode 100644 index 000000000..e25fdfda0 --- /dev/null +++ b/src/backend/bisheng/mcp_manage/clients/sse.py @@ -0,0 +1,29 @@ +from contextlib import asynccontextmanager + +from mcp.client.sse import sse_client + +from bisheng.mcp_manage.clients.base import BaseMcpClient + + +class SseClient(BaseMcpClient): + """ + SSE client for connecting to the mcp server. + """ + + def __init__(self, url: str, **kwargs): + """ + Initialize the SSE client. + + :param url: The URL of the SSE server. + """ + super().__init__() + self.url = url + self.kwargs = kwargs + + @asynccontextmanager + async def get_transport(self): + """ + Initialize the SSE client. + """ + async with sse_client(url=self.url, **self.kwargs) as (read, write): + yield read, write diff --git a/src/backend/bisheng/mcp_manage/clients/stdio.py b/src/backend/bisheng/mcp_manage/clients/stdio.py new file mode 100644 index 000000000..1e94eeb1b --- /dev/null +++ b/src/backend/bisheng/mcp_manage/clients/stdio.py @@ -0,0 +1,28 @@ +from contextlib import asynccontextmanager + +from mcp.client.stdio import stdio_client, StdioServerParameters + +from bisheng.mcp_manage.clients.base import BaseMcpClient + + +class StdioClient(BaseMcpClient): + """ + SSE client for connecting to the mcp server. + """ + + def __init__(self, **kwargs: dict): + """ + Initialize the SSE client. + + :param url: The URL of the SSE server. + """ + super().__init__() + self.server_params = StdioServerParameters(**kwargs) + + @asynccontextmanager + async def get_transport(self): + """ + Initialize the SSE client. + """ + async with stdio_client(server=self.server_params) as (read, write): + yield read, write diff --git a/src/backend/bisheng/mcp_manage/constant.py b/src/backend/bisheng/mcp_manage/constant.py new file mode 100644 index 000000000..9d307e75e --- /dev/null +++ b/src/backend/bisheng/mcp_manage/constant.py @@ -0,0 +1,5 @@ +from enum import Enum + +class McpClientType(Enum): + SSE = 'sse' + STDIO = 'stdio' diff --git a/src/backend/bisheng/mcp_manage/langchain/__init__.py b/src/backend/bisheng/mcp_manage/langchain/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/src/backend/bisheng/mcp_manage/langchain/tool.py b/src/backend/bisheng/mcp_manage/langchain/tool.py new file mode 100644 index 000000000..6cb920d18 --- /dev/null +++ b/src/backend/bisheng/mcp_manage/langchain/tool.py @@ -0,0 +1,88 @@ +import asyncio +import concurrent +import concurrent.futures +from typing import Any, Type + +from langchain_core.tools import StructuredTool +from pydantic import Field, create_model, BaseModel, ConfigDict + +from bisheng.mcp_manage.clients.base import BaseMcpClient + + +def convert_openai_params_to_model(params: list[dict]) -> Type[BaseModel]: + model_params = {} + for one in params: + field_type = one['schema']['type'] + if field_type == 'number': + field_type = float + elif field_type == 'integer': + field_type = int + elif field_type == 'string': + field_type = str + elif field_type == 'boolean': + field_type = bool + elif field_type == 'array': + field_type = list + elif field_type in {'object', 'dict'}: + param_object_param = {} + for param in one['schema']['properties'].keys(): + field_type = one['schema']['properties'][param]['type'] + if field_type == 'number': + field_type = float + elif field_type == 'integer': + field_type = int + elif field_type == 'string': + field_type = str + elif field_type == 'boolean': + field_type = bool + elif field_type == 'array': + field_type = list + param_object_param[param] = ( + field_type, + Field(description=one['schema']['properties'][param]['description'])) + param_model = create_model( + param, + __module__='bisheng.mcp_manage.langchain.tool', + **param_object_param) + field_type = param_model + else: + raise Exception(f'schema type is not support: {field_type}') + model_params[one['name']] = (field_type, Field(description=one['description'])) + return create_model('InputArgs', + __module__='bisheng.mcp_manage.langchain.tool', + __base__=BaseModel, + **model_params) + + +class McpTool(BaseModel): + model_config = ConfigDict(arbitrary_types_allowed=True) + + name: str + description: str + mcp_client: BaseMcpClient + mcp_tool_name: str + + def run(self, *args, **kwargs: Any) -> Any: + # todo call async method better when using in event pool + with concurrent.futures.ThreadPoolExecutor() as pool: + future = pool.submit(asyncio.run, self.arun(*args, **kwargs)) + resp = future.result() + return resp + + async def arun(self, *args, **kwargs: Any) -> Any: + """Use the tool asynchronously.""" + resp = await self.mcp_client.call_tool(self.mcp_tool_name, kwargs) + return resp + + @classmethod + def get_mcp_tool(cls, name: str, description: str, mcp_client: BaseMcpClient, + mcp_tool_name: str, arg_schema: Any, **kwargs) -> StructuredTool: + """Get a tool from the class.""" + c = cls(name=name, description=description, mcp_client=mcp_client, + mcp_tool_name=mcp_tool_name) + return StructuredTool(name=c.name, + description=c.description, + func=c.run, + coroutine=c.arun, + args_schema=arg_schema, + **kwargs) diff --git a/src/backend/bisheng/mcp_manage/manager.py b/src/backend/bisheng/mcp_manage/manager.py new file mode 100644 index 000000000..b0403d12a --- /dev/null +++ b/src/backend/bisheng/mcp_manage/manager.py @@ -0,0 +1,51 @@ +import json + +from bisheng.mcp_manage.clients.base import BaseMcpClient +from bisheng.mcp_manage.clients.sse import SseClient +from bisheng.mcp_manage.clients.stdio import StdioClient +from bisheng.mcp_manage.constant import McpClientType + + +class ClientManager: + + @classmethod + async def connect_mcp_from_json(cls, client_json: dict | str): + """ 获取对应配置下的的mcp连接 """ + return cls.sync_connect_mcp_from_json(client_json) + + @classmethod + def sync_connect_mcp_from_json(cls, client_json: dict | str) -> BaseMcpClient: + """ 获取对应配置下的的mcp连接 """ + if isinstance(client_json, str): + client_json = json.loads(client_json) + + mcp_servers = client_json['mcpServers'] + client_type = McpClientType.SSE.value + client_kwargs = {} + + for _, kwargs in mcp_servers.items(): + if 'command' in kwargs: + client_type = McpClientType.STDIO.value + kwargs.pop('name', '') + kwargs.pop('description', '') + client_kwargs = kwargs + break + return cls.sync_connect_mcp(client_type, **client_kwargs) + + + @classmethod + async def connect_mcp(cls, client_type: str, **kwargs) -> BaseMcpClient: + """ 获取对应url的mcp连接 """ + # 初始化对应的client + return cls.sync_connect_mcp(client_type, **kwargs) + + @classmethod + def sync_connect_mcp(cls, client_type: str, **kwargs) -> BaseMcpClient: + # 初始化对应的client + if client_type == McpClientType.SSE.value: + client = SseClient(**kwargs) + elif client_type == McpClientType.STDIO.value: + client = StdioClient(**kwargs) + else: + raise ValueError(f'client_type {client_type} not supported') + return client diff --git a/src/backend/bisheng/patches/fastapi_jwt_auth.patch b/src/backend/bisheng/patches/fastapi_jwt_auth.patch new file mode 100644 index 000000000..f8278ce14 --- /dev/null +++ b/src/backend/bisheng/patches/fastapi_jwt_auth.patch @@ -0,0 +1,22 @@ +12c12 +< authjwt_token_location: Optional[Sequence[StrictStr]] = {'headers'} +--- +> authjwt_token_location: Optional[List[StrictStr]] = {'headers'} +21c21 +< authjwt_decode_audience: Optional[Union[StrictStr,Sequence[StrictStr]]] = None +--- +> authjwt_decode_audience: Optional[Union[StrictStr,List[StrictStr]]] = None +23c23 +< authjwt_denylist_token_checks: Optional[Sequence[StrictStr]] = {'access','refresh'} +--- +> authjwt_denylist_token_checks: Optional[List[StrictStr]] = {'access','refresh'} +45c45 +< authjwt_csrf_methods: Optional[Sequence[StrictStr]] = {'POST','PUT','PATCH','DELETE'} +--- +> authjwt_csrf_methods: Optional[List[StrictStr]] = {'POST','PUT','PATCH','DELETE'} +84,85c84,85 +< min_anystr_length = 1 +< anystr_strip_whitespace = True +--- +> str_min_length = 1 +> str_strip_whitespace = True diff --git a/src/backend/bisheng/processing/process.py b/src/backend/bisheng/processing/process.py index 29cc668c2..52f8591a9 100644 --- a/src/backend/bisheng/processing/process.py +++ b/src/backend/bisheng/processing/process.py @@ -146,7 +146,7 @@ def generate_result(langchain_object: Union[Chain, VectorStore], inputs: dict): class Result(BaseModel): - result: Any + result: Any = None session_id: str diff --git a/src/backend/bisheng/restructure/assistants/agent.py b/src/backend/bisheng/restructure/assistants/agent.py deleted file mode 100644 index d1e469307..000000000 --- a/src/backend/bisheng/restructure/assistants/agent.py +++ /dev/null @@ -1,145 +0,0 @@ -import typing as tp -from pathlib import Path - -from langchain_core.messages import AIMessage -from langchain_core.runnables import RunnableConfig -from langchain_core.tools import BaseTool, Tool - -from bisheng.api.services.assistant_agent import AssistantAgent -from bisheng.api.utils import build_flow_no_yield -from bisheng.api.v1.schemas import InputRequest -from bisheng.cache.flow import InMemoryCache -from bisheng.cache.utils import CACHE_DIR -from bisheng.database.models.assistant import Assistant, AssistantLink, AssistantLinkDao -from bisheng.database.models.flow import FlowDao, FlowStatus -from bisheng.database.models.gpts_tools import GptsTools, GptsToolsDao -from bisheng.database.models.knowledge import KnowledgeDao -from bisheng.utils.logger import logger - - -class RtcAssistantAgent(AssistantAgent): - PersistPath = Path(CACHE_DIR) / 'assistant' - MemoryCache = InMemoryCache() - - def __init__(self, assistant: Assistant): - super().__init__(assistant, '') - - async def async_init(self): - await self.init_llm() - await self.init_abilities() - await self.init_agent() - - @classmethod - async def create(cls, assistant: Assistant) -> "RtcAssistantAgent": - if ins := cls.load(assistant.id): - return ins - else: - ins = cls(assistant) - await ins.async_init() - ins.save() - return ins - - @classmethod - def load(cls, assistant_id: str): - return cls.MemoryCache.get(assistant_id) - - def save(self): - self.MemoryCache.set(self.assistant.id, self) - - async def init_abilities(self): - """初始化能力""" - # link要么是tool(非零),要么是flow(非空),要么是knowledge(非零) - logger.info(f"start initialize agent abilities...") - links = AssistantLinkDao.get_assistant_link(assistant_id=self.assistant.id) - tool_links, flows_links, knowledge_links = self.split_tool_flow_knowledge(links) - tools = await self.init_tools(tool_links) - flows = await self.init_flows(flows_links) - knowledge = await self.init_knowledge(knowledge_links) - self.tools = tools + flows + knowledge # init_agent用 - - def split_tool_flow_knowledge(self, links: tp.List[AssistantLink]): - """ - link要么是tool(非零),要么是flow(非空),要么是knowledge(非零) - 应考虑分表 todo - """ - tools = [] - flows = [] - knowledge = [] - for link in links: - if link.tool_id: - tools.append(link) - elif link.flow_id: - flows.append(link) - elif link.knowledge_id: - knowledge.append(link) - return tools, flows, knowledge - - def split_preset_personal(self, tools: tp.List[GptsTools]): - """区分预定义tool和自定义tool""" - preset_tools = [] - personal_tools = [] - for one in tools: - if one.is_preset: - preset_tools.append(one) - else: - personal_tools.append(one) - return preset_tools, personal_tools - - async def init_tools(self, tool_links: tp.List[AssistantLink]) -> tp.List[BaseTool]: - """初始化工具""" - logger.info(f"start initialize agent abilities: tools...") - tool_chains: tp.List[BaseTool] = [] - tool_ids = [one.tool_id for one in tool_links] - instances = GptsToolsDao.get_list_by_ids(tool_ids) - preset_tools, personal_tools = self.split_preset_personal(instances) - if preset_tools: - preset_langchain = await self.init_preset_tools(preset_tools) - logger.info('act=build_preset_tools size={} return_tools={}', len(preset_tools), len(preset_langchain)) - tool_chains.extend(preset_langchain) - if personal_tools: - personal_langchain = await self.init_personal_tools(personal_tools) - logger.info('act=build_personal_tools size={} return_tools={}', len(personal_tools), - len(personal_langchain)) - tool_chains.extend(personal_langchain) - return tool_chains - - async def init_flows(self, flow_links: tp.List[AssistantLink]) -> tp.List[BaseTool]: - """初始化技能""" - logger.info(f"start initialize agent abilities: flows...") - flow_chains = [] - flow_ids = [one.flow_id for one in flow_links] - flow_data = FlowDao.get_flow_by_ids(flow_ids) - for datum in flow_data: - if datum.status != FlowStatus.ONLINE.value: - logger.warning('act=init_tools skip not online flow_id: {}', datum.id) - continue - tool_description = f'{datum.name}:{datum.description}' - fake_chat_id = self.assistant.id - graph = await build_flow_no_yield(graph_data=datum.data, artifacts={}, process_file=True, flow_id=datum.id, - chat_id=fake_chat_id) - built_obj = await graph.abuild() - logger.info('act=init_flow_tool build_end') - tool_name = f'flow_{datum.id}' - chain = Tool(name=tool_name, func=built_obj, coroutine=built_obj.acall, description=tool_description, - args_schema=InputRequest) - flow_chains.append(chain) - return flow_chains - - async def init_knowledge(self, knowledge_links: tp.List[AssistantLink]) -> tp.List[BaseTool]: - """初始化知识库""" - logger.info(f"start initialize agent abilities: knowledge...") - knowledge_chains = [] - knowledge_ids = [one.knowledge_id for one in knowledge_links] - knowledge_data = KnowledgeDao.get_list_by_ids(knowledge_ids) - for datum in knowledge_data: - knowledge_tool = await self.init_knowledge_tool(datum) - knowledge_chains.extend(knowledge_tool) - return knowledge_chains - - async def run_agent(self, inputs: list): - """运行智能体对话""" - result = await self.agent.ainvoke(inputs, config=RunnableConfig()) - # result包含了history,最后一个是最后一次回答 - if isinstance(result[-1], AIMessage): - return result[-1].content - return "" diff --git a/src/backend/bisheng/restructure/assistants/routers.py b/src/backend/bisheng/restructure/assistants/routers.py deleted file mode 100644 index 972378326..000000000 --- a/src/backend/bisheng/restructure/assistants/routers.py +++ /dev/null @@ -1,5 +0,0 @@ -from fastapi import APIRouter -from bisheng.restructure.assistants.views import router - -assistant_router = APIRouter(prefix='/api/v1') -assistant_router.include_router(router) diff --git a/src/backend/bisheng/restructure/assistants/schemas.py b/src/backend/bisheng/restructure/assistants/schemas.py deleted file mode 100644 index 63c79864e..000000000 --- a/src/backend/bisheng/restructure/assistants/schemas.py +++ /dev/null @@ -1,16 +0,0 @@ -from pydantic import BaseModel - - -class ChatInput(BaseModel): - assistant_id: str - message: str - chat_id: str = None - user_id: str = None - - -class StreamMsg(BaseModel): - event: str - data: str = "" - - def __str__(self) -> str: - return f'event: {self.event}\ndata: {self.data}\n\n' diff --git a/src/backend/bisheng/restructure/assistants/services.py b/src/backend/bisheng/restructure/assistants/services.py deleted file mode 100644 index 1f72b587d..000000000 --- a/src/backend/bisheng/restructure/assistants/services.py +++ /dev/null @@ -1,62 +0,0 @@ -from langchain_core.messages import AIMessage, HumanMessage - -from bisheng.database.models.message import ChatMessage, ChatMessageDao -from bisheng.restructure.assistants.agent import RtcAssistantAgent -from bisheng.settings import settings -from bisheng.utils.logger import logger - -DEFAULT_SIZE = 20 - - -class MsgCategory: - Question = 'question' - Answer = 'answer' - - -class MsgFrom: - Human = 1 - Bot = 2 - - -def get_chat_history(chat_id: str, size: int = DEFAULT_SIZE): - chat_history = [] - messages = ChatMessageDao.get_messages_by_chat_id(chat_id, ['question', 'answer'], size) - for one in messages: - # bug fix When constructing multi-turn dialogues, the input and response of - # the user and the assistant were reversed, leading to incorrect question-and-answer sequences. - if one.category == MsgCategory.Question: - chat_history.append(HumanMessage(content=one.message)) - elif one.category == MsgCategory.Answer: - chat_history.append(AIMessage(content=one.message)) - logger.info(f"loaded {len(chat_history)} chat history for chat_id {chat_id}") - return chat_history - - -async def chat_by_agent(agent: RtcAssistantAgent, message: str, chat_history: list): - new_message = HumanMessage(content=message) - if chat_history: - chat_history.append(new_message) - inputs = chat_history - else: - inputs = [new_message] - logger.info(f"start calling langchain agent...") - answer = await agent.run_agent(inputs) - # todo: 后续优化代码解释器的实现方案,保证输出的文件可以公开访问 - # 获取minio的share地址,把share域名去掉, 为毕昇的部署方案特殊处理下 - if gpts_tool_conf := settings.get_from_db('gpts').get('tools'): - if bisheng_code_conf := gpts_tool_conf.get("bisheng_code_interpreter"): - answer = answer.replace(f"http://{bisheng_code_conf['minio']['MINIO_SHAREPOIN']}", "") - return answer - - -def record_message(chat_id: str, user_id: str, msg_from: str, message: str, category: str): - return ChatMessageDao.insert_one(ChatMessage( - is_bot=msg_from == MsgFrom.Bot, - source=False, - message=message, - category=category, - type=msg_from, - flow_id=chat_id, # todo 智能体不是单独的flow - chat_id=chat_id, - user_id=user_id, - )) diff --git a/src/backend/bisheng/restructure/assistants/views.py b/src/backend/bisheng/restructure/assistants/views.py deleted file mode 100644 index 8335d617d..000000000 --- a/src/backend/bisheng/restructure/assistants/views.py +++ /dev/null @@ -1,45 +0,0 @@ -from uuid import uuid4 -from bisheng.database.models.assistant import AssistantDao, AssistantStatus -from bisheng.restructure.assistants.agent import RtcAssistantAgent -from bisheng.restructure.assistants.schemas import StreamMsg, ChatInput -from bisheng.restructure.assistants.services import MsgCategory, MsgFrom, chat_by_agent, get_chat_history, \ - record_message -from bisheng.restructure.logger import log_trace -from bisheng.utils.logger import logger -from fastapi import APIRouter -from fastapi.responses import ORJSONResponse, StreamingResponse - -router = APIRouter(prefix='/assistant') - - -@router.post('/sse', status_code=200, response_class=StreamingResponse) -@log_trace -async def chat(chat_input: ChatInput): - message = chat_input.message - user_id = chat_input.user_id - - async def _event_stream(): - if chat_id := chat_input.chat_id: - chat_history = get_chat_history(chat_id) - else: - chat_history = [] - chat_id = uuid4().hex - yield str(StreamMsg(event='chat', data=chat_id)) - assistant = AssistantDao.get_one_assistant(chat_input.assistant_id) - if not assistant: - yield str(StreamMsg(event='error', data="该助手已被删除")) - return - if assistant.status != AssistantStatus.ONLINE.value: - yield str(StreamMsg(event='error', data="当前助手未上线,无法直接对话")) - return - record_message(chat_id, user_id, MsgFrom.Human, message, MsgCategory.Question) - gpt_agent = await RtcAssistantAgent.create(assistant) - answer = await chat_by_agent(gpt_agent, message, chat_history) - record_message(chat_id, user_id, MsgFrom.Bot, answer, MsgCategory.Answer) - yield str(StreamMsg(event='message', data=answer)) - - try: - return StreamingResponse(_event_stream(), media_type='text/event-stream') - except Exception as exc: - logger.error(exc) - return ORJSONResponse(status_code=500, content=str(exc)) diff --git a/src/backend/bisheng/restructure/logger.py b/src/backend/bisheng/restructure/logger.py deleted file mode 100644 index 1db31e5f1..000000000 --- a/src/backend/bisheng/restructure/logger.py +++ /dev/null @@ -1,17 +0,0 @@ -# log工具 -from functools import wraps -from uuid import uuid4 - -from bisheng.utils.logger import logger - - -def log_trace(func): - """添加trace_id用于log追踪""" - - @wraps(func) - async def wrapper(*args, **kwargs): - trace_id = uuid4().hex - with logger.contextualize(trace_id=trace_id): - return await func(*args, **kwargs) - - return wrapper diff --git a/src/backend/bisheng/restructure/register.py b/src/backend/bisheng/restructure/register.py deleted file mode 100644 index 16598a4ef..000000000 --- a/src/backend/bisheng/restructure/register.py +++ /dev/null @@ -1,7 +0,0 @@ -from fastapi import FastAPI - -from bisheng.restructure.routers import router as router_restructure - - -def register_restructure(app: FastAPI): - app.include_router(router_restructure) diff --git a/src/backend/bisheng/restructure/routers.py b/src/backend/bisheng/restructure/routers.py deleted file mode 100644 index 148ff2ccf..000000000 --- a/src/backend/bisheng/restructure/routers.py +++ /dev/null @@ -1,6 +0,0 @@ -from fastapi import APIRouter - -from bisheng.restructure.assistants.routers import assistant_router - -router = APIRouter() -router.include_router(assistant_router) diff --git a/src/backend/bisheng/services/database/models/api_key/model.py b/src/backend/bisheng/services/database/models/api_key/model.py index 3a91c0ae2..c77682b59 100644 --- a/src/backend/bisheng/services/database/models/api_key/model.py +++ b/src/backend/bisheng/services/database/models/api_key/model.py @@ -2,7 +2,7 @@ from datetime import datetime from typing import Optional from uuid import UUID, uuid4 -from pydantic import validator +from pydantic import validator, field_validator from sqlmodel import Field, SQLModel @@ -42,7 +42,8 @@ class ApiKeyRead(ApiKeyBase): api_key: str = Field() user_id: UUID = Field() - @validator('api_key', always=True) + @field_validator('api_key', mode='before') + @classmethod def mask_api_key(cls, v): # This validator will always run, and will mask the API key return f"{v[:8]}{'*' * (len(v) - 8)}" diff --git a/src/backend/bisheng/services/database/models/flow/model.py b/src/backend/bisheng/services/database/models/flow/model.py index 8f20770fb..95c0de23d 100644 --- a/src/backend/bisheng/services/database/models/flow/model.py +++ b/src/backend/bisheng/services/database/models/flow/model.py @@ -20,7 +20,8 @@ class FlowBase(SQLModel): folder: Optional[str] = Field(default=None, nullable=True) @field_validator('data') - def validate_json(v): + @classmethod + def validate_json(cls, v): if not v: return v if not isinstance(v, dict): @@ -42,6 +43,7 @@ class FlowBase(SQLModel): return dt.isoformat() @field_validator('updated_at', mode='before') + @classmethod def validate_dt(cls, v): if v is None: return v diff --git a/src/backend/bisheng/services/settings/auth.py b/src/backend/bisheng/services/settings/auth.py index b8206620d..7880c1681 100644 --- a/src/backend/bisheng/services/settings/auth.py +++ b/src/backend/bisheng/services/settings/auth.py @@ -6,8 +6,8 @@ from bisheng.services.settings.constants import DEFAULT_SUPERUSER, DEFAULT_SUPER from bisheng.services.settings.utils import read_secret_from_file, write_secret_to_file from loguru import logger from passlib.context import CryptContext -from pydantic import Field, validator -from pydantic_settings import BaseSettings +from pydantic import Field, field_validator +from pydantic_settings import SettingsConfigDict, BaseSettings class AuthSettings(BaseSettings): @@ -35,35 +35,14 @@ class AuthSettings(BaseSettings): SUPERUSER_PASSWORD: str = DEFAULT_SUPERUSER_PASSWORD pwd_context: CryptContext = CryptContext(schemes=['bcrypt'], deprecated='auto') - - class Config: - validate_assignment = True - extra = 'ignore' - env_prefix = 'bisheng_' + model_config = SettingsConfigDict(validate_assignment=True, extra='ignore', env_prefix='bisheng_') def reset_credentials(self): self.SUPERUSER = DEFAULT_SUPERUSER self.SUPERUSER_PASSWORD = DEFAULT_SUPERUSER_PASSWORD - # If autologin is true, then we need to set the credentials to - # the default values - # so we need to validate the superuser and superuser_password - # fields - @validator('SUPERUSER', 'SUPERUSER_PASSWORD', pre=True) - def validate_superuser(cls, value, values): - if values.get('AUTO_LOGIN'): - if value != DEFAULT_SUPERUSER: - value = DEFAULT_SUPERUSER - logger.debug('Resetting superuser to default value') - if values.get('SUPERUSER_PASSWORD') != DEFAULT_SUPERUSER_PASSWORD: - values['SUPERUSER_PASSWORD'] = DEFAULT_SUPERUSER_PASSWORD - logger.debug('Resetting superuser password to default value') - - return value - - return value - - @validator('SECRET_KEY', pre=True) + @field_validator('SECRET_KEY', mode='before') + @classmethod def get_secret_key(cls, value, values): config_dir = values.get('CONFIG_DIR') diff --git a/src/backend/bisheng/services/settings/base.py b/src/backend/bisheng/services/settings/base.py index 15704dc6c..407580038 100644 --- a/src/backend/bisheng/services/settings/base.py +++ b/src/backend/bisheng/services/settings/base.py @@ -62,7 +62,8 @@ class Settings(BaseSettings): ] = 'https://api.langflow.store/flows/trigger/ec611a61-8460-4438-b187-a4f65e5559d4' LIKE_WEBHOOK_URL: Optional[str] = 'https://api.langflow.store/flows/trigger/64275852-ec00-45c1-984e-3bff814732da' - @validator('CONFIG_DIR', pre=True, allow_reuse=True) + @field_validator('CONFIG_DIR', mode="before") + @classmethod def set_langflow_dir(cls, value): if not value: from platformdirs import user_cache_dir @@ -85,7 +86,8 @@ class Settings(BaseSettings): return str(value) - @validator('DATABASE_URL', pre=True) + @field_validator('DATABASE_URL', mode='before') + @classmethod def set_database_url(cls, value, values): if not value: logger.debug('No database_url provided, trying LANGFLOW_DATABASE_URL env variable') @@ -118,6 +120,7 @@ class Settings(BaseSettings): return value @field_validator('COMPONENTS_PATH', mode='before') + @classmethod def set_components_path(cls, value): if os.getenv('LANGFLOW_COMPONENTS_PATH'): logger.debug('Adding LANGFLOW_COMPONENTS_PATH to components_path') diff --git a/src/backend/bisheng/services/store/schema.py b/src/backend/bisheng/services/store/schema.py index dbceaa08e..2cecae4e1 100644 --- a/src/backend/bisheng/services/store/schema.py +++ b/src/backend/bisheng/services/store/schema.py @@ -1,17 +1,17 @@ from typing import List, Optional from uuid import UUID -from pydantic import BaseModel, validator +from pydantic import field_validator, BaseModel class TagResponse(BaseModel): id: UUID - name: Optional[str] + name: Optional[str] = None class UsersLikesResponse(BaseModel): - likes_count: Optional[int] - liked_by_user: Optional[bool] + likes_count: Optional[int] = None + liked_by_user: Optional[bool] = None class CreateComponentResponse(BaseModel): @@ -19,7 +19,7 @@ class CreateComponentResponse(BaseModel): class TagsIdResponse(BaseModel): - tags_id: Optional[TagResponse] + tags_id: Optional[TagResponse] = None class ListComponentResponse(BaseModel): @@ -37,7 +37,8 @@ class ListComponentResponse(BaseModel): private: Optional[bool] = None # tags comes as a TagsIdResponse but we want to return a list of TagResponse - @validator('tags', pre=True) + @field_validator('tags', mode="before") + @classmethod def tags_to_list(cls, v): # Check if all values are have id and name # if so, return v else transform to TagResponse @@ -52,24 +53,24 @@ class ListComponentResponse(BaseModel): class ListComponentResponseModel(BaseModel): count: Optional[int] = 0 authorized: bool - results: Optional[List[ListComponentResponse]] + results: Optional[List[ListComponentResponse]] = None class DownloadComponentResponse(BaseModel): id: UUID - name: Optional[str] - description: Optional[str] - data: Optional[dict] - is_component: Optional[bool] + name: Optional[str] = None + description: Optional[str] = None + data: Optional[dict] = None + is_component: Optional[bool] = None metadata: Optional[dict] = {} class StoreComponentCreate(BaseModel): name: str - description: Optional[str] + description: Optional[str] = None data: dict - tags: Optional[List[str]] + tags: Optional[List[str]] = None parent: Optional[UUID] = None - is_component: Optional[bool] + is_component: Optional[bool] = None last_tested_version: Optional[str] = None private: Optional[bool] = True diff --git a/src/backend/bisheng/settings.py b/src/backend/bisheng/settings.py index 08aade354..f822a8b76 100644 --- a/src/backend/bisheng/settings.py +++ b/src/backend/bisheng/settings.py @@ -5,16 +5,15 @@ from typing import Dict, List, Optional, Union import yaml from cryptography.fernet import Fernet -from langchain.pydantic_v1 import BaseSettings, root_validator, validator from loguru import logger -from pydantic import BaseModel, Field +from pydantic import ConfigDict, BaseModel, Field, field_validator, model_validator from sqlmodel import select class LoggerConf(BaseModel): level: str = 'DEBUG' format: str = '[{level.name} process-{process.id}-{thread.id} {name}:{line}] - trace={extra[trace_id]} {message}' # noqa - handlers: List[Dict] = [] + handlers: List[Dict] = Field(default_factory=list, description='日志处理器') @classmethod def parse_logger_sink(cls, sink: str) -> str: @@ -26,7 +25,7 @@ class LoggerConf(BaseModel): env_keys[one] = os.getenv(one, '') return sink.format(**env_keys) - @validator('handlers', pre=True) + @field_validator('handlers') @classmethod def set_handlers(cls, value): if value is None: @@ -88,10 +87,7 @@ class WorkflowConf(BaseModel): class Settings(BaseModel): - class Config: - validate_assignment = True - arbitrary_types_allowed = True - extra = 'ignore' + model_config = ConfigDict(validate_assignment=True, arbitrary_types_allowed=True, extra='ignore') chains: dict = {} agents: dict = {} @@ -133,7 +129,8 @@ class Settings(BaseModel): object_storage: ObjectStore = {} workflow_conf: WorkflowConf = WorkflowConf() - @validator('database_url', pre=True) + @field_validator('database_url') + @classmethod def set_database_url(cls, value): if not value: logger.debug('No database_url provided, trying bisheng_DATABASE_URL env variable') @@ -155,7 +152,8 @@ class Settings(BaseModel): return value - @root_validator() + @model_validator(mode='before') + @classmethod def set_redis_url(cls, values): if 'redis_url' in values: if isinstance(values['redis_url'], dict): @@ -174,7 +172,8 @@ class Settings(BaseModel): values['redis_url'] = new_redis_url return values - @root_validator() + @model_validator(mode='before') + @classmethod def set_celery_redis_url(cls, values): if 'celery_redis_url' in values: if isinstance(values['celery_redis_url'], dict): @@ -193,7 +192,8 @@ class Settings(BaseModel): values['celery_redis_url'] = new_redis_url return values - @root_validator() + @model_validator(mode='before') + @classmethod def validate_lists(cls, values): for key, value in values.items(): if key != 'dev' and not value: diff --git a/src/backend/bisheng/template/field/base.py b/src/backend/bisheng/template/field/base.py index 5be3b5ef2..7ecfb17d8 100644 --- a/src/backend/bisheng/template/field/base.py +++ b/src/backend/bisheng/template/field/base.py @@ -30,7 +30,7 @@ class TemplateField(BaseModel): suffixes: list[str] = [] fileTypes: list[str] = [] - file_types: list[str] = Field(default=[], serialization_alias='fileTypes') + file_types: list[str] = Field(default_factory=list, serialization_alias='fileTypes') """List of file types associated with the field. Default is an empty list. (duplicate)""" file_path: Optional[str] = '' diff --git a/src/backend/bisheng/utils/__init__.py b/src/backend/bisheng/utils/__init__.py index 0441bc625..8bfd9e2ce 100644 --- a/src/backend/bisheng/utils/__init__.py +++ b/src/backend/bisheng/utils/__init__.py @@ -1,3 +1,4 @@ +import hashlib import uuid @@ -5,4 +6,10 @@ def generate_uuid() -> str: """ 生成uuid的字符串 """ - return uuid.uuid4().hex \ No newline at end of file + return uuid.uuid4().hex + + +def md5_hash(original_string: str): + md5 = hashlib.md5() + md5.update(original_string.encode('utf-8')) + return md5.hexdigest() diff --git a/src/backend/bisheng/utils/docx_temp.py b/src/backend/bisheng/utils/docx_temp.py index 8ba6d9c46..8c3b07370 100644 --- a/src/backend/bisheng/utils/docx_temp.py +++ b/src/backend/bisheng/utils/docx_temp.py @@ -220,7 +220,14 @@ class DocxTemplateRender(object): for i, row in enumerate(table.rows): for j, cell in enumerate(row.cells): if k1 in cell.text: - table.rows[i].cells[j].text = cell.text.replace(k1, v1) + for one in cell.paragraphs: + if k1 in one.text: + new_text = one.text.replace(k1, v1) + one.runs[0].text = new_text + for r_index, r in enumerate(one.runs): + if r_index == 0: + continue + r.text = '' for p in doc.paragraphs: # p.text = p.text.replace(k1, v1) diff --git a/src/backend/bisheng/utils/http_middleware.py b/src/backend/bisheng/utils/http_middleware.py index e2c847520..7797bc9d7 100644 --- a/src/backend/bisheng/utils/http_middleware.py +++ b/src/backend/bisheng/utils/http_middleware.py @@ -12,7 +12,10 @@ class CustomMiddleware(BaseHTTPMiddleware): async def dispatch(self, request: Request, call_next: RequestResponseEndpoint): # You can modify the request before passing it to the next middleware or endpoint - trace_id = str(uuid4().hex) + if request.headers.get('x-trace-id'): + trace_id = request.headers.get('x-trace-id') + else: + trace_id = str(uuid4().hex) start_time = time() with logger.contextualize(trace_id=trace_id): logger.info(f'{request.method} {request.url.path}') diff --git a/src/backend/bisheng/utils/util.py b/src/backend/bisheng/utils/util.py index 6c5b6a5c9..3822000f1 100644 --- a/src/backend/bisheng/utils/util.py +++ b/src/backend/bisheng/utils/util.py @@ -5,6 +5,8 @@ from functools import wraps from typing import Dict, Optional from urllib.parse import urlparse +from pydantic import BaseModel + from bisheng.template.frontend_node.constants import FORCE_SHOW_FIELDS from bisheng.utils import constants from docstring_parser import parse # type: ignore @@ -70,8 +72,8 @@ def build_template_from_class(name: str, type_to_cls_dict: Dict, add_function: b variables = {'_type': _type} - if '__fields__' in _class.__dict__: - for class_field_items, value in _class.__fields__.items(): + if getattr(_class, 'model_fields', None): + for class_field_items, value in _class.model_fields.items(): if class_field_items in ['callback_manager']: continue variables[class_field_items] = {} @@ -87,6 +89,16 @@ def build_template_from_class(name: str, type_to_cls_dict: Dict, add_function: b variables[class_field_items]['placeholder'] = ( docs.params[class_field_items] if class_field_items in docs.params else '') + else: + for name, param in inspect.signature(_class.__init__).parameters.items(): + if name == 'self': + continue + variables[name] = {} + variables[name]['default'] = get_default_factory(module=_class.__base__.__module__, function=str(param.annotation)) + variables[name]['annotation'] = str(param.annotation) + variables[name]['required'] = False + + base_classes = get_base_classes(_class) # Adding function to base classes to allow # the output to be a function @@ -97,6 +109,7 @@ def build_template_from_class(name: str, type_to_cls_dict: Dict, add_function: b 'description': docs.short_description or '', 'base_classes': base_classes, } + return None def build_template_from_method( @@ -211,13 +224,13 @@ def format_dict(d, name: Optional[str] = None): Returns: A new dictionary with the desired modifications applied. """ - + need_remove_key = [] # Process remaining keys for key, value in d.items(): if key == '_type': continue - _type = value['type'] + _type = value['type'] if 'type' in value else value['annotation'] if not isinstance(_type, str): _type = type_to_string(_type) @@ -293,6 +306,13 @@ def format_dict(d, name: Optional[str] = None): value['options'] = constants.ANTHROPIC_MODELS value['list'] = True value['value'] = constants.ANTHROPIC_MODELS[0] + + if 'value' in value and type(value['value']) == set: + value['value'] = list(value['value']) + if 'value' in value and inspect.isfunction(value['value']): + need_remove_key.append(key) + for one in need_remove_key: + del d[one] return d diff --git a/src/backend/bisheng/worker/workflow/redis_callback.py b/src/backend/bisheng/worker/workflow/redis_callback.py index 11e505616..d7c5f2a90 100644 --- a/src/backend/bisheng/worker/workflow/redis_callback.py +++ b/src/backend/bisheng/worker/workflow/redis_callback.py @@ -1,3 +1,4 @@ +import os import asyncio import json import time @@ -152,7 +153,7 @@ class RedisCallback(BaseCallback): break yield chat_response break - elif status_info['status'] == WorkflowStatus.WAITING.value and time.time() - status_info['time'] > 10: + elif status_info['status'] in [WorkflowStatus.WAITING.value, WorkflowStatus.INPUT_OVER.value] and time.time() - status_info['time'] > 10: # 10秒内没有收到状态更新,说明workflow没有启动,可能是celery worker线程数已满 self.set_workflow_status(WorkflowStatus.FAILED.value, 'workflow task execute busy') yield self.build_chat_response(WorkflowEventType.Error.value, 'over', @@ -200,8 +201,12 @@ class RedisCallback(BaseCallback): for key_info in old_message['input_schema']['value']: user_input_message += f"{key_info['value']}:{user_input.get(key_info['key'], '')}\n" else: - # 说明对话框输入 + # 说明对话框输入, 需要加下上传的文件信息, 和输入节点的数据结构有关 user_input_message = user_input[old_message['input_schema']['key']] + dialog_files_content = user_input.get('dialog_files_content', []) + for one in dialog_files_content: + user_input_message += f"\n{os.path.basename(one).split('?')[0]}" + self.save_chat_message(ChatResponse( message=user_input_message, category='question', @@ -259,7 +264,8 @@ class RedisCallback(BaseCallback): is_bot=chat_response.is_bot, source=chat_response.source, - message=json.dumps(chat_response.message, ensure_ascii=False), + message=chat_response.message if isinstance(chat_response.message, str) else json.dumps( + chat_response.message, ensure_ascii=False), extra=chat_response.extra, category=chat_response.category, files=json.dumps(chat_response.files, ensure_ascii=False) diff --git a/src/backend/bisheng/workflow/callback/event.py b/src/backend/bisheng/workflow/callback/event.py index e98abb18a..c9843400d 100644 --- a/src/backend/bisheng/workflow/callback/event.py +++ b/src/backend/bisheng/workflow/callback/event.py @@ -63,4 +63,4 @@ class StreamMsgData(BaseModel): class StreamMsgOverData(StreamMsgData): - source_documents: Optional[List[Any]] = Field([], description='Source documents') + source_documents: Optional[List[Any]] = Field(default=[], description='Source documents') diff --git a/src/backend/bisheng/workflow/callback/llm_callback.py b/src/backend/bisheng/workflow/callback/llm_callback.py index 77724611c..6de211a58 100644 --- a/src/backend/bisheng/workflow/callback/llm_callback.py +++ b/src/backend/bisheng/workflow/callback/llm_callback.py @@ -52,12 +52,13 @@ class LLMNodeCallbackHandler(BaseCallbackHandler): async def on_tool_end(self, output: str, **kwargs: Any) -> Any: """Run when tool ends running.""" logger.debug(f'on_tool_end output={output} kwargs={kwargs}') + result = output if isinstance(output, str) else getattr(output, 'content', output) if self.tool_list is not None: self.tool_list.append({ 'type': 'end', 'run_id': kwargs.get('run_id').hex, 'name': kwargs['name'], - 'output': output, + 'output': result, }) if kwargs['name'] == 'sql_agent': self.output = True @@ -97,7 +98,12 @@ class LLMNodeCallbackHandler(BaseCallbackHandler): return if not self.output: return - msg = response.generations[0][0].text + msg = response.generations[0][0].message + # ChatTongYi vl model special text + if isinstance(msg.content, list): + msg = ''.join([one.get('text', '') for one in msg.content]) + else: + msg = msg.content if not msg: logger.warning('LLM output is empty') return diff --git a/src/backend/bisheng/workflow/common/node.py b/src/backend/bisheng/workflow/common/node.py index cafabaff1..226697e42 100644 --- a/src/backend/bisheng/workflow/common/node.py +++ b/src/backend/bisheng/workflow/common/node.py @@ -1,7 +1,8 @@ +import copy from enum import Enum from typing import Optional, Any, List -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, field_validator class NodeType(Enum): @@ -24,7 +25,7 @@ class NodeType(Enum): class NodeParams(BaseModel): key: str = Field(default="", description="变量的key") label: Optional[str] = Field("", description="变量描述文本") - value: Optional[Any] = Field(description="变量的值") + value: Optional[Any] = Field(None, description="变量的值") # 变量类型 -> 数据格式的详情参考 https://dataelem.feishu.cn/wiki/IfBvwwvfFiHjuQkjFJgcxzoGnxb type: Optional[str] = Field("", description="变量类型") @@ -49,11 +50,20 @@ class BaseNodeData(BaseModel): group_params: Optional[List[NodeGroupParams]] = Field(default=None, description="Node group params") tab: Optional[dict] = Field({}, description="tab config") tool_key: Optional[str] = Field("", description="unique tool id, only for tool node") - v: str = Field(default="", description="node version") + v: Optional[int] = Field(default=0, description="node version") + + @field_validator('v', mode='before') + @classmethod + def convert_v_to_int(cls, v: str | int | None) -> int: + if isinstance(v, str): + return int(v) + elif v is None: + return 0 + return v def get_variable_info(self, variable_key: str) -> NodeParams | None: for group_info in self.group_params: for one in group_info.params: if one.key == variable_key: - return one - + return copy.deepcopy(one) + return None diff --git a/src/backend/bisheng/workflow/graph/graph_engine.py b/src/backend/bisheng/workflow/graph/graph_engine.py index 2a0bd0737..24f95b8fe 100644 --- a/src/backend/bisheng/workflow/graph/graph_engine.py +++ b/src/backend/bisheng/workflow/graph/graph_engine.py @@ -20,7 +20,7 @@ from typing_extensions import TypedDict class TempState(TypedDict): # not use, only for langgraph state graph - flag: Annotated[bool, operator.add] + flag: Annotated[bool, operator.and_] class GraphEngine: diff --git a/src/backend/bisheng/workflow/graph/graph_state.py b/src/backend/bisheng/workflow/graph/graph_state.py index 712a53c70..6e0f362b1 100644 --- a/src/backend/bisheng/workflow/graph/graph_state.py +++ b/src/backend/bisheng/workflow/graph/graph_state.py @@ -9,10 +9,10 @@ class GraphState(BaseModel): """ 所有节点的 全局状态管理 """ # 存储聊天历史 - history_memory: Optional[ConversationBufferWindowMemory] + history_memory: Optional[ConversationBufferWindowMemory] = None # 全局变量池 - variables_pool: Dict[str, Dict[str, Any]] = Field(default={}, description='全局变量池: {node_id: {key: value}}') + variables_pool: Dict[str, Dict[str, Any]] = Field(default_factory=dict, description='全局变量池: {node_id: {key: value}}') def get_history_memory(self, count: int) -> str: """ 获取聊天历史记录 diff --git a/src/backend/bisheng/workflow/graph/workflow.py b/src/backend/bisheng/workflow/graph/workflow.py index f782329bd..6a4ed2823 100644 --- a/src/backend/bisheng/workflow/graph/workflow.py +++ b/src/backend/bisheng/workflow/graph/workflow.py @@ -53,7 +53,6 @@ class Workflow: """ # 执行workflow if input_data is not None: - self.save_user_input_history(input_data) self.graph_engine.continue_run(input_data) else: # 首次运行时间 diff --git a/src/backend/bisheng/workflow/nodes/agent/agent.py b/src/backend/bisheng/workflow/nodes/agent/agent.py index 478a6a0cd..1afda559b 100644 --- a/src/backend/bisheng/workflow/nodes/agent/agent.py +++ b/src/backend/bisheng/workflow/nodes/agent/agent.py @@ -38,6 +38,8 @@ class AgentNode(BaseNode): self._user_prompt = PromptTemplateParser(template=self.node_params['user_prompt']) self._user_variables = self._user_prompt.extract() + self._image_prompt = self.node_params.get('image_prompt', []) + self._batch_variable_list = [] self._system_prompt_list = [] self._user_prompt_list = [] @@ -74,12 +76,10 @@ class AgentNode(BaseNode): self._sql_address = f'mysql+pymysql://{self._sql_agent["db_username"]}:{self._sql_agent["db_password"]}@{self._sql_agent["db_address"]}/{self._sql_agent["db_name"]}?charset=utf8mb4' # agent - self._agent_executor_type = 'get_react_agent_executor' + self._agent_executor_type = 'React' self._agent = None def _init_agent(self, system_prompt: str): - if self._agent: - return # 获取配置的助手模型列表 assistant_llm = LLMService.get_assistant_llm() if not assistant_llm.llm_list: @@ -350,19 +350,23 @@ class AgentNode(BaseNode): tool_list=tool_invoke_list, cancel_llm_end=True) config = RunnableConfig(callbacks=[llm_callback]) - logger.debug(f'user_prompt: {user}, history: {chat_history}') + human_message = HumanMessage(content=[{ + 'type': 'text', + 'text': user + }]) + human_message = self.contact_file_into_prompt(human_message, self._image_prompt) + chat_history.append(human_message) + logger.debug(f'agent invoke chat_history: {chat_history}') if self._agent_executor_type == 'ReAct': result = self._agent.invoke({ - 'input': user, - 'chat_history': chat_history - }, - config=config) + 'input': chat_history[-1].content, + 'chat_history': chat_history[:-1], + }, config=config) output = result['agent_outcome'].return_values['output'] if isinstance(output, dict): output = list(output.values())[0] return output, llm_callback.reasoning_content else: - chat_history.append(HumanMessage(content=user)) result = self._agent.invoke(chat_history, config=config) return result[-1].content, llm_callback.reasoning_content diff --git a/src/backend/bisheng/workflow/nodes/base.py b/src/backend/bisheng/workflow/nodes/base.py index 3be975636..d8accc14c 100644 --- a/src/backend/bisheng/workflow/nodes/base.py +++ b/src/backend/bisheng/workflow/nodes/base.py @@ -1,14 +1,18 @@ +import base64 import copy import uuid from abc import ABC, abstractmethod from typing import Any, Dict, List +from langchain_core.messages import HumanMessage + from bisheng.utils.exceptions import IgnoreException from bisheng.workflow.callback.base_callback import BaseCallback from bisheng.workflow.callback.event import NodeEndData, NodeStartData from bisheng.workflow.common.node import BaseNodeData, NodeType from bisheng.workflow.edges.edges import EdgeBase from bisheng.workflow.graph.graph_state import GraphState +from bisheng.workflow.nodes.prompt_template import PromptTemplateParser class BaseNode(ABC): @@ -126,6 +130,48 @@ class BaseNode(ABC): next_nodes.append(one.target) return next_nodes + def parse_msg_with_variables(self, msg: str) -> (str, list[str]): + """ + params: + msg: user input msg with node variables + return: + 0: new msg after replaced variable + 1: list of variables node_id.xxxx + """ + msg_template = PromptTemplateParser(template=msg) + variables = msg_template.extract() + if len(variables) > 0: + var_map = {} + for one in variables: + var_map[one] = self.get_other_node_variable(one) + msg = msg_template.format(var_map) + return msg, variables + + def get_file_base64_data(self, file_path: str) -> str: + with open(file_path, "rb") as f: + file_data = f.read() + base64_data = base64.b64encode(file_data).decode('utf-8') + return base64_data + + def contact_file_into_prompt(self, human_message: HumanMessage, variable_list: List[str]) -> HumanMessage: + if not variable_list: + if isinstance(human_message.content, list): + human_message.content = human_message.content[0].get('text') + return human_message + for image_variable in variable_list: + image_value = self.get_other_node_variable(image_variable) + if not image_value: + continue + for file_path in image_value: + base64_image = self.get_file_base64_data(file_path) + human_message.content.append({ + "type": "image", + "source_type": "base64", + "mime_type": "image/jpeg", + "data": base64_image, + }) + return human_message + def run(self, state: dict) -> Any: """ Run node entry diff --git a/src/backend/bisheng/workflow/nodes/condition/conidition_case.py b/src/backend/bisheng/workflow/nodes/condition/conidition_case.py index d4a58e215..6ab001f52 100644 --- a/src/backend/bisheng/workflow/nodes/condition/conidition_case.py +++ b/src/backend/bisheng/workflow/nodes/condition/conidition_case.py @@ -1,7 +1,7 @@ import re from typing import List, Optional, Dict -from pydantic import BaseModel, Field +from pydantic import ConfigDict, BaseModel, Field from loguru import logger from bisheng.workflow.nodes.base import BaseNode @@ -60,8 +60,7 @@ class ConditionOne(BaseModel): class ConditionCases(BaseModel): - class Config: - arbitrary_types_allowed = True + model_config = ConfigDict(arbitrary_types_allowed=True) id: str = Field(..., description='Unique id for case') operator: Optional[str] = Field('and', description='Operator for case') diff --git a/src/backend/bisheng/workflow/nodes/input/input.py b/src/backend/bisheng/workflow/nodes/input/input.py index 487803d1e..c13418e4c 100644 --- a/src/backend/bisheng/workflow/nodes/input/input.py +++ b/src/backend/bisheng/workflow/nodes/input/input.py @@ -3,38 +3,53 @@ import shutil import tempfile from typing import Any +from loguru import logger + from bisheng.api.services.knowledge_imp import decide_vectorstores, read_chunk_text from bisheng.api.services.llm import LLMService from bisheng.api.utils import md5_hash from bisheng.cache.utils import file_download from bisheng.chat.types import IgnoreException from bisheng.workflow.nodes.base import BaseNode -from loguru import logger class InputNode(BaseNode): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - + # 节点当前版本 + self._current_v = 2 # 记录是对话还是表单 self._tab = self.node_data.tab['value'] # 记录这个变量是什么类型的 self._node_params_map = {} new_node_params = {} + # 对话框里输入文件的最大长度,超过这个长度会被截断 + self._dialog_files_length = int(self.node_params.get('dialog_files_content_size', 15000)) + # save image file path + self._dialog_images_files = [] + if self.is_dialog_input(): new_node_params['user_input'] = self.node_params['user_input'] + new_node_params['dialog_files_content'] = self.node_params.get('dialog_files_content', []) else: for value_info in self.node_params['form_input']: new_node_params[value_info['key']] = value_info['value'] self._node_params_map[value_info['key']] = value_info self.node_params = new_node_params - self._file_ids = [] + self._image_ext = ['png', 'jpg', 'jpeg', 'bmp'] + + self._embedding = None + self._vector_client = None + self._es_client = None def is_dialog_input(self): """ 是否是对话形式的输入 """ + if self.node_data.v < self._current_v: + raise IgnoreException(f'{self.name} -- workflow node is update') + if self._tab == 'dialog_input': return True elif self._tab == 'form_input': @@ -43,12 +58,29 @@ class InputNode(BaseNode): def get_input_schema(self) -> Any: if self.is_dialog_input(): - return self.node_data.get_variable_info('user_input') - return self.node_data.get_variable_info('form_input') + user_input_info = self.node_data.get_variable_info('user_input') + user_input_info.value = [ + self.node_data.get_variable_info('dialog_files_content'), + self.node_data.get_variable_info('dialog_file_accept') + ] + return user_input_info + form_input_info = self.node_data.get_variable_info('form_input') + for one in form_input_info.value: + one['value'], _ = self.parse_msg_with_variables(one['value']) + return form_input_info def _run(self, unique_id: str): if self.is_dialog_input(): - return {'user_input': self.node_params['user_input']} + # 对话框形式的输入 + dislog_files_content, self._dialog_images_files = self.parse_dialog_files() + res = { + 'user_input': self.node_params['user_input'], + 'dialog_files_content': dislog_files_content, + 'dialog_image_files': self._dialog_images_files + } + self.graph_state.save_context(content=f'{res["dialog_files_content"]}\n{res["user_input"]}', + msg_sender='human') + return res ret = {} # 表单形式的需要去处理对应的文件上传 @@ -56,21 +88,91 @@ class InputNode(BaseNode): ret[key] = value key_info = self._node_params_map[key] if key_info['type'] == 'file': - new_params = self.parse_upload_file(key,key_info, value) - if key_info['multiple']: - new_params = {key_info['key']: new_params[key_info['key']]} + new_params = self.parse_upload_file(key, key_info, value) ret.update(new_params) return ret def parse_log(self, unique_id: str, result: dict) -> Any: ret = [] - for k,v in result.items(): + for k, v in result.items(): if self._node_params_map.get(k) and self._node_params_map[k]['type'] == 'file': continue ret.append({"key": f'{self.id}.{k}', "value": v, "type": "variable"}) return [ret] + def parse_dialog_files(self) -> (str, list[str]): + """ 获取对话框里上传的文件内容 """ + file_length = 0 + dialog_files_content = "" + image_files_path = [] + if not self.node_params.get('dialog_files_content'): + return dialog_files_content, image_files_path + for file_id in self.node_params['dialog_files_content']: + file_name, file_path, chunks, metadatas = self.get_upload_file_path_content(file_id) + file_ext = file_name.split('.')[-1].lower() + if file_ext in self._image_ext: + image_files_path.append(file_path) + + if file_length >= self._dialog_files_length: + continue + file_content = "\n".join(chunks) + file_content = file_content[:self._dialog_files_length - file_length] + file_length += len(file_content) + + dialog_files_content += f"[file name]: {file_name}\n[file content begin]\n{file_content}\n[file content end]\n" + return dialog_files_content, image_files_path + + def get_upload_file_path_content(self, file_url: str) -> (str, str, list, list): + """ + params: + file_url: upload to minio share url + return: + 0: file name + 1: file path in system + 2: chunks list + 3: metadata list + """ + # 1、获取默认的embedding模型 + if self._embedding is None: + embedding = LLMService.get_knowledge_default_embedding() + if not embedding: + raise Exception('没有配置默认的embedding模型') + self._embedding = embedding + + if self._vector_client is None: + # 2、初始化milvus和es实例 + milvus_collection_name = self.get_milvus_collection_name(getattr(self._embedding, 'model_id')) + self._vector_client = decide_vectorstores(milvus_collection_name, 'Milvus', self._embedding) + self._es_client = decide_vectorstores(self.tmp_collection_name, 'ElasticKeywordsSearch', + self._embedding) + + file_id = md5_hash(f'{file_url}') + filepath, file_name = file_download(file_url) + + # save original file path, because uns will convert file to pdf + original_file_path = os.path.join(tempfile.gettempdir(), f'{file_id}.{file_name.split(".")[-1]}') + shutil.copyfile(filepath, original_file_path) + texts = [] + metadatas = [] + try: + texts, metadatas, _, _ = read_chunk_text(filepath, file_name, + ['\n\n', '\n'], + ['after', 'after'], 1000, 500) + for metadata in metadatas: + metadata.update({ + 'file_id': file_id, + 'knowledge_id': self.workflow_id, + 'extra': '', + 'bbox': '', # 临时文件不能溯源,因为没有持久化存储源文件 + }) + except Exception as e: + logger.exception('parse input node file error') + if str(e).find('类型不支持') == -1: + raise e + + return file_name, original_file_path, texts, metadatas + def parse_upload_file(self, key: str, key_info: dict, value: str) -> dict | None: """ 将文件上传到milvus后 @@ -83,59 +185,58 @@ class InputNode(BaseNode): return { key_info['key']: None, key_info['file_content']: None, - key_info['file_path']: None + key_info['file_path']: None, + key_info['image_file']: None } - - # 1、获取默认的embedding模型 - embeddings = LLMService.get_knowledge_default_embedding() - if not embeddings: - raise Exception('没有配置默认的embedding模型') - - # 2、初始化milvus和es实例 - milvus_collection_name = self.get_milvus_collection_name(getattr(embeddings, 'model_id')) - vector_client = decide_vectorstores(milvus_collection_name, 'Milvus', embeddings) - es_client = decide_vectorstores(self.tmp_collection_name, 'ElasticKeywordsSearch', - embeddings) - - # 3、解析文件 + # 解析文件 all_metadata = [] - texts = [] - original_file_path = '' + all_file_content = '' + original_file_path = [] file_id = md5_hash(f'{key}:{value[0]}') - self._file_ids.append(file_id) + image_files_path = [] + file_content_max_size = int(key_info.get('file_content_size', 15000)) + file_content_length = 0 for one_file_url in value: - filepath, file_name = file_download(one_file_url) - if not original_file_path: - original_file_path = os.path.join(tempfile.gettempdir(), f'{file_id}.{file_name.split(".")[-1]}') - shutil.copyfile(filepath, original_file_path) - texts, metadatas, parse_type, partitions = read_chunk_text(filepath, file_name, - ['\n\n', '\n'], - ['after', 'after'], 1000, 500) - if len(texts) == 0: - raise ValueError('文件解析为空') + file_name, file_path, texts, metadatas = self.get_upload_file_path_content(one_file_url) + original_file_path.append(file_path) + file_ext = file_name.split('.')[-1].lower() + if file_ext in self._image_ext: + image_files_path.append(file_path) - for metadata in metadatas: - metadata.update({ + if file_content_length < file_content_max_size: + file_content = "\n".join(texts) + file_content = file_content[:file_content_max_size - file_content_length] + file_content_length += len(file_content) + all_file_content += f"[file name]: {file_name}\n[file content begin]\n{file_content}\n[file content end]\n" + + if not texts: + continue + + # 同一个变量对应的文件,放在一个file_id里 + for one in metadatas: + one.update({ 'file_id': file_id, 'knowledge_id': self.workflow_id, 'extra': '', 'bbox': '', # 临时文件不能溯源,因为没有持久化存储源文件 }) - # 4、上传到milvus和es - logger.info(f'workflow_add_vectordb file={key} file_name={file_name}') + + # 上传到milvus和es + logger.debug(f'workflow_add_vectordb file={key} file_name={file_name}') # 存入milvus - vector_client.add_texts(texts=texts, metadatas=metadatas) + self._vector_client.add_texts(texts=texts, metadatas=metadatas) - logger.info(f'workflow_add_es file={key} file_name={file_name}') + logger.debug(f'workflow_add_es file={key} file_name={file_name}') # 存入es - es_client.add_texts(texts=texts, metadatas=metadatas) + self._es_client.add_texts(texts=texts, metadatas=metadatas) - logger.info(f'workflow_record_file_metadata file={key} file_name={file_name}') + logger.debug(f'workflow_record_file_metadata file={key} file_name={file_name}') all_metadata.append(metadatas[0]) # 记录文件metadata,其他节点根据metadata数据去检索对应的文件 return { key_info['key']: all_metadata, - key_info['file_content']: "\n".join(texts), - key_info['file_path']: original_file_path + key_info['file_content']: all_file_content, + key_info['file_path']: original_file_path, + key_info['image_file']: image_files_path } diff --git a/src/backend/bisheng/workflow/nodes/llm/llm.py b/src/backend/bisheng/workflow/nodes/llm/llm.py index 8a50fab2e..4a7d0f6b5 100644 --- a/src/backend/bisheng/workflow/nodes/llm/llm.py +++ b/src/backend/bisheng/workflow/nodes/llm/llm.py @@ -19,6 +19,8 @@ class LLMNode(BaseNode): # 是否输出结果给用户 self._output_user = self.node_params.get('output_user', False) + self._image_prompt = self.node_params.get('image_prompt', []) + # 初始化prompt self._system_prompt = PromptTemplateParser(template=self.node_params['system_prompt']) self._system_variables = self._system_prompt.extract() @@ -113,7 +115,15 @@ class LLMNode(BaseNode): inputs = [] if system: inputs.append(SystemMessage(content=system)) - inputs.append(HumanMessage(content=user)) + + human_message = HumanMessage(content=[{ + 'type': 'text', + 'text': user + }]) + human_message = self.contact_file_into_prompt(human_message, self._image_prompt) + inputs.append(human_message) + + logger.debug(f'llm invoke chat_history: {inputs} {self._image_prompt}') result = self._llm.invoke(inputs, config=config) diff --git a/src/backend/bisheng/workflow/nodes/output/output.py b/src/backend/bisheng/workflow/nodes/output/output.py index 4b6415f01..d48560433 100644 --- a/src/backend/bisheng/workflow/nodes/output/output.py +++ b/src/backend/bisheng/workflow/nodes/output/output.py @@ -1,3 +1,4 @@ +import json from typing import Any from bisheng.utils.minio_client import MinioClient @@ -21,8 +22,12 @@ class OutputNode(BaseNode): self._handled_output_result = self._output_result # user input msg - self._output_msg = self.node_params['output_msg']['msg'] - self._output_files = self.node_params['output_msg']['files'] + if 'output_msg' in self.node_params: + _original_output_msg = self.node_params['output_msg'] + else: + _original_output_msg = self.node_params['message'] + self._output_msg = _original_output_msg['msg'] + self._output_files = _original_output_msg['files'] # 替换变量后消息内容 self._parsed_output_msg = '' @@ -35,6 +40,7 @@ class OutputNode(BaseNode): def handle_input(self, user_input: dict) -> Any: # 需要存入state, + self.graph_state.save_context(content=json.dumps(user_input, ensure_ascii=False), msg_sender='human') self._handled_output_result = user_input['output_result'] self.graph_state.set_variable(self.id, 'output_result', user_input['output_result']) @@ -59,6 +65,7 @@ class OutputNode(BaseNode): self.parse_output_msg() self.send_output_msg(unique_id) res = { + 'message': self._parsed_output_msg, 'output_result': self._handled_output_result } return res diff --git a/src/backend/bisheng/workflow/nodes/output/output_fake.py b/src/backend/bisheng/workflow/nodes/output/output_fake.py index 719e74b80..12f40240e 100644 --- a/src/backend/bisheng/workflow/nodes/output/output_fake.py +++ b/src/backend/bisheng/workflow/nodes/output/output_fake.py @@ -1,14 +1,12 @@ from bisheng.workflow.callback.event import NodeEndData -from pydantic import BaseModel, Field +from pydantic import ConfigDict, BaseModel, Field from bisheng.workflow.nodes.output.output import OutputNode class OutputFakeNode(BaseModel): """ 用来处理output的中断,判断是否需要用户的输入 """ - - class Config: - arbitrary_types_allowed = True + model_config = ConfigDict(arbitrary_types_allowed=True) id: str output_node: OutputNode diff --git a/src/backend/bisheng/workflow/nodes/qa_retriever/qa_retriever.py b/src/backend/bisheng/workflow/nodes/qa_retriever/qa_retriever.py index 645cdba0e..0c5625d9b 100644 --- a/src/backend/bisheng/workflow/nodes/qa_retriever/qa_retriever.py +++ b/src/backend/bisheng/workflow/nodes/qa_retriever/qa_retriever.py @@ -50,6 +50,7 @@ class QARetrieverNode(BaseNode): result_str = json.loads(result['result'][0].metadata['extra'])['answer'] else: result_str = '' + self.graph_state.set_variable(self.id, '$retrieved_result$', None) return { 'retrieved_result': result_str diff --git a/src/backend/bisheng/workflow/nodes/rag/rag.py b/src/backend/bisheng/workflow/nodes/rag/rag.py index e98f35791..4fdd9df6f 100644 --- a/src/backend/bisheng/workflow/nodes/rag/rag.py +++ b/src/backend/bisheng/workflow/nodes/rag/rag.py @@ -193,8 +193,6 @@ class RagNode(BaseNode): self._qa_prompt = ChatPromptTemplate.from_messages(messages_general) def init_milvus(self): - if self._milvus: - return if self._knowledge_type == 'knowledge': node_type = 'MilvusWithPermissionCheck' params = { @@ -228,8 +226,6 @@ class RagNode(BaseNode): self._milvus = instantiate_vectorstore(node_type, class_object=class_obj, params=params) def init_es(self): - if self._es: - return if self._knowledge_type == 'knowledge': node_type = 'ElasticsearchWithPermissionCheck' params = { diff --git a/src/backend/bisheng/workflow/nodes/tool/tool.py b/src/backend/bisheng/workflow/nodes/tool/tool.py index dee28b2a0..397743308 100644 --- a/src/backend/bisheng/workflow/nodes/tool/tool.py +++ b/src/backend/bisheng/workflow/nodes/tool/tool.py @@ -1,6 +1,7 @@ from typing import Any from bisheng.api.services.assistant_agent import AssistantAgent +from bisheng.database.constants import ToolPresetType from bisheng.database.models.gpts_tools import GptsToolsDao from bisheng.workflow.nodes.base import BaseNode from bisheng.workflow.nodes.prompt_template import PromptTemplateParser @@ -14,10 +15,12 @@ class ToolNode(BaseNode): self._tool_info = GptsToolsDao.get_tool_by_tool_key(tool_key=self._tool_key) if not self._tool_info: raise Exception(f"工具{self._tool_key}不存在") - if self._tool_info.is_preset: + if self._tool_info.is_preset == ToolPresetType.PRESET.value: self._tool = AssistantAgent.sync_init_preset_tools(tool_list=[self._tool_info], llm=None)[0] - else: + elif self._tool_info.is_preset == ToolPresetType.API.value: self._tool = AssistantAgent.sync_init_personal_tools([self._tool_info])[0] + else: + self._tool = AssistantAgent.sync_init_mcp_tools([self._tool_info])[0] def _run(self, unique_id: str): tool_input = self.parse_tool_input() @@ -47,7 +50,11 @@ class ToolNode(BaseNode): for key, val in self.node_params.items(): if key == "output": continue - ret[key] = self.parse_template_msg(val) + new_val = self.parse_template_msg(val) + if new_val == '' or new_val is None: + continue + ret[key] = new_val + return ret def parse_template_msg(self, msg: str): diff --git a/src/backend/entrypoint.sh b/src/backend/entrypoint.sh index c5a9affca..31357455f 100644 --- a/src/backend/entrypoint.sh +++ b/src/backend/entrypoint.sh @@ -1,4 +1,4 @@ nohup uvicorn bisheng.main:app --host 0.0.0.0 --port 7860 --no-access-log --workers 8 & -# -c 是指定celery的并发数 -celery -A bisheng.worker.main worker -l info -c 16 +# -c 是指定celery的并发数,线程数 +celery -A bisheng.worker.main worker -l info -c 100 -P threads diff --git a/src/backend/pyproject.toml b/src/backend/pyproject.toml index 119eff782..96af0f0cb 100644 --- a/src/backend/pyproject.toml +++ b/src/backend/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "bisheng" -version = "1.1.1" +version = "1.1.0" description = "A Python package with a built-in web application" authors = ["Dataelement "] maintainers = [ @@ -18,22 +18,25 @@ include = ["./bisheng/*", "bisheng/**/*"] bisheng = "bisheng.__main__:main" [tool.poetry.dependencies] -bisheng_langchain = "1.1.1" +python = ">=3.10,<3.11" +bisheng_langchain = "1.2.0.dev1" bisheng_pyautogen = "0.3.2" -langchain = "^0.2.16" +langchain = "^0.3.23" +langchain-community = "^0.3.21" langchain_experimental = "*" -langchain_openai = "^0.1.25" -langchain-community = "^0.2.17" -langgraph = "^0.2.37" -openai = "^1.51.2" -langchain-google-genai = "^1.0.10" -langchain-anthropic = "^0.1.4" +langchain_openai = "^0.3.12" +langchain-ollama = "^0.3.0" +langchain-google-genai = "^2.1.0" +langchain-anthropic = "^0.3.10" +langchain-deepseek = "^0.1.3" +langgraph = "^0.3.27" +openai = "^1.68.2" minio = "7.2.0" loguru = "^0.7.1" fastapi_jwt_auth = "^0.5.0" captcha = "^0.5.0" rsa = "4.9" -pydantic = "^1.10.21" +pydantic = "^2.7.2" numexpr = "^2.8.6" pypinyin = "0.50.0" celery = { extras = ["redis"], version = "^5.3.6" } @@ -41,9 +44,8 @@ redis = "^5.0.0" jieba = "^0.42.1" PyMuPDF = "^1.23.8" shapely = "^2.0.1" -python = ">=3.9,<3.11" -fastapi = "^0.108.0" -uvicorn = "^0.22.0" +fastapi = "^0.115.0" +uvicorn = "^0.23.1" beautifulsoup4 = "^4.12.2" google-search-results = "^2.4.1" google-api-python-client = "^2.79.0" @@ -70,7 +72,7 @@ cohere = "^4.11.0" python-multipart = "^0.0.6" sqlmodel = "^0.0.14" pymysql = "^0.10.1" -pymilvus = "2.4.0" +pymilvus = "2.4.10" elasticsearch = "^8.9.0" orjson = "^3.9.1" multiprocess = "^0.70.14" @@ -89,11 +91,14 @@ matplotlib = "3.8.4" cchardet = "^2.1.7" llama-index = "0.9.48" tenacity = "<8.4.0" -bisheng-ragas = "^1.0.1" +bisheng-ragas = "^1.0.2" qianfan = "^0.4.4" dashscope = "^1.20.3" blobfile = "^3.0.0" -langchain-ollama = "0.1.3" +mcp = "^1.6.0" +numpy = "^1.26.2" +pyjwt = "^1.7.1" +tencentcloud-sdk-python = "^3.0.1373" [tool.poetry.dev-dependencies] black = "^23.1.0" diff --git a/src/bisheng-langchain/bisheng_langchain/agents/chatglm_functions_agent/base.py b/src/bisheng-langchain/bisheng_langchain/agents/chatglm_functions_agent/base.py index 00a1af080..e38f349b2 100644 --- a/src/bisheng-langchain/bisheng_langchain/agents/chatglm_functions_agent/base.py +++ b/src/bisheng-langchain/bisheng_langchain/agents/chatglm_functions_agent/base.py @@ -2,6 +2,8 @@ import json import re from typing import Any, Dict, List, Optional, Sequence, Tuple, Union +from pydantic import model_validator, Field + from bisheng_langchain.chat_models.host_llm import HostChatGLM from langchain.agents.agent import Agent, AgentOutputParser, BaseSingleActionAgent from langchain.agents.structured_chat.output_parser import StructuredChatOutputParserWithRetries @@ -14,7 +16,6 @@ from langchain.schema import AgentAction, AgentFinish, BasePromptTemplate from langchain.schema.language_model import BaseLanguageModel from langchain.schema.messages import ChatMessage from langchain.tools import BaseTool, StructuredTool -from langchain_core.pydantic_v1 import Field, root_validator HUMAN_MESSAGE_TEMPLATE = '{input}\n\n{agent_scratchpad}' @@ -81,13 +82,15 @@ class ChatglmFunctionsAgent(BaseSingleActionAgent): """Get allowed tools.""" return list([t.name for t in self.tools]) - @root_validator + @model_validator(mode='before') + @classmethod def validate_llm(cls, values: dict) -> dict: if not isinstance(values['llm'], HostChatGLM): raise ValueError('Only supported with ChatGLM3 models.') return values - @root_validator + @model_validator(mode='before') + @classmethod def validate_prompt(cls, values: dict) -> dict: prompt: BasePromptTemplate = values['prompt'] if 'agent_scratchpad' not in prompt.input_variables: diff --git a/src/bisheng-langchain/bisheng_langchain/agents/llm_functions_agent/base.py b/src/bisheng-langchain/bisheng_langchain/agents/llm_functions_agent/base.py index e6b9d8a17..d6c332aac 100644 --- a/src/bisheng-langchain/bisheng_langchain/agents/llm_functions_agent/base.py +++ b/src/bisheng-langchain/bisheng_langchain/agents/llm_functions_agent/base.py @@ -3,6 +3,8 @@ import json from json import JSONDecodeError from typing import Any, List, Optional, Sequence, Tuple, Union +from pydantic import model_validator + from bisheng_langchain.chat_models.host_llm import HostQwenChat from bisheng_langchain.chat_models.proxy_llm import ProxyChatLLM from langchain.agents import BaseSingleActionAgent @@ -16,7 +18,6 @@ from langchain.schema.messages import AIMessage, BaseMessage, FunctionMessage, S from langchain.tools import BaseTool from langchain.tools.convert_to_openai import format_tool_to_openai_function from langchain_core.agents import AgentActionMessageLog -from langchain_core.pydantic_v1 import root_validator from langchain_openai import ChatOpenAI @@ -135,7 +136,8 @@ class LLMFunctionsAgent(BaseSingleActionAgent): """Get allowed tools.""" return list([t.name for t in self.tools]) - @root_validator + @model_validator(mode='before') + @classmethod def validate_llm(cls, values: dict) -> dict: if ((not isinstance(values['llm'], ChatOpenAI)) and (not isinstance(values['llm'], HostQwenChat)) @@ -144,7 +146,8 @@ class LLMFunctionsAgent(BaseSingleActionAgent): 'Only supported with ChatOpenAI and HostQwenChat and ProxyChatLLM models.') return values - @root_validator + @model_validator(mode='before') + @classmethod def validate_prompt(cls, values: dict) -> dict: prompt: BasePromptTemplate = values['prompt'] if 'agent_scratchpad' not in prompt.input_variables: diff --git a/src/bisheng-langchain/bisheng_langchain/chains/qa_generation/base.py b/src/bisheng-langchain/bisheng_langchain/chains/qa_generation/base.py index 9b8e268b4..9525f97c0 100644 --- a/src/bisheng-langchain/bisheng_langchain/chains/qa_generation/base.py +++ b/src/bisheng-langchain/bisheng_langchain/chains/qa_generation/base.py @@ -9,7 +9,7 @@ from typing import Any, Dict, List, Optional from langchain_core.callbacks import CallbackManagerForChainRun from langchain_core.language_models import BaseLanguageModel from langchain_core.prompts import BasePromptTemplate, ChatPromptTemplate -from langchain_core.pydantic_v1 import Field +from pydantic import Field from langchain_text_splitters import RecursiveCharacterTextSplitter, TextSplitter from langchain.chains.base import Chain diff --git a/src/bisheng-langchain/bisheng_langchain/chains/transform.py b/src/bisheng-langchain/bisheng_langchain/chains/transform.py index f666f8720..a29416b1a 100644 --- a/src/bisheng-langchain/bisheng_langchain/chains/transform.py +++ b/src/bisheng-langchain/bisheng_langchain/chains/transform.py @@ -6,7 +6,7 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional from langchain.chains.base import Chain from langchain_core.callbacks import AsyncCallbackManagerForChainRun, CallbackManagerForChainRun -from langchain_core.pydantic_v1 import Field +from pydantic import Field logger = logging.getLogger(__name__) diff --git a/src/bisheng-langchain/bisheng_langchain/chat_models/__init__.py b/src/bisheng-langchain/bisheng_langchain/chat_models/__init__.py index 4de8343a3..aab038ee0 100644 --- a/src/bisheng-langchain/bisheng_langchain/chat_models/__init__.py +++ b/src/bisheng-langchain/bisheng_langchain/chat_models/__init__.py @@ -4,11 +4,10 @@ from .proxy_llm import ProxyChatLLM from .qwen import ChatQWen from .wenxin import ChatWenxin from .xunfeiai import ChatXunfeiAI -from .zhipuai import ChatZhipuAI -from .sensetime import SenseChat +from .sensetime import SenseChat __all__ = [ - 'ProxyChatLLM', 'ChatMinimaxAI', 'ChatWenxin', 'ChatZhipuAI', 'ChatXunfeiAI', 'HostChatGLM', + 'ProxyChatLLM', 'ChatMinimaxAI', 'ChatWenxin', 'ChatXunfeiAI', 'HostChatGLM', 'HostBaichuanChat', 'HostLlama2Chat', 'HostQwenChat', 'CustomLLMChat', 'ChatQWen', 'SenseChat', 'HostYuanChat', 'HostYiChat', 'HostQwen1_5Chat' ] diff --git a/src/bisheng-langchain/bisheng_langchain/chat_models/host_llm.py b/src/bisheng-langchain/bisheng_langchain/chat_models/host_llm.py index 03dd315b7..b6c2b19a2 100644 --- a/src/bisheng-langchain/bisheng_langchain/chat_models/host_llm.py +++ b/src/bisheng-langchain/bisheng_langchain/chat_models/host_llm.py @@ -7,6 +7,8 @@ import sys from typing import TYPE_CHECKING, Any, Callable, Dict, List, Mapping, Optional, Tuple, Union import requests +from pydantic import ConfigDict, model_validator, Field + from bisheng_langchain.utils.requests import Requests from langchain.callbacks.manager import AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun from langchain.chat_models.base import BaseChatModel @@ -15,7 +17,6 @@ from langchain.schema.messages import (AIMessage, BaseMessage, ChatMessage, Func HumanMessage, SystemMessage) from langchain.utils import get_from_dict_or_env from langchain_core.language_models.llms import create_base_retry_decorator -from langchain_core.pydantic_v1 import Field, root_validator # from requests.exceptions import HTTPError @@ -138,13 +139,10 @@ class BaseHostChatLLM(BaseChatModel): verbose: Optional[bool] = False decoupled: Optional[bool] = False + model_config = ConfigDict(validate_by_name=True) - class Config: - """Configuration for this pydantic object.""" - - allow_population_by_field_name = True - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['host_base_url'] = get_from_dict_or_env(values, 'host_base_url', 'HostBaseUrl') diff --git a/src/bisheng-langchain/bisheng_langchain/chat_models/minimax.py b/src/bisheng-langchain/bisheng_langchain/chat_models/minimax.py index 78df85ac4..b5238f954 100644 --- a/src/bisheng-langchain/bisheng_langchain/chat_models/minimax.py +++ b/src/bisheng-langchain/bisheng_langchain/chat_models/minimax.py @@ -12,7 +12,7 @@ from langchain.schema import ChatGeneration, ChatResult from langchain.schema.messages import (AIMessage, BaseMessage, ChatMessage, FunctionMessage, HumanMessage, SystemMessage) from langchain.utils import get_from_dict_or_env -from langchain_core.pydantic_v1 import Field, root_validator +from pydantic import ConfigDict, model_validator, Field from tenacity import (before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential) @@ -141,13 +141,10 @@ class ChatMinimaxAI(BaseChatModel): when tiktoken is called, you can specify a model name to use here.""" verbose: Optional[bool] = False + model_config = ConfigDict(validate_by_name=True) - class Config: - """Configuration for this pydantic object.""" - - allow_population_by_field_name = True - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['minimaxai_api_key'] = get_from_dict_or_env(values, 'minimaxai_api_key', diff --git a/src/bisheng-langchain/bisheng_langchain/chat_models/proxy_llm.py b/src/bisheng-langchain/bisheng_langchain/chat_models/proxy_llm.py index bfaab40cf..54888cd54 100644 --- a/src/bisheng-langchain/bisheng_langchain/chat_models/proxy_llm.py +++ b/src/bisheng-langchain/bisheng_langchain/chat_models/proxy_llm.py @@ -6,6 +6,8 @@ import logging import sys from typing import TYPE_CHECKING, Any, Callable, Dict, List, Mapping, Optional, Tuple, Union +from pydantic import ConfigDict, model_validator, Field + from bisheng_langchain.utils import requests from langchain.callbacks.manager import AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun from langchain.chat_models.base import BaseChatModel @@ -13,7 +15,6 @@ from langchain.schema import ChatGeneration, ChatResult from langchain.schema.messages import (AIMessage, BaseMessage, ChatMessage, FunctionMessage, HumanMessage, SystemMessage) from langchain.utils import get_from_dict_or_env -from langchain_core.pydantic_v1 import Field, root_validator from tenacity import (before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential) @@ -137,13 +138,10 @@ class ProxyChatLLM(BaseChatModel): when using one of the many model providers that expose an OpenAI-like API but with different models. In those cases, in order to avoid erroring when tiktoken is called, you can specify a model name to use here.""" + model_config = ConfigDict(validate_by_name=True) - class Config: - """Configuration for this pydantic object.""" - - allow_population_by_field_name = True - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['elemai_api_key'] = get_from_dict_or_env(values, 'elemai_api_key', 'ELEMAI_API_KEY') diff --git a/src/bisheng-langchain/bisheng_langchain/chat_models/qwen.py b/src/bisheng-langchain/bisheng_langchain/chat_models/qwen.py index 9e596e81d..e9f208cf5 100644 --- a/src/bisheng-langchain/bisheng_langchain/chat_models/qwen.py +++ b/src/bisheng-langchain/bisheng_langchain/chat_models/qwen.py @@ -7,6 +7,8 @@ import logging import sys from typing import TYPE_CHECKING, Any, Callable, Dict, List, Mapping, Optional, Tuple, Union +from pydantic import ConfigDict, model_validator, Field + from bisheng_langchain.utils.requests import Requests # import requests from langchain.callbacks.manager import AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun @@ -15,7 +17,6 @@ from langchain.schema import ChatGeneration, ChatResult from langchain.schema.messages import (AIMessage, BaseMessage, ChatMessage, FunctionMessage, HumanMessage, SystemMessage, ToolMessage) from langchain.utils import get_from_dict_or_env -from langchain_core.pydantic_v1 import Field, root_validator from tenacity import (before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential) @@ -165,13 +166,10 @@ class ChatQWen(BaseChatModel): API but with different models. In those cases, in order to avoid erroring when tiktoken is called, you can specify a model name to use here.""" verbose: Optional[bool] = False + model_config = ConfigDict(validate_by_name=True) - class Config: - """Configuration for this pydantic object.""" - - allow_population_by_field_name = True - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['api_key'] = get_from_dict_or_env(values, 'api_key', 'QWEN_API_KEY') diff --git a/src/bisheng-langchain/bisheng_langchain/chat_models/sensetime.py b/src/bisheng-langchain/bisheng_langchain/chat_models/sensetime.py index 095c00b86..ea582e85b 100644 --- a/src/bisheng-langchain/bisheng_langchain/chat_models/sensetime.py +++ b/src/bisheng-langchain/bisheng_langchain/chat_models/sensetime.py @@ -7,6 +7,8 @@ import time from typing import Any, Dict, List, Mapping, Optional, Tuple, Union import jwt +from pydantic import ConfigDict, model_validator, Field + from bisheng_langchain.utils.requests import Requests from langchain.callbacks.manager import AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun from langchain.chat_models.base import BaseChatModel @@ -14,7 +16,6 @@ from langchain.schema import ChatGeneration, ChatResult from langchain.schema.messages import (AIMessage, BaseMessage, ChatMessage, FunctionMessage, HumanMessage, SystemMessage) from langchain.utils import get_from_dict_or_env -from langchain_core.pydantic_v1 import Field, root_validator from tenacity import (before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential) @@ -157,13 +158,10 @@ class SenseChat(BaseChatModel): API but with different models. In those cases, in order to avoid erroring when tiktoken is called, you can specify a model name to use here.""" verbose: Optional[bool] = False + model_config = ConfigDict(validate_by_name=True) - class Config: - """Configuration for this pydantic object.""" - - allow_population_by_field_name = True - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" diff --git a/src/bisheng-langchain/bisheng_langchain/chat_models/wenxin.py b/src/bisheng-langchain/bisheng_langchain/chat_models/wenxin.py index b8e2f1e7f..3170d2b21 100644 --- a/src/bisheng-langchain/bisheng_langchain/chat_models/wenxin.py +++ b/src/bisheng-langchain/bisheng_langchain/chat_models/wenxin.py @@ -12,7 +12,7 @@ from langchain.schema import ChatGeneration, ChatResult from langchain.schema.messages import (AIMessage, BaseMessage, ChatMessage, FunctionMessage, HumanMessage, SystemMessage) from langchain.utils import get_from_dict_or_env -from langchain_core.pydantic_v1 import Field, root_validator +from pydantic import ConfigDict, model_validator, Field from tenacity import (before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential) @@ -138,13 +138,10 @@ class ChatWenxin(BaseChatModel): API but with different models. In those cases, in order to avoid erroring when tiktoken is called, you can specify a model name to use here.""" verbose: Optional[bool] = False + model_config = ConfigDict(validate_by_name=True) - class Config: - """Configuration for this pydantic object.""" - - allow_population_by_field_name = True - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['wenxin_api_key'] = get_from_dict_or_env(values, 'wenxin_api_key', 'WENXIN_API_KEY') diff --git a/src/bisheng-langchain/bisheng_langchain/chat_models/xunfeiai.py b/src/bisheng-langchain/bisheng_langchain/chat_models/xunfeiai.py index 99f6cd4ef..a3ec1eaea 100644 --- a/src/bisheng-langchain/bisheng_langchain/chat_models/xunfeiai.py +++ b/src/bisheng-langchain/bisheng_langchain/chat_models/xunfeiai.py @@ -12,7 +12,7 @@ from langchain.schema import ChatGeneration, ChatResult from langchain.schema.messages import (AIMessage, BaseMessage, ChatMessage, FunctionMessage, HumanMessage, SystemMessage) from langchain.utils import get_from_dict_or_env -from langchain_core.pydantic_v1 import Field, root_validator +from pydantic import ConfigDict, model_validator, Field from tenacity import (before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential) @@ -141,13 +141,10 @@ class ChatXunfeiAI(BaseChatModel): API but with different models. In those cases, in order to avoid erroring when tiktoken is called, you can specify a model name to use here.""" verbose: Optional[bool] = False + model_config = ConfigDict(validate_by_name=True) - class Config: - """Configuration for this pydantic object.""" - - allow_population_by_field_name = True - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['xunfeiai_appid'] = get_from_dict_or_env( diff --git a/src/bisheng-langchain/bisheng_langchain/chat_models/zhipuai.py b/src/bisheng-langchain/bisheng_langchain/chat_models/zhipuai.py index b5b6254db..dda1a8b10 100644 --- a/src/bisheng-langchain/bisheng_langchain/chat_models/zhipuai.py +++ b/src/bisheng-langchain/bisheng_langchain/chat_models/zhipuai.py @@ -12,7 +12,7 @@ from langchain.schema import ChatGeneration, ChatResult from langchain.schema.messages import (AIMessage, BaseMessage, ChatMessage, FunctionMessage, HumanMessage, SystemMessage) from langchain.utils import get_from_dict_or_env -from langchain_core.pydantic_v1 import Field, root_validator +from pydantic import ConfigDict, model_validator, Field from tenacity import (before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential) @@ -149,13 +149,10 @@ class ChatZhipuAI(BaseChatModel): API but with different models. In those cases, in order to avoid erroring when tiktoken is called, you can specify a model name to use here.""" verbose: Optional[bool] = False + model_config = ConfigDict(validate_by_name=True) - class Config: - """Configuration for this pydantic object.""" - - allow_population_by_field_name = True - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['zhipuai_api_key'] = get_from_dict_or_env(values, 'zhipuai_api_key', diff --git a/src/bisheng-langchain/bisheng_langchain/embeddings/host_embedding.py b/src/bisheng-langchain/bisheng_langchain/embeddings/host_embedding.py index 57da49b1d..11beff202 100644 --- a/src/bisheng-langchain/bisheng_langchain/embeddings/host_embedding.py +++ b/src/bisheng-langchain/bisheng_langchain/embeddings/host_embedding.py @@ -6,7 +6,7 @@ from typing import Any, Callable, Dict, List, Optional, Tuple, Union import requests from langchain.embeddings.base import Embeddings from langchain.utils import get_from_dict_or_env -from langchain_core.pydantic_v1 import BaseModel, Field, root_validator +from pydantic import model_validator, BaseModel, Field from tenacity import (before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential) @@ -42,7 +42,7 @@ class HostEmbeddings(BaseModel, Embeddings): """host embedding models. """ - client: Optional[Any] #: :meta private: + client: Optional[Any] = None #: :meta private: """Model name to use.""" model: str = 'embedding-host' host_base_url: str = None @@ -64,7 +64,8 @@ class HostEmbeddings(BaseModel, Embeddings): url_ep: Optional[str] = None - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['host_base_url'] = get_from_dict_or_env(values, 'host_base_url', 'HostBaseUrl') @@ -164,7 +165,8 @@ class CustomHostEmbedding(HostEmbeddings): model: str = Field('custom-embedding', alias='model') embedding_ctx_length: int = 512 - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['host_base_url'] = get_from_dict_or_env(values, 'host_base_url', 'HostBaseUrl') diff --git a/src/bisheng-langchain/bisheng_langchain/embeddings/huggingfacegte.py b/src/bisheng-langchain/bisheng_langchain/embeddings/huggingfacegte.py index 780a0468a..4cef1f76c 100644 --- a/src/bisheng-langchain/bisheng_langchain/embeddings/huggingfacegte.py +++ b/src/bisheng-langchain/bisheng_langchain/embeddings/huggingfacegte.py @@ -2,7 +2,7 @@ from typing import Any, Dict, List, Optional import requests from langchain_core.embeddings import Embeddings -from langchain_core.pydantic_v1 import BaseModel, Extra, Field +from pydantic import BaseModel, Extra, Field DEFAULT_Multilingual_MODEL = "thenlper/gte-large-zh" @@ -27,7 +27,7 @@ class HuggingFaceGteEmbeddings(BaseModel, Embeddings): ) """ - client: Any #: :meta private: + client: Any = None #: :meta private: model_name: str = DEFAULT_Multilingual_MODEL """Model name to use.""" cache_folder: Optional[str] = None diff --git a/src/bisheng-langchain/bisheng_langchain/embeddings/huggingfacemultilingual.py b/src/bisheng-langchain/bisheng_langchain/embeddings/huggingfacemultilingual.py index af704e6c9..d9d66e39d 100644 --- a/src/bisheng-langchain/bisheng_langchain/embeddings/huggingfacemultilingual.py +++ b/src/bisheng-langchain/bisheng_langchain/embeddings/huggingfacemultilingual.py @@ -2,7 +2,7 @@ from typing import Any, Dict, List, Optional import requests from langchain_core.embeddings import Embeddings -from langchain_core.pydantic_v1 import BaseModel, Extra, Field +from pydantic import BaseModel, Extra, Field DEFAULT_Multilingual_MODEL = "intfloat/multilingual-e5-large" @@ -27,7 +27,7 @@ class HuggingFaceMultilingualEmbeddings(BaseModel, Embeddings): ) """ - client: Any #: :meta private: + client: Any = None #: :meta private: model_name: str = DEFAULT_Multilingual_MODEL """Model name to use.""" cache_folder: Optional[str] = None diff --git a/src/bisheng-langchain/bisheng_langchain/embeddings/wenxin.py b/src/bisheng-langchain/bisheng_langchain/embeddings/wenxin.py index 23c4ef8f8..65f5905cf 100644 --- a/src/bisheng-langchain/bisheng_langchain/embeddings/wenxin.py +++ b/src/bisheng-langchain/bisheng_langchain/embeddings/wenxin.py @@ -7,7 +7,7 @@ from typing import Any, Callable, Dict, List, Optional, Tuple, Union # import numpy as np from langchain.embeddings.base import Embeddings from langchain.utils import get_from_dict_or_env -from langchain_core.pydantic_v1 import BaseModel, Extra, Field, root_validator +from pydantic import ConfigDict, model_validator, BaseModel, Field from requests.exceptions import HTTPError from tenacity import (before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential) @@ -55,7 +55,7 @@ class WenxinEmbeddings(BaseModel, Embeddings): """ - client: Optional[Any] #: :meta private: + client: Optional[Any] = None #: :meta private: model: str = 'embedding-v1' deployment: Optional[str] = 'default' @@ -72,13 +72,10 @@ class WenxinEmbeddings(BaseModel, Embeddings): model_kwargs: Optional[Dict[str, Any]] = Field(default_factory=dict) """Holds any model parameters valid for `create` call not explicitly specified.""" + model_config = ConfigDict(extra='forbid') - class Config: - """Configuration for this pydantic object.""" - - extra = Extra.forbid - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values['wenxin_api_key'] = get_from_dict_or_env(values, 'wenxin_api_key', 'WENXIN_API_KEY') diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/agent_types/llm_functions_agent.py b/src/bisheng-langchain/bisheng_langchain/gpts/agent_types/llm_functions_agent.py index 12ac7cc07..7a6d582b0 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/agent_types/llm_functions_agent.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/agent_types/llm_functions_agent.py @@ -8,7 +8,7 @@ from langchain_core.language_models.base import LanguageModelLike from langchain_core.messages import FunctionMessage, SystemMessage, ToolMessage from langgraph.graph import END from langgraph.graph.message import MessageGraph -from langgraph.prebuilt import ToolExecutor, ToolInvocation +from langgraph.prebuilt import ToolNode from langgraph.utils.runnable import RunnableCallable @@ -35,7 +35,7 @@ def get_openai_functions_agent_executor(tools: list[BaseTool], llm: LanguageMode llm_with_tools = llm agent = _get_messages | llm_with_tools - tool_executor = ToolExecutor(tools) + tool_nodes = ToolNode(tools=tools) # Define the function that determines whether to continue or not def should_continue(messages): @@ -55,63 +55,11 @@ def get_openai_functions_agent_executor(tools: list[BaseTool], llm: LanguageMode # Define the function to execute tools async def acall_tool(messages): - actions: list[ToolInvocation] = [] - # Based on the continue condition - # we know the last message involves a function call - last_message = messages[-1] - for tool_call in last_message.additional_kwargs['tool_calls']: - function = tool_call['function'] - function_name = function['name'] - try: - _tool_input = json.loads(function['arguments'] or '{}') - except Exception as e: - raise Exception(f"Error parsing arguments for function: {function_name}. arguments: {function['arguments']}. error: {str(e)}") - # We construct an ToolInvocation from the function_call - actions.append(ToolInvocation( - tool=function_name, - tool_input=_tool_input, - )) - # We call the tool_executor and get back a response - responses = await tool_executor.abatch(actions, **kwargs) - # We use the response to create a ToolMessage - tool_messages = [ - LiberalToolMessage( - tool_call_id=tool_call['id'], - content=response, - additional_kwargs={'name': tool_call['function']['name']}, - ) - for tool_call, response in zip(last_message.additional_kwargs['tool_calls'], responses) - ] + tool_messages = await tool_nodes.ainvoke(messages, None, store=None) return tool_messages def call_tool(messages): - actions: list[ToolInvocation] = [] - # Based on the continue condition - # we know the last message involves a function call - last_message = messages[-1] - for tool_call in last_message.additional_kwargs['tool_calls']: - function = tool_call['function'] - function_name = function['name'] - try: - _tool_input = json.loads(function['arguments'] or '{}') - except Exception as e: - raise Exception(f"Error parsing arguments for function: {function_name}. arguments: {function['arguments']}. error: {str(e)}") - # We construct an ToolInvocation from the function_call - actions.append(ToolInvocation( - tool=function_name, - tool_input=_tool_input, - )) - # We call the tool_executor and get back a response - responses = tool_executor.batch(actions, **kwargs) - # We use the response to create a ToolMessage - tool_messages = [ - LiberalToolMessage( - tool_call_id=tool_call['id'], - content=response, - additional_kwargs={'name': tool_call['function']['name']}, - ) - for tool_call, response in zip(last_message.additional_kwargs['tool_calls'], responses) - ] + tool_messages = tool_nodes.invoke(messages, config=None, store=None) return tool_messages workflow = MessageGraph() @@ -185,7 +133,7 @@ def get_qwen_local_functions_agent_executor( else: llm_with_tools = llm agent = _get_messages | llm_with_tools - tool_executor = ToolExecutor(tools) + tool_nodes = ToolNode(tools=tools) # Define the function that determines whether to continue or not def should_continue(messages): @@ -199,27 +147,7 @@ def get_qwen_local_functions_agent_executor( # Define the function to execute tools async def call_tool(messages): - actions: list[ToolInvocation] = [] - # Based on the continue condition - # we know the last message involves a function call - last_message = messages[-1] - # only one function - function = last_message.additional_kwargs['function_call'] - function_name = function['name'] - try: - _tool_input = json.loads(function['arguments'] or '{}') - except Exception as e: - raise Exception( - f"Error parsing arguments for function: {function_name}. arguments: {function['arguments']}. error: {str(e)}") - # We construct an ToolInvocation from the function_call - actions.append(ToolInvocation( - tool=function_name, - tool_input=_tool_input, - )) - # We call the tool_executor and get back a response - responses = await tool_executor.abatch(actions, **kwargs) - # We use the response to create a ToolMessage - tool_messages = [LiberalFunctionMessage(content=responses[0], name=function_name)] + tool_messages = await tool_nodes.ainvoke(messages, config=None, store=None) return tool_messages workflow = MessageGraph() diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/agent_types/llm_react_agent.py b/src/bisheng-langchain/bisheng_langchain/gpts/agent_types/llm_react_agent.py index 0d0b0f905..37efa4cf1 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/agent_types/llm_react_agent.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/agent_types/llm_react_agent.py @@ -9,7 +9,7 @@ from langchain_core.language_models import LanguageModelLike from langchain_core.messages import BaseMessage from langgraph.graph import END, StateGraph from langgraph.graph.state import CompiledStateGraph -from langgraph.prebuilt.tool_executor import ToolExecutor +from langgraph.prebuilt import ToolNode from langgraph.utils.runnable import RunnableCallable @@ -64,10 +64,7 @@ def create_agent_executor(agent_runnable, tools, input_schema=None) -> CompiledS The `CompiledStateGraph` object. """ - if isinstance(tools, ToolExecutor): - tool_executor = tools - else: - tool_executor = ToolExecutor(tools) + tool_executor = ToolNode(tools=tools) state = _get_agent_state(input_schema) diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/base.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/base.py index f62841d2c..0e5ed7d61 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/base.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/base.py @@ -1,7 +1,8 @@ -from typing import Any, Dict, Tuple, Type, Union +from typing import Any, Dict, Tuple, Type, Union, Optional + +from pydantic import ConfigDict, model_validator, BaseModel, Field from bisheng_langchain.utils.requests import Requests, RequestsWrapper -from langchain_core.pydantic_v1 import BaseModel, Extra, Field, root_validator from langchain_core.tools import BaseTool, Tool from loguru import logger @@ -12,7 +13,7 @@ class ApiArg(BaseModel): class MultArgsSchemaTool(Tool): - def _to_args_and_kwargs(self, tool_input: Union[str, Dict]) -> Tuple[Tuple, Dict]: + def _to_args_and_kwargs(self, tool_input: Union[str, Dict], tool_call_id: Optional[str]) -> Tuple[Tuple, Dict]: # For backwards compatibility, if run_input is a string, # pass as a positional argument. if isinstance(tool_input, str): @@ -32,13 +33,10 @@ class APIToolBase(BaseModel): params: Dict[str, Any] = {} input_key: str = 'keyword' args_schema: Type[BaseModel] = ApiArg + model_config = ConfigDict(extra="forbid") - class Config: - """Configuration for this pydantic object.""" - - extra = Extra.forbid - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" timeout = values.get('request_timeout', 30) diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/firecrawl.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/firecrawl.py index 0b474862f..fa5b90b80 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/firecrawl.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/firecrawl.py @@ -3,7 +3,7 @@ from typing import Any, Dict, Type import requests -from langchain_core.pydantic_v1 import BaseModel, Field, root_validator +from pydantic import BaseModel, Field from bisheng_langchain.gpts.tools.api_tools.base import (APIToolBase, MultArgsSchemaTool) diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/flow.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/flow.py index 7d7a3e372..60942dff2 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/flow.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/flow.py @@ -1,5 +1,5 @@ from loguru import logger -from langchain_core.pydantic_v1 import BaseModel, Field +from pydantic import BaseModel, Field from typing import Any from .base import APIToolBase from .base import MultArgsSchemaTool diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/macro_data.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/macro_data.py index 78c7923fc..29f38f8dc 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/macro_data.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/macro_data.py @@ -4,15 +4,15 @@ from typing import Any import pandas as pd import requests -from langchain.pydantic_v1 import BaseModel, Field +from pydantic import BaseModel, Field from langchain_core.tools import BaseTool from .base import MultArgsSchemaTool class QueryArg(BaseModel): - start_date: str = Field(default='', description='开始月份, 使用YYYY-MM-DD 方式表示', example='2023-01-01') - end_date: str = Field(default='', description='结束月份,使用YYYY-MM-DD 方式表示', example='2023-05-01') + start_date: str = Field(default='', description='开始月份, 使用YYYY-MM-DD 方式表示', examples=['2023-01-01']) + end_date: str = Field(default='', description='结束月份,使用YYYY-MM-DD 方式表示', examples=['2023-05-01']) class MacroData(BaseModel): diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/openapi.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/openapi.py index 4465edf9e..dc69c09e0 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/openapi.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/openapi.py @@ -9,9 +9,9 @@ from .base import APIToolBase, Field, MultArgsSchemaTool class OpenApiTools(APIToolBase): - api_key: Optional[str] - api_location: Optional[str] - parameter_name: Optional[str] + api_key: Optional[str] = None + api_location: Optional[str] = None + parameter_name: Optional[str] = None def get_real_path(self, path_params: dict | None): path = self.params['path'] diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/sina.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/sina.py index acb69d71a..3481178b5 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/sina.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/sina.py @@ -7,7 +7,7 @@ import re from datetime import datetime from typing import List, Type -from langchain_core.pydantic_v1 import BaseModel, Field +from pydantic import BaseModel, Field from loguru import logger from .base import APIToolBase diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/tianyancha.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/tianyancha.py index 11610777e..d5cd013ac 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/tianyancha.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/api_tools/tianyancha.py @@ -3,8 +3,9 @@ from __future__ import annotations from typing import Any, Dict, Type +from pydantic import model_validator, BaseModel, Field + from bisheng_langchain.utils.requests import Requests, RequestsWrapper -from langchain_core.pydantic_v1 import BaseModel, Field, root_validator from .base import APIToolBase @@ -19,7 +20,8 @@ class CompanyInfo(APIToolBase): api_key: str = None args_schema: Type[BaseModel] = InputArgs - @root_validator(pre=True) + @model_validator(mode='before') + @classmethod def build_header(cls, values: Dict[str, Any]) -> Dict[str, Any]: """Build headers that were passed in.""" if not values.get('api_key'): @@ -30,7 +32,8 @@ class CompanyInfo(APIToolBase): values['headers'] = headers return values - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" timeout = values.get('request_timeout', 30) diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/bing_search/tool.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/bing_search/tool.py index 9b9cf05f9..72d8c753e 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/bing_search/tool.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/bing_search/tool.py @@ -2,7 +2,7 @@ from typing import Optional, Type -from langchain.pydantic_v1 import BaseModel, Field +from pydantic import BaseModel, Field from langchain_community.utilities.bing_search import BingSearchAPIWrapper from langchain_core.callbacks import CallbackManagerForToolRun from langchain_core.tools import BaseTool @@ -43,7 +43,7 @@ class BingSearchResults(BaseTool): "Input should be a search query. Output is a JSON array of the query results" ) num_results: int = 5 - args_schema = BingSearchInput + args_schema: Type[BaseModel] = BingSearchInput api_wrapper: BingSearchAPIWrapper def _run( diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/calculator/tool.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/calculator/tool.py index c3043b0c4..c751a001d 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/calculator/tool.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/calculator/tool.py @@ -2,7 +2,7 @@ import math from math import * import sympy -from langchain.pydantic_v1 import BaseModel, Field +from pydantic import BaseModel, Field from langchain.tools import tool from sympy import * @@ -10,7 +10,7 @@ from sympy import * class CalculatorInput(BaseModel): expression: str = Field( description="The input to this tool should be a mathematical expression using only Python's built-in mathematical operators.", - example='200*7', + examples=['200*7'], ) diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/code_interpreter/tool.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/code_interpreter/tool.py index 4c0dc27fa..1567e2dc5 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/code_interpreter/tool.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/code_interpreter/tool.py @@ -15,7 +15,7 @@ from uuid import uuid4 import matplotlib from langchain_community.tools import Tool -from langchain_core.pydantic_v1 import BaseModel, Field +from pydantic import BaseModel, Field from loguru import logger CODE_BLOCK_PATTERN = r"```(\w*)\n(.*?)\n```" @@ -239,7 +239,7 @@ class CodeInterpreterToolArguments(BaseModel): python_code: str = Field( ..., - example="print('Hello World')", + examples=["print('Hello World')"], description=( 'The pure python script to be evaluated. ' 'The contents will be in main.py. ' diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/dalle_image_generator/tool.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/dalle_image_generator/tool.py index 74ed8616e..0fcd888bb 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/dalle_image_generator/tool.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/dalle_image_generator/tool.py @@ -2,13 +2,11 @@ import logging import os from typing import Any, Dict, Mapping, Optional, Tuple, Type, Union -from langchain.pydantic_v1 import BaseModel, Field -from langchain_community.utilities.dalle_image_generator import DallEAPIWrapper from langchain_community.utils.openai import is_openai_v1 from langchain_core.callbacks import CallbackManagerForToolRun -from langchain_core.pydantic_v1 import BaseModel, Extra, Field, root_validator from langchain_core.tools import BaseTool from langchain_core.utils import get_from_dict_or_env, get_pydantic_field_names +from pydantic import ConfigDict, model_validator, BaseModel, Field from bisheng_langchain.utils.azure_dalle_image_generator import AzureDallEWrapper @@ -26,7 +24,7 @@ class DallEAPIWrapper(BaseModel): 2. save your OPENAI_API_KEY in an environment variable """ - client: Any #: :meta private: + client: Any = None #: :meta private: async_client: Any = Field(default=None, exclude=True) #: :meta private: model_name: str = Field(default="dall-e-2", alias="model") model_kwargs: Dict[str, Any] = Field(default_factory=dict) @@ -59,13 +57,10 @@ class DallEAPIWrapper(BaseModel): http_async_client: Union[Any, None] = None """Optional httpx.AsyncClient. Only used for async invocations. Must specify http_client as well if you'd like a custom client for sync invocations.""" + model_config = ConfigDict(extra='forbid') - class Config: - """Configuration for this pydantic object.""" - - extra = Extra.forbid - - @root_validator(pre=True) + @model_validator(mode='before') + @classmethod def build_extra(cls, values: Dict[str, Any]) -> Dict[str, Any]: """Build extra kwargs from additional params that were passed in.""" all_required_field_names = get_pydantic_field_names(cls) @@ -91,7 +86,8 @@ class DallEAPIWrapper(BaseModel): values["model_kwargs"] = extra return values - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" values["openai_api_key"] = get_from_dict_or_env(values, "openai_api_key", "OPENAI_API_KEY") diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/get_current_time/tool.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/get_current_time/tool.py index f23215666..2a666454b 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/get_current_time/tool.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/get_current_time/tool.py @@ -1,7 +1,7 @@ from datetime import datetime import pytz -from langchain.pydantic_v1 import BaseModel, Field +from pydantic import BaseModel, Field from langchain.tools import tool diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/dingding.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/dingding.py index 25bd27361..9b15700fd 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/dingding.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/dingding.py @@ -1,8 +1,7 @@ from typing import Any, Optional, Type import requests -from langchain_core.pydantic_v1 import BaseModel, Field, root_validator -from loguru import logger +from pydantic import BaseModel, Field from bisheng_langchain.gpts.tools.api_tools.base import (APIToolBase, MultArgsSchemaTool) diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/email.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/email.py index 7569ad8b5..70245cfc8 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/email.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/email.py @@ -1,11 +1,9 @@ -import os import smtplib -from email.mime.application import MIMEApplication from email.mime.multipart import MIMEMultipart from email.mime.text import MIMEText -from typing import Any, Optional +from typing import Any -from langchain_core.pydantic_v1 import BaseModel, Field, root_validator +from pydantic import BaseModel, Field from bisheng_langchain.gpts.tools.api_tools.base import (APIToolBase, MultArgsSchemaTool) @@ -17,7 +15,7 @@ class InputArgs(BaseModel): content: str = Field(description="邮件正文内容") -class EmailMessageTool(APIToolBase): +class EmailMessageTool(BaseModel): email_account: str = Field(description="发件人邮箱") email_password: str = Field(description="邮箱授权码/密码") @@ -27,9 +25,9 @@ class EmailMessageTool(APIToolBase): def send_email( self, - receiver, - subject, - content, + receiver: str = None, + subject: str = None, + content: str = None, ): """ 发送电子邮件函数 diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/feishu.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/feishu.py index 2bcfee83e..f613f2804 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/feishu.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/feishu.py @@ -1,29 +1,28 @@ from typing import Any, Optional, Type import requests -from langchain_core.pydantic_v1 import BaseModel, Field, root_validator -from loguru import logger +from pydantic import BaseModel, Field from bisheng_langchain.gpts.tools.api_tools.base import (APIToolBase, MultArgsSchemaTool) class InputArgs(BaseModel): - message: Optional[str] = Field(description="需要发送的钉钉消息") - receive_id: Optional[str] = Field(description="接收的ID") - receive_id_type: Optional[str] = Field(description="接收的ID类型") - container_id: Optional[str] = Field(description="container_id") - start_time: Optional[str] = Field(description="start_time") - end_time: Optional[str] = Field(description="end_time") + message: Optional[str] = Field(None, description="需要发送的钉钉消息") + receive_id: Optional[str] = Field(None, description="接收的ID") + receive_id_type: Optional[str] = Field(None, description="接收的ID类型") + container_id: Optional[str] = Field(None, description="container_id") + start_time: Optional[str] = Field(None, description="start_time") + end_time: Optional[str] = Field(None, description="end_time") # page_token: Optional[str] = Field(description="page_token") - container_id_type: Optional[str] = Field(description="container_id_type") + container_id_type: Optional[str] = Field(None, description="container_id_type") page_size: Optional[int] = Field(default=20,description="page_size") - page_token: Optional[str] = Field(description="page_token") + page_token: Optional[str] = Field(None, description="page_token") sort_type: Optional[str] = Field(description="sort_type",default="ByCreateTimeAsc") class FeishuMessageTool(BaseModel): - API_BASE_URL = "https://open.feishu.cn/open-apis" + API_BASE_URL: str = "https://open.feishu.cn/open-apis" app_id: str = Field(description="app_id") app_secret: str = Field(description="app_secret") diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/wechat.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/wechat.py index 9332ae699..9b3fc35ea 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/wechat.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/message/wechat.py @@ -1,8 +1,7 @@ -from typing import Any, Optional, Type +from typing import Any import requests -from langchain_core.pydantic_v1 import BaseModel, Field, root_validator -from loguru import logger +from pydantic import BaseModel, Field from bisheng_langchain.gpts.tools.api_tools.base import (APIToolBase, MultArgsSchemaTool) diff --git a/src/bisheng-langchain/bisheng_langchain/gpts/tools/sql_agent/tool.py b/src/bisheng-langchain/bisheng_langchain/gpts/tools/sql_agent/tool.py index ccd665018..17ee1a081 100644 --- a/src/bisheng-langchain/bisheng_langchain/gpts/tools/sql_agent/tool.py +++ b/src/bisheng-langchain/bisheng_langchain/gpts/tools/sql_agent/tool.py @@ -1,236 +1,77 @@ -from typing import Type, Optional, TypedDict, Annotated, Any, Literal +from typing import Type, Optional from langchain_community.agent_toolkits import SQLDatabaseToolkit from langchain_community.utilities import SQLDatabase from langchain_core.callbacks import CallbackManagerForToolRun from langchain_core.language_models import BaseLanguageModel -from langchain_core.messages import AnyMessage, AIMessage, ToolMessage -from langchain_core.prompts import ChatPromptTemplate -from langchain_core.runnables import RunnableLambda, RunnableWithFallbacks -from langchain_core.tools import BaseTool, tool -from langgraph.constants import END, START -from langgraph.graph import add_messages, StateGraph -from langgraph.prebuilt import ToolNode -from pydantic import BaseModel, Field +from langchain_core.messages import HumanMessage +from langchain_core.tools import BaseTool +from langgraph.graph.graph import CompiledGraph +from langgraph.prebuilt import create_react_agent +from pydantic import BaseModel, Field, ConfigDict +_agent_system_prompt = """You are an autonomous agent that answers user questions by querying an SQL database through the provided tools. -class State(TypedDict): - messages: Annotated[list[AnyMessage], add_messages] +When a new question arrives, follow the steps *in order*: +1. ALWAYS call `sql_db_list_tables` first. + Purpose: discover what tables are available. Never skip this step. -def handle_tool_error(state) -> dict: - error = state.get("error") - tool_calls = state["messages"][-1].tool_calls - return { - "messages": [ - ToolMessage( - content=f"Error: {repr(error)}\n please fix your mistakes.", - tool_call_id=tc["id"], - ) - for tc in tool_calls - ] - } +2. Choose the table(s) that are probably relevant, then call `sql_db_schema` + once for each of those tables to obtain their schemas. +3. Write one syntactically-correct {dialect} SELECT statement. + Guidelines for this query: + - Return no more than 50 rows **unless** the user explicitly requests another limit. + - Select only the columns needed to answer the question; avoid `SELECT *`. + - If helpful, add `ORDER BY` on a meaningful column so the most interesting rows appear first. + - ABSOLUTELY NO data-modification statements (INSERT, UPDATE, DELETE, DROP, …). + - Double-check the SQL before executing. -def create_tool_node_with_fallback(tools: list) -> RunnableWithFallbacks[Any, dict]: - """ - Create a ToolNode with a fallback to handle errors and surface them to the agent. - """ - return ToolNode(tools).with_fallbacks( - [RunnableLambda(handle_tool_error)], exception_key="error" - ) +4. Execute the query with the execution tool `sql_db_query`. + If execution fails, inspect the error, revise the SQL, and try again. + Repeat until the query runs successfully or you are certain the request + cannot be satisfied. +5. Read the resulting rows and craft a concise, direct answer for the user. + If the result set is empty, explain that no matching data was found. -class SubmitFinalAnswer(BaseModel): - """Submit the final answer to the user based on the query results.""" +6. Include the final SQL query in your answer unless the user asks you not to. - final_answer: str = Field(..., description="The final answer to the user") +Remember: +- List tables → fetch schemas → write & verify SELECT → execute → answer. +- Never skip steps 1 or 2. +- Never perform DML. +- Keep answers focused on the user's question.""" -class QueryDBTool(BaseTool): - name = "db_query_tool" - description = """Execute a SQL query against the database and get back the result. - If the query is not correct, an error message will be returned. - If an error is returned, rewrite the query, check the query, and try again.""" - - db: SQLDatabase - - def _run(self, query: str, run_manager: Optional[CallbackManagerForToolRun] = None): - result = self.db.run_no_throw(query) - if not result: - return "Error: Query failed. Please rewrite your query and try again." - return result class SqlAgentAPIWrapper(BaseModel): + model_config = ConfigDict(arbitrary_types_allowed=True) + llm: BaseLanguageModel = Field(description="llm to use for sql agent") sql_address: str = Field(description="sql database address for SQLDatabase uri") - db: Optional[SQLDatabase] - list_tables_tool: Optional[BaseTool] - get_schema_tool: Optional[BaseTool] - db_query_tool: Optional[BaseTool] - query_check: Optional[Any] - query_gen: Optional[Any] - workflow: Optional[StateGraph] - app: Optional[Any] - schema_llm: Optional[Any] - query_check_llm: Optional[Any] - query_gen_llm: Optional[Any] - - class Config: - arbitrary_types_allowed = True + db: Optional[SQLDatabase] = None + agent: Optional[CompiledGraph] = None def __init__(self, **kwargs): super().__init__(**kwargs) self.llm = kwargs.get('llm') - - # todo 修改sql agent实现逻辑。此处逻辑只支持bishengLLM组件。原因是因为目前的实现必须实例化多个llm对象,每个llm对象绑定不同的tool - self.schema_llm = self.llm.__class__(model_id=self.llm.model_id, model_name=self.llm.model_name) - self.query_check_llm = self.llm.__class__(model_id=self.llm.model_id, model_name=self.llm.model_name) - self.query_gen_llm = self.llm.__class__(model_id=self.llm.model_id, model_name=self.llm.model_name) self.sql_address = kwargs.get('sql_address') self.db = SQLDatabase.from_uri(self.sql_address) toolkit = SQLDatabaseToolkit(db=self.db, llm=self.llm) tools = toolkit.get_tools() - self.list_tables_tool = next(tool for tool in tools if tool.name == "sql_db_list_tables") - self.get_schema_tool = next(tool for tool in tools if tool.name == "sql_db_schema") - self.db_query_tool = QueryDBTool(db=self.db) - - self.query_check = self.init_query_check() - self.query_gen = self.init_query_gen() - - # Define a new graph - self.workflow = StateGraph(State) - self.init_workflow() - self.app = self.workflow.compile(checkpointer=False, debug=True) - - def init_workflow(self): - self.workflow.add_node("first_tool_call", self.first_tool_call) - self.workflow.add_node( - "list_tables_tool", create_tool_node_with_fallback([self.list_tables_tool]) + self.agent = create_react_agent( + self.llm, + tools, + prompt=_agent_system_prompt.format(dialect=self.db.dialect), + checkpointer=False, ) - self.workflow.add_node("get_schema_tool", create_tool_node_with_fallback([self.get_schema_tool])) - - model_get_schema = self.schema_llm.bind_tools( - [self.get_schema_tool] - ) - self.workflow.add_node( - "model_get_schema", - lambda state: { - "messages": [model_get_schema.invoke(state["messages"])], - }, - ) - - self.workflow.add_node("query_gen", self.query_gen_node) - self.workflow.add_node("correct_query", self.model_check_query) - - self.workflow.add_node("execute_query", create_tool_node_with_fallback([self.db_query_tool])) - - self.workflow.add_edge(START, "first_tool_call") - self.workflow.add_edge("first_tool_call", "list_tables_tool") - self.workflow.add_edge("list_tables_tool", "model_get_schema") - self.workflow.add_edge("model_get_schema", "get_schema_tool") - self.workflow.add_edge("get_schema_tool", "query_gen") - self.workflow.add_conditional_edges( - "query_gen", - self.should_continue, - ) - self.workflow.add_edge("correct_query", "execute_query") - self.workflow.add_edge("execute_query", "query_gen") - - @staticmethod - def should_continue(state: State) -> Literal[END, "correct_query", "query_gen"]: - messages = state["messages"] - last_message = messages[-1] - # If there is a tool call, then we finish - if getattr(last_message, "tool_calls", None): - return END - if last_message.content.startswith("Error:"): - return "query_gen" - else: - return "correct_query" - - def init_query_check(self): - query_check_system = """You are a SQL expert with a strong attention to detail. - Double check the SQLite query for common mistakes, including: - - Using NOT IN with NULL values - - Using UNION when UNION ALL should have been used - - Using BETWEEN for exclusive ranges - - Data type mismatch in predicates - - Properly quoting identifiers - - Using the correct number of arguments for functions - - Casting to the correct data type - - Using the proper columns for joins - - If there are any of the above mistakes, rewrite the query. If there are no mistakes, just reproduce the original query. - - You will call the appropriate tool to execute the query after running this check.""" - - query_check_prompt = ChatPromptTemplate.from_messages( - [("system", query_check_system), ("placeholder", "{messages}")] - ) - query_check = query_check_prompt | self.query_check_llm.bind_tools( - [self.db_query_tool] - ) - return query_check - - def first_tool_call(self, state: State) -> dict[str, list[AIMessage]]: - return { - "messages": [ - AIMessage( - content="", - tool_calls=[ - { - "name": "sql_db_list_tables", - "args": {}, - "id": "tool_abcd123", - } - ], - ) - ] - } - - def model_check_query(self, state: State) -> dict[str, list[AIMessage]]: - """ - Use this tool to double-check if your query is correct before executing it. - """ - return {"messages": [self.query_check.invoke({"messages": [state["messages"][-1]]})]} - - def init_query_gen(self): - # Add a node for a model to generate a query based on the question and schema - query_gen_system = """You are a SQL expert with a strong attention to detail.Given an input question, output a syntactically correct SQL query to run, then look at the results of the query and return the answer.DO NOT call any tool besides SubmitFinalAnswer to submit the final answer.When generating the query:Output the SQL query that answers the input question without a tool call.Unless the user specifies a specific number of examples they wish to obtain, always limit your query to at most 10 results.You can order the results by a relevant column to return the most interesting examples in the database.Never query for all the columns from a specific table, only ask for the relevant columns given the question.If you get an error while executing a query, rewrite the query and try again.If you get an empty result set, you should try to rewrite the query to get a non-empty result set. NEVER make stuff up if you don't have enough information to answer the query... just say you don't have enough information.If you have enough information to answer the input question, simply invoke the appropriate tool to submit the final answer to the user.DO NOT make any DML statements (INSERT, UPDATE, DELETE, DROP etc.) to the database.""" - query_gen_prompt = ChatPromptTemplate.from_messages( - [("system", query_gen_system), ("placeholder", "{messages}")] - ) - query_gen = query_gen_prompt | self.query_gen_llm.bind_tools( - [SubmitFinalAnswer] - ) - return query_gen - - def query_gen_node(self, state: State) -> Any: - message = self.query_gen.invoke(state) - - # Sometimes, the LLM will hallucinate and call the wrong tool. We need to catch this and return an error message. - tool_messages = [] - if message.tool_calls: - for tc in message.tool_calls: - if tc["name"] != "SubmitFinalAnswer": - tool_messages.append( - ToolMessage( - content=f"Error: The wrong tool was called: {tc['name']}. Please fix your mistakes. Remember to only call SubmitFinalAnswer to submit the final answer. Generated queries should be outputted WITHOUT a tool call.", - tool_call_id=tc["id"], - ) - ) - else: - tool_messages = [] - return {"messages": [message] + tool_messages} - def run(self, query: str) -> str: - messages = self.app.invoke({"messages": [("user", query)]}, config={ - 'recursion_limit': 50 - }) - return messages["messages"][-1].tool_calls[0]["args"]["final_answer"] + messages = self.agent.invoke({"messages": [HumanMessage(content=query)]}) + return messages["messages"][-1].content def arun(self, query: str) -> str: return self.run(query) @@ -241,8 +82,8 @@ class SqlAgentInput(BaseModel): class SqlAgentTool(BaseTool): - name = "sql_agent" - description = "回答与 SQL 数据库有关的问题。给定用户问题,将从数据库中获取可用的表以及对应 DDL,生成 SQL 查询语句并进行执行,最终得到执行结果。" + name: str = "sql_agent" + description: str = "回答与 SQL 数据库有关的问题。给定用户问题,将从数据库中获取可用的表以及对应 DDL,生成 SQL 查询语句并进行执行,最终得到执行结果。" args_schema: Type[BaseModel] = SqlAgentInput api_wrapper: SqlAgentAPIWrapper diff --git a/src/bisheng-langchain/bisheng_langchain/input_output/input.py b/src/bisheng-langchain/bisheng_langchain/input_output/input.py index f5e684418..48717da5b 100644 --- a/src/bisheng-langchain/bisheng_langchain/input_output/input.py +++ b/src/bisheng-langchain/bisheng_langchain/input_output/input.py @@ -1,12 +1,12 @@ from typing import List, Optional -from pydantic import BaseModel, Extra +from pydantic import ConfigDict, BaseModel class InputNode(BaseModel): """Input组件,用来控制输入""" - input: Optional[List[str]] + input: Optional[List[str]] = None def text(self): return self.input @@ -15,14 +15,10 @@ class InputNode(BaseModel): class VariableNode(BaseModel): """用来设置变量""" # key - variables: Optional[List[str]] + variables: Optional[List[str]] = None # vaulues variable_value: Optional[List[str]] = [] - - class Config: - """Configuration for this pydantic object.""" - - extra = Extra.forbid + model_config = ConfigDict(extra="forbid") def text(self): if self.variable_value: @@ -36,9 +32,9 @@ class VariableNode(BaseModel): class InputFileNode(BaseModel): - file_path: Optional[str] - file_name: Optional[str] - file_type: Optional[str] # tips for file + file_path: Optional[str] = None + file_name: Optional[str] = None + file_type: Optional[str] = None # tips for file """Output组件,用来控制输出""" def text(self): diff --git a/src/bisheng-langchain/bisheng_langchain/input_output/output.py b/src/bisheng-langchain/bisheng_langchain/input_output/output.py index 4ba3d1521..4352b86ba 100644 --- a/src/bisheng-langchain/bisheng_langchain/input_output/output.py +++ b/src/bisheng-langchain/bisheng_langchain/input_output/output.py @@ -5,7 +5,7 @@ from venv import logger from bisheng_langchain.chains import LoaderOutputChain from langchain.callbacks.manager import AsyncCallbackManagerForChainRun, CallbackManagerForChainRun from langchain.chains.base import Chain -from pydantic import BaseModel, Extra +from pydantic import ConfigDict, BaseModel _TEXT_COLOR_MAPPING = { 'blue': '36;1', @@ -52,11 +52,7 @@ class Report(Chain): input_key: str = 'report_name' #: :meta private: output_key: str = 'text' #: :meta private: - - class Config: - """Configuration for this pydantic object.""" - extra = Extra.forbid - arbitrary_types_allowed = True + model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True) @property def input_keys(self) -> List[str]: diff --git a/src/bisheng-langchain/bisheng_langchain/memory/redis.py b/src/bisheng-langchain/bisheng_langchain/memory/redis.py index 29b958d05..195f82e4e 100644 --- a/src/bisheng-langchain/bisheng_langchain/memory/redis.py +++ b/src/bisheng-langchain/bisheng_langchain/memory/redis.py @@ -5,8 +5,7 @@ import redis from langchain.memory.chat_memory import BaseChatMemory from langchain_core.messages import (AIMessage, BaseMessage, HumanMessage, get_buffer_string, message_to_dict, messages_from_dict) -from langchain_core.pydantic_v1 import root_validator -from pydantic import Field +from pydantic import Field, model_validator class ConversationRedisMemory(BaseChatMemory): @@ -20,7 +19,8 @@ class ConversationRedisMemory(BaseChatMemory): redis_prefix: str = 'redis_buffer_' ttl: Optional[int] = None - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: redis_url = values.get('redis_url') if not redis_url: diff --git a/src/bisheng-langchain/bisheng_langchain/rag/bisheng_rag_chain.py b/src/bisheng-langchain/bisheng_langchain/rag/bisheng_rag_chain.py index 06ac20f46..35b30116d 100644 --- a/src/bisheng-langchain/bisheng_langchain/rag/bisheng_rag_chain.py +++ b/src/bisheng-langchain/bisheng_langchain/rag/bisheng_rag_chain.py @@ -10,7 +10,7 @@ from langchain_core.callbacks import (AsyncCallbackManagerForChainRun, CallbackM from langchain_core.language_models import BaseLanguageModel from langchain_core.prompts import (ChatPromptTemplate, HumanMessagePromptTemplate, SystemMessagePromptTemplate) -from langchain_core.pydantic_v1 import Extra, Field +from pydantic import ConfigDict, Field from .bisheng_rag_tool import BishengRAGTool @@ -52,13 +52,7 @@ class BishengRetrievalQA(Chain): """Return the source documents or not.""" bisheng_rag_tool: BishengRAGTool = Field(default_factory=BishengRAGTool, description='RAG tool') - - class Config: - """Configuration for this pydantic object.""" - - extra = Extra.forbid - arbitrary_types_allowed = True - allow_population_by_field_name = True + model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True, validate_by_name=True) @property def input_keys(self) -> List[str]: diff --git a/src/bisheng-langchain/bisheng_langchain/rag/bisheng_rag_tool.py b/src/bisheng-langchain/bisheng_langchain/rag/bisheng_rag_tool.py index 75a192d8c..8fdf7c15d 100644 --- a/src/bisheng-langchain/bisheng_langchain/rag/bisheng_rag_tool.py +++ b/src/bisheng-langchain/bisheng_langchain/rag/bisheng_rag_tool.py @@ -15,7 +15,7 @@ from langchain.chains.combine_documents import create_stuff_documents_chain from langchain_core.callbacks import CallbackManagerForChainRun from langchain_core.language_models.base import LanguageModelLike from langchain_core.prompts import ChatPromptTemplate -from langchain_core.pydantic_v1 import BaseModel, Field +from pydantic import BaseModel, Field from langchain_core.runnables import RunnableConfig from langchain_core.tools import BaseTool, Tool from langchain_core.vectorstores import VectorStoreRetriever @@ -24,7 +24,7 @@ from loguru import logger class MultArgsSchemaTool(Tool): - def _to_args_and_kwargs(self, tool_input: Union[str, Dict]) -> Tuple[Tuple, Dict]: + def _to_args_and_kwargs(self, tool_input: Union[str, Dict], tool_call_id: Optional[str]) -> Tuple[Tuple, Dict]: # For backwards compatibility, if run_input is a string, # pass as a positional argument. if isinstance(tool_input, str): diff --git a/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/baseline_vector_retriever.py b/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/baseline_vector_retriever.py index ad38081bf..1a6a30e01 100644 --- a/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/baseline_vector_retriever.py +++ b/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/baseline_vector_retriever.py @@ -2,7 +2,7 @@ from typing import Any, List, Optional from langchain.text_splitter import TextSplitter from langchain_core.documents import Document -from langchain_core.pydantic_v1 import Field +from pydantic import Field from langchain_core.retrievers import BaseRetriever from loguru import logger diff --git a/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/keyword_retriever.py b/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/keyword_retriever.py index 3f0a827c7..49c0305d4 100644 --- a/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/keyword_retriever.py +++ b/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/keyword_retriever.py @@ -2,7 +2,7 @@ from typing import Any, List, Optional from langchain.text_splitter import TextSplitter from langchain_core.documents import Document -from langchain_core.pydantic_v1 import Field +from pydantic import Field from langchain_core.retrievers import BaseRetriever from loguru import logger diff --git a/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/mix_retriever.py b/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/mix_retriever.py index 0ed3ec098..adad44117 100644 --- a/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/mix_retriever.py +++ b/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/mix_retriever.py @@ -3,7 +3,7 @@ from typing import Any, List, Optional from bisheng_langchain.vectorstores import ElasticKeywordsSearch from langchain.text_splitter import TextSplitter from langchain_core.documents import Document -from langchain_core.pydantic_v1 import Field +from pydantic import Field from langchain_core.retrievers import BaseRetriever diff --git a/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/smaller_chunks_retriever.py b/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/smaller_chunks_retriever.py index b39168441..41afad65b 100644 --- a/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/smaller_chunks_retriever.py +++ b/src/bisheng-langchain/bisheng_langchain/rag/init_retrievers/smaller_chunks_retriever.py @@ -3,7 +3,7 @@ from typing import Any, List, Optional from langchain.text_splitter import TextSplitter from langchain_core.documents import Document -from langchain_core.pydantic_v1 import Field +from pydantic import Field from langchain_core.retrievers import BaseRetriever @@ -15,7 +15,7 @@ class SmallerChunksVectorRetriever(BaseRetriever): parent_splitter: Optional[TextSplitter] = None """The text splitter to use to create parent documents. If none, then the parent documents will be the raw documents passed in.""" - id_key = 'doc_id' + id_key: str = 'doc_id' def add_documents( self, diff --git a/src/bisheng-langchain/bisheng_langchain/rag/test/test_smaller_chunks.py b/src/bisheng-langchain/bisheng_langchain/rag/test/test_smaller_chunks.py index eb3e4f48e..72766ef45 100644 --- a/src/bisheng-langchain/bisheng_langchain/rag/test/test_smaller_chunks.py +++ b/src/bisheng-langchain/bisheng_langchain/rag/test/test_smaller_chunks.py @@ -6,7 +6,7 @@ from typing import Any, Dict, Iterable, List, Optional import httpx from bisheng_langchain.vectorstores.milvus import Milvus from langchain_core.documents import Document -from langchain_core.pydantic_v1 import Field +from pydantic import Field from langchain_core.retrievers import BaseRetriever from langchain_core.vectorstores import VectorStore from loguru import logger diff --git a/src/bisheng-langchain/bisheng_langchain/retrievers/ensemble.py b/src/bisheng-langchain/bisheng_langchain/retrievers/ensemble.py index 24b0ff444..8f217b2cd 100644 --- a/src/bisheng-langchain/bisheng_langchain/retrievers/ensemble.py +++ b/src/bisheng-langchain/bisheng_langchain/retrievers/ensemble.py @@ -6,13 +6,13 @@ multiple retrievers by using weighted Reciprocal Rank Fusion from typing import Any, Dict, List from langchain_core.documents import Document -from langchain_core.pydantic_v1 import root_validator from langchain_core.retrievers import BaseRetriever from langchain.callbacks.manager import ( AsyncCallbackManagerForRetrieverRun, CallbackManagerForRetrieverRun, ) +from pydantic import model_validator class EnsembleRetriever(BaseRetriever): @@ -33,7 +33,8 @@ class EnsembleRetriever(BaseRetriever): weights: List[float] c: int = 60 - @root_validator(pre=True) + @model_validator(mode='before') + @classmethod def set_weights(cls, values: Dict[str, Any]) -> Dict[str, Any]: if not values.get("weights"): n_retrievers = len(values["retrievers"]) diff --git a/src/bisheng-langchain/bisheng_langchain/utils/azure_dalle_image_generator.py b/src/bisheng-langchain/bisheng_langchain/utils/azure_dalle_image_generator.py index 3d303736b..e347ef81f 100644 --- a/src/bisheng-langchain/bisheng_langchain/utils/azure_dalle_image_generator.py +++ b/src/bisheng-langchain/bisheng_langchain/utils/azure_dalle_image_generator.py @@ -3,7 +3,7 @@ from typing import Callable, Dict, Optional, Union import openai from langchain_community.utilities.dalle_image_generator import DallEAPIWrapper -from langchain_core.pydantic_v1 import Field, SecretStr, root_validator +from pydantic import Field, SecretStr, model_validator from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env @@ -56,7 +56,8 @@ class AzureDallEWrapper(DallEAPIWrapper): chunk_size: int = 2048 """Maximum number of texts to embed in each batch""" - @root_validator() + @model_validator(mode='before') + @classmethod def validate_environment(cls, values: Dict) -> Dict: """Validate that api key and python package exists in environment.""" # Check OPENAI_KEY for backwards compatibility. diff --git a/src/bisheng-langchain/bisheng_langchain/utils/requests.py b/src/bisheng-langchain/bisheng_langchain/utils/requests.py index 4a007a120..f13155af2 100644 --- a/src/bisheng-langchain/bisheng_langchain/utils/requests.py +++ b/src/bisheng-langchain/bisheng_langchain/utils/requests.py @@ -5,7 +5,7 @@ from typing import Any, AsyncGenerator, Dict, Optional, Tuple, Union import aiohttp import requests from loguru import logger -from pydantic import BaseModel, Extra +from pydantic import ConfigDict, BaseModel class Requests(BaseModel): @@ -19,12 +19,7 @@ class Requests(BaseModel): aiosession: Optional[aiohttp.ClientSession] = None auth: Optional[Any] = None request_timeout: Union[float, Tuple[float, float]] = 120 - - class Config: - """Configuration for this pydantic object.""" - - extra = Extra.forbid - arbitrary_types_allowed = True + model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True) def get(self, url: str, **kwargs: Any) -> requests.Response: """GET the URL and return the text.""" @@ -140,12 +135,7 @@ class TextRequestsWrapper(BaseModel): aiosession: Optional[aiohttp.ClientSession] = None auth: Optional[Any] = None request_timeout: Union[float, Tuple[float, float]] = 120 - - class Config: - """Configuration for this pydantic object.""" - - extra = Extra.forbid - arbitrary_types_allowed = True + model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True) @property def requests(self) -> Requests: diff --git a/src/bisheng-langchain/bisheng_langchain/vectorstores/retriever.py b/src/bisheng-langchain/bisheng_langchain/vectorstores/retriever.py index 0bc4437dc..866d1a68e 100644 --- a/src/bisheng-langchain/bisheng_langchain/vectorstores/retriever.py +++ b/src/bisheng-langchain/bisheng_langchain/vectorstores/retriever.py @@ -6,7 +6,7 @@ from venv import logger import requests from langchain.schema.document import Document from langchain.vectorstores.base import VectorStore, VectorStoreRetriever -from langchain_core.pydantic_v1 import Field, root_validator +from pydantic import ConfigDict, model_validator, Field if TYPE_CHECKING: from langchain.callbacks.manager import ( @@ -25,13 +25,10 @@ class VectorStoreFilterRetriever(VectorStoreRetriever): 'mmr', ) access_url: str = None + model_config = ConfigDict(arbitrary_types_allowed=True) - class Config: - """Configuration for this pydantic object.""" - - arbitrary_types_allowed = True - - @root_validator() + @model_validator(mode='before') + @classmethod def validate_search_type(cls, values: Dict) -> Dict: """Validate search type.""" search_type = values['search_type'] diff --git a/src/bisheng-langchain/requirements.txt b/src/bisheng-langchain/requirements.txt index 2ad5b56e6..c4b6636cc 100644 --- a/src/bisheng-langchain/requirements.txt +++ b/src/bisheng-langchain/requirements.txt @@ -1,4 +1,4 @@ -langchain==0.2.* +langchain==0.3.* zhipuai websocket-client elasticsearch @@ -10,8 +10,8 @@ pydantic pymupdf==1.23.8 shapely==2.0.2 filetype==1.2.0 -langgraph==0.2.* -openai==1.51.* -langchain_openai>=0.1.25 +langgraph==0.3.* +openai==1.* +langchain_openai==0.3.* llama-index==0.9.48 bisheng-ragas==1.* \ No newline at end of file diff --git a/src/bisheng-langchain/version.txt b/src/bisheng-langchain/version.txt index 88a730bea..24e56e03c 100644 --- a/src/bisheng-langchain/version.txt +++ b/src/bisheng-langchain/version.txt @@ -1 +1 @@ -v0.3.6.dev1 +v1.2.1 \ No newline at end of file diff --git a/src/frontend/client/src/components/Chat/Messages/Content/Files.tsx b/src/frontend/client/src/components/Chat/Messages/Content/Files.tsx index 5b4d2d656..d2011c756 100644 --- a/src/frontend/client/src/components/Chat/Messages/Content/Files.tsx +++ b/src/frontend/client/src/components/Chat/Messages/Content/Files.tsx @@ -5,7 +5,8 @@ import Image from './Image'; const Files = ({ message }: { message?: TMessage }) => { const imageFiles = useMemo(() => { - return message?.files?.filter((file) => file.type?.startsWith('image/')) || []; + const images = message?.files?.filter((file) => file.type?.startsWith('image/')) || []; + return images.map(file => ({ ...file, filepath: file.filepath?.replace(/^https?:\/\/[^\/]+/, __APP_ENV__.BASE_URL) })) }, [message?.files]); const otherFiles = useMemo(() => { @@ -28,8 +29,8 @@ const Files = ({ message }: { message?: TMessage }) => { height: `${file.height ?? 1920}px`, width: `${file.height ?? 1080}px`, }} - // n={imageFiles.length} - // i={i} + // n={imageFiles.length} + // i={i} /> ))} diff --git a/src/frontend/client/src/components/Chat/Messages/MultiMessage.tsx b/src/frontend/client/src/components/Chat/Messages/MultiMessage.tsx index a1eca7fef..8818fb6b3 100644 --- a/src/frontend/client/src/components/Chat/Messages/MultiMessage.tsx +++ b/src/frontend/client/src/components/Chat/Messages/MultiMessage.tsx @@ -61,7 +61,7 @@ export default function MultiMessage({ setSiblingIdx={setSiblingIdxRev} /> ); - } else if (message.content) { + } else if (message.content && message.content.length) { return ( { + return path.replace(/^\/workspace\/tmp-dir/, '/tmp-dir'); + }, } }, }, diff --git a/src/frontend/platform/package.json b/src/frontend/platform/package.json index 2e377a47d..ebf669194 100644 --- a/src/frontend/platform/package.json +++ b/src/frontend/platform/package.json @@ -1,6 +1,6 @@ { "name": "bisheng", - "version": "1.1.1", + "version": "1.2.0", "private": true, "dependencies": { "@headlessui/react": "^2.0.4", diff --git a/src/frontend/platform/public/assets/api/output.png b/src/frontend/platform/public/assets/api/output.png new file mode 100644 index 000000000..48ec975a4 Binary files /dev/null and b/src/frontend/platform/public/assets/api/output.png differ diff --git a/src/frontend/platform/public/locales/en/bs.json b/src/frontend/platform/public/locales/en/bs.json index 6955c607f..db91ea0ab 100644 --- a/src/frontend/platform/public/locales/en/bs.json +++ b/src/frontend/platform/public/locales/en/bs.json @@ -238,7 +238,7 @@ "searchLabels": "Search labels", "confirmed": "Confirmed", "confirm": "Confirm", - "runNewWorkflow": "Run New Workflow", + "runNewWorkflow": "Run New Chat", "chatEndMessage": "This chat has ended" }, "model": { @@ -904,8 +904,6 @@ "10910": "Current knowledge base version does not support segment modification. Please create a new knowledge base to modify segments.", "10800": "Service provider name already exists, please modify", "10801": "Model cannot be duplicated", - "10802": "Failed to add service provider, all models failed to initialize", - "10803": "Failed to add service provider, some models failed to initialize", "10700": "Tag already exists", "10701": "Tag not found", "10600": "Incorrect username or password", @@ -927,7 +925,7 @@ "all": "All", "confirmButton": "Confirm", "add": "Add", - "back": "Back", + "back": "Return", "create": "Create", "delete": "Delete", "deleteSuccess": "Delete successful", diff --git a/src/frontend/platform/public/locales/en/knowledge.json b/src/frontend/platform/public/locales/en/knowledge.json index ac1a5965f..60431b334 100644 --- a/src/frontend/platform/public/locales/en/knowledge.json +++ b/src/frontend/platform/public/locales/en/knowledge.json @@ -1,7 +1,7 @@ { "fileManagement": "File Management", "chunkManagement": "Chunk Management", - "back": "Exit", + "back": "Return", "previewHint": "Click the button on the left to preview the results", "uploadHint": "Please complete file upload first", "nameRequired": "Knowledge Base Name Cannot Be Empty", @@ -88,7 +88,7 @@ "pleaseEnterQuestion": "Please enter a question", "questionAndAnswerCannotBeEmpty": "Question and answer cannot be empty", "max100CharactersForSimilarQuestion": "Max 100 characters for similar question", - "max1000CharactersForAnswer": "Max 1000 characters for answer", + "max1000CharactersForAnswer": "Max 10000 characters for answer", "updateQa": "Update QA", "createQa": "Create QA", "similarQuestions": "Similar Questions", diff --git a/src/frontend/platform/public/locales/zh/bs.json b/src/frontend/platform/public/locales/zh/bs.json index 7e4f4060a..c2da52388 100644 --- a/src/frontend/platform/public/locales/zh/bs.json +++ b/src/frontend/platform/public/locales/zh/bs.json @@ -234,7 +234,7 @@ "searchLabels": "搜索标签", "confirmed": "已确认", "confirm": "确认", - "runNewWorkflow": "运行新工作流", + "runNewWorkflow": "开启新会话", "chatEndMessage": "本轮会话已结束" }, "model": { @@ -566,9 +566,9 @@ }, "tools": { "addTool": "添加工具", - "createCustomTool": "自定义工具", + "createCustomTool": "API工具", "builtinTools": "内置工具", - "customTools": "自定义工具", + "customTools": "API工具", "search": "搜索", "empty": "空空如也", "manageCustomTools": "在此页面管理您的自定义工具,对自定义工具创建、编辑等等", @@ -901,8 +901,6 @@ "10910": "当前知识库版本不支持修改分段,请创建新知识库后进行分段修改", "10800": "模型服务提供方名称重复,请修改", "10801": "模型不可重复", - "10802": "添加模型服务提供方失败,模型全部初始化失败", - "10803": "添加模型服务提供方失败,部分模型初始化失败", "10700": "标签已存在", "10701": "未找到对应的标签", "10600": "账号或密码错误", @@ -914,7 +912,7 @@ "10606": "用户组和角色不能为空", "10610": "用户组内还有用户,不能删除", "10920": "未配置QA知识库相似问模型", - "10930": "该问题已被标注过", + "10930": "该问题已存在", "10527": "工作流等待用户输入超时", "10528": "{{type}}节点执行超过最大次数", "10531": "{{type}}功能已升级,需删除后重新拖入", diff --git a/src/frontend/platform/public/locales/zh/flow.json b/src/frontend/platform/public/locales/zh/flow.json index a5889cce0..03f974739 100644 --- a/src/frontend/platform/public/locales/zh/flow.json +++ b/src/frontend/platform/public/locales/zh/flow.json @@ -130,7 +130,7 @@ "atLeastOneFormItem": "至少添加一个表单项", "editFormItem": "修改表单项", "addFormItem": "添加表单项", - "userInputLabel": "用户输入框展示内容", + "userInputLabel": "用户输入内容变量", "userInputPlaceholder": "此处为空时,需要用户手动输入意见;预置文本时,可允许用户在预置文本的基础上修改并提交。", "optionsCannotBeEmpty": "选项不可为空", "noInteraction": "无交互", diff --git a/src/frontend/platform/public/locales/zh/knowledge.json b/src/frontend/platform/public/locales/zh/knowledge.json index d4dd016da..b3d08d413 100644 --- a/src/frontend/platform/public/locales/zh/knowledge.json +++ b/src/frontend/platform/public/locales/zh/knowledge.json @@ -1,7 +1,7 @@ { "fileManagement": "文件管理", "chunkManagement": "分段管理", - "back": "退出", + "back": "返回", "previewHint": "左侧点击按钮预览结果", "uploadHint": "请先完成文件上传", "nameRequired": "知识库名称不可为空", @@ -88,11 +88,13 @@ "pleaseEnterQuestion": "请先输入问题", "questionAndAnswerCannotBeEmpty": "问题、答案不能为空", "max100CharactersForSimilarQuestion": "相似问最多100个字", - "max1000CharactersForAnswer": "答案最多1000个字", + "max1000CharactersForAnswer": "答案最多10000个字", "updateQa": "更新 QA", "createQa": "创建 QA", "similarQuestions": "相似问题", "aiGenerate": "AI生成", + "errorMsg": "文件解析成功,其中{{value}}条问题上传失败,请检查文件后再试", + "successMsg": "文件上传成功", "cancel2": "取消", "confirm": "确认" } \ No newline at end of file diff --git a/src/frontend/platform/public/models/data.json b/src/frontend/platform/public/models/data.json new file mode 100644 index 000000000..bcbfc93d0 --- /dev/null +++ b/src/frontend/platform/public/models/data.json @@ -0,0 +1,109 @@ +{ + "openai": [{ + "name": "model 1", + "model_name": "gpt-4o", + "model_type": "llm" + }, + { + "name": "model 2", + "model_name": "text-embedding-ada-002", + "model_type": "embedding" + } + ], + "azure_openai": [{ + "name": "model 1", + "model_name": "gpt-4o", + "model_type": "llm" + }, + { + "name": "model 2", + "model_name": "text-embedding-ada-002", + "model_type": "embedding" + } + ], + "qwen": [{ + "name": "model 1", + "model_name": "qwen2.5-72b-instruct", + "model_type": "llm" + }, + { + "name": "model 2", + "model_name": "qwen-max", + "model_type": "llm" + }, + { + "name": "model 3", + "model_name": "qwq-plus", + "model_type": "llm" + }, + { + "name": "model 4", + "model_name": "text-embedding-v3", + "model_type": "embedding" + } + ], + "deepseek": [{ + "name": "model 1", + "model_name": "deepseek-chat", + "model_type": "llm" + }, + { + "name": "model 2", + "model_name": "deepseek-reasoner", + "model_type": "llm" + } + ], + "qianfan": [{ + "name": "model 1", + "model_name": "ernie-4.0-8k", + "model_type": "llm" + }], + "tencent": [{ + "name": "model 1", + "model_name": "hunyuan-turbos-latest", + "model_type": "llm" + }, + { + "name": "model 2", + "model_name": "hunyuan-t1-latest", + "model_type": "llm" + } + ], + "moonshot": [{ + "name": "model 1", + "model_name": "moonshot-v1-32k", + "model_type": "llm" + }], + "zhipu": [{ + "name": "model 1", + "model_name": "glm-4-plus", + "model_type": "llm" + }], + "minimax": [{ + "name": "model 1", + "model_name": "MiniMax-Text-01", + "model_type": "llm" + }], + "volcengine": [{ + "name": "model 1", + "model_name": "deepseek-v3", + "model_type": "llm" + }, + { + "name": "model 2", + "model_name": "deepseek-r1", + "model_type": "llm" + } + ], + "silicon": [{ + "name": "model 1", + "model_name": "deepseek-ai/DeepSeek-R1", + "model_type": "llm" + }, + { + "name": "model 2", + "model_name": "deepseek-ai/DeepSeek-V3", + "model_type": "llm" + } + ] +} \ No newline at end of file diff --git a/src/frontend/platform/src/CustomNodes/GenericNode/index.tsx b/src/frontend/platform/src/CustomNodes/GenericNode/index.tsx index 99b9e2bd4..e876b6b19 100644 --- a/src/frontend/platform/src/CustomNodes/GenericNode/index.tsx +++ b/src/frontend/platform/src/CustomNodes/GenericNode/index.tsx @@ -76,6 +76,8 @@ export default function GenericNode({ data, positionAbsoluteX, positionAbsoluteY return ( <> + +
-
{data.node.description}
+
{data.node.description}
<> {Object.keys(data.node.template) .filter((t) => t.charAt(0) !== "_") diff --git a/src/frontend/platform/src/components/Pro/security/AssistantSetting.tsx b/src/frontend/platform/src/components/Pro/security/AssistantSetting.tsx index 264a5aea9..b3d12460b 100644 --- a/src/frontend/platform/src/components/Pro/security/AssistantSetting.tsx +++ b/src/frontend/platform/src/components/Pro/security/AssistantSetting.tsx @@ -6,12 +6,21 @@ import { useToast } from "@/components/bs-ui/toast/use-toast"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/bs-ui/tooltip"; import { getSensitiveApi, sensitiveSaveApi } from "@/controllers/API/pro"; import { CircleHelp } from "lucide-react"; -import { useEffect, useState } from "react"; +import { forwardRef, useEffect, useImperativeHandle, useState } from "react"; import { useTranslation } from "react-i18next"; import FormSet from "./FormSet"; import FormView from "./FormView"; -export default function AssistantSetting({ id, type }) { +interface AssistantSettingProps { + id: string | number; + type: string; +} + +export interface AssistantSettingRef { + create: (id: number) => Promise; +} + +const AssistantSetting = forwardRef(({ id, type }, ref) => { const { t } = useTranslation(); const [open, setOpen] = useState(false); @@ -25,7 +34,7 @@ export default function AssistantSetting({ id, type }) { // load useEffect(() => { - if (id !== 3) { + if (id && id !== 3) { getSensitiveApi(id, type).then(res => { const { is_check, auto_reply, words, words_type } = res; setForm({ @@ -38,22 +47,47 @@ export default function AssistantSetting({ id, type }) { } }, [id, type]); - const handleFormChange = async (_form) => { + // 验证表单 + const validateForm = (formData: typeof form) => { + if (!formData.isCheck) return true const errors = []; - if (_form.wordsType.length === 0) errors.push(t('build.errors.selectAtLeastOneWordType')); - if (_form.autoReply === '') errors.push(t('build.errors.autoReplyNotEmpty')); - if (errors.length) { - return toast({ title: t('prompt'), variant: 'error', description: errors.join(', ') }); + if (formData.wordsType.length === 0) { + errors.push(t('build.errors.selectAtLeastOneWordType')); } + if (formData.autoReply === '') { + errors.push(t('build.errors.autoReplyNotEmpty')); + } + if (errors.length) { + toast({ title: t('prompt'), variant: 'error', description: errors.join(', ') }); + return false + } + return true + }; + + // 暴露 save 方法给 ref + useImperativeHandle(ref, () => ({ + create: async (saveId?: number) => { + if (!form.isCheck) return true + if (!validateForm(form)) return false + const res = await sensitiveSaveApi({ ...form, id: saveId, type }); + return true; + } + })); + + const handleFormChange = async (_form) => { + if (!validateForm(_form)) return true setForm(_form); - await sensitiveSaveApi({ ..._form, id, type }); - message({ title: t('prompt'), variant: 'success', description: t('build.saveSuccess') }); + if (id) { + await sensitiveSaveApi({ ..._form, id, type }); + message({ title: t('prompt'), variant: 'success', description: t('build.saveSuccess') }); + return false + } }; const onOff = (bln) => { setForm({ ...form, isCheck: bln }); - sensitiveSaveApi({ ...form, isCheck: bln, id, type }); + id && sensitiveSaveApi({ ...form, isCheck: bln, id, type }); if (bln) setOpen(true); }; @@ -99,4 +133,9 @@ export default function AssistantSetting({ id, type }) { ); -} +}) + + +AssistantSetting.displayName = "AssistantSetting"; + +export default AssistantSetting; \ No newline at end of file diff --git a/src/frontend/platform/src/components/bs-comp/apiComponent/ApiAccessFlow.tsx b/src/frontend/platform/src/components/bs-comp/apiComponent/ApiAccessFlow.tsx index db926d8d2..46799cfb7 100644 --- a/src/frontend/platform/src/components/bs-comp/apiComponent/ApiAccessFlow.tsx +++ b/src/frontend/platform/src/components/bs-comp/apiComponent/ApiAccessFlow.tsx @@ -19,7 +19,7 @@ import { useParams } from 'react-router-dom'; const ApiAccessFlow = () => { const { t } = useTranslation() - const {id} = useParams() + const { id } = useParams() // const { flow, getTweak, tabsState } = useContext(TabsContext); // const curl_code = getCurlCode(flow, getTweak, tabsState); // const pythonCode = getPythonApiCode(flow, getTweak, tabsState); @@ -52,6 +52,8 @@ url = "${location.origin}/api/v2/workflow/invoke" payload = json.dumps({ "workflow_id": "${id}", + "stream": False, # 为空或者不传,都会请求流式返回工作流事件。本示例为了直观展示返回结果,所以改 +为非流式请求,真实场景下为了用户体验建议请求流式。 }) headers = { @@ -448,7 +450,7 @@ print(response.text)# 输出工作流的响应`

等待输入事件-对话框形式

-

当工作流返回 event="input" 且 input_type="dialog_input"时,表示后端希望前端在对话框中接收用户输入。

+

当工作流返回 event="input" 且 input_type="dialog_input"时,表示后端希望前端在对话框中接收用户输入以及上传文件(非必须)。

下一次请求 /invoke 接口必带的关键字段是 node_id,message_id,session_id 以及对话框输入。

事件数据示例

@@ -462,7 +464,7 @@ print(response.text)# 输出工作流的响应` -
+
处理逻辑:

  • 绘制对话框,接收用户输入内容
  • -
  • 携带 node_id、session_id、message_id 再次请求 /workflow/invoke,示例如下:
  • +
  • 携带 node_id、session_id、message_id 再次请求 /workflow/invoke
  • +
  • 如果用户没有在对话框内上传文件,请求示例如下
+
+
    +
  • 如果用户在对话框内上传了文件
  • +
  • 如果有文件类型,调用毕昇文件上传接口获取到文件url,示例如下:
  • +
+
+ + {`import requests +def upload_file(local_path: str): + server = "http://ip:port" + url = server + '/api/v1/knowledge/upload' + headers = {} + files = {'file': open(local_path, 'rb')} + res = requests.post(url, headers=headers, files=files) + file_path = res.json()['data'].get('file_path', '') + return file_path + + financeA = upload_file("caibao.pdf") + financeB = upload_file("caibao2.pdf")`} + + +
+
    +
  • 成功获取用户的输入和上传文件的url后,拼接为如下格式的接口入参
  • +
+
+ + {`payload = json.dumps({ + "workflow_id": "c90bb7f2-b7d1-49bf-9fb6-3ab60ff8e414", + "session_id": "d4347ab8e8cd48c48ac9920dbb5a9b35_async_task_id", # 上次返回的 session_id + "message_id": "385140", + "input": { + "input_2775b": { # 这里对应返回事件里的 node_id + # input_schema.value中元素的 key 以及对应要传入的值 + "user_input": "你好", + # 上传文件后获取到的文件url列表 + "dialog_files_content": ["minio://127.0.0.1:9000/xxxx"] + } + } +}) +`} +

等待输入事件-表单形式

-

当工作流返回 event="input" 且 input_type="form_input"时,表示后端希望前端渲染一个表单,让用户填写内容。

-

下一次请求 /invoke 接口必带的关键字段是 node_id, message_id, session_id 以及用户填写的表单值。

+

当工作流返回 event="input" 且 input_type="form_input"时,后端希望前端渲染一个表单,让用户填写内容。

+

下一次请求 /invoke 接口必带的字段是 node_id, message_id, session_id 以及用户填写的表单值。

事件数据示例

@@ -628,7 +690,7 @@ def upload_file(local_path: str): "message_id": "xxxxx", "input": { "input_xxx": { # 事件里的 node_id - # key是input_schme.value中元素的 key 以及对应要传入的值 + # key是input_schema.value中元素的 key 以及对应要传入的值 "text_input": "用户输入的内容", "file": ["minio://127.0.0.1:9000/xxxx"] # 用户上传文件获取到的文件url, 允许多选就是多个url "category": "选项2" # 将选项内容赋值给变量。当允许多选时,多个选项内容通过逗号分隔。 @@ -769,7 +831,7 @@ def upload_file(local_path: str): "message_id": "消息的唯一ID", "input": { "output_123": { # 事件里的节点ID - # key是input_schme.value中元素的key + # key是input_schema.value中元素的key "output_result": "用户输入的内容" } } @@ -867,7 +929,7 @@ def upload_file(local_path: str): "message_id": "xxxxxx", "input": { "output_xxx": { # 事件里的节点ID - # key是input_schme.value中元素的key + # key是input_schema.value中元素的key "output_result": "e2107f75" # 用户选择选项对应的id } } diff --git a/src/frontend/platform/src/components/bs-comp/chatComponent/ChatInput.tsx b/src/frontend/platform/src/components/bs-comp/chatComponent/ChatInput.tsx index af6a52166..c61340f2c 100644 --- a/src/frontend/platform/src/components/bs-comp/chatComponent/ChatInput.tsx +++ b/src/frontend/platform/src/components/bs-comp/chatComponent/ChatInput.tsx @@ -239,7 +239,7 @@ export default function ChatInput({ clear, form, questions, inputForm, wsUrl, on noAccess: false, liked: 0, create_time: formatDate(new Date(), 'yyyy-MM-ddTHH:mm:ss') - }, data.type === 'end_cover') + }, data.type === 'end_cover' && data.category === 'anwser') if (!msgClosedRef.current) msgClosedRef.current = true } else if (data.type === "close") { diff --git a/src/frontend/platform/src/components/bs-comp/chatComponent/MessageBs.tsx b/src/frontend/platform/src/components/bs-comp/chatComponent/MessageBs.tsx index 18be25343..848a657a6 100644 --- a/src/frontend/platform/src/components/bs-comp/chatComponent/MessageBs.tsx +++ b/src/frontend/platform/src/components/bs-comp/chatComponent/MessageBs.tsx @@ -59,7 +59,7 @@ export const ReasoningLog = ({ loading, msg = '' }) => { } -export default function MessageBs({ mark = false, logo, data, onUnlike = () => { }, onSource, onMarkClick }: { logo: string, data: ChatMessageType, onUnlike?: any, onSource?: any }) { +export default function MessageBs({ debug, mark = false, logo, data, onUnlike = () => { }, onSource, onMarkClick }: { logo: string, data: ChatMessageType, onUnlike?: any, onSource?: any }) { const avatarColor = colorList[ (data.sender?.split('').reduce((num, s) => num + s.charCodeAt(), 0) || 0) % colorList.length ] @@ -161,14 +161,14 @@ export default function MessageBs({ mark = false, logo, data, onUnlike = () => { messageId: data.id, message: data.message || data.thought, })} /> - + >} } diff --git a/src/frontend/platform/src/components/bs-comp/chatComponent/MessageButtons.tsx b/src/frontend/platform/src/components/bs-comp/chatComponent/MessageButtons.tsx index 37ec88248..3cf25feaf 100644 --- a/src/frontend/platform/src/components/bs-comp/chatComponent/MessageButtons.tsx +++ b/src/frontend/platform/src/components/bs-comp/chatComponent/MessageButtons.tsx @@ -38,7 +38,7 @@ export default function MessageButtons({ mark = false, id, onCopy, data, onUnlik copyTrackingApi(id) } - return
+ return
{mark && }
-
: (!Array.isArray(data.message.data) &&
+
: (!Array.isArray(data.message.data) &&
-
+ {!debug &&
{!running && handleResend(false)} />} {!running && handleResend(true)} />} {appConfig.dialogQuickSearch && } -
+
}
) } diff --git a/src/frontend/platform/src/components/bs-comp/chatComponent/index.tsx b/src/frontend/platform/src/components/bs-comp/chatComponent/index.tsx index 69228a3ae..8ddd179c6 100644 --- a/src/frontend/platform/src/components/bs-comp/chatComponent/index.tsx +++ b/src/frontend/platform/src/components/bs-comp/chatComponent/index.tsx @@ -3,6 +3,7 @@ import MessagePanne from "./MessagePanne"; export default function ChatComponent({ stop = false, + debug = false, logo = '', clear = false, questions = [], @@ -17,7 +18,7 @@ export default function ChatComponent({ }) { return
- +
}; diff --git a/src/frontend/platform/src/components/bs-comp/chatComponent/messageStore.ts b/src/frontend/platform/src/components/bs-comp/chatComponent/messageStore.ts index 51b69b0b6..8c834ea3c 100644 --- a/src/frontend/platform/src/components/bs-comp/chatComponent/messageStore.ts +++ b/src/frontend/platform/src/components/bs-comp/chatComponent/messageStore.ts @@ -183,12 +183,16 @@ export const useMessageStore = create((set, get) => ({ // run log类型存在嵌套情况,使用 extra 匹配 currentMessage; 否则取最近 let currentMessageIndex = 0 for (let i = messages.length - 1; i >= 0; i--) { + if (messages[i].isSend) break; if (isRunLog && messages[i].extra === wsdata.extra) { currentMessageIndex = i; break; } else if (!isRunLog && !runLogsTypes.includes(messages[i].category)) { currentMessageIndex = i; break; + } else if (wsdata.type === 'end_cover' && messages[i].category === 'tool') { + currentMessageIndex = i; + break; } } const currentMessage = messages[currentMessageIndex] @@ -203,6 +207,14 @@ export const useMessageStore = create((set, get) => ({ } else { message = currentMessage.message + (wsdata.message || '') } + + // 敏感词特殊处理 + if (wsdata.type === 'end_cover' && currentMessage.category === 'tool') { + messages.forEach((msg) => { + msg.end = true // 闭合所有会话 + }) + cover = false + } const newCurrentMessage = { ...currentMessage, ...wsdata, @@ -233,14 +245,20 @@ export const useMessageStore = create((set, get) => ({ // } // 删除重复消息 const prevMessage = messages[currentMessageIndex - 1]; + + // hack + if (wsdata.type === 'end_cover' && !prevMessage.isSend) { + cover = true + } + // 有思考不覆盖 只覆盖message,保留思考 if (prevMessage?.reasoning_log) { if ((prevMessage && prevMessage.message === newCurrentMessage.message && prevMessage.thought === newCurrentMessage.thought) || cover) { - const removedMsg = messages.pop() - prevMessage.message = removedMsg.message + const removedMsg = messages.pop() + prevMessage.message = removedMsg.message } } else { if ((prevMessage diff --git a/src/frontend/platform/src/components/bs-comp/sheets/ToolsSheet.tsx b/src/frontend/platform/src/components/bs-comp/sheets/ToolsSheet.tsx index 930bedbab..feee7a368 100644 --- a/src/frontend/platform/src/components/bs-comp/sheets/ToolsSheet.tsx +++ b/src/frontend/platform/src/components/bs-comp/sheets/ToolsSheet.tsx @@ -1,10 +1,12 @@ +import { LoadIcon } from "@/components/bs-icons/loading"; import { Accordion } from "@/components/bs-ui/accordion"; import { Button } from "@/components/bs-ui/button"; import { SearchInput } from "@/components/bs-ui/input"; import { Sheet, SheetContent, SheetTitle, SheetTrigger } from "@/components/bs-ui/sheet"; import { getAssistantToolsApi } from "@/controllers/API/assistant"; +import { useMcpRefrensh } from "@/pages/BuildPage/tools"; import ToolItem from "@/pages/BuildPage/tools/ToolItem"; -import { Star, User } from "lucide-react"; +import { CpuIcon, Star, User } from "lucide-react"; import { useEffect, useMemo, useState } from "react"; import { useTranslation } from "react-i18next"; @@ -15,11 +17,15 @@ export default function ToolsSheet({ select, onSelect, children }) { const [keyword, setKeyword] = useState('') const [allData, setAllData] = useState([]) - useEffect(() => { + + const loadMData = () => { getAssistantToolsApi(type).then(res => { setAllData(res) setKeyword('') }) + } + useEffect(() => { + loadMData() }, [type]) const options = useMemo(() => { @@ -30,6 +36,7 @@ export default function ToolsSheet({ select, onSelect, children }) { }); }, [keyword, allData]) + const { loading, refresh } = useMcpRefrensh() return ( !open && setKeyword('')}> @@ -40,12 +47,6 @@ export default function ToolsSheet({ select, onSelect, children }) {
{t('build.addTool')} setKeyword(e.target.value)} /> -
{t('tools.customTools')}
+
setType("mcp")} + > + + MCP工具 +
+
+ {type === 'custom' && } + {type === 'mcp' && } + {type === 'mcp' && } +
{ options.length ? options.map(el => ( diff --git a/src/frontend/platform/src/components/bs-ui/alert.tsx b/src/frontend/platform/src/components/bs-ui/alert.tsx index f610caff0..3bd0fa4a8 100644 --- a/src/frontend/platform/src/components/bs-ui/alert.tsx +++ b/src/frontend/platform/src/components/bs-ui/alert.tsx @@ -37,7 +37,7 @@ const AlertTitle = React.forwardRef< >(({ className, ...props }, ref) => (
)) diff --git a/src/frontend/platform/src/components/bs-ui/input/avator.tsx b/src/frontend/platform/src/components/bs-ui/input/avator.tsx index ecbd556ab..9cf019d71 100644 --- a/src/frontend/platform/src/components/bs-ui/input/avator.tsx +++ b/src/frontend/platform/src/components/bs-ui/input/avator.tsx @@ -37,7 +37,7 @@ export default function Avator({ return
{ - value ? : children + value ? : children }
resetCols(defaultValue, options)) - const selectOptionsRef = useRef(defaultValue) + const selectOptionsRef = useRef([...defaultValue]) const handleHover = (option, isLeaf, colIndex) => { setIsHover(true) // // setValues([]) // 从新选择清空值 @@ -167,20 +167,23 @@ export default function Cascader({ error = false, selectClass = '', close = fals {/* 123 */} -
- { - cols.map((_options, index) => { - return
handleHover(op, isLeaf, index)} - onClick={handleClick} - key={index} - /> - }) - } - + {cols.length + ?
+ { + cols.map((_options, index) => { + return
handleHover(op, isLeaf, index)} + onClick={handleClick} + key={index} + /> + }) + } + + :
空
+ } }; diff --git a/src/frontend/platform/src/components/bs-ui/select/index.tsx b/src/frontend/platform/src/components/bs-ui/select/index.tsx index 6fb18f97f..644594ba2 100644 --- a/src/frontend/platform/src/components/bs-ui/select/index.tsx +++ b/src/frontend/platform/src/components/bs-ui/select/index.tsx @@ -18,7 +18,7 @@ const SelectTrigger = React.forwardRef< span]:line-clamp-1 data-[placeholder]:text-gray-400", + "group flex h-9 w-full items-center justify-between whitespace-nowrap rounded-md border border-input bg-search-input px-3 py-2 text-sm shadow-sm ring-offset-background placeholder:text-muted-foreground focus:outline-none focus:ring-1 focus:ring-ring disabled:cursor-not-allowed disabled:opacity-50 [&>span]:line-clamp-1 data-[placeholder]:text-gray-500", className )} {...props} diff --git a/src/frontend/platform/src/components/bs-ui/toast/index.tsx b/src/frontend/platform/src/components/bs-ui/toast/index.tsx index b84f4a1dc..8021d1ab0 100644 --- a/src/frontend/platform/src/components/bs-ui/toast/index.tsx +++ b/src/frontend/platform/src/components/bs-ui/toast/index.tsx @@ -23,7 +23,7 @@ export function Toaster() {
{title && {title}} {description && ( - {description} + {description} )}
{action} diff --git a/src/frontend/platform/src/components/bs-ui/tooltip/index.tsx b/src/frontend/platform/src/components/bs-ui/tooltip/index.tsx index da992bd65..ff5eadbbf 100644 --- a/src/frontend/platform/src/components/bs-ui/tooltip/index.tsx +++ b/src/frontend/platform/src/components/bs-ui/tooltip/index.tsx @@ -19,7 +19,7 @@ const TooltipContent = React.forwardRef< ref={ref} sideOffset={sideOffset} className={cname( - "z-50 overflow-hidden rounded-md bg-primary/80 px-3 py-1.5 text-xs text-primary-foreground animate-in fade-in-0 zoom-in-95 data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=closed]:zoom-out-95 data-[side=bottom]:slide-in-from-top-2 data-[side=left]:slide-in-from-right-2 data-[side=right]:slide-in-from-left-2 data-[side=top]:slide-in-from-bottom-2", + "z-50 overflow-hidden rounded-md bg-primary/80 px-3 py-1.5 text-xs text-primary-foreground animate-in fade-in-0 zoom-in-95 data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=closed]:zoom-out-95 data-[side=bottom]:slide-in-from-top-2 data-[side=left]:slide-in-from-right-2 data-[side=right]:slide-in-from-left-2 data-[side=top]:slide-in-from-bottom-2,data-[side=top-right]:slide-in-from-bottom-2 translate-x-4", className )} {...props} @@ -37,7 +37,7 @@ export const QuestionTooltip = ({ className = '', content }) => ( -
{content}
+
{content}
diff --git a/src/frontend/platform/src/components/bs-ui/tooltip/tip.tsx b/src/frontend/platform/src/components/bs-ui/tooltip/tip.tsx index b49013811..081eaf799 100644 --- a/src/frontend/platform/src/components/bs-ui/tooltip/tip.tsx +++ b/src/frontend/platform/src/components/bs-ui/tooltip/tip.tsx @@ -9,7 +9,7 @@ export default function Tip({ delayDuration = 200, }: { content: string; - side: "top" | "right" | "bottom" | "left"; + side: "top" | "right" | "bottom" | "left"|"top-right"; asChild?: boolean; children: React.ReactNode; styleClasses?: string; @@ -25,7 +25,7 @@ export default function Tip({ avoidCollisions={false} sticky="always" > -
{content}
+
{content}
); diff --git a/src/frontend/platform/src/components/bs-ui/upload/button.tsx b/src/frontend/platform/src/components/bs-ui/upload/button.tsx new file mode 100644 index 000000000..ba979da1e --- /dev/null +++ b/src/frontend/platform/src/components/bs-ui/upload/button.tsx @@ -0,0 +1,217 @@ +import { useState, useRef, useCallback } from 'react'; +import axios, { AxiosProgressEvent, CancelTokenSource } from 'axios'; +import { Button } from "../button"; +// TODO 待测试组件 + +interface UploadButtonProps { + uploadUrl: string; + allowedFileTypes?: string[]; + maxFileSize?: number; // 单位:MB + multiple?: boolean; + formParams?: Record; + className?: string; + children?: React.ReactNode; + onSuccess?: (response: any) => void; + onError?: (error: Error) => void; + onProgress?: (percentage: number) => void; + onBeforeUpload?: (files: File[]) => boolean; +} + +export default function UploadButton({ + uploadUrl, + allowedFileTypes = [], + maxFileSize = 10, // 默认10MB + multiple = false, + formParams = {}, + className, + children = 'Upload', + onSuccess, + onError, + onProgress, + onBeforeUpload +}: UploadButtonProps) { + const fileInputRef = useRef(null); + const [isUploading, setIsUploading] = useState(false); + const [previews, setPreviews] = useState([]); + const [dragActive, setDragActive] = useState(false); + const cancelToken = useRef(); + + // 生成预览图 + const generatePreviews = useCallback((files: File[]) => { + const imageFiles = files.filter(file => file.type.startsWith('image/')); + const previewPromises = imageFiles.map(file => + new Promise((resolve) => { + const reader = new FileReader(); + reader.onload = (e) => resolve(e.target?.result as string); + reader.readAsDataURL(file); + }) + ); + + Promise.all(previewPromises).then(urls => { + setPreviews(prev => [...prev, ...urls]); + }); + }, []); + + // 处理文件校验 + const validateFiles = (files: File[]) => { + // 文件类型校验 + if (allowedFileTypes.length > 0) { + const isValid = files.every(file => + allowedFileTypes.includes(file.type) + ); + if (!isValid) { + throw new Error(`只支持以下文件类型: ${allowedFileTypes.join(', ')}`); + } + } + + // 文件大小校验 + if (maxFileSize > 0) { + const sizeValid = files.every(file => + file.size <= maxFileSize * 1024 * 1024 + ); + if (!sizeValid) { + throw new Error(`文件大小不能超过 ${maxFileSize}MB`); + } + } + }; + + const handleUpload = async (files: File[]) => { + try { + // 校验文件 + validateFiles(files); + + // 上传前回调 + if (onBeforeUpload && !onBeforeUpload(files)) return; + + setIsUploading(true); + const formData = new FormData(); + + // 添加表单参数 + Object.entries(formParams).forEach(([key, value]) => { + formData.append(key, value); + }); + + // 添加文件 + files.forEach(file => { + formData.append('files', file); + }); + + // 生成预览 + generatePreviews(files); + + // 创建取消令牌 + cancelToken.current = axios.CancelToken.source(); + + const response = await axios.post(uploadUrl, formData, { + onUploadProgress: (progressEvent: AxiosProgressEvent) => { + if (progressEvent.total) { + const percent = Math.round( + (progressEvent.loaded * 100) / progressEvent.total + ); + onProgress?.(percent); + } + }, + cancelToken: cancelToken.current.token + }); + + onSuccess?.(response.data); + } catch (err) { + if (!axios.isCancel(err)) { + onError?.(err as Error); + } + } finally { + setIsUploading(false); + } + }; + + // 拖拽处理 + const handleDrag = (e: React.DragEvent) => { + e.preventDefault(); + e.stopPropagation(); + if (e.type === 'dragover') { + setDragActive(true); + } else if (e.type === 'dragleave') { + setDragActive(false); + } + }; + + const handleDrop = (e: React.DragEvent) => { + e.preventDefault(); + e.stopPropagation(); + setDragActive(false); + + if (e.dataTransfer.files && e.dataTransfer.files.length > 0) { + const files = Array.from(e.dataTransfer.files); + handleUpload(files); + } + }; + + // 文件选择处理 + const handleFileChange = (e: React.ChangeEvent) => { + const files = Array.from(e.target.files || []); + if (files.length > 0) { + handleUpload(files); + } + }; + + return ( +
+ {/* 拖拽区域 */} +
+ 拖拽文件到此区域或点击下方按钮上传 + + {/* 预览区域 */} + {previews.length > 0 && ( +
+ {previews.map((src, index) => ( + {`预览-${index}`} + ))} +
+ )} +
+ + + + +
+ ); +} \ No newline at end of file diff --git a/src/frontend/platform/src/components/bs-ui/upload/simple.tsx b/src/frontend/platform/src/components/bs-ui/upload/simple.tsx index d93cb37fd..6744d07d6 100644 --- a/src/frontend/platform/src/components/bs-ui/upload/simple.tsx +++ b/src/frontend/platform/src/components/bs-ui/upload/simple.tsx @@ -5,7 +5,7 @@ import { useToast } from "../toast/use-toast"; import { cname } from "../utils"; import axios from "@/controllers/request"; -export default function SimpleUpload({ filekey, uploadUrl, accept, className = '', onUpload, onProgress, onError, onSuccess }) { +export default function SimpleUpload({ filekey, uploadUrl, accept, className = '', preCheck, onUpload, onProgress, onError, onSuccess }) { const { t } = useTranslation(); const { toast } = useToast() @@ -24,6 +24,26 @@ export default function SimpleUpload({ filekey, uploadUrl, accept, className = ' }); if (!files.length) return + // 执行预校验(如果提供了preCheck函数) + if (preCheck) { + try { + const checkResult = await preCheck(files[0]); + if (checkResult?.valid === false) { + toast({ + title: t('prompt'), + description: checkResult.message || t('code.preCheckFailed'), + }); + return; + } + } catch (error) { + toast({ + title: t('prompt'), + description: error.message || t('code.preCheckError'), + }); + return; + } + } + const formData = new FormData(); formData.append(filekey, files[0]); diff --git a/src/frontend/platform/src/components/inputComponent/index.tsx b/src/frontend/platform/src/components/inputComponent/index.tsx index 5ba7cdf26..23ee14c67 100644 --- a/src/frontend/platform/src/components/inputComponent/index.tsx +++ b/src/frontend/platform/src/components/inputComponent/index.tsx @@ -33,6 +33,7 @@ export default function InputComponent({ value={myValue} // maxLength={maxLength} className={classNames( + "whitespace-normal", disabled ? " input-disable " : "", password && !pwdVisible && myValue !== "" ? " text-clip password " diff --git a/src/frontend/platform/src/components/inputFileComponent/index.tsx b/src/frontend/platform/src/components/inputFileComponent/index.tsx index 8adc159ee..86ca4a602 100644 --- a/src/frontend/platform/src/components/inputFileComponent/index.tsx +++ b/src/frontend/platform/src/components/inputFileComponent/index.tsx @@ -1,3 +1,4 @@ +import { locationContext } from "@/contexts/locationContext"; import { FileSearch2 } from "lucide-react"; import { useContext, useEffect, useState } from "react"; import { alertContext } from "../../contexts/alertContext"; @@ -8,7 +9,6 @@ import { FileComponentType } from "../../types/components"; import { LoadIcon } from "../bs-icons/loading"; import { Button } from "../bs-ui/button"; import { useToast } from "../bs-ui/toast/use-toast"; -import { locationContext } from "@/contexts/locationContext"; export default function InputFileComponent({ value, @@ -52,13 +52,9 @@ export default function InputFileComponent({ const checkFileSize = (file) => { const maxSize = (appConfig.uploadFileMaxSize || 50) * 1024 * 1024; if (file.size > maxSize) { - toast({ - variant: 'error', - description: `上传文件大小不能超过 ${appConfig.uploadFileMaxSize} MB` - }) - setLoading(false); - return true + return `文件:${file.name} 超过 ${appConfig.uploadFileMaxSize} MB,已移除` } + return '' } const handleButtonClick = () => { @@ -76,7 +72,14 @@ export default function InputFileComponent({ // Get the selected file const file = (e.target as HTMLInputElement).files?.[0]; - if (checkFileSize(file)) return + const errorMsg = checkFileSize(file) + if (errorMsg) { + toast({ + variant: 'error', + description: errorMsg + }) + return setLoading(false); + } // Check if the file type is correct // if (file && checkFileType(file.name)) { // Upload the file @@ -136,15 +139,31 @@ export default function InputFileComponent({ setLoading(true); // Get the selected files - const files = (e.target as HTMLInputElement).files; + const _files = (e.target as HTMLInputElement).files; - if (files && files.length > 0) { - const fileNames = Array.from(files).map(file => file.name); // Extract file names + if (_files && _files.length > 0) { const filePaths = []; // This will hold the file paths after successful upload - for (let i = 0; i < files.length; i++) { - if (checkFileSize(files[i])) return + + const errorMsgs = [] + const files = [] + for (let i = 0; i < _files.length; i++) { + const errorMsg = checkFileSize(_files[i]) + errorMsg ? errorMsgs.push(errorMsg) : files.push(_files[i]) } + if (errorMsgs.length) { + toast({ + variant: 'error', + description: errorMsgs + }) + // 文件都不符合要求 结束上传 + if (errorMsgs.length === _files.length) { + return setLoading(false); + } + } + + const fileNames = Array.from(files).map(file => file.name); // Extract file names + // Perform the upload for each file const uploadPromises = Array.from(files).map(file => { return isSSO @@ -159,19 +178,17 @@ export default function InputFileComponent({ setLoading(false); throw new Error(res); // Exit the upload if error occurs } - const { file_path } = res; - filePaths.push(file_path); // Store file paths + return res.file_path }) : uploadFile(file, flow.id).then((data) => { console.log("File uploaded successfully"); - const { file_path } = data; - filePaths.push(file_path); // Store file paths + return data.file_path }); }); // Wait for all file uploads to finish Promise.all(uploadPromises) - .then(() => { + .then((filePaths) => { // After all files are uploaded successfully, update the state setMyValue(fileNames.join(",")); // Join file names with commas onChange(fileNames); // Pass an array of file names diff --git a/src/frontend/platform/src/contexts/locationContext.tsx b/src/frontend/platform/src/contexts/locationContext.tsx index 8edff4395..b5ca0b993 100644 --- a/src/frontend/platform/src/contexts/locationContext.tsx +++ b/src/frontend/platform/src/contexts/locationContext.tsx @@ -78,7 +78,7 @@ export function LocationProvider({ children }: { children: ReactNode }) { officeUrl: res.office_url, dialogTips: res.dialog_tips, dialogQuickSearch: res.dialog_quick_search, - websocketHost: res.websocket_url || window.location.host, + websocketHost: res.websocket_url, isPro: !!res.pro, chatPrompt: !!res.application_usage_tips, noFace: !res.show_github_and_help, diff --git a/src/frontend/platform/src/contexts/tabsContext.tsx b/src/frontend/platform/src/contexts/tabsContext.tsx index 903546c62..3711f3a95 100644 --- a/src/frontend/platform/src/contexts/tabsContext.tsx +++ b/src/frontend/platform/src/contexts/tabsContext.tsx @@ -49,10 +49,11 @@ export function TabsProvider({ children }: { children: ReactNode }) { async function saveFlow(flow: FlowType) { // save api - const newFlow = await captureAndAlertRequestErrorHoc(updateFlowApi(flow)) + const {data, ...info} = flow + const newFlow = await captureAndAlertRequestErrorHoc(updateFlowApi(info)) if (!newFlow) return null; console.log('action :>> ', 'save'); - setFlow(newFlow) + setFlow((flow) => ({...newFlow, data: flow.data})) setTabsState((prev) => { return { ...prev, diff --git a/src/frontend/platform/src/controllers/API/assistant.ts b/src/frontend/platform/src/controllers/API/assistant.ts index 74e908cbd..51f112fe6 100644 --- a/src/frontend/platform/src/controllers/API/assistant.ts +++ b/src/frontend/platform/src/controllers/API/assistant.ts @@ -80,15 +80,30 @@ export const getChatOnlineApi = async (page, keyword, tag_id) => { // 获取工具集合 -export const getAssistantToolsApi = async (type: 'all' | 'default' | 'custom'): Promise => { +// 之前的is_preset字段改为了枚举 +// 0:自定义api工具 +// 1:预置工具 +// 2:mcp工具 +export const getAssistantToolsApi = async (type: 'all' | 'default' | 'custom' | 'mcp'): Promise => { const queryStr = { all: '', - default: '?is_preset=true', - custom: '?is_preset=false' + default: '?is_preset=1', + custom: '?is_preset=0', + mcp: '?is_preset=2' } return await axios.get(`/api/v1/assistant/tool_list${queryStr[type]}`) }; +// 获取mcp服务集合 +export const getAssistantMcpApi = async (): Promise => { + return getAssistantToolsApi('mcp') +} + +// 刷新mcp服务 +export const refreshAssistantMcpApi = async (): Promise => { + return await axios.post(`/api/v1/assistant/mcp/refresh`) +} + // 修改内置工具配置 export const updateAssistantToolApi = async (tool_id, extra) => { return await axios.post(`/api/v1/assistant/tool/config`, { tool_id, extra }) diff --git a/src/frontend/platform/src/controllers/API/finetune.ts b/src/frontend/platform/src/controllers/API/finetune.ts index 894f6d61c..3fbb1c0ff 100644 --- a/src/frontend/platform/src/controllers/API/finetune.ts +++ b/src/frontend/platform/src/controllers/API/finetune.ts @@ -127,11 +127,25 @@ export const getModelListApi = async (): Promise => { // 添加模型 export const addLLmServer = async (data: any) => { + // 删除前端生成的id (id: string ) + data.models = data.models.map((item) => { + const { id, ...other } = item + return typeof id === 'string' ? { + ...other + } : item + }) return await axios.post(`/api/v1/llm`, data) }; // 修改模型 export const updateLLmServer = async (data: any) => { + // 删除前端生成的id (id: string ) + data.models = data.models.map((item) => { + const { id, ...other } = item + return typeof id === 'string' ? { + ...other + } : item + }) return await axios.put(`/api/v1/llm`, data) } diff --git a/src/frontend/platform/src/controllers/API/flow.ts b/src/frontend/platform/src/controllers/API/flow.ts index f636ab6da..5950504a6 100644 --- a/src/frontend/platform/src/controllers/API/flow.ts +++ b/src/frontend/platform/src/controllers/API/flow.ts @@ -343,3 +343,22 @@ export async function runTestCase(data: { question_list, version_list, node_id, return await axios.post(`/api/v1/flows/compare`, data); } +/** + * 聊天窗上传文件 + */ +export async function uploadChatFile(v, file: File, onProgress): Promise { + const formData = new FormData(); + formData.append("file", file); + return await axios.post(`/api/v1/knowledge/upload`, formData, { + headers: { + "Content-Type": "multipart/form-data" + }, + onUploadProgress: (progressEvent) => { + // Calculate progress percentage + if (progressEvent.total) { + const progress = Math.round((progressEvent.loaded * 100) / progressEvent.total); + onProgress(progress); + } + } + }); +} diff --git a/src/frontend/platform/src/controllers/API/index.ts b/src/frontend/platform/src/controllers/API/index.ts index d8dd4c221..4d0c8531a 100644 --- a/src/frontend/platform/src/controllers/API/index.ts +++ b/src/frontend/platform/src/controllers/API/index.ts @@ -305,6 +305,46 @@ export async function getQaList(id, data: { page, pageSize, keyword }) { }); } + +/** + * 导出QA文件 + */ +export async function getQaFile(id): Promise<{ file_list: string[] }> { + return await axios.get(`/api/v1/knowledge/qa/export/${id}`); +} + +/** + * 导入QA文件 + */ +export async function postImportQaFile(id, params): Promise<{ + result: { + answers: string, + questions: string[], + }[] +}> { + const { url } = params; + return await axios.post(`/api/v1/knowledge/qa/import/${id}`, { + file_list: [url], + }); +} + +/** + * 预览QA文件 + */ +export async function getQaFilePreview(id, params): Promise<{ + result: { + answers: string, + questions: string[], + }[] +}> { + const { url, size, offset } = params; + return await axios.post(`/api/v1/knowledge/qa/preview/${id}`, { + file_url: url, + size, + offset, + }); +} + /** * 修改qa状态 */ diff --git a/src/frontend/platform/src/controllers/API/pro.ts b/src/frontend/platform/src/controllers/API/pro.ts index 9971ee468..53cc7343c 100644 --- a/src/frontend/platform/src/controllers/API/pro.ts +++ b/src/frontend/platform/src/controllers/API/pro.ts @@ -68,7 +68,7 @@ export const saveGroupApi = async (data: any): Promise => { group_name, assistant, skill, - workFlows + work_flows: workFlows }); }; diff --git a/src/frontend/platform/src/controllers/API/tools.ts b/src/frontend/platform/src/controllers/API/tools.ts index fce663345..5e52b1ff4 100644 --- a/src/frontend/platform/src/controllers/API/tools.ts +++ b/src/frontend/platform/src/controllers/API/tools.ts @@ -57,6 +57,25 @@ export const downloadToolSchema = async (data: { download_url: string } | { file return await axios.post(`/api/v1/assistant/tool_schema`, data); }; +/** + * 解析mcp服务器配置接口 + */ +export const getMcpServeByConfig = async (data: { file_content: string }): Promise => { + return await axios.post(`/api/v1/assistant/mcp/tool_schema`, data); +} + +/** + * mcp测试接口 + */ +export const testMcpApi = async (data: { file_content: string }) => { + return await axios({ + method: 'post', + url: '/api/v1/assistant/mcp/tool_test', + data + }) +} + + /** * 工具测试接口 * @returns diff --git a/src/frontend/platform/src/controllers/API/workflow.ts b/src/frontend/platform/src/controllers/API/workflow.ts index f9a3f474f..0b19f5ecb 100644 --- a/src/frontend/platform/src/controllers/API/workflow.ts +++ b/src/frontend/platform/src/controllers/API/workflow.ts @@ -163,7 +163,7 @@ const workflowTemplate = [ "name": "输入", "description": "接收用户在会话页面的输入,支持 2 种形式:对话框输入,表单输入。", "type": "input", - "v": "1", + "v": "2", "tab": { "value": "dialog_input", "options": [ @@ -186,10 +186,40 @@ const workflowTemplate = [ { "key": "user_input", "global": "key", - "label": "用户输入内容", + "label": "输入文本内容", "type": "var", "tab": "dialog_input" }, + { + "key": "dialog_files_content", + "global": "key", + "label": "上传文件内容", + "type": "var", + "tab": "dialog_input" + }, + { + "key": "dialog_files_content_size", + "label": "文件内容长度上限", + "type": "char_number", + "min": 0, + "value": 15000, + "tab": "dialog_input" + }, + { + "key": "dialog_file_accept", + "label": "上传文件类型", + "type": "select_fileaccept", + "value": "all", + "tab": "dialog_input" + }, + { + "key": "dialog_image_files", + "global": "key", + "label": "上传图片文件", + "type": "var", + "tab": "dialog_input", + "help": "提取上传文件中的图片文件,当助手或大模型节点使用多模态大模型时,可传入此图片。" + }, { "key": "form_input", "global": "item:form_input", @@ -207,13 +237,14 @@ const workflowTemplate = [ "name": "输出", "description": "可向用户发送消息,并且支持进行更丰富的交互,例如请求用户批准进行某项敏感操作、允许用户在模型输出内容的基础上直接修改并提交。", "type": "output", - "v": "1", + "v": "2", "group_params": [ { "params": [ { - "key": "output_msg", + "key": "message", "label": "消息内容", + "global": "key", "type": "var_textarea_file", "required": true, "placeholder": "输入需要发送给用户的消息,例如“接下来我将执行 XX 操作,请您确认”,“以下是我的初版草稿,您可以在其基础上进行修改”", @@ -243,7 +274,7 @@ const workflowTemplate = [ "name": "大模型", "description": "调用大模型回答用户问题或者处理任务。", "type": "llm", - "v": "1", + "v": "2", "tab": { "value": "single", "options": [ @@ -284,7 +315,7 @@ const workflowTemplate = [ "type": "bisheng_model", "value": "", "required": true, - "placeholder": "请选择模型" + "placeholder": "请在模型管理中配置 LLM 模型" }, { "key": "temperature", @@ -316,7 +347,14 @@ const workflowTemplate = [ "test": "var", "value": "", "required": true - } + }, + { + "key": "image_prompt", + "label": "视觉", + "type": "image_prompt", + "value": [], + "help": "当使用多模态大模型时,可通过此功能传入图片,结合图像内容进行问答" + }, ] }, { @@ -346,7 +384,7 @@ const workflowTemplate = [ "name": "助手", "description": "AI 自主进行任务规划,选择合适的知识库、数据库或工具进行调用。", "type": "agent", - "v": "1", + "v": "2", "tab": { "value": "single", "options": [ @@ -387,7 +425,7 @@ const workflowTemplate = [ "type": "agent_model", "required": true, "value": "", - "placeholder": "请选择模型" + "placeholder": "请在系统模型设置中配置助手推理模型" }, { "key": "temperature", @@ -437,7 +475,14 @@ const workflowTemplate = [ "value": 50 }, "help": "带入模型上下文的历史消息条数,为 0 时代表不包含上下文信息。" - } + }, + { + "key": "image_prompt", + "label": "视觉", + "type": "image_prompt", + "value": "", + "help": "当使用多模态大模型时,可通过此功能传入图片,结合图像内容进行问答" + }, ] }, { @@ -638,7 +683,7 @@ const workflowTemplate = [ "type": "bisheng_model", "value": "", "required": true, - "placeholder": "请选择模型" + "placeholder": "请在模型管理中配置 LLM 模型" }, { "key": "temperature", @@ -853,7 +898,7 @@ const workflowTemplateEN = [ "name": "Input", "description": "Receive user input on the session page, supports two forms: dialog input, form input.", "type": "input", - "v": "1", + "v": "2", "tab": { "value": "dialog_input", "options": [ @@ -876,10 +921,40 @@ const workflowTemplateEN = [ { "key": "user_input", "global": "key", - "label": "User Input Content", + "label": "Enter text content", "type": "var", "tab": "dialog_input" }, + { + "key": "dialog_files_content", + "global": "key", + "label": "Upload file content", + "type": "var", + "tab": "dialog_input" + }, + { + "key": "dialog_files_content_size", + "label": "Maximum length of file content (words)", + "type": "number", + "min": 0, + "value": 15000, + "tab": "dialog_input" + }, + { + "key": "dialog_file_accept", + "label": "Upload file type", + "type": "select_fileaccept", + "value": "all", + "tab": "dialog_input" + }, + { + "key": "dialog_image_files", + "global": "key", + "label": "Upload image files", + "type": "var", + "tab": "dialog_input", + "help": "Extract the image file from the uploaded file. When the assistant or large model node uses the MultiModal Machine Learning large model, this image can be passed in." + }, { "global": "item:form_input", "key": "form_input", @@ -897,13 +972,14 @@ const workflowTemplateEN = [ "name": "Output", "description": "Send messages to users and support richer interactions, such as requesting user approval for sensitive operations or allowing users to directly modify and submit model-generated content.", "type": "output", - "v": "1", + "v": "2", "group_params": [ { "params": [ { - "key": "output_msg", + "key": "message", "label": "Message Content", + "global": "key", "type": "var_textarea_file", "required": true, "placeholder": "Enter the message to send to the user, e.g., 'I will perform XX operation next, please confirm', or 'Here is my draft, feel free to modify it'.", @@ -933,7 +1009,7 @@ const workflowTemplateEN = [ "name": "LLM", "description": "Invoke a large language model to answer user questions or process tasks.", "type": "llm", - "v": "1", + "v": "2", "tab": { "value": "single", "options": [ @@ -973,7 +1049,7 @@ const workflowTemplateEN = [ "type": "bisheng_model", "value": "", "required": true, - "placeholder": "Select a model" + "placeholder": "Please configure LLM model in model management" }, { "key": "temperature", @@ -1002,7 +1078,14 @@ const workflowTemplateEN = [ "test": "var", "value": "", "required": true - } + }, + { + "key": "image_prompt", + "label": "Visual", + "type": "image_prompt", + "value": [], + "help": "When using MultiModal Machine Learning large models, you can use this function to pass in images and combine them with image content for Q & A" + }, ] }, { @@ -1031,7 +1114,7 @@ const workflowTemplateEN = [ "name": "Agent", "description": "AI autonomously plans tasks and selects appropriate knowledge bases or tools for invocation.", "type": "agent", - "v": "1", + "v": "2", "tab": { "value": "single", "options": [ @@ -1071,7 +1154,7 @@ const workflowTemplateEN = [ "type": "agent_model", "required": true, "value": "", - "placeholder": "Select a model" + "placeholder": "Please configure the assistant inference model in Model Management - System Model Settings" }, { "key": "temperature", @@ -1115,7 +1198,14 @@ const workflowTemplateEN = [ "value": 50 }, "help": "Include historical chat records." - } + }, + { + "key": "image_prompt", + "label": "Visual", + "type": "image_prompt", + "value": "", + "help": "When using MultiModal Machine Learning large models, you can use this function to pass in images and combine them with image content for Q & A" + }, ] }, { @@ -1315,7 +1405,7 @@ const workflowTemplateEN = [ "type": "bisheng_model", "value": "", "required": true, - "placeholder": "Select a model" + "placeholder": "Please configure LLM model in model management" }, { "key": "temperature", diff --git a/src/frontend/platform/src/controllers/request.ts b/src/frontend/platform/src/controllers/request.ts index 8b9b3daa0..fa74287fe 100644 --- a/src/frontend/platform/src/controllers/request.ts +++ b/src/frontend/platform/src/controllers/request.ts @@ -27,13 +27,12 @@ customAxios.interceptors.response.use(function (response) { const i18Msg = i18next.t(`errors.${response.data.status_code}`) const errorMessage = i18Msg === `errors.${response.data.status_code}` ? response.data.status_message : i18Msg - // 特殊状态码 - if ([10802, 10803].includes(response.data.status_code)) { - return { ...response.data.data, code: response.data.status_code, msg: errorMessage }; - } // 无权访问 if (response.data.status_code === 403) { - location.href = __APP_ENV__.BASE_URL + '/403' + // 修改不跳转 + if (response.config.method === 'get') { + location.href = __APP_ENV__.BASE_URL + '/403' + } return Promise.reject(errorMessage); } // 异地登录 diff --git a/src/frontend/platform/src/modals/EditNodeModal/index.tsx b/src/frontend/platform/src/modals/EditNodeModal/index.tsx index 07dd8aff2..2364b56bf 100644 --- a/src/frontend/platform/src/modals/EditNodeModal/index.tsx +++ b/src/frontend/platform/src/modals/EditNodeModal/index.tsx @@ -154,7 +154,10 @@ export default function EditNodeModal({ data }: { data: NodeDataType }) {
- {data.node?.description} +

+ {data.node?.description} +

+
{/* */} List diff --git a/src/frontend/platform/src/pages/BuildPage/CreateApp.tsx b/src/frontend/platform/src/pages/BuildPage/CreateApp.tsx index ba567ca92..d95e19935 100644 --- a/src/frontend/platform/src/pages/BuildPage/CreateApp.tsx +++ b/src/frontend/platform/src/pages/BuildPage/CreateApp.tsx @@ -37,9 +37,10 @@ const CreateApp = forwardRef(({ onSave }, ref) => { const [loading, setLoading] = useState(false); const { t } = useTranslation('flow'); const { appConfig } = useContext(locationContext) + const securityRef = useRef(null); // 应用id (edit) - const appidRef = useRef(''); + const [appId, setAppId] = useState(''); // State for errors const [errors, setErrors] = useState({}); @@ -64,7 +65,7 @@ ${t('build.exampleTwo', { ns: 'bs' })} setErrors({}) setOpen(true); tempDataRef.current = null; - appidRef.current = ''; + setAppId(''); }, // edit edit(type: AppType, flow: any) { @@ -74,7 +75,7 @@ ${t('build.exampleTwo', { ns: 'bs' })} setErrors({}) setOpen(true); tempDataRef.current = null; - appidRef.current = flow.flow_id; + setAppId(flow.id); }, })); @@ -180,9 +181,18 @@ ${t('build.exampleTwo', { ns: 'bs' })} navigate('/assistant/' + res.id) } } else { - // 工作流 - const res = await captureAndAlertRequestErrorHoc(createWorkflowApi(formData.name, formData.desc, formData.url)) - if (res) navigate('/flow/' + res.id) + if (appId) return navigate('/flow/' + appId) // 避免重复创建 + // 创建工作流 + const workflow = await captureAndAlertRequestErrorHoc(createWorkflowApi(formData.name, formData.desc, formData.url)) + if (workflow) { + const navigateToFlow = (id) => navigate(`/flow/${id}`); + // 非Pro版本直接跳转 + if (!appConfig.isPro) return navigateToFlow(workflow.id) + + setAppId(workflow.id) + const securityCreated = await securityRef.current.create(workflow.id) + if (securityCreated) navigateToFlow(workflow.id) + } } } setLoading(false); @@ -244,7 +254,7 @@ ${t('build.exampleTwo', { ns: 'bs' })}
{/* 工作流安全审查 */} {appConfig.isPro && - + } diff --git a/src/frontend/platform/src/pages/BuildPage/apps.tsx b/src/frontend/platform/src/pages/BuildPage/apps.tsx index 248e41158..66621d9ba 100644 --- a/src/frontend/platform/src/pages/BuildPage/apps.tsx +++ b/src/frontend/platform/src/pages/BuildPage/apps.tsx @@ -127,7 +127,11 @@ export default function apps() { }) } + const { toast } = useToast() const handleSetting = (data) => { + if (!data.write) { + return toast({ variant: 'warning', description: '无编辑权限' }) + } if (data.flow_type === 5) { // 上线状态下,助手不能进入编辑 navigate(`/assistant/${data.id}`) @@ -215,7 +219,8 @@ export default function apps() { id={item.id} logo={item.logo} type={TypeNames[item.flow_type]} - edit={item.write} + edit + // edit={item.write} title={item.name} isAdmin={user.role === 'admin'} description={item.description} diff --git a/src/frontend/platform/src/pages/BuildPage/assistant/editAssistant/TestChat.tsx b/src/frontend/platform/src/pages/BuildPage/assistant/editAssistant/TestChat.tsx index a6d4e52bc..d661837f8 100644 --- a/src/frontend/platform/src/pages/BuildPage/assistant/editAssistant/TestChat.tsx +++ b/src/frontend/platform/src/pages/BuildPage/assistant/editAssistant/TestChat.tsx @@ -47,6 +47,7 @@ export default function TestChat({ assisId, guideQuestion, onClear }) { {t('build.debugPreview')}
void; - onUpload: (url: string) => void; + onUpload: (url: string, relative_path?: string) => void; }) => { const { toast } = useToast(); @@ -34,7 +34,7 @@ export const IconUploadSection = ({ // 压缩后的文件(result 是 Blob 类型) const compressedFile = new File([result], file.name, { type: result.type }); uploadFileWithProgress(compressedFile, (progress) => { }, 'icon', '').then(res => { - onUpload(res.file_path) + onUpload(res.file_path, res.relative_path) }); }, error(err) { diff --git a/src/frontend/platform/src/pages/BuildPage/bench/index.tsx b/src/frontend/platform/src/pages/BuildPage/bench/index.tsx index 657d85b77..b77224c15 100644 --- a/src/frontend/platform/src/pages/BuildPage/bench/index.tsx +++ b/src/frontend/platform/src/pages/BuildPage/bench/index.tsx @@ -33,10 +33,12 @@ export interface ChatConfigForm { sidebarIcon: { enabled: boolean; image: string; + relative_path: string; }; assistantIcon: { enabled: boolean; image: string; + relative_path: string; }; sidebarSlogan: string; welcomeMessage: string; @@ -83,11 +85,11 @@ export default function index() { navigate('/build/apps') } }, [user]) - - const uploadAvator = (fileUrl: string, type: 'sidebar' | 'assistant') => { + + const uploadAvator = (fileUrl: string, type: 'sidebar' | 'assistant', relativePath?: string) => { setFormData(prev => ({ ...prev, - [`${type}Icon`]: { ...prev[`${type}Icon`], image: fileUrl } + [`${type}Icon`]: { ...prev[`${type}Icon`], image: fileUrl, relative_path: relativePath } })); }; @@ -128,14 +130,14 @@ export default function index() { enabled={formData.sidebarIcon.enabled} image={formData.sidebarIcon.image} onToggle={(enabled) => toggleFeature('sidebarIcon', enabled)} - onUpload={(fileUrl) => uploadAvator(fileUrl, 'sidebar')} + onUpload={(fileUrl, relativePath) => uploadAvator(fileUrl, 'sidebar', relativePath)} /> toggleFeature('assistantIcon', enabled)} - onUpload={(fileUrl) => uploadAvator(fileUrl, 'assistant')} + onUpload={(fileUrl, relativePath) => uploadAvator(fileUrl, 'assistant', relativePath)} /> @@ -297,8 +299,8 @@ export default function index() { const useChatConfig = () => { const [formData, setFormData] = useState({ menuShow: true, - sidebarIcon: { enabled: true, image: '' }, - assistantIcon: { enabled: true, image: '' }, + sidebarIcon: { enabled: true, image: '', relative_path: '' }, + assistantIcon: { enabled: true, image: '', relative_path: '' }, sidebarSlogan: '', welcomeMessage: '', functionDescription: '', diff --git a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/Chat.tsx b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/Chat.tsx index cce1df2e3..b00d3c5e3 100644 --- a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/Chat.tsx +++ b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/Chat.tsx @@ -5,6 +5,7 @@ import { LoadingIcon } from "@/components/bs-icons/loading"; export default function Chat({ stop = false, + debug, autoRun, logo = '', clear = false, @@ -20,7 +21,7 @@ export default function Chat({ return
- + setLoading(false)} >
{/* {loading &&
diff --git a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatFiles.tsx b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatFiles.tsx new file mode 100644 index 000000000..0da2b6d80 --- /dev/null +++ b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatFiles.tsx @@ -0,0 +1,203 @@ +import { useToast } from "@/components/bs-ui/toast/use-toast"; +import { generateUUID } from "@/components/bs-ui/utils"; +import Loading from "@/components/ui/loading"; +import { locationContext } from "@/contexts/locationContext"; +import { uploadChatFile } from "@/controllers/API/flow"; +import { getFileExtension } from "@/util/utils"; +import { FileIcon, PaperclipIcon, X } from "lucide-react"; +import { useContext, useMemo, useRef, useState } from "react"; + +// @accepts '.png,.jpg' +export default function ChatFiles({ v, accepts, onChange }) { + const [files, setFiles] = useState([]); + const filesRef = useRef([]); + const remainingUploadsRef = useRef(0); + const { appConfig } = useContext(locationContext); + // const fileAccepts = useMemo(() => appConfig.libAccepts.map((ext) => `.${ext}`), [appConfig.libAccepts]); + const { toast } = useToast(); + + const fileInputRef = useRef(null); + const fileSizeLimit = appConfig.uploadFileMaxSize * 1024 * 1024; // File size limit in bytes + + const handleFileChange = (e) => { + const selectedFiles = Array.from(e.target.files); + const validFiles = []; + const invalidFiles = []; + + fileInputRef.current.value = '' + // Validate files based on file extensions + selectedFiles.forEach((file) => { + if (file.size <= fileSizeLimit) { + validFiles.push({ id: generateUUID(6), file }); + } else { + invalidFiles.push({ id: generateUUID(6), file }); + } + }); + + // Show invalid file toast + if (invalidFiles.length > 0) { + toast({ + variant: 'info', + description: invalidFiles.map(file => `文件:${file.file.name}超过${appConfig.uploadFileMaxSize}M,已移除`), + }); + } + + if (!validFiles.length) return; + + // Trigger onChange with null to indicate uploading state + onChange(null); + + // Add valid files to state with initial progress + const filesWithProgress = validFiles.map(({ file, id }) => { + return { + name: file.name, + size: file.size, + type: file.type, + isUploading: true, + progress: 0, // Set initial progress to 0 + id, // Use the generated id + file // Keep original file object for later use + }; + }); + + setFiles(prevFiles => { + const res = [...prevFiles, ...filesWithProgress]; + filesRef.current = res; + return res; + }); + + // Keep track of the number of remaining uploads + remainingUploadsRef.current = validFiles.length; + + // Create an array of promises to handle multiple file uploads concurrently + const uploadPromises = validFiles.map(({ file, id }) => { + return uploadChatFile(v, file, (progress) => { + // Update progress for each file individually + setFiles((prevFiles) => { + const updatedFiles = prevFiles.map(f => { + if (f.id === id) { + return { ...f, progress }; // Update progress for the specific file + } + return f; + }); + filesRef.current = updatedFiles; + return updatedFiles; + }); + }).then(response => { + const filePath = response.file_path; // Assuming the response contains the file ID + filesRef.current = filesRef.current.map(f => { + if (f.id === id) { + return { ...f, isUploading: false, filePath, progress: 100 }; // Set progress to 100 when uploaded + } + return f; + }); + setFiles(filesRef.current); + + remainingUploadsRef.current -= 1; // Decrease the remaining uploads count + if (remainingUploadsRef.current === 0) { + // Once all files are uploaded, trigger onChange with the file IDs + const uploadedFileIds = filesRef.current.filter(f => f.id).map(f => ({ path: f.filePath, name: f.name })); + onChange(uploadedFileIds); // Pass the file IDs to onChange + } + }).catch(() => { + // Handle upload failure + toast({ + variant: 'error', + description: `文件上传失败: ${file.name}`, + }); + handleFileRemove(file.name); + remainingUploadsRef.current -= 1; // Decrease the remaining uploads count + if (remainingUploadsRef.current === 0) { + // If no files remain, trigger onChange immediately + const uploadedFileIds = filesRef.current.filter(f => f.id).map(f => ({ path: f.filePath, name: f.name })); + onChange(uploadedFileIds); + } + }); + }); + + // Wait for all files to finish uploading + Promise.all(uploadPromises).then(() => { + // Once all files are uploaded, trigger onChange with the file IDs + const uploadedFileIds = filesRef.current.filter(f => f.id).map(f => ({ path: f.filePath, name: f.name })); + onChange(uploadedFileIds); // Pass the file IDs to onChange + }); + }; + + const handleFileRemove = (fileName) => { + const res = filesRef.current.filter(file => file.name !== fileName); + filesRef.current = res + setFiles(res); + + // If we manually remove a file during upload, we decrease the remaining upload counter + remainingUploadsRef.current = Math.max(remainingUploadsRef.current - 1, 0); + + if (remainingUploadsRef.current === 0) { + // If no files remain, trigger onChange immediately + const uploadedFileIds = filesRef.current.filter(f => f.id).map(f => ({ id: f.id, name: f.name })); + onChange(uploadedFileIds); // Trigger onChange with uploaded file IDs + } + }; + + const formatFileSize = (size) => { + let fileSize = typeof size === 'string' ? parseFloat(size) : size; + const units = ['B', 'KB', 'MB', 'GB']; + let index = 0; + + while (fileSize >= 1024 && index < units.length - 1) { + fileSize /= 1024; + index++; + } + + return `${fileSize.toFixed(2)} ${units[index]}`; + }; + + return ( +
+ {/* Displaying files */} + {!!files.length &&
+ {files.map((file, index) => ( +
+ {/* Remove button */} + handleFileRemove(file.name)} + className="hidden group-hover:block absolute -right-1 -top-1 bg-gray-50 border-2 border-gray-300 text-gray-600 rounded-full cursor-pointer" + > + + + + {/* File Icon */} +
+ {file.isUploading ? : } +
+ + {/* File details */} +
+
+ {file.name} +
+ {file.isUploading ? file.progress === 100 + ?
解析中...
+ :
上传中... {file.progress}%
+ :
{getFileExtension(file.name)} {formatFileSize(file.size)}
} +
+
+ ))} +
} + + {/* File Upload Button */} +
fileInputRef.current.click()}> + +
+ + {/* File Input */} + +
+ ); +} diff --git a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatInput.tsx b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatInput.tsx index fd2b24571..5e8faa248 100644 --- a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatInput.tsx +++ b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatInput.tsx @@ -4,15 +4,21 @@ import { Button } from "@/components/bs-ui/button"; import { Textarea } from "@/components/bs-ui/input"; import { useToast } from "@/components/bs-ui/toast/use-toast"; import { locationContext } from "@/contexts/locationContext"; -import { useContext, useEffect, useRef, useState } from "react"; +import { useContext, useEffect, useMemo, useRef, useState } from "react"; import { useTranslation } from "react-i18next"; // import GuideQuestions from "./GuideQuestions"; // import { useMessageStore } from "./messageStore"; +import Tip from "@/components/bs-ui/tooltip/tip"; import { RefreshCw } from "lucide-react"; -import GuideQuestions from "./GuideQuestions"; -import InputForm from "./InputForm"; -import { useMessageStore } from "./messageStore"; import useFlowStore from "../flowStore"; +import ChatFiles from "./ChatFiles"; +import GuideQuestions from "./GuideQuestions"; +import { useMessageStore } from "./messageStore"; + +export const FileTypes = { + IMAGE: ['.PNG', '.JPEG', '.JPG', '.BMP'], + FILE: ['.PDF', '.TXT', '.MD', '.HTML', '.XLS', '.XLSX', '.DOC', '.DOCX', '.PPT', '.PPTX'], +} export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, onLoad }) { const { toast } = useToast() @@ -23,11 +29,25 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on const questionsRef = useRef(null) const inputNodeIdRef = useRef('') // 当前输入框节点id const messageIdRef = useRef('') // 当前输入框节点messageId - const [inputForm, setInputForm] = useState(null) // input表单 + const [accepts, setAccepts] = useState('*') // 接受文件类型 const [showWhenLocked, setShowWhenLocked] = useState(false) // 强制开启表单按钮,不限制于input锁定 - const { messages, hisMessages, chatId, createSendMsg, createWsMsg, streamWsMsg, insetSeparator, destory, insetNodeRun, setShowGuideQuestion } = useMessageStore() + const { + messages, + hisMessages, + chatId, + createSendMsg, + createWsMsg, + overWsMsg, + inputForm, + setInputForm, + streamWsMsg, + insetSeparator, + destory, + insetNodeRun, + setShowGuideQuestion + } = useMessageStore() console.log('ui messages :>> ', messages); const currentChatIdRef = useRef(null) @@ -35,7 +55,7 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on const continueRef = useRef(false) // 停止状态 const [stop, setStop] = useState({ - show: false, + show: true, disable: false }) /** @@ -85,6 +105,7 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on }, []) const handleSendClick = async () => { + if (fileUploading) return // 解除锁定状态下 form 按钮开放的状态 // setShowWhenLocked(false) // 关闭引导词 @@ -92,9 +113,15 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on // 收起表单 // formShow && setFormShow(false) // setFormShow(false) - - const value = inputRef.current.value - if (value.trim() === '') return + const [filePath, fileNames] = getFileIds().reduce((acc, cur) => { + acc[0].push(cur.path) + acc[1].push(cur.name) + return acc + }, [[], []]) + // 文件拼接入消息 + const _value = inputRef.current.value + if (_value.trim() === '' && filePath.length === 0) return + const value = fileNames.length > 0 ? fileNames.join('\n') + '\n' + _value : _value; const event = new Event('input', { bubbles: true, cancelable: true }); inputRef.current.value = '' @@ -104,6 +131,7 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on const wsMsg = onBeforSend('input', { nodeId: inputNodeIdRef.current, msg: value, + files: filePath, category: "question", extra: '', message_id: messageIdRef.current, @@ -123,25 +151,6 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on } } - const handleSendForm = async ([data, msg]) => { - setInputForm(null) - createSendMsg(msg) - await createWebSocket() - sendWsMsg({ - action: 'input', - data: { - [inputNodeIdRef.current]: { - data, - message: msg, - message_id: messageIdRef.current, - category: 'question', - extra: '', - source: 0 - } - } - }) - } - const sendWsMsg = async (msg) => { try { wsRef.current.send(JSON.stringify(msg)) @@ -155,6 +164,7 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on } const wsRef = useRef(null) + const reRunStateRef = useRef(false) const createWebSocket = () => { // 单例 if (wsRef.current) return Promise.resolve('ok'); @@ -172,12 +182,19 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on }; ws.onmessage = (event) => { const data = JSON.parse(event.data); - console.log('result message data :>> ', data); + // 过滤一些不需要的数据 + if ((data.category === 'end_cover' && data.type !== 'end_cover')) { + return + } if (data.type === 'begin') { setStop({ show: true, disable: false }) - } else if (data.type === 'close') { - setStop({ show: false, disable: false }) + } else if (data.type === 'close' && data.category === 'processing') { + if (!reRunStateRef.current) { + // 重试时阻止关闭stop + setStop({ show: false, disable: false }) + } + reRunStateRef.current = false } // const errorMsg = data.category === 'error' ? data.intermediate_steps : '' @@ -209,6 +226,11 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on } setInputLock({ locked: true, reason: '' }) } + event.reason && addNotification({ + type: 'error', + title: '运行异常', + description: event.reason + }) }; ws.onerror = (ev) => { wsRef.current = null @@ -233,14 +255,24 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on } const setRunCache = useFlowStore(state => state.setRunCache) + const addNotification = useFlowStore((state) => state.addNotification); // 接受 ws 消息 const handleWsMessage = (data) => { if (data.category === 'error') { const { code, message } = data.message setInputLock({ locked: true, reason: '' }) + + // 记录 + const errorMsg = code == 500 ? message : t(`errors.${code}`, { type: message }) + addNotification({ + type: 'error', + title: '运行异常', + description: errorMsg + }) + return toast({ variant: 'error', - description: code == 500 ? message : t(`errors.${code}`, { type: message }) + description: errorMsg }); } else if (data.category === 'node_run') { inputNodeIdRef.current = data.message.node_id @@ -257,6 +289,18 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on const { node_id, input_schema } = data.message inputNodeIdRef.current = node_id messageIdRef.current = data.message_id + // 限制文件类型 + if (input_schema.tab === 'dialog_input') { + const schemaItem = input_schema.value?.find(el => el.key === 'dialog_file_accept') + const fileAccept = schemaItem?.value + if (fileAccept === 'image') { + setAccepts(FileTypes.IMAGE.join(',')) + } else if (fileAccept === 'file') { + setAccepts(FileTypes.FILE.join(',')) + } else { + setAccepts(FileTypes.IMAGE.join(',') + ',' + FileTypes.FILE.join(',')) + } + } // 待用户输入 input_schema.tab === 'form_input' ? setInputForm(input_schema) : setInputLock({ locked: false, reason: '' }) return @@ -264,9 +308,14 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on return questionsRef.current.updateQuestions(data.message.guide_question.filter(q => q)) } else if (data.category === 'stream_msg') { streamWsMsg(data) + } else if (data.category === 'end_cover' && data.type === 'end_cover') { + setInputLock({ locked: true, reason: '' }) + sendWsMsg({ "action": "stop" }); + return overWsMsg(data) + // return handleRestartClick() } - if (data.type === 'close') { + if (data.type === 'close' && data.category === 'processing') { insetSeparator(t('chat.chatEndMessage')) setInputLock({ locked: true, reason: '' }) // 重启会话按钮,接收close确认后端处理结束后重启会话 @@ -303,9 +352,12 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on } const handleOutPutEvent = async (e) => { const { nodeId, data, message } = e.detail + const { flow_id, chat_id } = onBeforSend('flowInfo', {}) await createWebSocket() sendWsMsg({ action: 'input', + flow_id, + chat_id, data: { [nodeId]: { data, @@ -319,9 +371,33 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on } }) } + const handleSendForm = async (e) => { + const { data, msg } = e.detail + setInputForm(null) + createSendMsg(msg) + await createWebSocket() + const { flow_id, chat_id } = onBeforSend('flowInfo', {}) + sendWsMsg({ + action: 'input', + flow_id, + chat_id, + data: { + [inputNodeIdRef.current]: { + data, + message: msg, + message_id: messageIdRef.current, + category: 'question', + extra: '', + source: 0 + } + } + }) + } + document.addEventListener('inputFormEvent', handleSendForm) document.addEventListener('outputMsgEvent', handleOutPutEvent) document.addEventListener('userResendMsgEvent', handleCustomEvent) return () => { + document.removeEventListener('inputFormEvent', handleSendForm) document.removeEventListener('outputMsgEvent', handleOutPutEvent) document.removeEventListener('userResendMsgEvent', handleCustomEvent) } @@ -354,12 +430,15 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on const [restarted, setRestarted] = useState(false) const handleRestartClick = () => { sendWsMsg({ "action": "stop" }); + setInputForm(null) setRestarted(true) const chatId = currentChatIdRef.current.startsWith('test') ? '' : currentChatIdRef.current restartCallBackRef.current[chatId] = () => { createWebSocket().then(() => { setRestarted(false) - sendWsMsg(onBeforSend('init_data', {})) + onBeforSend('refresh_flow', {}).then((data) => { + sendWsMsg(data) + }) }) } // wsRef.current?.close() @@ -370,18 +449,24 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on // sendWsMsg(onBeforSend('init_data', {})) // }) // }, 300); + if (stop.show) { + reRunStateRef.current = true + } } + const placholder = useMemo(() => { + // if (inputForm) { + // return ' 点击刷新按钮可开启新对话' + // } + const reason = inputLock.reason || ' ' + return inputLock.locked ? reason : t('chat.inputPlaceholder') + }, [inputForm, inputLock]) + + // 文件上传状态 + const { fileUploading, getFileIds, loadingChange } = useFileLoading(inputLock.locked) + return
- {/* form */} - { - inputForm &&
-
- -
-
- } {/* 引导问题 */} {/* restart */}
- + + +
{/* form switch */}
@@ -401,32 +488,27 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on >
}
+ {/* 附件 */} + {!inputLock.locked && } {/* send */}
{ !inputLock.locked && handleSendClick() }}> - + onClick={() => { !inputLock.locked && !fileUploading && handleSendClick() }}> +
{/* stop & 重置 */} -
- {stop.show ? null - // - : +
+ {!stop.show && }
{/* question */} @@ -437,7 +519,7 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on style={{ height: 56 }} disabled={inputLock.locked} onInput={handleTextAreaHeight} - placeholder={inputLock.locked ? inputLock.reason : t('chat.inputPlaceholder')} + placeholder={placholder} className={"resize-none py-4 pr-10 text-md min-h-6 max-h-[200px] scrollbar-hide dark:bg-[#131415] text-gray-800" + (form && ' pl-10')} onKeyDown={(event) => { if (event.key === "Enter" && !event.shiftKey) { @@ -450,3 +532,30 @@ export default function ChatInput({ autoRun, clear, form, wsUrl, onBeforSend, on

{appConfig.dialogTips}

}; + + + +const useFileLoading = (locked) => { + const [loading, setLoading] = useState(false); + const filesRef = useRef([]) + useEffect(() => { + if (locked) filesRef.current = [] + }, [locked]) + return { + fileUploading: loading, + getFileIds: () => filesRef.current, + loadingChange(files: string[] | null) { + if (files) { + setLoading(false) + filesRef.current = files + } else { + setLoading(true) + filesRef.current = [] + } + }, + clear() { + setLoading(false) + filesRef.current = [] + } + } +} \ No newline at end of file diff --git a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatMessages.tsx b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatMessages.tsx index 3fa545a74..a3f86d92d 100644 --- a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatMessages.tsx +++ b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatMessages.tsx @@ -9,21 +9,30 @@ import { useTranslation } from "react-i18next"; // import RunLog from "./RunLog"; // import Separator from "./Separator"; import Separator from "@/components/bs-comp/chatComponent/Separator"; +import InputForm from "./InputForm"; import MessageBs from "./MessageBs"; import MessageBsChoose from "./MessageBsChoose"; import MessageNodeRun from "./MessageNodeRun"; import { useMessageStore } from "./messageStore"; import MessageUser from "./MessageUser"; -export default function ChatMessages({ mark = false, logo, useName, disableBtn = false, guideWord, loadMore, onMarkClick }) { +export default function ChatMessages({ + debug, + mark = false, + logo, + useName, + guideWord, + loadMore, + onMarkClick = undefined +}) { const { t } = useTranslation() - const { chatId, messages, hisMessages } = useMessageStore() + const { chatId, messages, inputForm } = useMessageStore() // 反馈 const thumbRef = useRef(null) // 溯 const sourceRef = useRef(null -) + ) // 自动滚动 const messagesRef = useRef(null) const scrollLockRef = useRef(false) @@ -81,7 +90,7 @@ export default function ChatMessages({ mark = false, logo, useName, disableBtn = // 成对的qa msg const findQa = (msgs, index) => { const item = msgs[index] - if (['stream_msg', 'answer'].includes(item.category)) { + if (['stream_msg', 'answer', 'output_msg'].includes(item.category)) { const a = item.message.msg || item.message let q = '' while (index > -1) { @@ -97,7 +106,7 @@ export default function ChatMessages({ mark = false, logo, useName, disableBtn = let a = '' while (msgs[++index]) { const aItem = msgs[index] - if (['stream_msg', 'answer'].includes(aItem.category)) { + if (['stream_msg', 'answer', 'output_msg'].includes(aItem.category)) { a = aItem.message.msg || aItem.message break } @@ -114,19 +123,26 @@ export default function ChatMessages({ mark = false, logo, useName, disableBtn = case 'input': return null case 'question': - return { onMarkClick('question', msg.id, findQa(messagesList, index)) }} />; + return { onMarkClick?.('question', msg.id, findQa(messagesList, index)) }} + />; case 'guide_word': case 'output_msg': case 'stream_msg': + case "answer": return { thumbRef.current?.openModal(chatId) }} onSource={(data) => { sourceRef.current?.openModal(data) }} - onMarkClick={() => onMarkClick('answer', msg.message_id, findQa(messagesList, index))} + onMarkClick={() => onMarkClick?.('answer', msg.message_id, findQa(messagesList, index))} />; case 'separator': return ; @@ -141,6 +157,7 @@ export default function ChatMessages({ mark = false, logo, useName, disableBtn = } }) } + {inputForm && }
diff --git a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatPane.tsx b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatPane.tsx index 64baa4266..fb41ec21e 100644 --- a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatPane.tsx +++ b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/ChatPane.tsx @@ -1,15 +1,39 @@ +import { getFlowApi } from "@/controllers/API/flow"; import { useEffect } from "react"; import Chat from "./Chat"; import { useMessageStore } from "./messageStore"; -export default function ChatPane({ autoRun = false, chatId, flow, wsUrl = '' }: { autoRun?: boolean, chatId: string, flow: any, wsUrl?: string }) { +export default function ChatPane({ debug = false, autoRun = false, chatId, flow, wsUrl = '' }: { debug?: boolean, autoRun?: boolean, chatId: string, flow: any, wsUrl?: string }) { const { changeChatId } = useMessageStore() useEffect(() => { changeChatId(chatId) }, [chatId]) - const getMessage = (action, { nodeId, msg, category, extra, source, message_id }) => { + const getMessage = (action, { nodeId, msg, category, extra, files, source, message_id }) => { + if (action === 'refresh_flow') { + return getFlowApi(flow.id, 'v1').then(f => { + const { data, ...other } = f + const { edges, nodes, viewport } = data + return { + action: 'init_data', + chat_id: chatId.startsWith('test') ? undefined : chatId, + flow_id: flow.id, + data: { + ...other, + edges, + nodes, + viewport + } + } + }) + } + if (action === 'flowInfo') { + return { + flow_id: flow.id, + chat_id: chatId.startsWith('test') ? undefined : chatId, + } + } if (action === 'getInputForm') { const node = flow.nodes.find(node => node.id === nodeId) if (node.data.tab.value === 'input') return null @@ -38,10 +62,13 @@ export default function ChatPane({ autoRun = false, chatId, flow, wsUrl = '' }: ) return { action, + flow_id: flow.id, + chat_id: chatId.startsWith('test') ? undefined : chatId, data: { [nodeId]: { data: { - [variable]: msg + [variable]: msg, + dialog_files_content: files }, message: msg, message_id, @@ -62,6 +89,7 @@ export default function ChatPane({ autoRun = false, chatId, flow, wsUrl = '' }: } return { @@ -25,7 +26,8 @@ export const ChatTest = forwardRef((props, ref) => { setOpen(true); // 通过 `run` 方法打开 `Sheet` setSmall(false); - setFlow(flow); + // 克隆flow用于调试 + setFlow(cloneDeep(flow)); setChatId(`test_${generateUUID(16)}`); }, 0); }, @@ -72,7 +74,7 @@ export const ChatTest = forwardRef((props, ref) => { if (!open) return null; - const host = appConfig.websocketHost || ''; + const host = appConfig.websocketHost || window.location.host; return (
{
e.stopPropagation()}> - +
{!small &&
void }) => { +const InputForm = ({ data }: { data: WorkflowNodeParam }) => { const { t } = useTranslation() const formDataRef = useRef(data.value.reduce((map, item) => { @@ -56,91 +57,99 @@ const InputForm = ({ data, onSubmit }: { data: WorkflowNodeParam, onSubmit: (dat variant: 'warning' }) } - onSubmit([valuesObject, stringObject]) + const myEvent = new CustomEvent('inputFormEvent', { + detail: { + data: valuesObject, + msg: stringObject + } + }); + document.dispatchEvent(myEvent); } const [multiVal, setMultiVal] = useState([]) - return
-
- { - data.value.map((item, i) => ( -
- {item.required && *} - {item.value} - {/* {item.required ? " *" : ""} */} -
- {(() => { - switch (item.type) { - case FormItemType.Text: - return ( - handleChange(item, val)} - /> - ) - case FormItemType.Select: - return ( - item.multiple ? - ({ - label: el.text, - value: el.text - })) - } - placeholder={'请选择'} - onChange={(v) => { - setMultiVal(prev => ({ ...prev, [item.key]: v })); - handleChange(item, v.join(',')) - }} - > - {/* {children?.(reload)} */} - - : - ) - case FormItemType.File: - return ( - updataFileName(item, name)} - fileTypes={["png", "jpg", "jpeg", "doc", "docx", "ppt", "pptx", "xls", "xlsx", "txt", "md", "html", "pdf"]} - suffixes={['xxx']} - onFileChange={(val) => handleChange(item, val)} - /> - ) - default: - return null - } - })()} + return
+
+
+ { + data.value.map((item, i) => ( +
+ {item.required && *} + {item.value} + {/* {item.required ? " *" : ""} */} +
+ {(() => { + switch (item.type) { + case FormItemType.Text: + return ( + handleChange(item, val)} + /> + ) + case FormItemType.Select: + return ( + item.multiple ? + ({ + label: el.text, + value: el.text + })) + } + placeholder={'请选择'} + onChange={(v) => { + setMultiVal(prev => ({ ...prev, [item.key]: v })); + handleChange(item, v.join(',')) + }} + > + {/* {children?.(reload)} */} + + : + ) + case FormItemType.File: + return ( + updataFileName(item, name)} + // fileTypes={FileTypes[item.file_type.toUpperCase()]} + suffixes={FileTypes[item.file_type.toUpperCase()]} + onFileChange={(val) => handleChange(item, val)} + /> + ) + default: + return null + } + })()} +
-
- )) - } + )) + } + +
-
}; diff --git a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/MessageBs.tsx b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/MessageBs.tsx index 006b908f4..b2008a591 100644 --- a/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/MessageBs.tsx +++ b/src/frontend/platform/src/pages/BuildPage/flow/FlowChat/MessageBs.tsx @@ -1,11 +1,14 @@ import MessageButtons from "@/components/bs-comp/chatComponent/MessageButtons"; import SourceEntry from "@/components/bs-comp/chatComponent/SourceEntry"; +import { ToastIcon } from "@/components/bs-icons"; import { AvatarIcon } from "@/components/bs-icons/avatar"; import { LoadIcon, LoadingIcon } from "@/components/bs-icons/loading"; +import { cname } from "@/components/bs-ui/utils"; import { CodeBlock } from "@/modals/formModal/chatMessage/codeBlock"; import { WorkflowMessage } from "@/types/flow"; import { formatStrTime } from "@/util/utils"; import { copyText } from "@/utils"; +import { ChevronDown } from "lucide-react"; import { useMemo, useRef, useState } from "react"; import ReactMarkdown from "react-markdown"; import rehypeMathjax from "rehype-mathjax"; @@ -13,9 +16,6 @@ import remarkGfm from "remark-gfm"; import remarkMath from "remark-math"; import ChatFile from "./ChatFileFile"; import { useMessageStore } from "./messageStore"; -import { ChevronDown } from "lucide-react"; -import { cname } from "@/components/bs-ui/utils"; -import { ToastIcon } from "@/components/bs-icons"; // 颜色列表 const colorList = [ @@ -61,7 +61,8 @@ const ReasoningLog = ({ loading, msg = '' }) => {
} -export default function MessageBs({ mark = false, logo, data, onUnlike = () => { }, disableBtn = false, onSource, onMarkClick }: { logo: string, data: WorkflowMessage, onUnlike?: any, onSource?: any }) { +export default function MessageBs({ debug, mark = false, logo, data, onUnlike = () => { }, onSource, onMarkClick }: + { debug?: boolean, ogo: string, data: WorkflowMessage, onUnlike?: any, onSource?: any }) { const avatarColor = colorList[ (data.sender?.split('').reduce((num, s) => num + s.charCodeAt(), 0) || 0) % colorList.length ] @@ -167,7 +168,7 @@ export default function MessageBs({ mark = false, logo, data, onUnlike = () => { message, })} /> - {!disableBtn && { + return typeof data.files === 'string' ? [] : data.files + }, [data.files]) + + // hack + if (typeof data.files === 'string') return null + return
@@ -135,7 +142,7 @@ export default function MessageBsChoose({ type = 'choose', logo, data }: { type?
{mkdown}
{/* files */}
- {data.files?.map((file) =>
handleDownloadFile(file)} > @@ -152,10 +159,10 @@ export default function MessageBsChoose({ type = 'choose', logo, data }: { type? {type === 'input' ?