From 2d78ec4c039f251b8968891fd3678b06bfb248fd Mon Sep 17 00:00:00 2001 From: James Hallowell <5460955+JamesHallowell@users.noreply.github.com> Date: Wed, 22 Jul 2026 23:16:00 +0100 Subject: [PATCH] Simplify some of the endpoint traits --- examples/endpoints.rs | 2 +- src/engine/mod.rs | 6 +++--- src/performer/endpoints/event.rs | 10 +++++++--- src/performer/endpoints/stream.rs | 16 +++++++++++++--- src/performer/endpoints/value.rs | 25 +++++++++++++++++-------- src/performer/mod.rs | 7 +++++-- tests/endpoints.rs | 12 ++++++------ tests/engine.rs | 4 ++-- 8 files changed, 54 insertions(+), 28 deletions(-) diff --git a/examples/endpoints.rs b/examples/endpoints.rs index 899c265..6d63995 100644 --- a/examples/endpoints.rs +++ b/examples/endpoints.rs @@ -79,7 +79,7 @@ fn main() -> Result<(), Box> { performer.advance(); assert_eq!( performer.get::(value_out_dynamic), - Ok(ValueRef::Int32(14)) + ValueRef::Int32(14) ); /* diff --git a/src/engine/mod.rs b/src/engine/mod.rs index 13764c4..6e03dc2 100644 --- a/src/engine/mod.rs +++ b/src/engine/mod.rs @@ -8,7 +8,7 @@ use { crate::{ endpoint::{EndpointHandle, EndpointInfo}, ffi::EnginePtr, - performer::{Endpoint, EndpointError, EndpointType, OutputEvent, Performer}, + performer::{Endpoint, EndpointError, MakeEndpoint, OutputEvent, Performer}, program::Program, }, std::{ @@ -194,7 +194,7 @@ impl Engine { /// Returns an endpoint handle. pub fn endpoint(&mut self, id: impl AsRef) -> Result, EndpointError> where - T: EndpointType, + T: MakeEndpoint, { let id = id.as_ref(); @@ -212,7 +212,7 @@ impl Engine { self.state.endpoints.insert(handle, info.clone()); - EndpointType::make(handle, info) + MakeEndpoint::make(handle, info) } /// Returns the details of the program loaded into the engine. diff --git a/src/performer/endpoints/event.rs b/src/performer/endpoints/event.rs index 2903979..9bf2386 100644 --- a/src/performer/endpoints/event.rs +++ b/src/performer/endpoints/event.rs @@ -1,6 +1,6 @@ use crate::{ endpoint::{EndpointDirection, EndpointHandle, EndpointInfo}, - performer::{Endpoint, EndpointError, EndpointType, Performer}, + performer::{Endpoint, EndpointError, GetHandle, MakeEndpoint, Performer}, value::ValueRef, }; @@ -16,7 +16,7 @@ pub struct OutputEvent { handle: EndpointHandle, } -impl EndpointType for InputEvent { +impl MakeEndpoint for InputEvent { fn make( handle: EndpointHandle, endpoint: EndpointInfo, @@ -31,13 +31,15 @@ impl EndpointType for InputEvent { Ok(Endpoint(InputEvent { handle })) } +} +impl GetHandle for InputEvent { fn handle(&self) -> EndpointHandle { self.handle } } -impl EndpointType for OutputEvent { +impl MakeEndpoint for OutputEvent { fn make(handle: EndpointHandle, endpoint: EndpointInfo) -> Result, EndpointError> where Self: Sized, @@ -52,7 +54,9 @@ impl EndpointType for OutputEvent { Ok(Endpoint(Self { handle })) } +} +impl GetHandle for OutputEvent { fn handle(&self) -> EndpointHandle { self.handle } diff --git a/src/performer/endpoints/stream.rs b/src/performer/endpoints/stream.rs index 27ba64b..177daab 100644 --- a/src/performer/endpoints/stream.rs +++ b/src/performer/endpoints/stream.rs @@ -1,7 +1,7 @@ use { crate::{ endpoint::{EndpointDirection, EndpointHandle, EndpointInfo}, - performer::{Endpoint, EndpointError, EndpointType, Performer}, + performer::{Endpoint, EndpointError, GetHandle, MakeEndpoint, Performer}, value::types::{IsScalar, Type}, }, std::marker::PhantomData, @@ -27,7 +27,7 @@ where _marker: PhantomData, } -impl EndpointType for InputStream +impl MakeEndpoint for InputStream where T: StreamType, { @@ -42,13 +42,18 @@ where _marker: PhantomData, })) } +} +impl GetHandle for InputStream +where + T: StreamType, +{ fn handle(&self) -> EndpointHandle { self.handle } } -impl EndpointType for OutputStream +impl MakeEndpoint for OutputStream where T: StreamType, { @@ -63,7 +68,12 @@ where _marker: PhantomData, })) } +} +impl GetHandle for OutputStream +where + T: StreamType, +{ fn handle(&self) -> EndpointHandle { self.handle } diff --git a/src/performer/endpoints/value.rs b/src/performer/endpoints/value.rs index 3e47c92..9ac9b3d 100644 --- a/src/performer/endpoints/value.rs +++ b/src/performer/endpoints/value.rs @@ -1,7 +1,7 @@ use { crate::{ endpoint::{EndpointDirection, EndpointHandle, EndpointInfo}, - performer::{endpoints::Endpoint, EndpointError, EndpointType, Performer}, + performer::{endpoints::Endpoint, EndpointError, GetHandle, MakeEndpoint, Performer}, value::{Value, ValueRef}, }, std::{any::TypeId, marker::PhantomData}, @@ -21,7 +21,7 @@ pub struct OutputValue { _marker: PhantomData, } -impl EndpointType for InputValue +impl MakeEndpoint for InputValue where T: 'static, { @@ -36,13 +36,15 @@ where _marker: PhantomData, })) } +} +impl GetHandle for InputValue { fn handle(&self) -> EndpointHandle { self.handle } } -impl EndpointType for OutputValue +impl MakeEndpoint for OutputValue where T: 'static, { @@ -57,7 +59,9 @@ where _marker: PhantomData, })) } +} +impl GetHandle for OutputValue { fn handle(&self) -> EndpointHandle { self.handle } @@ -225,7 +229,7 @@ impl GetOutputValue for bool { } impl GetOutputValue for Value { - type Output<'a> = Result, ()>; + type Output<'a> = ValueRef<'a>; fn get_output_value( performer: &mut Performer, @@ -237,11 +241,16 @@ impl GetOutputValue for Value { .endpoints .get(&endpoint.handle) .and_then(|endpoint| endpoint.as_value()) - .map(|value_endpoint| value_endpoint.ty().as_ref()) - .expect("failed to determine endpoint type"); + .map(|value_endpoint| value_endpoint.ty().as_ref()); - ptr.copy_output_value(endpoint.handle, buffer); + debug_assert!(ty.is_some(), "endpoint should exist and be a value type"); - Ok(ValueRef::new_from_slice(ty, &buffer[..ty.size()])) + match ty { + Some(ty) => { + ptr.copy_output_value(endpoint.handle, buffer); + ValueRef::new_from_slice(ty, &buffer[..ty.size()]) + } + None => ValueRef::Void, + } } } diff --git a/src/performer/mod.rs b/src/performer/mod.rs index 39ac997..5204486 100644 --- a/src/performer/mod.rs +++ b/src/performer/mod.rs @@ -74,7 +74,7 @@ impl Performer { /// Returns information about a given endpoint. pub fn endpoint_info(&self, Endpoint(endpoint): Endpoint) -> Option<&EndpointInfo> where - T: EndpointType, + T: GetHandle, { self.endpoints.get(&endpoint.handle()) } @@ -171,14 +171,17 @@ pub enum EndpointError { } #[doc(hidden)] -pub trait EndpointType: sealed::Sealed { +pub trait MakeEndpoint: sealed::Sealed { fn make( handle: EndpointHandle, endpoint: EndpointInfo, ) -> Result, EndpointError> where Self: Sized; +} +#[doc(hidden)] +pub trait GetHandle: sealed::Sealed { fn handle(&self) -> EndpointHandle; } diff --git a/tests/endpoints.rs b/tests/endpoints.rs index d85c1d9..f228ae1 100644 --- a/tests/endpoints.rs +++ b/tests/endpoints.rs @@ -132,7 +132,7 @@ fn can_read_and_write_complex32_numbers() { performer.advance(); - let result: Complex32 = performer.get::(output).unwrap().try_into().unwrap(); + let result: Complex32 = performer.get::(output).try_into().unwrap(); assert_eq!( result, @@ -175,7 +175,7 @@ fn can_read_and_write_complex64_numbers() { performer.advance(); - let result: Complex64 = performer.get::(output).unwrap().try_into().unwrap(); + let result: Complex64 = performer.get::(output).try_into().unwrap(); assert_eq!( result, @@ -212,7 +212,7 @@ fn can_read_structs() { performer.advance(); - let value = performer.get::(output).unwrap(); + let value = performer.get::(output); let object = value.as_object().unwrap(); assert_eq!(object.field("a").unwrap(), ValueRef::Bool(true)); @@ -247,7 +247,7 @@ fn can_read_and_write_arrays() { performer.advance(); - let value = performer.get::(output).unwrap(); + let value = performer.get::(output); let array = value.as_array().unwrap(); assert_eq!(array.len(), 4); @@ -511,7 +511,7 @@ fn read_and_write_vectors() { performer.set::(input, [1, 2, 3, 4].into()).unwrap(); performer.advance(); - let value = performer.get::(output).unwrap(); + let value = performer.get::(output); let array = value.as_array().unwrap(); let elems: Vec<_> = array.elems().collect(); @@ -770,7 +770,7 @@ fn string_endpoints() { performer.advance(); - let value = if let ValueRef::String(string) = performer.get::(out).unwrap() { + let value = if let ValueRef::String(string) = performer.get::(out) { string } else { panic!("expected string"); diff --git a/tests/engine.rs b/tests/engine.rs index 5f87001..33876b2 100644 --- a/tests/engine.rs +++ b/tests/engine.rs @@ -232,7 +232,7 @@ fn loading_external_variables_struct() { performer.advance(); - let result: Complex32 = performer.get(out).unwrap().try_into().unwrap(); + let result: Complex32 = performer.get(out).try_into().unwrap(); assert_eq!(result.real, 42.0); assert_eq!(result.imag, 21.0); } @@ -262,7 +262,7 @@ fn loading_external_variables_array() { performer.advance(); - let value = performer.get(out).unwrap(); + let value = performer.get(out); let array = value.as_array().unwrap(); assert_eq!(array.get(0), Some(ValueRef::Int32(1)));