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 python;
41mod template;
42
43pub use template::{
44 ChatTemplate, ChatTemplateValue, ContextMixins, deepseek_formatter_for, kimi_k3_formatter_for,
45 may_be_fix_tool_schema, native_formatter_for,
46};
47
48#[derive(serde::Serialize, serde::Deserialize, Clone, Debug, PartialEq, Eq, Hash)]
53#[serde(rename_all = "snake_case")]
54pub enum PromptContextMixin {
55 OaiChat,
57
58 Llama3DateTime,
60}
61
62pub fn thinking_bool_from_args(args: Option<&HashMap<String, serde_json::Value>>) -> Option<bool> {
70 let args = args?;
71 for key in ["thinking", "enable_thinking"] {
72 if let Some(v) = args.get(key).and_then(|x| x.as_bool()) {
73 return Some(v);
74 }
75 }
76 None
77}
78
79#[derive(Debug)]
80pub enum TokenInput {
81 Single(Vec<u32>),
82 Batch(Vec<Vec<u32>>),
83}
84
85#[derive(Debug)]
86pub enum TextInput {
87 Single(String),
88 Batch(Vec<String>),
89}
90
91#[derive(Debug)]
92pub enum PromptInput {
93 Tokens(TokenInput),
94 Text(TextInput),
95}
96
97#[derive(Debug, Clone, PartialEq, Eq)]
99pub struct RenderedSegment {
100 pub text: String,
101 pub allow_special: bool,
102}
103
104impl RenderedSegment {
105 pub fn new(text: impl Into<String>, allow_special: bool) -> Self {
106 Self {
107 text: text.into(),
108 allow_special,
109 }
110 }
111
112 pub fn as_encode_segment(&self) -> dynamo_tokenizers::EncodeSegment<'_> {
113 dynamo_tokenizers::EncodeSegment::new(&self.text, self.allow_special)
114 }
115}
116
117#[derive(Debug, Clone, PartialEq, Eq)]
123pub struct RenderedPrompt {
124 text: String,
125 segments: Option<Vec<RenderedSegment>>,
126 pending_segments: usize,
135}
136
137impl RenderedPrompt {
138 pub fn text(text: String) -> Self {
139 Self {
140 text,
141 segments: None,
142 pending_segments: 0,
143 }
144 }
145
146 pub fn segmented(segments: Vec<RenderedSegment>) -> Self {
147 Self::segmented_with_pending(segments, 0)
148 }
149
150 pub fn segmented_with_pending(segments: Vec<RenderedSegment>, pending_segments: usize) -> Self {
158 assert!(
159 pending_segments <= segments.len(),
160 "pending_segments ({pending_segments}) exceeds rendered segments ({})",
161 segments.len()
162 );
163 let text = segments
164 .iter()
165 .map(|segment| segment.text.as_str())
166 .collect();
167 Self {
168 text,
169 segments: Some(segments),
170 pending_segments,
171 }
172 }
173
174 pub fn as_str(&self) -> &str {
175 &self.text
176 }
177
178 pub fn segments(&self) -> Option<&[RenderedSegment]> {
179 self.segments.as_deref()
180 }
181
182 pub fn encode_segments(&self) -> Option<Vec<dynamo_tokenizers::EncodeSegment<'_>>> {
183 Some(
184 self.segments()?
185 .iter()
186 .map(RenderedSegment::as_encode_segment)
187 .collect(),
188 )
189 }
190
191 pub fn pending_segments(&self) -> usize {
196 self.pending_segments
197 }
198
199 pub fn pending_encode_segments(&self) -> Option<Vec<dynamo_tokenizers::EncodeSegment<'_>>> {
204 let segments = self.segments()?;
205 let start = segments.len() - self.pending_segments;
206 Some(
207 segments[start..]
208 .iter()
209 .map(RenderedSegment::as_encode_segment)
210 .collect(),
211 )
212 }
213
214 pub fn into_text(self) -> String {
215 self.text
216 }
217}
218
219#[derive(Debug, Clone, PartialEq, Eq)]
225pub enum PromptRenderError {
226 InvalidRequest(String),
227}
228
229impl PromptRenderError {
230 pub fn invalid_request(message: impl Into<String>) -> Self {
231 Self::InvalidRequest(message.into())
232 }
233}
234
235impl std::fmt::Display for PromptRenderError {
236 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
237 match self {
238 Self::InvalidRequest(message) => f.write_str(message),
239 }
240 }
241}
242
243impl std::error::Error for PromptRenderError {}
244
245pub trait OAIChatLikeRequest {
252 fn model(&self) -> String;
253 fn messages(&self) -> Value;
254 fn typed_messages(&self) -> Option<&[dynamo_protocols::types::ChatCompletionRequestMessage]> {
260 None
261 }
262 fn tools(&self) -> Option<Value> {
263 None
264 }
265 fn tool_choice(&self) -> Option<Value> {
266 None
267 }
268 fn response_format(&self) -> Option<Value> {
269 None
270 }
271
272 fn reasoning_effort(&self) -> Option<Value> {
275 None
276 }
277
278 fn should_add_generation_prompt(&self) -> bool;
279
280 fn chat_template_args(&self) -> Option<&HashMap<String, serde_json::Value>> {
282 None
283 }
284
285 fn prompt_input_type(&self) -> PromptInput {
287 PromptInput::Text(TextInput::Single(String::new()))
288 }
289
290 fn extract_tokens(&self) -> Option<TokenInput> {
292 None
293 }
294
295 fn extract_text(&self) -> Option<TextInput> {
296 None
297 }
298
299 fn mm_processor_kwargs(&self) -> Option<&serde_json::Value> {
300 None
301 }
302}
303
304pub(crate) fn messages_to_json(req: &dyn OAIChatLikeRequest) -> Result<serde_json::Value> {
307 if let Some(messages) = req.typed_messages() {
308 serde_json::to_value(messages)
309 } else {
310 serde_json::to_value(req.messages())
311 }
312 .context("Failed to convert messages to JSON")
313}
314
315pub trait OAIPromptFormatter: Send + Sync + 'static {
316 fn supports_add_generation_prompt(&self) -> bool;
317 fn render(&self, req: &dyn OAIChatLikeRequest) -> Result<String>;
318
319 fn media_message_order(&self, _request: &dyn OAIChatLikeRequest) -> Option<Vec<usize>> {
323 None
324 }
325
326 fn render_prompt(&self, req: &dyn OAIChatLikeRequest) -> Result<RenderedPrompt> {
327 self.render(req).map(RenderedPrompt::text)
328 }
329}
330
331pub(crate) fn reject_unsupported_partial_assistant(messages: &serde_json::Value) -> Result<()> {
338 let has_partial =
339 messages.as_array().into_iter().flatten().any(|message| {
340 message.get("partial").and_then(serde_json::Value::as_bool) == Some(true)
341 });
342 if has_partial {
343 return Err(PromptRenderError::invalid_request(
344 "assistant `partial: true` is not supported by this model's prompt formatter",
345 )
346 .into());
347 }
348 Ok(())
349}
350
351pub(crate) fn reject_unsupported_message_tools(
354 messages: &serde_json::Value,
355 supported_tool_roles: &[&str],
356) -> Result<()> {
357 let offending = messages.as_array().into_iter().flatten().find(|message| {
358 let declares_tools = message
359 .get("tools")
360 .is_some_and(|tools| !tools.is_null() && !tools.as_array().is_some_and(Vec::is_empty));
361 let role_is_supported = message
362 .get("role")
363 .and_then(serde_json::Value::as_str)
364 .is_some_and(|role| supported_tool_roles.contains(&role));
365 declares_tools && !role_is_supported
366 });
367
368 if let Some(message) = offending {
369 let role = message
370 .get("role")
371 .and_then(serde_json::Value::as_str)
372 .unwrap_or("<missing>");
373 return Err(PromptRenderError::invalid_request(format!(
374 "message-level `tools` on role {role:?} are not supported by this model's prompt \
375 formatter"
376 ))
377 .into());
378 }
379 Ok(())
380}
381
382#[derive(Clone)]
383pub enum PromptFormatter {
384 OAI(Arc<dyn OAIPromptFormatter>),
385}
386
387#[derive(Debug, Default)]
389pub struct NoOpFormatter;
390
391impl OAIPromptFormatter for NoOpFormatter {
392 fn supports_add_generation_prompt(&self) -> bool {
393 false
394 }
395
396 fn render(&self, req: &dyn OAIChatLikeRequest) -> Result<String> {
397 let messages = req.messages();
398 let messages_json = serde_json::to_value(&messages)?;
399 reject_unsupported_partial_assistant(&messages_json)?;
400 reject_unsupported_message_tools(&messages_json, &[])?;
401
402 let first_message = messages
403 .get_item_by_index(0)
404 .map_err(|_| anyhow::Error::msg("No message at index 0 or messages array is empty"))?;
405
406 let content = first_message
407 .get_attr("content")
408 .map_err(|_| anyhow::Error::msg("First message has no 'content' field"))?;
409
410 let content_str = content
411 .as_str()
412 .ok_or_else(|| anyhow::Error::msg("Message content is not a string"))?
413 .to_string();
414 Ok(content_str)
415 }
416}
417
418impl PromptFormatter {
419 pub fn no_op() -> Self {
420 Self::OAI(Arc::new(NoOpFormatter))
421 }
422}
423
424#[cfg(test)]
425mod rendered_prompt_tests {
426 use super::{
427 NoOpFormatter, OAIPromptFormatter, PromptRenderError, RenderedPrompt, RenderedSegment,
428 };
429
430 #[test]
431 fn owned_segments_borrow_into_tokenizer_segments() {
432 let prompt = RenderedPrompt::segmented(vec![
433 RenderedSegment::new("<|open|>", true),
434 RenderedSegment::new("user text", false),
435 ]);
436
437 let segments = prompt.encode_segments().expect("segmented prompt");
438 assert_eq!(segments[0].text, "<|open|>");
439 assert!(segments[0].allow_special);
440 assert_eq!(segments[1].text, "user text");
441 assert!(!segments[1].allow_special);
442 assert_eq!(prompt.as_str(), "<|open|>user text");
443 }
444
445 #[test]
446 fn no_op_formatter_rejects_unsupported_partial_assistant() {
447 let request: dynamo_protocols::types::CreateChatCompletionRequest =
448 serde_json::from_value(serde_json::json!({
449 "model": "test",
450 "messages": [
451 {"role": "user", "content": "Continue"},
452 {"role": "assistant", "content": "prefix", "partial": true}
453 ]
454 }))
455 .unwrap();
456
457 let error = NoOpFormatter.render(&request).unwrap_err();
458 assert!(matches!(
459 error.downcast_ref::<PromptRenderError>(),
460 Some(PromptRenderError::InvalidRequest(message))
461 if message.contains("`partial: true` is not supported")
462 ));
463 }
464
465 #[test]
466 fn no_op_formatter_rejects_message_level_tools() {
467 let request: dynamo_protocols::types::CreateChatCompletionRequest =
468 serde_json::from_value(serde_json::json!({
469 "model": "test",
470 "messages": [
471 {"role": "system", "tools": [{"name": "lookup"}]},
472 {"role": "user", "content": "Continue"}
473 ]
474 }))
475 .unwrap();
476
477 let error = NoOpFormatter.render(&request).unwrap_err();
478 assert!(matches!(
479 error.downcast_ref::<PromptRenderError>(),
480 Some(PromptRenderError::InvalidRequest(message))
481 if message.contains("message-level `tools`")
482 ));
483 }
484
485 #[test]
486 fn rendered_prompt_pending_segments_default_to_zero() {
487 let prompt = RenderedPrompt::segmented(vec![RenderedSegment::new("a", false)]);
488 assert_eq!(prompt.pending_segments(), 0);
489 assert!(prompt.pending_encode_segments().unwrap().is_empty());
490 assert_eq!(RenderedPrompt::text("a".to_string()).pending_segments(), 0);
491 assert!(
492 RenderedPrompt::text("a".to_string())
493 .pending_encode_segments()
494 .is_none()
495 );
496 }
497
498 #[test]
499 fn rendered_prompt_exposes_trailing_pending_segments() {
500 let prompt = RenderedPrompt::segmented_with_pending(
501 vec![
502 RenderedSegment::new("prompt", false),
503 RenderedSegment::new("<|open|>", true),
504 RenderedSegment::new("response", false),
505 RenderedSegment::new("<|sep|>", true),
506 ],
507 3,
508 );
509 assert_eq!(prompt.as_str(), "prompt<|open|>response<|sep|>");
510 assert_eq!(prompt.pending_segments(), 3);
511 let pending: Vec<&str> = prompt
512 .pending_encode_segments()
513 .unwrap()
514 .iter()
515 .map(|segment| segment.text)
516 .collect();
517 assert_eq!(pending, ["<|open|>", "response", "<|sep|>"]);
518 }
519
520 #[test]
521 #[should_panic(expected = "pending_segments")]
522 fn rendered_prompt_rejects_more_pending_than_segments() {
523 RenderedPrompt::segmented_with_pending(vec![RenderedSegment::new("a", false)], 2);
524 }
525}