Skip to main content

gwr_engine/port/
mod.rs

1// Copyright (c) 2024 Graphcore Ltd. All rights reserved.
2
3//! Port
4
5use std::cell::RefCell;
6use std::fmt;
7use std::pin::Pin;
8use std::rc::Rc;
9use std::task::{Context, Poll, Waker};
10
11use futures::Future;
12use futures::future::FusedFuture;
13use gwr_track::connect;
14use gwr_track::entity::{Entity, GetEntity};
15use gwr_track::tracker::aka::Aka;
16
17use crate::engine::Engine;
18use crate::port::monitor::Monitor;
19use crate::sim_error;
20use crate::time::clock::Clock;
21use crate::traits::SimObject;
22use crate::types::{SimError, SimResult};
23
24pub mod monitor;
25
26pub type PortStateResult<T> = Result<Rc<PortState<T>>, SimError>;
27pub type PortGetResult<T> = Result<PortGet<T>, SimError>;
28pub type PortStartGetResult<T> = Result<PortStartGet<T>, SimError>;
29pub type PortPutResult<T> = Result<PortPut<T>, SimError>;
30pub type PortTryPutResult<T> = Result<PortTryPut<T>, SimError>;
31
32pub struct PortState<T>
33where
34    T: SimObject,
35{
36    value: RefCell<Option<T>>,
37    put_released: RefCell<bool>,
38    waiting_get: RefCell<Option<Waker>>,
39    waiting_put: RefCell<Option<Waker>>,
40    pub in_port_entity: Rc<Entity>,
41    monitor: Option<Rc<Monitor>>,
42}
43
44impl<T> PortState<T>
45where
46    T: SimObject,
47{
48    fn new(
49        engine: &Engine,
50        clock: &Clock,
51        in_port_entity: Rc<Entity>,
52        window_size_ticks: Option<u64>,
53    ) -> Self {
54        let monitor = window_size_ticks.map(|window_size_ticks| {
55            Monitor::new_and_register(engine, &in_port_entity, clock, window_size_ticks)
56        });
57        Self {
58            value: RefCell::new(None),
59            put_released: RefCell::new(true),
60            waiting_get: RefCell::new(None),
61            waiting_put: RefCell::new(None),
62            in_port_entity,
63            monitor,
64        }
65    }
66}
67
68pub struct InPort<T>
69where
70    T: SimObject,
71{
72    entity: Rc<Entity>,
73    state: Rc<PortState<T>>,
74    connected: RefCell<bool>,
75}
76
77impl<T> fmt::Display for InPort<T>
78where
79    T: SimObject,
80{
81    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
82        self.entity.fmt(f)
83    }
84}
85
86impl<T> InPort<T>
87where
88    T: SimObject,
89{
90    #[must_use]
91    pub fn new(engine: &Engine, clock: &Clock, parent: &Rc<Entity>, name: &str) -> Self {
92        Self::new_with_renames(engine, clock, parent, name, None)
93    }
94
95    #[must_use]
96    pub fn new_with_renames(
97        engine: &Engine,
98        clock: &Clock,
99        parent: &Rc<Entity>,
100        name: &str,
101        aka: Option<&Aka>,
102    ) -> Self {
103        let entity = Rc::new(Entity::new_with_renames(parent, name, aka));
104        let monitor_window_size = entity.tracker.monitoring_window_size_for(entity.id);
105        Self {
106            entity: entity.clone(),
107            state: Rc::new(PortState::new(engine, clock, entity, monitor_window_size)),
108            connected: RefCell::new(false),
109        }
110    }
111
112    pub fn state(&self) -> PortStateResult<T> {
113        if *self.connected.borrow() {
114            return sim_error!("{self} already connected");
115        }
116
117        *self.connected.borrow_mut() = true;
118        Ok(self.state.clone())
119    }
120
121    #[must_use]
122    pub fn has_value(&self) -> bool {
123        self.state.value.borrow().is_some()
124    }
125
126    #[must_use = "Futures do nothing unless you `.await` or otherwise use them"]
127    pub fn get(&mut self) -> PortGetResult<T> {
128        if !*self.connected.borrow() {
129            return sim_error!("{self} not connected");
130        }
131
132        Ok(PortGet {
133            state: self.state.clone(),
134            done: false,
135        })
136    }
137
138    /// Must be matched with a `finish_get` to allow the OutPort to continue.
139    #[must_use = "Futures do nothing unless you `.await` or otherwise use them"]
140    pub fn start_get(&mut self) -> PortStartGetResult<T> {
141        if !*self.connected.borrow() {
142            return sim_error!("{self} not connected");
143        }
144
145        Ok(PortStartGet {
146            state: self.state.clone(),
147            done: false,
148        })
149    }
150
151    /// Must be matched with a `start_get ` to consume the value.
152    pub fn finish_get(&mut self) {
153        *self.state.put_released.borrow_mut() = true;
154        if let Some(waker) = self.state.waiting_put.borrow_mut().take() {
155            waker.wake();
156        }
157    }
158}
159
160pub struct OutPort<T>
161where
162    T: SimObject,
163{
164    entity: Rc<Entity>,
165    state: Option<Rc<PortState<T>>>,
166}
167
168impl<T> GetEntity for OutPort<T>
169where
170    T: SimObject,
171{
172    fn entity(&self) -> &Rc<Entity> {
173        &self.entity
174    }
175}
176
177impl<T> fmt::Display for OutPort<T>
178where
179    T: SimObject,
180{
181    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
182        self.entity.fmt(f)
183    }
184}
185
186impl<T> OutPort<T>
187where
188    T: SimObject,
189{
190    #[must_use]
191    pub fn new(parent: &Rc<Entity>, name: &str) -> Self {
192        Self::new_with_renames(parent, name, None)
193    }
194
195    #[must_use]
196    pub fn new_with_renames(parent: &Rc<Entity>, name: &str, aka: Option<&Aka>) -> Self {
197        let entity = Rc::new(Entity::new_with_renames(parent, name, aka));
198        Self {
199            entity,
200            state: None,
201        }
202    }
203
204    pub fn connect(&mut self, port_state: PortStateResult<T>) -> SimResult {
205        let port_state = port_state?;
206
207        connect!(self.entity ; port_state.in_port_entity);
208        match self.state {
209            Some(_) => {
210                return sim_error!("{self} already connected");
211            }
212            None => {
213                self.state = Some(port_state);
214            }
215        }
216        Ok(())
217    }
218
219    #[must_use = "Futures do nothing unless you `.await` or otherwise use them"]
220    pub fn put(&mut self, value: T) -> PortPutResult<T> {
221        let state = match self.state.as_ref() {
222            Some(s) => s.clone(),
223            None => return sim_error!("{self} not connected"),
224        };
225        Ok(PortPut {
226            state,
227            value: Some(value),
228            done: false,
229        })
230    }
231
232    #[must_use = "Futures do nothing unless you `.await` or otherwise use them"]
233    pub fn try_put(&mut self) -> PortTryPutResult<T> {
234        let state = match self.state.as_ref() {
235            Some(s) => s.clone(),
236            None => return sim_error!("{self} not connected"),
237        };
238        Ok(PortTryPut { state, done: false })
239    }
240}
241
242pub struct PortPut<T>
243where
244    T: SimObject,
245{
246    state: Rc<PortState<T>>,
247    value: Option<T>,
248    done: bool,
249}
250
251impl<T> Future for PortPut<T>
252where
253    T: SimObject,
254{
255    type Output = ();
256
257    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
258        match self.value.take() {
259            Some(value) => {
260                // The state is designed to be shared between one put/get pair so it should
261                // not be possible for the value in the state to be set at this point.
262                assert!(self.state.value.borrow().is_none());
263
264                *self.state.value.borrow_mut() = Some(value);
265                *self.state.put_released.borrow_mut() = false;
266                if let Some(waker) = self.state.waiting_get.borrow_mut().take() {
267                    waker.wake();
268                }
269                *self.state.waiting_put.borrow_mut() = Some(cx.waker().clone());
270                Poll::Pending
271            }
272            None => {
273                if *self.state.put_released.borrow() {
274                    // Getter has consumed the value and released the putter.
275                    self.done = true;
276                    Poll::Ready(())
277                } else {
278                    // Stay pending as the task was woken before the getter has removed
279                    // the value and released the putter.
280                    *self.state.waiting_put.borrow_mut() = Some(cx.waker().clone());
281                    Poll::Pending
282                }
283            }
284        }
285    }
286}
287
288impl<T> FusedFuture for PortPut<T>
289where
290    T: SimObject,
291{
292    fn is_terminated(&self) -> bool {
293        self.done
294    }
295}
296
297pub struct PortTryPut<T>
298where
299    T: SimObject,
300{
301    state: Rc<PortState<T>>,
302    done: bool,
303}
304
305impl<T> Future for PortTryPut<T>
306where
307    T: SimObject,
308{
309    type Output = ();
310
311    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
312        if self.state.waiting_get.borrow().is_some() {
313            self.done = true;
314            Poll::Ready(())
315        } else {
316            *self.state.waiting_put.borrow_mut() = Some(cx.waker().clone());
317            Poll::Pending
318        }
319    }
320}
321
322impl<T> FusedFuture for PortTryPut<T>
323where
324    T: SimObject,
325{
326    fn is_terminated(&self) -> bool {
327        self.done
328    }
329}
330
331pub struct PortGet<T>
332where
333    T: SimObject,
334{
335    state: Rc<PortState<T>>,
336    done: bool,
337}
338
339impl<T> Future for PortGet<T>
340where
341    T: SimObject,
342{
343    type Output = T;
344
345    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
346        let value = self.state.value.borrow_mut().take();
347        if let Some(value) = value {
348            self.done = true;
349            self.state.waiting_get.borrow_mut().take();
350            *self.state.put_released.borrow_mut() = true;
351
352            // Track the object through the port monitor if there is one
353            if let Some(monitor) = self.state.monitor.as_ref() {
354                monitor.sample(&value);
355            }
356
357            if let Some(waker) = self.state.waiting_put.borrow_mut().take() {
358                waker.wake();
359            }
360            Poll::Ready(value)
361        } else {
362            if let Some(waker) = self.state.waiting_put.borrow_mut().take() {
363                waker.wake();
364            }
365
366            *self.state.waiting_get.borrow_mut() = Some(cx.waker().clone());
367            Poll::Pending
368        }
369    }
370}
371
372impl<T> FusedFuture for PortGet<T>
373where
374    T: SimObject,
375{
376    fn is_terminated(&self) -> bool {
377        self.done
378    }
379}
380
381pub struct PortStartGet<T>
382where
383    T: SimObject,
384{
385    state: Rc<PortState<T>>,
386    done: bool,
387}
388
389impl<T> Future for PortStartGet<T>
390where
391    T: SimObject,
392{
393    type Output = T;
394
395    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
396        let value = self.state.value.borrow_mut().take();
397        if let Some(value) = value {
398            self.done = true;
399            self.state.waiting_get.borrow_mut().take();
400
401            // Track the object through the port monitor if there is one
402            if let Some(monitor) = self.state.monitor.as_ref() {
403                monitor.sample(&value);
404            }
405
406            Poll::Ready(value)
407        } else {
408            *self.state.waiting_get.borrow_mut() = Some(cx.waker().clone());
409            Poll::Pending
410        }
411    }
412}
413
414impl<T> FusedFuture for PortStartGet<T>
415where
416    T: SimObject,
417{
418    fn is_terminated(&self) -> bool {
419        self.done
420    }
421}
422
423#[cfg(test)]
424mod tests {
425    use std::sync::Arc;
426    use std::sync::atomic::{AtomicUsize, Ordering};
427    use std::task::{Wake, Waker};
428
429    use futures::future::FusedFuture;
430    use futures::task::noop_waker;
431    use gwr_track::Tracker;
432    use gwr_track::entity::Entity;
433    use gwr_track::tracker::dev_null_tracker;
434
435    use super::*;
436    use crate::traits::TotalBytes;
437
438    struct TestContext {
439        // Just kept to ensure it isn't dropped
440        _tracker: Tracker,
441        engine: Engine,
442        clock: Clock,
443    }
444
445    fn test_context() -> TestContext {
446        let tracker = dev_null_tracker();
447        let mut engine = Engine::new(&tracker);
448        let clock = engine.default_clock();
449
450        TestContext {
451            _tracker: tracker,
452            engine,
453            clock,
454        }
455    }
456
457    fn test_state<T: SimObject>() -> Rc<PortState<T>> {
458        let context = test_context();
459        let entity = Rc::new(Entity::new(context.engine.top(), "rx"));
460
461        Rc::new(PortState::new(
462            &context.engine,
463            &context.clock,
464            entity,
465            None,
466        ))
467    }
468
469    fn monitored_test_state<T: SimObject>() -> Rc<PortState<T>> {
470        let context = test_context();
471        let entity = Rc::new(Entity::new(context.engine.top(), "rx"));
472
473        Rc::new(PortState::new(
474            &context.engine,
475            &context.clock,
476            entity,
477            Some(1),
478        ))
479    }
480
481    struct WakeCounter {
482        wakes_count: Arc<AtomicUsize>,
483    }
484
485    impl Wake for WakeCounter {
486        fn wake(self: Arc<Self>) {
487            self.wakes_count.fetch_add(1, Ordering::SeqCst);
488        }
489
490        fn wake_by_ref(self: &Arc<Self>) {
491            self.wakes_count.fetch_add(1, Ordering::SeqCst);
492        }
493    }
494
495    fn counting_waker() -> (Arc<AtomicUsize>, Waker) {
496        let wakes_count = Arc::new(AtomicUsize::new(0));
497        let waker = Waker::from(Arc::new(WakeCounter {
498            wakes_count: wakes_count.clone(),
499        }));
500
501        (wakes_count, waker)
502    }
503
504    #[test]
505    fn wake_counter_counts_wake_and_wake_by_ref() {
506        let (wakes_count, waker) = counting_waker();
507
508        waker.wake_by_ref();
509        assert_eq!(wakes_count.load(Ordering::SeqCst), 1);
510
511        waker.wake();
512        assert_eq!(wakes_count.load(Ordering::SeqCst), 2);
513    }
514
515    #[test]
516    fn in_port_state_can_only_connect_once() {
517        let context = test_context();
518        let in_port =
519            InPort::<i32>::new(&context.engine, &context.clock, context.engine.top(), "rx");
520
521        assert!(in_port.state().is_ok());
522
523        let err = in_port
524            .state()
525            .err()
526            .expect("second state call should fail");
527        assert!(format!("{err}").contains("already connected"));
528    }
529
530    #[test]
531    fn out_port_connect_can_only_connect_once() {
532        let context = test_context();
533        let mut out_port = OutPort::<i32>::new(context.engine.top(), "tx");
534        let first_in_port =
535            InPort::new(&context.engine, &context.clock, context.engine.top(), "rx1");
536        let second_in_port =
537            InPort::new(&context.engine, &context.clock, context.engine.top(), "rx2");
538
539        out_port.connect(first_in_port.state()).unwrap();
540
541        let err = out_port.connect(second_in_port.state()).unwrap_err();
542        assert!(format!("{err}").contains("already connected"));
543    }
544
545    #[test]
546    fn out_port_entity_returns_port_entity() {
547        let context = test_context();
548        let out_port = OutPort::<i32>::new(context.engine.top(), "tx");
549
550        assert!(Rc::ptr_eq(out_port.entity(), &out_port.entity));
551    }
552
553    #[test]
554    fn start_get_requires_connection_and_finish_get_wakes_putter() {
555        let context = test_context();
556        let mut in_port =
557            InPort::<i32>::new(&context.engine, &context.clock, context.engine.top(), "rx");
558
559        assert!(in_port.start_get().is_err());
560        assert!(in_port.state().is_ok());
561        assert!(in_port.start_get().is_ok());
562
563        let waker = noop_waker();
564        *in_port.state.waiting_put.borrow_mut() = Some(waker);
565        in_port.finish_get();
566
567        assert!(in_port.state.waiting_put.borrow().is_none());
568    }
569
570    #[test]
571    fn finish_get_without_waiting_putter_is_a_noop() {
572        let context = test_context();
573        let mut in_port =
574            InPort::<i32>::new(&context.engine, &context.clock, context.engine.top(), "rx");
575
576        in_port.finish_get();
577
578        assert!(in_port.state.waiting_put.borrow().is_none());
579    }
580
581    #[test]
582    fn port_put_waits_until_value_is_consumed_before_terminating() {
583        let state = test_state::<i32>();
584        let put = PortPut {
585            state: state.clone(),
586            value: Some(123),
587            done: false,
588        };
589        let mut put = Box::pin(put);
590        let waker = noop_waker();
591        let mut cx = Context::from_waker(&waker);
592
593        assert_eq!(put.as_mut().poll(&mut cx), Poll::Pending);
594        assert!(!put.is_terminated());
595        assert_eq!(*state.value.borrow(), Some(123));
596        assert!(state.waiting_put.borrow().is_some());
597
598        assert_eq!(put.as_mut().poll(&mut cx), Poll::Pending);
599        assert!(!put.is_terminated());
600
601        assert_eq!(state.value.borrow_mut().take(), Some(123));
602        *state.put_released.borrow_mut() = true;
603
604        assert_eq!(put.as_mut().poll(&mut cx), Poll::Ready(()));
605        assert!(put.is_terminated());
606    }
607
608    #[test]
609    fn port_put_waits_for_start_get_to_finish_before_terminating() {
610        let state = test_state::<i32>();
611        let put = PortPut {
612            state: state.clone(),
613            value: Some(123),
614            done: false,
615        };
616        let mut put = Box::pin(put);
617        let start_get = PortStartGet {
618            state: state.clone(),
619            done: false,
620        };
621        let mut start_get = Box::pin(start_get);
622        let waker = noop_waker();
623        let mut cx = Context::from_waker(&waker);
624
625        assert_eq!(put.as_mut().poll(&mut cx), Poll::Pending);
626        assert_eq!(start_get.as_mut().poll(&mut cx), Poll::Ready(123));
627        assert!(state.value.borrow().is_none());
628
629        assert_eq!(put.as_mut().poll(&mut cx), Poll::Pending);
630        assert!(!put.is_terminated());
631
632        *state.put_released.borrow_mut() = true;
633
634        assert_eq!(put.as_mut().poll(&mut cx), Poll::Ready(()));
635        assert!(put.is_terminated());
636    }
637
638    #[test]
639    fn port_try_put_waits_for_getter_then_completes() {
640        let state = test_state::<i32>();
641        let try_put = PortTryPut {
642            state: state.clone(),
643            done: false,
644        };
645        let mut try_put = Box::pin(try_put);
646        let waker = noop_waker();
647        let mut cx = Context::from_waker(&waker);
648
649        assert_eq!(try_put.as_mut().poll(&mut cx), Poll::Pending);
650        assert!(!try_put.is_terminated());
651        assert!(state.waiting_put.borrow().is_some());
652
653        *state.waiting_get.borrow_mut() = Some(noop_waker());
654
655        assert_eq!(try_put.as_mut().poll(&mut cx), Poll::Ready(()));
656        assert!(try_put.is_terminated());
657    }
658
659    #[test]
660    fn connected_out_port_creates_try_put_future() {
661        let context = test_context();
662        let mut out_port = OutPort::<i32>::new(context.engine.top(), "tx");
663        let in_port = InPort::new(&context.engine, &context.clock, context.engine.top(), "rx");
664
665        out_port.connect(in_port.state()).unwrap();
666
667        assert!(out_port.try_put().is_ok());
668    }
669
670    #[test]
671    fn port_get_waits_then_returns_value_and_reports_termination() {
672        let state = test_state::<i32>();
673        let get = PortGet {
674            state: state.clone(),
675            done: false,
676        };
677        let mut get = Box::pin(get);
678        let waker = noop_waker();
679        let mut cx = Context::from_waker(&waker);
680
681        assert_eq!(get.as_mut().poll(&mut cx), Poll::Pending);
682        assert!(!get.is_terminated());
683        assert!(state.waiting_get.borrow().is_some());
684
685        *state.value.borrow_mut() = Some(456);
686        *state.waiting_put.borrow_mut() = Some(noop_waker());
687
688        assert_eq!(get.as_mut().poll(&mut cx), Poll::Ready(456));
689        assert!(get.is_terminated());
690        assert!(state.waiting_put.borrow().is_none());
691    }
692
693    #[test]
694    fn port_get_pending_wakes_waiting_putter() {
695        let state = test_state::<i32>();
696        let get = PortGet {
697            state: state.clone(),
698            done: false,
699        };
700        let mut get = Box::pin(get);
701        let waker = noop_waker();
702        let mut cx = Context::from_waker(&waker);
703        *state.waiting_put.borrow_mut() = Some(noop_waker());
704
705        assert_eq!(get.as_mut().poll(&mut cx), Poll::Pending);
706
707        assert!(state.waiting_put.borrow().is_none());
708        assert!(state.waiting_get.borrow().is_some());
709    }
710
711    #[test]
712    fn port_get_samples_monitored_values() {
713        let state = monitored_test_state::<i32>();
714        let monitor = state
715            .monitor
716            .as_ref()
717            .expect("monitored state should create a monitor");
718        let get = PortGet {
719            state: state.clone(),
720            done: false,
721        };
722        let mut get = Box::pin(get);
723        let waker = noop_waker();
724        let mut cx = Context::from_waker(&waker);
725        *state.value.borrow_mut() = Some(456);
726
727        assert_eq!(get.as_mut().poll(&mut cx), Poll::Ready(456));
728        assert_eq!(monitor.bytes_in_window(), 456_i32.total_bytes());
729    }
730
731    #[test]
732    fn port_start_get_waits_then_returns_value_without_finishing_put() {
733        let state = test_state::<i32>();
734        let start_get = PortStartGet {
735            state: state.clone(),
736            done: false,
737        };
738        let mut start_get = Box::pin(start_get);
739        let waker = noop_waker();
740        let mut cx = Context::from_waker(&waker);
741
742        assert_eq!(start_get.as_mut().poll(&mut cx), Poll::Pending);
743        assert!(!start_get.is_terminated());
744        assert!(state.waiting_get.borrow().is_some());
745
746        let (waiting_put_wakes, waiting_put_waker) = counting_waker();
747        *state.waiting_put.borrow_mut() = Some(waiting_put_waker.clone());
748        *state.value.borrow_mut() = Some(789);
749
750        assert_eq!(start_get.as_mut().poll(&mut cx), Poll::Ready(789));
751        assert!(start_get.is_terminated());
752        assert!(state.waiting_get.borrow().is_none());
753        assert_eq!(waiting_put_wakes.load(Ordering::SeqCst), 0);
754    }
755
756    #[test]
757    fn port_start_get_samples_monitored_values() {
758        let state = monitored_test_state::<i32>();
759        let monitor = state
760            .monitor
761            .as_ref()
762            .expect("monitored state should create a monitor");
763        let start_get = PortStartGet {
764            state: state.clone(),
765            done: false,
766        };
767        let mut start_get = Box::pin(start_get);
768        let waker = noop_waker();
769        let mut cx = Context::from_waker(&waker);
770        *state.value.borrow_mut() = Some(789);
771
772        assert_eq!(start_get.as_mut().poll(&mut cx), Poll::Ready(789));
773        assert_eq!(monitor.bytes_in_window(), 789_i32.total_bytes());
774    }
775}