use std::hash::{Hash, Hasher};
use chrono::{DateTime, Utc};
use optionstratlib::OptionStyle;
use optionstratlib::prelude::Positive;
use serde::{Deserialize, Serialize};
use crate::error::ConfigError;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct InstrumentKey {
pub underlying: String,
pub expiration_utc: DateTime<Utc>,
pub strike: Positive,
pub style: OptionStyle,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct ContractSpecFingerprint {
pub contract_multiplier: u32,
pub settlement: SettlementStyle,
pub exercise: ExerciseStyle,
pub quote_currency: String,
pub venue_product_code: String,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[repr(u8)]
pub enum SettlementStyle {
Cash,
Physical,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[repr(u8)]
pub enum ExerciseStyle {
European,
American,
}
#[derive(Debug, Clone)]
pub struct Instrument {
pub key: InstrumentKey,
pub provider: ProviderId,
pub native_symbol: String,
pub stream_symbol: Option<String>,
pub spec: ContractSpecFingerprint,
}
impl PartialEq for Instrument {
fn eq(&self, other: &Self) -> bool {
self.key == other.key
}
}
impl Eq for Instrument {}
impl Hash for Instrument {
fn hash<H: Hasher>(&self, state: &mut H) {
self.key.hash(state);
}
}
pub const RESERVED_PROVIDER_IDS: [&str; 6] =
["deribit", "tastytrade", "dxlink", "ig", "alpaca", "ibkr"];
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
#[serde(try_from = "String", into = "String")]
pub struct ProviderId(String);
impl ProviderId {
#[must_use = "a validated provider id must be used"]
pub fn new(id: impl Into<String>) -> Result<Self, ConfigError> {
let id = id.into();
if is_valid_provider_id(&id) {
Ok(Self(id))
} else {
Err(ConfigError::InvalidValue {
field: "provider id".to_owned(),
reason: format!(
"`{id}` must match ^[a-z][a-z0-9]*(?:[_-][a-z0-9]+)*$ \
(2-32 chars, no leading/trailing/adjacent `-`/`_`)"
),
})
}
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
#[must_use]
pub fn is_reserved(&self) -> bool {
RESERVED_PROVIDER_IDS.contains(&self.0.as_str())
}
}
impl std::fmt::Display for ProviderId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
impl TryFrom<String> for ProviderId {
type Error = ConfigError;
fn try_from(value: String) -> Result<Self, Self::Error> {
Self::new(value)
}
}
impl From<ProviderId> for String {
fn from(value: ProviderId) -> Self {
value.0
}
}
fn is_valid_provider_id(id: &str) -> bool {
let len = id.chars().count();
if !(2..=32).contains(&len) {
return false;
}
let mut chars = id.chars();
match chars.next() {
Some(c) if c.is_ascii_lowercase() => {}
_ => return false,
}
let mut prev_was_separator = false;
for c in chars {
if c == '-' || c == '_' {
if prev_was_separator {
return false; }
prev_was_separator = true;
} else if c.is_ascii_lowercase() || c.is_ascii_digit() {
prev_was_separator = false;
} else {
return false; }
}
!prev_was_separator
}
#[cfg(test)]
mod tests {
use super::*;
#[track_caller]
fn pid(id: &str) -> ProviderId {
match ProviderId::new(id) {
Ok(p) => p,
Err(e) => panic!("expected a valid provider id `{id}`, got: {e}"),
}
}
#[track_caller]
fn utc(secs: i64) -> DateTime<Utc> {
match DateTime::<Utc>::from_timestamp(secs, 0) {
Some(t) => t,
None => panic!("invalid test timestamp: {secs}"),
}
}
#[track_caller]
fn strike(value: f64) -> Positive {
match Positive::new(value) {
Ok(p) => p,
Err(e) => panic!("invalid test strike `{value}`: {e}"),
}
}
fn sample_key() -> InstrumentKey {
InstrumentKey {
underlying: "BTC".to_owned(),
expiration_utc: utc(1_700_000_000),
strike: strike(60_000.0),
style: OptionStyle::Call,
}
}
fn sample_spec(multiplier: u32) -> ContractSpecFingerprint {
ContractSpecFingerprint {
contract_multiplier: multiplier,
settlement: SettlementStyle::Cash,
exercise: ExerciseStyle::European,
quote_currency: "USD".to_owned(),
venue_product_code: "BTC".to_owned(),
}
}
fn hash_of<T: Hash>(value: &T) -> u64 {
use std::collections::hash_map::DefaultHasher;
let mut hasher = DefaultHasher::new();
value.hash(&mut hasher);
hasher.finish()
}
#[test]
fn test_instrument_eq_collapses_same_key_different_provider_symbol_spec() {
let key = sample_key();
let rest = Instrument {
key: key.clone(),
provider: pid("deribit"),
native_symbol: "BTC-27JUN25-60000-C".to_owned(),
stream_symbol: None,
spec: sample_spec(1),
};
let stream = Instrument {
key,
provider: pid("dxlink"),
native_symbol: ".BTC250627C60000".to_owned(),
stream_symbol: Some("dxfeed-symbol".to_owned()),
spec: sample_spec(100),
};
assert_eq!(rest, stream);
assert_eq!(hash_of(&rest), hash_of(&stream));
}
#[test]
fn test_instrument_ne_when_keys_differ() {
let base = Instrument {
key: sample_key(),
provider: pid("deribit"),
native_symbol: "a".to_owned(),
stream_symbol: None,
spec: sample_spec(1),
};
let other_key = InstrumentKey {
strike: strike(61_000.0),
..sample_key()
};
let other = Instrument {
key: other_key,
provider: pid("deribit"),
native_symbol: "a".to_owned(),
stream_symbol: None,
spec: sample_spec(1),
};
assert_ne!(base, other);
}
#[test]
fn test_instrument_key_ne_across_underlying() {
let a = sample_key();
let b = InstrumentKey {
underlying: "ETH".to_owned(),
..sample_key()
};
assert_ne!(a, b);
}
#[test]
fn test_instrument_key_ne_across_expiration() {
let a = sample_key();
let b = InstrumentKey {
expiration_utc: utc(1_700_086_400),
..sample_key()
};
assert_ne!(a, b);
}
#[test]
fn test_instrument_key_ne_across_strike() {
let a = sample_key();
let b = InstrumentKey {
strike: strike(60_500.0),
..sample_key()
};
assert_ne!(a, b);
}
#[test]
fn test_instrument_key_ne_across_style() {
let a = sample_key();
let b = InstrumentKey {
style: OptionStyle::Put,
..sample_key()
};
assert_ne!(a, b);
}
#[test]
fn test_instrument_key_eq_and_hash_equal_when_all_fields_match() {
let a = sample_key();
let b = sample_key();
assert_eq!(a, b);
assert_eq!(hash_of(&a), hash_of(&b));
}
#[test]
fn test_contract_spec_fingerprint_eq_by_value() {
assert_eq!(sample_spec(100), sample_spec(100));
assert_ne!(sample_spec(100), sample_spec(1));
}
#[test]
fn test_reserved_provider_ids_membership_is_exactly_six() {
assert_eq!(RESERVED_PROVIDER_IDS.len(), 6);
for id in ["deribit", "tastytrade", "dxlink", "ig", "alpaca", "ibkr"] {
assert!(pid(id).is_reserved(), "`{id}` should be reserved");
}
}
#[test]
fn test_provider_id_is_reserved_false_for_custom_id() {
assert!(!pid("my-broker").is_reserved());
}
#[test]
fn test_provider_id_accepts_all_built_in_ids() {
for id in RESERVED_PROVIDER_IDS {
assert!(ProviderId::new(id).is_ok(), "`{id}` should be valid");
}
}
#[test]
fn test_provider_id_accepts_documented_examples() {
for id in ["my-broker", "my_broker", "td-ameritrade"] {
assert!(ProviderId::new(id).is_ok(), "`{id}` should be valid");
}
}
#[test]
fn test_provider_id_rejects_uppercase() {
assert!(ProviderId::new("Deribit").is_err());
}
#[test]
fn test_provider_id_rejects_leading_digit() {
assert!(ProviderId::new("1broker").is_err());
}
#[test]
fn test_provider_id_rejects_empty() {
assert!(ProviderId::new("").is_err());
}
#[test]
fn test_provider_id_rejects_too_short() {
assert!(ProviderId::new("a").is_err());
}
#[test]
fn test_provider_id_rejects_too_long() {
assert!(ProviderId::new("a".repeat(33)).is_err());
}
#[test]
fn test_provider_id_accepts_max_length() {
assert!(ProviderId::new("a".repeat(32)).is_ok());
}
#[test]
fn test_provider_id_rejects_leading_separator() {
assert!(ProviderId::new("-broker").is_err());
assert!(ProviderId::new("_broker").is_err());
}
#[test]
fn test_provider_id_rejects_trailing_separator() {
assert!(ProviderId::new("broker-").is_err());
assert!(ProviderId::new("broker_").is_err());
}
#[test]
fn test_provider_id_rejects_adjacent_separators() {
assert!(ProviderId::new("a--b").is_err());
assert!(ProviderId::new("a__b").is_err());
assert!(ProviderId::new("a-_b").is_err());
assert!(ProviderId::new("a_-b").is_err());
}
#[test]
fn test_provider_id_rejects_non_ascii() {
assert!(ProviderId::new("brok\u{00e9}r").is_err());
}
#[test]
fn test_provider_id_new_error_names_field_provider_id() {
match ProviderId::new("Bad") {
Err(ConfigError::InvalidValue { field, .. }) => assert_eq!(field, "provider id"),
other => panic!("expected InvalidValue on provider id, got {other:?}"),
}
}
#[test]
fn test_provider_id_as_str_returns_inner() {
assert_eq!(pid("deribit").as_str(), "deribit");
}
#[test]
fn test_provider_id_ordering_delegates_to_inner_string() {
assert!(pid("alpaca") < pid("deribit"));
}
#[test]
fn test_provider_id_try_from_string_revalidates() {
assert!(ProviderId::try_from("deribit".to_owned()).is_ok());
assert!(ProviderId::try_from("Deribit".to_owned()).is_err());
}
#[test]
fn test_provider_id_into_string_returns_inner() {
let raw: String = pid("my-broker").into();
assert_eq!(raw, "my-broker");
}
#[test]
fn test_provider_id_serde_roundtrips_through_string() {
#[derive(Serialize, Deserialize)]
struct Wrap {
id: ProviderId,
}
let rendered = match toml::to_string(&Wrap {
id: pid("my-broker"),
}) {
Ok(s) => s,
Err(e) => panic!("serialize failed: {e}"),
};
assert!(rendered.contains("my-broker"));
match toml::from_str::<Wrap>(&rendered) {
Ok(w) => assert_eq!(w.id.as_str(), "my-broker"),
Err(e) => panic!("deserialize failed: {e}"),
}
}
#[test]
fn test_provider_id_serde_rejects_malformed_on_deserialize() {
#[derive(Deserialize)]
struct Wrap {
#[allow(dead_code)]
id: ProviderId,
}
assert!(toml::from_str::<Wrap>("id = \"BadId\"\n").is_err());
}
}