pub mod data;
mod interval;
pub mod source;
mod store;
pub use data::{Data, DataBasics, DataQuery, kline, word_trend};
pub use interval::Interval;
pub use source::Source;
use tokio::sync::Mutex;
use std::{
collections::HashMap,
fmt,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
};
use surrealdb::{Surreal, engine::local::Db};
use crate::{
core::{CoreError, CoreResult},
provider::store::Store,
};
type AtomicMap<S> = HashMap<S, HashMap<Interval, Arc<AtomicBool>>>;
pub trait ProviderTrait: Send + Sync + fmt::Debug + 'static {
type _Source: Source;
fn provide(
&self,
query: &DataQuery<<Self::_Source as Source>::_Data>,
) -> impl Future<Output = CoreResult<Vec<<Self::_Source as Source>::_Data>>> + Send;
fn provide_or_empty(
&self,
query: &DataQuery<<Self::_Source as Source>::_Data>,
) -> impl Future<Output = CoreResult<Vec<<Self::_Source as Source>::_Data>>> + Send {
async move {
match self.provide(query).await {
Ok(data) => Ok(data),
Err(CoreError::UnavailableData) => Ok(vec![]),
Err(e) => Err(e),
}
}
}
fn name(&self) -> String;
}
#[derive(Debug)]
pub struct Provider<S: Source> {
source: S,
shared_atomics: Arc<Mutex<AtomicMap<<S::_Data as Data>::Specifier>>>,
local_atomics: Mutex<AtomicMap<<S::_Data as Data>::Specifier>>,
store: Store<S::_Data>,
}
impl<S: Source> Provider<S> {
pub fn new(source: S, db: Surreal<Db>) -> Self {
let store = Store::new(db);
let shared_atomics = Arc::new(Mutex::new(AtomicMap::new()));
let local_atomics = Mutex::new(AtomicMap::new());
Self {
source,
shared_atomics,
local_atomics,
store,
}
}
async fn get_flag(
&self,
specifier: &<S::_Data as Data>::Specifier,
interval: Interval,
) -> Arc<AtomicBool> {
let mut local_atomics_guard = self.local_atomics.lock().await;
let interval_map = local_atomics_guard.entry(specifier.clone()).or_default();
if let Some(atomic) = interval_map.get(&interval) {
atomic.clone()
} else {
let mut shared_atomics_guard = self.shared_atomics.lock().await;
let new_atomic = shared_atomics_guard
.entry(specifier.clone())
.or_default()
.entry(interval)
.or_insert_with(|| Arc::new(AtomicBool::new(false)));
interval_map.insert(interval, new_atomic.clone());
new_atomic.clone()
}
}
async fn fetch_data(
&self,
local_data: Vec<S::_Data>,
query: &DataQuery<S::_Data>,
) -> CoreResult<()> {
let begin = if let Some(last_data) = local_data.last() {
last_data.basics().end
} else {
*query.begin()
};
tracing::trace!("Source query begin date: {begin:#?}");
let fresh_data = self
.source
.fetch(*query.interval(), query.specifier().clone(), begin)
.await?;
if fresh_data.is_empty() {
return Err(CoreError::UnavailableData);
}
self.store
.inject(query.specifier(), *query.interval(), fresh_data)
.await?;
Ok(())
}
}
impl<S: Source> Clone for Provider<S> {
fn clone(&self) -> Self {
Self {
source: self.source.clone(),
shared_atomics: self.shared_atomics.clone(),
local_atomics: Mutex::new(AtomicMap::new()),
store: self.store.clone(),
}
}
}
impl<S: Source> ProviderTrait for Provider<S> {
type _Source = S;
async fn provide(
&self,
query: &DataQuery<<Self::_Source as Source>::_Data>,
) -> CoreResult<Vec<<Self::_Source as Source>::_Data>> {
let mut source_calls_safety = 0;
tracing::debug!("Provider : Getting data");
tracing::trace!("Query: {query:#?}");
loop {
tracing::trace!("Provider data loop round {}", source_calls_safety);
let (local_data, complete) = self.store.try_get_data(query).await?;
if complete {
tracing::debug!("Data acquired");
tracing::trace!("data size: {}", local_data.len());
return Ok(local_data);
}
if source_calls_safety > self.source.max_requests() {
return Err(CoreError::param_error("Max data requests reached"));
}
tracing::debug!("Find the flag");
let flag = self.get_flag(query.specifier(), *query.interval()).await;
tracing::trace!("Try to take the flag");
match flag.compare_exchange_weak(false, true, Ordering::Acquire, Ordering::Relaxed) {
Ok(_) => {
tracing::debug!("Getting data from source");
let res = self.fetch_data(local_data, query).await;
flag.store(false, Ordering::Release);
res?;
source_calls_safety += 1;
}
Err(_) => {
tracing::trace!("Data already updating");
}
}
}
}
fn name(&self) -> String {
self.source.name()
}
}
#[cfg(test)]
pub mod test {
use crate::provider::source::test::EmptySource;
use super::*;
#[derive(Debug)]
pub struct MockProvider<D: Data> {
providing: fn(&DataQuery<D>) -> CoreResult<Vec<D>>,
}
impl<D: Data> MockProvider<D> {
pub fn new(f: fn(&DataQuery<D>) -> CoreResult<Vec<D>>) -> Self {
Self { providing: f }
}
}
impl<D: Data> ProviderTrait for MockProvider<D> {
type _Source = EmptySource<D>;
async fn provide(
&self,
query: &DataQuery<<Self::_Source as Source>::_Data>,
) -> CoreResult<Vec<<Self::_Source as Source>::_Data>> {
(self.providing)(query)
}
fn name(&self) -> String {
"MockProvider".to_string()
}
}
}
#[cfg(test)]
mod tests {
use core::panic;
use std::{
sync::{
Arc,
atomic::{AtomicBool, AtomicU32, Ordering},
},
time::Duration,
};
use chrono::{TimeZone as _, Utc};
use surrealdb::{Surreal, engine::local::Mem};
use tokio::time::sleep;
use crate::{
core::CoreResult,
provider::{
Data, DataBasics, DataQuery, Interval, Provider, ProviderTrait as _, Source,
data::test::{TestData, TestDataSpecifier, data_vec_over},
source::test::MockSource,
},
};
#[tokio::test]
async fn get_data() {
let db = Surreal::new::<Mem>(()).await.unwrap();
db.use_ns("test").use_db("test").await.unwrap();
let mocked_source = MockSource::new(
|interval, specific, begin| {
data_vec_over(
interval,
begin,
begin + interval.time_delta().checked_mul(10).unwrap(),
TestData {
content: format!("{specific:?}"),
basics: DataBasics::default(),
},
MockSource::<TestData>::static_name(),
)
},
1,
);
let unusable_source = MockSource::new(
|_, _, _| {
panic!("Unusable source");
},
0,
);
let provider_1 = Provider::new(mocked_source.clone(), db.clone());
let provider_2 = Provider::new(unusable_source, db);
let data_1 = provider_1
.provide(&DataQuery::new(
TestDataSpecifier("test".to_string()),
Interval::OneMinute,
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 1, 1, 0, 3, 0).unwrap(),
))
.await
.unwrap();
let data_2 = provider_2
.provide(&DataQuery::new(
TestDataSpecifier("test".to_string()),
Interval::OneMinute,
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 1, 1, 0, 3, 0).unwrap(),
))
.await
.unwrap();
assert_eq!(data_1.len(), 3);
assert_eq!(
data_1[1],
TestData {
content: format!("{:?}", TestDataSpecifier("test".to_string())),
basics: DataBasics {
begin: Utc.with_ymd_and_hms(2024, 1, 1, 0, 1, 0).unwrap(),
end: Utc.with_ymd_and_hms(2024, 1, 1, 0, 2, 0).unwrap(),
update_date: data_1[1].basics().update_date,
source: mocked_source.name(),
}
}
);
assert_eq!(data_1, data_2);
}
#[derive(Debug, Clone)]
struct MonoCallTestSource {
flag: Arc<AtomicBool>,
counter: Arc<AtomicU32>,
}
impl MonoCallTestSource {
pub fn new() -> Self {
Self {
flag: Arc::new(AtomicBool::new(false)),
counter: Arc::new(AtomicU32::new(0)),
}
}
}
impl Source for MonoCallTestSource {
type _Data = TestData;
fn name(&self) -> String {
"MonoCallTestSource".to_string()
}
fn max_requests(&self) -> u32 {
3
}
async fn fetch(
&self,
interval: Interval,
specific: <Self::_Data as super::Data>::Specifier,
begin: chrono::DateTime<Utc>,
) -> CoreResult<Vec<Self::_Data>> {
match self
.flag
.compare_exchange(false, true, Ordering::Acquire, Ordering::Relaxed)
{
Ok(_) => {
sleep(Duration::from_millis(50)).await;
self.counter.fetch_add(1, Ordering::SeqCst);
self.flag.store(false, Ordering::Release);
if self.counter.to_owned().load(Ordering::SeqCst) > self.max_requests() {
panic!("Too many requests");
}
data_vec_over(
interval,
begin,
begin + interval.time_delta().checked_mul(2).unwrap(),
TestData::new(format!("{specific:?}"), DataBasics::default()),
self.name(),
)
}
Err(_) => panic!("Multiple requests at the same time"),
}
}
}
#[tokio::test]
async fn parallel_requests() {
let db = Surreal::new::<Mem>(()).await.unwrap();
db.use_ns("test").use_db("test").await.unwrap();
let mocked_source = MonoCallTestSource::new();
let provider = Provider::new(mocked_source.clone(), db.clone());
let query = DataQuery::<TestData>::new(
TestDataSpecifier("test".to_string()),
Interval::OneMinute,
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 1, 1, 0, 6, 0).unwrap(),
);
let expected_data = data_vec_over(
Interval::OneMinute,
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 1, 1, 0, 6, 0).unwrap(),
TestData::new("test".to_string(), DataBasics::default()),
provider.name(),
)
.unwrap();
for _ in (0..10).enumerate() {
tokio::spawn(test_provider(
provider.clone(),
query.clone(),
expected_data.clone(),
6,
));
}
}
async fn test_provider(
provider: Provider<MonoCallTestSource>,
query: DataQuery<TestData>,
expected_data: Vec<TestData>,
data_len: usize,
) {
let data = provider.provide(&query).await.unwrap();
assert_eq!(data.len(), data_len);
assert_eq!(data, expected_data);
}
}