Skip to main content

potato_type/prompt/
settings.rs

1use crate::anthropic::v1::request::AnthropicSettings;
2use crate::error::TypeError;
3use crate::{
4    google::v1::generate::request::GeminiSettings, openai::v1::chat::settings::OpenAIChatSettings,
5};
6use crate::{Provider, SettingsType};
7use potato_util::PyHelperFuncs;
8use pyo3::prelude::*;
9use pyo3::IntoPyObjectExt;
10use serde::{Deserialize, Serialize};
11use serde_json::Value;
12
13#[pyclass(from_py_object)]
14#[derive(Debug, Serialize, Deserialize, Clone, PartialEq)]
15#[serde(untagged)]
16#[allow(clippy::large_enum_variant)]
17pub enum ModelSettings {
18    OpenAIChat(OpenAIChatSettings),
19    GoogleChat(GeminiSettings),
20    AnthropicChat(AnthropicSettings),
21}
22
23impl Default for ModelSettings {
24    fn default() -> Self {
25        ModelSettings::OpenAIChat(OpenAIChatSettings::default())
26    }
27}
28
29#[pymethods]
30impl ModelSettings {
31    #[new]
32    pub fn new(settings: &Bound<'_, PyAny>) -> Result<Self, TypeError> {
33        potatohead_macro::try_extract_py_object!(
34            settings,
35            OpenAIChatSettings => ModelSettings::OpenAIChat,
36            GeminiSettings => ModelSettings::GoogleChat,
37            AnthropicSettings => ModelSettings::AnthropicChat,
38        );
39
40        // If none matched, return error
41        Err(TypeError::InvalidModelSettings)
42    }
43
44    #[getter]
45    pub fn settings<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
46        match self {
47            ModelSettings::OpenAIChat(settings) => {
48                Ok(Py::new(py, settings.clone())?.into_bound_py_any(py)?)
49            }
50            ModelSettings::GoogleChat(settings) => {
51                Ok(Py::new(py, settings.clone())?.into_bound_py_any(py)?)
52            }
53            ModelSettings::AnthropicChat(settings) => {
54                Ok(Py::new(py, settings.clone())?.into_bound_py_any(py)?)
55            }
56        }
57    }
58
59    pub fn model_dump_json(&self) -> String {
60        serde_json::to_string(self).unwrap()
61    }
62
63    pub fn model_dump<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, TypeError> {
64        match self {
65            ModelSettings::OpenAIChat(settings) => Ok(settings.model_dump(py)?),
66            ModelSettings::GoogleChat(settings) => Ok(settings.model_dump(py)?),
67            ModelSettings::AnthropicChat(settings) => Ok(settings.model_dump(py)?),
68        }
69    }
70
71    pub fn settings_type(&self) -> SettingsType {
72        SettingsType::ModelSettings
73    }
74
75    pub fn __str__(&self) -> String {
76        PyHelperFuncs::__str__(self)
77    }
78}
79
80impl ModelSettings {
81    pub fn validate_provider(&self, provider: &Provider) -> Result<(), TypeError> {
82        match provider {
83            Provider::OpenAI => match self {
84                ModelSettings::OpenAIChat(_) => Ok(()),
85                _ => Err(TypeError::InvalidModelSettings),
86            },
87            Provider::Gemini => match self {
88                ModelSettings::GoogleChat(_) => Ok(()),
89                _ => Err(TypeError::InvalidModelSettings),
90            },
91            Provider::Vertex => match self {
92                ModelSettings::GoogleChat(_) => Ok(()),
93                _ => Err(TypeError::InvalidModelSettings),
94            },
95            Provider::Google => match self {
96                ModelSettings::GoogleChat(_) => Ok(()),
97                _ => Err(TypeError::InvalidModelSettings),
98            },
99            Provider::Anthropic => match self {
100                ModelSettings::AnthropicChat(_) => Ok(()),
101                _ => Err(TypeError::InvalidModelSettings),
102            },
103            Provider::GoogleAdk => match self {
104                ModelSettings::GoogleChat(_) => Ok(()),
105                _ => Err(TypeError::InvalidModelSettings),
106            },
107            Provider::Undefined => match self {
108                ModelSettings::OpenAIChat(_) => Ok(()),
109                ModelSettings::GoogleChat(_) => Ok(()),
110                ModelSettings::AnthropicChat(_) => Ok(()),
111            },
112        }
113    }
114
115    pub fn provider_default_settings(provider: &Provider) -> Self {
116        match provider {
117            Provider::OpenAI => ModelSettings::OpenAIChat(OpenAIChatSettings::default()),
118            Provider::Gemini | Provider::Google | Provider::Vertex | Provider::GoogleAdk => {
119                ModelSettings::GoogleChat(GeminiSettings::default())
120            }
121            _ => ModelSettings::OpenAIChat(OpenAIChatSettings::default()), // Fallback to OpenAI settings
122        }
123    }
124
125    pub fn get_openai_settings(&self) -> Option<OpenAIChatSettings> {
126        match self {
127            ModelSettings::OpenAIChat(settings) => {
128                let mut cloned_settings = settings.clone();
129                // set extra body to None
130                cloned_settings.extra_body = None;
131                Some(cloned_settings)
132            }
133            _ => None,
134        }
135    }
136
137    pub fn get_gemini_settings(&self) -> Option<GeminiSettings> {
138        match self {
139            ModelSettings::GoogleChat(settings) => {
140                let mut cloned_settings = settings.clone();
141                // set extra body to None
142                cloned_settings.extra_body = None;
143                Some(cloned_settings)
144            }
145            _ => None,
146        }
147    }
148
149    pub fn get_anthropic_settings(&self) -> AnthropicSettings {
150        match self {
151            ModelSettings::AnthropicChat(settings) => {
152                let mut cloned_settings = settings.clone();
153                // set extra body to None
154                cloned_settings.extra_body = None;
155                cloned_settings
156            }
157            _ => AnthropicSettings::default(),
158        }
159    }
160
161    pub fn extra_body(&self) -> Option<&Value> {
162        match self {
163            ModelSettings::OpenAIChat(settings) => settings.extra_body.as_ref(),
164            ModelSettings::GoogleChat(settings) => settings.extra_body.as_ref(),
165            ModelSettings::AnthropicChat(settings) => settings.extra_body.as_ref(),
166        }
167    }
168
169    pub fn provider(&self) -> Provider {
170        match self {
171            ModelSettings::OpenAIChat(_) => Provider::OpenAI,
172            ModelSettings::GoogleChat(_) => Provider::Gemini,
173            ModelSettings::AnthropicChat(_) => Provider::Anthropic,
174        }
175    }
176}
177
178#[cfg(test)]
179mod tests {
180    use super::*;
181
182    #[test]
183    fn test_validate_provider_google_adk_accepts_google_chat() {
184        let settings = ModelSettings::GoogleChat(GeminiSettings::default());
185        assert!(settings.validate_provider(&Provider::GoogleAdk).is_ok());
186    }
187
188    #[test]
189    fn test_validate_provider_google_adk_rejects_openai_chat() {
190        let settings = ModelSettings::OpenAIChat(OpenAIChatSettings::default());
191        assert!(settings.validate_provider(&Provider::GoogleAdk).is_err());
192    }
193
194    #[test]
195    fn test_provider_default_settings_google_adk_returns_google_chat() {
196        let settings = ModelSettings::provider_default_settings(&Provider::GoogleAdk);
197        assert!(matches!(settings, ModelSettings::GoogleChat(_)));
198    }
199
200    #[test]
201    fn test_provider_default_settings_google_adk_passes_own_validation() {
202        // Ensures the default settings returned for GoogleAdk are compatible with its
203        // validate_provider check — previously this would return Err(InvalidModelSettings).
204        let settings = ModelSettings::provider_default_settings(&Provider::GoogleAdk);
205        assert!(settings.validate_provider(&Provider::GoogleAdk).is_ok());
206    }
207}