fix(s3select): enforce query and resource limits (#5028)

* fix(s3select): enforce query and resource limits

* fix(s3select): close query resource limit gaps

* fix(s3select): preserve timeout and stream invariants

* fix(s3select): enforce staged query limits

* fix(s3select): preserve policy error compatibility

* fix(s3select): bound error source traversal
This commit is contained in:
GatewayJ
2026-07-25 18:44:53 +08:00
committed by GitHub
parent 2dc4d0b651
commit 0364523dad
13 changed files with 3228 additions and 222 deletions
File diff suppressed because it is too large Load Diff
+104 -5
View File
@@ -12,18 +12,26 @@
// See the License for the specific language governing permissions and
// limitations under the License.
use std::sync::Arc;
use std::{
sync::{Arc, LazyLock},
time::Duration,
};
use async_trait::async_trait;
use derive_builder::Builder;
use rustfs_s3select_api::{
QueryResult,
query::{
Query, dispatcher::QueryDispatcher, execution::QueryStateMachineRef, logical_planner::Plan, session::SessionCtxFactory,
Query,
dispatcher::QueryDispatcher,
execution::QueryStateMachineRef,
logical_planner::Plan,
session::{DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES as DEFAULT_MEMORY_LIMIT_BYTES, SessionCtxFactory},
},
server::dbms::{DatabaseManagerSystem, QueryHandle},
};
use s3s::dto::SelectObjectContentInput;
use tokio::sync::Semaphore;
use crate::{
dispatcher::manager::SimpleQueryDispatcherBuilder,
@@ -34,6 +42,16 @@ use crate::{
};
const ENV_RUSTFS_S3SELECT_TARGET_PARTITIONS: &str = "RUSTFS_S3SELECT_TARGET_PARTITIONS";
const ENV_RUSTFS_S3SELECT_MEMORY_LIMIT_BYTES: &str = "RUSTFS_S3SELECT_MEMORY_LIMIT_BYTES";
const ENV_RUSTFS_S3SELECT_QUERY_TIMEOUT_SECS: &str = "RUSTFS_S3SELECT_QUERY_TIMEOUT_SECS";
const ENV_RUSTFS_S3SELECT_MAX_CONCURRENT_QUERIES: &str = "RUSTFS_S3SELECT_MAX_CONCURRENT_QUERIES";
pub(crate) const DEFAULT_QUERY_TIMEOUT_SECS: u64 = 300;
pub(crate) const DEFAULT_MAX_CONCURRENT_QUERIES: usize = 4;
const MAX_QUERY_TIMEOUT_SECS: u64 = 24 * 60 * 60;
const TEST_MAX_CONCURRENT_QUERIES: usize = 1024;
static QUERY_ADMISSION: LazyLock<Arc<Semaphore>> =
LazyLock::new(|| Arc::new(Semaphore::new(S3SelectRuntimeConfig::from_env().max_concurrent_queries)));
#[derive(Builder)]
pub struct RustFSms<D: QueryDispatcher> {
@@ -79,9 +97,23 @@ where
}
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
struct S3SelectRuntimeConfig {
target_partitions: usize,
memory_limit_bytes: usize,
query_timeout: Duration,
max_concurrent_queries: usize,
}
impl Default for S3SelectRuntimeConfig {
fn default() -> Self {
Self {
target_partitions: 0,
memory_limit_bytes: DEFAULT_MEMORY_LIMIT_BYTES,
query_timeout: Duration::from_secs(DEFAULT_QUERY_TIMEOUT_SECS),
max_concurrent_queries: DEFAULT_MAX_CONCURRENT_QUERIES,
}
}
}
impl S3SelectRuntimeConfig {
@@ -90,14 +122,55 @@ impl S3SelectRuntimeConfig {
target_partitions: target_partitions_from_env_value(
std::env::var(ENV_RUSTFS_S3SELECT_TARGET_PARTITIONS).ok().as_deref(),
),
memory_limit_bytes: bounded_usize_from_env_value(
std::env::var(ENV_RUSTFS_S3SELECT_MEMORY_LIMIT_BYTES).ok().as_deref(),
DEFAULT_MEMORY_LIMIT_BYTES,
usize::MAX,
),
query_timeout: s3_select_query_timeout(),
max_concurrent_queries: bounded_usize_from_env_value(
std::env::var(ENV_RUSTFS_S3SELECT_MAX_CONCURRENT_QUERIES).ok().as_deref(),
DEFAULT_MAX_CONCURRENT_QUERIES,
Semaphore::MAX_PERMITS,
),
}
}
}
pub fn s3_select_query_timeout() -> Duration {
Duration::from_secs(bounded_u64_from_env_value(
std::env::var(ENV_RUSTFS_S3SELECT_QUERY_TIMEOUT_SECS).ok().as_deref(),
DEFAULT_QUERY_TIMEOUT_SECS,
MAX_QUERY_TIMEOUT_SECS,
))
}
fn target_partitions_from_env_value(value: Option<&str>) -> usize {
value.and_then(|value| value.parse::<usize>().ok()).unwrap_or(0)
}
fn bounded_usize_from_env_value(value: Option<&str>, default: usize, max: usize) -> usize {
value
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| (1..=max).contains(value))
.unwrap_or(default)
}
fn bounded_u64_from_env_value(value: Option<&str>, default: u64, max: u64) -> u64 {
value
.and_then(|value| value.parse::<u64>().ok())
.filter(|value| (1..=max).contains(value))
.unwrap_or(default)
}
fn query_admission(is_test: bool) -> Arc<Semaphore> {
if is_test {
Arc::new(Semaphore::new(TEST_MAX_CONCURRENT_QUERIES))
} else {
Arc::clone(&QUERY_ADMISSION)
}
}
pub async fn make_rustfsms(input: Arc<SelectObjectContentInput>, is_test: bool) -> QueryResult<impl DatabaseManagerSystem> {
// init Function Manager, we can define some UDF if need
let func_manager = SimpleFunctionMetadataManager::default();
@@ -116,8 +189,11 @@ pub async fn make_rustfsms(input: Arc<SelectObjectContentInput>, is_test: bool)
.with_func_manager(Arc::new(func_manager))
.with_default_table_provider(default_table_provider)
.with_session_factory(session_factory)
.with_memory_limit_bytes(runtime_config.memory_limit_bytes)
.with_parser(parser)
.with_query_execution_factory(query_execution_factory)
.with_query_admission(query_admission(is_test))
.with_query_timeout(runtime_config.query_timeout)
.build()?;
let mut builder = RustFSmsBuilder::default();
@@ -143,8 +219,11 @@ pub async fn make_rustfsms_with_components(
.with_func_manager(func_manager)
.with_default_table_provider(default_table_provider)
.with_session_factory(session_factory)
.with_memory_limit_bytes(runtime_config.memory_limit_bytes)
.with_parser(parser)
.with_query_execution_factory(query_execution_factory)
.with_query_admission(query_admission(is_test))
.with_query_timeout(runtime_config.query_timeout)
.build()?;
let mut builder = RustFSmsBuilder::default();
@@ -167,7 +246,10 @@ mod tests {
use crate::get_global_db;
use super::{S3SelectRuntimeConfig, target_partitions_from_env_value};
use super::{
DEFAULT_MAX_CONCURRENT_QUERIES, DEFAULT_MEMORY_LIMIT_BYTES, DEFAULT_QUERY_TIMEOUT_SECS, MAX_QUERY_TIMEOUT_SECS,
S3SelectRuntimeConfig, bounded_u64_from_env_value, bounded_usize_from_env_value, target_partitions_from_env_value,
};
#[test]
fn parses_target_partitions_from_env_value() {
@@ -179,7 +261,24 @@ mod tests {
#[test]
fn default_runtime_config_uses_datafusion_default_partitions() {
assert_eq!(S3SelectRuntimeConfig::default().target_partitions, 0);
let config = S3SelectRuntimeConfig::default();
assert_eq!(config.target_partitions, 0);
assert_eq!(config.memory_limit_bytes, DEFAULT_MEMORY_LIMIT_BYTES);
assert_eq!(config.query_timeout.as_secs(), DEFAULT_QUERY_TIMEOUT_SECS);
assert_eq!(config.max_concurrent_queries, DEFAULT_MAX_CONCURRENT_QUERIES);
}
#[test]
fn resource_limits_reject_invalid_and_out_of_range_values() {
assert_eq!(bounded_usize_from_env_value(Some("1024"), 64, 2048), 1024);
assert_eq!(bounded_usize_from_env_value(Some("0"), 64, 2048), 64);
assert_eq!(bounded_usize_from_env_value(Some("4096"), 64, 2048), 64);
assert_eq!(bounded_usize_from_env_value(Some("invalid"), 64, 2048), 64);
assert_eq!(bounded_u64_from_env_value(Some("30"), 300, MAX_QUERY_TIMEOUT_SECS), 30);
assert_eq!(bounded_u64_from_env_value(Some("0"), 300, MAX_QUERY_TIMEOUT_SECS), 300);
assert_eq!(bounded_u64_from_env_value(Some("86401"), 300, MAX_QUERY_TIMEOUT_SECS), 300);
assert_eq!(bounded_u64_from_env_value(None, 300, MAX_QUERY_TIMEOUT_SECS), 300);
}
#[tokio::test]
+5 -1
View File
@@ -19,7 +19,7 @@ use datafusion::arrow::datatypes::DataType;
use datafusion::common::{Result as DFResult, TableReference};
use datafusion::datasource::TableProvider;
use datafusion::logical_expr::var_provider::is_system_variables;
use datafusion::logical_expr::{AggregateUDF, HigherOrderUDF, ScalarUDF, TableSource, WindowUDF};
use datafusion::logical_expr::{AggregateUDF, HigherOrderUDF, ScalarUDF, TableSource, WindowUDF, planner::ExprPlanner};
use datafusion::variable::VarType;
use datafusion::{config::ConfigOptions, sql::planner::ContextProvider};
use rustfs_s3select_api::query::{function::FuncMetaManagerRef, session::SessionCtx};
@@ -81,6 +81,10 @@ impl ContextProviderExtension for MetadataProvider {
}
impl ContextProvider for MetadataProvider {
fn get_expr_planners(&self) -> &[Arc<dyn ExprPlanner>] {
self.session.inner().expr_planners()
}
fn get_function_meta(&self, name: &str) -> Option<Arc<ScalarUDF>> {
self.func_manager
.udf(name)
+263 -2
View File
@@ -12,11 +12,18 @@
// See the License for the specific language governing permissions and
// limitations under the License.
use std::ops::ControlFlow;
use async_recursion::async_recursion;
use async_trait::async_trait;
use datafusion::sql::{planner::SqlToRel, sqlparser::ast::Statement};
use datafusion::sql::{
planner::SqlToRel,
sqlparser::ast::{
GroupByExpr, ObjectNamePart, OrderByKind, Query, Select, SelectFlavor, SetExpr, Statement, TableFactor, Visit, Visitor,
},
};
use rustfs_s3select_api::{
QueryError, QueryResult,
QueryError, QueryResult, S3SelectPolicyError,
query::{
ast::ExtStatement,
logical_planner::{LogicalPlanner, Plan, QueryPlan},
@@ -60,6 +67,7 @@ impl<'a, S: ContextProviderExtension + Send + Sync + 'a> SqlPlanner<'a, S> {
async fn df_sql_to_plan(&self, stmt: Statement, _session: &SessionCtx) -> QueryResult<Plan> {
match stmt {
Statement::Query(_) => {
validate_s3_select_statement(&stmt)?;
let df_plan = self.df_planner.sql_statement_to_plan(stmt)?;
let plan = Plan::Query(QueryPlan {
df_plan,
@@ -72,3 +80,256 @@ impl<'a, S: ContextProviderExtension + Send + Sync + 'a> SqlPlanner<'a, S> {
}
}
}
fn validate_s3_select_statement(statement: &Statement) -> QueryResult<()> {
let Statement::Query(query) = statement else {
return Err(unsupported_structure("only SELECT queries are supported"));
};
if query.with.is_some()
|| query.order_by.as_ref().is_some_and(|order_by| {
order_by.interpolate.is_some()
|| match &order_by.kind {
OrderByKind::Expressions(expressions) => expressions.iter().any(|expression| expression.with_fill.is_some()),
OrderByKind::All(_) => true,
}
})
|| query.fetch.is_some()
|| !query.locks.is_empty()
|| query.for_clause.is_some()
|| query.settings.is_some()
|| query.format_clause.is_some()
|| !query.pipe_operators.is_empty()
{
return Err(unsupported_structure("the query contains an unsupported clause"));
}
if let Some(limit_clause) = query.limit_clause.as_ref()
&& !matches!(
limit_clause,
datafusion::sql::sqlparser::ast::LimitClause::LimitOffset {
limit: Some(_),
offset: None,
limit_by,
} if limit_by.is_empty()
)
{
return Err(unsupported_structure("only LIMIT without OFFSET is supported"));
}
if let Some(datafusion::sql::sqlparser::ast::LimitClause::LimitOffset { limit: Some(limit), .. }) =
query.limit_clause.as_ref()
&& limit.to_string().parse::<u64>().is_err()
{
return Err(unsupported_structure("LIMIT must be a non-negative integer"));
}
let mut detector = SubqueryDetector { visited_root: false };
if query.visit(&mut detector).is_break() {
return Err(unsupported_structure("subqueries are not supported"));
}
let SetExpr::Select(select) = query.body.as_ref() else {
return Err(unsupported_structure("set operations and nested queries are not supported"));
};
validate_select(select)
}
fn validate_select(select: &Select) -> QueryResult<()> {
if !select.optimizer_hints.is_empty()
|| select.distinct.is_some()
|| select.select_modifiers.is_some()
|| select.top.is_some()
|| select.exclude.is_some()
|| select.into.is_some()
|| !select.lateral_views.is_empty()
|| select.prewhere.is_some()
|| !select.connect_by.is_empty()
|| !select.cluster_by.is_empty()
|| !select.distribute_by.is_empty()
|| !select.sort_by.is_empty()
|| select.having.is_some()
|| !select.named_window.is_empty()
|| select.qualify.is_some()
|| select.value_table_mode.is_some()
|| select.flavor != SelectFlavor::Standard
|| !matches!(&select.group_by, GroupByExpr::Expressions(_, modifiers) if modifiers.is_empty())
{
return Err(unsupported_structure("the SELECT contains an unsupported clause"));
}
let [table] = select.from.as_slice() else {
return Err(unsupported_structure("exactly one S3Object source is required"));
};
if !table.joins.is_empty() {
return Err(unsupported_structure("JOIN is not supported"));
}
let TableFactor::Table {
name,
alias,
args,
with_hints,
version,
with_ordinality,
partitions,
sample,
index_hints,
..
} = &table.relation
else {
return Err(unsupported_structure("subqueries and table functions are not supported"));
};
if args.is_some()
|| !with_hints.is_empty()
|| version.is_some()
|| *with_ordinality
|| !partitions.is_empty()
|| sample.is_some()
|| !index_hints.is_empty()
|| alias.as_ref().is_some_and(|alias| !alias.columns.is_empty())
{
return Err(unsupported_structure("the S3Object source contains unsupported modifiers"));
}
let ([ObjectNamePart::Identifier(table_name)] | [ObjectNamePart::Identifier(table_name), ObjectNamePart::Identifier(_)]) =
name.0.as_slice()
else {
return Err(unsupported_structure("the source must be S3Object"));
};
let is_s3_object = if table_name.quote_style.is_some() {
table_name.value == "S3Object"
} else {
table_name.value.eq_ignore_ascii_case("S3Object")
};
if !is_s3_object {
return Err(unsupported_structure("the source must be S3Object"));
}
Ok(())
}
fn unsupported_structure(message: &str) -> QueryError {
S3SelectPolicyError::UnsupportedSqlStructure {
message: message.to_string(),
}
.into()
}
struct SubqueryDetector {
visited_root: bool,
}
impl Visitor for SubqueryDetector {
type Break = ();
fn pre_visit_query(&mut self, _query: &Query) -> ControlFlow<Self::Break> {
if self.visited_root {
ControlFlow::Break(())
} else {
self.visited_root = true;
ControlFlow::Continue(())
}
}
}
#[cfg(test)]
mod tests {
use super::validate_s3_select_statement;
use crate::sql::parser::ExtParser;
use datafusion::sql::sqlparser::ast::Statement;
use rustfs_s3select_api::{S3SelectPolicyError, query::ast::ExtStatement};
fn parse_statement(sql: &str) -> Statement {
let mut statements = ExtParser::parse_sql(sql).expect("SQL should parse");
let ExtStatement::SqlStatement(statement) = statements.pop_front().expect("one SQL statement");
*statement
}
#[test]
fn accepts_s3_select_query_shape() {
let statement = parse_statement("SELECT s.id FROM S3Object AS s WHERE s.id = '1' LIMIT 10");
assert!(validate_s3_select_statement(&statement).is_ok());
}
#[test]
fn accepts_json_sub_path_source() {
let statement = parse_statement("SELECT e.name FROM S3Object.employees AS e");
assert!(validate_s3_select_statement(&statement).is_ok());
}
#[test]
fn accepts_group_by_and_order_by() {
let statement = parse_statement("SELECT department, COUNT(*) FROM S3Object GROUP BY department ORDER BY department");
assert!(validate_s3_select_statement(&statement).is_ok());
}
#[test]
fn rejects_join() {
let statement = parse_statement("SELECT * FROM S3Object a JOIN S3Object b ON a.id = b.id");
assert!(matches!(
validate_s3_select_statement(&statement),
Err(ref err) if matches!(
err.s3_select_policy_error(),
Some(S3SelectPolicyError::UnsupportedSqlStructure { message }) if message == "JOIN is not supported"
)
));
}
#[test]
fn rejects_subquery() {
let statement = parse_statement("SELECT * FROM S3Object WHERE id IN (SELECT id FROM S3Object)");
assert!(matches!(
validate_s3_select_statement(&statement),
Err(ref err) if matches!(
err.s3_select_policy_error(),
Some(S3SelectPolicyError::UnsupportedSqlStructure { message }) if message == "subqueries are not supported"
)
));
}
#[test]
fn rejects_subquery_in_order_by() {
let statement = parse_statement("SELECT id FROM S3Object ORDER BY (SELECT id FROM S3Object)");
assert!(matches!(
validate_s3_select_statement(&statement),
Err(ref err) if matches!(
err.s3_select_policy_error(),
Some(S3SelectPolicyError::UnsupportedSqlStructure { message }) if message == "subqueries are not supported"
)
));
}
#[test]
fn rejects_non_s3_object_source() {
let statement = parse_statement("SELECT * FROM other_table");
assert!(matches!(
validate_s3_select_statement(&statement),
Err(ref err) if matches!(
err.s3_select_policy_error(),
Some(S3SelectPolicyError::UnsupportedSqlStructure { message }) if message == "the source must be S3Object"
)
));
}
#[test]
fn rejects_unsupported_select_clauses() {
for sql in [
"SELECT DISTINCT id FROM S3Object",
"SELECT * FROM S3Object OFFSET 1",
"SELECT * FROM S3Object UNION SELECT * FROM S3Object",
] {
let statement = parse_statement(sql);
assert!(
matches!(
validate_s3_select_statement(&statement),
Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::UnsupportedSqlStructure { .. }))
),
"query should be rejected: {sql}"
);
}
}
}
@@ -15,7 +15,10 @@
#[cfg(test)]
mod integration_tests {
use crate::{create_fresh_db, get_global_db, instance::make_rustfsms};
use datafusion::arrow::array::{Array, StringArray};
use datafusion::arrow::{
array::{Array, Int64Array, StringArray},
record_batch::RecordBatch,
};
use rustfs_s3select_api::{
QueryError,
query::{Context, Query},
@@ -26,6 +29,52 @@ mod integration_tests {
};
use std::sync::Arc;
fn assert_ages_descending(output: &[RecordBatch]) {
let ages: Vec<i64> = output
.iter()
.flat_map(|batch| {
batch
.column(1)
.as_any()
.downcast_ref::<Int64Array>()
.expect("age column should be Int64")
.values()
.iter()
.copied()
.collect::<Vec<_>>()
})
.collect();
assert_eq!(ages, vec![40, 38, 35, 32, 30, 28, 26, 25, 24, 22]);
}
fn assert_department_counts(output: &[RecordBatch]) {
let mut counts: Vec<(&str, i64)> = output
.iter()
.flat_map(|batch| {
let departments = batch
.column(0)
.as_any()
.downcast_ref::<StringArray>()
.expect("department column should be Utf8");
let counts = batch
.column(1)
.as_any()
.downcast_ref::<Int64Array>()
.expect("count column should be Int64");
departments.iter().zip(counts.iter()).map(|(department, count)| {
(
department.expect("department should not be null"),
count.expect("count should not be null"),
)
})
})
.collect();
counts.sort_unstable();
assert_eq!(counts, vec![("Finance", 3), ("HR", 2), ("IT", 3), ("Marketing", 2)]);
}
fn create_test_input(sql: &str) -> SelectObjectContentInput {
SelectObjectContentInput {
bucket: "test-bucket".to_string(),
@@ -292,21 +341,21 @@ mod integration_tests {
#[tokio::test]
async fn test_select_with_aggregation() {
let sql = "SELECT department, COUNT(*) as count FROM S3Object GROUP BY department";
let sql = "SELECT department, COUNT(*) FROM S3Object GROUP BY department";
let input = create_test_input(sql);
let db = get_global_db(input.clone(), true).await.unwrap();
let query = Query::new(Context { input: Arc::new(input) }, sql.to_string());
let result = db.execute(&query).await;
// Aggregation queries might fail due to lack of actual data, which is acceptable
match result {
Ok(_) => {
// If successful, that's great
}
Err(_) => {
// Expected to fail due to no actual data source
}
}
let output = db
.execute(&query)
.await
.expect("execute grouped CSV query")
.result()
.chunk_result()
.await
.expect("collect grouped CSV output");
assert_department_counts(&output);
}
#[tokio::test]
@@ -381,13 +430,22 @@ mod integration_tests {
#[tokio::test]
async fn test_query_with_order_by() {
let sql = "SELECT name, age FROM S3Object ORDER BY age DESC";
let sql = "SELECT name, CAST(age AS BIGINT) AS age FROM S3Object ORDER BY CAST(age AS BIGINT) DESC";
let input = create_test_input(sql);
let db = get_global_db(input.clone(), true).await.unwrap();
let query = Query::new(Context { input: Arc::new(input) }, sql.to_string());
let result = db.execute(&query).await;
assert!(result.is_ok());
let output = db
.execute(&query)
.await
.expect("execute ordered CSV query")
.result()
.chunk_result()
.await
.expect("collect ordered CSV output");
assert_eq!(output.iter().map(|batch| batch.num_rows()).sum::<usize>(), 10);
assert_ages_descending(&output);
}
#[tokio::test]
@@ -582,21 +640,21 @@ mod integration_tests {
#[tokio::test]
async fn test_select_with_aggregation_json() {
let sql = "SELECT department, COUNT(*) as count FROM S3Object GROUP BY department";
let sql = "SELECT department, COUNT(*) FROM S3Object GROUP BY department";
let input = create_test_json_input(sql);
let db = get_global_db(input.clone(), true).await.unwrap();
let query = Query::new(Context { input: Arc::new(input) }, sql.to_string());
let result = db.execute(&query).await;
// Aggregation queries may fail due to lack of actual data, which is acceptable
match result {
Ok(_) => {
// If successful, that's great
}
Err(_) => {
// Expected to fail due to no actual data source
}
}
let output = db
.execute(&query)
.await
.expect("execute grouped JSON query")
.result()
.chunk_result()
.await
.expect("collect grouped JSON output");
assert_department_counts(&output);
}
#[tokio::test]
@@ -672,8 +730,17 @@ mod integration_tests {
let db = get_global_db(input.clone(), true).await.unwrap();
let query = Query::new(Context { input: Arc::new(input) }, sql.to_string());
let result = db.execute(&query).await;
assert!(result.is_ok());
let output = db
.execute(&query)
.await
.expect("execute ordered JSON query")
.result()
.chunk_result()
.await
.expect("collect ordered JSON output");
assert_eq!(output.iter().map(|batch| batch.num_rows()).sum::<usize>(), 10);
assert_ages_descending(&output);
}
#[tokio::test]