Skip to main content

solver/models/
bilateral_hawkes_order_flow_imbalance.rs

1use super::control::{ControlProblem, StateDerivatives};
2use super::market_making::MarketMakingControl;
3use super::normal_cdf;
4use crate::numeric::finite_difference::discretization::{DimensionKind, Transport};
5use crate::numeric::finite_difference::pde::PdeProblem;
6
7// Re-export so existing paths compile.
8pub use super::TerminalCondition;
9
10/// Bilateral Hawkes model with price-impact-driven adverse selection.
11///
12/// # State Space: `[q, lambda_plus, lambda_minus]`
13///
14/// | dim | symbol       | description                                      |
15/// |-----|--------------|--------------------------------------------------|
16/// | 0   | q            | Inventory level (integer, jump-controlled)       |
17/// | 1   | lambda_plus  | Buy market-order intensity (ODE drift, no noise) |
18/// | 2   | lambda_minus | Sell market-order intensity (ODE drift, no noise)|
19///
20/// # Dynamics
21///
22/// `d(lambda+) = beta*(mu - lambda+) dt + alpha dN+`
23/// `d(lambda-) = beta*(mu - lambda-) dt + alpha dN-`
24/// `dS = sigma dW + eta_ofi (dN+ - dN-)`
25///
26/// Each buy market order (dN+) moves the mid price UP by `eta_ofi`; each sell market
27/// order (dN-) moves it DOWN by `eta_ofi`. The market maker accounts for this
28/// price impact when setting optimal spreads.
29///
30/// # Optimal Spreads Under Price Impact
31///
32/// Deriving the FOC from the CARA HJB with per-fill wealth effect `(q±1)*eta_ofi` gives:
33///
34/// $$ \delta_{bid}^* = \delta_{base} - \frac{\partial V}{\partial q} + (q+1)\eta_{OFI} $$
35/// $$ \delta_{ask}^* = \delta_{base} + \frac{\partial V}{\partial q} - (q-1)\eta_{OFI} $$
36///
37/// Economic interpretation:
38/// - Long inventory (q > 0): ask narrows (selling into a rising price is favourable),
39///   bid widens (avoid accumulating more when sell MOs push price down).
40/// - Short inventory (q < 0): ask widens (avoid shorting further into a rising price),
41///   bid narrows (cover cheaply when sell MOs push price down).
42///
43/// Setting `eta_ofi = 0` recovers the plain `BilateralHawkes` (BHK) model exactly.
44///
45/// In the simulation pass `impact_factor = -eta_ofi` to `SimulatedDataSource` so the
46/// price process matches the assumed dynamics `dS = sigma dW + eta_ofi (dN+ - dN-)`.
47///
48/// # Hawkes Approximation
49///
50/// Same fluid (mean-field) approximation as `BilateralHawkes`:
51///
52/// $$ d(\lambda_+)/dt \approx \beta(\mu - \lambda_+) + \alpha \lambda_+ $$
53/// $$ d(\lambda_-)/dt \approx \beta(\mu - \lambda_-) + \alpha \lambda_- $$
54///
55#[derive(Clone, Debug)]
56pub struct BilateralHawkesOrderFlowImbalance {
57    /// Risk-aversion parameter.
58    pub gamma: f64,
59    /// Price volatility.
60    pub sigma: f64,
61    /// Fill-rate decay (Avellaneda-Stoikov kappa).
62    pub kappa: f64,
63    // --- Hawkes parameters (shared for both sides) ---
64    /// Jump size added to the excited side's intensity.
65    pub alpha: f64,
66    /// Mean-reversion speed.
67    pub beta: f64,
68    /// Baseline intensity (mean-reversion target).
69    pub mu: f64,
70    // --- Price-impact parameter ---
71    /// Price impact per fill event (eta_ofi in dS = sigma*dW + eta_ofi*(dN+ - dN-)).
72    /// Higher eta_ofi = stronger q-dependent spread correction.
73    /// Setting eta_ofi = 0 recovers the plain BilateralHawkes (BHK) model.
74    pub eta_ofi: f64,
75    /// Grid step for both lambda+ and lambda- dimensions.
76    pub lambda_step: f64,
77    /// Grid step for the inventory dimension.
78    pub dq: f64,
79    /// Minimum inventory (lower hard boundary). Defaults to -infinity.
80    pub q_min: f64,
81    /// Maximum inventory (upper hard boundary). Defaults to +infinity.
82    pub q_max: f64,
83    /// Terminal condition for the value function at T. Defaults to Zero.
84    pub terminal_condition: TerminalCondition,
85    /// Optional terminal liquidation half-spread used when terminal_condition is
86    /// LiquidationCost. If None, uses the AS base spread.
87    pub terminal_liquidation_half_spread: Option<f64>,
88}
89
90impl BilateralHawkesOrderFlowImbalance {
91    pub fn new(
92        gamma: f64,
93        sigma: f64,
94        kappa: f64,
95        alpha: f64,
96        beta: f64,
97        mu: f64,
98        eta_ofi: f64,
99    ) -> Self {
100        Self {
101            gamma,
102            sigma,
103            kappa,
104            alpha,
105            beta,
106            mu,
107            eta_ofi,
108            lambda_step: 1.0,
109            dq: 1.0,
110            q_min: f64::NEG_INFINITY,
111            q_max: f64::INFINITY,
112            terminal_condition: TerminalCondition::Zero,
113            terminal_liquidation_half_spread: None,
114        }
115    }
116
117    pub fn with_terminal_condition(mut self, terminal_condition: TerminalCondition) -> Self {
118        self.terminal_condition = terminal_condition;
119        self
120    }
121
122    pub fn with_terminal_liquidation_half_spread(mut self, half_spread: f64) -> Self {
123        self.terminal_liquidation_half_spread = Some(half_spread.max(0.0));
124        self
125    }
126
127    pub fn with_inventory_bounds(mut self, q_min: f64, q_max: f64) -> Self {
128        self.q_min = q_min;
129        self.q_max = q_max;
130        self
131    }
132
133    pub fn with_lambda_step(mut self, lambda_step: f64) -> Self {
134        self.lambda_step = lambda_step.abs().max(1e-8);
135        self
136    }
137
138    pub fn with_dq(mut self, dq: f64) -> Self {
139        self.dq = dq.abs().max(1e-8);
140        self
141    }
142
143    /// Computes optimal bid/ask half-spreads with price-impact-driven inventory correction.
144    ///
145    /// When eta_ofi > 0 each fill carries an adverse-selection wealth effect `(q±1)*eta_ofi`.
146    /// The FOC on the CARA HJB gives:
147    ///   delta_bid = base_spread - dV/dq_fwd + (q+1)*eta_ofi
148    ///   delta_ask = base_spread + dV/dq_bwd - (q-1)*eta_ofi
149    ///
150    /// Long inventory (q > 0): ask narrows (sell into rising price), bid widens.
151    /// Short inventory (q < 0): ask widens (avoid adding short), bid narrows.
152    /// eta_ofi = 0 recovers the plain BHK spread formula identically.
153    pub fn get_spreads(&self, state: &[f64; 3], derivs: &StateDerivatives<3>) -> (f64, f64) {
154        let q = state[0];
155        let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa).ln();
156        let delta_bid = (base_spread - derivs.fwd[0] + (q + 1.0) * self.eta_ofi).clamp(-10.0, 10.0);
157        let delta_ask = (base_spread + derivs.bwd[0] - (q - 1.0) * self.eta_ofi).clamp(-10.0, 10.0);
158        (delta_bid, delta_ask)
159    }
160
161    /// The base arrival intensity is state-dependent: returns the average of
162    /// the buy/sell intensity states for intensity-to-spread conversion.
163    pub fn fill_rate_base(&self, state: &[f64; 3]) -> f64 {
164        ((state[1] + state[2]) * 0.5).max(1e-10)
165    }
166
167    /// The fill-rate decay parameter used for intensity-to-spread conversion.
168    pub fn fill_rate_decay(&self) -> f64 {
169        self.kappa
170    }
171}
172
173impl ControlProblem<3> for BilateralHawkesOrderFlowImbalance {
174    type Control = MarketMakingControl;
175
176    fn optimize(&self, _t: f64, state: &[f64; 3], derivs: &StateDerivatives<3>) -> Self::Control {
177        let q = state[0];
178        let lambda_plus = state[1].max(0.0);
179        let lambda_minus = state[2].max(0.0);
180
181        let (d_bid, d_ask) = self.get_spreads(state, derivs);
182
183        let max_mult = 20.0_f64;
184        let lambda_bid_fill = lambda_minus * (-self.kappa * d_bid).exp().min(max_mult);
185        let lambda_ask_fill = lambda_plus * (-self.kappa * d_ask).exp().min(max_mult);
186
187        let lambda_bid_fill = if q >= self.q_max {
188            0.0
189        } else {
190            lambda_bid_fill
191        };
192        let lambda_ask_fill = if q <= self.q_min {
193            0.0
194        } else {
195            lambda_ask_fill
196        };
197
198        MarketMakingControl::new(lambda_bid_fill, lambda_ask_fill)
199    }
200
201    fn running_reward(&self, _t: f64, _state: &[f64; 3], control: &Self::Control) -> f64 {
202        (control.bid_intensity + control.ask_intensity) / (self.gamma + self.kappa)
203    }
204
205    fn bsde_driver(
206        &self,
207        _t: f64,
208        state: &[f64; 3],
209        control: &Self::Control,
210        _derivs: &StateDerivatives<3>,
211        dt: f64,
212    ) -> f64 {
213        // Full-value problem: the two Hawkes intensity drifts are already
214        // simulated forward, so the backward driver adds only the running
215        // reward plus the local inventory-risk source. The reward is bounded at
216        // `1/dt` per fill side to match the forward fill-probability clamp.
217        let q = state[0];
218        let local = -0.5 * self.gamma * self.sigma.powi(2) * q.powi(2);
219        let rate_cap = 1.0 / dt.max(1e-12);
220        let reward = (control.bid_intensity.min(rate_cap) + control.ask_intensity.min(rate_cap))
221            / (self.gamma + self.kappa);
222        reward + local
223    }
224
225    fn generator(
226        &self,
227        _t: f64,
228        state: &[f64; 3],
229        _control: &Self::Control,
230        derivs: &StateDerivatives<3>,
231    ) -> f64 {
232        let q = state[0];
233        let lambda_plus = state[1].max(0.0);
234        let lambda_minus = state[2].max(0.0);
235        let net_drift_plus = self.beta * (self.mu - lambda_plus) + self.alpha * lambda_plus;
236        let net_drift_minus = self.beta * (self.mu - lambda_minus) + self.alpha * lambda_minus;
237
238        -0.5 * self.gamma * self.sigma.powi(2) * q.powi(2)
239            + net_drift_plus * derivs.grad[1]
240            + net_drift_minus * derivs.grad[2]
241    }
242
243    fn terminal(&self, state: &[f64; 3]) -> f64 {
244        match self.terminal_condition {
245            TerminalCondition::Zero => 0.0,
246            TerminalCondition::LiquidationCost => {
247                let q = state[0];
248                let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa).ln();
249                let half_spread = self.terminal_liquidation_half_spread.unwrap_or(base_spread);
250                -q.abs() * half_spread
251            }
252        }
253    }
254
255    fn discount_rate(&self, _state: &[f64; 3]) -> f64 {
256        0.0
257    }
258
259    fn constant_discount_rate(&self) -> Option<f64> {
260        Some(0.0)
261    }
262
263    fn next_step(&self, _t: f64, state: &[f64; 3], dt: f64, noise: &[f64; 3]) -> [f64; 3] {
264        let lambda_plus = state[1].max(0.0);
265        let lambda_minus = state[2].max(0.0);
266        let drift_plus = self.beta * (self.mu - lambda_plus) + self.alpha * lambda_plus;
267        let drift_minus = self.beta * (self.mu - lambda_minus) + self.alpha * lambda_minus;
268        let mut next = [
269            state[0],
270            (lambda_plus + drift_plus * dt).max(0.0),
271            (lambda_minus + drift_minus * dt).max(0.0),
272        ];
273
274        let u = normal_cdf(noise[0]);
275        let lambda_plus = state[1].max(0.0);
276        let lambda_minus = state[2].max(0.0);
277        let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa).ln();
278        let lambda_bid_fill = lambda_minus * (-self.kappa * base_spread).exp();
279        let lambda_ask_fill = lambda_plus * (-self.kappa * base_spread).exp();
280        let p_bid = (lambda_bid_fill * dt).clamp(0.0, 1.0);
281        let p_ask = (lambda_ask_fill * dt).clamp(0.0, 1.0);
282        if u < p_bid {
283            next[0] += 1.0;
284            next[2] += self.alpha;
285        } else if u > 1.0 - p_ask {
286            next[0] -= 1.0;
287            next[1] += self.alpha;
288        }
289
290        next[1] = next[1].max(0.0);
291        next[2] = next[2].max(0.0);
292        next
293    }
294
295    fn next_step_controlled(
296        &self,
297        _t: f64,
298        state: &[f64; 3],
299        control: &Self::Control,
300        dt: f64,
301        noise: &[f64; 3],
302    ) -> [f64; 3] {
303        let lambda_plus = state[1].max(0.0);
304        let lambda_minus = state[2].max(0.0);
305        let drift_plus = self.beta * (self.mu - lambda_plus) + self.alpha * lambda_plus;
306        let drift_minus = self.beta * (self.mu - lambda_minus) + self.alpha * lambda_minus;
307        let mut next = [
308            state[0],
309            (lambda_plus + drift_plus * dt).max(0.0),
310            (lambda_minus + drift_minus * dt).max(0.0),
311        ];
312
313        // The optimal control already carries the fill intensities, so the
314        // Bernoulli fill events use them instead of a frozen proxy quote.
315        let u = normal_cdf(noise[0]);
316        let p_bid = (control.bid_intensity * dt).clamp(0.0, 1.0);
317        let p_ask = (control.ask_intensity * dt).clamp(0.0, 1.0);
318        if u < p_bid {
319            next[0] += 1.0;
320            next[2] += self.alpha;
321        } else if u > 1.0 - p_ask {
322            next[0] -= 1.0;
323            next[1] += self.alpha;
324        }
325
326        next[1] = next[1].max(0.0);
327        next[2] = next[2].max(0.0);
328        next
329    }
330
331    fn is_diffusion_dimension(&self, _dim: usize) -> bool {
332        false
333    }
334
335    fn gradient_step(&self, dim: usize) -> f64 {
336        match dim {
337            1 | 2 => self.lambda_step.abs().max(1e-8),
338            _ => 1.0,
339        }
340    }
341}
342
343impl PdeProblem<3> for BilateralHawkesOrderFlowImbalance {
344    fn dimension_kind(&self, dim: usize) -> DimensionKind {
345        match dim {
346            0 => DimensionKind::DiscreteJump,
347            _ => DimensionKind::DeterministicDrift,
348        }
349    }
350
351    fn transport(
352        &self,
353        _t: f64,
354        state: &[f64; 3],
355        control: &Self::Control,
356        derivs: &StateDerivatives<3>,
357    ) -> Transport<3> {
358        let q = state[0];
359        let lambda_plus = state[1].max(0.0);
360        let lambda_minus = state[2].max(0.0);
361        let net_drift_plus = self.beta * (self.mu - lambda_plus) + self.alpha * lambda_plus;
362        let net_drift_minus = self.beta * (self.mu - lambda_minus) + self.alpha * lambda_minus;
363        let step = self.lambda_step.abs().max(1e-12);
364
365        let rate_plus = net_drift_plus.abs() / step;
366        let rate_minus = net_drift_minus.abs() / step;
367        let (lp_plus, lp_minus) = if net_drift_plus >= 0.0 {
368            (rate_plus, 0.0)
369        } else {
370            (0.0, rate_plus)
371        };
372        let (lm_plus, lm_minus) = if net_drift_minus >= 0.0 {
373            (rate_minus, 0.0)
374        } else {
375            (0.0, rate_minus)
376        };
377
378        let hamiltonian =
379            (control.bid_intensity + control.ask_intensity) / (self.gamma + self.kappa);
380        let jump_transport =
381            control.bid_intensity * derivs.fwd[0] - control.ask_intensity * derivs.bwd[0];
382        let local = -0.5 * self.gamma * self.sigma.powi(2) * q.powi(2);
383
384        Transport::new(
385            [control.bid_intensity, lp_plus, lm_plus],
386            [control.ask_intensity, lp_minus, lm_minus],
387            hamiltonian + local - jump_transport,
388        )
389    }
390}
391
392#[cfg(test)]
393mod tests {
394    use super::*;
395
396    fn zero_derivs() -> StateDerivatives<3> {
397        StateDerivatives::with_directional([0.0; 3], [0.0; 3], [0.0; 3], [0.0; 3])
398    }
399
400    fn derivs_with(fwd0: f64, bwd0: f64) -> StateDerivatives<3> {
401        let mut d = zero_derivs();
402        d.fwd[0] = fwd0;
403        d.bwd[0] = bwd0;
404        d
405    }
406
407    fn zero_state() -> [f64; 3] {
408        [0.0, 1.0, 1.0]
409    }
410
411    fn default_model() -> BilateralHawkesOrderFlowImbalance {
412        BilateralHawkesOrderFlowImbalance::new(0.5, 0.2, 1.5, 0.5, 2.0, 1.0, 0.5)
413            .with_lambda_step(0.1)
414            .with_dq(1.0)
415    }
416
417    #[test]
418    fn model_creation() {
419        let m = default_model();
420        assert_eq!(m.gamma, 0.5);
421        assert_eq!(m.sigma, 0.2);
422        assert_eq!(m.kappa, 1.5);
423        assert_eq!(m.alpha, 0.5);
424        assert_eq!(m.beta, 2.0);
425        assert_eq!(m.mu, 1.0);
426        assert_eq!(m.eta_ofi, 0.5);
427    }
428
429    #[test]
430    fn bsde_driver_excludes_intensity_transport() {
431        let m = default_model();
432        let state = [2.0, 1.0, 1.0];
433        let control = ControlProblem::optimize(&m, 0.0, &state, &zero_derivs());
434
435        let mut derivs = zero_derivs();
436        derivs.grad[1] = 1.0;
437        derivs.grad[2] = 1.0;
438
439        let dt = 0.01;
440        let q = state[0];
441        let local = -0.5 * m.gamma * m.sigma.powi(2) * q.powi(2);
442        let rate_cap = 1.0 / dt;
443        let reward = (control.bid_intensity.min(rate_cap) + control.ask_intensity.min(rate_cap))
444            / (m.gamma + m.kappa);
445        let expected = reward + local;
446
447        assert_eq!(
448            ControlProblem::bsde_driver(&m, 0.0, &state, &control, &derivs, dt),
449            expected
450        );
451        assert!(
452            ControlProblem::bsde_driver(&m, 0.0, &state, &control, &derivs, dt)
453                != ControlProblem::driver(&m, 0.0, &state, &control, &derivs),
454            "full driver adds intensity transport, bsde driver must not"
455        );
456    }
457
458    #[test]
459    fn next_step_controlled_uses_optimal_intensities() {
460        let m = default_model();
461        let state = [1.0, 1.0, 1.0];
462
463        let ctrl = MarketMakingControl::new(0.0, 0.0);
464        let next =
465            ControlProblem::next_step_controlled(&m, 0.0, &state, &ctrl, 0.01, &[0.0, 0.0, 0.0]);
466        assert_eq!(next[0], 1.0);
467
468        let ctrl = MarketMakingControl::new(1000.0, 0.0);
469        let next =
470            ControlProblem::next_step_controlled(&m, 0.0, &state, &ctrl, 0.01, &[0.0, 0.0, 0.0]);
471        assert_eq!(next[0], 2.0);
472    }
473
474    #[test]
475    fn dimension_kinds_match_state_space() {
476        let m = default_model();
477        assert_eq!(
478            PdeProblem::dimension_kind(&m, 0),
479            DimensionKind::DiscreteJump
480        );
481        assert_eq!(
482            PdeProblem::dimension_kind(&m, 1),
483            DimensionKind::DeterministicDrift
484        );
485        assert_eq!(
486            PdeProblem::dimension_kind(&m, 2),
487            DimensionKind::DeterministicDrift
488        );
489    }
490
491    #[test]
492    fn zero_inventory_gives_symmetric_spreads() {
493        let m = default_model();
494        let state = [0.0, 1.0, 1.0];
495        let (bid, ask) = m.get_spreads(&state, &zero_derivs());
496        let base = (1.0 / m.gamma) * (1.0 + m.gamma / m.kappa).ln();
497        assert!(
498            (bid - ask).abs() < 1e-12,
499            "bid and ask should be equal at q=0"
500        );
501        assert!((bid - (base + m.eta_ofi)).abs() < 1e-12);
502    }
503
504    #[test]
505    fn long_inventory_widens_bid_narrows_ask() {
506        let m = default_model(); // xi = 0.5
507        let q = 5.0;
508        let state_long = [q, 1.0, 1.0];
509        let state_zero = [0.0, 1.0, 1.0];
510        let (bid_long, ask_long) = m.get_spreads(&state_long, &zero_derivs());
511        let (bid_zero, ask_zero) = m.get_spreads(&state_zero, &zero_derivs());
512        assert!(bid_long > bid_zero, "long position widens bid");
513        assert!(ask_long < ask_zero, "long position narrows ask");
514    }
515
516    #[test]
517    fn short_inventory_narrows_bid_widens_ask() {
518        let m = default_model(); // xi = 0.5
519        let q = -5.0;
520        let state_short = [q, 1.0, 1.0];
521        let state_zero = [0.0, 1.0, 1.0];
522        let (bid_short, ask_short) = m.get_spreads(&state_short, &zero_derivs());
523        let (bid_zero, ask_zero) = m.get_spreads(&state_zero, &zero_derivs());
524        assert!(bid_short < bid_zero, "short position narrows bid");
525        assert!(ask_short > ask_zero, "short position widens ask");
526    }
527
528    #[test]
529    fn larger_xi_gives_larger_spread_corrections() {
530        let m1 = BilateralHawkesOrderFlowImbalance::new(0.5, 0.2, 1.5, 0.5, 2.0, 1.0, 0.5)
531            .with_lambda_step(0.1);
532        let m2 = BilateralHawkesOrderFlowImbalance::new(0.5, 0.2, 1.5, 0.5, 2.0, 1.0, 1.0)
533            .with_lambda_step(0.1);
534        let m0 = BilateralHawkesOrderFlowImbalance::new(0.5, 0.2, 1.5, 0.5, 2.0, 1.0, 0.0)
535            .with_lambda_step(0.1);
536
537        let state = [5.0, 1.0, 1.0]; // q=5 (long)
538        let derivs = zero_derivs();
539
540        let (bid0, ask0) = m0.get_spreads(&state, &derivs);
541        let (bid1, ask1) = m1.get_spreads(&state, &derivs);
542        let (bid2, ask2) = m2.get_spreads(&state, &derivs);
543
544        assert!(bid1 > bid0);
545        assert!(bid2 > bid1);
546        assert!(ask1 < ask0);
547        assert!(ask2 < ask1);
548    }
549
550    #[test]
551    fn terminal_condition_zero() {
552        let m = default_model().with_terminal_condition(TerminalCondition::Zero);
553        let state = [5.0, 1.0, 1.0];
554        assert_eq!(ControlProblem::terminal(&m, &state), 0.0);
555    }
556
557    #[test]
558    fn terminal_condition_liquidation_cost() {
559        let m = default_model()
560            .with_terminal_condition(TerminalCondition::LiquidationCost)
561            .with_terminal_liquidation_half_spread(0.1);
562
563        let state_q5 = [5.0, 1.0, 1.0];
564        let state_q0 = [0.0, 1.0, 1.0];
565        let state_qm3 = [-3.0, 1.0, 1.0];
566
567        assert_eq!(ControlProblem::terminal(&m, &state_q5), -5.0 * 0.1);
568        assert_eq!(ControlProblem::terminal(&m, &state_q0), 0.0);
569        assert_eq!(ControlProblem::terminal(&m, &state_qm3), -3.0 * 0.1);
570    }
571
572    #[test]
573    fn optimize_output_nonnegative() {
574        let m = default_model();
575        let state = zero_state();
576        let derivs = zero_derivs();
577        let output = ControlProblem::optimize(&m, 0.0, &state, &derivs);
578
579        assert!(output.bid_intensity >= 0.0);
580        assert!(output.ask_intensity >= 0.0);
581    }
582
583    #[test]
584    fn inventory_bounds_enforcement() {
585        let m = default_model().with_inventory_bounds(-2.0, 2.0);
586        let state = [2.0, 2.0, 1.0]; // at upper bound
587        let derivs = derivs_with(0.0, -1.0); // negative gradient encourages sells
588
589        let output = ControlProblem::optimize(&m, 0.0, &state, &derivs);
590        // At upper bound, can't increase inventory further: no bid fill.
591        assert_eq!(output.bid_intensity, 0.0);
592    }
593
594    #[test]
595    fn hawkes_dynamics_stability() {
596        let m = default_model();
597
598        let rho = m.alpha / m.beta;
599        let stationary_intensity = m.mu / (1.0 - rho);
600
601        let mut state = [0.0, stationary_intensity, stationary_intensity];
602
603        for _ in 0..1000 {
604            state = ControlProblem::next_step(&m, 0.0, &state, 0.01, &[0.0, 0.0, 0.0]);
605        }
606
607        assert!(state[1] > 0.0 && state[1] < 10.0 * stationary_intensity);
608        assert!(state[2] > 0.0 && state[2] < 10.0 * stationary_intensity);
609    }
610
611    #[test]
612    fn next_step_consistency() {
613        let m = default_model();
614        let state = [0.0, 1.5, 1.0];
615        let dt = 0.01;
616
617        let next = ControlProblem::next_step(&m, 0.0, &state, dt, &[0.0; 3]);
618
619        assert_eq!(next[0], state[0]);
620
621        let drift_plus = m.beta * (m.mu - state[1]) + m.alpha * state[1];
622        let expected_lambda_plus = state[1] + drift_plus * dt;
623        assert!((next[1] - expected_lambda_plus).abs() < 1e-10);
624    }
625
626    #[test]
627    fn next_step_can_fill_inventory() {
628        let m = default_model();
629        let state = [0.0, 2.0, 2.0];
630        let dt = 0.01;
631
632        // Very negative noise -> uniform(0,1) near 0 -> bid fill -> q increases.
633        let up = ControlProblem::next_step(&m, 0.0, &state, dt, &[-10.0, 0.0, 0.0]);
634        // Very positive noise -> uniform near 1 -> ask fill -> q decreases.
635        let down = ControlProblem::next_step(&m, 0.0, &state, dt, &[10.0, 0.0, 0.0]);
636
637        assert!(up[0] > state[0], "bid fill should increase inventory");
638        assert!(down[0] < state[0], "ask fill should decrease inventory");
639    }
640
641    #[test]
642    fn gradient_step_returns_correct_values() {
643        let m = default_model();
644
645        assert_eq!(ControlProblem::gradient_step(&m, 0), 1.0);
646        assert_eq!(ControlProblem::gradient_step(&m, 1), 0.1);
647        assert_eq!(ControlProblem::gradient_step(&m, 2), 0.1);
648    }
649
650    #[test]
651    fn is_diffusion_dimension() {
652        let m = default_model();
653
654        assert!(!ControlProblem::is_diffusion_dimension(&m, 0));
655        assert!(!ControlProblem::is_diffusion_dimension(&m, 1));
656        assert!(!ControlProblem::is_diffusion_dimension(&m, 2));
657    }
658
659    #[test]
660    fn constant_discount_rate() {
661        let m = default_model();
662        assert_eq!(ControlProblem::constant_discount_rate(&m), Some(0.0));
663    }
664
665    #[test]
666    fn spread_formula_matches_analytical_expression() {
667        let m = default_model(); // xi = 0.5
668        let base = (1.0 / m.gamma) * (1.0 + m.gamma / m.kappa).ln();
669        for &q in &[-5.0_f64, -1.0, 0.0, 1.0, 5.0] {
670            let state = [q, 1.0, 1.0];
671            let (bid, ask) = m.get_spreads(&state, &zero_derivs());
672            let expected_bid = base + (q + 1.0) * m.eta_ofi;
673            let expected_ask = base - (q - 1.0) * m.eta_ofi;
674            assert!((bid - expected_bid).abs() < 1e-12, "bid mismatch at q={q}");
675            assert!((ask - expected_ask).abs() < 1e-12, "ask mismatch at q={q}");
676        }
677    }
678
679    #[test]
680    fn compare_with_balanced_no_imbalance() {
681        let m = default_model();
682        let state = [0.0, 1.0, 1.0];
683
684        let (bid, ask) = m.get_spreads(&state, &zero_derivs());
685        assert!((bid - ask).abs() < 1e-12);
686    }
687
688    #[test]
689    fn extreme_imbalance_clamping() {
690        let m = BilateralHawkesOrderFlowImbalance::new(0.5, 0.2, 1.5, 0.5, 2.0, 1.0, 100.0)
691            .with_lambda_step(0.1);
692
693        let state = [0.0, 10.0, 0.0]; // extreme imbalance
694        let derivs = derivs_with(100.0, 100.0);
695
696        let (bid, ask) = m.get_spreads(&state, &derivs);
697
698        assert!((-10.0..=10.0).contains(&bid));
699        assert!((-10.0..=10.0).contains(&ask));
700    }
701
702    #[test]
703    fn inventory_control_respects_bounds() {
704        let m = default_model().with_inventory_bounds(-5.0, 5.0);
705        let state_at_max = [5.0, 2.0, 1.0];
706        let state_at_min = [-5.0, 1.0, 2.0];
707        let derivs = zero_derivs();
708
709        let output_max = ControlProblem::optimize(&m, 0.0, &state_at_max, &derivs);
710        let output_min = ControlProblem::optimize(&m, 0.0, &state_at_min, &derivs);
711
712        // At max inventory, cannot increase: bid fill must be zero.
713        assert_eq!(output_max.bid_intensity, 0.0);
714
715        // At min inventory, cannot decrease: ask fill must be zero.
716        assert_eq!(output_min.ask_intensity, 0.0);
717    }
718
719    #[test]
720    fn xi_zero_recovers_baseline() {
721        let m = BilateralHawkesOrderFlowImbalance::new(0.5, 0.2, 1.5, 0.5, 2.0, 1.0, 0.0)
722            .with_lambda_step(0.1);
723        let state = [0.0, 2.0, 1.0];
724        let (bid, ask) = m.get_spreads(&state, &zero_derivs());
725        let base = (1.0 / m.gamma) * (1.0 + m.gamma / m.kappa).ln();
726        assert!((bid - base).abs() < 1e-12);
727        assert!((ask - base).abs() < 1e-12);
728    }
729
730    #[test]
731    fn model_builder_methods() {
732        let m = BilateralHawkesOrderFlowImbalance::new(0.5, 0.2, 1.5, 0.5, 2.0, 1.0, 0.5)
733            .with_lambda_step(0.2)
734            .with_dq(0.5)
735            .with_inventory_bounds(-10.0, 10.0)
736            .with_terminal_condition(TerminalCondition::LiquidationCost)
737            .with_terminal_liquidation_half_spread(0.05);
738
739        assert_eq!(m.lambda_step, 0.2);
740        assert_eq!(m.dq, 0.5);
741        assert_eq!(m.q_min, -10.0);
742        assert_eq!(m.q_max, 10.0);
743        assert_eq!(m.terminal_liquidation_half_spread, Some(0.05));
744    }
745}