1use crate::anthropic::v1::request::{AnthropicSettings, MessageParam as AnthropicMessage};
2use crate::error::TypeError;
3use crate::google::v1::generate::request::{GeminiContent, GeminiSettings};
4use crate::openai::v1::chat::request::ChatMessage as OpenAIChatMessage;
5use crate::openai::v1::chat::settings::OpenAIChatSettings;
6use crate::prompt::builder::{to_provider_request, ProviderRequest};
7use crate::prompt::settings::ModelSettings;
8use crate::prompt::types::parse_response_to_json;
9use crate::prompt::types::ResponseType;
10use crate::prompt::types::Role;
11use crate::prompt::{AnthropicMessageList, GeminiContentList, MessageNum, OpenAIMessageList};
12use crate::tools::AgentToolDefinition;
13use crate::traits::MessageFactory;
14use crate::SettingsType;
15use crate::{Provider, SaveName};
16use potato_util::utils::extract_string_value;
17use potato_util::PyHelperFuncs;
18use potatohead_macro::try_extract_message;
19use pyo3::prelude::*;
20use pyo3::types::{PyDict, PyList, PyString, PyTuple};
21use pythonize::pythonize;
22use serde::{Deserialize, Deserializer, Serialize};
23use serde_json::Value;
24use std::collections::BTreeSet;
25use std::path::PathBuf;
26
27fn deserialize_string_or_vec<'de, D>(deserializer: D) -> Result<Vec<String>, D::Error>
30where
31 D: Deserializer<'de>,
32{
33 #[derive(Deserialize)]
34 #[serde(untagged)]
35 enum StringOrVec {
36 Single(String),
37 List(Vec<String>),
38 }
39
40 match StringOrVec::deserialize(deserializer)? {
41 StringOrVec::Single(s) => Ok(vec![s]),
42 StringOrVec::List(v) => Ok(v),
43 }
44}
45
46#[derive(Debug, Deserialize)]
61pub struct GenericPromptConfig {
62 model: String,
63 provider: String,
64 #[serde(deserialize_with = "deserialize_string_or_vec")]
65 messages: Vec<String>,
66 #[serde(default)]
67 system_instructions: Option<Vec<String>>,
68 #[serde(default)]
69 settings: Option<Value>,
70 response_format: Option<Value>,
71}
72
73fn create_message_for_provider(
74 content: String,
75 provider: &Provider,
76 role: &str,
77) -> Result<MessageNum, TypeError> {
78 match provider {
79 Provider::OpenAI => {
80 OpenAIChatMessage::from_text(content, role).map(MessageNum::OpenAIMessageV1)
81 }
82 Provider::Anthropic => {
83 AnthropicMessage::from_text(content, role).map(MessageNum::AnthropicMessageV1)
84 }
85 Provider::Gemini | Provider::Google | Provider::Vertex | Provider::GoogleAdk => {
86 GeminiContent::from_text(content, role).map(MessageNum::GeminiContentV1)
87 }
88 _ => Err(TypeError::Error(format!(
89 "Unsupported provider for message creation: {:?}",
90 provider
91 ))),
92 }
93}
94
95fn parse_single_message(
96 message: &Bound<'_, PyAny>,
97 provider: &Provider,
98 default_role: &str,
99) -> Result<MessageNum, TypeError> {
100 if message.is_instance_of::<PyString>() {
102 let text = message.extract::<String>()?;
103 return create_message_for_provider(text, provider, default_role);
104 }
105
106 try_extract_message!(
108 message,
109 OpenAIChatMessage => MessageNum::OpenAIMessageV1,
110 AnthropicMessage => MessageNum::AnthropicMessageV1,
111 GeminiContent => MessageNum::GeminiContentV1,
112 );
113
114 Err(TypeError::InvalidMessageTypeInList(
115 message.get_type().name()?.to_string(),
116 ))
117}
118
119fn parse_messages(
120 messages: &Bound<'_, PyAny>,
121 provider: &Provider,
122 default_role: &str,
123) -> Result<Vec<MessageNum>, TypeError> {
124 let mut messages =
126 if !messages.is_instance_of::<PyList>() && !messages.is_instance_of::<PyTuple>() {
127 vec![parse_single_message(messages, provider, default_role)?]
128 } else {
129 messages
131 .try_iter()?
132 .map(|item| {
133 let item = item?;
134 parse_single_message(&item, provider, default_role)
135 })
136 .collect::<Result<Vec<_>, _>>()?
137 };
138
139 if provider == &Provider::Anthropic
142 && (default_role == Role::System.as_str()
143 || default_role == Role::Assistant.as_str()
144 || default_role == Role::Developer.as_str())
145 {
146 for msg in messages.iter_mut() {
147 msg.anthropic_message_to_system_message()?;
148 }
149 }
150
151 Ok(messages)
152}
153
154fn get_system_role(provider: &Provider) -> &'static str {
155 match provider {
156 Provider::OpenAI => Role::Developer.into(),
157 Provider::Gemini | Provider::Vertex | Provider::Google | Provider::GoogleAdk => {
158 Role::Model.into()
159 }
160 Provider::Anthropic => Role::System.into(),
161 _ => Role::System.into(),
162 }
163}
164
165pub fn create_system_message_for_provider(
168 content: String,
169 provider: &Provider,
170) -> Result<MessageNum, TypeError> {
171 let role = get_system_role(provider);
172 let mut msg = create_message_for_provider(content, provider, role)?;
173 if provider == &Provider::Anthropic {
175 msg.anthropic_message_to_system_message()?;
176 }
177 Ok(msg)
178}
179
180pub fn extract_system_instructions(
182 system_instruction: Option<&Bound<'_, PyAny>>,
183 provider: &Provider,
184) -> Result<Option<Vec<MessageNum>>, TypeError> {
185 let system_instructions = if let Some(sys_inst) = system_instruction {
186 Some(parse_messages(
187 sys_inst,
188 provider,
189 get_system_role(provider),
190 )?)
191 } else {
192 None
193 };
194
195 Ok(system_instructions)
196}
197
198#[pyclass(from_py_object)]
199#[derive(Debug, Serialize, Clone, PartialEq)]
200pub struct Prompt {
201 pub request: ProviderRequest,
202
203 #[pyo3(get)]
204 pub model: String,
205
206 #[pyo3(get)]
207 pub provider: Provider,
208
209 pub version: String,
210
211 #[pyo3(get)]
212 #[serde(default)]
213 pub parameters: Vec<String>,
214
215 #[serde(default)]
216 pub response_type: ResponseType,
217}
218
219fn extract_model_settings(model_settings: &Bound<'_, PyAny>) -> Result<ModelSettings, TypeError> {
221 let settings_type = model_settings
222 .call_method0("settings_type")?
223 .extract::<SettingsType>()?;
224
225 match settings_type {
226 SettingsType::OpenAIChat => model_settings
227 .extract::<OpenAIChatSettings>()
228 .map(ModelSettings::OpenAIChat),
229 SettingsType::GoogleChat => model_settings
230 .extract::<GeminiSettings>()
231 .map(ModelSettings::GoogleChat),
232 SettingsType::Anthropic => model_settings
233 .extract::<AnthropicSettings>()
234 .map(ModelSettings::AnthropicChat),
235 SettingsType::ModelSettings => model_settings.extract::<ModelSettings>(),
236 }
237 .map_err(Into::into)
238}
239
240#[derive(Debug, Deserialize)]
241#[serde(untagged)]
242enum PromptFormat {
243 Generic(GenericPromptConfig),
244 Full(Box<PromptInternal>),
245}
246
247#[derive(Debug, Deserialize)]
248struct PromptInternal {
249 request: ProviderRequest,
250 model: String,
251 provider: Provider,
252 version: String,
253 #[serde(default)]
254 parameters: Vec<String>,
255 #[serde(default)]
256 response_type: ResponseType,
257}
258
259impl<'de> Deserialize<'de> for Prompt {
260 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
261 where
262 D: serde::Deserializer<'de>,
263 {
264 let format = PromptFormat::deserialize(deserializer)?;
265
266 match format {
267 PromptFormat::Generic(config) => Self::from_generic_config(config)
268 .map_err(|e| serde::de::Error::custom(e.to_string())),
269 PromptFormat::Full(internal) => Ok(Prompt {
270 request: internal.request,
271 model: internal.model,
272 provider: internal.provider,
273 version: internal.version,
274 parameters: internal.parameters,
275 response_type: internal.response_type,
276 }),
277 }
278 }
279}
280
281#[pymethods]
282impl Prompt {
283 #[new]
299 #[pyo3(signature = (messages, model, provider, system_instructions=None, model_settings=None, output_type=None))]
300 pub fn new(
301 py: Python<'_>,
302 messages: &Bound<'_, PyAny>,
303 model: &str,
304 provider: &Bound<'_, PyAny>,
305 system_instructions: Option<&Bound<'_, PyAny>>,
306 model_settings: Option<&Bound<'_, PyAny>>,
307 output_type: Option<&Bound<'_, PyAny>>, ) -> Result<Self, TypeError> {
309 let model_settings = model_settings
311 .as_ref()
312 .map(|s| extract_model_settings(s))
313 .transpose()?;
314
315 let provider = Provider::extract_provider(provider)?;
317
318 let messages = parse_messages(messages, &provider, Role::User.into())?;
321 let system_instructions = if let Some(sys_inst) = system_instructions {
322 parse_messages(sys_inst, &provider, get_system_role(&provider))?
323 } else {
324 vec![]
325 };
326
327 let (response_type, response_json_schema) = match output_type {
329 Some(output_type) => {
330 parse_response_to_json(py, output_type)?
332 }
333 None => (ResponseType::Null, None),
334 };
335
336 Self::new_rs(
337 messages,
338 model,
339 provider,
340 system_instructions,
341 model_settings,
342 response_json_schema,
343 response_type,
344 )
345 }
346
347 #[getter]
348 pub fn model_settings<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
349 self.request.model_settings(py)
350 }
351
352 #[getter]
353 pub fn model_identifier(&self) -> String {
354 format!("{}:{}", self.provider.as_str(), self.model)
355 }
356
357 #[pyo3(signature = (path = None))]
358 pub fn save_prompt(&self, path: Option<PathBuf>) -> Result<PathBuf, TypeError> {
359 let save_path = path.unwrap_or_else(|| PathBuf::from(SaveName::Prompt));
360 PyHelperFuncs::save_to_json(self, &save_path)?;
361 Ok(save_path)
362 }
363
364 #[staticmethod]
365 pub fn from_path(path: PathBuf) -> Result<Self, TypeError> {
366 let content = std::fs::read_to_string(&path)?;
367
368 let extension = path
369 .extension()
370 .and_then(|ext| ext.to_str())
371 .ok_or_else(|| TypeError::Error(format!("Invalid file path: {:?}", path)))?;
372
373 let mut prompt: Prompt = match extension.to_lowercase().as_str() {
374 "json" => serde_json::from_str(&content)?,
375 "yaml" | "yml" => serde_yaml::from_str(&content)?,
376 _ => {
377 return Err(TypeError::Error(format!(
378 "Unsupported file extension '{}'. Expected .json, .yaml, or .yml",
379 extension
380 )))
381 }
382 };
383
384 if prompt.parameters.is_empty() {
385 let system_instructions: Vec<MessageNum> = prompt
386 .request
387 .system_instructions()
388 .iter()
389 .map(|msg| (*msg).clone())
390 .collect();
391 let parameters =
392 Self::extract_variables(prompt.request.messages(), &system_instructions);
393 prompt.parameters = parameters;
394 }
395
396 Ok(prompt)
397 }
398
399 #[staticmethod]
400 pub fn model_validate_json(json_string: String) -> Result<Self, TypeError> {
401 let json_value: Value = serde_json::from_str(&json_string)?;
402 let model: Self = serde_json::from_value(json_value)?;
403
404 Ok(model)
405 }
406
407 pub fn model_dump_json(&self) -> String {
408 serde_json::to_string(self).unwrap()
409 }
410
411 pub fn __str__(&self) -> String {
412 PyHelperFuncs::__str__(self)
413 }
414
415 #[getter]
416 pub fn all_messages<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyList>, TypeError> {
418 self.request.get_all_py_messages(py)
419 }
420
421 #[getter]
422 pub fn messages<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyList>, TypeError> {
424 self.request.get_py_messages(py)
425 }
426
427 #[getter]
428 pub fn message<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
430 self.request.get_py_message(py)
431 }
432
433 #[getter]
434 pub fn openai_messages(&self) -> Result<OpenAIMessageList, TypeError> {
437 if self.provider != Provider::OpenAI {
438 return Err(TypeError::Error(
439 "Prompt provider is not OpenAI".to_string(),
440 ));
441 }
442 let messages = self
443 .request
444 .messages()
445 .iter()
446 .filter(|msg| msg.is_user_message())
447 .filter_map(|msg| match msg {
448 MessageNum::OpenAIMessageV1(m) => Some(m.clone()),
449 _ => None,
450 })
451 .collect::<Vec<_>>();
452 Ok(OpenAIMessageList { messages })
453 }
454
455 #[getter]
456 pub fn openai_message(&self) -> Result<OpenAIChatMessage, TypeError> {
459 if self.provider != Provider::OpenAI {
460 return Err(TypeError::Error(
461 "Prompt provider is not OpenAI".to_string(),
462 ));
463 }
464 self.request.get_openai_message()
465 }
466
467 #[getter]
468 pub fn gemini_messages(&self) -> Result<GeminiContentList, TypeError> {
471 if !self.is_google_provider() {
472 return Err(TypeError::Error(
473 "Prompt provider is not Google, Gemini, or Vertex".to_string(),
474 ));
475 }
476 let messages = self
477 .request
478 .messages()
479 .iter()
480 .filter(|msg| msg.is_user_message())
481 .filter_map(|msg| match msg {
482 MessageNum::GeminiContentV1(m) => Some(m.clone()),
483 _ => None,
484 })
485 .collect::<Vec<_>>();
486
487 Ok(GeminiContentList { messages })
488 }
489
490 #[getter]
491 pub fn gemini_message(&self) -> Result<GeminiContent, TypeError> {
494 if !self.is_google_provider() {
495 return Err(TypeError::Error(
496 "Prompt provider is not Google, Gemini, or Vertex".to_string(),
497 ));
498 }
499 self.request.get_gemini_message()
500 }
501
502 #[getter]
503 pub fn anthropic_messages(&self) -> Result<AnthropicMessageList, TypeError> {
504 if self.provider != Provider::Anthropic {
505 return Err(TypeError::Error(
506 "Prompt provider is not Anthropic".to_string(),
507 ));
508 }
509 let messages = self
510 .request
511 .messages()
512 .iter()
513 .filter(|msg| msg.is_user_message())
514 .filter_map(|msg| match msg {
515 MessageNum::AnthropicMessageV1(m) => Some(m.clone()),
516 _ => None,
517 })
518 .collect::<Vec<_>>();
519
520 Ok(AnthropicMessageList { messages })
521 }
522
523 #[getter]
524 pub fn anthropic_message(&self) -> Result<AnthropicMessage, TypeError> {
526 if self.provider != Provider::Anthropic {
527 return Err(TypeError::Error(
528 "Prompt provider is not Anthropic".to_string(),
529 ));
530 }
531 self.request.get_anthropic_message()
532 }
533
534 #[getter]
535 pub fn system_instructions<'py>(
536 &self,
537 py: Python<'py>,
538 ) -> Result<Bound<'py, PyList>, TypeError> {
539 self.request.get_py_system_instructions(py)
540 }
541
542 #[pyo3(signature = (name=None, value=None, **kwargs))]
550 pub fn bind(
551 &self,
552 name: Option<&str>,
553 value: Option<&Bound<'_, PyAny>>,
554 kwargs: Option<&Bound<'_, PyDict>>,
555 ) -> Result<Self, TypeError> {
556 let mut new_prompt = self.clone();
557
558 if let (Some(name), Some(value)) = (name, value) {
559 let var_value = extract_string_value(value)?;
560 for message in new_prompt.request.messages_mut() {
561 message.bind_mut(name, &var_value)?;
562 }
563 }
564
565 if let Some(kwargs) = kwargs {
566 for (key, val) in kwargs.iter() {
567 let var_name = key.extract::<String>()?;
568 let var_value = extract_string_value(&val)?;
569
570 for message in new_prompt.request.messages_mut() {
571 message.bind_mut(&var_name, &var_value)?;
572 }
573 }
574 }
575
576 if name.is_none() && kwargs.is_none_or(|k| k.is_empty()) {
577 return Err(TypeError::Error(
578 "Must provide either (name, value) or keyword arguments for binding".to_string(),
579 ));
580 }
581
582 Ok(new_prompt)
583 }
584
585 #[pyo3(signature = (name=None, value=None, **kwargs))]
592 pub fn bind_mut(
593 &mut self,
594 name: Option<&str>,
595 value: Option<&Bound<'_, PyAny>>,
596 kwargs: Option<&Bound<'_, PyDict>>,
597 ) -> Result<(), TypeError> {
598 if let (Some(name), Some(value)) = (name, value) {
599 let var_value = extract_string_value(value)?;
600 for message in self.request.messages_mut() {
601 message.bind_mut(name, &var_value)?;
602 }
603 }
604
605 if let Some(kwargs) = kwargs {
606 for (key, val) in kwargs.iter() {
607 let var_name = key.extract::<String>()?;
608 let var_value = extract_string_value(&val)?;
609
610 for message in self.request.messages_mut() {
611 message.bind_mut(&var_name, &var_value)?;
612 }
613 }
614 }
615
616 if name.is_none() && kwargs.is_none_or(|k| k.is_empty()) {
617 return Err(TypeError::Error(
618 "Must provide either (name, value) or keyword arguments for binding".to_string(),
619 ));
620 }
621
622 Ok(())
623 }
624
625 #[getter]
626 pub fn response_json_schema_pretty(&self) -> Option<String> {
627 Some(PyHelperFuncs::__str__(
628 self.request.response_json_schema().as_ref()?,
629 ))
630 }
631
632 #[getter]
633 #[pyo3(name = "response_json_schema")]
634 pub fn response_json_schema_py(&self) -> Option<String> {
635 Some(self.request.response_json_schema().as_ref()?.to_string())
636 }
637
638 pub fn model_dump<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
639 let request = &self.request.to_json()?;
640 Ok(pythonize(py, request)?)
641 }
642}
643
644impl Prompt {
645 pub fn from_generic_config(config: GenericPromptConfig) -> Result<Self, TypeError> {
648 if config.messages.is_empty() {
650 return Err(TypeError::Error(
651 "Prompt has no messages. Generic prompt format requires at least one message."
652 .to_string(),
653 ));
654 }
655
656 let provider = Provider::from_string(&config.provider)?;
658
659 let messages: Vec<MessageNum> = config
661 .messages
662 .into_iter()
663 .map(|msg| create_message_for_provider(msg, &provider, Role::User.as_str()))
664 .collect::<Result<Vec<_>, _>>()?;
665
666 let system_instructions = if let Some(sys_inst) = config.system_instructions {
668 sys_inst
669 .into_iter()
670 .map(|msg| create_message_for_provider(msg, &provider, get_system_role(&provider)))
671 .collect::<Result<Vec<_>, _>>()?
672 } else {
673 Vec::new()
674 };
675
676 let model_settings = if let Some(settings) = config.settings {
678 Some(Self::settings_from_value(settings, &provider)?)
679 } else {
680 None
681 };
682
683 Self::new_rs(
685 messages,
686 &config.model,
687 provider,
688 system_instructions,
689 model_settings,
690 config.response_format,
691 ResponseType::Null,
692 )
693 }
694
695 fn settings_from_value(value: Value, provider: &Provider) -> Result<ModelSettings, TypeError> {
698 match provider {
699 Provider::OpenAI => {
700 let settings: OpenAIChatSettings = serde_json::from_value(value)?;
701 Ok(ModelSettings::OpenAIChat(settings))
702 }
703 Provider::Anthropic => {
704 let settings: AnthropicSettings = serde_json::from_value(value)?;
705 Ok(ModelSettings::AnthropicChat(settings))
706 }
707 Provider::Gemini | Provider::Google | Provider::Vertex | Provider::GoogleAdk => {
708 let settings: GeminiSettings = serde_json::from_value(value)?;
709 Ok(ModelSettings::GoogleChat(settings))
710 }
711 _ => Err(TypeError::Error(format!(
712 "Settings not supported for provider: {:?}",
713 provider
714 ))),
715 }
716 }
717
718 pub fn response_json_schema(&self) -> Option<&Value> {
719 self.request.response_json_schema()
720 }
721
722 pub fn new_rs(
723 messages: Vec<MessageNum>,
724 model: &str,
725 provider: Provider,
726 system_instructions: Vec<MessageNum>,
727 model_settings: Option<ModelSettings>,
728 response_json_schema: Option<Value>,
729 response_type: ResponseType,
730 ) -> Result<Self, TypeError> {
731 let model = model.to_string();
732 let version = potato_util::version();
734 let model_settings = match model_settings {
736 Some(settings) => {
737 settings.validate_provider(&provider)?;
739 settings
740 }
741 None => ModelSettings::provider_default_settings(&provider),
742 };
743
744 let parameters = Self::extract_variables(&messages, &system_instructions);
746
747 let request = to_provider_request(
749 messages,
750 system_instructions,
751 model.clone(),
752 model_settings,
753 response_json_schema,
754 )?;
755
756 Ok(Self {
757 request,
758 version,
759 parameters,
760 response_type,
761 model,
762 provider,
763 })
764 }
765
766 fn is_google_provider(&self) -> bool {
767 matches!(
768 self.provider,
769 Provider::Google | Provider::Gemini | Provider::Vertex | Provider::GoogleAdk
770 )
771 }
772
773 pub fn add_tools(&mut self, tools: Vec<AgentToolDefinition>) -> Result<(), TypeError> {
774 self.request.add_tools(tools)
775 }
776
777 pub fn extract_variables(
778 messages: &[MessageNum],
779 system_instructions: &[MessageNum],
780 ) -> Vec<String> {
781 let mut variables = BTreeSet::new();
782
783 for msg in system_instructions {
785 variables.extend(msg.extract_variables());
786 }
787
788 for msg in messages {
790 variables.extend(msg.extract_variables());
791 }
792
793 variables.into_iter().collect()
794 }
795
796 pub fn model_dump_value(&self) -> Value {
797 serde_json::to_value(self).unwrap_or(Value::Null)
799 }
800
801 pub fn to_request_json(&self) -> Result<Value, TypeError> {
802 let json_value = serde_json::to_value(self)?;
804
805 Ok(json_value)
806 }
807
808 pub fn set_response_json_schema(
809 &mut self,
810 response_json_schema: Option<Value>,
811 response_type: ResponseType,
812 ) {
813 self.request.set_response_json_schema(response_json_schema);
814 self.response_type = response_type;
815 }
816}
817
818#[cfg(test)]
820mod tests {
821 use super::*;
822 use crate::anthropic::v1::request::{
823 Base64ImageSource, Base64PDFSource, ContentBlockParam, DocumentBlockParam, ImageBlockParam,
824 MessageParam, PlainTextSource, TextBlockParam, UrlImageSource, UrlPDFSource,
825 };
826 use crate::google::{DataNum, GeminiContent, Part};
827 use crate::openai::v1::chat::request::{
828 ChatMessage as OpenAIChatMessage, ContentPart, FileContentPart, ImageContentPart,
829 TextContentPart,
830 };
831 use crate::prompt::types::Score;
832 use crate::StructuredOutput;
833
834 fn create_openai_chat_message() -> OpenAIChatMessage {
835 let text_part = TextContentPart::new("What company is this logo from?".to_string());
836 let text_content_part = ContentPart::Text(text_part);
837 OpenAIChatMessage {
838 role: "user".to_string(),
839 content: vec![text_content_part],
840 name: None,
841 }
842 }
843
844 fn create_system_openai_chat_message() -> OpenAIChatMessage {
845 let text_part = TextContentPart::new("system_prompt".to_string());
846 let text_content_part = ContentPart::Text(text_part);
847 OpenAIChatMessage {
848 role: "developer".to_string(),
849 content: vec![text_content_part],
850 name: None,
851 }
852 }
853
854 fn create_openai_image_message() -> OpenAIChatMessage {
855 let image_part = ImageContentPart::new("https://iili.io/3Hs4FMg.png".to_string(), None);
856 let image_content_part = ContentPart::ImageUrl(image_part);
857 OpenAIChatMessage {
858 role: "user".to_string(),
859 content: vec![image_content_part],
860 name: None,
861 }
862 }
863
864 fn create_openai_file_message() -> OpenAIChatMessage {
865 let file_part = FileContentPart::new(
866 Some("filedata".to_string()),
867 Some("fileid".to_string()),
868 Some("filename".to_string()),
869 );
870 let file_content_part = ContentPart::FileContent(file_part);
871 OpenAIChatMessage {
872 role: "user".to_string(),
873 content: vec![file_content_part],
874 name: None,
875 }
876 }
877
878 fn create_anthropic_text_message() -> MessageParam {
879 let text_block =
880 TextBlockParam::new_rs("What company is this logo from?".to_string(), None, None);
881 MessageParam {
882 role: "user".to_string(),
883 content: vec![ContentBlockParam {
884 inner: crate::anthropic::v1::request::ContentBlock::Text(text_block),
885 }],
886 }
887 }
888
889 fn create_anthropic_system_message() -> MessageParam {
890 let text_block = TextBlockParam::new_rs("system_prompt".to_string(), None, None);
891 MessageParam {
892 role: "assistant".to_string(),
893 content: vec![ContentBlockParam {
894 inner: crate::anthropic::v1::request::ContentBlock::Text(text_block),
895 }],
896 }
897 }
898
899 fn create_anthropic_base64_image_message() -> MessageParam {
900 let image_source =
901 Base64ImageSource::new("image/png".to_string(), "base64data".to_string()).unwrap();
902 let image_block = ImageBlockParam {
903 source: crate::anthropic::v1::request::ImageSource::Base64(image_source),
904 cache_control: None,
905 r#type: "image".to_string(),
906 };
907 MessageParam {
908 role: "user".to_string(),
909 content: vec![ContentBlockParam {
910 inner: crate::anthropic::v1::request::ContentBlock::Image(image_block),
911 }],
912 }
913 }
914
915 fn create_anthropic_url_image_message() -> MessageParam {
916 let image_source = UrlImageSource::new("https://iili.io/3Hs4FMg.png".to_string());
917 let image_block = ImageBlockParam {
918 source: crate::anthropic::v1::request::ImageSource::Url(image_source),
919 cache_control: None,
920 r#type: "image".to_string(),
921 };
922 MessageParam {
923 role: "user".to_string(),
924 content: vec![ContentBlockParam {
925 inner: crate::anthropic::v1::request::ContentBlock::Image(image_block),
926 }],
927 }
928 }
929
930 fn create_anthropic_base64_pdf_message() -> MessageParam {
931 let pdf_source = Base64PDFSource::new("base64pdfdata".to_string()).unwrap();
932 let document_block = DocumentBlockParam {
933 source: crate::anthropic::v1::request::DocumentSource::Base64(pdf_source),
934 cache_control: None,
935 title: Some("test_document.pdf".to_string()),
936 context: None,
937 r#type: "document".to_string(),
938 citations: None,
939 };
940 MessageParam {
941 role: "user".to_string(),
942 content: vec![ContentBlockParam {
943 inner: crate::anthropic::v1::request::ContentBlock::Document(document_block),
944 }],
945 }
946 }
947
948 fn create_anthropic_url_pdf_message() -> MessageParam {
949 let pdf_source = UrlPDFSource::new("https://example.com/document.pdf".to_string());
950 let document_block = DocumentBlockParam {
951 source: crate::anthropic::v1::request::DocumentSource::Url(pdf_source),
952 cache_control: None,
953 title: Some("test_document.pdf".to_string()),
954 context: None,
955 r#type: "document".to_string(),
956 citations: None,
957 };
958 MessageParam {
959 role: "user".to_string(),
960 content: vec![ContentBlockParam {
961 inner: crate::anthropic::v1::request::ContentBlock::Document(document_block),
962 }],
963 }
964 }
965
966 fn create_anthropic_plain_text_document_message() -> MessageParam {
967 let text_source = PlainTextSource::new("Plain text document content".to_string());
968 let document_block = DocumentBlockParam {
969 source: crate::anthropic::v1::request::DocumentSource::Text(text_source),
970 cache_control: None,
971 title: Some("text_document.txt".to_string()),
972 context: Some("Context for the document".to_string()),
973 r#type: "document".to_string(),
974 citations: None,
975 };
976 MessageParam {
977 role: "user".to_string(),
978 content: vec![ContentBlockParam {
979 inner: crate::anthropic::v1::request::ContentBlock::Document(document_block),
980 }],
981 }
982 }
983
984 #[test]
985 fn test_task_list_add_and_get() {
986 let text_part = TextContentPart::new("Test prompt. ${param1} ${param2}".to_string());
987 let content_part = ContentPart::Text(text_part);
988 let message = OpenAIChatMessage {
989 role: "user".to_string(),
990 content: vec![content_part],
991 name: None,
992 };
993
994 let prompt = Prompt::new_rs(
995 vec![MessageNum::OpenAIMessageV1(message)],
996 "gpt-4o",
997 Provider::OpenAI,
998 vec![],
999 None,
1000 None,
1001 ResponseType::Null,
1002 )
1003 .unwrap();
1004
1005 assert_eq!(prompt.request.messages().len(), 1);
1007
1008 assert!(prompt.parameters.len() == 2);
1010
1011 let mut parameters = prompt.parameters.clone();
1013 parameters.sort();
1014
1015 assert_eq!(parameters[0], "param1");
1016 assert_eq!(parameters[1], "param2");
1017
1018 let bound_msg = prompt.request.messages()[0]
1020 .bind("param1", "Value1")
1021 .unwrap();
1022 let bound_msg = bound_msg.bind("param2", "Value2").unwrap();
1023
1024 match bound_msg.clone() {
1026 MessageNum::OpenAIMessageV1(msg) => {
1027 if let ContentPart::Text(text_part) = &msg.content[0] {
1028 assert_eq!(text_part.text, "Test prompt. Value1 Value2");
1029 } else {
1030 panic!("Expected TextContentPart");
1031 }
1032 }
1033 _ => panic!("Expected OpenAIMessageV1"),
1034 }
1035 }
1036
1037 #[test]
1038 fn test_image_prompt() {
1039 let text_message = create_openai_chat_message();
1040 let image_message = create_openai_image_message();
1041
1042 let system_text_part = TextContentPart::new("system_prompt".to_string());
1043 let system_text_content_part = ContentPart::Text(system_text_part);
1044
1045 let system_text_message = OpenAIChatMessage {
1046 role: "assistant".to_string(),
1047 content: vec![system_text_content_part],
1048 name: None,
1049 };
1050
1051 let prompt = Prompt::new_rs(
1052 vec![
1053 MessageNum::OpenAIMessageV1(text_message),
1054 MessageNum::OpenAIMessageV1(image_message),
1055 ],
1056 "gpt-4o",
1057 Provider::OpenAI,
1058 vec![MessageNum::OpenAIMessageV1(system_text_message)],
1059 None,
1060 None,
1061 ResponseType::Null,
1062 )
1063 .unwrap();
1064
1065 if let MessageNum::OpenAIMessageV1(msg) = &prompt.request.messages()[1] {
1067 if let ContentPart::Text(text_part) = &msg.content[0] {
1068 assert_eq!(text_part.text, "What company is this logo from?");
1069 } else {
1070 panic!("Expected TextContentPart for the first user message");
1071 }
1072 } else {
1073 panic!("Expected OpenAIMessageV1 for the first user message");
1074 }
1075
1076 if let MessageNum::OpenAIMessageV1(msg) = &prompt.request.messages()[2] {
1078 if let ContentPart::ImageUrl(image_url) = &msg.content[0] {
1079 assert_eq!(image_url.image_url.url, "https://iili.io/3Hs4FMg.png");
1080 assert_eq!(image_url.r#type, "image_url");
1081 } else {
1082 panic!("Expected ContentPart::Image for the second user message");
1083 }
1084 } else {
1085 panic!("Expected OpenAIMessageV1 for the second user message");
1086 }
1087 }
1088
1089 #[test]
1090 fn test_document_prompt() {
1091 let text_message = create_openai_chat_message();
1092 let file_message = create_openai_file_message();
1093 let system_message = create_system_openai_chat_message();
1094
1095 let prompt = Prompt::new_rs(
1096 vec![
1097 MessageNum::OpenAIMessageV1(text_message),
1098 MessageNum::OpenAIMessageV1(file_message),
1099 ],
1100 "gpt-4o",
1101 Provider::OpenAI,
1102 vec![MessageNum::OpenAIMessageV1(system_message)],
1103 None,
1104 None,
1105 ResponseType::Null,
1106 )
1107 .unwrap();
1108
1109 if let MessageNum::OpenAIMessageV1(msg) = &prompt.request.messages()[2] {
1111 if let ContentPart::FileContent(file_content) = &msg.content[0] {
1112 assert_eq!(file_content.file.file_id.as_ref().unwrap(), "fileid");
1113 assert_eq!(file_content.file.filename.as_ref().unwrap(), "filename");
1114 } else {
1115 panic!("Expected ContentPart::FileContent for the second user message");
1116 }
1117 } else {
1118 panic!("Expected OpenAIMessageV1 for the first user message");
1119 }
1120 }
1121
1122 #[test]
1123 fn test_response_format_score() {
1124 let text_message = create_openai_chat_message();
1125 let prompt = Prompt::new_rs(
1126 vec![MessageNum::OpenAIMessageV1(text_message)],
1127 "gpt-4o",
1128 Provider::OpenAI,
1129 vec![],
1130 None,
1131 Some(Score::get_structured_output_schema()),
1132 ResponseType::Null,
1133 )
1134 .unwrap();
1135
1136 assert!(prompt.response_json_schema().is_some());
1138 }
1139
1140 #[test]
1141 fn test_anthropic_text_message_binding() {
1142 let text_block =
1143 TextBlockParam::new_rs("Test prompt. ${param1} ${param2}".to_string(), None, None);
1144 let message = MessageParam {
1145 role: "user".to_string(),
1146 content: vec![ContentBlockParam {
1147 inner: crate::anthropic::v1::request::ContentBlock::Text(text_block),
1148 }],
1149 };
1150
1151 let prompt = Prompt::new_rs(
1152 vec![MessageNum::AnthropicMessageV1(message)],
1153 "claude-3-5-sonnet-20241022",
1154 Provider::Anthropic,
1155 vec![],
1156 None,
1157 None,
1158 ResponseType::Null,
1159 )
1160 .unwrap();
1161
1162 assert_eq!(prompt.request.messages().len(), 1);
1163 assert_eq!(prompt.parameters.len(), 2);
1164
1165 let mut parameters = prompt.parameters.clone();
1166 parameters.sort();
1167 assert_eq!(parameters[0], "param1");
1168 assert_eq!(parameters[1], "param2");
1169
1170 let bound_msg = prompt.request.messages()[0]
1172 .bind("param1", "Value1")
1173 .unwrap();
1174 let bound_msg = bound_msg.bind("param2", "Value2").unwrap();
1175
1176 match bound_msg {
1177 MessageNum::AnthropicMessageV1(msg) => {
1178 if let crate::anthropic::v1::request::ContentBlock::Text(text_block) =
1179 &msg.content[0].inner
1180 {
1181 assert_eq!(text_block.text, "Test prompt. Value1 Value2");
1182 } else {
1183 panic!("Expected TextBlockParam");
1184 }
1185 }
1186 _ => panic!("Expected AnthropicMessageV1"),
1187 }
1188 }
1189
1190 #[test]
1191 fn test_anthropic_url_image_prompt() {
1192 let text_message = create_anthropic_text_message();
1193 let image_message = create_anthropic_url_image_message();
1194 let system_message = create_anthropic_system_message();
1195
1196 let prompt = Prompt::new_rs(
1197 vec![
1198 MessageNum::AnthropicMessageV1(text_message),
1199 MessageNum::AnthropicMessageV1(image_message),
1200 ],
1201 "claude-3-5-sonnet-20241022",
1202 Provider::Anthropic,
1203 vec![MessageNum::AnthropicMessageV1(system_message)],
1204 None,
1205 None,
1206 ResponseType::Null,
1207 )
1208 .unwrap();
1209
1210 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[0] {
1212 if let crate::anthropic::v1::request::ContentBlock::Text(text_block) =
1213 &msg.content[0].inner
1214 {
1215 assert_eq!(text_block.text, "What company is this logo from?");
1216 } else {
1217 panic!("Expected TextBlock for first message");
1218 }
1219 } else {
1220 panic!("Expected AnthropicMessageV1");
1221 }
1222
1223 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1225 if let crate::anthropic::v1::request::ContentBlock::Image(image_block) =
1226 &msg.content[0].inner
1227 {
1228 match &image_block.source {
1229 crate::anthropic::v1::request::ImageSource::Url(url_source) => {
1230 assert_eq!(url_source.url, "https://iili.io/3Hs4FMg.png");
1231 assert_eq!(url_source.r#type, "url");
1232 }
1233 _ => panic!("Expected URL image source"),
1234 }
1235 assert_eq!(image_block.r#type, "image");
1236 } else {
1237 panic!("Expected ImageBlock for second message");
1238 }
1239 } else {
1240 panic!("Expected AnthropicMessageV1");
1241 }
1242 }
1243
1244 #[test]
1245 fn test_anthropic_base64_image_prompt() {
1246 let text_message = create_anthropic_text_message();
1247 let image_message = create_anthropic_base64_image_message();
1248
1249 let prompt = Prompt::new_rs(
1250 vec![
1251 MessageNum::AnthropicMessageV1(text_message),
1252 MessageNum::AnthropicMessageV1(image_message),
1253 ],
1254 "claude-3-5-sonnet-20241022",
1255 Provider::Anthropic,
1256 vec![],
1257 None,
1258 None,
1259 ResponseType::Null,
1260 )
1261 .unwrap();
1262
1263 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1265 if let crate::anthropic::v1::request::ContentBlock::Image(image_block) =
1266 &msg.content[0].inner
1267 {
1268 match &image_block.source {
1269 crate::anthropic::v1::request::ImageSource::Base64(base64_source) => {
1270 assert_eq!(base64_source.media_type, "image/png");
1271 assert_eq!(base64_source.data, "base64data");
1272 assert_eq!(base64_source.r#type, "base64");
1273 }
1274 _ => panic!("Expected Base64 image source"),
1275 }
1276 } else {
1277 panic!("Expected ImageBlock");
1278 }
1279 } else {
1280 panic!("Expected AnthropicMessageV1");
1281 }
1282 }
1283
1284 #[test]
1286 fn test_anthropic_base64_pdf_document_prompt() {
1287 let text_message = create_anthropic_text_message();
1288 let pdf_message = create_anthropic_base64_pdf_message();
1289 let system_message = create_anthropic_system_message();
1290
1291 let prompt = Prompt::new_rs(
1292 vec![
1293 MessageNum::AnthropicMessageV1(text_message),
1294 MessageNum::AnthropicMessageV1(pdf_message),
1295 ],
1296 "claude-3-5-sonnet-20241022",
1297 Provider::Anthropic,
1298 vec![MessageNum::AnthropicMessageV1(system_message)],
1299 None,
1300 None,
1301 ResponseType::Null,
1302 )
1303 .unwrap();
1304
1305 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1307 if let crate::anthropic::v1::request::ContentBlock::Document(document_block) =
1308 &msg.content[0].inner
1309 {
1310 match &document_block.source {
1311 crate::anthropic::v1::request::DocumentSource::Base64(pdf_source) => {
1312 assert_eq!(pdf_source.media_type, "application/pdf");
1313 assert_eq!(pdf_source.data, "base64pdfdata");
1314 assert_eq!(pdf_source.r#type, "base64");
1315 }
1316 _ => panic!("Expected Base64 PDF source"),
1317 }
1318 assert_eq!(document_block.r#type, "document");
1319 assert_eq!(document_block.title.as_ref().unwrap(), "test_document.pdf");
1320 } else {
1321 panic!("Expected DocumentBlock");
1322 }
1323 } else {
1324 panic!("Expected AnthropicMessageV1");
1325 }
1326 }
1327
1328 #[test]
1330 fn test_anthropic_url_pdf_document_prompt() {
1331 let text_message = create_anthropic_text_message();
1332 let pdf_message = create_anthropic_url_pdf_message();
1333
1334 let prompt = Prompt::new_rs(
1335 vec![
1336 MessageNum::AnthropicMessageV1(text_message),
1337 MessageNum::AnthropicMessageV1(pdf_message),
1338 ],
1339 "claude-3-5-sonnet-20241022",
1340 Provider::Anthropic,
1341 vec![],
1342 None,
1343 None,
1344 ResponseType::Null,
1345 )
1346 .unwrap();
1347
1348 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1350 if let crate::anthropic::v1::request::ContentBlock::Document(document_block) =
1351 &msg.content[0].inner
1352 {
1353 match &document_block.source {
1354 crate::anthropic::v1::request::DocumentSource::Url(url_source) => {
1355 assert_eq!(url_source.url, "https://example.com/document.pdf");
1356 assert_eq!(url_source.r#type, "url");
1357 }
1358 _ => panic!("Expected URL PDF source"),
1359 }
1360 } else {
1361 panic!("Expected DocumentBlock");
1362 }
1363 } else {
1364 panic!("Expected AnthropicMessageV1");
1365 }
1366 }
1367
1368 #[test]
1370 fn test_anthropic_plain_text_document_prompt() {
1371 let text_message = create_anthropic_text_message();
1372 let text_doc_message = create_anthropic_plain_text_document_message();
1373
1374 let prompt = Prompt::new_rs(
1375 vec![
1376 MessageNum::AnthropicMessageV1(text_message),
1377 MessageNum::AnthropicMessageV1(text_doc_message),
1378 ],
1379 "claude-3-5-sonnet-20241022",
1380 Provider::Anthropic,
1381 vec![],
1382 None,
1383 None,
1384 ResponseType::Null,
1385 )
1386 .unwrap();
1387
1388 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1390 if let crate::anthropic::v1::request::ContentBlock::Document(document_block) =
1391 &msg.content[0].inner
1392 {
1393 match &document_block.source {
1394 crate::anthropic::v1::request::DocumentSource::Text(text_source) => {
1395 assert_eq!(text_source.media_type, "text/plain");
1396 assert_eq!(text_source.data, "Plain text document content");
1397 assert_eq!(text_source.r#type, "text");
1398 }
1399 _ => panic!("Expected Text document source"),
1400 }
1401 assert_eq!(
1402 document_block.context.as_ref().unwrap(),
1403 "Context for the document"
1404 );
1405 } else {
1406 panic!("Expected DocumentBlock");
1407 }
1408 } else {
1409 panic!("Expected AnthropicMessageV1");
1410 }
1411 }
1412
1413 #[test]
1415 fn test_anthropic_mixed_content_prompt() {
1416 let text_message = create_anthropic_text_message();
1417 let pdf_message = create_anthropic_base64_pdf_message();
1418 let text_doc_message = create_anthropic_plain_text_document_message();
1419 let system_message = create_anthropic_system_message();
1420
1421 let prompt = Prompt::new_rs(
1422 vec![
1423 MessageNum::AnthropicMessageV1(text_message),
1424 MessageNum::AnthropicMessageV1(pdf_message),
1425 MessageNum::AnthropicMessageV1(text_doc_message),
1426 ],
1427 "claude-3-5-sonnet-20241022",
1428 Provider::Anthropic,
1429 vec![MessageNum::AnthropicMessageV1(system_message)],
1430 None,
1431 None,
1432 ResponseType::Null,
1433 )
1434 .unwrap();
1435
1436 assert_eq!(prompt.request.messages().len(), 3);
1437 assert_eq!(prompt.request.system_instructions().len(), 1);
1438 assert_eq!(prompt.provider, Provider::Anthropic);
1439 assert_eq!(prompt.model, "claude-3-5-sonnet-20241022");
1440 }
1441
1442 #[test]
1444 fn test_gemini_chat_message() {
1445 let text = Part::from_text("Test prompt. ${param1} ${param2}".to_string());
1446 let message = GeminiContent {
1447 role: "user".to_string(),
1448 parts: vec![text],
1449 };
1450
1451 let prompt = Prompt::new_rs(
1452 vec![MessageNum::GeminiContentV1(message)],
1453 "gemini-1.5-pro",
1454 Provider::Google,
1455 vec![],
1456 None,
1457 None,
1458 ResponseType::Null,
1459 )
1460 .unwrap();
1461
1462 assert_eq!(prompt.request.messages().len(), 1);
1463 assert_eq!(prompt.parameters.len(), 2);
1464
1465 let mut parameters = prompt.parameters.clone();
1466 parameters.sort();
1467 assert_eq!(parameters[0], "param1");
1468 assert_eq!(parameters[1], "param2");
1469
1470 let bound_msg = prompt.request.messages()[0]
1472 .bind("param1", "Value1")
1473 .unwrap();
1474 let bound_msg = bound_msg.bind("param2", "Value2").unwrap();
1475
1476 match bound_msg {
1477 MessageNum::GeminiContentV1(msg) => {
1478 if let DataNum::Text(text_part) = &msg.parts[0].data {
1479 assert_eq!(text_part, "Test prompt. Value1 Value2");
1480 } else {
1481 panic!("Expected Text Part");
1482 }
1483 }
1484 _ => panic!("Expected GeminiContentV1"),
1485 }
1486 }
1487}