use crate::pipeline::apply::event::PositionedEvent;
use crate::pipeline::apply::feed::ChangeFeed;
use crate::pipeline::apply::opts::ApplyOpts;
use crate::pipeline::apply::runtime::{
apply_changes_with, apply_relation_changes_with, apply_transformed_sink_events,
write_relations_with, write_rows_with, ApplyContext,
};
use crate::pipeline::apply::transform::BatchTransformer;
use crate::pipeline::pipeline::Pipeline;
use anyhow::{Context, Result};
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{Duration, Instant};
use surreal_sync_core::SurrealSink;
use surreal_sync_core::{Change, Relation, RelationChange, Row};
use tracing::debug;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ControlSignal {
SchemaRefresh,
AdHocSnapshot {
tables: Vec<String>,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum StopReason {
Cancelled,
Deadline,
Until,
Finished,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum CheckpointPolicy {
#[default]
PersistAfterAdvance,
AdvanceOnly,
IntervalWhenDrained {
interval: Duration,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RuntimeExit {
Stopped(StopReason),
}
#[async_trait::async_trait]
pub trait AdhocApply: Send + Sync {
async fn write_rows(&self, rows: Vec<Row>) -> Result<()>;
async fn write_relations(&self, relations: Vec<Relation>) -> Result<()>;
async fn apply_changes(&self, changes: Vec<Change>) -> Result<()>;
async fn apply_relation_changes(&self, changes: Vec<RelationChange>) -> Result<()>;
fn apply_opts(&self) -> &ApplyOpts;
}
struct AdhocApplyImpl<'a, S, T> {
sink: &'a S,
transformer: Arc<T>,
apply_opts: &'a ApplyOpts,
}
#[async_trait::async_trait]
impl<'a, S, T> AdhocApply for AdhocApplyImpl<'a, S, T>
where
S: SurrealSink,
T: BatchTransformer + 'static,
{
async fn write_rows(&self, rows: Vec<Row>) -> Result<()> {
write_rows_with(
self.sink,
Arc::clone(&self.transformer),
rows,
self.apply_opts,
)
.await
}
async fn write_relations(&self, relations: Vec<Relation>) -> Result<()> {
write_relations_with(
self.sink,
Arc::clone(&self.transformer),
relations,
self.apply_opts,
)
.await
}
async fn apply_changes(&self, changes: Vec<Change>) -> Result<()> {
apply_changes_with(
self.sink,
Arc::clone(&self.transformer),
changes,
self.apply_opts,
)
.await
}
async fn apply_relation_changes(&self, changes: Vec<RelationChange>) -> Result<()> {
apply_relation_changes_with(
self.sink,
Arc::clone(&self.transformer),
changes,
self.apply_opts,
)
.await
}
fn apply_opts(&self) -> &ApplyOpts {
self.apply_opts
}
}
#[async_trait::async_trait]
pub trait SourceDriver: Send {
type Position: Clone + Send + Sync + 'static;
async fn poll_work(&mut self) -> Result<Vec<PositionedEvent<Self::Position>>>;
async fn advance_watermark(&mut self, position: Self::Position) -> Result<()>;
fn is_finished(&self) -> bool {
false
}
async fn between_events(&mut self) -> Result<Vec<ControlSignal>> {
Ok(Vec::new())
}
async fn on_schema_refresh(&mut self) -> Result<()> {
Ok(())
}
async fn on_adhoc_snapshot(
&mut self,
_tables: &[String],
_apply: &dyn AdhocApply,
) -> Result<()> {
Ok(())
}
fn stop_reason(&self) -> Option<StopReason> {
None
}
fn checkpoint_policy(&self) -> CheckpointPolicy {
CheckpointPolicy::PersistAfterAdvance
}
async fn persist_checkpoint(&mut self, _position: Self::Position) -> Result<()> {
Ok(())
}
async fn read_progress_for_persist(&mut self) -> Result<Option<Self::Position>> {
Ok(None)
}
fn note_sunk_events(&mut self, _count: u64) {}
}
pub struct ChangeFeedDriver<F> {
pub inner: F,
}
impl<F> ChangeFeedDriver<F> {
pub fn new(inner: F) -> Self {
Self { inner }
}
pub fn inner(&self) -> &F {
&self.inner
}
pub fn inner_mut(&mut self) -> &mut F {
&mut self.inner
}
pub fn into_inner(self) -> F {
self.inner
}
}
#[async_trait::async_trait]
impl<F> SourceDriver for ChangeFeedDriver<F>
where
F: ChangeFeed,
{
type Position = F::Position;
async fn poll_work(&mut self) -> Result<Vec<PositionedEvent<Self::Position>>> {
let changes = self.inner.poll_changes().await?;
Ok(changes.into_iter().map(PositionedEvent::from).collect())
}
async fn advance_watermark(&mut self, position: Self::Position) -> Result<()> {
self.inner.advance_watermark(position).await
}
fn is_finished(&self) -> bool {
self.inner.is_finished()
}
}
pub struct ChangeFeedRef<'a, F: ChangeFeed> {
inner: &'a mut F,
}
impl<'a, F: ChangeFeed> ChangeFeedRef<'a, F> {
pub fn new(inner: &'a mut F) -> Self {
Self { inner }
}
}
#[async_trait::async_trait]
impl<'a, F> SourceDriver for ChangeFeedRef<'a, F>
where
F: ChangeFeed,
{
type Position = F::Position;
async fn poll_work(&mut self) -> Result<Vec<PositionedEvent<Self::Position>>> {
let changes = self.inner.poll_changes().await?;
Ok(changes.into_iter().map(PositionedEvent::from).collect())
}
async fn advance_watermark(&mut self, position: Self::Position) -> Result<()> {
self.inner.advance_watermark(position).await
}
fn is_finished(&self) -> bool {
self.inner.is_finished()
}
}
#[derive(Debug, Clone, Default)]
pub struct SourceRuntimeOpts {
pub deadline: Option<Instant>,
pub cancelled: bool,
}
impl SourceRuntimeOpts {
pub fn new() -> Self {
Self::default()
}
pub fn with_deadline(mut self, deadline: Instant) -> Self {
self.deadline = Some(deadline);
self
}
pub fn with_cancelled(mut self, cancelled: bool) -> Self {
self.cancelled = cancelled;
self
}
}
pub async fn run_source_runtime<D, S>(
driver: &mut D,
sink: &S,
pipeline: &Pipeline,
apply_opts: &ApplyOpts,
runtime_opts: &SourceRuntimeOpts,
) -> Result<RuntimeExit>
where
D: SourceDriver,
S: SurrealSink,
{
run_source_runtime_with(
driver,
sink,
Arc::new(pipeline.clone()),
apply_opts,
runtime_opts,
)
.await
}
pub async fn run_source_runtime_with<D, S, T>(
driver: &mut D,
sink: &S,
transformer: Arc<T>,
apply_opts: &ApplyOpts,
runtime_opts: &SourceRuntimeOpts,
) -> Result<RuntimeExit>
where
D: SourceDriver,
S: SurrealSink,
T: BatchTransformer + 'static,
{
let mut ctx = ApplyContext::new(sink, Arc::clone(&transformer), apply_opts);
let mut sinking: Option<PendingSink<'_, D::Position>> = None;
loop {
if let Some(reason) = effective_stop_reason(driver, runtime_opts) {
finish_pending_sink(&mut ctx, driver, &mut sinking).await?;
ctx.flush_for_driver(driver).await?;
return Ok(RuntimeExit::Stopped(reason));
}
let signals = driver.between_events().await.context("between_events")?;
if !signals.is_empty() {
finish_pending_sink(&mut ctx, driver, &mut sinking).await?;
handle_control_signals(driver, &mut ctx, sink, &transformer, apply_opts, signals)
.await?;
}
loop {
if let Some(reason) = effective_stop_reason(driver, runtime_opts) {
finish_pending_sink(&mut ctx, driver, &mut sinking).await?;
ctx.flush_for_driver(driver).await?;
return Ok(RuntimeExit::Stopped(reason));
}
ctx.poll_join_ready_public().await?;
try_launch_sink(sink, &mut ctx, &mut sinking)?;
if ctx.window_occupancy() >= apply_opts.max_in_flight {
break;
}
while ctx.buffer_len() < apply_opts.batch_size && !driver.is_finished() {
let polled = driver.poll_work().await.context("poll_work")?;
if polled.is_empty() {
break;
}
for pe in polled {
ctx.push_buffered_event(pe);
}
}
let started = if ctx.buffer_len() >= apply_opts.batch_size {
ctx.try_start_full_batch()
} else if ctx.buffer_len() > 0
&& (driver.is_finished() || ctx.should_flush_partial_public())
{
ctx.try_start_partial_batch()
} else {
false
};
if !started {
break;
}
ctx.poll_join_ready_public().await?;
try_launch_sink(sink, &mut ctx, &mut sinking)?;
}
let has_transform = ctx.in_flight_count() > 0;
let has_sink = sinking.is_some();
let has_buffer = ctx.buffer_len() > 0;
let finished = driver.is_finished();
if !has_transform && !has_sink && !has_buffer {
if finished {
ctx.flush_for_driver(driver).await?;
return Ok(RuntimeExit::Stopped(StopReason::Finished));
}
if let Some(reason) = effective_stop_reason(driver, runtime_opts) {
ctx.flush_for_driver(driver).await?;
return Ok(RuntimeExit::Stopped(reason));
}
ctx.try_interval_persist_public(driver).await?;
tokio::time::sleep(apply_opts.batch_max_wait.min(Duration::from_millis(10))).await;
continue;
}
if !has_transform && !has_sink {
tokio::time::sleep(apply_opts.batch_max_wait.min(Duration::from_millis(10))).await;
continue;
}
if has_sink && ctx.window_occupancy() < apply_opts.max_in_flight && !finished {
let idle = apply_opts.batch_max_wait.min(Duration::from_millis(10));
if has_transform {
tokio::select! {
biased;
outcome = ctx.wait_one_completion_public() => {
outcome?;
try_launch_sink(sink, &mut ctx, &mut sinking)?;
}
result = poll_pending_sink(&mut sinking) => {
complete_pending_sink(&mut ctx, driver, &mut sinking, result).await?;
ctx.try_interval_persist_public(driver).await?;
try_launch_sink(sink, &mut ctx, &mut sinking)?;
}
_ = tokio::time::sleep(idle) => {}
}
} else {
tokio::select! {
result = poll_pending_sink(&mut sinking) => {
complete_pending_sink(&mut ctx, driver, &mut sinking, result).await?;
ctx.try_interval_persist_public(driver).await?;
try_launch_sink(sink, &mut ctx, &mut sinking)?;
}
_ = tokio::time::sleep(idle) => {}
}
}
continue;
}
tokio::select! {
biased;
outcome = ctx.wait_one_completion_public(), if has_transform => {
outcome?;
try_launch_sink(sink, &mut ctx, &mut sinking)?;
}
result = poll_pending_sink(&mut sinking), if has_sink => {
complete_pending_sink(&mut ctx, driver, &mut sinking, result).await?;
ctx.try_interval_persist_public(driver).await?;
try_launch_sink(sink, &mut ctx, &mut sinking)?;
}
}
}
}
struct PendingSinkMeta<P> {
batch_id: u64,
last_position: P,
event_count: u64,
sunk: u64,
}
struct PendingSink<'s, P> {
meta: PendingSinkMeta<P>,
drive: Pin<Box<dyn Future<Output = Result<SinkDrive>> + Send + 's>>,
}
fn try_launch_sink<'s, S, T, P>(
sink: &'s S,
ctx: &mut ApplyContext<'_, S, T, P>,
sinking: &mut Option<PendingSink<'s, P>>,
) -> Result<()>
where
S: SurrealSink,
T: BatchTransformer + 'static,
P: Clone + Send + Sync + 'static,
{
if sinking.is_some() {
return Ok(());
}
let Some(batch) = ctx.prepare_ordered_sink() else {
return Ok(());
};
match batch.result {
Ok(events) => {
let event_count = batch.event_count;
let drive = Box::pin(async move {
apply_transformed_sink_events(sink, &events).await?;
Ok(SinkDrive::Applied)
});
*sinking = Some(PendingSink {
meta: PendingSinkMeta {
batch_id: batch.batch_id,
last_position: batch.last_position,
event_count,
sunk: event_count,
},
drive,
});
}
Err(e) => {
let drive = Box::pin(async move { Ok(SinkDrive::TransformFailed(e)) });
*sinking = Some(PendingSink {
meta: PendingSinkMeta {
batch_id: batch.batch_id,
last_position: batch.last_position,
event_count: batch.event_count,
sunk: 0,
},
drive,
});
}
}
Ok(())
}
enum SinkDrive {
Applied,
TransformFailed(anyhow::Error),
}
async fn poll_pending_sink<'s, P>(sinking: &mut Option<PendingSink<'s, P>>) -> Result<SinkDrive> {
match sinking.as_mut() {
Some(pending) => pending.drive.as_mut().await,
None => std::future::pending().await,
}
}
async fn complete_pending_sink<D, S, T>(
ctx: &mut ApplyContext<'_, S, T, D::Position>,
driver: &mut D,
sinking: &mut Option<PendingSink<'_, D::Position>>,
result: Result<SinkDrive>,
) -> Result<()>
where
D: SourceDriver,
S: SurrealSink,
T: BatchTransformer + 'static,
{
let pending = sinking.take().expect("sink slot");
match result {
Ok(SinkDrive::Applied) => {
ctx.finish_sink_ok_driver(driver, pending.meta.last_position, pending.meta.sunk)
.await
}
Ok(SinkDrive::TransformFailed(e)) => {
ctx.finish_sink_err_driver(
driver,
pending.meta.batch_id,
pending.meta.last_position,
pending.meta.event_count,
e,
)
.await
}
Err(e) => {
ctx.finish_sink_err_driver(
driver,
pending.meta.batch_id,
pending.meta.last_position,
pending.meta.event_count,
e,
)
.await
}
}
}
async fn finish_pending_sink<D, S, T>(
ctx: &mut ApplyContext<'_, S, T, D::Position>,
driver: &mut D,
sinking: &mut Option<PendingSink<'_, D::Position>>,
) -> Result<()>
where
D: SourceDriver,
S: SurrealSink,
T: BatchTransformer + 'static,
{
if sinking.is_none() {
return Ok(());
}
let result = poll_pending_sink(sinking).await;
complete_pending_sink(ctx, driver, sinking, result).await
}
fn effective_stop_reason<D: SourceDriver>(
driver: &D,
runtime_opts: &SourceRuntimeOpts,
) -> Option<StopReason> {
if runtime_opts.cancelled {
return Some(StopReason::Cancelled);
}
if let Some(deadline) = runtime_opts.deadline {
if Instant::now() >= deadline {
return Some(StopReason::Deadline);
}
}
driver.stop_reason()
}
async fn handle_control_signals<D, S, T>(
driver: &mut D,
ctx: &mut ApplyContext<'_, S, T, D::Position>,
sink: &S,
transformer: &Arc<T>,
apply_opts: &ApplyOpts,
signals: Vec<ControlSignal>,
) -> Result<()>
where
D: SourceDriver,
S: SurrealSink,
T: BatchTransformer + 'static,
{
ctx.flush_for_driver(driver).await?;
for signal in signals {
match signal {
ControlSignal::SchemaRefresh => {
debug!("SourceDriver schema refresh");
driver
.on_schema_refresh()
.await
.context("on_schema_refresh")?;
}
ControlSignal::AdHocSnapshot { tables } => {
debug!(?tables, "SourceDriver ad-hoc snapshot");
let apply = AdhocApplyImpl {
sink,
transformer: Arc::clone(transformer),
apply_opts,
};
driver
.on_adhoc_snapshot(&tables, &apply)
.await
.context("on_adhoc_snapshot")?;
}
}
}
Ok(())
}