Skip to main content

behest_runtime/
component_factory.rs

1//! Adapters that wrap existing runtime traits as [`Component`]s.
2//!
3//! `behest` has many long-lived trait abstractions
4//! (`ChatProvider`, `SessionStore`, etc.) that predate the
5//! [`Component`] trait. To make them composable with the
6//! [`ComponentRegistry`](super::registry::ComponentRegistry) without
7//! forcing every existing implementation to provide a `Component::init`
8//! method, this module provides ready-made wrapper factories that
9//! take an already-constructed `Arc<T>` and expose it as a registered
10//! component.
11//!
12//! These wrappers are the canonical M3 deliverable: the existing trait
13//! surface stays unchanged, while the runtime becomes composable on
14//! top.
15//!
16//! # Example
17//!
18//! ```no_run
19//! use std::sync::Arc;
20//! use behest_runtime::component_factory::ChatProviderComponent;
21//! use behest_runtime::registry::ComponentRegistry;
22//!
23//! # async fn build(registry: &ComponentRegistry) {
24//! // Caller-supplied provider, pre-constructed:
25//! let provider: Arc<dyn behest_provider::ChatProvider> = todo!();
26//! let factory = ChatProviderComponent::new("primary", provider);
27//! // ... register with ComponentRegistry via register_factory().
28//! # }
29//! ```
30
31#![allow(clippy::pedantic)]
32
33use std::any::Any;
34use std::fmt;
35use std::sync::Arc;
36
37use async_trait::async_trait;
38use futures_util::future::BoxFuture;
39use schemars::JsonSchema;
40use serde::{Deserialize, Serialize};
41
42use super::component::{AnyComponent, AnyComponentError, Component, ComponentContext};
43use super::extension::ExtensionError;
44use super::registry::{ComponentDescriptor, ComponentFactory, RegistryError};
45use behest_provider::{ChatProvider, EmbeddingProvider};
46
47/// Configuration for a chat provider component wrapper. Currently
48/// unused; the wrapper takes a pre-constructed provider. Exists to
49/// satisfy the [`Component::Config`] bound so the wrapper can be
50/// registered in a [`ComponentRegistry`](super::registry::ComponentRegistry).
51#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema)]
52pub struct EmptyConfig;
53
54/// Error type for wrapper components: lifecycle errors only, since
55/// the wrapped provider was constructed externally.
56#[derive(Debug, thiserror::Error)]
57#[error("provider wrapper error: {0}")]
58pub struct WrapperError(pub String);
59
60impl WrapperError {
61    /// Construct a wrapper error from a displayable value.
62    #[must_use]
63    pub fn new(msg: impl Into<String>) -> Self {
64        Self(msg.into())
65    }
66}
67
68/// A [`Component`] that wraps an existing [`ChatProvider`].
69pub struct ChatProviderComponent {
70    name: String,
71    provider: Arc<dyn ChatProvider>,
72}
73
74impl ChatProviderComponent {
75    /// Construct a wrapper around a pre-built chat provider.
76    #[must_use]
77    pub fn new(name: impl Into<String>, provider: Arc<dyn ChatProvider>) -> Self {
78        Self {
79            name: name.into(),
80            provider,
81        }
82    }
83
84    /// Borrow the inner provider.
85    #[must_use]
86    pub fn provider(&self) -> &Arc<dyn ChatProvider> {
87        &self.provider
88    }
89}
90
91impl fmt::Debug for ChatProviderComponent {
92    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
93        f.debug_struct("ChatProviderComponent")
94            .field("name", &self.name)
95            .field("provider_id", &self.provider.id())
96            .finish()
97    }
98}
99
100#[async_trait]
101impl Component for ChatProviderComponent {
102    const NAME: &'static str = "ChatProviderComponent";
103    type Config = EmptyConfig;
104    type Error = WrapperError;
105
106    async fn init(_cfg: &Self::Config, _ctx: &ComponentContext) -> Result<Self, Self::Error> {
107        Err(WrapperError::new(
108            "wrapper components must be constructed via new()",
109        ))
110    }
111
112    async fn start(&self) -> Result<(), Self::Error> {
113        Ok(())
114    }
115
116    async fn stop(&self) -> Result<(), Self::Error> {
117        Ok(())
118    }
119
120    async fn health(&self) -> behest_core::health::HealthStatus {
121        behest_core::health::HealthStatus::healthy()
122    }
123}
124
125/// A [`Component`] that wraps an existing [`EmbeddingProvider`].
126pub struct EmbeddingProviderComponent {
127    name: String,
128    provider: Arc<dyn EmbeddingProvider>,
129}
130
131impl EmbeddingProviderComponent {
132    /// Construct a wrapper around a pre-built embedding provider.
133    #[must_use]
134    pub fn new(name: impl Into<String>, provider: Arc<dyn EmbeddingProvider>) -> Self {
135        Self {
136            name: name.into(),
137            provider,
138        }
139    }
140
141    /// Borrow the inner provider.
142    #[must_use]
143    pub fn provider(&self) -> &Arc<dyn EmbeddingProvider> {
144        &self.provider
145    }
146}
147
148impl fmt::Debug for EmbeddingProviderComponent {
149    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
150        f.debug_struct("EmbeddingProviderComponent")
151            .field("name", &self.name)
152            .field("provider_id", &self.provider.id())
153            .finish()
154    }
155}
156
157#[async_trait]
158impl Component for EmbeddingProviderComponent {
159    const NAME: &'static str = "EmbeddingProviderComponent";
160    type Config = EmptyConfig;
161    type Error = WrapperError;
162
163    async fn init(_cfg: &Self::Config, _ctx: &ComponentContext) -> Result<Self, Self::Error> {
164        Err(WrapperError::new(
165            "wrapper components must be constructed via new()",
166        ))
167    }
168
169    async fn start(&self) -> Result<(), Self::Error> {
170        Ok(())
171    }
172
173    async fn stop(&self) -> Result<(), Self::Error> {
174        Ok(())
175    }
176}
177
178/// [`AnyComponent`] adapter for [`ChatProviderComponent`].
179pub struct ChatProviderAny {
180    inner: Arc<ChatProviderComponent>,
181}
182
183impl ChatProviderAny {
184    /// Wrap a [`ChatProviderComponent`] as a type-erased [`AnyComponent`].
185    #[must_use]
186    pub fn new(inner: Arc<ChatProviderComponent>) -> Self {
187        Self { inner }
188    }
189}
190
191#[async_trait]
192impl AnyComponent for ChatProviderAny {
193    fn name(&self) -> &'static str {
194        ChatProviderComponent::NAME
195    }
196
197    fn as_any_arc(&self) -> Arc<dyn Any + Send + Sync> {
198        self.inner.clone()
199    }
200
201    fn start(&self) -> BoxFuture<'_, Result<(), AnyComponentError>> {
202        let name = self.inner.name.clone();
203        let inner = self.inner.clone();
204        Box::pin(async move {
205            inner
206                .start()
207                .await
208                .map_err(|e| AnyComponentError::Component {
209                    name,
210                    message: e.to_string(),
211                })
212        })
213    }
214
215    fn stop(&self) -> BoxFuture<'_, Result<(), AnyComponentError>> {
216        let name = self.inner.name.clone();
217        let inner = self.inner.clone();
218        Box::pin(async move {
219            inner
220                .stop()
221                .await
222                .map_err(|e| AnyComponentError::Component {
223                    name,
224                    message: e.to_string(),
225                })
226        })
227    }
228
229    fn health(&self) -> BoxFuture<'_, behest_core::health::HealthStatus> {
230        let inner = self.inner.clone();
231        Box::pin(async move { inner.health().await })
232    }
233
234    fn pre_replace(&self) -> BoxFuture<'_, Result<(), AnyComponentError>> {
235        let name = self.inner.name.clone();
236        let inner = self.inner.clone();
237        Box::pin(async move {
238            inner
239                .pre_replace_hook()
240                .await
241                .map_err(|e| AnyComponentError::Component {
242                    name,
243                    message: e.to_string(),
244                })
245        })
246    }
247
248    fn post_replace(&self) -> BoxFuture<'_, Result<(), AnyComponentError>> {
249        let name = self.inner.name.clone();
250        let inner = self.inner.clone();
251        Box::pin(async move {
252            inner
253                .post_replace_hook()
254                .await
255                .map_err(|e| AnyComponentError::Component {
256                    name,
257                    message: e.to_string(),
258                })
259        })
260    }
261}
262
263/// [`AnyComponent`] adapter for [`EmbeddingProviderComponent`].
264pub struct EmbeddingProviderAny {
265    inner: Arc<EmbeddingProviderComponent>,
266}
267
268impl EmbeddingProviderAny {
269    /// Wrap an [`EmbeddingProviderComponent`] as a type-erased [`AnyComponent`].
270    #[must_use]
271    pub fn new(inner: Arc<EmbeddingProviderComponent>) -> Self {
272        Self { inner }
273    }
274}
275
276#[async_trait]
277impl AnyComponent for EmbeddingProviderAny {
278    fn name(&self) -> &'static str {
279        EmbeddingProviderComponent::NAME
280    }
281
282    fn as_any_arc(&self) -> Arc<dyn Any + Send + Sync> {
283        self.inner.clone()
284    }
285
286    fn start(&self) -> BoxFuture<'_, Result<(), AnyComponentError>> {
287        let name = self.inner.name.clone();
288        let inner = self.inner.clone();
289        Box::pin(async move {
290            inner
291                .start()
292                .await
293                .map_err(|e| AnyComponentError::Component {
294                    name,
295                    message: e.to_string(),
296                })
297        })
298    }
299
300    fn stop(&self) -> BoxFuture<'_, Result<(), AnyComponentError>> {
301        let name = self.inner.name.clone();
302        let inner = self.inner.clone();
303        Box::pin(async move {
304            inner
305                .stop()
306                .await
307                .map_err(|e| AnyComponentError::Component {
308                    name,
309                    message: e.to_string(),
310                })
311        })
312    }
313
314    fn health(&self) -> BoxFuture<'_, behest_core::health::HealthStatus> {
315        let inner = self.inner.clone();
316        Box::pin(async move { inner.health().await })
317    }
318
319    fn pre_replace(&self) -> BoxFuture<'_, Result<(), AnyComponentError>> {
320        let name = self.inner.name.clone();
321        let inner = self.inner.clone();
322        Box::pin(async move {
323            inner
324                .pre_replace_hook()
325                .await
326                .map_err(|e| AnyComponentError::Component {
327                    name,
328                    message: e.to_string(),
329                })
330        })
331    }
332
333    fn post_replace(&self) -> BoxFuture<'_, Result<(), AnyComponentError>> {
334        let name = self.inner.name.clone();
335        let inner = self.inner.clone();
336        Box::pin(async move {
337            inner
338                .post_replace_hook()
339                .await
340                .map_err(|e| AnyComponentError::Component {
341                    name,
342                    message: e.to_string(),
343                })
344        })
345    }
346}
347
348/// Factory for [`ChatProviderComponent`] that takes a pre-built
349/// `Arc<dyn ChatProvider>` and registers it under the given name.
350pub struct ChatProviderFactory {
351    descriptor: ComponentDescriptor,
352    provider: Arc<dyn ChatProvider>,
353}
354
355impl ChatProviderFactory {
356    /// Construct a factory. The `name` is the user-assigned instance
357    /// name in the registry.
358    #[must_use]
359    pub fn new(name: impl Into<String>, provider: Arc<dyn ChatProvider>) -> Self {
360        let name = name.into();
361        let descriptor = ComponentDescriptor {
362            name: name.clone(),
363            depends_on: Vec::new(),
364            config: serde_json::json!({}),
365        };
366        Self {
367            descriptor,
368            provider,
369        }
370    }
371}
372
373#[async_trait]
374impl ComponentFactory for ChatProviderFactory {
375    fn name(&self) -> &str {
376        &self.descriptor.name
377    }
378
379    fn kind(&self) -> &'static str {
380        ChatProviderComponent::NAME
381    }
382
383    fn depends_on(&self) -> Vec<String> {
384        self.descriptor.depends_on.clone()
385    }
386
387    async fn build(
388        self: Box<Self>,
389        _config: serde_json::Value,
390        _ctx: &ComponentContext,
391    ) -> Result<Box<dyn AnyComponent>, RegistryError> {
392        let inner = Arc::new(ChatProviderComponent::new(
393            self.descriptor.name.clone(),
394            self.provider,
395        ));
396        Ok(Box::new(ChatProviderAny::new(inner)))
397    }
398}
399
400/// Factory for [`EmbeddingProviderComponent`].
401pub struct EmbeddingProviderFactory {
402    descriptor: ComponentDescriptor,
403    provider: Arc<dyn EmbeddingProvider>,
404}
405
406impl EmbeddingProviderFactory {
407    /// Construct a factory.
408    #[must_use]
409    pub fn new(name: impl Into<String>, provider: Arc<dyn EmbeddingProvider>) -> Self {
410        let name = name.into();
411        let descriptor = ComponentDescriptor {
412            name: name.clone(),
413            depends_on: Vec::new(),
414            config: serde_json::json!({}),
415        };
416        Self {
417            descriptor,
418            provider,
419        }
420    }
421}
422
423#[async_trait]
424impl ComponentFactory for EmbeddingProviderFactory {
425    fn name(&self) -> &str {
426        &self.descriptor.name
427    }
428
429    fn kind(&self) -> &'static str {
430        EmbeddingProviderComponent::NAME
431    }
432
433    fn depends_on(&self) -> Vec<String> {
434        self.descriptor.depends_on.clone()
435    }
436
437    async fn build(
438        self: Box<Self>,
439        _config: serde_json::Value,
440        _ctx: &ComponentContext,
441    ) -> Result<Box<dyn AnyComponent>, RegistryError> {
442        let inner = Arc::new(EmbeddingProviderComponent::new(
443            self.descriptor.name.clone(),
444            self.provider,
445        ));
446        Ok(Box::new(EmbeddingProviderAny::new(inner)))
447    }
448}
449
450/// Convenience: re-export a name-keyed error so callers don't have to
451/// import [`ExtensionError`] separately.
452pub type FactoryError = ExtensionError;
453
454#[cfg(test)]
455mod tests {
456    use super::*;
457    use crate::lifecycle::ShutdownToken;
458    use crate::registry::ComponentRegistry;
459    use async_trait::async_trait;
460    use behest_core::error::ProviderError;
461    use behest_core::health::HealthStatus;
462    use behest_provider::{
463        ChatRequest, ChatResponse, ChatStream, EmbeddingRequest, EmbeddingResponse, FinishReason,
464        Message, ProviderCapabilities, ProviderId, ProviderResult, TokenUsage,
465    };
466
467    struct StubChat;
468    #[async_trait]
469    impl ChatProvider for StubChat {
470        fn id(&self) -> ProviderId {
471            ProviderId::new("stub")
472        }
473        fn capabilities(&self) -> ProviderCapabilities {
474            ProviderCapabilities::default()
475        }
476        async fn complete(&self, _r: ChatRequest) -> ProviderResult<ChatResponse> {
477            Ok(ChatResponse {
478                provider: self.id(),
479                model: behest_provider::ModelName::new("stub-model"),
480                message: Message::user_text(""),
481                finish_reason: FinishReason::Stop,
482                usage: Some(TokenUsage::new(0, 0)),
483                raw: None,
484            })
485        }
486        async fn stream(&self, _r: ChatRequest) -> ProviderResult<ChatStream> {
487            Err(ProviderError::Unsupported {
488                provider: self.id(),
489                feature: "stream".into(),
490            })
491        }
492    }
493
494    struct StubEmbedding;
495    #[async_trait]
496    impl EmbeddingProvider for StubEmbedding {
497        fn id(&self) -> ProviderId {
498            ProviderId::new("stub-emb")
499        }
500        fn capabilities(&self) -> ProviderCapabilities {
501            ProviderCapabilities::default()
502        }
503        async fn embed(&self, _r: EmbeddingRequest) -> ProviderResult<EmbeddingResponse> {
504            Ok(EmbeddingResponse {
505                provider: self.id(),
506                model: behest_provider::ModelName::new("stub-emb-model"),
507                embeddings: Vec::new(),
508                usage: None,
509                raw: None,
510            })
511        }
512    }
513
514    #[tokio::test]
515    async fn chat_provider_wrapper_registers_with_registry() {
516        let registry = ComponentRegistry::new();
517        let factory = ChatProviderFactory::new("primary", Arc::new(StubChat));
518        registry
519            .register_factory(
520                ComponentDescriptor {
521                    name: "primary".into(),
522                    depends_on: Vec::new(),
523                    config: serde_json::json!({}),
524                },
525                Box::new(factory),
526            )
527            .unwrap_or_else(|e| panic!("{e}"));
528        registry.init_all().await.unwrap_or_else(|e| panic!("{e}"));
529        registry.start_all().await.unwrap_or_else(|e| panic!("{e}"));
530        let c = registry
531            .get::<ChatProviderComponent>("primary")
532            .unwrap_or_else(|e| panic!("{e}"));
533        assert_eq!(c.name, "primary");
534        assert_eq!(c.provider.id().as_str(), "stub");
535    }
536
537    #[tokio::test]
538    async fn embedding_provider_wrapper_registers_with_registry() {
539        let registry = ComponentRegistry::new();
540        let factory = EmbeddingProviderFactory::new("primary", Arc::new(StubEmbedding));
541        registry
542            .register_factory(
543                ComponentDescriptor {
544                    name: "primary".into(),
545                    depends_on: Vec::new(),
546                    config: serde_json::json!({}),
547                },
548                Box::new(factory),
549            )
550            .unwrap_or_else(|e| panic!("{e}"));
551        registry.init_all().await.unwrap_or_else(|e| panic!("{e}"));
552        let c = registry
553            .get::<EmbeddingProviderComponent>("primary")
554            .unwrap_or_else(|e| panic!("{e}"));
555        assert_eq!(c.name, "primary");
556    }
557
558    #[tokio::test]
559    async fn chat_provider_init_returns_wrapper_error() {
560        let shutdown = ShutdownToken::new();
561        let ctx = ComponentContext::new(shutdown);
562        let cfg = EmptyConfig;
563        let result = ChatProviderComponent::init(&cfg, &ctx).await;
564        assert!(result.is_err());
565    }
566
567    #[tokio::test]
568    async fn default_lifecycle_is_noop_for_provider_wrappers() {
569        let c = ChatProviderComponent::new("x", Arc::new(StubChat));
570        let _ = c.start().await;
571        let _ = c.stop().await;
572        let h = c.health().await;
573        assert_eq!(h, HealthStatus::healthy());
574    }
575}