mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-30 01:19:34 +00:00
Add WeChat image messaging and QR login polish
This commit is contained in:
@@ -295,7 +295,6 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
} else {
|
||||
tool_defs
|
||||
};
|
||||
|
||||
// Update context for this iteration
|
||||
reason_ctx.available_tools = tool_defs;
|
||||
reason_ctx.system_prompt = Some(if force_text {
|
||||
|
||||
+102
-60
@@ -31,7 +31,58 @@ fn requires_preexisting_uuid_thread(channel: &str) -> bool {
|
||||
matches!(channel, "gateway" | "test")
|
||||
}
|
||||
|
||||
fn validate_inbound_text_for_message(
|
||||
safety: &crate::safety::SafetyLayer,
|
||||
content: &str,
|
||||
attachments: &[crate::channels::IncomingAttachment],
|
||||
) -> crate::safety::ValidationResult {
|
||||
if content.trim().is_empty() && !attachments.is_empty() {
|
||||
crate::safety::ValidationResult::ok()
|
||||
} else {
|
||||
safety.validate_input(content)
|
||||
}
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
fn reject_unsafe_inbound_user_message(
|
||||
&self,
|
||||
message: &IncomingMessage,
|
||||
content: &str,
|
||||
) -> Option<SubmissionResult> {
|
||||
let validation =
|
||||
validate_inbound_text_for_message(self.safety(), content, &message.attachments);
|
||||
if !validation.is_valid {
|
||||
let details = validation
|
||||
.errors
|
||||
.iter()
|
||||
.map(|e| format!("{}: {}", e.field, e.message))
|
||||
.collect::<Vec<_>>()
|
||||
.join("; ");
|
||||
return Some(SubmissionResult::error(format!(
|
||||
"Input rejected by safety validation: {details}",
|
||||
)));
|
||||
}
|
||||
|
||||
let violations = self.safety().check_policy(content);
|
||||
if violations
|
||||
.iter()
|
||||
.any(|rule| rule.action == crate::safety::PolicyAction::Block)
|
||||
{
|
||||
return Some(SubmissionResult::error("Input rejected by safety policy."));
|
||||
}
|
||||
|
||||
if let Some(warning) = self.safety().scan_inbound_for_secrets(content) {
|
||||
tracing::warn!(
|
||||
user = %message.user_id,
|
||||
channel = %message.channel,
|
||||
"Inbound message blocked: contains leaked secret"
|
||||
);
|
||||
return Some(SubmissionResult::error(warning));
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
/// Hydrate a historical thread from DB into memory if not already present.
|
||||
///
|
||||
/// Called before `resolve_thread` so that the session manager finds the
|
||||
@@ -226,34 +277,11 @@ impl Agent {
|
||||
}
|
||||
|
||||
// Run the same safety checks that the normal path applies
|
||||
// (validation, policy, secret scan) so that blocked content
|
||||
// is never stored in pending_messages or serialized.
|
||||
let validation = self.safety().validate_input(content);
|
||||
if !validation.is_valid {
|
||||
let details = validation
|
||||
.errors
|
||||
.iter()
|
||||
.map(|e| format!("{}: {}", e.field, e.message))
|
||||
.collect::<Vec<_>>()
|
||||
.join("; ");
|
||||
return Ok(SubmissionResult::error(format!(
|
||||
"Input rejected by safety validation: {details}",
|
||||
)));
|
||||
}
|
||||
let violations = self.safety().check_policy(content);
|
||||
if violations
|
||||
.iter()
|
||||
.any(|rule| rule.action == crate::safety::PolicyAction::Block)
|
||||
// so blocked content is never stored in pending_messages.
|
||||
if let Some(rejection) =
|
||||
self.reject_unsafe_inbound_user_message(message, content)
|
||||
{
|
||||
return Ok(SubmissionResult::error("Input rejected by safety policy."));
|
||||
}
|
||||
if let Some(warning) = self.safety().scan_inbound_for_secrets(content) {
|
||||
tracing::warn!(
|
||||
user = %message.user_id,
|
||||
channel = %message.channel,
|
||||
"Queued message blocked: contains leaked secret"
|
||||
);
|
||||
return Ok(SubmissionResult::error(warning));
|
||||
return Ok(rejection);
|
||||
}
|
||||
|
||||
if !thread.queue_message(content.to_string()) {
|
||||
@@ -307,39 +335,11 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
// Safety validation for user input
|
||||
let validation = self.safety().validate_input(content);
|
||||
if !validation.is_valid {
|
||||
let details = validation
|
||||
.errors
|
||||
.iter()
|
||||
.map(|e| format!("{}: {}", e.field, e.message))
|
||||
.collect::<Vec<_>>()
|
||||
.join("; ");
|
||||
return Ok(SubmissionResult::error(format!(
|
||||
"Input rejected by safety validation: {}",
|
||||
details
|
||||
)));
|
||||
}
|
||||
|
||||
let violations = self.safety().check_policy(content);
|
||||
if violations
|
||||
.iter()
|
||||
.any(|rule| rule.action == crate::safety::PolicyAction::Block)
|
||||
{
|
||||
return Ok(SubmissionResult::error("Input rejected by safety policy."));
|
||||
}
|
||||
|
||||
// Scan inbound messages for secrets (API keys, tokens).
|
||||
// Catching them here prevents the LLM from echoing them back, which
|
||||
// would trigger the outbound leak detector and create error loops.
|
||||
if let Some(warning) = self.safety().scan_inbound_for_secrets(content) {
|
||||
tracing::warn!(
|
||||
user = %message.user_id,
|
||||
channel = %message.channel,
|
||||
"Inbound message blocked: contains leaked secret"
|
||||
);
|
||||
return Ok(SubmissionResult::error(warning));
|
||||
// Validate inbound content before the turn is created. Attachment-only
|
||||
// messages are allowed to pass through so multimodal channels can send
|
||||
// an empty text body alongside real image/document payloads.
|
||||
if let Some(rejection) = self.reject_unsafe_inbound_user_message(message, content) {
|
||||
return Ok(rejection);
|
||||
}
|
||||
|
||||
// Handle explicit commands (starting with /) directly
|
||||
@@ -1880,6 +1880,9 @@ fn rebuild_chat_messages_from_db(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::channels::{AttachmentKind, IncomingAttachment};
|
||||
use crate::config::SafetyConfig;
|
||||
use crate::safety::SafetyLayer;
|
||||
|
||||
#[test]
|
||||
fn test_rebuild_chat_messages_user_assistant_only() {
|
||||
@@ -2017,6 +2020,45 @@ mod tests {
|
||||
assert_eq!(result[7].content, "Written");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_inbound_text_rejects_empty_text_without_attachments() {
|
||||
let safety = SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 10_000,
|
||||
injection_check_enabled: true,
|
||||
});
|
||||
|
||||
let result = validate_inbound_text_for_message(&safety, "", &[]);
|
||||
assert!(!result.is_valid);
|
||||
assert_eq!(result.errors.len(), 1);
|
||||
assert_eq!(result.errors[0].field, "input");
|
||||
assert_eq!(result.errors[0].message, "Input cannot be empty");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_inbound_text_allows_empty_text_when_attachments_exist() {
|
||||
let safety = SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 10_000,
|
||||
injection_check_enabled: true,
|
||||
});
|
||||
|
||||
let attachments = vec![IncomingAttachment {
|
||||
id: "image-1".to_string(),
|
||||
kind: AttachmentKind::Image,
|
||||
mime_type: "image/jpeg".to_string(),
|
||||
filename: Some("photo.jpg".to_string()),
|
||||
size_bytes: Some(128),
|
||||
source_url: Some("https://example.com/photo.jpg".to_string()),
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
data: vec![1, 2, 3],
|
||||
duration_secs: None,
|
||||
}];
|
||||
|
||||
let result = validate_inbound_text_for_message(&safety, "", &attachments);
|
||||
assert!(result.is_valid);
|
||||
assert!(result.errors.is_empty());
|
||||
}
|
||||
|
||||
fn make_db_msg(role: &str, content: &str) -> crate::history::ConversationMessage {
|
||||
crate::history::ConversationMessage {
|
||||
id: uuid::Uuid::new_v4(),
|
||||
|
||||
@@ -0,0 +1,313 @@
|
||||
use std::time::Duration;
|
||||
|
||||
use aes::Aes128;
|
||||
use aes::cipher::{BlockDecrypt, KeyInit, generic_array::GenericArray};
|
||||
use base64::Engine as _;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::channels::wasm::capabilities::ChannelCapabilities;
|
||||
use crate::channels::wasm::host::{Attachment, ChannelHostState};
|
||||
|
||||
const AES_BLOCK_SIZE: usize = 16;
|
||||
const MAX_ATTACHMENT_BYTES: usize = 20 * 1024 * 1024;
|
||||
const WECHAT_CHANNEL_NAME: &str = "wechat";
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct WechatAttachmentExtras {
|
||||
#[serde(default)]
|
||||
wechat_aes_key: Option<String>,
|
||||
}
|
||||
|
||||
pub(crate) async fn hydrate_attachment_for_channel(
|
||||
channel_name: &str,
|
||||
capabilities: &ChannelCapabilities,
|
||||
attachment: &mut Attachment,
|
||||
) {
|
||||
if !should_hydrate_wechat_attachment(channel_name, attachment) {
|
||||
return;
|
||||
}
|
||||
|
||||
let Some(source_url) = attachment.source_url.as_deref() else {
|
||||
return;
|
||||
};
|
||||
let Some(encoded_aes_key) = wechat_aes_key(&attachment.extras_json) else {
|
||||
tracing::warn!(
|
||||
channel = %channel_name,
|
||||
attachment_id = %attachment.id,
|
||||
"Skipping WeChat image hydration: missing AES key metadata"
|
||||
);
|
||||
return;
|
||||
};
|
||||
|
||||
match download_wechat_attachment_bytes(channel_name, capabilities, source_url).await {
|
||||
Ok(ciphertext) => match decrypt_wechat_image_bytes(&ciphertext, &encoded_aes_key) {
|
||||
Ok(plaintext) => {
|
||||
attachment.size_bytes = Some(plaintext.len() as u64);
|
||||
attachment.mime_type = detect_image_mime(&plaintext).to_string();
|
||||
attachment.data = plaintext;
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
channel = %channel_name,
|
||||
attachment_id = %attachment.id,
|
||||
error = %error,
|
||||
"Failed to decrypt WeChat image attachment"
|
||||
);
|
||||
}
|
||||
},
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
channel = %channel_name,
|
||||
attachment_id = %attachment.id,
|
||||
error = %error,
|
||||
"Failed to download WeChat image attachment"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn should_hydrate_wechat_attachment(channel_name: &str, attachment: &Attachment) -> bool {
|
||||
channel_name == WECHAT_CHANNEL_NAME
|
||||
&& attachment.data.is_empty()
|
||||
&& attachment.mime_type.starts_with("image/")
|
||||
}
|
||||
|
||||
fn wechat_aes_key(extras_json: &str) -> Option<String> {
|
||||
if extras_json.trim().is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
serde_json::from_str::<WechatAttachmentExtras>(extras_json)
|
||||
.ok()
|
||||
.and_then(|extras| extras.wechat_aes_key)
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
}
|
||||
|
||||
async fn download_wechat_attachment_bytes(
|
||||
channel_name: &str,
|
||||
capabilities: &ChannelCapabilities,
|
||||
source_url: &str,
|
||||
) -> Result<Vec<u8>, String> {
|
||||
let host_state = ChannelHostState::new(channel_name, capabilities.clone());
|
||||
host_state.check_http_allowed(source_url, "GET")?;
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.connect_timeout(Duration::from_secs(10))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| format!("Failed to build HTTP client: {e}"))?;
|
||||
|
||||
let response = client
|
||||
.get(source_url)
|
||||
.timeout(Duration::from_secs(15))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("WeChat CDN download failed: {e}"))?;
|
||||
|
||||
if response.status() != reqwest::StatusCode::OK {
|
||||
return Err(format!(
|
||||
"WeChat CDN download returned {}",
|
||||
response.status()
|
||||
));
|
||||
}
|
||||
|
||||
let bytes = response
|
||||
.bytes()
|
||||
.await
|
||||
.map_err(|e| format!("Failed to read WeChat CDN response body: {e}"))?
|
||||
.to_vec();
|
||||
|
||||
if bytes.is_empty() {
|
||||
return Err("WeChat CDN download returned an empty body".to_string());
|
||||
}
|
||||
if bytes.len() > MAX_ATTACHMENT_BYTES {
|
||||
return Err(format!(
|
||||
"WeChat image attachment exceeds {MAX_ATTACHMENT_BYTES} bytes"
|
||||
));
|
||||
}
|
||||
|
||||
Ok(bytes)
|
||||
}
|
||||
|
||||
fn decrypt_wechat_image_bytes(ciphertext: &[u8], encoded_aes_key: &str) -> Result<Vec<u8>, String> {
|
||||
let key = parse_aes_key(encoded_aes_key)?;
|
||||
decrypt_aes_ecb_pkcs7(ciphertext, &key)
|
||||
}
|
||||
|
||||
fn parse_aes_key(encoded: &str) -> Result<Vec<u8>, String> {
|
||||
let decoded = if encoded.len() == 32 && encoded.bytes().all(|byte| byte.is_ascii_hexdigit()) {
|
||||
decode_hex(encoded)?
|
||||
} else {
|
||||
base64::engine::general_purpose::STANDARD
|
||||
.decode(encoded)
|
||||
.map_err(|e| format!("Failed to decode WeChat AES key: {e}"))?
|
||||
};
|
||||
|
||||
if decoded.len() == AES_BLOCK_SIZE {
|
||||
return Ok(decoded);
|
||||
}
|
||||
|
||||
if decoded.len() == 32 && decoded.iter().all(|byte| byte.is_ascii_hexdigit()) {
|
||||
return decode_hex(
|
||||
std::str::from_utf8(&decoded)
|
||||
.map_err(|e| format!("WeChat AES key hex payload is not valid UTF-8: {e}"))?,
|
||||
);
|
||||
}
|
||||
|
||||
Err(format!(
|
||||
"WeChat AES key must decode to 16 bytes or a 32-char hex string, got {} bytes",
|
||||
decoded.len()
|
||||
))
|
||||
}
|
||||
|
||||
fn decode_hex(input: &str) -> Result<Vec<u8>, String> {
|
||||
if !input.len().is_multiple_of(2) {
|
||||
return Err("hex input length must be even".to_string());
|
||||
}
|
||||
let mut bytes = Vec::with_capacity(input.len() / 2);
|
||||
let chars: Vec<u8> = input.as_bytes().to_vec();
|
||||
for idx in (0..chars.len()).step_by(2) {
|
||||
let high = from_hex_digit(chars[idx])?;
|
||||
let low = from_hex_digit(chars[idx + 1])?;
|
||||
bytes.push((high << 4) | low);
|
||||
}
|
||||
Ok(bytes)
|
||||
}
|
||||
|
||||
fn from_hex_digit(value: u8) -> Result<u8, String> {
|
||||
match value {
|
||||
b'0'..=b'9' => Ok(value - b'0'),
|
||||
b'a'..=b'f' => Ok(value - b'a' + 10),
|
||||
b'A'..=b'F' => Ok(value - b'A' + 10),
|
||||
_ => Err(format!("invalid hex digit '{}'", value as char)),
|
||||
}
|
||||
}
|
||||
|
||||
fn decrypt_aes_ecb_pkcs7(ciphertext: &[u8], key: &[u8]) -> Result<Vec<u8>, String> {
|
||||
if !ciphertext.len().is_multiple_of(AES_BLOCK_SIZE) {
|
||||
return Err("ciphertext length is not a multiple of 16 bytes".to_string());
|
||||
}
|
||||
|
||||
let cipher = Aes128::new_from_slice(key).map_err(|e| format!("Invalid AES key: {e}"))?;
|
||||
let mut plaintext = ciphertext.to_vec();
|
||||
for chunk in plaintext.chunks_exact_mut(AES_BLOCK_SIZE) {
|
||||
cipher.decrypt_block(GenericArray::from_mut_slice(chunk));
|
||||
}
|
||||
|
||||
let pad_len = *plaintext
|
||||
.last()
|
||||
.ok_or_else(|| "ciphertext decrypted to an empty buffer".to_string())?
|
||||
as usize;
|
||||
if pad_len == 0 || pad_len > AES_BLOCK_SIZE || pad_len > plaintext.len() {
|
||||
return Err("invalid PKCS7 padding".to_string());
|
||||
}
|
||||
if !plaintext[plaintext.len() - pad_len..]
|
||||
.iter()
|
||||
.all(|byte| *byte as usize == pad_len)
|
||||
{
|
||||
return Err("invalid PKCS7 padding bytes".to_string());
|
||||
}
|
||||
plaintext.truncate(plaintext.len() - pad_len);
|
||||
Ok(plaintext)
|
||||
}
|
||||
|
||||
fn detect_image_mime(bytes: &[u8]) -> &'static str {
|
||||
if bytes.starts_with(&[0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A]) {
|
||||
"image/png"
|
||||
} else if bytes.starts_with(&[0xFF, 0xD8, 0xFF]) {
|
||||
"image/jpeg"
|
||||
} else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") {
|
||||
"image/gif"
|
||||
} else if bytes.len() >= 12 && &bytes[..4] == b"RIFF" && &bytes[8..12] == b"WEBP" {
|
||||
"image/webp"
|
||||
} else {
|
||||
"image/jpeg"
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn encrypt_aes_ecb_pkcs7(plaintext: &[u8], key: &[u8]) -> Result<Vec<u8>, String> {
|
||||
use aes::cipher::BlockEncrypt;
|
||||
|
||||
let cipher = Aes128::new_from_slice(key).map_err(|e| format!("Invalid AES key: {e}"))?;
|
||||
let mut padded = plaintext.to_vec();
|
||||
let pad_len = AES_BLOCK_SIZE - (padded.len() % AES_BLOCK_SIZE);
|
||||
padded.extend(std::iter::repeat_n(pad_len as u8, pad_len));
|
||||
|
||||
for chunk in padded.chunks_exact_mut(AES_BLOCK_SIZE) {
|
||||
cipher.encrypt_block(GenericArray::from_mut_slice(chunk));
|
||||
}
|
||||
|
||||
Ok(padded)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
Attachment, decrypt_wechat_image_bytes, detect_image_mime, encrypt_aes_ecb_pkcs7,
|
||||
hydrate_attachment_for_channel, should_hydrate_wechat_attachment,
|
||||
};
|
||||
use crate::channels::wasm::ChannelCapabilities;
|
||||
use base64::Engine as _;
|
||||
|
||||
fn make_attachment() -> Attachment {
|
||||
Attachment {
|
||||
id: "wechat-image-1".to_string(),
|
||||
mime_type: "image/jpeg".to_string(),
|
||||
filename: Some("wechat-image.jpg".to_string()),
|
||||
size_bytes: None,
|
||||
source_url: Some(
|
||||
"https://novac2c.cdn.weixin.qq.com/c2c/download?encrypted_query_param=test"
|
||||
.to_string(),
|
||||
),
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
extras_json: String::new(),
|
||||
data: Vec::new(),
|
||||
duration_secs: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn encode_test_extras_json(aes_key: &str) -> String {
|
||||
serde_json::json!({ "wechat_aes_key": aes_key }).to_string()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decrypt_wechat_image_bytes_round_trips() {
|
||||
let key = [7u8; 16];
|
||||
let plaintext = vec![0xFF, 0xD8, 0xFF, 0xDB, 0x00, 0x11];
|
||||
let ciphertext = encrypt_aes_ecb_pkcs7(&plaintext, &key).unwrap();
|
||||
let encoded_key = base64::engine::general_purpose::STANDARD.encode(key);
|
||||
let decrypted = decrypt_wechat_image_bytes(&ciphertext, &encoded_key).unwrap();
|
||||
assert_eq!(decrypted, plaintext);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detect_image_mime_prefers_magic_bytes() {
|
||||
assert_eq!(detect_image_mime(&[0xFF, 0xD8, 0xFF, 0x00]), "image/jpeg");
|
||||
assert_eq!(
|
||||
detect_image_mime(&[0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A]),
|
||||
"image/png"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn wechat_attachment_hydration_only_applies_to_wechat_images() {
|
||||
let mut attachment = make_attachment();
|
||||
attachment.extras_json = encode_test_extras_json("ZmFrZS1rZXk=");
|
||||
assert!(should_hydrate_wechat_attachment("wechat", &attachment));
|
||||
assert!(!should_hydrate_wechat_attachment("telegram", &attachment));
|
||||
|
||||
attachment.mime_type = "application/pdf".to_string();
|
||||
assert!(!should_hydrate_wechat_attachment("wechat", &attachment));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn hydration_skips_when_metadata_is_missing() {
|
||||
let mut attachment = make_attachment();
|
||||
let caps = ChannelCapabilities::for_channel("wechat");
|
||||
hydrate_attachment_for_channel("wechat", &caps, &mut attachment).await;
|
||||
assert!(attachment.data.is_empty());
|
||||
assert_eq!(attachment.size_bytes, None);
|
||||
}
|
||||
}
|
||||
@@ -35,6 +35,8 @@ pub struct Attachment {
|
||||
pub storage_key: Option<String>,
|
||||
/// Extracted text content (e.g., OCR result, PDF text, audio transcript).
|
||||
pub extracted_text: Option<String>,
|
||||
/// Extensible metadata from the channel payload.
|
||||
pub extras_json: String,
|
||||
/// Raw file bytes (for small files downloaded by the channel).
|
||||
pub data: Vec<u8>,
|
||||
/// Duration in seconds (for audio/video).
|
||||
@@ -995,6 +997,7 @@ mod tests {
|
||||
source_url: None,
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
extras_json: String::new(),
|
||||
data: Vec::new(),
|
||||
duration_secs: None,
|
||||
}
|
||||
|
||||
@@ -78,6 +78,7 @@
|
||||
//! }
|
||||
//! ```
|
||||
|
||||
mod attachment_hydration;
|
||||
mod bundled;
|
||||
mod capabilities;
|
||||
mod error;
|
||||
|
||||
@@ -159,6 +159,11 @@ impl WasmChannelRouter {
|
||||
self.channels.read().await.get(channel_name).cloned()
|
||||
}
|
||||
|
||||
/// Get a registered channel directly by name.
|
||||
pub async fn get_channel_by_name(&self, channel_name: &str) -> Option<Arc<WasmChannel>> {
|
||||
self.channels.read().await.get(channel_name).cloned()
|
||||
}
|
||||
|
||||
/// Validate a secret for a channel.
|
||||
pub async fn validate_secret(&self, channel_name: &str, provided: &str) -> bool {
|
||||
let secrets = self.secrets.read().await;
|
||||
@@ -710,6 +715,10 @@ mod tests {
|
||||
// Should not find non-existent path
|
||||
let not_found = router.get_channel_for_path("/webhook/telegram").await;
|
||||
assert!(not_found.is_none());
|
||||
|
||||
let found_by_name = router.get_channel_by_name("slack").await;
|
||||
assert!(found_by_name.is_some());
|
||||
assert_eq!(found_by_name.unwrap().channel_name(), "slack");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
+145
-78
@@ -573,6 +573,7 @@ impl near::agent::channel_host::Host for ChannelStoreData {
|
||||
source_url: a.source_url,
|
||||
storage_key: a.storage_key,
|
||||
extracted_text: a.extracted_text,
|
||||
extras_json: a.extras_json,
|
||||
data,
|
||||
duration_secs,
|
||||
}
|
||||
@@ -1181,22 +1182,32 @@ impl WasmChannel {
|
||||
)
|
||||
}
|
||||
|
||||
fn log_on_start_host_state(&self, host_state: &mut ChannelHostState) {
|
||||
fn log_host_state_entries(channel_name: &str, host_state: &mut ChannelHostState) {
|
||||
for entry in host_state.take_logs() {
|
||||
match entry.level {
|
||||
crate::tools::wasm::LogLevel::Trace => {
|
||||
tracing::trace!(channel = %channel_name, "{}", entry.message);
|
||||
}
|
||||
crate::tools::wasm::LogLevel::Debug => {
|
||||
tracing::debug!(channel = %channel_name, "{}", entry.message);
|
||||
}
|
||||
crate::tools::wasm::LogLevel::Info => {
|
||||
tracing::info!(channel = %channel_name, "{}", entry.message);
|
||||
}
|
||||
crate::tools::wasm::LogLevel::Error => {
|
||||
tracing::error!(channel = %self.name, "{}", entry.message);
|
||||
tracing::error!(channel = %channel_name, "{}", entry.message);
|
||||
}
|
||||
crate::tools::wasm::LogLevel::Warn => {
|
||||
tracing::warn!(channel = %self.name, "{}", entry.message);
|
||||
}
|
||||
_ => {
|
||||
tracing::debug!(channel = %self.name, "{}", entry.message);
|
||||
tracing::warn!(channel = %channel_name, "{}", entry.message);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn log_on_start_host_state(&self, host_state: &mut ChannelHostState) {
|
||||
Self::log_host_state_entries(&self.name, host_state);
|
||||
}
|
||||
|
||||
async fn execute_on_start_with_state(
|
||||
&self,
|
||||
) -> Result<(Result<ChannelConfig, WasmChannelError>, ChannelHostState), WasmChannelError> {
|
||||
@@ -1480,18 +1491,20 @@ impl WasmChannel {
|
||||
|
||||
// Call on_poll using the generated typed interface
|
||||
let channel_iface = instance.near_agent_channel();
|
||||
channel_iface
|
||||
let poll_result = channel_iface
|
||||
.call_on_poll(&mut store)
|
||||
.map_err(|e| Self::map_wasm_error(e, &prepared.name, prepared.limits.fuel))?;
|
||||
.map_err(|e| Self::map_wasm_error(e, &prepared.name, prepared.limits.fuel));
|
||||
|
||||
let mut host_state =
|
||||
Self::extract_host_state(&mut store, &prepared.name, &capabilities);
|
||||
|
||||
// Commit pending workspace writes to the persistent store
|
||||
let pending_writes = host_state.take_pending_writes();
|
||||
workspace_store.commit_writes(&pending_writes);
|
||||
if poll_result.is_ok() {
|
||||
// Commit pending workspace writes only after a successful callback.
|
||||
let pending_writes = host_state.take_pending_writes();
|
||||
workspace_store.commit_writes(&pending_writes);
|
||||
}
|
||||
|
||||
Ok(((), host_state))
|
||||
Ok((poll_result, host_state))
|
||||
})
|
||||
.await
|
||||
.map_err(|e| WasmChannelError::ExecutionPanicked {
|
||||
@@ -1503,7 +1516,10 @@ impl WasmChannel {
|
||||
|
||||
let channel_name = self.name.clone();
|
||||
match result {
|
||||
Ok(Ok(((), mut host_state))) => {
|
||||
Ok(Ok((poll_result, mut host_state))) => {
|
||||
Self::log_host_state_entries(&channel_name, &mut host_state);
|
||||
poll_result?;
|
||||
|
||||
// Process emitted messages
|
||||
let emitted = host_state.take_emitted_messages();
|
||||
self.process_emitted_messages(emitted).await?;
|
||||
@@ -2181,6 +2197,16 @@ impl WasmChannel {
|
||||
};
|
||||
|
||||
for emitted in messages {
|
||||
let EmittedMessage {
|
||||
user_id,
|
||||
user_name,
|
||||
content,
|
||||
thread_id,
|
||||
metadata_json,
|
||||
attachments,
|
||||
..
|
||||
} = emitted;
|
||||
|
||||
// Check rate limit — acquire and release the write lock before send().await
|
||||
{
|
||||
let mut rate_limiter = self.rate_limiter.write().await;
|
||||
@@ -2198,55 +2224,41 @@ impl WasmChannel {
|
||||
let (resolved_user_id, is_owner_sender) = resolve_message_scope(
|
||||
&self.owner_scope_id,
|
||||
self.owner_actor_id.as_deref(),
|
||||
&emitted.user_id,
|
||||
&user_id,
|
||||
);
|
||||
|
||||
// Convert to IncomingMessage
|
||||
let mut msg = IncomingMessage::new(&self.name, &resolved_user_id, &emitted.content)
|
||||
let mut msg = IncomingMessage::new(&self.name, &resolved_user_id, &content)
|
||||
.with_owner_id(&self.owner_scope_id)
|
||||
.with_sender_id(&emitted.user_id);
|
||||
.with_sender_id(&user_id);
|
||||
|
||||
if let Some(name) = emitted.user_name {
|
||||
if let Some(name) = user_name {
|
||||
msg = msg.with_user_name(name);
|
||||
}
|
||||
|
||||
if let Some(thread_id) = emitted.thread_id {
|
||||
if let Some(thread_id) = thread_id {
|
||||
msg = msg.with_thread(thread_id);
|
||||
}
|
||||
|
||||
// Convert attachments
|
||||
if !emitted.attachments.is_empty() {
|
||||
let incoming_attachments = emitted
|
||||
.attachments
|
||||
.iter()
|
||||
.map(|a| crate::channels::IncomingAttachment {
|
||||
id: a.id.clone(),
|
||||
kind: crate::channels::AttachmentKind::from_mime_type(&a.mime_type),
|
||||
mime_type: a.mime_type.clone(),
|
||||
filename: a.filename.clone(),
|
||||
size_bytes: a.size_bytes,
|
||||
source_url: a.source_url.clone(),
|
||||
storage_key: a.storage_key.clone(),
|
||||
extracted_text: a.extracted_text.clone(),
|
||||
data: a.data.clone(),
|
||||
duration_secs: a.duration_secs,
|
||||
})
|
||||
.collect();
|
||||
if !attachments.is_empty() {
|
||||
let incoming_attachments =
|
||||
convert_emitted_attachments(&self.name, &self.capabilities, attachments).await;
|
||||
msg = msg.with_attachments(incoming_attachments);
|
||||
}
|
||||
|
||||
// Parse metadata JSON
|
||||
msg = apply_emitted_metadata(msg, &emitted.metadata_json);
|
||||
msg = apply_emitted_metadata(msg, &metadata_json);
|
||||
if is_owner_sender {
|
||||
// Store for owner-target routing (chat_id etc.).
|
||||
self.update_broadcast_metadata(&emitted.metadata_json).await;
|
||||
self.update_broadcast_metadata(&metadata_json).await;
|
||||
}
|
||||
|
||||
// Send to stream — no locks held across this await
|
||||
tracing::info!(
|
||||
channel = %self.name,
|
||||
user_id = %emitted.user_id,
|
||||
content_len = emitted.content.len(),
|
||||
user_id = %user_id,
|
||||
content_len = content.len(),
|
||||
attachment_count = msg.attachments.len(),
|
||||
"Sending emitted message to agent"
|
||||
);
|
||||
@@ -2331,6 +2343,7 @@ impl WasmChannel {
|
||||
&& let Err(e) = Self::dispatch_emitted_messages(
|
||||
EmitDispatchContext {
|
||||
channel_name: &channel_name,
|
||||
capabilities: &capabilities,
|
||||
owner_scope_id: &owner_scope_id,
|
||||
owner_actor_id: owner_actor_id.as_deref(),
|
||||
message_tx: &message_tx,
|
||||
@@ -2416,18 +2429,20 @@ impl WasmChannel {
|
||||
|
||||
// Call on_poll using the generated typed interface
|
||||
let channel_iface = instance.near_agent_channel();
|
||||
channel_iface
|
||||
let poll_result = channel_iface
|
||||
.call_on_poll(&mut store)
|
||||
.map_err(|e| Self::map_wasm_error(e, &prepared.name, prepared.limits.fuel))?;
|
||||
.map_err(|e| Self::map_wasm_error(e, &prepared.name, prepared.limits.fuel));
|
||||
|
||||
let mut host_state =
|
||||
Self::extract_host_state(&mut store, &prepared.name, &capabilities);
|
||||
|
||||
// Commit pending workspace writes to the persistent store
|
||||
let pending_writes = host_state.take_pending_writes();
|
||||
workspace_store.commit_writes(&pending_writes);
|
||||
if poll_result.is_ok() {
|
||||
// Commit pending workspace writes only after a successful callback.
|
||||
let pending_writes = host_state.take_pending_writes();
|
||||
workspace_store.commit_writes(&pending_writes);
|
||||
}
|
||||
|
||||
Ok(host_state)
|
||||
Ok((poll_result, host_state))
|
||||
})
|
||||
.await
|
||||
.map_err(|e| WasmChannelError::ExecutionPanicked {
|
||||
@@ -2438,7 +2453,10 @@ impl WasmChannel {
|
||||
.await;
|
||||
|
||||
match result {
|
||||
Ok(Ok(mut host_state)) => {
|
||||
Ok(Ok((poll_result, mut host_state))) => {
|
||||
Self::log_host_state_entries(channel_name, &mut host_state);
|
||||
poll_result?;
|
||||
|
||||
let emitted = host_state.take_emitted_messages();
|
||||
tracing::debug!(
|
||||
channel = %channel_name,
|
||||
@@ -2484,6 +2502,16 @@ impl WasmChannel {
|
||||
};
|
||||
|
||||
for emitted in messages {
|
||||
let EmittedMessage {
|
||||
user_id,
|
||||
user_name,
|
||||
content,
|
||||
thread_id,
|
||||
metadata_json,
|
||||
attachments,
|
||||
..
|
||||
} = emitted;
|
||||
|
||||
// Check rate limit — acquire and release the write lock before send().await
|
||||
{
|
||||
let mut limiter = dispatch.rate_limiter.write().await;
|
||||
@@ -2498,54 +2526,40 @@ impl WasmChannel {
|
||||
}
|
||||
}
|
||||
|
||||
let (resolved_user_id, is_owner_sender) = resolve_message_scope(
|
||||
dispatch.owner_scope_id,
|
||||
dispatch.owner_actor_id,
|
||||
&emitted.user_id,
|
||||
);
|
||||
let (resolved_user_id, is_owner_sender) =
|
||||
resolve_message_scope(dispatch.owner_scope_id, dispatch.owner_actor_id, &user_id);
|
||||
|
||||
// Convert to IncomingMessage
|
||||
let mut msg =
|
||||
IncomingMessage::new(dispatch.channel_name, &resolved_user_id, &emitted.content)
|
||||
.with_owner_id(dispatch.owner_scope_id)
|
||||
.with_sender_id(&emitted.user_id);
|
||||
let mut msg = IncomingMessage::new(dispatch.channel_name, &resolved_user_id, &content)
|
||||
.with_owner_id(dispatch.owner_scope_id)
|
||||
.with_sender_id(&user_id);
|
||||
|
||||
if let Some(name) = emitted.user_name {
|
||||
if let Some(name) = user_name {
|
||||
msg = msg.with_user_name(name);
|
||||
}
|
||||
|
||||
if let Some(thread_id) = emitted.thread_id {
|
||||
if let Some(thread_id) = thread_id {
|
||||
msg = msg.with_thread(thread_id);
|
||||
}
|
||||
|
||||
// Convert attachments
|
||||
if !emitted.attachments.is_empty() {
|
||||
let incoming_attachments = emitted
|
||||
.attachments
|
||||
.iter()
|
||||
.map(|a| crate::channels::IncomingAttachment {
|
||||
id: a.id.clone(),
|
||||
kind: crate::channels::AttachmentKind::from_mime_type(&a.mime_type),
|
||||
mime_type: a.mime_type.clone(),
|
||||
filename: a.filename.clone(),
|
||||
size_bytes: a.size_bytes,
|
||||
source_url: a.source_url.clone(),
|
||||
storage_key: a.storage_key.clone(),
|
||||
extracted_text: a.extracted_text.clone(),
|
||||
data: a.data.clone(),
|
||||
duration_secs: a.duration_secs,
|
||||
})
|
||||
.collect();
|
||||
if !attachments.is_empty() {
|
||||
let incoming_attachments = convert_emitted_attachments(
|
||||
dispatch.channel_name,
|
||||
dispatch.capabilities,
|
||||
attachments,
|
||||
)
|
||||
.await;
|
||||
msg = msg.with_attachments(incoming_attachments);
|
||||
}
|
||||
|
||||
msg = apply_emitted_metadata(msg, &emitted.metadata_json);
|
||||
msg = apply_emitted_metadata(msg, &metadata_json);
|
||||
if is_owner_sender {
|
||||
// Store for owner-target routing (chat_id etc.)
|
||||
do_update_broadcast_metadata(
|
||||
dispatch.channel_name,
|
||||
dispatch.owner_scope_id,
|
||||
&emitted.metadata_json,
|
||||
&metadata_json,
|
||||
dispatch.last_broadcast_metadata,
|
||||
dispatch.settings_store,
|
||||
)
|
||||
@@ -2555,8 +2569,8 @@ impl WasmChannel {
|
||||
// Send to stream — no locks held across this await
|
||||
tracing::info!(
|
||||
channel = %dispatch.channel_name,
|
||||
user_id = %emitted.user_id,
|
||||
content_len = emitted.content.len(),
|
||||
user_id = %user_id,
|
||||
content_len = content.len(),
|
||||
attachment_count = msg.attachments.len(),
|
||||
"Sending polled message to agent"
|
||||
);
|
||||
@@ -2581,6 +2595,7 @@ impl WasmChannel {
|
||||
|
||||
struct EmitDispatchContext<'a> {
|
||||
channel_name: &'a str,
|
||||
capabilities: &'a ChannelCapabilities,
|
||||
owner_scope_id: &'a str,
|
||||
owner_actor_id: Option<&'a str>,
|
||||
message_tx: &'a RwLock<Option<mpsc::Sender<IncomingMessage>>>,
|
||||
@@ -3243,6 +3258,38 @@ async fn resolve_channel_host_credentials(
|
||||
/// Maximum total attachment size (50 MB).
|
||||
const MAX_TOTAL_ATTACHMENT_BYTES: u64 = 50 * 1024 * 1024;
|
||||
|
||||
async fn convert_emitted_attachments(
|
||||
channel_name: &str,
|
||||
capabilities: &ChannelCapabilities,
|
||||
attachments: Vec<crate::channels::wasm::host::Attachment>,
|
||||
) -> Vec<crate::channels::IncomingAttachment> {
|
||||
let mut hydrated = attachments;
|
||||
for attachment in &mut hydrated {
|
||||
crate::channels::wasm::attachment_hydration::hydrate_attachment_for_channel(
|
||||
channel_name,
|
||||
capabilities,
|
||||
attachment,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
hydrated
|
||||
.into_iter()
|
||||
.map(|attachment| crate::channels::IncomingAttachment {
|
||||
id: attachment.id,
|
||||
kind: crate::channels::AttachmentKind::from_mime_type(&attachment.mime_type),
|
||||
mime_type: attachment.mime_type,
|
||||
filename: attachment.filename,
|
||||
size_bytes: attachment.size_bytes,
|
||||
source_url: attachment.source_url,
|
||||
storage_key: attachment.storage_key,
|
||||
extracted_text: attachment.extracted_text,
|
||||
data: attachment.data,
|
||||
duration_secs: attachment.duration_secs,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Detect MIME type from file extension using the `mime_guess` crate.
|
||||
fn mime_from_extension(path: &str) -> String {
|
||||
mime_guess::from_path(path)
|
||||
@@ -3455,6 +3502,8 @@ mod tests {
|
||||
|
||||
let (tx, mut rx) = tokio::sync::mpsc::channel(10);
|
||||
let message_tx = Arc::new(tokio::sync::RwLock::new(Some(tx)));
|
||||
let capabilities =
|
||||
crate::channels::wasm::capabilities::ChannelCapabilities::for_channel("test-channel");
|
||||
|
||||
let rate_limiter = Arc::new(tokio::sync::RwLock::new(
|
||||
crate::channels::wasm::host::ChannelEmitRateLimiter::new(
|
||||
@@ -3471,6 +3520,7 @@ mod tests {
|
||||
let result = WasmChannel::dispatch_emitted_messages(
|
||||
EmitDispatchContext {
|
||||
channel_name: "test-channel",
|
||||
capabilities: &capabilities,
|
||||
owner_scope_id: "default",
|
||||
owner_actor_id: None,
|
||||
message_tx: &message_tx,
|
||||
@@ -3503,6 +3553,8 @@ mod tests {
|
||||
|
||||
// No sender available (channel not started)
|
||||
let message_tx = Arc::new(tokio::sync::RwLock::new(None));
|
||||
let capabilities =
|
||||
crate::channels::wasm::capabilities::ChannelCapabilities::for_channel("test-channel");
|
||||
let rate_limiter = Arc::new(tokio::sync::RwLock::new(
|
||||
crate::channels::wasm::host::ChannelEmitRateLimiter::new(
|
||||
crate::channels::wasm::capabilities::EmitRateLimitConfig::default(),
|
||||
@@ -3516,6 +3568,7 @@ mod tests {
|
||||
let result = WasmChannel::dispatch_emitted_messages(
|
||||
EmitDispatchContext {
|
||||
channel_name: "test-channel",
|
||||
capabilities: &capabilities,
|
||||
owner_scope_id: "default",
|
||||
owner_actor_id: None,
|
||||
message_tx: &message_tx,
|
||||
@@ -4506,6 +4559,8 @@ mod tests {
|
||||
|
||||
let (tx, mut rx) = tokio::sync::mpsc::channel(10);
|
||||
let message_tx = Arc::new(tokio::sync::RwLock::new(Some(tx)));
|
||||
let capabilities =
|
||||
crate::channels::wasm::capabilities::ChannelCapabilities::for_channel("test-channel");
|
||||
|
||||
let rate_limiter = Arc::new(tokio::sync::RwLock::new(
|
||||
crate::channels::wasm::host::ChannelEmitRateLimiter::new(
|
||||
@@ -4522,6 +4577,7 @@ mod tests {
|
||||
source_url: Some("https://api.telegram.org/file/photo123".to_string()),
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
extras_json: String::new(),
|
||||
data: Vec::new(),
|
||||
duration_secs: None,
|
||||
},
|
||||
@@ -4533,6 +4589,7 @@ mod tests {
|
||||
source_url: None,
|
||||
storage_key: Some("store/doc456".to_string()),
|
||||
extracted_text: Some("Report contents...".to_string()),
|
||||
extras_json: String::new(),
|
||||
data: Vec::new(),
|
||||
duration_secs: None,
|
||||
},
|
||||
@@ -4545,6 +4602,7 @@ mod tests {
|
||||
let result = WasmChannel::dispatch_emitted_messages(
|
||||
EmitDispatchContext {
|
||||
channel_name: "test-channel",
|
||||
capabilities: &capabilities,
|
||||
owner_scope_id: "default",
|
||||
owner_actor_id: None,
|
||||
message_tx: &message_tx,
|
||||
@@ -4591,6 +4649,8 @@ mod tests {
|
||||
|
||||
let (tx, mut rx) = tokio::sync::mpsc::channel(10);
|
||||
let message_tx = Arc::new(tokio::sync::RwLock::new(Some(tx)));
|
||||
let capabilities =
|
||||
crate::channels::wasm::capabilities::ChannelCapabilities::for_channel("telegram");
|
||||
let rate_limiter = Arc::new(tokio::sync::RwLock::new(
|
||||
crate::channels::wasm::host::ChannelEmitRateLimiter::new(
|
||||
crate::channels::wasm::capabilities::EmitRateLimitConfig::default(),
|
||||
@@ -4606,6 +4666,7 @@ mod tests {
|
||||
let result = WasmChannel::dispatch_emitted_messages(
|
||||
EmitDispatchContext {
|
||||
channel_name: "telegram",
|
||||
capabilities: &capabilities,
|
||||
owner_scope_id: "owner-scope",
|
||||
owner_actor_id: Some("telegram-owner"),
|
||||
message_tx: &message_tx,
|
||||
@@ -4634,6 +4695,8 @@ mod tests {
|
||||
|
||||
let (tx, mut rx) = tokio::sync::mpsc::channel(10);
|
||||
let message_tx = Arc::new(tokio::sync::RwLock::new(Some(tx)));
|
||||
let capabilities =
|
||||
crate::channels::wasm::capabilities::ChannelCapabilities::for_channel("telegram");
|
||||
let rate_limiter = Arc::new(tokio::sync::RwLock::new(
|
||||
crate::channels::wasm::host::ChannelEmitRateLimiter::new(
|
||||
crate::channels::wasm::capabilities::EmitRateLimitConfig::default(),
|
||||
@@ -4648,6 +4711,7 @@ mod tests {
|
||||
let result = WasmChannel::dispatch_emitted_messages(
|
||||
EmitDispatchContext {
|
||||
channel_name: "telegram",
|
||||
capabilities: &capabilities,
|
||||
owner_scope_id: "owner-scope",
|
||||
owner_actor_id: Some("telegram-owner"),
|
||||
message_tx: &message_tx,
|
||||
@@ -4717,6 +4781,8 @@ mod tests {
|
||||
|
||||
let (tx, mut rx) = tokio::sync::mpsc::channel(10);
|
||||
let message_tx = Arc::new(tokio::sync::RwLock::new(Some(tx)));
|
||||
let capabilities =
|
||||
crate::channels::wasm::capabilities::ChannelCapabilities::for_channel("test-channel");
|
||||
|
||||
let rate_limiter = Arc::new(tokio::sync::RwLock::new(
|
||||
crate::channels::wasm::host::ChannelEmitRateLimiter::new(
|
||||
@@ -4730,6 +4796,7 @@ mod tests {
|
||||
let result = WasmChannel::dispatch_emitted_messages(
|
||||
EmitDispatchContext {
|
||||
channel_name: "test-channel",
|
||||
capabilities: &capabilities,
|
||||
owner_scope_id: "default",
|
||||
owner_actor_id: None,
|
||||
message_tx: &message_tx,
|
||||
|
||||
@@ -4089,6 +4089,10 @@ impl ExtensionManager {
|
||||
|
||||
let webhook_path = format!("/webhook/{}", name);
|
||||
let existing_channel = match router.get_channel_for_path(&webhook_path).await {
|
||||
Some(ch) => Some(ch),
|
||||
None => router.get_channel_by_name(name).await,
|
||||
};
|
||||
let existing_channel = match existing_channel {
|
||||
Some(ch) => ch,
|
||||
None => {
|
||||
return Ok(ActivateResult {
|
||||
|
||||
Reference in New Issue
Block a user