493 lines
17 KiB
Python
493 lines
17 KiB
Python
"""Application entry point"""
|
|
|
|
import argparse
|
|
import logging
|
|
import json
|
|
import os
|
|
import signal
|
|
from enum import Enum
|
|
from typing import Callable
|
|
|
|
import torch
|
|
from asr import Asr, AsrRawResult
|
|
from asr_filter import AsrFilter, AsrFilterResult
|
|
from asr_eventizer import AsrEventizer
|
|
from video_references import UnresolvedReferences, VideoReferenceBuilder
|
|
from reference_resolver import ReferenceResolver
|
|
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
|
|
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"
|
|
ASR_FILTER = "asr_filter"
|
|
ASR_EVENTS = "asr_events"
|
|
VIDEO_REFERENCES = "video_references"
|
|
REFERENCE_RESOLVER = "reference_resolver"
|
|
STRUCTURE_BUILDER = "structure_builder"
|
|
STRUCTURE_REFINER = "structure_refiner"
|
|
MARKDOWN_BUILDER = "markdown_builder"
|
|
PDF_BUILDER = "pdf_builder"
|
|
|
|
#
|
|
# Utility
|
|
#
|
|
def check_cuda() -> None:
|
|
"""Checks if CUDA is available."""
|
|
if not torch.cuda.is_available():
|
|
raise RuntimeError(
|
|
"CUDA is unavaiable! Refusing to start, it's pointless."
|
|
)
|
|
|
|
def setup_arguments() -> argparse.Namespace:
|
|
"""Parses CLI arguments and returns namespace. It may terminate the
|
|
application on invalid arguments.
|
|
|
|
Returns:
|
|
- argparse namespace
|
|
"""
|
|
voice_models = Asr.get_models_list()
|
|
ai_api_key = os.environ.get("AI_API_KEY", default=...)
|
|
|
|
parser = argparse.ArgumentParser(
|
|
prog="sumka",
|
|
description="Summarizes large video/audio files into convenient format",
|
|
)
|
|
parser.add_argument(
|
|
"--asr-model",
|
|
choices=voice_models,
|
|
default="turbo" if "turbo" in voice_models else voice_models[-1]
|
|
)
|
|
parser.add_argument(
|
|
"--asr-language",
|
|
choices=["ru", "en"],
|
|
default="ru"
|
|
)
|
|
parser.add_argument(
|
|
"--asr-ai-model",
|
|
type=str,
|
|
default="google/gemini-3.1-flash-lite"
|
|
)
|
|
parser.add_argument(
|
|
"--asr-ai-base-url",
|
|
type=str,
|
|
default="https://api.proxyapi.ru/v1"
|
|
)
|
|
parser.add_argument(
|
|
"--asr-ai-api-key",
|
|
type=str,
|
|
default=ai_api_key
|
|
)
|
|
parser.add_argument(
|
|
"--video-ref-ai-model",
|
|
type=str,
|
|
default="google/gemini-3.1-flash-lite"
|
|
)
|
|
parser.add_argument(
|
|
"--video-ref-ai-base-url",
|
|
type=str,
|
|
default="https://api.proxyapi.ru/v1"
|
|
)
|
|
parser.add_argument(
|
|
"--video-ref-ai-api-key",
|
|
type=str,
|
|
default=ai_api_key
|
|
)
|
|
parser.add_argument(
|
|
"--resolver-ai-model",
|
|
type=str,
|
|
default="google/gemini-3.5-flash"
|
|
)
|
|
parser.add_argument(
|
|
"--resolver-ai-base-url",
|
|
type=str,
|
|
default="https://api.proxyapi.ru/v1"
|
|
)
|
|
parser.add_argument(
|
|
"--resolver-ai-api-key",
|
|
type=str,
|
|
default=ai_api_key
|
|
)
|
|
parser.add_argument(
|
|
"--resolver-frame-interval",
|
|
type=float,
|
|
default=1.0
|
|
)
|
|
parser.add_argument(
|
|
"--resolver-max-frames",
|
|
type=int,
|
|
default=9
|
|
)
|
|
parser.add_argument(
|
|
"--structure-ai-model",
|
|
type=str,
|
|
default="google/gemini-3.5-flash"
|
|
)
|
|
parser.add_argument(
|
|
"--structure-ai-base-url",
|
|
type=str,
|
|
default="https://api.proxyapi.ru/v1"
|
|
)
|
|
parser.add_argument(
|
|
"--structure-ai-api-key",
|
|
type=str,
|
|
default=ai_api_key
|
|
)
|
|
parser.add_argument(
|
|
"--refiner-ai-model",
|
|
type=str,
|
|
default="google/gemini-3.5-flash"
|
|
)
|
|
parser.add_argument(
|
|
"--refiner-ai-base-url",
|
|
type=str,
|
|
default="https://api.proxyapi.ru/v1"
|
|
)
|
|
parser.add_argument(
|
|
"--refiner-ai-api-key",
|
|
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')
|
|
args = parser.parse_args()
|
|
if args.report_interval <= 0:
|
|
parser.error("--report-interval must be positive")
|
|
return args
|
|
|
|
#
|
|
# Workflow
|
|
#
|
|
def on_media_separation(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
|
# there may be no data for this step
|
|
if input_data:
|
|
raise RuntimeError("There must be no input data for Media Separation")
|
|
# do not split if there's `audio.mp3`
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping media separation")
|
|
return (Step.VOICE_RECOGNITION, None)
|
|
|
|
# filesnames to look for
|
|
VIDEO_INPUTS = [
|
|
"input.mp4",
|
|
"input.mkv",
|
|
"input.avi"
|
|
]
|
|
AUDIO_INPUTS = [
|
|
"input.mp3",
|
|
"input.m4a",
|
|
"input.wav"
|
|
]
|
|
# execute video split
|
|
for v in VIDEO_INPUTS:
|
|
if os.path.isfile(v):
|
|
logging.info(f"Splitting {v} to audio.mp3 and video.mp4")
|
|
ffmpeg_split_video(v)
|
|
return (Step.VOICE_RECOGNITION, None)
|
|
# execute audio conversion
|
|
for a in AUDIO_INPUTS:
|
|
if os.path.isfile(a):
|
|
logging.info(f"Converting {a} to audio.mp3")
|
|
ffmpeg_to_mp3(a)
|
|
return (Step.VOICE_RECOGNITION, None)
|
|
logging.error(f"No supported `input.*` files found")
|
|
return (None, None)
|
|
|
|
def on_voice_recognition(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
|
# do not perform recognition if output file exists
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping voice recognition")
|
|
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
|
return (Step.ASR_FILTER, json.load(f))
|
|
logging.info(f"Loading ASR model `{ARGS.asr_model}`, language `{ARGS.asr_language}`")
|
|
asr = Asr(ARGS.asr_model)
|
|
logging.info(f"Speech recognition...")
|
|
result = asr.recognize("audio.mp3", language=ARGS.asr_language)
|
|
logging.info(f"Speech recognition done")
|
|
return (Step.ASR_FILTER, result.model_dump(mode="json"))
|
|
|
|
def on_asr_filter(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
|
# do not filter if output file exists
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping ASR filter")
|
|
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
|
return (Step.ASR_EVENTS, json.load(f))
|
|
# bad request
|
|
if input_data is None:
|
|
logging.error("Can't filter raw ASR output without input_data")
|
|
return (None, None)
|
|
logging.info("Filtering raw ASR output...")
|
|
filter = AsrFilter()
|
|
result = filter.filter(AsrRawResult(**input_data))
|
|
return (Step.ASR_EVENTS, result.model_dump(mode="json"))
|
|
|
|
def on_asr_events(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
|
# do not filter if output file exists
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping ASR eventizing")
|
|
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
|
return (Step.VIDEO_REFERENCES, json.load(f))
|
|
# bad request
|
|
if input_data is None:
|
|
logging.error("Can't create audio events without input_data")
|
|
return (None, None)
|
|
logging.info("Creating the agent")
|
|
agent = Agent(
|
|
model=ARGS.asr_ai_model,
|
|
base_url=ARGS.asr_ai_base_url,
|
|
api_key=ARGS.asr_ai_api_key
|
|
)
|
|
logging.info("Creating the windowizer")
|
|
windowizer = Windowizer()
|
|
logging.info("Creating audio events...")
|
|
eventizer = AsrEventizer(agent, windowizer)
|
|
result = eventizer.eventize(AsrFilterResult(**input_data))
|
|
return (Step.VIDEO_REFERENCES, result.model_dump(mode="json"))
|
|
|
|
def on_video_references(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
|
# do not filter if output file exists
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping video references")
|
|
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
|
return (Step.REFERENCE_RESOLVER, json.load(f))
|
|
if not os.path.isfile("video.mp4"):
|
|
logging.info("No video, skipping")
|
|
return (Step.REFERENCE_RESOLVER, None)
|
|
if input_data is None:
|
|
logging.error("Can't build video references without input_data")
|
|
return (None, None)
|
|
logging.info("Creating the agent")
|
|
agent = Agent(
|
|
model=ARGS.video_ref_ai_model,
|
|
base_url=ARGS.video_ref_ai_base_url,
|
|
api_key=ARGS.video_ref_ai_api_key
|
|
)
|
|
logging.info("Building video references...")
|
|
builder = VideoReferenceBuilder(agent, Windowizer())
|
|
result = builder.build(Timeline(**input_data))
|
|
return (Step.REFERENCE_RESOLVER, result.model_dump(mode="json"))
|
|
|
|
def on_reference_resolver(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
|
# do not filter if output file exists
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping reference resolver")
|
|
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
|
return (Step.STRUCTURE_BUILDER, json.load(f))
|
|
if not os.path.isfile("video.mp4"):
|
|
logging.info("No video, skipping")
|
|
with open("audio_events.json", "rb") as f:
|
|
return (Step.STRUCTURE_BUILDER, json.load(f))
|
|
if input_data is None:
|
|
logging.error("Can't resolve video references without input_data")
|
|
return (None, None)
|
|
with open("audio_events.json", "rb") as f:
|
|
timeline = Timeline(**json.load(f))
|
|
logging.info("Creating the agent")
|
|
agent = Agent(
|
|
model=ARGS.resolver_ai_model,
|
|
base_url=ARGS.resolver_ai_base_url,
|
|
api_key=ARGS.resolver_ai_api_key
|
|
)
|
|
logging.info("Resolving video references...")
|
|
resolver = ReferenceResolver(
|
|
agent,
|
|
frame_interval=ARGS.resolver_frame_interval,
|
|
max_frames=ARGS.resolver_max_frames,
|
|
)
|
|
result = resolver.resolve(timeline, UnresolvedReferences(**input_data))
|
|
return (Step.STRUCTURE_BUILDER, result.model_dump(mode="json"))
|
|
|
|
def on_structure_builder(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
|
# don't if done
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping structure builder")
|
|
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
|
return (Step.STRUCTURE_REFINER, json.load(f))
|
|
if input_data is None:
|
|
logging.error("Can't build document structure without input_data")
|
|
return (None, None)
|
|
logging.info("Creating the agent")
|
|
agent = Agent(
|
|
model=ARGS.structure_ai_model,
|
|
base_url=ARGS.structure_ai_base_url,
|
|
api_key=ARGS.structure_ai_api_key
|
|
)
|
|
logging.info("Creating the windowizer")
|
|
windowizer = Windowizer()
|
|
logging.info("Building document structure...")
|
|
builder = StructureBuilder(agent, windowizer)
|
|
result = builder.build(Timeline(**input_data))
|
|
return (Step.STRUCTURE_REFINER, result.model_dump(mode="json"))
|
|
|
|
def on_structure_refiner(current_step: Step, input_data: dict | None) -> tuple[Step | None, dict | None]:
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping structure refiner")
|
|
with open(WORKFLOW_DATA[current_step][0], "rb") as f:
|
|
return (Step.MARKDOWN_BUILDER, json.load(f))
|
|
if input_data is None:
|
|
logging.error("Can't refine document structure without input_data")
|
|
return (None, None)
|
|
logging.info("Creating the agent")
|
|
agent = Agent(
|
|
model=ARGS.refiner_ai_model,
|
|
base_url=ARGS.refiner_ai_base_url,
|
|
api_key=ARGS.refiner_ai_api_key
|
|
)
|
|
logging.info("Refining document structure...")
|
|
refiner = StructureRefiner(agent, Windowizer())
|
|
result = refiner.refine(Structure(**input_data))
|
|
return (Step.MARKDOWN_BUILDER, result.model_dump(mode="json"))
|
|
|
|
def on_markdown_builder(current_step: Step, input_data: dict | None) -> tuple[Step | None, str | None]:
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping Markdown builder")
|
|
return (Step.PDF_BUILDER, None)
|
|
if input_data is None:
|
|
logging.error("Can't build Markdown without input_data")
|
|
return (None, None)
|
|
with open("events.json", "rb") as f:
|
|
timeline = Timeline(**json.load(f))
|
|
logging.info("Building Markdown...")
|
|
builder = MarkdownBuilder(timeline)
|
|
return (Step.PDF_BUILDER, builder.build(Structure(**input_data)))
|
|
|
|
def on_pdf_builder(current_step: Step, input_data: dict | None) -> tuple[None, None]:
|
|
if os.path.isfile(WORKFLOW_DATA[current_step][0]):
|
|
logging.info("Skipping PDF builder")
|
|
return (None, None)
|
|
if not os.path.isfile("output.md"):
|
|
logging.error("Can't build PDF without output.md")
|
|
return (None, None)
|
|
logging.info("Building PDF...")
|
|
PdfBuilder().build("output.md", WORKFLOW_DATA[current_step][0])
|
|
logging.info("PDF is saved to %s", WORKFLOW_DATA[current_step][0])
|
|
return (None, None)
|
|
|
|
#
|
|
# Main
|
|
#
|
|
WORKFLOW_DATA: dict[Step, tuple[str, Callable[[Step, dict | None], tuple[Step | None, dict | str | None]] | None]] = {
|
|
Step.MEDIA_SEPARATION: ("audio.mp3", on_media_separation),
|
|
Step.VOICE_RECOGNITION: ("asr_raw.json", on_voice_recognition),
|
|
Step.ASR_FILTER: ("asr.json", on_asr_filter),
|
|
Step.ASR_EVENTS: ("audio_events.json", on_asr_events),
|
|
Step.VIDEO_REFERENCES: ("unresolved.json", on_video_references),
|
|
Step.REFERENCE_RESOLVER: ("events.json", on_reference_resolver),
|
|
Step.STRUCTURE_BUILDER: ("structure.json", on_structure_builder),
|
|
Step.STRUCTURE_REFINER: ("structure_refined.json", on_structure_refiner),
|
|
Step.MARKDOWN_BUILDER: ("output.md", on_markdown_builder),
|
|
Step.PDF_BUILDER: ("output.pdf", on_pdf_builder),
|
|
}
|
|
"""Information about workflow.
|
|
|
|
Step.CODE: (
|
|
"path/to/result.json",
|
|
(cur_step: Step, input_data: dict | None)
|
|
-> (next_step: Step | None, output_data: dict | None)
|
|
)
|
|
"""
|
|
|
|
|
|
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:
|
|
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:
|
|
if isinstance(ret, str):
|
|
f.write(ret)
|
|
else:
|
|
intermediate_result = ret
|
|
json.dump(ret, f, ensure_ascii=False, indent=4)
|
|
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__":
|
|
main()
|