rig_core/providers/azure/
extension.rs1use serde::Serialize;
14use serde_json::Value;
15
16use crate::completion::provider_options::reply_field;
17use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
18use crate::message::Api;
19use crate::providers::openai::extension::{
20 ChatOptions, CompletionTokensDetails, PromptTokensDetails,
21};
22
23#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
25pub struct AzureExt;
26
27impl ProviderExtension for AzureExt {
28 const PROVIDER: &'static str = super::PROVIDER_NAME;
29 type Options = AzureOptions;
30 type Extras = AzureExtras;
31}
32
33#[non_exhaustive]
38#[derive(Clone, Debug, Default, PartialEq, Serialize)]
39pub struct AzureOptions {
40 #[serde(rename = "openai.chat")]
42 pub chat: AzureChat,
43}
44
45#[non_exhaustive]
47#[derive(Clone, Debug, Default, PartialEq, Serialize)]
48pub struct AzureChat {
49 #[serde(flatten)]
51 pub openai: ChatOptions,
52 #[serde(skip_serializing_if = "Vec::is_empty")]
56 pub data_sources: Vec<Value>,
57}
58
59impl AzureOptions {
60 pub fn new() -> Self {
62 Self::default()
63 }
64
65 pub fn chat(mut self, chat: ChatOptions) -> Self {
67 self.chat.openai = chat;
68 self
69 }
70
71 pub fn data_source(mut self, source: Value) -> Self {
73 self.chat.data_sources.push(source);
74 self
75 }
76
77 fn with_openai(mut self, set: impl FnOnce(ChatOptions) -> ChatOptions) -> Self {
79 self.chat.openai = set(std::mem::take(&mut self.chat.openai));
80 self
81 }
82
83 pub fn logprobs(self, logprobs: bool) -> Self {
86 self.with_openai(|chat| chat.logprobs(logprobs))
87 }
88
89 pub fn top_logprobs(self, count: u8) -> Self {
92 self.with_openai(|chat| chat.top_logprobs(count))
93 }
94
95 pub fn frequency_penalty(self, penalty: f64) -> Self {
97 self.with_openai(|chat| chat.frequency_penalty(penalty))
98 }
99
100 pub fn presence_penalty(self, penalty: f64) -> Self {
102 self.with_openai(|chat| chat.presence_penalty(penalty))
103 }
104
105 pub fn logit_bias(self, token: u32, bias: i32) -> Self {
107 self.with_openai(|chat| chat.logit_bias(token, bias))
108 }
109
110 pub fn prediction(self, content: impl Into<String>) -> Self {
112 self.with_openai(|chat| chat.prediction(content))
113 }
114}
115
116impl ExtensionOptions for AzureOptions {
117 type Ext = AzureExt;
118}
119
120#[non_exhaustive]
123#[derive(Clone, Debug, Default, PartialEq)]
124pub struct AzureExtras {
125 pub service_tier: Option<String>,
127 pub system_fingerprint: Option<String>,
129 pub prompt_tokens_details: Option<PromptTokensDetails>,
131 pub completion_tokens_details: Option<CompletionTokensDetails>,
133 pub prompt_filter_results: Option<Vec<Value>>,
135 pub content_filter_results: Option<Value>,
137}
138
139impl ReplyExtras for AzureExtras {
140 fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
141 Ok(Self {
142 service_tier: reply_field(raw, "/service_tier")?,
143 system_fingerprint: reply_field(raw, "/system_fingerprint")?,
144 prompt_tokens_details: reply_field(raw, "/usage/prompt_tokens_details")?,
145 completion_tokens_details: reply_field(raw, "/usage/completion_tokens_details")?,
146 prompt_filter_results: reply_field(raw, "/prompt_filter_results")?,
147 content_filter_results: reply_field(raw, "/choices/0/content_filter_results")?,
148 })
149 }
150}
151
152#[cfg(test)]
153mod tests;