use std::path::{Path, PathBuf};
use std::sync::Mutex;
use onnx_runtime_ep_api::{
DeviceBuffer, EpConfig, EpContext, EpError, EpId, ExecutionProvider, Fence, Kernel,
KernelMatch, Result as EpResult,
};
use onnx_runtime_ir::{
Attribute, DataType, DeviceId, DeviceType, Graph, Node, NodeId, Shape,
TensorLayout, ValueId, static_shape,
};
use onnx_runtime_loader::{
EpContextDumpConfig, Model, ep_context_nodes as loader_ep_context_nodes, load_model,
};
use onnx_runtime_session::{
CompiledPartition, InferenceSession, SessionError, dump_session_ep_context,
load_ep_context_nodes,
};
struct MockCompiledEp {
keys: Vec<String>,
loaded: Mutex<Vec<Vec<u8>>>,
save_blob: Vec<u8>,
save_version: String,
}
impl MockCompiledEp {
fn new() -> Self {
Self {
keys: vec!["MOCK".to_string()],
loaded: Mutex::new(Vec::new()),
save_blob: Vec::new(),
save_version: String::new(),
}
}
fn compiling(blob: &[u8], version: &str) -> Self {
Self {
keys: vec!["MOCK".to_string()],
loaded: Mutex::new(Vec::new()),
save_blob: blob.to_vec(),
save_version: version.to_string(),
}
}
fn loaded(&self) -> Vec<Vec<u8>> {
self.loaded.lock().unwrap().clone()
}
}
impl ExecutionProvider for MockCompiledEp {
fn name(&self) -> &str {
"mock_compiled_ep"
}
fn device_type(&self) -> DeviceType {
DeviceType::Custom(0)
}
fn device_id(&self) -> DeviceId {
DeviceId::new(DeviceType::Custom(0), 0)
}
fn initialize(&mut self, _config: &EpConfig) -> EpResult<()> {
Ok(())
}
fn shutdown(&mut self) -> EpResult<()> {
Ok(())
}
fn supports_op(&self, _op: &Node, _shapes: &[Shape], _layouts: &[TensorLayout]) -> KernelMatch {
KernelMatch::Unsupported
}
fn get_kernel(
&self,
_op: &Node,
_shapes: &[Vec<usize>],
_opset: u64,
) -> EpResult<Box<dyn Kernel>> {
Err(EpError::NoEpForOp {
op_type: "<mock>".to_string(),
})
}
fn allocate(&self, _size: usize, _alignment: usize) -> EpResult<DeviceBuffer> {
Err(EpError::NotInitialized)
}
fn deallocate(&self, _buffer: DeviceBuffer) -> EpResult<()> {
Ok(())
}
fn copy(&self, _src: &DeviceBuffer, _dst: &mut DeviceBuffer, _size: usize) -> EpResult<()> {
Ok(())
}
fn copy_async(
&self,
_src: &DeviceBuffer,
_dst: &mut DeviceBuffer,
_size: usize,
) -> EpResult<Fence> {
Ok(Fence::default())
}
fn sync(&self) -> EpResult<()> {
Ok(())
}
fn context_source_keys(&self) -> Vec<String> {
self.keys.clone()
}
fn save_context(&self) -> EpResult<EpContext> {
Ok(EpContext::new(
self.name(),
self.save_version.clone(),
self.save_blob.clone(),
Vec::new(),
"mock-device",
))
}
fn load_context(&self, ctx: &EpContext) -> EpResult<()> {
self.loaded.lock().unwrap().push(ctx.data.clone());
Ok(())
}
}
fn s_attr(v: &str) -> Attribute {
Attribute::String(v.as_bytes().to_vec())
}
fn i_attr(v: i64) -> Attribute {
Attribute::Int(v)
}
fn embedded_blob_attr(bytes: &[u8]) -> Attribute {
Attribute::String(bytes.to_vec())
}
fn add_epctx_node(
g: &mut Graph,
input: ValueId,
tag: &str,
attrs: Vec<(&str, Attribute)>,
) -> NodeId {
let out = g.create_named_value(
format!("Y_{tag}"),
DataType::Float32,
static_shape([2usize, 8]),
);
let mut node = Node::new(NodeId(0), "EPContext", vec![Some(input)], vec![out]);
node.domain = "com.microsoft".to_string();
for (k, v) in attrs {
node.attributes.insert(k.to_string(), v);
}
let id = g.insert_node(node);
g.add_output(out);
id
}
fn graph_with_input() -> (Graph, ValueId) {
let mut g = Graph::new();
g.opset_imports
.insert("com.microsoft".to_string(), 1);
let x = g.create_named_value("X", DataType::Float32, static_shape([2usize, 4]));
g.add_input(x);
(g, x)
}
fn eps(mock: &MockCompiledEp) -> [(EpId, &dyn ExecutionProvider); 1] {
[(EpId(0), mock)]
}
#[test]
fn embed_mode_round_trip_dispatches_exact_bytes() {
let payload: Vec<u8> = vec![0x00, 0x01, 0x80, 0xFE, 0xFF, b'v', b'1', 0xC3, 0x28];
let (mut g, x) = graph_with_input();
add_epctx_node(
&mut g,
x,
"a",
vec![
("embed_mode", i_attr(1)),
("main_context", i_attr(1)),
("source", s_attr("MOCK")),
("ep_sdk_version", s_attr("9.9.9")),
("ep_cache_context", embedded_blob_attr(&payload)),
],
);
let mock = MockCompiledEp::new();
let placement = load_ep_context_nodes(&g, Path::new("."), &eps(&mock)).expect("dispatch");
assert_eq!(placement.handled.len(), 1, "one EPContext node handled");
assert_eq!(
mock.loaded(),
vec![payload],
"load_context received the exact inline blob bytes"
);
}
#[test]
fn external_mode_round_trip_loads_file_bytes() {
let dir: PathBuf = Path::new(env!("CARGO_TARGET_TMPDIR")).join("epctx_external");
std::fs::create_dir_all(&dir).unwrap();
let payload: Vec<u8> = vec![0xDE, 0xAD, 0xBE, 0xEF, 0x00, 0x7F, 0x80];
let bin = dir.join("ctx_blob.bin");
std::fs::write(&bin, &payload).unwrap();
let (mut g, x) = graph_with_input();
add_epctx_node(
&mut g,
x,
"ext",
vec![
("embed_mode", i_attr(0)),
("main_context", i_attr(1)),
("source", s_attr("MOCK")),
("ep_cache_context", s_attr("ctx_blob.bin")),
],
);
let mock = MockCompiledEp::new();
let placement = load_ep_context_nodes(&g, &dir, &eps(&mock)).expect("dispatch");
assert_eq!(placement.handled.len(), 1);
assert_eq!(
mock.loaded(),
vec![payload],
"load_context received the external file's exact bytes"
);
}
#[test]
fn unclaimed_source_surfaces_no_ep_for_context() {
let (mut g, x) = graph_with_input();
add_epctx_node(
&mut g,
x,
"qnn",
vec![
("embed_mode", i_attr(1)),
("source", s_attr("QNN")),
("ep_cache_context", embedded_blob_attr(b"blob")),
],
);
let mock = MockCompiledEp::new();
let err = load_ep_context_nodes(&g, Path::new("."), &eps(&mock)).expect_err("must fail");
match err {
SessionError::Ep(EpError::NoEpForContext { source_key }) => {
assert_eq!(source_key.as_deref(), Some("QNN"));
}
other => panic!("expected NoEpForContext {{ source_key: QNN }}, got {other:?}"),
}
assert!(mock.loaded().is_empty(), "no context restored on failure");
}
#[test]
fn main_context_dedup_and_reference_resolution() {
let shared: Vec<u8> = b"one-packed-binary-holding-two-graphs".to_vec();
let (mut g, x) = graph_with_input();
add_epctx_node(
&mut g,
x,
"p0",
vec![
("embed_mode", i_attr(1)),
("main_context", i_attr(1)),
("source", s_attr("MOCK")),
("partition_name", s_attr("p0")),
("ep_cache_context", embedded_blob_attr(&shared)),
],
);
add_epctx_node(
&mut g,
x,
"p1",
vec![
("embed_mode", i_attr(1)),
("main_context", i_attr(1)),
("source", s_attr("MOCK")),
("partition_name", s_attr("p1")),
("ep_cache_context", embedded_blob_attr(&shared)),
],
);
add_epctx_node(
&mut g,
x,
"ref",
vec![
("embed_mode", i_attr(1)),
("main_context", i_attr(0)),
("source", s_attr("MOCK")),
("partition_name", s_attr("p1")),
],
);
let mock = MockCompiledEp::new();
let placement = load_ep_context_nodes(&g, Path::new("."), &eps(&mock)).expect("dispatch");
assert_eq!(placement.handled.len(), 3);
assert_eq!(
mock.loaded(),
vec![shared],
"identical payload loaded once; reference resolved without a second load"
);
}
#[test]
fn dangling_reference_is_an_error() {
let (mut g, x) = graph_with_input();
add_epctx_node(
&mut g,
x,
"ref",
vec![
("main_context", i_attr(0)),
("source", s_attr("MOCK")),
("partition_name", s_attr("missing")),
],
);
let mock = MockCompiledEp::new();
let err = load_ep_context_nodes(&g, Path::new("."), &eps(&mock)).expect_err("must fail");
match err {
SessionError::DanglingEpContext {
source_key,
partition_name,
} => {
assert_eq!(source_key.as_deref(), Some("MOCK"));
assert_eq!(partition_name.as_deref(), Some("missing"));
}
other => panic!("expected DanglingEpContext, got {other:?}"),
}
assert!(mock.loaded().is_empty());
}
#[test]
fn duplicate_source_key_across_eps_is_rejected() {
let (mut g, x) = graph_with_input();
add_epctx_node(
&mut g,
x,
"a",
vec![
("source", s_attr("MOCK")),
("ep_cache_context", embedded_blob_attr(b"x")),
],
);
let a = MockCompiledEp::new();
let b = MockCompiledEp::new();
let eps: [(EpId, &dyn ExecutionProvider); 2] = [(EpId(0), &a), (EpId(1), &b)];
let err = load_ep_context_nodes(&g, Path::new("."), &eps).expect_err("must fail");
assert!(
matches!(
err,
SessionError::Ep(EpError::DuplicateContextSource { .. })
),
"expected DuplicateContextSource, got {err:?}"
);
}
#[test]
fn session_build_rejects_unclaimed_ep_context_node() {
let (mut g, x) = graph_with_input();
add_epctx_node(
&mut g,
x,
"a",
vec![
("embed_mode", i_attr(1)),
("source", s_attr("SomeCompiledEP")),
("ep_cache_context", embedded_blob_attr(b"compiled-blob")),
],
);
let err = match InferenceSession::from_graph(g) {
Ok(_) => panic!("CPU-only session must not claim a compiled-EP context node"),
Err(e) => e,
};
assert!(
matches!(err, SessionError::Ep(EpError::NoEpForContext { .. })),
"expected NoEpForContext, got {err:?}"
);
}
fn build_partition_graph() -> (Graph, Vec<NodeId>) {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = g.create_named_value("X", DataType::Float32, static_shape([2usize, 4]));
g.add_input(x);
let h = g.create_named_value("H", DataType::Float32, static_shape([2usize, 4]));
let id1 = g.insert_node(Node::new(NodeId(0), "Relu", vec![Some(x)], vec![h]));
let y = g.create_named_value("Y", DataType::Float32, static_shape([2usize, 4]));
let id2 = g.insert_node(Node::new(NodeId(0), "Relu", vec![Some(h)], vec![y]));
g.add_output(y);
(g, vec![id1, id2])
}
fn input_names(g: &Graph, id: NodeId) -> Vec<String> {
g.node(id)
.inputs
.iter()
.map(|slot| match slot {
Some(v) => g.value(*v).name.clone().unwrap_or_default(),
None => String::new(),
})
.collect()
}
fn output_names(g: &Graph, id: NodeId) -> Vec<String> {
g.node(id)
.outputs
.iter()
.map(|v| g.value(*v).name.clone().unwrap_or_default())
.collect()
}
fn compiled_blob() -> Vec<u8> {
vec![0x00, 0x01, 0x80, 0xFE, 0xFF, b'c', b'x', 0xC3, 0x28, 0x00, 0x7F]
}
#[test]
fn dump_embed_round_trip_is_byte_exact() {
let dir = Path::new(env!("CARGO_TARGET_TMPDIR")).join("epctx_dump_embed");
std::fs::create_dir_all(&dir).unwrap();
let orig = dir.join("mymodel.onnx");
let payload = compiled_blob();
let (g, covered) = build_partition_graph();
let ep = MockCompiledEp::compiling(&payload, "7.7.7");
let model = Model::new(&g);
let config = EpContextDumpConfig {
enable: true,
file_path: None,
embed_mode: 1,
};
let parts = [CompiledPartition {
ep: &ep,
partition_name: "part0".to_string(),
covered_nodes: covered,
}];
let out = dump_session_ep_context(&model, &orig, &parts, &config).expect("dump");
assert_eq!(out, dir.join("mymodel_ctx.onnx"), "default <stem>_ctx.onnx path");
assert!(out.exists(), "context model written");
let g2 = load_model(&out).expect("reload ctx model");
let ids: Vec<NodeId> = loader_ep_context_nodes(&g2).map(|n| n.node).collect();
assert_eq!(ids.len(), 1, "the partition collapsed to one EPContext node");
let ep_id = ids[0];
assert_eq!(input_names(&g2, ep_id), vec!["X".to_string()], "boundary input");
assert_eq!(output_names(&g2, ep_id), vec!["Y".to_string()], "boundary output");
assert!(
g2.nodes.values().all(|n| n.op_type == "EPContext"),
"only the EPContext node remains"
);
let mock = MockCompiledEp::new();
let placement = load_ep_context_nodes(&g2, &dir, &eps(&mock)).expect("consume");
assert_eq!(placement.handled.len(), 1);
assert_eq!(
mock.loaded(),
vec![payload],
"the embedded blob round-tripped byte-exact"
);
}
#[test]
fn dump_external_round_trip_via_sidecar_bin() {
let dir = Path::new(env!("CARGO_TARGET_TMPDIR")).join("epctx_dump_external");
std::fs::create_dir_all(&dir).unwrap();
let orig = dir.join("net.onnx");
let payload = compiled_blob();
let (g, covered) = build_partition_graph();
let ep = MockCompiledEp::compiling(&payload, "3.1.4");
let model = Model::new(&g);
let config = EpContextDumpConfig {
enable: true,
file_path: None,
embed_mode: 0,
};
let parts = [CompiledPartition {
ep: &ep,
partition_name: "part0".to_string(),
covered_nodes: covered,
}];
let out = dump_session_ep_context(&model, &orig, &parts, &config).expect("dump");
assert_eq!(out, dir.join("net_ctx.onnx"));
let sidecar = dir.join("net_ctx_p0_MOCK_part0.bin");
assert!(sidecar.exists(), "external sidecar written next to ctx model");
assert_eq!(std::fs::read(&sidecar).unwrap(), payload, "sidecar holds the blob");
let g2 = load_model(&out).expect("reload ctx model");
let mock = MockCompiledEp::new();
let placement = load_ep_context_nodes(&g2, &dir, &eps(&mock)).expect("consume");
assert_eq!(placement.handled.len(), 1);
assert_eq!(
mock.loaded(),
vec![payload],
"the external blob round-tripped byte-exact"
);
}
#[test]
fn dump_honours_explicit_output_path() {
let dir = Path::new(env!("CARGO_TARGET_TMPDIR")).join("epctx_dump_explicit");
std::fs::create_dir_all(&dir).unwrap();
let orig = dir.join("src.onnx");
let explicit = dir.join("chosen_name.onnx");
let payload = compiled_blob();
let (g, covered) = build_partition_graph();
let ep = MockCompiledEp::compiling(&payload, "1.0.0");
let model = Model::new(&g);
let config = EpContextDumpConfig {
enable: true,
file_path: Some(explicit.clone()),
embed_mode: 1,
};
let parts = [CompiledPartition {
ep: &ep,
partition_name: "p".to_string(),
covered_nodes: covered,
}];
let out = dump_session_ep_context(&model, &orig, &parts, &config).expect("dump");
assert_eq!(out, explicit);
assert!(explicit.exists());
}
#[test]
fn dump_disabled_config_is_a_no_op() {
let dir = Path::new(env!("CARGO_TARGET_TMPDIR")).join("epctx_dump_disabled");
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let orig = dir.join("mymodel.onnx");
let payload = compiled_blob();
let (g, covered) = build_partition_graph();
let ep = MockCompiledEp::compiling(&payload, "7.7.7");
let model = Model::new(&g);
let config = EpContextDumpConfig {
enable: false,
file_path: None,
embed_mode: 0,
};
let parts = [CompiledPartition {
ep: &ep,
partition_name: "part0".to_string(),
covered_nodes: covered,
}];
let out = dump_session_ep_context(&model, &orig, &parts, &config).expect("dump");
assert_eq!(out, dir.join("mymodel_ctx.onnx"), "returns the would-be path");
assert!(!out.exists(), "disabled config writes no ctx model");
let entries: Vec<_> = std::fs::read_dir(&dir).unwrap().collect();
assert!(entries.is_empty(), "disabled config writes no files at all");
}
#[test]
fn builder_options_drive_export_byte_exact() {
let dir = Path::new(env!("CARGO_TARGET_TMPDIR")).join("epctx_builder_e2e");
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let orig = dir.join("orig.onnx");
let (g, _covered_in_g) = build_partition_graph();
let bytes = onnx_runtime_loader::encode_model(&Model::new(&g)).expect("encode");
std::fs::write(&orig, &bytes).unwrap();
let ctx_out = dir.join("explicit_ctx.onnx");
let session = InferenceSession::builder()
.model(&orig)
.option("ep.context_enable", "1")
.option("ep.context_file_path", ctx_out.to_str().unwrap())
.option("ep.context_embed_mode", "1")
.build()
.expect("build session");
let cfg = session.ep_context_config();
assert!(cfg.enable);
assert_eq!(cfg.file_path.as_deref(), Some(ctx_out.as_path()));
assert_eq!(cfg.embed_mode, 1);
let covered: Vec<NodeId> = session
.graph()
.nodes
.iter()
.filter(|(_, n)| n.op_type == "Relu")
.map(|(id, _)| id)
.collect();
assert_eq!(covered.len(), 2, "two Relus form the partition");
let payload = compiled_blob();
let ep = MockCompiledEp::compiling(&payload, "5.5.5");
let parts = [CompiledPartition {
ep: &ep,
partition_name: "part0".to_string(),
covered_nodes: covered,
}];
let out = session.export_ep_context(&orig, &parts).expect("export");
assert_eq!(out, ctx_out, "export writes to the configured file_path");
assert!(out.exists(), "context model written");
let g2 = load_model(&out).expect("reload ctx model");
let ids: Vec<NodeId> = loader_ep_context_nodes(&g2).map(|n| n.node).collect();
assert_eq!(ids.len(), 1, "the partition collapsed to one EPContext node");
let mock = MockCompiledEp::new();
let placement = load_ep_context_nodes(&g2, &dir, &eps(&mock)).expect("consume");
assert_eq!(placement.handled.len(), 1);
assert_eq!(
mock.loaded(),
vec![payload],
"the exported blob round-tripped byte-exact through the builder path"
);
}
#[test]
fn builder_disabled_export_writes_nothing() {
let dir = Path::new(env!("CARGO_TARGET_TMPDIR")).join("epctx_builder_disabled");
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let orig = dir.join("orig.onnx");
let (g, _c) = build_partition_graph();
let bytes = onnx_runtime_loader::encode_model(&Model::new(&g)).expect("encode");
std::fs::write(&orig, &bytes).unwrap();
let session = InferenceSession::builder()
.model(&orig)
.build()
.expect("build session");
assert!(!session.ep_context_config().enable);
let covered: Vec<NodeId> = session
.graph()
.nodes
.iter()
.filter(|(_, n)| n.op_type == "Relu")
.map(|(id, _)| id)
.collect();
let ep = MockCompiledEp::compiling(&compiled_blob(), "0.0.0");
let parts = [CompiledPartition {
ep: &ep,
partition_name: "p".to_string(),
covered_nodes: covered,
}];
let out = session.export_ep_context(&orig, &parts).expect("export no-op");
assert_eq!(out, dir.join("orig_ctx.onnx"), "returns the would-be path");
assert!(!out.exists(), "disabled config writes no ctx model");
let names: Vec<String> = std::fs::read_dir(&dir)
.unwrap()
.map(|e| e.unwrap().file_name().to_string_lossy().into_owned())
.collect();
assert_eq!(names, vec!["orig.onnx".to_string()], "no extra files written");
}