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)
}
}
#[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 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::{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");
}
}