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;
const BLOCK_TYPE_CALL: &str = "call";
const BLOCK_TYPE_RESULT: &str = "result";
#[derive(Debug, Default)]
pub(crate) struct ToolCallRows(Vec<ToolCallRow>);
impl ToolCallRows {
pub(super) fn push_call(&mut self, projection: ToolCallProjection) {
self.0.push(projection.into());
}
pub(super) fn push_result(&mut self, projection: ToolResultProjection) {
self.0.push(projection.into());
}
#[cfg(test)]
pub(super) fn as_slice(&self) -> &[ToolCallRow] {
&self.0
}
}
#[derive(Debug)]
pub(super) struct ToolCallRow {
pub(super) partition: String,
pub(super) position: u64,
pub(super) turn_id: Option<String>,
pub(super) tool_call_id: String,
pub(super) block_type: &'static str,
pub(super) name: String,
pub(super) arguments: Option<String>,
pub(super) result: Option<String>,
pub(super) first_party: Option<bool>,
pub(super) internal_only: bool,
pub(super) trust: String,
}
pub(super) struct ToolCallProjection {
pub(super) partition: String,
pub(super) internal_only: bool,
pub(super) fact: polyc_facts::ToolCallFact,
}
impl From<ToolCallProjection> for ToolCallRow {
fn from(projection: ToolCallProjection) -> Self {
let ToolCallProjection {
partition,
internal_only,
fact,
} = projection;
let polyc_facts::ToolCallFact {
position,
turn_id,
tool_call_id,
name,
arguments,
trust,
} = fact;
Self {
partition,
position,
turn_id,
tool_call_id,
block_type: BLOCK_TYPE_CALL,
name,
arguments: Some(json_string(&arguments)),
result: None,
first_party: None,
internal_only,
trust,
}
}
}
pub(super) struct ToolResultProjection {
pub(super) partition: String,
pub(super) internal_only: bool,
pub(super) fact: polyc_facts::ToolResultFact,
}
impl From<ToolResultProjection> for ToolCallRow {
fn from(projection: ToolResultProjection) -> Self {
let ToolResultProjection {
partition,
internal_only,
fact,
} = projection;
let polyc_facts::ToolResultFact {
position,
turn_id,
tool_call_id,
name,
result,
first_party,
trust,
} = fact;
Self {
partition,
position,
turn_id,
tool_call_id,
block_type: BLOCK_TYPE_RESULT,
name,
arguments: None,
result: Some(json_string(&result)),
first_party: Some(first_party),
internal_only,
trust,
}
}
}
#[must_use]
pub(crate) fn tool_calls_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("tool_call_id", DataType::Utf8, false),
Field::new("block_type", DataType::Utf8, false),
Field::new("name", DataType::Utf8, false),
Field::new("arguments", DataType::Utf8, true),
Field::new("result", DataType::Utf8, true),
Field::new("first_party", DataType::Boolean, true),
Field::new("internal_only", DataType::Boolean, false),
Field::new("trust", DataType::Utf8, false),
]))
}
pub(crate) fn decode_tool_calls_batch(rows: &ToolCallRows) -> 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 tool_call_id_b = StringBuilder::with_capacity(rows.len(), rows.len() * 16);
let mut block_type_b = StringBuilder::with_capacity(rows.len(), rows.len() * 8);
let mut name_b = StringBuilder::with_capacity(rows.len(), rows.len() * 16);
let mut arguments_b = StringBuilder::with_capacity(rows.len(), rows.len() * 32);
let mut result_b = StringBuilder::with_capacity(rows.len(), rows.len() * 32);
let mut first_party_b = BooleanBuilder::with_capacity(rows.len());
let mut internal_only_b = BooleanBuilder::with_capacity(rows.len());
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(),
}
tool_call_id_b.append_value(&row.tool_call_id);
block_type_b.append_value(row.block_type);
name_b.append_value(&row.name);
match &row.arguments {
Some(args) => arguments_b.append_value(args),
None => arguments_b.append_null(),
}
match &row.result {
Some(result) => result_b.append_value(result),
None => result_b.append_null(),
}
match row.first_party {
Some(first_party) => first_party_b.append_value(first_party),
None => first_party_b.append_null(),
}
internal_only_b.append_value(row.internal_only);
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(tool_call_id_b.finish()),
Arc::new(block_type_b.finish()),
Arc::new(name_b.finish()),
Arc::new(arguments_b.finish()),
Arc::new(result_b.finish()),
Arc::new(first_party_b.finish()),
Arc::new(internal_only_b.finish()),
Arc::new(trust_b.finish()),
];
RecordBatch::try_new(tool_calls_schema(), columns)
}
fn json_string(value: &serde_json::Value) -> String {
serde_json::to_string(value).unwrap_or_else(|_| "null".to_string())
}
#[cfg(test)]
mod tests {
use arrow::array::{Array as _, BooleanArray, StringArray, UInt64Array};
use serde_json::json;
use super::*;
#[test]
fn schema_shape_is_unchanged_by_the_split() {
let schema = tool_calls_schema();
let expected = [
("partition", DataType::Utf8, false),
("position", DataType::UInt64, false),
("turn_id", DataType::Utf8, true),
("tool_call_id", DataType::Utf8, false),
("block_type", DataType::Utf8, false),
("name", DataType::Utf8, false),
("arguments", DataType::Utf8, true),
("result", DataType::Utf8, true),
("first_party", DataType::Boolean, true),
("internal_only", DataType::Boolean, 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 call_and_result_rows_build_without_message_types_or_output() {
let mut rows = ToolCallRows::default();
rows.push_call(ToolCallProjection {
partition: "conv-a".to_string(),
internal_only: true,
fact: polyc_facts::ToolCallFact {
position: 5,
turn_id: Some("turn-a".to_string()),
tool_call_id: "call-1".to_string(),
name: "search".to_string(),
arguments: json!({"q": "a"}),
trust: "trusted_user".to_string(),
},
});
rows.push_result(ToolResultProjection {
partition: "conv-a".to_string(),
internal_only: false,
fact: polyc_facts::ToolResultFact {
position: 6,
turn_id: Some("turn-a".to_string()),
tool_call_id: "call-1".to_string(),
name: "search".to_string(),
result: json!({"hits": 1}),
first_party: true,
trust: "trusted_model".to_string(),
},
});
let batch = decode_tool_calls_batch(&rows).expect("tool-call batch");
assert_eq!(batch.num_rows(), 2);
assert_eq!(batch.schema(), tool_calls_schema());
assert_identity_columns(&batch);
assert_payload_columns(&batch);
}
fn assert_identity_columns(batch: &RecordBatch) {
let partition = batch
.column(0)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!(
(partition.value(0), partition.value(1)),
("conv-a", "conv-a")
);
let position = batch
.column(1)
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
assert_eq!((position.value(0), position.value(1)), (5, 6));
let turn_id = batch
.column(2)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!((turn_id.value(0), turn_id.value(1)), ("turn-a", "turn-a"));
let call_id = batch
.column(3)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!((call_id.value(0), call_id.value(1)), ("call-1", "call-1"));
let kind = batch
.column(4)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!((kind.value(0), kind.value(1)), ("call", "result"));
let name = batch
.column(5)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!((name.value(0), name.value(1)), ("search", "search"));
}
fn assert_payload_columns(batch: &RecordBatch) {
let arguments = batch
.column(6)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!(arguments.value(0), r#"{"q":"a"}"#);
assert!(arguments.is_null(1));
let result = batch
.column(7)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert!(result.is_null(0));
assert_eq!(result.value(1), r#"{"hits":1}"#);
let first_party = batch
.column(8)
.as_any()
.downcast_ref::<BooleanArray>()
.unwrap();
assert!(first_party.is_null(0));
assert!(first_party.value(1));
let internal_only = batch
.column(9)
.as_any()
.downcast_ref::<BooleanArray>()
.unwrap();
assert!(internal_only.value(0));
assert!(!internal_only.value(1));
let trust = batch
.column(10)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!(
(trust.value(0), trust.value(1)),
("trusted_user", "trusted_model")
);
}
}