diff --git a/backend/apps/chat/models/chat_model.py b/backend/apps/chat/models/chat_model.py index f169cfb17..fe6b4d0a4 100644 --- a/backend/apps/chat/models/chat_model.py +++ b/backend/apps/chat/models/chat_model.py @@ -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 @@ -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 @@ -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): @@ -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='密码') @@ -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): @@ -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): diff --git a/backend/apps/chat/task/llm.py b/backend/apps/chat/task/llm.py index e9ef80420..c3429c59a 100644 --- a/backend/apps/chat/task/llm.py +++ b/backend/apps/chat/task/llm.py @@ -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) diff --git a/backend/apps/mcp/mcp.py b/backend/apps/mcp/mcp.py index 2be45074c..9a05b138e 100644 --- a/backend/apps/mcp/mcp.py +++ b/backend/apps/mcp/mcp.py @@ -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 @@ -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") @@ -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)