Skip to main content

dynamo_renderer/template/
tokcfg.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//based on: https://github.com/EricLBuehler/mistral.rs/blob/d970bb5feb863acf8e8ec90de97e18221fb959f1/mistralrs-core/src/pipeline/chat_template.rs
5
6use 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/// Support older tool use patterns where the tool use template was separate from the default/chat template.
37/// Modern patterns use a single template with a `tool_use` key, e.g.
38///
39/// ```jinja
40/// {%- if tools is not none and tool_choice is not none %}
41/// ```
42#[derive(Debug, Deserialize)]
43pub struct ChatTemplateValue(
44    #[serde(with = "either::serde_untagged")] pub Either<String, Vec<HashMap<String, String>>>,
45);
46
47/// If present, pad_token is usually a single value. Deepseek R1 and it's distill's use a map.
48#[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)]
56/// Template for chat models including bos/eos/unk as well as the chat template.
57pub struct ChatTemplate {
58    pub bos_token: Option<BeginEndUnkTok>,
59    pub eos_token: Option<BeginEndUnkTok>,
60    pub unk_token: Option<BeginEndUnkTok>,
61
62    /// Jinja format [chat templating] for chat completion.
63    ///
64    /// [chat templating]: https://huggingface.co/docs/transformers/chat_templating
65    pub chat_template: Option<ChatTemplateValue>,
66
67    // future
68    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
116/// Python `str` of a non-integer float (`1e-06`, `1e+16`, `nan`); minijinja never uses
117/// an exponent (`0.000001`, `10000000000000000.0`) and prints `NaN`. `None` for every
118/// other value.
119fn 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
126/// Output formatter for `{{ x }}`. HF renders through Python, which prints a float
127/// with `str`; see [`python_float_str`]. Other values are unchanged. `~` concatenation
128/// is a VM op that can't be overridden, and filters that stringify values themselves
129/// (`join`, `replace`, ...) use minijinja's spelling, so `'x' ~ 1e-6` and
130/// `[1e-6] | join` still render `0.000001`.
131pub 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
138/// The `string` filter with Python `str` floats, matching [`python_formatter`];
139/// other values go to the builtin.
140pub 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
147/// Mirrors HF transformers' `tojson` filter, not stock Jinja2's. Transformers
148/// overrides Jinja's HTML-safe `tojson` with plain
149/// `json.dumps(x, ensure_ascii=False)` in its chat-template environment, and
150/// vLLM/SGLang render through that — so chat templates (and the models trained
151/// on their output) expect Python separators and **no** HTML escaping
152/// (`'`, `<`, `>`, `&` stay literal). serde_json leaves non-ASCII unescaped by
153/// default, matching `ensure_ascii=False`.
154pub 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        // Python `json.dumps(indent=n)` separators are `(",", ": ")` with the
158        // item separator followed by newline + indent — PrettyFormatter matches.
159        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
177/// Parse a JSON string into a structured value.
178///
179/// HuggingFace/transformers chat-template environments expose this filter, and several
180/// published templates depend on it — e.g. Step-3.7-Flash's `tool_use` block applies it to
181/// a tool call's `arguments`, which arrive as a JSON *string*, to iterate the decoded
182/// object. Without it minijinja aborts the render with `unknown filter: fromjson`, which
183/// fails every multi-turn tool-call request.
184///
185/// Values that are not strings pass through untouched, so templates that apply the filter
186/// defensively to an already-decoded value keep rendering.
187pub 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        // Already-decoded values must survive a defensive `| fromjson`.
227        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    /// Regression: renders the shape of Step-3.7-Flash's `tool_use` block, where a tool
238    /// call's `arguments` arrive as a JSON string. Before the filter existed this failed
239    /// with `unknown filter: fromjson`, 500-ing every multi-turn tool-call request.
240    #[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}