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::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,
provider: 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>,
}
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>,
#[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: ProviderRequest,
model: String,
provider: Provider,
version: String,
#[serde(default)]
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: internal.request,
model: internal.model,
provider: internal.provider,
version: internal.version,
parameters: internal.parameters,
response_type: internal.response_type,
}),
}
}
}
#[pymethods]
impl Prompt {
#[new]
#[pyo3(signature = (messages, model, provider, system_instructions=None, model_settings=None, output_type=None))]
pub fn new(
py: Python<'_>,
messages: &Bound<'_, PyAny>,
model: &str,
provider: &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::extract_provider(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>) -> PyResult<PathBuf> {
let save_path = path.unwrap_or_else(|| PathBuf::from(SaveName::Prompt));
PyHelperFuncs::save_to_json(self, &save_path)?;
Ok(save_path)
}
#[staticmethod]
pub fn from_path(path: PathBuf) -> 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;
}
Ok(prompt)
}
#[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(())
}
#[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 {
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::from_string(&config.provider)?;
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(
messages: Vec<MessageNum>,
model: &str,
provider: Provider,
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),
};
let parameters = Self::extract_variables(&messages, &system_instructions);
let request = to_provider_request(
messages,
system_instructions,
model.clone(),
model_settings,
response_json_schema,
)?;
Ok(Self {
request,
version,
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 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;
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"),
}
}
}