1use crate::credential_schema::CredentialFormSchema;
18use crate::error::{AgentLoopError, LlmErrorKind, Result};
19use crate::openresponses_protocol::{CompactOutputItem, CompactRequest, CompactResponse};
20use crate::tool_types::{ToolCall, ToolDefinition};
21use async_trait::async_trait;
22use chrono::{DateTime, Utc};
23use futures::Stream;
24use serde::{Deserialize, Serialize};
25use std::collections::HashMap;
26use std::pin::Pin;
27use std::sync::Arc;
28
29pub type LlmResponseStream = Pin<Box<dyn Stream<Item = Result<LlmStreamEvent>> + Send>>;
35
36#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
42#[serde(tag = "type", rename_all = "snake_case")]
43pub enum ProviderOpaqueContext {
44 OpenResponsesCompact { output: Vec<CompactOutputItem> },
46}
47
48impl std::fmt::Debug for ProviderOpaqueContext {
49 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
50 match self {
51 Self::OpenResponsesCompact { output } => f
52 .debug_struct("OpenResponsesCompact")
53 .field("item_count", &output.len())
54 .finish_non_exhaustive(),
55 }
56 }
57}
58
59#[derive(Debug, Clone, PartialEq, Eq)]
65pub struct LlmStreamError {
66 pub code: Option<String>,
68 pub status: Option<u16>,
70 pub message: String,
72}
73
74impl LlmStreamError {
75 pub fn new(message: impl Into<String>) -> Self {
76 Self {
77 code: None,
78 status: None,
79 message: message.into(),
80 }
81 }
82
83 pub fn provider(
85 code: Option<impl Into<String>>,
86 status: Option<u16>,
87 message: impl Into<String>,
88 ) -> Self {
89 Self {
90 code: code.map(Into::into),
91 status,
92 message: message.into(),
93 }
94 }
95
96 pub fn kind(&self) -> LlmErrorKind {
98 if let Some(code) = self.code.as_deref()
99 && let Some(kind) = LlmErrorKind::from_provider_code(code)
100 {
101 return kind;
102 }
103 if let Some(status) = self.status {
104 return LlmErrorKind::from_provider_status(status, &self.message);
105 }
106 LlmErrorKind::from_error_text(&self.message)
107 }
108}
109
110impl std::error::Error for LlmStreamError {}
111
112impl std::fmt::Display for LlmStreamError {
113 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
114 match (&self.code, self.status) {
115 (Some(code), Some(status)) => write!(f, "{code} ({status}): {}", self.message),
116 (Some(code), None) => write!(f, "{code}: {}", self.message),
117 (None, Some(status)) => write!(f, "({status}): {}", self.message),
118 (None, None) => f.write_str(&self.message),
119 }
120 }
121}
122
123impl From<String> for LlmStreamError {
124 fn from(message: String) -> Self {
125 Self::new(message)
126 }
127}
128
129impl From<&str> for LlmStreamError {
130 fn from(message: &str) -> Self {
131 Self::new(message)
132 }
133}
134
135#[derive(Debug, Clone)]
137pub enum LlmStreamEvent {
138 TextDelta(String),
140 ThinkingDelta(String),
142 ThinkingSignature(String),
145 ReasonItem {
151 provider: String,
153 model: Option<String>,
155 item_id: String,
157 encrypted_content: Option<String>,
159 summary: Vec<String>,
161 token_count: Option<u32>,
163 },
164 ToolCalls(Vec<ToolCall>),
166 MessagePhase(crate::execution_phase::ExecutionPhase),
176 Done(Box<LlmCompletionMetadata>),
178 Error(LlmStreamError),
180}
181
182#[derive(Debug, Clone)]
193pub struct DiscoveredModel {
194 pub model_id: String,
196 pub display_name: Option<String>,
198 pub created_at: Option<DateTime<Utc>>,
200 pub owned_by: Option<String>,
202 pub capabilities: Vec<String>,
206 pub discovered_profile: Option<crate::model::ModelProfile>,
209}
210
211#[derive(Debug, Clone, Default)]
224pub struct LlmCompletionMetadata {
225 pub total_tokens: Option<u32>,
227 pub prompt_tokens: Option<u32>,
229 pub completion_tokens: Option<u32>,
231 pub cache_read_tokens: Option<u32>,
233 pub cache_creation_tokens: Option<u32>,
235 pub provider_cost_usd: Option<f64>,
239 pub model: Option<String>,
241 pub finish_reason: Option<String>,
243 pub retry_metadata: Option<crate::llm_retry::RetryMetadata>,
245 pub response_id: Option<String>,
248 pub phase: Option<String>,
252}
253
254pub fn disjoint_prompt_tokens(reported_input: u32, cache_read: Option<u32>) -> u32 {
264 reported_input.saturating_sub(cache_read.unwrap_or(0))
265}
266
267#[async_trait]
289pub trait ChatDriver: Send + Sync {
290 async fn chat_completion_stream(
292 &self,
293 messages: Vec<LlmMessage>,
294 config: &LlmCallConfig,
295 ) -> Result<LlmResponseStream>;
296
297 async fn chat_completion(
299 &self,
300 messages: Vec<LlmMessage>,
301 config: &LlmCallConfig,
302 ) -> Result<LlmResponse> {
303 use futures::StreamExt;
304
305 let mut stream = self.chat_completion_stream(messages, config).await?;
306 let mut text = String::new();
307 let mut thinking = String::new();
308 let mut thinking_signature: Option<String> = None;
309 let mut tool_calls = Vec::new();
310 let mut metadata = LlmCompletionMetadata::default();
311
312 while let Some(event) = stream.next().await {
313 match event? {
314 LlmStreamEvent::TextDelta(delta) => text.push_str(&delta),
315 LlmStreamEvent::ThinkingDelta(delta) => thinking.push_str(&delta),
316 LlmStreamEvent::ThinkingSignature(sig) => thinking_signature = Some(sig),
317 LlmStreamEvent::ReasonItem {
318 encrypted_content, ..
319 } => {
320 if let Some(sig) = encrypted_content {
321 thinking_signature = Some(sig);
322 }
323 }
324 LlmStreamEvent::ToolCalls(calls) => tool_calls = calls,
325 LlmStreamEvent::MessagePhase(_) => {}
328 LlmStreamEvent::Done(meta) => metadata = *meta,
329 LlmStreamEvent::Error(err) => {
330 return Err(crate::error::AgentLoopError::llm_kind(
331 err.kind(),
332 err.to_string(),
333 ));
334 }
335 }
336 }
337
338 Ok(LlmResponse {
339 text,
340 thinking: if thinking.is_empty() {
341 None
342 } else {
343 Some(thinking)
344 },
345 thinking_signature,
346 tool_calls: if tool_calls.is_empty() {
347 None
348 } else {
349 Some(tool_calls)
350 },
351 metadata,
352 })
353 }
354
355 async fn list_models(&self) -> Result<Option<Vec<DiscoveredModel>>> {
363 Ok(None)
365 }
366
367 fn supports_compact(&self) -> bool {
376 false
378 }
379
380 fn supports_stateful_responses(&self) -> bool {
386 false
387 }
388
389 fn effective_context_window(&self, _model: &str) -> Option<usize> {
395 None
396 }
397
398 fn supports_parallel_tool_calls(&self, _model: &str) -> bool {
409 false
410 }
411
412 async fn compact(&self, _request: CompactRequest) -> Result<Option<CompactResponse>> {
432 Ok(None)
434 }
435}
436
437#[async_trait]
439impl ChatDriver for Box<dyn ChatDriver> {
440 async fn chat_completion_stream(
441 &self,
442 messages: Vec<LlmMessage>,
443 config: &LlmCallConfig,
444 ) -> Result<LlmResponseStream> {
445 (**self).chat_completion_stream(messages, config).await
446 }
447
448 async fn chat_completion(
449 &self,
450 messages: Vec<LlmMessage>,
451 config: &LlmCallConfig,
452 ) -> Result<LlmResponse> {
453 (**self).chat_completion(messages, config).await
454 }
455
456 async fn list_models(&self) -> Result<Option<Vec<DiscoveredModel>>> {
457 (**self).list_models().await
458 }
459
460 fn supports_compact(&self) -> bool {
461 (**self).supports_compact()
462 }
463
464 fn supports_stateful_responses(&self) -> bool {
465 (**self).supports_stateful_responses()
466 }
467
468 fn effective_context_window(&self, model: &str) -> Option<usize> {
469 (**self).effective_context_window(model)
470 }
471
472 fn supports_parallel_tool_calls(&self, model: &str) -> bool {
473 (**self).supports_parallel_tool_calls(model)
474 }
475
476 async fn compact(&self, request: CompactRequest) -> Result<Option<CompactResponse>> {
477 (**self).compact(request).await
478 }
479}
480
481#[derive(Debug, Clone)]
487pub struct LlmMessage {
488 pub role: LlmMessageRole,
489 pub content: LlmMessageContent,
490 pub tool_calls: Option<Vec<ToolCall>>,
491 pub tool_call_id: Option<String>,
492 pub phase: Option<crate::execution_phase::ExecutionPhase>,
497 pub thinking: Option<String>,
500 pub thinking_signature: Option<String>,
503}
504
505impl LlmMessage {
506 pub fn text(role: LlmMessageRole, content: impl Into<String>) -> Self {
508 Self {
509 role,
510 content: LlmMessageContent::Text(content.into()),
511 tool_calls: None,
512 tool_call_id: None,
513 phase: None,
514 thinking: None,
515 thinking_signature: None,
516 }
517 }
518
519 pub fn parts(role: LlmMessageRole, parts: Vec<LlmContentPart>) -> Self {
521 Self {
522 role,
523 content: LlmMessageContent::Parts(parts),
524 tool_calls: None,
525 tool_call_id: None,
526 phase: None,
527 thinking: None,
528 thinking_signature: None,
529 }
530 }
531
532 pub fn content_as_text(&self) -> String {
534 self.content.to_text()
535 }
536
537 pub fn prepend_text_prefix(&mut self, prefix: &str) {
542 match &mut self.content {
543 LlmMessageContent::Text(text) => {
544 *text = format!("{}{}", prefix, text);
545 }
546 LlmMessageContent::Parts(parts) => {
547 for part in parts.iter_mut() {
548 if let LlmContentPart::Text { text } = part {
549 *text = format!("{}{}", prefix, text);
550 return;
551 }
552 }
553 parts.insert(
555 0,
556 LlmContentPart::Text {
557 text: prefix.to_string(),
558 },
559 );
560 }
561 }
562 }
563}
564
565pub fn fold_system_messages(messages: &[LlmMessage]) -> Option<String> {
576 let mut system: Option<String> = None;
577 for msg in messages {
578 if msg.role == LlmMessageRole::System {
579 let text = msg.content.to_text();
580 system = Some(match system.take() {
581 Some(existing) if !existing.is_empty() => format!("{existing}\n\n{text}"),
582 _ => text,
583 });
584 }
585 }
586 system
587}
588
589#[derive(Debug, Clone)]
591pub enum LlmMessageContent {
592 Text(String),
594 Parts(Vec<LlmContentPart>),
596}
597
598impl LlmMessageContent {
599 pub fn to_text(&self) -> String {
601 match self {
602 LlmMessageContent::Text(s) => s.clone(),
603 LlmMessageContent::Parts(parts) => parts
604 .iter()
605 .filter_map(|p| match p {
606 LlmContentPart::Text { text } => Some(text.clone()),
607 _ => None,
608 })
609 .collect::<Vec<_>>()
610 .join(""),
611 }
612 }
613
614 pub fn is_text(&self) -> bool {
616 matches!(self, LlmMessageContent::Text(_))
617 }
618
619 pub fn is_parts(&self) -> bool {
621 matches!(self, LlmMessageContent::Parts(_))
622 }
623}
624
625impl From<String> for LlmMessageContent {
626 fn from(s: String) -> Self {
627 LlmMessageContent::Text(s)
628 }
629}
630
631impl From<&str> for LlmMessageContent {
632 fn from(s: &str) -> Self {
633 LlmMessageContent::Text(s.to_string())
634 }
635}
636
637#[derive(Debug, Clone)]
639pub enum LlmContentPart {
640 Text { text: String },
642 Image { url: String },
644 Audio { url: String },
646}
647
648impl LlmContentPart {
649 pub fn text(text: impl Into<String>) -> Self {
651 LlmContentPart::Text { text: text.into() }
652 }
653
654 pub fn image(url: impl Into<String>) -> Self {
656 LlmContentPart::Image { url: url.into() }
657 }
658
659 pub fn audio(url: impl Into<String>) -> Self {
661 LlmContentPart::Audio { url: url.into() }
662 }
663}
664
665#[derive(Debug, Clone, PartialEq, Eq)]
667pub enum LlmMessageRole {
668 System,
669 User,
670 Assistant,
671 Tool,
672}
673
674#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
684pub struct ToolSearchConfig {
685 pub enabled: bool,
687 pub threshold: usize,
690}
691
692#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
694#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
695#[serde(rename_all = "snake_case")]
696pub enum PromptCacheStrategy {
697 #[default]
699 Auto,
700}
701
702#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
707#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
708pub struct PromptCacheConfig {
709 pub enabled: bool,
711 #[serde(default)]
713 pub strategy: PromptCacheStrategy,
714 #[serde(default, skip_serializing_if = "Option::is_none")]
721 pub gemini_cached_content: Option<String>,
722}
723
724#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
734#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
735#[serde(tag = "kind", rename_all = "snake_case")]
736pub enum OpenRouterRoutingPreset {
737 CheapestWithTools,
739 LowestLatencyReview,
741 ZdrOnly,
743 ByokFirst,
745 NoDataCollection,
747 StrictJson,
749 ReasoningRequired,
751 MaxPrice {
754 #[serde(default, skip_serializing_if = "Option::is_none")]
756 prompt_usd_per_million: Option<f64>,
757 #[serde(default, skip_serializing_if = "Option::is_none")]
759 completion_usd_per_million: Option<f64>,
760 },
761}
762
763#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
771#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
772#[serde(rename_all = "snake_case")]
773pub enum OpenRouterCapacityStrategy {
774 #[default]
776 SharedCapacity,
777 ByokFirst,
781 ByokOnly,
786}
787
788#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
796#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
797#[serde(rename_all = "snake_case")]
798pub enum OpenRouterServerToolKind {
799 WebSearch,
800 WebFetch,
801 Datetime,
802 ImageGeneration,
803 ApplyPatch,
804 Fusion,
805 Advisor,
806 Subagent,
807}
808
809impl OpenRouterServerToolKind {
810 pub const ALL: [OpenRouterServerToolKind; 8] = [
812 Self::WebSearch,
813 Self::WebFetch,
814 Self::Datetime,
815 Self::ImageGeneration,
816 Self::ApplyPatch,
817 Self::Fusion,
818 Self::Advisor,
819 Self::Subagent,
820 ];
821
822 pub fn name(&self) -> &'static str {
824 match self {
825 Self::WebSearch => "web_search",
826 Self::WebFetch => "web_fetch",
827 Self::Datetime => "datetime",
828 Self::ImageGeneration => "image_generation",
829 Self::ApplyPatch => "apply_patch",
830 Self::Fusion => "fusion",
831 Self::Advisor => "advisor",
832 Self::Subagent => "subagent",
833 }
834 }
835
836 pub fn display_name(&self) -> &'static str {
838 match self {
839 Self::WebSearch => "Web Search",
840 Self::WebFetch => "Web Fetch",
841 Self::Datetime => "Date & Time",
842 Self::ImageGeneration => "Image Generation",
843 Self::ApplyPatch => "Apply Patch",
844 Self::Fusion => "Fusion",
845 Self::Advisor => "Advisor",
846 Self::Subagent => "Subagent",
847 }
848 }
849
850 pub fn wire_type(&self) -> String {
853 format!("openrouter:{}", self.name())
854 }
855
856 pub fn from_name(name: &str) -> Option<Self> {
858 Self::ALL.into_iter().find(|kind| kind.name() == name)
859 }
860}
861
862#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
866#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
867pub struct OpenRouterServerTool {
868 pub kind: OpenRouterServerToolKind,
869 #[serde(default, skip_serializing_if = "Option::is_none")]
870 #[cfg_attr(feature = "openapi", schema(value_type = Option<Object>))]
871 pub parameters: Option<serde_json::Value>,
872}
873
874impl OpenRouterServerTool {
875 pub fn new(kind: OpenRouterServerToolKind) -> Self {
877 Self {
878 kind,
879 parameters: None,
880 }
881 }
882
883 pub fn with_parameters(kind: OpenRouterServerToolKind, parameters: serde_json::Value) -> Self {
885 Self {
886 kind,
887 parameters: Some(parameters),
888 }
889 }
890}
891
892#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
895#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
896pub struct OpenRouterRoutingConfig {
897 #[serde(default, skip_serializing_if = "Vec::is_empty")]
899 pub models: Vec<String>,
900 #[serde(default, skip_serializing_if = "Option::is_none")]
903 pub route: Option<OpenRouterRoute>,
904 #[serde(default, skip_serializing_if = "Option::is_none")]
906 pub provider: Option<OpenRouterProviderRouting>,
907 #[serde(default, skip_serializing_if = "Option::is_none")]
909 pub plugins: Option<OpenRouterPluginConfig>,
910 #[serde(default, skip_serializing_if = "Option::is_none")]
914 pub capacity_strategy: Option<OpenRouterCapacityStrategy>,
915 #[serde(default, skip_serializing_if = "Vec::is_empty")]
919 pub presets: Vec<OpenRouterRoutingPreset>,
920 #[serde(default, skip_serializing_if = "Vec::is_empty")]
923 pub server_tools: Vec<OpenRouterServerTool>,
924}
925
926impl OpenRouterRoutingConfig {
927 pub fn is_empty(&self) -> bool {
928 self.models.is_empty()
929 && self.route.is_none()
930 && self.provider.is_none()
931 && self.plugins.as_ref().is_none_or(|p| p.is_empty())
932 && matches!(
933 self.capacity_strategy,
934 None | Some(OpenRouterCapacityStrategy::SharedCapacity)
935 )
936 && self.presets.is_empty()
937 && self.server_tools.is_empty()
938 }
939
940 pub fn fallback_models(models: impl IntoIterator<Item = impl Into<String>>) -> Self {
942 let models = models.into_iter().map(Into::into).collect::<Vec<_>>();
943 let route = (!models.is_empty()).then_some(OpenRouterRoute::Fallback);
944 Self {
945 models,
946 route,
947 provider: None,
948 plugins: None,
949 capacity_strategy: None,
950 presets: vec![],
951 server_tools: vec![],
952 }
953 }
954
955 pub fn validate_for_primary_model(
956 &self,
957 primary_model: &str,
958 ) -> std::result::Result<(), String> {
959 if self.route == Some(OpenRouterRoute::Fallback) && self.models.is_empty() {
960 return Err(
961 "OpenRouter fallback routing requires at least one model in `models`".to_string(),
962 );
963 }
964
965 if let Some(first_model) = self.models.first()
966 && first_model != primary_model
967 {
968 return Err(format!(
969 "OpenRouter routing models[0] ('{first_model}') must match primary model ('{primary_model}')"
970 ));
971 }
972
973 Ok(())
974 }
975
976 pub fn apply_capacity_strategy(&self) -> std::result::Result<Self, String> {
986 match self.capacity_strategy {
987 None | Some(OpenRouterCapacityStrategy::SharedCapacity) => Ok(self.clone()),
988 Some(OpenRouterCapacityStrategy::ByokFirst) => {
989 let mut result = self.clone();
990 let provider = result.provider.get_or_insert_with(Default::default);
991 if provider.allow_fallbacks.is_none() {
992 provider.allow_fallbacks = Some(true);
993 }
994 Ok(result)
995 }
996 Some(OpenRouterCapacityStrategy::ByokOnly) => {
997 let only_is_empty = self.provider.as_ref().is_none_or(|p| p.only.is_empty());
998 if only_is_empty {
999 return Err(
1000 "OpenRouter BYOK-only strategy requires provider.only to list at least \
1001 one upstream provider slug. Configure the provider list to match the \
1002 BYOK providers registered in your OpenRouter workspace."
1003 .to_string(),
1004 );
1005 }
1006 let mut result = self.clone();
1007 let provider = result.provider.get_or_insert_with(Default::default);
1008 provider.allow_fallbacks = Some(false);
1009 Ok(result)
1010 }
1011 }
1012 }
1013
1014 pub fn apply_presets(&self) -> std::result::Result<Self, String> {
1024 if self.presets.is_empty() {
1025 return Ok(self.clone());
1026 }
1027
1028 let mut derived = OpenRouterProviderRouting::default();
1029
1030 for preset in &self.presets {
1031 match preset {
1032 OpenRouterRoutingPreset::CheapestWithTools => {
1033 derived.require_parameters = Some(true);
1034 derived.sort = Some(OpenRouterProviderSort::Simple(
1035 OpenRouterProviderSortBy::Price,
1036 ));
1037 }
1038 OpenRouterRoutingPreset::LowestLatencyReview => {
1039 derived.sort = Some(OpenRouterProviderSort::Simple(
1040 OpenRouterProviderSortBy::Throughput,
1041 ));
1042 }
1043 OpenRouterRoutingPreset::ZdrOnly => {
1044 derived.zdr = Some(true);
1045 }
1046 OpenRouterRoutingPreset::ByokFirst => {
1047 if derived.allow_fallbacks.is_none() {
1048 derived.allow_fallbacks = Some(true);
1049 }
1050 }
1051 OpenRouterRoutingPreset::NoDataCollection => {
1052 derived.data_collection = Some(OpenRouterDataCollection::Deny);
1053 }
1054 OpenRouterRoutingPreset::StrictJson
1055 | OpenRouterRoutingPreset::ReasoningRequired => {
1056 derived.require_parameters = Some(true);
1057 }
1058 OpenRouterRoutingPreset::MaxPrice {
1059 prompt_usd_per_million,
1060 completion_usd_per_million,
1061 } => {
1062 if prompt_usd_per_million.is_some_and(|v| v < 0.0)
1063 || completion_usd_per_million.is_some_and(|v| v < 0.0)
1064 {
1065 return Err(
1066 "MaxPrice preset values must be non-negative USD per million tokens"
1067 .to_string(),
1068 );
1069 }
1070 if prompt_usd_per_million.is_some() || completion_usd_per_million.is_some() {
1071 let mp = derived.max_price.get_or_insert_with(Default::default);
1072 if let Some(p) = prompt_usd_per_million {
1073 mp.prompt = Some(p / 1_000_000.0);
1074 }
1075 if let Some(c) = completion_usd_per_million {
1076 mp.completion = Some(c / 1_000_000.0);
1077 }
1078 }
1079 }
1080 }
1081 }
1082
1083 let merged = merge_provider_routing(derived, self.provider.clone().unwrap_or_default());
1085
1086 let mut result = self.clone();
1087 result.presets = vec![];
1088 result.provider = if merged.is_empty() {
1089 None
1090 } else {
1091 Some(merged)
1092 };
1093 Ok(result)
1094 }
1095}
1096
1097fn merge_provider_routing(
1101 derived: OpenRouterProviderRouting,
1102 explicit: OpenRouterProviderRouting,
1103) -> OpenRouterProviderRouting {
1104 OpenRouterProviderRouting {
1105 order: if !explicit.order.is_empty() {
1106 explicit.order
1107 } else {
1108 derived.order
1109 },
1110 only: if !explicit.only.is_empty() {
1111 explicit.only
1112 } else {
1113 derived.only
1114 },
1115 ignore: if !explicit.ignore.is_empty() {
1116 explicit.ignore
1117 } else {
1118 derived.ignore
1119 },
1120 allow_fallbacks: explicit.allow_fallbacks.or(derived.allow_fallbacks),
1121 require_parameters: explicit.require_parameters.or(derived.require_parameters),
1122 data_collection: explicit.data_collection.or(derived.data_collection),
1123 zdr: explicit.zdr.or(derived.zdr),
1124 enforce_distillable_text: explicit
1125 .enforce_distillable_text
1126 .or(derived.enforce_distillable_text),
1127 quantizations: if !explicit.quantizations.is_empty() {
1128 explicit.quantizations
1129 } else {
1130 derived.quantizations
1131 },
1132 sort: explicit.sort.or(derived.sort),
1133 max_price: explicit.max_price.or(derived.max_price),
1134 }
1135}
1136
1137#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
1139#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1140#[serde(rename_all = "snake_case")]
1141pub enum OpenRouterRoute {
1142 Fallback,
1143}
1144
1145#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
1147#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1148pub struct OpenRouterProviderRouting {
1149 #[serde(default, skip_serializing_if = "Vec::is_empty")]
1151 pub order: Vec<String>,
1152 #[serde(default, skip_serializing_if = "Vec::is_empty")]
1154 pub only: Vec<String>,
1155 #[serde(default, skip_serializing_if = "Vec::is_empty")]
1157 pub ignore: Vec<String>,
1158 #[serde(default, skip_serializing_if = "Option::is_none")]
1160 pub allow_fallbacks: Option<bool>,
1161 #[serde(default, skip_serializing_if = "Option::is_none")]
1163 pub require_parameters: Option<bool>,
1164 #[serde(default, skip_serializing_if = "Option::is_none")]
1166 pub data_collection: Option<OpenRouterDataCollection>,
1167 #[serde(default, skip_serializing_if = "Option::is_none")]
1169 pub zdr: Option<bool>,
1170 #[serde(default, skip_serializing_if = "Option::is_none")]
1172 pub enforce_distillable_text: Option<bool>,
1173 #[serde(default, skip_serializing_if = "Vec::is_empty")]
1175 pub quantizations: Vec<String>,
1176 #[serde(default, skip_serializing_if = "Option::is_none")]
1178 pub sort: Option<OpenRouterProviderSort>,
1179 #[serde(default, skip_serializing_if = "Option::is_none")]
1181 pub max_price: Option<OpenRouterMaxPrice>,
1182}
1183
1184impl OpenRouterProviderRouting {
1185 pub fn is_empty(&self) -> bool {
1186 self.order.is_empty()
1187 && self.only.is_empty()
1188 && self.ignore.is_empty()
1189 && self.allow_fallbacks.is_none()
1190 && self.require_parameters.is_none()
1191 && self.data_collection.is_none()
1192 && self.zdr.is_none()
1193 && self.enforce_distillable_text.is_none()
1194 && self.quantizations.is_empty()
1195 && self.sort.is_none()
1196 && self.max_price.is_none()
1197 }
1198}
1199
1200#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
1202#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1203#[serde(rename_all = "snake_case")]
1204pub enum OpenRouterDataCollection {
1205 Allow,
1206 Deny,
1207}
1208
1209#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
1211#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1212#[serde(untagged)]
1213pub enum OpenRouterProviderSort {
1214 Simple(OpenRouterProviderSortBy),
1215 Advanced(OpenRouterProviderSortOptions),
1216}
1217
1218#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
1220#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1221#[serde(rename_all = "snake_case")]
1222pub enum OpenRouterProviderSortBy {
1223 Price,
1224 Throughput,
1225 Latency,
1226}
1227
1228#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
1230#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1231pub struct OpenRouterProviderSortOptions {
1232 pub by: OpenRouterProviderSortBy,
1233 #[serde(default, skip_serializing_if = "Option::is_none")]
1234 pub partition: Option<OpenRouterSortPartition>,
1235}
1236
1237#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
1239#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1240#[serde(rename_all = "snake_case")]
1241pub enum OpenRouterSortPartition {
1242 Model,
1243 None,
1244}
1245
1246#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
1249#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1250pub struct OpenRouterMaxPrice {
1251 #[serde(default, skip_serializing_if = "Option::is_none")]
1252 pub prompt: Option<f64>,
1253 #[serde(default, skip_serializing_if = "Option::is_none")]
1254 pub completion: Option<f64>,
1255 #[serde(default, skip_serializing_if = "Option::is_none")]
1256 pub request: Option<f64>,
1257 #[serde(default, skip_serializing_if = "Option::is_none")]
1258 pub image: Option<f64>,
1259}
1260
1261#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
1267#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1268pub struct OpenRouterWebSearchPlugin {
1269 #[serde(default, skip_serializing_if = "Option::is_none")]
1271 pub max_results: Option<u32>,
1272 #[serde(default, skip_serializing_if = "Option::is_none")]
1274 pub search_prompt: Option<String>,
1275}
1276
1277#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
1282#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1283pub struct OpenRouterFilePlugin {}
1284
1285#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
1290#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
1291pub struct OpenRouterPluginConfig {
1292 #[serde(default, skip_serializing_if = "Option::is_none")]
1294 pub web: Option<OpenRouterWebSearchPlugin>,
1295 #[serde(default, skip_serializing_if = "Option::is_none")]
1297 pub file: Option<OpenRouterFilePlugin>,
1298}
1299
1300impl OpenRouterPluginConfig {
1301 pub fn is_empty(&self) -> bool {
1302 self.web.is_none() && self.file.is_none()
1303 }
1304}
1305
1306pub const OPENROUTER_HTTP_REFERER_METADATA_KEY: &str = "openrouter.http_referer";
1308pub const OPENROUTER_X_TITLE_METADATA_KEY: &str = "openrouter.x_title";
1310
1311#[derive(Debug, Clone)]
1313pub struct LlmCallConfig {
1314 pub model: String,
1315 pub temperature: Option<f32>,
1316 pub max_tokens: Option<u32>,
1317 pub tools: Vec<ToolDefinition>,
1318 pub reasoning_effort: Option<String>,
1320 pub speed: Option<String>,
1324 pub verbosity: Option<String>,
1328 pub metadata: HashMap<String, String>,
1332 pub previous_response_id: Option<String>,
1335 pub provider_opaque_context: Option<ProviderOpaqueContext>,
1341 pub tool_search: Option<ToolSearchConfig>,
1343 pub prompt_cache: Option<PromptCacheConfig>,
1345 pub openrouter_routing: Option<OpenRouterRoutingConfig>,
1347 pub parallel_tool_calls: Option<bool>,
1354 pub volatile_suffix_len: usize,
1364}
1365
1366impl LlmCallConfig {
1367 pub fn resolved_parallel_tool_calls(&self, supported: bool) -> Option<bool> {
1377 if supported {
1378 self.parallel_tool_calls
1379 } else {
1380 None
1381 }
1382 }
1383}
1384
1385#[derive(Debug, Clone)]
1390pub struct LlmResponse {
1391 pub text: String,
1392 pub thinking: Option<String>,
1394 pub thinking_signature: Option<String>,
1396 pub tool_calls: Option<Vec<ToolCall>>,
1397 pub metadata: LlmCompletionMetadata,
1398}
1399
1400pub struct LlmCallConfigBuilder {
1406 config: LlmCallConfig,
1407}
1408
1409impl LlmCallConfigBuilder {
1410 pub fn from_config(config: LlmCallConfig) -> Self {
1412 Self { config }
1413 }
1414
1415 pub fn reasoning_effort(mut self, effort: impl Into<String>) -> Self {
1417 self.config.reasoning_effort = Some(effort.into());
1418 self
1419 }
1420
1421 pub fn speed(mut self, speed: impl Into<String>) -> Self {
1423 self.config.speed = Some(speed.into());
1424 self
1425 }
1426
1427 pub fn verbosity(mut self, verbosity: impl Into<String>) -> Self {
1429 self.config.verbosity = Some(verbosity.into());
1430 self
1431 }
1432
1433 pub fn model(mut self, model: impl Into<String>) -> Self {
1435 self.config.model = model.into();
1436 self
1437 }
1438
1439 pub fn temperature(mut self, temp: f32) -> Self {
1441 self.config.temperature = Some(temp);
1442 self
1443 }
1444
1445 pub fn max_tokens(mut self, tokens: u32) -> Self {
1447 self.config.max_tokens = Some(tokens);
1448 self
1449 }
1450
1451 pub fn tools(mut self, tools: Vec<ToolDefinition>) -> Self {
1453 self.config.tools = tools;
1454 self
1455 }
1456
1457 pub fn metadata(mut self, metadata: HashMap<String, String>) -> Self {
1462 self.config.metadata = metadata;
1463 self
1464 }
1465
1466 pub fn with_metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
1468 self.config.metadata.insert(key.into(), value.into());
1469 self
1470 }
1471
1472 pub fn previous_response_id(mut self, id: Option<String>) -> Self {
1474 self.config.previous_response_id = id;
1475 self
1476 }
1477
1478 pub fn provider_opaque_context(mut self, context: Option<ProviderOpaqueContext>) -> Self {
1480 self.config.provider_opaque_context = context;
1481 self
1482 }
1483
1484 pub fn tool_search(mut self, config: ToolSearchConfig) -> Self {
1486 self.config.tool_search = Some(config);
1487 self
1488 }
1489
1490 pub fn prompt_cache(mut self, config: PromptCacheConfig) -> Self {
1492 self.config.prompt_cache = Some(config);
1493 self
1494 }
1495
1496 pub fn openrouter_routing(mut self, config: OpenRouterRoutingConfig) -> Self {
1498 self.config.openrouter_routing = (!config.is_empty()).then_some(config);
1499 self
1500 }
1501
1502 pub fn parallel_tool_calls(mut self, parallel_tool_calls: Option<bool>) -> Self {
1504 self.config.parallel_tool_calls = parallel_tool_calls;
1505 self
1506 }
1507
1508 pub fn volatile_suffix_len(mut self, len: usize) -> Self {
1512 self.config.volatile_suffix_len = len;
1513 self
1514 }
1515
1516 pub fn build(self) -> LlmCallConfig {
1518 self.config
1519 }
1520}
1521
1522pub use crate::provider::DriverId;
1531
1532#[derive(Debug, Clone, Default, PartialEq, Eq)]
1538pub struct ProviderMetadata {
1539 pub refresh_token: Option<String>,
1541 pub account_id: Option<String>,
1543 pub extra: Option<serde_json::Value>,
1545}
1546
1547#[derive(Debug, Clone)]
1549pub struct ProviderConfig {
1550 pub provider_type: DriverId,
1552 pub api_key: Option<String>,
1554 pub base_url: Option<String>,
1556 pub metadata: ProviderMetadata,
1558}
1559
1560impl ProviderConfig {
1561 pub fn new(provider_type: DriverId) -> Self {
1563 Self {
1564 provider_type,
1565 api_key: None,
1566 base_url: None,
1567 metadata: ProviderMetadata::default(),
1568 }
1569 }
1570
1571 pub fn with_api_key(mut self, api_key: impl Into<String>) -> Self {
1573 self.api_key = Some(api_key.into());
1574 self
1575 }
1576
1577 pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
1579 self.base_url = Some(base_url.into());
1580 self
1581 }
1582
1583 pub fn with_metadata(mut self, metadata: ProviderMetadata) -> Self {
1585 self.metadata = metadata;
1586 self
1587 }
1588}
1589
1590#[derive(Debug, Clone)]
1596pub struct DriverConfig {
1597 pub provider_type: DriverId,
1599 pub api_key: Option<String>,
1605 pub credentials: std::collections::BTreeMap<String, String>,
1611 pub base_url: Option<String>,
1613 pub metadata: ProviderMetadata,
1615}
1616
1617impl DriverConfig {
1618 pub fn from_provider_config(config: &ProviderConfig) -> Self {
1624 Self {
1625 provider_type: config.provider_type.clone(),
1626 credentials: crate::credential_schema::parse_credential_document(
1627 config.api_key.as_deref(),
1628 ),
1629 api_key: config.api_key.clone(),
1630 base_url: config.base_url.clone(),
1631 metadata: config.metadata.clone(),
1632 }
1633 }
1634
1635 pub fn credential(&self, name: &str) -> Option<&str> {
1637 self.credentials
1638 .get(name)
1639 .map(String::as_str)
1640 .filter(|s| !s.is_empty())
1641 }
1642}
1643
1644pub type BoxedChatDriver = Box<dyn ChatDriver>;
1649
1650#[derive(Debug, Clone)]
1656pub struct EmbedRequest {
1657 pub texts: Vec<String>,
1659 pub model: String,
1661}
1662
1663#[derive(Debug, Clone)]
1665pub struct EmbedResponse {
1666 pub embeddings: Vec<Vec<f32>>,
1668 pub usage_tokens: Option<u32>,
1671}
1672
1673#[derive(Debug, thiserror::Error)]
1675pub enum EmbeddingsDriverError {
1676 #[error("embeddings provider returned an error: {0}")]
1677 Provider(String),
1678 #[error("embeddings request failed: {0}")]
1679 Transport(String),
1680}
1681
1682#[async_trait]
1688pub trait EmbeddingsDriver: Send + Sync {
1689 async fn embed(
1691 &self,
1692 request: EmbedRequest,
1693 ) -> std::result::Result<EmbedResponse, EmbeddingsDriverError>;
1694}
1695
1696pub type BoxedEmbeddingsDriver = Box<dyn EmbeddingsDriver>;
1698
1699pub type EmbeddingsDriverFactory =
1701 Arc<dyn Fn(&DriverConfig) -> BoxedEmbeddingsDriver + Send + Sync>;
1702
1703pub type DriverFactory = Arc<dyn Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync>;
1712
1713#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
1719#[serde(rename_all = "snake_case")]
1720pub enum ServiceKind {
1721 Chat,
1723 Embeddings,
1725 Realtime,
1727 Images,
1729 Rerank,
1731}
1732
1733impl std::fmt::Display for ServiceKind {
1734 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1735 let s = match self {
1736 ServiceKind::Chat => "chat",
1737 ServiceKind::Embeddings => "embeddings",
1738 ServiceKind::Realtime => "realtime",
1739 ServiceKind::Images => "images",
1740 ServiceKind::Rerank => "rerank",
1741 };
1742 f.write_str(s)
1743 }
1744}
1745
1746#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1759pub enum DriverOAuthFlow {
1760 OpenRouterPkce,
1768}
1769
1770#[derive(Debug, Clone)]
1775pub struct DriverOAuthConfig {
1776 pub authorize_url: String,
1778 pub token_url: String,
1780 pub flow: DriverOAuthFlow,
1782}
1783
1784impl DriverOAuthConfig {
1785 pub fn openrouter() -> Self {
1787 Self {
1788 authorize_url: "https://openrouter.ai/auth".to_string(),
1789 token_url: "https://openrouter.ai/api/v1/auth/keys".to_string(),
1790 flow: DriverOAuthFlow::OpenRouterPkce,
1791 }
1792 }
1793}
1794
1795#[derive(Clone)]
1802pub struct DriverDescriptor {
1803 pub id: DriverId,
1805 pub display_name: String,
1807 pub services: Vec<ServiceKind>,
1809 pub credential_schema: CredentialFormSchema,
1811 pub oauth: Option<DriverOAuthConfig>,
1814 pub chat: Option<DriverFactory>,
1816 pub embeddings: Option<EmbeddingsDriverFactory>,
1818}
1819
1820impl DriverDescriptor {
1821 pub fn chat_only<F>(id: impl Into<DriverId>, factory: F) -> Self
1826 where
1827 F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
1828 {
1829 let id = id.into();
1830 Self {
1831 display_name: default_display_name(&id),
1832 credential_schema: default_credential_schema(&id),
1833 services: vec![ServiceKind::Chat],
1834 oauth: None,
1835 chat: Some(Arc::new(factory)),
1836 embeddings: None,
1837 id,
1838 }
1839 }
1840
1841 pub fn supports(&self, service: ServiceKind) -> bool {
1843 self.services.contains(&service)
1844 }
1845}
1846
1847impl std::fmt::Debug for DriverDescriptor {
1848 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1849 f.debug_struct("DriverDescriptor")
1850 .field("id", &self.id)
1851 .field("display_name", &self.display_name)
1852 .field("services", &self.services)
1853 .field("oauth", &self.oauth.is_some())
1854 .field("chat", &self.chat.is_some())
1855 .field("embeddings", &self.embeddings.is_some())
1856 .finish()
1857 }
1858}
1859
1860fn default_display_name(id: &DriverId) -> String {
1861 match id {
1862 DriverId::OpenAI => "OpenAI".to_string(),
1863 DriverId::OpenRouter => "OpenRouter".to_string(),
1864 DriverId::AzureOpenAI => "Azure OpenAI".to_string(),
1865 DriverId::OpenAICompletions => "OpenAI (Chat Completions)".to_string(),
1866 DriverId::Anthropic => "Anthropic".to_string(),
1867 DriverId::Gemini => "Google Gemini".to_string(),
1868 DriverId::Bedrock => "AWS Bedrock".to_string(),
1869 DriverId::Mai => "Microsoft MAI".to_string(),
1870 DriverId::Fireworks => "Fireworks AI".to_string(),
1871 DriverId::Meta => "Meta Model API".to_string(),
1872 DriverId::LlmSim => "LLM Simulator".to_string(),
1873 DriverId::External(id) => id.to_string(),
1874 }
1875}
1876
1877fn default_credential_schema(id: &DriverId) -> CredentialFormSchema {
1878 match id {
1879 DriverId::LlmSim | DriverId::External(_) => CredentialFormSchema::empty(),
1881 _ => CredentialFormSchema::api_key(String::new()),
1882 }
1883}
1884
1885#[derive(Clone, Default)]
1905pub struct DriverRegistry {
1906 descriptors: HashMap<DriverId, DriverDescriptor>,
1907}
1908
1909impl DriverRegistry {
1910 pub fn new() -> Self {
1912 Self {
1913 descriptors: HashMap::new(),
1914 }
1915 }
1916
1917 pub fn register_descriptor(&mut self, descriptor: DriverDescriptor) {
1923 if self.descriptors.contains_key(&descriptor.id) {
1924 panic!(
1925 "driver already registered for provider '{}'; \
1926 use register_descriptor_or_replace to overwrite intentionally",
1927 descriptor.id
1928 );
1929 }
1930 self.descriptors.insert(descriptor.id.clone(), descriptor);
1931 }
1932
1933 pub fn register_descriptor_or_replace(&mut self, descriptor: DriverDescriptor) {
1935 self.descriptors.insert(descriptor.id.clone(), descriptor);
1936 }
1937
1938 pub fn register<F>(&mut self, provider_type: impl Into<DriverId>, factory: F)
1944 where
1945 F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
1946 {
1947 self.register_descriptor(DriverDescriptor::chat_only(provider_type, factory));
1948 }
1949
1950 pub fn register_or_replace<F>(&mut self, provider_type: impl Into<DriverId>, factory: F)
1955 where
1956 F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
1957 {
1958 self.register_descriptor_or_replace(DriverDescriptor::chat_only(provider_type, factory));
1959 }
1960
1961 pub fn register_external<F>(&mut self, id: impl Into<Arc<str>>, factory: F)
1966 where
1967 F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
1968 {
1969 self.register(DriverId::external(id), factory);
1970 }
1971
1972 pub fn create_chat_driver(&self, config: &ProviderConfig) -> Result<BoxedChatDriver> {
1981 let requires_api_key = !matches!(
1985 config.provider_type,
1986 DriverId::LlmSim | DriverId::External(_) | DriverId::Mai
1987 );
1988 if requires_api_key && config.api_key.is_none() {
1989 return Err(AgentLoopError::llm(
1990 "API key is required. Configure the API key in provider settings.",
1991 ));
1992 }
1993
1994 let descriptor = self.descriptors.get(&config.provider_type).ok_or_else(|| {
1996 AgentLoopError::driver_not_registered(config.provider_type.to_string())
1997 })?;
1998 let factory = descriptor.chat.as_ref().ok_or_else(|| {
1999 AgentLoopError::llm(format!(
2000 "Provider driver '{}' does not implement the chat service.",
2001 config.provider_type
2002 ))
2003 })?;
2004
2005 let driver_config = DriverConfig::from_provider_config(config);
2007 Ok(factory(&driver_config))
2008 }
2009
2010 pub fn has_driver(&self, provider_type: &DriverId) -> bool {
2012 self.descriptors.contains_key(provider_type)
2013 }
2014
2015 pub fn descriptor(&self, provider_type: &DriverId) -> Option<&DriverDescriptor> {
2017 self.descriptors.get(provider_type)
2018 }
2019
2020 pub fn supports(&self, provider_type: &DriverId, service: ServiceKind) -> bool {
2022 self.descriptors
2023 .get(provider_type)
2024 .is_some_and(|d| d.supports(service))
2025 }
2026
2027 pub fn providers_for(&self, service: ServiceKind) -> Vec<DriverId> {
2029 self.descriptors
2030 .values()
2031 .filter(|d| d.supports(service))
2032 .map(|d| d.id.clone())
2033 .collect()
2034 }
2035
2036 pub fn registered_providers(&self) -> Vec<DriverId> {
2038 self.descriptors.keys().cloned().collect()
2039 }
2040
2041 pub fn create_embeddings_driver(
2049 &self,
2050 config: &ProviderConfig,
2051 ) -> std::result::Result<BoxedEmbeddingsDriver, EmbeddingsDriverError> {
2052 let requires_api_key = !matches!(
2053 config.provider_type,
2054 DriverId::LlmSim | DriverId::External(_)
2055 );
2056 if requires_api_key && config.api_key.is_none() {
2057 return Err(EmbeddingsDriverError::Provider(
2058 "API key is required. Configure the API key in provider settings.".to_string(),
2059 ));
2060 }
2061 let descriptor = self.descriptors.get(&config.provider_type).ok_or_else(|| {
2062 EmbeddingsDriverError::Provider(format!(
2063 "No driver registered for provider '{}'",
2064 config.provider_type
2065 ))
2066 })?;
2067 let factory = descriptor.embeddings.as_ref().ok_or_else(|| {
2068 EmbeddingsDriverError::Provider(format!(
2069 "Provider driver '{}' does not implement the embeddings service.",
2070 config.provider_type
2071 ))
2072 })?;
2073 let driver_config = DriverConfig::from_provider_config(config);
2074 Ok(factory(&driver_config))
2075 }
2076}
2077
2078const MAX_TOOL_RESULT_BYTES: usize = 64 * 1024;
2083
2084const TRUNCATION_SUFFIX: &str =
2085 "\n\n[Output truncated — exceeded 64 KiB limit. Try quiet flags, pipes, or redirect to file.]";
2086
2087pub fn truncate_tool_result(text: String) -> String {
2088 if text.len() <= MAX_TOOL_RESULT_BYTES {
2089 return text;
2090 }
2091 let content_budget = MAX_TOOL_RESULT_BYTES.saturating_sub(TRUNCATION_SUFFIX.len());
2092 let mut end = content_budget;
2093 while end > 0 && !text.is_char_boundary(end) {
2094 end -= 1;
2095 }
2096 let mut truncated = text[..end].to_string();
2097 truncated.push_str(TRUNCATION_SUFFIX);
2098 truncated
2099}
2100
2101#[cfg(test)]
2106mod tests {
2107 use super::*;
2108
2109 #[test]
2110 fn test_disjoint_prompt_tokens_subtracts_cached_subset() {
2111 assert_eq!(disjoint_prompt_tokens(1000, Some(800)), 200);
2114 assert_eq!(disjoint_prompt_tokens(1000, None), 1000);
2116 assert_eq!(disjoint_prompt_tokens(1000, Some(0)), 1000);
2117 assert_eq!(disjoint_prompt_tokens(800, Some(1000)), 0);
2119 }
2120
2121 #[test]
2122 fn test_chat_driver_defaults_are_conservative_and_boxed_capabilities_forward() {
2123 struct DefaultDriver;
2125 #[async_trait]
2126 impl ChatDriver for DefaultDriver {
2127 async fn chat_completion_stream(
2128 &self,
2129 _messages: Vec<LlmMessage>,
2130 _config: &LlmCallConfig,
2131 ) -> Result<LlmResponseStream> {
2132 unreachable!()
2133 }
2134 }
2135 assert!(!DefaultDriver.supports_parallel_tool_calls("any-model"));
2136 assert!(!DefaultDriver.supports_stateful_responses());
2137
2138 struct StatefulDriver;
2139 #[async_trait]
2140 impl ChatDriver for StatefulDriver {
2141 async fn chat_completion_stream(
2142 &self,
2143 _messages: Vec<LlmMessage>,
2144 _config: &LlmCallConfig,
2145 ) -> Result<LlmResponseStream> {
2146 unreachable!()
2147 }
2148
2149 fn supports_stateful_responses(&self) -> bool {
2150 true
2151 }
2152 }
2153 let boxed: BoxedChatDriver = Box::new(StatefulDriver);
2154 assert!(boxed.supports_stateful_responses());
2155 }
2156
2157 #[test]
2158 fn test_fold_system_messages_none_when_absent() {
2159 let messages = vec![
2160 LlmMessage::text(LlmMessageRole::User, "hi"),
2161 LlmMessage::text(LlmMessageRole::Assistant, "ok"),
2162 ];
2163 assert_eq!(fold_system_messages(&messages), None);
2164 }
2165
2166 #[test]
2167 fn test_fold_system_messages_single() {
2168 let messages = vec![
2169 LlmMessage::text(LlmMessageRole::System, "AGENT-PROMPT"),
2170 LlmMessage::text(LlmMessageRole::User, "hi"),
2171 ];
2172 assert_eq!(
2173 fold_system_messages(&messages),
2174 Some("AGENT-PROMPT".to_string())
2175 );
2176 }
2177
2178 #[test]
2179 fn test_fold_system_messages_accumulates_in_order() {
2180 let messages = vec![
2184 LlmMessage::text(LlmMessageRole::System, "A"),
2185 LlmMessage::text(LlmMessageRole::User, "hi"),
2186 LlmMessage::text(LlmMessageRole::Assistant, "ok"),
2187 LlmMessage::text(LlmMessageRole::System, "B"),
2188 ];
2189 assert_eq!(fold_system_messages(&messages), Some("A\n\nB".to_string()));
2190 }
2191
2192 #[test]
2193 fn test_fold_system_messages_concatenates_parts() {
2194 let messages = vec![LlmMessage::parts(
2195 LlmMessageRole::System,
2196 vec![
2197 LlmContentPart::text("foo"),
2198 LlmContentPart::image("data:image/png;base64,xxx"),
2199 LlmContentPart::text("bar"),
2200 ],
2201 )];
2202 assert_eq!(fold_system_messages(&messages), Some("foobar".to_string()));
2203 }
2204
2205 #[test]
2206 fn test_openrouter_fallback_models_empty_is_empty() {
2207 let routing = OpenRouterRoutingConfig::fallback_models(std::iter::empty::<String>());
2208
2209 assert!(routing.is_empty());
2210 assert_eq!(routing.route, None);
2211 }
2212
2213 #[test]
2214 fn test_openrouter_routing_validates_primary_model() {
2215 let routing = OpenRouterRoutingConfig::fallback_models([
2216 "openai/gpt-5-mini",
2217 "anthropic/claude-sonnet-4.5",
2218 ]);
2219
2220 assert!(
2221 routing
2222 .validate_for_primary_model("openai/gpt-5-mini")
2223 .is_ok()
2224 );
2225 let err = routing
2226 .validate_for_primary_model("anthropic/claude-sonnet-4.5")
2227 .unwrap_err();
2228 assert!(err.contains("models[0]"));
2229 }
2230
2231 #[test]
2232 fn test_openrouter_routing_rejects_fallback_without_models() {
2233 let routing = OpenRouterRoutingConfig {
2234 route: Some(OpenRouterRoute::Fallback),
2235 ..Default::default()
2236 };
2237
2238 let err = routing
2239 .validate_for_primary_model("openai/gpt-5-mini")
2240 .unwrap_err();
2241 assert!(err.contains("requires at least one model"));
2242 }
2243
2244 #[test]
2245 fn test_openrouter_routing_serializes_request_fields() {
2246 let routing = OpenRouterRoutingConfig {
2247 models: vec![
2248 "openai/gpt-5-mini".to_string(),
2249 "anthropic/claude-sonnet-4.5".to_string(),
2250 ],
2251 route: Some(OpenRouterRoute::Fallback),
2252 provider: Some(OpenRouterProviderRouting {
2253 order: vec!["anthropic".to_string(), "openai".to_string()],
2254 allow_fallbacks: Some(false),
2255 require_parameters: Some(true),
2256 data_collection: Some(OpenRouterDataCollection::Deny),
2257 zdr: Some(true),
2258 sort: Some(OpenRouterProviderSort::Advanced(
2259 OpenRouterProviderSortOptions {
2260 by: OpenRouterProviderSortBy::Throughput,
2261 partition: Some(OpenRouterSortPartition::None),
2262 },
2263 )),
2264 max_price: Some(OpenRouterMaxPrice {
2265 prompt: Some(1.0),
2266 completion: Some(2.0),
2267 ..Default::default()
2268 }),
2269 ..Default::default()
2270 }),
2271 ..Default::default()
2272 };
2273
2274 let json = serde_json::to_value(routing).unwrap();
2275
2276 assert_eq!(
2277 json,
2278 serde_json::json!({
2279 "models": [
2280 "openai/gpt-5-mini",
2281 "anthropic/claude-sonnet-4.5"
2282 ],
2283 "route": "fallback",
2284 "provider": {
2285 "order": ["anthropic", "openai"],
2286 "allow_fallbacks": false,
2287 "require_parameters": true,
2288 "data_collection": "deny",
2289 "zdr": true,
2290 "sort": {
2291 "by": "throughput",
2292 "partition": "none"
2293 },
2294 "max_price": {
2295 "prompt": 1.0,
2296 "completion": 2.0
2297 }
2298 }
2299 })
2300 );
2301 }
2302
2303 #[test]
2304 fn test_provider_type_parsing() {
2305 assert_eq!("openai".parse::<DriverId>().unwrap(), DriverId::OpenAI);
2306 assert_eq!(
2307 "openrouter".parse::<DriverId>().unwrap(),
2308 DriverId::OpenRouter
2309 );
2310 assert_eq!(
2311 "openai_completions".parse::<DriverId>().unwrap(),
2312 DriverId::OpenAICompletions
2313 );
2314 assert_eq!(
2315 "azure_openai".parse::<DriverId>().unwrap(),
2316 DriverId::AzureOpenAI
2317 );
2318 assert_eq!(
2319 "anthropic".parse::<DriverId>().unwrap(),
2320 DriverId::Anthropic
2321 );
2322 assert_eq!("gemini".parse::<DriverId>().unwrap(), DriverId::Gemini);
2323 assert_eq!(
2325 "ollama".parse::<DriverId>().unwrap(),
2326 DriverId::external("ollama")
2327 );
2328 assert_eq!(
2329 "custom".parse::<DriverId>().unwrap(),
2330 DriverId::external("custom")
2331 );
2332 }
2333
2334 #[test]
2335 fn test_external_provider_id_is_case_insensitive() {
2336 assert_eq!("OpenAI".parse::<DriverId>().unwrap(), DriverId::OpenAI);
2339 assert_eq!(
2340 "Ollama".parse::<DriverId>().unwrap(),
2341 "ollama".parse::<DriverId>().unwrap()
2342 );
2343 assert_eq!(DriverId::external("OpenAI-Codex").as_str(), "openai-codex");
2344 assert_eq!(
2346 DriverId::external("MyProvider"),
2347 "myprovider".parse::<DriverId>().unwrap()
2348 );
2349 }
2350
2351 #[test]
2352 fn test_provider_type_display() {
2353 assert_eq!(DriverId::OpenAI.to_string(), "openai");
2354 assert_eq!(DriverId::OpenRouter.to_string(), "openrouter");
2355 assert_eq!(DriverId::AzureOpenAI.to_string(), "azure_openai");
2356 assert_eq!(
2357 DriverId::OpenAICompletions.to_string(),
2358 "openai_completions"
2359 );
2360 assert_eq!(DriverId::Anthropic.to_string(), "anthropic");
2361 assert_eq!(DriverId::Gemini.to_string(), "gemini");
2362 }
2363
2364 #[test]
2365 fn test_provider_config_builder() {
2366 let config = ProviderConfig::new(DriverId::Anthropic)
2367 .with_api_key("test-key")
2368 .with_base_url("https://custom.api.com");
2369
2370 assert_eq!(config.provider_type, DriverId::Anthropic);
2371 assert_eq!(config.api_key, Some("test-key".to_string()));
2372 assert_eq!(config.base_url, Some("https://custom.api.com".to_string()));
2373 }
2374
2375 #[test]
2376 fn test_driver_registry_requires_api_key() {
2377 let mut registry = DriverRegistry::new();
2379 registry.register(DriverId::OpenAI, |_config| {
2380 struct MockDriver;
2382 #[async_trait]
2383 impl ChatDriver for MockDriver {
2384 async fn chat_completion_stream(
2385 &self,
2386 _messages: Vec<LlmMessage>,
2387 _config: &LlmCallConfig,
2388 ) -> Result<LlmResponseStream> {
2389 unimplemented!()
2390 }
2391 }
2392 Box::new(MockDriver)
2393 });
2394
2395 let config = ProviderConfig::new(DriverId::OpenAI);
2397 let result = registry.create_chat_driver(&config);
2398 assert!(result.is_err());
2399
2400 let config_with_key = ProviderConfig::new(DriverId::OpenAI).with_api_key("test-key");
2402 let result = registry.create_chat_driver(&config_with_key);
2403 assert!(result.is_ok());
2404 }
2405
2406 #[test]
2407 fn test_driver_registry_returns_error_for_unregistered_provider() {
2408 let registry = DriverRegistry::new();
2409 let config = ProviderConfig::new(DriverId::Anthropic).with_api_key("test-key");
2410
2411 let result = registry.create_chat_driver(&config);
2412
2413 if let Err(AgentLoopError::DriverNotRegistered(provider)) = result {
2415 assert_eq!(provider, "anthropic");
2416 } else {
2417 panic!("Expected DriverNotRegistered error");
2418 }
2419 }
2420
2421 #[test]
2422 fn test_driver_registry_registration() {
2423 let mut registry = DriverRegistry::new();
2424
2425 assert!(!registry.has_driver(&DriverId::OpenAI));
2426 assert!(!registry.has_driver(&DriverId::Anthropic));
2427
2428 registry.register(DriverId::OpenAI, |_config| {
2429 struct MockDriver;
2430 #[async_trait]
2431 impl ChatDriver for MockDriver {
2432 async fn chat_completion_stream(
2433 &self,
2434 _messages: Vec<LlmMessage>,
2435 _config: &LlmCallConfig,
2436 ) -> Result<LlmResponseStream> {
2437 unimplemented!()
2438 }
2439 }
2440 Box::new(MockDriver)
2441 });
2442
2443 assert!(registry.has_driver(&DriverId::OpenAI));
2444 assert!(!registry.has_driver(&DriverId::Anthropic));
2445 }
2446
2447 #[test]
2448 fn test_register_external_and_create_driver_without_api_key() {
2449 struct MockDriver;
2450 #[async_trait]
2451 impl ChatDriver for MockDriver {
2452 async fn chat_completion_stream(
2453 &self,
2454 _messages: Vec<LlmMessage>,
2455 _config: &LlmCallConfig,
2456 ) -> Result<LlmResponseStream> {
2457 unimplemented!()
2458 }
2459 }
2460
2461 let mut registry = DriverRegistry::new();
2462 registry.register_external("openai-codex", |config| {
2463 assert_eq!(config.provider_type, DriverId::external("openai-codex"));
2465 Box::new(MockDriver)
2466 });
2467
2468 assert!(registry.has_driver(&DriverId::external("openai-codex")));
2469
2470 let config = ProviderConfig::new(DriverId::external("openai-codex")).with_metadata(
2472 ProviderMetadata {
2473 refresh_token: Some("rt".into()),
2474 ..Default::default()
2475 },
2476 );
2477 assert!(registry.create_chat_driver(&config).is_ok());
2478 }
2479
2480 #[test]
2481 fn test_register_defaults_to_chat_only_descriptor() {
2482 struct MockDriver;
2483 #[async_trait]
2484 impl ChatDriver for MockDriver {
2485 async fn chat_completion_stream(
2486 &self,
2487 _messages: Vec<LlmMessage>,
2488 _config: &LlmCallConfig,
2489 ) -> Result<LlmResponseStream> {
2490 unimplemented!()
2491 }
2492 }
2493
2494 let mut registry = DriverRegistry::new();
2495 registry.register(DriverId::Anthropic, |_config| Box::new(MockDriver));
2496
2497 let descriptor = registry.descriptor(&DriverId::Anthropic).unwrap();
2498 assert_eq!(descriptor.display_name, "Anthropic");
2499 assert_eq!(descriptor.services, vec![ServiceKind::Chat]);
2500 assert!(descriptor.chat.is_some());
2501 assert_eq!(descriptor.credential_schema.fields.len(), 1);
2503 assert_eq!(descriptor.credential_schema.fields[0].name, "api_key");
2504 assert!(descriptor.credential_schema.fields[0].required);
2505
2506 registry.register(DriverId::LlmSim, |_config| Box::new(MockDriver));
2508 let sim = registry.descriptor(&DriverId::LlmSim).unwrap();
2509 assert!(sim.credential_schema.fields.is_empty());
2510 }
2511
2512 #[test]
2513 fn test_descriptor_services_and_lookup() {
2514 struct MockDriver;
2515 #[async_trait]
2516 impl ChatDriver for MockDriver {
2517 async fn chat_completion_stream(
2518 &self,
2519 _messages: Vec<LlmMessage>,
2520 _config: &LlmCallConfig,
2521 ) -> Result<LlmResponseStream> {
2522 unimplemented!()
2523 }
2524 }
2525
2526 let mut registry = DriverRegistry::new();
2527 registry.register_descriptor(DriverDescriptor {
2528 services: vec![ServiceKind::Chat, ServiceKind::Realtime],
2529 ..DriverDescriptor::chat_only(DriverId::OpenAI, |_config| Box::new(MockDriver))
2530 });
2531 registry.register(DriverId::Anthropic, |_config| Box::new(MockDriver));
2532
2533 assert!(registry.supports(&DriverId::OpenAI, ServiceKind::Chat));
2534 assert!(registry.supports(&DriverId::OpenAI, ServiceKind::Realtime));
2535 assert!(!registry.supports(&DriverId::Anthropic, ServiceKind::Realtime));
2536 assert!(!registry.supports(&DriverId::Gemini, ServiceKind::Chat));
2537
2538 let realtime = registry.providers_for(ServiceKind::Realtime);
2539 assert_eq!(realtime, vec![DriverId::OpenAI]);
2540 let mut chat = registry.providers_for(ServiceKind::Chat);
2541 chat.sort_by_key(|p| p.to_string());
2542 assert_eq!(chat, vec![DriverId::Anthropic, DriverId::OpenAI]);
2543 }
2544
2545 #[test]
2546 fn test_create_chat_driver_fails_without_chat_factory() {
2547 let mut registry = DriverRegistry::new();
2548 registry.register_descriptor(DriverDescriptor {
2549 id: DriverId::external("embeddings-only"),
2550 display_name: "Embeddings Only".to_string(),
2551 services: vec![ServiceKind::Embeddings],
2552 credential_schema: CredentialFormSchema::empty(),
2553 oauth: None,
2554 chat: None,
2555 embeddings: None,
2556 });
2557
2558 let config = ProviderConfig::new(DriverId::external("embeddings-only"));
2559 let err = match registry.create_chat_driver(&config) {
2560 Ok(_) => panic!("expected error for missing chat factory"),
2561 Err(err) => err,
2562 };
2563 assert!(
2564 err.to_string()
2565 .contains("does not implement the chat service"),
2566 "unexpected error: {err}"
2567 );
2568 }
2569
2570 #[test]
2571 #[should_panic(expected = "already registered")]
2572 fn test_register_duplicate_panics() {
2573 struct MockDriver;
2574 #[async_trait]
2575 impl ChatDriver for MockDriver {
2576 async fn chat_completion_stream(
2577 &self,
2578 _messages: Vec<LlmMessage>,
2579 _config: &LlmCallConfig,
2580 ) -> Result<LlmResponseStream> {
2581 unimplemented!()
2582 }
2583 }
2584
2585 let mut registry = DriverRegistry::new();
2586 registry.register(DriverId::OpenAI, |_config| Box::new(MockDriver));
2587 registry.register(DriverId::OpenAI, |_config| Box::new(MockDriver));
2589 }
2590
2591 #[test]
2592 fn test_register_or_replace_overwrites() {
2593 struct MockDriver;
2594 #[async_trait]
2595 impl ChatDriver for MockDriver {
2596 async fn chat_completion_stream(
2597 &self,
2598 _messages: Vec<LlmMessage>,
2599 _config: &LlmCallConfig,
2600 ) -> Result<LlmResponseStream> {
2601 unimplemented!()
2602 }
2603 }
2604
2605 let mut registry = DriverRegistry::new();
2606 registry.register(DriverId::LlmSim, |_config| Box::new(MockDriver));
2607 registry.register_or_replace(DriverId::LlmSim, |_config| Box::new(MockDriver));
2609 assert!(registry.has_driver(&DriverId::LlmSim));
2610 }
2611
2612 #[test]
2613 fn test_prepend_text_prefix_simple_text() {
2614 let mut msg = LlmMessage::text(LlmMessageRole::User, "Hello bot");
2615 msg.prepend_text_prefix("[Alice] ");
2616 assert_eq!(msg.content_as_text(), "[Alice] Hello bot");
2617 }
2618
2619 #[test]
2620 fn test_prepend_text_prefix_parts() {
2621 let mut msg = LlmMessage::parts(
2622 LlmMessageRole::User,
2623 vec![
2624 LlmContentPart::Text {
2625 text: "Hello".to_string(),
2626 },
2627 LlmContentPart::Image {
2628 url: "data:image/png;base64,abc".to_string(),
2629 },
2630 ],
2631 );
2632 msg.prepend_text_prefix("[Bob] ");
2633 match &msg.content {
2634 LlmMessageContent::Parts(parts) => {
2635 if let LlmContentPart::Text { text } = &parts[0] {
2636 assert_eq!(text, "[Bob] Hello");
2637 } else {
2638 panic!("Expected text part");
2639 }
2640 }
2641 _ => panic!("Expected parts content"),
2642 }
2643 }
2644
2645 #[test]
2646 fn test_prepend_text_prefix_parts_no_text() {
2647 let mut msg = LlmMessage::parts(
2648 LlmMessageRole::User,
2649 vec![LlmContentPart::Image {
2650 url: "data:image/png;base64,abc".to_string(),
2651 }],
2652 );
2653 msg.prepend_text_prefix("[Eve] ");
2654 match &msg.content {
2655 LlmMessageContent::Parts(parts) => {
2656 assert_eq!(parts.len(), 2);
2657 if let LlmContentPart::Text { text } = &parts[0] {
2658 assert_eq!(text, "[Eve] ");
2659 } else {
2660 panic!("Expected prepended text part");
2661 }
2662 }
2663 _ => panic!("Expected parts content"),
2664 }
2665 }
2666
2667 #[test]
2668 fn test_openrouter_plugin_config_is_empty() {
2669 assert!(OpenRouterPluginConfig::default().is_empty());
2670 assert!(
2671 !OpenRouterPluginConfig {
2672 web: Some(OpenRouterWebSearchPlugin::default()),
2673 file: None,
2674 }
2675 .is_empty()
2676 );
2677 assert!(
2678 !OpenRouterPluginConfig {
2679 web: None,
2680 file: Some(OpenRouterFilePlugin {}),
2681 }
2682 .is_empty()
2683 );
2684 }
2685
2686 #[test]
2687 fn test_openrouter_routing_is_empty_with_plugins() {
2688 let with_plugins = OpenRouterRoutingConfig {
2689 plugins: Some(OpenRouterPluginConfig {
2690 web: Some(OpenRouterWebSearchPlugin::default()),
2691 file: None,
2692 }),
2693 ..Default::default()
2694 };
2695 assert!(!with_plugins.is_empty());
2696
2697 let empty_plugins = OpenRouterRoutingConfig {
2698 plugins: Some(OpenRouterPluginConfig::default()),
2699 ..Default::default()
2700 };
2701 assert!(empty_plugins.is_empty());
2702 }
2703
2704 #[test]
2705 fn test_openrouter_web_search_plugin_serialization() {
2706 let plugin = OpenRouterWebSearchPlugin {
2707 max_results: Some(10),
2708 search_prompt: Some("search for Rust crates".to_string()),
2709 };
2710 let json = serde_json::to_value(&plugin).unwrap();
2711 assert_eq!(json["max_results"], 10);
2712 assert_eq!(json["search_prompt"], "search for Rust crates");
2713 }
2714
2715 #[test]
2716 fn test_openrouter_web_search_plugin_omits_none_fields() {
2717 let plugin = OpenRouterWebSearchPlugin::default();
2718 let json = serde_json::to_value(&plugin).unwrap();
2719 assert!(json.get("max_results").is_none());
2720 assert!(json.get("search_prompt").is_none());
2721 }
2722
2723 #[test]
2724 fn test_capacity_strategy_shared_capacity_is_noop() {
2725 let base = OpenRouterRoutingConfig {
2726 models: vec!["openai/gpt-5-mini".to_string()],
2727 capacity_strategy: Some(OpenRouterCapacityStrategy::SharedCapacity),
2728 ..Default::default()
2729 };
2730 let result = base.apply_capacity_strategy().unwrap();
2731 assert_eq!(
2732 result.capacity_strategy,
2733 Some(OpenRouterCapacityStrategy::SharedCapacity)
2734 );
2735 assert!(result.provider.is_none());
2736 }
2737
2738 #[test]
2739 fn test_capacity_strategy_none_is_noop() {
2740 let base = OpenRouterRoutingConfig {
2741 models: vec!["openai/gpt-5-mini".to_string()],
2742 capacity_strategy: None,
2743 ..Default::default()
2744 };
2745 let result = base.apply_capacity_strategy().unwrap();
2746 assert!(result.provider.is_none());
2747 }
2748
2749 #[test]
2750 fn test_capacity_strategy_byok_first_sets_allow_fallbacks() {
2751 let base = OpenRouterRoutingConfig {
2752 models: vec!["openai/gpt-5-mini".to_string()],
2753 capacity_strategy: Some(OpenRouterCapacityStrategy::ByokFirst),
2754 ..Default::default()
2755 };
2756 let result = base.apply_capacity_strategy().unwrap();
2757 let provider = result.provider.as_ref().expect("provider set by ByokFirst");
2758 assert_eq!(provider.allow_fallbacks, Some(true));
2759 }
2760
2761 #[test]
2762 fn test_capacity_strategy_byok_first_preserves_explicit_allow_fallbacks() {
2763 let base = OpenRouterRoutingConfig {
2765 models: vec!["openai/gpt-5-mini".to_string()],
2766 capacity_strategy: Some(OpenRouterCapacityStrategy::ByokFirst),
2767 provider: Some(OpenRouterProviderRouting {
2768 allow_fallbacks: Some(false),
2769 ..Default::default()
2770 }),
2771 ..Default::default()
2772 };
2773 let result = base.apply_capacity_strategy().unwrap();
2774 let provider = result.provider.as_ref().unwrap();
2775 assert_eq!(provider.allow_fallbacks, Some(false));
2776 }
2777
2778 #[test]
2779 fn test_capacity_strategy_byok_only_requires_provider_only() {
2780 let base = OpenRouterRoutingConfig {
2781 models: vec!["openai/gpt-5-mini".to_string()],
2782 capacity_strategy: Some(OpenRouterCapacityStrategy::ByokOnly),
2783 ..Default::default()
2784 };
2785 let err = base.apply_capacity_strategy().unwrap_err();
2786 assert!(
2787 err.contains("provider.only"),
2788 "error should mention provider.only: {err}"
2789 );
2790 }
2791
2792 #[test]
2793 fn test_capacity_strategy_byok_only_disables_fallbacks() {
2794 let base = OpenRouterRoutingConfig {
2795 models: vec!["openai/gpt-5-mini".to_string()],
2796 capacity_strategy: Some(OpenRouterCapacityStrategy::ByokOnly),
2797 provider: Some(OpenRouterProviderRouting {
2798 only: vec!["my-byok-provider".to_string()],
2799 ..Default::default()
2800 }),
2801 ..Default::default()
2802 };
2803 let result = base.apply_capacity_strategy().unwrap();
2804 let provider = result.provider.as_ref().unwrap();
2805 assert_eq!(provider.allow_fallbacks, Some(false));
2806 assert_eq!(provider.only, vec!["my-byok-provider"]);
2807 }
2808
2809 #[test]
2810 fn test_capacity_strategy_byok_only_not_empty_in_is_empty() {
2811 let with_strategy = OpenRouterRoutingConfig {
2812 capacity_strategy: Some(OpenRouterCapacityStrategy::ByokOnly),
2813 ..Default::default()
2814 };
2815 assert!(!with_strategy.is_empty());
2816
2817 let byok_first = OpenRouterRoutingConfig {
2818 capacity_strategy: Some(OpenRouterCapacityStrategy::ByokFirst),
2819 ..Default::default()
2820 };
2821 assert!(!byok_first.is_empty());
2822
2823 let shared = OpenRouterRoutingConfig {
2824 capacity_strategy: Some(OpenRouterCapacityStrategy::SharedCapacity),
2825 ..Default::default()
2826 };
2827 assert!(shared.is_empty());
2828 }
2829
2830 #[test]
2837 fn test_preset_no_presets_is_noop() {
2838 let base = OpenRouterRoutingConfig {
2839 models: vec!["openai/gpt-5-mini".to_string()],
2840 ..Default::default()
2841 };
2842 let result = base.apply_presets().unwrap();
2843 assert_eq!(result, base);
2844 }
2845
2846 #[test]
2847 fn test_preset_cheapest_with_tools_sets_require_parameters_and_sort_price() {
2848 let base = OpenRouterRoutingConfig {
2849 presets: vec![OpenRouterRoutingPreset::CheapestWithTools],
2850 ..Default::default()
2851 };
2852 let result = base.apply_presets().unwrap();
2853 assert!(result.presets.is_empty(), "presets cleared after apply");
2854 let provider = result.provider.expect("provider set by preset");
2855 assert_eq!(provider.require_parameters, Some(true));
2856 assert_eq!(
2857 provider.sort,
2858 Some(OpenRouterProviderSort::Simple(
2859 OpenRouterProviderSortBy::Price
2860 ))
2861 );
2862 }
2863
2864 #[test]
2865 fn test_preset_lowest_latency_review_sets_sort_throughput() {
2866 let base = OpenRouterRoutingConfig {
2867 presets: vec![OpenRouterRoutingPreset::LowestLatencyReview],
2868 ..Default::default()
2869 };
2870 let result = base.apply_presets().unwrap();
2871 let provider = result.provider.expect("provider set by preset");
2872 assert_eq!(
2873 provider.sort,
2874 Some(OpenRouterProviderSort::Simple(
2875 OpenRouterProviderSortBy::Throughput
2876 ))
2877 );
2878 }
2879
2880 #[test]
2881 fn test_preset_zdr_only_sets_zdr() {
2882 let base = OpenRouterRoutingConfig {
2883 presets: vec![OpenRouterRoutingPreset::ZdrOnly],
2884 ..Default::default()
2885 };
2886 let result = base.apply_presets().unwrap();
2887 let provider = result.provider.expect("provider set");
2888 assert_eq!(provider.zdr, Some(true));
2889 }
2890
2891 #[test]
2892 fn test_preset_byok_first_sets_allow_fallbacks() {
2893 let base = OpenRouterRoutingConfig {
2894 presets: vec![OpenRouterRoutingPreset::ByokFirst],
2895 ..Default::default()
2896 };
2897 let result = base.apply_presets().unwrap();
2898 let provider = result.provider.expect("provider set");
2899 assert_eq!(provider.allow_fallbacks, Some(true));
2900 }
2901
2902 #[test]
2903 fn test_preset_no_data_collection_sets_data_collection_deny() {
2904 let base = OpenRouterRoutingConfig {
2905 presets: vec![OpenRouterRoutingPreset::NoDataCollection],
2906 ..Default::default()
2907 };
2908 let result = base.apply_presets().unwrap();
2909 let provider = result.provider.expect("provider set");
2910 assert_eq!(
2911 provider.data_collection,
2912 Some(OpenRouterDataCollection::Deny)
2913 );
2914 }
2915
2916 #[test]
2917 fn test_preset_strict_json_sets_require_parameters() {
2918 let base = OpenRouterRoutingConfig {
2919 presets: vec![OpenRouterRoutingPreset::StrictJson],
2920 ..Default::default()
2921 };
2922 let result = base.apply_presets().unwrap();
2923 let provider = result.provider.expect("provider set");
2924 assert_eq!(provider.require_parameters, Some(true));
2925 }
2926
2927 #[test]
2928 fn test_preset_reasoning_required_sets_require_parameters() {
2929 let base = OpenRouterRoutingConfig {
2930 presets: vec![OpenRouterRoutingPreset::ReasoningRequired],
2931 ..Default::default()
2932 };
2933 let result = base.apply_presets().unwrap();
2934 let provider = result.provider.expect("provider set");
2935 assert_eq!(provider.require_parameters, Some(true));
2936 }
2937
2938 #[test]
2939 fn test_preset_max_price_converts_usd_per_million() {
2940 let base = OpenRouterRoutingConfig {
2941 presets: vec![OpenRouterRoutingPreset::MaxPrice {
2942 prompt_usd_per_million: Some(5.0),
2943 completion_usd_per_million: Some(15.0),
2944 }],
2945 ..Default::default()
2946 };
2947 let result = base.apply_presets().unwrap();
2948 let provider = result.provider.expect("provider set");
2949 let max_price = provider.max_price.expect("max_price set");
2950 let prompt = max_price.prompt.expect("prompt set");
2952 assert!((prompt - 5.0 / 1_000_000.0).abs() < f64::EPSILON);
2953 let completion = max_price.completion.expect("completion set");
2954 assert!((completion - 15.0 / 1_000_000.0).abs() < f64::EPSILON);
2955 }
2956
2957 #[test]
2958 fn test_preset_max_price_rejects_negative_values() {
2959 let base = OpenRouterRoutingConfig {
2960 presets: vec![OpenRouterRoutingPreset::MaxPrice {
2961 prompt_usd_per_million: Some(-1.0),
2962 completion_usd_per_million: None,
2963 }],
2964 ..Default::default()
2965 };
2966 let err = base.apply_presets().unwrap_err();
2967 assert!(
2968 err.contains("non-negative"),
2969 "error should mention non-negative: {err}"
2970 );
2971 }
2972
2973 #[test]
2974 fn test_preset_max_price_both_none_no_provider_field() {
2975 let base = OpenRouterRoutingConfig {
2976 presets: vec![OpenRouterRoutingPreset::MaxPrice {
2977 prompt_usd_per_million: None,
2978 completion_usd_per_million: None,
2979 }],
2980 ..Default::default()
2981 };
2982 let result = base.apply_presets().unwrap();
2983 assert!(
2984 result.provider.is_none(),
2985 "MaxPrice with no dimensions should not produce a provider field"
2986 );
2987 }
2988
2989 #[test]
2990 fn test_preset_explicit_provider_overrides_preset() {
2991 let base = OpenRouterRoutingConfig {
2992 presets: vec![OpenRouterRoutingPreset::CheapestWithTools],
2993 provider: Some(OpenRouterProviderRouting {
2994 sort: Some(OpenRouterProviderSort::Simple(
2996 OpenRouterProviderSortBy::Throughput,
2997 )),
2998 ..Default::default()
2999 }),
3000 ..Default::default()
3001 };
3002 let result = base.apply_presets().unwrap();
3003 let provider = result.provider.expect("provider set");
3004 assert_eq!(
3006 provider.sort,
3007 Some(OpenRouterProviderSort::Simple(
3008 OpenRouterProviderSortBy::Throughput
3009 ))
3010 );
3011 assert_eq!(provider.require_parameters, Some(true));
3013 }
3014
3015 #[test]
3016 fn test_preset_multiple_presets_combined() {
3017 let base = OpenRouterRoutingConfig {
3018 presets: vec![
3019 OpenRouterRoutingPreset::ZdrOnly,
3020 OpenRouterRoutingPreset::NoDataCollection,
3021 OpenRouterRoutingPreset::LowestLatencyReview,
3022 ],
3023 ..Default::default()
3024 };
3025 let result = base.apply_presets().unwrap();
3026 let provider = result.provider.expect("provider set");
3027 assert_eq!(provider.zdr, Some(true));
3028 assert_eq!(
3029 provider.data_collection,
3030 Some(OpenRouterDataCollection::Deny)
3031 );
3032 assert_eq!(
3033 provider.sort,
3034 Some(OpenRouterProviderSort::Simple(
3035 OpenRouterProviderSortBy::Throughput
3036 ))
3037 );
3038 }
3039
3040 #[test]
3041 fn test_preset_later_preset_overrides_sort() {
3042 let base = OpenRouterRoutingConfig {
3043 presets: vec![
3044 OpenRouterRoutingPreset::CheapestWithTools, OpenRouterRoutingPreset::LowestLatencyReview, ],
3047 ..Default::default()
3048 };
3049 let result = base.apply_presets().unwrap();
3050 let provider = result.provider.expect("provider set");
3051 assert_eq!(
3053 provider.sort,
3054 Some(OpenRouterProviderSort::Simple(
3055 OpenRouterProviderSortBy::Throughput
3056 ))
3057 );
3058 assert_eq!(provider.require_parameters, Some(true));
3060 }
3061
3062 #[test]
3063 fn test_preset_non_empty_in_is_empty() {
3064 let with_preset = OpenRouterRoutingConfig {
3065 presets: vec![OpenRouterRoutingPreset::ZdrOnly],
3066 ..Default::default()
3067 };
3068 assert!(!with_preset.is_empty());
3069
3070 let without = OpenRouterRoutingConfig::default();
3071 assert!(without.is_empty());
3072 }
3073}