use docling_core::tree::{Formatting, ItemTree, TreeKind, TreeTrack};
use docling_core::{DoclingDocument, Node};
use crate::backend::markdown::escape_text;
use crate::backend::DeclarativeBackend;
use crate::error::ConversionError;
use crate::source::SourceDocument;
pub struct WebVttBackend;
#[derive(Default, Clone)]
struct Meta {
voice: Option<String>,
formatting: Option<Formatting>,
}
struct Run {
text: String,
meta: Meta,
}
enum Comp {
Text {
text: String,
terminator: bool,
},
Span {
tag: Tag,
annotation: Option<String>,
children: Vec<Comp>,
},
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Tag {
Bold,
Italic,
Underline,
Voice,
Transparent,
}
const MAX_SPAN_DEPTH: usize = 100;
struct Cue {
start: f64,
end: f64,
identifier: Option<String>,
payload: Vec<Comp>,
}
impl DeclarativeBackend for WebVttBackend {
fn convert(&self, source: &SourceDocument) -> Result<DoclingDocument, ConversionError> {
let content = source.text()?.replace("\r\n", "\n").replace('\r', "\n");
let mut doc = DoclingDocument::new(&source.name);
let mut tree = ItemTree::default();
let (header, body) = content.split_once('\n').unwrap_or((&content, ""));
if let Some(rest) = header.strip_prefix("WEBVTT") {
let title = rest.trim();
if !title.is_empty() {
doc.push(Node::Heading {
level: 1,
text: escape_text(title),
});
tree.add(None, None, text_kind("title", title, None));
}
}
for block in blank_line_blocks(body) {
if block.starts_with("NOTE")
|| block.starts_with("STYLE")
|| block.starts_with("REGION")
{
continue;
}
let Some(cue) = parse_cue_block(block) else {
continue;
};
let mut paras: Vec<Vec<Run>> = vec![Vec::new()];
extract_components(&cue.payload, &mut Vec::new(), &mut paras);
for para in ¶s {
if para.is_empty() {
continue;
}
let text = para.iter().map(serialize_run).collect::<Vec<_>>().join(" ");
if !text.is_empty() {
doc.push(Node::Paragraph { text });
}
let track = |run: &Run| TreeTrack {
start_time: cue.start,
end_time: cue.end,
identifier: cue.identifier.clone(),
voice: run.meta.voice.clone().filter(|v| !v.is_empty()),
};
if let [run] = para.as_slice() {
let id = tree.add(
None,
None,
text_kind("text", &run.text, run.meta.formatting),
);
tree.items[id].source = Some(track(run));
} else {
let group = tree.add(
None,
None,
TreeKind::Group {
label: "inline".into(),
name: "WebVTT cue span".into(),
},
);
for run in para {
let id = tree.add(
Some(group),
None,
text_kind("text", &run.text, run.meta.formatting),
);
tree.items[id].source = Some(track(run));
}
}
}
}
doc.tree = Some(tree);
Ok(doc)
}
}
fn text_kind(label: &str, text: &str, formatting: Option<Formatting>) -> TreeKind {
TreeKind::Text {
label: label.into(),
text: text.into(),
orig: None,
formatting,
hyperlink: None,
level: None,
list: None,
}
}
fn blank_line_blocks(body: &str) -> Vec<&str> {
let mut blocks = Vec::new();
let mut start: Option<usize> = None;
let mut offset = 0;
for line in body.split_inclusive('\n') {
let end = offset + line.len();
if line.trim().is_empty() {
if let Some(s) = start.take() {
blocks.push(body[s..offset].trim());
}
} else if start.is_none() {
start = Some(offset);
}
offset = end;
}
if let Some(s) = start {
blocks.push(body[s..].trim());
}
blocks
}
fn parse_cue_block(block: &str) -> Option<Cue> {
let lines: Vec<&str> = block.lines().collect();
let first = *lines.first()?;
let (identifier, timing_line, cue_lines) = if !first.contains("-->") && lines.len() > 1 {
(Some(first.to_string()), lines[1], &lines[2..])
} else {
(None, first, &lines[1..])
};
let mut parts = timing_line.split("-->");
let start = parts.next()?.trim();
let end = parts.next()?.trim();
if parts.next().is_some() {
return None;
}
let end = end.split([' ', '\t']).next().unwrap_or("");
let (start, end) = (parse_timestamp(start)?, parse_timestamp(end)?);
if end <= start {
return None;
}
let mut cue_text = cue_lines.join("\n").trim().to_string();
if cue_text.starts_with("<v") && !cue_text.contains("</v>") {
cue_text.push_str("</v>");
}
Some(Cue {
start,
end,
identifier,
payload: parse_components(&cue_text),
})
}
fn parse_timestamp(raw: &str) -> Option<f64> {
let (rest, millis) = raw.split_once('.')?;
if millis.len() != 3 || !millis.bytes().all(|b| b.is_ascii_digit()) {
return None;
}
let fields: Vec<&str> = rest.split(':').collect();
let (hours, minutes, seconds) = match fields.as_slice() {
[m, s] => ("", *m, *s),
[h, m, s] if h.len() >= 2 => (*h, *m, *s),
_ => return None,
};
let two_digits = |f: &str| {
(f.len() == 2 && f.bytes().all(|b| b.is_ascii_digit()) && f.as_bytes()[0] <= b'5')
.then(|| f.parse::<u64>().ok())
.flatten()
};
let hours = if hours.is_empty() {
0
} else if hours.bytes().all(|b| b.is_ascii_digit()) {
hours.parse::<u64>().ok()?
} else {
return None;
};
let (minutes, seconds) = (two_digits(minutes)?, two_digits(seconds)?);
let millis: u64 = millis.parse().ok()?;
Some((hours * 3600 + minutes * 60 + seconds) as f64 + millis as f64 / 1000.0)
}
fn is_timestamp_tag(tag: &str) -> bool {
parse_timestamp(tag).is_some()
}
fn parse_components(cue_text: &str) -> Vec<Comp> {
let mut stack: Vec<(Option<Tag>, Option<String>, Vec<Comp>)> = Vec::new();
let mut root: Vec<Comp> = Vec::new();
let mut buf = String::new();
let mut rest = cue_text;
let flush = |buf: &mut String, target: &mut Vec<Comp>| {
if buf.is_empty() {
return;
}
let text = std::mem::take(buf);
let pieces: Vec<&str> = text.split('\n').collect();
let n = pieces.len();
for (i, line) in pieces.iter().enumerate() {
if !line.is_empty() {
target.push(Comp::Text {
text: line.to_string(),
terminator: i + 1 < n,
});
}
}
};
while let Some(lt) = rest.find('<') {
let Some(gt_rel) = rest[lt..].find('>') else {
break;
};
let tag = &rest[lt + 1..lt + gt_rel];
let after = &rest[lt + gt_rel + 1..];
if is_timestamp_tag(tag) {
buf.push_str(&rest[..lt]);
rest = after;
continue;
}
let Some((closing, kind, annotation)) = parse_tag(tag) else {
buf.push_str(&rest[..lt]);
rest = after;
continue;
};
buf.push_str(&rest[..lt]);
flush(&mut buf, stack.last_mut().map_or(&mut root, |s| &mut s.2));
rest = after;
if closing {
if let Some((t, annotation, children)) = stack.pop() {
let target = stack.last_mut().map_or(&mut root, |s| &mut s.2);
match t {
Some(t) if t != kind => {
return Vec::new();
}
Some(tag) => target.push(Comp::Span {
tag,
annotation,
children,
}),
None => target.extend(children),
}
}
} else {
let tag = (stack.len() < MAX_SPAN_DEPTH).then_some(kind);
stack.push((tag, annotation, Vec::new()));
}
}
buf.push_str(rest);
flush(&mut buf, stack.last_mut().map_or(&mut root, |s| &mut s.2));
root
}
fn parse_tag(tag: &str) -> Option<(bool, Tag, Option<String>)> {
let (closing, body) = match tag.strip_prefix('/') {
Some(b) => (true, b),
None => (false, tag),
};
let name_len = body
.find(|c: char| !c.is_ascii_alphabetic())
.unwrap_or(body.len());
let (name, mut after) = body.split_at(name_len);
let kind = match name {
"i" => Tag::Italic,
"b" => Tag::Bold,
"u" => Tag::Underline,
"v" => Tag::Voice,
"c" | "lang" => Tag::Transparent,
_ => return None,
};
while let Some(cls) = after.strip_prefix('.') {
let len = cls
.find(['\t', '\n', '\r', ' ', '&', '<', '>', '.'])
.unwrap_or(cls.len());
if len == 0 {
return None;
}
after = &cls[len..];
}
let annotation = match after.chars().next() {
None => None,
Some(' ' | '\t') => {
let a = &after[1..];
if a.contains(['\n', '\r', '&']) {
return None;
}
let a = a.trim();
(!a.is_empty()).then(|| a.to_string())
}
Some(_) => return None,
};
Some((closing, kind, annotation))
}
fn extract_components(comps: &[Comp], parents: &mut Vec<Meta>, paras: &mut Vec<Vec<Run>>) {
for comp in comps {
let mut meta = parents.last().cloned().unwrap_or_default();
match comp {
Comp::Text { text, terminator } => {
paras.last_mut().expect("paragraph").push(Run {
text: text.clone(),
meta,
});
if *terminator {
paras.push(Vec::new());
}
}
Comp::Span {
tag,
annotation,
children,
} => {
match tag {
Tag::Bold => {
meta.formatting.get_or_insert_with(Formatting::default).bold = true
}
Tag::Italic => {
meta.formatting
.get_or_insert_with(Formatting::default)
.italic = true
}
Tag::Underline => {
meta.formatting
.get_or_insert_with(Formatting::default)
.underline = true
}
Tag::Voice => meta.voice = annotation.clone(),
Tag::Transparent => {}
}
parents.push(meta);
extract_components(children, parents, paras);
parents.pop();
}
}
}
}
fn serialize_run(run: &Run) -> String {
let mut s = escape_text(&run.text);
let fmt = run.meta.formatting.unwrap_or_default();
if fmt.bold {
s = format!("**{s}**");
}
if fmt.italic {
s = format!("*{s}*");
}
s
}
#[cfg(test)]
mod tests {
use super::*;
use crate::format::InputFormat;
fn convert(vtt: &str) -> DoclingDocument {
let src = SourceDocument::from_bytes("t", InputFormat::Vtt, vtt.as_bytes().to_vec());
WebVttBackend.convert(&src).unwrap()
}
fn md(vtt: &str) -> String {
convert(vtt).export_to_markdown()
}
#[test]
fn strips_voice_and_skips_notes() {
let out = md("WEBVTT\n\nNOTE hi\n\n00:01.000 --> 00:02.000\n<v Roger>Hello world\n");
assert_eq!(out.trim(), "Hello world");
}
#[test]
fn text_after_multiline_span_stays_in_reading_order() {
let out =
md("WEBVTT\n\n00:00:01.000 --> 00:00:05.000\n<v Bob>Hello\nthere</v> and afterwards\n");
assert_eq!(out.trim(), "Hello\n\nthere and afterwards");
}
#[test]
fn cue_timestamp_tags_are_stripped_from_the_run() {
let out = md("WEBVTT\n\n00:00:00.030 --> 00:00:02.669\nthe<00:00:00.389> quick<00:00:00.750> brown\n");
assert_eq!(out.trim(), "the quick brown");
}
#[test]
fn cr_and_crlf_terminators_parse() {
assert_eq!(
md("WEBVTT\r\r00:00:00.000 --> 00:00:01.000\rHello world\r").trim(),
"Hello world"
);
assert_eq!(
md("WEBVTT\r\n\r\n00:00:00.000 --> 00:00:01.000\r\nHello\r\nworld\r\n").trim(),
"Hello\n\nworld"
);
}
#[test]
fn nested_spans_serialize_with_inline_join() {
let out = md("WEBVTT\n\n00:01.000 --> 00:02.000\n\
a <i>b <lang es>c</lang></i> d\n");
assert_eq!(out.trim(), "a *b * *c* d");
}
#[test]
fn tree_carries_track_source_and_formatting() {
let doc = convert(
"WEBVTT Kitchen talk\n\nid-1\n01:02:03.500 --> 01:02:04.750 line:0\n\
<v Chef>Hello <b>there</b></v>\nBye\n",
);
let json = doc.export_to_json();
let v: serde_json::Value = serde_json::from_str(&json).unwrap();
let texts = v["texts"].as_array().unwrap();
assert_eq!(texts[0]["label"], "title");
assert_eq!(texts[0]["text"], "Kitchen talk");
assert!(texts[0].get("source").is_none());
assert_eq!(v["groups"][0]["name"], "WebVTT cue span");
assert_eq!(v["groups"][0]["label"], "inline");
assert_eq!(texts[1]["parent"]["$ref"], "#/groups/0");
assert_eq!(texts[1]["text"], "Hello ");
assert_eq!(
texts[1]["source"],
serde_json::json!([{
"kind": "track",
"start_time": 3723.5,
"end_time": 3724.75,
"identifier": "id-1",
"voice": "Chef",
}])
);
assert!(texts[1].get("formatting").is_none());
assert_eq!(texts[2]["text"], "there");
assert_eq!(texts[2]["formatting"]["bold"], true);
assert_eq!(texts[2]["source"][0]["voice"], "Chef");
assert_eq!(texts[3]["text"], "Bye");
assert_eq!(texts[3]["parent"]["$ref"], "#/groups/0");
assert!(texts[3]["source"][0].get("voice").is_none());
assert_eq!(v["body"]["children"].as_array().unwrap().len(), 2);
}
#[test]
fn malformed_timings_drop_the_block() {
assert_eq!(md("WEBVTT\n\n00:01.000 -> 00:02.000\nlost\n").trim(), "");
assert_eq!(md("WEBVTT\n\n00:02.000 --> 00:01.000\nlost\n").trim(), "");
assert_eq!(
md("WEBVTT\n\n00:01.000 --> 00:02.000 align:start\n<v Ann>kept\n").trim(),
"kept"
);
}
}