Skip to main content

rig_core/completion/message/
native.rs

1//! Where an assistant turn came from and the provider items it was decoded
2//! from. An [`Origin`] names the wire, provider and requested model of a
3//! turn; a [`Native`] holds one provider item verbatim beside the canonical
4//! block it decoded to, with a [`Fingerprint`] of that block so an edited
5//! block stops replaying its stale item.
6//!
7//! ```
8//! use rig_core::message::{AssistantContent, Text};
9//!
10//! let block = AssistantContent::Text(Text::new("hi"))
11//!     .with_native(serde_json::json!({"type": "text", "text": "hi", "citations": []}));
12//! assert!(block.native_item().is_some());
13//! ```
14
15use std::borrow::Cow;
16
17use serde::{Deserialize, Serialize};
18
19/// The wire format a turn was produced by, for example
20/// `"anthropic.messages"`, `"openai.responses"` or `"openai.chat"`.
21#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
22#[serde(transparent)]
23pub struct Api(Cow<'static, str>);
24
25impl Api {
26    /// The API named `name`, in a const context.
27    pub const fn from_static(name: &'static str) -> Self {
28        Self(Cow::Borrowed(name))
29    }
30
31    /// The API's name.
32    pub fn as_str(&self) -> &str {
33        &self.0
34    }
35}
36
37impl From<&'static str> for Api {
38    fn from(name: &'static str) -> Self {
39        Self::from_static(name)
40    }
41}
42
43impl From<String> for Api {
44    fn from(name: String) -> Self {
45        Self(Cow::Owned(name))
46    }
47}
48
49impl std::fmt::Display for Api {
50    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
51        f.write_str(&self.0)
52    }
53}
54
55/// Which wire, provider and model produced an assistant turn.
56///
57/// `model` is the model the request named. Replay compares it, with `api`
58/// and `provider`, against the target: only an exact match replays the
59/// turn's provider items.
60#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
61pub struct Origin {
62    /// The wire format.
63    pub api: Api,
64    /// The provider descriptor name (`"anthropic"`).
65    pub provider: String,
66    /// The model the request named.
67    pub model: String,
68    /// The model the provider reported, when it reported one.
69    #[serde(default, skip_serializing_if = "Option::is_none")]
70    pub response_model: Option<String>,
71    /// The provider's response id, when it sent one.
72    #[serde(default, skip_serializing_if = "Option::is_none")]
73    pub response_id: Option<String>,
74    /// The fingerprint of the request's tools and system prompt. A wire
75    /// that binds its items to them ([`ReplayTarget::binds_context`])
76    /// replays a turn made under another context as if from another model.
77    ///
78    /// [`ReplayTarget::binds_context`]: crate::completion::ReplayTarget::binds_context
79    #[serde(default, skip_serializing_if = "Option::is_none")]
80    pub context: Option<Fingerprint>,
81}
82
83impl Origin {
84    /// A turn from `model` on `provider` over `api`, with no response
85    /// metadata.
86    pub fn new(api: impl Into<Api>, provider: impl Into<String>, model: impl Into<String>) -> Self {
87        Self {
88            api: api.into(),
89            provider: provider.into(),
90            model: model.into(),
91            response_model: None,
92            response_id: None,
93            context: None,
94        }
95    }
96
97    /// Whether this turn came from exactly the wire, provider and model
98    /// `target` names. The one sameness rule replay uses.
99    pub fn same_model(&self, api: &Api, provider: &str, model: &str) -> bool {
100        &self.api == api && self.provider == provider && self.model == model
101    }
102}
103
104/// How an assistant turn ended.
105///
106/// A turn that ended in [`Self::Error`] or [`Self::Aborted`] is kept in
107/// history but never replayed to a model.
108#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
109#[serde(rename_all = "snake_case")]
110#[non_exhaustive]
111pub enum StopReason {
112    /// The model finished.
113    Stop,
114    /// The output-token limit cut the turn.
115    Length,
116    /// The model stopped to call tools.
117    ToolUse,
118    /// The provider failed the turn or refused it, with its explanation.
119    Error(String),
120    /// The caller cancelled the turn.
121    Aborted(String),
122}
123
124impl StopReason {
125    /// Whether the turn is incomplete and must not be replayed.
126    pub fn is_failure(&self) -> bool {
127        matches!(self, Self::Error(_) | Self::Aborted(_))
128    }
129}
130
131/// A 64-bit FNV-1a hash of a value's JSON serialization with every object's
132/// keys in sorted order and every whole number written as an integer, so a
133/// store that reorders keys (Postgres `jsonb`, sorted-key dumps) or writes
134/// `20.0` as `20` never changes it. It is stored as 16 hex digits, since a
135/// store that reads JSON numbers as doubles would round a 64-bit number.
136#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
137pub struct Fingerprint(u64);
138
139impl Serialize for Fingerprint {
140    fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
141        serializer.serialize_str(&format!("{:016x}", self.0))
142    }
143}
144
145impl<'de> Deserialize<'de> for Fingerprint {
146    fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
147        let text = std::borrow::Cow::<'de, str>::deserialize(deserializer)?;
148        u64::from_str_radix(&text, 16)
149            .map(Self)
150            .map_err(|_| serde::de::Error::custom("a fingerprint is 16 hex digits"))
151    }
152}
153
154impl Fingerprint {
155    /// The fingerprint of `value`'s JSON bytes, keys sorted and whole
156    /// numbers as integers.
157    pub fn of(value: &impl Serialize) -> Self {
158        fn sorted(value: serde_json::Value) -> serde_json::Value {
159            match value {
160                serde_json::Value::Object(fields) => {
161                    let mut fields: Vec<_> = fields.into_iter().collect();
162                    fields.sort_by(|(left, _), (right, _)| left.cmp(right));
163                    serde_json::Value::Object(
164                        fields
165                            .into_iter()
166                            .map(|(key, value)| (key, sorted(value)))
167                            .collect(),
168                    )
169                }
170                serde_json::Value::Array(values) => {
171                    serde_json::Value::Array(values.into_iter().map(sorted).collect())
172                }
173                serde_json::Value::Number(number) => number
174                    .as_f64()
175                    .filter(|float| number.is_f64() && float.fract() == 0.0)
176                    .and_then(|float| format!("{float:.0}").parse::<serde_json::Number>().ok())
177                    .map_or(serde_json::Value::Number(number), serde_json::Value::Number),
178                value => value,
179            }
180        }
181        struct Fnv(u64);
182        impl std::io::Write for Fnv {
183            fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
184                for byte in bytes {
185                    self.0 ^= u64::from(*byte);
186                    self.0 = self.0.wrapping_mul(0x0000_0100_0000_01b3);
187                }
188                Ok(bytes.len())
189            }
190
191            fn flush(&mut self) -> std::io::Result<()> {
192                Ok(())
193            }
194        }
195        let mut hash = Fnv(0xcbf2_9ce4_8422_2325);
196        // Message types always serialize; a failure would only leave a
197        // fingerprint no native item matches, which replays canonically.
198        if let Ok(value) = serde_json::to_value(value) {
199            let _ = serde_json::to_writer(&mut hash, &sorted(value));
200        }
201        Self(hash.0)
202    }
203}
204
205/// One provider item, verbatim in its API's JSON shape, beside the
206/// canonical block (or message) it was decoded to.
207///
208/// `fingerprint` is the canonical form's at decode time. When the block is
209/// edited its fingerprint changes, the item is stale, and encoders rebuild
210/// the wire item from the canonical fields instead.
211#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
212pub struct Native {
213    /// The provider item.
214    pub item: serde_json::Value,
215    /// The fingerprint of the canonical form the item was decoded to.
216    pub fingerprint: Fingerprint,
217}
218
219/// A block's stored `native`, or `None` when it cannot be read, such as one
220/// fingerprinted by an earlier projection: the block then loads with its
221/// canonical fields and replays from them.
222pub(crate) fn lenient<'de, D: serde::Deserializer<'de>>(
223    deserializer: D,
224) -> Result<Option<Native>, D::Error> {
225    let value = Option::<serde_json::Value>::deserialize(deserializer)?;
226    Ok(value.and_then(|value| Native::deserialize(value).ok()))
227}
228
229/// A provider item with no canonical meaning, such as a hosted-tool step or
230/// a compaction record. Only the API that produced it reads it.
231#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
232pub struct Opaque {
233    /// The provider item.
234    pub item: serde_json::Value,
235    /// Whether the item goes back to the model that produced it. A
236    /// client-executed call nothing answers is kept but never sent.
237    pub replay: bool,
238}
239
240impl Opaque {
241    /// The item's `type` field, when it is an object that has one.
242    pub fn kind(&self) -> Option<&str> {
243        self.item.get("type").and_then(serde_json::Value::as_str)
244    }
245}
246
247#[cfg(test)]
248mod tests;