rig_core/completion/message/
citation.rs1use std::ops::Range;
26
27use serde::{Deserialize, Serialize};
28
29use super::{Fingerprint, Text};
30use crate::wire::{SpanUnit, WireCitation, WireSpan};
31
32#[non_exhaustive]
34#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
35pub struct Citation {
36 #[serde(default, skip_serializing_if = "Option::is_none")]
38 pub span: Option<Span>,
39 #[serde(default)]
41 pub sources: Vec<Source>,
42}
43
44impl Citation {
45 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#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
59pub struct Span {
60 start: usize,
61 end: usize,
62}
63
64impl Span {
65 pub fn start(&self) -> usize {
67 self.start
68 }
69
70 pub fn end(&self) -> usize {
72 self.end
73 }
74
75 pub fn range(&self) -> Range<usize> {
77 self.start..self.end
78 }
79}
80
81#[non_exhaustive]
83#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
84pub struct Source {
85 pub location: SourceLocation,
87 #[serde(default, skip_serializing_if = "Option::is_none")]
89 pub title: Option<String>,
90 #[serde(default, skip_serializing_if = "Option::is_none")]
92 pub cited_text: Option<String>,
93 #[serde(default, skip_serializing_if = "Option::is_none")]
95 pub confidence: Option<f32>,
96}
97
98impl Source {
99 pub fn new(location: SourceLocation) -> Self {
101 Self {
102 location,
103 title: None,
104 cited_text: None,
105 confidence: None,
106 }
107 }
108
109 pub fn title(mut self, title: impl Into<String>) -> Self {
111 self.title = Some(title.into());
112 self
113 }
114
115 pub fn cited_text(mut self, text: impl Into<String>) -> Self {
117 self.cited_text = Some(text.into());
118 self
119 }
120
121 pub fn confidence(mut self, confidence: f32) -> Self {
123 self.confidence = Some(confidence);
124 self
125 }
126}
127
128#[non_exhaustive]
130#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
131#[serde(tag = "type", rename_all = "snake_case")]
132pub enum SourceLocation {
133 Document {
135 #[serde(default, skip_serializing_if = "Option::is_none")]
137 index: Option<u32>,
138 #[serde(default, skip_serializing_if = "Option::is_none")]
140 id: Option<String>,
141 #[serde(default, skip_serializing_if = "Option::is_none")]
143 within: Option<DocumentRange>,
144 },
145 Url {
147 url: String,
149 },
150 File {
152 file_id: String,
154 #[serde(default, skip_serializing_if = "Option::is_none")]
156 filename: Option<String>,
157 #[serde(default, skip_serializing_if = "Option::is_none")]
159 container_id: Option<String>,
160 },
161 SearchResult {
163 index: u32,
165 source: String,
167 #[serde(default, skip_serializing_if = "Option::is_none")]
169 blocks: Option<Range<u32>>,
170 },
171 ToolOutput {
173 id: String,
175 },
176}
177
178#[non_exhaustive]
180#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
181#[serde(tag = "unit", rename_all = "snake_case")]
182pub enum DocumentRange {
183 Chars(Range<u64>),
185 Pages(Range<u32>),
187 Blocks(Range<u32>),
189}
190
191#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
193pub(super) struct Citations {
194 fingerprint: Fingerprint,
195 list: Vec<Citation>,
196}
197
198pub(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
210fn 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 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 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 pub fn clear_citations(&mut self) {
250 self.citations = None;
251 }
252
253 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 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 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
294pub(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
321fn 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
348fn 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
356fn 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;