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}