use foundation_arrow::arrow_array::RecordBatch;
use foundation_arrow::{FromArrow, ToArrow};
use foundation_macros::{ArrowSchema, FromArrow, ToArrow};
use serde::{Deserialize, Serialize};
use crate::types::{Messages, ModelOutput, SessionRecord};
#[derive(Debug)]
pub enum SerError {
Json(String),
Arrow(String),
}
impl core::fmt::Display for SerError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
SerError::Json(m) => write!(f, "json serialization: {m}"),
SerError::Arrow(m) => write!(f, "arrow serialization: {m}"),
}
}
}
impl std::error::Error for SerError {}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, ArrowSchema, ToArrow, FromArrow)]
pub struct SessionRecordRow {
pub id: String,
pub record_type: String,
pub role: String,
pub title: Option<String>,
pub summary: Option<String>,
pub model: Option<String>,
pub input_tokens: Option<u32>,
pub output_tokens: Option<u32>,
pub created_at: u64,
pub content: String,
}
impl SessionRecordRow {
pub fn from_record(record: &SessionRecord) -> Result<Self, SerError> {
let content = serde_json::to_string(record).map_err(|e| SerError::Json(e.to_string()))?;
let mut row = Self {
id: String::new(),
record_type: record_type_of(record).to_string(),
role: String::new(),
title: None,
summary: None,
model: None,
input_tokens: None,
output_tokens: None,
created_at: 0,
content,
};
match record {
SessionRecord::Conversation { message } => {
row.id = message.id().to_string();
row.created_at = message.id().timestamp();
match message {
Messages::User { role, .. } => row.role = role.as_wire().to_string(),
Messages::ToolResult { .. } => row.role = "tool".to_string(),
Messages::Assistant {
model,
usage,
content,
..
} => {
row.role = "assistant".to_string();
row.model = Some(model.to_string());
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
{
row.input_tokens = Some(usage.input.max(0.0).round() as u32);
row.output_tokens = Some(usage.output.max(0.0).round() as u32);
}
if let ModelOutput::ToolCall { name, .. } = content {
row.title = Some(name.clone());
}
}
}
}
SessionRecord::WorkingMemory { facts, .. } => {
row.summary = facts.first().map(|f| f.fact.clone());
row.title = Some(format!("{} fact(s)", facts.len()));
}
SessionRecord::Observation { observations, .. } => {
row.summary = observations.first().map(|o| o.content.clone());
row.title = Some(format!("{} observation(s)", observations.len()));
}
SessionRecord::Reflection { reflections, .. } => {
row.summary = reflections.first().map(|r| r.summary.clone());
row.title = Some(format!("{} reflection(s)", reflections.len()));
}
SessionRecord::FailedAction { error, .. } => {
row.summary = Some(error.to_string());
row.title = Some("failed_action".to_string());
}
SessionRecord::Summary {
message_count,
usage,
} => {
row.title = Some(format!("summary ({message_count} msgs)"));
#[allow(clippy::cast_possible_truncation)]
{
row.input_tokens = Some(usage.input.min(u64::from(u32::MAX)) as u32);
row.output_tokens = Some(usage.output.min(u64::from(u32::MAX)) as u32);
}
}
SessionRecord::Retracted { id, reason, .. } => {
row.id = id.to_string();
row.title = Some("retracted".to_string());
row.summary = Some(reason.clone());
}
}
Ok(row)
}
pub fn into_record(self) -> Result<SessionRecord, SerError> {
serde_json::from_str(&self.content).map_err(|e| SerError::Json(e.to_string()))
}
}
pub fn to_record_batch(records: &[SessionRecord]) -> Result<RecordBatch, SerError> {
let rows: Vec<SessionRecordRow> = records
.iter()
.map(SessionRecordRow::from_record)
.collect::<Result<_, _>>()?;
SessionRecordRow::to_arrow_batch(&rows).map_err(|e| SerError::Arrow(e.to_string()))
}
pub fn from_record_batch(batch: &RecordBatch) -> Result<Vec<SessionRecord>, SerError> {
let rows =
SessionRecordRow::from_arrow_batch(batch).map_err(|e| SerError::Arrow(e.to_string()))?;
rows.into_iter()
.map(SessionRecordRow::into_record)
.collect()
}
fn record_type_of(record: &SessionRecord) -> &'static str {
match record {
SessionRecord::Conversation { .. } => "conversation",
SessionRecord::WorkingMemory { .. } => "working_memory",
SessionRecord::Observation { .. } => "observation",
SessionRecord::Reflection { .. } => "reflection",
SessionRecord::FailedAction { .. } => "failed_action",
SessionRecord::Summary { .. } => "summary",
SessionRecord::Retracted { .. } => "retracted",
}
}
#[must_use]
pub fn sum_output_tokens(rows: &[SessionRecordRow]) -> u64 {
rows.iter()
.filter_map(|r| r.output_tokens)
.map(u64::from)
.sum()
}
#[must_use]
pub fn sum_input_tokens(rows: &[SessionRecordRow]) -> u64 {
rows.iter()
.filter_map(|r| r.input_tokens)
.map(u64::from)
.sum()
}
#[must_use]
pub fn filter_by_type<'a>(
rows: &'a [SessionRecordRow],
record_type: &str,
) -> Vec<&'a SessionRecordRow> {
rows.iter()
.filter(|r| r.record_type == record_type)
.collect()
}