potato_type/prompt/
settings.rs1use 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 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()), }
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 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 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 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 let settings = ModelSettings::provider_default_settings(&Provider::GoogleAdk);
205 assert!(settings.validate_provider(&Provider::GoogleAdk).is_ok());
206 }
207}