use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub enum LlmProvider {
OpenAI,
Anthropic,
Custom { endpoint: String },
Mock,
}
impl LlmProvider {
pub fn default_endpoint(&self) -> Option<&str> {
match self {
LlmProvider::OpenAI => Some("https://api.openai.com/v1/chat/completions"),
LlmProvider::Anthropic => Some("https://api.anthropic.com/v1/messages"),
LlmProvider::Custom { endpoint } => Some(endpoint.as_str()),
LlmProvider::Mock => None,
}
}
pub fn is_mock(&self) -> bool {
matches!(self, LlmProvider::Mock)
}
pub fn name(&self) -> &str {
match self {
LlmProvider::OpenAI => "openai",
LlmProvider::Anthropic => "anthropic",
LlmProvider::Custom { .. } => "custom",
LlmProvider::Mock => "mock",
}
}
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
pub struct LlmProposerConfig {
provider: LlmProvider,
model: String,
prompt_template: String,
max_tokens: usize,
temperature: f32,
seed: Option<u64>,
version: String,
mock_responses: Option<Vec<String>>,
}
impl LlmProposerConfig {
pub fn new(
provider: LlmProvider,
model: impl Into<String>,
prompt_template: impl Into<String>,
max_tokens: usize,
temperature: f32,
version: impl Into<String>,
) -> Self {
LlmProposerConfig {
provider,
model: model.into(),
prompt_template: prompt_template.into(),
max_tokens,
temperature,
seed: None,
version: version.into(),
mock_responses: None,
}
}
pub fn provider(&self) -> &LlmProvider {
&self.provider
}
pub fn model(&self) -> &str {
&self.model
}
pub fn prompt_template(&self) -> &str {
&self.prompt_template
}
pub fn max_tokens(&self) -> usize {
self.max_tokens
}
pub fn temperature(&self) -> f32 {
self.temperature
}
pub fn seed(&self) -> Option<u64> {
self.seed
}
pub fn version(&self) -> &str {
&self.version
}
pub fn mock_responses(&self) -> Option<&Vec<String>> {
self.mock_responses.as_ref()
}
pub fn openai(
model: impl Into<String>,
prompt_template: impl Into<String>,
max_tokens: usize,
temperature: f32,
version: impl Into<String>,
) -> Self {
Self::new(
LlmProvider::OpenAI,
model,
prompt_template,
max_tokens,
temperature,
version,
)
}
pub fn anthropic(
model: impl Into<String>,
prompt_template: impl Into<String>,
max_tokens: usize,
temperature: f32,
version: impl Into<String>,
) -> Self {
Self::new(
LlmProvider::Anthropic,
model,
prompt_template,
max_tokens,
temperature,
version,
)
}
pub fn custom(
endpoint: impl Into<String>,
model: impl Into<String>,
prompt_template: impl Into<String>,
max_tokens: usize,
temperature: f32,
version: impl Into<String>,
) -> Self {
Self::new(
LlmProvider::Custom {
endpoint: endpoint.into(),
},
model,
prompt_template,
max_tokens,
temperature,
version,
)
}
pub fn mock(
model: impl Into<String>,
responses: Vec<String>,
version: impl Into<String>,
) -> Self {
LlmProposerConfig {
provider: LlmProvider::Mock,
model: model.into(),
prompt_template: "mock".to_string(), max_tokens: 0, temperature: 0.0, seed: None,
version: version.into(),
mock_responses: Some(responses),
}
}
pub fn with_seed(mut self, seed: u64) -> Self {
self.seed = Some(seed);
self
}
pub fn is_mock(&self) -> bool {
self.provider.is_mock()
}
pub fn endpoint(&self) -> Option<&str> {
self.provider.default_endpoint()
}
pub fn mock_response_count(&self) -> Option<usize> {
self.mock_responses.as_ref().map(|r| r.len())
}
pub fn validate(&self) -> Result<(), String> {
if self.model.is_empty() {
return Err("model name cannot be empty".to_string());
}
if !self.is_mock() {
if self.prompt_template.is_empty() {
return Err("prompt_template cannot be empty".to_string());
}
if !self.prompt_template.contains("{hypothesis}") {
return Err(
"prompt_template should contain {hypothesis} placeholder for LlmProposer"
.to_string(),
);
}
if !self.prompt_template.contains("{evidence}") {
return Err(
"prompt_template should contain {evidence} placeholder for LlmProposer"
.to_string(),
);
}
}
match &self.provider {
LlmProvider::OpenAI => {
if self.temperature < 0.0 || self.temperature > 2.0 {
return Err(format!(
"OpenAI temperature {} out of range [0.0, 2.0]",
self.temperature
));
}
}
LlmProvider::Anthropic => {
if self.temperature < 0.0 || self.temperature > 1.0 {
return Err(format!(
"Anthropic temperature {} out of range [0.0, 1.0]",
self.temperature
));
}
}
LlmProvider::Custom { .. } => {
if self.temperature < 0.0 || self.temperature > 2.0 {
return Err(format!(
"temperature {} out of range [0.0, 2.0]",
self.temperature
));
}
}
LlmProvider::Mock => {
}
}
if self.max_tokens == 0 && !self.is_mock() {
return Err("max_tokens must be positive (except in mock mode)".to_string());
}
if self.is_mock()
&& (self.mock_responses.is_none()
|| self.mock_responses.as_ref().is_none_or(|r| r.is_empty()))
{
return Err("mock mode requires at least one response in mock_responses".to_string());
}
if let LlmProvider::Custom { endpoint } = &self.provider {
if endpoint.is_empty() {
return Err("custom endpoint cannot be empty".to_string());
}
if !endpoint.starts_with("http://") && !endpoint.starts_with("https://") {
return Err(
"custom endpoint must start with http:// or https:// (got: {})"
.to_string()
.replace("{}", endpoint),
);
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_openai_config_creation() {
let config = LlmProposerConfig::openai("gpt-4", "You are helpful", 256, 0.7, "0.1.0");
assert_eq!(config.model, "gpt-4");
assert_eq!(config.max_tokens, 256);
assert!(!config.is_mock());
assert_eq!(
config.endpoint(),
Some("https://api.openai.com/v1/chat/completions")
);
}
#[test]
fn test_mock_config_creation() {
let responses = vec!["resp1".to_string(), "resp2".to_string()];
let config = LlmProposerConfig::mock("mock-model", responses, "0.1.0");
assert!(config.is_mock());
assert_eq!(config.mock_response_count(), Some(2));
assert_eq!(config.endpoint(), None);
}
#[test]
fn test_custom_endpoint_config() {
let config = LlmProposerConfig::custom(
"http://localhost:8000/v1/completions",
"local-model",
"prompt",
100,
0.5,
"0.1.0",
);
assert_eq!(
config.endpoint(),
Some("http://localhost:8000/v1/completions")
);
}
#[test]
fn test_config_validation() {
let config = LlmProposerConfig::openai(
"gpt-4",
"Given hypothesis: {hypothesis}, evidence: {evidence}",
256,
0.7,
"0.1.0",
);
assert!(config.validate().is_ok());
let config = LlmProposerConfig::openai("gpt-4", "evidence: {evidence}", 256, 0.7, "0.1.0");
assert!(config.validate().is_err());
let config =
LlmProposerConfig::openai("gpt-4", "hypothesis: {hypothesis}", 256, 0.7, "0.1.0");
assert!(config.validate().is_err());
let config = LlmProposerConfig::anthropic(
"claude-opus",
"Given hypothesis: {hypothesis}, evidence: {evidence}",
256,
1.5,
"0.1.0",
);
assert!(config.validate().is_err());
let config = LlmProposerConfig::mock("mock", vec![], "0.1.0");
assert!(config.validate().is_err());
let config = LlmProposerConfig::mock("mock", vec!["resp1".to_string()], "0.1.0");
assert!(config.validate().is_ok());
}
#[test]
fn test_seed_builder() {
let config = LlmProposerConfig::openai("gpt-4", "prompt", 256, 0.7, "0.1.0").with_seed(42);
assert_eq!(config.seed, Some(42));
}
#[test]
fn test_provider_names() {
assert_eq!(LlmProvider::OpenAI.name(), "openai");
assert_eq!(LlmProvider::Anthropic.name(), "anthropic");
assert_eq!(
LlmProvider::Custom {
endpoint: "http://localhost:8000".to_string()
}
.name(),
"custom"
);
assert_eq!(LlmProvider::Mock.name(), "mock");
}
}