use std::fmt;
use std::time::Duration;
use crate::errors::Error;
pub(crate) const API_KEY_ENV: &str = "TYPESAFE_API_KEY";
pub(crate) const BASE_URL_ENV: &str = "TYPESAFE_BASE_URL";
pub(crate) const DEFAULT_MODEL_ENV: &str = "TYPESAFE_DEFAULT_MODEL";
pub(crate) const LOG_LEVEL_ENV: &str = "TYPESAFE_LOG_LEVEL";
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, PartialOrd, Ord)]
pub enum LogLevel {
Debug,
Info,
Warn,
Error,
#[default]
Off,
}
impl LogLevel {
pub fn parse(name: &str) -> Option<Self> {
match name.trim().to_ascii_lowercase().as_str() {
"debug" => Some(Self::Debug),
"info" => Some(Self::Info),
"warn" | "warning" => Some(Self::Warn),
"error" => Some(Self::Error),
"off" => Some(Self::Off),
_ => None,
}
}
pub(crate) fn resolve(explicit: Option<Self>) -> Self {
if let Some(level) = explicit {
return level;
}
env_trimmed(LOG_LEVEL_ENV)
.and_then(|value| Self::parse(&value))
.unwrap_or_default()
}
}
impl fmt::Display for LogLevel {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let name = match self {
Self::Debug => "debug",
Self::Info => "info",
Self::Warn => "warn",
Self::Error => "error",
Self::Off => "off",
};
f.write_str(name)
}
}
pub(crate) fn env_trimmed(name: &str) -> Option<String> {
std::env::var(name)
.ok()
.map(|value| value.trim().to_owned())
.filter(|value| !value.is_empty())
}
fn resolve_string(value: Option<&str>, env: &str, default: &str) -> String {
if let Some(value) = value {
return value.trim().to_owned();
}
env_trimmed(env).unwrap_or_else(|| default.to_owned())
}
pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result<String, Error> {
let key = api_key
.map(str::trim)
.map(str::to_owned)
.or_else(|| env_trimmed(API_KEY_ENV))
.unwrap_or_default();
if key.is_empty() {
return Err(Error::Config(format!(
"No API key was provided. Pass api_key or set the {API_KEY_ENV} environment variable."
)));
}
if !key.is_ascii() {
return Err(Error::Config(
"API key must contain only printable ASCII characters without whitespace.".into(),
));
}
if key.chars().any(|c| c.is_whitespace() || c.is_control()) {
return Err(Error::Config(
"API key must contain only printable ASCII characters without whitespace.".into(),
));
}
if !key.chars().all(|c| {
let c = c as u8;
(0x21..=0x7e).contains(&c)
}) {
return Err(Error::Config(
"API key must contain only printable ASCII characters without whitespace.".into(),
));
}
Ok(key)
}
pub(crate) fn resolve_base_url(base_url: Option<&str>) -> String {
resolve_string(base_url, BASE_URL_ENV, crate::DEFAULT_BASE_URL)
.trim_end_matches('/')
.to_owned()
}
pub(crate) fn resolve_default_model(default_model: Option<&str>) -> String {
resolve_string(default_model, DEFAULT_MODEL_ENV, crate::DEFAULT_MODEL)
}
pub(crate) fn validate_timeout(timeout: Duration) -> Result<Duration, Error> {
if timeout.is_zero() {
return Err(Error::Config(
"timeout must be a positive duration of time.".into(),
));
}
Ok(timeout)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_log_level() {
assert_eq!(LogLevel::parse("debug"), Some(LogLevel::Debug));
assert_eq!(LogLevel::parse("INFO"), Some(LogLevel::Info));
assert_eq!(LogLevel::parse(" warn "), Some(LogLevel::Warn));
assert_eq!(LogLevel::parse("warning"), Some(LogLevel::Warn));
assert_eq!(LogLevel::parse("error"), Some(LogLevel::Error));
assert_eq!(LogLevel::parse("off"), Some(LogLevel::Off));
assert_eq!(LogLevel::parse("nope"), None);
assert_eq!(LogLevel::parse(""), None);
}
#[test]
fn default_log_level_is_off() {
assert_eq!(LogLevel::default(), LogLevel::Off);
}
#[test]
fn api_key_validation() {
let err = resolve_api_key(None).unwrap_err();
assert!(matches!(err, Error::Config(_)));
assert!(err.to_string().contains("No API key"));
for bad in ["key with space", "key\ttab", "key\nnewline", "ünïcödé"] {
let err = resolve_api_key(Some(bad)).unwrap_err();
assert!(matches!(err, Error::Config(_)));
assert!(!err.to_string().contains(bad), "leaked key: {err}");
}
assert_eq!(
resolve_api_key(Some(" sk-live-abc123 ")).unwrap(),
"sk-live-abc123"
);
}
#[test]
fn timeout_validation() {
assert!(validate_timeout(Duration::from_millis(1)).is_ok());
let err = validate_timeout(Duration::ZERO).unwrap_err();
assert!(matches!(err, Error::Config(_)));
}
#[test]
fn base_url_strips_trailing_slash() {
assert_eq!(
resolve_base_url(Some("https://example.com/")),
"https://example.com"
);
assert_eq!(
resolve_base_url(Some("https://example.com//")),
"https://example.com"
);
}
}