mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-31 08:39:24 +00:00
Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9cb64dd8a7 | ||
|
|
6198c98673 | ||
|
|
dd0a0e10ab | ||
|
|
1d5777824c | ||
|
|
adf4e25c8f |
@@ -122,18 +122,32 @@ impl RelayClient {
|
|||||||
/// instance_url in chat-api. IronClaw only passes an optional CSRF nonce
|
/// instance_url in chat-api. IronClaw only passes an optional CSRF nonce
|
||||||
/// for validating the callback — no URLs.
|
/// for validating the callback — no URLs.
|
||||||
pub async fn initiate_oauth(&self, state_nonce: Option<&str>) -> Result<String, RelayError> {
|
pub async fn initiate_oauth(&self, state_nonce: Option<&str>) -> Result<String, RelayError> {
|
||||||
|
let url = format!("{}/oauth/slack/auth", self.base_url);
|
||||||
|
tracing::debug!(relay_url = %url, "RelayClient::initiate_oauth: sending request");
|
||||||
let mut query: Vec<(&str, &str)> = vec![];
|
let mut query: Vec<(&str, &str)> = vec![];
|
||||||
if let Some(nonce) = state_nonce {
|
if let Some(nonce) = state_nonce {
|
||||||
query.push(("state_nonce", nonce));
|
query.push(("state_nonce", nonce));
|
||||||
}
|
}
|
||||||
let resp = self
|
let resp = self
|
||||||
.http
|
.http
|
||||||
.get(format!("{}/oauth/slack/auth", self.base_url))
|
.get(&url)
|
||||||
.bearer_auth(self.api_key.expose_secret())
|
.bearer_auth(self.api_key.expose_secret())
|
||||||
.query(&query)
|
.query(&query)
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
.map_err(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
relay_url = %url,
|
||||||
|
error = %e,
|
||||||
|
"RelayClient::initiate_oauth: network request failed"
|
||||||
|
);
|
||||||
|
RelayError::Network(e.to_string())
|
||||||
|
})?;
|
||||||
|
tracing::debug!(
|
||||||
|
relay_url = %url,
|
||||||
|
status = %resp.status(),
|
||||||
|
"RelayClient::initiate_oauth: received response"
|
||||||
|
);
|
||||||
|
|
||||||
let status = resp.status();
|
let status = resp.status();
|
||||||
if status.is_redirection() {
|
if status.is_redirection() {
|
||||||
@@ -224,20 +238,39 @@ impl RelayClient {
|
|||||||
method: &str,
|
method: &str,
|
||||||
body: serde_json::Value,
|
body: serde_json::Value,
|
||||||
) -> Result<serde_json::Value, RelayError> {
|
) -> Result<serde_json::Value, RelayError> {
|
||||||
|
let url = format!("{}/proxy/{}/{}", self.base_url, provider, method);
|
||||||
|
tracing::debug!(
|
||||||
|
relay_url = %url,
|
||||||
|
provider = %provider,
|
||||||
|
method = %method,
|
||||||
|
"RelayClient::proxy_provider: sending request"
|
||||||
|
);
|
||||||
let query: Vec<(&str, &str)> = vec![("team_id", team_id)];
|
let query: Vec<(&str, &str)> = vec![("team_id", team_id)];
|
||||||
let resp = self
|
let resp = self
|
||||||
.http
|
.http
|
||||||
.post(format!("{}/proxy/{}/{}", self.base_url, provider, method))
|
.post(&url)
|
||||||
.bearer_auth(self.api_key.expose_secret())
|
.bearer_auth(self.api_key.expose_secret())
|
||||||
.query(&query)
|
.query(&query)
|
||||||
.json(&body)
|
.json(&body)
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
.map_err(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
relay_url = %url,
|
||||||
|
error = %e,
|
||||||
|
"RelayClient::proxy_provider: network request failed"
|
||||||
|
);
|
||||||
|
RelayError::Network(e.to_string())
|
||||||
|
})?;
|
||||||
|
|
||||||
if !resp.status().is_success() {
|
if !resp.status().is_success() {
|
||||||
let status = resp.status().as_u16();
|
let status = resp.status().as_u16();
|
||||||
let body = resp.text().await.unwrap_or_default();
|
let body = resp.text().await.unwrap_or_default();
|
||||||
|
tracing::warn!(
|
||||||
|
relay_url = %url,
|
||||||
|
status = status,
|
||||||
|
"RelayClient::proxy_provider: channel-relay returned error"
|
||||||
|
);
|
||||||
return Err(RelayError::Api {
|
return Err(RelayError::Api {
|
||||||
status,
|
status,
|
||||||
message: body,
|
message: body,
|
||||||
@@ -255,23 +288,45 @@ impl RelayClient {
|
|||||||
/// 32-byte secret. Called once at activation time; the result is cached in the
|
/// 32-byte secret. Called once at activation time; the result is cached in the
|
||||||
/// extension manager so subsequent calls to `relay_signing_secret()` use it.
|
/// extension manager so subsequent calls to `relay_signing_secret()` use it.
|
||||||
pub async fn get_signing_secret(&self, team_id: &str) -> Result<Vec<u8>, RelayError> {
|
pub async fn get_signing_secret(&self, team_id: &str) -> Result<Vec<u8>, RelayError> {
|
||||||
|
let url = format!("{}/relay/signing-secret", self.base_url);
|
||||||
|
tracing::debug!(
|
||||||
|
relay_url = %url,
|
||||||
|
"RelayClient::get_signing_secret: fetching signing secret"
|
||||||
|
);
|
||||||
let resp = self
|
let resp = self
|
||||||
.http
|
.http
|
||||||
.get(format!("{}/relay/signing-secret", self.base_url))
|
.get(&url)
|
||||||
.bearer_auth(self.api_key.expose_secret())
|
.bearer_auth(self.api_key.expose_secret())
|
||||||
.query(&[("team_id", team_id)])
|
.query(&[("team_id", team_id)])
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
.map_err(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
relay_url = %url,
|
||||||
|
error = %e,
|
||||||
|
"RelayClient::get_signing_secret: network request failed"
|
||||||
|
);
|
||||||
|
RelayError::Network(e.to_string())
|
||||||
|
})?;
|
||||||
|
|
||||||
if !resp.status().is_success() {
|
if !resp.status().is_success() {
|
||||||
let status = resp.status().as_u16();
|
let status = resp.status().as_u16();
|
||||||
let body = resp.text().await.unwrap_or_default();
|
let body = resp.text().await.unwrap_or_default();
|
||||||
|
tracing::warn!(
|
||||||
|
relay_url = %url,
|
||||||
|
status = status,
|
||||||
|
body = %body,
|
||||||
|
"RelayClient::get_signing_secret: channel-relay returned error"
|
||||||
|
);
|
||||||
return Err(RelayError::Api {
|
return Err(RelayError::Api {
|
||||||
status,
|
status,
|
||||||
message: body,
|
message: body,
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
tracing::debug!(
|
||||||
|
relay_url = %url,
|
||||||
|
"RelayClient::get_signing_secret: received successful response"
|
||||||
|
);
|
||||||
|
|
||||||
let body: serde_json::Value = resp
|
let body: serde_json::Value = resp
|
||||||
.json()
|
.json()
|
||||||
|
|||||||
@@ -1177,11 +1177,31 @@ async fn slack_relay_oauth_callback_handler(
|
|||||||
|
|
||||||
// Store team_id in settings
|
// Store team_id in settings
|
||||||
let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME);
|
let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME);
|
||||||
let _ = store
|
tracing::info!(
|
||||||
|
relay = DEFAULT_RELAY_NAME,
|
||||||
|
owner_id = %state.owner_id,
|
||||||
|
team_id_key = %team_id_key,
|
||||||
|
"relay OAuth callback: storing team_id in settings"
|
||||||
|
);
|
||||||
|
store
|
||||||
.set_setting(&state.owner_id, &team_id_key, &serde_json::json!(team_id))
|
.set_setting(&state.owner_id, &team_id_key, &serde_json::json!(team_id))
|
||||||
.await;
|
.await
|
||||||
|
.map_err(|e| {
|
||||||
|
tracing::error!(
|
||||||
|
relay = DEFAULT_RELAY_NAME,
|
||||||
|
owner_id = %state.owner_id,
|
||||||
|
error = %e,
|
||||||
|
"relay OAuth callback: failed to persist team_id to settings store"
|
||||||
|
);
|
||||||
|
format!("Failed to persist relay team_id: {e}")
|
||||||
|
})?;
|
||||||
|
|
||||||
// Activate the relay channel
|
// Activate the relay channel
|
||||||
|
tracing::info!(
|
||||||
|
relay = DEFAULT_RELAY_NAME,
|
||||||
|
owner_id = %state.owner_id,
|
||||||
|
"relay OAuth callback: activating relay channel"
|
||||||
|
);
|
||||||
ext_mgr
|
ext_mgr
|
||||||
.activate_stored_relay(DEFAULT_RELAY_NAME, &state.owner_id)
|
.activate_stored_relay(DEFAULT_RELAY_NAME, &state.owner_id)
|
||||||
.await
|
.await
|
||||||
@@ -2181,6 +2201,11 @@ async fn extensions_activate_handler(
|
|||||||
AuthenticatedUser(user): AuthenticatedUser,
|
AuthenticatedUser(user): AuthenticatedUser,
|
||||||
Path(name): Path<String>,
|
Path(name): Path<String>,
|
||||||
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
user_id = %user.user_id,
|
||||||
|
"extensions_activate_handler: received activate request"
|
||||||
|
);
|
||||||
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
||||||
StatusCode::NOT_IMPLEMENTED,
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
"Extension manager not available (secrets store required)".to_string(),
|
"Extension manager not available (secrets store required)".to_string(),
|
||||||
@@ -2188,6 +2213,10 @@ async fn extensions_activate_handler(
|
|||||||
|
|
||||||
match ext_mgr.activate(&name, &user.user_id).await {
|
match ext_mgr.activate(&name, &user.user_id).await {
|
||||||
Ok(result) => {
|
Ok(result) => {
|
||||||
|
tracing::info!(
|
||||||
|
extension = %name,
|
||||||
|
"extensions_activate_handler: activation succeeded"
|
||||||
|
);
|
||||||
// Activation loaded the WASM module. Check if the tool needs
|
// Activation loaded the WASM module. Check if the tool needs
|
||||||
// OAuth scope expansion (e.g., adding google-docs when gmail
|
// OAuth scope expansion (e.g., adding google-docs when gmail
|
||||||
// already has a token but missing the documents scope).
|
// already has a token but missing the documents scope).
|
||||||
@@ -2206,6 +2235,13 @@ async fn extensions_activate_handler(
|
|||||||
crate::extensions::ExtensionError::AuthRequired
|
crate::extensions::ExtensionError::AuthRequired
|
||||||
);
|
);
|
||||||
|
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
error = %activate_err,
|
||||||
|
needs_auth = needs_auth,
|
||||||
|
"extensions_activate_handler: activation failed, attempting auth fallback"
|
||||||
|
);
|
||||||
|
|
||||||
if !needs_auth {
|
if !needs_auth {
|
||||||
return Ok(Json(ActionResponse::fail(activate_err.to_string())));
|
return Ok(Json(ActionResponse::fail(activate_err.to_string())));
|
||||||
}
|
}
|
||||||
@@ -2213,10 +2249,21 @@ async fn extensions_activate_handler(
|
|||||||
// Activation failed due to auth; try authenticating first.
|
// Activation failed due to auth; try authenticating first.
|
||||||
match ext_mgr.auth(&name, &user.user_id).await {
|
match ext_mgr.auth(&name, &user.user_id).await {
|
||||||
Ok(auth_result) if auth_result.is_authenticated() => {
|
Ok(auth_result) if auth_result.is_authenticated() => {
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
"extensions_activate_handler: auth reports authenticated, retrying activate"
|
||||||
|
);
|
||||||
// Auth succeeded, retry activation.
|
// Auth succeeded, retry activation.
|
||||||
match ext_mgr.activate(&name, &user.user_id).await {
|
match ext_mgr.activate(&name, &user.user_id).await {
|
||||||
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
||||||
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
error = %e,
|
||||||
|
"extensions_activate_handler: retry after auth still failed"
|
||||||
|
);
|
||||||
|
Ok(Json(ActionResponse::fail(e.to_string())))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Ok(auth_result) => {
|
Ok(auth_result) => {
|
||||||
|
|||||||
+2
-1
@@ -127,7 +127,8 @@ async fn list_settings(
|
|||||||
}
|
}
|
||||||
|
|
||||||
let display_value = if value.len() > 60 {
|
let display_value = if value.len() > 60 {
|
||||||
format!("{}...", &value[..57])
|
let end = crate::util::floor_char_boundary(&value, 57);
|
||||||
|
format!("{}...", &value[..end])
|
||||||
} else {
|
} else {
|
||||||
value
|
value
|
||||||
};
|
};
|
||||||
|
|||||||
+15
-1
@@ -256,7 +256,8 @@ fn truncate_content(s: &str, max_len: usize) -> String {
|
|||||||
if s.len() <= max_len {
|
if s.len() <= max_len {
|
||||||
s.to_string()
|
s.to_string()
|
||||||
} else {
|
} else {
|
||||||
format!("{}...", &s[..max_len])
|
let end = crate::util::floor_char_boundary(s, max_len);
|
||||||
|
format!("{}...", &s[..end])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -292,4 +293,17 @@ mod tests {
|
|||||||
assert_eq!(truncate_content("hello", 10), "hello");
|
assert_eq!(truncate_content("hello", 10), "hello");
|
||||||
assert_eq!(truncate_content("hello world", 5), "hello...");
|
assert_eq!(truncate_content("hello world", 5), "hello...");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_truncate_content_multibyte_does_not_panic() {
|
||||||
|
// \u{00e9} is precomposed 'é' (2 bytes in UTF-8)
|
||||||
|
let s = "caf\u{00e9} au lait"; // "café au lait", é starts at byte 3
|
||||||
|
let result = truncate_content(s, 4); // byte 4 is inside 2-byte é
|
||||||
|
assert_eq!(result, "caf...");
|
||||||
|
|
||||||
|
// 4-byte emoji: slicing mid-emoji must not panic
|
||||||
|
let emoji = "Hi \u{1F600} there"; // 😀 is 4 bytes, starts at byte 3
|
||||||
|
let result = truncate_content(emoji, 4); // byte 4 is inside 😀
|
||||||
|
assert_eq!(result, "Hi ...");
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -192,6 +192,9 @@ pub struct JobContext {
|
|||||||
/// but subsequent tools (e.g., `json`) may need the full output. This
|
/// but subsequent tools (e.g., `json`) may need the full output. This
|
||||||
/// stash stores the complete, unsanitized output so tools can reference
|
/// stash stores the complete, unsanitized output so tools can reference
|
||||||
/// previous results by ID via `$tool_call_id` parameter syntax.
|
/// previous results by ID via `$tool_call_id` parameter syntax.
|
||||||
|
///
|
||||||
|
/// Also used for cross-tool implicit state (keys prefixed with `__`) such
|
||||||
|
/// as `__routine_last_name` for fallback recovery in routine tool chains.
|
||||||
#[serde(skip)]
|
#[serde(skip)]
|
||||||
pub tool_output_stash: Arc<tokio::sync::RwLock<HashMap<String, String>>>,
|
pub tool_output_stash: Arc<tokio::sync::RwLock<HashMap<String, String>>>,
|
||||||
/// User's preferred timezone (IANA name, e.g. "America/New_York"). Defaults to "UTC".
|
/// User's preferred timezone (IANA name, e.g. "America/New_York"). Defaults to "UTC".
|
||||||
|
|||||||
+389
-30
@@ -659,6 +659,66 @@ impl ExtensionManager {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Resolve the relay URL override for an extension from settings.
|
||||||
|
///
|
||||||
|
/// Returns `Some(url)` if a non-empty per-extension `relay_url` override is
|
||||||
|
/// set for the given extension; otherwise returns `None` and callers should
|
||||||
|
/// fall back to the env-level `RelayConfig`.
|
||||||
|
///
|
||||||
|
/// Uses `self.user_id` (owner scope) for consistency with `configure()`,
|
||||||
|
/// which also writes setting_path fields under the owner scope.
|
||||||
|
///
|
||||||
|
/// The override is validated: only `http` / `https` schemes are accepted
|
||||||
|
/// and the URL must not contain userinfo (embedded credentials). This
|
||||||
|
/// prevents a malicious override from exfiltrating the instance-wide relay
|
||||||
|
/// API key to an attacker-controlled host.
|
||||||
|
async fn effective_relay_url(&self, name: &str) -> Option<String> {
|
||||||
|
if let Some(ref store) = self.store {
|
||||||
|
let key = format!("extensions.{name}.relay_url");
|
||||||
|
if let Ok(Some(v)) = store.get_setting(&self.user_id, &key).await {
|
||||||
|
let url = v
|
||||||
|
.as_str()
|
||||||
|
.map(|s| s.trim().to_string())
|
||||||
|
.filter(|s| !s.is_empty());
|
||||||
|
if let Some(ref u) = url {
|
||||||
|
// Validate the override to prevent API-key exfiltration:
|
||||||
|
// only allow http(s) with no embedded credentials.
|
||||||
|
match url::Url::parse(u) {
|
||||||
|
Ok(parsed)
|
||||||
|
if (parsed.scheme() == "http" || parsed.scheme() == "https")
|
||||||
|
&& parsed.username().is_empty()
|
||||||
|
&& parsed.password().is_none() =>
|
||||||
|
{
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
relay_url_host = %parsed.host_str().unwrap_or("unknown"),
|
||||||
|
"effective_relay_url: using per-extension override from settings"
|
||||||
|
);
|
||||||
|
return url;
|
||||||
|
}
|
||||||
|
Ok(parsed) => {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
scheme = %parsed.scheme(),
|
||||||
|
has_userinfo = !parsed.username().is_empty() || parsed.password().is_some(),
|
||||||
|
"effective_relay_url: rejecting override — \
|
||||||
|
only http/https without embedded credentials is allowed"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
error = %e,
|
||||||
|
"effective_relay_url: rejecting override — invalid URL"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
None
|
||||||
|
}
|
||||||
|
|
||||||
/// Get the shared relay event sender for the webhook endpoint.
|
/// Get the shared relay event sender for the webhook endpoint.
|
||||||
pub fn relay_event_tx(
|
pub fn relay_event_tx(
|
||||||
&self,
|
&self,
|
||||||
@@ -892,6 +952,46 @@ impl ExtensionManager {
|
|||||||
false
|
false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Check whether a stored `team_id` setting exists for the given relay extension.
|
||||||
|
///
|
||||||
|
/// Unlike [`is_relay_channel`], this does **not** consult the in-memory
|
||||||
|
/// `installed_relay_extensions` set — it only looks at the persistent settings
|
||||||
|
/// store. This distinction matters for `auth_channel_relay`: an extension can
|
||||||
|
/// be *installed* (present in the in-memory set) but not yet *authenticated*
|
||||||
|
/// (no OAuth completed, no team_id stored).
|
||||||
|
async fn has_stored_team_id(&self, name: &str, _user_id: &str) -> bool {
|
||||||
|
if let Some(ref store) = self.store {
|
||||||
|
let key = format!("relay:{}:team_id", name);
|
||||||
|
// Use owner scope (self.user_id) for consistency: the OAuth callback
|
||||||
|
// stores team_id under state.owner_id which maps to self.user_id.
|
||||||
|
match store.get_setting(&self.user_id, &key).await {
|
||||||
|
Ok(Some(v)) => {
|
||||||
|
let has_id = v.as_str().is_some_and(|s| !s.is_empty());
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
has_team_id = has_id,
|
||||||
|
"has_stored_team_id: checked store"
|
||||||
|
);
|
||||||
|
return has_id;
|
||||||
|
}
|
||||||
|
Ok(None) => {
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
"has_stored_team_id: no team_id setting found"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
error = %e,
|
||||||
|
"has_stored_team_id: failed to read from settings store"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
false
|
||||||
|
}
|
||||||
|
|
||||||
/// Restore persisted relay channels after startup.
|
/// Restore persisted relay channels after startup.
|
||||||
///
|
///
|
||||||
/// Loads the persisted active channel list, filters to relay types (those with
|
/// Loads the persisted active channel list, filters to relay types (those with
|
||||||
@@ -1418,7 +1518,7 @@ impl ExtensionManager {
|
|||||||
let errors = self.activation_errors.read().await;
|
let errors = self.activation_errors.read().await;
|
||||||
for name in installed.iter() {
|
for name in installed.iter() {
|
||||||
let active = active_names.contains(name);
|
let active = active_names.contains(name);
|
||||||
let authenticated = self.is_relay_channel(name, user_id).await;
|
let authenticated = self.has_stored_team_id(name, user_id).await;
|
||||||
let activation_error = errors.get(name).cloned();
|
let activation_error = errors.get(name).cloned();
|
||||||
let registry_entry = self
|
let registry_entry = self
|
||||||
.registry
|
.registry
|
||||||
@@ -4191,20 +4291,69 @@ impl ExtensionManager {
|
|||||||
name: &str,
|
name: &str,
|
||||||
user_id: &str,
|
user_id: &str,
|
||||||
) -> Result<AuthResult, ExtensionError> {
|
) -> Result<AuthResult, ExtensionError> {
|
||||||
// Check if already authenticated (team_id setting exists)
|
tracing::debug!(
|
||||||
if self.is_relay_channel(name, user_id).await {
|
extension = %name,
|
||||||
|
user_id = %user_id,
|
||||||
|
"auth_channel_relay: starting"
|
||||||
|
);
|
||||||
|
|
||||||
|
// Check if already authenticated by looking for a stored team_id.
|
||||||
|
// We intentionally skip the `installed_relay_extensions` in-memory set
|
||||||
|
// here because that set only tracks *installed* extensions — an extension
|
||||||
|
// can be installed (via registry) but not yet authenticated (no OAuth
|
||||||
|
// completed). Checking just `is_relay_channel()` would short-circuit
|
||||||
|
// to "authenticated" even when no team_id exists, preventing the OAuth
|
||||||
|
// flow from being offered to the user.
|
||||||
|
if self.has_stored_team_id(name, user_id).await {
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
"auth_channel_relay: already authenticated (team_id in store)"
|
||||||
|
);
|
||||||
return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay));
|
return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
"auth_channel_relay: no stored team_id, initiating OAuth"
|
||||||
|
);
|
||||||
|
|
||||||
// Use relay config captured at startup
|
// Use relay config captured at startup
|
||||||
let relay_config = self.relay_config()?;
|
let relay_config = self.relay_config().map_err(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
error = %e,
|
||||||
|
"auth_channel_relay: relay config not available — \
|
||||||
|
CHANNEL_RELAY_URL and CHANNEL_RELAY_API_KEY must be set"
|
||||||
|
);
|
||||||
|
e
|
||||||
|
})?;
|
||||||
|
|
||||||
|
// Allow per-extension URL override from settings
|
||||||
|
let effective_url = self
|
||||||
|
.effective_relay_url(name)
|
||||||
|
.await
|
||||||
|
.unwrap_or_else(|| relay_config.url.clone());
|
||||||
|
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
relay_url = %effective_url,
|
||||||
|
"auth_channel_relay: creating relay client for OAuth"
|
||||||
|
);
|
||||||
|
|
||||||
let client = crate::channels::relay::RelayClient::new(
|
let client = crate::channels::relay::RelayClient::new(
|
||||||
relay_config.url.clone(),
|
effective_url.clone(),
|
||||||
relay_config.api_key.clone(),
|
relay_config.api_key.clone(),
|
||||||
relay_config.request_timeout_secs,
|
relay_config.request_timeout_secs,
|
||||||
)
|
)
|
||||||
.map_err(|e| ExtensionError::Config(e.to_string()))?;
|
.map_err(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
relay_url = %effective_url,
|
||||||
|
error = %e,
|
||||||
|
"auth_channel_relay: failed to create relay HTTP client"
|
||||||
|
);
|
||||||
|
ExtensionError::Config(e.to_string())
|
||||||
|
})?;
|
||||||
|
|
||||||
// Generate CSRF nonce — IronClaw validates this on the callback to ensure
|
// Generate CSRF nonce — IronClaw validates this on the callback to ensure
|
||||||
// the OAuth completion is legitimate. Channel-relay embeds it in the signed
|
// the OAuth completion is legitimate. Channel-relay embeds it in the signed
|
||||||
@@ -4216,18 +4365,44 @@ impl ExtensionManager {
|
|||||||
self.secrets
|
self.secrets
|
||||||
.create(user_id, CreateSecretParams::new(&state_key, &state_nonce))
|
.create(user_id, CreateSecretParams::new(&state_key, &state_nonce))
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}")))?;
|
.map_err(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
error = %e,
|
||||||
|
"auth_channel_relay: failed to store OAuth state nonce"
|
||||||
|
);
|
||||||
|
ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}"))
|
||||||
|
})?;
|
||||||
|
|
||||||
// Channel-relay derives all URLs from trusted instance_url in chat-api.
|
// Channel-relay derives all URLs from trusted instance_url in chat-api.
|
||||||
// We only pass the nonce for CSRF validation on the callback.
|
// We only pass the nonce for CSRF validation on the callback.
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
relay_url = %effective_url,
|
||||||
|
"auth_channel_relay: calling initiate_oauth on channel-relay"
|
||||||
|
);
|
||||||
match client.initiate_oauth(Some(&state_nonce)).await {
|
match client.initiate_oauth(Some(&state_nonce)).await {
|
||||||
Ok(auth_url) => Ok(AuthResult::awaiting_authorization(
|
Ok(auth_url) => {
|
||||||
name,
|
tracing::info!(
|
||||||
ExtensionKind::ChannelRelay,
|
extension = %name,
|
||||||
auth_url,
|
"auth_channel_relay: OAuth URL obtained, awaiting user authorization"
|
||||||
"redirect".to_string(),
|
);
|
||||||
)),
|
Ok(AuthResult::awaiting_authorization(
|
||||||
Err(e) => Err(ExtensionError::AuthFailed(e.to_string())),
|
name,
|
||||||
|
ExtensionKind::ChannelRelay,
|
||||||
|
auth_url,
|
||||||
|
"redirect".to_string(),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
relay_url = %effective_url,
|
||||||
|
error = %e,
|
||||||
|
"auth_channel_relay: initiate_oauth call to channel-relay failed"
|
||||||
|
);
|
||||||
|
Err(ExtensionError::AuthFailed(e.to_string()))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -4237,40 +4412,112 @@ impl ExtensionManager {
|
|||||||
name: &str,
|
name: &str,
|
||||||
user_id: &str,
|
user_id: &str,
|
||||||
) -> Result<ActivateResult, ExtensionError> {
|
) -> Result<ActivateResult, ExtensionError> {
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
user_id = %user_id,
|
||||||
|
"activate_channel_relay: starting"
|
||||||
|
);
|
||||||
|
|
||||||
let team_id_key = format!("relay:{}:team_id", name);
|
let team_id_key = format!("relay:{}:team_id", name);
|
||||||
|
|
||||||
// Get team_id from settings (stored by the OAuth callback)
|
// Get team_id from settings (stored by the OAuth callback)
|
||||||
let team_id = if let Some(ref store) = self.store {
|
let team_id = if let Some(ref store) = self.store {
|
||||||
store
|
match store.get_setting(user_id, &team_id_key).await {
|
||||||
.get_setting(user_id, &team_id_key)
|
Ok(Some(v)) => {
|
||||||
.await
|
let id = v.as_str().map(|s| s.to_string()).unwrap_or_default();
|
||||||
.ok()
|
tracing::debug!(
|
||||||
.flatten()
|
extension = %name,
|
||||||
.and_then(|v| v.as_str().map(|s| s.to_string()))
|
team_id_empty = id.is_empty(),
|
||||||
.unwrap_or_default()
|
"activate_channel_relay: loaded team_id from store"
|
||||||
|
);
|
||||||
|
id
|
||||||
|
}
|
||||||
|
Ok(None) => {
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
setting_key = %team_id_key,
|
||||||
|
"activate_channel_relay: no team_id in settings store"
|
||||||
|
);
|
||||||
|
String::new()
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
error = %e,
|
||||||
|
"activate_channel_relay: failed to read team_id from settings store"
|
||||||
|
);
|
||||||
|
String::new()
|
||||||
|
}
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
"activate_channel_relay: no settings store available"
|
||||||
|
);
|
||||||
String::new()
|
String::new()
|
||||||
};
|
};
|
||||||
|
|
||||||
if team_id.is_empty() {
|
if team_id.is_empty() {
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
"activate_channel_relay: team_id is empty, returning AuthRequired"
|
||||||
|
);
|
||||||
return Err(ExtensionError::AuthRequired);
|
return Err(ExtensionError::AuthRequired);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use relay config captured at startup
|
// Use relay config captured at startup
|
||||||
let relay_config = self.relay_config()?;
|
let relay_config = self.relay_config().map_err(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
error = %e,
|
||||||
|
"activate_channel_relay: relay config not available"
|
||||||
|
);
|
||||||
|
e
|
||||||
|
})?;
|
||||||
|
|
||||||
|
// Allow per-extension URL override from settings
|
||||||
|
let effective_url = self
|
||||||
|
.effective_relay_url(name)
|
||||||
|
.await
|
||||||
|
.unwrap_or_else(|| relay_config.url.clone());
|
||||||
|
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
relay_url = %effective_url,
|
||||||
|
"activate_channel_relay: relay config loaded"
|
||||||
|
);
|
||||||
|
|
||||||
let instance_id = self.relay_instance_id(relay_config, user_id);
|
let instance_id = self.relay_instance_id(relay_config, user_id);
|
||||||
|
|
||||||
let client = crate::channels::relay::RelayClient::new(
|
let client = crate::channels::relay::RelayClient::new(
|
||||||
relay_config.url.clone(),
|
effective_url.clone(),
|
||||||
relay_config.api_key.clone(),
|
relay_config.api_key.clone(),
|
||||||
relay_config.request_timeout_secs,
|
relay_config.request_timeout_secs,
|
||||||
)
|
)
|
||||||
.map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
|
.map_err(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
relay_url = %effective_url,
|
||||||
|
error = %e,
|
||||||
|
"activate_channel_relay: failed to create relay HTTP client"
|
||||||
|
);
|
||||||
|
ExtensionError::ActivationFailed(e.to_string())
|
||||||
|
})?;
|
||||||
|
|
||||||
// Fetch the per-instance signing secret from channel-relay.
|
// Fetch the per-instance signing secret from channel-relay.
|
||||||
// This must succeed — there is no fallback.
|
// This must succeed — there is no fallback.
|
||||||
|
tracing::debug!(
|
||||||
|
extension = %name,
|
||||||
|
relay_url = %effective_url,
|
||||||
|
"activate_channel_relay: fetching signing secret from channel-relay"
|
||||||
|
);
|
||||||
let signing_secret = client.get_signing_secret(&team_id).await.map_err(|e| {
|
let signing_secret = client.get_signing_secret(&team_id).await.map_err(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
relay_url = %effective_url,
|
||||||
|
error = %e,
|
||||||
|
"activate_channel_relay: failed to fetch signing secret from channel-relay"
|
||||||
|
);
|
||||||
ExtensionError::Config(format!("Failed to fetch relay signing secret: {e}"))
|
ExtensionError::Config(format!("Failed to fetch relay signing secret: {e}"))
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
@@ -4289,16 +4536,29 @@ impl ExtensionManager {
|
|||||||
// Hot-add to channel manager
|
// Hot-add to channel manager
|
||||||
let cm_guard = self.relay_channel_manager.read().await;
|
let cm_guard = self.relay_channel_manager.read().await;
|
||||||
let channel_mgr = cm_guard.as_ref().ok_or_else(|| {
|
let channel_mgr = cm_guard.as_ref().ok_or_else(|| {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
"activate_channel_relay: channel manager not initialized"
|
||||||
|
);
|
||||||
ExtensionError::ActivationFailed("Channel manager not initialized".to_string())
|
ExtensionError::ActivationFailed("Channel manager not initialized".to_string())
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
channel_mgr
|
channel_mgr.hot_add(Box::new(channel)).await.map_err(|e| {
|
||||||
.hot_add(Box::new(channel))
|
tracing::warn!(
|
||||||
.await
|
extension = %name,
|
||||||
.map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
|
error = %e,
|
||||||
|
"activate_channel_relay: hot_add to channel manager failed"
|
||||||
|
);
|
||||||
|
ExtensionError::ActivationFailed(e.to_string())
|
||||||
|
})?;
|
||||||
|
|
||||||
if let Ok(mut cache) = self.relay_signing_secret_cache.lock() {
|
if let Ok(mut cache) = self.relay_signing_secret_cache.lock() {
|
||||||
*cache = Some(signing_secret);
|
*cache = Some(signing_secret);
|
||||||
|
} else {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
"activate_channel_relay: failed to cache signing secret (mutex poisoned)"
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Store the event sender so the web gateway's relay webhook endpoint can push events
|
// Store the event sender so the web gateway's relay webhook endpoint can push events
|
||||||
@@ -4316,6 +4576,12 @@ impl ExtensionManager {
|
|||||||
self.broadcast_extension_status(name, "active", Some(&status_msg))
|
self.broadcast_extension_status(name, "active", Some(&status_msg))
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
|
tracing::info!(
|
||||||
|
extension = %name,
|
||||||
|
instance_id = %instance_id,
|
||||||
|
"activate_channel_relay: relay channel activated successfully"
|
||||||
|
);
|
||||||
|
|
||||||
Ok(ActivateResult {
|
Ok(ActivateResult {
|
||||||
name: name.to_string(),
|
name: name.to_string(),
|
||||||
kind: ExtensionKind::ChannelRelay,
|
kind: ExtensionKind::ChannelRelay,
|
||||||
@@ -4595,6 +4861,41 @@ impl ExtensionManager {
|
|||||||
}
|
}
|
||||||
Ok(ExtensionSetupSchema { secrets, fields })
|
Ok(ExtensionSetupSchema { secrets, fields })
|
||||||
}
|
}
|
||||||
|
ExtensionKind::ChannelRelay => {
|
||||||
|
let relay_url_key = format!("extensions.{name}.relay_url");
|
||||||
|
let current_url = if let Some(ref store) = self.store {
|
||||||
|
match store.get_setting(&self.user_id, &relay_url_key).await {
|
||||||
|
Ok(value_opt) => value_opt
|
||||||
|
.and_then(|v| v.as_str().map(|s| s.to_string()))
|
||||||
|
.filter(|s| !s.is_empty()),
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
extension = %name,
|
||||||
|
setting_key = %relay_url_key,
|
||||||
|
error = %e,
|
||||||
|
"get_setup_schema: failed to read relay_url from settings"
|
||||||
|
);
|
||||||
|
None
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let env_url = self.relay_config.as_ref().map(|c| c.url.as_str());
|
||||||
|
Ok(ExtensionSetupSchema {
|
||||||
|
secrets: Vec::new(),
|
||||||
|
fields: vec![crate::channels::web::types::SetupFieldInfo {
|
||||||
|
name: "relay_url".to_string(),
|
||||||
|
prompt: format!(
|
||||||
|
"Channel-relay service URL (leave empty to use env default{})",
|
||||||
|
env_url.map(|u| format!(": {u}")).unwrap_or_default()
|
||||||
|
),
|
||||||
|
optional: true,
|
||||||
|
provided: current_url.is_some(),
|
||||||
|
input_type: crate::tools::wasm::ToolSetupFieldInputType::Text,
|
||||||
|
}],
|
||||||
|
})
|
||||||
|
}
|
||||||
_ => Ok(ExtensionSetupSchema {
|
_ => Ok(ExtensionSetupSchema {
|
||||||
secrets: Vec::new(),
|
secrets: Vec::new(),
|
||||||
fields: Vec::new(),
|
fields: Vec::new(),
|
||||||
@@ -4997,7 +5298,17 @@ impl ExtensionManager {
|
|||||||
names.insert(server.token_secret_name());
|
names.insert(server.token_secret_name());
|
||||||
(names, Vec::new())
|
(names, Vec::new())
|
||||||
}
|
}
|
||||||
ExtensionKind::ChannelRelay => (std::collections::HashSet::new(), Vec::new()),
|
ExtensionKind::ChannelRelay => {
|
||||||
|
let relay_fields = vec![crate::tools::wasm::ToolFieldSetupSchema {
|
||||||
|
name: "relay_url".to_string(),
|
||||||
|
prompt: "Channel-relay service URL override".to_string(),
|
||||||
|
optional: true,
|
||||||
|
setting_path: Some(format!("extensions.{name}.relay_url")),
|
||||||
|
input_type: crate::tools::wasm::ToolSetupFieldInputType::Text,
|
||||||
|
restart_required: false,
|
||||||
|
}];
|
||||||
|
(std::collections::HashSet::new(), relay_fields)
|
||||||
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
let allowed_fields: std::collections::HashSet<String> =
|
let allowed_fields: std::collections::HashSet<String> =
|
||||||
@@ -5088,13 +5399,28 @@ impl ExtensionManager {
|
|||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
let trimmed = field_value.trim();
|
let trimmed = field_value.trim();
|
||||||
|
let field_def = setup_field_defs.get(field_name);
|
||||||
|
|
||||||
|
// Empty value on an optional field with a setting_path: clear the
|
||||||
|
// stored override so the system reverts to the env/default value.
|
||||||
if trimmed.is_empty() {
|
if trimmed.is_empty() {
|
||||||
|
if let Some(def) = field_def
|
||||||
|
&& def.optional
|
||||||
|
{
|
||||||
|
stored_fields.remove(field_name);
|
||||||
|
if let Some(setting_path) = &def.setting_path {
|
||||||
|
Self::validate_setup_setting_path(name, setting_path)?;
|
||||||
|
if let Some(store) = self.store.as_ref() {
|
||||||
|
let _ = store.delete_setting(&self.user_id, setting_path).await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
stored_fields.insert(field_name.clone(), trimmed.to_string());
|
stored_fields.insert(field_name.clone(), trimmed.to_string());
|
||||||
|
|
||||||
if let Some(field_def) = setup_field_defs.get(field_name) {
|
if let Some(field_def) = field_def {
|
||||||
if field_def.restart_required {
|
if field_def.restart_required {
|
||||||
restart_required = true;
|
restart_required = true;
|
||||||
}
|
}
|
||||||
@@ -7058,6 +7384,39 @@ mod tests {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Regression: installed-but-not-authenticated relay must NOT short-circuit
|
||||||
|
/// `auth_channel_relay()` to "authenticated". Previously, `auth_channel_relay`
|
||||||
|
/// called `is_relay_channel()` which checked the in-memory
|
||||||
|
/// `installed_relay_extensions` set; that returned `true` even when no team_id
|
||||||
|
/// existed in the store, so the OAuth URL was never offered.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_auth_channel_relay_installed_without_team_id_is_not_authenticated() {
|
||||||
|
let dir = tempfile::tempdir().expect("temp dir");
|
||||||
|
let mgr = make_test_manager(None, dir.path().to_path_buf());
|
||||||
|
|
||||||
|
// Mark as installed (simulates clicking Install in the UI)
|
||||||
|
mgr.installed_relay_extensions
|
||||||
|
.write()
|
||||||
|
.await
|
||||||
|
.insert("slack-relay".to_string());
|
||||||
|
|
||||||
|
// Without a stored team_id, auth should NOT return authenticated.
|
||||||
|
// It should fail because relay config is missing (no CHANNEL_RELAY_URL),
|
||||||
|
// but the key assertion is that it does NOT return Ok(authenticated).
|
||||||
|
let result = mgr.auth_channel_relay("slack-relay", "test").await;
|
||||||
|
match result {
|
||||||
|
Ok(ref auth_result) if auth_result.is_authenticated() => {
|
||||||
|
panic!(
|
||||||
|
"auth_channel_relay returned authenticated for installed-but-no-team-id relay; \
|
||||||
|
expected either an OAuth URL or a config error"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
// Config error (no relay URL) or awaiting_authorization — both are correct
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_remove_relay_shuts_down_via_relay_channel_manager() {
|
async fn test_remove_relay_shuts_down_via_relay_channel_manager() {
|
||||||
// Regression: remove() only checked channel_runtime for shutdown, missing
|
// Regression: remove() only checked channel_runtime for shutdown, missing
|
||||||
|
|||||||
@@ -451,7 +451,7 @@ impl NearAiChatProvider {
|
|||||||
provider: "nearai_chat".to_string(),
|
provider: "nearai_chat".to_string(),
|
||||||
reason: format!(
|
reason: format!(
|
||||||
"No model names found in response: {}",
|
"No model names found in response: {}",
|
||||||
&response_text[..response_text.len().min(300)]
|
&response_text[..crate::util::floor_char_boundary(&response_text, 300)]
|
||||||
),
|
),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -650,6 +650,23 @@ pub(crate) fn routine_update_parameters_schema() -> Value {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const ROUTINE_LAST_NAME_STASH_KEY: &str = "__routine_last_name";
|
||||||
|
|
||||||
|
async fn stash_last_routine_name(ctx: &JobContext, name: &str) {
|
||||||
|
ctx.tool_output_stash
|
||||||
|
.write()
|
||||||
|
.await
|
||||||
|
.insert(ROUTINE_LAST_NAME_STASH_KEY.to_string(), name.to_string());
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn restore_last_routine_name(ctx: &JobContext) -> Option<String> {
|
||||||
|
ctx.tool_output_stash
|
||||||
|
.read()
|
||||||
|
.await
|
||||||
|
.get(ROUTINE_LAST_NAME_STASH_KEY)
|
||||||
|
.cloned()
|
||||||
|
}
|
||||||
|
|
||||||
fn nested_object<'a>(params: &'a Value, field: &str) -> Option<&'a Map<String, Value>> {
|
fn nested_object<'a>(params: &'a Value, field: &str) -> Option<&'a Map<String, Value>> {
|
||||||
params.get(field).and_then(Value::as_object)
|
params.get(field).and_then(Value::as_object)
|
||||||
}
|
}
|
||||||
@@ -1093,6 +1110,7 @@ impl Tool for RoutineCreateTool {
|
|||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
let normalized = parse_routine_create_request(¶ms)?;
|
let normalized = parse_routine_create_request(¶ms)?;
|
||||||
|
stash_last_routine_name(ctx, &normalized.name).await;
|
||||||
let trigger = build_routine_trigger(&normalized.trigger);
|
let trigger = build_routine_trigger(&normalized.trigger);
|
||||||
let action =
|
let action =
|
||||||
build_routine_action(&normalized.name, &normalized.prompt, &normalized.execution);
|
build_routine_action(&normalized.name, &normalized.prompt, &normalized.execution);
|
||||||
@@ -1274,6 +1292,7 @@ impl Tool for RoutineUpdateTool {
|
|||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
let name = require_str(¶ms, "name")?;
|
let name = require_str(¶ms, "name")?;
|
||||||
|
stash_last_routine_name(ctx, name).await;
|
||||||
|
|
||||||
let mut routine = self
|
let mut routine = self
|
||||||
.store
|
.store
|
||||||
@@ -1411,11 +1430,24 @@ impl Tool for RoutineDeleteTool {
|
|||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
let name = require_str(¶ms, "name")?;
|
let name = if let Some(name) = params.get("name").and_then(|v| v.as_str()) {
|
||||||
|
if name.trim().is_empty() {
|
||||||
|
return Err(ToolError::InvalidParameters(
|
||||||
|
"'name' parameter cannot be empty".to_string(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
name.to_string()
|
||||||
|
} else {
|
||||||
|
restore_last_routine_name(ctx).await.ok_or_else(|| {
|
||||||
|
ToolError::InvalidParameters(
|
||||||
|
"missing 'name' parameter and no previous routine target to infer".to_string(),
|
||||||
|
)
|
||||||
|
})?
|
||||||
|
};
|
||||||
|
|
||||||
let routine = self
|
let routine = self
|
||||||
.store
|
.store
|
||||||
.get_routine_by_name(&ctx.user_id, name)
|
.get_routine_by_name(&ctx.user_id, &name)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("DB error: {e}")))?
|
.map_err(|e| ToolError::ExecutionFailed(format!("DB error: {e}")))?
|
||||||
.ok_or_else(|| ToolError::ExecutionFailed(format!("routine '{}' not found", name)))?;
|
.ok_or_else(|| ToolError::ExecutionFailed(format!("routine '{}' not found", name)))?;
|
||||||
@@ -1430,7 +1462,7 @@ impl Tool for RoutineDeleteTool {
|
|||||||
self.engine.refresh_event_cache().await;
|
self.engine.refresh_event_cache().await;
|
||||||
|
|
||||||
let result = serde_json::json!({
|
let result = serde_json::json!({
|
||||||
"name": name,
|
"name": &name,
|
||||||
"deleted": deleted,
|
"deleted": deleted,
|
||||||
});
|
});
|
||||||
|
|
||||||
|
|||||||
+19
-1
@@ -117,6 +117,11 @@ impl McpClient {
|
|||||||
/// The config must use HTTP transport (the default); for stdio/UDS use `new_with_transport`.
|
/// The config must use HTTP transport (the default); for stdio/UDS use `new_with_transport`.
|
||||||
///
|
///
|
||||||
/// Returns an error if the config uses a non-HTTP transport.
|
/// Returns an error if the config uses a non-HTTP transport.
|
||||||
|
///
|
||||||
|
/// **Note:** The session manager is NOT wired into the transport. For
|
||||||
|
/// production use, prefer `create_client_from_config()` which constructs
|
||||||
|
/// the transport with session tracking.
|
||||||
|
#[cfg(test)]
|
||||||
pub fn new_with_config(config: McpServerConfig) -> Result<Self, ToolError> {
|
pub fn new_with_config(config: McpServerConfig) -> Result<Self, ToolError> {
|
||||||
if !matches!(
|
if !matches!(
|
||||||
config.effective_transport(),
|
config.effective_transport(),
|
||||||
@@ -214,7 +219,14 @@ impl McpClient {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Attach a session manager for Streamable HTTP session tracking.
|
/// Attach a session manager to the **client** only.
|
||||||
|
///
|
||||||
|
/// **Warning:** This does NOT wire the session manager into the underlying
|
||||||
|
/// `HttpMcpTransport`, so the transport will not capture `Mcp-Session-Id`
|
||||||
|
/// from responses. For production use, construct the transport with
|
||||||
|
/// `HttpMcpTransport::with_session_manager()` and pass it to
|
||||||
|
/// `new_with_transport()` instead. See `create_client_from_config()`.
|
||||||
|
#[cfg(test)]
|
||||||
pub fn with_session_manager(mut self, session_manager: Arc<McpSessionManager>) -> Self {
|
pub fn with_session_manager(mut self, session_manager: Arc<McpSessionManager>) -> Self {
|
||||||
self.session_manager = Some(session_manager);
|
self.session_manager = Some(session_manager);
|
||||||
self
|
self
|
||||||
@@ -235,6 +247,12 @@ impl McpClient {
|
|||||||
self.session_manager.is_some()
|
self.session_manager.is_some()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Get the underlying transport (test-only).
|
||||||
|
#[cfg(test)]
|
||||||
|
pub(crate) fn transport(&self) -> &Arc<dyn McpTransport> {
|
||||||
|
&self.transport
|
||||||
|
}
|
||||||
|
|
||||||
/// Get the next request ID.
|
/// Get the next request ID.
|
||||||
fn next_request_id(&self) -> u64 {
|
fn next_request_id(&self) -> u64 {
|
||||||
self.next_id.fetch_add(1, Ordering::SeqCst)
|
self.next_id.fetch_add(1, Ordering::SeqCst)
|
||||||
|
|||||||
+101
-16
@@ -7,6 +7,7 @@ use std::sync::Arc;
|
|||||||
|
|
||||||
use crate::secrets::SecretsStore;
|
use crate::secrets::SecretsStore;
|
||||||
use crate::tools::mcp::config::{EffectiveTransport, McpServerConfig};
|
use crate::tools::mcp::config::{EffectiveTransport, McpServerConfig};
|
||||||
|
use crate::tools::mcp::http_transport::HttpMcpTransport;
|
||||||
use crate::tools::mcp::{McpClient, McpProcessManager, McpSessionManager, McpTransport};
|
use crate::tools::mcp::{McpClient, McpProcessManager, McpSessionManager, McpTransport};
|
||||||
|
|
||||||
/// Error returned when MCP client creation fails.
|
/// Error returned when MCP client creation fails.
|
||||||
@@ -78,33 +79,37 @@ pub async fn create_client_from_config(
|
|||||||
Err(McpFactoryError::UnixNotSupported { name: server_name })
|
Err(McpFactoryError::UnixNotSupported { name: server_name })
|
||||||
}
|
}
|
||||||
EffectiveTransport::Http => {
|
EffectiveTransport::Http => {
|
||||||
|
// Authenticated (OAuth) path: tokens exist or server requires auth.
|
||||||
if let Some(ref secrets) = secrets {
|
if let Some(ref secrets) = secrets {
|
||||||
let has_tokens =
|
let has_tokens =
|
||||||
crate::tools::mcp::is_authenticated(&server, secrets, user_id).await;
|
crate::tools::mcp::is_authenticated(&server, secrets, user_id).await;
|
||||||
|
|
||||||
if has_tokens || server.requires_auth() {
|
if has_tokens || server.requires_auth() {
|
||||||
Ok(McpClient::new_authenticated(
|
return Ok(McpClient::new_authenticated(
|
||||||
server,
|
server,
|
||||||
Arc::clone(session_manager),
|
Arc::clone(session_manager),
|
||||||
Arc::clone(secrets),
|
Arc::clone(secrets),
|
||||||
user_id,
|
user_id,
|
||||||
))
|
));
|
||||||
} else {
|
|
||||||
Ok(McpClient::new_with_config(server)
|
|
||||||
.map_err(|e| McpFactoryError::InvalidConfig {
|
|
||||||
name: server_name.clone(),
|
|
||||||
reason: e.to_string(),
|
|
||||||
})?
|
|
||||||
.with_session_manager(Arc::clone(session_manager)))
|
|
||||||
}
|
}
|
||||||
} else {
|
|
||||||
Ok(McpClient::new_with_config(server)
|
|
||||||
.map_err(|e| McpFactoryError::InvalidConfig {
|
|
||||||
name: server_name,
|
|
||||||
reason: e.to_string(),
|
|
||||||
})?
|
|
||||||
.with_session_manager(Arc::clone(session_manager)))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Non-OAuth HTTP: wire the session manager into the *transport* so
|
||||||
|
// it captures `Mcp-Session-Id` from responses. Passing it only to
|
||||||
|
// the client (via `with_session_manager`) is not enough — the
|
||||||
|
// transport must know about it to read/write the header.
|
||||||
|
let transport = Arc::new(
|
||||||
|
HttpMcpTransport::new(server.url.clone(), server.name.clone())
|
||||||
|
.with_session_manager(Arc::clone(session_manager)),
|
||||||
|
);
|
||||||
|
Ok(McpClient::new_with_transport(
|
||||||
|
server.name.clone(),
|
||||||
|
transport,
|
||||||
|
Some(Arc::clone(session_manager)),
|
||||||
|
secrets,
|
||||||
|
user_id,
|
||||||
|
Some(server),
|
||||||
|
))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -134,4 +139,84 @@ mod tests {
|
|||||||
"non-OAuth HTTP clients must carry a session manager"
|
"non-OAuth HTTP clients must carry a session manager"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Regression test: the factory must wire the session manager into the
|
||||||
|
/// *transport*, not just the client. Otherwise the transport never
|
||||||
|
/// captures `Mcp-Session-Id` from responses and subsequent requests
|
||||||
|
/// lack the header, causing the server to reject them.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_factory_non_oauth_http_transport_captures_session_id() {
|
||||||
|
use axum::http::header::HeaderName;
|
||||||
|
use axum::{Router, http::StatusCode, response::IntoResponse, routing::post};
|
||||||
|
use tokio::net::TcpListener;
|
||||||
|
|
||||||
|
const SESSION_ID: &str = "test-session-abc123";
|
||||||
|
|
||||||
|
async fn session_echo() -> impl IntoResponse {
|
||||||
|
let body = serde_json::json!({
|
||||||
|
"jsonrpc": "2.0",
|
||||||
|
"id": 1,
|
||||||
|
"result": {}
|
||||||
|
})
|
||||||
|
.to_string();
|
||||||
|
(
|
||||||
|
StatusCode::OK,
|
||||||
|
[(
|
||||||
|
HeaderName::from_static("mcp-session-id"),
|
||||||
|
SESSION_ID.to_string(),
|
||||||
|
)],
|
||||||
|
body,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
let app = Router::new().route("/", post(session_echo));
|
||||||
|
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||||
|
let addr = listener.local_addr().unwrap();
|
||||||
|
let url = format!("http://127.0.0.1:{}", addr.port());
|
||||||
|
|
||||||
|
tokio::spawn(async move {
|
||||||
|
axum::serve(listener, app).await.unwrap();
|
||||||
|
});
|
||||||
|
|
||||||
|
let server = McpServerConfig::new("session-test", &url);
|
||||||
|
let session_manager = Arc::new(McpSessionManager::new());
|
||||||
|
let process_manager = Arc::new(McpProcessManager::new());
|
||||||
|
|
||||||
|
let client = create_client_from_config(
|
||||||
|
server,
|
||||||
|
&session_manager,
|
||||||
|
&process_manager,
|
||||||
|
None,
|
||||||
|
"test-user",
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect("factory should succeed for HTTP config");
|
||||||
|
|
||||||
|
// Pre-create a session entry so that update_session_id has something to update.
|
||||||
|
// In production, the MCP initialize handshake calls get_or_create before responses arrive.
|
||||||
|
session_manager.get_or_create("session-test", &url).await;
|
||||||
|
|
||||||
|
// Send a request through the client's transport to trigger session capture.
|
||||||
|
use crate::tools::mcp::protocol::McpRequest;
|
||||||
|
let request = McpRequest {
|
||||||
|
jsonrpc: "2.0".to_string(),
|
||||||
|
id: Some(1),
|
||||||
|
method: "test".to_string(),
|
||||||
|
params: Some(serde_json::json!({})),
|
||||||
|
};
|
||||||
|
let headers = std::collections::HashMap::new();
|
||||||
|
client
|
||||||
|
.transport()
|
||||||
|
.send(&request, &headers)
|
||||||
|
.await
|
||||||
|
.expect("request should succeed");
|
||||||
|
|
||||||
|
// Verify the session manager captured the session ID from the response.
|
||||||
|
let captured = session_manager.get_session_id("session-test").await;
|
||||||
|
assert_eq!(
|
||||||
|
captured.as_deref(),
|
||||||
|
Some(SESSION_ID),
|
||||||
|
"transport must capture Mcp-Session-Id into session manager"
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -494,6 +494,34 @@ mod tests {
|
|||||||
assert_eq!(echoed["authorization"], "Bearer oauth-token");
|
assert_eq!(echoed["authorization"], "Bearer oauth-token");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Regression test for #1436: 202 Accepted responses for notifications
|
||||||
|
/// were parsed as JSON, causing "Failed to parse MCP response" errors
|
||||||
|
/// that broke the MCP session handshake.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_wire_202_accepted_for_notification() {
|
||||||
|
use axum::{Router, http::StatusCode, routing::post};
|
||||||
|
use tokio::net::TcpListener;
|
||||||
|
|
||||||
|
async fn accept_notification() -> StatusCode {
|
||||||
|
StatusCode::ACCEPTED
|
||||||
|
}
|
||||||
|
|
||||||
|
let app = Router::new().route("/", post(accept_notification));
|
||||||
|
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||||
|
let addr = listener.local_addr().unwrap();
|
||||||
|
let url = format!("http://127.0.0.1:{}", addr.port());
|
||||||
|
|
||||||
|
tokio::spawn(async move {
|
||||||
|
axum::serve(listener, app).await.unwrap();
|
||||||
|
});
|
||||||
|
|
||||||
|
let transport = HttpMcpTransport::new(&url, "test-202");
|
||||||
|
let request = McpRequest::initialized_notification();
|
||||||
|
let response = transport.send(&request, &HashMap::new()).await.unwrap();
|
||||||
|
assert!(response.result.is_none());
|
||||||
|
assert!(response.error.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_wire_custom_auth_preserved_when_no_per_request_auth() {
|
async fn test_wire_custom_auth_preserved_when_no_per_request_auth() {
|
||||||
let (url, _handle) = spawn_echo_server().await;
|
let (url, _handle) = spawn_echo_server().await;
|
||||||
|
|||||||
@@ -205,7 +205,44 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// -----------------------------------------------------------------------
|
// -----------------------------------------------------------------------
|
||||||
// Test 5: routine_manual_create_defaults_to_tools_enabled
|
// Test 5: routine_update_fail_delete_fallback
|
||||||
|
// -----------------------------------------------------------------------
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn routine_update_fail_delete_fallback() {
|
||||||
|
let trace = LlmTrace::from_file(concat!(
|
||||||
|
env!("CARGO_MANIFEST_DIR"),
|
||||||
|
"/tests/fixtures/llm_traces/tools/routine_update_fail_delete_fallback.json"
|
||||||
|
))
|
||||||
|
.expect("failed to load routine_update_fail_delete_fallback.json");
|
||||||
|
|
||||||
|
let rig = TestRigBuilder::new()
|
||||||
|
.with_trace(trace.clone())
|
||||||
|
.with_auto_approve_tools(true)
|
||||||
|
.build()
|
||||||
|
.await;
|
||||||
|
|
||||||
|
rig.send_message("Try converting a routine trigger, then recover by deleting it")
|
||||||
|
.await;
|
||||||
|
let responses = rig.wait_for_responses(1, Duration::from_secs(15)).await;
|
||||||
|
|
||||||
|
rig.verify_trace_expects(&trace, &responses);
|
||||||
|
|
||||||
|
let completed = rig.tool_calls_completed();
|
||||||
|
assert!(
|
||||||
|
completed.iter().any(|(n, ok)| n == "routine_update" && !ok),
|
||||||
|
"routine_update should fail in this regression path: {completed:?}"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
completed.iter().any(|(n, ok)| n == "routine_delete" && *ok),
|
||||||
|
"routine_delete should recover successfully via preserved routine identity: {completed:?}"
|
||||||
|
);
|
||||||
|
|
||||||
|
rig.shutdown();
|
||||||
|
}
|
||||||
|
|
||||||
|
// -----------------------------------------------------------------------
|
||||||
|
// Test 6: routine_manual_create_defaults_to_tools_enabled
|
||||||
// -----------------------------------------------------------------------
|
// -----------------------------------------------------------------------
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
@@ -246,7 +283,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// -----------------------------------------------------------------------
|
// -----------------------------------------------------------------------
|
||||||
// Test 6: routine_manual_create_explicit_no_tools
|
// Test 7: routine_manual_create_explicit_no_tools
|
||||||
// -----------------------------------------------------------------------
|
// -----------------------------------------------------------------------
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
@@ -287,7 +324,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// -----------------------------------------------------------------------
|
// -----------------------------------------------------------------------
|
||||||
// Test 7: routine_history
|
// Test 8: routine_history
|
||||||
// -----------------------------------------------------------------------
|
// -----------------------------------------------------------------------
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
|
|||||||
@@ -0,0 +1,70 @@
|
|||||||
|
{
|
||||||
|
"model_name": "test-routine-update-fail-delete-fallback",
|
||||||
|
"expects": {
|
||||||
|
"tools_used": ["routine_create", "routine_update", "routine_delete"],
|
||||||
|
"tool_results_contain": {
|
||||||
|
"routine_update": "Cannot update schedule or timezone on a non-cron routine.",
|
||||||
|
"routine_delete": "temp-routine"
|
||||||
|
},
|
||||||
|
"min_responses": 1
|
||||||
|
},
|
||||||
|
"steps": [
|
||||||
|
{
|
||||||
|
"response": {
|
||||||
|
"type": "tool_calls",
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"id": "call_rc_fallback",
|
||||||
|
"name": "routine_create",
|
||||||
|
"arguments": {
|
||||||
|
"name": "temp-routine",
|
||||||
|
"trigger_type": "manual",
|
||||||
|
"prompt": "Temporary routine for fallback test."
|
||||||
|
}
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"input_tokens": 120,
|
||||||
|
"output_tokens": 40
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"response": {
|
||||||
|
"type": "tool_calls",
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"id": "call_ru_fallback",
|
||||||
|
"name": "routine_update",
|
||||||
|
"arguments": {
|
||||||
|
"name": "temp-routine",
|
||||||
|
"schedule": "0 */10 * * * *"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"input_tokens": 200,
|
||||||
|
"output_tokens": 30
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"response": {
|
||||||
|
"type": "tool_calls",
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"id": "call_rd_fallback",
|
||||||
|
"name": "routine_delete",
|
||||||
|
"arguments": {}
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"input_tokens": 300,
|
||||||
|
"output_tokens": 20
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"response": {
|
||||||
|
"type": "text",
|
||||||
|
"content": "I recovered from the failed update and cleaned up the original routine.",
|
||||||
|
"input_tokens": 380,
|
||||||
|
"output_tokens": 25
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user