use std::sync::Arc;
use arrow::array::{ArrayRef, BooleanBuilder, StringBuilder, UInt64Builder};
use arrow::datatypes::{DataType, Field, Schema, SchemaRef};
use arrow::error::ArrowError;
use arrow::record_batch::RecordBatch;
#[derive(Debug, Default)]
pub(crate) struct MessageRows(Vec<MessageRow>);
impl MessageRows {
pub(super) fn push(&mut self, projection: MessageProjection) {
self.0.push(projection.into());
}
#[cfg(test)]
pub(super) fn as_slice(&self) -> &[MessageRow] {
&self.0
}
}
#[derive(Debug)]
pub(super) struct MessageRow {
pub(super) partition: String,
pub(super) position: u64,
pub(super) turn_id: Option<String>,
pub(super) role: String,
pub(super) internal_only: bool,
pub(super) text: String,
pub(super) trust: String,
}
pub(super) struct MessageProjection {
pub(super) partition: String,
pub(super) role: String,
pub(super) internal_only: bool,
pub(super) fact: polyc_facts::TextFact,
}
impl From<MessageProjection> for MessageRow {
fn from(projection: MessageProjection) -> Self {
let MessageProjection {
partition,
role,
internal_only,
fact,
} = projection;
let polyc_facts::TextFact {
position,
turn_id,
text,
trust,
} = fact;
Self {
partition,
position,
turn_id,
role,
internal_only,
text,
trust,
}
}
}
#[must_use]
pub(crate) fn messages_schema() -> SchemaRef {
Arc::new(Schema::new(vec![
Field::new("partition", DataType::Utf8, false),
Field::new("position", DataType::UInt64, false),
Field::new("turn_id", DataType::Utf8, true),
Field::new("role", DataType::Utf8, false),
Field::new("internal_only", DataType::Boolean, false),
Field::new("text", DataType::Utf8, false),
Field::new("trust", DataType::Utf8, false),
]))
}
pub(crate) fn decode_messages_batch(rows: &MessageRows) -> Result<RecordBatch, ArrowError> {
let rows = &rows.0;
let mut partition_b = StringBuilder::with_capacity(rows.len(), rows.len() * 8);
let mut position_b = UInt64Builder::with_capacity(rows.len());
let mut turn_id_b = StringBuilder::with_capacity(rows.len(), rows.len() * 36);
let mut role_b = StringBuilder::with_capacity(rows.len(), rows.len() * 12);
let mut internal_only_b = BooleanBuilder::with_capacity(rows.len());
let mut text_b = StringBuilder::with_capacity(rows.len(), rows.len() * 64);
let mut trust_b = StringBuilder::with_capacity(rows.len(), rows.len() * 20);
for row in rows {
partition_b.append_value(&row.partition);
position_b.append_value(row.position);
match &row.turn_id {
Some(id) => turn_id_b.append_value(id),
None => turn_id_b.append_null(),
}
role_b.append_value(&row.role);
internal_only_b.append_value(row.internal_only);
text_b.append_value(&row.text);
trust_b.append_value(&row.trust);
}
let columns: Vec<ArrayRef> = vec![
Arc::new(partition_b.finish()),
Arc::new(position_b.finish()),
Arc::new(turn_id_b.finish()),
Arc::new(role_b.finish()),
Arc::new(internal_only_b.finish()),
Arc::new(text_b.finish()),
Arc::new(trust_b.finish()),
];
RecordBatch::try_new(messages_schema(), columns)
}
#[cfg(test)]
mod tests {
use arrow::array::{Array as _, StringArray};
use super::*;
#[test]
fn schema_shape_is_unchanged_by_the_split() {
let schema = messages_schema();
let expected = [
("partition", DataType::Utf8, false),
("position", DataType::UInt64, false),
("turn_id", DataType::Utf8, true),
("role", DataType::Utf8, false),
("internal_only", DataType::Boolean, false),
("text", DataType::Utf8, false),
("trust", DataType::Utf8, false),
];
assert_eq!(schema.fields().len(), expected.len());
for (field, (name, data_type, nullable)) in schema.fields().iter().zip(expected) {
assert_eq!(field.name(), name);
assert_eq!(field.data_type(), &data_type);
assert_eq!(field.is_nullable(), nullable);
}
}
#[test]
fn text_rows_build_without_tool_call_output() {
let mut rows = MessageRows::default();
rows.push(MessageProjection {
partition: "conv-a".to_string(),
role: "user".to_string(),
internal_only: true,
fact: polyc_facts::TextFact {
position: 5,
turn_id: Some("turn-a".to_string()),
text: "hello".to_string(),
trust: "trusted_user".to_string(),
},
});
let batch = decode_messages_batch(&rows).expect("message batch");
assert_eq!(batch.num_rows(), 1);
assert_eq!(batch.schema(), messages_schema());
assert_eq!(
batch
.column(0)
.as_any()
.downcast_ref::<StringArray>()
.unwrap()
.value(0),
"conv-a"
);
assert_eq!(
batch
.column(5)
.as_any()
.downcast_ref::<StringArray>()
.unwrap()
.value(0),
"hello"
);
assert!(
batch
.column(2)
.as_any()
.downcast_ref::<StringArray>()
.unwrap()
.is_valid(0)
);
assert_eq!(
batch
.column(6)
.as_any()
.downcast_ref::<StringArray>()
.unwrap()
.value(0),
"trusted_user"
);
}
}