use std::sync::{
Arc, LazyLock, Mutex,
atomic::{AtomicI32, Ordering},
};
use dashmap::DashMap;
use super::safe_ptr;
use crate::{Error, ErrorKind, ErrorOrigin, Result, cc_client::client::CcClient, raw};
type ContextMap = DashMap<i32, Arc<Mutex<CcClient>>>;
static CONTEXTS: LazyLock<ContextMap> = LazyLock::new(DashMap::new);
type RegisteredCtxMap = DashMap<usize, i32>;
static REGISTERED_CTXS: LazyLock<RegisteredCtxMap> = LazyLock::new(DashMap::new);
pub struct ContextManager;
impl ContextManager {
pub fn add_context(ctx: *mut raw::TEEC_Context, client: CcClient) -> Result<()> {
let mut ctx_nn = safe_ptr::deref_mut(ctx)?;
let ctx_ref = unsafe { ctx_nn.as_mut() };
let ctx_id: i32 = ctx_ref.imp.fd;
let client_arc = Arc::new(Mutex::new(client));
if CONTEXTS.contains_key(&ctx_id) {
log::warn!(
"ContextManager: context {} already exists, rejecting duplicate",
ctx_id
);
return Err(Error::new(ErrorKind::BadState).with_origin(ErrorOrigin::API));
}
CONTEXTS.insert(ctx_id, client_arc);
REGISTERED_CTXS.insert(ctx as usize, ctx_id);
Ok(())
}
pub fn registered_id(ctx: *mut raw::TEEC_Context) -> Option<i32> {
REGISTERED_CTXS
.get(&(ctx as usize))
.map(|entry| *entry.value())
}
pub fn is_registered(ctx: *mut raw::TEEC_Context) -> bool {
!ctx.is_null() && REGISTERED_CTXS.contains_key(&(ctx as usize))
}
pub fn remove_context(ctx: *mut raw::TEEC_Context) {
let Some(registered_id) = Self::registered_id(ctx) else {
log::warn!("ContextManager: remove_context: ctx 未登记,忽略");
return;
};
if let Ok(mut ctx_nn) = safe_ptr::deref_mut(ctx) {
let ctx_ref = unsafe { ctx_nn.as_mut() };
let ctx_id = ctx_ref.imp.fd;
if ctx_id != registered_id {
log::warn!(
"ContextManager: remove_context: imp.fd {ctx_id} 与登记 id {registered_id} 不一致,忽略"
);
return;
}
if let Some((_, client_arc)) = CONTEXTS.remove(&ctx_id)
&& let Ok(mut client) = client_arc.lock()
{
client.close();
}
ctx_ref.imp.fd = -1;
REGISTERED_CTXS.remove(&(ctx as usize));
}
}
pub fn get_client(ctx: *mut raw::TEEC_Context) -> Result<Arc<Mutex<CcClient>>> {
let registered_id = REGISTERED_CTXS
.get(&(ctx as usize))
.map(|entry| *entry.value())
.ok_or_else(|| Error::new(ErrorKind::BadParameters).with_origin(ErrorOrigin::API))?;
let ctx_nn = safe_ptr::deref(ctx)?;
let ctx_ref = unsafe { ctx_nn.as_ref() };
let ctx_id = ctx_ref.imp.fd;
if ctx_id != registered_id {
return Err(Error::new(ErrorKind::BadParameters).with_origin(ErrorOrigin::API));
}
CONTEXTS
.get(&ctx_id)
.map(|entry| entry.value().clone())
.ok_or_else(|| Error::new(ErrorKind::Generic).with_origin(ErrorOrigin::API))
}
}
static CONTEXT_ID_COUNTER: LazyLock<AtomicI32> = LazyLock::new(|| AtomicI32::new(0));
pub(crate) fn initialize_context_impl(ctx: *mut raw::TEEC_Context) -> Result<()> {
let mut ctx_nn = safe_ptr::deref_mut(ctx)?;
let ctx_ref = unsafe { ctx_nn.as_mut() };
let client = CcClient::init().map_err(|e| {
log::warn!("TEEC_InitializeContext:初始化机密通信上下文失败:{e}");
Error::new(ErrorKind::Communication)
})?;
let id = CONTEXT_ID_COUNTER.fetch_add(1, Ordering::SeqCst);
ctx_ref.imp.fd = id;
ctx_ref.imp.reg_mem = true;
ctx_ref.imp.memref_null = true;
ContextManager::add_context(ctx, client)
}
#[cfg(test)]
mod context_tests {
use super::*;
use std::ptr;
fn create_test_context(id: i32) -> raw::TEEC_Context {
raw::TEEC_Context {
imp: raw::TEEC_Context__Imp {
fd: id,
memref_null: false,
reg_mem: false,
},
}
}
#[test]
fn test_remove_context_null_ptr() {
ContextManager::remove_context(ptr::null_mut());
}
#[test]
fn test_get_client_null_ptr() {
let result = ContextManager::get_client(ptr::null_mut());
assert!(result.is_err(), "空指针应该导致失败");
}
#[test]
fn test_get_client_unregistered() {
let mut ctx = create_test_context(999);
let result = ContextManager::get_client(&mut ctx as *mut raw::TEEC_Context);
assert!(result.is_err(), "未注册的上下文应该导致失败");
}
#[test]
fn test_context_lifecycle() {
let mut ctx = create_test_context(1000);
let ctx_addr = &mut ctx as *mut raw::TEEC_Context as usize;
REGISTERED_CTXS.insert(ctx_addr, 1000);
assert_eq!(ctx.imp.fd, 1000);
ContextManager::remove_context(&mut ctx as *mut raw::TEEC_Context);
assert_eq!(ctx.imp.fd, -1, "移除后 ID 应该被重置为 -1");
assert!(
!REGISTERED_CTXS.contains_key(&ctx_addr),
"移除后应清除地址登记"
);
}
#[test]
fn test_remove_context_unregistered_noop() {
let mut ctx = create_test_context(1001);
ContextManager::remove_context(&mut ctx as *mut raw::TEEC_Context);
assert_eq!(ctx.imp.fd, 1001, "未登记的上下文不应被修改");
}
#[test]
fn test_remove_context_fd_tamper_rejected() {
let mut ctx = create_test_context(2002);
let ctx_addr = &mut ctx as *mut raw::TEEC_Context as usize;
REGISTERED_CTXS.insert(ctx_addr, 2002);
ctx.imp.fd = 1337;
ContextManager::remove_context(&mut ctx as *mut raw::TEEC_Context);
assert_eq!(ctx.imp.fd, 1337, "fd 篡改时不应写入结构体");
assert!(
REGISTERED_CTXS.contains_key(&ctx_addr),
"fd 篡改时不应清除登记"
);
REGISTERED_CTXS.remove(&ctx_addr);
}
#[test]
fn test_get_client_registered_addr_fd_mismatch_rejected() {
let mut ctx = create_test_context(2000);
let ctx_addr = &mut ctx as *mut raw::TEEC_Context as usize;
REGISTERED_CTXS.insert(ctx_addr, 2000);
ctx.imp.fd = 1999;
let result = ContextManager::get_client(&mut ctx as *mut raw::TEEC_Context);
assert!(
result.is_err(),
"imp.fd 与地址登记的 context_id 不一致应被拒绝"
);
REGISTERED_CTXS.remove(&ctx_addr);
}
#[test]
fn test_remove_context_clears_addr_registration() {
let mut ctx = create_test_context(2001);
let ctx_addr = &mut ctx as *mut raw::TEEC_Context as usize;
REGISTERED_CTXS.insert(ctx_addr, 2001);
ContextManager::remove_context(&mut ctx as *mut raw::TEEC_Context);
assert!(
!REGISTERED_CTXS.contains_key(&ctx_addr),
"移除上下文后应清除地址登记"
);
}
}