Skip to main content

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}