1use crate::compact::{CompactOutputItem, CompactRequest, CompactResponse};
18use crate::credential_schema::CredentialFormSchema;
19use crate::error::{AgentLoopError, LlmErrorKind, Result};
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 {
46 output: Vec<CompactOutputItem>,
47 #[serde(default, skip_serializing_if = "Option::is_none")]
48 reasoning_state: Option<crate::reasoning_updates::ReasoningState>,
49 },
50}
51
52impl std::fmt::Debug for ProviderOpaqueContext {
53 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
54 match self {
55 Self::OpenResponsesCompact { output, .. } => f
56 .debug_struct("OpenResponsesCompact")
57 .field("item_count", &output.len())
58 .finish_non_exhaustive(),
59 }
60 }
61}
62
63#[derive(Debug, Clone, PartialEq, Eq)]
69pub struct LlmStreamError {
70 pub code: Option<String>,
72 pub status: Option<u16>,
74 pub message: String,
76}
77
78impl LlmStreamError {
79 pub fn new(message: impl Into<String>) -> Self {
80 Self {
81 code: None,
82 status: None,
83 message: message.into(),
84 }
85 }
86
87 pub fn provider(
89 code: Option<impl Into<String>>,
90 status: Option<u16>,
91 message: impl Into<String>,
92 ) -> Self {
93 Self {
94 code: code.map(Into::into),
95 status,
96 message: message.into(),
97 }
98 }
99
100 pub fn kind(&self) -> LlmErrorKind {
102 if let Some(code) = self.code.as_deref()
103 && let Some(kind) = LlmErrorKind::from_provider_code(code)
104 {
105 return kind;
106 }
107 if let Some(status) = self.status {
108 return LlmErrorKind::from_provider_status(status, &self.message);
109 }
110 LlmErrorKind::from_error_text(&self.message)
111 }
112}
113
114impl std::error::Error for LlmStreamError {}
115
116impl std::fmt::Display for LlmStreamError {
117 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
118 match (&self.code, self.status) {
119 (Some(code), Some(status)) => write!(f, "{code} ({status}): {}", self.message),
120 (Some(code), None) => write!(f, "{code}: {}", self.message),
121 (None, Some(status)) => write!(f, "({status}): {}", self.message),
122 (None, None) => f.write_str(&self.message),
123 }
124 }
125}
126
127impl From<String> for LlmStreamError {
128 fn from(message: String) -> Self {
129 Self::new(message)
130 }
131}
132
133impl From<&str> for LlmStreamError {
134 fn from(message: &str) -> Self {
135 Self::new(message)
136 }
137}
138
139#[derive(Debug, Clone)]
145#[non_exhaustive]
146pub enum LlmStreamEvent {
147 TextDelta(String),
149 ReasoningDelta { delta: String, summary: bool },
156 ReasoningItem(crate::reasoning::ReasoningContentPart),
162 ToolCalls(Vec<ToolCall>),
164 NativeToolCall(crate::native_async::NativeToolCall),
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)]
226#[non_exhaustive]
227pub struct LlmCompletionMetadata {
228 pub total_tokens: Option<u32>,
230 pub prompt_tokens: Option<u32>,
232 pub completion_tokens: Option<u32>,
234 pub cache_read_tokens: Option<u32>,
236 pub cache_creation_tokens: Option<u32>,
238 pub provider_cost_usd: Option<f64>,
242 pub model: Option<String>,
244 pub finish_reason: Option<String>,
246 pub retry_metadata: Option<crate::llm_retry::RetryMetadata>,
248 pub response_id: Option<String>,
251 pub phase: Option<String>,
255 pub cache_diagnostics: Option<serde_json::Value>,
263}
264
265pub fn disjoint_prompt_tokens(reported_input: u32, cache_read: Option<u32>) -> u32 {
275 reported_input.saturating_sub(cache_read.unwrap_or(0))
276}
277
278#[async_trait]
300pub trait ChatDriver: Send + Sync {
301 fn native_async_driver(
303 &self,
304 _model: &str,
305 _tools: std::collections::BTreeMap<String, Option<serde_json::Value>>,
306 _continuation: Option<crate::native_async::Delivery>,
307 ) -> Option<Arc<dyn ChatDriver>> {
308 None
309 }
310 async fn chat_completion_stream(
312 &self,
313 endpoint: &crate::runtime_provider::ProviderEndpoint,
314 messages: Vec<LlmMessage>,
315 config: &LlmCallConfig,
316 ) -> Result<LlmResponseStream>;
317
318 async fn chat_completion(
320 &self,
321 endpoint: &crate::runtime_provider::ProviderEndpoint,
322 messages: Vec<LlmMessage>,
323 config: &LlmCallConfig,
324 ) -> Result<LlmResponse> {
325 use futures::StreamExt;
326
327 let mut stream = self
328 .chat_completion_stream(endpoint, messages, config)
329 .await?;
330 let mut text = String::new();
331 let mut reasoning: Vec<crate::reasoning::ReasoningContentPart> = Vec::new();
332 let mut tool_calls = Vec::new();
333 let mut metadata = LlmCompletionMetadata::default();
334
335 while let Some(event) = stream.next().await {
336 match event? {
337 LlmStreamEvent::TextDelta(delta) => text.push_str(&delta),
338 LlmStreamEvent::ReasoningDelta { .. } => {}
341 LlmStreamEvent::ReasoningItem(item) => reasoning.push(item),
342 LlmStreamEvent::ToolCalls(calls) => tool_calls = calls,
343 LlmStreamEvent::NativeToolCall(_) => {
344 return Err(crate::error::AgentLoopError::config(
345 "native async/custom calls require a streaming coordinator",
346 ));
347 }
348 LlmStreamEvent::MessagePhase(_) => {}
351 LlmStreamEvent::Done(meta) => metadata = *meta,
352 LlmStreamEvent::Error(err) => {
353 return Err(crate::error::AgentLoopError::llm_kind(
354 err.kind(),
355 err.to_string(),
356 ));
357 }
358 }
359 }
360
361 Ok(LlmResponse {
362 text,
363 reasoning,
364 tool_calls: if tool_calls.is_empty() {
365 None
366 } else {
367 Some(tool_calls)
368 },
369 metadata,
370 })
371 }
372
373 fn supports_native_non_streaming(&self) -> bool {
381 false
382 }
383
384 async fn chat_completion_non_streaming(
391 &self,
392 endpoint: &crate::runtime_provider::ProviderEndpoint,
393 messages: Vec<LlmMessage>,
394 config: &LlmCallConfig,
395 ) -> Result<LlmResponse> {
396 self.chat_completion(endpoint, messages, config).await
397 }
398
399 async fn list_models(
407 &self,
408 _endpoint: &crate::runtime_provider::ProviderEndpoint,
409 ) -> Result<Option<Vec<DiscoveredModel>>> {
410 Ok(None)
412 }
413
414 fn supports_compact(&self) -> bool {
423 false
425 }
426
427 fn supports_stateful_responses(&self) -> bool {
433 false
434 }
435
436 fn effective_context_window(&self, _model: &str) -> Option<usize> {
442 None
443 }
444
445 fn supports_parallel_tool_calls(&self, _model: &str) -> bool {
456 false
457 }
458
459 async fn compact(
479 &self,
480 _endpoint: &crate::runtime_provider::ProviderEndpoint,
481 _request: CompactRequest,
482 ) -> Result<Option<CompactResponse>> {
483 Ok(None)
485 }
486}
487
488#[async_trait]
490impl ChatDriver for Box<dyn ChatDriver> {
491 fn native_async_driver(
492 &self,
493 model: &str,
494 tools: std::collections::BTreeMap<String, Option<serde_json::Value>>,
495 continuation: Option<crate::native_async::Delivery>,
496 ) -> Option<Arc<dyn ChatDriver>> {
497 (**self).native_async_driver(model, tools, continuation)
498 }
499 async fn chat_completion_stream(
500 &self,
501 endpoint: &crate::runtime_provider::ProviderEndpoint,
502 messages: Vec<LlmMessage>,
503 config: &LlmCallConfig,
504 ) -> Result<LlmResponseStream> {
505 (**self)
506 .chat_completion_stream(endpoint, messages, config)
507 .await
508 }
509
510 async fn chat_completion(
511 &self,
512 endpoint: &crate::runtime_provider::ProviderEndpoint,
513 messages: Vec<LlmMessage>,
514 config: &LlmCallConfig,
515 ) -> Result<LlmResponse> {
516 (**self).chat_completion(endpoint, messages, config).await
517 }
518
519 fn supports_native_non_streaming(&self) -> bool {
520 (**self).supports_native_non_streaming()
521 }
522
523 async fn chat_completion_non_streaming(
524 &self,
525 endpoint: &crate::runtime_provider::ProviderEndpoint,
526 messages: Vec<LlmMessage>,
527 config: &LlmCallConfig,
528 ) -> Result<LlmResponse> {
529 (**self)
530 .chat_completion_non_streaming(endpoint, messages, config)
531 .await
532 }
533
534 async fn list_models(
535 &self,
536 endpoint: &crate::runtime_provider::ProviderEndpoint,
537 ) -> Result<Option<Vec<DiscoveredModel>>> {
538 (**self).list_models(endpoint).await
539 }
540
541 fn supports_compact(&self) -> bool {
542 (**self).supports_compact()
543 }
544
545 fn supports_stateful_responses(&self) -> bool {
546 (**self).supports_stateful_responses()
547 }
548
549 fn effective_context_window(&self, model: &str) -> Option<usize> {
550 (**self).effective_context_window(model)
551 }
552
553 fn supports_parallel_tool_calls(&self, model: &str) -> bool {
554 (**self).supports_parallel_tool_calls(model)
555 }
556
557 async fn compact(
558 &self,
559 endpoint: &crate::runtime_provider::ProviderEndpoint,
560 request: CompactRequest,
561 ) -> Result<Option<CompactResponse>> {
562 (**self).compact(endpoint, request).await
563 }
564}
565
566#[derive(Debug, Clone)]
572pub struct LlmMessage {
573 pub native_tool_calls: Vec<crate::native_async::NativeToolCall>,
575 pub role: LlmMessageRole,
576 pub content: LlmMessageContent,
577 pub tool_calls: Option<Vec<ToolCall>>,
578 pub tool_call_id: Option<String>,
579 pub phase: Option<crate::execution_phase::ExecutionPhase>,
584 pub reasoning: Vec<crate::reasoning::ReasoningContentPart>,
591 pub configuration_update: Option<crate::model::ReasoningEffort>,
594}
595
596impl LlmMessage {
597 pub fn text(role: LlmMessageRole, content: impl Into<String>) -> Self {
599 Self {
600 native_tool_calls: Vec::new(),
601 role,
602 content: LlmMessageContent::Text(content.into()),
603 tool_calls: None,
604 tool_call_id: None,
605 phase: None,
606 reasoning: Vec::new(),
607 configuration_update: None,
608 }
609 }
610
611 pub fn parts(role: LlmMessageRole, parts: Vec<LlmContentPart>) -> Self {
613 Self {
614 native_tool_calls: Vec::new(),
615 role,
616 content: LlmMessageContent::Parts(parts),
617 tool_calls: None,
618 tool_call_id: None,
619 phase: None,
620 reasoning: Vec::new(),
621 configuration_update: None,
622 }
623 }
624
625 pub fn content_as_text(&self) -> String {
627 self.content.to_text()
628 }
629
630 pub fn prepend_text_prefix(&mut self, prefix: &str) {
635 match &mut self.content {
636 LlmMessageContent::Text(text) => {
637 *text = format!("{}{}", prefix, text);
638 }
639 LlmMessageContent::Parts(parts) => {
640 for part in parts.iter_mut() {
641 if let LlmContentPart::Text { text } = part {
642 *text = format!("{}{}", prefix, text);
643 return;
644 }
645 }
646 parts.insert(
648 0,
649 LlmContentPart::Text {
650 text: prefix.to_string(),
651 },
652 );
653 }
654 }
655 }
656}
657
658pub fn fold_system_messages(messages: &[LlmMessage]) -> Option<String> {
669 let mut system: Option<String> = None;
670 for msg in messages {
671 if msg.role == LlmMessageRole::System {
672 let text = msg.content.to_text();
673 system = Some(match system.take() {
674 Some(existing) if !existing.is_empty() => format!("{existing}\n\n{text}"),
675 _ => text,
676 });
677 }
678 }
679 system
680}
681
682#[derive(Debug, Clone)]
684pub enum LlmMessageContent {
685 Text(String),
687 Parts(Vec<LlmContentPart>),
689}
690
691impl LlmMessageContent {
692 pub fn to_text(&self) -> String {
694 match self {
695 LlmMessageContent::Text(s) => s.clone(),
696 LlmMessageContent::Parts(parts) => parts
697 .iter()
698 .filter_map(|p| match p {
699 LlmContentPart::Text { text } => Some(text.clone()),
700 _ => None,
701 })
702 .collect::<Vec<_>>()
703 .join(""),
704 }
705 }
706
707 pub fn is_text(&self) -> bool {
709 matches!(self, LlmMessageContent::Text(_))
710 }
711
712 pub fn is_parts(&self) -> bool {
714 matches!(self, LlmMessageContent::Parts(_))
715 }
716}
717
718impl From<String> for LlmMessageContent {
719 fn from(s: String) -> Self {
720 LlmMessageContent::Text(s)
721 }
722}
723
724impl From<&str> for LlmMessageContent {
725 fn from(s: &str) -> Self {
726 LlmMessageContent::Text(s.to_string())
727 }
728}
729
730#[derive(Debug, Clone)]
735#[non_exhaustive]
736pub enum LlmContentPart {
737 Text { text: String },
739 Image { url: String },
741 Audio { url: String },
743 File {
745 url: String,
746 filename: Option<String>,
747 },
748}
749
750impl LlmContentPart {
751 pub fn text(text: impl Into<String>) -> Self {
753 LlmContentPart::Text { text: text.into() }
754 }
755
756 pub fn image(url: impl Into<String>) -> Self {
758 LlmContentPart::Image { url: url.into() }
759 }
760
761 pub fn audio(url: impl Into<String>) -> Self {
763 LlmContentPart::Audio { url: url.into() }
764 }
765
766 pub fn file(url: impl Into<String>, filename: Option<String>) -> Self {
768 LlmContentPart::File {
769 url: url.into(),
770 filename,
771 }
772 }
773}
774
775#[derive(Debug, Clone, PartialEq, Eq)]
777pub enum LlmMessageRole {
778 System,
779 User,
780 Assistant,
781 Tool,
782}
783
784#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
794pub struct ToolSearchConfig {
795 pub enabled: bool,
797 pub threshold: usize,
800}
801
802#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
804#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
805#[serde(rename_all = "snake_case")]
806pub enum PromptCacheStrategy {
807 #[default]
809 Auto,
810 Explicit,
813}
814
815#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
820#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
821pub struct PromptCacheConfig {
822 pub enabled: bool,
824 #[serde(default)]
826 pub strategy: PromptCacheStrategy,
827 #[serde(default, skip_serializing_if = "Option::is_none")]
834 pub gemini_cached_content: Option<String>,
835}
836
837#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
844#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
845pub struct CacheDiagnosticsConfig {
846 pub enabled: bool,
848 #[serde(default, skip_serializing_if = "Option::is_none")]
854 pub previous_message_id: Option<String>,
855}
856
857#[derive(Debug, Clone, Default)]
865#[non_exhaustive]
866pub struct LlmCallConfig {
867 pub reasoning_state: Option<crate::reasoning_updates::ReasoningState>,
869 pub model: String,
870 pub temperature: Option<f32>,
871 pub max_tokens: Option<u32>,
872 pub tools: Vec<ToolDefinition>,
873 pub reasoning_effort: Option<crate::model::ReasoningEffort>,
879 pub speed: Option<String>,
883 pub verbosity: Option<String>,
887 pub metadata: HashMap<String, String>,
891 pub previous_response_id: Option<String>,
894 pub provider_opaque_context: Option<ProviderOpaqueContext>,
900 pub tool_search: Option<ToolSearchConfig>,
902 pub prompt_cache: Option<PromptCacheConfig>,
904 pub driver_options: HashMap<String, serde_json::Value>,
908 pub parallel_tool_calls: Option<bool>,
915 pub volatile_suffix_len: usize,
925 pub extra_headers: Vec<(String, String)>,
934 pub cache_diagnostics: Option<CacheDiagnosticsConfig>,
936}
937
938impl LlmCallConfig {
939 pub fn new(model: impl Into<String>) -> Self {
941 Self {
942 model: model.into(),
943 ..Default::default()
944 }
945 }
946
947 pub fn resolved_parallel_tool_calls(&self, supported: bool) -> Option<bool> {
957 if supported {
958 self.parallel_tool_calls
959 } else {
960 None
961 }
962 }
963}
964
965#[derive(Debug, Clone)]
970pub struct LlmResponse {
971 pub text: String,
972 pub reasoning: Vec<crate::reasoning::ReasoningContentPart>,
974 pub tool_calls: Option<Vec<ToolCall>>,
975 pub metadata: LlmCompletionMetadata,
976}
977
978pub struct LlmCallConfigBuilder {
984 config: LlmCallConfig,
985}
986
987impl LlmCallConfigBuilder {
988 pub fn from_config(config: LlmCallConfig) -> Self {
990 Self { config }
991 }
992
993 pub fn reasoning_effort(mut self, effort: crate::model::ReasoningEffort) -> Self {
995 self.config.reasoning_effort = Some(effort);
996 self
997 }
998
999 pub fn speed(mut self, speed: impl Into<String>) -> Self {
1001 self.config.speed = Some(speed.into());
1002 self
1003 }
1004
1005 pub fn verbosity(mut self, verbosity: impl Into<String>) -> Self {
1007 self.config.verbosity = Some(verbosity.into());
1008 self
1009 }
1010
1011 pub fn model(mut self, model: impl Into<String>) -> Self {
1013 self.config.model = model.into();
1014 self
1015 }
1016
1017 pub fn temperature(mut self, temp: f32) -> Self {
1019 self.config.temperature = Some(temp);
1020 self
1021 }
1022
1023 pub fn max_tokens(mut self, tokens: u32) -> Self {
1025 self.config.max_tokens = Some(tokens);
1026 self
1027 }
1028
1029 pub fn tools(mut self, tools: Vec<ToolDefinition>) -> Self {
1031 self.config.tools = tools;
1032 self
1033 }
1034
1035 pub fn metadata(mut self, metadata: HashMap<String, String>) -> Self {
1040 self.config.metadata = metadata;
1041 self
1042 }
1043
1044 pub fn with_metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
1046 self.config.metadata.insert(key.into(), value.into());
1047 self
1048 }
1049
1050 pub fn previous_response_id(mut self, id: Option<String>) -> Self {
1052 self.config.previous_response_id = id;
1053 self
1054 }
1055
1056 pub fn provider_opaque_context(mut self, context: Option<ProviderOpaqueContext>) -> Self {
1058 self.config.provider_opaque_context = context;
1059 self
1060 }
1061
1062 pub fn tool_search(mut self, config: ToolSearchConfig) -> Self {
1064 self.config.tool_search = Some(config);
1065 self
1066 }
1067
1068 pub fn prompt_cache(mut self, config: PromptCacheConfig) -> Self {
1070 self.config.prompt_cache = Some(config);
1071 self
1072 }
1073
1074 pub fn driver_option(mut self, key: impl Into<String>, value: serde_json::Value) -> Self {
1078 self.config.driver_options.insert(key.into(), value);
1079 self
1080 }
1081
1082 pub fn parallel_tool_calls(mut self, parallel_tool_calls: Option<bool>) -> Self {
1084 self.config.parallel_tool_calls = parallel_tool_calls;
1085 self
1086 }
1087
1088 pub fn volatile_suffix_len(mut self, len: usize) -> Self {
1092 self.config.volatile_suffix_len = len;
1093 self
1094 }
1095
1096 pub fn extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
1098 self.config.extra_headers = headers;
1099 self
1100 }
1101
1102 pub fn extra_header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
1104 self.config.extra_headers.push((name.into(), value.into()));
1105 self
1106 }
1107
1108 pub fn cache_diagnostics(mut self, config: CacheDiagnosticsConfig) -> Self {
1110 self.config.cache_diagnostics = Some(config);
1111 self
1112 }
1113
1114 pub fn build(self) -> LlmCallConfig {
1116 self.config
1117 }
1118}
1119
1120pub use crate::provider::DriverId;
1129
1130#[derive(Clone, Default, PartialEq, Eq)]
1136pub struct ProviderMetadata {
1137 pub refresh_token: Option<String>,
1139 pub account_id: Option<String>,
1141 pub extra: Option<serde_json::Value>,
1143}
1144
1145impl std::fmt::Debug for ProviderMetadata {
1146 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1147 f.debug_struct("ProviderMetadata")
1148 .field(
1149 "refresh_token",
1150 &self.refresh_token.as_ref().map(|_| "<configured>"),
1151 )
1152 .field("account_id", &self.account_id)
1153 .field("extra", &self.extra.as_ref().map(|_| "<configured>"))
1154 .finish()
1155 }
1156}
1157
1158#[derive(Clone)]
1164#[non_exhaustive]
1165pub struct ProviderConfig {
1166 pub provider: crate::runtime_provider::ProviderKey,
1168 pub provider_type: DriverId,
1170 pub api_key: Option<String>,
1172 pub base_url: Option<String>,
1174 pub metadata: ProviderMetadata,
1176 pub request_options: crate::provider::ProviderRequestOptions,
1179}
1180
1181impl ProviderConfig {
1182 pub fn new(provider_type: DriverId) -> Self {
1184 let provider = crate::runtime_provider::ProviderKey::new(provider_type.as_str());
1185 Self {
1186 provider,
1187 provider_type,
1188 api_key: None,
1189 base_url: None,
1190 metadata: ProviderMetadata::default(),
1191 request_options: Default::default(),
1192 }
1193 }
1194
1195 pub fn for_provider(
1198 provider: impl Into<crate::runtime_provider::ProviderKey>,
1199 provider_type: DriverId,
1200 ) -> Self {
1201 Self {
1202 provider: provider.into(),
1203 provider_type,
1204 api_key: None,
1205 base_url: None,
1206 metadata: ProviderMetadata::default(),
1207 request_options: Default::default(),
1208 }
1209 }
1210
1211 pub fn with_api_key(mut self, api_key: impl Into<String>) -> Self {
1213 self.api_key = Some(api_key.into());
1214 self
1215 }
1216
1217 pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
1219 self.base_url = Some(base_url.into());
1220 self
1221 }
1222
1223 pub fn with_metadata(mut self, metadata: ProviderMetadata) -> Self {
1225 self.metadata = metadata;
1226 self
1227 }
1228
1229 pub fn with_request_options(
1231 mut self,
1232 request_options: crate::provider::ProviderRequestOptions,
1233 ) -> Self {
1234 self.request_options = request_options;
1235 self
1236 }
1237}
1238
1239impl std::fmt::Debug for ProviderConfig {
1240 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1241 f.debug_struct("ProviderConfig")
1242 .field("provider", &self.provider)
1243 .field("provider_type", &self.provider_type)
1244 .field("auth", &self.api_key.as_ref().map(|_| "<configured>"))
1245 .field("base_url", &self.base_url.as_ref().map(|_| "<configured>"))
1246 .field(
1247 "metadata",
1248 &self.metadata.extra.as_ref().map(|_| "<configured>"),
1249 )
1250 .finish()
1251 }
1252}
1253
1254#[derive(Clone)]
1260pub struct DriverConfig {
1261 pub provider: crate::runtime_provider::ProviderKey,
1263 pub provider_type: DriverId,
1265 pub api_key: Option<String>,
1271 pub credentials: std::collections::BTreeMap<String, String>,
1277 pub base_url: Option<String>,
1279 pub metadata: ProviderMetadata,
1281}
1282
1283impl DriverConfig {
1284 pub fn from_provider_config(config: &ProviderConfig) -> Self {
1290 Self {
1291 provider: config.provider.clone(),
1292 provider_type: config.provider_type.clone(),
1293 credentials: crate::credential_schema::parse_credential_document(
1294 config.api_key.as_deref(),
1295 ),
1296 api_key: config.api_key.clone(),
1297 base_url: config.base_url.clone(),
1298 metadata: config.metadata.clone(),
1299 }
1300 }
1301
1302 pub fn credential(&self, name: &str) -> Option<&str> {
1304 self.credentials
1305 .get(name)
1306 .map(String::as_str)
1307 .filter(|s| !s.is_empty())
1308 }
1309}
1310
1311impl std::fmt::Debug for DriverConfig {
1312 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1313 f.debug_struct("DriverConfig")
1314 .field("provider", &self.provider)
1315 .field("provider_type", &self.provider_type)
1316 .field("auth", &self.api_key.as_ref().map(|_| "<configured>"))
1317 .field(
1318 "credential_fields",
1319 &self.credentials.keys().collect::<Vec<_>>(),
1320 )
1321 .field("base_url", &self.base_url.as_ref().map(|_| "<configured>"))
1322 .finish()
1323 }
1324}
1325
1326pub type BoxedChatDriver = Box<dyn ChatDriver>;
1328
1329#[derive(Debug, Clone)]
1335pub struct EmbedRequest {
1336 pub texts: Vec<String>,
1338 pub model: String,
1340}
1341
1342#[derive(Debug, Clone)]
1344pub struct EmbedResponse {
1345 pub embeddings: Vec<Vec<f32>>,
1347 pub usage_tokens: Option<u32>,
1350 pub actual_cost_usd: Option<f64>,
1355}
1356
1357#[derive(Debug, thiserror::Error)]
1359pub enum EmbeddingsDriverError {
1360 #[error("embeddings provider returned an error: {0}")]
1361 Provider(String),
1362 #[error("embeddings request failed: {0}")]
1363 Transport(String),
1364}
1365
1366#[async_trait]
1372pub trait EmbeddingsDriver: Send + Sync {
1373 async fn embed(
1375 &self,
1376 endpoint: &crate::runtime_provider::ProviderEndpoint,
1377 request: EmbedRequest,
1378 ) -> std::result::Result<EmbedResponse, EmbeddingsDriverError>;
1379}
1380
1381#[async_trait]
1382impl EmbeddingsDriver for Box<dyn EmbeddingsDriver> {
1383 async fn embed(
1384 &self,
1385 endpoint: &crate::runtime_provider::ProviderEndpoint,
1386 request: EmbedRequest,
1387 ) -> std::result::Result<EmbedResponse, EmbeddingsDriverError> {
1388 (**self).embed(endpoint, request).await
1389 }
1390}
1391
1392pub type BoxedEmbeddingsDriver = Box<dyn EmbeddingsDriver>;
1394
1395pub type EmbeddingsDriverFactory =
1397 Arc<dyn Fn(&DriverConfig) -> BoxedEmbeddingsDriver + Send + Sync>;
1398
1399pub type DriverFactory = Arc<dyn Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync>;
1408
1409struct CredentialGateDriver {
1416 inner: BoxedChatDriver,
1417 message: String,
1418}
1419
1420impl CredentialGateDriver {
1421 fn error(&self) -> AgentLoopError {
1422 AgentLoopError::llm_kind(LlmErrorKind::Authentication, self.message.clone())
1423 }
1424}
1425
1426#[async_trait]
1427impl ChatDriver for CredentialGateDriver {
1428 async fn chat_completion_stream(
1429 &self,
1430 _endpoint: &crate::runtime_provider::ProviderEndpoint,
1431 _messages: Vec<LlmMessage>,
1432 _config: &LlmCallConfig,
1433 ) -> Result<LlmResponseStream> {
1434 Err(self.error())
1435 }
1436
1437 async fn list_models(
1438 &self,
1439 _endpoint: &crate::runtime_provider::ProviderEndpoint,
1440 ) -> Result<Option<Vec<DiscoveredModel>>> {
1441 Err(self.error())
1442 }
1443
1444 fn supports_compact(&self) -> bool {
1445 self.inner.supports_compact()
1446 }
1447
1448 fn supports_stateful_responses(&self) -> bool {
1449 self.inner.supports_stateful_responses()
1450 }
1451
1452 fn effective_context_window(&self, model: &str) -> Option<usize> {
1453 self.inner.effective_context_window(model)
1454 }
1455
1456 fn supports_parallel_tool_calls(&self, model: &str) -> bool {
1457 self.inner.supports_parallel_tool_calls(model)
1458 }
1459
1460 async fn compact(
1461 &self,
1462 _endpoint: &crate::runtime_provider::ProviderEndpoint,
1463 _request: CompactRequest,
1464 ) -> Result<Option<CompactResponse>> {
1465 Err(self.error())
1466 }
1467}
1468
1469struct RequestOptionsDriver {
1477 inner: Arc<dyn ChatDriver>,
1478 options: crate::provider::ProviderRequestOptions,
1479}
1480
1481impl RequestOptionsDriver {
1482 fn wrap(
1485 driver: BoxedChatDriver,
1486 options: &crate::provider::ProviderRequestOptions,
1487 ) -> BoxedChatDriver {
1488 if options.is_empty() {
1489 return driver;
1490 }
1491 Box::new(Self {
1492 inner: Arc::from(driver),
1493 options: options.clone(),
1494 })
1495 }
1496
1497 fn apply(&self, config: &LlmCallConfig) -> LlmCallConfig {
1498 let mut config = config.clone();
1499 config.extra_headers.extend(self.options.header_pairs());
1500 if self.options.cache_diagnostics {
1501 config.cache_diagnostics = Some(CacheDiagnosticsConfig {
1502 enabled: true,
1503 previous_message_id: config.previous_response_id.clone(),
1507 });
1508 }
1509 config
1510 }
1511}
1512
1513#[async_trait]
1514impl ChatDriver for RequestOptionsDriver {
1515 fn native_async_driver(
1516 &self,
1517 model: &str,
1518 tools: std::collections::BTreeMap<String, Option<serde_json::Value>>,
1519 continuation: Option<crate::native_async::Delivery>,
1520 ) -> Option<Arc<dyn ChatDriver>> {
1521 Some(Arc::new(Self {
1522 inner: self.inner.native_async_driver(model, tools, continuation)?,
1523 options: self.options.clone(),
1524 }))
1525 }
1526 async fn chat_completion_stream(
1527 &self,
1528 endpoint: &crate::runtime_provider::ProviderEndpoint,
1529 messages: Vec<LlmMessage>,
1530 config: &LlmCallConfig,
1531 ) -> Result<LlmResponseStream> {
1532 self.inner
1533 .chat_completion_stream(endpoint, messages, &self.apply(config))
1534 .await
1535 }
1536
1537 async fn chat_completion(
1538 &self,
1539 endpoint: &crate::runtime_provider::ProviderEndpoint,
1540 messages: Vec<LlmMessage>,
1541 config: &LlmCallConfig,
1542 ) -> Result<LlmResponse> {
1543 self.inner
1544 .chat_completion(endpoint, messages, &self.apply(config))
1545 .await
1546 }
1547
1548 fn supports_native_non_streaming(&self) -> bool {
1549 self.inner.supports_native_non_streaming()
1550 }
1551
1552 async fn chat_completion_non_streaming(
1553 &self,
1554 endpoint: &crate::runtime_provider::ProviderEndpoint,
1555 messages: Vec<LlmMessage>,
1556 config: &LlmCallConfig,
1557 ) -> Result<LlmResponse> {
1558 self.inner
1559 .chat_completion_non_streaming(endpoint, messages, &self.apply(config))
1560 .await
1561 }
1562
1563 async fn list_models(
1564 &self,
1565 endpoint: &crate::runtime_provider::ProviderEndpoint,
1566 ) -> Result<Option<Vec<DiscoveredModel>>> {
1567 self.inner.list_models(endpoint).await
1568 }
1569
1570 fn supports_compact(&self) -> bool {
1571 self.inner.supports_compact()
1572 }
1573
1574 fn supports_stateful_responses(&self) -> bool {
1575 self.inner.supports_stateful_responses()
1576 }
1577
1578 fn effective_context_window(&self, model: &str) -> Option<usize> {
1579 self.inner.effective_context_window(model)
1580 }
1581
1582 fn supports_parallel_tool_calls(&self, model: &str) -> bool {
1583 self.inner.supports_parallel_tool_calls(model)
1584 }
1585
1586 async fn compact(
1587 &self,
1588 endpoint: &crate::runtime_provider::ProviderEndpoint,
1589 request: CompactRequest,
1590 ) -> Result<Option<CompactResponse>> {
1591 self.inner.compact(endpoint, request).await
1592 }
1593}
1594
1595pub use everruns_model_profiles::ServiceKind;
1604
1605#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1618pub enum DriverOAuthFlow {
1619 OpenRouterPkce,
1627}
1628
1629#[derive(Debug, Clone)]
1634pub struct DriverOAuthConfig {
1635 pub authorize_url: String,
1637 pub token_url: String,
1639 pub flow: DriverOAuthFlow,
1641}
1642
1643impl DriverOAuthConfig {
1644 pub fn openrouter() -> Self {
1646 Self {
1647 authorize_url: "https://openrouter.ai/auth".to_string(),
1648 token_url: "https://openrouter.ai/api/v1/auth/keys".to_string(),
1649 flow: DriverOAuthFlow::OpenRouterPkce,
1650 }
1651 }
1652}
1653
1654#[derive(Clone)]
1661pub struct DriverDescriptor {
1662 pub id: DriverId,
1664 pub display_name: String,
1666 pub services: Vec<ServiceKind>,
1668 pub credential_schema: CredentialFormSchema,
1670 pub base_url_env: Option<String>,
1681 pub oauth: Option<DriverOAuthConfig>,
1684 pub chat: Option<DriverFactory>,
1686 pub embeddings: Option<EmbeddingsDriverFactory>,
1688}
1689
1690impl DriverDescriptor {
1691 pub fn chat_only<F>(id: impl Into<DriverId>, factory: F) -> Self
1696 where
1697 F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
1698 {
1699 let id = id.into();
1700 Self {
1701 display_name: default_display_name(&id),
1702 credential_schema: default_credential_schema(&id),
1703 base_url_env: None,
1704 services: vec![ServiceKind::Chat],
1705 oauth: None,
1706 chat: Some(Arc::new(factory)),
1707 embeddings: None,
1708 id,
1709 }
1710 }
1711
1712 pub fn with_base_url_env(mut self, base_url_env: impl Into<String>) -> Self {
1714 self.base_url_env = Some(base_url_env.into());
1715 self
1716 }
1717
1718 pub fn supports(&self, service: ServiceKind) -> bool {
1720 self.services.contains(&service)
1721 }
1722
1723 pub fn declared_env_vars(&self) -> Vec<String> {
1730 self.credential_schema
1731 .fields
1732 .iter()
1733 .flat_map(|field| field.env.iter().cloned())
1734 .chain(self.base_url_env.clone())
1735 .collect()
1736 }
1737
1738 pub fn base_url_from_env<F>(&self, lookup: F) -> Option<String>
1741 where
1742 F: Fn(&str) -> Option<String>,
1743 {
1744 self.base_url_env
1745 .as_deref()
1746 .and_then(lookup)
1747 .filter(|value| !value.trim().is_empty())
1748 }
1749}
1750
1751impl std::fmt::Debug for DriverDescriptor {
1752 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1753 f.debug_struct("DriverDescriptor")
1754 .field("id", &self.id)
1755 .field("display_name", &self.display_name)
1756 .field("services", &self.services)
1757 .field("oauth", &self.oauth.is_some())
1758 .field("chat", &self.chat.is_some())
1759 .field("embeddings", &self.embeddings.is_some())
1760 .finish()
1761 }
1762}
1763
1764fn default_display_name(id: &DriverId) -> String {
1765 id.as_str().replace(['_', '-'], " ")
1766}
1767
1768fn default_credential_schema(id: &DriverId) -> CredentialFormSchema {
1769 if id == &DriverId::LlmSim {
1770 CredentialFormSchema::empty()
1771 } else {
1772 CredentialFormSchema {
1777 fields: vec![
1778 crate::credential_schema::FormField::password("api_key", "API Key").required(),
1779 ],
1780 instructions_markdown: String::new(),
1781 }
1782 }
1783}
1784
1785#[derive(Clone, Default)]
1805pub struct DriverRegistry {
1806 descriptors: HashMap<DriverId, DriverDescriptor>,
1807 providers: crate::runtime_provider::RuntimeProviderRegistry,
1808}
1809
1810impl DriverRegistry {
1811 pub fn new() -> Self {
1813 Self {
1814 descriptors: HashMap::new(),
1815 providers: crate::runtime_provider::RuntimeProviderRegistry::new(),
1816 }
1817 }
1818
1819 pub fn register_provider(
1821 &mut self,
1822 provider: crate::runtime_provider::RuntimeProvider,
1823 ) -> Result<()> {
1824 self.providers.register(provider)
1825 }
1826
1827 pub fn replace_provider(
1829 &mut self,
1830 provider: crate::runtime_provider::RuntimeProvider,
1831 ) -> Option<Arc<crate::runtime_provider::RuntimeProvider>> {
1832 self.providers.replace(provider)
1833 }
1834
1835 pub fn provider(
1837 &self,
1838 id: &crate::runtime_provider::ProviderKey,
1839 ) -> Option<Arc<crate::runtime_provider::RuntimeProvider>> {
1840 self.providers.get(id)
1841 }
1842
1843 pub fn register_descriptor(&mut self, descriptor: DriverDescriptor) {
1849 if self.descriptors.contains_key(&descriptor.id) {
1850 panic!(
1851 "driver already registered for provider '{}'; \
1852 use register_descriptor_or_replace to overwrite intentionally",
1853 descriptor.id
1854 );
1855 }
1856 self.descriptors.insert(descriptor.id.clone(), descriptor);
1857 }
1858
1859 pub fn register_descriptor_or_replace(&mut self, descriptor: DriverDescriptor) {
1861 self.descriptors.insert(descriptor.id.clone(), descriptor);
1862 }
1863
1864 pub fn register<F>(&mut self, provider_type: impl Into<DriverId>, factory: F)
1870 where
1871 F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
1872 {
1873 self.register_descriptor(DriverDescriptor::chat_only(provider_type, factory));
1874 }
1875
1876 pub fn register_or_replace<F>(&mut self, provider_type: impl Into<DriverId>, factory: F)
1881 where
1882 F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
1883 {
1884 self.register_descriptor_or_replace(DriverDescriptor::chat_only(provider_type, factory));
1885 }
1886
1887 pub fn register_external<F>(&mut self, id: impl AsRef<str>, factory: F)
1892 where
1893 F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
1894 {
1895 let mut descriptor = DriverDescriptor::chat_only(DriverId::external(id), factory);
1896 descriptor.credential_schema = CredentialFormSchema::empty();
1897 self.register_descriptor(descriptor);
1898 }
1899
1900 pub fn create_chat_driver(&self, config: &ProviderConfig) -> Result<BoxedChatDriver> {
1910 if let Some(provider) = self.providers.get(&config.provider) {
1911 return Ok(RequestOptionsDriver::wrap(
1912 (*provider).clone().into_boxed_driver(),
1913 &config.request_options,
1914 ));
1915 }
1916 let descriptor = self.descriptors.get(&config.provider_type).ok_or_else(|| {
1917 AgentLoopError::driver_not_registered(config.provider_type.to_string())
1918 })?;
1919 let factory = descriptor.chat.as_ref().ok_or_else(|| {
1921 AgentLoopError::llm(format!(
1922 "Provider driver '{}' does not implement the chat service.",
1923 config.provider_type
1924 ))
1925 })?;
1926
1927 let driver_config = DriverConfig::from_provider_config(config);
1929 let driver = factory(&driver_config);
1930 let mut credential_fields = driver_config.credentials.clone();
1931 if let Some(serde_json::Value::Object(extra)) = &driver_config.metadata.extra {
1932 for (name, value) in extra {
1933 if let Some(value) = value.as_str() {
1934 credential_fields
1935 .entry(name.clone())
1936 .or_insert_with(|| value.to_string());
1937 }
1938 }
1939 }
1940 let credential_errors = descriptor.credential_schema.validate(&credential_fields);
1941 if credential_errors.is_empty() {
1942 Ok(RequestOptionsDriver::wrap(driver, &config.request_options))
1943 } else {
1944 let message = if descriptor.credential_schema.fields.len() == 1
1945 && descriptor.credential_schema.fields[0].name == "api_key"
1946 {
1947 "API key is required. Configure the API key in provider settings.".to_string()
1948 } else {
1949 format!(
1950 "Provider credentials are required. Configure provider settings: {}",
1951 credential_errors.join(" ")
1952 )
1953 };
1954 Ok(Box::new(CredentialGateDriver {
1955 inner: driver,
1956 message,
1957 }))
1958 }
1959 }
1960
1961 pub fn has_driver(&self, provider_type: &DriverId) -> bool {
1963 self.descriptors.contains_key(provider_type)
1964 }
1965
1966 pub fn descriptor(&self, provider_type: &DriverId) -> Option<&DriverDescriptor> {
1968 self.descriptors.get(provider_type)
1969 }
1970
1971 pub fn supports(&self, provider_type: &DriverId, service: ServiceKind) -> bool {
1973 self.descriptors
1974 .get(provider_type)
1975 .is_some_and(|d| d.supports(service))
1976 }
1977
1978 pub fn providers_for(&self, service: ServiceKind) -> Vec<DriverId> {
1980 self.descriptors
1981 .values()
1982 .filter(|d| d.supports(service))
1983 .map(|d| d.id.clone())
1984 .collect()
1985 }
1986
1987 pub fn registered_providers(&self) -> Vec<DriverId> {
1989 self.descriptors.keys().cloned().collect()
1990 }
1991
1992 pub fn registered_provider_ids(&self) -> Vec<String> {
1994 self.providers.ids()
1995 }
1996
1997 pub fn create_embeddings_driver(
2005 &self,
2006 config: &ProviderConfig,
2007 ) -> std::result::Result<BoxedEmbeddingsDriver, EmbeddingsDriverError> {
2008 let requires_api_key = config.provider_type != DriverId::LlmSim;
2009 if requires_api_key && config.api_key.is_none() {
2010 return Err(EmbeddingsDriverError::Provider(
2011 "API key is required. Configure the API key in provider settings.".to_string(),
2012 ));
2013 }
2014 let descriptor = self.descriptors.get(&config.provider_type).ok_or_else(|| {
2015 EmbeddingsDriverError::Provider(format!(
2016 "No driver registered for provider '{}'",
2017 config.provider_type
2018 ))
2019 })?;
2020 let factory = descriptor.embeddings.as_ref().ok_or_else(|| {
2021 EmbeddingsDriverError::Provider(format!(
2022 "Provider driver '{}' does not implement the embeddings service.",
2023 config.provider_type
2024 ))
2025 })?;
2026 let driver_config = DriverConfig::from_provider_config(config);
2027 Ok(factory(&driver_config))
2028 }
2029}
2030
2031const MAX_TOOL_RESULT_BYTES: usize = 64 * 1024;
2036
2037const TRUNCATION_SUFFIX: &str =
2038 "\n\n[Output truncated — exceeded 64 KiB limit. Try quiet flags, pipes, or redirect to file.]";
2039
2040pub fn truncate_tool_result(text: String) -> String {
2041 if text.len() <= MAX_TOOL_RESULT_BYTES {
2042 return text;
2043 }
2044 let content_budget = MAX_TOOL_RESULT_BYTES.saturating_sub(TRUNCATION_SUFFIX.len());
2045 let mut end = content_budget;
2046 while end > 0 && !text.is_char_boundary(end) {
2047 end -= 1;
2048 }
2049 let mut truncated = text[..end].to_string();
2050 truncated.push_str(TRUNCATION_SUFFIX);
2051 truncated
2052}
2053
2054#[cfg(test)]
2059mod tests {
2060 use super::*;
2061 use crate::runtime_provider::ProviderEndpoint;
2062
2063 #[test]
2064 fn test_disjoint_prompt_tokens_subtracts_cached_subset() {
2065 assert_eq!(disjoint_prompt_tokens(1000, Some(800)), 200);
2068 assert_eq!(disjoint_prompt_tokens(1000, None), 1000);
2070 assert_eq!(disjoint_prompt_tokens(1000, Some(0)), 1000);
2071 assert_eq!(disjoint_prompt_tokens(800, Some(1000)), 0);
2073 }
2074
2075 fn bare_call_config() -> LlmCallConfig {
2078 LlmCallConfig {
2079 model: "claude-opus-4-8".to_string(),
2080 temperature: None,
2081 max_tokens: None,
2082 tools: vec![],
2083 reasoning_effort: None,
2084 speed: None,
2085 verbosity: None,
2086 metadata: HashMap::new(),
2087 previous_response_id: None,
2088 provider_opaque_context: None,
2089 tool_search: None,
2090 prompt_cache: None,
2091 driver_options: Default::default(),
2092 parallel_tool_calls: None,
2093 volatile_suffix_len: 0,
2094 extra_headers: Vec::new(),
2095 cache_diagnostics: None,
2096 reasoning_state: None,
2097 }
2098 }
2099
2100 #[test]
2101 fn provider_config_debug_redacts_runtime_values() {
2102 let config = ProviderConfig::new(DriverId::OpenAI)
2103 .with_api_key("secret-key")
2104 .with_base_url("https://user:password@example.test/v1?token=secret")
2105 .with_metadata(ProviderMetadata {
2106 refresh_token: Some("refresh-secret".into()),
2107 account_id: Some("account-1".into()),
2108 extra: Some(serde_json::json!({ "client_secret": "metadata-secret" })),
2109 });
2110 let debug = format!("{config:?}");
2111 assert!(debug.contains("ProviderConfig"));
2112 assert!(debug.contains("openai"));
2113 assert!(debug.contains("<configured>"));
2114 for secret in [
2115 "secret-key",
2116 "password",
2117 "token=secret",
2118 "refresh-secret",
2119 "metadata-secret",
2120 ] {
2121 assert!(!debug.contains(secret), "debug output exposed {secret}");
2122 }
2123 }
2124
2125 #[test]
2126 fn system_messages_fold_only_system_text_in_transcript_order() {
2127 use LlmMessageRole::{Assistant, System, Tool, User};
2128 for (messages, expected) in [
2129 (vec![], None),
2130 (
2131 vec![
2132 LlmMessage::text(User, "user"),
2133 LlmMessage::text(Assistant, "answer"),
2134 LlmMessage::text(Tool, "result"),
2135 ],
2136 None,
2137 ),
2138 (vec![LlmMessage::text(System, "")], Some("")),
2139 (
2140 vec![
2141 LlmMessage::text(System, "rules"),
2142 LlmMessage::text(User, "question"),
2143 ],
2144 Some("rules"),
2145 ),
2146 (
2147 vec![
2148 LlmMessage::text(System, "first"),
2149 LlmMessage::text(User, "question"),
2150 LlmMessage::text(System, "second"),
2151 LlmMessage::text(Assistant, "answer"),
2152 LlmMessage::text(System, "third"),
2153 ],
2154 Some("first\n\nsecond\n\nthird"),
2155 ),
2156 (
2157 vec![
2158 LlmMessage::parts(
2159 System,
2160 vec![
2161 LlmContentPart::text("foo"),
2162 LlmContentPart::image("image"),
2163 LlmContentPart::audio("audio"),
2164 LlmContentPart::text("bar"),
2165 ],
2166 ),
2167 LlmMessage::text(System, "next"),
2168 ],
2169 Some("foobar\n\nnext"),
2170 ),
2171 ] {
2172 assert_eq!(fold_system_messages(&messages).as_deref(), expected);
2173 }
2174 }
2175
2176 #[test]
2177 fn prefix_preserves_all_media_and_changes_only_the_first_text_part() {
2178 let mut plain = LlmMessage::text(LlmMessageRole::User, "Hello");
2179 plain.prepend_text_prefix("[Alice] ");
2180 assert!(
2181 matches!(plain.content, LlmMessageContent::Text(ref text) if text == "[Alice] Hello")
2182 );
2183 for (parts, expected) in [
2184 (vec![], vec![("text", "[Alice] ")]),
2185 (
2186 vec![
2187 LlmContentPart::image("image"),
2188 LlmContentPart::audio("audio"),
2189 ],
2190 vec![("text", "[Alice] "), ("image", "image"), ("audio", "audio")],
2191 ),
2192 (
2193 vec![
2194 LlmContentPart::text("Hello"),
2195 LlmContentPart::image("image"),
2196 ],
2197 vec![("text", "[Alice] Hello"), ("image", "image")],
2198 ),
2199 (
2200 vec![
2201 LlmContentPart::image("image"),
2202 LlmContentPart::text("Hello"),
2203 LlmContentPart::audio("audio"),
2204 LlmContentPart::text("later"),
2205 ],
2206 vec![
2207 ("image", "image"),
2208 ("text", "[Alice] Hello"),
2209 ("audio", "audio"),
2210 ("text", "later"),
2211 ],
2212 ),
2213 (
2214 vec![LlmContentPart::text(""), LlmContentPart::text("later")],
2215 vec![("text", "[Alice] "), ("text", "later")],
2216 ),
2217 ] {
2218 let mut message = LlmMessage::parts(LlmMessageRole::Tool, parts);
2219 message.tool_call_id = Some("call-1".into());
2220 message.prepend_text_prefix("[Alice] ");
2221 let LlmMessageContent::Parts(parts) = &message.content else {
2222 panic!("parts must remain parts")
2223 };
2224 let actual: Vec<_> = parts
2225 .iter()
2226 .map(|part| match part {
2227 LlmContentPart::Text { text } => ("text", text.as_str()),
2228 LlmContentPart::Image { url } => ("image", url.as_str()),
2229 LlmContentPart::Audio { url } => ("audio", url.as_str()),
2230 LlmContentPart::File { url, .. } => ("file", url.as_str()),
2231 })
2232 .collect();
2233 assert_eq!(actual, expected);
2234 assert_eq!(message.role, LlmMessageRole::Tool);
2235 assert_eq!(message.tool_call_id.as_deref(), Some("call-1"));
2236 }
2237 }
2238 struct FixtureDriver(&'static str);
2239
2240 #[async_trait]
2241 impl ChatDriver for FixtureDriver {
2242 async fn chat_completion_stream(
2243 &self,
2244 _: &ProviderEndpoint,
2245 _: Vec<LlmMessage>,
2246 _: &LlmCallConfig,
2247 ) -> Result<LlmResponseStream> {
2248 Ok(Box::pin(futures::stream::iter([
2249 Ok(LlmStreamEvent::TextDelta(self.0.into())),
2250 Ok(LlmStreamEvent::Done(Box::default())),
2251 ])))
2252 }
2253 async fn list_models(&self, _: &ProviderEndpoint) -> Result<Option<Vec<DiscoveredModel>>> {
2254 Ok(Some(vec![DiscoveredModel {
2255 model_id: self.0.into(),
2256 display_name: None,
2257 created_at: None,
2258 owned_by: None,
2259 capabilities: vec!["chat".into()],
2260 discovered_profile: None,
2261 }]))
2262 }
2263 async fn compact(
2264 &self,
2265 _: &ProviderEndpoint,
2266 request: CompactRequest,
2267 ) -> Result<Option<CompactResponse>> {
2268 Ok(Some(CompactResponse {
2269 output: vec![crate::compact::CompactOutputItem::Compaction {
2270 encrypted_content: request.model,
2271 }],
2272 usage: None,
2273 }))
2274 }
2275 fn supports_compact(&self) -> bool {
2276 true
2277 }
2278 fn supports_stateful_responses(&self) -> bool {
2279 true
2280 }
2281 fn effective_context_window(&self, model: &str) -> Option<usize> {
2282 (model == "known").then_some(12345)
2283 }
2284 fn supports_parallel_tool_calls(&self, model: &str) -> bool {
2285 model == "known"
2286 }
2287 }
2288
2289 fn compact_fixture() -> CompactRequest {
2290 CompactRequest {
2291 reasoning_state: None,
2292 model: "compact-model".into(),
2293 input: vec![],
2294 previous_response_id: None,
2295 instructions: None,
2296 }
2297 }
2298
2299 #[tokio::test]
2300 async fn default_and_boxed_drivers_preserve_optional_operations_and_model_capabilities() {
2301 struct DefaultDriver;
2302 #[async_trait]
2303 impl ChatDriver for DefaultDriver {
2304 async fn chat_completion_stream(
2305 &self,
2306 _: &ProviderEndpoint,
2307 _: Vec<LlmMessage>,
2308 _: &LlmCallConfig,
2309 ) -> Result<LlmResponseStream> {
2310 Ok(Box::pin(futures::stream::empty()))
2311 }
2312 }
2313 let endpoint = ProviderEndpoint::default();
2314 assert!(!DefaultDriver.supports_compact());
2315 assert!(!DefaultDriver.supports_stateful_responses());
2316 assert!(!DefaultDriver.supports_parallel_tool_calls("known"));
2317 assert_eq!(DefaultDriver.effective_context_window("known"), None);
2318 assert!(
2319 DefaultDriver
2320 .list_models(&endpoint)
2321 .await
2322 .unwrap()
2323 .is_none()
2324 );
2325 assert!(
2326 DefaultDriver
2327 .compact(&endpoint, compact_fixture())
2328 .await
2329 .unwrap()
2330 .is_none()
2331 );
2332 let boxed: BoxedChatDriver = Box::new(FixtureDriver("boxed"));
2333 assert!(boxed.supports_compact());
2334 assert!(boxed.supports_stateful_responses());
2335 for (model, expected) in [("known", true), ("unknown", false)] {
2336 assert_eq!(boxed.supports_parallel_tool_calls(model), expected);
2337 assert_eq!(
2338 boxed.effective_context_window(model),
2339 expected.then_some(12345)
2340 );
2341 }
2342 assert_eq!(
2343 boxed
2344 .chat_completion(&endpoint, vec![], &bare_call_config())
2345 .await
2346 .unwrap()
2347 .text,
2348 "boxed"
2349 );
2350 }
2351
2352 #[tokio::test]
2353 async fn registry_replacement_changes_factory_and_preserves_other_descriptors() {
2354 let mut registry = DriverRegistry::new();
2355 assert!(registry.registered_providers().is_empty());
2356 registry.register(DriverId::LlmSim, |_| Box::new(FixtureDriver("first")));
2357 registry.register_descriptor(DriverDescriptor {
2358 display_name: "OpenAI custom".into(),
2359 services: vec![ServiceKind::Chat, ServiceKind::Realtime],
2360 ..DriverDescriptor::chat_only(DriverId::OpenAI, |_| Box::new(FixtureDriver("other")))
2361 });
2362 let config = ProviderConfig::new(DriverId::LlmSim);
2363 let endpoint = ProviderEndpoint::default();
2364 assert_eq!(
2365 registry
2366 .create_chat_driver(&config)
2367 .unwrap()
2368 .chat_completion(&endpoint, vec![], &bare_call_config())
2369 .await
2370 .unwrap()
2371 .text,
2372 "first"
2373 );
2374 registry.register_or_replace(DriverId::LlmSim, |_| Box::new(FixtureDriver("replacement")));
2375 assert_eq!(
2376 registry
2377 .create_chat_driver(&config)
2378 .unwrap()
2379 .chat_completion(&endpoint, vec![], &bare_call_config())
2380 .await
2381 .unwrap()
2382 .text,
2383 "replacement"
2384 );
2385 assert!(registry.has_driver(&DriverId::LlmSim));
2386 assert!(!registry.has_driver(&DriverId::Anthropic));
2387 assert_eq!(
2388 registry.providers_for(ServiceKind::Realtime),
2389 vec![DriverId::OpenAI]
2390 );
2391 let mut chat = registry.providers_for(ServiceKind::Chat);
2392 chat.sort_by_key(|id| id.to_string());
2393 assert_eq!(chat, vec![DriverId::LlmSim, DriverId::OpenAI]);
2394 assert!(registry.supports(&DriverId::OpenAI, ServiceKind::Realtime));
2395 assert!(!registry.supports(&DriverId::LlmSim, ServiceKind::Realtime));
2396 assert!(!registry.supports(&DriverId::Gemini, ServiceKind::Chat));
2397 assert_eq!(
2398 registry.descriptor(&DriverId::OpenAI).unwrap().display_name,
2399 "OpenAI custom"
2400 );
2401 assert_eq!(
2402 registry
2403 .create_chat_driver(
2404 &ProviderConfig::new(DriverId::OpenAI).with_api_key("synthetic-key")
2405 )
2406 .unwrap()
2407 .chat_completion(&endpoint, vec![], &bare_call_config())
2408 .await
2409 .unwrap()
2410 .text,
2411 "other"
2412 );
2413 let defaults = DriverDescriptor::chat_only(DriverId::Anthropic, |_| {
2414 Box::new(FixtureDriver("default"))
2415 });
2416 assert_eq!(defaults.display_name, "anthropic");
2417 let sim = registry.descriptor(&DriverId::LlmSim).unwrap();
2418 assert!(sim.credential_schema.fields.is_empty());
2419 assert_eq!(sim.services, vec![ServiceKind::Chat]);
2420 assert!(sim.chat.is_some());
2421 let real = registry.descriptor(&DriverId::OpenAI).unwrap();
2422 assert_eq!(real.credential_schema.fields.len(), 1);
2423 assert_eq!(real.credential_schema.fields[0].name, "api_key");
2424 assert!(real.credential_schema.fields[0].required);
2425 assert!(registry.descriptor(&DriverId::Gemini).is_none());
2426 }
2427
2428 #[test]
2429 #[should_panic(expected = "already registered")]
2430 fn duplicate_registration_rejects_an_existing_driver() {
2431 let mut registry = DriverRegistry::new();
2432 registry.register(DriverId::OpenAI, |_| Box::new(FixtureDriver("first")));
2433 registry.register(DriverId::OpenAI, |_| Box::new(FixtureDriver("second")));
2434 }
2435
2436 #[tokio::test]
2437 async fn factory_receives_complete_config_and_external_metadata_auth_remains_keyless() {
2438 let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
2439 let capture = seen.clone();
2440 let mut registry = DriverRegistry::new();
2441 registry.register_external("CUSTOM", move |config| {
2442 capture.lock().unwrap().push(config.clone());
2443 Box::new(FixtureDriver("external"))
2444 });
2445 let metadata = ProviderMetadata {
2446 refresh_token: Some("refresh".into()),
2447 account_id: Some("account".into()),
2448 extra: Some(serde_json::json!({"region":"west"})),
2449 };
2450 for key in [None, Some("synthetic-key")] {
2451 let mut config =
2452 ProviderConfig::for_provider("connection", DriverId::external("custom"))
2453 .with_base_url("https://gateway.example/v1")
2454 .with_metadata(metadata.clone());
2455 if let Some(key) = key {
2456 config = config.with_api_key(key);
2457 }
2458 let response = registry
2459 .create_chat_driver(&config)
2460 .unwrap()
2461 .chat_completion(&ProviderEndpoint::default(), vec![], &bare_call_config())
2462 .await
2463 .unwrap();
2464 assert_eq!(response.text, "external");
2465 let received = seen.lock().unwrap().pop().unwrap();
2466 assert_eq!(received.provider.as_str(), "connection");
2467 assert_eq!(received.provider_type, DriverId::external("custom"));
2468 assert_eq!(received.api_key.as_deref(), key);
2469 assert_eq!(received.credential("api_key"), key);
2470 assert_eq!(received.credentials.len(), usize::from(key.is_some()));
2471 assert_eq!(
2472 received.base_url.as_deref(),
2473 Some("https://gateway.example/v1")
2474 );
2475 assert_eq!(received.metadata, metadata);
2476 }
2477 assert!(
2478 registry
2479 .descriptor(&DriverId::external("custom"))
2480 .unwrap()
2481 .credential_schema
2482 .fields
2483 .is_empty()
2484 );
2485 }
2486
2487 #[test]
2488 fn registry_distinguishes_missing_driver_from_missing_chat_service() {
2489 let mut registry = DriverRegistry::new();
2490 assert!(
2491 matches!(registry.create_chat_driver(&ProviderConfig::new(DriverId::Anthropic)), Err(AgentLoopError::DriverNotRegistered(id)) if id == "anthropic")
2492 );
2493 registry.register_descriptor(DriverDescriptor {
2494 id: DriverId::external("embeddings-only"),
2495 display_name: "Embeddings Only".into(),
2496 services: vec![ServiceKind::Embeddings],
2497 credential_schema: CredentialFormSchema::empty(),
2498 base_url_env: None,
2499 oauth: None,
2500 chat: None,
2501 embeddings: None,
2502 });
2503 match registry
2504 .create_chat_driver(&ProviderConfig::new(DriverId::external("embeddings-only")))
2505 {
2506 Err(AgentLoopError::Llm(error)) => assert_eq!(
2507 error.message,
2508 "Provider driver 'embeddings-only' does not implement the chat service."
2509 ),
2510 _ => panic!("expected a missing-chat-service error"),
2511 }
2512 }
2513
2514 #[tokio::test]
2515 async fn credential_gate_rejects_every_io_operation_before_dispatch() {
2516 struct ForbiddenDriver;
2517 #[async_trait]
2518 impl ChatDriver for ForbiddenDriver {
2519 async fn chat_completion_stream(
2520 &self,
2521 _: &ProviderEndpoint,
2522 _: Vec<LlmMessage>,
2523 _: &LlmCallConfig,
2524 ) -> Result<LlmResponseStream> {
2525 panic!("unauthenticated stream dispatch")
2526 }
2527 async fn list_models(
2528 &self,
2529 _: &ProviderEndpoint,
2530 ) -> Result<Option<Vec<DiscoveredModel>>> {
2531 panic!("unauthenticated model dispatch")
2532 }
2533 async fn compact(
2534 &self,
2535 _: &ProviderEndpoint,
2536 _: CompactRequest,
2537 ) -> Result<Option<CompactResponse>> {
2538 panic!("unauthenticated compact dispatch")
2539 }
2540 }
2541 let mut registry = DriverRegistry::new();
2542 registry.register(DriverId::OpenAI, |config| {
2543 if config.api_key.is_some() {
2544 Box::new(FixtureDriver("authenticated"))
2545 } else {
2546 Box::new(ForbiddenDriver)
2547 }
2548 });
2549 let driver = registry
2550 .create_chat_driver(&ProviderConfig::new(DriverId::OpenAI))
2551 .unwrap();
2552 let endpoint = ProviderEndpoint::default();
2553 let stream_error = match driver
2554 .chat_completion_stream(&endpoint, vec![], &bare_call_config())
2555 .await
2556 {
2557 Err(error) => error,
2558 Ok(_) => panic!("expected authentication error"),
2559 };
2560 for error in [
2561 stream_error,
2562 driver
2563 .chat_completion(&endpoint, vec![], &bare_call_config())
2564 .await
2565 .unwrap_err(),
2566 driver.list_models(&endpoint).await.unwrap_err(),
2567 driver
2568 .compact(&endpoint, compact_fixture())
2569 .await
2570 .unwrap_err(),
2571 ] {
2572 assert_eq!(error.llm_error_kind(), Some(LlmErrorKind::Authentication));
2573 assert_eq!(
2574 error.to_string(),
2575 "LLM error: API key is required. Configure the API key in provider settings."
2576 );
2577 }
2578 let driver = registry
2579 .create_chat_driver(
2580 &ProviderConfig::new(DriverId::OpenAI).with_api_key("synthetic-key"),
2581 )
2582 .unwrap();
2583 assert_eq!(
2584 driver
2585 .chat_completion(&endpoint, vec![], &bare_call_config())
2586 .await
2587 .unwrap()
2588 .text,
2589 "authenticated"
2590 );
2591 assert_eq!(
2592 driver.list_models(&endpoint).await.unwrap().unwrap()[0].model_id,
2593 "authenticated"
2594 );
2595 assert_eq!(
2596 serde_json::to_value(
2597 driver
2598 .compact(&endpoint, compact_fixture())
2599 .await
2600 .unwrap()
2601 .unwrap()
2602 .output
2603 )
2604 .unwrap(),
2605 serde_json::json!([{"type":"compaction","encrypted_content":"compact-model"}])
2606 );
2607 }
2608
2609 #[tokio::test]
2610 async fn request_options_preserve_calls_and_apply_headers_and_diagnostics_independently() {
2611 struct CapturingDriver(Arc<std::sync::Mutex<Vec<LlmCallConfig>>>);
2612 impl CapturingDriver {
2613 fn capture(
2614 &self,
2615 endpoint: &ProviderEndpoint,
2616 messages: &[LlmMessage],
2617 config: &LlmCallConfig,
2618 ) {
2619 assert_eq!(
2620 endpoint.url("probe").as_deref(),
2621 Some("https://gateway.example/v1/probe")
2622 );
2623 assert_eq!(messages.len(), 1);
2624 assert_eq!(messages[0].role, LlmMessageRole::User);
2625 assert_eq!(messages[0].content_as_text(), "request text");
2626 self.0.lock().unwrap().push(config.clone());
2627 }
2628 }
2629 #[async_trait]
2630 impl ChatDriver for CapturingDriver {
2631 async fn chat_completion_stream(
2632 &self,
2633 endpoint: &ProviderEndpoint,
2634 messages: Vec<LlmMessage>,
2635 config: &LlmCallConfig,
2636 ) -> Result<LlmResponseStream> {
2637 self.capture(endpoint, &messages, config);
2638 FixtureDriver("stream")
2639 .chat_completion_stream(endpoint, messages, config)
2640 .await
2641 }
2642 async fn chat_completion(
2643 &self,
2644 endpoint: &ProviderEndpoint,
2645 messages: Vec<LlmMessage>,
2646 config: &LlmCallConfig,
2647 ) -> Result<LlmResponse> {
2648 self.capture(endpoint, &messages, config);
2649 FixtureDriver("completion")
2650 .chat_completion(endpoint, messages, config)
2651 .await
2652 }
2653 }
2654 let provider = crate::Provider::new("fixture", FixtureDriver("endpoint"))
2655 .base_url("https://gateway.example/v1");
2656 for (headers, diagnostics) in [(false, false), (true, false), (false, true), (true, true)] {
2657 let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
2658 let options = crate::provider::ProviderRequestOptions {
2659 headers: if headers {
2660 vec![crate::provider::ProviderRequestHeader {
2661 name: "x-base".into(),
2662 value: "connection".into(),
2663 }]
2664 } else {
2665 vec![]
2666 },
2667 cache_diagnostics: diagnostics,
2668 };
2669 let driver =
2670 RequestOptionsDriver::wrap(Box::new(CapturingDriver(seen.clone())), &options);
2671 let mut config = bare_call_config();
2672 config.model = "requested-model".into();
2673 config.temperature = Some(0.25);
2674 config.max_tokens = Some(42);
2675 config
2676 .metadata
2677 .insert("session_id".into(), "session-one".into());
2678 config.previous_response_id = Some("response-one".into());
2679 config.extra_headers = vec![("x-base".into(), "original".into())];
2680 config.cache_diagnostics = Some(CacheDiagnosticsConfig {
2681 enabled: false,
2682 previous_message_id: Some("existing".into()),
2683 });
2684 let mut stream = driver
2685 .chat_completion_stream(
2686 provider.endpoint(),
2687 vec![LlmMessage::text(LlmMessageRole::User, "request text")],
2688 &config,
2689 )
2690 .await
2691 .unwrap();
2692 use futures::StreamExt;
2693 assert!(
2694 matches!(stream.next().await.unwrap().unwrap(), LlmStreamEvent::TextDelta(text) if text == "stream")
2695 );
2696 assert!(matches!(
2697 stream.next().await.unwrap().unwrap(),
2698 LlmStreamEvent::Done(_)
2699 ));
2700 assert!(stream.next().await.is_none());
2701 assert_eq!(
2702 driver
2703 .chat_completion(
2704 provider.endpoint(),
2705 vec![LlmMessage::text(LlmMessageRole::User, "request text")],
2706 &config
2707 )
2708 .await
2709 .unwrap()
2710 .text,
2711 "completion"
2712 );
2713 let mut expected_headers = vec![("x-base".into(), "original".into())];
2714 if headers {
2715 expected_headers.push(("x-base".into(), "connection".into()));
2716 }
2717 let observed = seen.lock().unwrap();
2718 assert_eq!(observed.len(), 2);
2719 for received in observed.iter() {
2720 assert_eq!(received.extra_headers, expected_headers);
2721 let diagnostic = received.cache_diagnostics.as_ref().unwrap();
2722 assert_eq!(diagnostic.enabled, diagnostics);
2723 assert_eq!(
2724 diagnostic.previous_message_id.as_deref(),
2725 Some(if diagnostics {
2726 "response-one"
2727 } else {
2728 "existing"
2729 })
2730 );
2731 assert_eq!(received.model, "requested-model");
2732 assert_eq!(received.temperature, Some(0.25));
2733 assert_eq!(received.max_tokens, Some(42));
2734 assert_eq!(received.metadata, config.metadata);
2735 assert_eq!(received.previous_response_id, config.previous_response_id);
2736 }
2737 assert_eq!(
2738 config.extra_headers,
2739 vec![("x-base".into(), "original".into())]
2740 );
2741 assert!(!config.cache_diagnostics.as_ref().unwrap().enabled);
2742 assert_eq!(
2743 config
2744 .cache_diagnostics
2745 .as_ref()
2746 .unwrap()
2747 .previous_message_id
2748 .as_deref(),
2749 Some("existing")
2750 );
2751 }
2752 let options = crate::provider::ProviderRequestOptions {
2753 headers: vec![],
2754 cache_diagnostics: true,
2755 };
2756 let wrapped = RequestOptionsDriver::wrap(Box::new(FixtureDriver("forwarded")), &options);
2757 assert!(wrapped.supports_compact());
2758 assert!(wrapped.supports_stateful_responses());
2759 for (model, expected) in [("known", true), ("unknown", false)] {
2760 assert_eq!(wrapped.supports_parallel_tool_calls(model), expected);
2761 assert_eq!(
2762 wrapped.effective_context_window(model),
2763 expected.then_some(12345)
2764 );
2765 }
2766 assert_eq!(
2767 wrapped
2768 .list_models(provider.endpoint())
2769 .await
2770 .unwrap()
2771 .unwrap()[0]
2772 .model_id,
2773 "forwarded"
2774 );
2775 assert_eq!(
2776 serde_json::to_value(
2777 wrapped
2778 .compact(provider.endpoint(), compact_fixture())
2779 .await
2780 .unwrap()
2781 .unwrap()
2782 .output
2783 )
2784 .unwrap(),
2785 serde_json::json!([{"type":"compaction","encrypted_content":"compact-model"}])
2786 );
2787 }
2788}