diff --git a/Cargo.toml b/Cargo.toml index b4a0de0..d030ce5 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "a3s-boot" -version = "0.1.2" +version = "0.1.3" edition = "2021" authors = ["A3S Lab"] license = "MIT" @@ -28,6 +28,14 @@ events = ["dep:a3s-event"] file-upload = ["dep:bytes", "dep:multer"] health = [] http-client = ["dep:reqwest"] +ilink = [ + "dep:async-trait", + "dep:base64", + "dep:rand", + "dep:reqwest", + "dep:url", + "dep:zeroize", +] logging = [] macros = ["dep:a3s-boot-macros"] openapi-schemas = ["dep:schemars"] @@ -51,8 +59,10 @@ a3s-acl = { version = "0.2.1", optional = true } a3s-boot-macros = { version = "0.1.2", path = "macros", optional = true } a3s-event = { version = "0.3.0", default-features = false, optional = true } a3s-lane = { version = "0.5.1", default-features = false, optional = true } +async-trait = { version = "0.1", optional = true } async-nats = { version = "0.49.1", default-features = false, optional = true } axum = { version = "0.8", features = ["ws"], optional = true } +base64 = { version = "0.22", optional = true } bytes = { version = "1", optional = true } chrono = { version = "0.4", optional = true } cron = { version = "0.17.0", optional = true } @@ -65,8 +75,9 @@ lapin = { version = "4.10.0", default-features = false, features = ["tokio"], op multer = { version = "3", optional = true } percent-encoding = "2" prost = { version = "0.14.4", optional = true } +rand = { version = "0.8", optional = true } redis = { version = "1.3.0", default-features = false, features = ["tokio-comp"], optional = true } -reqwest = { version = "0.12", default-features = false, features = ["rustls-tls"], optional = true } +reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "stream"], optional = true } rskafka = { version = "0.6.0", default-features = false, optional = true } rumqttc = { version = "0.25.1", default-features = false, optional = true } schemars = { version = "1", optional = true } @@ -78,6 +89,8 @@ thiserror = "2" tokio = { version = "1", features = ["fs", "io-util", "net", "rt", "sync", "time"], optional = true } tonic = { version = "0.14.6", default-features = false, features = ["codegen", "transport"], optional = true } tonic-prost = { version = "0.14.6", optional = true } +url = { version = "2", optional = true } +zeroize = { version = "1", features = ["derive"], optional = true } [dev-dependencies] tokio = { version = "1", features = ["macros", "rt", "sync", "time"] } diff --git a/README.md b/README.md index e08fff5..9c7e7f0 100644 --- a/README.md +++ b/README.md @@ -123,6 +123,7 @@ are opt-in. | Scheduling | `schedule` | Cron, interval, and timeout jobs | | Observability | `logging`, `health` | Structured logging and health indicators | | HTTP utilities | `http-client`, `compression` | Outbound HTTP and gzip responses | +| Channels | `ilink` | Tencent Weixin iLink QR login, polling, messaging, and lifecycle client | | Content | `file-upload`, `static` | Multipart uploads and static files | | Context | `request-context` | Task-local access to the current request | | OpenAPI | `openapi-schemas` | `schemars`-based component schemas | @@ -141,7 +142,7 @@ is not a claim of support for every production backend. ```toml [dependencies] -a3s-boot = "0.1.2" +a3s-boot = "0.1.3" tokio = { version = "1", features = ["macros", "rt-multi-thread"] } ``` @@ -149,14 +150,14 @@ For a core-only build without Axum, macros, or shutdown signal handling: ```toml [dependencies] -a3s-boot = { version = "0.1.2", default-features = false } +a3s-boot = { version = "0.1.3", default-features = false } ``` Enable only the optional modules an application uses: ```toml [dependencies] -a3s-boot = { version = "0.1.2", features = ["auth", "security", "openapi-schemas"] } +a3s-boot = { version = "0.1.3", features = ["auth", "security", "openapi-schemas"] } serde = { version = "1", features = ["derive"] } tokio = { version = "1", features = ["macros", "rt-multi-thread"] } ``` @@ -322,6 +323,27 @@ Transport implementations share typed payload handling, scoped providers, validation, guards, interceptors, pipes, exception filters, and client APIs. Protocol delivery and durability semantics still depend on the selected backend. +### Weixin iLink + +The optional `ilink` feature provides the native Rust protocol boundary used by +the Tencent Weixin channel. `IlinkModule` exports a typed `IlinkClient` +provider; the client owns QR login requests, authenticated headers, strict +server URL validation, update polling, text replies, typing calls, and channel +start/stop notifications. + +```rust +use a3s_boot::ilink::IlinkModule; + +let module = IlinkModule::weixin("A3S/0.10.1"); +``` + +The wire defaults are compatible with Tencent `openclaw-weixin` v2.4.6: +`iLink-App-Id: bot`, `bot_type=3`, and packed client version `2.4.6`. The +product-specific `bot_agent` remains `A3S/` so upstream diagnostics do +not misidentify the caller. Boot deliberately does not own browser APIs, +credential persistence, owner authorization, or agent/session commands; those +policies stay in the host application. + ## Architecture The application core is independent of its HTTP server and message broker: diff --git a/ROADMAP.md b/ROADMAP.md index 6a05d2c..3af0409 100644 --- a/ROADMAP.md +++ b/ROADMAP.md @@ -222,6 +222,11 @@ Implemented today: topics plus optional gRPC unary request/reply and event calls. Transport error envelopes round-trip through the same `BootError` HTTP exception mapping used by HTTP routes. +- Optional native Tencent Weixin iLink support with an injectable client, + QR-login and redirect handling, authenticated update polling, text replies, + typing calls, lifecycle notifications, bounded responses, and strict URL + validation. Browser APIs, credential storage, and product authorization + remain host-application responsibilities. - ACL-backed typed configuration modules with `ConfigModule`, named/global provider exports, environment/default function support, and validation hooks. - Provider-backed outbound HTTP clients with `HttpModule`, `HttpService`, diff --git a/src/ilink/auth.rs b/src/ilink/auth.rs new file mode 100644 index 0000000..3e3e859 --- /dev/null +++ b/src/ilink/auth.rs @@ -0,0 +1,93 @@ +//! Secret handling and client-version encoding for the iLink wire protocol. + +use std::fmt; + +use base64::Engine as _; +use serde::{Deserialize, Deserializer, Serialize, Serializer}; +use thiserror::Error; +use zeroize::{Zeroize, ZeroizeOnDrop}; + +const MAX_SECRET_BYTES: usize = 64 * 1024; + +#[derive(Clone, PartialEq, Eq, Zeroize, ZeroizeOnDrop)] +pub struct SecretValue(String); + +impl SecretValue { + pub fn new(value: impl Into) -> Result { + let value = value.into(); + if value.is_empty() { + return Err(SecretValueError::Empty); + } + if value.len() > MAX_SECRET_BYTES { + return Err(SecretValueError::TooLarge); + } + Ok(Self(value)) + } + + pub fn expose(&self) -> &str { + &self.0 + } +} + +impl fmt::Debug for SecretValue { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("SecretValue([REDACTED])") + } +} + +impl Serialize for SecretValue { + fn serialize(&self, serializer: S) -> Result + where + S: Serializer, + { + serializer.serialize_str(self.expose()) + } +} + +impl<'de> Deserialize<'de> for SecretValue { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + let value = String::deserialize(deserializer)?; + Self::new(value).map_err(serde::de::Error::custom) + } +} + +#[derive(Clone, Copy, Debug, Error, PartialEq, Eq)] +pub enum SecretValueError { + #[error("secret value is empty")] + Empty, + #[error("secret value exceeds the protocol size limit")] + TooLarge, +} + +#[derive(Clone, Debug, Error, PartialEq, Eq)] +pub enum ClientVersionError { + #[error("client version must contain exactly three numeric components")] + InvalidShape, + #[error("client version component exceeds 255")] + ComponentOutOfRange, +} + +pub(super) fn pack_client_version(version: &str) -> Result { + let components = version.split('.').collect::>(); + if components.len() != 3 || components.iter().any(|component| component.is_empty()) { + return Err(ClientVersionError::InvalidShape); + } + let mut parsed = [0u32; 3]; + for (index, component) in components.into_iter().enumerate() { + let value = component + .parse::() + .map_err(|_| ClientVersionError::InvalidShape)?; + if value > u8::MAX as u32 { + return Err(ClientVersionError::ComponentOutOfRange); + } + parsed[index] = value; + } + Ok((parsed[0] << 16) | (parsed[1] << 8) | parsed[2]) +} + +pub(super) fn random_wechat_uin() -> String { + base64::engine::general_purpose::STANDARD.encode(rand::random::().to_string()) +} diff --git a/src/ilink/client.rs b/src/ilink/client.rs new file mode 100644 index 0000000..cb6d0db --- /dev/null +++ b/src/ilink/client.rs @@ -0,0 +1,172 @@ +//! iLink client identity, authentication, and protocol errors. + +use std::fmt; +use std::time::Duration; + +use reqwest::header::{HeaderMap, HeaderName, HeaderValue, AUTHORIZATION, CONTENT_TYPE}; +use thiserror::Error; + +use super::auth::{pack_client_version, random_wechat_uin, ClientVersionError, SecretValue}; +use super::url_policy::{IlinkUrlError, ValidatedBaseUrl}; + +pub(super) const STALE_TOKEN_ERROR_CODE: i64 = -14; +pub(super) const DEFAULT_API_TIMEOUT: Duration = Duration::from_secs(15); +pub(super) const DEFAULT_CONFIG_TIMEOUT: Duration = Duration::from_secs(10); +pub(super) const DEFAULT_QR_POLL_TIMEOUT: Duration = Duration::from_secs(35); +pub(super) const MAX_LONG_POLL_TIMEOUT: Duration = Duration::from_secs(60); +pub(super) const MAX_RESPONSE_BYTES: usize = 1024 * 1024; + +const AUTHORIZATION_TYPE: HeaderName = HeaderName::from_static("authorizationtype"); +const WECHAT_UIN: HeaderName = HeaderName::from_static("x-wechat-uin"); +const ILINK_APP_ID: HeaderName = HeaderName::from_static("ilink-app-id"); +const ILINK_APP_CLIENT_VERSION: HeaderName = HeaderName::from_static("ilink-app-clientversion"); + +#[derive(Clone, Debug)] +pub(super) struct IlinkClientIdentity { + app_id: String, + packed_client_version: u32, + pub(super) bot_type: String, + pub(super) channel_version: String, + pub(super) bot_agent: String, +} + +impl IlinkClientIdentity { + pub(super) fn new( + app_id: impl Into, + bot_type: impl Into, + client_version: &str, + bot_agent: impl Into, + ) -> Result { + let app_id = bounded_ascii(app_id.into(), "app id", 128)?; + let bot_type = bounded_ascii(bot_type.into(), "bot type", 16)?; + let channel_version = bounded_ascii(client_version.to_string(), "client version", 32)?; + let bot_agent = bounded_ascii(bot_agent.into(), "bot agent", 256)?; + let packed_client_version = pack_client_version(client_version)?; + Ok(Self { + app_id, + packed_client_version, + bot_type, + channel_version, + bot_agent, + }) + } + + pub(super) fn base_info(&self) -> super::types::BaseInfo { + super::types::BaseInfo { + channel_version: Some(self.channel_version.clone()), + bot_agent: Some(self.bot_agent.clone()), + } + } + + pub(super) fn application_headers(&self) -> Result { + let mut headers = HeaderMap::new(); + headers.insert(ILINK_APP_ID, safe_header_value(&self.app_id, "app id")?); + headers.insert( + ILINK_APP_CLIENT_VERSION, + safe_header_value(&self.packed_client_version.to_string(), "client version")?, + ); + Ok(headers) + } + + pub(super) fn post_headers( + &self, + token: Option<&SecretValue>, + ) -> Result { + let mut headers = self.application_headers()?; + headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); + headers.insert( + AUTHORIZATION_TYPE, + HeaderValue::from_static("ilink_bot_token"), + ); + headers.insert( + WECHAT_UIN, + safe_header_value(&random_wechat_uin(), "Weixin UIN")?, + ); + if let Some(token) = token { + headers.insert( + AUTHORIZATION, + safe_header_value(&format!("Bearer {}", token.expose()), "authorization")?, + ); + } + Ok(headers) + } +} + +#[derive(Clone)] +pub struct IlinkAuth { + pub(super) base_url: ValidatedBaseUrl, + pub(super) bot_token: SecretValue, +} + +impl IlinkAuth { + pub fn new(base_url: ValidatedBaseUrl, bot_token: SecretValue) -> Self { + Self { + base_url, + bot_token, + } + } +} + +impl fmt::Debug for IlinkAuth { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("IlinkAuth") + .field("base_url", &self.base_url) + .field("bot_token", &self.bot_token) + .finish() + } +} + +fn bounded_ascii(value: String, field: &'static str, max: usize) -> Result { + if value.is_empty() + || value.len() > max + || !value.is_ascii() + || value.chars().any(char::is_control) + { + return Err(IlinkError::InvalidConfiguration(field)); + } + Ok(value) +} + +fn safe_header_value(value: &str, field: &'static str) -> Result { + HeaderValue::from_str(value).map_err(|_| IlinkError::InvalidConfiguration(field)) +} + +pub(super) fn ensure_api_success( + operation: &'static str, + ret: Option, + errcode: Option, +) -> Result<(), IlinkError> { + let code = errcode + .filter(|code| *code != 0) + .or_else(|| ret.filter(|code| *code != 0)); + match code { + None => Ok(()), + Some(STALE_TOKEN_ERROR_CODE) => Err(IlinkError::StaleCredential), + Some(code) => Err(IlinkError::Protocol { operation, code }), + } +} + +#[derive(Clone, Debug, Error, PartialEq, Eq)] +pub enum IlinkError { + #[error("iLink configuration field is invalid: {0}")] + InvalidConfiguration(&'static str), + #[error("iLink URL policy rejected the request")] + Url(#[from] IlinkUrlError), + #[error("iLink client version is invalid")] + ClientVersion(#[from] ClientVersionError), + #[error("iLink request timed out")] + Timeout, + #[error("iLink transport failed")] + Transport, + #[error("iLink returned HTTP status {0}")] + HttpStatus(u16), + #[error("iLink response exceeds the size limit")] + ResponseTooLarge, + #[error("iLink returned an invalid response for {0}")] + InvalidResponse(&'static str), + #[error("iLink credential is stale")] + StaleCredential, + #[error("iLink operation {operation} failed with code {code}")] + Protocol { operation: &'static str, code: i64 }, +} diff --git a/src/ilink/login.rs b/src/ilink/login.rs new file mode 100644 index 0000000..ffbc424 --- /dev/null +++ b/src/ilink/login.rs @@ -0,0 +1,105 @@ +//! QR-code login protocol operations. + +use super::auth::SecretValue; +use super::client::{IlinkError, DEFAULT_QR_POLL_TIMEOUT}; +use super::transport::IlinkClient; +use super::types::{CreateQrRequest, CreateQrResponse, PollQrResponse}; +use super::url_policy::ValidatedBaseUrl; + +const MAX_QR_IMAGE_CONTENT_BYTES: usize = 256 * 1024; +const MAX_LOCAL_TOKEN_COUNT: usize = 10; + +impl IlinkClient { + pub(super) async fn create_qr_request( + &self, + local_tokens: &[SecretValue], + ) -> Result { + if local_tokens.len() > MAX_LOCAL_TOKEN_COUNT { + return Err(IlinkError::InvalidConfiguration("local token list")); + } + let mut url = self.qr_base_url.join("ilink/bot/get_bot_qrcode")?; + url.query_pairs_mut() + .append_pair("bot_type", &self.identity.bot_type); + let request = self + .http + .post(url) + .headers(self.identity.post_headers(None)?); + let response: CreateQrResponse = self + .post_json_without_timeout( + request, + &CreateQrRequest { + local_token_list: local_tokens.to_vec(), + }, + "create_qr", + ) + .await?; + if response.qrcode_img_content.expose().len() > MAX_QR_IMAGE_CONTENT_BYTES + || response.qrcode_img_content.expose().contains('\0') + { + return Err(IlinkError::InvalidResponse("create_qr")); + } + Ok(response) + } + + pub(super) async fn poll_qr_request( + &self, + base_url: &ValidatedBaseUrl, + qrcode: &SecretValue, + verify_code: Option<&SecretValue>, + ) -> Result { + let mut url = base_url.join("ilink/bot/get_qrcode_status")?; + { + let mut query = url.query_pairs_mut(); + query.append_pair("qrcode", qrcode.expose()); + if let Some(verify_code) = verify_code { + query.append_pair("verify_code", verify_code.expose()); + } + } + let request = self + .http + .get(url) + .headers(self.identity.application_headers()?); + let response: PollQrResponse = match self + .get_json(request, DEFAULT_QR_POLL_TIMEOUT, "poll_qr") + .await + { + Ok(response) => response, + Err(error) if retriable_poll_error(&error) => return Ok(PollQrResponse::waiting()), + Err(error) => return Err(error), + }; + match response.status { + super::types::QrCodeStatus::Unknown => { + return Err(IlinkError::InvalidResponse("poll_qr")); + } + super::types::QrCodeStatus::ScanedButRedirect => { + let redirect_host = response + .redirect_host + .as_deref() + .ok_or(IlinkError::InvalidResponse("poll_qr"))?; + self.host_policy.validate_redirect_host(redirect_host)?; + } + super::types::QrCodeStatus::Confirmed => { + if response.bot_token.is_none() + || response.ilink_bot_id.is_none() + || response.ilink_user_id.is_none() + { + return Err(IlinkError::InvalidResponse("poll_qr")); + } + let account_base_url = response + .baseurl + .as_deref() + .ok_or(IlinkError::InvalidResponse("poll_qr"))?; + self.host_policy.validate(account_base_url)?; + } + _ => {} + } + Ok(response) + } +} + +fn retriable_poll_error(error: &IlinkError) -> bool { + matches!( + error, + IlinkError::Timeout | IlinkError::Transport | IlinkError::HttpStatus(408 | 429 | 500..=599) + ) +} diff --git a/src/ilink/messages.rs b/src/ilink/messages.rs new file mode 100644 index 0000000..d1d27ad --- /dev/null +++ b/src/ilink/messages.rs @@ -0,0 +1,71 @@ +//! Outbound iLink message protocol operations. + +use super::auth::SecretValue; +use super::client::{ensure_api_success, IlinkAuth, IlinkError, DEFAULT_API_TIMEOUT}; +use super::transport::IlinkClient; +use super::types::{ + MessageItem, OutboundWeixinMessage, SendMessageRequest, SendMessageResponse, TextItem, + MESSAGE_ITEM_TYPE_TEXT, MESSAGE_STATE_FINISH, MESSAGE_TYPE_BOT, +}; + +const MAX_TEXT_BYTES: usize = 16 * 1024; +const MAX_CLIENT_ID_BYTES: usize = 128; +const MAX_RUN_ID_BYTES: usize = 128; + +impl IlinkClient { + pub(super) async fn send_text_request( + &self, + auth: &IlinkAuth, + recipient: &SecretValue, + context_token: Option<&SecretValue>, + client_id: &str, + run_id: Option<&str>, + text: &str, + ) -> Result { + validate_bounded_text(text, "message text", MAX_TEXT_BYTES)?; + validate_bounded_text(client_id, "client id", MAX_CLIENT_ID_BYTES)?; + if let Some(run_id) = run_id { + validate_bounded_text(run_id, "run id", MAX_RUN_ID_BYTES)?; + } + let request = self.authenticated_post(auth, "ilink/bot/sendmessage")?; + let response: SendMessageResponse = self + .post_json( + request, + &SendMessageRequest { + msg: OutboundWeixinMessage { + from_user_id: String::new(), + to_user_id: recipient.clone(), + client_id: client_id.to_string(), + message_type: MESSAGE_TYPE_BOT, + message_state: MESSAGE_STATE_FINISH, + item_list: vec![MessageItem { + item_type: Some(MESSAGE_ITEM_TYPE_TEXT), + text_item: Some(TextItem { + text: Some(text.to_string()), + }), + ..MessageItem::default() + }], + context_token: context_token.cloned(), + run_id: run_id.map(str::to_string), + }, + base_info: self.identity.base_info(), + }, + DEFAULT_API_TIMEOUT, + "send_message", + ) + .await?; + ensure_api_success("send_message", response.ret, None)?; + Ok(response) + } +} + +fn validate_bounded_text( + value: &str, + field: &'static str, + max_bytes: usize, +) -> Result<(), IlinkError> { + if value.is_empty() || value.len() > max_bytes || value.contains('\0') { + return Err(IlinkError::InvalidConfiguration(field)); + } + Ok(()) +} diff --git a/src/ilink/mod.rs b/src/ilink/mod.rs new file mode 100644 index 0000000..cd57c29 --- /dev/null +++ b/src/ilink/mod.rs @@ -0,0 +1,145 @@ +//! Tencent Weixin iLink protocol support. +//! +//! This module implements the HTTP/JSON contract used by Tencent's +//! `openclaw-weixin` SDK. Product applications remain responsible for storing +//! credentials, exposing user-facing APIs, and deciding what remote actions a +//! bound Weixin account may perform. + +mod auth; +mod client; +mod login; +mod messages; +mod transport; +mod types; +mod updates; +mod url_policy; + +use std::sync::Arc; + +use crate::{BootError, Module, ProviderDefinition, ProviderToken, Result as BootResult}; + +pub use auth::{ClientVersionError, SecretValue, SecretValueError}; +pub use client::{IlinkAuth, IlinkError}; +pub use transport::{IlinkClient, IlinkLoginTransport, IlinkMessagingTransport}; +pub use types::{ + CreateQrResponse, GetConfigResponse, GetUpdatesResponse, NotifyResponse, PollQrResponse, + QrCodeStatus, SendMessageResponse, SendTypingResponse, WeixinMessage, MESSAGE_STATE_FINISH, + MESSAGE_TYPE_USER, +}; +pub use url_policy::{IlinkUrlError, ValidatedBaseUrl}; + +/// Tencent's public iLink application identity used by the official SDK. +pub const WEIXIN_ILINK_APP_ID: &str = "bot"; + +/// Tencent's iLink bot type used by the official Weixin channel. +pub const WEIXIN_ILINK_BOT_TYPE: &str = "3"; + +/// Wire-contract version verified against Tencent `openclaw-weixin` v2.4.6. +pub const WEIXIN_ILINK_PROTOCOL_VERSION: &str = "2.4.6"; + +/// Fixed Tencent endpoint used to create and initially poll login QR codes. +pub const WEIXIN_ILINK_BASE_URL: &str = "https://ilinkai.weixin.qq.com"; + +const PRIMARY_ILINK_HOST: &str = "ilinkai.weixin.qq.com"; +const TENCENT_ILINK_HOST_SUFFIX: &str = "qq.com"; + +/// Configuration used to build an [`IlinkClient`]. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct IlinkClientOptions { + app_id: String, + bot_type: String, + protocol_version: String, + bot_agent: String, + base_url: String, + allowed_hosts: Vec, + allowed_host_suffixes: Vec, +} + +impl IlinkClientOptions { + /// Build the Tencent-compatible defaults for a product-specific bot agent. + /// + /// `bot_agent` identifies the host product for observability. It does not + /// replace Tencent's fixed `iLink-App-Id` protocol header. + pub fn weixin(bot_agent: impl Into) -> Self { + Self { + app_id: WEIXIN_ILINK_APP_ID.to_string(), + bot_type: WEIXIN_ILINK_BOT_TYPE.to_string(), + protocol_version: WEIXIN_ILINK_PROTOCOL_VERSION.to_string(), + bot_agent: bot_agent.into(), + base_url: WEIXIN_ILINK_BASE_URL.to_string(), + allowed_hosts: vec![PRIMARY_ILINK_HOST.to_string()], + allowed_host_suffixes: vec![TENCENT_ILINK_HOST_SUFFIX.to_string()], + } + } + + /// Add an exact HTTPS host that Tencent may return for account routing. + pub fn with_allowed_host(mut self, host: impl Into) -> Self { + let host = host.into(); + if !self + .allowed_hosts + .iter() + .any(|current| current.eq_ignore_ascii_case(&host)) + { + self.allowed_hosts.push(host); + } + self + } +} + +impl IlinkClient { + /// Build a concrete iLink client from validated options. + pub fn from_options(options: IlinkClientOptions) -> Result { + let identity = client::IlinkClientIdentity::new( + options.app_id, + options.bot_type, + &options.protocol_version, + options.bot_agent, + )?; + let host_policy = url_policy::IlinkHostPolicy::production_with_suffixes( + &options.allowed_hosts, + &options.allowed_host_suffixes, + )?; + Self::new(identity, host_policy, &options.base_url) + } + + /// Build the Tencent-compatible Weixin client used by product hosts. + pub fn weixin(bot_agent: impl Into) -> Result { + Self::from_options(IlinkClientOptions::weixin(bot_agent)) + } +} + +/// A3S Boot module that exports one validated [`IlinkClient`] provider. +#[derive(Clone, Debug)] +pub struct IlinkModule { + options: IlinkClientOptions, +} + +impl IlinkModule { + pub fn new(options: IlinkClientOptions) -> Self { + Self { options } + } + + pub fn weixin(bot_agent: impl Into) -> Self { + Self::new(IlinkClientOptions::weixin(bot_agent)) + } +} + +impl Module for IlinkModule { + fn name(&self) -> &'static str { + "a3s-boot-ilink" + } + + fn providers(&self) -> BootResult> { + let client = IlinkClient::from_options(self.options.clone()).map_err(|error| { + BootError::Internal(format!("failed to configure iLink client: {error}")) + })?; + Ok(vec![ProviderDefinition::from_arc(Arc::new(client))]) + } + + fn exports(&self) -> BootResult> { + Ok(vec![ProviderToken::of::()]) + } +} + +#[cfg(test)] +mod tests; diff --git a/src/ilink/tests.rs b/src/ilink/tests.rs new file mode 100644 index 0000000..586242b --- /dev/null +++ b/src/ilink/tests.rs @@ -0,0 +1,818 @@ +//! Contract tests derived from Tencent openclaw-weixin v2.4.6 wire behavior. + +use std::collections::HashMap; +use std::time::Duration; + +use axum::extract::{OriginalUri, Query, State}; +use axum::http::{HeaderMap, StatusCode}; +use axum::response::Redirect; +use axum::routing::{get, post}; +use axum::{Json, Router}; +use base64::Engine as _; +use serde_json::{json, Value}; +use tokio::sync::mpsc; + +use super::auth::{pack_client_version, SecretValue}; +use super::client::{IlinkAuth, IlinkClientIdentity, IlinkError}; +use super::transport::{IlinkClient, IlinkLoginTransport, IlinkMessagingTransport}; +use super::types::{GetUpdatesResponse, PollQrResponse, QrCodeStatus}; +use super::updates::validate_updates_response; +use super::url_policy::IlinkHostPolicy; +use super::{ + WEIXIN_ILINK_APP_ID, WEIXIN_ILINK_BASE_URL, WEIXIN_ILINK_BOT_TYPE, + WEIXIN_ILINK_PROTOCOL_VERSION, +}; + +#[derive(Debug)] +struct CapturedRequest { + operation: &'static str, + headers: HeaderMap, + query: HashMap, + body: Option, +} + +#[derive(Clone)] +struct CaptureState { + sender: mpsc::UnboundedSender, +} + +struct MockServer { + origin: String, + task: tokio::task::JoinHandle<()>, +} + +impl Drop for MockServer { + fn drop(&mut self) { + self.task.abort(); + } +} + +async fn spawn_mock_server(app: Router) -> MockServer { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind mock iLink server"); + let origin = format!("http://{}/", listener.local_addr().expect("mock address")); + let task = tokio::spawn(async move { + axum::serve(listener, app) + .await + .expect("serve mock iLink requests"); + }); + MockServer { origin, task } +} + +fn test_transport(origin: &str) -> IlinkClient { + let identity = IlinkClientIdentity::new("a3s-test-app", "a3s", "1.0.11", "A3S/0.9.7") + .expect("valid test identity"); + let policy = IlinkHostPolicy::for_test_origin(origin).expect("loopback test origin"); + IlinkClient::new(identity, policy, origin).expect("test transport") +} + +fn test_auth(transport: &IlinkClient, origin: &str) -> IlinkAuth { + IlinkAuth { + base_url: transport + .validate_account_base_url(origin) + .expect("validated account base URL"), + bot_token: SecretValue::new("bot-token-canary").expect("bot token"), + } +} + +async fn next_capture(receiver: &mut mpsc::UnboundedReceiver) -> CapturedRequest { + tokio::time::timeout(Duration::from_secs(1), receiver.recv()) + .await + .expect("mock request timeout") + .expect("mock request channel closed") +} + +fn assert_application_headers(headers: &HeaderMap) { + assert_eq!( + headers + .get("ilink-app-id") + .and_then(|value| value.to_str().ok()), + Some("a3s-test-app") + ); + assert_eq!( + headers + .get("ilink-app-clientversion") + .and_then(|value| value.to_str().ok()), + Some("65547") + ); +} + +fn assert_post_headers(headers: &HeaderMap, authenticated: bool) { + assert_application_headers(headers); + assert_eq!( + headers + .get("authorizationtype") + .and_then(|value| value.to_str().ok()), + Some("ilink_bot_token") + ); + assert_eq!( + headers + .get("content-type") + .and_then(|value| value.to_str().ok()), + Some("application/json") + ); + let encoded_uin = headers + .get("x-wechat-uin") + .and_then(|value| value.to_str().ok()) + .expect("random Weixin UIN header"); + let decoded_uin = base64::engine::general_purpose::STANDARD + .decode(encoded_uin) + .expect("base64 Weixin UIN"); + std::str::from_utf8(&decoded_uin) + .expect("UTF-8 Weixin UIN") + .parse::() + .expect("decimal uint32 Weixin UIN"); + let authorization = headers + .get("authorization") + .and_then(|value| value.to_str().ok()); + if authenticated { + assert_eq!(authorization, Some("Bearer bot-token-canary")); + } else { + assert_eq!(authorization, None); + } +} + +async fn capture_create_qr( + State(state): State, + headers: HeaderMap, + Query(query): Query>, + Json(body): Json, +) -> Json { + state + .sender + .send(CapturedRequest { + operation: "create_qr", + headers, + query, + body: Some(body), + }) + .expect("capture create QR request"); + Json(json!({ + "qrcode": "qr-canary", + "qrcode_img_content": "data:image/png;base64,cXItY2FuYXJ5" + })) +} + +async fn capture_poll_qr( + State(state): State, + headers: HeaderMap, + Query(query): Query>, +) -> Json { + state + .sender + .send(CapturedRequest { + operation: "poll_qr", + headers, + query, + body: None, + }) + .expect("capture poll QR request"); + Json(json!({ "status": "wait" })) +} + +async fn capture_get_updates( + State(state): State, + headers: HeaderMap, + Json(body): Json, +) -> Json { + state + .sender + .send(CapturedRequest { + operation: "get_updates", + headers, + query: HashMap::new(), + body: Some(body), + }) + .expect("capture get updates request"); + Json(json!({ + "ret": 0, + "msgs": [], + "get_updates_buf": "next-cursor-canary", + "longpolling_timeout_ms": 35_000 + })) +} + +async fn capture_send_message( + State(state): State, + headers: HeaderMap, + Json(body): Json, +) -> Json { + state + .sender + .send(CapturedRequest { + operation: "send_message", + headers, + query: HashMap::new(), + body: Some(body), + }) + .expect("capture send message request"); + Json(json!({ "ret": 0 })) +} + +async fn capture_control_request( + State(state): State, + OriginalUri(uri): OriginalUri, + headers: HeaderMap, + Json(body): Json, +) -> Json { + let operation = match uri.path() { + "/ilink/bot/getconfig" => "get_config", + "/ilink/bot/sendtyping" => "send_typing", + "/ilink/bot/msg/notifystart" => "notify_start", + "/ilink/bot/msg/notifystop" => "notify_stop", + path => panic!("unexpected control path: {path}"), + }; + state + .sender + .send(CapturedRequest { + operation, + headers, + query: HashMap::new(), + body: Some(body), + }) + .expect("capture control request"); + Json(json!({ "ret": 0, "typing_ticket": "typing-ticket-canary" })) +} + +#[test] +fn weixin_ilink_contract_packs_client_versions() { + assert_eq!(pack_client_version("1.0.11").unwrap(), 0x0001_000b); + assert_eq!(pack_client_version("2.4.6").unwrap(), 0x0002_0406); + assert!(pack_client_version("1.2").is_err()); + assert!(pack_client_version("1.256.0").is_err()); + assert!(pack_client_version("not-a-version").is_err()); +} + +#[test] +fn weixin_ilink_defaults_match_tencent_sdk_v2_4_6() { + let client = IlinkClient::weixin("A3S/0.10.1").expect("Tencent-compatible client"); + let headers = client + .identity + .application_headers() + .expect("application headers"); + + assert_eq!(WEIXIN_ILINK_APP_ID, "bot"); + assert_eq!(WEIXIN_ILINK_BOT_TYPE, "3"); + assert_eq!(WEIXIN_ILINK_PROTOCOL_VERSION, "2.4.6"); + assert_eq!(WEIXIN_ILINK_BASE_URL, "https://ilinkai.weixin.qq.com"); + assert_eq!(client.identity.bot_type, "3"); + assert_eq!(client.identity.channel_version, "2.4.6"); + assert_eq!(client.identity.bot_agent, "A3S/0.10.1"); + assert_eq!( + headers + .get("ilink-app-id") + .and_then(|value| value.to_str().ok()), + Some("bot") + ); + assert_eq!( + headers + .get("ilink-app-clientversion") + .and_then(|value| value.to_str().ok()), + Some("132102") + ); +} + +#[test] +fn weixin_ilink_module_exports_the_boot_protocol_client() { + let module = crate::TestingModule::builder() + .import(super::IlinkModule::weixin("A3S/0.10.1")) + .compile() + .expect("compile iLink module"); + let client = module.get::().expect("resolve iLink client"); + + assert_eq!(client.identity.bot_type, WEIXIN_ILINK_BOT_TYPE); + assert_eq!( + client.identity.channel_version, + WEIXIN_ILINK_PROTOCOL_VERSION + ); +} + +#[test] +fn weixin_ilink_contract_redacts_secret_debug_output() { + let secret = SecretValue::new("canary-bot-token").unwrap(); + + let rendered = format!("{secret:?}"); + + assert_eq!(rendered, "SecretValue([REDACTED])"); + assert!(!rendered.contains("canary-bot-token")); + assert_eq!(secret.expose(), "canary-bot-token"); +} + +#[test] +fn weixin_ilink_contract_accepts_known_qr_states_and_contains_unknown_states() { + let redirected: PollQrResponse = serde_json::from_value(serde_json::json!({ + "status": "scaned_but_redirect", + "redirect_host": "https://ilinkai.weixin.qq.com" + })) + .unwrap(); + assert_eq!(redirected.status, QrCodeStatus::ScanedButRedirect); + + let unknown: PollQrResponse = serde_json::from_value(serde_json::json!({ + "status": "future_state" + })) + .unwrap(); + assert_eq!(unknown.status, QrCodeStatus::Unknown); +} + +#[test] +fn weixin_ilink_contract_deserializes_text_updates_without_exposing_tokens_in_debug() { + let response: GetUpdatesResponse = serde_json::from_value(serde_json::json!({ + "ret": 0, + "msgs": [{ + "seq": 7, + "message_id": 42, + "from_user_id": "owner-canary", + "message_type": 1, + "item_list": [{ + "type": 1, + "text_item": { "text": "进度" } + }], + "context_token": "context-canary" + }], + "get_updates_buf": "cursor-canary", + "longpolling_timeout_ms": 35000 + })) + .unwrap(); + + assert_eq!(response.messages.len(), 1); + assert_eq!(response.messages[0].message_id, Some(42)); + assert_eq!(response.messages[0].text(), Some("进度")); + assert_eq!(response.long_polling_timeout_ms, Some(35_000)); + let rendered = format!("{response:?}"); + assert!(!rendered.contains("owner-canary")); + assert!(!rendered.contains("context-canary")); + assert!(!rendered.contains("cursor-canary")); +} + +#[test] +fn weixin_ilink_contract_enforces_production_base_url_policy() { + let policy = IlinkHostPolicy::production(["ilinkai.weixin.qq.com"]).unwrap(); + + assert!(policy.validate("https://ilinkai.weixin.qq.com").is_ok()); + assert!(policy + .validate("https://ilinkai.weixin.qq.com:443/region/") + .is_ok()); + for rejected in [ + "http://ilinkai.weixin.qq.com", + "https://127.0.0.1", + "https://user@ilinkai.weixin.qq.com", + "https://ilinkai.weixin.qq.com:8443", + "https://ilinkai.weixin.qq.com.attacker.example", + "https://attacker.example/?next=ilinkai.weixin.qq.com", + "https://ilinkai.weixin.qq.com/#fragment", + ] { + assert!( + policy.validate(rejected).is_err(), + "unexpectedly accepted {rejected}" + ); + } + assert!(policy + .validate_redirect_host("ilinkai.weixin.qq.com") + .is_ok()); + for rejected in [ + "https://ilinkai.weixin.qq.com", + "user@ilinkai.weixin.qq.com", + "ilinkai.weixin.qq.com:8443", + "ilinkai.weixin.qq.com.attacker.example", + "127.0.0.1", + ] { + assert!( + policy.validate_redirect_host(rejected).is_err(), + "unexpectedly accepted redirect host {rejected}" + ); + } +} + +#[test] +fn weixin_ilink_defaults_allow_tencent_idc_hosts_without_suffix_confusion() { + let client = IlinkClient::weixin("A3S/0.10.1").expect("Tencent-compatible client"); + + assert!(client + .validate_account_base_url("https://shard.weixin.qq.com/account/") + .is_ok()); + assert!(client.validate_redirect_host("shard.weixin.qq.com").is_ok()); + for rejected in [ + "https://qq.com.attacker.example/", + "https://weixin.qq.com.attacker.example/", + "https://127.0.0.1/", + ] { + assert!( + client.validate_account_base_url(rejected).is_err(), + "unexpectedly accepted {rejected}" + ); + } +} + +#[test] +fn weixin_ilink_contract_bounds_update_arrays_text_and_server_timeout() { + let too_many_messages: GetUpdatesResponse = serde_json::from_value(json!({ + "ret": 0, + "msgs": (0..257).map(|index| json!({ "message_id": index })).collect::>(), + "get_updates_buf": "cursor-canary" + })) + .unwrap(); + assert_eq!( + validate_updates_response(&too_many_messages), + Err(IlinkError::InvalidResponse("get_updates")) + ); + + let too_many_items: GetUpdatesResponse = serde_json::from_value(json!({ + "ret": 0, + "msgs": [{ + "message_id": 1, + "item_list": (0..33).map(|_| json!({ "type": 0 })).collect::>() + }], + "get_updates_buf": "cursor-canary" + })) + .unwrap(); + assert_eq!( + validate_updates_response(&too_many_items), + Err(IlinkError::InvalidResponse("get_updates")) + ); + + let oversized_text: GetUpdatesResponse = serde_json::from_value(json!({ + "ret": 0, + "msgs": [{ + "message_id": 1, + "item_list": [{ + "type": 1, + "text_item": { "text": "x".repeat(16 * 1024 + 1) } + }] + }], + "get_updates_buf": "cursor-canary" + })) + .unwrap(); + assert_eq!( + validate_updates_response(&oversized_text), + Err(IlinkError::InvalidResponse("get_updates")) + ); + + let unbounded_timeout: GetUpdatesResponse = serde_json::from_value(json!({ + "ret": 0, + "msgs": [], + "get_updates_buf": "cursor-canary", + "longpolling_timeout_ms": 60_001 + })) + .unwrap(); + assert_eq!( + validate_updates_response(&unbounded_timeout), + Err(IlinkError::InvalidResponse("get_updates")) + ); +} + +#[tokio::test] +async fn weixin_ilink_contract_sends_qr_update_and_text_requests() { + let (sender, mut receiver) = mpsc::unbounded_channel(); + let app = Router::new() + .route("/ilink/bot/get_bot_qrcode", post(capture_create_qr)) + .route("/ilink/bot/get_qrcode_status", get(capture_poll_qr)) + .route("/ilink/bot/getupdates", post(capture_get_updates)) + .route("/ilink/bot/sendmessage", post(capture_send_message)) + .route("/ilink/bot/getconfig", post(capture_control_request)) + .route("/ilink/bot/sendtyping", post(capture_control_request)) + .route("/ilink/bot/msg/notifystart", post(capture_control_request)) + .route("/ilink/bot/msg/notifystop", post(capture_control_request)) + .with_state(CaptureState { sender }); + let server = spawn_mock_server(app).await; + let transport = test_transport(&server.origin); + let auth = test_auth(&transport, &server.origin); + + let local_token = SecretValue::new("local-token-canary").expect("local token"); + let created = transport + .create_qr(&[local_token]) + .await + .expect("create QR response"); + assert_eq!(created.qrcode.expose(), "qr-canary"); + let qr_base_url = transport + .validate_account_base_url(&server.origin) + .expect("QR polling base URL"); + let qr_code = SecretValue::new("qr-canary").expect("QR code"); + let verify_code = SecretValue::new("123456").expect("verification code"); + let polled = transport + .poll_qr(&qr_base_url, &qr_code, Some(&verify_code)) + .await + .expect("poll QR response"); + assert_eq!(polled.status, QrCodeStatus::Wait); + let updates = transport + .get_updates(&auth, "cursor-canary", Duration::from_secs(35)) + .await + .expect("get updates response"); + assert_eq!( + updates.update_cursor.as_ref().map(SecretValue::expose), + Some("next-cursor-canary") + ); + let recipient = SecretValue::new("owner-canary").expect("owner ID"); + let context_token = SecretValue::new("context-canary").expect("context token"); + transport + .send_text( + &auth, + &recipient, + Some(&context_token), + "client-id-canary", + Some("run-id-canary"), + "任务仍在执行", + ) + .await + .expect("send message response"); + let config = transport + .get_config(&auth, Some(&recipient), Some(&context_token)) + .await + .expect("get config response"); + assert_eq!( + config.typing_ticket.as_ref().map(SecretValue::expose), + Some("typing-ticket-canary") + ); + let typing_ticket = SecretValue::new("typing-ticket-canary").expect("typing ticket"); + transport + .send_typing(&auth, &recipient, &typing_ticket, 1) + .await + .expect("send typing response"); + transport + .notify_start(&auth) + .await + .expect("notify start response"); + transport + .notify_stop(&auth) + .await + .expect("notify stop response"); + + let create = next_capture(&mut receiver).await; + assert_eq!(create.operation, "create_qr"); + assert_post_headers(&create.headers, false); + assert_eq!( + create.query.get("bot_type").map(String::as_str), + Some("a3s") + ); + assert_eq!( + create.body, + Some(json!({ "local_token_list": ["local-token-canary"] })) + ); + + let poll = next_capture(&mut receiver).await; + assert_eq!(poll.operation, "poll_qr"); + assert_application_headers(&poll.headers); + for absent in [ + "authorizationtype", + "authorization", + "content-type", + "x-wechat-uin", + ] { + assert!(!poll.headers.contains_key(absent)); + } + assert_eq!( + poll.query.get("qrcode").map(String::as_str), + Some("qr-canary") + ); + assert_eq!( + poll.query.get("verify_code").map(String::as_str), + Some("123456") + ); + assert_eq!(poll.body, None); + + let get_updates = next_capture(&mut receiver).await; + assert_eq!(get_updates.operation, "get_updates"); + assert_post_headers(&get_updates.headers, true); + assert_eq!( + get_updates.body, + Some(json!({ + "get_updates_buf": "cursor-canary", + "base_info": { + "channel_version": "1.0.11", + "bot_agent": "A3S/0.9.7" + } + })) + ); + + let send_message = next_capture(&mut receiver).await; + assert_eq!(send_message.operation, "send_message"); + assert_post_headers(&send_message.headers, true); + assert_eq!( + send_message.body, + Some(json!({ + "msg": { + "from_user_id": "", + "to_user_id": "owner-canary", + "client_id": "client-id-canary", + "message_type": 2, + "message_state": 2, + "item_list": [{ + "type": 1, + "text_item": { "text": "任务仍在执行" } + }], + "context_token": "context-canary", + "run_id": "run-id-canary" + }, + "base_info": { + "channel_version": "1.0.11", + "bot_agent": "A3S/0.9.7" + } + })) + ); + + let get_config = next_capture(&mut receiver).await; + assert_eq!(get_config.operation, "get_config"); + assert_post_headers(&get_config.headers, true); + assert_eq!( + get_config.body, + Some(json!({ + "base_info": { + "channel_version": "1.0.11", + "bot_agent": "A3S/0.9.7" + }, + "ilink_user_id": "owner-canary", + "context_token": "context-canary" + })) + ); + + let send_typing = next_capture(&mut receiver).await; + assert_eq!(send_typing.operation, "send_typing"); + assert_post_headers(&send_typing.headers, true); + assert_eq!( + send_typing.body, + Some(json!({ + "ilink_user_id": "owner-canary", + "typing_ticket": "typing-ticket-canary", + "status": 1, + "base_info": { + "channel_version": "1.0.11", + "bot_agent": "A3S/0.9.7" + } + })) + ); + + for operation in ["notify_start", "notify_stop"] { + let notify = next_capture(&mut receiver).await; + assert_eq!(notify.operation, operation); + assert_post_headers(¬ify.headers, true); + assert_eq!( + notify.body, + Some(json!({ + "base_info": { + "channel_version": "1.0.11", + "bot_agent": "A3S/0.9.7" + } + })) + ); + } +} + +#[tokio::test] +async fn weixin_ilink_contract_fails_closed_on_unknown_qr_state() { + let app = Router::new().route( + "/ilink/bot/get_qrcode_status", + get(|| async { Json(json!({ "status": "future_state" })) }), + ); + let server = spawn_mock_server(app).await; + let transport = test_transport(&server.origin); + let qr_base_url = transport + .validate_account_base_url(&server.origin) + .expect("QR base URL"); + let qr_code = SecretValue::new("qr-canary").expect("QR code"); + + let error = transport + .poll_qr(&qr_base_url, &qr_code, None) + .await + .expect_err("unknown QR state must fail closed"); + + assert_eq!(error, IlinkError::InvalidResponse("poll_qr")); +} + +#[tokio::test] +async fn weixin_ilink_contract_maps_stale_credentials() { + let app = Router::new().route( + "/ilink/bot/getupdates", + post(|| async { Json(json!({ "errcode": -14, "msgs": [] })) }), + ); + let server = spawn_mock_server(app).await; + let transport = test_transport(&server.origin); + let auth = test_auth(&transport, &server.origin); + + let error = transport + .get_updates(&auth, "cursor-canary", Duration::from_secs(35)) + .await + .expect_err("stale credential must fail closed"); + + assert_eq!(error, IlinkError::StaleCredential); +} + +#[tokio::test] +async fn weixin_ilink_contract_treats_poll_gateway_errors_as_waiting() { + let app = Router::new().route( + "/ilink/bot/get_qrcode_status", + get(|| async { StatusCode::SERVICE_UNAVAILABLE }), + ); + let server = spawn_mock_server(app).await; + let transport = test_transport(&server.origin); + let qr_base_url = transport + .validate_account_base_url(&server.origin) + .expect("QR base URL"); + let qr_code = SecretValue::new("qr-canary").expect("QR code"); + + let response = transport + .poll_qr(&qr_base_url, &qr_code, None) + .await + .expect("gateway failures are transient while QR login is active"); + + assert_eq!(response.status, QrCodeStatus::Wait); +} + +#[tokio::test] +async fn weixin_ilink_contract_treats_update_timeout_as_an_empty_poll() { + let app = Router::new().route( + "/ilink/bot/getupdates", + post(|| async { + tokio::time::sleep(Duration::from_millis(50)).await; + Json(json!({ + "ret": 0, + "msgs": [], + "get_updates_buf": "server-cursor" + })) + }), + ); + let server = spawn_mock_server(app).await; + let transport = test_transport(&server.origin); + let auth = test_auth(&transport, &server.origin); + + let response = transport + .get_updates(&auth, "cursor-canary", Duration::from_millis(10)) + .await + .expect("long-poll timeout is normal control flow"); + + assert!(response.messages.is_empty()); + assert_eq!( + response.update_cursor.as_ref().map(SecretValue::expose), + Some("cursor-canary") + ); +} + +#[tokio::test] +async fn weixin_ilink_contract_rejects_redirect_status_oversize_and_timeout() { + let app = Router::new() + .route( + "/redirect", + get(|| async { Redirect::temporary("/target") }), + ) + .route("/status", get(|| async { StatusCode::SERVICE_UNAVAILABLE })) + .route( + "/large", + get(|| async { Json(json!({ "payload": "too large" })) }), + ) + .route( + "/slow", + get(|| async { + tokio::time::sleep(Duration::from_millis(100)).await; + Json(json!({})) + }), + ); + let server = spawn_mock_server(app).await; + let mut transport = test_transport(&server.origin); + + let redirect_url = transport + .qr_base_url + .join("redirect") + .expect("redirect URL"); + let redirect_error = transport + .get_json::( + transport.http.get(redirect_url), + Duration::from_secs(1), + "redirect", + ) + .await + .expect_err("redirect must not be followed"); + assert_eq!(redirect_error, IlinkError::HttpStatus(307)); + + let status_url = transport.qr_base_url.join("status").expect("status URL"); + let status_error = transport + .get_json::( + transport.http.get(status_url), + Duration::from_secs(1), + "status", + ) + .await + .expect_err("non-success status must fail"); + assert_eq!(status_error, IlinkError::HttpStatus(503)); + + transport.max_response_bytes = 8; + let large_url = transport.qr_base_url.join("large").expect("large URL"); + let large_error = transport + .get_json::( + transport.http.get(large_url), + Duration::from_secs(1), + "large", + ) + .await + .expect_err("oversized response must fail"); + assert_eq!(large_error, IlinkError::ResponseTooLarge); + + let slow_url = transport.qr_base_url.join("slow").expect("slow URL"); + let timeout_error = transport + .get_json::( + transport.http.get(slow_url), + Duration::from_millis(10), + "slow", + ) + .await + .expect_err("slow response must time out"); + assert_eq!(timeout_error, IlinkError::Timeout); +} diff --git a/src/ilink/transport.rs b/src/ilink/transport.rs new file mode 100644 index 0000000..4f2bc95 --- /dev/null +++ b/src/ilink/transport.rs @@ -0,0 +1,311 @@ +//! HTTP transport and object-safe iLink client boundaries. + +use std::time::Duration; + +use async_trait::async_trait; +use futures_util::StreamExt; +use reqwest::redirect::Policy; +use serde::de::DeserializeOwned; + +use super::auth::SecretValue; +use super::client::{ + IlinkAuth, IlinkClientIdentity, IlinkError, MAX_LONG_POLL_TIMEOUT, MAX_RESPONSE_BYTES, +}; +use super::types::{ + CreateQrResponse, GetUpdatesResponse, NotifyResponse, PollQrResponse, SendMessageResponse, +}; +use super::types::{GetConfigResponse, SendTypingResponse}; +use super::url_policy::{IlinkHostPolicy, ValidatedBaseUrl}; + +#[async_trait] +pub trait IlinkLoginTransport: Send + Sync { + fn qr_base_url(&self) -> ValidatedBaseUrl; + + fn validate_account_base_url(&self, base_url: &str) -> Result; + + fn validate_redirect_host(&self, redirect_host: &str) -> Result; + + async fn create_qr(&self, local_tokens: &[SecretValue]) + -> Result; + + async fn poll_qr( + &self, + base_url: &ValidatedBaseUrl, + qrcode: &SecretValue, + verify_code: Option<&SecretValue>, + ) -> Result; +} + +#[async_trait] +pub trait IlinkMessagingTransport: Send + Sync { + fn validate_account_base_url(&self, base_url: &str) -> Result; + + async fn get_updates( + &self, + auth: &IlinkAuth, + update_cursor: &str, + long_poll_timeout: Duration, + ) -> Result; + + async fn send_text( + &self, + auth: &IlinkAuth, + recipient: &SecretValue, + context_token: Option<&SecretValue>, + client_id: &str, + run_id: Option<&str>, + text: &str, + ) -> Result; + + async fn get_config( + &self, + auth: &IlinkAuth, + owner_id: Option<&SecretValue>, + context_token: Option<&SecretValue>, + ) -> Result; + + async fn send_typing( + &self, + auth: &IlinkAuth, + owner_id: &SecretValue, + typing_ticket: &SecretValue, + status: i32, + ) -> Result; + + async fn notify_start(&self, auth: &IlinkAuth) -> Result; + + async fn notify_stop(&self, auth: &IlinkAuth) -> Result; +} + +pub struct IlinkClient { + pub(super) http: reqwest::Client, + pub(super) identity: IlinkClientIdentity, + pub(super) host_policy: IlinkHostPolicy, + pub(super) qr_base_url: ValidatedBaseUrl, + pub(super) max_response_bytes: usize, +} + +impl IlinkClient { + pub(super) fn new( + identity: IlinkClientIdentity, + host_policy: IlinkHostPolicy, + qr_base_url: &str, + ) -> Result { + let qr_base_url = host_policy.validate(qr_base_url)?; + let http = reqwest::Client::builder() + .redirect(Policy::none()) + .connect_timeout(Duration::from_secs(10)) + .no_proxy() + .build() + .map_err(|_| IlinkError::InvalidConfiguration("HTTP client"))?; + Ok(Self { + http, + identity, + host_policy, + qr_base_url, + max_response_bytes: MAX_RESPONSE_BYTES, + }) + } + + pub(super) fn validate_account_base_url( + &self, + base_url: &str, + ) -> Result { + self.host_policy.validate(base_url).map_err(Into::into) + } + + pub(super) fn validate_redirect_host( + &self, + redirect_host: &str, + ) -> Result { + self.host_policy + .validate_redirect_host(redirect_host) + .map_err(Into::into) + } + + pub(super) async fn get_json( + &self, + request: reqwest::RequestBuilder, + timeout: Duration, + operation: &'static str, + ) -> Result + where + T: DeserializeOwned, + { + self.execute_json(request.timeout(timeout), operation).await + } + + pub(super) async fn post_json( + &self, + request: reqwest::RequestBuilder, + body: &B, + timeout: Duration, + operation: &'static str, + ) -> Result + where + B: serde::Serialize + Sync, + T: DeserializeOwned, + { + self.execute_json(request.json(body).timeout(timeout), operation) + .await + } + + pub(super) async fn post_json_without_timeout( + &self, + request: reqwest::RequestBuilder, + body: &B, + operation: &'static str, + ) -> Result + where + B: serde::Serialize + Sync, + T: DeserializeOwned, + { + self.execute_json(request.json(body), operation).await + } + + async fn execute_json( + &self, + request: reqwest::RequestBuilder, + operation: &'static str, + ) -> Result + where + T: DeserializeOwned, + { + let response = request.send().await.map_err(map_reqwest_error)?; + if !response.status().is_success() { + return Err(IlinkError::HttpStatus(response.status().as_u16())); + } + if response + .content_length() + .is_some_and(|length| length > self.max_response_bytes as u64) + { + return Err(IlinkError::ResponseTooLarge); + } + let mut bytes = Vec::new(); + let mut stream = response.bytes_stream(); + while let Some(chunk) = stream.next().await { + let chunk = chunk.map_err(map_reqwest_error)?; + if bytes.len().saturating_add(chunk.len()) > self.max_response_bytes { + return Err(IlinkError::ResponseTooLarge); + } + bytes.extend_from_slice(&chunk); + } + serde_json::from_slice(&bytes).map_err(|_| IlinkError::InvalidResponse(operation)) + } + + pub(super) fn authenticated_post( + &self, + auth: &IlinkAuth, + endpoint: &str, + ) -> Result { + let url = auth.base_url.join(endpoint)?; + Ok(self + .http + .post(url) + .headers(self.identity.post_headers(Some(&auth.bot_token))?)) + } + + pub(super) fn bounded_long_poll_timeout(timeout: Duration) -> Result { + if timeout.is_zero() || timeout > MAX_LONG_POLL_TIMEOUT { + return Err(IlinkError::InvalidConfiguration("long poll timeout")); + } + Ok(timeout) + } +} + +fn map_reqwest_error(error: reqwest::Error) -> IlinkError { + if error.is_timeout() { + IlinkError::Timeout + } else { + IlinkError::Transport + } +} + +#[async_trait] +impl IlinkLoginTransport for IlinkClient { + fn qr_base_url(&self) -> ValidatedBaseUrl { + self.qr_base_url.clone() + } + + fn validate_account_base_url(&self, base_url: &str) -> Result { + self.validate_account_base_url(base_url) + } + + fn validate_redirect_host(&self, redirect_host: &str) -> Result { + self.validate_redirect_host(redirect_host) + } + + async fn create_qr( + &self, + local_tokens: &[SecretValue], + ) -> Result { + self.create_qr_request(local_tokens).await + } + + async fn poll_qr( + &self, + base_url: &ValidatedBaseUrl, + qrcode: &SecretValue, + verify_code: Option<&SecretValue>, + ) -> Result { + self.poll_qr_request(base_url, qrcode, verify_code).await + } +} + +#[async_trait] +impl IlinkMessagingTransport for IlinkClient { + fn validate_account_base_url(&self, base_url: &str) -> Result { + self.validate_account_base_url(base_url) + } + + async fn get_updates( + &self, + auth: &IlinkAuth, + update_cursor: &str, + long_poll_timeout: Duration, + ) -> Result { + self.get_updates_request(auth, update_cursor, long_poll_timeout) + .await + } + + async fn send_text( + &self, + auth: &IlinkAuth, + recipient: &SecretValue, + context_token: Option<&SecretValue>, + client_id: &str, + run_id: Option<&str>, + text: &str, + ) -> Result { + self.send_text_request(auth, recipient, context_token, client_id, run_id, text) + .await + } + + async fn get_config( + &self, + auth: &IlinkAuth, + owner_id: Option<&SecretValue>, + context_token: Option<&SecretValue>, + ) -> Result { + self.get_config_request(auth, owner_id, context_token).await + } + + async fn send_typing( + &self, + auth: &IlinkAuth, + owner_id: &SecretValue, + typing_ticket: &SecretValue, + status: i32, + ) -> Result { + self.send_typing_request(auth, owner_id, typing_ticket, status) + .await + } + + async fn notify_start(&self, auth: &IlinkAuth) -> Result { + self.notify_request(auth, true).await + } + + async fn notify_stop(&self, auth: &IlinkAuth) -> Result { + self.notify_request(auth, false).await + } +} diff --git a/src/ilink/types.rs b/src/ilink/types.rs new file mode 100644 index 0000000..23ff184 --- /dev/null +++ b/src/ilink/types.rs @@ -0,0 +1,292 @@ +//! JSON wire types used by the iLink protocol. + +use std::fmt; + +use serde::{Deserialize, Deserializer, Serialize}; + +use super::auth::SecretValue; + +pub const MESSAGE_TYPE_USER: i32 = 1; +pub(super) const MESSAGE_TYPE_BOT: i32 = 2; +pub(super) const MESSAGE_ITEM_TYPE_TEXT: i32 = 1; +pub const MESSAGE_STATE_FINISH: i32 = 2; + +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +pub(super) struct BaseInfo { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) channel_version: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) bot_agent: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +pub(super) struct CreateQrRequest { + pub(super) local_token_list: Vec, +} + +#[derive(Clone, Debug, PartialEq, Eq, Deserialize)] +pub struct CreateQrResponse { + pub qrcode: SecretValue, + pub qrcode_img_content: SecretValue, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum QrCodeStatus { + Wait, + Scaned, + Confirmed, + Expired, + ScanedButRedirect, + NeedVerifycode, + VerifyCodeBlocked, + BindedRedirect, + #[serde(other)] + Unknown, +} + +#[derive(Clone, Debug, PartialEq, Eq, Deserialize)] +pub struct PollQrResponse { + pub status: QrCodeStatus, + #[serde(default, deserialize_with = "deserialize_optional_secret")] + pub bot_token: Option, + #[serde(default, deserialize_with = "deserialize_optional_secret")] + pub ilink_bot_id: Option, + #[serde(default, deserialize_with = "deserialize_optional_secret")] + pub ilink_user_id: Option, + #[serde(default)] + pub baseurl: Option, + #[serde(default)] + pub redirect_host: Option, +} + +impl PollQrResponse { + pub(super) fn waiting() -> Self { + Self { + status: QrCodeStatus::Wait, + bot_token: None, + ilink_bot_id: None, + ilink_user_id: None, + baseurl: None, + redirect_host: None, + } + } +} + +#[derive(Clone, Default, PartialEq, Eq, Serialize)] +pub(super) struct GetUpdatesRequest { + pub(super) get_updates_buf: String, + #[serde(default)] + pub(super) base_info: BaseInfo, +} + +impl fmt::Debug for GetUpdatesRequest { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("GetUpdatesRequest") + .field("has_update_cursor", &!self.get_updates_buf.is_empty()) + .field("base_info", &self.base_info) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq, Deserialize)] +pub struct GetUpdatesResponse { + #[serde(default)] + pub(super) ret: Option, + #[serde(default)] + pub(super) errcode: Option, + #[serde(default)] + pub(super) errmsg: Option, + #[serde(default, rename = "msgs")] + pub messages: Vec, + #[serde( + default, + rename = "get_updates_buf", + deserialize_with = "deserialize_optional_secret" + )] + pub update_cursor: Option, + #[serde(default, rename = "longpolling_timeout_ms")] + pub long_polling_timeout_ms: Option, +} + +impl fmt::Debug for GetUpdatesResponse { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("GetUpdatesResponse") + .field("ret", &self.ret) + .field("errcode", &self.errcode) + .field("message_count", &self.messages.len()) + .field("has_update_cursor", &self.update_cursor.is_some()) + .field("long_polling_timeout_ms", &self.long_polling_timeout_ms) + .finish() + } +} + +#[derive(Clone, PartialEq, Eq, Deserialize)] +pub struct WeixinMessage { + #[serde(default)] + pub seq: Option, + #[serde(default)] + pub message_id: Option, + #[serde(default, deserialize_with = "deserialize_optional_secret")] + pub from_user_id: Option, + #[serde(default, deserialize_with = "deserialize_optional_secret")] + pub to_user_id: Option, + #[serde(default)] + pub client_id: Option, + #[serde(default)] + pub create_time_ms: Option, + #[serde(default)] + pub update_time_ms: Option, + #[serde(default, deserialize_with = "deserialize_optional_secret")] + pub session_id: Option, + #[serde(default, deserialize_with = "deserialize_optional_secret")] + pub group_id: Option, + #[serde(default)] + pub message_type: Option, + #[serde(default)] + pub message_state: Option, + #[serde(default)] + pub(super) item_list: Vec, + #[serde(default, deserialize_with = "deserialize_optional_secret")] + pub context_token: Option, + #[serde(default)] + pub run_id: Option, +} + +impl WeixinMessage { + pub fn text(&self) -> Option<&str> { + self.item_list.iter().find_map(|item| { + (item.item_type == Some(MESSAGE_ITEM_TYPE_TEXT)) + .then(|| item.text_item.as_ref()?.text.as_deref()) + .flatten() + }) + } +} + +impl fmt::Debug for WeixinMessage { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("WeixinMessage") + .field("seq", &self.seq) + .field("message_id", &self.message_id) + .field("message_type", &self.message_type) + .field("message_state", &self.message_state) + .field("item_count", &self.item_list.len()) + .field("has_context_token", &self.context_token.is_some()) + .finish() + } +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +pub(super) struct MessageItem { + #[serde(default, rename = "type", skip_serializing_if = "Option::is_none")] + pub(super) item_type: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) create_time_ms: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) update_time_ms: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) is_completed: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) msg_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) text_item: Option, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +pub(super) struct TextItem { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) text: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +pub(super) struct SendMessageRequest { + pub(super) msg: OutboundWeixinMessage, + pub(super) base_info: BaseInfo, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +pub(super) struct OutboundWeixinMessage { + pub(super) from_user_id: String, + pub(super) to_user_id: SecretValue, + pub(super) client_id: String, + pub(super) message_type: i32, + pub(super) message_state: i32, + pub(super) item_list: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub(super) context_token: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub(super) run_id: Option, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)] +pub struct SendMessageResponse { + #[serde(default)] + pub(super) ret: Option, + #[serde(default)] + pub(super) errmsg: Option, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)] +pub(super) struct GetConfigRequest { + #[serde(default)] + pub(super) base_info: BaseInfo, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) ilink_user_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(super) context_token: Option, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)] +pub struct GetConfigResponse { + #[serde(default)] + pub(super) ret: Option, + #[serde(default)] + pub(super) errmsg: Option, + #[serde(default, deserialize_with = "deserialize_optional_secret")] + pub(super) typing_ticket: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +pub(super) struct SendTypingRequest { + pub(super) ilink_user_id: SecretValue, + pub(super) typing_ticket: SecretValue, + pub(super) status: i32, + pub(super) base_info: BaseInfo, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)] +pub struct SendTypingResponse { + #[serde(default)] + pub(super) ret: Option, + #[serde(default)] + pub(super) errmsg: Option, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)] +pub(super) struct NotifyRequest { + #[serde(default)] + pub(super) base_info: BaseInfo, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)] +pub struct NotifyResponse { + #[serde(default)] + pub(super) ret: Option, + #[serde(default)] + pub(super) errmsg: Option, +} + +fn deserialize_optional_secret<'de, D>(deserializer: D) -> Result, D::Error> +where + D: Deserializer<'de>, +{ + let value = Option::::deserialize(deserializer)?; + value + .filter(|value| !value.is_empty()) + .map(SecretValue::new) + .transpose() + .map_err(serde::de::Error::custom) +} diff --git a/src/ilink/updates.rs b/src/ilink/updates.rs new file mode 100644 index 0000000..b906268 --- /dev/null +++ b/src/ilink/updates.rs @@ -0,0 +1,180 @@ +//! Update polling, typing, and lifecycle notification operations. + +use std::time::Duration; + +use super::auth::SecretValue; +use super::client::{ensure_api_success, IlinkAuth, IlinkError, DEFAULT_CONFIG_TIMEOUT}; +use super::transport::IlinkClient; +use super::types::{GetConfigRequest, GetConfigResponse, SendTypingRequest, SendTypingResponse}; +use super::types::{GetUpdatesRequest, GetUpdatesResponse, NotifyRequest, NotifyResponse}; + +const MAX_UPDATE_MESSAGES: usize = 256; +const MAX_MESSAGE_ITEMS: usize = 32; +const MAX_INBOUND_TEXT_BYTES: usize = 16 * 1024; +const MAX_INBOUND_ID_BYTES: usize = 256; +const MAX_SERVER_LONG_POLL_MS: u64 = 60_000; + +impl IlinkClient { + pub(super) async fn get_updates_request( + &self, + auth: &IlinkAuth, + update_cursor: &str, + long_poll_timeout: Duration, + ) -> Result { + if update_cursor.len() > 64 * 1024 { + return Err(IlinkError::InvalidConfiguration("update cursor")); + } + let request = self.authenticated_post(auth, "ilink/bot/getupdates")?; + let response: GetUpdatesResponse = match self + .post_json( + request, + &GetUpdatesRequest { + get_updates_buf: update_cursor.to_string(), + base_info: self.identity.base_info(), + }, + Self::bounded_long_poll_timeout(long_poll_timeout)?, + "get_updates", + ) + .await + { + Ok(response) => response, + Err(IlinkError::Timeout) => { + let update_cursor = if update_cursor.is_empty() { + None + } else { + Some( + SecretValue::new(update_cursor.to_string()) + .map_err(|_| IlinkError::InvalidConfiguration("update cursor"))?, + ) + }; + return Ok(GetUpdatesResponse { + ret: Some(0), + errcode: None, + errmsg: None, + messages: Vec::new(), + update_cursor, + long_polling_timeout_ms: None, + }); + } + Err(error) => return Err(error), + }; + ensure_api_success("get_updates", response.ret, response.errcode)?; + validate_updates_response(&response)?; + Ok(response) + } + + pub(super) async fn get_config_request( + &self, + auth: &IlinkAuth, + owner_id: Option<&SecretValue>, + context_token: Option<&SecretValue>, + ) -> Result { + let request = self.authenticated_post(auth, "ilink/bot/getconfig")?; + let response: GetConfigResponse = self + .post_json( + request, + &GetConfigRequest { + base_info: self.identity.base_info(), + ilink_user_id: owner_id.cloned(), + context_token: context_token.cloned(), + }, + DEFAULT_CONFIG_TIMEOUT, + "get_config", + ) + .await?; + ensure_api_success("get_config", response.ret, None)?; + Ok(response) + } + + pub(super) async fn send_typing_request( + &self, + auth: &IlinkAuth, + owner_id: &SecretValue, + typing_ticket: &SecretValue, + status: i32, + ) -> Result { + if !matches!(status, 1 | 2) { + return Err(IlinkError::InvalidConfiguration("typing status")); + } + let request = self.authenticated_post(auth, "ilink/bot/sendtyping")?; + let response: SendTypingResponse = self + .post_json( + request, + &SendTypingRequest { + ilink_user_id: owner_id.clone(), + typing_ticket: typing_ticket.clone(), + status, + base_info: self.identity.base_info(), + }, + DEFAULT_CONFIG_TIMEOUT, + "send_typing", + ) + .await?; + ensure_api_success("send_typing", response.ret, None)?; + Ok(response) + } + + pub(super) async fn notify_request( + &self, + auth: &IlinkAuth, + start: bool, + ) -> Result { + let (endpoint, operation) = if start { + ("ilink/bot/msg/notifystart", "notify_start") + } else { + ("ilink/bot/msg/notifystop", "notify_stop") + }; + let request = self.authenticated_post(auth, endpoint)?; + let response: NotifyResponse = self + .post_json( + request, + &NotifyRequest { + base_info: self.identity.base_info(), + }, + DEFAULT_CONFIG_TIMEOUT, + operation, + ) + .await?; + ensure_api_success(operation, response.ret, None)?; + Ok(response) + } +} + +pub(super) fn validate_updates_response(response: &GetUpdatesResponse) -> Result<(), IlinkError> { + if response.messages.len() > MAX_UPDATE_MESSAGES + || response + .long_polling_timeout_ms + .is_some_and(|timeout| timeout == 0 || timeout > MAX_SERVER_LONG_POLL_MS) + { + return Err(IlinkError::InvalidResponse("get_updates")); + } + for message in &response.messages { + if message.item_list.len() > MAX_MESSAGE_ITEMS + || message + .client_id + .as_deref() + .is_some_and(|value| value.len() > MAX_INBOUND_ID_BYTES || value.contains('\0')) + || message + .run_id + .as_deref() + .is_some_and(|value| value.len() > MAX_INBOUND_ID_BYTES || value.contains('\0')) + { + return Err(IlinkError::InvalidResponse("get_updates")); + } + for item in &message.item_list { + if item + .msg_id + .as_deref() + .is_some_and(|value| value.len() > MAX_INBOUND_ID_BYTES || value.contains('\0')) + || item + .text_item + .as_ref() + .and_then(|text| text.text.as_deref()) + .is_some_and(|text| text.len() > MAX_INBOUND_TEXT_BYTES || text.contains('\0')) + { + return Err(IlinkError::InvalidResponse("get_updates")); + } + } + } + Ok(()) +} diff --git a/src/ilink/url_policy.rs b/src/ilink/url_policy.rs new file mode 100644 index 0000000..1daa85a --- /dev/null +++ b/src/ilink/url_policy.rs @@ -0,0 +1,221 @@ +//! Strict URL validation for server-selected iLink endpoints. + +use std::collections::HashSet; +use std::fmt; + +use thiserror::Error; +use url::{Host, Url}; + +#[derive(Clone)] +pub(super) struct IlinkHostPolicy { + allowed_hosts: HashSet, + allowed_host_suffixes: HashSet, + allow_insecure_loopback: bool, +} + +impl IlinkHostPolicy { + pub(super) fn production( + allowed_hosts: impl IntoIterator>, + ) -> Result { + let mut normalized = HashSet::new(); + for host in allowed_hosts { + let host = normalize_allowed_host(host.as_ref())?; + normalized.insert(host); + } + if normalized.is_empty() { + return Err(IlinkUrlError::EmptyHostPolicy); + } + Ok(Self { + allowed_hosts: normalized, + allowed_host_suffixes: HashSet::new(), + allow_insecure_loopback: false, + }) + } + + pub(super) fn production_with_suffixes( + allowed_hosts: impl IntoIterator>, + allowed_host_suffixes: impl IntoIterator>, + ) -> Result { + let mut policy = Self::production(allowed_hosts)?; + for suffix in allowed_host_suffixes { + policy + .allowed_host_suffixes + .insert(normalize_allowed_host(suffix.as_ref())?); + } + Ok(policy) + } + + pub(super) fn for_test_origin(origin: &str) -> Result { + let url = Url::parse(origin).map_err(|_| IlinkUrlError::InvalidUrl)?; + let host = url + .host_str() + .ok_or(IlinkUrlError::MissingHost)? + .to_ascii_lowercase(); + if url.scheme() != "http" || !is_loopback_host(&host) { + return Err(IlinkUrlError::InsecureScheme); + } + Ok(Self { + allowed_hosts: [host].into_iter().collect(), + allowed_host_suffixes: HashSet::new(), + allow_insecure_loopback: true, + }) + } + + pub(super) fn validate(&self, candidate: &str) -> Result { + let mut url = Url::parse(candidate).map_err(|_| IlinkUrlError::InvalidUrl)?; + if !url.username().is_empty() || url.password().is_some() { + return Err(IlinkUrlError::UserInfoNotAllowed); + } + if url.query().is_some() { + return Err(IlinkUrlError::QueryNotAllowed); + } + if url.fragment().is_some() { + return Err(IlinkUrlError::FragmentNotAllowed); + } + let host = url + .host_str() + .ok_or(IlinkUrlError::MissingHost)? + .to_ascii_lowercase(); + let production_https = url.scheme() == "https" && url.port().is_none_or(|port| port == 443); + let test_http = + self.allow_insecure_loopback && url.scheme() == "http" && is_loopback_host(&host); + if !production_https && !test_http { + return if url.scheme() != "https" && !test_http { + Err(IlinkUrlError::InsecureScheme) + } else { + Err(IlinkUrlError::PortNotAllowed) + }; + } + match url.host() { + Some(Host::Domain(_)) => {} + Some(Host::Ipv4(_)) | Some(Host::Ipv6(_)) if test_http => {} + Some(Host::Ipv4(_)) | Some(Host::Ipv6(_)) => { + return Err(IlinkUrlError::IpLiteralNotAllowed) + } + None => return Err(IlinkUrlError::MissingHost), + } + if !self.host_is_allowed(&host) { + return Err(IlinkUrlError::HostNotAllowed); + } + if !url.path().ends_with('/') { + let path = format!("{}/", url.path()); + url.set_path(&path); + } + Ok(ValidatedBaseUrl(url)) + } + + fn host_is_allowed(&self, host: &str) -> bool { + self.allowed_hosts.contains(host) + || self + .allowed_host_suffixes + .iter() + .any(|suffix| host == suffix || host.ends_with(&format!(".{suffix}"))) + } + + pub(super) fn validate_redirect_host( + &self, + candidate: &str, + ) -> Result { + let probe = + Url::parse(&format!("https://{candidate}/")).map_err(|_| IlinkUrlError::InvalidUrl)?; + if !probe.username().is_empty() + || probe.password().is_some() + || probe.path() != "/" + || probe.query().is_some() + || probe.fragment().is_some() + { + return Err(IlinkUrlError::InvalidUrl); + } + let host = probe + .host_str() + .ok_or(IlinkUrlError::MissingHost)? + .to_ascii_lowercase(); + let scheme = if self.allow_insecure_loopback && is_loopback_host(&host) { + "http" + } else { + "https" + }; + self.validate(&format!("{scheme}://{candidate}/")) + } +} + +fn normalize_allowed_host(host: &str) -> Result { + let host = host.trim().trim_end_matches('.').to_ascii_lowercase(); + if host.is_empty() + || host.contains('/') + || host.contains(':') + || host.parse::().is_ok() + { + return Err(IlinkUrlError::InvalidAllowedHost); + } + let probe = + Url::parse(&format!("https://{host}")).map_err(|_| IlinkUrlError::InvalidAllowedHost)?; + if !matches!(probe.host(), Some(Host::Domain(_))) { + return Err(IlinkUrlError::InvalidAllowedHost); + } + Ok(host) +} + +fn is_loopback_host(host: &str) -> bool { + host == "localhost" + || host + .parse::() + .is_ok_and(|address| address.is_loopback()) +} + +#[derive(Clone, PartialEq, Eq)] +pub struct ValidatedBaseUrl(Url); + +impl ValidatedBaseUrl { + pub(super) fn join(&self, endpoint: &str) -> Result { + if endpoint.starts_with('/') || endpoint.contains("..") { + return Err(IlinkUrlError::InvalidEndpoint); + } + self.0 + .join(endpoint) + .map_err(|_| IlinkUrlError::InvalidEndpoint) + } + + #[doc(hidden)] + pub fn insecure_loopback_for_tests(origin: &str) -> Result { + IlinkHostPolicy::for_test_origin(origin)?.validate(origin) + } +} + +impl fmt::Debug for ValidatedBaseUrl { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ValidatedBaseUrl") + .field("scheme", &self.0.scheme()) + .field("host", &self.0.host_str()) + .finish_non_exhaustive() + } +} + +#[derive(Clone, Copy, Debug, Error, PartialEq, Eq)] +pub enum IlinkUrlError { + #[error("iLink URL is invalid")] + InvalidUrl, + #[error("iLink URL must use HTTPS")] + InsecureScheme, + #[error("iLink URL host is missing")] + MissingHost, + #[error("iLink URL user information is not allowed")] + UserInfoNotAllowed, + #[error("iLink URL port is not allowed")] + PortNotAllowed, + #[error("iLink URL IP literals are not allowed")] + IpLiteralNotAllowed, + #[error("iLink URL host is not approved")] + HostNotAllowed, + #[error("iLink URL query is not allowed")] + QueryNotAllowed, + #[error("iLink URL fragment is not allowed")] + FragmentNotAllowed, + #[error("iLink endpoint is invalid")] + InvalidEndpoint, + #[error("iLink host policy must not be empty")] + EmptyHostPolicy, + #[error("iLink allowed host is invalid")] + InvalidAllowedHost, +} diff --git a/src/lib.rs b/src/lib.rs index a3877e7..c7b4ff5 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -34,6 +34,8 @@ mod health; mod http; #[cfg(feature = "http-client")] mod http_client; +#[cfg(feature = "ilink")] +pub mod ilink; #[cfg(feature = "logging")] mod logging; mod module;