use std::sync::Arc;
use async_trait::async_trait;
use dataflow_rs::engine::error::DataflowError;
use dataflow_rs::engine::task_context::TaskContext;
use serde_json::Value;
use super::connector_handler::{ConnectorHandler, Produced};
use super::connector_helpers::{
ConnectorCall, json_type_name, require_op, resolve_required_str, resolve_value,
to_connect_error, to_exec_error,
};
use super::schema::{FieldKind, FieldSchema};
use super::templated_input::TemplatedInput;
use crate::connector::ConnectorRegistry;
use crate::connector::cache_backend::{CachePool, CachePurpose};
use crate::engine::HandlerError;
pub struct CacheWriteHandler {
pub cache_pool: Arc<CachePool>,
pub registry: Arc<ConnectorRegistry>,
}
pub struct CacheWrite {
key: String,
value: Value,
ttl: Option<u64>,
}
#[async_trait]
impl ConnectorHandler for CacheWriteHandler {
const NAME: &'static str = "cache_write";
type Kind = crate::connector::kind::Cache;
type Input = TemplatedInput;
type Parsed = CacheWrite;
fn registry(&self) -> &Arc<ConnectorRegistry> {
&self.registry
}
fn parse(
&self,
call: &ConnectorCall<'_>,
input: &TemplatedInput,
ctx: &TaskContext<'_>,
) -> Result<Self::Parsed, HandlerError> {
let key = resolve_required_str(input, "key", call.name, ctx)?;
let value = match input.get("value") {
Some(v) => resolve_value(v, ctx),
None => {
return Err(
DataflowError::Validation(format!("{} requires 'value'", call.name)).into(),
);
}
};
Ok(CacheWrite {
key,
value,
ttl: resolve_ttl_secs(input, call.name, ctx)?,
})
}
fn gate(
_parsed: &Self::Parsed,
conn: &crate::connector::CacheConnectorConfig,
connector: &str,
) -> Result<(), HandlerError> {
Ok(require_op(conn.operations.write, "write", connector)?)
}
async fn run(
&self,
write: Self::Parsed,
conn: &crate::connector::CacheConnectorConfig,
call: &ConnectorCall<'_>,
_input: &TemplatedInput,
_ctx: &mut TaskContext<'_>,
) -> Result<Produced, HandlerError> {
let backend = self
.cache_pool
.get_backend(CachePurpose::Workflow, call.connector, conn)
.await
.map_err(to_connect_error)?;
let value_str = serde_json::to_string(&write.value).map_err(|e| {
DataflowError::Validation(format!("Failed to serialize value for cache: {e}"))
})?;
match write.ttl {
Some(ttl) => backend
.set_ex(&write.key, &value_str, ttl)
.await
.map_err(to_exec_error)?,
None => backend
.set(&write.key, &value_str)
.await
.map_err(to_exec_error)?,
}
tracing::debug!(key = %write.key, ttl = ?write.ttl, "Wrote value to cache");
Ok(Produced::nothing())
}
}
fn resolve_ttl_secs(
input: &TemplatedInput,
name: &str,
ctx: &TaskContext<'_>,
) -> Result<Option<u64>, DataflowError> {
let Some(raw) = input.get("ttl_secs") else {
return Ok(None);
};
match resolve_value(raw, ctx) {
Value::Null => Ok(None),
Value::Number(n) => {
if let Some(u) = n.as_u64() {
Ok(Some(u))
} else if let Some(f) = n.as_f64()
&& f >= 0.0
&& f.fract() == 0.0
{
Ok(Some(f as u64))
} else {
Err(DataflowError::Validation(format!(
"{name} 'ttl_secs' must be a non-negative whole number of seconds, got {n}"
)))
}
}
other => Err(DataflowError::Validation(format!(
"{name} 'ttl_secs' must resolve to a number, got {}",
json_type_name(&other)
))),
}
}
pub(super) const CACHE_WRITE_FIELDS: &[FieldSchema] = &[
FieldSchema {
name: "connector",
description: "Name of the cache connector to write to.",
kind: FieldKind::String,
required: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "key",
description: "Cache key to set. JSONLogic: a literal, or an expression over the message.",
kind: FieldKind::String,
required: true,
template_at: &[""],
..FieldSchema::DEFAULT
},
FieldSchema {
name: "value",
description: "Value to store. May be any JSON value. Accepts {\"var\": \"path\"} to read the value from the message.",
kind: FieldKind::Any,
required: true,
resolvable: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "ttl_secs",
description: "Time-to-live in seconds. Omit for no expiry. JSONLogic: a literal, or an expression over the message.",
kind: FieldKind::Number,
template_at: &[""],
..FieldSchema::DEFAULT
},
];