use super::KdbConnection;
use crate::nodes::{FutStream, RunParams, StreamOperators};
use crate::types::*;
use chrono::NaiveDateTime;
use futures::StreamExt;
use kdb_plus_fixed::ipc::{ConnectionMethod, K, QStream};
use kdb_plus_fixed::qtype;
use std::pin::Pin;
use std::rc::Rc;
pub trait KdbSerialize: Sized {
fn to_kdb_row(&self) -> K;
}
#[must_use]
pub fn kdb_write<T>(
connection: KdbConnection,
table_name: impl Into<String>,
upstream: &Rc<dyn Stream<Burst<T>>>,
) -> Rc<dyn Node>
where
T: Element + Send + KdbSerialize + 'static,
{
let table_name = table_name.into();
let consumer = Box::new(
move |_ctx: RunParams, source: Pin<Box<dyn FutStream<Burst<T>>>>| {
kdb_write_consumer(connection, table_name, source)
},
);
upstream.consume_async(consumer)
}
async fn kdb_write_consumer<T>(
connection: KdbConnection,
table_name: String,
mut source: Pin<Box<dyn FutStream<Burst<T>>>>,
) -> anyhow::Result<()>
where
T: Element + Send + KdbSerialize + 'static,
{
let creds = connection.credentials_string();
let mut socket = QStream::connect(
ConnectionMethod::TCP,
&connection.host,
connection.port,
&creds,
)
.await?;
while let Some((time, batch)) = source.next().await {
let naive: NaiveDateTime = time.into();
let ts_str = naive.format("%Y.%m.%dD%H:%M:%S%.9f").to_string();
let rows: Vec<K> = batch
.into_iter()
.map(|record| record.to_kdb_row())
.collect();
let columns = k_rows_to_columns(rows)?;
if columns.is_empty() {
continue;
}
let n = columns[0].len();
let ts_frag = if n == 1 {
format!("enlist {ts_str}")
} else {
std::iter::repeat_n(ts_str.as_str(), n)
.collect::<Vec<_>>()
.join(" ")
};
let mut col_frags: Vec<String> = vec![ts_frag];
for col in columns {
col_frags.push(format_kdb_column_q(&col)?);
}
let q = format!(
"insert[`{table_name}; ({cols})]",
cols = col_frags.join("; ")
);
let response = socket.send_sync_message(&q.as_str()).await?;
if response.get_type() == qtype::ERROR {
anyhow::bail!(
"KDB insert error: {}",
response.get_error_string().unwrap_or("unknown")
);
}
}
Ok(())
}
fn k_rows_to_columns(rows: Vec<K>) -> anyhow::Result<Vec<Vec<K>>> {
let serialized: Vec<Vec<K>> = rows
.into_iter()
.map(|row| {
let row_type = row.get_type();
row.as_vec::<K>().map(|v| v.to_vec()).map_err(|_| {
anyhow::anyhow!(
"kdb_write: KdbSerialize::to_kdb_row must return a compound list, \
got K type {row_type}"
)
})
})
.collect::<anyhow::Result<Vec<_>>>()?;
if serialized.is_empty() {
return Ok(Vec::new());
}
let n_cols = serialized[0].len();
if let Some(bad) = serialized.iter().position(|r| r.len() != n_cols) {
anyhow::bail!(
"kdb_write: ragged burst — row {bad} has {} columns, expected {n_cols} \
(all rows must share the table schema)",
serialized[bad].len()
);
}
let n = serialized.len();
let mut columns: Vec<Vec<K>> = (0..n_cols).map(|_| Vec::with_capacity(n)).collect();
for row in serialized {
for (col_idx, val) in row.into_iter().enumerate() {
columns[col_idx].push(val);
}
}
Ok(columns)
}
fn q_float64_token(v: f64) -> String {
if v.is_nan() {
"0n".to_string()
} else if v == f64::INFINITY {
"0w".to_string()
} else if v == f64::NEG_INFINITY {
"-0w".to_string()
} else {
format!("{v}")
}
}
fn q_float32_token(v: f32) -> String {
if v.is_nan() {
"0Ne".to_string()
} else if v == f32::INFINITY {
"0we".to_string()
} else if v == f32::NEG_INFINITY {
"-0we".to_string()
} else {
format!("{v}")
}
}
fn q_string_literal(s: &str) -> String {
format!("\"{}\"", s.replace('\\', "\\\\").replace('"', "\\\""))
}
fn format_kdb_column_q(atoms: &[K]) -> anyhow::Result<String> {
if atoms.is_empty() {
return Ok("()".to_string());
}
let col_type = atoms[0].get_type();
match col_type {
qtype::SYMBOL_ATOM => {
let syms: anyhow::Result<Vec<String>> = atoms
.iter()
.map(|k| Ok(k.get_symbol()?.to_string()))
.collect();
let syms = syms?;
if syms.len() == 1 {
Ok(format!("enlist `${}", q_string_literal(&syms[0])))
} else {
let parts: Vec<String> = syms.iter().map(|s| q_string_literal(s)).collect();
Ok(format!("`$({})", parts.join(";")))
}
}
qtype::FLOAT_ATOM => {
let vals: anyhow::Result<Vec<f64>> = atoms.iter().map(|k| Ok(k.get_float()?)).collect();
let vals = vals?;
if vals.iter().all(|v| v.is_finite()) {
if vals.len() == 1 {
Ok(format!("enlist {}f", vals[0]))
} else {
let parts: Vec<_> = vals.iter().map(|v| format!("{v}")).collect();
Ok(format!("{}f", parts.join(" ")))
}
} else {
let parts: Vec<String> = vals.iter().map(|v| q_float64_token(*v)).collect();
if parts.len() == 1 {
Ok(format!("enlist {}", parts[0]))
} else {
Ok(format!("`float$({})", parts.join(";")))
}
}
}
qtype::LONG_ATOM => {
let vals: anyhow::Result<Vec<i64>> = atoms.iter().map(|k| Ok(k.get_long()?)).collect();
let vals = vals?;
if vals.len() == 1 {
Ok(format!("enlist {}j", vals[0]))
} else {
let parts: Vec<_> = vals.iter().map(|v| format!("{v}")).collect();
Ok(format!("{}j", parts.join(" ")))
}
}
qtype::INT_ATOM => {
let vals: anyhow::Result<Vec<i32>> = atoms.iter().map(|k| Ok(k.get_int()?)).collect();
let vals = vals?;
if vals.len() == 1 {
Ok(format!("enlist {}i", vals[0]))
} else {
let parts: Vec<_> = vals.iter().map(|v| format!("{v}")).collect();
Ok(format!("{}i", parts.join(" ")))
}
}
qtype::BOOL_ATOM => {
let vals: anyhow::Result<Vec<bool>> = atoms.iter().map(|k| Ok(k.get_bool()?)).collect();
let vals = vals?;
if vals.len() == 1 {
Ok(format!("enlist {}b", if vals[0] { 1 } else { 0 }))
} else {
let parts: Vec<_> = vals
.iter()
.map(|v| format!("{}", if *v { 1 } else { 0 }))
.collect();
Ok(format!("{}b", parts.join(" ")))
}
}
qtype::REAL_ATOM => {
let vals: anyhow::Result<Vec<f32>> = atoms.iter().map(|k| Ok(k.get_real()?)).collect();
let vals = vals?;
if vals.iter().all(|v| v.is_finite()) {
if vals.len() == 1 {
Ok(format!("enlist {}e", vals[0]))
} else {
let parts: Vec<_> = vals.iter().map(|v| format!("{v}")).collect();
Ok(format!("{}e", parts.join(" ")))
}
} else {
let parts: Vec<String> = vals.iter().map(|v| q_float32_token(*v)).collect();
if parts.len() == 1 {
Ok(format!("enlist {}", parts[0]))
} else {
Ok(format!("`real$({})", parts.join(";")))
}
}
}
other => anyhow::bail!("unsupported KDB column type {other} in kdb_write"),
}
}
pub trait KdbWriteOperators<T: Element> {
#[must_use]
fn kdb_write(self: &Rc<Self>, conn: KdbConnection, table: &str) -> Rc<dyn Node>;
}
impl<T: Element + Send + KdbSerialize + 'static> KdbWriteOperators<T> for dyn Stream<Burst<T>> {
fn kdb_write(self: &Rc<Self>, conn: KdbConnection, table: &str) -> Rc<dyn Node> {
kdb_write(conn, table, self)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::burst;
use kdb_plus_fixed::qtype;
#[test]
fn test_kdb_serialize_trait() {
#[derive(Debug, Clone, Default)]
struct TestRecord {
sym: String,
price: f64,
size: i64,
}
impl KdbSerialize for TestRecord {
fn to_kdb_row(&self) -> K {
K::new_compound_list(vec![
K::new_symbol(self.sym.clone()),
K::new_float(self.price),
K::new_long(self.size),
])
}
}
let record = TestRecord {
sym: "AAPL".to_string(),
price: 185.50,
size: 100,
};
let row = record.to_kdb_row();
assert_eq!(row.get_type(), qtype::COMPOUND_LIST);
}
#[test]
fn symbols_with_special_chars_format_as_string_cast() {
let one = format_kdb_column_q(&[K::new_symbol("BTC-USD".to_string())]).unwrap();
assert_eq!(one, "enlist `$\"BTC-USD\"");
let many = format_kdb_column_q(&[
K::new_symbol("BTC-USD".to_string()),
K::new_symbol("BTC-100K-3D-YES".to_string()),
])
.unwrap();
assert_eq!(many, "`$(\"BTC-USD\";\"BTC-100K-3D-YES\")");
assert_eq!(
format_kdb_column_q(&[K::new_symbol("AAPL".to_string())]).unwrap(),
"enlist `$\"AAPL\""
);
}
#[test]
fn non_finite_floats_use_q_null_and_infinity() {
assert_eq!(
format_kdb_column_q(&[K::new_float(f64::NAN)]).unwrap(),
"enlist 0n"
);
assert_eq!(
format_kdb_column_q(&[K::new_float(f64::INFINITY)]).unwrap(),
"enlist 0w"
);
assert_eq!(
format_kdb_column_q(&[K::new_float(f64::NEG_INFINITY)]).unwrap(),
"enlist -0w"
);
assert_eq!(
format_kdb_column_q(&[K::new_float(1.5), K::new_float(f64::NAN)]).unwrap(),
"`float$(1.5;0n)"
);
assert_eq!(
format_kdb_column_q(&[K::new_float(2.5)]).unwrap(),
"enlist 2.5f"
);
}
#[test]
fn non_finite_reals_use_q_null_and_infinity() {
assert_eq!(
format_kdb_column_q(&[K::new_real(f32::NAN)]).unwrap(),
"enlist 0Ne"
);
assert_eq!(
format_kdb_column_q(&[K::new_real(f32::INFINITY)]).unwrap(),
"enlist 0we"
);
assert_eq!(
format_kdb_column_q(&[K::new_real(1.0), K::new_real(f32::NEG_INFINITY)]).unwrap(),
"`real$(1;-0we)"
);
}
#[test]
fn non_compound_row_is_an_error_not_a_silent_drop() {
let err =
k_rows_to_columns(vec![K::new_float(1.0)]).expect_err("a bare atom is not a valid row");
assert!(
format!("{err}").contains("must return a compound list"),
"unexpected error: {err}"
);
}
#[test]
fn ragged_burst_is_an_error_not_a_panic() {
let rows = vec![
K::new_compound_list(vec![K::new_symbol("A".to_string()), K::new_float(1.0)]),
K::new_compound_list(vec![K::new_symbol("B".to_string())]),
];
let err = k_rows_to_columns(rows).expect_err("ragged rows must error");
assert!(
format!("{err}").contains("ragged burst"),
"unexpected error: {err}"
);
}
#[test]
fn well_formed_burst_transposes_to_columns() {
let rows = vec![
K::new_compound_list(vec![K::new_symbol("A".to_string()), K::new_float(1.0)]),
K::new_compound_list(vec![K::new_symbol("B".to_string()), K::new_float(2.0)]),
];
let columns = k_rows_to_columns(rows).unwrap();
assert_eq!(columns.len(), 2); assert_eq!(columns[0].len(), 2); assert_eq!(columns[1].len(), 2);
}
#[test]
fn test_kdb_write_node_creation() {
use crate::nodes::constant;
#[derive(Debug, Clone, Default)]
struct TestTrade {
sym: String,
price: f64,
}
impl KdbSerialize for TestTrade {
fn to_kdb_row(&self) -> K {
K::new_compound_list(vec![
K::new_symbol(self.sym.clone()),
K::new_float(self.price),
])
}
}
let conn = KdbConnection::new("localhost", 5000);
let trade = TestTrade {
sym: "TEST".to_string(),
price: 100.0,
};
let stream = constant(burst![trade]);
let _node = kdb_write(conn, "test_table", &stream);
}
}