use crate::callback::invoke_callback;
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 msg.is_null() {
return windows_sys::Win32::Foundation::S_OK;
}
let should_continue = unsafe {
invoke_callback::<F, _>(ctx, |callback| {
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,
};
callback(&progress)
})
};
match should_continue {
Some(true) => windows_sys::Win32::Foundation::S_OK,
_ => 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, "获取镜像列表失败"));
}
if count > 0 && raw_images.is_null() {
return Err(WslcError::UnexpectedSdkResult(format!(
"WslcListSessionImages 返回成功状态并声称有 {count} 个镜像,却未给出数组指针"
)));
}
let images =
unsafe { ComArray::from_raw(raw_images, count as usize) }.ok_or_else(|| {
WslcError::UnexpectedSdkResult("WslcListSessionImages 未返回镜像数组".to_string())
})?;
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,
format!("拉取镜像失败,实际拉取地址: '{resolved_uri}'"),
)?;
}
log::info!(
"镜像拉取完成,请求 URI: '{}',实际拉取地址: '{}',所属会话: '{}'",
uri,
resolved_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, format!("导入镜像失败,名称: '{image_name}'"))?;
}
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,
format!("载入镜像失败,所属会话: '{}'", session.name()),
)?;
}
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,
format!("从句柄载入镜像失败,所属会话: '{}'", session.name()),
)
}
}
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,
format!("从句柄导入镜像失败,名称: '{image_name}'"),
)
}
}
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, format!("镜像打标签失败,目标: '{repo}:{tag}'"))
}
}
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, format!("推送镜像失败,目标: '{image}'"))?;
}
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,
format!("删除镜像失败,名称或 ID: '{name_or_id}'"),
)?;
}
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"));
}
}