use crate::imp::core::abi::ErrorCode;
use crate::imp::host::error::HostError;
use crate::imp::host::runtime::RuntimeState;
use std::net::{IpAddr, ToSocketAddrs};
use wasmtime::{Caller, Linker};
fn is_blocked_ip(ip: &IpAddr) -> bool {
match ip {
IpAddr::V4(a) => {
let o = a.octets();
a.is_loopback()
|| a.is_private()
|| a.is_link_local()
|| a.is_unspecified()
|| a.is_broadcast()
|| a.is_documentation()
|| (o[0] == 100 && (o[1] & 0xC0) == 0x40)
}
IpAddr::V6(a) => {
a.is_loopback()
|| a.is_unspecified()
|| (a.segments()[0] & 0xfe00) == 0xfc00 || (a.segments()[0] & 0xffc0) == 0xfe80 || a
.to_ipv4_mapped()
.map(|v4| is_blocked_ip(&IpAddr::V4(v4)))
.unwrap_or(false)
}
}
}
fn jwks_url_is_allowed(url: &str) -> bool {
let parsed = match reqwest::Url::parse(url) {
Ok(u) => u,
Err(_) => return false,
};
let insecure_ok = std::env::var_os("DIGSTORE_ALLOW_INSECURE_JWKS").is_some();
match parsed.scheme() {
"https" => {}
"http" if insecure_ok => return parsed.host_str().is_some(),
_ => return false,
}
let host = match parsed.host_str() {
Some(h) => h,
None => return false,
};
if insecure_ok {
return true;
}
let port = parsed.port_or_known_default().unwrap_or(443);
match (host, port).to_socket_addrs() {
Ok(addrs) => {
let addrs: Vec<_> = addrs.collect();
!addrs.is_empty() && addrs.iter().all(|sa| !is_blocked_ip(&sa.ip()))
}
Err(_) => false,
}
}
pub fn register(linker: &mut Linker<RuntimeState>) -> Result<(), HostError> {
let m = "dig_host";
linker
.func_wrap(
m,
"host_get_current_time",
|caller: Caller<'_, RuntimeState>| -> i64 {
caller.data().host.clock.now_unix_secs() as i64
},
)
.map_err(|e| HostError::Wasmtime(e.to_string()))?;
linker
.func_wrap(
m,
"host_random_bytes",
|mut caller: Caller<'_, RuntimeState>, count: i32| -> i32 {
if count < 0 {
return ErrorCode::InvalidParameter as i32;
}
let max = caller.data().host.config.max_random_bytes as usize;
let state = &mut caller.data_mut().host;
match state.rng.fill(count as usize, max) {
Some(bytes) => match state.return_buffer.set(&bytes) {
Ok(n) => n as i32,
Err(_) => ErrorCode::GeneralError as i32,
},
None => ErrorCode::InvalidParameter as i32,
}
},
)
.map_err(|e| HostError::Wasmtime(e.to_string()))?;
linker
.func_wrap(
m,
"host_get_public_key",
|mut caller: Caller<'_, RuntimeState>| -> i32 {
let pk = caller.data().host.keys.bls_public.0; match caller.data_mut().host.return_buffer.set(&pk) {
Ok(n) => n as i32,
Err(_) => ErrorCode::GeneralError as i32,
}
},
)
.map_err(|e| HostError::Wasmtime(e.to_string()))?;
const TAG_LEN: usize = crate::imp::core::ATTEST_DST.len();
const CHALLENGE_LEN: usize = TAG_LEN + 32 + 32 + 8;
const SESSION_TTL_SECS: u64 = 300;
linker
.func_wrap(
m,
"host_create_attestation",
|mut caller: Caller<'_, RuntimeState>, challenge_ptr: i32| -> i32 {
let mem = match caller.get_export("memory").and_then(|e| e.into_memory()) {
Some(mem) => mem,
None => return ErrorCode::GeneralError as i32,
};
let data = mem.data(&caller);
let start = challenge_ptr as usize;
let end = match start.checked_add(CHALLENGE_LEN) {
Some(e) if e <= data.len() => e,
_ => return ErrorCode::InvalidParameter as i32,
};
let challenge = data[start..end].to_vec();
let state = &mut caller.data_mut().host;
let sig = match state.attestation.attest(&challenge) {
Ok(s) => s,
Err(_) => return ErrorCode::AttestationFailed as i32,
};
let pk = state.attestation.public_key();
let mut resp = Vec::with_capacity(48 + 32 + 96);
resp.extend_from_slice(&pk.0);
resp.extend_from_slice(&state.instance_id.0);
resp.extend_from_slice(&sig.0);
state.last_signature = Some(sig);
match state.return_buffer.set(&resp) {
Ok(n) => n as i32,
Err(_) => ErrorCode::GeneralError as i32,
}
},
)
.map_err(|e| HostError::Wasmtime(e.to_string()))?;
linker
.func_wrap(
m,
"host_establish_session",
|mut caller: Caller<'_, RuntimeState>, challenge_ptr: i32| -> i32 {
let mem = match caller.get_export("memory").and_then(|e| e.into_memory()) {
Some(mem) => mem,
None => return ErrorCode::GeneralError as i32,
};
let data = mem.data(&caller);
let start = challenge_ptr as usize;
let end = match start.checked_add(CHALLENGE_LEN) {
Some(e) if e <= data.len() => e,
_ => return ErrorCode::InvalidParameter as i32,
};
let challenge = &data[start..end];
let mut nonce = [0u8; 32];
let mut store_id = [0u8; 32];
nonce.copy_from_slice(&challenge[TAG_LEN..TAG_LEN + 32]);
store_id.copy_from_slice(&challenge[TAG_LEN + 32..TAG_LEN + 64]);
let now = caller.data().host.clock.now_unix_secs();
caller
.data_mut()
.host
.sessions
.establish(nonce, store_id, now, SESSION_TTL_SECS);
0
},
)
.map_err(|e| HostError::Wasmtime(e.to_string()))?;
linker
.func_wrap(
m,
"host_verify_session",
|caller: Caller<'_, RuntimeState>| -> i32 {
let now = caller.data().host.clock.now_unix_secs();
let sessions = &caller.data().host.sessions;
if sessions.is_valid(now) {
1
} else if sessions.is_expired_at(now) {
ErrorCode::SessionExpired as i32
} else {
0
}
},
)
.map_err(|e| HostError::Wasmtime(e.to_string()))?;
linker
.func_wrap(
m,
"jwks_fetch",
|mut caller: Caller<'_, RuntimeState>, url_ptr: i32, url_len: i32| -> i32 {
let now = caller.data().host.clock.now_unix_secs();
let sessions = &caller.data().host.sessions;
if !sessions.is_valid(now) {
return if sessions.is_expired_at(now) {
ErrorCode::SessionExpired as i32
} else {
ErrorCode::NoSession as i32
};
}
let timeout_secs = caller.data().host.http_timeout_secs;
let mem = match caller.get_export("memory").and_then(|e| e.into_memory()) {
Some(mem) => mem,
None => return ErrorCode::GeneralError as i32,
};
let data = mem.data(&caller);
let start = url_ptr as usize;
let end = match start.checked_add(url_len as usize) {
Some(e) if e <= data.len() => e,
_ => return ErrorCode::InvalidParameter as i32,
};
let url = match std::str::from_utf8(&data[start..end]) {
Ok(u) => u.to_string(),
Err(_) => return ErrorCode::InvalidParameter as i32,
};
if !jwks_url_is_allowed(&url) {
return ErrorCode::InvalidParameter as i32;
}
let resp = match reqwest::blocking::Client::new()
.get(&url)
.timeout(std::time::Duration::from_secs(timeout_secs))
.send()
{
Ok(r) => r,
Err(e) if e.is_timeout() => return ErrorCode::Timeout as i32,
Err(_) => return ErrorCode::NetworkError as i32,
};
let body = match resp.bytes() {
Ok(b) => b,
Err(_) => return ErrorCode::NetworkError as i32,
};
match caller.data_mut().host.return_buffer.set(&body) {
Ok(n) => n as i32,
Err(_) => ErrorCode::GeneralError as i32,
}
},
)
.map_err(|e| HostError::Wasmtime(e.to_string()))?;
linker
.func_wrap(
m,
"host_read_return_buffer",
|mut caller: Caller<'_, RuntimeState>, dest_ptr: i32| -> i32 {
let mem = match caller.get_export("memory").and_then(|e| e.into_memory()) {
Some(mem) => mem,
None => return ErrorCode::GeneralError as i32,
};
let buf = caller.data().host.return_buffer.as_slice().to_vec();
let data = mem.data_mut(&mut caller);
let start = dest_ptr as usize;
let end = match start.checked_add(buf.len()) {
Some(e) => e,
None => return ErrorCode::InvalidParameter as i32,
};
if end > data.len() {
return ErrorCode::BufferTooSmall as i32;
}
data[start..end].copy_from_slice(&buf);
buf.len() as i32
},
)
.map_err(|e| HostError::Wasmtime(e.to_string()))?;
Ok(())
}