328 lines
10 KiB
Python
Executable File
328 lines
10 KiB
Python
Executable File
#!/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()
|