mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
Add WASM sandbox secure API extension
Extends the WASM sandbox with HTTP API capabilities, secrets management, tool aliasing, and leak detection. Key security principle: WASM never sees credentials, injection happens at host boundary. New modules: - secrets: AES-256-GCM encrypted storage with HKDF key derivation - leak_detector: Aho-Corasick + regex pattern matching for secret exfiltration - capabilities: Extended capability system (HTTP, ToolInvoke, Secrets) - allowlist: HTTP endpoint validation with glob patterns - credential_injector: Host-boundary credential injection - rate_limiter: Sliding window per-tool rate limiting - storage: WASM binary storage with BLAKE3 integrity verification Leak detection happens at two points: 1. Before HTTP request (prevents exfiltration via URL/headers/body) 2. After response (prevents exposure in outputs returned to WASM) Co-Authored-By: Claude Opus 4.5 <[email protected]>
This commit is contained in:
co-authored by
Claude Opus 4.5
parent
45bbfa026d
commit
32bfd24154
Generated
+142
@@ -11,6 +11,41 @@ dependencies = [
|
||||
"gimli",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aead"
|
||||
version = "0.5.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0"
|
||||
dependencies = [
|
||||
"crypto-common",
|
||||
"generic-array",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aes"
|
||||
version = "0.8.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cipher",
|
||||
"cpufeatures",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aes-gcm"
|
||||
version = "0.10.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "831010a0f742e1209b3bcea8fab6a8e149051ba6099432c8cb2cc117dec3ead1"
|
||||
dependencies = [
|
||||
"aead",
|
||||
"aes",
|
||||
"cipher",
|
||||
"ctr",
|
||||
"ghash",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ahash"
|
||||
version = "0.7.8"
|
||||
@@ -129,6 +164,12 @@ version = "1.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1"
|
||||
|
||||
[[package]]
|
||||
name = "arrayref"
|
||||
version = "0.3.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "76a2e8124351fda1ef8aaaa3bbd7ebbcb486bbcd4225aca0aa0d84bb2db8fecb"
|
||||
|
||||
[[package]]
|
||||
name = "arrayvec"
|
||||
version = "0.7.6"
|
||||
@@ -246,6 +287,20 @@ dependencies = [
|
||||
"wyz",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "blake3"
|
||||
version = "1.8.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2468ef7d57b3fb7e16b576e8377cdbde2320c60e1491e961d11da40fc4f02a2d"
|
||||
dependencies = [
|
||||
"arrayref",
|
||||
"arrayvec",
|
||||
"cc",
|
||||
"cfg-if",
|
||||
"constant_time_eq",
|
||||
"cpufeatures",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "block-buffer"
|
||||
version = "0.10.4"
|
||||
@@ -409,6 +464,16 @@ dependencies = [
|
||||
"windows-link",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "cipher"
|
||||
version = "0.4.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad"
|
||||
dependencies = [
|
||||
"crypto-common",
|
||||
"inout",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "clap"
|
||||
version = "4.5.56"
|
||||
@@ -464,6 +529,12 @@ version = "1.0.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b05b61dc5112cbb17e4b6cd61790d9845d13888356391624cbe7e41efeac1e75"
|
||||
|
||||
[[package]]
|
||||
name = "constant_time_eq"
|
||||
version = "0.4.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3d52eff69cd5e647efe296129160853a42795992097e8af39800e1060caeea9b"
|
||||
|
||||
[[package]]
|
||||
name = "core-foundation"
|
||||
version = "0.10.1"
|
||||
@@ -646,9 +717,19 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a"
|
||||
dependencies = [
|
||||
"generic-array",
|
||||
"rand_core 0.6.4",
|
||||
"typenum",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ctr"
|
||||
version = "0.9.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0369ee1ad671834580515889b80f2ea915f23b8be8d0daa4bbaf2ac5c7590835"
|
||||
dependencies = [
|
||||
"cipher",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "darling"
|
||||
version = "0.21.3"
|
||||
@@ -1068,6 +1149,16 @@ dependencies = [
|
||||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ghash"
|
||||
version = "0.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f0d8a4362ccb29cb0b265253fb0a2728f592895ee6854fd9bc13f2ffda266ff1"
|
||||
dependencies = [
|
||||
"opaque-debug",
|
||||
"polyval",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "gimli"
|
||||
version = "0.31.1"
|
||||
@@ -1131,6 +1222,15 @@ version = "0.4.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70"
|
||||
|
||||
[[package]]
|
||||
name = "hkdf"
|
||||
version = "0.12.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7b5f8eb2ad728638ea2c7d47a21db23b7b58a72ed6a38256b8a1849f15fbbdf7"
|
||||
dependencies = [
|
||||
"hmac",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "hmac"
|
||||
version = "0.12.1"
|
||||
@@ -1447,6 +1547,15 @@ dependencies = [
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "inout"
|
||||
version = "0.1.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "879f10e63c20629ecabbb64a8010319738c66a5cd0c29b02d63d272b03751d01"
|
||||
dependencies = [
|
||||
"generic-array",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ipnet"
|
||||
version = "2.11.0"
|
||||
@@ -1674,18 +1783,22 @@ dependencies = [
|
||||
name = "near-agent"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"aes-gcm",
|
||||
"aho-corasick",
|
||||
"anyhow",
|
||||
"async-trait",
|
||||
"axum",
|
||||
"blake3",
|
||||
"chrono",
|
||||
"clap",
|
||||
"deadpool-postgres",
|
||||
"dotenvy",
|
||||
"futures",
|
||||
"hkdf",
|
||||
"pgvector",
|
||||
"postgres-types",
|
||||
"pretty_assertions",
|
||||
"rand 0.8.5",
|
||||
"refinery",
|
||||
"regex",
|
||||
"reqwest",
|
||||
@@ -1694,6 +1807,7 @@ dependencies = [
|
||||
"secrecy",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2",
|
||||
"testcontainers-modules",
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
@@ -1775,6 +1889,12 @@ version = "1.70.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe"
|
||||
|
||||
[[package]]
|
||||
name = "opaque-debug"
|
||||
version = "0.3.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381"
|
||||
|
||||
[[package]]
|
||||
name = "openssl-probe"
|
||||
version = "0.2.1"
|
||||
@@ -1888,6 +2008,18 @@ version = "0.3.32"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7edddbd0b52d732b21ad9a5fab5c704c14cd949e5e9a1ec5929a24fded1b904c"
|
||||
|
||||
[[package]]
|
||||
name = "polyval"
|
||||
version = "0.6.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures",
|
||||
"opaque-debug",
|
||||
"universal-hash",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "postcard"
|
||||
version = "1.1.3"
|
||||
@@ -3459,6 +3591,16 @@ version = "0.2.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853"
|
||||
|
||||
[[package]]
|
||||
name = "universal-hash"
|
||||
version = "0.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea"
|
||||
dependencies = [
|
||||
"crypto-common",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "untrusted"
|
||||
version = "0.9.0"
|
||||
|
||||
@@ -67,6 +67,13 @@ pgvector = { version = "0.4", features = ["postgres"] }
|
||||
# WASM sandbox for untrusted tool execution
|
||||
wasmtime = { version = "28", features = ["component-model"] }
|
||||
|
||||
# Cryptography for secrets management
|
||||
aes-gcm = "0.10"
|
||||
hkdf = "0.12"
|
||||
sha2 = "0.10"
|
||||
blake3 = "1"
|
||||
rand = "0.8"
|
||||
|
||||
[dev-dependencies]
|
||||
tokio-test = "0.4"
|
||||
testcontainers-modules = { version = "0.11", features = ["postgres"] }
|
||||
|
||||
@@ -0,0 +1,322 @@
|
||||
-- WASM Secure API Extension
|
||||
-- V2: Secrets management, WASM tool storage, capabilities, and leak detection
|
||||
|
||||
-- ==================== Secrets ====================
|
||||
-- Encrypted secret storage for credential injection into WASM HTTP requests.
|
||||
-- WASM tools NEVER see plaintext secrets; injection happens at host boundary.
|
||||
|
||||
CREATE TABLE secrets (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
user_id TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
|
||||
-- AES-256-GCM encrypted value (nonce || ciphertext || tag)
|
||||
encrypted_value BYTEA NOT NULL,
|
||||
-- Per-secret key derivation salt (for HKDF)
|
||||
key_salt BYTEA NOT NULL,
|
||||
|
||||
-- Optional metadata
|
||||
provider TEXT, -- e.g., "openai", "anthropic", "stripe"
|
||||
expires_at TIMESTAMPTZ,
|
||||
last_used_at TIMESTAMPTZ,
|
||||
usage_count BIGINT NOT NULL DEFAULT 0,
|
||||
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
|
||||
CONSTRAINT unique_secret_per_user UNIQUE (user_id, name)
|
||||
);
|
||||
|
||||
CREATE INDEX idx_secrets_user ON secrets(user_id);
|
||||
CREATE INDEX idx_secrets_provider ON secrets(provider) WHERE provider IS NOT NULL;
|
||||
CREATE INDEX idx_secrets_expires ON secrets(expires_at) WHERE expires_at IS NOT NULL;
|
||||
|
||||
-- Trigger to update updated_at
|
||||
CREATE TRIGGER update_secrets_updated_at
|
||||
BEFORE UPDATE ON secrets
|
||||
FOR EACH ROW
|
||||
EXECUTE FUNCTION update_updated_at_column();
|
||||
|
||||
-- ==================== WASM Tools ====================
|
||||
-- Store compiled WASM binaries with integrity verification.
|
||||
|
||||
CREATE TABLE wasm_tools (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
user_id TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
version TEXT NOT NULL DEFAULT '1.0.0',
|
||||
|
||||
description TEXT NOT NULL,
|
||||
wasm_binary BYTEA NOT NULL,
|
||||
-- BLAKE3 hash for integrity verification on load
|
||||
binary_hash BYTEA NOT NULL,
|
||||
parameters_schema JSONB NOT NULL,
|
||||
|
||||
-- Provenance
|
||||
source_url TEXT,
|
||||
-- Trust levels: 'system' (built-in), 'verified' (audited), 'user' (untrusted)
|
||||
trust_level TEXT NOT NULL DEFAULT 'user',
|
||||
|
||||
-- Status: 'active', 'disabled', 'quarantined'
|
||||
status TEXT NOT NULL DEFAULT 'active',
|
||||
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
|
||||
CONSTRAINT unique_wasm_tool_version UNIQUE (user_id, name, version)
|
||||
);
|
||||
|
||||
CREATE INDEX idx_wasm_tools_user ON wasm_tools(user_id);
|
||||
CREATE INDEX idx_wasm_tools_name ON wasm_tools(user_id, name);
|
||||
CREATE INDEX idx_wasm_tools_status ON wasm_tools(status);
|
||||
CREATE INDEX idx_wasm_tools_trust ON wasm_tools(trust_level);
|
||||
|
||||
CREATE TRIGGER update_wasm_tools_updated_at
|
||||
BEFORE UPDATE ON wasm_tools
|
||||
FOR EACH ROW
|
||||
EXECUTE FUNCTION update_updated_at_column();
|
||||
|
||||
-- ==================== Tool Capabilities ====================
|
||||
-- Fine-grained capability configuration per WASM tool.
|
||||
-- Follows principle of least privilege.
|
||||
|
||||
CREATE TABLE tool_capabilities (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
wasm_tool_id UUID NOT NULL REFERENCES wasm_tools(id) ON DELETE CASCADE,
|
||||
|
||||
-- HTTP capability: allowed endpoint patterns
|
||||
-- Each pattern is: {"host": "api.example.com", "path_prefix": "/v1/", "methods": ["GET", "POST"]}
|
||||
http_allowlist JSONB NOT NULL DEFAULT '[]',
|
||||
|
||||
-- Secrets this tool can use (injected at host boundary)
|
||||
-- Tool never sees the actual secret values
|
||||
allowed_secrets TEXT[] NOT NULL DEFAULT '{}',
|
||||
|
||||
-- Tool invocation aliases (indirection layer)
|
||||
-- Maps alias name to real tool name, e.g., {"search": "brave_search"}
|
||||
tool_aliases JSONB NOT NULL DEFAULT '{}',
|
||||
|
||||
-- Rate limiting
|
||||
requests_per_minute INT NOT NULL DEFAULT 60,
|
||||
requests_per_hour INT NOT NULL DEFAULT 1000,
|
||||
|
||||
-- Request/response size limits
|
||||
max_request_body_bytes BIGINT NOT NULL DEFAULT 1048576, -- 1 MB
|
||||
max_response_body_bytes BIGINT NOT NULL DEFAULT 10485760, -- 10 MB
|
||||
|
||||
-- Workspace access (path prefixes tool can read)
|
||||
workspace_read_prefixes TEXT[] NOT NULL DEFAULT '{}',
|
||||
|
||||
-- Timeout for HTTP requests (seconds)
|
||||
http_timeout_secs INT NOT NULL DEFAULT 30,
|
||||
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
|
||||
CONSTRAINT unique_capabilities_per_tool UNIQUE (wasm_tool_id)
|
||||
);
|
||||
|
||||
CREATE INDEX idx_tool_capabilities_tool ON tool_capabilities(wasm_tool_id);
|
||||
|
||||
CREATE TRIGGER update_tool_capabilities_updated_at
|
||||
BEFORE UPDATE ON tool_capabilities
|
||||
FOR EACH ROW
|
||||
EXECUTE FUNCTION update_updated_at_column();
|
||||
|
||||
-- ==================== Leak Detection Patterns ====================
|
||||
-- Patterns for detecting secret leakage in tool outputs.
|
||||
-- Scanned before returning data to WASM or LLM.
|
||||
|
||||
CREATE TABLE leak_detection_patterns (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
|
||||
-- Regex pattern for detection
|
||||
pattern TEXT NOT NULL,
|
||||
|
||||
-- Severity: 'critical', 'high', 'medium', 'low'
|
||||
severity TEXT NOT NULL DEFAULT 'high',
|
||||
|
||||
-- Action: 'block' (fail request), 'redact' (mask secret), 'warn' (log only)
|
||||
action TEXT NOT NULL DEFAULT 'block',
|
||||
|
||||
enabled BOOLEAN NOT NULL DEFAULT true,
|
||||
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
);
|
||||
|
||||
CREATE INDEX idx_leak_patterns_enabled ON leak_detection_patterns(enabled) WHERE enabled = true;
|
||||
|
||||
-- Pre-populate with common API key patterns
|
||||
INSERT INTO leak_detection_patterns (name, pattern, severity, action) VALUES
|
||||
-- OpenAI (sk-proj-... or sk-... followed by alphanumeric)
|
||||
('openai_api_key', 'sk-(?:proj-)?[a-zA-Z0-9]{20,}(?:T3BlbkFJ[a-zA-Z0-9_-]*)?', 'critical', 'block'),
|
||||
|
||||
-- Anthropic (sk-ant-api followed by 90+ chars)
|
||||
('anthropic_api_key', 'sk-ant-api[a-zA-Z0-9_-]{90,}', 'critical', 'block'),
|
||||
|
||||
-- AWS Access Key ID (starts with AKIA)
|
||||
('aws_access_key', 'AKIA[0-9A-Z]{16}', 'critical', 'block'),
|
||||
|
||||
-- AWS Secret Access Key (40 char base64-ish)
|
||||
('aws_secret_key', '(?<![A-Za-z0-9/+=])[A-Za-z0-9/+=]{40}(?![A-Za-z0-9/+=])', 'high', 'block'),
|
||||
|
||||
-- GitHub tokens (gh[pousr]_...)
|
||||
('github_token', 'gh[pousr]_[A-Za-z0-9_]{36,}', 'critical', 'block'),
|
||||
|
||||
-- GitHub fine-grained PAT
|
||||
('github_fine_grained_pat', 'github_pat_[a-zA-Z0-9]{22}_[a-zA-Z0-9]{59}', 'critical', 'block'),
|
||||
|
||||
-- Stripe keys (sk_live_... or sk_test_...)
|
||||
('stripe_api_key', 'sk_(?:live|test)_[a-zA-Z0-9]{24,}', 'critical', 'block'),
|
||||
|
||||
-- NEAR AI session tokens
|
||||
('nearai_session', 'sess_[a-zA-Z0-9]{32,}', 'critical', 'block'),
|
||||
|
||||
-- Generic Bearer tokens in headers
|
||||
('bearer_token', 'Bearer\s+[a-zA-Z0-9_-]{20,}', 'high', 'redact'),
|
||||
|
||||
-- PEM private keys
|
||||
('pem_private_key', '-----BEGIN\s+(?:RSA\s+)?PRIVATE\s+KEY-----', 'critical', 'block'),
|
||||
|
||||
-- SSH private keys
|
||||
('ssh_private_key', '-----BEGIN\s+(?:OPENSSH|EC|DSA)\s+PRIVATE\s+KEY-----', 'critical', 'block'),
|
||||
|
||||
-- Google API keys
|
||||
('google_api_key', 'AIza[0-9A-Za-z_-]{35}', 'high', 'block'),
|
||||
|
||||
-- Slack tokens
|
||||
('slack_token', 'xox[baprs]-[0-9a-zA-Z-]{10,}', 'high', 'block'),
|
||||
|
||||
-- Discord tokens
|
||||
('discord_token', '[MN][A-Za-z\d]{23,}\.[\w-]{6}\.[\w-]{27}', 'high', 'block'),
|
||||
|
||||
-- Twilio (starts with SK)
|
||||
('twilio_api_key', 'SK[a-fA-F0-9]{32}', 'high', 'block'),
|
||||
|
||||
-- SendGrid
|
||||
('sendgrid_api_key', 'SG\.[a-zA-Z0-9_-]{22}\.[a-zA-Z0-9_-]{43}', 'high', 'block'),
|
||||
|
||||
-- Mailchimp
|
||||
('mailchimp_api_key', '[a-f0-9]{32}-us[0-9]{1,2}', 'medium', 'block'),
|
||||
|
||||
-- Generic high-entropy strings (potential secrets) - careful with false positives
|
||||
('high_entropy_hex', '(?<![a-fA-F0-9])[a-fA-F0-9]{64}(?![a-fA-F0-9])', 'medium', 'warn');
|
||||
|
||||
-- ==================== Rate Limit State ====================
|
||||
-- Track rate limit consumption per tool per user.
|
||||
|
||||
CREATE TABLE tool_rate_limit_state (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
wasm_tool_id UUID NOT NULL REFERENCES wasm_tools(id) ON DELETE CASCADE,
|
||||
user_id TEXT NOT NULL,
|
||||
|
||||
-- Sliding window counters
|
||||
minute_window_start TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
minute_count INT NOT NULL DEFAULT 0,
|
||||
hour_window_start TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
hour_count INT NOT NULL DEFAULT 0,
|
||||
|
||||
CONSTRAINT unique_rate_limit_per_tool_user UNIQUE (wasm_tool_id, user_id)
|
||||
);
|
||||
|
||||
CREATE INDEX idx_rate_limit_tool ON tool_rate_limit_state(wasm_tool_id);
|
||||
CREATE INDEX idx_rate_limit_user ON tool_rate_limit_state(user_id);
|
||||
|
||||
-- ==================== Secret Usage Audit Log ====================
|
||||
-- Audit trail for secret access (credential injection events).
|
||||
|
||||
CREATE TABLE secret_usage_log (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
secret_id UUID NOT NULL REFERENCES secrets(id) ON DELETE CASCADE,
|
||||
wasm_tool_id UUID REFERENCES wasm_tools(id) ON DELETE SET NULL,
|
||||
user_id TEXT NOT NULL,
|
||||
|
||||
-- What endpoint was the secret injected for
|
||||
target_host TEXT NOT NULL,
|
||||
target_path TEXT,
|
||||
|
||||
-- Result of the operation
|
||||
success BOOLEAN NOT NULL,
|
||||
error_message TEXT,
|
||||
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
);
|
||||
|
||||
CREATE INDEX idx_secret_usage_secret ON secret_usage_log(secret_id);
|
||||
CREATE INDEX idx_secret_usage_tool ON secret_usage_log(wasm_tool_id);
|
||||
CREATE INDEX idx_secret_usage_user ON secret_usage_log(user_id);
|
||||
CREATE INDEX idx_secret_usage_created ON secret_usage_log(created_at DESC);
|
||||
|
||||
-- Partition by month for large deployments (optional, commented out)
|
||||
-- CREATE TABLE secret_usage_log_y2024m01 PARTITION OF secret_usage_log
|
||||
-- FOR VALUES FROM ('2024-01-01') TO ('2024-02-01');
|
||||
|
||||
-- ==================== Leak Detection Events ====================
|
||||
-- Log when potential secret leaks are detected and blocked.
|
||||
|
||||
CREATE TABLE leak_detection_events (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
pattern_id UUID REFERENCES leak_detection_patterns(id) ON DELETE SET NULL,
|
||||
wasm_tool_id UUID REFERENCES wasm_tools(id) ON DELETE SET NULL,
|
||||
user_id TEXT NOT NULL,
|
||||
|
||||
-- Where the leak was detected
|
||||
source TEXT NOT NULL, -- 'http_response', 'tool_output', 'log_message'
|
||||
action_taken TEXT NOT NULL, -- 'blocked', 'redacted', 'warned'
|
||||
|
||||
-- Redacted context (no actual secrets stored)
|
||||
context_preview TEXT, -- First 100 chars with secret masked
|
||||
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
);
|
||||
|
||||
CREATE INDEX idx_leak_events_pattern ON leak_detection_events(pattern_id);
|
||||
CREATE INDEX idx_leak_events_tool ON leak_detection_events(wasm_tool_id);
|
||||
CREATE INDEX idx_leak_events_user ON leak_detection_events(user_id);
|
||||
CREATE INDEX idx_leak_events_created ON leak_detection_events(created_at DESC);
|
||||
|
||||
-- ==================== Views ====================
|
||||
|
||||
-- View: Tools with their capabilities
|
||||
CREATE VIEW wasm_tools_with_capabilities AS
|
||||
SELECT
|
||||
t.id,
|
||||
t.user_id,
|
||||
t.name,
|
||||
t.version,
|
||||
t.description,
|
||||
t.trust_level,
|
||||
t.status,
|
||||
t.created_at,
|
||||
t.updated_at,
|
||||
c.http_allowlist,
|
||||
c.allowed_secrets,
|
||||
c.tool_aliases,
|
||||
c.requests_per_minute,
|
||||
c.requests_per_hour,
|
||||
c.workspace_read_prefixes
|
||||
FROM wasm_tools t
|
||||
LEFT JOIN tool_capabilities c ON c.wasm_tool_id = t.id;
|
||||
|
||||
-- View: Active leak detection patterns
|
||||
CREATE VIEW active_leak_patterns AS
|
||||
SELECT id, name, pattern, severity, action
|
||||
FROM leak_detection_patterns
|
||||
WHERE enabled = true;
|
||||
|
||||
-- View: Recent leak events summary
|
||||
CREATE VIEW recent_leak_events AS
|
||||
SELECT
|
||||
le.created_at,
|
||||
le.source,
|
||||
le.action_taken,
|
||||
lp.name as pattern_name,
|
||||
lp.severity,
|
||||
wt.name as tool_name,
|
||||
le.user_id
|
||||
FROM leak_detection_events le
|
||||
LEFT JOIN leak_detection_patterns lp ON lp.id = le.pattern_id
|
||||
LEFT JOIN wasm_tools wt ON wt.id = le.wasm_tool_id
|
||||
WHERE le.created_at > NOW() - INTERVAL '24 hours'
|
||||
ORDER BY le.created_at DESC;
|
||||
@@ -16,6 +16,7 @@ pub struct Config {
|
||||
pub agent: AgentConfig,
|
||||
pub safety: SafetyConfig,
|
||||
pub wasm: WasmConfig,
|
||||
pub secrets: SecretsConfig,
|
||||
}
|
||||
|
||||
impl Config {
|
||||
@@ -31,6 +32,7 @@ impl Config {
|
||||
agent: AgentConfig::from_env()?,
|
||||
safety: SafetyConfig::from_env()?,
|
||||
wasm: WasmConfig::from_env()?,
|
||||
secrets: SecretsConfig::from_env()?,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -250,6 +252,52 @@ pub struct WasmConfig {
|
||||
pub cache_dir: Option<PathBuf>,
|
||||
}
|
||||
|
||||
/// Secrets management configuration.
|
||||
#[derive(Clone, Default)]
|
||||
pub struct SecretsConfig {
|
||||
/// Master key for encrypting secrets (loaded from SECRETS_MASTER_KEY env var).
|
||||
/// Must be at least 32 bytes for AES-256-GCM.
|
||||
pub master_key: Option<SecretString>,
|
||||
/// Whether secrets management is enabled.
|
||||
pub enabled: bool,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for SecretsConfig {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("SecretsConfig")
|
||||
.field("master_key", &self.master_key.is_some())
|
||||
.field("enabled", &self.enabled)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl SecretsConfig {
|
||||
fn from_env() -> Result<Self, ConfigError> {
|
||||
let master_key = optional_env("SECRETS_MASTER_KEY")?.map(SecretString::from);
|
||||
let enabled = master_key.is_some();
|
||||
|
||||
// Validate master key length if provided
|
||||
if let Some(ref key) = master_key {
|
||||
if key.expose_secret().len() < 32 {
|
||||
return Err(ConfigError::InvalidValue {
|
||||
key: "SECRETS_MASTER_KEY".to_string(),
|
||||
message: "must be at least 32 bytes for AES-256-GCM".to_string(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
master_key,
|
||||
enabled,
|
||||
})
|
||||
}
|
||||
|
||||
/// Get the master key if configured.
|
||||
pub fn master_key(&self) -> Option<&SecretString> {
|
||||
self.master_key.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for WasmConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
|
||||
@@ -48,6 +48,7 @@ pub mod evaluation;
|
||||
pub mod history;
|
||||
pub mod llm;
|
||||
pub mod safety;
|
||||
pub mod secrets;
|
||||
pub mod tools;
|
||||
pub mod workspace;
|
||||
|
||||
|
||||
@@ -0,0 +1,710 @@
|
||||
//! Secret leak detection for WASM sandbox.
|
||||
//!
|
||||
//! Scans data at the sandbox boundary to prevent secret exfiltration.
|
||||
//! Uses Aho-Corasick for fast multi-pattern matching plus regex for
|
||||
//! complex patterns.
|
||||
//!
|
||||
//! # Security Model
|
||||
//!
|
||||
//! Leak detection happens at TWO points:
|
||||
//!
|
||||
//! 1. **Before outbound requests** - Prevents WASM from exfiltrating secrets
|
||||
//! by encoding them in URLs, headers, or request bodies
|
||||
//! 2. **After responses/outputs** - Prevents accidental exposure in logs,
|
||||
//! tool outputs, or data returned to WASM
|
||||
//!
|
||||
//! # Architecture
|
||||
//!
|
||||
//! ```text
|
||||
//! ┌─────────────────────────────────────────────────────────────────────────────┐
|
||||
//! │ WASM HTTP Request Flow │
|
||||
//! │ │
|
||||
//! │ WASM ──► Allowlist ──► Leak Scan ──► Credential ──► Execute ──► Response │
|
||||
//! │ Validator (request) Injector Request │ │
|
||||
//! │ ▼ │
|
||||
//! │ WASM ◀── Leak Scan ◀── Response │
|
||||
//! │ (response) │
|
||||
//! └─────────────────────────────────────────────────────────────────────────────┘
|
||||
//!
|
||||
//! ┌─────────────────────────────────────────────────────────────────────────────┐
|
||||
//! │ Scan Result Actions │
|
||||
//! │ │
|
||||
//! │ LeakDetector.scan() ──► LeakScanResult │
|
||||
//! │ │ │
|
||||
//! │ ├─► clean: pass through │
|
||||
//! │ ├─► warn: log, pass │
|
||||
//! │ ├─► redact: mask secret │
|
||||
//! │ └─► block: reject entirely │
|
||||
//! └─────────────────────────────────────────────────────────────────────────────┘
|
||||
//! ```
|
||||
|
||||
use std::ops::Range;
|
||||
|
||||
use aho_corasick::AhoCorasick;
|
||||
use regex::Regex;
|
||||
|
||||
/// Action to take when a leak is detected.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum LeakAction {
|
||||
/// Block the output entirely (for critical secrets).
|
||||
Block,
|
||||
/// Redact the secret, replacing it with [REDACTED].
|
||||
Redact,
|
||||
/// Log a warning but allow the output.
|
||||
Warn,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for LeakAction {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
LeakAction::Block => write!(f, "block"),
|
||||
LeakAction::Redact => write!(f, "redact"),
|
||||
LeakAction::Warn => write!(f, "warn"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Severity of a detected leak.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
|
||||
pub enum LeakSeverity {
|
||||
Low,
|
||||
Medium,
|
||||
High,
|
||||
Critical,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for LeakSeverity {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
LeakSeverity::Low => write!(f, "low"),
|
||||
LeakSeverity::Medium => write!(f, "medium"),
|
||||
LeakSeverity::High => write!(f, "high"),
|
||||
LeakSeverity::Critical => write!(f, "critical"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// A pattern for detecting secret leaks.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct LeakPattern {
|
||||
pub name: String,
|
||||
pub regex: Regex,
|
||||
pub severity: LeakSeverity,
|
||||
pub action: LeakAction,
|
||||
}
|
||||
|
||||
/// A detected potential secret leak.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct LeakMatch {
|
||||
pub pattern_name: String,
|
||||
pub severity: LeakSeverity,
|
||||
pub action: LeakAction,
|
||||
/// Location in the scanned content.
|
||||
pub location: Range<usize>,
|
||||
/// A preview of the match with the secret partially masked.
|
||||
pub masked_preview: String,
|
||||
}
|
||||
|
||||
/// Result of scanning content for leaks.
|
||||
#[derive(Debug)]
|
||||
pub struct LeakScanResult {
|
||||
/// All detected potential leaks.
|
||||
pub matches: Vec<LeakMatch>,
|
||||
/// Whether any match requires blocking.
|
||||
pub should_block: bool,
|
||||
/// Content with secrets redacted (if redaction was applied).
|
||||
pub redacted_content: Option<String>,
|
||||
}
|
||||
|
||||
impl LeakScanResult {
|
||||
/// Check if content is clean (no leaks detected).
|
||||
pub fn is_clean(&self) -> bool {
|
||||
self.matches.is_empty()
|
||||
}
|
||||
|
||||
/// Get the highest severity found.
|
||||
pub fn max_severity(&self) -> Option<LeakSeverity> {
|
||||
self.matches.iter().map(|m| m.severity).max()
|
||||
}
|
||||
}
|
||||
|
||||
/// Detector for secret leaks in output data.
|
||||
pub struct LeakDetector {
|
||||
patterns: Vec<LeakPattern>,
|
||||
/// For fast prefix matching of known patterns
|
||||
prefix_matcher: Option<AhoCorasick>,
|
||||
known_prefixes: Vec<(String, usize)>, // (prefix, pattern_index)
|
||||
}
|
||||
|
||||
impl LeakDetector {
|
||||
/// Create a new detector with default patterns.
|
||||
pub fn new() -> Self {
|
||||
Self::with_patterns(default_patterns())
|
||||
}
|
||||
|
||||
/// Create a detector with custom patterns.
|
||||
pub fn with_patterns(patterns: Vec<LeakPattern>) -> Self {
|
||||
// Build prefix matcher for patterns that start with a known prefix
|
||||
let mut prefixes = Vec::new();
|
||||
for (idx, pattern) in patterns.iter().enumerate() {
|
||||
if let Some(prefix) = extract_literal_prefix(pattern.regex.as_str()) {
|
||||
if prefix.len() >= 3 {
|
||||
prefixes.push((prefix, idx));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let prefix_matcher = if !prefixes.is_empty() {
|
||||
let prefix_strings: Vec<&str> = prefixes.iter().map(|(s, _)| s.as_str()).collect();
|
||||
AhoCorasick::builder()
|
||||
.ascii_case_insensitive(false)
|
||||
.build(&prefix_strings)
|
||||
.ok()
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
Self {
|
||||
patterns,
|
||||
prefix_matcher,
|
||||
known_prefixes: prefixes,
|
||||
}
|
||||
}
|
||||
|
||||
/// Scan content for potential secret leaks.
|
||||
pub fn scan(&self, content: &str) -> LeakScanResult {
|
||||
let mut matches = Vec::new();
|
||||
let mut should_block = false;
|
||||
let mut redact_ranges = Vec::new();
|
||||
|
||||
// Use prefix matcher for quick elimination
|
||||
let candidate_indices: Vec<usize> = if let Some(ref matcher) = self.prefix_matcher {
|
||||
let mut indices = Vec::new();
|
||||
for mat in matcher.find_iter(content) {
|
||||
let pattern_idx = self.known_prefixes[mat.pattern().as_usize()].1;
|
||||
if !indices.contains(&pattern_idx) {
|
||||
indices.push(pattern_idx);
|
||||
}
|
||||
}
|
||||
// Also include patterns without prefixes
|
||||
for (idx, _) in self.patterns.iter().enumerate() {
|
||||
if !self.known_prefixes.iter().any(|(_, i)| *i == idx) && !indices.contains(&idx) {
|
||||
indices.push(idx);
|
||||
}
|
||||
}
|
||||
indices
|
||||
} else {
|
||||
(0..self.patterns.len()).collect()
|
||||
};
|
||||
|
||||
// Check candidate patterns
|
||||
for idx in candidate_indices {
|
||||
let pattern = &self.patterns[idx];
|
||||
for mat in pattern.regex.find_iter(content) {
|
||||
let matched_text = mat.as_str();
|
||||
let location = mat.start()..mat.end();
|
||||
|
||||
let leak_match = LeakMatch {
|
||||
pattern_name: pattern.name.clone(),
|
||||
severity: pattern.severity,
|
||||
action: pattern.action,
|
||||
location: location.clone(),
|
||||
masked_preview: mask_secret(matched_text),
|
||||
};
|
||||
|
||||
if pattern.action == LeakAction::Block {
|
||||
should_block = true;
|
||||
}
|
||||
|
||||
if pattern.action == LeakAction::Redact {
|
||||
redact_ranges.push(location.clone());
|
||||
}
|
||||
|
||||
matches.push(leak_match);
|
||||
}
|
||||
}
|
||||
|
||||
// Sort by location for proper redaction
|
||||
matches.sort_by_key(|m| m.location.start);
|
||||
redact_ranges.sort_by_key(|r| r.start);
|
||||
|
||||
// Build redacted content if needed
|
||||
let redacted_content = if !redact_ranges.is_empty() {
|
||||
Some(apply_redactions(content, &redact_ranges))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
LeakScanResult {
|
||||
matches,
|
||||
should_block,
|
||||
redacted_content,
|
||||
}
|
||||
}
|
||||
|
||||
/// Scan content and return cleaned version based on action.
|
||||
///
|
||||
/// Returns `Err` if content should be blocked, `Ok(content)` otherwise.
|
||||
pub fn scan_and_clean(&self, content: &str) -> Result<String, LeakDetectionError> {
|
||||
let result = self.scan(content);
|
||||
|
||||
if result.should_block {
|
||||
// Find the blocking match for error message
|
||||
let blocking_match = result
|
||||
.matches
|
||||
.iter()
|
||||
.find(|m| m.action == LeakAction::Block);
|
||||
return Err(LeakDetectionError::SecretLeakBlocked {
|
||||
pattern: blocking_match
|
||||
.map(|m| m.pattern_name.clone())
|
||||
.unwrap_or_default(),
|
||||
preview: blocking_match
|
||||
.map(|m| m.masked_preview.clone())
|
||||
.unwrap_or_default(),
|
||||
});
|
||||
}
|
||||
|
||||
// Log warnings
|
||||
for m in &result.matches {
|
||||
if m.action == LeakAction::Warn {
|
||||
tracing::warn!(
|
||||
pattern = %m.pattern_name,
|
||||
severity = %m.severity,
|
||||
preview = %m.masked_preview,
|
||||
"Potential secret leak detected (warning only)"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// Return redacted content if any, otherwise original
|
||||
Ok(result
|
||||
.redacted_content
|
||||
.unwrap_or_else(|| content.to_string()))
|
||||
}
|
||||
|
||||
/// Scan an outbound HTTP request for potential secret leakage.
|
||||
///
|
||||
/// This MUST be called before executing any HTTP request from WASM
|
||||
/// to prevent exfiltration of secrets via URL, headers, or body.
|
||||
///
|
||||
/// Returns `Err` if any part contains a blocked secret pattern.
|
||||
pub fn scan_http_request(
|
||||
&self,
|
||||
url: &str,
|
||||
headers: &[(String, String)],
|
||||
body: Option<&[u8]>,
|
||||
) -> Result<(), LeakDetectionError> {
|
||||
// Scan URL (most common exfiltration vector)
|
||||
self.scan_and_clean(url)?;
|
||||
|
||||
// Scan each header value
|
||||
for (name, value) in headers {
|
||||
self.scan_and_clean(value).map_err(|e| {
|
||||
LeakDetectionError::SecretLeakBlocked {
|
||||
pattern: format!("header:{}", name),
|
||||
preview: e.to_string(),
|
||||
}
|
||||
})?;
|
||||
}
|
||||
|
||||
// Scan body if present and valid UTF-8
|
||||
if let Some(body_bytes) = body {
|
||||
if let Ok(body_str) = std::str::from_utf8(body_bytes) {
|
||||
self.scan_and_clean(body_str)?;
|
||||
}
|
||||
// Binary bodies are not scanned (could add hex pattern detection later)
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Add a custom pattern at runtime.
|
||||
pub fn add_pattern(&mut self, pattern: LeakPattern) {
|
||||
self.patterns.push(pattern);
|
||||
// Note: prefix_matcher won't be updated; rebuild if needed
|
||||
}
|
||||
|
||||
/// Get the number of patterns.
|
||||
pub fn pattern_count(&self) -> usize {
|
||||
self.patterns.len()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for LeakDetector {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// Error from leak detection.
|
||||
#[derive(Debug, Clone, thiserror::Error)]
|
||||
pub enum LeakDetectionError {
|
||||
#[error("Secret leak blocked: pattern '{pattern}' matched '{preview}'")]
|
||||
SecretLeakBlocked { pattern: String, preview: String },
|
||||
}
|
||||
|
||||
/// Mask a secret for safe display.
|
||||
///
|
||||
/// Shows first 4 and last 4 characters, masks the middle.
|
||||
fn mask_secret(secret: &str) -> String {
|
||||
let len = secret.len();
|
||||
if len <= 8 {
|
||||
return "*".repeat(len);
|
||||
}
|
||||
|
||||
let prefix: String = secret.chars().take(4).collect();
|
||||
let suffix: String = secret.chars().skip(len - 4).collect();
|
||||
let middle_len = len - 8;
|
||||
format!("{}{}{}", prefix, "*".repeat(middle_len.min(8)), suffix)
|
||||
}
|
||||
|
||||
/// Apply redaction ranges to content.
|
||||
fn apply_redactions(content: &str, ranges: &[Range<usize>]) -> String {
|
||||
if ranges.is_empty() {
|
||||
return content.to_string();
|
||||
}
|
||||
|
||||
let mut result = String::with_capacity(content.len());
|
||||
let mut last_end = 0;
|
||||
|
||||
for range in ranges {
|
||||
if range.start > last_end {
|
||||
result.push_str(&content[last_end..range.start]);
|
||||
}
|
||||
result.push_str("[REDACTED]");
|
||||
last_end = range.end;
|
||||
}
|
||||
|
||||
if last_end < content.len() {
|
||||
result.push_str(&content[last_end..]);
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
/// Extract a literal prefix from a regex pattern (if one exists).
|
||||
fn extract_literal_prefix(pattern: &str) -> Option<String> {
|
||||
let mut prefix = String::new();
|
||||
|
||||
for ch in pattern.chars() {
|
||||
match ch {
|
||||
// These start special regex constructs
|
||||
'[' | '(' | '.' | '*' | '+' | '?' | '{' | '|' | '^' | '$' => break,
|
||||
// Escape sequence
|
||||
'\\' => break,
|
||||
// Regular character
|
||||
_ => prefix.push(ch),
|
||||
}
|
||||
}
|
||||
|
||||
if prefix.len() >= 3 {
|
||||
Some(prefix)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
/// Default leak detection patterns.
|
||||
fn default_patterns() -> Vec<LeakPattern> {
|
||||
vec![
|
||||
// OpenAI API keys
|
||||
LeakPattern {
|
||||
name: "openai_api_key".to_string(),
|
||||
regex: Regex::new(r"sk-(?:proj-)?[a-zA-Z0-9]{20,}(?:T3BlbkFJ[a-zA-Z0-9_-]*)?").unwrap(),
|
||||
severity: LeakSeverity::Critical,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// Anthropic API keys
|
||||
LeakPattern {
|
||||
name: "anthropic_api_key".to_string(),
|
||||
regex: Regex::new(r"sk-ant-api[a-zA-Z0-9_-]{90,}").unwrap(),
|
||||
severity: LeakSeverity::Critical,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// AWS Access Key ID
|
||||
LeakPattern {
|
||||
name: "aws_access_key".to_string(),
|
||||
regex: Regex::new(r"AKIA[0-9A-Z]{16}").unwrap(),
|
||||
severity: LeakSeverity::Critical,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// GitHub tokens
|
||||
LeakPattern {
|
||||
name: "github_token".to_string(),
|
||||
regex: Regex::new(r"gh[pousr]_[A-Za-z0-9_]{36,}").unwrap(),
|
||||
severity: LeakSeverity::Critical,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// GitHub fine-grained PAT
|
||||
LeakPattern {
|
||||
name: "github_fine_grained_pat".to_string(),
|
||||
regex: Regex::new(r"github_pat_[a-zA-Z0-9]{22}_[a-zA-Z0-9]{59}").unwrap(),
|
||||
severity: LeakSeverity::Critical,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// Stripe keys
|
||||
LeakPattern {
|
||||
name: "stripe_api_key".to_string(),
|
||||
regex: Regex::new(r"sk_(?:live|test)_[a-zA-Z0-9]{24,}").unwrap(),
|
||||
severity: LeakSeverity::Critical,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// NEAR AI session tokens
|
||||
LeakPattern {
|
||||
name: "nearai_session".to_string(),
|
||||
regex: Regex::new(r"sess_[a-zA-Z0-9]{32,}").unwrap(),
|
||||
severity: LeakSeverity::Critical,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// PEM private keys
|
||||
LeakPattern {
|
||||
name: "pem_private_key".to_string(),
|
||||
regex: Regex::new(r"-----BEGIN\s+(?:RSA\s+)?PRIVATE\s+KEY-----").unwrap(),
|
||||
severity: LeakSeverity::Critical,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// SSH private keys
|
||||
LeakPattern {
|
||||
name: "ssh_private_key".to_string(),
|
||||
regex: Regex::new(r"-----BEGIN\s+(?:OPENSSH|EC|DSA)\s+PRIVATE\s+KEY-----").unwrap(),
|
||||
severity: LeakSeverity::Critical,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// Google API keys
|
||||
LeakPattern {
|
||||
name: "google_api_key".to_string(),
|
||||
regex: Regex::new(r"AIza[0-9A-Za-z_-]{35}").unwrap(),
|
||||
severity: LeakSeverity::High,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// Slack tokens
|
||||
LeakPattern {
|
||||
name: "slack_token".to_string(),
|
||||
regex: Regex::new(r"xox[baprs]-[0-9a-zA-Z-]{10,}").unwrap(),
|
||||
severity: LeakSeverity::High,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// Twilio API keys
|
||||
LeakPattern {
|
||||
name: "twilio_api_key".to_string(),
|
||||
regex: Regex::new(r"SK[a-fA-F0-9]{32}").unwrap(),
|
||||
severity: LeakSeverity::High,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// SendGrid API keys
|
||||
LeakPattern {
|
||||
name: "sendgrid_api_key".to_string(),
|
||||
regex: Regex::new(r"SG\.[a-zA-Z0-9_-]{22}\.[a-zA-Z0-9_-]{43}").unwrap(),
|
||||
severity: LeakSeverity::High,
|
||||
action: LeakAction::Block,
|
||||
},
|
||||
// Bearer tokens (redact instead of block, might be intentional)
|
||||
LeakPattern {
|
||||
name: "bearer_token".to_string(),
|
||||
regex: Regex::new(r"Bearer\s+[a-zA-Z0-9_-]{20,}").unwrap(),
|
||||
severity: LeakSeverity::High,
|
||||
action: LeakAction::Redact,
|
||||
},
|
||||
// Authorization header with key
|
||||
LeakPattern {
|
||||
name: "auth_header".to_string(),
|
||||
regex: Regex::new(r"(?i)authorization:\s*[a-zA-Z]+\s+[a-zA-Z0-9_-]{20,}").unwrap(),
|
||||
severity: LeakSeverity::High,
|
||||
action: LeakAction::Redact,
|
||||
},
|
||||
// High entropy hex (potential secrets, warn only)
|
||||
// Uses word boundary since look-around isn't supported in the regex crate.
|
||||
// This catches standalone 64-char hex strings (like SHA256 hashes used as secrets).
|
||||
LeakPattern {
|
||||
name: "high_entropy_hex".to_string(),
|
||||
regex: Regex::new(r"\b[a-fA-F0-9]{64}\b").unwrap(),
|
||||
severity: LeakSeverity::Medium,
|
||||
action: LeakAction::Warn,
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::safety::leak_detector::{LeakDetector, LeakSeverity};
|
||||
|
||||
#[test]
|
||||
fn test_detect_openai_key() {
|
||||
let detector = LeakDetector::new();
|
||||
let content = "API key: sk-proj-abc123def456ghi789jkl012mno345pqrT3BlbkFJtest123";
|
||||
|
||||
let result = detector.scan(content);
|
||||
assert!(!result.is_clean());
|
||||
assert!(result.should_block);
|
||||
assert!(
|
||||
result
|
||||
.matches
|
||||
.iter()
|
||||
.any(|m| m.pattern_name == "openai_api_key")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_detect_github_token() {
|
||||
let detector = LeakDetector::new();
|
||||
let content = "token: ghp_xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx";
|
||||
|
||||
let result = detector.scan(content);
|
||||
assert!(!result.is_clean());
|
||||
assert!(
|
||||
result
|
||||
.matches
|
||||
.iter()
|
||||
.any(|m| m.pattern_name == "github_token")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_detect_aws_key() {
|
||||
let detector = LeakDetector::new();
|
||||
let content = "AWS_ACCESS_KEY_ID=AKIAIOSFODNN7EXAMPLE";
|
||||
|
||||
let result = detector.scan(content);
|
||||
assert!(!result.is_clean());
|
||||
assert!(
|
||||
result
|
||||
.matches
|
||||
.iter()
|
||||
.any(|m| m.pattern_name == "aws_access_key")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_detect_pem_key() {
|
||||
let detector = LeakDetector::new();
|
||||
let content = "-----BEGIN RSA PRIVATE KEY-----\nMIIEowIBAAKCAQEA...";
|
||||
|
||||
let result = detector.scan(content);
|
||||
assert!(!result.is_clean());
|
||||
assert!(
|
||||
result
|
||||
.matches
|
||||
.iter()
|
||||
.any(|m| m.pattern_name == "pem_private_key")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clean_content() {
|
||||
let detector = LeakDetector::new();
|
||||
let content = "Hello world! This is just regular text with no secrets.";
|
||||
|
||||
let result = detector.scan(content);
|
||||
assert!(result.is_clean());
|
||||
assert!(!result.should_block);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_redact_bearer_token() {
|
||||
let detector = LeakDetector::new();
|
||||
let content = "Authorization: Bearer eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9_longtokenvalue";
|
||||
|
||||
let result = detector.scan(content);
|
||||
assert!(!result.is_clean());
|
||||
assert!(!result.should_block); // Bearer is redact, not block
|
||||
|
||||
let redacted = result.redacted_content.unwrap();
|
||||
assert!(redacted.contains("[REDACTED]"));
|
||||
assert!(!redacted.contains("eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_scan_and_clean_blocks() {
|
||||
let detector = LeakDetector::new();
|
||||
let content = "sk-proj-test1234567890abcdefghij";
|
||||
|
||||
let result = detector.scan_and_clean(content);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_scan_and_clean_passes_clean() {
|
||||
let detector = LeakDetector::new();
|
||||
let content = "Just regular text";
|
||||
|
||||
let result = detector.scan_and_clean(content);
|
||||
assert!(result.is_ok());
|
||||
assert_eq!(result.unwrap(), content);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mask_secret() {
|
||||
use crate::safety::leak_detector::mask_secret;
|
||||
|
||||
assert_eq!(mask_secret("short"), "*****");
|
||||
assert_eq!(mask_secret("sk-test1234567890abcdef"), "sk-t********cdef");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_multiple_matches() {
|
||||
let detector = LeakDetector::new();
|
||||
let content = "Keys: AKIAIOSFODNN7EXAMPLE and ghp_xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx";
|
||||
|
||||
let result = detector.scan(content);
|
||||
assert_eq!(result.matches.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_severity_ordering() {
|
||||
assert!(LeakSeverity::Critical > LeakSeverity::High);
|
||||
assert!(LeakSeverity::High > LeakSeverity::Medium);
|
||||
assert!(LeakSeverity::Medium > LeakSeverity::Low);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_scan_http_request_clean() {
|
||||
let detector = LeakDetector::new();
|
||||
|
||||
let result = detector.scan_http_request(
|
||||
"https://api.example.com/data",
|
||||
&[("Content-Type".to_string(), "application/json".to_string())],
|
||||
Some(b"{\"query\": \"hello\"}"),
|
||||
);
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_scan_http_request_blocks_secret_in_url() {
|
||||
let detector = LeakDetector::new();
|
||||
|
||||
// Attempt to exfiltrate AWS key in URL
|
||||
let result = detector.scan_http_request(
|
||||
"https://evil.com/steal?key=AKIAIOSFODNN7EXAMPLE",
|
||||
&[],
|
||||
None,
|
||||
);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_scan_http_request_blocks_secret_in_header() {
|
||||
let detector = LeakDetector::new();
|
||||
|
||||
// Attempt to exfiltrate in custom header
|
||||
let result = detector.scan_http_request(
|
||||
"https://api.example.com/data",
|
||||
&[("X-Custom".to_string(), "ghp_xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx".to_string())],
|
||||
None,
|
||||
);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_scan_http_request_blocks_secret_in_body() {
|
||||
let detector = LeakDetector::new();
|
||||
|
||||
// Attempt to exfiltrate in request body
|
||||
let body = b"{\"stolen\": \"sk-proj-test1234567890abcdefghij\"}";
|
||||
let result = detector.scan_http_request(
|
||||
"https://api.example.com/webhook",
|
||||
&[],
|
||||
Some(body),
|
||||
);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
}
|
||||
@@ -5,11 +5,17 @@
|
||||
//! - Sanitizing tool outputs before they reach the LLM
|
||||
//! - Validating inputs before processing
|
||||
//! - Enforcing safety policies
|
||||
//! - Detecting secret leakage in outputs
|
||||
|
||||
mod leak_detector;
|
||||
mod policy;
|
||||
mod sanitizer;
|
||||
mod validator;
|
||||
|
||||
pub use leak_detector::{
|
||||
LeakAction, LeakDetectionError, LeakDetector, LeakMatch, LeakPattern, LeakScanResult,
|
||||
LeakSeverity,
|
||||
};
|
||||
pub use policy::{Policy, PolicyRule, Severity};
|
||||
pub use sanitizer::{InjectionWarning, SanitizedOutput, Sanitizer};
|
||||
pub use validator::{ValidationResult, Validator};
|
||||
|
||||
@@ -0,0 +1,250 @@
|
||||
//! Cryptographic operations for secret storage.
|
||||
//!
|
||||
//! Uses AES-256-GCM for authenticated encryption with per-secret key derivation.
|
||||
//!
|
||||
//! # Key Derivation
|
||||
//!
|
||||
//! ```text
|
||||
//! master_key (from env) ─┬─► HKDF-SHA256 ─► derived_key (per secret)
|
||||
//! │
|
||||
//! per-secret salt ───────┘
|
||||
//! ```
|
||||
//!
|
||||
//! Each secret has its own randomly-generated salt, so even if two secrets
|
||||
//! have the same plaintext, they'll have different ciphertexts.
|
||||
|
||||
use aes_gcm::{
|
||||
Aes256Gcm, KeyInit, Nonce,
|
||||
aead::{Aead, AeadCore, OsRng},
|
||||
};
|
||||
use hkdf::Hkdf;
|
||||
use secrecy::{ExposeSecret, SecretString};
|
||||
use sha2::Sha256;
|
||||
|
||||
use crate::secrets::types::{DecryptedSecret, SecretError};
|
||||
|
||||
/// Size of the AES-256 key in bytes.
|
||||
const KEY_SIZE: usize = 32;
|
||||
|
||||
/// Size of the GCM nonce in bytes.
|
||||
const NONCE_SIZE: usize = 12;
|
||||
|
||||
/// Size of the per-secret salt for key derivation.
|
||||
const SALT_SIZE: usize = 32;
|
||||
|
||||
/// Size of the GCM authentication tag.
|
||||
const TAG_SIZE: usize = 16;
|
||||
|
||||
/// Cryptographic operations for secrets.
|
||||
///
|
||||
/// Holds the master key and provides encrypt/decrypt operations.
|
||||
/// The master key is kept in secure memory and zeroed on drop.
|
||||
pub struct SecretsCrypto {
|
||||
master_key: SecretString,
|
||||
}
|
||||
|
||||
impl SecretsCrypto {
|
||||
/// Create a new crypto instance from a master key.
|
||||
///
|
||||
/// The master key should be at least 32 bytes of high-entropy data,
|
||||
/// typically loaded from an environment variable or secure vault.
|
||||
pub fn new(master_key: SecretString) -> Result<Self, SecretError> {
|
||||
// Validate master key length
|
||||
if master_key.expose_secret().len() < KEY_SIZE {
|
||||
return Err(SecretError::InvalidMasterKey);
|
||||
}
|
||||
Ok(Self { master_key })
|
||||
}
|
||||
|
||||
/// Generate a random salt for a new secret.
|
||||
pub fn generate_salt() -> Vec<u8> {
|
||||
let mut salt = vec![0u8; SALT_SIZE];
|
||||
rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut salt);
|
||||
salt
|
||||
}
|
||||
|
||||
/// Encrypt a secret value.
|
||||
///
|
||||
/// Returns (encrypted_value, salt) where:
|
||||
/// - encrypted_value = nonce || ciphertext || tag
|
||||
/// - salt = random bytes used for key derivation
|
||||
pub fn encrypt(&self, plaintext: &[u8]) -> Result<(Vec<u8>, Vec<u8>), SecretError> {
|
||||
let salt = Self::generate_salt();
|
||||
let derived_key = self.derive_key(&salt)?;
|
||||
|
||||
let cipher = Aes256Gcm::new_from_slice(&derived_key).map_err(|e| {
|
||||
SecretError::EncryptionFailed(format!("Failed to create cipher: {}", e))
|
||||
})?;
|
||||
|
||||
// Generate random nonce
|
||||
let nonce = Aes256Gcm::generate_nonce(&mut OsRng);
|
||||
|
||||
// Encrypt
|
||||
let ciphertext = cipher
|
||||
.encrypt(&nonce, plaintext)
|
||||
.map_err(|e| SecretError::EncryptionFailed(format!("Encryption failed: {}", e)))?;
|
||||
|
||||
// Combine: nonce || ciphertext (which includes tag)
|
||||
let mut encrypted = Vec::with_capacity(NONCE_SIZE + ciphertext.len());
|
||||
encrypted.extend_from_slice(&nonce);
|
||||
encrypted.extend_from_slice(&ciphertext);
|
||||
|
||||
Ok((encrypted, salt))
|
||||
}
|
||||
|
||||
/// Decrypt a secret value.
|
||||
///
|
||||
/// Takes the encrypted_value (nonce || ciphertext || tag) and the salt
|
||||
/// that was used during encryption.
|
||||
pub fn decrypt(
|
||||
&self,
|
||||
encrypted_value: &[u8],
|
||||
salt: &[u8],
|
||||
) -> Result<DecryptedSecret, SecretError> {
|
||||
if encrypted_value.len() < NONCE_SIZE + TAG_SIZE {
|
||||
return Err(SecretError::DecryptionFailed(
|
||||
"Encrypted value too short".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let derived_key = self.derive_key(salt)?;
|
||||
|
||||
let cipher = Aes256Gcm::new_from_slice(&derived_key).map_err(|e| {
|
||||
SecretError::DecryptionFailed(format!("Failed to create cipher: {}", e))
|
||||
})?;
|
||||
|
||||
// Split: nonce || ciphertext
|
||||
let (nonce_bytes, ciphertext) = encrypted_value.split_at(NONCE_SIZE);
|
||||
let nonce = Nonce::from_slice(nonce_bytes);
|
||||
|
||||
// Decrypt
|
||||
let plaintext = cipher
|
||||
.decrypt(nonce, ciphertext)
|
||||
.map_err(|e| SecretError::DecryptionFailed(format!("Decryption failed: {}", e)))?;
|
||||
|
||||
DecryptedSecret::from_bytes(plaintext)
|
||||
}
|
||||
|
||||
/// Derive a per-secret key using HKDF-SHA256.
|
||||
fn derive_key(&self, salt: &[u8]) -> Result<[u8; KEY_SIZE], SecretError> {
|
||||
let master_bytes = self.master_key.expose_secret().as_bytes();
|
||||
|
||||
// HKDF extract + expand
|
||||
let hk = Hkdf::<Sha256>::new(Some(salt), master_bytes);
|
||||
|
||||
let mut derived = [0u8; KEY_SIZE];
|
||||
hk.expand(b"near-agent-secrets-v1", &mut derived)
|
||||
.map_err(|_| SecretError::EncryptionFailed("HKDF expansion failed".to_string()))?;
|
||||
|
||||
Ok(derived)
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for SecretsCrypto {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("SecretsCrypto")
|
||||
.field("master_key", &"[REDACTED]")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use secrecy::SecretString;
|
||||
|
||||
use crate::secrets::crypto::SecretsCrypto;
|
||||
|
||||
fn test_crypto() -> SecretsCrypto {
|
||||
// 32-byte test key
|
||||
let key = "0123456789abcdef0123456789abcdef";
|
||||
SecretsCrypto::new(SecretString::from(key.to_string())).unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_encrypt_decrypt_roundtrip() {
|
||||
let crypto = test_crypto();
|
||||
let plaintext = b"my_super_secret_api_key_12345";
|
||||
|
||||
let (encrypted, salt) = crypto.encrypt(plaintext).unwrap();
|
||||
|
||||
// Encrypted should be larger than plaintext (nonce + tag)
|
||||
assert!(encrypted.len() > plaintext.len());
|
||||
|
||||
let decrypted = crypto.decrypt(&encrypted, &salt).unwrap();
|
||||
assert_eq!(decrypted.expose().as_bytes(), plaintext);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_different_salts_different_ciphertext() {
|
||||
let crypto = test_crypto();
|
||||
let plaintext = b"same_secret";
|
||||
|
||||
let (encrypted1, salt1) = crypto.encrypt(plaintext).unwrap();
|
||||
let (encrypted2, salt2) = crypto.encrypt(plaintext).unwrap();
|
||||
|
||||
// Same plaintext, different salts = different ciphertext
|
||||
assert_ne!(salt1, salt2);
|
||||
assert_ne!(encrypted1, encrypted2);
|
||||
|
||||
// But both decrypt to the same value
|
||||
let decrypted1 = crypto.decrypt(&encrypted1, &salt1).unwrap();
|
||||
let decrypted2 = crypto.decrypt(&encrypted2, &salt2).unwrap();
|
||||
assert_eq!(decrypted1.expose(), decrypted2.expose());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_wrong_salt_fails() {
|
||||
let crypto = test_crypto();
|
||||
let plaintext = b"secret";
|
||||
|
||||
let (encrypted, _salt) = crypto.encrypt(plaintext).unwrap();
|
||||
let wrong_salt = SecretsCrypto::generate_salt();
|
||||
|
||||
let result = crypto.decrypt(&encrypted, &wrong_salt);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tampered_ciphertext_fails() {
|
||||
let crypto = test_crypto();
|
||||
let plaintext = b"secret";
|
||||
|
||||
let (mut encrypted, salt) = crypto.encrypt(plaintext).unwrap();
|
||||
|
||||
// Tamper with the ciphertext
|
||||
if let Some(byte) = encrypted.last_mut() {
|
||||
*byte ^= 0xFF;
|
||||
}
|
||||
|
||||
let result = crypto.decrypt(&encrypted, &salt);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_master_key_too_short() {
|
||||
let short_key = "tooshort";
|
||||
let result = SecretsCrypto::new(SecretString::from(short_key.to_string()));
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_plaintext() {
|
||||
let crypto = test_crypto();
|
||||
let plaintext = b"";
|
||||
|
||||
let (encrypted, salt) = crypto.encrypt(plaintext).unwrap();
|
||||
let decrypted = crypto.decrypt(&encrypted, &salt).unwrap();
|
||||
assert!(decrypted.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_large_plaintext() {
|
||||
let crypto = test_crypto();
|
||||
// 1 MB of data
|
||||
let plaintext = vec![0x42u8; 1024 * 1024];
|
||||
|
||||
let (encrypted, salt) = crypto.encrypt(&plaintext).unwrap();
|
||||
let decrypted = crypto.decrypt(&encrypted, &salt).unwrap();
|
||||
assert_eq!(decrypted.expose().as_bytes(), plaintext.as_slice());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
//! Secrets management for secure credential storage and injection.
|
||||
//!
|
||||
//! This module provides:
|
||||
//! - AES-256-GCM encrypted secret storage
|
||||
//! - Per-secret key derivation (HKDF-SHA256)
|
||||
//! - PostgreSQL persistence
|
||||
//! - Access control for WASM tools
|
||||
//!
|
||||
//! # Security Model
|
||||
//!
|
||||
//! ```text
|
||||
//! ┌─────────────────────────────────────────────────────────────────────────────┐
|
||||
//! │ Secret Lifecycle │
|
||||
//! │ │
|
||||
//! │ User stores secret ──► Encrypt with AES-256-GCM ──► Store in PostgreSQL │
|
||||
//! │ (per-secret key via HKDF) │
|
||||
//! │ │
|
||||
//! │ WASM requests HTTP ──► Host checks allowlist ──► Decrypt secret ──► │
|
||||
//! │ & allowed_secrets (in memory only) │
|
||||
//! │ │ │
|
||||
//! │ ▼ │
|
||||
//! │ Inject into request ──► Execute HTTP call │
|
||||
//! │ (WASM never sees value) │
|
||||
//! │ │ │
|
||||
//! │ ▼ │
|
||||
//! │ Leak detector scans ──► Return response to WASM │
|
||||
//! │ response for secrets │
|
||||
//! └─────────────────────────────────────────────────────────────────────────────┘
|
||||
//! ```
|
||||
//!
|
||||
//! # Example
|
||||
//!
|
||||
//! ```ignore
|
||||
//! use near_agent::secrets::{SecretsStore, PostgresSecretsStore, SecretsCrypto, CreateSecretParams};
|
||||
//! use secrecy::SecretString;
|
||||
//!
|
||||
//! // Initialize crypto with master key from environment
|
||||
//! let master_key = SecretString::from(std::env::var("SECRETS_MASTER_KEY")?);
|
||||
//! let crypto = Arc::new(SecretsCrypto::new(master_key)?);
|
||||
//!
|
||||
//! // Create store
|
||||
//! let store = PostgresSecretsStore::new(pool, crypto);
|
||||
//!
|
||||
//! // Store a secret
|
||||
//! store.create("user_123", CreateSecretParams::new("openai_key", "sk-...")).await?;
|
||||
//!
|
||||
//! // Check if secret exists (WASM can call this)
|
||||
//! let exists = store.exists("user_123", "openai_key").await?;
|
||||
//!
|
||||
//! // Decrypt for injection (host boundary only)
|
||||
//! let decrypted = store.get_decrypted("user_123", "openai_key").await?;
|
||||
//! ```
|
||||
|
||||
mod crypto;
|
||||
mod store;
|
||||
mod types;
|
||||
|
||||
pub use crypto::SecretsCrypto;
|
||||
pub use store::{PostgresSecretsStore, SecretsStore};
|
||||
pub use types::{
|
||||
CreateSecretParams, CredentialLocation, CredentialMapping, DecryptedSecret, Secret,
|
||||
SecretError, SecretRef,
|
||||
};
|
||||
|
||||
#[cfg(test)]
|
||||
pub use store::testing::InMemorySecretsStore;
|
||||
@@ -0,0 +1,586 @@
|
||||
//! Secret storage with PostgreSQL persistence.
|
||||
//!
|
||||
//! Provides CRUD operations for encrypted secrets. The store handles:
|
||||
//! - Encryption/decryption via SecretsCrypto
|
||||
//! - Expiration checking
|
||||
//! - Usage tracking
|
||||
//! - Access control (which secrets a tool can use)
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use chrono::Utc;
|
||||
use deadpool_postgres::Pool;
|
||||
use secrecy::ExposeSecret;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::secrets::crypto::SecretsCrypto;
|
||||
use crate::secrets::types::{CreateSecretParams, DecryptedSecret, Secret, SecretError, SecretRef};
|
||||
|
||||
/// Trait for secret storage operations.
|
||||
///
|
||||
/// Allows for different implementations (PostgreSQL, in-memory for testing).
|
||||
#[async_trait]
|
||||
pub trait SecretsStore: Send + Sync {
|
||||
/// Store a new secret.
|
||||
async fn create(
|
||||
&self,
|
||||
user_id: &str,
|
||||
params: CreateSecretParams,
|
||||
) -> Result<Secret, SecretError>;
|
||||
|
||||
/// Get a secret by name (encrypted form).
|
||||
async fn get(&self, user_id: &str, name: &str) -> Result<Secret, SecretError>;
|
||||
|
||||
/// Get and decrypt a secret.
|
||||
async fn get_decrypted(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<DecryptedSecret, SecretError>;
|
||||
|
||||
/// Check if a secret exists.
|
||||
async fn exists(&self, user_id: &str, name: &str) -> Result<bool, SecretError>;
|
||||
|
||||
/// List all secret references for a user (no values).
|
||||
async fn list(&self, user_id: &str) -> Result<Vec<SecretRef>, SecretError>;
|
||||
|
||||
/// Delete a secret.
|
||||
async fn delete(&self, user_id: &str, name: &str) -> Result<bool, SecretError>;
|
||||
|
||||
/// Update secret usage tracking.
|
||||
async fn record_usage(&self, secret_id: Uuid) -> Result<(), SecretError>;
|
||||
|
||||
/// Check if a secret is accessible by a tool (based on allowed_secrets).
|
||||
async fn is_accessible(
|
||||
&self,
|
||||
user_id: &str,
|
||||
secret_name: &str,
|
||||
allowed_secrets: &[String],
|
||||
) -> Result<bool, SecretError>;
|
||||
}
|
||||
|
||||
/// PostgreSQL implementation of SecretsStore.
|
||||
pub struct PostgresSecretsStore {
|
||||
pool: Pool,
|
||||
crypto: Arc<SecretsCrypto>,
|
||||
}
|
||||
|
||||
impl PostgresSecretsStore {
|
||||
/// Create a new store with the given database pool and crypto instance.
|
||||
pub fn new(pool: Pool, crypto: Arc<SecretsCrypto>) -> Self {
|
||||
Self { pool, crypto }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SecretsStore for PostgresSecretsStore {
|
||||
async fn create(
|
||||
&self,
|
||||
user_id: &str,
|
||||
params: CreateSecretParams,
|
||||
) -> Result<Secret, SecretError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
// Encrypt the secret value
|
||||
let plaintext = params.value.expose_secret().as_bytes();
|
||||
let (encrypted_value, key_salt) = self.crypto.encrypt(plaintext)?;
|
||||
|
||||
let id = Uuid::new_v4();
|
||||
let now = Utc::now();
|
||||
|
||||
let row = client
|
||||
.query_one(
|
||||
r#"
|
||||
INSERT INTO secrets (id, user_id, name, encrypted_value, key_salt, provider, expires_at, created_at, updated_at)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $8)
|
||||
ON CONFLICT (user_id, name) DO UPDATE SET
|
||||
encrypted_value = EXCLUDED.encrypted_value,
|
||||
key_salt = EXCLUDED.key_salt,
|
||||
provider = EXCLUDED.provider,
|
||||
expires_at = EXCLUDED.expires_at,
|
||||
updated_at = NOW()
|
||||
RETURNING id, user_id, name, encrypted_value, key_salt, provider, expires_at,
|
||||
last_used_at, usage_count, created_at, updated_at
|
||||
"#,
|
||||
&[
|
||||
&id,
|
||||
&user_id,
|
||||
¶ms.name,
|
||||
&encrypted_value,
|
||||
&key_salt,
|
||||
¶ms.provider,
|
||||
¶ms.expires_at,
|
||||
&now,
|
||||
],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
Ok(row_to_secret(&row))
|
||||
}
|
||||
|
||||
async fn get(&self, user_id: &str, name: &str) -> Result<Secret, SecretError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
let row = client
|
||||
.query_opt(
|
||||
r#"
|
||||
SELECT id, user_id, name, encrypted_value, key_salt, provider, expires_at,
|
||||
last_used_at, usage_count, created_at, updated_at
|
||||
FROM secrets
|
||||
WHERE user_id = $1 AND name = $2
|
||||
"#,
|
||||
&[&user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
match row {
|
||||
Some(r) => {
|
||||
let secret = row_to_secret(&r);
|
||||
|
||||
// Check expiration
|
||||
if let Some(expires_at) = secret.expires_at {
|
||||
if expires_at < Utc::now() {
|
||||
return Err(SecretError::Expired);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(secret)
|
||||
}
|
||||
None => Err(SecretError::NotFound(name.to_string())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_decrypted(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<DecryptedSecret, SecretError> {
|
||||
let secret = self.get(user_id, name).await?;
|
||||
self.crypto
|
||||
.decrypt(&secret.encrypted_value, &secret.key_salt)
|
||||
}
|
||||
|
||||
async fn exists(&self, user_id: &str, name: &str) -> Result<bool, SecretError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
let row = client
|
||||
.query_one(
|
||||
"SELECT EXISTS(SELECT 1 FROM secrets WHERE user_id = $1 AND name = $2)",
|
||||
&[&user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
Ok(row.get(0))
|
||||
}
|
||||
|
||||
async fn list(&self, user_id: &str) -> Result<Vec<SecretRef>, SecretError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
let rows = client
|
||||
.query(
|
||||
"SELECT name, provider FROM secrets WHERE user_id = $1 ORDER BY name",
|
||||
&[&user_id],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
Ok(rows
|
||||
.into_iter()
|
||||
.map(|r| SecretRef {
|
||||
name: r.get(0),
|
||||
provider: r.get(1),
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn delete(&self, user_id: &str, name: &str) -> Result<bool, SecretError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
let result = client
|
||||
.execute(
|
||||
"DELETE FROM secrets WHERE user_id = $1 AND name = $2",
|
||||
&[&user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
Ok(result > 0)
|
||||
}
|
||||
|
||||
async fn record_usage(&self, secret_id: Uuid) -> Result<(), SecretError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
client
|
||||
.execute(
|
||||
r#"
|
||||
UPDATE secrets
|
||||
SET last_used_at = NOW(), usage_count = usage_count + 1
|
||||
WHERE id = $1
|
||||
"#,
|
||||
&[&secret_id],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| SecretError::Database(e.to_string()))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn is_accessible(
|
||||
&self,
|
||||
user_id: &str,
|
||||
secret_name: &str,
|
||||
allowed_secrets: &[String],
|
||||
) -> Result<bool, SecretError> {
|
||||
// First check if the secret exists
|
||||
if !self.exists(user_id, secret_name).await? {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
// Check if secret is in the allowed list
|
||||
// Supports glob patterns: "openai_*" matches "openai_api_key"
|
||||
for pattern in allowed_secrets {
|
||||
if pattern == secret_name {
|
||||
return Ok(true);
|
||||
}
|
||||
|
||||
// Simple glob: * matches any suffix
|
||||
if let Some(prefix) = pattern.strip_suffix('*') {
|
||||
if secret_name.starts_with(prefix) {
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(false)
|
||||
}
|
||||
}
|
||||
|
||||
fn row_to_secret(row: &tokio_postgres::Row) -> Secret {
|
||||
Secret {
|
||||
id: row.get("id"),
|
||||
user_id: row.get("user_id"),
|
||||
name: row.get("name"),
|
||||
encrypted_value: row.get("encrypted_value"),
|
||||
key_salt: row.get("key_salt"),
|
||||
provider: row.get("provider"),
|
||||
expires_at: row.get("expires_at"),
|
||||
last_used_at: row.get("last_used_at"),
|
||||
usage_count: row.get("usage_count"),
|
||||
created_at: row.get("created_at"),
|
||||
updated_at: row.get("updated_at"),
|
||||
}
|
||||
}
|
||||
|
||||
/// In-memory implementation for testing.
|
||||
#[cfg(test)]
|
||||
pub mod testing {
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use chrono::Utc;
|
||||
use secrecy::ExposeSecret;
|
||||
use tokio::sync::RwLock;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::secrets::crypto::SecretsCrypto;
|
||||
use crate::secrets::store::SecretsStore;
|
||||
use crate::secrets::types::{
|
||||
CreateSecretParams, DecryptedSecret, Secret, SecretError, SecretRef,
|
||||
};
|
||||
|
||||
pub struct InMemorySecretsStore {
|
||||
secrets: RwLock<HashMap<(String, String), Secret>>,
|
||||
crypto: Arc<SecretsCrypto>,
|
||||
}
|
||||
|
||||
impl InMemorySecretsStore {
|
||||
pub fn new(crypto: Arc<SecretsCrypto>) -> Self {
|
||||
Self {
|
||||
secrets: RwLock::new(HashMap::new()),
|
||||
crypto,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SecretsStore for InMemorySecretsStore {
|
||||
async fn create(
|
||||
&self,
|
||||
user_id: &str,
|
||||
params: CreateSecretParams,
|
||||
) -> Result<Secret, SecretError> {
|
||||
let plaintext = params.value.expose_secret().as_bytes();
|
||||
let (encrypted_value, key_salt) = self.crypto.encrypt(plaintext)?;
|
||||
|
||||
let now = Utc::now();
|
||||
let secret = Secret {
|
||||
id: Uuid::new_v4(),
|
||||
user_id: user_id.to_string(),
|
||||
name: params.name.clone(),
|
||||
encrypted_value,
|
||||
key_salt,
|
||||
provider: params.provider,
|
||||
expires_at: params.expires_at,
|
||||
last_used_at: None,
|
||||
usage_count: 0,
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
};
|
||||
|
||||
self.secrets
|
||||
.write()
|
||||
.await
|
||||
.insert((user_id.to_string(), params.name), secret.clone());
|
||||
Ok(secret)
|
||||
}
|
||||
|
||||
async fn get(&self, user_id: &str, name: &str) -> Result<Secret, SecretError> {
|
||||
self.secrets
|
||||
.read()
|
||||
.await
|
||||
.get(&(user_id.to_string(), name.to_string()))
|
||||
.cloned()
|
||||
.ok_or_else(|| SecretError::NotFound(name.to_string()))
|
||||
}
|
||||
|
||||
async fn get_decrypted(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<DecryptedSecret, SecretError> {
|
||||
let secret = self.get(user_id, name).await?;
|
||||
self.crypto
|
||||
.decrypt(&secret.encrypted_value, &secret.key_salt)
|
||||
}
|
||||
|
||||
async fn exists(&self, user_id: &str, name: &str) -> Result<bool, SecretError> {
|
||||
Ok(self
|
||||
.secrets
|
||||
.read()
|
||||
.await
|
||||
.contains_key(&(user_id.to_string(), name.to_string())))
|
||||
}
|
||||
|
||||
async fn list(&self, user_id: &str) -> Result<Vec<SecretRef>, SecretError> {
|
||||
Ok(self
|
||||
.secrets
|
||||
.read()
|
||||
.await
|
||||
.iter()
|
||||
.filter(|((uid, _), _)| uid == user_id)
|
||||
.map(|((_, _), s)| SecretRef {
|
||||
name: s.name.clone(),
|
||||
provider: s.provider.clone(),
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn delete(&self, user_id: &str, name: &str) -> Result<bool, SecretError> {
|
||||
Ok(self
|
||||
.secrets
|
||||
.write()
|
||||
.await
|
||||
.remove(&(user_id.to_string(), name.to_string()))
|
||||
.is_some())
|
||||
}
|
||||
|
||||
async fn record_usage(&self, _secret_id: Uuid) -> Result<(), SecretError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn is_accessible(
|
||||
&self,
|
||||
user_id: &str,
|
||||
secret_name: &str,
|
||||
allowed_secrets: &[String],
|
||||
) -> Result<bool, SecretError> {
|
||||
if !self.exists(user_id, secret_name).await? {
|
||||
return Ok(false);
|
||||
}
|
||||
for pattern in allowed_secrets {
|
||||
if pattern == secret_name {
|
||||
return Ok(true);
|
||||
}
|
||||
if let Some(prefix) = pattern.strip_suffix('*') {
|
||||
if secret_name.starts_with(prefix) {
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(false)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use secrecy::SecretString;
|
||||
|
||||
use crate::secrets::crypto::SecretsCrypto;
|
||||
use crate::secrets::store::SecretsStore;
|
||||
use crate::secrets::store::testing::InMemorySecretsStore;
|
||||
use crate::secrets::types::CreateSecretParams;
|
||||
|
||||
fn test_store() -> InMemorySecretsStore {
|
||||
let key = "0123456789abcdef0123456789abcdef";
|
||||
let crypto = Arc::new(SecretsCrypto::new(SecretString::from(key.to_string())).unwrap());
|
||||
InMemorySecretsStore::new(crypto)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_create_and_get() {
|
||||
let store = test_store();
|
||||
let params = CreateSecretParams::new("api_key", "sk-test-12345");
|
||||
|
||||
store.create("user1", params).await.unwrap();
|
||||
|
||||
let decrypted = store.get_decrypted("user1", "api_key").await.unwrap();
|
||||
assert_eq!(decrypted.expose(), "sk-test-12345");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_exists() {
|
||||
let store = test_store();
|
||||
let params = CreateSecretParams::new("my_secret", "value");
|
||||
|
||||
assert!(!store.exists("user1", "my_secret").await.unwrap());
|
||||
store.create("user1", params).await.unwrap();
|
||||
assert!(store.exists("user1", "my_secret").await.unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_delete() {
|
||||
let store = test_store();
|
||||
let params = CreateSecretParams::new("to_delete", "value");
|
||||
|
||||
store.create("user1", params).await.unwrap();
|
||||
assert!(store.exists("user1", "to_delete").await.unwrap());
|
||||
|
||||
store.delete("user1", "to_delete").await.unwrap();
|
||||
assert!(!store.exists("user1", "to_delete").await.unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_list() {
|
||||
let store = test_store();
|
||||
|
||||
store
|
||||
.create("user1", CreateSecretParams::new("key1", "v1"))
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.create(
|
||||
"user1",
|
||||
CreateSecretParams::new("key2", "v2").with_provider("openai"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.create("user2", CreateSecretParams::new("key3", "v3"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let list = store.list("user1").await.unwrap();
|
||||
assert_eq!(list.len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_is_accessible() {
|
||||
let store = test_store();
|
||||
store
|
||||
.create("user1", CreateSecretParams::new("openai_key", "sk-test"))
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.create("user1", CreateSecretParams::new("stripe_key", "sk-live"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Exact match
|
||||
let allowed = vec!["openai_key".to_string()];
|
||||
assert!(
|
||||
store
|
||||
.is_accessible("user1", "openai_key", &allowed)
|
||||
.await
|
||||
.unwrap()
|
||||
);
|
||||
assert!(
|
||||
!store
|
||||
.is_accessible("user1", "stripe_key", &allowed)
|
||||
.await
|
||||
.unwrap()
|
||||
);
|
||||
|
||||
// Glob pattern
|
||||
let allowed = vec!["openai_*".to_string()];
|
||||
assert!(
|
||||
store
|
||||
.is_accessible("user1", "openai_key", &allowed)
|
||||
.await
|
||||
.unwrap()
|
||||
);
|
||||
assert!(
|
||||
!store
|
||||
.is_accessible("user1", "stripe_key", &allowed)
|
||||
.await
|
||||
.unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_user_isolation() {
|
||||
let store = test_store();
|
||||
|
||||
store
|
||||
.create(
|
||||
"user1",
|
||||
CreateSecretParams::new("shared_name", "user1_value"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.create(
|
||||
"user2",
|
||||
CreateSecretParams::new("shared_name", "user2_value"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let v1 = store.get_decrypted("user1", "shared_name").await.unwrap();
|
||||
let v2 = store.get_decrypted("user2", "shared_name").await.unwrap();
|
||||
|
||||
assert_eq!(v1.expose(), "user1_value");
|
||||
assert_eq!(v2.expose(), "user2_value");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,281 @@
|
||||
//! Secret types for credential management.
|
||||
//!
|
||||
//! WASM tools NEVER see plaintext secrets. This module provides types
|
||||
//! for secure storage and reference without exposing actual values.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use chrono::{DateTime, Utc};
|
||||
use secrecy::{ExposeSecret, SecretString};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use uuid::Uuid;
|
||||
|
||||
/// A stored secret with encrypted value.
|
||||
///
|
||||
/// The plaintext is never stored; only the encrypted form exists in the database.
|
||||
#[derive(Clone)]
|
||||
pub struct Secret {
|
||||
pub id: Uuid,
|
||||
pub user_id: String,
|
||||
pub name: String,
|
||||
/// AES-256-GCM encrypted value (nonce || ciphertext || tag).
|
||||
pub encrypted_value: Vec<u8>,
|
||||
/// Per-secret salt for key derivation.
|
||||
pub key_salt: Vec<u8>,
|
||||
/// Optional provider hint (e.g., "openai", "stripe").
|
||||
pub provider: Option<String>,
|
||||
/// When this secret expires (None = never).
|
||||
pub expires_at: Option<DateTime<Utc>>,
|
||||
/// Last time this secret was used for injection.
|
||||
pub last_used_at: Option<DateTime<Utc>>,
|
||||
/// Total number of times this secret has been used.
|
||||
pub usage_count: i64,
|
||||
pub created_at: DateTime<Utc>,
|
||||
pub updated_at: DateTime<Utc>,
|
||||
}
|
||||
|
||||
impl fmt::Debug for Secret {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("Secret")
|
||||
.field("id", &self.id)
|
||||
.field("user_id", &self.user_id)
|
||||
.field("name", &self.name)
|
||||
.field("encrypted_value", &"[REDACTED]")
|
||||
.field("key_salt", &"[REDACTED]")
|
||||
.field("provider", &self.provider)
|
||||
.field("expires_at", &self.expires_at)
|
||||
.field("last_used_at", &self.last_used_at)
|
||||
.field("usage_count", &self.usage_count)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// A reference to a secret by name, without exposing the value.
|
||||
///
|
||||
/// WASM tools receive these references and can check if secrets exist,
|
||||
/// but they cannot read the actual values.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SecretRef {
|
||||
pub name: String,
|
||||
pub provider: Option<String>,
|
||||
}
|
||||
|
||||
impl SecretRef {
|
||||
pub fn new(name: impl Into<String>) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
provider: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_provider(mut self, provider: impl Into<String>) -> Self {
|
||||
self.provider = Some(provider.into());
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// A decrypted secret value, held in secure memory.
|
||||
///
|
||||
/// This type:
|
||||
/// - Zeros memory on drop
|
||||
/// - Never appears in Debug output
|
||||
/// - Only exists briefly during credential injection
|
||||
pub struct DecryptedSecret {
|
||||
value: SecretString,
|
||||
}
|
||||
|
||||
impl DecryptedSecret {
|
||||
/// Create a new decrypted secret from raw bytes.
|
||||
///
|
||||
/// The bytes are converted to a UTF-8 string. For binary secrets,
|
||||
/// consider base64 encoding before storage.
|
||||
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self, SecretError> {
|
||||
// Convert to string, then wrap in SecretString
|
||||
let s = String::from_utf8(bytes).map_err(|_| SecretError::InvalidUtf8)?;
|
||||
Ok(Self {
|
||||
value: SecretString::from(s),
|
||||
})
|
||||
}
|
||||
|
||||
/// Expose the secret value for injection.
|
||||
///
|
||||
/// This is the ONLY way to access the plaintext. Use sparingly
|
||||
/// and ensure the exposed value isn't logged or persisted.
|
||||
pub fn expose(&self) -> &str {
|
||||
self.value.expose_secret()
|
||||
}
|
||||
|
||||
/// Get the length of the secret without exposing it.
|
||||
pub fn len(&self) -> usize {
|
||||
self.value.expose_secret().len()
|
||||
}
|
||||
|
||||
/// Check if the secret is empty.
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.len() == 0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for DecryptedSecret {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(f, "DecryptedSecret([REDACTED, {} bytes])", self.len())
|
||||
}
|
||||
}
|
||||
|
||||
impl Clone for DecryptedSecret {
|
||||
fn clone(&self) -> Self {
|
||||
Self {
|
||||
value: SecretString::from(self.value.expose_secret().to_string()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Errors that can occur during secret operations.
|
||||
#[derive(Debug, Clone, thiserror::Error)]
|
||||
pub enum SecretError {
|
||||
#[error("Secret not found: {0}")]
|
||||
NotFound(String),
|
||||
|
||||
#[error("Secret has expired")]
|
||||
Expired,
|
||||
|
||||
#[error("Decryption failed: {0}")]
|
||||
DecryptionFailed(String),
|
||||
|
||||
#[error("Encryption failed: {0}")]
|
||||
EncryptionFailed(String),
|
||||
|
||||
#[error("Invalid master key")]
|
||||
InvalidMasterKey,
|
||||
|
||||
#[error("Secret value is not valid UTF-8")]
|
||||
InvalidUtf8,
|
||||
|
||||
#[error("Database error: {0}")]
|
||||
Database(String),
|
||||
|
||||
#[error("Secret access denied for tool")]
|
||||
AccessDenied,
|
||||
}
|
||||
|
||||
/// Parameters for creating a new secret.
|
||||
#[derive(Debug)]
|
||||
pub struct CreateSecretParams {
|
||||
pub name: String,
|
||||
pub value: SecretString,
|
||||
pub provider: Option<String>,
|
||||
pub expires_at: Option<DateTime<Utc>>,
|
||||
}
|
||||
|
||||
impl CreateSecretParams {
|
||||
pub fn new(name: impl Into<String>, value: impl Into<String>) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
value: SecretString::from(value.into()),
|
||||
provider: None,
|
||||
expires_at: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_provider(mut self, provider: impl Into<String>) -> Self {
|
||||
self.provider = Some(provider.into());
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_expiry(mut self, expires_at: DateTime<Utc>) -> Self {
|
||||
self.expires_at = Some(expires_at);
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// Where a credential should be injected in an HTTP request.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub enum CredentialLocation {
|
||||
/// Inject as Authorization header (e.g., "Bearer {secret}")
|
||||
AuthorizationBearer,
|
||||
/// Inject as Authorization header with Basic auth
|
||||
AuthorizationBasic { username: String },
|
||||
/// Inject as a custom header
|
||||
Header {
|
||||
name: String,
|
||||
prefix: Option<String>,
|
||||
},
|
||||
/// Inject as a query parameter
|
||||
QueryParam { name: String },
|
||||
}
|
||||
|
||||
impl Default for CredentialLocation {
|
||||
fn default() -> Self {
|
||||
Self::AuthorizationBearer
|
||||
}
|
||||
}
|
||||
|
||||
/// Mapping from a secret name to where it should be injected.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CredentialMapping {
|
||||
/// Name of the secret to use.
|
||||
pub secret_name: String,
|
||||
/// Where to inject the credential.
|
||||
pub location: CredentialLocation,
|
||||
/// Host patterns this credential applies to (glob syntax).
|
||||
pub host_patterns: Vec<String>,
|
||||
}
|
||||
|
||||
impl CredentialMapping {
|
||||
pub fn bearer(secret_name: impl Into<String>, host_pattern: impl Into<String>) -> Self {
|
||||
Self {
|
||||
secret_name: secret_name.into(),
|
||||
location: CredentialLocation::AuthorizationBearer,
|
||||
host_patterns: vec![host_pattern.into()],
|
||||
}
|
||||
}
|
||||
|
||||
pub fn header(
|
||||
secret_name: impl Into<String>,
|
||||
header_name: impl Into<String>,
|
||||
host_pattern: impl Into<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
secret_name: secret_name.into(),
|
||||
location: CredentialLocation::Header {
|
||||
name: header_name.into(),
|
||||
prefix: None,
|
||||
},
|
||||
host_patterns: vec![host_pattern.into()],
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::secrets::types::{CreateSecretParams, DecryptedSecret, SecretRef};
|
||||
|
||||
#[test]
|
||||
fn test_secret_ref_creation() {
|
||||
let r = SecretRef::new("my_api_key").with_provider("openai");
|
||||
assert_eq!(r.name, "my_api_key");
|
||||
assert_eq!(r.provider, Some("openai".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_decrypted_secret_redaction() {
|
||||
let secret = DecryptedSecret::from_bytes(b"super_secret_value".to_vec()).unwrap();
|
||||
let debug_str = format!("{:?}", secret);
|
||||
assert!(!debug_str.contains("super_secret_value"));
|
||||
assert!(debug_str.contains("REDACTED"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_decrypted_secret_expose() {
|
||||
let secret = DecryptedSecret::from_bytes(b"test_value".to_vec()).unwrap();
|
||||
assert_eq!(secret.expose(), "test_value");
|
||||
assert_eq!(secret.len(), 10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_params() {
|
||||
let params = CreateSecretParams::new("key", "value").with_provider("stripe");
|
||||
assert_eq!(params.name, "key");
|
||||
assert_eq!(params.provider, Some("stripe".to_string()));
|
||||
}
|
||||
}
|
||||
+73
-1
@@ -9,7 +9,8 @@ use crate::llm::ToolDefinition;
|
||||
use crate::tools::builtin::{EchoTool, HttpTool, JsonTool, TimeTool};
|
||||
use crate::tools::tool::Tool;
|
||||
use crate::tools::wasm::{
|
||||
Capabilities, ResourceLimits, WasmError, WasmToolRuntime, WasmToolWrapper,
|
||||
Capabilities, ResourceLimits, WasmError, WasmStorageError, WasmToolRuntime, WasmToolStore,
|
||||
WasmToolWrapper,
|
||||
};
|
||||
|
||||
/// Registry of available tools.
|
||||
@@ -152,6 +153,77 @@ impl ToolRegistry {
|
||||
tracing::info!(name = reg.name, "Registered WASM tool");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Register a WASM tool from database storage.
|
||||
///
|
||||
/// Loads the WASM binary with integrity verification and configures capabilities.
|
||||
///
|
||||
/// # Example
|
||||
///
|
||||
/// ```ignore
|
||||
/// let store = PostgresWasmToolStore::new(pool);
|
||||
/// let runtime = Arc::new(WasmToolRuntime::new(WasmRuntimeConfig::default())?);
|
||||
///
|
||||
/// registry.register_wasm_from_storage(
|
||||
/// &store,
|
||||
/// &runtime,
|
||||
/// "user_123",
|
||||
/// "my_tool",
|
||||
/// ).await?;
|
||||
/// ```
|
||||
pub async fn register_wasm_from_storage(
|
||||
&self,
|
||||
store: &dyn WasmToolStore,
|
||||
runtime: &Arc<WasmToolRuntime>,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<(), WasmRegistrationError> {
|
||||
// Load tool with integrity verification
|
||||
let tool_with_binary = store
|
||||
.get_with_binary(user_id, name)
|
||||
.await
|
||||
.map_err(WasmRegistrationError::Storage)?;
|
||||
|
||||
// Load capabilities
|
||||
let stored_caps = store
|
||||
.get_capabilities(tool_with_binary.tool.id)
|
||||
.await
|
||||
.map_err(WasmRegistrationError::Storage)?;
|
||||
|
||||
let capabilities = stored_caps.map(|c| c.to_capabilities()).unwrap_or_default();
|
||||
|
||||
// Register the tool
|
||||
self.register_wasm(WasmToolRegistration {
|
||||
name: &tool_with_binary.tool.name,
|
||||
wasm_bytes: &tool_with_binary.wasm_binary,
|
||||
runtime,
|
||||
capabilities,
|
||||
limits: None,
|
||||
description: Some(&tool_with_binary.tool.description),
|
||||
schema: Some(tool_with_binary.tool.parameters_schema.clone()),
|
||||
})
|
||||
.await
|
||||
.map_err(WasmRegistrationError::Wasm)?;
|
||||
|
||||
tracing::info!(
|
||||
name = tool_with_binary.tool.name,
|
||||
user_id = user_id,
|
||||
trust_level = %tool_with_binary.tool.trust_level,
|
||||
"Registered WASM tool from storage"
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Error when registering a WASM tool from storage.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum WasmRegistrationError {
|
||||
#[error("Storage error: {0}")]
|
||||
Storage(#[from] WasmStorageError),
|
||||
|
||||
#[error("WASM error: {0}")]
|
||||
Wasm(#[from] WasmError),
|
||||
}
|
||||
|
||||
/// Configuration for registering a WASM tool.
|
||||
|
||||
@@ -0,0 +1,357 @@
|
||||
//! HTTP endpoint allowlist validation.
|
||||
//!
|
||||
//! Validates that HTTP requests from WASM tools only go to allowed endpoints.
|
||||
//! This is the first line of defense against unauthorized API access.
|
||||
//!
|
||||
//! # Validation Flow
|
||||
//!
|
||||
//! ```text
|
||||
//! WASM HTTP request ──► Parse URL ──► Check allowlist ──► Allow/Deny
|
||||
//! │ │
|
||||
//! │ ├─► Host match?
|
||||
//! │ ├─► Path prefix match?
|
||||
//! │ └─► Method allowed?
|
||||
//! │
|
||||
//! └─► Validate URL format
|
||||
//! ```
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use crate::tools::wasm::capabilities::EndpointPattern;
|
||||
|
||||
/// Result of allowlist validation.
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum AllowlistResult {
|
||||
/// Request is allowed.
|
||||
Allowed,
|
||||
/// Request is denied with reason.
|
||||
Denied(DenyReason),
|
||||
}
|
||||
|
||||
impl AllowlistResult {
|
||||
pub fn is_allowed(&self) -> bool {
|
||||
matches!(self, AllowlistResult::Allowed)
|
||||
}
|
||||
}
|
||||
|
||||
/// Reason why a request was denied.
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum DenyReason {
|
||||
/// URL could not be parsed.
|
||||
InvalidUrl(String),
|
||||
/// Host is not in the allowlist.
|
||||
HostNotAllowed(String),
|
||||
/// Path does not match any allowed prefix.
|
||||
PathNotAllowed { host: String, path: String },
|
||||
/// HTTP method is not allowed for this endpoint.
|
||||
MethodNotAllowed { method: String, host: String },
|
||||
/// Allowlist is empty (no endpoints configured).
|
||||
EmptyAllowlist,
|
||||
/// URL scheme is not HTTPS.
|
||||
InsecureScheme(String),
|
||||
}
|
||||
|
||||
impl fmt::Display for DenyReason {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
DenyReason::InvalidUrl(url) => write!(f, "Invalid URL: {}", url),
|
||||
DenyReason::HostNotAllowed(host) => write!(f, "Host not in allowlist: {}", host),
|
||||
DenyReason::PathNotAllowed { host, path } => {
|
||||
write!(f, "Path not allowed for host {}: {}", host, path)
|
||||
}
|
||||
DenyReason::MethodNotAllowed { method, host } => {
|
||||
write!(f, "Method {} not allowed for host {}", method, host)
|
||||
}
|
||||
DenyReason::EmptyAllowlist => write!(f, "No endpoints in allowlist"),
|
||||
DenyReason::InsecureScheme(scheme) => {
|
||||
write!(f, "Insecure scheme: {} (only HTTPS allowed)", scheme)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Validates HTTP requests against an allowlist.
|
||||
pub struct AllowlistValidator {
|
||||
patterns: Vec<EndpointPattern>,
|
||||
/// Whether to require HTTPS (default: true).
|
||||
require_https: bool,
|
||||
}
|
||||
|
||||
impl AllowlistValidator {
|
||||
/// Create a new validator with the given patterns.
|
||||
pub fn new(patterns: Vec<EndpointPattern>) -> Self {
|
||||
Self {
|
||||
patterns,
|
||||
require_https: true,
|
||||
}
|
||||
}
|
||||
|
||||
/// Allow HTTP (insecure) requests. Use with caution.
|
||||
pub fn allow_http(mut self) -> Self {
|
||||
self.require_https = false;
|
||||
self
|
||||
}
|
||||
|
||||
/// Check if a request is allowed.
|
||||
pub fn validate(&self, url: &str, method: &str) -> AllowlistResult {
|
||||
// Check for empty allowlist
|
||||
if self.patterns.is_empty() {
|
||||
return AllowlistResult::Denied(DenyReason::EmptyAllowlist);
|
||||
}
|
||||
|
||||
// Parse the URL
|
||||
let parsed = match parse_url(url) {
|
||||
Ok(p) => p,
|
||||
Err(e) => return AllowlistResult::Denied(DenyReason::InvalidUrl(e)),
|
||||
};
|
||||
|
||||
// Check HTTPS requirement
|
||||
if self.require_https && parsed.scheme != "https" {
|
||||
return AllowlistResult::Denied(DenyReason::InsecureScheme(parsed.scheme.clone()));
|
||||
}
|
||||
|
||||
// Find a matching pattern
|
||||
for pattern in &self.patterns {
|
||||
if pattern.matches(&parsed.host, &parsed.path, method) {
|
||||
return AllowlistResult::Allowed;
|
||||
}
|
||||
}
|
||||
|
||||
// No pattern matched, figure out why for better error messages
|
||||
let host_matches: Vec<_> = self
|
||||
.patterns
|
||||
.iter()
|
||||
.filter(|p| p.host_matches(&parsed.host))
|
||||
.collect();
|
||||
|
||||
if host_matches.is_empty() {
|
||||
AllowlistResult::Denied(DenyReason::HostNotAllowed(parsed.host))
|
||||
} else {
|
||||
// Host matches but path/method doesn't
|
||||
let path_matches: Vec<_> = host_matches
|
||||
.iter()
|
||||
.filter(|p| {
|
||||
p.path_prefix.is_none()
|
||||
|| parsed
|
||||
.path
|
||||
.starts_with(p.path_prefix.as_deref().unwrap_or(""))
|
||||
})
|
||||
.collect();
|
||||
|
||||
if path_matches.is_empty() {
|
||||
AllowlistResult::Denied(DenyReason::PathNotAllowed {
|
||||
host: parsed.host,
|
||||
path: parsed.path,
|
||||
})
|
||||
} else {
|
||||
AllowlistResult::Denied(DenyReason::MethodNotAllowed {
|
||||
method: method.to_string(),
|
||||
host: parsed.host,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if any pattern would allow this host.
|
||||
pub fn host_allowed(&self, host: &str) -> bool {
|
||||
self.patterns.iter().any(|p| p.host_matches(host))
|
||||
}
|
||||
|
||||
/// Get all allowed hosts (for debugging/logging).
|
||||
pub fn allowed_hosts(&self) -> Vec<&str> {
|
||||
self.patterns.iter().map(|p| p.host.as_str()).collect()
|
||||
}
|
||||
}
|
||||
|
||||
/// Parsed URL components.
|
||||
struct ParsedUrl {
|
||||
scheme: String,
|
||||
host: String,
|
||||
path: String,
|
||||
}
|
||||
|
||||
/// Simple URL parser (avoids pulling in a full URL crate).
|
||||
fn parse_url(url: &str) -> Result<ParsedUrl, String> {
|
||||
// Find scheme
|
||||
let (scheme, rest) = url
|
||||
.split_once("://")
|
||||
.ok_or_else(|| "Missing scheme (expected http:// or https://)".to_string())?;
|
||||
|
||||
let scheme = scheme.to_lowercase();
|
||||
if scheme != "http" && scheme != "https" {
|
||||
return Err(format!("Unsupported scheme: {}", scheme));
|
||||
}
|
||||
|
||||
// Split host from path
|
||||
let (host_and_port, path) = match rest.find('/') {
|
||||
Some(idx) => (&rest[..idx], &rest[idx..]),
|
||||
None => (rest, "/"),
|
||||
};
|
||||
|
||||
// Remove port from host
|
||||
let host = match host_and_port.rfind(':') {
|
||||
Some(idx) => {
|
||||
// Make sure this isn't an IPv6 address
|
||||
if host_and_port.starts_with('[') {
|
||||
// IPv6: [::1]:8080 or [::1]
|
||||
if let Some(bracket_idx) = host_and_port.find(']') {
|
||||
// Extract the IPv6 address without brackets
|
||||
&host_and_port[1..bracket_idx]
|
||||
} else {
|
||||
return Err("Invalid IPv6 address".to_string());
|
||||
}
|
||||
} else {
|
||||
&host_and_port[..idx]
|
||||
}
|
||||
}
|
||||
None => host_and_port,
|
||||
};
|
||||
|
||||
// Validate host
|
||||
if host.is_empty() {
|
||||
return Err("Empty host".to_string());
|
||||
}
|
||||
|
||||
Ok(ParsedUrl {
|
||||
scheme,
|
||||
host: host.to_lowercase(),
|
||||
path: path.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tools::wasm::allowlist::{AllowlistValidator, DenyReason};
|
||||
use crate::tools::wasm::capabilities::EndpointPattern;
|
||||
|
||||
fn validator_with_patterns() -> AllowlistValidator {
|
||||
AllowlistValidator::new(vec![
|
||||
EndpointPattern::host("api.openai.com").with_path_prefix("/v1/"),
|
||||
EndpointPattern::host("api.anthropic.com")
|
||||
.with_path_prefix("/v1/messages")
|
||||
.with_methods(vec!["POST".to_string()]),
|
||||
EndpointPattern::host("*.example.com"),
|
||||
])
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_allowed_request() {
|
||||
let validator = validator_with_patterns();
|
||||
|
||||
let result = validator.validate("https://api.openai.com/v1/chat/completions", "POST");
|
||||
assert!(result.is_allowed());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_denied_wrong_host() {
|
||||
let validator = validator_with_patterns();
|
||||
|
||||
let result = validator.validate("https://evil.com/steal/data", "GET");
|
||||
assert!(!result.is_allowed());
|
||||
|
||||
if let super::AllowlistResult::Denied(reason) = result {
|
||||
assert!(matches!(reason, DenyReason::HostNotAllowed(_)));
|
||||
} else {
|
||||
panic!("Expected denied");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_denied_wrong_path() {
|
||||
let validator = validator_with_patterns();
|
||||
|
||||
let result = validator.validate("https://api.openai.com/v2/different", "GET");
|
||||
assert!(!result.is_allowed());
|
||||
|
||||
if let super::AllowlistResult::Denied(reason) = result {
|
||||
assert!(matches!(reason, DenyReason::PathNotAllowed { .. }));
|
||||
} else {
|
||||
panic!("Expected denied");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_denied_wrong_method() {
|
||||
let validator = validator_with_patterns();
|
||||
|
||||
// Anthropic endpoint only allows POST
|
||||
let result = validator.validate("https://api.anthropic.com/v1/messages", "GET");
|
||||
assert!(!result.is_allowed());
|
||||
|
||||
if let super::AllowlistResult::Denied(reason) = result {
|
||||
assert!(matches!(reason, DenyReason::MethodNotAllowed { .. }));
|
||||
} else {
|
||||
panic!("Expected denied");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_wildcard_host() {
|
||||
let validator = validator_with_patterns();
|
||||
|
||||
let result = validator.validate("https://api.example.com/anything", "GET");
|
||||
assert!(result.is_allowed());
|
||||
|
||||
let result = validator.validate("https://sub.api.example.com/anything", "GET");
|
||||
assert!(result.is_allowed());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_require_https() {
|
||||
let validator = validator_with_patterns();
|
||||
|
||||
let result = validator.validate("http://api.openai.com/v1/chat", "GET");
|
||||
assert!(!result.is_allowed());
|
||||
|
||||
if let super::AllowlistResult::Denied(reason) = result {
|
||||
assert!(matches!(reason, DenyReason::InsecureScheme(_)));
|
||||
} else {
|
||||
panic!("Expected denied");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_allow_http() {
|
||||
let validator = validator_with_patterns().allow_http();
|
||||
|
||||
let result = validator.validate("http://api.example.com/test", "GET");
|
||||
assert!(result.is_allowed());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_allowlist() {
|
||||
let validator = AllowlistValidator::new(vec![]);
|
||||
|
||||
let result = validator.validate("https://anything.com/", "GET");
|
||||
assert!(!result.is_allowed());
|
||||
|
||||
if let super::AllowlistResult::Denied(reason) = result {
|
||||
assert!(matches!(reason, DenyReason::EmptyAllowlist));
|
||||
} else {
|
||||
panic!("Expected denied");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_invalid_url() {
|
||||
let validator = validator_with_patterns();
|
||||
|
||||
let result = validator.validate("not-a-url", "GET");
|
||||
assert!(!result.is_allowed());
|
||||
|
||||
if let super::AllowlistResult::Denied(reason) = result {
|
||||
assert!(matches!(reason, DenyReason::InvalidUrl(_)));
|
||||
} else {
|
||||
panic!("Expected denied");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_url_with_port() {
|
||||
let validator =
|
||||
AllowlistValidator::new(vec![EndpointPattern::host("localhost")]).allow_http();
|
||||
|
||||
let result = validator.validate("http://localhost:8080/api", "GET");
|
||||
assert!(result.is_allowed());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,423 @@
|
||||
//! Extended capabilities for WASM sandbox.
|
||||
//!
|
||||
//! Defines the capability system that controls what a WASM tool can do.
|
||||
//! All capabilities are opt-in; tools have NO access by default.
|
||||
//!
|
||||
//! # Capability Types
|
||||
//!
|
||||
//! - **Workspace**: Read files from the agent's workspace
|
||||
//! - **HTTP**: Make HTTP requests to allowlisted endpoints
|
||||
//! - **ToolInvoke**: Call other tools via aliases
|
||||
//! - **Secrets**: Check if secrets exist (never read values)
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::secrets::CredentialMapping;
|
||||
|
||||
/// All capabilities that can be granted to a WASM tool.
|
||||
///
|
||||
/// By default, all capabilities are `None` (disabled).
|
||||
/// Each must be explicitly granted.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct Capabilities {
|
||||
/// Read files from workspace.
|
||||
pub workspace_read: Option<WorkspaceCapability>,
|
||||
/// Make HTTP requests.
|
||||
pub http: Option<HttpCapability>,
|
||||
/// Invoke other tools.
|
||||
pub tool_invoke: Option<ToolInvokeCapability>,
|
||||
/// Check if secrets exist.
|
||||
pub secrets: Option<SecretsCapability>,
|
||||
}
|
||||
|
||||
impl Capabilities {
|
||||
/// Create capabilities with no permissions.
|
||||
pub fn none() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
/// Enable workspace read with the given allowed prefixes.
|
||||
pub fn with_workspace_read(mut self, prefixes: Vec<String>) -> Self {
|
||||
self.workspace_read = Some(WorkspaceCapability {
|
||||
allowed_prefixes: prefixes,
|
||||
reader: None,
|
||||
});
|
||||
self
|
||||
}
|
||||
|
||||
/// Enable HTTP requests with the given configuration.
|
||||
pub fn with_http(mut self, http: HttpCapability) -> Self {
|
||||
self.http = Some(http);
|
||||
self
|
||||
}
|
||||
|
||||
/// Enable tool invocation with the given aliases.
|
||||
pub fn with_tool_invoke(mut self, aliases: HashMap<String, String>) -> Self {
|
||||
self.tool_invoke = Some(ToolInvokeCapability {
|
||||
aliases,
|
||||
rate_limit: RateLimitConfig::default(),
|
||||
});
|
||||
self
|
||||
}
|
||||
|
||||
/// Enable secret existence checks.
|
||||
pub fn with_secrets(mut self, allowed: Vec<String>) -> Self {
|
||||
self.secrets = Some(SecretsCapability {
|
||||
allowed_names: allowed,
|
||||
});
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// Workspace read capability configuration.
|
||||
#[derive(Clone, Default)]
|
||||
pub struct WorkspaceCapability {
|
||||
/// Allowed path prefixes (e.g., ["context/", "daily/"]).
|
||||
/// Empty means all paths allowed (within safety constraints).
|
||||
pub allowed_prefixes: Vec<String>,
|
||||
/// Function to actually read from workspace.
|
||||
/// This is injected by the runtime to avoid coupling to workspace impl.
|
||||
pub reader: Option<Arc<dyn WorkspaceReader>>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for WorkspaceCapability {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("WorkspaceCapability")
|
||||
.field("allowed_prefixes", &self.allowed_prefixes)
|
||||
.field("reader", &self.reader.is_some())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Trait for reading from workspace (allows mocking in tests).
|
||||
pub trait WorkspaceReader: Send + Sync {
|
||||
fn read(&self, path: &str) -> Option<String>;
|
||||
}
|
||||
|
||||
/// HTTP request capability configuration.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct HttpCapability {
|
||||
/// Allowed endpoint patterns.
|
||||
pub allowlist: Vec<EndpointPattern>,
|
||||
/// Credential mappings (secret name -> injection location).
|
||||
pub credentials: HashMap<String, CredentialMapping>,
|
||||
/// Rate limiting configuration.
|
||||
pub rate_limit: RateLimitConfig,
|
||||
/// Maximum request body size in bytes.
|
||||
pub max_request_bytes: usize,
|
||||
/// Maximum response body size in bytes.
|
||||
pub max_response_bytes: usize,
|
||||
/// Request timeout.
|
||||
pub timeout: Duration,
|
||||
}
|
||||
|
||||
impl Default for HttpCapability {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
allowlist: Vec::new(),
|
||||
credentials: HashMap::new(),
|
||||
rate_limit: RateLimitConfig::default(),
|
||||
max_request_bytes: 1024 * 1024, // 1 MB
|
||||
max_response_bytes: 10 * 1024 * 1024, // 10 MB
|
||||
timeout: Duration::from_secs(30),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl HttpCapability {
|
||||
/// Create a new HTTP capability with an allowlist.
|
||||
pub fn new(allowlist: Vec<EndpointPattern>) -> Self {
|
||||
Self {
|
||||
allowlist,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
/// Add a credential mapping.
|
||||
pub fn with_credential(mut self, name: impl Into<String>, mapping: CredentialMapping) -> Self {
|
||||
self.credentials.insert(name.into(), mapping);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set rate limiting.
|
||||
pub fn with_rate_limit(mut self, rate_limit: RateLimitConfig) -> Self {
|
||||
self.rate_limit = rate_limit;
|
||||
self
|
||||
}
|
||||
|
||||
/// Set request timeout.
|
||||
pub fn with_timeout(mut self, timeout: Duration) -> Self {
|
||||
self.timeout = timeout;
|
||||
self
|
||||
}
|
||||
|
||||
/// Set max request body size.
|
||||
pub fn with_max_request_bytes(mut self, bytes: usize) -> Self {
|
||||
self.max_request_bytes = bytes;
|
||||
self
|
||||
}
|
||||
|
||||
/// Set max response body size.
|
||||
pub fn with_max_response_bytes(mut self, bytes: usize) -> Self {
|
||||
self.max_response_bytes = bytes;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// Pattern for matching allowed HTTP endpoints.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct EndpointPattern {
|
||||
/// Hostname pattern (e.g., "api.example.com", "*.example.com").
|
||||
pub host: String,
|
||||
/// Path prefix (e.g., "/v1/", "/api/").
|
||||
pub path_prefix: Option<String>,
|
||||
/// Allowed HTTP methods (empty = all methods allowed).
|
||||
pub methods: Vec<String>,
|
||||
}
|
||||
|
||||
impl EndpointPattern {
|
||||
/// Create a pattern for a specific host.
|
||||
pub fn host(host: impl Into<String>) -> Self {
|
||||
Self {
|
||||
host: host.into(),
|
||||
path_prefix: None,
|
||||
methods: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Add a path prefix constraint.
|
||||
pub fn with_path_prefix(mut self, prefix: impl Into<String>) -> Self {
|
||||
self.path_prefix = Some(prefix.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// Restrict to specific HTTP methods.
|
||||
pub fn with_methods(mut self, methods: Vec<String>) -> Self {
|
||||
self.methods = methods;
|
||||
self
|
||||
}
|
||||
|
||||
/// Check if this pattern matches a URL and method.
|
||||
pub fn matches(&self, url_host: &str, url_path: &str, method: &str) -> bool {
|
||||
// Check host
|
||||
if !self.host_matches(url_host) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Check path prefix
|
||||
if let Some(ref prefix) = self.path_prefix {
|
||||
if !url_path.starts_with(prefix) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
// Check method
|
||||
if !self.methods.is_empty() {
|
||||
let method_upper = method.to_uppercase();
|
||||
if !self
|
||||
.methods
|
||||
.iter()
|
||||
.any(|m| m.to_uppercase() == method_upper)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
true
|
||||
}
|
||||
|
||||
/// Check if host pattern matches (public for allowlist validation).
|
||||
pub fn host_matches(&self, url_host: &str) -> bool {
|
||||
if self.host == url_host {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Support wildcard: *.example.com matches sub.example.com
|
||||
if let Some(suffix) = self.host.strip_prefix("*.") {
|
||||
if url_host.ends_with(suffix) && url_host.len() > suffix.len() {
|
||||
// Ensure there's a dot before the suffix (or it's the whole thing)
|
||||
let prefix = &url_host[..url_host.len() - suffix.len()];
|
||||
if prefix.ends_with('.') || prefix.is_empty() {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// Tool invocation capability.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct ToolInvokeCapability {
|
||||
/// Mapping from alias to real tool name.
|
||||
/// WASM calls tools by alias, never by real name.
|
||||
pub aliases: HashMap<String, String>,
|
||||
/// Rate limiting for tool calls.
|
||||
pub rate_limit: RateLimitConfig,
|
||||
}
|
||||
|
||||
impl ToolInvokeCapability {
|
||||
/// Create with a set of aliases.
|
||||
pub fn new(aliases: HashMap<String, String>) -> Self {
|
||||
Self {
|
||||
aliases,
|
||||
rate_limit: RateLimitConfig::default(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve an alias to a real tool name.
|
||||
pub fn resolve_alias(&self, alias: &str) -> Option<&str> {
|
||||
self.aliases.get(alias).map(|s| s.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
/// Secrets capability (existence check only).
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct SecretsCapability {
|
||||
/// Secret names this tool can check existence of.
|
||||
/// Supports glob: "openai_*" matches "openai_key", "openai_org".
|
||||
pub allowed_names: Vec<String>,
|
||||
}
|
||||
|
||||
impl SecretsCapability {
|
||||
/// Check if a secret name is allowed.
|
||||
pub fn is_allowed(&self, name: &str) -> bool {
|
||||
for pattern in &self.allowed_names {
|
||||
if pattern == name {
|
||||
return true;
|
||||
}
|
||||
if let Some(prefix) = pattern.strip_suffix('*') {
|
||||
if name.starts_with(prefix) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// Rate limiting configuration.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct RateLimitConfig {
|
||||
/// Maximum requests per minute.
|
||||
pub requests_per_minute: u32,
|
||||
/// Maximum requests per hour.
|
||||
pub requests_per_hour: u32,
|
||||
}
|
||||
|
||||
impl Default for RateLimitConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
requests_per_minute: 60,
|
||||
requests_per_hour: 1000,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl RateLimitConfig {
|
||||
/// Create a restrictive rate limit.
|
||||
pub fn restrictive() -> Self {
|
||||
Self {
|
||||
requests_per_minute: 10,
|
||||
requests_per_hour: 100,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a permissive rate limit.
|
||||
pub fn permissive() -> Self {
|
||||
Self {
|
||||
requests_per_minute: 120,
|
||||
requests_per_hour: 5000,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tools::wasm::capabilities::{Capabilities, EndpointPattern, SecretsCapability};
|
||||
|
||||
#[test]
|
||||
fn test_capabilities_default_is_none() {
|
||||
let caps = Capabilities::default();
|
||||
assert!(caps.workspace_read.is_none());
|
||||
assert!(caps.http.is_none());
|
||||
assert!(caps.tool_invoke.is_none());
|
||||
assert!(caps.secrets.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_endpoint_pattern_exact_host() {
|
||||
let pattern = EndpointPattern::host("api.example.com");
|
||||
|
||||
assert!(pattern.matches("api.example.com", "/", "GET"));
|
||||
assert!(!pattern.matches("other.example.com", "/", "GET"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_endpoint_pattern_wildcard_host() {
|
||||
let pattern = EndpointPattern::host("*.example.com");
|
||||
|
||||
assert!(pattern.matches("api.example.com", "/", "GET"));
|
||||
assert!(pattern.matches("sub.api.example.com", "/", "GET"));
|
||||
assert!(!pattern.matches("example.com", "/", "GET"));
|
||||
assert!(!pattern.matches("notexample.com", "/", "GET"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_endpoint_pattern_path_prefix() {
|
||||
let pattern = EndpointPattern::host("api.example.com").with_path_prefix("/v1/");
|
||||
|
||||
assert!(pattern.matches("api.example.com", "/v1/users", "GET"));
|
||||
assert!(pattern.matches("api.example.com", "/v1/", "GET"));
|
||||
assert!(!pattern.matches("api.example.com", "/v2/users", "GET"));
|
||||
assert!(!pattern.matches("api.example.com", "/", "GET"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_endpoint_pattern_methods() {
|
||||
let pattern = EndpointPattern::host("api.example.com")
|
||||
.with_methods(vec!["GET".to_string(), "POST".to_string()]);
|
||||
|
||||
assert!(pattern.matches("api.example.com", "/", "GET"));
|
||||
assert!(pattern.matches("api.example.com", "/", "get")); // case insensitive
|
||||
assert!(pattern.matches("api.example.com", "/", "POST"));
|
||||
assert!(!pattern.matches("api.example.com", "/", "DELETE"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_secrets_capability_exact_match() {
|
||||
let cap = SecretsCapability {
|
||||
allowed_names: vec!["openai_key".to_string()],
|
||||
};
|
||||
|
||||
assert!(cap.is_allowed("openai_key"));
|
||||
assert!(!cap.is_allowed("anthropic_key"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_secrets_capability_glob() {
|
||||
let cap = SecretsCapability {
|
||||
allowed_names: vec!["openai_*".to_string()],
|
||||
};
|
||||
|
||||
assert!(cap.is_allowed("openai_key"));
|
||||
assert!(cap.is_allowed("openai_org"));
|
||||
assert!(!cap.is_allowed("anthropic_key"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_capabilities_builder() {
|
||||
let caps = Capabilities::none()
|
||||
.with_workspace_read(vec!["context/".to_string()])
|
||||
.with_secrets(vec!["test_*".to_string()]);
|
||||
|
||||
assert!(caps.workspace_read.is_some());
|
||||
assert!(caps.secrets.is_some());
|
||||
assert!(caps.http.is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,429 @@
|
||||
//! Credential injection for WASM HTTP requests.
|
||||
//!
|
||||
//! Injects secrets into HTTP requests at the host boundary.
|
||||
//! WASM tools NEVER see the actual credential values.
|
||||
//!
|
||||
//! # Injection Flow
|
||||
//!
|
||||
//! ```text
|
||||
//! WASM requests HTTP ──► Host receives request ──► Match credentials by host
|
||||
//! │
|
||||
//! ┌───────────────────┘
|
||||
//! ▼
|
||||
//! Decrypt secret from store
|
||||
//! │
|
||||
//! ▼
|
||||
//! Inject into request:
|
||||
//! ├─► Authorization header (Bearer/Basic)
|
||||
//! ├─► Custom header (X-API-Key, etc.)
|
||||
//! └─► Query parameter
|
||||
//! │
|
||||
//! ▼
|
||||
//! Execute HTTP request
|
||||
//! ```
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::secrets::{
|
||||
CredentialLocation, CredentialMapping, DecryptedSecret, SecretError, SecretsStore,
|
||||
};
|
||||
|
||||
/// Error during credential injection.
|
||||
#[derive(Debug, Clone, thiserror::Error)]
|
||||
pub enum InjectionError {
|
||||
#[error("Secret not found: {0}")]
|
||||
SecretNotFound(String),
|
||||
|
||||
#[error("Secret access denied: {0}")]
|
||||
AccessDenied(String),
|
||||
|
||||
#[error("Secret has expired: {0}")]
|
||||
SecretExpired(String),
|
||||
|
||||
#[error("Decryption failed: {0}")]
|
||||
DecryptionFailed(String),
|
||||
|
||||
#[error("No matching credential for host: {0}")]
|
||||
NoMatchingCredential(String),
|
||||
}
|
||||
|
||||
impl From<SecretError> for InjectionError {
|
||||
fn from(e: SecretError) -> Self {
|
||||
match e {
|
||||
SecretError::NotFound(name) => InjectionError::SecretNotFound(name),
|
||||
SecretError::Expired => InjectionError::SecretExpired("unknown".to_string()),
|
||||
SecretError::AccessDenied => InjectionError::AccessDenied("unknown".to_string()),
|
||||
SecretError::DecryptionFailed(msg) => InjectionError::DecryptionFailed(msg),
|
||||
_ => InjectionError::DecryptionFailed(e.to_string()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Result of credential injection.
|
||||
#[derive(Debug)]
|
||||
pub struct InjectedCredentials {
|
||||
/// Headers to add to the request.
|
||||
pub headers: HashMap<String, String>,
|
||||
/// Query parameters to add.
|
||||
pub query_params: HashMap<String, String>,
|
||||
}
|
||||
|
||||
impl InjectedCredentials {
|
||||
pub fn empty() -> Self {
|
||||
Self {
|
||||
headers: HashMap::new(),
|
||||
query_params: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.headers.is_empty() && self.query_params.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
/// Injects credentials into HTTP requests.
|
||||
pub struct CredentialInjector {
|
||||
mappings: HashMap<String, CredentialMapping>,
|
||||
allowed_secrets: Vec<String>,
|
||||
}
|
||||
|
||||
impl CredentialInjector {
|
||||
/// Create a new injector with the given mappings.
|
||||
pub fn new(mappings: HashMap<String, CredentialMapping>, allowed_secrets: Vec<String>) -> Self {
|
||||
Self {
|
||||
mappings,
|
||||
allowed_secrets,
|
||||
}
|
||||
}
|
||||
|
||||
/// Find credentials that should be injected for a given host.
|
||||
pub fn find_credentials_for_host(&self, host: &str) -> Vec<&CredentialMapping> {
|
||||
self.mappings
|
||||
.values()
|
||||
.filter(|mapping| {
|
||||
mapping
|
||||
.host_patterns
|
||||
.iter()
|
||||
.any(|pattern| host_matches_pattern(host, pattern))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Inject credentials for an HTTP request.
|
||||
///
|
||||
/// Returns the headers and query params to add to the request.
|
||||
pub async fn inject(
|
||||
&self,
|
||||
user_id: &str,
|
||||
host: &str,
|
||||
store: &dyn SecretsStore,
|
||||
) -> Result<InjectedCredentials, InjectionError> {
|
||||
let matching_mappings = self.find_credentials_for_host(host);
|
||||
|
||||
if matching_mappings.is_empty() {
|
||||
// No credentials needed for this host
|
||||
return Ok(InjectedCredentials::empty());
|
||||
}
|
||||
|
||||
let mut result = InjectedCredentials::empty();
|
||||
|
||||
for mapping in matching_mappings {
|
||||
// Check if secret is in allowed list
|
||||
if !self.is_secret_allowed(&mapping.secret_name) {
|
||||
return Err(InjectionError::AccessDenied(mapping.secret_name.clone()));
|
||||
}
|
||||
|
||||
// Get the decrypted secret
|
||||
let secret = store
|
||||
.get_decrypted(user_id, &mapping.secret_name)
|
||||
.await
|
||||
.map_err(|e| match e {
|
||||
SecretError::NotFound(name) => InjectionError::SecretNotFound(name),
|
||||
SecretError::Expired => {
|
||||
InjectionError::SecretExpired(mapping.secret_name.clone())
|
||||
}
|
||||
_ => InjectionError::DecryptionFailed(e.to_string()),
|
||||
})?;
|
||||
|
||||
// Inject based on location
|
||||
inject_credential(&mut result, &mapping.location, &secret);
|
||||
}
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// Check if a secret name is in the allowed list.
|
||||
fn is_secret_allowed(&self, name: &str) -> bool {
|
||||
for pattern in &self.allowed_secrets {
|
||||
if pattern == name {
|
||||
return true;
|
||||
}
|
||||
if let Some(prefix) = pattern.strip_suffix('*') {
|
||||
if name.starts_with(prefix) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// Inject a single credential into the result.
|
||||
fn inject_credential(
|
||||
result: &mut InjectedCredentials,
|
||||
location: &CredentialLocation,
|
||||
secret: &DecryptedSecret,
|
||||
) {
|
||||
match location {
|
||||
CredentialLocation::AuthorizationBearer => {
|
||||
result.headers.insert(
|
||||
"Authorization".to_string(),
|
||||
format!("Bearer {}", secret.expose()),
|
||||
);
|
||||
}
|
||||
CredentialLocation::AuthorizationBasic { username } => {
|
||||
let credentials = format!("{}:{}", username, secret.expose());
|
||||
let encoded = base64_encode(credentials.as_bytes());
|
||||
result
|
||||
.headers
|
||||
.insert("Authorization".to_string(), format!("Basic {}", encoded));
|
||||
}
|
||||
CredentialLocation::Header { name, prefix } => {
|
||||
let value = match prefix {
|
||||
Some(p) => format!("{}{}", p, secret.expose()),
|
||||
None => secret.expose().to_string(),
|
||||
};
|
||||
result.headers.insert(name.clone(), value);
|
||||
}
|
||||
CredentialLocation::QueryParam { name } => {
|
||||
result
|
||||
.query_params
|
||||
.insert(name.clone(), secret.expose().to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if a host matches a pattern (supports wildcards).
|
||||
fn host_matches_pattern(host: &str, pattern: &str) -> bool {
|
||||
if pattern == host {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Support wildcard: *.example.com matches sub.example.com
|
||||
if let Some(suffix) = pattern.strip_prefix("*.") {
|
||||
if host.ends_with(suffix) && host.len() > suffix.len() {
|
||||
let prefix = &host[..host.len() - suffix.len()];
|
||||
if prefix.ends_with('.') || prefix.is_empty() {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
false
|
||||
}
|
||||
|
||||
/// Simple base64 encoding (avoids extra dependency).
|
||||
fn base64_encode(input: &[u8]) -> String {
|
||||
const ALPHABET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
|
||||
|
||||
let mut result = String::new();
|
||||
let mut i = 0;
|
||||
|
||||
while i < input.len() {
|
||||
let b0 = input[i];
|
||||
let b1 = if i + 1 < input.len() { input[i + 1] } else { 0 };
|
||||
let b2 = if i + 2 < input.len() { input[i + 2] } else { 0 };
|
||||
|
||||
result.push(ALPHABET[(b0 >> 2) as usize] as char);
|
||||
result.push(ALPHABET[(((b0 & 0x03) << 4) | (b1 >> 4)) as usize] as char);
|
||||
|
||||
if i + 1 < input.len() {
|
||||
result.push(ALPHABET[(((b1 & 0x0f) << 2) | (b2 >> 6)) as usize] as char);
|
||||
} else {
|
||||
result.push('=');
|
||||
}
|
||||
|
||||
if i + 2 < input.len() {
|
||||
result.push(ALPHABET[(b2 & 0x3f) as usize] as char);
|
||||
} else {
|
||||
result.push('=');
|
||||
}
|
||||
|
||||
i += 3;
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use secrecy::SecretString;
|
||||
|
||||
use crate::secrets::{
|
||||
CreateSecretParams, CredentialLocation, CredentialMapping, InMemorySecretsStore,
|
||||
SecretsCrypto, SecretsStore,
|
||||
};
|
||||
use crate::tools::wasm::credential_injector::{
|
||||
CredentialInjector, base64_encode, host_matches_pattern,
|
||||
};
|
||||
|
||||
fn test_store() -> InMemorySecretsStore {
|
||||
let key = "0123456789abcdef0123456789abcdef";
|
||||
let crypto = Arc::new(SecretsCrypto::new(SecretString::from(key.to_string())).unwrap());
|
||||
InMemorySecretsStore::new(crypto)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_host_matches_exact() {
|
||||
assert!(host_matches_pattern("api.openai.com", "api.openai.com"));
|
||||
assert!(!host_matches_pattern("api.openai.com", "other.com"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_host_matches_wildcard() {
|
||||
assert!(host_matches_pattern("api.example.com", "*.example.com"));
|
||||
assert!(host_matches_pattern("sub.api.example.com", "*.example.com"));
|
||||
assert!(!host_matches_pattern("example.com", "*.example.com"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_base64_encode() {
|
||||
assert_eq!(base64_encode(b"hello"), "aGVsbG8=");
|
||||
assert_eq!(base64_encode(b"user:pass"), "dXNlcjpwYXNz");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_inject_bearer() {
|
||||
let store = test_store();
|
||||
store
|
||||
.create("user1", CreateSecretParams::new("openai_key", "sk-test123"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut mappings = HashMap::new();
|
||||
mappings.insert(
|
||||
"openai".to_string(),
|
||||
CredentialMapping {
|
||||
secret_name: "openai_key".to_string(),
|
||||
location: CredentialLocation::AuthorizationBearer,
|
||||
host_patterns: vec!["api.openai.com".to_string()],
|
||||
},
|
||||
);
|
||||
|
||||
let injector = CredentialInjector::new(mappings, vec!["openai_key".to_string()]);
|
||||
let result = injector
|
||||
.inject("user1", "api.openai.com", &store)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
result.headers.get("Authorization"),
|
||||
Some(&"Bearer sk-test123".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_inject_custom_header() {
|
||||
let store = test_store();
|
||||
store
|
||||
.create("user1", CreateSecretParams::new("api_key", "secret123"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut mappings = HashMap::new();
|
||||
mappings.insert(
|
||||
"custom".to_string(),
|
||||
CredentialMapping {
|
||||
secret_name: "api_key".to_string(),
|
||||
location: CredentialLocation::Header {
|
||||
name: "X-API-Key".to_string(),
|
||||
prefix: None,
|
||||
},
|
||||
host_patterns: vec!["*.example.com".to_string()],
|
||||
},
|
||||
);
|
||||
|
||||
let injector = CredentialInjector::new(mappings, vec!["api_key".to_string()]);
|
||||
let result = injector
|
||||
.inject("user1", "api.example.com", &store)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
result.headers.get("X-API-Key"),
|
||||
Some(&"secret123".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_inject_basic_auth() {
|
||||
let store = test_store();
|
||||
store
|
||||
.create("user1", CreateSecretParams::new("password", "mypassword"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut mappings = HashMap::new();
|
||||
mappings.insert(
|
||||
"basic".to_string(),
|
||||
CredentialMapping {
|
||||
secret_name: "password".to_string(),
|
||||
location: CredentialLocation::AuthorizationBasic {
|
||||
username: "myuser".to_string(),
|
||||
},
|
||||
host_patterns: vec!["api.service.com".to_string()],
|
||||
},
|
||||
);
|
||||
|
||||
let injector = CredentialInjector::new(mappings, vec!["password".to_string()]);
|
||||
let result = injector
|
||||
.inject("user1", "api.service.com", &store)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// myuser:mypassword base64 encoded
|
||||
let expected = format!("Basic {}", base64_encode(b"myuser:mypassword"));
|
||||
assert_eq!(result.headers.get("Authorization"), Some(&expected));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_no_credentials_for_host() {
|
||||
let store = test_store();
|
||||
|
||||
let injector = CredentialInjector::new(HashMap::new(), vec![]);
|
||||
let result = injector
|
||||
.inject("user1", "unknown.com", &store)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(result.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_access_denied_for_secret() {
|
||||
let store = test_store();
|
||||
store
|
||||
.create("user1", CreateSecretParams::new("secret_key", "value"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut mappings = HashMap::new();
|
||||
mappings.insert(
|
||||
"test".to_string(),
|
||||
CredentialMapping {
|
||||
secret_name: "secret_key".to_string(),
|
||||
location: CredentialLocation::AuthorizationBearer,
|
||||
host_patterns: vec!["api.test.com".to_string()],
|
||||
},
|
||||
);
|
||||
|
||||
// Empty allowed list = nothing allowed
|
||||
let injector = CredentialInjector::new(mappings, vec![]);
|
||||
let result = injector.inject("user1", "api.test.com", &store).await;
|
||||
|
||||
assert!(result.is_err());
|
||||
}
|
||||
}
|
||||
+253
-42
@@ -2,10 +2,29 @@
|
||||
//!
|
||||
//! Implements a minimal, security-focused host API following VMLogic patterns
|
||||
//! from NEAR blockchain. The principle is: deny by default, grant minimal capabilities.
|
||||
//!
|
||||
//! # Extended API (V2)
|
||||
//!
|
||||
//! In addition to the basic log/time/workspace functions, the host now provides:
|
||||
//!
|
||||
//! - **http_request**: Make HTTP requests to allowlisted endpoints with credential injection
|
||||
//! - **tool_invoke**: Call other tools via aliases
|
||||
//! - **secret_exists**: Check if a secret exists (never read values)
|
||||
//!
|
||||
//! # Security Architecture
|
||||
//!
|
||||
//! ```text
|
||||
//! WASM Tool ──▶ Host Function ──▶ Allowlist ──▶ Credential ──▶ Execute
|
||||
//! (untrusted) (boundary) Validator Injector Request
|
||||
//! │
|
||||
//! ▼
|
||||
//! ◀────── Leak Detector ◀────── Response
|
||||
//! (sanitized, no secrets)
|
||||
//! ```
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use crate::tools::wasm::capabilities::Capabilities;
|
||||
use crate::tools::wasm::error::WasmError;
|
||||
|
||||
/// Maximum log entries per execution (prevents log spam attacks).
|
||||
@@ -44,46 +63,10 @@ pub struct LogEntry {
|
||||
pub timestamp_millis: u64,
|
||||
}
|
||||
|
||||
/// Capabilities that can be granted to a WASM tool.
|
||||
///
|
||||
/// By default, tools have NO capabilities. Each must be explicitly granted.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct Capabilities {
|
||||
/// If Some, tool can read from workspace at these paths.
|
||||
/// Empty vec means workspace access granted but no paths allowed yet.
|
||||
/// None means workspace access completely disabled.
|
||||
pub workspace_read: Option<WorkspaceCapability>,
|
||||
}
|
||||
|
||||
/// Workspace read capability configuration.
|
||||
#[derive(Clone, Default)]
|
||||
pub struct WorkspaceCapability {
|
||||
/// Allowed path prefixes (e.g., ["context/", "daily/"]).
|
||||
/// Empty means all paths allowed (within safety constraints).
|
||||
pub allowed_prefixes: Vec<String>,
|
||||
/// Function to actually read from workspace.
|
||||
/// This is injected by the runtime to avoid coupling to workspace impl.
|
||||
pub reader: Option<Arc<dyn WorkspaceReader>>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for WorkspaceCapability {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("WorkspaceCapability")
|
||||
.field("allowed_prefixes", &self.allowed_prefixes)
|
||||
.field("reader", &self.reader.is_some())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Trait for reading from workspace (allows mocking in tests).
|
||||
pub trait WorkspaceReader: Send + Sync {
|
||||
fn read(&self, path: &str) -> Option<String>;
|
||||
}
|
||||
|
||||
/// Host state maintained during WASM execution.
|
||||
///
|
||||
/// This is the "VMLogic" equivalent, it tracks all side effects and enforces limits.
|
||||
#[derive(Debug)]
|
||||
/// Extended in V2 to support HTTP requests, tool invocation, and secret checks.
|
||||
pub struct HostState {
|
||||
/// Collected log entries.
|
||||
logs: Vec<LogEntry>,
|
||||
@@ -93,6 +76,25 @@ pub struct HostState {
|
||||
capabilities: Capabilities,
|
||||
/// Count of log entries dropped due to rate limiting.
|
||||
logs_dropped: usize,
|
||||
/// User ID for secret/credential lookups.
|
||||
user_id: Option<String>,
|
||||
/// HTTP request count for rate limiting within this execution.
|
||||
http_request_count: u32,
|
||||
/// Tool invoke count for rate limiting within this execution.
|
||||
tool_invoke_count: u32,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for HostState {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("HostState")
|
||||
.field("logs_count", &self.logs.len())
|
||||
.field("logging_enabled", &self.logging_enabled)
|
||||
.field("logs_dropped", &self.logs_dropped)
|
||||
.field("user_id", &self.user_id)
|
||||
.field("http_request_count", &self.http_request_count)
|
||||
.field("tool_invoke_count", &self.tool_invoke_count)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl HostState {
|
||||
@@ -103,6 +105,22 @@ impl HostState {
|
||||
logging_enabled: true,
|
||||
capabilities,
|
||||
logs_dropped: 0,
|
||||
user_id: None,
|
||||
http_request_count: 0,
|
||||
tool_invoke_count: 0,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a new host state with user context.
|
||||
pub fn new_with_user(capabilities: Capabilities, user_id: impl Into<String>) -> Self {
|
||||
Self {
|
||||
logs: Vec::new(),
|
||||
logging_enabled: true,
|
||||
capabilities,
|
||||
logs_dropped: 0,
|
||||
user_id: Some(user_id.into()),
|
||||
http_request_count: 0,
|
||||
tool_invoke_count: 0,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -111,6 +129,16 @@ impl HostState {
|
||||
Self::new(Capabilities::default())
|
||||
}
|
||||
|
||||
/// Get the user ID if set.
|
||||
pub fn user_id(&self) -> Option<&str> {
|
||||
self.user_id.as_deref()
|
||||
}
|
||||
|
||||
/// Get the capabilities.
|
||||
pub fn capabilities(&self) -> &Capabilities {
|
||||
&self.capabilities
|
||||
}
|
||||
|
||||
/// Log a message from WASM.
|
||||
///
|
||||
/// Returns Ok(()) if logged, Err if rate limited or too long.
|
||||
@@ -204,6 +232,114 @@ impl HostState {
|
||||
pub fn logs_dropped(&self) -> usize {
|
||||
self.logs_dropped
|
||||
}
|
||||
|
||||
/// Check if a secret exists (does not expose value).
|
||||
///
|
||||
/// Returns false if:
|
||||
/// - Secrets capability not granted
|
||||
/// - Secret name not in allowed list
|
||||
/// - User ID not set
|
||||
pub fn secret_exists(&self, name: &str) -> bool {
|
||||
let capability = match &self.capabilities.secrets {
|
||||
Some(cap) => cap,
|
||||
None => return false,
|
||||
};
|
||||
|
||||
// Check if name is allowed
|
||||
capability.is_allowed(name)
|
||||
}
|
||||
|
||||
/// Check if HTTP capability is available for a given URL and method.
|
||||
///
|
||||
/// Returns an error message if not allowed.
|
||||
pub fn check_http_allowed(&self, url: &str, method: &str) -> Result<(), String> {
|
||||
let capability = self
|
||||
.capabilities
|
||||
.http
|
||||
.as_ref()
|
||||
.ok_or_else(|| "HTTP capability not granted".to_string())?;
|
||||
|
||||
// Use the allowlist validator
|
||||
use crate::tools::wasm::allowlist::AllowlistValidator;
|
||||
|
||||
let validator = AllowlistValidator::new(capability.allowlist.clone());
|
||||
let result = validator.validate(url, method);
|
||||
|
||||
if result.is_allowed() {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(format!("HTTP request not allowed: {:?}", result))
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if tool invocation is allowed for an alias.
|
||||
///
|
||||
/// Returns the real tool name if allowed, error otherwise.
|
||||
pub fn check_tool_invoke_allowed(&self, alias: &str) -> Result<String, String> {
|
||||
let capability = self
|
||||
.capabilities
|
||||
.tool_invoke
|
||||
.as_ref()
|
||||
.ok_or_else(|| "Tool invocation capability not granted".to_string())?;
|
||||
|
||||
capability
|
||||
.resolve_alias(alias)
|
||||
.map(|s| s.to_string())
|
||||
.ok_or_else(|| format!("Unknown tool alias: {}", alias))
|
||||
}
|
||||
|
||||
/// Increment HTTP request counter and check rate limit.
|
||||
///
|
||||
/// Returns error if rate limit exceeded.
|
||||
pub fn record_http_request(&mut self) -> Result<(), String> {
|
||||
// Verify HTTP capability exists
|
||||
let _capability = self
|
||||
.capabilities
|
||||
.http
|
||||
.as_ref()
|
||||
.ok_or_else(|| "HTTP capability not granted".to_string())?;
|
||||
|
||||
self.http_request_count += 1;
|
||||
|
||||
// Simple per-execution rate limit (additional to global rate limiter)
|
||||
// This prevents a single execution from making too many requests
|
||||
const MAX_REQUESTS_PER_EXECUTION: u32 = 50;
|
||||
if self.http_request_count > MAX_REQUESTS_PER_EXECUTION {
|
||||
return Err(format!(
|
||||
"Too many HTTP requests in single execution (max {})",
|
||||
MAX_REQUESTS_PER_EXECUTION
|
||||
));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Increment tool invoke counter and check rate limit.
|
||||
///
|
||||
/// Returns error if rate limit exceeded.
|
||||
pub fn record_tool_invoke(&mut self) -> Result<(), String> {
|
||||
self.tool_invoke_count += 1;
|
||||
|
||||
const MAX_INVOKES_PER_EXECUTION: u32 = 20;
|
||||
if self.tool_invoke_count > MAX_INVOKES_PER_EXECUTION {
|
||||
return Err(format!(
|
||||
"Too many tool invocations in single execution (max {})",
|
||||
MAX_INVOKES_PER_EXECUTION
|
||||
));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Get HTTP request count for this execution.
|
||||
pub fn http_request_count(&self) -> u32 {
|
||||
self.http_request_count
|
||||
}
|
||||
|
||||
/// Get tool invoke count for this execution.
|
||||
pub fn tool_invoke_count(&self) -> u32 {
|
||||
self.tool_invoke_count
|
||||
}
|
||||
}
|
||||
|
||||
/// Validate a workspace path for security.
|
||||
@@ -243,12 +379,15 @@ fn validate_workspace_path(path: &str) -> Result<(), WasmError> {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tools::wasm::host::{
|
||||
Capabilities, HostState, LogLevel, MAX_LOG_ENTRIES, MAX_LOG_MESSAGE_BYTES,
|
||||
WorkspaceCapability, WorkspaceReader, validate_workspace_path,
|
||||
};
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::tools::wasm::capabilities::{
|
||||
Capabilities, SecretsCapability, WorkspaceCapability, WorkspaceReader,
|
||||
};
|
||||
use crate::tools::wasm::host::{
|
||||
HostState, LogLevel, MAX_LOG_ENTRIES, MAX_LOG_MESSAGE_BYTES, validate_workspace_path,
|
||||
};
|
||||
|
||||
struct MockReader {
|
||||
content: String,
|
||||
}
|
||||
@@ -330,6 +469,7 @@ mod tests {
|
||||
allowed_prefixes: vec![],
|
||||
reader: Some(reader),
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let state = HostState::new(capabilities);
|
||||
@@ -348,6 +488,7 @@ mod tests {
|
||||
allowed_prefixes: vec!["context/".to_string()],
|
||||
reader: Some(reader),
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let state = HostState::new(capabilities);
|
||||
@@ -392,4 +533,74 @@ mod tests {
|
||||
assert!(validate_workspace_path("projects/alpha/notes.md").is_ok());
|
||||
assert!(validate_workspace_path("MEMORY.md").is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_secret_exists_no_capability() {
|
||||
let state = HostState::minimal();
|
||||
assert!(!state.secret_exists("any_secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_secret_exists_with_capability() {
|
||||
let capabilities = Capabilities {
|
||||
secrets: Some(SecretsCapability {
|
||||
allowed_names: vec!["openai_*".to_string(), "exact_name".to_string()],
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let state = HostState::new(capabilities);
|
||||
|
||||
// Glob match
|
||||
assert!(state.secret_exists("openai_key"));
|
||||
assert!(state.secret_exists("openai_org"));
|
||||
|
||||
// Exact match
|
||||
assert!(state.secret_exists("exact_name"));
|
||||
|
||||
// Not allowed
|
||||
assert!(!state.secret_exists("stripe_key"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_http_request_rate_limit() {
|
||||
// Create state with HTTP capability enabled
|
||||
let capabilities = Capabilities {
|
||||
http: Some(crate::tools::wasm::capabilities::HttpCapability::default()),
|
||||
..Default::default()
|
||||
};
|
||||
let mut state = HostState::new(capabilities);
|
||||
|
||||
// Should allow up to 50 requests
|
||||
for _ in 0..50 {
|
||||
assert!(state.record_http_request().is_ok());
|
||||
}
|
||||
|
||||
// 51st should fail
|
||||
assert!(state.record_http_request().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_invoke_rate_limit() {
|
||||
// Create state with tool invoke capability enabled
|
||||
let capabilities = Capabilities {
|
||||
tool_invoke: Some(crate::tools::wasm::capabilities::ToolInvokeCapability::default()),
|
||||
..Default::default()
|
||||
};
|
||||
let mut state = HostState::new(capabilities);
|
||||
|
||||
// Should allow up to 20 invocations
|
||||
for _ in 0..20 {
|
||||
assert!(state.record_tool_invoke().is_ok());
|
||||
}
|
||||
|
||||
// 21st should fail
|
||||
assert!(state.record_tool_invoke().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_new_with_user() {
|
||||
let state = HostState::new_with_user(Capabilities::default(), "user123");
|
||||
assert_eq!(state.user_id(), Some("user123"));
|
||||
}
|
||||
}
|
||||
|
||||
+50
-20
@@ -10,26 +10,23 @@
|
||||
//!
|
||||
//! - **Memory limits**: Memory growth is bounded via ResourceLimiter.
|
||||
//!
|
||||
//! - **Minimal host API**: Only log, time, and optional workspace read.
|
||||
//! - **Extended host API (V2)**: log, time, workspace, HTTP, tool invoke, secrets
|
||||
//!
|
||||
//! - **Capability-based security**: Features are opt-in via Capabilities.
|
||||
//!
|
||||
//! # Architecture
|
||||
//! # Architecture (V2)
|
||||
//!
|
||||
//! ```text
|
||||
//! ┌─────────────────────────────────────────────────────────────────────┐
|
||||
//! │ Tool Registration │
|
||||
//! │ WASM bytes → Validate → Compile (AOT) → PreparedModule (cached) │
|
||||
//! └─────────────────────────────────────────────────────────────────────┘
|
||||
//! │
|
||||
//! ▼
|
||||
//! ┌─────────────────────────────────────────────────────────────────────┐
|
||||
//! │ Tool Execution │
|
||||
//! │ JSON params → WasmToolWrapper → Fresh Instance → Execute → Result │
|
||||
//! │ ↓ ↓ │
|
||||
//! │ ResourceLimiter HostState │
|
||||
//! │ (memory, fuel) (log, time, workspace) │
|
||||
//! └─────────────────────────────────────────────────────────────────────┘
|
||||
//! ┌─────────────────────────────────────────────────────────────────────────────┐
|
||||
//! │ WASM Tool Execution │
|
||||
//! │ │
|
||||
//! │ WASM Tool ──▶ Host Function ──▶ Allowlist ──▶ Credential ──▶ Execute │
|
||||
//! │ (untrusted) (boundary) Validator Injector Request │
|
||||
//! │ │ │
|
||||
//! │ ▼ │
|
||||
//! │ ◀────── Leak Detector ◀────── Response │
|
||||
//! │ (sanitized, no secrets) │
|
||||
//! └─────────────────────────────────────────────────────────────────────────────┘
|
||||
//! ```
|
||||
//!
|
||||
//! # Security Constraints
|
||||
@@ -40,17 +37,22 @@
|
||||
//! | Memory exhaustion | ResourceLimiter, 10MB default |
|
||||
//! | Infinite loops | Epoch interruption + tokio timeout |
|
||||
//! | Filesystem access | No WASI FS, only host workspace_read |
|
||||
//! | Network access | No network host functions |
|
||||
//! | Network access | Allowlisted endpoints only |
|
||||
//! | Credential exposure | Injection at host boundary only |
|
||||
//! | Secret exfiltration | Leak detector scans all outputs |
|
||||
//! | Log spam | Max 1000 entries, 4KB per message |
|
||||
//! | Path traversal | Validate paths (no `..`, no `/` prefix) |
|
||||
//! | Trap recovery | Discard instance, never reuse |
|
||||
//! | Side channels | Fresh instance per execution |
|
||||
//! | Rate abuse | Per-tool rate limiting |
|
||||
//! | WASM tampering | BLAKE3 hash verification on load |
|
||||
//! | Direct tool access | Tool aliasing (indirection layer) |
|
||||
//!
|
||||
//! # Example
|
||||
//!
|
||||
//! ```ignore
|
||||
//! use near_agent::tools::wasm::{WasmToolRuntime, WasmRuntimeConfig, WasmToolWrapper};
|
||||
//! use near_agent::tools::wasm::host::Capabilities;
|
||||
//! use near_agent::tools::wasm::Capabilities;
|
||||
//! use std::sync::Arc;
|
||||
//!
|
||||
//! // Create runtime
|
||||
@@ -60,24 +62,52 @@
|
||||
//! let wasm_bytes = std::fs::read("my_tool.wasm")?;
|
||||
//! let prepared = runtime.prepare("my_tool", &wasm_bytes, None).await?;
|
||||
//!
|
||||
//! // Create wrapper with minimal capabilities
|
||||
//! let tool = WasmToolWrapper::new(runtime, prepared, Capabilities::default());
|
||||
//! // Create wrapper with HTTP capability
|
||||
//! let capabilities = Capabilities::none()
|
||||
//! .with_http(HttpCapability::new(vec![
|
||||
//! EndpointPattern::host("api.openai.com").with_path_prefix("/v1/"),
|
||||
//! ]));
|
||||
//! let tool = WasmToolWrapper::new(runtime, prepared, capabilities);
|
||||
//!
|
||||
//! // Execute (implements Tool trait)
|
||||
//! let output = tool.execute(serde_json::json!({"input": "test"}), &ctx).await?;
|
||||
//! ```
|
||||
|
||||
mod allowlist;
|
||||
mod capabilities;
|
||||
mod credential_injector;
|
||||
mod error;
|
||||
mod host;
|
||||
mod limits;
|
||||
mod rate_limiter;
|
||||
mod runtime;
|
||||
mod storage;
|
||||
mod wrapper;
|
||||
|
||||
// Core types
|
||||
pub use error::{TrapCode, TrapInfo, WasmError};
|
||||
pub use host::{Capabilities, HostState, LogEntry, LogLevel, WorkspaceCapability, WorkspaceReader};
|
||||
pub use host::{HostState, LogEntry, LogLevel};
|
||||
pub use limits::{
|
||||
DEFAULT_FUEL_LIMIT, DEFAULT_MEMORY_LIMIT, DEFAULT_TIMEOUT, FuelConfig, ResourceLimits,
|
||||
WasmResourceLimiter,
|
||||
};
|
||||
pub use runtime::{PreparedModule, WasmRuntimeConfig, WasmToolRuntime};
|
||||
pub use wrapper::WasmToolWrapper;
|
||||
|
||||
// Capabilities (V2)
|
||||
pub use capabilities::{
|
||||
Capabilities, EndpointPattern, HttpCapability, RateLimitConfig, SecretsCapability,
|
||||
ToolInvokeCapability, WorkspaceCapability, WorkspaceReader,
|
||||
};
|
||||
|
||||
// Security components (V2)
|
||||
pub use allowlist::{AllowlistResult, AllowlistValidator, DenyReason};
|
||||
pub use credential_injector::{CredentialInjector, InjectedCredentials, InjectionError};
|
||||
pub use rate_limiter::{LimitType, RateLimitError, RateLimitResult, RateLimiter};
|
||||
|
||||
// Storage (V2)
|
||||
pub use storage::{
|
||||
PostgresWasmToolStore, StoreToolParams, StoredCapabilities, StoredWasmTool,
|
||||
StoredWasmToolWithBinary, ToolStatus, TrustLevel, WasmStorageError, WasmToolStore,
|
||||
compute_binary_hash, verify_binary_integrity,
|
||||
};
|
||||
|
||||
@@ -0,0 +1,422 @@
|
||||
//! Rate limiting for WASM tool operations.
|
||||
//!
|
||||
//! Provides per-tool rate limiting for HTTP requests and tool invocations.
|
||||
//! Uses a sliding window algorithm for smooth rate enforcement.
|
||||
//!
|
||||
//! # Rate Limit Algorithm
|
||||
//!
|
||||
//! Uses a simplified sliding window counter:
|
||||
//! - Track request counts for current minute and hour windows
|
||||
//! - Reset counters when window expires
|
||||
//! - Increment counter and check against limits
|
||||
//!
|
||||
//! # Persistence
|
||||
//!
|
||||
//! Rate limit state can be persisted to PostgreSQL for cross-process
|
||||
//! rate limiting (useful for distributed deployments).
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use crate::tools::wasm::capabilities::RateLimitConfig;
|
||||
|
||||
/// Result of a rate limit check.
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum RateLimitResult {
|
||||
/// Request is allowed.
|
||||
Allowed {
|
||||
/// Remaining requests in the current minute.
|
||||
remaining_minute: u32,
|
||||
/// Remaining requests in the current hour.
|
||||
remaining_hour: u32,
|
||||
},
|
||||
/// Request is rate limited.
|
||||
Limited {
|
||||
/// When the rate limit will reset.
|
||||
retry_after: Duration,
|
||||
/// Which limit was exceeded.
|
||||
limit_type: LimitType,
|
||||
},
|
||||
}
|
||||
|
||||
impl RateLimitResult {
|
||||
pub fn is_allowed(&self) -> bool {
|
||||
matches!(self, RateLimitResult::Allowed { .. })
|
||||
}
|
||||
}
|
||||
|
||||
/// Which rate limit was exceeded.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum LimitType {
|
||||
PerMinute,
|
||||
PerHour,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for LimitType {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
LimitType::PerMinute => write!(f, "per-minute"),
|
||||
LimitType::PerHour => write!(f, "per-hour"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// State for a single rate limit window.
|
||||
#[derive(Debug, Clone)]
|
||||
struct WindowState {
|
||||
window_start: Instant,
|
||||
count: u32,
|
||||
}
|
||||
|
||||
impl WindowState {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
window_start: Instant::now(),
|
||||
count: 0,
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if the window has expired and reset if needed.
|
||||
fn maybe_reset(&mut self, window_duration: Duration) {
|
||||
if self.window_start.elapsed() >= window_duration {
|
||||
self.window_start = Instant::now();
|
||||
self.count = 0;
|
||||
}
|
||||
}
|
||||
|
||||
/// Time until window resets.
|
||||
fn time_until_reset(&self, window_duration: Duration) -> Duration {
|
||||
let elapsed = self.window_start.elapsed();
|
||||
if elapsed >= window_duration {
|
||||
Duration::ZERO
|
||||
} else {
|
||||
window_duration - elapsed
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Rate limit state for a single tool.
|
||||
#[derive(Debug)]
|
||||
struct ToolRateLimitState {
|
||||
minute_window: WindowState,
|
||||
hour_window: WindowState,
|
||||
}
|
||||
|
||||
impl ToolRateLimitState {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
minute_window: WindowState::new(),
|
||||
hour_window: WindowState::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// In-memory rate limiter for WASM tools.
|
||||
pub struct RateLimiter {
|
||||
/// State per (user_id, tool_name).
|
||||
state: RwLock<HashMap<(String, String), ToolRateLimitState>>,
|
||||
}
|
||||
|
||||
impl RateLimiter {
|
||||
/// Create a new rate limiter.
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
state: RwLock::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if a request is allowed and record it if so.
|
||||
pub async fn check_and_record(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tool_name: &str,
|
||||
config: &RateLimitConfig,
|
||||
) -> RateLimitResult {
|
||||
let key = (user_id.to_string(), tool_name.to_string());
|
||||
|
||||
let mut state = self.state.write().await;
|
||||
let tool_state = state.entry(key).or_insert_with(ToolRateLimitState::new);
|
||||
|
||||
// Reset windows if expired
|
||||
tool_state
|
||||
.minute_window
|
||||
.maybe_reset(Duration::from_secs(60));
|
||||
tool_state
|
||||
.hour_window
|
||||
.maybe_reset(Duration::from_secs(3600));
|
||||
|
||||
// Check minute limit
|
||||
if tool_state.minute_window.count >= config.requests_per_minute {
|
||||
return RateLimitResult::Limited {
|
||||
retry_after: tool_state
|
||||
.minute_window
|
||||
.time_until_reset(Duration::from_secs(60)),
|
||||
limit_type: LimitType::PerMinute,
|
||||
};
|
||||
}
|
||||
|
||||
// Check hour limit
|
||||
if tool_state.hour_window.count >= config.requests_per_hour {
|
||||
return RateLimitResult::Limited {
|
||||
retry_after: tool_state
|
||||
.hour_window
|
||||
.time_until_reset(Duration::from_secs(3600)),
|
||||
limit_type: LimitType::PerHour,
|
||||
};
|
||||
}
|
||||
|
||||
// Record the request
|
||||
tool_state.minute_window.count += 1;
|
||||
tool_state.hour_window.count += 1;
|
||||
|
||||
RateLimitResult::Allowed {
|
||||
remaining_minute: config.requests_per_minute - tool_state.minute_window.count,
|
||||
remaining_hour: config.requests_per_hour - tool_state.hour_window.count,
|
||||
}
|
||||
}
|
||||
|
||||
/// Check without recording (for preview/estimation).
|
||||
pub async fn check(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tool_name: &str,
|
||||
config: &RateLimitConfig,
|
||||
) -> RateLimitResult {
|
||||
let key = (user_id.to_string(), tool_name.to_string());
|
||||
|
||||
let mut state = self.state.write().await;
|
||||
let tool_state = state.entry(key).or_insert_with(ToolRateLimitState::new);
|
||||
|
||||
// Reset windows if expired
|
||||
tool_state
|
||||
.minute_window
|
||||
.maybe_reset(Duration::from_secs(60));
|
||||
tool_state
|
||||
.hour_window
|
||||
.maybe_reset(Duration::from_secs(3600));
|
||||
|
||||
// Check minute limit
|
||||
if tool_state.minute_window.count >= config.requests_per_minute {
|
||||
return RateLimitResult::Limited {
|
||||
retry_after: tool_state
|
||||
.minute_window
|
||||
.time_until_reset(Duration::from_secs(60)),
|
||||
limit_type: LimitType::PerMinute,
|
||||
};
|
||||
}
|
||||
|
||||
// Check hour limit
|
||||
if tool_state.hour_window.count >= config.requests_per_hour {
|
||||
return RateLimitResult::Limited {
|
||||
retry_after: tool_state
|
||||
.hour_window
|
||||
.time_until_reset(Duration::from_secs(3600)),
|
||||
limit_type: LimitType::PerHour,
|
||||
};
|
||||
}
|
||||
|
||||
RateLimitResult::Allowed {
|
||||
remaining_minute: config.requests_per_minute - tool_state.minute_window.count,
|
||||
remaining_hour: config.requests_per_hour - tool_state.hour_window.count,
|
||||
}
|
||||
}
|
||||
|
||||
/// Get current usage for a tool.
|
||||
pub async fn get_usage(&self, user_id: &str, tool_name: &str) -> Option<(u32, u32)> {
|
||||
let key = (user_id.to_string(), tool_name.to_string());
|
||||
let state = self.state.read().await;
|
||||
|
||||
state
|
||||
.get(&key)
|
||||
.map(|s| (s.minute_window.count, s.hour_window.count))
|
||||
}
|
||||
|
||||
/// Clear rate limit state for a tool (for testing or manual reset).
|
||||
pub async fn clear(&self, user_id: &str, tool_name: &str) {
|
||||
let key = (user_id.to_string(), tool_name.to_string());
|
||||
self.state.write().await.remove(&key);
|
||||
}
|
||||
|
||||
/// Clear all rate limit state.
|
||||
pub async fn clear_all(&self) {
|
||||
self.state.write().await.clear();
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for RateLimiter {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// Error when rate limited.
|
||||
#[derive(Debug, Clone, thiserror::Error)]
|
||||
#[error("Rate limited ({limit_type}), retry after {retry_after:?}")]
|
||||
pub struct RateLimitError {
|
||||
pub retry_after: Duration,
|
||||
pub limit_type: LimitType,
|
||||
}
|
||||
|
||||
impl From<RateLimitResult> for Result<(), RateLimitError> {
|
||||
fn from(result: RateLimitResult) -> Self {
|
||||
match result {
|
||||
RateLimitResult::Allowed { .. } => Ok(()),
|
||||
RateLimitResult::Limited {
|
||||
retry_after,
|
||||
limit_type,
|
||||
} => Err(RateLimitError {
|
||||
retry_after,
|
||||
limit_type,
|
||||
}),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tools::wasm::capabilities::RateLimitConfig;
|
||||
use crate::tools::wasm::rate_limiter::{LimitType, RateLimitResult, RateLimiter};
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_allowed_within_limits() {
|
||||
let limiter = RateLimiter::new();
|
||||
let config = RateLimitConfig {
|
||||
requests_per_minute: 10,
|
||||
requests_per_hour: 100,
|
||||
};
|
||||
|
||||
let result = limiter.check_and_record("user1", "tool1", &config).await;
|
||||
|
||||
match result {
|
||||
RateLimitResult::Allowed {
|
||||
remaining_minute,
|
||||
remaining_hour,
|
||||
} => {
|
||||
assert_eq!(remaining_minute, 9);
|
||||
assert_eq!(remaining_hour, 99);
|
||||
}
|
||||
_ => panic!("Expected allowed"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_minute_limit_exceeded() {
|
||||
let limiter = RateLimiter::new();
|
||||
let config = RateLimitConfig {
|
||||
requests_per_minute: 2,
|
||||
requests_per_hour: 100,
|
||||
};
|
||||
|
||||
// Use up the minute limit
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
|
||||
// Third request should be limited
|
||||
let result = limiter.check_and_record("user1", "tool1", &config).await;
|
||||
|
||||
match result {
|
||||
RateLimitResult::Limited {
|
||||
limit_type,
|
||||
retry_after,
|
||||
} => {
|
||||
assert_eq!(limit_type, LimitType::PerMinute);
|
||||
assert!(retry_after.as_secs() <= 60);
|
||||
}
|
||||
_ => panic!("Expected limited"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_hour_limit_exceeded() {
|
||||
let limiter = RateLimiter::new();
|
||||
let config = RateLimitConfig {
|
||||
requests_per_minute: 100, // High minute limit
|
||||
requests_per_hour: 2, // Low hour limit
|
||||
};
|
||||
|
||||
// Use up the hour limit
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
|
||||
// Third request should be limited
|
||||
let result = limiter.check_and_record("user1", "tool1", &config).await;
|
||||
|
||||
match result {
|
||||
RateLimitResult::Limited { limit_type, .. } => {
|
||||
assert_eq!(limit_type, LimitType::PerHour);
|
||||
}
|
||||
_ => panic!("Expected limited"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_user_isolation() {
|
||||
let limiter = RateLimiter::new();
|
||||
let config = RateLimitConfig {
|
||||
requests_per_minute: 1,
|
||||
requests_per_hour: 10,
|
||||
};
|
||||
|
||||
// User1 uses their limit
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
let result1 = limiter.check_and_record("user1", "tool1", &config).await;
|
||||
|
||||
// User2 should still have their limit
|
||||
let result2 = limiter.check_and_record("user2", "tool1", &config).await;
|
||||
|
||||
assert!(!result1.is_allowed());
|
||||
assert!(result2.is_allowed());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tool_isolation() {
|
||||
let limiter = RateLimiter::new();
|
||||
let config = RateLimitConfig {
|
||||
requests_per_minute: 1,
|
||||
requests_per_hour: 10,
|
||||
};
|
||||
|
||||
// Tool1 uses its limit
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
let result1 = limiter.check_and_record("user1", "tool1", &config).await;
|
||||
|
||||
// Tool2 should still have its limit
|
||||
let result2 = limiter.check_and_record("user1", "tool2", &config).await;
|
||||
|
||||
assert!(!result1.is_allowed());
|
||||
assert!(result2.is_allowed());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_get_usage() {
|
||||
let limiter = RateLimiter::new();
|
||||
let config = RateLimitConfig::default();
|
||||
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
|
||||
let usage = limiter.get_usage("user1", "tool1").await;
|
||||
assert_eq!(usage, Some((3, 3)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_clear() {
|
||||
let limiter = RateLimiter::new();
|
||||
let config = RateLimitConfig {
|
||||
requests_per_minute: 1,
|
||||
requests_per_hour: 10,
|
||||
};
|
||||
|
||||
limiter.check_and_record("user1", "tool1", &config).await;
|
||||
let result1 = limiter.check_and_record("user1", "tool1", &config).await;
|
||||
assert!(!result1.is_allowed());
|
||||
|
||||
limiter.clear("user1", "tool1").await;
|
||||
|
||||
let result2 = limiter.check_and_record("user1", "tool1", &config).await;
|
||||
assert!(result2.is_allowed());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,616 @@
|
||||
//! WASM binary storage with integrity verification.
|
||||
//!
|
||||
//! Stores compiled WASM tools in PostgreSQL with BLAKE3 hash verification.
|
||||
//! On load, the hash is verified to detect tampering.
|
||||
//!
|
||||
//! # Storage Flow
|
||||
//!
|
||||
//! ```text
|
||||
//! WASM bytes ──► BLAKE3 hash ──► Store in PostgreSQL
|
||||
//! │ (binary + hash)
|
||||
//! │
|
||||
//! └──► Later: Load ──► Verify hash ──► Return bytes
|
||||
//! ```
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use chrono::{DateTime, Utc};
|
||||
use deadpool_postgres::Pool;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::tools::wasm::capabilities::{
|
||||
Capabilities, EndpointPattern, HttpCapability, RateLimitConfig, SecretsCapability,
|
||||
ToolInvokeCapability,
|
||||
};
|
||||
|
||||
/// Trust level for a WASM tool.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum TrustLevel {
|
||||
/// Built-in system tool (highest trust).
|
||||
System,
|
||||
/// Audited and verified tool.
|
||||
Verified,
|
||||
/// User-uploaded tool (untrusted).
|
||||
User,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for TrustLevel {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
TrustLevel::System => write!(f, "system"),
|
||||
TrustLevel::Verified => write!(f, "verified"),
|
||||
TrustLevel::User => write!(f, "user"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::str::FromStr for TrustLevel {
|
||||
type Err = String;
|
||||
|
||||
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||
match s.to_lowercase().as_str() {
|
||||
"system" => Ok(TrustLevel::System),
|
||||
"verified" => Ok(TrustLevel::Verified),
|
||||
"user" => Ok(TrustLevel::User),
|
||||
_ => Err(format!("Unknown trust level: {}", s)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Status of a WASM tool.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ToolStatus {
|
||||
/// Tool is active and can be used.
|
||||
Active,
|
||||
/// Tool is disabled (manually or due to errors).
|
||||
Disabled,
|
||||
/// Tool is quarantined (suspected malicious).
|
||||
Quarantined,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ToolStatus {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
ToolStatus::Active => write!(f, "active"),
|
||||
ToolStatus::Disabled => write!(f, "disabled"),
|
||||
ToolStatus::Quarantined => write!(f, "quarantined"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::str::FromStr for ToolStatus {
|
||||
type Err = String;
|
||||
|
||||
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||
match s.to_lowercase().as_str() {
|
||||
"active" => Ok(ToolStatus::Active),
|
||||
"disabled" => Ok(ToolStatus::Disabled),
|
||||
"quarantined" => Ok(ToolStatus::Quarantined),
|
||||
_ => Err(format!("Unknown status: {}", s)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// A stored WASM tool.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct StoredWasmTool {
|
||||
pub id: Uuid,
|
||||
pub user_id: String,
|
||||
pub name: String,
|
||||
pub version: String,
|
||||
pub description: String,
|
||||
pub parameters_schema: serde_json::Value,
|
||||
pub source_url: Option<String>,
|
||||
pub trust_level: TrustLevel,
|
||||
pub status: ToolStatus,
|
||||
pub created_at: DateTime<Utc>,
|
||||
pub updated_at: DateTime<Utc>,
|
||||
}
|
||||
|
||||
/// Full tool data including binary (not returned by default for efficiency).
|
||||
#[derive(Debug)]
|
||||
pub struct StoredWasmToolWithBinary {
|
||||
pub tool: StoredWasmTool,
|
||||
pub wasm_binary: Vec<u8>,
|
||||
pub binary_hash: Vec<u8>,
|
||||
}
|
||||
|
||||
/// Capabilities stored in the database.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct StoredCapabilities {
|
||||
pub id: Uuid,
|
||||
pub wasm_tool_id: Uuid,
|
||||
pub http_allowlist: Vec<EndpointPattern>,
|
||||
pub allowed_secrets: Vec<String>,
|
||||
pub tool_aliases: HashMap<String, String>,
|
||||
pub requests_per_minute: u32,
|
||||
pub requests_per_hour: u32,
|
||||
pub max_request_body_bytes: i64,
|
||||
pub max_response_body_bytes: i64,
|
||||
pub workspace_read_prefixes: Vec<String>,
|
||||
pub http_timeout_secs: i32,
|
||||
}
|
||||
|
||||
impl StoredCapabilities {
|
||||
/// Convert to runtime Capabilities struct.
|
||||
pub fn to_capabilities(&self) -> Capabilities {
|
||||
let mut caps = Capabilities::default();
|
||||
|
||||
// Workspace read
|
||||
if !self.workspace_read_prefixes.is_empty() {
|
||||
caps = caps.with_workspace_read(self.workspace_read_prefixes.clone());
|
||||
}
|
||||
|
||||
// HTTP capability
|
||||
if !self.http_allowlist.is_empty() {
|
||||
caps.http = Some(HttpCapability {
|
||||
allowlist: self.http_allowlist.clone(),
|
||||
credentials: HashMap::new(), // Loaded separately
|
||||
rate_limit: RateLimitConfig {
|
||||
requests_per_minute: self.requests_per_minute,
|
||||
requests_per_hour: self.requests_per_hour,
|
||||
},
|
||||
max_request_bytes: self.max_request_body_bytes as usize,
|
||||
max_response_bytes: self.max_response_body_bytes as usize,
|
||||
timeout: std::time::Duration::from_secs(self.http_timeout_secs as u64),
|
||||
});
|
||||
}
|
||||
|
||||
// Tool invoke capability
|
||||
if !self.tool_aliases.is_empty() {
|
||||
caps.tool_invoke = Some(ToolInvokeCapability {
|
||||
aliases: self.tool_aliases.clone(),
|
||||
rate_limit: RateLimitConfig {
|
||||
requests_per_minute: self.requests_per_minute,
|
||||
requests_per_hour: self.requests_per_hour,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
// Secrets capability
|
||||
if !self.allowed_secrets.is_empty() {
|
||||
caps.secrets = Some(SecretsCapability {
|
||||
allowed_names: self.allowed_secrets.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
caps
|
||||
}
|
||||
}
|
||||
|
||||
/// Error from WASM storage operations.
|
||||
#[derive(Debug, Clone, thiserror::Error)]
|
||||
pub enum WasmStorageError {
|
||||
#[error("Tool not found: {0}")]
|
||||
NotFound(String),
|
||||
|
||||
#[error("Tool is disabled")]
|
||||
Disabled,
|
||||
|
||||
#[error("Tool is quarantined")]
|
||||
Quarantined,
|
||||
|
||||
#[error("Binary integrity check failed: hash mismatch")]
|
||||
IntegrityCheckFailed,
|
||||
|
||||
#[error("Database error: {0}")]
|
||||
Database(String),
|
||||
|
||||
#[error("Invalid data: {0}")]
|
||||
InvalidData(String),
|
||||
}
|
||||
|
||||
/// Trait for WASM tool storage.
|
||||
#[async_trait]
|
||||
pub trait WasmToolStore: Send + Sync {
|
||||
/// Store a new WASM tool.
|
||||
async fn store(&self, params: StoreToolParams) -> Result<StoredWasmTool, WasmStorageError>;
|
||||
|
||||
/// Get tool metadata (without binary).
|
||||
async fn get(&self, user_id: &str, name: &str) -> Result<StoredWasmTool, WasmStorageError>;
|
||||
|
||||
/// Get tool with binary (verifies integrity).
|
||||
async fn get_with_binary(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<StoredWasmToolWithBinary, WasmStorageError>;
|
||||
|
||||
/// Get tool capabilities.
|
||||
async fn get_capabilities(
|
||||
&self,
|
||||
tool_id: Uuid,
|
||||
) -> Result<Option<StoredCapabilities>, WasmStorageError>;
|
||||
|
||||
/// List all tools for a user.
|
||||
async fn list(&self, user_id: &str) -> Result<Vec<StoredWasmTool>, WasmStorageError>;
|
||||
|
||||
/// Update tool status.
|
||||
async fn update_status(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
status: ToolStatus,
|
||||
) -> Result<(), WasmStorageError>;
|
||||
|
||||
/// Delete a tool.
|
||||
async fn delete(&self, user_id: &str, name: &str) -> Result<bool, WasmStorageError>;
|
||||
}
|
||||
|
||||
/// Parameters for storing a new tool.
|
||||
pub struct StoreToolParams {
|
||||
pub user_id: String,
|
||||
pub name: String,
|
||||
pub version: String,
|
||||
pub description: String,
|
||||
pub wasm_binary: Vec<u8>,
|
||||
pub parameters_schema: serde_json::Value,
|
||||
pub source_url: Option<String>,
|
||||
pub trust_level: TrustLevel,
|
||||
}
|
||||
|
||||
/// Compute BLAKE3 hash of WASM binary.
|
||||
pub fn compute_binary_hash(binary: &[u8]) -> Vec<u8> {
|
||||
let hash = blake3::hash(binary);
|
||||
hash.as_bytes().to_vec()
|
||||
}
|
||||
|
||||
/// Verify binary integrity against stored hash.
|
||||
pub fn verify_binary_integrity(binary: &[u8], expected_hash: &[u8]) -> bool {
|
||||
let actual_hash = compute_binary_hash(binary);
|
||||
actual_hash == expected_hash
|
||||
}
|
||||
|
||||
/// PostgreSQL implementation of WasmToolStore.
|
||||
pub struct PostgresWasmToolStore {
|
||||
pool: Pool,
|
||||
}
|
||||
|
||||
impl PostgresWasmToolStore {
|
||||
pub fn new(pool: Pool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl WasmToolStore for PostgresWasmToolStore {
|
||||
async fn store(&self, params: StoreToolParams) -> Result<StoredWasmTool, WasmStorageError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
let binary_hash = compute_binary_hash(¶ms.wasm_binary);
|
||||
let id = Uuid::new_v4();
|
||||
let now = Utc::now();
|
||||
|
||||
let row = client
|
||||
.query_one(
|
||||
r#"
|
||||
INSERT INTO wasm_tools (
|
||||
id, user_id, name, version, description, wasm_binary, binary_hash,
|
||||
parameters_schema, source_url, trust_level, status, created_at, updated_at
|
||||
)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, 'active', $11, $11)
|
||||
ON CONFLICT (user_id, name, version) DO UPDATE SET
|
||||
description = EXCLUDED.description,
|
||||
wasm_binary = EXCLUDED.wasm_binary,
|
||||
binary_hash = EXCLUDED.binary_hash,
|
||||
parameters_schema = EXCLUDED.parameters_schema,
|
||||
source_url = EXCLUDED.source_url,
|
||||
updated_at = NOW()
|
||||
RETURNING id, user_id, name, version, description, parameters_schema,
|
||||
source_url, trust_level, status, created_at, updated_at
|
||||
"#,
|
||||
&[
|
||||
&id,
|
||||
¶ms.user_id,
|
||||
¶ms.name,
|
||||
¶ms.version,
|
||||
¶ms.description,
|
||||
¶ms.wasm_binary,
|
||||
&binary_hash,
|
||||
¶ms.parameters_schema,
|
||||
¶ms.source_url,
|
||||
¶ms.trust_level.to_string(),
|
||||
&now,
|
||||
],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
row_to_tool(&row)
|
||||
}
|
||||
|
||||
async fn get(&self, user_id: &str, name: &str) -> Result<StoredWasmTool, WasmStorageError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
let row = client
|
||||
.query_opt(
|
||||
r#"
|
||||
SELECT id, user_id, name, version, description, parameters_schema,
|
||||
source_url, trust_level, status, created_at, updated_at
|
||||
FROM wasm_tools
|
||||
WHERE user_id = $1 AND name = $2 AND status = 'active'
|
||||
ORDER BY version DESC
|
||||
LIMIT 1
|
||||
"#,
|
||||
&[&user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
match row {
|
||||
Some(r) => {
|
||||
let tool = row_to_tool(&r)?;
|
||||
match tool.status {
|
||||
ToolStatus::Active => Ok(tool),
|
||||
ToolStatus::Disabled => Err(WasmStorageError::Disabled),
|
||||
ToolStatus::Quarantined => Err(WasmStorageError::Quarantined),
|
||||
}
|
||||
}
|
||||
None => Err(WasmStorageError::NotFound(name.to_string())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_with_binary(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<StoredWasmToolWithBinary, WasmStorageError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
let row = client
|
||||
.query_opt(
|
||||
r#"
|
||||
SELECT id, user_id, name, version, description, wasm_binary, binary_hash,
|
||||
parameters_schema, source_url, trust_level, status, created_at, updated_at
|
||||
FROM wasm_tools
|
||||
WHERE user_id = $1 AND name = $2 AND status = 'active'
|
||||
ORDER BY version DESC
|
||||
LIMIT 1
|
||||
"#,
|
||||
&[&user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
match row {
|
||||
Some(r) => {
|
||||
let wasm_binary: Vec<u8> = r.get("wasm_binary");
|
||||
let binary_hash: Vec<u8> = r.get("binary_hash");
|
||||
|
||||
// Verify integrity
|
||||
if !verify_binary_integrity(&wasm_binary, &binary_hash) {
|
||||
tracing::error!(
|
||||
user_id = user_id,
|
||||
name = name,
|
||||
"WASM binary integrity check failed"
|
||||
);
|
||||
return Err(WasmStorageError::IntegrityCheckFailed);
|
||||
}
|
||||
|
||||
let tool = row_to_tool(&r)?;
|
||||
|
||||
match tool.status {
|
||||
ToolStatus::Active => Ok(StoredWasmToolWithBinary {
|
||||
tool,
|
||||
wasm_binary,
|
||||
binary_hash,
|
||||
}),
|
||||
ToolStatus::Disabled => Err(WasmStorageError::Disabled),
|
||||
ToolStatus::Quarantined => Err(WasmStorageError::Quarantined),
|
||||
}
|
||||
}
|
||||
None => Err(WasmStorageError::NotFound(name.to_string())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_capabilities(
|
||||
&self,
|
||||
tool_id: Uuid,
|
||||
) -> Result<Option<StoredCapabilities>, WasmStorageError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
let row = client
|
||||
.query_opt(
|
||||
r#"
|
||||
SELECT id, wasm_tool_id, http_allowlist, allowed_secrets, tool_aliases,
|
||||
requests_per_minute, requests_per_hour, max_request_body_bytes,
|
||||
max_response_body_bytes, workspace_read_prefixes, http_timeout_secs
|
||||
FROM tool_capabilities
|
||||
WHERE wasm_tool_id = $1
|
||||
"#,
|
||||
&[&tool_id],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
match row {
|
||||
Some(r) => {
|
||||
let http_allowlist_json: serde_json::Value = r.get("http_allowlist");
|
||||
let tool_aliases_json: serde_json::Value = r.get("tool_aliases");
|
||||
|
||||
let http_allowlist: Vec<EndpointPattern> =
|
||||
serde_json::from_value(http_allowlist_json).unwrap_or_default();
|
||||
let tool_aliases: HashMap<String, String> =
|
||||
serde_json::from_value(tool_aliases_json).unwrap_or_default();
|
||||
|
||||
Ok(Some(StoredCapabilities {
|
||||
id: r.get("id"),
|
||||
wasm_tool_id: r.get("wasm_tool_id"),
|
||||
http_allowlist,
|
||||
allowed_secrets: r.get("allowed_secrets"),
|
||||
tool_aliases,
|
||||
requests_per_minute: r.get::<_, i32>("requests_per_minute") as u32,
|
||||
requests_per_hour: r.get::<_, i32>("requests_per_hour") as u32,
|
||||
max_request_body_bytes: r.get("max_request_body_bytes"),
|
||||
max_response_body_bytes: r.get("max_response_body_bytes"),
|
||||
workspace_read_prefixes: r.get("workspace_read_prefixes"),
|
||||
http_timeout_secs: r.get("http_timeout_secs"),
|
||||
}))
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list(&self, user_id: &str) -> Result<Vec<StoredWasmTool>, WasmStorageError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
let rows = client
|
||||
.query(
|
||||
r#"
|
||||
SELECT DISTINCT ON (name) id, user_id, name, version, description,
|
||||
parameters_schema, source_url, trust_level, status, created_at, updated_at
|
||||
FROM wasm_tools
|
||||
WHERE user_id = $1
|
||||
ORDER BY name, version DESC
|
||||
"#,
|
||||
&[&user_id],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
rows.into_iter().map(|r| row_to_tool(&r)).collect()
|
||||
}
|
||||
|
||||
async fn update_status(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
status: ToolStatus,
|
||||
) -> Result<(), WasmStorageError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
let result = client
|
||||
.execute(
|
||||
"UPDATE wasm_tools SET status = $1, updated_at = NOW() WHERE user_id = $2 AND name = $3",
|
||||
&[&status.to_string(), &user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
if result == 0 {
|
||||
return Err(WasmStorageError::NotFound(name.to_string()));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn delete(&self, user_id: &str, name: &str) -> Result<bool, WasmStorageError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
let result = client
|
||||
.execute(
|
||||
"DELETE FROM wasm_tools WHERE user_id = $1 AND name = $2",
|
||||
&[&user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmStorageError::Database(e.to_string()))?;
|
||||
|
||||
Ok(result > 0)
|
||||
}
|
||||
}
|
||||
|
||||
fn row_to_tool(row: &tokio_postgres::Row) -> Result<StoredWasmTool, WasmStorageError> {
|
||||
let trust_level_str: String = row.get("trust_level");
|
||||
let status_str: String = row.get("status");
|
||||
|
||||
Ok(StoredWasmTool {
|
||||
id: row.get("id"),
|
||||
user_id: row.get("user_id"),
|
||||
name: row.get("name"),
|
||||
version: row.get("version"),
|
||||
description: row.get("description"),
|
||||
parameters_schema: row.get("parameters_schema"),
|
||||
source_url: row.get("source_url"),
|
||||
trust_level: trust_level_str
|
||||
.parse()
|
||||
.map_err(WasmStorageError::InvalidData)?,
|
||||
status: status_str
|
||||
.parse()
|
||||
.map_err(WasmStorageError::InvalidData)?,
|
||||
created_at: row.get("created_at"),
|
||||
updated_at: row.get("updated_at"),
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tools::wasm::storage::{
|
||||
ToolStatus, TrustLevel, compute_binary_hash, verify_binary_integrity,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn test_compute_hash() {
|
||||
let binary = b"(module)";
|
||||
let hash = compute_binary_hash(binary);
|
||||
assert_eq!(hash.len(), 32); // BLAKE3 produces 32-byte hash
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_verify_integrity_success() {
|
||||
let binary = b"test wasm binary content";
|
||||
let hash = compute_binary_hash(binary);
|
||||
assert!(verify_binary_integrity(binary, &hash));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_verify_integrity_failure() {
|
||||
let binary = b"test wasm binary content";
|
||||
let hash = compute_binary_hash(binary);
|
||||
let tampered = b"tampered wasm binary content";
|
||||
assert!(!verify_binary_integrity(tampered, &hash));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trust_level_parse() {
|
||||
assert_eq!("system".parse::<TrustLevel>().unwrap(), TrustLevel::System);
|
||||
assert_eq!(
|
||||
"verified".parse::<TrustLevel>().unwrap(),
|
||||
TrustLevel::Verified
|
||||
);
|
||||
assert_eq!("user".parse::<TrustLevel>().unwrap(), TrustLevel::User);
|
||||
assert!("invalid".parse::<TrustLevel>().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_status_parse() {
|
||||
assert_eq!("active".parse::<ToolStatus>().unwrap(), ToolStatus::Active);
|
||||
assert_eq!(
|
||||
"disabled".parse::<ToolStatus>().unwrap(),
|
||||
ToolStatus::Disabled
|
||||
);
|
||||
assert_eq!(
|
||||
"quarantined".parse::<ToolStatus>().unwrap(),
|
||||
ToolStatus::Quarantined
|
||||
);
|
||||
assert!("invalid".parse::<ToolStatus>().is_err());
|
||||
}
|
||||
}
|
||||
@@ -12,8 +12,9 @@ use wasmtime::component::{Component, Linker, Val};
|
||||
|
||||
use crate::context::JobContext;
|
||||
use crate::tools::tool::{Tool, ToolError, ToolOutput};
|
||||
use crate::tools::wasm::capabilities::Capabilities;
|
||||
use crate::tools::wasm::error::WasmError;
|
||||
use crate::tools::wasm::host::{Capabilities, HostState, LogLevel};
|
||||
use crate::tools::wasm::host::{HostState, LogLevel};
|
||||
use crate::tools::wasm::limits::{ResourceLimits, WasmResourceLimiter};
|
||||
use crate::tools::wasm::runtime::{PreparedModule, WasmToolRuntime};
|
||||
|
||||
@@ -380,10 +381,11 @@ impl std::fmt::Debug for WasmToolWrapper {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tools::wasm::host::Capabilities;
|
||||
use crate::tools::wasm::runtime::{WasmRuntimeConfig, WasmToolRuntime};
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::tools::wasm::capabilities::Capabilities;
|
||||
use crate::tools::wasm::runtime::{WasmRuntimeConfig, WasmToolRuntime};
|
||||
|
||||
#[test]
|
||||
fn test_wrapper_creation() {
|
||||
// This test verifies the runtime can be created
|
||||
@@ -399,5 +401,8 @@ mod tests {
|
||||
fn test_capabilities_default() {
|
||||
let caps = Capabilities::default();
|
||||
assert!(caps.workspace_read.is_none());
|
||||
assert!(caps.http.is_none());
|
||||
assert!(caps.tool_invoke.is_none());
|
||||
assert!(caps.secrets.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,6 +2,12 @@
|
||||
//
|
||||
// Defines the contract between sandboxed tools and the host runtime.
|
||||
// Tools export the `tool` interface; the host provides the `host` interface.
|
||||
//
|
||||
// Security Model:
|
||||
// - WASM tools are untrusted and run in a sandbox
|
||||
// - All capabilities are opt-in (default: no access)
|
||||
// - Secrets are NEVER exposed to WASM; credentials are injected at host boundary
|
||||
// - All outputs are scanned for secret leakage before returning to WASM
|
||||
|
||||
package near:agent;
|
||||
|
||||
@@ -33,6 +39,65 @@ interface host {
|
||||
/// Path must be relative (no leading /) and cannot contain "..".
|
||||
/// Returns None if the file doesn't exist or capability not granted.
|
||||
workspace-read: func(path: string) -> option<string>;
|
||||
|
||||
// ==================== HTTP Capability ====================
|
||||
|
||||
/// Response from an HTTP request.
|
||||
record http-response {
|
||||
/// HTTP status code.
|
||||
status: u16,
|
||||
/// Response headers as JSON object string.
|
||||
headers-json: string,
|
||||
/// Response body bytes.
|
||||
body: list<u8>,
|
||||
}
|
||||
|
||||
/// Make an HTTP request (if capability granted).
|
||||
///
|
||||
/// Security:
|
||||
/// - Only allowed endpoints (host/path patterns) can be accessed
|
||||
/// - Credentials are injected by the host; WASM never sees them
|
||||
/// - Response is scanned for leaked secrets before returning
|
||||
/// - Rate-limited per tool
|
||||
///
|
||||
/// Returns Err with error message if:
|
||||
/// - Endpoint not in allowlist
|
||||
/// - Rate limit exceeded
|
||||
/// - Request/response size limit exceeded
|
||||
/// - Network error
|
||||
/// - Timeout
|
||||
/// - Secret leak detected in response
|
||||
http-request: func(
|
||||
method: string,
|
||||
url: string,
|
||||
headers-json: string,
|
||||
body: option<list<u8>>
|
||||
) -> result<http-response, string>;
|
||||
|
||||
// ==================== Tool Invocation Capability ====================
|
||||
|
||||
/// Invoke another tool by alias (if capability granted).
|
||||
///
|
||||
/// Security:
|
||||
/// - WASM calls tools by alias, not real name (indirection layer)
|
||||
/// - Only aliased tools can be invoked
|
||||
/// - Rate-limited per tool
|
||||
/// - Output is scanned for leaked secrets before returning
|
||||
///
|
||||
/// Returns the tool output as JSON string, or Err with error message.
|
||||
tool-invoke: func(alias: string, params-json: string) -> result<string, string>;
|
||||
|
||||
// ==================== Secrets Capability ====================
|
||||
|
||||
/// Check if a secret exists (if capability granted).
|
||||
///
|
||||
/// Security:
|
||||
/// - WASM can only check existence, NEVER read values
|
||||
/// - Only allowed secret names can be checked
|
||||
/// - Actual credentials are injected by host during HTTP requests
|
||||
///
|
||||
/// Returns true if the secret exists and is accessible to this tool.
|
||||
secret-exists: func(name: string) -> bool;
|
||||
}
|
||||
|
||||
/// Tool interface that sandboxed tools must implement.
|
||||
|
||||
Reference in New Issue
Block a user