use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use serde::{de::DeserializeOwned, Serialize};
use yrs::updates::decoder::Decode;
use yrs::updates::encoder::Encode;
use yrs::{
Array, Doc, GetString, Map, MapRef, ReadTxn, StateVector, Subscription, Text, Transact, Update,
WriteTxn,
};
use crate::error::{CrdtError, CrdtResult};
use crate::reactive::Watcher;
const FIELDS_MAP: &str = "_fields";
const SEP: char = '\u{1}';
pub struct CrdtDoc {
doc: Doc,
fields: MapRef,
actor: u64,
revision: Arc<AtomicU64>,
pending: Arc<Mutex<Vec<Vec<u8>>>>,
applying_remote: Arc<AtomicBool>,
_update_sub: Subscription,
}
impl CrdtDoc {
pub fn new(actor_id: u64) -> Self {
let doc = Doc::with_client_id(actor_id);
let fields = doc.transact_mut().get_or_insert_map(FIELDS_MAP);
let revision = Arc::new(AtomicU64::new(0));
let pending = Arc::new(Mutex::new(Vec::new()));
let applying_remote = Arc::new(AtomicBool::new(false));
let revision_for_obs = Arc::clone(&revision);
let pending_for_obs = Arc::clone(&pending);
let remote_for_obs = Arc::clone(&applying_remote);
let update_sub = doc
.observe_update_v1(move |_txn, event| {
revision_for_obs.fetch_add(1, Ordering::Relaxed);
if !remote_for_obs.load(Ordering::Relaxed) {
pending_for_obs.lock().unwrap().push(event.update.clone());
}
})
.expect("no transaction is active during construction");
Self {
doc,
fields,
actor: actor_id,
revision,
pending,
applying_remote,
_update_sub: update_sub,
}
}
pub fn take_local_updates(&self) -> Vec<Vec<u8>> {
std::mem::take(&mut *self.pending.lock().unwrap())
}
pub fn actor_id(&self) -> u64 {
self.actor
}
pub fn revision(&self) -> u64 {
self.revision.load(Ordering::Relaxed)
}
pub fn watch(&self) -> Watcher {
Watcher::new(Arc::clone(&self.revision))
}
pub fn set_register<T: Serialize>(&self, field: &str, value: &T) -> CrdtResult<()> {
let json = serde_json::to_string(value).map_err(|source| CrdtError::ValueCodec {
field: field.to_string(),
source,
})?;
let mut txn = self.doc.transact_mut();
self.fields.insert(&mut txn, field, json);
Ok(())
}
pub fn register_keys(&self) -> Vec<String> {
let txn = self.doc.transact();
self.fields
.iter(&txn)
.map(|(k, _)| k.to_string())
.filter(|k| !k.starts_with(SEP))
.collect()
}
pub fn remove_register(&self, field: &str) {
let mut txn = self.doc.transact_mut();
self.fields.remove(&mut txn, field);
}
pub fn get_register<T: DeserializeOwned>(&self, field: &str) -> CrdtResult<Option<T>> {
let txn = self.doc.transact();
match self.fields.get(&txn, field) {
None => Ok(None),
Some(out) => {
let json = out.to_string(&txn);
let value = serde_json::from_str(&json).map_err(|source| CrdtError::ValueCodec {
field: field.to_string(),
source,
})?;
Ok(Some(value))
}
}
}
fn counter_key(field: &str, actor: u64) -> String {
format!("{SEP}c{SEP}{field}{SEP}{actor}")
}
fn counter_prefix(field: &str) -> String {
format!("{SEP}c{SEP}{field}{SEP}")
}
pub fn increment(&self, field: &str, delta: i64) {
let key = Self::counter_key(field, self.actor);
let mut txn = self.doc.transact_mut();
let current: i64 = match self.fields.get(&txn, &key) {
Some(out) => out.to_string(&txn).parse().unwrap_or(0),
None => 0,
};
self.fields.insert(&mut txn, key, (current + delta).to_string());
}
pub fn counter(&self, field: &str) -> i64 {
let prefix = Self::counter_prefix(field);
let txn = self.doc.transact();
self.fields
.iter(&txn)
.filter(|(k, _)| k.starts_with(&prefix))
.map(|(_, v)| v.to_string(&txn).parse::<i64>().unwrap_or(0))
.sum()
}
fn text_name(field: &str) -> String {
format!("text{SEP}{field}")
}
pub fn text_push(&self, field: &str, content: &str) {
let mut txn = self.doc.transact_mut();
let text = txn.get_or_insert_text(Self::text_name(field).as_str());
text.push(&mut txn, content);
}
pub fn text_insert(&self, field: &str, index: u32, content: &str) {
let mut txn = self.doc.transact_mut();
let text = txn.get_or_insert_text(Self::text_name(field).as_str());
text.insert(&mut txn, index, content);
}
pub fn text_remove(&self, field: &str, index: u32, len: u32) {
let mut txn = self.doc.transact_mut();
let text = txn.get_or_insert_text(Self::text_name(field).as_str());
text.remove_range(&mut txn, index, len);
}
pub fn text(&self, field: &str) -> String {
let txn = self.doc.transact();
match txn.get_text(Self::text_name(field).as_str()) {
Some(t) => t.get_string(&txn),
None => String::new(),
}
}
pub fn text_len(&self, field: &str) -> u32 {
let txn = self.doc.transact();
txn.get_text(Self::text_name(field).as_str())
.map(|t| t.len(&txn))
.unwrap_or(0)
}
fn set_name(field: &str) -> String {
format!("set{SEP}{field}")
}
pub fn set_add(&self, field: &str, element: &str) {
let mut txn = self.doc.transact_mut();
let arr = txn.get_or_insert_array(Self::set_name(field).as_str());
arr.push_back(&mut txn, element.to_string());
}
pub fn set_remove(&self, field: &str, element: &str) {
let mut txn = self.doc.transact_mut();
let arr = txn.get_or_insert_array(Self::set_name(field).as_str());
let mut matches = Vec::new();
for (i, out) in arr.iter(&txn).enumerate() {
if out.to_string(&txn) == element {
matches.push(i as u32);
}
}
for &i in matches.iter().rev() {
arr.remove_range(&mut txn, i, 1);
}
}
pub fn set_contains(&self, field: &str, element: &str) -> bool {
let txn = self.doc.transact();
match txn.get_array(Self::set_name(field).as_str()) {
Some(arr) => arr.iter(&txn).any(|out| out.to_string(&txn) == element),
None => false,
}
}
pub fn set_members(&self, field: &str) -> Vec<String> {
let txn = self.doc.transact();
let mut members = std::collections::BTreeSet::new();
if let Some(arr) = txn.get_array(Self::set_name(field).as_str()) {
for out in arr.iter(&txn) {
members.insert(out.to_string(&txn));
}
}
members.into_iter().collect()
}
pub fn estimated_state_size(&self) -> usize {
self.encode_full().len()
}
pub fn state_vector(&self) -> Vec<u8> {
self.doc.transact().state_vector().encode_v1()
}
pub fn encode_update_since(&self, their_state_vector: &[u8]) -> CrdtResult<Vec<u8>> {
let sv = StateVector::decode_v1(their_state_vector)
.map_err(|e| CrdtError::DecodeStateVector(e.to_string()))?;
Ok(self.doc.transact().encode_state_as_update_v1(&sv))
}
pub fn encode_full(&self) -> Vec<u8> {
self.doc
.transact()
.encode_state_as_update_v1(&StateVector::default())
}
pub fn apply_update(&self, update: &[u8]) -> CrdtResult<()> {
let update =
Update::decode_v1(update).map_err(|e| CrdtError::DecodeUpdate(e.to_string()))?;
self.applying_remote.store(true, Ordering::Relaxed);
let result = {
let mut txn = self.doc.transact_mut();
txn.apply_update(update)
.map_err(|e| CrdtError::ApplyUpdate(e.to_string()))
};
self.applying_remote.store(false, Ordering::Relaxed);
result
}
pub fn ops_behind(&self, peer_state_vector: &[u8]) -> CrdtResult<usize> {
let peer = StateVector::decode_v1(peer_state_vector)
.map_err(|e| CrdtError::DecodeStateVector(e.to_string()))?;
let txn = self.doc.transact();
let mine = txn.state_vector();
let mut missing = 0usize;
for (client, peer_clock) in peer.iter() {
let my_clock = mine.get(client);
if *peer_clock > my_clock {
missing += (*peer_clock - my_clock) as usize;
}
}
Ok(missing)
}
}