#![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_EXECUTION_ENTRY, CONVERSATION_SECURITY_ENTRY, EXECUTION_MODEL_CALL,
EXECUTION_USAGE, SECURITY_ATTRIBUTION,
};
const FLEET_USAGE_SQL: &str = polyc_query_model::statements::FLEET_USAGE_SQL;
fn schema(table: polyc_projection::family::TableId) -> arrow::datatypes::SchemaRef {
use polyc_projection::family::LogicalType;
let family = if table.family_str() == CONVERSATION_EXECUTION_ENTRY.family_str() {
CONVERSATION_EXECUTION_ENTRY
} else {
CONVERSATION_SECURITY_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))
}
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 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", "conv-a",
])),
Arc::new(fixed_incarnation(4)),
Arc::new(UInt64Array::from(vec![0_u64, 1, 0, 2])),
Arc::new(StringArray::from(vec![
"turn-1", "turn-2", "turn-3", "turn-4",
])),
Arc::new(UInt64Array::from(vec![100_u64, 50, 10, 30])),
Arc::new(UInt64Array::from(vec![20_u64, 10, 5, 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-a", "conv-a",
])),
Arc::new(fixed_incarnation(4)),
Arc::new(UInt64Array::from(vec![0_u64, 1, 2, 3])),
Arc::new(StringArray::from(vec![
"turn-1", "turn-1", "turn-2", "turn-4",
])),
Arc::new(StringArray::from(vec!["p", "p", "p", "p"])),
Arc::new(StringArray::from(vec!["m", "m", "m", "m"])),
Arc::new(UInt64Array::from(vec![500_u64, 1_000, 2_000, 3_000])),
],
)
.unwrap();
ctx.register_table(
"model_call",
Arc::new(MemTable::try_new(model_call_schema, vec![vec![model_call]]).unwrap()),
)
.unwrap();
let attribution_schema = schema(SECURITY_ATTRIBUTION);
let attribution = RecordBatch::try_new(
Arc::clone(&attribution_schema),
vec![
Arc::new(StringArray::from(vec![
"conv-a", "conv-a", "conv-a", "conv-a",
])),
Arc::new(fixed_incarnation(4)),
Arc::new(UInt64Array::from(vec![0_u64, 1, 0, 0])),
Arc::new(StringArray::from(vec![
"turn-1", "turn-1", "turn-2", "turn-4",
])),
Arc::new(StringArray::from(vec![
"persona-a",
"persona-a-updated",
"persona-b",
"persona-b",
])),
Arc::new(StringArray::from(vec![
"initiator",
"initiator",
"initiator",
"initiator",
])),
],
)
.unwrap();
ctx.register_table(
"attribution",
Arc::new(MemTable::try_new(attribution_schema, vec![vec![attribution]]).unwrap()),
)
.unwrap();
ctx
}
fn with_since_ms(threshold: u64) -> String {
FLEET_USAGE_SQL.replace("$1", &threshold.to_string())
}
#[tokio::test]
async fn fleet_usage_sql_aggregates_by_conversation_and_caller_never_one_row_per_turn() {
let ctx = context();
let df = ctx.sql(&with_since_ms(0)).await.unwrap();
let batches = df.collect().await.unwrap();
let rows: usize = batches.iter().map(RecordBatch::num_rows).sum();
assert_eq!(
rows, 3,
"one row per (conversation, caller) pair, not one row per turn — \
turn-2 and turn-4 (both conv-a/persona-b) must collapse into one row"
);
assert_eq!(
batches[0]
.schema()
.fields()
.iter()
.map(|field| field.name().as_str())
.collect::<Vec<_>>(),
vec![
"partition",
"caller_persona_id",
"input_tokens",
"output_tokens",
"turn_count",
"last_active_ms",
],
"must match ProjectedRowsShape::FleetUsage's own expected-column list, in order"
);
let mut by_key: std::collections::HashMap<(String, String), (u64, u64, u64, u64)> =
std::collections::HashMap::new();
for batch in &batches {
let partition = batch
.column_by_name("partition")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let caller = batch
.column_by_name("caller_persona_id")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let input_tokens = batch
.column_by_name("input_tokens")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
let turn_count = batch
.column_by_name("turn_count")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
let last_active_ms = batch
.column_by_name("last_active_ms")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
for row in 0..batch.num_rows() {
by_key.insert(
(
partition.value(row).to_owned(),
caller.value(row).to_owned(),
),
(
input_tokens.value(row),
turn_count.value(row),
last_active_ms.value(row),
0,
),
);
}
}
let (input_tokens, turn_count, last_active_ms, _) =
by_key[&("conv-a".to_owned(), "persona-a-updated".to_owned())];
assert_eq!(input_tokens, 100, "turn-1's own usage, never doubled");
assert_eq!(turn_count, 1);
assert_eq!(
last_active_ms, 1_000,
"the LATER (higher-position) model_call clock wins"
);
let (input_tokens, turn_count, last_active_ms, _) =
by_key[&("conv-a".to_owned(), "persona-b".to_owned())];
assert_eq!(
input_tokens, 80,
"turn-2 (50) and turn-4 (30) sum into ONE row, not two"
);
assert_eq!(turn_count, 2);
assert_eq!(last_active_ms, 3_000, "the later of turn-2/turn-4's clocks");
let (input_tokens, turn_count, _, _) = by_key[&("conv-b".to_owned(), String::new())];
assert_eq!(
input_tokens, 10,
"a turn with no attribution row reads as an empty caller id — \
Control's own UNATTRIBUTED_PERSONA_ID bucket handles the mapping"
);
assert_eq!(turn_count, 1);
}
#[tokio::test]
async fn fleet_usage_sql_since_ms_excludes_earlier_turns_and_unattributed_zero_clocks() {
let ctx = context();
let df = ctx.sql(&with_since_ms(1_500)).await.unwrap();
let batches = df.collect().await.unwrap();
let rows: usize = batches.iter().map(RecordBatch::num_rows).sum();
assert_eq!(
rows, 1,
"only the conv-a/persona-b group survives the 1_500ms floor"
);
let batch = &batches[0];
let partition = batch
.column_by_name("partition")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let caller = batch
.column_by_name("caller_persona_id")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let input_tokens = batch
.column_by_name("input_tokens")
.unwrap()
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
assert_eq!(partition.value(0), "conv-a");
assert_eq!(caller.value(0), "persona-b");
assert_eq!(
input_tokens.value(0),
80,
"both turn-2 and turn-4 clear the 1_500ms floor"
);
}