Skip to main content

alpaca_data/cache/
client.rs

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