make claude hold a weak reference back to gateway

This commit is contained in:
Chris Beck
2025-12-07 16:28:00 -07:00
parent 739075f2e2
commit 70d09ffc98
2 changed files with 45 additions and 24 deletions
+24 -12
View File
@@ -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
+21 -12
View File
@@ -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())),
} }