Skip to main content

dynamo_protocols/types/
completion.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3//
4// Re-exports upstream async-openai completion types and defines
5// inference-serving extensions.
6
7use 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
18// Re-export response type from upstream (identical)
19pub use async_openai::types::completions::CreateCompletionResponse;
20
21/// Custom deserializer for the echo parameter that only accepts booleans.
22/// Rejects integers and strings with clear error messages.
23fn 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/// Completion request with inference-serving extensions.
89///
90/// Extends upstream `CreateCompletionRequest` with:
91/// - `prompt_embeds`: base64-encoded PyTorch tensor for pre-computed embeddings
92/// - `echo`: strict bool validation (rejects integers/strings)
93/// - `stream_options`: uses our extended `ChatCompletionStreamOptions` (with `continuous_usage_stats`)
94#[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    /// Base64-encoded PyTorch tensor containing pre-computed embeddings.
107    /// At least one of prompt or prompt_embeds is required.
108    #[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    /// Echo back the prompt in addition to the completion.
127    /// Strict bool validation -- rejects integers and strings.
128    #[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
147/// Parsed server side events stream until an \[DONE\] is received from server.
148pub 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}