use byteorder::{ReadBytesExt, WriteBytesExt};
use std::borrow::Cow;
use std::hash::Hash;
use std::io::{Read, Write};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Once, OnceLock, RwLock, RwLockWriteGuard};
use std::thread::LocalKey;
use std::{
cell::{Cell, RefCell},
collections::hash_map::Entry,
ops::{Deref, DerefMut},
};
use ahash::{HashMap, HashMapExt, HashSet, HashSetExt};
use append_only_vec::AppendOnlyVec;
use byteorder::LittleEndian;
use smartstring::alias::String;
use crate::atom::{
DerivativeFunction, EvaluationInfo, NamespacedSymbol, NormalizationFunction,
SeriesExpansionFunction, SymbolAttribute, SymbolBuilder, UserData,
};
use crate::domains::finite_field::Zp64;
use crate::poly::PolyVariable;
use crate::printer::PrintFunction;
use crate::warn;
use crate::{
LicenseManager,
atom::{Atom, Symbol},
coefficient::Coefficient,
domains::{
finite_field::FiniteFieldCore,
float::{Complex, Float, Real},
},
};
pub(crate) const SYMBOLICA_MAGIC: u32 = 0x37871367;
pub(crate) const EXPORT_FORMAT_VERSION: u16 = 4;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct FiniteFieldIndex(pub(crate) usize);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) struct VariableListIndex(pub(crate) usize);
pub struct StateMap {
pub(crate) symbols: HashMap<u32, Symbol>,
pub(crate) finite_fields: HashMap<FiniteFieldIndex, FiniteFieldIndex>,
pub(crate) variables_lists: HashMap<u64, Arc<Vec<PolyVariable>>>,
}
pub trait HasStateMap {
fn get_state_map(&self) -> &StateMap;
}
impl HasStateMap for StateMap {
fn get_state_map(&self) -> &StateMap {
self
}
}
impl StateMap {
pub fn is_empty(&self) -> bool {
self.symbols.is_empty() && self.finite_fields.is_empty() && self.variables_lists.is_empty()
}
}
pub(crate) struct SymbolData {
pub(crate) name: String,
pub(crate) namespace: Cow<'static, str>,
pub(crate) file: Cow<'static, str>,
pub(crate) line: usize,
pub(crate) custom_normalization: Option<NormalizationFunction>,
pub(crate) custom_print: Option<PrintFunction>,
pub(crate) custom_derivative: Option<DerivativeFunction>,
pub(crate) custom_series: Option<Box<SeriesExpansionFunction>>,
pub(crate) custom_evaluation: Option<EvaluationInfo>,
pub(crate) aliases: Vec<std::string::String>,
pub(crate) tags: Vec<std::string::String>,
pub(crate) user_data: UserData,
}
fn builtin_constant_evaluation(symbol: Symbol) -> Option<EvaluationInfo> {
match symbol.get_id() {
Symbol::E_ID => Some(EvaluationInfo::constant(|_tags, prec| {
Ok(Complex::new(
Float::with_val(prec, 1).exp(),
Float::new(prec),
))
})),
Symbol::PI_ID => Some(EvaluationInfo::constant(|_tags, prec| {
Ok(Complex::new(
Float::with_val(prec, crate::domains::backend::float::Constant::Pi),
Float::new(prec),
))
})),
_ => None,
}
}
impl SymbolData {
fn default_from_symbol(name: &str) -> Self {
Self {
name: format!("symbolica::{}", name).into(),
file: file!().into(),
namespace: "symbolica".into(),
line: line!() as usize,
custom_normalization: None,
custom_print: None,
custom_derivative: None,
custom_series: None,
custom_evaluation: None,
aliases: vec![],
tags: vec![],
user_data: UserData::None,
}
}
}
pub struct StateInitializer {
init: fn(),
name: &'static str,
dependencies: &'static [&'static str],
}
impl StateInitializer {
pub const fn new(
name: &'static str,
init: fn(),
dependencies: &'static [&'static str],
) -> Self {
Self {
init,
name,
dependencies,
}
}
}
#[macro_export]
macro_rules! initialize {
($init:expr $(, $deps:expr)* $(,)?) => {
$crate::_inventory::submit! {
$crate::state::StateInitializer::new(
env!("CARGO_CRATE_NAME"),
$init,
&["symbolica", $($deps),*],
)
}
};
}
#[cfg(not(doctest))]
inventory::submit! {
StateInitializer::new(
"symbolica",
|| { },
&[]
)
}
inventory::collect!(StateInitializer);
static STATE: OnceLock<RwLock<State>> = OnceLock::new();
static STATE_INITIALIZER: Once = Once::new();
static ID_TO_STR: AppendOnlyVec<(Symbol, SymbolData)> = AppendOnlyVec::new();
static FINITE_FIELDS: AppendOnlyVec<Zp64> = AppendOnlyVec::new();
static VARIABLE_LISTS: AppendOnlyVec<Arc<Vec<PolyVariable>>> = AppendOnlyVec::new();
static SYMBOL_OFFSET: AtomicUsize = AtomicUsize::new(0);
thread_local!(
static WORKSPACE: Workspace = const { Workspace::new() }
);
thread_local!(
static RUNNING_STATE_INITIALIZER: Cell<bool> = const { Cell::new(false) }
);
pub struct State {
str_to_id: HashMap<String, Symbol>,
builtin_symbols: HashSet<String>,
}
impl Default for State {
fn default() -> Self {
Self::new()
}
}
impl State {
pub(crate) const ARG: Symbol =
Symbol::raw_fn(0, 0, false, false, false, false, false, false, false, false);
pub(crate) const COEFF: Symbol =
Symbol::raw_fn(1, 0, false, false, false, false, true, false, false, false);
pub(crate) const EXP: Symbol =
Symbol::raw_fn(2, 0, false, false, false, false, false, false, false, false);
pub(crate) const LOG: Symbol =
Symbol::raw_fn(3, 0, false, false, false, false, false, false, false, false);
pub(crate) const SIN: Symbol =
Symbol::raw_fn(4, 0, false, false, false, false, false, false, false, false);
pub(crate) const COS: Symbol =
Symbol::raw_fn(5, 0, false, false, false, false, false, false, false, false);
pub(crate) const SQRT: Symbol =
Symbol::raw_fn(6, 0, false, false, false, false, false, false, false, false);
pub(crate) const CONJ: Symbol =
Symbol::raw_fn(7, 0, false, false, false, false, false, false, false, false);
pub(crate) const ABS: Symbol =
Symbol::raw_fn(8, 0, false, false, false, false, false, true, false, true);
pub(crate) const IF: Symbol =
Symbol::raw_fn(9, 0, false, false, false, false, false, false, false, false);
pub(crate) const DERIVATIVE: Symbol = Symbol::raw_fn(
10, 0, false, false, false, false, false, false, false, false,
);
pub(crate) const E: Symbol =
Symbol::raw_fn(11, 0, false, false, false, false, true, true, false, true);
pub(crate) const PI: Symbol =
Symbol::raw_fn(12, 0, false, false, false, false, true, true, false, true);
pub(crate) const SEP: Symbol =
Symbol::raw_fn(13, 0, false, false, false, false, true, true, true, true);
pub(crate) const OPT: Symbol = Symbol::raw_fn(
14, 0, false, false, false, false, false, false, false, false,
);
pub(crate) const ALT: Symbol = Symbol::raw_fn(
15, 0, false, false, false, false, false, false, false, false,
);
pub const BUILTIN_SYMBOL_NAMES: [&'static str; 16] = [
"arg",
"coeff",
"exp",
"log",
"sin",
"cos",
"sqrt",
"conj",
"abs",
"if",
"der",
Symbol::E_STR,
Symbol::PI_STR,
Symbol::SEP_STR,
"opt",
"alt",
];
pub const BUILTIN_NAMES_AND_ALIASES: [&'static str; 18] = [
"arg",
"coeff",
"exp",
"log",
"sin",
"cos",
"sqrt",
"conj",
"abs",
"if",
"der",
Symbol::E_STR,
Symbol::PI_STR,
Symbol::SEP_STR,
"opt",
"alt",
"euler_e",
"pi",
];
pub const BUILTIN_SYMBOLS: [Symbol; 16] = [
Self::ARG,
Self::COEFF,
Self::EXP,
Self::LOG,
Self::SIN,
Self::COS,
Self::SQRT,
Self::CONJ,
Self::ABS,
Self::IF,
Self::DERIVATIVE,
Self::E,
Self::PI,
Self::SEP,
Self::OPT,
Self::ALT,
];
pub(crate) fn is_builtin_name<S: AsRef<str>>(&self, str: S) -> bool {
self.builtin_symbols.contains(str.as_ref())
}
pub fn is_builtin<S: AsRef<str>>(str: S) -> bool {
Self::get_global_state()
.read()
.unwrap()
.builtin_symbols
.contains(str.as_ref())
}
fn initialize_builtin_symbols(&mut self) {
let offset = SYMBOL_OFFSET.load(Ordering::Relaxed);
for (name, symbol) in Self::BUILTIN_SYMBOL_NAMES.iter().zip(Self::BUILTIN_SYMBOLS) {
let mut data = SymbolData::default_from_symbol(name);
data.custom_evaluation = builtin_constant_evaluation(symbol);
self.builtin_symbols.insert((*name).into());
if symbol == Self::E {
data.aliases = vec!["symbolica::euler_e".to_owned()];
self.builtin_symbols.insert("euler_e".into());
} else if symbol == Self::PI {
data.aliases = vec!["symbolica::pi".to_owned()];
self.builtin_symbols.insert("pi".into());
}
let id = ID_TO_STR.push((symbol, data)) - offset;
assert_eq!(symbol.get_id() as usize, id);
let data = &ID_TO_STR[id + offset].1;
self.str_to_id.insert(data.name.clone(), symbol);
for alias in &data.aliases {
self.str_to_id.insert(alias.clone().into(), symbol);
}
}
}
fn new() -> State {
let mut state = State {
str_to_id: HashMap::new(),
builtin_symbols: HashSet::new(),
};
state.initialize_builtin_symbols();
state
}
fn initialize_state() {
struct ReentryGuard;
impl ReentryGuard {
fn enter(running: &'static LocalKey<Cell<bool>>) -> Self {
running.with(|running| {
assert!(
!running.replace(true),
"nested state initializer execution is not supported"
);
});
Self
}
}
impl Drop for ReentryGuard {
fn drop(&mut self) {
RUNNING_STATE_INITIALIZER.with(|running| running.set(false));
}
}
STATE.get_or_init(|| RwLock::new(State::new()));
if RUNNING_STATE_INITIALIZER.with(|running| running.get()) {
return;
}
let mut initializing = false;
STATE_INITIALIZER.call_once(|| {
let _guard = ReentryGuard::enter(&RUNNING_STATE_INITIALIZER);
initializing = true;
#[cfg(test)]
{
STATE.get().unwrap().write().unwrap().initialize_test();
}
let mut initializers: Vec<_> =
inventory::iter::<StateInitializer>.into_iter().collect();
initializers.sort_by_key(|initializer| initializer.name);
let mut initializers_by_name = HashMap::with_capacity(initializers.len());
for initializer in &initializers {
if initializers_by_name
.insert(initializer.name, *initializer)
.is_some()
{
panic!(
"Multiple state initializers registered for crate `{}`",
initializer.name
);
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum VisitState {
Visiting,
Visited,
}
fn visit_initializer<'a>(
current: &'a StateInitializer,
initializer: &HashMap<&'static str, &'a StateInitializer>,
visit_state: &mut HashMap<&'static str, VisitState>,
ordered_initializers: &mut Vec<&'a StateInitializer>,
) {
match visit_state.get(current.name) {
Some(VisitState::Visited) => return,
Some(VisitState::Visiting) => {
panic!(
"Cyclic state initializer dependency involving crate `{}`",
current.name
);
}
None => {}
}
visit_state.insert(current.name, VisitState::Visiting);
for dependency in current.dependencies {
let dependency_initializer = initializer.get(dependency).unwrap_or_else(|| {
panic!(
"State initializer for crate `{}` depends on missing crate `{}`",
current.name, dependency
)
});
visit_initializer(
dependency_initializer,
initializer,
visit_state,
ordered_initializers,
);
}
visit_state.insert(current.name, VisitState::Visited);
ordered_initializers.push(current);
}
let mut ordered_initializers = Vec::with_capacity(initializers.len());
let mut visit_state = HashMap::with_capacity(initializers.len());
for initializer in &initializers {
visit_initializer(
initializer,
&initializers_by_name,
&mut visit_state,
&mut ordered_initializers,
);
}
for initializer in ordered_initializers {
(initializer.init)();
}
});
if !initializing {
LicenseManager::check();
}
}
#[inline]
pub(crate) fn get_global_state() -> &'static RwLock<State> {
Self::initialize_state();
STATE.get().unwrap()
}
#[cfg(test)]
fn initialize_test(&mut self) {
use crate::{atom::SymbolAttribute, wrap_symbol};
for i in 0..30 {
let _ = self.get_symbol(wrap_symbol!(format!("symbolica::v{}", i)));
}
for i in 0..30 {
let _ = self.get_symbol(wrap_symbol!(format!("symbolica::f{}", i)));
}
for i in 0..5 {
let _ = self.get_symbol_with_attributes(
wrap_symbol!(format!("symbolica::fs{}", i)),
&[SymbolAttribute::Symmetric],
None,
None,
None,
None,
None,
vec![],
vec![],
None,
);
}
for i in 0..5 {
let _ = self.get_symbol_with_attributes(
wrap_symbol!(format!("symbolica::fc{}", i)),
&[SymbolAttribute::Cyclesymmetric],
None,
None,
None,
None,
None,
vec![],
vec![],
None,
);
}
for i in 0..5 {
let _ = self.get_symbol_with_attributes(
wrap_symbol!(format!("symbolica::fa{}", i)),
&[SymbolAttribute::Antisymmetric],
None,
None,
None,
None,
None,
vec![],
vec![],
None,
);
}
for i in 0..5 {
let _ = self.get_symbol_with_attributes(
wrap_symbol!(format!("symbolica::fl{}", i)),
&[SymbolAttribute::Linear],
None,
None,
None,
None,
None,
vec![],
vec![],
None,
);
}
for i in 0..5 {
let _ = self.get_symbol_with_attributes(
wrap_symbol!(format!("symbolica::fsl{}", i)),
&[SymbolAttribute::Symmetric, SymbolAttribute::Linear],
None,
None,
None,
None,
None,
vec![],
vec![],
None,
);
}
}
pub unsafe fn reset() {
let mut state = Self::get_global_state().write().unwrap();
state.str_to_id.clear();
SYMBOL_OFFSET.store(ID_TO_STR.len(), Ordering::Relaxed);
state.initialize_builtin_symbols();
#[cfg(test)]
{
state.initialize_test();
}
}
#[inline(always)]
#[allow(dead_code)]
pub(crate) unsafe fn symbol_from_id(id: u32) -> Symbol {
if ID_TO_STR.len() == 0 {
Self::initialize_state();
}
ID_TO_STR[id as usize].0
}
pub fn symbol_iter() -> impl Iterator<Item = (Symbol, &'static str)> {
if ID_TO_STR.len() == 0 {
Self::initialize_state();
}
ID_TO_STR
.iter()
.skip(SYMBOL_OFFSET.load(Ordering::Relaxed))
.map(|s| (s.0, s.1.name.as_str()))
}
pub(crate) fn is_fixed_builtin(id: Symbol) -> bool {
id.get_id() < Self::BUILTIN_SYMBOL_NAMES.len() as u32
}
pub(crate) fn fetch_symbol(&self, name: &str) -> Option<Symbol> {
self.str_to_id.get(name).cloned()
}
pub(crate) fn get_next_symbol_index(&self) -> u32 {
let offset = SYMBOL_OFFSET.load(Ordering::Relaxed);
(ID_TO_STR.len() - offset) as u32
}
pub(crate) fn get_wildcard_level(str: &str) -> u8 {
let mut wildcard_level = 0;
for x in str.chars().rev() {
if x != '_' {
break;
}
wildcard_level += 1;
}
wildcard_level
}
pub(crate) fn get_symbol(&mut self, name: NamespacedSymbol) -> Result<Symbol, String> {
match self.str_to_id.entry(name.symbol.into()) {
Entry::Occupied(o) => Ok(*o.get()),
Entry::Vacant(v) => {
let offset = SYMBOL_OFFSET.load(Ordering::Relaxed);
if ID_TO_STR.len() - offset == u32::MAX as usize - 1 {
panic!("Too many variables defined");
}
let wildcard_level = State::get_wildcard_level(v.key());
let id = ID_TO_STR.len() - offset;
let new_symbol = Symbol::raw_var(id as u32, wildcard_level);
let id_ret = ID_TO_STR.push((
new_symbol,
SymbolData {
name: v.key().clone(),
file: name.file,
namespace: name.namespace,
line: name.line,
aliases: vec![],
custom_normalization: None,
custom_print: None,
custom_derivative: None,
custom_series: None,
custom_evaluation: None,
tags: vec![],
user_data: UserData::None,
},
)) - offset;
assert_eq!(id, id_ret);
v.insert(new_symbol);
Ok(new_symbol)
}
}
}
pub(crate) fn get_state_mut() -> RwLockWriteGuard<'static, State> {
Self::initialize_state();
STATE.get().unwrap().write().unwrap()
}
pub(crate) fn get_symbol_with_attributes(
&mut self,
name: NamespacedSymbol,
attributes: &[SymbolAttribute],
normalization_function: Option<NormalizationFunction>,
print_function: Option<PrintFunction>,
derivative_function: Option<DerivativeFunction>,
series_function: Option<Box<SeriesExpansionFunction>>,
evaluation_function: Option<EvaluationInfo>,
tags: Vec<std::string::String>,
mut aliases: Vec<std::string::String>,
user_data: Option<UserData>,
) -> Result<Symbol, String> {
for alias in &mut aliases {
if !alias.contains("::") {
*alias = format!("{}::{}", name.namespace, alias);
} else if !alias.starts_with(name.namespace.as_ref()) {
return Err(format!(
"Alias {alias} defined in different namespace from main symbol namespace {}",
name.namespace
)
.into());
}
}
match self.str_to_id.entry(name.symbol.into()) {
Entry::Occupied(o) => {
let r = *o.get();
let new_id = Symbol::raw_fn(
r.get_id(),
r.get_wildcard_level(),
attributes.contains(&SymbolAttribute::Symmetric),
attributes.contains(&SymbolAttribute::Antisymmetric),
attributes.contains(&SymbolAttribute::Cyclesymmetric),
attributes.contains(&SymbolAttribute::Linear),
attributes.contains(&SymbolAttribute::Scalar),
attributes.contains(&SymbolAttribute::Real),
attributes.contains(&SymbolAttribute::Integer),
attributes.contains(&SymbolAttribute::Positive),
);
if r.get_wildcard_level() == new_id.get_wildcard_level()
&& r.is_symmetric() == new_id.is_symmetric()
&& r.is_antisymmetric() == new_id.is_antisymmetric()
&& r.is_cyclesymmetric() == new_id.is_cyclesymmetric()
&& r.is_linear() == new_id.is_linear()
&& r.is_scalar() == new_id.is_scalar()
&& r.is_real() == new_id.is_real()
&& r.is_integer() == new_id.is_integer()
&& r.is_positive() == new_id.is_positive()
&& normalization_function.is_none()
&& print_function.is_none()
&& derivative_function.is_none()
&& series_function.is_none()
&& evaluation_function.is_none()
&& tags == r.get_tags()
&& aliases == r.get_aliases()
&& user_data.as_ref().unwrap_or(&UserData::None)
== &ID_TO_STR[r.get_id() as usize].1.user_data
{
Ok(r)
} else {
let data = &ID_TO_STR[r.get_id() as usize].1;
let mut diff_attr = String::new();
if r.is_antisymmetric() != new_id.is_antisymmetric() {
diff_attr.push_str(&format!(
"\t- antisymmetric: {} vs {}\n",
r.is_antisymmetric(),
new_id.is_antisymmetric()
));
}
if r.is_symmetric() != new_id.is_symmetric() {
diff_attr.push_str(&format!(
"\t- symmetric: {} vs {}\n",
r.is_symmetric(),
new_id.is_symmetric()
));
}
if r.is_cyclesymmetric() != new_id.is_cyclesymmetric() {
diff_attr.push_str(&format!(
"\t- cyclesymmetric: {} vs {}\n",
r.is_cyclesymmetric(),
new_id.is_cyclesymmetric()
));
}
if r.is_linear() != new_id.is_linear() {
diff_attr.push_str(&format!(
"\t- linear: {} vs {}\n",
r.is_linear(),
new_id.is_linear()
));
}
if r.is_scalar() != new_id.is_scalar() {
diff_attr.push_str(&format!(
"\t- scalar: {} vs {}\n",
r.is_scalar(),
new_id.is_scalar()
));
}
if r.is_real() != new_id.is_real() {
diff_attr.push_str(&format!(
"\t- real: {} vs {}\n",
r.is_real(),
new_id.is_real()
));
}
if r.is_integer() != new_id.is_integer() {
diff_attr.push_str(&format!(
"\t- integer: {} vs {}\n",
r.is_integer(),
new_id.is_integer()
));
}
if r.is_positive() != new_id.is_positive() {
diff_attr.push_str(&format!(
"\t- positive: {} vs {}\n",
r.is_positive(),
new_id.is_positive()
));
}
if tags != r.get_tags() {
diff_attr.push_str(&format!(
"\t- tags: {:?} vs {:?}\n",
r.get_tags(),
tags
));
}
if aliases != r.get_aliases() {
diff_attr.push_str(&format!(
"\t- aliases: {:?} vs {:?}\n",
r.get_aliases(),
aliases
));
}
if normalization_function.is_some() {
diff_attr.push_str("\t- new normalization function specified.\n");
}
if print_function.is_some() {
diff_attr.push_str("\t- new print function specified.\n");
}
if derivative_function.is_some() {
diff_attr.push_str("\t- new derivative function specified.\n");
}
if series_function.is_some() {
diff_attr.push_str("\t- new series function specified.\n");
}
if evaluation_function.is_some() {
diff_attr.push_str("\t- new evaluation function specified.\n");
}
if user_data.as_ref().unwrap_or(&UserData::None) != &data.user_data {
diff_attr.push_str(&format!(
"\t- new user data specified: {:?} vs {:?}\n",
data.user_data, user_data
));
}
if data.file.is_empty() {
Err(format!(
"Symbol {} redefined with new attributes:\n{}",
data.name, diff_attr
)
.into())
} else {
Err(format!("Symbol {} redefined with new attributes:\n{}\nThe first definition occurred here: {}:{}.", data.name, diff_attr, data.file, data.line).into())
}
}
}
Entry::Vacant(v) => {
let offset = SYMBOL_OFFSET.load(Ordering::Relaxed);
if ID_TO_STR.len() - offset == u32::MAX as usize - 1 {
panic!("Too many variables defined");
}
let id = ID_TO_STR.len() - offset;
let wildcard_level = State::get_wildcard_level(v.key());
let new_symbol = Symbol::raw_fn(
id as u32,
wildcard_level,
attributes.contains(&SymbolAttribute::Symmetric),
attributes.contains(&SymbolAttribute::Antisymmetric),
attributes.contains(&SymbolAttribute::Cyclesymmetric),
attributes.contains(&SymbolAttribute::Linear),
attributes.contains(&SymbolAttribute::Scalar),
attributes.contains(&SymbolAttribute::Real),
attributes.contains(&SymbolAttribute::Integer),
attributes.contains(&SymbolAttribute::Positive),
);
let id_ret = ID_TO_STR.push((
new_symbol,
SymbolData {
name: v.key().clone(),
file: name.file,
namespace: name.namespace.clone(),
line: name.line,
custom_normalization: normalization_function,
custom_print: print_function,
custom_derivative: derivative_function,
custom_series: series_function,
custom_evaluation: evaluation_function,
tags,
aliases: aliases.clone(),
user_data: user_data.unwrap_or(UserData::None),
},
)) - offset;
assert_eq!(id, id_ret);
v.insert(new_symbol);
if new_symbol.get_namespace() == "symbolica" {
self.builtin_symbols
.insert(new_symbol.get_stripped_name().into());
}
for alias in aliases {
match self.str_to_id.entry(alias.into()) {
Entry::Occupied(o) => {
let old_symbol = o.get();
let old_data = old_symbol.get_global_data();
if old_data.file.is_empty() {
return Err(
format!("Alias {} already defined before", o.key()).into()
);
} else {
return Err(format!(
"Alias {} already defined here: {}:{}.",
old_data.name, old_data.file, old_data.line
)
.into());
}
}
Entry::Vacant(v) => {
if new_symbol.get_namespace() == "symbolica" {
self.builtin_symbols.insert(v.key()[11..].into());
}
v.insert(new_symbol);
}
}
}
Ok(new_symbol)
}
}
}
#[inline]
pub(crate) fn get_name(id: Symbol) -> &'static str {
if ID_TO_STR.len() == 0 {
Self::initialize_state();
}
&ID_TO_STR[id.get_id() as usize + SYMBOL_OFFSET.load(Ordering::Relaxed)]
.1
.name
}
#[inline]
pub(crate) fn get_symbol_namespace(id: Symbol) -> &'static str {
if ID_TO_STR.len() == 0 {
Self::initialize_state();
}
ID_TO_STR[id.get_id() as usize + SYMBOL_OFFSET.load(Ordering::Relaxed)]
.1
.namespace
.as_ref()
}
#[inline]
pub(crate) fn get_symbol_data(id: Symbol) -> &'static SymbolData {
if ID_TO_STR.len() == 0 {
Self::initialize_state();
}
&ID_TO_STR[id.get_id() as usize + SYMBOL_OFFSET.load(Ordering::Relaxed)].1
}
#[inline]
pub(crate) fn get_normalization_function(id: Symbol) -> Option<&'static NormalizationFunction> {
if ID_TO_STR.len() == 0 {
Self::initialize_state();
}
ID_TO_STR[id.get_id() as usize + SYMBOL_OFFSET.load(Ordering::Relaxed)]
.1
.custom_normalization
.as_ref()
}
pub(crate) fn get_finite_field(fi: FiniteFieldIndex) -> &'static Zp64 {
&FINITE_FIELDS[fi.0]
}
pub(crate) fn get_or_insert_finite_field(f: Zp64) -> FiniteFieldIndex {
Self::get_global_state()
.write()
.unwrap()
.get_or_insert_finite_field_impl(f)
}
pub(crate) fn get_or_insert_finite_field_impl(&mut self, f: Zp64) -> FiniteFieldIndex {
for (i, f2) in FINITE_FIELDS.iter().enumerate() {
if f.get_prime() == f2.get_prime() {
return FiniteFieldIndex(i);
}
}
let index = FINITE_FIELDS.push(f);
FiniteFieldIndex(index)
}
pub(crate) fn get_variable_list(fi: VariableListIndex) -> Arc<Vec<PolyVariable>> {
VARIABLE_LISTS[fi.0].clone()
}
pub(crate) fn get_or_insert_variable_list(f: Arc<Vec<PolyVariable>>) -> VariableListIndex {
Self::get_global_state()
.write()
.unwrap()
.get_or_insert_variable_list_impl(f)
}
pub(crate) fn get_or_insert_variable_list_impl(
&mut self,
f: Arc<Vec<PolyVariable>>,
) -> VariableListIndex {
for (i, f2) in VARIABLE_LISTS.iter().enumerate() {
if f2 == &f {
return VariableListIndex(i);
}
}
let index = VARIABLE_LISTS.push(f);
VariableListIndex(index)
}
#[inline(always)]
pub fn export<W: Write>(dest: &mut W) -> Result<(), std::io::Error> {
if ID_TO_STR.len() == 0 {
Self::initialize_state();
}
dest.write_u32::<LittleEndian>(SYMBOLICA_MAGIC)?;
dest.write_u16::<LittleEndian>(EXPORT_FORMAT_VERSION)?;
dest.write_u64::<LittleEndian>(
ID_TO_STR.len() as u64 - SYMBOL_OFFSET.load(Ordering::Relaxed) as u64,
)?;
for (s, _) in State::symbol_iter() {
s.export(dest)?;
}
dest.write_u64::<LittleEndian>(FINITE_FIELDS.len() as u64)?;
for x in FINITE_FIELDS.iter() {
dest.write_u64::<LittleEndian>(x.get_prime())?;
}
dest.write_u64::<LittleEndian>(VARIABLE_LISTS.len() as u64)?;
for x in VARIABLE_LISTS.iter() {
dest.write_u64::<LittleEndian>(x.len() as u64)?;
for y in x.iter() {
match y {
PolyVariable::Symbol(s) => {
dest.write_u8(0)?;
dest.write_u32::<LittleEndian>(s.get_id())?;
}
PolyVariable::Temporary(u) => {
dest.write_u8(1)?;
dest.write_u64::<LittleEndian>(*u as u64)?;
}
PolyVariable::Function(v, t) => {
dest.write_u8(2)?;
dest.write_u32::<LittleEndian>(v.get_id())?;
t.as_view().write(dest.by_ref())?;
}
PolyVariable::Power(t) => {
dest.write_u8(3)?;
t.as_view().write(dest.by_ref())?;
}
}
}
}
Ok(())
}
#[inline(always)]
pub fn import<R: Read>(
source: &mut R,
conflict_fn: Option<Box<dyn Fn(&str) -> String>>,
) -> Result<StateMap, std::io::Error> {
let magic = source.read_u32::<LittleEndian>()?;
if magic != SYMBOLICA_MAGIC {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"Invalid magic number: the file is not exported from Symbolica",
));
}
let version = source.read_u16::<LittleEndian>()?;
if version != EXPORT_FORMAT_VERSION {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"Invalid export format version: expected {} but got {}",
EXPORT_FORMAT_VERSION, version
),
));
}
let mut state_map = StateMap {
symbols: HashMap::default(),
finite_fields: HashMap::default(),
variables_lists: HashMap::default(),
};
let n_symbols = source.read_u64::<LittleEndian>()?;
for x in 0..n_symbols {
let (mut name, namespace, attributes, tags, extra_data, aliases, is_exportable) =
Symbol::import_impl(source)?;
loop {
let num_symbols = ID_TO_STR.len();
match SymbolBuilder::new(NamespacedSymbol {
symbol: name.clone().into(),
namespace: namespace.to_string().into(),
file: "".into(),
line: 0,
})
.with_attributes(attributes.clone())
.with_tags(tags.clone())
.with_user_data(extra_data.clone())
.with_aliases(aliases.clone())
.build()
{
Ok(id) => {
if !is_exportable && num_symbols != ID_TO_STR.len() {
warn!(
"Imported symbol {name} was previously defined with user-defined functions, but the imported version does not have any."
);
}
if x as u32 != id.get_id() {
state_map.symbols.insert(x as u32, id);
}
break;
}
Err(e) => {
if let Some(f) = &conflict_fn {
let new_name = f(&name);
let mut old_wildcard_level = 0;
for x in name.chars().rev() {
if x != '_' {
break;
}
old_wildcard_level += 1;
}
let mut new_wildcard_level = 0;
for x in new_name.chars().rev() {
if x != '_' {
break;
}
new_wildcard_level += 1;
}
if old_wildcard_level == new_wildcard_level {
name = new_name.to_string();
}
} else {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("Symbol conflict: {e}"),
));
}
}
}
}
}
let n_finite_fields = source.read_u64::<LittleEndian>()?;
for x in 0..n_finite_fields {
let prime = source.read_u64::<LittleEndian>()?;
let id = State::get_or_insert_finite_field(Zp64::new(prime));
if x != id.0 as u64 {
state_map
.finite_fields
.insert(FiniteFieldIndex(x as usize), id);
}
}
let n_variable_lists = source.read_u64::<LittleEndian>()?;
for x in 0..n_variable_lists {
let n_vars = source.read_u64::<LittleEndian>()?;
let mut variables = vec![];
for _ in 0..n_vars {
match source.read_u8()? {
0 => {
let id = source.read_u32::<LittleEndian>()?;
if let Some(new_id) = state_map.symbols.get(&id) {
variables.push(PolyVariable::Symbol(*new_id));
} else {
variables.push(PolyVariable::Symbol(ID_TO_STR[id as usize].0))
}
}
1 => {
let u = source.read_u64::<LittleEndian>()?;
variables.push(PolyVariable::Temporary(u as usize))
}
2 => {
let id = source.read_u32::<LittleEndian>()?;
let symb = if let Some(new_id) = state_map.symbols.get(&id) {
*new_id
} else {
ID_TO_STR[id as usize].0
};
let mut f = Atom::new();
f.read(&mut *source)?;
let f_r = f.as_view().rename(&state_map);
variables.push(PolyVariable::Function(symb, f_r));
}
3 => {
let mut f = Atom::new();
f.read(&mut *source)?;
let f_r = f.as_view().rename(&state_map);
variables.push(PolyVariable::Power(f_r));
}
_ => {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"Invalid variable type",
));
}
}
}
let vars = Arc::new(variables);
let new_id = State::get_or_insert_variable_list(vars.clone());
if x != new_id.0 as u64 {
state_map.variables_lists.insert(x, vars);
}
}
Ok(state_map)
}
}
pub struct Workspace {
atom_buffer: RefCell<Vec<Atom>>,
}
impl Workspace {
const ATOM_BUFFER_MAX: usize = 30;
const ATOM_CACHE_SIZE_MAX: usize = 20_000_000;
const fn new() -> Self {
Workspace {
atom_buffer: RefCell::new(Vec::new()),
}
}
#[inline]
pub fn get_local() -> &'static LocalKey<Workspace> {
LicenseManager::check();
&WORKSPACE
}
#[inline]
pub fn new_atom(&self) -> RecycledAtom {
if let Ok(mut a) = self.atom_buffer.try_borrow_mut() {
if let Some(b) = a.pop() {
b.into()
} else {
Atom::default().into()
}
} else {
Atom::default().into() }
}
#[inline]
pub fn new_var(&self, id: Symbol) -> RecycledAtom {
let mut owned = self.new_atom();
owned.to_var(id);
owned
}
#[inline]
pub fn new_num<T: Into<Coefficient>>(&self, num: T) -> RecycledAtom {
let mut owned = self.new_atom();
owned.to_num(num);
owned
}
pub fn return_atom(&self, atom: Atom) {
if let Ok(mut a) = self.atom_buffer.try_borrow_mut() {
a.push(atom);
}
}
}
#[derive(PartialEq, Eq, Debug, Hash, Clone)]
pub struct RecycledAtom(Atom);
impl From<Atom> for RecycledAtom {
fn from(a: Atom) -> Self {
RecycledAtom(a)
}
}
impl std::fmt::Display for RecycledAtom {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.0.fmt(f)
}
}
impl Default for RecycledAtom {
fn default() -> Self {
Self::new()
}
}
impl RecycledAtom {
#[inline]
pub fn new() -> RecycledAtom {
Workspace::get_local().with(|ws| ws.new_atom())
}
pub fn wrap(atom: Atom) -> RecycledAtom {
RecycledAtom(atom)
}
#[inline]
pub fn new_var(id: Symbol) -> RecycledAtom {
let mut owned = Self::new();
owned.to_var(id);
owned
}
#[inline]
pub fn new_num<T: Into<Coefficient>>(num: T) -> RecycledAtom {
let mut owned = Self::new();
owned.to_num(num);
owned
}
pub fn into_inner(mut self) -> Atom {
std::mem::replace(&mut self.0, Atom::Zero)
}
}
impl Deref for RecycledAtom {
type Target = Atom;
fn deref(&self) -> &Atom {
&self.0
}
}
impl DerefMut for RecycledAtom {
fn deref_mut(&mut self) -> &mut Atom {
&mut self.0
}
}
impl AsRef<Atom> for RecycledAtom {
fn as_ref(&self) -> &Atom {
self.deref()
}
}
impl Drop for RecycledAtom {
#[inline]
fn drop(&mut self) {
if let Atom::Zero = self.0 {
return;
}
if self.0.get_capacity() > Workspace::ATOM_CACHE_SIZE_MAX {
return;
}
let _ = WORKSPACE.try_with(
#[inline(always)]
|ws| {
if let Ok(mut a) = ws.atom_buffer.try_borrow_mut()
&& a.len() < Workspace::ATOM_BUFFER_MAX
{
a.push(std::mem::replace(&mut self.0, Atom::Zero));
}
},
);
}
}
#[cfg(test)]
mod tests {
use crate::{
atom::{Atom, AtomCore, AtomView, InlineVar, Symbol},
parse, symbol,
};
use super::State;
#[test]
fn state_export_import() {
let mut export = vec![];
State::export(&mut export).unwrap();
let i = State::import(&mut export.as_slice(), None).unwrap();
assert!(i.is_empty());
}
#[test]
fn export_symbol_data() {
let s = symbol!(
"symbolica::symbol_data::a",
data = crate::state::UserData::Atom(parse!("z"))
);
let s1 = s.to_atom();
let mut export = vec![];
s1.export(&mut export).unwrap();
let a = Atom::import(&mut export.as_slice(), None)
.unwrap()
.get_symbol()
.unwrap();
assert_eq!(a.get_data(), &crate::state::UserData::Atom(parse!("z")));
}
#[test]
fn custom_normalization() {
let _real_log = symbol!(
"custom_normalization_real_log",
norm = |input, out| {
if let AtomView::Fun(f) = input {
if f.get_nargs() == 1 {
let arg = f.iter().next().unwrap();
if let AtomView::Pow(p) = arg {
let (b, e) = p.get_base_exp();
if b == InlineVar::new(Symbol::E).as_view() {
out.set_from_view(&e);
}
}
}
}
}
);
let e = parse!("custom_normalization_real_log(exp(x))");
assert_eq!(e, parse!("x"));
}
}