use crate::com_memory::ComAnsiString;
use crate::error::WslcError;
use crate::session::WslcSessionHandle;
use std::ffi::CString;
use wslcsdk_sys::types::WslcIdentityTokenType;
use wslcsdk_sys::*;
#[derive(Clone, PartialEq, Eq)]
pub struct AuthTokenResult {
pub identity_token: String,
pub token_type: WslcIdentityTokenType,
}
impl std::fmt::Debug for AuthTokenResult {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AuthTokenResult")
.field("identity_token", &"[已脱敏]")
.field("token_type", &self.token_type)
.finish()
}
}
impl Drop for AuthTokenResult {
fn drop(&mut self) {
unsafe {
for byte in self.identity_token.as_mut_vec() {
std::ptr::write_volatile(byte, 0);
}
}
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct WslcRegistryManager;
impl WslcRegistryManager {
pub fn authenticate(
session: &WslcSessionHandle,
server_address: &str,
username: &str,
password: &str,
) -> Result<AuthTokenResult, WslcError> {
let c_server = CString::new(server_address)
.map_err(|e| WslcError::NulError(format!("服务器地址非法: {e}")))?;
let c_user =
CString::new(username).map_err(|e| WslcError::NulError(format!("用户名非法: {e}")))?;
let c_pass =
CString::new(password).map_err(|e| WslcError::NulError(format!("密码非法: {e}")))?;
let mut token_ptr: *mut i8 = std::ptr::null_mut();
let mut token_type = WslcIdentityTokenType::Unknown;
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe {
WslcSessionAuthenticate(
session.as_raw(),
c_server.as_ptr(),
c_user.as_ptr(),
c_pass.as_ptr(),
&mut token_ptr,
&mut token_type,
&mut err_msg,
)
};
let mut pass_bytes = c_pass.into_bytes_with_nul();
for b in &mut pass_bytes {
unsafe {
std::ptr::write_volatile(b, 0);
}
}
unsafe {
WslcError::check(
hr,
err_msg,
format!("镜像仓库身份认证失败,服务器: '{server_address}'"),
)?;
}
if token_ptr.is_null() {
return Err(WslcError::UnexpectedSdkResult(
"WslcSessionAuthenticate 返回成功状态却未给出令牌指针".to_string(),
));
}
let token = unsafe { ComAnsiString::from_raw(token_ptr) }.ok_or_else(|| {
WslcError::UnexpectedSdkResult(
"WslcSessionAuthenticate 返回成功状态却未给出令牌字符串".to_string(),
)
})?;
Ok(AuthTokenResult {
identity_token: token
.as_str()
.map_err(|e| {
WslcError::InvalidConfiguration(format!("仓库返回的认证令牌非合法 UTF-8: {e}"))
})?
.to_string(),
token_type,
})
}
}
const DEFAULT_REGISTRY: &str = "docker.io";
const DEFAULT_MIRROR_ENV: &str = "WSLC_REGISTRY_MIRROR";
const MIRROR_ENV_PREFIX: &str = "WSLC_REGISTRY_MIRROR_";
pub fn resolve_image_reference(image: &str) -> Result<String, WslcError> {
resolve_image_reference_with(image, |key| std::env::var(key).ok())
}
pub(crate) fn resolve_image_reference_with<F>(image: &str, mut env: F) -> Result<String, WslcError>
where
F: FnMut(&str) -> Option<String>,
{
let parsed = parse_image_reference(image);
let registry = parsed.registry.unwrap_or(DEFAULT_REGISTRY);
let Some(mirror) = mirror_for_registry(registry, &mut env)? else {
return Ok(image.to_string());
};
Ok(format!(
"{}/{}",
mirror.trim_end_matches('/'),
parsed.repository_with_tag
))
}
fn mirror_for_registry<F>(registry: &str, env: &mut F) -> Result<Option<String>, WslcError>
where
F: FnMut(&str) -> Option<String>,
{
let registry_key: String = registry
.chars()
.map(|ch| {
if ch.is_ascii_alphanumeric() {
ch.to_ascii_uppercase()
} else {
'_'
}
})
.collect();
let exact_key = format!("{MIRROR_ENV_PREFIX}{registry_key}");
let mirror = env(&exact_key).or_else(|| {
if registry == DEFAULT_REGISTRY {
env(DEFAULT_MIRROR_ENV)
} else {
None
}
});
match mirror.map(|v| v.trim().to_string()) {
Some(v) if v.is_empty() => Err(WslcError::InvalidConfiguration(
"镜像加速器环境变量未配置具体地址".to_string(),
)),
Some(v) => {
let normalized = strip_mirror_scheme(&v)?;
validate_mirror_address(&normalized)?;
Ok(Some(normalized))
}
None => Ok(None),
}
}
fn strip_mirror_scheme(value: &str) -> Result<String, WslcError> {
let Some((scheme, rest)) = value.split_once("://") else {
return Ok(value.to_string());
};
if !matches!(scheme, "http" | "https") {
return Err(WslcError::InvalidConfiguration(format!(
"镜像加速器仅支持 http 或 https 方案,实际为: {scheme}://"
)));
}
if rest.contains("://") {
return Err(WslcError::InvalidConfiguration(format!(
"镜像加速器地址含多个 scheme 分隔符,疑似畸形配置: {value}"
)));
}
Ok(rest.to_string())
}
fn validate_mirror_address(value: &str) -> Result<(), WslcError> {
const ALLOWED_EXTRA: &[char] = &['.', '-', ':', '/'];
if value.is_empty() {
return Err(WslcError::InvalidConfiguration(
"镜像加速器地址剥离方案前缀后为空".to_string(),
));
}
if let Some(bad) = value
.chars()
.find(|c| !c.is_ascii_alphanumeric() && !ALLOWED_EXTRA.contains(c))
{
return Err(WslcError::InvalidConfiguration(format!(
"镜像加速器地址含非法字符 {bad:?}:{value};\
仅允许字母数字与 `.` `-` `:` `/`(分别用于主机名、端口与路径前缀)"
)));
}
Ok(())
}
struct ParsedImage<'a> {
registry: Option<&'a str>,
repository_with_tag: &'a str,
}
fn parse_image_reference(image: &str) -> ParsedImage<'_> {
let (first, rest) = image.split_once('/').unwrap_or((image, ""));
let has_registry = first == "localhost" || first.contains('.') || first.contains(':');
if has_registry && !rest.is_empty() {
ParsedImage {
registry: Some(first),
repository_with_tag: rest,
}
} else {
ParsedImage {
registry: None,
repository_with_tag: image,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_mirror_applies_to_default_registry() {
let res = resolve_image_reference_with("ubuntu:latest", |k| {
if k == DEFAULT_MIRROR_ENV {
Some("mirror.example.com".to_string())
} else {
None
}
});
assert_eq!(res.expect("解析失败"), "mirror.example.com/ubuntu:latest");
}
#[test]
fn test_mirror_applies_to_named_registry() {
let res = resolve_image_reference_with("ghcr.io/org/repo:1.0", |k| {
if k == "WSLC_REGISTRY_MIRROR_GHCR_IO" {
Some("ghcr-mirror.example.com".to_string())
} else {
None
}
});
assert_eq!(
res.expect("解析失败"),
"ghcr-mirror.example.com/org/repo:1.0"
);
}
#[test]
fn test_default_mirror_env_is_not_applied_to_other_registries() {
let res = resolve_image_reference_with("ghcr.io/org/repo:1.0", |k| {
if k == DEFAULT_MIRROR_ENV {
Some("mirror.example.com".to_string())
} else {
None
}
});
assert_eq!(res.expect("解析失败"), "ghcr.io/org/repo:1.0");
}
#[test]
fn test_mirror_with_scheme_and_trailing_slash_is_normalized() {
let res = resolve_image_reference_with("ubuntu:latest", |k| {
if k == DEFAULT_MIRROR_ENV {
Some("https://mirror.example.com/".to_string())
} else {
None
}
});
assert_eq!(res.expect("解析失败"), "mirror.example.com/ubuntu:latest");
}
#[test]
fn test_blank_mirror_is_rejected() {
let res = resolve_image_reference_with("ubuntu:latest", |k| {
if k == DEFAULT_MIRROR_ENV {
Some(" ".to_string())
} else {
None
}
});
match res {
Err(WslcError::InvalidConfiguration(_)) => {}
other => panic!("预期返回 InvalidConfiguration,实际为: {other:?}"),
}
}
#[test]
fn test_parse_image_reference_registry_detection() {
let plain = parse_image_reference("ubuntu:latest");
assert_eq!(plain.registry, None);
assert_eq!(plain.repository_with_tag, "ubuntu:latest");
let domain = parse_image_reference("ghcr.io/org/repo:1.0");
assert_eq!(domain.registry, Some("ghcr.io"));
assert_eq!(domain.repository_with_tag, "org/repo:1.0");
let with_port = parse_image_reference("127.0.0.1:5000/team/app:v1");
assert_eq!(with_port.registry, Some("127.0.0.1:5000"));
assert_eq!(with_port.repository_with_tag, "team/app:v1");
let local = parse_image_reference("localhost/team/app");
assert_eq!(local.registry, Some("localhost"));
assert_eq!(local.repository_with_tag, "team/app");
let tag_only = parse_image_reference("ubuntu:22.04");
assert_eq!(tag_only.registry, None);
assert_eq!(tag_only.repository_with_tag, "ubuntu:22.04");
let bare = parse_image_reference("alpine");
assert_eq!(bare.registry, None);
assert_eq!(bare.repository_with_tag, "alpine");
}
#[test]
fn test_no_mirror_configured_returns_input_verbatim() {
let res = resolve_image_reference_with("ubuntu:latest", |_| None);
assert_eq!(res.expect("解析失败"), "ubuntu:latest");
}
fn mirror_env(value: &str) -> impl FnMut(&str) -> Option<String> + '_ {
move |_| Some(value.to_string())
}
#[test]
fn test_mirror_with_injection_characters_is_rejected() {
assert!(
resolve_image_reference_with("alpine:latest", mirror_env("user@evil.com")).is_err(),
"含 `@` 的加速器地址必须被拒绝,否则可注入用户信息段"
);
assert!(
resolve_image_reference_with("alpine:latest", mirror_env("evil.com?a=b")).is_err(),
"含 `?` 的加速器地址必须被拒绝"
);
assert!(
resolve_image_reference_with("alpine:latest", mirror_env("evil.com#f")).is_err(),
"含 `#` 的加速器地址必须被拒绝"
);
assert!(
resolve_image_reference_with("alpine:latest", mirror_env("evil .com")).is_err(),
"含空格的加速器地址必须被拒绝"
);
assert!(
resolve_image_reference_with("alpine:latest", mirror_env("evil\n.com")).is_err(),
"含换行的加速器地址必须被拒绝,否则可污染日志输出"
);
assert!(
resolve_image_reference_with("alpine:latest", mirror_env("evil\t.com")).is_err(),
"含制表符的加速器地址必须被拒绝"
);
assert_eq!(
resolve_image_reference_with("alpine:latest", mirror_env(" evil.example.com\n"))
.expect("首尾空白应被 trim 归一化后接受"),
"evil.example.com/alpine:latest"
);
assert!(
resolve_image_reference_with("alpine:latest", mirror_env(r"evil.com\path")).is_err(),
"含反斜杠的加速器地址必须被拒绝"
);
}
#[test]
fn test_mirror_with_legal_address_forms_is_accepted() {
assert_eq!(
resolve_image_reference_with("alpine:latest", mirror_env("mirror.example.com"))
.expect("合法主机名应被接受"),
"mirror.example.com/alpine:latest"
);
assert_eq!(
resolve_image_reference_with("alpine:latest", mirror_env("127.0.0.1:5000"))
.expect("合法 host:port 应被接受"),
"127.0.0.1:5000/alpine:latest"
);
assert_eq!(
resolve_image_reference_with(
"alpine:latest",
mirror_env("registry.example.com/docker")
)
.expect("合法路径前缀应被接受"),
"registry.example.com/docker/alpine:latest"
);
assert_eq!(
resolve_image_reference_with("alpine:latest", mirror_env("my-mirror-01.example.com"))
.expect("含连字符的合法主机名应被接受"),
"my-mirror-01.example.com/alpine:latest"
);
}
#[test]
fn test_mirror_rejects_non_http_schemes() {
for bad in ["file://evil/path", "ftp://evil.com", "gopher://evil.com"] {
let res = resolve_image_reference_with("alpine:latest", mirror_env(bad));
assert!(
matches!(res, Err(WslcError::InvalidConfiguration(_))),
"方案 {bad:?} 应被拒绝,实际为: {res:?}"
);
}
assert!(
resolve_image_reference_with("alpine:latest", mirror_env("http://m.example.com"))
.is_ok()
);
assert!(
resolve_image_reference_with("alpine:latest", mirror_env("https://m.example.com"))
.is_ok()
);
}
#[test]
fn test_mirror_with_multiple_scheme_separators_is_rejected() {
let res = resolve_image_reference_with(
"alpine:latest",
mirror_env("https://a.example.com://evil.com"),
);
assert!(
matches!(res, Err(WslcError::InvalidConfiguration(_))),
"含多个 scheme 分隔符的值必须报错而非被静默截断,实际为: {res:?}"
);
}
#[test]
fn test_mirror_empty_after_scheme_strip_is_rejected() {
let res = resolve_image_reference_with("alpine:latest", mirror_env("https://"));
assert!(
matches!(res, Err(WslcError::InvalidConfiguration(_))),
"剥离 scheme 后为空的地址必须报错,实际为: {res:?}"
);
}
#[test]
fn test_validate_mirror_address_boundaries() {
assert!(validate_mirror_address("").is_err());
assert!(validate_mirror_address("a").is_ok());
assert!(validate_mirror_address("a.b-c.d:1/x").is_ok());
assert!(validate_mirror_address("evil.com/镜像").is_err());
assert!(validate_mirror_address("evil.com/日本").is_err());
}
#[test]
fn test_debug_output_redacts_identity_token() {
let result = AuthTokenResult {
identity_token: "dG9rZW4tc2VjcmV0LXZhbHVl".to_string(),
token_type: WslcIdentityTokenType::Unknown,
};
let debug = format!("{result:?}");
assert!(
!debug.contains("dG9rZW4tc2VjcmV0LXZhbHVl"),
"Debug 输出绝不得包含令牌明文,实际为: {debug}"
);
assert!(
debug.contains("[已脱敏]"),
"令牌位置应以脱敏占位符呈现,实际为: {debug}"
);
assert!(
debug.contains("token_type"),
"令牌类型应保留在 Debug 输出中,实际为: {debug}"
);
}
}