Feat/1.2.0 (#1308)

This commit is contained in:
GuoQing Zhang
2025-05-19 21:48:40 +08:00
committed by GitHub
319 changed files with 7288 additions and 3083 deletions
+206 -16
View File
@@ -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
+5 -1
View File
@@ -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"]
+1
View File
@@ -1,3 +1,4 @@
# 毕昇后端代码
* Dockerfile 使用 poetry 进行 Python 依赖管理
+5
View File
@@ -25,3 +25,8 @@ class UnAuthorizedError(BaseErrorCode):
class NotFoundError(BaseErrorCode):
Code: int = 404
Msg: str = '资源不存在'
class ServerError(BaseErrorCode):
Code: int = 500
Msg: str = '服务器错误'
+1 -1
View File
@@ -24,7 +24,7 @@ class KnowledgeSimilarError(BaseErrorCode):
class KnowledgeQAError(BaseErrorCode):
Code: int = 10930
Msg: str = '该问题已被标注过'
Msg: str = '该问题已存在'
class KnowledgeCPError(BaseErrorCode):
+35 -22
View File
@@ -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):
@@ -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])
@@ -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)
+1 -1
View File
@@ -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):
+16 -2
View File
@@ -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
@@ -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('{<file_title>'):
return chunk.split(cls.chunk_split)[-1]
return chunk.split('</file_abstract>\n<paragraph_content>')[-1].rstrip(
'</paragraph_content>}')
chunk = chunk.split('<paragraph_content>')[-1]
chunk = chunk.split('</paragraph_content>')[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 <think>.*</think> tag content
title = re.sub('<think>.*</think>', '', 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
+35 -7
View File
@@ -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,
@@ -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]:
+198
View File
@@ -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
+1 -1
View File
@@ -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
@@ -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
@@ -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:
+70 -79
View File
@@ -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})
+5 -5
View File
@@ -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列表'),
+5 -3
View File
@@ -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': []}
+6 -4
View File
@@ -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'))
+82 -59
View File
@@ -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,
+7 -7
View File
@@ -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(),
+5 -6
View File
@@ -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:
+3 -3
View File
@@ -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()):
+16 -16
View File
@@ -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()
+4 -4
View File
@@ -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()
+232 -44
View File
@@ -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})
+42 -89
View File
@@ -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)
+9 -4
View File
@@ -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)
+1 -1
View File
@@ -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
@@ -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
@@ -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):
@@ -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
+10 -2
View File
@@ -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')
+81 -71
View File
@@ -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
+5 -5
View File
@@ -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:
+8 -6
View File
@@ -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:
+8 -8
View File
@@ -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列表')):
+5 -5
View File
@@ -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:
+10 -18
View File
@@ -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,
+2 -4
View File
@@ -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)
+2 -2
View File
@@ -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,
+4 -4
View File
@@ -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'),
+4 -4
View File
@@ -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)
+2 -2
View File
@@ -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的用户信息来做权限校验
+4 -8
View File
@@ -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),
+6 -3
View File
@@ -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 = []
# 非流式返回累计的事件列表
+1 -1
View File
@@ -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)
+6
View File
@@ -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)
@@ -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
+34 -12
View File
@@ -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,
+3 -1
View File
@@ -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."""
+4 -1
View File
@@ -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('<think>.*</think>', '', 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')
+13
View File
@@ -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' # 答案
@@ -27838,7 +27838,7 @@
{
"update_time": "2025-03-03 20:06:19",
"parameters": null,
"name": "「多助手并行+穿行报告生成」",
"name": "「多助手并行+串行报告生成」",
"description": "",
"flow_id": "497051c652b94cc3a2ffb482c520b1bd",
"api_parameters": null,
+1 -1
View File
@@ -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')
@@ -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):
@@ -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地址")
+30 -5
View File
@@ -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)
@@ -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')))
+10 -12
View File
@@ -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):
+9 -11
View File
@@ -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):
@@ -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):
+20 -16
View File
@@ -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(',')
+20 -19
View File
@@ -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())
@@ -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:
@@ -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()
+10 -10
View File
@@ -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):
@@ -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):
@@ -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):
@@ -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:
@@ -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:
@@ -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
@@ -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))
@@ -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)
+22 -20
View File
@@ -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]:
@@ -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):
@@ -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):
@@ -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):
+12 -14
View File
@@ -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):
+13 -15
View File
@@ -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()
@@ -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:
+10 -12
View File
@@ -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):
+8 -12
View File
@@ -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()
@@ -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')))
+11 -15
View File
@@ -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):
@@ -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):
+28 -30
View File
@@ -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):
@@ -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):
@@ -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):
@@ -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
+1 -1
View File
@@ -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)
@@ -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:
@@ -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]
@@ -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()
+21 -12
View File
@@ -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
@@ -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 = {
+3 -2
View File
@@ -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
@@ -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

Some files were not shown because too many files have changed in this diff Show More