1mod specific_price;
2
3#[cfg(feature = "schemars")]
4use alloc::borrow::ToOwned;
5use alloc::vec::Vec;
6#[cfg(any(feature = "borsh", feature = "schemars"))]
7use alloc::{format, string::ToString};
8use core::marker::PhantomData;
9
10use super::Aggregate;
11use crate::proxy::WeightedSource;
12use crate::Price;
13use specific_price::SpecificPrice;
14
15fn median<T>(sorted_weighted_items: &[(T, u32)]) -> (usize, usize) {
24 let total_weight = sorted_weighted_items
25 .iter()
26 .map(|(_, weight)| u128::from(*weight))
27 .sum::<u128>();
28
29 if total_weight == 0 {
30 let high = sorted_weighted_items.len() / 2;
31 let low = sorted_weighted_items.len().saturating_sub(1) / 2;
32 return (low, high);
33 }
34
35 let low_target = total_weight.div_ceil(2);
36 let high_target = total_weight / 2 + 1;
37 let find_target = |target| {
38 let mut cumulative = 0u128;
39 let Some(index) = sorted_weighted_items.iter().position(|(_, weight)| {
40 cumulative += u128::from(*weight);
41 cumulative >= target
42 }) else {
43 unreachable!("positive total weight must reach its target");
44 };
45 index
46 };
47
48 (find_target(low_target), find_target(high_target))
49}
50
51#[cfg(kani)]
52pub(crate) fn median_indices_for_proof<T>(sorted_weighted_items: &[(T, u32)]) -> (usize, usize) {
53 median(sorted_weighted_items)
54}
55
56pub trait MedianVariant {
57 fn median<T>(sorted_weighted_items: &[(T, u32)]) -> usize;
58}
59
60serialize! {
61 #[derive(Debug, Clone, PartialEq, Eq)]
62 pub struct Low;
63}
64
65impl MedianVariant for Low {
66 fn median<T>(sorted_weighted_items: &[(T, u32)]) -> usize {
67 let (lo, hi) = median(sorted_weighted_items);
68 lo.min(hi)
69 }
70}
71
72serialize! {
73 #[derive(Debug, Clone, PartialEq, Eq)]
74 pub struct High;
75}
76
77impl MedianVariant for High {
78 fn median<T>(sorted_weighted_items: &[(T, u32)]) -> usize {
79 let (lo, hi) = median(sorted_weighted_items);
80 lo.max(hi)
81 }
82}
83
84pub type MedianLow<S> = Median<Low, S>;
85pub type MedianHigh<S> = Median<High, S>;
86
87serialize! {
88 #[derive(Debug, Clone, PartialEq, Eq)]
89 pub struct Median<V: MedianVariant, S> {
90 #[cfg_attr(feature = "serde", serde(skip))]
91 #[cfg_attr(feature = "borsh", borsh(skip))]
92 _variant: PhantomData<V>,
93 pub sources: Vec<WeightedSource<S>>,
94 pub min_sources: u32,
99 }
100}
101
102impl<V: MedianVariant, S> Median<V, S> {
103 pub fn new(sources: impl IntoIterator<Item = WeightedSource<S>>) -> Self {
104 Self {
105 _variant: PhantomData,
106 sources: sources.into_iter().collect(),
107 min_sources: 1,
108 }
109 }
110}
111
112impl<V: MedianVariant, S> Aggregate<S> for Median<V, S> {
113 fn aggregate<I>(&self, prices: I) -> Result<Price, super::Error>
114 where
115 I: IntoIterator<Item = Option<Price>>,
116 I::IntoIter: ExactSizeIterator<Item = Option<Price>>,
117 {
118 let prices = prices.into_iter();
119 let actual = prices.len();
120
121 if actual != self.sources.len() {
122 return Err(super::Error::LengthMismatch {
123 expected: self.sources.len(),
124 actual,
125 });
126 }
127
128 let mut values = Vec::with_capacity(actual.saturating_mul(2));
129 let mut valid_sources = 0usize;
130 for (price, source) in prices.zip(&self.sources) {
131 if let Some(price) = price.filter(|_| source.weight > 0) {
132 valid_sources += 1;
133 let (lower, upper) = SpecificPrice::split(&price);
134 values.push((lower, source.weight));
135 values.push((upper, source.weight));
136 }
137 }
138
139 let min_sources = self.min_sources.max(1);
140 if valid_sources < min_sources as usize {
141 return Err(super::Error::TooFewValidSources {
142 expected: min_sources as usize,
143 actual: valid_sources,
144 });
145 }
146
147 values.sort_unstable();
148
149 Ok(values.swap_remove(V::median(&values)).0.into())
150 }
151}
152
153#[cfg(test)]
154mod tests {
155 use alloc::vec;
156
157 use crate::proxy::aggregator::method::Error;
158
159 use super::*;
160
161 fn price(value: i64, conf: u64, publish_time_s: u64) -> Price {
162 Price {
163 price: value,
164 conf,
165 expo: -6,
166 publish_time_ns: templar_primitives::Nanoseconds::from_secs(publish_time_s),
167 }
168 }
169
170 fn median_low(weights: &[u32], min_sources: u32) -> MedianLow<&'static str> {
171 MedianLow {
172 _variant: PhantomData,
173 sources: weights
174 .iter()
175 .map(|weight| WeightedSource::new("source", *weight))
176 .collect(),
177 min_sources,
178 }
179 }
180
181 #[test]
182 fn aggregate_empty_returns_too_few_valid_sources() {
183 let error = MedianLow::<&'static str>::new([])
184 .aggregate(vec![])
185 .unwrap_err();
186 assert!(matches!(
187 error,
188 Error::TooFewValidSources {
189 expected: 1,
190 actual: 0,
191 }
192 ));
193 }
194
195 #[test]
196 fn aggregate_single_price_no_conf() {
197 let result = median_low(&[1], 1)
198 .aggregate(vec![Some(price(1_000_000, 0, 0))])
199 .unwrap();
200 assert_eq!(result.price, 1_000_000);
201 }
202
203 #[test]
204 fn aggregate_median_of_three() {
205 let prices = vec![
206 Some(price(1_000_000, 0, 0)),
207 Some(price(2_000_000, 0, 0)),
208 Some(price(3_000_000, 0, 0)),
209 ];
210 let result = median_low(&[1, 1, 1], 1).aggregate(prices).unwrap();
211 assert_eq!(result.price, 2_000_000);
212 }
213
214 #[test]
215 fn aggregate_min_sources_not_met_returns_error() {
216 let prices = vec![Some(price(1_000_000, 0, 0)), Some(price(2_000_000, 0, 0))];
217 let error = median_low(&[1, 1], 3).aggregate(prices).unwrap_err();
218 assert!(matches!(
219 error,
220 Error::TooFewValidSources {
221 expected: 3,
222 actual: 2,
223 }
224 ));
225 }
226
227 #[test]
228 fn aggregate_min_sources_exactly_met() {
229 let prices = vec![Some(price(1_000_000, 0, 0)), Some(price(2_000_000, 0, 0))];
230 assert!(median_low(&[1, 1], 2).aggregate(prices).is_ok());
231 }
232
233 #[test]
234 fn raw_weighted_median_handles_zero_and_simple_edges() {
235 assert_eq!(median(&[("a", 0_u32), ("b", 0_u32)]), (0, 1));
236 assert_eq!(median(&[("a", 1_u32)]), (0, 0));
237 assert_eq!(median(&[("a", 1_u32), ("b", 1_u32), ("c", 1_u32)]), (1, 1));
238 assert_eq!(
239 median(&[("a", 1_u32), ("b", 100_u32), ("c", 1_u32)]),
240 (1, 1)
241 );
242 assert_eq!(
243 median(&[("a", 0_u32), ("b", 0_u32), ("c", 0_u32), ("d", 0_u32)]),
244 (1, 2)
245 );
246 }
247
248 #[test]
249 fn raw_weighted_median_handles_large_cumulative_weight_without_u32_overflow() {
250 let list = [
251 ("a", u32::MAX - 10),
252 ("b", 20),
253 ("c", 10),
254 ("d", u32::MAX - 5),
255 ];
256
257 assert_eq!(Low::median(&list), 1);
258 assert_eq!(High::median(&list), 1);
259 }
260
261 #[rstest::rstest]
262 #[case(&[("a", 1)], "a")]
263 #[case(&[("a", 1), ("b", 1), ("c", 1)], "b")]
264 #[case(&[("a", 1), ("b", 1), ("c", 1), ("d", 1)], "b")]
265 #[case(&[("a", 2), ("b", 1), ("c", 1), ("d", 1)], "b")]
266 #[case(&[("a", 1), ("b", 1), ("c", 1), ("d", 2)], "c")]
267 #[case(&[("a", 10), ("b", 2), ("c", 6), ("d", 2)], "a")]
268 #[case(&[("a", 1), ("b", 10000), ("c", 1)], "b")]
269 #[case(&[("a", 2), ("b", 1), ("c", 1)], "a")]
270 #[case(&[("a", u32::MAX), ("b", u32::MAX), ("c", u32::MAX)], "b")]
271 #[case(&[("a", u32::MAX), ("b", 0), ("c", u32::MAX)], "a")]
272 #[case(&[("a", 0), ("b", 0), ("c", 0), ("d", 0)], "b")]
273 #[case(&[("a", 0), ("b", 0), ("c", 0), ("d", 0), ("e", 0)], "c")]
274 #[case(&[("a", 0), ("b", 0), ("c", 0), ("d", 1)], "d")]
275 #[case(&[("a", 0), ("b", 1), ("c", 0), ("d", 1)], "b")]
276 fn weighted_median_low(#[case] list: &[(&str, u32)], #[case] expected: &str) {
277 let item = list[Low::median(list)].0;
278 assert_eq!(item, expected);
279 }
280
281 #[rstest::rstest]
282 #[case(&[("a", 0), ("b", 0)], "b")]
283 #[case(&[("a", 0), ("b", 0), ("c", 0), ("d", 0)], "c")]
284 #[case(&[("a", 0), ("b", 0), ("c", 0), ("d", 0), ("e", 0)], "c")]
285 fn weighted_median_high_all_zero_uses_upper_middle(
286 #[case] list: &[(&str, u32)],
287 #[case] expected: &str,
288 ) {
289 let item = list[High::median(list)].0;
290 assert_eq!(item, expected);
291 }
292
293 #[test]
294 fn zero_weight_sources_do_not_satisfy_quorum() {
295 let prices = vec![
296 Some(price(100, 0, 0)),
297 Some(price(1, 0, 0)),
298 Some(price(1, 0, 0)),
299 ];
300 assert!(matches!(
301 median_low(&[1, 0, 0], 3).aggregate(prices),
302 Err(Error::TooFewValidSources {
303 expected: 3,
304 actual: 1
305 })
306 ));
307 }
308
309 #[test]
310 fn hal_36_high_median_ignores_minority_and_zero_weight_items() {
311 assert_eq!(High::median(&[("a", 2), ("b", 1)]), 0);
312 assert_eq!(High::median(&[("a", 1), ("b", 0)]), 0);
313 }
314
315 #[test]
316 fn weighted_medians_match_repeated_weight_reference() {
317 for len in 1_usize..=6 {
318 let combinations = (0..len).fold(1_usize, |combinations, _| combinations * 5);
319 for encoded in 0..combinations {
320 let mut digits = encoded;
321 let items = (0..len)
322 .map(|index| {
323 let Ok(weight) = u32::try_from(digits % 5) else {
324 unreachable!("modulo-five remainder fits u32");
325 };
326 digits /= 5;
327 (index, weight)
328 })
329 .collect::<Vec<_>>();
330 let repeated = items
331 .iter()
332 .flat_map(|(index, weight)| {
333 core::iter::repeat_n(
334 *index,
335 usize::try_from(*weight)
336 .unwrap_or_else(|_| unreachable!("weight fits usize")),
337 )
338 })
339 .collect::<Vec<_>>();
340
341 if repeated.is_empty() {
342 continue;
343 }
344
345 assert_eq!(
346 Low::median(&items),
347 repeated[(repeated.len() - 1) / 2],
348 "{items:?}"
349 );
350 assert_eq!(
351 High::median(&items),
352 repeated[repeated.len() / 2],
353 "{items:?}"
354 );
355 }
356 }
357 }
358}