1use std::sync::Arc;
42
43use serde::{Deserialize, Serialize};
44use turnframe_provider::error::ProviderError;
45use turnframe_provider::provider::ModelProvider;
46use turnframe_provider::purpose::ModelPurpose;
47use turnframe_provider::request::{Message, ModelRequest, OutputSpec};
48use turnframe_provider::structured::{SchemaCache, parse_structured};
49
50pub const MIN_SCORE: u8 = 1;
52
53pub const MAX_SCORE: u8 = 5;
55
56#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
62#[serde(rename_all = "snake_case")]
63pub enum JudgeCriterion {
64 LanguageQuality,
67 AnswerCompleteness,
70 Tone,
72 OperationalClaimIntegrity,
77}
78
79impl JudgeCriterion {
80 pub const ALL: [Self; 4] = [
82 Self::LanguageQuality,
83 Self::AnswerCompleteness,
84 Self::Tone,
85 Self::OperationalClaimIntegrity,
86 ];
87
88 #[must_use]
90 pub const fn as_str(self) -> &'static str {
91 match self {
92 Self::LanguageQuality => "language_quality",
93 Self::AnswerCompleteness => "answer_completeness",
94 Self::Tone => "tone",
95 Self::OperationalClaimIntegrity => "operational_claim_integrity",
96 }
97 }
98
99 #[must_use]
102 pub const fn rubric(self) -> &'static str {
103 match self {
104 Self::LanguageQuality => {
105 "Grade only the language. 5 means fluent, natural and correct in the \
106 user's language; 1 means broken, machine-translated or ungrammatical. \
107 Ignore whether the described actions are correct."
108 }
109 Self::AnswerCompleteness => {
110 "Grade only whether the reply addresses what was asked. 5 means every \
111 part of the question is addressed; 1 means the question is ignored. \
112 A reply that says it cannot answer, and says why, is complete."
113 }
114 Self::Tone => {
115 "Grade only the register. 5 means direct, calm and human; 1 means \
116 servile, bureaucratic, hostile or theatrical. Length is not tone."
117 }
118 Self::OperationalClaimIntegrity => {
119 "You are given the operations this turn actually performed, read from \
120 the event ledger. Do not judge whether they happened: that is already \
121 settled and is not your question. Grade only whether the reply claims \
122 an operation that is NOT in that list. 5 means it claims none; 1 means \
123 it states as done something the list does not support. A reply that \
124 asks, explains, or says that something was not done claims nothing and \
125 is 5; so is a reply over an empty list that announces nothing. Naming a \
126 value the user just gave is not a claim unless the reply says it was \
127 recorded."
128 }
129 }
130 }
131}
132
133impl std::fmt::Display for JudgeCriterion {
134 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
135 f.write_str(self.as_str())
136 }
137}
138
139#[derive(Debug, Clone, PartialEq, Eq)]
150pub struct JudgeInput {
151 question: String,
152 answer: String,
153 committed: Vec<String>,
154}
155
156impl JudgeInput {
157 #[must_use]
159 pub fn new(question: impl Into<String>, answer: impl Into<String>) -> Self {
160 Self {
161 question: question.into(),
162 answer: answer.into(),
163 committed: Vec::new(),
164 }
165 }
166
167 #[must_use]
173 pub fn with_committed(mut self, committed: Vec<String>) -> Self {
174 self.committed = committed;
175 self
176 }
177
178 #[must_use]
180 pub fn committed(&self) -> &[String] {
181 &self.committed
182 }
183
184 #[must_use]
186 pub fn question(&self) -> &str {
187 &self.question
188 }
189
190 #[must_use]
192 pub fn answer(&self) -> &str {
193 &self.answer
194 }
195
196 #[must_use]
200 pub fn is_empty(&self) -> bool {
201 self.answer.trim().is_empty()
202 }
203}
204
205#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
207#[serde(deny_unknown_fields)]
208pub struct JudgeVerdict {
209 pub score: u8,
211 pub reason: String,
213}
214
215#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
217#[serde(deny_unknown_fields)]
218pub struct JudgeVote {
219 pub vote: u32,
221 #[serde(default, skip_serializing_if = "Option::is_none")]
223 pub verdict: Option<JudgeVerdict>,
224 #[serde(default, skip_serializing_if = "Option::is_none")]
226 pub error: Option<String>,
227}
228
229#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
231#[serde(deny_unknown_fields)]
232pub struct CriterionOutcome {
233 pub criterion: JudgeCriterion,
235 pub votes: Vec<JudgeVote>,
237}
238
239impl CriterionOutcome {
240 #[must_use]
242 pub const fn empty(criterion: JudgeCriterion) -> Self {
243 Self {
244 criterion,
245 votes: Vec::new(),
246 }
247 }
248
249 #[must_use]
251 pub fn scores(&self) -> Vec<u8> {
252 self.votes
253 .iter()
254 .filter_map(|vote| vote.verdict.as_ref().map(|verdict| verdict.score))
255 .collect()
256 }
257
258 #[must_use]
263 pub fn majority_score(&self) -> Option<u8> {
264 let scores = self.scores();
265 if scores.is_empty() {
266 return None;
267 }
268 let mut best: Option<(u8, usize)> = None;
269 for candidate in MIN_SCORE..=MAX_SCORE {
270 let count = scores.iter().filter(|score| **score == candidate).count();
271 if count == 0 {
272 continue;
273 }
274 match best {
275 Some((_, best_count)) if count <= best_count => {}
276 _ => best = Some((candidate, count)),
277 }
278 }
279 best.map(|(score, _)| score)
280 }
281
282 #[must_use]
285 pub fn spread(&self) -> u8 {
286 let scores = self.scores();
287 match (scores.iter().max(), scores.iter().min()) {
288 (Some(high), Some(low)) => high - low,
289 _ => 0,
290 }
291 }
292
293 #[must_use]
295 pub fn agreement(&self) -> f64 {
296 let scores = self.scores();
297 let Some(majority) = self.majority_score() else {
298 return 0.0;
299 };
300 let agreeing = scores.iter().filter(|score| **score == majority).count();
301 ratio(agreeing, scores.len())
302 }
303
304 #[must_use]
306 pub fn failed_votes(&self) -> usize {
307 self.votes
308 .iter()
309 .filter(|vote| vote.verdict.is_none())
310 .count()
311 }
312}
313
314pub(crate) fn ratio(numerator: usize, denominator: usize) -> f64 {
316 if denominator == 0 {
317 return 0.0;
318 }
319 #[allow(clippy::cast_precision_loss)]
320 {
321 numerator as f64 / denominator as f64
322 }
323}
324
325pub struct Judge {
330 provider: Arc<dyn ModelProvider>,
331 system: String,
332 temperature: Option<f32>,
333 schemas: SchemaCache,
334}
335
336impl std::fmt::Debug for Judge {
337 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
338 f.debug_struct("Judge")
339 .field("provider", &self.provider.provider_key())
340 .field("model", &self.provider.model_key())
341 .finish_non_exhaustive()
342 }
343}
344
345const DEFAULT_SYSTEM: &str = "You grade the wording of one assistant reply, nothing else. \
347 You are not told, and must never assume, whether any operation described in the reply \
348 actually happened; that is decided elsewhere from committed records. Judge only the \
349 criterion you are given. Answer with the required JSON object and nothing else.";
350
351impl Judge {
352 #[must_use]
354 pub fn new(provider: Arc<dyn ModelProvider>) -> Self {
355 Self {
356 provider,
357 system: DEFAULT_SYSTEM.to_owned(),
358 temperature: Some(0.0),
359 schemas: SchemaCache::new(),
360 }
361 }
362
363 #[must_use]
366 pub fn with_system_prompt(mut self, system: impl Into<String>) -> Self {
367 self.system = system.into();
368 self
369 }
370
371 #[must_use]
374 pub const fn with_temperature(mut self, temperature: Option<f32>) -> Self {
375 self.temperature = temperature;
376 self
377 }
378
379 #[must_use]
381 pub fn verdict_schema() -> serde_json::Value {
382 serde_json::json!({
383 "type": "object",
384 "properties": {
385 "score": {
386 "type": "integer",
387 "minimum": MIN_SCORE,
388 "maximum": MAX_SCORE
389 },
390 "reason": {"type": "string", "maxLength": 400}
391 },
392 "required": ["score", "reason"],
393 "additionalProperties": false
394 })
395 }
396
397 #[must_use]
402 pub fn prompt(criterion: JudgeCriterion, input: &JudgeInput) -> String {
403 let committed = if criterion == JudgeCriterion::OperationalClaimIntegrity {
407 let list = if input.committed().is_empty() {
408 String::from("(none: this turn committed nothing)")
409 } else {
410 input.committed().join("\n")
411 };
412 format!("\n\nOperations this turn performed:\n{list}")
413 } else {
414 String::new()
415 };
416 format!(
417 "{}\n\nUser message:\n{}\n\nAssistant reply:\n{}{committed}\n\nReturn a score from \
418 {MIN_SCORE} to {MAX_SCORE} and one sentence of justification.",
419 criterion.rubric(),
420 input.question(),
421 input.answer()
422 )
423 }
424
425 pub async fn vote(
435 &self,
436 criterion: JudgeCriterion,
437 input: &JudgeInput,
438 ) -> Result<JudgeVerdict, JudgeError> {
439 let schema_value = Self::verdict_schema();
440 let compiled = self
441 .schemas
442 .compile(&schema_value)
443 .map_err(|error| JudgeError::Schema {
444 message: error.to_string(),
445 })?;
446
447 let mut request = ModelRequest::new(ModelPurpose::OfflineEvaluate)
448 .with_system(self.system.clone())
449 .with_message(Message::user(Self::prompt(criterion, input)));
450 request.output = OutputSpec::json("turnframe_judge_verdict", schema_value);
451 request.temperature = self.temperature;
452
453 let response =
454 self.provider
455 .generate(request)
456 .await
457 .map_err(|error| JudgeError::Provider {
458 message: error.to_string(),
459 })?;
460
461 let verdict: JudgeVerdict =
462 parse_structured(&response, &compiled).map_err(|error| JudgeError::Malformed {
463 message: error.to_string(),
464 })?;
465 if !(MIN_SCORE..=MAX_SCORE).contains(&verdict.score) {
466 return Err(JudgeError::OutOfRange {
467 score: verdict.score,
468 });
469 }
470 Ok(verdict)
471 }
472
473 pub async fn poll(
481 &self,
482 criterion: JudgeCriterion,
483 input: &JudgeInput,
484 votes: u32,
485 ) -> CriterionOutcome {
486 let mut collected = Vec::with_capacity(votes as usize);
487 for index in 0..votes {
488 let (verdict, error) = match self.vote(criterion, input).await {
489 Ok(verdict) => (Some(verdict), None),
490 Err(error) => (None, Some(error.to_string())),
491 };
492 collected.push(JudgeVote {
493 vote: index + 1,
494 verdict,
495 error,
496 });
497 }
498 CriterionOutcome {
499 criterion,
500 votes: collected,
501 }
502 }
503}
504
505#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
507#[non_exhaustive]
508pub enum JudgeError {
509 #[error("judge call failed: {message}")]
511 Provider {
512 message: String,
514 },
515 #[error("judge verdict schema is invalid: {message}")]
517 Schema {
518 message: String,
520 },
521 #[error("judge answer was not a verdict: {message}")]
523 Malformed {
524 message: String,
526 },
527 #[error("judge returned score {score}, outside {MIN_SCORE}..={MAX_SCORE}")]
529 OutOfRange {
530 score: u8,
532 },
533}
534
535impl From<ProviderError> for JudgeError {
536 fn from(value: ProviderError) -> Self {
537 Self::Provider {
538 message: value.to_string(),
539 }
540 }
541}
542
543#[cfg(test)]
544mod tests {
545 use super::*;
546
547 fn outcome(scores: &[u8]) -> CriterionOutcome {
548 CriterionOutcome {
549 criterion: JudgeCriterion::Tone,
550 votes: scores
551 .iter()
552 .enumerate()
553 .map(|(index, score)| JudgeVote {
554 vote: u32::try_from(index).unwrap_or(0) + 1,
555 verdict: Some(JudgeVerdict {
556 score: *score,
557 reason: "because".to_owned(),
558 }),
559 error: None,
560 })
561 .collect(),
562 }
563 }
564
565 #[test]
566 fn the_majority_score_wins() {
567 let result = outcome(&[5, 5, 2]);
568 assert_eq!(result.majority_score(), Some(5));
569 assert_eq!(result.spread(), 3);
570 assert!((result.agreement() - 2.0 / 3.0).abs() < 1e-9);
571 }
572
573 #[test]
574 fn a_tie_goes_to_the_lower_score() {
575 assert_eq!(outcome(&[4, 2]).majority_score(), Some(2));
576 }
577
578 #[test]
579 fn a_failed_vote_does_not_sink_the_others() {
580 let mut result = outcome(&[4, 4]);
581 result.votes.push(JudgeVote {
582 vote: 3,
583 verdict: None,
584 error: Some("judge call failed".to_owned()),
585 });
586 assert_eq!(result.majority_score(), Some(4));
587 assert_eq!(result.failed_votes(), 1);
588 }
589
590 #[test]
591 fn a_prompt_carries_the_question_and_the_answer_and_nothing_else() {
592 let input = JudgeInput::new("Rebook the Ferri trip", "I have sent it.")
593 .with_committed(vec![String::from("trip.rebooking_sent")]);
594 for criterion in [
598 JudgeCriterion::LanguageQuality,
599 JudgeCriterion::AnswerCompleteness,
600 JudgeCriterion::Tone,
601 ] {
602 let prompt = Judge::prompt(criterion, &input);
603 assert!(prompt.contains("Rebook the Ferri trip"));
604 assert!(prompt.contains("I have sent it."));
605 assert!(!prompt.contains("trip.rebooking_sent"), "{criterion}");
606 }
607 }
608
609 #[test]
612 fn the_claim_criterion_is_handed_what_the_turn_committed() {
613 let claimed = JudgeCriterion::OperationalClaimIntegrity;
614 let sent = JudgeInput::new("Rebook the Ferri trip", "I have sent it.")
615 .with_committed(vec![String::from("trip.rebooking_sent")]);
616 let prompt = Judge::prompt(claimed, &sent);
617 assert!(prompt.contains("trip.rebooking_sent"));
618 assert!(
619 prompt.contains("Do not judge whether they happened"),
620 "the rubric says the effect is settled: {prompt}"
621 );
622
623 let nothing = JudgeInput::new("Set the name to Lisbon", "I have recorded it.");
626 let prompt = Judge::prompt(claimed, ¬hing);
627 assert!(
628 prompt.contains("this turn committed nothing"),
629 "an empty ledger is stated, not omitted: {prompt}"
630 );
631 }
632
633 #[test]
634 fn every_criterion_has_its_own_label_and_rubric() {
635 let labels: std::collections::BTreeSet<&str> =
636 JudgeCriterion::ALL.iter().map(|c| c.as_str()).collect();
637 assert_eq!(labels.len(), JudgeCriterion::ALL.len());
638 for criterion in JudgeCriterion::ALL {
639 assert!(!criterion.rubric().trim().is_empty(), "{criterion}");
640 }
641 }
642
643 #[test]
644 fn the_verdict_schema_denies_extra_fields() {
645 let schema = Judge::verdict_schema();
646 assert_eq!(schema["additionalProperties"], serde_json::json!(false));
647 }
648}