use anyhow::Result;
use minijinja::value::Value;
use std::collections::HashMap;
use std::sync::Arc;
pub use dynamo_tokenizers;
pub mod deepseek;
pub mod inkling;
pub mod kimi_k3;
mod template;
pub use template::{
ChatTemplate, ChatTemplateValue, ContextMixins, deepseek_formatter_for, kimi_k3_formatter_for,
may_be_fix_tool_schema, native_formatter_for,
};
#[derive(serde::Serialize, serde::Deserialize, Clone, Debug, PartialEq, Eq, Hash)]
#[serde(rename_all = "snake_case")]
pub enum PromptContextMixin {
OaiChat,
Llama3DateTime,
}
pub fn thinking_bool_from_args(args: Option<&HashMap<String, serde_json::Value>>) -> Option<bool> {
let args = args?;
for key in ["thinking", "enable_thinking"] {
if let Some(v) = args.get(key).and_then(|x| x.as_bool()) {
return Some(v);
}
}
None
}
#[derive(Debug)]
pub enum TokenInput {
Single(Vec<u32>),
Batch(Vec<Vec<u32>>),
}
#[derive(Debug)]
pub enum TextInput {
Single(String),
Batch(Vec<String>),
}
#[derive(Debug)]
pub enum PromptInput {
Tokens(TokenInput),
Text(TextInput),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RenderedSegment {
pub text: String,
pub allow_special: bool,
}
impl RenderedSegment {
pub fn new(text: impl Into<String>, allow_special: bool) -> Self {
Self {
text: text.into(),
allow_special,
}
}
pub fn as_encode_segment(&self) -> dynamo_tokenizers::EncodeSegment<'_> {
dynamo_tokenizers::EncodeSegment::new(&self.text, self.allow_special)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RenderedPrompt {
text: String,
segments: Option<Vec<RenderedSegment>>,
}
impl RenderedPrompt {
pub fn text(text: String) -> Self {
Self {
text,
segments: None,
}
}
pub fn segmented(segments: Vec<RenderedSegment>) -> Self {
let text = segments
.iter()
.map(|segment| segment.text.as_str())
.collect();
Self {
text,
segments: Some(segments),
}
}
pub fn as_str(&self) -> &str {
&self.text
}
pub fn segments(&self) -> Option<&[RenderedSegment]> {
self.segments.as_deref()
}
pub fn encode_segments(&self) -> Option<Vec<dynamo_tokenizers::EncodeSegment<'_>>> {
Some(
self.segments()?
.iter()
.map(RenderedSegment::as_encode_segment)
.collect(),
)
}
pub fn into_text(self) -> String {
self.text
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PromptRenderError {
InvalidRequest(String),
}
impl PromptRenderError {
pub fn invalid_request(message: impl Into<String>) -> Self {
Self::InvalidRequest(message.into())
}
}
impl std::fmt::Display for PromptRenderError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::InvalidRequest(message) => f.write_str(message),
}
}
}
impl std::error::Error for PromptRenderError {}
pub trait OAIChatLikeRequest {
fn model(&self) -> String;
fn messages(&self) -> Value;
fn typed_messages(&self) -> Option<&[dynamo_protocols::types::ChatCompletionRequestMessage]> {
None
}
fn tools(&self) -> Option<Value> {
None
}
fn tool_choice(&self) -> Option<Value> {
None
}
fn response_format(&self) -> Option<Value> {
None
}
fn reasoning_effort(&self) -> Option<Value> {
None
}
fn should_add_generation_prompt(&self) -> bool;
fn chat_template_args(&self) -> Option<&HashMap<String, serde_json::Value>> {
None
}
fn prompt_input_type(&self) -> PromptInput {
PromptInput::Text(TextInput::Single(String::new()))
}
fn extract_tokens(&self) -> Option<TokenInput> {
None
}
fn extract_text(&self) -> Option<TextInput> {
None
}
fn mm_processor_kwargs(&self) -> Option<&serde_json::Value> {
None
}
}
pub trait OAIPromptFormatter: Send + Sync + 'static {
fn supports_add_generation_prompt(&self) -> bool;
fn render(&self, req: &dyn OAIChatLikeRequest) -> Result<String>;
fn render_prompt(&self, req: &dyn OAIChatLikeRequest) -> Result<RenderedPrompt> {
self.render(req).map(RenderedPrompt::text)
}
}
pub(crate) fn reject_unsupported_partial_assistant(messages: &serde_json::Value) -> Result<()> {
let has_partial =
messages.as_array().into_iter().flatten().any(|message| {
message.get("partial").and_then(serde_json::Value::as_bool) == Some(true)
});
if has_partial {
return Err(PromptRenderError::invalid_request(
"assistant `partial: true` is not supported by this model's prompt formatter",
)
.into());
}
Ok(())
}
pub(crate) fn reject_unsupported_message_tools(
messages: &serde_json::Value,
supported_tool_roles: &[&str],
) -> Result<()> {
let offending = messages.as_array().into_iter().flatten().find(|message| {
let declares_tools = message
.get("tools")
.is_some_and(|tools| !tools.is_null() && !tools.as_array().is_some_and(Vec::is_empty));
let role_is_supported = message
.get("role")
.and_then(serde_json::Value::as_str)
.is_some_and(|role| supported_tool_roles.contains(&role));
declares_tools && !role_is_supported
});
if let Some(message) = offending {
let role = message
.get("role")
.and_then(serde_json::Value::as_str)
.unwrap_or("<missing>");
return Err(PromptRenderError::invalid_request(format!(
"message-level `tools` on role {role:?} are not supported by this model's prompt \
formatter"
))
.into());
}
Ok(())
}
#[derive(Clone)]
pub enum PromptFormatter {
OAI(Arc<dyn OAIPromptFormatter>),
}
#[derive(Debug, Default)]
pub struct NoOpFormatter;
impl OAIPromptFormatter for NoOpFormatter {
fn supports_add_generation_prompt(&self) -> bool {
false
}
fn render(&self, req: &dyn OAIChatLikeRequest) -> Result<String> {
let messages = req.messages();
let messages_json = serde_json::to_value(&messages)?;
reject_unsupported_partial_assistant(&messages_json)?;
reject_unsupported_message_tools(&messages_json, &[])?;
let first_message = messages
.get_item_by_index(0)
.map_err(|_| anyhow::Error::msg("No message at index 0 or messages array is empty"))?;
let content = first_message
.get_attr("content")
.map_err(|_| anyhow::Error::msg("First message has no 'content' field"))?;
let content_str = content
.as_str()
.ok_or_else(|| anyhow::Error::msg("Message content is not a string"))?
.to_string();
Ok(content_str)
}
}
impl PromptFormatter {
pub fn no_op() -> Self {
Self::OAI(Arc::new(NoOpFormatter))
}
}
#[cfg(test)]
mod rendered_prompt_tests {
use super::{
NoOpFormatter, OAIPromptFormatter, PromptRenderError, RenderedPrompt, RenderedSegment,
};
#[test]
fn owned_segments_borrow_into_tokenizer_segments() {
let prompt = RenderedPrompt::segmented(vec![
RenderedSegment::new("<|open|>", true),
RenderedSegment::new("user text", false),
]);
let segments = prompt.encode_segments().expect("segmented prompt");
assert_eq!(segments[0].text, "<|open|>");
assert!(segments[0].allow_special);
assert_eq!(segments[1].text, "user text");
assert!(!segments[1].allow_special);
assert_eq!(prompt.as_str(), "<|open|>user text");
}
#[test]
fn no_op_formatter_rejects_unsupported_partial_assistant() {
let request: dynamo_protocols::types::CreateChatCompletionRequest =
serde_json::from_value(serde_json::json!({
"model": "test",
"messages": [
{"role": "user", "content": "Continue"},
{"role": "assistant", "content": "prefix", "partial": true}
]
}))
.unwrap();
let error = NoOpFormatter.render(&request).unwrap_err();
assert!(matches!(
error.downcast_ref::<PromptRenderError>(),
Some(PromptRenderError::InvalidRequest(message))
if message.contains("`partial: true` is not supported")
));
}
#[test]
fn no_op_formatter_rejects_message_level_tools() {
let request: dynamo_protocols::types::CreateChatCompletionRequest =
serde_json::from_value(serde_json::json!({
"model": "test",
"messages": [
{"role": "system", "tools": [{"name": "lookup"}]},
{"role": "user", "content": "Continue"}
]
}))
.unwrap();
let error = NoOpFormatter.render(&request).unwrap_err();
assert!(matches!(
error.downcast_ref::<PromptRenderError>(),
Some(PromptRenderError::InvalidRequest(message))
if message.contains("message-level `tools`")
));
}
}