1use std::{any::Any, sync::Arc};
4
5use sim_kernel::{
6 AbiVersion, Args, Callable, ClassRef, Cx, DefaultFactory, Dependency, Export, Expr, Factory,
7 Lib, LibManifest, LibTarget, Linker, Object, RawArgs, Result, Symbol, Value, Version,
8};
9
10use super::{agent_fixtures, runtime};
11
12pub fn stats_disparate_impact_claim_symbol() -> Symbol {
14 Symbol::qualified("stats", "disparate-impact-claim")
15}
16
17pub fn stats_mean_claim_symbol() -> Symbol {
19 Symbol::qualified("stats", "mean-claim")
20}
21
22pub fn stats_variance_claim_symbol() -> Symbol {
24 Symbol::qualified("stats", "variance-claim")
25}
26
27pub fn stats_entropy_claim_symbol() -> Symbol {
29 Symbol::qualified("stats", "entropy-claim")
30}
31
32pub fn stats_claims_symbol() -> Symbol {
34 Symbol::qualified("stats", "claims")
35}
36
37pub fn stats_kmeans_symbol() -> Symbol {
39 Symbol::qualified("stats", "kmeans")
40}
41
42pub fn stats_gmm_symbol() -> Symbol {
44 Symbol::qualified("stats", "gmm")
45}
46
47pub fn stats_exact_binary_interval_symbol() -> Symbol {
49 Symbol::qualified("stats", "exact-binary-interval")
50}
51pub fn stats_paired_bootstrap_symbol() -> Symbol {
53 Symbol::qualified("stats", "paired-bootstrap")
54}
55pub fn stats_clustered_bootstrap_symbol() -> Symbol {
57 Symbol::qualified("stats", "clustered-bootstrap")
58}
59pub fn stats_registered_look_symbol() -> Symbol {
61 Symbol::qualified("stats", "registered-look-interval")
62}
63pub fn stats_isotonic_symbol() -> Symbol {
65 Symbol::qualified("stats", "isotonic")
66}
67
68fn function_symbols() -> [Symbol; 16] {
69 [
70 stats_mean_claim_symbol(),
71 stats_variance_claim_symbol(),
72 stats_entropy_claim_symbol(),
73 stats_disparate_impact_claim_symbol(),
74 stats_claims_symbol(),
75 stats_kmeans_symbol(),
76 stats_gmm_symbol(),
77 stats_exact_binary_interval_symbol(),
78 stats_paired_bootstrap_symbol(),
79 stats_clustered_bootstrap_symbol(),
80 stats_registered_look_symbol(),
81 stats_isotonic_symbol(),
82 agent_fixtures::fixture_symbols()[0].clone(),
83 agent_fixtures::fixture_symbols()[1].clone(),
84 agent_fixtures::fixture_symbols()[2].clone(),
85 agent_fixtures::fixture_symbols()[3].clone(),
86 ]
87}
88
89#[derive(Clone)]
90struct StatsFunction {
91 symbol: Symbol,
92}
93
94impl Object for StatsFunction {
95 fn display(&self, _cx: &mut Cx) -> Result<String> {
96 Ok(format!("#<function {}>", self.symbol))
97 }
98
99 fn as_any(&self) -> &dyn Any {
100 self
101 }
102}
103
104impl sim_kernel::ObjectCompat for StatsFunction {
105 fn class(&self, cx: &mut Cx) -> Result<ClassRef> {
106 if let Some(value) = cx
107 .registry()
108 .class_by_symbol(&Symbol::qualified("core", "Function"))
109 {
110 return Ok(value.clone());
111 }
112 DefaultFactory.class_stub(
113 sim_kernel::CORE_FUNCTION_CLASS_ID,
114 Symbol::qualified("core", "Function"),
115 )
116 }
117
118 fn as_expr(&self, _cx: &mut Cx) -> Result<Expr> {
119 Ok(Expr::Symbol(self.symbol.clone()))
120 }
121
122 fn as_callable(&self) -> Option<&dyn Callable> {
123 Some(self)
124 }
125}
126
127impl Callable for StatsFunction {
128 fn call(&self, cx: &mut Cx, args: Args) -> Result<Value> {
129 if self.symbol == stats_mean_claim_symbol() {
130 return runtime::call_stats_mean_claim(cx, args);
131 }
132 if self.symbol == stats_variance_claim_symbol() {
133 return runtime::call_stats_variance_claim(cx, args);
134 }
135 if self.symbol == stats_entropy_claim_symbol() {
136 return runtime::call_stats_entropy_claim(cx, args);
137 }
138 if self.symbol == stats_disparate_impact_claim_symbol() {
139 return runtime::call_stats_disparate_impact_claim(cx, args);
140 }
141 if self.symbol == stats_claims_symbol() {
142 return runtime::call_stats_claims(cx, args);
143 }
144 if self.symbol == stats_kmeans_symbol() {
145 return super::runtime_clustering::call_kmeans_values(cx, args.into_vec());
146 }
147 if self.symbol == stats_gmm_symbol() {
148 return super::runtime_clustering::call_gmm_values(cx, args.into_vec());
149 }
150 if super::runtime_decision::is_symbol(&self.symbol) {
151 return super::runtime_decision::call(cx, &self.symbol, args.into_vec());
152 }
153 if let Some(value) = agent_fixtures::call_fixture(cx, &self.symbol, args)? {
154 return Ok(value);
155 }
156 unreachable!("unregistered stats function {}", self.symbol)
157 }
158
159 fn call_exprs(&self, cx: &mut Cx, args: RawArgs) -> Result<Value> {
160 if self.symbol == stats_kmeans_symbol() {
161 return super::runtime_clustering::call_kmeans_exprs(cx, args.into_exprs());
162 }
163 if self.symbol == stats_gmm_symbol() {
164 return super::runtime_clustering::call_gmm_exprs(cx, args.into_exprs());
165 }
166 let values = args
167 .into_exprs()
168 .into_iter()
169 .map(|expr| cx.eval_expr(expr))
170 .collect::<Result<Vec<_>>>()?;
171 self.call(cx, Args::new(values))
172 }
173}
174
175pub struct StatsNumbersLib;
177
178impl StatsNumbersLib {
179 pub fn new() -> Self {
181 Self
182 }
183}
184
185impl Default for StatsNumbersLib {
186 fn default() -> Self {
187 Self::new()
188 }
189}
190
191impl Lib for StatsNumbersLib {
192 fn manifest(&self) -> LibManifest {
193 LibManifest {
194 id: Symbol::qualified("numbers", "stats"),
195 version: Version(env!("CARGO_PKG_VERSION").to_owned()),
196 abi: AbiVersion { major: 0, minor: 1 },
197 target: LibTarget::HostRegistered,
198 requires: Vec::<Dependency>::new(),
199 capabilities: Vec::new(),
200 exports: function_symbols()
201 .into_iter()
202 .map(|symbol| Export::Function {
203 symbol,
204 function_id: None,
205 })
206 .collect(),
207 }
208 }
209
210 fn load(&self, _cx: &mut sim_kernel::LoadCx, linker: &mut Linker<'_>) -> Result<()> {
211 for symbol in function_symbols() {
212 linker.function_value(
213 symbol.clone(),
214 DefaultFactory.opaque(Arc::new(StatsFunction { symbol }))?,
215 )?;
216 }
217 Ok(())
218 }
219}