use serde::Serialize;
use serde_json::Value;
use crate::completion::provider_options::reply_field;
use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
use crate::message::Api;
use crate::providers::openai::extension::{
ChatOptions, CompletionTokensDetails, PromptTokensDetails,
};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct AzureExt;
impl ProviderExtension for AzureExt {
const PROVIDER: &'static str = super::PROVIDER_NAME;
type Options = AzureOptions;
type Extras = AzureExtras;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct AzureOptions {
#[serde(rename = "openai.chat")]
pub chat: AzureChat,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct AzureChat {
#[serde(flatten)]
pub openai: ChatOptions,
#[serde(skip_serializing_if = "Vec::is_empty")]
pub data_sources: Vec<Value>,
}
impl AzureOptions {
pub fn new() -> Self {
Self::default()
}
pub fn chat(mut self, chat: ChatOptions) -> Self {
self.chat.openai = chat;
self
}
pub fn data_source(mut self, source: Value) -> Self {
self.chat.data_sources.push(source);
self
}
fn with_openai(mut self, set: impl FnOnce(ChatOptions) -> ChatOptions) -> Self {
self.chat.openai = set(std::mem::take(&mut self.chat.openai));
self
}
pub fn logprobs(self, logprobs: bool) -> Self {
self.with_openai(|chat| chat.logprobs(logprobs))
}
pub fn top_logprobs(self, count: u8) -> Self {
self.with_openai(|chat| chat.top_logprobs(count))
}
pub fn frequency_penalty(self, penalty: f64) -> Self {
self.with_openai(|chat| chat.frequency_penalty(penalty))
}
pub fn presence_penalty(self, penalty: f64) -> Self {
self.with_openai(|chat| chat.presence_penalty(penalty))
}
pub fn logit_bias(self, token: u32, bias: i32) -> Self {
self.with_openai(|chat| chat.logit_bias(token, bias))
}
pub fn prediction(self, content: impl Into<String>) -> Self {
self.with_openai(|chat| chat.prediction(content))
}
}
impl ExtensionOptions for AzureOptions {
type Ext = AzureExt;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq)]
pub struct AzureExtras {
pub service_tier: Option<String>,
pub system_fingerprint: Option<String>,
pub prompt_tokens_details: Option<PromptTokensDetails>,
pub completion_tokens_details: Option<CompletionTokensDetails>,
pub prompt_filter_results: Option<Vec<Value>>,
pub content_filter_results: Option<Value>,
}
impl ReplyExtras for AzureExtras {
fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
Ok(Self {
service_tier: reply_field(raw, "/service_tier")?,
system_fingerprint: reply_field(raw, "/system_fingerprint")?,
prompt_tokens_details: reply_field(raw, "/usage/prompt_tokens_details")?,
completion_tokens_details: reply_field(raw, "/usage/completion_tokens_details")?,
prompt_filter_results: reply_field(raw, "/prompt_filter_results")?,
content_filter_results: reply_field(raw, "/choices/0/content_filter_results")?,
})
}
}
#[cfg(test)]
mod tests;