use std::sync::Arc;
use std::time::Instant;
use async_trait::async_trait;
use dataflow_rs::TemplateCompiler;
use dataflow_rs::engine::error::DataflowError;
use dataflow_rs::engine::functions::AsyncFunctionHandler;
use dataflow_rs::engine::task_context::TaskContext;
use dataflow_rs::engine::task_outcome::TaskOutcome;
use serde_json::{Map, Value};
use tokio::sync::Semaphore;
use super::error::{Category, Failure};
use super::limits::Limits;
use super::runtime::{LoadedComponent, WasmRuntime};
use crate::engine::FunctionEntry;
use crate::engine::functions::connector_helpers::{apply_output, resolve_value};
use crate::engine::functions::templated_input::TemplatedInput;
use crate::plugin::manifest::OUTPUT_FIELD;
#[derive(Clone)]
pub struct PluginFunctionHandler {
pub entry: Arc<FunctionEntry>,
pub loaded: Arc<LoadedComponent>,
pub runtime: Arc<WasmRuntime>,
pub limits: Limits,
resolved: Arc<FunctionEntry>,
permits: Arc<Semaphore>,
plugin_id: String,
}
impl PluginFunctionHandler {
pub fn new(
entry: Arc<FunctionEntry>,
loaded: Arc<LoadedComponent>,
runtime: Arc<WasmRuntime>,
limits: Limits,
) -> Self {
let plugin_id = entry
.plugin
.as_ref()
.map(|p| p.id.clone())
.unwrap_or_default();
let mut resolved = (*entry).clone();
if let Some(fields) = resolved.input_fields.as_mut() {
for field in fields {
field.resolvable = false;
field.template_at = &[];
}
}
Self {
entry,
loaded,
runtime,
limits,
resolved: Arc::new(resolved),
permits: Arc::new(Semaphore::new(limits.max_concurrency as usize)),
plugin_id,
}
}
pub fn name(&self) -> &str {
&self.entry.name
}
async fn run(&self, ctx: &mut TaskContext<'_>, input: &TemplatedInput) -> Result<(), Failure> {
let raw = input.raw();
let obj = raw
.as_object()
.ok_or_else(|| Failure::host(Category::CallerInput, "input must be a JSON object"))?;
let output = match obj.get(OUTPUT_FIELD) {
Some(Value::String(path)) if !path.trim().is_empty() => path.clone(),
Some(_) => {
return Err(Failure::host(
Category::CallerInput,
"'output' must be a non-empty dotted context path",
));
}
None => match self.entry.writes {
crate::engine::functions::schema::WriteShape::OutputPath {
default_root: Some(root),
} => root.to_string(),
_ => {
return Err(Failure::host(
Category::CallerInput,
"the task names no 'output' and the function declares no default root",
));
}
},
};
let mut guest = Map::new();
for field in self.entry.input_fields.as_deref().unwrap_or(&[]) {
if field.name == OUTPUT_FIELD {
continue;
}
let Some(value) = obj.get(&field.name) else {
continue;
};
let value = match input.template_value(&field.name, ctx) {
Some(Ok(evaluated)) => evaluated,
Some(Err(e)) => {
return Err(Failure::host(
Category::CallerInput,
format!("'{}' did not evaluate: {e}", field.name),
));
}
None if field.resolvable => resolve_value(value, ctx),
None => value.clone(),
};
guest.insert(field.name.clone(), value);
}
let guest = Value::Object(guest);
let problems = self.resolved.validate_input(&guest, "task");
if !problems.is_empty() {
let text = problems
.iter()
.map(|p| format!("{} ({})", p.message, p.code))
.collect::<Vec<_>>()
.join("; ");
return Err(Failure::host(Category::CallerInput, text));
}
let text = serde_json::to_string(&guest).map_err(|e| {
Failure::host(Category::CallerInput, "input is not serialisable").with_detail(e)
})?;
if text.len() > self.limits.max_request_bytes {
return Err(Failure::host(
Category::RequestSize,
format!(
"the input is {} bytes, over the {} byte limit",
text.len(),
self.limits.max_request_bytes
),
));
}
let queued = Instant::now();
let permit = tokio::time::timeout(self.limits.timeout, self.permits.acquire())
.await
.ok()
.and_then(Result::ok)
.ok_or_else(|| {
Failure::host(
Category::Permit,
format!(
"no concurrency permit within {:?} ({} already running)",
self.limits.timeout, self.limits.max_concurrency
),
)
})?;
crate::metrics::record_plugin_queue_time(
&self.plugin_id,
self.entry.label(),
queued.elapsed().as_secs_f64(),
);
let json = self
.runtime
.invoke(&self.loaded, &self.limits, &self.entry.name, &text)
.await
.map_err(super::runtime::Invocation::into_failure)?;
drop(permit);
let value: Value = serde_json::from_str(&json).map_err(|e| {
Failure::host(
Category::BadResult,
"the plugin returned something that is not JSON",
)
.with_detail(e)
})?;
apply_output(ctx, &output, value);
Ok(())
}
}
#[async_trait]
impl AsyncFunctionHandler for PluginFunctionHandler {
type Input = TemplatedInput;
fn parse_input_with(&self, input: &Value) -> dataflow_rs::Result<TemplatedInput> {
let problems = self.entry.validate_input(input, "task");
if !problems.is_empty() {
return Err(DataflowError::Validation(
problems
.iter()
.map(|p| format!("{} ({})", p.message, p.code))
.collect::<Vec<_>>()
.join("; "),
));
}
Ok(TemplatedInput::from(input.clone()))
}
fn compile_input_with(
&self,
input: &mut TemplatedInput,
c: &TemplateCompiler,
) -> dataflow_rs::Result<()> {
input.compile_fields(
&self.entry.name,
self.entry.input_fields.as_deref().unwrap_or(&[]),
c,
)
}
async fn execute(
&self,
ctx: &mut TaskContext<'_>,
input: &TemplatedInput,
) -> dataflow_rs::Result<TaskOutcome> {
let started = Instant::now();
let result = self.run(ctx, input).await;
let secs = started.elapsed().as_secs_f64();
match result {
Ok(()) => {
crate::metrics::record_plugin_invocation(
&self.plugin_id,
self.entry.label(),
"ok",
secs,
);
Ok(TaskOutcome::Success)
}
Err(failure) => {
crate::metrics::record_plugin_invocation(
&self.plugin_id,
self.entry.label(),
"error",
secs,
);
crate::metrics::record_plugin_failure(
&self.plugin_id,
self.entry.label(),
failure.category.as_str(),
);
if let Some(detail) = &failure.detail {
tracing::warn!(
plugin = %self.plugin_id,
function = %self.entry.name,
digest = %self.loaded.digest,
workflow = ctx.workflow_id().unwrap_or("-"),
task = ctx.task_id().unwrap_or("-"),
category = failure.category.as_str(),
detail = %detail,
"Plugin invocation failed"
);
}
Err(failure
.into_handler_error()
.prefixed(&self.entry.name)
.into())
}
}
}
}