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
319fn join_reasoning_content_segments_in_messages(messages: &mut serde_json::Value) {
325 let Some(msgs) = messages.as_array_mut() else {
326 return;
327 };
328
329 for msg in msgs.iter_mut() {
330 if msg.get("role").and_then(|r| r.as_str()) != Some("assistant") {
331 continue;
332 }
333 if let Some(reasoning) = msg.get_mut("reasoning_content")
334 && let Some(segments) = reasoning.as_array()
335 {
336 let joined = segments
337 .iter()
338 .filter_map(|s| s.as_str())
339 .filter(|s| !s.is_empty())
340 .collect::<Vec<_>>()
341 .join("\n");
342 *reasoning = joined.into();
343 }
344 }
345}
346
347impl OAIChatLikeRequest for dynamo_protocols::types::CreateChatCompletionRequest {
353 fn model(&self) -> String {
354 self.model.clone()
355 }
356
357 fn messages(&self) -> Value {
358 let messages_json = serde_json::to_value(&self.messages).unwrap();
359 Value::from_serialize(&messages_json)
360 }
361
362 fn typed_messages(&self) -> Option<&[dynamo_protocols::types::ChatCompletionRequestMessage]> {
363 Some(self.messages.as_slice())
364 }
365
366 fn tools(&self) -> Option<Value> {
367 if self.tools.is_none() {
368 None
369 } else {
370 Some(may_be_fix_tool_schema(
371 serde_json::to_value(&self.tools).unwrap(),
372 )?)
373 }
374 }
375
376 fn tool_choice(&self) -> Option<Value> {
377 if self.tool_choice.is_none() {
378 None
379 } else {
380 Some(Value::from_serialize(&self.tool_choice))
381 }
382 }
383
384 fn response_format(&self) -> Option<Value> {
385 self.response_format.as_ref().map(Value::from_serialize)
386 }
387
388 fn reasoning_effort(&self) -> Option<Value> {
389 self.reasoning_effort.as_ref().map(Value::from_serialize)
390 }
391
392 fn should_add_generation_prompt(&self) -> bool {
393 true
395 }
396
397 fn extract_text(&self) -> Option<TextInput> {
398 Some(TextInput::Single(String::new()))
399 }
400
401 fn mm_processor_kwargs(&self) -> Option<&serde_json::Value> {
402 self.mm_processor_kwargs.as_ref()
403 }
404}
405
406fn merge_message_content(
410 target: serde_json::Value,
411 source: serde_json::Value,
412) -> serde_json::Value {
413 use serde_json::Value;
414 let text_part = |text: String| json!({"type": "text", "text": text});
415 match (target, source) {
416 (Value::String(mut target), Value::String(source)) => {
417 if !target.is_empty() && !source.is_empty() {
418 target.push_str("\n\n");
419 }
420 target.push_str(&source);
421 Value::String(target)
422 }
423 (Value::Array(mut target), Value::Array(source)) => {
424 target.extend(source);
425 Value::Array(target)
426 }
427 (Value::Array(mut target), Value::String(source)) => {
428 if !source.is_empty() {
429 target.push(text_part(source));
430 }
431 Value::Array(target)
432 }
433 (Value::String(target), Value::Array(source)) => {
434 let mut parts = Vec::with_capacity(source.len() + 1);
435 if !target.is_empty() {
436 parts.push(text_part(target));
437 }
438 parts.extend(source);
439 Value::Array(parts)
440 }
441 (Value::Null, source) => source,
442 (target, _) => target,
444 }
445}
446
447fn append_message_content(target: &mut serde_json::Value, source: serde_json::Value) {
449 let Some(target) = target.as_object_mut() else {
450 return;
451 };
452 let merged = merge_message_content(
453 target.remove("content").unwrap_or(serde_json::Value::Null),
454 source,
455 );
456 target.insert("content".to_string(), merged);
457}
458
459fn take_message_content(message: &mut serde_json::Value) -> serde_json::Value {
460 message
461 .get_mut("content")
462 .map(serde_json::Value::take)
463 .unwrap_or(serde_json::Value::Null)
464}
465
466fn normalize_system_messages(messages: &mut serde_json::Value, rules: SystemNormalization) {
470 let serde_json::Value::Array(list) = messages else {
471 return;
472 };
473 let role_is =
474 |m: &serde_json::Value, r: &str| m.get("role").and_then(|v| v.as_str()) == Some(r);
475
476 if rules.demote_nonleading_system {
477 let leading = list.iter().take_while(|m| role_is(m, "system")).count();
480 if leading > 1 {
481 for mut trailing in list.drain(1..leading).collect::<Vec<_>>() {
482 let content = take_message_content(&mut trailing);
483 append_message_content(&mut list[0], content);
484 }
485 }
486
487 let leading = list.iter().take_while(|m| role_is(m, "system")).count();
490 for m in list.iter_mut().skip(leading) {
491 if role_is(m, "system")
492 && let Some(m) = m.as_object_mut()
493 {
494 m.insert("role".to_string(), json!("user"));
495 }
496 }
497 }
498
499 if rules.coalesce_consecutive_users {
500 let mut coalesced: Vec<serde_json::Value> = Vec::with_capacity(list.len());
501 for mut m in list.drain(..) {
502 if role_is(&m, "user") && coalesced.last().is_some_and(|p| role_is(p, "user")) {
503 let content = take_message_content(&mut m);
504 append_message_content(coalesced.last_mut().unwrap(), content);
505 } else {
506 coalesced.push(m);
507 }
508 }
509 *list = coalesced;
510 }
511}
512
513impl OAIPromptFormatter for HfTokenizerConfigJsonFormatter {
514 fn supports_add_generation_prompt(&self) -> bool {
515 self.supports_add_generation_prompt
516 }
517
518 fn render(&self, req: &dyn OAIChatLikeRequest) -> Result<String> {
519 let mixins = Value::from_dyn_object(self.mixins.clone());
520
521 let tools = req.tools();
522 let tools = if self.exclude_tools_when_tool_choice_none {
525 match req.tool_choice() {
526 Some(ref tc) if tc.as_str() == Some("none") => None,
527 _ => tools,
528 }
529 } else {
530 tools
531 };
532 let has_tools = tools.as_ref().and_then(|v| v.len()).is_some_and(|l| l > 0);
534 let add_generation_prompt = req.should_add_generation_prompt();
535
536 tracing::trace!(
537 "Rendering prompt with tools: {:?}, add_generation_prompt: {}",
538 has_tools,
539 add_generation_prompt
540 );
541
542 let (
545 template_name,
546 template_handles_tool_calls_args_string,
547 template_handles_reasoning,
548 template_requires_reasoning_string,
549 system_normalization,
550 ) = if has_tools {
551 (
552 "tool_use",
553 self.tool_use_template_handles_tool_calls_arguments_string,
554 self.tool_use_template_handles_reasoning,
555 self.tool_use_template_requires_reasoning_string,
556 self.tool_use_system_normalization,
557 )
558 } else {
559 (
560 "default",
561 self.default_template_handles_tool_calls_arguments_string,
562 self.default_template_handles_reasoning,
563 self.default_template_requires_reasoning_string,
564 self.default_system_normalization,
565 )
566 };
567
568 let mut messages_for_template = crate::messages_to_json(req)?;
569
570 crate::reject_unsupported_partial_assistant(&messages_for_template)?;
571 crate::reject_unsupported_message_tools(&messages_for_template, &[])?;
572
573 if system_normalization.is_required() {
574 normalize_system_messages(&mut messages_for_template, system_normalization);
575 }
576
577 messages_for_template = may_be_fix_msg_content(
578 messages_for_template,
579 self.requires_content_arrays,
580 self.image_placeholder_template,
581 );
582
583 if !template_handles_tool_calls_args_string {
590 normalize_tool_calls_arguments_in_messages(&mut messages_for_template);
591 }
592 normalize_function_call_arguments_in_messages(&mut messages_for_template);
596
597 if !template_handles_reasoning {
602 inject_reasoning_content_into_messages(&mut messages_for_template);
603 } else if template_requires_reasoning_string {
604 join_reasoning_content_segments_in_messages(&mut messages_for_template);
605 }
606
607 let ctx = context! {
608 messages => messages_for_template,
609 tools => tools,
610 bos_token => self.config.bos_tok(),
611 eos_token => self.config.eos_tok(),
612 unk_token => self.config.unk_tok(),
613 add_generation_prompt => add_generation_prompt,
614 ..mixins
615 };
616
617 let ctx = if let Some(args) = req.chat_template_args() {
619 let extra = Value::from_serialize(args);
620 context! { ..ctx, ..extra }
621 } else {
622 ctx
623 };
624
625 let tmpl: minijinja::Template<'_, '_> = self.env.get_template(template_name)?;
626 Ok(tmpl.render(&ctx)?)
627 }
628}
629
630#[cfg(test)]
631mod tests {
632 use super::*;
633
634 use dynamo_protocols::types::ChatCompletionRequestMessage as Msg;
635 use dynamo_protocols::types::CreateChatCompletionRequest as NvCreateChatCompletionRequest;
638 use minijinja::{Environment, context};
639
640 use super::super::tokcfg::ChatTemplate as SysChatTemplate;
643 use super::super::{
644 ContextMixins as SysMixins, HfTokenizerConfigJsonFormatter as SysFormatter,
645 };
646
647 fn formatter_for(template: &str) -> SysFormatter {
648 let ct: SysChatTemplate = serde_json::from_value(json!({
650 "chat_template": template,
651 "bos_token": "<s>",
652 "eos_token": "</s>",
653 "unk_token": "<unk>",
654 }))
655 .unwrap();
656 SysFormatter::new(ct, SysMixins::new(&[])).unwrap()
657 }
658
659 fn formatter_for_templates(default: &str, tool_use: &str) -> SysFormatter {
660 let ct: SysChatTemplate = serde_json::from_value(json!({
661 "chat_template": [
662 {"default": default},
663 {"tool_use": tool_use},
664 ],
665 "bos_token": "<s>",
666 "eos_token": "</s>",
667 "unk_token": "<unk>",
668 }))
669 .unwrap();
670 SysFormatter::new(ct, SysMixins::new(&[])).unwrap()
671 }
672
673 fn try_formatter_for(template: &str) -> Option<SysFormatter> {
674 let ct: SysChatTemplate = serde_json::from_value(json!({
675 "chat_template": template,
676 "bos_token": "<s>",
677 "eos_token": "</s>",
678 "unk_token": "<unk>",
679 }))
680 .ok()?;
681 SysFormatter::new(ct, SysMixins::new(&[])).ok()
682 }
683
684 fn render_shape(f: &SysFormatter, messages: serde_json::Value) -> Result<String> {
685 let req: NvCreateChatCompletionRequest =
686 serde_json::from_value(json!({ "model": "test", "messages": messages })).unwrap();
687 f.render(&req)
688 }
689
690 fn render_shape_with_tools(f: &SysFormatter, messages: serde_json::Value) -> Result<String> {
691 let req: NvCreateChatCompletionRequest = serde_json::from_value(json!({
692 "model": "test",
693 "messages": messages,
694 "tools": [{
695 "type": "function",
696 "function": {"name": "noop", "parameters": {}}
697 }]
698 }))
699 .unwrap();
700 f.render(&req)
701 }
702
703 struct RawMessagesRequest(Value);
704
705 impl OAIChatLikeRequest for RawMessagesRequest {
706 fn model(&self) -> String {
707 "test".to_string()
708 }
709
710 fn messages(&self) -> Value {
711 self.0.clone()
712 }
713
714 fn should_add_generation_prompt(&self) -> bool {
715 true
716 }
717 }
718
719 fn render_raw_shape(f: &SysFormatter, messages: serde_json::Value) -> Result<String> {
720 f.render(&RawMessagesRequest(Value::from_serialize(&messages)))
721 }
722
723 #[test]
724 fn content_normalization_preserves_unchanged_values() {
725 let messages = json!([
726 {"role": "user", "content": " 中文 <special>\n", "name": "user"},
727 {"role": "assistant", "content": null},
728 {"role": "user", "content": []},
729 {"role": "assistant", "tool_calls": []}
730 ]);
731 assert_eq!(
732 may_be_fix_msg_content(messages.clone(), false, Some("")),
733 messages
734 );
735 }
736
737 const PERMISSIVE_TMPL: &str = concat!(
738 "{%- for m in messages -%}",
739 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
740 "{%- endfor -%}"
741 );
742
743 #[test]
744 fn jinja_templates_reject_message_level_tools() {
745 let f = formatter_for(PERMISSIVE_TMPL);
746 let error = render_shape(
747 &f,
748 json!([
749 {"role": "system", "tools": [{"name": "lookup", "parameters": {"type": "object"}}]},
750 {"role": "user", "content": "hi"}
751 ]),
752 )
753 .unwrap_err();
754 assert!(matches!(
755 error.downcast_ref::<crate::PromptRenderError>(),
756 Some(crate::PromptRenderError::InvalidRequest(message))
757 if message.contains("message-level `tools`")
758 ));
759
760 let error = render_raw_shape(
761 &f,
762 json!([{
763 "role": "user",
764 "content": "hi",
765 "tools": [{"name": "lookup", "parameters": {"type": "object"}}]
766 }]),
767 )
768 .unwrap_err();
769 assert!(matches!(
770 error.downcast_ref::<crate::PromptRenderError>(),
771 Some(crate::PromptRenderError::InvalidRequest(message))
772 if message.contains("message-level `tools`")
773 ));
774
775 let rendered = render_shape(
776 &f,
777 json!([
778 {"role": "system", "content": "You are helpful.", "tools": []},
779 {"role": "user", "content": "hi"}
780 ]),
781 )
782 .unwrap();
783 assert!(rendered.contains("<|im_start|>system\nYou are helpful.<|im_end|>"));
784 }
785
786 #[test]
787 fn jinja_templates_reject_unsupported_partial_assistant() {
788 let f = formatter_for(PERMISSIVE_TMPL);
789 let error = render_shape(
790 &f,
791 json!([
792 {"role": "user", "content": "Continue"},
793 {"role": "assistant", "content": "prefix", "partial": true}
794 ]),
795 )
796 .unwrap_err();
797 assert!(matches!(
798 error.downcast_ref::<crate::PromptRenderError>(),
799 Some(crate::PromptRenderError::InvalidRequest(message))
800 if message.contains("`partial: true` is not supported")
801 ));
802
803 let rendered = render_shape(
804 &f,
805 json!([{"role": "assistant", "content": "ordinary", "partial": false}]),
806 )
807 .unwrap();
808 assert!(rendered.contains("ordinary"));
809 }
810 const STRICT_LEADING_TMPL: &str = concat!(
812 "{%- for m in messages -%}",
813 "{%- if m.role == 'system' and not loop.first -%}",
814 "{{ raise_exception('System message must be at the beginning.') }}",
815 "{%- endif -%}",
816 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
817 "{%- endfor -%}"
818 );
819 const ALTERNATION_TMPL: &str = concat!(
821 "{%- set ns = namespace(prev='') -%}",
822 "{%- for m in messages -%}",
823 "{%- if m.role == 'user' and ns.prev == 'user' -%}",
824 "{{ raise_exception('Conversation roles must alternate.') }}",
825 "{%- endif -%}",
826 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
827 "{%- set ns.prev = m.role -%}",
828 "{%- endfor -%}"
829 );
830 const STRICT_BOTH_TMPL: &str = concat!(
832 "{%- set ns = namespace(prev='') -%}",
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 "{%- if m.role == 'user' and ns.prev == 'user' -%}",
838 "{{ raise_exception('Conversation roles must alternate.') }}",
839 "{%- endif -%}",
840 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
841 "{%- set ns.prev = m.role -%}",
842 "{%- endfor -%}"
843 );
844 const DEFAULT_NONE_GATED_TMPL: &str = concat!(
847 "{%- set strict = tools is not none -%}",
848 "{%- for m in messages -%}",
849 "{%- if strict and m.role == 'system' and not loop.first -%}",
850 "{{ raise_exception('System message must be at the beginning.') }}",
851 "{%- endif -%}",
852 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
853 "{%- endfor -%}"
854 );
855 const TOOL_NONEMPTY_GATED_TMPL: &str = concat!(
856 "{%- set strict = tools|length > 0 -%}",
857 "{%- for m in messages -%}",
858 "{%- if strict and m.role == 'system' and not loop.first -%}",
859 "{{ raise_exception('System message must be at the beginning.') }}",
860 "{%- endif -%}",
861 "<|im_start|>{{ m.role }}\n{{ m.content }}<|im_end|>\n",
862 "{%- endfor -%}"
863 );
864 const STRICT_ARRAY_TMPL: &str = concat!(
866 "{%- for m in messages -%}",
867 "{%- if m.role == 'system' and not loop.first -%}",
868 "{{ raise_exception('System message must be at the beginning.') }}",
869 "{%- endif -%}",
870 "<|im_start|>{{ m.role }}\n",
871 "{%- if m.content is not string -%}",
872 "{%- for part in m.content -%}{{ part.text }}{%- endfor -%}",
873 "{%- endif -%}",
874 "<|im_end|>\n",
875 "{%- endfor -%}"
876 );
877
878 fn claude_shape() -> serde_json::Value {
880 json!([
881 {"role": "system", "content": "You are Claude Code."},
882 {"role": "user", "content": "hello"},
883 {"role": "system", "content": "mid-conversation reminder"},
884 ])
885 }
886
887 fn all_restrictions() -> SystemNormalization {
888 SystemNormalization {
889 demote_nonleading_system: true,
890 coalesce_consecutive_users: true,
891 }
892 }
893
894 #[test]
895 fn permissive_template_is_not_flagged_and_renders_untouched() {
896 let f = formatter_for(PERMISSIVE_TMPL);
897 assert!(!f.default_system_normalization.is_required());
898 assert!(!f.tool_use_system_normalization.is_required());
899 let out = render_shape(&f, claude_shape()).unwrap();
900 assert!(out.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
901 }
902
903 #[test]
904 fn strict_leading_template_demotes_mid_system_but_keeps_user_turns_apart() {
905 let f = formatter_for(STRICT_LEADING_TMPL);
906 assert!(f.default_system_normalization.demote_nonleading_system);
907 assert!(!f.default_system_normalization.coalesce_consecutive_users);
909
910 let out = render_shape(&f, claude_shape()).unwrap();
912 assert_eq!(out.matches("<|im_start|>system").count(), 1);
913 assert!(out.contains("<|im_start|>user\nhello<|im_end|>"));
914 assert!(out.contains("<|im_start|>user\nmid-conversation reminder<|im_end|>"));
915 }
916
917 #[test]
918 fn alternation_template_coalesces_users_but_keeps_mid_system() {
919 let f = formatter_for(ALTERNATION_TMPL);
920 assert!(f.default_system_normalization.coalesce_consecutive_users);
921 assert!(!f.default_system_normalization.demote_nonleading_system);
923
924 let out = render_shape(&f, claude_shape()).unwrap();
925 assert!(out.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
926
927 let out = render_shape(
928 &f,
929 json!([
930 {"role": "system", "content": "s"},
931 {"role": "user", "content": "hello"},
932 {"role": "user", "content": "again"},
933 ]),
934 )
935 .unwrap();
936 assert_eq!(out.matches("<|im_start|>user").count(), 1);
937 assert!(out.contains("<|im_start|>user\nhello\n\nagain<|im_end|>"));
938 }
939
940 #[test]
941 fn strict_both_template_demotes_then_coalesces() {
942 let f = formatter_for(STRICT_BOTH_TMPL);
943 assert!(f.default_system_normalization.demote_nonleading_system);
944 assert!(f.default_system_normalization.coalesce_consecutive_users);
945
946 let out = render_shape(&f, claude_shape()).unwrap();
947 assert_eq!(out.matches("<|im_start|>system").count(), 1);
948 assert!(out.contains("<|im_start|>user\nhello\n\nmid-conversation reminder<|im_end|>"));
949 }
950
951 #[test]
952 fn system_normalization_flag_is_selected_per_template() {
953 let f = formatter_for_templates(PERMISSIVE_TMPL, STRICT_LEADING_TMPL);
954 assert!(!f.default_system_normalization.is_required());
955 assert!(f.tool_use_system_normalization.is_required());
956
957 let no_tools = render_shape(&f, claude_shape()).unwrap();
958 assert!(no_tools.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
959 let with_tools = render_shape_with_tools(&f, claude_shape()).unwrap();
960 assert_eq!(with_tools.matches("<|im_start|>system").count(), 1);
961 assert!(with_tools.contains("<|im_start|>user\nmid-conversation reminder<|im_end|>"));
962
963 let f = formatter_for_templates(STRICT_LEADING_TMPL, PERMISSIVE_TMPL);
964 assert!(f.default_system_normalization.is_required());
965 assert!(!f.tool_use_system_normalization.is_required());
966 let with_tools = render_shape_with_tools(&f, claude_shape()).unwrap();
967 assert!(with_tools.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
968 }
969
970 #[test]
971 fn system_normalization_probe_uses_runtime_tools_shape() {
972 let f = formatter_for_templates(DEFAULT_NONE_GATED_TMPL, TOOL_NONEMPTY_GATED_TMPL);
973 assert!(!f.default_system_normalization.is_required());
974 assert!(f.tool_use_system_normalization.is_required());
975
976 let no_tools = render_shape(&f, claude_shape()).unwrap();
977 assert!(no_tools.contains("<|im_start|>system\nmid-conversation reminder<|im_end|>"));
978
979 let with_tools = render_shape_with_tools(&f, claude_shape()).unwrap();
980 assert_eq!(with_tools.matches("<|im_start|>system").count(), 1);
981 assert!(with_tools.contains("<|im_start|>user\nmid-conversation reminder<|im_end|>"));
982 }
983
984 #[test]
985 fn system_normalization_precedes_required_content_array_conversion() {
986 let f = formatter_for(STRICT_ARRAY_TMPL);
987 assert!(f.requires_content_arrays);
988 assert!(f.default_system_normalization.demote_nonleading_system);
989
990 let out = render_shape(
991 &f,
992 json!([
993 {"role": "system", "content": "A"},
994 {"role": "system", "content": "B"},
995 {"role": "user", "content": "hello"},
996 ]),
997 )
998 .unwrap();
999 assert!(out.contains("A\n\nB"));
1000 }
1001
1002 #[test]
1003 fn normalize_preserves_multimodal_user_content_and_fields() {
1004 let mut m = json!([
1005 {
1006 "role": "user",
1007 "name": "kept",
1008 "content": [
1009 {"type": "text", "text": "look"},
1010 {"type": "image"},
1011 ],
1012 },
1013 {"role": "system", "content": "remember"},
1014 ]);
1015 normalize_system_messages(&mut m, all_restrictions());
1016 assert_eq!(
1017 m,
1018 json!([{
1019 "role": "user",
1020 "name": "kept",
1021 "content": [
1022 {"type": "text", "text": "look"},
1023 {"type": "image"},
1024 {"type": "text", "text": "remember"},
1025 ],
1026 }])
1027 );
1028 }
1029
1030 #[test]
1033 fn coalesce_preserves_multimodal_content_of_the_merged_turn() {
1034 let mut m = json!([
1035 {"role": "user", "content": "look"},
1036 {"role": "user", "content": [
1037 {"type": "text", "text": "at this"},
1038 {"type": "image_url", "image_url": {"url": "http://img"}},
1039 ]},
1040 ]);
1041 normalize_system_messages(&mut m, all_restrictions());
1042 assert_eq!(
1043 m,
1044 json!([{
1045 "role": "user",
1046 "content": [
1047 {"type": "text", "text": "look"},
1048 {"type": "text", "text": "at this"},
1049 {"type": "image_url", "image_url": {"url": "http://img"}},
1050 ],
1051 }])
1052 );
1053 }
1054
1055 #[test]
1056 fn normalize_merges_leading_run_and_coalesces() {
1057 let mut m = json!([
1058 {"role": "system", "content": "A"},
1059 {"role": "system", "content": "B"},
1060 {"role": "user", "content": "hi"},
1061 {"role": "system", "content": "reminder"},
1062 ]);
1063 normalize_system_messages(&mut m, all_restrictions());
1064 assert_eq!(
1065 m,
1066 json!([
1067 {"role": "system", "content": "A\n\nB"},
1068 {"role": "user", "content": "hi\n\nreminder"},
1069 ])
1070 );
1071 }
1072
1073 #[test]
1076 fn each_restriction_applies_only_its_own_rewrite() {
1077 let shape = json!([
1078 {"role": "system", "content": "A"},
1079 {"role": "system", "content": "B"},
1080 {"role": "user", "content": "hi"},
1081 {"role": "system", "content": "reminder"},
1082 ]);
1083
1084 let mut demote_only = shape.clone();
1085 normalize_system_messages(
1086 &mut demote_only,
1087 SystemNormalization {
1088 demote_nonleading_system: true,
1089 coalesce_consecutive_users: false,
1090 },
1091 );
1092 assert_eq!(
1093 demote_only,
1094 json!([
1095 {"role": "system", "content": "A\n\nB"},
1096 {"role": "user", "content": "hi"},
1097 {"role": "user", "content": "reminder"},
1098 ])
1099 );
1100
1101 let mut coalesce_only = shape.clone();
1102 normalize_system_messages(
1103 &mut coalesce_only,
1104 SystemNormalization {
1105 demote_nonleading_system: false,
1106 coalesce_consecutive_users: true,
1107 },
1108 );
1109 assert_eq!(coalesce_only, shape);
1110 }
1111
1112 #[test]
1113 fn normalize_preserves_array_system_content() {
1114 let mut m = json!([
1115 {"role": "user", "content": "hi"},
1116 {"role": "system", "content": [{"type": "text", "text": "one"},
1117 {"type": "text", "text": "two"}]},
1118 ]);
1119 normalize_system_messages(&mut m, all_restrictions());
1120 assert_eq!(
1121 m,
1122 json!([{"role": "user", "content": [
1123 {"type": "text", "text": "hi"},
1124 {"type": "text", "text": "one"},
1125 {"type": "text", "text": "two"},
1126 ]}])
1127 );
1128 }
1129
1130 #[test]
1139 #[ignore]
1140 fn adaptive_system_corpus_audit() {
1141 let dir =
1142 std::env::var("TEMPLATE_CORPUS").expect("set TEMPLATE_CORPUS to the templates dir");
1143 let manifest: serde_json::Value =
1144 serde_json::from_str(&std::fs::read_to_string(format!("{dir}/manifest.json")).unwrap())
1145 .unwrap();
1146
1147 let sys = |c: &str| json!({"role": "system", "content": c});
1150 let usr = |c: &str| json!({"role": "user", "content": c});
1151 let asst = |c: &str| json!({"role": "assistant", "content": c});
1152 let shapes: Vec<(&str, serde_json::Value)> = vec![
1153 ("turn1", json!([sys("s"), usr("u"), sys("mid")])),
1154 (
1155 "multiturn",
1156 json!([sys("s"), usr("u"), sys("mid"), asst("a"), usr("u2")]),
1157 ),
1158 (
1159 "mid_after_asst",
1160 json!([sys("s"), usr("u"), asst("a"), sys("mid"), usr("u2")]),
1161 ),
1162 ("double_leading", json!([sys("s0"), sys("s1"), usr("u")])),
1163 ("consec_user", json!([sys("s"), usr("u0"), usr("u1")])),
1164 (
1165 "tail_reminder",
1166 json!([
1167 sys("s"),
1168 usr("u"),
1169 asst("a"),
1170 usr("u2"),
1171 sys("mid"),
1172 usr("u3")
1173 ]),
1174 ),
1175 ("leading_only_baseline", json!([sys("s"), usr("u")])),
1176 ];
1177
1178 let mut total = 0usize;
1179 let mut flagged = 0usize;
1180 let mut demote_only = 0usize;
1181 let mut coalesce = 0usize;
1182 let mut failures: Vec<String> = Vec::new();
1183 for (file, meta) in manifest.as_object().unwrap() {
1184 let tmpl = std::fs::read_to_string(format!("{dir}/{file}.jinja")).unwrap();
1185 let model = meta["model"].as_str().unwrap_or(file);
1186 let f = match try_formatter_for(&tmpl) {
1189 Some(f) => f,
1190 None => {
1191 eprintln!("[skip-compile] {model}");
1192 continue;
1193 }
1194 };
1195 if render_shape(&f, json!([sys("s"), usr("u")])).is_err() {
1198 eprintln!("[skip-baseline] {model}");
1199 continue;
1200 }
1201 total += 1;
1202 let rules = f.default_system_normalization;
1203 let flag = rules.is_required();
1204 if flag {
1205 flagged += 1;
1206 }
1207 if rules.demote_nonleading_system {
1208 demote_only += usize::from(!rules.coalesce_consecutive_users);
1209 }
1210 if rules.coalesce_consecutive_users {
1211 coalesce += 1;
1212 }
1213 for (name, shape) in &shapes {
1214 if render_shape(&f, shape.clone()).is_err() {
1215 failures.push(format!("{model} | shape={name} | flag={flag}"));
1216 }
1217 }
1218 if flag {
1219 eprintln!(
1220 "[ok] demote={} coalesce={} {model}",
1221 rules.demote_nonleading_system, rules.coalesce_consecutive_users
1222 );
1223 }
1224 }
1225 eprintln!(
1226 "\naudited {total} templates ({flagged} flagged: {demote_only} demote-only, \
1227 {coalesce} coalescing); {} shape failures",
1228 failures.len()
1229 );
1230 for f in &failures {
1231 eprintln!(" FAIL {f}");
1232 }
1233 assert!(
1234 failures.is_empty(),
1235 "{} template/shape combinations did not render (probe insufficient or normalization insufficient)",
1236 failures.len()
1237 );
1238 }
1239
1240 #[test]
1248 fn test_render_long_conversation_does_not_overflow_stack() {
1249 let handle = std::thread::Builder::new()
1250 .stack_size(2 * 1024 * 1024)
1251 .spawn(|| {
1252 let template_string = concat!(
1253 "{%- set ns = namespace(items=[]) -%}",
1254 "{%- for m in messages -%}",
1255 "{%- set ns.items = ns.items + [m] -%}",
1256 "{%- endfor -%}",
1257 "COUNT={{ ns.items | length }}"
1258 );
1259 let chat_template: ChatTemplate =
1260 serde_json::from_value(serde_json::json!({ "chat_template": template_string }))
1261 .unwrap();
1262 let formatter =
1263 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[]))
1264 .unwrap();
1265
1266 let n = 3000;
1267 let messages: Vec<serde_json::Value> = (0..n)
1268 .map(|i| serde_json::json!({"role": "user", "content": format!("turn {i}")}))
1269 .collect();
1270 let request: NvCreateChatCompletionRequest =
1271 serde_json::from_value(serde_json::json!({
1272 "model": "test",
1273 "messages": messages,
1274 }))
1275 .unwrap();
1276
1277 let rendered = formatter.render(&request).unwrap();
1279 assert_eq!(rendered.trim(), format!("COUNT={n}"));
1280 })
1281 .unwrap();
1282 handle.join().unwrap();
1283 }
1284
1285 #[test]
1296 #[ignore]
1297 fn dump_gptoss_tool_prompt() {
1298 use super::tokcfg::ChatTemplate;
1299 use super::{ContextMixins, HfTokenizerConfigJsonFormatter};
1300
1301 let path = std::env::var("GPTOSS_CHAT_TEMPLATE").expect(
1302 "set GPTOSS_CHAT_TEMPLATE to the tokenizer_config.json, chat_template.jinja, or model dir path",
1303 );
1304 let input_path = std::path::Path::new(&path);
1305 let file_path = if input_path.is_dir() {
1306 input_path.join("tokenizer_config.json")
1308 } else {
1309 input_path.to_path_buf()
1310 };
1311 let raw = std::fs::read_to_string(&file_path).expect("read chat template file");
1312 let template_string: String = match serde_json::from_str::<serde_json::Value>(&raw) {
1319 Ok(v) if v.get("chat_template").is_some() => v["chat_template"]
1320 .as_str()
1321 .expect("chat_template field must be a string")
1322 .to_string(),
1323 _ => {
1324 let sibling = std::path::Path::new(&path)
1325 .parent()
1326 .map(|d| d.join("chat_template.jinja"));
1327 match sibling {
1328 Some(p) if p.exists() => {
1329 eprintln!(
1330 "[info] {path} had no chat_template field; using {}",
1331 p.display()
1332 );
1333 std::fs::read_to_string(&p).expect("read sibling chat_template.jinja")
1334 }
1335 _ => raw,
1336 }
1337 }
1338 };
1339
1340 assert!(
1343 template_string.contains("{%") || template_string.contains("{{"),
1344 "resolved template has no Jinja tags — GPTOSS_CHAT_TEMPLATE ({path}) is probably \
1345 tokenizer_config.json with no chat_template field and no sibling chat_template.jinja. \
1346 Point it at the chat_template.jinja file."
1347 );
1348
1349 let chat_template: ChatTemplate =
1350 serde_json::from_value(serde_json::json!({ "chat_template": template_string }))
1351 .unwrap();
1352
1353 let formatter =
1354 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
1355
1356 let request: NvCreateChatCompletionRequest = serde_json::from_str(
1358 r#"{
1359 "model": "openai/gpt-oss-120b",
1360 "messages": [{"role":"user","content":"Search the repo for the string \"countHook\"."}],
1361 "tools": [
1362 {"type":"function","function":{"name":"grep","description":"search files","parameters":{"type":"object","properties":{"pattern":{"type":"string"},"path":{"type":"string"}},"required":["pattern"]}}},
1363 {"type":"function","function":{"name":"read","description":"read a file","parameters":{"type":"object","properties":{"filePath":{"type":"string"}},"required":["filePath"]}}}
1364 ]
1365 }"#,
1366 )
1367 .unwrap();
1368
1369 let rendered = formatter.render(&request).unwrap();
1370 eprintln!("================ RENDERED gpt-oss PROMPT (tools declared) ================");
1371 eprintln!("{rendered}");
1372 eprintln!("================ END RENDERED PROMPT ================");
1373 eprintln!("[diagnostics] does the rendered prompt contain…");
1374 for needle in [
1375 "commentary",
1376 "Calls to these tools",
1377 "functions",
1378 "# Tools",
1379 "<|channel|>",
1380 "constrain",
1381 "analysis",
1382 ] {
1383 eprintln!(
1384 " {:>22}: {}",
1385 format!("{needle:?}"),
1386 rendered.contains(needle)
1387 );
1388 }
1389 }
1390
1391 #[test]
1393 fn test_convert_media_url_to_placeholder_single_type() {
1394 let mut content_array = vec![
1395 serde_json::json!({"type": "text", "text": "Check this image:"}),
1396 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1397 serde_json::json!({"type": "text", "text": "What do you see?"}),
1398 ];
1399
1400 let conversions = &[("image_url", "image")];
1401 convert_media_url_to_placeholder(&mut content_array, conversions);
1402
1403 assert_eq!(content_array.len(), 3);
1404 assert_eq!(content_array[0]["type"], "text");
1406 assert_eq!(content_array[0]["text"], "Check this image:");
1407 assert_eq!(content_array[1]["type"], "image");
1409 assert!(content_array[1].get("image_url").is_none());
1410 assert_eq!(content_array[2]["type"], "text");
1412 assert_eq!(content_array[2]["text"], "What do you see?");
1413 }
1414
1415 #[test]
1417 fn test_convert_media_url_to_placeholder_multiple_same_type() {
1418 let mut content_array = vec![
1419 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image1.jpg"}}),
1420 serde_json::json!({"type": "text", "text": "vs"}),
1421 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image2.jpg"}}),
1422 ];
1423
1424 let conversions = &[("image_url", "image")];
1425 convert_media_url_to_placeholder(&mut content_array, conversions);
1426
1427 assert_eq!(content_array.len(), 3);
1428 assert_eq!(content_array[0]["type"], "image");
1429 assert_eq!(content_array[1]["type"], "text");
1430 assert_eq!(content_array[2]["type"], "image");
1431 }
1432
1433 #[test]
1435 fn test_convert_media_url_to_placeholder_selective_conversion() {
1436 let mut content_array = vec![
1437 serde_json::json!({"type": "audio_url", "audio_url": {"url": "https://example.com/audio.mp3"}}),
1438 serde_json::json!({"type": "video_url", "video_url": {"url": "https://example.com/video.mp4"}}),
1439 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1440 ];
1441
1442 let conversions = &[("image_url", "image")];
1444 convert_media_url_to_placeholder(&mut content_array, conversions);
1445
1446 assert_eq!(content_array.len(), 3);
1447 assert_eq!(content_array[0]["type"], "audio_url");
1449 assert!(content_array[0].get("audio_url").is_some());
1450 assert_eq!(content_array[1]["type"], "video_url");
1451 assert!(content_array[1].get("video_url").is_some());
1452 assert_eq!(content_array[2]["type"], "image");
1454 assert!(content_array[2].get("image_url").is_none());
1455 }
1456
1457 #[test]
1459 fn test_convert_media_url_to_placeholder_multiple_types() {
1460 let mut content_array = vec![
1461 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1462 serde_json::json!({"type": "text", "text": "and listen to"}),
1463 serde_json::json!({"type": "audio_url", "audio_url": {"url": "https://example.com/audio.mp3"}}),
1464 serde_json::json!({"type": "text", "text": "and watch"}),
1465 serde_json::json!({"type": "video_url", "video_url": {"url": "https://example.com/video.mp4"}}),
1466 ];
1467
1468 let conversions = &[
1470 ("image_url", "image"),
1471 ("audio_url", "audio"),
1472 ("video_url", "video"),
1473 ];
1474 convert_media_url_to_placeholder(&mut content_array, conversions);
1475
1476 assert_eq!(content_array.len(), 5);
1477 assert_eq!(content_array[0]["type"], "image");
1478 assert!(content_array[0].get("image_url").is_none());
1479 assert_eq!(content_array[1]["type"], "text");
1480 assert_eq!(content_array[2]["type"], "audio");
1481 assert!(content_array[2].get("audio_url").is_none());
1482 assert_eq!(content_array[3]["type"], "text");
1483 assert_eq!(content_array[4]["type"], "video");
1484 assert!(content_array[4].get("video_url").is_none());
1485 }
1486
1487 #[test]
1489 fn test_convert_media_url_to_placeholder_no_conversions() {
1490 let mut content_array = vec![
1491 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1492 serde_json::json!({"type": "text", "text": "hello"}),
1493 ];
1494
1495 let conversions: &[(&str, &str)] = &[];
1496 convert_media_url_to_placeholder(&mut content_array, conversions);
1497
1498 assert_eq!(content_array.len(), 2);
1499 assert_eq!(content_array[0]["type"], "image_url");
1501 assert!(content_array[0].get("image_url").is_some());
1502 assert_eq!(content_array[1]["type"], "text");
1503 }
1504
1505 #[test]
1508 fn test_default_media_type_conversions_only_converts_image_url() {
1509 let mut content_array = vec![
1510 serde_json::json!({"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}),
1511 serde_json::json!({"type": "video_url", "video_url": {"url": "https://example.com/video.mp4"}}),
1512 serde_json::json!({"type": "audio_url", "audio_url": {"url": "https://example.com/audio.mp3"}}),
1513 serde_json::json!({"type": "text", "text": "hello"}),
1514 ];
1515
1516 convert_media_url_to_placeholder(&mut content_array, DEFAULT_MEDIA_TYPE_CONVERSIONS);
1518
1519 assert_eq!(content_array.len(), 4);
1520
1521 assert_eq!(content_array[0]["type"], "image");
1523 assert!(content_array[0].get("image_url").is_none());
1524
1525 assert_eq!(content_array[1]["type"], "video");
1527 assert!(content_array[1].get("video_url").is_none());
1528
1529 assert_eq!(content_array[2]["type"], "audio");
1531 assert!(content_array[2].get("audio_url").is_none());
1532
1533 assert_eq!(content_array[3]["type"], "text");
1535 assert_eq!(content_array[3]["text"], "hello");
1536 }
1537
1538 #[test]
1539 fn test_may_be_fix_tool_schema_missing_type_and_properties() {
1540 let json_str = r#"{
1541 "model": "gpt-4o",
1542 "messages": [],
1543 "tools": [
1544 {
1545 "type": "function",
1546 "function": {
1547 "name": "get_weather",
1548 "description": "Get the current weather in a given location",
1549 "parameters": {},
1550 "strict": null
1551 }
1552 }
1553 ]
1554 }"#;
1555
1556 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1557 let tools = serde_json::to_value(request.tools()).unwrap();
1558
1559 assert!(tools[0]["function"]["parameters"]["type"] == "object");
1560 assert!(
1561 tools[0]["function"]["parameters"]["properties"]
1562 == serde_json::Value::Object(Default::default())
1563 );
1564 }
1565
1566 #[test]
1567 fn test_may_be_fix_tool_schema_missing_type() {
1568 let json_str = r#"{
1569 "model": "gpt-4o",
1570 "messages": [],
1571 "tools": [
1572 {
1573 "type": "function",
1574 "function": {
1575 "name": "get_weather",
1576 "description": "Get the current weather in a given location",
1577 "parameters": {
1578 "properties": {
1579 "location": {
1580 "type": "string",
1581 "description": "City and state, e.g., 'San Francisco, CA'"
1582 }
1583 }
1584 },
1585 "strict": null
1586 }
1587 }
1588 ]
1589 }"#;
1590 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1591
1592 let tools = serde_json::to_value(request.tools()).unwrap();
1593
1594 assert_eq!(tools[0]["function"]["parameters"]["type"], "object");
1595
1596 let mut expected_properties = serde_json::Map::new();
1597 let mut location = serde_json::Map::new();
1598 location.insert(
1599 "type".to_string(),
1600 serde_json::Value::String("string".to_string()),
1601 );
1602 location.insert(
1603 "description".to_string(),
1604 serde_json::Value::String("City and state, e.g., 'San Francisco, CA'".to_string()),
1605 );
1606 expected_properties.insert("location".to_string(), serde_json::Value::Object(location));
1607
1608 assert_eq!(
1609 tools[0]["function"]["parameters"]["properties"],
1610 serde_json::Value::Object(expected_properties)
1611 );
1612 }
1613
1614 #[test]
1615 fn test_may_be_fix_tool_schema_missing_properties() {
1616 let json_str = r#"{
1617 "model": "gpt-4o",
1618 "messages": [],
1619 "tools": [
1620 {
1621 "type": "function",
1622 "function": {
1623 "name": "get_weather",
1624 "description": "Get the current weather in a given location",
1625 "parameters": {"type": "object"},
1626 "strict": null
1627 }
1628 }
1629 ]
1630 }"#;
1631
1632 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1633 let tools = serde_json::to_value(request.tools()).unwrap();
1634
1635 assert_eq!(
1636 tools[0]["function"]["parameters"]["properties"],
1637 serde_json::Value::Object(Default::default())
1638 );
1639 assert_eq!(tools[0]["function"]["parameters"]["type"], "object");
1640 }
1641
1642 #[test]
1643 fn test_may_be_fix_tool_schema_missing_description() {
1644 let json_str = r#"{
1648 "model": "gpt-4o",
1649 "messages": [],
1650 "tools": [
1651 {
1652 "type": "function",
1653 "function": {
1654 "name": "noop",
1655 "parameters": {
1656 "type": "object",
1657 "properties": { "x": { "type": "string" } },
1658 "required": ["x"],
1659 "additionalProperties": false
1660 },
1661 "strict": null
1662 }
1663 }
1664 ]
1665 }"#;
1666
1667 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1668 let tools = serde_json::to_value(request.tools()).unwrap();
1669
1670 assert_eq!(
1671 tools[0]["function"]["description"],
1672 serde_json::Value::String(String::new())
1673 );
1674 }
1675
1676 #[test]
1677 fn test_may_be_fix_tool_schema_null_description() {
1678 let json_str = r#"{
1680 "model": "gpt-4o",
1681 "messages": [],
1682 "tools": [
1683 {
1684 "type": "function",
1685 "function": {
1686 "name": "noop",
1687 "description": null,
1688 "parameters": {"type": "object", "properties": {}},
1689 "strict": null
1690 }
1691 }
1692 ]
1693 }"#;
1694
1695 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1696 let tools = serde_json::to_value(request.tools()).unwrap();
1697
1698 assert_eq!(
1699 tools[0]["function"]["description"],
1700 serde_json::Value::String(String::new())
1701 );
1702 }
1703
1704 #[test]
1705 fn test_may_be_fix_tool_schema_preserves_description() {
1706 let json_str = r#"{
1708 "model": "gpt-4o",
1709 "messages": [],
1710 "tools": [
1711 {
1712 "type": "function",
1713 "function": {
1714 "name": "get_weather",
1715 "description": "Get the current weather in a given location",
1716 "parameters": {"type": "object", "properties": {}},
1717 "strict": null
1718 }
1719 }
1720 ]
1721 }"#;
1722
1723 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1724 let tools = serde_json::to_value(request.tools()).unwrap();
1725
1726 assert_eq!(
1727 tools[0]["function"]["description"],
1728 "Get the current weather in a given location"
1729 );
1730 }
1731
1732 #[test]
1734 fn test_may_be_fix_msg_content_user_multipart() {
1735 let json_str = r#"{
1736 "model": "gpt-4o",
1737 "messages": [
1738 {
1739 "role": "user",
1740 "content": [
1741 {"type": "text", "text": "part 1"},
1742 {"type": "text", "text": "part 2"}
1743 ]
1744 }
1745 ]
1746 }"#;
1747
1748 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1749 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1750
1751 let messages =
1753 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1754
1755 assert_eq!(
1757 messages[0]["content"],
1758 serde_json::Value::String("part 1\npart 2".to_string())
1759 );
1760 }
1761
1762 #[test]
1765 fn test_may_be_fix_msg_content_mixed_messages() {
1766 let json_str = r#"{
1767 "model": "gpt-4o",
1768 "messages": [
1769 {
1770 "role": "system",
1771 "content": "You are a helpful assistant"
1772 },
1773 {
1774 "role": "user",
1775 "content": [
1776 {"type": "text", "text": "Hello"},
1777 {"type": "text", "text": "World"}
1778 ]
1779 },
1780 {
1781 "role": "assistant",
1782 "content": "Hi there!"
1783 },
1784 {
1785 "role": "user",
1786 "content": [
1787 {"type": "text", "text": "Another"},
1788 {"type": "text", "text": "multi-part"},
1789 {"type": "text", "text": "message"}
1790 ]
1791 }
1792 ]
1793 }"#;
1794
1795 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1796 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1797
1798 let messages =
1800 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1801
1802 assert_eq!(
1804 messages[0]["content"],
1805 serde_json::Value::String("You are a helpful assistant".to_string())
1806 );
1807
1808 assert_eq!(
1810 messages[1]["content"],
1811 serde_json::Value::String("Hello\nWorld".to_string())
1812 );
1813
1814 assert_eq!(
1816 messages[2]["content"],
1817 serde_json::Value::String("Hi there!".to_string())
1818 );
1819
1820 assert_eq!(
1822 messages[3]["content"],
1823 serde_json::Value::String("Another\nmulti-part\nmessage".to_string())
1824 );
1825 }
1826
1827 #[test]
1829 fn test_may_be_fix_msg_content_empty_array() {
1830 let json_str = r#"{
1831 "model": "gpt-4o",
1832 "messages": [
1833 {
1834 "role": "user",
1835 "content": []
1836 }
1837 ]
1838 }"#;
1839
1840 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1841 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1842
1843 let messages =
1845 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1846
1847 assert!(messages[0]["content"].is_array());
1849 assert_eq!(messages[0]["content"].as_array().unwrap().len(), 0);
1850 }
1851
1852 #[test]
1859 fn test_may_be_fix_msg_content_empty_array_with_placeholder_template() {
1860 let json_str = r#"{
1861 "model": "phi-3-vision",
1862 "messages": [
1863 {
1864 "role": "user",
1865 "content": []
1866 }
1867 ]
1868 }"#;
1869
1870 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1871 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1872
1873 let messages = serde_json::to_value(may_be_fix_msg_content(
1876 messages_raw,
1877 false,
1878 Some("<|image_{n}|>"),
1879 ))
1880 .unwrap();
1881
1882 assert!(
1883 messages[0]["content"].is_array(),
1884 "empty array should be preserved as `[]`, not flattened to `\"\"`"
1885 );
1886 assert_eq!(messages[0]["content"].as_array().unwrap().len(), 0);
1887 }
1888
1889 #[test]
1891 fn test_may_be_fix_msg_content_single_text() {
1892 let json_str = r#"{
1893 "model": "gpt-4o",
1894 "messages": [
1895 {
1896 "role": "user",
1897 "content": "Simple text message"
1898 }
1899 ]
1900 }"#;
1901
1902 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1903 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1904
1905 let messages =
1907 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1908
1909 assert_eq!(
1911 messages[0]["content"],
1912 serde_json::Value::String("Simple text message".to_string())
1913 );
1914 }
1915
1916 #[test]
1919 fn test_may_be_fix_msg_content_mixed_types() {
1920 let json_str = r#"{
1921 "model": "gpt-4o",
1922 "messages": [
1923 {
1924 "role": "user",
1925 "content": [
1926 {"type": "text", "text": "Check this image:"},
1927 {"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}},
1928 {"type": "text", "text": "What do you see?"}
1929 ]
1930 }
1931 ]
1932 }"#;
1933
1934 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1935 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1936
1937 let messages =
1939 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
1940
1941 assert!(messages[0]["content"].is_array());
1944 let content_array = messages[0]["content"].as_array().unwrap();
1945 assert_eq!(content_array.len(), 3);
1946 assert_eq!(content_array[0]["type"], "text");
1947 assert_eq!(content_array[1]["type"], "image");
1948 assert!(content_array[1].get("image_url").is_none());
1949 assert_eq!(content_array[2]["type"], "text");
1950 }
1951
1952 #[test]
1958 fn test_may_be_fix_msg_content_flattens_phi3_style() {
1959 let json_str = r#"{
1960 "model": "phi-3-vision",
1961 "messages": [
1962 {
1963 "role": "user",
1964 "content": [
1965 {"type": "text", "text": "First "},
1966 {"type": "image_url", "image_url": {"url": "https://example.com/a.jpg"}},
1967 {"type": "text", "text": " then "},
1968 {"type": "image_url", "image_url": {"url": "https://example.com/b.jpg"}},
1969 {"type": "text", "text": "?"}
1970 ]
1971 }
1972 ]
1973 }"#;
1974 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
1975 let messages_raw = serde_json::to_value(request.messages()).unwrap();
1976
1977 let messages = serde_json::to_value(may_be_fix_msg_content(
1978 messages_raw,
1979 false,
1980 Some("<|image_{n}|>"),
1981 ))
1982 .unwrap();
1983
1984 let content = messages[0]["content"].as_str().expect("content flattened");
1985 assert_eq!(content, "First <|image_1|> then <|image_2|>?");
1986 }
1987
1988 #[test]
1990 fn test_may_be_fix_msg_content_flattens_llava_style() {
1991 let json_str = r#"{
1992 "model": "llava-1.5-7b-hf",
1993 "messages": [
1994 {
1995 "role": "user",
1996 "content": [
1997 {"type": "text", "text": "Describe: "},
1998 {"type": "image_url", "image_url": {"url": "https://example.com/x.jpg"}}
1999 ]
2000 }
2001 ]
2002 }"#;
2003 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2004 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2005
2006 let messages =
2007 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, Some("<image>")))
2008 .unwrap();
2009
2010 let content = messages[0]["content"].as_str().expect("content flattened");
2011 assert_eq!(content, "Describe: <image>");
2012 }
2013
2014 #[test]
2020 fn test_may_be_fix_msg_content_flattens_empty_placeholder() {
2021 let json_str = r#"{
2022 "model": "nvidia/NVIDIA-Nemotron-Parse-v1.2",
2023 "messages": [
2024 {
2025 "role": "user",
2026 "content": [
2027 {"type": "text", "text": "</s><s><predict_bbox><predict_classes><output_markdown><predict_no_text_in_pic>"},
2028 {"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}}
2029 ]
2030 }
2031 ]
2032 }"#;
2033 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2034 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2035
2036 let messages =
2037 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, Some(""))).unwrap();
2038
2039 let content = messages[0]["content"].as_str().expect("content flattened");
2040 assert_eq!(
2041 content,
2042 "</s><s><predict_bbox><predict_classes><output_markdown><predict_no_text_in_pic>"
2043 );
2044 }
2045
2046 #[test]
2053 fn test_render_nemotron_parse_passthrough() {
2054 use super::super::tokcfg::ChatTemplate;
2055 use super::{ContextMixins, HfTokenizerConfigJsonFormatter};
2056
2057 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
2058 "chat_template": "{% for message in messages %}{{ message['content'] }}{% endfor %}"
2059 }))
2060 .unwrap();
2061 let formatter =
2062 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
2063
2064 for prompt in [
2065 "</s><s><predict_bbox><predict_classes><output_markdown><predict_no_text_in_pic>",
2066 "</s><s><predict_bbox><predict_classes><output_markdown><predict_text_in_pic>",
2067 ] {
2068 let request: NvCreateChatCompletionRequest =
2069 serde_json::from_value(serde_json::json!({
2070 "model": "nvidia/NVIDIA-Nemotron-Parse-v1.2",
2071 "messages": [{
2072 "role": "user",
2073 "content": [
2074 {"type": "text", "text": prompt},
2075 {"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}}
2076 ]
2077 }]
2078 }))
2079 .unwrap();
2080
2081 let rendered = formatter.render(&request).unwrap();
2082 assert_eq!(
2083 rendered, prompt,
2084 "rendered prompt must be the control tokens only, with no JSON-serialized image array"
2085 );
2086 }
2087 }
2088
2089 #[test]
2092 fn test_may_be_fix_msg_content_non_text_only() {
2093 let json_str = r#"{
2094 "model": "gpt-4o",
2095 "messages": [
2096 {
2097 "role": "user",
2098 "content": [
2099 {"type": "image_url", "image_url": {"url": "https://example.com/image1.jpg"}},
2100 {"type": "image_url", "image_url": {"url": "https://example.com/image2.jpg"}}
2101 ]
2102 }
2103 ]
2104 }"#;
2105
2106 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2107 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2108
2109 let messages =
2111 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
2112
2113 assert!(messages[0]["content"].is_array());
2115 let content_array = messages[0]["content"].as_array().unwrap();
2116 assert_eq!(content_array.len(), 2);
2117 assert_eq!(content_array[0]["type"], "image");
2118 assert_eq!(content_array[1]["type"], "image");
2119 }
2120
2121 #[test]
2122 fn test_none_tools_safe_for_all_templates() {
2123 use super::tokcfg::ChatTemplate;
2124 use super::{ContextMixins, HfTokenizerConfigJsonFormatter};
2125
2126 let length_template = r#"
2130{%- if tools is iterable and tools | length > 0 %}
2131Tools available: {{ tools | length }}
2132{%- else %}
2133No tools
2134{%- endif %}
2135"#;
2136
2137 let no_tool_template = r#"
2140{%- if tools is not none %}
2141TOOL MODE
2142{%- else %}
2143NORMAL MODE
2144{%- endif %}
2145"#;
2146
2147 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
2148 "chat_template": [
2149 {"safe_length": length_template},
2150 {"no_tool": no_tool_template}
2151 ]
2152 }))
2153 .unwrap();
2154
2155 let formatter =
2156 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
2157
2158 let ctx = context! { tools => Option::<Value>::None };
2159
2160 let result1 = formatter
2161 .env
2162 .get_template("safe_length")
2163 .unwrap()
2164 .render(&ctx);
2165 println!("Safe length template with no tools => None: {:?}", result1);
2166 assert!(
2167 result1.is_ok(),
2168 "Jinja template with and conditional and length filter should handle None: {:?}",
2169 result1
2170 );
2171 assert!(
2172 result1.unwrap().contains("No tools"),
2173 "Should show 'No tools'"
2174 );
2175
2176 let result2 = formatter.env.get_template("no_tool").unwrap().render(&ctx);
2177 println!("Default template with no tools => None: {:?}", result2);
2178 assert!(
2179 result2.is_ok(),
2180 "Jinja template with if tools is not none conditional should handle None: {:?}",
2181 result2
2182 );
2183 assert!(result2.unwrap().contains("NORMAL MODE"));
2184 }
2185
2186 #[test]
2188 fn test_may_be_fix_msg_content_multiple_content_types() {
2189 let json_str = r#"{
2191 "model": "gpt-4o",
2192 "messages": [
2193 {
2194 "role": "user",
2195 "content": [
2196 {"type": "text", "text": "Listen to this:"},
2197 {"type": "audio_url", "audio_url": {"url": "https://example.com/audio.mp3"}},
2198 {"type": "text", "text": "And look at:"},
2199 {"type": "image_url", "image_url": {"url": "https://example.com/img.jpg"}},
2200 {"type": "text", "text": "What do you think?"}
2201 ]
2202 }
2203 ]
2204 }"#;
2205
2206 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2207 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2208 let messages =
2209 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
2210
2211 assert!(messages[0]["content"].is_array());
2213 let content_array = messages[0]["content"].as_array().unwrap();
2214 assert_eq!(content_array.len(), 5);
2215 assert_eq!(content_array[0]["type"], "text");
2216 assert_eq!(content_array[1]["type"], "audio");
2217 assert_eq!(content_array[2]["type"], "text");
2218 assert_eq!(content_array[3]["type"], "image");
2219 assert_eq!(content_array[4]["type"], "text");
2220
2221 let json_str = r#"{
2223 "model": "gpt-4o",
2224 "messages": [
2225 {
2226 "role": "user",
2227 "content": [
2228 {"type": "text", "text": "Check this:"},
2229 {"type": "video_url", "video_url": {"url": "https://example.com/vid.mp4"}},
2230 {"type": "text", "text": "Interesting?"}
2231 ]
2232 }
2233 ]
2234 }"#;
2235
2236 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2237 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2238 let messages =
2239 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
2240
2241 assert!(messages[0]["content"].is_array());
2243 assert_eq!(messages[0]["content"].as_array().unwrap().len(), 3);
2244 }
2245
2246 #[test]
2247 fn test_normalize_tool_arguments_tojson() {
2248 let tmpl = r#"{{ messages[0].tool_calls[0].function.arguments | tojson }}"#;
2249
2250 let mut messages = serde_json::Value::Array(vec![serde_json::json!({
2252 "role": "assistant",
2253 "tool_calls": [{
2254 "type": "function",
2255 "function": {
2256 "name": "get_current_weather",
2257 "arguments": "{\"format\":\"celsius\",\"location\":\"San Francisco, CA\"}"
2258 }
2259 }]
2260 })]);
2261
2262 normalize_tool_calls_arguments_in_messages(&mut messages);
2263
2264 let mut env = Environment::new();
2265 env.add_filter("tojson", super::super::tokcfg::tojson);
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_eq!(
2276 out,
2277 r#"{"format": "celsius", "location": "San Francisco, CA"}"#
2278 );
2279 }
2280
2281 #[test]
2282 fn test_normalize_tool_arguments_items_loop() {
2283 let tmpl = r#"{% for k, v in messages[0].tool_calls[0].function.arguments|items %}{{k}}={{v}};{% endfor %}"#;
2284
2285 let mut messages = serde_json::Value::Array(vec![serde_json::json!({
2286 "role": "assistant",
2287 "tool_calls": [{
2288 "type": "function",
2289 "function": {
2290 "name": "f",
2291 "arguments": "{\"a\":1,\"b\":\"x\"}"
2292 }
2293 }]
2294 })]);
2295
2296 normalize_tool_calls_arguments_in_messages(&mut messages);
2297
2298 let mut env = Environment::new();
2299 env.add_template("t", tmpl).unwrap();
2300 let out = env
2301 .get_template("t")
2302 .unwrap()
2303 .render(context! { messages => messages.as_array().unwrap() })
2304 .unwrap();
2305
2306 assert!(out == "a=1;b=x;" || out == "b=x;a=1;");
2307 }
2308
2309 #[test]
2310 fn test_normalize_tool_arguments_legacy_function_call() {
2311 let mut messages = serde_json::Value::Array(vec![serde_json::json!({
2313 "role": "assistant",
2314 "function_call": {
2315 "name": "get_weather",
2316 "arguments": "{\"location\":\"NYC\"}"
2317 }
2318 })]);
2319
2320 normalize_function_call_arguments_in_messages(&mut messages);
2321
2322 assert_eq!(
2323 messages[0]["function_call"]["arguments"],
2324 serde_json::json!({"location": "NYC"})
2325 );
2326 }
2327
2328 #[test]
2329 fn test_normalize_tool_arguments_malformed_json_passthrough() {
2330 let mut messages = serde_json::Value::Array(vec![serde_json::json!({
2332 "role": "assistant",
2333 "tool_calls": [{
2334 "type": "function",
2335 "function": {
2336 "name": "f",
2337 "arguments": "not valid json at all"
2338 }
2339 }]
2340 })]);
2341
2342 normalize_tool_calls_arguments_in_messages(&mut messages);
2343
2344 assert_eq!(
2345 messages[0]["tool_calls"][0]["function"]["arguments"],
2346 serde_json::Value::String("not valid json at all".to_string())
2347 );
2348 }
2349
2350 #[test]
2351 fn test_normalize_tool_arguments_with_multimodal_content() {
2352 let json_str = r#"{
2353 "model": "gpt-4o",
2354 "messages": [
2355 {
2356 "role": "user",
2357 "content": [
2358 {"type": "text", "text": "Check this:"},
2359 {"type": "video_url", "video_url": {"url": "https://example.com/vid.mp4"}},
2360 {"type": "text", "text": "Interesting?"}
2361 ]
2362 },
2363 {
2364 "role": "assistant",
2365 "tool_calls": [{
2366 "id": "call_123",
2367 "type": "function",
2368 "function": {
2369 "name": "analyze_video",
2370 "arguments": "{\"url\":\"https://example.com/vid.mp4\",\"format\":\"mp4\"}"
2371 }
2372 }]
2373 }
2374 ]
2375 }"#;
2376
2377 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2378 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2379
2380 let mut messages =
2382 serde_json::to_value(may_be_fix_msg_content(messages_raw, false, None)).unwrap();
2383
2384 normalize_tool_calls_arguments_in_messages(&mut messages);
2385
2386 assert!(messages[0]["content"].is_array());
2388 assert_eq!(messages[0]["content"].as_array().unwrap().len(), 3);
2389
2390 assert!(messages[1]["tool_calls"][0]["function"]["arguments"].is_object());
2392 assert_eq!(
2393 messages[1]["tool_calls"][0]["function"]["arguments"]["url"],
2394 "https://example.com/vid.mp4"
2395 );
2396 }
2397
2398 #[test]
2401 fn test_minimax_m3_history_tool_call_float_arguments_match_hf() {
2402 let template = r#"{%- set ns_token = ']<]minimax[>[' -%}
2403{%- macro to_xml(val, ns) -%}
2404{%- if val is mapping -%}
2405{%- for k, v in val.items() if v is not none -%}
2406{{ ns }}<{{ k }}>{{ to_xml(v, ns) }}{{ ns }}</{{ k }}>
2407{%- endfor -%}
2408{%- elif val is iterable and val is not string -%}
2409{%- for item in val -%}
2410{{ ns }}<item>{{ to_xml(item, ns) }}{{ ns }}</item>
2411{%- endfor -%}
2412{%- elif val is none -%}
2413{%- elif val is boolean -%}
2414{{ val | tojson }}
2415{%- else -%}
2416{{ val }}
2417{%- endif -%}
2418{%- endmacro -%}
2419{%- for message in messages if message.tool_calls -%}
2420{%- for tool_call in message.tool_calls -%}
2421{%- if tool_call.function -%}
2422{%- set tool_call = tool_call.function -%}
2423{%- endif -%}
2424{{- ns_token + '<invoke name="' + tool_call.name + '">' }}
2425{%- set _args = tool_call.arguments -%}
2426{%- for k, v in _args.items() if v is not none %}
2427{{- ns_token + '<' + k + '>' -}}
2428{{- to_xml(v, ns_token) -}}
2429{{- ns_token + '</' + k + '>' }}
2430{%- endfor -%}
2431{{- ns_token + '</invoke>' ~ '\n' }}
2432{%- endfor -%}
2433{%- endfor -%}"#;
2434 let rendered = render_shape(
2435 &formatter_for(template),
2436 json!([
2437 {"role": "user", "content": "u"},
2438 {"role": "assistant", "content": "", "tool_calls": [{
2439 "id": "c1",
2440 "type": "function",
2441 "function": {
2442 "name": "fit",
2443 "arguments": r#"{"tolerance": 1e-07, "bounds": [0.00001, 1e16]}"#
2444 }
2445 }]}
2446 ]),
2447 )
2448 .unwrap();
2449 assert_eq!(
2450 rendered,
2451 "]<]minimax[>[<invoke name=\"fit\">]<]minimax[>[<tolerance>1e-07]<]minimax[>[</tolerance>]<]minimax[>[<bounds>]<]minimax[>[<item>1e-05]<]minimax[>[</item>]<]minimax[>[<item>1e+16]<]minimax[>[</item>]<]minimax[>[</bounds>]<]minimax[>[</invoke>\n"
2452 );
2453 }
2454
2455 #[test]
2458 fn test_qwen3_coder_history_tool_call_float_arguments_match_hf() {
2459 let template = r#"{%- for message in messages if message.tool_calls -%}
2460{%- for tool_call in message.tool_calls %}
2461 {%- if tool_call.function is defined %}
2462 {%- set tool_call = tool_call.function %}
2463 {%- endif %}
2464 {%- for args_name, args_value in tool_call.arguments|items %}
2465 {{- '<parameter=' + args_name + '>\n' }}
2466 {%- set args_value = args_value | tojson | safe if args_value is mapping or (args_value is sequence and args_value is not string) else args_value | string %}
2467 {{- args_value }}
2468 {{- '\n</parameter>\n' }}
2469 {%- endfor %}
2470{%- endfor %}
2471{%- endfor %}"#;
2472 let rendered = render_shape(
2473 &formatter_for(template),
2474 json!([
2475 {"role": "user", "content": "u"},
2476 {"role": "assistant", "content": "", "tool_calls": [{
2477 "id": "c1",
2478 "type": "function",
2479 "function": {
2480 "name": "fit",
2481 "arguments": r#"{"tolerance": 1e-07, "bounds": [0.00001, 1e16]}"#
2482 }
2483 }]}
2484 ]),
2485 )
2486 .unwrap();
2487 assert_eq!(
2488 rendered,
2489 "<parameter=tolerance>\n1e-07\n</parameter>\n<parameter=bounds>\n[1e-05, 1e+16]\n</parameter>\n"
2490 );
2491 }
2492
2493 #[test]
2495 fn test_may_be_fix_msg_content_string_to_array() {
2496 let json_str = r#"{
2497 "model": "gpt-4o",
2498 "messages": [
2499 {
2500 "role": "user",
2501 "content": "Hello, how are you?"
2502 }
2503 ]
2504 }"#;
2505
2506 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2507 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2508
2509 let messages =
2511 serde_json::to_value(may_be_fix_msg_content(messages_raw, true, None)).unwrap();
2512
2513 assert!(messages[0]["content"].is_array());
2515 let content_array = messages[0]["content"].as_array().unwrap();
2516 assert_eq!(content_array.len(), 1);
2517 assert_eq!(content_array[0]["type"], "text");
2518 assert_eq!(content_array[0]["text"], "Hello, how are you?");
2519 }
2520
2521 #[test]
2523 fn test_may_be_fix_msg_content_array_preserved_with_multimodal() {
2524 let json_str = r#"{
2525 "model": "gpt-4o",
2526 "messages": [
2527 {
2528 "role": "user",
2529 "content": [
2530 {"type": "text", "text": "part 1"},
2531 {"type": "text", "text": "part 2"}
2532 ]
2533 }
2534 ]
2535 }"#;
2536
2537 let request: NvCreateChatCompletionRequest = serde_json::from_str(json_str).unwrap();
2538 let messages_raw = serde_json::to_value(request.messages()).unwrap();
2539
2540 let messages =
2542 serde_json::to_value(may_be_fix_msg_content(messages_raw, true, None)).unwrap();
2543
2544 assert!(messages[0]["content"].is_array());
2546 let content_array = messages[0]["content"].as_array().unwrap();
2547 assert_eq!(content_array.len(), 2);
2548 assert_eq!(content_array[0]["text"], "part 1");
2549 assert_eq!(content_array[1]["text"], "part 2");
2550 }
2551
2552 fn user() -> Msg {
2553 Msg::User(Default::default())
2554 }
2555 fn tool() -> Msg {
2556 Msg::Tool(Default::default())
2557 }
2558
2559 fn dummy_state(messages: Vec<Msg>) -> NvCreateChatCompletionRequest {
2560 let json = serde_json::json!({
2561 "model": "test-model",
2562 "messages": messages
2563 });
2564 serde_json::from_value(json).unwrap()
2565 }
2566
2567 #[test]
2568 fn add_after_user() {
2569 let s = dummy_state(vec![user()]);
2570 assert!(s.should_add_generation_prompt());
2571 }
2572
2573 #[test]
2574 fn add_after_tool() {
2575 let s = dummy_state(vec![tool()]);
2576 assert!(s.should_add_generation_prompt());
2577 }
2578
2579 #[test]
2580 fn add_when_empty() {
2581 let s = dummy_state(vec![]);
2582 assert!(s.should_add_generation_prompt());
2583 }
2584
2585 fn tool_aware_formatter(
2587 exclude_tools_when_tool_choice_none: bool,
2588 ) -> HfTokenizerConfigJsonFormatter {
2589 let template = r#"
2590{%- if tools is iterable and tools | length > 0 %}
2591TOOL_MODE tools={{ tools | length }}
2592{%- else %}
2593NORMAL_MODE
2594{%- endif %}
2595{{ messages[0].content }}"#;
2596
2597 let chat_template: super::tokcfg::ChatTemplate =
2598 serde_json::from_value(serde_json::json!({ "chat_template": template })).unwrap();
2599
2600 HfTokenizerConfigJsonFormatter::with_options(
2601 chat_template,
2602 ContextMixins::new(&[]),
2603 exclude_tools_when_tool_choice_none,
2604 )
2605 .unwrap()
2606 }
2607
2608 fn gemma4_tool_template_for_tests() -> &'static str {
2609 r#"
2610{{ bos_token }}
2611{%- set loop_messages = messages -%}
2612{%- set ns_turn = namespace(last_user_idx=-1) -%}
2613{%- for i in range(loop_messages | length) -%}
2614 {%- if loop_messages[i]['role'] == 'user' -%}
2615 {%- set ns_turn.last_user_idx = i -%}
2616 {%- endif -%}
2617{%- endfor -%}
2618{%- for message in loop_messages -%}
2619 {%- set role = 'model' if message['role'] == 'assistant' else message['role'] -%}
2620 {{- '<|turn>' + role + '\n' }}
2621
2622 {%- if message.get('reasoning') and loop.index0 > ns_turn.last_user_idx and message.get('tool_calls') -%}
2623 {{- '<|channel>thought\n' + message['reasoning'] + '\n<channel|>'}}
2624 {%- endif -%}
2625
2626 {%- if message['tool_calls'] -%}
2627 {%- for tool_call in message['tool_calls'] -%}
2628 {%- set function = tool_call['function'] -%}
2629 {{- '<|tool_call>call:' + function['name'] + '{' -}}
2630 {%- if function['arguments'] is mapping -%}
2631 {%- set ns_args = namespace(found_first=false) -%}
2632 {%- for key, value in function['arguments'] | dictsort -%}
2633 {%- if ns_args.found_first %},{% endif -%}
2634 {%- set ns_args.found_first = true -%}
2635 {{- key -}}:{{- value -}}
2636 {%- endfor -%}
2637 {%- elif function['arguments'] is string -%}
2638 {{- function['arguments'] -}}
2639 {%- endif -%}
2640 {{- '}<tool_call|>' -}}
2641 {%- endfor -%}
2642 {%- endif -%}
2643
2644 {%- if message['content'] is string -%}
2645 {{- message['content'] -}}
2646 {%- endif -%}
2647 {{- '<turn|>\n' -}}
2648{%- endfor -%}
2649"#
2650 }
2651
2652 fn make_gemma4_tool_formatter_for_tests() -> HfTokenizerConfigJsonFormatter {
2653 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
2654 "chat_template": gemma4_tool_template_for_tests()
2655 }))
2656 .unwrap();
2657 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap()
2658 }
2659
2660 fn request_with_tool_choice(tool_choice: &str) -> NvCreateChatCompletionRequest {
2662 serde_json::from_value(serde_json::json!({
2663 "model": "test",
2664 "messages": [{"role": "user", "content": "hello"}],
2665 "tools": [{
2666 "type": "function",
2667 "function": {
2668 "name": "get_weather",
2669 "description": "Get weather",
2670 "parameters": {"type": "object", "properties": {"location": {"type": "string"}}}
2671 }
2672 }],
2673 "tool_choice": tool_choice
2674 }))
2675 .unwrap()
2676 }
2677
2678 #[test]
2679 fn test_exclude_tools_strips_when_tool_choice_none() {
2680 let formatter = tool_aware_formatter(true);
2681 let request = request_with_tool_choice("none");
2682 let result = formatter.render(&request).unwrap();
2683 assert!(
2684 result.contains("NORMAL_MODE"),
2685 "With exclude_tools=true and tool_choice=none, tools should be stripped. Got: {}",
2686 result
2687 );
2688 }
2689
2690 #[test]
2691 fn test_exclude_tools_keeps_when_tool_choice_auto() {
2692 let formatter = tool_aware_formatter(true);
2693 let request = request_with_tool_choice("auto");
2694 let result = formatter.render(&request).unwrap();
2695 assert!(
2696 result.contains("TOOL_MODE"),
2697 "With tool_choice=auto, tools should be included. Got: {}",
2698 result
2699 );
2700 }
2701
2702 #[test]
2703 fn test_no_exclude_tools_keeps_when_tool_choice_none() {
2704 let formatter = tool_aware_formatter(false);
2705 let request = request_with_tool_choice("none");
2706 let result = formatter.render(&request).unwrap();
2707 assert!(
2708 result.contains("TOOL_MODE"),
2709 "With exclude_tools=false and tool_choice=none, tools should NOT be stripped. Got: {}",
2710 result
2711 );
2712 }
2713
2714 #[test]
2715 fn test_inject_reasoning_content_segments_with_tool_calls() {
2716 let mut messages = serde_json::json!([
2718 {
2719 "role": "user",
2720 "content": "What is sqrt(144) and sqrt(256)?"
2721 },
2722 {
2723 "role": "assistant",
2724 "content": "Let me calculate those.",
2725 "reasoning_content": ["I need to compute sqrt(144)", "Now sqrt(256)", ""],
2726 "tool_calls": [
2727 {
2728 "id": "call_0",
2729 "type": "function",
2730 "function": {
2731 "name": "calculator",
2732 "arguments": "{\"expr\": \"sqrt(144)\"}"
2733 }
2734 },
2735 {
2736 "id": "call_1",
2737 "type": "function",
2738 "function": {
2739 "name": "calculator",
2740 "arguments": "{\"expr\": \"sqrt(256)\"}"
2741 }
2742 }
2743 ]
2744 }
2745 ]);
2746
2747 inject_reasoning_content_into_messages(&mut messages);
2748
2749 let assistant = &messages[1];
2750
2751 assert!(
2753 assistant.get("reasoning_content").is_none(),
2754 "reasoning_content should be removed after injection"
2755 );
2756
2757 let content = assistant["content"].as_str().unwrap();
2759 assert!(
2760 content.starts_with("<think>I need to compute sqrt(144)</think>"),
2761 "content should start with first reasoning segment, got: {}",
2762 content
2763 );
2764 assert!(
2765 content.contains("<think>Now sqrt(256)</think>"),
2766 "content should contain second reasoning segment"
2767 );
2768 assert!(
2770 !content.contains("<think></think>"),
2771 "empty segments should be skipped"
2772 );
2773 assert!(
2775 content.ends_with("Let me calculate those."),
2776 "original content should be at the end, got: {}",
2777 content
2778 );
2779
2780 assert!(assistant.get("tool_calls").is_some());
2782 assert_eq!(assistant["tool_calls"].as_array().unwrap().len(), 2);
2783 }
2784
2785 #[test]
2786 fn test_gemma4_template_renders_reasoning_content_segments_around_tool_calls() {
2787 let formatter = make_gemma4_tool_formatter_for_tests();
2788 assert!(
2789 formatter.tool_use_template_handles_reasoning,
2790 "Gemma4 template adaptation should make reasoning_content native"
2791 );
2792
2793 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
2794 "model": "gemma4-test",
2795 "messages": [
2796 {"role": "user", "content": "inspect two things"},
2797 {
2798 "role": "assistant",
2799 "content": null,
2800 "reasoning_content": [
2801 "Think before the first call.",
2802 "Think before the second call.",
2803 "Think after both calls."
2804 ],
2805 "tool_calls": [
2806 {
2807 "id": "call_0",
2808 "type": "function",
2809 "function": {
2810 "name": "first_tool",
2811 "arguments": "{\"path\":\".\"}"
2812 }
2813 },
2814 {
2815 "id": "call_1",
2816 "type": "function",
2817 "function": {
2818 "name": "second_tool",
2819 "arguments": "{\"path\":\"/tmp\"}"
2820 }
2821 }
2822 ]
2823 }
2824 ]
2825 }))
2826 .unwrap();
2827
2828 let rendered = formatter.render(&request).unwrap();
2829
2830 let expected = concat!(
2831 "<|channel>thought\nThink before the first call.\n<channel|>",
2832 "<|tool_call>call:first_tool{path:.}<tool_call|>",
2833 "<|channel>thought\nThink before the second call.\n<channel|>",
2834 "<|tool_call>call:second_tool{path:/tmp}<tool_call|>",
2835 "<|channel>thought\nThink after both calls.\n<channel|>"
2836 );
2837 assert!(
2838 rendered.contains(expected),
2839 "Gemma4 reasoning segments should stay adjacent to their tool calls, got: {rendered}"
2840 );
2841 assert!(!rendered.contains("<think>"));
2842 assert!(!rendered.contains("reasoning_content"));
2843 }
2844
2845 #[test]
2846 fn test_gemma4_template_renders_reasoning_content_without_tool_calls() {
2847 let formatter = make_gemma4_tool_formatter_for_tests();
2848 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
2849 "model": "gemma4-test",
2850 "messages": [
2851 {"role": "user", "content": "answer directly"},
2852 {
2853 "role": "assistant",
2854 "content": "Direct answer.",
2855 "reasoning_content": "Private thought."
2856 }
2857 ]
2858 }))
2859 .unwrap();
2860
2861 let rendered = formatter.render(&request).unwrap();
2862
2863 assert!(
2864 rendered.contains("<|channel>thought\nPrivate thought.\n<channel|>Direct answer."),
2865 "Gemma4 reasoning_content should render in the thought channel, got: {rendered}"
2866 );
2867 assert!(!rendered.contains("<think>"));
2868 assert!(!rendered.contains("reasoning_content"));
2869 }
2870
2871 #[test]
2877 fn test_reasoning_flag_is_per_template_not_global() {
2878 const PLAIN_DEFAULT: &str = "{{ bos_token }}{%- for message in messages -%}\
2881 {{ message['role'] }}: {{ message['content'] }}\n{%- endfor -%}";
2882
2883 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
2884 "chat_template": [
2885 {"default": PLAIN_DEFAULT},
2886 {"tool_use": gemma4_tool_template_for_tests()},
2887 ]
2888 }))
2889 .unwrap();
2890 let formatter =
2891 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
2892
2893 assert!(
2897 formatter.tool_use_template_handles_reasoning,
2898 "adapted gemma4 tool_use template should handle reasoning natively"
2899 );
2900 assert!(
2901 !formatter.default_template_handles_reasoning,
2902 "plain default template does not reference reasoning_content"
2903 );
2904
2905 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
2908 "model": "gemma4-test",
2909 "messages": [
2910 {"role": "user", "content": "answer directly"},
2911 {
2912 "role": "assistant",
2913 "content": "Direct answer.",
2914 "reasoning_content": "Private thought."
2915 }
2916 ]
2917 }))
2918 .unwrap();
2919
2920 let rendered = formatter.render(&request).unwrap();
2921 assert!(
2922 rendered.contains("<think>Private thought.</think>Direct answer."),
2923 "reasoning must be injected on the no-tool default path, got: {rendered}"
2924 );
2925 }
2926
2927 fn reasoning_tool_call_turn(reasoning: serde_json::Value) -> serde_json::Value {
2929 let call = |id: &str, expr: &str| {
2930 json!({"id": id, "type": "function",
2931 "function": {"name": "calc", "arguments": json!({"expr": expr}).to_string()}})
2932 };
2933 json!([
2934 {"role": "user", "content": "sqrt(144) + sqrt(256)?"},
2935 {"role": "assistant", "content": null, "reasoning_content": reasoning,
2936 "tool_calls": [call("call_0", "sqrt(144)"), call("call_1", "sqrt(256)")]},
2937 {"role": "tool", "tool_call_id": "call_0", "content": "12"},
2938 {"role": "tool", "tool_call_id": "call_1", "content": "16"}
2939 ])
2940 }
2941
2942 #[test]
2946 fn test_string_reasoning_template_joins_reasoning_content_segments() {
2947 const MINIMAX_REASONING_TMPL: &str = r#"{%- for message in messages -%}
2948{%- if message.role == 'assistant' -%}
2949{{- ']~b]ai' ~ '\n' -}}
2950{%- set reasoning_content = '' -%}
2951{%- if message.reasoning_content is string -%}
2952{%- set reasoning_content = message.reasoning_content -%}
2953{%- endif -%}
2954{%- if reasoning_content -%}
2955{{- '<think>' ~ '\n' ~ reasoning_content ~ '\n' ~ '</think>' ~ '\n\n' -}}
2956{%- endif -%}
2957{%- for tool_call in message.tool_calls -%}
2958{{- '<invoke name="' ~ tool_call.function.name ~ '">' -}}
2959{%- endfor -%}
2960{{- '[e~[\n' -}}
2961{%- else -%}
2962{{- ']~b]' ~ message.role ~ '\n' ~ message.content ~ '[e~[\n' -}}
2963{%- endif -%}
2964{%- endfor -%}"#;
2965 let f = formatter_for(MINIMAX_REASONING_TMPL);
2966 let expected = concat!(
2968 "]~b]user\nsqrt(144) + sqrt(256)?[e~[\n",
2969 "]~b]ai\n<think>\nCheck both.\nThen add.\n</think>\n\n",
2970 "<invoke name=\"calc\"><invoke name=\"calc\">[e~[\n",
2971 "]~b]tool\n12[e~[\n]~b]tool\n16[e~[\n",
2972 );
2973 for render in [render_shape, render_shape_with_tools] {
2975 for reasoning in [
2976 json!(["Check both.", "Then add.", ""]),
2977 json!("Check both.\nThen add."),
2978 ] {
2979 assert_eq!(
2980 render(&f, reasoning_tool_call_turn(reasoning)).unwrap(),
2981 expected
2982 );
2983 }
2984 }
2985 }
2986
2987 #[test]
2988 fn test_inject_reasoning_content_text_variant() {
2989 let mut messages = serde_json::json!([
2990 {
2991 "role": "assistant",
2992 "content": "The answer is 42.",
2993 "reasoning_content": "Let me think about this carefully."
2994 }
2995 ]);
2996
2997 inject_reasoning_content_into_messages(&mut messages);
2998
2999 let assistant = &messages[0];
3000 assert!(assistant.get("reasoning_content").is_none());
3001 let content = assistant["content"].as_str().unwrap();
3002 assert_eq!(
3003 content,
3004 "<think>Let me think about this carefully.</think>The answer is 42."
3005 );
3006 }
3007
3008 #[test]
3009 fn test_inject_reasoning_content_null_content() {
3010 let mut messages = serde_json::json!([
3012 {
3013 "role": "assistant",
3014 "content": null,
3015 "reasoning_content": "Thinking...",
3016 "tool_calls": [{"id": "call_0", "type": "function", "function": {"name": "f", "arguments": "{}"}}]
3017 }
3018 ]);
3019
3020 inject_reasoning_content_into_messages(&mut messages);
3021
3022 let content = messages[0]["content"].as_str().unwrap();
3023 assert_eq!(content, "<think>Thinking...</think>");
3024 assert!(messages[0].get("reasoning_content").is_none());
3025 }
3026
3027 #[test]
3028 fn test_inject_reasoning_content_skips_non_assistant() {
3029 let mut messages = serde_json::json!([
3030 {
3031 "role": "user",
3032 "content": "hello",
3033 "reasoning_content": "should not be touched"
3034 }
3035 ]);
3036
3037 inject_reasoning_content_into_messages(&mut messages);
3038
3039 assert!(messages[0].get("reasoning_content").is_some());
3041 }
3042
3043 fn make_test_formatter() -> HfTokenizerConfigJsonFormatter {
3045 use super::tokcfg::ChatTemplate;
3046 use super::{ContextMixins, HfTokenizerConfigJsonFormatter};
3047
3048 let template = r#"{%- for message in messages %}{{ message.role }}: {{ message.content }}
3051{%- endfor %}
3052{%- if add_generation_prompt %}assistant:{%- endif %}"#;
3053
3054 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
3055 "chat_template": template
3056 }))
3057 .unwrap();
3058
3059 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap()
3060 }
3061
3062 #[test]
3065 fn test_reasoning_content_text_roundtrip_render() {
3066 use super::OAIPromptFormatter;
3067 let formatter = make_test_formatter();
3068
3069 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
3070 "model": "test-model",
3071 "messages": [
3072 {"role": "user", "content": "What is sqrt(144)?"},
3073 {
3074 "role": "assistant",
3075 "content": "The answer is 12.",
3076 "reasoning_content": "I need to compute the square root of 144."
3077 },
3078 {"role": "user", "content": "Are you sure?"}
3079 ]
3080 }))
3081 .unwrap();
3082
3083 let rendered = formatter.render(&request).unwrap();
3084
3085 assert!(
3086 rendered.contains("<think>I need to compute the square root of 144.</think>"),
3087 "reasoning_content must appear as <think> block, got: {}",
3088 rendered
3089 );
3090 assert!(
3091 rendered.contains("The answer is 12."),
3092 "original content must be preserved"
3093 );
3094 assert!(
3095 !rendered.contains("reasoning_content"),
3096 "raw reasoning_content field should not leak into prompt"
3097 );
3098 }
3099
3100 #[test]
3104 fn test_reasoning_content_agentic_tool_call_roundtrip_render() {
3105 use super::OAIPromptFormatter;
3106 let formatter = make_test_formatter();
3107
3108 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
3109 "model": "test-model",
3110 "messages": [
3111 {"role": "user", "content": "What is sqrt(144) + sqrt(256)?"},
3112 {
3113 "role": "assistant",
3114 "content": null,
3115 "reasoning_content": "I need to compute both square roots. Let me start with sqrt(144).",
3116 "tool_calls": [{
3117 "id": "call_0",
3118 "type": "function",
3119 "function": {
3120 "name": "calculator",
3121 "arguments": "{\"expr\": \"sqrt(144)\"}"
3122 }
3123 }]
3124 },
3125 {
3126 "role": "tool",
3127 "tool_call_id": "call_0",
3128 "content": "12"
3129 },
3130 {
3131 "role": "assistant",
3132 "content": "sqrt(144) = 12 and sqrt(256) = 16, so the answer is 28.",
3133 "reasoning_content": "Got 12 for sqrt(144). Now sqrt(256) = 16. Sum is 28."
3134 },
3135 {"role": "user", "content": "Thanks!"}
3136 ]
3137 }))
3138 .unwrap();
3139
3140 let rendered = formatter.render(&request).unwrap();
3141
3142 assert!(
3144 rendered.contains("<think>I need to compute both square roots"),
3145 "first turn reasoning must be in prompt, got: {}",
3146 rendered
3147 );
3148 assert!(
3150 rendered.contains("<think>Got 12 for sqrt(144)"),
3151 "second turn reasoning must be in prompt"
3152 );
3153 assert!(
3154 rendered.contains("the answer is 28"),
3155 "final answer content must be preserved"
3156 );
3157 assert!(
3159 !rendered.contains("reasoning_content"),
3160 "raw reasoning_content field should not leak into prompt"
3161 );
3162 }
3163
3164 #[test]
3166 fn test_reasoning_injected_when_template_ignores_it() {
3167 use super::OAIPromptFormatter;
3168 let formatter = make_test_formatter();
3169
3170 assert!(!formatter.default_template_handles_reasoning);
3172 assert!(!formatter.tool_use_template_handles_reasoning);
3173
3174 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
3175 "model": "test-model",
3176 "messages": [
3177 {"role": "user", "content": "Hello"},
3178 {
3179 "role": "assistant",
3180 "content": "Hi.",
3181 "reasoning_content": "The user said hello."
3182 },
3183 {"role": "user", "content": "Bye"}
3184 ]
3185 }))
3186 .unwrap();
3187
3188 let rendered = formatter.render(&request).unwrap();
3189 assert!(
3190 rendered.contains("<think>The user said hello.</think>"),
3191 "injection must happen when template ignores reasoning_content, got: {}",
3192 rendered
3193 );
3194 }
3195
3196 #[test]
3198 fn test_reasoning_not_injected_when_template_handles_it() {
3199 use super::tokcfg::ChatTemplate;
3200 use super::{ContextMixins, HfTokenizerConfigJsonFormatter, OAIPromptFormatter};
3201
3202 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>
3204{%- endif %}{{ message.role }}: {{ message.content }}
3205{%- endfor %}
3206{%- if add_generation_prompt %}assistant:{%- endif %}"#;
3207
3208 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
3209 "chat_template": template
3210 }))
3211 .unwrap();
3212
3213 let formatter =
3214 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
3215
3216 assert!(formatter.default_template_handles_reasoning);
3218 assert!(formatter.tool_use_template_handles_reasoning);
3219
3220 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
3221 "model": "test-model",
3222 "messages": [
3223 {"role": "user", "content": "Hello"},
3224 {
3225 "role": "assistant",
3226 "content": "Hi.",
3227 "reasoning_content": "The user said hello."
3228 },
3229 {"role": "user", "content": "Bye"}
3230 ]
3231 }))
3232 .unwrap();
3233
3234 let rendered = formatter.render(&request).unwrap();
3235
3236 assert!(
3238 rendered.contains("<think>The user said hello.</think>"),
3239 "template must render reasoning_content natively, got: {}",
3240 rendered
3241 );
3242 let think_count = rendered.matches("<think>").count();
3244 assert_eq!(
3245 think_count, 1,
3246 "must have exactly one <think> block (from template), got {} in: {}",
3247 think_count, rendered
3248 );
3249 }
3250
3251 const QWEN3_THINKING_TEMPLATE: &str = r##"{%- if tools %}
3255 {{- '<|im_start|>system\n' }}
3256 {%- if messages[0].role == 'system' %}
3257 {{- messages[0].content + '\n\n' }}
3258 {%- endif %}
3259 {{- "# 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>" }}
3260 {%- for tool in tools %}
3261 {{- "\n" }}
3262 {{- tool | tojson }}
3263 {%- endfor %}
3264 {{- "\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" }}
3265{%- else %}
3266 {%- if messages[0].role == 'system' %}
3267 {{- '<|im_start|>system\n' + messages[0].content + '<|im_end|>\n' }}
3268 {%- endif %}
3269{%- endif %}
3270{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
3271{%- for message in messages[::-1] %}
3272 {%- set index = (messages|length - 1) - loop.index0 %}
3273 {%- 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>')) %}
3274 {%- set ns.multi_step_tool = false %}
3275 {%- set ns.last_query_index = index %}
3276 {%- endif %}
3277{%- endfor %}
3278{%- for message in messages %}
3279 {%- if message.content is string %}
3280 {%- set content = message.content %}
3281 {%- else %}
3282 {%- set content = '' %}
3283 {%- endif %}
3284 {%- if (message.role == "user") or (message.role == "system" and not loop.first) %}
3285 {{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
3286 {%- elif message.role == "assistant" %}
3287 {%- set reasoning_content = '' %}
3288 {%- if message.reasoning_content is string %}
3289 {%- set reasoning_content = message.reasoning_content %}
3290 {%- else %}
3291 {%- if '</think>' in content %}
3292 {%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
3293 {%- set content = content.split('</think>')[-1].lstrip('\n') %}
3294 {%- endif %}
3295 {%- endif %}
3296 {%- if loop.index0 > ns.last_query_index %}
3297 {%- if loop.last or (not loop.last and reasoning_content) %}
3298 {{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content.strip('\n') + '\n</think>\n\n' + content.lstrip('\n') }}
3299 {%- else %}
3300 {{- '<|im_start|>' + message.role + '\n' + content }}
3301 {%- endif %}
3302 {%- else %}
3303 {{- '<|im_start|>' + message.role + '\n' + content }}
3304 {%- endif %}
3305 {%- if message.tool_calls %}
3306 {%- for tool_call in message.tool_calls %}
3307 {%- if (loop.first and content) or (not loop.first) %}
3308 {{- '\n' }}
3309 {%- endif %}
3310 {%- if tool_call.function %}
3311 {%- set tool_call = tool_call.function %}
3312 {%- endif %}
3313 {{- '<tool_call>\n{"name": "' }}
3314 {{- tool_call.name }}
3315 {{- '", "arguments": ' }}
3316 {%- if tool_call.arguments is string %}
3317 {{- tool_call.arguments }}
3318 {%- else %}
3319 {{- tool_call.arguments | tojson }}
3320 {%- endif %}
3321 {{- '}\n</tool_call>' }}
3322 {%- endfor %}
3323 {%- endif %}
3324 {{- '<|im_end|>\n' }}
3325 {%- elif message.role == "tool" %}
3326 {%- if loop.first or (messages[loop.index0 - 1].role != "tool") %}
3327 {{- '<|im_start|>user' }}
3328 {%- endif %}
3329 {{- '\n<tool_response>\n' }}
3330 {{- content }}
3331 {{- '\n</tool_response>' }}
3332 {%- if loop.last or (messages[loop.index0 + 1].role != "tool") %}
3333 {{- '<|im_end|>\n' }}
3334 {%- endif %}
3335 {%- endif %}
3336{%- endfor %}
3337{%- if add_generation_prompt %}
3338 {{- '<|im_start|>assistant\n<think>\n' }}
3339{%- endif %}"##;
3340
3341 fn qwen3_thinking_formatter() -> HfTokenizerConfigJsonFormatter {
3342 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
3343 "chat_template": QWEN3_THINKING_TEMPLATE,
3344 }))
3345 .unwrap();
3346 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap()
3347 }
3348
3349 #[test]
3350 fn test_qwen3_thinking_template_flags_detected() {
3351 let formatter = qwen3_thinking_formatter();
3352 assert!(
3353 formatter.tool_use_template_handles_reasoning,
3354 "template references reasoning_content directly"
3355 );
3356 assert!(
3359 formatter.default_template_handles_tool_calls_arguments_string,
3360 "default template branches on `arguments is string`"
3361 );
3362 assert!(
3363 formatter.tool_use_template_handles_tool_calls_arguments_string,
3364 "tool_use template branches on `arguments is string`"
3365 );
3366 }
3367
3368 const QWEN38_REJECTS_STRING_ARGS_TEMPLATE: &str = r##"{%- for message in messages %}
3372 {%- if message.role == "assistant" and message.tool_calls %}
3373 {%- for tool_call in message.tool_calls %}
3374 {%- if tool_call.function %}
3375 {%- set tool_call = tool_call.function %}
3376 {%- endif %}
3377 {{- '<tool_call>\n<function=' + tool_call.name + '>\n' }}
3378 {%- if tool_call.arguments is mapping %}
3379 {%- for args_name, args_value in tool_call.arguments|items %}
3380 {{- '<parameter=' + args_name + '>\n' + args_value + '\n</parameter>\n' }}
3381 {%- endfor %}
3382 {%- elif tool_call.arguments is string %}
3383 {%- if tool_call.arguments|trim %}
3384 {{- raise_exception('Tool call arguments were passed as a JSON string.') }}
3385 {%- endif %}
3386 {%- endif %}
3387 {{- '</function>\n</tool_call>' }}
3388 {%- endfor %}
3389 {%- else %}
3390 {{- '<|im_start|>' + message.role + '\n' + message.content + '<|im_end|>\n' }}
3391 {%- endif %}
3392{%- endfor %}"##;
3393
3394 #[test]
3398 fn test_template_rejecting_string_arguments_gets_objects() {
3399 let chat_template: ChatTemplate = serde_json::from_value(serde_json::json!({
3400 "chat_template": QWEN38_REJECTS_STRING_ARGS_TEMPLATE,
3401 }))
3402 .unwrap();
3403 let formatter =
3404 HfTokenizerConfigJsonFormatter::new(chat_template, ContextMixins::new(&[])).unwrap();
3405 assert!(!formatter.default_template_handles_tool_calls_arguments_string);
3406 assert!(!formatter.tool_use_template_handles_tool_calls_arguments_string);
3407
3408 let request: NvCreateChatCompletionRequest = serde_json::from_value(serde_json::json!({
3409 "model": "qwen3.8",
3410 "messages": [
3411 {"role": "user", "content": "What's the weather in San Francisco?"},
3412 {"role": "assistant", "content": "", "tool_calls": [{
3413 "id": "call_sf",
3414 "type": "function",
3415 "function": {"name": "get_weather", "arguments": "{\"location\": \"San Francisco\"}"}
3416 }]},
3417 {"role": "tool", "tool_call_id": "call_sf", "content": "Foggy"}
3418 ],
3419 }))
3420 .unwrap();
3421 let rendered = formatter.render(&request).unwrap();
3422 assert!(
3423 rendered.contains("<parameter=location>\nSan Francisco\n</parameter>"),
3424 "{rendered}"
3425 );
3426 }
3427
3428 #[test]
3438 fn test_qwen3_thinking_append_only_across_tool_use_turn() {
3439 let formatter = qwen3_thinking_formatter();
3440
3441 let tools = serde_json::json!([{
3442 "type": "function",
3443 "function": {
3444 "name": "get_weather",
3445 "description": "Get the current weather for a location",
3446 "parameters": {
3447 "type": "object",
3448 "properties": {
3449 "location": {"type": "string"},
3450 "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}
3451 },
3452 "required": ["location"]
3453 }
3454 }
3455 }]);
3456
3457 let turn1_request: NvCreateChatCompletionRequest =
3459 serde_json::from_value(serde_json::json!({
3460 "model": "qwen3-thinking",
3461 "messages": [
3462 {"role": "system", "content": "You are a helpful assistant."},
3463 {"role": "user", "content": "What's the weather in San Francisco?"},
3464 ],
3465 "tools": tools,
3466 }))
3467 .unwrap();
3468 let p1 = formatter.render(&turn1_request).unwrap();
3469
3470 let model_emitted = "I'll call get_weather for SF.\n\
3474 </think>\n\n\
3475 <tool_call>\n\
3476 {\"name\": \"get_weather\", \"arguments\": {\"location\": \"San Francisco\", \"unit\": \"celsius\"}}\n\
3477 </tool_call><|im_end|>\n";
3478 let wire_after_t1 = format!("{p1}{model_emitted}");
3479
3480 let turn2_request: NvCreateChatCompletionRequest =
3483 serde_json::from_value(serde_json::json!({
3484 "model": "qwen3-thinking",
3485 "messages": [
3486 {"role": "system", "content": "You are a helpful assistant."},
3487 {"role": "user", "content": "What's the weather in San Francisco?"},
3488 {
3489 "role": "assistant",
3490 "content": "",
3491 "reasoning_content": "I'll call get_weather for SF.",
3492 "tool_calls": [{
3493 "id": "call_sf",
3494 "type": "function",
3495 "function": {
3496 "name": "get_weather",
3497 "arguments": "{\"location\": \"San Francisco\", \"unit\": \"celsius\"}"
3498 }
3499 }]
3500 },
3501 {
3502 "role": "tool",
3503 "tool_call_id": "call_sf",
3504 "content": "{\"temp\": 18, \"conditions\": \"Foggy\"}"
3505 }
3506 ],
3507 "tools": tools,
3508 }))
3509 .unwrap();
3510 let p2 = formatter.render(&turn2_request).unwrap();
3511
3512 if !p2.starts_with(&wire_after_t1) {
3513 let div = wire_after_t1
3515 .as_bytes()
3516 .iter()
3517 .zip(p2.as_bytes())
3518 .position(|(a, b)| a != b)
3519 .unwrap_or_else(|| wire_after_t1.len().min(p2.len()));
3520 let lo = div.saturating_sub(40);
3521 panic!(
3522 "turn-2 prompt is NOT a prefix-extension of [turn-1 + model bytes]\n \
3523 diverges at byte {div}\n \
3524 wire ends: ...{}|{}\n \
3525 t2 has: ...{}|{}",
3526 String::from_utf8_lossy(&wire_after_t1.as_bytes()[lo..div]),
3527 String::from_utf8_lossy(
3528 &wire_after_t1.as_bytes()[div..(div + 60).min(wire_after_t1.len())]
3529 ),
3530 String::from_utf8_lossy(&p2.as_bytes()[lo..div]),
3531 String::from_utf8_lossy(&p2.as_bytes()[div..(div + 60).min(p2.len())]),
3532 );
3533 }
3534
3535 let suffix = &p2[wire_after_t1.len()..];
3538 assert!(
3539 suffix.contains("<tool_response>"),
3540 "appended bytes must include the tool response, got: {suffix}"
3541 );
3542 assert!(
3543 suffix.ends_with("<|im_start|>assistant\n<think>\n"),
3544 "appended bytes must end with the next generation prompt, got: {suffix}"
3545 );
3546 }
3547
3548 #[test]
3552 fn test_qwen3_thinking_renders_reasoning_content_segments_as_string() {
3553 let formatter = qwen3_thinking_formatter();
3554 assert!(formatter.default_template_requires_reasoning_string);
3555 assert!(formatter.tool_use_template_requires_reasoning_string);
3556
3557 for render in [render_shape, render_shape_with_tools] {
3558 let segments = json!(["Check both.", "Then add.", ""]);
3559 let rendered = render(&formatter, reasoning_tool_call_turn(segments)).unwrap();
3560 assert!(
3561 rendered.contains(
3562 "<|im_start|>assistant\n<think>\nCheck both.\nThen add.\n</think>\n\n<tool_call>"
3563 ),
3564 "{rendered}"
3565 );
3566 let string = json!("Check both.\nThen add.");
3567 assert_eq!(
3568 rendered,
3569 render(&formatter, reasoning_tool_call_turn(string)).unwrap()
3570 );
3571 }
3572 }
3573}