1#![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#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema)]
52pub struct EmptyConfig;
53
54#[derive(Debug, thiserror::Error)]
57#[error("provider wrapper error: {0}")]
58pub struct WrapperError(pub String);
59
60impl WrapperError {
61 #[must_use]
63 pub fn new(msg: impl Into<String>) -> Self {
64 Self(msg.into())
65 }
66}
67
68pub struct ChatProviderComponent {
70 name: String,
71 provider: Arc<dyn ChatProvider>,
72}
73
74impl ChatProviderComponent {
75 #[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 #[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
125pub struct EmbeddingProviderComponent {
127 name: String,
128 provider: Arc<dyn EmbeddingProvider>,
129}
130
131impl EmbeddingProviderComponent {
132 #[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 #[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
178pub struct ChatProviderAny {
180 inner: Arc<ChatProviderComponent>,
181}
182
183impl ChatProviderAny {
184 #[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
263pub struct EmbeddingProviderAny {
265 inner: Arc<EmbeddingProviderComponent>,
266}
267
268impl EmbeddingProviderAny {
269 #[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
348pub struct ChatProviderFactory {
351 descriptor: ComponentDescriptor,
352 provider: Arc<dyn ChatProvider>,
353}
354
355impl ChatProviderFactory {
356 #[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
400pub struct EmbeddingProviderFactory {
402 descriptor: ComponentDescriptor,
403 provider: Arc<dyn EmbeddingProvider>,
404}
405
406impl EmbeddingProviderFactory {
407 #[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
450pub 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}