Skip to main content

alpaca_data/cache/
client.rs

1use std::collections::{HashMap, HashSet};
2use std::sync::Arc;
3use std::time::SystemTime;
4
5use chrono::{DateTime, Utc};
6use tokio::sync::RwLock;
7
8use crate::cache::state::{
9    BarsMap, CacheState, StockBarsRequest, collect_cached_hits, normalize_option_symbols,
10    normalize_stock_symbols,
11};
12use crate::cache::stats::CacheStats;
13use crate::options::{self, OptionsFeed, SnapshotsRequest as OptionSnapshotsRequest};
14use crate::stocks::{self, DataFeed, SnapshotsRequest as StockSnapshotsRequest};
15use crate::{Client, Error};
16
17#[derive(Clone)]
18pub struct CachedClientConfig {
19    pub stocks_feed: Arc<dyn Fn() -> DataFeed + Send + Sync>,
20    pub options_feed: OptionsFeed,
21}
22
23impl Default for CachedClientConfig {
24    fn default() -> Self {
25        Self {
26            stocks_feed: Arc::new(|| stocks::preferred_feed(false)),
27            options_feed: options::preferred_feed(),
28        }
29    }
30}
31
32pub struct CachedClient {
33    raw: Client,
34    config: CachedClientConfig,
35    state: RwLock<CacheState>,
36}
37
38impl CachedClient {
39    #[must_use]
40    pub fn new(raw: Client) -> Self {
41        Self::with_config(raw, CachedClientConfig::default())
42    }
43
44    #[must_use]
45    pub fn with_config(raw: Client, config: CachedClientConfig) -> Self {
46        Self {
47            raw,
48            config,
49            state: RwLock::new(CacheState::default()),
50        }
51    }
52
53    #[must_use]
54    pub fn raw(&self) -> &Client {
55        &self.raw
56    }
57
58    pub async fn stocks<S: AsRef<str>>(
59        &self,
60        symbols: &[S],
61    ) -> Result<HashMap<String, stocks::Snapshot>, Error> {
62        let requested = normalize_stock_symbols(symbols);
63        if requested.is_empty() {
64            return Ok(HashMap::new());
65        }
66
67        let resolved = unique_resolved_symbols(&requested);
68        let (mut hits, missing) = {
69            let state = self.state.read().await;
70            collect_cached_hits(&resolved, &state.stocks.values, &state.stocks.empty)
71        };
72
73        if !missing.is_empty() {
74            let fetched = self.fetch_stocks(&missing).await?;
75            let mut state = self.state.write().await;
76            for symbol in &missing {
77                if let Some(snapshot) = fetched.get(symbol) {
78                    hits.insert(symbol.clone(), snapshot.clone());
79                }
80            }
81            state
82                .stocks
83                .reconcile(&missing, &fetched, SystemTime::now());
84        }
85
86        Ok(requested
87            .into_iter()
88            .filter_map(|(original, resolved)| {
89                hits.get(&resolved)
90                    .cloned()
91                    .map(|snapshot| (original, snapshot))
92            })
93            .collect())
94    }
95
96    pub async fn stock(&self, symbol: &str) -> Option<stocks::Snapshot> {
97        self.stocks(&[symbol])
98            .await
99            .ok()?
100            .into_iter()
101            .next()
102            .map(|(_, snapshot)| snapshot)
103    }
104
105    pub async fn options<S: AsRef<str>>(
106        &self,
107        contracts: &[S],
108    ) -> Result<HashMap<String, options::Snapshot>, Error> {
109        let requested = normalize_option_symbols(contracts);
110        if requested.is_empty() {
111            return Ok(HashMap::new());
112        }
113
114        let (mut hits, missing) = {
115            let state = self.state.read().await;
116            collect_cached_hits(&requested, &state.options.values, &state.options.empty)
117        };
118
119        if !missing.is_empty() {
120            let fetched = self.fetch_options(&missing).await?;
121            let mut state = self.state.write().await;
122            for contract in &missing {
123                if let Some(snapshot) = fetched.get(contract) {
124                    hits.insert(contract.clone(), snapshot.clone());
125                }
126            }
127            state
128                .options
129                .reconcile(&missing, &fetched, SystemTime::now());
130        }
131
132        Ok(requested
133            .into_iter()
134            .filter_map(|contract| hits.remove_entry(&contract))
135            .collect())
136    }
137
138    pub async fn option(&self, contract: &str) -> Option<options::Snapshot> {
139        self.options(&[contract]).await.ok()?.remove(contract)
140    }
141
142    pub async fn watch_stocks(&self, symbols: &[String]) {
143        let normalized = normalize_stock_symbols(symbols);
144        let mut state = self.state.write().await;
145        for (_, symbol) in normalized {
146            state.stocks.subscribed.insert(symbol);
147        }
148    }
149
150    pub async fn watch_options(&self, contracts: &[String]) {
151        let normalized = normalize_option_symbols(contracts);
152        let mut state = self.state.write().await;
153        for contract in normalized {
154            state.options.subscribed.insert(contract);
155        }
156    }
157
158    pub async fn refresh_stocks(&self) -> Result<usize, Error> {
159        let symbols = {
160            let state = self.state.read().await;
161            state.stocks.subscribed.iter().cloned().collect::<Vec<_>>()
162        };
163        if symbols.is_empty() {
164            return Ok(0);
165        }
166
167        let fetched = self.fetch_stocks(&symbols).await?;
168        let mut state = self.state.write().await;
169        let count = state
170            .stocks
171            .reconcile(&symbols, &fetched, SystemTime::now());
172        Ok(count)
173    }
174
175    pub async fn refresh_options(&self) -> Result<usize, Error> {
176        let contracts = {
177            let state = self.state.read().await;
178            state.options.subscribed.iter().cloned().collect::<Vec<_>>()
179        };
180        if contracts.is_empty() {
181            return Ok(0);
182        }
183
184        let fetched = self.fetch_options(&contracts).await?;
185        let mut state = self.state.write().await;
186        let count = state
187            .options
188            .reconcile(&contracts, &fetched, SystemTime::now());
189        Ok(count)
190    }
191
192    pub async fn watch_bars(&self, request: StockBarsRequest) {
193        let request = request.normalized();
194        let mut state = self.state.write().await;
195        state
196            .bars
197            .requests
198            .entry(request.key.clone())
199            .and_modify(|current| current.merge_from(&request))
200            .or_insert(request);
201    }
202
203    pub async fn bars(&self, key: &str) -> Result<HashMap<String, Vec<stocks::BarPoint>>, Error> {
204        let request = self.bars_request(key).await?;
205        let missing = {
206            let state = self.state.read().await;
207            let cached = state.bars.values.get(key);
208            let empty = state.bars.empty.get(key);
209
210            request
211                .symbols
212                .iter()
213                .filter(|symbol| {
214                    !cached.is_some_and(|bars| bars.contains_key(*symbol))
215                        && !empty.is_some_and(|values| values.contains(*symbol))
216                })
217                .cloned()
218                .collect::<Vec<_>>()
219        };
220
221        if missing.is_empty() {
222            let state = self.state.read().await;
223            return Ok(state.bars.values.get(key).cloned().unwrap_or_default());
224        }
225
226        self.fetch_missing_bars(key, &request, &missing).await
227    }
228
229    pub async fn bar(&self, key: &str, symbol: &str) -> Option<Vec<stocks::BarPoint>> {
230        let resolved = stocks::display_stock_symbol(symbol);
231        {
232            let state = self.state.read().await;
233            if let Some(values) = state.bars.values.get(key)
234                && let Some(bars) = values.get(&resolved)
235            {
236                return Some(bars.clone());
237            }
238            if state
239                .bars
240                .empty
241                .get(key)
242                .is_some_and(|symbols| symbols.contains(&resolved))
243            {
244                return None;
245            }
246        }
247
248        self.bars(key).await.ok()?.get(&resolved).cloned()
249    }
250
251    pub async fn refresh_bars(&self, key: &str) -> Result<usize, Error> {
252        let request = self.bars_request(key).await?;
253        let fetched = self.fetch_bars_request(&request, &request.symbols).await?;
254        let count = fetched.len();
255
256        let missing: HashSet<String> = request
257            .symbols
258            .iter()
259            .filter(|symbol| !fetched.contains_key(*symbol))
260            .cloned()
261            .collect();
262
263        let mut state = self.state.write().await;
264        state.bars.values.insert(key.to_string(), fetched);
265        state.bars.empty.insert(key.to_string(), missing);
266        state
267            .bars
268            .updated_at
269            .insert(key.to_string(), SystemTime::now());
270        Ok(count)
271    }
272
273    pub async fn clear_options(&self) {
274        let mut state = self.state.write().await;
275        state.options.subscribed.clear();
276        state.options.values.clear();
277        state.options.empty.clear();
278        state.options.updated_at = None;
279    }
280
281    pub async fn stats(&self) -> CacheStats {
282        let state = self.state.read().await;
283        CacheStats {
284            subscribed_symbols: state.stocks.subscribed.len(),
285            subscribed_contracts: state.options.subscribed.len(),
286            subscribed_bar_requests: state.bars.requests.len(),
287            cached_stocks: state.stocks.values.len(),
288            cached_options: state.options.values.len(),
289            unavailable_stocks: state.stocks.empty.len(),
290            unavailable_options: state.options.empty.len(),
291            cached_bar_symbols: state.bars.values.values().map(HashMap::len).sum(),
292            stocks_updated_at: format_timestamp(state.stocks.updated_at),
293            options_updated_at: format_timestamp(state.options.updated_at),
294            bars_updated_at: state
295                .bars
296                .updated_at
297                .iter()
298                .map(|(key, value)| {
299                    (
300                        key.clone(),
301                        format_timestamp(Some(*value)).unwrap_or_default(),
302                    )
303                })
304                .collect(),
305        }
306    }
307
308    async fn fetch_stocks(
309        &self,
310        symbols: &[String],
311    ) -> Result<HashMap<String, stocks::Snapshot>, Error> {
312        self.raw
313            .stocks()
314            .snapshots(StockSnapshotsRequest {
315                symbols: symbols.to_vec(),
316                feed: Some((self.config.stocks_feed)()),
317                currency: None,
318            })
319            .await
320    }
321
322    async fn fetch_options(
323        &self,
324        contracts: &[String],
325    ) -> Result<HashMap<String, options::Snapshot>, Error> {
326        self.raw
327            .options()
328            .snapshots_all(OptionSnapshotsRequest {
329                symbols: contracts.to_vec(),
330                feed: Some(self.config.options_feed),
331                limit: Some(1000),
332                page_token: None,
333            })
334            .await
335            .map(|response| response.snapshots)
336    }
337
338    async fn bars_request(&self, key: &str) -> Result<StockBarsRequest, Error> {
339        let key = key.trim();
340        if key.is_empty() {
341            return Err(Error::InvalidRequest(
342                "bars key is invalid: must not be empty".to_owned(),
343            ));
344        }
345
346        let state = self.state.read().await;
347        state
348            .bars
349            .requests
350            .get(key)
351            .cloned()
352            .ok_or_else(|| Error::InvalidRequest(format!("bars key is unknown: {key}")))
353    }
354
355    async fn fetch_missing_bars(
356        &self,
357        key: &str,
358        request: &StockBarsRequest,
359        missing: &[String],
360    ) -> Result<HashMap<String, Vec<stocks::BarPoint>>, Error> {
361        let fetched = self.fetch_bars_request(request, missing).await?;
362        let missing_empty: HashSet<String> = missing
363            .iter()
364            .filter(|symbol| !fetched.contains_key(*symbol))
365            .cloned()
366            .collect();
367
368        let mut state = self.state.write().await;
369        let key = key.to_string();
370        {
371            let cached = state.bars.values.entry(key.clone()).or_default();
372            for (symbol, bars) in &fetched {
373                cached.insert(symbol.clone(), bars.clone());
374            }
375        }
376        {
377            let empty = state.bars.empty.entry(key.clone()).or_default();
378            for symbol in missing {
379                if missing_empty.contains(symbol) {
380                    empty.insert(symbol.clone());
381                } else {
382                    empty.remove(symbol);
383                }
384            }
385        }
386        state.bars.updated_at.insert(key.clone(), SystemTime::now());
387
388        Ok(state.bars.values.get(&key).cloned().unwrap_or_default())
389    }
390
391    async fn fetch_bars_request(
392        &self,
393        request: &StockBarsRequest,
394        symbols: &[String],
395    ) -> Result<BarsMap, Error> {
396        if symbols.is_empty() {
397            return Ok(HashMap::new());
398        }
399
400        let mut merged = HashMap::new();
401        let chunk_size = request.chunk_size.max(1);
402        let daily = request.timeframe == stocks::TimeFrame::day_1();
403
404        for chunk in symbols.chunks(chunk_size) {
405            let response = self
406                .raw
407                .stocks()
408                .bars_all(stocks::BarsRequest {
409                    symbols: chunk.to_vec(),
410                    timeframe: request.timeframe.clone(),
411                    start: request.start.clone(),
412                    end: request.end.clone(),
413                    limit: Some(request.limit),
414                    adjustment: request.adjustment.clone(),
415                    feed: request.feed,
416                    sort: None,
417                    asof: None,
418                    currency: request.currency.clone(),
419                    page_token: None,
420                })
421                .await?;
422
423            for (symbol, bars) in response.bars {
424                merged.insert(
425                    symbol,
426                    bars.into_iter().map(|bar| bar.point(daily)).collect(),
427                );
428            }
429        }
430
431        Ok(merged)
432    }
433}
434
435fn unique_resolved_symbols(requested: &[(String, String)]) -> Vec<String> {
436    let mut resolved = Vec::new();
437    let mut seen = HashSet::new();
438    for (_, symbol) in requested {
439        if seen.insert(symbol.clone()) {
440            resolved.push(symbol.clone());
441        }
442    }
443    resolved
444}
445
446fn format_timestamp(value: Option<SystemTime>) -> Option<String> {
447    value.map(|value| {
448        DateTime::<Utc>::from(value)
449            .format("%Y-%m-%d %H:%M:%S")
450            .to_string()
451    })
452}