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, QueryBudget, QueryFailure, acquire_conn, decode_failure, encode_failure,
reject_mongo_connector, require_op_allowed, resolve_bind_params, resolve_row_format,
to_connect_error,
};
use super::schema::{FieldKind, FieldSchema};
use super::templated_input::TemplatedInput;
use crate::connector::ConnectorRegistry;
use crate::connector::pool_cache::SqlPoolCache;
use crate::engine::HandlerError;
const NAME: &str = <DbReadHandler as ConnectorHandler>::NAME;
pub struct DbReadHandler {
pub pool_cache: Arc<SqlPoolCache>,
pub registry: Arc<ConnectorRegistry>,
pub max_rows: usize,
}
pub struct DbRead {
query: String,
params: Vec<Value>,
format: crate::connector::sql_decode::RowFormat,
}
impl DbRead {
pub(super) fn parse_statement(
call: &ConnectorCall<'_>,
input: &TemplatedInput,
ctx: &TaskContext<'_>,
) -> Result<Self, HandlerError> {
Ok(Self {
query: call.require_str(input, "query")?.to_string(),
params: resolve_bind_params(input, call.name, ctx)?,
format: crate::connector::sql_decode::RowFormat::default(),
})
}
pub(super) fn parse_read(
call: &ConnectorCall<'_>,
input: &TemplatedInput,
ctx: &TaskContext<'_>,
) -> Result<Self, HandlerError> {
let read = Self {
format: resolve_row_format(input, call.name, ctx)?,
..Self::parse_statement(call, input, ctx)?
};
require_read_only(&read.query, call.name)?;
Ok(read)
}
pub(super) fn query(&self) -> &str {
&self.query
}
pub(super) fn params(&self) -> &[Value] {
&self.params
}
pub(super) fn format(&self) -> crate::connector::sql_decode::RowFormat {
self.format
}
}
#[async_trait]
impl ConnectorHandler for DbReadHandler {
const NAME: &'static str = "db_read";
type Kind = crate::connector::kind::Db;
type Input = TemplatedInput;
type Parsed = DbRead;
fn registry(&self) -> &Arc<ConnectorRegistry> {
&self.registry
}
fn parse(
&self,
call: &ConnectorCall<'_>,
input: &TemplatedInput,
ctx: &TaskContext<'_>,
) -> Result<Self::Parsed, HandlerError> {
DbRead::parse_read(call, input, ctx)
}
fn gate(
_parsed: &Self::Parsed,
conn: &crate::connector::DbConnectorConfig,
connector: &str,
) -> Result<(), HandlerError> {
require_op_allowed(&conn.operations, "read", connector)?;
Ok(reject_mongo_connector(
<Self as ConnectorHandler>::NAME,
connector,
conn,
)?)
}
async fn run(
&self,
read: Self::Parsed,
db_config: &crate::connector::DbConnectorConfig,
call: &ConnectorCall<'_>,
_input: &TemplatedInput,
_ctx: &mut TaskContext<'_>,
) -> Result<Produced, HandlerError> {
let pool = self
.pool_cache
.get_pool(call.connector, db_config)
.await
.map_err(to_connect_error)?;
let max_rows = self.max_rows;
let params = read.params();
let format = read.format();
let query = read.query();
let budget = QueryBudget::start(db_config.query_timeout_ms);
let scalars: Vec<crate::connector::sql_encode::Scalar> =
params.iter().map(Into::into).collect();
let json = crate::connector::pool_cache::dispatch_sql_pool!(
&pool, p, rows_to_json, bind, typed_args, _write_result => {
let mut conn = acquire_conn(&budget, call.name, p).await?;
let bound = budget
.run(call.name, async {
typed_args(&mut conn, query, Some(&scalars))
.await
.map_err(|e| QueryFailure::Classified(encode_failure(NAME, e)))
})
.await?;
let rows = budget.run(call.name, async {
use futures::TryStreamExt;
let sqlx_query = match bound {
crate::connector::sql_encode::Bound::Typed(args) => {
sqlx::query_with(sqlx::AssertSqlSafe(query), args)
}
crate::connector::sql_encode::Bound::Fallback { cache } => {
bind(sqlx::query(sqlx::AssertSqlSafe(query)), params).persistent(cache)
}
};
let mut stream = sqlx_query.fetch(&mut *conn);
let mut rows = Vec::new();
while let Some(row) = stream.try_next().await? {
if rows.len() >= max_rows {
return Err(
crate::engine::functions::connector_helpers::QueryFailure::Limit(
format!(
"{} result exceeds query.max_limit ({max_rows} rows) \
— add a LIMIT to the query or raise the cap",
call.name
),
),
);
}
rows.push(row);
}
Ok(rows)
})
.await?;
rows_to_json(&rows, format).map_err(|e| decode_failure(NAME, e))?
}
);
Ok(Value::Array(json).into())
}
}
pub(super) const DB_READ_FIELDS: &[FieldSchema] = &[
FieldSchema {
name: "connector",
description: "Name of the SQL connector to query.",
kind: FieldKind::String,
required: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "query",
description: "Read statement — SELECT, WITH, VALUES or TABLE; a write belongs in db_write, which has its own 'raw_write' connector gate. Bind placeholders are the backend's own spelling: ? for SQLite and MySQL, $1, $2, ... for PostgreSQL.",
kind: FieldKind::String,
required: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "params",
description: "Array of values to bind to query placeholders, in order. Accepts {\"var\": \"path\"} to read the value from the message.",
kind: FieldKind::Array,
resolvable: true,
..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.",
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.",
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
},
];
fn require_read_only(query: &str, handler_name: &str) -> Result<(), HandlerError> {
use crate::sql_lex::{READ_STATEMENTS, ReadOnlyViolation};
let message = match crate::sql_lex::read_only_violation(query) {
None => return Ok(()),
Some(ReadOnlyViolation::Empty) => format!("{handler_name} 'query' has no statement to run"),
Some(ReadOnlyViolation::NotARead { keyword }) => format!(
"{handler_name} runs read statements only, but this one starts with \
'{keyword}' — use db_write for INSERT/UPDATE/DELETE (it has its own \
'raw_write' connector gate). Reads start with {}",
READ_STATEMENTS.join(", ")
),
Some(ReadOnlyViolation::ModifyingCte { keyword }) => format!(
"{handler_name} runs read statements only, but this one carries a \
data-modifying '{keyword}' common table expression — use db_write \
(it has its own 'raw_write' connector gate)"
),
};
Err(DataflowError::Validation(message).into())
}
#[cfg(test)]
mod tests {
use super::*;
fn check(sql: &str) -> Result<(), String> {
require_read_only(sql, "db_read").map_err(|e| {
let e: DataflowError = e.into();
e.to_string()
})
}
#[test]
fn reads_are_admitted() {
for sql in [
"SELECT id FROM users WHERE id = $1",
" \n select 1",
"-- a comment\nSELECT 1",
"/* block */ SELECT 1",
"(SELECT 1) UNION (SELECT 2)",
"WITH recent AS (SELECT * FROM orders) SELECT * FROM recent",
"VALUES (1), (2)",
"TABLE users",
"SELECT id FROM jobs ORDER BY id FOR UPDATE SKIP LOCKED",
"SELECT id FROM notes WHERE body = 'delete from users'",
"SELECT \"delete\" FROM t",
"SELECT total AS deleted FROM t",
"SELECT CAST(a AS text) FROM t",
] {
assert!(
check(sql).is_ok(),
"must be admitted: {sql} — {:?}",
check(sql)
);
}
}
#[test]
fn writes_are_refused() {
for sql in [
"DELETE FROM audit_log WHERE id > 0 RETURNING id",
"delete from audit_log",
"INSERT INTO t (a) VALUES (1)",
"UPDATE t SET a = 1",
"TRUNCATE t",
"DROP TABLE t",
"PRAGMA journal_mode = WAL",
"EXPLAIN ANALYZE DELETE FROM t",
" -- lead in\n DELETE FROM t",
] {
let err = check(sql).expect_err(&format!("must be refused: {sql}"));
assert!(err.contains("read statements only"), "{sql}: {err}");
}
}
#[test]
fn a_data_modifying_cte_is_refused() {
for sql in [
"WITH gone AS (DELETE FROM t RETURNING id) SELECT * FROM gone",
"WITH added AS (INSERT INTO t (a) VALUES (1) RETURNING id) SELECT * FROM added",
"with m as materialized (update t set a = 1 returning id) select * from m",
] {
let err = check(sql).expect_err(&format!("must be refused: {sql}"));
assert!(err.contains("data-modifying"), "{sql}: {err}");
}
}
#[test]
fn quoted_text_is_not_syntax() {
assert!(check("SELECT 1 /* AS (DELETE */").is_ok());
assert!(check("SELECT 'x AS (DELETE FROM t)' AS s").is_ok());
assert!(check("SELECT $tag$ AS (DELETE FROM t) $tag$ AS s").is_ok());
assert!(check("SELECT * FROM t WHERE a = $1 AND b = $2").is_ok());
}
#[test]
fn an_empty_statement_is_refused() {
let err = check(" -- nothing here\n").expect_err("empty");
assert!(err.contains("no statement"), "{err}");
}
}