use std::sync::Arc;
use arrow::array::{ArrayRef, StringBuilder, UInt64Builder};
use arrow::datatypes::{DataType, Field, Schema, SchemaRef};
use arrow::error::ArrowError;
use arrow::record_batch::RecordBatch;
use polyc_proto::proto::polychrome::persona::v1::UsageRollup;
#[allow(unused_imports)]
use super::wallets::PersonaWalletRow;
#[derive(Debug, Clone)]
pub(crate) struct PersonaUsageRow {
pub persona_id: String,
pub committed_turns: u64,
pub input_tokens: u64,
pub output_tokens: u64,
pub last_active_ms: u64,
pub conversation_ids: String,
pub updated_at_ms: u64,
}
impl PersonaUsageRow {
#[must_use]
pub(crate) fn new(persona_id: &str, rollup: &UsageRollup) -> Self {
Self {
persona_id: persona_id.to_string(),
committed_turns: rollup.committed_turns,
input_tokens: rollup.input_tokens,
output_tokens: rollup.output_tokens,
last_active_ms: rollup.last_active_ms,
conversation_ids: serde_json::to_string(&rollup.conversation_ids)
.unwrap_or_else(|_| "[]".to_string()),
updated_at_ms: rollup.updated_at_ms,
}
}
}
#[must_use]
pub(crate) fn persona_usage_schema() -> SchemaRef {
Arc::new(Schema::new(vec![
Field::new("persona_id", DataType::Utf8, false),
Field::new("committed_turns", DataType::UInt64, false),
Field::new("input_tokens", DataType::UInt64, false),
Field::new("output_tokens", DataType::UInt64, false),
Field::new("last_active_ms", DataType::UInt64, false),
Field::new("conversation_ids", DataType::Utf8, false),
Field::new("updated_at_ms", DataType::UInt64, false),
]))
}
pub(crate) fn decode_persona_usage_batch(
rows: &[PersonaUsageRow],
) -> Result<RecordBatch, ArrowError> {
let mut persona_id_b = StringBuilder::with_capacity(rows.len(), rows.len() * 36);
let mut committed_turns_b = UInt64Builder::with_capacity(rows.len());
let mut input_tokens_b = UInt64Builder::with_capacity(rows.len());
let mut output_tokens_b = UInt64Builder::with_capacity(rows.len());
let mut last_active_ms_b = UInt64Builder::with_capacity(rows.len());
let mut conversation_ids_b = StringBuilder::with_capacity(rows.len(), rows.len() * 32);
let mut updated_at_ms_b = UInt64Builder::with_capacity(rows.len());
for row in rows {
persona_id_b.append_value(&row.persona_id);
committed_turns_b.append_value(row.committed_turns);
input_tokens_b.append_value(row.input_tokens);
output_tokens_b.append_value(row.output_tokens);
last_active_ms_b.append_value(row.last_active_ms);
conversation_ids_b.append_value(&row.conversation_ids);
updated_at_ms_b.append_value(row.updated_at_ms);
}
let columns: Vec<ArrayRef> = vec![
Arc::new(persona_id_b.finish()),
Arc::new(committed_turns_b.finish()),
Arc::new(input_tokens_b.finish()),
Arc::new(output_tokens_b.finish()),
Arc::new(last_active_ms_b.finish()),
Arc::new(conversation_ids_b.finish()),
Arc::new(updated_at_ms_b.finish()),
];
RecordBatch::try_new(persona_usage_schema(), columns)
}
#[cfg(test)]
mod tests {
use arrow::array::Array as _;
use super::*;
#[test]
fn persona_usage_round_trips_to_its_row() {
let rollup = UsageRollup {
persona_id: "persona-1".to_string(),
committed_turns: 7,
input_tokens: 100,
output_tokens: 200,
last_active_ms: 9_000,
conversation_ids: vec!["conv-a".to_string(), "conv-b".to_string()],
updated_at_ms: 9_500,
..Default::default()
};
let row = PersonaUsageRow::new("persona-1", &rollup);
assert_eq!(row.conversation_ids, r#"["conv-a","conv-b"]"#);
let batch = decode_persona_usage_batch(&[row]).expect("batch build");
assert_eq!(batch.num_rows(), 1);
assert_eq!(batch.schema(), persona_usage_schema());
let committed_turns = batch
.column(1)
.as_any()
.downcast_ref::<arrow::array::UInt64Array>()
.unwrap()
.value(0);
assert_eq!(committed_turns, 7);
}
#[test]
fn persona_usage_schema_shape() {
let schema = persona_usage_schema();
let names: Vec<&str> = schema.fields().iter().map(|f| f.name().as_str()).collect();
assert_eq!(
names,
vec![
"persona_id",
"committed_turns",
"input_tokens",
"output_tokens",
"last_active_ms",
"conversation_ids",
"updated_at_ms",
]
);
for field in schema.fields() {
assert!(!field.is_nullable(), "{} must be non-null", field.name());
}
}
}