Skip to main content

rig_core/completion/message/
citation.rs

1//! What a provider says supports an answer text. A [`Citation`] lives on the
2//! [`Text`] it cites, with a byte [`Span`] into that text or none for the
3//! whole block, and the [`Source`]s behind it. The list is bound to the text
4//! it was resolved against: once the text is edited, [`Text::citations`]
5//! reads empty.
6//!
7//! Decoders never build a span: they hand the fold a
8//! [`WireCitation`] whose span names its unit, and
9//! the fold resolves it to bytes when the text block closes.
10//!
11//! ```
12//! use rig_core::message::{Citation, Source, SourceLocation, Text};
13//!
14//! let text = Text::new("Dock Seven is open.");
15//! let span = text.span(0..10).ok_or("on a character boundary")?;
16//! let mut citation = Citation::new([Source::new(SourceLocation::Url {
17//!     url: "https://example.com/docks".to_owned(),
18//! })]);
19//! citation.span = Some(span);
20//! let text = text.with_citations([citation]);
21//! assert_eq!(text.cited(&text.citations()[0]), Some("Dock Seven"));
22//! # Ok::<(), &str>(())
23//! ```
24
25use std::ops::Range;
26
27use serde::{Deserialize, Serialize};
28
29use super::{Fingerprint, Text};
30use crate::wire::{SpanUnit, WireCitation, WireSpan};
31
32/// A claim in a text and the sources that support it.
33#[non_exhaustive]
34#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
35pub struct Citation {
36    /// The bytes of the text it covers; `None` for the whole block.
37    #[serde(default, skip_serializing_if = "Option::is_none")]
38    pub span: Option<Span>,
39    /// What supports the claim.
40    #[serde(default)]
41    pub sources: Vec<Source>,
42}
43
44impl Citation {
45    /// A citation of the whole block by `sources`.
46    pub fn new(sources: impl IntoIterator<Item = Source>) -> Self {
47        Self {
48            span: None,
49            sources: sources.into_iter().collect(),
50        }
51    }
52}
53
54/// A byte range of a text, on character boundaries. The fold and
55/// [`Text::span`] make one, and stored history deserializes one;
56/// [`Text::with_citations`] drops any that is off a character boundary of
57/// its text.
58#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
59pub struct Span {
60    start: usize,
61    end: usize,
62}
63
64impl Span {
65    /// The first byte.
66    pub fn start(&self) -> usize {
67        self.start
68    }
69
70    /// The byte after the last.
71    pub fn end(&self) -> usize {
72        self.end
73    }
74
75    /// `start..end`.
76    pub fn range(&self) -> Range<usize> {
77        self.start..self.end
78    }
79}
80
81/// One source of a citation.
82#[non_exhaustive]
83#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
84pub struct Source {
85    /// Where the source is.
86    pub location: SourceLocation,
87    /// The source's title.
88    #[serde(default, skip_serializing_if = "Option::is_none")]
89    pub title: Option<String>,
90    /// The passage of the source the provider quoted.
91    #[serde(default, skip_serializing_if = "Option::is_none")]
92    pub cited_text: Option<String>,
93    /// The provider's confidence, 0 to 1.
94    #[serde(default, skip_serializing_if = "Option::is_none")]
95    pub confidence: Option<f32>,
96}
97
98impl Source {
99    /// A source at `location` with no title, quote or confidence.
100    pub fn new(location: SourceLocation) -> Self {
101        Self {
102            location,
103            title: None,
104            cited_text: None,
105            confidence: None,
106        }
107    }
108
109    /// Set the title.
110    pub fn title(mut self, title: impl Into<String>) -> Self {
111        self.title = Some(title.into());
112        self
113    }
114
115    /// Set the quoted passage of the source.
116    pub fn cited_text(mut self, text: impl Into<String>) -> Self {
117        self.cited_text = Some(text.into());
118        self
119    }
120
121    /// Set the provider's confidence.
122    pub fn confidence(mut self, confidence: f32) -> Self {
123        self.confidence = Some(confidence);
124        self
125    }
126}
127
128/// Where a cited source is.
129#[non_exhaustive]
130#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
131#[serde(tag = "type", rename_all = "snake_case")]
132pub enum SourceLocation {
133    /// A document of the request, by position or id, and where in it.
134    Document {
135        /// The document's position in the request.
136        #[serde(default, skip_serializing_if = "Option::is_none")]
137        index: Option<u32>,
138        /// The document's id.
139        #[serde(default, skip_serializing_if = "Option::is_none")]
140        id: Option<String>,
141        /// The part of the document cited.
142        #[serde(default, skip_serializing_if = "Option::is_none")]
143        within: Option<DocumentRange>,
144    },
145    /// A web page.
146    Url {
147        /// The page's address.
148        url: String,
149    },
150    /// A file the provider stores.
151    File {
152        /// The provider's file id.
153        file_id: String,
154        /// The file's name.
155        #[serde(default, skip_serializing_if = "Option::is_none")]
156        filename: Option<String>,
157        /// The container holding the file.
158        #[serde(default, skip_serializing_if = "Option::is_none")]
159        container_id: Option<String>,
160    },
161    /// A search result the request supplied.
162    SearchResult {
163        /// The result's position in the request.
164        index: u32,
165        /// The result's source, as the request named it.
166        source: String,
167        /// The content blocks of the result cited.
168        #[serde(default, skip_serializing_if = "Option::is_none")]
169        blocks: Option<Range<u32>>,
170    },
171    /// A tool's output.
172    ToolOutput {
173        /// The tool call or output id.
174        id: String,
175    },
176}
177
178/// The part of a document a citation covers.
179#[non_exhaustive]
180#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
181#[serde(tag = "unit", rename_all = "snake_case")]
182pub enum DocumentRange {
183    /// Characters of the document's text.
184    Chars(Range<u64>),
185    /// Pages, from 1, the end exclusive.
186    Pages(Range<u32>),
187    /// Content blocks or chunks.
188    Blocks(Range<u32>),
189}
190
191/// The citations of one text and the fingerprint of the text they fit.
192#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
193pub(super) struct Citations {
194    fingerprint: Fingerprint,
195    list: Vec<Citation>,
196}
197
198/// A stored `citations` value, or `None` when it cannot be read.
199pub(super) fn lenient<'de, D: serde::Deserializer<'de>>(
200    deserializer: D,
201) -> Result<Option<Citations>, D::Error> {
202    let value = Option::<serde_json::Value>::deserialize(deserializer)?;
203    Ok(value.and_then(|value| Citations::deserialize(value).ok()))
204}
205
206fn fingerprint(text: &str) -> Fingerprint {
207    Fingerprint::of(&text)
208}
209
210/// Whether `citation`'s span lies on character boundaries of `text`.
211fn fits(text: &str, citation: &Citation) -> bool {
212    citation
213        .span
214        .is_none_or(|span| text.get(span.range()).is_some())
215}
216
217impl Text {
218    /// The citations, while `text` is what they were resolved against;
219    /// empty once it has been edited.
220    pub fn citations(&self) -> &[Citation] {
221        match &self.citations {
222            Some(citations) if citations.fingerprint == fingerprint(&self.text) => &citations.list,
223            _ => &[],
224        }
225    }
226
227    /// This text with `citations`, bound to its current text. A citation
228    /// whose span is off a character boundary or past the end is dropped
229    /// with a warning.
230    pub fn with_citations(mut self, citations: impl IntoIterator<Item = Citation>) -> Self {
231        let list = citations
232            .into_iter()
233            .filter(|citation| {
234                let fits = fits(&self.text, citation);
235                if !fits {
236                    tracing::warn!(
237                        span = ?citation.span,
238                        "dropped a citation whose span does not fit its text"
239                    );
240                }
241                fits
242            })
243            .collect();
244        self.bind(list);
245        self
246    }
247
248    /// Remove every citation.
249    pub fn clear_citations(&mut self) {
250        self.citations = None;
251    }
252
253    /// The text `citation` covers: its span, or the whole block. It reads
254    /// any citation against the current text and does not check that the
255    /// citation belongs to it, so a citation kept from before an edit gives
256    /// whatever text now sits at its span, or `None` past the end. Read
257    /// citations through [`Text::citations`], which hides them once the
258    /// text changes.
259    pub fn cited(&self, citation: &Citation) -> Option<&str> {
260        match citation.span {
261            Some(span) => self.text.get(span.range()),
262            None => Some(&self.text),
263        }
264    }
265
266    /// A span over `range`, in bytes of this text; `None` when either end
267    /// is off a character boundary or past the end, or `range` is reversed.
268    pub fn span(&self, range: Range<usize>) -> Option<Span> {
269        self.text.get(range.clone()).map(|_| Span {
270            start: range.start,
271            end: range.end,
272        })
273    }
274
275    fn bind(&mut self, list: Vec<Citation>) {
276        self.citations = (!list.is_empty()).then(|| Citations {
277            fingerprint: fingerprint(&self.text),
278            list,
279        });
280    }
281
282    /// Stored citations whose spans do not fit the stored text are dropped.
283    pub(super) fn checked(mut self) -> Self {
284        if let Some(citations) = &mut self.citations {
285            citations.list.retain(|citation| fits(&self.text, citation));
286            if citations.list.is_empty() {
287                self.citations = None;
288            }
289        }
290        self
291    }
292}
293
294/// Attach `wire` to `text`, the block the provider's item at `index` became:
295/// each span resolved to bytes of the text, after `kept`, the citations the
296/// block holds already. A span that does not resolve, or covers other text
297/// than the provider quoted, drops its citation with a warning; a citation
298/// never fails a reply. The completion fold is the only caller.
299pub(crate) fn attach(
300    text: &mut Text,
301    kept: Vec<Citation>,
302    wire: Vec<WireCitation>,
303    provider: &str,
304    index: usize,
305) {
306    let mut list = kept;
307    for citation in wire {
308        let WireCitation { span, sources } = citation;
309        let span = match span.map(|span| resolve(&text.text, &span)).transpose() {
310            Ok(span) => span,
311            Err(reason) => {
312                tracing::warn!(provider, index, reason, "dropped a citation");
313                continue;
314            }
315        };
316        list.push(Citation { span, sources });
317    }
318    text.bind(list);
319}
320
321/// The byte span `span` names in `text`.
322fn resolve(text: &str, span: &WireSpan) -> Result<Span, &'static str> {
323    let wire = |offset: u64| usize::try_from(offset).ok();
324    let (Some(start), Some(end)) = (wire(span.start), wire(span.end)) else {
325        return Err("the span is past the end of the text");
326    };
327    let byte = |offset: usize| match span.unit {
328        SpanUnit::Bytes => Some(offset),
329        SpanUnit::Chars => char_offset(text, offset),
330        SpanUnit::Utf16 => utf16_offset(text, offset),
331    };
332    let (Some(start), Some(end)) = (byte(start), byte(end)) else {
333        return Err("the span is past the end of the text or splits a character");
334    };
335    let Some(covered) = text.get(start..end) else {
336        return Err("the span is reversed, past the end of the text or splits a character");
337    };
338    if span
339        .quoted
340        .as_deref()
341        .is_some_and(|quoted| quoted != covered)
342    {
343        return Err("the span covers other text than the provider quoted");
344    }
345    Ok(Span { start, end })
346}
347
348/// The byte offset of the `chars`th character of `text`.
349fn char_offset(text: &str, chars: usize) -> Option<usize> {
350    text.char_indices()
351        .map(|(at, _)| at)
352        .chain(std::iter::once(text.len()))
353        .nth(chars)
354}
355
356/// The byte offset of the `units`th UTF-16 code unit of `text`, if it
357/// starts a character.
358fn utf16_offset(text: &str, units: usize) -> Option<usize> {
359    let mut seen = 0;
360    for (at, character) in text.char_indices() {
361        if seen == units {
362            return Some(at);
363        }
364        if seen > units {
365            return None;
366        }
367        seen += character.len_utf16();
368    }
369    (seen == units).then_some(text.len())
370}
371
372#[cfg(test)]
373mod tests;