use std::collections::HashMap;
use std::sync::OnceLock;
use super::super::ggml_type::GgmlType;
use super::error::ApexError;
use super::fingerprint::{vendor_config_content, ApexConfigRef};
#[derive(Debug)]
pub struct MudlerConfig {
pub map: HashMap<String, GgmlType>,
pub source_path: &'static str,
}
impl MudlerConfig {
pub fn parse(content: &str, source_path: &'static str) -> Result<Self, ApexError> {
let mut map: HashMap<String, GgmlType> = HashMap::new();
for (lineno, raw) in content.lines().enumerate() {
let line = raw.trim();
if line.is_empty() || line.starts_with('#') {
continue;
}
let (name, tyname) =
line.split_once('=')
.ok_or_else(|| ApexError::MudlerConfigParse {
source_path: source_path.to_string(),
line_number: lineno + 1,
detail: format!("missing `=` separator in `{line}`"),
})?;
let name = name.trim().to_string();
let tyname = tyname.trim();
let ggml = GgmlType::from_name(tyname).ok_or_else(|| ApexError::MudlerConfigParse {
source_path: source_path.to_string(),
line_number: lineno + 1,
detail: format!("unknown GgmlType token `{tyname}`"),
})?;
if let Some(prev) = map.insert(name.clone(), ggml) {
if prev != ggml {
return Err(ApexError::MudlerConfigParse {
source_path: source_path.to_string(),
line_number: lineno + 1,
detail: format!(
"tensor `{name}` reassigned: {} → {}",
prev.name(),
ggml.name()
),
});
}
}
}
Ok(Self { map, source_path })
}
pub fn target_for(&self, tensor_name: &str) -> Result<GgmlType, ApexError> {
if let Some(&t) = self.map.get(tensor_name) {
return Ok(t);
}
for (key, &t) in &self.map {
if tensor_name.len() > key.len() + 1
&& tensor_name.starts_with(key)
&& tensor_name.as_bytes().get(key.len()) == Some(&b'.')
{
return Ok(t);
}
}
Err(ApexError::TensorNotInMudlerConfig {
source_path: self.source_path.to_string(),
tensor_name: tensor_name.to_string(),
})
}
pub fn contains_match(&self, tensor_name: &str) -> bool {
if self.map.contains_key(tensor_name) {
return true;
}
for key in self.map.keys() {
if tensor_name.len() > key.len() + 1
&& tensor_name.starts_with(key)
&& tensor_name.as_bytes().get(key.len()) == Some(&b'.')
{
return true;
}
}
false
}
}
fn cache_slot(path: &'static str) -> &'static OnceLock<Result<MudlerConfig, ApexError>> {
use std::sync::Mutex;
static SLOTS: Mutex<
Vec<(
&'static str,
&'static OnceLock<Result<MudlerConfig, ApexError>>,
)>,
> = Mutex::new(Vec::new());
let mut slots = SLOTS.lock().unwrap();
if let Some((_, slot)) = slots.iter().find(|(p, _)| *p == path) {
return slot;
}
let slot: &'static OnceLock<Result<MudlerConfig, ApexError>> =
Box::leak(Box::new(OnceLock::new()));
slots.push((path, slot));
slot
}
pub fn load_mudler_config(entry: &ApexConfigRef) -> Result<&'static MudlerConfig, ApexError> {
let static_path: &'static str = super::fingerprint::VENDOR_CONFIGS
.iter()
.find(|(p, _)| *p == entry.mudler_config_path)
.map(|(p, _)| *p)
.ok_or_else(|| ApexError::FingerprintConfigMissing {
fingerprint: entry.fingerprint.clone(),
mudler_config_path: entry.mudler_config_path.clone(),
})?;
let content =
vendor_config_content(static_path).ok_or_else(|| ApexError::FingerprintConfigMissing {
fingerprint: entry.fingerprint.clone(),
mudler_config_path: entry.mudler_config_path.clone(),
})?;
let slot = cache_slot(static_path);
let result = slot.get_or_init(|| MudlerConfig::parse(content, static_path));
match result {
Ok(cfg) => Ok(cfg),
Err(e) => Err(e.clone()),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_gemma4_balanced_layer_5_routed_expert() {
let content = super::super::fingerprint::vendor_config_content(
"vendor/apex-quant/configs/gemma4_26b_balanced.txt",
)
.expect("baked vendor content present");
let cfg = MudlerConfig::parse(content, "test").unwrap();
assert_eq!(
cfg.target_for("blk.5.ffn_gate_exps").unwrap(),
GgmlType::Q5_K
);
assert_eq!(
cfg.target_for("blk.5.ffn_gate_exps.weight").unwrap(),
GgmlType::Q5_K
);
assert_eq!(
cfg.target_for("blk.0.attn_q.weight").unwrap(),
GgmlType::Q6_K
);
assert_eq!(
cfg.target_for("blk.0.ffn_gate_shexp.weight").unwrap(),
GgmlType::Q8_0
);
}
#[test]
fn parse_carnice_mtp_quality_layer_5() {
let content = super::super::fingerprint::vendor_config_content(
"vendor/apex-quant/configs/carnice_qwen36_mtp_quality.txt",
)
.expect("baked vendor content present");
let cfg = MudlerConfig::parse(content, "test").unwrap();
assert_eq!(
cfg.target_for("blk.5.ffn_gate_exps").unwrap(),
GgmlType::Q5_K
);
assert_eq!(cfg.target_for("blk.5.attn_q").unwrap(), GgmlType::Q6_K);
assert_eq!(
cfg.target_for("blk.5.ffn_gate_shexp").unwrap(),
GgmlType::Q8_0
);
}
#[test]
fn parse_handles_mixed_case_ggml_names() {
let content = "blk.0.attn_q=Q5_K\nblk.0.attn_k=q5_K\nblk.0.attn_v=q5_k";
let cfg = MudlerConfig::parse(content, "test").unwrap();
assert_eq!(cfg.target_for("blk.0.attn_q").unwrap(), GgmlType::Q5_K);
assert_eq!(cfg.target_for("blk.0.attn_k").unwrap(), GgmlType::Q5_K);
assert_eq!(cfg.target_for("blk.0.attn_v").unwrap(), GgmlType::Q5_K);
}
#[test]
fn parse_skips_blank_and_comment_lines() {
let content = "
# this is a comment
blk.0.attn_q=Q5_K
## indented comment
blk.0.attn_k=Q5_K
";
let cfg = MudlerConfig::parse(content, "test").unwrap();
assert_eq!(cfg.map.len(), 2);
}
#[test]
fn parse_errors_on_missing_separator() {
let content = "blk.0.attn_q Q5_K";
let err = MudlerConfig::parse(content, "test/path").unwrap_err();
match err {
ApexError::MudlerConfigParse {
source_path,
line_number,
..
} => {
assert_eq!(source_path, "test/path");
assert_eq!(line_number, 1);
}
other => panic!("expected MudlerConfigParse, got {other:?}"),
}
}
#[test]
fn parse_errors_on_unknown_ggml_type() {
let content = "blk.0.attn_q=Q9_K";
let err = MudlerConfig::parse(content, "test").unwrap_err();
assert!(matches!(err, ApexError::MudlerConfigParse { .. }));
}
#[test]
fn missing_tensor_returns_typed_error() {
let cfg = MudlerConfig::parse("blk.0.attn_q=Q5_K", "test").unwrap();
let err = cfg.target_for("blk.99.fake").unwrap_err();
assert!(matches!(err, ApexError::TensorNotInMudlerConfig { .. }));
}
}