mod tracking;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, RwLock, Weak};
use ciborium::Value as CborValue;
use tokio::sync::{broadcast, watch};
use vantage_core::Result;
use vantage_types::Record;
use crate::dio::{Dio, DioEvent, DioInner, Generation, cbor_scalar_string};
use crate::ops::{ChangeFlash, FlashKind};
#[derive(Debug, Clone)]
pub enum ServoStatus {
Tracking,
Pending,
Failed(String),
}
pub(crate) struct ServoState {
_tally: crate::stats::Tally,
id: RwLock<Option<String>>,
baseline: RwLock<Option<Record<CborValue>>>,
data: RwLock<Record<CborValue>>,
status: RwLock<ServoStatus>,
generation: AtomicU64,
generation_tx: watch::Sender<Generation>,
}
impl ServoState {
fn bump_generation(&self) {
let next = self.generation.fetch_add(1, Ordering::SeqCst) + 1;
let _ = self.generation_tx.send_replace(Generation(next));
}
fn absorb(&self, incoming: Option<Record<CborValue>>) {
{
let mut baseline = self.baseline.write().unwrap();
let mut data = self.data.write().unwrap();
tracking::absorb(&mut baseline, &mut data, incoming);
}
self.bump_generation();
}
fn set_status(&self, status: ServoStatus) {
*self.status.write().unwrap() = status;
self.bump_generation();
}
}
pub struct Servo {
dio: Dio,
state: Arc<ServoState>,
_guard: ServoGuard,
}
struct ServoGuard {
task: tokio::task::JoinHandle<()>,
}
impl Drop for ServoGuard {
fn drop(&mut self) {
self.task.abort();
}
}
impl Servo {
pub fn id(&self) -> Option<String> {
self.state.id.read().unwrap().clone()
}
pub fn set(&self, field: impl Into<String>, value: impl Into<CborValue>) {
self.state
.data
.write()
.unwrap()
.insert(field.into(), value.into());
self.state.bump_generation();
}
pub fn get(&self, field: &str) -> Option<CborValue> {
self.state.data.read().unwrap().get(field).cloned()
}
pub fn record(&self) -> Record<CborValue> {
self.state.data.read().unwrap().clone()
}
pub fn baseline(&self) -> Option<Record<CborValue>> {
self.state.baseline.read().unwrap().clone()
}
pub fn error(&self) -> Record<CborValue> {
let baseline = self.state.baseline.read().unwrap();
let data = self.state.data.read().unwrap();
tracking::error_of(baseline.as_ref(), &data)
}
pub fn dirty(&self, field: &str) -> bool {
self.error().get(field).is_some()
}
pub fn is_dirty(&self) -> bool {
!self.error().is_empty()
}
pub fn revert(&self, field: &str) {
{
let baseline = self.state.baseline.read().unwrap();
let mut data = self.state.data.write().unwrap();
match baseline.as_ref().and_then(|b| b.get(field)).cloned() {
Some(measured) => {
data.insert(field.to_string(), measured);
}
None => {
data.shift_remove(field);
}
}
}
self.state.bump_generation();
}
pub fn revert_all(&self) {
{
let baseline = self.state.baseline.read().unwrap();
let mut data = self.state.data.write().unwrap();
*data = baseline.clone().unwrap_or_default();
}
self.state.bump_generation();
}
pub fn status(&self) -> ServoStatus {
self.state.status.read().unwrap().clone()
}
pub fn subscribe(&self) -> watch::Receiver<Generation> {
self.state.generation_tx.subscribe()
}
pub async fn flash(&self) -> Result<Option<ChangeFlash>> {
let flash = {
let baseline = self.state.baseline.read().unwrap();
let data = self.state.data.read().unwrap();
match baseline.as_ref() {
Some(base) => {
let error = tracking::error_of(Some(base), &data);
if error.is_empty() {
return Ok(None);
}
let id = self
.state
.id
.read()
.unwrap()
.clone()
.expect("a servo with a baseline is bound to an id");
ChangeFlash::new(FlashKind::Patch, Some(id), error).with_before(base.clone())
}
None => {
if data.is_empty() {
return Ok(None);
}
let id = match self.state.id.read().unwrap().clone() {
Some(id) => id,
None => self.id_from_record(&data)?,
};
ChangeFlash::insert(id, data.clone())
}
}
};
let after = flash.after().expect("insert/patch always has an after");
{
*self.state.id.write().unwrap() = flash.id().map(str::to_string);
*self.state.baseline.write().unwrap() = Some(after.clone());
*self.state.data.write().unwrap() = after;
*self.state.status.write().unwrap() = ServoStatus::Pending;
}
self.state.bump_generation();
match self.dio.flash(flash.clone()).await {
Ok(()) => {
self.state.set_status(ServoStatus::Tracking);
Ok(Some(flash))
}
Err(e) => {
self.state.set_status(ServoStatus::Failed(e.to_string()));
Err(e)
}
}
}
pub async fn delete(&self) -> Result<ChangeFlash> {
let (id, before) = {
let id =
self.state.id.read().unwrap().clone().ok_or_else(|| {
vantage_core::error!("an unsaved servo has no record to delete")
})?;
(id, self.state.baseline.read().unwrap().clone())
};
let mut flash = ChangeFlash::delete(id);
if let Some(b) = before {
flash = flash.with_before(b);
}
self.dio.flash(flash.clone()).await?;
{
*self.state.baseline.write().unwrap() = None;
self.state.data.write().unwrap().clear();
}
self.state.bump_generation();
Ok(flash)
}
pub(crate) fn absorb(&self, incoming: Option<Record<CborValue>>) {
self.state.absorb(incoming);
}
fn id_from_record(&self, data: &Record<CborValue>) -> Result<String> {
let id_column = self
.dio
.master()
.get_id_column()
.unwrap_or("id")
.to_string();
let id = data
.get(&id_column)
.map(cbor_scalar_string)
.filter(|s| !s.is_empty());
id.ok_or_else(|| {
vantage_core::error!(
"flashing a new record requires its id field",
id_column = id_column
)
})
}
}
async fn track_loop(
state: Arc<ServoState>,
dio_weak: Weak<DioInner>,
mut bus: broadcast::Receiver<DioEvent>,
) {
loop {
if dio_weak.upgrade().is_none() {
return;
}
let bound = |id: &str| state.id.read().unwrap().as_deref() == Some(id);
match bus.recv().await {
Ok(DioEvent::RecordChanged { id })
| Ok(DioEvent::RecordInserted { id })
| Ok(DioEvent::RecordRemoved { id })
if bound(&id) =>
{
absorb_from_cache(&state, &dio_weak).await;
}
Ok(DioEvent::DatasetChanged) => {
absorb_from_cache(&state, &dio_weak).await;
}
Ok(DioEvent::WritePending { id, kind }) if bound(&id) => {
if matches!(
kind,
crate::FlashKind::Patch | crate::FlashKind::Replace | crate::FlashKind::Insert
) {
state.set_status(ServoStatus::Pending);
}
}
Ok(DioEvent::WriteReverted { id, error, kind }) if bound(&id) => {
if matches!(
kind,
crate::FlashKind::Patch | crate::FlashKind::Replace | crate::FlashKind::Insert
) {
state.set_status(ServoStatus::Failed(error));
}
absorb_from_cache(&state, &dio_weak).await;
}
Ok(_) => {}
Err(broadcast::error::RecvError::Lagged(_)) => {
absorb_from_cache(&state, &dio_weak).await;
}
Err(broadcast::error::RecvError::Closed) => return,
}
}
}
async fn absorb_from_cache(state: &Arc<ServoState>, dio_weak: &Weak<DioInner>) {
let Some(inner) = dio_weak.upgrade() else {
return;
};
let Some(id) = state.id.read().unwrap().clone() else {
return;
};
match inner.cache.get_value(&id).await {
Ok(value) => state.absorb(value),
Err(e) => tracing::error!(error = %e, "servo measurement read failed"),
}
}
pub(crate) fn spawn_servo(dio: &Dio, id: Option<String>) -> Servo {
let (generation_tx, _rx) = watch::channel(Generation::default());
let state = Arc::new(ServoState {
_tally: crate::stats::Tally::servo(),
id: RwLock::new(id),
baseline: RwLock::new(None),
data: RwLock::new(Record::new()),
status: RwLock::new(ServoStatus::Tracking),
generation: AtomicU64::new(0),
generation_tx,
});
let bus_rx = dio.inner.event_bus.subscribe();
let dio_weak = Arc::downgrade(&dio.inner);
let task_state = state.clone();
let task = dio
.inner
.lens
.runtime
.spawn(track_loop(task_state, dio_weak, bus_rx));
Servo {
dio: dio.clone(),
state,
_guard: ServoGuard { task },
}
}