use std::collections::BTreeSet;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::{ChatMessage, Role};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Fidelity {
ByteLossless,
ValueLossless,
Semantic,
}
impl Fidelity {
pub fn tolerates_residue(self) -> bool {
matches!(self, Self::Semantic)
}
}
pub fn messages_equal(a: &ChatMessage, b: &ChatMessage) -> bool {
if a.role != b.role || a.content != b.content || a.tool_call_id != b.tool_call_id {
return false;
}
let (a_calls, b_calls) = (a.tool_calls(), b.tool_calls());
a_calls.len() == b_calls.len()
&& a_calls.iter().zip(b_calls).all(|(a_call, b_call)| {
a_call.id == b_call.id
&& a_call.function.name == b_call.function.name
&& a_call.function.parsed_arguments().ok()
== b_call.function.parsed_arguments().ok()
})
}
pub fn messages_equal_multimodal(a: &ChatMessage, b: &ChatMessage) -> bool {
if !messages_equal(a, b) || a.name != b.name {
return false;
}
let empty = Vec::new();
let a_parts = a.content_parts.as_ref().unwrap_or(&empty);
let b_parts = b.content_parts.as_ref().unwrap_or(&empty);
a_parts.len() == b_parts.len()
&& a_parts
.iter()
.zip(b_parts)
.all(|(a_part, b_part)| normalize_part(a_part) == normalize_part(b_part))
}
fn normalize_part(part: &Value) -> (String, Option<String>, Option<Vec<u8>>) {
let kind = part
.get("type")
.and_then(Value::as_str)
.unwrap_or("")
.to_owned();
let url = part
.get("image_url")
.and_then(|value| value.get("url"))
.and_then(Value::as_str);
match url {
Some(url) if url.starts_with("data:") => {
let rest = &url["data:".len()..];
let (metadata, data) = rest.split_once(',').unwrap_or((rest, ""));
let mime = metadata
.strip_suffix(";base64")
.unwrap_or(metadata)
.to_owned();
(kind, Some(mime), decode_base64(data))
}
Some(url) => (kind, Some(url.to_owned()), None),
None => (kind, None, None),
}
}
fn decode_base64(input: &str) -> Option<Vec<u8>> {
fn digit(byte: u8) -> Option<u8> {
match byte {
b'A'..=b'Z' => Some(byte - b'A'),
b'a'..=b'z' => Some(byte - b'a' + 26),
b'0'..=b'9' => Some(byte - b'0' + 52),
b'+' => Some(62),
b'/' => Some(63),
_ => None,
}
}
let bytes = input
.bytes()
.filter(|byte| *byte != b'\n' && *byte != b'\r')
.collect::<Vec<_>>();
let mut output = Vec::with_capacity(bytes.len() / 4 * 3 + 3);
let mut chunk = [0_u8; 4];
let mut chunk_len = 0;
let mut padding = 0;
for byte in bytes {
if byte == b'=' {
padding += 1;
chunk[chunk_len] = 0;
} else {
chunk[chunk_len] = digit(byte)?;
}
chunk_len += 1;
if chunk_len == 4 {
let value = ((chunk[0] as u32) << 18)
| ((chunk[1] as u32) << 12)
| ((chunk[2] as u32) << 6)
| chunk[3] as u32;
output.push((value >> 16) as u8);
if padding < 2 {
output.push((value >> 8) as u8);
}
if padding < 1 {
output.push(value as u8);
}
chunk_len = 0;
padding = 0;
}
}
Some(output)
}
pub fn core_messages(messages: &[ChatMessage]) -> Vec<ChatMessage> {
messages
.iter()
.filter(|message| message.role != Role::System)
.cloned()
.collect()
}
pub fn replay_excluded(message: &ChatMessage) -> bool {
message.metadata.get("compacted_out").map(String::as_str) == Some("true")
|| message
.metadata
.get("pi_exclude_from_context")
.map(String::as_str)
== Some("true")
}
pub fn replay_eligible(messages: &[ChatMessage]) -> Vec<ChatMessage> {
messages
.iter()
.filter(|message| !replay_excluded(message))
.cloned()
.collect()
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct FidelityResidue {
pub compacted_out_excluded: usize,
pub other_dropped_messages: usize,
pub dropped_metadata_keys: BTreeSet<String>,
}
impl FidelityResidue {
pub fn is_semantically_lossless(&self) -> bool {
self.other_dropped_messages == 0 && self.dropped_metadata_keys.is_empty()
}
pub fn count(&self) -> usize {
self.compacted_out_excluded + self.other_dropped_messages + self.dropped_metadata_keys.len()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct FidelityMetric {
pub matched: usize,
pub total: usize,
pub residue: FidelityResidue,
}
impl FidelityMetric {
pub fn percent(&self) -> f64 {
if self.total == 0 {
100.0
} else {
self.matched as f64 / self.total as f64 * 100.0
}
}
pub fn pct(&self) -> f64 {
self.percent()
}
pub fn is_semantically_lossless(&self) -> bool {
self.matched + self.residue.compacted_out_excluded == self.total
&& self.residue.is_semantically_lossless()
}
}
pub fn measure_fidelity(source: &[ChatMessage], reloaded: &[ChatMessage]) -> FidelityMetric {
let mut residue = FidelityResidue::default();
let mut matched = 0;
let mut reload_index = 0;
for source_message in source {
let found = (reload_index..reloaded.len())
.find(|index| messages_equal_multimodal(source_message, &reloaded[*index]));
match found {
Some(index) => {
matched += 1;
for (key, value) in &source_message.metadata {
if reloaded[index].metadata.get(key) != Some(value) {
residue.dropped_metadata_keys.insert(key.clone());
}
}
reload_index = index + 1;
}
None if replay_excluded(source_message) => residue.compacted_out_excluded += 1,
None => residue.other_dropped_messages += 1,
}
}
FidelityMetric {
matched,
total: source.len(),
residue,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn semantic_losslessness_ignores_only_explicit_replay_exclusions() {
let mut excluded = ChatMessage::user("old");
excluded
.metadata
.insert("compacted_out".into(), "true".into());
let kept = ChatMessage::user("new");
let metric = measure_fidelity(&[excluded, kept.clone()], &[kept]);
assert_eq!(metric.percent(), 50.0);
assert!(metric.is_semantically_lossless());
}
#[test]
fn changed_metadata_values_are_semantic_residue() {
let mut source = ChatMessage::user("hello");
source.metadata.insert("model".into(), "alpha".into());
let mut reloaded = source.clone();
reloaded.metadata.insert("model".into(), "beta".into());
let metric = measure_fidelity(&[source], &[reloaded]);
assert_eq!(metric.matched, 1);
assert_eq!(
metric.residue.dropped_metadata_keys,
BTreeSet::from(["model".into()])
);
assert!(!metric.is_semantically_lossless());
}
}