rig_core/providers/groq/
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;
19
20#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
22pub struct GroqExt;
23
24impl ProviderExtension for GroqExt {
25 const PROVIDER: &'static str = super::PROVIDER_NAME;
26 type Options = GroqOptions;
27 type Extras = GroqExtras;
28}
29
30#[non_exhaustive]
32#[derive(Clone, Debug, Default, PartialEq, Serialize)]
33pub struct GroqOptions {
34 #[serde(rename = "*")]
36 pub shared: GroqShared,
37}
38
39#[non_exhaustive]
41#[derive(Clone, Debug, Default, PartialEq, Serialize)]
42pub struct GroqShared {
43 #[serde(skip_serializing_if = "Option::is_none")]
45 pub reasoning_format: Option<ReasoningFormat>,
46 #[serde(skip_serializing_if = "Option::is_none")]
49 pub include_reasoning: Option<bool>,
50 #[serde(skip_serializing_if = "Option::is_none")]
52 pub search_settings: Option<SearchSettings>,
53 #[serde(skip_serializing_if = "Option::is_none")]
55 pub citation_options: Option<CitationOptions>,
56}
57
58#[non_exhaustive]
60#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
61#[serde(rename_all = "lowercase")]
62pub enum ReasoningFormat {
63 Hidden,
65 Raw,
67 Parsed,
69}
70
71#[non_exhaustive]
73#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
74#[serde(rename_all = "lowercase")]
75pub enum CitationOptions {
76 Enabled,
78 Disabled,
80}
81
82#[non_exhaustive]
84#[derive(Clone, Debug, Default, PartialEq, Serialize)]
85pub struct SearchSettings {
86 #[serde(skip_serializing_if = "Vec::is_empty")]
88 pub exclude_domains: Vec<String>,
89 #[serde(skip_serializing_if = "Vec::is_empty")]
91 pub include_domains: Vec<String>,
92 #[serde(skip_serializing_if = "Option::is_none")]
94 pub country: Option<String>,
95}
96
97impl SearchSettings {
98 pub fn new() -> Self {
100 Self::default()
101 }
102
103 pub fn exclude_domains(mut self, domains: impl IntoIterator<Item = impl Into<String>>) -> Self {
105 self.exclude_domains = domains.into_iter().map(Into::into).collect();
106 self
107 }
108
109 pub fn include_domains(mut self, domains: impl IntoIterator<Item = impl Into<String>>) -> Self {
111 self.include_domains = domains.into_iter().map(Into::into).collect();
112 self
113 }
114
115 pub fn country(mut self, country: impl Into<String>) -> Self {
117 self.country = Some(country.into());
118 self
119 }
120}
121
122impl GroqOptions {
123 pub fn new() -> Self {
125 Self::default()
126 }
127
128 pub fn reasoning_format(mut self, format: ReasoningFormat) -> Self {
130 self.shared.reasoning_format = Some(format);
131 self.shared.include_reasoning = None;
132 self
133 }
134
135 pub fn include_reasoning(mut self, include: bool) -> Self {
138 self.shared.include_reasoning = Some(include);
139 self.shared.reasoning_format = None;
140 self
141 }
142
143 pub fn search_settings(mut self, settings: SearchSettings) -> Self {
145 self.shared.search_settings = Some(settings);
146 self
147 }
148
149 pub fn citation_options(mut self, citations: CitationOptions) -> Self {
151 self.shared.citation_options = Some(citations);
152 self
153 }
154}
155
156impl ExtensionOptions for GroqOptions {
157 type Ext = GroqExt;
158}
159
160#[non_exhaustive]
162#[derive(Clone, Debug, Default, PartialEq)]
163pub struct GroqExtras {
164 pub x_groq: Option<Value>,
166 pub queue_time: Option<f64>,
168 pub prompt_time: Option<f64>,
170 pub completion_time: Option<f64>,
172 pub total_time: Option<f64>,
174 pub usage_breakdown: Option<Value>,
176 pub service_tier: Option<String>,
178 pub executed_tools: Option<Vec<Value>>,
180 pub system_fingerprint: Option<String>,
182}
183
184impl ReplyExtras for GroqExtras {
185 fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
186 Ok(Self {
187 x_groq: reply_field(raw, "/x_groq")?,
188 queue_time: reply_field(raw, "/usage/queue_time")?,
189 prompt_time: reply_field(raw, "/usage/prompt_time")?,
190 completion_time: reply_field(raw, "/usage/completion_time")?,
191 total_time: reply_field(raw, "/usage/total_time")?,
192 usage_breakdown: reply_field(raw, "/usage_breakdown")?,
193 service_tier: reply_field(raw, "/service_tier")?,
194 executed_tools: reply_field(raw, "/choices/0/message/executed_tools")?,
195 system_fingerprint: reply_field(raw, "/system_fingerprint")?,
196 })
197 }
198}
199
200#[cfg(test)]
201mod tests;