use serde::de::{MapAccess, Visitor};
use serde::ser::SerializeMap;
use serde::{Deserializer, Serializer};
use std::collections::BTreeSet;
use std::fmt;
use std::sync::Mutex;
use super::AttrValue;
const MAX_INTERNED: usize = 4096;
static INTERNED: Mutex<BTreeSet<&'static str>> = Mutex::new(BTreeSet::new());
#[derive(Debug)]
pub(crate) struct InternLimitExceeded;
impl fmt::Display for InternLimitExceeded {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"doctree interner exceeded {MAX_INTERNED} distinct strings; \
refusing to intern more (this input likely wasn't produced by \
this crate's own encoder)"
)
}
}
impl std::error::Error for InternLimitExceeded {}
pub(crate) fn intern(s: &str) -> Result<&'static str, InternLimitExceeded> {
let mut interned = INTERNED.lock().unwrap();
if let Some(&existing) = interned.get(s) {
return Ok(existing);
}
if interned.len() >= MAX_INTERNED {
return Err(InternLimitExceeded);
}
let leaked: &'static str = Box::leak(s.to_owned().into_boxed_str());
interned.insert(leaked);
Ok(leaked)
}
pub(crate) fn serialize_str<S>(value: &&'static str, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.serialize_str(value)
}
pub(crate) fn serialize_extra<S>(
extra: &[(&'static str, AttrValue)],
serializer: S,
) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut map = serializer.serialize_map(Some(extra.len()))?;
for (key, value) in extra {
map.serialize_entry(key, value)?;
}
map.end()
}
struct ExtraVisitor;
impl<'de> Visitor<'de> for ExtraVisitor {
type Value = Vec<(&'static str, AttrValue)>;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a map of attribute keys to values")
}
fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
where
A: MapAccess<'de>,
{
let mut out = Vec::with_capacity(map.size_hint().unwrap_or(0));
while let Some((key, value)) = map.next_entry::<String, AttrValue>()? {
let key = intern(&key).map_err(serde::de::Error::custom)?;
out.push((key, value));
}
Ok(out)
}
}
pub(crate) fn deserialize_extra<'de, D>(
deserializer: D,
) -> Result<Vec<(&'static str, AttrValue)>, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_map(ExtraVisitor)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn intern_returns_pointer_equal_str_on_repeat_calls() {
let a = intern("paragraph").unwrap();
let b = intern("paragraph").unwrap();
assert_eq!(a, "paragraph");
assert!(std::ptr::eq(a, b));
}
#[test]
fn intern_distinguishes_different_strings() {
let a = intern("section").unwrap();
let b = intern("title").unwrap();
assert_ne!(a, b);
}
}