#![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_FINANCIAL_ENTRY, FINANCIAL_OUTBOUND_PAYMENTS, FINANCIAL_PAYMENTS,
FINANCIAL_REFUSALS,
};
const PAYMENTS_SQL: &str = "SELECT partition, position, turn_id, direction, reference, \
amount_base_units, asset, recipient, method, tool_call_id, payer_kind, timestamp_unix \
FROM payments WHERE subject = $1 \
UNION ALL \
SELECT partition, position, turn_id, direction, reference, amount_base_units, asset, \
recipient, method, tool_call_id, payer_kind, timestamp_unix FROM outbound_payments \
WHERE subject = $1 \
ORDER BY timestamp_unix DESC, partition DESC, position DESC LIMIT 500";
const RECENT_REFUSALS_SQL: &str = "SELECT partition, position, turn_id, reason, reason_detail, \
merchant_host, requested_base_units, permitted_base_units, tool_call_id, timestamp_unix \
FROM refusals WHERE subject = $1 \
ORDER BY timestamp_unix DESC, partition DESC, position DESC LIMIT 500";
fn schema(table: polyc_projection::family::TableId) -> arrow::datatypes::SchemaRef {
use polyc_projection::family::LogicalType;
let declared = CONVERSATION_FINANCIAL_ENTRY
.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))
}
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()
}
#[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 payments_schema = schema(FINANCIAL_PAYMENTS);
let payments = RecordBatch::try_new(
Arc::clone(&payments_schema),
vec![
Arc::new(StringArray::from(vec!["conv-shared", "conv-shared"])),
Arc::new(fixed_incarnation(2)),
Arc::new(UInt64Array::from(vec![0_u64, 1])),
Arc::new(StringArray::from(vec!["persona-a", "persona-b"])),
Arc::new(UInt64Array::from(vec![1_000_u64, 2_000])),
Arc::new(StringArray::from(vec!["USDC", "USDC"])),
Arc::new(StringArray::from(vec!["ref-a-in", "ref-b-in"])),
Arc::new(StringArray::from(vec!["0xRa", "0xRb"])),
Arc::new(StringArray::from(vec!["tempo", "tempo"])),
Arc::new(StringArray::from(vec!["call-a-in", "call-b-in"])),
Arc::new(StringArray::from(vec!["turn-a1", "turn-b1"])),
Arc::new(StringArray::from(vec!["inbound", "inbound"])),
Arc::new(StringArray::from(vec!["linked_wallet", "linked_wallet"])),
Arc::new(UInt64Array::from(vec![1_000_u64, 2_000])),
],
)
.unwrap();
ctx.register_table(
"payments",
Arc::new(MemTable::try_new(payments_schema, vec![vec![payments]]).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-shared", "conv-shared"])),
Arc::new(fixed_incarnation(2)),
Arc::new(UInt64Array::from(vec![2_u64, 3])),
Arc::new(StringArray::from(vec!["persona-a", "persona-b"])),
Arc::new(UInt64Array::from(vec![500_u64, 700])),
Arc::new(StringArray::from(vec!["0xToken", "0xToken"])),
Arc::new(StringArray::from(vec!["ref-a-out", "ref-b-out"])),
Arc::new(StringArray::from(vec!["0xRc", "0xRd"])),
Arc::new(StringArray::from(vec!["tempo", "tempo"])),
Arc::new(StringArray::from(vec!["call-a-out", "call-b-out"])),
Arc::new(StringArray::from(vec!["turn-a2", "turn-b2"])),
Arc::new(StringArray::from(vec!["outbound", "outbound"])),
Arc::new(StringArray::from(vec!["deployment", "deployment"])),
Arc::new(UInt64Array::from(vec![3_000_u64, 4_000])),
],
)
.unwrap();
ctx.register_table(
"outbound_payments",
Arc::new(MemTable::try_new(outbound_schema, vec![vec![outbound]]).unwrap()),
)
.unwrap();
let refusals_schema = schema(FINANCIAL_REFUSALS);
let refusals = RecordBatch::try_new(
Arc::clone(&refusals_schema),
vec![
Arc::new(StringArray::from(vec!["conv-shared", "conv-shared"])),
Arc::new(fixed_incarnation(2)),
Arc::new(UInt64Array::from(vec![0_u64, 1])),
Arc::new(StringArray::from(vec!["persona-a", "persona-b"])),
Arc::new(StringArray::from(vec!["over_spend_cap", "over_spend_cap"])),
Arc::new(StringArray::from(vec!["m-a.example", "m-b.example"])),
Arc::new(UInt64Array::from(vec![500_u64, 600])),
Arc::new(UInt64Array::from(vec![100_u64, 200])),
Arc::new(StringArray::from(vec!["call-a-blocked", "call-b-blocked"])),
Arc::new(UInt64Array::from(vec![1_500_u64, 2_500])),
Arc::new(StringArray::from(vec!["turn-a3", "turn-b3"])),
Arc::new(StringArray::from(vec!["detail-a", "detail-b"])),
],
)
.unwrap();
ctx.register_table(
"refusals",
Arc::new(MemTable::try_new(refusals_schema, vec![vec![refusals]]).unwrap()),
)
.unwrap();
ctx
}
#[tokio::test]
async fn payments_sql_subject_filter_excludes_the_co_participant_and_binds_both_union_arms() {
let ctx = context();
let sql = PAYMENTS_SQL.replace("$1", "'persona-a'");
let df = ctx.sql(&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's own inbound AND outbound row — never persona-b's, \
proving `$1` filters BOTH arms of the UNION ALL"
);
let mut references = Vec::new();
for batch in &batches {
let reference = batch
.column_by_name("reference")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
for row in 0..batch.num_rows() {
references.push(reference.value(row).to_owned());
}
}
assert!(
references.contains(&"ref-a-in".to_owned()) && references.contains(&"ref-a-out".to_owned()),
"persona-a's own rows must both be present: {references:?}"
);
assert!(
!references.contains(&"ref-b-in".to_owned())
&& !references.contains(&"ref-b-out".to_owned()),
"persona-b's rows must never appear on persona-a's own read: {references:?}"
);
assert_eq!(
references,
vec!["ref-a-out".to_owned(), "ref-a-in".to_owned()]
);
}
#[tokio::test]
async fn recent_refusals_sql_subject_filter_excludes_the_co_participant() {
let ctx = context();
let sql = RECENT_REFUSALS_SQL.replace("$1", "'persona-a'");
let df = ctx.sql(&sql).await.unwrap();
let batches = df.collect().await.unwrap();
let rows: usize = batches.iter().map(RecordBatch::num_rows).sum();
assert_eq!(rows, 1, "only persona-a's own blocked attempt");
let batch = &batches[0];
let tool_call_id = batch
.column_by_name("tool_call_id")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
assert_eq!(
tool_call_id.value(0),
"call-a-blocked",
"never persona-b's blocked attempt in the same conversation"
);
}