1use std::collections::{BTreeMap, BTreeSet};
8
9use chrono::{Duration, NaiveDateTime, NaiveTime};
10use qs_core::types::Side;
11use serde::{Deserialize, Serialize};
12use thiserror::Error;
13
14pub const MAX_POLICIES: usize = 64;
16pub const MAX_GROUPS: usize = 64;
17pub const MAX_GROUP_SYMBOLS: usize = 256;
19pub const MAX_GROUP_ID_BYTES: usize = 64;
21
22const RISK_EPSILON: f64 = 1e-9;
24
25#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
27#[serde(deny_unknown_fields)]
28pub struct CorrelationGroup {
29 pub id: String,
30 pub symbols: BTreeSet<String>,
31}
32
33#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
37#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
38pub enum RiskPolicy {
39 MaxOpenPositions { limit: usize },
41 MaxOpenPerSymbol { limit: usize },
43 GroupRiskCap { group: String, max_group_risk: f64 },
45 DailyLossHalt {
47 max_loss: LossLimit,
48 reset_at_utc: NaiveTime,
49 },
50 KillSwitch {
52 max_drawdown_percent: f64,
53 action: HaltAction,
54 },
55}
56
57impl RiskPolicy {
58 pub const fn name(&self) -> &'static str {
60 match self {
61 Self::MaxOpenPositions { .. } => "max_open_positions",
62 Self::MaxOpenPerSymbol { .. } => "max_open_per_symbol",
63 Self::GroupRiskCap { .. } => "group_risk_cap",
64 Self::DailyLossHalt { .. } => "daily_loss_halt",
65 Self::KillSwitch { .. } => "kill_switch",
66 }
67 }
68}
69
70#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
72#[serde(rename_all = "snake_case", deny_unknown_fields)]
73pub enum LossLimit {
74 AccountPercent(f64),
76 Amount(f64),
78 RiskMultiples(f64),
80}
81
82#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
84#[serde(rename_all = "snake_case")]
85pub enum HaltAction {
86 Halt,
88 HaltAndCloseAll,
90}
91
92#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
94#[serde(rename_all = "snake_case")]
95pub enum IntentKind {
96 Entry,
97 ScaleIn,
98}
99
100#[derive(Debug, Clone, PartialEq)]
102pub struct ExposureIntent<'a> {
103 pub symbol: &'a str,
104 pub side: Side,
105 pub kind: IntentKind,
106 pub requested_risk: Option<f64>,
108}
109
110#[derive(Debug, Clone, PartialEq)]
112pub struct ExposureFact {
113 pub symbol: String,
114 pub side: Side,
115 pub risk: Option<f64>,
117}
118
119#[derive(Debug, Clone, Copy, PartialEq)]
121pub struct PortfolioFacts<'a> {
122 pub now: NaiveDateTime,
123 pub balance: f64,
125 pub drawdown_fraction: Option<f64>,
127 pub day_realized_r: f64,
129 pub open: &'a [ExposureFact],
130 pub pending: &'a [ExposureFact],
131 pub reserved: &'a [ExposureFact],
133}
134
135#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
137#[serde(tag = "verdict", rename_all = "snake_case")]
138pub enum Verdict {
139 Approve,
140 Reject { policy: String, reason: String },
141}
142
143impl Verdict {
144 fn reject(policy: &str, reason: impl Into<String>) -> Self {
145 Self::Reject {
146 policy: policy.to_owned(),
147 reason: reason.into(),
148 }
149 }
150
151 pub const fn is_approved(&self) -> bool {
152 matches!(self, Self::Approve)
153 }
154}
155
156#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
158#[serde(rename_all = "snake_case")]
159pub enum HaltCommand {
160 CancelAllPending,
161 CloseAll,
162}
163
164#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
166pub struct HaltInterval {
167 pub policy: String,
168 pub from: NaiveDateTime,
169 pub to: Option<NaiveDateTime>,
171}
172
173#[derive(Debug, Clone, PartialEq, Error)]
175pub enum RiskConfigError {
176 #[error("too many {what}: {count} exceeds {max}")]
177 TooMany {
178 what: &'static str,
179 count: usize,
180 max: usize,
181 },
182 #[error(
183 "group id '{id}' must be 1 to {MAX_GROUP_ID_BYTES} bytes of ASCII letters, digits, '_', '-', or '.'"
184 )]
185 InvalidGroupId { id: String },
186 #[error("group '{id}' is declared more than once")]
187 DuplicateGroup { id: String },
188 #[error("group '{id}' must list between 1 and {MAX_GROUP_SYMBOLS} non-empty symbols")]
189 InvalidGroupSymbols { id: String },
190 #[error("{policy} refers to undeclared group '{group}'")]
191 UnknownGroup { policy: &'static str, group: String },
192 #[error("{policy}: {field} must be {requirement}, got {value}")]
193 InvalidValue {
194 policy: &'static str,
195 field: &'static str,
196 requirement: &'static str,
197 value: String,
198 },
199}
200
201#[derive(Debug, Clone, PartialEq)]
202struct DailyState {
203 day_start: NaiveDateTime,
204 day_start_balance: f64,
205 halted: bool,
206}
207
208#[derive(Debug, Clone, PartialEq)]
210pub struct PortfolioSupervisor {
211 policies: Vec<RiskPolicy>,
212 groups: BTreeMap<String, BTreeSet<String>>,
213 daily: Option<DailyState>,
214 tripped_kill_switches: BTreeSet<usize>,
216 close_all_issued: bool,
217 intervals: Vec<HaltInterval>,
218}
219
220impl PortfolioSupervisor {
221 pub fn new(
223 policies: Vec<RiskPolicy>,
224 groups: Vec<CorrelationGroup>,
225 ) -> Result<Self, RiskConfigError> {
226 if policies.len() > MAX_POLICIES {
227 return Err(RiskConfigError::TooMany {
228 what: "policies",
229 count: policies.len(),
230 max: MAX_POLICIES,
231 });
232 }
233 if groups.len() > MAX_GROUPS {
234 return Err(RiskConfigError::TooMany {
235 what: "groups",
236 count: groups.len(),
237 max: MAX_GROUPS,
238 });
239 }
240 let mut declared = BTreeMap::new();
241 for group in groups {
242 if !valid_group_id(&group.id) {
243 return Err(RiskConfigError::InvalidGroupId { id: group.id });
244 }
245 if group.symbols.is_empty()
246 || group.symbols.len() > MAX_GROUP_SYMBOLS
247 || group.symbols.iter().any(|symbol| symbol.trim().is_empty())
248 {
249 return Err(RiskConfigError::InvalidGroupSymbols { id: group.id });
250 }
251 if declared.contains_key(&group.id) {
252 return Err(RiskConfigError::DuplicateGroup { id: group.id });
253 }
254 declared.insert(group.id, group.symbols);
255 }
256 for policy in &policies {
257 validate_policy(policy, &declared)?;
258 }
259 let resets = policies
260 .iter()
261 .filter_map(|policy| match policy {
262 RiskPolicy::DailyLossHalt { reset_at_utc, .. } => Some(*reset_at_utc),
263 _ => None,
264 })
265 .collect::<BTreeSet<_>>();
266 if resets.len() > 1 {
267 return Err(RiskConfigError::InvalidValue {
268 policy: "daily_loss_halt",
269 field: "reset_at_utc",
270 requirement: "the same instant for every daily loss policy",
271 value: resets
272 .iter()
273 .map(ToString::to_string)
274 .collect::<Vec<_>>()
275 .join(", "),
276 });
277 }
278 Ok(Self {
279 policies,
280 groups: declared,
281 daily: None,
282 tripped_kill_switches: BTreeSet::new(),
283 close_all_issued: false,
284 intervals: Vec::new(),
285 })
286 }
287
288 pub fn policies(&self) -> &[RiskPolicy] {
289 &self.policies
290 }
291
292 pub fn caps_group_risk(&self) -> bool {
294 self.policies
295 .iter()
296 .any(|policy| matches!(policy, RiskPolicy::GroupRiskCap { .. }))
297 }
298
299 pub fn day_start(&self) -> Option<NaiveDateTime> {
301 self.daily.as_ref().map(|daily| daily.day_start)
302 }
303
304 pub fn halted(&self) -> bool {
306 !self.tripped_kill_switches.is_empty()
307 || self.daily.as_ref().is_some_and(|daily| daily.halted)
308 }
309
310 pub fn begin(&mut self, now: NaiveDateTime, balance: f64) {
312 let Some(reset) = self.daily_reset() else {
313 return;
314 };
315 let day_start = latest_reset_at_or_before(now, reset);
316 match &mut self.daily {
317 Some(daily) if daily.day_start == day_start => {}
318 Some(daily) => {
319 if daily.halted {
320 close_interval(&mut self.intervals, "daily_loss_halt", day_start);
321 }
322 *daily = DailyState {
323 day_start,
324 day_start_balance: balance,
325 halted: false,
326 };
327 }
328 None => {
329 self.daily = Some(DailyState {
330 day_start,
331 day_start_balance: balance,
332 halted: false,
333 });
334 }
335 }
336 }
337
338 pub fn on_boundary(&mut self, facts: &PortfolioFacts<'_>) -> Vec<HaltCommand> {
340 self.begin(facts.now, facts.balance);
341 let mut commands = Vec::new();
342 let mut halt_started = false;
343 for (index, policy) in self.policies.iter().enumerate() {
344 match policy {
345 RiskPolicy::DailyLossHalt { max_loss, .. } => {
346 let Some(daily) = self.daily.as_mut() else {
347 continue;
348 };
349 if daily.halted {
350 continue;
351 }
352 let breached = match max_loss {
353 LossLimit::AccountPercent(percent) => {
354 daily.day_start_balance - facts.balance
355 >= daily.day_start_balance * percent / 100.0 - RISK_EPSILON
356 }
357 LossLimit::Amount(amount) => {
358 daily.day_start_balance - facts.balance >= amount - RISK_EPSILON
359 }
360 LossLimit::RiskMultiples(multiples) => {
361 -facts.day_realized_r >= multiples - RISK_EPSILON
362 }
363 };
364 if breached {
365 daily.halted = true;
366 halt_started = true;
367 self.intervals.push(HaltInterval {
368 policy: policy.name().to_owned(),
369 from: facts.now,
370 to: None,
371 });
372 }
373 }
374 RiskPolicy::KillSwitch {
375 max_drawdown_percent,
376 action,
377 } => {
378 if self.tripped_kill_switches.contains(&index) {
380 continue;
381 }
382 let Some(drawdown) = facts.drawdown_fraction else {
383 continue;
384 };
385 if drawdown * 100.0 >= max_drawdown_percent - RISK_EPSILON {
386 let first_trip = self.tripped_kill_switches.is_empty();
387 self.tripped_kill_switches.insert(index);
388 halt_started |= first_trip;
389 self.intervals.push(HaltInterval {
390 policy: policy.name().to_owned(),
391 from: facts.now,
392 to: None,
393 });
394 if *action == HaltAction::HaltAndCloseAll && !self.close_all_issued {
395 self.close_all_issued = true;
396 commands.push(HaltCommand::CloseAll);
397 }
398 }
399 }
400 _ => {}
401 }
402 }
403 if halt_started {
404 commands.insert(0, HaltCommand::CancelAllPending);
405 }
406 commands
407 }
408
409 pub fn review(&self, facts: &PortfolioFacts<'_>, intent: &ExposureIntent<'_>) -> Verdict {
411 if !self.tripped_kill_switches.is_empty() {
412 return Verdict::reject(
413 "kill_switch",
414 "new exposure is halted for the rest of the run",
415 );
416 }
417 if let Some(daily) = self.daily.as_ref().filter(|daily| daily.halted) {
418 return Verdict::reject(
419 "daily_loss_halt",
420 format!(
421 "new exposure is halted until the reset after {}",
422 daily.day_start
423 ),
424 );
425 }
426 let counted = || facts.open.iter().chain(facts.pending).chain(facts.reserved);
427 for policy in &self.policies {
428 match policy {
429 RiskPolicy::MaxOpenPositions { limit } if intent.kind == IntentKind::Entry => {
430 let count = counted().count();
431 if count >= *limit {
432 return Verdict::reject(
433 policy.name(),
434 format!(
435 "{count} positions open, pending, or approved of limit {limit}"
436 ),
437 );
438 }
439 }
440 RiskPolicy::MaxOpenPerSymbol { limit } if intent.kind == IntentKind::Entry => {
441 let count = counted()
442 .filter(|fact| fact.symbol == intent.symbol)
443 .count();
444 if count >= *limit {
445 return Verdict::reject(
446 policy.name(),
447 format!(
448 "{count} positions open, pending, or approved on {} of limit {limit}",
449 intent.symbol
450 ),
451 );
452 }
453 }
454 RiskPolicy::GroupRiskCap {
455 group,
456 max_group_risk,
457 } => {
458 let symbols = &self.groups[group];
459 if !symbols.contains(intent.symbol) {
460 continue;
461 }
462 let Some(requested) = intent.requested_risk else {
463 return Verdict::reject(
464 policy.name(),
465 format!(
466 "risk_unmeasurable: the request's risk is unknown before the fill, so group '{group}' cannot be capped"
467 ),
468 );
469 };
470 let mut carried = 0.0;
471 for fact in counted().filter(|fact| symbols.contains(&fact.symbol)) {
472 match fact.risk {
473 Some(risk) => carried += risk,
474 None => {
475 return Verdict::reject(
476 policy.name(),
477 format!(
478 "group_risk_unmeasurable: a {} position in group '{group}' has unknown risk",
479 fact.symbol
480 ),
481 );
482 }
483 }
484 }
485 if carried + requested > max_group_risk + RISK_EPSILON {
486 return Verdict::reject(
487 policy.name(),
488 format!(
489 "group '{group}' carries {carried} and the request adds {requested}, above the cap of {max_group_risk}"
490 ),
491 );
492 }
493 }
494 _ => {}
495 }
496 }
497 Verdict::Approve
498 }
499
500 pub fn finish(self) -> Vec<HaltInterval> {
502 self.intervals
503 }
504
505 pub fn intervals(&self) -> &[HaltInterval] {
507 &self.intervals
508 }
509
510 fn daily_reset(&self) -> Option<NaiveTime> {
511 self.policies.iter().find_map(|policy| match policy {
512 RiskPolicy::DailyLossHalt { reset_at_utc, .. } => Some(*reset_at_utc),
513 _ => None,
514 })
515 }
516}
517
518fn close_interval(intervals: &mut [HaltInterval], policy: &str, at: NaiveDateTime) {
519 if let Some(interval) = intervals
520 .iter_mut()
521 .rev()
522 .find(|interval| interval.policy == policy && interval.to.is_none())
523 {
524 interval.to = Some(at);
525 }
526}
527
528fn latest_reset_at_or_before(now: NaiveDateTime, reset: NaiveTime) -> NaiveDateTime {
529 let today = now.date().and_time(reset);
530 if today <= now {
531 today
532 } else {
533 today - Duration::days(1)
534 }
535}
536
537fn valid_group_id(id: &str) -> bool {
538 !id.is_empty()
539 && id.len() <= MAX_GROUP_ID_BYTES
540 && id
541 .bytes()
542 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.'))
543}
544
545fn validate_policy(
546 policy: &RiskPolicy,
547 groups: &BTreeMap<String, BTreeSet<String>>,
548) -> Result<(), RiskConfigError> {
549 let invalid = |field: &'static str, requirement: &'static str, value: String| {
550 Err(RiskConfigError::InvalidValue {
551 policy: policy.name(),
552 field,
553 requirement,
554 value,
555 })
556 };
557 match policy {
558 RiskPolicy::MaxOpenPositions { limit } | RiskPolicy::MaxOpenPerSymbol { limit } => {
559 if *limit == 0 {
560 return invalid("limit", "at least 1", limit.to_string());
561 }
562 }
563 RiskPolicy::GroupRiskCap {
564 group,
565 max_group_risk,
566 } => {
567 if !groups.contains_key(group) {
568 return Err(RiskConfigError::UnknownGroup {
569 policy: policy.name(),
570 group: group.clone(),
571 });
572 }
573 if !(max_group_risk.is_finite() && *max_group_risk > 0.0) {
574 return invalid(
575 "max_group_risk",
576 "finite and positive",
577 max_group_risk.to_string(),
578 );
579 }
580 }
581 RiskPolicy::DailyLossHalt { max_loss, .. } => {
582 let (field, value) = match max_loss {
583 LossLimit::AccountPercent(value) => ("account_percent", *value),
584 LossLimit::Amount(value) => ("amount", *value),
585 LossLimit::RiskMultiples(value) => ("risk_multiples", *value),
586 };
587 let valid = value.is_finite()
588 && value > 0.0
589 && (!matches!(max_loss, LossLimit::AccountPercent(_)) || value <= 100.0);
590 if !valid {
591 return invalid(
592 field,
593 "finite and positive, and a percent at most 100",
594 value.to_string(),
595 );
596 }
597 }
598 RiskPolicy::KillSwitch {
599 max_drawdown_percent,
600 ..
601 } => {
602 if !(max_drawdown_percent.is_finite()
603 && *max_drawdown_percent > 0.0
604 && *max_drawdown_percent <= 100.0)
605 {
606 return invalid(
607 "max_drawdown_percent",
608 "greater than 0 and at most 100",
609 max_drawdown_percent.to_string(),
610 );
611 }
612 }
613 }
614 Ok(())
615}
616
617#[cfg(test)]
618mod tests;