返回 VideoClaw
project_helpers.py
根目录 / video-claw / video-claw / backend / api / services / project_helpers.py
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
142 lines PYTHON