use crate::core::error::Error;
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
const MAX_QUERY_LENGTH: usize = 100_000;
const MAX_SYSTEM_PROMPT_LENGTH: usize = 10_000;
const MIN_TIMEOUT_SECS: u64 = 1;
const MAX_TIMEOUT_SECS: u64 = 3600; const MAX_TOKENS_LIMIT: usize = 200_000;
const MAX_TOOL_NAME_LENGTH: usize = 100;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Config {
#[serde(skip_serializing_if = "Option::is_none")]
pub system_prompt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub mcp_config_path: Option<PathBuf>,
#[serde(skip_serializing_if = "Option::is_none")]
pub allowed_tools: Option<Vec<String>>,
#[serde(default)]
pub stream_format: StreamFormat,
#[serde(default)]
pub non_interactive: bool,
#[serde(default)]
pub verbose: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub timeout_secs: Option<u64>,
#[serde(default)]
pub continue_session: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub resume_session_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub append_system_prompt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub disallowed_tools: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_turns: Option<u32>,
#[serde(default = "default_skip_permissions")]
pub skip_permissions: bool,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize, Default, PartialEq)]
#[serde(rename_all = "lowercase")]
pub enum StreamFormat {
#[default]
Text,
Json,
StreamJson,
}
fn default_skip_permissions() -> bool {
true
}
impl Default for Config {
fn default() -> Self {
Self {
system_prompt: None,
model: None,
mcp_config_path: None,
allowed_tools: None,
stream_format: StreamFormat::default(),
non_interactive: true,
verbose: false,
max_tokens: None,
timeout_secs: Some(30), continue_session: false,
resume_session_id: None,
append_system_prompt: None,
disallowed_tools: None,
max_turns: None,
skip_permissions: default_skip_permissions(),
}
}
}
impl Config {
pub fn builder() -> ConfigBuilder {
ConfigBuilder::new()
}
pub fn validate(&self) -> Result<(), Error> {
if let Some(prompt) = &self.system_prompt {
if prompt.len() > MAX_SYSTEM_PROMPT_LENGTH {
return Err(Error::InvalidInput(format!(
"System prompt exceeds maximum length of {} characters (got {})",
MAX_SYSTEM_PROMPT_LENGTH,
prompt.len()
)));
}
if contains_malicious_patterns(prompt) {
return Err(Error::InvalidInput(
"System prompt contains potentially malicious content".to_string(),
));
}
}
if let Some(timeout) = self.timeout_secs {
if timeout < MIN_TIMEOUT_SECS || timeout > MAX_TIMEOUT_SECS {
return Err(Error::InvalidInput(format!(
"Timeout must be between {} and {} seconds (got {})",
MIN_TIMEOUT_SECS, MAX_TIMEOUT_SECS, timeout
)));
}
}
if let Some(max_tokens) = self.max_tokens {
if max_tokens == 0 || max_tokens > MAX_TOKENS_LIMIT {
return Err(Error::InvalidInput(format!(
"Max tokens must be between 1 and {} (got {})",
MAX_TOKENS_LIMIT, max_tokens
)));
}
}
if let Some(tools) = &self.allowed_tools {
for tool in tools {
if tool.is_empty() || tool.len() > MAX_TOOL_NAME_LENGTH {
return Err(Error::InvalidInput(format!(
"Tool name length must be between 1 and {} characters (got '{}')",
MAX_TOOL_NAME_LENGTH, tool
)));
}
if let Err(e) = crate::core::types::ToolPermission::parse_granular(tool) {
return Err(Error::InvalidInput(format!(
"Invalid tool permission format: '{}'. Error: {}",
tool, e
)));
}
}
}
if let Some(tools) = &self.disallowed_tools {
for tool in tools {
if tool.is_empty() || tool.len() > MAX_TOOL_NAME_LENGTH {
return Err(Error::InvalidInput(format!(
"Disallowed tool name length must be between 1 and {} characters (got '{}')",
MAX_TOOL_NAME_LENGTH, tool
)));
}
if let Err(e) = crate::core::types::ToolPermission::parse_granular(tool) {
return Err(Error::InvalidInput(format!(
"Invalid disallowed tool permission format: '{}'. Error: {}",
tool, e
)));
}
}
}
if let Some(path) = &self.mcp_config_path {
if path.as_os_str().is_empty() {
return Err(Error::InvalidInput(
"MCP config path cannot be empty".to_string(),
));
}
}
if let Some(turns) = self.max_turns {
if turns == 0 {
return Err(Error::InvalidInput(
"Max turns must be greater than 0".to_string(),
));
}
}
if let (Some(allowed), Some(disallowed)) = (&self.allowed_tools, &self.disallowed_tools) {
for tool in disallowed {
if allowed.contains(tool) {
return Err(Error::InvalidInput(format!(
"Tool '{}' cannot be both allowed and disallowed",
tool
)));
}
}
}
if self.system_prompt.is_some() && self.append_system_prompt.is_some() {
return Err(Error::InvalidInput(
"Cannot use both system_prompt and append_system_prompt simultaneously".to_string(),
));
}
if let Some(prompt) = &self.append_system_prompt {
if prompt.len() > MAX_SYSTEM_PROMPT_LENGTH {
return Err(Error::InvalidInput(format!(
"Append system prompt exceeds maximum length of {} characters (got {})",
MAX_SYSTEM_PROMPT_LENGTH,
prompt.len()
)));
}
if contains_malicious_patterns(prompt) {
return Err(Error::InvalidInput(
"Append system prompt contains potentially malicious content".to_string(),
));
}
}
if let Some(session_id) = &self.resume_session_id {
if session_id.is_empty() {
return Err(Error::InvalidInput(
"Resume session ID cannot be empty".to_string(),
));
}
if session_id.len() > 100 {
return Err(Error::InvalidInput(
"Resume session ID exceeds maximum length of 100 characters".to_string(),
));
}
if !session_id
.chars()
.all(|c| c.is_alphanumeric() || c == '_' || c == '-')
{
return Err(Error::InvalidInput(
"Resume session ID contains invalid characters. Only alphanumeric, underscore, and hyphen are allowed".to_string(),
));
}
}
Ok(())
}
}
pub struct ConfigBuilder {
config: Config,
}
impl Default for ConfigBuilder {
fn default() -> Self {
Self::new()
}
}
impl ConfigBuilder {
pub fn new() -> Self {
Self {
config: Config::default(),
}
}
#[must_use]
pub fn system_prompt(mut self, prompt: impl Into<String>) -> Self {
self.config.system_prompt = Some(prompt.into());
self
}
#[must_use]
pub fn model(mut self, model: impl Into<String>) -> Self {
self.config.model = Some(model.into());
self
}
#[must_use]
pub fn mcp_config(mut self, path: impl Into<PathBuf>) -> Self {
self.config.mcp_config_path = Some(path.into());
self
}
#[must_use]
pub fn allowed_tools(mut self, tools: Vec<String>) -> Self {
self.config.allowed_tools = Some(tools);
self
}
#[must_use]
pub fn stream_format(mut self, format: StreamFormat) -> Self {
self.config.stream_format = format;
self
}
#[must_use]
pub fn non_interactive(mut self, non_interactive: bool) -> Self {
self.config.non_interactive = non_interactive;
self
}
#[must_use]
pub fn max_tokens(mut self, max_tokens: usize) -> Self {
self.config.max_tokens = Some(max_tokens);
self
}
#[must_use]
pub fn timeout_secs(mut self, timeout_secs: u64) -> Self {
self.config.timeout_secs = Some(timeout_secs);
self
}
#[must_use]
pub fn verbose(mut self, verbose: bool) -> Self {
self.config.verbose = verbose;
self
}
#[must_use]
pub fn continue_session(mut self) -> Self {
self.config.continue_session = true;
self
}
#[must_use]
pub fn resume_session(mut self, session_id: String) -> Self {
self.config.resume_session_id = Some(session_id);
self
}
#[must_use]
pub fn append_system_prompt(mut self, prompt: impl Into<String>) -> Self {
self.config.append_system_prompt = Some(prompt.into());
self
}
#[must_use]
pub fn disallowed_tools(mut self, tools: Vec<String>) -> Self {
self.config.disallowed_tools = Some(tools);
self
}
#[must_use]
pub fn max_turns(mut self, turns: u32) -> Self {
self.config.max_turns = Some(turns);
self
}
#[must_use]
pub fn skip_permissions(mut self, skip: bool) -> Self {
self.config.skip_permissions = skip;
self
}
pub fn build(self) -> Result<Config, Error> {
self.config.validate()?;
Ok(self.config)
}
}
pub fn validate_query(query: &str) -> Result<(), Error> {
if query.is_empty() {
return Err(Error::InvalidInput("Query cannot be empty".to_string()));
}
if query.len() > MAX_QUERY_LENGTH {
return Err(Error::InvalidInput(format!(
"Query exceeds maximum length of {} characters (got {})",
MAX_QUERY_LENGTH,
query.len()
)));
}
if contains_malicious_patterns(query) {
return Err(Error::InvalidInput(
"Query contains potentially malicious content".to_string(),
));
}
Ok(())
}
fn contains_malicious_patterns(text: &str) -> bool {
let malicious_patterns = [
"<script",
"javascript:",
"onclick=",
"onerror=",
"$(",
"${",
"`",
"&&",
"||",
";",
"|",
">",
"<",
"../",
"..\\",
"' OR ",
"\" OR ",
"'; DROP",
"\0",
];
let lower_text = text.to_lowercase();
malicious_patterns
.iter()
.any(|pattern| lower_text.contains(pattern))
}