use std::os::raw::c_char;
use std::{ffi::NulError, fmt::Display};
use arrow_schema::ArrowError;
use crate::constants;
pub type AdbcStatusCode = u8;
#[derive(Debug, Copy, Clone, PartialEq, Eq)]
pub enum Status {
Ok,
Unknown,
NotImplemented,
NotFound,
AlreadyExists,
InvalidArguments,
InvalidState,
InvalidData,
Integrity,
Internal,
IO,
Cancelled,
Timeout,
Unauthenticated,
Unauthorized,
}
#[derive(Debug, PartialEq, Eq, Clone)]
pub struct Error {
pub message: String,
pub status: Status,
pub vendor_code: i32,
pub sqlstate: [c_char; 5], pub details: Option<Vec<(String, Vec<u8>)>>,
}
pub type Result<T> = std::result::Result<T, Error>;
impl Error {
pub fn with_message_and_status(message: impl Into<String>, status: Status) -> Self {
Self {
message: message.into(),
status,
vendor_code: 0,
sqlstate: [0; 5],
details: None,
}
}
}
impl Display for Error {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let safe_ascii = |c: c_char| -> char {
if c == 0 {
'0'
} else if (32..=126).contains(&c) {
char::from(c as u8)
} else {
'\u{FFFD}'
}
};
write!(
f,
"{:?}: {} (sqlstate: {}{}{}{}{}, vendor_code: {})",
self.status,
self.message,
safe_ascii(self.sqlstate[0]),
safe_ascii(self.sqlstate[1]),
safe_ascii(self.sqlstate[2]),
safe_ascii(self.sqlstate[3]),
safe_ascii(self.sqlstate[4]),
self.vendor_code
)
}
}
impl std::error::Error for Error {}
impl From<ArrowError> for Error {
fn from(value: ArrowError) -> Self {
Self {
message: value.to_string(),
status: Status::Internal,
vendor_code: 0,
sqlstate: [0; 5],
details: None,
}
}
}
impl From<NulError> for Error {
fn from(value: NulError) -> Self {
Self {
message: format!(
"Interior null byte was found at position {}",
value.nul_position()
),
status: Status::InvalidData,
vendor_code: 0,
sqlstate: [0; 5],
details: None,
}
}
}
impl From<std::str::Utf8Error> for Error {
fn from(value: std::str::Utf8Error) -> Self {
Self {
message: format!("Error while decoding UTF-8: {value}"),
status: Status::Internal,
vendor_code: 0,
sqlstate: [0; 5],
details: None,
}
}
}
impl From<std::ffi::IntoStringError> for Error {
fn from(value: std::ffi::IntoStringError) -> Self {
let error = value.utf8_error();
error.into()
}
}
impl TryFrom<AdbcStatusCode> for Status {
type Error = Error;
fn try_from(value: AdbcStatusCode) -> Result<Self> {
match value {
constants::ADBC_STATUS_OK => Ok(Status::Ok),
constants::ADBC_STATUS_UNKNOWN => Ok(Status::Unknown),
constants::ADBC_STATUS_NOT_IMPLEMENTED => Ok(Status::NotImplemented),
constants::ADBC_STATUS_NOT_FOUND => Ok(Status::NotFound),
constants::ADBC_STATUS_ALREADY_EXISTS => Ok(Status::AlreadyExists),
constants::ADBC_STATUS_INVALID_ARGUMENT => Ok(Status::InvalidArguments),
constants::ADBC_STATUS_INVALID_STATE => Ok(Status::InvalidState),
constants::ADBC_STATUS_INVALID_DATA => Ok(Status::InvalidData),
constants::ADBC_STATUS_INTEGRITY => Ok(Status::Integrity),
constants::ADBC_STATUS_INTERNAL => Ok(Status::Internal),
constants::ADBC_STATUS_IO => Ok(Status::IO),
constants::ADBC_STATUS_CANCELLED => Ok(Status::Cancelled),
constants::ADBC_STATUS_TIMEOUT => Ok(Status::Timeout),
constants::ADBC_STATUS_UNAUTHENTICATED => Ok(Status::Unauthenticated),
constants::ADBC_STATUS_UNAUTHORIZED => Ok(Status::Unauthorized),
v => Err(Error::with_message_and_status(
format!("Unknown status code: {v}"),
Status::InvalidData,
)),
}
}
}
impl From<Status> for AdbcStatusCode {
fn from(value: Status) -> Self {
match value {
Status::Ok => constants::ADBC_STATUS_OK,
Status::Unknown => constants::ADBC_STATUS_UNKNOWN,
Status::NotImplemented => constants::ADBC_STATUS_NOT_IMPLEMENTED,
Status::NotFound => constants::ADBC_STATUS_NOT_FOUND,
Status::AlreadyExists => constants::ADBC_STATUS_ALREADY_EXISTS,
Status::InvalidArguments => constants::ADBC_STATUS_INVALID_ARGUMENT,
Status::InvalidState => constants::ADBC_STATUS_INVALID_STATE,
Status::InvalidData => constants::ADBC_STATUS_INVALID_DATA,
Status::Integrity => constants::ADBC_STATUS_INTEGRITY,
Status::Internal => constants::ADBC_STATUS_INTERNAL,
Status::IO => constants::ADBC_STATUS_IO,
Status::Cancelled => constants::ADBC_STATUS_CANCELLED,
Status::Timeout => constants::ADBC_STATUS_TIMEOUT,
Status::Unauthenticated => constants::ADBC_STATUS_UNAUTHENTICATED,
Status::Unauthorized => constants::ADBC_STATUS_UNAUTHORIZED,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn display_unset_sqlstate() {
let err = Error::with_message_and_status("something failed", Status::Unknown);
let msg = err.to_string();
assert_eq!(
msg,
"Unknown: something failed (sqlstate: 00000, vendor_code: 0)"
);
}
#[test]
fn display_ascii_sqlstate() {
let err = Error {
message: "constraint violation".into(),
status: Status::Integrity,
vendor_code: 42,
sqlstate: [
b'2' as c_char,
b'3' as c_char,
b'5' as c_char,
b'0' as c_char,
b'5' as c_char,
],
details: None,
};
let msg = err.to_string();
assert_eq!(
msg,
"Integrity: constraint violation (sqlstate: 23505, vendor_code: 42)"
);
}
#[test]
fn display_non_printable_sqlstate() {
let err = Error {
message: "bad state".into(),
status: Status::Internal,
vendor_code: 0,
sqlstate: [1, 2, 3, 4, 5],
details: None,
};
let msg = err.to_string();
assert_eq!(msg, "Internal: bad state (sqlstate: �����, vendor_code: 0)");
}
#[test]
fn display_mixed_sqlstate() {
let err = Error {
message: "mixed".into(),
status: Status::InvalidData,
vendor_code: 7,
sqlstate: [b'H' as c_char, b'V' as c_char, 0, 0, 1],
details: None,
};
let msg = err.to_string();
assert_eq!(msg, "InvalidData: mixed (sqlstate: HV00�, vendor_code: 7)");
}
}