Du kannst nicht mehr als 25 Themen auswählen Themen müssen mit entweder einem Buchstaben oder einer Ziffer beginnen. Sie können Bindestriche („-“) enthalten und bis zu 35 Zeichen lang sein.

canvas_app.py 13KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315
  1. #
  2. # Copyright 2024 The InfiniFlow Authors. All Rights Reserved.
  3. #
  4. # Licensed under the Apache License, Version 2.0 (the "License");
  5. # you may not use this file except in compliance with the License.
  6. # You may obtain a copy of the License at
  7. #
  8. # http://www.apache.org/licenses/LICENSE-2.0
  9. #
  10. # Unless required by applicable law or agreed to in writing, software
  11. # distributed under the License is distributed on an "AS IS" BASIS,
  12. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
  13. # See the License for the specific language governing permissions and
  14. # limitations under the License.
  15. #
  16. import json
  17. import traceback
  18. from flask import request, Response
  19. from flask_login import login_required, current_user
  20. from api.db.services.canvas_service import CanvasTemplateService, UserCanvasService
  21. from api.db.services.user_canvas_version import UserCanvasVersionService
  22. from api.settings import RetCode
  23. from api.utils import get_uuid
  24. from api.utils.api_utils import get_json_result, server_error_response, validate_request, get_data_error_result
  25. from agent.canvas import Canvas
  26. from peewee import MySQLDatabase, PostgresqlDatabase
  27. from api.db.db_models import APIToken
  28. import time
  29. @manager.route('/templates', methods=['GET']) # noqa: F821
  30. @login_required
  31. def templates():
  32. return get_json_result(data=[c.to_dict() for c in CanvasTemplateService.get_all()])
  33. @manager.route('/list', methods=['GET']) # noqa: F821
  34. @login_required
  35. def canvas_list():
  36. return get_json_result(data=sorted([c.to_dict() for c in \
  37. UserCanvasService.query(user_id=current_user.id)], key=lambda x: x["update_time"]*-1)
  38. )
  39. @manager.route('/rm', methods=['POST']) # noqa: F821
  40. @validate_request("canvas_ids")
  41. @login_required
  42. def rm():
  43. for i in request.json["canvas_ids"]:
  44. if not UserCanvasService.query(user_id=current_user.id,id=i):
  45. return get_json_result(
  46. data=False, message='Only owner of canvas authorized for this operation.',
  47. code=RetCode.OPERATING_ERROR)
  48. UserCanvasService.delete_by_id(i)
  49. return get_json_result(data=True)
  50. @manager.route('/set', methods=['POST']) # noqa: F821
  51. @validate_request("dsl", "title")
  52. @login_required
  53. def save():
  54. req = request.json
  55. req["user_id"] = current_user.id
  56. if not isinstance(req["dsl"], str):
  57. req["dsl"] = json.dumps(req["dsl"], ensure_ascii=False)
  58. req["dsl"] = json.loads(req["dsl"])
  59. if "id" not in req:
  60. if UserCanvasService.query(user_id=current_user.id, title=req["title"].strip()):
  61. return get_data_error_result(message=f"{req['title'].strip()} already exists.")
  62. req["id"] = get_uuid()
  63. if not UserCanvasService.save(**req):
  64. return get_data_error_result(message="Fail to save canvas.")
  65. else:
  66. if not UserCanvasService.query(user_id=current_user.id, id=req["id"]):
  67. return get_json_result(
  68. data=False, message='Only owner of canvas authorized for this operation.',
  69. code=RetCode.OPERATING_ERROR)
  70. UserCanvasService.update_by_id(req["id"], req)
  71. # save version
  72. UserCanvasVersionService.insert( user_canvas_id=req["id"], dsl=req["dsl"], title="{0}_{1}".format(req["title"], time.strftime("%Y_%m_%d_%H_%M_%S")))
  73. UserCanvasVersionService.delete_all_versions(req["id"])
  74. return get_json_result(data=req)
  75. @manager.route('/get/<canvas_id>', methods=['GET']) # noqa: F821
  76. @login_required
  77. def get(canvas_id):
  78. e, c = UserCanvasService.get_by_id(canvas_id)
  79. if not e:
  80. return get_data_error_result(message="canvas not found.")
  81. return get_json_result(data=c.to_dict())
  82. @manager.route('/getsse/<canvas_id>', methods=['GET']) # type: ignore # noqa: F821
  83. def getsse(canvas_id):
  84. token = request.headers.get('Authorization').split()
  85. if len(token) != 2:
  86. return get_data_error_result(message='Authorization is not valid!"')
  87. token = token[1]
  88. objs = APIToken.query(beta=token)
  89. if not objs:
  90. return get_data_error_result(message='Authentication error: API key is invalid!"')
  91. e, c = UserCanvasService.get_by_id(canvas_id)
  92. if not e:
  93. return get_data_error_result(message="canvas not found.")
  94. return get_json_result(data=c.to_dict())
  95. @manager.route('/completion', methods=['POST']) # noqa: F821
  96. @validate_request("id")
  97. @login_required
  98. def run():
  99. req = request.json
  100. stream = req.get("stream", True)
  101. e, cvs = UserCanvasService.get_by_id(req["id"])
  102. if not e:
  103. return get_data_error_result(message="canvas not found.")
  104. if not UserCanvasService.query(user_id=current_user.id, id=req["id"]):
  105. return get_json_result(
  106. data=False, message='Only owner of canvas authorized for this operation.',
  107. code=RetCode.OPERATING_ERROR)
  108. if not isinstance(cvs.dsl, str):
  109. cvs.dsl = json.dumps(cvs.dsl, ensure_ascii=False)
  110. final_ans = {"reference": [], "content": ""}
  111. message_id = req.get("message_id", get_uuid())
  112. try:
  113. canvas = Canvas(cvs.dsl, current_user.id)
  114. if "message" in req:
  115. canvas.messages.append({"role": "user", "content": req["message"], "id": message_id})
  116. canvas.add_user_input(req["message"])
  117. except Exception as e:
  118. return server_error_response(e)
  119. if stream:
  120. def sse():
  121. nonlocal answer, cvs
  122. try:
  123. for ans in canvas.run(stream=True):
  124. if ans.get("running_status"):
  125. yield "data:" + json.dumps({"code": 0, "message": "",
  126. "data": {"answer": ans["content"],
  127. "running_status": True}},
  128. ensure_ascii=False) + "\n\n"
  129. continue
  130. for k in ans.keys():
  131. final_ans[k] = ans[k]
  132. ans = {"answer": ans["content"], "reference": ans.get("reference", [])}
  133. yield "data:" + json.dumps({"code": 0, "message": "", "data": ans}, ensure_ascii=False) + "\n\n"
  134. canvas.messages.append({"role": "assistant", "content": final_ans["content"], "id": message_id})
  135. canvas.history.append(("assistant", final_ans["content"]))
  136. if not canvas.path[-1]:
  137. canvas.path.pop(-1)
  138. if final_ans.get("reference"):
  139. canvas.reference.append(final_ans["reference"])
  140. cvs.dsl = json.loads(str(canvas))
  141. UserCanvasService.update_by_id(req["id"], cvs.to_dict())
  142. except Exception as e:
  143. cvs.dsl = json.loads(str(canvas))
  144. if not canvas.path[-1]:
  145. canvas.path.pop(-1)
  146. UserCanvasService.update_by_id(req["id"], cvs.to_dict())
  147. traceback.print_exc()
  148. yield "data:" + json.dumps({"code": 500, "message": str(e),
  149. "data": {"answer": "**ERROR**: " + str(e), "reference": []}},
  150. ensure_ascii=False) + "\n\n"
  151. yield "data:" + json.dumps({"code": 0, "message": "", "data": True}, ensure_ascii=False) + "\n\n"
  152. resp = Response(sse(), mimetype="text/event-stream")
  153. resp.headers.add_header("Cache-control", "no-cache")
  154. resp.headers.add_header("Connection", "keep-alive")
  155. resp.headers.add_header("X-Accel-Buffering", "no")
  156. resp.headers.add_header("Content-Type", "text/event-stream; charset=utf-8")
  157. return resp
  158. for answer in canvas.run(stream=False):
  159. if answer.get("running_status"):
  160. continue
  161. final_ans["content"] = "\n".join(answer["content"]) if "content" in answer else ""
  162. canvas.messages.append({"role": "assistant", "content": final_ans["content"], "id": message_id})
  163. if final_ans.get("reference"):
  164. canvas.reference.append(final_ans["reference"])
  165. cvs.dsl = json.loads(str(canvas))
  166. UserCanvasService.update_by_id(req["id"], cvs.to_dict())
  167. return get_json_result(data={"answer": final_ans["content"], "reference": final_ans.get("reference", [])})
  168. @manager.route('/reset', methods=['POST']) # noqa: F821
  169. @validate_request("id")
  170. @login_required
  171. def reset():
  172. req = request.json
  173. try:
  174. e, user_canvas = UserCanvasService.get_by_id(req["id"])
  175. if not e:
  176. return get_data_error_result(message="canvas not found.")
  177. if not UserCanvasService.query(user_id=current_user.id, id=req["id"]):
  178. return get_json_result(
  179. data=False, message='Only owner of canvas authorized for this operation.',
  180. code=RetCode.OPERATING_ERROR)
  181. canvas = Canvas(json.dumps(user_canvas.dsl), current_user.id)
  182. canvas.reset()
  183. req["dsl"] = json.loads(str(canvas))
  184. UserCanvasService.update_by_id(req["id"], {"dsl": req["dsl"]})
  185. return get_json_result(data=req["dsl"])
  186. except Exception as e:
  187. return server_error_response(e)
  188. @manager.route('/input_elements', methods=['GET']) # noqa: F821
  189. @login_required
  190. def input_elements():
  191. cvs_id = request.args.get("id")
  192. cpn_id = request.args.get("component_id")
  193. try:
  194. e, user_canvas = UserCanvasService.get_by_id(cvs_id)
  195. if not e:
  196. return get_data_error_result(message="canvas not found.")
  197. if not UserCanvasService.query(user_id=current_user.id, id=cvs_id):
  198. return get_json_result(
  199. data=False, message='Only owner of canvas authorized for this operation.',
  200. code=RetCode.OPERATING_ERROR)
  201. canvas = Canvas(json.dumps(user_canvas.dsl), current_user.id)
  202. return get_json_result(data=canvas.get_component_input_elements(cpn_id))
  203. except Exception as e:
  204. return server_error_response(e)
  205. @manager.route('/debug', methods=['POST']) # noqa: F821
  206. @validate_request("id", "component_id", "params")
  207. @login_required
  208. def debug():
  209. req = request.json
  210. for p in req["params"]:
  211. assert p.get("key")
  212. try:
  213. e, user_canvas = UserCanvasService.get_by_id(req["id"])
  214. if not e:
  215. return get_data_error_result(message="canvas not found.")
  216. if not UserCanvasService.query(user_id=current_user.id, id=req["id"]):
  217. return get_json_result(
  218. data=False, message='Only owner of canvas authorized for this operation.',
  219. code=RetCode.OPERATING_ERROR)
  220. canvas = Canvas(json.dumps(user_canvas.dsl), current_user.id)
  221. canvas.get_component(req["component_id"])["obj"]._param.debug_inputs = req["params"]
  222. df = canvas.get_component(req["component_id"])["obj"].debug()
  223. return get_json_result(data=df.to_dict(orient="records"))
  224. except Exception as e:
  225. return server_error_response(e)
  226. @manager.route('/test_db_connect', methods=['POST']) # noqa: F821
  227. @validate_request("db_type", "database", "username", "host", "port", "password")
  228. @login_required
  229. def test_db_connect():
  230. req = request.json
  231. try:
  232. if req["db_type"] in ["mysql", "mariadb"]:
  233. db = MySQLDatabase(req["database"], user=req["username"], host=req["host"], port=req["port"],
  234. password=req["password"])
  235. elif req["db_type"] == 'postgresql':
  236. db = PostgresqlDatabase(req["database"], user=req["username"], host=req["host"], port=req["port"],
  237. password=req["password"])
  238. elif req["db_type"] == 'mssql':
  239. import pyodbc
  240. connection_string = (
  241. f"DRIVER={{ODBC Driver 17 for SQL Server}};"
  242. f"SERVER={req['host']},{req['port']};"
  243. f"DATABASE={req['database']};"
  244. f"UID={req['username']};"
  245. f"PWD={req['password']};"
  246. )
  247. db = pyodbc.connect(connection_string)
  248. cursor = db.cursor()
  249. cursor.execute("SELECT 1")
  250. cursor.close()
  251. else:
  252. return server_error_response("Unsupported database type.")
  253. if req["db_type"] != 'mssql':
  254. db.connect()
  255. db.close()
  256. return get_json_result(data="Database Connection Successful!")
  257. except Exception as e:
  258. return server_error_response(e)
  259. #api get list version dsl of canvas
  260. @manager.route('/getlistversion/<canvas_id>', methods=['GET']) # noqa: F821
  261. @login_required
  262. def getlistversion(canvas_id):
  263. try:
  264. list =sorted([c.to_dict() for c in UserCanvasVersionService.list_by_canvas_id(canvas_id)], key=lambda x: x["update_time"]*-1)
  265. return get_json_result(data=list)
  266. except Exception as e:
  267. return get_data_error_result(message=f"Error getting history files: {e}")
  268. #api get version dsl of canvas
  269. @manager.route('/getversion/<version_id>', methods=['GET']) # noqa: F821
  270. @login_required
  271. def getversion( version_id):
  272. try:
  273. e, version = UserCanvasVersionService.get_by_id(version_id)
  274. if version:
  275. return get_json_result(data=version.to_dict())
  276. except Exception as e:
  277. return get_json_result(data=f"Error getting history file: {e}")