add support for multiple system prompt files
This commit is contained in:
@@ -30,14 +30,14 @@ pub struct ClaudeConfig {
|
|||||||
/// Path to file containing the Claude API key.
|
/// Path to file containing the Claude API key.
|
||||||
#[conf(long, env)]
|
#[conf(long, env)]
|
||||||
pub api_key_file: PathBuf,
|
pub api_key_file: PathBuf,
|
||||||
/// Path to file containing the system prompt.
|
/// Paths to files containing system prompt components (in order, last one is cached).
|
||||||
#[conf(long, env)]
|
#[conf(repeat, long, env)]
|
||||||
pub system_prompt_file: PathBuf,
|
pub system_prompt_files: Vec<PathBuf>,
|
||||||
/// Claude API URL.
|
/// Claude API URL.
|
||||||
#[conf(long, env, default_value = "https://api.anthropic.com/v1/messages")]
|
#[conf(long, env, default_value = "https://api.anthropic.com/v1/messages")]
|
||||||
pub claude_api_url: String,
|
pub claude_api_url: String,
|
||||||
/// Claude model to use.
|
/// Claude model to use.
|
||||||
#[conf(long, env, default_value = "claude-opus-4-20250514")]
|
#[conf(long, env, default_value = "claude-sonnet-4-5-20250929")]
|
||||||
pub claude_model: String,
|
pub claude_model: String,
|
||||||
/// Maximum tokens in the response.
|
/// Maximum tokens in the response.
|
||||||
#[conf(long, env, default_value = "1024")]
|
#[conf(long, env, default_value = "1024")]
|
||||||
@@ -54,8 +54,8 @@ pub enum ClaudeError {
|
|||||||
#[error("failed to read API key file: {0}")]
|
#[error("failed to read API key file: {0}")]
|
||||||
ApiKeyRead(std::io::Error),
|
ApiKeyRead(std::io::Error),
|
||||||
/// Failed to read system prompt file.
|
/// Failed to read system prompt file.
|
||||||
#[error("failed to read system prompt file: {0}")]
|
#[error("failed to read system prompt file {0:?}: {1}")]
|
||||||
SystemPromptRead(std::io::Error),
|
SystemPromptRead(PathBuf, std::io::Error),
|
||||||
/// HTTP request failed.
|
/// HTTP request failed.
|
||||||
#[error("HTTP request failed: {0}")]
|
#[error("HTTP request failed: {0}")]
|
||||||
Request(#[from] reqwest::Error),
|
Request(#[from] reqwest::Error),
|
||||||
|
|||||||
@@ -75,7 +75,7 @@ pub struct ClaudeWorker {
|
|||||||
config: ClaudeConfig,
|
config: ClaudeConfig,
|
||||||
client: reqwest::Client,
|
client: reqwest::Client,
|
||||||
api_key: String,
|
api_key: String,
|
||||||
system_prompt: String,
|
system_prompts: Vec<String>,
|
||||||
// FIXME: use this and append to system prompt within <summary> </summary> tags
|
// FIXME: use this and append to system prompt within <summary> </summary> tags
|
||||||
#[allow(dead_code)]
|
#[allow(dead_code)]
|
||||||
summary: String,
|
summary: String,
|
||||||
@@ -88,7 +88,7 @@ pub struct ClaudeWorker {
|
|||||||
impl ClaudeWorker {
|
impl ClaudeWorker {
|
||||||
/// Create a new Claude worker.
|
/// Create a new Claude worker.
|
||||||
///
|
///
|
||||||
/// Reads the API key and system prompt from the configured files.
|
/// Reads the API key and system prompts from the configured files.
|
||||||
pub fn new(
|
pub fn new(
|
||||||
config: ClaudeConfig,
|
config: ClaudeConfig,
|
||||||
tool_executor: Weak<dyn ToolExecutor>,
|
tool_executor: Weak<dyn ToolExecutor>,
|
||||||
@@ -100,14 +100,20 @@ impl ClaudeWorker {
|
|||||||
.trim()
|
.trim()
|
||||||
.to_owned();
|
.to_owned();
|
||||||
|
|
||||||
let system_prompt = std::fs::read_to_string(&config.system_prompt_file)
|
let system_prompts: Vec<String> = config
|
||||||
.map_err(ClaudeError::SystemPromptRead)?;
|
.system_prompt_files
|
||||||
|
.iter()
|
||||||
|
.map(|path| {
|
||||||
|
std::fs::read_to_string(path)
|
||||||
|
.map_err(|e| ClaudeError::SystemPromptRead(path.clone(), e))
|
||||||
|
})
|
||||||
|
.collect::<Result<_, _>>()?;
|
||||||
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
config,
|
config,
|
||||||
client: reqwest::Client::new(),
|
client: reqwest::Client::new(),
|
||||||
api_key,
|
api_key,
|
||||||
system_prompt,
|
system_prompts,
|
||||||
summary: String::new(),
|
summary: String::new(),
|
||||||
messages: Default::default(),
|
messages: Default::default(),
|
||||||
tool_executor,
|
tool_executor,
|
||||||
@@ -211,6 +217,21 @@ impl ClaudeWorker {
|
|||||||
info!("Claude request: {}", text);
|
info!("Claude request: {}", text);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Build system content blocks, caching only the last one
|
||||||
|
let system: Vec<SystemContent> = self
|
||||||
|
.system_prompts
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.map(|(i, text)| {
|
||||||
|
let content = SystemContent::text(text);
|
||||||
|
if i == self.system_prompts.len() - 1 {
|
||||||
|
content.cached()
|
||||||
|
} else {
|
||||||
|
content
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
for iteration in 0..max_iterations {
|
for iteration in 0..max_iterations {
|
||||||
// Check for stop before making API call
|
// Check for stop before making API call
|
||||||
self.check_stop()?;
|
self.check_stop()?;
|
||||||
@@ -218,7 +239,7 @@ impl ClaudeWorker {
|
|||||||
let request_body = MessagesRequest {
|
let request_body = MessagesRequest {
|
||||||
model: &self.config.claude_model,
|
model: &self.config.claude_model,
|
||||||
max_tokens: self.config.claude_max_tokens,
|
max_tokens: self.config.claude_max_tokens,
|
||||||
system: &self.system_prompt,
|
system: &system,
|
||||||
messages: &self.messages,
|
messages: &self.messages,
|
||||||
tools: tools.clone(),
|
tools: tools.clone(),
|
||||||
};
|
};
|
||||||
@@ -314,12 +335,52 @@ impl ClaudeWorker {
|
|||||||
struct MessagesRequest<'a> {
|
struct MessagesRequest<'a> {
|
||||||
model: &'a str,
|
model: &'a str,
|
||||||
max_tokens: u32,
|
max_tokens: u32,
|
||||||
system: &'a str,
|
system: &'a [SystemContent],
|
||||||
messages: &'a [MessageContent],
|
messages: &'a [MessageContent],
|
||||||
#[serde(skip_serializing_if = "Vec::is_empty")]
|
#[serde(skip_serializing_if = "Vec::is_empty")]
|
||||||
tools: Vec<Tool>,
|
tools: Vec<Tool>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// A content block in the system prompt array.
|
||||||
|
#[derive(Clone, Serialize)]
|
||||||
|
struct SystemContent {
|
||||||
|
#[serde(rename = "type")]
|
||||||
|
content_type: &'static str,
|
||||||
|
text: String,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
cache_control: Option<CacheControl>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl SystemContent {
|
||||||
|
fn text(text: impl Into<String>) -> Self {
|
||||||
|
Self {
|
||||||
|
content_type: "text",
|
||||||
|
text: text.into(),
|
||||||
|
cache_control: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn cached(mut self) -> Self {
|
||||||
|
self.cache_control = Some(CacheControl::ephemeral());
|
||||||
|
self
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Cache control directive for prompt caching.
|
||||||
|
#[derive(Clone, Serialize)]
|
||||||
|
struct CacheControl {
|
||||||
|
#[serde(rename = "type")]
|
||||||
|
cache_type: &'static str,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl CacheControl {
|
||||||
|
fn ephemeral() -> Self {
|
||||||
|
Self {
|
||||||
|
cache_type: "ephemeral",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// A message in the conversation (can have multiple content blocks).
|
/// A message in the conversation (can have multiple content blocks).
|
||||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||||
struct MessageContent {
|
struct MessageContent {
|
||||||
|
|||||||
Reference in New Issue
Block a user