Skip to main content

rig_core/providers/azure/
extension.rs

1//! Azure OpenAI's typed request options and reply extras, on Chat
2//! Completions: OpenAI's Chat fields and Azure's own data sources.
3//!
4//! ```
5//! use rig_core::completion::CompletionRequest;
6//! use rig_core::providers::azure::extension::{AzureOptions};
7//!
8//! let options = AzureOptions::new().logprobs(true);
9//! let request = CompletionRequest::new("hi").provider_option(options);
10//! # let _ = request;
11//! ```
12
13use 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/// Azure OpenAI's extension marker.
24#[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/// Azure OpenAI's request options, sent on Chat Completions.
34///
35/// The Chat field setters here write the same field as
36/// [`chat`](Self::chat), which replaces every such field set before it.
37#[non_exhaustive]
38#[derive(Clone, Debug, Default, PartialEq, Serialize)]
39pub struct AzureOptions {
40    /// Sent on Chat Completions only.
41    #[serde(rename = "openai.chat")]
42    pub chat: AzureChat,
43}
44
45/// The Chat Completions fields Azure takes.
46#[non_exhaustive]
47#[derive(Clone, Debug, Default, PartialEq, Serialize)]
48pub struct AzureChat {
49    /// OpenAI's Chat fields.
50    #[serde(flatten)]
51    pub openai: ChatOptions,
52    /// The data sources the model answers from
53    /// (<https://learn.microsoft.com/en-us/azure/ai-foundry/openai/references/on-your-data>),
54    /// each sent as given.
55    #[serde(skip_serializing_if = "Vec::is_empty")]
56    pub data_sources: Vec<Value>,
57}
58
59impl AzureOptions {
60    /// No option set.
61    pub fn new() -> Self {
62        Self::default()
63    }
64
65    /// Send OpenAI's Chat fields `chat`.
66    pub fn chat(mut self, chat: ChatOptions) -> Self {
67        self.chat.openai = chat;
68        self
69    }
70
71    /// Answer from `source` too, such as an Azure AI Search index.
72    pub fn data_source(mut self, source: Value) -> Self {
73        self.chat.data_sources.push(source);
74        self
75    }
76
77    /// Apply `set` to OpenAI's Chat fields.
78    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    /// Whether the reply carries log probabilities, as
84    /// [`ChatOptions::logprobs`].
85    pub fn logprobs(self, logprobs: bool) -> Self {
86        self.with_openai(|chat| chat.logprobs(logprobs))
87    }
88
89    /// Report the `count` most likely tokens per position, as
90    /// [`ChatOptions::top_logprobs`].
91    pub fn top_logprobs(self, count: u8) -> Self {
92        self.with_openai(|chat| chat.top_logprobs(count))
93    }
94
95    /// Set the frequency penalty, as [`ChatOptions::frequency_penalty`].
96    pub fn frequency_penalty(self, penalty: f64) -> Self {
97        self.with_openai(|chat| chat.frequency_penalty(penalty))
98    }
99
100    /// Set the presence penalty, as [`ChatOptions::presence_penalty`].
101    pub fn presence_penalty(self, penalty: f64) -> Self {
102        self.with_openai(|chat| chat.presence_penalty(penalty))
103    }
104
105    /// Bias token id `token` by `bias`, as [`ChatOptions::logit_bias`].
106    pub fn logit_bias(self, token: u32, bias: i32) -> Self {
107        self.with_openai(|chat| chat.logit_bias(token, bias))
108    }
109
110    /// Predict the output as `content`, as [`ChatOptions::prediction`].
111    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/// Azure OpenAI's Chat reply fields. Each is `None` when the reply lacks
121/// it.
122#[non_exhaustive]
123#[derive(Clone, Debug, Default, PartialEq)]
124pub struct AzureExtras {
125    /// The service tier that served the request.
126    pub service_tier: Option<String>,
127    /// The backend configuration's fingerprint.
128    pub system_fingerprint: Option<String>,
129    /// Prompt token details.
130    pub prompt_tokens_details: Option<PromptTokensDetails>,
131    /// Completion token details.
132    pub completion_tokens_details: Option<CompletionTokensDetails>,
133    /// The content filter's verdicts on the prompt, one per prompt.
134    pub prompt_filter_results: Option<Vec<Value>>,
135    /// The content filter's verdict on the first choice.
136    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;