use std::env;
use std::path::{Path, PathBuf};
use open_agent::ApiProtocol;
use serde::Deserialize;
use thiserror::Error;
use toml::Value;
mod backend;
pub use backend::{BackendKind, LlmConfig, ReasoningEffort};
pub const DEFAULT_MAX_REVIEW_ROUNDS: u32 = 3;
#[derive(Debug, Deserialize)]
#[serde(default)]
pub struct Config {
pub max_review_rounds: u32,
pub llm: Vec<LlmConfig>,
}
impl Default for Config {
fn default() -> Self {
Self {
max_review_rounds: DEFAULT_MAX_REVIEW_ROUNDS,
llm: Vec::new(),
}
}
}
impl Config {
pub fn providers(&self) -> Vec<&LlmConfig> {
self.llm.iter().filter(|p| p.enabled).collect()
}
}
#[derive(Debug, Error)]
pub enum ConfigError {
#[error("could not read {0}: {1}")]
Io(PathBuf, std::io::Error),
#[error("could not parse {0}: {1}")]
Parse(PathBuf, String),
#[error("environment variable `{0}` is not set (referenced by `{1}`)")]
EnvVarUnset(String, String),
#[error("environment variable `{0}` is not valid UTF-8 (referenced by `{1}`)")]
EnvVarNotUnicode(String, String),
#[error(
"[[llm]] #{} in file order: temperature {temperature} is outside the allowed range 0.0..=2.0",
index + 1
)]
Temperature { index: usize, temperature: f32 },
#[error("[[llm]] #{} in file order: max_concurrent must be at least 1", index + 1)]
ZeroConcurrency { index: usize },
#[error("[[llm]] #{} in file order: timeout_secs must be at least 1", index + 1)]
ZeroTimeout { index: usize },
#[error("[[llm]] #{} in file order: max_tokens must be at least 1 when set", index + 1)]
ZeroMaxTokens { index: usize },
#[error("max_review_rounds must be at least 1")]
ZeroReviewRounds,
#[error(
"[[llm]] #{} in file order: unknown protocol `{value}`; expected `openai` or `anthropic`",
index + 1
)]
UnknownProtocol { index: usize, value: String },
#[error(
"[[llm]] #{} in file order: unknown backend `{value}`; expected `http` or `codex`",
index + 1
)]
UnknownBackend { index: usize, value: String },
#[error(
"[[llm]] #{} in file order: unknown reasoning_effort `{value}`; expected `minimal`, `low`, `medium`, `high`, or `xhigh`",
index + 1
)]
UnknownReasoningEffort { index: usize, value: String },
#[error(
"[[llm]] #{} in file order: backend `{backend}` does not support `{field}`",
index + 1
)]
BackendField {
index: usize,
backend: &'static str,
field: &'static str,
},
#[error(
"[[llm]] #{} in file order: backend `{backend}` requires `{field}`",
index + 1
)]
BackendMissingField {
index: usize,
backend: &'static str,
field: &'static str,
},
#[error(
"{0} declares no `[[llm]]` provider; drep 2.x has no deterministic-only mode. \
Run `drep init` to write one."
)]
NoProviders(PathBuf),
#[error(
"every `[[llm]]` provider in {0} has `enabled = false`; drep 2.x has no \
deterministic-only mode. Re-enable one, or run `drep init` to write another."
)]
NoEnabledProviders(PathBuf),
}
pub fn default_config_path() -> PathBuf {
PathBuf::from("drep.toml")
}
pub fn parse_protocol(raw: Option<&str>) -> Option<ApiProtocol> {
match raw {
None => Some(ApiProtocol::default()),
Some(name) => ApiProtocol::from_wire(name),
}
}
pub fn load(path: &Path) -> Result<Config, ConfigError> {
let content =
std::fs::read_to_string(path).map_err(|err| ConfigError::Io(path.to_path_buf(), err))?;
let mut tree: Value = toml::from_str(&content).map_err(|err: toml::de::Error| {
ConfigError::Parse(path.to_path_buf(), err.message().to_owned())
})?;
let disabled = disabled_provider_indices(&tree);
expand_env_except(&mut tree, path, &disabled)?;
let explicit_fields = backend::explicit_fields(&tree);
let config: Config = tree.try_into().map_err(|err: toml::de::Error| {
ConfigError::Parse(path.to_path_buf(), err.message().to_owned())
})?;
validate(&config, path, &explicit_fields)?;
Ok(config)
}
fn validate(
config: &Config,
path: &Path,
explicit_fields: &[backend::ExplicitFields],
) -> Result<(), ConfigError> {
if config.max_review_rounds == 0 {
return Err(ConfigError::ZeroReviewRounds);
}
if config.llm.is_empty() {
return Err(ConfigError::NoProviders(path.to_path_buf()));
}
if config.providers().is_empty() {
return Err(ConfigError::NoEnabledProviders(path.to_path_buf()));
}
for (index, llm) in config.llm.iter().enumerate().filter(|(_, l)| l.enabled) {
backend::validate(
llm,
explicit_fields.get(index).copied().unwrap_or_default(),
index,
)?;
if llm.max_concurrent == 0 {
return Err(ConfigError::ZeroConcurrency { index });
}
if llm.timeout_secs == 0 {
return Err(ConfigError::ZeroTimeout { index });
}
if llm.max_tokens == Some(0) {
return Err(ConfigError::ZeroMaxTokens { index });
}
if llm.backend != BackendKind::Http {
continue;
}
if let Some(t) = llm.temperature
&& !(0.0..=2.0).contains(&t)
{
return Err(ConfigError::Temperature {
index,
temperature: t,
});
}
if let Some(raw) = llm.protocol.as_deref()
&& parse_protocol(Some(raw)).is_none()
{
return Err(ConfigError::UnknownProtocol {
index,
value: raw.to_owned(),
});
}
}
Ok(())
}
fn disabled_provider_indices(tree: &Value) -> std::collections::BTreeSet<usize> {
let default_enabled = LlmConfig::default().enabled;
tree.get("llm")
.and_then(Value::as_array)
.map(|entries| {
entries
.iter()
.enumerate()
.filter(|(_, entry)| {
!entry
.get("enabled")
.and_then(Value::as_bool)
.unwrap_or(default_enabled)
})
.map(|(index, _)| index)
.collect()
})
.unwrap_or_default()
}
fn expand_env_except(
tree: &mut Value,
source: &Path,
skip: &std::collections::BTreeSet<usize>,
) -> Result<(), ConfigError> {
if skip.is_empty() {
return expand_env_in(tree, source);
}
let Some(table) = tree.as_table_mut() else {
return expand_env_in(tree, source);
};
for (key, value) in table.iter_mut() {
if key != "llm" {
expand_env_in(value, source)?;
continue;
}
let Some(entries) = value.as_array_mut() else {
expand_env_in(value, source)?;
continue;
};
for (index, entry) in entries.iter_mut().enumerate() {
if !skip.contains(&index) {
expand_env_in(entry, source)?;
}
}
}
Ok(())
}
fn expand_env_in(value: &mut Value, source: &Path) -> Result<(), ConfigError> {
match value {
Value::String(s) => {
*s = expand_string(s, source)?;
}
Value::Table(table) => {
for (_, inner) in table.iter_mut() {
expand_env_in(inner, source)?;
}
}
Value::Array(items) => {
for inner in items.iter_mut() {
expand_env_in(inner, source)?;
}
}
_ => {}
}
Ok(())
}
pub fn env_var_refs(s: &str) -> Vec<String> {
let mut refs = Vec::new();
let mut rest = s;
while let Some((_, after_open)) = rest.split_once("${") {
let Some((name, after_close)) = after_open.split_once('}') else {
break;
};
refs.push(name.to_owned());
rest = after_close;
}
refs
}
pub fn required_env_var_refs(value: &Value) -> Vec<String> {
let disabled = disabled_provider_indices(value);
if disabled.is_empty() {
return env_var_refs_in(value);
}
let mut seen = std::collections::BTreeSet::new();
let mut out = Vec::new();
let Some(table) = value.as_table() else {
return env_var_refs_in(value);
};
for (key, inner) in table {
if key != "llm" {
collect_env_refs(inner, &mut seen, &mut out);
continue;
}
let Some(entries) = inner.as_array() else {
collect_env_refs(inner, &mut seen, &mut out);
continue;
};
for (index, entry) in entries.iter().enumerate() {
if !disabled.contains(&index) {
collect_env_refs(entry, &mut seen, &mut out);
}
}
}
out
}
pub fn env_var_refs_in(value: &Value) -> Vec<String> {
let mut seen = std::collections::BTreeSet::new();
let mut out = Vec::new();
collect_env_refs(value, &mut seen, &mut out);
out
}
fn collect_env_refs(
value: &Value,
seen: &mut std::collections::BTreeSet<String>,
out: &mut Vec<String>,
) {
match value {
Value::String(s) => {
for name in env_var_refs(s) {
if seen.insert(name.clone()) {
out.push(name);
}
}
}
Value::Table(table) => {
for (_, inner) in table {
collect_env_refs(inner, seen, out);
}
}
Value::Array(items) => {
for inner in items {
collect_env_refs(inner, seen, out);
}
}
_ => {}
}
}
fn expand_string(s: &str, source: &Path) -> Result<String, ConfigError> {
let mut out = String::with_capacity(s.len());
let mut chars = s.chars().peekable();
while let Some(c) = chars.next() {
if c != '$' {
out.push(c);
continue;
}
if chars.peek() != Some(&'{') {
out.push(c);
continue;
}
chars.next();
let mut name = String::new();
let mut closed = false;
for next in chars.by_ref() {
if next == '}' {
closed = true;
break;
}
name.push(next);
}
if !closed {
return Err(ConfigError::Parse(
source.to_path_buf(),
format!("unterminated `${{` in `{s}`"),
));
}
if name.is_empty() {
return Err(ConfigError::Parse(
source.to_path_buf(),
format!("empty environment variable reference in `{s}`"),
));
}
let value = match env::var(&name) {
Ok(value) => value,
Err(env::VarError::NotPresent) => {
return Err(ConfigError::EnvVarUnset(name, source.display().to_string()));
}
Err(env::VarError::NotUnicode(_)) => {
return Err(ConfigError::EnvVarNotUnicode(
name,
source.display().to_string(),
));
}
};
out.push_str(&value);
}
Ok(out)
}
#[cfg(test)]
mod tests;