use std::fmt;
use std::str::FromStr;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ToolProfile {
Core,
#[default]
Full,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum AlwaysLoad {
#[default]
Core,
None,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "snake_case")]
pub enum ResponseFormat {
#[default]
Json,
Text,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default)]
pub struct SessionOptions {
pub tools: ToolProfile,
pub max_tokens: Option<usize>,
pub project: Option<String>,
pub always_load: AlwaysLoad,
pub format: ResponseFormat,
}
macro_rules! str_enum {
($ty:ty { $($variant:ident => $text:literal),+ }) => {
impl fmt::Display for $ty {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self { $(Self::$variant => $text),+ })
}
}
impl FromStr for $ty {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
$($text => Ok(Self::$variant),)+
other => Err(format!(
"unknown value '{other}' (expected one of: {})",
[$($text),+].join(", ")
)),
}
}
}
};
}
str_enum!(ToolProfile { Core => "core", Full => "full" });
str_enum!(AlwaysLoad { Core => "core", None => "none" });
str_enum!(ResponseFormat { Json => "json", Text => "text" });
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn options_round_trip_and_default_when_absent() {
let opts = SessionOptions {
tools: ToolProfile::Core,
max_tokens: Some(2000),
project: Some("demo".into()),
always_load: AlwaysLoad::None,
format: ResponseFormat::Text,
};
let json = serde_json::to_string(&opts).unwrap();
assert_eq!(serde_json::from_str::<SessionOptions>(&json).unwrap(), opts);
assert_eq!(
serde_json::from_str::<SessionOptions>("{}").unwrap(),
SessionOptions::default()
);
}
#[test]
fn profiles_parse_from_flag_text() {
assert_eq!("core".parse::<ToolProfile>().unwrap(), ToolProfile::Core);
assert_eq!("none".parse::<AlwaysLoad>().unwrap(), AlwaysLoad::None);
assert!("everything".parse::<ToolProfile>().is_err());
}
}