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
This commit is contained in:
stefan-gorules
2025-09-11 09:58:16 +02:00
committed by GitHub
parent 91509821be
commit 013c1312f8
7 changed files with 338 additions and 137 deletions
+5
View File
@@ -6,6 +6,7 @@ use zen_engine::ZEN_CONFIG;
pub struct ZenConfig {
pub nodes_in_context: Option<bool>,
pub function_timeout_millis: Option<u32>,
pub http_auth: Option<bool>,
}
#[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);
}
}
+7 -1
View File
@@ -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 }
+2
View File
@@ -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),
}
}
}
+3 -1
View File
@@ -73,7 +73,7 @@ where
})
}
fn make_error<Error>(&self, error: Error) -> NodeError
pub(crate) fn make_error<Error>(&self, error: Error) -> NodeError
where
Error: Into<Box<dyn std::error::Error>>,
{
@@ -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,
}
}
+50 -27
View File
@@ -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<Self::NodeData, Self::TraceData>) -> 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<HandlerResponse>,
}
struct FunctionContext<'a> {
context: &'a NodeContext<FunctionContent, FunctionV2Trace>,
function: &'a Function,
start: Instant,
}
trait FunctionErrorExt<T> {
async fn function_context(self, ctx: &FunctionContext) -> Result<T, NodeError>;
}
impl<T> FunctionErrorExt<T> for Result<T, FunctionError> {
async fn function_context(self, c: &FunctionContext<'_>) -> Result<T, NodeError> {
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))
}
}
}
}
@@ -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<Self> {
rquickjs_serde::from_value(value).or_throw(&ctx)
}
}
#[derive(Debug)]
struct CachedProvider<Provider>(Arc<Provider>)
where
Provider: reqsign::ProvideCredential + Debug;
#[async_trait]
impl<Provider> reqsign::ProvideCredential for CachedProvider<Provider>
where
Provider: reqsign::ProvideCredential + Debug,
{
type Credential = Provider::Credential;
async fn provide_credential(
&self,
ctx: &reqsign::Context,
) -> reqsign::Result<Option<Self::Credential>> {
self.0.provide_credential(ctx).await
}
}
impl<Provider> Clone for CachedProvider<Provider>
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<str>,
}
#[derive(Clone)]
struct AwsRegion(pub Arc<str>);
impl<'de> Deserialize<'de> for AwsRegion {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let region_str = Option::<Arc<str>>::deserialize(deserializer)?;
match region_str {
Some(region) => Ok(AwsRegion(region)),
None => {
static AWS_REGION_ENV: OnceLock<Option<Arc<str>>> = 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<Body>) -> anyhow::Result<Request> {
static CACHED_PROVIDER: OnceLock<CachedProvider<aws::DefaultCredentialProvider>> =
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<str>,
}
impl GcpIamAuth {
pub async fn build_request(&self, http_request: HttpRequest<Body>) -> anyhow::Result<Request> {
static CACHED_PROVIDER: OnceLock<CachedProvider<google::DefaultCredentialProvider>> =
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<Body>) -> anyhow::Result<Request> {
static CACHED_PROVIDER: OnceLock<CachedProvider<azure::DefaultCredentialProvider>> =
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")
}
}
@@ -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<HttpResponse<'js>> {
static HTTP_CLIENT: OnceLock<reqwest::Client> = 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<JsValue>,
}
impl<'js> FromJs<'js> for HttpConfig {
fn from_js(ctx: &Ctx<'js>, value: Value<'js>) -> rquickjs::Result<Self> {
let object = value.into_object().or_throw(ctx)?;
let headers_obj: Option<Object<'js>> = 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<'js>> = 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<Value<'js>> = 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<String, StringPrimitive>,
#[serde(default)]
params: HashMap<String, StringPrimitive>,
data: Option<serde_json::Value>,
auth: Option<HttpConfigAuth>,
}
#[derive(Debug, Clone)]
pub(crate) struct StringPrimitive(pub String);
impl<'de> Deserialize<'de> for StringPrimitive {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
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))
}
}