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
20pub trait RequestBuilder {
22 type Request: Serialize;
23
24 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#[derive(Debug, Clone, Copy)]
36pub enum RequestType {
37 OpenAIChatV1,
38 AnthropicMessageV1,
39 GeminiContentV1,
40}
41
42impl MessageNum {
43 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 MessageNum::RawV1(_) => RequestType::OpenAIChatV1,
52 }
53 }
54}
55
56#[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 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 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 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 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 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 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 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 let request_type = messages
337 .first()
338 .ok_or_else(|| TypeError::Error("Prompt has no messages".to_string()))?
339 .request_type();
340
341 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 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}