make claude hold a weak reference back to gateway
This commit is contained in:
@@ -7,6 +7,7 @@ use conf::Conf;
|
|||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
use std::path::PathBuf;
|
use std::path::PathBuf;
|
||||||
|
use std::sync::Weak;
|
||||||
use tracing::info;
|
use tracing::info;
|
||||||
|
|
||||||
/// The Anthropic API version header value. This is a stable version identifier,
|
/// The Anthropic API version header value. This is a stable version identifier,
|
||||||
@@ -60,6 +61,9 @@ pub enum ClaudeError {
|
|||||||
/// Too many tool use iterations.
|
/// Too many tool use iterations.
|
||||||
#[error("exceeded maximum tool use iterations ({0})")]
|
#[error("exceeded maximum tool use iterations ({0})")]
|
||||||
TooManyIterations(u32),
|
TooManyIterations(u32),
|
||||||
|
/// Tool executor is no longer available.
|
||||||
|
#[error("tool executor is gone")]
|
||||||
|
ToolExecutorGone,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Claude API client.
|
/// Claude API client.
|
||||||
@@ -68,6 +72,7 @@ pub struct ClaudeApi {
|
|||||||
client: reqwest::Client,
|
client: reqwest::Client,
|
||||||
api_key: String,
|
api_key: String,
|
||||||
system_prompt: String,
|
system_prompt: String,
|
||||||
|
tool_executor: Weak<dyn ToolExecutor>,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Request body for the Claude Messages API.
|
/// Request body for the Claude Messages API.
|
||||||
@@ -159,7 +164,13 @@ impl ClaudeApi {
|
|||||||
/// Create a new Claude API client from configuration.
|
/// Create a new Claude API client from configuration.
|
||||||
///
|
///
|
||||||
/// Reads the API key and system prompt from the configured files.
|
/// Reads the API key and system prompt from the configured files.
|
||||||
pub fn new(config: ClaudeConfig) -> Result<Self, ClaudeError> {
|
/// The tool executor is held as a weak reference, so it will not prevent
|
||||||
|
/// the executor from being dropped. If the executor is dropped during a
|
||||||
|
/// request, tool use will fail with `ToolExecutorGone`.
|
||||||
|
pub fn new(
|
||||||
|
config: ClaudeConfig,
|
||||||
|
tool_executor: Weak<dyn ToolExecutor>,
|
||||||
|
) -> Result<Self, ClaudeError> {
|
||||||
let api_key = std::fs::read_to_string(&config.api_key_file)
|
let api_key = std::fs::read_to_string(&config.api_key_file)
|
||||||
.map_err(ClaudeError::ApiKeyRead)?
|
.map_err(ClaudeError::ApiKeyRead)?
|
||||||
.trim()
|
.trim()
|
||||||
@@ -175,19 +186,21 @@ impl ClaudeApi {
|
|||||||
client,
|
client,
|
||||||
api_key,
|
api_key,
|
||||||
system_prompt,
|
system_prompt,
|
||||||
|
tool_executor,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Send a request to the Claude API and return the response text.
|
/// Send a request to the Claude API and return the response text.
|
||||||
/// If a tool executor is provided, handles tool use in a loop.
|
/// If a tool executor has been installed, handles tool use in a loop.
|
||||||
pub async fn request(
|
pub async fn request(&self, prompt: &str) -> Result<String, ClaudeError> {
|
||||||
&self,
|
|
||||||
prompt: &str,
|
|
||||||
tool_executor: Option<&dyn ToolExecutor>,
|
|
||||||
) -> Result<String, ClaudeError> {
|
|
||||||
let max_iterations = self.config.claude_max_iterations;
|
let max_iterations = self.config.claude_max_iterations;
|
||||||
|
|
||||||
let tools = tool_executor.map(|te| te.tools()).unwrap_or_default();
|
// Get tools from the executor if still alive
|
||||||
|
let executor = self.tool_executor.upgrade();
|
||||||
|
let tools = executor
|
||||||
|
.as_ref()
|
||||||
|
.map(|te| te.tools())
|
||||||
|
.unwrap_or_default();
|
||||||
let mut messages = vec![MessageContent::user(prompt)];
|
let mut messages = vec![MessageContent::user(prompt)];
|
||||||
|
|
||||||
info!("Claude request: {}", prompt);
|
info!("Claude request: {}", prompt);
|
||||||
@@ -224,10 +237,9 @@ impl ClaudeApi {
|
|||||||
|
|
||||||
// Check if we need to handle tool use
|
// Check if we need to handle tool use
|
||||||
if response.stop_reason == "tool_use" {
|
if response.stop_reason == "tool_use" {
|
||||||
let Some(executor) = tool_executor else {
|
// Try to get a strong reference to the executor
|
||||||
return Err(ClaudeError::ToolError(
|
let Some(executor) = self.tool_executor.upgrade() else {
|
||||||
"tool use requested but no executor provided".to_owned(),
|
return Err(ClaudeError::ToolExecutorGone);
|
||||||
));
|
|
||||||
};
|
};
|
||||||
|
|
||||||
// Add assistant's response to messages
|
// Add assistant's response to messages
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ use http::{Method, Request, Response, StatusCode};
|
|||||||
use http_body::Body;
|
use http_body::Body;
|
||||||
use http_body_util::BodyExt;
|
use http_body_util::BodyExt;
|
||||||
use prometheus_http_client::{AlertStatus, ExtractLabels};
|
use prometheus_http_client::{AlertStatus, ExtractLabels};
|
||||||
use std::{fmt::Write, net::SocketAddr, path::PathBuf, sync::Arc, sync::Mutex, time::Duration};
|
use std::{fmt::Write, net::SocketAddr, path::PathBuf, sync::Arc, sync::Mutex, sync::OnceLock, sync::Weak, time::Duration};
|
||||||
use tokio::{
|
use tokio::{
|
||||||
join,
|
join,
|
||||||
sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel},
|
sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel},
|
||||||
@@ -230,7 +230,8 @@ pub struct Gateway {
|
|||||||
/// Handler for admin messages that don't start with `/`
|
/// Handler for admin messages that don't start with `/`
|
||||||
message_handler: Option<Box<dyn MessageHandler>>,
|
message_handler: Option<Box<dyn MessageHandler>>,
|
||||||
/// Claude API client for AI-powered responses.
|
/// Claude API client for AI-powered responses.
|
||||||
claude: Option<ClaudeApi>,
|
/// Initialized after Arc creation so it can hold a weak reference back to Gateway.
|
||||||
|
claude: OnceLock<Box<ClaudeApi>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Gateway {
|
impl Gateway {
|
||||||
@@ -251,12 +252,9 @@ impl Gateway {
|
|||||||
|
|
||||||
let log_handler = LogHandler::new(config.log_handler.clone(), signal_alert_mq_tx.clone());
|
let log_handler = LogHandler::new(config.log_handler.clone(), signal_alert_mq_tx.clone());
|
||||||
|
|
||||||
let claude = config
|
let claude_config = config.claude.clone();
|
||||||
.claude
|
|
||||||
.clone()
|
|
||||||
.map(|cc| ClaudeApi::new(cc).expect("Invalid claude config"));
|
|
||||||
|
|
||||||
Arc::new(Self {
|
let gateway = Arc::new(Self {
|
||||||
config,
|
config,
|
||||||
signal_alert_mq_tx,
|
signal_alert_mq_tx,
|
||||||
signal_alert_mq_rx: Mutex::new(Some(signal_alert_mq_rx)),
|
signal_alert_mq_rx: Mutex::new(Some(signal_alert_mq_rx)),
|
||||||
@@ -264,8 +262,20 @@ impl Gateway {
|
|||||||
prometheus,
|
prometheus,
|
||||||
log_handler,
|
log_handler,
|
||||||
message_handler,
|
message_handler,
|
||||||
claude,
|
claude: OnceLock::new(),
|
||||||
})
|
});
|
||||||
|
|
||||||
|
// Initialize Claude with a weak reference back to the gateway
|
||||||
|
if let Some(cc) = claude_config {
|
||||||
|
let claude = ClaudeApi::new(cc, Arc::downgrade(&gateway) as Weak<dyn ToolExecutor>)
|
||||||
|
.expect("Invalid claude config");
|
||||||
|
gateway
|
||||||
|
.claude
|
||||||
|
.set(Box::new(claude))
|
||||||
|
.unwrap_or_else(|_| panic!("claude OnceLock was already set"));
|
||||||
|
}
|
||||||
|
|
||||||
|
gateway
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Run the gateway main loop, reconnecting to signal-cli on errors.
|
/// Run the gateway main loop, reconnecting to signal-cli on errors.
|
||||||
@@ -678,12 +688,11 @@ impl Gateway {
|
|||||||
GatewayCommand::Claude { prompt } => {
|
GatewayCommand::Claude { prompt } => {
|
||||||
let claude = self
|
let claude = self
|
||||||
.claude
|
.claude
|
||||||
.as_ref()
|
.get()
|
||||||
.ok_or_else(|| (501u16, "claude was not configured".into()))?;
|
.ok_or_else(|| (501u16, "claude was not configured".into()))?;
|
||||||
|
|
||||||
let prompt_text = prompt.join(" ");
|
let prompt_text = prompt.join(" ");
|
||||||
// Pass self as the tool executor so Claude can query prometheus
|
match claude.request(&prompt_text).await {
|
||||||
match claude.request(&prompt_text, Some(self)).await {
|
|
||||||
Ok(response) => Ok(AdminMessageResponse::new(response)),
|
Ok(response) => Ok(AdminMessageResponse::new(response)),
|
||||||
Err(err) => Err((500, err.to_string().into())),
|
Err(err) => Err((500, err.to_string().into())),
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user