use crate::{DisabledSessionFileSystemFactory, SessionFileSystemFactory};
use everruns_core::{
Capability, CapabilityRegistry, EgressService, UtilityLlmService,
tool_context::ToolContextExtensions,
};
use everruns_provider::driver_registry::DriverRegistry;
use std::sync::{Arc, RwLock};
pub struct HostComposition {
capability_registry: RwLock<Arc<CapabilityRegistry>>,
driver_registry: DriverRegistry,
egress_service: Arc<dyn EgressService>,
utility_llm_service: Arc<dyn UtilityLlmService>,
session_file_system_factory: Arc<dyn SessionFileSystemFactory>,
extensions: ToolContextExtensions,
}
impl HostComposition {
pub fn new(capability_registry: CapabilityRegistry, driver_registry: DriverRegistry) -> Self {
Self {
capability_registry: RwLock::new(Arc::new(capability_registry)),
driver_registry,
egress_service: Arc::new(everruns_core::DisabledEgressService),
utility_llm_service: Arc::new(everruns_core::DisabledUtilityLlmService),
session_file_system_factory: Arc::new(DisabledSessionFileSystemFactory),
extensions: ToolContextExtensions::default(),
}
}
pub fn builder() -> HostCompositionBuilder {
HostCompositionBuilder::new()
}
pub fn capability_registry(&self) -> Arc<CapabilityRegistry> {
self.read_registry().clone()
}
pub fn register_capability(
&self,
capability: Arc<dyn Capability>,
) -> Result<(), everruns_capability::CapabilityError> {
self.update_registry(|registry| registry.try_register_arc(capability))
}
pub fn register_capability_overriding(&self, capability: Arc<dyn Capability>) {
let _ = self.update_registry(|registry| {
registry.register_arc(capability);
Ok::<(), everruns_capability::CapabilityError>(())
});
}
pub fn is_capability_registered(&self, id: &str) -> bool {
self.read_registry().has(id)
}
fn read_registry(&self) -> std::sync::RwLockReadGuard<'_, Arc<CapabilityRegistry>> {
self.capability_registry
.read()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
fn update_registry<E>(
&self,
mutate: impl FnOnce(&mut CapabilityRegistry) -> Result<(), E>,
) -> Result<(), E> {
let mut guard = self
.capability_registry
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let mut next = (**guard).clone();
mutate(&mut next)?;
*guard = Arc::new(next);
Ok(())
}
pub fn driver_registry(&self) -> &DriverRegistry {
&self.driver_registry
}
pub fn driver_registry_mut(&mut self) -> &mut DriverRegistry {
&mut self.driver_registry
}
pub fn egress_service(&self) -> Arc<dyn EgressService> {
self.egress_service.clone()
}
pub fn utility_llm_service(&self) -> Arc<dyn UtilityLlmService> {
self.utility_llm_service.clone()
}
pub fn session_file_system_factory(&self) -> Arc<dyn SessionFileSystemFactory> {
self.session_file_system_factory.clone()
}
pub fn extension<T: std::any::Any + Send + Sync>(&self) -> Option<Arc<T>> {
self.extensions.get::<T>()
}
}
impl Clone for HostComposition {
fn clone(&self) -> Self {
Self {
capability_registry: RwLock::new(self.capability_registry()),
driver_registry: self.driver_registry.clone(),
egress_service: self.egress_service.clone(),
utility_llm_service: self.utility_llm_service.clone(),
session_file_system_factory: self.session_file_system_factory.clone(),
extensions: self.extensions.clone(),
}
}
}
impl Default for HostComposition {
fn default() -> Self {
Self::new(CapabilityRegistry::new(), DriverRegistry::new())
}
}
impl std::fmt::Debug for HostComposition {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HostComposition")
.field("capabilities", &self.capability_registry())
.field("drivers", &self.driver_registry.registered_providers())
.field("egress_service", &self.egress_service.name())
.field("utility_llm_service", &self.utility_llm_service.name())
.field(
"session_file_system_factory",
&self.session_file_system_factory.name(),
)
.field("extensions", &self.extensions)
.finish()
}
}
pub struct HostCompositionBuilder {
composition: HostComposition,
}
impl HostCompositionBuilder {
pub fn new() -> Self {
Self {
composition: HostComposition::default(),
}
}
pub fn capability_registry(self, registry: CapabilityRegistry) -> Self {
*self
.composition
.capability_registry
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = Arc::new(registry);
self
}
pub fn capability(self, capability: impl Capability + 'static) -> Self {
self.composition
.register_capability_overriding(Arc::new(capability));
self
}
pub fn driver_registry(mut self, registry: DriverRegistry) -> Self {
self.composition.driver_registry = registry;
self
}
pub fn egress_service(mut self, service: Arc<dyn EgressService>) -> Self {
self.composition.egress_service = service;
self
}
pub fn utility_llm_service(mut self, service: Arc<dyn UtilityLlmService>) -> Self {
self.composition.utility_llm_service = service;
self
}
pub fn session_file_system_factory(
mut self,
factory: Arc<dyn SessionFileSystemFactory>,
) -> Self {
self.composition.session_file_system_factory = factory;
self
}
pub fn extension<T: std::any::Any + Send + Sync>(mut self, value: Arc<T>) -> Self {
self.composition.extensions.insert(value);
self
}
pub fn build(self) -> HostComposition {
self.composition
}
}
impl Default for HostCompositionBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use async_trait::async_trait;
use everruns_builtins::HumanIntentCapability;
use everruns_core::CapabilityStatus;
struct StubChatDriver;
#[async_trait]
impl everruns_provider::driver_registry::ChatDriver for StubChatDriver {
async fn chat_completion_stream(
&self,
_endpoint: &everruns_provider::runtime_provider::ProviderEndpoint,
_messages: Vec<everruns_provider::driver_registry::LlmMessage>,
_config: &everruns_provider::driver_registry::LlmCallConfig,
) -> everruns_provider::error::Result<everruns_provider::driver_registry::LlmResponseStream>
{
Ok(Box::pin(futures::stream::empty()))
}
}
#[test]
fn composition_builder_registers_capabilities_and_drivers() {
let mut drivers = DriverRegistry::new();
let mut descriptor = everruns_provider::driver_registry::DriverDescriptor::chat_only(
everruns_provider::provider::DriverId::LlmSim,
|_config| {
Box::new(StubChatDriver) as everruns_provider::driver_registry::BoxedChatDriver
},
);
descriptor.display_name = "Stub".into();
drivers.register_descriptor_or_replace(descriptor);
let composition = HostComposition::builder()
.driver_registry(drivers.clone())
.capability(HumanIntentCapability)
.build();
assert!(composition.capability_registry().has("human_intent"));
assert!(
composition
.driver_registry()
.has_driver(&everruns_provider::provider::DriverId::LlmSim)
);
}
#[test]
fn composition_registries_stay_mutable_after_build() {
let composition = HostComposition::default();
composition.register_capability_overriding(Arc::new(HumanIntentCapability));
let info = everruns_core::CapabilityInfo::from_core(
composition
.capability_registry()
.get("human_intent")
.expect("human_intent registered")
.as_ref(),
);
assert_eq!(info.status, CapabilityStatus::Available);
}
}