use crate::anthropic::v1::request::{AnthropicSettings, MessageParam as AnthropicMessage};
use crate::error::TypeError;
use crate::google::v1::generate::request::{GeminiContent, GeminiSettings};
use crate::openai::v1::chat::request::ChatMessage as OpenAIChatMessage;
use crate::openai::v1::chat::settings::OpenAIChatSettings;
use crate::prompt::builder::{to_provider_request, ProviderRequest};
use crate::prompt::settings::ModelSettings;
use crate::prompt::types::parse_response_to_json;
use crate::prompt::types::ResponseType;
use crate::prompt::types::Role;
use crate::prompt::{AnthropicMessageList, GeminiContentList, MessageNum, OpenAIMessageList};
use crate::tools::AgentToolDefinition;
use crate::traits::MessageFactory;
use crate::SettingsType;
use crate::{Provider, SaveName};
use potato_util::utils::extract_string_value;
use potato_util::PyHelperFuncs;
use potatohead_macro::try_extract_message;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyList, PyString, PyTuple};
use pythonize::pythonize;
use serde::{Deserialize, Deserializer, Serialize};
use serde_json::Value;
use std::collections::BTreeSet;
use std::path::{Path, PathBuf};
fn deserialize_string_or_vec<'de, D>(deserializer: D) -> Result<Vec<String>, D::Error>
where
D: Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(untagged)]
enum StringOrVec {
Single(String),
List(Vec<String>),
}
match StringOrVec::deserialize(deserializer)? {
StringOrVec::Single(s) => Ok(vec![s]),
StringOrVec::List(v) => Ok(v),
}
}
#[derive(Debug, Deserialize)]
pub struct GenericPromptConfig {
model: String,
#[serde(default)]
provider: Option<String>,
#[serde(deserialize_with = "deserialize_string_or_vec")]
messages: Vec<String>,
#[serde(default)]
system_instructions: Option<Vec<String>>,
#[serde(default)]
settings: Option<Value>,
response_format: Option<Value>,
}
const PROMPT_FILE_EXTENSIONS: [&str; 3] = ["yaml", "yml", "json"];
fn push_attempted_path(attempted_paths: &mut Vec<PathBuf>, path: PathBuf) {
if !attempted_paths.iter().any(|existing| existing == &path) {
attempted_paths.push(path);
}
}
fn format_candidate_paths(paths: &[PathBuf]) -> String {
paths
.iter()
.map(|path| path.display().to_string())
.collect::<Vec<_>>()
.join(", ")
}
fn create_message_for_provider(
content: String,
provider: &Provider,
role: &str,
) -> Result<MessageNum, TypeError> {
match provider {
Provider::OpenAI => {
OpenAIChatMessage::from_text(content, role).map(MessageNum::OpenAIMessageV1)
}
Provider::Anthropic => {
AnthropicMessage::from_text(content, role).map(MessageNum::AnthropicMessageV1)
}
Provider::Gemini | Provider::Google | Provider::Vertex | Provider::GoogleAdk => {
GeminiContent::from_text(content, role).map(MessageNum::GeminiContentV1)
}
_ => Err(TypeError::Error(format!(
"Unsupported provider for message creation: {:?}",
provider
))),
}
}
fn parse_single_message(
message: &Bound<'_, PyAny>,
provider: &Provider,
default_role: &str,
) -> Result<MessageNum, TypeError> {
if message.is_instance_of::<PyString>() {
let text = message.extract::<String>()?;
return create_message_for_provider(text, provider, default_role);
}
try_extract_message!(
message,
OpenAIChatMessage => MessageNum::OpenAIMessageV1,
AnthropicMessage => MessageNum::AnthropicMessageV1,
GeminiContent => MessageNum::GeminiContentV1,
);
Err(TypeError::InvalidMessageTypeInList(
message.get_type().name()?.to_string(),
))
}
fn parse_messages(
messages: &Bound<'_, PyAny>,
provider: &Provider,
default_role: &str,
) -> Result<Vec<MessageNum>, TypeError> {
let mut messages =
if !messages.is_instance_of::<PyList>() && !messages.is_instance_of::<PyTuple>() {
vec![parse_single_message(messages, provider, default_role)?]
} else {
messages
.try_iter()?
.map(|item| {
let item = item?;
parse_single_message(&item, provider, default_role)
})
.collect::<Result<Vec<_>, _>>()?
};
if provider == &Provider::Anthropic
&& (default_role == Role::System.as_str()
|| default_role == Role::Assistant.as_str()
|| default_role == Role::Developer.as_str())
{
for msg in messages.iter_mut() {
msg.anthropic_message_to_system_message()?;
}
}
Ok(messages)
}
fn get_system_role(provider: &Provider) -> &'static str {
match provider {
Provider::OpenAI => Role::Developer.into(),
Provider::Gemini | Provider::Vertex | Provider::Google | Provider::GoogleAdk => {
Role::Model.into()
}
Provider::Anthropic => Role::System.into(),
_ => Role::System.into(),
}
}
pub fn create_system_message_for_provider(
content: String,
provider: &Provider,
) -> Result<MessageNum, TypeError> {
let role = get_system_role(provider);
let mut msg = create_message_for_provider(content, provider, role)?;
if provider == &Provider::Anthropic {
msg.anthropic_message_to_system_message()?;
}
Ok(msg)
}
pub fn extract_system_instructions(
system_instruction: Option<&Bound<'_, PyAny>>,
provider: &Provider,
) -> Result<Option<Vec<MessageNum>>, TypeError> {
let system_instructions = if let Some(sys_inst) = system_instruction {
Some(parse_messages(
sys_inst,
provider,
get_system_role(provider),
)?)
} else {
None
};
Ok(system_instructions)
}
#[pyclass(from_py_object)]
#[derive(Debug, Serialize, Clone, PartialEq)]
pub struct Prompt {
pub request: ProviderRequest,
#[pyo3(get)]
pub model: String,
#[pyo3(get)]
pub provider: Provider,
pub version: String,
#[pyo3(get)]
#[serde(default)]
pub parameters: Vec<String>,
#[pyo3(get)]
#[serde(default)]
pub media_parameters: Vec<String>,
#[serde(default)]
pub response_type: ResponseType,
}
fn extract_model_settings(model_settings: &Bound<'_, PyAny>) -> Result<ModelSettings, TypeError> {
let settings_type = model_settings
.call_method0("settings_type")?
.extract::<SettingsType>()?;
match settings_type {
SettingsType::OpenAIChat => model_settings
.extract::<OpenAIChatSettings>()
.map(ModelSettings::OpenAIChat),
SettingsType::GoogleChat => model_settings
.extract::<GeminiSettings>()
.map(ModelSettings::GoogleChat),
SettingsType::Anthropic => model_settings
.extract::<AnthropicSettings>()
.map(ModelSettings::AnthropicChat),
SettingsType::ModelSettings => model_settings.extract::<ModelSettings>(),
}
.map_err(Into::into)
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
enum PromptFormat {
Generic(GenericPromptConfig),
Full(Box<PromptInternal>),
}
#[derive(Debug, Deserialize)]
struct PromptInternal {
request: Value,
model: String,
provider: Provider,
version: String,
#[serde(default)]
parameters: Vec<String>,
#[serde(default)]
media_parameters: Vec<String>,
#[serde(default)]
response_type: ResponseType,
}
impl<'de> Deserialize<'de> for Prompt {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let format = PromptFormat::deserialize(deserializer)?;
match format {
PromptFormat::Generic(config) => Self::from_generic_config(config)
.map_err(|e| serde::de::Error::custom(e.to_string())),
PromptFormat::Full(internal) => Ok(Prompt {
request: provider_request_from_value(internal.provider.clone(), internal.request)
.map_err(|e| serde::de::Error::custom(e.to_string()))?,
model: internal.model,
provider: internal.provider,
version: internal.version,
parameters: internal.parameters,
media_parameters: internal.media_parameters,
response_type: internal.response_type,
}),
}
}
}
fn provider_request_from_value(
provider: Provider,
value: Value,
) -> Result<ProviderRequest, TypeError> {
match provider {
Provider::OpenAI => Ok(ProviderRequest::OpenAIV1(serde_json::from_value(value)?)),
Provider::Anthropic => Ok(ProviderRequest::AnthropicV1(serde_json::from_value(value)?)),
Provider::Gemini | Provider::Google | Provider::Vertex | Provider::GoogleAdk => {
Ok(ProviderRequest::GeminiV1(serde_json::from_value(value)?))
}
Provider::Undefined => Err(TypeError::UnsupportedProviderForRequestCreation),
}
}
#[pymethods]
impl Prompt {
#[new]
#[pyo3(signature = (messages, model, provider=None, system_instructions=None, model_settings=None, output_type=None))]
pub fn new(
py: Python<'_>,
messages: &Bound<'_, PyAny>,
model: &str,
provider: Option<&Bound<'_, PyAny>>,
system_instructions: Option<&Bound<'_, PyAny>>,
model_settings: Option<&Bound<'_, PyAny>>,
output_type: Option<&Bound<'_, PyAny>>, ) -> Result<Self, TypeError> {
let model_settings = model_settings
.as_ref()
.map(|s| extract_model_settings(s))
.transpose()?;
let provider = Provider::resolve_from_py(provider)?;
let messages = parse_messages(messages, &provider, Role::User.into())?;
let system_instructions = if let Some(sys_inst) = system_instructions {
parse_messages(sys_inst, &provider, get_system_role(&provider))?
} else {
vec![]
};
let (response_type, response_json_schema) = match output_type {
Some(output_type) => {
parse_response_to_json(py, output_type)?
}
None => (ResponseType::Null, None),
};
Self::new_rs(
messages,
model,
provider,
system_instructions,
model_settings,
response_json_schema,
response_type,
)
}
#[getter]
pub fn model_settings<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
self.request.model_settings(py)
}
#[getter]
pub fn model_identifier(&self) -> String {
format!("{}:{}", self.provider.as_str(), self.model)
}
#[pyo3(signature = (path = None))]
pub fn save_prompt(&self, path: Option<PathBuf>) -> Result<PathBuf, TypeError> {
let save_path = path.unwrap_or_else(|| PathBuf::from(SaveName::Prompt));
PyHelperFuncs::save_to_json(self, &save_path)?;
Ok(save_path.with_extension("json"))
}
#[staticmethod]
pub fn from_path(path: PathBuf) -> Result<Self, TypeError> {
Self::load_prompt_from_path(path.as_path(), None)
}
#[staticmethod]
pub fn model_validate_json(json_string: String) -> Result<Self, TypeError> {
let json_value: Value = serde_json::from_str(&json_string)?;
let model: Self = serde_json::from_value(json_value)?;
Ok(model)
}
pub fn model_dump_json(&self) -> String {
serde_json::to_string(self).unwrap()
}
pub fn __str__(&self) -> String {
PyHelperFuncs::__str__(self)
}
#[getter]
pub fn all_messages<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyList>, TypeError> {
self.request.get_all_py_messages(py)
}
#[getter]
pub fn messages<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyList>, TypeError> {
self.request.get_py_messages(py)
}
#[getter]
pub fn message<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
self.request.get_py_message(py)
}
#[getter]
pub fn openai_messages(&self) -> Result<OpenAIMessageList, TypeError> {
if self.provider != Provider::OpenAI {
return Err(TypeError::Error(
"Prompt provider is not OpenAI".to_string(),
));
}
let messages = self
.request
.messages()
.iter()
.filter(|msg| msg.is_user_message())
.filter_map(|msg| match msg {
MessageNum::OpenAIMessageV1(m) => Some(m.clone()),
_ => None,
})
.collect::<Vec<_>>();
Ok(OpenAIMessageList { messages })
}
#[getter]
pub fn openai_message(&self) -> Result<OpenAIChatMessage, TypeError> {
if self.provider != Provider::OpenAI {
return Err(TypeError::Error(
"Prompt provider is not OpenAI".to_string(),
));
}
self.request.get_openai_message()
}
#[getter]
pub fn gemini_messages(&self) -> Result<GeminiContentList, TypeError> {
if !self.is_google_provider() {
return Err(TypeError::Error(
"Prompt provider is not Google, Gemini, or Vertex".to_string(),
));
}
let messages = self
.request
.messages()
.iter()
.filter(|msg| msg.is_user_message())
.filter_map(|msg| match msg {
MessageNum::GeminiContentV1(m) => Some(m.clone()),
_ => None,
})
.collect::<Vec<_>>();
Ok(GeminiContentList { messages })
}
#[getter]
pub fn gemini_message(&self) -> Result<GeminiContent, TypeError> {
if !self.is_google_provider() {
return Err(TypeError::Error(
"Prompt provider is not Google, Gemini, or Vertex".to_string(),
));
}
self.request.get_gemini_message()
}
#[getter]
pub fn anthropic_messages(&self) -> Result<AnthropicMessageList, TypeError> {
if self.provider != Provider::Anthropic {
return Err(TypeError::Error(
"Prompt provider is not Anthropic".to_string(),
));
}
let messages = self
.request
.messages()
.iter()
.filter(|msg| msg.is_user_message())
.filter_map(|msg| match msg {
MessageNum::AnthropicMessageV1(m) => Some(m.clone()),
_ => None,
})
.collect::<Vec<_>>();
Ok(AnthropicMessageList { messages })
}
#[getter]
pub fn anthropic_message(&self) -> Result<AnthropicMessage, TypeError> {
if self.provider != Provider::Anthropic {
return Err(TypeError::Error(
"Prompt provider is not Anthropic".to_string(),
));
}
self.request.get_anthropic_message()
}
#[getter]
pub fn system_instructions<'py>(
&self,
py: Python<'py>,
) -> Result<Bound<'py, PyList>, TypeError> {
self.request.get_py_system_instructions(py)
}
#[pyo3(signature = (name=None, value=None, **kwargs))]
pub fn bind(
&self,
name: Option<&str>,
value: Option<&Bound<'_, PyAny>>,
kwargs: Option<&Bound<'_, PyDict>>,
) -> Result<Self, TypeError> {
let mut new_prompt = self.clone();
if let (Some(name), Some(value)) = (name, value) {
let var_value = extract_string_value(value)?;
for message in new_prompt.request.messages_mut() {
message.bind_mut(name, &var_value)?;
}
}
if let Some(kwargs) = kwargs {
for (key, val) in kwargs.iter() {
let var_name = key.extract::<String>()?;
let var_value = extract_string_value(&val)?;
for message in new_prompt.request.messages_mut() {
message.bind_mut(&var_name, &var_value)?;
}
}
}
if name.is_none() && kwargs.is_none_or(|k| k.is_empty()) {
return Err(TypeError::Error(
"Must provide either (name, value) or keyword arguments for binding".to_string(),
));
}
Ok(new_prompt)
}
#[pyo3(signature = (name=None, value=None, **kwargs))]
pub fn bind_mut(
&mut self,
name: Option<&str>,
value: Option<&Bound<'_, PyAny>>,
kwargs: Option<&Bound<'_, PyDict>>,
) -> Result<(), TypeError> {
if let (Some(name), Some(value)) = (name, value) {
let var_value = extract_string_value(value)?;
for message in self.request.messages_mut() {
message.bind_mut(name, &var_value)?;
}
}
if let Some(kwargs) = kwargs {
for (key, val) in kwargs.iter() {
let var_name = key.extract::<String>()?;
let var_value = extract_string_value(&val)?;
for message in self.request.messages_mut() {
message.bind_mut(&var_name, &var_value)?;
}
}
}
if name.is_none() && kwargs.is_none_or(|k| k.is_empty()) {
return Err(TypeError::Error(
"Must provide either (name, value) or keyword arguments for binding".to_string(),
));
}
Ok(())
}
pub fn bind_media(
&self,
name: &str,
media: &crate::prompt::media::MediaRef,
) -> Result<Self, TypeError> {
let mut new_prompt = self.clone();
new_prompt.bind_media_mut(name, media)?;
Ok(new_prompt)
}
pub fn bind_media_mut(
&mut self,
name: &str,
media: &crate::prompt::media::MediaRef,
) -> Result<(), TypeError> {
let token = format!("${{media:{name}}}");
let mut found = false;
for message in self.request.messages_mut() {
if message.bind_media_mut(&token, media, &self.provider)? {
found = true;
}
}
if !found {
return Err(TypeError::MediaPlaceholderNotFound {
name: name.to_string(),
});
}
Ok(())
}
#[getter]
pub fn response_json_schema_pretty(&self) -> Option<String> {
Some(PyHelperFuncs::__str__(
self.request.response_json_schema().as_ref()?,
))
}
#[getter]
#[pyo3(name = "response_json_schema")]
pub fn response_json_schema_py(&self) -> Option<String> {
Some(self.request.response_json_schema().as_ref()?.to_string())
}
pub fn model_dump<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
let request = &self.request.to_json()?;
Ok(pythonize(py, request)?)
}
}
impl Prompt {
fn read_prompt_file(path: &Path) -> Result<Self, TypeError> {
let content = std::fs::read_to_string(path)?;
let extension = path
.extension()
.and_then(|ext| ext.to_str())
.ok_or_else(|| TypeError::Error(format!("Invalid file path: {:?}", path)))?;
let mut prompt: Prompt = match extension.to_lowercase().as_str() {
"json" => serde_json::from_str(&content)?,
"yaml" | "yml" => serde_yaml::from_str(&content)?,
_ => {
return Err(TypeError::Error(format!(
"Unsupported file extension '{}'. Expected .json, .yaml, or .yml",
extension
)))
}
};
if prompt.parameters.is_empty() {
let system_instructions: Vec<MessageNum> = prompt
.request
.system_instructions()
.iter()
.map(|msg| (*msg).clone())
.collect();
let parameters =
Self::extract_variables(prompt.request.messages(), &system_instructions);
prompt.parameters = parameters;
}
if prompt.media_parameters.is_empty() {
let system_instructions: Vec<MessageNum> = prompt
.request
.system_instructions()
.iter()
.map(|msg| (*msg).clone())
.collect();
let media_parameters =
Self::extract_media_variables(prompt.request.messages(), &system_instructions);
prompt.media_parameters = media_parameters;
}
Ok(prompt)
}
fn resolve_prompt_candidate(
requested_path: &Path,
candidate_path: PathBuf,
attempted_paths: &mut Vec<PathBuf>,
) -> Result<Option<PathBuf>, TypeError> {
push_attempted_path(attempted_paths, candidate_path.clone());
if candidate_path.is_file() {
return Ok(Some(candidate_path));
}
if requested_path.extension().is_some() {
return Ok(None);
}
let mut matches = Vec::new();
for extension in PROMPT_FILE_EXTENSIONS {
let extension_candidate = candidate_path.with_extension(extension);
push_attempted_path(attempted_paths, extension_candidate.clone());
if extension_candidate.is_file() {
matches.push(extension_candidate);
}
}
match matches.len() {
0 => Ok(None),
1 => Ok(matches.into_iter().next()),
_ => Err(TypeError::AmbiguousPromptPath {
requested_path: requested_path.display().to_string(),
candidate_paths: format_candidate_paths(&matches),
}),
}
}
fn resolve_prompt_path(path: &Path, base_dir: Option<&Path>) -> Result<PathBuf, TypeError> {
let mut attempted_paths = Vec::new();
if path.is_absolute() {
return Self::resolve_prompt_candidate(path, path.to_path_buf(), &mut attempted_paths)?
.ok_or_else(|| TypeError::PromptPathNotFound {
requested_path: path.display().to_string(),
attempted_paths: format_candidate_paths(&attempted_paths),
});
}
let mut candidate_roots = Vec::new();
if let Some(base_dir) = base_dir {
candidate_roots.push(base_dir.to_path_buf());
}
let current_dir = std::env::current_dir()?;
if !candidate_roots.iter().any(|root| root == ¤t_dir) {
candidate_roots.push(current_dir);
}
for root in candidate_roots {
let candidate_path = root.join(path);
if let Some(resolved_path) =
Self::resolve_prompt_candidate(path, candidate_path, &mut attempted_paths)?
{
return Ok(resolved_path);
}
}
Err(TypeError::PromptPathNotFound {
requested_path: path.display().to_string(),
attempted_paths: format_candidate_paths(&attempted_paths),
})
}
fn load_prompt_from_path(path: &Path, base_dir: Option<&Path>) -> Result<Self, TypeError> {
let resolved_path = Self::resolve_prompt_path(path, base_dir)?;
Self::read_prompt_file(&resolved_path)
}
pub fn from_path_with_base(
path: impl AsRef<Path>,
base_dir: impl AsRef<Path>,
) -> Result<Self, TypeError> {
Self::load_prompt_from_path(path.as_ref(), Some(base_dir.as_ref()))
}
pub fn from_generic_config(config: GenericPromptConfig) -> Result<Self, TypeError> {
if config.messages.is_empty() {
return Err(TypeError::Error(
"Prompt has no messages. Generic prompt format requires at least one message."
.to_string(),
));
}
let provider = Provider::resolve(config.provider.as_deref())?;
let messages: Vec<MessageNum> = config
.messages
.into_iter()
.map(|msg| create_message_for_provider(msg, &provider, Role::User.as_str()))
.collect::<Result<Vec<_>, _>>()?;
let system_instructions = if let Some(sys_inst) = config.system_instructions {
sys_inst
.into_iter()
.map(|msg| create_message_for_provider(msg, &provider, get_system_role(&provider)))
.collect::<Result<Vec<_>, _>>()?
} else {
Vec::new()
};
let model_settings = if let Some(settings) = config.settings {
Some(Self::settings_from_value(settings, &provider)?)
} else {
None
};
Self::new_rs(
messages,
&config.model,
provider,
system_instructions,
model_settings,
config.response_format,
ResponseType::Null,
)
}
fn settings_from_value(value: Value, provider: &Provider) -> Result<ModelSettings, TypeError> {
match provider {
Provider::OpenAI => {
let settings: OpenAIChatSettings = serde_json::from_value(value)?;
Ok(ModelSettings::OpenAIChat(settings))
}
Provider::Anthropic => {
let settings: AnthropicSettings = serde_json::from_value(value)?;
Ok(ModelSettings::AnthropicChat(settings))
}
Provider::Gemini | Provider::Google | Provider::Vertex | Provider::GoogleAdk => {
let settings: GeminiSettings = serde_json::from_value(value)?;
Ok(ModelSettings::GoogleChat(settings))
}
_ => Err(TypeError::Error(format!(
"Settings not supported for provider: {:?}",
provider
))),
}
}
pub fn response_json_schema(&self) -> Option<&Value> {
self.request.response_json_schema()
}
pub fn new_rs(
mut messages: Vec<MessageNum>,
model: &str,
provider: Provider,
mut system_instructions: Vec<MessageNum>,
model_settings: Option<ModelSettings>,
response_json_schema: Option<Value>,
response_type: ResponseType,
) -> Result<Self, TypeError> {
let model = model.to_string();
let version = potato_util::version();
let model_settings = match model_settings {
Some(settings) => {
settings.validate_provider(&provider)?;
settings
}
None => ModelSettings::provider_default_settings(&provider),
};
if system_instructions
.iter()
.any(|msg| !msg.extract_media_variables().is_empty())
{
return Err(TypeError::MediaInSystemMessage);
}
for msg in messages.iter_mut() {
msg.split_media_placeholders()?;
}
for msg in system_instructions.iter_mut() {
msg.split_media_placeholders()?;
}
let parameters = Self::extract_variables(&messages, &system_instructions);
let media_parameters = Self::extract_media_variables(&messages, &system_instructions);
let request = to_provider_request(
messages,
system_instructions,
model.clone(),
model_settings,
response_json_schema,
)?;
Ok(Self {
request,
version,
parameters,
media_parameters,
response_type,
model,
provider,
})
}
fn is_google_provider(&self) -> bool {
matches!(
self.provider,
Provider::Google | Provider::Gemini | Provider::Vertex | Provider::GoogleAdk
)
}
pub fn add_tools(&mut self, tools: Vec<AgentToolDefinition>) -> Result<(), TypeError> {
self.request.add_tools(tools)
}
pub fn extract_variables(
messages: &[MessageNum],
system_instructions: &[MessageNum],
) -> Vec<String> {
let mut variables = BTreeSet::new();
for msg in system_instructions {
variables.extend(msg.extract_variables());
}
for msg in messages {
variables.extend(msg.extract_variables());
}
variables.into_iter().collect()
}
pub fn extract_media_variables(
messages: &[MessageNum],
system_instructions: &[MessageNum],
) -> Vec<String> {
let mut variables = BTreeSet::new();
for msg in system_instructions {
variables.extend(msg.extract_media_variables());
}
for msg in messages {
variables.extend(msg.extract_media_variables());
}
variables.into_iter().collect()
}
pub fn model_dump_value(&self) -> Value {
serde_json::to_value(self).unwrap_or(Value::Null)
}
pub fn to_request_json(&self) -> Result<Value, TypeError> {
let json_value = serde_json::to_value(self)?;
Ok(json_value)
}
pub fn set_response_json_schema(
&mut self,
response_json_schema: Option<Value>,
response_type: ResponseType,
) {
self.request.set_response_json_schema(response_json_schema);
self.response_type = response_type;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::anthropic::v1::request::{
Base64ImageSource, Base64PDFSource, ContentBlockParam, DocumentBlockParam, ImageBlockParam,
MessageParam, PlainTextSource, TextBlockParam, UrlImageSource, UrlPDFSource,
};
use crate::google::{DataNum, GeminiContent, Part};
use crate::openai::v1::chat::request::{
ChatMessage as OpenAIChatMessage, ContentPart, FileContentPart, ImageContentPart,
TextContentPart,
};
use crate::prompt::types::Score;
use crate::StructuredOutput;
use std::fs;
use std::path::{Path, PathBuf};
fn create_temp_prompt_dir() -> PathBuf {
let dir = std::env::temp_dir().join(format!(
"potatohead-prompt-tests-{}",
potato_util::create_uuid7()
));
fs::create_dir_all(&dir).unwrap();
dir
}
fn write_generic_prompt(path: &Path, provider: &str, model: &str, message: &str) {
let content =
format!("model: {model}\nprovider: {provider}\nmessages:\n - \"{message}\"\n");
fs::write(path, content).unwrap();
}
fn create_openai_chat_message() -> OpenAIChatMessage {
let text_part = TextContentPart::new("What company is this logo from?".to_string());
let text_content_part = ContentPart::Text(text_part);
OpenAIChatMessage {
role: "user".to_string(),
content: vec![text_content_part],
name: None,
}
}
fn create_system_openai_chat_message() -> OpenAIChatMessage {
let text_part = TextContentPart::new("system_prompt".to_string());
let text_content_part = ContentPart::Text(text_part);
OpenAIChatMessage {
role: "developer".to_string(),
content: vec![text_content_part],
name: None,
}
}
fn create_openai_image_message() -> OpenAIChatMessage {
let image_part = ImageContentPart::new("https://iili.io/3Hs4FMg.png".to_string(), None);
let image_content_part = ContentPart::ImageUrl(image_part);
OpenAIChatMessage {
role: "user".to_string(),
content: vec![image_content_part],
name: None,
}
}
fn create_openai_file_message() -> OpenAIChatMessage {
let file_part = FileContentPart::new(
Some("filedata".to_string()),
Some("fileid".to_string()),
Some("filename".to_string()),
);
let file_content_part = ContentPart::FileContent(file_part);
OpenAIChatMessage {
role: "user".to_string(),
content: vec![file_content_part],
name: None,
}
}
fn create_anthropic_text_message() -> MessageParam {
let text_block =
TextBlockParam::new_rs("What company is this logo from?".to_string(), None, None);
MessageParam {
role: "user".to_string(),
content: vec![ContentBlockParam {
inner: crate::anthropic::v1::request::ContentBlock::Text(text_block),
}],
}
}
fn create_anthropic_system_message() -> MessageParam {
let text_block = TextBlockParam::new_rs("system_prompt".to_string(), None, None);
MessageParam {
role: "assistant".to_string(),
content: vec![ContentBlockParam {
inner: crate::anthropic::v1::request::ContentBlock::Text(text_block),
}],
}
}
fn create_anthropic_base64_image_message() -> MessageParam {
let image_source =
Base64ImageSource::new("image/png".to_string(), "base64data".to_string()).unwrap();
let image_block = ImageBlockParam {
source: crate::anthropic::v1::request::ImageSource::Base64(image_source),
cache_control: None,
r#type: "image".to_string(),
};
MessageParam {
role: "user".to_string(),
content: vec![ContentBlockParam {
inner: crate::anthropic::v1::request::ContentBlock::Image(image_block),
}],
}
}
fn create_anthropic_url_image_message() -> MessageParam {
let image_source = UrlImageSource::new("https://iili.io/3Hs4FMg.png".to_string());
let image_block = ImageBlockParam {
source: crate::anthropic::v1::request::ImageSource::Url(image_source),
cache_control: None,
r#type: "image".to_string(),
};
MessageParam {
role: "user".to_string(),
content: vec![ContentBlockParam {
inner: crate::anthropic::v1::request::ContentBlock::Image(image_block),
}],
}
}
fn create_anthropic_base64_pdf_message() -> MessageParam {
let pdf_source = Base64PDFSource::new("base64pdfdata".to_string()).unwrap();
let document_block = DocumentBlockParam {
source: crate::anthropic::v1::request::DocumentSource::Base64(pdf_source),
cache_control: None,
title: Some("test_document.pdf".to_string()),
context: None,
r#type: "document".to_string(),
citations: None,
};
MessageParam {
role: "user".to_string(),
content: vec![ContentBlockParam {
inner: crate::anthropic::v1::request::ContentBlock::Document(document_block),
}],
}
}
fn create_anthropic_url_pdf_message() -> MessageParam {
let pdf_source = UrlPDFSource::new("https://example.com/document.pdf".to_string());
let document_block = DocumentBlockParam {
source: crate::anthropic::v1::request::DocumentSource::Url(pdf_source),
cache_control: None,
title: Some("test_document.pdf".to_string()),
context: None,
r#type: "document".to_string(),
citations: None,
};
MessageParam {
role: "user".to_string(),
content: vec![ContentBlockParam {
inner: crate::anthropic::v1::request::ContentBlock::Document(document_block),
}],
}
}
fn create_anthropic_plain_text_document_message() -> MessageParam {
let text_source = PlainTextSource::new("Plain text document content".to_string());
let document_block = DocumentBlockParam {
source: crate::anthropic::v1::request::DocumentSource::Text(text_source),
cache_control: None,
title: Some("text_document.txt".to_string()),
context: Some("Context for the document".to_string()),
r#type: "document".to_string(),
citations: None,
};
MessageParam {
role: "user".to_string(),
content: vec![ContentBlockParam {
inner: crate::anthropic::v1::request::ContentBlock::Document(document_block),
}],
}
}
#[test]
fn test_task_list_add_and_get() {
let text_part = TextContentPart::new("Test prompt. ${param1} ${param2}".to_string());
let content_part = ContentPart::Text(text_part);
let message = OpenAIChatMessage {
role: "user".to_string(),
content: vec![content_part],
name: None,
};
let prompt = Prompt::new_rs(
vec![MessageNum::OpenAIMessageV1(message)],
"gpt-4o",
Provider::OpenAI,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap();
assert_eq!(prompt.request.messages().len(), 1);
assert!(prompt.parameters.len() == 2);
let mut parameters = prompt.parameters.clone();
parameters.sort();
assert_eq!(parameters[0], "param1");
assert_eq!(parameters[1], "param2");
let bound_msg = prompt.request.messages()[0]
.bind("param1", "Value1")
.unwrap();
let bound_msg = bound_msg.bind("param2", "Value2").unwrap();
match bound_msg.clone() {
MessageNum::OpenAIMessageV1(msg) => {
if let ContentPart::Text(text_part) = &msg.content[0] {
assert_eq!(text_part.text, "Test prompt. Value1 Value2");
} else {
panic!("Expected TextContentPart");
}
}
_ => panic!("Expected OpenAIMessageV1"),
}
}
#[test]
fn test_image_prompt() {
let text_message = create_openai_chat_message();
let image_message = create_openai_image_message();
let system_text_part = TextContentPart::new("system_prompt".to_string());
let system_text_content_part = ContentPart::Text(system_text_part);
let system_text_message = OpenAIChatMessage {
role: "assistant".to_string(),
content: vec![system_text_content_part],
name: None,
};
let prompt = Prompt::new_rs(
vec![
MessageNum::OpenAIMessageV1(text_message),
MessageNum::OpenAIMessageV1(image_message),
],
"gpt-4o",
Provider::OpenAI,
vec![MessageNum::OpenAIMessageV1(system_text_message)],
None,
None,
ResponseType::Null,
)
.unwrap();
if let MessageNum::OpenAIMessageV1(msg) = &prompt.request.messages()[1] {
if let ContentPart::Text(text_part) = &msg.content[0] {
assert_eq!(text_part.text, "What company is this logo from?");
} else {
panic!("Expected TextContentPart for the first user message");
}
} else {
panic!("Expected OpenAIMessageV1 for the first user message");
}
if let MessageNum::OpenAIMessageV1(msg) = &prompt.request.messages()[2] {
if let ContentPart::ImageUrl(image_url) = &msg.content[0] {
assert_eq!(image_url.image_url.url, "https://iili.io/3Hs4FMg.png");
assert_eq!(image_url.r#type, "image_url");
} else {
panic!("Expected ContentPart::Image for the second user message");
}
} else {
panic!("Expected OpenAIMessageV1 for the second user message");
}
}
#[test]
fn test_document_prompt() {
let text_message = create_openai_chat_message();
let file_message = create_openai_file_message();
let system_message = create_system_openai_chat_message();
let prompt = Prompt::new_rs(
vec![
MessageNum::OpenAIMessageV1(text_message),
MessageNum::OpenAIMessageV1(file_message),
],
"gpt-4o",
Provider::OpenAI,
vec![MessageNum::OpenAIMessageV1(system_message)],
None,
None,
ResponseType::Null,
)
.unwrap();
if let MessageNum::OpenAIMessageV1(msg) = &prompt.request.messages()[2] {
if let ContentPart::FileContent(file_content) = &msg.content[0] {
assert_eq!(file_content.file.file_id.as_ref().unwrap(), "fileid");
assert_eq!(file_content.file.filename.as_ref().unwrap(), "filename");
} else {
panic!("Expected ContentPart::FileContent for the second user message");
}
} else {
panic!("Expected OpenAIMessageV1 for the first user message");
}
}
#[test]
fn test_response_format_score() {
let text_message = create_openai_chat_message();
let prompt = Prompt::new_rs(
vec![MessageNum::OpenAIMessageV1(text_message)],
"gpt-4o",
Provider::OpenAI,
vec![],
None,
Some(Score::get_structured_output_schema()),
ResponseType::Null,
)
.unwrap();
assert!(prompt.response_json_schema().is_some());
}
#[test]
fn test_anthropic_text_message_binding() {
let text_block =
TextBlockParam::new_rs("Test prompt. ${param1} ${param2}".to_string(), None, None);
let message = MessageParam {
role: "user".to_string(),
content: vec![ContentBlockParam {
inner: crate::anthropic::v1::request::ContentBlock::Text(text_block),
}],
};
let prompt = Prompt::new_rs(
vec![MessageNum::AnthropicMessageV1(message)],
"claude-3-5-sonnet-20241022",
Provider::Anthropic,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap();
assert_eq!(prompt.request.messages().len(), 1);
assert_eq!(prompt.parameters.len(), 2);
let mut parameters = prompt.parameters.clone();
parameters.sort();
assert_eq!(parameters[0], "param1");
assert_eq!(parameters[1], "param2");
let bound_msg = prompt.request.messages()[0]
.bind("param1", "Value1")
.unwrap();
let bound_msg = bound_msg.bind("param2", "Value2").unwrap();
match bound_msg {
MessageNum::AnthropicMessageV1(msg) => {
if let crate::anthropic::v1::request::ContentBlock::Text(text_block) =
&msg.content[0].inner
{
assert_eq!(text_block.text, "Test prompt. Value1 Value2");
} else {
panic!("Expected TextBlockParam");
}
}
_ => panic!("Expected AnthropicMessageV1"),
}
}
#[test]
fn test_anthropic_url_image_prompt() {
let text_message = create_anthropic_text_message();
let image_message = create_anthropic_url_image_message();
let system_message = create_anthropic_system_message();
let prompt = Prompt::new_rs(
vec![
MessageNum::AnthropicMessageV1(text_message),
MessageNum::AnthropicMessageV1(image_message),
],
"claude-3-5-sonnet-20241022",
Provider::Anthropic,
vec![MessageNum::AnthropicMessageV1(system_message)],
None,
None,
ResponseType::Null,
)
.unwrap();
if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[0] {
if let crate::anthropic::v1::request::ContentBlock::Text(text_block) =
&msg.content[0].inner
{
assert_eq!(text_block.text, "What company is this logo from?");
} else {
panic!("Expected TextBlock for first message");
}
} else {
panic!("Expected AnthropicMessageV1");
}
if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
if let crate::anthropic::v1::request::ContentBlock::Image(image_block) =
&msg.content[0].inner
{
match &image_block.source {
crate::anthropic::v1::request::ImageSource::Url(url_source) => {
assert_eq!(url_source.url, "https://iili.io/3Hs4FMg.png");
assert_eq!(url_source.r#type, "url");
}
_ => panic!("Expected URL image source"),
}
assert_eq!(image_block.r#type, "image");
} else {
panic!("Expected ImageBlock for second message");
}
} else {
panic!("Expected AnthropicMessageV1");
}
}
#[test]
fn test_anthropic_base64_image_prompt() {
let text_message = create_anthropic_text_message();
let image_message = create_anthropic_base64_image_message();
let prompt = Prompt::new_rs(
vec![
MessageNum::AnthropicMessageV1(text_message),
MessageNum::AnthropicMessageV1(image_message),
],
"claude-3-5-sonnet-20241022",
Provider::Anthropic,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap();
if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
if let crate::anthropic::v1::request::ContentBlock::Image(image_block) =
&msg.content[0].inner
{
match &image_block.source {
crate::anthropic::v1::request::ImageSource::Base64(base64_source) => {
assert_eq!(base64_source.media_type, "image/png");
assert_eq!(base64_source.data, "base64data");
assert_eq!(base64_source.r#type, "base64");
}
_ => panic!("Expected Base64 image source"),
}
} else {
panic!("Expected ImageBlock");
}
} else {
panic!("Expected AnthropicMessageV1");
}
}
#[test]
fn test_anthropic_base64_pdf_document_prompt() {
let text_message = create_anthropic_text_message();
let pdf_message = create_anthropic_base64_pdf_message();
let system_message = create_anthropic_system_message();
let prompt = Prompt::new_rs(
vec![
MessageNum::AnthropicMessageV1(text_message),
MessageNum::AnthropicMessageV1(pdf_message),
],
"claude-3-5-sonnet-20241022",
Provider::Anthropic,
vec![MessageNum::AnthropicMessageV1(system_message)],
None,
None,
ResponseType::Null,
)
.unwrap();
if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
if let crate::anthropic::v1::request::ContentBlock::Document(document_block) =
&msg.content[0].inner
{
match &document_block.source {
crate::anthropic::v1::request::DocumentSource::Base64(pdf_source) => {
assert_eq!(pdf_source.media_type, "application/pdf");
assert_eq!(pdf_source.data, "base64pdfdata");
assert_eq!(pdf_source.r#type, "base64");
}
_ => panic!("Expected Base64 PDF source"),
}
assert_eq!(document_block.r#type, "document");
assert_eq!(document_block.title.as_ref().unwrap(), "test_document.pdf");
} else {
panic!("Expected DocumentBlock");
}
} else {
panic!("Expected AnthropicMessageV1");
}
}
#[test]
fn test_anthropic_url_pdf_document_prompt() {
let text_message = create_anthropic_text_message();
let pdf_message = create_anthropic_url_pdf_message();
let prompt = Prompt::new_rs(
vec![
MessageNum::AnthropicMessageV1(text_message),
MessageNum::AnthropicMessageV1(pdf_message),
],
"claude-3-5-sonnet-20241022",
Provider::Anthropic,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap();
if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
if let crate::anthropic::v1::request::ContentBlock::Document(document_block) =
&msg.content[0].inner
{
match &document_block.source {
crate::anthropic::v1::request::DocumentSource::Url(url_source) => {
assert_eq!(url_source.url, "https://example.com/document.pdf");
assert_eq!(url_source.r#type, "url");
}
_ => panic!("Expected URL PDF source"),
}
} else {
panic!("Expected DocumentBlock");
}
} else {
panic!("Expected AnthropicMessageV1");
}
}
#[test]
fn test_anthropic_plain_text_document_prompt() {
let text_message = create_anthropic_text_message();
let text_doc_message = create_anthropic_plain_text_document_message();
let prompt = Prompt::new_rs(
vec![
MessageNum::AnthropicMessageV1(text_message),
MessageNum::AnthropicMessageV1(text_doc_message),
],
"claude-3-5-sonnet-20241022",
Provider::Anthropic,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap();
if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
if let crate::anthropic::v1::request::ContentBlock::Document(document_block) =
&msg.content[0].inner
{
match &document_block.source {
crate::anthropic::v1::request::DocumentSource::Text(text_source) => {
assert_eq!(text_source.media_type, "text/plain");
assert_eq!(text_source.data, "Plain text document content");
assert_eq!(text_source.r#type, "text");
}
_ => panic!("Expected Text document source"),
}
assert_eq!(
document_block.context.as_ref().unwrap(),
"Context for the document"
);
} else {
panic!("Expected DocumentBlock");
}
} else {
panic!("Expected AnthropicMessageV1");
}
}
#[test]
fn test_anthropic_mixed_content_prompt() {
let text_message = create_anthropic_text_message();
let pdf_message = create_anthropic_base64_pdf_message();
let text_doc_message = create_anthropic_plain_text_document_message();
let system_message = create_anthropic_system_message();
let prompt = Prompt::new_rs(
vec![
MessageNum::AnthropicMessageV1(text_message),
MessageNum::AnthropicMessageV1(pdf_message),
MessageNum::AnthropicMessageV1(text_doc_message),
],
"claude-3-5-sonnet-20241022",
Provider::Anthropic,
vec![MessageNum::AnthropicMessageV1(system_message)],
None,
None,
ResponseType::Null,
)
.unwrap();
assert_eq!(prompt.request.messages().len(), 3);
assert_eq!(prompt.request.system_instructions().len(), 1);
assert_eq!(prompt.provider, Provider::Anthropic);
assert_eq!(prompt.model, "claude-3-5-sonnet-20241022");
}
#[test]
fn test_gemini_chat_message() {
let text = Part::from_text("Test prompt. ${param1} ${param2}".to_string());
let message = GeminiContent {
role: "user".to_string(),
parts: vec![text],
};
let prompt = Prompt::new_rs(
vec![MessageNum::GeminiContentV1(message)],
"gemini-1.5-pro",
Provider::Google,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap();
assert_eq!(prompt.request.messages().len(), 1);
assert_eq!(prompt.parameters.len(), 2);
let mut parameters = prompt.parameters.clone();
parameters.sort();
assert_eq!(parameters[0], "param1");
assert_eq!(parameters[1], "param2");
let bound_msg = prompt.request.messages()[0]
.bind("param1", "Value1")
.unwrap();
let bound_msg = bound_msg.bind("param2", "Value2").unwrap();
match bound_msg {
MessageNum::GeminiContentV1(msg) => {
if let DataNum::Text(text_part) = &msg.parts[0].data {
assert_eq!(text_part, "Test prompt. Value1 Value2");
} else {
panic!("Expected Text Part");
}
}
_ => panic!("Expected GeminiContentV1"),
}
}
#[test]
fn test_from_path_with_base_resolves_missing_extension() {
let temp_dir = create_temp_prompt_dir();
let prompt_path = temp_dir.join("prompt.yaml");
write_generic_prompt(&prompt_path, "openai", "gpt-4o", "Hello ${name}");
let prompt = Prompt::from_path_with_base("prompt", &temp_dir).unwrap();
assert_eq!(prompt.model, "gpt-4o");
assert_eq!(prompt.provider, Provider::OpenAI);
assert_eq!(prompt.parameters, vec!["name".to_string()]);
fs::remove_dir_all(temp_dir).unwrap();
}
#[test]
fn test_from_path_reports_ambiguous_matches() {
let temp_dir = create_temp_prompt_dir();
write_generic_prompt(&temp_dir.join("prompt.yaml"), "openai", "gpt-4o", "Hello");
fs::write(
temp_dir.join("prompt.json"),
r#"{"model":"gpt-4o","provider":"openai","messages":["Hello"]}"#,
)
.unwrap();
let error = Prompt::from_path_with_base("prompt", &temp_dir).unwrap_err();
match error {
TypeError::AmbiguousPromptPath {
requested_path,
candidate_paths,
} => {
assert_eq!(requested_path, "prompt");
assert!(candidate_paths.contains("prompt.yaml"));
assert!(candidate_paths.contains("prompt.json"));
}
other => panic!("expected AmbiguousPromptPath, got {other:?}"),
}
fs::remove_dir_all(temp_dir).unwrap();
}
#[test]
fn test_from_path_reports_attempted_paths_when_missing() {
let temp_dir = create_temp_prompt_dir();
let missing_path = temp_dir.join("missing_prompt");
let error = Prompt::from_path(missing_path.clone()).unwrap_err();
match error {
TypeError::PromptPathNotFound {
requested_path,
attempted_paths,
} => {
assert_eq!(requested_path, missing_path.display().to_string());
assert!(attempted_paths.contains("missing_prompt"));
assert!(attempted_paths.contains("missing_prompt.yaml"));
assert!(attempted_paths.contains("missing_prompt.yml"));
assert!(attempted_paths.contains("missing_prompt.json"));
}
other => panic!("expected PromptPathNotFound, got {other:?}"),
}
fs::remove_dir_all(temp_dir).unwrap();
}
#[test]
fn test_save_prompt_returns_written_json_path() {
let temp_dir = create_temp_prompt_dir();
let text_part = TextContentPart::new("Hello ${name}".to_string());
let content_part = ContentPart::Text(text_part);
let message = OpenAIChatMessage {
role: "user".to_string(),
content: vec![content_part],
name: None,
};
let prompt = Prompt::new_rs(
vec![MessageNum::OpenAIMessageV1(message)],
"gpt-4o",
Provider::OpenAI,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap();
let saved_path = prompt
.save_prompt(Some(temp_dir.join("saved_prompt")))
.unwrap();
let loaded_prompt = Prompt::from_path(saved_path.clone()).unwrap();
assert_eq!(
saved_path.extension().and_then(|ext| ext.to_str()),
Some("json")
);
assert!(saved_path.is_file());
assert_eq!(loaded_prompt.model, "gpt-4o");
assert_eq!(loaded_prompt.provider, Provider::OpenAI);
fs::remove_dir_all(temp_dir).unwrap();
}
}
#[cfg(test)]
mod media_binding_tests {
use super::*;
use crate::anthropic::v1::request::{
ContentBlock, ContentBlockParam, DocumentSource, ImageSource, MessageParam, TextBlockParam,
};
use crate::google::v1::generate::request::{DataNum, GeminiContent, Part};
use crate::openai::v1::chat::request::{ChatMessage, ContentPart, TextContentPart};
use crate::prompt::media::{MediaKind, MediaRef};
use crate::prompt::types::MessageNum;
use crate::Provider;
use base64::Engine;
fn make_user_anthropic(text: &str) -> MessageNum {
MessageNum::AnthropicMessageV1(MessageParam {
content: vec![ContentBlockParam {
inner: ContentBlock::Text(TextBlockParam::new_rs(text.to_string(), None, None)),
}],
role: "user".to_string(),
})
}
fn make_user_openai(text: &str) -> MessageNum {
MessageNum::OpenAIMessageV1(ChatMessage {
role: "user".to_string(),
content: vec![ContentPart::Text(TextContentPart::new(text.to_string()))],
name: None,
})
}
fn make_user_gemini(text: &str) -> MessageNum {
MessageNum::GeminiContentV1(GeminiContent {
role: "user".to_string(),
parts: vec![Part {
data: DataNum::Text(text.to_string()),
..Default::default()
}],
})
}
fn make_system_openai(text: &str) -> MessageNum {
MessageNum::OpenAIMessageV1(ChatMessage {
role: "developer".to_string(),
content: vec![ContentPart::Text(TextContentPart::new(text.to_string()))],
name: None,
})
}
fn make_system_gemini(text: &str) -> MessageNum {
MessageNum::GeminiContentV1(GeminiContent {
role: "model".to_string(),
parts: vec![Part {
data: DataNum::Text(text.to_string()),
..Default::default()
}],
})
}
fn build_prompt(provider: Provider, model: &str, msg: MessageNum) -> Prompt {
Prompt::new_rs(
vec![msg],
model,
provider,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap()
}
#[test]
fn anthropic_splitter_isolates_token() {
let p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("hello ${media:chart} world"),
);
let blocks = match &p.request.messages()[0] {
MessageNum::AnthropicMessageV1(m) => &m.content,
_ => panic!(),
};
assert_eq!(blocks.len(), 3);
let texts: Vec<String> = blocks
.iter()
.filter_map(|b| match &b.inner {
ContentBlock::Text(t) => Some(t.text.clone()),
_ => None,
})
.collect();
assert_eq!(texts, vec!["hello ", "${media:chart}", " world"]);
}
#[test]
fn openai_splitter_isolates_token() {
let p = build_prompt(
Provider::OpenAI,
"gpt-4o",
make_user_openai("a ${media:x} b"),
);
let parts = match &p.request.messages()[0] {
MessageNum::OpenAIMessageV1(m) => &m.content,
_ => panic!(),
};
assert_eq!(parts.len(), 3);
}
#[test]
fn gemini_splitter_isolates_token() {
let p = build_prompt(
Provider::Gemini,
"gemini-2.0-flash",
make_user_gemini("foo ${media:y}"),
);
let parts = match &p.request.messages()[0] {
MessageNum::GeminiContentV1(m) => &m.parts,
_ => panic!(),
};
assert_eq!(parts.len(), 2);
}
#[test]
fn splitter_handles_token_at_start_and_end() {
let p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${media:a}${media:b}"),
);
let blocks = match &p.request.messages()[0] {
MessageNum::AnthropicMessageV1(m) => &m.content,
_ => panic!(),
};
assert_eq!(blocks.len(), 2);
}
#[test]
fn parameters_and_media_parameters_disjoint() {
let p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${greeting} on ${media:chart} for ${media:doc}"),
);
assert_eq!(p.parameters, vec!["greeting"]);
let mut media = p.media_parameters.clone();
media.sort();
assert_eq!(media, vec!["chart".to_string(), "doc".to_string()]);
}
#[test]
fn duplicate_media_token_dedups() {
let p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${media:x} foo ${media:x}"),
);
assert_eq!(p.media_parameters, vec!["x".to_string()]);
}
#[test]
fn anthropic_image_url() {
let mut p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${media:chart}"),
);
p.bind_media_mut(
"chart",
&MediaRef::image_url("https://x/y.png".into(), None),
)
.unwrap();
let blocks = match &p.request.messages()[0] {
MessageNum::AnthropicMessageV1(m) => &m.content,
_ => panic!(),
};
match &blocks[0].inner {
ContentBlock::Image(b) => match &b.source {
ImageSource::Url(s) => assert_eq!(s.url, "https://x/y.png"),
_ => panic!("expected url source"),
},
_ => panic!("expected image block"),
}
}
#[test]
fn anthropic_image_bytes_roundtrip_serde() {
let mut p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${media:chart}"),
);
p.bind_media_mut(
"chart",
&MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"FAKEPNG"),
)
.unwrap();
let json = serde_json::to_value(&p).unwrap();
let restored: Prompt = serde_json::from_value(json).unwrap();
assert_eq!(p, restored);
}
#[test]
fn anthropic_document_base64() {
let mut p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${media:doc}"),
);
p.bind_media_mut(
"doc",
&MediaRef::from_bytes(MediaKind::Document, "application/pdf".into(), b"%PDF"),
)
.unwrap();
let blocks = match &p.request.messages()[0] {
MessageNum::AnthropicMessageV1(m) => &m.content,
_ => panic!(),
};
match &blocks[0].inner {
ContentBlock::Document(b) => match &b.source {
DocumentSource::Base64(s) => {
assert_eq!(s.media_type, "application/pdf");
assert_eq!(
s.data,
base64::engine::general_purpose::STANDARD.encode(b"%PDF")
);
}
_ => panic!(),
},
_ => panic!(),
}
}
#[test]
fn anthropic_document_url() {
let mut p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${media:doc}"),
);
p.bind_media_mut(
"doc",
&MediaRef::document_url("https://x/y.pdf".into(), None),
)
.unwrap();
let blocks = match &p.request.messages()[0] {
MessageNum::AnthropicMessageV1(m) => &m.content,
_ => panic!(),
};
assert!(matches!(
&blocks[0].inner,
ContentBlock::Document(b)
if matches!(&b.source, DocumentSource::Url(s) if s.url == "https://x/y.pdf")
));
}
#[test]
fn openai_image_bytes_becomes_data_url() {
let mut p = build_prompt(
Provider::OpenAI,
"gpt-4o",
make_user_openai("${media:chart}"),
);
p.bind_media_mut(
"chart",
&MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"FAKE"),
)
.unwrap();
let parts = match &p.request.messages()[0] {
MessageNum::OpenAIMessageV1(m) => &m.content,
_ => panic!(),
};
match &parts[0] {
ContentPart::ImageUrl(p) => {
assert!(p.image_url.url.starts_with("data:image/png;base64,"));
}
_ => panic!("expected image_url"),
}
}
#[test]
fn openai_image_url_passthrough() {
let mut p = build_prompt(Provider::OpenAI, "gpt-4o", make_user_openai("${media:c}"));
p.bind_media_mut("c", &MediaRef::image_url("https://x/y.png".into(), None))
.unwrap();
let parts = match &p.request.messages()[0] {
MessageNum::OpenAIMessageV1(m) => &m.content,
_ => panic!(),
};
assert!(matches!(
&parts[0],
ContentPart::ImageUrl(p) if p.image_url.url == "https://x/y.png"
));
}
#[test]
fn openai_document_url_rejected() {
let mut p = build_prompt(Provider::OpenAI, "gpt-4o", make_user_openai("${media:doc}"));
let err = p
.bind_media_mut(
"doc",
&MediaRef::document_url("https://x/y.pdf".into(), None),
)
.unwrap_err();
assert!(matches!(err, TypeError::UnsupportedMediaForProvider { .. }));
}
#[test]
fn openai_document_bytes_becomes_file_content() {
let mut p = build_prompt(Provider::OpenAI, "gpt-4o", make_user_openai("${media:doc}"));
p.bind_media_mut(
"doc",
&MediaRef::from_bytes(MediaKind::Document, "application/pdf".into(), b"%PDF"),
)
.unwrap();
let parts = match &p.request.messages()[0] {
MessageNum::OpenAIMessageV1(m) => &m.content,
_ => panic!(),
};
match &parts[0] {
ContentPart::FileContent(p) => {
assert_eq!(p.r#type, "file");
assert!(p
.file
.file_data
.as_deref()
.unwrap()
.starts_with("data:application/pdf;base64,"));
}
_ => panic!("expected file content"),
}
}
#[test]
fn gemini_inline_data_from_bytes() {
let mut p = build_prompt(
Provider::Gemini,
"gemini-2.0-flash",
make_user_gemini("${media:chart}"),
);
p.bind_media_mut(
"chart",
&MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"FAKE"),
)
.unwrap();
let parts = match &p.request.messages()[0] {
MessageNum::GeminiContentV1(m) => &m.parts,
_ => panic!(),
};
match &parts[0].data {
DataNum::InlineData(b) => assert_eq!(b.mime_type, "image/png"),
_ => panic!(),
}
}
#[test]
fn gemini_https_url_rejected() {
let mut p = build_prompt(
Provider::Gemini,
"gemini-2.0-flash",
make_user_gemini("${media:c}"),
);
let err = p
.bind_media_mut(
"c",
&MediaRef::image_url("https://x/y.png".into(), Some("image/png".into())),
)
.unwrap_err();
assert!(matches!(err, TypeError::UnsupportedMediaForProvider { .. }));
}
#[test]
fn gemini_gs_url_accepted() {
let mut p = build_prompt(
Provider::Gemini,
"gemini-2.0-flash",
make_user_gemini("${media:c}"),
);
p.bind_media_mut(
"c",
&MediaRef::image_url("gs://b/c.png".into(), Some("image/png".into())),
)
.unwrap();
let parts = match &p.request.messages()[0] {
MessageNum::GeminiContentV1(m) => &m.parts,
_ => panic!(),
};
match &parts[0].data {
DataNum::FileData(f) => {
assert_eq!(f.file_uri, "gs://b/c.png");
assert_eq!(f.mime_type, "image/png");
}
_ => panic!(),
}
}
#[test]
fn gemini_url_without_mime_rejected() {
let mut p = build_prompt(
Provider::Gemini,
"gemini-2.0-flash",
make_user_gemini("${media:c}"),
);
let err = p
.bind_media_mut("c", &MediaRef::image_url("gs://b/c.png".into(), None))
.unwrap_err();
assert!(matches!(err, TypeError::InvalidMediaType(_)));
}
#[test]
fn missing_placeholder_errors() {
let mut p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${media:foo}"),
);
let err = p
.bind_media_mut(
"bar",
&MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"X"),
)
.unwrap_err();
assert!(matches!(err, TypeError::MediaPlaceholderNotFound { .. }));
}
#[test]
fn media_in_system_message_rejected_at_construction() {
let sys = MessageNum::AnthropicSystemMessageV1(TextBlockParam::new_rs(
"system ${media:x}".to_string(),
None,
None,
));
let result = Prompt::new_rs(
vec![make_user_anthropic("hi")],
"claude-sonnet-4-5",
Provider::Anthropic,
vec![sys],
None,
None,
ResponseType::Null,
);
assert!(matches!(result, Err(TypeError::MediaInSystemMessage)));
}
#[test]
fn media_in_openai_system_message_rejected_at_construction() {
let result = Prompt::new_rs(
vec![make_user_openai("hi")],
"gpt-4o",
Provider::OpenAI,
vec![make_system_openai("system ${media:x}")],
None,
None,
ResponseType::Null,
);
assert!(matches!(result, Err(TypeError::MediaInSystemMessage)));
}
#[test]
fn media_in_gemini_system_message_rejected_at_construction() {
let result = Prompt::new_rs(
vec![make_user_gemini("hi")],
"gemini-2.0-flash",
Provider::Gemini,
vec![make_system_gemini("system ${media:x}")],
None,
None,
ResponseType::Null,
);
assert!(matches!(result, Err(TypeError::MediaInSystemMessage)));
}
#[test]
fn bind_does_not_touch_media_token() {
let mut p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${name} and ${media:name}"),
);
for m in p.request.messages_mut() {
m.bind_mut("name", "Steven").unwrap();
}
let blocks = match &p.request.messages()[0] {
MessageNum::AnthropicMessageV1(m) => &m.content,
_ => panic!(),
};
match &blocks.last().unwrap().inner {
ContentBlock::Text(t) => assert_eq!(t.text, "${media:name}"),
_ => panic!(),
}
}
#[test]
fn bind_then_bind_media_independent() {
let mut p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${greeting} ${media:img}"),
);
for m in p.request.messages_mut() {
m.bind_mut("greeting", "Hello").unwrap();
}
p.bind_media_mut(
"img",
&MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"X"),
)
.unwrap();
let blocks = match &p.request.messages()[0] {
MessageNum::AnthropicMessageV1(m) => &m.content,
_ => panic!(),
};
let mut saw_text = false;
let mut saw_image = false;
for b in blocks {
match &b.inner {
ContentBlock::Text(t) if t.text.contains("Hello") => saw_text = true,
ContentBlock::Image(_) => saw_image = true,
_ => {}
}
}
assert!(saw_text && saw_image);
}
#[test]
fn multiple_media_bindings_in_sequence() {
let mut p = build_prompt(
Provider::Anthropic,
"claude-sonnet-4-5",
make_user_anthropic("${media:a} ${media:b}"),
);
p.bind_media_mut(
"a",
&MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"A"),
)
.unwrap();
p.bind_media_mut(
"b",
&MediaRef::from_bytes(MediaKind::Document, "application/pdf".into(), b"%PDF"),
)
.unwrap();
let blocks = match &p.request.messages()[0] {
MessageNum::AnthropicMessageV1(m) => &m.content,
_ => panic!(),
};
let mut saw_image = false;
let mut saw_doc = false;
for b in blocks {
match &b.inner {
ContentBlock::Image(_) => saw_image = true,
ContentBlock::Document(_) => saw_doc = true,
_ => {}
}
}
assert!(saw_image && saw_doc);
}
#[test]
fn placeholder_replaced_across_multiple_messages() {
let mut p = Prompt::new_rs(
vec![
make_user_anthropic("${media:x}"),
make_user_anthropic("again ${media:x}"),
],
"claude-sonnet-4-5",
Provider::Anthropic,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap();
p.bind_media_mut(
"x",
&MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"X"),
)
.unwrap();
for msg in p.request.messages() {
let has_image = match msg {
MessageNum::AnthropicMessageV1(m) => m
.content
.iter()
.any(|b| matches!(b.inner, ContentBlock::Image(_))),
_ => false,
};
assert!(has_image);
}
}
}