use serde::de::{self, Deserializer, MapAccess, SeqAccess, Visitor};
use serde::{Deserialize, Serialize, Serializer};
use std::collections::HashMap;
use std::convert::Infallible;
use std::fmt;
use std::marker::PhantomData;
use std::str::FromStr;
#[allow(clippy::trivially_copy_pass_by_ref)]
pub(crate) fn is_false(value: &bool) -> bool {
!value
}
pub fn serialize_map_sorted<S, V>(
map: &HashMap<String, V>,
serializer: S,
) -> Result<S::Ok, S::Error>
where
S: Serializer,
V: Serialize,
{
let mut entries: Vec<_> = map.iter().collect();
entries.sort_by_key(|(key, _)| *key);
serializer.collect_map(entries)
}
pub fn serialize_optional_map_sorted<S, V>(
map: &Option<HashMap<String, V>>,
serializer: S,
) -> Result<S::Ok, S::Error>
where
S: Serializer,
V: Serialize,
{
match map {
Some(map) => serialize_map_sorted(map, serializer),
None => serializer.serialize_none(),
}
}
pub fn to_canonical_string(value: &serde_json::Value) -> String {
fn sorted(value: &serde_json::Value) -> serde_json::Value {
match value {
serde_json::Value::Object(map) => {
let mut entries: Vec<_> = map.iter().collect();
entries.sort_by_key(|(key, _)| *key);
let mut out = serde_json::Map::new();
for (key, value) in entries {
out.insert(key.clone(), sorted(value));
}
serde_json::Value::Object(out)
}
serde_json::Value::Array(items) => {
serde_json::Value::Array(items.iter().map(sorted).collect())
}
serde_json::Value::Null
| serde_json::Value::Bool(_)
| serde_json::Value::Number(_)
| serde_json::Value::String(_) => value.clone(),
}
}
sorted(value).to_string()
}
pub fn merge(a: serde_json::Value, b: serde_json::Value) -> serde_json::Value {
match (a, b) {
(serde_json::Value::Object(mut a_map), serde_json::Value::Object(b_map)) => {
b_map.into_iter().for_each(|(key, value)| {
a_map.insert(key, value);
});
serde_json::Value::Object(a_map)
}
(a, _) => a,
}
}
pub(crate) fn merge_params(
existing: Option<serde_json::Value>,
params: Option<serde_json::Value>,
) -> Option<serde_json::Value> {
Some(match (existing, params?) {
(Some(existing), params) => merge(existing, params),
(None, params) => params,
})
}
#[cfg_attr(not(any(feature = "image", feature = "audio")), allow(dead_code))]
pub fn merge_inplace(a: &mut serde_json::Value, b: serde_json::Value) {
if let (serde_json::Value::Object(a_map), serde_json::Value::Object(b_map)) = (a, b) {
b_map.into_iter().for_each(|(key, value)| {
a_map.insert(key, value);
});
}
}
pub fn value_to_json_string(value: &serde_json::Value) -> String {
match value {
serde_json::Value::String(s) => s.clone(),
other => other.to_string(),
}
}
pub fn serialize_json_value(value: &serde_json::Value) -> String {
value.to_string()
}
pub fn deserialize_json_string_or_value<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
where
D: Deserializer<'de>,
{
let value = Option::<serde_json::Value>::deserialize(deserializer)?;
Ok(match value {
None | Some(serde_json::Value::Null) => None,
Some(v) => Some(value_to_json_string(&v)),
})
}
pub fn parse_tool_arguments(arguments: &str) -> serde_json::Result<serde_json::Value> {
if arguments.trim().is_empty() {
return Ok(serde_json::Value::Object(serde_json::Map::new()));
}
serde_json::from_str(arguments)
}
pub fn parse_partial_object(text: &str) -> Option<serde_json::Map<String, serde_json::Value>> {
const MAX_ATTEMPTS: usize = 256;
let mut stack = Vec::new();
let mut in_string = false;
let mut escaped = false;
let mut cuts: Vec<(usize, Vec<u8>, bool)> = Vec::new();
let bytes = text.as_bytes();
for (at, byte) in bytes.iter().enumerate() {
if in_string {
match (escaped, *byte) {
(true, _) => escaped = false,
(false, b'\\') => escaped = true,
(false, b'"') => {
in_string = false;
cuts.push((at + 1, stack.clone(), false));
}
(false, _) => {}
}
continue;
}
match byte {
b'"' => in_string = true,
b'{' | b'[' => {
stack.push(*byte);
cuts.push((at + 1, stack.clone(), false));
}
b'}' | b']' => {
stack.pop();
cuts.push((at + 1, stack.clone(), false));
}
b if b.is_ascii_alphanumeric() || *b == b'.' || *b == b'-' || *b == b'+' => {
let next = bytes.get(at + 1).copied();
if !next.is_some_and(|n| {
n.is_ascii_alphanumeric() || n == b'.' || n == b'-' || n == b'+'
}) {
cuts.push((at + 1, stack.clone(), false));
}
}
_ => {}
}
}
if in_string {
let at = text.len() - usize::from(escaped);
cuts.push((at, stack.clone(), true));
}
cuts.iter()
.rev()
.take(MAX_ATTEMPTS)
.find_map(|(at, open, quote)| {
let mut candidate = text.get(..*at)?.to_owned();
if *quote {
candidate.push('"');
}
for container in open.iter().rev() {
candidate.push(if *container == b'{' { '}' } else { ']' });
}
match serde_json::from_str(&candidate) {
Ok(serde_json::Value::Object(object)) => Some(object),
_ => None,
}
})
}
pub fn parse_partial_arguments(text: &str) -> serde_json::Map<String, serde_json::Value> {
if text.trim().is_empty() {
return serde_json::Map::new();
}
let repaired = repair_json(text);
let complete = serde_json::from_str(text)
.ok()
.or_else(|| serde_json::from_str(&repaired).ok());
if let Some(value) = complete {
return match value {
serde_json::Value::Object(object) => object,
_ => serde_json::Map::new(),
};
}
let cut = without_partial_escape(&repaired);
parse_partial_object(cut)
.or_else(|| parse_partial_object(text))
.unwrap_or_default()
}
fn repair_json(text: &str) -> String {
let mut repaired = String::with_capacity(text.len());
let mut in_string = false;
let mut chars = text.chars().peekable();
while let Some(c) = chars.next() {
if !in_string {
in_string = c == '"';
repaired.push(c);
continue;
}
match c {
'"' => {
in_string = false;
repaired.push(c);
}
'\\' => match chars.peek() {
Some('"' | '\\' | '/' | 'b' | 'f' | 'n' | 'r' | 't' | 'u') => {
repaired.push(c);
if let Some(next) = chars.next() {
repaired.push(next);
}
}
None => repaired.push(c),
Some(_) => repaired.push_str("\\\\"),
},
'\n' => repaired.push_str("\\n"),
'\r' => repaired.push_str("\\r"),
'\t' => repaired.push_str("\\t"),
c if c.is_control() => {
use std::fmt::Write as _;
let _ = write!(repaired, "\\u{:04x}", u32::from(c));
}
c => repaired.push(c),
}
}
repaired
}
fn without_partial_escape(text: &str) -> &str {
let bytes = text.as_bytes();
let tail = bytes.len().saturating_sub(5);
let Some(at) = (tail..bytes.len())
.rev()
.find(|at| bytes.get(*at..*at + 2) == Some(b"\\u"))
else {
return text;
};
let digits = bytes.get(at + 2..).unwrap_or_default();
let escaping = bytes
.get(..at)
.unwrap_or_default()
.iter()
.rev()
.take_while(|byte| **byte == b'\\')
.count();
if digits.len() < 4 && digits.iter().all(u8::is_ascii_hexdigit) && escaping % 2 == 0 {
text.get(..at).unwrap_or(text)
} else {
text
}
}
pub mod stringified_json {
use super::parse_tool_arguments;
use serde::{self, Deserialize, Deserializer, Serializer};
pub fn serialize<S>(value: &serde_json::Value, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let s = value.to_string();
serializer.serialize_str(&s)
}
pub fn deserialize<'de, D>(deserializer: D) -> Result<serde_json::Value, D::Error>
where
D: Deserializer<'de>,
{
let s = String::deserialize(deserializer)?;
if s.trim().is_empty() {
return Ok(serde_json::Value::Object(serde_json::Map::new()));
}
serde_json::from_str(&s).map_err(serde::de::Error::custom)
}
pub fn deserialize_maybe_stringified<'de, D>(
deserializer: D,
) -> Result<serde_json::Value, D::Error>
where
D: Deserializer<'de>,
{
match serde_json::Value::deserialize(deserializer)? {
serde_json::Value::String(s) => {
parse_tool_arguments(&s).map_err(serde::de::Error::custom)
}
other => Ok(other),
}
}
}
pub fn string_or_vec<'de, T, D>(deserializer: D) -> Result<Vec<T>, D::Error>
where
T: Deserialize<'de> + FromStr<Err = Infallible>,
D: Deserializer<'de>,
{
struct StringOrVec<T>(PhantomData<fn() -> T>);
impl<'de, T> Visitor<'de> for StringOrVec<T>
where
T: Deserialize<'de> + FromStr<Err = Infallible>,
{
type Value = Vec<T>;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a string, sequence, or null")
}
fn visit_str<E>(self, value: &str) -> Result<Vec<T>, E>
where
E: de::Error,
{
let item = FromStr::from_str(value).map_err(de::Error::custom)?;
Ok(vec![item])
}
fn visit_seq<A>(self, seq: A) -> Result<Vec<T>, A::Error>
where
A: SeqAccess<'de>,
{
Deserialize::deserialize(de::value::SeqAccessDeserializer::new(seq))
}
fn visit_map<M>(self, map: M) -> Result<Vec<T>, M::Error>
where
M: MapAccess<'de>,
{
let item = Deserialize::deserialize(de::value::MapAccessDeserializer::new(map))?;
Ok(vec![item])
}
fn visit_none<E>(self) -> Result<Vec<T>, E>
where
E: de::Error,
{
Ok(vec![])
}
fn visit_unit<E>(self) -> Result<Vec<T>, E>
where
E: de::Error,
{
Ok(vec![])
}
}
deserializer.deserialize_any(StringOrVec(PhantomData))
}
pub fn null_or_default<'de, T, D>(deserializer: D) -> Result<T, D::Error>
where
T: Deserialize<'de> + Default,
D: Deserializer<'de>,
{
Ok(Option::<T>::deserialize(deserializer)?.unwrap_or_default())
}
#[cfg(test)]
mod tests;
pub trait Lenient {
fn str(&self, key: &str) -> Option<&str>;
fn u64(&self, key: &str) -> Option<u64>;
fn i64(&self, key: &str) -> Option<i64>;
fn f64(&self, key: &str) -> Option<f64>;
fn bool(&self, key: &str) -> Option<bool>;
fn obj(&self, key: &str) -> Option<&serde_json::Map<String, serde_json::Value>>;
fn arr(&self, key: &str) -> &[serde_json::Value];
fn at(&self, pointer: &str) -> Option<&serde_json::Value>;
fn as_u64_lenient(&self) -> Option<u64>;
}
impl Lenient for serde_json::Value {
fn str(&self, key: &str) -> Option<&str> {
self.get(key)?.as_str()
}
fn u64(&self, key: &str) -> Option<u64> {
self.get(key)?.as_u64_lenient()
}
fn i64(&self, key: &str) -> Option<i64> {
match self.get(key)? {
serde_json::Value::Number(number) => number.as_i64(),
serde_json::Value::String(text) => text.trim().parse().ok(),
_ => None,
}
}
fn f64(&self, key: &str) -> Option<f64> {
match self.get(key)? {
serde_json::Value::Number(number) => number.as_f64(),
serde_json::Value::String(text) => text.trim().parse().ok(),
_ => None,
}
}
fn bool(&self, key: &str) -> Option<bool> {
self.get(key)?.as_bool()
}
fn obj(&self, key: &str) -> Option<&serde_json::Map<String, serde_json::Value>> {
self.get(key)?.as_object()
}
fn arr(&self, key: &str) -> &[serde_json::Value] {
self.get(key)
.and_then(serde_json::Value::as_array)
.map_or(&[], Vec::as_slice)
}
fn at(&self, pointer: &str) -> Option<&serde_json::Value> {
self.pointer(pointer).filter(|value| !value.is_null())
}
fn as_u64_lenient(&self) -> Option<u64> {
match self {
serde_json::Value::Number(number) => number.as_u64().or_else(|| {
number
.as_f64()
.filter(|float| *float >= 0.0 && float.fract() == 0.0)
.map(|float| float as u64)
}),
serde_json::Value::String(text) => text.trim().parse().ok(),
_ => None,
}
}
}