1use std::collections::HashMap;
9use std::fmt;
10use std::sync::Arc;
11
12use async_trait::async_trait;
13use futures::StreamExt;
14use serde::{Deserialize, Serialize};
15
16use crate::driver_registry::{BoxedChatDriver, ChatDriver};
17use crate::error::Result;
18
19#[derive(Debug, Clone, PartialEq, Eq, Hash)]
21pub struct ProviderKey(String);
22
23impl ProviderKey {
24 pub fn new(id: impl AsRef<str>) -> Self {
25 Self(id.as_ref().trim().to_ascii_lowercase())
26 }
27
28 pub fn as_str(&self) -> &str {
29 &self.0
30 }
31}
32
33impl fmt::Display for ProviderKey {
34 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
35 f.write_str(&self.0)
36 }
37}
38
39impl From<&str> for ProviderKey {
40 fn from(value: &str) -> Self {
41 Self::new(value)
42 }
43}
44
45impl From<String> for ProviderKey {
46 fn from(value: String) -> Self {
47 Self::new(value)
48 }
49}
50
51impl Serialize for ProviderKey {
52 fn serialize<S: serde::Serializer>(
53 &self,
54 serializer: S,
55 ) -> std::result::Result<S::Ok, S::Error> {
56 serializer.serialize_str(self.as_str())
57 }
58}
59
60impl<'de> Deserialize<'de> for ProviderKey {
61 fn deserialize<D: serde::Deserializer<'de>>(
62 deserializer: D,
63 ) -> std::result::Result<Self, D::Error> {
64 String::deserialize(deserializer).map(Self::new)
65 }
66}
67
68pub struct ProviderAuthRequest<'a> {
73 pub method: &'a str,
74 pub url: &'a str,
75 pub headers: &'a [(String, String)],
76 pub body: &'a [u8],
77}
78
79#[async_trait]
81pub trait ProviderAuth: Send + Sync {
82 async fn headers(&self, request: ProviderAuthRequest<'_>) -> Result<Vec<(String, String)>>;
83 fn as_any(&self) -> &dyn std::any::Any;
84}
85
86pub struct BearerAuth {
88 key: String,
89}
90
91impl BearerAuth {
92 pub fn new(key: impl Into<String>) -> Self {
93 Self { key: key.into() }
94 }
95}
96
97#[async_trait]
98impl ProviderAuth for BearerAuth {
99 async fn headers(&self, _request: ProviderAuthRequest<'_>) -> Result<Vec<(String, String)>> {
100 Ok(vec![(
101 "authorization".to_string(),
102 format!("Bearer {}", self.key),
103 )])
104 }
105 fn as_any(&self) -> &dyn std::any::Any {
106 self
107 }
108}
109
110pub struct StaticHeaderAuth {
112 name: String,
113 value: String,
114}
115
116impl StaticHeaderAuth {
117 pub fn new(name: impl Into<String>, value: impl Into<String>) -> Self {
118 Self {
119 name: name.into().to_ascii_lowercase(),
120 value: value.into(),
121 }
122 }
123}
124
125#[async_trait]
126impl ProviderAuth for StaticHeaderAuth {
127 async fn headers(&self, _request: ProviderAuthRequest<'_>) -> Result<Vec<(String, String)>> {
128 Ok(vec![(self.name.clone(), self.value.clone())])
129 }
130 fn as_any(&self) -> &dyn std::any::Any {
131 self
132 }
133}
134
135#[derive(Clone, Default)]
140pub struct ProviderEndpoint {
141 base_url: Option<String>,
142 headers: Vec<(String, String)>,
143 auth: Option<Arc<dyn ProviderAuth>>,
144}
145
146impl ProviderEndpoint {
147 pub fn from_parts(base_url: impl Into<String>, auth: impl ProviderAuth + 'static) -> Self {
149 Self {
150 base_url: Some(base_url.into().trim_end_matches('/').to_string()),
151 headers: Vec::new(),
152 auth: Some(Arc::new(auth)),
153 }
154 }
155
156 pub fn base_url(&self) -> Option<&str> {
157 self.base_url.as_deref()
158 }
159
160 pub fn url(&self, path: &str) -> Option<String> {
161 self.base_url.as_ref().map(|base| {
162 let base = base.trim_end_matches('/');
163 if path.is_empty() {
164 return base.to_string();
165 }
166 if let Ok(mut url) = url::Url::parse(base) {
167 let (path, fragment) = path
170 .split_once('#')
171 .map_or((path, None), |(p, f)| (p, Some(f)));
172 let (path, query) = path
173 .split_once('?')
174 .map_or((path, None), |(p, q)| (p, Some(q)));
175 let suffix = format!("/{}", path.trim_start_matches('/'));
176 let existing = url.path().trim_end_matches('/');
179 if !existing.ends_with(&suffix) {
180 url.set_path(&format!("{existing}{suffix}"));
181 }
182 if let Some(query) = query {
183 let combined = match url.query() {
184 Some(existing) if !existing.is_empty() && !query.is_empty() => {
185 format!("{existing}&{query}")
186 }
187 Some(existing) if query.is_empty() => existing.to_string(),
188 _ => query.to_string(),
189 };
190 url.set_query(Some(&combined));
191 }
192 if let Some(fragment) = fragment {
193 url.set_fragment(Some(fragment));
194 }
195 url.to_string()
196 } else {
197 format!("{base}/{}", path.trim_start_matches('/'))
198 }
199 })
200 }
201
202 pub async fn resolve(
203 &self,
204 method: &str,
205 url: impl Into<String>,
206 body: &[u8],
207 ) -> Result<ResolvedProviderRequest> {
208 let url = url.into();
209 let mut headers = self.headers.clone();
210 if let Some(auth) = &self.auth {
211 let auth_headers = auth
212 .headers(ProviderAuthRequest {
213 method,
214 url: &url,
215 headers: &headers,
216 body,
217 })
218 .await?;
219 for (name, value) in auth_headers {
220 headers.retain(|(existing, _)| !existing.eq_ignore_ascii_case(&name));
221 headers.push((name.to_ascii_lowercase(), value));
222 }
223 }
224 Ok(ResolvedProviderRequest { url, headers })
225 }
226
227 pub fn auth<T: ProviderAuth + 'static>(&self) -> Option<&T> {
230 self.auth.as_deref()?.as_any().downcast_ref()
231 }
232}
233
234impl fmt::Debug for ProviderEndpoint {
235 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
236 f.debug_struct("ProviderEndpoint")
237 .field("base_url", &self.base_url.as_ref().map(|_| "<configured>"))
238 .field("auth", &self.auth.as_ref().map(|_| "<configured>"))
239 .field(
240 "headers",
241 &self
242 .headers
243 .iter()
244 .map(|(name, _)| name.as_str())
245 .collect::<Vec<_>>(),
246 )
247 .finish()
248 }
249}
250
251#[derive(Clone)]
253pub struct ResolvedProviderRequest {
254 pub url: String,
255 pub headers: Vec<(String, String)>,
256}
257
258impl fmt::Debug for ResolvedProviderRequest {
259 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
260 f.debug_struct("ResolvedProviderRequest")
261 .field("url", &redacted_url(&self.url))
262 .field(
263 "headers",
264 &self
265 .headers
266 .iter()
267 .map(|(name, _)| name.as_str())
268 .collect::<Vec<_>>(),
269 )
270 .finish()
271 }
272}
273
274fn redacted_url(value: &str) -> String {
275 let Ok(mut url) = url::Url::parse(value) else {
276 return "<configured>".to_string();
277 };
278 let _ = url.set_username("");
279 let _ = url.set_password(None);
280 url.set_query(None);
281 url.set_fragment(None);
282 url.to_string()
283}
284
285#[derive(Clone)]
287pub struct RuntimeProvider {
288 id: ProviderKey,
289 driver: Option<Arc<dyn ChatDriver>>,
290 decisions: Option<Arc<dyn crate::decision_driver::DecisionDriver>>,
291 embeddings: Option<Arc<dyn crate::driver_registry::EmbeddingsDriver>>,
292 endpoint: ProviderEndpoint,
293 driver_id: Option<crate::provider::DriverId>,
299}
300
301pub type Provider = RuntimeProvider;
303
304impl RuntimeProvider {
305 pub fn new(id: impl Into<ProviderKey>, driver: impl ChatDriver + 'static) -> Self {
306 Self::from_driver(id, Arc::new(driver))
307 }
308
309 pub fn from_driver(id: impl Into<ProviderKey>, driver: Arc<dyn ChatDriver>) -> Self {
310 Self {
311 id: id.into(),
312 driver: Some(driver),
313 decisions: None,
314 embeddings: None,
315 endpoint: ProviderEndpoint::default(),
316 driver_id: None,
317 }
318 }
319
320 pub fn base_url(mut self, url: impl Into<String>) -> Self {
321 self.endpoint.base_url = Some(url.into().trim_end_matches('/').to_string());
322 self
323 }
324
325 pub fn auth(mut self, auth: impl ProviderAuth + 'static) -> Self {
326 self.endpoint.auth = Some(Arc::new(auth));
327 self
328 }
329
330 pub fn auth_arc(mut self, auth: Arc<dyn ProviderAuth>) -> Self {
331 self.endpoint.auth = Some(auth);
332 self
333 }
334
335 pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
336 self.endpoint
337 .headers
338 .push((name.into().to_ascii_lowercase(), value.into()));
339 self
340 }
341
342 pub fn with_driver_id(mut self, driver_id: crate::provider::DriverId) -> Self {
345 self.driver_id = Some(driver_id);
346 self
347 }
348
349 pub fn id(&self) -> &ProviderKey {
350 &self.id
351 }
352
353 pub fn driver_id(&self) -> crate::provider::DriverId {
356 self.driver_id
357 .clone()
358 .unwrap_or_else(|| crate::provider::DriverId::external(self.id.as_str()))
359 }
360
361 pub fn driver(&self) -> Result<&Arc<dyn ChatDriver>> {
362 self.driver
363 .as_ref()
364 .ok_or_else(|| self.unsupported_service("chat"))
365 }
366
367 fn unsupported_service(&self, service: &str) -> crate::error::AgentLoopError {
368 crate::error::AgentLoopError::Configuration(format!(
369 "Provider '{}' does not implement the {service} service",
370 self.id
371 ))
372 }
373
374 pub fn services(id: impl Into<ProviderKey>) -> Self {
376 Self {
377 id: id.into(),
378 driver: None,
379 decisions: None,
380 embeddings: None,
381 endpoint: ProviderEndpoint::default(),
382 driver_id: None,
383 }
384 }
385
386 pub fn with_decisions(
387 mut self,
388 driver: impl crate::decision_driver::DecisionDriver + 'static,
389 ) -> Self {
390 self.decisions = Some(Arc::new(driver));
391 self
392 }
393
394 pub fn with_embeddings(
395 mut self,
396 driver: impl crate::driver_registry::EmbeddingsDriver + 'static,
397 ) -> Self {
398 self.embeddings = Some(Arc::new(driver));
399 self
400 }
401
402 pub fn supports_service(&self, service: crate::ServiceKind) -> bool {
403 match service {
404 crate::ServiceKind::Chat => self.driver.is_some(),
405 crate::ServiceKind::Decisions => self.decisions.is_some(),
406 crate::ServiceKind::Embeddings => self.embeddings.is_some(),
407 _ => false,
408 }
409 }
410
411 pub async fn evaluate_decisions(
412 &self,
413 request: crate::decisions::DecisionRequest,
414 ) -> Result<crate::decisions::DecisionOutcome> {
415 if request.provider.as_ref().is_some_and(|key| key != &self.id) {
416 return Err(crate::error::AgentLoopError::Configuration(
417 "A bound provider cannot select another account".into(),
418 ));
419 }
420 let driver = self
421 .decisions
422 .as_ref()
423 .ok_or_else(|| self.unsupported_service("decisions"))?;
424 driver.capabilities().check(driver.id(), &request)?;
425 driver
426 .evaluate(&self.endpoint, request)
427 .await
428 .map_err(|error| error.with_provider(self.id.as_str()))
429 }
430
431 pub async fn embed(
432 &self,
433 request: crate::driver_registry::EmbedRequest,
434 ) -> std::result::Result<
435 crate::driver_registry::EmbedResponse,
436 crate::driver_registry::EmbeddingsDriverError,
437 > {
438 let driver = self.embeddings.as_ref().ok_or_else(|| {
439 crate::driver_registry::EmbeddingsDriverError::Provider(
440 self.unsupported_service("embeddings").to_string(),
441 )
442 })?;
443 driver.embed(&self.endpoint, request).await
444 }
445
446 pub fn endpoint(&self) -> &ProviderEndpoint {
447 &self.endpoint
448 }
449
450 fn check_response_format(&self, config: &crate::driver_registry::LlmCallConfig) -> Result<()> {
453 if config.response_format.is_some()
454 && !self.driver()?.supports_response_format(&config.model)
455 {
456 return Err(crate::error::AgentLoopError::Configuration(format!(
457 "Structured output (response_format) is not supported by provider '{}' for model '{}'",
458 self.id, config.model
459 )));
460 }
461 Ok(())
462 }
463
464 pub async fn chat_completion_stream(
465 &self,
466 messages: Vec<crate::driver_registry::Message>,
467 config: &crate::driver_registry::LlmCallConfig,
468 ) -> Result<crate::driver_registry::LlmResponseStream> {
469 self.check_response_format(config)?;
470 let id = self.id.to_string();
471 let limits = config.limits;
472 let (stream, spent) = crate::turn_collector::connect_within(
476 &limits,
477 self.driver()?
478 .chat_completion_stream(&self.endpoint, messages, config),
479 )
480 .await
481 .map_err(|error| error.with_provider(&id))?;
482 let stream: crate::driver_registry::LlmResponseStream =
483 Box::pin(stream.map(move |result| result.map_err(|error| error.with_provider(&id))));
484 let stream = crate::turn_collector::limit_stream(stream, limits.after(spent));
488 Ok(Box::pin(stream.filter(|event| {
492 std::future::ready(!matches!(event, Ok(crate::driver_registry::LlmStreamEvent::TextDelta(delta)) if delta.is_empty()))
493 })))
494 }
495
496 pub async fn chat_completion(
497 &self,
498 messages: Vec<crate::driver_registry::Message>,
499 config: &crate::driver_registry::LlmCallConfig,
500 ) -> Result<crate::driver_registry::LlmResponse> {
501 self.check_response_format(config)?;
502 self.driver()?
503 .chat_completion(&self.endpoint, messages, config)
504 .await
505 .map_err(|error| error.with_provider(self.id.as_str()))
506 }
507
508 pub fn supports_native_non_streaming(&self) -> bool {
509 self.driver
510 .as_ref()
511 .is_some_and(|driver| driver.supports_native_non_streaming())
512 }
513
514 pub async fn chat_completion_non_streaming(
515 &self,
516 messages: Vec<crate::driver_registry::Message>,
517 config: &crate::driver_registry::LlmCallConfig,
518 ) -> Result<crate::driver_registry::LlmResponse> {
519 self.check_response_format(config)?;
520 if !config.limits.is_unbounded() {
521 let stream = self.chat_completion_stream(messages, config).await?;
525 return Ok(
526 crate::turn_collector::collect_turn(stream, &config.limits, |_| {})
527 .await?
528 .into_response(),
529 );
530 }
531 self.driver()?
532 .chat_completion_non_streaming(&self.endpoint, messages, config)
533 .await
534 .map_err(|error| error.with_provider(self.id.as_str()))
535 }
536
537 pub async fn list_models(
538 &self,
539 ) -> Result<Option<Vec<crate::driver_registry::DiscoveredModel>>> {
540 let Some(driver) = &self.driver else {
541 return Ok(None);
542 };
543 driver
544 .list_models(&self.endpoint)
545 .await
546 .map_err(|error| error.with_provider(self.id.as_str()))
547 }
548
549 pub async fn models(
556 &self,
557 ) -> Result<Option<Vec<crate::model_discovery::DiscoveredProviderModel>>> {
558 let Some(models) = self.list_models().await? else {
559 return Ok(None);
560 };
561 Ok(Some(crate::model_discovery::normalize_and_enrich(
562 &self.driver_id(),
563 models,
564 )))
565 }
566
567 pub fn into_boxed_driver(self) -> BoxedChatDriver {
568 Box::new(ProviderBoundDriver(self))
569 }
570
571 pub fn into_embeddings_driver(
572 self,
573 ) -> std::result::Result<
574 crate::driver_registry::BoxedEmbeddingsDriver,
575 crate::driver_registry::EmbeddingsDriverError,
576 > {
577 if self.embeddings.is_none() {
578 return Err(crate::driver_registry::EmbeddingsDriverError::Provider(
579 "Provider does not support embeddings".into(),
580 ));
581 }
582 Ok(Box::new(ProviderOwnedEmbeddingsDriver(self)))
583 }
584
585 pub fn bind_embeddings(
586 self,
587 driver: crate::driver_registry::BoxedEmbeddingsDriver,
588 ) -> crate::driver_registry::BoxedEmbeddingsDriver {
589 Box::new(ProviderBoundEmbeddingsDriver {
590 id: self.id,
591 endpoint: self.endpoint,
592 driver,
593 })
594 }
595}
596
597impl fmt::Debug for RuntimeProvider {
598 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
599 f.debug_struct("Provider")
600 .field("id", &self.id)
601 .field("endpoint", &self.endpoint)
602 .finish_non_exhaustive()
603 }
604}
605
606struct ProviderBoundDriver(RuntimeProvider);
607
608struct ProviderBoundEmbeddingsDriver {
609 id: ProviderKey,
610 endpoint: ProviderEndpoint,
611 driver: crate::driver_registry::BoxedEmbeddingsDriver,
612}
613
614#[async_trait]
615impl crate::driver_registry::EmbeddingsDriver for ProviderBoundEmbeddingsDriver {
616 async fn embed(
617 &self,
618 _endpoint: &ProviderEndpoint,
619 request: crate::driver_registry::EmbedRequest,
620 ) -> std::result::Result<
621 crate::driver_registry::EmbedResponse,
622 crate::driver_registry::EmbeddingsDriverError,
623 > {
624 self.driver
625 .embed(&self.endpoint, request)
626 .await
627 .map_err(|error| {
628 crate::driver_registry::EmbeddingsDriverError::Provider(format!(
629 "provider '{}': {error}",
630 self.id
631 ))
632 })
633 }
634}
635
636#[async_trait]
637impl ChatDriver for ProviderBoundDriver {
638 fn native_async_driver(
639 &self,
640 model: &str,
641 tools: std::collections::BTreeMap<String, Option<serde_json::Value>>,
642 continuation: Option<crate::native_async::Delivery>,
643 ) -> Option<Arc<dyn ChatDriver>> {
644 let driver = self
645 .0
646 .driver
647 .as_ref()?
648 .native_async_driver(model, tools, continuation)?;
649 Some(Arc::new(ProviderBoundDriver(RuntimeProvider {
650 id: self.0.id.clone(),
651 endpoint: self.0.endpoint.clone(),
652 driver: Some(driver),
653 decisions: self.0.decisions.clone(),
654 embeddings: self.0.embeddings.clone(),
655 driver_id: self.0.driver_id.clone(),
656 })))
657 }
658 async fn chat_completion_stream(
659 &self,
660 _endpoint: &ProviderEndpoint,
661 messages: Vec<crate::driver_registry::Message>,
662 config: &crate::driver_registry::LlmCallConfig,
663 ) -> Result<crate::driver_registry::LlmResponseStream> {
664 self.0.chat_completion_stream(messages, config).await
665 }
666
667 async fn list_models(
668 &self,
669 _endpoint: &ProviderEndpoint,
670 ) -> Result<Option<Vec<crate::driver_registry::DiscoveredModel>>> {
671 self.0.list_models().await
672 }
673
674 fn supports_native_non_streaming(&self) -> bool {
675 self.0.supports_native_non_streaming()
676 }
677
678 async fn chat_completion_non_streaming(
679 &self,
680 _endpoint: &ProviderEndpoint,
681 messages: Vec<crate::driver_registry::Message>,
682 config: &crate::driver_registry::LlmCallConfig,
683 ) -> Result<crate::driver_registry::LlmResponse> {
684 self.0.chat_completion_non_streaming(messages, config).await
685 }
686
687 fn supports_compact(&self) -> bool {
688 self.0
689 .driver
690 .as_ref()
691 .map(|driver| driver.supports_compact())
692 .unwrap_or(false)
693 }
694
695 fn supports_stateful_responses(&self) -> bool {
696 self.0
697 .driver
698 .as_ref()
699 .map(|driver| driver.supports_stateful_responses())
700 .unwrap_or(false)
701 }
702
703 fn effective_context_window(&self, model: &str) -> Option<usize> {
704 self.0
705 .driver
706 .as_ref()
707 .map(|driver| driver.effective_context_window(model))
708 .unwrap_or(None)
709 }
710
711 fn supports_parallel_tool_calls(&self, model: &str) -> bool {
712 self.0
713 .driver
714 .as_ref()
715 .map(|driver| driver.supports_parallel_tool_calls(model))
716 .unwrap_or(false)
717 }
718
719 fn supports_response_format(&self, model: &str) -> bool {
720 self.0
721 .driver
722 .as_ref()
723 .map(|driver| driver.supports_response_format(model))
724 .unwrap_or(false)
725 }
726
727 fn provider_managed_reduction_option(
728 &self,
729 _endpoint: &ProviderEndpoint,
730 model: &str,
731 budget_tokens: usize,
732 ) -> Option<(String, serde_json::Value)> {
733 self.0.driver.as_ref()?.provider_managed_reduction_option(
734 self.0.endpoint(),
735 model,
736 budget_tokens,
737 )
738 }
739
740 fn provider_managed_reduction_fallback_reason(
741 &self,
742 _endpoint: &ProviderEndpoint,
743 config: &crate::driver_registry::LlmCallConfig,
744 ) -> Option<&'static str> {
745 self.0
746 .driver
747 .as_ref()?
748 .provider_managed_reduction_fallback_reason(self.0.endpoint(), config)
749 }
750
751 fn validate_provider_opaque_context(
752 &self,
753 context: &crate::driver_registry::ProviderOpaqueContext,
754 ) -> bool {
755 self.0
756 .driver
757 .as_ref()
758 .map(|driver| driver.validate_provider_opaque_context(context))
759 .unwrap_or(false)
760 }
761
762 async fn compact(
763 &self,
764 _endpoint: &ProviderEndpoint,
765 request: crate::compact::CompactRequest,
766 ) -> Result<Option<crate::compact::CompactResponse>> {
767 self.0
768 .driver()?
769 .compact(self.0.endpoint(), request)
770 .await
771 .map_err(|error| error.with_provider(self.0.id.as_str()))
772 }
773}
774
775#[derive(Clone, Default)]
777pub struct RuntimeProviderRegistry {
778 providers: HashMap<ProviderKey, Arc<RuntimeProvider>>,
779}
780
781pub type ProviderRegistry = RuntimeProviderRegistry;
783
784impl RuntimeProviderRegistry {
785 pub fn new() -> Self {
786 Self::default()
787 }
788
789 pub fn register(&mut self, provider: RuntimeProvider) -> Result<()> {
790 if self.providers.contains_key(provider.id()) {
791 return Err(crate::error::AgentLoopError::Configuration(format!(
792 "provider '{}' is already registered; use replace() to overwrite intentionally",
793 provider.id()
794 )));
795 }
796 self.providers
797 .insert(provider.id.clone(), Arc::new(provider));
798 Ok(())
799 }
800
801 pub fn replace(&mut self, provider: RuntimeProvider) -> Option<Arc<RuntimeProvider>> {
802 self.providers
803 .insert(provider.id.clone(), Arc::new(provider))
804 }
805
806 pub fn get(&self, id: &ProviderKey) -> Option<Arc<RuntimeProvider>> {
807 self.providers.get(id).cloned()
808 }
809
810 pub fn ids(&self) -> Vec<String> {
811 let mut ids = self
812 .providers
813 .keys()
814 .map(ToString::to_string)
815 .collect::<Vec<_>>();
816 ids.sort();
817 ids
818 }
819}
820
821impl fmt::Debug for RuntimeProviderRegistry {
822 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
823 f.debug_struct("ProviderRegistry")
824 .field("providers", &self.ids())
825 .finish()
826 }
827}
828
829#[async_trait]
830impl crate::decisions::DecisionsService for RuntimeProvider {
831 fn is_configured(&self) -> bool {
832 self.decisions.is_some()
833 }
834 async fn evaluate(
835 &self,
836 request: crate::decisions::DecisionRequest,
837 ) -> Result<crate::decisions::DecisionOutcome> {
838 if request
839 .provider
840 .as_ref()
841 .is_some_and(|provider| provider != &self.id)
842 {
843 return Err(crate::error::AgentLoopError::Configuration(
844 "A bound provider cannot select another account".into(),
845 ));
846 }
847 self.evaluate_decisions(request).await
848 }
849 fn name(&self) -> &'static str {
850 "ProviderDecisions"
851 }
852}
853
854struct ProviderOwnedEmbeddingsDriver(RuntimeProvider);
855#[async_trait]
856impl crate::driver_registry::EmbeddingsDriver for ProviderOwnedEmbeddingsDriver {
857 async fn embed(
858 &self,
859 _endpoint: &ProviderEndpoint,
860 request: crate::driver_registry::EmbedRequest,
861 ) -> std::result::Result<
862 crate::driver_registry::EmbedResponse,
863 crate::driver_registry::EmbeddingsDriverError,
864 > {
865 self.0.embed(request).await
866 }
867}
868
869#[cfg(test)]
870mod tests {
871 use super::*;
872 use std::sync::atomic::{AtomicBool, Ordering};
873 use std::time::Duration;
874
875 struct Noop;
876 #[async_trait]
877 impl ChatDriver for Noop {
878 async fn chat_completion_stream(
879 &self,
880 _endpoint: &ProviderEndpoint,
881 _messages: Vec<crate::Message>,
882 _config: &crate::LlmCallConfig,
883 ) -> Result<crate::LlmResponseStream> {
884 unreachable!("configuration-only fixture must not execute")
885 }
886 }
887
888 struct Catalog;
889 #[async_trait]
890 impl ChatDriver for Catalog {
891 async fn chat_completion_stream(
892 &self,
893 _endpoint: &ProviderEndpoint,
894 _messages: Vec<crate::Message>,
895 _config: &crate::LlmCallConfig,
896 ) -> Result<crate::LlmResponseStream> {
897 unreachable!("catalog-only fixture must not execute")
898 }
899
900 async fn list_models(
901 &self,
902 _endpoint: &ProviderEndpoint,
903 ) -> Result<Option<Vec<crate::driver_registry::DiscoveredModel>>> {
904 Ok(Some(vec![crate::driver_registry::DiscoveredModel {
905 model_id: "gpt-5.6-terra".to_string(),
906 display_name: None,
907 created_at: None,
908 owned_by: None,
909 capabilities: vec!["chat".to_string()],
910 discovered_profile: None,
911 }]))
912 }
913 }
914
915 struct NativeNonStreaming {
916 native_called: Arc<AtomicBool>,
917 connect_delay: Duration,
918 event_delay: Duration,
919 response: String,
920 }
921
922 #[async_trait]
923 impl ChatDriver for NativeNonStreaming {
924 fn supports_native_non_streaming(&self) -> bool {
925 true
926 }
927
928 async fn chat_completion_stream(
929 &self,
930 _endpoint: &ProviderEndpoint,
931 _messages: Vec<crate::Message>,
932 _config: &crate::LlmCallConfig,
933 ) -> Result<crate::LlmResponseStream> {
934 tokio::time::sleep(self.connect_delay).await;
935 let events = vec![
936 (
937 Duration::ZERO,
938 Ok(crate::LlmStreamEvent::TextDelta(self.response.clone())),
939 ),
940 (
941 self.event_delay,
942 Ok(crate::LlmStreamEvent::Done(Box::default())),
943 ),
944 ];
945 Ok(Box::pin(futures::stream::iter(events).then(
946 |(delay, event)| async move {
947 tokio::time::sleep(delay).await;
948 event
949 },
950 )))
951 }
952
953 async fn chat_completion_non_streaming(
954 &self,
955 _endpoint: &ProviderEndpoint,
956 _messages: Vec<crate::Message>,
957 _config: &crate::LlmCallConfig,
958 ) -> Result<crate::LlmResponse> {
959 self.native_called.store(true, Ordering::SeqCst);
960 Ok(crate::LlmResponse {
961 text: self.response.clone(),
962 reasoning: vec![],
963 tool_calls: None,
964 metadata: Default::default(),
965 })
966 }
967 }
968
969 #[test]
970 fn the_driver_kind_falls_back_to_the_runtime_key() {
971 assert_eq!(
972 RuntimeProvider::new("openai", Noop).driver_id(),
973 crate::provider::DriverId::OpenAI
974 );
975 }
976
977 #[test]
978 fn a_declared_driver_kind_survives_a_caller_chosen_key() {
979 let provider = RuntimeProvider::new("my-gateway", Noop)
980 .with_driver_id(crate::provider::DriverId::OpenAI);
981 assert_eq!(provider.id().as_str(), "my-gateway");
982 assert_eq!(provider.driver_id(), crate::provider::DriverId::OpenAI);
983 }
984
985 #[tokio::test]
986 async fn models_enriches_bare_ids_through_the_declared_driver_kind() {
987 let catalog = RuntimeProvider::new("my-gateway", Catalog)
988 .with_driver_id(crate::provider::DriverId::OpenAI)
989 .models()
990 .await
991 .expect("catalog request")
992 .expect("driver offers a catalog");
993 assert_eq!(catalog.len(), 1);
994 assert!(catalog[0].display_name.is_some());
996 }
997
998 #[tokio::test]
999 async fn a_driver_without_a_catalog_reports_no_catalog() {
1000 assert!(
1001 RuntimeProvider::new("openai", Noop)
1002 .models()
1003 .await
1004 .expect("catalog request")
1005 .is_none()
1006 );
1007 }
1008
1009 #[tokio::test]
1010 async fn bounded_non_streaming_calls_enforce_the_total_timeout_before_headers() {
1011 let native_called = Arc::new(AtomicBool::new(false));
1012 let provider = RuntimeProvider::new(
1013 "bounded",
1014 NativeNonStreaming {
1015 native_called: Arc::clone(&native_called),
1016 connect_delay: Duration::from_millis(100),
1017 event_delay: Duration::ZERO,
1018 response: "done".into(),
1019 },
1020 );
1021 let mut config = crate::LlmCallConfig::new("model");
1022 config.limits =
1023 crate::turn_collector::TurnLimits::default().with_total(Duration::from_millis(10));
1024
1025 let error = provider
1026 .chat_completion_non_streaming(vec![], &config)
1027 .await
1028 .expect_err("the call must time out");
1029
1030 assert!(error.to_string().contains("did not finish"));
1031 assert!(!native_called.load(Ordering::SeqCst));
1032 }
1033
1034 #[tokio::test]
1035 async fn bounded_non_streaming_calls_enforce_the_total_timeout_during_the_body() {
1036 let native_called = Arc::new(AtomicBool::new(false));
1037 let provider = RuntimeProvider::new(
1038 "bounded",
1039 NativeNonStreaming {
1040 native_called: Arc::clone(&native_called),
1041 connect_delay: Duration::ZERO,
1042 event_delay: Duration::from_millis(100),
1043 response: "partial".into(),
1044 },
1045 );
1046 let mut config = crate::LlmCallConfig::new("model");
1047 config.limits =
1048 crate::turn_collector::TurnLimits::default().with_total(Duration::from_millis(10));
1049
1050 let error = provider
1051 .chat_completion_non_streaming(vec![], &config)
1052 .await
1053 .expect_err("the body must time out");
1054
1055 assert!(error.to_string().contains("did not finish"));
1056 assert!(!native_called.load(Ordering::SeqCst));
1057 }
1058
1059 #[tokio::test]
1060 async fn bounded_non_streaming_calls_reject_oversized_responses() {
1061 let native_called = Arc::new(AtomicBool::new(false));
1062 let provider = RuntimeProvider::new(
1063 "bounded",
1064 NativeNonStreaming {
1065 native_called: Arc::clone(&native_called),
1066 connect_delay: Duration::ZERO,
1067 event_delay: Duration::ZERO,
1068 response: "too large".into(),
1069 },
1070 );
1071 let mut config = crate::LlmCallConfig::new("model");
1072 config.limits = crate::turn_collector::TurnLimits::default().with_max_response_bytes(3);
1073
1074 let error = provider
1075 .chat_completion_non_streaming(vec![], &config)
1076 .await
1077 .expect_err("the response must exceed the cap");
1078
1079 assert!(error.to_string().contains("3-byte limit"));
1080 assert!(!native_called.load(Ordering::SeqCst));
1081 }
1082
1083 #[test]
1084 fn endpoint_deduplicates_only_complete_path_segments() {
1085 for (base, path, expected) in [
1086 (
1087 "https://service.example/v1/",
1088 "chat",
1089 "https://service.example/v1/chat",
1090 ),
1091 (
1092 "https://service.example/v1/chat",
1093 "/chat",
1094 "https://service.example/v1/chat",
1095 ),
1096 (
1097 "https://service.example/v1/notchat",
1098 "chat",
1099 "https://service.example/v1/notchat/chat",
1100 ),
1101 ("https://chat", "chat", "https://chat/chat"),
1102 (
1103 "https://service.example/v1beta",
1104 "models/gemini-2.5-flash:streamGenerateContent?alt=sse",
1105 "https://service.example/v1beta/models/gemini-2.5-flash:streamGenerateContent?alt=sse",
1106 ),
1107 (
1108 "https://service.example/v1?token=a%2Fb#base",
1109 "chat?alt=sse#operation",
1110 "https://service.example/v1/chat?token=a%2Fb&alt=sse#operation",
1111 ),
1112 (
1113 "https://service.example/v1/chat?token=x",
1114 "chat?alt=sse",
1115 "https://service.example/v1/chat?token=x&alt=sse",
1116 ),
1117 (
1118 "https://service.example/v1?token=x",
1119 "chat",
1120 "https://service.example/v1/chat?token=x",
1121 ),
1122 (
1123 "https://service.example/v1/chat?token=x",
1124 "chat",
1125 "https://service.example/v1/chat?token=x",
1126 ),
1127 (
1128 "https://service.example/v1/notchat/completions",
1129 "chat/completions",
1130 "https://service.example/v1/notchat/completions/chat/completions",
1131 ),
1132 (
1133 "https://service.example/v1/chat/",
1134 "",
1135 "https://service.example/v1/chat",
1136 ),
1137 ] {
1138 let endpoint = ProviderEndpoint {
1139 base_url: Some(base.into()),
1140 ..Default::default()
1141 };
1142 assert_eq!(
1143 endpoint.url(path).as_deref(),
1144 Some(expected),
1145 "{base} + {path}"
1146 );
1147 }
1148 assert_eq!(ProviderEndpoint::default().url("chat"), None);
1149 }
1150
1151 #[test]
1152 fn provider_key_deserialization_is_canonical() {
1153 let key: ProviderKey = serde_json::from_str(r#"" Gateway-PROD ""#).unwrap();
1154 assert_eq!(key.as_str(), "gateway-prod");
1155 assert_eq!(serde_json::to_string(&key).unwrap(), r#""gateway-prod""#);
1156 }
1157
1158 #[tokio::test]
1159 async fn debug_redacts_auth_and_header_values() {
1160 let provider = RuntimeProvider::new("Gateway", Noop)
1161 .base_url("https://example.test/")
1162 .header("x-secret", "hidden-service-value")
1163 .auth(BearerAuth::new("hidden-key"));
1164 let debug = format!("{provider:?}");
1165 assert!(debug.contains("gateway"));
1166 assert!(debug.contains("x-secret"));
1167 assert!(!debug.contains("hidden-service-value"));
1168 assert!(!debug.contains("hidden-key"));
1169 let request = provider
1170 .endpoint()
1171 .resolve(
1172 "POST",
1173 "https://user:password@example.test/chat?token=query-secret#fragment-secret",
1174 b"{}",
1175 )
1176 .await
1177 .unwrap();
1178 assert_eq!(
1179 format!("{request:?}"),
1180 "ResolvedProviderRequest { url: \"https://example.test/chat\", headers: [\"x-secret\", \"authorization\"] }"
1181 );
1182 }
1183
1184 #[test]
1185 fn duplicate_registration_is_explicit() {
1186 let mut registry = RuntimeProviderRegistry::new();
1187 registry.register(RuntimeProvider::new("a", Noop)).unwrap();
1188 let error = registry
1189 .register(RuntimeProvider::new("A", Noop).base_url("https://rejected.example"))
1190 .unwrap_err();
1191 let crate::AgentLoopError::Configuration(message) = error else {
1192 panic!("duplicate must be a configuration error")
1193 };
1194 assert_eq!(
1195 message,
1196 "provider 'a' is already registered; use replace() to overwrite intentionally"
1197 );
1198 assert_eq!(registry.ids(), vec!["a"]);
1199 let original = registry.get(&ProviderKey::new(" A ")).unwrap();
1200 assert!(
1201 original.endpoint().base_url().is_none(),
1202 "rejected duplicate must not replace original"
1203 );
1204 let old = registry
1205 .replace(RuntimeProvider::new("a", Noop).base_url("https://replacement.example"));
1206 assert!(Arc::ptr_eq(&old.unwrap(), &original));
1207 assert_eq!(
1208 registry
1209 .get(&ProviderKey::new("a"))
1210 .unwrap()
1211 .endpoint()
1212 .base_url(),
1213 Some("https://replacement.example")
1214 );
1215 assert!(
1216 original.endpoint().base_url().is_none(),
1217 "existing handle retains old provider"
1218 );
1219 assert!(registry.get(&ProviderKey::new("missing")).is_none());
1220 assert!(registry.replace(RuntimeProvider::new("z", Noop)).is_none());
1221 assert_eq!(registry.ids(), ["a", "z"]);
1222 }
1223
1224 #[tokio::test]
1225 async fn one_protocol_serves_distinct_provider_identities() {
1226 let protocol: Arc<dyn ChatDriver> = Arc::new(Noop);
1227 let first = Provider::from_driver("first", protocol.clone())
1228 .base_url("https://first.example/v1")
1229 .header("x-service", "first")
1230 .auth(BearerAuth::new("first-key"));
1231 let second = Provider::from_driver("second", protocol.clone())
1232 .base_url("https://second.example/v1")
1233 .header("x-service", "second")
1234 .auth(BearerAuth::new("second-key"));
1235
1236 assert!(Arc::ptr_eq(
1237 first.driver().unwrap(),
1238 second.driver().unwrap()
1239 ));
1240 let first_request = first
1241 .endpoint()
1242 .resolve("POST", first.endpoint().url("chat").unwrap(), b"{}")
1243 .await
1244 .unwrap();
1245 let second_request = second
1246 .endpoint()
1247 .resolve("POST", second.endpoint().url("chat").unwrap(), b"{}")
1248 .await
1249 .unwrap();
1250 assert_eq!(first_request.url, "https://first.example/v1/chat");
1251 assert_eq!(second_request.url, "https://second.example/v1/chat");
1252 assert_eq!(
1253 first_request.headers,
1254 [
1255 ("x-service".into(), "first".into()),
1256 ("authorization".into(), "Bearer first-key".into())
1257 ]
1258 );
1259 assert_eq!(
1260 second_request.headers,
1261 [
1262 ("x-service".into(), "second".into()),
1263 ("authorization".into(), "Bearer second-key".into())
1264 ]
1265 );
1266 }
1267
1268 #[tokio::test]
1269 async fn refreshable_auth_is_resolved_for_each_request() {
1270 struct Rotating(std::sync::atomic::AtomicUsize);
1271 #[async_trait]
1272 impl ProviderAuth for Rotating {
1273 async fn headers(
1274 &self,
1275 request: ProviderAuthRequest<'_>,
1276 ) -> Result<Vec<(String, String)>> {
1277 assert_eq!(request.method, "POST");
1278 assert_eq!(request.url, "https://service.example/chat");
1279 assert_eq!(
1280 request.headers,
1281 [
1282 ("AUTHORIZATION".into(), "stale".into()),
1283 ("x-static".into(), "preserve".into())
1284 ]
1285 );
1286 let token = self.0.fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1;
1287 Ok(vec![
1288 ("authorization".into(), format!("Bearer token-{token}")),
1289 (
1290 "x-signed-body".into(),
1291 String::from_utf8_lossy(request.body).into_owned(),
1292 ),
1293 ])
1294 }
1295 fn as_any(&self) -> &dyn std::any::Any {
1296 self
1297 }
1298 }
1299
1300 let endpoint = ProviderEndpoint {
1301 base_url: Some("https://service.example".into()),
1302 headers: vec![
1303 ("AUTHORIZATION".into(), "stale".into()),
1304 ("x-static".into(), "preserve".into()),
1305 ],
1306 auth: Some(Arc::new(Rotating(std::sync::atomic::AtomicUsize::new(0)))),
1307 };
1308 let first = endpoint
1309 .resolve("POST", "https://service.example/chat", b"one")
1310 .await
1311 .unwrap();
1312 let second = endpoint
1313 .resolve("POST", "https://service.example/chat", b"two")
1314 .await
1315 .unwrap();
1316 assert_eq!(
1317 first.headers,
1318 [
1319 ("x-static".into(), "preserve".into()),
1320 ("authorization".into(), "Bearer token-1".into()),
1321 ("x-signed-body".into(), "one".into())
1322 ]
1323 );
1324 assert_eq!(
1325 second.headers,
1326 [
1327 ("x-static".into(), "preserve".into()),
1328 ("authorization".into(), "Bearer token-2".into()),
1329 ("x-signed-body".into(), "two".into())
1330 ]
1331 );
1332 }
1333
1334 #[tokio::test]
1335 async fn provider_identity_prefixes_start_and_stream_errors() {
1336 struct Failing {
1337 fail_to_start: bool,
1338 }
1339 #[async_trait]
1340 impl ChatDriver for Failing {
1341 async fn chat_completion_stream(
1342 &self,
1343 _endpoint: &ProviderEndpoint,
1344 _messages: Vec<crate::Message>,
1345 _config: &crate::LlmCallConfig,
1346 ) -> Result<crate::LlmResponseStream> {
1347 if self.fail_to_start {
1348 return Err(crate::AgentLoopError::llm("request failed"));
1349 }
1350 Ok(Box::pin(futures::stream::once(async {
1351 Err(crate::AgentLoopError::llm("stream failed"))
1352 })))
1353 }
1354 }
1355
1356 let config = crate::LlmCallConfig {
1357 model: "model".into(),
1358 ..Default::default()
1359 };
1360 let start = Provider::new(
1361 "customer-gateway",
1362 Failing {
1363 fail_to_start: true,
1364 },
1365 );
1366 let error = match start.chat_completion_stream(Vec::new(), &config).await {
1367 Ok(_) => panic!("the test driver should fail before returning a stream"),
1368 Err(error) => error,
1369 };
1370 let crate::AgentLoopError::Llm(error) = error else {
1371 panic!("LLM error variant must survive")
1372 };
1373 assert_eq!(error.message, "provider 'customer-gateway': request failed");
1374
1375 let stream = Provider::new(
1376 "customer-gateway",
1377 Failing {
1378 fail_to_start: false,
1379 },
1380 );
1381 let error = stream
1382 .chat_completion_stream(Vec::new(), &config)
1383 .await
1384 .unwrap()
1385 .next()
1386 .await
1387 .unwrap()
1388 .unwrap_err();
1389 let crate::AgentLoopError::Llm(error) = error else {
1390 panic!("LLM error variant must survive")
1391 };
1392 assert_eq!(error.message, "provider 'customer-gateway': stream failed");
1393 }
1394}
1395
1396#[cfg(test)]
1397#[path = "runtime_provider_service_tests.rs"]
1398mod service_tests;