use std::collections::HashMap;
use crate::error::{Error, Result};
const NS_SPREADSHEETML: &str = "http://schemas.openxmlformats.org/spreadsheetml/2006/main";
#[derive(Debug, Clone, Default)]
pub struct SharedStringsTable {
strings: Vec<String>,
index: HashMap<String, usize>,
reference_count: usize,
}
impl SharedStringsTable {
pub fn new() -> Self {
Self::default()
}
pub fn add(&mut self, s: impl Into<String>) -> usize {
let s = s.into();
self.reference_count += 1;
if let Some(&i) = self.index.get(&s) {
return i;
}
let i = self.strings.len();
self.index.insert(s.clone(), i);
self.strings.push(s);
i
}
pub fn index_of(&self, s: &str) -> Option<usize> {
self.index.get(s).copied()
}
pub fn get(&self, index: usize) -> Option<&str> {
self.strings.get(index).map(|s| s.as_str())
}
pub fn unique_count(&self) -> usize {
self.strings.len()
}
pub fn count(&self) -> usize {
self.reference_count
}
pub fn is_empty(&self) -> bool {
self.strings.is_empty()
}
pub fn iter(&self) -> impl Iterator<Item = (usize, &str)> {
self.strings
.iter()
.enumerate()
.map(|(i, s)| (i, s.as_str()))
}
pub fn to_xml(&self) -> String {
let mut out = String::with_capacity(256 + self.strings.len() * 32);
out.push_str("<?xml version=\"1.0\" encoding=\"UTF-8\" standalone=\"yes\"?>\r\n");
out.push_str(&format!(
"<sst xmlns=\"{}\" count=\"{}\" uniqueCount=\"{}\">",
NS_SPREADSHEETML,
self.reference_count,
self.strings.len()
));
for s in &self.strings {
out.push_str("<si><t>");
out.push_str(&xml_escape_text(s));
out.push_str("</t></si>");
}
out.push_str("</sst>");
out
}
pub fn from_xml(xml: &str) -> Result<Self> {
use quick_xml::events::Event;
use quick_xml::reader::Reader;
let mut sst = SharedStringsTable::new();
let mut rd = Reader::from_str(xml);
let mut buf = Vec::new();
let mut in_si = false;
let mut in_r = false;
let mut in_t = false;
let mut current = String::new();
loop {
match rd.read_event_into(&mut buf) {
Ok(Event::Start(e)) => match e.name().as_ref() {
b"si" => {
in_si = true;
current.clear();
}
b"r" if in_si => {
in_r = true;
}
b"t" if in_si => {
in_t = true;
}
_ => {}
},
Ok(Event::End(e)) => match e.name().as_ref() {
b"si" => {
if in_si {
sst.strings.push(current.clone());
let idx = sst.strings.len() - 1;
sst.index.insert(current.clone(), idx);
sst.reference_count += 1;
current.clear();
}
in_si = false;
}
b"t" => {
in_t = false;
}
b"r" => {
in_r = false;
}
_ => {}
},
Ok(Event::Text(t)) if in_si && in_t => {
let text_str = std::str::from_utf8(t.as_ref()).unwrap_or("");
current.push_str(text_str);
}
Ok(Event::GeneralRef(r)) if in_si && in_t => {
if let Some(ch) = r
.resolve_char_ref()
.map_err(|e| Error::Xml(format!("sharedStrings char ref: {e}")))?
{
current.push(ch);
} else {
let name = r
.decode()
.map_err(|e| Error::Xml(format!("sharedStrings entity decode: {e}")))?;
let ch = match name.as_ref() {
"lt" => '<',
"gt" => '>',
"amp" => '&',
"quot" => '"',
"apos" => '\'',
other => {
return Err(Error::Xml(format!(
"sharedStrings unknown entity: &{other};"
)))
}
};
current.push(ch);
}
}
Ok(Event::CData(t)) if in_si && in_t => {
if let Ok(s) = std::str::from_utf8(&t) {
current.push_str(s);
}
}
Ok(Event::Eof) => break,
Ok(_) => {}
Err(e) => {
return Err(Error::Xml(format!("sharedStrings parse: {e}")));
}
}
buf.clear();
}
let _ = in_r;
Ok(sst)
}
}
fn xml_escape_text(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for c in s.chars() {
match c {
'&' => out.push_str("&"),
'<' => out.push_str("<"),
'>' => out.push_str(">"),
_ => out.push(c),
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn add_dedup() {
let mut sst = SharedStringsTable::new();
let i1 = sst.add("Hello");
let i2 = sst.add("World");
let i3 = sst.add("Hello"); assert_eq!(i1, 0);
assert_eq!(i2, 1);
assert_eq!(i3, 0);
assert_eq!(sst.unique_count(), 2);
assert_eq!(sst.count(), 3); }
#[test]
fn get_and_iter() {
let mut sst = SharedStringsTable::new();
sst.add("Hello");
sst.add("World");
assert_eq!(sst.get(0), Some("Hello"));
assert_eq!(sst.get(1), Some("World"));
assert_eq!(sst.get(2), None);
let all: Vec<(usize, &str)> = sst.iter().collect();
assert_eq!(all, vec![(0, "Hello"), (1, "World")]);
}
#[test]
fn round_trip() {
let mut sst = SharedStringsTable::new();
sst.add("Hello");
sst.add("World");
sst.add("中文测试");
let xml = sst.to_xml();
let sst2 = SharedStringsTable::from_xml(&xml).unwrap();
assert_eq!(sst2.unique_count(), 3);
assert_eq!(sst2.get(0), Some("Hello"));
assert_eq!(sst2.get(1), Some("World"));
assert_eq!(sst2.get(2), Some("中文测试"));
}
#[test]
fn xml_escape_special_chars() {
let mut sst = SharedStringsTable::new();
sst.add("a<b>&c");
let xml = sst.to_xml();
assert!(xml.contains("a<b>&c"));
let sst2 = SharedStringsTable::from_xml(&xml).unwrap();
assert_eq!(sst2.get(0), Some("a<b>&c"));
}
#[test]
fn empty_table() {
let sst = SharedStringsTable::new();
assert!(sst.is_empty());
assert_eq!(sst.unique_count(), 0);
let xml = sst.to_xml();
assert!(xml.contains("count=\"0\""));
assert!(xml.contains("uniqueCount=\"0\""));
}
}