use trussed::{
backend::BackendId,
virt::{self, StoreConfig},
};
use trussed_core::{syscall, try_syscall, types::ShortData};
use runner::Backends;
type Client<'a> = virt::Client<'a, Backends>;
mod extensions {
use serde::{Deserialize, Serialize};
use trussed_core::{
serde_extensions::{Extension, ExtensionClient, ExtensionResult},
types::ShortData,
Error,
};
pub struct TestExtension;
impl Extension for TestExtension {
type Request = TestRequest;
type Reply = TestReply;
}
#[derive(Deserialize, Serialize)]
pub enum TestRequest {
GetCalls(GetCallsRequest),
Reverse(ReverseRequest),
}
#[derive(Deserialize, Serialize)]
pub struct GetCallsRequest;
impl From<GetCallsRequest> for TestRequest {
fn from(request: GetCallsRequest) -> Self {
Self::GetCalls(request)
}
}
#[derive(Deserialize, Serialize)]
pub struct ReverseRequest {
pub s: ShortData,
}
impl From<ReverseRequest> for TestRequest {
fn from(request: ReverseRequest) -> Self {
Self::Reverse(request)
}
}
#[derive(Deserialize, Serialize)]
pub enum TestReply {
GetCalls(GetCallsReply),
Reverse(ReverseReply),
}
#[derive(Deserialize, Serialize)]
pub struct GetCallsReply {
pub calls: u64,
}
impl TryFrom<TestReply> for GetCallsReply {
type Error = Error;
fn try_from(reply: TestReply) -> Result<Self, Self::Error> {
match reply {
TestReply::GetCalls(reply) => Ok(reply),
_ => Err(Error::InternalError),
}
}
}
#[derive(Deserialize, Serialize)]
pub struct ReverseReply {
pub s: ShortData,
}
impl TryFrom<TestReply> for ReverseReply {
type Error = Error;
fn try_from(reply: TestReply) -> Result<Self, Self::Error> {
match reply {
TestReply::Reverse(reply) => Ok(reply),
_ => Err(Error::InternalError),
}
}
}
pub trait TestClient: ExtensionClient<TestExtension> {
fn test_calls(&mut self) -> ExtensionResult<'_, TestExtension, GetCallsReply, Self> {
self.extension(GetCallsRequest)
}
fn reverse(
&mut self,
s: ShortData,
) -> ExtensionResult<'_, TestExtension, ReverseReply, Self> {
self.extension(ReverseRequest { s })
}
}
impl<C: ExtensionClient<TestExtension>> TestClient for C {}
pub struct SampleExtension;
impl Extension for SampleExtension {
type Request = SampleRequest;
type Reply = SampleReply;
}
#[derive(Deserialize, Serialize)]
pub enum SampleRequest {
GetCalls(GetCallsRequest),
Truncate(TruncateRequest),
}
impl From<GetCallsRequest> for SampleRequest {
fn from(request: GetCallsRequest) -> Self {
Self::GetCalls(request)
}
}
#[derive(Deserialize, Serialize)]
pub struct TruncateRequest {
pub s: ShortData,
}
impl From<TruncateRequest> for SampleRequest {
fn from(request: TruncateRequest) -> Self {
Self::Truncate(request)
}
}
#[derive(Deserialize, Serialize)]
pub enum SampleReply {
GetCalls(GetCallsReply),
Truncate(TruncateReply),
}
impl TryFrom<SampleReply> for GetCallsReply {
type Error = Error;
fn try_from(reply: SampleReply) -> Result<Self, Self::Error> {
match reply {
SampleReply::GetCalls(reply) => Ok(reply),
_ => Err(Error::InternalError),
}
}
}
#[derive(Deserialize, Serialize)]
pub struct TruncateReply {
pub s: ShortData,
}
impl TryFrom<SampleReply> for TruncateReply {
type Error = Error;
fn try_from(reply: SampleReply) -> Result<Self, Self::Error> {
match reply {
SampleReply::Truncate(reply) => Ok(reply),
_ => Err(Error::InternalError),
}
}
}
pub trait SampleClient: ExtensionClient<SampleExtension> {
fn sample_calls(&mut self) -> ExtensionResult<'_, SampleExtension, GetCallsReply, Self> {
self.extension(GetCallsRequest)
}
fn truncate(
&mut self,
s: ShortData,
) -> ExtensionResult<'_, SampleExtension, TruncateReply, Self> {
self.extension(TruncateRequest { s })
}
}
impl<C: ExtensionClient<SampleExtension>> SampleClient for C {}
}
mod backends {
use super::extensions::{
GetCallsReply, ReverseReply, SampleExtension, SampleReply, SampleRequest, TestExtension,
TestReply, TestRequest, TruncateReply,
};
use trussed::{
backend::Backend, platform::Platform, serde_extensions::ExtensionImpl,
service::ServiceResources, types::CoreContext,
};
use trussed_core::{types::ShortData, Error};
#[derive(Default)]
pub struct TestContext {
calls: u64,
}
#[derive(Default)]
pub struct TestBackend;
impl Backend for TestBackend {
type Context = TestContext;
}
impl ExtensionImpl<TestExtension> for TestBackend {
fn extension_request<P: Platform>(
&mut self,
_core_ctx: &mut CoreContext,
backend_ctx: &mut TestContext,
request: &TestRequest,
_resources: &mut ServiceResources<P>,
) -> Result<TestReply, Error> {
match request {
TestRequest::GetCalls(_) => Ok(TestReply::GetCalls(GetCallsReply {
calls: backend_ctx.calls,
})),
TestRequest::Reverse(request) => {
backend_ctx.calls += 1;
let mut s = ShortData::new();
for byte in request.s.iter().rev() {
s.push(*byte).unwrap();
}
Ok(TestReply::Reverse(ReverseReply { s }))
}
}
}
}
#[derive(Default)]
pub struct SampleContext {
calls: u64,
}
#[derive(Default)]
pub struct SampleBackend;
impl Backend for SampleBackend {
type Context = SampleContext;
}
impl ExtensionImpl<SampleExtension> for SampleBackend {
fn extension_request<P: Platform>(
&mut self,
_core_ctx: &mut CoreContext,
backend_ctx: &mut SampleContext,
request: &SampleRequest,
_resources: &mut ServiceResources<P>,
) -> Result<SampleReply, Error> {
match request {
SampleRequest::GetCalls(_) => Ok(SampleReply::GetCalls(GetCallsReply {
calls: backend_ctx.calls,
})),
SampleRequest::Truncate(request) => {
backend_ctx.calls += 1;
let mut s = ShortData::new();
for byte in request.s.iter().take(3) {
s.push(*byte).unwrap();
}
Ok(SampleReply::Truncate(TruncateReply { s }))
}
}
}
}
impl ExtensionImpl<TestExtension> for SampleBackend {
fn extension_request<P: Platform>(
&mut self,
_core_ctx: &mut CoreContext,
backend_ctx: &mut SampleContext,
request: &TestRequest,
_resources: &mut ServiceResources<P>,
) -> Result<TestReply, Error> {
match request {
TestRequest::GetCalls(_) => Ok(TestReply::GetCalls(GetCallsReply {
calls: backend_ctx.calls,
})),
TestRequest::Reverse(request) => {
backend_ctx.calls += 1;
let mut s = ShortData::new();
for byte in request.s.iter().rev() {
s.push(*byte).unwrap();
}
Ok(TestReply::Reverse(ReverseReply { s }))
}
}
}
}
}
mod runner {
use super::{
backends::{SampleBackend, TestBackend},
extensions::{SampleExtension, TestExtension},
};
pub mod id {
pub enum Backend {
Test,
Sample,
}
#[derive(trussed_derive::ExtensionId)]
pub enum Extension {
Test = 37,
Sample = 42,
}
}
use trussed::backend::BackendId;
use trussed_derive::ExtensionDispatch;
#[derive(Default, ExtensionDispatch)]
#[dispatch(backend_id = "id::Backend", extension_id = "id::Extension")]
#[extensions(Test = "TestExtension", Sample = "SampleExtension")]
pub struct Backends {
#[extensions("Test")]
test: TestBackend,
#[extensions("Test", "Sample")]
sample: SampleBackend,
}
pub const BACKENDS_TEST1: &[BackendId<id::Backend>] =
&[BackendId::Custom(id::Backend::Test), BackendId::Core];
pub const BACKENDS_TEST2: &[BackendId<id::Backend>] =
&[BackendId::Core, BackendId::Custom(id::Backend::Test)];
pub const BACKENDS_SAMPLE1: &[BackendId<id::Backend>] =
&[BackendId::Custom(id::Backend::Sample), BackendId::Core];
pub const BACKENDS_SAMPLE2: &[BackendId<id::Backend>] =
&[BackendId::Core, BackendId::Custom(id::Backend::Sample)];
pub const BACKENDS_MIXED: &[BackendId<id::Backend>] = &[
BackendId::Custom(id::Backend::Test),
BackendId::Custom(id::Backend::Sample),
];
}
pub fn run<F: FnOnce(&mut Client<'_>)>(backends: &'static [BackendId<runner::id::Backend>], f: F) {
virt::with_platform(StoreConfig::ram(), |platform| {
platform.run_client_with_backends(
"test",
runner::Backends::default(),
backends,
|mut client| f(&mut client),
)
})
}
#[test]
fn test_extension() {
use extensions::TestClient as _;
let msg = ShortData::from(&[0x01, 0x02, 0x03]);
let rev = ShortData::from(&[0x03, 0x02, 0x01]);
run(&[], |client| {
assert!(try_syscall!(client.reverse(msg.clone())).is_err());
});
run(runner::BACKENDS_TEST1, |client| {
assert_eq!(syscall!(client.test_calls()).calls, 0);
assert_eq!(syscall!(client.reverse(msg.clone())).s, rev);
assert_eq!(syscall!(client.test_calls()).calls, 1);
assert_eq!(syscall!(client.test_calls()).calls, 1);
assert_eq!(syscall!(client.reverse(msg.clone())).s, rev);
assert_eq!(syscall!(client.test_calls()).calls, 2);
});
run(runner::BACKENDS_TEST2, |client| {
assert_eq!(syscall!(client.test_calls()).calls, 0);
assert_eq!(syscall!(client.reverse(msg.clone())).s, rev);
assert_eq!(syscall!(client.test_calls()).calls, 1);
});
}
#[test]
fn sample_extension() {
use extensions::SampleClient as _;
use extensions::TestClient as _;
let msg = ShortData::from(&[1, 2, 3, 4]);
let rev = ShortData::from(&[4, 3, 2, 1]);
let trunc = ShortData::from(&[1, 2, 3]);
run(&[], |client| {
assert!(try_syscall!(client.truncate(msg.clone())).is_err());
});
run(runner::BACKENDS_SAMPLE1, |client| {
assert_eq!(syscall!(client.sample_calls()).calls, 0);
assert_eq!(syscall!(client.test_calls()).calls, 0);
assert_eq!(syscall!(client.reverse(msg.clone())).s, rev);
assert_eq!(syscall!(client.truncate(msg.clone())).s, trunc);
assert_eq!(syscall!(client.sample_calls()).calls, 2);
assert_eq!(syscall!(client.test_calls()).calls, 2);
assert_eq!(syscall!(client.sample_calls()).calls, 2);
assert_eq!(syscall!(client.truncate(msg.clone())).s, trunc);
assert_eq!(syscall!(client.sample_calls()).calls, 3);
});
run(runner::BACKENDS_SAMPLE2, |client| {
assert_eq!(syscall!(client.sample_calls()).calls, 0);
assert_eq!(syscall!(client.truncate(msg.clone())).s, trunc);
assert_eq!(syscall!(client.sample_calls()).calls, 1);
});
}
#[test]
fn mixed_extension() {
use extensions::SampleClient as _;
use extensions::TestClient as _;
let msg = ShortData::from(&[1, 2, 3, 4]);
let rev = ShortData::from(&[4, 3, 2, 1]);
let trunc = ShortData::from(&[1, 2, 3]);
run(runner::BACKENDS_MIXED, |client| {
assert_eq!(syscall!(client.sample_calls()).calls, 0);
assert_eq!(syscall!(client.test_calls()).calls, 0);
assert_eq!(syscall!(client.reverse(msg.clone())).s, rev);
assert_eq!(syscall!(client.truncate(msg.clone())).s, trunc);
assert_eq!(syscall!(client.sample_calls()).calls, 1);
assert_eq!(syscall!(client.test_calls()).calls, 1);
assert_eq!(syscall!(client.truncate(msg.clone())).s, trunc);
assert_eq!(syscall!(client.sample_calls()).calls, 2);
});
}