1use std::fmt::Write as _;
4
5use serde::{Deserialize, Serialize};
6
7use super::user::UserMove;
8
9#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
11pub struct Exchange {
12 pub said: Option<UserMove>,
14 pub reply: String,
16 pub failed: Option<String>,
18 pub way_forward: bool,
20 pub asks: Vec<String>,
22 pub offers: Vec<String>,
24 pub refused: Vec<String>,
26 pub not_understood: usize,
28}
29
30#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
32pub struct ConversationScore {
33 pub reached: bool,
35 pub turns: usize,
37 pub dead_ends: usize,
39 pub loops: usize,
41 pub not_understood: usize,
43 pub refused: usize,
45 pub offers_refused: usize,
47 pub failed_turns: usize,
49}
50
51impl ConversationScore {
52 #[must_use]
54 pub const fn violations(&self) -> usize {
55 self.dead_ends + self.offers_refused
56 }
57}
58
59#[must_use]
61pub fn score(exchanges: &[Exchange], reached: bool) -> ConversationScore {
62 let answered = |exchange: &&Exchange| exchange.failed.is_none();
63 ConversationScore {
64 reached,
65 turns: exchanges.len(),
66 dead_ends: exchanges
67 .iter()
68 .filter(answered)
69 .filter(|exchange| !exchange.way_forward)
70 .count(),
71 loops: exchanges
72 .windows(2)
73 .filter(|pair| pair[1].failed.is_none() && !pair[1].asks.is_empty())
74 .filter(|pair| pair[0].asks == pair[1].asks)
75 .count(),
76 not_understood: exchanges
77 .iter()
78 .map(|exchange| exchange.not_understood)
79 .sum(),
80 refused: exchanges
81 .iter()
82 .map(|exchange| exchange.refused.len())
83 .sum(),
84 offers_refused: exchanges
85 .windows(2)
86 .map(|pair| {
87 pair[1]
88 .refused
89 .iter()
90 .filter(|operation| pair[0].offers.contains(operation))
91 .count()
92 })
93 .sum(),
94 failed_turns: exchanges
95 .iter()
96 .filter(|exchange| !answered(exchange))
97 .count(),
98 }
99}
100
101#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
103pub struct Conversation {
104 pub goal: String,
106 pub manner: String,
108 pub sample: u32,
110 pub exchanges: Vec<Exchange>,
112 pub ended: Ending,
114 pub score: ConversationScore,
116}
117
118impl Conversation {
119 #[must_use]
121 pub fn transcript(&self) -> String {
122 let mut out = format!(
123 "{} ({}, sample {}): {}",
124 self.goal,
125 self.manner,
126 self.sample,
127 if self.score.reached {
128 "reached"
129 } else {
130 "not reached"
131 }
132 );
133 for exchange in &self.exchanges {
134 if let Some(said) = &exchange.said {
135 let _ = write!(out, "\n user: {said}");
136 }
137 match &exchange.failed {
138 Some(code) => {
139 let _ = write!(out, "\n (the turn failed: {code})");
140 }
141 None => {
142 let _ = write!(out, "\n assistant: {}", exchange.reply.replace('\n', " "));
143 }
144 }
145 }
146 let _ = write!(out, "\n ({})", self.ended);
147 out
148 }
149
150 fn badness(&self) -> (usize, bool, usize) {
152 (
153 self.score.violations(),
154 !self.score.reached,
155 self.score.turns,
156 )
157 }
158}
159
160#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
162#[serde(tag = "kind", rename_all = "snake_case")]
163pub enum Ending {
164 Done,
166 TurnLimit,
168 UserFailed {
170 message: String,
172 },
173 Unprepared {
175 message: String,
177 },
178}
179
180impl std::fmt::Display for Ending {
181 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
182 match self {
183 Self::Done => f.write_str("the user was done"),
184 Self::TurnLimit => f.write_str("the turn limit came first"),
185 Self::UserFailed { message } => write!(f, "the simulated user failed: {message}"),
186 Self::Unprepared { message } => write!(f, "the world was not prepared: {message}"),
187 }
188 }
189}
190
191type Class = fn(&ConversationScore) -> usize;
193
194#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
196pub struct SimulationReport {
197 pub conversations: Vec<Conversation>,
199}
200
201impl SimulationReport {
202 fn measured(&self) -> impl Iterator<Item = &Conversation> {
204 self.conversations
205 .iter()
206 .filter(|conversation| !matches!(conversation.ended, Ending::Unprepared { .. }))
207 }
208
209 #[must_use]
211 pub fn measured_count(&self) -> usize {
212 self.measured().count()
213 }
214
215 #[must_use]
217 pub fn violations(&self) -> usize {
218 self.measured()
219 .map(|conversation| conversation.score.violations())
220 .sum()
221 }
222
223 #[must_use]
225 pub fn worst(&self, count: usize) -> Vec<&Conversation> {
226 let mut all: Vec<&Conversation> = self.measured().collect();
227 all.sort_by_key(|conversation| std::cmp::Reverse(conversation.badness()));
228 all.truncate(count);
229 all
230 }
231
232 #[must_use]
234 pub fn summary(&self) -> String {
235 let measured: Vec<&Conversation> = self.measured().collect();
236 let total = measured.len();
237 let sum = |class: Class| -> usize {
238 measured
239 .iter()
240 .map(|conversation| class(&conversation.score))
241 .sum()
242 };
243 let per = |count: usize| {
244 if total == 0 {
245 0.0
246 } else {
247 count as f64 / total as f64
248 }
249 };
250 let reached = measured.iter().filter(|c| c.score.reached).count();
251 let mut out = format!(
252 "{total} conversation(s), {} unprepared\nreached: {reached}/{total} ({:.0}%)",
253 self.conversations.len() - total,
254 per(reached) * 100.0
255 );
256 let classes: [(&str, Class); 7] = [
257 ("turns", |s| s.turns),
258 ("dead ends", |s| s.dead_ends),
259 ("loops", |s| s.loops),
260 ("not understood", |s| s.not_understood),
261 ("refused", |s| s.refused),
262 ("offers refused", |s| s.offers_refused),
263 ("failed turns", |s| s.failed_turns),
264 ];
265 for (name, class) in classes {
266 let count = sum(class);
267 let _ = write!(
268 out,
269 "\n{name}: {count} ({:.2} per conversation)",
270 per(count)
271 );
272 }
273 let _ = write!(out, "\nguarantee violations: {}", self.violations());
274 for conversation in self.worst(3) {
275 let _ = write!(out, "\n\n{}", conversation.transcript());
276 }
277 out
278 }
279
280 pub fn to_json(&self) -> serde_json::Result<String> {
286 serde_json::to_string_pretty(self)
287 }
288}
289
290#[cfg(test)]
291mod tests {
292 use super::*;
293
294 fn exchange(way_forward: bool) -> Exchange {
295 Exchange {
296 said: Some(UserMove::say("hello")),
297 reply: "Hi.".to_owned(),
298 way_forward,
299 ..Exchange::default()
300 }
301 }
302
303 #[test]
304 fn a_reply_with_no_way_forward_is_a_dead_end_and_a_violation() {
305 let score = score(&[exchange(true), exchange(false)], true);
306 assert_eq!(score.dead_ends, 1);
307 assert_eq!(score.violations(), 1);
308 }
309
310 #[test]
311 fn the_same_ask_twice_in_a_row_is_a_loop() {
312 let asking = |what: &str| Exchange {
313 asks: vec![what.to_owned()],
314 ..exchange(true)
315 };
316 let score = score(
317 &[
318 asking("trip/1: name"),
319 asking("trip/1: name"),
320 asking("trip/1: date"),
321 ],
322 false,
323 );
324 assert_eq!(score.loops, 1);
325 assert_eq!(
326 score.violations(),
327 0,
328 "a loop is measured, not a broken guarantee"
329 );
330 }
331
332 #[test]
333 fn an_offer_refused_on_the_next_turn_is_a_violation() {
334 let offering = Exchange {
335 offers: vec!["sample.send".to_owned()],
336 ..exchange(true)
337 };
338 let refusing = Exchange {
339 refused: vec!["sample.send".to_owned(), "sample.other".to_owned()],
340 ..exchange(true)
341 };
342 let score = score(&[offering, refusing], false);
343 assert_eq!((score.refused, score.offers_refused), (2, 1));
344 assert_eq!(score.violations(), 1);
345 }
346
347 #[test]
348 fn a_failed_turn_is_counted_apart_from_dead_ends() {
349 let failed = Exchange {
350 failed: Some("internal".to_owned()),
351 ..exchange(false)
352 };
353 let score = score(&[failed], false);
354 assert_eq!((score.failed_turns, score.dead_ends), (1, 0));
355 }
356
357 #[test]
358 fn the_worst_conversations_come_first() {
359 let conversation = |goal: &str, reached: bool, dead_ends: usize| Conversation {
360 goal: goal.to_owned(),
361 manner: "plain".to_owned(),
362 sample: 1,
363 exchanges: Vec::new(),
364 ended: Ending::Done,
365 score: ConversationScore {
366 reached,
367 dead_ends,
368 ..ConversationScore::default()
369 },
370 };
371 let report = SimulationReport {
372 conversations: vec![
373 conversation("fine", true, 0),
374 conversation("missed", false, 0),
375 conversation("broken", true, 2),
376 ],
377 };
378 let worst: Vec<&str> = report.worst(3).iter().map(|c| c.goal.as_str()).collect();
379 assert_eq!(worst, ["broken", "missed", "fine"]);
380 assert_eq!(report.violations(), 2);
381 }
382}