use core::ffi::c_char;
use core::pin::Pin;
use core::ptr::NonNull;
use core::slice;
use core::time::Duration;
use std::ffi::CString;
use crate::alloc::Allocator;
use crate::block::Block;
use crate::builder::BlockBuilder;
use crate::codec::{Codec, Compression};
use crate::error::{Error, ErrorKind, Result, check};
use crate::io::Io;
use crate::query::{QueryOpts, RawQueryOpts, cstring};
use crate::sys;
#[derive(Clone, Debug, Default)]
pub struct ClientOpts {
client_name: Option<String>,
database: Option<String>,
user: Option<String>,
password: Option<String>,
pub client_version_major: u64,
pub client_version_minor: u64,
pub client_version_patch: u64,
pub compression: Compression,
pub read_buffer_bytes: usize,
}
impl ClientOpts {
pub fn new() -> Self {
Self::default()
}
pub fn client_name(mut self, s: &str) -> Self {
self.client_name = Some(s.to_owned());
self
}
pub fn database(mut self, s: &str) -> Self {
self.database = Some(s.to_owned());
self
}
pub fn user(mut self, s: &str) -> Self {
self.user = Some(s.to_owned());
self
}
pub fn password(mut self, s: &str) -> Self {
self.password = Some(s.to_owned());
self
}
pub fn client_version(mut self, major: u64, minor: u64, patch: u64) -> Self {
self.client_version_major = major;
self.client_version_minor = minor;
self.client_version_patch = patch;
self
}
pub fn compression(mut self, compression: Compression) -> Self {
self.compression = compression;
self
}
pub(crate) fn to_raw(&self, codec: Option<*const sys::chc_codec>) -> Result<RawClientOpts> {
let mut owned = Vec::with_capacity(4);
let mut field = |label, value: &Option<String>| -> Result<*const c_char> {
let Some(value) = value else {
return Ok(core::ptr::null());
};
owned.push(cstring(label, value)?);
Ok(owned.last().expect("just pushed").as_ptr())
};
let client_name = field("client name", &self.client_name)?;
let database = field("database", &self.database)?;
let user = field("user", &self.user)?;
let password = field("password", &self.password)?;
Ok(RawClientOpts {
_owned: owned,
raw: sys::chc_client_opts {
client_name,
client_version_major: self.client_version_major,
client_version_minor: self.client_version_minor,
client_version_patch: self.client_version_patch,
database,
user,
password,
compression: self.compression as i32,
codec: codec.unwrap_or(core::ptr::null()),
read_buffer_bytes: self.read_buffer_bytes,
},
})
}
pub(crate) fn validate_codec(&self, codec: Option<Pin<&Codec>>) -> Result<()> {
if self.compression == Compression::None {
return Ok(());
}
let codec = codec.ok_or_else(|| {
Error::new(
ErrorKind::Usage,
format!("{:?} compression requires a codec", self.compression),
)
})?;
if codec.supports(self.compression) {
Ok(())
} else {
Err(Error::new(
ErrorKind::Usage,
format!("codec does not support {:?} compression", self.compression),
))
}
}
}
pub(crate) struct RawClientOpts {
_owned: Vec<CString>,
raw: sys::chc_client_opts,
}
impl RawClientOpts {
#[inline]
pub(crate) fn as_ptr(&self) -> *const sys::chc_client_opts {
&self.raw
}
}
#[derive(Debug, Clone)]
pub struct ServerInfo {
pub name: String,
pub timezone: String,
pub display_name: String,
pub version_major: u64,
pub version_minor: u64,
pub version_patch: u64,
pub revision: u64,
}
impl ServerInfo {
pub(crate) fn from_raw(raw: &sys::chc_server_info) -> Self {
Self {
name: cstr_array_to_string(&raw.name),
timezone: cstr_array_to_string(&raw.timezone),
display_name: cstr_array_to_string(&raw.display_name),
version_major: raw.version_major,
version_minor: raw.version_minor,
version_patch: raw.version_patch,
revision: raw.revision,
}
}
}
fn cstr_array_to_string(buf: &[c_char]) -> String {
let end = buf.iter().position(|&b| b == 0).unwrap_or(buf.len());
let bytes: &[u8] = unsafe { slice::from_raw_parts(buf.as_ptr().cast::<u8>(), end) };
String::from_utf8_lossy(bytes).into_owned()
}
pub struct Client<'fd> {
raw: NonNull<sys::chc_client>,
alloc: Box<Allocator>,
_codec: Option<Pin<Box<Codec>>>,
io: Pin<Box<dyn Io + Send + 'fd>>,
}
impl<'fd> Client<'fd> {
pub fn init<I: Io + Send + 'fd>(
opts: &ClientOpts,
alloc: Allocator,
mut io: Pin<Box<I>>,
codec: Option<Pin<Box<Codec>>>,
) -> Result<Self> {
opts.validate_codec(codec.as_ref().map(|codec| codec.as_ref()))?;
let codec_ptr = codec.as_ref().map(|c| c.as_ref().as_ptr());
let raw_opts = opts.to_raw(codec_ptr)?;
let alloc = Box::new(alloc);
let mut out: *mut sys::chc_client = core::ptr::null_mut();
let mut exc: *mut sys::chc_exception = core::ptr::null_mut();
let mut err = sys::chc_err::zeroed();
let rc = unsafe {
sys::chc_client_init(
&mut out,
raw_opts.as_ptr(),
alloc.as_ptr(),
io.as_mut().io_ptr(),
&mut exc,
&mut err,
)
};
if let Some(e) = take_handshake_exception(exc, *alloc) {
return Err(e);
}
check(rc, &err)?;
Ok(Self {
raw: NonNull::new(out).expect("chc_client_init returned OK with NULL"),
alloc,
_codec: codec,
io,
})
}
pub fn server_info(&self) -> Option<ServerInfo> {
let p = unsafe { sys::chc_client_server_info(self.raw.as_ptr().cast_const()) };
(!p.is_null()).then(|| ServerInfo::from_raw(unsafe { &*p }))
}
pub fn set_read_timeout(&mut self, timeout: Option<Duration>) -> Result<()> {
self.io.as_mut().set_read_timeout(timeout)
}
pub fn send_query(&mut self, sql: &str, query_id: Option<&str>) -> Result<()> {
let (qid, qid_len) = query_id
.map(|q| (q.as_ptr().cast::<c_char>(), q.len()))
.unwrap_or((core::ptr::null(), 0));
let mut err = sys::chc_err::zeroed();
let rc = unsafe {
sys::chc_client_send_query(
self.raw.as_ptr(),
sql.as_ptr().cast::<c_char>(),
sql.len(),
qid,
qid_len,
&mut err,
)
};
check(rc, &err)
}
pub fn send_query_with(&mut self, sql: &str, opts: &QueryOpts<'_>) -> Result<()> {
let raw_opts = RawQueryOpts::new(opts)?;
let mut err = sys::chc_err::zeroed();
let rc = unsafe {
sys::chc_client_send_query_ex(
self.raw.as_ptr(),
sql.as_ptr().cast::<c_char>(),
sql.len(),
raw_opts.as_ptr(),
&mut err,
)
};
check(rc, &err)
}
pub fn send_data(&mut self, builder: Option<&BlockBuilder<'_>>) -> Result<()> {
let bb_ptr = builder.map(|b| b.as_ptr()).unwrap_or(core::ptr::null());
let mut err = sys::chc_err::zeroed();
let rc = unsafe { sys::chc_client_send_data(self.raw.as_ptr(), bb_ptr, &mut err) };
check(rc, &err)
}
pub fn send_cancel(&mut self) -> Result<()> {
let mut err = sys::chc_err::zeroed();
let rc = unsafe { sys::chc_client_send_cancel(self.raw.as_ptr(), &mut err) };
check(rc, &err)
}
pub fn send_ping(&mut self) -> Result<()> {
let mut err = sys::chc_err::zeroed();
let rc = unsafe { sys::chc_client_send_ping(self.raw.as_ptr(), &mut err) };
check(rc, &err)
}
pub fn recv_event(&mut self) -> Result<Event> {
let mut raw = sys::chc_packet::zeroed();
let mut err = sys::chc_err::zeroed();
let rc = unsafe { sys::chc_client_recv_packet(self.raw.as_ptr(), &mut raw, &mut err) };
if let Err(e) = check(rc, &err) {
unsafe { sys::chc_packet_clear(self.raw.as_ptr(), &mut raw) };
return Err(e);
}
let event = Event::from_raw(&mut raw, *self.alloc);
unsafe { sys::chc_packet_clear(self.raw.as_ptr(), &mut raw) };
event
}
}
impl<'fd> Drop for Client<'fd> {
fn drop(&mut self) {
unsafe { sys::chc_client_close(self.raw.as_ptr()) };
}
}
unsafe impl<'fd> Send for Client<'fd> {}
pub struct Exception {
raw: NonNull<sys::chc_exception>,
alloc: Allocator,
}
impl Exception {
pub(crate) unsafe fn from_raw(raw: NonNull<sys::chc_exception>, alloc: Allocator) -> Self {
Self { raw, alloc }
}
pub fn code(&self) -> i32 {
unsafe { (*self.raw.as_ptr()).code }
}
pub fn name(&self) -> &[u8] {
let r = unsafe { self.raw.as_ref() };
cstr_bytes(r.name, r.name_len)
}
pub fn display_text(&self) -> &[u8] {
let r = unsafe { self.raw.as_ref() };
cstr_bytes(r.display_text, r.display_text_len)
}
pub fn stack_trace(&self) -> &[u8] {
let r = unsafe { self.raw.as_ref() };
cstr_bytes(r.stack_trace, r.stack_trace_len)
}
}
impl Drop for Exception {
fn drop(&mut self) {
unsafe { sys::chc_exception_free(self.raw.as_ptr(), self.alloc.as_ptr()) };
}
}
impl core::fmt::Debug for Exception {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("Exception")
.field("code", &self.code())
.field("name", &String::from_utf8_lossy(self.name()))
.field(
"display_text",
&String::from_utf8_lossy(self.display_text()),
)
.field("stack_trace", &String::from_utf8_lossy(self.stack_trace()))
.finish()
}
}
impl core::fmt::Display for Exception {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(
f,
"{} (code {}): {}",
String::from_utf8_lossy(self.name()),
self.code(),
String::from_utf8_lossy(self.display_text()),
)
}
}
impl std::error::Error for Exception {}
unsafe impl Send for Exception {}
impl From<Exception> for Error {
fn from(exc: Exception) -> Self {
Self {
kind: ErrorKind::Server,
server_code: exc.code(),
message: String::from_utf8_lossy(exc.display_text()).into_owned(),
server_name: String::from_utf8_lossy(exc.name()).into_owned(),
}
}
}
pub(crate) fn take_handshake_exception(
exc: *mut sys::chc_exception,
alloc: Allocator,
) -> Option<Error> {
NonNull::new(exc).map(|p| unsafe { Exception::from_raw(p, alloc) }.into())
}
fn cstr_bytes<'a>(ptr: *mut c_char, len: usize) -> &'a [u8] {
if ptr.is_null() || len == 0 {
return &[];
}
debug_assert!(
len <= isize::MAX as usize,
"clickhouse-c published exception field len = {len}",
);
unsafe { slice::from_raw_parts(ptr.cast::<u8>(), len) }
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(i32)]
pub enum PacketKind {
Data = sys::CHC_PKT_DATA,
Exception = sys::CHC_PKT_EXCEPTION,
Progress = sys::CHC_PKT_PROGRESS,
Pong = sys::CHC_PKT_PONG,
EndOfStream = sys::CHC_PKT_END_OF_STREAM,
ProfileInfo = sys::CHC_PKT_PROFILE_INFO,
Totals = sys::CHC_PKT_TOTALS,
Extremes = sys::CHC_PKT_EXTREMES,
Log = sys::CHC_PKT_LOG,
TableColumns = sys::CHC_PKT_TABLE_COLUMNS,
ProfileEvents = sys::CHC_PKT_PROFILE_EVENTS,
TimezoneUpdate = sys::CHC_PKT_TIMEZONE_UPDATE,
}
impl PacketKind {
pub(crate) fn from_raw(k: sys::chc_packet_kind) -> Option<Self> {
Some(match k {
sys::CHC_PKT_DATA => Self::Data,
sys::CHC_PKT_EXCEPTION => Self::Exception,
sys::CHC_PKT_PROGRESS => Self::Progress,
sys::CHC_PKT_PONG => Self::Pong,
sys::CHC_PKT_END_OF_STREAM => Self::EndOfStream,
sys::CHC_PKT_PROFILE_INFO => Self::ProfileInfo,
sys::CHC_PKT_TOTALS => Self::Totals,
sys::CHC_PKT_EXTREMES => Self::Extremes,
sys::CHC_PKT_LOG => Self::Log,
sys::CHC_PKT_TABLE_COLUMNS => Self::TableColumns,
sys::CHC_PKT_PROFILE_EVENTS => Self::ProfileEvents,
sys::CHC_PKT_TIMEZONE_UPDATE => Self::TimezoneUpdate,
_ => return None,
})
}
}
pub enum Event {
Data(Block),
Totals(Block),
Extremes(Block),
Log(Block),
ProfileEvents(Block),
Exception(Exception),
Progress(Progress),
ProfileInfo(ProfileInfo),
Pong,
EndOfStream,
TableColumns,
TimezoneUpdate,
}
impl Event {
pub(crate) fn from_raw(raw: &mut sys::chc_packet, alloc: Allocator) -> Result<Self> {
let Some(kind) = PacketKind::from_raw(raw.kind) else {
return Err(Error::new(
ErrorKind::Protocol,
format!("unknown server packet {}", raw.kind),
));
};
Ok(match kind {
PacketKind::Data => Self::Data(take_block(raw, alloc)?),
PacketKind::Totals => Self::Totals(take_block(raw, alloc)?),
PacketKind::Extremes => Self::Extremes(take_block(raw, alloc)?),
PacketKind::Log => Self::Log(take_block(raw, alloc)?),
PacketKind::ProfileEvents => Self::ProfileEvents(take_block(raw, alloc)?),
PacketKind::Exception => Self::Exception(take_exception(raw, alloc)?),
PacketKind::Progress => {
Self::Progress(Progress::from_raw(unsafe { &raw.payload.progress }))
}
PacketKind::ProfileInfo => {
Self::ProfileInfo(ProfileInfo::from_raw(unsafe { &raw.payload.profile }))
}
PacketKind::Pong => Self::Pong,
PacketKind::EndOfStream => Self::EndOfStream,
PacketKind::TableColumns => Self::TableColumns,
PacketKind::TimezoneUpdate => Self::TimezoneUpdate,
})
}
}
fn take_block(raw: &mut sys::chc_packet, alloc: Allocator) -> Result<Block> {
let p = unsafe { raw.payload.block };
raw.payload.block = core::ptr::null_mut();
unsafe { Block::from_raw(p, alloc) }
.ok_or_else(|| Error::new(ErrorKind::Protocol, "block packet missing block"))
}
fn take_exception(raw: &mut sys::chc_packet, alloc: Allocator) -> Result<Exception> {
let p = NonNull::new(unsafe { raw.payload.exception })
.ok_or_else(|| Error::new(ErrorKind::Protocol, "exception packet missing exception"))?;
raw.payload.exception = core::ptr::null_mut();
Ok(unsafe { Exception::from_raw(p, alloc) })
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Progress {
pub rows: u64,
pub bytes: u64,
pub total_rows: u64,
pub total_bytes: u64,
pub written_rows: u64,
pub written_bytes: u64,
pub elapsed_ns: u64,
}
impl Progress {
fn from_raw(raw: &sys::chc_packet_progress) -> Self {
Self {
rows: raw.rows,
bytes: raw.bytes,
total_rows: raw.total_rows,
total_bytes: raw.total_bytes,
written_rows: raw.written_rows,
written_bytes: raw.written_bytes,
elapsed_ns: raw.elapsed_ns,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ProfileInfo {
pub rows: u64,
pub blocks: u64,
pub bytes: u64,
pub rows_before_limit: u64,
pub applied_limit: bool,
pub calculated_rows_before_limit: bool,
}
impl ProfileInfo {
fn from_raw(raw: &sys::chc_packet_profile) -> Self {
Self {
rows: raw.rows,
blocks: raw.blocks,
bytes: raw.bytes,
rows_before_limit: raw.rows_before_limit,
applied_limit: raw.applied_limit != 0,
calculated_rows_before_limit: raw.calculated_rows_before_limit != 0,
}
}
}
#[cfg(test)]
mod tests {
use std::ffi::CStr;
use core::ffi::c_char;
use super::{ClientOpts, Event, PacketKind, cstr_bytes};
use crate::sys;
use crate::{Allocator, Compression, ErrorKind};
#[test]
fn compression_without_a_codec_is_a_usage_error() {
let err = ClientOpts::new()
.compression(Compression::Lz4)
.validate_codec(None)
.expect_err("missing codec");
assert_eq!(err.kind, ErrorKind::Usage);
}
#[cfg(feature = "lz4")]
#[test]
fn compression_requires_matching_codec() {
use crate::Codec;
let codec = Codec::lz4();
let mismatch = ClientOpts::new()
.compression(Compression::Zstd)
.validate_codec(Some(codec.as_ref()))
.expect_err("mismatched codec");
assert_eq!(mismatch.kind, ErrorKind::Usage);
}
#[cfg(feature = "lz4")]
#[test]
fn matching_codec_passes_validation() {
use crate::Codec;
let codec = Codec::lz4();
ClientOpts::new()
.compression(Compression::Lz4)
.validate_codec(Some(codec.as_ref()))
.expect("lz4 codec for lz4 compression");
}
#[test]
fn handshake_fields_reach_raw_opts() {
let opts = ClientOpts::new()
.client_name("probe")
.database("db")
.user("reader")
.password("secret")
.client_version(1, 2, 3);
let raw = opts.to_raw(None).expect("no interior NUL");
let field = |p: *const core::ffi::c_char| unsafe { CStr::from_ptr(p) }.to_owned();
let raw = unsafe { &*raw.as_ptr() };
assert_eq!(field(raw.client_name).to_bytes(), b"probe");
assert_eq!(field(raw.database).to_bytes(), b"db");
assert_eq!(field(raw.user).to_bytes(), b"reader");
assert_eq!(field(raw.password).to_bytes(), b"secret");
assert_eq!(
(
raw.client_version_major,
raw.client_version_minor,
raw.client_version_patch,
),
(1, 2, 3),
);
assert!(raw.codec.is_null());
}
#[test]
fn default_opts_leave_every_string_null() {
let raw = ClientOpts::new().to_raw(None).expect("no strings");
let raw = unsafe { &*raw.as_ptr() };
assert!(raw.client_name.is_null());
assert!(raw.database.is_null());
assert!(raw.user.is_null());
assert!(raw.password.is_null());
}
#[test]
fn unknown_packet_kind_is_none() {
assert!(PacketKind::from_raw(i32::MAX).is_none());
assert!(PacketKind::from_raw(-1).is_none());
}
#[test]
fn unknown_packet_is_a_protocol_error() {
let mut raw = sys::chc_packet::zeroed();
raw.kind = i32::MAX;
let err = Event::from_raw(&mut raw, Allocator::stdlib())
.err()
.expect("unknown packet kind accepted");
assert_eq!(err.kind, ErrorKind::Protocol);
assert!(err.message.contains("unknown server packet"), "{err}");
}
const KINDS: [(sys::chc_packet_kind, bool); 12] = [
(sys::CHC_PKT_DATA, true),
(sys::CHC_PKT_TOTALS, true),
(sys::CHC_PKT_EXTREMES, true),
(sys::CHC_PKT_LOG, true),
(sys::CHC_PKT_PROFILE_EVENTS, true),
(sys::CHC_PKT_EXCEPTION, true),
(sys::CHC_PKT_PROGRESS, false),
(sys::CHC_PKT_PROFILE_INFO, false),
(sys::CHC_PKT_PONG, false),
(sys::CHC_PKT_END_OF_STREAM, false),
(sys::CHC_PKT_TABLE_COLUMNS, false),
(sys::CHC_PKT_TIMEZONE_UPDATE, false),
];
#[test]
fn each_packet_kind_converts_to_its_variant() {
for (kind, carries_payload) in KINDS {
let mut raw = sys::chc_packet::zeroed();
raw.kind = kind;
let event = Event::from_raw(&mut raw, Allocator::stdlib());
assert_eq!(
matches!(&event, Err(err) if err.kind == ErrorKind::Protocol),
carries_payload,
"kind {kind}",
);
assert_eq!(matches!(event, Ok(Event::Pong)), kind == sys::CHC_PKT_PONG);
assert_eq!(
matches!(event, Ok(Event::EndOfStream)),
kind == sys::CHC_PKT_END_OF_STREAM,
);
assert_eq!(
matches!(event, Ok(Event::TableColumns)),
kind == sys::CHC_PKT_TABLE_COLUMNS,
);
assert_eq!(
matches!(event, Ok(Event::TimezoneUpdate)),
kind == sys::CHC_PKT_TIMEZONE_UPDATE,
);
assert_eq!(
matches!(event, Ok(Event::Progress(_))),
kind == sys::CHC_PKT_PROGRESS,
);
assert_eq!(
matches!(event, Ok(Event::ProfileInfo(_))),
kind == sys::CHC_PKT_PROFILE_INFO,
);
}
}
#[test]
fn every_c_packet_kind_has_a_variant() {
for (kind, _) in KINDS {
assert!(PacketKind::from_raw(kind).is_some(), "kind {kind}");
}
}
#[test]
fn empty_exception_fields_read_as_empty_slices() {
assert!(cstr_bytes(core::ptr::null_mut(), 7).is_empty());
let mut byte = b'x' as c_char;
assert!(cstr_bytes(&mut byte, 0).is_empty());
}
}