chat.py 36 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825
  1. import asyncio
  2. import datetime
  3. import io
  4. import json
  5. import time
  6. import uuid
  7. import fitz
  8. from fastapi import HTTPException
  9. from sqlalchemy import or_
  10. from Log import logger
  11. from app.config.agent_base_url import RG_CHAT_DIALOG, DF_CHAT_AGENT, DF_CHAT_PARAMETERS, RG_CHAT_SESSIONS, \
  12. DF_CHAT_WORKFLOW, DF_UPLOAD_FILE, RG_ORIGINAL_URL, RG_CHAT_UPDATE_URL, DF_WORKFLOW_DRAFT, DF_WORKFLOW_PUBLISH
  13. from app.config.config import settings
  14. from app.config.const import *
  15. from app.models import DialogModel, ApiTokenModel, UserTokenModel, ComplexChatSessionDao, ChatDataRequest, \
  16. ComplexChatDao, KnowledgeModel, UserModel, KnowledgeUserModel
  17. from app.models.v2.session_model import ChatSessionDao, ChatData
  18. from app.service.v2.app_driver.chat_agent import ChatAgent
  19. from app.service.v2.app_driver.chat_data import ChatBaseApply
  20. from app.service.v2.app_driver.chat_dialog import ChatDialog
  21. from app.service.v2.app_driver.chat_workflow import ChatWorkflow
  22. from docx import Document
  23. from dashscope import get_tokenizer # dashscope版本 >= 1.14.0
  24. async def update_session_log(db, session_id: str, message: dict, conversation_id: str):
  25. await ChatSessionDao(db).update_session_by_id(
  26. session_id=session_id,
  27. session=None,
  28. message=message,
  29. conversation_id=conversation_id
  30. )
  31. async def add_session_log(db, session_id: str, question: str, chat_id: str, user_id, event_type: str,
  32. conversation_id: str, agent_type, query: dict = None):
  33. try:
  34. session = await ChatSessionDao(db).update_or_insert_by_id(
  35. session_id=session_id,
  36. name=question[:255],
  37. agent_id=chat_id,
  38. agent_type=agent_type,
  39. tenant_id=user_id,
  40. message={"role": "user", "content": question, "query": query},
  41. conversation_id=conversation_id,
  42. event_type=event_type
  43. )
  44. return session
  45. except Exception as e:
  46. logger.error(e)
  47. return None
  48. async def get_app_token(db, app_id):
  49. app_token = db.query(UserTokenModel).filter_by(id=app_id).first()
  50. if app_token:
  51. return app_token.access_token
  52. return ""
  53. async def get_chat_token(db, app_id):
  54. app_token = db.query(ApiTokenModel).filter_by(app_id=app_id).first()
  55. if app_token:
  56. return app_token.token
  57. return ""
  58. async def get_workflow_token(db):
  59. user_token = db.query(UserTokenModel).filter(UserTokenModel.id == workflow_server).first()
  60. return user_token.access_token if user_token else ""
  61. async def add_chat_token(db, data):
  62. try:
  63. api_token = ApiTokenModel(**data)
  64. db.add(api_token)
  65. db.commit()
  66. except Exception as e:
  67. logger.error(e)
  68. async def get_chat_info(db, chat_id: str):
  69. return db.query(DialogModel).filter_by(id=chat_id, status=Dialog_STATSU_ON).first()
  70. async def get_chat_object(mode):
  71. if mode == workflow_chat:
  72. url = settings.dify_base_url + DF_CHAT_WORKFLOW
  73. return ChatWorkflow(), url
  74. else:
  75. url = settings.dify_base_url + DF_CHAT_AGENT
  76. return ChatAgent(), url
  77. async def get_user_kb(db, user_id: int, kb_ids: list) -> list:
  78. res = []
  79. user = db.query(UserModel).filter(UserModel.id == user_id).first()
  80. if user is None:
  81. return res
  82. query = db.query(KnowledgeModel)
  83. if user.permission != "admin":
  84. klg_list = [j.id for i in user.groups for j in i.knowledges]
  85. for i in db.query(KnowledgeUserModel).filter(KnowledgeUserModel.user_id == user_id,
  86. KnowledgeUserModel.status == 1).all():
  87. if i.kb_id not in klg_list:
  88. klg_list.append(i.kb_id)
  89. query = query.filter(or_(KnowledgeModel.id.in_(klg_list), KnowledgeModel.tenant_id == str(user_id)))
  90. kb_list = query.all()
  91. for kb in kb_list:
  92. if kb.id in kb_ids:
  93. if kb.permission == "team":
  94. res.append(kb.id)
  95. elif kb.tenant_id == str(user_id):
  96. res.append(kb.id)
  97. return res
  98. else:
  99. return kb_ids
  100. async def service_chat_dialog(db, chat_id: str, question: str, session_id: str, user_id: int, mode: str, kb_ids: list):
  101. conversation_id = ""
  102. token = await get_chat_token(db, rg_api_token)
  103. url = settings.fwr_base_url + RG_CHAT_DIALOG.format(chat_id)
  104. kb_id = await get_user_kb(db, user_id, kb_ids)
  105. if not kb_id:
  106. yield "data: " + json.dumps({"message": smart_message_error,
  107. "error": "\n**ERROR**: The agent has no knowledge base to work with!",
  108. "status": http_400},
  109. ensure_ascii=False) + "\n\n"
  110. return
  111. chat = ChatDialog()
  112. session = await add_session_log(db, session_id, question, chat_id, user_id, mode, session_id, RG_TYPE)
  113. if session:
  114. conversation_id = session.conversation_id
  115. message = {"role": "assistant", "answer": "", "reference": {}}
  116. try:
  117. async for ans in chat.chat_completions(url, await chat.complex_request_data(question, kb_id, conversation_id),
  118. await chat.get_headers(token)):
  119. data = {}
  120. error = ""
  121. status = http_200
  122. if ans.get("code", None) == 102:
  123. error = ans.get("message", "error!")
  124. status = http_400
  125. event = smart_message_error
  126. else:
  127. if isinstance(ans.get("data"), bool) and ans.get("data") is True:
  128. event = smart_message_end
  129. else:
  130. data = ans.get("data", {})
  131. # conversation_id = data.get("session_id", "")
  132. if "session_id" in data:
  133. del data["session_id"]
  134. message = data
  135. event = smart_message_cover
  136. message_str = "data: " + json.dumps(
  137. {"event": event, "data": data, "error": error, "status": status, "session_id": session_id},
  138. ensure_ascii=False) + "\n\n"
  139. for i in range(0, len(message_str), max_chunk_size):
  140. chunk = message_str[i:i + max_chunk_size]
  141. # print(chunk)
  142. yield chunk # 发送分块消息
  143. except Exception as e:
  144. logger.error(e)
  145. try:
  146. yield "data: " + json.dumps({"message": smart_message_error,
  147. "error": "\n**ERROR**: " + str(e), "status": http_500},
  148. ensure_ascii=False) + "\n\n"
  149. except:
  150. ...
  151. finally:
  152. message["role"] = "assistant"
  153. await update_session_log(db, session_id, message, conversation_id)
  154. async def data_process(data):
  155. if isinstance(data, str):
  156. return data.replace("dify", "smart")
  157. elif isinstance(data, dict):
  158. for k in list(data.keys()):
  159. if isinstance(k, str) and "dify" in k:
  160. new_k = k.replace("dify", "smart")
  161. data[new_k] = await data_process(data[k])
  162. del data[k]
  163. else:
  164. data[k] = await data_process(data[k])
  165. return data
  166. elif isinstance(data, list):
  167. for i in range(len(data)):
  168. data[i] = await data_process(data[i])
  169. return data
  170. else:
  171. return data
  172. async def service_chat_workflow(db, chat_id: str, chat_data: ChatData, session_id: str, user_id, mode: str):
  173. conversation_id = ""
  174. answer_event = ""
  175. answer_agent = ""
  176. answer_workflow = ""
  177. download_url = ""
  178. message_id = ""
  179. task_id = ""
  180. error = ""
  181. files = []
  182. node_list = []
  183. token = await get_chat_token(db, chat_id)
  184. chat, url = await get_chat_object(mode)
  185. if hasattr(chat_data, "query"):
  186. query = chat_data.query
  187. else:
  188. query = "start new conversation"
  189. session = await add_session_log(db, session_id, query if query else "start new conversation", chat_id, user_id,
  190. mode, conversation_id, DF_TYPE, chat_data.to_dict())
  191. if session:
  192. conversation_id = session.conversation_id
  193. try:
  194. async for ans in chat.chat_completions(url,
  195. await chat.request_data(query, conversation_id, str(user_id), chat_data),
  196. await chat.get_headers(token)):
  197. data = {}
  198. status = http_200
  199. conversation_id = ans.get("conversation_id")
  200. task_id = ans.get("task_id")
  201. if ans.get("event") == message_error:
  202. error = ans.get("message", "参数异常!")
  203. status = http_400
  204. event = smart_message_error
  205. elif ans.get("event") == message_agent:
  206. data = {"answer": ans.get("answer", ""), "id": ans.get("message_id", "")}
  207. answer_agent += ans.get("answer", "")
  208. message_id = ans.get("message_id", "")
  209. event = smart_message_stream
  210. elif ans.get("event") == message_event:
  211. data = {"answer": ans.get("answer", ""), "id": ans.get("message_id", "")}
  212. answer_event += ans.get("answer", "")
  213. message_id = ans.get("message_id", "")
  214. event = smart_message_stream
  215. elif ans.get("event") == message_file:
  216. data = {"url": ans.get("url", ""), "id": ans.get("id", ""),
  217. "type": ans.get("type", "")}
  218. files.append(data)
  219. event = smart_message_file
  220. elif ans.get("event") in [workflow_started, node_started, node_finished]:
  221. data = ans.get("data", {})
  222. data["inputs"] = await data_process(data.get("inputs", {}))
  223. data["outputs"] = await data_process(data.get("outputs", {}))
  224. data["files"] = await data_process(data.get("files", []))
  225. data["process_data"] = ""
  226. if data.get("status") == "failed":
  227. status = http_500
  228. error = data.get("error", "")
  229. node_list.append(ans)
  230. event = [smart_workflow_started, smart_node_started, smart_node_finished][
  231. [workflow_started, node_started, node_finished].index(ans.get("event"))]
  232. elif ans.get("event") == workflow_finished:
  233. data = ans.get("data", {})
  234. answer_workflow = data.get("outputs", {}).get("output", data.get("outputs", {}).get("answer"))
  235. download_url = data.get("outputs", {}).get("download_url")
  236. event = smart_workflow_finished
  237. if data.get("status") == "failed":
  238. status = http_500
  239. error = data.get("error", "")
  240. node_list.append(ans)
  241. elif ans.get("event") == message_end:
  242. event = smart_message_end
  243. else:
  244. continue
  245. yield "data: " + json.dumps(
  246. {"event": event, "data": data, "error": error, "status": status, "task_id": task_id,
  247. "session_id": session_id},
  248. ensure_ascii=False) + "\n\n"
  249. except Exception as e:
  250. logger.error(e)
  251. try:
  252. yield "data: " + json.dumps({"message": smart_message_error,
  253. "error": "\n**ERROR**: " + str(e), "status": http_500},
  254. ensure_ascii=False) + "\n\n"
  255. except:
  256. ...
  257. finally:
  258. await update_session_log(db, session_id, {"role": "assistant",
  259. "answer": answer_event or answer_agent or answer_workflow or error,
  260. "download_url": download_url,
  261. "node_list": node_list, "task_id": task_id, "id": message_id,
  262. "error": error}, conversation_id)
  263. async def service_chat_basic(db, chat_id: str, chat_data: ChatData, session_id: str, user_id, mode: str):
  264. if chat_id == basic_report_talk:
  265. complex_chat = await ComplexChatDao(db).get_complex_chat_by_mode(chat_data.report_mode)
  266. if complex_chat:
  267. ...
  268. async def service_chat_parameters(db, chat_id, user_id):
  269. chat_info = db.query(DialogModel).filter_by(id=chat_id).first()
  270. if not chat_info:
  271. return {}
  272. return chat_info.parameters
  273. async def service_chat_sessions(db, chat_id, name):
  274. token = await get_chat_token(db, rg_api_token)
  275. # print(token)
  276. if not token:
  277. return {}
  278. url = settings.fwr_base_url + RG_CHAT_SESSIONS.format(chat_id)
  279. chat = ChatDialog()
  280. return await chat.chat_sessions(url, {"name": name}, await chat.get_headers(token))
  281. async def service_chat_sessions_list(db, chat_id, current, page_size, user_id, keyword):
  282. total, session_list = await ChatSessionDao(db).get_session_list(
  283. user_id=user_id,
  284. agent_id=chat_id,
  285. keyword=keyword,
  286. page=current,
  287. page_size=page_size
  288. )
  289. return json.dumps({"total": total, "rows": [session.to_dict() for session in session_list]})
  290. async def service_chat_session_log(db, session_id):
  291. session_log = await ChatSessionDao(db).get_session_by_id(session_id)
  292. if not session_log:
  293. return {}
  294. log_info = session_log.log_to_json()
  295. if session_log.event_type == complex_chat:
  296. total, message_list = await ComplexChatSessionDao(db).get_session_list(session_id)
  297. log_info["message"] = [message.log_to_json() for message in message_list[::-1]]
  298. return json.dumps(log_info)
  299. async def service_chat_upload(db, chat_id, file, user_id):
  300. files = []
  301. token = await get_chat_token(db, chat_id)
  302. if not token:
  303. return files
  304. url = settings.dify_base_url + DF_UPLOAD_FILE
  305. chat = ChatBaseApply()
  306. for f in file:
  307. try:
  308. file_content = await f.read()
  309. file_upload = await chat.chat_upload(url, {"file": (f.filename, file_content)}, {"user": str(user_id)},
  310. {'Authorization': f'Bearer {token}'})
  311. try:
  312. tokens = await read_file(file_content, f.filename, f.content_type)
  313. file_upload["tokens"] = tokens
  314. except:
  315. ...
  316. files.append(file_upload)
  317. except Exception as e:
  318. logger.error(e)
  319. return json.dumps(files) if files else ""
  320. async def get_str_token(input_str):
  321. # 获取tokenizer对象,目前只支持通义千问系列模型
  322. tokenizer = get_tokenizer('qwen-turbo')
  323. # 将字符串切分成token并转换为token id
  324. tokens = tokenizer.encode(input_str)
  325. return len(tokens)
  326. async def read_pdf(pdf_stream):
  327. text = ""
  328. with fitz.open(stream=pdf_stream, filetype="pdf") as pdf_document:
  329. for page in pdf_document:
  330. text += page.get_text()
  331. return text
  332. async def read_word(word_stream):
  333. # 使用 python-docx 打开 Word 文件流
  334. doc = Document(io.BytesIO(word_stream))
  335. # 提取每个段落的文本
  336. text = ""
  337. for para in doc.paragraphs:
  338. text += para.text
  339. return text
  340. async def read_file(file, filename, content_type):
  341. text = ""
  342. if content_type == "application/pdf" or filename.endswith('.pdf'):
  343. # 提取 PDF 内容
  344. text = await read_pdf(file)
  345. elif content_type == "application/vnd.openxmlformats-officedocument.wordprocessingml.document" or filename.endswith(
  346. '.docx'):
  347. text = await read_word(file)
  348. return await get_str_token(text)
  349. async def service_chunk_retrieval(query, knowledge_id, top_k, similarity_threshold, api_key):
  350. # print(query)
  351. try:
  352. request_data = json.loads(query)
  353. payload = {
  354. "question": request_data.get("query", ""),
  355. "dataset_ids": request_data.get("dataset_ids", []),
  356. "page_size": top_k,
  357. "similarity_threshold": similarity_threshold if similarity_threshold else 0.2
  358. }
  359. except json.JSONDecodeError as e:
  360. fixed_json = query.replace("'", '"')
  361. try:
  362. request_data = json.loads(fixed_json)
  363. payload = {
  364. "question": request_data.get("query", ""),
  365. "dataset_ids": request_data.get("dataset_ids", []),
  366. "page_size": top_k,
  367. "similarity_threshold": similarity_threshold if similarity_threshold else 0.2
  368. }
  369. except Exception:
  370. payload = {
  371. "question": query,
  372. "dataset_ids": [knowledge_id],
  373. "page_size": top_k,
  374. "similarity_threshold": similarity_threshold if similarity_threshold else 0.2
  375. }
  376. # print(payload)
  377. url = settings.fwr_base_url + RG_ORIGINAL_URL
  378. chat = ChatBaseApply()
  379. response = await chat.chat_post(url, payload, await chat.get_headers(api_key))
  380. if not response:
  381. raise HTTPException(status_code=500, detail="服务异常!")
  382. records = [
  383. {
  384. "content": chunk["content"],
  385. "score": chunk["similarity"],
  386. "title": chunk.get("document_keyword", "Unknown Document"),
  387. "metadata": {"document_id": chunk["document_id"],
  388. "path": f"{settings.fwr_base_url}/document/{chunk['document_id']}?ext={chunk.get('document_keyword').split('.')[-1]}&prefix=document",
  389. 'highlight': chunk.get("highlight"), "image_id": chunk.get("image_id"),
  390. "positions": chunk.get("positions"), }
  391. }
  392. for chunk in response.get("data", {}).get("chunks", [])
  393. ]
  394. # print(len(records))
  395. # print(records)
  396. return records
  397. async def service_base_chunk_retrieval(query, knowledge_id, top_k, similarity_threshold, api_key):
  398. # request_data = json.loads(query)
  399. payload = {
  400. "question": query,
  401. "dataset_ids": [knowledge_id],
  402. "page_size": top_k,
  403. "similarity_threshold": similarity_threshold
  404. }
  405. url = settings.fwr_base_url + RG_ORIGINAL_URL
  406. # url = "http://192.168.20.116:11080/" + RG_ORIGINAL_URL
  407. chat = ChatBaseApply()
  408. response = await chat.chat_post(url, payload, await chat.get_headers(api_key))
  409. if not response:
  410. raise HTTPException(status_code=500, detail="服务异常!")
  411. records = [
  412. {
  413. "content": chunk["content"],
  414. "score": chunk["similarity"],
  415. "title": chunk.get("document_keyword", "Unknown Document"),
  416. "metadata": {"document_id": chunk["document_id"]}
  417. }
  418. for chunk in response.get("data", {}).get("chunks", [])
  419. ]
  420. return records
  421. async def add_complex_log(db, message_id, chat_id, session_id, chat_mode, query, user_id, mode, agent_type,
  422. message_type, conversation_id="", node_data=None, query_data=None):
  423. if not node_data:
  424. node_data = []
  425. if not query_data:
  426. query_data = {}
  427. # print(node_data)
  428. # print("--------------------------------------------------------")
  429. # print(query_data)
  430. try:
  431. complex_log = ComplexChatSessionDao(db)
  432. if not conversation_id:
  433. session = await complex_log.get_session_by_session_id(session_id, chat_id)
  434. if session:
  435. conversation_id = session.conversation_id
  436. await complex_log.create_session(message_id,
  437. chat_id=chat_id,
  438. session_id=session_id,
  439. chat_mode=chat_mode,
  440. message_type=message_type,
  441. content=query,
  442. event_type=mode,
  443. tenant_id=user_id,
  444. conversation_id=conversation_id,
  445. node_data=json.dumps(node_data),
  446. query=json.dumps(query_data),
  447. agent_type=agent_type)
  448. return conversation_id, True
  449. except Exception as e:
  450. logger.error(e)
  451. return conversation_id, False
  452. async def add_query_files(db, message_id):
  453. query = {}
  454. complex_log = await ComplexChatSessionDao(db).get_session_by_id(message_id)
  455. if complex_log:
  456. query = json.loads(complex_log.query)
  457. return query.get("files", [])
  458. async def service_complex_chat(db, chat_id, mode, user_id, chat_request: ChatDataRequest):
  459. answer_event = ""
  460. answer_agent = ""
  461. answer_dialog = ""
  462. answer_workflow = ""
  463. download_url = ""
  464. message_id = ""
  465. task_id = ""
  466. error = ""
  467. node_list = []
  468. reference = {}
  469. conversation_id = ""
  470. query_data = chat_request.to_dict()
  471. new_message_id = str(uuid.uuid4())
  472. inputs = {"is_deep": chat_request.isDeep}
  473. files = chat_request.files
  474. if chat_request.chatMode == complex_content_optimization_chat:
  475. inputs["type"] = chat_request.optimizeType
  476. elif chat_request.chatMode == complex_dialog_chat:
  477. if not files and chat_request.parentId:
  478. files = await add_query_files(db, chat_request.parentId)
  479. if chat_request.chatMode != complex_content_optimization_chat:
  480. await add_session_log(db, chat_request.sessionId, chat_request.query if chat_request.query else "未命名会话",
  481. chat_id, user_id,
  482. mode, "", DF_TYPE)
  483. conversation_id, message = await add_complex_log(db, new_message_id, chat_id, chat_request.sessionId,
  484. chat_request.chatMode, chat_request.query, user_id, mode,
  485. DF_TYPE, 1, query_data=query_data)
  486. if not message:
  487. yield "data: " + json.dumps({"message": smart_message_error,
  488. "error": "\n**ERROR**: 创建会话失败!", "status": http_500},
  489. ensure_ascii=False) + "\n\n"
  490. return
  491. query_data["parentId"] = new_message_id
  492. try:
  493. if chat_request.chatMode == complex_knowledge_chat or chat_request.chatMode == complex_knowledge_chat_deep:
  494. if not conversation_id:
  495. session = await service_chat_sessions(db, chat_id, chat_request.query)
  496. # print(session)
  497. if not session or session.get("code") != 0:
  498. yield "data: " + json.dumps(
  499. {"message": smart_message_error, "error": "\n**ERROR**: chat agent error", "status": http_500})
  500. return
  501. conversation_id = session.get("data", {}).get("id")
  502. token = await get_chat_token(db, rg_api_token)
  503. url = settings.fwr_base_url + RG_CHAT_DIALOG.format(chat_id)
  504. chat = ChatDialog()
  505. try:
  506. async for ans in chat.chat_completions(url, await chat.complex_request_data(chat_request.query,
  507. chat_request.knowledgeId,
  508. conversation_id),
  509. await chat.get_headers(token)):
  510. data = {}
  511. error = ""
  512. status = http_200
  513. if ans.get("code", None) == 102:
  514. error = ans.get("message", "error!")
  515. status = http_400
  516. event = smart_message_error
  517. else:
  518. if isinstance(ans.get("data"), bool) and ans.get("data") is True:
  519. event = smart_message_end
  520. else:
  521. data = ans.get("data", {})
  522. # conversation_id = data.get("session_id", "")
  523. if "session_id" in data:
  524. del data["session_id"]
  525. data["prompt"] = ""
  526. if not message_id:
  527. message_id = data.get("id", "")
  528. answer_dialog = data.get("answer", "")
  529. reference = data.get("reference", {})
  530. event = smart_message_cover
  531. message_str = "data: " + json.dumps(
  532. {"event": event, "data": data, "error": error, "status": status, "message_id": message_id,
  533. "parent_id": new_message_id,
  534. "session_id": chat_request.sessionId},
  535. ensure_ascii=False) + "\n\n"
  536. for i in range(0, len(message_str), max_chunk_size):
  537. chunk = message_str[i:i + max_chunk_size]
  538. # print(chunk)
  539. yield chunk # 发送分块消息
  540. except Exception as e:
  541. logger.error(e)
  542. try:
  543. yield "data: " + json.dumps({"message": smart_message_error,
  544. "error": "\n**ERROR**: " + str(e), "status": http_500},
  545. ensure_ascii=False) + "\n\n"
  546. except:
  547. ...
  548. else:
  549. token = await get_chat_token(db, chat_id)
  550. chat, url = await get_chat_object(mode)
  551. async for ans in chat.chat_completions(url,
  552. await chat.complex_request_data(chat_request.query, conversation_id,
  553. str(user_id), files=files,
  554. inputs=inputs),
  555. await chat.get_headers(token)):
  556. # print(ans)
  557. data = {}
  558. status = http_200
  559. conversation_id = ans.get("conversation_id")
  560. task_id = ans.get("task_id")
  561. if ans.get("event") == message_error:
  562. error = ans.get("message", "参数异常!")
  563. status = http_400
  564. event = smart_message_error
  565. elif ans.get("event") == message_agent:
  566. data = {"answer": ans.get("answer", ""), "id": ans.get("message_id", "")}
  567. answer_agent += ans.get("answer", "")
  568. message_id = ans.get("message_id", "")
  569. event = smart_message_stream
  570. elif ans.get("event") == message_event:
  571. data = {"answer": ans.get("answer", ""), "id": ans.get("message_id", "")}
  572. answer_event += ans.get("answer", "")
  573. message_id = ans.get("message_id", "")
  574. event = smart_message_stream
  575. elif ans.get("event") == message_file:
  576. data = {"url": ans.get("url", ""), "id": ans.get("id", ""),
  577. "type": ans.get("type", "")}
  578. files.append(data)
  579. event = smart_message_file
  580. elif ans.get("event") in [workflow_started, node_started, node_finished]:
  581. data = ans.get("data", {})
  582. data["inputs"] = await data_process(data.get("inputs", {}))
  583. data["outputs"] = await data_process(data.get("outputs", {}))
  584. data["files"] = await data_process(data.get("files", []))
  585. data["process_data"] = ""
  586. if data.get("status") == "failed":
  587. status = http_500
  588. error = data.get("error", "")
  589. node_list.append(ans)
  590. event = [smart_workflow_started, smart_node_started, smart_node_finished][
  591. [workflow_started, node_started, node_finished].index(ans.get("event"))]
  592. elif ans.get("event") == workflow_finished:
  593. data = ans.get("data", {})
  594. answer_workflow = data.get("outputs", {}).get("output", data.get("outputs", {}).get("answer"))
  595. download_url = data.get("outputs", {}).get("download_url")
  596. event = smart_workflow_finished
  597. if data.get("status") == "failed":
  598. status = http_500
  599. error = data.get("error", "")
  600. node_list.append(ans)
  601. elif ans.get("event") == message_end:
  602. event = smart_message_end
  603. else:
  604. continue
  605. yield "data: " + json.dumps(
  606. {"event": event, "data": data, "error": error, "status": status, "task_id": task_id,
  607. "message_id": message_id,
  608. "parent_id": new_message_id,
  609. "session_id": chat_request.sessionId},
  610. ensure_ascii=False) + "\n\n"
  611. except Exception as e:
  612. logger.error(e)
  613. try:
  614. yield "data: " + json.dumps({"message": smart_message_error,
  615. "error": "\n**ERROR**: " + str(e), "status": http_500},
  616. ensure_ascii=False) + "\n\n"
  617. except:
  618. ...
  619. finally:
  620. # await update_session_log(db, session_id, {"role": "assistant",
  621. # "answer": answer_event or answer_agent or answer_workflow or error,
  622. # "download_url": download_url,
  623. # "node_list": node_list, "task_id": task_id, "id": message_id,
  624. # "error": error}, conversation_id)
  625. if message_id:
  626. await add_complex_log(db, message_id, chat_id, chat_request.sessionId, chat_request.chatMode,
  627. answer_event or answer_agent or answer_workflow or answer_dialog or error, user_id,
  628. mode, DF_TYPE, 2, conversation_id, node_data=node_list or reference,
  629. query_data=query_data)
  630. async def service_complex_upload(db, chat_id, file, user_id):
  631. files = []
  632. token = await get_chat_token(db, chat_id)
  633. if not token:
  634. return files
  635. url = settings.dify_base_url + DF_UPLOAD_FILE
  636. chat = ChatBaseApply()
  637. for f in file:
  638. try:
  639. file_content = await f.read()
  640. file_upload = await chat.chat_upload(url, {"file": (f.filename, file_content)}, {"user": str(user_id)},
  641. {'Authorization': f'Bearer {token}'})
  642. # try:
  643. # tokens = await read_file(file_content, f.filename, f.content_type)
  644. # file_upload["tokens"] = tokens
  645. # except:
  646. # ...
  647. files.append(file_upload)
  648. except Exception as e:
  649. logger.error(e)
  650. return json.dumps(files) if files else ""
  651. async def service_complex_model(db, chat_type, model_type, model_name, model_provider):
  652. if chat_type == 1 and model_type == 1:
  653. return await set_dialog_model(db, complex_knowledge_chat, model_name,model_provider)
  654. elif chat_type == 1 and model_type == 2:
  655. return await set_dialog_model(db, complex_knowledge_chat_deep, model_name, model_provider)
  656. else:
  657. if model_type == 1:
  658. chats = [complex_dialog_chat, complex_network_chat, complex_mindmap_chat, complex_content_optimization_chat]
  659. else:
  660. chats = [complex_dialog_chat, complex_network_chat]
  661. return await set_workflow_model(db,chats
  662. , # , complex_network_chat, complex_mindmap_chat, complex_content_optimization_chat
  663. model_type, model_name, model_provider)
  664. async def set_dialog_model(db, chat_mode, model_name, model_provider):
  665. chat = await ComplexChatDao(db).get_complex_chat_by_mode(chat_mode)
  666. if chat:
  667. access_token = await get_chat_token(db, rg_api_token)
  668. url = settings.fwr_base_url + RG_CHAT_UPDATE_URL.format(chat.id)
  669. chat_base = ChatBaseApply()
  670. payload = {
  671. "name": chat.name,
  672. "llm": {
  673. "model_name": model_name
  674. }
  675. }
  676. response = await chat_base.chat_put(url, payload, await chat_base.get_headers(access_token))
  677. # print(response)
  678. if not response:
  679. return "服务异常,修改失败!"
  680. await ComplexChatDao(db).update_complex_chat_by_id(chat.id, {"chat_model": model_name, "chat_model_ds": model_name, "chat_provider": model_provider, "update_date": datetime.datetime.now()})
  681. return ""
  682. async def set_workflow_model(db, chat_modes, model_type, model_name, model_provider):
  683. chat_base = ChatBaseApply()
  684. token = await get_workflow_token(db)
  685. for chat_mode in chat_modes:
  686. chat = await ComplexChatDao(db).get_complex_chat_by_mode(chat_mode)
  687. if chat:
  688. get_draft_url = settings.dify_base_url + DF_WORKFLOW_DRAFT.format(chat.id)
  689. draft_data = await chat_base.chat_get(get_draft_url, {}, await chat_base.get_headers(token))
  690. if draft_data:
  691. graph = draft_data.get("graph")
  692. for node in graph.get("nodes"):
  693. if node.get("data", {}).get("type") == "llm":
  694. if model_type == 1 and "深度搜索" not in node.get("data", {}).get("title"):
  695. node["data"]["model"]["name"] = model_name
  696. node["data"]["model"]["provider"] = model_provider
  697. elif model_type == 2 and "深度搜索" in node.get("data", {}).get("title"):
  698. node["data"]["model"]["name"] = model_name
  699. node["data"]["model"]["provider"] = model_provider
  700. draft_data_query = {"conversation_variables": draft_data.get("conversation_variables"),
  701. "environment_variables": draft_data.get("environment_variables"),
  702. "hash": draft_data.get("hash"),
  703. "features": draft_data.get("features"),
  704. "graph": graph}
  705. set_draft_data = await chat_base.chat_post(get_draft_url, draft_data_query, await chat_base.get_headers(token))
  706. if set_draft_data and set_draft_data.get("result") == "success":
  707. publish_url = settings.dify_base_url + DF_WORKFLOW_PUBLISH.format(chat.id)
  708. publish_data = await chat_base.chat_post(publish_url, {}, await chat_base.get_headers(token))
  709. if publish_data and publish_data.get("result") == "success":
  710. update_kwargs = {"chat_provider": model_provider, "update_date": datetime.datetime.now()}
  711. if model_type == 1:
  712. update_kwargs["chat_model"] = model_name
  713. else:
  714. update_kwargs["chat_model_ds"] = model_name
  715. await ComplexChatDao(db).update_complex_chat_by_id(chat.id, update_kwargs)
  716. async def service_get_complex_model(db):
  717. res = {}
  718. for complexs in await ComplexChatDao(db).aget_complex_chat():
  719. if complexs.chat_mode == complex_knowledge_chat:
  720. res["dialog"] = {"modelName": complexs.chat_model, "modelProvider": complexs.chat_provider}
  721. elif complexs.chat_mode == complex_knowledge_chat_deep:
  722. res["dialog_ds"] = {"modelName": complexs.chat_model_ds, "modelProvider": complexs.chat_provider}
  723. else:
  724. res["workflow"] = {"modelName": complexs.chat_model, "modelProvider": complexs.chat_provider}
  725. res["workflow_ds"] = {"modelName": complexs.chat_model_ds, "modelProvider": complexs.chat_provider}
  726. return json.dumps(res)
  727. if __name__ == "__main__":
  728. q = json.dumps({"query": "设备", "dataset_ids": ["fc68db52f43111efb94a0242ac120004"]})
  729. top_k = 2
  730. similarity_threshold = 0.5
  731. api_key = "ragflow-Y4MGYwY2JlZjM2YjExZWY4ZWU5MDI0Mm"
  732. # a = service_chunk_retrieval(q, top_k, similarity_threshold, api_key)
  733. # print(a)
  734. async def a():
  735. b = await service_chunk_retrieval(q, top_k, similarity_threshold, api_key)
  736. print(b)
  737. asyncio.run(a())