use crate::handles::slice_to_cow_utf8;
use super::{
as_handle::AsHandle,
buffer::{clamp_small_int, mut_buf_ptr},
SqlChar,
};
use odbc_sys::{SqlReturn, SQLSTATE_SIZE};
use std::fmt;
#[cfg(not(feature = "narrow"))]
use odbc_sys::SQLGetDiagRecW as sql_get_diag_rec;
#[cfg(feature = "narrow")]
use odbc_sys::SQLGetDiagRec as sql_get_diag_rec;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct State(pub [u8; SQLSTATE_SIZE]);
impl State {
pub const INVALID_STATE_TRANSACTION: State = State(*b"25000");
pub const INVALID_ATTRIBUTE_VALUE: State = State(*b"HY024");
pub const INVALID_SQL_DATA_TYPE: State = State(*b"HY004");
pub fn from_chars_with_nul(code: &[SqlChar; SQLSTATE_SIZE + 1]) -> Self {
let mut ascii = [0; SQLSTATE_SIZE];
for (index, letter) in code[..SQLSTATE_SIZE].iter().copied().enumerate() {
ascii[index] = letter as u8;
}
State(ascii)
}
pub fn as_str(&self) -> &str {
std::str::from_utf8(&self.0).unwrap()
}
}
#[derive(Debug, Clone, Copy)]
pub struct DiagnosticResult {
pub state: State,
pub native_error: i32,
}
pub fn diagnostics(
handle: &dyn AsHandle,
rec_number: i16,
message_text: &mut Vec<SqlChar>,
) -> Option<DiagnosticResult> {
assert!(rec_number > 0);
let cap = message_text.capacity();
message_text.resize(cap, 0);
let mut text_length = 0;
let mut state = [0; SQLSTATE_SIZE + 1];
let mut native_error = 0;
let ret = unsafe {
sql_get_diag_rec(
handle.handle_type(),
handle.as_handle(),
rec_number,
state.as_mut_ptr(),
&mut native_error,
mut_buf_ptr(message_text),
clamp_small_int(message_text.len()),
&mut text_length,
)
};
let result = DiagnosticResult {
state: State::from_chars_with_nul(&state),
native_error,
};
let mut text_length: usize = text_length.try_into().unwrap();
match ret {
SqlReturn::SUCCESS | SqlReturn::SUCCESS_WITH_INFO => {
if text_length > message_text.len() {
message_text.resize(text_length + 1, 0);
diagnostics(handle, rec_number, message_text)
} else {
while text_length > 0 && message_text[text_length - 1] == 0 {
text_length -= 1;
}
message_text.resize(text_length, 0);
Some(result)
}
}
SqlReturn::NO_DATA => None,
SqlReturn::ERROR => panic!("rec_number argument of diagnostics must be > 0."),
unexpected => panic!("SQLGetDiagRec returned: {:?}", unexpected),
}
}
#[derive(Default)]
pub struct Record {
pub state: State,
pub native_error: i32,
pub message: Vec<SqlChar>,
}
impl Record {
pub fn fill_from(&mut self, handle: &dyn AsHandle, record_number: i16) -> bool {
match diagnostics(handle, record_number, &mut self.message) {
Some(result) => {
self.state = result.state;
self.native_error = result.native_error;
true
}
None => false,
}
}
}
impl fmt::Display for Record {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let message = slice_to_cow_utf8(&self.message);
write!(
f,
"State: {}, Native error: {}, Message: {}",
self.state.as_str(),
self.native_error,
message,
)
}
}
impl fmt::Debug for Record {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(self, f)
}
}
#[cfg(test)]
mod tests {
use crate::handles::diagnostics::State;
use super::Record;
#[cfg(not(feature = "narrow"))]
fn to_vec_sql_char(text: &str) -> Vec<u16> {
text.encode_utf16().collect()
}
#[cfg(feature = "narrow")]
fn to_vec_sql_char(text: &str) -> Vec<u8> {
text.bytes().collect()
}
#[test]
fn formatting() {
let message = to_vec_sql_char("[Microsoft][ODBC Driver Manager] Function sequence error");
let rec = Record {
state: State(*b"HY010"),
message,
..Record::default()
};
assert_eq!(
format!("{}", rec),
"State: HY010, Native error: 0, Message: [Microsoft][ODBC Driver Manager] \
Function sequence error"
);
}
}