use std::fmt::Write as _;
use serde::{Deserialize, Serialize};
use super::user::UserMove;
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct Exchange {
pub said: Option<UserMove>,
pub reply: String,
pub failed: Option<String>,
pub way_forward: bool,
pub asks: Vec<String>,
pub offers: Vec<String>,
pub refused: Vec<String>,
pub not_understood: usize,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ConversationScore {
pub reached: bool,
pub turns: usize,
pub dead_ends: usize,
pub loops: usize,
pub not_understood: usize,
pub refused: usize,
pub offers_refused: usize,
pub failed_turns: usize,
}
impl ConversationScore {
#[must_use]
pub const fn violations(&self) -> usize {
self.dead_ends + self.offers_refused
}
}
#[must_use]
pub fn score(exchanges: &[Exchange], reached: bool) -> ConversationScore {
let answered = |exchange: &&Exchange| exchange.failed.is_none();
ConversationScore {
reached,
turns: exchanges.len(),
dead_ends: exchanges
.iter()
.filter(answered)
.filter(|exchange| !exchange.way_forward)
.count(),
loops: exchanges
.windows(2)
.filter(|pair| pair[1].failed.is_none() && !pair[1].asks.is_empty())
.filter(|pair| pair[0].asks == pair[1].asks)
.count(),
not_understood: exchanges
.iter()
.map(|exchange| exchange.not_understood)
.sum(),
refused: exchanges
.iter()
.map(|exchange| exchange.refused.len())
.sum(),
offers_refused: exchanges
.windows(2)
.map(|pair| {
pair[1]
.refused
.iter()
.filter(|operation| pair[0].offers.contains(operation))
.count()
})
.sum(),
failed_turns: exchanges
.iter()
.filter(|exchange| !answered(exchange))
.count(),
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Conversation {
pub goal: String,
pub manner: String,
pub sample: u32,
pub exchanges: Vec<Exchange>,
pub ended: Ending,
pub score: ConversationScore,
}
impl Conversation {
#[must_use]
pub fn transcript(&self) -> String {
let mut out = format!(
"{} ({}, sample {}): {}",
self.goal,
self.manner,
self.sample,
if self.score.reached {
"reached"
} else {
"not reached"
}
);
for exchange in &self.exchanges {
if let Some(said) = &exchange.said {
let _ = write!(out, "\n user: {said}");
}
match &exchange.failed {
Some(code) => {
let _ = write!(out, "\n (the turn failed: {code})");
}
None => {
let _ = write!(out, "\n assistant: {}", exchange.reply.replace('\n', " "));
}
}
}
let _ = write!(out, "\n ({})", self.ended);
out
}
fn badness(&self) -> (usize, bool, usize) {
(
self.score.violations(),
!self.score.reached,
self.score.turns,
)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum Ending {
Done,
TurnLimit,
UserFailed {
message: String,
},
Unprepared {
message: String,
},
}
impl std::fmt::Display for Ending {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Done => f.write_str("the user was done"),
Self::TurnLimit => f.write_str("the turn limit came first"),
Self::UserFailed { message } => write!(f, "the simulated user failed: {message}"),
Self::Unprepared { message } => write!(f, "the world was not prepared: {message}"),
}
}
}
type Class = fn(&ConversationScore) -> usize;
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct SimulationReport {
pub conversations: Vec<Conversation>,
}
impl SimulationReport {
fn measured(&self) -> impl Iterator<Item = &Conversation> {
self.conversations
.iter()
.filter(|conversation| !matches!(conversation.ended, Ending::Unprepared { .. }))
}
#[must_use]
pub fn measured_count(&self) -> usize {
self.measured().count()
}
#[must_use]
pub fn violations(&self) -> usize {
self.measured()
.map(|conversation| conversation.score.violations())
.sum()
}
#[must_use]
pub fn worst(&self, count: usize) -> Vec<&Conversation> {
let mut all: Vec<&Conversation> = self.measured().collect();
all.sort_by_key(|conversation| std::cmp::Reverse(conversation.badness()));
all.truncate(count);
all
}
#[must_use]
pub fn summary(&self) -> String {
let measured: Vec<&Conversation> = self.measured().collect();
let total = measured.len();
let sum = |class: Class| -> usize {
measured
.iter()
.map(|conversation| class(&conversation.score))
.sum()
};
let per = |count: usize| {
if total == 0 {
0.0
} else {
count as f64 / total as f64
}
};
let reached = measured.iter().filter(|c| c.score.reached).count();
let mut out = format!(
"{total} conversation(s), {} unprepared\nreached: {reached}/{total} ({:.0}%)",
self.conversations.len() - total,
per(reached) * 100.0
);
let classes: [(&str, Class); 7] = [
("turns", |s| s.turns),
("dead ends", |s| s.dead_ends),
("loops", |s| s.loops),
("not understood", |s| s.not_understood),
("refused", |s| s.refused),
("offers refused", |s| s.offers_refused),
("failed turns", |s| s.failed_turns),
];
for (name, class) in classes {
let count = sum(class);
let _ = write!(
out,
"\n{name}: {count} ({:.2} per conversation)",
per(count)
);
}
let _ = write!(out, "\nguarantee violations: {}", self.violations());
for conversation in self.worst(3) {
let _ = write!(out, "\n\n{}", conversation.transcript());
}
out
}
pub fn to_json(&self) -> serde_json::Result<String> {
serde_json::to_string_pretty(self)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn exchange(way_forward: bool) -> Exchange {
Exchange {
said: Some(UserMove::say("hello")),
reply: "Hi.".to_owned(),
way_forward,
..Exchange::default()
}
}
#[test]
fn a_reply_with_no_way_forward_is_a_dead_end_and_a_violation() {
let score = score(&[exchange(true), exchange(false)], true);
assert_eq!(score.dead_ends, 1);
assert_eq!(score.violations(), 1);
}
#[test]
fn the_same_ask_twice_in_a_row_is_a_loop() {
let asking = |what: &str| Exchange {
asks: vec![what.to_owned()],
..exchange(true)
};
let score = score(
&[
asking("trip/1: name"),
asking("trip/1: name"),
asking("trip/1: date"),
],
false,
);
assert_eq!(score.loops, 1);
assert_eq!(
score.violations(),
0,
"a loop is measured, not a broken guarantee"
);
}
#[test]
fn an_offer_refused_on_the_next_turn_is_a_violation() {
let offering = Exchange {
offers: vec!["sample.send".to_owned()],
..exchange(true)
};
let refusing = Exchange {
refused: vec!["sample.send".to_owned(), "sample.other".to_owned()],
..exchange(true)
};
let score = score(&[offering, refusing], false);
assert_eq!((score.refused, score.offers_refused), (2, 1));
assert_eq!(score.violations(), 1);
}
#[test]
fn a_failed_turn_is_counted_apart_from_dead_ends() {
let failed = Exchange {
failed: Some("internal".to_owned()),
..exchange(false)
};
let score = score(&[failed], false);
assert_eq!((score.failed_turns, score.dead_ends), (1, 0));
}
#[test]
fn the_worst_conversations_come_first() {
let conversation = |goal: &str, reached: bool, dead_ends: usize| Conversation {
goal: goal.to_owned(),
manner: "plain".to_owned(),
sample: 1,
exchanges: Vec::new(),
ended: Ending::Done,
score: ConversationScore {
reached,
dead_ends,
..ConversationScore::default()
},
};
let report = SimulationReport {
conversations: vec![
conversation("fine", true, 0),
conversation("missed", false, 0),
conversation("broken", true, 2),
],
};
let worst: Vec<&str> = report.worst(3).iter().map(|c| c.goal.as_str()).collect();
assert_eq!(worst, ["broken", "missed", "fine"]);
assert_eq!(report.violations(), 2);
}
}