use std::collections::HashMap;
use chrono::{DateTime, Local};
use either::Either;
use minijinja::{Error, ErrorKind, Value, value::Kwargs};
use serde::{Deserialize, Serialize};
#[allow(dead_code)]
#[derive(Debug, Deserialize)]
pub struct AddedTokensDecoder {
__type: Option<String>,
pub content: String,
lstrip: bool,
normalized: bool,
rstrip: bool,
single_word: bool,
special: Option<bool>,
}
pub fn raise_exception(msg: String) -> Result<String, minijinja::Error> {
Err(minijinja::Error::new(ErrorKind::InvalidOperation, msg))
}
#[derive(Debug, Deserialize)]
pub struct BeginEndUnkTok(
#[serde(with = "either::serde_untagged")] pub Either<String, AddedTokensDecoder>,
);
#[derive(Debug, Deserialize)]
pub struct ChatTemplateValue(
#[serde(with = "either::serde_untagged")] pub Either<String, Vec<HashMap<String, String>>>,
);
#[allow(dead_code)]
#[derive(Debug, Deserialize)]
pub struct PadTokenValue(
#[serde(with = "either::serde_untagged")] pub Either<String, AddedTokensDecoder>,
);
#[allow(dead_code)]
#[derive(Debug, Deserialize, Default)]
pub struct ChatTemplate {
pub bos_token: Option<BeginEndUnkTok>,
pub eos_token: Option<BeginEndUnkTok>,
pub unk_token: Option<BeginEndUnkTok>,
pub chat_template: Option<ChatTemplateValue>,
add_bos_token: Option<bool>,
add_eos_token: Option<bool>,
added_tokens_decoder: Option<HashMap<String, AddedTokensDecoder>>,
additional_special_tokens: Option<Vec<String>>,
clean_up_tokenization_spaces: Option<bool>,
device_map: Option<String>,
legacy: Option<bool>,
model_max_length: Option<f64>,
pad_token: Option<PadTokenValue>,
sp_model_kwargs: Option<HashMap<String, String>>,
spaces_between_special_tokens: Option<bool>,
tokenizer_class: Option<String>,
truncation_size: Option<String>,
use_default_system_prompt: Option<bool>,
}
impl ChatTemplate {
pub fn eos_tok(&self) -> Option<String> {
match self.eos_token.as_ref()?.0 {
Either::Left(ref lit) => Some(lit.clone()),
Either::Right(ref added) => Some(added.content.clone()),
}
}
pub fn bos_tok(&self) -> Option<String> {
match self.bos_token.as_ref()?.0 {
Either::Left(ref lit) => Some(lit.clone()),
Either::Right(ref added) => Some(added.content.clone()),
}
}
pub fn unk_tok(&self) -> Option<String> {
match self.unk_token.as_ref()?.0 {
Either::Left(ref lit) => Some(lit.clone()),
Either::Right(ref added) => Some(added.content.clone()),
}
}
}
#[allow(dead_code)]
#[derive(Debug, Deserialize)]
pub struct GenerationConfig {
#[serde(with = "either::serde_untagged")]
bos_token_id: Either<u32, Vec<u32>>,
#[serde(with = "either::serde_untagged")]
eos_token_id: Either<u32, Vec<u32>>,
}
struct PyJsonFormatter;
impl serde_json::ser::Formatter for PyJsonFormatter {
fn begin_array_value<W>(&mut self, writer: &mut W, first: bool) -> std::io::Result<()>
where
W: ?Sized + std::io::Write,
{
if !first {
writer.write_all(b", ")?;
}
Ok(())
}
fn begin_object_key<W>(&mut self, writer: &mut W, first: bool) -> std::io::Result<()>
where
W: ?Sized + std::io::Write,
{
if !first {
writer.write_all(b", ")?;
}
Ok(())
}
fn begin_object_value<W>(&mut self, writer: &mut W) -> std::io::Result<()>
where
W: ?Sized + std::io::Write,
{
writer.write_all(b": ")
}
}
pub fn tojson(value: Value, kwargs: Kwargs) -> Result<Value, Error> {
let mut buf = Vec::new();
let result = if let Ok(indent) = kwargs.get("indent") {
let repeat = b" ".repeat(indent);
let formatter = serde_json::ser::PrettyFormatter::with_indent(&repeat);
let mut serializer = serde_json::Serializer::with_formatter(&mut buf, formatter);
value.serialize(&mut serializer)
} else {
let mut serializer = serde_json::Serializer::with_formatter(&mut buf, PyJsonFormatter);
value.serialize(&mut serializer)
};
result.map_err(|err| {
Error::new(ErrorKind::BadSerialization, "cannot serialize to JSON").with_source(err)
})?;
String::from_utf8(buf)
.map_err(|err| {
Error::new(ErrorKind::BadSerialization, "cannot serialize to JSON").with_source(err)
})
.map(Value::from_safe_string)
}
pub fn fromjson(value: Value) -> Result<Value, Error> {
let Some(text) = value.as_str() else {
return Ok(value);
};
let parsed: serde_json::Value = serde_json::from_str(text).map_err(|err| {
Error::new(ErrorKind::InvalidOperation, "cannot parse JSON").with_source(err)
})?;
Ok(Value::from_serialize(&parsed))
}
pub fn strftime_now(format_str: &str) -> Result<Value, Error> {
let local: DateTime<Local> = Local::now();
Ok(Value::from_safe_string(
local.format(format_str).to_string(),
))
}
#[cfg(test)]
mod fromjson_tests {
use super::*;
#[test]
fn parses_json_object_string() {
let out = fromjson(Value::from(r#"{"location":"San Francisco","n":3}"#)).unwrap();
assert_eq!(
out.get_attr("location").unwrap().as_str(),
Some("San Francisco")
);
assert_eq!(out.get_attr("n").unwrap().to_string(), "3");
}
#[test]
fn parses_json_array_string() {
let out = fromjson(Value::from(r#"[1,2,3]"#)).unwrap();
assert_eq!(out.len(), Some(3));
}
#[test]
fn passes_through_non_string() {
let already = Value::from_serialize(serde_json::json!({"a": 1}));
let out = fromjson(already).unwrap();
assert_eq!(out.get_attr("a").unwrap().to_string(), "1");
}
#[test]
fn errors_on_malformed_json() {
assert!(fromjson(Value::from("{not json")).is_err());
}
#[test]
fn renders_tool_use_block_with_json_string_arguments() {
let mut env = minijinja::Environment::new();
env.add_filter("fromjson", fromjson);
env.add_template(
"tool_use",
"{% for tc in tool_calls %}{% set a = tc.function.arguments | fromjson %}\
CALL {{ tc.function.name }} loc={{ a.location }} unit={{ a.unit }}{% endfor %}",
)
.unwrap();
let rendered = env
.get_template("tool_use")
.unwrap()
.render(minijinja::context! { tool_calls => serde_json::json!([{
"function": {
"name": "get_weather",
"arguments": "{\"location\":\"San Francisco\",\"unit\":\"F\"}"
}
}])})
.expect("tool_use template must render once `fromjson` is registered");
assert_eq!(rendered, "CALL get_weather loc=San Francisco unit=F");
}
}