1use super::*;
5
6use crate::{OAIChatLikeRequest, TextInput};
7use minijinja::{context, value::Value};
8use serde_json::json;
9use std::result::Result::Ok;
10
11pub fn may_be_fix_tool_schema(tools: serde_json::Value) -> Option<Value> {
15 let mut updated_tools = Vec::new();
19 if let Some(arr) = tools.as_array() {
20 for tool in arr {
21 let mut tool = tool.clone();
22 if let Some(function) = tool.get_mut("function") {
23 if let Some(obj) = function.as_object_mut()
28 && !matches!(obj.get("description"), Some(serde_json::Value::String(_)))
29 {
30 obj.insert(
31 "description".to_string(),
32 serde_json::Value::String(String::new()),
33 );
34 }
35 }
36 if let Some(function) = tool.get_mut("function")
37 && let Some(parameters) = function.get_mut("parameters")
38 {
39 if parameters.is_object() {
41 let mut needs_type = false;
42 let mut needs_properties = false;
43 let is_empty = parameters
44 .as_object()
45 .map(|o| o.is_empty())
46 .unwrap_or(false);
47
48 if is_empty {
50 needs_type = true;
51 needs_properties = true;
52 } else {
53 if let Some(obj) = parameters.as_object() {
55 if !obj.contains_key("type") {
56 needs_type = true;
57 }
58 if !obj.contains_key("properties") {
59 needs_properties = true;
60 }
61 }
62 }
63
64 if (needs_type || needs_properties)
65 && let Some(obj) = parameters.as_object_mut()
66 {
67 if needs_type {
68 obj.insert(
69 "type".to_string(),
70 serde_json::Value::String("object".to_string()),
71 );
72 }
73 if needs_properties {
74 obj.insert(
75 "properties".to_string(),
76 serde_json::Value::Object(Default::default()),
77 );
78 }
79 }
80 }
81 }
82 updated_tools.push(tool);
83 }
84 }
85 Some(Value::from_serialize(&updated_tools))
86}
87
88const DEFAULT_MEDIA_TYPE_CONVERSIONS: &[(&str, &str)] = &[
91 ("image_url", "image"),
92 ("video_url", "video"),
93 ("audio_url", "audio"),
94];
95
96fn convert_media_url_to_placeholder(
98 content_array: &mut [serde_json::Value],
99 conversions: &[(&str, &str)],
100) {
101 for part in content_array {
102 let part_type = part.get("type").and_then(|t| t.as_str()).unwrap_or("");
103 if let Some((_, target_type)) = conversions.iter().find(|(src, _)| *src == part_type) {
104 *part = serde_json::json!({"type": target_type});
105 }
106 }
107}
108
109fn may_be_fix_msg_content(
110 mut messages: serde_json::Value,
111 preserve_arrays: bool,
112 image_placeholder_template: Option<&str>,
113) -> serde_json::Value {
114 let Some(arr) = messages.as_array_mut() else {
115 return messages;
116 };
117 for msg in arr {
118 let Some(content) = msg.get_mut("content") else {
119 continue;
120 };
121 match content {
122 serde_json::Value::String(_) if preserve_arrays => {
123 let text = content.take();
124 *content = serde_json::Value::Array(vec![serde_json::Value::Object(
125 serde_json::Map::from_iter([
126 ("type".into(), serde_json::Value::String("text".into())),
127 ("text".into(), text),
128 ]),
129 )]);
130 }
131 serde_json::Value::Array(parts) => {
132 convert_media_url_to_placeholder(parts, DEFAULT_MEDIA_TYPE_CONVERSIONS);
133 let text_only = !parts.is_empty()
136 && parts
137 .iter()
138 .all(|part| part.get("type").and_then(|v| v.as_str()) == Some("text"));
139 if text_only && !preserve_arrays {
140 let text = parts
141 .iter()
142 .filter_map(|part| part.get("text")?.as_str())
143 .collect::<Vec<_>>()
144 .join("\n");
145 *content = serde_json::Value::String(text);
146 } else if !preserve_arrays
147 && !parts.is_empty()
148 && let Some(template) = image_placeholder_template
149 {
150 *content = serde_json::Value::String(flatten_mixed_content(parts, template));
151 }
152 }
153 _ => {}
154 }
155 }
156 messages
157}
158
159fn flatten_mixed_content(parts: &[serde_json::Value], placeholder_tpl: &str) -> String {
182 let mut out = String::new();
183 let mut img_idx: u32 = 1;
184 for part in parts {
185 let type_str = part.get("type").and_then(|t| t.as_str()).unwrap_or("");
186 if type_str == "text" {
187 if let Some(text) = part.get("text").and_then(|t| t.as_str()) {
188 out.push_str(text);
189 }
190 } else if !type_str.is_empty() {
191 let placeholder = placeholder_tpl.replace("{n}", &img_idx.to_string());
192 out.push_str(&placeholder);
193 img_idx += 1;
194 }
195 }
196 out
197}
198
199fn normalize_tool_calls_arguments_in_messages(messages: &mut serde_json::Value) {
200 let Some(msgs) = messages.as_array_mut() else {
205 return;
206 };
207
208 for msg in msgs.iter_mut() {
209 if let Some(tool_calls) = msg.get_mut("tool_calls").and_then(|v| v.as_array_mut()) {
210 for tc in tool_calls {
211 if let Some(function) = tc.get_mut("function").and_then(|v| v.as_object_mut())
212 && let Some(args) = function.get_mut("arguments")
213 && let Some(s) = args.as_str()
214 && let Ok(parsed) = serde_json::from_str(s)
215 {
216 *args = parsed;
217 }
218 }
219 }
220 }
221}
222
223fn normalize_function_call_arguments_in_messages(messages: &mut serde_json::Value) {
224 let Some(msgs) = messages.as_array_mut() else {
229 return;
230 };
231
232 for msg in msgs.iter_mut() {
233 if let Some(function_call) = msg.get_mut("function_call").and_then(|v| v.as_object_mut())
234 && let Some(args) = function_call.get_mut("arguments")
235 && let Some(s) = args.as_str()
236 && let Ok(parsed) = serde_json::from_str(s)
237 {
238 *args = parsed;
239 }
240 }
241}
242
243fn inject_reasoning_content_into_messages(messages: &mut serde_json::Value) {
257 let Some(msgs) = messages.as_array_mut() else {
258 return;
259 };
260
261 for msg in msgs.iter_mut() {
262 if msg.get("role").and_then(|r| r.as_str()) != Some("assistant") {
263 continue;
264 }
265
266 let reasoning = match msg.get("reasoning_content") {
267 Some(serde_json::Value::String(s)) if !s.is_empty() => {
268 format!("<think>{}</think>", s)
269 }
270 Some(serde_json::Value::Array(segments)) => {
271 let mut result = String::new();
272 for seg in segments {
273 if let Some(s) = seg.as_str()
274 && !s.is_empty()
275 {
276 result.push_str("<think>");
277 result.push_str(s);
278 result.push_str("</think>");
279 }
280 }
281 if result.is_empty() {
282 continue;
283 }
284 result
285 }
286 _ => continue,
287 };
288
289 match msg.get("content") {
290 Some(serde_json::Value::String(s)) if !s.is_empty() => {
292 msg["content"] = serde_json::Value::String(format!("{}{}", reasoning, s));
293 }
294 None | Some(serde_json::Value::Null) | Some(serde_json::Value::String(_)) => {
295 msg["content"] = serde_json::Value::String(reasoning);
296 }
297 Some(serde_json::Value::Array(_)) => {
299 let think_part = serde_json::json!({
300 "type": "text",
301 "text": reasoning
302 });
303 if let Some(arr) = msg.get_mut("content").and_then(|v| v.as_array_mut()) {
304 arr.insert(0, think_part);
305 }
306 }
307 _ => continue,
309 }
310
311 if let Some(obj) = msg.as_object_mut() {
314 obj.remove("reasoning_content");
315 }
316 }
317}
318
319impl OAIChatLikeRequest for dynamo_protocols::types::CreateChatCompletionRequest {
325 fn model(&self) -> String {
326 self.model.clone()
327 }
328
329 fn messages(&self) -> Value {
330 let messages_json = serde_json::to_value(&self.messages).unwrap();
331 Value::from_serialize(&messages_json)
332 }
333
334 fn typed_messages(&self) -> Option<&[dynamo_protocols::types::ChatCompletionRequestMessage]> {
335 Some(self.messages.as_slice())
336 }
337
338 fn tools(&self) -> Option<Value> {
339 if self.tools.is_none() {
340 None
341 } else {
342 Some(may_be_fix_tool_schema(
343 serde_json::to_value(&self.tools).unwrap(),
344 )?)
345 }
346 }
347
348 fn tool_choice(&self) -> Option<Value> {
349 if self.tool_choice.is_none() {
350 None
351 } else {
352 Some(Value::from_serialize(&self.tool_choice))
353 }
354 }
355
356 fn response_format(&self) -> Option<Value> {
357 self.response_format.as_ref().map(Value::from_serialize)
358 }
359
360 fn reasoning_effort(&self) -> Option<Value> {
361 self.reasoning_effort.as_ref().map(Value::from_serialize)
362 }
363
364 fn should_add_generation_prompt(&self) -> bool {
365 true
367 }
368
369 fn extract_text(&self) -> Option<TextInput> {
370 Some(TextInput::Single(String::new()))
371 }
372
373 fn mm_processor_kwargs(&self) -> Option<&serde_json::Value> {
374 self.mm_processor_kwargs.as_ref()
375 }
376}
377
378fn merge_message_content(
382 target: serde_json::Value,
383 source: serde_json::Value,
384) -> serde_json::Value {
385 use serde_json::Value;
386 let text_part = |text: String| json!({"type": "text", "text": text});
387 match (target, source) {
388 (Value::String(mut target), Value::String(source)) => {
389 if !target.is_empty() && !source.is_empty() {
390 target.push_str("\n\n");
391 }
392 target.push_str(&source);
393 Value::String(target)
394 }
395 (Value::Array(mut target), Value::Array(source)) => {
396 target.extend(source);
397 Value::Array(target)
398 }
399 (Value::Array(mut target), Value::String(source)) => {
400 if !source.is_empty() {
401 target.push(text_part(source));
402 }
403 Value::Array(target)
404 }
405 (Value::String(target), Value::Array(source)) => {
406 let mut parts = Vec::with_capacity(source.len() + 1);
407 if !target.is_empty() {
408 parts.push(text_part(target));
409 }
410 parts.extend(source);
411 Value::Array(parts)
412 }
413 (Value::Null, source) => source,
414 (target, _) => target,
416 }
417}
418
419fn append_message_content(target: &mut serde_json::Value, source: serde_json::Value) {
421 let Some(target) = target.as_object_mut() else {
422 return;
423 };
424 let merged = merge_message_content(
425 target.remove("content").unwrap_or(serde_json::Value::Null),
426 source,
427 );
428 target.insert("content".to_string(), merged);
429}
430
431fn take_message_content(message: &mut serde_json::Value) -> serde_json::Value {
432 message
433 .get_mut("content")
434 .map(serde_json::Value::take)
435 .unwrap_or(serde_json::Value::Null)
436}
437
438fn normalize_system_messages(messages: &mut serde_json::Value, rules: SystemNormalization) {
442 let serde_json::Value::Array(list) = messages else {
443 return;
444 };
445 let role_is =
446 |m: &serde_json::Value, r: &str| m.get("role").and_then(|v| v.as_str()) == Some(r);
447
448 if rules.demote_nonleading_system {
449 let leading = list.iter().take_while(|m| role_is(m, "system")).count();
452 if leading > 1 {
453 for mut trailing in list.drain(1..leading).collect::<Vec<_>>() {
454 let content = take_message_content(&mut trailing);
455 append_message_content(&mut list[0], content);
456 }
457 }
458
459 let leading = list.iter().take_while(|m| role_is(m, "system")).count();
462 for m in list.iter_mut().skip(leading) {
463 if role_is(m, "system")
464 && let Some(m) = m.as_object_mut()
465 {
466 m.insert("role".to_string(), json!("user"));
467 }
468 }
469 }
470
471 if rules.coalesce_consecutive_users {
472 let mut coalesced: Vec<serde_json::Value> = Vec::with_capacity(list.len());
473 for mut m in list.drain(..) {
474 if role_is(&m, "user") && coalesced.last().is_some_and(|p| role_is(p, "user")) {
475 let content = take_message_content(&mut m);
476 append_message_content(coalesced.last_mut().unwrap(), content);
477 } else {
478 coalesced.push(m);
479 }
480 }
481 *list = coalesced;
482 }
483}
484
485impl OAIPromptFormatter for HfTokenizerConfigJsonFormatter {
486 fn supports_add_generation_prompt(&self) -> bool {
487 self.supports_add_generation_prompt
488 }
489
490 fn render(&self, req: &dyn OAIChatLikeRequest) -> Result<String> {
491 let mixins = Value::from_dyn_object(self.mixins.clone());
492
493 let tools = req.tools();
494 let tools = if self.exclude_tools_when_tool_choice_none {
497 match req.tool_choice() {
498 Some(ref tc) if tc.as_str() == Some("none") => None,
499 _ => tools,
500 }
501 } else {
502 tools
503 };
504 let has_tools = tools.as_ref().and_then(|v| v.len()).is_some_and(|l| l > 0);
506 let add_generation_prompt = req.should_add_generation_prompt();
507
508 tracing::trace!(
509 "Rendering prompt with tools: {:?}, add_generation_prompt: {}",
510 has_tools,
511 add_generation_prompt
512 );
513
514 let (
517 template_name,
518 template_handles_tool_calls_args_string,
519 template_handles_reasoning,
520 system_normalization,
521 ) = if has_tools {
522 (
523 "tool_use",
524 self.tool_use_template_handles_tool_calls_arguments_string,
525 self.tool_use_template_handles_reasoning,
526 self.tool_use_system_normalization,
527 )
528 } else {
529 (
530 "default",
531 self.default_template_handles_tool_calls_arguments_string,
532 self.default_template_handles_reasoning,
533 self.default_system_normalization,
534 )
535 };
536
537 let mut messages_for_template = crate::messages_to_json(req)?;
538
539 crate::reject_unsupported_partial_assistant(&messages_for_template)?;
540 crate::reject_unsupported_message_tools(&messages_for_template, &[])?;
541
542 if system_normalization.is_required() {
543 normalize_system_messages(&mut messages_for_template, system_normalization);
544 }
545
546 messages_for_template = may_be_fix_msg_content(
547 messages_for_template,
548 self.requires_content_arrays,
549 self.image_placeholder_template,
550 );
551
552 if !template_handles_tool_calls_args_string {
559 normalize_tool_calls_arguments_in_messages(&mut messages_for_template);
560 }
561 normalize_function_call_arguments_in_messages(&mut messages_for_template);
565
566 if !template_handles_reasoning {
571 inject_reasoning_content_into_messages(&mut messages_for_template);
572 }
573
574 let ctx = context! {
575 messages => messages_for_template,
576 tools => tools,
577 bos_token => self.config.bos_tok(),
578 eos_token => self.config.eos_tok(),
579 unk_token => self.config.unk_tok(),
580 add_generation_prompt => add_generation_prompt,
581 ..mixins
582 };
583
584 let ctx = if let Some(args) = req.chat_template_args() {
586 let extra = Value::from_serialize(args);
587 context! { ..ctx, ..extra }
588 } else {
589 ctx
590 };
591
592 let tmpl: minijinja::Template<'_, '_> = self.env.get_template(template_name)?;
593 Ok(tmpl.render(&ctx)?)
594 }
595}
596
597#[cfg(test)]
598mod tests {
599 use super::*;
600
601 use dynamo_protocols::types::ChatCompletionRequestMessage as Msg;
602 use dynamo_protocols::types::CreateChatCompletionRequest as NvCreateChatCompletionRequest;
605 use minijinja::{Environment, context};
606
607 use super::super::tokcfg::ChatTemplate as SysChatTemplate;
610 use super::super::{
611 ContextMixins as SysMixins, HfTokenizerConfigJsonFormatter as SysFormatter,
612 };
613
614 fn formatter_for(template: &str) -> SysFormatter {
615 let ct: SysChatTemplate = serde_json::from_value(json!({
617 "chat_template": template,
618 "bos_token": "<s>",
619 "eos_token": "</s>",
620 "unk_token": "<unk>",
621 }))
622 .unwrap();
623 SysFormatter::new(ct, SysMixins::new(&[])).unwrap()
624 }
625
626 fn formatter_for_templates(default: &str, tool_use: &str) -> SysFormatter {
627 let ct: SysChatTemplate = serde_json::from_value(json!({
628 "chat_template": [
629 {"default": default},
630 {"tool_use": tool_use},
631 ],
632 "bos_token": "<s>",
633 "eos_token": "</s>",
634 "unk_token": "<unk>",
635 }))
636 .unwrap();
637 SysFormatter::new(ct, SysMixins::new(&[])).unwrap()
638 }
639
640 fn try_formatter_for(template: &str) -> Option<SysFormatter> {
641 let ct: SysChatTemplate = serde_json::from_value(json!({
642 "chat_template": template,
643 "bos_token": "<s>",
644 "eos_token": "</s>",
645 "unk_token": "<unk>",
646 }))
647 .ok()?;
648 SysFormatter::new(ct, SysMixins::new(&[])).ok()
649 }
650
651 fn render_shape(f: &SysFormatter, messages: serde_json::Value) -> Result<String> {
652 let req: NvCreateChatCompletionRequest =
653 serde_json::from_value(json!({ "model": "test", "messages": messages })).unwrap();
654 f.render(&req)
655 }
656
657 fn render_shape_with_tools(f: &SysFormatter, messages: serde_json::Value) -> Result<String> {
658 let req: NvCreateChatCompletionRequest = serde_json::from_value(json!({
659 "model": "test",
660 "messages": messages,
661 "tools": [{
662 "type": "function",
663 "function": {"name": "noop", "parameters": {}}
664 }]
665 }))
666 .unwrap();
667 f.render(&req)
668 }
669
670 struct RawMessagesRequest(Value);
671
672 impl OAIChatLikeRequest for RawMessagesRequest {
673 fn model(&self) -> String {
674 "test".to_string()
675 }
676
677 fn messages(&self) -> Value {
678 self.0.clone()
679 }
680
681 fn should_add_generation_prompt(&self) -> bool {
682 true
683 }
684 }
685
686 fn render_raw_shape(f: &SysFormatter, messages: serde_json::Value) -> Result<String> {
687 f.render(&RawMessagesRequest(Value::from_serialize(&messages)))
688 }
689
690 #[test]
691 fn content_normalization_preserves_unchanged_values() {
692 let messages = json!([
693 {"role": "user", "content": " 中文 <special>\n", "name": "user"},
694 {"role": "assistant", "content": null},
695 {"role": "user", "content": []},
696 {"role": "assistant", "tool_calls": []}
697 ]);
698 assert_eq!(
699 may_be_fix_msg_content(messages.clone(), false, Some("")),
700 messages
701 );
702 }
703
704 const PERMISSIVE_TMPL: &str = concat!(
705 "{%- for m in messages -%}",
706 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
707 "{%- endfor -%}"
708 );
709
710 #[test]
711 fn jinja_templates_reject_message_level_tools() {
712 let f = formatter_for(PERMISSIVE_TMPL);
713 let error = render_shape(
714 &f,
715 json!([
716 {"role": "system", "tools": [{"name": "lookup", "parameters": {"type": "object"}}]},
717 {"role": "user", "content": "hi"}
718 ]),
719 )
720 .unwrap_err();
721 assert!(matches!(
722 error.downcast_ref::<crate::PromptRenderError>(),
723 Some(crate::PromptRenderError::InvalidRequest(message))
724 if message.contains("message-level `tools`")
725 ));
726
727 let error = render_raw_shape(
728 &f,
729 json!([{
730 "role": "user",
731 "content": "hi",
732 "tools": [{"name": "lookup", "parameters": {"type": "object"}}]
733 }]),
734 )
735 .unwrap_err();
736 assert!(matches!(
737 error.downcast_ref::<crate::PromptRenderError>(),
738 Some(crate::PromptRenderError::InvalidRequest(message))
739 if message.contains("message-level `tools`")
740 ));
741
742 let rendered = render_shape(
743 &f,
744 json!([
745 {"role": "system", "content": "You are helpful.", "tools": []},
746 {"role": "user", "content": "hi"}
747 ]),
748 )
749 .unwrap();
750 assert!(rendered.contains("<|im_start|>system\nYou are helpful.<|im_end|>"));
751 }
752
753 #[test]
754 fn jinja_templates_reject_unsupported_partial_assistant() {
755 let f = formatter_for(PERMISSIVE_TMPL);
756 let error = render_shape(
757 &f,
758 json!([
759 {"role": "user", "content": "Continue"},
760 {"role": "assistant", "content": "prefix", "partial": true}
761 ]),
762 )
763 .unwrap_err();
764 assert!(matches!(
765 error.downcast_ref::<crate::PromptRenderError>(),
766 Some(crate::PromptRenderError::InvalidRequest(message))
767 if message.contains("`partial: true` is not supported")
768 ));
769
770 let rendered = render_shape(
771 &f,
772 json!([{"role": "assistant", "content": "ordinary", "partial": false}]),
773 )
774 .unwrap();
775 assert!(rendered.contains("ordinary"));
776 }
777 const STRICT_LEADING_TMPL: &str = concat!(
779 "{%- for m in messages -%}",
780 "{%- if m.role == 'system' and not loop.first -%}",
781 "{{ raise_exception('System message must be at the beginning.') }}",
782 "{%- endif -%}",
783 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
784 "{%- endfor -%}"
785 );
786 const ALTERNATION_TMPL: &str = concat!(
788 "{%- set ns = namespace(prev='') -%}",
789 "{%- for m in messages -%}",
790 "{%- if m.role == 'user' and ns.prev == 'user' -%}",
791 "{{ raise_exception('Conversation roles must alternate.') }}",
792 "{%- endif -%}",
793 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
794 "{%- set ns.prev = m.role -%}",
795 "{%- endfor -%}"
796 );
797 const STRICT_BOTH_TMPL: &str = concat!(
799 "{%- set ns = namespace(prev='') -%}",
800 "{%- for m in messages -%}",
801 "{%- if m.role == 'system' and not loop.first -%}",
802 "{{ raise_exception('System message must be at the beginning.') }}",
803 "{%- endif -%}",
804 "{%- if m.role == 'user' and ns.prev == 'user' -%}",
805 "{{ raise_exception('Conversation roles must alternate.') }}",
806 "{%- endif -%}",
807 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
808 "{%- set ns.prev = m.role -%}",
809 "{%- endfor -%}"
810 );
811 const DEFAULT_NONE_GATED_TMPL: &str = concat!(
814 "{%- set strict = tools is not none -%}",
815 "{%- for m in messages -%}",
816 "{%- if strict and m.role == 'system' and not loop.first -%}",
817 "{{ raise_exception('System message must be at the beginning.') }}",
818 "{%- endif -%}",
819 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
820 "{%- endfor -%}"
821 );
822 const TOOL_NONEMPTY_GATED_TMPL: &str = concat!(
823 "{%- set strict = tools|length > 0 -%}",
824 "{%- for m in messages -%}",
825 "{%- if strict and m.role == 'system' and not loop.first -%}",
826 "{{ raise_exception('System message must be at the beginning.') }}",
827 "{%- endif -%}",
828 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
829 "{%- endfor -%}"
830 );
831 const STRICT_ARRAY_TMPL: &str = concat!(
833 "{%- for m in messages -%}",
834 "{%- if m.role == 'system' and not loop.first -%}",
835 "{{ raise_exception('System message must be at the beginning.') }}",
836 "{%- endif -%}",
837 "<|im_start|>{{ m.role }}\n",
838 "{%- if m.content is not string -%}",
839 "{%- for part in m.content -%}{{ part.text }}{%- endfor -%}",
840 "{%- endif -%}",
841 "<|im_end|>\n",
842 "{%- endfor -%}"
843 );
844
845 fn claude_shape() -> serde_json::Value {
847 json!([
848 {"role": "system", "content": "You are Claude Code."},
849 {"role": "user", "content": "hello"},
850 {"role": "system", "content": "mid-conversation reminder"},
851 ])
852 }
853
854 fn all_restrictions() -> SystemNormalization {
855 SystemNormalization {
856 demote_nonleading_system: true,
857 coalesce_consecutive_users: true,
858 }
859 }
860
861 #[test]
862 fn permissive_template_is_not_flagged_and_renders_untouched() {
863 let f = formatter_for(PERMISSIVE_TMPL);
864 assert!(!f.default_system_normalization.is_required());
865 assert!(!f.tool_use_system_normalization.is_required());
866 let out = render_shape(&f, claude_shape()).unwrap();
867 assert!(out.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
868 }
869
870 #[test]
871 fn strict_leading_template_demotes_mid_system_but_keeps_user_turns_apart() {
872 let f = formatter_for(STRICT_LEADING_TMPL);
873 assert!(f.default_system_normalization.demote_nonleading_system);
874 assert!(!f.default_system_normalization.coalesce_consecutive_users);
876
877 let out = render_shape(&f, claude_shape()).unwrap();
879 assert_eq!(out.matches("<|im_start|>system").count(), 1);
880 assert!(out.contains("<|im_start|>user\nhello<|im_end|>"));
881 assert!(out.contains("<|im_start|>user\nmid-conversation reminder<|im_end|>"));
882 }
883
884 #[test]
885 fn alternation_template_coalesces_users_but_keeps_mid_system() {
886 let f = formatter_for(ALTERNATION_TMPL);
887 assert!(f.default_system_normalization.coalesce_consecutive_users);
888 assert!(!f.default_system_normalization.demote_nonleading_system);
890
891 let out = render_shape(&f, claude_shape()).unwrap();
892 assert!(out.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
893
894 let out = render_shape(
895 &f,
896 json!([
897 {"role": "system", "content": "s"},
898 {"role": "user", "content": "hello"},
899 {"role": "user", "content": "again"},
900 ]),
901 )
902 .unwrap();
903 assert_eq!(out.matches("<|im_start|>user").count(), 1);
904 assert!(out.contains("<|im_start|>user\nhello\n\nagain<|im_end|>"));
905 }
906
907 #[test]
908 fn strict_both_template_demotes_then_coalesces() {
909 let f = formatter_for(STRICT_BOTH_TMPL);
910 assert!(f.default_system_normalization.demote_nonleading_system);
911 assert!(f.default_system_normalization.coalesce_consecutive_users);
912
913 let out = render_shape(&f, claude_shape()).unwrap();
914 assert_eq!(out.matches("<|im_start|>system").count(), 1);
915 assert!(out.contains("<|im_start|>user\nhello\n\nmid-conversation reminder<|im_end|>"));
916 }
917
918 #[test]
919 fn system_normalization_flag_is_selected_per_template() {
920 let f = formatter_for_templates(PERMISSIVE_TMPL, STRICT_LEADING_TMPL);
921 assert!(!f.default_system_normalization.is_required());
922 assert!(f.tool_use_system_normalization.is_required());
923
924 let no_tools = render_shape(&f, claude_shape()).unwrap();
925 assert!(no_tools.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
926 let with_tools = render_shape_with_tools(&f, claude_shape()).unwrap();
927 assert_eq!(with_tools.matches("<|im_start|>system").count(), 1);
928 assert!(with_tools.contains("<|im_start|>user\nmid-conversation reminder<|im_end|>"));
929
930 let f = formatter_for_templates(STRICT_LEADING_TMPL, PERMISSIVE_TMPL);
931 assert!(f.default_system_normalization.is_required());
932 assert!(!f.tool_use_system_normalization.is_required());
933 let with_tools = render_shape_with_tools(&f, claude_shape()).unwrap();
934 assert!(with_tools.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
935 }
936
937 #[test]
938 fn system_normalization_probe_uses_runtime_tools_shape() {
939 let f = formatter_for_templates(DEFAULT_NONE_GATED_TMPL, TOOL_NONEMPTY_GATED_TMPL);
940 assert!(!f.default_system_normalization.is_required());
941 assert!(f.tool_use_system_normalization.is_required());
942
943 let no_tools = render_shape(&f, claude_shape()).unwrap();
944 assert!(no_tools.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
945
946 let with_tools = render_shape_with_tools(&f, claude_shape()).unwrap();
947 assert_eq!(with_tools.matches("<|im_start|>system").count(), 1);
948 assert!(with_tools.contains("<|im_start|>user\nmid-conversation reminder<|im_end|>"));
949 }
950
951 #[test]
952 fn system_normalization_precedes_required_content_array_conversion() {
953 let f = formatter_for(STRICT_ARRAY_TMPL);
954 assert!(f.requires_content_arrays);
955 assert!(f.default_system_normalization.demote_nonleading_system);
956
957 let out = render_shape(
958 &f,
959 json!([
960 {"role": "system", "content": "A"},
961 {"role": "system", "content": "B"},
962 {"role": "user", "content": "hello"},
963 ]),
964 )
965 .unwrap();
966 assert!(out.contains("A\n\nB"));
967 }
968
969 #[test]
970 fn normalize_preserves_multimodal_user_content_and_fields() {
971 let mut m = json!([
972 {
973 "role": "user",
974 "name": "kept",
975 "content": [
976 {"type": "text", "text": "look"},
977 {"type": "image"},
978 ],
979 },
980 {"role": "system", "content": "remember"},
981 ]);
982 normalize_system_messages(&mut m, all_restrictions());
983 assert_eq!(
984 m,
985 json!([{
986 "role": "user",
987 "name": "kept",
988 "content": [
989 {"type": "text", "text": "look"},
990 {"type": "image"},
991 {"type": "text", "text": "remember"},
992 ],
993 }])
994 );
995 }
996
997 #[test]
1000 fn coalesce_preserves_multimodal_content_of_the_merged_turn() {
1001 let mut m = json!([
1002 {"role": "user", "content": "look"},
1003 {"role": "user", "content": [
1004 {"type": "text", "text": "at this"},
1005 {"type": "image_url", "image_url": {"url": "http://img"}},
1006 ]},
1007 ]);
1008 normalize_system_messages(&mut m, all_restrictions());
1009 assert_eq!(
1010 m,
1011 json!([{
1012 "role": "user",
1013 "content": [
1014 {"type": "text", "text": "look"},
1015 {"type": "text", "text": "at this"},
1016 {"type": "image_url", "image_url": {"url": "http://img"}},
1017 ],
1018 }])
1019 );
1020 }
1021
1022 #[test]
1023 fn normalize_merges_leading_run_and_coalesces() {
1024 let mut m = json!([
1025 {"role": "system", "content": "A"},
1026 {"role": "system", "content": "B"},
1027 {"role": "user", "content": "hi"},
1028 {"role": "system", "content": "reminder"},
1029 ]);
1030 normalize_system_messages(&mut m, all_restrictions());
1031 assert_eq!(
1032 m,
1033 json!([
1034 {"role": "system", "content": "A\n\nB"},
1035 {"role": "user", "content": "hi\n\nreminder"},
1036 ])
1037 );
1038 }
1039
1040 #[test]
1043 fn each_restriction_applies_only_its_own_rewrite() {
1044 let shape = json!([
1045 {"role": "system", "content": "A"},
1046 {"role": "system", "content": "B"},
1047 {"role": "user", "content": "hi"},
1048 {"role": "system", "content": "reminder"},
1049 ]);
1050
1051 let mut demote_only = shape.clone();
1052 normalize_system_messages(
1053 &mut demote_only,
1054 SystemNormalization {
1055 demote_nonleading_system: true,
1056 coalesce_consecutive_users: false,
1057 },
1058 );
1059 assert_eq!(
1060 demote_only,
1061 json!([
1062 {"role": "system", "content": "A\n\nB"},
1063 {"role": "user", "content": "hi"},
1064 {"role": "user", "content": "reminder"},
1065 ])
1066 );
1067
1068 let mut coalesce_only = shape.clone();
1069 normalize_system_messages(
1070 &mut coalesce_only,
1071 SystemNormalization {
1072 demote_nonleading_system: false,
1073 coalesce_consecutive_users: true,
1074 },
1075 );
1076 assert_eq!(coalesce_only, shape);
1077 }
1078
1079 #[test]
1080 fn normalize_preserves_array_system_content() {
1081 let mut m = json!([
1082 {"role": "user", "content": "hi"},
1083 {"role": "system", "content": [{"type": "text", "text": "one"},
1084 {"type": "text", "text": "two"}]},
1085 ]);
1086 normalize_system_messages(&mut m, all_restrictions());
1087 assert_eq!(
1088 m,
1089 json!([{"role": "user", "content": [
1090 {"type": "text", "text": "hi"},
1091 {"type": "text", "text": "one"},
1092 {"type": "text", "text": "two"},
1093 ]}])
1094 );
1095 }
1096
1097 #[test]
1106 #[ignore]
1107 fn adaptive_system_corpus_audit() {
1108 let dir =
1109 std::env::var("TEMPLATE_CORPUS").expect("set TEMPLATE_CORPUS to the templates dir");
1110 let manifest: serde_json::Value =
1111 serde_json::from_str(&std::fs::read_to_string(format!("{dir}/manifest.json")).unwrap())
1112 .unwrap();
1113
1114 let sys = |c: &str| json!({"role": "system", "content": c});
1117 let usr = |c: &str| json!({"role": "user", "content": c});
1118 let asst = |c: &str| json!({"role": "assistant", "content": c});
1119 let shapes: Vec<(&str, serde_json::Value)> = vec![
1120 ("turn1", json!([sys("s"), usr("u"), sys("mid")])),
1121 (
1122 "multiturn",
1123 json!([sys("s"), usr("u"), sys("mid"), asst("a"), usr("u2")]),
1124 ),
1125 (
1126 "mid_after_asst",
1127 json!([sys("s"), usr("u"), asst("a"), sys("mid"), usr("u2")]),
1128 ),
1129 ("double_leading", json!([sys("s0"), sys("s1"), usr("u")])),
1130 ("consec_user", json!([sys("s"), usr("u0"), usr("u1")])),
1131 (
1132 "tail_reminder",
1133 json!([
1134 sys("s"),
1135 usr("u"),
1136 asst("a"),
1137 usr("u2"),
1138 sys("mid"),
1139 usr("u3")
1140 ]),
1141 ),
1142 ("leading_only_baseline", json!([sys("s"), usr("u")])),
1143 ];
1144
1145 let mut total = 0usize;
1146 let mut flagged = 0usize;
1147 let mut demote_only = 0usize;
1148 let mut coalesce = 0usize;
1149 let mut failures: Vec<String> = Vec::new();
1150 for (file, meta) in manifest.as_object().unwrap() {
1151 let tmpl = std::fs::read_to_string(format!("{dir}/{file}.jinja")).unwrap();
1152 let model = meta["model"].as_str().unwrap_or(file);
1153 let f = match try_formatter_for(&tmpl) {
1156 Some(f) => f,
1157 None => {
1158 eprintln!("[skip-compile] {model}");
1159 continue;
1160 }
1161 };
1162 if render_shape(&f, json!([sys("s"), usr("u")])).is_err() {
1165 eprintln!("[skip-baseline] {model}");
1166 continue;
1167 }
1168 total += 1;
1169 let rules = f.default_system_normalization;
1170 let flag = rules.is_required();
1171 if flag {
1172 flagged += 1;
1173 }
1174 if rules.demote_nonleading_system {
1175 demote_only += usize::from(!rules.coalesce_consecutive_users);
1176 }
1177 if rules.coalesce_consecutive_users {
1178 coalesce += 1;
1179 }
1180 for (name, shape) in &shapes {
1181 if render_shape(&f, shape.clone()).is_err() {
1182 failures.push(format!("{model} | shape={name} | flag={flag}"));
1183 }
1184 }
1185 if flag {
1186 eprintln!(
1187 "[ok] demote={} coalesce={} {model}",
1188 rules.demote_nonleading_system, rules.coalesce_consecutive_users
1189 );
1190 }
1191 }
1192 eprintln!(
1193 "\naudited {total} templates ({flagged} flagged: {demote_only} demote-only, \
1194 {coalesce} coalescing); {} shape failures",
1195 failures.len()
1196 );
1197 for f in &failures {
1198 eprintln!(" FAIL {f}");
1199 }
1200 assert!(
1201 failures.is_empty(),
1202 "{} template/shape combinations did not render (probe insufficient or normalization insufficient)",
1203 failures.len()
1204 );
1205 }
1206
1207 #[test]
1215 fn test_render_long_conversation_does_not_overflow_stack() {
1216 let handle = std::thread::Builder::new()
1217 .stack_size(2 * 1024 * 1024)
1218 .spawn(|| {
1219 let template_string = concat!(
1220 "{%- set ns = namespace(items=[]) -%}",
1221 "{%- for m in messages -%}",
1222 "{%- set ns.items = ns.items + [m] -%}",
1223 "{%- endfor -%}",
1224 "COUNT={{ ns.items | length }}"
1225 );
1226 let chat_template: ChatTemplate =
1227 serde_json::from_value(serde_json::json!({ "chat_template": template_string }))
1228 .unwrap();
1229 let formatter =
1230 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[]))
1231 .unwrap();
1232
1233 let n = 3000;
1234 let messages: Vec<serde_json::Value> = (0..n)
1235 .map(|i| serde_json::json!({"role": "user", "content": format!("turn {i}")}))
1236 .collect();
1237 let request: NvCreateChatCompletionRequest =
1238 serde_json::from_value(serde_json::json!({
1239 "model": "test",
1240 "messages": messages,
1241 }))
1242 .unwrap();
1243
1244 let rendered = formatter.render(&request).unwrap();
1246 assert_eq!(rendered.trim(), format!("COUNT={n}"));
1247 })
1248 .unwrap();
1249 handle.join().unwrap();
1250 }
1251
1252 #[test]
1263 #[ignore]
1264 fn dump_gptoss_tool_prompt() {
1265 use super::tokcfg::ChatTemplate;
1266 use super::{ContextMixins, HfTokenizerConfigJsonFormatter};
1267
1268 let path = std::env::var("GPTOSS_CHAT_TEMPLATE").expect(
1269 "set GPTOSS_CHAT_TEMPLATE to the tokenizer_config.json, chat_template.jinja, or model dir path",
1270 );
1271 let input_path = std::path::Path::new(&path);
1272 let file_path = if input_path.is_dir() {
1273 input_path.join("tokenizer_config.json")
1275 } else {
1276 input_path.to_path_buf()
1277 };
1278 let raw = std::fs::read_to_string(&file_path).expect("read chat template file");
1279 let template_string: String = match serde_json::from_str::<serde_json::Value>(&raw) {
1286 Ok(v) if v.get("chat_template").is_some() => v["chat_template"]
1287 .as_str()
1288 .expect("chat_template field must be a string")
1289 .to_string(),
1290 _ => {
1291 let sibling = std::path::Path::new(&path)
1292 .parent()
1293 .map(|d| d.join("chat_template.jinja"));
1294 match sibling {
1295 Some(p) if p.exists() => {
1296 eprintln!(
1297 "[info] {path} had no chat_template field; using {}",
1298 p.display()
1299 );
1300 std::fs::read_to_string(&p).expect("read sibling chat_template.jinja")
1301 }
1302 _ => raw,
1303 }
1304 }
1305 };
1306
1307 assert!(
1310 template_string.contains("{%") || template_string.contains("{{"),
1311 "resolved template has no Jinja tags — GPTOSS_CHAT_TEMPLATE ({path}) is probably \
1312 tokenizer_config.json with no chat_template field and no sibling chat_template.jinja. \
1313 Point it at the chat_template.jinja file."
1314 );
1315
1316 let chat_template: ChatTemplate =
1317 serde_json::from_value(serde_json::json!({ "chat_template": template_string }))
1318 .unwrap();
1319
1320 let formatter =
1321 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
1322
1323 let request: NvCreateChatCompletionRequest = serde_json::from_str(
1325 r#"{
1326 "model": "openai/gpt-oss-120b",
1327 "messages": [{"role":"user","content":"Search the repo for the string \"countHook\"."}],
1328 "tools": [
1329 {"type":"function","function":{"name":"grep","description":"search files","parameters":{"type":"object","properties":{"pattern":{"type":"string"},"path":{"type":"string"}},"required":["pattern"]}}},
1330 {"type":"function","function":{"name":"read","description":"read a file","parameters":{"type":"object","properties":{"filePath":{"type":"string"}},"required":["filePath"]}}}
1331 ]
1332 }"#,
1333 )
1334 .unwrap();
1335
1336 let rendered = formatter.render(&request).unwrap();
1337 eprintln!("================ RENDERED gpt-oss PROMPT (tools declared) ================");
1338 eprintln!("{rendered}");
1339 eprintln!("================ END RENDERED PROMPT ================");
1340 eprintln!("[diagnostics] does the rendered prompt contain…");
1341 for needle in [
1342 "commentary",
1343 "Calls to these tools",
1344 "functions",
1345 "# Tools",
1346 "<|channel|>",
1347 "constrain",
1348 "analysis",
1349 ] {
1350 eprintln!(
1351 " {:>22}: {}",
1352 format!("{needle:?}"),
1353 rendered.contains(needle)
1354 );
1355 }
1356 }
1357
1358 #[test]
1360 fn test_convert_media_url_to_placeholder_single_type() {
1361 let mut content_array = vec![
1362 serde_json::json!({"type": "text", "text": "Check this image:"}),
1363 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1364 serde_json::json!({"type": "text", "text": "What do you see?"}),
1365 ];
1366
1367 let conversions = &[("image_url", "image")];
1368 convert_media_url_to_placeholder(&mut content_array, conversions);
1369
1370 assert_eq!(content_array.len(), 3);
1371 assert_eq!(content_array[0]["type"], "text");
1373 assert_eq!(content_array[0]["text"], "Check this image:");
1374 assert_eq!(content_array[1]["type"], "image");
1376 assert!(content_array[1].get("image_url").is_none());
1377 assert_eq!(content_array[2]["type"], "text");
1379 assert_eq!(content_array[2]["text"], "What do you see?");
1380 }
1381
1382 #[test]
1384 fn test_convert_media_url_to_placeholder_multiple_same_type() {
1385 let mut content_array = vec![
1386 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image1.jpg"}}),
1387 serde_json::json!({"type": "text", "text": "vs"}),
1388 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image2.jpg"}}),
1389 ];
1390
1391 let conversions = &[("image_url", "image")];
1392 convert_media_url_to_placeholder(&mut content_array, conversions);
1393
1394 assert_eq!(content_array.len(), 3);
1395 assert_eq!(content_array[0]["type"], "image");
1396 assert_eq!(content_array[1]["type"], "text");
1397 assert_eq!(content_array[2]["type"], "image");
1398 }
1399
1400 #[test]
1402 fn test_convert_media_url_to_placeholder_selective_conversion() {
1403 let mut content_array = vec![
1404 serde_json::json!({"type": "audio_url", "audio_url": {"url": "https://example.com/audio.mp3"}}),
1405 serde_json::json!({"type": "video_url", "video_url": {"url": "https://example.com/video.mp4"}}),
1406 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1407 ];
1408
1409 let conversions = &[("image_url", "image")];
1411 convert_media_url_to_placeholder(&mut content_array, conversions);
1412
1413 assert_eq!(content_array.len(), 3);
1414 assert_eq!(content_array[0]["type"], "audio_url");
1416 assert!(content_array[0].get("audio_url").is_some());
1417 assert_eq!(content_array[1]["type"], "video_url");
1418 assert!(content_array[1].get("video_url").is_some());
1419 assert_eq!(content_array[2]["type"], "image");
1421 assert!(content_array[2].get("image_url").is_none());
1422 }
1423
1424 #[test]
1426 fn test_convert_media_url_to_placeholder_multiple_types() {
1427 let mut content_array = vec![
1428 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1429 serde_json::json!({"type": "text", "text": "and listen to"}),
1430 serde_json::json!({"type": "audio_url", "audio_url": {"url": "https://example.com/audio.mp3"}}),
1431 serde_json::json!({"type": "text", "text": "and watch"}),
1432 serde_json::json!({"type": "video_url", "video_url": {"url": "https://example.com/video.mp4"}}),
1433 ];
1434
1435 let conversions = &[
1437 ("image_url", "image"),
1438 ("audio_url", "audio"),
1439 ("video_url", "video"),
1440 ];
1441 convert_media_url_to_placeholder(&mut content_array, conversions);
1442
1443 assert_eq!(content_array.len(), 5);
1444 assert_eq!(content_array[0]["type"], "image");
1445 assert!(content_array[0].get("image_url").is_none());
1446 assert_eq!(content_array[1]["type"], "text");
1447 assert_eq!(content_array[2]["type"], "audio");
1448 assert!(content_array[2].get("audio_url").is_none());
1449 assert_eq!(content_array[3]["type"], "text");
1450 assert_eq!(content_array[4]["type"], "video");
1451 assert!(content_array[4].get("video_url").is_none());
1452 }
1453
1454 #[test]
1456 fn test_convert_media_url_to_placeholder_no_conversions() {
1457 let mut content_array = vec![
1458 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1459 serde_json::json!({"type": "text", "text": "hello"}),
1460 ];
1461
1462 let conversions: &[(&str, &str)] = &[];
1463 convert_media_url_to_placeholder(&mut content_array, conversions);
1464
1465 assert_eq!(content_array.len(), 2);
1466 assert_eq!(content_array[0]["type"], "image_url");
1468 assert!(content_array[0].get("image_url").is_some());
1469 assert_eq!(content_array[1]["type"], "text");
1470 }
1471
1472 #[test]
1475 fn test_default_media_type_conversions_only_converts_image_url() {
1476 let mut content_array = vec![
1477 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1478 serde_json::json!({"type": "video_url", "video_url": {"url": "https://example.com/video.mp4"}}),
1479 serde_json::json!({"type": "audio_url", "audio_url": {"url": "https://example.com/audio.mp3"}}),
1480 serde_json::json!({"type": "text", "text": "hello"}),
1481 ];
1482
1483 convert_media_url_to_placeholder(&mut content_array, DEFAULT_MEDIA_TYPE_CONVERSIONS);
1485
1486 assert_eq!(content_array.len(), 4);
1487
1488 assert_eq!(content_array[0]["type"], "image");
1490 assert!(content_array[0].get("image_url").is_none());
1491
1492 assert_eq!(content_array[1]["type"], "video");
1494 assert!(content_array[1].get("video_url").is_none());
1495
1496 assert_eq!(content_array[2]["type"], "audio");
1498 assert!(content_array[2].get("audio_url").is_none());
1499
1500 assert_eq!(content_array[3]["type"], "text");
1502 assert_eq!(content_array[3]["text"], "hello");
1503 }
1504
1505 #[test]
1506 fn test_may_be_fix_tool_schema_missing_type_and_properties() {
1507 let json_str = r#"{
1508 "model": "gpt-4o",
1509 "messages": [],
1510 "tools": [
1511 {
1512 "type": "function",
1513 "function": {
1514 "name": "get_weather",
1515 "description": "Get the current weather in a given location",
1516 "parameters": {},
1517 "strict": null
1518 }
1519 }
1520 ]
1521 }"#;
1522
1523 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1524 let tools = serde_json::to_value(request.tools()).unwrap();
1525
1526 assert!(tools[0]["function"]["parameters"]["type"] == "object");
1527 assert!(
1528 tools[0]["function"]["parameters"]["properties"]
1529 == serde_json::Value::Object(Default::default())
1530 );
1531 }
1532
1533 #[test]
1534 fn test_may_be_fix_tool_schema_missing_type() {
1535 let json_str = r#"{
1536 "model": "gpt-4o",
1537 "messages": [],
1538 "tools": [
1539 {
1540 "type": "function",
1541 "function": {
1542 "name": "get_weather",
1543 "description": "Get the current weather in a given location",
1544 "parameters": {
1545 "properties": {
1546 "location": {
1547 "type": "string",
1548 "description": "City and state, e.g., 'San Francisco, CA'"
1549 }
1550 }
1551 },
1552 "strict": null
1553 }
1554 }
1555 ]
1556 }"#;
1557 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1558
1559 let tools = serde_json::to_value(request.tools()).unwrap();
1560
1561 assert_eq!(tools[0]["function"]["parameters"]["type"], "object");
1562
1563 let mut expected_properties = serde_json::Map::new();
1564 let mut location = serde_json::Map::new();
1565 location.insert(
1566 "type".to_string(),
1567 serde_json::Value::String("string".to_string()),
1568 );
1569 location.insert(
1570 "description".to_string(),
1571 serde_json::Value::String("City and state, e.g., 'San Francisco, CA'".to_string()),
1572 );
1573 expected_properties.insert("location".to_string(), serde_json::Value::Object(location));
1574
1575 assert_eq!(
1576 tools[0]["function"]["parameters"]["properties"],
1577 serde_json::Value::Object(expected_properties)
1578 );
1579 }
1580
1581 #[test]
1582 fn test_may_be_fix_tool_schema_missing_properties() {
1583 let json_str = r#"{
1584 "model": "gpt-4o",
1585 "messages": [],
1586 "tools": [
1587 {
1588 "type": "function",
1589 "function": {
1590 "name": "get_weather",
1591 "description": "Get the current weather in a given location",
1592 "parameters": {"type": "object"},
1593 "strict": null
1594 }
1595 }
1596 ]
1597 }"#;
1598
1599 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1600 let tools = serde_json::to_value(request.tools()).unwrap();
1601
1602 assert_eq!(
1603 tools[0]["function"]["parameters"]["properties"],
1604 serde_json::Value::Object(Default::default())
1605 );
1606 assert_eq!(tools[0]["function"]["parameters"]["type"], "object");
1607 }
1608
1609 #[test]
1610 fn test_may_be_fix_tool_schema_missing_description() {
1611 let json_str = r#"{
1615 "model": "gpt-4o",
1616 "messages": [],
1617 "tools": [
1618 {
1619 "type": "function",
1620 "function": {
1621 "name": "noop",
1622 "parameters": {
1623 "type": "object",
1624 "properties": { "x": { "type": "string" } },
1625 "required": ["x"],
1626 "additionalProperties": false
1627 },
1628 "strict": null
1629 }
1630 }
1631 ]
1632 }"#;
1633
1634 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1635 let tools = serde_json::to_value(request.tools()).unwrap();
1636
1637 assert_eq!(
1638 tools[0]["function"]["description"],
1639 serde_json::Value::String(String::new())
1640 );
1641 }
1642
1643 #[test]
1644 fn test_may_be_fix_tool_schema_null_description() {
1645 let json_str = r#"{
1647 "model": "gpt-4o",
1648 "messages": [],
1649 "tools": [
1650 {
1651 "type": "function",
1652 "function": {
1653 "name": "noop",
1654 "description": null,
1655 "parameters": {"type": "object", "properties": {}},
1656 "strict": null
1657 }
1658 }
1659 ]
1660 }"#;
1661
1662 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1663 let tools = serde_json::to_value(request.tools()).unwrap();
1664
1665 assert_eq!(
1666 tools[0]["function"]["description"],
1667 serde_json::Value::String(String::new())
1668 );
1669 }
1670
1671 #[test]
1672 fn test_may_be_fix_tool_schema_preserves_description() {
1673 let json_str = r#"{
1675 "model": "gpt-4o",
1676 "messages": [],
1677 "tools": [
1678 {
1679 "type": "function",
1680 "function": {
1681 "name": "get_weather",
1682 "description": "Get the current weather in a given location",
1683 "parameters": {"type": "object", "properties": {}},
1684 "strict": null
1685 }
1686 }
1687 ]
1688 }"#;
1689
1690 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1691 let tools = serde_json::to_value(request.tools()).unwrap();
1692
1693 assert_eq!(
1694 tools[0]["function"]["description"],
1695 "Get the current weather in a given location"
1696 );
1697 }
1698
1699 #[test]
1701 fn test_may_be_fix_msg_content_user_multipart() {
1702 let json_str = r#"{
1703 "model": "gpt-4o",
1704 "messages": [
1705 {
1706 "role": "user",
1707 "content": [
1708 {"type": "text", "text": "part 1"},
1709 {"type": "text", "text": "part 2"}
1710 ]
1711 }
1712 ]
1713 }"#;
1714
1715 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1716 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1717
1718 let messages =
1720 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1721
1722 assert_eq!(
1724 messages[0]["content"],
1725 serde_json::Value::String("part 1\npart 2".to_string())
1726 );
1727 }
1728
1729 #[test]
1732 fn test_may_be_fix_msg_content_mixed_messages() {
1733 let json_str = r#"{
1734 "model": "gpt-4o",
1735 "messages": [
1736 {
1737 "role": "system",
1738 "content": "You are a helpful assistant"
1739 },
1740 {
1741 "role": "user",
1742 "content": [
1743 {"type": "text", "text": "Hello"},
1744 {"type": "text", "text": "World"}
1745 ]
1746 },
1747 {
1748 "role": "assistant",
1749 "content": "Hi there!"
1750 },
1751 {
1752 "role": "user",
1753 "content": [
1754 {"type": "text", "text": "Another"},
1755 {"type": "text", "text": "multi-part"},
1756 {"type": "text", "text": "message"}
1757 ]
1758 }
1759 ]
1760 }"#;
1761
1762 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1763 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1764
1765 let messages =
1767 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1768
1769 assert_eq!(
1771 messages[0]["content"],
1772 serde_json::Value::String("You are a helpful assistant".to_string())
1773 );
1774
1775 assert_eq!(
1777 messages[1]["content"],
1778 serde_json::Value::String("Hello\nWorld".to_string())
1779 );
1780
1781 assert_eq!(
1783 messages[2]["content"],
1784 serde_json::Value::String("Hi there!".to_string())
1785 );
1786
1787 assert_eq!(
1789 messages[3]["content"],
1790 serde_json::Value::String("Another\nmulti-part\nmessage".to_string())
1791 );
1792 }
1793
1794 #[test]
1796 fn test_may_be_fix_msg_content_empty_array() {
1797 let json_str = r#"{
1798 "model": "gpt-4o",
1799 "messages": [
1800 {
1801 "role": "user",
1802 "content": []
1803 }
1804 ]
1805 }"#;
1806
1807 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1808 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1809
1810 let messages =
1812 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1813
1814 assert!(messages[0]["content"].is_array());
1816 assert_eq!(messages[0]["content"].as_array().unwrap().len(), 0);
1817 }
1818
1819 #[test]
1826 fn test_may_be_fix_msg_content_empty_array_with_placeholder_template() {
1827 let json_str = r#"{
1828 "model": "phi-3-vision",
1829 "messages": [
1830 {
1831 "role": "user",
1832 "content": []
1833 }
1834 ]
1835 }"#;
1836
1837 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1838 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1839
1840 let messages = serde_json::to_value(may_be_fix_msg_content(
1843 messages_raw,
1844 false,
1845 Some("<|image_{n}|>"),
1846 ))
1847 .unwrap();
1848
1849 assert!(
1850 messages[0]["content"].is_array(),
1851 "empty array should be preserved as `[]`, not flattened to `\"\"`"
1852 );
1853 assert_eq!(messages[0]["content"].as_array().unwrap().len(), 0);
1854 }
1855
1856 #[test]
1858 fn test_may_be_fix_msg_content_single_text() {
1859 let json_str = r#"{
1860 "model": "gpt-4o",
1861 "messages": [
1862 {
1863 "role": "user",
1864 "content": "Simple text message"
1865 }
1866 ]
1867 }"#;
1868
1869 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1870 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1871
1872 let messages =
1874 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1875
1876 assert_eq!(
1878 messages[0]["content"],
1879 serde_json::Value::String("Simple text message".to_string())
1880 );
1881 }
1882
1883 #[test]
1886 fn test_may_be_fix_msg_content_mixed_types() {
1887 let json_str = r#"{
1888 "model": "gpt-4o",
1889 "messages": [
1890 {
1891 "role": "user",
1892 "content": [
1893 {"type": "text", "text": "Check this image:"},
1894 {"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}},
1895 {"type": "text", "text": "What do you see?"}
1896 ]
1897 }
1898 ]
1899 }"#;
1900
1901 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1902 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1903
1904 let messages =
1906 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1907
1908 assert!(messages[0]["content"].is_array());
1911 let content_array = messages[0]["content"].as_array().unwrap();
1912 assert_eq!(content_array.len(), 3);
1913 assert_eq!(content_array[0]["type"], "text");
1914 assert_eq!(content_array[1]["type"], "image");
1915 assert!(content_array[1].get("image_url").is_none());
1916 assert_eq!(content_array[2]["type"], "text");
1917 }
1918
1919 #[test]
1925 fn test_may_be_fix_msg_content_flattens_phi3_style() {
1926 let json_str = r#"{
1927 "model": "phi-3-vision",
1928 "messages": [
1929 {
1930 "role": "user",
1931 "content": [
1932 {"type": "text", "text": "First "},
1933 {"type": "image_url", "image_url": {"url": "https://example.com/a.jpg"}},
1934 {"type": "text", "text": " then "},
1935 {"type": "image_url", "image_url": {"url": "https://example.com/b.jpg"}},
1936 {"type": "text", "text": "?"}
1937 ]
1938 }
1939 ]
1940 }"#;
1941 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1942 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1943
1944 let messages = serde_json::to_value(may_be_fix_msg_content(
1945 messages_raw,
1946 false,
1947 Some("<|image_{n}|>"),
1948 ))
1949 .unwrap();
1950
1951 let content = messages[0]["content"].as_str().expect("content flattened");
1952 assert_eq!(content, "First <|image_1|> then <|image_2|>?");
1953 }
1954
1955 #[test]
1957 fn test_may_be_fix_msg_content_flattens_llava_style() {
1958 let json_str = r#"{
1959 "model": "llava-1.5-7b-hf",
1960 "messages": [
1961 {
1962 "role": "user",
1963 "content": [
1964 {"type": "text", "text": "Describe: "},
1965 {"type": "image_url", "image_url": {"url": "https://example.com/x.jpg"}}
1966 ]
1967 }
1968 ]
1969 }"#;
1970 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1971 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1972
1973 let messages =
1974 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, Some("<image>")))
1975 .unwrap();
1976
1977 let content = messages[0]["content"].as_str().expect("content flattened");
1978 assert_eq!(content, "Describe: <image>");
1979 }
1980
1981 #[test]
1987 fn test_may_be_fix_msg_content_flattens_empty_placeholder() {
1988 let json_str = r#"{
1989 "model": "nvidia/NVIDIA-Nemotron-Parse-v1.2",
1990 "messages": [
1991 {
1992 "role": "user",
1993 "content": [
1994 {"type": "text", "text": "</s><s><predict_bbox><predict_classes><output_markdown><predict_no_text_in_pic>"},
1995 {"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}}
1996 ]
1997 }
1998 ]
1999 }"#;
2000 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2001 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2002
2003 let messages =
2004 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, Some(""))).unwrap();
2005
2006 let content = messages[0]["content"].as_str().expect("content flattened");
2007 assert_eq!(
2008 content,
2009 "</s><s><predict_bbox><predict_classes><output_markdown><predict_no_text_in_pic>"
2010 );
2011 }
2012
2013 #[test]
2020 fn test_render_nemotron_parse_passthrough() {
2021 use super::super::tokcfg::ChatTemplate;
2022 use super::{ContextMixins, HfTokenizerConfigJsonFormatter};
2023
2024 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
2025 "chat_template": "{% for message in messages %}{{ message['content'] }}{% endfor %}"
2026 }))
2027 .unwrap();
2028 let formatter =
2029 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
2030
2031 for prompt in [
2032 "</s><s><predict_bbox><predict_classes><output_markdown><predict_no_text_in_pic>",
2033 "</s><s><predict_bbox><predict_classes><output_markdown><predict_text_in_pic>",
2034 ] {
2035 let request: NvCreateChatCompletionRequest =
2036 serde_json::from_value(serde_json::json!({
2037 "model": "nvidia/NVIDIA-Nemotron-Parse-v1.2",
2038 "messages": [{
2039 "role": "user",
2040 "content": [
2041 {"type": "text", "text": prompt},
2042 {"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}}
2043 ]
2044 }]
2045 }))
2046 .unwrap();
2047
2048 let rendered = formatter.render(&request).unwrap();
2049 assert_eq!(
2050 rendered, prompt,
2051 "rendered prompt must be the control tokens only, with no JSON-serialized image array"
2052 );
2053 }
2054 }
2055
2056 #[test]
2059 fn test_may_be_fix_msg_content_non_text_only() {
2060 let json_str = r#"{
2061 "model": "gpt-4o",
2062 "messages": [
2063 {
2064 "role": "user",
2065 "content": [
2066 {"type": "image_url", "image_url": {"url": "https://example.com/image1.jpg"}},
2067 {"type": "image_url", "image_url": {"url": "https://example.com/image2.jpg"}}
2068 ]
2069 }
2070 ]
2071 }"#;
2072
2073 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2074 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2075
2076 let messages =
2078 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
2079
2080 assert!(messages[0]["content"].is_array());
2082 let content_array = messages[0]["content"].as_array().unwrap();
2083 assert_eq!(content_array.len(), 2);
2084 assert_eq!(content_array[0]["type"], "image");
2085 assert_eq!(content_array[1]["type"], "image");
2086 }
2087
2088 #[test]
2089 fn test_none_tools_safe_for_all_templates() {
2090 use super::tokcfg::ChatTemplate;
2091 use super::{ContextMixins, HfTokenizerConfigJsonFormatter};
2092
2093 let length_template = r#"
2097{%- if tools is iterable and tools | length > 0 %}
2098Tools available: {{ tools | length }}
2099{%- else %}
2100No tools
2101{%- endif %}
2102"#;
2103
2104 let no_tool_template = r#"
2107{%- if tools is not none %}
2108TOOL MODE
2109{%- else %}
2110NORMAL MODE
2111{%- endif %}
2112"#;
2113
2114 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
2115 "chat_template": [
2116 {"safe_length": length_template},
2117 {"no_tool": no_tool_template}
2118 ]
2119 }))
2120 .unwrap();
2121
2122 let formatter =
2123 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
2124
2125 let ctx = context! { tools => Option::<Value>::None };
2126
2127 let result1 = formatter
2128 .env
2129 .get_template("safe_length")
2130 .unwrap()
2131 .render(&ctx);
2132 println!("Safe length template with no tools => None: {:?}", result1);
2133 assert!(
2134 result1.is_ok(),
2135 "Jinja template with and conditional and length filter should handle None: {:?}",
2136 result1
2137 );
2138 assert!(
2139 result1.unwrap().contains("No tools"),
2140 "Should show 'No tools'"
2141 );
2142
2143 let result2 = formatter.env.get_template("no_tool").unwrap().render(&ctx);
2144 println!("Default template with no tools => None: {:?}", result2);
2145 assert!(
2146 result2.is_ok(),
2147 "Jinja template with if tools is not none conditional should handle None: {:?}",
2148 result2
2149 );
2150 assert!(result2.unwrap().contains("NORMAL MODE"));
2151 }
2152
2153 #[test]
2155 fn test_may_be_fix_msg_content_multiple_content_types() {
2156 let json_str = r#"{
2158 "model": "gpt-4o",
2159 "messages": [
2160 {
2161 "role": "user",
2162 "content": [
2163 {"type": "text", "text": "Listen to this:"},
2164 {"type": "audio_url", "audio_url": {"url": "https://example.com/audio.mp3"}},
2165 {"type": "text", "text": "And look at:"},
2166 {"type": "image_url", "image_url": {"url": "https://example.com/img.jpg"}},
2167 {"type": "text", "text": "What do you think?"}
2168 ]
2169 }
2170 ]
2171 }"#;
2172
2173 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2174 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2175 let messages =
2176 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
2177
2178 assert!(messages[0]["content"].is_array());
2180 let content_array = messages[0]["content"].as_array().unwrap();
2181 assert_eq!(content_array.len(), 5);
2182 assert_eq!(content_array[0]["type"], "text");
2183 assert_eq!(content_array[1]["type"], "audio");
2184 assert_eq!(content_array[2]["type"], "text");
2185 assert_eq!(content_array[3]["type"], "image");
2186 assert_eq!(content_array[4]["type"], "text");
2187
2188 let json_str = r#"{
2190 "model": "gpt-4o",
2191 "messages": [
2192 {
2193 "role": "user",
2194 "content": [
2195 {"type": "text", "text": "Check this:"},
2196 {"type": "video_url", "video_url": {"url": "https://example.com/vid.mp4"}},
2197 {"type": "text", "text": "Interesting?"}
2198 ]
2199 }
2200 ]
2201 }"#;
2202
2203 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2204 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2205 let messages =
2206 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
2207
2208 assert!(messages[0]["content"].is_array());
2210 assert_eq!(messages[0]["content"].as_array().unwrap().len(), 3);
2211 }
2212
2213 #[test]
2214 fn test_normalize_tool_arguments_tojson() {
2215 let tmpl = r#"{{ messages[0].tool_calls[0].function.arguments | tojson }}"#;
2216
2217 let mut messages = serde_json::Value::Array(vec![serde_json::json!({
2219 "role": "assistant",
2220 "tool_calls": [{
2221 "type": "function",
2222 "function": {
2223 "name": "get_current_weather",
2224 "arguments": "{\"format\":\"celsius\",\"location\":\"San Francisco, CA\"}"
2225 }
2226 }]
2227 })]);
2228
2229 normalize_tool_calls_arguments_in_messages(&mut messages);
2230
2231 let mut env = Environment::new();
2232 env.add_filter("tojson", super::super::tokcfg::tojson);
2233 env.add_template("t", tmpl).unwrap();
2234 let out = env
2235 .get_template("t")
2236 .unwrap()
2237 .render(context! { messages => messages.as_array().unwrap() })
2238 .unwrap();
2239
2240 assert_eq!(
2243 out,
2244 r#"{"format": "celsius", "location": "San Francisco, CA"}"#
2245 );
2246 }
2247
2248 #[test]
2249 fn test_normalize_tool_arguments_items_loop() {
2250 let tmpl = r#"{% for k, v in messages[0].tool_calls[0].function.arguments|items %}{{k}}={{v}};{% endfor %}"#;
2251
2252 let mut messages = serde_json::Value::Array(vec![serde_json::json!({
2253 "role": "assistant",
2254 "tool_calls": [{
2255 "type": "function",
2256 "function": {
2257 "name": "f",
2258 "arguments": "{\"a\":1,\"b\":\"x\"}"
2259 }
2260 }]
2261 })]);
2262
2263 normalize_tool_calls_arguments_in_messages(&mut messages);
2264
2265 let mut env = Environment::new();
2266 env.add_template("t", tmpl).unwrap();
2267 let out = env
2268 .get_template("t")
2269 .unwrap()
2270 .render(context! { messages => messages.as_array().unwrap() })
2271 .unwrap();
2272
2273 assert!(out == "a=1;b=x;" || out == "b=x;a=1;");
2274 }
2275
2276 #[test]
2277 fn test_normalize_tool_arguments_legacy_function_call() {
2278 let mut messages = serde_json::Value::Array(vec![serde_json::json!({
2280 "role": "assistant",
2281 "function_call": {
2282 "name": "get_weather",
2283 "arguments": "{\"location\":\"NYC\"}"
2284 }
2285 })]);
2286
2287 normalize_function_call_arguments_in_messages(&mut messages);
2288
2289 assert_eq!(
2290 messages[0]["function_call"]["arguments"],
2291 serde_json::json!({"location": "NYC"})
2292 );
2293 }
2294
2295 #[test]
2296 fn test_normalize_tool_arguments_malformed_json_passthrough() {
2297 let mut messages = serde_json::Value::Array(vec![serde_json::json!({
2299 "role": "assistant",
2300 "tool_calls": [{
2301 "type": "function",
2302 "function": {
2303 "name": "f",
2304 "arguments": "not valid json at all"
2305 }
2306 }]
2307 })]);
2308
2309 normalize_tool_calls_arguments_in_messages(&mut messages);
2310
2311 assert_eq!(
2312 messages[0]["tool_calls"][0]["function"]["arguments"],
2313 serde_json::Value::String("not valid json at all".to_string())
2314 );
2315 }
2316
2317 #[test]
2318 fn test_normalize_tool_arguments_with_multimodal_content() {
2319 let json_str = r#"{
2320 "model": "gpt-4o",
2321 "messages": [
2322 {
2323 "role": "user",
2324 "content": [
2325 {"type": "text", "text": "Check this:"},
2326 {"type": "video_url", "video_url": {"url": "https://example.com/vid.mp4"}},
2327 {"type": "text", "text": "Interesting?"}
2328 ]
2329 },
2330 {
2331 "role": "assistant",
2332 "tool_calls": [{
2333 "id": "call_123",
2334 "type": "function",
2335 "function": {
2336 "name": "analyze_video",
2337 "arguments": "{\"url\":\"https://example.com/vid.mp4\",\"format\":\"mp4\"}"
2338 }
2339 }]
2340 }
2341 ]
2342 }"#;
2343
2344 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2345 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2346
2347 let mut messages =
2349 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
2350
2351 normalize_tool_calls_arguments_in_messages(&mut messages);
2352
2353 assert!(messages[0]["content"].is_array());
2355 assert_eq!(messages[0]["content"].as_array().unwrap().len(), 3);
2356
2357 assert!(messages[1]["tool_calls"][0]["function"]["arguments"].is_object());
2359 assert_eq!(
2360 messages[1]["tool_calls"][0]["function"]["arguments"]["url"],
2361 "https://example.com/vid.mp4"
2362 );
2363 }
2364
2365 #[test]
2367 fn test_may_be_fix_msg_content_string_to_array() {
2368 let json_str = r#"{
2369 "model": "gpt-4o",
2370 "messages": [
2371 {
2372 "role": "user",
2373 "content": "Hello, how are you?"
2374 }
2375 ]
2376 }"#;
2377
2378 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2379 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2380
2381 let messages =
2383 serde_json::to_value(may_be_fix_msg_content(messages_raw, true, None)).unwrap();
2384
2385 assert!(messages[0]["content"].is_array());
2387 let content_array = messages[0]["content"].as_array().unwrap();
2388 assert_eq!(content_array.len(), 1);
2389 assert_eq!(content_array[0]["type"], "text");
2390 assert_eq!(content_array[0]["text"], "Hello, how are you?");
2391 }
2392
2393 #[test]
2395 fn test_may_be_fix_msg_content_array_preserved_with_multimodal() {
2396 let json_str = r#"{
2397 "model": "gpt-4o",
2398 "messages": [
2399 {
2400 "role": "user",
2401 "content": [
2402 {"type": "text", "text": "part 1"},
2403 {"type": "text", "text": "part 2"}
2404 ]
2405 }
2406 ]
2407 }"#;
2408
2409 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2410 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2411
2412 let messages =
2414 serde_json::to_value(may_be_fix_msg_content(messages_raw, true, None)).unwrap();
2415
2416 assert!(messages[0]["content"].is_array());
2418 let content_array = messages[0]["content"].as_array().unwrap();
2419 assert_eq!(content_array.len(), 2);
2420 assert_eq!(content_array[0]["text"], "part 1");
2421 assert_eq!(content_array[1]["text"], "part 2");
2422 }
2423
2424 fn user() -> Msg {
2425 Msg::User(Default::default())
2426 }
2427 fn tool() -> Msg {
2428 Msg::Tool(Default::default())
2429 }
2430
2431 fn dummy_state(messages: Vec<Msg>) -> NvCreateChatCompletionRequest {
2432 let json = serde_json::json!({
2433 "model": "test-model",
2434 "messages": messages
2435 });
2436 serde_json::from_value(json).unwrap()
2437 }
2438
2439 #[test]
2440 fn add_after_user() {
2441 let s = dummy_state(vec![user()]);
2442 assert!(s.should_add_generation_prompt());
2443 }
2444
2445 #[test]
2446 fn add_after_tool() {
2447 let s = dummy_state(vec![tool()]);
2448 assert!(s.should_add_generation_prompt());
2449 }
2450
2451 #[test]
2452 fn add_when_empty() {
2453 let s = dummy_state(vec![]);
2454 assert!(s.should_add_generation_prompt());
2455 }
2456
2457 fn tool_aware_formatter(
2459 exclude_tools_when_tool_choice_none: bool,
2460 ) -> HfTokenizerConfigJsonFormatter {
2461 let template = r#"
2462{%- if tools is iterable and tools | length > 0 %}
2463TOOL_MODE tools={{ tools | length }}
2464{%- else %}
2465NORMAL_MODE
2466{%- endif %}
2467{{ messages[0].content }}"#;
2468
2469 let chat_template: super::tokcfg::ChatTemplate =
2470 serde_json::from_value(serde_json::json!({ "chat_template": template })).unwrap();
2471
2472 HfTokenizerConfigJsonFormatter::with_options(
2473 chat_template,
2474 ContextMixins::new(&[]),
2475 exclude_tools_when_tool_choice_none,
2476 )
2477 .unwrap()
2478 }
2479
2480 fn gemma4_tool_template_for_tests() -> &'static str {
2481 r#"
2482{{ bos_token }}
2483{%- set loop_messages = messages -%}
2484{%- set ns_turn = namespace(last_user_idx=-1) -%}
2485{%- for i in range(loop_messages | length) -%}
2486 {%- if loop_messages[i]['role'] == 'user' -%}
2487 {%- set ns_turn.last_user_idx = i -%}
2488 {%- endif -%}
2489{%- endfor -%}
2490{%- for message in loop_messages -%}
2491 {%- set role = 'model' if message['role'] == 'assistant' else message['role'] -%}
2492 {{- '<|turn>' + role + '\n' }}
2493
2494 {%- if message.get('reasoning') and loop.index0 > ns_turn.last_user_idx and message.get('tool_calls') -%}
2495 {{- '<|channel>thought\n' + message['reasoning'] + '\n<channel|>'}}
2496 {%- endif -%}
2497
2498 {%- if message['tool_calls'] -%}
2499 {%- for tool_call in message['tool_calls'] -%}
2500 {%- set function = tool_call['function'] -%}
2501 {{- '<|tool_call>call:' + function['name'] + '{' -}}
2502 {%- if function['arguments'] is mapping -%}
2503 {%- set ns_args = namespace(found_first=false) -%}
2504 {%- for key, value in function['arguments'] | dictsort -%}
2505 {%- if ns_args.found_first %},{% endif -%}
2506 {%- set ns_args.found_first = true -%}
2507 {{- key -}}:{{- value -}}
2508 {%- endfor -%}
2509 {%- elif function['arguments'] is string -%}
2510 {{- function['arguments'] -}}
2511 {%- endif -%}
2512 {{- '}<tool_call|>' -}}
2513 {%- endfor -%}
2514 {%- endif -%}
2515
2516 {%- if message['content'] is string -%}
2517 {{- message['content'] -}}
2518 {%- endif -%}
2519 {{- '<turn|>\n' -}}
2520{%- endfor -%}
2521"#
2522 }
2523
2524 fn make_gemma4_tool_formatter_for_tests() -> HfTokenizerConfigJsonFormatter {
2525 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
2526 "chat_template": gemma4_tool_template_for_tests()
2527 }))
2528 .unwrap();
2529 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap()
2530 }
2531
2532 fn request_with_tool_choice(tool_choice: &str) -> NvCreateChatCompletionRequest {
2534 serde_json::from_value(serde_json::json!({
2535 "model": "test",
2536 "messages": [{"role": "user", "content": "hello"}],
2537 "tools": [{
2538 "type": "function",
2539 "function": {
2540 "name": "get_weather",
2541 "description": "Get weather",
2542 "parameters": {"type": "object", "properties": {"location": {"type": "string"}}}
2543 }
2544 }],
2545 "tool_choice": tool_choice
2546 }))
2547 .unwrap()
2548 }
2549
2550 #[test]
2551 fn test_exclude_tools_strips_when_tool_choice_none() {
2552 let formatter = tool_aware_formatter(true);
2553 let request = request_with_tool_choice("none");
2554 let result = formatter.render(&request).unwrap();
2555 assert!(
2556 result.contains("NORMAL_MODE"),
2557 "With exclude_tools=true and tool_choice=none, tools should be stripped. Got: {}",
2558 result
2559 );
2560 }
2561
2562 #[test]
2563 fn test_exclude_tools_keeps_when_tool_choice_auto() {
2564 let formatter = tool_aware_formatter(true);
2565 let request = request_with_tool_choice("auto");
2566 let result = formatter.render(&request).unwrap();
2567 assert!(
2568 result.contains("TOOL_MODE"),
2569 "With tool_choice=auto, tools should be included. Got: {}",
2570 result
2571 );
2572 }
2573
2574 #[test]
2575 fn test_no_exclude_tools_keeps_when_tool_choice_none() {
2576 let formatter = tool_aware_formatter(false);
2577 let request = request_with_tool_choice("none");
2578 let result = formatter.render(&request).unwrap();
2579 assert!(
2580 result.contains("TOOL_MODE"),
2581 "With exclude_tools=false and tool_choice=none, tools should NOT be stripped. Got: {}",
2582 result
2583 );
2584 }
2585
2586 #[test]
2587 fn test_inject_reasoning_content_segments_with_tool_calls() {
2588 let mut messages = serde_json::json!([
2590 {
2591 "role": "user",
2592 "content": "What is sqrt(144) and sqrt(256)?"
2593 },
2594 {
2595 "role": "assistant",
2596 "content": "Let me calculate those.",
2597 "reasoning_content": ["I need to compute sqrt(144)", "Now sqrt(256)", ""],
2598 "tool_calls": [
2599 {
2600 "id": "call_0",
2601 "type": "function",
2602 "function": {
2603 "name": "calculator",
2604 "arguments": "{\"expr\": \"sqrt(144)\"}"
2605 }
2606 },
2607 {
2608 "id": "call_1",
2609 "type": "function",
2610 "function": {
2611 "name": "calculator",
2612 "arguments": "{\"expr\": \"sqrt(256)\"}"
2613 }
2614 }
2615 ]
2616 }
2617 ]);
2618
2619 inject_reasoning_content_into_messages(&mut messages);
2620
2621 let assistant = &messages[1];
2622
2623 assert!(
2625 assistant.get("reasoning_content").is_none(),
2626 "reasoning_content should be removed after injection"
2627 );
2628
2629 let content = assistant["content"].as_str().unwrap();
2631 assert!(
2632 content.starts_with("<think>I need to compute sqrt(144)</think>"),
2633 "content should start with first reasoning segment, got: {}",
2634 content
2635 );
2636 assert!(
2637 content.contains("<think>Now sqrt(256)</think>"),
2638 "content should contain second reasoning segment"
2639 );
2640 assert!(
2642 !content.contains("<think></think>"),
2643 "empty segments should be skipped"
2644 );
2645 assert!(
2647 content.ends_with("Let me calculate those."),
2648 "original content should be at the end, got: {}",
2649 content
2650 );
2651
2652 assert!(assistant.get("tool_calls").is_some());
2654 assert_eq!(assistant["tool_calls"].as_array().unwrap().len(), 2);
2655 }
2656
2657 #[test]
2658 fn test_gemma4_template_renders_reasoning_content_segments_around_tool_calls() {
2659 let formatter = make_gemma4_tool_formatter_for_tests();
2660 assert!(
2661 formatter.tool_use_template_handles_reasoning,
2662 "Gemma4 template adaptation should make reasoning_content native"
2663 );
2664
2665 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
2666 "model": "gemma4-test",
2667 "messages": [
2668 {"role": "user", "content": "inspect two things"},
2669 {
2670 "role": "assistant",
2671 "content": null,
2672 "reasoning_content": [
2673 "Think before the first call.",
2674 "Think before the second call.",
2675 "Think after both calls."
2676 ],
2677 "tool_calls": [
2678 {
2679 "id": "call_0",
2680 "type": "function",
2681 "function": {
2682 "name": "first_tool",
2683 "arguments": "{\"path\":\".\"}"
2684 }
2685 },
2686 {
2687 "id": "call_1",
2688 "type": "function",
2689 "function": {
2690 "name": "second_tool",
2691 "arguments": "{\"path\":\"/tmp\"}"
2692 }
2693 }
2694 ]
2695 }
2696 ]
2697 }))
2698 .unwrap();
2699
2700 let rendered = formatter.render(&request).unwrap();
2701
2702 let expected = concat!(
2703 "<|channel>thought\nThink before the first call.\n<channel|>",
2704 "<|tool_call>call:first_tool{path:.}<tool_call|>",
2705 "<|channel>thought\nThink before the second call.\n<channel|>",
2706 "<|tool_call>call:second_tool{path:/tmp}<tool_call|>",
2707 "<|channel>thought\nThink after both calls.\n<channel|>"
2708 );
2709 assert!(
2710 rendered.contains(expected),
2711 "Gemma4 reasoning segments should stay adjacent to their tool calls, got: {rendered}"
2712 );
2713 assert!(!rendered.contains("<think>"));
2714 assert!(!rendered.contains("reasoning_content"));
2715 }
2716
2717 #[test]
2718 fn test_gemma4_template_renders_reasoning_content_without_tool_calls() {
2719 let formatter = make_gemma4_tool_formatter_for_tests();
2720 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
2721 "model": "gemma4-test",
2722 "messages": [
2723 {"role": "user", "content": "answer directly"},
2724 {
2725 "role": "assistant",
2726 "content": "Direct answer.",
2727 "reasoning_content": "Private thought."
2728 }
2729 ]
2730 }))
2731 .unwrap();
2732
2733 let rendered = formatter.render(&request).unwrap();
2734
2735 assert!(
2736 rendered.contains("<|channel>thought\nPrivate thought.\n<channel|>Direct answer."),
2737 "Gemma4 reasoning_content should render in the thought channel, got: {rendered}"
2738 );
2739 assert!(!rendered.contains("<think>"));
2740 assert!(!rendered.contains("reasoning_content"));
2741 }
2742
2743 #[test]
2749 fn test_reasoning_flag_is_per_template_not_global() {
2750 const PLAIN_DEFAULT: &str = "{{ bos_token }}{%- for message in messages -%}\
2753 {{ message['role'] }}: {{ message['content'] }}\n{%- endfor -%}";
2754
2755 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
2756 "chat_template": [
2757 {"default": PLAIN_DEFAULT},
2758 {"tool_use": gemma4_tool_template_for_tests()},
2759 ]
2760 }))
2761 .unwrap();
2762 let formatter =
2763 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
2764
2765 assert!(
2769 formatter.tool_use_template_handles_reasoning,
2770 "adapted gemma4 tool_use template should handle reasoning natively"
2771 );
2772 assert!(
2773 !formatter.default_template_handles_reasoning,
2774 "plain default template does not reference reasoning_content"
2775 );
2776
2777 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
2780 "model": "gemma4-test",
2781 "messages": [
2782 {"role": "user", "content": "answer directly"},
2783 {
2784 "role": "assistant",
2785 "content": "Direct answer.",
2786 "reasoning_content": "Private thought."
2787 }
2788 ]
2789 }))
2790 .unwrap();
2791
2792 let rendered = formatter.render(&request).unwrap();
2793 assert!(
2794 rendered.contains("<think>Private thought.</think>Direct answer."),
2795 "reasoning must be injected on the no-tool default path, got: {rendered}"
2796 );
2797 }
2798
2799 #[test]
2800 fn test_inject_reasoning_content_text_variant() {
2801 let mut messages = serde_json::json!([
2802 {
2803 "role": "assistant",
2804 "content": "The answer is 42.",
2805 "reasoning_content": "Let me think about this carefully."
2806 }
2807 ]);
2808
2809 inject_reasoning_content_into_messages(&mut messages);
2810
2811 let assistant = &messages[0];
2812 assert!(assistant.get("reasoning_content").is_none());
2813 let content = assistant["content"].as_str().unwrap();
2814 assert_eq!(
2815 content,
2816 "<think>Let me think about this carefully.</think>The answer is 42."
2817 );
2818 }
2819
2820 #[test]
2821 fn test_inject_reasoning_content_null_content() {
2822 let mut messages = serde_json::json!([
2824 {
2825 "role": "assistant",
2826 "content": null,
2827 "reasoning_content": "Thinking...",
2828 "tool_calls": [{"id": "call_0", "type": "function", "function": {"name": "f", "arguments": "{}"}}]
2829 }
2830 ]);
2831
2832 inject_reasoning_content_into_messages(&mut messages);
2833
2834 let content = messages[0]["content"].as_str().unwrap();
2835 assert_eq!(content, "<think>Thinking...</think>");
2836 assert!(messages[0].get("reasoning_content").is_none());
2837 }
2838
2839 #[test]
2840 fn test_inject_reasoning_content_skips_non_assistant() {
2841 let mut messages = serde_json::json!([
2842 {
2843 "role": "user",
2844 "content": "hello",
2845 "reasoning_content": "should not be touched"
2846 }
2847 ]);
2848
2849 inject_reasoning_content_into_messages(&mut messages);
2850
2851 assert!(messages[0].get("reasoning_content").is_some());
2853 }
2854
2855 fn make_test_formatter() -> HfTokenizerConfigJsonFormatter {
2857 use super::tokcfg::ChatTemplate;
2858 use super::{ContextMixins, HfTokenizerConfigJsonFormatter};
2859
2860 let template = r#"{%- for message in messages %}{{ message.role }}: {{ message.content }}
2863{%- endfor %}
2864{%- if add_generation_prompt %}assistant:{%- endif %}"#;
2865
2866 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
2867 "chat_template": template
2868 }))
2869 .unwrap();
2870
2871 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap()
2872 }
2873
2874 #[test]
2877 fn test_reasoning_content_text_roundtrip_render() {
2878 use super::OAIPromptFormatter;
2879 let formatter = make_test_formatter();
2880
2881 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
2882 "model": "test-model",
2883 "messages": [
2884 {"role": "user", "content": "What is sqrt(144)?"},
2885 {
2886 "role": "assistant",
2887 "content": "The answer is 12.",
2888 "reasoning_content": "I need to compute the square root of 144."
2889 },
2890 {"role": "user", "content": "Are you sure?"}
2891 ]
2892 }))
2893 .unwrap();
2894
2895 let rendered = formatter.render(&request).unwrap();
2896
2897 assert!(
2898 rendered.contains("<think>I need to compute the square root of 144.</think>"),
2899 "reasoning_content must appear as <think> block, got: {}",
2900 rendered
2901 );
2902 assert!(
2903 rendered.contains("The answer is 12."),
2904 "original content must be preserved"
2905 );
2906 assert!(
2907 !rendered.contains("reasoning_content"),
2908 "raw reasoning_content field should not leak into prompt"
2909 );
2910 }
2911
2912 #[test]
2916 fn test_reasoning_content_agentic_tool_call_roundtrip_render() {
2917 use super::OAIPromptFormatter;
2918 let formatter = make_test_formatter();
2919
2920 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
2921 "model": "test-model",
2922 "messages": [
2923 {"role": "user", "content": "What is sqrt(144) + sqrt(256)?"},
2924 {
2925 "role": "assistant",
2926 "content": null,
2927 "reasoning_content": "I need to compute both square roots. Let me start with sqrt(144).",
2928 "tool_calls": [{
2929 "id": "call_0",
2930 "type": "function",
2931 "function": {
2932 "name": "calculator",
2933 "arguments": "{\"expr\": \"sqrt(144)\"}"
2934 }
2935 }]
2936 },
2937 {
2938 "role": "tool",
2939 "tool_call_id": "call_0",
2940 "content": "12"
2941 },
2942 {
2943 "role": "assistant",
2944 "content": "sqrt(144) = 12 and sqrt(256) = 16, so the answer is 28.",
2945 "reasoning_content": "Got 12 for sqrt(144). Now sqrt(256) = 16. Sum is 28."
2946 },
2947 {"role": "user", "content": "Thanks!"}
2948 ]
2949 }))
2950 .unwrap();
2951
2952 let rendered = formatter.render(&request).unwrap();
2953
2954 assert!(
2956 rendered.contains("<think>I need to compute both square roots"),
2957 "first turn reasoning must be in prompt, got: {}",
2958 rendered
2959 );
2960 assert!(
2962 rendered.contains("<think>Got 12 for sqrt(144)"),
2963 "second turn reasoning must be in prompt"
2964 );
2965 assert!(
2966 rendered.contains("the answer is 28"),
2967 "final answer content must be preserved"
2968 );
2969 assert!(
2971 !rendered.contains("reasoning_content"),
2972 "raw reasoning_content field should not leak into prompt"
2973 );
2974 }
2975
2976 #[test]
2978 fn test_reasoning_injected_when_template_ignores_it() {
2979 use super::OAIPromptFormatter;
2980 let formatter = make_test_formatter();
2981
2982 assert!(!formatter.default_template_handles_reasoning);
2984 assert!(!formatter.tool_use_template_handles_reasoning);
2985
2986 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
2987 "model": "test-model",
2988 "messages": [
2989 {"role": "user", "content": "Hello"},
2990 {
2991 "role": "assistant",
2992 "content": "Hi.",
2993 "reasoning_content": "The user said hello."
2994 },
2995 {"role": "user", "content": "Bye"}
2996 ]
2997 }))
2998 .unwrap();
2999
3000 let rendered = formatter.render(&request).unwrap();
3001 assert!(
3002 rendered.contains("<think>The user said hello.</think>"),
3003 "injection must happen when template ignores reasoning_content, got: {}",
3004 rendered
3005 );
3006 }
3007
3008 #[test]
3010 fn test_reasoning_not_injected_when_template_handles_it() {
3011 use super::tokcfg::ChatTemplate;
3012 use super::{ContextMixins, HfTokenizerConfigJsonFormatter, OAIPromptFormatter};
3013
3014 let template = r#"{%- for message in messages %}{%- if message.role == "assistant" and message.reasoning_content is defined and message.reasoning_content %}<think>{{ message.reasoning_content }}</think>
3016{%- endif %}{{ message.role }}: {{ message.content }}
3017{%- endfor %}
3018{%- if add_generation_prompt %}assistant:{%- endif %}"#;
3019
3020 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
3021 "chat_template": template
3022 }))
3023 .unwrap();
3024
3025 let formatter =
3026 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
3027
3028 assert!(formatter.default_template_handles_reasoning);
3030 assert!(formatter.tool_use_template_handles_reasoning);
3031
3032 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
3033 "model": "test-model",
3034 "messages": [
3035 {"role": "user", "content": "Hello"},
3036 {
3037 "role": "assistant",
3038 "content": "Hi.",
3039 "reasoning_content": "The user said hello."
3040 },
3041 {"role": "user", "content": "Bye"}
3042 ]
3043 }))
3044 .unwrap();
3045
3046 let rendered = formatter.render(&request).unwrap();
3047
3048 assert!(
3050 rendered.contains("<think>The user said hello.</think>"),
3051 "template must render reasoning_content natively, got: {}",
3052 rendered
3053 );
3054 let think_count = rendered.matches("<think>").count();
3056 assert_eq!(
3057 think_count, 1,
3058 "must have exactly one <think> block (from template), got {} in: {}",
3059 think_count, rendered
3060 );
3061 }
3062
3063 const QWEN3_THINKING_TEMPLATE: &str = r##"{%- if tools %}
3067 {{- '<|im_start|>system\n' }}
3068 {%- if messages[0].role == 'system' %}
3069 {{- messages[0].content + '\n\n' }}
3070 {%- endif %}
3071 {{- "# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>" }}
3072 {%- for tool in tools %}
3073 {{- "\n" }}
3074 {{- tool | tojson }}
3075 {%- endfor %}
3076 {{- "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call><|im_end|>\n" }}
3077{%- else %}
3078 {%- if messages[0].role == 'system' %}
3079 {{- '<|im_start|>system\n' + messages[0].content + '<|im_end|>\n' }}
3080 {%- endif %}
3081{%- endif %}
3082{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
3083{%- for message in messages[::-1] %}
3084 {%- set index = (messages|length - 1) - loop.index0 %}
3085 {%- if ns.multi_step_tool and message.role == "user" and message.content is string and not(message.content.startswith('<tool_response>') and message.content.endswith('</tool_response>')) %}
3086 {%- set ns.multi_step_tool = false %}
3087 {%- set ns.last_query_index = index %}
3088 {%- endif %}
3089{%- endfor %}
3090{%- for message in messages %}
3091 {%- if message.content is string %}
3092 {%- set content = message.content %}
3093 {%- else %}
3094 {%- set content = '' %}
3095 {%- endif %}
3096 {%- if (message.role == "user") or (message.role == "system" and not loop.first) %}
3097 {{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
3098 {%- elif message.role == "assistant" %}
3099 {%- set reasoning_content = '' %}
3100 {%- if message.reasoning_content is string %}
3101 {%- set reasoning_content = message.reasoning_content %}
3102 {%- else %}
3103 {%- if '</think>' in content %}
3104 {%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
3105 {%- set content = content.split('</think>')[-1].lstrip('\n') %}
3106 {%- endif %}
3107 {%- endif %}
3108 {%- if loop.index0 > ns.last_query_index %}
3109 {%- if loop.last or (not loop.last and reasoning_content) %}
3110 {{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content.strip('\n') + '\n</think>\n\n' + content.lstrip('\n') }}
3111 {%- else %}
3112 {{- '<|im_start|>' + message.role + '\n' + content }}
3113 {%- endif %}
3114 {%- else %}
3115 {{- '<|im_start|>' + message.role + '\n' + content }}
3116 {%- endif %}
3117 {%- if message.tool_calls %}
3118 {%- for tool_call in message.tool_calls %}
3119 {%- if (loop.first and content) or (not loop.first) %}
3120 {{- '\n' }}
3121 {%- endif %}
3122 {%- if tool_call.function %}
3123 {%- set tool_call = tool_call.function %}
3124 {%- endif %}
3125 {{- '<tool_call>\n{"name": "' }}
3126 {{- tool_call.name }}
3127 {{- '", "arguments": ' }}
3128 {%- if tool_call.arguments is string %}
3129 {{- tool_call.arguments }}
3130 {%- else %}
3131 {{- tool_call.arguments | tojson }}
3132 {%- endif %}
3133 {{- '}\n</tool_call>' }}
3134 {%- endfor %}
3135 {%- endif %}
3136 {{- '<|im_end|>\n' }}
3137 {%- elif message.role == "tool" %}
3138 {%- if loop.first or (messages[loop.index0 - 1].role != "tool") %}
3139 {{- '<|im_start|>user' }}
3140 {%- endif %}
3141 {{- '\n<tool_response>\n' }}
3142 {{- content }}
3143 {{- '\n</tool_response>' }}
3144 {%- if loop.last or (messages[loop.index0 + 1].role != "tool") %}
3145 {{- '<|im_end|>\n' }}
3146 {%- endif %}
3147 {%- endif %}
3148{%- endfor %}
3149{%- if add_generation_prompt %}
3150 {{- '<|im_start|>assistant\n<think>\n' }}
3151{%- endif %}"##;
3152
3153 fn qwen3_thinking_formatter() -> HfTokenizerConfigJsonFormatter {
3154 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
3155 "chat_template": QWEN3_THINKING_TEMPLATE,
3156 }))
3157 .unwrap();
3158 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap()
3159 }
3160
3161 #[test]
3162 fn test_qwen3_thinking_template_flags_detected() {
3163 let formatter = qwen3_thinking_formatter();
3164 assert!(
3165 formatter.tool_use_template_handles_reasoning,
3166 "template references reasoning_content directly"
3167 );
3168 assert!(
3171 formatter.default_template_handles_tool_calls_arguments_string,
3172 "default template branches on `arguments is string`"
3173 );
3174 assert!(
3175 formatter.tool_use_template_handles_tool_calls_arguments_string,
3176 "tool_use template branches on `arguments is string`"
3177 );
3178 }
3179
3180 const QWEN38_REJECTS_STRING_ARGS_TEMPLATE: &str = r##"{%- for message in messages %}
3184 {%- if message.role == "assistant" and message.tool_calls %}
3185 {%- for tool_call in message.tool_calls %}
3186 {%- if tool_call.function %}
3187 {%- set tool_call = tool_call.function %}
3188 {%- endif %}
3189 {{- '<tool_call>\n<function=' + tool_call.name + '>\n' }}
3190 {%- if tool_call.arguments is mapping %}
3191 {%- for args_name, args_value in tool_call.arguments|items %}
3192 {{- '<parameter=' + args_name + '>\n' + args_value + '\n</parameter>\n' }}
3193 {%- endfor %}
3194 {%- elif tool_call.arguments is string %}
3195 {%- if tool_call.arguments|trim %}
3196 {{- raise_exception('Tool call arguments were passed as a JSON string.') }}
3197 {%- endif %}
3198 {%- endif %}
3199 {{- '</function>\n</tool_call>' }}
3200 {%- endfor %}
3201 {%- else %}
3202 {{- '<|im_start|>' + message.role + '\n' + message.content + '<|im_end|>\n' }}
3203 {%- endif %}
3204{%- endfor %}"##;
3205
3206 #[test]
3210 fn test_template_rejecting_string_arguments_gets_objects() {
3211 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
3212 "chat_template": QWEN38_REJECTS_STRING_ARGS_TEMPLATE,
3213 }))
3214 .unwrap();
3215 let formatter =
3216 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
3217 assert!(!formatter.default_template_handles_tool_calls_arguments_string);
3218 assert!(!formatter.tool_use_template_handles_tool_calls_arguments_string);
3219
3220 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
3221 "model": "qwen3.8",
3222 "messages": [
3223 {"role": "user", "content": "What's the weather in San Francisco?"},
3224 {"role": "assistant", "content": "", "tool_calls": [{
3225 "id": "call_sf",
3226 "type": "function",
3227 "function": {"name": "get_weather", "arguments": "{\"location\": \"San Francisco\"}"}
3228 }]},
3229 {"role": "tool", "tool_call_id": "call_sf", "content": "Foggy"}
3230 ],
3231 }))
3232 .unwrap();
3233 let rendered = formatter.render(&request).unwrap();
3234 assert!(
3235 rendered.contains("<parameter=location>\nSan Francisco\n</parameter>"),
3236 "{rendered}"
3237 );
3238 }
3239
3240 #[test]
3250 fn test_qwen3_thinking_append_only_across_tool_use_turn() {
3251 let formatter = qwen3_thinking_formatter();
3252
3253 let tools = serde_json::json!([{
3254 "type": "function",
3255 "function": {
3256 "name": "get_weather",
3257 "description": "Get the current weather for a location",
3258 "parameters": {
3259 "type": "object",
3260 "properties": {
3261 "location": {"type": "string"},
3262 "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}
3263 },
3264 "required": ["location"]
3265 }
3266 }
3267 }]);
3268
3269 let turn1_request: NvCreateChatCompletionRequest =
3271 serde_json::from_value(serde_json::json!({
3272 "model": "qwen3-thinking",
3273 "messages": [
3274 {"role": "system", "content": "You are a helpful assistant."},
3275 {"role": "user", "content": "What's the weather in San Francisco?"},
3276 ],
3277 "tools": tools,
3278 }))
3279 .unwrap();
3280 let p1 = formatter.render(&turn1_request).unwrap();
3281
3282 let model_emitted = "I'll call get_weather for SF.\n\
3286 </think>\n\n\
3287 <tool_call>\n\
3288 {\"name\": \"get_weather\", \"arguments\": {\"location\": \"San Francisco\", \"unit\": \"celsius\"}}\n\
3289 </tool_call><|im_end|>\n";
3290 let wire_after_t1 = format!("{p1}{model_emitted}");
3291
3292 let turn2_request: NvCreateChatCompletionRequest =
3295 serde_json::from_value(serde_json::json!({
3296 "model": "qwen3-thinking",
3297 "messages": [
3298 {"role": "system", "content": "You are a helpful assistant."},
3299 {"role": "user", "content": "What's the weather in San Francisco?"},
3300 {
3301 "role": "assistant",
3302 "content": "",
3303 "reasoning_content": "I'll call get_weather for SF.",
3304 "tool_calls": [{
3305 "id": "call_sf",
3306 "type": "function",
3307 "function": {
3308 "name": "get_weather",
3309 "arguments": "{\"location\": \"San Francisco\", \"unit\": \"celsius\"}"
3310 }
3311 }]
3312 },
3313 {
3314 "role": "tool",
3315 "tool_call_id": "call_sf",
3316 "content": "{\"temp\": 18, \"conditions\": \"Foggy\"}"
3317 }
3318 ],
3319 "tools": tools,
3320 }))
3321 .unwrap();
3322 let p2 = formatter.render(&turn2_request).unwrap();
3323
3324 if !p2.starts_with(&wire_after_t1) {
3325 let div = wire_after_t1
3327 .as_bytes()
3328 .iter()
3329 .zip(p2.as_bytes())
3330 .position(|(a, b)| a != b)
3331 .unwrap_or_else(|| wire_after_t1.len().min(p2.len()));
3332 let lo = div.saturating_sub(40);
3333 panic!(
3334 "turn-2 prompt is NOT a prefix-extension of [turn-1 + model bytes]\n \
3335 diverges at byte {div}\n \
3336 wire ends: ...{}|{}\n \
3337 t2 has: ...{}|{}",
3338 String::from_utf8_lossy(&wire_after_t1.as_bytes()[lo..div]),
3339 String::from_utf8_lossy(
3340 &wire_after_t1.as_bytes()[div..(div + 60).min(wire_after_t1.len())]
3341 ),
3342 String::from_utf8_lossy(&p2.as_bytes()[lo..div]),
3343 String::from_utf8_lossy(&p2.as_bytes()[div..(div + 60).min(p2.len())]),
3344 );
3345 }
3346
3347 let suffix = &p2[wire_after_t1.len()..];
3350 assert!(
3351 suffix.contains("<tool_response>"),
3352 "appended bytes must include the tool response, got: {suffix}"
3353 );
3354 assert!(
3355 suffix.ends_with("<|im_start|>assistant\n<think>\n"),
3356 "appended bytes must end with the next generation prompt, got: {suffix}"
3357 );
3358 }
3359}