templar_proxy_oracle_kernel/proxy/aggregator/method/median/
mod.rs

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
15/// Calculates the lower and upper weighted medians of a sorted list.
16///
17/// Zero-weight items do not contribute to either target when total weight is
18/// positive. An all-zero list uses the positional-median fallback.
19///
20/// # Panics
21///
22/// If the list is empty.
23fn 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        /// Minimum number of sources required for the aggregation to produce a result.
95        ///
96        /// For example, if the proxy has a Pyth source and a RedStone source, and `min_sources` is set to `2`,
97        /// the aggregation will only produce a result if both oracles provide a price.
98        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}