use pgwire::api::results::Response;
use pgwire::error::{ErrorInfo, PgWireError, PgWireResult};
use crate::control::security::identity::AuthenticatedIdentity;
use crate::control::server::shared::planning_overrides::parse_bool_session_value;
use super::super::types::sqlstate_error;
use super::core::NodeDbPgHandler;
enum TransactionCmd {
SetReadOnly,
SetReadWrite,
AcceptIsolation,
RejectIsolation(String),
}
fn classify_transaction_cmd(upper: &str, sql: &str) -> TransactionCmd {
if upper.contains("ISOLATION LEVEL") {
if upper.contains("READ COMMITTED") {
return TransactionCmd::AcceptIsolation;
}
let level = if upper.contains("SERIALIZABLE") {
Some("SERIALIZABLE")
} else if upper.contains("REPEATABLE READ") {
Some("REPEATABLE READ")
} else if upper.contains("READ UNCOMMITTED") {
Some("READ UNCOMMITTED")
} else {
None
};
let message = match level {
Some(lvl) => format!(
"SET TRANSACTION ISOLATION LEVEL {lvl} is not supported; \
NodeDB enforces Snapshot Isolation"
),
None => format!(
"unsupported SET TRANSACTION option: {}",
sql.split_whitespace().skip(2).collect::<Vec<_>>().join(" ")
),
};
return TransactionCmd::RejectIsolation(message);
}
if upper.contains("READ ONLY") {
return TransactionCmd::SetReadOnly;
}
if upper.contains("READ WRITE") {
return TransactionCmd::SetReadWrite;
}
TransactionCmd::RejectIsolation(format!(
"unsupported SET TRANSACTION option: {}",
sql.split_whitespace().skip(2).collect::<Vec<_>>().join(" ")
))
}
impl NodeDbPgHandler {
pub(super) fn handle_set(
&self,
identity: &AuthenticatedIdentity,
addr: &std::net::SocketAddr,
sql: &str,
) -> PgWireResult<Vec<Response>> {
use crate::control::server::shared::session::parse_set_command;
use pgwire::api::results::Tag;
let upper = sql.to_uppercase();
if upper.starts_with("SET TRANSACTION") || upper.starts_with("SET SESSION CHARACTERISTICS")
{
match classify_transaction_cmd(&upper, sql) {
TransactionCmd::SetReadOnly => {
self.sessions.set_parameter(
addr,
"transaction_access_mode".into(),
"read_only".into(),
);
return Ok(vec![Response::Execution(Tag::new("SET"))]);
}
TransactionCmd::SetReadWrite => {
self.sessions.set_parameter(
addr,
"transaction_access_mode".into(),
"read_write".into(),
);
return Ok(vec![Response::Execution(Tag::new("SET"))]);
}
TransactionCmd::AcceptIsolation => {
return Ok(vec![Response::Execution(Tag::new("SET"))]);
}
TransactionCmd::RejectIsolation(message) => {
return Err(sqlstate_error(
nodedb_types::error::sqlstate::FEATURE_NOT_SUPPORTED,
&message,
));
}
}
}
if upper.starts_with("SET ROLE ") || upper == "SET ROLE" {
return Err(sqlstate_error(
nodedb_types::error::sqlstate::FEATURE_NOT_SUPPORTED,
"SET ROLE is not supported: a session's role set is identity-bound \
at CREATE USER time. Use GRANT/REVOKE ROLE TO <user> to change \
a user's roles, or reconnect with a different user.",
));
}
if upper.starts_with("SET SESSION AUTHORIZATION") {
return Err(sqlstate_error(
nodedb_types::error::sqlstate::FEATURE_NOT_SUPPORTED,
"SET SESSION AUTHORIZATION is not supported: identity is bound at \
connection time. Reconnect as the target user.",
));
}
let (key, value) = match parse_set_command(sql) {
Some(kv) => kv,
None => {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"42601".to_owned(),
format!("syntax error in SET command: {sql}"),
))));
}
};
match key.as_str() {
"tenant" => {
return self.handle_set_tenant_name_or_id(identity, addr, &value);
}
"nodedb.tenant_id" => {
return self.handle_set_tenant_by_id(identity, addr, &value);
}
"role" => {
return Err(sqlstate_error(
nodedb_types::error::sqlstate::FEATURE_NOT_SUPPORTED,
"SET ROLE is not supported: a session's role set is identity-bound \
at CREATE USER time. Use GRANT/REVOKE ROLE TO <user> to change \
a user's roles, or reconnect with a different user.",
));
}
"session_authorization" => {
return Err(sqlstate_error(
nodedb_types::error::sqlstate::FEATURE_NOT_SUPPORTED,
"SET SESSION AUTHORIZATION is not supported: identity is bound at \
connection time. Reconnect as the target user.",
));
}
_ => {}
}
if key == "nodedb.consistency" {
match value.as_str() {
"strong" | "eventual" => {}
s if s.starts_with("bounded_staleness") => {}
_ => {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!(
"invalid value for nodedb.consistency: '{value}'. Valid: strong, bounded_staleness(<ms>), eventual"
),
))));
}
}
}
if key == crate::control::server::shared::session::read_consistency::PARAM_KEY
&& crate::control::server::shared::session::read_consistency::parse_value(&value)
.is_none()
{
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!(
"invalid value for {}: '{value}'. Valid: strong, bounded_staleness:<secs>, eventual",
crate::control::server::shared::session::read_consistency::PARAM_KEY
),
))));
}
if key == crate::control::server::shared::session::cross_shard_mode::PARAM_KEY
&& crate::control::server::shared::session::cross_shard_mode::parse_value(&value)
.is_none()
{
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!(
"invalid value for {}: '{value}'. Valid values: 'strict', 'best_effort_non_atomic'",
crate::control::server::shared::session::cross_shard_mode::PARAM_KEY
),
))));
}
if key == "nodedb.force_shuffle_join" && parse_bool_session_value(&value).is_none() {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!(
"invalid value for nodedb.force_shuffle_join: '{value}'. \
Valid: on, off, true, false, 1, 0"
),
))));
}
if key == "nodedb.shuffle_num_parts" && value.parse::<u32>().is_err() {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!(
"invalid value for nodedb.shuffle_num_parts: '{value}'. \
Must be a non-negative integer (0 = cluster default)"
),
))));
}
if key == "nodedb.force_shuffle_agg" && parse_bool_session_value(&value).is_none() {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!(
"invalid value for nodedb.force_shuffle_agg: '{value}'. \
Valid: on, off, true, false, 1, 0"
),
))));
}
if key == "nodedb.shuffle_agg_num_parts" && value.parse::<u32>().is_err() {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!(
"invalid value for nodedb.shuffle_agg_num_parts: '{value}'. \
Must be a non-negative integer (0 = cluster default)"
),
))));
}
if key == "nodedb.broadcast_threshold_bytes" && value.parse::<usize>().is_err() {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!(
"invalid value for nodedb.broadcast_threshold_bytes: '{value}'. \
Must be a non-negative integer (bytes; 0 = always shuffle \
when both sides are analyzed)"
),
))));
}
if key == "nodedb.shuffle_agg_threshold" && value.parse::<usize>().is_err() {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!(
"invalid value for nodedb.shuffle_agg_threshold: '{value}'. \
Must be a non-negative integer (distinct-group count; the GROUP \
BY is auto-shuffled when its estimated group cardinality exceeds \
this value)"
),
))));
}
if key == "nodedb.auth_session" {
use crate::control::security::session_handle::{ClientFingerprint, ResolveOutcome};
let caller_fp = ClientFingerprint::from_peer(identity.tenant_id, addr);
let conn_key = addr.to_string();
match self
.state
.session_handles
.resolve(&value, &conn_key, &caller_fp)
{
ResolveOutcome::RateLimited => {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"FATAL".to_owned(),
"53300".to_owned(),
"session handle resolve rate limit exceeded on this \
connection — closing"
.to_owned(),
))));
}
ResolveOutcome::Resolved(_) | ResolveOutcome::Miss => {
}
}
}
if !crate::control::server::shared::session::is_known_settable_runtime_parameter(&key) {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"42704".to_owned(),
format!("unrecognized configuration parameter \"{key}\""),
))));
}
self.sessions.set_parameter(addr, key, value);
Ok(vec![Response::Execution(Tag::new("SET"))])
}
fn apply_tenant_override(
&self,
identity: &AuthenticatedIdentity,
addr: &std::net::SocketAddr,
new_tenant: Option<crate::types::TenantId>,
source: &str,
) -> PgWireResult<Vec<Response>> {
use crate::control::security::audit::AuditEvent;
use pgwire::api::results::Tag;
if !identity.is_superuser {
return Err(sqlstate_error(
"42501",
"only superuser may change session tenant; a regular user's \
tenant is identity-bound at CREATE USER time",
));
}
if self.sessions.transaction_state(addr)
!= crate::control::server::shared::session::TransactionState::Idle
{
return Err(sqlstate_error(
"25001",
"cannot change session tenant inside an active transaction \
(COMMIT or ROLLBACK first)",
));
}
let prior = self.sessions.get_effective_tenant_id(addr);
self.sessions.set_effective_tenant_id(addr, new_tenant);
let detail = match new_tenant {
Some(t) => format!(
"{source}: tenant switched from {} to {}",
prior.unwrap_or(identity.tenant_id),
t
),
None => format!(
"{source}: tenant reset to identity-bound {}",
identity.tenant_id
),
};
self.state.audit_record(
AuditEvent::PrivilegeChange,
Some(identity.tenant_id),
&identity.username,
&detail,
);
Ok(vec![Response::Execution(Tag::new("SET"))])
}
pub(super) fn handle_set_tenant_name_or_id(
&self,
identity: &AuthenticatedIdentity,
addr: &std::net::SocketAddr,
value: &str,
) -> PgWireResult<Vec<Response>> {
if value.eq_ignore_ascii_case("default") {
return self.apply_tenant_override(identity, addr, None, "SET TENANT = DEFAULT");
}
let resolved = if let Ok(id) = value.parse::<u64>() {
crate::types::TenantId::new(id)
} else {
let catalog = self.state.credentials.catalog();
let stored = catalog
.find_tenant_by_name(value)
.map_err(|e| sqlstate_error("XX000", &format!("catalog read: {e}")))?
.ok_or_else(|| sqlstate_error("42704", &format!("tenant '{value}' not found")))?;
crate::types::TenantId::new(stored.tenant_id)
};
self.apply_tenant_override(identity, addr, Some(resolved), "SET TENANT")
}
pub(super) fn handle_set_tenant_by_id(
&self,
identity: &AuthenticatedIdentity,
addr: &std::net::SocketAddr,
value: &str,
) -> PgWireResult<Vec<Response>> {
if value.eq_ignore_ascii_case("default") {
return self.apply_tenant_override(
identity,
addr,
None,
"SET nodedb.tenant_id = DEFAULT",
);
}
let id: u64 = value.parse().map_err(|_| {
PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22023".to_owned(),
format!("invalid value for nodedb.tenant_id: '{value}'. Must be an integer."),
)))
})?;
self.apply_tenant_override(
identity,
addr,
Some(crate::types::TenantId::new(id)),
"SET nodedb.tenant_id",
)
}
pub(crate) fn handle_reset_tenant(
&self,
identity: &AuthenticatedIdentity,
addr: &std::net::SocketAddr,
) -> PgWireResult<Vec<Response>> {
use pgwire::api::results::Tag;
if !identity.is_superuser {
return Err(sqlstate_error("42501", "only superuser may RESET TENANT"));
}
if self.sessions.transaction_state(addr)
!= crate::control::server::shared::session::TransactionState::Idle
{
return Err(sqlstate_error(
"25001",
"cannot RESET TENANT inside an active transaction",
));
}
self.sessions.set_effective_tenant_id(addr, None);
Ok(vec![Response::Execution(Tag::new("RESET"))])
}
}
#[cfg(test)]
mod tests {
use super::{TransactionCmd, classify_transaction_cmd};
#[test]
fn tenant_id_above_u32_max_parses_as_u64() {
let big = "4294967296"; assert!(big.parse::<u64>().is_ok(), "should parse as u64");
assert!(big.parse::<u32>().is_err(), "should NOT parse as u32");
}
fn run(sql: &str) -> TransactionCmd {
let upper = sql.to_uppercase();
classify_transaction_cmd(&upper, sql)
}
fn is_accept(cmd: TransactionCmd) -> bool {
matches!(
cmd,
TransactionCmd::SetReadOnly
| TransactionCmd::SetReadWrite
| TransactionCmd::AcceptIsolation
)
}
fn rejection_code(cmd: TransactionCmd) -> Option<String> {
match cmd {
TransactionCmd::RejectIsolation(msg) => Some(msg),
_ => None,
}
}
#[test]
fn set_transaction_read_only() {
assert!(is_accept(run("SET TRANSACTION READ ONLY")));
assert!(matches!(
run("SET TRANSACTION READ ONLY"),
TransactionCmd::SetReadOnly
));
}
#[test]
fn set_transaction_read_write() {
assert!(matches!(
run("SET TRANSACTION READ WRITE"),
TransactionCmd::SetReadWrite
));
}
#[test]
fn set_transaction_read_committed() {
assert!(matches!(
run("SET TRANSACTION ISOLATION LEVEL READ COMMITTED"),
TransactionCmd::AcceptIsolation
));
}
#[test]
fn set_transaction_serializable() {
let msg = rejection_code(run("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE"))
.expect("expected rejection");
assert!(
msg.contains("SERIALIZABLE"),
"message should name the level: {msg}"
);
assert!(
msg.contains("Snapshot Isolation"),
"message should mention Snapshot Isolation: {msg}"
);
}
#[test]
fn set_transaction_repeatable_read() {
let msg = rejection_code(run("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ"))
.expect("expected rejection");
assert!(msg.contains("REPEATABLE READ"), "{msg}");
}
#[test]
fn set_transaction_read_uncommitted() {
let msg = rejection_code(run("SET TRANSACTION ISOLATION LEVEL READ UNCOMMITTED"))
.expect("expected rejection");
assert!(msg.contains("READ UNCOMMITTED"), "{msg}");
}
#[test]
fn set_session_characteristics_serializable() {
let sql = "SET SESSION CHARACTERISTICS AS TRANSACTION ISOLATION LEVEL SERIALIZABLE";
let msg = rejection_code(run(sql)).expect("expected rejection");
assert!(msg.contains("SERIALIZABLE"), "{msg}");
}
#[test]
fn set_transaction_unknown_option() {
let msg = rejection_code(run("SET TRANSACTION DEFERRABLE"))
.expect("expected rejection for unknown option");
assert!(
msg.contains("unsupported"),
"message should say unsupported: {msg}"
);
}
}