use crate::detail::ffi::{
FOUNDRY_LOCAL_PARAM_DO_SAMPLE, FOUNDRY_LOCAL_PARAM_EARLY_STOPPING,
FOUNDRY_LOCAL_PARAM_FREQUENCY_PENALTY, FOUNDRY_LOCAL_PARAM_MAX_OUTPUT_TOKENS,
FOUNDRY_LOCAL_PARAM_PRESENCE_PENALTY, FOUNDRY_LOCAL_PARAM_SEED,
FOUNDRY_LOCAL_PARAM_TEMPERATURE, FOUNDRY_LOCAL_PARAM_TOOL_CHOICE, FOUNDRY_LOCAL_PARAM_TOP_K,
FOUNDRY_LOCAL_PARAM_TOP_P,
};
use crate::item::Item;
use crate::item_queue::ItemQueue;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub enum ToolChoice {
#[default]
Auto,
None,
Required,
}
impl ToolChoice {
fn as_param(self) -> &'static str {
match self {
ToolChoice::Auto => "auto",
ToolChoice::None => "none",
ToolChoice::Required => "required",
}
}
}
#[derive(Debug, Clone, PartialEq, Default)]
pub struct SearchOptions {
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<i32>,
pub max_output_tokens: Option<i32>,
pub frequency_penalty: Option<f32>,
pub presence_penalty: Option<f32>,
pub seed: Option<i64>,
pub early_stopping: Option<bool>,
pub do_sample: Option<bool>,
}
#[derive(Debug, Clone, PartialEq, Default)]
pub struct RequestOptions {
pub search: SearchOptions,
pub tool_choice: Option<ToolChoice>,
pub additional_options: Vec<(String, String)>,
}
impl RequestOptions {
pub(crate) fn to_pairs(&self) -> Vec<(String, String)> {
let mut pairs: Vec<(String, String)> = self.additional_options.clone();
let s = &self.search;
if let Some(v) = s.temperature {
pairs.push((FOUNDRY_LOCAL_PARAM_TEMPERATURE.to_string(), v.to_string()));
}
if let Some(v) = s.top_p {
pairs.push((FOUNDRY_LOCAL_PARAM_TOP_P.to_string(), v.to_string()));
}
if let Some(v) = s.top_k {
pairs.push((FOUNDRY_LOCAL_PARAM_TOP_K.to_string(), v.to_string()));
}
if let Some(v) = s.max_output_tokens {
pairs.push((
FOUNDRY_LOCAL_PARAM_MAX_OUTPUT_TOKENS.to_string(),
v.to_string(),
));
}
if let Some(v) = s.frequency_penalty {
pairs.push((
FOUNDRY_LOCAL_PARAM_FREQUENCY_PENALTY.to_string(),
v.to_string(),
));
}
if let Some(v) = s.presence_penalty {
pairs.push((
FOUNDRY_LOCAL_PARAM_PRESENCE_PENALTY.to_string(),
v.to_string(),
));
}
if let Some(v) = s.seed {
pairs.push((FOUNDRY_LOCAL_PARAM_SEED.to_string(), v.to_string()));
}
if let Some(v) = s.early_stopping {
pairs.push((
FOUNDRY_LOCAL_PARAM_EARLY_STOPPING.to_string(),
bool_str(v).to_string(),
));
}
if let Some(v) = s.do_sample {
pairs.push((
FOUNDRY_LOCAL_PARAM_DO_SAMPLE.to_string(),
bool_str(v).to_string(),
));
}
if let Some(tc) = self.tool_choice {
pairs.push((
FOUNDRY_LOCAL_PARAM_TOOL_CHOICE.to_string(),
tc.as_param().to_string(),
));
}
pairs
}
}
fn bool_str(v: bool) -> &'static str {
if v {
"true"
} else {
"false"
}
}
#[derive(Debug, Clone, Default)]
pub struct Request {
pub items: Vec<Item>,
pub input_queue: Option<ItemQueue>,
pub options: Option<RequestOptions>,
}
impl Request {
pub fn new() -> Self {
Self::default()
}
pub fn from_items(items: impl Into<Vec<Item>>) -> Self {
Self {
items: items.into(),
input_queue: None,
options: None,
}
}
pub fn with_item(mut self, item: Item) -> Self {
self.items.push(item);
self
}
pub fn with_input_queue(mut self, queue: ItemQueue) -> Self {
self.input_queue = Some(queue);
self
}
pub fn with_options(mut self, options: RequestOptions) -> Self {
self.options = Some(options);
self
}
pub(crate) fn option_pairs(&self) -> Vec<(String, String)> {
self.options
.as_ref()
.map(RequestOptions::to_pairs)
.unwrap_or_default()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn typed_fields_override_additional_options() {
let opts = RequestOptions {
search: SearchOptions {
temperature: Some(0.5),
max_output_tokens: Some(128),
do_sample: Some(false),
..Default::default()
},
tool_choice: Some(ToolChoice::Required),
additional_options: vec![("temperature".into(), "9.9".into())],
};
let pairs = opts.to_pairs();
assert_eq!(pairs[0], ("temperature".to_string(), "9.9".to_string()));
assert!(pairs.iter().any(|(k, v)| k == "temperature" && v == "0.5"));
assert!(pairs.iter().rposition(|(k, _)| k == "temperature").unwrap() > 0);
assert!(pairs
.iter()
.any(|(k, v)| k == "max_output_tokens" && v == "128"));
assert!(pairs.iter().any(|(k, v)| k == "do_sample" && v == "false"));
assert!(pairs
.iter()
.any(|(k, v)| k == "tool_choice" && v == "required"));
}
#[test]
fn builder_assembles_request() {
let req = Request::from_items(vec![Item::text("hi")])
.with_item(Item::text("there"))
.with_options(RequestOptions::default());
assert_eq!(req.items.len(), 2);
assert!(req.options.is_some());
assert!(req.input_queue.is_none());
}
}