diff --git a/Cargo.toml b/Cargo.toml index 5ecf949..47d0df3 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -30,3 +30,5 @@ unwrap_used = "deny" cast_possible_truncation = "deny" cast_possible_wrap = "deny" cast_sign_loss = "deny" + +map_err_ignore = "deny" diff --git a/src/client.rs b/src/client.rs index 8826783..3f99639 100644 --- a/src/client.rs +++ b/src/client.rs @@ -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), @@ -166,13 +173,7 @@ fn decode_complete(payload: &[u8], rows: Vec>) -> Result { - 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")); } @@ -181,18 +182,12 @@ fn decode_complete(payload: &[u8], rows: Vec>) -> Result { 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() => { @@ -210,6 +205,13 @@ fn decode_complete(payload: &[u8], rows: Vec>) -> Result Result { + 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]) -> Result<(), ClientError> { if rows.is_empty() { Ok(()) diff --git a/src/executor/expression.rs b/src/executor/expression.rs index c4ea3e1..ab2dc9d 100644 --- a/src/executor/expression.rs +++ b/src/executor/expression.rs @@ -27,7 +27,7 @@ pub fn evaluate_expression( pub(super) fn execute_values(rows: Vec>) -> ExecutorResult { 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) }); diff --git a/src/protocol.rs b/src/protocol.rs index eb58ff1..11f4bc5 100644 --- a/src/protocol.rs +++ b/src/protocol.rs @@ -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)] @@ -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()); @@ -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, 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 { @@ -218,7 +237,7 @@ pub(crate) fn encode_row(values: &[Value]) -> Result, 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()); } @@ -258,15 +277,16 @@ pub(crate) fn decode_row(payload: &[u8]) -> Result, 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()? { @@ -318,9 +338,9 @@ impl<'a> Decoder<'a> { } fn array(&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 { diff --git a/src/relational/cursor.rs b/src/relational/cursor.rs index 9472b54..b8069a1 100644 --- a/src/relational/cursor.rs +++ b/src/relational/cursor.rs @@ -15,10 +15,11 @@ fn encode_table_key(table_key: TableKey) -> [u8; TABLE_KEY_SIZE] { } fn decode_table_key(key: &[u8]) -> StorageResult { - 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) } @@ -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 { - 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) } diff --git a/src/relational/tuple.rs b/src/relational/tuple.rs index 2c500f8..61cdc64 100644 --- a/src/relational/tuple.rs +++ b/src/relational/tuple.rs @@ -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())?; @@ -450,7 +450,7 @@ fn write_value_ref(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)?; diff --git a/src/sql_parser/parser/mod.rs b/src/sql_parser/parser/mod.rs index 05f1086..e3e7df0 100644 --- a/src/sql_parser/parser/mod.rs +++ b/src/sql_parser/parser/mod.rs @@ -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 } @@ -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"; diff --git a/src/storage/btree/rebalance.rs b/src/storage/btree/rebalance.rs index 9d1760f..8118eb3 100644 --- a/src/storage/btree/rebalance.rs +++ b/src/storage/btree/rebalance.rs @@ -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()); } diff --git a/src/storage/btree/rebalance_repair.rs b/src/storage/btree/rebalance_repair.rs index 5a370af..882fbd6 100644 --- a/src/storage/btree/rebalance_repair.rs +++ b/src/storage/btree/rebalance_repair.rs @@ -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) diff --git a/src/storage/btree/split.rs b/src/storage/btree/split.rs index f05b7d5..034326a 100644 --- a/src/storage/btree/split.rs +++ b/src/storage/btree/split.rs @@ -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 }, )) diff --git a/src/storage/log_manager/frame.rs b/src/storage/log_manager/frame.rs index 221a390..b1c8612 100644 --- a/src/storage/log_manager/frame.rs +++ b/src/storage/log_manager/frame.rs @@ -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)?; @@ -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 = @@ -163,8 +163,9 @@ pub(super) fn scan_transaction_frame( 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; @@ -361,12 +362,12 @@ pub(super) fn transaction_frame_len(buf: &'_ [u8]) -> Result(kind: &LogRecordKind<'a>) -> Result { 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)) @@ -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 } })?; diff --git a/src/storage/log_manager/mod.rs b/src/storage/log_manager/mod.rs index 21867af..ebf7145 100644 --- a/src/storage/log_manager/mod.rs +++ b/src/storage/log_manager/mod.rs @@ -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) @@ -687,9 +687,8 @@ impl RecoveryLogRecordKind { } fn page_image_array(image: &[u8]) -> Result, 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)) } diff --git a/src/storage/page/core.rs b/src/storage/page/core.rs index a0b2a88..618e180 100644 --- a/src/storage/page/core.rs +++ b/src/storage/page/core.rs @@ -120,7 +120,8 @@ impl Freeblock { } fn encode_page_u16(value: usize) -> PageResult { - u16::try_from(value).map_err(|_| PageError::CellTooLarge { len: value, max: u16::MAX as usize }) + u16::try_from(value) + .map_err(|_out_of_range| PageError::CellTooLarge { len: value, max: u16::MAX as usize }) } #[derive(Debug, Clone, Copy)] diff --git a/src/storage/page/interior.rs b/src/storage/page/interior.rs index 8072fa0..0a47e43 100644 --- a/src/storage/page/interior.rs +++ b/src/storage/page/interior.rs @@ -177,8 +177,9 @@ where first_overflow_page_id: Option, inline_payload: &[u8], ) -> PageResult { - let encoded_key_len = u16::try_from(key_len) - .map_err(|_| PageError::CellTooLarge { len: key_len, max: u16::MAX as usize })?; + let encoded_key_len = u16::try_from(key_len).map_err(|_out_of_range| { + PageError::CellTooLarge { len: key_len, max: u16::MAX as usize } + })?; let Some(expected_inline_len) = format::inline_payload_len(key_len, first_overflow_page_id) else { return Err(PageError::CellTooLarge { len: key_len, max: u16::MAX as usize }); diff --git a/src/storage/page/leaf.rs b/src/storage/page/leaf.rs index 6884503..8f61fee 100644 --- a/src/storage/page/leaf.rs +++ b/src/storage/page/leaf.rs @@ -119,12 +119,15 @@ fn validate_payload_parts( let payload_len = key_len .checked_add(value_len) .ok_or(PageError::CellTooLarge { len: usize::MAX, max: u16::MAX as usize })?; - let encoded_key_len = u16::try_from(key_len) - .map_err(|_| PageError::CellTooLarge { len: payload_len, max: u16::MAX as usize })?; - let encoded_value_len = u16::try_from(value_len) - .map_err(|_| PageError::CellTooLarge { len: payload_len, max: u16::MAX as usize })?; - let encoded_payload_len = u16::try_from(payload_len) - .map_err(|_| PageError::CellTooLarge { len: payload_len, max: u16::MAX as usize })?; + let encoded_key_len = u16::try_from(key_len).map_err(|_out_of_range| { + PageError::CellTooLarge { len: payload_len, max: u16::MAX as usize } + })?; + let encoded_value_len = u16::try_from(value_len).map_err(|_out_of_range| { + PageError::CellTooLarge { len: payload_len, max: u16::MAX as usize } + })?; + let encoded_payload_len = u16::try_from(payload_len).map_err(|_out_of_range| { + PageError::CellTooLarge { len: payload_len, max: u16::MAX as usize } + })?; let Some(expected_inline_len) = format::inline_payload_len(payload_len, first_overflow_page_id) else { return Err(PageError::CellTooLarge { len: payload_len, max: u16::MAX as usize }); diff --git a/src/storage/page_cache.rs b/src/storage/page_cache.rs index 49a8883..fa05c58 100644 --- a/src/storage/page_cache.rs +++ b/src/storage/page_cache.rs @@ -271,7 +271,7 @@ impl PageCache { self.inner.runtime.read_page(new_page_id, &mut data)?; { - let mut frame_data = frame.data.try_borrow_mut().map_err(|_| { + let mut frame_data = frame.data.try_borrow_mut().map_err(|_borrow_conflict| { PageCacheError::PageMutableBorrowConflict { page_id: old_page_id.unwrap_or(new_page_id), } @@ -307,7 +307,7 @@ impl PageCache { let page = frame .data .try_borrow() - .map_err(|_| PageCacheError::PageImmutableBorrowConflict { page_id })?; + .map_err(|_borrow_conflict| PageCacheError::PageImmutableBorrowConflict { page_id })?; self.inner .runtime .flush_wal_through(frame.lsn.get()) @@ -325,7 +325,7 @@ impl PageCache { let pin = self.fetch_page(restore.page_id)?; let frame = &self.inner.frames[pin.frame_id]; { - let mut data = frame.data.try_borrow_mut().map_err(|_| { + let mut data = frame.data.try_borrow_mut().map_err(|_borrow_conflict| { PageCacheError::PageMutableBorrowConflict { page_id: restore.page_id } })?; *data = restore.image; @@ -369,10 +369,9 @@ impl PinGuard { /// fails while a write guard is active. pub(crate) fn read(&self) -> PageCacheResult> { let frame = &self.page_cache.frames[self.frame_id]; - let page = frame - .data - .try_borrow() - .map_err(|_| PageCacheError::PageImmutableBorrowConflict { page_id: self.page_id })?; + let page = frame.data.try_borrow().map_err(|_borrow_conflict| { + PageCacheError::PageImmutableBorrowConflict { page_id: self.page_id } + })?; Ok(PageReadGuard { page }) } @@ -383,10 +382,9 @@ impl PinGuard { /// caller later decides not to mutate the page bytes. pub(crate) fn write(&self) -> PageCacheResult> { let frame = &self.page_cache.frames[self.frame_id]; - let page = frame - .data - .try_borrow_mut() - .map_err(|_| PageCacheError::PageMutableBorrowConflict { page_id: self.page_id })?; + let page = frame.data.try_borrow_mut().map_err(|_borrow_conflict| { + PageCacheError::PageMutableBorrowConflict { page_id: self.page_id } + })?; let before = *page; let was_dirty = frame.dirty.get(); frame.dirty.set(true);