use rskit_errors::{AppError, AppResult, ErrorCode};
use serde::{Deserialize, Serialize};
use crate::time::TimeRange;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SubtitleEntry {
pub range: TimeRange,
pub text: String,
pub style: Option<SubtitleStyle>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SubtitleStyle {
pub font_family: Option<String>,
pub font_size: Option<u16>,
pub color: Option<String>,
pub background: Option<String>,
pub bold: bool,
pub italic: bool,
pub position: SubtitlePosition,
}
#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize)]
pub enum SubtitlePosition {
#[default]
Bottom,
Top,
Center,
Custom {
x: u32,
y: u32,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SubtitleTrack {
pub entries: Vec<SubtitleEntry>,
pub language: Option<String>,
pub default_style: Option<SubtitleStyle>,
}
impl SubtitleTrack {
pub fn new() -> Self {
Self {
entries: Vec::new(),
language: None,
default_style: None,
}
}
#[must_use]
pub fn add(mut self, range: TimeRange, text: impl Into<String>) -> Self {
self.entries.push(SubtitleEntry {
range,
text: text.into(),
style: None,
});
self
}
#[must_use]
pub fn with_language(mut self, lang: impl Into<String>) -> Self {
self.language = Some(lang.into());
self
}
pub fn from_srt(content: &str) -> AppResult<Self> {
let mut entries = Vec::new();
let content = content
.strip_prefix('\u{feff}')
.unwrap_or(content)
.replace("\r\n", "\n");
let blocks: Vec<&str> = content
.split("\n\n")
.filter(|b| !b.trim().is_empty())
.collect();
for block in blocks {
let lines: Vec<&str> = block.trim().lines().collect();
if lines.is_empty() {
continue;
}
let time_idx = lines.iter().position(|l| l.contains(" --> "));
let Some(time_idx) = time_idx else {
continue;
};
let time_line = lines[time_idx];
let parts: Vec<&str> = time_line.split(" --> ").collect();
if parts.len() != 2 {
continue;
}
let start = parse_srt_time(parts[0].trim()).ok_or_else(|| {
AppError::new(
ErrorCode::InvalidFormat,
format!("invalid SRT time: {}", parts[0]),
)
})?;
let end = parse_srt_time(parts[1].trim()).ok_or_else(|| {
AppError::new(
ErrorCode::InvalidFormat,
format!("invalid SRT time: {}", parts[1]),
)
})?;
let text_lines = &lines[time_idx + 1..];
if text_lines.is_empty() {
continue;
}
let text = strip_html_tags(&text_lines.join("\n"));
entries.push(SubtitleEntry {
range: TimeRange::from_millis(start, end),
text,
style: None,
});
}
Ok(Self {
entries,
language: None,
default_style: None,
})
}
pub fn from_vtt(content: &str) -> AppResult<Self> {
let content = content
.strip_prefix('\u{feff}')
.unwrap_or(content)
.replace("\r\n", "\n");
let content = content
.strip_prefix("WEBVTT")
.unwrap_or(&content)
.trim_start();
let mut entries = Vec::new();
let blocks: Vec<&str> = content
.split("\n\n")
.filter(|b| !b.trim().is_empty())
.collect();
for block in blocks {
let lines: Vec<&str> = block.trim().lines().collect();
if lines.is_empty() {
continue;
}
let time_idx = lines.iter().position(|l| l.contains(" --> "));
let Some(time_idx) = time_idx else {
continue;
};
let time_line = lines[time_idx];
let parts: Vec<&str> = time_line.split(" --> ").collect();
if parts.len() != 2 {
continue;
}
let start_str = parts[0].trim();
let end_str = parts[1].split_whitespace().next().unwrap_or("");
let start = parse_vtt_time(start_str).ok_or_else(|| {
AppError::new(
ErrorCode::InvalidFormat,
format!("invalid VTT time: {start_str}"),
)
})?;
let end = parse_vtt_time(end_str).ok_or_else(|| {
AppError::new(
ErrorCode::InvalidFormat,
format!("invalid VTT time: {end_str}"),
)
})?;
let text_lines = &lines[time_idx + 1..];
if text_lines.is_empty() {
continue;
}
let raw_text = text_lines.join("\n");
let text = decode_html_entities(&strip_html_tags(&raw_text));
entries.push(SubtitleEntry {
range: TimeRange::from_millis(start, end),
text,
style: None,
});
}
Ok(Self {
entries,
language: None,
default_style: None,
})
}
pub fn to_srt(&self) -> String {
let mut out = String::new();
for (i, entry) in self.entries.iter().enumerate() {
out.push_str(&format!("{}\n", i + 1));
out.push_str(&format!(
"{} --> {}\n",
format_srt_time(entry.range.start.as_millis()),
format_srt_time(entry.range.end.as_millis()),
));
out.push_str(&entry.text);
out.push_str("\n\n");
}
out
}
pub fn to_vtt(&self) -> String {
let mut out = String::from("WEBVTT\n\n");
for entry in &self.entries {
out.push_str(&format!(
"{} --> {}\n",
format_vtt_time(entry.range.start.as_millis()),
format_vtt_time(entry.range.end.as_millis()),
));
out.push_str(&entry.text);
out.push_str("\n\n");
}
out
}
pub fn shift(&mut self, offset: i64) {
for entry in &mut self.entries {
entry.range = entry.range.shift(offset);
}
}
pub fn in_range(&self, range: &TimeRange) -> Self {
Self {
entries: self
.entries
.iter()
.filter(|e| e.range.overlaps(range))
.cloned()
.collect(),
language: self.language.clone(),
default_style: self.default_style.clone(),
}
}
}
impl Default for SubtitleTrack {
fn default() -> Self {
Self::new()
}
}
fn parse_srt_time(s: &str) -> Option<u64> {
let s = s.replace(',', ".");
parse_time_dotted(&s)
}
fn format_srt_time(ms: u64) -> String {
let millis = ms % 1000;
let total_secs = ms / 1000;
let secs = total_secs % 60;
let total_mins = total_secs / 60;
let mins = total_mins % 60;
let hours = total_mins / 60;
format!("{hours:02}:{mins:02}:{secs:02},{millis:03}")
}
fn parse_vtt_time(s: &str) -> Option<u64> {
parse_time_dotted(s)
}
fn format_vtt_time(ms: u64) -> String {
let millis = ms % 1000;
let total_secs = ms / 1000;
let secs = total_secs % 60;
let total_mins = total_secs / 60;
let mins = total_mins % 60;
let hours = total_mins / 60;
format!("{hours:02}:{mins:02}:{secs:02}.{millis:03}")
}
fn parse_time_dotted(s: &str) -> Option<u64> {
let (main, frac) = if let Some((m, f)) = s.split_once('.') {
(m, f.parse::<u64>().ok()?)
} else {
(s, 0)
};
let parts: Vec<&str> = main.split(':').collect();
let (h, m, sec) = match parts.len() {
3 => (
parts[0].parse::<u64>().ok()?,
parts[1].parse::<u64>().ok()?,
parts[2].parse::<u64>().ok()?,
),
2 => (
0,
parts[0].parse::<u64>().ok()?,
parts[1].parse::<u64>().ok()?,
),
_ => return None,
};
Some(h * 3_600_000 + m * 60_000 + sec * 1000 + frac)
}
fn strip_html_tags(s: &str) -> String {
let mut result = String::with_capacity(s.len());
let mut in_tag = false;
for ch in s.chars() {
match ch {
'<' => in_tag = true,
'>' => in_tag = false,
_ if !in_tag => result.push(ch),
_ => {}
}
}
result
}
fn decode_html_entities(s: &str) -> String {
s.replace("&", "&")
.replace("<", "<")
.replace(">", ">")
.replace(""", "\"")
.replace("'", "'")
.replace("'", "'")
.replace(" ", " ")
.replace("​", "")
}
#[cfg(test)]
mod tests {
use crate::time::TimeRange;
use super::*;
#[test]
fn malformed_srt_and_vtt_blocks_are_skipped_or_rejected() {
let srt =
"\u{feff}not a number\nno timestamp\n\n3\n00:00:01,000 --> 00:00:02,000\n<b>ok</b>";
assert!(SubtitleTrack::from_srt("1\nbad --> 00:00:02,000\nbad").is_err());
assert!(SubtitleTrack::from_srt("1\n00:00:01,000 --> bad\nbad").is_err());
let track = SubtitleTrack::from_srt(srt).unwrap();
assert_eq!(track.entries.len(), 1);
assert_eq!(track.entries[0].text, "ok");
assert!(SubtitleTrack::from_vtt("WEBVTT\n\nbad --> 00:00:02.000\nbad").is_err());
assert!(SubtitleTrack::from_vtt("WEBVTT\n\n00:00:01.000 --> bad\nbad").is_err());
let vtt = "WEBVTT\n\nNOTE no timestamp\n\ncue\n00:00:01.000 --> 00:00:02.000 align:start\n<c>& hi</c>";
let track = SubtitleTrack::from_vtt(vtt).unwrap();
assert_eq!(track.entries.len(), 1);
assert_eq!(track.entries[0].text, "& hi");
}
#[test]
fn subtitle_defaults_formatters_and_helpers_cover_edge_cases() {
let mut track = SubtitleTrack::default()
.with_language("en")
.add(TimeRange::from_millis(1_000, 2_500), "hello");
assert_eq!(track.language.as_deref(), Some("en"));
assert!(track.default_style.is_none());
assert!(track.to_srt().contains("00:00:01,000 --> 00:00:02,500"));
assert!(track.to_vtt().contains("00:00:01.000 --> 00:00:02.500"));
track.shift(-500);
assert_eq!(track.entries[0].range.start.as_millis(), 500);
assert_eq!(
track
.in_range(&TimeRange::from_millis(0, 600))
.entries
.len(),
1
);
assert_eq!(parse_srt_time("00:00:01,250"), Some(1_250));
assert_eq!(parse_vtt_time("01:02.003"), Some(62_003));
assert_eq!(parse_time_dotted("bad"), None);
assert_eq!(strip_html_tags("<b>a</b><i>b</i>"), "ab");
}
}