use super::{BoostedModel, ModelFormat};
use crate::error::Result;
use std::fmt;
use std::sync::OnceLock;
pub struct EmbeddedModel {
bytes: &'static [u8],
format: ModelFormat,
model: OnceLock<BoostedModel>,
}
impl EmbeddedModel {
#[must_use]
pub const fn new(bytes: &'static [u8], format: ModelFormat) -> Self {
Self {
bytes,
format,
model: OnceLock::new(),
}
}
#[must_use]
pub const fn bytes(&self) -> &'static [u8] {
self.bytes
}
#[must_use]
pub const fn format(&self) -> ModelFormat {
self.format
}
pub fn get(&self) -> Result<&BoostedModel> {
if let Some(model) = self.model.get() {
return Ok(model);
}
let model = BoostedModel::decode(self.bytes, self.format)?;
Ok(self.model.get_or_init(|| model))
}
}
impl fmt::Debug for EmbeddedModel {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("EmbeddedModel")
.field("bytes", &self.bytes.len())
.field("format", &self.format)
.field("decoded", &self.model.get().is_some())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::HessboostError;
const DART: &[u8] = include_bytes!("../../tests/data/saved/0.2.0/dart.bin");
#[test]
fn decodes_once_and_matches_decode() {
static MODEL: EmbeddedModel = EmbeddedModel::new(DART, ModelFormat::Binary);
let first = MODEL.get().unwrap();
assert!(std::ptr::eq(first, MODEL.get().unwrap()));
let decoded = BoostedModel::decode(DART, ModelFormat::Binary).unwrap();
assert_eq!(
first.encode(ModelFormat::Binary).unwrap(),
decoded.encode(ModelFormat::Binary).unwrap()
);
}
#[test]
fn bad_bytes_error_on_every_call() {
static MODEL: EmbeddedModel = EmbeddedModel::new(DART, ModelFormat::Json);
for _ in 0..2 {
assert!(matches!(MODEL.get(), Err(HessboostError::Json(_))));
}
}
}