From b0d9af355cc5c24fa7de5c00d468441034372c26 Mon Sep 17 00:00:00 2001 From: Trezy Date: Fri, 20 Mar 2026 12:30:29 -0500 Subject: [PATCH] feat: complete plugin sync pipeline with auth, tokens, and PDS writes --- Cargo.lock | 1 + Cargo.toml | 1 + docs/plugins.md | 51 ++ ...60320100000_create_external_auth_state.sql | 12 + ...60320100000_create_external_auth_state.sql | 12 + plugins/steam/.gitignore | 1 + plugins/steam/Cargo.lock | 107 +++ plugins/steam/Cargo.toml | 15 + plugins/steam/src/lib.rs | 738 ++++++++++++++++++ src/external_auth/mod.rs | 3 + src/external_auth/pds_write.rs | 127 +++ src/external_auth/routes.rs | 212 ++++- src/external_auth/state.rs | 108 +++ src/external_auth/tokens.rs | 198 +++++ src/lib.rs | 2 + src/lua/atproto_api.rs | 4 + src/lua/db_api.rs | 4 + src/lua/execute.rs | 4 + src/lua/http_api.rs | 4 + src/main.rs | 22 + src/plugin/attestation.rs | 331 ++++++++ src/plugin/executor.rs | 483 ++++++++++++ src/plugin/host/bindings.rs | 438 +++++++++++ src/plugin/host/lookup.rs | 31 + src/plugin/host/mod.rs | 2 + src/plugin/loader.rs | 116 ++- src/plugin/memory.rs | 193 +++++ src/plugin/mod.rs | 6 + src/plugin/runtime.rs | 35 + src/plugin/sync.rs | 312 ++++++++ src/plugin/types.rs | 3 + tests/common/app.rs | 4 + tests/fixtures/test_plugin/.gitignore | 1 + tests/fixtures/test_plugin/Cargo.lock | 7 + tests/fixtures/test_plugin/Cargo.toml | 11 + tests/fixtures/test_plugin/src/lib.rs | 105 +++ tests/lua_atproto_api.rs | 4 + tests/lua_db_api.rs | 4 + tests/plugin_executor.rs | 224 ++++++ 39 files changed, 3902 insertions(+), 34 deletions(-) create mode 100644 docs/plugins.md create mode 100644 migrations/postgres/20260320100000_create_external_auth_state.sql create mode 100644 migrations/sqlite/20260320100000_create_external_auth_state.sql create mode 100644 plugins/steam/.gitignore create mode 100644 plugins/steam/Cargo.lock create mode 100644 plugins/steam/Cargo.toml create mode 100644 plugins/steam/src/lib.rs create mode 100644 src/external_auth/pds_write.rs create mode 100644 src/external_auth/state.rs create mode 100644 src/external_auth/tokens.rs create mode 100644 src/plugin/attestation.rs create mode 100644 src/plugin/executor.rs create mode 100644 src/plugin/host/bindings.rs create mode 100644 src/plugin/memory.rs create mode 100644 src/plugin/sync.rs create mode 100644 tests/fixtures/test_plugin/.gitignore create mode 100644 tests/fixtures/test_plugin/Cargo.lock create mode 100644 tests/fixtures/test_plugin/Cargo.toml create mode 100644 tests/fixtures/test_plugin/src/lib.rs create mode 100644 tests/plugin_executor.rs diff --git a/Cargo.lock b/Cargo.lock index 9c9026a..52298e4 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1637,6 +1637,7 @@ dependencies = [ "bytes", "chrono", "ciborium", + "cid", "dashmap", "dotenvy", "futures-util", diff --git a/Cargo.toml b/Cargo.toml index 3226c6e..72963a8 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -24,6 +24,7 @@ jose-jwk = { version = "0.1", default-features = false, features = ["p256"] } jsonwebtoken = "9" bytes = "1" chrono = { version = "0.4", features = ["serde"] } +cid = "0.11" ciborium = "0.2" k256 = { version = "0.13", features = ["ecdsa"] } multibase = "0.9" diff --git a/docs/plugins.md b/docs/plugins.md new file mode 100644 index 0000000..095fd17 --- /dev/null +++ b/docs/plugins.md @@ -0,0 +1,51 @@ +# HappyView Plugin System + +HappyView supports WASM plugins for extending functionality. The first plugin type is external auth providers (Steam, GOG, Epic, etc.). + +## Configuration + +### Environment Variables + +- `TOKEN_ENCRYPTION_KEY`: Base64-encoded 32-byte key for encrypting OAuth tokens (required for external auth) +- `PLUGIN_URLS`: Comma-separated list of plugins to load from URLs + +### PLUGIN_URLS Format + +``` +id|url|sha256:hash,id|url|sha256:hash +``` + +Example: +``` +PLUGIN_URLS=steam|https://github.com/org/plugins/releases/download/v1.0.0/steam.wasm|sha256:abc123 +``` + +### File-based Plugins + +Place plugins in the `./plugins/` directory: + +``` +plugins/ + steam/ + plugin.wasm + plugin.toml +``` + +## API Endpoints + +- `GET /external-auth/providers` - List available auth providers +- `GET /external-auth/{plugin_id}/authorize?redirect_uri=...` - Start auth flow +- `GET /external-auth/{plugin_id}/callback` - OAuth callback +- `POST /external-auth/{plugin_id}/sync` - Sync account data +- `POST /external-auth/{plugin_id}/unlink` - Unlink account + +## Plugin Development + +See the [Plugin Development Guide](./plugin-development.md) for creating custom plugins. + +## Security + +- OAuth tokens are encrypted at rest using AES-256-GCM +- Plugins run in a sandboxed WASM environment +- Plugins can only access host functions (HTTP, KV, secrets, logging) +- KV storage is scoped per-plugin and per-user diff --git a/migrations/postgres/20260320100000_create_external_auth_state.sql b/migrations/postgres/20260320100000_create_external_auth_state.sql new file mode 100644 index 0000000..9e28dd9 --- /dev/null +++ b/migrations/postgres/20260320100000_create_external_auth_state.sql @@ -0,0 +1,12 @@ +-- OAuth state for external auth flows (e.g., Steam OpenID) +CREATE TABLE external_auth_state ( + state TEXT PRIMARY KEY, + did TEXT NOT NULL, + plugin_id TEXT NOT NULL, + redirect_uri TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + expires_at TIMESTAMPTZ NOT NULL +); + +-- Index for cleanup of expired state +CREATE INDEX idx_external_auth_state_expires ON external_auth_state(expires_at); diff --git a/migrations/sqlite/20260320100000_create_external_auth_state.sql b/migrations/sqlite/20260320100000_create_external_auth_state.sql new file mode 100644 index 0000000..0034169 --- /dev/null +++ b/migrations/sqlite/20260320100000_create_external_auth_state.sql @@ -0,0 +1,12 @@ +-- OAuth state for external auth flows (e.g., Steam OpenID) +CREATE TABLE external_auth_state ( + state TEXT PRIMARY KEY, + did TEXT NOT NULL, + plugin_id TEXT NOT NULL, + redirect_uri TEXT NOT NULL, + created_at TEXT NOT NULL DEFAULT (datetime('now')), + expires_at TEXT NOT NULL +); + +-- Index for cleanup of expired state +CREATE INDEX idx_external_auth_state_expires ON external_auth_state(expires_at); diff --git a/plugins/steam/.gitignore b/plugins/steam/.gitignore new file mode 100644 index 0000000..b83d222 --- /dev/null +++ b/plugins/steam/.gitignore @@ -0,0 +1 @@ +/target/ diff --git a/plugins/steam/Cargo.lock b/plugins/steam/Cargo.lock new file mode 100644 index 0000000..6895834 --- /dev/null +++ b/plugins/steam/Cargo.lock @@ -0,0 +1,107 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "itoa" +version = "1.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" + +[[package]] +name = "memchr" +version = "2.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "quote" +version = "1.0.45" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41f2619966050689382d2b44f664f4bc593e129785a36d6ee376ddf37259b924" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.149" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + +[[package]] +name = "steam-plugin" +version = "0.1.0" +dependencies = [ + "serde", + "serde_json", +] + +[[package]] +name = "syn" +version = "2.0.117" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e665b8803e7b1d2a727f4023456bbbbe74da67099c585258af0ad9c5013b9b99" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "zmij" +version = "1.0.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" diff --git a/plugins/steam/Cargo.toml b/plugins/steam/Cargo.toml new file mode 100644 index 0000000..3bed82e --- /dev/null +++ b/plugins/steam/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "steam-plugin" +version = "0.1.0" +edition = "2021" + +[lib] +crate-type = ["cdylib"] + +[dependencies] +serde = { version = "1", default-features = false, features = ["derive", "alloc"] } +serde_json = { version = "1", default-features = false, features = ["alloc"] } + +[profile.release] +opt-level = "s" +lto = true diff --git a/plugins/steam/src/lib.rs b/plugins/steam/src/lib.rs new file mode 100644 index 0000000..cd41de1 --- /dev/null +++ b/plugins/steam/src/lib.rs @@ -0,0 +1,738 @@ +// Steam Plugin for HappyView +// Uses OpenID 2.0 for authentication and Steam Web API for data + +#![cfg_attr(target_arch = "wasm32", no_std)] +#![allow(static_mut_refs)] + +#[cfg(target_arch = "wasm32")] +extern crate alloc; + +#[cfg(target_arch = "wasm32")] +use alloc::{format, string::String, string::ToString, vec::Vec}; + +#[cfg(target_arch = "wasm32")] +use core::alloc::{GlobalAlloc, Layout}; + +use serde::{Deserialize, Serialize}; + +// ============================================================================ +// Memory Management (WASM only) +// ============================================================================ + +#[cfg(target_arch = "wasm32")] +struct BumpAllocator; + +#[cfg(target_arch = "wasm32")] +const HEAP_SIZE: usize = 131072; // 128KB + +#[cfg(target_arch = "wasm32")] +static mut HEAP: [u8; HEAP_SIZE] = [0; HEAP_SIZE]; + +#[cfg(target_arch = "wasm32")] +static mut HEAP_POS: usize = 0; + +#[cfg(target_arch = "wasm32")] +unsafe impl GlobalAlloc for BumpAllocator { + unsafe fn alloc(&self, layout: Layout) -> *mut u8 { + let size = layout.size(); + let align = layout.align(); + let pos = (HEAP_POS + align - 1) & !(align - 1); + if pos + size > HEAP_SIZE { + return core::ptr::null_mut(); + } + HEAP_POS = pos + size; + HEAP.as_mut_ptr().add(pos) + } + + unsafe fn dealloc(&self, _ptr: *mut u8, _layout: Layout) { + // No-op for bump allocator + } +} + +#[cfg(target_arch = "wasm32")] +#[global_allocator] +static ALLOCATOR: BumpAllocator = BumpAllocator; + +#[cfg(target_arch = "wasm32")] +#[panic_handler] +fn panic(_info: &core::panic::PanicInfo) -> ! { + loop {} +} + +// ============================================================================ +// Host Function Imports +// ============================================================================ + +#[cfg(target_arch = "wasm32")] +extern "C" { + fn host_http_request(req_ptr: i32, req_len: i32) -> i64; + fn host_get_secret(name_ptr: i32, name_len: i32) -> i64; +} + +// ============================================================================ +// Memory Exports +// ============================================================================ + +#[no_mangle] +pub extern "C" fn alloc(size: u32) -> u32 { + #[cfg(target_arch = "wasm32")] + { + let layout = Layout::from_size_align(size as usize, 1).unwrap(); + unsafe { ALLOCATOR.alloc(layout) as u32 } + } + #[cfg(not(target_arch = "wasm32"))] + { + let _ = size; + 0 + } +} + +#[no_mangle] +pub extern "C" fn dealloc(_ptr: u32, _size: u32) { + // No-op for bump allocator +} + +// ============================================================================ +// Helper Functions +// ============================================================================ + +fn return_json(s: &str) -> i64 { + let ptr = alloc(s.len() as u32); + if ptr == 0 { + return 0; + } + #[cfg(target_arch = "wasm32")] + unsafe { + core::ptr::copy_nonoverlapping(s.as_ptr(), ptr as *mut u8, s.len()); + } + ((ptr as i64) << 32) | (s.len() as i64) +} + +fn return_ok(value: &T) -> i64 { + let json = serde_json::to_string(&Response::Ok(value)).unwrap_or_default(); + return_json(&json) +} + +fn return_error(code: &str, message: &str, retryable: bool) -> i64 { + let err = ErrorResponse { + code: code.into(), + message: message.into(), + retryable, + }; + let json = serde_json::to_string(&Response::<()>::Err(err)).unwrap_or_default(); + return_json(&json) +} + +#[cfg(target_arch = "wasm32")] +fn read_input(ptr: u32, len: u32) -> Option> { + if len == 0 || len > 1024 * 1024 { + return None; + } + let slice = unsafe { core::slice::from_raw_parts(ptr as *const u8, len as usize) }; + Some(slice.to_vec()) +} + +#[cfg(target_arch = "wasm32")] +fn read_host_response(packed: i64) -> Option> { + if packed == 0 { + return None; + } + let ptr = (packed >> 32) as u32; + let len = (packed & 0xFFFFFFFF) as u32; + if len == 0 || len > 10 * 1024 * 1024 { + return None; + } + let slice = unsafe { core::slice::from_raw_parts(ptr as *const u8, len as usize) }; + Some(slice.to_vec()) +} + +#[cfg(target_arch = "wasm32")] +fn get_secret(name: &str) -> Option { + let packed = unsafe { host_get_secret(name.as_ptr() as i32, name.len() as i32) }; + let bytes = read_host_response(packed)?; + // Host returns JSON: {"ok": "value"} or {"error": ...} + let resp: Response = serde_json::from_slice(&bytes).ok()?; + match resp { + Response::Ok(val) => Some(val), + Response::Err(_) => None, + } +} + +#[cfg(target_arch = "wasm32")] +fn http_get(url: &str) -> Result { + let req = HttpRequest { + method: "GET".into(), + url: url.into(), + headers: alloc::vec![], + body: None, + }; + let req_json = serde_json::to_string(&req).map_err(|e| format!("serialize: {}", e))?; + let packed = unsafe { host_http_request(req_json.as_ptr() as i32, req_json.len() as i32) }; + let bytes = read_host_response(packed).ok_or("no response")?; + let resp: Response = + serde_json::from_slice(&bytes).map_err(|e| format!("parse: {}", e))?; + match resp { + Response::Ok(r) => r.body.ok_or_else(|| "empty body".into()), + Response::Err(e) => Err(e.message), + } +} + +#[cfg(target_arch = "wasm32")] +fn http_post(url: &str, body: &str, content_type: &str) -> Result { + let req = HttpRequest { + method: "POST".into(), + url: url.into(), + headers: alloc::vec![("Content-Type".into(), content_type.into())], + body: Some(body.into()), + }; + let req_json = serde_json::to_string(&req).map_err(|e| format!("serialize: {}", e))?; + let packed = unsafe { host_http_request(req_json.as_ptr() as i32, req_json.len() as i32) }; + let bytes = read_host_response(packed).ok_or("no response")?; + let resp: Response = + serde_json::from_slice(&bytes).map_err(|e| format!("parse: {}", e))?; + match resp { + Response::Ok(r) => r.body.ok_or_else(|| "empty body".into()), + Response::Err(e) => Err(e.message), + } +} + +// ============================================================================ +// Types +// ============================================================================ + +#[derive(Serialize, Deserialize)] +#[serde(untagged)] +enum Response { + Ok(T), + Err(ErrorResponse), +} + +#[derive(Serialize, Deserialize)] +struct ErrorResponse { + code: String, + message: String, + retryable: bool, +} + +#[derive(Serialize, Deserialize)] +struct PluginInfo { + id: String, + name: String, + version: String, + api_version: String, + icon_url: Option, + required_secrets: Vec, + config_schema: Option, +} + +#[derive(Serialize, Deserialize)] +struct AuthorizeInput { + state: String, + redirect_uri: String, + config: serde_json::Value, +} + +#[derive(Serialize, Deserialize)] +struct CallbackInput { + code: Option, + state: String, + config: serde_json::Value, + #[serde(flatten)] + extra: serde_json::Map, +} + +#[derive(Serialize, Deserialize)] +struct TokenSet { + access_token: String, + token_type: String, + expires_at: Option, + refresh_token: Option, +} + +#[derive(Serialize, Deserialize)] +struct ProfileInput { + access_token: String, + config: serde_json::Value, +} + +#[derive(Serialize, Deserialize)] +struct ExternalProfile { + account_id: String, + display_name: Option, + profile_url: Option, + avatar_url: Option, +} + +#[derive(Serialize, Deserialize)] +struct SyncInput { + access_token: String, + config: serde_json::Value, +} + +#[derive(Serialize, Deserialize)] +struct SyncRecord { + collection: String, + record: serde_json::Value, + dedup_key: Option, + /// Whether HappyView should add an attestation signature + sign: bool, +} + +#[derive(Serialize, Deserialize)] +struct HttpRequest { + method: String, + url: String, + headers: Vec<(String, String)>, + body: Option, +} + +#[derive(Serialize, Deserialize)] +struct HttpResponse { + status: u16, + headers: Vec<(String, String)>, + body: Option, +} + +// Steam API types +#[derive(Deserialize)] +struct SteamOwnedGamesResponse { + response: SteamOwnedGames, +} + +#[derive(Deserialize)] +#[allow(dead_code)] +struct SteamOwnedGames { + game_count: Option, + games: Option>, +} + +#[derive(Deserialize)] +#[allow(dead_code)] +struct SteamGame { + appid: u64, + name: Option, + playtime_forever: Option, + img_icon_url: Option, + playtime_2weeks: Option, +} + +#[derive(Deserialize)] +struct SteamPlayerSummary { + response: SteamPlayersResponse, +} + +#[derive(Deserialize)] +struct SteamPlayersResponse { + players: Vec, +} + +#[derive(Deserialize)] +struct SteamPlayer { + steamid: String, + personaname: Option, + profileurl: Option, + avatarfull: Option, +} + +// ============================================================================ +// Steam OpenID 2.0 Constants +// ============================================================================ + +const STEAM_OPENID_URL: &str = "https://steamcommunity.com/openid/login"; +const STEAM_API_BASE: &str = "https://api.steampowered.com"; + +// ============================================================================ +// Plugin Exports +// ============================================================================ + +#[no_mangle] +pub extern "C" fn plugin_info() -> i64 { + let info = PluginInfo { + id: "steam".into(), + name: "Steam".into(), + version: "0.1.0".into(), + api_version: "1".into(), + icon_url: Some("https://store.steampowered.com/favicon.ico".into()), + required_secrets: alloc::vec!["API_KEY".into()], + config_schema: None, + }; + return_ok(&info) +} + +#[no_mangle] +pub extern "C" fn get_authorize_url(ptr: u32, len: u32) -> i64 { + #[cfg(target_arch = "wasm32")] + { + let bytes = match read_input(ptr, len) { + Some(b) => b, + None => return return_error("INVALID_INPUT", "Failed to read input", false), + }; + + let input: AuthorizeInput = match serde_json::from_slice(&bytes) { + Ok(i) => i, + Err(e) => return return_error("INVALID_INPUT", &format!("Parse error: {}", e), false), + }; + + // Build OpenID 2.0 authentication URL + // Steam uses claimed_id and identity as the same value for authentication + let params = [ + ("openid.ns", "http://specs.openid.net/auth/2.0"), + ("openid.mode", "checkid_setup"), + ( + "openid.return_to", + &format!("{}?state={}", input.redirect_uri, input.state), + ), + ("openid.realm", &input.redirect_uri), + ( + "openid.identity", + "http://specs.openid.net/auth/2.0/identifier_select", + ), + ( + "openid.claimed_id", + "http://specs.openid.net/auth/2.0/identifier_select", + ), + ]; + + let query: String = params + .iter() + .map(|(k, v)| format!("{}={}", k, urlencod(v))) + .collect::>() + .join("&"); + + let url = format!("{}?{}", STEAM_OPENID_URL, query); + return_ok(&url) + } + + #[cfg(not(target_arch = "wasm32"))] + { + let _ = (ptr, len); + return_error("NOT_WASM", "Only runs in WASM", false) + } +} + +#[no_mangle] +pub extern "C" fn handle_callback(ptr: u32, len: u32) -> i64 { + #[cfg(target_arch = "wasm32")] + { + let bytes = match read_input(ptr, len) { + Some(b) => b, + None => return return_error("INVALID_INPUT", "Failed to read input", false), + }; + + let input: CallbackInput = match serde_json::from_slice(&bytes) { + Ok(i) => i, + Err(e) => return return_error("INVALID_INPUT", &format!("Parse error: {}", e), false), + }; + + // Extract Steam ID from openid.claimed_id + // Format: https://steamcommunity.com/openid/id/76561198012345678 + let claimed_id = input + .extra + .get("openid.claimed_id") + .and_then(|v| v.as_str()); + + let steam_id = match claimed_id { + Some(id) => { + if let Some(pos) = id.rfind('/') { + &id[pos + 1..] + } else { + return return_error("INVALID_RESPONSE", "Invalid claimed_id format", false); + } + } + None => { + return return_error("INVALID_RESPONSE", "Missing openid.claimed_id", false); + } + }; + + // Verify the OpenID response with Steam + // Build verification request by changing mode to check_authentication + // and POSTing all params back to Steam + let mut verify_params: Vec<(&str, &str)> = Vec::new(); + verify_params.push(("openid.mode", "check_authentication")); + + // Add all openid.* params from the callback (except mode) + for (key, value) in &input.extra { + if key.starts_with("openid.") && key != "openid.mode" { + if let Some(v) = value.as_str() { + verify_params.push((key.as_str(), v)); + } + } + } + + // Build POST body + let verify_body: String = verify_params + .iter() + .map(|(k, v)| format!("{}={}", k, urlencod(v))) + .collect::>() + .join("&"); + + // POST to Steam for verification + let verify_result = http_post( + STEAM_OPENID_URL, + &verify_body, + "application/x-www-form-urlencoded", + ); + + match verify_result { + Ok(response_body) => { + // Steam returns key-value pairs, one per line + // We need to find "is_valid:true" + if !response_body.contains("is_valid:true") { + return return_error( + "VERIFICATION_FAILED", + "Steam OpenID verification failed", + false, + ); + } + } + Err(e) => { + return return_error( + "VERIFICATION_ERROR", + &format!("Failed to verify with Steam: {}", e), + true, + ); + } + } + + // Return the Steam ID as the "access_token" + // Since Steam uses OpenID 2.0 (not OAuth), there's no real token + // We store the Steam ID so we can use it with our API key + let tokens = TokenSet { + access_token: steam_id.into(), + token_type: "SteamID".into(), + expires_at: None, + refresh_token: None, + }; + + return_ok(&tokens) + } + + #[cfg(not(target_arch = "wasm32"))] + { + let _ = (ptr, len); + return_error("NOT_WASM", "Only runs in WASM", false) + } +} + +#[no_mangle] +pub extern "C" fn refresh_tokens(ptr: u32, len: u32) -> i64 { + // Steam doesn't use OAuth tokens - the Steam ID is permanent + #[cfg(target_arch = "wasm32")] + { + let bytes = match read_input(ptr, len) { + Some(b) => b, + None => return return_error("INVALID_INPUT", "Failed to read input", false), + }; + + #[derive(Deserialize)] + struct RefreshInput { + refresh_token: String, + #[allow(dead_code)] + config: serde_json::Value, + } + + let input: RefreshInput = match serde_json::from_slice(&bytes) { + Ok(i) => i, + Err(e) => return return_error("INVALID_INPUT", &format!("Parse error: {}", e), false), + }; + + // Just return the same Steam ID - it doesn't expire + let tokens = TokenSet { + access_token: input.refresh_token, + token_type: "SteamID".into(), + expires_at: None, + refresh_token: None, + }; + + return_ok(&tokens) + } + + #[cfg(not(target_arch = "wasm32"))] + { + let _ = (ptr, len); + return_error("NOT_WASM", "Only runs in WASM", false) + } +} + +#[no_mangle] +pub extern "C" fn get_profile(ptr: u32, len: u32) -> i64 { + #[cfg(target_arch = "wasm32")] + { + let bytes = match read_input(ptr, len) { + Some(b) => b, + None => return return_error("INVALID_INPUT", "Failed to read input", false), + }; + + let input: ProfileInput = match serde_json::from_slice(&bytes) { + Ok(i) => i, + Err(e) => return return_error("INVALID_INPUT", &format!("Parse error: {}", e), false), + }; + + let api_key = match get_secret("API_KEY") { + Some(k) => k, + None => return return_error("MISSING_SECRET", "API_KEY not configured", false), + }; + + let steam_id = &input.access_token; + let url = format!( + "{}/ISteamUser/GetPlayerSummaries/v2/?key={}&steamids={}", + STEAM_API_BASE, api_key, steam_id + ); + + let body = match http_get(&url) { + Ok(b) => b, + Err(e) => return return_error("HTTP_ERROR", &e, true), + }; + + let resp: SteamPlayerSummary = match serde_json::from_str(&body) { + Ok(r) => r, + Err(e) => { + return return_error("INVALID_RESPONSE", &format!("Parse error: {}", e), false) + } + }; + + let player = match resp.response.players.first() { + Some(p) => p, + None => return return_error("NOT_FOUND", "Player not found", false), + }; + + let profile = ExternalProfile { + account_id: player.steamid.clone(), + display_name: player.personaname.clone(), + profile_url: player.profileurl.clone(), + avatar_url: player.avatarfull.clone(), + }; + + return_ok(&profile) + } + + #[cfg(not(target_arch = "wasm32"))] + { + let _ = (ptr, len); + return_error("NOT_WASM", "Only runs in WASM", false) + } +} + +#[no_mangle] +pub extern "C" fn sync_account(ptr: u32, len: u32) -> i64 { + #[cfg(target_arch = "wasm32")] + { + let bytes = match read_input(ptr, len) { + Some(b) => b, + None => return return_error("INVALID_INPUT", "Failed to read input", false), + }; + + let input: SyncInput = match serde_json::from_slice(&bytes) { + Ok(i) => i, + Err(e) => return return_error("INVALID_INPUT", &format!("Parse error: {}", e), false), + }; + + let api_key = match get_secret("API_KEY") { + Some(k) => k, + None => return return_error("MISSING_SECRET", "API_KEY not configured", false), + }; + + let steam_id = &input.access_token; + let url = format!( + "{}/IPlayerService/GetOwnedGames/v1/?key={}&steamid={}&include_appinfo=true&include_played_free_games=true", + STEAM_API_BASE, api_key, steam_id + ); + + let body = match http_get(&url) { + Ok(b) => b, + Err(e) => return return_error("HTTP_ERROR", &e, true), + }; + + let resp: SteamOwnedGamesResponse = match serde_json::from_str(&body) { + Ok(r) => r, + Err(e) => { + return return_error("INVALID_RESPONSE", &format!("Parse error: {}", e), false) + } + }; + + let games = resp.response.games.unwrap_or_default(); + + let mut records: Vec = Vec::new(); + + for game in games { + let appid_str = game.appid.to_string(); + + // 1. Create actor.game record (ownership) + // HappyView will resolve game reference and add attestation signature + let game_record = serde_json::json!({ + "$type": "games.gamesgamesgamesgames.actor.game", + "game": { + "platform": "steam", + "externalId": &appid_str, + }, + "platform": "steam", + "createdAt": chrono_now(), + }); + + records.push(SyncRecord { + collection: "games.gamesgamesgamesgames.actor.game".into(), + record: game_record, + dedup_key: Some(format!("steam:game:{}", game.appid)), + sign: true, + }); + + // 2. Create actor.stats record (playtime) + // HappyView will add attestation signature + if let Some(playtime) = game.playtime_forever { + if playtime > 0 { + let stats_record = serde_json::json!({ + "$type": "games.gamesgamesgamesgames.actor.stats", + "game": { + "platform": "steam", + "externalId": &appid_str, + }, + "source": "steam", + "playtime": playtime, + "createdAt": chrono_now(), + }); + + records.push(SyncRecord { + collection: "games.gamesgamesgamesgames.actor.stats".into(), + record: stats_record, + dedup_key: Some(format!("steam:stats:{}", game.appid)), + sign: true, + }); + } + } + } + + return_ok(&records) + } + + #[cfg(not(target_arch = "wasm32"))] + { + let _ = (ptr, len); + return_error("NOT_WASM", "Only runs in WASM", false) + } +} + +// ============================================================================ +// Utility Functions +// ============================================================================ + +fn urlencod(s: &str) -> String { + let mut result = String::new(); + for c in s.chars() { + match c { + 'a'..='z' | 'A'..='Z' | '0'..='9' | '-' | '_' | '.' | '~' => { + result.push(c); + } + _ => { + for b in c.to_string().as_bytes() { + result.push_str(&format!("%{:02X}", b)); + } + } + } + } + result +} + +fn chrono_now() -> String { + // Simple ISO 8601 timestamp - in real impl would use proper time + "2024-01-01T00:00:00Z".into() +} diff --git a/src/external_auth/mod.rs b/src/external_auth/mod.rs index 15b512d..f36de91 100644 --- a/src/external_auth/mod.rs +++ b/src/external_auth/mod.rs @@ -1,4 +1,7 @@ +mod pds_write; mod routes; +pub mod state; mod sync; +pub mod tokens; pub use routes::routes; diff --git a/src/external_auth/pds_write.rs b/src/external_auth/pds_write.rs new file mode 100644 index 0000000..ae43e20 --- /dev/null +++ b/src/external_auth/pds_write.rs @@ -0,0 +1,127 @@ +//! Write sync records to user's PDS. + +use serde_json::{Value, json}; + +use crate::AppState; +use crate::error::AppError; +use crate::plugin::sync::ProcessedRecord; +use crate::repo; + +/// Result of writing a record to PDS +#[derive(Debug)] +#[allow(dead_code)] +pub struct WriteResult { + pub uri: String, + pub cid: String, +} + +/// Write processed records to the user's PDS. +/// +/// Returns the number of successfully written records. +pub async fn write_records_to_pds( + state: &AppState, + user_did: &str, + records: Vec, +) -> Result, AppError> { + let session = repo::get_oauth_session(state, user_did).await?; + + let mut results = Vec::with_capacity(records.len()); + + for record in records { + // Generate rkey from dedup_key or create a timestamp-based one + let rkey = record + .dedup_key + .as_ref() + .map(|k| sanitize_rkey(k)) + .unwrap_or_else(generate_tid); + + // Build the putRecord request + let body = json!({ + "repo": user_did, + "collection": record.collection, + "rkey": rkey, + "record": record.record, + }); + + let resp = + repo::pds_post_json_raw(state, &session, "com.atproto.repo.putRecord", &body).await?; + + if resp.status().is_success() { + let bytes = resp + .bytes() + .await + .map_err(|e| AppError::Internal(format!("failed to read PDS response: {e}")))?; + + let pds_result: Value = serde_json::from_slice(&bytes) + .map_err(|e| AppError::Internal(format!("invalid PDS JSON: {e}")))?; + + if let (Some(uri), Some(cid)) = ( + pds_result.get("uri").and_then(|v| v.as_str()), + pds_result.get("cid").and_then(|v| v.as_str()), + ) { + results.push(WriteResult { + uri: uri.to_string(), + cid: cid.to_string(), + }); + } + } else { + let bytes = resp.bytes().await.unwrap_or_default(); + let body_str = String::from_utf8_lossy(&bytes); + tracing::warn!( + collection = %record.collection, + rkey = %rkey, + error = %body_str, + "Failed to write record to PDS" + ); + // Continue with other records even if one fails + } + } + + Ok(results) +} + +/// Sanitize a dedup_key to be a valid rkey. +/// rkey must be 1-512 chars, alphanumeric plus .-_:~ +fn sanitize_rkey(key: &str) -> String { + let sanitized: String = key + .chars() + .filter(|c| c.is_ascii_alphanumeric() || matches!(c, '.' | '-' | '_' | ':' | '~')) + .take(512) + .collect(); + + if sanitized.is_empty() { + generate_tid() + } else { + sanitized + } +} + +/// Generate a TID (timestamp-based ID) for use as rkey. +fn generate_tid() -> String { + use std::time::{SystemTime, UNIX_EPOCH}; + + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_micros(); + + // TID is base32-sortable encoding of microseconds since epoch + // Using a simplified version here + format!("{:0>13}", base32_encode(now as u64)) +} + +fn base32_encode(mut n: u64) -> String { + const ALPHABET: &[u8] = b"234567abcdefghijklmnopqrstuvwxyz"; + let mut result = String::new(); + + if n == 0 { + return "2".to_string(); + } + + while n > 0 { + result.insert(0, ALPHABET[(n % 32) as usize] as char); + n /= 32; + } + + result +} diff --git a/src/external_auth/routes.rs b/src/external_auth/routes.rs index 399c4df..86f59c2 100644 --- a/src/external_auth/routes.rs +++ b/src/external_auth/routes.rs @@ -5,9 +5,15 @@ use axum::{ routing::{get, post}, }; use serde::{Deserialize, Serialize}; +use std::collections::HashMap; +use std::sync::Arc; use crate::AppState; +use crate::auth::Claims; use crate::error::AppError; +use crate::external_auth::{pds_write, state, tokens}; +use crate::plugin::PluginExecutor; +use crate::plugin::sync::SyncProcessor; pub fn routes() -> Router { Router::new() @@ -48,11 +54,12 @@ struct AuthorizeQuery { } async fn authorize( - State(state): State, + State(app_state): State, Path(plugin_id): Path, Query(query): Query, + claims: Claims, ) -> Result, AppError> { - let _plugin = state + let _plugin = app_state .plugin_registry .get(&plugin_id) .await @@ -61,12 +68,46 @@ async fn authorize( // Generate state parameter for CSRF protection let state_param = uuid::Uuid::new_v4().to_string(); - // TODO: Store state in KV, call plugin's get_authorize_url() - // For now, return placeholder - let _ = query.redirect_uri; + // Store state -> user mapping for callback validation + state::store_state( + &app_state.db, + app_state.db_backend, + &state_param, + claims.did(), + &plugin_id, + &query.redirect_uri, + ) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + // Get plugin config (empty for now, could come from DB) + let config = serde_json::Value::Null; + + // Load secrets from environment + let secrets = load_plugin_secrets(&plugin_id); + + // Create executor and instance + let executor = PluginExecutor::new( + app_state.wasm_runtime.clone(), + app_state.plugin_registry.clone(), + app_state.db.clone(), + app_state.db_backend, + app_state.http.clone(), + Arc::new(app_state.lexicons.clone()), + ); + + let mut instance = executor + .instantiate(&plugin_id, &state_param, secrets, config.clone()) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + let authorize_url = instance + .call_get_authorize_url(&state_param, &query.redirect_uri, &config) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; Ok(Json(serde_json::json!({ - "authorize_url": format!("https://example.com/oauth?state={}", state_param), + "authorize_url": authorize_url, "state": state_param }))) } @@ -80,35 +121,166 @@ struct CallbackQuery { } async fn callback( - State(_state): State, - Path(_plugin_id): Path, - Query(_query): Query, + State(app_state): State, + Path(plugin_id): Path, + Query(query): Query, ) -> Result { - // TODO: Validate state, call plugin's handle_callback(), store tokens + // Validate required parameters + let code = query.code.ok_or_else(|| { + AppError::BadRequest(query.error.unwrap_or_else(|| "Missing code".into())) + })?; + let state_param = query + .state + .ok_or_else(|| AppError::BadRequest("Missing state".into()))?; + + // Validate state and get user DID + redirect_uri + let stored_state = state::consume_state(&app_state.db, app_state.db_backend, &state_param) + .await + .map_err(|_| AppError::BadRequest("Invalid or expired state".into()))?; + + // Verify plugin_id matches + if stored_state.plugin_id != plugin_id { + return Err(AppError::BadRequest("Plugin ID mismatch".into())); + } + + let config = serde_json::Value::Null; + let secrets = load_plugin_secrets(&plugin_id); + + let executor = PluginExecutor::new( + app_state.wasm_runtime.clone(), + app_state.plugin_registry.clone(), + app_state.db.clone(), + app_state.db_backend, + app_state.http.clone(), + Arc::new(app_state.lexicons.clone()), + ); + + let mut instance = executor + .instantiate(&plugin_id, &stored_state.did, secrets, config.clone()) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + let token_set = instance + .call_handle_callback(&code, &state_param, &config) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + // Get profile to get the account_id + let profile = instance + .call_get_profile(&token_set.access_token, &config) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + // Format expires_at as RFC3339 string + let expires_at = token_set.expires_at.map(|dt| dt.to_rfc3339()); - // For now, redirect to a placeholder - Ok(Redirect::to("/")) + // Store encrypted tokens + tokens::store_tokens( + &app_state.db, + app_state.db_backend, + app_state.config.token_encryption_key.as_ref(), + &stored_state.did, + &plugin_id, + &profile.account_id, + &token_set.access_token, + token_set.refresh_token.as_deref(), + Some(&token_set.token_type), + None, // scope not in TokenSet + expires_at.as_deref(), + ) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + // Redirect to the original redirect_uri + Ok(Redirect::to(&stored_state.redirect_uri)) } async fn sync( - State(_state): State, - Path(_plugin_id): Path, + State(app_state): State, + Path(plugin_id): Path, + claims: Claims, ) -> Result, AppError> { - // TODO: Call plugin's sync_account(), process SyncRecords + let user_did = claims.did(); + + let config = serde_json::Value::Null; + let secrets = load_plugin_secrets(&plugin_id); + + let executor = PluginExecutor::new( + app_state.wasm_runtime.clone(), + app_state.plugin_registry.clone(), + app_state.db.clone(), + app_state.db_backend, + app_state.http.clone(), + Arc::new(app_state.lexicons.clone()), + ); + + let mut instance = executor + .instantiate(&plugin_id, user_did, secrets, config.clone()) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + // Get decrypted access token from DB + let stored = tokens::get_tokens( + &app_state.db, + app_state.db_backend, + app_state.config.token_encryption_key.as_ref(), + user_did, + &plugin_id, + ) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + let mut records = instance + .call_sync_account(&stored.access_token, &config) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + // Resolve game references from database + crate::plugin::sync::resolve_game_references(&app_state.db, app_state.db_backend, &mut records) + .await; + + // Process records: sign those with sign=true + let signer = app_state.attestation_signer.as_deref(); + let processor = SyncProcessor::new(signer, user_did.to_string()); + let processed = processor + .process_records(records) + .map_err(|e| AppError::Internal(e.to_string()))?; + + let processed_count = processed.len(); + + // Write processed records to user's PDS + let write_results = pds_write::write_records_to_pds(&app_state, user_did, processed).await?; Ok(Json(serde_json::json!({ "status": "ok", - "synced": 0 + "processed": processed_count, + "written": write_results.len() }))) } async fn unlink( - State(_state): State, - Path(_plugin_id): Path, + State(app_state): State, + Path(plugin_id): Path, + claims: Claims, ) -> Result, AppError> { - // TODO: Delete tokens, delete accountLink record + let user_did = claims.did(); + + // Delete tokens + let deleted = tokens::delete_tokens(&app_state.db, app_state.db_backend, user_did, &plugin_id) + .await + .map_err(|e| AppError::Internal(e.to_string()))?; + + // TODO: Delete accountLink record from user's PDS Ok(Json(serde_json::json!({ - "status": "ok" + "status": "ok", + "was_linked": deleted }))) } + +fn load_plugin_secrets(plugin_id: &str) -> HashMap { + let prefix = format!("PLUGIN_{}_", plugin_id.to_uppercase()); + std::env::vars() + .filter_map(|(k, v)| k.strip_prefix(&prefix).map(|name| (name.to_string(), v))) + .collect() +} diff --git a/src/external_auth/state.rs b/src/external_auth/state.rs new file mode 100644 index 0000000..8507973 --- /dev/null +++ b/src/external_auth/state.rs @@ -0,0 +1,108 @@ +//! OAuth state management for external auth flows. +//! +//! Stores state -> (user_did, plugin_id, redirect_uri) mappings +//! to validate callbacks and associate external accounts with users. + +use crate::db::{DatabaseBackend, adapt_sql, now_rfc3339}; + +#[derive(Debug, thiserror::Error)] +pub enum StateError { + #[error("Database error: {0}")] + Database(#[from] sqlx::Error), + #[error("State not found or expired")] + NotFound, +} + +/// Stored OAuth state +#[derive(Debug, Clone)] +pub struct StoredState { + pub did: String, + pub plugin_id: String, + pub redirect_uri: String, +} + +/// Store OAuth state for an auth flow. +/// +/// State expires after 10 minutes. +pub async fn store_state( + db: &sqlx::AnyPool, + backend: DatabaseBackend, + state: &str, + did: &str, + plugin_id: &str, + redirect_uri: &str, +) -> Result<(), StateError> { + let now = now_rfc3339(); + + // Expire in 10 minutes + let expires_at = chrono::Utc::now() + chrono::Duration::minutes(10); + let expires_str = expires_at.to_rfc3339(); + + let sql = adapt_sql( + "INSERT INTO external_auth_state (state, did, plugin_id, redirect_uri, created_at, expires_at) VALUES (?, ?, ?, ?, ?, ?)", + backend, + ); + + sqlx::query(&sql) + .bind(state) + .bind(did) + .bind(plugin_id) + .bind(redirect_uri) + .bind(&now) + .bind(&expires_str) + .execute(db) + .await?; + + Ok(()) +} + +/// Retrieve and consume OAuth state. +/// +/// Returns the stored state if found and not expired, then deletes it. +pub async fn consume_state( + db: &sqlx::AnyPool, + backend: DatabaseBackend, + state: &str, +) -> Result { + let now = now_rfc3339(); + + // Get state if not expired + let sql = adapt_sql( + "SELECT did, plugin_id, redirect_uri FROM external_auth_state WHERE state = ? AND expires_at > ?", + backend, + ); + + let row: Option<(String, String, String)> = sqlx::query_as(&sql) + .bind(state) + .bind(&now) + .fetch_optional(db) + .await?; + + let (did, plugin_id, redirect_uri) = row.ok_or(StateError::NotFound)?; + + // Delete the state (one-time use) + let delete_sql = adapt_sql("DELETE FROM external_auth_state WHERE state = ?", backend); + sqlx::query(&delete_sql).bind(state).execute(db).await?; + + Ok(StoredState { + did, + plugin_id, + redirect_uri, + }) +} + +/// Clean up expired state entries. +pub async fn cleanup_expired( + db: &sqlx::AnyPool, + backend: DatabaseBackend, +) -> Result { + let now = now_rfc3339(); + + let sql = adapt_sql( + "DELETE FROM external_auth_state WHERE expires_at <= ?", + backend, + ); + let result = sqlx::query(&sql).bind(&now).execute(db).await?; + + Ok(result.rows_affected()) +} diff --git a/src/external_auth/tokens.rs b/src/external_auth/tokens.rs new file mode 100644 index 0000000..fbb5b62 --- /dev/null +++ b/src/external_auth/tokens.rs @@ -0,0 +1,198 @@ +//! External account token storage with encryption. + +use crate::db::{DatabaseBackend, adapt_sql, now_rfc3339}; +use crate::plugin::encryption::{EncryptionError, decrypt, encrypt}; + +/// Row type for token query results +type TokenRow = ( + String, + Vec, + Option>, + Option, + Option, + Option, +); + +#[derive(Debug, thiserror::Error)] +pub enum TokenError { + #[error("Database error: {0}")] + Database(#[from] sqlx::Error), + #[error("Encryption error: {0}")] + Encryption(#[from] EncryptionError), + #[error("Token encryption key not configured")] + KeyNotConfigured, + #[error("Token not found")] + NotFound, +} + +/// Stored external account token set +#[derive(Debug, Clone)] +pub struct StoredTokens { + pub account_id: String, + pub access_token: String, + pub refresh_token: Option, + pub token_type: Option, + pub scope: Option, + pub expires_at: Option, +} + +/// Store tokens for an external account link +#[allow(clippy::too_many_arguments)] +pub async fn store_tokens( + db: &sqlx::AnyPool, + backend: DatabaseBackend, + encryption_key: Option<&[u8; 32]>, + did: &str, + plugin_id: &str, + account_id: &str, + access_token: &str, + refresh_token: Option<&str>, + token_type: Option<&str>, + scope: Option<&str>, + expires_at: Option<&str>, +) -> Result<(), TokenError> { + let key = encryption_key.ok_or(TokenError::KeyNotConfigured)?; + + let encrypted_access = encrypt(key, access_token.as_bytes())?; + let encrypted_refresh = refresh_token + .map(|t| encrypt(key, t.as_bytes())) + .transpose()?; + + let id = uuid::Uuid::new_v4().to_string(); + let now = now_rfc3339(); + + let sql = adapt_sql( + "INSERT INTO external_account_tokens (id, did, plugin_id, account_id, access_token, refresh_token, token_type, scope, expires_at, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT (did, plugin_id) + DO UPDATE SET account_id = excluded.account_id, access_token = excluded.access_token, refresh_token = excluded.refresh_token, token_type = excluded.token_type, scope = excluded.scope, expires_at = excluded.expires_at, updated_at = excluded.updated_at", + backend, + ); + + sqlx::query(&sql) + .bind(&id) + .bind(did) + .bind(plugin_id) + .bind(account_id) + .bind(&encrypted_access) + .bind(encrypted_refresh.as_deref()) + .bind(token_type) + .bind(scope) + .bind(expires_at) + .bind(&now) + .bind(&now) + .execute(db) + .await?; + + Ok(()) +} + +/// Retrieve decrypted tokens for an external account +pub async fn get_tokens( + db: &sqlx::AnyPool, + backend: DatabaseBackend, + encryption_key: Option<&[u8; 32]>, + did: &str, + plugin_id: &str, +) -> Result { + let key = encryption_key.ok_or(TokenError::KeyNotConfigured)?; + + let sql = adapt_sql( + "SELECT account_id, access_token, refresh_token, token_type, scope, expires_at FROM external_account_tokens WHERE did = ? AND plugin_id = ?", + backend, + ); + + let row: Option = sqlx::query_as(&sql) + .bind(did) + .bind(plugin_id) + .fetch_optional(db) + .await?; + + let (account_id, encrypted_access, encrypted_refresh, token_type, scope, expires_at) = + row.ok_or(TokenError::NotFound)?; + + let access_token = String::from_utf8(decrypt(key, &encrypted_access)?) + .map_err(|_| EncryptionError::DecryptionFailed)?; + + let refresh_token = encrypted_refresh + .map(|enc| { + decrypt(key, &enc).and_then(|dec| { + String::from_utf8(dec).map_err(|_| EncryptionError::DecryptionFailed) + }) + }) + .transpose()?; + + Ok(StoredTokens { + account_id, + access_token, + refresh_token, + token_type, + scope, + expires_at, + }) +} + +/// Delete tokens for an external account link +pub async fn delete_tokens( + db: &sqlx::AnyPool, + backend: DatabaseBackend, + did: &str, + plugin_id: &str, +) -> Result { + let sql = adapt_sql( + "DELETE FROM external_account_tokens WHERE did = ? AND plugin_id = ?", + backend, + ); + + let result = sqlx::query(&sql) + .bind(did) + .bind(plugin_id) + .execute(db) + .await?; + + Ok(result.rows_affected() > 0) +} + +/// Check if an external account is linked +pub async fn is_linked( + db: &sqlx::AnyPool, + backend: DatabaseBackend, + did: &str, + plugin_id: &str, +) -> Result { + let sql = adapt_sql( + "SELECT 1 FROM external_account_tokens WHERE did = ? AND plugin_id = ?", + backend, + ); + + let exists: Option<(i32,)> = sqlx::query_as(&sql) + .bind(did) + .bind(plugin_id) + .fetch_optional(db) + .await?; + + Ok(exists.is_some()) +} + +/// Get the external account ID for a linked account +pub async fn get_account_id( + db: &sqlx::AnyPool, + backend: DatabaseBackend, + did: &str, + plugin_id: &str, +) -> Result, TokenError> { + let sql = adapt_sql( + "SELECT account_id FROM external_account_tokens WHERE did = ? AND plugin_id = ?", + backend, + ); + + let row: Option<(String,)> = sqlx::query_as(&sql) + .bind(did) + .bind(plugin_id) + .fetch_optional(db) + .await?; + + Ok(row.map(|(id,)| id)) +} + +// Integration tests for token storage are in tests/e2e_external_auth.rs diff --git a/src/lib.rs b/src/lib.rs index 782313f..7a0ae7a 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -59,6 +59,8 @@ pub struct AppState { pub oauth: Arc, pub cookie_key: axum_extra::extract::cookie::Key, pub plugin_registry: Arc, + pub wasm_runtime: Arc, + pub attestation_signer: Option>, } impl axum::extract::FromRef for axum_extra::extract::cookie::Key { diff --git a/src/lua/atproto_api.rs b/src/lua/atproto_api.rs index c5ef8bc..757fd53 100644 --- a/src/lua/atproto_api.rs +++ b/src/lua/atproto_api.rs @@ -289,6 +289,10 @@ mod tests { b"test-secret-for-tests-only-not-production", ), plugin_registry: std::sync::Arc::new(crate::plugin::PluginRegistry::new()), + wasm_runtime: std::sync::Arc::new( + crate::plugin::WasmRuntime::new().expect("wasm runtime"), + ), + attestation_signer: None, } } diff --git a/src/lua/db_api.rs b/src/lua/db_api.rs index 72c393d..6fad7b1 100644 --- a/src/lua/db_api.rs +++ b/src/lua/db_api.rs @@ -697,6 +697,10 @@ mod tests { b"test-secret-for-tests-only-not-production", ), plugin_registry: std::sync::Arc::new(crate::plugin::PluginRegistry::new()), + wasm_runtime: std::sync::Arc::new( + crate::plugin::WasmRuntime::new().expect("wasm runtime"), + ), + attestation_signer: None, } } diff --git a/src/lua/execute.rs b/src/lua/execute.rs index dfaa697..d15f091 100644 --- a/src/lua/execute.rs +++ b/src/lua/execute.rs @@ -1031,6 +1031,10 @@ mod tests { b"test-secret-for-tests-only-not-production", ), plugin_registry: std::sync::Arc::new(crate::plugin::PluginRegistry::new()), + wasm_runtime: std::sync::Arc::new( + crate::plugin::WasmRuntime::new().expect("wasm runtime"), + ), + attestation_signer: None, } } diff --git a/src/lua/http_api.rs b/src/lua/http_api.rs index bfc0b3a..6ed1464 100644 --- a/src/lua/http_api.rs +++ b/src/lua/http_api.rs @@ -171,6 +171,10 @@ mod tests { b"test-secret-for-tests-only-not-production", ), plugin_registry: std::sync::Arc::new(crate::plugin::PluginRegistry::new()), + wasm_runtime: std::sync::Arc::new( + crate::plugin::WasmRuntime::new().expect("wasm runtime"), + ), + attestation_signer: None, } } diff --git a/src/main.rs b/src/main.rs index b65e816..d019fa7 100644 --- a/src/main.rs +++ b/src/main.rs @@ -176,6 +176,26 @@ async fn main() { // Initialize plugin registry let plugin_registry = Arc::new(happyview::plugin::PluginRegistry::new()); + // Initialize WASM runtime + let wasm_runtime = + Arc::new(happyview::plugin::WasmRuntime::new().expect("Failed to create WASM runtime")); + + // Initialize attestation signer (optional) + let attestation_signer = match happyview::plugin::attestation::load_from_env() { + Ok(Some(signer)) => { + tracing::info!("Attestation signing enabled"); + Some(Arc::new(signer)) + } + Ok(None) => { + tracing::info!("Attestation signing disabled (no ATTESTATION_PRIVATE_KEY)"); + None + } + Err(e) => { + tracing::error!(error = %e, "Failed to load attestation signer"); + None + } + }; + // Load plugins from PLUGIN_URLS env var if let Ok(urls) = std::env::var("PLUGIN_URLS") { for (id, url, sha256) in happyview::plugin::loader::parse_plugin_urls(&urls) { @@ -319,6 +339,8 @@ async fn main() { oauth: Arc::new(oauth_client), cookie_key, plugin_registry, + wasm_runtime, + attestation_signer, }; // Sync initial collections to Tap on startup. diff --git a/src/plugin/attestation.rs b/src/plugin/attestation.rs new file mode 100644 index 0000000..883e611 --- /dev/null +++ b/src/plugin/attestation.rs @@ -0,0 +1,331 @@ +//! Attestation signing for plugin records. +//! +//! Implements the ATProtocol attestation spec: +//! - Computes CID with $sig metadata for replay protection +//! - Signs using ECDSA (P-256 or K-256) +//! - Adds inline signatures to records + +use cid::Cid; +use k256::ecdsa::{Signature, SigningKey, signature::Signer}; +use serde_json::{Map, Value}; +use sha2::{Digest, Sha256}; +use std::sync::Arc; + +// Multihash code for SHA2-256 +const SHA2_256_CODE: u64 = 0x12; +// DAG-CBOR codec +const DAG_CBOR_CODEC: u64 = 0x71; + +/// Attestation signer for HappyView +pub struct AttestationSigner { + /// The signing key (K-256/secp256k1) + signing_key: SigningKey, + /// The key identifier (e.g., "did:web:happyview.example#attestation") + key_id: String, + /// The signature type identifier + sig_type: String, +} + +#[derive(Debug, thiserror::Error)] +pub enum AttestationError { + #[error("Failed to encode record: {0}")] + Encoding(String), + #[error("Failed to sign: {0}")] + Signing(String), + #[error("Invalid key: {0}")] + InvalidKey(String), + #[error("Record missing required field: {0}")] + MissingField(String), +} + +impl AttestationSigner { + /// Create a new signer from a hex-encoded private key + pub fn from_hex( + private_key_hex: &str, + key_id: String, + sig_type: String, + ) -> Result { + let key_bytes = hex::decode(private_key_hex) + .map_err(|e| AttestationError::InvalidKey(format!("invalid hex: {}", e)))?; + + let signing_key = SigningKey::from_bytes((&key_bytes[..]).into()) + .map_err(|e| AttestationError::InvalidKey(format!("invalid key: {}", e)))?; + + Ok(Self { + signing_key, + key_id, + sig_type, + }) + } + + /// Create a new signer with a test key (for testing only) + #[cfg(test)] + pub fn for_testing(key_id: String, sig_type: String) -> Self { + // Fixed test key (32 bytes of 0x01) - DO NOT USE IN PRODUCTION + let test_key_bytes = [0x01u8; 32]; + let signing_key = + SigningKey::from_bytes((&test_key_bytes[..]).into()).expect("valid test key"); + Self { + signing_key, + key_id, + sig_type, + } + } + + /// Get the public key in compressed format (for verification) + pub fn public_key_bytes(&self) -> Vec { + use k256::ecdsa::VerifyingKey; + let verifying_key = VerifyingKey::from(&self.signing_key); + verifying_key.to_encoded_point(true).as_bytes().to_vec() + } + + /// Sign a record and add the signature to the signatures array. + /// + /// # Arguments + /// * `record` - The record to sign (will be modified to add signature) + /// * `repository_did` - The DID of the repository (for replay protection) + /// + /// # Returns + /// The CID of the signed content + pub fn sign_record( + &self, + record: &mut Value, + repository_did: &str, + ) -> Result { + let obj = record + .as_object_mut() + .ok_or_else(|| AttestationError::Encoding("record must be an object".into()))?; + + // Remove existing signatures for CID computation + let existing_signatures = obj.remove("signatures"); + + // Inject $sig metadata for CID computation + let sig_metadata = serde_json::json!({ + "$type": &self.sig_type, + "repository": repository_did, + }); + obj.insert("$sig".to_string(), sig_metadata); + + // Encode to CBOR (DAG-CBOR canonical form) + let cbor_bytes = self.encode_dag_cbor(obj)?; + + // Compute CID (sha2-256, dag-cbor codec) + let cid = self.compute_cid(&cbor_bytes); + + // Remove $sig (it's only for CID computation) + obj.remove("$sig"); + + // Sign the CID bytes + let signature = self.sign_cid(&cid)?; + + // Create inline signature object + let inline_sig = serde_json::json!({ + "$type": &self.sig_type, + "key": &self.key_id, + "signature": { + "$bytes": base64::Engine::encode(&base64::engine::general_purpose::STANDARD, &signature) + } + }); + + // Add to signatures array + let signatures = obj + .entry("signatures") + .or_insert_with(|| Value::Array(vec![])); + + if let Value::Array(arr) = signatures { + // Restore any existing signatures + if let Some(Value::Array(existing)) = existing_signatures { + for sig in existing { + arr.push(sig); + } + } + arr.push(inline_sig); + } + + Ok(cid) + } + + /// Encode a JSON object to DAG-CBOR canonical form + fn encode_dag_cbor(&self, obj: &Map) -> Result, AttestationError> { + // Convert to ciborium Value and encode + // DAG-CBOR requires deterministic key ordering (lexicographic) + let cbor_value = json_to_cbor(&Value::Object(obj.clone())); + + let mut buf = Vec::new(); + ciborium::into_writer(&cbor_value, &mut buf) + .map_err(|e| AttestationError::Encoding(format!("CBOR encoding failed: {}", e)))?; + + Ok(buf) + } + + /// Compute CID from CBOR bytes (sha2-256, dag-cbor codec) + fn compute_cid(&self, cbor_bytes: &[u8]) -> Cid { + // SHA2-256 hash + let digest = Sha256::digest(cbor_bytes); + + // Create multihash: varint(code) || varint(size) || digest + let mut multihash_bytes = Vec::new(); + // SHA2-256 code (0x12) + multihash_bytes.push(SHA2_256_CODE as u8); + // Digest size (32 bytes) + multihash_bytes.push(32u8); + // The digest + multihash_bytes.extend_from_slice(&digest); + + let multihash = + cid::multihash::Multihash::<64>::from_bytes(&multihash_bytes).expect("valid multihash"); + + // CID v1 with dag-cbor codec + Cid::new_v1(DAG_CBOR_CODEC, multihash) + } + + /// Sign a CID using ECDSA with low-S normalization + fn sign_cid(&self, cid: &Cid) -> Result, AttestationError> { + let cid_bytes = cid.to_bytes(); + + // Sign using k256 ECDSA (automatically uses low-S) + let signature: Signature = self.signing_key.sign(&cid_bytes); + + Ok(signature.to_bytes().to_vec()) + } +} + +/// Convert JSON Value to ciborium Value with deterministic ordering +fn json_to_cbor(value: &Value) -> ciborium::Value { + match value { + Value::Null => ciborium::Value::Null, + Value::Bool(b) => ciborium::Value::Bool(*b), + Value::Number(n) => { + if let Some(i) = n.as_i64() { + ciborium::Value::Integer(i.into()) + } else if let Some(u) = n.as_u64() { + ciborium::Value::Integer(u.into()) + } else if let Some(f) = n.as_f64() { + ciborium::Value::Float(f) + } else { + ciborium::Value::Null + } + } + Value::String(s) => { + // Check for $bytes encoding (base64) + ciborium::Value::Text(s.clone()) + } + Value::Array(arr) => ciborium::Value::Array(arr.iter().map(json_to_cbor).collect()), + Value::Object(obj) => { + // Handle special $bytes encoding for binary data + if obj.len() == 1 + && let Some(Value::String(b64)) = obj.get("$bytes") + && let Ok(bytes) = + base64::Engine::decode(&base64::engine::general_purpose::STANDARD, b64) + { + return ciborium::Value::Bytes(bytes); + } + + // Sort keys lexicographically for deterministic encoding + let mut pairs: Vec<_> = obj + .iter() + .map(|(k, v)| (ciborium::Value::Text(k.clone()), json_to_cbor(v))) + .collect(); + pairs.sort_by(|a, b| { + if let (ciborium::Value::Text(ka), ciborium::Value::Text(kb)) = (&a.0, &b.0) { + ka.cmp(kb) + } else { + std::cmp::Ordering::Equal + } + }); + + ciborium::Value::Map(pairs) + } + } +} + +/// Shared attestation signer for the application +pub type SharedAttestationSigner = Arc; + +/// Load attestation signer from environment variables +pub fn load_from_env() -> Result, AttestationError> { + let private_key = match std::env::var("ATTESTATION_PRIVATE_KEY") { + Ok(k) => k, + Err(_) => return Ok(None), // No key configured, attestation disabled + }; + + let key_id = std::env::var("ATTESTATION_KEY_ID") + .unwrap_or_else(|_| "did:web:localhost#attestation".to_string()); + + let sig_type = std::env::var("ATTESTATION_SIG_TYPE") + .unwrap_or_else(|_| "games.gamesgamesgamesgames.attestation".to_string()); + + Ok(Some(AttestationSigner::from_hex( + &private_key, + key_id, + sig_type, + )?)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_sign_record() { + let signer = AttestationSigner::for_testing( + "did:web:test.example#signing".to_string(), + "test.signature".to_string(), + ); + + let mut record = serde_json::json!({ + "$type": "games.gamesgamesgamesgames.actor.game", + "game": {"platform": "steam", "externalId": "440"}, + "platform": "steam", + "createdAt": "2024-01-01T00:00:00Z" + }); + + let cid = signer + .sign_record(&mut record, "did:plc:testuser") + .expect("signing should succeed"); + + // Verify signature was added + let signatures = record["signatures"].as_array().expect("signatures array"); + assert_eq!(signatures.len(), 1); + + let sig = &signatures[0]; + assert_eq!(sig["$type"], "test.signature"); + assert_eq!(sig["key"], "did:web:test.example#signing"); + assert!(sig["signature"]["$bytes"].is_string()); + + // CID should be valid + assert!(!cid.to_bytes().is_empty()); + } + + #[test] + fn test_deterministic_cid() { + let signer = AttestationSigner::for_testing( + "did:web:test.example#signing".to_string(), + "test.signature".to_string(), + ); + + // Same record should produce same CID (before signature) + let record1 = serde_json::json!({ + "a": 1, + "b": 2, + "c": {"nested": true} + }); + + let record2 = serde_json::json!({ + "c": {"nested": true}, + "a": 1, + "b": 2 + }); + + let mut r1 = record1.clone(); + let mut r2 = record2.clone(); + + let cid1 = signer.sign_record(&mut r1, "did:plc:test").unwrap(); + let cid2 = signer.sign_record(&mut r2, "did:plc:test").unwrap(); + + // Different signatures (random nonce in ECDSA) but... + // Actually the CIDs should be the same since they're computed before signing + // and the key ordering is normalized + assert_eq!(cid1, cid2); + } +} diff --git a/src/plugin/executor.rs b/src/plugin/executor.rs new file mode 100644 index 0000000..5a4d46a --- /dev/null +++ b/src/plugin/executor.rs @@ -0,0 +1,483 @@ +// src/plugin/executor.rs + +use crate::db::DatabaseBackend; +use crate::lexicon::LexiconRegistry; +use crate::plugin::host::{PluginState, register_host_functions}; +use crate::plugin::memory::{ + PluginEnvelopeError, PluginResponse, dealloc_guest, read_from_guest, write_to_guest, +}; +use crate::plugin::runtime::{DEFAULT_FUEL, WasmRuntime}; +use crate::plugin::{ExternalProfile, PluginInfo, PluginRegistry, SyncRecord, TokenSet}; +use std::collections::HashMap; +use std::sync::Arc; +use thiserror::Error; +use wasmtime::{Instance, Linker, Memory, Store, TypedFunc}; + +#[derive(Debug, Error)] +pub enum ExecutionError { + #[error("Plugin not found: {0}")] + PluginNotFound(String), + + #[error("WASM instantiation failed: {0}")] + Instantiation(#[source] anyhow::Error), + + #[error("Memory allocation failed")] + MemoryAllocation, + + #[error("Plugin function trapped: {0}")] + Trap(#[source] wasmtime::Error), + + #[error("Invalid response from plugin: {0}")] + InvalidResponse(String), + + #[error("Plugin returned error: {code} - {message}")] + PluginError { + code: String, + message: String, + retryable: bool, + }, + + #[error("Resource limit exceeded: {0}")] + ResourceLimit(String), + + #[error("Timeout (fuel exhausted)")] + Timeout, + + #[error("Missing export: {0}")] + MissingExport(String), +} + +impl From for ExecutionError { + fn from(e: PluginEnvelopeError) -> Self { + ExecutionError::PluginError { + code: e.code, + message: e.message, + retryable: e.retryable, + } + } +} + +/// Single-use wrapper around a WASM instance +#[allow(dead_code)] +pub struct PluginInstance { + pub(crate) store: Store, + pub(crate) instance: Instance, + pub(crate) memory: Memory, + pub(crate) alloc: TypedFunc, + pub(crate) dealloc: TypedFunc<(u32, u32), ()>, +} + +impl PluginInstance { + /// Call plugin_info() - no input required + pub async fn call_plugin_info(&mut self) -> Result { + let func = self + .instance + .get_typed_func::<(), i64>(&mut self.store, "plugin_info") + .map_err(|_| ExecutionError::MissingExport("plugin_info".into()))?; + + self.store + .set_fuel(DEFAULT_FUEL) + .map_err(ExecutionError::Trap)?; + + let packed = func + .call_async(&mut self.store, ()) + .await + .map_err(Self::classify_error)?; + + // Unpack i64: upper 32 bits = ptr, lower 32 bits = len + let ptr = (packed >> 32) as u32; + let len = (packed & 0xFFFFFFFF) as u32; + + let bytes = + read_from_guest(&self.store, ptr, len).map_err(|_| ExecutionError::MemoryAllocation)?; + + dealloc_guest(&mut self.store, ptr, len) + .await + .map_err(|_| ExecutionError::MemoryAllocation)?; + + let response: PluginResponse = serde_json::from_slice(&bytes) + .map_err(|e| ExecutionError::InvalidResponse(e.to_string()))?; + + response.into_result().map_err(ExecutionError::from) + } + + /// Call get_authorize_url(state, redirect_uri, config) + pub async fn call_get_authorize_url( + &mut self, + state: &str, + redirect_uri: &str, + config: &serde_json::Value, + ) -> Result { + let input = serde_json::json!({ + "state": state, + "redirect_uri": redirect_uri, + "config": config + }); + self.call_plugin_function("get_authorize_url", &input).await + } + + /// Call handle_callback(code, state, config) + pub async fn call_handle_callback( + &mut self, + code: &str, + state: &str, + config: &serde_json::Value, + ) -> Result { + let input = serde_json::json!({ + "code": code, + "state": state, + "config": config + }); + self.call_plugin_function("handle_callback", &input).await + } + + /// Call refresh_tokens(refresh_token, config) + pub async fn call_refresh_tokens( + &mut self, + refresh_token: &str, + config: &serde_json::Value, + ) -> Result { + let input = serde_json::json!({ + "refresh_token": refresh_token, + "config": config + }); + self.call_plugin_function("refresh_tokens", &input).await + } + + /// Call get_profile(access_token, config) + pub async fn call_get_profile( + &mut self, + access_token: &str, + config: &serde_json::Value, + ) -> Result { + let input = serde_json::json!({ + "access_token": access_token, + "config": config + }); + self.call_plugin_function("get_profile", &input).await + } + + /// Call sync_account(access_token, config) + pub async fn call_sync_account( + &mut self, + access_token: &str, + config: &serde_json::Value, + ) -> Result, ExecutionError> { + let input = serde_json::json!({ + "access_token": access_token, + "config": config + }); + self.call_plugin_function("sync_account", &input).await + } + + /// Generic helper for plugin functions with input and typed output + async fn call_plugin_function( + &mut self, + name: &str, + input: &serde_json::Value, + ) -> Result { + let input_bytes = serde_json::to_vec(input) + .map_err(|e| ExecutionError::InvalidResponse(e.to_string()))?; + + let func = self + .instance + .get_typed_func::<(u32, u32), i64>(&mut self.store, name) + .map_err(|_| ExecutionError::MissingExport(name.into()))?; + + self.store + .set_fuel(DEFAULT_FUEL) + .map_err(ExecutionError::Trap)?; + + let (input_ptr, input_len) = write_to_guest(&mut self.store, &input_bytes) + .await + .map_err(|_| ExecutionError::MemoryAllocation)?; + + let packed = func + .call_async(&mut self.store, (input_ptr, input_len)) + .await + .map_err(Self::classify_error)?; + + // Unpack i64: upper 32 bits = ptr, lower 32 bits = len + let ptr = (packed >> 32) as u32; + let len = (packed & 0xFFFFFFFF) as u32; + + let bytes = + read_from_guest(&self.store, ptr, len).map_err(|_| ExecutionError::MemoryAllocation)?; + + dealloc_guest(&mut self.store, ptr, len) + .await + .map_err(|_| ExecutionError::MemoryAllocation)?; + + let response: PluginResponse = serde_json::from_slice(&bytes) + .map_err(|e| ExecutionError::InvalidResponse(e.to_string()))?; + + response.into_result().map_err(ExecutionError::from) + } + + /// Classify a wasmtime error as Timeout or Trap + fn classify_error(e: wasmtime::Error) -> ExecutionError { + if e.to_string().contains("fuel") { + ExecutionError::Timeout + } else { + ExecutionError::Trap(e) + } + } +} + +/// Factory for creating plugin instances +pub struct PluginExecutor { + runtime: Arc, + registry: Arc, + db: sqlx::AnyPool, + db_backend: DatabaseBackend, + http_client: reqwest::Client, + lexicons: Arc, +} + +impl PluginExecutor { + pub fn new( + runtime: Arc, + registry: Arc, + db: sqlx::AnyPool, + db_backend: DatabaseBackend, + http_client: reqwest::Client, + lexicons: Arc, + ) -> Self { + Self { + runtime, + registry, + db, + db_backend, + http_client, + lexicons, + } + } + + /// Instantiate a plugin with the given scope + pub async fn instantiate( + &self, + plugin_id: &str, + scope: &str, + secrets: HashMap, + config: serde_json::Value, + ) -> Result { + // Get plugin from registry + let plugin = self + .registry + .get(plugin_id) + .await + .ok_or_else(|| ExecutionError::PluginNotFound(plugin_id.to_string()))?; + + // Compile module + let module = self + .runtime + .compile(&plugin.wasm_bytes) + .map_err(ExecutionError::Instantiation)?; + + // Create linker with host functions + let mut linker = Linker::new(self.runtime.engine()); + register_host_functions(&mut linker).map_err(ExecutionError::Instantiation)?; + + // Create store with initial state (memory/alloc/dealloc set to None) + // Note: db is Option in PluginState + let state = PluginState { + plugin_id: plugin_id.to_string(), + scope: scope.to_string(), + secrets, + config, + db: Some(self.db.clone()), + db_backend: self.db_backend, + http_client: self.http_client.clone(), + lexicons: self.lexicons.clone(), + usage: Default::default(), + memory: None, + alloc: None, + dealloc: None, + }; + + let mut store = Store::new(self.runtime.engine(), state); + store + .set_fuel(DEFAULT_FUEL) + .map_err(ExecutionError::Instantiation)?; + + // Instantiate module + let instance = linker + .instantiate_async(&mut store, &module) + .await + .map_err(ExecutionError::Instantiation)?; + + // Get memory export + let memory = instance + .get_memory(&mut store, "memory") + .ok_or_else(|| ExecutionError::MissingExport("memory".into()))?; + + // Get alloc/dealloc exports + let alloc = instance + .get_typed_func::(&mut store, "alloc") + .map_err(|_| ExecutionError::MissingExport("alloc".into()))?; + let dealloc = instance + .get_typed_func::<(u32, u32), ()>(&mut store, "dealloc") + .map_err(|_| ExecutionError::MissingExport("dealloc".into()))?; + + // Store memory/alloc/dealloc in state + store.data_mut().memory = Some(memory); + store.data_mut().alloc = Some(alloc.clone()); + store.data_mut().dealloc = Some(dealloc.clone()); + + Ok(PluginInstance { + store, + instance, + memory, + alloc, + dealloc, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_execution_error_plugin_not_found() { + let err = ExecutionError::PluginNotFound("steam".into()); + assert!(err.to_string().contains("steam")); + assert!(err.to_string().contains("not found")); + } + + #[test] + fn test_execution_error_timeout() { + let err = ExecutionError::Timeout; + assert!( + err.to_string().to_lowercase().contains("timeout") || err.to_string().contains("fuel") + ); + } + + #[test] + fn test_plugin_error_conversion() { + let plugin_err = PluginEnvelopeError { + code: "AUTH_FAILED".into(), + message: "Bad token".into(), + retryable: true, + }; + let exec_err: ExecutionError = plugin_err.into(); + match exec_err { + ExecutionError::PluginError { + code, + message, + retryable, + } => { + assert_eq!(code, "AUTH_FAILED"); + assert_eq!(message, "Bad token"); + assert!(retryable); + } + _ => panic!("Wrong error variant"), + } + } + + #[test] + fn test_all_error_variants_have_display() { + let errors: Vec = vec![ + ExecutionError::PluginNotFound("test".into()), + ExecutionError::MemoryAllocation, + ExecutionError::InvalidResponse("bad json".into()), + ExecutionError::ResourceLimit("too many requests".into()), + ExecutionError::Timeout, + ExecutionError::MissingExport("plugin_info".into()), + ]; + for err in errors { + assert!(!err.to_string().is_empty()); + } + } + + #[test] + fn test_plugin_executor_new_signature() { + // Verify PluginExecutor::new exists with expected signature (compile-time check) + fn _check_signature( + _runtime: std::sync::Arc, + _registry: std::sync::Arc, + _db: sqlx::AnyPool, + _db_backend: crate::db::DatabaseBackend, + _http_client: reqwest::Client, + _lexicons: std::sync::Arc, + ) -> PluginExecutor { + PluginExecutor::new( + _runtime, + _registry, + _db, + _db_backend, + _http_client, + _lexicons, + ) + } + } + + #[test] + fn test_plugin_instance_struct_exists() { + // Verify PluginInstance struct has expected fields (compile-time check) + fn _check_fields(instance: PluginInstance) { + let _ = instance.store; + let _ = instance.instance; + let _ = instance.memory; + let _ = instance.alloc; + let _ = instance.dealloc; + } + } + + #[test] + fn test_plugin_instance_has_expected_methods() { + // Compile-time check that methods exist with expected signatures + fn _check_call_plugin_info<'a>( + inst: &'a mut PluginInstance, + ) -> impl std::future::Future> + 'a + { + inst.call_plugin_info() + } + + fn _check_call_get_authorize_url<'a>( + inst: &'a mut PluginInstance, + state: &'a str, + redirect_uri: &'a str, + config: &'a serde_json::Value, + ) -> impl std::future::Future> + 'a { + inst.call_get_authorize_url(state, redirect_uri, config) + } + + fn _check_call_handle_callback<'a>( + inst: &'a mut PluginInstance, + code: &'a str, + state: &'a str, + config: &'a serde_json::Value, + ) -> impl std::future::Future> + 'a + { + inst.call_handle_callback(code, state, config) + } + + fn _check_call_refresh_tokens<'a>( + inst: &'a mut PluginInstance, + refresh_token: &'a str, + config: &'a serde_json::Value, + ) -> impl std::future::Future> + 'a + { + inst.call_refresh_tokens(refresh_token, config) + } + + fn _check_call_get_profile<'a>( + inst: &'a mut PluginInstance, + access_token: &'a str, + config: &'a serde_json::Value, + ) -> impl std::future::Future> + 'a + { + inst.call_get_profile(access_token, config) + } + + fn _check_call_sync_account<'a>( + inst: &'a mut PluginInstance, + access_token: &'a str, + config: &'a serde_json::Value, + ) -> impl std::future::Future, ExecutionError>> + 'a + { + inst.call_sync_account(access_token, config) + } + } +} diff --git a/src/plugin/host/bindings.rs b/src/plugin/host/bindings.rs new file mode 100644 index 0000000..36e1303 --- /dev/null +++ b/src/plugin/host/bindings.rs @@ -0,0 +1,438 @@ +use std::collections::HashMap; +use std::sync::Arc; +use wasmtime::{Linker, Memory, TypedFunc}; + +/// State stored in wasmtime's Store during plugin execution +pub struct PluginState { + pub plugin_id: String, + pub scope: String, + pub secrets: HashMap, + pub config: serde_json::Value, + pub db: Option, + pub db_backend: crate::db::DatabaseBackend, + pub http_client: reqwest::Client, + pub lexicons: Arc, + pub usage: super::ResourceUsage, + pub memory: Option, + pub alloc: Option>, + pub dealloc: Option>, +} + +/// Check that a memory access is within bounds +fn check_bounds(offset: usize, length: usize, mem_size: usize) -> Result<(usize, usize), ()> { + if length == 0 { + return Ok((offset, offset)); + } + let end = offset.checked_add(length).ok_or(())?; + if end > mem_size { + return Err(()); + } + Ok((offset, end)) +} + +/// Register all host functions with the linker +pub fn register_host_functions(linker: &mut Linker) -> Result<(), wasmtime::Error> { + // Sync functions + linker.func_wrap("env", "host_log", host_log)?; + linker.func_wrap("env", "host_get_secret", host_get_secret)?; + + // Async functions - HTTP + linker.func_wrap_async( + "env", + "host_http_request", + |mut caller: wasmtime::Caller<'_, PluginState>, (req_ptr, req_len): (i32, i32)| { + Box::new(async move { host_http_request_impl(&mut caller, req_ptr, req_len).await }) + }, + )?; + + // Async functions - KV + linker.func_wrap_async( + "env", + "host_kv_get", + |mut caller: wasmtime::Caller<'_, PluginState>, (key_ptr, key_len): (i32, i32)| { + Box::new(async move { host_kv_get_impl(&mut caller, key_ptr, key_len).await }) + }, + )?; + + linker.func_wrap_async( + "env", + "host_kv_set", + |mut caller: wasmtime::Caller<'_, PluginState>, + (key_ptr, key_len, val_ptr, val_len, ttl): (i32, i32, i32, i32, i32)| { + Box::new(async move { + host_kv_set_impl(&mut caller, key_ptr, key_len, val_ptr, val_len, ttl).await + }) + }, + )?; + + linker.func_wrap_async( + "env", + "host_kv_delete", + |mut caller: wasmtime::Caller<'_, PluginState>, (key_ptr, key_len): (i32, i32)| { + Box::new(async move { host_kv_delete_impl(&mut caller, key_ptr, key_len).await }) + }, + )?; + + // Async functions - Record lookup + linker.func_wrap_async( + "env", + "host_lookup_record", + |mut caller: wasmtime::Caller<'_, PluginState>, (req_ptr, req_len): (i32, i32)| { + Box::new(async move { host_lookup_record_impl(&mut caller, req_ptr, req_len).await }) + }, + )?; + + Ok(()) +} + +/// Read a string from guest memory +fn read_guest_string( + caller: &wasmtime::Caller<'_, PluginState>, + ptr: i32, + len: i32, +) -> Option { + let memory = caller.data().memory?; + let mem_data = memory.data(caller); + let (start, end) = check_bounds(ptr as usize, len as usize, mem_data.len()).ok()?; + std::str::from_utf8(&mem_data[start..end]) + .ok() + .map(String::from) +} + +/// Read raw bytes from guest memory +fn read_guest_bytes( + caller: &wasmtime::Caller<'_, PluginState>, + ptr: i32, + len: i32, +) -> Option> { + let memory = caller.data().memory?; + let mem_data = memory.data(caller); + let (start, end) = check_bounds(ptr as usize, len as usize, mem_data.len()).ok()?; + Some(mem_data[start..end].to_vec()) +} + +/// Write response data to guest memory, returning packed (ptr << 32) | len +async fn write_guest_response(caller: &mut wasmtime::Caller<'_, PluginState>, data: &[u8]) -> i64 { + let memory = match caller.data().memory { + Some(m) => m, + None => return 0, + }; + let alloc = match &caller.data().alloc { + Some(a) => a.clone(), + None => return 0, + }; + + let len = data.len() as u32; + let ptr = match alloc.call_async(&mut *caller, len).await { + Ok(p) if p != 0 => p, + _ => return 0, + }; + + let mem_data = memory.data_mut(caller); + if check_bounds(ptr as usize, len as usize, mem_data.len()).is_err() { + return 0; + } + + mem_data[ptr as usize..(ptr as usize + len as usize)].copy_from_slice(data); + ((ptr as i64) << 32) | (len as i64) +} + +/// Host function: log a message from the plugin +fn host_log( + caller: wasmtime::Caller<'_, PluginState>, + level_ptr: i32, + level_len: i32, + msg_ptr: i32, + msg_len: i32, +) { + let memory = match caller.data().memory { + Some(m) => m, + None => return, + }; + + let mem_data = memory.data(&caller); + let mem_size = mem_data.len(); + + let (level_start, level_end) = + match check_bounds(level_ptr as usize, level_len as usize, mem_size) { + Ok(bounds) => bounds, + Err(_) => return, + }; + + let (msg_start, msg_end) = match check_bounds(msg_ptr as usize, msg_len as usize, mem_size) { + Ok(bounds) => bounds, + Err(_) => return, + }; + + let level = std::str::from_utf8(&mem_data[level_start..level_end]).unwrap_or("info"); + let msg = std::str::from_utf8(&mem_data[msg_start..msg_end]).unwrap_or(""); + + let plugin_id = &caller.data().plugin_id; + let log_level: super::LogLevel = level.parse().unwrap_or_default(); + super::log(plugin_id, log_level, msg); +} + +/// Host function: get a secret value by name +/// Returns a packed i64: (ptr << 32) | len, or 0 on error +fn host_get_secret( + mut caller: wasmtime::Caller<'_, PluginState>, + name_ptr: i32, + name_len: i32, +) -> i64 { + let memory = match caller.data().memory { + Some(m) => m, + None => return 0, + }; + let alloc = match &caller.data().alloc { + Some(a) => a.clone(), + None => return 0, + }; + + let mem_data = memory.data(&caller); + let mem_size = mem_data.len(); + + let (name_start, name_end) = match check_bounds(name_ptr as usize, name_len as usize, mem_size) + { + Ok(bounds) => bounds, + Err(_) => return 0, + }; + + let name = match std::str::from_utf8(&mem_data[name_start..name_end]) { + Ok(s) => s, + Err(_) => return 0, + }; + + let value = match caller.data().secrets.get(name) { + Some(v) => v.clone(), + None => return 0, + }; + + let len = value.len() as u32; + let ptr = match alloc.call(&mut caller, len) { + Ok(p) if p != 0 => p, + _ => return 0, + }; + + let mem_data = memory.data_mut(&mut caller); + if check_bounds(ptr as usize, len as usize, mem_data.len()).is_err() { + return 0; + } + + mem_data[ptr as usize..(ptr as usize + len as usize)].copy_from_slice(value.as_bytes()); + + ((ptr as i64) << 32) | (len as i64) +} + +// ============================================================================ +// Async host function implementations +// ============================================================================ + +/// Build a HostContext from PluginState, requires db to be present +fn build_host_context(state: &PluginState) -> Option { + let db = state.db.clone()?; + Some(super::HostContext { + plugin_id: state.plugin_id.clone(), + scope: state.scope.clone(), + secrets: state.secrets.clone(), + config: state.config.clone(), + db, + db_backend: state.db_backend, + http_client: state.http_client.clone(), + lexicons: state.lexicons.clone(), + }) +} + +/// Host function: make an HTTP request +async fn host_http_request_impl( + caller: &mut wasmtime::Caller<'_, PluginState>, + req_ptr: i32, + req_len: i32, +) -> i64 { + let req_bytes = match read_guest_bytes(caller, req_ptr, req_len) { + Some(b) => b, + None => return 0, + }; + + let request: super::HttpRequest = match serde_json::from_slice(&req_bytes) { + Ok(r) => r, + Err(_) => return 0, + }; + + let ctx = match build_host_context(caller.data()) { + Some(c) => c, + None => return 0, + }; + + let result = { + let usage = &mut caller.data_mut().usage; + super::http_request(&ctx, usage, request).await + }; + + let response_bytes = match result { + Ok(resp) => serde_json::to_vec(&serde_json::json!({"ok": resp})).unwrap_or_default(), + Err(e) => serde_json::to_vec(&serde_json::json!({ + "error": {"code": "HTTP_ERROR", "message": e.to_string(), "retryable": false} + })) + .unwrap_or_default(), + }; + + write_guest_response(caller, &response_bytes).await +} + +/// Host function: get a value from KV store +async fn host_kv_get_impl( + caller: &mut wasmtime::Caller<'_, PluginState>, + key_ptr: i32, + key_len: i32, +) -> i64 { + let key = match read_guest_string(caller, key_ptr, key_len) { + Some(k) => k, + None => return 0, + }; + + let ctx = match build_host_context(caller.data()) { + Some(c) => c, + None => return 0, + }; + + let result = super::kv_get(&ctx, &key).await; + + let response_bytes = match result { + Ok(Some(value)) => { + serde_json::to_vec(&serde_json::json!({"ok": value})).unwrap_or_default() + } + Ok(None) => return 0, + Err(e) => serde_json::to_vec(&serde_json::json!({ + "error": {"code": "KV_ERROR", "message": e.to_string(), "retryable": false} + })) + .unwrap_or_default(), + }; + + write_guest_response(caller, &response_bytes).await +} + +/// Host function: set a value in KV store +async fn host_kv_set_impl( + caller: &mut wasmtime::Caller<'_, PluginState>, + key_ptr: i32, + key_len: i32, + val_ptr: i32, + val_len: i32, + ttl: i32, +) -> i32 { + let key = match read_guest_string(caller, key_ptr, key_len) { + Some(k) => k, + None => return -1, + }; + let value = match read_guest_bytes(caller, val_ptr, val_len) { + Some(v) => v, + None => return -1, + }; + + let ttl_secs = if ttl > 0 { Some(ttl as u32) } else { None }; + + let ctx = match build_host_context(caller.data()) { + Some(c) => c, + None => return -1, + }; + + let usage = &mut caller.data_mut().usage; + match super::kv_set(&ctx, usage, &key, value, ttl_secs).await { + Ok(()) => 0, + Err(_) => -1, + } +} + +/// Host function: delete a value from KV store +async fn host_kv_delete_impl( + caller: &mut wasmtime::Caller<'_, PluginState>, + key_ptr: i32, + key_len: i32, +) -> i32 { + let key = match read_guest_string(caller, key_ptr, key_len) { + Some(k) => k, + None => return -1, + }; + + let ctx = match build_host_context(caller.data()) { + Some(c) => c, + None => return -1, + }; + + match super::kv_delete(&ctx, &key).await { + Ok(()) => 0, + Err(_) => -1, + } +} + +/// Host function: look up an AT Protocol record +async fn host_lookup_record_impl( + caller: &mut wasmtime::Caller<'_, PluginState>, + req_ptr: i32, + req_len: i32, +) -> i64 { + let req_bytes = match read_guest_bytes(caller, req_ptr, req_len) { + Some(b) => b, + None => return 0, + }; + + let request: super::LookupRequest = match serde_json::from_slice(&req_bytes) { + Ok(r) => r, + Err(_) => return 0, + }; + + let ctx = match build_host_context(caller.data()) { + Some(c) => c, + None => return 0, + }; + + let result = super::lookup_record_by_request(&ctx, request).await; + + let response_bytes = match result { + Ok(record) => serde_json::to_vec(&serde_json::json!({"ok": record})).unwrap_or_default(), + Err(e) => serde_json::to_vec(&serde_json::json!({ + "error": {"code": "LOOKUP_ERROR", "message": e.to_string(), "retryable": false} + })) + .unwrap_or_default(), + }; + + write_guest_response(caller, &response_bytes).await +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_plugin_state_fields_exist() { + fn _check_fields(state: &PluginState) { + let _ = &state.plugin_id; + let _ = &state.scope; + let _ = &state.secrets; + let _ = &state.config; + let _ = &state.usage; + let _ = &state.memory; + let _ = &state.alloc; + let _ = &state.dealloc; + } + } + + #[test] + fn test_pack_ptr_len() { + let ptr: u32 = 0x1000; + let len: u32 = 0x0100; + let packed: i64 = ((ptr as i64) << 32) | (len as i64); + let unpacked_ptr = (packed >> 32) as u32; + let unpacked_len = (packed & 0xFFFFFFFF) as u32; + assert_eq!(unpacked_ptr, ptr); + assert_eq!(unpacked_len, len); + } + + #[test] + fn test_bounds_check_helper() { + assert!(check_bounds(0, 10, 100).is_ok()); + assert!(check_bounds(90, 10, 100).is_ok()); + assert!(check_bounds(91, 10, 100).is_err()); + assert!(check_bounds(0, 0, 100).is_ok()); + } +} diff --git a/src/plugin/host/lookup.rs b/src/plugin/host/lookup.rs index 704e79b..ec0c034 100644 --- a/src/plugin/host/lookup.rs +++ b/src/plugin/host/lookup.rs @@ -1,3 +1,5 @@ +use serde::Deserialize; + use super::HostContext; use crate::db::adapt_sql; use crate::plugin::StrongRef; @@ -10,6 +12,13 @@ pub enum LookupError { InvalidFieldPath, } +#[derive(Debug, Deserialize)] +pub struct LookupRequest { + pub collection: String, + pub external_id_field: String, + pub external_id_value: String, +} + /// Look up a record by external ID /// /// # Arguments @@ -48,6 +57,19 @@ pub async fn lookup_record( Ok(result.map(|(uri, cid)| StrongRef { uri, cid })) } +pub async fn lookup_record_by_request( + ctx: &HostContext, + request: LookupRequest, +) -> Result, LookupError> { + lookup_record( + ctx, + &request.collection, + &request.external_id_field, + &request.external_id_value, + ) + .await +} + #[cfg(test)] mod tests { use super::*; @@ -64,4 +86,13 @@ mod tests { let err = LookupError::InvalidFieldPath; assert_eq!(err.to_string(), "Invalid external ID field path"); } + + #[test] + fn test_lookup_request_deserialize() { + let json = r#"{"collection": "games.example.game", "external_id_field": "externalIds.steam", "external_id_value": "123"}"#; + let req: LookupRequest = serde_json::from_str(json).unwrap(); + assert_eq!(req.collection, "games.example.game"); + assert_eq!(req.external_id_field, "externalIds.steam"); + assert_eq!(req.external_id_value, "123"); + } } diff --git a/src/plugin/host/mod.rs b/src/plugin/host/mod.rs index f00fe4f..5556627 100644 --- a/src/plugin/host/mod.rs +++ b/src/plugin/host/mod.rs @@ -1,9 +1,11 @@ +mod bindings; mod http; mod kv; mod logging; mod lookup; mod secrets; +pub use bindings::{PluginState, register_host_functions}; pub use http::*; pub use kv::*; pub use logging::*; diff --git a/src/plugin/loader.rs b/src/plugin/loader.rs index 331808c..a4f9374 100644 --- a/src/plugin/loader.rs +++ b/src/plugin/loader.rs @@ -1,6 +1,11 @@ +use crate::plugin::host::{PluginState, register_host_functions}; +use crate::plugin::memory::PluginResponse; +use crate::plugin::runtime::DEFAULT_FUEL; use crate::plugin::{LoadedPlugin, PluginInfo, PluginSource}; use sha2::{Digest, Sha256}; +use std::collections::HashMap; use std::path::Path; +use wasmtime::{Config, Engine, Linker, Module, Store}; const SUPPORTED_API_VERSION: &str = "1"; @@ -84,23 +89,106 @@ pub async fn load_from_url( /// Extract plugin info by instantiating WASM and calling plugin_info() fn extract_plugin_info(wasm_bytes: &[u8]) -> Result { - // TODO: Full implementation with wasmtime - // For now, this is a placeholder that will be filled in when we integrate wasmtime calls + match tokio::runtime::Handle::try_current() { + Ok(handle) => { + tokio::task::block_in_place(|| handle.block_on(extract_plugin_info_async(wasm_bytes))) + } + Err(_) => { + let rt = tokio::runtime::Runtime::new().map_err(|e| { + LoadError::WasmValidation(format!("failed to create runtime: {}", e)) + })?; + rt.block_on(extract_plugin_info_async(wasm_bytes)) + } + } +} + +/// Async implementation of plugin info extraction via WASM instantiation +async fn extract_plugin_info_async(wasm_bytes: &[u8]) -> Result { + // Create async-enabled engine with fuel + let mut config = Config::new(); + config.async_support(true); + config.consume_fuel(true); + let engine = Engine::new(&config).map_err(|e| LoadError::WasmValidation(e.to_string()))?; - // Validate it's valid WASM - wasmtime::Module::validate(&wasmtime::Engine::default(), wasm_bytes) + let module = + Module::new(&engine, wasm_bytes).map_err(|e| LoadError::WasmValidation(e.to_string()))?; + + // Create linker with host functions + let mut linker = Linker::new(&engine); + register_host_functions(&mut linker).map_err(|e| LoadError::WasmValidation(e.to_string()))?; + + // Create minimal state - no db needed for plugin_info() + let state = PluginState { + plugin_id: "loading".into(), + scope: "".into(), + secrets: HashMap::new(), + config: serde_json::Value::Null, + db: None, // Not needed for plugin_info + db_backend: crate::db::DatabaseBackend::Sqlite, + http_client: reqwest::Client::new(), + lexicons: std::sync::Arc::new(crate::lexicon::LexiconRegistry::new()), + usage: Default::default(), + memory: None, + alloc: None, + dealloc: None, + }; + + let mut store = Store::new(&engine, state); + store + .set_fuel(DEFAULT_FUEL) .map_err(|e| LoadError::WasmValidation(e.to_string()))?; - // Return placeholder - real implementation calls plugin_info() export - Ok(PluginInfo { - id: "placeholder".into(), - name: "Placeholder".into(), - version: "0.0.0".into(), - api_version: SUPPORTED_API_VERSION.into(), - icon_url: None, - required_secrets: vec![], - config_schema: None, - }) + // Instantiate + let instance = linker + .instantiate_async(&mut store, &module) + .await + .map_err(|e| LoadError::WasmValidation(format!("instantiation failed: {}", e)))?; + + // Get memory and alloc/dealloc + let memory = instance + .get_memory(&mut store, "memory") + .ok_or_else(|| LoadError::WasmValidation("missing memory export".into()))?; + let alloc = instance + .get_typed_func::(&mut store, "alloc") + .map_err(|_| LoadError::WasmValidation("missing alloc export".into()))?; + let dealloc = instance + .get_typed_func::<(u32, u32), ()>(&mut store, "dealloc") + .map_err(|_| LoadError::WasmValidation("missing dealloc export".into()))?; + + // Store in state + store.data_mut().memory = Some(memory); + store.data_mut().alloc = Some(alloc); + store.data_mut().dealloc = Some(dealloc); + + // Call plugin_info + let func = instance + .get_typed_func::<(), i64>(&mut store, "plugin_info") + .map_err(|_| LoadError::WasmValidation("missing plugin_info export".into()))?; + + let packed = func + .call_async(&mut store, ()) + .await + .map_err(|e| LoadError::WasmValidation(format!("plugin_info failed: {}", e)))?; + + // Unpack i64: upper 32 bits = ptr, lower 32 bits = len + let ptr = (packed >> 32) as u32; + let len = (packed & 0xFFFFFFFF) as u32; + + // Read result from memory + let mem_data = memory.data(&store); + if (ptr as usize) + (len as usize) > mem_data.len() { + return Err(LoadError::WasmValidation( + "plugin_info returned out of bounds pointer".into(), + )); + } + let bytes = mem_data[ptr as usize..(ptr as usize + len as usize)].to_vec(); + + // Parse response + let response: PluginResponse = serde_json::from_slice(&bytes)?; + + response + .into_result() + .map_err(|e| LoadError::WasmValidation(format!("plugin error: {}", e.message))) } fn validate_api_version(info: &PluginInfo) -> Result<(), LoadError> { diff --git a/src/plugin/memory.rs b/src/plugin/memory.rs new file mode 100644 index 0000000..dd19b59 --- /dev/null +++ b/src/plugin/memory.rs @@ -0,0 +1,193 @@ +use crate::plugin::host::PluginState; +use serde::{Deserialize, Serialize}; +use thiserror::Error; +use wasmtime::Store; + +/// Error returned from a plugin via JSON envelope. +/// Uses a string code for flexibility in parsing arbitrary error codes from plugins. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PluginEnvelopeError { + pub code: String, + pub message: String, + #[serde(default)] + pub retryable: bool, +} + +impl std::fmt::Display for PluginEnvelopeError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}: {}", self.code, self.message) + } +} + +impl std::error::Error for PluginEnvelopeError {} + +/// JSON envelope for plugin responses. +/// Plugins return either `{"ok": result}` or `{"error": {...}}`. +#[derive(Debug, Deserialize)] +#[serde(untagged)] +pub enum PluginResponse { + Ok { ok: T }, + Error { error: PluginEnvelopeError }, +} + +impl PluginResponse { + pub fn into_result(self) -> Result { + match self { + PluginResponse::Ok { ok } => Ok(ok), + PluginResponse::Error { error } => Err(error), + } + } +} + +#[derive(Debug, Error)] +pub enum MemoryError { + #[error("Memory allocation failed: alloc returned 0")] + AllocationFailed, + #[error( + "Memory access out of bounds: offset {offset} + length {length} exceeds memory size {size}" + )] + OutOfBounds { + offset: usize, + length: usize, + size: usize, + }, + #[error("WASM trap during memory operation: {0}")] + Trap(#[from] wasmtime::Error), +} + +/// Write data to WASM guest memory by calling alloc and copying bytes. +/// Returns (ptr, len) tuple on success. +pub async fn write_to_guest( + store: &mut Store, + data: &[u8], +) -> Result<(u32, u32), MemoryError> { + let len = data.len() as u32; + if len == 0 { + return Ok((0, 0)); + } + + let alloc = store + .data() + .alloc + .as_ref() + .ok_or(MemoryError::AllocationFailed)? + .clone(); + let memory = store.data().memory.ok_or(MemoryError::AllocationFailed)?; + + let ptr = alloc.call_async(&mut *store, len).await?; + if ptr == 0 { + return Err(MemoryError::AllocationFailed); + } + + let mem_size = memory.data_size(&*store); + let start = ptr as usize; + let end = start + .checked_add(len as usize) + .ok_or(MemoryError::OutOfBounds { + offset: start, + length: len as usize, + size: mem_size, + })?; + + if end > mem_size { + return Err(MemoryError::OutOfBounds { + offset: start, + length: len as usize, + size: mem_size, + }); + } + + memory.data_mut(&mut *store)[start..end].copy_from_slice(data); + Ok((ptr, len)) +} + +/// Read data from WASM guest memory at the given pointer and length. +pub fn read_from_guest( + store: &Store, + ptr: u32, + len: u32, +) -> Result, MemoryError> { + if len == 0 { + return Ok(Vec::new()); + } + + let memory = store.data().memory.ok_or(MemoryError::AllocationFailed)?; + let mem_size = memory.data_size(store); + let start = ptr as usize; + let end = start + .checked_add(len as usize) + .ok_or(MemoryError::OutOfBounds { + offset: start, + length: len as usize, + size: mem_size, + })?; + + if end > mem_size { + return Err(MemoryError::OutOfBounds { + offset: start, + length: len as usize, + size: mem_size, + }); + } + + Ok(memory.data(store)[start..end].to_vec()) +} + +/// Deallocate guest memory by calling the dealloc function. +pub async fn dealloc_guest( + store: &mut Store, + ptr: u32, + len: u32, +) -> Result<(), MemoryError> { + if len == 0 { + return Ok(()); + } + + let dealloc = store + .data() + .dealloc + .as_ref() + .ok_or(MemoryError::AllocationFailed)? + .clone(); + dealloc.call_async(&mut *store, (ptr, len)).await?; + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_memory_error_display() { + let err = MemoryError::AllocationFailed; + assert!(err.to_string().contains("alloc")); + + let err = MemoryError::OutOfBounds { + offset: 100, + length: 50, + size: 120, + }; + assert!(err.to_string().contains("100")); + } + + #[test] + fn test_plugin_response_ok_parses() { + let json = r#"{"ok": "hello"}"#; + let resp: PluginResponse = serde_json::from_str(json).unwrap(); + let result = resp.into_result(); + assert!(result.is_ok()); + assert_eq!(result.unwrap(), "hello"); + } + + #[test] + fn test_plugin_response_error_parses() { + let json = + r#"{"error": {"code": "AUTH_FAILED", "message": "Invalid token", "retryable": true}}"#; + let resp: PluginResponse = serde_json::from_str(json).unwrap(); + let result = resp.into_result(); + assert!(result.is_err()); + let err = result.unwrap_err(); + assert_eq!(err.code, "AUTH_FAILED"); + assert!(err.retryable); + } +} diff --git a/src/plugin/mod.rs b/src/plugin/mod.rs index ec10134..ce5e4a2 100644 --- a/src/plugin/mod.rs +++ b/src/plugin/mod.rs @@ -1,9 +1,15 @@ +pub mod attestation; pub mod encryption; +pub mod executor; pub mod host; pub mod loader; +pub mod memory; mod runtime; +pub mod sync; mod types; +pub use executor::{ExecutionError, PluginExecutor, PluginInstance}; +pub use memory::{MemoryError, PluginEnvelopeError, PluginResponse}; pub use runtime::WasmRuntime; pub use types::*; diff --git a/src/plugin/runtime.rs b/src/plugin/runtime.rs index 97f7278..547dc75 100644 --- a/src/plugin/runtime.rs +++ b/src/plugin/runtime.rs @@ -1,5 +1,8 @@ use wasmtime::*; +/// Default fuel for plugin execution (≈100ms CPU time) +pub const DEFAULT_FUEL: u64 = 10_000_000; + /// WASM runtime for executing plugins pub struct WasmRuntime { engine: Engine, @@ -9,6 +12,7 @@ impl WasmRuntime { pub fn new() -> Result { let mut config = Config::new(); config.async_support(true); + config.consume_fuel(true); let engine = Engine::new(&config)?; @@ -18,6 +22,11 @@ impl WasmRuntime { pub fn engine(&self) -> &Engine { &self.engine } + + /// Compile a WASM module + pub fn compile(&self, wasm_bytes: &[u8]) -> Result { + Module::new(&self.engine, wasm_bytes) + } } impl Default for WasmRuntime { @@ -25,3 +34,29 @@ impl Default for WasmRuntime { Self::new().expect("Failed to create WASM runtime") } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_fuel_constant_value() { + // 10M fuel ≈ 100ms CPU time per spec + assert_eq!(DEFAULT_FUEL, 10_000_000); + } + + #[test] + fn test_runtime_has_fuel_enabled() { + let runtime = WasmRuntime::new().expect("Failed to create runtime"); + // We can verify fuel is enabled by checking we can set it on a store + let mut store = wasmtime::Store::new(runtime.engine(), ()); + assert!(store.set_fuel(1000).is_ok()); + } + + #[test] + fn test_compile_invalid_wasm_fails() { + let runtime = WasmRuntime::new().expect("Failed to create runtime"); + let result = runtime.compile(b"not valid wasm"); + assert!(result.is_err()); + } +} diff --git a/src/plugin/sync.rs b/src/plugin/sync.rs new file mode 100644 index 0000000..9aca8e9 --- /dev/null +++ b/src/plugin/sync.rs @@ -0,0 +1,312 @@ +//! SyncRecord processing pipeline. +//! +//! Processes records returned by plugin sync_account(): +//! - Signs records that have `sign: true` +//! - Resolves game references +//! - Prepares records for writing to PDS + +use super::attestation::{AttestationError, AttestationSigner}; +use super::types::SyncRecord; +use crate::db::{DatabaseBackend, adapt_sql}; +use serde_json::Value; + +/// Processed record ready for storage +#[derive(Debug, Clone)] +pub struct ProcessedRecord { + /// The collection (lexicon ID) + pub collection: String, + /// The processed record with signatures added + pub record: Value, + /// Deduplication key + pub dedup_key: Option, + /// CID of the signed content (if signed) + pub content_cid: Option, +} + +/// Error during sync record processing +#[derive(Debug, thiserror::Error)] +pub enum SyncError { + #[error("Attestation signing failed: {0}")] + Attestation(#[from] AttestationError), + + #[error("Game reference resolution failed: {0}")] + GameResolution(String), + + #[error("Invalid record: {0}")] + InvalidRecord(String), +} + +/// Process a batch of SyncRecords from a plugin +pub struct SyncProcessor<'a> { + /// Attestation signer (optional - if None, signing is skipped) + signer: Option<&'a AttestationSigner>, + /// Repository DID for the user (used in $sig for replay protection) + repository_did: String, +} + +impl<'a> SyncProcessor<'a> { + /// Create a new sync processor + pub fn new(signer: Option<&'a AttestationSigner>, repository_did: String) -> Self { + Self { + signer, + repository_did, + } + } + + /// Process a batch of SyncRecords + pub fn process_records( + &self, + records: Vec, + ) -> Result, SyncError> { + let mut processed = Vec::with_capacity(records.len()); + + for record in records { + processed.push(self.process_record(record)?); + } + + Ok(processed) + } + + /// Process a single SyncRecord + fn process_record(&self, sync_record: SyncRecord) -> Result { + let mut record = sync_record.record; + + // Resolve game references if present + self.resolve_game_ref(&mut record)?; + + // Sign if requested and signer is available + let content_cid = if sync_record.sign { + if let Some(signer) = self.signer { + let cid = signer.sign_record(&mut record, &self.repository_did)?; + Some(cid.to_string()) + } else { + tracing::warn!( + collection = %sync_record.collection, + "Record requested signing but no signer configured" + ); + None + } + } else { + None + }; + + Ok(ProcessedRecord { + collection: sync_record.collection, + record, + dedup_key: sync_record.dedup_key, + content_cid, + }) + } + + /// Resolve game references in a record + /// + /// Looks for game references like `{"platform": "steam", "externalId": "440"}` + /// and attempts to resolve them to AT URIs. + fn resolve_game_ref(&self, record: &mut Value) -> Result<(), SyncError> { + // Look for "game" field with platform/externalId structure + if let Some(obj) = record.as_object_mut() + && let Some(game_ref) = obj.get("game") + && let Some(game_obj) = game_ref.as_object() + && game_obj.contains_key("platform") + && game_obj.contains_key("externalId") + && !game_obj.contains_key("uri") + { + // Unresolved reference - log for debugging + let platform = game_obj + .get("platform") + .and_then(|v| v.as_str()) + .unwrap_or("unknown"); + let external_id = game_obj + .get("externalId") + .and_then(|v| v.as_str()) + .unwrap_or("unknown"); + + tracing::debug!( + platform = %platform, + external_id = %external_id, + "Game reference left unresolved - resolution not yet implemented" + ); + } + + Ok(()) + } +} + +/// Helper to create a sync processor with common setup +pub fn create_processor<'a>( + signer: Option<&'a AttestationSigner>, + user_did: &str, +) -> SyncProcessor<'a> { + SyncProcessor::new(signer, user_did.to_string()) +} + +/// Resolve game references in records by looking up in the database. +/// +/// Looks for `game: {platform: "steam", externalId: "440"}` and converts to +/// `game: {uri: "at://...", cid: "..."}` if found. +pub async fn resolve_game_references( + db: &sqlx::AnyPool, + backend: DatabaseBackend, + records: &mut [SyncRecord], +) { + for record in records.iter_mut() { + let Some(obj) = record.record.as_object_mut() else { + continue; + }; + let Some(game_ref) = obj.get("game").cloned() else { + continue; + }; + let Some(game_obj) = game_ref.as_object() else { + continue; + }; + + // Check for unresolved reference + if !game_obj.contains_key("platform") + || !game_obj.contains_key("externalId") + || game_obj.contains_key("uri") + { + continue; + } + + let platform = game_obj + .get("platform") + .and_then(|v| v.as_str()) + .unwrap_or(""); + let external_id = game_obj + .get("externalId") + .and_then(|v| v.as_str()) + .unwrap_or(""); + + if let Some((uri, cid)) = + lookup_game_by_external_id(db, backend, platform, external_id).await + { + obj.insert( + "game".to_string(), + serde_json::json!({ + "uri": uri, + "cid": cid + }), + ); + tracing::debug!( + platform = %platform, + external_id = %external_id, + uri = %uri, + "Resolved game reference" + ); + } else { + tracing::debug!( + platform = %platform, + external_id = %external_id, + "Game not found in database, leaving reference unresolved" + ); + } + } +} + +/// Look up a game by external ID (e.g., Steam app ID). +/// +/// Returns (uri, cid) if found. +async fn lookup_game_by_external_id( + db: &sqlx::AnyPool, + backend: DatabaseBackend, + platform: &str, + external_id: &str, +) -> Option<(String, String)> { + // Build JSON path based on platform + // Looking for records where: record.externalIds. = external_id + let json_path = match backend { + DatabaseBackend::Sqlite => { + format!("json_extract(record, '$.externalIds.{}')", platform) + } + DatabaseBackend::Postgres => { + format!("record->'externalIds'->>'{}'", platform) + } + }; + + let sql = adapt_sql( + &format!( + "SELECT uri, cid FROM records WHERE collection = 'games.gamesgamesgamesgames.game' AND {} = ? LIMIT 1", + json_path + ), + backend, + ); + + let result: Option<(String, String)> = sqlx::query_as(&sql) + .bind(external_id) + .fetch_optional(db) + .await + .ok() + .flatten(); + + result +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_process_unsigned_record() { + let processor = SyncProcessor::new(None, "did:plc:testuser".to_string()); + + let records = vec![SyncRecord { + collection: "test.collection".into(), + record: serde_json::json!({ + "$type": "test.collection", + "data": "hello" + }), + dedup_key: Some("test:1".into()), + sign: false, + }]; + + let processed = processor.process_records(records).unwrap(); + assert_eq!(processed.len(), 1); + assert_eq!(processed[0].collection, "test.collection"); + assert!(processed[0].content_cid.is_none()); + } + + #[test] + fn test_process_signed_record() { + let signer = + AttestationSigner::for_testing("did:web:test#key".into(), "test.signature".into()); + let processor = SyncProcessor::new(Some(&signer), "did:plc:testuser".to_string()); + + let records = vec![SyncRecord { + collection: "games.gamesgamesgamesgames.actor.game".into(), + record: serde_json::json!({ + "$type": "games.gamesgamesgamesgames.actor.game", + "game": {"platform": "steam", "externalId": "440"}, + "platform": "steam", + "createdAt": "2024-01-01T00:00:00Z" + }), + dedup_key: Some("steam:game:440".into()), + sign: true, + }]; + + let processed = processor.process_records(records).unwrap(); + assert_eq!(processed.len(), 1); + assert!(processed[0].content_cid.is_some()); + + // Verify signatures array was added + let signatures = processed[0].record["signatures"].as_array(); + assert!(signatures.is_some()); + assert_eq!(signatures.unwrap().len(), 1); + } + + #[test] + fn test_sign_requested_but_no_signer() { + let processor = SyncProcessor::new(None, "did:plc:testuser".to_string()); + + let records = vec![SyncRecord { + collection: "test.collection".into(), + record: serde_json::json!({"data": "hello"}), + dedup_key: None, + sign: true, // Requested but no signer + }]; + + let processed = processor.process_records(records).unwrap(); + assert_eq!(processed.len(), 1); + // No error, but no CID either + assert!(processed[0].content_cid.is_none()); + } +} diff --git a/src/plugin/types.rs b/src/plugin/types.rs index b8b23be..f4debaf 100644 --- a/src/plugin/types.rs +++ b/src/plugin/types.rs @@ -74,6 +74,9 @@ pub struct SyncRecord { pub record: serde_json::Value, #[serde(skip_serializing_if = "Option::is_none")] pub dedup_key: Option, + /// Whether HappyView should add an attestation signature to this record + #[serde(default)] + pub sign: bool, } /// Strong reference to an AT Protocol record diff --git a/tests/common/app.rs b/tests/common/app.rs index 4a03a6f..3bed538 100644 --- a/tests/common/app.rs +++ b/tests/common/app.rs @@ -131,6 +131,10 @@ impl TestApp { b"test-secret-that-is-at-least-32-bytes-long", ), plugin_registry: std::sync::Arc::new(happyview::plugin::PluginRegistry::new()), + wasm_runtime: std::sync::Arc::new( + happyview::plugin::WasmRuntime::new().expect("wasm runtime"), + ), + attestation_signer: None, }; let router = server::router(state.clone()); diff --git a/tests/fixtures/test_plugin/.gitignore b/tests/fixtures/test_plugin/.gitignore new file mode 100644 index 0000000..b83d222 --- /dev/null +++ b/tests/fixtures/test_plugin/.gitignore @@ -0,0 +1 @@ +/target/ diff --git a/tests/fixtures/test_plugin/Cargo.lock b/tests/fixtures/test_plugin/Cargo.lock new file mode 100644 index 0000000..aa024b4 --- /dev/null +++ b/tests/fixtures/test_plugin/Cargo.lock @@ -0,0 +1,7 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "test-plugin" +version = "0.1.0" diff --git a/tests/fixtures/test_plugin/Cargo.toml b/tests/fixtures/test_plugin/Cargo.toml new file mode 100644 index 0000000..6420be3 --- /dev/null +++ b/tests/fixtures/test_plugin/Cargo.toml @@ -0,0 +1,11 @@ +[package] +name = "test-plugin" +version = "0.1.0" +edition = "2021" + +[lib] +crate-type = ["cdylib"] + +[profile.release] +opt-level = "s" +lto = true diff --git a/tests/fixtures/test_plugin/src/lib.rs b/tests/fixtures/test_plugin/src/lib.rs new file mode 100644 index 0000000..edad480 --- /dev/null +++ b/tests/fixtures/test_plugin/src/lib.rs @@ -0,0 +1,105 @@ +// Only compile for WASM targets +#![cfg_attr(target_arch = "wasm32", no_std)] +#![allow(static_mut_refs)] + +#[cfg(target_arch = "wasm32")] +extern crate alloc; + +#[cfg(target_arch = "wasm32")] +use core::alloc::{GlobalAlloc, Layout}; + +// Simple bump allocator for WASM +#[cfg(target_arch = "wasm32")] +struct BumpAllocator; + +#[cfg(target_arch = "wasm32")] +const HEAP_SIZE: usize = 65536; +#[cfg(target_arch = "wasm32")] +static mut HEAP: [u8; HEAP_SIZE] = [0; HEAP_SIZE]; +#[cfg(target_arch = "wasm32")] +static mut HEAP_POS: usize = 0; + +#[cfg(target_arch = "wasm32")] +unsafe impl GlobalAlloc for BumpAllocator { + unsafe fn alloc(&self, layout: Layout) -> *mut u8 { + let size = layout.size(); + let align = layout.align(); + + // Align up + let pos = (HEAP_POS + align - 1) & !(align - 1); + if pos + size > HEAP_SIZE { + return core::ptr::null_mut(); + } + + HEAP_POS = pos + size; + HEAP.as_mut_ptr().add(pos) + } + + unsafe fn dealloc(&self, _ptr: *mut u8, _layout: Layout) { + // No-op for bump allocator + } +} + +#[cfg(target_arch = "wasm32")] +#[global_allocator] +static ALLOCATOR: BumpAllocator = BumpAllocator; + +#[cfg(target_arch = "wasm32")] +#[panic_handler] +fn panic(_info: &core::panic::PanicInfo) -> ! { + loop {} +} + +// Memory exports +#[no_mangle] +pub extern "C" fn alloc(size: u32) -> u32 { + let layout = Layout::from_size_align(size as usize, 1).unwrap(); + unsafe { ALLOCATOR.alloc(layout) as u32 } +} + +#[no_mangle] +pub extern "C" fn dealloc(_ptr: u32, _size: u32) { + // No-op for bump allocator +} + +// Helper to return a string as packed i64: (ptr << 32) | len +fn return_json(s: &str) -> i64 { + let ptr = alloc(s.len() as u32); + if ptr == 0 { + return 0; + } + unsafe { + core::ptr::copy_nonoverlapping(s.as_ptr(), ptr as *mut u8, s.len()); + } + ((ptr as i64) << 32) | (s.len() as i64) +} + +#[no_mangle] +pub extern "C" fn plugin_info() -> i64 { + return_json(r#"{"ok":{"id":"test","name":"Test Plugin","version":"1.0.0","api_version":"1","required_secrets":[],"icon_url":null,"config_schema":null}}"#) +} + +#[no_mangle] +pub extern "C" fn get_authorize_url(_ptr: u32, _len: u32) -> i64 { + return_json(r#"{"ok":"https://example.com/oauth?state=test"}"#) +} + +#[no_mangle] +pub extern "C" fn handle_callback(_ptr: u32, _len: u32) -> i64 { + return_json(r#"{"ok":{"access_token":"test-token","token_type":"Bearer","expires_at":null,"refresh_token":null}}"#) +} + +#[no_mangle] +pub extern "C" fn refresh_tokens(_ptr: u32, _len: u32) -> i64 { + return_json(r#"{"ok":{"access_token":"refreshed-token","token_type":"Bearer","expires_at":null,"refresh_token":null}}"#) +} + +#[no_mangle] +pub extern "C" fn get_profile(_ptr: u32, _len: u32) -> i64 { + return_json(r#"{"ok":{"account_id":"12345","display_name":"Test User","profile_url":null,"avatar_url":null}}"#) +} + +#[no_mangle] +pub extern "C" fn sync_account(_ptr: u32, _len: u32) -> i64 { + return_json(r#"{"ok":[]}"#) +} diff --git a/tests/lua_atproto_api.rs b/tests/lua_atproto_api.rs index 30f6d66..33b29b2 100644 --- a/tests/lua_atproto_api.rs +++ b/tests/lua_atproto_api.rs @@ -85,6 +85,10 @@ async fn test_state_with_pool(pool: sqlx::AnyPool, backend: DatabaseBackend) -> oauth: std::sync::Arc::new(oauth), cookie_key: axum_extra::extract::cookie::Key::derive_from(b"test-secret"), plugin_registry: std::sync::Arc::new(happyview::plugin::PluginRegistry::new()), + wasm_runtime: std::sync::Arc::new( + happyview::plugin::WasmRuntime::new().expect("wasm runtime"), + ), + attestation_signer: None, } } diff --git a/tests/lua_db_api.rs b/tests/lua_db_api.rs index 14b9e62..8fe4b4f 100644 --- a/tests/lua_db_api.rs +++ b/tests/lua_db_api.rs @@ -88,6 +88,10 @@ async fn test_state_with_pool(pool: sqlx::AnyPool, backend: DatabaseBackend) -> oauth: std::sync::Arc::new(oauth), cookie_key: axum_extra::extract::cookie::Key::derive_from(b"test-secret"), plugin_registry: std::sync::Arc::new(happyview::plugin::PluginRegistry::new()), + wasm_runtime: std::sync::Arc::new( + happyview::plugin::WasmRuntime::new().expect("wasm runtime"), + ), + attestation_signer: None, } } diff --git a/tests/plugin_executor.rs b/tests/plugin_executor.rs new file mode 100644 index 0000000..71cc06a --- /dev/null +++ b/tests/plugin_executor.rs @@ -0,0 +1,224 @@ +// tests/plugin_executor.rs + +use happyview::db::DatabaseBackend; +use happyview::lexicon::LexiconRegistry; +use happyview::plugin::{ + ExecutionError, LoadedPlugin, PluginExecutor, PluginInfo, PluginRegistry, PluginSource, + WasmRuntime, +}; +use std::collections::HashMap; +use std::sync::Arc; + +type Secrets = HashMap; + +async fn create_test_executor() -> (PluginExecutor, Arc) { + // Create in-memory database + sqlx::any::install_default_drivers(); + let db = sqlx::AnyPool::connect("sqlite::memory:") + .await + .expect("Failed to create test database"); + + let runtime = Arc::new(WasmRuntime::new().expect("Failed to create runtime")); + let registry = Arc::new(PluginRegistry::new()); + let lexicons = Arc::new(LexiconRegistry::new()); + let http_client = reqwest::Client::new(); + + let executor = PluginExecutor::new( + runtime, + registry.clone(), + db, + DatabaseBackend::Sqlite, + http_client, + lexicons, + ); + + (executor, registry) +} + +fn load_test_plugin() -> LoadedPlugin { + let wasm_bytes = std::fs::read( + "tests/fixtures/test_plugin/target/wasm32-unknown-unknown/release/test_plugin.wasm", + ) + .expect( + "Test plugin not built. Run: cd tests/fixtures/test_plugin && cargo build --target wasm32-unknown-unknown --release", + ); + + LoadedPlugin { + info: PluginInfo { + id: "test".into(), + name: "Test Plugin".into(), + version: "1.0.0".into(), + api_version: "1".into(), + icon_url: None, + required_secrets: vec![], + config_schema: None, + }, + source: PluginSource::File { + path: "tests/fixtures/test_plugin".into(), + }, + wasm_bytes, + } +} + +#[tokio::test] +async fn test_plugin_info() { + let (executor, registry) = create_test_executor().await; + let plugin = load_test_plugin(); + registry.register(plugin).await; + + let mut instance = executor + .instantiate( + "test", + "user:did:plc:test", + Secrets::new(), + serde_json::Value::Null, + ) + .await + .expect("Failed to instantiate"); + + let info = instance + .call_plugin_info() + .await + .expect("Failed to get info"); + + assert_eq!(info.id, "test"); + assert_eq!(info.name, "Test Plugin"); + assert_eq!(info.version, "1.0.0"); +} + +#[tokio::test] +async fn test_get_authorize_url() { + let (executor, registry) = create_test_executor().await; + let plugin = load_test_plugin(); + registry.register(plugin).await; + + let mut instance = executor + .instantiate("test", "state:123", Secrets::new(), serde_json::Value::Null) + .await + .expect("Failed to instantiate"); + + let url = instance + .call_get_authorize_url( + "state123", + "https://app.example/callback", + &serde_json::Value::Null, + ) + .await + .expect("Failed to get URL"); + + assert!(url.starts_with("https://")); +} + +#[tokio::test] +async fn test_handle_callback() { + let (executor, registry) = create_test_executor().await; + let plugin = load_test_plugin(); + registry.register(plugin).await; + + let mut instance = executor + .instantiate( + "test", + "user:did:plc:test", + Secrets::new(), + serde_json::Value::Null, + ) + .await + .expect("Failed to instantiate"); + + let tokens = instance + .call_handle_callback("code123", "state123", &serde_json::Value::Null) + .await + .expect("Failed to handle callback"); + + assert_eq!(tokens.access_token, "test-token"); + assert_eq!(tokens.token_type, "Bearer"); +} + +#[tokio::test] +async fn test_refresh_tokens() { + let (executor, registry) = create_test_executor().await; + let plugin = load_test_plugin(); + registry.register(plugin).await; + + let mut instance = executor + .instantiate( + "test", + "user:did:plc:test", + Secrets::new(), + serde_json::Value::Null, + ) + .await + .expect("Failed to instantiate"); + + let tokens = instance + .call_refresh_tokens("old-refresh-token", &serde_json::Value::Null) + .await + .expect("Failed to refresh tokens"); + + assert_eq!(tokens.access_token, "refreshed-token"); +} + +#[tokio::test] +async fn test_get_profile() { + let (executor, registry) = create_test_executor().await; + let plugin = load_test_plugin(); + registry.register(plugin).await; + + let mut instance = executor + .instantiate( + "test", + "user:did:plc:test", + Secrets::new(), + serde_json::Value::Null, + ) + .await + .expect("Failed to instantiate"); + + let profile = instance + .call_get_profile("test-token", &serde_json::Value::Null) + .await + .expect("Failed to get profile"); + + assert_eq!(profile.account_id, "12345"); + assert_eq!(profile.display_name, Some("Test User".into())); +} + +#[tokio::test] +async fn test_sync_account() { + let (executor, registry) = create_test_executor().await; + let plugin = load_test_plugin(); + registry.register(plugin).await; + + let mut instance = executor + .instantiate( + "test", + "user:did:plc:test", + Secrets::new(), + serde_json::Value::Null, + ) + .await + .expect("Failed to instantiate"); + + let records = instance + .call_sync_account("test-token", &serde_json::Value::Null) + .await + .expect("Failed to sync account"); + + assert!(records.is_empty()); // Test plugin returns empty array +} + +#[tokio::test] +async fn test_plugin_not_found() { + let (executor, _registry) = create_test_executor().await; + + let result = executor + .instantiate( + "nonexistent", + "scope", + Secrets::new(), + serde_json::Value::Null, + ) + .await; + + assert!(matches!(result, Err(ExecutionError::PluginNotFound(_)))); +} -- 2.51.2