Improved integration ability
This commit is contained in:
105
main.py
105
main.py
@@ -1,12 +1,10 @@
|
||||
"""Application entry point"""
|
||||
|
||||
import traceback
|
||||
import argparse
|
||||
import logging
|
||||
import json
|
||||
import time
|
||||
import os
|
||||
from dataclasses import asdict
|
||||
import signal
|
||||
from enum import Enum
|
||||
from typing import Callable
|
||||
|
||||
@@ -20,6 +18,7 @@ from structure_builder import Structure, StructureBuilder
|
||||
from structure_refiner import StructureRefiner
|
||||
from markdown_builder import MarkdownBuilder
|
||||
from pdf_builder import PdfBuilder
|
||||
from reporting import StatusReporter
|
||||
from windowizer import Windowizer
|
||||
|
||||
from agent import Agent
|
||||
@@ -27,6 +26,15 @@ from utils import Timeline, ffmpeg_split_video, ffmpeg_to_mp3
|
||||
|
||||
ARGS: argparse.Namespace
|
||||
|
||||
|
||||
class TerminationRequested(Exception):
|
||||
"""Raised when the container receives SIGTERM."""
|
||||
|
||||
def __init__(self, signal_number: int) -> None:
|
||||
super().__init__(f"Received signal {signal_number}")
|
||||
self.signal_number = signal_number
|
||||
|
||||
|
||||
class Step(Enum):
|
||||
MEDIA_SEPARATION = "media_separation"
|
||||
VOICE_RECOGNITION = "voice_recognition"
|
||||
@@ -158,8 +166,17 @@ def setup_arguments() -> argparse.Namespace:
|
||||
type=str,
|
||||
default=ai_api_key
|
||||
)
|
||||
parser.add_argument(
|
||||
"--report-interval",
|
||||
type=float,
|
||||
default=30.0,
|
||||
help="seconds between report.txt heartbeat records",
|
||||
)
|
||||
parser.add_argument("-v", action='store_true')
|
||||
return parser.parse_args()
|
||||
args = parser.parse_args()
|
||||
if args.report_interval <= 0:
|
||||
parser.error("--report-interval must be positive")
|
||||
return args
|
||||
|
||||
#
|
||||
# Workflow
|
||||
@@ -393,23 +410,28 @@ Step.CODE: (
|
||||
)
|
||||
"""
|
||||
|
||||
def main() -> None:
|
||||
"""Application entry point"""
|
||||
global ARGS
|
||||
check_cuda()
|
||||
ARGS = setup_arguments()
|
||||
logging.basicConfig(level=logging.DEBUG if ARGS.v else logging.INFO)
|
||||
|
||||
def run_workflow(reporter: StatusReporter) -> list[str]:
|
||||
"""Execute the workflow and return paths to its final outputs."""
|
||||
step_to_do = Step.MEDIA_SEPARATION
|
||||
intermediate_result: dict | None = None
|
||||
step_numbers = {
|
||||
step: index
|
||||
for index, step in enumerate(WORKFLOW_DATA, start=1)
|
||||
}
|
||||
steps_total = len(WORKFLOW_DATA)
|
||||
|
||||
# execute steps while possible
|
||||
while step_to_do:
|
||||
current_step = step_to_do
|
||||
step_index = step_numbers[current_step]
|
||||
reporter.step_started(current_step.value, step_index, steps_total)
|
||||
step_data = WORKFLOW_DATA[step_to_do]
|
||||
output_file_path = step_data[0]
|
||||
func = step_data[1]
|
||||
if not func:
|
||||
logging.error(f"Step {step_to_do} has no function, stopping")
|
||||
break
|
||||
logging.info(f"Executing step {step_to_do}")
|
||||
raise RuntimeError(f"Step {step_to_do} has no function")
|
||||
logging.info("Executing step %s", step_to_do.value)
|
||||
step_to_do, ret = func(step_to_do, intermediate_result)
|
||||
if ret is not None:
|
||||
with open(output_file_path, "w", encoding="utf-8") as f:
|
||||
@@ -418,14 +440,53 @@ def main() -> None:
|
||||
else:
|
||||
intermediate_result = ret
|
||||
json.dump(ret, f, ensure_ascii=False, indent=4)
|
||||
logging.info(f"Intermediate results are saved {output_file_path}")
|
||||
# final report
|
||||
logging.info(f"Workflow was interrupted at step {step_to_do}")
|
||||
logging.info("Intermediate results are saved to %s", output_file_path)
|
||||
if step_to_do is None and current_step is not Step.PDF_BUILDER:
|
||||
raise RuntimeError(f"Workflow stopped at step {current_step.value}")
|
||||
reporter.step_completed(current_step.value, step_index, steps_total)
|
||||
|
||||
return [path for path in ("output.md", "output.pdf") if os.path.isfile(path)]
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Application entry point."""
|
||||
global ARGS
|
||||
ARGS = setup_arguments()
|
||||
logging.basicConfig(level=logging.DEBUG if ARGS.v else logging.INFO)
|
||||
reporter = StatusReporter(heartbeat_interval=ARGS.report_interval)
|
||||
reporter.start()
|
||||
|
||||
previous_sigterm_handler = signal.getsignal(signal.SIGTERM)
|
||||
|
||||
def handle_sigterm(signal_number: int, _frame: object) -> None:
|
||||
raise TerminationRequested(signal_number)
|
||||
|
||||
signal.signal(signal.SIGTERM, handle_sigterm)
|
||||
try:
|
||||
check_cuda()
|
||||
outputs = run_workflow(reporter)
|
||||
except TerminationRequested as error:
|
||||
logging.warning("Workflow was cancelled by SIGTERM")
|
||||
reporter.finish("cancelled", signal=error.signal_number)
|
||||
raise SystemExit(128 + error.signal_number) from None
|
||||
except KeyboardInterrupt:
|
||||
logging.warning("Workflow was cancelled by the user")
|
||||
reporter.finish("cancelled")
|
||||
raise
|
||||
except BaseException as error:
|
||||
logging.error("Workflow failed: %s: %s", type(error).__name__, error)
|
||||
reporter.finish(
|
||||
"failed",
|
||||
error=type(error).__name__,
|
||||
message=str(error),
|
||||
)
|
||||
raise
|
||||
else:
|
||||
reporter.finish("succeeded", outputs=outputs)
|
||||
logging.info("Workflow completed successfully")
|
||||
finally:
|
||||
signal.signal(signal.SIGTERM, previous_sigterm_handler)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
main()
|
||||
except SystemExit:
|
||||
raise
|
||||
except:
|
||||
traceback.print_exc()
|
||||
main()
|
||||
|
||||
Reference in New Issue
Block a user