use std::collections::HashMap;
use std::error::Error;
use serde::{Deserialize, Serialize};
use crate::core::RequestOptions;
use crate::OpenAIObject;
use crate::resource::APIResource;
#[derive(Debug, Clone)]
pub struct Completions {
pub client: Option<APIResource>,
}
impl Completions {
pub fn new() -> Self {
Completions {
client: None,
}
}
pub async fn create(&self, body: CompletionCreateParams) -> Result<Completion, Box<dyn Error>> {
let stream = body.stream.unwrap_or(false);
self.client.as_ref().unwrap().borrow().post(
"/completions",
Some( RequestOptions {
body: Some(body),
stream: Some(stream),
..Default::default()
})
).await
}
}
#[derive(Default, Debug, Deserialize, Serialize)]
pub struct Completion {
pub id: String,
pub choices: Vec<CompletionChoice>,
pub created: u64,
pub model: String,
pub object: OpenAIObject,
pub system_fingerprint: Option<String>,
pub usage: Option<CompletionUsage>,
}
#[derive(Default, Debug, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum FinishReason {
#[default]
Stop,
Length,
ContentFilter,
}
#[derive(Default, Debug, Deserialize, Serialize)]
pub struct CompletionChoice {
pub finish_reason: FinishReason,
pub index: u32,
pub logprobs: Option<Logprobs>,
pub text: String,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct Logprobs {
pub text_offset: Option<Vec<u32>>,
pub token_logprobs: Option<Vec<f32>>,
pub tokens: Option<Vec<String>>,
pub top_logprobs: Option<Vec<HashMap<String, f32>>>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct CompletionUsage {
pub completion_tokens: u32,
pub prompt_tokens: u32,
pub total_tokens: u32,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(untagged)]
pub enum CompletionCreate {
NonStreaming(CompletionCreateParams),
Streaming(CompletionCreateParams),
}
impl Default for CompletionCreate {
fn default() -> Self {
CompletionCreate::NonStreaming(CompletionCreateParams::default())
}
}
#[derive(Default, Debug, Clone, Deserialize, Serialize)]
pub struct CompletionCreateParams {
pub model: String,
pub prompt: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub best_of: Option<u32>,
pub echo: Option<bool>,
pub frequency_penalty: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logit_bias: Option<HashMap<String, f32>>,
pub logprobs: Option<u32>,
pub max_tokens: Option<u32>,
pub n: Option<u32>,
pub presence_penalty: Option<f32>,
pub seed: Option<u32>,
pub stop: Option<serde_json::Value>,
pub stream: Option<bool>,
pub stream_options: Option<StreamOptions>,
pub suffix: Option<String>,
pub temperature: Option<f32>,
pub top_p: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct StreamOptions {
}
impl From<CompletionCreateParams> for CompletionCreate {
fn from(params: CompletionCreateParams) -> Self {
CompletionCreate::NonStreaming(params)
}
}