use std::collections::HashMap;
use std::fmt;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use optionstratlib::ExpirationDate;
use tokio::sync::mpsc;
use tokio::sync::mpsc::error::TrySendError;
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use crate::chain::{
ChainFetch, DepthLadder, GreeksRow, Instrument, InstrumentKey, MarketUpdate, ProviderId,
QuoteUpdate,
};
use crate::error::ProviderError;
pub(crate) mod deribit;
#[cfg_attr(not(any(feature = "tastytrade", feature = "dxlink")), allow(dead_code))]
pub(crate) mod dxfeed_decode;
#[cfg(feature = "tastytrade")]
pub(crate) mod tastytrade;
#[cfg(feature = "alpaca")]
pub(crate) mod alpaca;
#[cfg(feature = "dxlink")]
pub(crate) mod dxlink;
#[cfg(feature = "ig")]
pub(crate) mod ig;
#[cfg(feature = "ibkr")]
pub(crate) mod ibkr;
#[async_trait]
pub trait Provider: Send + Sync {
fn id(&self) -> ProviderId;
fn capabilities(&self) -> ProviderCapabilities;
async fn discover(&self) -> Result<Vec<UnderlyingRef>, ProviderError>;
async fn fetch_chain(
&self,
underlying: &str,
expiration: &ExpirationDate,
) -> Result<ChainFetch, ProviderError>;
async fn subscribe(
&self,
req: SubscriptionRequest,
sink: MarketUpdateSink,
) -> Result<SubscriptionHandle, ProviderError>;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct ProviderCapabilities {
pub chain: ChainCapability,
pub depth: bool,
pub greeks: GreeksCapability,
pub option_stream: OptionStreamCapability,
pub underlying_stream: bool,
pub chain_poll: ChainPollCapability,
pub trades_tape: bool,
pub auth: AuthKind,
}
impl ProviderCapabilities {
#[must_use]
pub fn builder() -> ProviderCapabilitiesBuilder {
ProviderCapabilitiesBuilder::default()
}
}
#[derive(Debug, Clone, Default)]
pub struct ProviderCapabilitiesBuilder {
chain: ChainCapability,
depth: bool,
greeks: GreeksCapability,
option_stream: OptionStreamCapability,
underlying_stream: bool,
chain_poll: ChainPollCapability,
trades_tape: bool,
auth: AuthKind,
}
impl ProviderCapabilitiesBuilder {
#[must_use = "builders do nothing unless .build() is called"]
pub fn chain(mut self, chain: ChainCapability) -> Self {
self.chain = chain;
self
}
#[must_use = "builders do nothing unless .build() is called"]
pub fn depth(mut self, depth: bool) -> Self {
self.depth = depth;
self
}
#[must_use = "builders do nothing unless .build() is called"]
pub fn greeks(mut self, greeks: GreeksCapability) -> Self {
self.greeks = greeks;
self
}
#[must_use = "builders do nothing unless .build() is called"]
pub fn option_stream(mut self, option_stream: OptionStreamCapability) -> Self {
self.option_stream = option_stream;
self
}
#[must_use = "builders do nothing unless .build() is called"]
pub fn underlying_stream(mut self, underlying_stream: bool) -> Self {
self.underlying_stream = underlying_stream;
self
}
#[must_use = "builders do nothing unless .build() is called"]
pub fn chain_poll(mut self, chain_poll: ChainPollCapability) -> Self {
self.chain_poll = chain_poll;
self
}
#[must_use = "builders do nothing unless .build() is called"]
pub fn trades_tape(mut self, trades_tape: bool) -> Self {
self.trades_tape = trades_tape;
self
}
#[must_use = "builders do nothing unless .build() is called"]
pub fn auth(mut self, auth: AuthKind) -> Self {
self.auth = auth;
self
}
#[must_use]
pub fn build(self) -> ProviderCapabilities {
ProviderCapabilities {
chain: self.chain,
depth: self.depth,
greeks: self.greeks,
option_stream: self.option_stream,
underlying_stream: self.underlying_stream,
chain_poll: self.chain_poll,
trades_tape: self.trades_tape,
auth: self.auth,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[repr(u8)]
#[non_exhaustive]
pub enum ChainCapability {
Native,
Assemble,
Partial,
#[default]
None,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[repr(u8)]
#[non_exhaustive]
pub enum GreeksCapability {
Provided,
ComputedLocally,
#[default]
None,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum OptionStreamCapability {
#[default]
None,
SymbolOnly {
verified: bool,
},
ChainQuotes {
verified: bool,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum ChainPollCapability {
#[default]
None,
Poll {
interval_hint_secs: u32,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[repr(u8)]
#[non_exhaustive]
pub enum AuthKind {
#[default]
None,
Token,
KeySecret,
UserPass,
OAuth,
}
#[derive(Debug, Clone)]
pub struct UnderlyingRef {
pub underlying: String,
pub expirations: Vec<ExpirationDate>,
}
impl UnderlyingRef {
#[must_use]
pub fn new(underlying: impl Into<String>) -> Self {
Self {
underlying: underlying.into(),
expirations: Vec::new(),
}
}
#[must_use]
pub fn with_expirations(
underlying: impl Into<String>,
expirations: Vec<ExpirationDate>,
) -> Self {
Self {
underlying: underlying.into(),
expirations,
}
}
}
#[derive(Debug, Clone)]
pub struct SubscriptionRequest {
pub underlying: String,
pub expiration_utc: DateTime<Utc>,
pub instruments: Vec<Instrument>,
pub cancel: CancellationToken,
}
impl SubscriptionRequest {
#[must_use]
pub fn new(
underlying: impl Into<String>,
expiration_utc: DateTime<Utc>,
instruments: Vec<Instrument>,
cancel: CancellationToken,
) -> Self {
Self {
underlying: underlying.into(),
expiration_utc,
instruments,
cancel,
}
}
}
#[must_use = "dropping the handle immediately cancels the subscription"]
pub struct SubscriptionHandle {
cancel: Option<Box<dyn FnOnce() + Send>>,
join: Option<JoinHandle<()>>,
}
impl SubscriptionHandle {
#[must_use = "dropping the handle immediately cancels the subscription"]
pub fn new<F>(on_cancel: F) -> Self
where
F: FnOnce() + Send + 'static,
{
Self {
cancel: Some(Box::new(on_cancel)),
join: None,
}
}
#[must_use = "dropping the handle immediately cancels the subscription"]
pub fn spawned(cancel: CancellationToken, join: JoinHandle<()>) -> Self {
Self {
cancel: Some(Box::new(move || cancel.cancel())),
join: Some(join),
}
}
pub fn take_join_handle(&mut self) -> Option<JoinHandle<()>> {
let join = self.join.take();
if join.is_some() {
self.cancel = None;
}
join
}
pub fn abort(mut self) {
if let Some(cancel) = self.cancel.take() {
cancel();
}
}
}
impl fmt::Debug for SubscriptionHandle {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SubscriptionHandle")
.field("active", &self.cancel.is_some())
.field("supervised", &self.join.is_some())
.finish()
}
}
impl Drop for SubscriptionHandle {
fn drop(&mut self) {
if let Some(cancel) = self.cancel.take() {
cancel();
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum SendState {
Open,
Closed,
}
#[derive(Debug)]
pub struct MarketUpdateSink {
tx_control: mpsc::Sender<MarketUpdate>,
tx_coalesced: mpsc::Sender<MarketUpdate>,
staging: ProducerStaging,
}
impl MarketUpdateSink {
#[must_use]
pub fn new(
tx_control: mpsc::Sender<MarketUpdate>,
tx_coalesced: mpsc::Sender<MarketUpdate>,
) -> Self {
Self {
tx_control,
tx_coalesced,
staging: ProducerStaging::new(),
}
}
pub async fn send(&mut self, update: MarketUpdate) -> SendState {
match update {
control @ (MarketUpdate::Chain(_) | MarketUpdate::Health(_, _)) => {
self.send_control(control).await
}
coalesced @ (MarketUpdate::Quote(_)
| MarketUpdate::Greeks(_)
| MarketUpdate::Depth(_)) => self.publish_coalesced(coalesced),
}
}
pub(crate) async fn send_control(&mut self, update: MarketUpdate) -> SendState {
match self.tx_control.send(update).await {
Ok(()) => SendState::Open,
Err(_) => SendState::Closed,
}
}
pub(crate) fn publish_coalesced(&mut self, update: MarketUpdate) -> SendState {
self.staging.publish(&self.tx_coalesced, update)
}
pub(crate) fn flush(&mut self) -> SendState {
self.staging.flush(&self.tx_coalesced)
}
pub(crate) fn has_pending(&self) -> bool {
self.staging.has_pending()
}
pub(crate) fn epoch(&mut self) {
self.staging.clear();
}
pub(crate) fn is_closed(&self) -> bool {
self.tx_control.is_closed() || self.tx_coalesced.is_closed()
}
}
#[cfg(test)]
impl MarketUpdateSink {
pub(crate) fn staged_len(&self) -> usize {
self.staging.slots.len()
}
}
#[derive(Debug, Default)]
struct StagedInstrument {
quote: Option<QuoteUpdate>,
greeks: Option<GreeksRow>,
depth: Option<DepthLadder>,
}
impl StagedInstrument {
fn has_any(&self) -> bool {
self.quote.is_some() || self.greeks.is_some() || self.depth.is_some()
}
fn flush_into(&mut self, tx: &mpsc::Sender<MarketUpdate>) -> FlushStep {
if self.quote.is_some() {
match reserve_send(tx, &mut self.quote, MarketUpdate::Quote) {
FlushStep::Drained => {}
blocked => return blocked,
}
}
if self.greeks.is_some() {
match reserve_send(tx, &mut self.greeks, MarketUpdate::Greeks) {
FlushStep::Drained => {}
blocked => return blocked,
}
}
if self.depth.is_some() {
match reserve_send(tx, &mut self.depth, MarketUpdate::Depth) {
FlushStep::Drained => {}
blocked => return blocked,
}
}
FlushStep::Drained
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum FlushStep {
Drained,
Full,
Closed,
}
fn reserve_send<T>(
tx: &mpsc::Sender<MarketUpdate>,
slot: &mut Option<T>,
wrap: fn(T) -> MarketUpdate,
) -> FlushStep {
match tx.try_reserve() {
Ok(permit) => {
if let Some(value) = slot.take() {
permit.send(wrap(value));
}
FlushStep::Drained
}
Err(TrySendError::Full(())) => FlushStep::Full,
Err(TrySendError::Closed(())) => FlushStep::Closed,
}
}
#[derive(Debug, Default)]
struct ProducerStaging {
slots: HashMap<InstrumentKey, StagedInstrument>,
}
impl ProducerStaging {
fn new() -> Self {
Self {
slots: HashMap::new(),
}
}
fn publish(&mut self, tx: &mpsc::Sender<MarketUpdate>, update: MarketUpdate) -> SendState {
if self.flush(tx) == SendState::Closed {
return SendState::Closed;
}
match tx.try_send(update) {
Ok(()) => SendState::Open,
Err(TrySendError::Full(update)) => {
self.stage(update);
SendState::Open
}
Err(TrySendError::Closed(_)) => SendState::Closed,
}
}
fn has_pending(&self) -> bool {
self.slots.values().any(StagedInstrument::has_any)
}
fn flush(&mut self, tx: &mpsc::Sender<MarketUpdate>) -> SendState {
let mut closed = false;
let mut full = false;
self.slots.retain(|_key, slot| {
if !closed && !full {
match slot.flush_into(tx) {
FlushStep::Drained => {}
FlushStep::Full => full = true,
FlushStep::Closed => closed = true,
}
}
slot.has_any()
});
if closed {
SendState::Closed
} else {
SendState::Open
}
}
fn clear(&mut self) {
self.slots.clear();
}
fn stage(&mut self, update: MarketUpdate) {
match update {
MarketUpdate::Quote(quote) => {
if let Some(slot) = self.slot_mut("e.instrument.key) {
slot.quote = Some(quote);
}
}
MarketUpdate::Greeks(greeks) => {
if let Some(slot) = self.slot_mut(&greeks.instrument.key) {
slot.greeks = Some(greeks);
}
}
MarketUpdate::Depth(depth) => {
if let Some(slot) = self.slot_mut(&depth.instrument.key) {
slot.depth = Some(depth);
}
}
MarketUpdate::Chain(_) | MarketUpdate::Health(_, _) => {}
}
}
fn slot_mut(&mut self, key: &InstrumentKey) -> Option<&mut StagedInstrument> {
if !self.slots.contains_key(key) {
let _ = self.slots.insert(key.clone(), StagedInstrument::default());
}
self.slots.get_mut(key)
}
}
#[cfg(test)]
mod tests {
use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{Context, Poll, Waker};
use optionstratlib::OptionStyle;
use optionstratlib::chains::chain::OptionChain;
use optionstratlib::prelude::Positive;
use super::*;
use crate::chain::{
AliasCatalog, ContractSpecFingerprint, ExerciseStyle, ExpirySource, InstrumentKey,
SettlementStyle,
};
#[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 pos(value: f64) -> Positive {
match Positive::new(value) {
Ok(p) => p,
Err(e) => panic!("invalid test positive `{value}`: {e}"),
}
}
fn sample_key() -> InstrumentKey {
InstrumentKey {
underlying: "BTC".to_owned(),
expiration_utc: utc(1_700_000_000),
strike: pos(60_000.0),
style: OptionStyle::Call,
}
}
fn sample_instrument(provider: &str, native: &str, stream: Option<&str>) -> Instrument {
Instrument {
key: sample_key(),
provider: pid(provider),
native_symbol: native.to_owned(),
stream_symbol: stream.map(str::to_owned),
spec: ContractSpecFingerprint {
contract_multiplier: 1,
settlement: SettlementStyle::Cash,
exercise: ExerciseStyle::European,
quote_currency: "USD".to_owned(),
venue_product_code: "BTC".to_owned(),
},
}
}
fn block_on<F: Future>(future: F) -> F::Output {
let mut future = std::pin::pin!(future);
let waker = Waker::noop();
let mut cx = Context::from_waker(waker);
match future.as_mut().poll(&mut cx) {
Poll::Ready(output) => output,
Poll::Pending => panic!("test future parked; port futures must resolve on first poll"),
}
}
fn assert_send_sync<T: Send + Sync>() {}
struct FakeProvider {
id: ProviderId,
capabilities: ProviderCapabilities,
chain: ChainFetch,
underlyings: Vec<UnderlyingRef>,
}
#[async_trait]
impl Provider for FakeProvider {
fn id(&self) -> ProviderId {
self.id.clone()
}
fn capabilities(&self) -> ProviderCapabilities {
self.capabilities
}
async fn discover(&self) -> Result<Vec<UnderlyingRef>, ProviderError> {
Ok(self.underlyings.clone())
}
async fn fetch_chain(
&self,
_underlying: &str,
_expiration: &ExpirationDate,
) -> Result<ChainFetch, ProviderError> {
Ok(self.chain.clone())
}
async fn subscribe(
&self,
_req: SubscriptionRequest,
_sink: MarketUpdateSink,
) -> Result<SubscriptionHandle, ProviderError> {
Err(ProviderError::Unsupported("streaming"))
}
}
fn fake_provider() -> FakeProvider {
let mut aliases = AliasCatalog::new();
aliases.insert(sample_instrument(
"fake",
"BTC-27JUN25-60000-C",
Some("dxfeed-sym"),
));
FakeProvider {
id: pid("fake"),
capabilities: ProviderCapabilities::builder()
.chain(ChainCapability::Assemble)
.greeks(GreeksCapability::Provided)
.chain_poll(ChainPollCapability::Poll {
interval_hint_secs: 2,
})
.build(),
chain: ChainFetch::new(
OptionChain::new("BTC", pos(60_000.0), "2025-06-27".to_owned(), None, None),
ExpirySource::new("BTC", utc(1_700_000_000), pid("fake")),
aliases,
),
underlyings: vec![UnderlyingRef::new("BTC")],
}
}
#[test]
fn test_provider_capabilities_builder_sets_every_field() {
let caps = ProviderCapabilities::builder()
.chain(ChainCapability::Assemble)
.depth(true)
.greeks(GreeksCapability::ComputedLocally)
.option_stream(OptionStreamCapability::ChainQuotes { verified: false })
.underlying_stream(true)
.chain_poll(ChainPollCapability::Poll {
interval_hint_secs: 5,
})
.trades_tape(true)
.auth(AuthKind::UserPass)
.build();
assert_eq!(caps.chain, ChainCapability::Assemble);
assert!(caps.depth);
assert_eq!(caps.greeks, GreeksCapability::ComputedLocally);
assert_eq!(
caps.option_stream,
OptionStreamCapability::ChainQuotes { verified: false }
);
assert!(caps.underlying_stream);
assert_eq!(
caps.chain_poll,
ChainPollCapability::Poll {
interval_hint_secs: 5
}
);
assert!(caps.trades_tape);
assert_eq!(caps.auth, AuthKind::UserPass);
}
#[test]
fn test_provider_capabilities_builder_defaults_are_least_capable() {
let caps = ProviderCapabilities::builder().build();
assert_eq!(caps.chain, ChainCapability::None);
assert!(!caps.depth);
assert_eq!(caps.greeks, GreeksCapability::None);
assert_eq!(caps.option_stream, OptionStreamCapability::None);
assert!(!caps.underlying_stream);
assert_eq!(caps.chain_poll, ChainPollCapability::None);
assert!(!caps.trades_tape);
assert_eq!(caps.auth, AuthKind::None);
}
#[test]
fn test_capability_enum_defaults_are_none() {
assert_eq!(ChainCapability::default(), ChainCapability::None);
assert_eq!(GreeksCapability::default(), GreeksCapability::None);
assert_eq!(
OptionStreamCapability::default(),
OptionStreamCapability::None
);
assert_eq!(ChainPollCapability::default(), ChainPollCapability::None);
assert_eq!(AuthKind::default(), AuthKind::None);
}
#[test]
fn test_alias_catalog_round_trips_native_and_stream_symbols() {
let mut catalog = AliasCatalog::new();
catalog.insert(sample_instrument(
"deribit",
"BTC-27JUN25-60000-C",
Some("dxfeed-sym"),
));
assert_eq!(
catalog.resolve_symbol("BTC-27JUN25-60000-C"),
Some(&sample_key())
);
assert_eq!(catalog.resolve_symbol("dxfeed-sym"), Some(&sample_key()));
match catalog.instrument(&sample_key(), &pid("deribit")) {
Some(found) => assert_eq!(found.native_symbol, "BTC-27JUN25-60000-C"),
None => panic!("expected the deribit alias for the leg"),
}
}
#[test]
fn test_chain_fetch_carries_alias_catalog_forward_unchanged() {
let provider = fake_provider();
let fetch = block_on(provider.fetch_chain("BTC", &ExpirationDate::Days(pos(30.0))));
match fetch {
Ok(chain_fetch) => {
assert_eq!(
chain_fetch.aliases.resolve_symbol("dxfeed-sym"),
Some(&sample_key())
);
assert!(
chain_fetch
.aliases
.instrument(&sample_key(), &pid("fake"))
.is_some()
);
assert_eq!(chain_fetch.expiry_source.underlying, "BTC");
}
Err(e) => panic!("fetch_chain should succeed for the fake provider, got: {e}"),
}
}
#[test]
fn test_provider_trait_object_is_send_sync() {
assert_send_sync::<Box<dyn Provider>>();
}
#[test]
fn test_fake_provider_is_object_safe_and_reports_capabilities() {
let provider: Box<dyn Provider> = Box::new(fake_provider());
assert_eq!(provider.id().as_str(), "fake");
assert_eq!(provider.capabilities().chain, ChainCapability::Assemble);
}
#[test]
fn test_fake_provider_discover_lists_underlyings() {
let provider = fake_provider();
match block_on(provider.discover()) {
Ok(underlyings) => match underlyings.first() {
Some(first) => assert_eq!(first.underlying, "BTC"),
None => panic!("expected at least one underlying"),
},
Err(e) => panic!("discover should succeed for the fake provider, got: {e}"),
}
}
#[test]
fn test_fake_provider_subscribe_is_unsupported_for_poll_only_shape() {
let provider = fake_provider();
let (tx, _rx) = mpsc::channel::<MarketUpdate>(1);
let sink = MarketUpdateSink::new(tx.clone(), tx);
let request = SubscriptionRequest::new(
"BTC",
utc(1_700_000_000),
Vec::new(),
CancellationToken::new(),
);
match block_on(provider.subscribe(request, sink)) {
Err(ProviderError::Unsupported(what)) => assert_eq!(what, "streaming"),
other => panic!("expected Unsupported(\"streaming\"), got {other:?}"),
}
}
#[test]
fn test_subscription_handle_drop_cancels_subscription() {
let cancelled = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&cancelled);
let handle = SubscriptionHandle::new(move || flag.store(true, Ordering::SeqCst));
assert!(!cancelled.load(Ordering::SeqCst));
drop(handle);
assert!(cancelled.load(Ordering::SeqCst));
}
#[test]
fn test_subscription_handle_abort_cancels_once() {
let count = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&count);
let handle = SubscriptionHandle::new(move || flag.store(true, Ordering::SeqCst));
handle.abort();
assert!(count.load(Ordering::SeqCst));
}
#[test]
fn test_subscription_request_new_sets_fields() {
let cancel = CancellationToken::new();
let request = SubscriptionRequest::new(
"BTC",
utc(1_700_000_000),
vec![sample_instrument("deribit", "native", None)],
cancel.clone(),
);
assert_eq!(request.underlying, "BTC");
assert_eq!(request.expiration_utc, utc(1_700_000_000));
assert_eq!(request.instruments.len(), 1);
assert!(!request.cancel.is_cancelled());
cancel.cancel();
assert!(request.cancel.is_cancelled());
}
#[test]
fn test_underlying_ref_new_has_no_expirations() {
let underlying = UnderlyingRef::new("SPY");
assert_eq!(underlying.underlying, "SPY");
assert!(underlying.expirations.is_empty());
}
#[test]
fn test_underlying_ref_with_expirations_carries_them() {
let underlying =
UnderlyingRef::with_expirations("BTC", vec![ExpirationDate::Days(pos(7.0))]);
assert_eq!(underlying.expirations.len(), 1);
}
}