use crate::com_memory::ComArray;
use crate::error::WslcError;
use crate::session::{WslcSessionHandle, path_to_wide_null};
use core::ffi::c_void;
use serde::{Deserialize, Serialize};
use std::ffi::{CStr, CString};
use std::fmt::Write as _;
use std::os::windows::raw::HANDLE;
use std::path::Path;
use windows_sys::core::HRESULT;
use wslcsdk_sys::types::{
WslcImageInfo, WslcImageProgressMessage, WslcImportImageOptions, WslcLoadImageOptions,
WslcPullImageOptions, WslcPushImageOptions, WslcTagImageOptions,
};
use wslcsdk_sys::*;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ImageInfo {
pub name: String,
pub sha256: String,
pub size_bytes: i64,
pub created_unix_time: u64,
}
#[derive(Debug, Clone)]
pub struct ImageProgress<'a> {
pub id: &'a str,
pub status: WslcImageProgressStatus,
pub current_bytes: u64,
pub total_bytes: u64,
}
impl ImageProgress<'_> {
pub fn to_owned(&self) -> OwnedImageProgress {
OwnedImageProgress {
id: self.id.to_string(),
status: self.status,
current_bytes: self.current_bytes,
total_bytes: self.total_bytes,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct OwnedImageProgress {
pub id: String,
pub status: WslcImageProgressStatus,
pub current_bytes: u64,
pub total_bytes: u64,
}
fn to_hex(bytes: &[u8]) -> String {
let mut out = String::with_capacity(bytes.len() * 2);
for &b in bytes {
let _ = write!(out, "{b:02x}");
}
out
}
unsafe extern "system" fn progress_trampoline<F>(
msg: *const WslcImageProgressMessage,
ctx: *mut c_void,
) -> HRESULT
where
F: FnMut(&ImageProgress<'_>) -> bool,
{
if ctx.is_null() || msg.is_null() {
return windows_sys::Win32::Foundation::S_OK;
}
let res = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| unsafe {
let raw = &*msg;
let id_str = if !raw.id.is_null() {
CStr::from_ptr(raw.id).to_str().unwrap_or_default()
} else {
""
};
let progress = ImageProgress {
id: id_str,
status: raw.status,
current_bytes: raw.detail.current_bytes,
total_bytes: raw.detail.total_bytes,
};
let mutex = &*(ctx as *const std::sync::Mutex<F>);
if let Ok(mut callback) = mutex.lock() {
if callback(&progress) {
windows_sys::Win32::Foundation::S_OK
} else {
windows_sys::Win32::Foundation::E_ABORT
}
} else {
windows_sys::Win32::Foundation::E_ABORT
}
}));
res.unwrap_or(windows_sys::Win32::Foundation::E_ABORT)
}
#[derive(Debug, Clone, Copy, Default)]
pub struct WslcImageManager;
impl WslcImageManager {
pub fn list_images(session: &WslcSessionHandle) -> Result<Vec<ImageInfo>, WslcError> {
let mut raw_images: *mut WslcImageInfo = std::ptr::null_mut();
let mut count: u32 = 0;
let hr = unsafe { WslcListSessionImages(session.as_raw(), &mut raw_images, &mut count) };
if hr < 0 {
return Err(WslcError::from_hresult(hr, "获取镜像列表失败"));
}
let images = unsafe { ComArray::from_raw(raw_images, count as usize) }
.expect("非空长度下的官方输出数组不应构造失败");
let slice = images.as_slice();
if slice.is_empty() {
return Ok(Vec::new());
}
let mut list = Vec::with_capacity(slice.len());
for item in slice {
let name_end = item
.name
.iter()
.position(|&b| b == 0)
.unwrap_or(item.name.len());
let name = String::from_utf8_lossy(&item.name[..name_end]).to_string();
let sha_hex = to_hex(&item.sha256);
list.push(ImageInfo {
name,
sha256: sha_hex,
size_bytes: item.size_bytes,
created_unix_time: item.created_unix_time,
});
}
Ok(list)
}
pub fn pull_image<F>(
session: &WslcSessionHandle,
uri: &str,
registry_auth: Option<&str>,
on_progress: Option<F>,
) -> Result<(), WslcError>
where
F: FnMut(&ImageProgress<'_>) -> bool + Send,
{
let resolved_uri = crate::registry::resolve_image_reference(uri)?;
let c_uri = CString::new(resolved_uri.as_str())
.map_err(|e| WslcError::NulError(format!("URI 包含非法空字节: {e}")))?;
let c_auth = match registry_auth {
Some(a) => Some(
CString::new(a)
.map_err(|e| WslcError::NulError(format!("鉴权信息包含非法空字节: {e}")))?,
),
None => None,
};
let (cb, ctx) = match on_progress {
Some(f) => {
let boxed = Box::into_raw(Box::new(std::sync::Mutex::new(f)));
(Some(progress_trampoline::<F> as _), boxed as *mut c_void)
}
None => (None, std::ptr::null_mut()),
};
let options = WslcPullImageOptions {
uri: c_uri.as_ptr(),
progress_callback: cb,
progress_callback_context: ctx,
registry_auth: c_auth.as_ref().map_or(std::ptr::null(), |s| s.as_ptr()),
};
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe { WslcPullSessionImage(session.as_raw(), &options, &mut err_msg) };
if !ctx.is_null() {
let _ = unsafe { Box::from_raw(ctx as *mut std::sync::Mutex<F>) };
}
unsafe { WslcError::check(hr, err_msg) }?;
log::info!(
"镜像拉取完成,URI: '{}',所属会话: '{}'",
uri,
session.name()
);
Ok(())
}
pub fn import_image_from_file(
session: &WslcSessionHandle,
image_name: &str,
path: impl AsRef<Path>,
) -> Result<(), WslcError> {
let c_name = CString::new(image_name)
.map_err(|e| WslcError::NulError(format!("镜像名称非法: {e}")))?;
let wide_path = path_to_wide_null(path)?;
let options = WslcImportImageOptions::default();
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe {
WslcImportSessionImageFromFile(
session.as_raw(),
c_name.as_ptr(),
wide_path.as_ptr(),
&options,
&mut err_msg,
)
};
unsafe { WslcError::check(hr, err_msg) }?;
log::info!(
"镜像导入完成,名称: '{}',所属会话: '{}'",
image_name,
session.name()
);
Ok(())
}
pub fn load_image_from_file(
session: &WslcSessionHandle,
path: impl AsRef<Path>,
) -> Result<(), WslcError> {
let wide_path = path_to_wide_null(path)?;
let options = WslcLoadImageOptions::default();
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe {
WslcLoadSessionImageFromFile(
session.as_raw(),
wide_path.as_ptr(),
&options,
&mut err_msg,
)
};
unsafe { WslcError::check(hr, err_msg) }?;
log::info!("镜像载入完成,所属会话: '{}'", session.name());
Ok(())
}
pub unsafe fn load_image_from_handle(
session: &WslcSessionHandle,
content_handle: HANDLE,
bytes: u64,
) -> Result<(), WslcError> {
let options = WslcLoadImageOptions::default();
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe {
WslcLoadSessionImage(
session.as_raw(),
content_handle,
bytes,
&options,
&mut err_msg,
)
};
unsafe { WslcError::check(hr, err_msg) }
}
pub unsafe fn import_image_from_handle(
session: &WslcSessionHandle,
image_name: &str,
content_handle: HANDLE,
bytes: u64,
) -> Result<(), WslcError> {
let c_name = CString::new(image_name)
.map_err(|e| WslcError::NulError(format!("镜像名称非法: {e}")))?;
let options = WslcImportImageOptions::default();
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe {
WslcImportSessionImage(
session.as_raw(),
c_name.as_ptr(),
content_handle,
bytes,
&options,
&mut err_msg,
)
};
unsafe { WslcError::check(hr, err_msg) }
}
pub fn tag_image(
session: &WslcSessionHandle,
image: &str,
repo: &str,
tag: &str,
) -> Result<(), WslcError> {
let c_img =
CString::new(image).map_err(|e| WslcError::NulError(format!("镜像名称非法: {e}")))?;
let c_repo =
CString::new(repo).map_err(|e| WslcError::NulError(format!("仓库名称非法: {e}")))?;
let c_tag =
CString::new(tag).map_err(|e| WslcError::NulError(format!("标签名称非法: {e}")))?;
let options = WslcTagImageOptions {
image: c_img.as_ptr(),
repo: c_repo.as_ptr(),
tag: c_tag.as_ptr(),
};
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe { WslcTagSessionImage(session.as_raw(), &options, &mut err_msg) };
unsafe { WslcError::check(hr, err_msg) }
}
pub fn push_image_with_progress<F>(
session: &WslcSessionHandle,
image: &str,
registry_auth: Option<&str>,
on_progress: Option<F>,
) -> Result<(), WslcError>
where
F: FnMut(&ImageProgress<'_>) -> bool + Send,
{
let c_img =
CString::new(image).map_err(|e| WslcError::NulError(format!("镜像名称非法: {e}")))?;
let c_auth = match registry_auth {
Some(a) => Some(
CString::new(a).map_err(|e| WslcError::NulError(format!("鉴权信息非法: {e}")))?,
),
None => None,
};
let (cb, ctx) = match on_progress {
Some(f) => {
let boxed = Box::into_raw(Box::new(std::sync::Mutex::new(f)));
(Some(progress_trampoline::<F> as _), boxed as *mut c_void)
}
None => (None, std::ptr::null_mut()),
};
let options = WslcPushImageOptions {
image: c_img.as_ptr(),
registry_auth: c_auth.as_ref().map_or(std::ptr::null(), |s| s.as_ptr()),
progress_callback: cb,
progress_callback_context: ctx,
};
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe { WslcPushSessionImage(session.as_raw(), &options, &mut err_msg) };
if !ctx.is_null() {
let _ = unsafe { Box::from_raw(ctx as *mut std::sync::Mutex<F>) };
}
unsafe { WslcError::check(hr, err_msg) }?;
log::info!(
"镜像推送完成,镜像: '{}',所属会话: '{}'",
image,
session.name()
);
Ok(())
}
pub fn push_image(
session: &WslcSessionHandle,
image: &str,
registry_auth: Option<&str>,
) -> Result<(), WslcError> {
Self::push_image_with_progress(
session,
image,
registry_auth,
None::<fn(&ImageProgress<'_>) -> bool>,
)
}
pub fn delete_image(session: &WslcSessionHandle, name_or_id: &str) -> Result<(), WslcError> {
let c_str = CString::new(name_or_id)
.map_err(|e| WslcError::NulError(format!("镜像名称或 ID 非法: {e}")))?;
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe { WslcDeleteSessionImage(session.as_raw(), c_str.as_ptr(), &mut err_msg) };
unsafe { WslcError::check(hr, err_msg) }?;
log::info!(
"镜像删除成功,目标: '{}',所属会话: '{}'",
name_or_id,
session.name()
);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_image_progress_to_owned_roundtrip() {
let borrowed = ImageProgress {
id: "sha256:1234",
status: WslcImageProgressStatus::Downloading,
current_bytes: 1024,
total_bytes: 2048,
};
let owned = borrowed.to_owned();
assert_eq!(owned.id, "sha256:1234");
assert_eq!(owned.status, WslcImageProgressStatus::Downloading);
assert_eq!(owned.current_bytes, 1024);
assert_eq!(owned.total_bytes, 2048);
}
#[test]
fn test_image_progress_to_owned_is_independent() {
let mut owned = {
let borrowed = ImageProgress {
id: "layer-a",
status: WslcImageProgressStatus::Pulling,
current_bytes: 0,
total_bytes: 0,
};
borrowed.to_owned()
};
owned.id.push_str("-mutated");
assert_eq!(owned.id, "layer-a-mutated");
}
#[test]
fn test_image_info_serde_roundtrip() {
let info = ImageInfo {
name: "alpine:latest".to_string(),
sha256: "a".repeat(64),
size_bytes: 3_500_000,
created_unix_time: 1_700_000_000,
};
let json = serde_json::to_string(&info).expect("序列化失败");
let back: ImageInfo = serde_json::from_str(&json).expect("反序列化失败");
assert_eq!(info, back);
}
#[test]
fn test_to_hex_produces_lowercase_fixed_width() {
assert_eq!(to_hex(&[]), "");
assert_eq!(to_hex(&[0x00]), "00");
assert_eq!(to_hex(&[0xff]), "ff");
assert_eq!(to_hex(&[0x0f, 0xf0]), "0ff0");
let all: Vec<u8> = (0..=255u8).collect();
let hex = to_hex(&all);
assert_eq!(hex.len(), 512);
assert!(
hex.chars()
.all(|c| c.is_ascii_hexdigit() && !c.is_uppercase())
);
assert!(hex.starts_with("000102"));
assert!(hex.ends_with("fdfeff"));
}
}