diff --git a/src/lua/execute.rs b/src/lua/execute.rs index 1417284..2f9771d 100644 --- a/src/lua/execute.rs +++ b/src/lua/execute.rs @@ -15,6 +15,7 @@ use crate::repo; use super::context; use super::db_api; +use super::http_api; use super::record; use super::sandbox; @@ -110,6 +111,28 @@ pub async fn execute_procedure_script( return Err(AppError::Internal(error_message)); } + if let Err(e) = http_api::register_http_api(&lua, state_arc.clone()) { + let error_message = format!("failed to register http API: {e}"); + log_event( + &state.db, + EventLog { + event_type: "script.error".to_string(), + severity: Severity::Error, + actor_did: Some(claims.did().to_string()), + subject: Some(method.to_string()), + detail: serde_json::json!({ + "error": error_message, + "script_source": script_source, + "input": input_json, + "caller_did": claims.did(), + "method": method, + }), + }, + ) + .await; + return Err(AppError::Internal(error_message)); + } + if let Err(e) = record::register_record_api(&lua, state_arc, claims_arc, session_arc) { let error_message = format!("failed to register Record API: {e}"); log_event( @@ -318,7 +341,7 @@ pub async fn execute_query_script( let state_arc = Arc::new(state.clone()); - if let Err(e) = db_api::register_db_api(&lua, state_arc) { + if let Err(e) = db_api::register_db_api(&lua, state_arc.clone()) { let error_message = format!("failed to register db API: {e}"); log_event( &state.db, @@ -338,6 +361,26 @@ pub async fn execute_query_script( return Err(AppError::Internal(error_message)); } + if let Err(e) = http_api::register_http_api(&lua, state_arc) { + let error_message = format!("failed to register http API: {e}"); + log_event( + &state.db, + EventLog { + event_type: "script.error".to_string(), + severity: Severity::Error, + actor_did: None, + subject: Some(method.to_string()), + detail: serde_json::json!({ + "error": error_message, + "script_source": script_source, + "method": method, + }), + }, + ) + .await; + return Err(AppError::Internal(error_message)); + } + if let Err(e) = context::set_query_context(&lua, method, params, collection) { let error_message = format!("failed to set context: {e}"); log_event( diff --git a/src/lua/http_api.rs b/src/lua/http_api.rs new file mode 100644 index 0000000..0019a0d --- /dev/null +++ b/src/lua/http_api.rs @@ -0,0 +1,285 @@ +use mlua::{Lua, Result as LuaResult}; +use reqwest::Method; +use std::sync::Arc; + +use crate::AppState; + +/// Register the `http` table with async HTTP request functions. +pub fn register_http_api(lua: &Lua, state: Arc) -> LuaResult<()> { + let http_table = lua.create_table()?; + + let methods = [ + ("get", Method::GET), + ("post", Method::POST), + ("put", Method::PUT), + ("patch", Method::PATCH), + ("delete", Method::DELETE), + ("head", Method::HEAD), + ]; + + for (name, method) in methods { + let state_clone = state.clone(); + let func = + lua.create_async_function(move |lua, (url, opts): (String, Option)| { + let state = state_clone.clone(); + let method = method.clone(); + async move { + let mut builder = state.http.request(method.clone(), &url); + + if let Some(ref opts) = opts { + if let Ok(headers_table) = opts.get::("headers") { + for pair in headers_table.pairs::() { + let (key, value) = pair?; + builder = builder.header(key, value); + } + } + + if method != Method::GET + && method != Method::HEAD + && let Ok(body) = opts.get::("body") + { + builder = builder.body(body); + } + } + + let response = builder + .send() + .await + .map_err(|e| mlua::Error::runtime(format!("HTTP request failed: {e}")))?; + + let status = response.status().as_u16(); + + let headers_table = lua.create_table()?; + for (key, value) in response.headers() { + if let Ok(v) = value.to_str() { + headers_table.set(key.as_str().to_lowercase(), v.to_string())?; + } + } + + let body = if method == Method::HEAD { + String::new() + } else { + response.text().await.map_err(|e| { + mlua::Error::runtime(format!("HTTP read body failed: {e}")) + })? + }; + + let result = lua.create_table()?; + result.set("status", status)?; + result.set("body", body)?; + result.set("headers", headers_table)?; + + Ok(mlua::Value::Table(result)) + } + })?; + http_table.set(name, func)?; + } + + lua.globals().set("http", http_table)?; + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::Config; + use crate::lexicon::LexiconRegistry; + use tokio::sync::watch; + + fn test_state() -> AppState { + let config = Config { + host: "127.0.0.1".into(), + port: 3000, + database_url: String::new(), + aip_url: String::new(), + aip_public_url: String::new(), + tap_url: String::new(), + tap_admin_password: None, + relay_url: String::new(), + plc_url: String::new(), + static_dir: String::new(), + event_log_retention_days: 30, + }; + let (tx, _) = watch::channel(vec![]); + AppState { + config, + http: reqwest::Client::new(), + db: sqlx::PgPool::connect_lazy("postgres://localhost/fake").unwrap(), + lexicons: LexiconRegistry::new(), + collections_tx: tx, + } + } + + fn setup(state: &AppState) -> Lua { + let lua = Lua::new(); + register_http_api(&lua, Arc::new(state.clone())).unwrap(); + lua + } + + #[tokio::test] + async fn get_returns_status_and_body() { + let mock = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/test")) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_string("hello")) + .mount(&mock) + .await; + + let state = test_state(); + let lua = setup(&state); + let chunk = format!(r#"return http.get("{}/test")"#, mock.uri()); + let result: mlua::Table = lua.load(chunk).eval_async().await.unwrap(); + assert_eq!(result.get::("status").unwrap(), 200); + assert_eq!(result.get::("body").unwrap(), "hello"); + } + + #[tokio::test] + async fn get_returns_headers() { + let mock = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/h")) + .respond_with( + wiremock::ResponseTemplate::new(200) + .insert_header("X-Custom", "test-value") + .set_body_string(""), + ) + .mount(&mock) + .await; + + let state = test_state(); + let lua = setup(&state); + let chunk = format!(r#"return http.get("{}/h")"#, mock.uri()); + let result: mlua::Table = lua.load(chunk).eval_async().await.unwrap(); + let headers: mlua::Table = result.get("headers").unwrap(); + assert_eq!(headers.get::("x-custom").unwrap(), "test-value"); + } + + #[tokio::test] + async fn post_sends_body_and_headers() { + let mock = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("POST")) + .and(wiremock::matchers::path("/post")) + .and(wiremock::matchers::header( + "content-type", + "application/json", + )) + .and(wiremock::matchers::body_string(r#"{"k":"v"}"#)) + .respond_with(wiremock::ResponseTemplate::new(201).set_body_string("created")) + .mount(&mock) + .await; + + let state = test_state(); + let lua = setup(&state); + let chunk = format!( + r#"return http.post("{}/post", {{ + body = '{{"k":"v"}}', + headers = {{ ["content-type"] = "application/json" }} + }})"#, + mock.uri() + ); + let result: mlua::Table = lua.load(chunk).eval_async().await.unwrap(); + assert_eq!(result.get::("status").unwrap(), 201); + assert_eq!(result.get::("body").unwrap(), "created"); + } + + #[tokio::test] + async fn head_returns_empty_body() { + let mock = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("HEAD")) + .and(wiremock::matchers::path("/head")) + .respond_with(wiremock::ResponseTemplate::new(204)) + .mount(&mock) + .await; + + let state = test_state(); + let lua = setup(&state); + let chunk = format!(r#"return http.head("{}/head")"#, mock.uri()); + let result: mlua::Table = lua.load(chunk).eval_async().await.unwrap(); + assert_eq!(result.get::("status").unwrap(), 204); + assert_eq!(result.get::("body").unwrap(), ""); + } + + #[tokio::test] + async fn put_sends_body() { + let mock = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("PUT")) + .and(wiremock::matchers::path("/put")) + .and(wiremock::matchers::body_string("updated")) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_string("ok")) + .mount(&mock) + .await; + + let state = test_state(); + let lua = setup(&state); + let chunk = format!( + r#"return http.put("{}/put", {{ body = "updated" }})"#, + mock.uri() + ); + let result: mlua::Table = lua.load(chunk).eval_async().await.unwrap(); + assert_eq!(result.get::("status").unwrap(), 200); + } + + #[tokio::test] + async fn delete_works() { + let mock = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("DELETE")) + .and(wiremock::matchers::path("/del")) + .respond_with(wiremock::ResponseTemplate::new(204).set_body_string("")) + .mount(&mock) + .await; + + let state = test_state(); + let lua = setup(&state); + let chunk = format!(r#"return http.delete("{}/del")"#, mock.uri()); + let result: mlua::Table = lua.load(chunk).eval_async().await.unwrap(); + assert_eq!(result.get::("status").unwrap(), 204); + } + + #[tokio::test] + async fn patch_sends_body() { + let mock = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("PATCH")) + .and(wiremock::matchers::path("/patch")) + .and(wiremock::matchers::body_string("patched")) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_string("ok")) + .mount(&mock) + .await; + + let state = test_state(); + let lua = setup(&state); + let chunk = format!( + r#"return http.patch("{}/patch", {{ body = "patched" }})"#, + mock.uri() + ); + let result: mlua::Table = lua.load(chunk).eval_async().await.unwrap(); + assert_eq!(result.get::("status").unwrap(), 200); + } + + #[tokio::test] + async fn get_without_opts_works() { + let mock = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/simple")) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_string("ok")) + .mount(&mock) + .await; + + let state = test_state(); + let lua = setup(&state); + let chunk = format!(r#"return http.get("{}/simple")"#, mock.uri()); + let result: mlua::Table = lua.load(chunk).eval_async().await.unwrap(); + assert_eq!(result.get::("status").unwrap(), 200); + assert_eq!(result.get::("body").unwrap(), "ok"); + } + + #[tokio::test] + async fn invalid_url_returns_error() { + let state = test_state(); + let lua = setup(&state); + let result: Result = lua + .load(r#"return http.get("http://0.0.0.0:1/nope")"#) + .eval_async() + .await; + assert!(result.is_err()); + } +} diff --git a/src/lua/mod.rs b/src/lua/mod.rs index c10a318..8954354 100644 --- a/src/lua/mod.rs +++ b/src/lua/mod.rs @@ -1,6 +1,7 @@ mod context; mod db_api; mod execute; +mod http_api; mod record; pub(crate) mod sandbox; mod tid;