1use super::traits::{ControlOutput, Gradients, Model};
2
3use market_model::process::HestonProcess;
4
5pub use super::TerminalCondition;
7
8#[derive(Clone, Copy, Debug)]
21pub struct Heston {
22 pub gamma: f64,
24 pub kappa_m: f64, pub a: f64, pub v_kappa: f64, pub v_theta: f64, pub v_xi: f64, pub rho: f64,
36 pub price_scale: f64,
38 pub dv: f64, pub dq: f64,
42 pub q_min: f64,
44 pub q_max: f64,
46 pub terminal_condition: TerminalCondition,
48}
49
50impl Heston {
51 pub fn new(gamma: f64, kappa_m: f64, a: f64) -> Self {
52 Self {
53 gamma,
54 kappa_m,
55 a,
56 v_kappa: 2.0,
57 v_theta: 0.25,
58 v_xi: 0.3,
59 rho: 0.0,
60 price_scale: 1.0,
61 dv: 1.0,
62 dq: 1.0,
63 q_min: f64::NEG_INFINITY,
64 q_max: f64::INFINITY,
65 terminal_condition: TerminalCondition::Zero,
66 }
67 }
68
69 pub fn with_terminal_condition(mut self, terminal_condition: TerminalCondition) -> Self {
70 self.terminal_condition = terminal_condition;
71 self
72 }
73
74 pub fn with_inventory_bounds(mut self, q_min: f64, q_max: f64) -> Self {
75 self.q_min = q_min;
76 self.q_max = q_max;
77 self
78 }
79
80 pub fn with_variance_params(mut self, v_kappa: f64, v_theta: f64, v_xi: f64) -> Self {
81 self.v_kappa = v_kappa;
82 self.v_theta = v_theta;
83 self.v_xi = v_xi;
84 self
85 }
86
87 pub fn with_rho(mut self, rho: f64) -> Self {
88 self.rho = rho;
89 self
90 }
91
92 pub fn with_price_scale(mut self, price_scale: f64) -> Self {
93 self.price_scale = price_scale.abs().max(1e-8);
94 self
95 }
96
97 pub fn with_dv(mut self, dv: f64) -> Self {
98 self.dv = dv.abs().max(1e-8);
99 self
100 }
101
102 pub fn with_dq(mut self, dq: f64) -> Self {
103 self.dq = dq.abs().max(1e-8);
104 self
105 }
106
107 pub fn with_grid_steps(mut self, dx: &[f64; 2]) -> Self {
110 self.dq = dx[0].abs().max(1e-8);
111 self.dv = dx[1].abs().max(1e-8);
112 self
113 }
114
115 pub fn build_process(&self) -> HestonProcess {
119 HestonProcess::new(
120 0.0,
121 self.v_kappa,
122 self.v_theta,
123 self.v_xi,
124 self.rho,
125 self.price_scale,
126 self.v_theta,
127 )
128 }
129
130 pub fn get_spreads(&self, _q: f64, _v: f64, grads: &Gradients<2>) -> (f64, f64) {
131 let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa_m).ln();
138
139 let dv_dq_buy = grads.fwd[0];
140 let dv_dq_sell = grads.bwd[0];
141
142 let mut delta_bid = base_spread - dv_dq_buy;
143 let mut delta_ask = base_spread + dv_dq_sell;
144
145 let min_spread = -5.0;
146 let max_spread = 10.0;
147 delta_bid = delta_bid.max(min_spread).min(max_spread);
148 delta_ask = delta_ask.max(min_spread).min(max_spread);
149
150 (delta_bid, delta_ask)
151 }
152}
153
154impl Model<2> for Heston {
155 type Process = market_model::process::HestonProcess;
156
157 fn process(&self) -> market_model::process::HestonProcess {
158 self.build_process()
159 }
160
161 fn optimize(&self, state: &[f64; 2], grads: &Gradients<2>) -> ControlOutput<2> {
162 let q = state[0];
163 let v = state[1];
164
165 let (d_bid, d_ask) = self.get_spreads(q, v, grads);
167
168 let lambda_bid = self.a * (-self.kappa_m * d_bid).exp();
169 let lambda_ask = self.a * (-self.kappa_m * d_ask).exp();
170
171 let lambda_bid = if q >= self.q_max { 0.0 } else { lambda_bid };
175 let lambda_ask = if q <= self.q_min { 0.0 } else { lambda_ask };
176
177 let hamiltonian_val = (lambda_bid + lambda_ask) / (self.gamma + self.kappa_m);
180
181 let risk_penalty = -0.5 * self.gamma * self.price_scale.powi(2) * v * q.powi(2);
184
185 let drift_correction = self.dq * (lambda_bid * grads.fwd[0] - lambda_ask * grads.bwd[0]);
188
189 let mu_v = self.v_kappa * (self.v_theta - v)
198 - self.gamma * self.price_scale * self.rho * self.v_xi * v * q;
199 let sigma2_v = self.v_xi.powi(2) * v;
200
201 let diff_term = sigma2_v / (2.0 * self.dv.powi(2));
202 let drift_term_abs = mu_v.abs() / self.dv;
203
204 let mut lambda_v_plus = diff_term;
206 let mut lambda_v_minus = diff_term;
207
208 if mu_v > 0.0 {
209 lambda_v_plus += drift_term_abs;
210 } else {
211 lambda_v_minus += drift_term_abs;
212 }
213
214 ControlOutput {
215 lambda_plus: [lambda_bid, lambda_v_plus],
216 lambda_minus: [lambda_ask, lambda_v_minus],
217 flow: hamiltonian_val + risk_penalty - drift_correction,
218 }
219 }
220
221 fn terminal(&self, state: &[f64; 2]) -> f64 {
222 match self.terminal_condition {
223 TerminalCondition::Zero => 0.0,
224 TerminalCondition::LiquidationCost => {
225 let q = state[0];
228 let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa_m).ln();
229 -q.abs() * base_spread
230 }
231 }
232 }
233
234 fn constant_discount_rate(&self) -> Option<f64> {
235 Some(0.0)
236 }
237
238 fn next_step(&self, current_state: &[f64; 2], dt: f64, noise: &[f64; 2]) -> [f64; 2] {
239 let q = current_state[0];
240 let v = current_state[1];
241 let sqrt_dt = dt.sqrt();
242
243 let next_q = q;
246
247 let v_pos = v.max(1e-5);
254 let mu_v = self.v_kappa * (self.v_theta - v_pos)
255 - self.gamma * self.price_scale * self.rho * self.v_xi * v_pos * q;
256 let drift = mu_v * dt;
257 let diffusion = self.v_xi * v_pos.sqrt() * sqrt_dt * noise[1];
258
259 let mut next_v = v_pos + drift + diffusion;
260
261 if next_v < 0.0 {
262 next_v = -next_v;
263 }
264
265 [next_q, next_v]
266 }
267
268 fn is_diffusion_dimension(&self, dim: usize) -> bool {
269 dim == 1 }
271
272 fn is_integer_dimension(&self, dim: usize) -> bool {
273 dim == 0 }
275
276 fn gradient_step(&self, dim: usize) -> f64 {
277 if dim == 1 {
278 self.dv.abs().max(1e-8)
279 } else {
280 1.0
281 }
282 }
283
284 fn fill_rate_base(&self, _state: &[f64; 2]) -> f64 {
285 self.a
286 }
287
288 fn fill_rate_decay(&self) -> f64 {
289 self.kappa_m
290 }
291}
292
293#[cfg(test)]
294mod tests {
295 use super::super::traits::{Gradients, Model};
296 use super::*;
297 use market_model::process::StochasticProcess;
298 use market_model::process::heston::HestonProcess;
299 use ndarray::array;
300
301 fn zero_grads() -> Gradients<2> {
302 Gradients {
303 fwd: [0.0; 2],
304 bwd: [0.0; 2],
305 }
306 }
307
308 fn default_model() -> Heston {
309 Heston {
310 gamma: 0.5,
311 kappa_m: 1.5,
312 a: 10.0,
313 v_kappa: 2.0,
314 v_theta: 0.25,
315 v_xi: 0.3,
316 rho: 0.0,
317 price_scale: 1.0,
318 dv: 0.02,
319 dq: 1.0,
320 q_min: f64::NEG_INFINITY,
321 q_max: f64::INFINITY,
322 terminal_condition: TerminalCondition::Zero,
323 }
324 }
325
326 #[test]
327 fn zero_gradient_gives_base_spread() {
328 let m = default_model();
329 let grads = zero_grads();
330 let (bid, ask) = m.get_spreads(0.0, 0.25, &grads);
331 let base = (1.0 / m.gamma) * (1.0 + m.gamma / m.kappa_m).ln();
332 assert!((bid - base).abs() < 1e-12);
333 assert!((ask - base).abs() < 1e-12);
334 }
335
336 #[test]
337 fn base_spread_independent_of_variance_at_zero_gradient() {
338 let m = default_model();
339 let grads = zero_grads();
340 let (bid_low, ask_low) = m.get_spreads(0.0, 0.05, &grads);
341 let (bid_high, ask_high) = m.get_spreads(0.0, 0.80, &grads);
342 assert!((bid_low - bid_high).abs() < 1e-12);
343 assert!((ask_low - ask_high).abs() < 1e-12);
344 }
345
346 #[test]
347 fn risk_penalty_scales_with_variance() {
348 let m = default_model();
349 let grads = zero_grads();
350 let ctrl_low = m.optimize(&[2.0, 0.05], &grads);
351 let ctrl_high = m.optimize(&[2.0, 0.50], &grads);
352 assert!(
353 ctrl_high.flow < ctrl_low.flow,
354 "Higher variance should decrease flow: low={}, high={}",
355 ctrl_low.flow,
356 ctrl_high.flow
357 );
358 }
359
360 #[test]
361 fn variance_transport_uses_upwinding() {
362 let m = default_model();
363 let grads = zero_grads();
364 let ctrl = m.optimize(&[0.0, 0.05], &grads);
365 assert!(
366 ctrl.lambda_plus[1] > ctrl.lambda_minus[1],
367 "Upwind: positive drift should give lambda_plus > lambda_minus"
368 );
369
370 let ctrl = m.optimize(&[0.0, 0.80], &grads);
371 assert!(
372 ctrl.lambda_minus[1] > ctrl.lambda_plus[1],
373 "Upwind: negative drift should give lambda_minus > lambda_plus"
374 );
375 }
376
377 #[test]
378 fn rho_affects_variance_drift() {
379 let grads = zero_grads();
380 let state = [2.0, 0.25]; let ctrl0 = Heston {
383 rho: 0.0,
384 ..default_model()
385 }
386 .optimize(&state, &grads);
387 let ctrl_pos = Heston {
388 rho: 0.5,
389 ..default_model()
390 }
391 .optimize(&state, &grads);
392 let ctrl_neg = Heston {
393 rho: -0.5,
394 ..default_model()
395 }
396 .optimize(&state, &grads);
397
398 assert!(
399 ctrl_pos.lambda_minus[1] > ctrl0.lambda_minus[1],
400 "Positive rho with positive q should increase downward variance transport"
401 );
402 assert!(
403 ctrl_neg.lambda_plus[1] > ctrl0.lambda_plus[1],
404 "Negative rho with positive q should increase upward variance transport"
405 );
406 }
407
408 #[test]
409 fn rho_has_no_effect_at_zero_inventory() {
410 let grads = zero_grads();
411 let state = [0.0, 0.25]; let ctrl0 = Heston {
414 rho: 0.0,
415 ..default_model()
416 }
417 .optimize(&state, &grads);
418 let ctrl_pos = Heston {
419 rho: 0.8,
420 ..default_model()
421 }
422 .optimize(&state, &grads);
423
424 assert!((ctrl0.lambda_plus[1] - ctrl_pos.lambda_plus[1]).abs() < 1e-12);
425 assert!((ctrl0.lambda_minus[1] - ctrl_pos.lambda_minus[1]).abs() < 1e-12);
426 }
427
428 #[test]
429 fn cir_next_step_stays_positive() {
430 let m = default_model();
431 let state = [0.0, 0.001];
432 let next = m.next_step(&state, 0.01, &[0.0, -5.0]);
433 assert!(
434 next[1] > 0.0,
435 "CIR reflection should keep variance positive: {}",
436 next[1]
437 );
438 }
439
440 #[test]
441 fn next_step_mean_reverts_variance() {
442 let m = default_model();
443 let state = [0.0, 0.05];
445 let mut sum = 0.0;
446 for i in 0..1000 {
447 let noise = [0.0, (i as f64 * 0.01).sin()];
448 let next = m.next_step(&state, 0.01, &noise);
449 sum += next[1] - state[1];
450 }
451 assert!(
452 sum / 1000.0 > 0.0,
453 "Variance should drift upward when below theta"
454 );
455 }
456
457 #[test]
458 fn builder_methods_work() {
459 let m = Heston::new(0.5, 1.5, 10.0)
460 .with_variance_params(3.0, 0.1, 0.5)
461 .with_rho(-0.3)
462 .with_dv(0.05)
463 .with_dq(2.0);
464 assert_eq!(m.v_kappa, 3.0);
465 assert_eq!(m.rho, -0.3);
466 assert_eq!(m.dv, 0.05);
467 }
468
469 #[test]
470 fn with_grid_steps_sets_both() {
471 let m = Heston::new(0.5, 1.5, 10.0).with_grid_steps(&[2.0, 0.05]);
472 assert_eq!(m.dq, 2.0);
473 assert_eq!(m.dv, 0.05);
474 }
475
476 #[test]
477 fn risk_penalty_scales_with_price_scale_squared() {
478 let grads = zero_grads();
479 let state = [3.0, 0.25];
480
481 let m1 = Heston {
482 rho: 0.0,
483 price_scale: 1.0,
484 ..default_model()
485 };
486 let m10 = Heston {
487 rho: 0.0,
488 price_scale: 10.0,
489 ..default_model()
490 };
491
492 let c1 = m1.optimize(&state, &grads);
493 let c10 = m10.optimize(&state, &grads);
494
495 let penalty1 = -0.5 * m1.gamma * m1.price_scale.powi(2) * state[1] * state[0].powi(2);
496 let penalty10 = -0.5 * m10.gamma * m10.price_scale.powi(2) * state[1] * state[0].powi(2);
497 let expected_delta = penalty10 - penalty1;
498 let observed_delta = c10.flow - c1.flow;
499
500 assert!(
501 (observed_delta - expected_delta).abs() < 1e-10,
502 "Risk penalty scaling mismatch: observed_delta={}, expected_delta={}",
503 observed_delta,
504 expected_delta
505 );
506 }
507
508 #[test]
509 fn cross_term_scales_linearly_with_price_scale() {
510 let grads = zero_grads();
511 let state = [2.0, 0.25];
512 let base = Heston {
513 rho: -0.7,
514 v_xi: 0.8,
515 dv: 0.05,
516 price_scale: 1.0,
517 ..default_model()
518 };
519 let scaled = Heston {
520 price_scale: 5.0,
521 ..base
522 };
523
524 let c1 = base.optimize(&state, &grads);
525 let c5 = scaled.optimize(&state, &grads);
526
527 let diff1 = c1.lambda_plus[1] - c1.lambda_minus[1];
528 let diff5 = c5.lambda_plus[1] - c5.lambda_minus[1];
529
530 assert!(
531 diff1 > 0.0 && diff5 > 0.0,
532 "Expected positive variance-drift asymmetry from cross term, got diff1={}, diff5={}",
533 diff1,
534 diff5
535 );
536
537 let ratio = diff5 / diff1;
538 assert!(
539 (ratio - 5.0).abs() < 1e-10,
540 "Cross-term scaling should be linear in price_scale: ratio={}",
541 ratio
542 );
543 }
544
545 #[test]
546 fn process_diffusion_covariance_matches_multiplicative_form() {
547 let process = HestonProcess::new(0.0, 1.0, 0.25, 0.8, -0.7, 100.0, 0.36);
548 let s0 = 100.0;
549 let v = 0.36;
550 let x = array![s0, v];
551
552 let l = process.diffusion(&x, 0.0);
553 let cov = l.dot(&l.t());
554
555 let var_s = cov[[0, 0]];
556 let cov_sv = cov[[0, 1]];
557
558 let expected_var_s = s0.powi(2) * v;
559 let expected_cov_sv = s0 * process.rho * process.sigma * v;
560
561 assert!(
562 (var_s - expected_var_s).abs() < 1e-10,
563 "Var(dS)/dt mismatch: observed={}, expected={}",
564 var_s,
565 expected_var_s
566 );
567 assert!(
568 (cov_sv - expected_cov_sv).abs() < 1e-10,
569 "Cov(dS,dv)/dt mismatch: observed={}, expected={}",
570 cov_sv,
571 expected_cov_sv
572 );
573 }
574}