1use anyhow::{Context, Result};
28use minijinja::value::Value;
29use std::collections::HashMap;
30use std::sync::Arc;
31
32pub use dynamo_tokenizers;
36
37pub mod deepseek;
38pub mod inkling;
39pub mod kimi_k3;
40mod template;
41
42pub use template::{
43 ChatTemplate, ChatTemplateValue, ContextMixins, deepseek_formatter_for, kimi_k3_formatter_for,
44 may_be_fix_tool_schema, native_formatter_for,
45};
46
47#[derive(serde::Serialize, serde::Deserialize, Clone, Debug, PartialEq, Eq, Hash)]
52#[serde(rename_all = "snake_case")]
53pub enum PromptContextMixin {
54 OaiChat,
56
57 Llama3DateTime,
59}
60
61pub fn thinking_bool_from_args(args: Option<&HashMap<String, serde_json::Value>>) -> Option<bool> {
69 let args = args?;
70 for key in ["thinking", "enable_thinking"] {
71 if let Some(v) = args.get(key).and_then(|x| x.as_bool()) {
72 return Some(v);
73 }
74 }
75 None
76}
77
78#[derive(Debug)]
79pub enum TokenInput {
80 Single(Vec<u32>),
81 Batch(Vec<Vec<u32>>),
82}
83
84#[derive(Debug)]
85pub enum TextInput {
86 Single(String),
87 Batch(Vec<String>),
88}
89
90#[derive(Debug)]
91pub enum PromptInput {
92 Tokens(TokenInput),
93 Text(TextInput),
94}
95
96#[derive(Debug, Clone, PartialEq, Eq)]
98pub struct RenderedSegment {
99 pub text: String,
100 pub allow_special: bool,
101}
102
103impl RenderedSegment {
104 pub fn new(text: impl Into<String>, allow_special: bool) -> Self {
105 Self {
106 text: text.into(),
107 allow_special,
108 }
109 }
110
111 pub fn as_encode_segment(&self) -> dynamo_tokenizers::EncodeSegment<'_> {
112 dynamo_tokenizers::EncodeSegment::new(&self.text, self.allow_special)
113 }
114}
115
116#[derive(Debug, Clone, PartialEq, Eq)]
122pub struct RenderedPrompt {
123 text: String,
124 segments: Option<Vec<RenderedSegment>>,
125}
126
127impl RenderedPrompt {
128 pub fn text(text: String) -> Self {
129 Self {
130 text,
131 segments: None,
132 }
133 }
134
135 pub fn segmented(segments: Vec<RenderedSegment>) -> Self {
136 let text = segments
137 .iter()
138 .map(|segment| segment.text.as_str())
139 .collect();
140 Self {
141 text,
142 segments: Some(segments),
143 }
144 }
145
146 pub fn as_str(&self) -> &str {
147 &self.text
148 }
149
150 pub fn segments(&self) -> Option<&[RenderedSegment]> {
151 self.segments.as_deref()
152 }
153
154 pub fn encode_segments(&self) -> Option<Vec<dynamo_tokenizers::EncodeSegment<'_>>> {
155 Some(
156 self.segments()?
157 .iter()
158 .map(RenderedSegment::as_encode_segment)
159 .collect(),
160 )
161 }
162
163 pub fn into_text(self) -> String {
164 self.text
165 }
166}
167
168#[derive(Debug, Clone, PartialEq, Eq)]
174pub enum PromptRenderError {
175 InvalidRequest(String),
176}
177
178impl PromptRenderError {
179 pub fn invalid_request(message: impl Into<String>) -> Self {
180 Self::InvalidRequest(message.into())
181 }
182}
183
184impl std::fmt::Display for PromptRenderError {
185 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
186 match self {
187 Self::InvalidRequest(message) => f.write_str(message),
188 }
189 }
190}
191
192impl std::error::Error for PromptRenderError {}
193
194pub trait OAIChatLikeRequest {
201 fn model(&self) -> String;
202 fn messages(&self) -> Value;
203 fn typed_messages(&self) -> Option<&[dynamo_protocols::types::ChatCompletionRequestMessage]> {
209 None
210 }
211 fn tools(&self) -> Option<Value> {
212 None
213 }
214 fn tool_choice(&self) -> Option<Value> {
215 None
216 }
217 fn response_format(&self) -> Option<Value> {
218 None
219 }
220
221 fn reasoning_effort(&self) -> Option<Value> {
224 None
225 }
226
227 fn should_add_generation_prompt(&self) -> bool;
228
229 fn chat_template_args(&self) -> Option<&HashMap<String, serde_json::Value>> {
231 None
232 }
233
234 fn prompt_input_type(&self) -> PromptInput {
236 PromptInput::Text(TextInput::Single(String::new()))
237 }
238
239 fn extract_tokens(&self) -> Option<TokenInput> {
241 None
242 }
243
244 fn extract_text(&self) -> Option<TextInput> {
245 None
246 }
247
248 fn mm_processor_kwargs(&self) -> Option<&serde_json::Value> {
249 None
250 }
251}
252
253pub(crate) fn messages_to_json(req: &dyn OAIChatLikeRequest) -> Result<serde_json::Value> {
256 if let Some(messages) = req.typed_messages() {
257 serde_json::to_value(messages)
258 } else {
259 serde_json::to_value(req.messages())
260 }
261 .context("Failed to convert messages to JSON")
262}
263
264pub trait OAIPromptFormatter: Send + Sync + 'static {
265 fn supports_add_generation_prompt(&self) -> bool;
266 fn render(&self, req: &dyn OAIChatLikeRequest) -> Result<String>;
267
268 fn media_message_order(&self, _request: &dyn OAIChatLikeRequest) -> Option<Vec<usize>> {
272 None
273 }
274
275 fn render_prompt(&self, req: &dyn OAIChatLikeRequest) -> Result<RenderedPrompt> {
276 self.render(req).map(RenderedPrompt::text)
277 }
278}
279
280pub(crate) fn reject_unsupported_partial_assistant(messages: &serde_json::Value) -> Result<()> {
287 let has_partial =
288 messages.as_array().into_iter().flatten().any(|message| {
289 message.get("partial").and_then(serde_json::Value::as_bool) == Some(true)
290 });
291 if has_partial {
292 return Err(PromptRenderError::invalid_request(
293 "assistant `partial: true` is not supported by this model's prompt formatter",
294 )
295 .into());
296 }
297 Ok(())
298}
299
300pub(crate) fn reject_unsupported_message_tools(
303 messages: &serde_json::Value,
304 supported_tool_roles: &[&str],
305) -> Result<()> {
306 let offending = messages.as_array().into_iter().flatten().find(|message| {
307 let declares_tools = message
308 .get("tools")
309 .is_some_and(|tools| !tools.is_null() && !tools.as_array().is_some_and(Vec::is_empty));
310 let role_is_supported = message
311 .get("role")
312 .and_then(serde_json::Value::as_str)
313 .is_some_and(|role| supported_tool_roles.contains(&role));
314 declares_tools && !role_is_supported
315 });
316
317 if let Some(message) = offending {
318 let role = message
319 .get("role")
320 .and_then(serde_json::Value::as_str)
321 .unwrap_or("<missing>");
322 return Err(PromptRenderError::invalid_request(format!(
323 "message-level `tools` on role {role:?} are not supported by this model's prompt \
324 formatter"
325 ))
326 .into());
327 }
328 Ok(())
329}
330
331#[derive(Clone)]
332pub enum PromptFormatter {
333 OAI(Arc<dyn OAIPromptFormatter>),
334}
335
336#[derive(Debug, Default)]
338pub struct NoOpFormatter;
339
340impl OAIPromptFormatter for NoOpFormatter {
341 fn supports_add_generation_prompt(&self) -> bool {
342 false
343 }
344
345 fn render(&self, req: &dyn OAIChatLikeRequest) -> Result<String> {
346 let messages = req.messages();
347 let messages_json = serde_json::to_value(&messages)?;
348 reject_unsupported_partial_assistant(&messages_json)?;
349 reject_unsupported_message_tools(&messages_json, &[])?;
350
351 let first_message = messages
352 .get_item_by_index(0)
353 .map_err(|_| anyhow::Error::msg("No message at index 0 or messages array is empty"))?;
354
355 let content = first_message
356 .get_attr("content")
357 .map_err(|_| anyhow::Error::msg("First message has no 'content' field"))?;
358
359 let content_str = content
360 .as_str()
361 .ok_or_else(|| anyhow::Error::msg("Message content is not a string"))?
362 .to_string();
363 Ok(content_str)
364 }
365}
366
367impl PromptFormatter {
368 pub fn no_op() -> Self {
369 Self::OAI(Arc::new(NoOpFormatter))
370 }
371}
372
373#[cfg(test)]
374mod rendered_prompt_tests {
375 use super::{
376 NoOpFormatter, OAIPromptFormatter, PromptRenderError, RenderedPrompt, RenderedSegment,
377 };
378
379 #[test]
380 fn owned_segments_borrow_into_tokenizer_segments() {
381 let prompt = RenderedPrompt::segmented(vec![
382 RenderedSegment::new("<|open|>", true),
383 RenderedSegment::new("user text", false),
384 ]);
385
386 let segments = prompt.encode_segments().expect("segmented prompt");
387 assert_eq!(segments[0].text, "<|open|>");
388 assert!(segments[0].allow_special);
389 assert_eq!(segments[1].text, "user text");
390 assert!(!segments[1].allow_special);
391 assert_eq!(prompt.as_str(), "<|open|>user text");
392 }
393
394 #[test]
395 fn no_op_formatter_rejects_unsupported_partial_assistant() {
396 let request: dynamo_protocols::types::CreateChatCompletionRequest =
397 serde_json::from_value(serde_json::json!({
398 "model": "test",
399 "messages": [
400 {"role": "user", "content": "Continue"},
401 {"role": "assistant", "content": "prefix", "partial": true}
402 ]
403 }))
404 .unwrap();
405
406 let error = NoOpFormatter.render(&request).unwrap_err();
407 assert!(matches!(
408 error.downcast_ref::<PromptRenderError>(),
409 Some(PromptRenderError::InvalidRequest(message))
410 if message.contains("`partial: true` is not supported")
411 ));
412 }
413
414 #[test]
415 fn no_op_formatter_rejects_message_level_tools() {
416 let request: dynamo_protocols::types::CreateChatCompletionRequest =
417 serde_json::from_value(serde_json::json!({
418 "model": "test",
419 "messages": [
420 {"role": "system", "tools": [{"name": "lookup"}]},
421 {"role": "user", "content": "Continue"}
422 ]
423 }))
424 .unwrap();
425
426 let error = NoOpFormatter.render(&request).unwrap_err();
427 assert!(matches!(
428 error.downcast_ref::<PromptRenderError>(),
429 Some(PromptRenderError::InvalidRequest(message))
430 if message.contains("message-level `tools`")
431 ));
432 }
433}