openai_interface/completions/request.rs
1use std::collections::HashMap;
2
3use serde::{Deserialize, Serialize};
4use url::Url;
5
6use crate::{
7 errors::OapiError,
8 rest::post::{Post, PostNoStream, PostStream},
9};
10
11#[derive(Debug, Serialize, Deserialize, Default, Clone)]
12pub struct CompletionRequest {
13 /// ID of the model to use. Note that not all models are supported for completion.
14 pub model: String,
15 /// The prompt(s) to generate completions for, encoded as a string, array of
16 /// strings, array of tokens, or array of token arrays.
17 /// Note that <|endoftext|> is the document separator that the model sees during
18 /// training, so if a prompt is not specified the model will generate as if from the
19 /// beginning of a new document.
20 pub prompt: Prompt,
21 /// Generates `best_of` completions server-side and returns the "best" (the one with
22 /// the highest log probability per token). Results cannot be streamed.
23 ///
24 /// When used with `n`, `best_of` controls the number of candidate completions and
25 /// `n` specifies how many to return – `best_of` must be greater than `n`.
26 ///
27 /// **Note:** Because this parameter generates many completions, it can quickly
28 /// consume your token quota. Use carefully and ensure that you have reasonable
29 /// settings for `max_tokens` and `stop`.
30 #[serde(skip_serializing_if = "Option::is_none")]
31 pub best_of: Option<usize>,
32 /// Echo back the prompt in addition to the completion
33 #[serde(skip_serializing_if = "Option::is_none")]
34 pub echo: Option<bool>,
35 /// Number between -2.0 and 2.0. Positive values penalize new tokens based on their
36 /// existing frequency in the text so far, decreasing the model's likelihood to
37 /// repeat the same line verbatim.
38 ///
39 /// [more info about frequency/presence penalties](https://platform.openai.com/docs/guides/text-generation)
40 #[serde(skip_serializing_if = "Option::is_none")]
41 pub frequency_penalty: Option<f32>,
42 /// Modify the likelihood of specified tokens appearing in the completion.
43 ///
44 /// Accepts a JSON object that maps tokens (specified by their token ID in the GPT
45 /// tokenizer) to an associated bias value from -100 to 100. You can use this
46 /// [tokenizer tool](/tokenizer?view=bpe) to convert text to token IDs.
47 /// Mathematically, the bias is added to the logits generated by the model prior to
48 /// sampling. The exact effect will vary per model, but values between -1 and 1
49 /// should decrease or increase likelihood of selection; values like -100 or 100
50 /// should result in a ban or exclusive selection of the relevant token.
51 ///
52 /// As an example, you can pass `{"50256": -100}` to prevent the <|end-of-stream|> token
53 /// from being generated.
54 #[serde(skip_serializing_if = "Option::is_none")]
55 pub logit_bias: Option<HashMap<String, isize>>,
56 /// Include the log probabilities on the `logprobs` most likely output tokens, as
57 /// well the chosen tokens. For example, if `logprobs` is 5, the API will return a
58 /// list of the 5 most likely tokens. The API will always return the `logprob` of
59 /// the sampled token, so there may be up to `logprobs+1` elements in the response.
60 ///
61 /// The maximum value for `logprobs` is 5.
62 #[serde(skip_serializing_if = "Option::is_none")]
63 pub logprobs: Option<usize>,
64 /// The maximum number of [tokens](/tokenizer) that can be generated in the
65 /// completion.
66 ///
67 /// The token count of your prompt plus `max_tokens` cannot exceed the model's
68 /// context length.
69 /// [Example Python code](https://cookbook.openai.com/examples/how_to_count_tokens_with_tiktoken)
70 /// for counting tokens.
71 #[serde(skip_serializing_if = "Option::is_none")]
72 pub max_tokens: Option<usize>,
73 /// How many completions to generate for each prompt.
74 ///
75 /// **Note:** Because this parameter generates many completions, it can quickly
76 /// consume your token quota. Use carefully and ensure that you have reasonable
77 /// settings for `max_tokens` and `stop`.
78 #[serde(skip_serializing_if = "Option::is_none")]
79 pub n: Option<usize>,
80 /// Number between -2.0 and 2.0. Positive values penalize new tokens based on
81 /// whether they appear in the text so far, increasing the model's likelihood to
82 /// talk about new topics.
83 ///
84 /// [See more information about frequency and presence penalties.](https://platform.openai.com/docs/guides/text-generation)
85 #[serde(skip_serializing_if = "Option::is_none")]
86 pub presence_penalty: Option<f32>,
87 /// If specified, our system will make a best effort to sample deterministically,
88 /// such that repeated requests with the same `seed` and parameters should return
89 /// the same result.
90 ///
91 /// Determinism is not guaranteed, and you should refer to the `system_fingerprint`
92 /// response parameter to monitor changes in the backend.
93 #[serde(skip_serializing_if = "Option::is_none")]
94 pub seed: Option<usize>,
95 /// Up to 4 sequences where the API will stop generating further tokens. The
96 /// returned text will not contain the stop sequence.
97 ///
98 /// Note: Not supported with latest reasoning models `o3` and `o4-mini`.
99 #[serde(skip_serializing_if = "Option::is_none")]
100 pub stop: Option<StopKeywords>,
101 /// Whether to stream back partial progress. If set to `true`, tokens will be sent as
102 /// data-only
103 /// [server-sent events](https://developer.mozilla.org/en-US/docs/Web/API/Server-sent_events/Using_server-sent_events#Event_stream_format)
104 /// as they become available, with the stream terminated by a `data: [DONE]`
105 /// message.
106 /// [Example Python code](https://cookbook.openai.com/examples/how_to_stream_completions).
107 #[serde(skip_serializing_if = "Option::is_none")]
108 pub stream: Option<bool>,
109 /// Options for streaming response. Only set this when you set `stream: true`.
110 #[serde(skip_serializing_if = "Option::is_none")]
111 pub stream_options: Option<StreamOptions>,
112 /// The suffix that comes after a completion of inserted text.
113 ///
114 /// This parameter is only supported for `gpt-3.5-turbo-instruct`.
115 /// vLLM rejects it outright on `/v1/completions`.
116 #[serde(skip_serializing_if = "Option::is_none")]
117 pub suffix: Option<String>,
118 /// What sampling temperature to use, between 0 and 2. Higher values like 0.8 will
119 /// make the output more random, while lower values like 0.2 will make it more
120 /// focused and deterministic.
121 ///
122 /// It is generally recommended to alter this or `top_p` but not both.
123 #[serde(skip_serializing_if = "Option::is_none")]
124 pub temperature: Option<f32>,
125 /// An alternative to sampling with temperature, called nucleus sampling,
126 /// where the model considers the results of the tokens with `top_p`
127 /// probability mass. So 0.1 means only the tokens comprising the top 10%
128 /// probability mass are considered.
129 ///
130 /// It is generally recommended to alter this or `temperature` but not both.
131 #[serde(skip_serializing_if = "Option::is_none")]
132 pub top_p: Option<f32>,
133 /// A unique identifier representing your end-user, which can help OpenAI to monitor
134 /// and detect abuse.
135 /// [Learn more from OpenAI](https://platform.openai.com/docs/guides/safety-best-practices#end-user-ids).
136 #[serde(skip_serializing_if = "Option::is_none")]
137 pub user: Option<String>,
138 /// vLLM: extra sampling parameters (`min_p`, `repetition_penalty`,
139 /// `stop_token_ids`, `prompt_logprobs`, ...) that OpenAI's API does not
140 /// define. Flattened into the top level of the request body.
141 ///
142 /// The chat-only vLLM parameters (`chat_template_kwargs`,
143 /// `structured_outputs`, ...) are deliberately absent: this endpoint
144 /// does not accept them.
145 #[cfg(feature = "vllm")]
146 #[serde(flatten, default, skip_serializing_if = "Option::is_none")]
147 pub vllm_sampling: Option<crate::vllm::SamplingParams>,
148 /// Add additional JSON properties to the request.
149 ///
150 /// Flattened into the JSON body, matching the official SDK's `extra_body`
151 /// semantics: the entries are sent as top-level request parameters.
152 #[serde(flatten, default, skip_serializing_if = "Option::is_none")]
153 pub extra_body_map: Option<serde_json::Map<String, serde_json::Value>>,
154}
155
156#[derive(Debug, Serialize, Deserialize, Clone)]
157#[serde(untagged)]
158pub enum Prompt {
159 /// String
160 PromptString(String),
161 /// Array of strings
162 PromptStringArray(Vec<String>),
163 /// Array of tokens
164 TokensArray(Vec<usize>),
165 /// Array of arrays of tokens
166 TokenArraysArray(Vec<Vec<usize>>),
167}
168impl Default for Prompt {
169 fn default() -> Self {
170 Self::PromptString("".to_string())
171 }
172}
173
174#[derive(Debug, Serialize, Deserialize, Clone)]
175pub struct StreamOptions {
176 /// If set, an additional chunk will be streamed before the `data: [DONE]` message.
177 ///
178 /// The `usage` field on this chunk shows the token usage statistics for the entire
179 /// request, and the `choices` field will always be an empty array.
180 ///
181 /// All other chunks will also include a `usage` field, but with a null value.
182 /// **NOTE:** If the stream is interrupted, you may not receive the final usage
183 /// chunk which contains the total token usage for the request.
184 pub include_usage: bool,
185}
186
187#[derive(Debug, Serialize, Deserialize, Clone)]
188#[serde(untagged)]
189pub enum StopKeywords {
190 Word(String),
191 Words(Vec<String>),
192}
193
194impl CompletionRequest {
195 /// Whether this request asks for a streamed response. Defaults to
196 /// `false` when [`CompletionRequest::stream`] is `None`.
197 pub fn is_streaming(&self) -> bool {
198 self.stream.unwrap_or(false)
199 }
200}
201
202impl Post for CompletionRequest {
203 fn is_streaming(&self) -> bool {
204 CompletionRequest::is_streaming(self)
205 }
206
207 /// Builds the URL for the request.
208 ///
209 /// `base_url` should be like <https://api.openai.com/v1>
210 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
211 let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
212 url.path_segments_mut()
213 .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
214 .push("completions");
215
216 Ok(url.to_string())
217 }
218}
219
220impl PostNoStream for CompletionRequest {
221 type Response = super::response::Completion;
222}
223
224impl PostStream for CompletionRequest {
225 type Response = super::response::Completion;
226}
227
228#[cfg(test)]
229mod tests {
230 use futures_util::StreamExt;
231
232 use super::*;
233
234 const QWEN_MODEL: &str = "qwen-coder-turbo";
235 const QWEN_URL: &str = "https://dashscope.aliyuncs.com/compatible-mode/v1";
236
237 fn qwen_api_key() -> Option<String> {
238 std::env::var("QWEN_API_KEY")
239 .ok()
240 .map(|key| key.trim().to_string())
241 .filter(|key| !key.is_empty())
242 }
243
244 #[tokio::test]
245 async fn test_qwen_completions_no_stream() -> Result<(), anyhow::Error> {
246 let Some(api_key) = qwen_api_key() else {
247 println!("Skipping: set QWEN_API_KEY to run this test");
248 return Ok(());
249 };
250
251 let request_body = CompletionRequest {
252 model: QWEN_MODEL.to_string(),
253 prompt: Prompt::PromptString(
254 r#"
255 package main
256
257 import (
258 "fmt"
259 "strings"
260 "net/http"
261 "io/ioutil"
262 )
263
264 func main() {
265
266 url := "https://api.deepseek.com/chat/completions"
267 method := "POST"
268
269 payload := strings.NewReader(`{
270 "messages": [
271 {
272 "content": "You are a helpful assistant",
273 "role": "system"
274 },
275 {
276 "content": "Hi",
277 "role": "user"
278 }
279 ],
280 "model": "deepseek-chat",
281 "frequency_penalty": 0,
282 "max_tokens": 4096,
283 "presence_penalty": 0,
284 "response_format": {
285 "type": "text"
286 },
287 "stop": null,
288 "stream": false,
289 "stream_options": null,
290 "temperature": 1,
291 "top_p": 1,
292 "tools": null,
293 "tool_choice": "none",
294 "logprobs": false,
295 "top_logprobs": null
296 }`)
297
298 client := &http.Client {
299 }
300 req, err := http.NewRequest(method, url, payload)
301
302 if err != nil {
303 fmt.Println(err)
304 return
305 }
306 req.Header.Add("Content-Type", "application/json")
307 req.Header.Add("Accept", "application/json")
308 req.Header.Add("Authorization", "Bearer <TOKEN>")
309
310 res, err := client.Do(req)
311 if err != nil {
312 fmt.Println(err)
313 return
314 }
315 defer res.Body.Close()
316"#
317 .to_string(),
318 ),
319 suffix: Some(
320 r#"
321 if err != nil {
322 fmt.Println(err)
323 return
324 }
325 fmt.Println(string(body))
326}
327"#
328 .to_string(),
329 ),
330 stream: Some(false),
331 ..Default::default()
332 };
333
334 let result = request_body
335 .get_response_string(
336 &crate::rest::default_client(),
337 QWEN_URL,
338 &crate::rest::RequestOptions::bearer(&api_key),
339 )
340 .await?;
341 println!("{}", result);
342
343 Ok(())
344 }
345
346 #[tokio::test]
347 async fn test_qwen_completions_stream() -> Result<(), anyhow::Error> {
348 let Some(api_key) = qwen_api_key() else {
349 println!("Skipping: set QWEN_API_KEY to run this test");
350 return Ok(());
351 };
352
353 let request_body = CompletionRequest {
354 model: QWEN_MODEL.to_string(),
355 prompt: Prompt::PromptString(
356 r#"
357 package main
358
359 import (
360 "fmt"
361 "strings"
362 "net/http"
363 "io/ioutil"
364 )
365
366 func main() {
367
368 url := "https://api.deepseek.com/chat/completions"
369 method := "POST"
370
371 payload := strings.NewReader(`{
372 "messages": [
373 {
374 "content": "You are a helpful assistant",
375 "role": "system"
376 },
377 {
378 "content": "Hi",
379 "role": "user"
380 }
381 ],
382 "model": "deepseek-chat",
383 "frequency_penalty": 0,
384 "max_tokens": 4096,
385 "presence_penalty": 0,
386 "response_format": {
387 "type": "text"
388 },
389 "stop": null,
390 "stream": true,
391 "stream_options": null,
392 "temperature": 1,
393 "top_p": 1,
394 "tools": null,
395 "tool_choice": "none",
396 "logprobs": false,
397 "top_logprobs": null
398 }`)
399
400 client := &http.Client {
401 }
402 req, err := http.NewRequest(method, url, payload)
403
404 if err != nil {
405 fmt.Println(err)
406 return
407 }
408 req.Header.Add("Content-Type", "application/json")
409 req.Header.Add("Accept", "application/json")
410 req.Header.Add("Authorization", "Bearer <TOKEN>")
411
412 res, err := client.Do(req)
413 if err != nil {
414 fmt.Println(err)
415 return
416 }
417 defer res.Body.Close()
418 "#
419 .to_string(),
420 ),
421 suffix: Some(
422 r#"
423 if err != nil {
424 fmt.Println(err)
425 return
426 }
427 fmt.Println(string(body))
428 }
429 "#
430 .to_string(),
431 ),
432 stream: Some(true),
433 ..Default::default()
434 };
435
436 let mut stream = request_body
437 .get_stream_response_string(
438 &crate::rest::default_client(),
439 QWEN_URL,
440 &crate::rest::RequestOptions::bearer(&api_key),
441 )
442 .await?;
443
444 while let Some(chunk) = stream.next().await {
445 match chunk {
446 Ok(data) => {
447 println!("Received chunk: {:?}", data);
448 }
449 Err(e) => {
450 eprintln!("Error receiving chunk: {:?}", e);
451 break;
452 }
453 }
454 }
455
456 Ok(())
457 }
458}