mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-09-02 09:39:37 +00:00
feat: add channel-relay integration for Slack (#790)
* fix(ci): secrets can't be used in step if conditions [skip-regression-check] (#787) GitHub Actions step-level `if:` doesn't have access to `secrets` context. Replace `if: secrets.X != ''` with `continue-on-error: true` and let the Set token step handle the fallback. Co-authored-by: Claude Sonnet 4.6 <[email protected]> * feat: add channel-relay integration for Slack via external relay service - Add RelayChannel and RelayClient for connecting to channel-relay SSE streams - Add RelayConfig with env-based configuration (CHANNEL_RELAY_URL, CHANNEL_RELAY_API_KEY) - Add channel-relay extension lifecycle: install, OAuth auth, activate with hot-add - Add proxy message sending through channel-relay for Slack chat.postMessage - Add extension registry entry for Slack relay with OAuth auth hint - Add relay integration test with mock SSE server - Wire relay channel into app startup with reconnect on stored credentials - Add AuthRequired extension error variant for cleaner auth flow detection [skip-regression-check] * chore: apply cargo fmt * fix: remove remaining Telegram test references in relay channel * fix: address PR #790 review feedback — parser handle leak, CSRF, circuit breaker - Fix parser handle leak on reconnect by sharing Arc<RwLock> instead of creating a local copy in start() (shutdown now aborts the correct task) - Add CSRF state nonce to OAuth flow: generate in auth_channel_relay, validate in slack_relay_oauth_callback_handler, one-time use - Remove dead proxy_slack method, update integration test to use proxy_provider - Add reconnect circuit breaker (max_consecutive_failures, default 50) - Fix stale docs (Telegram refs), extract event_types constants --------- Co-authored-by: Henry Park <[email protected]> Co-authored-by: Claude Sonnet 4.6 <[email protected]>
This commit is contained in:
co-authored by
Henry Park
Claude Sonnet 4.6
parent
3a841b30d8
commit
b0214fef41
@@ -56,6 +56,17 @@ impl ChannelManager {
|
||||
/// the agent loop.
|
||||
pub async fn hot_add(&self, channel: Box<dyn Channel>) -> Result<(), ChannelError> {
|
||||
let name = channel.name().to_string();
|
||||
|
||||
// Shut down any existing channel with the same name to avoid parallel consumers.
|
||||
// The old forwarding task will stop when the channel's stream ends after shutdown.
|
||||
{
|
||||
let channels = self.channels.read().await;
|
||||
if let Some(existing) = channels.get(&name) {
|
||||
tracing::debug!(channel = %name, "Shutting down existing channel before hot-add replacement");
|
||||
let _ = existing.shutdown().await;
|
||||
}
|
||||
}
|
||||
|
||||
let stream = channel.start().await?;
|
||||
|
||||
// Register for respond/broadcast/send_status
|
||||
@@ -337,4 +348,30 @@ mod tests {
|
||||
let msg = stream.next().await.expect("stream ended");
|
||||
assert_eq!(msg.content, "background alert");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_hot_add_replaces_existing_channel() {
|
||||
// Regression: hot_add must shut down the existing channel before replacing it,
|
||||
// to prevent duplicate SSE consumers from running in parallel.
|
||||
let manager = ChannelManager::new();
|
||||
let (stub1, _tx1) = StubChannel::new("relay");
|
||||
manager.add(Box::new(stub1)).await;
|
||||
let mut stream = manager.start_all().await.expect("start_all");
|
||||
|
||||
// Hot-add a replacement channel with the same name
|
||||
let (stub2, tx2) = StubChannel::new("relay");
|
||||
manager.hot_add(Box::new(stub2)).await.expect("hot_add");
|
||||
|
||||
// Send through the new channel — should arrive in the merged stream
|
||||
tx2.send(IncomingMessage::new("relay", "u1", "from new"))
|
||||
.await
|
||||
.expect("send");
|
||||
let msg = stream.next().await.expect("stream");
|
||||
assert_eq!(msg.content, "from new");
|
||||
|
||||
// Verify only one channel entry exists
|
||||
let channels = manager.channels.read().await;
|
||||
assert_eq!(channels.len(), 1);
|
||||
assert!(channels.contains_key("relay"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -30,6 +30,7 @@
|
||||
mod channel;
|
||||
mod http;
|
||||
mod manager;
|
||||
pub mod relay;
|
||||
mod repl;
|
||||
mod signal;
|
||||
pub mod wasm;
|
||||
|
||||
@@ -0,0 +1,642 @@
|
||||
//! Channel trait implementation for channel-relay SSE streams.
|
||||
//!
|
||||
//! `RelayChannel` connects to a channel-relay service via SSE, converts
|
||||
//! incoming events to `IncomingMessage`s, and sends responses via the
|
||||
//! relay's provider-specific proxy API (Slack).
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tokio::sync::{RwLock, mpsc};
|
||||
|
||||
use crate::channels::relay::client::{RelayClient, RelayError};
|
||||
use crate::channels::{Channel, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate};
|
||||
use crate::error::ChannelError;
|
||||
|
||||
/// Default channel name for the Slack relay integration.
|
||||
pub const DEFAULT_RELAY_NAME: &str = "slack-relay";
|
||||
|
||||
/// The messaging provider backing a relay channel.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum RelayProvider {
|
||||
Slack,
|
||||
}
|
||||
|
||||
impl RelayProvider {
|
||||
/// Provider string used in proxy API routes and metadata.
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Slack => "slack",
|
||||
}
|
||||
}
|
||||
|
||||
/// The default channel name for this provider.
|
||||
pub fn channel_name(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Slack => DEFAULT_RELAY_NAME,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Channel implementation that connects to a channel-relay SSE stream.
|
||||
pub struct RelayChannel {
|
||||
client: RelayClient,
|
||||
provider: RelayProvider,
|
||||
stream_token: Arc<RwLock<String>>,
|
||||
team_id: String,
|
||||
instance_id: String,
|
||||
user_id: String,
|
||||
/// SSE stream long-poll timeout in seconds.
|
||||
stream_timeout_secs: u64,
|
||||
/// Initial exponential backoff in milliseconds.
|
||||
backoff_initial_ms: u64,
|
||||
/// Maximum exponential backoff in milliseconds.
|
||||
backoff_max_ms: u64,
|
||||
/// Handle to the reconnect task for clean shutdown.
|
||||
reconnect_handle: RwLock<Option<tokio::task::JoinHandle<()>>>,
|
||||
/// Handle to the SSE parser task for clean shutdown.
|
||||
parser_handle: Arc<RwLock<Option<tokio::task::JoinHandle<()>>>>,
|
||||
/// Maximum consecutive reconnect failures before giving up.
|
||||
max_consecutive_failures: u64,
|
||||
}
|
||||
|
||||
impl RelayChannel {
|
||||
/// Create a new relay channel for Slack (default provider).
|
||||
pub fn new(
|
||||
client: RelayClient,
|
||||
stream_token: String,
|
||||
team_id: String,
|
||||
instance_id: String,
|
||||
user_id: String,
|
||||
) -> Self {
|
||||
Self::new_with_provider(
|
||||
client,
|
||||
RelayProvider::Slack,
|
||||
stream_token,
|
||||
team_id,
|
||||
instance_id,
|
||||
user_id,
|
||||
)
|
||||
}
|
||||
|
||||
/// Create a new relay channel with a specific provider.
|
||||
pub fn new_with_provider(
|
||||
client: RelayClient,
|
||||
provider: RelayProvider,
|
||||
stream_token: String,
|
||||
team_id: String,
|
||||
instance_id: String,
|
||||
user_id: String,
|
||||
) -> Self {
|
||||
Self {
|
||||
client,
|
||||
provider,
|
||||
stream_token: Arc::new(RwLock::new(stream_token)),
|
||||
team_id,
|
||||
instance_id,
|
||||
user_id,
|
||||
stream_timeout_secs: 86400,
|
||||
backoff_initial_ms: 1000,
|
||||
backoff_max_ms: 60000,
|
||||
reconnect_handle: RwLock::new(None),
|
||||
parser_handle: Arc::new(RwLock::new(None)),
|
||||
max_consecutive_failures: 50,
|
||||
}
|
||||
}
|
||||
|
||||
/// Set backoff/timeout parameters from relay config values.
|
||||
pub fn with_timeouts(
|
||||
mut self,
|
||||
stream_timeout_secs: u64,
|
||||
backoff_initial_ms: u64,
|
||||
backoff_max_ms: u64,
|
||||
) -> Self {
|
||||
self.stream_timeout_secs = stream_timeout_secs;
|
||||
self.backoff_initial_ms = backoff_initial_ms;
|
||||
self.backoff_max_ms = backoff_max_ms;
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the maximum number of consecutive reconnect failures before giving up.
|
||||
pub fn with_max_failures(mut self, max: u64) -> Self {
|
||||
self.max_consecutive_failures = max;
|
||||
self
|
||||
}
|
||||
|
||||
/// Build a provider-appropriate proxy body for sending a message.
|
||||
fn build_send_body(
|
||||
&self,
|
||||
channel_id: &str,
|
||||
text: &str,
|
||||
thread_id: Option<&str>,
|
||||
) -> (String, serde_json::Value) {
|
||||
match self.provider {
|
||||
RelayProvider::Slack => {
|
||||
let mut body = serde_json::json!({
|
||||
"channel": channel_id,
|
||||
"text": text,
|
||||
});
|
||||
if let Some(tid) = thread_id {
|
||||
body["thread_ts"] = serde_json::Value::String(tid.to_string());
|
||||
}
|
||||
("chat.postMessage".to_string(), body)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Send a message via the provider proxy.
|
||||
async fn proxy_send(
|
||||
&self,
|
||||
team_id: &str,
|
||||
method: &str,
|
||||
body: serde_json::Value,
|
||||
) -> Result<serde_json::Value, RelayError> {
|
||||
self.client
|
||||
.proxy_provider(
|
||||
self.provider.as_str(),
|
||||
team_id,
|
||||
method,
|
||||
body,
|
||||
Some(&self.instance_id),
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Channel for RelayChannel {
|
||||
fn name(&self) -> &str {
|
||||
self.provider.channel_name()
|
||||
}
|
||||
|
||||
async fn start(&self) -> Result<MessageStream, ChannelError> {
|
||||
let channel_name = self.name().to_string();
|
||||
let token = self.stream_token.read().await.clone();
|
||||
let (stream, initial_parser_handle) = self
|
||||
.client
|
||||
.connect_stream(&token, self.stream_timeout_secs)
|
||||
.await
|
||||
.map_err(|e| ChannelError::StartupFailed {
|
||||
name: channel_name.clone(),
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
*self.parser_handle.write().await = Some(initial_parser_handle);
|
||||
|
||||
let (tx, rx) = mpsc::channel(64);
|
||||
|
||||
// Spawn the stream reader + reconnect task
|
||||
let client = self.client.clone();
|
||||
let stream_token = Arc::clone(&self.stream_token);
|
||||
let instance_id = self.instance_id.clone();
|
||||
let user_id = self.user_id.clone();
|
||||
let team_id = self.team_id.clone();
|
||||
let stream_timeout_secs = self.stream_timeout_secs;
|
||||
let backoff_initial_ms = self.backoff_initial_ms;
|
||||
let backoff_max_ms = self.backoff_max_ms;
|
||||
let max_consecutive_failures = self.max_consecutive_failures;
|
||||
let parser_handle = Arc::clone(&self.parser_handle);
|
||||
let provider_str = self.provider.as_str().to_string();
|
||||
let relay_name = channel_name.clone();
|
||||
|
||||
let handle = tokio::spawn(async move {
|
||||
use futures::StreamExt;
|
||||
|
||||
let mut current_stream = stream;
|
||||
let mut backoff_ms = backoff_initial_ms;
|
||||
let mut consecutive_failures: u64 = 0;
|
||||
|
||||
loop {
|
||||
// Read events from the current stream
|
||||
while let Some(event) = current_stream.next().await {
|
||||
// Reset backoff and failure count on successful event
|
||||
backoff_ms = backoff_initial_ms;
|
||||
consecutive_failures = 0;
|
||||
|
||||
// Validate required fields
|
||||
if event.sender_id.is_empty()
|
||||
|| event.channel_id.is_empty()
|
||||
|| event.provider_scope.is_empty()
|
||||
{
|
||||
tracing::debug!(
|
||||
event_type = %event.event_type,
|
||||
sender_id = %event.sender_id,
|
||||
channel_id = %event.channel_id,
|
||||
"Relay: skipping event with missing required fields"
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Skip non-message events
|
||||
if !event.is_message() {
|
||||
tracing::debug!(
|
||||
event_type = %event.event_type,
|
||||
"Relay: skipping non-message event"
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
event_type = %event.event_type,
|
||||
sender = %event.sender_id,
|
||||
channel = %event.channel_id,
|
||||
provider = %provider_str,
|
||||
"Relay: received message from {}", provider_str
|
||||
);
|
||||
|
||||
let msg = IncomingMessage::new(&relay_name, &event.sender_id, event.text())
|
||||
.with_user_name(event.display_name())
|
||||
.with_metadata(serde_json::json!({
|
||||
"team_id": event.team_id(),
|
||||
"channel_id": event.channel_id,
|
||||
"sender_id": event.sender_id,
|
||||
"sender_name": event.display_name(),
|
||||
"event_type": event.event_type,
|
||||
"thread_id": event.thread_id,
|
||||
"provider": event.provider,
|
||||
}));
|
||||
|
||||
let msg = if let Some(ref thread_id) = event.thread_id {
|
||||
msg.with_thread(thread_id)
|
||||
} else {
|
||||
msg.with_thread(&event.channel_id)
|
||||
};
|
||||
|
||||
if tx.send(msg).await.is_err() {
|
||||
tracing::info!("Relay channel receiver dropped, stopping");
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
// Stream ended, attempt reconnect with backoff
|
||||
consecutive_failures += 1;
|
||||
if consecutive_failures >= max_consecutive_failures {
|
||||
tracing::error!(
|
||||
channel = %relay_name,
|
||||
failures = consecutive_failures,
|
||||
"Relay channel giving up after {} consecutive failures",
|
||||
consecutive_failures
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
tracing::warn!(
|
||||
backoff_ms = backoff_ms,
|
||||
failures = consecutive_failures,
|
||||
"Relay SSE stream ended, reconnecting..."
|
||||
);
|
||||
tokio::time::sleep(std::time::Duration::from_millis(backoff_ms)).await;
|
||||
backoff_ms = (backoff_ms * 2).min(backoff_max_ms);
|
||||
|
||||
// Try to reconnect
|
||||
let token = stream_token.read().await.clone();
|
||||
match client.connect_stream(&token, stream_timeout_secs).await {
|
||||
Ok((new_stream, new_parser)) => {
|
||||
tracing::info!("Relay SSE stream reconnected");
|
||||
current_stream = new_stream;
|
||||
// Abort old parser before replacing
|
||||
if let Some(old) = parser_handle.write().await.take() {
|
||||
old.abort();
|
||||
}
|
||||
*parser_handle.write().await = Some(new_parser);
|
||||
}
|
||||
Err(RelayError::TokenExpired) => {
|
||||
// Attempt token renewal
|
||||
tracing::info!("Relay stream token expired, renewing...");
|
||||
match client.renew_token(&instance_id, &user_id).await {
|
||||
Ok(new_token) => {
|
||||
*stream_token.write().await = new_token.clone();
|
||||
match client.connect_stream(&new_token, stream_timeout_secs).await {
|
||||
Ok((new_stream, new_parser)) => {
|
||||
tracing::info!(
|
||||
"Relay SSE stream reconnected with new token"
|
||||
);
|
||||
current_stream = new_stream;
|
||||
if let Some(old) = parser_handle.write().await.take() {
|
||||
old.abort();
|
||||
}
|
||||
*parser_handle.write().await = Some(new_parser);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(
|
||||
error = %e,
|
||||
"Failed to reconnect after token renewal"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(
|
||||
error = %e,
|
||||
"Failed to renew relay stream token"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(error = %e, "Failed to reconnect relay SSE stream");
|
||||
}
|
||||
}
|
||||
|
||||
// Check if the team is still valid (skip when team_id is unknown,
|
||||
// e.g. when no DB store was available at activation time)
|
||||
if !team_id.is_empty() {
|
||||
match client.list_connections(&instance_id).await {
|
||||
Ok(conns) => {
|
||||
let has_team =
|
||||
conns.iter().any(|c| c.team_id == team_id && c.connected);
|
||||
if !has_team {
|
||||
tracing::warn!(
|
||||
team_id = %team_id,
|
||||
"Team no longer connected, stopping relay channel"
|
||||
);
|
||||
return;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
error = %e,
|
||||
"Could not verify team connection, will retry next iteration"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
*self.reconnect_handle.write().await = Some(handle);
|
||||
|
||||
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
|
||||
Ok(Box::pin(stream))
|
||||
}
|
||||
|
||||
async fn respond(
|
||||
&self,
|
||||
msg: &IncomingMessage,
|
||||
response: OutgoingResponse,
|
||||
) -> Result<(), ChannelError> {
|
||||
let channel_name = self.name().to_string();
|
||||
let metadata = &msg.metadata;
|
||||
let team_id = metadata
|
||||
.get("team_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or(&self.team_id);
|
||||
let channel_id = metadata
|
||||
.get("channel_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| ChannelError::SendFailed {
|
||||
name: channel_name.clone(),
|
||||
reason: "Missing channel_id in message metadata".to_string(),
|
||||
})?;
|
||||
|
||||
// Determine thread_id from response or metadata
|
||||
let thread_id = response
|
||||
.thread_id
|
||||
.as_deref()
|
||||
.or_else(|| metadata.get("thread_id").and_then(|v| v.as_str()));
|
||||
|
||||
let (method, body) = self.build_send_body(channel_id, &response.content, thread_id);
|
||||
|
||||
self.proxy_send(team_id, &method, body)
|
||||
.await
|
||||
.map_err(|e| ChannelError::SendFailed {
|
||||
name: channel_name,
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Status updates are not forwarded to messaging providers to avoid noise.
|
||||
async fn send_status(
|
||||
&self,
|
||||
_status: StatusUpdate,
|
||||
_metadata: &serde_json::Value,
|
||||
) -> Result<(), ChannelError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn broadcast(
|
||||
&self,
|
||||
target: &str,
|
||||
response: OutgoingResponse,
|
||||
) -> Result<(), ChannelError> {
|
||||
let channel_name = self.name().to_string();
|
||||
|
||||
// Determine thread_id from response or metadata
|
||||
let thread_id = response
|
||||
.thread_id
|
||||
.as_deref()
|
||||
.or_else(|| response.metadata.get("thread_ts").and_then(|v| v.as_str()));
|
||||
|
||||
let (method, body) = self.build_send_body(target, &response.content, thread_id);
|
||||
|
||||
self.proxy_send(&self.team_id, &method, body)
|
||||
.await
|
||||
.map_err(|e| ChannelError::SendFailed {
|
||||
name: channel_name,
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn health_check(&self) -> Result<(), ChannelError> {
|
||||
self.client
|
||||
.list_connections(&self.instance_id)
|
||||
.await
|
||||
.map_err(|_| ChannelError::HealthCheckFailed {
|
||||
name: self.name().to_string(),
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn conversation_context(&self, metadata: &serde_json::Value) -> HashMap<String, String> {
|
||||
let mut ctx = HashMap::new();
|
||||
|
||||
if let Some(sender) = metadata.get("sender_name").and_then(|v| v.as_str()) {
|
||||
ctx.insert("sender".to_string(), sender.to_string());
|
||||
}
|
||||
if let Some(sender_id) = metadata.get("sender_id").and_then(|v| v.as_str()) {
|
||||
ctx.insert("sender_uuid".to_string(), sender_id.to_string());
|
||||
}
|
||||
if let Some(channel_id) = metadata.get("channel_id").and_then(|v| v.as_str()) {
|
||||
ctx.insert("group".to_string(), channel_id.to_string());
|
||||
}
|
||||
ctx.insert("platform".to_string(), self.provider.as_str().to_string());
|
||||
|
||||
ctx
|
||||
}
|
||||
|
||||
async fn shutdown(&self) -> Result<(), ChannelError> {
|
||||
if let Some(handle) = self.reconnect_handle.write().await.take() {
|
||||
handle.abort();
|
||||
}
|
||||
if let Some(handle) = self.parser_handle.write().await.take() {
|
||||
handle.abort();
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn test_client() -> RelayClient {
|
||||
RelayClient::new(
|
||||
"http://localhost:3001".into(),
|
||||
secrecy::SecretString::from("key".to_string()),
|
||||
30,
|
||||
)
|
||||
.expect("client")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn relay_channel_name() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
assert_eq!(channel.name(), DEFAULT_RELAY_NAME);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn conversation_context_extracts_metadata() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
|
||||
let metadata = serde_json::json!({
|
||||
"sender_name": "bob",
|
||||
"sender_id": "U123",
|
||||
"channel_id": "C456",
|
||||
});
|
||||
let ctx = channel.conversation_context(&metadata);
|
||||
assert_eq!(ctx.get("sender"), Some(&"bob".to_string()));
|
||||
assert_eq!(ctx.get("sender_uuid"), Some(&"U123".to_string()));
|
||||
assert_eq!(ctx.get("platform"), Some(&"slack".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn metadata_shape_includes_event_type_and_sender_name() {
|
||||
// Regression: metadata JSON must include event_type and sender_name
|
||||
// for downstream routing (DM vs channel) and conversation_context().
|
||||
let metadata = serde_json::json!({
|
||||
"team_id": "T123",
|
||||
"channel_id": "C456",
|
||||
"sender_id": "U789",
|
||||
"sender_name": "alice",
|
||||
"event_type": "direct_message",
|
||||
"thread_id": null,
|
||||
"provider": "slack",
|
||||
});
|
||||
// event_type must be present for DM-vs-channel routing
|
||||
assert_eq!(
|
||||
metadata.get("event_type").and_then(|v| v.as_str()),
|
||||
Some("direct_message")
|
||||
);
|
||||
// sender_name must be present for conversation_context
|
||||
assert_eq!(
|
||||
metadata.get("sender_name").and_then(|v| v.as_str()),
|
||||
Some("alice")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn with_timeouts_sets_values() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
)
|
||||
.with_timeouts(43200, 2000, 120000);
|
||||
|
||||
assert_eq!(channel.stream_timeout_secs, 43200);
|
||||
assert_eq!(channel.backoff_initial_ms, 2000);
|
||||
assert_eq!(channel.backoff_max_ms, 120000);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_send_body_slack() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
let (method, body) = channel.build_send_body("C456", "hello", Some("1234567.890"));
|
||||
assert_eq!(method, "chat.postMessage");
|
||||
assert_eq!(body["channel"], "C456");
|
||||
assert_eq!(body["text"], "hello");
|
||||
assert_eq!(body["thread_ts"], "1234567.890");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parser_handle_is_shared_arc() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
// parser_handle should be an Arc — cloning should give a second reference
|
||||
let handle_clone = Arc::clone(&channel.parser_handle);
|
||||
// Both point to the same allocation
|
||||
assert!(Arc::ptr_eq(&channel.parser_handle, &handle_clone));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn with_max_failures_sets_value() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
)
|
||||
.with_max_failures(10);
|
||||
|
||||
assert_eq!(channel.max_consecutive_failures, 10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_max_failures_is_50() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
assert_eq!(channel.max_consecutive_failures, 50);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_team_id_accepted_at_construction() {
|
||||
// Regression: empty team_id (when no DB store is available) must not
|
||||
// prevent channel construction or cause immediate shutdown.
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
String::new(), // empty team_id
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
assert_eq!(channel.team_id, "");
|
||||
// The reconnect loop now skips team validation when team_id is empty,
|
||||
// so the channel remains alive.
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,549 @@
|
||||
//! HTTP client for the channel-relay service.
|
||||
//!
|
||||
//! Wraps reqwest for all channel-relay API calls: OAuth initiation,
|
||||
//! SSE streaming, token renewal, and Slack API proxy.
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use futures::Stream;
|
||||
use secrecy::{ExposeSecret, SecretString};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
/// Known relay event types.
|
||||
pub mod event_types {
|
||||
pub const MESSAGE: &str = "message";
|
||||
pub const DIRECT_MESSAGE: &str = "direct_message";
|
||||
pub const MENTION: &str = "mention";
|
||||
}
|
||||
|
||||
/// A parsed SSE event from the channel-relay stream.
|
||||
///
|
||||
/// Field names match the channel-relay `ChannelEvent` struct exactly.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ChannelEvent {
|
||||
/// Unique event ID.
|
||||
#[serde(default)]
|
||||
pub id: String,
|
||||
/// Event type enum from channel-relay (e.g., "direct_message", "message", "mention").
|
||||
pub event_type: String,
|
||||
/// Provider (e.g., "slack").
|
||||
#[serde(default)]
|
||||
pub provider: String,
|
||||
/// Team/workspace ID (called `provider_scope` in channel-relay).
|
||||
#[serde(alias = "team_id", default)]
|
||||
pub provider_scope: String,
|
||||
/// Channel or DM conversation ID.
|
||||
#[serde(default)]
|
||||
pub channel_id: String,
|
||||
/// Sender user ID.
|
||||
#[serde(default)]
|
||||
pub sender_id: String,
|
||||
/// Sender display name.
|
||||
#[serde(default)]
|
||||
pub sender_name: Option<String>,
|
||||
/// Message text content (called `content` in channel-relay).
|
||||
#[serde(alias = "text", default)]
|
||||
pub content: Option<String>,
|
||||
/// Thread ID (for threaded replies, called `thread_id` in channel-relay).
|
||||
#[serde(alias = "thread_ts", default)]
|
||||
pub thread_id: Option<String>,
|
||||
/// Full raw event data.
|
||||
#[serde(default)]
|
||||
pub raw: serde_json::Value,
|
||||
/// Event timestamp (ISO 8601 from channel-relay).
|
||||
#[serde(default)]
|
||||
pub timestamp: Option<String>,
|
||||
}
|
||||
|
||||
impl ChannelEvent {
|
||||
/// Get the team_id (provider_scope).
|
||||
pub fn team_id(&self) -> &str {
|
||||
&self.provider_scope
|
||||
}
|
||||
|
||||
/// Get the message text content.
|
||||
pub fn text(&self) -> &str {
|
||||
self.content.as_deref().unwrap_or("")
|
||||
}
|
||||
|
||||
/// Get the sender name or fallback to sender_id.
|
||||
pub fn display_name(&self) -> &str {
|
||||
self.sender_name.as_deref().unwrap_or(&self.sender_id)
|
||||
}
|
||||
|
||||
/// Check if this is a message-like event that should be forwarded to the agent.
|
||||
pub fn is_message(&self) -> bool {
|
||||
matches!(
|
||||
self.event_type.as_str(),
|
||||
event_types::MESSAGE | event_types::DIRECT_MESSAGE | event_types::MENTION
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// Connection info returned by list_connections.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct Connection {
|
||||
pub provider: String,
|
||||
pub team_id: String,
|
||||
pub team_name: Option<String>,
|
||||
pub connected: bool,
|
||||
}
|
||||
|
||||
/// HTTP client for the channel-relay service.
|
||||
#[derive(Clone)]
|
||||
pub struct RelayClient {
|
||||
http: reqwest::Client,
|
||||
base_url: String,
|
||||
api_key: SecretString,
|
||||
}
|
||||
|
||||
impl RelayClient {
|
||||
/// Create a new relay client.
|
||||
pub fn new(
|
||||
base_url: String,
|
||||
api_key: SecretString,
|
||||
request_timeout_secs: u64,
|
||||
) -> Result<Self, RelayError> {
|
||||
let http = reqwest::Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(request_timeout_secs))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| RelayError::Network(format!("Failed to build HTTP client: {e}")))?;
|
||||
|
||||
Ok(Self {
|
||||
http,
|
||||
base_url: base_url.trim_end_matches('/').to_string(),
|
||||
api_key,
|
||||
})
|
||||
}
|
||||
|
||||
/// Initiate Slack OAuth flow via channel-relay.
|
||||
///
|
||||
/// Calls `GET /oauth/slack/auth` with `redirect(Policy::none())` and
|
||||
/// returns the `Location` header (Slack OAuth URL) without following it.
|
||||
pub async fn initiate_oauth(
|
||||
&self,
|
||||
instance_id: &str,
|
||||
user_id: &str,
|
||||
callback_url: &str,
|
||||
) -> Result<String, RelayError> {
|
||||
let resp = self
|
||||
.http
|
||||
.get(format!("{}/oauth/slack/auth", self.base_url))
|
||||
.header("X-API-Key", self.api_key.expose_secret())
|
||||
.query(&[
|
||||
("instance_id", instance_id),
|
||||
("user_id", user_id),
|
||||
("callback", callback_url),
|
||||
])
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
||||
|
||||
let status = resp.status();
|
||||
if status.is_redirection() {
|
||||
let location = resp
|
||||
.headers()
|
||||
.get(reqwest::header::LOCATION)
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.map(|s| s.to_string())
|
||||
.ok_or_else(|| {
|
||||
RelayError::Protocol("Redirect response missing Location header".to_string())
|
||||
})?;
|
||||
Ok(location)
|
||||
} else if status.is_success() {
|
||||
// Some relay implementations return the URL in JSON body instead
|
||||
let body: serde_json::Value = resp
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| RelayError::Protocol(e.to_string()))?;
|
||||
body.get("auth_url")
|
||||
.or_else(|| body.get("url"))
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string())
|
||||
.ok_or_else(|| RelayError::Protocol("Response missing auth_url field".to_string()))
|
||||
} else {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
Err(RelayError::Api {
|
||||
status: status.as_u16(),
|
||||
message: body,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Connect to the SSE event stream.
|
||||
///
|
||||
/// Returns a stream of parsed `ChannelEvent`s and the `JoinHandle` of the
|
||||
/// background SSE parser task. The caller is responsible for reconnection
|
||||
/// logic on stream end/error and for aborting the handle on shutdown.
|
||||
pub async fn connect_stream(
|
||||
&self,
|
||||
stream_token: &str,
|
||||
stream_timeout_secs: u64,
|
||||
) -> Result<(ChannelEventStream, tokio::task::JoinHandle<()>), RelayError> {
|
||||
let resp = self
|
||||
.http
|
||||
.get(format!("{}/stream", self.base_url))
|
||||
.query(&[("token", stream_token)])
|
||||
.timeout(std::time::Duration::from_secs(stream_timeout_secs))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
||||
|
||||
let status = resp.status();
|
||||
if status == reqwest::StatusCode::UNAUTHORIZED {
|
||||
return Err(RelayError::TokenExpired);
|
||||
}
|
||||
if !status.is_success() {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
return Err(RelayError::Api {
|
||||
status: status.as_u16(),
|
||||
message: body,
|
||||
});
|
||||
}
|
||||
|
||||
// Spawn a background task that reads the SSE stream and sends parsed events
|
||||
let (tx, rx) = mpsc::channel(64);
|
||||
let byte_stream = resp.bytes_stream();
|
||||
let handle = tokio::spawn(parse_sse_stream(byte_stream, tx));
|
||||
|
||||
Ok((ChannelEventStream { rx }, handle))
|
||||
}
|
||||
|
||||
/// Renew an expired stream token.
|
||||
///
|
||||
/// Calls `POST /stream/renew` with API key auth, returns a new stream token.
|
||||
pub async fn renew_token(
|
||||
&self,
|
||||
instance_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<String, RelayError> {
|
||||
let resp = self
|
||||
.http
|
||||
.post(format!("{}/stream/renew", self.base_url))
|
||||
.header("X-API-Key", self.api_key.expose_secret())
|
||||
.json(&serde_json::json!({
|
||||
"instance_id": instance_id,
|
||||
"user_id": user_id,
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
||||
|
||||
let status = resp.status();
|
||||
if !status.is_success() {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
return Err(RelayError::Api {
|
||||
status: status.as_u16(),
|
||||
message: body,
|
||||
});
|
||||
}
|
||||
|
||||
let body: serde_json::Value = resp
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| RelayError::Protocol(e.to_string()))?;
|
||||
body.get("stream_token")
|
||||
.or_else(|| body.get("token"))
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string())
|
||||
.ok_or_else(|| RelayError::Protocol("Response missing stream_token field".to_string()))
|
||||
}
|
||||
|
||||
/// Proxy an API call through channel-relay for any provider.
|
||||
///
|
||||
/// Calls `POST /proxy/{provider}/{method}?team_id=X&instance_id=Y` with the given JSON body.
|
||||
pub async fn proxy_provider(
|
||||
&self,
|
||||
provider: &str,
|
||||
team_id: &str,
|
||||
method: &str,
|
||||
body: serde_json::Value,
|
||||
instance_id: Option<&str>,
|
||||
) -> Result<serde_json::Value, RelayError> {
|
||||
let mut query: Vec<(&str, &str)> = vec![("team_id", team_id)];
|
||||
if let Some(iid) = instance_id {
|
||||
query.push(("instance_id", iid));
|
||||
}
|
||||
let resp = self
|
||||
.http
|
||||
.post(format!("{}/proxy/{}/{}", self.base_url, provider, method))
|
||||
.header("X-API-Key", self.api_key.expose_secret())
|
||||
.query(&query)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
||||
|
||||
if !resp.status().is_success() {
|
||||
let status = resp.status().as_u16();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
return Err(RelayError::Api {
|
||||
status,
|
||||
message: body,
|
||||
});
|
||||
}
|
||||
|
||||
resp.json()
|
||||
.await
|
||||
.map_err(|e| RelayError::Protocol(e.to_string()))
|
||||
}
|
||||
|
||||
/// List active connections for an instance.
|
||||
pub async fn list_connections(&self, instance_id: &str) -> Result<Vec<Connection>, RelayError> {
|
||||
let resp = self
|
||||
.http
|
||||
.get(format!("{}/connections", self.base_url))
|
||||
.header("X-API-Key", self.api_key.expose_secret())
|
||||
.query(&[("instance_id", instance_id)])
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
||||
|
||||
if !resp.status().is_success() {
|
||||
let status = resp.status().as_u16();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
return Err(RelayError::Api {
|
||||
status,
|
||||
message: body,
|
||||
});
|
||||
}
|
||||
|
||||
resp.json()
|
||||
.await
|
||||
.map_err(|e| RelayError::Protocol(e.to_string()))
|
||||
}
|
||||
}
|
||||
|
||||
/// Async stream of parsed channel events from SSE.
|
||||
pub struct ChannelEventStream {
|
||||
rx: mpsc::Receiver<ChannelEvent>,
|
||||
}
|
||||
|
||||
impl Stream for ChannelEventStream {
|
||||
type Item = ChannelEvent;
|
||||
|
||||
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
self.rx.poll_recv(cx)
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse SSE format from a reqwest bytes stream.
|
||||
///
|
||||
/// SSE format:
|
||||
/// ```text
|
||||
/// event: message
|
||||
/// data: {"key": "value"}
|
||||
///
|
||||
/// ```
|
||||
/// Blank line terminates an event.
|
||||
async fn parse_sse_stream(
|
||||
byte_stream: impl futures::Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Send + 'static,
|
||||
tx: mpsc::Sender<ChannelEvent>,
|
||||
) {
|
||||
use futures::StreamExt;
|
||||
|
||||
let mut buffer = Vec::<u8>::new();
|
||||
let mut event_type = String::new();
|
||||
let mut data_lines = Vec::new();
|
||||
|
||||
let mut byte_stream = std::pin::pin!(byte_stream);
|
||||
while let Some(chunk_result) = byte_stream.next().await {
|
||||
let chunk = match chunk_result {
|
||||
Ok(c) => c,
|
||||
Err(e) => {
|
||||
tracing::debug!(error = %e, "SSE stream chunk error");
|
||||
break;
|
||||
}
|
||||
};
|
||||
|
||||
buffer.extend_from_slice(&chunk);
|
||||
|
||||
// Process complete lines (decode UTF-8 only on full lines to avoid
|
||||
// corruption when multi-byte characters span chunk boundaries)
|
||||
while let Some(newline_pos) = buffer.iter().position(|&b| b == b'\n') {
|
||||
let line = String::from_utf8_lossy(&buffer[..newline_pos])
|
||||
.trim_end_matches('\r')
|
||||
.to_string();
|
||||
buffer.drain(..=newline_pos);
|
||||
|
||||
if line.is_empty() {
|
||||
// Blank line = end of event
|
||||
if !data_lines.is_empty() {
|
||||
let data = data_lines.join("\n");
|
||||
if let Ok(mut event) = serde_json::from_str::<ChannelEvent>(&data) {
|
||||
if event.event_type.is_empty() && !event_type.is_empty() {
|
||||
event.event_type = event_type.clone();
|
||||
}
|
||||
if tx.send(event).await.is_err() {
|
||||
return; // receiver dropped
|
||||
}
|
||||
} else {
|
||||
tracing::debug!(
|
||||
event_type = %event_type,
|
||||
data_len = data.len(),
|
||||
"Failed to parse SSE event data as ChannelEvent"
|
||||
);
|
||||
}
|
||||
}
|
||||
event_type.clear();
|
||||
data_lines.clear();
|
||||
} else if let Some(value) = line.strip_prefix("event:") {
|
||||
event_type = value.trim().to_string();
|
||||
} else if let Some(value) = line.strip_prefix("data:") {
|
||||
data_lines.push(value.trim().to_string());
|
||||
}
|
||||
// Ignore other fields (id:, retry:, comments)
|
||||
}
|
||||
}
|
||||
|
||||
tracing::debug!("SSE stream ended");
|
||||
}
|
||||
|
||||
/// Errors from relay client operations.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum RelayError {
|
||||
#[error("Network error: {0}")]
|
||||
Network(String),
|
||||
|
||||
#[error("API error (HTTP {status}): {message}")]
|
||||
Api { status: u16, message: String },
|
||||
|
||||
#[error("Protocol error: {0}")]
|
||||
Protocol(String),
|
||||
|
||||
#[error("Stream token expired")]
|
||||
TokenExpired,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn channel_event_deserialize_minimal() {
|
||||
let json = r#"{"event_type": "message", "content": "hello"}"#;
|
||||
let event: ChannelEvent = serde_json::from_str(json).expect("parse failed");
|
||||
assert_eq!(event.event_type, "message");
|
||||
assert_eq!(event.text(), "hello");
|
||||
assert!(event.provider_scope.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn channel_event_deserialize_relay_format() {
|
||||
// Matches the actual channel-relay ChannelEvent serialization format.
|
||||
let json = r#"{
|
||||
"id": "evt_123",
|
||||
"event_type": "direct_message",
|
||||
"provider": "slack",
|
||||
"provider_scope": "T123",
|
||||
"channel_id": "D456",
|
||||
"sender_id": "U789",
|
||||
"sender_name": "bob",
|
||||
"content": "hi there",
|
||||
"thread_id": "1234567890.123456",
|
||||
"raw": {},
|
||||
"timestamp": "2026-03-09T21:00:00Z"
|
||||
}"#;
|
||||
let event: ChannelEvent = serde_json::from_str(json).expect("parse failed");
|
||||
assert_eq!(event.provider, "slack");
|
||||
assert_eq!(event.team_id(), "T123");
|
||||
assert_eq!(event.display_name(), "bob");
|
||||
assert_eq!(event.thread_id, Some("1234567890.123456".to_string()));
|
||||
assert!(event.is_message());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn channel_event_is_message() {
|
||||
let make = |et: &str| ChannelEvent {
|
||||
id: String::new(),
|
||||
event_type: et.to_string(),
|
||||
provider: String::new(),
|
||||
provider_scope: String::new(),
|
||||
channel_id: String::new(),
|
||||
sender_id: String::new(),
|
||||
sender_name: None,
|
||||
content: None,
|
||||
thread_id: None,
|
||||
raw: serde_json::Value::Null,
|
||||
timestamp: None,
|
||||
};
|
||||
assert!(make("message").is_message());
|
||||
assert!(make("direct_message").is_message());
|
||||
assert!(make("mention").is_message());
|
||||
assert!(!make("reaction").is_message());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn connection_deserialize() {
|
||||
let json = r#"{"provider": "slack", "team_id": "T123", "team_name": "My Team", "connected": true}"#;
|
||||
let conn: Connection = serde_json::from_str(json).expect("parse failed");
|
||||
assert_eq!(conn.provider, "slack");
|
||||
assert!(conn.connected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn relay_error_display() {
|
||||
let err = RelayError::Network("timeout".into());
|
||||
assert_eq!(err.to_string(), "Network error: timeout");
|
||||
|
||||
let err = RelayError::Api {
|
||||
status: 401,
|
||||
message: "unauthorized".into(),
|
||||
};
|
||||
assert_eq!(err.to_string(), "API error (HTTP 401): unauthorized");
|
||||
|
||||
let err = RelayError::TokenExpired;
|
||||
assert_eq!(err.to_string(), "Stream token expired");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_type_constants_match_is_message() {
|
||||
let make = |et: &str| ChannelEvent {
|
||||
id: String::new(),
|
||||
event_type: et.to_string(),
|
||||
provider: String::new(),
|
||||
provider_scope: String::new(),
|
||||
channel_id: String::new(),
|
||||
sender_id: String::new(),
|
||||
sender_name: None,
|
||||
content: None,
|
||||
thread_id: None,
|
||||
raw: serde_json::Value::Null,
|
||||
timestamp: None,
|
||||
};
|
||||
assert!(make(event_types::MESSAGE).is_message());
|
||||
assert!(make(event_types::DIRECT_MESSAGE).is_message());
|
||||
assert!(make(event_types::MENTION).is_message());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn parse_sse_handles_multibyte_utf8_across_chunks() {
|
||||
// The crab emoji (🦀) is 4 bytes: [0xF0, 0x9F, 0xA6, 0x80].
|
||||
// Split it across two chunks to verify no U+FFFD corruption.
|
||||
let event_json = r#"{"event_type":"message","content":"hello 🦀 world","provider_scope":"T1","channel_id":"C1","sender_id":"U1"}"#;
|
||||
let full = format!("event: message\ndata: {}\n\n", event_json);
|
||||
let bytes = full.as_bytes();
|
||||
|
||||
// Find the crab emoji and split mid-character
|
||||
let crab_pos = bytes
|
||||
.windows(4)
|
||||
.position(|w| w == [0xF0, 0x9F, 0xA6, 0x80])
|
||||
.expect("crab emoji not found");
|
||||
let split_at = crab_pos + 2; // split in the middle of the 4-byte emoji
|
||||
|
||||
let chunk1 = bytes::Bytes::copy_from_slice(&bytes[..split_at]);
|
||||
let chunk2 = bytes::Bytes::copy_from_slice(&bytes[split_at..]);
|
||||
|
||||
let chunks: Vec<Result<bytes::Bytes, reqwest::Error>> = vec![Ok(chunk1), Ok(chunk2)];
|
||||
let stream = futures::stream::iter(chunks);
|
||||
|
||||
let (tx, mut rx) = mpsc::channel(8);
|
||||
parse_sse_stream(stream, tx).await;
|
||||
|
||||
let event = rx.recv().await.expect("should receive event");
|
||||
assert_eq!(event.text(), "hello 🦀 world");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
//! Channel-relay integration for connecting to external messaging platforms
|
||||
//! (Slack) via the channel-relay service.
|
||||
//!
|
||||
//! The relay service handles OAuth, credential storage, webhook ingestion,
|
||||
//! and SSE event streaming. IronClaw consumes the SSE stream and sends
|
||||
//! messages via the relay's proxy API.
|
||||
|
||||
pub mod channel;
|
||||
pub mod client;
|
||||
|
||||
pub use channel::{DEFAULT_RELAY_NAME, RelayChannel};
|
||||
pub use client::RelayClient;
|
||||
@@ -46,6 +46,14 @@ pub async fn extensions_list_handler(
|
||||
} else {
|
||||
"configured".to_string()
|
||||
})
|
||||
} else if ext.kind == crate::extensions::ExtensionKind::ChannelRelay {
|
||||
Some(if ext.active {
|
||||
"active".to_string()
|
||||
} else if ext.authenticated {
|
||||
"configured".to_string()
|
||||
} else {
|
||||
"installed".to_string()
|
||||
})
|
||||
} else {
|
||||
None
|
||||
};
|
||||
@@ -103,6 +111,7 @@ pub async fn extensions_install_handler(
|
||||
"mcp_server" => Some(crate::extensions::ExtensionKind::McpServer),
|
||||
"wasm_tool" => Some(crate::extensions::ExtensionKind::WasmTool),
|
||||
"wasm_channel" => Some(crate::extensions::ExtensionKind::WasmChannel),
|
||||
"channel_relay" => Some(crate::extensions::ExtensionKind::ChannelRelay),
|
||||
_ => None,
|
||||
});
|
||||
|
||||
@@ -115,62 +124,6 @@ pub async fn extensions_install_handler(
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn extensions_activate_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(name): Path<String>,
|
||||
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
||||
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
||||
StatusCode::NOT_IMPLEMENTED,
|
||||
"Extension manager not available (secrets store required)".to_string(),
|
||||
))?;
|
||||
|
||||
match ext_mgr.activate(&name).await {
|
||||
Ok(result) => {
|
||||
// Activation just loads the WASM module. Auth (OAuth/manual) is
|
||||
// triggered separately via save_setup_secrets or the auth endpoint.
|
||||
Ok(Json(ActionResponse::ok(result.message)))
|
||||
}
|
||||
Err(activate_err) => {
|
||||
let err_str = activate_err.to_string();
|
||||
let needs_auth = err_str.contains("authentication")
|
||||
|| err_str.contains("401")
|
||||
|| err_str.contains("Unauthorized");
|
||||
|
||||
if !needs_auth {
|
||||
return Ok(Json(ActionResponse::fail(err_str)));
|
||||
}
|
||||
|
||||
// Activation failed due to auth; try authenticating first.
|
||||
match ext_mgr.auth(&name, None).await {
|
||||
Ok(auth_result) if auth_result.is_authenticated() => {
|
||||
// Auth succeeded, retry activation.
|
||||
match ext_mgr.activate(&name).await {
|
||||
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
||||
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
|
||||
}
|
||||
}
|
||||
Ok(auth_result) => {
|
||||
// Auth in progress (OAuth URL or awaiting manual token).
|
||||
let mut resp = ActionResponse::fail(
|
||||
auth_result
|
||||
.instructions()
|
||||
.map(String::from)
|
||||
.unwrap_or_else(|| format!("'{}' requires authentication.", name)),
|
||||
);
|
||||
resp.auth_url = auth_result.auth_url().map(String::from);
|
||||
resp.awaiting_token = Some(auth_result.is_awaiting_token());
|
||||
resp.instructions = auth_result.instructions().map(String::from);
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Err(auth_err) => Ok(Json(ActionResponse::fail(format!(
|
||||
"Authentication failed: {}",
|
||||
auth_err
|
||||
)))),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn extensions_remove_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(name): Path<String>,
|
||||
|
||||
@@ -97,6 +97,7 @@ impl GatewayChannel {
|
||||
skill_registry: None,
|
||||
skill_catalog: None,
|
||||
chat_rate_limiter: server::RateLimiter::new(30, 60),
|
||||
oauth_rate_limiter: server::RateLimiter::new(10, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
@@ -133,6 +134,7 @@ impl GatewayChannel {
|
||||
skill_registry: self.state.skill_registry.clone(),
|
||||
skill_catalog: self.state.skill_catalog.clone(),
|
||||
chat_rate_limiter: server::RateLimiter::new(30, 60),
|
||||
oauth_rate_limiter: server::RateLimiter::new(10, 60),
|
||||
registry_entries: self.state.registry_entries.clone(),
|
||||
cost_guard: self.state.cost_guard.clone(),
|
||||
routine_engine: Arc::clone(&self.state.routine_engine),
|
||||
|
||||
+389
-6
@@ -28,6 +28,7 @@ use uuid::Uuid;
|
||||
use crate::agent::SessionManager;
|
||||
use crate::bootstrap::ironclaw_base_dir;
|
||||
use crate::channels::IncomingMessage;
|
||||
use crate::channels::relay::DEFAULT_RELAY_NAME;
|
||||
use crate::channels::web::auth::{AuthState, auth_middleware};
|
||||
use crate::channels::web::handlers::jobs::{
|
||||
job_files_list_handler, job_files_read_handler, jobs_cancel_handler, jobs_detail_handler,
|
||||
@@ -164,6 +165,8 @@ pub struct GatewayState {
|
||||
pub scheduler: Option<crate::tools::builtin::SchedulerSlot>,
|
||||
/// Rate limiter for chat endpoints (30 messages per 60 seconds).
|
||||
pub chat_rate_limiter: RateLimiter,
|
||||
/// Rate limiter for OAuth callback endpoints (10 requests per 60 seconds).
|
||||
pub oauth_rate_limiter: RateLimiter,
|
||||
/// Registry catalog entries for the available extensions API.
|
||||
/// Populated at startup from `registry/` manifests, independent of extension manager.
|
||||
pub registry_entries: Vec<crate::extensions::RegistryEntry>,
|
||||
@@ -200,7 +203,11 @@ pub async fn start_server(
|
||||
// Public routes (no auth)
|
||||
let public = Router::new()
|
||||
.route("/api/health", get(health_handler))
|
||||
.route("/oauth/callback", get(oauth_callback_handler));
|
||||
.route("/oauth/callback", get(oauth_callback_handler))
|
||||
.route(
|
||||
"/oauth/slack/callback",
|
||||
get(slack_relay_oauth_callback_handler),
|
||||
);
|
||||
|
||||
// Protected routes (require auth)
|
||||
let auth_state = AuthState { token: auth_token };
|
||||
@@ -606,6 +613,208 @@ async fn oauth_callback_handler(
|
||||
axum::response::Html(html).into_response()
|
||||
}
|
||||
|
||||
/// OAuth callback for Slack via channel-relay.
|
||||
///
|
||||
/// This is a PUBLIC route (no Bearer token required) because channel-relay
|
||||
/// redirects the user's browser here after Slack OAuth completes.
|
||||
/// Query params: `stream_token`, `provider`, `team_id`.
|
||||
async fn slack_relay_oauth_callback_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Query(params): Query<std::collections::HashMap<String, String>>,
|
||||
) -> impl IntoResponse {
|
||||
// Rate limit
|
||||
if !state.oauth_rate_limiter.check() {
|
||||
return axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Too Many Requests</h2>\
|
||||
<p>Please try again later.</p>\
|
||||
</body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
|
||||
// Validate stream_token: required, non-empty, max 2048 bytes
|
||||
let stream_token = match params.get("stream_token") {
|
||||
Some(t) if !t.is_empty() && t.len() <= 2048 => t.clone(),
|
||||
Some(t) if t.len() > 2048 => {
|
||||
return axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Error</h2><p>Invalid callback parameters.</p></body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
_ => {
|
||||
return axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Error</h2><p>Invalid callback parameters.</p></body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
};
|
||||
|
||||
// Validate team_id format: empty or T followed by alphanumeric (max 20 chars)
|
||||
let team_id = params.get("team_id").cloned().unwrap_or_default();
|
||||
if !team_id.is_empty() {
|
||||
let valid_team_id = team_id.len() <= 21
|
||||
&& team_id.starts_with('T')
|
||||
&& team_id[1..].chars().all(|c| c.is_ascii_alphanumeric());
|
||||
if !valid_team_id {
|
||||
return axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Error</h2><p>Invalid callback parameters.</p></body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
}
|
||||
|
||||
// Validate provider: must be "slack" (only supported provider)
|
||||
let provider = params
|
||||
.get("provider")
|
||||
.cloned()
|
||||
.unwrap_or_else(|| "slack".into());
|
||||
if provider != "slack" {
|
||||
return axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Error</h2><p>Invalid callback parameters.</p></body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
|
||||
let ext_mgr = match state.extension_manager.as_ref() {
|
||||
Some(mgr) => mgr,
|
||||
None => {
|
||||
return axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Error</h2><p>Extension manager not available.</p></body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
};
|
||||
|
||||
// Validate CSRF state parameter
|
||||
let state_param = match params.get("state") {
|
||||
Some(s) if !s.is_empty() && s.len() <= 128 => s.clone(),
|
||||
_ => {
|
||||
return axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Error</h2><p>Invalid or expired authorization.</p></body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
};
|
||||
|
||||
let state_key = format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME);
|
||||
let stored_state = match ext_mgr
|
||||
.secrets()
|
||||
.get_decrypted(&state.user_id, &state_key)
|
||||
.await
|
||||
{
|
||||
Ok(secret) => secret.expose().to_string(),
|
||||
Err(_) => {
|
||||
return axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Error</h2><p>Invalid or expired authorization.</p></body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
};
|
||||
|
||||
if state_param != stored_state {
|
||||
return axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Error</h2><p>Invalid or expired authorization.</p></body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
|
||||
// Delete the nonce (one-time use)
|
||||
let _ = ext_mgr.secrets().delete(&state.user_id, &state_key).await;
|
||||
|
||||
let result: Result<(), String> = async {
|
||||
// Store the stream token as a secret
|
||||
let token_key = format!("relay:{}:stream_token", DEFAULT_RELAY_NAME);
|
||||
let _ = ext_mgr.secrets().delete(&state.user_id, &token_key).await;
|
||||
ext_mgr
|
||||
.secrets()
|
||||
.create(
|
||||
&state.user_id,
|
||||
crate::secrets::CreateSecretParams {
|
||||
name: token_key,
|
||||
value: secrecy::SecretString::from(stream_token),
|
||||
provider: Some(provider.clone()),
|
||||
expires_at: None,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.map_err(|e| format!("Failed to store stream token: {}", e))?;
|
||||
|
||||
// Store team_id in settings
|
||||
if let Some(ref store) = state.store {
|
||||
let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME);
|
||||
let _ = store
|
||||
.set_setting(&state.user_id, &team_id_key, &serde_json::json!(team_id))
|
||||
.await;
|
||||
}
|
||||
|
||||
// Activate the relay channel
|
||||
ext_mgr
|
||||
.activate_stored_relay(DEFAULT_RELAY_NAME)
|
||||
.await
|
||||
.map_err(|e| format!("Failed to activate relay channel: {}", e))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
.await;
|
||||
|
||||
let (success, message) = match &result {
|
||||
Ok(()) => (true, "Slack connected successfully!".to_string()),
|
||||
Err(e) => {
|
||||
tracing::error!(error = %e, "Slack relay OAuth callback failed");
|
||||
(
|
||||
false,
|
||||
"Connection failed. Check server logs for details.".to_string(),
|
||||
)
|
||||
}
|
||||
};
|
||||
|
||||
// Broadcast SSE event to notify the web UI
|
||||
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||
extension_name: DEFAULT_RELAY_NAME.to_string(),
|
||||
success,
|
||||
message: message.clone(),
|
||||
});
|
||||
|
||||
if success {
|
||||
axum::response::Html(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Slack Connected!</h2>\
|
||||
<p>You can close this tab and return to IronClaw.</p>\
|
||||
<script>window.close()</script>\
|
||||
</body></html>"
|
||||
.to_string(),
|
||||
)
|
||||
.into_response()
|
||||
} else {
|
||||
axum::response::Html(format!(
|
||||
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
|
||||
<h2>Connection Failed</h2>\
|
||||
<p>{}</p>\
|
||||
</body></html>",
|
||||
message
|
||||
))
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
|
||||
// --- Chat handlers ---
|
||||
|
||||
/// Convert web gateway `ImageData` to `IncomingAttachment` objects.
|
||||
@@ -1639,13 +1848,13 @@ async fn extensions_activate_handler(
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Err(activate_err) => {
|
||||
let err_str = activate_err.to_string();
|
||||
let needs_auth = err_str.contains("authentication")
|
||||
|| err_str.contains("401")
|
||||
|| err_str.contains("Unauthorized");
|
||||
let needs_auth = matches!(
|
||||
&activate_err,
|
||||
crate::extensions::ExtensionError::AuthRequired
|
||||
);
|
||||
|
||||
if !needs_auth {
|
||||
return Ok(Json(ActionResponse::fail(err_str)));
|
||||
return Ok(Json(ActionResponse::fail(activate_err.to_string())));
|
||||
}
|
||||
|
||||
// Activation failed due to auth; try authenticating first.
|
||||
@@ -2481,6 +2690,7 @@ mod tests {
|
||||
skill_catalog: None,
|
||||
scheduler: None,
|
||||
chat_rate_limiter: RateLimiter::new(30, 60),
|
||||
oauth_rate_limiter: RateLimiter::new(10, 60),
|
||||
registry_entries: vec![],
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
@@ -2803,4 +3013,177 @@ mod tests {
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
|
||||
// --- Slack relay OAuth CSRF tests ---
|
||||
|
||||
fn test_relay_oauth_router(state: Arc<GatewayState>) -> Router {
|
||||
Router::new()
|
||||
.route(
|
||||
"/oauth/slack/callback",
|
||||
get(slack_relay_oauth_callback_handler),
|
||||
)
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
fn test_secrets_store() -> Arc<dyn crate::secrets::SecretsStore + Send + Sync> {
|
||||
Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
|
||||
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
|
||||
"test-key-at-least-32-chars-long!!".to_string(),
|
||||
))
|
||||
.expect("crypto"),
|
||||
)))
|
||||
}
|
||||
|
||||
fn test_ext_mgr(
|
||||
secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync>,
|
||||
) -> Arc<ExtensionManager> {
|
||||
let tool_registry = Arc::new(ToolRegistry::new());
|
||||
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
|
||||
let mcp_pm = Arc::new(crate::tools::mcp::process::McpProcessManager::new());
|
||||
Arc::new(ExtensionManager::new(
|
||||
mcp_sm,
|
||||
mcp_pm,
|
||||
secrets,
|
||||
tool_registry,
|
||||
None,
|
||||
None,
|
||||
std::path::PathBuf::from("/tmp/wasm_tools"),
|
||||
std::path::PathBuf::from("/tmp/wasm_channels"),
|
||||
None,
|
||||
"test".to_string(),
|
||||
None,
|
||||
vec![],
|
||||
))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_relay_oauth_callback_missing_state_param() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
let secrets = test_secrets_store();
|
||||
let ext_mgr = test_ext_mgr(secrets);
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let app = test_relay_oauth_router(state);
|
||||
|
||||
// Callback without state param should be rejected
|
||||
let req = axum::http::Request::builder()
|
||||
.uri("/oauth/slack/callback?stream_token=tok123&team_id=T123&provider=slack")
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
assert!(
|
||||
html.contains("Invalid or expired authorization"),
|
||||
"Expected CSRF error, got: {}",
|
||||
&html[..html.len().min(300)]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_relay_oauth_callback_wrong_state_param() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
let secrets = test_secrets_store();
|
||||
|
||||
// Store a valid nonce
|
||||
secrets
|
||||
.create(
|
||||
"test",
|
||||
crate::secrets::CreateSecretParams::new(
|
||||
format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME),
|
||||
"correct-nonce-value",
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("store nonce");
|
||||
|
||||
let ext_mgr = test_ext_mgr(secrets);
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let app = test_relay_oauth_router(state);
|
||||
|
||||
// Callback with wrong state param
|
||||
let req = axum::http::Request::builder()
|
||||
.uri("/oauth/slack/callback?stream_token=tok123&team_id=T123&provider=slack&state=wrong-nonce")
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
assert!(
|
||||
html.contains("Invalid or expired authorization"),
|
||||
"Expected CSRF error for wrong nonce, got: {}",
|
||||
&html[..html.len().min(300)]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_relay_oauth_callback_correct_state_proceeds() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
let secrets = test_secrets_store();
|
||||
let nonce = "valid-test-nonce-12345";
|
||||
|
||||
// Store the correct nonce
|
||||
secrets
|
||||
.create(
|
||||
"test",
|
||||
crate::secrets::CreateSecretParams::new(
|
||||
format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME),
|
||||
nonce,
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("store nonce");
|
||||
|
||||
let ext_mgr = test_ext_mgr(secrets.clone());
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let app = test_relay_oauth_router(state);
|
||||
|
||||
// Callback with correct state param — will pass CSRF check
|
||||
// but may fail downstream (no real relay service) — that's OK,
|
||||
// we just verify it doesn't return a CSRF error.
|
||||
let req = axum::http::Request::builder()
|
||||
.uri(format!(
|
||||
"/oauth/slack/callback?stream_token=tok123&team_id=T123&provider=slack&state={}",
|
||||
nonce
|
||||
))
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
// Should NOT contain the CSRF error message
|
||||
assert!(
|
||||
!html.contains("Invalid or expired authorization"),
|
||||
"Should have passed CSRF check, got: {}",
|
||||
&html[..html.len().min(300)]
|
||||
);
|
||||
|
||||
// Verify the nonce was consumed (deleted)
|
||||
let state_key = format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME);
|
||||
let exists = secrets.exists("test", &state_key).await.unwrap_or(true);
|
||||
assert!(!exists, "CSRF nonce should be deleted after use");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2350,8 +2350,8 @@ function renderExtensionCard(ext) {
|
||||
activeLabel.textContent = ext.active ? 'Active' : 'Installed';
|
||||
actions.appendChild(activeLabel);
|
||||
|
||||
// MCP servers may be installed but inactive — show Activate button
|
||||
if (ext.kind === 'mcp_server' && !ext.active) {
|
||||
// MCP servers and channel-relay extensions may be installed but inactive — show Activate button
|
||||
if ((ext.kind === 'mcp_server' || ext.kind === 'channel_relay') && !ext.active) {
|
||||
const activateBtn = document.createElement('button');
|
||||
activateBtn.className = 'btn-ext activate';
|
||||
activateBtn.textContent = 'Activate';
|
||||
|
||||
@@ -82,6 +82,7 @@ impl TestGatewayBuilder {
|
||||
skill_catalog: None,
|
||||
scheduler: None,
|
||||
chat_rate_limiter: RateLimiter::new(30, 60),
|
||||
oauth_rate_limiter: RateLimiter::new(10, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
|
||||
@@ -509,6 +509,7 @@ mod tests {
|
||||
skill_registry: None,
|
||||
skill_catalog: None,
|
||||
chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60),
|
||||
oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
|
||||
Reference in New Issue
Block a user