diff --git a/devops/docker/compose/.env.example b/devops/docker/compose/.env.example index 2d4d9cad9..f2cf2b858 100644 --- a/devops/docker/compose/.env.example +++ b/devops/docker/compose/.env.example @@ -50,6 +50,7 @@ TROGON_GATEWAY_LOCAL_GITHUB_WEBHOOK_SECRET=local-dev-secret # --- Slack Source --- # SLACK_PRIMARY_SIGNING_SECRET= +# SLACK_PRIMARY_APP_TOKEN= # --- Logging --- # RUST_LOG=info diff --git a/devops/docker/compose/services/trogon-gateway/README.md b/devops/docker/compose/services/trogon-gateway/README.md index 6263ccdb1..6aafa19a9 100644 --- a/devops/docker/compose/services/trogon-gateway/README.md +++ b/devops/docker/compose/services/trogon-gateway/README.md @@ -84,6 +84,23 @@ bot_token = { env = "DISCORD_BOT_TOKEN" } Discord does not use the HTTP ingress or ngrok. It opens an outbound WebSocket connection to Discord and publishes every gateway event to NATS. +## Slack Socket Mode + +Slack integrations can use either the existing HTTP webhook transport or Socket +Mode. Configure exactly one transport per integration. Socket Mode does not use +the HTTP ingress or ngrok; it opens an outbound WebSocket to Slack and publishes +Events API payloads, interactive payloads, and slash commands to the same NATS +subjects as the webhook transport. + +```toml +[sources.slack.integrations.primary] +subject_prefix = "slack-primary" +stream_name = "SLACK_PRIMARY" + +[sources.slack.integrations.primary.socket_mode] +app_token = { env = "SLACK_PRIMARY_APP_TOKEN" } +``` + ## Telegram webhooks Set `webhook_secret` under `[sources.telegram.integrations..webhook]`. To let diff --git a/devops/docker/compose/services/trogon-gateway/gateway.toml b/devops/docker/compose/services/trogon-gateway/gateway.toml index 996e50c4a..06b10e561 100644 --- a/devops/docker/compose/services/trogon-gateway/gateway.toml +++ b/devops/docker/compose/services/trogon-gateway/gateway.toml @@ -119,3 +119,6 @@ webhook_secret = { env = "TROGON_GATEWAY_LOCAL_GITHUB_WEBHOOK_SECRET" } # [sources.slack.integrations.primary.webhook] # signing_secret = { env = "SLACK_PRIMARY_SIGNING_SECRET" } # timestamp_max_drift_secs = 300 +# +# [sources.slack.integrations.primary.socket_mode] +# app_token = { env = "SLACK_PRIMARY_APP_TOKEN" } diff --git a/rsworkspace/Cargo.lock b/rsworkspace/Cargo.lock index c5eeb37ff..970af586a 100644 --- a/rsworkspace/Cargo.lock +++ b/rsworkspace/Cargo.lock @@ -2774,7 +2774,11 @@ checksum = "8f72a05e828585856dacd553fba484c242c46e391fb0e58917c942ee9202915c" dependencies = [ "futures-util", "log", + "rustls", + "rustls-native-certs", + "rustls-pki-types", "tokio", + "tokio-rustls", "tungstenite", ] @@ -3130,6 +3134,7 @@ dependencies = [ "tempfile", "time", "tokio", + "tokio-tungstenite", "tower", "tracing", "tracing-subscriber", @@ -3248,6 +3253,8 @@ dependencies = [ "httparse", "log", "rand 0.9.2", + "rustls", + "rustls-pki-types", "sha1", "thiserror 2.0.18", ] diff --git a/rsworkspace/crates/trogon-gateway/Cargo.toml b/rsworkspace/crates/trogon-gateway/Cargo.toml index a8903960d..4867dbaff 100644 --- a/rsworkspace/crates/trogon-gateway/Cargo.toml +++ b/rsworkspace/crates/trogon-gateway/Cargo.toml @@ -29,6 +29,7 @@ serde_json = { workspace = true } sha2 = "0.11" subtle = "2.6" tokio = { workspace = true, features = ["full"] } +tokio-tungstenite = { workspace = true, features = ["rustls-tls-native-roots"] } tracing = { workspace = true } twilight-gateway = { workspace = true } twilight-model = { workspace = true } diff --git a/rsworkspace/crates/trogon-gateway/src/config.rs b/rsworkspace/crates/trogon-gateway/src/config.rs index 4aeac9a0b..1c69f541b 100644 --- a/rsworkspace/crates/trogon-gateway/src/config.rs +++ b/rsworkspace/crates/trogon-gateway/src/config.rs @@ -13,7 +13,10 @@ use crate::source::linear::config::LinearWebhookSecret; use crate::source::microsoft_graph::MicrosoftGraphClientState; use crate::source::notion::NotionVerificationToken; use crate::source::sentry::SentryClientSecret; -use crate::source::slack::config::SlackSigningSecret; +use crate::source::slack::config::{ + SlackAppToken, SlackSigningSecret, SlackSocketModeConfig as SlackSocketModeSourceConfig, SlackTransportConfig, + SlackWebhookConfig as SlackWebhookSourceConfig, +}; use crate::source::telegram::config::{ TelegramBotToken, TelegramPublicWebhookUrl, TelegramWebhookRegistrationConfig, TelegramWebhookSecret, }; @@ -70,6 +73,17 @@ impl std::error::Error for DurationTooLong {} const SENTRY_MAX_ACK_TIMEOUT_SECS: u64 = 1; +#[derive(Debug)] +struct SlackTransportConflict; + +impl fmt::Display for SlackTransportConflict { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("configure exactly one of webhook or socket_mode") + } +} + +impl std::error::Error for SlackTransportConflict {} + #[derive(Clone, Debug, serde::Deserialize)] #[serde(untagged)] enum SecretInput { @@ -434,7 +448,7 @@ struct DiscordConfig { struct SlackConfig { status: Option, #[config(default = {})] - integrations: BTreeMap>, + integrations: BTreeMap, } #[derive(Config)] @@ -512,6 +526,18 @@ struct SourceIntegrationInput { webhook: Option, } +#[derive(serde::Deserialize)] +#[serde(deny_unknown_fields)] +struct SlackIntegrationInput { + status: Option, + subject_prefix: Option, + stream_name: Option, + stream_max_age_secs: Option, + nats_ack_timeout_secs: Option, + webhook: Option, + socket_mode: Option, +} + #[derive(serde::Deserialize)] #[serde(deny_unknown_fields)] struct GithubWebhookConfig { @@ -525,6 +551,12 @@ struct SlackWebhookConfig { timestamp_max_drift_secs: Option, } +#[derive(serde::Deserialize)] +#[serde(deny_unknown_fields)] +struct SlackSocketModeConfig { + app_token: Option, +} + #[derive(serde::Deserialize)] #[serde(deny_unknown_fields)] struct TelegramWebhookConfig { @@ -877,25 +909,6 @@ fn resolve_slack_integrations( if !resolve_integration_source_status("slack", &id, integration.status.as_deref(), errors) { continue; } - let Some(webhook) = integration.webhook else { - continue; - }; - let Some(secret) = require_integration_value("slack", &id, "signing_secret", webhook.signing_secret, errors) - else { - continue; - }; - let signing_secret = match SlackSigningSecret::new(secret) { - Ok(secret) => secret, - Err(error) => { - errors.push(ConfigValidationError::invalid_integration( - "slack", - &id, - "signing_secret", - error, - )); - continue; - } - }; let Some((subject_prefix, stream_name, stream_max_age, nats_ack_timeout)) = resolve_common_integration_fields( CommonIntegrationFieldsInput { source: "slack", @@ -912,37 +925,111 @@ fn resolve_slack_integrations( ) else { continue; }; - let timestamp_max_drift = match NonZeroDuration::from_secs( - webhook - .timestamp_max_drift_secs - .unwrap_or(DEFAULT_SLACK_TIMESTAMP_MAX_DRIFT_SECS), - ) { - Ok(duration) => duration, - Err(error) => { - errors.push(ConfigValidationError::invalid_integration( - "slack", - &id, - "timestamp_max_drift_secs", - error, - )); - continue; - } + let transport = match resolve_slack_transport(&id, integration.webhook, integration.socket_mode, errors) { + Some(transport) => transport, + None => continue, }; integrations.push(SourceIntegration::new( id, crate::source::slack::SlackConfig { - signing_secret, subject_prefix, stream_name, stream_max_age, nats_ack_timeout, - timestamp_max_drift, + transport, }, )); } integrations } +fn resolve_slack_transport( + id: &SourceIntegrationId, + webhook: Option, + socket_mode: Option, + errors: &mut Vec, +) -> Option { + match (webhook, socket_mode) { + (Some(webhook), None) => resolve_slack_webhook_transport(id, webhook, errors), + (None, Some(socket_mode)) => resolve_slack_socket_mode_transport(id, socket_mode, errors), + (Some(_), Some(_)) => { + errors.push(ConfigValidationError::invalid_integration( + "slack", + id, + "transport", + SlackTransportConflict, + )); + None + } + (None, None) => None, + } +} + +fn resolve_slack_webhook_transport( + id: &SourceIntegrationId, + webhook: SlackWebhookConfig, + errors: &mut Vec, +) -> Option { + let secret = require_integration_value("slack", id, "signing_secret", webhook.signing_secret, errors)?; + let signing_secret = match SlackSigningSecret::new(secret) { + Ok(secret) => secret, + Err(error) => { + errors.push(ConfigValidationError::invalid_integration( + "slack", + id, + "signing_secret", + error, + )); + return None; + } + }; + let timestamp_max_drift = match NonZeroDuration::from_secs( + webhook + .timestamp_max_drift_secs + .unwrap_or(DEFAULT_SLACK_TIMESTAMP_MAX_DRIFT_SECS), + ) { + Ok(duration) => duration, + Err(error) => { + errors.push(ConfigValidationError::invalid_integration( + "slack", + id, + "timestamp_max_drift_secs", + error, + )); + return None; + } + }; + + Some(SlackTransportConfig::Webhook(SlackWebhookSourceConfig { + signing_secret, + timestamp_max_drift, + })) +} + +fn resolve_slack_socket_mode_transport( + id: &SourceIntegrationId, + socket_mode: SlackSocketModeConfig, + errors: &mut Vec, +) -> Option { + let token = require_integration_value("slack", id, "app_token", socket_mode.app_token, errors)?; + let app_token = match SlackAppToken::new(token) { + Ok(token) => token, + Err(error) => { + errors.push(ConfigValidationError::invalid_integration( + "slack", + id, + "app_token", + error, + )); + return None; + } + }; + + Some(SlackTransportConfig::SocketMode(SlackSocketModeSourceConfig { + app_token, + })) +} + fn resolve_telegram_integrations( section: TelegramConfig, errors: &mut Vec, @@ -1869,6 +1956,15 @@ signing_secret = "{secret}" ) } + fn slack_socket_mode_toml(token: &str) -> String { + format!( + r#" +[sources.slack.integrations.primary.socket_mode] +app_token = "{token}" +"# + ) + } + fn telegram_toml(secret: &str) -> String { format!( r#" @@ -2274,6 +2370,68 @@ TROGON_SOURCE_DISCORD_BOT_TOKEN = "Bot my-bot-token" let f = write_toml(&slack_toml("slack-signing-secret")); let cfg = load(Some(f.path())).expect("load failed"); assert!(!cfg.slack.is_empty()); + assert!(cfg.slack[0].config.webhook().is_some()); + assert!(cfg.slack[0].config.socket_mode().is_none()); + } + + #[test] + fn slack_socket_mode_resolves_with_valid_app_token() { + let f = write_toml(&slack_socket_mode_toml("xapp-test-token")); + let cfg = load(Some(f.path())).expect("load failed"); + assert!(!cfg.slack.is_empty()); + assert!(cfg.slack[0].config.webhook().is_none()); + assert!(cfg.slack[0].config.socket_mode().is_some()); + } + + #[test] + fn slack_socket_mode_missing_app_token_is_invalid() { + let toml = r#" +[sources.slack.integrations.primary.socket_mode] +"#; + let f = write_toml(toml); + let result = load(Some(f.path())); + assert!( + matches!(result, Err(ConfigError::Validation(ref errs)) if errs.iter().any(|e| e.contains("slack/primary: missing app_token"))) + ); + } + + #[test] + fn slack_disabled_socket_mode_integration_is_skipped() { + let toml = r#" +[sources.slack.integrations.primary] +status = "disabled" + +[sources.slack.integrations.primary.socket_mode] +app_token = "xapp-test-token" +"#; + let f = write_toml(toml); + let cfg = load(Some(f.path())).expect("load failed"); + assert!(cfg.slack.is_empty()); + } + + #[test] + fn slack_socket_mode_rejects_non_app_token() { + let f = write_toml(&slack_socket_mode_toml("xoxb-bot-token")); + let result = load(Some(f.path())); + assert!( + matches!(result, Err(ConfigError::Validation(ref errs)) if errs.iter().any(|e| e.contains("slack/primary: invalid app_token: must start with xapp-"))) + ); + } + + #[test] + fn slack_integration_rejects_webhook_and_socket_mode_together() { + let toml = r#" +[sources.slack.integrations.primary.webhook] +signing_secret = "slack-secret" + +[sources.slack.integrations.primary.socket_mode] +app_token = "xapp-test-token" +"#; + let f = write_toml(toml); + let result = load(Some(f.path())); + assert!( + matches!(result, Err(ConfigError::Validation(ref errs)) if errs.iter().any(|e| e.contains("slack/primary: invalid transport: configure exactly one of webhook or socket_mode"))) + ); } #[test] diff --git a/rsworkspace/crates/trogon-gateway/src/http.rs b/rsworkspace/crates/trogon-gateway/src/http.rs index 470853c1e..6e71e6e1a 100644 --- a/rsworkspace/crates/trogon-gateway/src/http.rs +++ b/rsworkspace/crates/trogon-gateway/src/http.rs @@ -27,14 +27,23 @@ where publisher.clone(), |p, cfg| crate::source::github::router(p, cfg), ); - app = mount_webhook_integrations( - app, - "slack", - "/sources/slack", - &config.slack, - publisher.clone(), - |p, cfg| crate::source::slack::router(p, cfg), - ); + for integration in &config.slack { + if integration.config.webhook().is_none() { + continue; + } + let path = format!("/sources/slack/{}", integration.id); + app = app.nest( + &path, + crate::source::slack::router(publisher.clone(), &integration.config), + ); + let integration_id = integration.id.as_str(); + info!( + source = "slack", + integration = integration_id, + path, + "mounted source integration" + ); + } app = mount_webhook_integrations( app, "telegram", @@ -265,4 +274,28 @@ webhook_secret = "other-secret" assert_eq!(messages.len(), 1); assert_eq!(messages[0].subject, "github-acme-main.push"); } + + #[tokio::test] + async fn slack_socket_mode_integration_does_not_mount_webhook_route() { + let toml = r#" +[sources.slack.integrations.primary.socket_mode] +app_token = "xapp-test-token" +"#; + let f = write_toml(toml); + let cfg = load(Some(f.path())).expect("load failed"); + let app = mount_sources(cfg, wrap_publisher(MockJetStreamPublisher::new())); + + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/sources/slack/primary/webhook") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::NOT_FOUND); + } } diff --git a/rsworkspace/crates/trogon-gateway/src/main.rs b/rsworkspace/crates/trogon-gateway/src/main.rs index b92129412..732935ffa 100644 --- a/rsworkspace/crates/trogon-gateway/src/main.rs +++ b/rsworkspace/crates/trogon-gateway/src/main.rs @@ -168,6 +168,30 @@ async fn serve(resolved: config::ResolvedConfig) -> Result<(), Box Ok(()), + result = crate::source::slack::socket_mode::run(p, &slack_cfg) => { + result.map_err(|error| error.to_string()) + } + }; + ("slack-socket-mode", result) + }); + info!( + source = "slack", + integration = %integration_id, + "socket mode runner spawned" + ); + } + let app = trogon_std::telemetry::http::instrument_router(http::mount_sources(resolved, publisher)); let addr = SocketAddr::from(([0, 0, 0, 0], port)); diff --git a/rsworkspace/crates/trogon-gateway/src/source/slack/config.rs b/rsworkspace/crates/trogon-gateway/src/source/slack/config.rs index 8826c04e6..890104d08 100644 --- a/rsworkspace/crates/trogon-gateway/src/source/slack/config.rs +++ b/rsworkspace/crates/trogon-gateway/src/source/slack/config.rs @@ -23,18 +23,99 @@ impl fmt::Debug for SlackSigningSecret { } } -pub struct SlackConfig { +#[derive(Debug)] +pub enum SlackAppTokenError { + Empty(EmptySecret), + MissingPrefix, +} + +impl fmt::Display for SlackAppTokenError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Empty(error) => write!(f, "{error}"), + Self::MissingPrefix => f.write_str("must start with xapp-"), + } + } +} + +impl std::error::Error for SlackAppTokenError { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + match self { + Self::Empty(error) => Some(error), + Self::MissingPrefix => None, + } + } +} + +#[derive(Clone)] +pub struct SlackAppToken(SecretString); + +impl SlackAppToken { + pub fn new(s: impl AsRef) -> Result { + let secret = SecretString::new(s).map_err(SlackAppTokenError::Empty)?; + if !secret.as_str().starts_with("xapp-") { + return Err(SlackAppTokenError::MissingPrefix); + } + Ok(Self(secret)) + } + + pub fn as_str(&self) -> &str { + self.0.as_str() + } +} + +impl fmt::Debug for SlackAppToken { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("SlackAppToken(****)") + } +} + +#[derive(Clone)] +pub struct SlackWebhookConfig { pub signing_secret: SlackSigningSecret, + pub timestamp_max_drift: NonZeroDuration, +} + +#[derive(Clone)] +pub struct SlackSocketModeConfig { + pub app_token: SlackAppToken, +} + +#[derive(Clone)] +pub enum SlackTransportConfig { + Webhook(SlackWebhookConfig), + SocketMode(SlackSocketModeConfig), +} + +#[derive(Clone)] +pub struct SlackConfig { pub subject_prefix: NatsToken, pub stream_name: NatsToken, pub stream_max_age: StreamMaxAge, pub nats_ack_timeout: NonZeroDuration, - pub timestamp_max_drift: NonZeroDuration, + pub transport: SlackTransportConfig, +} + +impl SlackConfig { + pub fn webhook(&self) -> Option<&SlackWebhookConfig> { + match &self.transport { + SlackTransportConfig::Webhook(config) => Some(config), + SlackTransportConfig::SocketMode(_) => None, + } + } + + pub fn socket_mode(&self) -> Option<&SlackSocketModeConfig> { + match &self.transport { + SlackTransportConfig::Webhook(_) => None, + SlackTransportConfig::SocketMode(config) => Some(config), + } + } } #[cfg(test)] mod tests { use super::*; + use std::error::Error; #[test] fn slack_signing_secret_roundtrips() { @@ -47,4 +128,30 @@ mod tests { let secret = SlackSigningSecret::new("super-secret").unwrap(); assert_eq!(format!("{secret:?}"), "SlackSigningSecret(****)"); } + + #[test] + fn slack_app_token_roundtrips() { + let token = SlackAppToken::new("xapp-test-token").unwrap(); + assert_eq!(token.as_str(), "xapp-test-token"); + } + + #[test] + fn slack_app_token_debug_redacts() { + let token = SlackAppToken::new("xapp-test-token").unwrap(); + assert_eq!(format!("{token:?}"), "SlackAppToken(****)"); + } + + #[test] + fn slack_app_token_requires_app_prefix() { + let error = SlackAppToken::new("xoxb-not-app-token").unwrap_err(); + assert_eq!(error.to_string(), "must start with xapp-"); + assert!(error.source().is_none()); + } + + #[test] + fn slack_app_token_rejects_empty_token() { + let error = SlackAppToken::new("").unwrap_err(); + assert_eq!(error.to_string(), "secret must not be empty"); + assert!(error.source().is_some()); + } } diff --git a/rsworkspace/crates/trogon-gateway/src/source/slack/mod.rs b/rsworkspace/crates/trogon-gateway/src/source/slack/mod.rs index 9d29a0d1f..eecc3c6d8 100644 --- a/rsworkspace/crates/trogon-gateway/src/source/slack/mod.rs +++ b/rsworkspace/crates/trogon-gateway/src/source/slack/mod.rs @@ -1,13 +1,15 @@ //! # trogon-source-slack //! -//! Slack Events API webhook receiver that publishes events to NATS JetStream. +//! Slack Events API receiver that publishes events to NATS JetStream. //! //! ## How it works //! -//! 1. Slack sends `POST /webhook` with `X-Slack-Signature` and -//! `X-Slack-Request-Timestamp` headers plus a JSON payload. -//! 2. The server validates the HMAC-SHA256 signature against `SLACK_SIGNING_SECRET`. -//! 3. `url_verification` challenges are answered inline (no NATS publish). +//! 1. Slack delivers payloads by HTTP webhook or Socket Mode, depending on the +//! configured integration transport. +//! 2. HTTP webhooks are validated with HMAC-SHA256 signatures. Socket Mode uses +//! Slack's pre-authenticated WebSocket connection and acknowledges envelopes +//! after durable NATS publish. +//! 3. HTTP `url_verification` challenges are answered inline (no NATS publish). //! 4. `event_callback` payloads are published to NATS JetStream on //! `slack.{event.type}` subjects (e.g. `slack.message`, `slack.app_mention`). //! 5. The JetStream stream (`SLACK` by default, capturing `slack.>`) is created @@ -37,6 +39,7 @@ pub mod config; pub mod constants; pub mod server; pub mod signature; +pub mod socket_mode; pub use config::SlackConfig; pub use server::{provision, router}; diff --git a/rsworkspace/crates/trogon-gateway/src/source/slack/server.rs b/rsworkspace/crates/trogon-gateway/src/source/slack/server.rs index c3c2aef88..7a0e76333 100644 --- a/rsworkspace/crates/trogon-gateway/src/source/slack/server.rs +++ b/rsworkspace/crates/trogon-gateway/src/source/slack/server.rs @@ -3,6 +3,7 @@ use std::time::Duration; use super::config::SlackConfig; use super::config::SlackSigningSecret; +use super::config::SlackWebhookConfig; use super::constants::{ CONTENT_TYPE_FORM, HEADER_SIGNATURE, HEADER_TIMESTAMP, HTTP_BODY_SIZE_MAX, NATS_HEADER_EVENT_ID, NATS_HEADER_EVENT_TYPE, NATS_HEADER_PAYLOAD_KIND, NATS_HEADER_REJECT_REASON, NATS_HEADER_TEAM_ID, @@ -34,34 +35,293 @@ fn outcome_to_status(outcome: PublishOutcome) -> StatusCode } } -async fn publish_unroutable( - publisher: &ClaimCheckPublisher, - subject_prefix: &NatsToken, - reason: &str, - body: Bytes, - ack_timeout: NonZeroDuration, -) { - let subject = format!("{}.unroutable", subject_prefix); - let mut headers = async_nats::HeaderMap::new(); - headers.insert(NATS_HEADER_REJECT_REASON, reason); - headers.insert(NATS_HEADER_PAYLOAD_KIND, "unroutable"); - - let outcome = publisher - .publish_event(subject, headers, body, ack_timeout.into()) - .await; - outcome.log_on_error("slack.unroutable"); +#[derive(Clone)] +pub struct SlackBridge { + publisher: ClaimCheckPublisher, + subject_prefix: NatsToken, + nats_ack_timeout: NonZeroDuration, } #[derive(Clone)] struct AppState { - publisher: ClaimCheckPublisher, + bridge: SlackBridge, clock: C, signing_secret: SlackSigningSecret, - subject_prefix: NatsToken, - nats_ack_timeout: NonZeroDuration, timestamp_max_drift: NonZeroDuration, } +impl SlackBridge { + pub fn new(publisher: ClaimCheckPublisher, config: &SlackConfig) -> Self { + Self { + publisher, + subject_prefix: config.subject_prefix.clone(), + nats_ack_timeout: config.nats_ack_timeout, + } + } + + pub async fn handle_json_body(&self, body: Bytes) -> (StatusCode, String) { + let Ok(payload) = serde_json::from_slice::(&body) else { + warn!("Invalid JSON payload"); + self.publish_unroutable("invalid_json", body).await; + return (StatusCode::BAD_REQUEST, String::new()); + }; + + self.handle_json_payload(&body, &payload).await + } + + pub async fn handle_socket_interaction( + &self, + payload: &serde_json::Value, + ) -> Result { + let payload_json = serde_json::to_string(payload)?; + let status = self + .handle_interaction(&payload_json, Bytes::from(payload_json.clone())) + .await; + Ok(status) + } + + pub async fn handle_socket_slash_command( + &self, + payload: &serde_json::Value, + ) -> Result { + let raw_body = Bytes::from(serde_json::to_vec(payload)?); + Ok(self.handle_slash_command_value(payload, raw_body).await) + } + + pub async fn publish_socket_unroutable(&self, reason: &str, body: Bytes) { + self.publish_unroutable(reason, body).await; + } + + async fn publish_unroutable(&self, reason: &str, body: Bytes) { + let subject = format!("{}.unroutable", self.subject_prefix); + let mut headers = async_nats::HeaderMap::new(); + headers.insert(NATS_HEADER_REJECT_REASON, reason); + headers.insert(NATS_HEADER_PAYLOAD_KIND, "unroutable"); + + let outcome = self + .publisher + .publish_event(subject, headers, body, self.nats_ack_timeout.into()) + .await; + outcome.log_on_error("slack.unroutable"); + } + + async fn handle_json_payload(&self, body: &Bytes, payload: &serde_json::Value) -> (StatusCode, String) { + let payload_type = payload.get("type").and_then(|v| v.as_str()).unwrap_or_default(); + + if payload_type == "url_verification" { + let challenge = payload + .get("challenge") + .and_then(|v| v.as_str()) + .unwrap_or_default() + .to_owned(); + info!("Responding to Slack URL verification challenge"); + return (StatusCode::OK, challenge); + } + + if payload_type != "event_callback" { + warn!(payload_type, "Unhandled payload type"); + self.publish_unroutable("unhandled_payload_type", body.clone()).await; + return (StatusCode::OK, String::new()); + } + + let Some(event_type) = payload + .get("event") + .and_then(|v| v.get("type")) + .and_then(|v| v.as_str()) + else { + warn!("Missing event.type in event_callback payload"); + self.publish_unroutable("missing_event_type", body.clone()).await; + return (StatusCode::BAD_REQUEST, String::new()); + }; + let event_type = event_type.to_owned(); + + let Some(event_id) = payload.get("event_id").and_then(|v| v.as_str()) else { + warn!("Missing event_id in event_callback payload"); + self.publish_unroutable("missing_event_id", body.clone()).await; + return (StatusCode::BAD_REQUEST, String::new()); + }; + let event_id = event_id.to_owned(); + + let team_id = payload + .get("team_id") + .and_then(|v| v.as_str()) + .unwrap_or("unknown") + .to_owned(); + + let subject = format!("{}.event.{}", self.subject_prefix, event_type); + + let span = tracing::Span::current(); + span.record("event_type", &event_type); + span.record("event_id", &event_id); + span.record("subject", &subject); + + let mut nats_headers = async_nats::HeaderMap::new(); + nats_headers.insert(async_nats::header::NATS_MESSAGE_ID, event_id.as_str()); + nats_headers.insert(NATS_HEADER_EVENT_TYPE, event_type.as_str()); + nats_headers.insert(NATS_HEADER_EVENT_ID, event_id.as_str()); + nats_headers.insert(NATS_HEADER_TEAM_ID, team_id.as_str()); + nats_headers.insert(NATS_HEADER_PAYLOAD_KIND, "event"); + + let outcome = self + .publisher + .publish_event(subject, nats_headers, body.clone(), self.nats_ack_timeout.into()) + .await; + + (outcome_to_status(outcome), String::new()) + } + + async fn handle_form_payload(&self, body: &Bytes) -> (StatusCode, String) { + let form_str = match std::str::from_utf8(body) { + Ok(s) => s, + Err(_) => { + warn!("Invalid UTF-8 in form payload"); + self.publish_unroutable("invalid_utf8_form", body.clone()).await; + return (StatusCode::BAD_REQUEST, String::new()); + } + }; + + let fields: Vec<(String, String)> = form_urlencoded::parse(form_str.as_bytes()) + .map(|(k, v)| (k.into_owned(), v.into_owned())) + .collect(); + + let find_field = |name: &str| fields.iter().find(|(k, _)| k == name).map(|(_, v)| v.as_str()); + + if let Some(payload_json) = find_field("payload") { + let status = self.handle_interaction(payload_json, body.clone()).await; + return (status, String::new()); + } + + if let Some(command) = find_field("command") { + let status = self.handle_slash_command(command, &fields, body.clone()).await; + return (status, String::new()); + } + + warn!("Unrecognized form payload"); + self.publish_unroutable("unrecognized_form", body.clone()).await; + (StatusCode::BAD_REQUEST, String::new()) + } + + async fn handle_interaction(&self, payload_json: &str, raw_body: Bytes) -> StatusCode { + let Ok(payload) = serde_json::from_str::(payload_json) else { + warn!("Invalid JSON in interaction payload field"); + self.publish_unroutable("invalid_interaction_json", raw_body).await; + return StatusCode::BAD_REQUEST; + }; + + self.publish_interaction_payload(&payload, Bytes::from(payload_json.to_owned())) + .await + } + + async fn publish_interaction_payload(&self, payload: &serde_json::Value, body: Bytes) -> StatusCode { + let Some(interaction_type) = payload.get("type").and_then(|v| v.as_str()) else { + warn!("Missing type in interaction payload"); + self.publish_unroutable("missing_interaction_type", body).await; + return StatusCode::BAD_REQUEST; + }; + let interaction_type = interaction_type.to_owned(); + + let Some(trigger_id) = payload.get("trigger_id").and_then(|v| v.as_str()) else { + warn!("Missing trigger_id in interaction payload"); + self.publish_unroutable("missing_interaction_trigger_id", body).await; + return StatusCode::BAD_REQUEST; + }; + let trigger_id = trigger_id.to_owned(); + + let team_id = payload + .get("team") + .and_then(|v| v.get("id")) + .and_then(|v| v.as_str()) + .unwrap_or("unknown") + .to_owned(); + + let subject = format!("{}.interaction.{}", self.subject_prefix, interaction_type); + + let span = tracing::Span::current(); + span.record("event_type", &interaction_type); + span.record("event_id", &trigger_id); + span.record("subject", &subject); + + info!(interaction_type, "Received Slack interaction"); + + let mut nats_headers = async_nats::HeaderMap::new(); + nats_headers.insert(async_nats::header::NATS_MESSAGE_ID, trigger_id.as_str()); + nats_headers.insert(NATS_HEADER_EVENT_TYPE, interaction_type.as_str()); + nats_headers.insert(NATS_HEADER_TEAM_ID, team_id.as_str()); + nats_headers.insert(NATS_HEADER_PAYLOAD_KIND, "interaction"); + + let outcome = self + .publisher + .publish_event(subject, nats_headers, body, self.nats_ack_timeout.into()) + .await; + + outcome_to_status(outcome) + } + + async fn handle_slash_command(&self, command: &str, fields: &[(String, String)], raw_body: Bytes) -> StatusCode { + let team_id = fields + .iter() + .find(|(k, _)| k == "team_id") + .map(|(_, v)| v.as_str()) + .unwrap_or("unknown"); + + let Some(trigger_id) = fields.iter().find(|(k, _)| k == "trigger_id").map(|(_, v)| v.as_str()) else { + warn!(command, "Missing trigger_id in slash command payload"); + self.publish_unroutable("missing_command_trigger_id", raw_body).await; + return StatusCode::BAD_REQUEST; + }; + + self.publish_slash_command(command, team_id, trigger_id, raw_body).await + } + + async fn handle_slash_command_value(&self, payload: &serde_json::Value, raw_body: Bytes) -> StatusCode { + let command = payload.get("command").and_then(|v| v.as_str()).unwrap_or_default(); + if command.is_empty() { + warn!("Missing command in slash command payload"); + self.publish_unroutable("missing_command", raw_body).await; + return StatusCode::BAD_REQUEST; + } + let team_id = payload.get("team_id").and_then(|v| v.as_str()).unwrap_or("unknown"); + let Some(trigger_id) = payload.get("trigger_id").and_then(|v| v.as_str()) else { + warn!(command, "Missing trigger_id in slash command payload"); + self.publish_unroutable("missing_command_trigger_id", raw_body).await; + return StatusCode::BAD_REQUEST; + }; + + self.publish_slash_command(command, team_id, trigger_id, raw_body).await + } + + async fn publish_slash_command( + &self, + command: &str, + team_id: &str, + trigger_id: &str, + raw_body: Bytes, + ) -> StatusCode { + let command_name = command.trim_start_matches('/'); + let subject = format!("{}.command.{}", self.subject_prefix, command_name); + + let span = tracing::Span::current(); + span.record("event_type", command); + span.record("event_id", trigger_id); + span.record("subject", &subject); + + info!(command, "Received Slack slash command"); + + let mut nats_headers = async_nats::HeaderMap::new(); + nats_headers.insert(async_nats::header::NATS_MESSAGE_ID, trigger_id); + nats_headers.insert(NATS_HEADER_EVENT_TYPE, command); + nats_headers.insert(NATS_HEADER_TEAM_ID, team_id); + nats_headers.insert(NATS_HEADER_PAYLOAD_KIND, "command"); + + let outcome = self + .publisher + .publish_event(subject, nats_headers, raw_body, self.nats_ack_timeout.into()) + .await; + + outcome_to_status(outcome) + } +} + pub async fn provision(js: &C, config: &SlackConfig) -> Result<(), C::Error> { js.get_or_create_stream(async_nats::jetstream::stream::Config { name: config.stream_name.to_string(), @@ -81,21 +341,23 @@ pub fn router( publisher: ClaimCheckPublisher, config: &SlackConfig, ) -> Router { - router_with_clock(publisher, config, SystemClock) + let webhook = config + .webhook() + .expect("Slack webhook router requires webhook transport config"); + router_with_clock(publisher, config, webhook, SystemClock) } fn router_with_clock( publisher: ClaimCheckPublisher, config: &SlackConfig, + webhook: &SlackWebhookConfig, clock: C, ) -> Router { let state = AppState { - publisher, + bridge: SlackBridge::new(publisher, config), clock, - signing_secret: config.signing_secret.clone(), - subject_prefix: config.subject_prefix.clone(), - nats_ack_timeout: config.nats_ack_timeout, - timestamp_max_drift: config.timestamp_max_drift, + signing_secret: webhook.signing_secret.clone(), + timestamp_max_drift: webhook.timestamp_max_drift, }; Router::new() @@ -164,294 +426,12 @@ async fn handle_webhook_inner( - state: &AppState, - body: &Bytes, -) -> (StatusCode, String) { - let Ok(payload) = serde_json::from_slice::(body) else { - warn!("Invalid JSON payload"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "invalid_json", - body.clone(), - state.nats_ack_timeout, - ) - .await; - return (StatusCode::BAD_REQUEST, String::new()); - }; - - let payload_type = payload.get("type").and_then(|v| v.as_str()).unwrap_or_default(); - - if payload_type == "url_verification" { - let challenge = payload - .get("challenge") - .and_then(|v| v.as_str()) - .unwrap_or_default() - .to_owned(); - info!("Responding to Slack URL verification challenge"); - return (StatusCode::OK, challenge); - } - - if payload_type != "event_callback" { - warn!(payload_type, "Unhandled payload type"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "unhandled_payload_type", - body.clone(), - state.nats_ack_timeout, - ) - .await; - return (StatusCode::OK, String::new()); - } - - let Some(event_type) = payload - .get("event") - .and_then(|v| v.get("type")) - .and_then(|v| v.as_str()) - else { - warn!("Missing event.type in event_callback payload"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "missing_event_type", - body.clone(), - state.nats_ack_timeout, - ) - .await; - return (StatusCode::BAD_REQUEST, String::new()); - }; - let event_type = event_type.to_owned(); - - let Some(event_id) = payload.get("event_id").and_then(|v| v.as_str()) else { - warn!("Missing event_id in event_callback payload"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "missing_event_id", - body.clone(), - state.nats_ack_timeout, - ) - .await; - return (StatusCode::BAD_REQUEST, String::new()); - }; - let event_id = event_id.to_owned(); - - let team_id = payload - .get("team_id") - .and_then(|v| v.as_str()) - .unwrap_or("unknown") - .to_owned(); - - // TODO: File attachments (e.g. `message` events with `files` array) contain - // private URLs that require a bot token to download. A downstream NATS consumer - // should handle fetching and storing file content — the source stays a dumb pipe. - let subject = format!("{}.event.{}", state.subject_prefix, event_type); - - let span = tracing::Span::current(); - span.record("event_type", &event_type); - span.record("event_id", &event_id); - span.record("subject", &subject); - - let mut nats_headers = async_nats::HeaderMap::new(); - nats_headers.insert(async_nats::header::NATS_MESSAGE_ID, event_id.as_str()); - nats_headers.insert(NATS_HEADER_EVENT_TYPE, event_type.as_str()); - nats_headers.insert(NATS_HEADER_EVENT_ID, event_id.as_str()); - nats_headers.insert(NATS_HEADER_TEAM_ID, team_id.as_str()); - nats_headers.insert(NATS_HEADER_PAYLOAD_KIND, "event"); - - let outcome = state - .publisher - .publish_event(subject, nats_headers, body.clone(), state.nats_ack_timeout.into()) - .await; - - (outcome_to_status(outcome), String::new()) -} - -async fn handle_form_payload( - state: &AppState, - body: &Bytes, -) -> (StatusCode, String) { - let form_str = match std::str::from_utf8(body) { - Ok(s) => s, - Err(_) => { - warn!("Invalid UTF-8 in form payload"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "invalid_utf8_form", - body.clone(), - state.nats_ack_timeout, - ) - .await; - return (StatusCode::BAD_REQUEST, String::new()); - } - }; - - let fields: Vec<(String, String)> = form_urlencoded::parse(form_str.as_bytes()) - .map(|(k, v)| (k.into_owned(), v.into_owned())) - .collect(); - - let find_field = |name: &str| fields.iter().find(|(k, _)| k == name).map(|(_, v)| v.as_str()); - - if let Some(payload_json) = find_field("payload") { - return handle_interaction(state, payload_json, body).await; - } - - if let Some(command) = find_field("command") { - return handle_slash_command(state, command, &fields, body).await; - } - - warn!("Unrecognized form payload"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "unrecognized_form", - body.clone(), - state.nats_ack_timeout, - ) - .await; - (StatusCode::BAD_REQUEST, String::new()) -} - -async fn handle_interaction( - state: &AppState, - payload_json: &str, - raw_body: &Bytes, -) -> (StatusCode, String) { - let Ok(payload) = serde_json::from_str::(payload_json) else { - warn!("Invalid JSON in interaction payload field"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "invalid_interaction_json", - raw_body.clone(), - state.nats_ack_timeout, - ) - .await; - return (StatusCode::BAD_REQUEST, String::new()); - }; - - let Some(interaction_type) = payload.get("type").and_then(|v| v.as_str()) else { - warn!("Missing type in interaction payload"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "missing_interaction_type", - raw_body.clone(), - state.nats_ack_timeout, - ) - .await; - return (StatusCode::BAD_REQUEST, String::new()); - }; - let interaction_type = interaction_type.to_owned(); - - let Some(trigger_id) = payload.get("trigger_id").and_then(|v| v.as_str()) else { - warn!("Missing trigger_id in interaction payload"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "missing_interaction_trigger_id", - raw_body.clone(), - state.nats_ack_timeout, - ) - .await; - return (StatusCode::BAD_REQUEST, String::new()); - }; - let trigger_id = trigger_id.to_owned(); - - let team_id = payload - .get("team") - .and_then(|v| v.get("id")) - .and_then(|v| v.as_str()) - .unwrap_or("unknown") - .to_owned(); - - let subject = format!("{}.interaction.{}", state.subject_prefix, interaction_type); - - let span = tracing::Span::current(); - span.record("event_type", &interaction_type); - span.record("event_id", &trigger_id); - span.record("subject", &subject); - - info!(interaction_type, "Received Slack interaction"); - - let mut nats_headers = async_nats::HeaderMap::new(); - nats_headers.insert(async_nats::header::NATS_MESSAGE_ID, trigger_id.as_str()); - nats_headers.insert(NATS_HEADER_EVENT_TYPE, interaction_type.as_str()); - nats_headers.insert(NATS_HEADER_TEAM_ID, team_id.as_str()); - nats_headers.insert(NATS_HEADER_PAYLOAD_KIND, "interaction"); - - let outcome = state - .publisher - .publish_event( - subject, - nats_headers, - Bytes::from(payload_json.to_owned()), - state.nats_ack_timeout.into(), - ) - .await; - - (outcome_to_status(outcome), String::new()) -} - -async fn handle_slash_command( - state: &AppState, - command: &str, - fields: &[(String, String)], - raw_body: &Bytes, -) -> (StatusCode, String) { - let command_name = command.trim_start_matches('/'); - - let team_id = fields - .iter() - .find(|(k, _)| k == "team_id") - .map(|(_, v)| v.as_str()) - .unwrap_or("unknown"); - - let Some(trigger_id) = fields.iter().find(|(k, _)| k == "trigger_id").map(|(_, v)| v.as_str()) else { - warn!(command, "Missing trigger_id in slash command payload"); - publish_unroutable( - &state.publisher, - &state.subject_prefix, - "missing_command_trigger_id", - raw_body.clone(), - state.nats_ack_timeout, - ) - .await; - return (StatusCode::BAD_REQUEST, String::new()); - }; - - let subject = format!("{}.command.{}", state.subject_prefix, command_name); - - let span = tracing::Span::current(); - span.record("event_type", command); - span.record("event_id", trigger_id); - span.record("subject", &subject); - - info!(command, "Received Slack slash command"); - - let mut nats_headers = async_nats::HeaderMap::new(); - nats_headers.insert(async_nats::header::NATS_MESSAGE_ID, trigger_id); - nats_headers.insert(NATS_HEADER_EVENT_TYPE, command); - nats_headers.insert(NATS_HEADER_TEAM_ID, team_id); - nats_headers.insert(NATS_HEADER_PAYLOAD_KIND, "command"); - - let outcome = state - .publisher - .publish_event(subject, nats_headers, raw_body.clone(), state.nats_ack_timeout.into()) - .await; - - (outcome_to_status(outcome), String::new()) -} - #[cfg(test)] mod tests { use super::*; @@ -501,12 +481,14 @@ mod tests { fn test_config() -> SlackConfig { SlackConfig { - signing_secret: SlackSigningSecret::new(TEST_SECRET).unwrap(), subject_prefix: NatsToken::new("slack").unwrap(), stream_name: NatsToken::new("SLACK").unwrap(), stream_max_age: StreamMaxAge::from_secs(3600).unwrap(), nats_ack_timeout: NonZeroDuration::from_secs(10).unwrap(), - timestamp_max_drift: NonZeroDuration::from_secs(300).unwrap(), + transport: super::super::config::SlackTransportConfig::Webhook(SlackWebhookConfig { + signing_secret: SlackSigningSecret::new(TEST_SECRET).unwrap(), + timestamp_max_drift: NonZeroDuration::from_secs(300).unwrap(), + }), } } @@ -515,9 +497,11 @@ mod tests { } fn mock_app(publisher: MockJetStreamPublisher) -> Router { + let config = test_config(); router_with_clock( wrap_publisher(publisher), - &test_config(), + &config, + config.webhook().unwrap(), FixedEpochClock::from_secs(TEST_NOW), ) } @@ -798,13 +782,21 @@ mod tests { async fn subject_uses_configured_prefix() { let _guard = tracing_guard(); let publisher = MockJetStreamPublisher::new(); + let config = SlackConfig { + subject_prefix: NatsToken::new("custom").unwrap(), + stream_name: NatsToken::new("SLACK").unwrap(), + stream_max_age: StreamMaxAge::from_secs(3600).unwrap(), + nats_ack_timeout: NonZeroDuration::from_secs(10).unwrap(), + transport: super::super::config::SlackTransportConfig::Webhook(SlackWebhookConfig { + signing_secret: SlackSigningSecret::new(TEST_SECRET).unwrap(), + timestamp_max_drift: NonZeroDuration::from_secs(300).unwrap(), + }), + }; let state = AppState { - publisher: wrap_publisher(publisher.clone()), + bridge: SlackBridge::new(wrap_publisher(publisher.clone()), &config), clock: FixedEpochClock::from_secs(TEST_NOW), signing_secret: SlackSigningSecret::new(TEST_SECRET).unwrap(), - subject_prefix: NatsToken::new("custom").unwrap(), - nats_ack_timeout: NonZeroDuration::from_secs(10).unwrap(), timestamp_max_drift: NonZeroDuration::from_secs(300).unwrap(), }; @@ -941,18 +933,20 @@ mod tests { async fn ack_failure_returns_500() { let _guard = tracing_guard(); let publisher = AckFailPublisher::failing(); + let config = test_config(); let state = AppState { - publisher: ClaimCheckPublisher::new( - publisher, - MockObjectStore::new(), - "test-bucket".to_string(), - MaxPayload::from_server_limit(usize::MAX), + bridge: SlackBridge::new( + ClaimCheckPublisher::new( + publisher, + MockObjectStore::new(), + "test-bucket".to_string(), + MaxPayload::from_server_limit(usize::MAX), + ), + &config, ), clock: FixedEpochClock::from_secs(TEST_NOW), signing_secret: SlackSigningSecret::new(TEST_SECRET).unwrap(), - subject_prefix: NatsToken::new("slack").unwrap(), - nats_ack_timeout: NonZeroDuration::from_secs(10).unwrap(), timestamp_max_drift: NonZeroDuration::from_secs(300).unwrap(), }; @@ -976,18 +970,29 @@ mod tests { async fn ack_timeout_returns_500() { let _guard = tracing_guard(); let publisher = AckFailPublisher::hanging(); + let config = SlackConfig { + subject_prefix: NatsToken::new("slack").unwrap(), + stream_name: NatsToken::new("SLACK").unwrap(), + stream_max_age: StreamMaxAge::from_secs(3600).unwrap(), + nats_ack_timeout: NonZeroDuration::from_millis(10).unwrap(), + transport: super::super::config::SlackTransportConfig::Webhook(SlackWebhookConfig { + signing_secret: SlackSigningSecret::new(TEST_SECRET).unwrap(), + timestamp_max_drift: NonZeroDuration::from_secs(300).unwrap(), + }), + }; let state = AppState { - publisher: ClaimCheckPublisher::new( - publisher, - MockObjectStore::new(), - "test-bucket".to_string(), - MaxPayload::from_server_limit(usize::MAX), + bridge: SlackBridge::new( + ClaimCheckPublisher::new( + publisher, + MockObjectStore::new(), + "test-bucket".to_string(), + MaxPayload::from_server_limit(usize::MAX), + ), + &config, ), clock: FixedEpochClock::from_secs(TEST_NOW), signing_secret: SlackSigningSecret::new(TEST_SECRET).unwrap(), - subject_prefix: NatsToken::new("slack").unwrap(), - nats_ack_timeout: NonZeroDuration::from_millis(10).unwrap(), timestamp_max_drift: NonZeroDuration::from_secs(300).unwrap(), }; diff --git a/rsworkspace/crates/trogon-gateway/src/source/slack/socket_mode.rs b/rsworkspace/crates/trogon-gateway/src/source/slack/socket_mode.rs new file mode 100644 index 000000000..013c7b670 --- /dev/null +++ b/rsworkspace/crates/trogon-gateway/src/source/slack/socket_mode.rs @@ -0,0 +1,1041 @@ +use std::fmt; +use std::time::Duration; + +use bytes::Bytes; +use futures_util::{Sink, SinkExt, Stream, StreamExt}; +use serde::Deserialize; +use tokio_tungstenite::tungstenite::{Error as WebSocketError, Message}; +use tracing::{info, warn}; +use trogon_nats::jetstream::{ClaimCheckPublisher, JetStreamPublisher, ObjectStorePut}; + +use super::config::{SlackConfig, SlackSocketModeConfig}; +use super::server::SlackBridge; + +const APPS_CONNECTIONS_OPEN_URL: &str = "https://slack.com/api/apps.connections.open"; +const RECONNECT_INITIAL_DELAY: Duration = Duration::from_secs(1); +const RECONNECT_MAX_DELAY: Duration = Duration::from_secs(30); + +#[derive(Debug)] +pub enum SocketModeError { + MissingSocketModeConfig, + Http(reqwest::Error), + Api(String), + WebSocket(tokio_tungstenite::tungstenite::Error), + Json(serde_json::Error), +} + +impl fmt::Display for SocketModeError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::MissingSocketModeConfig => f.write_str("slack socket_mode config is missing"), + Self::Http(error) => write!(f, "Slack Socket Mode HTTP request failed: {error}"), + Self::Api(error) => write!(f, "Slack apps.connections.open failed: {error}"), + Self::WebSocket(error) => write!(f, "Slack Socket Mode WebSocket failed: {error}"), + Self::Json(error) => write!(f, "Slack Socket Mode JSON parsing failed: {error}"), + } + } +} + +impl std::error::Error for SocketModeError { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + match self { + Self::Http(error) => Some(error), + Self::WebSocket(error) => Some(error), + Self::Json(error) => Some(error), + Self::MissingSocketModeConfig | Self::Api(_) => None, + } + } +} + +impl From for SocketModeError { + fn from(error: reqwest::Error) -> Self { + Self::Http(error) + } +} + +impl From for SocketModeError { + fn from(error: tokio_tungstenite::tungstenite::Error) -> Self { + Self::WebSocket(error) + } +} + +impl From for SocketModeError { + fn from(error: serde_json::Error) -> Self { + Self::Json(error) + } +} + +#[derive(Deserialize)] +struct OpenConnectionResponse { + ok: bool, + url: Option, + error: Option, +} + +#[derive(Deserialize)] +struct SocketEnvelope { + #[serde(rename = "type")] + kind: String, + #[serde(default)] + envelope_id: Option, + #[serde(default)] + payload: Option, + #[serde(default)] + reason: Option, +} + +#[cfg(not(coverage))] +pub async fn run( + publisher: ClaimCheckPublisher, + config: &SlackConfig, +) -> Result<(), SocketModeError> { + let socket_mode = config + .socket_mode() + .ok_or(SocketModeError::MissingSocketModeConfig)? + .clone(); + let bridge = SlackBridge::new(publisher, config); + let client = reqwest::Client::new(); + let mut reconnect_delay = RECONNECT_INITIAL_DELAY; + + loop { + match connect_once(&client, APPS_CONNECTIONS_OPEN_URL, &bridge, &socket_mode).await { + Ok(()) => reconnect_delay = RECONNECT_INITIAL_DELAY, + Err(error) => warn!(error = %error, "Slack Socket Mode connection failed"), + } + + tokio::time::sleep(reconnect_delay).await; + reconnect_delay = next_reconnect_delay(reconnect_delay); + } +} + +fn next_reconnect_delay(current: Duration) -> Duration { + current.saturating_mul(2).min(RECONNECT_MAX_DELAY) +} + +async fn connect_once( + client: &reqwest::Client, + open_url: &str, + bridge: &SlackBridge, + config: &SlackSocketModeConfig, +) -> Result<(), SocketModeError> { + let websocket_url = open_socket_url(client, open_url, config).await?; + info!("connecting to Slack Socket Mode"); + let (ws, _) = tokio_tungstenite::connect_async(&websocket_url).await?; + let (mut sender, mut receiver) = ws.split(); + process_socket_messages(bridge, &mut receiver, &mut sender).await +} + +async fn process_socket_messages( + bridge: &SlackBridge, + receiver: &mut Incoming, + sender: &mut Outgoing, +) -> Result<(), SocketModeError> +where + P: JetStreamPublisher, + S: ObjectStorePut, + Incoming: Stream> + Unpin, + Outgoing: Sink + Unpin, +{ + while let Some(message) = receiver.next().await { + match message? { + Message::Text(text) => { + let (reconnect, ack) = handle_text_frame(bridge, text.as_str()).await?; + if let Some(ack) = ack { + sender.send(Message::Text(ack.into())).await?; + } + if reconnect { + return Ok(()); + } + } + Message::Close(_) => return Ok(()), + Message::Ping(payload) => sender.send(Message::Pong(payload)).await?, + Message::Binary(_) | Message::Pong(_) | Message::Frame(_) => {} + } + } + + Ok(()) +} + +async fn open_socket_url( + client: &reqwest::Client, + open_url: &str, + config: &SlackSocketModeConfig, +) -> Result { + let response = client + .post(open_url) + .header( + reqwest::header::AUTHORIZATION, + format!("Bearer {}", config.app_token.as_str()), + ) + .header(reqwest::header::CONTENT_TYPE, "application/x-www-form-urlencoded") + .send() + .await? + .error_for_status()? + .json::() + .await?; + + if response.ok { + return response + .url + .ok_or_else(|| SocketModeError::Api("missing url".to_string())); + } + + Err(SocketModeError::Api( + response.error.unwrap_or_else(|| "unknown error".to_string()), + )) +} + +async fn handle_text_frame( + bridge: &SlackBridge, + text: &str, +) -> Result<(bool, Option), SocketModeError> { + let envelope: SocketEnvelope = serde_json::from_str(text)?; + + match envelope.kind.as_str() { + "hello" => { + info!("Slack Socket Mode connection established"); + Ok((false, None)) + } + "disconnect" => { + let reason = envelope.reason.as_deref().unwrap_or("unknown"); + warn!(reason, "Slack Socket Mode disconnect requested"); + Ok((true, None)) + } + "events_api" | "interactive" | "slash_commands" => handle_payload_envelope(bridge, envelope, text) + .await + .map(|ack| (false, ack)), + other => { + warn!(kind = other, "Unhandled Slack Socket Mode envelope type"); + bridge + .publish_socket_unroutable("unhandled_socket_mode_type", Bytes::copy_from_slice(text.as_bytes())) + .await; + Ok((false, envelope.envelope_id.map(ack_frame))) + } + } +} + +async fn handle_payload_envelope( + bridge: &SlackBridge, + envelope: SocketEnvelope, + raw_text: &str, +) -> Result, SocketModeError> { + let Some(envelope_id) = envelope.envelope_id else { + warn!(kind = envelope.kind, "Missing Slack Socket Mode envelope_id"); + bridge + .publish_socket_unroutable("missing_envelope_id", Bytes::copy_from_slice(raw_text.as_bytes())) + .await; + return Ok(None); + }; + + let Some(payload) = envelope.payload else { + warn!(kind = envelope.kind, "Missing Slack Socket Mode payload"); + bridge + .publish_socket_unroutable( + "missing_socket_mode_payload", + Bytes::copy_from_slice(raw_text.as_bytes()), + ) + .await; + return Ok(None); + }; + + let status = match envelope.kind.as_str() { + "events_api" => { + let body = Bytes::from(serde_json::to_vec(&payload)?); + bridge.handle_json_body(body).await.0 + } + "interactive" => bridge.handle_socket_interaction(&payload).await?, + "slash_commands" => bridge.handle_socket_slash_command(&payload).await?, + other => { + warn!(kind = other, "Unhandled Slack Socket Mode payload envelope type"); + bridge + .publish_socket_unroutable( + "unhandled_socket_mode_type", + Bytes::copy_from_slice(raw_text.as_bytes()), + ) + .await; + return Ok(None); + } + }; + + if status.is_success() { + Ok(Some(ack_frame(envelope_id))) + } else { + Ok(None) + } +} + +fn ack_frame(envelope_id: String) -> String { + serde_json::json!({ "envelope_id": envelope_id }).to_string() +} + +#[cfg(test)] +mod tests { + use super::*; + use axum::Router; + use axum::extract::State; + use axum::extract::ws::{Message as AxumMessage, WebSocketUpgrade}; + use axum::http::StatusCode; + use axum::response::Response; + use axum::routing::{any, post}; + use std::error::Error; + use std::future::IntoFuture; + use std::net::SocketAddr; + use tokio::net::TcpListener; + use tokio::sync::mpsc; + use trogon_nats::NatsToken; + use trogon_nats::jetstream::{MaxPayload, MockJetStreamPublisher, MockObjectStore, StreamMaxAge}; + use trogon_std::NonZeroDuration; + + fn wrap_publisher( + publisher: MockJetStreamPublisher, + ) -> ClaimCheckPublisher { + ClaimCheckPublisher::new( + publisher, + MockObjectStore::new(), + "test-bucket".to_string(), + MaxPayload::from_server_limit(usize::MAX), + ) + } + + fn socket_config() -> SlackConfig { + SlackConfig { + subject_prefix: NatsToken::new("slack").unwrap(), + stream_name: NatsToken::new("SLACK").unwrap(), + stream_max_age: StreamMaxAge::from_secs(3600).unwrap(), + nats_ack_timeout: NonZeroDuration::from_secs(10).unwrap(), + transport: super::super::config::SlackTransportConfig::SocketMode(SlackSocketModeConfig { + app_token: super::super::config::SlackAppToken::new("xapp-test-token").unwrap(), + }), + } + } + + #[cfg(not(coverage))] + fn webhook_config() -> SlackConfig { + SlackConfig { + subject_prefix: NatsToken::new("slack").unwrap(), + stream_name: NatsToken::new("SLACK").unwrap(), + stream_max_age: StreamMaxAge::from_secs(3600).unwrap(), + nats_ack_timeout: NonZeroDuration::from_secs(10).unwrap(), + transport: super::super::config::SlackTransportConfig::Webhook(super::super::config::SlackWebhookConfig { + signing_secret: super::super::config::SlackSigningSecret::new("slack-secret").unwrap(), + timestamp_max_drift: NonZeroDuration::from_secs(300).unwrap(), + }), + } + } + + fn bridge(publisher: MockJetStreamPublisher) -> SlackBridge { + SlackBridge::new(wrap_publisher(publisher), &socket_config()) + } + + #[test] + fn reconnect_delay_doubles_until_cap() { + assert_eq!(next_reconnect_delay(Duration::from_secs(1)), Duration::from_secs(2)); + assert_eq!(next_reconnect_delay(Duration::from_secs(20)), RECONNECT_MAX_DELAY); + assert_eq!(next_reconnect_delay(RECONNECT_MAX_DELAY), RECONNECT_MAX_DELAY); + } + + #[cfg(not(coverage))] + #[tokio::test] + async fn run_requires_socket_mode_config() { + let error = run(wrap_publisher(MockJetStreamPublisher::new()), &webhook_config()) + .await + .unwrap_err(); + + assert!(matches!(error, SocketModeError::MissingSocketModeConfig)); + assert_eq!(error.to_string(), "slack socket_mode config is missing"); + assert!(error.source().is_none()); + } + + #[test] + fn socket_mode_error_sources_are_exposed() { + let error = SocketModeError::MissingSocketModeConfig; + assert_eq!(error.to_string(), "slack socket_mode config is missing"); + assert!(error.source().is_none()); + + let json_error = serde_json::from_str::("not-json").unwrap_err(); + let error = SocketModeError::from(json_error); + assert!(error.to_string().contains("JSON parsing failed")); + assert!(error.source().is_some()); + + let error = SocketModeError::WebSocket(tokio_tungstenite::tungstenite::Error::ConnectionClosed); + assert!(error.to_string().contains("WebSocket failed")); + assert!(error.source().is_some()); + + let error = SocketModeError::Api("invalid_auth".to_string()); + assert_eq!(error.to_string(), "Slack apps.connections.open failed: invalid_auth"); + assert!(error.source().is_none()); + } + + #[tokio::test] + async fn socket_events_api_payload_publishes_and_acks() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "events_api", + "envelope_id": "env-1", + "payload": { + "type": "event_callback", + "event_id": "Ev01ABC123", + "team_id": "T01ABC", + "event": { "type": "message", "text": "hello" } + } + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher.clone()), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, Some(r#"{"envelope_id":"env-1"}"#.to_string())); + let messages = publisher.published_messages(); + assert_eq!(messages.len(), 1); + assert_eq!(messages[0].subject, "slack.event.message"); + assert_eq!( + messages[0] + .headers + .get(super::super::constants::NATS_HEADER_PAYLOAD_KIND) + .map(|v| v.as_str()), + Some("event"), + ); + } + + #[tokio::test] + async fn socket_interactive_payload_publishes_and_acks() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "interactive", + "envelope_id": "env-2", + "payload": { + "type": "block_actions", + "trigger_id": "trigger123", + "team": { "id": "T01ABC" } + } + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher.clone()), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, Some(r#"{"envelope_id":"env-2"}"#.to_string())); + let messages = publisher.published_messages(); + assert_eq!(messages[0].subject, "slack.interaction.block_actions"); + assert_eq!( + messages[0] + .headers + .get(super::super::constants::NATS_HEADER_PAYLOAD_KIND) + .map(|v| v.as_str()), + Some("interaction"), + ); + } + + #[tokio::test] + async fn socket_slash_command_payload_publishes_and_acks() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "slash_commands", + "envelope_id": "env-3", + "payload": { + "command": "/trogon", + "team_id": "T01ABC", + "trigger_id": "trigger456" + } + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher.clone()), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, Some(r#"{"envelope_id":"env-3"}"#.to_string())); + let messages = publisher.published_messages(); + assert_eq!(messages[0].subject, "slack.command.trogon"); + assert_eq!( + messages[0] + .headers + .get(super::super::constants::NATS_HEADER_PAYLOAD_KIND) + .map(|v| v.as_str()), + Some("command"), + ); + } + + #[tokio::test] + async fn failed_publish_does_not_ack() { + let publisher = MockJetStreamPublisher::new(); + publisher.fail_next_js_publish(); + let envelope = serde_json::json!({ + "type": "events_api", + "envelope_id": "env-4", + "payload": { + "type": "event_callback", + "event_id": "Ev01ABC123", + "team_id": "T01ABC", + "event": { "type": "message" } + } + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, None); + } + + #[tokio::test] + async fn missing_event_id_publishes_unroutable_and_does_not_ack() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "events_api", + "envelope_id": "env-5", + "payload": { + "type": "event_callback", + "team_id": "T01ABC", + "event": { "type": "message" } + } + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher.clone()), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, None); + assert_eq!(publisher.published_messages()[0].subject, "slack.unroutable"); + } + + #[tokio::test] + async fn hello_frame_does_not_ack_or_reconnect() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "hello" + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, None); + } + + #[tokio::test] + async fn unknown_frame_publishes_unroutable_and_acks_when_enveloped() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "unsupported", + "envelope_id": "env-unsupported", + "payload": { "value": true } + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher.clone()), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, Some(r#"{"envelope_id":"env-unsupported"}"#.to_string())); + let messages = publisher.published_messages(); + assert_eq!(messages[0].subject, "slack.unroutable"); + assert_eq!( + messages[0] + .headers + .get(super::super::constants::NATS_HEADER_REJECT_REASON) + .map(|v| v.as_str()), + Some("unhandled_socket_mode_type"), + ); + } + + #[tokio::test] + async fn payload_frame_missing_envelope_id_publishes_unroutable_without_ack() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "events_api", + "payload": { + "type": "event_callback", + "event_id": "Ev01ABC123", + "event": { "type": "message" } + } + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher.clone()), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, None); + assert_eq!( + publisher.published_messages()[0] + .headers + .get(super::super::constants::NATS_HEADER_REJECT_REASON) + .map(|v| v.as_str()), + Some("missing_envelope_id"), + ); + } + + #[tokio::test] + async fn payload_frame_missing_payload_publishes_unroutable_without_ack() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "events_api", + "envelope_id": "env-no-payload" + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher.clone()), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, None); + assert_eq!( + publisher.published_messages()[0] + .headers + .get(super::super::constants::NATS_HEADER_REJECT_REASON) + .map(|v| v.as_str()), + Some("missing_socket_mode_payload"), + ); + } + + #[tokio::test] + async fn payload_envelope_with_unexpected_kind_publishes_unroutable_without_ack() { + let publisher = MockJetStreamPublisher::new(); + let raw_text = r#"{"type":"unexpected","envelope_id":"env-unexpected","payload":{}}"#; + let envelope = SocketEnvelope { + kind: "unexpected".to_string(), + envelope_id: Some("env-unexpected".to_string()), + payload: Some(serde_json::json!({})), + reason: None, + }; + + let ack = handle_payload_envelope(&bridge(publisher.clone()), envelope, raw_text) + .await + .unwrap(); + + assert_eq!(ack, None); + assert_eq!( + publisher.published_messages()[0] + .headers + .get(super::super::constants::NATS_HEADER_REJECT_REASON) + .map(|v| v.as_str()), + Some("unhandled_socket_mode_type"), + ); + } + + #[tokio::test] + async fn socket_slash_command_missing_command_publishes_unroutable_without_ack() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "slash_commands", + "envelope_id": "env-missing-command", + "payload": { + "team_id": "T01ABC", + "trigger_id": "trigger456" + } + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher.clone()), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, None); + assert_eq!( + publisher.published_messages()[0] + .headers + .get(super::super::constants::NATS_HEADER_REJECT_REASON) + .map(|v| v.as_str()), + Some("missing_command"), + ); + } + + #[tokio::test] + async fn socket_slash_command_missing_trigger_publishes_unroutable_without_ack() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "slash_commands", + "envelope_id": "env-missing-trigger", + "payload": { + "command": "/trogon", + "team_id": "T01ABC" + } + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher.clone()), &envelope.to_string()) + .await + .unwrap(); + + assert!(!reconnect); + assert_eq!(ack, None); + assert_eq!( + publisher.published_messages()[0] + .headers + .get(super::super::constants::NATS_HEADER_REJECT_REASON) + .map(|v| v.as_str()), + Some("missing_command_trigger_id"), + ); + } + + #[derive(Clone)] + struct OpenState { + ws_url: String, + seen_auth: mpsc::UnboundedSender, + } + + #[derive(Clone)] + struct OpenBodyState { + body: &'static str, + } + + #[derive(Clone)] + struct TextWsState { + text: String, + seen_ack: mpsc::UnboundedSender, + } + + async fn open_handler(State(state): State, headers: axum::http::HeaderMap) -> String { + let auth = headers + .get(reqwest::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(); + state.seen_auth.send(auth).unwrap(); + serde_json::json!({ "ok": true, "url": state.ws_url }).to_string() + } + + async fn open_body_handler(State(state): State) -> String { + state.body.to_string() + } + + async fn open_status_handler() -> (StatusCode, &'static str) { + (StatusCode::INTERNAL_SERVER_ERROR, "failed") + } + + async fn ws_handler(ws: WebSocketUpgrade, State(sender): State>) -> Response { + ws.on_upgrade(move |mut socket| async move { + sender.send(()).unwrap(); + socket.close().await.unwrap(); + }) + } + + async fn text_ws_handler(ws: WebSocketUpgrade, State(state): State) -> Response { + ws.on_upgrade(move |mut socket| async move { + socket.send(AxumMessage::Text(state.text.into())).await.unwrap(); + if let Some(Ok(AxumMessage::Text(ack))) = socket.recv().await { + state.seen_ack.send(ack.to_string()).unwrap(); + } + socket.close().await.unwrap(); + }) + } + + async fn disconnect_ws_handler(ws: WebSocketUpgrade) -> Response { + ws.on_upgrade(move |mut socket| async move { + socket + .send(AxumMessage::Text(r#"{"type":"disconnect"}"#.into())) + .await + .unwrap(); + }) + } + + async fn control_ws_handler(ws: WebSocketUpgrade) -> Response { + ws.on_upgrade(move |mut socket| async move { + socket + .send(AxumMessage::Ping(Bytes::from_static(b"ping"))) + .await + .unwrap(); + socket + .send(AxumMessage::Binary(Bytes::from_static(b"payload"))) + .await + .unwrap(); + socket.close().await.unwrap(); + }) + } + + async fn spawn_server(app: Router) -> SocketAddr { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(axum::serve(listener, app).into_future()); + addr + } + + #[tokio::test] + async fn apps_connections_open_response_url_is_used() { + let (ws_seen_tx, mut ws_seen_rx) = mpsc::unbounded_channel(); + let ws_addr = spawn_server(Router::new().route("/socket", any(ws_handler)).with_state(ws_seen_tx)).await; + let (auth_tx, mut auth_rx) = mpsc::unbounded_channel(); + let open_addr = spawn_server( + Router::new() + .route("/apps.connections.open", post(open_handler)) + .with_state(OpenState { + ws_url: format!("ws://{ws_addr}/socket"), + seen_auth: auth_tx, + }), + ) + .await; + + let config = socket_config(); + let client = reqwest::Client::new(); + let url = open_socket_url( + &client, + &format!("http://{open_addr}/apps.connections.open"), + config.socket_mode().unwrap(), + ) + .await + .unwrap(); + + assert_eq!(url, format!("ws://{ws_addr}/socket")); + assert_eq!(auth_rx.recv().await.as_deref(), Some("Bearer xapp-test-token")); + + let bridge = bridge(MockJetStreamPublisher::new()); + let result = tokio::time::timeout( + Duration::from_secs(2), + connect_once( + &client, + &format!("http://{open_addr}/apps.connections.open"), + &bridge, + config.socket_mode().unwrap(), + ), + ) + .await; + + result.unwrap().unwrap(); + assert!(ws_seen_rx.recv().await.is_some()); + } + + #[tokio::test] + async fn apps_connections_open_missing_url_is_api_error() { + let open_addr = spawn_server( + Router::new() + .route("/apps.connections.open", post(open_body_handler)) + .with_state(OpenBodyState { body: r#"{"ok":true}"# }), + ) + .await; + + let config = socket_config(); + let error = open_socket_url( + &reqwest::Client::new(), + &format!("http://{open_addr}/apps.connections.open"), + config.socket_mode().unwrap(), + ) + .await + .unwrap_err(); + + assert!(matches!(error, SocketModeError::Api(ref message) if message == "missing url")); + } + + #[tokio::test] + async fn apps_connections_open_error_is_api_error() { + let open_addr = spawn_server( + Router::new() + .route("/apps.connections.open", post(open_body_handler)) + .with_state(OpenBodyState { + body: r#"{"ok":false,"error":"invalid_auth"}"#, + }), + ) + .await; + + let config = socket_config(); + let error = open_socket_url( + &reqwest::Client::new(), + &format!("http://{open_addr}/apps.connections.open"), + config.socket_mode().unwrap(), + ) + .await + .unwrap_err(); + + assert!(matches!(error, SocketModeError::Api(ref message) if message == "invalid_auth")); + } + + #[tokio::test] + async fn apps_connections_open_http_error_exposes_source() { + let open_addr = spawn_server(Router::new().route("/apps.connections.open", post(open_status_handler))).await; + + let config = socket_config(); + let error = open_socket_url( + &reqwest::Client::new(), + &format!("http://{open_addr}/apps.connections.open"), + config.socket_mode().unwrap(), + ) + .await + .unwrap_err(); + + assert!(matches!(error, SocketModeError::Http(_))); + assert!(error.to_string().contains("HTTP request failed")); + assert!(error.source().is_some()); + } + + #[tokio::test] + async fn connect_once_websocket_error_exposes_source() { + let open_addr = spawn_server( + Router::new() + .route("/apps.connections.open", post(open_body_handler)) + .with_state(OpenBodyState { + body: r#"{"ok":true,"url":"ws://127.0.0.1:9/socket"}"#, + }), + ) + .await; + + let config = socket_config(); + let bridge = bridge(MockJetStreamPublisher::new()); + let error = tokio::time::timeout( + Duration::from_secs(2), + connect_once( + &reqwest::Client::new(), + &format!("http://{open_addr}/apps.connections.open"), + &bridge, + config.socket_mode().unwrap(), + ), + ) + .await + .unwrap() + .unwrap_err(); + + assert!(matches!(error, SocketModeError::WebSocket(_))); + assert!(error.to_string().contains("WebSocket failed")); + assert!(error.source().is_some()); + } + + #[tokio::test] + async fn connect_once_sends_ack_for_successful_text_frame() { + let envelope = serde_json::json!({ + "type": "events_api", + "envelope_id": "env-connect", + "payload": { + "type": "event_callback", + "event_id": "EvConnect", + "team_id": "T01ABC", + "event": { "type": "message" } + } + }); + let (ack_tx, mut ack_rx) = mpsc::unbounded_channel(); + let ws_addr = spawn_server( + Router::new() + .route("/socket", any(text_ws_handler)) + .with_state(TextWsState { + text: envelope.to_string(), + seen_ack: ack_tx, + }), + ) + .await; + let open_addr = spawn_server( + Router::new() + .route("/apps.connections.open", post(open_body_handler)) + .with_state(OpenBodyState { + body: Box::leak( + serde_json::json!({ "ok": true, "url": format!("ws://{ws_addr}/socket") }) + .to_string() + .into_boxed_str(), + ), + }), + ) + .await; + + let config = socket_config(); + let bridge = bridge(MockJetStreamPublisher::new()); + connect_once( + &reqwest::Client::new(), + &format!("http://{open_addr}/apps.connections.open"), + &bridge, + config.socket_mode().unwrap(), + ) + .await + .unwrap(); + + assert_eq!(ack_rx.recv().await.as_deref(), Some(r#"{"envelope_id":"env-connect"}"#),); + } + + #[tokio::test] + async fn connect_once_reconnects_on_disconnect_frame() { + let ws_addr = spawn_server(Router::new().route("/socket", any(disconnect_ws_handler))).await; + let open_addr = spawn_server( + Router::new() + .route("/apps.connections.open", post(open_body_handler)) + .with_state(OpenBodyState { + body: Box::leak( + serde_json::json!({ "ok": true, "url": format!("ws://{ws_addr}/socket") }) + .to_string() + .into_boxed_str(), + ), + }), + ) + .await; + + let config = socket_config(); + let bridge = bridge(MockJetStreamPublisher::new()); + connect_once( + &reqwest::Client::new(), + &format!("http://{open_addr}/apps.connections.open"), + &bridge, + config.socket_mode().unwrap(), + ) + .await + .unwrap(); + } + + #[tokio::test] + async fn connect_once_handles_control_frames() { + let ws_addr = spawn_server(Router::new().route("/socket", any(control_ws_handler))).await; + let open_addr = spawn_server( + Router::new() + .route("/apps.connections.open", post(open_body_handler)) + .with_state(OpenBodyState { + body: Box::leak( + serde_json::json!({ "ok": true, "url": format!("ws://{ws_addr}/socket") }) + .to_string() + .into_boxed_str(), + ), + }), + ) + .await; + + let config = socket_config(); + let bridge = bridge(MockJetStreamPublisher::new()); + connect_once( + &reqwest::Client::new(), + &format!("http://{open_addr}/apps.connections.open"), + &bridge, + config.socket_mode().unwrap(), + ) + .await + .unwrap(); + } + + #[tokio::test] + async fn process_socket_messages_returns_when_stream_ends() { + let mut receiver = futures_util::stream::empty::>(); + let mut sender = futures_util::sink::drain::() + .sink_map_err(|never: std::convert::Infallible| -> WebSocketError { match never {} }); + + process_socket_messages(&bridge(MockJetStreamPublisher::new()), &mut receiver, &mut sender) + .await + .unwrap(); + } + + #[tokio::test] + async fn disconnect_frame_completes_without_ack() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "disconnect", + "reason": "refresh_requested" + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher), &envelope.to_string()) + .await + .unwrap(); + + assert!(reconnect); + assert_eq!(ack, None); + } + + #[tokio::test] + async fn disconnect_frame_defaults_missing_reason_without_ack() { + let publisher = MockJetStreamPublisher::new(); + let envelope = serde_json::json!({ + "type": "disconnect" + }); + + let (reconnect, ack) = handle_text_frame(&bridge(publisher), &envelope.to_string()) + .await + .unwrap(); + + assert!(reconnect); + assert_eq!(ack, None); + } +}