#![allow(
clippy::arbitrary_source_item_ordering,
reason = "define_data_row_copy must expand before the impl that uses DATA_ROW_COPY_SQL, \
which places the generated helper function before the impl"
)]
#![allow(
clippy::big_endian_bytes,
reason = "PostgreSQL binary COPY format specifies big-endian encoding throughout"
)]
use chrono::{DateTime, Utc};
use futures::future::try_join_all;
use uuid::Uuid;
use crate::postgres::retry::{RetryPolicy, retry};
use crate::types::error::Result;
use crate::types::model::{DataRow, VariantDataRow};
use super::PostgresRepository;
use sqlx::PgPool;
macro_rules! define_data_row_copy {
(($first_column:literal, $first_write:expr) $(, ($column:literal, $write:expr))* $(,)?) => {
const DATA_ROW_COPY_COLUMN_LIST: &[&str] = &[$first_column $(, $column)*];
const DATA_ROW_COPY_SQL: &str = concat!(
"COPY horizon_public.data_row (",
$first_column,
$(", ", $column),*,
") FROM STDIN WITH (FORMAT binary)"
);
#[allow(
clippy::single_call_fn,
reason = "Macro-generated writer keeps COPY column order and encoding in one place"
)]
fn append_data_row_copy_tuple(buf: &mut Vec<u8>, row: &VariantDataRow, now: DateTime<Utc>) {
#[allow(
clippy::as_conversions,
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
reason = "COPY column count is a small compile-time constant well within i16"
)]
let field_count = DATA_ROW_COPY_COLUMN_LIST.len() as i16;
buf.extend_from_slice(&field_count.to_be_bytes());
($first_write)(buf, row, now);
$(
($write)(buf, row, now);
)*
}
};
}
define_data_row_copy! {
("data_stream_id", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_field(buf, Some(row.data_stream_id.as_bytes()));
}),
("datetime", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_timestamptz(buf, row.datetime);
}),
("vector", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_float8_array_field(buf, row.vector.as_deref());
}),
("data_type", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_field(buf, Some(row.data_type.as_bytes()));
}),
("specification_id", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_field(buf, Some(row.specification_id.as_bytes()));
}),
("vector_start_bound", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_f64(buf, row.vector_start_bound);
}),
("vector_end_bound", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_f64(buf, row.vector_end_bound);
}),
("created_datetime", |buf: &mut Vec<u8>, row: &VariantDataRow, now: DateTime<Utc>| {
put_timestamptz(buf, row.created_datetime.unwrap_or(now));
}),
("modified_datetime", |buf: &mut Vec<u8>, _row: &VariantDataRow, now: DateTime<Utc>| {
put_timestamptz(buf, now);
}),
("variant", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_text(buf, row.variant.as_deref());
}),
("payload_int", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_i32(buf, row.payload_int);
}),
("payload_int_array", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_int4_array_field(buf, row.payload_int_array.as_deref());
}),
("payload_float32", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_f32(buf, row.payload_float32);
}),
("payload_float32_array", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_float4_array_field(buf, row.payload_float32_array.as_deref());
}),
("payload_float64", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_f64(buf, row.payload_float64);
}),
("payload_float64_array", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_float8_array_field(buf, row.payload_float64_array.as_deref());
}),
("payload_string", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_text(buf, row.payload_string.as_deref());
}),
("payload_struct", |buf: &mut Vec<u8>, row: &VariantDataRow, _now: DateTime<Utc>| {
put_optional_jsonb(buf, row.payload_struct.as_ref());
}),
}
impl PostgresRepository {
async fn binary_parallel_pass(
&self,
row_slice: &[VariantDataRow],
worker_count_option: Option<usize>,
) -> Result<u64> {
let worker_count = worker_count_option.unwrap_or(8);
let part = row_slice.len().div_ceil(worker_count).max(1);
let policy = RetryPolicy::default().with_max_retries(2);
let count_vec = try_join_all(row_slice.chunks(part).map(|chunk| {
retry(policy, move || {
Self::copy_binary_partition(&self.pool, chunk)
})
}))
.await?;
Ok(count_vec.into_iter().sum())
}
#[allow(
clippy::single_call_fn,
reason = "Used to breakup binary_parallel_pass function for readability"
)]
async fn copy_binary_partition(pool: &PgPool, row_slice: &[VariantDataRow]) -> Result<u64> {
let mut transaction = pool.begin().await?;
let copied = {
let mut copy = transaction.copy_in_raw(DATA_ROW_COPY_SQL).await?;
let mut buf: Vec<u8> = Vec::with_capacity(16 * 1024 * 1024);
buf.extend_from_slice(b"PGCOPY\n\xFF\r\n\0");
buf.extend_from_slice(&0_i32.to_be_bytes()); buf.extend_from_slice(&0_i32.to_be_bytes());
for row in row_slice {
let now = Utc::now();
append_data_row_copy_tuple(&mut buf, row, now);
if buf.len() >= 16 * 1024 * 1024 {
copy.send(buf.as_slice()).await?;
buf.clear();
}
}
buf.extend_from_slice(&(-1_i16).to_be_bytes());
copy.send(buf.as_slice()).await?;
copy.finish().await?
};
transaction.commit().await?;
Ok(copied)
}
pub async fn insert_data_row(&self, data_row: &DataRow) -> Result<DataRow> {
let variant_data_row = VariantDataRow::from(data_row);
self.insert_variant_data_row(&variant_data_row)
.await?
.try_into()
}
pub async fn insert_variant_data_row(
&self,
data_row: &VariantDataRow,
) -> Result<VariantDataRow> {
data_row.validate_payload()?;
let now = Utc::now();
Ok(sqlx::query_as!(
VariantDataRow,
r#"
INSERT INTO horizon_public.data_row
(data_stream_id, datetime, vector, data_type, specification_id,
vector_start_bound, vector_end_bound, created_datetime, modified_datetime,
variant, payload_int, payload_int_array, payload_float32,
payload_float32_array, payload_float64, payload_float64_array,
payload_string, payload_struct)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18)
RETURNING created_datetime, modified_datetime, data_stream_id,
datetime, vector, data_type, specification_id,
vector_start_bound, vector_end_bound, variant,
payload_int, payload_int_array, payload_float32, payload_float32_array,
payload_float64, payload_float64_array, payload_string, payload_struct
"#,
data_row.data_stream_id,
data_row.datetime,
data_row.vector.as_deref(),
data_row.data_type,
data_row.specification_id,
data_row.vector_start_bound,
data_row.vector_end_bound,
data_row.created_datetime.unwrap_or(now),
now,
data_row.variant,
data_row.payload_int,
data_row.payload_int_array.as_deref(),
data_row.payload_float32,
data_row.payload_float32_array.as_deref(),
data_row.payload_float64,
data_row.payload_float64_array.as_deref(),
data_row.payload_string,
data_row.payload_struct,
)
.fetch_one(&self.pool)
.await?)
}
pub async fn insert_data_row_batch(&self, row_slice: &[DataRow]) -> Result<u64> {
let variant_data_row_list: Vec<VariantDataRow> =
row_slice.iter().map(VariantDataRow::from).collect();
self.insert_variant_data_row_batch(&variant_data_row_list)
.await
}
pub async fn insert_variant_data_row_batch(&self, row_slice: &[VariantDataRow]) -> Result<u64> {
for row in row_slice {
row.validate_payload()?;
}
self.binary_parallel_pass(row_slice, None).await
}
pub async fn list_data_rows(&self, data_stream_id: Uuid) -> Result<Vec<DataRow>> {
self.list_variant_data_rows(data_stream_id)
.await?
.into_iter()
.map(DataRow::try_from)
.collect()
}
pub async fn list_variant_data_rows(
&self,
data_stream_id: Uuid,
) -> Result<Vec<VariantDataRow>> {
retry(RetryPolicy::default(), || async move {
Ok(sqlx::query_as!(
VariantDataRow,
r#"
SELECT dr.created_datetime, dr.modified_datetime, dr.data_stream_id,
dr.datetime, dr.vector, dr.data_type, dr.specification_id,
dr.vector_start_bound, dr.vector_end_bound, dr.variant,
dr.payload_int, dr.payload_int_array, dr.payload_float32,
dr.payload_float32_array, dr.payload_float64, dr.payload_float64_array,
dr.payload_string, dr.payload_struct
FROM horizon_public.data_row dr
JOIN horizon_public.data_stream ds ON dr.data_stream_id = ds.id
WHERE dr.data_stream_id = $1
ORDER BY dr.datetime
"#,
data_stream_id,
)
.fetch_all(&self.pool)
.await?)
})
.await
}
}
#[allow(
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
clippy::as_conversions,
reason = "Field lengths are bounded by PostgreSQL protocol limits, well within i32 range"
)]
fn put_field(buf: &mut Vec<u8>, bytes: Option<&[u8]>) {
match bytes {
Some(b) => {
buf.extend_from_slice(&(b.len() as i32).to_be_bytes());
buf.extend_from_slice(b);
}
None => buf.extend_from_slice(&(-1_i32).to_be_bytes()),
}
}
#[allow(
clippy::single_call_fn,
reason = "Mirrors put_float8_array_field for float4 payload columns"
)]
#[allow(
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
clippy::as_conversions,
clippy::arithmetic_side_effects,
reason = "Vector lengths are bounded by application constraints, well within i32 range"
)]
fn put_float4_array_field(buf: &mut Vec<u8>, values: &[f32]) {
const FLOAT4_OID: i32 = 700;
let array_len = (20 + values.len() * 8) as i32;
buf.extend_from_slice(&array_len.to_be_bytes());
buf.extend_from_slice(&1_i32.to_be_bytes());
buf.extend_from_slice(&0_i32.to_be_bytes());
buf.extend_from_slice(&FLOAT4_OID.to_be_bytes());
buf.extend_from_slice(&(values.len() as i32).to_be_bytes());
buf.extend_from_slice(&1_i32.to_be_bytes());
for &v in values {
buf.extend_from_slice(&4_i32.to_be_bytes());
buf.extend_from_slice(&v.to_be_bytes());
}
}
#[allow(
clippy::single_call_fn,
reason = "Re-usable function that could be used for other array fields in the future"
)]
#[allow(
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
clippy::as_conversions,
clippy::arithmetic_side_effects,
reason = "Vector lengths are bounded by application constraints, well within i32 range"
)]
fn put_float8_array_field(buf: &mut Vec<u8>, values: &[f64]) {
const FLOAT8_OID: i32 = 701;
let array_len = (20 + values.len() * 12) as i32;
buf.extend_from_slice(&array_len.to_be_bytes()); buf.extend_from_slice(&1_i32.to_be_bytes()); buf.extend_from_slice(&0_i32.to_be_bytes()); buf.extend_from_slice(&FLOAT8_OID.to_be_bytes()); buf.extend_from_slice(&(values.len() as i32).to_be_bytes()); buf.extend_from_slice(&1_i32.to_be_bytes()); for &v in values {
buf.extend_from_slice(&8_i32.to_be_bytes()); buf.extend_from_slice(&v.to_be_bytes());
}
}
#[allow(
clippy::single_call_fn,
reason = "Mirrors put_float8_array_field for int4 payload columns"
)]
#[allow(
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
clippy::as_conversions,
clippy::arithmetic_side_effects,
reason = "Vector lengths are bounded by application constraints, well within i32 range"
)]
fn put_int4_array_field(buf: &mut Vec<u8>, values: &[i32]) {
const INT4_OID: i32 = 23;
let array_len = (20 + values.len() * 8) as i32;
buf.extend_from_slice(&array_len.to_be_bytes());
buf.extend_from_slice(&1_i32.to_be_bytes());
buf.extend_from_slice(&0_i32.to_be_bytes());
buf.extend_from_slice(&INT4_OID.to_be_bytes());
buf.extend_from_slice(&(values.len() as i32).to_be_bytes());
buf.extend_from_slice(&1_i32.to_be_bytes());
for &v in values {
buf.extend_from_slice(&4_i32.to_be_bytes());
buf.extend_from_slice(&v.to_be_bytes());
}
}
#[allow(
clippy::single_call_fn,
reason = "Keeps COPY field writers consistent across payload scalar types"
)]
fn put_optional_f32(buf: &mut Vec<u8>, value: Option<f32>) {
match value {
Some(v) => put_field(buf, Some(&v.to_be_bytes())),
None => put_field(buf, None),
}
}
fn put_optional_f64(buf: &mut Vec<u8>, value: Option<f64>) {
match value {
Some(v) => put_field(buf, Some(&v.to_be_bytes())),
None => put_field(buf, None),
}
}
#[allow(
clippy::single_call_fn,
reason = "Keeps COPY field writers consistent across payload array types"
)]
fn put_optional_float4_array_field(buf: &mut Vec<u8>, values: Option<&[f32]>) {
match values {
Some(value_list) => put_float4_array_field(buf, value_list),
None => put_field(buf, None),
}
}
fn put_optional_float8_array_field(buf: &mut Vec<u8>, values: Option<&[f64]>) {
match values {
Some(value_list) => put_float8_array_field(buf, value_list),
None => put_field(buf, None),
}
}
#[allow(
clippy::single_call_fn,
reason = "Keeps COPY field writers consistent across payload scalar types"
)]
fn put_optional_i32(buf: &mut Vec<u8>, value: Option<i32>) {
match value {
Some(v) => put_field(buf, Some(&v.to_be_bytes())),
None => put_field(buf, None),
}
}
#[allow(
clippy::single_call_fn,
reason = "Keeps COPY field writers consistent across payload array types"
)]
fn put_optional_int4_array_field(buf: &mut Vec<u8>, values: Option<&[i32]>) {
match values {
Some(value_list) => put_int4_array_field(buf, value_list),
None => put_field(buf, None),
}
}
#[allow(
clippy::single_call_fn,
reason = "Keeps COPY field writers consistent across payload scalar types"
)]
fn put_optional_jsonb(buf: &mut Vec<u8>, value: Option<&serde_json::Value>) {
match value {
Some(json) => {
let text = json.to_string();
let mut bytes = Vec::with_capacity(text.len().saturating_add(1));
bytes.push(1);
bytes.extend_from_slice(text.as_bytes());
put_field(buf, Some(&bytes));
}
None => put_field(buf, None),
}
}
fn put_optional_text(buf: &mut Vec<u8>, value: Option<&str>) {
put_field(buf, value.map(str::as_bytes));
}
#[allow(
clippy::arithmetic_side_effects,
reason = "Timestamps are always after the PG epoch (2000-01-01); subtraction cannot underflow"
)]
fn put_timestamptz(buf: &mut Vec<u8>, dt: DateTime<Utc>) {
const PG_EPOCH_OFFSET_MICROS: i64 = 946_684_800_000_000;
let micros = dt.timestamp_micros() - PG_EPOCH_OFFSET_MICROS;
put_field(buf, Some(µs.to_be_bytes()));
}
#[cfg(test)]
#[allow(
clippy::indexing_slicing,
clippy::unwrap_used,
reason = "tests use indexing and unwrap for brevity"
)]
mod tests {
use super::*;
use crate::types::model::{DataRowPayload, name_to_uuid};
#[test]
fn data_row_copy_sql_matches_column_list() {
let expected = format!(
"COPY horizon_public.data_row ({}) FROM STDIN WITH (FORMAT binary)",
DATA_ROW_COPY_COLUMN_LIST.join(", ")
);
assert_eq!(DATA_ROW_COPY_SQL, expected);
assert!(!DATA_ROW_COPY_COLUMN_LIST.is_empty());
}
#[test]
fn append_data_row_copy_tuple_writes_field_count() {
let row = VariantDataRow::from_payload(
Uuid::nil(),
Utc::now(),
"btr".to_owned(),
name_to_uuid("unspecified"),
DataRowPayload::Float64(1.0),
);
let mut buf = Vec::new();
append_data_row_copy_tuple(&mut buf, &row, Utc::now());
let field_count = i16::from_be_bytes([buf[0], buf[1]]);
assert_eq!(
field_count,
i16::try_from(DATA_ROW_COPY_COLUMN_LIST.len()).unwrap()
);
}
}