diff --git a/subtitle_translator.py b/subtitle_translator.py new file mode 100755 index 0000000..ac8f83e --- /dev/null +++ b/subtitle_translator.py @@ -0,0 +1,327 @@ +#!/usr/bin/env -S uv run --script +# /// script +# dependencies = ["srt", "httpx"] +# /// + +import argparse +import json +import logging +import os +from pathlib import Path +import time +from typing import Iterator +import srt +import httpx + +api_key = os.getenv("OPENROUTER_API_KEY") +if not api_key: + raise RuntimeError("OPENROUTER_API_KEY environment variable not set") +client = httpx.Client( + headers={"Authorization": f"Bearer {api_key}"}, + base_url="https://openrouter.ai/api/v1", + timeout=120, +) + + +# retry is a decorator to retry the function on failure +def retry(max_retries: int = 3, delay: float = 1.0): + def decorator(func): + def wrapper(*args, **kwargs): + for attempt in range(max_retries): + try: + return func(*args, **kwargs) + except Exception as e: + logging.debug(f"Error: {e}, retrying {attempt + 1}/{max_retries}...") + if attempt < max_retries - 1: + time.sleep(delay) + raise RuntimeError("Max retries exceeded") + + return wrapper + + return decorator + + +def combine_entries(entries: list[srt.Subtitle], max_entries: float = 2) -> list[srt.Subtitle]: + """ + Merge adjacent subtitles in groups of 'factor' if: + - They are no more than 1s apart (start of next - end of prev <= 1s) + - The total word count is fewer than 8 + """ + if not entries: + return [] + + merged = [] + i = 0 + idx = 1 + n = len(entries) + while i < n: + group = [entries[i]] + j = 1 + while j < int(max_entries) and (i + j) < n: + prev = group[-1] + curr = entries[i + j] + gap = (curr.start - prev.end).total_seconds() + total_words = sum(len(s.content.split()) for s in group) + len(curr.content.split()) + if gap <= 1 and total_words < 8: + group.append(curr) + j += 1 + else: + break + start_time = group[0].start + end_time = group[-1].end + combined_text = " ".join(x.content.replace("\n", " ") for x in group) + words = combined_text.split() + formatted_text = format_subtitle_text(words) + merged.append(srt.Subtitle(index=idx, start=start_time, end=end_time, content=formatted_text, proprietary="")) + idx += 1 + i += len(group) + return merged + + +def format_subtitle_text(words: list[str], max_words_per_line: int = 10) -> str: + """ + Format a list of words into subtitle text with appropriate line breaks. + + Args: + words: List of words to format + max_words_per_line: Maximum words allowed per line + + Returns: + Formatted subtitle text with line breaks + """ + if not words: + return "" + + if len(words) <= max_words_per_line: + return " ".join(words) + + lines = [] + i = 0 + + while i < len(words): + remaining = len(words) - i + + # Handle orphan words (single word on last line) + if remaining <= 1 and lines: + lines[-1] += " " + words[i] + break + + # Calculate words for this line + words_for_line = min(max_words_per_line, remaining) + lines.append(" ".join(words[i : i + words_for_line])) + i += words_for_line + + return "\n".join(lines) + + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter) + p.add_argument("input", type=Path, help="Input SRT file") + p.add_argument("output", type=Path, help="Output SRT file") + p.add_argument("--translate", type=str, help="Language code for translation") + p.add_argument( + "--condense", + type=int, + help="Merge adjacent short subtitles in groups of N if close together", + ) + return p.parse_args() + + +def clean_entries(subs: list[srt.Subtitle]) -> list[srt.Subtitle]: + """ + Remove all occurrences of '\\h' from subtitle content. + """ + for sub in subs: + content = sub.content.replace(r"\h", " ") + content = content.replace(r"\n", "\n") + content = "\n".join(line.strip() for line in content.splitlines() if line.strip()) + sub.content = content + return subs + + +def remove_overlaps(subs: list[srt.Subtitle]) -> list[srt.Subtitle]: + """ + Remove duplicate lines between consecutive subtitle entries. + + If the beginning of a subtitle entry matches the end of the previous entry, + those duplicate lines are removed from the current subtitle. + + Args: + subs: List of subtitle entries + + Returns: + List of cleaned subtitle entries with duplicates removed + """ + if not subs: + return [] + + cleaned = [] + prev_lines = [] + + for sub in subs: + # Split content into non-empty lines + lines = [line.strip() for line in sub.content.splitlines()] + lines = [line for line in lines if line] + + if not lines: + continue + + # Check for overlap with previous subtitle + overlap = 0 + max_overlap = min(len(prev_lines), len(lines)) + for i in range(max_overlap, 0, -1): + if prev_lines[-i:] == lines[:i]: + overlap = i + break + + # Keep only non-overlapping lines + non_dup_lines = lines[overlap:] + + if non_dup_lines: + # Create new subtitle with cleaned content + new_sub = srt.Subtitle(index=sub.index, start=sub.start, end=sub.end, content="\n".join(non_dup_lines), proprietary=sub.proprietary) + cleaned.append(new_sub) + prev_lines = lines + + # Reindex subtitles + for i, sub in enumerate(cleaned, 1): + sub.index = i + + return cleaned + + +def chunkify[T](it: list[T], size: int) -> Iterator[list[T]]: + """Yield successive n-sized chunks from it.""" + for i in range(0, len(it), size): + yield it[i : i + size] + + +TOOL_SCHEMA_SAVE_SUBTITLES = { + "type": "function", + "function": { + "name": "save_subtitles", + "description": "Save translated subtitle entries", + "parameters": { + "type": "object", + "properties": { + "entries": { + "type": "array", + "items": { + "type": "object", + "properties": {"id": {"type": "integer"}, "translated": {"type": "string"}}, + "required": ["id", "translated"], + }, + } + }, + "required": ["entries"], + }, + }, +} + + +def translate_entries(entries: list[srt.Subtitle], lang: str) -> list[srt.Subtitle]: + logging.info(f"Translating subtitles to {lang} via OpenRouter") + + @retry(max_retries=5) + def translate_chunk(chunk: list[srt.Subtitle]) -> list[srt.Subtitle]: + prompt = f""" +You are an expert subtitle translator. +Translate the following subtitle entries from English to {lang}. +For each entry, only translate the 'content' field. Do not change the 'id'. +Do NOT translate proper nouns, names, or technical terms; keep them in English. +Do NOT add, remove, or merge entries. Do not include any explanations or comments. + +Return exactly the same number of entries as provided. For each input entry, there must be one output entry with the same 'id'. +If an entry is empty, return it as empty. + +Use your save_subtitles tool to return the results. +""" + chunk_json = json.dumps( + [ + { + "id": s.index, + "content": s.content.strip(), + } + for s in chunk + ], + ensure_ascii=False, + ) + payload = { + "model": "google/gemini-2.5-flash-lite", + "messages": [ + {"role": "system", "content": prompt}, + {"role": "user", "content": chunk_json}, + ], + "tools": [TOOL_SCHEMA_SAVE_SUBTITLES], + "tool_choice": {"type": "function", "function": {"name": TOOL_SCHEMA_SAVE_SUBTITLES["function"]["name"]}}, + } + res = client.post(url="/chat/completions", json=payload) + if res.is_error: + logging.error(f"Error translating subtitles: {res.text}") + res.raise_for_status() + + data = res.json() + + tool_calls = data["choices"][0]["message"]["tool_calls"] + + for call in tool_calls: + if call["function"]["name"] != TOOL_SCHEMA_SAVE_SUBTITLES["function"]["name"]: + continue + parsed = json.loads(call["function"]["arguments"]) + translated_entries = parsed["entries"] + if len(translated_entries) != len(chunk): + raise ValueError(f"Expected {len(chunk)} translated entries, got {len(translated_entries)}") + # Reassemble subtitles with translated content + id_to_sub = {s.index: s for s in chunk} + result = [] + for entry in translated_entries: + orig = id_to_sub.get(entry["id"]) + if orig: + orig.content = entry["translated"].strip() + result.append(orig) + return result + + combined = [] + chunks = chunkify(entries, 20) + for i, chunk in enumerate(chunks, start=1): + logging.info(f"Translating chunk {i}/{len(chunks)}") + translated = translate_chunk(chunk) + combined.extend(translated) + # Reindex + for i, sub in enumerate(combined, 1): + sub.index = i + return combined + + +def translate(subtitle_path: Path, lang: str, save_path: Path, condense: int | None = None) -> None: + """ + Translate the subtitles in the given SRT file to the specified language. + Returns a new SRT file with translated subtitles. + """ + entries = list(srt.parse(subtitle_path.read_text())) + entries = clean_entries(entries) + entries = remove_overlaps(entries) + if condense: + entries = combine_entries(entries, max_entries=condense) + + if not entries: + raise ValueError("No valid subtitle entries found to translate") + + translated = translate_entries(entries, lang=lang) + save_path.write_text(srt.compose(translated)) + + +def main(): + logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s: %(message)s") + args = parse_args() + + translate( + subtitle_path=args.input, + lang=args.translate, + save_path=args.output, + condense=args.condense, + ) + + +if __name__ == "__main__": + main()