Skip to main content

potato_type/prompt/
builder.rs

1use crate::anthropic::v1::request::AnthropicMessageRequestV1;
2use crate::anthropic::MessageParam;
3use crate::google::v1::generate::request::GeminiGenerateContentRequestV1;
4use crate::google::GeminiContent;
5use crate::openai::v1::chat::request::OpenAIChatCompletionRequestV1;
6use crate::openai::ChatMessage;
7use crate::prompt::types::MessageNum;
8use crate::prompt::ModelSettings;
9use crate::tools::{AgentToolDefinition, ToolCall};
10use crate::traits::RequestAdapter;
11use crate::{Provider, TypeError};
12use potatohead_macro::dispatch_trait_method;
13use pyo3::types::PyList;
14use pyo3::types::PyListMethods;
15use pyo3::Python;
16use pyo3::{Bound, PyAny};
17use serde::{Deserialize, Serialize};
18use serde_json::Value;
19
20/// Trait for converting a list of messages into a provider-specific request
21pub trait RequestBuilder {
22    type Request: Serialize;
23
24    /// Build a request from messages and settings
25    fn build_request(
26        messages: Vec<MessageNum>,
27        system_instructions: Vec<MessageNum>,
28        model: String,
29        settings: ModelSettings,
30        response_format: Option<Value>,
31    ) -> Result<Self::Request, TypeError>;
32}
33
34/// Type marker for request routing
35#[derive(Debug, Clone, Copy)]
36pub enum RequestType {
37    OpenAIChatV1,
38    AnthropicMessageV1,
39    GeminiContentV1,
40}
41
42impl MessageNum {
43    /// Determine the request type from this message variant
44    pub fn request_type(&self) -> RequestType {
45        match self {
46            MessageNum::OpenAIMessageV1(_) => RequestType::OpenAIChatV1,
47            MessageNum::AnthropicMessageV1(_) => RequestType::AnthropicMessageV1,
48            MessageNum::GeminiContentV1(_) => RequestType::GeminiContentV1,
49            MessageNum::AnthropicSystemMessageV1(_) => RequestType::AnthropicMessageV1,
50            // RawV1 is a fallback — should not appear in initial messages list
51            MessageNum::RawV1(_) => RequestType::OpenAIChatV1,
52        }
53    }
54}
55
56/// Unified enum for provider-specific requests
57/// This serves as a central access point for accessing request attributes from within a Prompt
58#[derive(Debug, Serialize, Deserialize, Clone, PartialEq)]
59#[serde(untagged)]
60pub enum ProviderRequest {
61    OpenAIV1(OpenAIChatCompletionRequestV1),
62    AnthropicV1(AnthropicMessageRequestV1),
63    GeminiV1(GeminiGenerateContentRequestV1),
64}
65
66impl ProviderRequest {
67    pub fn insert_message(&mut self, message: MessageNum, idx: Option<usize>) {
68        self.messages_mut().insert(idx.unwrap_or(0), message);
69    }
70
71    pub fn push_message(&mut self, message: MessageNum) {
72        self.messages_mut().push(message);
73    }
74
75    pub fn messages(&self) -> &[MessageNum] {
76        dispatch_trait_method!(self, RequestAdapter, messages())
77    }
78
79    pub fn system_instructions(&self) -> Vec<&MessageNum> {
80        dispatch_trait_method!(self, RequestAdapter, system_instructions())
81    }
82
83    pub fn messages_mut(&mut self) -> &mut Vec<MessageNum> {
84        dispatch_trait_method!(mut self, RequestAdapter, messages_mut())
85    }
86
87    pub fn add_tools(&mut self, tools: Vec<AgentToolDefinition>) -> Result<(), TypeError> {
88        dispatch_trait_method!(mut self, RequestAdapter, add_tools(tools))
89    }
90
91    pub fn prepend_system_instructions(
92        &mut self,
93        instructions: Vec<MessageNum>,
94    ) -> Result<(), TypeError> {
95        dispatch_trait_method!(mut self, RequestAdapter, preprend_system_instructions(instructions))
96    }
97
98    /// Returns the messages as a Python list
99    pub(crate) fn get_py_messages<'py>(
100        &self,
101        py: Python<'py>,
102    ) -> Result<Bound<'py, PyList>, TypeError> {
103        let py_messages = PyList::empty(py);
104
105        for msg in self.messages() {
106            if msg.is_user_message() {
107                py_messages.append(msg.to_bound_py_object(py)?)?;
108            }
109        }
110
111        Ok(py_messages)
112    }
113
114    pub(crate) fn get_all_py_messages<'py>(
115        &self,
116        py: Python<'py>,
117    ) -> Result<Bound<'py, PyList>, TypeError> {
118        let py_messages = PyList::empty(py);
119
120        for msg in self.messages() {
121            py_messages.append(msg.to_bound_py_object(py)?)?;
122        }
123
124        Ok(py_messages)
125    }
126
127    /// Returns the last message in the request as a Python object
128    pub(crate) fn get_py_message<'py>(
129        &self,
130        py: Python<'py>,
131    ) -> Result<Bound<'py, PyAny>, TypeError> {
132        let last = self
133            .messages()
134            .iter()
135            .rev()
136            .find(|msg| msg.is_user_message())
137            .ok_or_else(|| {
138                TypeError::Error("No messages in request to convert to Python object".to_string())
139            })?;
140
141        last.to_bound_py_object(py)
142    }
143
144    pub(crate) fn get_openai_message(&self) -> Result<ChatMessage, TypeError> {
145        let last = self
146            .messages()
147            .iter()
148            .rev()
149            .find(|msg| msg.is_user_message())
150            .ok_or_else(|| {
151                TypeError::Error("No messages in request to convert to Python object".to_string())
152            })?;
153
154        match last {
155            MessageNum::OpenAIMessageV1(msg) => Ok(msg.clone()),
156            _ => Err(TypeError::Error(
157                "Last message is not an OpenAI ChatMessage".to_string(),
158            )),
159        }
160    }
161
162    /// Returns the messages as Anthropic MessageParam Python objects
163    pub(crate) fn get_gemini_message(&self) -> Result<GeminiContent, TypeError> {
164        let last = self
165            .messages()
166            .iter()
167            .rev()
168            .find(|msg| msg.is_user_message())
169            .ok_or_else(|| {
170                TypeError::Error("No messages in request to convert to Python object".to_string())
171            })?;
172
173        match last {
174            MessageNum::GeminiContentV1(msg) => Ok(msg.clone()),
175            _ => Err(TypeError::Error(
176                "Last message is not a GeminiContent".to_string(),
177            )),
178        }
179    }
180
181    /// Returns the messages as Anthropic MessageParam Python objects
182    pub(crate) fn get_anthropic_message(&self) -> Result<MessageParam, TypeError> {
183        let last = self
184            .messages()
185            .iter()
186            .rev()
187            .find(|msg| msg.is_user_message())
188            .ok_or_else(|| {
189                TypeError::Error("No messages in request to convert to Python object".to_string())
190            })?;
191
192        Ok(match last {
193            MessageNum::AnthropicMessageV1(msg) => msg.clone(),
194            _ => {
195                return Err(TypeError::Error(
196                    "Last message is not an Anthropic MessageParam".to_string(),
197                ))
198            }
199        })
200    }
201
202    pub(crate) fn get_py_system_instructions<'py>(
203        &self,
204        py: Python<'py>,
205    ) -> Result<Bound<'py, PyList>, TypeError> {
206        dispatch_trait_method!(self, RequestAdapter, get_py_system_instructions(py))
207    }
208
209    pub(crate) fn model_settings<'py>(
210        &self,
211        py: Python<'py>,
212    ) -> Result<Bound<'py, PyAny>, TypeError> {
213        dispatch_trait_method!(self, RequestAdapter, model_settings(py))
214    }
215
216    pub fn response_json_schema(&self) -> Option<&Value> {
217        dispatch_trait_method!(self, RequestAdapter, response_json_schema())
218    }
219
220    pub fn has_structured_output(&self) -> bool {
221        self.response_json_schema().is_some()
222    }
223
224    /// Retrieve the JSON request body for the specified provider
225    /// This method will first attempt to match the provider type,
226    /// returning an error if there is a mismatch.
227    pub fn to_request(&self, provider: &Provider) -> Result<Value, TypeError> {
228        let is_matched = dispatch_trait_method!(self, RequestAdapter, match_provider(provider));
229
230        if !is_matched {
231            return Err(TypeError::Error(
232                "ProviderRequest does not match the specified provider".to_string(),
233            ));
234        }
235        dispatch_trait_method!(self, RequestAdapter, to_request_body())
236    }
237
238    /// Serialize to JSON for API requests
239    pub fn to_json(&self) -> Result<Value, TypeError> {
240        Ok(serde_json::to_value(self)?)
241    }
242
243    pub fn set_response_json_schema(&mut self, response_json_schema: Option<Value>) {
244        dispatch_trait_method!(mut self, RequestAdapter, set_response_json_schema(response_json_schema))
245    }
246
247    /// Injects a tool result message into the conversation in the appropriate provider format.
248    ///
249    /// - OpenAI: Pushes a RawV1 message with role="tool" and tool_call_id
250    /// - Anthropic: Pushes a user MessageParam with a tool_result content block
251    /// - Gemini: Pushes a user GeminiContent with a functionResponse part
252    pub fn add_tool_result(
253        &mut self,
254        call: &ToolCall,
255        result: &serde_json::Value,
256    ) -> Result<(), TypeError> {
257        match self {
258            ProviderRequest::OpenAIV1(_) => {
259                let call_id = call
260                    .call_id
261                    .clone()
262                    .unwrap_or_else(|| "unknown".to_string());
263                let content = result.to_string();
264                let raw = serde_json::json!({
265                    "role": "tool",
266                    "tool_call_id": call_id,
267                    "content": content,
268                });
269                self.messages_mut().push(MessageNum::RawV1(raw));
270                Ok(())
271            }
272            ProviderRequest::AnthropicV1(_) => {
273                use crate::anthropic::v1::request::{
274                    ContentBlock, ContentBlockParam, TextBlockParam, ToolResultBlockParam,
275                    ToolResultContentEnum,
276                };
277                let tool_use_id = call
278                    .call_id
279                    .clone()
280                    .unwrap_or_else(|| call.tool_name.clone());
281                let result_text = result.to_string();
282                let text_block = TextBlockParam::new_rs(result_text, None, None);
283                let tool_result = ToolResultBlockParam {
284                    tool_use_id,
285                    is_error: None,
286                    cache_control: None,
287                    r#type: "tool_result".to_string(),
288                    content: Some(ToolResultContentEnum::Text(vec![text_block])),
289                };
290                let content_block = ContentBlockParam {
291                    inner: ContentBlock::ToolResult(tool_result),
292                };
293                let message = MessageParam {
294                    role: "user".to_string(),
295                    content: vec![content_block],
296                };
297                self.messages_mut()
298                    .push(MessageNum::AnthropicMessageV1(message));
299                Ok(())
300            }
301            ProviderRequest::GeminiV1(_) => {
302                use crate::google::v1::generate::request::{
303                    DataNum, FunctionResponse, GeminiContent, Part,
304                };
305                let name = call.tool_name.clone();
306                let mut response_map = std::collections::HashMap::new();
307                response_map.insert("output".to_string(), result.clone());
308                let func_response = FunctionResponse {
309                    name,
310                    response: response_map,
311                };
312                let part = Part {
313                    data: DataNum::FunctionResponse(func_response),
314                    ..Default::default()
315                };
316                let content = GeminiContent {
317                    role: "user".to_string(),
318                    parts: vec![part],
319                };
320                self.messages_mut()
321                    .push(MessageNum::GeminiContentV1(content));
322                Ok(())
323            }
324        }
325    }
326}
327
328pub fn to_provider_request(
329    messages: Vec<MessageNum>,
330    system_instructions: Vec<MessageNum>,
331    model: String,
332    model_settings: ModelSettings,
333    response_json_schema: Option<Value>,
334) -> Result<ProviderRequest, TypeError> {
335    // Determine request type from first message
336    let request_type = messages
337        .first()
338        .ok_or_else(|| TypeError::Error("Prompt has no messages".to_string()))?
339        .request_type();
340
341    // Validate all messages are same type
342    for msg in &messages {
343        if msg.request_type() as u8 != request_type as u8 {
344            return Err(TypeError::Error(
345                "All messages must be of the same provider type".to_string(),
346            ));
347        }
348    }
349
350    // Build appropriate request based on type
351    match request_type {
352        RequestType::OpenAIChatV1 => OpenAIChatCompletionRequestV1::build_provider_enum(
353            messages,
354            system_instructions,
355            model,
356            model_settings,
357            response_json_schema,
358        ),
359        RequestType::AnthropicMessageV1 => AnthropicMessageRequestV1::build_provider_enum(
360            messages,
361            system_instructions,
362            model,
363            model_settings,
364            response_json_schema,
365        ),
366        RequestType::GeminiContentV1 => GeminiGenerateContentRequestV1::build_provider_enum(
367            messages,
368            system_instructions,
369            model,
370            model_settings,
371            response_json_schema,
372        ),
373    }
374}