94fdcdd07a
Signed-off-by: Adrian Cole <adrian@tetrate.io>
523 lines
20 KiB
Rust
523 lines
20 KiB
Rust
use std::path::PathBuf;
|
|
use std::sync::Arc;
|
|
use std::time::Duration;
|
|
|
|
use futures::StreamExt;
|
|
use tokio_util::sync::CancellationToken;
|
|
|
|
use crate::agents::{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::execution::manager::AgentManager;
|
|
use crate::model::ModelConfig;
|
|
use crate::session::SessionType;
|
|
use crate::session::{EnabledExtensionsState, ExtensionState, Session};
|
|
|
|
use super::pairing::PairingStore;
|
|
use super::{Gateway, GatewayConfig, IncomingMessage, OutgoingMessage, PairingState, PlatformUser};
|
|
|
|
#[derive(Clone)]
|
|
pub struct GatewayHandler {
|
|
agent_manager: Arc<AgentManager>,
|
|
pairing_store: Arc<PairingStore>,
|
|
gateway: Arc<dyn Gateway>,
|
|
config: GatewayConfig,
|
|
}
|
|
|
|
impl GatewayHandler {
|
|
pub fn new(
|
|
agent_manager: Arc<AgentManager>,
|
|
pairing_store: Arc<PairingStore>,
|
|
gateway: Arc<dyn Gateway>,
|
|
config: GatewayConfig,
|
|
) -> Self {
|
|
Self {
|
|
agent_manager,
|
|
pairing_store,
|
|
gateway,
|
|
config,
|
|
}
|
|
}
|
|
|
|
pub async fn handle_message(&self, message: IncomingMessage) -> anyhow::Result<()> {
|
|
let pairing = self.pairing_store.get(&message.user).await?;
|
|
|
|
match pairing {
|
|
PairingState::Unpaired => {
|
|
if let Some(gateway_type) = self.try_consume_code(message.text.trim()).await? {
|
|
if gateway_type == self.config.gateway_type {
|
|
self.complete_pairing(&message.user).await?;
|
|
} else {
|
|
self.gateway
|
|
.send_message(
|
|
&message.user,
|
|
OutgoingMessage::Text {
|
|
body: "⚠️ That code is for a different gateway.".into(),
|
|
},
|
|
)
|
|
.await?;
|
|
}
|
|
} else {
|
|
self.gateway
|
|
.send_message(
|
|
&message.user,
|
|
OutgoingMessage::Text {
|
|
body: "Welcome! Enter your pairing code to connect to goose."
|
|
.into(),
|
|
},
|
|
)
|
|
.await?;
|
|
}
|
|
}
|
|
PairingState::PendingCode { code, expires_at } => {
|
|
let now = chrono::Utc::now().timestamp();
|
|
if now > expires_at {
|
|
self.pairing_store
|
|
.set(&message.user, PairingState::Unpaired)
|
|
.await?;
|
|
self.gateway
|
|
.send_message(
|
|
&message.user,
|
|
OutgoingMessage::Text {
|
|
body: "Your pairing code expired. Please request a new one.".into(),
|
|
},
|
|
)
|
|
.await?;
|
|
} else if message.text.trim().eq_ignore_ascii_case(&code) {
|
|
self.complete_pairing(&message.user).await?;
|
|
} else {
|
|
self.gateway
|
|
.send_message(
|
|
&message.user,
|
|
OutgoingMessage::Text {
|
|
body: "Invalid code. Please try again.".into(),
|
|
},
|
|
)
|
|
.await?;
|
|
}
|
|
}
|
|
PairingState::Paired { session_id, .. } => {
|
|
self.relay_to_session(&message, &session_id).await?;
|
|
}
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
async fn try_consume_code(&self, text: &str) -> anyhow::Result<Option<String>> {
|
|
let normalized = text.to_uppercase().replace(['-', ' '], "");
|
|
if normalized.len() == 6
|
|
&& normalized
|
|
.chars()
|
|
.all(|c| "ABCDEFGHJKLMNPQRSTUVWXYZ23456789".contains(c))
|
|
{
|
|
return self.pairing_store.consume_pending_code(&normalized).await;
|
|
}
|
|
Ok(None)
|
|
}
|
|
|
|
async fn complete_pairing(&self, user: &PlatformUser) -> anyhow::Result<()> {
|
|
let working_dir = gateway_working_dir(&user.platform, &user.user_id);
|
|
std::fs::create_dir_all(&working_dir)?;
|
|
|
|
let session_name = format!(
|
|
"{}/{}",
|
|
user.platform,
|
|
user.display_name.as_deref().unwrap_or(&user.user_id)
|
|
);
|
|
|
|
let config = Config::global();
|
|
let session = self
|
|
.agent_manager
|
|
.session_manager()
|
|
.create_session(
|
|
working_dir,
|
|
session_name,
|
|
SessionType::Gateway,
|
|
config.get_goose_mode().unwrap_or_default(),
|
|
)
|
|
.await?;
|
|
|
|
let manager = self.agent_manager.session_manager();
|
|
|
|
// Store the current provider and model config on the session so the agent
|
|
// can be restored after LRU eviction, matching the start_agent flow.
|
|
let mut update = manager.update(&session.id);
|
|
if let Ok(provider) = config.get_goose_provider() {
|
|
update = update.provider_name(provider);
|
|
}
|
|
if let Ok(model_name) = config.get_goose_model() {
|
|
if let Ok(model_config) = ModelConfig::new(&model_name) {
|
|
update = update.model_config(model_config);
|
|
}
|
|
}
|
|
|
|
// Store default extensions so load_extensions_from_session works.
|
|
let extensions = get_enabled_extensions();
|
|
let extensions_state = EnabledExtensionsState::new(extensions);
|
|
let mut extension_data = session.extension_data.clone();
|
|
if let Err(e) = extensions_state.to_extension_data(&mut extension_data) {
|
|
tracing::warn!(error = %e, "failed to initialize gateway session extensions");
|
|
} else {
|
|
update = update.extension_data(extension_data);
|
|
}
|
|
|
|
update.apply().await?;
|
|
|
|
let now = chrono::Utc::now().timestamp();
|
|
self.pairing_store
|
|
.set(
|
|
user,
|
|
PairingState::Paired {
|
|
session_id: session.id.clone(),
|
|
paired_at: now,
|
|
},
|
|
)
|
|
.await?;
|
|
|
|
self.gateway
|
|
.send_message(
|
|
user,
|
|
OutgoingMessage::Text {
|
|
body: "Paired! You can now chat with goose.".into(),
|
|
},
|
|
)
|
|
.await?;
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Sync the session's provider, model, and extensions with the current
|
|
/// global config so gateway sessions always reflect what the user has
|
|
/// configured in the desktop app. Returns `true` if extensions changed
|
|
/// (which means the caller must recreate the agent so stale extension
|
|
/// processes are torn down).
|
|
async fn sync_session_config(&self, session: &Session) -> anyhow::Result<bool> {
|
|
let config = Config::global();
|
|
let manager = self.agent_manager.session_manager();
|
|
|
|
// --- current global config ---
|
|
let current_provider = config.get_goose_provider().ok();
|
|
let current_model_name = config.get_goose_model().ok();
|
|
let current_extensions = get_enabled_extensions();
|
|
let current_mode = config.get_goose_mode().unwrap_or_default();
|
|
|
|
// --- what the session has ---
|
|
let session_extensions: Vec<ExtensionConfig> =
|
|
EnabledExtensionsState::from_extension_data(&session.extension_data)
|
|
.map(|s| s.extensions)
|
|
.unwrap_or_default();
|
|
|
|
let provider_changed = current_provider.as_deref() != session.provider_name.as_deref();
|
|
let model_changed = current_model_name.as_deref()
|
|
!= session.model_config.as_ref().map(|m| m.model_name.as_str());
|
|
let extensions_changed = current_extensions != session_extensions;
|
|
let mode_changed = current_mode != session.goose_mode;
|
|
|
|
if !provider_changed && !model_changed && !extensions_changed && !mode_changed {
|
|
return Ok(false);
|
|
}
|
|
|
|
tracing::info!(
|
|
session_id = %session.id,
|
|
provider_changed,
|
|
model_changed,
|
|
extensions_changed,
|
|
mode_changed,
|
|
"syncing gateway session with current config"
|
|
);
|
|
|
|
let mut update = manager.update(&session.id);
|
|
|
|
if let Some(ref provider) = current_provider {
|
|
update = update.provider_name(provider);
|
|
}
|
|
if let Some(ref model_name) = current_model_name {
|
|
if let Ok(model_config) = ModelConfig::new(model_name) {
|
|
update = update.model_config(model_config);
|
|
}
|
|
}
|
|
|
|
if extensions_changed {
|
|
let extensions_state = EnabledExtensionsState::new(current_extensions);
|
|
let mut extension_data = session.extension_data.clone();
|
|
if let Err(e) = extensions_state.to_extension_data(&mut extension_data) {
|
|
tracing::warn!(error = %e, "failed to update gateway session extensions");
|
|
} else {
|
|
update = update.extension_data(extension_data);
|
|
}
|
|
}
|
|
|
|
if mode_changed {
|
|
update = update.goose_mode(current_mode);
|
|
}
|
|
|
|
update.apply().await?;
|
|
Ok(extensions_changed)
|
|
}
|
|
|
|
async fn relay_to_session(
|
|
&self,
|
|
message: &IncomingMessage,
|
|
session_id: &str,
|
|
) -> anyhow::Result<()> {
|
|
self.gateway
|
|
.send_message(&message.user, OutgoingMessage::Typing)
|
|
.await?;
|
|
|
|
let session = self
|
|
.agent_manager
|
|
.session_manager()
|
|
.get_session(session_id, false)
|
|
.await?;
|
|
|
|
// Sync provider/model/extensions with the user's current desktop config.
|
|
// If extensions changed we must tear down the old agent so stale
|
|
// extension processes don't linger.
|
|
let extensions_changed = self.sync_session_config(&session).await?;
|
|
if extensions_changed {
|
|
let _ = self.agent_manager.remove_session(session_id).await;
|
|
}
|
|
|
|
let agent = self
|
|
.agent_manager
|
|
.get_or_create_agent(session_id.to_string())
|
|
.await?;
|
|
|
|
// Re-read the session after sync so restore picks up the new values.
|
|
let session = self
|
|
.agent_manager
|
|
.session_manager()
|
|
.get_session(session_id, false)
|
|
.await?;
|
|
|
|
// Ensure provider is configured (handles first use and LRU eviction).
|
|
if let Err(e) = agent.restore_provider_from_session(&session).await {
|
|
self.gateway
|
|
.send_message(
|
|
&message.user,
|
|
OutgoingMessage::Text {
|
|
body: format!("⚠️ Failed to configure provider: {e}"),
|
|
},
|
|
)
|
|
.await?;
|
|
return Ok(());
|
|
}
|
|
|
|
// Load extensions (skips any already loaded on the agent).
|
|
agent.load_extensions_from_session(&session).await;
|
|
|
|
let cancel = CancellationToken::new();
|
|
let user_message = Message::user().with_text(&message.text);
|
|
|
|
// Cap tool-calling loops so the agent doesn't run away doing
|
|
// dozens of tool calls before responding. After this many
|
|
// LLM→tool round-trips the agent will stop and reply with
|
|
// whatever it has.
|
|
const GATEWAY_MAX_TURNS: u32 = 5;
|
|
|
|
let session_config = SessionConfig {
|
|
id: session_id.to_string(),
|
|
schedule_id: None,
|
|
max_turns: Some(GATEWAY_MAX_TURNS),
|
|
retry_config: None,
|
|
};
|
|
|
|
let mut stream = match agent
|
|
.reply(user_message, session_config, Some(cancel))
|
|
.await
|
|
{
|
|
Ok(s) => s,
|
|
Err(e) => {
|
|
self.gateway
|
|
.send_message(
|
|
&message.user,
|
|
OutgoingMessage::Text {
|
|
body: format!("⚠️ Failed to start agent: {e}"),
|
|
},
|
|
)
|
|
.await?;
|
|
return Ok(());
|
|
}
|
|
};
|
|
|
|
// Telegram stops showing "typing…" after ~5 seconds. Re-send the
|
|
// indicator every 4 s so the user always sees activity while the
|
|
// agent is working (tool calls, LLM round-trips, etc.).
|
|
let typing_cancel = CancellationToken::new();
|
|
let typing_gateway = self.gateway.clone();
|
|
let typing_user = message.user.clone();
|
|
let typing_handle = tokio::spawn({
|
|
let cancel = typing_cancel.clone();
|
|
async move {
|
|
let mut interval = tokio::time::interval(Duration::from_secs(4));
|
|
interval.tick().await; // first tick is immediate, skip it
|
|
loop {
|
|
tokio::select! {
|
|
_ = cancel.cancelled() => break,
|
|
_ = interval.tick() => {
|
|
if let Err(e) = typing_gateway
|
|
.send_message(&typing_user, OutgoingMessage::Typing)
|
|
.await
|
|
{
|
|
tracing::debug!(error = %e, "failed to re-send typing indicator");
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
});
|
|
|
|
// Buffer text within a single assistant message so we send one
|
|
// Telegram message per LLM turn rather than per-chunk. When a
|
|
// ToolRequest appears in the same message we flush the buffer
|
|
// first — the user sees "Let me check…" immediately, then the
|
|
// typing indicator while the tool runs, then the next response.
|
|
let mut pending_text = String::new();
|
|
let mut sent_any = false;
|
|
let mut event_count: u64 = 0;
|
|
|
|
while let Some(event) = stream.next().await {
|
|
event_count += 1;
|
|
match event {
|
|
Ok(AgentEvent::Message(ref msg)) => {
|
|
tracing::debug!(
|
|
session_id,
|
|
role = ?msg.role,
|
|
content_items = msg.content.len(),
|
|
"gateway stream: message event #{event_count}"
|
|
);
|
|
if msg.role == rmcp::model::Role::Assistant {
|
|
for content in &msg.content {
|
|
match content {
|
|
MessageContent::Text(t) => {
|
|
if !t.text.is_empty() {
|
|
pending_text.push_str(&t.text);
|
|
}
|
|
}
|
|
MessageContent::ToolRequest(req) => {
|
|
// Flush any accumulated text before
|
|
// the tool runs — the user sees the
|
|
// assistant's intent immediately.
|
|
if !pending_text.is_empty() {
|
|
let _ = self
|
|
.gateway
|
|
.send_message(
|
|
&message.user,
|
|
OutgoingMessage::Text {
|
|
body: std::mem::take(&mut pending_text),
|
|
},
|
|
)
|
|
.await;
|
|
sent_any = true;
|
|
}
|
|
if let Ok(call) = &req.tool_call {
|
|
tracing::debug!(
|
|
session_id,
|
|
tool = %call.name,
|
|
"gateway stream: tool request"
|
|
);
|
|
let _ = self
|
|
.gateway
|
|
.send_message(&message.user, OutgoingMessage::Typing)
|
|
.await;
|
|
}
|
|
}
|
|
MessageContent::ToolResponse(resp) => {
|
|
tracing::debug!(
|
|
session_id,
|
|
id = %resp.id,
|
|
success = resp.tool_result.is_ok(),
|
|
"gateway stream: tool response"
|
|
);
|
|
}
|
|
_ => {}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
Ok(AgentEvent::McpNotification(_)) => {
|
|
tracing::debug!(
|
|
session_id,
|
|
"gateway stream: mcp notification #{event_count}"
|
|
);
|
|
}
|
|
Ok(AgentEvent::ModelChange {
|
|
ref model,
|
|
ref mode,
|
|
}) => {
|
|
tracing::debug!(
|
|
session_id,
|
|
model,
|
|
mode,
|
|
"gateway stream: model change #{event_count}"
|
|
);
|
|
}
|
|
Ok(AgentEvent::HistoryReplaced(_)) => {
|
|
tracing::debug!(
|
|
session_id,
|
|
"gateway stream: history replaced #{event_count}"
|
|
);
|
|
}
|
|
Err(e) => {
|
|
tracing::error!(session_id, error = %e, "gateway stream: error at event #{event_count}");
|
|
// Stop typing indicator before sending error.
|
|
typing_cancel.cancel();
|
|
let _ = typing_handle.await;
|
|
self.gateway
|
|
.send_message(
|
|
&message.user,
|
|
OutgoingMessage::Text {
|
|
body: format!("⚠️ Agent error: {e}"),
|
|
},
|
|
)
|
|
.await?;
|
|
return Ok(());
|
|
}
|
|
}
|
|
}
|
|
|
|
// Stream finished — stop the typing indicator.
|
|
typing_cancel.cancel();
|
|
let _ = typing_handle.await;
|
|
|
|
tracing::debug!(
|
|
session_id,
|
|
event_count,
|
|
pending_text_len = pending_text.len(),
|
|
sent_any,
|
|
"gateway stream: finished"
|
|
);
|
|
|
|
// Send any remaining buffered text (this is typically the final
|
|
// assistant response after the last tool round-trip).
|
|
if !pending_text.is_empty() {
|
|
self.gateway
|
|
.send_message(&message.user, OutgoingMessage::Text { body: pending_text })
|
|
.await?;
|
|
} else if !sent_any {
|
|
// Nothing was ever sent — let the user know.
|
|
self.gateway
|
|
.send_message(
|
|
&message.user,
|
|
OutgoingMessage::Text {
|
|
body: "(No response)".to_string(),
|
|
},
|
|
)
|
|
.await?;
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
}
|
|
|
|
fn gateway_working_dir(platform: &str, user_id: &str) -> PathBuf {
|
|
Paths::config_dir()
|
|
.join("gateway")
|
|
.join(platform)
|
|
.join(user_id)
|
|
}
|