#![allow(clippy::unwrap_used)]
use std::sync::Arc;
use arrow::array::{RecordBatch, StringArray, UInt64Array};
use datafusion::datasource::MemTable;
use datafusion::prelude::SessionContext;
use polyc_projection::family::{
CONVERSATION_CORE_ENTRY, CONVERSATION_EXECUTION_ENTRY, CONVERSATION_FINANCIAL_ENTRY,
CONVERSATION_TURNS, EXECUTION_MODEL_CALL, EXECUTION_USAGE, FINANCIAL_OUTBOUND_PAYMENTS,
FINANCIAL_PAYMENTS,
};
const CONVERSATIONS_SQL: &str = polyc_query_model::statements::DASHBOARD_CONVERSATIONS_SQL;
const SPEND_SQL: &str = polyc_query_model::statements::DASHBOARD_SPEND_SQL;
fn schema(table: polyc_projection::family::TableId) -> arrow::datatypes::SchemaRef {
use polyc_projection::family::LogicalType;
let family = if table.family_str() == CONVERSATION_CORE_ENTRY.family_str() {
CONVERSATION_CORE_ENTRY
} else if table.family_str() == CONVERSATION_EXECUTION_ENTRY.family_str() {
CONVERSATION_EXECUTION_ENTRY
} else {
CONVERSATION_FINANCIAL_ENTRY
};
let declared = family.table(table).expect("table declared");
let fields: Vec<arrow::datatypes::Field> = declared
.fields()
.iter()
.map(|field| {
let data_type = 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(), data_type, field.nullable())
})
.collect();
Arc::new(arrow::datatypes::Schema::new(fields))
}
#[allow(
clippy::too_many_lines,
reason = "one fixture batch per table; splitting it hides which columns each table carries"
)]
fn context() -> SessionContext {
let ctx = SessionContext::new();
let turns_schema = schema(CONVERSATION_TURNS);
let turns = RecordBatch::try_new(
Arc::clone(&turns_schema),
vec![
Arc::new(StringArray::from(vec!["conv-a", "conv-a", "conv-b"])),
Arc::new(fixed_incarnation(3)),
Arc::new(StringArray::from(vec!["turn-1", "turn-2", "turn-3"])),
Arc::new(UInt64Array::from(vec![0_u64, 1, 0])),
Arc::new(UInt64Array::from(vec![0_u64, 1, 0])),
Arc::new(UInt64Array::from(vec![0_u64, 1, 0])),
],
)
.unwrap();
ctx.register_table(
"turns",
Arc::new(MemTable::try_new(turns_schema, vec![vec![turns]]).unwrap()),
)
.unwrap();
let usage_schema = schema(EXECUTION_USAGE);
let usage = RecordBatch::try_new(
Arc::clone(&usage_schema),
vec![
Arc::new(StringArray::from(vec!["conv-a", "conv-a", "conv-b"])),
Arc::new(fixed_incarnation(3)),
Arc::new(UInt64Array::from(vec![0_u64, 1, 0])),
Arc::new(StringArray::from(vec!["turn-1", "turn-2", "turn-3"])),
Arc::new(UInt64Array::from(vec![100_u64, 50, 10])),
Arc::new(UInt64Array::from(vec![20_u64, 10, 5])),
],
)
.unwrap();
ctx.register_table(
"usage",
Arc::new(MemTable::try_new(usage_schema, vec![vec![usage]]).unwrap()),
)
.unwrap();
let model_call_schema = schema(EXECUTION_MODEL_CALL);
let model_call = RecordBatch::try_new(
Arc::clone(&model_call_schema),
vec![
Arc::new(StringArray::from(vec!["conv-a", "conv-a", "conv-b"])),
Arc::new(fixed_incarnation(3)),
Arc::new(UInt64Array::from(vec![0_u64, 1, 0])),
Arc::new(StringArray::from(vec!["turn-1", "turn-2", "turn-3"])),
Arc::new(StringArray::from(vec!["p", "p", "p"])),
Arc::new(StringArray::from(vec!["m", "m", "m"])),
Arc::new(UInt64Array::from(vec![1_000_u64, 2_000, 0])),
],
)
.unwrap();
ctx.register_table(
"model_call",
Arc::new(MemTable::try_new(model_call_schema, vec![vec![model_call]]).unwrap()),
)
.unwrap();
let outbound_schema = schema(FINANCIAL_OUTBOUND_PAYMENTS);
let outbound = RecordBatch::try_new(
Arc::clone(&outbound_schema),
vec![
Arc::new(StringArray::from(vec!["conv-a", "conv-a", "conv-b"])),
Arc::new(fixed_incarnation(3)),
Arc::new(UInt64Array::from(vec![0_u64, 1, 0])),
Arc::new(StringArray::from(vec![
"persona-a",
"persona-a",
"persona-a",
])),
Arc::new(UInt64Array::from(vec![1_000_u64, 500, 2_000])),
Arc::new(StringArray::from(vec!["0xToken", "0xToken", "0xToken"])),
Arc::new(StringArray::from(vec!["ref", "ref", "ref"])),
Arc::new(StringArray::from(vec!["0xR", "0xR", "0xR"])),
Arc::new(StringArray::from(vec!["tempo", "tempo", "tempo"])),
Arc::new(StringArray::from(vec!["call-1", "call-2", "call-3"])),
Arc::new(StringArray::from(vec!["", "", ""])),
Arc::new(StringArray::from(vec!["outbound", "outbound", "outbound"])),
Arc::new(StringArray::from(vec!["", "", ""])),
Arc::new(UInt64Array::from(vec![0_u64, 0, 0])),
],
)
.unwrap();
ctx.register_table(
"outbound_payments",
Arc::new(MemTable::try_new(outbound_schema, vec![vec![outbound]]).unwrap()),
)
.unwrap();
let payments_schema = schema(FINANCIAL_PAYMENTS);
let payments = RecordBatch::try_new(
Arc::clone(&payments_schema),
vec![
Arc::new(StringArray::from(vec!["conv-c"])),
Arc::new(fixed_incarnation(1)),
Arc::new(UInt64Array::from(vec![0_u64])),
Arc::new(StringArray::from(vec!["persona-b"])),
Arc::new(UInt64Array::from(vec![300_000_u64])),
Arc::new(StringArray::from(vec!["USD"])),
Arc::new(StringArray::from(vec!["ref"])),
Arc::new(StringArray::from(vec!["0xR"])),
Arc::new(StringArray::from(vec!["tempo"])),
Arc::new(StringArray::from(vec!["call-4"])),
Arc::new(StringArray::from(vec![""])),
Arc::new(StringArray::from(vec!["inbound"])),
Arc::new(StringArray::from(vec![""])),
Arc::new(UInt64Array::from(vec![0_u64])),
],
)
.unwrap();
ctx.register_table(
"payments",
Arc::new(MemTable::try_new(payments_schema, vec![vec![payments]]).unwrap()),
)
.unwrap();
ctx
}
fn fixed_incarnation(len: usize) -> arrow::array::FixedSizeBinaryArray {
let mut builder = arrow::array::FixedSizeBinaryBuilder::with_capacity(len, 32);
for _ in 0..len {
builder.append_value([0_u8; 32]).unwrap();
}
builder.finish()
}
#[tokio::test]
async fn conversations_sql_joins_every_table_without_fanout() {
let ctx = context();
let df = ctx.sql(CONVERSATIONS_SQL).await.unwrap();
let batches = df.collect().await.unwrap();
let rows: usize = batches.iter().map(RecordBatch::num_rows).sum();
assert_eq!(rows, 2, "one row per conversation with a committed turn");
let batch = &batches[0];
let ids = batch
.column_by_name("id")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let committed_turns = batch
.column_by_name("committed_turns")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
let input_tokens = batch
.column_by_name("input_tokens")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
let spend = batch
.column_by_name("spend_base_units")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
let a = ids
.iter()
.position(|id| id == Some("conv-a"))
.expect("conv-a present");
assert_eq!(
committed_turns.value(a),
2,
"conv-a has two committed turns"
);
assert_eq!(
input_tokens.value(a),
150,
"conv-a's usage sums 100 + 50 — a fanout bug would report 300 (2 turns × 150)"
);
assert_eq!(
spend.value(a),
1_500,
"conv-a's outbound spend sums 1000 + 500 across its two payments"
);
let b = ids
.iter()
.position(|id| id == Some("conv-b"))
.expect("conv-b present");
assert_eq!(committed_turns.value(b), 1);
assert_eq!(input_tokens.value(b), 10);
assert_eq!(spend.value(b), 2_000);
}
#[tokio::test]
async fn spend_sql_rolls_up_by_persona_across_conversations_and_never_sums_directions() {
let ctx = context();
let df = ctx.sql(SPEND_SQL).await.unwrap();
let batches = df.collect().await.unwrap();
let rows: usize = batches.iter().map(RecordBatch::num_rows).sum();
assert_eq!(
rows, 2,
"persona-a spent, persona-b was charged — two rollup rows"
);
let batch = &batches[0];
let personas = batch
.column_by_name("persona")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let spend = batch
.column_by_name("spend_base_units")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
let charged = batch
.column_by_name("charged_base_units")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
let conversation_count = batch
.column_by_name("conversation_count")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
let a = personas
.iter()
.position(|p| p == Some("persona-a"))
.expect("persona-a present");
assert_eq!(
spend.value(a),
3_500,
"persona-a's spend sums across BOTH its conversations"
);
assert_eq!(charged.value(a), 0, "persona-a was never charged");
assert_eq!(conversation_count.value(a), 2);
let b = personas
.iter()
.position(|p| p == Some("persona-b"))
.expect("persona-b present");
assert_eq!(spend.value(b), 0, "persona-b never spent");
assert_eq!(
charged.value(b),
300_000,
"persona-b's charge never lands in a spend figure"
);
assert_eq!(conversation_count.value(b), 1);
}
const PAYMENTS_UNION_SQL: &str = "SELECT position, turn_id, direction, reference, \
amount_base_units, asset, recipient, method, tool_call_id, subject, payer_kind, \
timestamp_unix FROM payments \
UNION ALL \
SELECT position, turn_id, direction, reference, amount_base_units, asset, recipient, \
method, tool_call_id, subject, payer_kind, timestamp_unix FROM outbound_payments \
ORDER BY position";
#[tokio::test]
async fn payments_union_sql_reads_both_tables_with_unscaled_base_units() {
let ctx = context();
let df = ctx.sql(PAYMENTS_UNION_SQL).await.unwrap();
let batches = df.collect().await.unwrap();
let rows: usize = batches.iter().map(RecordBatch::num_rows).sum();
assert_eq!(
rows, 4,
"three outbound_payments rows plus one payments row, unioned"
);
let mut by_reference_and_direction: std::collections::HashMap<(String, String), u64> =
std::collections::HashMap::new();
for batch in &batches {
let tool_call_id = batch
.column_by_name("tool_call_id")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let direction = batch
.column_by_name("direction")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let amount = batch
.column_by_name("amount_base_units")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
for row in 0..batch.num_rows() {
by_reference_and_direction.insert(
(
tool_call_id.value(row).to_owned(),
direction.value(row).to_owned(),
),
amount.value(row),
);
}
}
assert_eq!(
by_reference_and_direction[&("call-1".to_owned(), "outbound".to_owned())],
1_000,
"an outbound row's amount_base_units passes through unscaled"
);
assert_eq!(
by_reference_and_direction[&("call-2".to_owned(), "outbound".to_owned())],
500
);
assert_eq!(
by_reference_and_direction[&("call-3".to_owned(), "outbound".to_owned())],
2_000
);
assert_eq!(
by_reference_and_direction[&("call-4".to_owned(), "inbound".to_owned())],
300_000,
"an inbound row's amount_base_units is the family's normalized figure, \
not the receipt's original decimal-dollar string"
);
}