Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion backend/apps/chat/models/chat_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,10 +175,12 @@ class RenameChat(BaseModel):
brief: str = ''
brief_generate: bool = True


class SimpleChat(BaseModel):
id: int = None
brief: str = ''


class ChatItem(BaseModel):
id: Optional[int] = None
oid: Optional[int] = None
Expand All @@ -195,6 +197,7 @@ class ChatItem(BaseModel):
recommended_generate: Optional[bool] = False
latest_record_time: Optional[datetime] = None


class ChatInfo(BaseModel):
id: Optional[int] = None
create_time: datetime = None
Expand Down Expand Up @@ -362,6 +365,7 @@ def dynamic_user_question(self):
class ChatQuestion(AiModelQuestion):
chat_id: int
datasource_id: Optional[int] = None
custom_model: Optional[str | int] = None


class ChatMcp(ChatQuestion):
Expand All @@ -373,6 +377,10 @@ class McpDs(BaseModel):
oid: Optional[str] = Body(description='组织ID,如果不传则为最后一次登录SQLBot时所使用的组织ID', default=None)


class WsMcp(BaseModel):
oid: Optional[str | int] = Body(description='组织ID')


class ChatToken(BaseModel):
username: str = Body(description='用户名')
password: str = Body(description='密码')
Expand All @@ -383,7 +391,7 @@ class ChatStart(BaseModel):
password: str = Body(description='密码', default=None)
token: str = Body(description='token', default=None)
oid: Optional[str] = Body(
description='组织ID,仅当数据源ID为空时有效,如果不传则为最后一次登录SQLBot时所使用的组织ID', default=None)
description='组织ID,如果不传则为最后一次登录SQLBot时所使用的组织ID', default=None)


class ChatQuestionBase(BaseModel):
Expand All @@ -397,6 +405,7 @@ class McpQuestion(ChatQuestionBase):
lang: Optional[str] = Body(description='语言:zh-CN|zh-TW|en|ko-KR', default='zh-CN')
datasource_id: Optional[int | str] = Body(description='数据源ID,仅当当前对话没有确定数据源时有效', default=None)
return_img: Optional[bool] = Body(description='是否返回图表,默认为true开启, 关闭false则仅返回数据', default=True)
custom_model: Optional[str | int] = Body(description='模型ID', default=None)


class AxisObj(BaseModel):
Expand Down
4 changes: 4 additions & 0 deletions backend/apps/chat/task/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,6 +226,10 @@ async def create(cls, *args, **kwargs):
if any(str(model.id) == str(args[3].custom_model) for model in _ai_model_list):
specialized_model_id = args[3].custom_model
print("use custom model: id[" + specialized_model_id + "]")
if args[2] and args[2].custom_model:
if any(str(model.id) == str(args[2].custom_model) for model in _ai_model_list):
specialized_model_id = args[2].custom_model
print("use custom model: id[" + specialized_model_id + "]")
config: LLMConfig = await get_default_config(specialized_model_id)
instance = cls(*args, **kwargs, config=config)

Expand Down
14 changes: 7 additions & 7 deletions backend/apps/mcp/mcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,9 @@

from apps.chat.api.chat import create_chat, question_answer_inner
from apps.chat.models.chat_model import ChatMcp, CreateChat, ChatStart, McpQuestion, McpAssistant, ChatQuestion, \
ChatFinishStep, McpDs, ChatToken
ChatFinishStep, McpDs, ChatToken, WsMcp
from apps.datasource.crud.datasource import get_datasource_list
from apps.system.crud.aimodel_manage import get_ai_model_list_by_workspace
from apps.system.crud.user import authenticate, user_ws_options
from apps.system.crud.user import get_db_user
from apps.system.models.system_model import UserWsModel
Expand Down Expand Up @@ -156,11 +157,9 @@ async def datasource_list(session: SessionDep, trans: Trans, mcp_ds: McpDs):
return result


#
#
# @router.get("/model_list", operation_id="get_model_list")
# async def get_model_list(session: SessionDep):
# return session.query(AiModelDetail).all()
@router.post("/mcp_model_list", operation_id="mcp_model_list")
async def get_model_by_ws(session: SessionDep, mcp_oid: WsMcp):
return get_ai_model_list_by_workspace(session, mcp_oid.oid, False)


@router.post("/mcp_question", operation_id="mcp_question")
Expand Down Expand Up @@ -191,7 +190,8 @@ async def mcp_question(session: SessionDep, trans: Trans, chat: McpQuestion):
else:
raise HTTPException(status_code=400, detail="Invalid datasource ID")

mcp_chat = ChatMcp(token=chat.token, chat_id=chat.chat_id, question=chat.question, datasource_id=ds_id)
mcp_chat = ChatMcp(token=chat.token, chat_id=chat.chat_id, question=chat.question, datasource_id=ds_id,
custom_model=chat.custom_model)

return await question_answer_inner(session=session, current_user=session_user, request_question=mcp_chat,
in_chat=False, stream=chat.stream, return_img=chat.return_img)
Expand Down
Loading