use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use async_trait::async_trait;
use dataflow_rs::engine::error::DataflowError;
use dataflow_rs::engine::task_context::TaskContext;
use futures::TryStreamExt;
use mongodb::bson::Document;
use serde_json::{Map, Value};
use sqlx::AssertSqlSafe;
use super::connector_handler::{ConnectorHandler, Produced};
use super::connector_helpers::{
ConnectorCall, QueryBudget, QueryFailure, acquire_conn, build_entity_registry, decode_failure,
encode_failure, es_request, is_mongo, require_op_allowed, resolve_params, resolve_row_format,
timed_query, to_connect_error, to_exec_error,
};
use super::schema::{FieldKind, FieldSchema};
use super::templated_input::TemplatedInput;
use crate::connector::mongo_pool::MongoPoolCache;
use crate::connector::pool_cache::SqlPoolCache;
use crate::connector::{ConnectorConfig, ConnectorRegistry, EsConnectorConfig};
use crate::engine::HandlerError;
use crate::query::{self, GroupKey, SqlDialect};
use crate::storage::detect_backend;
const NAME: &str = <DataQueryHandler as ConnectorHandler>::NAME;
const MAX_INCLUDE_KEYS_PER_QUERY: usize = 1000;
pub struct DataQueryHandler {
pub pool_cache: Arc<SqlPoolCache>,
pub mongo_pool_cache: Arc<MongoPoolCache>,
pub http_client: reqwest::Client,
pub registry: Arc<ConnectorRegistry>,
pub limits: crate::config::QueryConfig,
}
pub struct DataQuery {
query: Value,
params: Map<String, Value>,
schema: Option<Value>,
format: crate::connector::sql_decode::RowFormat,
database: Option<String>,
}
#[async_trait]
impl ConnectorHandler for DataQueryHandler {
const NAME: &'static str = "data_query";
type Kind = crate::connector::DataBackend;
type Input = TemplatedInput;
type Parsed = DataQuery;
fn registry(&self) -> &Arc<ConnectorRegistry> {
&self.registry
}
fn parse(
&self,
call: &ConnectorCall<'_>,
input: &TemplatedInput,
ctx: &TaskContext<'_>,
) -> Result<Self::Parsed, HandlerError> {
let query = input
.get("query")
.ok_or_else(|| {
DataflowError::Validation(format!("{} requires 'query' field", call.name))
})?
.clone();
Ok(DataQuery {
query,
params: resolve_params(input, <Self as ConnectorHandler>::NAME, ctx),
format: resolve_row_format(input, call.name, ctx)?,
schema: input.get("schema").cloned(),
database: input
.get("database")
.and_then(Value::as_str)
.map(str::to_string),
})
}
fn gate(
_parsed: &Self::Parsed,
conn: &ConnectorConfig,
connector: &str,
) -> Result<(), HandlerError> {
if let Some(gates) = conn.operation_gates() {
require_op_allowed(gates, "read", connector)?;
}
Ok(())
}
async fn run(
&self,
parsed: Self::Parsed,
conn: &ConnectorConfig,
call: &ConnectorCall<'_>,
_input: &TemplatedInput,
_ctx: &mut TaskContext<'_>,
) -> Result<Produced, HandlerError> {
let registry = build_entity_registry(parsed.schema.as_ref(), conn, call.connector)?;
let query = &parsed.query;
let params = &parsed.params;
let result = match conn {
ConnectorConfig::Es(es) => {
let eq = query::translate_es(query, params, ®istry, &self.limits)
.map_err(DataflowError::from)?;
if eq.count {
run_es_count(&self.http_client, es, &eq).await?
} else {
run_es_search(&self.http_client, es, &eq).await?
}
}
ConnectorConfig::Db(db) if is_mongo(&db.connection_string) => {
let database = parsed.database.as_deref().ok_or_else(|| {
DataflowError::Validation(format!("{} requires 'database' field", call.name))
})?;
let mq = query::translate_mongo(query, params, ®istry, &self.limits)
.map_err(DataflowError::from)?;
let client = self
.mongo_pool_cache
.get_client(call.connector, db)
.await
.map_err(to_connect_error)?;
let coll = client
.database(database)
.collection::<Document>(&mq.collection);
if mq.count {
let n = timed_query(db.query_timeout_ms, call.name, async {
coll.count_documents(mq.filter)
.await
.map_err(|e| e.to_string())
})
.await?;
return Ok(count_result(n).into());
}
let docs: Vec<Document> = timed_query(db.query_timeout_ms, call.name, async {
let mut find = coll.find(mq.filter);
if let Some(p) = mq.projection {
find = find.projection(p);
}
if let Some(s) = mq.sort {
find = find.sort(s);
}
if let Some(sk) = mq.skip {
find = find.skip(sk);
}
find = find.limit(mq.limit as i64);
let cursor = find.await.map_err(|e| e.to_string())?;
cursor.try_collect().await.map_err(|e| e.to_string())
})
.await?;
super::mongo_common::docs_to_json(&docs, NAME)?
}
ConnectorConfig::Db(db) => {
let dialect: SqlDialect = detect_backend(&db.connection_string)
.map_err(to_exec_error)?
.into();
let plan = query::plan_sql(query, params, ®istry, dialect, &self.limits)
.map_err(DataflowError::from)?;
let pool = self
.pool_cache
.get_pool(call.connector, db)
.await
.map_err(to_connect_error)?;
if plan.count {
return Ok(run_sql_count(&pool, &plan, dialect, db.query_timeout_ms)
.await?
.into());
}
run_sql_with_includes(&pool, &plan, dialect, db.query_timeout_ms, parsed.format)
.await?
}
other => unreachable!(
"DataBackend admitted a '{}' connector",
other.connector_type()
),
};
Ok(result.into())
}
}
fn count_result(n: u64) -> Value {
serde_json::json!({ "count": n })
}
async fn run_sql_count(
pool: &crate::connector::pool_cache::SqlPool,
plan: &query::SqlPlan,
dialect: SqlDialect,
timeout_ms: Option<u64>,
) -> Result<Value, DataflowError> {
let (sql, values) = query::backend::sql::build_for(dialect, &plan.main);
let budget = QueryBudget::start(timeout_ms);
let rows: Vec<Value> = crate::connector::pool_cache::dispatch_sql_pool!(
pool, p, rows_to_json, _bind, typed_args, _write_result => {
let scalars = crate::connector::sql_encode::scalars_from_sea(&values.0);
let mut conn = acquire_conn(&budget, NAME, p).await?;
let bound = budget
.run(NAME, async {
typed_args(&mut conn, &sql, scalars.as_deref())
.await
.map_err(|e| QueryFailure::Classified(encode_failure(NAME, e)))
})
.await?;
let fallback = crate::connector::sql_encode::sea_args_for(p, values);
let q = match bound {
crate::connector::sql_encode::Bound::Typed(args) => {
sqlx::query_with(AssertSqlSafe(sql.as_str()), args)
}
crate::connector::sql_encode::Bound::Fallback { cache } => {
sqlx::query_with(AssertSqlSafe(sql.as_str()), fallback)
.persistent(cache)
}
};
let rows = budget.run(NAME, q.fetch_all(&mut *conn)).await?;
rows_to_json(&rows, crate::connector::sql_decode::RowFormat::default())
.map_err(|e| decode_failure(NAME, e))?
}
);
let n = rows
.first()
.and_then(|r| r.get(query::backend::sql::COUNT_COLUMN))
.and_then(Value::as_u64)
.ok_or_else(|| {
DataflowError::function_execution(
format!("{NAME}: the count query returned no count"),
None,
)
})?;
Ok(count_result(n))
}
async fn run_es_count(
client: &reqwest::Client,
es: &EsConnectorConfig,
eq: &query::backend::es::EsQuery,
) -> Result<Value, DataflowError> {
let url = format!("{}/{}/_count", es.url.trim_end_matches('/'), eq.index);
let req = es_request(client, es, reqwest::Method::POST, &url)
.await?
.json(&eq.body);
let (status, body) = super::connector_helpers::send_es(req, es.max_response_size).await?;
if !status.is_success() {
return Err(DataflowError::function_execution(
format!("Elasticsearch count failed ({status}): {body}"),
None,
));
}
let n = body.get("count").and_then(Value::as_u64).ok_or_else(|| {
DataflowError::function_execution(
format!("Elasticsearch count returned no count: {body}"),
None,
)
})?;
Ok(count_result(n))
}
async fn run_es_search(
client: &reqwest::Client,
es: &EsConnectorConfig,
eq: &query::backend::es::EsQuery,
) -> Result<Value, DataflowError> {
let url = format!("{}/{}/_search", es.url.trim_end_matches('/'), eq.index);
let req = es_request(client, es, reqwest::Method::POST, &url)
.await?
.json(&eq.body);
let (status, body) = super::connector_helpers::send_es(req, es.max_response_size).await?;
if !status.is_success() {
return Err(DataflowError::function_execution(
format!("Elasticsearch search failed ({status}): {body}"),
None,
));
}
let docs: Vec<Value> = body
.get("hits")
.and_then(|h| h.get("hits"))
.and_then(|h| h.as_array())
.map(|hits| {
hits.iter()
.map(|h| {
let mut source = h.get("_source").cloned().unwrap_or(Value::Null);
if eq.include_id
&& let Some(id) = h.get(query::backend::es::ES_DOCUMENT_KEY)
&& let Value::Object(map) = &mut source
{
map.insert(query::backend::es::ES_DOCUMENT_KEY.to_string(), id.clone());
}
source
})
.collect()
})
.unwrap_or_default();
Ok(Value::Array(docs))
}
async fn run_sql_with_includes(
pool: &crate::connector::pool_cache::SqlPool,
plan: &query::SqlPlan,
dialect: SqlDialect,
timeout_ms: Option<u64>,
format: crate::connector::sql_decode::RowFormat,
) -> Result<Value, DataflowError> {
let budget = QueryBudget::start(timeout_ms);
let (sql, values) = query::backend::sql::build_for(dialect, &plan.main);
let mut parents: Vec<Value> = crate::connector::pool_cache::dispatch_sql_pool!(
pool, p, rows_to_json, _bind, typed_args, _write_result => {
let scalars = crate::connector::sql_encode::scalars_from_sea(&values.0);
let mut conn = acquire_conn(&budget, NAME, p).await?;
let bound = budget
.run(NAME, async {
typed_args(&mut conn, &sql, scalars.as_deref())
.await
.map_err(|e| QueryFailure::Classified(encode_failure(NAME, e)))
})
.await?;
let fallback = crate::connector::sql_encode::sea_args_for(p, values);
let q = match bound {
crate::connector::sql_encode::Bound::Typed(args) => {
sqlx::query_with(AssertSqlSafe(sql.as_str()), args)
}
crate::connector::sql_encode::Bound::Fallback { cache } => {
sqlx::query_with(AssertSqlSafe(sql.as_str()), fallback)
.persistent(cache)
}
};
let rows = budget.run(NAME, q.fetch_all(&mut *conn)).await?;
rows_to_json(&rows, format).map_err(|e| decode_failure(NAME, e))?
}
);
for inc in &plan.includes {
let mut seen = HashSet::new();
let mut keys = Vec::new();
for p in &parents {
if let Some(k) = p.get(&inc.local)
&& let Some(gk) = GroupKey::from_json(k)
&& let Some(sv) = query::backend::sql::json_key_to_sea(k)
&& seen.insert(gk)
{
keys.push(sv);
}
}
let strip = inc.strip();
let mut groups: HashMap<GroupKey, Vec<Value>> = HashMap::new();
for chunk in keys.chunks(MAX_INCLUDE_KEYS_PER_QUERY) {
let (csql, cvalues) = query::backend::sql::build_include_select(inc, chunk, dialect);
let children: Vec<Value> = crate::connector::pool_cache::dispatch_sql_pool!(
pool, p, rows_to_json, _bind, typed_args, _write_result => {
let cscalars = crate::connector::sql_encode::scalars_from_sea(&cvalues.0);
let mut conn = acquire_conn(&budget, NAME, p).await?;
let bound = budget
.run(NAME, async {
typed_args(&mut conn, &csql, cscalars.as_deref())
.await
.map_err(|e| QueryFailure::Classified(encode_failure(NAME, e)))
})
.await?;
let fallback = crate::connector::sql_encode::sea_args_for(p, cvalues);
let q = match bound {
crate::connector::sql_encode::Bound::Typed(args) => {
sqlx::query_with(AssertSqlSafe(csql.as_str()), args)
}
crate::connector::sql_encode::Bound::Fallback { cache } => {
sqlx::query_with(AssertSqlSafe(csql.as_str()), fallback)
.persistent(cache)
}
};
let crows = budget.run(NAME, q.fetch_all(&mut *conn)).await?;
rows_to_json(&crows, format).map_err(|e| decode_failure(NAME, e))?
}
);
for mut child in children {
let Some(fk) = child.get(&inc.foreign).and_then(GroupKey::from_json) else {
continue;
};
if let Value::Object(m) = &mut child {
m.remove(query::backend::sql::INCLUDE_RANK_COLUMN);
for s in &strip {
m.remove(s);
}
}
groups.entry(fk).or_default().push(child);
}
}
for p in &mut parents {
let kids = p
.get(&inc.local)
.and_then(GroupKey::from_json)
.and_then(|k| groups.get(&k).cloned())
.unwrap_or_default();
if let Value::Object(m) = p {
m.insert(inc.field.clone(), Value::Array(kids));
}
}
}
if !plan.strip.is_empty() {
for p in &mut parents {
if let Value::Object(m) = p {
for s in &plan.strip {
m.remove(s);
}
}
}
}
Ok(Value::Array(parents))
}
pub(super) const DATA_QUERY_FIELDS: &[FieldSchema] = &[
FieldSchema {
name: "connector",
description: "Name of the db (SQL/MongoDB) or es (Elasticsearch) connector to query.",
kind: FieldKind::String,
required: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "query",
description: "Backend-neutral query envelope: \
source/filter/fields/sort/limit/skip/after/include/count. \
An include selection is {fields, sort, limit}; `sort` is required because \
the per-parent page is cut in the database. \"after\" is a keyset cursor — \
the previous page's last row, one value per sort key — which pages without \
an offset and so is not bounded by query.max_skip. \"count\": true answers \
{\"count\": n} instead of rows, and refuses the keys that shape a row set.",
kind: FieldKind::Object,
required: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "params",
description: "Object of named values folded into the filter's {\"param\": ..} nodes. \
A value of {\"var\": \"path\"} is read from the message context.",
kind: FieldKind::Object,
resolvable: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "schema",
description: "Inline entity schema (renames, type hints, allowlist, relations) enabling \
some/all/none and typed coercion. Undeclared entities and columns are \
rejected, so a query without one reaches nothing; pass \
{\"unmapped\": \"identity\"} for pre-1.0 pass-through.",
kind: FieldKind::Object,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "database",
description: "MongoDB database name. Optional here because the same task shape is \
valid against SQL and Elasticsearch, which need no database key; \
required — and checked at workflow activation — once the referenced \
connector is a MongoDB one (F52).",
kind: FieldKind::String,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "numeric_as",
description: "How an arbitrary-precision decimal column is rendered: \"number\" (default) or \"string\". A number is computable in JSONLogic and rounds beyond 2^53 or on most decimal fractions; a string keeps every digit, which is what a money column needs. SQL backends only.",
kind: FieldKind::String,
resolvable: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "binary_as",
description: "How a binary column is rendered: \"auto\" (default), \"hex\", \"base64\" or \"text\". Auto reads the bytes as text when they are valid UTF-8 and as hex when they are not, so its result shape depends on the data; name an encoding for a column that is genuinely binary. SQL backends only.",
kind: FieldKind::String,
resolvable: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "output",
description: "Dotted path in the message where rows are written. Defaults to \"data\".",
kind: FieldKind::String,
template_at: &[""],
..FieldSchema::DEFAULT
},
];