Coverage for apps/streampersona/streampersona_job.py: 61%

255 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-09 04:47 +0000

1""" 

2StreamPersona job to generate a video podcast. 

3It coordinates the execution of the different models. 

4""" 

5 

6import sys 

7import time 

8import json 

9import aiofiles 

10import aiofiles.os 

11import asyncio 

12import math 

13 

14from typing import override 

15from typing import Dict 

16from typing import Any 

17from typing import List 

18from typing import Tuple 

19from typing import Optional 

20from typing import cast 

21 

22from PIL import Image 

23 

24from pptx import Presentation 

25from pptx.exc import PackageNotFoundError 

26 

27from ppt_utils import pptx_to_images 

28from ppt_utils import get_num_slides 

29 

30from persona_prompts import IMG_PROMPT 

31from persona_prompts import IMG_NEG_PROMPT 

32from persona_prompts import VIDEO_PROMPT 

33from persona_prompts import VIDEO_NEG_PROMPT 

34 

35 

36# Local relative imports 

37sys.path.append("..") # noqa: E402 

38sys.path.append("../..") # noqa: E402 

39 

40from streamwise_job import StreamWiseJob 

41from streamwise_job import JobStatus 

42from streamwise_job import OutputMode 

43 

44from lmm_service_manager import LMMServiceManager 

45 

46from client import ServiceError 

47 

48from gen_video_chunked import GenVideoChunked 

49 

50from tts_utils import estimate_num_words_from_audio_duration 

51 

52from console_utils import bytes_to_human 

53 

54from file_utils import read_file_bytes 

55from file_utils import save_base64_as_binary 

56 

57from media_utils import get_audio_duration 

58from media_utils import save_video_audio 

59from media_utils import get_video_frames 

60from media_utils import get_video_file_info 

61from media_utils import get_audio_file_info 

62from media_utils import get_frame_with_text 

63from media_utils import concatenate_videos 

64 

65from k8s_utils import NoActiveContainerError 

66from k8s_utils import NoRunnableContainerError 

67from k8s_utils import ServiceNotFoundError 

68 

69from video import MAX_FT_DURATION_SECS 

70from video import FANTASYTALKING_FPS 

71 

72 

73MAX_LOG_TEXT = 100 

74MAX_IMG_LINE_CHARS = 50 

75 

76 

77def overlay_image_on_image( 

78 base_image: Image.Image, 

79 overlay_image: Image.Image, 

80 position: Tuple[str, str] = ("bottom", "right"), 

81 overlay_percentage: float = 0.25 

82) -> Image.Image: 

83 """ 

84 Overlay an image on top of another image at the specified position. 

85 Positions can be 'top', 'bottom', 'left', 'right', 'center'. 

86 """ 

87 base_width, base_height = base_image.size 

88 overlay_width, overlay_height = overlay_image.size 

89 overlay_image = overlay_image.copy().resize(( 

90 int(base_width * overlay_percentage), 

91 int(base_height * overlay_percentage) 

92 )) 

93 overlay_width, overlay_height = overlay_image.size 

94 

95 if position[0] == "top": 

96 y = 0 

97 elif position[0] == "bottom": 

98 y = base_height - overlay_height 

99 else: # center 

100 y = (base_height - overlay_height) // 2 

101 

102 if position[1] == "left": 

103 x = 0 

104 elif position[1] == "right": 

105 x = base_width - overlay_width 

106 else: # center 

107 x = (base_width - overlay_width) // 2 

108 

109 new_image = base_image.copy() 

110 new_image.paste( 

111 overlay_image, 

112 (x, y), 

113 overlay_image.convert("RGBA")) 

114 return new_image 

115 

116 

117class StreamPersonaJob(StreamWiseJob): 

118 """A job to generate a podcast with images, audio, and video.""" 

119 

120 def __init__( 

121 self, 

122 job_id: str, 

123 service_manager: LMMServiceManager, 

124 config: Dict[str, Any] = {}, 

125 ) -> None: 

126 super().__init__( 

127 "streampersona", 

128 job_id, 

129 service_manager, 

130 config) 

131 self.image: Optional[Image.Image] = None 

132 self.image_task: Optional[asyncio.Task[Image.Image]] = None 

133 

134 @override 

135 async def generate( 

136 self, 

137 job_config: Dict[str, Any], 

138 ) -> None: 

139 pptx_base64 = job_config.get("pptx_base64", None) 

140 if pptx_base64 is None: 

141 raise ValueError("Missing 'pptx_base64' in job config") 

142 await self.gen_persona(pptx_base64) 

143 

144 async def gen_slide_video( 

145 self, 

146 slide_number: int, 

147 slide_text: str, 

148 ) -> Optional[str]: 

149 """ 

150 Generate the video for one slide. 

151 """ 

152 t0 = time.time() 

153 voice = "am_adam" # TODO select voice 

154 audio_base64 = await self.gen.gen_audio( 

155 slide_text, 

156 voice=voice, 

157 task_id=f"{slide_number:03d}", 

158 deadline=self.get_slide_deadline(slide_number), 

159 ) 

160 if not audio_base64: 

161 raise ValueError(f"Cannot generate audio for slide {slide_number} with text '{slide_text}'") 

162 audio_duration = get_audio_duration(audio_base64) 

163 audio_path = f"{self.job_path}/{slide_number:03d}.wav" 

164 audio_len = bytes_to_human(len(audio_base64)) 

165 await save_base64_as_binary(audio_path, audio_base64) 

166 self.logger.info( 

167 f"[{slide_number}] Generated audio with {audio_len} and {audio_duration:.3f} seconds.") 

168 

169 # We need to wait for the image generated 

170 if self.image is None and self.image_task is not None: 

171 self.image = await self.image_task 

172 if self.image is None: 

173 raise ValueError("Image generation task completed with no image.") 

174 width, height = self.image.size 

175 image_path = f"{self.job_path}/persona.png" 

176 self.image.save(image_path) 

177 self.logger.info(f"Image with {width}x{height} pixels saved to '{image_path}'.") 

178 

179 output_mode = self.get_config_output_mode() 

180 if output_mode is OutputMode.AUDIO_ONLY: 

181 return None # We just show the slide, no persona video 

182 

183 if self.image is None: 

184 raise ValueError("Image is required for video generation but not available.") 

185 width, height = self.image.size 

186 video_prompt = VIDEO_PROMPT 

187 video_neg_prompt = VIDEO_NEG_PROMPT 

188 num_steps = self.get_num_steps() 

189 

190 # Generate video for slide 

191 if output_mode == OutputMode.VIDEO_AUDIO_UNSYNCED: 

192 slide_video_binary = await self.gen.gen_video( 

193 img=self.image, 

194 prompt=video_prompt, 

195 neg_prompt=video_neg_prompt, 

196 width=width, 

197 height=height, 

198 video_seconds=audio_duration, # Because of rounding, this may produce more frames than asked 

199 steps=num_steps, 

200 task_id=f"{slide_number:03d}", 

201 wait_request=True, 

202 deadline=self.get_slide_deadline(slide_number), 

203 ) 

204 elif audio_duration < MAX_FT_DURATION_SECS: 

205 slide_video_binary = await self.gen.gen_video_audio_from_img( 

206 img=self.image, 

207 audio_base64=audio_base64, 

208 prompt=video_prompt, 

209 neg_prompt=video_neg_prompt, 

210 width=width, 

211 height=height, 

212 steps=num_steps, 

213 task_id=f"{slide_number:03d}", 

214 deadline=self.get_slide_deadline(slide_number), 

215 ) 

216 else: 

217 # TODO 

218 # slide_video_binary = await self.gen_video_audio_from_img_chunks( 

219 gen_video_chunked = GenVideoChunked( 

220 video_id=slide_number, 

221 gen=self.gen, 

222 job_path=self.job_path, 

223 logger=self.logger) 

224 slide_video_binary = await gen_video_chunked.gen_video_chunked( 

225 audio_path=audio_path, 

226 image=self.image, 

227 prompt=video_prompt, 

228 neg_prompt=video_neg_prompt, 

229 width=width, 

230 height=height, 

231 num_steps=num_steps, 

232 upscaling=self.get_config_bool("upscaling"), 

233 debug=self.get_config_bool("debug_image"), 

234 deadline=self.get_slide_deadline(slide_number), 

235 ) 

236 

237 video_frames = await get_video_frames(slide_video_binary) 

238 video_file_info = get_video_file_info(slide_video_binary) 

239 video_info = video_file_info.get("video", {}) 

240 _video_fps = video_info.get("fps") 

241 video_fps: float = _video_fps if _video_fps is not None else FANTASYTALKING_FPS 

242 

243 video_audio_path = f"{self.job_path}/{slide_number:03d}_persona.mp4" 

244 video_audio_path = await save_video_audio( 

245 video_content=video_frames, 

246 audio_path=audio_path, 

247 out_video_path=video_audio_path, 

248 fps=video_fps) 

249 

250 self._log_video_info(f"[{slide_number}] Generated slide", video_audio_path) 

251 self.logger.info(f"[{slide_number}] Generated slide in {time.time() - t0:.3f} seconds.") 

252 

253 return video_audio_path 

254 

255 def _handle_slide_exception( 

256 self, 

257 slide_number: int, 

258 ex: Exception 

259 ) -> None: 

260 """Handle exceptions during slide processing.""" 

261 if isinstance(ex, ServiceError): 

262 self.logger.error(f"[{slide_number}] Service error: {ex}") 

263 elif isinstance(ex, (NoRunnableContainerError, NoActiveContainerError, ServiceNotFoundError)): 

264 self.logger.error(f"[{slide_number}] {ex}") 

265 else: 

266 self.logger.error(f"[{slide_number}] Error ({type(ex).__name__}): {ex}") 

267 

268 def _handle_package_not_found( 

269 self, 

270 ex: Exception 

271 ) -> None: 

272 self.logger.error( 

273 f"Package not found for {self.job_id}: {ex}. " 

274 "Is file encrypted or protected by password? " 

275 "Make it public." 

276 ) 

277 

278 async def gen_persona( 

279 self, 

280 pptx_base64: str, 

281 ) -> None: 

282 """ 

283 Generate a video with a character (persona) going over the slides. 

284 """ 

285 async with self.job_status_handler( 

286 extra_handlers={ 

287 PackageNotFoundError: self._handle_package_not_found, 

288 } 

289 ): 

290 if not pptx_base64: 

291 self.logger.error("Document is required.") 

292 await self.save_status(JobStatus.FAILED) 

293 raise ValueError("Missing 'pptx_base64' in request") 

294 self.logger.info(f"Generating persona for slides with {bytes_to_human(len(pptx_base64))}.") 

295 

296 # Save as PDF for debugging 

297 self.logger.info(f"Document base64 with {bytes_to_human(len(pptx_base64))}.") 

298 pptx_path = f"{self.job_path}/document.pptx" 

299 await save_base64_as_binary(pptx_path, pptx_base64) 

300 

301 # Save as text for debugging 

302 txt_path = f"{self.job_path}/document.txt" 

303 async with aiofiles.open(txt_path, "w", encoding="utf-8") as file: 

304 presentation = Presentation(pptx_path) 

305 slide_ix = 0 

306 for slide in presentation.slides: 

307 is_hidden = slide._element.get("show") == "0" 

308 if not is_hidden: 

309 await file.write(f"--- Slide {slide_ix + 1} ---\n") 

310 for shape in slide.shapes: 

311 if shape.has_text_frame: 

312 for paragraph in shape.text_frame.paragraphs: 

313 await file.write(paragraph.text + "\n") 

314 slide_ix += 1 

315 

316 await self.save_status(JobStatus.RUNNING) 

317 

318 # Estimate number of words per slide to fit into video duration 

319 video_duration_seconds = self.get_config_int("video_duration_seconds", 60) 

320 if video_duration_seconds <= 0: 

321 raise ValueError("video_duration_seconds must be positive") 

322 num_slides = get_num_slides(pptx_path) 

323 num_slides = max(1, num_slides) 

324 seconds_per_slide = video_duration_seconds / num_slides 

325 num_words_per_slide = estimate_num_words_from_audio_duration(seconds_per_slide) 

326 self.logger.info( 

327 f"Estimated {num_words_per_slide} words for each of the {num_slides} slides " 

328 f"to fit into a {video_duration_seconds} seconds video.") 

329 

330 persona_size_ratio = self.get_config_float("persona_size_ratio", 0.25) 

331 if not (0.05 <= persona_size_ratio <= 0.5): 

332 raise ValueError("persona_size_ratio must be between 5% and 50%") 

333 

334 await self.save_status(JobStatus.RUNNING) 

335 

336 # Save as images for generating slides video 

337 slide_image_paths = pptx_to_images( 

338 pptx_path, 

339 output_path=self.job_path, 

340 width=self.width, 

341 height=self.height, 

342 logger=self.logger, 

343 ) 

344 if not slide_image_paths: 

345 raise ValueError(f"No images extracted from PPTX {pptx_path}. Is it empty?") 

346 

347 # Persona image sketch parameters 

348 img_prompt = IMG_PROMPT 

349 img_neg_prompt = IMG_NEG_PROMPT 

350 persona_width = int(math.ceil(self.width * persona_size_ratio)) 

351 persona_height = int(math.ceil(self.height * persona_size_ratio)) 

352 

353 await self.save_status(JobStatus.RUNNING) 

354 

355 # Generate transcript for slides 

356 slide_video_tasks = {} 

357 async with aiofiles.open(f"{self.job_path}/slides_transcript.jsonl", "wb") as file: 

358 async for line_json in self.gen.gen_slides_transcript( 

359 pptx_base64=pptx_base64, 

360 max_words_per_slide=num_words_per_slide, 

361 task_id=self.job_id, 

362 ): 

363 line = (json.dumps(line_json) + "\n").encode("utf-8") 

364 await file.write(line) 

365 await file.flush() 

366 

367 line_type = line_json.get("type", "") 

368 if line_type == "persona": 

369 # gender = line_json.get("gender", "unknown") 

370 img_prompt = line_json.get("description", IMG_PROMPT) 

371 

372 # Generate persona image 

373 self.image_task = asyncio.create_task( 

374 self.gen.gen_image( 

375 img_prompt, 

376 neg_prompt=img_neg_prompt, 

377 width=persona_width, 

378 height=persona_height, 

379 steps=25, # TODO steps 

380 task_id="persona_image", 

381 deadline=self.get_submission_time(), 

382 )) 

383 elif line_type == "slide_transcript": 

384 slide_number = line_json.get("slide_number", -1) 

385 slide_text = line_json.get("transcript", "") 

386 MAX_LOG_TEXT_LENGTH = 60 

387 self.logger.info(f"Slide {slide_number}: {slide_text[:MAX_LOG_TEXT_LENGTH]}...") 

388 

389 # Generate slide video 

390 slide_video_task = asyncio.create_task( 

391 self.gen_slide_video( 

392 slide_number, 

393 slide_text, 

394 )) 

395 slide_video_tasks[slide_number] = slide_video_task 

396 else: 

397 self.logger.info(f"Unknown line type: {line_type}") 

398 

399 await self.save_status(JobStatus.RUNNING) 

400 

401 # Collect slide videos 

402 slide_video_paths: List[Optional[str]] = [] 

403 results = await asyncio.gather(*slide_video_tasks.values(), return_exceptions=True) 

404 for slide_number, result in enumerate(results): 

405 if isinstance(result, Exception): 

406 self._handle_slide_exception(slide_number, result) 

407 slide_video_paths.append(None) 

408 elif not result: 

409 self.logger.warning(f"[{slide_number}] No video generated. Skipping...") 

410 slide_video_paths.append(None) 

411 elif isinstance(result, str): 

412 slide_video_paths.append(result) 

413 self.logger.info(f"[{slide_number}] Video+audio saved to '{result}'.") 

414 else: 

415 self.logger.warning(f"[{slide_number}] Wrong result generated {type(result)}...") 

416 slide_video_paths.append(None) 

417 

418 if not slide_video_paths: 

419 raise ValueError("No slide videos generated. Cannot create final slides video.") 

420 

421 await self.save_status(JobStatus.RUNNING) 

422 

423 # Concatenate all slide videos into a final slides video 

424 self.logger.info( 

425 f"Overlaying {len(slide_image_paths)} slides with " 

426 f"{len(slide_video_paths)} persona videos...") 

427 

428 for slide_num, slide_video_path in enumerate(slide_video_paths): 

429 audio_path = f"{self.job_path}/{slide_num + 1:03d}.wav" 

430 if not await aiofiles.os.path.exists(audio_path): 

431 # TODO should we continue with empty audio? 

432 raise FileNotFoundError(f"Audio file for slide {slide_num} not found at '{audio_path}'.") 

433 

434 if slide_num < len(slide_image_paths): 

435 slide_image_path = slide_image_paths[slide_num] 

436 slide_image: Image.Image = Image.open(slide_image_path) 

437 slide_image = slide_image.resize((self.width, self.height)) 

438 else: 

439 self.logger.warning(f"Slide image for slide {slide_num} not found, generating blank slide.") 

440 slide_image = cast(Image.Image, get_frame_with_text( 

441 width=self.width, 

442 height=self.height, 

443 text=f"Slide {slide_num + 1} not available.", 

444 output_type="pil", 

445 background_color="white", 

446 font_color="black", 

447 )) 

448 

449 if slide_video_path: 

450 self.logger.info(f"[{slide_num}] Overlaying persona video on slide.") 

451 

452 slide_video_binary = await read_file_bytes(slide_video_path) 

453 slide_video_frames = await get_video_frames(slide_video_binary) 

454 for frame_index, slide_video_frame in enumerate(slide_video_frames): 

455 slide_video_frames[frame_index] = overlay_image_on_image( 

456 base_image=slide_image, 

457 overlay_image=slide_video_frame, 

458 position=("bottom", "right"), 

459 overlay_percentage=persona_size_ratio) 

460 

461 video_file_info = get_video_file_info(slide_video_binary) 

462 video_info = video_file_info.get("video", {}) 

463 _fps = video_info.get("fps") 

464 fps: float = _fps if _fps is not None else FANTASYTALKING_FPS 

465 else: 

466 fps = FANTASYTALKING_FPS # By default, use FT fps 

467 audio_info = get_audio_file_info(audio_path) 

468 audio_duration = audio_info["duration_seconds"] 

469 num_frames = int(math.ceil(fps * audio_duration)) 

470 if self.image: 

471 self.logger.info( 

472 f"[{slide_num}] No persona video, generating static video with " 

473 f"{audio_duration:.3f} seconds and {num_frames} frames.") 

474 slide_image_persona = overlay_image_on_image( 

475 base_image=slide_image, 

476 overlay_image=self.image, 

477 position=("bottom", "right"), 

478 overlay_percentage=persona_size_ratio) 

479 slide_video_frames = [slide_image_persona] * num_frames 

480 else: 

481 self.logger.info( 

482 f"[{slide_num}] No persona image, generating slide video with " 

483 f"{audio_duration:.3f} seconds and {num_frames} frames.") 

484 slide_video_frames = [slide_image] * num_frames 

485 

486 slide_video_path = f"{self.job_path}/{slide_num + 1:03d}.mp4" 

487 slide_video_path = await save_video_audio( 

488 video_content=slide_video_frames, 

489 audio_path=audio_path, 

490 out_video_path=slide_video_path, 

491 fps=fps) 

492 slide_video_paths[slide_num] = slide_video_path 

493 self.logger.info(f"[{slide_num}] Slide video+audio saved to '{slide_video_path}'.") 

494 

495 non_none_video_paths: List[str] = [p for p in slide_video_paths if p is not None] 

496 video_binary = await concatenate_videos(non_none_video_paths) 

497 video_path = f"{self.job_path}/{self.job_id}.mp4" 

498 async with aiofiles.open(video_path, "wb") as file: 

499 await file.write(video_binary) 

500 

501 def get_slide_deadline( 

502 self, 

503 slide_number: int, 

504 ) -> Optional[float]: 

505 """ 

506 Get the deadline for a slide. 

507 """ 

508 submission_time = self.get_submission_time() 

509 SECONDS_PER_SLIDE = 5.0 # TODO 

510 slide_deadline = submission_time + (slide_number * SECONDS_PER_SLIDE) 

511 return slide_deadline