use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::marker::PhantomData;
use surrealdb::Surreal;
use surrealdb::engine::local::Db;
use crate::core::{CoreError, CoreResult};
use crate::provider::data::{Data, DataBasics, DataQuery, DataSpecifier};
use crate::provider::interval::Interval;
#[derive(Debug, Serialize, Deserialize)]
struct DataSection {
pub begin: DateTime<Utc>,
pub end: DateTime<Utc>,
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(bound = "D: Data")]
struct DataWrapper<D: Data> {
begin: DateTime<Utc>,
end: DateTime<Utc>,
data: D,
}
impl<D: Data> From<D> for DataWrapper<D> {
fn from(data: D) -> Self {
Self {
begin: data.basics().begin,
end: data.basics().end,
data,
}
}
}
impl<D: Data> DataWrapper<D> {
fn into_data(self) -> D {
self.data
}
}
#[derive(Debug)]
pub struct Store<D: Data> {
db: Surreal<Db>,
_marker: PhantomData<D>,
}
impl<D: Data> Clone for Store<D> {
fn clone(&self) -> Self {
Self {
db: self.db.clone(),
_marker: PhantomData,
}
}
}
impl<D: Data> Store<D> {
pub fn new(db: Surreal<Db>) -> Self {
Self {
db,
_marker: PhantomData,
}
}
#[allow(dead_code)]
fn check_basics(&self, interval: Interval, basics: &DataBasics) -> CoreResult<()> {
let (begin, end) = interval.get_time_bounds(basics.begin)?;
if begin != basics.begin || end != basics.end {
Err(CoreError::ParamError(format!(
"Wrong time bounds in {basics:?} for interval {interval} (begin: {begin}, end: {end})"
)))
} else {
Ok(())
}
}
pub async fn inject(
&self,
specifier: &D::Specifier,
interval: Interval,
elems: Vec<D>,
) -> CoreResult<()> {
if elems.is_empty() {
tracing::warn!("No elements to inject");
return Ok(());
}
let first = elems.first().unwrap();
#[cfg(test)]
{
let mut begin = first.basics().begin;
for elem in &elems {
let basics = elem.basics();
self.check_basics(interval, basics)?;
if begin != basics.begin {
return Err(CoreError::ParamError(format!(
"Hole between {begin} and {}",
basics.begin
)));
}
begin = basics.end;
}
}
tracing::debug!(
"Injecting {} in {} ({})",
elems.len(),
specifier.name(),
interval
);
tracing::trace!("elems: {elems:#?}");
let begin = first.basics().begin;
let end = elems.last().unwrap().basics().end;
let mut response = self.db
.query("SELECT * FROM type::table($register) WHERE begin <= $begin AND end >= $begin ORDER BY begin ASC")
.query("SELECT * FROM type::table($register) WHERE begin <= $end AND end >= $end ORDER BY end DESC")
.query("DELETE type::table($register) WHERE begin >= $begin AND begin <= $end OR end >= $begin AND end <= $end OR begin <= $begin AND end >= $end")
.bind(("register", register_from(specifier, interval)))
.bind(("begin", begin))
.bind(("end", end))
.await?;
let begin_section: Vec<DataSection> = response.take(0)?;
let end_section: Vec<DataSection> = response.take(1)?;
if begin_section.len() > 1 || end_section.len() > 1 {
return Err(CoreError::db_error(
"Db corrupted, multiple sections found for a single date",
));
}
tracing::trace!(
"begin_section found locally: {:#?}, \nend_section found locally: {:#?}",
begin_section,
end_section
);
let new_begin = match begin_section.first() {
Some(section) => section.begin,
None => begin,
};
let new_end = match end_section.first() {
Some(section) => section.end,
None => end,
};
let section = DataSection {
begin: new_begin,
end: new_end,
};
let _: Option<DataSection> = self
.db
.upsert((
register_from(specifier, interval),
section.begin.timestamp(),
))
.content(section)
.await?;
for elem in elems {
let _: Option<DataWrapper<D>> = self
.db
.upsert((
table_from(specifier, interval),
elem.basics().begin.timestamp(),
))
.content(DataWrapper::from(elem))
.await?;
}
Ok(())
}
pub async fn try_get_data(&self, params: &DataQuery<D>) -> CoreResult<(Vec<D>, bool)> {
let mut response = self
.db
.query("SELECT * FROM type::table($register) WHERE begin <= $begin AND end >= $begin")
.bind((
"register",
register_from(params.specifier(), *params.interval()),
))
.bind(("begin", *params.begin()))
.await?;
let content: Vec<DataSection> = response.take(0)?;
if content.is_empty() {
return Ok((Vec::new(), false));
}
let section = content.first().unwrap();
let (last_elem_should_begin, last_elem_might_end) =
params.interval().get_time_bounds(section.end)?;
let complete = section.end == *params.end()
|| (last_elem_should_begin > *params.end() && *params.end() <= last_elem_might_end);
tracing::debug!(
"Asking DB for {} ({})",
params.specifier().name(),
params.interval()
);
tracing::trace!(
"table: {}, begin: {}, end: {}",
table_from(params.specifier(), *params.interval()),
params.begin(),
params.end()
);
let mut data_response = self
.db
.query("SELECT * FROM type::table($table) WHERE begin >= $begin AND end <= $end_section AND end <= $end ORDER BY begin ASC")
.bind(("begin", *params.begin()))
.bind(("end_section", section.end))
.bind(("end", *params.end()))
.bind(("table", table_from(params.specifier(), *params.interval())))
.await?;
let wrapped_elems: Vec<DataWrapper<D>> = data_response.take(0)?;
let elems: Vec<D> = wrapped_elems
.into_iter()
.map(DataWrapper::into_data)
.collect();
Ok((elems, complete))
}
}
fn table_from<S: DataSpecifier>(specifier: &S, interval: Interval) -> String {
format!("{}_{}", specifier.name(), interval)
}
fn register_from<S: DataSpecifier>(specifier: &S, interval: Interval) -> String {
format!("{}_{}_reg", specifier.name(), interval)
}
#[cfg(test)]
mod tests {
use chrono::{TimeZone as _, Utc};
use surrealdb::{Surreal, engine::local::Mem};
use crate::provider::{
Store,
data::test::{TestData, TestDataSpecifier},
interval::Interval,
store::{DataBasics, DataQuery},
};
fn default_wrapper() -> DataBasics {
DataBasics {
source: String::from("test"),
..Default::default()
}
}
#[tokio::test]
async fn simple() {
let db = Surreal::new::<Mem>(()).await.unwrap();
db.use_ns("test").use_db("test").await.unwrap();
let store = Store::<TestData>::new(db);
let specifier = TestDataSpecifier::default();
let jan = TestData {
content: String::from("test"),
basics: DataBasics {
begin: Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
end: Utc.with_ymd_and_hms(2024, 2, 1, 0, 0, 0).unwrap(),
..default_wrapper()
},
};
let feb = TestData {
content: String::from("test"),
basics: DataBasics {
begin: Utc.with_ymd_and_hms(2024, 2, 1, 0, 0, 0).unwrap(),
end: Utc.with_ymd_and_hms(2024, 3, 1, 0, 0, 0).unwrap(),
..default_wrapper()
},
};
let mar = TestData {
content: String::from("test"),
basics: DataBasics {
begin: Utc.with_ymd_and_hms(2024, 3, 1, 0, 0, 0).unwrap(),
end: Utc.with_ymd_and_hms(2024, 4, 1, 0, 0, 0).unwrap(),
..default_wrapper()
},
};
let data1 = vec![feb.clone(), mar.clone()];
store
.inject(&specifier, Interval::OneMonth, data1)
.await
.unwrap();
let data2 = vec![jan.clone(), feb.clone()];
store
.inject(&specifier, Interval::OneMonth, data2)
.await
.unwrap();
let (elems, complete) = store
.try_get_data(&DataQuery::new(
TestDataSpecifier::default(),
Interval::OneMonth,
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 4, 1, 0, 0, 0).unwrap(),
))
.await
.unwrap();
assert!(complete);
assert_eq!(elems, vec![jan, feb, mar]);
}
#[tokio::test]
async fn uncomplete() {
let db = Surreal::new::<Mem>(()).await.unwrap();
db.use_ns("test").use_db("test").await.unwrap();
let store = Store::<TestData>::new(db);
let specifier = TestDataSpecifier::default();
let feb = TestData {
content: String::from("test"),
basics: DataBasics {
begin: Utc.with_ymd_and_hms(2024, 2, 1, 0, 0, 0).unwrap(),
end: Utc.with_ymd_and_hms(2024, 3, 1, 0, 0, 0).unwrap(),
..default_wrapper()
},
};
let mar = TestData {
content: String::from("test"),
basics: DataBasics {
begin: Utc.with_ymd_and_hms(2024, 3, 1, 0, 0, 0).unwrap(),
end: Utc.with_ymd_and_hms(2024, 4, 1, 0, 0, 0).unwrap(),
..default_wrapper()
},
};
let elems = vec![feb.clone(), mar.clone()];
store
.inject(&specifier, Interval::OneMonth, elems)
.await
.unwrap();
let (elems, complete) = store
.try_get_data(&DataQuery::new(
TestDataSpecifier::default(),
Interval::OneMonth,
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 4, 1, 0, 0, 0).unwrap(),
))
.await
.unwrap();
assert!(!complete);
assert_eq!(elems, vec![]);
let (elems, complete) = store
.try_get_data(&DataQuery::new(
TestDataSpecifier::default(),
Interval::OneMonth,
Utc.with_ymd_and_hms(2024, 2, 1, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 5, 1, 0, 0, 0).unwrap(),
))
.await
.unwrap();
assert!(!complete);
assert_eq!(elems, vec![feb, mar]);
}
}