use std::sync::Arc;
use std::time::Duration;
use dataflow_rs::engine::error::DataflowError;
use dataflow_rs::engine::task_context::TaskContext;
use dataflow_rs::engine::task_outcome::TaskOutcome;
use serde_json::{Map, Value};
use crate::connector::{
CacheConnectorConfig, ConnectorConfig, ConnectorRegistry, DbConnectorConfig, EsConnectorConfig,
HttpOperationGates, OperationGates,
};
use crate::query::EntityRegistry;
pub fn build_entity_registry(
schema: Option<&Value>,
connector_config: &ConnectorConfig,
connector_name: &str,
) -> Result<EntityRegistry, DataflowError> {
let mut registry = match schema {
Some(s) => EntityRegistry::from_json(s)?,
None => EntityRegistry::default(),
};
if let Some(guards) = connector_config.dialect_guards() {
if !guards.schema_is_sufficient(!registry.is_empty(), registry.is_identity_mode()) {
return Err(crate::errors::connector_detail_error(format!(
"connector '{connector_name}' requires a declared schema \
(dialect.require_schema): supply \"schema\" with an \"entities\" map \
and without \"unmapped\": \"identity\""
)));
}
registry.restrict_to(&guards.allowed_entities);
}
Ok(registry)
}
pub fn require_op_allowed(
gates: &OperationGates,
op: &str,
connector_name: &str,
) -> Result<(), DataflowError> {
require_op(gates.allows(op), op, connector_name)
}
pub fn require_op(allowed: bool, op: &str, connector_name: &str) -> Result<(), DataflowError> {
if !allowed {
return Err(crate::errors::connector_detail_error(format!(
"operation '{op}' is disabled on connector '{connector_name}'"
)));
}
Ok(())
}
pub fn require_method_allowed(
gates: &HttpOperationGates,
method: &str,
connector_name: &str,
) -> Result<(), DataflowError> {
if !gates.allows_method(method) {
return Err(crate::errors::connector_detail_error(format!(
"HTTP method '{method}' is not allowed on connector \
'{connector_name}' (allowed: {})",
gates.methods.join(", ")
)));
}
Ok(())
}
pub async fn es_request(
client: &reqwest::Client,
es: &EsConnectorConfig,
method: reqwest::Method,
url: &str,
) -> Result<reqwest::RequestBuilder, DataflowError> {
if !es.allow_private_urls
&& let Err(msg) = crate::validation::validate_url_not_private(url).await
{
return Err(DataflowError::function_execution(
format!("SSRF protection: {msg}"),
None,
));
}
let mut req = client.request(method, url);
if let Some(auth) = &es.auth {
req = super::http_common::apply_auth(req, auth);
}
if let Some(ms) = es.request_timeout_ms {
req = req.timeout(Duration::from_millis(ms));
}
Ok(req)
}
pub async fn read_es_body(
resp: reqwest::Response,
max_size: usize,
) -> Result<Value, DataflowError> {
if let Some(len) = resp.content_length()
&& len as usize > max_size
{
return Err(DataflowError::function_execution(
format!(
"Elasticsearch response declared Content-Length {len} exceeds \
limit of {max_size} bytes"
),
None,
));
}
let bytes = resp.bytes().await.map_err(to_exec_error)?;
if bytes.len() > max_size {
return Err(DataflowError::function_execution(
format!(
"Elasticsearch response body is {} bytes, exceeding limit of {max_size} bytes",
bytes.len()
),
None,
));
}
serde_json::from_slice(&bytes).map_err(to_exec_error)
}
pub async fn send_es(
req: reqwest::RequestBuilder,
max_response_size: usize,
) -> Result<(reqwest::StatusCode, Value), DataflowError> {
let resp = req.send().await.map_err(to_exec_error)?;
let status = resp.status();
let body: Value = read_es_body(resp, max_response_size).await?;
Ok((status, body))
}
pub fn es_write_error(status: reqwest::StatusCode, body: &Value) -> DataflowError {
DataflowError::function_execution(
format!("Elasticsearch write failed ({status}): {body}"),
None,
)
}
pub struct ConnectorCall<'a> {
pub name: &'static str,
pub connector: &'a str,
pub channel: String,
pub output: &'a str,
}
impl<'a> ConnectorCall<'a> {
pub fn begin(
name: &'static str,
input: &'a Value,
ctx: &TaskContext<'_>,
) -> Result<Self, DataflowError> {
Ok(Self {
name,
connector: require_str_field(input, "connector", name)?,
channel: super::extract_channel(ctx.message()).to_string(),
output: extract_output_path(input),
})
}
pub fn require_str<'i>(&self, input: &'i Value, field: &str) -> Result<&'i str, DataflowError> {
require_str_field(input, field, self.name)
}
pub async fn resolve(
&self,
registry: &ConnectorRegistry,
op: Option<&str>,
) -> Result<Arc<ConnectorConfig>, DataflowError> {
let config = resolve_connector(registry, self.connector).await?;
if let Some(op) = op
&& let Some(gates) = config.operation_gates()
{
require_op_allowed(gates, op, self.connector)?;
}
Ok(config)
}
pub async fn run<F>(
&self,
registry: &ConnectorRegistry,
fut: F,
) -> dataflow_rs::Result<TaskOutcome>
where
F: std::future::Future<Output = dataflow_rs::Result<TaskOutcome>>,
{
guarded_handler(self.name, registry, self.connector, &self.channel, fut).await
}
}
pub async fn guarded_handler<F>(
fn_name: &'static str,
registry: &ConnectorRegistry,
connector: &str,
channel: &str,
fut: F,
) -> dataflow_rs::Result<TaskOutcome>
where
F: std::future::Future<Output = dataflow_rs::Result<TaskOutcome>>,
{
if !registry.circuit_breaker_enabled() {
return observed_handler_named(fn_name, connector, channel, fut).await;
}
let breaker = registry
.get_or_create_breaker(&format!("{channel}:{connector}"))
.await;
if !breaker.check() {
crate::metrics::record_circuit_breaker_rejection(connector, channel);
return Err(crate::errors::circuit_open_dataflow_error(
connector, channel,
));
}
let result = observed_handler_named(fn_name, connector, channel, fut).await;
match &result {
Ok(_) => breaker.record_success(),
Err(e) if e.retryable() => {
if breaker.record_failure() {
tracing::warn!(
connector = connector,
channel = channel,
"Circuit breaker tripped"
);
crate::metrics::record_circuit_breaker_trip(connector, channel);
}
}
Err(_) => {}
}
result
}
pub async fn observed_handler_named<F>(
fn_name: &'static str,
connector: &str,
channel: &str,
fut: F,
) -> dataflow_rs::Result<TaskOutcome>
where
F: std::future::Future<Output = dataflow_rs::Result<TaskOutcome>>,
{
let start = std::time::Instant::now();
let result = crate::engine::profile::record(fn_name, Some(connector), fut).await;
let status = if result.is_ok() { "ok" } else { "error" };
crate::metrics::record_connector_request(connector, channel, status);
crate::metrics::record_connector_duration(connector, channel, start.elapsed().as_secs_f64());
result
}
pub fn extract_output_path(input: &Value) -> &str {
input
.get("output")
.and_then(|v| v.as_str())
.unwrap_or("data")
}
pub fn to_exec_error(e: impl std::fmt::Display) -> DataflowError {
DataflowError::function_execution(e.to_string(), None)
}
pub fn to_connect_error(e: impl std::fmt::Display) -> DataflowError {
DataflowError::Io(e.to_string())
}
pub fn to_limit_error(message: impl std::fmt::Display) -> DataflowError {
DataflowError::Validation(message.to_string())
}
pub const LIMIT_MARKER: &str = "orion.limit: ";
pub fn require_str_field<'a>(
input: &'a Value,
field: &str,
handler_name: &str,
) -> Result<&'a str, DataflowError> {
input.get(field).and_then(|v| v.as_str()).ok_or_else(|| {
DataflowError::Validation(format!("{handler_name} requires '{field}' field"))
})
}
pub use crate::connector::is_mongo_url as is_mongo;
pub fn reject_mongo_connector(
function: &str,
connector_name: &str,
db_config: &crate::connector::DbConnectorConfig,
) -> Result<(), DataflowError> {
if is_mongo(&db_config.connection_string) {
return Err(DataflowError::Validation(format!(
"{function} requires a SQL connector, but '{connector_name}' is a MongoDB \
connector — use mongo_read or data_query for MongoDB"
)));
}
Ok(())
}
pub async fn resolve_connector(
registry: &ConnectorRegistry,
name: &str,
) -> Result<Arc<ConnectorConfig>, DataflowError> {
registry.get(name).await.ok_or_else(|| {
DataflowError::function_execution(format!("Connector '{name}' not found"), None)
})
}
pub fn require_db_connector<'a>(
config: &'a ConnectorConfig,
name: &str,
) -> Result<&'a DbConnectorConfig, DataflowError> {
match config {
ConnectorConfig::Db(c) => Ok(c),
_ => Err(crate::errors::connector_detail_error(format!(
"Connector '{name}' is not a database connector"
))),
}
}
pub fn require_http_connector<'a>(
config: &'a ConnectorConfig,
name: &str,
) -> Result<&'a crate::connector::HttpConnectorConfig, DataflowError> {
match config {
ConnectorConfig::Http(c) => Ok(c),
_ => Err(crate::errors::connector_detail_error(format!(
"Connector '{name}' is not an HTTP connector"
))),
}
}
pub fn require_kafka_connector<'a>(
config: &'a ConnectorConfig,
name: &str,
) -> Result<&'a crate::connector::KafkaConnectorConfig, DataflowError> {
match config {
ConnectorConfig::Kafka(c) => Ok(c),
_ => Err(crate::errors::connector_detail_error(format!(
"Connector '{name}' is not a Kafka connector"
))),
}
}
pub fn require_cache_connector<'a>(
config: &'a ConnectorConfig,
name: &str,
) -> Result<&'a CacheConnectorConfig, DataflowError> {
match config {
ConnectorConfig::Cache(c) => Ok(c),
_ => Err(crate::errors::connector_detail_error(format!(
"Connector '{name}' is not a cache connector"
))),
}
}
pub fn apply_output(ctx: &mut TaskContext<'_>, output_path: &str, new_value: Value) {
ctx.set_json(output_path, &new_value);
}
pub fn resolve_value(value: &Value, ctx: &TaskContext<'_>) -> Value {
match value {
Value::Object(o) => {
if o.len() == 1
&& let Some(spec) = o.get("var")
{
return resolve_var(spec, ctx);
}
Value::Object(
o.iter()
.map(|(k, v)| (k.clone(), resolve_value(v, ctx)))
.collect(),
)
}
Value::Array(a) => Value::Array(a.iter().map(|v| resolve_value(v, ctx)).collect()),
other => other.clone(),
}
}
fn resolve_var(spec: &Value, ctx: &TaskContext<'_>) -> Value {
let (path, default) = match spec {
Value::String(p) => (p.as_str(), Value::Null),
Value::Array(a) => match a.first().and_then(|v| v.as_str()) {
Some(p) => (p, a.get(1).cloned().unwrap_or(Value::Null)),
None => return Value::Null,
},
_ => return Value::Null,
};
ctx.get(path).map(Value::from).unwrap_or(default)
}
pub fn resolve_params(params: Option<&Value>, ctx: &TaskContext<'_>) -> Map<String, Value> {
match params.map(|p| resolve_value(p, ctx)) {
Some(Value::Object(map)) => map,
_ => Map::new(),
}
}
pub fn resolve_required_str(
input: &Value,
field: &str,
handler_name: &str,
ctx: &TaskContext<'_>,
) -> Result<String, DataflowError> {
let Some(raw) = input.get(field) else {
return Err(DataflowError::Validation(format!(
"{handler_name} requires '{field}' field"
)));
};
match resolve_value(raw, ctx) {
Value::String(s) => Ok(s),
Value::Number(n) => Ok(n.to_string()),
Value::Bool(b) => Ok(b.to_string()),
other => Err(DataflowError::Validation(format!(
"{handler_name} '{field}' must resolve to a string or number, got {}",
json_type_name(&other)
))),
}
}
pub fn resolve_bind_params(
input: &Value,
handler_name: &str,
ctx: &TaskContext<'_>,
) -> Result<Vec<Value>, DataflowError> {
match input.get("params") {
None | Some(Value::Null) => Ok(Vec::new()),
Some(raw) => match resolve_value(raw, ctx) {
Value::Array(a) => Ok(a),
other => Err(DataflowError::Validation(format!(
"{handler_name} 'params' must resolve to an array of bind values, got {}",
json_type_name(&other)
))),
},
}
}
pub fn json_type_name(v: &Value) -> &'static str {
match v {
Value::Null => "null",
Value::Bool(_) => "boolean",
Value::Number(_) => "number",
Value::String(_) => "string",
Value::Array(_) => "array",
Value::Object(_) => "object",
}
}
pub fn bind_json_params<'q>(
mut query: sqlx::query::Query<'q, sqlx::Any, sqlx::any::AnyArguments<'q>>,
params: &'q [Value],
) -> sqlx::query::Query<'q, sqlx::Any, sqlx::any::AnyArguments<'q>> {
for param in params {
query = match param {
Value::String(s) => query.bind(s.as_str()),
Value::Number(n) => {
if let Some(i) = n.as_i64() {
query.bind(i)
} else if let Some(f) = n.as_f64() {
query.bind(f)
} else {
query.bind(n.to_string())
}
}
Value::Bool(b) => query.bind(*b),
Value::Null => query.bind(None::<String>),
_ => query.bind(param.to_string()),
};
}
query
}
const DEFAULT_QUERY_TIMEOUT_MS: u64 = 30_000;
#[derive(Debug, Clone, Copy)]
pub struct QueryBudget {
deadline: tokio::time::Instant,
total_ms: u64,
}
impl QueryBudget {
pub fn start(timeout_ms: Option<u64>) -> Self {
let total_ms = timeout_ms.unwrap_or(DEFAULT_QUERY_TIMEOUT_MS);
Self {
deadline: tokio::time::Instant::now() + Duration::from_millis(total_ms),
total_ms,
}
}
pub async fn run<F, T, E>(&self, handler_name: &str, operation: F) -> Result<T, DataflowError>
where
F: std::future::Future<Output = Result<T, E>>,
E: std::fmt::Display,
{
let total_ms = self.total_ms;
tokio::time::timeout_at(self.deadline, operation)
.await
.map_err(|_| {
DataflowError::Timeout(format!("{handler_name} query timed out after {total_ms}ms"))
})?
.map_err(|e| {
let text = e.to_string();
if let Some(detail) = text.strip_prefix(LIMIT_MARKER) {
return to_limit_error(detail);
}
DataflowError::function_execution(
format!("{handler_name} query failed: {text}"),
None,
)
})
}
}
pub async fn timed_query<F, T, E>(
timeout_ms: Option<u64>,
handler_name: &str,
operation: F,
) -> Result<T, DataflowError>
where
F: std::future::Future<Output = Result<T, E>>,
E: std::fmt::Display,
{
QueryBudget::start(timeout_ms)
.run(handler_name, operation)
.await
}
#[cfg(test)]
mod tests {
use super::*;
use crate::connector::DialectGuards;
fn es_config(allow_private_urls: bool) -> EsConnectorConfig {
EsConnectorConfig {
max_response_size: 10 * 1024 * 1024,
url: "http://127.0.0.1:9200".to_string(),
auth: None,
request_timeout_ms: None,
allow_private_urls,
operations: OperationGates::default(),
dialect: DialectGuards::default(),
}
}
#[tokio::test]
async fn test_es_request_blocks_private_url() {
let client = reqwest::Client::new();
let result = es_request(
&client,
&es_config(false),
reqwest::Method::POST,
"http://127.0.0.1:9200/idx/_search",
)
.await;
let err = result.err().map(|e| e.to_string()).unwrap_or_default();
assert!(err.contains("SSRF protection"), "unexpected error: {err}");
}
#[tokio::test]
async fn test_es_request_allows_private_url_when_opted_in() {
let client = reqwest::Client::new();
let result = es_request(
&client,
&es_config(true),
reqwest::Method::POST,
"http://127.0.0.1:9200/idx/_search",
)
.await;
assert!(result.is_ok());
}
}
#[cfg(test)]
mod observability_tests {
const HANDLERS: [&str; 9] = [
"cache_read",
"cache_write",
"db_read",
"db_write",
"data_query",
"data_write",
"mongo_read",
"http_call",
"publish_kafka",
];
fn handler_source(handler: &str) -> String {
let dir = concat!(env!("CARGO_MANIFEST_DIR"), "/src/engine/functions");
std::fs::read_to_string(format!("{dir}/{handler}.rs")).expect("handler source")
}
#[test]
fn every_connector_handler_is_wrapped_in_the_observability_shell() {
let mut unwrapped = Vec::new();
for handler in HANDLERS {
let src = handler_source(handler);
if !["guarded_handler", "call.run("]
.iter()
.any(|w| src.contains(w))
{
unwrapped.push(handler);
}
assert!(
!src.contains("profile::record("),
"{handler} calls the raw profiler; go through ConnectorCall::run \
so its connector metrics are not conditional on the circuit breaker"
);
}
assert!(
unwrapped.is_empty(),
"these connector handlers emit no connector metrics: {unwrapped:?}"
);
}
#[test]
fn the_literal_prologue_precedes_message_dependent_resolution() {
const RESOLVERS: [&str; 4] = [
"resolve_required_str(",
"resolve_bind_params(",
"resolve_params(",
"resolve_value(",
];
for handler in HANDLERS {
let src = handler_source(handler);
let Some(begin) = src.find("ConnectorCall::begin(") else {
assert!(
["http_call", "publish_kafka"].contains(&handler),
"{handler} has no ConnectorCall prologue"
);
continue;
};
for resolver in RESOLVERS {
if let Some(at) = src.find(resolver) {
assert!(
begin < at,
"{handler}.rs calls {resolver} before ConnectorCall::begin, so a task \
missing 'connector' reports some other field first (proposal F58)"
);
}
}
}
}
#[test]
fn a_handler_names_itself_exactly_once() {
for handler in HANDLERS {
let src = handler_source(handler);
let quoted = format!("\"{handler}\"");
let occurrences = src.matches("ed).count();
assert_eq!(
occurrences, 1,
"{handler}.rs writes \"{handler}\" {occurrences} times; it should \
appear only in `const NAME` and be read back from there"
);
assert!(
src.contains(&format!("const NAME: &str = {quoted}")),
"{handler}.rs has no `const NAME` (proposal F48)"
);
}
}
}
#[cfg(test)]
mod error_taxonomy_tests {
use super::*;
#[test]
fn a_failure_to_connect_is_retryable_but_a_failed_query_is_not() {
assert!(
to_connect_error("connection refused").retryable(),
"an unreachable backend must be retryable, like the HTTP path"
);
assert!(
!to_exec_error("syntax error at or near \"SELCT\"").retryable(),
"a query the backend rejected is not worth retrying"
);
}
#[test]
fn a_limit_error_is_validation_not_execution() {
let err = to_limit_error("result exceeds query.max_limit — add a LIMIT");
assert!(
matches!(err, DataflowError::Validation(_)),
"expected Validation, got {err:?}"
);
assert!(!err.retryable(), "a limit does not fix itself on retry");
}
#[tokio::test]
async fn timed_query_classifies_a_marked_limit_and_strips_the_marker() {
let err = timed_query(Some(1_000), "db_read", async {
Err::<(), String>(format!("{LIMIT_MARKER}too many rows — add a LIMIT"))
})
.await
.expect_err("the operation failed");
assert!(
matches!(err, DataflowError::Validation(ref m) if m == "too many rows — add a LIMIT"),
"expected a stripped Validation, got {err:?}"
);
}
#[tokio::test]
async fn timed_query_leaves_an_ordinary_failure_as_execution() {
let err = timed_query(Some(1_000), "db_read", async {
Err::<(), String>("connection reset".to_string())
})
.await
.expect_err("the operation failed");
assert!(
matches!(err, DataflowError::FunctionExecution { .. }),
"expected FunctionExecution, got {err:?}"
);
}
}