#![allow(clippy::unwrap_used)]
use std::collections::BTreeMap;
use std::sync::Arc;
use arrow::array::{
ArrayRef, BooleanArray, FixedSizeBinaryBuilder, RecordBatch, StringArray, UInt64Array,
};
use datafusion::datasource::MemTable;
use datafusion::prelude::SessionContext;
use polyc_projection::family::{
CONVERSATION_TRACE_ENTRY, MEMORY_FACTS, MEMORY_INVALIDATIONS, MEMORY_PORTABLE_FACTS,
MEMORY_PORTABLE_INVALIDATIONS, OBSERVED_ROUTINES_ENTRY, OBSERVED_ROUTINES_TABLE,
PERSONA_MEMORY_ENTRY, TRACE_MESSAGES, TRACE_PARTICIPANT_ADDRESSES, TRACE_STEPS, TRACE_TURNS,
TableId,
};
use serde_json::Value;
const TRACE_SQL: &str = polyc_query_model::statements::COMPOSITE_TRACE_SQL;
const MEMORY_SQL: &str = polyc_query_model::statements::COMPOSITE_TRACE_MEMORY_SQL;
const ROUTINES_SQL: &str = polyc_query_model::statements::COMPOSITE_TRACE_ROUTINES_SQL;
const ADDRESSES_SQL: &str = polyc_query_model::statements::COMPOSITE_TRACE_ADDRESSES_SQL;
type Row<'a> = &'a [(&'a str, Value)];
fn schema(table: TableId) -> arrow::datatypes::SchemaRef {
use polyc_projection::family::LogicalType;
let declared = [
CONVERSATION_TRACE_ENTRY,
PERSONA_MEMORY_ENTRY,
OBSERVED_ROUTINES_ENTRY,
]
.into_iter()
.find_map(|family| family.table(table))
.expect("table declared");
Arc::new(arrow::datatypes::Schema::new(
declared
.fields()
.iter()
.map(|field| {
let ty = match field.logical_type() {
LogicalType::Utf8 => arrow::datatypes::DataType::Utf8,
LogicalType::FixedBytes { len } => {
arrow::datatypes::DataType::FixedSizeBinary(i32::try_from(len).unwrap())
}
LogicalType::UInt64 => arrow::datatypes::DataType::UInt64,
LogicalType::Boolean => arrow::datatypes::DataType::Boolean,
};
arrow::datatypes::Field::new(field.name(), ty, field.nullable())
})
.collect::<Vec<_>>(),
))
}
fn batch(table: TableId, rows: &[Row<'_>]) -> RecordBatch {
use arrow::datatypes::DataType;
let schema = schema(table);
let rows = rows
.iter()
.map(|row| row.iter().cloned().collect::<BTreeMap<_, _>>())
.collect::<Vec<_>>();
let columns = schema
.fields()
.iter()
.map(|field| -> ArrayRef {
match field.data_type() {
DataType::Utf8 => Arc::new(StringArray::from(
rows.iter()
.map(|row| {
row.get(field.name().as_str())
.and_then(Value::as_str)
.unwrap_or("")
})
.collect::<Vec<_>>(),
)),
DataType::UInt64 => Arc::new(UInt64Array::from(
rows.iter()
.map(|row| {
row.get(field.name().as_str())
.and_then(Value::as_u64)
.unwrap_or(0)
})
.collect::<Vec<_>>(),
)),
DataType::Boolean => Arc::new(BooleanArray::from(
rows.iter()
.map(|row| {
row.get(field.name().as_str())
.and_then(Value::as_bool)
.unwrap_or(false)
})
.collect::<Vec<_>>(),
)),
DataType::FixedSizeBinary(size) => {
let mut builder = FixedSizeBinaryBuilder::with_capacity(rows.len(), *size);
for _ in &rows {
builder
.append_value(vec![7_u8; usize::try_from(*size).unwrap()])
.unwrap();
}
Arc::new(builder.finish())
}
other => panic!("unsupported test type {other:?}"),
}
})
.collect();
RecordBatch::try_new(schema, columns).unwrap()
}
fn register(ctx: &SessionContext, table: TableId, rows: &[Row<'_>]) {
let schema = schema(table);
let batches = if rows.is_empty() {
vec![RecordBatch::new_empty(Arc::clone(&schema))]
} else {
vec![batch(table, rows)]
};
ctx.register_table(
table.as_str(),
Arc::new(MemTable::try_new(schema, vec![batches]).unwrap()),
)
.unwrap();
}
fn register_empty_trace_tables(ctx: &SessionContext, populated: &[TableId]) {
for declared in CONVERSATION_TRACE_ENTRY.tables() {
let table = declared.table();
if !populated.contains(&table) {
register(ctx, table, &[]);
}
}
}
fn bind(sql: &str, first: &str, second: Option<&str>, third: Option<&str>) -> String {
let sql = sql.replace("$1", &format!("'{first}'"));
let sql = second.map_or_else(
|| sql.clone(),
|value| sql.replace("$2", &format!("'{value}'")),
);
third.map_or_else(
|| sql.clone(),
|value| sql.replace("$3", &format!("'{value}'")),
)
}
#[tokio::test]
async fn trace_statement_returns_one_row_per_step_in_turn_order() {
let ctx = SessionContext::new();
register(
&ctx,
TRACE_TURNS,
&[
&[
("partition", Value::from("conv-x")),
("turn_id", Value::from("turn-b")),
("first_position", Value::from(20)),
("status", Value::from("complete")),
("caller_persona_id", Value::from("persona-x")),
],
&[
("partition", Value::from("conv-x")),
("turn_id", Value::from("turn-a")),
("first_position", Value::from(10)),
("status", Value::from("complete")),
("caller_persona_id", Value::from("persona-y")),
],
],
);
register(
&ctx,
TRACE_STEPS,
&[
&[
("partition", Value::from("conv-x")),
("position", Value::from(21)),
("ordinal", Value::from(0)),
("turn_id", Value::from("turn-b")),
("step_kind", Value::from("message")),
],
&[
("partition", Value::from("conv-x")),
("position", Value::from(11)),
("ordinal", Value::from(0)),
("turn_id", Value::from("turn-a")),
("step_kind", Value::from("message")),
],
],
);
register(
&ctx,
TRACE_MESSAGES,
&[
&[
("partition", Value::from("conv-x")),
("position", Value::from(21)),
("ordinal", Value::from(0)),
("turn_id", Value::from("turn-b")),
("role", Value::from("output")),
("text", Value::from("second")),
],
&[
("partition", Value::from("conv-x")),
("position", Value::from(11)),
("ordinal", Value::from(0)),
("turn_id", Value::from("turn-a")),
("role", Value::from("input")),
("text", Value::from("first")),
],
],
);
register_empty_trace_tables(&ctx, &[TRACE_TURNS, TRACE_STEPS, TRACE_MESSAGES]);
let batches = ctx
.sql(&bind(TRACE_SQL, "conv-x", None, None))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::<usize>(), 2);
let turn_ids = batches[0]
.column_by_name("turn_id")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!(turn_ids.value(0), "turn-a");
assert_eq!(turn_ids.value(1), "turn-b");
}
#[tokio::test]
async fn memory_statement_combines_owner_rows_and_only_portable_participant_rows() {
let ctx = SessionContext::new();
register(
&ctx,
MEMORY_FACTS,
&[
&[
("partition", Value::from("persona-x-mem")),
("position", Value::from(1)),
("fact_id", Value::from("own-direct")),
("text", Value::from("owned")),
("provenance_conversation_id", Value::from("x")),
("provenance_turn_id", Value::from("turn-a")),
("scope", Value::from("direct")),
],
&[
("partition", Value::from("persona-y-mem")),
("position", Value::from(2)),
("fact_id", Value::from("foreign-direct")),
("provenance_conversation_id", Value::from("x")),
("provenance_turn_id", Value::from("turn-a")),
("scope", Value::from("direct")),
],
],
);
register(
&ctx,
MEMORY_PORTABLE_FACTS,
&[
&[
("partition", Value::from("persona-y-mem")),
("position", Value::from(3)),
("fact_id", Value::from("participant-portable")),
("text", Value::from("portable")),
("provenance_conversation_id", Value::from("x")),
("provenance_turn_id", Value::from("turn-a")),
("scope", Value::from("portable")),
],
&[
("partition", Value::from("persona-z-mem")),
("position", Value::from(4)),
("fact_id", Value::from("other-conversation")),
("provenance_conversation_id", Value::from("z")),
("provenance_turn_id", Value::from("turn-z")),
("scope", Value::from("portable")),
],
],
);
register(&ctx, MEMORY_INVALIDATIONS, &[]);
register(&ctx, MEMORY_PORTABLE_INVALIDATIONS, &[]);
register(
&ctx,
TRACE_TURNS,
&[&[
("partition", Value::from("conv-x")),
("turn_id", Value::from("turn-a")),
("first_position", Value::from(1)),
("status", Value::from("complete")),
]],
);
let batches = ctx
.sql(&bind(
MEMORY_SQL,
"persona-x-mem",
Some("x"),
Some("conv-x"),
))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::<usize>(), 2);
let ids = batches[0]
.column_by_name("fact_id")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!(
(ids.value(0), ids.value(1)),
("own-direct", "participant-portable")
);
}
#[tokio::test]
async fn routine_statement_uses_the_matching_published_observation_rows() {
let ctx = SessionContext::new();
register(
&ctx,
OBSERVED_ROUTINES_TABLE,
&[
&[
("partition", Value::from("observed-routines-team")),
("position", Value::from(8)),
("observation_ordinal", Value::from(8)),
("namespace", Value::from("team")),
("name", Value::from("matching")),
("fire_conversation_id", Value::from("x")),
("phase", Value::from("ready")),
],
&[
("partition", Value::from("observed-routines-team")),
("position", Value::from(8)),
("observation_ordinal", Value::from(8)),
("namespace", Value::from("team")),
("name", Value::from("unrelated")),
("fire_conversation_id", Value::from("z")),
],
],
);
let batches = ctx
.sql(&bind(ROUTINES_SQL, "x", None, None))
.await
.unwrap()
.collect()
.await
.unwrap();
assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::<usize>(), 1);
let names = batches[0]
.column_by_name("name")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!(names.value(0), "matching");
}
#[tokio::test]
async fn addresses_statement_returns_one_row_per_address_step() {
let ctx = SessionContext::new();
register(
&ctx,
TRACE_PARTICIPANT_ADDRESSES,
&[
&[
("partition", Value::from("conv-x")),
("position", Value::from(20)),
("ordinal", Value::from(0)),
("step_kind", Value::from("payment_receipt")),
("address", Value::from("0xpayer")),
],
&[
("partition", Value::from("conv-y")),
("position", Value::from(5)),
("ordinal", Value::from(0)),
("step_kind", Value::from("payment_receipt")),
("address", Value::from("0xother")),
],
&[
("partition", Value::from("conv-x")),
("position", Value::from(10)),
("ordinal", Value::from(0)),
("step_kind", Value::from("wallet_link")),
("address", Value::from("0xwallet")),
],
],
);
let batches = ctx
.sql(&bind(ADDRESSES_SQL, "conv-x", None, None))
.await
.unwrap()
.collect()
.await
.unwrap();
let mut rows = Vec::new();
for batch in &batches {
let column = |name: &str| batch.column_by_name(name).unwrap().clone();
let position = column("position");
let position = position.as_any().downcast_ref::<UInt64Array>().unwrap();
let kind = column("step_kind");
let kind = kind.as_any().downcast_ref::<StringArray>().unwrap();
let address = column("address");
let address = address.as_any().downcast_ref::<StringArray>().unwrap();
for index in 0..batch.num_rows() {
rows.push((
position.value(index),
kind.value(index).to_owned(),
address.value(index).to_owned(),
));
}
}
assert_eq!(
rows,
vec![
(10, "wallet_link".to_owned(), "0xwallet".to_owned()),
(20, "payment_receipt".to_owned(), "0xpayer".to_owned()),
],
"one row per address step of the named conversation, in step order"
);
}