| 1 | import asyncio |
| 2 | import json |
| 3 | import queue |
| 4 | import threading |
| 5 | import time |
| 6 | from typing import Any, AsyncIterator, Callable, Dict, Optional |
| 7 | |
| 8 | from fastapi import Request |
| 9 | |
| 10 | from core.orchestrator import WorkflowStage |
| 11 | |
| 12 | STAGE_NAME_MAP = { |
| 13 | "script_generation": "剧本生成", |
| 14 | "character_design": "角色/场景设计", |
| 15 | "storyboard": "分镜设计", |
| 16 | "reference_generation": "参考图生成", |
| 17 | "video_generation": "视频生成", |
| 18 | "post_production": "后期剪辑", |
| 19 | } |
| 20 | |
| 21 | |
| 22 | def build_openclaw_message(stage: str, result: Dict[str, Any]) -> str: |
| 23 | openclaw_msg = result.get("openclaw_hint", "") |
| 24 | if not openclaw_msg and result.get("requires_intervention", False): |
| 25 | stage_name = STAGE_NAME_MAP.get(stage, stage) |
| 26 | openclaw_msg = f"{stage_name}完成,需要用户确认。请展示给用户并等待用户确认后才能调用 /continue。" |
| 27 | return openclaw_msg |
| 28 | |
| 29 | |
| 30 | def make_progress_channel(): |
| 31 | progress_events = queue.Queue() |
| 32 | event_trigger = asyncio.Event() |
| 33 | loop = asyncio.get_running_loop() |
| 34 | |
| 35 | def progress_callback(phase, step, percent, data=None): |
| 36 | event = {"phase": phase, "step": step, "percent": percent} |
| 37 | if data: |
| 38 | event["data"] = data |
| 39 | progress_events.put(event) |
| 40 | try: |
| 41 | loop.call_soon_threadsafe(event_trigger.set) |
| 42 | except RuntimeError: |
| 43 | pass |
| 44 | |
| 45 | return progress_events, event_trigger, progress_callback |
| 46 | |
| 47 | |
| 48 | def serialize_progress_event(progress: Dict[str, Any]) -> str: |
| 49 | event = { |
| 50 | "type": "progress", |
| 51 | "message": f"{progress['phase']}: {progress['step']}", |
| 52 | "phase": progress["phase"], |
| 53 | "step_desc": progress["step"], |
| 54 | "percent": progress["percent"], |
| 55 | } |
| 56 | if progress.get("data"): |
| 57 | event["data"] = progress["data"] |
| 58 | return json.dumps(event) + "\n" |
| 59 | |
| 60 | |
| 61 | async def stream_workflow_task( |
| 62 | *, |
| 63 | request: Request, |
| 64 | workflow_engine, |
| 65 | state, |
| 66 | stage: str, |
| 67 | input_data: Dict[str, Any], |
| 68 | cancellation_check: Callable[[], bool], |
| 69 | progress_callback: Callable[..., None], |
| 70 | progress_events, |
| 71 | event_trigger: asyncio.Event, |
| 72 | intervention: Optional[Dict[str, Any]] = None, |
| 73 | include_payload_summary: bool = False, |
| 74 | on_disconnect: Optional[Callable[[], None]] = None, |
| 75 | ) -> AsyncIterator[str]: |
| 76 | stage_enum = WorkflowStage(stage) |
| 77 | |
| 78 | try: |
| 79 | task = asyncio.create_task( |
| 80 | workflow_engine.execute_stage( |
| 81 | state, |
| 82 | stage_enum, |
| 83 | input_data, |
| 84 | cancellation_check=cancellation_check, |
| 85 | progress_callback=progress_callback, |
| 86 | intervention=intervention, |
| 87 | ) |
| 88 | ) |
| 89 | |
| 90 | while not task.done(): |
| 91 | try: |
| 92 | await asyncio.wait_for(event_trigger.wait(), timeout=15.0) |
| 93 | except asyncio.TimeoutError: |
| 94 | yield json.dumps({"type": "heartbeat", "time": time.time()}) + "\n" |
| 95 | |
| 96 | event_trigger.clear() |
| 97 | |
| 98 | while not progress_events.empty(): |
| 99 | try: |
| 100 | yield serialize_progress_event(progress_events.get_nowait()) |
| 101 | except queue.Empty: |
| 102 | break |
| 103 | |
| 104 | if await request.is_disconnected(): |
| 105 | workflow_engine.track_background_task(task) |
| 106 | return |
| 107 | |
| 108 | while not progress_events.empty(): |
| 109 | try: |
| 110 | yield serialize_progress_event(progress_events.get_nowait()) |
| 111 | await asyncio.sleep(0) |
| 112 | except queue.Empty: |
| 113 | break |
| 114 | |
| 115 | result = task.result() |
| 116 | status_snapshot = workflow_engine.persist_session_snapshot(state.session_id) |
| 117 | |
| 118 | payload = { |
| 119 | "type": "stage_complete", |
| 120 | "stage": stage, |
| 121 | "status": status_snapshot, |
| 122 | "requires_intervention": result.get("requires_intervention", False), |
| 123 | "openclaw": build_openclaw_message(stage, result), |
| 124 | } |
| 125 | if include_payload_summary: |
| 126 | payload["payload_summary"] = result.get("payload") |
| 127 | yield json.dumps(payload) + "\n" |
| 128 | |
| 129 | except Exception as e: |
| 130 | try: |
| 131 | workflow_engine.persist_session_snapshot(state.session_id) |
| 132 | except Exception: |
| 133 | pass |
| 134 | yield json.dumps({"type": "error", "content": str(e)}) + "\n" |
| 135 | |
| 136 | |
| 137 | def make_cancellation(workflow_engine, session_id: str): |
| 138 | workflow_engine.reset_stop_event(session_id) |
| 139 | session_stop = workflow_engine.get_stop_event(session_id) |
| 140 | request_stop = threading.Event() |
| 141 | return lambda: request_stop.is_set() or session_stop.is_set(), request_stop.set |
| 142 |