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}