Skip to main content

turnframe_eval/simulate/
score.rs

1//! A conversation scored by code, and a run reported as rates over its conversations.
2
3use std::fmt::Write as _;
4
5use serde::{Deserialize, Serialize};
6
7use super::user::UserMove;
8
9/// One turn of a conversation, as code reads it.
10#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
11pub struct Exchange {
12    /// What the person did.
13    pub said: Option<UserMove>,
14    /// What they read back.
15    pub reply: String,
16    /// The error code of a turn that failed outright.
17    pub failed: Option<String>,
18    /// Whether the reply ended on a question, an ask, a card or an offer.
19    pub way_forward: bool,
20    /// What the reply asked, each as its record and what it waits on.
21    pub asks: Vec<String>,
22    /// The operations the reply offered.
23    pub offers: Vec<String>,
24    /// The operations the domain refused this turn.
25    pub refused: Vec<String>,
26    /// How many parts of the message were not understood.
27    pub not_understood: usize,
28}
29
30/// The classes a conversation is scored on.
31#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
32pub struct ConversationScore {
33    /// Whether the goal's state was reached.
34    pub reached: bool,
35    /// Turns taken.
36    pub turns: usize,
37    /// Replies that ended on no question, no ask, no card and no offer.
38    pub dead_ends: usize,
39    /// Replies that asked what the reply before them asked.
40    pub loops: usize,
41    /// Parts of messages not understood.
42    pub not_understood: usize,
43    /// Acts the domain refused.
44    pub refused: usize,
45    /// Refused acts of an operation the reply before had offered.
46    pub offers_refused: usize,
47    /// Turns that failed outright.
48    pub failed_turns: usize,
49}
50
51impl ConversationScore {
52    /// Broken guarantees: a reply with no way forward, an offer the domain refused. Zero.
53    #[must_use]
54    pub const fn violations(&self) -> usize {
55        self.dead_ends + self.offers_refused
56    }
57}
58
59/// Scores `exchanges`, the goal `reached` or not.
60#[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/// One conversation: the goal, the manner, what was said and how it scored.
102#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
103pub struct Conversation {
104    /// The goal's id.
105    pub goal: String,
106    /// The manner played.
107    pub manner: String,
108    /// Which sample of the pair it is.
109    pub sample: u32,
110    /// The turns, in order.
111    pub exchanges: Vec<Exchange>,
112    /// Why the conversation stopped.
113    pub ended: Ending,
114    /// Its score.
115    pub score: ConversationScore,
116}
117
118impl Conversation {
119    /// The conversation as a person would read it back.
120    #[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    /// How badly it went, worst first when sorted descending.
151    fn badness(&self) -> (usize, bool, usize) {
152        (
153            self.score.violations(),
154            !self.score.reached,
155            self.score.turns,
156        )
157    }
158}
159
160/// Why a conversation stopped.
161#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
162#[serde(tag = "kind", rename_all = "snake_case")]
163pub enum Ending {
164    /// The person said they were done.
165    Done,
166    /// The turn limit came first.
167    TurnLimit,
168    /// The person could not decide: the simulator failed.
169    UserFailed {
170        /// Why.
171        message: String,
172    },
173    /// The world could not be prepared.
174    Unprepared {
175        /// Why.
176        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
191/// How one class is read off a score.
192type Class = fn(&ConversationScore) -> usize;
193
194/// A run: every conversation, reported as rates.
195#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
196pub struct SimulationReport {
197    /// The conversations, in goal, manner and sample order.
198    pub conversations: Vec<Conversation>,
199}
200
201impl SimulationReport {
202    /// The conversations that ran: an unprepared world measured nothing.
203    fn measured(&self) -> impl Iterator<Item = &Conversation> {
204        self.conversations
205            .iter()
206            .filter(|conversation| !matches!(conversation.ended, Ending::Unprepared { .. }))
207    }
208
209    /// How many conversations ran.
210    #[must_use]
211    pub fn measured_count(&self) -> usize {
212        self.measured().count()
213    }
214
215    /// Broken guarantees over the whole run: zero, or a defect.
216    #[must_use]
217    pub fn violations(&self) -> usize {
218        self.measured()
219            .map(|conversation| conversation.score.violations())
220            .sum()
221    }
222
223    /// The `count` worst conversations: guarantees broken, then goals missed, then length.
224    #[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    /// Each class as a rate over the conversations, then the worst transcripts.
233    #[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    /// The report as JSON.
281    ///
282    /// # Errors
283    ///
284    /// A serialization error, which a report of plain data does not produce.
285    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}