From 013c1312f8fcce73fcd1091d22feb67c7780359c Mon Sep 17 00:00:00 2001 From: stefan-gorules <127550877+stefan-gorules@users.noreply.github.com> Date: Thu, 11 Sep 2025 09:58:16 +0200 Subject: [PATCH] feat: http auth (#393) * temp: http auth * add better error handling * improve Cargo.toml format * add enable/disable for http auth * make aws region optional * add fallback for aws --- bindings/nodejs/src/config.rs | 5 + core/engine/Cargo.toml | 8 +- core/engine/src/config.rs | 2 + core/engine/src/nodes/context.rs | 4 +- core/engine/src/nodes/function/v2/mod.rs | 77 ++++--- .../src/nodes/function/v2/module/http/auth.rs | 186 +++++++++++++++++ .../v2/module/{http.rs => http/mod.rs} | 193 ++++++++---------- 7 files changed, 338 insertions(+), 137 deletions(-) create mode 100644 core/engine/src/nodes/function/v2/module/http/auth.rs rename core/engine/src/nodes/function/v2/module/{http.rs => http/mod.rs} (50%) diff --git a/bindings/nodejs/src/config.rs b/bindings/nodejs/src/config.rs index 7cb24526..dfdb0515 100644 --- a/bindings/nodejs/src/config.rs +++ b/bindings/nodejs/src/config.rs @@ -6,6 +6,7 @@ use zen_engine::ZEN_CONFIG; pub struct ZenConfig { pub nodes_in_context: Option, pub function_timeout_millis: Option, + pub http_auth: Option, } #[allow(dead_code)] @@ -20,4 +21,8 @@ pub fn override_config(config: ZenConfig) { .function_timeout_millis .store(val as u64, Ordering::Relaxed); } + + if let Some(val) = config.http_auth { + ZEN_CONFIG.http_auth.store(val, Ordering::Relaxed); + } } diff --git a/core/engine/Cargo.toml b/core/engine/Cargo.toml index af07e2d6..00ce7a10 100644 --- a/core/engine/Cargo.toml +++ b/core/engine/Cargo.toml @@ -23,13 +23,19 @@ json_dotpath = { workspace = true } rust_decimal = { workspace = true, features = ["maths-nopanic"] } fixedbitset = "0.5" tokio = { workspace = true, features = ["sync", "time"] } +http = { version = "1.3" } reqwest = { version = "0.12", features = ["json", "rustls-tls"], default-features = false } rquickjs = { version = "0.9", features = ["macro", "loader", "rust-alloc", "futures", "either", "properties"] } +rquickjs-serde = { version = "0.1" } jsonschema = "0.29" zen-types = { path = "../types", version = "0.50.2" } zen-expression = { path = "../expression", version = "0.50.2" } zen-tmpl = { path = "../template", version = "0.50.2" } -log = "0.4.27" + +# HTTP Auth IAM +async-trait = "0.1" +reqsign = { version = "0.17", features = ["aws", "azure", "google", "default-context"] } +sha2 = "0.10" [dev-dependencies] chrono = { workspace = true } diff --git a/core/engine/src/config.rs b/core/engine/src/config.rs index 7762c0a3..424eefe5 100644 --- a/core/engine/src/config.rs +++ b/core/engine/src/config.rs @@ -5,6 +5,7 @@ use std::sync::atomic::{AtomicBool, AtomicU64}; pub struct ZenConfig { pub nodes_in_context: AtomicBool, pub function_timeout_millis: AtomicU64, + pub http_auth: AtomicBool, } impl Default for ZenConfig { @@ -12,6 +13,7 @@ impl Default for ZenConfig { Self { nodes_in_context: AtomicBool::new(true), function_timeout_millis: AtomicU64::new(5_000), + http_auth: AtomicBool::new(true), } } } diff --git a/core/engine/src/nodes/context.rs b/core/engine/src/nodes/context.rs index 65a9284b..375e4c98 100644 --- a/core/engine/src/nodes/context.rs +++ b/core/engine/src/nodes/context.rs @@ -73,7 +73,7 @@ where }) } - fn make_error(&self, error: Error) -> NodeError + pub(crate) fn make_error(&self, error: Error) -> NodeError where Error: Into>, { @@ -319,6 +319,7 @@ pub struct NodeContextConfig { pub nodes_in_context: bool, pub max_depth: u8, pub function_timeout_millis: u64, + pub http_auth: bool, } impl Default for NodeContextConfig { @@ -327,6 +328,7 @@ impl Default for NodeContextConfig { trace: false, nodes_in_context: ZEN_CONFIG.nodes_in_context.load(Ordering::Relaxed), function_timeout_millis: ZEN_CONFIG.function_timeout_millis.load(Ordering::Relaxed), + http_auth: ZEN_CONFIG.http_auth.load(Ordering::Relaxed), max_depth: 5, } } diff --git a/core/engine/src/nodes/function/v2/mod.rs b/core/engine/src/nodes/function/v2/mod.rs index 016fa631..a9436e72 100644 --- a/core/engine/src/nodes/function/v2/mod.rs +++ b/core/engine/src/nodes/function/v2/mod.rs @@ -1,13 +1,13 @@ use std::ops::Deref; -use std::time::Duration; +use std::time::{Duration, Instant}; use crate::nodes::definition::NodeHandler; -use crate::nodes::function::v2::error::FunctionResult; +use crate::nodes::function::v2::error::{FunctionError, FunctionResult}; use crate::nodes::function::v2::function::{Function, HandlerResponse}; use crate::nodes::function::v2::module::console::Log; use crate::nodes::function::v2::serde::JsValue; use crate::nodes::result::NodeResult; -use crate::nodes::{NodeContext, NodeContextExt}; +use crate::nodes::{NodeContext, NodeError}; use ::serde::{Deserialize, Serialize}; use rquickjs::{async_with, CatchResultExt, Object}; use serde_json::json; @@ -28,8 +28,7 @@ impl NodeHandler for FunctionV2NodeHandler { type TraceData = FunctionV2Trace; async fn handle(&self, ctx: NodeContext) -> NodeResult { - let start = std::time::Instant::now(); - + let start = Instant::now(); if ctx.node.omit_nodes { ctx.input.dot_remove("$nodes"); } @@ -45,41 +44,35 @@ impl NodeHandler for FunctionV2NodeHandler { .set_interrupt_handler(Some(interrupt_handler)) .await; + let function_context = FunctionContext { + start, + context: &ctx, + function: &function, + }; + self.attach_globals(function, &ctx) .await - .node_context(&ctx)?; + .function_context(&function_context) + .await?; function .register_module(&module_name, ctx.node.source.deref()) .await - .node_context(&ctx)?; + .function_context(&function_context) + .await?; let response_result = function .call_handler(&module_name, JsValue(ctx.input.clone())) .await; - match response_result { - Ok(response) => { - function.runtime().set_interrupt_handler(None).await; - ctx.trace(|t| { - t.log = response.logs.clone(); - }); + function.runtime().set_interrupt_handler(None).await; - ctx.success(response.data) - } - Err(e) => { - let log = function.extract_logs().await; - ctx.trace(|t| { - t.log = log; - t.log.push(Log { - lines: vec![json!(e.to_string()).to_string()], - ms_since_run: start.elapsed().as_millis() as usize, - }); - }); + let response = response_result.function_context(&function_context).await?; + ctx.trace(|t| { + t.log = response.logs.clone(); + }); - ctx.error(e) - } - } + ctx.success(response.data) } } @@ -115,3 +108,33 @@ pub struct FunctionResponse { performance: String, data: Option, } + +struct FunctionContext<'a> { + context: &'a NodeContext, + function: &'a Function, + start: Instant, +} + +trait FunctionErrorExt { + async fn function_context(self, ctx: &FunctionContext) -> Result; +} + +impl FunctionErrorExt for Result { + async fn function_context(self, c: &FunctionContext<'_>) -> Result { + match self { + Ok(ok) => Ok(ok), + Err(err) => { + let log = c.function.extract_logs().await; + c.context.trace(|t| { + t.log = log; + t.log.push(Log { + lines: vec![json!(err.to_string()).to_string()], + ms_since_run: c.start.elapsed().as_millis() as usize, + }); + }); + + Err(c.context.make_error(err)) + } + } + } +} diff --git a/core/engine/src/nodes/function/v2/module/http/auth.rs b/core/engine/src/nodes/function/v2/module/http/auth.rs new file mode 100644 index 00000000..dd6ad7b7 --- /dev/null +++ b/core/engine/src/nodes/function/v2/module/http/auth.rs @@ -0,0 +1,186 @@ +use crate::nodes::function::v2::error::ResultExt; +use crate::nodes::function::v2::module::http::HttpConfig; +use ::http::Request as HttpRequest; +use anyhow::Context; +use async_trait::async_trait; +use http::HeaderValue; +use reqsign::{aws, azure, google}; +use reqwest::{Body, Request}; +use rquickjs::{Ctx, FromJs, Value}; +use serde::{Deserialize, Deserializer}; +use sha2::{Digest, Sha256}; +use std::fmt::Debug; +use std::ops::Deref; +use std::sync::{Arc, OnceLock}; + +#[derive(Deserialize, Clone)] +#[serde(tag = "type", rename_all = "camelCase")] +pub(crate) enum HttpConfigAuth { + #[serde(rename = "iam")] + Iam(IamAuth), +} + +#[derive(Deserialize, Clone)] +#[serde(tag = "provider", rename_all = "camelCase")] +pub(crate) enum IamAuth { + Aws(AwsIamAuth), + Azure(AzureIamAuth), + Gcp(GcpIamAuth), +} + +impl<'js> FromJs<'js> for HttpConfig { + fn from_js(ctx: &Ctx<'js>, value: Value<'js>) -> rquickjs::Result { + rquickjs_serde::from_value(value).or_throw(&ctx) + } +} + +#[derive(Debug)] +struct CachedProvider(Arc) +where + Provider: reqsign::ProvideCredential + Debug; + +#[async_trait] +impl reqsign::ProvideCredential for CachedProvider +where + Provider: reqsign::ProvideCredential + Debug, +{ + type Credential = Provider::Credential; + + async fn provide_credential( + &self, + ctx: &reqsign::Context, + ) -> reqsign::Result> { + self.0.provide_credential(ctx).await + } +} + +impl Clone for CachedProvider +where + Provider: reqsign::ProvideCredential + Debug, +{ + fn clone(&self) -> Self { + CachedProvider(self.0.clone()) + } +} + +#[derive(Deserialize, Clone)] +#[serde(rename_all = "camelCase")] +pub(crate) struct AwsIamAuth { + region: AwsRegion, + service: Arc, +} + +#[derive(Clone)] +struct AwsRegion(pub Arc); + +impl<'de> Deserialize<'de> for AwsRegion { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + let region_str = Option::>::deserialize(deserializer)?; + match region_str { + Some(region) => Ok(AwsRegion(region)), + None => { + static AWS_REGION_ENV: OnceLock>> = OnceLock::new(); + let aws_region_opt = AWS_REGION_ENV.get_or_init(|| { + std::env::var("AWS_REGION") + .or_else(|_| std::env::var("AWS_DEFAULT_REGION")) + .map(Arc::from) + .ok() + }); + let Some(aws_region) = aws_region_opt else { + return Err(serde::de::Error::custom( + "AWS_REGION environment variable is missing - region parameter is required", + )); + }; + + Ok(AwsRegion(aws_region.clone())) + } + } + } +} + +impl AwsIamAuth { + pub async fn build_request(&self, http_request: HttpRequest) -> anyhow::Result { + static CACHED_PROVIDER: OnceLock> = + OnceLock::new(); + let provider = CACHED_PROVIDER + .get_or_init(|| CachedProvider(Arc::new(aws::DefaultCredentialProvider::new()))) + .clone(); + + let signer = aws::default_signer(self.service.deref(), self.region.0.deref()) + .with_credential_provider(provider); + + let (mut parts, body) = http_request.into_parts(); + let payload_hash_opt = body + .as_bytes() + .map(|body_bytes| format!("{:x}", Sha256::digest(&body_bytes))); + + if let Some(payload_hash) = payload_hash_opt { + parts.headers.insert( + "x-amz-content-sha256", + HeaderValue::from_str(payload_hash.as_str())?, + ); + } + + signer + .sign(&mut parts, None) + .await + .context("Failed to sign request body")?; + + let new_http_request = HttpRequest::from_parts(parts, body); + Request::try_from(new_http_request).context("Failed to create request") + } +} + +#[derive(Deserialize, Clone)] +#[serde(rename_all = "camelCase")] +pub(crate) struct GcpIamAuth { + service: Arc, +} + +impl GcpIamAuth { + pub async fn build_request(&self, http_request: HttpRequest) -> anyhow::Result { + static CACHED_PROVIDER: OnceLock> = + OnceLock::new(); + let provider = CACHED_PROVIDER + .get_or_init(|| CachedProvider(Arc::new(google::DefaultCredentialProvider::new()))) + .clone(); + + let signer = + google::default_signer(self.service.deref()).with_credential_provider(provider); + let (mut parts, body) = http_request.into_parts(); + + signer + .sign(&mut parts, None) + .await + .context("Failed to sign request body")?; + let new_http_request = HttpRequest::from_parts(parts, body); + Request::try_from(new_http_request).context("Failed to create request") + } +} + +#[derive(Deserialize, Clone)] +#[serde(rename_all = "camelCase")] +pub(crate) struct AzureIamAuth; + +impl AzureIamAuth { + pub async fn build_request(&self, http_request: HttpRequest) -> anyhow::Result { + static CACHED_PROVIDER: OnceLock> = + OnceLock::new(); + let provider = CACHED_PROVIDER + .get_or_init(|| CachedProvider(Arc::new(azure::DefaultCredentialProvider::new()))) + .clone(); + + let signer = azure::default_signer().with_credential_provider(provider); + let (mut parts, body) = http_request.into_parts(); + + signer + .sign(&mut parts, None) + .await + .context("Failed to sign request body")?; + let new_http_request = HttpRequest::from_parts(parts, body); + Request::try_from(new_http_request).context("Failed to create request") + } +} diff --git a/core/engine/src/nodes/function/v2/module/http.rs b/core/engine/src/nodes/function/v2/module/http/mod.rs similarity index 50% rename from core/engine/src/nodes/function/v2/module/http.rs rename to core/engine/src/nodes/function/v2/module/http/mod.rs index 7b6731c1..7e9c9b61 100644 --- a/core/engine/src/nodes/function/v2/module/http.rs +++ b/core/engine/src/nodes/function/v2/module/http/mod.rs @@ -1,12 +1,18 @@ +mod auth; + use crate::nodes::function::v2::error::ResultExt; use crate::nodes::function::v2::module::export_default; +use crate::nodes::function::v2::module::http::auth::{HttpConfigAuth, IamAuth}; use crate::nodes::function::v2::serde::JsValue; -use reqwest::header::{HeaderMap, HeaderName}; -use reqwest::Method; +use crate::ZEN_CONFIG; +use ::http::Request as HttpRequest; +use ahash::HashMap; +use reqwest::{Body, Method, Request, Url}; use rquickjs::module::{Declarations, Exports, ModuleDef}; use rquickjs::prelude::{Async, Func, Opt}; -use rquickjs::{CatchResultExt, Ctx, FromJs, IntoAtom, IntoJs, Object, Value}; -use std::str::FromStr; +use rquickjs::{CatchResultExt, Ctx, IntoAtom, IntoJs, Object, Value}; +use serde::{Deserialize, Deserializer}; +use std::sync::atomic::Ordering; use std::sync::OnceLock; use zen_expression::variable::Variable; @@ -36,24 +42,55 @@ async fn execute_http<'js>( ) -> rquickjs::Result> { static HTTP_CLIENT: OnceLock = OnceLock::new(); let client = HTTP_CLIENT.get_or_init(|| reqwest::Client::new()).clone(); - let mut builder = client.request(method, url); - if let Some(data) = data { - builder = builder.json(&data.0); - } - if let Some(config) = config { - builder = builder - .headers(config.headers) - .query(config.params.as_slice()); - - if let Some(data) = config.data { - if !matches!(data.0, Variable::Null) { - builder = builder.json(&data.0); - } + let mut url = Url::parse(&url).or_throw(&ctx)?; + if let Some(config) = &config { + for (k, v) in &config.params { + url.query_pairs_mut().append_pair(k.as_str(), v.0.as_str()); } } - let response = builder.send().await.or_throw(&ctx)?; + let mut request_builder = HttpRequest::builder().method(method).uri(url.as_str()); + if let Some(config) = &config { + for (k, v) in &config.headers { + request_builder = request_builder.header(k.as_str(), v.0.as_str()); + } + } + + let auth_method = config + .as_ref() + .filter(|_| ZEN_CONFIG.http_auth.load(Ordering::Relaxed)) + .and_then(|c| c.auth.clone()); + + let request_data_opt = config + .and_then(|c| c.data) + .and_then(|_| data.map(|d| d.0.to_value())); + + let http_request = match request_data_opt { + None => request_builder.body(Body::default()).or_throw(&ctx)?, + Some(request_data) => { + let request_body_json = serde_json::to_vec(&request_data).or_throw(&ctx)?; + request_builder + .body(Body::from(request_body_json)) + .or_throw(&ctx)? + } + }; + + let request = match auth_method { + Some(HttpConfigAuth::Iam(IamAuth::Aws(config))) => { + config.build_request(http_request).await.or_throw(&ctx)? + } + Some(HttpConfigAuth::Iam(IamAuth::Azure(config))) => { + config.build_request(http_request).await.or_throw(&ctx)? + } + Some(HttpConfigAuth::Iam(IamAuth::Gcp(config))) => { + config.build_request(http_request).await.or_throw(&ctx)? + } + None => Request::try_from(http_request).or_throw(&ctx)?, + }; + + // Apply auth + let response = client.execute(request).await.or_throw(&ctx)?; let status = response.status().as_u16(); let header_object = Object::new(ctx.clone()).catch(&ctx).or_throw(&ctx)?; for (key, value) in response.headers() { @@ -72,96 +109,6 @@ async fn execute_http<'js>( }) } -#[derive(Default)] -pub(crate) struct HttpConfig { - headers: HeaderMap, - params: Vec<(String, String)>, - data: Option, -} - -impl<'js> FromJs<'js> for HttpConfig { - fn from_js(ctx: &Ctx<'js>, value: Value<'js>) -> rquickjs::Result { - let object = value.into_object().or_throw(ctx)?; - let headers_obj: Option> = object.get("headers").or_throw(ctx)?; - let headers = if let Some(headers_obj) = headers_obj { - let mut header_map = HeaderMap::with_capacity(headers_obj.len()); - for result in headers_obj.into_iter() { - let Ok((key, value)) = result else { - continue; - }; - - let value = JsValue::from_js(ctx, value)?; - let str_value = match value.0 { - Variable::Bool(b) => Some(b.to_string()), - Variable::Number(n) => Some(n.to_string()), - Variable::String(s) => Some(s.to_string()), - Variable::Null => None, - Variable::Array(_) => None, - Variable::Object(_) => None, - Variable::Dynamic(_) => None, - }; - - let key_value = key.to_string()?; - let key = HeaderName::from_str(key_value.as_str()).or_throw(&ctx)?; - if let Some(str_value) = str_value { - header_map.insert(key, str_value.parse().or_throw(&ctx)?); - } - } - - header_map - } else { - HeaderMap::default() - }; - - let params_obj: Option> = object.get("params").or_throw(ctx)?; - let params = if let Some(params_obj) = params_obj { - let mut params = Vec::with_capacity(params_obj.len()); - for result in params_obj.into_iter() { - let Ok((key, value)) = result else { - continue; - }; - - let value = JsValue::from_js(ctx, value)?; - let str_value = match value.0 { - Variable::Bool(b) => Some(b.to_string()), - Variable::Number(n) => Some(n.to_string()), - Variable::String(s) => Some(s.to_string()), - Variable::Null => None, - Variable::Array(_) => None, - Variable::Object(_) => None, - Variable::Dynamic(_) => None, - }; - - let key = key.to_string()?; - if let Some(str_value) = str_value { - params.push((key, str_value)); - } - } - - params - } else { - Vec::default() - }; - - let data_obj: Option> = object.get("data").ok(); - let data = if let Some(data_obj) = data_obj { - Some( - JsValue::from_js(&ctx, data_obj) - .catch(&ctx) - .or_throw(&ctx)?, - ) - } else { - None - }; - - Ok(Self { - headers, - params, - data, - }) - } -} - async fn get<'js>( ctx: Ctx<'js>, url: String, @@ -242,3 +189,33 @@ impl ModuleDef for HttpModule { }) } } + +#[derive(Deserialize)] +pub(crate) struct HttpConfig { + #[serde(default)] + headers: HashMap, + #[serde(default)] + params: HashMap, + data: Option, + auth: Option, +} + +#[derive(Debug, Clone)] +pub(crate) struct StringPrimitive(pub String); + +impl<'de> Deserialize<'de> for StringPrimitive { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + let value = serde_json::Value::deserialize(deserializer)?; + let data = match value { + serde_json::Value::Bool(b) => b.to_string(), + serde_json::Value::Number(n) => n.to_string(), + serde_json::Value::String(s) => s, + _ => return Err(serde::de::Error::custom("Value is not a string")), + }; + + Ok(Self(data)) + } +}