use anyhow::{Context, Result};
use minijinja::value::Value;
use dynamo_renderer::{
ChatTemplate, ChatTemplateValue, ContextMixins, OAIChatLikeRequest, PromptFormatter,
PromptInput, RenderedPrompt, RenderedSegment, TextInput, TokenInput, deepseek_formatter_for,
kimi_k3_formatter_for, may_be_fix_tool_schema,
};
use crate::model_card::{ModelDeploymentCard, PromptFormatterArtifact};
use crate::protocols::openai::{
chat_completions::NvCreateChatCompletionRequest, completions::NvCreateCompletionRequest,
};
pub trait MediaRequestExt {
fn media_io_kwargs(&self) -> Option<&serde_json::Value>;
}
fn parse_args_object(
s: &str,
) -> anyhow::Result<Option<serde_json::Map<String, serde_json::Value>>> {
let top: serde_json::Value = serde_json::from_str(s)?;
let top_obj = match top {
serde_json::Value::Object(m) => m,
_ => return Ok(None),
};
let mut out = serde_json::Map::with_capacity(top_obj.len());
for (k, cooked) in top_obj {
if !cooked.is_number() {
out.insert(k, cooked);
continue;
}
let cooked_str = cooked.to_string();
if let Some(raw_token) = extract_value_token(s, &k) {
let raw_trimmed = raw_token.trim();
if raw_trimmed != cooked_str.as_str() {
out.insert(k, serde_json::Value::String(raw_trimmed.to_string()));
continue;
}
}
out.insert(k, cooked);
}
Ok(Some(out))
}
fn extract_value_token<'a>(s: &'a str, key: &str) -> Option<&'a str> {
let key_json = serde_json::to_string(key).ok()?;
let needle = format!("{}:", key_json);
let start = s.find(needle.as_str())?;
let after_colon = s[start + needle.len()..].trim_start();
let end = token_end(after_colon)?;
let source_offset = s.len() - s[start + needle.len()..].len()
+ (after_colon.as_ptr() as usize - s[start + needle.len()..].as_ptr() as usize);
Some(&s[source_offset..source_offset + end])
}
fn token_end(s: &str) -> Option<usize> {
let s = s.trim_start();
let first = s.chars().next()?;
match first {
'"' => {
let mut i = 1;
let b = s.as_bytes();
while i < b.len() {
if b[i] == b'\\' {
i += 2;
} else if b[i] == b'"' {
return Some(i + 1);
} else {
i += 1;
}
}
None
}
'{' | '[' => {
let (open, close) = if first == '{' {
(b'{', b'}')
} else {
(b'[', b']')
};
let mut depth = 0i32;
let mut in_str = false;
let mut escape = false;
for (i, &b) in s.as_bytes().iter().enumerate() {
if escape {
escape = false;
continue;
}
if in_str {
if b == b'\\' {
escape = true;
} else if b == b'"' {
in_str = false;
}
} else {
match b {
b'"' => in_str = true,
b if b == open => depth += 1,
b if b == close => {
depth -= 1;
if depth == 0 {
return Some(i + 1);
}
}
_ => {}
}
}
}
None
}
_ => {
let end = s
.find(|c: char| c == ',' || c == '}' || c == ']' || c.is_whitespace())
.unwrap_or(s.len());
if end == 0 { None } else { Some(end) }
}
}
}
pub(crate) fn normalize_tool_call_arguments(
messages_json: &mut serde_json::Value,
) -> anyhow::Result<()> {
let Some(messages) = messages_json.as_array_mut() else {
return Ok(());
};
for message in messages {
let Some(tool_calls) = message
.get_mut("tool_calls")
.and_then(serde_json::Value::as_array_mut)
else {
continue;
};
for tc in tool_calls.iter_mut() {
let Some(args_str) = tc.pointer("/function/arguments").and_then(|v| v.as_str()) else {
continue;
};
if args_str.is_empty() {
if let Some(obj) = tc
.get_mut("function")
.and_then(serde_json::Value::as_object_mut)
{
obj.insert(
"arguments".to_string(),
serde_json::Value::Object(serde_json::Map::new()),
);
}
continue;
}
let value = match parse_args_object(args_str) {
Ok(Some(map)) => serde_json::Value::Object(map),
Ok(None) => {
tracing::warn!(
args_len = args_str.len(),
"tool_call arguments parsed to a non-object; \
substituting {{}} for template safety"
);
serde_json::Value::Object(serde_json::Map::new())
}
Err(e) => {
anyhow::bail!(
"tool_call arguments are not valid JSON (len={}): {e}",
args_str.len()
);
}
};
if let Some(obj) = tc
.get_mut("function")
.and_then(serde_json::Value::as_object_mut)
{
obj.insert("arguments".to_string(), value);
}
}
}
Ok(())
}
impl OAIChatLikeRequest for NvCreateChatCompletionRequest {
fn model(&self) -> String {
self.inner.model.clone()
}
fn messages(&self) -> Value {
let messages_json = serde_json::to_value(&self.inner.messages).unwrap();
Value::from_serialize(&messages_json)
}
fn typed_messages(&self) -> Option<&[dynamo_protocols::types::ChatCompletionRequestMessage]> {
Some(self.inner.messages.as_slice())
}
fn tools(&self) -> Option<Value> {
if self.inner.tools.is_none() {
None
} else {
Some(may_be_fix_tool_schema(
serde_json::to_value(&self.inner.tools).unwrap(),
)?)
}
}
fn tool_choice(&self) -> Option<Value> {
if self.inner.tool_choice.is_none() {
None
} else {
Some(Value::from_serialize(&self.inner.tool_choice))
}
}
fn response_format(&self) -> Option<Value> {
self.inner
.response_format
.as_ref()
.map(Value::from_serialize)
}
fn should_add_generation_prompt(&self) -> bool {
if self.common.continue_final_message == Some(true) {
return false;
}
self.common.add_generation_prompt.unwrap_or(true)
}
fn extract_text(&self) -> Option<TextInput> {
Some(TextInput::Single(String::new()))
}
fn chat_template_args(&self) -> Option<&std::collections::HashMap<String, serde_json::Value>> {
self.chat_template_args.as_ref()
}
fn mm_processor_kwargs(&self) -> Option<&serde_json::Value> {
self.inner.mm_processor_kwargs.as_ref()
}
}
impl MediaRequestExt for NvCreateChatCompletionRequest {
fn media_io_kwargs(&self) -> Option<&serde_json::Value> {
self.media_io_kwargs.as_ref()
}
}
impl OAIChatLikeRequest for NvCreateCompletionRequest {
fn model(&self) -> String {
self.inner.model.clone()
}
fn messages(&self) -> minijinja::value::Value {
let message = dynamo_protocols::types::ChatCompletionRequestMessage::User(
dynamo_protocols::types::ChatCompletionRequestUserMessage {
content: dynamo_protocols::types::ChatCompletionRequestUserMessageContent::Text(
crate::protocols::openai::completions::prompt_to_string(&self.inner.prompt),
),
name: None,
},
);
minijinja::value::Value::from_serialize(vec![message])
}
fn should_add_generation_prompt(&self) -> bool {
true
}
fn prompt_input_type(&self) -> PromptInput {
match &self.inner.prompt {
dynamo_protocols::types::Prompt::IntegerArray(_) => {
PromptInput::Tokens(TokenInput::Single(vec![]))
}
dynamo_protocols::types::Prompt::ArrayOfIntegerArray(_) => {
PromptInput::Tokens(TokenInput::Batch(vec![]))
}
dynamo_protocols::types::Prompt::String(_) => {
PromptInput::Text(TextInput::Single(String::new()))
}
dynamo_protocols::types::Prompt::StringArray(_) => {
PromptInput::Text(TextInput::Batch(vec![]))
}
}
}
fn extract_tokens(&self) -> Option<TokenInput> {
match &self.inner.prompt {
dynamo_protocols::types::Prompt::IntegerArray(tokens) => {
Some(TokenInput::Single(tokens.clone()))
}
dynamo_protocols::types::Prompt::ArrayOfIntegerArray(arrays) => {
Some(TokenInput::Batch(arrays.clone()))
}
_ => None,
}
}
fn extract_text(&self) -> Option<TextInput> {
match &self.inner.prompt {
dynamo_protocols::types::Prompt::String(text) => {
Some(TextInput::Single(text.to_string()))
}
dynamo_protocols::types::Prompt::StringArray(texts) => {
Some(TextInput::Batch(texts.to_vec()))
}
_ => None,
}
}
}
impl MediaRequestExt for NvCreateCompletionRequest {
fn media_io_kwargs(&self) -> Option<&serde_json::Value> {
None
}
}
pub fn prompt_formatter_from_mdc(mdc: &ModelDeploymentCard) -> Result<PromptFormatter> {
let model_type_lower = mdc
.model_info
.as_ref()
.and_then(|info| info.get_model_info().ok())
.map(|info| info.model_type().to_lowercase())
.filter(|s| !s.is_empty());
let display_name_lower = mdc.display_name.to_lowercase();
if let Some(formatter) = kimi_k3_formatter_for(
&model_type_lower,
&display_name_lower,
mdc.runtime_config.exclude_tools_when_tool_choice_none,
) {
return Ok(formatter);
}
if let Some(formatter) = deepseek_formatter_for(&model_type_lower, &display_name_lower) {
return Ok(formatter);
}
match mdc
.prompt_formatter
.as_ref()
.ok_or(anyhow::anyhow!("MDC does not contain a prompt formatter"))?
{
PromptFormatterArtifact::HfTokenizerConfigJson(checked_file) => {
let Some(file) = checked_file.path() else {
anyhow::bail!(
"HfTokenizerConfigJson for {} is a URL, cannot load",
mdc.display_name
);
};
let contents = std::fs::read_to_string(file).with_context(|| {
format!(
"prompt_formatter_from_mdc fs:read_to_string '{}'",
file.display()
)
})?;
let mut config: ChatTemplate = serde_json::from_str(&contents).inspect_err(|err| {
crate::log_json_err(&file.display().to_string(), &contents, err)
})?;
match mdc.chat_template_file.as_ref() {
Some(PromptFormatterArtifact::HfChatTemplateJinja {
file: checked_file, ..
}) => {
let Some(path) = checked_file.path() else {
anyhow::bail!(
"HfChatTemplateJinja for {} is a URL, cannot load",
mdc.display_name
);
};
let chat_template = std::fs::read_to_string(path)
.with_context(|| format!("fs:read_to_string '{}'", path.display()))?;
config.chat_template = Some(ChatTemplateValue(either::Left(chat_template)));
}
Some(PromptFormatterArtifact::HfChatTemplateJson {
file: checked_file, ..
}) => {
let Some(path) = checked_file.path() else {
anyhow::bail!(
"HfChatTemplateJson for {} is a URL, cannot load",
mdc.display_name
);
};
let raw = std::fs::read_to_string(path)
.with_context(|| format!("fs:read_to_string '{}'", path.display()))?;
let wrapper: serde_json::Value = serde_json::from_str(&raw)
.with_context(|| format!("Failed to parse '{}' as JSON", path.display()))?;
let field = wrapper.get("chat_template").ok_or_else(|| {
anyhow::anyhow!(
"'{}' does not contain a 'chat_template' field",
path.display()
)
})?;
let value = serde_json::from_value::<ChatTemplateValue>(field.clone())
.with_context(|| {
format!(
"Failed to deserialize 'chat_template' in '{}'",
path.display()
)
})?;
config.chat_template = Some(value);
}
_ => {}
}
PromptFormatter::from_parts(
config,
mdc.prompt_context
.clone()
.map_or(ContextMixins::default(), |x| ContextMixins::new(&x)),
mdc.runtime_config.exclude_tools_when_tool_choice_none,
)
}
PromptFormatterArtifact::HfChatTemplateJinja { .. }
| PromptFormatterArtifact::HfChatTemplateJson { .. } => Err(anyhow::anyhow!(
"prompt_formatter should not have type HfChatTemplate*"
)),
}
}
const CONTINUE_FINAL_MESSAGE_NOT_FOUND: &str =
"Unable to continue the final message because it was not found in the rendered chat template.";
pub(crate) const CONTINUE_FINAL_MESSAGE_TAG: &str = "CONTINUE_FINAL_MESSAGE_TAG ";
pub(crate) fn append_continue_final_message_tag(messages: &mut serde_json::Value) -> Result<()> {
let Some(last) = messages.as_array_mut().and_then(|arr| arr.last_mut()) else {
anyhow::bail!(CONTINUE_FINAL_MESSAGE_NOT_FOUND);
};
match last.get_mut("content") {
Some(serde_json::Value::String(text)) if !text.is_empty() => {
text.push_str(CONTINUE_FINAL_MESSAGE_TAG);
Ok(())
}
Some(serde_json::Value::Array(parts)) => {
for part in parts.iter_mut().rev() {
if let Some(serde_json::Value::String(text)) = part.get_mut("text") {
if text.is_empty() {
continue;
}
text.push_str(CONTINUE_FINAL_MESSAGE_TAG);
return Ok(());
}
}
anyhow::bail!(CONTINUE_FINAL_MESSAGE_NOT_FOUND);
}
_ => anyhow::bail!(CONTINUE_FINAL_MESSAGE_NOT_FOUND),
}
}
pub(crate) fn apply_continue_final_message(prompt: RenderedPrompt) -> Result<RenderedPrompt> {
let rendered = prompt.as_str();
let tag = CONTINUE_FINAL_MESSAGE_TAG;
let tag_name = tag.trim_end();
let Some(tag_loc) = rendered.rfind(tag_name) else {
anyhow::bail!(CONTINUE_FINAL_MESSAGE_NOT_FOUND);
};
let end = if rendered[tag_loc..].starts_with(tag) {
tag_loc
} else {
rendered[..tag_loc].trim_end().len()
};
Ok(truncate_rendered_prompt(prompt, end))
}
fn truncate_rendered_prompt(prompt: RenderedPrompt, end: usize) -> RenderedPrompt {
let Some(segments) = prompt.segments() else {
return RenderedPrompt::text(prompt.as_str()[..end].to_string());
};
let mut out = Vec::new();
let mut offset = 0usize;
for seg in segments {
let next = offset + seg.text.len();
if next <= end {
if !seg.text.is_empty() {
out.push(seg.clone());
}
offset = next;
if offset == end {
break;
}
continue;
}
if offset < end {
let keep = end - offset;
if keep > 0 && keep <= seg.text.len() && seg.text.is_char_boundary(keep) {
out.push(RenderedSegment {
text: seg.text[..keep].to_string(),
allow_special: seg.allow_special,
});
}
}
break;
}
if out.is_empty() {
RenderedPrompt::text(String::new())
} else {
RenderedPrompt::segmented(out)
}
}
#[cfg(test)]
mod tests {
use super::normalize_tool_call_arguments;
fn make_tool_call_messages(arguments: &str) -> serde_json::Value {
serde_json::json!([{
"role": "assistant",
"tool_calls": [{
"id": "call_1",
"type": "function",
"function": { "name": "f", "arguments": arguments }
}]
}])
}
#[test]
fn normalize_preserves_large_integer() {
let big_int = "123456789012345678901234567890";
let mut msgs = make_tool_call_messages(&format!("{{\"id\": {big_int}}}"));
normalize_tool_call_arguments(&mut msgs).unwrap();
let id_slot = &msgs[0]["tool_calls"][0]["function"]["arguments"]["id"];
let preserved = id_slot
.as_str()
.map(|s| s.to_string())
.unwrap_or_else(|| id_slot.to_string());
assert_eq!(
preserved, big_int,
"large integer must not be corrupted by f64 round-trip"
);
}
#[test]
fn normalize_empty_string_becomes_empty_object() {
let mut msgs = make_tool_call_messages("");
normalize_tool_call_arguments(&mut msgs).unwrap();
assert_eq!(
msgs[0]["tool_calls"][0]["function"]["arguments"],
serde_json::json!({}),
"empty arguments string must normalise to an empty object"
);
}
#[test]
fn normalize_malformed_json_returns_error() {
let mut msgs = make_tool_call_messages("not-valid-json{");
let result = normalize_tool_call_arguments(&mut msgs);
assert!(
result.is_err(),
"malformed arguments JSON must return Err, not substitute {{}}"
);
}
#[test]
fn normalize_valid_object_passes_through() {
let mut msgs = make_tool_call_messages(r#"{"city": "Paris", "unit": "celsius"}"#);
normalize_tool_call_arguments(&mut msgs).unwrap();
assert_eq!(
msgs[0]["tool_calls"][0]["function"]["arguments"]["city"],
serde_json::json!("Paris")
);
}
#[test]
fn glm52_historical_arguments_are_object() {
let mut msgs =
make_tool_call_messages(r#"{"location": "San Francisco", "unit": "celsius"}"#);
normalize_tool_call_arguments(&mut msgs).unwrap();
let args = &msgs[0]["tool_calls"][0]["function"]["arguments"];
assert!(
args.is_object(),
"GLM-5.2 (JsonObject mode): arguments must be a JSON object, got: {args}"
);
assert_eq!(args["location"], serde_json::json!("San Francisco"));
}
#[test]
fn gptoss_historical_arguments_remain_string() {
let msgs = make_tool_call_messages(r#"{"location": "San Francisco", "unit": "celsius"}"#);
let args = &msgs[0]["tool_calls"][0]["function"]["arguments"];
assert!(
args.is_string(),
"GPT-OSS (JsonString mode): arguments must remain a JSON string, got: {args}"
);
assert_eq!(
args.as_str().unwrap(),
r#"{"location": "San Francisco", "unit": "celsius"}"#
);
}
fn continue_rendered(
prompt: dynamo_renderer::RenderedPrompt,
) -> dynamo_renderer::RenderedPrompt {
super::apply_continue_final_message(prompt).unwrap()
}
#[test]
fn continue_final_message_truncates_after_last_assistant_text() {
use dynamo_renderer::RenderedPrompt;
let rendered = continue_rendered(RenderedPrompt::text(format!(
"user text<|im_end|>LLM-Native Interaction{}<|im_end|><|im_start|>assistant",
super::CONTINUE_FINAL_MESSAGE_TAG
)));
assert_eq!(
rendered.as_str(),
"user text<|im_end|>LLM-Native Interaction"
);
}
#[test]
fn continue_final_message_rstrips_when_template_trims_tag_space() {
use dynamo_renderer::RenderedPrompt;
let rendered = continue_rendered(RenderedPrompt::text(format!(
"hello world{}extra",
super::CONTINUE_FINAL_MESSAGE_TAG.trim_end()
)));
assert_eq!(rendered.as_str(), "hello world");
}
#[test]
fn continue_final_message_marker_survives_repeated_last_text() {
use dynamo_renderer::RenderedPrompt;
let rendered = continue_rendered(RenderedPrompt::text(format!(
"LLM-Native Interaction in the user turn. LLM-Native Interaction{}<|im_end|>",
super::CONTINUE_FINAL_MESSAGE_TAG
)));
assert_eq!(
rendered.as_str(),
"LLM-Native Interaction in the user turn. LLM-Native Interaction"
);
}
#[test]
fn continue_final_message_appends_tag_to_last_array_text_part() {
let mut messages = serde_json::json!([{
"role": "assistant",
"content": [
{"type": "text", "text": "ignored"},
{"type": "text", "text": "Design"}
]
}]);
super::append_continue_final_message_tag(&mut messages).unwrap();
assert_eq!(
messages[0]["content"][1]["text"].as_str().unwrap(),
format!("Design{}", super::CONTINUE_FINAL_MESSAGE_TAG)
);
assert_eq!(messages[0]["content"][0]["text"], "ignored");
}
#[test]
fn continue_final_message_appends_tag_to_developer_array_text_part() {
let mut messages = serde_json::json!([{
"role": "developer",
"content": [
{"type": "text", "text": "ignored"},
{"type": "text", "text": "Design"}
]
}]);
super::append_continue_final_message_tag(&mut messages).unwrap();
assert_eq!(
messages[0]["content"][1]["text"].as_str().unwrap(),
format!("Design{}", super::CONTINUE_FINAL_MESSAGE_TAG)
);
}
#[test]
fn continue_final_message_continues_final_user_message() {
use dynamo_renderer::RenderedPrompt;
let rendered = continue_rendered(RenderedPrompt::text(format!(
"hello world{}extra",
super::CONTINUE_FINAL_MESSAGE_TAG
)));
assert_eq!(rendered.as_str(), "hello world");
}
#[test]
fn continue_final_message_marker_survives_rewritten_final_turn() {
use dynamo_renderer::RenderedPrompt;
let rendered = continue_rendered(RenderedPrompt::text(format!(
"same|PREVIOUS|SAME{}|closed",
super::CONTINUE_FINAL_MESSAGE_TAG
)));
assert_eq!(rendered.as_str(), "same|PREVIOUS|SAME");
}
#[test]
fn continue_final_message_whitespace_only_does_not_match_structural_spaces() {
use dynamo_renderer::RenderedPrompt;
let mut messages = serde_json::json!([
{"role": "user", "content": "hello"},
{"role": "assistant", "content": " "}
]);
super::append_continue_final_message_tag(&mut messages).unwrap();
assert_eq!(
messages[1]["content"].as_str().unwrap(),
format!(" {}", super::CONTINUE_FINAL_MESSAGE_TAG)
);
let rendered = continue_rendered(RenderedPrompt::text(format!(
"hello world {}extra",
super::CONTINUE_FINAL_MESSAGE_TAG
)));
assert_eq!(rendered.as_str(), "hello world ");
}
#[test]
fn continue_final_message_errors_when_marker_is_missing() {
use dynamo_renderer::RenderedPrompt;
let err =
super::apply_continue_final_message(RenderedPrompt::text("unchanged".to_string()))
.unwrap_err();
assert!(
err.to_string()
.contains("not found in the rendered chat template"),
"unexpected error: {err}"
);
}
#[test]
fn continue_final_message_errors_when_final_content_is_empty() {
let mut messages = serde_json::json!([{"role": "assistant", "content": ""}]);
let err = super::append_continue_final_message_tag(&mut messages).unwrap_err();
assert!(
err.to_string()
.contains("not found in the rendered chat template"),
"unexpected error: {err}"
);
}
#[test]
fn continue_final_message_preserves_segment_boundaries() {
use dynamo_renderer::{RenderedPrompt, RenderedSegment};
let rendered = continue_rendered(RenderedPrompt::segmented(vec![
RenderedSegment {
text: "<|im_start|>assistant\n".to_string(),
allow_special: true,
},
RenderedSegment {
text: format!(
"LLM-Native Interaction{}",
super::CONTINUE_FINAL_MESSAGE_TAG
),
allow_special: false,
},
RenderedSegment {
text: "<|im_end|>".to_string(),
allow_special: true,
},
RenderedSegment {
text: "<|im_start|>assistant\n".to_string(),
allow_special: true,
},
]));
assert_eq!(
rendered.as_str(),
"<|im_start|>assistant\nLLM-Native Interaction"
);
let segments = rendered
.segments()
.expect("Kimi-style prompts must keep segment boundaries");
assert_eq!(segments.len(), 2);
assert!(segments[0].allow_special);
assert_eq!(segments[0].text, "<|im_start|>assistant\n");
assert!(!segments[1].allow_special);
assert_eq!(segments[1].text, "LLM-Native Interaction");
}
#[test]
fn continue_final_message_truncates_inside_an_ordinary_segment() {
use dynamo_renderer::{RenderedPrompt, RenderedSegment};
let rendered = continue_rendered(RenderedPrompt::segmented(vec![
RenderedSegment {
text: "<ctrl>".to_string(),
allow_special: true,
},
RenderedSegment {
text: format!("hello world{}extra", super::CONTINUE_FINAL_MESSAGE_TAG),
allow_special: false,
},
]));
assert_eq!(rendered.as_str(), "<ctrl>hello world");
let segments = rendered
.segments()
.expect("truncated prompt stays segmented");
assert_eq!(segments.len(), 2);
assert!(segments[0].allow_special);
assert_eq!(segments[1].text, "hello world");
assert!(!segments[1].allow_special);
}
}