mirror of
https://github.com/dataelement/bisheng.git
synced 2026-09-24 23:19:52 +08:00
Feat/1.2.0 (#1308)
This commit is contained in:
+206
-16
@@ -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
|
||||
|
||||
|
||||
@@ -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,3 +1,4 @@
|
||||
# 毕昇后端代码
|
||||
|
||||
* Dockerfile 使用 poetry 进行 Python 依赖管理
|
||||
|
||||
|
||||
@@ -25,3 +25,8 @@ class UnAuthorizedError(BaseErrorCode):
|
||||
class NotFoundError(BaseErrorCode):
|
||||
Code: int = 404
|
||||
Msg: str = '资源不存在'
|
||||
|
||||
|
||||
class ServerError(BaseErrorCode):
|
||||
Code: int = 500
|
||||
Msg: str = '服务器错误'
|
||||
|
||||
@@ -24,7 +24,7 @@ class KnowledgeSimilarError(BaseErrorCode):
|
||||
|
||||
class KnowledgeQAError(BaseErrorCode):
|
||||
Code: int = 10930
|
||||
Msg: str = '该问题已被标注过'
|
||||
Msg: str = '该问题已存在'
|
||||
|
||||
|
||||
class KnowledgeCPError(BaseErrorCode):
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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,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:
|
||||
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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列表'),
|
||||
|
||||
@@ -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': []}
|
||||
|
||||
|
||||
@@ -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'))
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()):
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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列表')):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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'),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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的用户信息来做权限校验
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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 = []
|
||||
# 非流式返回累计的事件列表
|
||||
|
||||
Vendored
+1
-1
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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地址")
|
||||
|
||||
@@ -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')))
|
||||
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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(',')
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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')))
|
||||
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user