use std::collections::HashSet;
use std::num::NonZeroUsize;
use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::Mutex;
use tracing::{debug, warn};
use crate::domain::detector::{Confidence, ContentTypeDetector, DetectedType};
pub struct MagikaDetector {
session: Arc<Mutex<Option<magika::Session>>>,
supported_extensions: Arc<HashSet<String>>,
intra_op_threads: Option<NonZeroUsize>,
}
impl MagikaDetector {
pub fn new(
supported_extensions: impl IntoIterator<Item = impl Into<String>>,
) -> magika::Result<Self> {
Self::with_config(supported_extensions, None)
}
pub fn with_config(
supported_extensions: impl IntoIterator<Item = impl Into<String>>,
intra_op_threads: Option<NonZeroUsize>,
) -> magika::Result<Self> {
Ok(Self {
session: Arc::new(Mutex::new(Some(Self::build_session(intra_op_threads)?))),
supported_extensions: Arc::new(
supported_extensions
.into_iter()
.map(|ext| ext.into().to_lowercase())
.collect(),
),
intra_op_threads,
})
}
fn build_session(intra_op_threads: Option<NonZeroUsize>) -> magika::Result<magika::Session> {
let mut builder = magika::Session::builder();
if let Some(threads) = intra_op_threads {
builder = builder.with_intra_threads(threads.get());
}
builder.build()
}
}
fn map_inference_result(
result: magika::Result<magika::FileType>,
supported_extensions: &HashSet<String>,
) -> Option<DetectedType> {
let file_type = result
.inspect_err(|e| warn!(error = %e, "magika inference failed"))
.ok()?;
let info = file_type.info();
let Some(extension) = info
.extensions
.iter()
.find(|ext| supported_extensions.contains(&ext.to_lowercase()))
else {
debug!(
label = info.label,
candidate_extensions = ?info.extensions,
"magika: detected label maps to no registered extension; falling back to hint"
);
return None;
};
let score = file_type.score();
let Some(confidence) = Confidence::new(score) else {
warn!(
score,
label = info.label,
"magika: inference returned a NaN score"
);
return None;
};
Some(DetectedType {
extension: extension.to_lowercase(),
confidence,
})
}
impl MagikaDetector {
#[tracing::instrument(
name = "magika.inference",
skip(self, infer),
fields(intra_op_threads = ?self.intra_op_threads)
)]
async fn run_inference(
&self,
infer: impl FnOnce(&mut magika::Session) -> magika::Result<magika::FileType> + Send + 'static,
) -> Option<DetectedType> {
let mut guard = Arc::clone(&self.session).lock_owned().await;
if guard.is_none() {
*guard = Self::rebuild_session(self.intra_op_threads).await;
if guard.is_none() {
return None;
}
}
let supported_extensions = Arc::clone(&self.supported_extensions);
tokio::task::spawn_blocking(move || {
let mut guard = guard;
let mut session = guard.take()?;
let result = map_inference_result(infer(&mut session), &supported_extensions);
*guard = Some(session);
result
})
.await
.unwrap_or_else(|join_err| {
warn!(error = %join_err, "magika inference task panicked; session will be rebuilt");
None
})
}
async fn rebuild_session(intra_op_threads: Option<NonZeroUsize>) -> Option<magika::Session> {
match tokio::task::spawn_blocking(move || Self::build_session(intra_op_threads)).await {
Ok(Ok(fresh_session)) => Some(fresh_session),
Ok(Err(e)) => {
warn!(error = %e, "failed to rebuild the magika session");
None
}
Err(join_err) => {
warn!(error = %join_err, "rebuilding the magika session panicked");
None
}
}
}
}
#[async_trait]
impl ContentTypeDetector for MagikaDetector {
async fn detect(&self, bytes: bytes::Bytes) -> Option<DetectedType> {
self.run_inference(move |session| session.identify_content_sync(bytes.as_ref()))
.await
}
async fn detect_path(&self, path: &std::path::Path) -> Option<DetectedType> {
let path = path.to_path_buf();
self.run_inference(move |session| session.identify_file_sync(&path))
.await
}
}
#[cfg(test)]
mod tests {
use super::*;
fn inferred(content_type: magika::ContentType, score: f32) -> magika::FileType {
magika::FileType::Inferred(magika::InferredType {
content_type: None,
inferred_type: content_type,
score,
})
}
#[test]
fn inference_error_is_swallowed_to_none() {
let err = magika::Error::IOError(std::io::Error::other("simulated inference failure"));
assert!(map_inference_result(Err(err), &HashSet::new()).is_none());
}
#[test]
fn label_present_in_supported_extensions_maps_to_detected_type() {
let supported: HashSet<String> = ["pdf", "html"].into_iter().map(str::to_owned).collect();
let detected =
map_inference_result(Ok(inferred(magika::ContentType::Pdf, 0.97)), &supported)
.expect("pdf is in supported_extensions, so this must resolve");
assert_eq!(detected.extension, "pdf");
assert!((detected.confidence.get() - 0.97).abs() < f32::EPSILON);
}
#[test]
fn label_absent_from_supported_extensions_returns_none() {
let supported: HashSet<String> = ["html"].into_iter().map(str::to_owned).collect();
assert!(
map_inference_result(Ok(inferred(magika::ContentType::Pdf, 0.97)), &supported)
.is_none()
);
}
#[test]
fn nan_score_returns_none() {
let supported: HashSet<String> = ["pdf"].into_iter().map(str::to_owned).collect();
assert!(
map_inference_result(Ok(inferred(magika::ContentType::Pdf, f32::NAN)), &supported)
.is_none()
);
}
#[test]
fn label_extension_is_lowercased() {
let supported: HashSet<String> = ["PDF".to_lowercase()].into_iter().collect();
let detected =
map_inference_result(Ok(inferred(magika::ContentType::Pdf, 0.5)), &supported)
.expect("lowercased supported_extensions must still match");
assert_eq!(detected.extension, "pdf");
}
}