You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

ws.py 3.9 kB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697
  1. # api/ws.py
  2. import asyncio
  3. import json
  4. from datetime import datetime
  5. from uuid import UUID
  6. from fastapi import APIRouter, Depends, WebSocket, WebSocketDisconnect
  7. from loguru import logger
  8. from ...datamodel import Run, RunStatus
  9. from ..deps import get_db, get_websocket_manager
  10. from ..managers import WebSocketManager
  11. router = APIRouter()
  12. @router.websocket("/runs/{run_id}")
  13. async def run_websocket(
  14. websocket: WebSocket,
  15. run_id: UUID,
  16. ws_manager: WebSocketManager = Depends(get_websocket_manager),
  17. db=Depends(get_db),
  18. ):
  19. """WebSocket endpoint for run communication"""
  20. # Verify run exists and is in valid state
  21. run_response = db.get(Run, filters={"id": run_id}, return_json=False)
  22. if not run_response.status or not run_response.data:
  23. logger.warning(f"Run not found: {run_id}")
  24. await websocket.close(code=4004, reason="Run not found")
  25. return
  26. run = run_response.data[0]
  27. if run.status not in [RunStatus.CREATED, RunStatus.ACTIVE]:
  28. await websocket.close(code=4003, reason="Run not in valid state")
  29. return
  30. # Connect websocket
  31. connected = await ws_manager.connect(websocket, run_id)
  32. if not connected:
  33. await websocket.close(code=4002, reason="Failed to establish connection")
  34. return
  35. try:
  36. logger.info(f"WebSocket connection established for run {run_id}")
  37. while True:
  38. try:
  39. raw_message = await websocket.receive_text()
  40. message = json.loads(raw_message)
  41. if message.get("type") == "start":
  42. # Handle start message
  43. logger.info(f"Received start request for run {run_id}")
  44. task = message.get("task")
  45. team_config = message.get("team_config")
  46. if task and team_config:
  47. # await ws_manager.start_stream(run_id, task, team_config)
  48. asyncio.create_task(ws_manager.start_stream(run_id, task, team_config))
  49. else:
  50. logger.warning(f"Invalid start message format for run {run_id}")
  51. await websocket.send_json(
  52. {
  53. "type": "error",
  54. "error": "Invalid start message format",
  55. "timestamp": datetime.utcnow().isoformat(),
  56. }
  57. )
  58. elif message.get("type") == "stop":
  59. logger.info(f"Received stop request for run {run_id}")
  60. reason = message.get("reason") or "User requested stop/cancellation"
  61. await ws_manager.stop_run(run_id, reason=reason)
  62. break
  63. elif message.get("type") == "ping":
  64. await websocket.send_json({"type": "pong", "timestamp": datetime.utcnow().isoformat()})
  65. elif message.get("type") == "input_response":
  66. # Handle input response from client
  67. response = message.get("response")
  68. if response is not None:
  69. await ws_manager.handle_input_response(run_id, response)
  70. else:
  71. logger.warning(f"Invalid input response format for run {run_id}")
  72. except json.JSONDecodeError:
  73. logger.warning(f"Invalid JSON received: {raw_message}")
  74. await websocket.send_json(
  75. {"type": "error", "error": "Invalid message format", "timestamp": datetime.utcnow().isoformat()}
  76. )
  77. except WebSocketDisconnect:
  78. logger.info(f"WebSocket disconnected for run {run_id}")
  79. except Exception as e:
  80. logger.error(f"WebSocket error: {str(e)}")
  81. finally:
  82. await ws_manager.disconnect(run_id)