use serde::Serialize;
use serde_json::Value;
use crate::completion::provider_options::reply_field;
use crate::completion::{ExtensionOptions, ProviderExtension, ReplyExtras};
use crate::message::Api;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct GroqExt;
impl ProviderExtension for GroqExt {
const PROVIDER: &'static str = super::PROVIDER_NAME;
type Options = GroqOptions;
type Extras = GroqExtras;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct GroqOptions {
#[serde(rename = "*")]
pub shared: GroqShared,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct GroqShared {
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_format: Option<ReasoningFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub include_reasoning: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub search_settings: Option<SearchSettings>,
#[serde(skip_serializing_if = "Option::is_none")]
pub citation_options: Option<CitationOptions>,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum ReasoningFormat {
Hidden,
Raw,
Parsed,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum CitationOptions {
Enabled,
Disabled,
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq, Serialize)]
pub struct SearchSettings {
#[serde(skip_serializing_if = "Vec::is_empty")]
pub exclude_domains: Vec<String>,
#[serde(skip_serializing_if = "Vec::is_empty")]
pub include_domains: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub country: Option<String>,
}
impl SearchSettings {
pub fn new() -> Self {
Self::default()
}
pub fn exclude_domains(mut self, domains: impl IntoIterator<Item = impl Into<String>>) -> Self {
self.exclude_domains = domains.into_iter().map(Into::into).collect();
self
}
pub fn include_domains(mut self, domains: impl IntoIterator<Item = impl Into<String>>) -> Self {
self.include_domains = domains.into_iter().map(Into::into).collect();
self
}
pub fn country(mut self, country: impl Into<String>) -> Self {
self.country = Some(country.into());
self
}
}
impl GroqOptions {
pub fn new() -> Self {
Self::default()
}
pub fn reasoning_format(mut self, format: ReasoningFormat) -> Self {
self.shared.reasoning_format = Some(format);
self.shared.include_reasoning = None;
self
}
pub fn include_reasoning(mut self, include: bool) -> Self {
self.shared.include_reasoning = Some(include);
self.shared.reasoning_format = None;
self
}
pub fn search_settings(mut self, settings: SearchSettings) -> Self {
self.shared.search_settings = Some(settings);
self
}
pub fn citation_options(mut self, citations: CitationOptions) -> Self {
self.shared.citation_options = Some(citations);
self
}
}
impl ExtensionOptions for GroqOptions {
type Ext = GroqExt;
}
#[non_exhaustive]
#[derive(Clone, Debug, Default, PartialEq)]
pub struct GroqExtras {
pub x_groq: Option<Value>,
pub queue_time: Option<f64>,
pub prompt_time: Option<f64>,
pub completion_time: Option<f64>,
pub total_time: Option<f64>,
pub usage_breakdown: Option<Value>,
pub service_tier: Option<String>,
pub executed_tools: Option<Vec<Value>>,
pub system_fingerprint: Option<String>,
}
impl ReplyExtras for GroqExtras {
fn from_reply(_api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
Ok(Self {
x_groq: reply_field(raw, "/x_groq")?,
queue_time: reply_field(raw, "/usage/queue_time")?,
prompt_time: reply_field(raw, "/usage/prompt_time")?,
completion_time: reply_field(raw, "/usage/completion_time")?,
total_time: reply_field(raw, "/usage/total_time")?,
usage_breakdown: reply_field(raw, "/usage_breakdown")?,
service_tier: reply_field(raw, "/service_tier")?,
executed_tools: reply_field(raw, "/choices/0/message/executed_tools")?,
system_fingerprint: reply_field(raw, "/system_fingerprint")?,
})
}
}
#[cfg(test)]
mod tests;