use std::sync::Arc;
use unb_runtime::ClientSession;
use serde_json::Value;
use wasm_bindgen::prelude::*;
use crate::endpoint::{dial_endpoints, Endpoint, EndpointSet, TransportKind};
#[wasm_bindgen]
pub struct Session {
session: Arc<ClientSession>,
}
#[wasm_bindgen]
impl Session {
pub async fn connect(url: String) -> Result<Session, JsError> {
let mut set = EndpointSet::new();
set.push(Endpoint {
kind: TransportKind::WebSocket,
address: url,
cert_hash: None,
});
let session = dial_endpoints(&set)
.await
.map_err(|error| JsError::new(&error.to_string()))?;
Ok(Session { session })
}
#[wasm_bindgen(js_name = connectEndpoints)]
pub async fn connect_endpoints(endpoints: JsValue) -> Result<Session, JsError> {
let value = to_value(endpoints)?;
let items = value
.as_array()
.ok_or_else(|| JsError::new("endpoints must be an array"))?;
let mut set = EndpointSet::new();
for item in items {
let kind = match item["kind"].as_str() {
Some("webtransport") => TransportKind::WebTransport,
Some("ws") | Some("websocket") => TransportKind::WebSocket,
other => return Err(JsError::new(&format!("unknown transport kind: {other:?}"))),
};
let address = item["address"]
.as_str()
.ok_or_else(|| JsError::new("endpoint address must be a string"))?
.to_string();
let cert_hash = match item["certHash"].as_str() {
Some(hex) => Some(parse_hex32(hex)?),
None => None,
};
set.push(Endpoint {
kind,
address,
cert_hash,
});
}
let session = dial_endpoints(&set)
.await
.map_err(|error| JsError::new(&error.to_string()))?;
Ok(Session { session })
}
pub async fn fetch(&self, request: JsValue) -> Result<String, JsError> {
let request = to_value(request)?;
if let Some(method) = request.get("method").and_then(Value::as_str) {
if method != "POST" {
return Err(JsError::new("application requests use POST"));
}
}
if !matches!(
request.get("kind").and_then(Value::as_str),
None | Some("request")
) {
return Err(JsError::new("fetch only supports request operations"));
}
let target = request["target"]
.as_str()
.ok_or_else(|| JsError::new("request target must be a string"))?;
let mut builder = http::Request::builder()
.method(http::Method::POST)
.uri(format!("/{}", target.trim_start_matches('/')));
if let Some(header_value) = request.get("headers") {
let values = header_value
.as_object()
.ok_or_else(|| JsError::new("request headers must be an object"))?;
for (name, value) in values {
if name.to_ascii_lowercase().starts_with("unb-") {
return Err(JsError::new(&format!(
"{name}: unb-* headers are reserved for framing metadata"
)));
}
let value = value.as_str().ok_or_else(|| {
JsError::new(&format!("{name}: header values must be strings"))
})?;
builder = builder.header(name.as_str(), value);
}
}
let body = request.get("body").cloned().unwrap_or(Value::Null);
let request = builder
.body(if body.is_null() {
bytes::Bytes::new()
} else {
bytes::Bytes::from(body.to_string())
})
.map_err(|error| JsError::new(&error.to_string()))?;
let response = self
.session
.fetch(request, FETCH_TIMEOUT)
.await
.map_err(|error| JsError::new(&error.to_string()))?;
let mut headers = serde_json::Map::new();
for (name, value) in response.headers() {
if name.as_str().starts_with("unb-") {
continue;
}
let value = value
.to_str()
.map_err(|error| JsError::new(&error.to_string()))?;
headers.insert(name.as_str().into(), Value::String(value.into()));
}
let body = if response.body().is_empty() {
Value::Null
} else {
serde_json::from_slice(response.body())
.map_err(|error| JsError::new(&error.to_string()))?
};
Ok(serde_json::json!({
"status": response.status().as_u16(),
"headers": headers,
"body": body
})
.to_string())
}
}
const FETCH_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
fn parse_hex32(hex: &str) -> Result<[u8; 32], JsError> {
let cleaned = hex.trim().replace(':', "");
if !cleaned.is_ascii() || cleaned.len() != 64 {
return Err(JsError::new("certHash must be 32 bytes of hex"));
}
let mut out = [0u8; 32];
for (index, byte) in out.iter_mut().enumerate() {
let pair = &cleaned[index * 2..index * 2 + 2];
*byte = u8::from_str_radix(pair, 16)
.map_err(|_| JsError::new("certHash must be 32 bytes of hex"))?;
}
Ok(out)
}
fn to_value(payload: JsValue) -> Result<Value, JsError> {
if payload.is_undefined() || payload.is_null() {
return Ok(Value::Null);
}
let text = js_sys::JSON::stringify(&payload)
.map_err(|_| JsError::new("payload is not JSON-serializable"))?
.as_string()
.unwrap_or_else(|| "null".into());
serde_json::from_str(&text).map_err(|error| JsError::new(&error.to_string()))
}