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;
use polyc_eventlog::Event;
use polyc_proto::events_decode::try_decode_event_payload;
use polyc_proto::kinds;
use polyc_proto::proto::polychrome::agent::v1::Message;
const BLOCK_TYPE_CALL: &str = "call";
const BLOCK_TYPE_RESULT: &str = "result";
const MESSAGE_CONTENT_KIND_BASES: &[&str] = &[kinds::USER_MSG, kinds::OUTPUT_MSG];
#[derive(Debug, Clone)]
pub(crate) struct MessageRow {
pub partition: String,
pub position: u64,
pub turn_id: Option<String>,
pub role: String,
pub internal_only: bool,
pub text: String,
pub trust: String,
}
#[derive(Debug, Clone)]
pub(crate) struct ToolCallRow {
pub partition: String,
pub position: u64,
pub turn_id: Option<String>,
pub tool_call_id: String,
pub block_type: &'static str,
pub name: String,
pub arguments: Option<String>,
pub result: Option<String>,
pub first_party: Option<bool>,
pub internal_only: bool,
pub trust: String,
}
#[derive(Debug, Clone, Default)]
pub(crate) struct MessageContentRows {
pub messages: Vec<MessageRow>,
pub tool_calls: Vec<ToolCallRow>,
}
#[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),
]))
}
#[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_messages_batch(rows: &[MessageRow]) -> Result<RecordBatch, ArrowError> {
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)
}
pub(crate) fn decode_tool_calls_batch(rows: &[ToolCallRow]) -> Result<RecordBatch, ArrowError> {
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())
}
#[must_use]
pub(crate) fn decode_message_content_events(
partition: &str,
events: &[(u64, Event)],
) -> MessageContentRows {
let mut out = MessageContentRows::default();
for (position, event) in events {
let (base, turn_uuid) = kinds::parse(&event.kind);
if !MESSAGE_CONTENT_KIND_BASES.contains(&base) {
continue;
}
let turn_id = turn_uuid.map(|id| id.to_string());
let message = match try_decode_event_payload::<Message>(&event.payload) {
Ok(message) => message,
Err(e) => {
if !event.payload.is_empty() {
tracing::warn!(
error = %e,
len = event.payload.len(),
table = "messages/tool_calls",
"corrupt message payload; skipping row"
);
}
continue;
}
};
let internal_only = message.internal_only;
let trust = event.trust.as_str();
let folded =
polyc_facts::fold_message_content(&message, *position, turn_id.as_deref(), trust);
for warning in folded.warnings {
tracing::warn!(table = "tool_calls", %warning, "message-content fold warning");
}
match folded.content {
polyc_facts::MessageContent::Text(t) => {
out.messages.push(MessageRow {
partition: partition.to_string(),
position: *position,
turn_id,
role: message.role,
internal_only,
text: t.text,
trust: t.trust,
});
}
polyc_facts::MessageContent::ToolCall(call) => {
out.tool_calls.push(ToolCallRow {
partition: partition.to_string(),
position: *position,
turn_id,
tool_call_id: call.tool_call_id,
block_type: BLOCK_TYPE_CALL,
name: call.name,
arguments: Some(json_string(&call.arguments)),
result: None,
first_party: None,
internal_only,
trust: call.trust,
});
}
polyc_facts::MessageContent::ToolResult(res) => {
out.tool_calls.push(ToolCallRow {
partition: partition.to_string(),
position: *position,
turn_id,
tool_call_id: res.tool_call_id,
block_type: BLOCK_TYPE_RESULT,
name: res.name,
arguments: None,
result: Some(json_string(&res.result)),
first_party: Some(res.first_party),
internal_only,
trust: res.trust,
});
}
polyc_facts::MessageContent::None => {}
}
}
out
}
#[cfg(test)]
mod tests {
use arrow::array::Array as _;
use buffa::Message as _;
use polyc_proto::proto::polychrome::agent::v1::{
Content, FunctionCallContent, FunctionResultContent, TextContent, ToolCallContent,
ToolResultContent, content, function_result_content, tool_call_content,
tool_result_content,
};
use serde_json::json;
use uuid::Uuid;
use super::*;
fn text_message(role: &str, text: &str, internal_only: bool) -> Message {
Message {
role: role.to_string(),
content: buffa::MessageField::some(Content {
r#type: Some(content::Type::Text(Box::new(TextContent {
text: text.to_string(),
..Default::default()
}))),
..Default::default()
}),
internal_only,
__buffa_unknown_fields: buffa::UnknownFields::default(),
}
}
fn tool_call_message(id: &str, name: &str, args: serde_json::Value) -> Message {
let arguments = serde_json::from_value::<buffa_types::google::protobuf::Struct>(args)
.map(buffa::MessageField::some)
.unwrap_or_default();
Message {
role: "model".to_string(),
content: buffa::MessageField::some(Content {
r#type: Some(content::Type::ToolCall(Box::new(ToolCallContent {
id: id.to_string(),
r#type: Some(tool_call_content::Type::FunctionCall(Box::new(
FunctionCallContent {
name: name.to_string(),
arguments,
..Default::default()
},
))),
..Default::default()
}))),
..Default::default()
}),
internal_only: false,
__buffa_unknown_fields: buffa::UnknownFields::default(),
}
}
fn tool_result_message(
id: &str,
name: &str,
result: serde_json::Value,
first_party: bool,
) -> Message {
let response = serde_json::from_value::<buffa_types::google::protobuf::Struct>(result)
.ok()
.map(|s| function_result_content::Result::Response(Box::new(s)));
Message {
role: "tool".to_string(),
content: buffa::MessageField::some(Content {
r#type: Some(content::Type::ToolResult(Box::new(ToolResultContent {
call_id: id.to_string(),
first_party,
r#type: Some(tool_result_content::Type::FunctionResult(Box::new(
FunctionResultContent {
name: name.to_string(),
result: response,
..Default::default()
},
))),
..Default::default()
}))),
..Default::default()
}),
internal_only: false,
__buffa_unknown_fields: buffa::UnknownFields::default(),
}
}
#[test]
fn messages_schema_shape() {
let schema = messages_schema();
let names: Vec<&str> = schema.fields().iter().map(|f| f.name().as_str()).collect();
assert_eq!(
names,
vec![
"partition",
"position",
"turn_id",
"role",
"internal_only",
"text",
"trust",
]
);
let expect = [
("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),
];
for (field, (name, ty, nullable)) in schema.fields().iter().zip(expect) {
assert_eq!(field.name(), name);
assert_eq!(field.data_type(), &ty);
assert_eq!(field.is_nullable(), nullable);
}
}
#[test]
fn tool_calls_schema_shape() {
let schema = tool_calls_schema();
let names: Vec<&str> = schema.fields().iter().map(|f| f.name().as_str()).collect();
assert_eq!(
names,
vec![
"partition",
"position",
"turn_id",
"tool_call_id",
"block_type",
"name",
"arguments",
"result",
"first_party",
"internal_only",
"trust",
]
);
let expect = [
("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),
];
for (field, (name, ty, nullable)) in schema.fields().iter().zip(expect) {
assert_eq!(field.name(), name);
assert_eq!(field.data_type(), &ty);
assert_eq!(field.is_nullable(), nullable);
}
}
#[test]
fn decode_messages_batch_round_trips() {
let rows = vec![MessageRow {
partition: "conv-a".to_string(),
position: 5,
turn_id: Some("turn-xyz".to_string()),
role: "user".to_string(),
internal_only: false,
text: "hello".to_string(),
trust: "trusted_user".to_string(),
}];
let batch = decode_messages_batch(&rows).expect("batch build");
assert_eq!(batch.num_rows(), 1);
assert_eq!(batch.schema(), messages_schema());
let text = batch
.column(5)
.as_any()
.downcast_ref::<arrow::array::StringArray>()
.unwrap();
assert_eq!(text.value(0), "hello");
}
#[test]
fn decode_tool_calls_batch_round_trips_a_call_and_a_result() {
let rows = vec![
ToolCallRow {
partition: "conv-a".to_string(),
position: 5,
turn_id: Some("turn-xyz".to_string()),
tool_call_id: "call-1".to_string(),
block_type: BLOCK_TYPE_CALL,
name: "search".to_string(),
arguments: Some(r#"{"q":"a"}"#.to_string()),
result: None,
first_party: None,
internal_only: false,
trust: "trusted_user".to_string(),
},
ToolCallRow {
partition: "conv-a".to_string(),
position: 6,
turn_id: Some("turn-xyz".to_string()),
tool_call_id: "call-1".to_string(),
block_type: BLOCK_TYPE_RESULT,
name: "search".to_string(),
arguments: None,
result: Some(r#"{"hits":1}"#.to_string()),
first_party: Some(true),
internal_only: false,
trust: "trusted_user".to_string(),
},
];
let batch = decode_tool_calls_batch(&rows).expect("batch build");
assert_eq!(batch.num_rows(), 2);
assert_eq!(batch.schema(), tool_calls_schema());
let block_type = batch
.column(4)
.as_any()
.downcast_ref::<arrow::array::StringArray>()
.unwrap();
assert_eq!(block_type.value(0), "call");
assert_eq!(block_type.value(1), "result");
let arguments = batch
.column(6)
.as_any()
.downcast_ref::<arrow::array::StringArray>()
.unwrap();
assert_eq!(arguments.value(0), r#"{"q":"a"}"#);
assert!(arguments.is_null(1));
let result = batch
.column(7)
.as_any()
.downcast_ref::<arrow::array::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::<arrow::array::BooleanArray>()
.unwrap();
assert!(first_party.is_null(0));
assert!(first_party.value(1));
}
#[test]
fn text_message_decodes_to_a_messages_row_with_role_and_internal_only() {
let turn = Uuid::from_u128(0x0195_abcd_ef01_2345_6789_abcd_ef01_2345);
let message = text_message("user", "hi there", false);
let events = vec![(
1,
Event::trusted(
kinds::tagged(kinds::USER_MSG, &turn),
message.encode_to_vec(),
),
)];
let rows = decode_message_content_events("conv-real", &events);
assert_eq!(rows.messages.len(), 1);
assert!(rows.tool_calls.is_empty());
let row = &rows.messages[0];
assert_eq!(row.partition, "conv-real");
assert_eq!(row.position, 1);
assert_eq!(row.turn_id, Some(turn.to_string()));
assert_eq!(row.role, "user");
assert!(!row.internal_only);
assert_eq!(row.text, "hi there");
assert_eq!(row.trust, polyc_eventlog::TrustTag::TrustedUser.as_str());
}
#[test]
fn internal_only_flag_is_carried_onto_the_row() {
let message = text_message("model", "ground truth note", true);
let events = vec![(1, Event::new(kinds::OUTPUT_MSG, message.encode_to_vec()))];
let rows = decode_message_content_events("conv-internal", &events);
assert_eq!(rows.messages.len(), 1);
assert!(rows.messages[0].internal_only);
}
#[test]
fn tool_call_message_decodes_to_a_tool_calls_row() {
let turn = Uuid::from_u128(0x0195_abcd_ef01_2345_6789_abcd_ef01_9999);
let message = tool_call_message("call-1", "search", json!({"q": "a"}));
let events = vec![(
2,
Event::new(
kinds::tagged(kinds::OUTPUT_MSG, &turn),
message.encode_to_vec(),
),
)];
let rows = decode_message_content_events("conv-real", &events);
assert!(rows.messages.is_empty());
assert_eq!(rows.tool_calls.len(), 1);
let row = &rows.tool_calls[0];
assert_eq!(row.tool_call_id, "call-1");
assert_eq!(row.block_type, "call");
assert_eq!(row.name, "search");
assert_eq!(row.arguments.as_deref(), Some(r#"{"q":"a"}"#));
assert_eq!(row.result, None);
assert_eq!(row.first_party, None);
}
#[test]
fn tool_result_message_decodes_to_a_tool_calls_row_paired_by_id() {
let turn = Uuid::from_u128(0x0195_abcd_ef01_2345_6789_abcd_ef01_9999);
let call = tool_call_message("call-1", "search", json!({"q": "a"}));
let result = tool_result_message("call-1", "search", json!({"hits": 1}), true);
let events = vec![
(
2,
Event::new(
kinds::tagged(kinds::OUTPUT_MSG, &turn),
call.encode_to_vec(),
),
),
(
3,
Event::new(
kinds::tagged(kinds::USER_MSG, &turn),
result.encode_to_vec(),
),
),
];
let rows = decode_message_content_events("conv-paired", &events);
assert_eq!(rows.tool_calls.len(), 2);
assert_eq!(rows.tool_calls[0].tool_call_id, "call-1");
assert_eq!(rows.tool_calls[0].block_type, "call");
assert_eq!(rows.tool_calls[1].tool_call_id, "call-1");
assert_eq!(rows.tool_calls[1].block_type, "result");
assert_eq!(
rows.tool_calls[1].result.as_deref(),
Some(r#"{"hits":1.0}"#)
);
assert_eq!(rows.tool_calls[1].first_party, Some(true));
}
#[test]
fn empty_payload_folds_to_no_rows() {
let events = vec![(1, Event::new(kinds::USER_MSG, Vec::new()))];
let rows = decode_message_content_events("conv-empty", &events);
assert!(rows.messages.is_empty());
assert!(rows.tool_calls.is_empty());
}
#[test]
fn undecodable_non_empty_payload_is_skipped() {
let events = vec![
(1, Event::new(kinds::USER_MSG, vec![0xFF, 0xFE, 0xFD])),
(2, Event::new(kinds::USER_MSG, Vec::new())),
];
let rows = decode_message_content_events("conv-corrupt", &events);
assert!(rows.messages.is_empty());
assert!(rows.tool_calls.is_empty());
}
#[test]
fn unrelated_kind_is_not_decoded() {
let message = text_message("user", "hi", false);
let events = vec![(1, Event::new(kinds::USAGE, message.encode_to_vec()))];
let rows = decode_message_content_events("conv-unrelated", &events);
assert!(rows.messages.is_empty());
assert!(rows.tool_calls.is_empty());
}
}