mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-31 00:29:24 +00:00
* feat: Add secure prompt-based skills system (Phase 1 MVP) Implement a skills system that extends the agent with prompt-level instructions from local directories. Skills declare activation criteria, tool permissions, and trust tiers that determine authority attenuation. Core security model: the minimum trust level of any active skill determines a tool ceiling -- tools above the ceiling are removed from the LLM's tool list entirely at the API level, preventing prompt-based manipulation. New modules: - skills/mod.rs: Core types (SkillTrust, SkillManifest, LoadedSkill) - skills/scanner.rs: Content scanner for manipulation detection - skills/registry.rs: Filesystem discovery and manifest parsing - skills/selector.rs: Deterministic two-phase prefilter (no LLM) - skills/attenuation.rs: Trust-based tool filtering Integration: - Agent loop selects skills per-turn and applies tool attenuation - Reasoning engine injects skill context with structural isolation - Config supports SKILLS_ENABLED, SKILLS_DIR, SKILLS_MAX_ACTIVE, SKILLS_MAX_CONTEXT_TOKENS environment variables - Disabled by default (SKILLS_ENABLED=false) 41 new tests covering all modules. Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: Address all adversarial review findings for skills system Security fixes: - Escape skill name/version in XML attributes to prevent trust spoofing - Escape prompt content to prevent </skill> tag breakout - Require integrity hash for Verified/Community tier skills - Validate skill names against [a-zA-Z0-9][a-zA-Z0-9._-]{0,63} - Add 64 KiB file size limit on prompt.md Bug fixes: - Use actual SkillsConfig from AgentDeps instead of SkillsConfig::default() - Add skills_config field to AgentDeps, wired through from main.rs Performance: - Pre-compile regex patterns at load time (cached on LoadedSkill) - Selector uses pre-compiled patterns instead of recompiling per message - Switch all std::fs to tokio::fs for non-blocking async I/O Hardening: - Cap keyword score at 30 points to prevent keyword stuffing attacks - Enforce max 20 keywords and 5 patterns per skill - Normalize line endings (CRLF/CR to LF) before hashing - Also includes cargo fmt formatting fixes for adjacent code Tests: 54 skills tests pass (up from 41), zero new clippy warnings. Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: Address medium/low severity findings from adversarial review Fixes all 18 medium/low severity findings identified by the security review: - mod.rs: Add MAX_TAGS_PER_SKILL cap (10) in enforce_limits(); use RegexBuilder with 64 KiB size_limit to prevent ReDoS; replace case-enumerated escape_skill_content with regex matching all case variants plus whitespace/null byte injection between </ and skill; document allowed_patterns as unenforced until Phase 2; document Marketplace URL validation as Phase 3 concern - registry.rs: Add MAX_MANIFEST_FILE_SIZE (16 KiB) check before reading; add symlink detection via symlink_metadata to reject symlinks in discover_local; add MAX_DISCOVERED_SKILLS (100) cap; validate prompt_hash format (sha256: + 64 hex chars); warn on name collision before overwriting; accept SkillSource parameter in load_skill instead of always using Local; add InvalidHashFormat, ManifestTooLarge, SymlinkDetected error variants - selector.rs: Add MAX_TAG_SCORE (15) cap parallel to keyword cap; warn when declared max_context_tokens diverges >2x from actual prompt size - scanner.rs: Add mixed-script homoglyph detection (Cyrillic, Greek, Armenian unicode ranges); document token-boundary bypass and semantic paraphrasing as known limitations - attenuation.rs: Document READ_ONLY_TOOLS maintenance requirements - agent_loop.rs: Surface scan warnings via structured tracing; add structured audit events for skill activation and tool attenuation 61 tests pass, 0 new clippy warnings. Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: Address 12 findings from second adversarial security review HIGH: - Escape opening <skill tags in prompt content (prevents fake skill block injection) - Scan manifest metadata fields (description, author, tags, reasons) not just prompt - Block trust downgrade on name collision (existing Local can't be replaced by Community) MEDIUM: - Eliminate TOCTOU gap: read files then check size instead of metadata-then-read - Reject file-level symlinks in load_skill (prompt.md, skill.toml) - Truncate and filter manifest.skill.tags (prevent unlimited tag scoring) - Cap regex pattern score at 40 (prevent 5x20=100 dominating keyword+tag) - Add doc comment about skill_list tool exposing metadata (sanitization required) - Move Community disclaimer inside <skill> tags (not outside structural boundary) - Filter keywords/tags shorter than 3 chars (prevent broad matching) LOW: - Enforce minimum token_cost of 1 (max_context_tokens=0 can't bypass budget) - Remove redundant try_exists checks in discover_local (let load_skill handle errors) 70 skills tests passing. Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: Add HTTP endpoint scoping for skills (Phase 1) Skills that declare an [http] section in skill.toml now have their HTTP requests constrained to declared endpoints at runtime. This addresses the gap where allowed_patterns was parsed but never enforced -- once the http tool was visible via attenuation, the LLM could reach any URL. Enforcement reuses EndpointPattern/AllowlistValidator from the WASM capability system. Semantics: if no active skill declares [http], all requests pass through (backward compat). If any skill declares [http], URLs must match at least one skill's allowlist (union). Community skills' [http] declarations are silently ignored (defense in depth). Shell commands using curl/wget are also validated against scopes. Scanner gains detection for known exfiltration domains (webhook.site, ngrok.io, etc.), overly broad wildcards, and credential/host mismatches. Closes #38 Co-Authored-By: Claude Opus 4.6 <[email protected]> * style: Apply cargo fmt to http_scoping.rs Co-Authored-By: Claude Opus 4.6 <[email protected]> * style: Apply cargo fmt across codebase Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: Add parameter-level permission enforcement for skills (Phase 2) Activates enforcement of `allowed_patterns` in skill.toml permissions. Previously these patterns were parsed but not enforced -- a Verified skill declaring `permissions.shell` with `allowed_patterns = [{command = "cargo *"}]` could still run any shell command. Now the enforcer validates tool parameters against declared glob patterns before execution. Key changes: - New `enforcer.rs` module with `SkillPermissionEnforcer`, `glob_to_regex()`, and `validate_tool_call()` with union semantics across active skills - Typed pattern enums (`ShellPattern`, `FilePathPattern`, `MemoryTargetPattern`) replace the previous `Vec<serde_json::Value>` in `ToolPermissionDeclaration` - Scanner gains `scan_permission_patterns()` detecting dangerous patterns (rm, sudo, curl, bare wildcards, command chaining, sensitive paths, identity files) - Registry blocks non-Local skills with critical permission pattern warnings - Agent loop threads enforcer into `execute_chat_tool` alongside HTTP scoping Trust interaction: Community patterns ignored, Verified enforced, Local without patterns unrestricted, Local with patterns enforced as guidance. Union semantics across skills -- tool call allowed if ANY skill's patterns permit it. 34 new tests. All 818 library tests pass. Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: Add worker permission enforcement and LLM behavioral analysis (Phase 3+4) Phase 3 - Worker-side permission enforcement: - Add SerializedToolPermission/SerializedPattern DTOs for HTTP boundary crossing - Extend JobDescription, ContainerHandle, and orchestrator API to carry permissions - CreateJobTool snapshots and forwards skill permissions to spawned workers - Worker runtime builds SkillPermissionEnforcer and checks before tool execution - Load-time token budget enforcement rejects prompts exceeding 2x declared budget - Deduplicate enforcer construction: from_active_skills() delegates to from_serialized() Phase 4 - LLM behavioral analysis: - BehavioralAnalyzer with cached, LLM-based semantic content analysis - Structured output parsing (FINDING|CATEGORY|SEVERITY|DESCRIPTION or CLEAN) - Content-hash caching with bounded size (MAX_CACHE_ENTRIES=256) - Graceful degradation when LLM unavailable - Integrated into load_skill() for non-Local skills; critical findings block loading Review fixes: - Real cache tests with CountingLlm mock (test_cache_hit, test_cache_miss, test_cache_bounded) - UTF-8-safe truncate() in worker runtime - Few-shot examples in behavioral analysis prompt - Documented max_context_tokens=0 opt-out and create_job() permission gap 848 tests passing, no new clippy warnings. Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: Address review feedback from serrrfirat on skills-phase2 - Fix truncate_cmd UTF-8 panic: use char-boundary-aware slicing - Remove redundant effective_tools branching in reasoning.rs - Document cache eviction as known limitation (arbitrary, not LRU) - Add safety comment on SkillTrust enum ordering (security-critical) - Simplify active_skills selection (prefilter_skills handles empty input) Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: address remaining skills review feedback * refactor: replace skills system with OpenClaw SKILL.md format + 2-state trust Replace the 5-gate, 3-tier trust hierarchy (scanner, behavioral analyzer, parameter-level enforcer, HTTP endpoint scoping) with a simplified 3-layer security model: gating -> attenuation -> Docker confinement. Key changes: - SKILL.md format (YAML frontmatter + markdown prompt) replaces skill.toml + prompt.md - 2-state trust (Installed/Trusted) replaces 3-tier (Community/Verified/Local) - New parser.rs for SKILL.md parsing with serde_yaml - New gating.rs for requirements checking (bins/env/config) - Simplified registry with 2-location discovery (workspace + user dirs) - Removed scanner, behavioral_analyzer, enforcer, http_scoping (~4,100 lines) - Removed skill_permissions propagation through job/orchestrator/worker pipeline - Added serde_yaml dependency for YAML frontmatter parsing Net: -5,298 lines, 59 skills tests pass, 907 total tests pass. Co-Authored-By: Claude Opus 4.6 <[email protected]> * feat: add in-app skill management tools and ClawHub catalog integration Add 4 chat-callable tools (skill_list, skill_search, skill_install, skill_remove) plus matching web gateway endpoints for managing skills at runtime. The catalog fetches from ClawHub's public registry API at runtime rather than bundling entries at compile time. Key changes: - SkillRegistry gains mutation methods (install_skill, remove_skill, reload, find_by_name) with Arc<RwLock> for concurrent access - New catalog module queries ClawHub /api/v1/search with in-memory caching (5-min TTL, configurable via CLAWHUB_REGISTRY env var) - skill_list and skill_search added to READ_ONLY_TOOLS for safe use under Installed trust ceiling - Web gateway gets /api/skills, /api/skills/search, /api/skills/install, and /api/skills/{name} DELETE endpoints Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: address PR #51 review feedback from ilblackdragon Security: - Add SSRF protection to fetch_skill_content: require HTTPS, reject private/loopback/link-local IPs and internal hostnames, disable redirects. Gateway install handler now reuses the same validation. - URL-encode slug in skill_download_url to prevent query injection. - Require X-Confirm-Action header on gateway skill install/remove endpoints (equivalent to chat tool requires_approval gate). Correctness: - Eliminate all block_in_place/block_on usage in skill tools and gateway handlers. Split install into prepare_install_to_disk (static async, no lock) + commit_install (sync, brief write lock). Same pattern for remove: validate_remove + delete_skill_files + commit_remove. - Write normalized content to disk in install_skill (was writing original un-normalized content, causing hash mismatch on re-read). - Fix token estimation from 0.75 to 0.25 tokens/byte (~4 chars per token) in registry.rs, selector.rs, and standalone loader. Dependencies: - Replace deprecated serde_yaml 0.9 with serde_yml 0.0.12. - Remove unused toml dependency. Co-Authored-By: Claude Opus 4.6 <[email protected]> --------- Co-authored-by: Claude Opus 4.6 <[email protected]>
495 lines
17 KiB
Rust
495 lines
17 KiB
Rust
//! WebSocket handler for bidirectional client communication.
|
|
//!
|
|
//! Provides the same event stream as SSE but also accepts incoming messages
|
|
//! (chat, approvals) over a single persistent connection.
|
|
//!
|
|
//! ```text
|
|
//! Client ──── WS frame: {"type":"message","content":"hello"} ──► Agent Loop
|
|
//! ◄─── WS frame: {"type":"event","event_type":"response","data":{...}} ── Broadcast
|
|
//! ──── WS frame: {"type":"ping"} ──────────────────────────────────────►
|
|
//! ◄─── WS frame: {"type":"pong"} ──────────────────────────────────────
|
|
//! ```
|
|
|
|
use std::sync::Arc;
|
|
use std::sync::atomic::{AtomicU64, Ordering};
|
|
|
|
use axum::extract::ws::{Message, WebSocket};
|
|
use futures::{SinkExt, StreamExt};
|
|
use tokio::sync::mpsc;
|
|
use uuid::Uuid;
|
|
|
|
use crate::agent::submission::Submission;
|
|
use crate::channels::IncomingMessage;
|
|
use crate::channels::web::server::GatewayState;
|
|
use crate::channels::web::types::{WsClientMessage, WsServerMessage};
|
|
|
|
/// Tracks active WebSocket connections.
|
|
pub struct WsConnectionTracker {
|
|
count: AtomicU64,
|
|
}
|
|
|
|
impl WsConnectionTracker {
|
|
pub fn new() -> Self {
|
|
Self {
|
|
count: AtomicU64::new(0),
|
|
}
|
|
}
|
|
|
|
pub fn connection_count(&self) -> u64 {
|
|
self.count.load(Ordering::Relaxed)
|
|
}
|
|
|
|
fn increment(&self) {
|
|
self.count.fetch_add(1, Ordering::Relaxed);
|
|
}
|
|
|
|
fn decrement(&self) {
|
|
self.count.fetch_sub(1, Ordering::Relaxed);
|
|
}
|
|
}
|
|
|
|
impl Default for WsConnectionTracker {
|
|
fn default() -> Self {
|
|
Self::new()
|
|
}
|
|
}
|
|
|
|
/// Handle an upgraded WebSocket connection.
|
|
///
|
|
/// Spawns two tasks:
|
|
/// - **sender**: forwards broadcast events to the WebSocket client
|
|
/// - **receiver**: reads client frames and routes them to the agent
|
|
///
|
|
/// When either task ends (client disconnect or broadcast closed), both are
|
|
/// cleaned up.
|
|
pub async fn handle_ws_connection(socket: WebSocket, state: Arc<GatewayState>) {
|
|
let (mut ws_sink, mut ws_stream) = socket.split();
|
|
|
|
// Track connection
|
|
if let Some(ref tracker) = state.ws_tracker {
|
|
tracker.increment();
|
|
}
|
|
let tracker_for_drop = state.ws_tracker.clone();
|
|
|
|
// Subscribe to broadcast events (same source as SSE).
|
|
// Reject if we've hit the connection limit.
|
|
let Some(raw_stream) = state.sse.subscribe_raw() else {
|
|
tracing::warn!("WebSocket rejected: too many connections");
|
|
// Decrement the WS tracker we already incremented above.
|
|
if let Some(ref tracker) = tracker_for_drop {
|
|
tracker.decrement();
|
|
}
|
|
return;
|
|
};
|
|
let mut event_stream = Box::pin(raw_stream);
|
|
|
|
// Channel for the sender task to receive messages from both
|
|
// the broadcast stream and any direct sends (like Pong)
|
|
let (direct_tx, mut direct_rx) = mpsc::channel::<WsServerMessage>(64);
|
|
|
|
// Sender task: forward broadcast events + direct messages to WS client
|
|
let sender_handle = tokio::spawn(async move {
|
|
loop {
|
|
let msg = tokio::select! {
|
|
event = event_stream.next() => {
|
|
match event {
|
|
Some(sse_event) => WsServerMessage::from_sse_event(&sse_event),
|
|
None => break, // Broadcast channel closed
|
|
}
|
|
}
|
|
direct = direct_rx.recv() => {
|
|
match direct {
|
|
Some(msg) => msg,
|
|
None => break, // Direct channel closed
|
|
}
|
|
}
|
|
};
|
|
|
|
let json = match serde_json::to_string(&msg) {
|
|
Ok(j) => j,
|
|
Err(_) => continue,
|
|
};
|
|
|
|
if ws_sink.send(Message::Text(json.into())).await.is_err() {
|
|
break; // Client disconnected
|
|
}
|
|
}
|
|
});
|
|
|
|
// Receiver task: read client frames and route to agent
|
|
let user_id = state.user_id.clone();
|
|
while let Some(Ok(frame)) = ws_stream.next().await {
|
|
match frame {
|
|
Message::Text(text) => {
|
|
let parsed: Result<WsClientMessage, _> = serde_json::from_str(&text);
|
|
match parsed {
|
|
Ok(client_msg) => {
|
|
handle_client_message(client_msg, &state, &user_id, &direct_tx).await;
|
|
}
|
|
Err(e) => {
|
|
let _ = direct_tx
|
|
.send(WsServerMessage::Error {
|
|
message: format!("Invalid message: {}", e),
|
|
})
|
|
.await;
|
|
}
|
|
}
|
|
}
|
|
Message::Close(_) => break,
|
|
// Ignore binary, ping/pong (axum handles protocol-level pings)
|
|
_ => {}
|
|
}
|
|
}
|
|
|
|
// Clean up: abort sender, decrement counter
|
|
sender_handle.abort();
|
|
if let Some(ref tracker) = tracker_for_drop {
|
|
tracker.decrement();
|
|
}
|
|
}
|
|
|
|
/// Route a parsed client message to the appropriate handler.
|
|
async fn handle_client_message(
|
|
msg: WsClientMessage,
|
|
state: &GatewayState,
|
|
user_id: &str,
|
|
direct_tx: &mpsc::Sender<WsServerMessage>,
|
|
) {
|
|
match msg {
|
|
WsClientMessage::Message { content, thread_id } => {
|
|
let mut incoming = IncomingMessage::new("gateway", user_id, &content);
|
|
if let Some(ref tid) = thread_id {
|
|
incoming = incoming.with_thread(tid);
|
|
}
|
|
|
|
let tx_guard = state.msg_tx.read().await;
|
|
if let Some(ref tx) = *tx_guard {
|
|
if tx.send(incoming).await.is_err() {
|
|
let _ = direct_tx
|
|
.send(WsServerMessage::Error {
|
|
message: "Channel closed".to_string(),
|
|
})
|
|
.await;
|
|
}
|
|
} else {
|
|
let _ = direct_tx
|
|
.send(WsServerMessage::Error {
|
|
message: "Channel not started".to_string(),
|
|
})
|
|
.await;
|
|
}
|
|
}
|
|
WsClientMessage::Approval {
|
|
request_id,
|
|
action,
|
|
thread_id,
|
|
} => {
|
|
let (approved, always) = match action.as_str() {
|
|
"approve" => (true, false),
|
|
"always" => (true, true),
|
|
"deny" => (false, false),
|
|
other => {
|
|
let _ = direct_tx
|
|
.send(WsServerMessage::Error {
|
|
message: format!("Unknown approval action: {}", other),
|
|
})
|
|
.await;
|
|
return;
|
|
}
|
|
};
|
|
|
|
let request_uuid = match Uuid::parse_str(&request_id) {
|
|
Ok(id) => id,
|
|
Err(_) => {
|
|
let _ = direct_tx
|
|
.send(WsServerMessage::Error {
|
|
message: "Invalid request_id (expected UUID)".to_string(),
|
|
})
|
|
.await;
|
|
return;
|
|
}
|
|
};
|
|
|
|
let approval = Submission::ExecApproval {
|
|
request_id: request_uuid,
|
|
approved,
|
|
always,
|
|
};
|
|
let content = match serde_json::to_string(&approval) {
|
|
Ok(c) => c,
|
|
Err(e) => {
|
|
let _ = direct_tx
|
|
.send(WsServerMessage::Error {
|
|
message: format!("Failed to serialize approval: {}", e),
|
|
})
|
|
.await;
|
|
return;
|
|
}
|
|
};
|
|
|
|
let mut msg = IncomingMessage::new("gateway", user_id, content);
|
|
if let Some(ref tid) = thread_id {
|
|
msg = msg.with_thread(tid);
|
|
}
|
|
let tx_guard = state.msg_tx.read().await;
|
|
if let Some(ref tx) = *tx_guard {
|
|
let _ = tx.send(msg).await;
|
|
}
|
|
}
|
|
WsClientMessage::AuthToken {
|
|
extension_name,
|
|
token,
|
|
} => {
|
|
if let Some(ref ext_mgr) = state.extension_manager {
|
|
match ext_mgr.auth(&extension_name, Some(&token)).await {
|
|
Ok(result) if result.status == "authenticated" => {
|
|
let msg = match ext_mgr.activate(&extension_name).await {
|
|
Ok(r) => format!(
|
|
"{} authenticated ({} tools loaded)",
|
|
extension_name,
|
|
r.tools_loaded.len()
|
|
),
|
|
Err(e) => format!(
|
|
"{} authenticated but activation failed: {}",
|
|
extension_name, e
|
|
),
|
|
};
|
|
crate::channels::web::server::clear_auth_mode(state).await;
|
|
state
|
|
.sse
|
|
.broadcast(crate::channels::web::types::SseEvent::AuthCompleted {
|
|
extension_name,
|
|
success: true,
|
|
message: msg,
|
|
});
|
|
}
|
|
Ok(result) => {
|
|
state
|
|
.sse
|
|
.broadcast(crate::channels::web::types::SseEvent::AuthRequired {
|
|
extension_name,
|
|
instructions: result.instructions,
|
|
auth_url: result.auth_url,
|
|
setup_url: result.setup_url,
|
|
});
|
|
}
|
|
Err(e) => {
|
|
let _ = direct_tx
|
|
.send(WsServerMessage::Error {
|
|
message: format!("Auth failed: {}", e),
|
|
})
|
|
.await;
|
|
}
|
|
}
|
|
} else {
|
|
let _ = direct_tx
|
|
.send(WsServerMessage::Error {
|
|
message: "Extension manager not available".to_string(),
|
|
})
|
|
.await;
|
|
}
|
|
}
|
|
WsClientMessage::AuthCancel { .. } => {
|
|
crate::channels::web::server::clear_auth_mode(state).await;
|
|
}
|
|
WsClientMessage::Ping => {
|
|
let _ = direct_tx.send(WsServerMessage::Pong).await;
|
|
}
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_ws_connection_tracker() {
|
|
let tracker = WsConnectionTracker::new();
|
|
assert_eq!(tracker.connection_count(), 0);
|
|
|
|
tracker.increment();
|
|
assert_eq!(tracker.connection_count(), 1);
|
|
|
|
tracker.increment();
|
|
assert_eq!(tracker.connection_count(), 2);
|
|
|
|
tracker.decrement();
|
|
assert_eq!(tracker.connection_count(), 1);
|
|
|
|
tracker.decrement();
|
|
assert_eq!(tracker.connection_count(), 0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_ws_connection_tracker_default() {
|
|
let tracker = WsConnectionTracker::default();
|
|
assert_eq!(tracker.connection_count(), 0);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_handle_client_message_ping() {
|
|
// Ping should produce a Pong on the direct channel
|
|
let (direct_tx, mut direct_rx) = mpsc::channel(16);
|
|
let state = make_test_state(None).await;
|
|
|
|
handle_client_message(WsClientMessage::Ping, &state, "user1", &direct_tx).await;
|
|
|
|
let response = direct_rx.recv().await.unwrap();
|
|
assert!(matches!(response, WsServerMessage::Pong));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_handle_client_message_sends_to_agent() {
|
|
// A Message should be forwarded to the agent's msg_tx
|
|
let (agent_tx, mut agent_rx) = mpsc::channel(16);
|
|
let state = make_test_state(Some(agent_tx)).await;
|
|
let (direct_tx, _direct_rx) = mpsc::channel(16);
|
|
|
|
handle_client_message(
|
|
WsClientMessage::Message {
|
|
content: "hello agent".to_string(),
|
|
thread_id: Some("t1".to_string()),
|
|
},
|
|
&state,
|
|
"user1",
|
|
&direct_tx,
|
|
)
|
|
.await;
|
|
|
|
let incoming = agent_rx.recv().await.unwrap();
|
|
assert_eq!(incoming.content, "hello agent");
|
|
assert_eq!(incoming.thread_id.as_deref(), Some("t1"));
|
|
assert_eq!(incoming.channel, "gateway");
|
|
assert_eq!(incoming.user_id, "user1");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_handle_client_message_no_channel() {
|
|
// When msg_tx is None, should send an error back
|
|
let state = make_test_state(None).await;
|
|
let (direct_tx, mut direct_rx) = mpsc::channel(16);
|
|
|
|
handle_client_message(
|
|
WsClientMessage::Message {
|
|
content: "hello".to_string(),
|
|
thread_id: None,
|
|
},
|
|
&state,
|
|
"user1",
|
|
&direct_tx,
|
|
)
|
|
.await;
|
|
|
|
let response = direct_rx.recv().await.unwrap();
|
|
match response {
|
|
WsServerMessage::Error { message } => {
|
|
assert!(message.contains("not started"));
|
|
}
|
|
_ => panic!("Expected Error variant"),
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_handle_client_approval_approve() {
|
|
let (agent_tx, mut agent_rx) = mpsc::channel(16);
|
|
let state = make_test_state(Some(agent_tx)).await;
|
|
let (direct_tx, _direct_rx) = mpsc::channel(16);
|
|
|
|
let request_id = Uuid::new_v4();
|
|
handle_client_message(
|
|
WsClientMessage::Approval {
|
|
request_id: request_id.to_string(),
|
|
action: "approve".to_string(),
|
|
thread_id: Some("thread-42".to_string()),
|
|
},
|
|
&state,
|
|
"user1",
|
|
&direct_tx,
|
|
)
|
|
.await;
|
|
|
|
let incoming = agent_rx.recv().await.unwrap();
|
|
// The content should be a serialized ExecApproval
|
|
assert!(incoming.content.contains("ExecApproval"));
|
|
// Thread should be forwarded onto the IncomingMessage.
|
|
assert_eq!(incoming.thread_id.as_deref(), Some("thread-42"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_handle_client_approval_invalid_action() {
|
|
let state = make_test_state(None).await;
|
|
let (direct_tx, mut direct_rx) = mpsc::channel(16);
|
|
|
|
handle_client_message(
|
|
WsClientMessage::Approval {
|
|
request_id: Uuid::new_v4().to_string(),
|
|
action: "maybe".to_string(),
|
|
thread_id: None,
|
|
},
|
|
&state,
|
|
"user1",
|
|
&direct_tx,
|
|
)
|
|
.await;
|
|
|
|
let response = direct_rx.recv().await.unwrap();
|
|
match response {
|
|
WsServerMessage::Error { message } => {
|
|
assert!(message.contains("Unknown approval action"));
|
|
}
|
|
_ => panic!("Expected Error variant"),
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_handle_client_approval_invalid_uuid() {
|
|
let state = make_test_state(None).await;
|
|
let (direct_tx, mut direct_rx) = mpsc::channel(16);
|
|
|
|
handle_client_message(
|
|
WsClientMessage::Approval {
|
|
request_id: "not-a-uuid".to_string(),
|
|
action: "approve".to_string(),
|
|
thread_id: None,
|
|
},
|
|
&state,
|
|
"user1",
|
|
&direct_tx,
|
|
)
|
|
.await;
|
|
|
|
let response = direct_rx.recv().await.unwrap();
|
|
match response {
|
|
WsServerMessage::Error { message } => {
|
|
assert!(message.contains("Invalid request_id"));
|
|
}
|
|
_ => panic!("Expected Error variant"),
|
|
}
|
|
}
|
|
|
|
/// Helper to create a GatewayState for testing.
|
|
async fn make_test_state(msg_tx: Option<mpsc::Sender<IncomingMessage>>) -> GatewayState {
|
|
use crate::channels::web::sse::SseManager;
|
|
|
|
GatewayState {
|
|
msg_tx: tokio::sync::RwLock::new(msg_tx),
|
|
sse: SseManager::new(),
|
|
workspace: None,
|
|
session_manager: None,
|
|
log_broadcaster: None,
|
|
extension_manager: None,
|
|
tool_registry: None,
|
|
store: None,
|
|
job_manager: None,
|
|
prompt_queue: None,
|
|
user_id: "test".to_string(),
|
|
shutdown_tx: tokio::sync::RwLock::new(None),
|
|
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
|
llm_provider: None,
|
|
skill_registry: None,
|
|
skill_catalog: None,
|
|
chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60),
|
|
}
|
|
}
|
|
}
|