From d6df493beb630f8066349d01ac9402119d928227 Mon Sep 17 00:00:00 2001 From: nachiketb Date: Wed, 30 Sep 2026 12:01:41 -0700 Subject: [PATCH 1/4] feat(protocol): add decision model request and response types Signed-off-by: nachiketb --- crates/protocol/src/decision.rs | 453 ++++++++++++++++++++++++++++++++ crates/protocol/src/lib.rs | 2 + 2 files changed, 455 insertions(+) create mode 100644 crates/protocol/src/decision.rs diff --git a/crates/protocol/src/decision.rs b/crates/protocol/src/decision.rs new file mode 100644 index 000000000..e49c32b2e --- /dev/null +++ b/crates/protocol/src/decision.rs @@ -0,0 +1,453 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Provider-neutral decision questions and answers, separate from LLM messages. +//! +//! Requests use checked construction and deserialization. Responses must be +//! constructed or decoded with their request to check answer kinds and rubrics. +//! Enums use snake-case `type` tags and a `data` payload in serialized form. + +use std::collections::{BTreeMap, BTreeSet}; + +use serde::{Deserialize, Deserializer, Serialize, de}; +use serde_json::Value; + +use crate::{ModelId, Usage}; + +/// Maximum absolute error in a distribution's sum. Values are never renormalized. +pub const DISTRIBUTION_SUM_TOLERANCE: f64 = 1e-6; + +/// A decision value or message violates the protocol contract. +#[derive(Debug, thiserror::Error)] +#[error("{0}")] +pub struct DecisionError(String); + +/// A finite probability in `[0, 1]`, serialized as a number. +#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] +#[serde(try_from = "f64", into = "f64")] +pub struct Probability(f64); + +impl Probability { + /// Rejects non-finite numbers and values outside `[0, 1]`. + pub fn new(value: f64) -> Result { + if !value.is_finite() || !(0.0..=1.0).contains(&value) { + return Err(DecisionError( + "probability must be finite and in [0, 1]".into(), + )); + } + Ok(Self(value)) + } + + /// Returns the probability as a number. + pub fn get(self) -> f64 { + self.0 + } +} + +impl TryFrom for Probability { + type Error = DecisionError; + + fn try_from(value: f64) -> Result { + Self::new(value) + } +} + +impl From for f64 { + fn from(value: Probability) -> Self { + value.get() + } +} + +/// A finite, nonnegative rubric position, including fractional positions. +/// The response constructor also checks the request's upper bound. +#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] +#[serde(try_from = "f64", into = "f64")] +pub struct ScoreValue(f64); + +impl ScoreValue { + /// Rejects non-finite and negative positions. + pub fn new(value: f64) -> Result { + if !value.is_finite() || value < 0.0 { + return Err(DecisionError("score must be finite and nonnegative".into())); + } + Ok(Self(value)) + } + + /// Returns the rubric position. + pub fn get(self) -> f64 { + self.0 + } +} + +impl TryFrom for ScoreValue { + type Error = DecisionError; + + fn try_from(value: f64) -> Result { + Self::new(value) + } +} + +impl From for f64 { + fn from(value: ScoreValue) -> Self { + value.get() + } +} + +/// Finite provider confidence; its scale and meaning are provider-specific. +#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] +#[serde(try_from = "f64", into = "f64")] +pub struct ProviderConfidence(f64); + +impl ProviderConfidence { + /// Rejects non-finite confidence values without imposing a probability scale. + pub fn new(value: f64) -> Result { + if !value.is_finite() { + return Err(DecisionError("provider confidence must be finite".into())); + } + Ok(Self(value)) + } + + /// Returns the provider's confidence value. + pub fn get(self) -> f64 { + self.0 + } +} + +impl TryFrom for ProviderConfidence { + type Error = DecisionError; + + fn try_from(value: f64) -> Result { + Self::new(value) + } +} + +impl From for f64 { + fn from(value: ProviderConfidence) -> Self { + value.get() + } +} + +/// Shared context evaluated against independent, named questions. +#[derive(Clone, Debug, PartialEq, Serialize)] +pub struct DecisionRequest { + /// Optional until a target is selected. + pub model: Option, + /// Conversation, application state, or other material to evaluate. + pub context: Value, + questions: BTreeMap, +} + +impl DecisionRequest { + /// Requires questions, nonempty unique choice options, and at least two score levels. + pub fn new( + model: Option, + context: Value, + questions: BTreeMap, + ) -> Result { + if questions.is_empty() { + return Err(DecisionError( + "request must contain at least one question".into(), + )); + } + for (id, question) in &questions { + match &question.kind { + DecisionKind::Choice { options } => { + let ids: BTreeSet<_> = options.iter().map(|option| &option.id).collect(); + if options.is_empty() || ids.len() != options.len() { + return Err(DecisionError(format!( + "question {id:?}: choice options must be nonempty with unique IDs" + ))); + } + } + DecisionKind::Score { levels } if levels.len() < 2 => { + return Err(DecisionError(format!( + "question {id:?}: score requires at least two levels" + ))); + } + _ => {} + } + } + Ok(Self { + model, + context, + questions, + }) + } + + /// Questions are immutable so their checked rubrics remain valid. + pub fn questions(&self) -> &BTreeMap { + &self.questions + } +} + +impl<'de> Deserialize<'de> for DecisionRequest { + fn deserialize>(deserializer: D) -> Result { + #[derive(Deserialize)] + struct RawRequest { + model: Option, + context: Value, + #[serde(deserialize_with = "unique_map")] + questions: BTreeMap, + } + let raw = RawRequest::deserialize(deserializer)?; + Self::new(raw.model, raw.context, raw.questions).map_err(de::Error::custom) + } +} + +/// Instructions and the expected answer shape for one question. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct DecisionQuestion { + /// Structured or textual instructions shared with the provider. + pub instructions: Value, + /// Checked against the answer when constructing a response. + pub kind: DecisionKind, +} + +/// The answer shape and any options or ordered rubric levels. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "type", content = "data", rename_all = "snake_case")] +pub enum DecisionKind { + /// A Boolean judgment or probability of true. + Boolean { + /// Meaning of a true answer, when needed. + true_description: Option, + /// Meaning of a false answer, when needed. + false_description: Option, + }, + /// Select one of the declared options. + Choice { + /// Nonempty options with unique IDs, preserving caller order. + options: Vec, + }, + /// A position on an ordered rubric, not an arbitrary numeric measurement. + Score { + /// At least two levels, ordered low to high and indexed from zero. + levels: Vec, + }, +} + +/// An identified choice with an optional structured description. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ChoiceOption { + /// Stable identifier used by choice answers and distributions. + pub id: String, + /// Meaning of this option, when its ID alone is insufficient. + pub description: Option, +} + +/// Complete answers checked against a request. +/// +/// Use [`Self::deserialize`] with the matching request to decode a response. +/// There is no request-free `Deserialize` implementation. +#[derive(Clone, Debug, PartialEq, Serialize)] +pub struct DecisionResponse { + /// Provider-reported response identifier. + pub id: Option, + /// Provider-reported model identifier. + pub model: Option, + answers: BTreeMap, + /// Available token counts; absent counts remain unknown. + pub usage: Usage, +} + +impl DecisionResponse { + /// Checks answer coverage, kinds, selected options, score bounds, and distributions. + /// Score/distribution consistency remains the provider's responsibility. + pub fn new( + request: &DecisionRequest, + id: Option, + model: Option, + answers: BTreeMap, + usage: Usage, + ) -> Result { + if !answers.keys().eq(request.questions.keys()) { + return Err(DecisionError( + "answer IDs must exactly match question IDs".into(), + )); + } + for (id, question) in &request.questions { + check_answer(&question.kind, &answers[id].value) + .map_err(|error| DecisionError(format!("question {id:?}: {error}")))?; + } + Ok(Self { + id, + model, + answers, + usage, + }) + } + + /// Decodes the canonical representation and checks it against its request. + pub fn deserialize<'de, D: Deserializer<'de>>( + request: &DecisionRequest, + deserializer: D, + ) -> Result { + #[derive(Deserialize)] + struct RawResponse { + id: Option, + model: Option, + #[serde(deserialize_with = "unique_map")] + answers: BTreeMap, + #[serde(default)] + usage: Usage, + } + let raw = RawResponse::deserialize(deserializer)?; + Self::new(request, raw.id, raw.model, raw.answers, raw.usage).map_err(de::Error::custom) + } + + /// Answers are immutable after their request-dependent checks. + pub fn answers(&self) -> &BTreeMap { + &self.answers + } +} + +/// An answer and separate, optional provider confidence. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct DecisionAnswer { + /// The estimate, checked against its question when building a response. + pub value: DecisionValue, + /// Provider-specific confidence, distinct from answer probabilities. + pub provider_confidence: Option, +} + +/// A typed estimate; missing distributions remain unknown. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "type", content = "data", rename_all = "snake_case")] +pub enum DecisionValue { + /// A Boolean judgment or probability, without an implicit threshold. + Boolean(BooleanEstimate), + /// One selected option with an optional complete distribution. + Choice { + /// Must name an option in the matching question. + selected: String, + /// Maps every declared option ID to its probability when available. + #[serde(default, deserialize_with = "optional_unique_map")] + probabilities: Option>, + }, + /// A fractional position in the matching request's rubric. + Score { + /// Must lie in `0..=N-1` for the request's N levels. + value: ScoreValue, + /// Follows the request's level order. Retain that request to interpret it. + probabilities: Option>, + }, +} + +/// Preserves Boolean-only answers without inventing probability or certainty. +#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "type", content = "data", rename_all = "snake_case")] +pub enum BooleanEstimate { + /// A Boolean judgment with no probability supplied. + Value(bool), + /// Probability of true; algorithms choose their own thresholds. + ProbabilityTrue(Probability), +} + +fn check_answer(kind: &DecisionKind, answer: &DecisionValue) -> Result<(), DecisionError> { + match (kind, answer) { + (DecisionKind::Boolean { .. }, DecisionValue::Boolean(_)) => Ok(()), + ( + DecisionKind::Choice { options }, + DecisionValue::Choice { + selected, + probabilities, + }, + ) => { + if !options.iter().any(|option| option.id == *selected) { + return Err(DecisionError( + "selected choice is not a declared option".into(), + )); + } + if let Some(probabilities) = probabilities { + if probabilities.len() != options.len() + || options + .iter() + .any(|option| !probabilities.contains_key(&option.id)) + { + return Err(DecisionError( + "distribution must cover every option exactly".into(), + )); + } + check_sum(probabilities.values())?; + } + Ok(()) + } + ( + DecisionKind::Score { levels }, + DecisionValue::Score { + value, + probabilities, + }, + ) => { + if value.get() > (levels.len() - 1) as f64 { + return Err(DecisionError("score exceeds the request's rubric".into())); + } + if let Some(probabilities) = probabilities { + if probabilities.len() != levels.len() { + return Err(DecisionError( + "distribution must cover every level exactly".into(), + )); + } + check_sum(probabilities.iter())?; + } + Ok(()) + } + _ => Err(DecisionError( + "answer kind does not match question kind".into(), + )), + } +} + +fn check_sum<'a>( + probabilities: impl Iterator, +) -> Result<(), DecisionError> { + let sum: f64 = probabilities.map(|probability| probability.get()).sum(); + if (sum - 1.0).abs() > DISTRIBUTION_SUM_TOLERANCE { + return Err(DecisionError(format!( + "distribution must sum to one within {DISTRIBUTION_SUM_TOLERANCE}" + ))); + } + Ok(()) +} + +// Reject duplicate JSON keys before a map could silently discard an answer or option. +fn unique_map<'de, D, T>(deserializer: D) -> Result, D::Error> +where + D: Deserializer<'de>, + T: Deserialize<'de>, +{ + struct Visitor(std::marker::PhantomData); + + impl<'de, T: Deserialize<'de>> de::Visitor<'de> for Visitor { + type Value = BTreeMap; + + fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result { + formatter.write_str("a map with unique IDs") + } + + fn visit_map>(self, mut access: A) -> Result { + let mut map = BTreeMap::new(); + while let Some((key, value)) = access.next_entry::()? { + match map.entry(key) { + std::collections::btree_map::Entry::Vacant(entry) => { + entry.insert(value); + } + std::collections::btree_map::Entry::Occupied(entry) => { + return Err(de::Error::custom(format!("duplicate ID {:?}", entry.key()))); + } + } + } + Ok(map) + } + } + deserializer.deserialize_map(Visitor(std::marker::PhantomData)) +} + +fn optional_unique_map<'de, D: Deserializer<'de>>( + deserializer: D, +) -> Result>, D::Error> { + #[derive(Deserialize)] + struct UniqueMap(#[serde(deserialize_with = "unique_map")] BTreeMap); + + Ok(Option::::deserialize(deserializer)?.map(|map| map.0)) +} diff --git a/crates/protocol/src/lib.rs b/crates/protocol/src/lib.rs index 0d76b4d17..b32dc6c43 100644 --- a/crates/protocol/src/lib.rs +++ b/crates/protocol/src/lib.rs @@ -7,6 +7,7 @@ pub mod category; pub mod client; pub mod codex_namespaces; +pub mod decision; pub mod envelope; pub mod format; pub mod llm; @@ -16,6 +17,7 @@ pub mod stream; pub use category::*; pub use client::*; +pub use decision::*; pub use envelope::*; pub use format::*; pub use llm::*; From 20c2a19f1361fb07b62457156684bd100458022c Mon Sep 17 00:00:00 2001 From: nachiketb Date: Wed, 30 Sep 2026 12:05:51 -0700 Subject: [PATCH 2/4] refactor(protocol): simplify decision types to public data Signed-off-by: nachiketb --- crates/protocol/src/decision.rs | 361 ++------------------------------ 1 file changed, 22 insertions(+), 339 deletions(-) diff --git a/crates/protocol/src/decision.rs b/crates/protocol/src/decision.rs index e49c32b2e..7d1ea28a8 100644 --- a/crates/protocol/src/decision.rs +++ b/crates/protocol/src/decision.rs @@ -3,195 +3,40 @@ //! Provider-neutral decision questions and answers, separate from LLM messages. //! -//! Requests use checked construction and deserialization. Responses must be -//! constructed or decoded with their request to check answer kinds and rubrics. //! Enums use snake-case `type` tags and a `data` payload in serialized form. +//! Fields are public; providers and callers are responsible for valid values. -use std::collections::{BTreeMap, BTreeSet}; +use std::collections::BTreeMap; -use serde::{Deserialize, Deserializer, Serialize, de}; +use serde::{Deserialize, Serialize}; use serde_json::Value; use crate::{ModelId, Usage}; -/// Maximum absolute error in a distribution's sum. Values are never renormalized. -pub const DISTRIBUTION_SUM_TOLERANCE: f64 = 1e-6; - -/// A decision value or message violates the protocol contract. -#[derive(Debug, thiserror::Error)] -#[error("{0}")] -pub struct DecisionError(String); - -/// A finite probability in `[0, 1]`, serialized as a number. +/// Answer probability on a `[0, 1]` scale. #[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] -#[serde(try_from = "f64", into = "f64")] -pub struct Probability(f64); +#[serde(transparent)] +pub struct Probability(pub f64); -impl Probability { - /// Rejects non-finite numbers and values outside `[0, 1]`. - pub fn new(value: f64) -> Result { - if !value.is_finite() || !(0.0..=1.0).contains(&value) { - return Err(DecisionError( - "probability must be finite and in [0, 1]".into(), - )); - } - Ok(Self(value)) - } - - /// Returns the probability as a number. - pub fn get(self) -> f64 { - self.0 - } -} - -impl TryFrom for Probability { - type Error = DecisionError; - - fn try_from(value: f64) -> Result { - Self::new(value) - } -} - -impl From for f64 { - fn from(value: Probability) -> Self { - value.get() - } -} - -/// A finite, nonnegative rubric position, including fractional positions. -/// The response constructor also checks the request's upper bound. +/// Position in the request's rubric, including fractional positions. #[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] -#[serde(try_from = "f64", into = "f64")] -pub struct ScoreValue(f64); - -impl ScoreValue { - /// Rejects non-finite and negative positions. - pub fn new(value: f64) -> Result { - if !value.is_finite() || value < 0.0 { - return Err(DecisionError("score must be finite and nonnegative".into())); - } - Ok(Self(value)) - } - - /// Returns the rubric position. - pub fn get(self) -> f64 { - self.0 - } -} +#[serde(transparent)] +pub struct ScoreValue(pub f64); -impl TryFrom for ScoreValue { - type Error = DecisionError; - - fn try_from(value: f64) -> Result { - Self::new(value) - } -} - -impl From for f64 { - fn from(value: ScoreValue) -> Self { - value.get() - } -} - -/// Finite provider confidence; its scale and meaning are provider-specific. +/// Provider confidence; its scale and meaning are provider-specific. #[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] -#[serde(try_from = "f64", into = "f64")] -pub struct ProviderConfidence(f64); - -impl ProviderConfidence { - /// Rejects non-finite confidence values without imposing a probability scale. - pub fn new(value: f64) -> Result { - if !value.is_finite() { - return Err(DecisionError("provider confidence must be finite".into())); - } - Ok(Self(value)) - } - - /// Returns the provider's confidence value. - pub fn get(self) -> f64 { - self.0 - } -} - -impl TryFrom for ProviderConfidence { - type Error = DecisionError; - - fn try_from(value: f64) -> Result { - Self::new(value) - } -} - -impl From for f64 { - fn from(value: ProviderConfidence) -> Self { - value.get() - } -} +#[serde(transparent)] +pub struct ProviderConfidence(pub f64); /// Shared context evaluated against independent, named questions. -#[derive(Clone, Debug, PartialEq, Serialize)] +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct DecisionRequest { /// Optional until a target is selected. pub model: Option, /// Conversation, application state, or other material to evaluate. pub context: Value, - questions: BTreeMap, -} - -impl DecisionRequest { - /// Requires questions, nonempty unique choice options, and at least two score levels. - pub fn new( - model: Option, - context: Value, - questions: BTreeMap, - ) -> Result { - if questions.is_empty() { - return Err(DecisionError( - "request must contain at least one question".into(), - )); - } - for (id, question) in &questions { - match &question.kind { - DecisionKind::Choice { options } => { - let ids: BTreeSet<_> = options.iter().map(|option| &option.id).collect(); - if options.is_empty() || ids.len() != options.len() { - return Err(DecisionError(format!( - "question {id:?}: choice options must be nonempty with unique IDs" - ))); - } - } - DecisionKind::Score { levels } if levels.len() < 2 => { - return Err(DecisionError(format!( - "question {id:?}: score requires at least two levels" - ))); - } - _ => {} - } - } - Ok(Self { - model, - context, - questions, - }) - } - - /// Questions are immutable so their checked rubrics remain valid. - pub fn questions(&self) -> &BTreeMap { - &self.questions - } -} - -impl<'de> Deserialize<'de> for DecisionRequest { - fn deserialize>(deserializer: D) -> Result { - #[derive(Deserialize)] - struct RawRequest { - model: Option, - context: Value, - #[serde(deserialize_with = "unique_map")] - questions: BTreeMap, - } - let raw = RawRequest::deserialize(deserializer)?; - Self::new(raw.model, raw.context, raw.questions).map_err(de::Error::custom) - } + /// Independent questions keyed by their IDs. + pub questions: BTreeMap, } /// Instructions and the expected answer shape for one question. @@ -199,7 +44,7 @@ impl<'de> Deserialize<'de> for DecisionRequest { pub struct DecisionQuestion { /// Structured or textual instructions shared with the provider. pub instructions: Value, - /// Checked against the answer when constructing a response. + /// Expected answer shape. pub kind: DecisionKind, } @@ -235,76 +80,24 @@ pub struct ChoiceOption { pub description: Option, } -/// Complete answers checked against a request. -/// -/// Use [`Self::deserialize`] with the matching request to decode a response. -/// There is no request-free `Deserialize` implementation. -#[derive(Clone, Debug, PartialEq, Serialize)] +/// Answers keyed by the matching request's question IDs. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct DecisionResponse { /// Provider-reported response identifier. pub id: Option, /// Provider-reported model identifier. pub model: Option, - answers: BTreeMap, + /// Typed answers corresponding to the request's questions. + pub answers: BTreeMap, /// Available token counts; absent counts remain unknown. + #[serde(default)] pub usage: Usage, } -impl DecisionResponse { - /// Checks answer coverage, kinds, selected options, score bounds, and distributions. - /// Score/distribution consistency remains the provider's responsibility. - pub fn new( - request: &DecisionRequest, - id: Option, - model: Option, - answers: BTreeMap, - usage: Usage, - ) -> Result { - if !answers.keys().eq(request.questions.keys()) { - return Err(DecisionError( - "answer IDs must exactly match question IDs".into(), - )); - } - for (id, question) in &request.questions { - check_answer(&question.kind, &answers[id].value) - .map_err(|error| DecisionError(format!("question {id:?}: {error}")))?; - } - Ok(Self { - id, - model, - answers, - usage, - }) - } - - /// Decodes the canonical representation and checks it against its request. - pub fn deserialize<'de, D: Deserializer<'de>>( - request: &DecisionRequest, - deserializer: D, - ) -> Result { - #[derive(Deserialize)] - struct RawResponse { - id: Option, - model: Option, - #[serde(deserialize_with = "unique_map")] - answers: BTreeMap, - #[serde(default)] - usage: Usage, - } - let raw = RawResponse::deserialize(deserializer)?; - Self::new(request, raw.id, raw.model, raw.answers, raw.usage).map_err(de::Error::custom) - } - - /// Answers are immutable after their request-dependent checks. - pub fn answers(&self) -> &BTreeMap { - &self.answers - } -} - /// An answer and separate, optional provider confidence. #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct DecisionAnswer { - /// The estimate, checked against its question when building a response. + /// The estimate for the matching question. pub value: DecisionValue, /// Provider-specific confidence, distinct from answer probabilities. pub provider_confidence: Option, @@ -321,7 +114,6 @@ pub enum DecisionValue { /// Must name an option in the matching question. selected: String, /// Maps every declared option ID to its probability when available. - #[serde(default, deserialize_with = "optional_unique_map")] probabilities: Option>, }, /// A fractional position in the matching request's rubric. @@ -342,112 +134,3 @@ pub enum BooleanEstimate { /// Probability of true; algorithms choose their own thresholds. ProbabilityTrue(Probability), } - -fn check_answer(kind: &DecisionKind, answer: &DecisionValue) -> Result<(), DecisionError> { - match (kind, answer) { - (DecisionKind::Boolean { .. }, DecisionValue::Boolean(_)) => Ok(()), - ( - DecisionKind::Choice { options }, - DecisionValue::Choice { - selected, - probabilities, - }, - ) => { - if !options.iter().any(|option| option.id == *selected) { - return Err(DecisionError( - "selected choice is not a declared option".into(), - )); - } - if let Some(probabilities) = probabilities { - if probabilities.len() != options.len() - || options - .iter() - .any(|option| !probabilities.contains_key(&option.id)) - { - return Err(DecisionError( - "distribution must cover every option exactly".into(), - )); - } - check_sum(probabilities.values())?; - } - Ok(()) - } - ( - DecisionKind::Score { levels }, - DecisionValue::Score { - value, - probabilities, - }, - ) => { - if value.get() > (levels.len() - 1) as f64 { - return Err(DecisionError("score exceeds the request's rubric".into())); - } - if let Some(probabilities) = probabilities { - if probabilities.len() != levels.len() { - return Err(DecisionError( - "distribution must cover every level exactly".into(), - )); - } - check_sum(probabilities.iter())?; - } - Ok(()) - } - _ => Err(DecisionError( - "answer kind does not match question kind".into(), - )), - } -} - -fn check_sum<'a>( - probabilities: impl Iterator, -) -> Result<(), DecisionError> { - let sum: f64 = probabilities.map(|probability| probability.get()).sum(); - if (sum - 1.0).abs() > DISTRIBUTION_SUM_TOLERANCE { - return Err(DecisionError(format!( - "distribution must sum to one within {DISTRIBUTION_SUM_TOLERANCE}" - ))); - } - Ok(()) -} - -// Reject duplicate JSON keys before a map could silently discard an answer or option. -fn unique_map<'de, D, T>(deserializer: D) -> Result, D::Error> -where - D: Deserializer<'de>, - T: Deserialize<'de>, -{ - struct Visitor(std::marker::PhantomData); - - impl<'de, T: Deserialize<'de>> de::Visitor<'de> for Visitor { - type Value = BTreeMap; - - fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result { - formatter.write_str("a map with unique IDs") - } - - fn visit_map>(self, mut access: A) -> Result { - let mut map = BTreeMap::new(); - while let Some((key, value)) = access.next_entry::()? { - match map.entry(key) { - std::collections::btree_map::Entry::Vacant(entry) => { - entry.insert(value); - } - std::collections::btree_map::Entry::Occupied(entry) => { - return Err(de::Error::custom(format!("duplicate ID {:?}", entry.key()))); - } - } - } - Ok(map) - } - } - deserializer.deserialize_map(Visitor(std::marker::PhantomData)) -} - -fn optional_unique_map<'de, D: Deserializer<'de>>( - deserializer: D, -) -> Result>, D::Error> { - #[derive(Deserialize)] - struct UniqueMap(#[serde(deserialize_with = "unique_map")] BTreeMap); - - Ok(Option::::deserialize(deserializer)?.map(|map| map.0)) -} From 067eb8be241b257cdd77139376e48d542111647b Mon Sep 17 00:00:00 2001 From: nachiketb Date: Wed, 30 Sep 2026 12:08:55 -0700 Subject: [PATCH 3/4] fix(protocol): check decision numbers during deserialization Signed-off-by: nachiketb --- crates/protocol/src/decision.rs | 37 ++++++++++++++++++++++++++++----- 1 file changed, 32 insertions(+), 5 deletions(-) diff --git a/crates/protocol/src/decision.rs b/crates/protocol/src/decision.rs index 7d1ea28a8..f56406725 100644 --- a/crates/protocol/src/decision.rs +++ b/crates/protocol/src/decision.rs @@ -4,11 +4,12 @@ //! Provider-neutral decision questions and answers, separate from LLM messages. //! //! Enums use snake-case `type` tags and a `data` payload in serialized form. -//! Fields are public; providers and callers are responsible for valid values. +//! Numeric bounds are checked during deserialization. Public fields allow direct +//! construction; callers are responsible for valid values in that case. use std::collections::BTreeMap; -use serde::{Deserialize, Serialize}; +use serde::{Deserialize, Deserializer, Serialize, de}; use serde_json::Value; use crate::{ModelId, Usage}; @@ -16,17 +17,43 @@ use crate::{ModelId, Usage}; /// Answer probability on a `[0, 1]` scale. #[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] #[serde(transparent)] -pub struct Probability(pub f64); +pub struct Probability(#[serde(deserialize_with = "deserialize_probability")] pub f64); /// Position in the request's rubric, including fractional positions. #[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] #[serde(transparent)] -pub struct ScoreValue(pub f64); +pub struct ScoreValue(#[serde(deserialize_with = "deserialize_score")] pub f64); /// Provider confidence; its scale and meaning are provider-specific. #[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] #[serde(transparent)] -pub struct ProviderConfidence(pub f64); +pub struct ProviderConfidence(#[serde(deserialize_with = "deserialize_confidence")] pub f64); + +fn deserialize_probability<'de, D: Deserializer<'de>>(deserializer: D) -> Result { + let value = f64::deserialize(deserializer)?; + if !(0.0..=1.0).contains(&value) { + return Err(de::Error::custom( + "probability must be finite and in [0, 1]", + )); + } + Ok(value) +} + +fn deserialize_score<'de, D: Deserializer<'de>>(deserializer: D) -> Result { + let value = f64::deserialize(deserializer)?; + if !value.is_finite() || value < 0.0 { + return Err(de::Error::custom("score must be finite and nonnegative")); + } + Ok(value) +} + +fn deserialize_confidence<'de, D: Deserializer<'de>>(deserializer: D) -> Result { + let value = f64::deserialize(deserializer)?; + if !value.is_finite() { + return Err(de::Error::custom("provider confidence must be finite")); + } + Ok(value) +} /// Shared context evaluated against independent, named questions. #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] From aa0192ac1dd99151fd101c63c258091c09840078 Mon Sep 17 00:00:00 2001 From: nachiketb Date: Wed, 30 Sep 2026 12:13:06 -0700 Subject: [PATCH 4/4] refactor(protocol): remove decision numeric decoding checks Signed-off-by: nachiketb --- crates/protocol/src/decision.rs | 37 +++++---------------------------- 1 file changed, 5 insertions(+), 32 deletions(-) diff --git a/crates/protocol/src/decision.rs b/crates/protocol/src/decision.rs index f56406725..7d1ea28a8 100644 --- a/crates/protocol/src/decision.rs +++ b/crates/protocol/src/decision.rs @@ -4,12 +4,11 @@ //! Provider-neutral decision questions and answers, separate from LLM messages. //! //! Enums use snake-case `type` tags and a `data` payload in serialized form. -//! Numeric bounds are checked during deserialization. Public fields allow direct -//! construction; callers are responsible for valid values in that case. +//! Fields are public; providers and callers are responsible for valid values. use std::collections::BTreeMap; -use serde::{Deserialize, Deserializer, Serialize, de}; +use serde::{Deserialize, Serialize}; use serde_json::Value; use crate::{ModelId, Usage}; @@ -17,43 +16,17 @@ use crate::{ModelId, Usage}; /// Answer probability on a `[0, 1]` scale. #[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] #[serde(transparent)] -pub struct Probability(#[serde(deserialize_with = "deserialize_probability")] pub f64); +pub struct Probability(pub f64); /// Position in the request's rubric, including fractional positions. #[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] #[serde(transparent)] -pub struct ScoreValue(#[serde(deserialize_with = "deserialize_score")] pub f64); +pub struct ScoreValue(pub f64); /// Provider confidence; its scale and meaning are provider-specific. #[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)] #[serde(transparent)] -pub struct ProviderConfidence(#[serde(deserialize_with = "deserialize_confidence")] pub f64); - -fn deserialize_probability<'de, D: Deserializer<'de>>(deserializer: D) -> Result { - let value = f64::deserialize(deserializer)?; - if !(0.0..=1.0).contains(&value) { - return Err(de::Error::custom( - "probability must be finite and in [0, 1]", - )); - } - Ok(value) -} - -fn deserialize_score<'de, D: Deserializer<'de>>(deserializer: D) -> Result { - let value = f64::deserialize(deserializer)?; - if !value.is_finite() || value < 0.0 { - return Err(de::Error::custom("score must be finite and nonnegative")); - } - Ok(value) -} - -fn deserialize_confidence<'de, D: Deserializer<'de>>(deserializer: D) -> Result { - let value = f64::deserialize(deserializer)?; - if !value.is_finite() { - return Err(de::Error::custom("provider confidence must be finite")); - } - Ok(value) -} +pub struct ProviderConfidence(pub f64); /// Shared context evaluated against independent, named questions. #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]