rig_core/test_utils/
relay.rs1use std::collections::VecDeque;
4use std::sync::Mutex;
5
6use crate::completion::{ModelRef, ProviderCapabilities};
7use crate::effect::{EffectFamily, EffectKind, FamilyDescriptor, HandlerDescriptor, family};
8use crate::error::{ErrorKind, ErrorReport};
9use crate::serve::{Dispatch, Reply, Serve};
10
11use super::{MockFrame, MockScript, MockStreamEvent};
12
13pub struct MockRelay {
18 label: ModelRef,
19 turns: Mutex<VecDeque<Vec<MockStreamEvent>>>,
20}
21
22impl MockRelay {
23 pub fn new(
25 label: impl Into<ModelRef>,
26 turns: impl IntoIterator<Item = impl IntoIterator<Item = MockStreamEvent>>,
27 ) -> Self {
28 Self {
29 label: label.into(),
30 turns: Mutex::new(
31 turns
32 .into_iter()
33 .map(|turn| turn.into_iter().collect())
34 .collect(),
35 ),
36 }
37 }
38}
39
40impl Serve for MockRelay {
41 type Family = family::Completion;
42
43 fn descriptor(&self) -> HandlerDescriptor {
44 HandlerDescriptor {
45 key: crate::effect::model_key(self.label.as_str()),
46 family: FamilyDescriptor::Completion {
47 model: self.label.clone(),
48 capabilities: ProviderCapabilities::default(),
49 },
50 layers: Vec::new(),
51 }
52 }
53
54 async fn serve(&self, kind: EffectKind, _dispatch: Dispatch) -> Reply {
55 if !matches!(kind, EffectKind::Completion { .. }) {
56 return Reply::Outcome(Err(ErrorReport::new(
57 ErrorKind::HandlerUnavailable,
58 format!(
59 "a {} handler cannot serve a `{}` effect",
60 EffectFamily::Completion,
61 kind.name()
62 ),
63 )));
64 }
65 let turn = self
66 .turns
67 .lock()
68 .unwrap_or_else(std::sync::PoisonError::into_inner)
69 .pop_front();
70 let Some(turn) = turn else {
71 return Reply::Outcome(Err(ErrorReport::new(
72 ErrorKind::Provider,
73 "mock relay has no scripted turn",
74 )));
75 };
76 let frames = turn.into_iter().map(MockFrame::Event);
77 Reply::Stream(crate::driver::relay_frames(
78 &MockScript::new(self.label.as_str()),
79 frames,
80 ))
81 }
82}