use log::{Level, LevelFilter, Log, Metadata, Record};
use std::sync::Once;
static INIT: Once = Once::new();
struct EpLogger {
level: LevelFilter,
}
impl Log for EpLogger {
fn enabled(&self, metadata: &Metadata) -> bool {
metadata.level() <= self.level
}
fn log(&self, record: &Record) {
if self.enabled(record.metadata()) {
let level_tag = match record.level() {
Level::Error => "ERROR",
Level::Warn => "WARN",
Level::Info => "INFO",
Level::Debug => "DEBUG",
Level::Trace => "TRACE",
};
eprintln!("[rust-mlx-ep] {level_tag}: {}", record.args());
}
}
fn flush(&self) {}
}
fn resolve_level() -> LevelFilter {
if let Ok(val) = std::env::var("RUST_LOG") {
let level_str = val
.split(',')
.find(|s| s.contains("onnxruntime_ep_mlx"))
.and_then(|s| s.split('=').nth(1))
.unwrap_or(&val);
if let Ok(f) = level_str.parse::<LevelFilter>() {
return f;
}
}
if std::env::var("ONNXRUNTIME_EP_MLX_TRACE")
.ok()
.filter(|s| !s.is_empty())
.is_some()
{
return LevelFilter::Debug;
}
if std::env::var("ONNXRUNTIME_EP_MLX_VERBOSE")
.map(|v| v == "1")
.unwrap_or(false)
{
return LevelFilter::Info;
}
LevelFilter::Warn
}
pub fn init() {
INIT.call_once(|| {
let level = resolve_level();
let logger = Box::new(EpLogger { level });
let _ = log::set_boxed_logger(logger);
log::set_max_level(level);
});
}