run_step.py 3.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105
  1. from datetime import datetime
  2. from typing import List
  3. from sqlalchemy.orm import Session
  4. from sqlmodel import select
  5. from app.exceptions.exception import ResourceNotFoundError, ValidateFailedError
  6. from app.models import RunStep
  7. import logging
  8. class RunStepService:
  9. @staticmethod
  10. def new_run_step(
  11. *, session: Session, type, status="in_progress", assistant_id, thread_id, run_id, step_details
  12. ) -> RunStep:
  13. run_step = RunStep(
  14. type=type,
  15. status=status,
  16. assistant_id=assistant_id,
  17. thread_id=thread_id,
  18. run_id=run_id,
  19. step_details=step_details,
  20. )
  21. session.add(run_step)
  22. session.commit()
  23. session.refresh(run_step)
  24. return run_step
  25. @staticmethod
  26. def get_run_step(*, run_step_id, session: Session) -> RunStep:
  27. run_step = session.execute(select(RunStep).where(RunStep.id == run_step_id)).scalars().one_or_none()
  28. if not run_step:
  29. raise ResourceNotFoundError(f"run_step {run_step_id} not found")
  30. return run_step
  31. @staticmethod
  32. def get_run_step_list(*, run_id, thread_id, session: Session) -> List[RunStep]:
  33. session.rollback()
  34. statement = select(RunStep).where(RunStep.run_id == run_id).where(RunStep.thread_id == thread_id)
  35. result = session.execute(statement).scalars().all()
  36. logging.info("run_id: %s", run_id)
  37. logging.info("run_id: %s", thread_id)
  38. logging.info("result: %s", result)
  39. return result
  40. @staticmethod
  41. def to_cancelled(*, session: Session, run_step_id) -> RunStep:
  42. run_step = RunStepService.get_run_step(run_step_id=run_step_id, session=session)
  43. RunStepService.check_status_in(run_step=run_step, status_list=["in_progress", "cancelled"])
  44. if run_step.status != "cancelled":
  45. run_step.status = "cancelled"
  46. run_step.cancelled_at = datetime.now()
  47. session.add(run_step)
  48. session.commit()
  49. session.refresh(run_step)
  50. return run_step
  51. @staticmethod
  52. def update_step_details(*, session: Session, run_step_id, step_details, completed=False) -> RunStep:
  53. run_step = RunStepService.get_run_step(run_step_id=run_step_id, session=session)
  54. RunStepService.check_status_in(run_step=run_step, status_list=["in_progress", "completed"])
  55. #run_step.step_details = step_details
  56. if isinstance(step_details, dict):
  57. print("step_details is a dict")
  58. new_step_details = dict(run_step.step_details or {})
  59. new_step_details.update(step_details)
  60. run_step.step_details = new_step_details
  61. print("step_details", step_details)
  62. print("run_step.step_details", run_step.step_details)
  63. else:
  64. run_step.step_details = step_details
  65. if completed and run_step.status != "completed":
  66. run_step.status = "completed"
  67. run_step.completed_at = datetime.now()
  68. session.add(run_step)
  69. session.commit()
  70. session.refresh(run_step)
  71. return run_step
  72. @staticmethod
  73. def to_failed(*, session: Session, run_step_id, last_error) -> RunStep:
  74. run_step = RunStepService.get_run_step(run_step_id=run_step_id, session=session)
  75. RunStepService.check_status_in(run_step=run_step, status_list=["in_progress", "failed"])
  76. if run_step.status != "failed":
  77. run_step.status = "failed"
  78. run_step.failed_at = datetime.now()
  79. run_step.last_error = {"code": "server_error", "message": str(last_error)}
  80. session.add(run_step)
  81. session.commit()
  82. session.refresh(run_step)
  83. return run_step
  84. @staticmethod
  85. def check_status_in(run_step, status_list):
  86. if run_step.status not in status_list:
  87. raise ValidateFailedError(f"invalid run_step {run_step.id} status {run_step.status}")