Skip to content

Commit 678e2ec

Browse files
committed
feat: MCP support model
1 parent b42e917 commit 678e2ec

3 files changed

Lines changed: 21 additions & 8 deletions

File tree

‎backend/apps/chat/models/chat_model.py‎

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -175,10 +175,12 @@ class RenameChat(BaseModel):
175175
brief: str = ''
176176
brief_generate: bool = True
177177

178+
178179
class SimpleChat(BaseModel):
179180
id: int = None
180181
brief: str = ''
181182

183+
182184
class ChatItem(BaseModel):
183185
id: Optional[int] = None
184186
oid: Optional[int] = None
@@ -195,6 +197,7 @@ class ChatItem(BaseModel):
195197
recommended_generate: Optional[bool] = False
196198
latest_record_time: Optional[datetime] = None
197199

200+
198201
class ChatInfo(BaseModel):
199202
id: Optional[int] = None
200203
create_time: datetime = None
@@ -362,6 +365,7 @@ def dynamic_user_question(self):
362365
class ChatQuestion(AiModelQuestion):
363366
chat_id: int
364367
datasource_id: Optional[int] = None
368+
custom_model: Optional[str | int] = None
365369

366370

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

375379

380+
class WsMcp(BaseModel):
381+
oid: Optional[str | int] = Body(description='组织ID')
382+
383+
376384
class ChatToken(BaseModel):
377385
username: str = Body(description='用户名')
378386
password: str = Body(description='密码')
@@ -383,7 +391,7 @@ class ChatStart(BaseModel):
383391
password: str = Body(description='密码', default=None)
384392
token: str = Body(description='token', default=None)
385393
oid: Optional[str] = Body(
386-
description='组织ID,仅当数据源ID为空时有效,如果不传则为最后一次登录SQLBot时所使用的组织ID', default=None)
394+
description='组织ID,如果不传则为最后一次登录SQLBot时所使用的组织ID', default=None)
387395

388396

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

401410

402411
class AxisObj(BaseModel):

‎backend/apps/chat/task/llm.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -226,6 +226,10 @@ async def create(cls, *args, **kwargs):
226226
if any(str(model.id) == str(args[3].custom_model) for model in _ai_model_list):
227227
specialized_model_id = args[3].custom_model
228228
print("use custom model: id[" + specialized_model_id + "]")
229+
if args[2] and args[2].custom_model:
230+
if any(str(model.id) == str(args[2].custom_model) for model in _ai_model_list):
231+
specialized_model_id = args[2].custom_model
232+
print("use custom model: id[" + specialized_model_id + "]")
229233
config: LLMConfig = await get_default_config(specialized_model_id)
230234
instance = cls(*args, **kwargs, config=config)
231235

‎backend/apps/mcp/mcp.py‎

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -13,8 +13,9 @@
1313

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

158159

159-
#
160-
#
161-
# @router.get("/model_list", operation_id="get_model_list")
162-
# async def get_model_list(session: SessionDep):
163-
# return session.query(AiModelDetail).all()
160+
@router.post("/mcp_model_list", operation_id="mcp_model_list")
161+
async def get_model_by_ws(session: SessionDep, mcp_oid: WsMcp):
162+
return get_ai_model_list_by_workspace(session, mcp_oid.oid, False)
164163

165164

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

194-
mcp_chat = ChatMcp(token=chat.token, chat_id=chat.chat_id, question=chat.question, datasource_id=ds_id)
193+
mcp_chat = ChatMcp(token=chat.token, chat_id=chat.chat_id, question=chat.question, datasource_id=ds_id,
194+
custom_model=chat.custom_model)
195195

196196
return await question_answer_inner(session=session, current_user=session_user, request_question=mcp_chat,
197197
in_chat=False, stream=chat.stream, return_img=chat.return_img)

0 commit comments

Comments
 (0)