thingvellir 0.0.14

a concurrent, shared-nothing abstraction that manages an assembly of things
Documentation
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();

        // Try loading data that does not exist!
        {
            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)
                ))
            ));
        }

        // Commit with some data:
        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(())))
            ));
        }

        // Try loading that data again.
        {
            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(())
    }
}