use core::fmt;
use core::str::FromStr;
use mongreldb_types::ids::{DatabaseId, QueryId};
use crate::prepared::StatementId;
#[repr(transparent)]
#[derive(
Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)]
pub struct SessionId(pub [u8; 16]);
impl SessionId {
pub const ZERO: Self = Self([0u8; 16]);
pub const fn from_bytes(bytes: [u8; 16]) -> Self {
Self(bytes)
}
pub const fn as_bytes(&self) -> &[u8; 16] {
&self.0
}
pub fn to_hex(self) -> String {
let mut out = String::with_capacity(32);
for byte in self.0 {
out.push(char::from_digit((byte >> 4) as u32, 16).expect("nibble"));
out.push(char::from_digit((byte & 0x0f) as u32, 16).expect("nibble"));
}
out
}
}
impl fmt::Display for SessionId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.to_hex())
}
}
impl fmt::Debug for SessionId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "SessionId({})", self.to_hex())
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum SessionIdParseError {
#[error("invalid session id length: expected 32 hex digits, got {0} chars")]
InvalidLength(usize),
#[error("invalid hex character `{0}` in session id")]
InvalidCharacter(char),
}
impl FromStr for SessionId {
type Err = SessionIdParseError;
fn from_str(text: &str) -> Result<Self, Self::Err> {
let compact: String = text.chars().filter(|c| *c != '-').collect();
if compact.chars().count() != 32 {
return Err(SessionIdParseError::InvalidLength(compact.chars().count()));
}
let mut bytes = [0u8; 16];
let mut chars = compact.chars();
for byte in &mut bytes {
let hi = chars.next().expect("length checked above");
let lo = chars.next().expect("length checked above");
let hi = hi
.to_digit(16)
.ok_or(SessionIdParseError::InvalidCharacter(hi))?;
let lo = lo
.to_digit(16)
.ok_or(SessionIdParseError::InvalidCharacter(lo))?;
*byte = ((hi << 4) | lo) as u8;
}
Ok(Self(bytes))
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum AuthenticatedIdentity {
Credentialless,
CatalogUser {
username: String,
user_id: u64,
created_version: u64,
},
ServicePrincipal {
name: String,
},
ExternalPrincipal {
provider: String,
subject: String,
username: String,
user_id: u64,
created_version: u64,
scopes: Vec<String>,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub enum IsolationLevel {
#[default]
Snapshot,
ReadCommitted,
Serializable,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub enum ParameterValue {
Null,
Bool(bool),
Integer(i64),
Float(f64),
Text(String),
Bytes(Vec<u8>),
}
impl ParameterValue {
pub fn type_name(&self) -> &'static str {
match self {
Self::Null => "NULL",
Self::Bool(_) => "BOOL",
Self::Integer(_) => "INT64",
Self::Float(_) => "FLOAT64",
Self::Text(_) => "TEXT",
Self::Bytes(_) => "BYTES",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum AdminCommand {
CreateUser {
username: String,
},
DropUser {
username: String,
},
GrantRole {
username: String,
role: String,
},
RevokeRole {
username: String,
role: String,
},
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub enum ExecuteCommand {
Sql {
text: String,
params: Vec<ParameterValue>,
},
ExecutePrepared {
statement_id: StatementId,
params: Vec<ParameterValue>,
},
Begin {
isolation: IsolationLevel,
},
Commit,
Rollback,
Cancel {
query_id: QueryId,
},
GetSchema {
table: String,
},
Admin(AdminCommand),
}
pub const DEFAULT_MAX_ROWS: u64 = 100_000;
pub const DEFAULT_MAX_BYTES: u64 = 64 * 1024 * 1024;
pub const DEFAULT_MAX_CANDIDATE_COUNT: u64 = 1_000_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub struct ResultLimits {
pub max_rows: Option<u64>,
pub max_bytes: Option<u64>,
pub max_candidate_count: Option<u64>,
}
impl ResultLimits {
pub fn effective_max_rows(&self) -> u64 {
self.max_rows.unwrap_or(DEFAULT_MAX_ROWS)
}
pub fn effective_max_bytes(&self) -> u64 {
self.max_bytes.unwrap_or(DEFAULT_MAX_BYTES)
}
pub fn effective_max_candidate_count(&self) -> u64 {
self.max_candidate_count
.unwrap_or(DEFAULT_MAX_CANDIDATE_COUNT)
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct ExecuteRequest {
pub request_id: [u8; 16],
pub query_id: QueryId,
pub session_id: Option<SessionId>,
pub database_id: DatabaseId,
pub principal: AuthenticatedIdentity,
pub command: ExecuteCommand,
pub deadline_unix_micros: Option<u64>,
pub result_limits: ResultLimits,
pub resource_group: Option<String>,
pub idempotency_key: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::assert_serde_round_trip;
fn sample_principal() -> AuthenticatedIdentity {
AuthenticatedIdentity::CatalogUser {
username: "alice".to_owned(),
user_id: 42,
created_version: 7,
}
}
fn sample_params() -> Vec<ParameterValue> {
vec![
ParameterValue::Null,
ParameterValue::Bool(true),
ParameterValue::Integer(-42),
ParameterValue::Float(2.5),
ParameterValue::Text("hello".to_owned()),
ParameterValue::Bytes(vec![0xde, 0xad, 0xbe, 0xef]),
]
}
#[test]
fn session_id_text_form_round_trips() {
let bytes = [
0x00, 0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0xfe, 0xdc, 0xba, 0x98, 0x76,
0x54, 0x32,
];
let id = SessionId::from_bytes(bytes);
assert_eq!(id.to_hex(), "000123456789abcdeffedcba98765432");
assert_eq!(id.to_string(), id.to_hex());
assert_eq!(
format!("{id:?}"),
"SessionId(000123456789abcdeffedcba98765432)"
);
assert_eq!(id.to_hex().parse::<SessionId>().unwrap(), id);
assert_eq!(
"00012345-6789-abcd-effe-dcba98765432"
.parse::<SessionId>()
.unwrap(),
id,
"hyphenated UUID form parses"
);
assert_eq!(
id.to_hex()
.to_ascii_uppercase()
.parse::<SessionId>()
.unwrap(),
id,
"uppercase hex parses"
);
assert_eq!(id.as_bytes(), &bytes);
assert_eq!(SessionId::ZERO.as_bytes(), &[0u8; 16]);
}
#[test]
fn session_id_parse_rejects_bad_input() {
assert_eq!(
"".parse::<SessionId>(),
Err(SessionIdParseError::InvalidLength(0))
);
assert_eq!(
"abcd".parse::<SessionId>(),
Err(SessionIdParseError::InvalidLength(4))
);
assert_eq!(
"g".repeat(32).parse::<SessionId>(),
Err(SessionIdParseError::InvalidCharacter('g'))
);
}
#[test]
fn session_id_serde_round_trip() {
assert_serde_round_trip(&SessionId::from_bytes([0x5a; 16]));
assert_serde_round_trip(&SessionId::ZERO);
}
#[test]
fn authenticated_identity_serde_round_trip() {
assert_serde_round_trip(&AuthenticatedIdentity::Credentialless);
assert_serde_round_trip(&sample_principal());
assert_serde_round_trip(&AuthenticatedIdentity::ServicePrincipal {
name: "replication".to_owned(),
});
}
#[test]
fn isolation_level_serde_round_trip_and_default() {
assert_eq!(IsolationLevel::default(), IsolationLevel::Snapshot);
for level in [
IsolationLevel::Snapshot,
IsolationLevel::ReadCommitted,
IsolationLevel::Serializable,
] {
assert_serde_round_trip(&level);
}
}
#[test]
fn parameter_value_serde_round_trip_and_type_names() {
for value in sample_params() {
assert_serde_round_trip(&value);
}
let names: Vec<&'static str> = sample_params()
.iter()
.map(ParameterValue::type_name)
.collect();
assert_eq!(names, ["NULL", "BOOL", "INT64", "FLOAT64", "TEXT", "BYTES"]);
}
#[test]
fn admin_command_serde_round_trip() {
for command in [
AdminCommand::CreateUser {
username: "alice".to_owned(),
},
AdminCommand::DropUser {
username: "alice".to_owned(),
},
AdminCommand::GrantRole {
username: "alice".to_owned(),
role: "analyst".to_owned(),
},
AdminCommand::RevokeRole {
username: "alice".to_owned(),
role: "analyst".to_owned(),
},
] {
assert_serde_round_trip(&command);
}
}
#[test]
fn execute_command_serde_round_trip_every_variant() {
for command in [
ExecuteCommand::Sql {
text: "SELECT * FROM t WHERE a = ?".to_owned(),
params: sample_params(),
},
ExecuteCommand::ExecutePrepared {
statement_id: StatementId::new(9),
params: sample_params(),
},
ExecuteCommand::Begin {
isolation: IsolationLevel::Serializable,
},
ExecuteCommand::Commit,
ExecuteCommand::Rollback,
ExecuteCommand::Cancel {
query_id: QueryId::new_random(),
},
ExecuteCommand::GetSchema {
table: "events".to_owned(),
},
ExecuteCommand::Admin(AdminCommand::GrantRole {
username: "alice".to_owned(),
role: "analyst".to_owned(),
}),
] {
assert_serde_round_trip(&command);
}
}
#[test]
fn result_limits_defaults_are_bounded() {
let limits = ResultLimits::default();
assert_eq!(limits.max_rows, None);
assert_eq!(limits.effective_max_rows(), DEFAULT_MAX_ROWS);
assert_eq!(limits.effective_max_bytes(), DEFAULT_MAX_BYTES);
assert_eq!(
limits.effective_max_candidate_count(),
DEFAULT_MAX_CANDIDATE_COUNT
);
let explicit = ResultLimits {
max_rows: Some(10),
max_bytes: Some(1024),
max_candidate_count: Some(50),
};
assert_eq!(explicit.effective_max_rows(), 10);
assert_eq!(explicit.effective_max_bytes(), 1024);
assert_eq!(explicit.effective_max_candidate_count(), 50);
assert_serde_round_trip(&limits);
assert_serde_round_trip(&explicit);
}
#[test]
fn execute_request_serde_round_trip() {
let full = ExecuteRequest {
request_id: [0x11; 16],
query_id: QueryId::new_random(),
session_id: Some(SessionId::from_bytes([0x22; 16])),
database_id: DatabaseId::new_random(),
principal: sample_principal(),
command: ExecuteCommand::Sql {
text: "SELECT 1".to_owned(),
params: sample_params(),
},
deadline_unix_micros: Some(1_758_000_000_000_000),
result_limits: ResultLimits {
max_rows: Some(500),
max_bytes: None,
max_candidate_count: Some(10_000),
},
resource_group: Some("interactive".to_owned()),
idempotency_key: Some("req-42".to_owned()),
};
assert_serde_round_trip(&full);
let minimal = ExecuteRequest {
session_id: None,
deadline_unix_micros: None,
result_limits: ResultLimits::default(),
resource_group: None,
idempotency_key: None,
principal: AuthenticatedIdentity::Credentialless,
command: ExecuteCommand::Commit,
..full.clone()
};
assert_serde_round_trip(&minimal);
}
}