Skip to main content

potato_type/prompt/
interface.rs

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
27/// Deserializes `messages` from either a single string or a list of strings.
28/// This allows YAML block scalars (`|`) to be used for multi-line prompts.
29fn 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/// Generic prompt configuration structure for user-friendly YAML/JSON format.
47/// This format allows users to write prompts in a more intuitive way:
48/// ```yaml
49/// model: gemini-1.5-pro
50/// provider: Google
51/// messages:
52///   - "Hello ${variable1}"
53///   - "This is ${variable2}"
54/// settings:
55///   generation_config:
56///     max_output_tokens: 1024
57///     temperature: 0.7
58/// ```
59/// `messages` also accepts a single block-scalar string (YAML `|`).
60#[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    // String conversion (most common case)
101    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 each message type using macro
107    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    // Single message
125    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            // List/tuple of messages
130            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    // Convert Anthropic system messages to TextBlockParam format
140    // optimize this later - maybe
141    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
165/// Create a single system `MessageNum` for the given provider from a plain Rust string.
166/// Useful for `AgentBuilder::system_prompt()` which runs in pure-Rust (no Python context).
167pub 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    // Anthropic system messages must be TextBlockParam
174    if provider == &Provider::Anthropic {
175        msg.anthropic_message_to_system_message()?;
176    }
177    Ok(msg)
178}
179
180/// Helper for extracting system instructions from optional parameter
181pub 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
219/// ModelSettings variant based on the type of settings provided.
220fn 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    /// Creates a new Prompt object.
284    /// Main parsing logic is as follows:
285    /// 1. Extract model settings if provided, otherwise use provider default settings.
286    /// 2. Message and system instructions are expected to be a variant of MessageNum (OpenAIChatMessage, AnthropicMessage or GeminiContent).
287    /// 3. On instantiation, message will be check if is_instance_of pystring. If pystring, provider will be used to map to appropriate message Text type
288    /// 4. If message is a pylist, each item will be checked for is_instance_of pystring or MessageNum variant and converted accordingly.
289    /// 5. If message is a single MessageNum variant, it will be extracted and wrapped in a vec.
290    /// 6. After messages are parsed, a full provider request struct will by built using to_provider_request function.
291    /// # Arguments:
292    /// * `message`: A single message or list of messages representing user input.
293    /// * `model`: The model identifier to use for the prompt.
294    /// * `provider`: The provider to use for the prompt.
295    /// * `system_instruction`: Optional system instruction message or list of messages.
296    /// * `model_settings`: Optional model settings to use for the prompt.
297    /// * `output_type`: Optional output type to enforce structured output.
298    #[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>>, // can be a pydantic model or one of Opsml's predefined outputs
308    ) -> Result<Self, TypeError> {
309        // 1. get model settings if provided
310        let model_settings = model_settings
311            .as_ref()
312            .map(|s| extract_model_settings(s))
313            .transpose()?;
314
315        // 2. extract provider
316        let provider = Provider::extract_provider(provider)?;
317
318        // 3. Parse user messages with "user" role
319        // We'll use this to figure out the type of request struct to create
320        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        // 4.  validate response_json_schema
328        let (response_type, response_json_schema) = match output_type {
329            Some(output_type) => {
330                // check if output_type is a pydantic model and extract the model json schema
331                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    /// Returns all messages as Python objects, including system instructions, user messages, and assistant messages.
417    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    /// Returns User messages as Python objects. This means, system instructions are excluded.
423    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    /// Returns the last User message as a Python object. This means, system instructions are excluded.
429    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    /// Returns the messages as OpenAI ChatMessage Python objects
435    /// This is a helper that provide strict typing when working with OpenAI prompts
436    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    /// Returns the last message as an OpenAI ChatMessage Python object
457    /// This is a helper that provide strict typing when working with OpenAI prompts
458    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    /// Returns the messages as Google GeminiContent Python objects
469    /// This is a helper that provide strict typing when working with Google/Gemini/Vertex prompts
470    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    /// Returns the last message as a Google GeminiContent Python object
492    /// This is a helper that provide strict typing when working with Google/Gemini/Vertex prompts
493    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    /// Returns the last message as an Anthropic MessageParam Python object
525    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    /// Binds a variable in the prompt to a value. This will return a new Prompt with the variable bound to the value.
543    /// This will iterate over all user messages and bind the variable in each message.
544    /// # Arguments:
545    /// * `name`: The name of the variable to bind.
546    /// * `value`: The value to bind the variable to.
547    /// # Returns:
548    /// * `Result<Self, PromptError>`: Returns a new Prompt with the variable bound to the value.
549    #[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    /// Binds a variable in the prompt to a value. This will mutate the current Prompt and bind the variable in each user message.
586    /// # Arguments:
587    /// * `name`: The name of the variable to bind.
588    /// * `value`: The value to bind the variable to.
589    /// # Returns:
590    /// * `Result<(), PromptError>`: Returns Ok(()) on success or an error if the binding fails.
591    #[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    /// Converts a generic prompt configuration to a Prompt instance.
646    /// This handles the user-friendly YAML/JSON format parsing.
647    pub fn from_generic_config(config: GenericPromptConfig) -> Result<Self, TypeError> {
648        // Validate that messages is not empty
649        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        // Parse provider string to Provider enum
657        let provider = Provider::from_string(&config.provider)?;
658
659        // Convert message strings to MessageNum based on provider
660        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        // Convert system instructions if present
667        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        // Convert settings to ModelSettings based on provider
677        let model_settings = if let Some(settings) = config.settings {
678            Some(Self::settings_from_value(settings, &provider)?)
679        } else {
680            None
681        };
682
683        // Create the prompt using new_rs
684        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    /// Converts a JSON Value to ModelSettings based on the provider.
696    /// This handles the provider-specific settings structure from generic configs.
697    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        // get version from crate
733        let version = potato_util::version();
734        // If model_settings is not provided, set model and provider to undefined if missing
735        let model_settings = match model_settings {
736            Some(settings) => {
737                // validates if provider and settings are compatible
738                settings.validate_provider(&provider)?;
739                settings
740            }
741            None => ModelSettings::provider_default_settings(&provider),
742        };
743
744        // extract named parameters in prompt
745        let parameters = Self::extract_variables(&messages, &system_instructions);
746
747        // Build the provider request
748        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        // Extract from system instructions
784        for msg in system_instructions {
785            variables.extend(msg.extract_variables());
786        }
787
788        // Extract from user messages
789        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        // Convert the Prompt to a JSON Value
798        serde_json::to_value(self).unwrap_or(Value::Null)
799    }
800
801    pub fn to_request_json(&self) -> Result<Value, TypeError> {
802        // Convert the Prompt to a JSON Value
803        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// tests
819#[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        // Check if the prompt was created successfully
1006        assert_eq!(prompt.request.messages().len(), 1);
1007
1008        // check prompt parameters
1009        assert!(prompt.parameters.len() == 2);
1010
1011        // sort parameters to ensure order does not affect the test
1012        let mut parameters = prompt.parameters.clone();
1013        parameters.sort();
1014
1015        assert_eq!(parameters[0], "param1");
1016        assert_eq!(parameters[1], "param2");
1017
1018        // bind parameter
1019        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        // Check if the bound message contains the correct values
1025        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        // Check the first user message
1066        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        // Check the second user message (ImageUrl)
1077        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        // Check the 2nd user message (file)
1110        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        // Check if the response json schema is set correctly
1137        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        // Test parameter binding
1171        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        // Check first message (text)
1211        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        // Check second message (image URL)
1224        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        // Check second message (base64 image)
1264        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: Anthropic PDF document (base64)
1285    #[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        // Check second message (PDF document)
1306        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: Anthropic URL PDF document
1329    #[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        // Check second message (URL PDF)
1349        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: Anthropic plain text document
1369    #[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        // Check second message (plain text document)
1389        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: Mixed Anthropic content (text + multiple documents)
1414    #[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    // gemini test
1443    #[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        // Test parameter binding
1471        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}