use std::convert::TryInto;
use std::marker::PhantomData;
use std::sync::Arc;
use scylla_driver::frame::value::MaybeUnset;
use scylla_driver::prepared_statement::PreparedStatement;
use scylla_driver::statement::Consistency;
use scylla_driver::Session;
use tokio::task::JoinSet;
use tokio::time::Instant;
#[cfg(feature = "cbor")]
use super::serializers::CborSerializer;
use super::{ToPrimaryKey, UpstreamSerializer};
use crate::{
CommitToUpstream, DataCommitRequest, DataLoadRequest, LoadFromUpstream, MutableUpstreamFactory,
ServiceData, UpstreamError,
};
const DEFAULT_PRIMARY_KEY_FIELD_NAME: &str = "key";
const DEFAULT_VALUE_FIELD_NAME: &str = "value";
pub struct ScyllaKvUpstreamFactory<S> {
session: Arc<Session>,
select_statement: Arc<PreparedStatement>,
insert_statement: Arc<PreparedStatement>,
serializer: S,
}
pub struct ScyllaKvUpstreamFactoryBuilder<S> {
session: Session,
keyspace_name: String,
table_name: String,
primary_key_field_name: String,
value_field_name: String,
create_table_if_not_exists: bool,
read_consistency: Consistency,
write_consistency: Consistency,
serializer: S,
}
impl<S> ScyllaKvUpstreamFactoryBuilder<S> {
pub fn with_serializer<K, T>(
session: Session,
keyspace_name: K,
table_name: T,
serializer: S,
) -> Self
where
K: Into<String>,
T: Into<String>,
{
Self {
session,
keyspace_name: keyspace_name.into(),
table_name: table_name.into(),
primary_key_field_name: DEFAULT_PRIMARY_KEY_FIELD_NAME.into(),
value_field_name: DEFAULT_VALUE_FIELD_NAME.into(),
create_table_if_not_exists: false,
read_consistency: Consistency::LocalQuorum,
write_consistency: Consistency::LocalQuorum,
serializer,
}
}
pub fn create_table_if_not_exists(mut self) -> Self {
self.create_table_if_not_exists = true;
self
}
pub fn with_field_names<K, V>(mut self, primary_key_field_name: K, value_field_name: V) -> Self
where
K: Into<String>,
V: Into<String>,
{
self.primary_key_field_name = primary_key_field_name.into();
self.value_field_name = value_field_name.into();
self
}
pub async fn build(self) -> Result<ScyllaKvUpstreamFactory<S>, UpstreamError> {
let session = Arc::new(self.session);
session.use_keyspace(self.keyspace_name, false).await?;
if self.create_table_if_not_exists {
session
.query(
format!(
"CREATE TABLE IF NOT EXISTS {} ({} text PRIMARY KEY, {} blob)",
self.table_name, &self.primary_key_field_name, &self.value_field_name,
),
(),
)
.await?;
}
let mut select_statement = session
.prepare(format!(
"SELECT {} FROM {} WHERE {} = ?",
&self.value_field_name, &self.table_name, &self.primary_key_field_name
))
.await?;
select_statement.set_consistency(self.read_consistency);
let mut insert_statement = session
.prepare(format!(
"INSERT INTO {} ({}, {}) VALUES(?, ?) USING TTL ?",
&self.table_name, &self.primary_key_field_name, &self.value_field_name,
))
.await?;
insert_statement.set_consistency(self.write_consistency);
Ok(ScyllaKvUpstreamFactory {
session,
select_statement: Arc::new(select_statement),
insert_statement: Arc::new(insert_statement),
serializer: self.serializer,
})
}
}
impl<K, V, S> MutableUpstreamFactory<K, V> for ScyllaKvUpstreamFactory<S>
where
K: ToPrimaryKey + Send + 'static,
V: ServiceData + Default,
S: UpstreamSerializer<V>,
{
type Upstream = ScyllaKvUpstream<K, V, S>;
fn create(&mut self) -> Self::Upstream {
ScyllaKvUpstream {
_phantom: PhantomData,
session: self.session.clone(),
select_statement: self.select_statement.clone(),
insert_statement: self.insert_statement.clone(),
serializer: self.serializer.clone(),
}
}
}
pub struct ScyllaKvUpstream<K, V, S> {
_phantom: PhantomData<(K, V)>,
session: Arc<Session>,
select_statement: Arc<PreparedStatement>,
insert_statement: Arc<PreparedStatement>,
serializer: S,
}
impl<K, V, S> LoadFromUpstream<K, V> for ScyllaKvUpstream<K, V, S>
where
K: ToPrimaryKey + Send + 'static,
V: ServiceData + Default,
S: UpstreamSerializer<V>,
{
fn load(&mut self, request: DataLoadRequest<K, V>) {
let session = self.session.clone();
let select_statement = self.select_statement.clone();
let key = request.key().to_primary_key();
let serializer = self.serializer.clone();
request.spawn_default(async move {
let result = session.execute(&select_statement, vec![key]).await?;
let row = result.first_row().map_err(|_| UpstreamError::KeyNotFound)?;
let column = &row.columns[0];
let value_blob = column
.as_ref()
.and_then(|cql_value| cql_value.as_blob())
.map(|blob| blob.as_slice())
.ok_or(UpstreamError::KeyNotFound)?;
let value = serializer
.deserialize(value_blob)
.map_err(UpstreamError::serialization_error)?;
Ok(value)
})
}
}
impl<K, V, S> CommitToUpstream<K, V> for ScyllaKvUpstream<K, V, S>
where
K: ToPrimaryKey + Send + 'static,
V: ServiceData,
S: UpstreamSerializer<V>,
{
fn commit(&mut self, request: DataCommitRequest<K, V>) {
let session = self.session.clone();
let insert_statement = self.insert_statement.clone();
let key = request.key().to_primary_key();
let expires_at = request.data().get_expires_at().cloned();
let serialized = match self.serializer.serialize(request.data()) {
Err(e) => return request.reject(UpstreamError::serialization_error(e)),
Ok(vec) => vec,
};
request.into_processing().spawn(async move {
let ttl: MaybeUnset<i32> = match expires_at {
Some(expires_at) => {
let now = Instant::now();
let expires_in_secs = expires_at.saturating_duration_since(now).as_secs();
MaybeUnset::Set(expires_in_secs.try_into().unwrap_or(i32::MAX))
}
None => MaybeUnset::Unset,
};
session
.execute(&insert_statement, (key, serialized, ttl))
.await?;
Ok(())
});
}
}
#[cfg(feature = "cbor")]
pub type ScyllaKvCborUpstreamFactoryBuilder = ScyllaKvUpstreamFactoryBuilder<CborSerializer>;
#[cfg(feature = "cbor")]
pub type ScyllaKvCborUpstreamFactory = ScyllaKvUpstreamFactory<CborSerializer>;
#[cfg(feature = "cbor")]
pub type ScyllaKvCborUpstream<K, V> = ScyllaKvUpstream<K, V, CborSerializer>;
#[cfg(feature = "cbor")]
impl ScyllaKvCborUpstreamFactoryBuilder {
pub fn new<K, T>(session: Session, keyspace_name: K, table_name: T) -> Self
where
K: Into<String>,
T: Into<String>,
{
ScyllaKvUpstreamFactoryBuilder::with_serializer(
session,
keyspace_name,
table_name,
CborSerializer {},
)
}
}
#[cfg(test)]
mod test {
use scylla_driver::SessionBuilder;
use super::{
ScyllaKvCborUpstream, ScyllaKvCborUpstreamFactory, ScyllaKvCborUpstreamFactoryBuilder,
};
use crate::{shard::InternalJoinSetResult, ServiceData, UpstreamError};
use crate::{
CommitToUpstream, DataCommitRequest, DataLoadRequest, LoadFromUpstream,
MutableUpstreamFactory,
};
use tokio::task::JoinSet;
use tokio::time::{sleep, Duration};
use serde::{Deserialize, Serialize};
async fn create_upstream_factory() -> ScyllaKvCborUpstreamFactory {
let session = SessionBuilder::new()
.known_node("127.0.0.1")
.build()
.await
.unwrap();
session.query("CREATE KEYSPACE IF NOT EXISTS test_thingvellir_scylla_upstream WITH replication = {'class': 'SimpleStrategy', 'replication_factor': '1'} AND durable_writes = true", ()).await.unwrap();
session
.query(
"DROP TABLE IF EXISTS test_thingvellir_scylla_upstream.data",
(),
)
.await
.unwrap();
ScyllaKvCborUpstreamFactoryBuilder::new(session, "test_thingvellir_scylla_upstream", "data")
.create_table_if_not_exists()
.with_field_names("key", "value")
.build()
.await
.unwrap()
}
#[derive(Serialize, Deserialize, Default, Debug)]
struct Data {
test: String,
}
impl ServiceData for Data {}
#[tokio::test]
async fn test_upstream_load_commit() -> Result<(), crate::UpstreamError> {
let mut factory = create_upstream_factory().await;
let mut upstream: ScyllaKvCborUpstream<i32, Data> = factory.create();
let mut join_set = JoinSet::new();
{
let request = DataLoadRequest::new(1, &mut join_set);
upstream.load(request);
sleep(Duration::from_millis(500)).await;
let result = match join_set.join_next().await {
Some(data) => data.ok().ok_or(()),
None => Result::Err(()),
};
assert!(matches!(
result,
Ok(InternalJoinSetResult::DataLoadResult(
2i32,
Result::Err(UpstreamError::KeyNotFound)
))
));
}
let mut new_data = Data {
test: "hello cbor!".into(),
};
{
let request = DataCommitRequest::new(2, &mut new_data, &mut join_set);
upstream.commit(request);
sleep(Duration::from_millis(500)).await;
let result = match join_set.join_next().await {
Some(data) => data.ok().ok_or(()),
None => Result::Err(()),
};
assert!(matches!(
result,
Ok(InternalJoinSetResult::DataCommitResult(2i32, Ok(())))
));
}
{
let request = DataLoadRequest::new(2, &mut join_set);
upstream.load(request);
sleep(Duration::from_millis(500)).await;
let result = match join_set.join_next().await {
Some(data) => data.ok().ok_or(()),
None => Result::Err(()),
};
match result.ok().unwrap() {
InternalJoinSetResult::DataLoadResult(key, Ok(data)) => {
assert_eq!(key, 2);
assert_eq!(data.test, new_data.test);
}
_ => unreachable!(),
}
}
Ok(())
}
}