Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -30,3 +30,5 @@ unwrap_used = "deny"
cast_possible_truncation = "deny"
cast_possible_wrap = "deny"
cast_sign_loss = "deny"

map_err_ignore = "deny"
32 changes: 17 additions & 15 deletions src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,13 @@ pub enum ClientError {
/// The server sent a valid frame in an invalid protocol state.
#[error("unexpected server message: {0}")]
UnexpectedMessage(&'static str),
/// An `EXPLAIN` response contained invalid UTF-8.
#[error("unexpected server message: EXPLAIN result is not UTF-8")]
InvalidExplainUtf8 {
/// The underlying UTF-8 decoding failure.
#[source]
source: std::str::Utf8Error,
},
/// The database name is empty or exceeds the protocol limit.
#[error("invalid database name: {0}")]
InvalidDatabaseName(&'static str),
Expand Down Expand Up @@ -166,13 +173,7 @@ fn decode_complete(payload: &[u8], rows: Vec<Vec<Value>>) -> Result<QueryResult,
};
match kind {
COMPLETE_ROWS => {
if data.len() != 8 {
return Err(ClientError::UnexpectedMessage("row completion count is invalid"));
}
let expected =
u64::from_be_bytes(data.try_into().map_err(|_| {
ClientError::UnexpectedMessage("row completion count is invalid")
})?);
let expected = decode_completion_count(data, "row completion count is invalid")?;
if expected != rows.len() as u64 {
return Err(ClientError::UnexpectedMessage("row completion count does not match"));
}
Expand All @@ -181,18 +182,12 @@ fn decode_complete(payload: &[u8], rows: Vec<Vec<Value>>) -> Result<QueryResult,
COMPLETE_EXPLAIN => {
require_no_rows(&rows)?;
let plan = std::str::from_utf8(data)
.map_err(|_| ClientError::UnexpectedMessage("EXPLAIN result is not UTF-8"))?;
.map_err(|source| ClientError::InvalidExplainUtf8 { source })?;
Ok(QueryResult::Explain(plan.to_owned()))
}
COMPLETE_ROWS_AFFECTED => {
require_no_rows(&rows)?;
if data.len() != 8 {
return Err(ClientError::UnexpectedMessage("rows-affected count is invalid"));
}
let count =
u64::from_be_bytes(data.try_into().map_err(|_| {
ClientError::UnexpectedMessage("rows-affected count is invalid")
})?);
let count = decode_completion_count(data, "rows-affected count is invalid")?;
Ok(QueryResult::RowsAffected(count))
}
COMPLETE_SCHEMA_AFFECTED if data.is_empty() => {
Expand All @@ -210,6 +205,13 @@ fn decode_complete(payload: &[u8], rows: Vec<Vec<Value>>) -> Result<QueryResult,
}
}

fn decode_completion_count(data: &[u8], invalid_message: &'static str) -> Result<u64, ClientError> {
let [b0, b1, b2, b3, b4, b5, b6, b7] = data else {
return Err(ClientError::UnexpectedMessage(invalid_message));
};
Ok(u64::from_be_bytes([*b0, *b1, *b2, *b3, *b4, *b5, *b6, *b7]))
}

fn require_no_rows(rows: &[Vec<Value>]) -> Result<(), ClientError> {
if rows.is_empty() {
Ok(())
Expand Down
2 changes: 1 addition & 1 deletion src/executor/expression.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ pub fn evaluate_expression(
pub(super) fn execute_values(rows: Vec<Vec<PlannedExpression>>) -> ExecutorResult<ExecutionOutput> {
let rows = rows.into_iter().enumerate().map(|(row_index, expressions)| {
let table_key = TableKey::try_from(row_index)
.map_err(|_| ExecutorError::ValuesRowIndexOutOfRange { row_index })?;
.map_err(|_out_of_range| ExecutorError::ValuesRowIndexOutOfRange { row_index })?;
let input = empty_record(table_key)?;
evaluate_expressions(&expressions, &input)
});
Expand Down
42 changes: 31 additions & 11 deletions src/protocol.rs
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,24 @@ pub enum ProtocolError {
/// A frame or message payload is malformed.
#[error("malformed protocol message: {0}")]
Malformed(&'static str),
/// A text field in a message payload is not valid UTF-8.
#[error("malformed protocol message: {context}")]
InvalidUtf8 {
/// Description of the malformed text field.
context: &'static str,
/// The underlying UTF-8 decoding failure.
#[source]
source: std::str::Utf8Error,
},
/// Memory for decoded row values could not be reserved.
#[error("failed to allocate storage for {value_count} row values")]
RowAllocationFailed {
/// Number of values declared by the row payload.
value_count: usize,
/// The underlying allocation failure.
#[source]
source: std::collections::TryReserveError,
},
}

#[derive(Debug, PartialEq, Eq)]
Expand Down Expand Up @@ -179,7 +197,7 @@ pub(crate) fn write_frame(
return Err(ProtocolError::FrameTooLarge { length: payload.len() });
}
let length = u32::try_from(payload.len())
.map_err(|_| ProtocolError::FrameTooLarge { length: payload.len() })?;
.map_err(|_out_of_range| ProtocolError::FrameTooLarge { length: payload.len() })?;
let mut header = [0_u8; HEADER_LEN];
header[..4].copy_from_slice(&MAGIC);
header[4..6].copy_from_slice(&VERSION.to_be_bytes());
Expand All @@ -202,14 +220,15 @@ pub(crate) fn decode_error(payload: &[u8]) -> Result<(ErrorCode, String), Protoc
return Err(ProtocolError::Malformed("error response has no error code"));
}
let code = ErrorCode::from_u16(u16::from_be_bytes([payload[0], payload[1]]));
let message = std::str::from_utf8(&payload[2..])
.map_err(|_| ProtocolError::Malformed("error message is not UTF-8"))?;
let message = std::str::from_utf8(&payload[2..]).map_err(|source| {
ProtocolError::InvalidUtf8 { context: "error message is not UTF-8", source }
})?;
Ok((code, message.to_owned()))
}

pub(crate) fn encode_row(values: &[Value]) -> Result<Vec<u8>, ProtocolError> {
let count = u32::try_from(values.len())
.map_err(|_| ProtocolError::Malformed("row has too many values"))?;
.map_err(|_out_of_range| ProtocolError::Malformed("row has too many values"))?;
let mut payload = Vec::new();
payload.extend_from_slice(&count.to_be_bytes());
for value in values {
Expand All @@ -218,7 +237,7 @@ pub(crate) fn encode_row(values: &[Value]) -> Result<Vec<u8>, ProtocolError> {
Value::String(value) => {
payload.push(0x01);
let length = u32::try_from(value.len())
.map_err(|_| ProtocolError::Malformed("text value is too large"))?;
.map_err(|_out_of_range| ProtocolError::Malformed("text value is too large"))?;
payload.extend_from_slice(&length.to_be_bytes());
payload.extend_from_slice(value.as_bytes());
}
Expand Down Expand Up @@ -258,15 +277,16 @@ pub(crate) fn decode_row(payload: &[u8]) -> Result<Vec<Value>, ProtocolError> {
let mut values = Vec::new();
values
.try_reserve_exact(count)
.map_err(|_| ProtocolError::Malformed("cannot allocate row values"))?;
.map_err(|source| ProtocolError::RowAllocationFailed { value_count: count, source })?;
for _ in 0..count {
values.push(match decoder.u8()? {
0x00 => Value::Null,
0x01 => {
let length = decoder.u32()? as usize;
let bytes = decoder.bytes(length)?;
let value = std::str::from_utf8(bytes)
.map_err(|_| ProtocolError::Malformed("text value is not UTF-8"))?;
let value = std::str::from_utf8(bytes).map_err(|source| {
ProtocolError::InvalidUtf8 { context: "text value is not UTF-8", source }
})?;
Value::String(value.to_owned())
}
0x02 => match decoder.u8()? {
Expand Down Expand Up @@ -318,9 +338,9 @@ impl<'a> Decoder<'a> {
}

fn array<const N: usize>(&mut self) -> Result<[u8; N], ProtocolError> {
self.bytes(N)?
.try_into()
.map_err(|_| ProtocolError::Malformed("message payload is truncated"))
let mut array = [0; N];
array.copy_from_slice(self.bytes(N)?);
Ok(array)
}

fn u8(&mut self) -> Result<u8, ProtocolError> {
Expand Down
18 changes: 10 additions & 8 deletions src/relational/cursor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,10 +15,11 @@ fn encode_table_key(table_key: TableKey) -> [u8; TABLE_KEY_SIZE] {
}

fn decode_table_key(key: &[u8]) -> StorageResult<TableKey> {
let bytes: [u8; TABLE_KEY_SIZE] = key.try_into().map_err(|_| PageError::CorruptCell {
slot_index: 0,
kind: CellCorruption::InvalidTableKeyLength { actual: key.len() },
})?;
let bytes: [u8; TABLE_KEY_SIZE] =
key.try_into().map_err(|_wrong_length| PageError::CorruptCell {
slot_index: 0,
kind: CellCorruption::InvalidTableKeyLength { actual: key.len() },
})?;
Ok(TableKey::from_be_bytes(bytes) ^ TableKey::MIN)
}

Expand All @@ -34,10 +35,11 @@ pub(crate) fn encode_index_entry_key(index_key: &[u8], table_key: TableKey) -> V
}

fn decode_index_table_key(value: &[u8]) -> StorageResult<TableKey> {
let bytes: [u8; TABLE_KEY_SIZE] = value.try_into().map_err(|_| PageError::CorruptCell {
slot_index: 0,
kind: CellCorruption::InvalidIndexTableKeyValueLength { actual: value.len() },
})?;
let bytes: [u8; TABLE_KEY_SIZE] =
value.try_into().map_err(|_wrong_length| PageError::CorruptCell {
slot_index: 0,
kind: CellCorruption::InvalidIndexTableKeyValueLength { actual: value.len() },
})?;
Ok(TableKey::from_be_bytes(bytes) ^ TableKey::MIN)
}

Expand Down
4 changes: 2 additions & 2 deletions src/relational/tuple.rs
Original file line number Diff line number Diff line change
Expand Up @@ -434,7 +434,7 @@ where
I::IntoIter: ExactSizeIterator,
{
let values = values.into_iter();
let value_count = u32::try_from(values.len()).map_err(|_| {
let value_count = u32::try_from(values.len()).map_err(|_out_of_range| {
io::Error::new(io::ErrorKind::InvalidInput, "tuple value count exceeds u32::MAX")
})?;
writer.write_all(&value_count.to_le_bytes())?;
Expand All @@ -450,7 +450,7 @@ fn write_value_ref<W: Write>(writer: &mut W, value: ValueRef<'_>) -> io::Result<
match value {
ValueRef::Null => write_tlv_header(writer, TAG_NULL, NULL_LENGTH),
ValueRef::String(value) => {
let len = u32::try_from(value.len()).map_err(|_| {
let len = u32::try_from(value.len()).map_err(|_out_of_range| {
io::Error::new(io::ErrorKind::InvalidInput, "string length exceeds u32::MAX")
})?;
write_tlv_header(writer, TAG_STRING, len)?;
Expand Down
28 changes: 25 additions & 3 deletions src/sql_parser/parser/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -242,9 +242,12 @@ impl<'a> Parser<'a> {
TokenKind::Identifier(id) => Expression::Identifier(id),
TokenKind::Asterisk => Expression::Wildcard,
TokenKind::LeftParen => {
let lhs = self
.expr_bp(0)
.map_err(|_| SQLError::new(SQLErrorKind::UnclosedParenthesis, token.offset))?;
let lhs = self.expr_bp(0).map_err(|error| match error.kind {
SQLErrorKind::UnexpectedEnd => {
SQLError::new(SQLErrorKind::UnclosedParenthesis, token.offset)
}
_ => error,
})?;
self.lexer.expect_token(TokenKind::RightParen)?;
lhs
}
Expand Down Expand Up @@ -406,6 +409,25 @@ mod parser_tests {
assert_eq!(Err(expected_err), parser.expr());
}

#[test]
fn test_parenthesized_expression_preserves_inner_error() {
let parser = Parser::new("(operand invalid_operator)");
let expected_err = SQLError::new(
SQLErrorKind::InvalidOperator { op: TokenKind::Identifier("invalid_operator") },
9,
);

assert_eq!(Err(expected_err), parser.expr());
}

#[test]
fn test_unclosed_parenthesized_expression_maps_unexpected_end() {
let parser = Parser::new("(operand +");
let expected_err = SQLError::new(SQLErrorKind::UnclosedParenthesis, 0);

assert_eq!(Err(expected_err), parser.expr());
}

#[test]
fn test_parse_inequality_operators() {
let s = "12 < 34";
Expand Down
5 changes: 3 additions & 2 deletions src/storage/btree/rebalance.rs
Original file line number Diff line number Diff line change
Expand Up @@ -147,8 +147,9 @@ impl TreeCursor {
if child_index == usize::from(slot_count) {
return Ok(interior.rightmost_child());
}
let slot_index = u16::try_from(child_index)
.map_err(|_| PageError::InvalidSlotIndex { slot_index: u16::MAX, slot_count })?;
let slot_index = u16::try_from(child_index).map_err(|_out_of_range| {
PageError::InvalidSlotIndex { slot_index: u16::MAX, slot_count }
})?;
if slot_index > slot_count {
return Err(PageError::InvalidSlotIndex { slot_index, slot_count }.into());
}
Expand Down
2 changes: 1 addition & 1 deletion src/storage/btree/rebalance_repair.rs
Original file line number Diff line number Diff line change
Expand Up @@ -123,7 +123,7 @@ impl TreeCursor {
let child_ref = if child_index + 1 == child_count {
ChildSlotRef::Rightmost
} else {
let slot_index = u16::try_from(child_index).map_err(|_| {
let slot_index = u16::try_from(child_index).map_err(|_out_of_range| {
PageError::InvalidSlotIndex { slot_index: u16::MAX, slot_count }
})?;
ChildSlotRef::Slot(slot_index)
Expand Down
2 changes: 1 addition & 1 deletion src/storage/btree/split.rs
Original file line number Diff line number Diff line change
Expand Up @@ -317,7 +317,7 @@ impl TreeCursor {
.ok_or(StorageError::Internal(InternalError::InvariantViolation(
InvariantViolation::LeafSplitTargetMissing,
)))?;
let target_slot_index = u16::try_from(target_slot_index).map_err(|_| {
let target_slot_index = u16::try_from(target_slot_index).map_err(|_out_of_range| {
StorageError::Internal(InternalError::InvariantViolation(
InvariantViolation::LeafSplitTargetSlotOutOfRange { slot_index: target_slot_index },
))
Expand Down
27 changes: 14 additions & 13 deletions src/storage/log_manager/frame.rs
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@ pub(super) fn serialize_transaction<'a, W: Write>(
) -> Result<(), LogManagerError> {
validate_record_txn_ids(txn_id, records)?;
let entry_count = u32::try_from(records.len())
.map_err(|_| LogManagerError::TooManyRecords { count: records.len() })?;
.map_err(|_out_of_range| LogManagerError::TooManyRecords { count: records.len() })?;
let payload_len = serialized_records_len(records)?;

write_frame_header(&mut writer, txn_id, entry_count, payload_len)?;
Expand Down Expand Up @@ -87,7 +87,7 @@ pub(super) fn deserialize_transaction(
let entry_count = cursor.read_u32()?;
let payload_len = cursor.read_u64()?;
let payload_len = usize::try_from(payload_len)
.map_err(|_| LogManagerError::PayloadLengthTooLarge { payload_len })?;
.map_err(|_out_of_range| LogManagerError::PayloadLengthTooLarge { payload_len })?;

let payload_start = cursor.position;
let payload_end =
Expand Down Expand Up @@ -163,8 +163,9 @@ pub(super) fn scan_transaction_frame<R: Read>(
return Ok(None);
};
let header = parse_header(&header_bytes)?;
let payload_len = usize::try_from(header.payload_len)
.map_err(|_| LogManagerError::PayloadLengthTooLarge { payload_len: header.payload_len })?;
let payload_len = usize::try_from(header.payload_len).map_err(|_out_of_range| {
LogManagerError::PayloadLengthTooLarge { payload_len: header.payload_len }
})?;
let mut remaining = payload_len;
let mut digest = CRC32.digest();
let mut actual_record_count = 0usize;
Expand Down Expand Up @@ -361,12 +362,12 @@ pub(super) fn transaction_frame_len(buf: &'_ [u8]) -> Result<usize, LogManagerEr
return Err(LogManagerError::TruncatedFrame { needed: HEADER_LEN, remaining: buf.len() });
}

let header_bytes: &[u8; HEADER_LEN] = (&buf[..HEADER_LEN]).try_into().map_err(|_| {
LogManagerError::TruncatedFrame { needed: HEADER_LEN, remaining: buf.len() }
let mut header_bytes = [0; HEADER_LEN];
header_bytes.copy_from_slice(&buf[..HEADER_LEN]);
let header = parse_header(&header_bytes)?;
let payload_len = usize::try_from(header.payload_len).map_err(|_out_of_range| {
LogManagerError::PayloadLengthTooLarge { payload_len: header.payload_len }
})?;
let header = parse_header(header_bytes)?;
let payload_len = usize::try_from(header.payload_len)
.map_err(|_| LogManagerError::PayloadLengthTooLarge { payload_len: header.payload_len })?;

HEADER_LEN
.checked_add(payload_len)
Expand Down Expand Up @@ -464,10 +465,10 @@ fn serialized_record_len<'a>(kind: &LogRecordKind<'a>) -> Result<u64, LogManager
LogRecordKind::PageUpdate { redo_data, undo_data, .. } => {
validate_page_image_len(redo_data)?;
validate_page_image_len(undo_data)?;
let redo_len = u32::try_from(redo_data.len()).map_err(|_| {
let redo_len = u32::try_from(redo_data.len()).map_err(|_out_of_range| {
LogManagerError::PayloadLengthTooLarge { payload_len: redo_data.len() as u64 }
})?;
let undo_len = u32::try_from(undo_data.len()).map_err(|_| {
let undo_len = u32::try_from(undo_data.len()).map_err(|_out_of_range| {
LogManagerError::PayloadLengthTooLarge { payload_len: undo_data.len() as u64 }
})?;
Ok(1 + 8 + 4 + 4 + u64::from(redo_len) + u64::from(undo_len))
Expand Down Expand Up @@ -523,10 +524,10 @@ fn write_log_record_payload<'a, W: Write>(
LogRecordKind::PageUpdate { page_id, redo_data, undo_data } => {
validate_page_image_len(redo_data)?;
validate_page_image_len(undo_data)?;
let redo_len = u32::try_from(redo_data.len()).map_err(|_| {
let redo_len = u32::try_from(redo_data.len()).map_err(|_out_of_range| {
LogManagerError::PayloadLengthTooLarge { payload_len: redo_data.len() as u64 }
})?;
let undo_len = u32::try_from(undo_data.len()).map_err(|_| {
let undo_len = u32::try_from(undo_data.len()).map_err(|_out_of_range| {
LogManagerError::PayloadLengthTooLarge { payload_len: undo_data.len() as u64 }
})?;

Expand Down
7 changes: 3 additions & 4 deletions src/storage/log_manager/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -441,7 +441,7 @@ impl LogManager {
validate_record_txn_ids(txn_id, records)?;

let record_count = u64::try_from(records.len())
.map_err(|_| LogManagerError::TooManyRecords { count: records.len() })?;
.map_err(|_out_of_range| LogManagerError::TooManyRecords { count: records.len() })?;
let lsn = self
.highest_appended_lsn
.unwrap_or(ZERO_LSN)
Expand Down Expand Up @@ -687,9 +687,8 @@ impl RecoveryLogRecordKind {
}

fn page_image_array(image: &[u8]) -> Result<Box<[u8; PAGE_SIZE]>, LogManagerError> {
let image = image.try_into().map_err(|_| LogManagerError::InvalidPageImageLength {
expected: PAGE_SIZE,
actual: image.len(),
let image = image.try_into().map_err(|_wrong_length| {
LogManagerError::InvalidPageImageLength { expected: PAGE_SIZE, actual: image.len() }
})?;
Ok(Box::new(image))
}
Expand Down
Loading
Loading