use std::path::Path;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LocomoSample {
pub id: String,
pub sessions: Vec<Session>,
pub qa: Vec<QaItem>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Session {
#[serde(default)]
pub session_id: String,
pub turns: Vec<Turn>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Turn {
pub speaker: String,
pub text: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QaItem {
pub question: String,
pub answers: Vec<String>,
#[serde(default = "default_category")]
pub category: u8,
}
fn default_category() -> u8 {
1
}
impl LocomoSample {
pub fn transcript(&self) -> String {
let mut lines = Vec::new();
for session in &self.sessions {
for turn in &session.turns {
lines.push(format!("{}: {}", turn.speaker, turn.text));
}
}
lines.join("\n")
}
pub fn turn_count(&self) -> usize {
self.sessions.iter().map(|s| s.turns.len()).sum()
}
}
pub fn parse_suite(raw: &str) -> Result<Vec<LocomoSample>, String> {
let trimmed = raw.trim_start();
if trimmed.starts_with('[') {
return serde_json::from_str(trimmed).map_err(|e| format!("invalid JSON array: {e}"));
}
let mut out = Vec::new();
for (i, line) in raw.lines().enumerate() {
let l = line.trim();
if l.is_empty() || l.starts_with('#') {
continue;
}
let sample: LocomoSample =
serde_json::from_str(l).map_err(|e| format!("line {}: {e}", i + 1))?;
out.push(sample);
}
if out.is_empty() {
return Err("suite contained no samples".to_string());
}
Ok(out)
}
pub fn load_suite(path: &Path) -> Result<Vec<LocomoSample>, String> {
let raw =
std::fs::read_to_string(path).map_err(|e| format!("reading {}: {e}", path.display()))?;
parse_suite(&raw)
}
pub const REFERENCE_SUITE: &str = include_str!("../../../data/locomo/reference-suite.ndjson");
pub fn reference_samples() -> Vec<LocomoSample> {
parse_suite(REFERENCE_SUITE).expect("bundled reference suite must be valid")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn reference_suite_parses_and_is_grounded() {
let samples = reference_samples();
assert!(!samples.is_empty());
for s in &samples {
assert!(!s.sessions.is_empty(), "{} has no sessions", s.id);
assert!(!s.qa.is_empty(), "{} has no QA", s.id);
let transcript = s.transcript().to_lowercase();
for qa in &s.qa {
assert!(!qa.answers.is_empty(), "QA without answers in {}", s.id);
let grounded = qa
.answers
.iter()
.any(|a| transcript.contains(&a.to_lowercase()));
assert!(
grounded,
"answer for '{}' not grounded in transcript of {}",
qa.question, s.id
);
}
}
}
#[test]
fn parses_json_array_form() {
let raw = r#"[{"id":"x","sessions":[{"session_id":"s","turns":[{"speaker":"A","text":"hi"}]}],"qa":[{"question":"q","answers":["hi"]}]}]"#;
let s = parse_suite(raw).unwrap();
assert_eq!(s.len(), 1);
assert_eq!(s[0].qa[0].category, 1, "default category applied");
}
#[test]
fn empty_suite_errors() {
assert!(parse_suite("# only a comment\n").is_err());
}
}