use std::sync::Arc;
use std::time::Duration;
use arc_swap::ArcSwap;
use super::profile;
pub struct EngineHandle(ArcSwap<dataflow_rs::Engine>);
impl EngineHandle {
pub fn new(engine: Arc<dataflow_rs::Engine>) -> Self {
Self(ArcSwap::new(engine))
}
pub fn load(&self) -> Arc<dataflow_rs::Engine> {
self.0.load_full()
}
pub fn store(&self, engine: Arc<dataflow_rs::Engine>) {
self.0.store(engine);
}
}
pub type EngineCallResult = (dataflow_rs::Result<()>, Option<dataflow_rs::ExecutionTrace>);
#[derive(Debug, Clone, Copy)]
pub struct TraceCapture {
pub max_snapshot_bytes: usize,
}
pub async fn run_for_channel(
engine: &Arc<dataflow_rs::Engine>,
channel: &str,
message: &mut dataflow_rs::Message,
timeout_ms: Option<u64>,
profile: Option<&Arc<profile::ProfileCollector>>,
capture: Option<TraceCapture>,
) -> Result<EngineCallResult, u64> {
let run = run_for_channel_inner(engine, channel, message, timeout_ms, capture);
if let Some(p) = profile {
profile::ORION_PROFILE.scope(p.clone(), run).await
} else {
run.await
}
}
async fn with_deadline<F>(timeout_ms: Option<u64>, fut: F) -> Result<F::Output, u64>
where
F: std::future::Future,
{
match timeout_ms {
Some(ms) => tokio::time::timeout(Duration::from_millis(ms), fut)
.await
.map_err(|_| ms),
None => Ok(fut.await),
}
}
fn trace_options(max_snapshot_bytes: usize) -> dataflow_rs::TraceOptions {
dataflow_rs::TraceOptions {
changes: true,
snapshot_audit_trail: dataflow_rs::AuditTrailScope::Own,
max_snapshot_bytes,
redact_paths: vec!["metadata.headers".to_string()],
..Default::default()
}
}
async fn run_for_channel_inner(
engine: &Arc<dataflow_rs::Engine>,
channel: &str,
message: &mut dataflow_rs::Message,
timeout_ms: Option<u64>,
capture: Option<TraceCapture>,
) -> Result<EngineCallResult, u64> {
if let Some(capture) = capture {
let inner = with_deadline(
timeout_ms,
engine.process_message_for_channel_with_trace_options(
channel,
message,
trace_options(capture.max_snapshot_bytes),
),
)
.await?;
Ok(match inner {
Ok(trace) => (Ok(()), Some(trace)),
Err(e) => (Err(e), None),
})
} else {
let inner = with_deadline(
timeout_ms,
engine.process_message_for_channel(channel, message),
)
.await?;
Ok((inner, None))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn a_deadline_applies_regardless_of_trace_capture() {
let slow = || async {
tokio::time::sleep(Duration::from_millis(200)).await;
42
};
assert_eq!(with_deadline(Some(20), slow()).await, Err(20));
assert_eq!(with_deadline(None, slow()).await, Ok(42));
assert_eq!(with_deadline(Some(5_000), slow()).await, Ok(42));
}
}