use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
#[derive(Debug)]
pub enum ProviderError {
NotFound {
key: String,
},
Backend {
detail: String,
},
}
impl std::fmt::Display for ProviderError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NotFound { key } => {
write!(f, "ContentProvider: key not found: {key}")
}
Self::Backend { detail } => {
write!(f, "ContentProvider: backend error: {detail}")
}
}
}
}
impl std::error::Error for ProviderError {}
pub type ProviderResult<T> = Result<T, ProviderError>;
pub trait ContentProvider {
fn fetch(&self, key: &str) -> ProviderResult<Vec<u8>>;
fn fetch_string(&self, key: &str) -> ProviderResult<String> {
let bytes = self.fetch(key)?;
String::from_utf8(bytes).map_err(|e| ProviderError::Backend {
detail: format!("invalid utf-8 in {key}: {e}"),
})
}
fn contains(&self, key: &str) -> bool {
self.fetch(key).is_ok()
}
}
#[derive(Debug, Clone)]
pub struct FsContentProvider {
root: PathBuf,
}
impl FsContentProvider {
#[must_use]
pub fn new<P: Into<PathBuf>>(root: P) -> Self {
Self { root: root.into() }
}
#[must_use]
pub fn root(&self) -> &Path {
&self.root
}
fn resolve(&self, key: &str) -> ProviderResult<PathBuf> {
if key.split('/').any(|seg| seg == "..") {
return Err(ProviderError::Backend {
detail: format!("rejected traversal key: {key}"),
});
}
Ok(self.root.join(key))
}
}
impl ContentProvider for FsContentProvider {
fn fetch(&self, key: &str) -> ProviderResult<Vec<u8>> {
let path = self.resolve(key)?;
match std::fs::read(&path) {
Ok(bytes) => Ok(bytes),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
Err(ProviderError::NotFound { key: key.into() })
}
Err(e) => Err(ProviderError::Backend {
detail: format!("read {}: {e}", path.display()),
}),
}
}
fn contains(&self, key: &str) -> bool {
self.resolve(key).is_ok_and(|p| p.exists())
}
}
#[derive(Debug, Clone, Default)]
pub struct MemoryContentProvider {
map: BTreeMap<String, Vec<u8>>,
}
impl MemoryContentProvider {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn insert<K: Into<String>, V: Into<Vec<u8>>>(
&mut self,
key: K,
value: V,
) -> Option<Vec<u8>> {
self.map.insert(key.into(), value.into())
}
#[must_use]
pub fn len(&self) -> usize {
self.map.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.map.is_empty()
}
}
impl ContentProvider for MemoryContentProvider {
fn fetch(&self, key: &str) -> ProviderResult<Vec<u8>> {
self.map
.get(key)
.cloned()
.ok_or_else(|| ProviderError::NotFound { key: key.into() })
}
fn contains(&self, key: &str) -> bool {
self.map.contains_key(key)
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
#[test]
fn memory_provider_round_trip() {
let mut mem = MemoryContentProvider::new();
assert!(mem.is_empty());
let _ = mem.insert("a.md", b"hello".to_vec());
assert_eq!(mem.len(), 1);
assert!(!mem.is_empty());
assert!(mem.contains("a.md"));
assert!(!mem.contains("missing"));
let bytes = mem.fetch("a.md").unwrap();
assert_eq!(bytes, b"hello");
let text = mem.fetch_string("a.md").unwrap();
assert_eq!(text, "hello");
}
#[test]
fn memory_provider_not_found_is_distinct() {
let mem = MemoryContentProvider::new();
match mem.fetch("nope") {
Err(ProviderError::NotFound { key }) => assert_eq!(key, "nope"),
other => panic!("expected NotFound, got {other:?}"),
}
}
#[test]
fn memory_provider_invalid_utf8_is_backend_error() {
let mut mem = MemoryContentProvider::new();
let _ = mem.insert("bad", vec![0xffu8, 0xfe, 0xfd]);
match mem.fetch_string("bad") {
Err(ProviderError::Backend { detail }) => {
assert!(detail.contains("invalid utf-8"));
}
other => panic!("expected Backend, got {other:?}"),
}
}
#[test]
fn provider_error_display() {
let nf = ProviderError::NotFound { key: "a".into() };
let be = ProviderError::Backend {
detail: "boom".into(),
};
assert!(format!("{nf}").contains("not found"));
assert!(format!("{be}").contains("backend"));
}
#[test]
fn fs_provider_reads_file() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("hello.md");
std::fs::write(&path, b"# Hello").unwrap();
let fs = FsContentProvider::new(dir.path());
assert_eq!(fs.root(), dir.path());
let bytes = fs.fetch("hello.md").unwrap();
assert_eq!(bytes, b"# Hello");
assert!(fs.contains("hello.md"));
assert!(!fs.contains("absent.md"));
}
#[test]
fn fs_provider_rejects_traversal() {
let dir = tempfile::tempdir().unwrap();
let fs = FsContentProvider::new(dir.path());
match fs.fetch("../etc/passwd") {
Err(ProviderError::Backend { detail }) => {
assert!(detail.contains("traversal"));
}
other => panic!("expected traversal rejection, got {other:?}"),
}
assert!(!fs.contains("../etc/passwd"));
}
#[test]
fn fs_provider_missing_is_not_found() {
let dir = tempfile::tempdir().unwrap();
let fs = FsContentProvider::new(dir.path());
match fs.fetch("nope.md") {
Err(ProviderError::NotFound { key }) => assert_eq!(key, "nope.md"),
other => panic!("expected NotFound, got {other:?}"),
}
}
#[test]
fn provider_error_debug() {
let nf = ProviderError::NotFound { key: "k".into() };
let s = format!("{nf:?}");
assert!(s.contains("NotFound"));
}
#[test]
fn fs_provider_root_accessor() {
let dir = tempfile::tempdir().unwrap();
let fs = FsContentProvider::new(dir.path());
assert_eq!(fs.root(), dir.path());
}
#[test]
fn fs_provider_clone() {
let dir = tempfile::tempdir().unwrap();
let fs = FsContentProvider::new(dir.path());
let cloned = fs.clone();
assert_eq!(cloned.root(), fs.root());
}
#[test]
fn fs_provider_fetch_string_decodes_utf8() {
let dir = tempfile::tempdir().unwrap();
std::fs::write(dir.path().join("a.md"), "héllo").unwrap();
let fs = FsContentProvider::new(dir.path());
assert_eq!(fs.fetch_string("a.md").unwrap(), "héllo");
}
#[test]
fn fs_provider_fetch_string_rejects_invalid_utf8() {
let dir = tempfile::tempdir().unwrap();
std::fs::write(dir.path().join("bad.md"), [0xffu8, 0xfe, 0xfd])
.unwrap();
let fs = FsContentProvider::new(dir.path());
match fs.fetch_string("bad.md") {
Err(ProviderError::Backend { detail }) => {
assert!(detail.contains("invalid utf-8"));
}
other => panic!("expected Backend, got {other:?}"),
}
}
#[test]
fn fs_provider_contains_when_present() {
let dir = tempfile::tempdir().unwrap();
std::fs::write(dir.path().join("a.md"), "x").unwrap();
let fs = FsContentProvider::new(dir.path());
assert!(fs.contains("a.md"));
}
#[test]
fn fs_provider_nested_traversal_rejected() {
let dir = tempfile::tempdir().unwrap();
let fs = FsContentProvider::new(dir.path());
match fs.fetch("a/../../b") {
Err(ProviderError::Backend { detail }) => {
assert!(detail.contains("traversal"));
}
other => panic!("expected traversal rejection, got {other:?}"),
}
}
#[test]
fn memory_provider_insert_returns_previous_value() {
let mut mem = MemoryContentProvider::new();
assert!(mem.insert("k", b"v1".to_vec()).is_none());
let prev = mem.insert("k", b"v2".to_vec());
assert_eq!(prev.as_deref(), Some(&b"v1"[..]));
assert_eq!(mem.fetch("k").unwrap(), b"v2");
}
#[test]
fn memory_provider_default_equivalent_to_new() {
let a = MemoryContentProvider::default();
let b = MemoryContentProvider::new();
assert_eq!(a.len(), b.len());
assert!(a.is_empty());
}
#[test]
fn provider_error_display_messages() {
let nf = ProviderError::NotFound { key: "x".into() };
assert_eq!(format!("{nf}"), "ContentProvider: key not found: x");
let be = ProviderError::Backend { detail: "y".into() };
assert_eq!(format!("{be}"), "ContentProvider: backend error: y");
}
#[test]
fn provider_error_is_std_error() {
let err: Box<dyn std::error::Error> =
Box::new(ProviderError::NotFound { key: "k".into() });
assert!(err.to_string().contains("not found"));
}
#[test]
fn memory_provider_contains_via_trait_object() {
let mut mem = MemoryContentProvider::new();
let _ = mem.insert("a", b"1".to_vec());
let provider: &dyn ContentProvider = &mem;
assert!(provider.contains("a"));
assert!(!provider.contains("missing"));
}
}