168 lines
5.2 KiB
Rust

use std::ffi::{CStr, CString};
use std::os::raw::c_char;
use std::path::Path;
use tokio::net::UnixStream;
use tonic::transport::{Endpoint, Uri};
use tower::service_fn;
// Generated by tonic-build
pub mod workload {
tonic::include_proto!("_");
}
use workload::spiffe_workload_api_client::SpiffeWorkloadApiClient;
use workload::X509svidRequest;
#[repr(C)]
pub struct SvidResponseC {
spiffe_id: *mut c_char,
x509_svid: *mut u8,
x509_svid_len: usize,
x509_svid_key: *mut u8,
x509_svid_key_len: usize,
bundle: *mut u8,
bundle_len: usize,
error: *mut c_char,
}
impl SvidResponseC {
fn with_error(err_msg: &str) -> Self {
let err_c = match CString::new(err_msg) {
Ok(c) => c.into_raw(),
Err(_) => std::ptr::null_mut(),
};
SvidResponseC {
spiffe_id: std::ptr::null_mut(),
x509_svid: std::ptr::null_mut(),
x509_svid_len: 0,
x509_svid_key: std::ptr::null_mut(),
x509_svid_key_len: 0,
bundle: std::ptr::null_mut(),
bundle_len: 0,
error: err_c,
}
}
}
async fn fetch_svid_async(socket_path: &str) -> Result<SvidResponseC, Box<dyn std::error::Error>> {
let path = Path::new(socket_path).to_path_buf();
let channel = Endpoint::try_from("http://[::]:50051")?
.connect_with_connector(service_fn(move |_: Uri| {
let path_clone = path.clone();
async move { UnixStream::connect(path_clone).await }
}))
.await?;
let mut client = SpiffeWorkloadApiClient::new(channel);
let request = tonic::Request::new(X509svidRequest {});
let mut stream = client.fetch_x509svid(request).await?.into_inner();
if let Some(response) = stream.message().await? {
if let Some(svid) = response.svids.first() {
let spiffe_id_c = CString::new(svid.spiffe_id.clone())?.into_raw();
let mut svid_box = svid.x509_svid.clone().into_boxed_slice();
let x509_svid = svid_box.as_mut_ptr();
let x509_svid_len = svid_box.len();
std::mem::forget(svid_box);
let mut key_box = svid.x509_svid_key.clone().into_boxed_slice();
let x509_svid_key = key_box.as_mut_ptr();
let x509_svid_key_len = key_box.len();
std::mem::forget(key_box);
let mut bundle_box = svid.bundle.clone().into_boxed_slice();
let bundle = bundle_box.as_mut_ptr();
let bundle_len = bundle_box.len();
std::mem::forget(bundle_box);
return Ok(SvidResponseC {
spiffe_id: spiffe_id_c,
x509_svid,
x509_svid_len,
x509_svid_key,
x509_svid_key_len,
bundle,
bundle_len,
error: std::ptr::null_mut(),
});
}
}
Err("No SVID received from Workload API".into())
}
#[no_mangle]
pub extern "C" fn fetch_svid(socket_path_ptr: *const c_char) -> *mut SvidResponseC {
if socket_path_ptr.is_null() {
let resp = Box::new(SvidResponseC::with_error("socket_path_ptr is null"));
return Box::into_raw(resp);
}
let socket_path = unsafe {
match CStr::from_ptr(socket_path_ptr).to_str() {
Ok(s) => s.to_string(),
Err(e) => {
let resp = Box::new(SvidResponseC::with_error(&format!("Invalid UTF-8 in socket path: {}", e)));
return Box::into_raw(resp);
}
}
};
let rt = match tokio::runtime::Builder::new_current_thread().enable_all().build() {
Ok(rt) => rt,
Err(e) => {
let resp = Box::new(SvidResponseC::with_error(&format!("Failed to build tokio runtime: {}", e)));
return Box::into_raw(resp);
}
};
let result = match rt.block_on(fetch_svid_async(&socket_path)) {
Ok(data) => data,
Err(e) => SvidResponseC::with_error(&e.to_string()),
};
Box::into_raw(Box::new(result))
}
#[no_mangle]
pub extern "C" fn free_svid(ptr: *mut SvidResponseC) {
if ptr.is_null() {
return;
}
unsafe {
let mut resp = Box::from_raw(ptr);
if !resp.spiffe_id.is_null() {
let _ = CString::from_raw(resp.spiffe_id);
resp.spiffe_id = std::ptr::null_mut();
}
if !resp.error.is_null() {
let _ = CString::from_raw(resp.error);
resp.error = std::ptr::null_mut();
}
if !resp.x509_svid.is_null() && resp.x509_svid_len > 0 {
let _ = Box::from_raw(std::ptr::slice_from_raw_parts_mut(resp.x509_svid, resp.x509_svid_len));
resp.x509_svid = std::ptr::null_mut();
resp.x509_svid_len = 0;
}
if !resp.x509_svid_key.is_null() && resp.x509_svid_key_len > 0 {
let _ = Box::from_raw(std::ptr::slice_from_raw_parts_mut(resp.x509_svid_key, resp.x509_svid_key_len));
resp.x509_svid_key = std::ptr::null_mut();
resp.x509_svid_key_len = 0;
}
if !resp.bundle.is_null() && resp.bundle_len > 0 {
let _ = Box::from_raw(std::ptr::slice_from_raw_parts_mut(resp.bundle, resp.bundle_len));
resp.bundle = std::ptr::null_mut();
resp.bundle_len = 0;
}
}
}