fix(telegram): prompt for tool approval in gateway sessions (#10613)

Signed-off-by: Abhijay Jain <Abhijay007j@gmail.com>
This commit is contained in:
Abhijay Jain
2026-07-28 01:07:22 +05:30
committed by GitHub
parent 971d217842
commit a074d8eb3e
3 changed files with 208 additions and 5 deletions
+191 -3
View File
@@ -1,16 +1,20 @@
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use futures::StreamExt;
use tokio::sync::Mutex;
use tokio_util::sync::CancellationToken;
use crate::agents::{AgentEvent, ExtensionConfig, SessionConfig};
use crate::agents::{Agent, AgentEvent, ExtensionConfig, SessionConfig};
use crate::config::extensions::get_enabled_extensions;
use crate::config::paths::Paths;
use crate::config::Config;
use crate::conversation::message::{Message, MessageContent};
use crate::conversation::message::{ActionRequiredData, Message, MessageContent};
use crate::execution::manager::AgentManager;
use crate::permission::permission_confirmation::PrincipalType;
use crate::permission::{Permission, PermissionConfirmation};
use crate::session::SessionType;
use crate::session::{EnabledExtensionsState, ExtensionState, Session};
@@ -36,12 +40,21 @@ fn resolve_gateway_max_turns(gateway_override: Option<u32>, global_max_turns: Op
.unwrap_or(DEFAULT_GATEWAY_MAX_TURNS)
}
struct PendingConfirmation {
agent: Arc<Agent>,
request_id: String,
}
#[derive(Clone)]
pub struct GatewayHandler {
agent_manager: Arc<AgentManager>,
pairing_store: Arc<PairingStore>,
gateway: Arc<dyn Gateway>,
config: GatewayConfig,
/// Tracks users who have a tool-confirmation prompt awaiting their reply.
pending_confirmations: Arc<Mutex<HashMap<PlatformUser, PendingConfirmation>>>,
/// Serializes `relay_to_session` per user; confirmation replies bypass this lock.
turn_locks: Arc<Mutex<HashMap<PlatformUser, Arc<Mutex<()>>>>>,
}
impl GatewayHandler {
@@ -56,6 +69,34 @@ impl GatewayHandler {
pairing_store,
gateway,
config,
pending_confirmations: Arc::new(Mutex::new(HashMap::new())),
turn_locks: Arc::new(Mutex::new(HashMap::new())),
}
}
pub async fn deny_pending_confirmations(&self) {
let pending: Vec<_> = self.pending_confirmations.lock().await.drain().collect();
for (_, confirmation) in pending {
confirmation
.agent
.handle_confirmation(
confirmation.request_id,
PermissionConfirmation {
principal_type: PrincipalType::Tool,
permission: Permission::DenyOnce,
},
)
.await;
}
}
async fn prune_turn_lock(&self, user: &PlatformUser) {
let mut locks = self.turn_locks.lock().await;
let in_use = locks
.get(user)
.is_some_and(|lock| Arc::strong_count(lock) > 1);
if !in_use {
locks.remove(user);
}
}
@@ -117,13 +158,81 @@ impl GatewayHandler {
}
}
PairingState::Paired { session_id, .. } => {
self.relay_to_session(&message, &session_id).await?;
if self
.pending_confirmations
.lock()
.await
.contains_key(&message.user)
{
self.handle_pending_confirmation(&message).await?;
} else {
let turn_lock = {
let mut locks = self.turn_locks.lock().await;
Arc::clone(
locks
.entry(message.user.clone())
.or_insert_with(|| Arc::new(Mutex::new(()))),
)
};
let turn_guard = turn_lock.lock().await;
let result = self.relay_to_session(&message, &session_id).await;
drop(turn_guard);
drop(turn_lock);
self.prune_turn_lock(&message.user).await;
result?;
}
}
}
Ok(())
}
async fn handle_pending_confirmation(&self, message: &IncomingMessage) -> anyhow::Result<()> {
let text = message.text.trim().to_lowercase();
let permission = match text.as_str() {
"approve" | "yes" | "y" => Some(Permission::AllowOnce),
"approve always" => Some(Permission::AlwaysAllow),
"deny" | "no" | "n" => Some(Permission::DenyOnce),
"deny always" => Some(Permission::AlwaysDeny),
_ => None,
};
let Some(permission) = permission else {
self.gateway
.send_message(
&message.user,
OutgoingMessage::Text {
body: "Please reply with: approve, approve always, deny, or deny always."
.into(),
},
)
.await?;
return Ok(());
};
let Some(pending) = self
.pending_confirmations
.lock()
.await
.remove(&message.user)
else {
return Ok(());
};
pending
.agent
.handle_confirmation(
pending.request_id,
PermissionConfirmation {
principal_type: PrincipalType::Tool,
permission,
},
)
.await;
Ok(())
}
async fn try_consume_code(&self, text: &str) -> anyhow::Result<Option<String>> {
let normalized = text.to_uppercase().replace(['-', ' '], "");
if normalized.len() == 6
@@ -474,6 +583,85 @@ impl GatewayHandler {
"gateway stream: tool response"
);
}
MessageContent::ActionRequired(action_required) => {
if let ActionRequiredData::ToolConfirmation {
id,
tool_name,
arguments,
prompt,
} = &action_required.data
{
// Flush pending text so the user sees context first.
if !pending_text.is_empty() {
let _ = self
.gateway
.send_message(
&message.user,
OutgoingMessage::Text {
body: std::mem::take(&mut pending_text),
},
)
.await;
}
let args_display = serde_json::to_string_pretty(arguments)
.unwrap_or_default();
let mut approval_text = format!(
"🔐 Approval required\n\nTool: {tool_name}\nArguments:\n{args_display}"
);
if let Some(p) = prompt {
approval_text.push_str(&format!("\n\n{p}"));
}
approval_text.push_str(
"\n\nReply with:\n\
• approve — allow once\n\
• approve always — always allow\n\
• deny — deny once\n\
• deny always — always deny",
);
self.pending_confirmations.lock().await.insert(
message.user.clone(),
PendingConfirmation {
agent: agent.clone(),
request_id: id.clone(),
},
);
let send_result = self
.gateway
.send_message(
&message.user,
OutgoingMessage::Text {
body: approval_text,
},
)
.await;
if let Err(e) = send_result {
tracing::error!(
session_id,
error = %e,
"failed to deliver tool approval prompt; denying tool call"
);
self.pending_confirmations
.lock()
.await
.remove(&message.user);
agent
.handle_confirmation(
id.clone(),
PermissionConfirmation {
principal_type: PrincipalType::Tool,
permission: Permission::DenyOnce,
},
)
.await;
} else {
sent_any = true;
}
}
}
_ => {}
}
}
+7 -1
View File
@@ -31,6 +31,7 @@ pub struct GatewayInstance {
pub gateway: Arc<dyn Gateway>,
pub cancel: CancellationToken,
pub handle: tokio::task::JoinHandle<()>,
handler: GatewayHandler,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -129,6 +130,7 @@ impl GatewayManager {
.remove(gateway_type)
.ok_or_else(|| anyhow::anyhow!("Gateway '{}' is not running", gateway_type))?;
instance.handler.deny_pending_confirmations().await;
instance.cancel.cancel();
let _ = instance.handle.await;
@@ -140,6 +142,7 @@ impl GatewayManager {
pub async fn remove_gateway(&self, gateway_type: &str) -> anyhow::Result<()> {
// Stop if running (ignore error if not running).
if let Some(instance) = self.gateways.write().await.remove(gateway_type) {
instance.handler.deny_pending_confirmations().await;
instance.cancel.cancel();
let _ = instance.handle.await;
}
@@ -195,13 +198,14 @@ impl GatewayManager {
gateway.clone(),
config.clone(),
);
let handler_for_task = handler.clone();
let gateway_clone = gateway.clone();
let cancel_clone = cancel.clone();
let gateway_type_for_task = gw_type.clone();
let handle = tokio::spawn(async move {
if let Err(e) = gateway_clone.start(handler, cancel_clone).await {
if let Err(e) = gateway_clone.start(handler_for_task, cancel_clone).await {
tracing::error!(gateway = %gateway_type_for_task, error = %e, "gateway stopped with error");
}
});
@@ -211,6 +215,7 @@ impl GatewayManager {
gateway,
cancel,
handle,
handler,
};
self.gateways.write().await.insert(gw_type, instance);
@@ -222,6 +227,7 @@ impl GatewayManager {
let instances: Vec<(String, GatewayInstance)> =
self.gateways.write().await.drain().collect();
for (gateway_type, instance) in instances {
instance.handler.deny_pending_confirmations().await;
instance.cancel.cancel();
let _ = instance.handle.await;
tracing::info!(gateway = %gateway_type, "gateway stopped");
+10 -1
View File
@@ -161,14 +161,23 @@ impl TelegramGateway {
"Telegram rejected HTML, falling back to plain text"
);
for plain_chunk in split_message(text, MAX_MESSAGE_LENGTH) {
self.client
let plain_resp = self
.client
.post(self.api_url("sendMessage"))
.json(&serde_json::json!({
"chat_id": chat_id,
"text": plain_chunk,
}))
.send()
.await?
.json::<TelegramResponse<serde_json::Value>>()
.await?;
if !plain_resp.ok {
anyhow::bail!(
"Telegram sendMessage failed: {}",
plain_resp.description.unwrap_or_default()
);
}
}
return Ok(());
}