use std::ops::Range;
use serde::{Deserialize, Serialize};
use super::{Fingerprint, Text};
use crate::wire::{SpanUnit, WireCitation, WireSpan};
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct Citation {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub span: Option<Span>,
#[serde(default)]
pub sources: Vec<Source>,
}
impl Citation {
pub fn new(sources: impl IntoIterator<Item = Source>) -> Self {
Self {
span: None,
sources: sources.into_iter().collect(),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct Span {
start: usize,
end: usize,
}
impl Span {
pub fn start(&self) -> usize {
self.start
}
pub fn end(&self) -> usize {
self.end
}
pub fn range(&self) -> Range<usize> {
self.start..self.end
}
}
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct Source {
pub location: SourceLocation,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub title: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cited_text: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub confidence: Option<f32>,
}
impl Source {
pub fn new(location: SourceLocation) -> Self {
Self {
location,
title: None,
cited_text: None,
confidence: None,
}
}
pub fn title(mut self, title: impl Into<String>) -> Self {
self.title = Some(title.into());
self
}
pub fn cited_text(mut self, text: impl Into<String>) -> Self {
self.cited_text = Some(text.into());
self
}
pub fn confidence(mut self, confidence: f32) -> Self {
self.confidence = Some(confidence);
self
}
}
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum SourceLocation {
Document {
#[serde(default, skip_serializing_if = "Option::is_none")]
index: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
within: Option<DocumentRange>,
},
Url {
url: String,
},
File {
file_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
filename: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
container_id: Option<String>,
},
SearchResult {
index: u32,
source: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
blocks: Option<Range<u32>>,
},
ToolOutput {
id: String,
},
}
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(tag = "unit", rename_all = "snake_case")]
pub enum DocumentRange {
Chars(Range<u64>),
Pages(Range<u32>),
Blocks(Range<u32>),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub(super) struct Citations {
fingerprint: Fingerprint,
list: Vec<Citation>,
}
pub(super) fn lenient<'de, D: serde::Deserializer<'de>>(
deserializer: D,
) -> Result<Option<Citations>, D::Error> {
let value = Option::<serde_json::Value>::deserialize(deserializer)?;
Ok(value.and_then(|value| Citations::deserialize(value).ok()))
}
fn fingerprint(text: &str) -> Fingerprint {
Fingerprint::of(&text)
}
fn fits(text: &str, citation: &Citation) -> bool {
citation
.span
.is_none_or(|span| text.get(span.range()).is_some())
}
impl Text {
pub fn citations(&self) -> &[Citation] {
match &self.citations {
Some(citations) if citations.fingerprint == fingerprint(&self.text) => &citations.list,
_ => &[],
}
}
pub fn with_citations(mut self, citations: impl IntoIterator<Item = Citation>) -> Self {
let list = citations
.into_iter()
.filter(|citation| {
let fits = fits(&self.text, citation);
if !fits {
tracing::warn!(
span = ?citation.span,
"dropped a citation whose span does not fit its text"
);
}
fits
})
.collect();
self.bind(list);
self
}
pub fn clear_citations(&mut self) {
self.citations = None;
}
pub fn cited(&self, citation: &Citation) -> Option<&str> {
match citation.span {
Some(span) => self.text.get(span.range()),
None => Some(&self.text),
}
}
pub fn span(&self, range: Range<usize>) -> Option<Span> {
self.text.get(range.clone()).map(|_| Span {
start: range.start,
end: range.end,
})
}
fn bind(&mut self, list: Vec<Citation>) {
self.citations = (!list.is_empty()).then(|| Citations {
fingerprint: fingerprint(&self.text),
list,
});
}
pub(super) fn checked(mut self) -> Self {
if let Some(citations) = &mut self.citations {
citations.list.retain(|citation| fits(&self.text, citation));
if citations.list.is_empty() {
self.citations = None;
}
}
self
}
}
pub(crate) fn attach(
text: &mut Text,
kept: Vec<Citation>,
wire: Vec<WireCitation>,
provider: &str,
index: usize,
) {
let mut list = kept;
for citation in wire {
let WireCitation { span, sources } = citation;
let span = match span.map(|span| resolve(&text.text, &span)).transpose() {
Ok(span) => span,
Err(reason) => {
tracing::warn!(provider, index, reason, "dropped a citation");
continue;
}
};
list.push(Citation { span, sources });
}
text.bind(list);
}
fn resolve(text: &str, span: &WireSpan) -> Result<Span, &'static str> {
let wire = |offset: u64| usize::try_from(offset).ok();
let (Some(start), Some(end)) = (wire(span.start), wire(span.end)) else {
return Err("the span is past the end of the text");
};
let byte = |offset: usize| match span.unit {
SpanUnit::Bytes => Some(offset),
SpanUnit::Chars => char_offset(text, offset),
SpanUnit::Utf16 => utf16_offset(text, offset),
};
let (Some(start), Some(end)) = (byte(start), byte(end)) else {
return Err("the span is past the end of the text or splits a character");
};
let Some(covered) = text.get(start..end) else {
return Err("the span is reversed, past the end of the text or splits a character");
};
if span
.quoted
.as_deref()
.is_some_and(|quoted| quoted != covered)
{
return Err("the span covers other text than the provider quoted");
}
Ok(Span { start, end })
}
fn char_offset(text: &str, chars: usize) -> Option<usize> {
text.char_indices()
.map(|(at, _)| at)
.chain(std::iter::once(text.len()))
.nth(chars)
}
fn utf16_offset(text: &str, units: usize) -> Option<usize> {
let mut seen = 0;
for (at, character) in text.char_indices() {
if seen == units {
return Some(at);
}
if seen > units {
return None;
}
seen += character.len_utf16();
}
(seen == units).then_some(text.len())
}
#[cfg(test)]
mod tests;