1use std::collections::HashMap;
8use std::pin::Pin;
9
10use derive_builder::Builder;
11use futures::Stream;
12use serde::{Deserialize, Serialize};
13
14use crate::error::OpenAIError;
15
16use super::{ChatCompletionStreamOptions, Prompt, Stop};
17
18pub use async_openai::types::completions::CreateCompletionResponse;
20
21fn deserialize_echo_bool<'de, D>(deserializer: D) -> Result<Option<bool>, D::Error>
24where
25 D: serde::Deserializer<'de>,
26{
27 struct StrictBoolVisitor;
28
29 impl<'de> serde::de::Visitor<'de> for StrictBoolVisitor {
30 type Value = Option<bool>;
31
32 fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
33 formatter.write_str("echo parameter to be a boolean (true or false) or null")
34 }
35
36 fn visit_some<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
37 where
38 D: serde::Deserializer<'de>,
39 {
40 deserializer.deserialize_any(BoolOnlyVisitor)
41 }
42
43 fn visit_none<E>(self) -> Result<Self::Value, E>
44 where
45 E: serde::de::Error,
46 {
47 Ok(None)
48 }
49
50 fn visit_unit<E>(self) -> Result<Self::Value, E>
51 where
52 E: serde::de::Error,
53 {
54 Ok(None)
55 }
56 }
57
58 struct BoolOnlyVisitor;
59
60 impl<'de> serde::de::Visitor<'de> for BoolOnlyVisitor {
61 type Value = Option<bool>;
62
63 fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
64 formatter.write_str("echo parameter to be a boolean (true or false) or null")
65 }
66
67 fn visit_bool<E>(self, value: bool) -> Result<Self::Value, E>
68 where
69 E: serde::de::Error,
70 {
71 Ok(Some(value))
72 }
73
74 fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
75 where
76 E: serde::de::Error,
77 {
78 Err(E::invalid_type(
79 serde::de::Unexpected::Str(value),
80 &"echo parameter to be a boolean (true or false) or null",
81 ))
82 }
83 }
84
85 deserializer.deserialize_option(StrictBoolVisitor)
86}
87
88#[derive(Clone, Serialize, Deserialize, Default, Debug, Builder, PartialEq)]
95#[builder(name = "CreateCompletionRequestArgs")]
96#[builder(pattern = "mutable")]
97#[builder(setter(into, strip_option), default)]
98#[builder(derive(Debug))]
99#[builder(build_fn(error = "OpenAIError"))]
100#[cfg_attr(feature = "protocol-schema", derive(utoipa::ToSchema))]
101#[cfg_attr(feature = "protocol-schema", schema(as = dynamo_protocols::completion::CreateCompletionRequest))]
102pub struct CreateCompletionRequest {
103 pub model: String,
104 #[cfg_attr(feature = "protocol-schema", schema(value_type = crate::schema::Prompt))]
105 pub prompt: Prompt,
106 #[serde(skip_serializing_if = "Option::is_none")]
109 pub prompt_embeds: Option<String>,
110 #[serde(skip_serializing_if = "Option::is_none")]
111 pub suffix: Option<String>,
112 #[serde(skip_serializing_if = "Option::is_none")]
113 pub max_tokens: Option<u32>,
114 #[serde(skip_serializing_if = "Option::is_none")]
115 pub temperature: Option<f32>,
116 #[serde(skip_serializing_if = "Option::is_none")]
117 pub top_p: Option<f32>,
118 #[serde(skip_serializing_if = "Option::is_none")]
119 pub n: Option<u8>,
120 #[serde(skip_serializing_if = "Option::is_none")]
121 pub stream: Option<bool>,
122 #[serde(skip_serializing_if = "Option::is_none")]
123 pub stream_options: Option<ChatCompletionStreamOptions>,
124 #[serde(skip_serializing_if = "Option::is_none")]
125 pub logprobs: Option<u8>,
126 #[serde(skip_serializing_if = "Option::is_none")]
129 #[serde(default, deserialize_with = "deserialize_echo_bool")]
130 pub echo: Option<bool>,
131 #[serde(skip_serializing_if = "Option::is_none")]
132 pub stop: Option<Stop>,
133 #[serde(skip_serializing_if = "Option::is_none")]
134 pub presence_penalty: Option<f32>,
135 #[serde(skip_serializing_if = "Option::is_none")]
136 pub frequency_penalty: Option<f32>,
137 #[serde(skip_serializing_if = "Option::is_none")]
138 pub best_of: Option<u8>,
139 #[serde(skip_serializing_if = "Option::is_none")]
140 pub logit_bias: Option<HashMap<String, serde_json::Value>>,
141 #[serde(skip_serializing_if = "Option::is_none")]
142 pub user: Option<String>,
143 #[serde(skip_serializing_if = "Option::is_none")]
144 pub seed: Option<i64>,
145}
146
147pub type CompletionResponseStream =
149 Pin<Box<dyn Stream<Item = Result<CreateCompletionResponse, OpenAIError>> + Send>>;
150
151#[cfg(test)]
152mod tests {
153 use super::*;
154
155 #[test]
156 fn echo_rejects_integer() {
157 let json = r#"{"model": "test_model", "prompt": "test", "echo": 1}"#;
158 let result: Result<CreateCompletionRequest, _> = serde_json::from_str(json);
159 assert!(result.is_err());
160 let err_msg = result.unwrap_err().to_string();
161 assert!(err_msg.contains("invalid type"));
162 assert!(err_msg.contains("integer"));
163 assert!(err_msg.contains("echo parameter"));
164 }
165
166 #[test]
167 fn echo_rejects_string() {
168 let json = r#"{"model": "test_model", "prompt": "test", "echo": "null"}"#;
169 let result: Result<CreateCompletionRequest, _> = serde_json::from_str(json);
170 assert!(result.is_err());
171 let err_msg = result.unwrap_err().to_string();
172 assert!(err_msg.contains("invalid type"));
173 assert!(err_msg.contains("string"));
174 assert!(err_msg.contains("echo parameter"));
175 }
176
177 #[test]
178 fn completion_choice_serializes_openai_shape() {
179 use crate::types::{Choice, CompletionFinishReason};
180
181 let choice = Choice {
182 text: "hello".to_string(),
183 index: 0,
184 logprobs: None,
185 finish_reason: Some(CompletionFinishReason::Stop),
186 };
187
188 let value = serde_json::to_value(choice).expect("serialize choice");
189
190 assert_eq!(value["finish_reason"], "stop");
191 assert_eq!(value["text"], "hello");
192 }
193
194 #[test]
195 fn stop_accepts_token_id_array() {
196 let json = r#"{"model": "test_model", "prompt": [1, 2, 3], "stop": [32, 34]}"#;
197 let request: CreateCompletionRequest = serde_json::from_str(json).unwrap();
198
199 assert_eq!(request.stop, Some(Stop::TokenIdArray(vec![32, 34])));
200 }
201
202 #[test]
203 fn stop_accepts_string_and_string_array() {
204 let one_stop = r#"{"model": "test_model", "prompt": "hello", "stop": " The"}"#;
205 let request: CreateCompletionRequest = serde_json::from_str(one_stop).unwrap();
206
207 assert_eq!(request.stop, Some(Stop::String(" The".to_string())));
208
209 let many_stops = r#"{"model": "test_model", "prompt": "hello", "stop": ["A", "B"]}"#;
210 let request: CreateCompletionRequest = serde_json::from_str(many_stops).unwrap();
211
212 assert_eq!(
213 request.stop,
214 Some(Stop::StringArray(vec!["A".to_string(), "B".to_string()]))
215 );
216 }
217
218 #[test]
219 fn stop_token_id_display_string_remains_string_stop() {
220 let json = r#"{"model": "test_model", "prompt": [1, 2, 3], "stop": "token_id:576"}"#;
221 let request: CreateCompletionRequest = serde_json::from_str(json).unwrap();
222
223 assert_eq!(request.stop, Some(Stop::String("token_id:576".to_string())));
224
225 let json = r#"{"model": "test_model", "prompt": [1, 2, 3], "stop": ["token_id:576"]}"#;
226 let request: CreateCompletionRequest = serde_json::from_str(json).unwrap();
227
228 assert_eq!(
229 request.stop,
230 Some(Stop::StringArray(vec!["token_id:576".to_string()]))
231 );
232 }
233
234 #[test]
235 fn builder_accepts_upstream_stop_configuration() {
236 let upstream_stop = async_openai::types::chat::StopConfiguration::String("END".to_string());
237
238 let request = CreateCompletionRequestArgs::default()
239 .model("test_model")
240 .prompt(Prompt::String("hello".to_string()))
241 .stop(upstream_stop)
242 .build()
243 .unwrap();
244
245 assert_eq!(request.stop, Some(Stop::String("END".to_string())));
246 }
247
248 #[test]
249 fn stop_rejects_single_token_id() {
250 let json = r#"{"model": "test_model", "prompt": [1, 2, 3], "stop": 576}"#;
251 let result: Result<CreateCompletionRequest, _> = serde_json::from_str(json);
252
253 assert!(result.is_err());
254 }
255}