use rkyv::{Archive, Deserialize as RkyvDeserialize, Serialize as RkyvSerialize};
use serde::{Deserialize, Serialize};
use crate::dictionary::context_id_map::ContextIdMap;
use crate::dictionary::schema::Schema;
const DEFAULT_WORD_COST: i16 = -10000;
const DEFAULT_LEFT_CONTEXT_ID: u16 = 1288;
const DEFAULT_RIGHT_CONTEXT_ID: u16 = 1288;
const DEFAULT_FIELD_VALUE: &str = "*";
pub const DICTIONARY_FORMAT_VERSION: u32 = 2;
const LEGACY_FORMAT_VERSION: u32 = 1;
fn legacy_format_version() -> u32 {
LEGACY_FORMAT_VERSION
}
#[derive(Clone, Serialize, Deserialize, Archive, RkyvSerialize, RkyvDeserialize)]
pub struct ModelInfo {
pub feature_count: usize,
pub label_count: usize,
pub max_left_context_id: usize,
pub max_right_context_id: usize,
pub connection_matrix_size: String,
pub version: String,
pub training_iterations: u64,
pub regularization: f64,
pub updated_at: u64,
}
#[derive(Clone, Serialize, Deserialize, Archive, RkyvSerialize, RkyvDeserialize)]
pub struct Metadata {
#[serde(default = "legacy_format_version")]
pub format_version: u32,
pub name: String, pub encoding: String, pub default_word_cost: i16, pub default_left_context_id: u16, pub default_right_context_id: u16, pub default_field_value: String, pub flexible_csv: bool, pub skip_invalid_cost_or_id: bool, pub normalize_details: bool, #[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub connection_id_mapping: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub context_id_map: Option<ContextIdMap>,
pub dictionary_schema: Schema, pub user_dictionary_schema: Schema, #[serde(skip_serializing_if = "Option::is_none")]
pub model_info: Option<ModelInfo>, }
impl Default for Metadata {
fn default() -> Self {
Metadata::new(
"default".to_string(),
"UTF-8".to_string(),
DEFAULT_WORD_COST,
DEFAULT_LEFT_CONTEXT_ID,
DEFAULT_RIGHT_CONTEXT_ID,
DEFAULT_FIELD_VALUE.to_string(),
false,
false,
false,
Schema::default(),
Schema::new(vec![
"surface".to_string(),
"reading".to_string(),
"pronunciation".to_string(),
]),
)
}
}
impl Metadata {
#[allow(clippy::too_many_arguments)]
pub fn new(
name: String,
encoding: String,
simple_word_cost: i16,
default_left_context_id: u16,
default_right_context_id: u16,
default_field_value: String,
flexible_csv: bool,
skip_invalid_cost_or_id: bool,
normalize_details: bool,
schema: Schema,
userdic_schema: Schema,
) -> Self {
Self {
format_version: DICTIONARY_FORMAT_VERSION,
encoding,
default_word_cost: simple_word_cost,
default_left_context_id,
default_right_context_id,
default_field_value,
dictionary_schema: schema,
name,
flexible_csv,
skip_invalid_cost_or_id,
normalize_details,
connection_id_mapping: false,
context_id_map: None,
user_dictionary_schema: userdic_schema,
model_info: None,
}
}
pub fn load(data: &[u8]) -> crate::LinderaResult<Self> {
if data.is_empty() {
return Err(crate::error::LinderaErrorKind::Io
.with_error(anyhow::anyhow!("Empty metadata data")));
}
serde_json::from_slice(data).map_err(|err| {
crate::error::LinderaErrorKind::Deserialize
.with_error(anyhow::anyhow!(err))
.add_context("Failed to deserialize metadata from JSON")
})
}
pub fn validate_format_version(&self) -> crate::LinderaResult<()> {
if self.format_version == DICTIONARY_FORMAT_VERSION {
return Ok(());
}
let hint = if self.format_version < DICTIONARY_FORMAT_VERSION {
"rebuild it with `lindera build`, or download a matching prebuilt dictionary with `lindera download`"
} else {
"upgrade Lindera to a version that understands this dictionary"
};
Err(crate::error::LinderaErrorKind::Deserialize.with_error(anyhow::anyhow!(
"Dictionary '{}' has format version {}, but this build of Lindera reads format version {}. To fix this, {hint}.",
self.name,
self.format_version,
DICTIONARY_FORMAT_VERSION,
)))
}
pub fn load_or_default(data: &[u8], default_fn: fn() -> Self) -> Self {
if data.is_empty() {
default_fn()
} else {
match Self::load(data) {
Ok(metadata) => metadata,
Err(_) => default_fn(),
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_metadata_default() {
let metadata = Metadata::default();
assert_eq!(metadata.name, "default");
}
#[test]
fn metadata_without_format_version_reads_as_legacy() {
let json = serde_json::to_value(Metadata::default()).unwrap();
let mut object = json.as_object().unwrap().clone();
object.remove("format_version");
let without = serde_json::to_vec(&object).unwrap();
let metadata = Metadata::load(&without).unwrap();
assert_eq!(metadata.format_version, LEGACY_FORMAT_VERSION);
}
#[test]
fn legacy_version_is_the_first_format_version() {
assert_eq!(LEGACY_FORMAT_VERSION, 1);
}
#[test]
fn validate_format_version_accepts_the_current_version() {
let metadata = Metadata::default();
assert_eq!(metadata.format_version, DICTIONARY_FORMAT_VERSION);
assert!(metadata.validate_format_version().is_ok());
}
#[test]
fn validate_format_version_rejects_an_older_dictionary() {
let metadata = Metadata {
name: "ipadic".to_string(),
format_version: DICTIONARY_FORMAT_VERSION - 1,
..Metadata::default()
};
let err = metadata.validate_format_version().unwrap_err().to_string();
assert!(err.contains("ipadic"), "{err}");
assert!(err.contains("lindera build"), "{err}");
}
#[test]
fn validate_format_version_rejects_a_newer_dictionary() {
let metadata = Metadata {
format_version: DICTIONARY_FORMAT_VERSION + 1,
..Metadata::default()
};
let err = metadata.validate_format_version().unwrap_err().to_string();
assert!(err.contains("upgrade Lindera"), "{err}");
}
#[test]
fn format_version_round_trips_through_json() {
let metadata = Metadata {
format_version: 7,
..Metadata::default()
};
let json = serde_json::to_vec(&metadata).unwrap();
assert_eq!(Metadata::load(&json).unwrap().format_version, 7);
}
#[test]
fn test_metadata_new() {
let schema = Schema::default();
let metadata = Metadata::new(
"TestDict".to_string(),
"UTF-8".to_string(),
-10000,
0,
0,
"*".to_string(),
false,
false,
false,
schema.clone(),
Schema::new(vec!["surface".to_string(), "reading".to_string()]),
);
assert_eq!(metadata.name, "TestDict");
}
#[test]
fn test_metadata_serialization() {
let metadata = Metadata::default();
let serialized = serde_json::to_string(&metadata).unwrap();
assert!(serialized.contains("default"));
assert!(serialized.contains("schema"));
assert!(serialized.contains("name"));
let deserialized: Metadata = serde_json::from_str(&serialized).unwrap();
assert_eq!(deserialized.name, "default");
}
}