everruns_host/
composition.rs1use crate::{DisabledSessionFileSystemFactory, SessionFileSystemFactory};
26use everruns_core::{
27 Capability, CapabilityRegistry, ClassifierService, EgressService, UtilityLlmService,
28 tool_context::ToolContextExtensions,
29};
30use everruns_provider::driver_registry::DriverRegistry;
31use std::sync::{Arc, RwLock};
32
33pub struct HostComposition {
54 capability_registry: RwLock<Arc<CapabilityRegistry>>,
58 driver_registry: DriverRegistry,
59 egress_service: Arc<dyn EgressService>,
60 utility_llm_service: Arc<dyn UtilityLlmService>,
61 classifier: Arc<dyn ClassifierService>,
62 session_file_system_factory: Arc<dyn SessionFileSystemFactory>,
63 extensions: ToolContextExtensions,
64}
65
66impl HostComposition {
67 pub fn new(capability_registry: CapabilityRegistry, driver_registry: DriverRegistry) -> Self {
69 Self {
70 capability_registry: RwLock::new(Arc::new(capability_registry)),
71 driver_registry,
72 egress_service: Arc::new(everruns_core::DisabledEgressService),
73 utility_llm_service: Arc::new(everruns_core::DisabledUtilityLlmService),
74 classifier: Arc::new(everruns_core::DisabledClassifierService),
75 session_file_system_factory: Arc::new(DisabledSessionFileSystemFactory),
76 extensions: ToolContextExtensions::default(),
77 }
78 }
79
80 pub fn builder() -> HostCompositionBuilder {
82 HostCompositionBuilder::new()
83 }
84
85 pub fn capability_registry(&self) -> Arc<CapabilityRegistry> {
91 self.read_registry().clone()
92 }
93
94 pub fn register_capability(
105 &self,
106 capability: Arc<dyn Capability>,
107 ) -> Result<(), everruns_capability::CapabilityError> {
108 self.update_registry(|registry| registry.try_register_arc(capability))
109 }
110
111 pub fn register_capability_overriding(&self, capability: Arc<dyn Capability>) {
118 let _ = self.update_registry(|registry| {
119 registry.register_arc(capability);
120 Ok::<(), everruns_capability::CapabilityError>(())
121 });
122 }
123
124 pub fn is_capability_registered(&self, id: &str) -> bool {
129 self.read_registry().has(id)
130 }
131
132 fn read_registry(&self) -> std::sync::RwLockReadGuard<'_, Arc<CapabilityRegistry>> {
133 self.capability_registry
137 .read()
138 .unwrap_or_else(|poisoned| poisoned.into_inner())
139 }
140
141 fn update_registry<E>(
142 &self,
143 mutate: impl FnOnce(&mut CapabilityRegistry) -> Result<(), E>,
144 ) -> Result<(), E> {
145 let mut guard = self
146 .capability_registry
147 .write()
148 .unwrap_or_else(|poisoned| poisoned.into_inner());
149 let mut next = (**guard).clone();
150 mutate(&mut next)?;
151 *guard = Arc::new(next);
152 Ok(())
153 }
154
155 pub fn driver_registry(&self) -> &DriverRegistry {
157 &self.driver_registry
158 }
159
160 pub fn driver_registry_mut(&mut self) -> &mut DriverRegistry {
162 &mut self.driver_registry
163 }
164
165 pub fn egress_service(&self) -> Arc<dyn EgressService> {
167 self.egress_service.clone()
168 }
169
170 pub fn utility_llm_service(&self) -> Arc<dyn UtilityLlmService> {
172 self.utility_llm_service.clone()
173 }
174
175 pub fn classifier(&self) -> Arc<dyn ClassifierService> {
178 self.classifier.clone()
179 }
180
181 pub fn session_file_system_factory(&self) -> Arc<dyn SessionFileSystemFactory> {
183 self.session_file_system_factory.clone()
184 }
185
186 pub fn extension<T: std::any::Any + Send + Sync>(&self) -> Option<Arc<T>> {
188 self.extensions.get::<T>()
189 }
190}
191
192impl Clone for HostComposition {
199 fn clone(&self) -> Self {
200 Self {
201 capability_registry: RwLock::new(self.capability_registry()),
202 driver_registry: self.driver_registry.clone(),
203 egress_service: self.egress_service.clone(),
204 utility_llm_service: self.utility_llm_service.clone(),
205 classifier: self.classifier.clone(),
206 session_file_system_factory: self.session_file_system_factory.clone(),
207 extensions: self.extensions.clone(),
208 }
209 }
210}
211
212impl Default for HostComposition {
213 fn default() -> Self {
214 Self::new(CapabilityRegistry::new(), DriverRegistry::new())
215 }
216}
217
218impl std::fmt::Debug for HostComposition {
219 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
220 f.debug_struct("HostComposition")
221 .field("capabilities", &self.capability_registry())
222 .field("drivers", &self.driver_registry.registered_providers())
223 .field("egress_service", &self.egress_service.name())
224 .field("utility_llm_service", &self.utility_llm_service.name())
225 .field("classifier", &self.classifier.name())
226 .field(
227 "session_file_system_factory",
228 &self.session_file_system_factory.name(),
229 )
230 .field("extensions", &self.extensions)
231 .finish()
232 }
233}
234
235pub struct HostCompositionBuilder {
237 composition: HostComposition,
238}
239
240impl HostCompositionBuilder {
241 pub fn new() -> Self {
243 Self {
244 composition: HostComposition::default(),
245 }
246 }
247
248 pub fn capability_registry(self, registry: CapabilityRegistry) -> Self {
250 *self
251 .composition
252 .capability_registry
253 .write()
254 .unwrap_or_else(|poisoned| poisoned.into_inner()) = Arc::new(registry);
255 self
256 }
257
258 pub fn capability(self, capability: impl Capability + 'static) -> Self {
260 self.composition
261 .register_capability_overriding(Arc::new(capability));
262 self
263 }
264
265 pub fn driver_registry(mut self, registry: DriverRegistry) -> Self {
267 self.composition.driver_registry = registry;
268 self
269 }
270
271 pub fn egress_service(mut self, service: Arc<dyn EgressService>) -> Self {
273 self.composition.egress_service = service;
274 self
275 }
276
277 pub fn utility_llm_service(mut self, service: Arc<dyn UtilityLlmService>) -> Self {
279 self.composition.utility_llm_service = service;
280 self
281 }
282
283 pub fn classifier(mut self, service: Arc<dyn ClassifierService>) -> Self {
285 self.composition.classifier = service;
286 self
287 }
288
289 pub fn session_file_system_factory(
291 mut self,
292 factory: Arc<dyn SessionFileSystemFactory>,
293 ) -> Self {
294 self.composition.session_file_system_factory = factory;
295 self
296 }
297
298 pub fn extension<T: std::any::Any + Send + Sync>(mut self, value: Arc<T>) -> Self {
300 self.composition.extensions.insert(value);
301 self
302 }
303
304 pub fn build(self) -> HostComposition {
306 self.composition
307 }
308}
309
310impl Default for HostCompositionBuilder {
311 fn default() -> Self {
312 Self::new()
313 }
314}
315
316#[cfg(test)]
317mod tests {
318 use super::*;
319 use async_trait::async_trait;
320 use everruns_builtins::HumanIntentCapability;
321 use everruns_core::CapabilityStatus;
322
323 struct StubChatDriver;
325
326 #[async_trait]
327 impl everruns_provider::driver_registry::ChatDriver for StubChatDriver {
328 async fn chat_completion_stream(
329 &self,
330 _endpoint: &everruns_provider::runtime_provider::ProviderEndpoint,
331 _messages: Vec<everruns_provider::driver_registry::LlmMessage>,
332 _config: &everruns_provider::driver_registry::LlmCallConfig,
333 ) -> everruns_provider::error::Result<everruns_provider::driver_registry::LlmResponseStream>
334 {
335 Ok(Box::pin(futures::stream::empty()))
336 }
337 }
338
339 #[test]
340 fn composition_builder_registers_capabilities_and_drivers() {
341 let mut drivers = DriverRegistry::new();
342 let mut descriptor = everruns_provider::driver_registry::DriverDescriptor::chat_only(
343 everruns_provider::provider::DriverId::LlmSim,
344 |_config| {
345 Box::new(StubChatDriver) as everruns_provider::driver_registry::BoxedChatDriver
346 },
347 );
348 descriptor.display_name = "Stub".into();
349 drivers.register_descriptor_or_replace(descriptor);
350
351 let composition = HostComposition::builder()
352 .driver_registry(drivers.clone())
353 .capability(HumanIntentCapability)
354 .build();
355
356 assert!(composition.capability_registry().has("human_intent"));
357 assert!(
358 composition
359 .driver_registry()
360 .has_driver(&everruns_provider::provider::DriverId::LlmSim)
361 );
362 }
363
364 #[test]
365 fn composition_registries_stay_mutable_after_build() {
366 let composition = HostComposition::default();
367 composition.register_capability_overriding(Arc::new(HumanIntentCapability));
368
369 let info = everruns_core::CapabilityInfo::from_core(
370 composition
371 .capability_registry()
372 .get("human_intent")
373 .expect("human_intent registered")
374 .as_ref(),
375 );
376 assert_eq!(info.status, CapabilityStatus::Available);
377 }
378}