dynamo_renderer/template/
tokcfg.rs1use std::collections::HashMap;
7
8use chrono::{DateTime, Local};
9use either::Either;
10use minijinja::{Error, ErrorKind, Output, State, Value, escape_formatter, value::Kwargs};
11use serde::{Deserialize, Serialize};
12
13use crate::python::{PyFloats, PyJsonFormatter, python_float_repr};
14
15#[allow(dead_code)]
16#[derive(Debug, Deserialize)]
17pub struct AddedTokensDecoder {
18 __type: Option<String>,
19 pub content: String,
20 lstrip: bool,
21 normalized: bool,
22 rstrip: bool,
23 single_word: bool,
24 special: Option<bool>,
25}
26
27pub fn raise_exception(msg: String) -> Result<String, minijinja::Error> {
28 Err(minijinja::Error::new(ErrorKind::InvalidOperation, msg))
29}
30
31#[derive(Debug, Deserialize)]
32pub struct BeginEndUnkTok(
33 #[serde(with = "either::serde_untagged")] pub Either<String, AddedTokensDecoder>,
34);
35
36#[derive(Debug, Deserialize)]
43pub struct ChatTemplateValue(
44 #[serde(with = "either::serde_untagged")] pub Either<String, Vec<HashMap<String, String>>>,
45);
46
47#[allow(dead_code)]
49#[derive(Debug, Deserialize)]
50pub struct PadTokenValue(
51 #[serde(with = "either::serde_untagged")] pub Either<String, AddedTokensDecoder>,
52);
53
54#[allow(dead_code)]
55#[derive(Debug, Deserialize, Default)]
56pub struct ChatTemplate {
58 pub bos_token: Option<BeginEndUnkTok>,
59 pub eos_token: Option<BeginEndUnkTok>,
60 pub unk_token: Option<BeginEndUnkTok>,
61
62 pub chat_template: Option<ChatTemplateValue>,
66
67 add_bos_token: Option<bool>,
69 add_eos_token: Option<bool>,
70 added_tokens_decoder: Option<HashMap<String, AddedTokensDecoder>>,
71 additional_special_tokens: Option<Vec<String>>,
72 clean_up_tokenization_spaces: Option<bool>,
73 device_map: Option<String>,
74 legacy: Option<bool>,
75 model_max_length: Option<f64>,
76 pad_token: Option<PadTokenValue>,
77 sp_model_kwargs: Option<HashMap<String, String>>,
78 spaces_between_special_tokens: Option<bool>,
79 tokenizer_class: Option<String>,
80 truncation_size: Option<String>,
81 use_default_system_prompt: Option<bool>,
82}
83
84impl ChatTemplate {
85 pub fn eos_tok(&self) -> Option<String> {
86 match self.eos_token.as_ref()?.0 {
87 Either::Left(ref lit) => Some(lit.clone()),
88 Either::Right(ref added) => Some(added.content.clone()),
89 }
90 }
91
92 pub fn bos_tok(&self) -> Option<String> {
93 match self.bos_token.as_ref()?.0 {
94 Either::Left(ref lit) => Some(lit.clone()),
95 Either::Right(ref added) => Some(added.content.clone()),
96 }
97 }
98
99 pub fn unk_tok(&self) -> Option<String> {
100 match self.unk_token.as_ref()?.0 {
101 Either::Left(ref lit) => Some(lit.clone()),
102 Either::Right(ref added) => Some(added.content.clone()),
103 }
104 }
105}
106
107#[allow(dead_code)]
108#[derive(Debug, Deserialize)]
109pub struct GenerationConfig {
110 #[serde(with = "either::serde_untagged")]
111 bos_token_id: Either<u32, Vec<u32>>,
112 #[serde(with = "either::serde_untagged")]
113 eos_token_id: Either<u32, Vec<u32>>,
114}
115
116fn python_float_str(value: &Value) -> Option<String> {
120 if !value.is_number() || value.is_integer() {
121 return None;
122 }
123 f64::try_from(value.clone()).ok().map(python_float_repr)
124}
125
126pub fn python_formatter(out: &mut Output, state: &State, value: &Value) -> Result<(), Error> {
132 match python_float_str(value) {
133 Some(text) => out.write_str(&text).map_err(Error::from),
134 None => escape_formatter(out, state, value),
135 }
136}
137
138pub fn python_string(state: &State, value: &Value) -> Result<Value, Error> {
141 match python_float_str(value) {
142 Some(text) => Ok(Value::from(text)),
143 None => minijinja::filters::string(state, value),
144 }
145}
146
147pub fn tojson(value: Value, kwargs: Kwargs) -> Result<Value, Error> {
155 let mut buf = Vec::new();
156 let result = if let Ok(indent) = kwargs.get("indent") {
157 let repeat = b" ".repeat(indent);
160 let formatter = PyFloats(serde_json::ser::PrettyFormatter::with_indent(&repeat));
161 let mut serializer = serde_json::Serializer::with_formatter(&mut buf, formatter);
162 value.serialize(&mut serializer)
163 } else {
164 let mut serializer = serde_json::Serializer::with_formatter(&mut buf, PyJsonFormatter);
165 value.serialize(&mut serializer)
166 };
167 result.map_err(|err| {
168 Error::new(ErrorKind::BadSerialization, "cannot serialize to JSON").with_source(err)
169 })?;
170 String::from_utf8(buf)
171 .map_err(|err| {
172 Error::new(ErrorKind::BadSerialization, "cannot serialize to JSON").with_source(err)
173 })
174 .map(Value::from_safe_string)
175}
176
177pub fn fromjson(value: Value) -> Result<Value, Error> {
188 let Some(text) = value.as_str() else {
189 return Ok(value);
190 };
191 let parsed: serde_json::Value = serde_json::from_str(text).map_err(|err| {
192 Error::new(ErrorKind::InvalidOperation, "cannot parse JSON").with_source(err)
193 })?;
194 Ok(Value::from_serialize(&parsed))
195}
196
197pub fn strftime_now(format_str: &str) -> Result<Value, Error> {
198 let local: DateTime<Local> = Local::now();
199 Ok(Value::from_safe_string(
200 local.format(format_str).to_string(),
201 ))
202}
203
204#[cfg(test)]
205mod fromjson_tests {
206 use super::*;
207
208 #[test]
209 fn parses_json_object_string() {
210 let out = fromjson(Value::from(r#"{"location":"San Francisco","n":3}"#)).unwrap();
211 assert_eq!(
212 out.get_attr("location").unwrap().as_str(),
213 Some("San Francisco")
214 );
215 assert_eq!(out.get_attr("n").unwrap().to_string(), "3");
216 }
217
218 #[test]
219 fn parses_json_array_string() {
220 let out = fromjson(Value::from(r#"[1,2,3]"#)).unwrap();
221 assert_eq!(out.len(), Some(3));
222 }
223
224 #[test]
225 fn passes_through_non_string() {
226 let already = Value::from_serialize(serde_json::json!({"a": 1}));
228 let out = fromjson(already).unwrap();
229 assert_eq!(out.get_attr("a").unwrap().to_string(), "1");
230 }
231
232 #[test]
233 fn errors_on_malformed_json() {
234 assert!(fromjson(Value::from("{not json")).is_err());
235 }
236
237 #[test]
241 fn renders_tool_use_block_with_json_string_arguments() {
242 let mut env = minijinja::Environment::new();
243 env.add_filter("fromjson", fromjson);
244 env.add_template(
245 "tool_use",
246 "{% for tc in tool_calls %}{% set a = tc.function.arguments | fromjson %}\
247CALL {{ tc.function.name }} loc={{ a.location }} unit={{ a.unit }}{% endfor %}",
248 )
249 .unwrap();
250 let rendered = env
251 .get_template("tool_use")
252 .unwrap()
253 .render(minijinja::context! { tool_calls => serde_json::json!([{
254 "function": {
255 "name": "get_weather",
256 "arguments": "{\"location\":\"San Francisco\",\"unit\":\"F\"}"
257 }
258 }])})
259 .expect("tool_use template must render once `fromjson` is registered");
260 assert_eq!(rendered, "CALL get_weather loc=San Francisco unit=F");
261 }
262}