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::{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)]
57pub struct GenericPromptConfig {
58 model: String,
59 #[serde(default)]
60 provider: Option<String>,
61 #[serde(deserialize_with = "deserialize_string_or_vec")]
62 messages: Vec<String>,
63 #[serde(default)]
64 system_instructions: Option<Vec<String>>,
65 #[serde(default)]
66 settings: Option<Value>,
67 response_format: Option<Value>,
68}
69
70const PROMPT_FILE_EXTENSIONS: [&str; 3] = ["yaml", "yml", "json"];
71
72fn push_attempted_path(attempted_paths: &mut Vec<PathBuf>, path: PathBuf) {
73 if !attempted_paths.iter().any(|existing| existing == &path) {
74 attempted_paths.push(path);
75 }
76}
77
78fn format_candidate_paths(paths: &[PathBuf]) -> String {
79 paths
80 .iter()
81 .map(|path| path.display().to_string())
82 .collect::<Vec<_>>()
83 .join(", ")
84}
85
86fn create_message_for_provider(
87 content: String,
88 provider: &Provider,
89 role: &str,
90) -> Result<MessageNum, TypeError> {
91 match provider {
92 Provider::OpenAI => {
93 OpenAIChatMessage::from_text(content, role).map(MessageNum::OpenAIMessageV1)
94 }
95 Provider::Anthropic => {
96 AnthropicMessage::from_text(content, role).map(MessageNum::AnthropicMessageV1)
97 }
98 Provider::Gemini | Provider::Google | Provider::Vertex | Provider::GoogleAdk => {
99 GeminiContent::from_text(content, role).map(MessageNum::GeminiContentV1)
100 }
101 _ => Err(TypeError::Error(format!(
102 "Unsupported provider for message creation: {:?}",
103 provider
104 ))),
105 }
106}
107
108fn parse_single_message(
109 message: &Bound<'_, PyAny>,
110 provider: &Provider,
111 default_role: &str,
112) -> Result<MessageNum, TypeError> {
113 if message.is_instance_of::<PyString>() {
115 let text = message.extract::<String>()?;
116 return create_message_for_provider(text, provider, default_role);
117 }
118
119 try_extract_message!(
121 message,
122 OpenAIChatMessage => MessageNum::OpenAIMessageV1,
123 AnthropicMessage => MessageNum::AnthropicMessageV1,
124 GeminiContent => MessageNum::GeminiContentV1,
125 );
126
127 Err(TypeError::InvalidMessageTypeInList(
128 message.get_type().name()?.to_string(),
129 ))
130}
131
132fn parse_messages(
133 messages: &Bound<'_, PyAny>,
134 provider: &Provider,
135 default_role: &str,
136) -> Result<Vec<MessageNum>, TypeError> {
137 let mut messages =
139 if !messages.is_instance_of::<PyList>() && !messages.is_instance_of::<PyTuple>() {
140 vec![parse_single_message(messages, provider, default_role)?]
141 } else {
142 messages
144 .try_iter()?
145 .map(|item| {
146 let item = item?;
147 parse_single_message(&item, provider, default_role)
148 })
149 .collect::<Result<Vec<_>, _>>()?
150 };
151
152 if provider == &Provider::Anthropic
155 && (default_role == Role::System.as_str()
156 || default_role == Role::Assistant.as_str()
157 || default_role == Role::Developer.as_str())
158 {
159 for msg in messages.iter_mut() {
160 msg.anthropic_message_to_system_message()?;
161 }
162 }
163
164 Ok(messages)
165}
166
167fn get_system_role(provider: &Provider) -> &'static str {
168 match provider {
169 Provider::OpenAI => Role::Developer.into(),
170 Provider::Gemini | Provider::Vertex | Provider::Google | Provider::GoogleAdk => {
171 Role::Model.into()
172 }
173 Provider::Anthropic => Role::System.into(),
174 _ => Role::System.into(),
175 }
176}
177
178pub fn create_system_message_for_provider(
181 content: String,
182 provider: &Provider,
183) -> Result<MessageNum, TypeError> {
184 let role = get_system_role(provider);
185 let mut msg = create_message_for_provider(content, provider, role)?;
186 if provider == &Provider::Anthropic {
188 msg.anthropic_message_to_system_message()?;
189 }
190 Ok(msg)
191}
192
193pub fn extract_system_instructions(
195 system_instruction: Option<&Bound<'_, PyAny>>,
196 provider: &Provider,
197) -> Result<Option<Vec<MessageNum>>, TypeError> {
198 let system_instructions = if let Some(sys_inst) = system_instruction {
199 Some(parse_messages(
200 sys_inst,
201 provider,
202 get_system_role(provider),
203 )?)
204 } else {
205 None
206 };
207
208 Ok(system_instructions)
209}
210
211#[pyclass(from_py_object)]
212#[derive(Debug, Serialize, Clone, PartialEq)]
213pub struct Prompt {
214 pub request: ProviderRequest,
215
216 #[pyo3(get)]
217 pub model: String,
218
219 #[pyo3(get)]
220 pub provider: Provider,
221
222 pub version: String,
223
224 #[pyo3(get)]
225 #[serde(default)]
226 pub parameters: Vec<String>,
227
228 #[pyo3(get)]
229 #[serde(default)]
230 pub media_parameters: Vec<String>,
231
232 #[serde(default)]
233 pub response_type: ResponseType,
234}
235
236fn extract_model_settings(model_settings: &Bound<'_, PyAny>) -> Result<ModelSettings, TypeError> {
238 let settings_type = model_settings
239 .call_method0("settings_type")?
240 .extract::<SettingsType>()?;
241
242 match settings_type {
243 SettingsType::OpenAIChat => model_settings
244 .extract::<OpenAIChatSettings>()
245 .map(ModelSettings::OpenAIChat),
246 SettingsType::GoogleChat => model_settings
247 .extract::<GeminiSettings>()
248 .map(ModelSettings::GoogleChat),
249 SettingsType::Anthropic => model_settings
250 .extract::<AnthropicSettings>()
251 .map(ModelSettings::AnthropicChat),
252 SettingsType::ModelSettings => model_settings.extract::<ModelSettings>(),
253 }
254 .map_err(Into::into)
255}
256
257#[derive(Debug, Deserialize)]
258#[serde(untagged)]
259enum PromptFormat {
260 Generic(GenericPromptConfig),
261 Full(Box<PromptInternal>),
262}
263
264#[derive(Debug, Deserialize)]
265struct PromptInternal {
266 request: Value,
267 model: String,
268 provider: Provider,
269 version: String,
270 #[serde(default)]
271 parameters: Vec<String>,
272 #[serde(default)]
273 media_parameters: Vec<String>,
274 #[serde(default)]
275 response_type: ResponseType,
276}
277
278impl<'de> Deserialize<'de> for Prompt {
279 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
280 where
281 D: serde::Deserializer<'de>,
282 {
283 let format = PromptFormat::deserialize(deserializer)?;
284
285 match format {
286 PromptFormat::Generic(config) => Self::from_generic_config(config)
287 .map_err(|e| serde::de::Error::custom(e.to_string())),
288 PromptFormat::Full(internal) => Ok(Prompt {
289 request: provider_request_from_value(internal.provider.clone(), internal.request)
290 .map_err(|e| serde::de::Error::custom(e.to_string()))?,
291 model: internal.model,
292 provider: internal.provider,
293 version: internal.version,
294 parameters: internal.parameters,
295 media_parameters: internal.media_parameters,
296 response_type: internal.response_type,
297 }),
298 }
299 }
300}
301
302fn provider_request_from_value(
303 provider: Provider,
304 value: Value,
305) -> Result<ProviderRequest, TypeError> {
306 match provider {
307 Provider::OpenAI => Ok(ProviderRequest::OpenAIV1(serde_json::from_value(value)?)),
308 Provider::Anthropic => Ok(ProviderRequest::AnthropicV1(serde_json::from_value(value)?)),
309 Provider::Gemini | Provider::Google | Provider::Vertex | Provider::GoogleAdk => {
310 Ok(ProviderRequest::GeminiV1(serde_json::from_value(value)?))
311 }
312 Provider::Undefined => Err(TypeError::UnsupportedProviderForRequestCreation),
313 }
314}
315
316#[pymethods]
317impl Prompt {
318 #[new]
334 #[pyo3(signature = (messages, model, provider=None, system_instructions=None, model_settings=None, output_type=None))]
335 pub fn new(
336 py: Python<'_>,
337 messages: &Bound<'_, PyAny>,
338 model: &str,
339 provider: Option<&Bound<'_, PyAny>>,
340 system_instructions: Option<&Bound<'_, PyAny>>,
341 model_settings: Option<&Bound<'_, PyAny>>,
342 output_type: Option<&Bound<'_, PyAny>>, ) -> Result<Self, TypeError> {
344 let model_settings = model_settings
346 .as_ref()
347 .map(|s| extract_model_settings(s))
348 .transpose()?;
349
350 let provider = Provider::resolve_from_py(provider)?;
352
353 let messages = parse_messages(messages, &provider, Role::User.into())?;
356 let system_instructions = if let Some(sys_inst) = system_instructions {
357 parse_messages(sys_inst, &provider, get_system_role(&provider))?
358 } else {
359 vec![]
360 };
361
362 let (response_type, response_json_schema) = match output_type {
364 Some(output_type) => {
365 parse_response_to_json(py, output_type)?
367 }
368 None => (ResponseType::Null, None),
369 };
370
371 Self::new_rs(
372 messages,
373 model,
374 provider,
375 system_instructions,
376 model_settings,
377 response_json_schema,
378 response_type,
379 )
380 }
381
382 #[getter]
383 pub fn model_settings<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
384 self.request.model_settings(py)
385 }
386
387 #[getter]
388 pub fn model_identifier(&self) -> String {
389 format!("{}:{}", self.provider.as_str(), self.model)
390 }
391
392 #[pyo3(signature = (path = None))]
393 pub fn save_prompt(&self, path: Option<PathBuf>) -> Result<PathBuf, TypeError> {
394 let save_path = path.unwrap_or_else(|| PathBuf::from(SaveName::Prompt));
395 PyHelperFuncs::save_to_json(self, &save_path)?;
396 Ok(save_path.with_extension("json"))
397 }
398
399 #[staticmethod]
400 pub fn from_path(path: PathBuf) -> Result<Self, TypeError> {
401 Self::load_prompt_from_path(path.as_path(), None)
402 }
403
404 #[staticmethod]
405 pub fn model_validate_json(json_string: String) -> Result<Self, TypeError> {
406 let json_value: Value = serde_json::from_str(&json_string)?;
407 let model: Self = serde_json::from_value(json_value)?;
408
409 Ok(model)
410 }
411
412 pub fn model_dump_json(&self) -> String {
413 serde_json::to_string(self).unwrap()
414 }
415
416 pub fn __str__(&self) -> String {
417 PyHelperFuncs::__str__(self)
418 }
419
420 #[getter]
421 pub fn all_messages<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyList>, TypeError> {
423 self.request.get_all_py_messages(py)
424 }
425
426 #[getter]
427 pub fn messages<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyList>, TypeError> {
429 self.request.get_py_messages(py)
430 }
431
432 #[getter]
433 pub fn message<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
435 self.request.get_py_message(py)
436 }
437
438 #[getter]
439 pub fn openai_messages(&self) -> Result<OpenAIMessageList, TypeError> {
442 if self.provider != Provider::OpenAI {
443 return Err(TypeError::Error(
444 "Prompt provider is not OpenAI".to_string(),
445 ));
446 }
447 let messages = self
448 .request
449 .messages()
450 .iter()
451 .filter(|msg| msg.is_user_message())
452 .filter_map(|msg| match msg {
453 MessageNum::OpenAIMessageV1(m) => Some(m.clone()),
454 _ => None,
455 })
456 .collect::<Vec<_>>();
457 Ok(OpenAIMessageList { messages })
458 }
459
460 #[getter]
461 pub fn openai_message(&self) -> Result<OpenAIChatMessage, TypeError> {
464 if self.provider != Provider::OpenAI {
465 return Err(TypeError::Error(
466 "Prompt provider is not OpenAI".to_string(),
467 ));
468 }
469 self.request.get_openai_message()
470 }
471
472 #[getter]
473 pub fn gemini_messages(&self) -> Result<GeminiContentList, TypeError> {
476 if !self.is_google_provider() {
477 return Err(TypeError::Error(
478 "Prompt provider is not Google, Gemini, or Vertex".to_string(),
479 ));
480 }
481 let messages = self
482 .request
483 .messages()
484 .iter()
485 .filter(|msg| msg.is_user_message())
486 .filter_map(|msg| match msg {
487 MessageNum::GeminiContentV1(m) => Some(m.clone()),
488 _ => None,
489 })
490 .collect::<Vec<_>>();
491
492 Ok(GeminiContentList { messages })
493 }
494
495 #[getter]
496 pub fn gemini_message(&self) -> Result<GeminiContent, TypeError> {
499 if !self.is_google_provider() {
500 return Err(TypeError::Error(
501 "Prompt provider is not Google, Gemini, or Vertex".to_string(),
502 ));
503 }
504 self.request.get_gemini_message()
505 }
506
507 #[getter]
508 pub fn anthropic_messages(&self) -> Result<AnthropicMessageList, TypeError> {
509 if self.provider != Provider::Anthropic {
510 return Err(TypeError::Error(
511 "Prompt provider is not Anthropic".to_string(),
512 ));
513 }
514 let messages = self
515 .request
516 .messages()
517 .iter()
518 .filter(|msg| msg.is_user_message())
519 .filter_map(|msg| match msg {
520 MessageNum::AnthropicMessageV1(m) => Some(m.clone()),
521 _ => None,
522 })
523 .collect::<Vec<_>>();
524
525 Ok(AnthropicMessageList { messages })
526 }
527
528 #[getter]
529 pub fn anthropic_message(&self) -> Result<AnthropicMessage, TypeError> {
531 if self.provider != Provider::Anthropic {
532 return Err(TypeError::Error(
533 "Prompt provider is not Anthropic".to_string(),
534 ));
535 }
536 self.request.get_anthropic_message()
537 }
538
539 #[getter]
540 pub fn system_instructions<'py>(
541 &self,
542 py: Python<'py>,
543 ) -> Result<Bound<'py, PyList>, TypeError> {
544 self.request.get_py_system_instructions(py)
545 }
546
547 #[pyo3(signature = (name=None, value=None, **kwargs))]
555 pub fn bind(
556 &self,
557 name: Option<&str>,
558 value: Option<&Bound<'_, PyAny>>,
559 kwargs: Option<&Bound<'_, PyDict>>,
560 ) -> Result<Self, TypeError> {
561 let mut new_prompt = self.clone();
562
563 if let (Some(name), Some(value)) = (name, value) {
564 let var_value = extract_string_value(value)?;
565 for message in new_prompt.request.messages_mut() {
566 message.bind_mut(name, &var_value)?;
567 }
568 }
569
570 if let Some(kwargs) = kwargs {
571 for (key, val) in kwargs.iter() {
572 let var_name = key.extract::<String>()?;
573 let var_value = extract_string_value(&val)?;
574
575 for message in new_prompt.request.messages_mut() {
576 message.bind_mut(&var_name, &var_value)?;
577 }
578 }
579 }
580
581 if name.is_none() && kwargs.is_none_or(|k| k.is_empty()) {
582 return Err(TypeError::Error(
583 "Must provide either (name, value) or keyword arguments for binding".to_string(),
584 ));
585 }
586
587 Ok(new_prompt)
588 }
589
590 #[pyo3(signature = (name=None, value=None, **kwargs))]
597 pub fn bind_mut(
598 &mut self,
599 name: Option<&str>,
600 value: Option<&Bound<'_, PyAny>>,
601 kwargs: Option<&Bound<'_, PyDict>>,
602 ) -> Result<(), TypeError> {
603 if let (Some(name), Some(value)) = (name, value) {
604 let var_value = extract_string_value(value)?;
605 for message in self.request.messages_mut() {
606 message.bind_mut(name, &var_value)?;
607 }
608 }
609
610 if let Some(kwargs) = kwargs {
611 for (key, val) in kwargs.iter() {
612 let var_name = key.extract::<String>()?;
613 let var_value = extract_string_value(&val)?;
614
615 for message in self.request.messages_mut() {
616 message.bind_mut(&var_name, &var_value)?;
617 }
618 }
619 }
620
621 if name.is_none() && kwargs.is_none_or(|k| k.is_empty()) {
622 return Err(TypeError::Error(
623 "Must provide either (name, value) or keyword arguments for binding".to_string(),
624 ));
625 }
626
627 Ok(())
628 }
629
630 pub fn bind_media(
631 &self,
632 name: &str,
633 media: &crate::prompt::media::MediaRef,
634 ) -> Result<Self, TypeError> {
635 let mut new_prompt = self.clone();
636 new_prompt.bind_media_mut(name, media)?;
637 Ok(new_prompt)
638 }
639
640 pub fn bind_media_mut(
641 &mut self,
642 name: &str,
643 media: &crate::prompt::media::MediaRef,
644 ) -> Result<(), TypeError> {
645 let token = format!("${{media:{name}}}");
646 let mut found = false;
647 for message in self.request.messages_mut() {
648 if message.bind_media_mut(&token, media, &self.provider)? {
649 found = true;
650 }
651 }
652 if !found {
653 return Err(TypeError::MediaPlaceholderNotFound {
654 name: name.to_string(),
655 });
656 }
657 Ok(())
658 }
659
660 #[getter]
661 pub fn response_json_schema_pretty(&self) -> Option<String> {
662 Some(PyHelperFuncs::__str__(
663 self.request.response_json_schema().as_ref()?,
664 ))
665 }
666
667 #[getter]
668 #[pyo3(name = "response_json_schema")]
669 pub fn response_json_schema_py(&self) -> Option<String> {
670 Some(self.request.response_json_schema().as_ref()?.to_string())
671 }
672
673 pub fn model_dump<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
674 let request = &self.request.to_json()?;
675 Ok(pythonize(py, request)?)
676 }
677}
678
679impl Prompt {
680 fn read_prompt_file(path: &Path) -> Result<Self, TypeError> {
681 let content = std::fs::read_to_string(path)?;
682
683 let extension = path
684 .extension()
685 .and_then(|ext| ext.to_str())
686 .ok_or_else(|| TypeError::Error(format!("Invalid file path: {:?}", path)))?;
687
688 let mut prompt: Prompt = match extension.to_lowercase().as_str() {
689 "json" => serde_json::from_str(&content)?,
690 "yaml" | "yml" => serde_yaml::from_str(&content)?,
691 _ => {
692 return Err(TypeError::Error(format!(
693 "Unsupported file extension '{}'. Expected .json, .yaml, or .yml",
694 extension
695 )))
696 }
697 };
698
699 if prompt.parameters.is_empty() {
700 let system_instructions: Vec<MessageNum> = prompt
701 .request
702 .system_instructions()
703 .iter()
704 .map(|msg| (*msg).clone())
705 .collect();
706 let parameters =
707 Self::extract_variables(prompt.request.messages(), &system_instructions);
708 prompt.parameters = parameters;
709 }
710
711 if prompt.media_parameters.is_empty() {
712 let system_instructions: Vec<MessageNum> = prompt
713 .request
714 .system_instructions()
715 .iter()
716 .map(|msg| (*msg).clone())
717 .collect();
718 let media_parameters =
719 Self::extract_media_variables(prompt.request.messages(), &system_instructions);
720 prompt.media_parameters = media_parameters;
721 }
722
723 Ok(prompt)
724 }
725
726 fn resolve_prompt_candidate(
727 requested_path: &Path,
728 candidate_path: PathBuf,
729 attempted_paths: &mut Vec<PathBuf>,
730 ) -> Result<Option<PathBuf>, TypeError> {
731 push_attempted_path(attempted_paths, candidate_path.clone());
732
733 if candidate_path.is_file() {
734 return Ok(Some(candidate_path));
735 }
736
737 if requested_path.extension().is_some() {
738 return Ok(None);
739 }
740
741 let mut matches = Vec::new();
742
743 for extension in PROMPT_FILE_EXTENSIONS {
744 let extension_candidate = candidate_path.with_extension(extension);
745 push_attempted_path(attempted_paths, extension_candidate.clone());
746
747 if extension_candidate.is_file() {
748 matches.push(extension_candidate);
749 }
750 }
751
752 match matches.len() {
753 0 => Ok(None),
754 1 => Ok(matches.into_iter().next()),
755 _ => Err(TypeError::AmbiguousPromptPath {
756 requested_path: requested_path.display().to_string(),
757 candidate_paths: format_candidate_paths(&matches),
758 }),
759 }
760 }
761
762 fn resolve_prompt_path(path: &Path, base_dir: Option<&Path>) -> Result<PathBuf, TypeError> {
763 let mut attempted_paths = Vec::new();
764
765 if path.is_absolute() {
766 return Self::resolve_prompt_candidate(path, path.to_path_buf(), &mut attempted_paths)?
767 .ok_or_else(|| TypeError::PromptPathNotFound {
768 requested_path: path.display().to_string(),
769 attempted_paths: format_candidate_paths(&attempted_paths),
770 });
771 }
772
773 let mut candidate_roots = Vec::new();
774 if let Some(base_dir) = base_dir {
775 candidate_roots.push(base_dir.to_path_buf());
776 }
777
778 let current_dir = std::env::current_dir()?;
779 if !candidate_roots.iter().any(|root| root == ¤t_dir) {
780 candidate_roots.push(current_dir);
781 }
782
783 for root in candidate_roots {
784 let candidate_path = root.join(path);
785 if let Some(resolved_path) =
786 Self::resolve_prompt_candidate(path, candidate_path, &mut attempted_paths)?
787 {
788 return Ok(resolved_path);
789 }
790 }
791
792 Err(TypeError::PromptPathNotFound {
793 requested_path: path.display().to_string(),
794 attempted_paths: format_candidate_paths(&attempted_paths),
795 })
796 }
797
798 fn load_prompt_from_path(path: &Path, base_dir: Option<&Path>) -> Result<Self, TypeError> {
799 let resolved_path = Self::resolve_prompt_path(path, base_dir)?;
800 Self::read_prompt_file(&resolved_path)
801 }
802
803 pub fn from_path_with_base(
804 path: impl AsRef<Path>,
805 base_dir: impl AsRef<Path>,
806 ) -> Result<Self, TypeError> {
807 Self::load_prompt_from_path(path.as_ref(), Some(base_dir.as_ref()))
808 }
809
810 pub fn from_generic_config(config: GenericPromptConfig) -> Result<Self, TypeError> {
813 if config.messages.is_empty() {
815 return Err(TypeError::Error(
816 "Prompt has no messages. Generic prompt format requires at least one message."
817 .to_string(),
818 ));
819 }
820
821 let provider = Provider::resolve(config.provider.as_deref())?;
823
824 let messages: Vec<MessageNum> = config
826 .messages
827 .into_iter()
828 .map(|msg| create_message_for_provider(msg, &provider, Role::User.as_str()))
829 .collect::<Result<Vec<_>, _>>()?;
830
831 let system_instructions = if let Some(sys_inst) = config.system_instructions {
833 sys_inst
834 .into_iter()
835 .map(|msg| create_message_for_provider(msg, &provider, get_system_role(&provider)))
836 .collect::<Result<Vec<_>, _>>()?
837 } else {
838 Vec::new()
839 };
840
841 let model_settings = if let Some(settings) = config.settings {
843 Some(Self::settings_from_value(settings, &provider)?)
844 } else {
845 None
846 };
847
848 Self::new_rs(
850 messages,
851 &config.model,
852 provider,
853 system_instructions,
854 model_settings,
855 config.response_format,
856 ResponseType::Null,
857 )
858 }
859
860 fn settings_from_value(value: Value, provider: &Provider) -> Result<ModelSettings, TypeError> {
863 match provider {
864 Provider::OpenAI => {
865 let settings: OpenAIChatSettings = serde_json::from_value(value)?;
866 Ok(ModelSettings::OpenAIChat(settings))
867 }
868 Provider::Anthropic => {
869 let settings: AnthropicSettings = serde_json::from_value(value)?;
870 Ok(ModelSettings::AnthropicChat(settings))
871 }
872 Provider::Gemini | Provider::Google | Provider::Vertex | Provider::GoogleAdk => {
873 let settings: GeminiSettings = serde_json::from_value(value)?;
874 Ok(ModelSettings::GoogleChat(settings))
875 }
876 _ => Err(TypeError::Error(format!(
877 "Settings not supported for provider: {:?}",
878 provider
879 ))),
880 }
881 }
882
883 pub fn response_json_schema(&self) -> Option<&Value> {
884 self.request.response_json_schema()
885 }
886
887 pub fn new_rs(
888 mut messages: Vec<MessageNum>,
889 model: &str,
890 provider: Provider,
891 mut system_instructions: Vec<MessageNum>,
892 model_settings: Option<ModelSettings>,
893 response_json_schema: Option<Value>,
894 response_type: ResponseType,
895 ) -> Result<Self, TypeError> {
896 let model = model.to_string();
897 let version = potato_util::version();
899 let model_settings = match model_settings {
901 Some(settings) => {
902 settings.validate_provider(&provider)?;
904 settings
905 }
906 None => ModelSettings::provider_default_settings(&provider),
907 };
908
909 if system_instructions
910 .iter()
911 .any(|msg| !msg.extract_media_variables().is_empty())
912 {
913 return Err(TypeError::MediaInSystemMessage);
914 }
915
916 for msg in messages.iter_mut() {
917 msg.split_media_placeholders()?;
918 }
919 for msg in system_instructions.iter_mut() {
920 msg.split_media_placeholders()?;
921 }
922
923 let parameters = Self::extract_variables(&messages, &system_instructions);
925 let media_parameters = Self::extract_media_variables(&messages, &system_instructions);
926
927 let request = to_provider_request(
929 messages,
930 system_instructions,
931 model.clone(),
932 model_settings,
933 response_json_schema,
934 )?;
935
936 Ok(Self {
937 request,
938 version,
939 parameters,
940 media_parameters,
941 response_type,
942 model,
943 provider,
944 })
945 }
946
947 fn is_google_provider(&self) -> bool {
948 matches!(
949 self.provider,
950 Provider::Google | Provider::Gemini | Provider::Vertex | Provider::GoogleAdk
951 )
952 }
953
954 pub fn add_tools(&mut self, tools: Vec<AgentToolDefinition>) -> Result<(), TypeError> {
955 self.request.add_tools(tools)
956 }
957
958 pub fn extract_variables(
959 messages: &[MessageNum],
960 system_instructions: &[MessageNum],
961 ) -> Vec<String> {
962 let mut variables = BTreeSet::new();
963
964 for msg in system_instructions {
966 variables.extend(msg.extract_variables());
967 }
968
969 for msg in messages {
971 variables.extend(msg.extract_variables());
972 }
973
974 variables.into_iter().collect()
975 }
976
977 pub fn extract_media_variables(
978 messages: &[MessageNum],
979 system_instructions: &[MessageNum],
980 ) -> Vec<String> {
981 let mut variables = BTreeSet::new();
982
983 for msg in system_instructions {
984 variables.extend(msg.extract_media_variables());
985 }
986
987 for msg in messages {
988 variables.extend(msg.extract_media_variables());
989 }
990
991 variables.into_iter().collect()
992 }
993
994 pub fn model_dump_value(&self) -> Value {
995 serde_json::to_value(self).unwrap_or(Value::Null)
997 }
998
999 pub fn to_request_json(&self) -> Result<Value, TypeError> {
1000 let json_value = serde_json::to_value(self)?;
1002
1003 Ok(json_value)
1004 }
1005
1006 pub fn set_response_json_schema(
1007 &mut self,
1008 response_json_schema: Option<Value>,
1009 response_type: ResponseType,
1010 ) {
1011 self.request.set_response_json_schema(response_json_schema);
1012 self.response_type = response_type;
1013 }
1014}
1015
1016#[cfg(test)]
1018mod tests {
1019 use super::*;
1020 use crate::anthropic::v1::request::{
1021 Base64ImageSource, Base64PDFSource, ContentBlockParam, DocumentBlockParam, ImageBlockParam,
1022 MessageParam, PlainTextSource, TextBlockParam, UrlImageSource, UrlPDFSource,
1023 };
1024 use crate::google::{DataNum, GeminiContent, Part};
1025 use crate::openai::v1::chat::request::{
1026 ChatMessage as OpenAIChatMessage, ContentPart, FileContentPart, ImageContentPart,
1027 TextContentPart,
1028 };
1029 use crate::prompt::types::Score;
1030 use crate::StructuredOutput;
1031 use std::fs;
1032 use std::path::{Path, PathBuf};
1033
1034 fn create_temp_prompt_dir() -> PathBuf {
1035 let dir = std::env::temp_dir().join(format!(
1036 "potatohead-prompt-tests-{}",
1037 potato_util::create_uuid7()
1038 ));
1039 fs::create_dir_all(&dir).unwrap();
1040 dir
1041 }
1042
1043 fn write_generic_prompt(path: &Path, provider: &str, model: &str, message: &str) {
1044 let content =
1045 format!("model: {model}\nprovider: {provider}\nmessages:\n - \"{message}\"\n");
1046 fs::write(path, content).unwrap();
1047 }
1048
1049 fn create_openai_chat_message() -> OpenAIChatMessage {
1050 let text_part = TextContentPart::new("What company is this logo from?".to_string());
1051 let text_content_part = ContentPart::Text(text_part);
1052 OpenAIChatMessage {
1053 role: "user".to_string(),
1054 content: vec![text_content_part],
1055 name: None,
1056 }
1057 }
1058
1059 fn create_system_openai_chat_message() -> OpenAIChatMessage {
1060 let text_part = TextContentPart::new("system_prompt".to_string());
1061 let text_content_part = ContentPart::Text(text_part);
1062 OpenAIChatMessage {
1063 role: "developer".to_string(),
1064 content: vec![text_content_part],
1065 name: None,
1066 }
1067 }
1068
1069 fn create_openai_image_message() -> OpenAIChatMessage {
1070 let image_part = ImageContentPart::new("https://iili.io/3Hs4FMg.png".to_string(), None);
1071 let image_content_part = ContentPart::ImageUrl(image_part);
1072 OpenAIChatMessage {
1073 role: "user".to_string(),
1074 content: vec![image_content_part],
1075 name: None,
1076 }
1077 }
1078
1079 fn create_openai_file_message() -> OpenAIChatMessage {
1080 let file_part = FileContentPart::new(
1081 Some("filedata".to_string()),
1082 Some("fileid".to_string()),
1083 Some("filename".to_string()),
1084 );
1085 let file_content_part = ContentPart::FileContent(file_part);
1086 OpenAIChatMessage {
1087 role: "user".to_string(),
1088 content: vec![file_content_part],
1089 name: None,
1090 }
1091 }
1092
1093 fn create_anthropic_text_message() -> MessageParam {
1094 let text_block =
1095 TextBlockParam::new_rs("What company is this logo from?".to_string(), None, None);
1096 MessageParam {
1097 role: "user".to_string(),
1098 content: vec![ContentBlockParam {
1099 inner: crate::anthropic::v1::request::ContentBlock::Text(text_block),
1100 }],
1101 }
1102 }
1103
1104 fn create_anthropic_system_message() -> MessageParam {
1105 let text_block = TextBlockParam::new_rs("system_prompt".to_string(), None, None);
1106 MessageParam {
1107 role: "assistant".to_string(),
1108 content: vec![ContentBlockParam {
1109 inner: crate::anthropic::v1::request::ContentBlock::Text(text_block),
1110 }],
1111 }
1112 }
1113
1114 fn create_anthropic_base64_image_message() -> MessageParam {
1115 let image_source =
1116 Base64ImageSource::new("image/png".to_string(), "base64data".to_string()).unwrap();
1117 let image_block = ImageBlockParam {
1118 source: crate::anthropic::v1::request::ImageSource::Base64(image_source),
1119 cache_control: None,
1120 r#type: "image".to_string(),
1121 };
1122 MessageParam {
1123 role: "user".to_string(),
1124 content: vec![ContentBlockParam {
1125 inner: crate::anthropic::v1::request::ContentBlock::Image(image_block),
1126 }],
1127 }
1128 }
1129
1130 fn create_anthropic_url_image_message() -> MessageParam {
1131 let image_source = UrlImageSource::new("https://iili.io/3Hs4FMg.png".to_string());
1132 let image_block = ImageBlockParam {
1133 source: crate::anthropic::v1::request::ImageSource::Url(image_source),
1134 cache_control: None,
1135 r#type: "image".to_string(),
1136 };
1137 MessageParam {
1138 role: "user".to_string(),
1139 content: vec![ContentBlockParam {
1140 inner: crate::anthropic::v1::request::ContentBlock::Image(image_block),
1141 }],
1142 }
1143 }
1144
1145 fn create_anthropic_base64_pdf_message() -> MessageParam {
1146 let pdf_source = Base64PDFSource::new("base64pdfdata".to_string()).unwrap();
1147 let document_block = DocumentBlockParam {
1148 source: crate::anthropic::v1::request::DocumentSource::Base64(pdf_source),
1149 cache_control: None,
1150 title: Some("test_document.pdf".to_string()),
1151 context: None,
1152 r#type: "document".to_string(),
1153 citations: None,
1154 };
1155 MessageParam {
1156 role: "user".to_string(),
1157 content: vec![ContentBlockParam {
1158 inner: crate::anthropic::v1::request::ContentBlock::Document(document_block),
1159 }],
1160 }
1161 }
1162
1163 fn create_anthropic_url_pdf_message() -> MessageParam {
1164 let pdf_source = UrlPDFSource::new("https://example.com/document.pdf".to_string());
1165 let document_block = DocumentBlockParam {
1166 source: crate::anthropic::v1::request::DocumentSource::Url(pdf_source),
1167 cache_control: None,
1168 title: Some("test_document.pdf".to_string()),
1169 context: None,
1170 r#type: "document".to_string(),
1171 citations: None,
1172 };
1173 MessageParam {
1174 role: "user".to_string(),
1175 content: vec![ContentBlockParam {
1176 inner: crate::anthropic::v1::request::ContentBlock::Document(document_block),
1177 }],
1178 }
1179 }
1180
1181 fn create_anthropic_plain_text_document_message() -> MessageParam {
1182 let text_source = PlainTextSource::new("Plain text document content".to_string());
1183 let document_block = DocumentBlockParam {
1184 source: crate::anthropic::v1::request::DocumentSource::Text(text_source),
1185 cache_control: None,
1186 title: Some("text_document.txt".to_string()),
1187 context: Some("Context for the document".to_string()),
1188 r#type: "document".to_string(),
1189 citations: None,
1190 };
1191 MessageParam {
1192 role: "user".to_string(),
1193 content: vec![ContentBlockParam {
1194 inner: crate::anthropic::v1::request::ContentBlock::Document(document_block),
1195 }],
1196 }
1197 }
1198
1199 #[test]
1200 fn test_task_list_add_and_get() {
1201 let text_part = TextContentPart::new("Test prompt. ${param1} ${param2}".to_string());
1202 let content_part = ContentPart::Text(text_part);
1203 let message = OpenAIChatMessage {
1204 role: "user".to_string(),
1205 content: vec![content_part],
1206 name: None,
1207 };
1208
1209 let prompt = Prompt::new_rs(
1210 vec![MessageNum::OpenAIMessageV1(message)],
1211 "gpt-4o",
1212 Provider::OpenAI,
1213 vec![],
1214 None,
1215 None,
1216 ResponseType::Null,
1217 )
1218 .unwrap();
1219
1220 assert_eq!(prompt.request.messages().len(), 1);
1222
1223 assert!(prompt.parameters.len() == 2);
1225
1226 let mut parameters = prompt.parameters.clone();
1228 parameters.sort();
1229
1230 assert_eq!(parameters[0], "param1");
1231 assert_eq!(parameters[1], "param2");
1232
1233 let bound_msg = prompt.request.messages()[0]
1235 .bind("param1", "Value1")
1236 .unwrap();
1237 let bound_msg = bound_msg.bind("param2", "Value2").unwrap();
1238
1239 match bound_msg.clone() {
1241 MessageNum::OpenAIMessageV1(msg) => {
1242 if let ContentPart::Text(text_part) = &msg.content[0] {
1243 assert_eq!(text_part.text, "Test prompt. Value1 Value2");
1244 } else {
1245 panic!("Expected TextContentPart");
1246 }
1247 }
1248 _ => panic!("Expected OpenAIMessageV1"),
1249 }
1250 }
1251
1252 #[test]
1253 fn test_image_prompt() {
1254 let text_message = create_openai_chat_message();
1255 let image_message = create_openai_image_message();
1256
1257 let system_text_part = TextContentPart::new("system_prompt".to_string());
1258 let system_text_content_part = ContentPart::Text(system_text_part);
1259
1260 let system_text_message = OpenAIChatMessage {
1261 role: "assistant".to_string(),
1262 content: vec![system_text_content_part],
1263 name: None,
1264 };
1265
1266 let prompt = Prompt::new_rs(
1267 vec![
1268 MessageNum::OpenAIMessageV1(text_message),
1269 MessageNum::OpenAIMessageV1(image_message),
1270 ],
1271 "gpt-4o",
1272 Provider::OpenAI,
1273 vec![MessageNum::OpenAIMessageV1(system_text_message)],
1274 None,
1275 None,
1276 ResponseType::Null,
1277 )
1278 .unwrap();
1279
1280 if let MessageNum::OpenAIMessageV1(msg) = &prompt.request.messages()[1] {
1282 if let ContentPart::Text(text_part) = &msg.content[0] {
1283 assert_eq!(text_part.text, "What company is this logo from?");
1284 } else {
1285 panic!("Expected TextContentPart for the first user message");
1286 }
1287 } else {
1288 panic!("Expected OpenAIMessageV1 for the first user message");
1289 }
1290
1291 if let MessageNum::OpenAIMessageV1(msg) = &prompt.request.messages()[2] {
1293 if let ContentPart::ImageUrl(image_url) = &msg.content[0] {
1294 assert_eq!(image_url.image_url.url, "https://iili.io/3Hs4FMg.png");
1295 assert_eq!(image_url.r#type, "image_url");
1296 } else {
1297 panic!("Expected ContentPart::Image for the second user message");
1298 }
1299 } else {
1300 panic!("Expected OpenAIMessageV1 for the second user message");
1301 }
1302 }
1303
1304 #[test]
1305 fn test_document_prompt() {
1306 let text_message = create_openai_chat_message();
1307 let file_message = create_openai_file_message();
1308 let system_message = create_system_openai_chat_message();
1309
1310 let prompt = Prompt::new_rs(
1311 vec![
1312 MessageNum::OpenAIMessageV1(text_message),
1313 MessageNum::OpenAIMessageV1(file_message),
1314 ],
1315 "gpt-4o",
1316 Provider::OpenAI,
1317 vec![MessageNum::OpenAIMessageV1(system_message)],
1318 None,
1319 None,
1320 ResponseType::Null,
1321 )
1322 .unwrap();
1323
1324 if let MessageNum::OpenAIMessageV1(msg) = &prompt.request.messages()[2] {
1326 if let ContentPart::FileContent(file_content) = &msg.content[0] {
1327 assert_eq!(file_content.file.file_id.as_ref().unwrap(), "fileid");
1328 assert_eq!(file_content.file.filename.as_ref().unwrap(), "filename");
1329 } else {
1330 panic!("Expected ContentPart::FileContent for the second user message");
1331 }
1332 } else {
1333 panic!("Expected OpenAIMessageV1 for the first user message");
1334 }
1335 }
1336
1337 #[test]
1338 fn test_response_format_score() {
1339 let text_message = create_openai_chat_message();
1340 let prompt = Prompt::new_rs(
1341 vec![MessageNum::OpenAIMessageV1(text_message)],
1342 "gpt-4o",
1343 Provider::OpenAI,
1344 vec![],
1345 None,
1346 Some(Score::get_structured_output_schema()),
1347 ResponseType::Null,
1348 )
1349 .unwrap();
1350
1351 assert!(prompt.response_json_schema().is_some());
1353 }
1354
1355 #[test]
1356 fn test_anthropic_text_message_binding() {
1357 let text_block =
1358 TextBlockParam::new_rs("Test prompt. ${param1} ${param2}".to_string(), None, None);
1359 let message = MessageParam {
1360 role: "user".to_string(),
1361 content: vec![ContentBlockParam {
1362 inner: crate::anthropic::v1::request::ContentBlock::Text(text_block),
1363 }],
1364 };
1365
1366 let prompt = Prompt::new_rs(
1367 vec![MessageNum::AnthropicMessageV1(message)],
1368 "claude-3-5-sonnet-20241022",
1369 Provider::Anthropic,
1370 vec![],
1371 None,
1372 None,
1373 ResponseType::Null,
1374 )
1375 .unwrap();
1376
1377 assert_eq!(prompt.request.messages().len(), 1);
1378 assert_eq!(prompt.parameters.len(), 2);
1379
1380 let mut parameters = prompt.parameters.clone();
1381 parameters.sort();
1382 assert_eq!(parameters[0], "param1");
1383 assert_eq!(parameters[1], "param2");
1384
1385 let bound_msg = prompt.request.messages()[0]
1387 .bind("param1", "Value1")
1388 .unwrap();
1389 let bound_msg = bound_msg.bind("param2", "Value2").unwrap();
1390
1391 match bound_msg {
1392 MessageNum::AnthropicMessageV1(msg) => {
1393 if let crate::anthropic::v1::request::ContentBlock::Text(text_block) =
1394 &msg.content[0].inner
1395 {
1396 assert_eq!(text_block.text, "Test prompt. Value1 Value2");
1397 } else {
1398 panic!("Expected TextBlockParam");
1399 }
1400 }
1401 _ => panic!("Expected AnthropicMessageV1"),
1402 }
1403 }
1404
1405 #[test]
1406 fn test_anthropic_url_image_prompt() {
1407 let text_message = create_anthropic_text_message();
1408 let image_message = create_anthropic_url_image_message();
1409 let system_message = create_anthropic_system_message();
1410
1411 let prompt = Prompt::new_rs(
1412 vec![
1413 MessageNum::AnthropicMessageV1(text_message),
1414 MessageNum::AnthropicMessageV1(image_message),
1415 ],
1416 "claude-3-5-sonnet-20241022",
1417 Provider::Anthropic,
1418 vec![MessageNum::AnthropicMessageV1(system_message)],
1419 None,
1420 None,
1421 ResponseType::Null,
1422 )
1423 .unwrap();
1424
1425 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[0] {
1427 if let crate::anthropic::v1::request::ContentBlock::Text(text_block) =
1428 &msg.content[0].inner
1429 {
1430 assert_eq!(text_block.text, "What company is this logo from?");
1431 } else {
1432 panic!("Expected TextBlock for first message");
1433 }
1434 } else {
1435 panic!("Expected AnthropicMessageV1");
1436 }
1437
1438 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1440 if let crate::anthropic::v1::request::ContentBlock::Image(image_block) =
1441 &msg.content[0].inner
1442 {
1443 match &image_block.source {
1444 crate::anthropic::v1::request::ImageSource::Url(url_source) => {
1445 assert_eq!(url_source.url, "https://iili.io/3Hs4FMg.png");
1446 assert_eq!(url_source.r#type, "url");
1447 }
1448 _ => panic!("Expected URL image source"),
1449 }
1450 assert_eq!(image_block.r#type, "image");
1451 } else {
1452 panic!("Expected ImageBlock for second message");
1453 }
1454 } else {
1455 panic!("Expected AnthropicMessageV1");
1456 }
1457 }
1458
1459 #[test]
1460 fn test_anthropic_base64_image_prompt() {
1461 let text_message = create_anthropic_text_message();
1462 let image_message = create_anthropic_base64_image_message();
1463
1464 let prompt = Prompt::new_rs(
1465 vec![
1466 MessageNum::AnthropicMessageV1(text_message),
1467 MessageNum::AnthropicMessageV1(image_message),
1468 ],
1469 "claude-3-5-sonnet-20241022",
1470 Provider::Anthropic,
1471 vec![],
1472 None,
1473 None,
1474 ResponseType::Null,
1475 )
1476 .unwrap();
1477
1478 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1480 if let crate::anthropic::v1::request::ContentBlock::Image(image_block) =
1481 &msg.content[0].inner
1482 {
1483 match &image_block.source {
1484 crate::anthropic::v1::request::ImageSource::Base64(base64_source) => {
1485 assert_eq!(base64_source.media_type, "image/png");
1486 assert_eq!(base64_source.data, "base64data");
1487 assert_eq!(base64_source.r#type, "base64");
1488 }
1489 _ => panic!("Expected Base64 image source"),
1490 }
1491 } else {
1492 panic!("Expected ImageBlock");
1493 }
1494 } else {
1495 panic!("Expected AnthropicMessageV1");
1496 }
1497 }
1498
1499 #[test]
1501 fn test_anthropic_base64_pdf_document_prompt() {
1502 let text_message = create_anthropic_text_message();
1503 let pdf_message = create_anthropic_base64_pdf_message();
1504 let system_message = create_anthropic_system_message();
1505
1506 let prompt = Prompt::new_rs(
1507 vec![
1508 MessageNum::AnthropicMessageV1(text_message),
1509 MessageNum::AnthropicMessageV1(pdf_message),
1510 ],
1511 "claude-3-5-sonnet-20241022",
1512 Provider::Anthropic,
1513 vec![MessageNum::AnthropicMessageV1(system_message)],
1514 None,
1515 None,
1516 ResponseType::Null,
1517 )
1518 .unwrap();
1519
1520 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1522 if let crate::anthropic::v1::request::ContentBlock::Document(document_block) =
1523 &msg.content[0].inner
1524 {
1525 match &document_block.source {
1526 crate::anthropic::v1::request::DocumentSource::Base64(pdf_source) => {
1527 assert_eq!(pdf_source.media_type, "application/pdf");
1528 assert_eq!(pdf_source.data, "base64pdfdata");
1529 assert_eq!(pdf_source.r#type, "base64");
1530 }
1531 _ => panic!("Expected Base64 PDF source"),
1532 }
1533 assert_eq!(document_block.r#type, "document");
1534 assert_eq!(document_block.title.as_ref().unwrap(), "test_document.pdf");
1535 } else {
1536 panic!("Expected DocumentBlock");
1537 }
1538 } else {
1539 panic!("Expected AnthropicMessageV1");
1540 }
1541 }
1542
1543 #[test]
1545 fn test_anthropic_url_pdf_document_prompt() {
1546 let text_message = create_anthropic_text_message();
1547 let pdf_message = create_anthropic_url_pdf_message();
1548
1549 let prompt = Prompt::new_rs(
1550 vec![
1551 MessageNum::AnthropicMessageV1(text_message),
1552 MessageNum::AnthropicMessageV1(pdf_message),
1553 ],
1554 "claude-3-5-sonnet-20241022",
1555 Provider::Anthropic,
1556 vec![],
1557 None,
1558 None,
1559 ResponseType::Null,
1560 )
1561 .unwrap();
1562
1563 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1565 if let crate::anthropic::v1::request::ContentBlock::Document(document_block) =
1566 &msg.content[0].inner
1567 {
1568 match &document_block.source {
1569 crate::anthropic::v1::request::DocumentSource::Url(url_source) => {
1570 assert_eq!(url_source.url, "https://example.com/document.pdf");
1571 assert_eq!(url_source.r#type, "url");
1572 }
1573 _ => panic!("Expected URL PDF source"),
1574 }
1575 } else {
1576 panic!("Expected DocumentBlock");
1577 }
1578 } else {
1579 panic!("Expected AnthropicMessageV1");
1580 }
1581 }
1582
1583 #[test]
1585 fn test_anthropic_plain_text_document_prompt() {
1586 let text_message = create_anthropic_text_message();
1587 let text_doc_message = create_anthropic_plain_text_document_message();
1588
1589 let prompt = Prompt::new_rs(
1590 vec![
1591 MessageNum::AnthropicMessageV1(text_message),
1592 MessageNum::AnthropicMessageV1(text_doc_message),
1593 ],
1594 "claude-3-5-sonnet-20241022",
1595 Provider::Anthropic,
1596 vec![],
1597 None,
1598 None,
1599 ResponseType::Null,
1600 )
1601 .unwrap();
1602
1603 if let MessageNum::AnthropicMessageV1(msg) = &prompt.request.messages()[1] {
1605 if let crate::anthropic::v1::request::ContentBlock::Document(document_block) =
1606 &msg.content[0].inner
1607 {
1608 match &document_block.source {
1609 crate::anthropic::v1::request::DocumentSource::Text(text_source) => {
1610 assert_eq!(text_source.media_type, "text/plain");
1611 assert_eq!(text_source.data, "Plain text document content");
1612 assert_eq!(text_source.r#type, "text");
1613 }
1614 _ => panic!("Expected Text document source"),
1615 }
1616 assert_eq!(
1617 document_block.context.as_ref().unwrap(),
1618 "Context for the document"
1619 );
1620 } else {
1621 panic!("Expected DocumentBlock");
1622 }
1623 } else {
1624 panic!("Expected AnthropicMessageV1");
1625 }
1626 }
1627
1628 #[test]
1630 fn test_anthropic_mixed_content_prompt() {
1631 let text_message = create_anthropic_text_message();
1632 let pdf_message = create_anthropic_base64_pdf_message();
1633 let text_doc_message = create_anthropic_plain_text_document_message();
1634 let system_message = create_anthropic_system_message();
1635
1636 let prompt = Prompt::new_rs(
1637 vec![
1638 MessageNum::AnthropicMessageV1(text_message),
1639 MessageNum::AnthropicMessageV1(pdf_message),
1640 MessageNum::AnthropicMessageV1(text_doc_message),
1641 ],
1642 "claude-3-5-sonnet-20241022",
1643 Provider::Anthropic,
1644 vec![MessageNum::AnthropicMessageV1(system_message)],
1645 None,
1646 None,
1647 ResponseType::Null,
1648 )
1649 .unwrap();
1650
1651 assert_eq!(prompt.request.messages().len(), 3);
1652 assert_eq!(prompt.request.system_instructions().len(), 1);
1653 assert_eq!(prompt.provider, Provider::Anthropic);
1654 assert_eq!(prompt.model, "claude-3-5-sonnet-20241022");
1655 }
1656
1657 #[test]
1659 fn test_gemini_chat_message() {
1660 let text = Part::from_text("Test prompt. ${param1} ${param2}".to_string());
1661 let message = GeminiContent {
1662 role: "user".to_string(),
1663 parts: vec![text],
1664 };
1665
1666 let prompt = Prompt::new_rs(
1667 vec![MessageNum::GeminiContentV1(message)],
1668 "gemini-1.5-pro",
1669 Provider::Google,
1670 vec![],
1671 None,
1672 None,
1673 ResponseType::Null,
1674 )
1675 .unwrap();
1676
1677 assert_eq!(prompt.request.messages().len(), 1);
1678 assert_eq!(prompt.parameters.len(), 2);
1679
1680 let mut parameters = prompt.parameters.clone();
1681 parameters.sort();
1682 assert_eq!(parameters[0], "param1");
1683 assert_eq!(parameters[1], "param2");
1684
1685 let bound_msg = prompt.request.messages()[0]
1687 .bind("param1", "Value1")
1688 .unwrap();
1689 let bound_msg = bound_msg.bind("param2", "Value2").unwrap();
1690
1691 match bound_msg {
1692 MessageNum::GeminiContentV1(msg) => {
1693 if let DataNum::Text(text_part) = &msg.parts[0].data {
1694 assert_eq!(text_part, "Test prompt. Value1 Value2");
1695 } else {
1696 panic!("Expected Text Part");
1697 }
1698 }
1699 _ => panic!("Expected GeminiContentV1"),
1700 }
1701 }
1702
1703 #[test]
1704 fn test_from_path_with_base_resolves_missing_extension() {
1705 let temp_dir = create_temp_prompt_dir();
1706 let prompt_path = temp_dir.join("prompt.yaml");
1707 write_generic_prompt(&prompt_path, "openai", "gpt-4o", "Hello ${name}");
1708
1709 let prompt = Prompt::from_path_with_base("prompt", &temp_dir).unwrap();
1710
1711 assert_eq!(prompt.model, "gpt-4o");
1712 assert_eq!(prompt.provider, Provider::OpenAI);
1713 assert_eq!(prompt.parameters, vec!["name".to_string()]);
1714
1715 fs::remove_dir_all(temp_dir).unwrap();
1716 }
1717
1718 #[test]
1719 fn test_from_path_reports_ambiguous_matches() {
1720 let temp_dir = create_temp_prompt_dir();
1721 write_generic_prompt(&temp_dir.join("prompt.yaml"), "openai", "gpt-4o", "Hello");
1722 fs::write(
1723 temp_dir.join("prompt.json"),
1724 r#"{"model":"gpt-4o","provider":"openai","messages":["Hello"]}"#,
1725 )
1726 .unwrap();
1727
1728 let error = Prompt::from_path_with_base("prompt", &temp_dir).unwrap_err();
1729
1730 match error {
1731 TypeError::AmbiguousPromptPath {
1732 requested_path,
1733 candidate_paths,
1734 } => {
1735 assert_eq!(requested_path, "prompt");
1736 assert!(candidate_paths.contains("prompt.yaml"));
1737 assert!(candidate_paths.contains("prompt.json"));
1738 }
1739 other => panic!("expected AmbiguousPromptPath, got {other:?}"),
1740 }
1741
1742 fs::remove_dir_all(temp_dir).unwrap();
1743 }
1744
1745 #[test]
1746 fn test_from_path_reports_attempted_paths_when_missing() {
1747 let temp_dir = create_temp_prompt_dir();
1748 let missing_path = temp_dir.join("missing_prompt");
1749
1750 let error = Prompt::from_path(missing_path.clone()).unwrap_err();
1751
1752 match error {
1753 TypeError::PromptPathNotFound {
1754 requested_path,
1755 attempted_paths,
1756 } => {
1757 assert_eq!(requested_path, missing_path.display().to_string());
1758 assert!(attempted_paths.contains("missing_prompt"));
1759 assert!(attempted_paths.contains("missing_prompt.yaml"));
1760 assert!(attempted_paths.contains("missing_prompt.yml"));
1761 assert!(attempted_paths.contains("missing_prompt.json"));
1762 }
1763 other => panic!("expected PromptPathNotFound, got {other:?}"),
1764 }
1765
1766 fs::remove_dir_all(temp_dir).unwrap();
1767 }
1768
1769 #[test]
1770 fn test_save_prompt_returns_written_json_path() {
1771 let temp_dir = create_temp_prompt_dir();
1772 let text_part = TextContentPart::new("Hello ${name}".to_string());
1773 let content_part = ContentPart::Text(text_part);
1774 let message = OpenAIChatMessage {
1775 role: "user".to_string(),
1776 content: vec![content_part],
1777 name: None,
1778 };
1779 let prompt = Prompt::new_rs(
1780 vec![MessageNum::OpenAIMessageV1(message)],
1781 "gpt-4o",
1782 Provider::OpenAI,
1783 vec![],
1784 None,
1785 None,
1786 ResponseType::Null,
1787 )
1788 .unwrap();
1789
1790 let saved_path = prompt
1791 .save_prompt(Some(temp_dir.join("saved_prompt")))
1792 .unwrap();
1793 let loaded_prompt = Prompt::from_path(saved_path.clone()).unwrap();
1794
1795 assert_eq!(
1796 saved_path.extension().and_then(|ext| ext.to_str()),
1797 Some("json")
1798 );
1799 assert!(saved_path.is_file());
1800 assert_eq!(loaded_prompt.model, "gpt-4o");
1801 assert_eq!(loaded_prompt.provider, Provider::OpenAI);
1802
1803 fs::remove_dir_all(temp_dir).unwrap();
1804 }
1805}
1806
1807#[cfg(test)]
1808mod media_binding_tests {
1809 use super::*;
1810 use crate::anthropic::v1::request::{
1811 ContentBlock, ContentBlockParam, DocumentSource, ImageSource, MessageParam, TextBlockParam,
1812 };
1813 use crate::google::v1::generate::request::{DataNum, GeminiContent, Part};
1814 use crate::openai::v1::chat::request::{ChatMessage, ContentPart, TextContentPart};
1815 use crate::prompt::media::{MediaKind, MediaRef};
1816 use crate::prompt::types::MessageNum;
1817 use crate::Provider;
1818 use base64::Engine;
1819
1820 fn make_user_anthropic(text: &str) -> MessageNum {
1821 MessageNum::AnthropicMessageV1(MessageParam {
1822 content: vec![ContentBlockParam {
1823 inner: ContentBlock::Text(TextBlockParam::new_rs(text.to_string(), None, None)),
1824 }],
1825 role: "user".to_string(),
1826 })
1827 }
1828
1829 fn make_user_openai(text: &str) -> MessageNum {
1830 MessageNum::OpenAIMessageV1(ChatMessage {
1831 role: "user".to_string(),
1832 content: vec![ContentPart::Text(TextContentPart::new(text.to_string()))],
1833 name: None,
1834 })
1835 }
1836
1837 fn make_user_gemini(text: &str) -> MessageNum {
1838 MessageNum::GeminiContentV1(GeminiContent {
1839 role: "user".to_string(),
1840 parts: vec![Part {
1841 data: DataNum::Text(text.to_string()),
1842 ..Default::default()
1843 }],
1844 })
1845 }
1846
1847 fn make_system_openai(text: &str) -> MessageNum {
1848 MessageNum::OpenAIMessageV1(ChatMessage {
1849 role: "developer".to_string(),
1850 content: vec![ContentPart::Text(TextContentPart::new(text.to_string()))],
1851 name: None,
1852 })
1853 }
1854
1855 fn make_system_gemini(text: &str) -> MessageNum {
1856 MessageNum::GeminiContentV1(GeminiContent {
1857 role: "model".to_string(),
1858 parts: vec![Part {
1859 data: DataNum::Text(text.to_string()),
1860 ..Default::default()
1861 }],
1862 })
1863 }
1864
1865 fn build_prompt(provider: Provider, model: &str, msg: MessageNum) -> Prompt {
1866 Prompt::new_rs(
1867 vec![msg],
1868 model,
1869 provider,
1870 vec![],
1871 None,
1872 None,
1873 ResponseType::Null,
1874 )
1875 .unwrap()
1876 }
1877
1878 #[test]
1879 fn anthropic_splitter_isolates_token() {
1880 let p = build_prompt(
1881 Provider::Anthropic,
1882 "claude-sonnet-4-5",
1883 make_user_anthropic("hello ${media:chart} world"),
1884 );
1885 let blocks = match &p.request.messages()[0] {
1886 MessageNum::AnthropicMessageV1(m) => &m.content,
1887 _ => panic!(),
1888 };
1889 assert_eq!(blocks.len(), 3);
1890 let texts: Vec<String> = blocks
1891 .iter()
1892 .filter_map(|b| match &b.inner {
1893 ContentBlock::Text(t) => Some(t.text.clone()),
1894 _ => None,
1895 })
1896 .collect();
1897 assert_eq!(texts, vec!["hello ", "${media:chart}", " world"]);
1898 }
1899
1900 #[test]
1901 fn openai_splitter_isolates_token() {
1902 let p = build_prompt(
1903 Provider::OpenAI,
1904 "gpt-4o",
1905 make_user_openai("a ${media:x} b"),
1906 );
1907 let parts = match &p.request.messages()[0] {
1908 MessageNum::OpenAIMessageV1(m) => &m.content,
1909 _ => panic!(),
1910 };
1911 assert_eq!(parts.len(), 3);
1912 }
1913
1914 #[test]
1915 fn gemini_splitter_isolates_token() {
1916 let p = build_prompt(
1917 Provider::Gemini,
1918 "gemini-2.0-flash",
1919 make_user_gemini("foo ${media:y}"),
1920 );
1921 let parts = match &p.request.messages()[0] {
1922 MessageNum::GeminiContentV1(m) => &m.parts,
1923 _ => panic!(),
1924 };
1925 assert_eq!(parts.len(), 2);
1926 }
1927
1928 #[test]
1929 fn splitter_handles_token_at_start_and_end() {
1930 let p = build_prompt(
1931 Provider::Anthropic,
1932 "claude-sonnet-4-5",
1933 make_user_anthropic("${media:a}${media:b}"),
1934 );
1935 let blocks = match &p.request.messages()[0] {
1936 MessageNum::AnthropicMessageV1(m) => &m.content,
1937 _ => panic!(),
1938 };
1939 assert_eq!(blocks.len(), 2);
1940 }
1941
1942 #[test]
1943 fn parameters_and_media_parameters_disjoint() {
1944 let p = build_prompt(
1945 Provider::Anthropic,
1946 "claude-sonnet-4-5",
1947 make_user_anthropic("${greeting} on ${media:chart} for ${media:doc}"),
1948 );
1949 assert_eq!(p.parameters, vec!["greeting"]);
1950 let mut media = p.media_parameters.clone();
1951 media.sort();
1952 assert_eq!(media, vec!["chart".to_string(), "doc".to_string()]);
1953 }
1954
1955 #[test]
1956 fn duplicate_media_token_dedups() {
1957 let p = build_prompt(
1958 Provider::Anthropic,
1959 "claude-sonnet-4-5",
1960 make_user_anthropic("${media:x} foo ${media:x}"),
1961 );
1962 assert_eq!(p.media_parameters, vec!["x".to_string()]);
1963 }
1964
1965 #[test]
1966 fn anthropic_image_url() {
1967 let mut p = build_prompt(
1968 Provider::Anthropic,
1969 "claude-sonnet-4-5",
1970 make_user_anthropic("${media:chart}"),
1971 );
1972 p.bind_media_mut(
1973 "chart",
1974 &MediaRef::image_url("https://x/y.png".into(), None),
1975 )
1976 .unwrap();
1977 let blocks = match &p.request.messages()[0] {
1978 MessageNum::AnthropicMessageV1(m) => &m.content,
1979 _ => panic!(),
1980 };
1981 match &blocks[0].inner {
1982 ContentBlock::Image(b) => match &b.source {
1983 ImageSource::Url(s) => assert_eq!(s.url, "https://x/y.png"),
1984 _ => panic!("expected url source"),
1985 },
1986 _ => panic!("expected image block"),
1987 }
1988 }
1989
1990 #[test]
1991 fn anthropic_image_bytes_roundtrip_serde() {
1992 let mut p = build_prompt(
1993 Provider::Anthropic,
1994 "claude-sonnet-4-5",
1995 make_user_anthropic("${media:chart}"),
1996 );
1997 p.bind_media_mut(
1998 "chart",
1999 &MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"FAKEPNG"),
2000 )
2001 .unwrap();
2002 let json = serde_json::to_value(&p).unwrap();
2003 let restored: Prompt = serde_json::from_value(json).unwrap();
2004 assert_eq!(p, restored);
2005 }
2006
2007 #[test]
2008 fn anthropic_document_base64() {
2009 let mut p = build_prompt(
2010 Provider::Anthropic,
2011 "claude-sonnet-4-5",
2012 make_user_anthropic("${media:doc}"),
2013 );
2014 p.bind_media_mut(
2015 "doc",
2016 &MediaRef::from_bytes(MediaKind::Document, "application/pdf".into(), b"%PDF"),
2017 )
2018 .unwrap();
2019 let blocks = match &p.request.messages()[0] {
2020 MessageNum::AnthropicMessageV1(m) => &m.content,
2021 _ => panic!(),
2022 };
2023 match &blocks[0].inner {
2024 ContentBlock::Document(b) => match &b.source {
2025 DocumentSource::Base64(s) => {
2026 assert_eq!(s.media_type, "application/pdf");
2027 assert_eq!(
2028 s.data,
2029 base64::engine::general_purpose::STANDARD.encode(b"%PDF")
2030 );
2031 }
2032 _ => panic!(),
2033 },
2034 _ => panic!(),
2035 }
2036 }
2037
2038 #[test]
2039 fn anthropic_document_url() {
2040 let mut p = build_prompt(
2041 Provider::Anthropic,
2042 "claude-sonnet-4-5",
2043 make_user_anthropic("${media:doc}"),
2044 );
2045 p.bind_media_mut(
2046 "doc",
2047 &MediaRef::document_url("https://x/y.pdf".into(), None),
2048 )
2049 .unwrap();
2050 let blocks = match &p.request.messages()[0] {
2051 MessageNum::AnthropicMessageV1(m) => &m.content,
2052 _ => panic!(),
2053 };
2054 assert!(matches!(
2055 &blocks[0].inner,
2056 ContentBlock::Document(b)
2057 if matches!(&b.source, DocumentSource::Url(s) if s.url == "https://x/y.pdf")
2058 ));
2059 }
2060
2061 #[test]
2062 fn openai_image_bytes_becomes_data_url() {
2063 let mut p = build_prompt(
2064 Provider::OpenAI,
2065 "gpt-4o",
2066 make_user_openai("${media:chart}"),
2067 );
2068 p.bind_media_mut(
2069 "chart",
2070 &MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"FAKE"),
2071 )
2072 .unwrap();
2073 let parts = match &p.request.messages()[0] {
2074 MessageNum::OpenAIMessageV1(m) => &m.content,
2075 _ => panic!(),
2076 };
2077 match &parts[0] {
2078 ContentPart::ImageUrl(p) => {
2079 assert!(p.image_url.url.starts_with("data:image/png;base64,"));
2080 }
2081 _ => panic!("expected image_url"),
2082 }
2083 }
2084
2085 #[test]
2086 fn openai_image_url_passthrough() {
2087 let mut p = build_prompt(Provider::OpenAI, "gpt-4o", make_user_openai("${media:c}"));
2088 p.bind_media_mut("c", &MediaRef::image_url("https://x/y.png".into(), None))
2089 .unwrap();
2090 let parts = match &p.request.messages()[0] {
2091 MessageNum::OpenAIMessageV1(m) => &m.content,
2092 _ => panic!(),
2093 };
2094 assert!(matches!(
2095 &parts[0],
2096 ContentPart::ImageUrl(p) if p.image_url.url == "https://x/y.png"
2097 ));
2098 }
2099
2100 #[test]
2101 fn openai_document_url_rejected() {
2102 let mut p = build_prompt(Provider::OpenAI, "gpt-4o", make_user_openai("${media:doc}"));
2103 let err = p
2104 .bind_media_mut(
2105 "doc",
2106 &MediaRef::document_url("https://x/y.pdf".into(), None),
2107 )
2108 .unwrap_err();
2109 assert!(matches!(err, TypeError::UnsupportedMediaForProvider { .. }));
2110 }
2111
2112 #[test]
2113 fn openai_document_bytes_becomes_file_content() {
2114 let mut p = build_prompt(Provider::OpenAI, "gpt-4o", make_user_openai("${media:doc}"));
2115 p.bind_media_mut(
2116 "doc",
2117 &MediaRef::from_bytes(MediaKind::Document, "application/pdf".into(), b"%PDF"),
2118 )
2119 .unwrap();
2120 let parts = match &p.request.messages()[0] {
2121 MessageNum::OpenAIMessageV1(m) => &m.content,
2122 _ => panic!(),
2123 };
2124 match &parts[0] {
2125 ContentPart::FileContent(p) => {
2126 assert_eq!(p.r#type, "file");
2127 assert!(p
2128 .file
2129 .file_data
2130 .as_deref()
2131 .unwrap()
2132 .starts_with("data:application/pdf;base64,"));
2133 }
2134 _ => panic!("expected file content"),
2135 }
2136 }
2137
2138 #[test]
2139 fn gemini_inline_data_from_bytes() {
2140 let mut p = build_prompt(
2141 Provider::Gemini,
2142 "gemini-2.0-flash",
2143 make_user_gemini("${media:chart}"),
2144 );
2145 p.bind_media_mut(
2146 "chart",
2147 &MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"FAKE"),
2148 )
2149 .unwrap();
2150 let parts = match &p.request.messages()[0] {
2151 MessageNum::GeminiContentV1(m) => &m.parts,
2152 _ => panic!(),
2153 };
2154 match &parts[0].data {
2155 DataNum::InlineData(b) => assert_eq!(b.mime_type, "image/png"),
2156 _ => panic!(),
2157 }
2158 }
2159
2160 #[test]
2161 fn gemini_https_url_rejected() {
2162 let mut p = build_prompt(
2163 Provider::Gemini,
2164 "gemini-2.0-flash",
2165 make_user_gemini("${media:c}"),
2166 );
2167 let err = p
2168 .bind_media_mut(
2169 "c",
2170 &MediaRef::image_url("https://x/y.png".into(), Some("image/png".into())),
2171 )
2172 .unwrap_err();
2173 assert!(matches!(err, TypeError::UnsupportedMediaForProvider { .. }));
2174 }
2175
2176 #[test]
2177 fn gemini_gs_url_accepted() {
2178 let mut p = build_prompt(
2179 Provider::Gemini,
2180 "gemini-2.0-flash",
2181 make_user_gemini("${media:c}"),
2182 );
2183 p.bind_media_mut(
2184 "c",
2185 &MediaRef::image_url("gs://b/c.png".into(), Some("image/png".into())),
2186 )
2187 .unwrap();
2188 let parts = match &p.request.messages()[0] {
2189 MessageNum::GeminiContentV1(m) => &m.parts,
2190 _ => panic!(),
2191 };
2192 match &parts[0].data {
2193 DataNum::FileData(f) => {
2194 assert_eq!(f.file_uri, "gs://b/c.png");
2195 assert_eq!(f.mime_type, "image/png");
2196 }
2197 _ => panic!(),
2198 }
2199 }
2200
2201 #[test]
2202 fn gemini_url_without_mime_rejected() {
2203 let mut p = build_prompt(
2204 Provider::Gemini,
2205 "gemini-2.0-flash",
2206 make_user_gemini("${media:c}"),
2207 );
2208 let err = p
2209 .bind_media_mut("c", &MediaRef::image_url("gs://b/c.png".into(), None))
2210 .unwrap_err();
2211 assert!(matches!(err, TypeError::InvalidMediaType(_)));
2212 }
2213
2214 #[test]
2215 fn missing_placeholder_errors() {
2216 let mut p = build_prompt(
2217 Provider::Anthropic,
2218 "claude-sonnet-4-5",
2219 make_user_anthropic("${media:foo}"),
2220 );
2221 let err = p
2222 .bind_media_mut(
2223 "bar",
2224 &MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"X"),
2225 )
2226 .unwrap_err();
2227 assert!(matches!(err, TypeError::MediaPlaceholderNotFound { .. }));
2228 }
2229
2230 #[test]
2231 fn media_in_system_message_rejected_at_construction() {
2232 let sys = MessageNum::AnthropicSystemMessageV1(TextBlockParam::new_rs(
2233 "system ${media:x}".to_string(),
2234 None,
2235 None,
2236 ));
2237 let result = Prompt::new_rs(
2238 vec![make_user_anthropic("hi")],
2239 "claude-sonnet-4-5",
2240 Provider::Anthropic,
2241 vec![sys],
2242 None,
2243 None,
2244 ResponseType::Null,
2245 );
2246 assert!(matches!(result, Err(TypeError::MediaInSystemMessage)));
2247 }
2248
2249 #[test]
2250 fn media_in_openai_system_message_rejected_at_construction() {
2251 let result = Prompt::new_rs(
2252 vec![make_user_openai("hi")],
2253 "gpt-4o",
2254 Provider::OpenAI,
2255 vec![make_system_openai("system ${media:x}")],
2256 None,
2257 None,
2258 ResponseType::Null,
2259 );
2260 assert!(matches!(result, Err(TypeError::MediaInSystemMessage)));
2261 }
2262
2263 #[test]
2264 fn media_in_gemini_system_message_rejected_at_construction() {
2265 let result = Prompt::new_rs(
2266 vec![make_user_gemini("hi")],
2267 "gemini-2.0-flash",
2268 Provider::Gemini,
2269 vec![make_system_gemini("system ${media:x}")],
2270 None,
2271 None,
2272 ResponseType::Null,
2273 );
2274 assert!(matches!(result, Err(TypeError::MediaInSystemMessage)));
2275 }
2276
2277 #[test]
2278 fn bind_does_not_touch_media_token() {
2279 let mut p = build_prompt(
2280 Provider::Anthropic,
2281 "claude-sonnet-4-5",
2282 make_user_anthropic("${name} and ${media:name}"),
2283 );
2284 for m in p.request.messages_mut() {
2285 m.bind_mut("name", "Steven").unwrap();
2286 }
2287 let blocks = match &p.request.messages()[0] {
2288 MessageNum::AnthropicMessageV1(m) => &m.content,
2289 _ => panic!(),
2290 };
2291 match &blocks.last().unwrap().inner {
2292 ContentBlock::Text(t) => assert_eq!(t.text, "${media:name}"),
2293 _ => panic!(),
2294 }
2295 }
2296
2297 #[test]
2298 fn bind_then_bind_media_independent() {
2299 let mut p = build_prompt(
2300 Provider::Anthropic,
2301 "claude-sonnet-4-5",
2302 make_user_anthropic("${greeting} ${media:img}"),
2303 );
2304 for m in p.request.messages_mut() {
2305 m.bind_mut("greeting", "Hello").unwrap();
2306 }
2307 p.bind_media_mut(
2308 "img",
2309 &MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"X"),
2310 )
2311 .unwrap();
2312 let blocks = match &p.request.messages()[0] {
2313 MessageNum::AnthropicMessageV1(m) => &m.content,
2314 _ => panic!(),
2315 };
2316 let mut saw_text = false;
2317 let mut saw_image = false;
2318 for b in blocks {
2319 match &b.inner {
2320 ContentBlock::Text(t) if t.text.contains("Hello") => saw_text = true,
2321 ContentBlock::Image(_) => saw_image = true,
2322 _ => {}
2323 }
2324 }
2325 assert!(saw_text && saw_image);
2326 }
2327
2328 #[test]
2329 fn multiple_media_bindings_in_sequence() {
2330 let mut p = build_prompt(
2331 Provider::Anthropic,
2332 "claude-sonnet-4-5",
2333 make_user_anthropic("${media:a} ${media:b}"),
2334 );
2335 p.bind_media_mut(
2336 "a",
2337 &MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"A"),
2338 )
2339 .unwrap();
2340 p.bind_media_mut(
2341 "b",
2342 &MediaRef::from_bytes(MediaKind::Document, "application/pdf".into(), b"%PDF"),
2343 )
2344 .unwrap();
2345 let blocks = match &p.request.messages()[0] {
2346 MessageNum::AnthropicMessageV1(m) => &m.content,
2347 _ => panic!(),
2348 };
2349 let mut saw_image = false;
2350 let mut saw_doc = false;
2351 for b in blocks {
2352 match &b.inner {
2353 ContentBlock::Image(_) => saw_image = true,
2354 ContentBlock::Document(_) => saw_doc = true,
2355 _ => {}
2356 }
2357 }
2358 assert!(saw_image && saw_doc);
2359 }
2360
2361 #[test]
2362 fn placeholder_replaced_across_multiple_messages() {
2363 let mut p = Prompt::new_rs(
2364 vec![
2365 make_user_anthropic("${media:x}"),
2366 make_user_anthropic("again ${media:x}"),
2367 ],
2368 "claude-sonnet-4-5",
2369 Provider::Anthropic,
2370 vec![],
2371 None,
2372 None,
2373 ResponseType::Null,
2374 )
2375 .unwrap();
2376 p.bind_media_mut(
2377 "x",
2378 &MediaRef::from_bytes(MediaKind::Image, "image/png".into(), b"X"),
2379 )
2380 .unwrap();
2381 for msg in p.request.messages() {
2382 let has_image = match msg {
2383 MessageNum::AnthropicMessageV1(m) => m
2384 .content
2385 .iter()
2386 .any(|b| matches!(b.inner, ContentBlock::Image(_))),
2387 _ => false,
2388 };
2389 assert!(has_image);
2390 }
2391 }
2392}