1use std::ops::Deref;
2
3use near_sdk::{near, require};
4use templar_primitives::number::Decimal;
5
6pub trait UsageCurve {
7 fn at(&self, usage_ratio: Decimal) -> Decimal;
8}
9
10#[derive(Clone, Debug, PartialEq, Eq)]
11#[near(serializers = [json, borsh])]
12pub enum InterestRateStrategy {
13 Linear(Linear),
14 #[cfg_attr(not(target_arch = "wasm32"), schemars(with = "PiecewiseParams"))]
15 Piecewise(Piecewise),
16 #[cfg_attr(not(target_arch = "wasm32"), schemars(with = "Exponential2Params"))]
17 Exponential2(Exponential2),
18}
19
20impl InterestRateStrategy {
21 pub const fn zero() -> Self {
22 Self::Linear(Linear {
23 base: Decimal::ZERO,
24 top: Decimal::ZERO,
25 })
26 }
27
28 #[must_use]
29 pub fn linear(base: Decimal, top: Decimal) -> Option<Self> {
30 Some(Self::Linear(Linear::new(base, top)?))
31 }
32
33 #[must_use]
34 pub fn piecewise(
35 base: Decimal,
36 optimal: Decimal,
37 rate_1: Decimal,
38 rate_2: Decimal,
39 ) -> Option<Self> {
40 Some(Self::Piecewise(Piecewise::new(
41 base, optimal, rate_1, rate_2,
42 )?))
43 }
44
45 #[must_use]
46 pub fn exponential2(base: Decimal, top: Decimal, eccentricity: Decimal) -> Option<Self> {
47 Some(Self::Exponential2(Exponential2::new(
48 base,
49 top,
50 eccentricity,
51 )?))
52 }
53}
54
55impl Deref for InterestRateStrategy {
56 type Target = dyn UsageCurve;
57
58 fn deref(&self) -> &Self::Target {
59 match self {
60 Self::Linear(linear) => linear as &dyn UsageCurve,
61 Self::Piecewise(piecewise) => piecewise as &dyn UsageCurve,
62 Self::Exponential2(exponential2) => exponential2 as &dyn UsageCurve,
63 }
64 }
65}
66
67#[derive(Debug, Clone, PartialEq, Eq)]
71#[near(serializers = [borsh, json])]
72pub struct Linear {
73 base: Decimal,
74 top: Decimal,
75}
76
77impl Linear {
78 pub fn new(base: Decimal, top: Decimal) -> Option<Self> {
79 (base <= top).then_some(Self { base, top })
80 }
81}
82
83impl UsageCurve for Linear {
84 fn at(&self, usage_ratio: Decimal) -> Decimal {
85 usage_ratio * (self.top - self.base) + self.base
86 }
87}
88
89#[derive(Debug, Clone, PartialEq, Eq)]
96#[near(serializers = [borsh, json])]
97#[serde(try_from = "PiecewiseParams", into = "PiecewiseParams")]
98pub struct Piecewise {
99 params: PiecewiseParams,
100 i_negative_rate_2_b: Decimal,
101}
102
103impl Piecewise {
104 pub fn new(base: Decimal, optimal: Decimal, rate_1: Decimal, rate_2: Decimal) -> Option<Self> {
105 if optimal > 1u32 {
106 return None;
107 }
108
109 if rate_1 > rate_2 {
110 return None;
111 }
112
113 Some(Self {
114 i_negative_rate_2_b: optimal * (rate_2 - rate_1) - base,
115 params: PiecewiseParams {
116 base,
117 optimal,
118 rate_1,
119 rate_2,
120 },
121 })
122 }
123}
124
125impl UsageCurve for Piecewise {
126 fn at(&self, usage_ratio: Decimal) -> Decimal {
127 require!(
128 usage_ratio <= Decimal::ONE,
129 "Invariant violation: Usage ratio cannot be over 100%.",
130 );
131
132 if usage_ratio < self.params.optimal {
133 self.params.rate_1 * usage_ratio + self.params.base
134 } else {
135 self.params.rate_2 * usage_ratio - self.i_negative_rate_2_b
136 }
137 }
138}
139
140#[derive(Debug, Clone, PartialEq, Eq)]
141#[near(serializers = [json, borsh])]
142pub struct PiecewiseParams {
143 base: Decimal,
144 optimal: Decimal,
145 rate_1: Decimal,
146 rate_2: Decimal,
147}
148
149impl TryFrom<PiecewiseParams> for Piecewise {
150 type Error = &'static str;
151
152 fn try_from(
153 PiecewiseParams {
154 base,
155 optimal,
156 rate_1,
157 rate_2,
158 }: PiecewiseParams,
159 ) -> Result<Self, Self::Error> {
160 Self::new(base, optimal, rate_1, rate_2).ok_or("Invalid Piecewise parameters")
161 }
162}
163
164impl From<Piecewise> for PiecewiseParams {
165 fn from(value: Piecewise) -> Self {
166 value.params
167 }
168}
169
170#[derive(Debug, Clone, PartialEq, Eq)]
174#[near(serializers = [borsh, json])]
175#[serde(try_from = "Exponential2Params", into = "Exponential2Params")]
176pub struct Exponential2 {
177 params: Exponential2Params,
178 i_factor: Decimal,
179}
180
181impl Exponential2 {
182 pub fn new(base: Decimal, top: Decimal, eccentricity: Decimal) -> Option<Self> {
185 if base > top {
186 return None;
187 }
188
189 if eccentricity > 24u32 || eccentricity.is_zero() {
190 return None;
191 }
192
193 #[allow(clippy::unwrap_used, reason = "Invariant checked above")]
194 Some(Self {
195 i_factor: (top - base) / (eccentricity.pow2().unwrap() - 1u32),
196 params: Exponential2Params {
197 base,
198 top,
199 eccentricity,
200 },
201 })
202 }
203}
204
205impl UsageCurve for Exponential2 {
206 fn at(&self, usage_ratio: Decimal) -> Decimal {
207 require!(
208 usage_ratio <= Decimal::ONE,
209 "Invariant violation: Usage ratio cannot be over 100%.",
210 );
211
212 #[allow(clippy::unwrap_used, reason = "Invariant checked above")]
213 (self.params.base
214 + self.i_factor * ((self.params.eccentricity * usage_ratio).pow2().unwrap() - 1u32))
215 }
216}
217
218#[derive(Debug, Clone, PartialEq, Eq)]
219#[near(serializers = [json, borsh])]
220pub struct Exponential2Params {
221 base: Decimal,
222 top: Decimal,
223 eccentricity: Decimal,
224}
225
226impl TryFrom<Exponential2Params> for Exponential2 {
227 type Error = &'static str;
228
229 fn try_from(
230 Exponential2Params {
231 base,
232 top,
233 eccentricity,
234 }: Exponential2Params,
235 ) -> Result<Self, Self::Error> {
236 Self::new(base, top, eccentricity).ok_or("Invalid Exponential2 parameters")
237 }
238}
239
240impl From<Exponential2> for Exponential2Params {
241 fn from(value: Exponential2) -> Self {
242 value.params
243 }
244}
245
246#[cfg(test)]
247mod tests {
248 use std::ops::Div;
249
250 use templar_primitives::dec;
251
252 #[test]
253 fn schema_matches_serde_wire_forms() {
254 let schema = near_sdk::serde_json::to_value(schemars::schema_for!(InterestRateStrategy))
255 .expect("strategy schema serializes");
256 let validator = jsonschema::draft7::new(&schema).expect("strategy schema is Draft 7");
257 let strategies = [
258 (
259 InterestRateStrategy::linear(Decimal::ZERO, Decimal::ONE)
260 .expect("valid linear strategy"),
261 near_sdk::serde_json::json!({"Linear":{"base":"0","top":"1"}}),
262 ),
263 (
264 InterestRateStrategy::piecewise(
265 Decimal::ZERO,
266 dec!("0.5"),
267 dec!("0.125"),
268 dec!("0.5"),
269 )
270 .expect("valid piecewise strategy"),
271 near_sdk::serde_json::json!({"Piecewise":{"base":"0","optimal":"0.5","rate_1":"0.125","rate_2":"0.5"}}),
272 ),
273 (
274 InterestRateStrategy::exponential2(Decimal::ZERO, Decimal::ONE, Decimal::ONE)
275 .expect("valid exponential strategy"),
276 near_sdk::serde_json::json!({"Exponential2":{"base":"0","top":"1","eccentricity":"1"}}),
277 ),
278 ];
279
280 for (strategy, expected_json) in strategies {
281 let serialized =
282 near_sdk::serde_json::to_value(&strategy).expect("strategy serialization succeeds");
283 assert_eq!(serialized, expected_json);
284 validator
285 .validate(&serialized)
286 .expect("serialized strategy is valid under its schema");
287 assert_eq!(
288 near_sdk::serde_json::from_value::<InterestRateStrategy>(serialized)
289 .expect("strategy deserialization succeeds"),
290 strategy
291 );
292 }
293
294 for internal_json in [
295 near_sdk::serde_json::json!({"Piecewise":{"params":{"base":"0","optimal":"0.5","rate_1":"0.125","rate_2":"0.5"},"i_negative_rate_2_b":"0.1875"}}),
296 near_sdk::serde_json::json!({"Exponential2":{"params":{"base":"0","top":"1","eccentricity":"1"},"i_factor":"1"}}),
297 ] {
298 validator
299 .validate(&internal_json)
300 .expect_err("internal runtime representation is not schema-valid");
301 near_sdk::serde_json::from_value::<InterestRateStrategy>(internal_json)
302 .expect_err("internal runtime representation is not serde-valid");
303 }
304 }
305
306 use super::*;
307
308 #[test]
309 fn piecewise() {
310 let s = Piecewise::new(Decimal::ZERO, dec!("0.9"), dec!("0.035"), dec!("0.6")).unwrap();
311
312 assert!(s.at(Decimal::ZERO).near_equal(Decimal::ZERO));
313 assert!(s.at(dec!("0.1")).near_equal(dec!("0.0035")));
314 assert!(s.at(dec!("0.5")).near_equal(dec!("0.0175")));
315 assert!(s.at(dec!("0.6")).near_equal(dec!("0.021")));
316 assert!(s.at(dec!("0.9")).near_equal(dec!("0.0315")));
317 assert!(s.at(dec!("0.95")).near_equal(dec!("0.0615")));
318 assert!(s.at(Decimal::ONE).near_equal(dec!("0.0915")));
319 }
320
321 #[test]
329 #[should_panic(expected = "overflow")]
332 fn piecewise_new_underflows_when_base_exceeds_cross_term() {
333 let _ = Piecewise::new(dec!("0.5"), dec!("0.9"), dec!("0.0"), dec!("0.1"));
335 }
336
337 #[test]
338 fn exponential2() {
339 let s = Exponential2::new(dec!("0.005"), dec!("0.08"), dec!("6")).unwrap();
340 assert!(s.at(Decimal::ZERO).near_equal(dec!("0.005")));
341 assert!(s
342 .at(dec!("0.25"))
343 .near_equal(dec!("0.00717669895803117868762306839097547161")));
344 assert!(s.at(Decimal::ONE_HALF).near_equal(Decimal::ONE.div(75u32)));
345 }
346}