use crate::{Message, Role};
const IM_START: &str = "<|im_start|>";
const IM_END: &str = "<|im_end|>";
fn role_tag(role: Role) -> &'static str {
match role {
Role::System => "system",
Role::User => "user",
Role::Assistant => "assistant",
}
}
pub(crate) fn render_chatml(messages: &[Message]) -> String {
let mut prompt = String::new();
for message in messages {
prompt.push_str(IM_START);
prompt.push_str(role_tag(message.role));
prompt.push('\n');
prompt.push_str(&message.content);
prompt.push_str(IM_END);
prompt.push('\n');
}
prompt.push_str(IM_START);
prompt.push_str("assistant\n");
prompt
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Message;
#[test]
fn renders_a_system_user_assistant_conversation_exactly() {
let messages = [
Message::system("You are a helpful assistant"),
Message::user("Hello"),
Message::assistant("Hi there"),
Message::user("Who are you"),
Message::assistant(" I am an assistant "),
Message::user("Another question"),
];
let expected = "<|im_start|>system\nYou are a helpful assistant<|im_end|>\n\
<|im_start|>user\nHello<|im_end|>\n\
<|im_start|>assistant\nHi there<|im_end|>\n\
<|im_start|>user\nWho are you<|im_end|>\n\
<|im_start|>assistant\n I am an assistant <|im_end|>\n\
<|im_start|>user\nAnother question<|im_end|>\n\
<|im_start|>assistant\n";
assert_eq!(render_chatml(&messages), expected);
}
#[test]
fn renders_edge_case_conversations_exactly() {
let cases: &[(&[Message], &str)] = &[
(&[], "<|im_start|>assistant\n"),
(
&[Message { role: Role::User, content: String::new() }],
"<|im_start|>user\n<|im_end|>\n<|im_start|>assistant\n",
),
(
&[Message { role: Role::System, content: "be terse".to_string() }],
"<|im_start|>system\nbe terse<|im_end|>\n<|im_start|>assistant\n",
),
(
&[
Message { role: Role::System, content: "be terse".to_string() },
Message { role: Role::User, content: "2+2?".to_string() },
],
"<|im_start|>system\nbe terse<|im_end|>\n<|im_start|>user\n2+2?<|im_end|>\n<|im_start|>assistant\n",
),
];
for (messages, expected) in cases {
assert_eq!(render_chatml(messages), *expected, "mismatch rendering {messages:?}");
}
}
#[test]
fn message_content_is_never_escaped_or_reinterpreted() {
let messages = [Message::user("line one\nline two <|im_end|> not a real stop tag")];
assert_eq!(
render_chatml(&messages),
"<|im_start|>user\nline one\nline two <|im_end|> not a real stop tag<|im_end|>\n<|im_start|>assistant\n"
);
}
}