use crate::compact::{CompactOutputItem, CompactRequest, CompactResponse};
use crate::credential_schema::CredentialFormSchema;
use crate::error::{AgentLoopError, LlmErrorKind, Result};
use crate::tool_types::{ToolCall, ToolDefinition};
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use futures::Stream;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::pin::Pin;
use std::sync::Arc;
pub type LlmResponseStream = Pin<Box<dyn Stream<Item = Result<LlmStreamEvent>> + Send>>;
#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ProviderOpaqueContext {
OpenResponsesCompact {
output: Vec<CompactOutputItem>,
#[serde(default, skip_serializing_if = "Option::is_none")]
reasoning_state: Option<crate::reasoning_updates::ReasoningState>,
},
}
impl std::fmt::Debug for ProviderOpaqueContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::OpenResponsesCompact { output, .. } => f
.debug_struct("OpenResponsesCompact")
.field("item_count", &output.len())
.finish_non_exhaustive(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LlmStreamError {
pub code: Option<String>,
pub status: Option<u16>,
pub message: String,
}
impl LlmStreamError {
pub fn new(message: impl Into<String>) -> Self {
Self {
code: None,
status: None,
message: message.into(),
}
}
pub fn provider(
code: Option<impl Into<String>>,
status: Option<u16>,
message: impl Into<String>,
) -> Self {
Self {
code: code.map(Into::into),
status,
message: message.into(),
}
}
pub fn kind(&self) -> LlmErrorKind {
if let Some(code) = self.code.as_deref()
&& let Some(kind) = LlmErrorKind::from_provider_code(code)
{
return kind;
}
if let Some(status) = self.status {
return LlmErrorKind::from_provider_status(status, &self.message);
}
LlmErrorKind::from_error_text(&self.message)
}
}
impl std::error::Error for LlmStreamError {}
impl std::fmt::Display for LlmStreamError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match (&self.code, self.status) {
(Some(code), Some(status)) => write!(f, "{code} ({status}): {}", self.message),
(Some(code), None) => write!(f, "{code}: {}", self.message),
(None, Some(status)) => write!(f, "({status}): {}", self.message),
(None, None) => f.write_str(&self.message),
}
}
}
impl From<String> for LlmStreamError {
fn from(message: String) -> Self {
Self::new(message)
}
}
impl From<&str> for LlmStreamError {
fn from(message: &str) -> Self {
Self::new(message)
}
}
#[derive(Debug, Clone)]
pub enum LlmStreamEvent {
TextDelta(String),
ReasoningDelta { delta: String, summary: bool },
ReasoningItem(crate::reasoning::ReasoningContentPart),
ToolCalls(Vec<ToolCall>),
NativeToolCall(crate::native_async::NativeToolCall),
MessagePhase(crate::execution_phase::ExecutionPhase),
Done(Box<LlmCompletionMetadata>),
Error(LlmStreamError),
}
#[derive(Debug, Clone)]
pub struct DiscoveredModel {
pub model_id: String,
pub display_name: Option<String>,
pub created_at: Option<DateTime<Utc>>,
pub owned_by: Option<String>,
pub capabilities: Vec<String>,
pub discovered_profile: Option<crate::model::ModelProfile>,
}
#[derive(Debug, Clone, Default)]
pub struct LlmCompletionMetadata {
pub total_tokens: Option<u32>,
pub prompt_tokens: Option<u32>,
pub completion_tokens: Option<u32>,
pub cache_read_tokens: Option<u32>,
pub cache_creation_tokens: Option<u32>,
pub provider_cost_usd: Option<f64>,
pub model: Option<String>,
pub finish_reason: Option<String>,
pub retry_metadata: Option<crate::llm_retry::RetryMetadata>,
pub response_id: Option<String>,
pub phase: Option<String>,
pub cache_diagnostics: Option<serde_json::Value>,
}
pub fn disjoint_prompt_tokens(reported_input: u32, cache_read: Option<u32>) -> u32 {
reported_input.saturating_sub(cache_read.unwrap_or(0))
}
#[async_trait]
pub trait ChatDriver: Send + Sync {
fn native_async_driver(
&self,
_model: &str,
_tools: std::collections::BTreeMap<String, Option<serde_json::Value>>,
_continuation: Option<crate::native_async::Delivery>,
) -> Option<Arc<dyn ChatDriver>> {
None
}
async fn chat_completion_stream(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponseStream>;
async fn chat_completion(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponse> {
use futures::StreamExt;
let mut stream = self
.chat_completion_stream(endpoint, messages, config)
.await?;
let mut text = String::new();
let mut reasoning: Vec<crate::reasoning::ReasoningContentPart> = Vec::new();
let mut tool_calls = Vec::new();
let mut metadata = LlmCompletionMetadata::default();
while let Some(event) = stream.next().await {
match event? {
LlmStreamEvent::TextDelta(delta) => text.push_str(&delta),
LlmStreamEvent::ReasoningDelta { .. } => {}
LlmStreamEvent::ReasoningItem(item) => reasoning.push(item),
LlmStreamEvent::ToolCalls(calls) => tool_calls = calls,
LlmStreamEvent::NativeToolCall(_) => {
return Err(crate::error::AgentLoopError::config(
"native async/custom calls require a streaming coordinator",
));
}
LlmStreamEvent::MessagePhase(_) => {}
LlmStreamEvent::Done(meta) => metadata = *meta,
LlmStreamEvent::Error(err) => {
return Err(crate::error::AgentLoopError::llm_kind(
err.kind(),
err.to_string(),
));
}
}
}
Ok(LlmResponse {
text,
reasoning,
tool_calls: if tool_calls.is_empty() {
None
} else {
Some(tool_calls)
},
metadata,
})
}
fn supports_native_non_streaming(&self) -> bool {
false
}
async fn chat_completion_non_streaming(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponse> {
self.chat_completion(endpoint, messages, config).await
}
async fn list_models(
&self,
_endpoint: &crate::runtime_provider::ProviderEndpoint,
) -> Result<Option<Vec<DiscoveredModel>>> {
Ok(None)
}
fn supports_compact(&self) -> bool {
false
}
fn supports_stateful_responses(&self) -> bool {
false
}
fn effective_context_window(&self, _model: &str) -> Option<usize> {
None
}
fn supports_parallel_tool_calls(&self, _model: &str) -> bool {
false
}
async fn compact(
&self,
_endpoint: &crate::runtime_provider::ProviderEndpoint,
_request: CompactRequest,
) -> Result<Option<CompactResponse>> {
Ok(None)
}
}
#[async_trait]
impl ChatDriver for Box<dyn ChatDriver> {
fn native_async_driver(
&self,
model: &str,
tools: std::collections::BTreeMap<String, Option<serde_json::Value>>,
continuation: Option<crate::native_async::Delivery>,
) -> Option<Arc<dyn ChatDriver>> {
(**self).native_async_driver(model, tools, continuation)
}
async fn chat_completion_stream(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponseStream> {
(**self)
.chat_completion_stream(endpoint, messages, config)
.await
}
async fn chat_completion(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponse> {
(**self).chat_completion(endpoint, messages, config).await
}
fn supports_native_non_streaming(&self) -> bool {
(**self).supports_native_non_streaming()
}
async fn chat_completion_non_streaming(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponse> {
(**self)
.chat_completion_non_streaming(endpoint, messages, config)
.await
}
async fn list_models(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
) -> Result<Option<Vec<DiscoveredModel>>> {
(**self).list_models(endpoint).await
}
fn supports_compact(&self) -> bool {
(**self).supports_compact()
}
fn supports_stateful_responses(&self) -> bool {
(**self).supports_stateful_responses()
}
fn effective_context_window(&self, model: &str) -> Option<usize> {
(**self).effective_context_window(model)
}
fn supports_parallel_tool_calls(&self, model: &str) -> bool {
(**self).supports_parallel_tool_calls(model)
}
async fn compact(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
request: CompactRequest,
) -> Result<Option<CompactResponse>> {
(**self).compact(endpoint, request).await
}
}
#[derive(Debug, Clone)]
pub struct LlmMessage {
pub native_tool_calls: Vec<crate::native_async::NativeToolCall>,
pub role: LlmMessageRole,
pub content: LlmMessageContent,
pub tool_calls: Option<Vec<ToolCall>>,
pub tool_call_id: Option<String>,
pub phase: Option<crate::execution_phase::ExecutionPhase>,
pub reasoning: Vec<crate::reasoning::ReasoningContentPart>,
pub configuration_update: Option<crate::model::ReasoningEffort>,
}
impl LlmMessage {
pub fn text(role: LlmMessageRole, content: impl Into<String>) -> Self {
Self {
native_tool_calls: Vec::new(),
role,
content: LlmMessageContent::Text(content.into()),
tool_calls: None,
tool_call_id: None,
phase: None,
reasoning: Vec::new(),
configuration_update: None,
}
}
pub fn parts(role: LlmMessageRole, parts: Vec<LlmContentPart>) -> Self {
Self {
native_tool_calls: Vec::new(),
role,
content: LlmMessageContent::Parts(parts),
tool_calls: None,
tool_call_id: None,
phase: None,
reasoning: Vec::new(),
configuration_update: None,
}
}
pub fn content_as_text(&self) -> String {
self.content.to_text()
}
pub fn prepend_text_prefix(&mut self, prefix: &str) {
match &mut self.content {
LlmMessageContent::Text(text) => {
*text = format!("{}{}", prefix, text);
}
LlmMessageContent::Parts(parts) => {
for part in parts.iter_mut() {
if let LlmContentPart::Text { text } = part {
*text = format!("{}{}", prefix, text);
return;
}
}
parts.insert(
0,
LlmContentPart::Text {
text: prefix.to_string(),
},
);
}
}
}
}
pub fn fold_system_messages(messages: &[LlmMessage]) -> Option<String> {
let mut system: Option<String> = None;
for msg in messages {
if msg.role == LlmMessageRole::System {
let text = msg.content.to_text();
system = Some(match system.take() {
Some(existing) if !existing.is_empty() => format!("{existing}\n\n{text}"),
_ => text,
});
}
}
system
}
#[derive(Debug, Clone)]
pub enum LlmMessageContent {
Text(String),
Parts(Vec<LlmContentPart>),
}
impl LlmMessageContent {
pub fn to_text(&self) -> String {
match self {
LlmMessageContent::Text(s) => s.clone(),
LlmMessageContent::Parts(parts) => parts
.iter()
.filter_map(|p| match p {
LlmContentPart::Text { text } => Some(text.clone()),
_ => None,
})
.collect::<Vec<_>>()
.join(""),
}
}
pub fn is_text(&self) -> bool {
matches!(self, LlmMessageContent::Text(_))
}
pub fn is_parts(&self) -> bool {
matches!(self, LlmMessageContent::Parts(_))
}
}
impl From<String> for LlmMessageContent {
fn from(s: String) -> Self {
LlmMessageContent::Text(s)
}
}
impl From<&str> for LlmMessageContent {
fn from(s: &str) -> Self {
LlmMessageContent::Text(s.to_string())
}
}
#[derive(Debug, Clone)]
pub enum LlmContentPart {
Text { text: String },
Image { url: String },
Audio { url: String },
File {
url: String,
filename: Option<String>,
},
}
impl LlmContentPart {
pub fn text(text: impl Into<String>) -> Self {
LlmContentPart::Text { text: text.into() }
}
pub fn image(url: impl Into<String>) -> Self {
LlmContentPart::Image { url: url.into() }
}
pub fn audio(url: impl Into<String>) -> Self {
LlmContentPart::Audio { url: url.into() }
}
pub fn file(url: impl Into<String>, filename: Option<String>) -> Self {
LlmContentPart::File {
url: url.into(),
filename,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum LlmMessageRole {
System,
User,
Assistant,
Tool,
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub struct ToolSearchConfig {
pub enabled: bool,
pub threshold: usize,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
#[serde(rename_all = "snake_case")]
pub enum PromptCacheStrategy {
#[default]
Auto,
Explicit,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
pub struct PromptCacheConfig {
pub enabled: bool,
#[serde(default)]
pub strategy: PromptCacheStrategy,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub gemini_cached_content: Option<String>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
pub struct CacheDiagnosticsConfig {
pub enabled: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub previous_message_id: Option<String>,
}
#[derive(Debug, Clone)]
pub struct LlmCallConfig {
pub reasoning_state: Option<crate::reasoning_updates::ReasoningState>,
pub model: String,
pub temperature: Option<f32>,
pub max_tokens: Option<u32>,
pub tools: Vec<ToolDefinition>,
pub reasoning_effort: Option<crate::model::ReasoningEffort>,
pub speed: Option<String>,
pub verbosity: Option<String>,
pub metadata: HashMap<String, String>,
pub previous_response_id: Option<String>,
pub provider_opaque_context: Option<ProviderOpaqueContext>,
pub tool_search: Option<ToolSearchConfig>,
pub prompt_cache: Option<PromptCacheConfig>,
pub driver_options: HashMap<String, serde_json::Value>,
pub parallel_tool_calls: Option<bool>,
pub volatile_suffix_len: usize,
pub extra_headers: Vec<(String, String)>,
pub cache_diagnostics: Option<CacheDiagnosticsConfig>,
}
impl LlmCallConfig {
pub fn resolved_parallel_tool_calls(&self, supported: bool) -> Option<bool> {
if supported {
self.parallel_tool_calls
} else {
None
}
}
}
#[derive(Debug, Clone)]
pub struct LlmResponse {
pub text: String,
pub reasoning: Vec<crate::reasoning::ReasoningContentPart>,
pub tool_calls: Option<Vec<ToolCall>>,
pub metadata: LlmCompletionMetadata,
}
pub struct LlmCallConfigBuilder {
config: LlmCallConfig,
}
impl LlmCallConfigBuilder {
pub fn from_config(config: LlmCallConfig) -> Self {
Self { config }
}
pub fn reasoning_effort(mut self, effort: crate::model::ReasoningEffort) -> Self {
self.config.reasoning_effort = Some(effort);
self
}
pub fn speed(mut self, speed: impl Into<String>) -> Self {
self.config.speed = Some(speed.into());
self
}
pub fn verbosity(mut self, verbosity: impl Into<String>) -> Self {
self.config.verbosity = Some(verbosity.into());
self
}
pub fn model(mut self, model: impl Into<String>) -> Self {
self.config.model = model.into();
self
}
pub fn temperature(mut self, temp: f32) -> Self {
self.config.temperature = Some(temp);
self
}
pub fn max_tokens(mut self, tokens: u32) -> Self {
self.config.max_tokens = Some(tokens);
self
}
pub fn tools(mut self, tools: Vec<ToolDefinition>) -> Self {
self.config.tools = tools;
self
}
pub fn metadata(mut self, metadata: HashMap<String, String>) -> Self {
self.config.metadata = metadata;
self
}
pub fn with_metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.config.metadata.insert(key.into(), value.into());
self
}
pub fn previous_response_id(mut self, id: Option<String>) -> Self {
self.config.previous_response_id = id;
self
}
pub fn provider_opaque_context(mut self, context: Option<ProviderOpaqueContext>) -> Self {
self.config.provider_opaque_context = context;
self
}
pub fn tool_search(mut self, config: ToolSearchConfig) -> Self {
self.config.tool_search = Some(config);
self
}
pub fn prompt_cache(mut self, config: PromptCacheConfig) -> Self {
self.config.prompt_cache = Some(config);
self
}
pub fn driver_option(mut self, key: impl Into<String>, value: serde_json::Value) -> Self {
self.config.driver_options.insert(key.into(), value);
self
}
pub fn parallel_tool_calls(mut self, parallel_tool_calls: Option<bool>) -> Self {
self.config.parallel_tool_calls = parallel_tool_calls;
self
}
pub fn volatile_suffix_len(mut self, len: usize) -> Self {
self.config.volatile_suffix_len = len;
self
}
pub fn extra_headers(mut self, headers: Vec<(String, String)>) -> Self {
self.config.extra_headers = headers;
self
}
pub fn extra_header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
self.config.extra_headers.push((name.into(), value.into()));
self
}
pub fn cache_diagnostics(mut self, config: CacheDiagnosticsConfig) -> Self {
self.config.cache_diagnostics = Some(config);
self
}
pub fn build(self) -> LlmCallConfig {
self.config
}
}
pub use crate::provider::DriverId;
#[derive(Clone, Default, PartialEq, Eq)]
pub struct ProviderMetadata {
pub refresh_token: Option<String>,
pub account_id: Option<String>,
pub extra: Option<serde_json::Value>,
}
impl std::fmt::Debug for ProviderMetadata {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ProviderMetadata")
.field(
"refresh_token",
&self.refresh_token.as_ref().map(|_| "<configured>"),
)
.field("account_id", &self.account_id)
.field("extra", &self.extra.as_ref().map(|_| "<configured>"))
.finish()
}
}
#[derive(Clone)]
pub struct ProviderConfig {
pub provider: crate::runtime_provider::ProviderKey,
pub provider_type: DriverId,
pub api_key: Option<String>,
pub base_url: Option<String>,
pub metadata: ProviderMetadata,
pub request_options: crate::provider::ProviderRequestOptions,
}
impl ProviderConfig {
pub fn new(provider_type: DriverId) -> Self {
let provider = crate::runtime_provider::ProviderKey::new(provider_type.as_str());
Self {
provider,
provider_type,
api_key: None,
base_url: None,
metadata: ProviderMetadata::default(),
request_options: Default::default(),
}
}
pub fn for_provider(
provider: impl Into<crate::runtime_provider::ProviderKey>,
provider_type: DriverId,
) -> Self {
Self {
provider: provider.into(),
provider_type,
api_key: None,
base_url: None,
metadata: ProviderMetadata::default(),
request_options: Default::default(),
}
}
pub fn with_api_key(mut self, api_key: impl Into<String>) -> Self {
self.api_key = Some(api_key.into());
self
}
pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
self.base_url = Some(base_url.into());
self
}
pub fn with_metadata(mut self, metadata: ProviderMetadata) -> Self {
self.metadata = metadata;
self
}
pub fn with_request_options(
mut self,
request_options: crate::provider::ProviderRequestOptions,
) -> Self {
self.request_options = request_options;
self
}
}
impl std::fmt::Debug for ProviderConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ProviderConfig")
.field("provider", &self.provider)
.field("provider_type", &self.provider_type)
.field("auth", &self.api_key.as_ref().map(|_| "<configured>"))
.field("base_url", &self.base_url.as_ref().map(|_| "<configured>"))
.field(
"metadata",
&self.metadata.extra.as_ref().map(|_| "<configured>"),
)
.finish()
}
}
#[derive(Clone)]
pub struct DriverConfig {
pub provider: crate::runtime_provider::ProviderKey,
pub provider_type: DriverId,
pub api_key: Option<String>,
pub credentials: std::collections::BTreeMap<String, String>,
pub base_url: Option<String>,
pub metadata: ProviderMetadata,
}
impl DriverConfig {
pub fn from_provider_config(config: &ProviderConfig) -> Self {
Self {
provider: config.provider.clone(),
provider_type: config.provider_type.clone(),
credentials: crate::credential_schema::parse_credential_document(
config.api_key.as_deref(),
),
api_key: config.api_key.clone(),
base_url: config.base_url.clone(),
metadata: config.metadata.clone(),
}
}
pub fn credential(&self, name: &str) -> Option<&str> {
self.credentials
.get(name)
.map(String::as_str)
.filter(|s| !s.is_empty())
}
}
impl std::fmt::Debug for DriverConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DriverConfig")
.field("provider", &self.provider)
.field("provider_type", &self.provider_type)
.field("auth", &self.api_key.as_ref().map(|_| "<configured>"))
.field(
"credential_fields",
&self.credentials.keys().collect::<Vec<_>>(),
)
.field("base_url", &self.base_url.as_ref().map(|_| "<configured>"))
.finish()
}
}
pub type BoxedChatDriver = Box<dyn ChatDriver>;
#[derive(Debug, Clone)]
pub struct EmbedRequest {
pub texts: Vec<String>,
pub model: String,
}
#[derive(Debug, Clone)]
pub struct EmbedResponse {
pub embeddings: Vec<Vec<f32>>,
pub usage_tokens: Option<u32>,
pub actual_cost_usd: Option<f64>,
}
#[derive(Debug, thiserror::Error)]
pub enum EmbeddingsDriverError {
#[error("embeddings provider returned an error: {0}")]
Provider(String),
#[error("embeddings request failed: {0}")]
Transport(String),
}
#[async_trait]
pub trait EmbeddingsDriver: Send + Sync {
async fn embed(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
request: EmbedRequest,
) -> std::result::Result<EmbedResponse, EmbeddingsDriverError>;
}
#[async_trait]
impl EmbeddingsDriver for Box<dyn EmbeddingsDriver> {
async fn embed(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
request: EmbedRequest,
) -> std::result::Result<EmbedResponse, EmbeddingsDriverError> {
(**self).embed(endpoint, request).await
}
}
pub type BoxedEmbeddingsDriver = Box<dyn EmbeddingsDriver>;
pub type EmbeddingsDriverFactory =
Arc<dyn Fn(&DriverConfig) -> BoxedEmbeddingsDriver + Send + Sync>;
pub type DriverFactory = Arc<dyn Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync>;
struct CredentialGateDriver {
inner: BoxedChatDriver,
message: String,
}
impl CredentialGateDriver {
fn error(&self) -> AgentLoopError {
AgentLoopError::llm_kind(LlmErrorKind::Authentication, self.message.clone())
}
}
#[async_trait]
impl ChatDriver for CredentialGateDriver {
async fn chat_completion_stream(
&self,
_endpoint: &crate::runtime_provider::ProviderEndpoint,
_messages: Vec<LlmMessage>,
_config: &LlmCallConfig,
) -> Result<LlmResponseStream> {
Err(self.error())
}
async fn list_models(
&self,
_endpoint: &crate::runtime_provider::ProviderEndpoint,
) -> Result<Option<Vec<DiscoveredModel>>> {
Err(self.error())
}
fn supports_compact(&self) -> bool {
self.inner.supports_compact()
}
fn supports_stateful_responses(&self) -> bool {
self.inner.supports_stateful_responses()
}
fn effective_context_window(&self, model: &str) -> Option<usize> {
self.inner.effective_context_window(model)
}
fn supports_parallel_tool_calls(&self, model: &str) -> bool {
self.inner.supports_parallel_tool_calls(model)
}
async fn compact(
&self,
_endpoint: &crate::runtime_provider::ProviderEndpoint,
_request: CompactRequest,
) -> Result<Option<CompactResponse>> {
Err(self.error())
}
}
struct RequestOptionsDriver {
inner: Arc<dyn ChatDriver>,
options: crate::provider::ProviderRequestOptions,
}
impl RequestOptionsDriver {
fn wrap(
driver: BoxedChatDriver,
options: &crate::provider::ProviderRequestOptions,
) -> BoxedChatDriver {
if options.is_empty() {
return driver;
}
Box::new(Self {
inner: Arc::from(driver),
options: options.clone(),
})
}
fn apply(&self, config: &LlmCallConfig) -> LlmCallConfig {
let mut config = config.clone();
config.extra_headers.extend(self.options.header_pairs());
if self.options.cache_diagnostics {
config.cache_diagnostics = Some(CacheDiagnosticsConfig {
enabled: true,
previous_message_id: config.previous_response_id.clone(),
});
}
config
}
}
#[async_trait]
impl ChatDriver for RequestOptionsDriver {
fn native_async_driver(
&self,
model: &str,
tools: std::collections::BTreeMap<String, Option<serde_json::Value>>,
continuation: Option<crate::native_async::Delivery>,
) -> Option<Arc<dyn ChatDriver>> {
Some(Arc::new(Self {
inner: self.inner.native_async_driver(model, tools, continuation)?,
options: self.options.clone(),
}))
}
async fn chat_completion_stream(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponseStream> {
self.inner
.chat_completion_stream(endpoint, messages, &self.apply(config))
.await
}
async fn chat_completion(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponse> {
self.inner
.chat_completion(endpoint, messages, &self.apply(config))
.await
}
fn supports_native_non_streaming(&self) -> bool {
self.inner.supports_native_non_streaming()
}
async fn chat_completion_non_streaming(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponse> {
self.inner
.chat_completion_non_streaming(endpoint, messages, &self.apply(config))
.await
}
async fn list_models(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
) -> Result<Option<Vec<DiscoveredModel>>> {
self.inner.list_models(endpoint).await
}
fn supports_compact(&self) -> bool {
self.inner.supports_compact()
}
fn supports_stateful_responses(&self) -> bool {
self.inner.supports_stateful_responses()
}
fn effective_context_window(&self, model: &str) -> Option<usize> {
self.inner.effective_context_window(model)
}
fn supports_parallel_tool_calls(&self, model: &str) -> bool {
self.inner.supports_parallel_tool_calls(model)
}
async fn compact(
&self,
endpoint: &crate::runtime_provider::ProviderEndpoint,
request: CompactRequest,
) -> Result<Option<CompactResponse>> {
self.inner.compact(endpoint, request).await
}
}
pub use everruns_model_profiles::ServiceKind;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DriverOAuthFlow {
OpenRouterPkce,
}
#[derive(Debug, Clone)]
pub struct DriverOAuthConfig {
pub authorize_url: String,
pub token_url: String,
pub flow: DriverOAuthFlow,
}
impl DriverOAuthConfig {
pub fn openrouter() -> Self {
Self {
authorize_url: "https://openrouter.ai/auth".to_string(),
token_url: "https://openrouter.ai/api/v1/auth/keys".to_string(),
flow: DriverOAuthFlow::OpenRouterPkce,
}
}
}
#[derive(Clone)]
pub struct DriverDescriptor {
pub id: DriverId,
pub display_name: String,
pub services: Vec<ServiceKind>,
pub credential_schema: CredentialFormSchema,
pub oauth: Option<DriverOAuthConfig>,
pub chat: Option<DriverFactory>,
pub embeddings: Option<EmbeddingsDriverFactory>,
}
impl DriverDescriptor {
pub fn chat_only<F>(id: impl Into<DriverId>, factory: F) -> Self
where
F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
{
let id = id.into();
Self {
display_name: default_display_name(&id),
credential_schema: default_credential_schema(&id),
services: vec![ServiceKind::Chat],
oauth: None,
chat: Some(Arc::new(factory)),
embeddings: None,
id,
}
}
pub fn supports(&self, service: ServiceKind) -> bool {
self.services.contains(&service)
}
}
impl std::fmt::Debug for DriverDescriptor {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DriverDescriptor")
.field("id", &self.id)
.field("display_name", &self.display_name)
.field("services", &self.services)
.field("oauth", &self.oauth.is_some())
.field("chat", &self.chat.is_some())
.field("embeddings", &self.embeddings.is_some())
.finish()
}
}
fn default_display_name(id: &DriverId) -> String {
id.as_str().replace(['_', '-'], " ")
}
fn default_credential_schema(id: &DriverId) -> CredentialFormSchema {
if id == &DriverId::LlmSim {
CredentialFormSchema::empty()
} else {
CredentialFormSchema::api_key(String::new())
}
}
#[derive(Clone, Default)]
pub struct DriverRegistry {
descriptors: HashMap<DriverId, DriverDescriptor>,
providers: crate::runtime_provider::RuntimeProviderRegistry,
}
impl DriverRegistry {
pub fn new() -> Self {
Self {
descriptors: HashMap::new(),
providers: crate::runtime_provider::RuntimeProviderRegistry::new(),
}
}
pub fn register_provider(
&mut self,
provider: crate::runtime_provider::RuntimeProvider,
) -> Result<()> {
self.providers.register(provider)
}
pub fn replace_provider(
&mut self,
provider: crate::runtime_provider::RuntimeProvider,
) -> Option<Arc<crate::runtime_provider::RuntimeProvider>> {
self.providers.replace(provider)
}
pub fn provider(
&self,
id: &crate::runtime_provider::ProviderKey,
) -> Option<Arc<crate::runtime_provider::RuntimeProvider>> {
self.providers.get(id)
}
pub fn register_descriptor(&mut self, descriptor: DriverDescriptor) {
if self.descriptors.contains_key(&descriptor.id) {
panic!(
"driver already registered for provider '{}'; \
use register_descriptor_or_replace to overwrite intentionally",
descriptor.id
);
}
self.descriptors.insert(descriptor.id.clone(), descriptor);
}
pub fn register_descriptor_or_replace(&mut self, descriptor: DriverDescriptor) {
self.descriptors.insert(descriptor.id.clone(), descriptor);
}
pub fn register<F>(&mut self, provider_type: impl Into<DriverId>, factory: F)
where
F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
{
self.register_descriptor(DriverDescriptor::chat_only(provider_type, factory));
}
pub fn register_or_replace<F>(&mut self, provider_type: impl Into<DriverId>, factory: F)
where
F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
{
self.register_descriptor_or_replace(DriverDescriptor::chat_only(provider_type, factory));
}
pub fn register_external<F>(&mut self, id: impl AsRef<str>, factory: F)
where
F: Fn(&DriverConfig) -> BoxedChatDriver + Send + Sync + 'static,
{
let mut descriptor = DriverDescriptor::chat_only(DriverId::external(id), factory);
descriptor.credential_schema = CredentialFormSchema::empty();
self.register_descriptor(descriptor);
}
pub fn create_chat_driver(&self, config: &ProviderConfig) -> Result<BoxedChatDriver> {
if let Some(provider) = self.providers.get(&config.provider) {
return Ok(RequestOptionsDriver::wrap(
(*provider).clone().into_boxed_driver(),
&config.request_options,
));
}
let descriptor = self.descriptors.get(&config.provider_type).ok_or_else(|| {
AgentLoopError::driver_not_registered(config.provider_type.to_string())
})?;
let factory = descriptor.chat.as_ref().ok_or_else(|| {
AgentLoopError::llm(format!(
"Provider driver '{}' does not implement the chat service.",
config.provider_type
))
})?;
let driver_config = DriverConfig::from_provider_config(config);
let driver = factory(&driver_config);
let mut credential_fields = driver_config.credentials.clone();
if let Some(serde_json::Value::Object(extra)) = &driver_config.metadata.extra {
for (name, value) in extra {
if let Some(value) = value.as_str() {
credential_fields
.entry(name.clone())
.or_insert_with(|| value.to_string());
}
}
}
let credential_errors = descriptor.credential_schema.validate(&credential_fields);
if credential_errors.is_empty() {
Ok(RequestOptionsDriver::wrap(driver, &config.request_options))
} else {
let message = if descriptor.credential_schema.fields.len() == 1
&& descriptor.credential_schema.fields[0].name == "api_key"
{
"API key is required. Configure the API key in provider settings.".to_string()
} else {
format!(
"Provider credentials are required. Configure provider settings: {}",
credential_errors.join(" ")
)
};
Ok(Box::new(CredentialGateDriver {
inner: driver,
message,
}))
}
}
pub fn has_driver(&self, provider_type: &DriverId) -> bool {
self.descriptors.contains_key(provider_type)
}
pub fn descriptor(&self, provider_type: &DriverId) -> Option<&DriverDescriptor> {
self.descriptors.get(provider_type)
}
pub fn supports(&self, provider_type: &DriverId, service: ServiceKind) -> bool {
self.descriptors
.get(provider_type)
.is_some_and(|d| d.supports(service))
}
pub fn providers_for(&self, service: ServiceKind) -> Vec<DriverId> {
self.descriptors
.values()
.filter(|d| d.supports(service))
.map(|d| d.id.clone())
.collect()
}
pub fn registered_providers(&self) -> Vec<DriverId> {
self.descriptors.keys().cloned().collect()
}
pub fn registered_provider_ids(&self) -> Vec<String> {
self.providers.ids()
}
pub fn create_embeddings_driver(
&self,
config: &ProviderConfig,
) -> std::result::Result<BoxedEmbeddingsDriver, EmbeddingsDriverError> {
let requires_api_key = config.provider_type != DriverId::LlmSim;
if requires_api_key && config.api_key.is_none() {
return Err(EmbeddingsDriverError::Provider(
"API key is required. Configure the API key in provider settings.".to_string(),
));
}
let descriptor = self.descriptors.get(&config.provider_type).ok_or_else(|| {
EmbeddingsDriverError::Provider(format!(
"No driver registered for provider '{}'",
config.provider_type
))
})?;
let factory = descriptor.embeddings.as_ref().ok_or_else(|| {
EmbeddingsDriverError::Provider(format!(
"Provider driver '{}' does not implement the embeddings service.",
config.provider_type
))
})?;
let driver_config = DriverConfig::from_provider_config(config);
Ok(factory(&driver_config))
}
}
const MAX_TOOL_RESULT_BYTES: usize = 64 * 1024;
const TRUNCATION_SUFFIX: &str =
"\n\n[Output truncated — exceeded 64 KiB limit. Try quiet flags, pipes, or redirect to file.]";
pub fn truncate_tool_result(text: String) -> String {
if text.len() <= MAX_TOOL_RESULT_BYTES {
return text;
}
let content_budget = MAX_TOOL_RESULT_BYTES.saturating_sub(TRUNCATION_SUFFIX.len());
let mut end = content_budget;
while end > 0 && !text.is_char_boundary(end) {
end -= 1;
}
let mut truncated = text[..end].to_string();
truncated.push_str(TRUNCATION_SUFFIX);
truncated
}
#[cfg(test)]
mod tests {
use super::*;
use crate::runtime_provider::ProviderEndpoint;
#[test]
fn test_disjoint_prompt_tokens_subtracts_cached_subset() {
assert_eq!(disjoint_prompt_tokens(1000, Some(800)), 200);
assert_eq!(disjoint_prompt_tokens(1000, None), 1000);
assert_eq!(disjoint_prompt_tokens(1000, Some(0)), 1000);
assert_eq!(disjoint_prompt_tokens(800, Some(1000)), 0);
}
fn bare_call_config() -> LlmCallConfig {
LlmCallConfig {
model: "claude-opus-4-8".to_string(),
temperature: None,
max_tokens: None,
tools: vec![],
reasoning_effort: None,
speed: None,
verbosity: None,
metadata: HashMap::new(),
previous_response_id: None,
provider_opaque_context: None,
tool_search: None,
prompt_cache: None,
driver_options: Default::default(),
parallel_tool_calls: None,
volatile_suffix_len: 0,
extra_headers: Vec::new(),
cache_diagnostics: None,
reasoning_state: None,
}
}
#[test]
fn provider_config_debug_redacts_runtime_values() {
let config = ProviderConfig::new(DriverId::OpenAI)
.with_api_key("secret-key")
.with_base_url("https://user:password@example.test/v1?token=secret")
.with_metadata(ProviderMetadata {
refresh_token: Some("refresh-secret".into()),
account_id: Some("account-1".into()),
extra: Some(serde_json::json!({ "client_secret": "metadata-secret" })),
});
let debug = format!("{config:?}");
assert!(debug.contains("ProviderConfig"));
assert!(debug.contains("openai"));
assert!(debug.contains("<configured>"));
for secret in [
"secret-key",
"password",
"token=secret",
"refresh-secret",
"metadata-secret",
] {
assert!(!debug.contains(secret), "debug output exposed {secret}");
}
}
#[test]
fn system_messages_fold_only_system_text_in_transcript_order() {
use LlmMessageRole::{Assistant, System, Tool, User};
for (messages, expected) in [
(vec![], None),
(
vec![
LlmMessage::text(User, "user"),
LlmMessage::text(Assistant, "answer"),
LlmMessage::text(Tool, "result"),
],
None,
),
(vec![LlmMessage::text(System, "")], Some("")),
(
vec![
LlmMessage::text(System, "rules"),
LlmMessage::text(User, "question"),
],
Some("rules"),
),
(
vec![
LlmMessage::text(System, "first"),
LlmMessage::text(User, "question"),
LlmMessage::text(System, "second"),
LlmMessage::text(Assistant, "answer"),
LlmMessage::text(System, "third"),
],
Some("first\n\nsecond\n\nthird"),
),
(
vec![
LlmMessage::parts(
System,
vec![
LlmContentPart::text("foo"),
LlmContentPart::image("image"),
LlmContentPart::audio("audio"),
LlmContentPart::text("bar"),
],
),
LlmMessage::text(System, "next"),
],
Some("foobar\n\nnext"),
),
] {
assert_eq!(fold_system_messages(&messages).as_deref(), expected);
}
}
#[test]
fn prefix_preserves_all_media_and_changes_only_the_first_text_part() {
let mut plain = LlmMessage::text(LlmMessageRole::User, "Hello");
plain.prepend_text_prefix("[Alice] ");
assert!(
matches!(plain.content, LlmMessageContent::Text(ref text) if text == "[Alice] Hello")
);
for (parts, expected) in [
(vec![], vec![("text", "[Alice] ")]),
(
vec![
LlmContentPart::image("image"),
LlmContentPart::audio("audio"),
],
vec![("text", "[Alice] "), ("image", "image"), ("audio", "audio")],
),
(
vec![
LlmContentPart::text("Hello"),
LlmContentPart::image("image"),
],
vec![("text", "[Alice] Hello"), ("image", "image")],
),
(
vec![
LlmContentPart::image("image"),
LlmContentPart::text("Hello"),
LlmContentPart::audio("audio"),
LlmContentPart::text("later"),
],
vec![
("image", "image"),
("text", "[Alice] Hello"),
("audio", "audio"),
("text", "later"),
],
),
(
vec![LlmContentPart::text(""), LlmContentPart::text("later")],
vec![("text", "[Alice] "), ("text", "later")],
),
] {
let mut message = LlmMessage::parts(LlmMessageRole::Tool, parts);
message.tool_call_id = Some("call-1".into());
message.prepend_text_prefix("[Alice] ");
let LlmMessageContent::Parts(parts) = &message.content else {
panic!("parts must remain parts")
};
let actual: Vec<_> = parts
.iter()
.map(|part| match part {
LlmContentPart::Text { text } => ("text", text.as_str()),
LlmContentPart::Image { url } => ("image", url.as_str()),
LlmContentPart::Audio { url } => ("audio", url.as_str()),
LlmContentPart::File { url, .. } => ("file", url.as_str()),
})
.collect();
assert_eq!(actual, expected);
assert_eq!(message.role, LlmMessageRole::Tool);
assert_eq!(message.tool_call_id.as_deref(), Some("call-1"));
}
}
struct FixtureDriver(&'static str);
#[async_trait]
impl ChatDriver for FixtureDriver {
async fn chat_completion_stream(
&self,
_: &ProviderEndpoint,
_: Vec<LlmMessage>,
_: &LlmCallConfig,
) -> Result<LlmResponseStream> {
Ok(Box::pin(futures::stream::iter([
Ok(LlmStreamEvent::TextDelta(self.0.into())),
Ok(LlmStreamEvent::Done(Box::default())),
])))
}
async fn list_models(&self, _: &ProviderEndpoint) -> Result<Option<Vec<DiscoveredModel>>> {
Ok(Some(vec![DiscoveredModel {
model_id: self.0.into(),
display_name: None,
created_at: None,
owned_by: None,
capabilities: vec!["chat".into()],
discovered_profile: None,
}]))
}
async fn compact(
&self,
_: &ProviderEndpoint,
request: CompactRequest,
) -> Result<Option<CompactResponse>> {
Ok(Some(CompactResponse {
output: vec![crate::compact::CompactOutputItem::Compaction {
encrypted_content: request.model,
}],
usage: None,
}))
}
fn supports_compact(&self) -> bool {
true
}
fn supports_stateful_responses(&self) -> bool {
true
}
fn effective_context_window(&self, model: &str) -> Option<usize> {
(model == "known").then_some(12345)
}
fn supports_parallel_tool_calls(&self, model: &str) -> bool {
model == "known"
}
}
fn compact_fixture() -> CompactRequest {
CompactRequest {
reasoning_state: None,
model: "compact-model".into(),
input: vec![],
previous_response_id: None,
instructions: None,
}
}
#[tokio::test]
async fn default_and_boxed_drivers_preserve_optional_operations_and_model_capabilities() {
struct DefaultDriver;
#[async_trait]
impl ChatDriver for DefaultDriver {
async fn chat_completion_stream(
&self,
_: &ProviderEndpoint,
_: Vec<LlmMessage>,
_: &LlmCallConfig,
) -> Result<LlmResponseStream> {
Ok(Box::pin(futures::stream::empty()))
}
}
let endpoint = ProviderEndpoint::default();
assert!(!DefaultDriver.supports_compact());
assert!(!DefaultDriver.supports_stateful_responses());
assert!(!DefaultDriver.supports_parallel_tool_calls("known"));
assert_eq!(DefaultDriver.effective_context_window("known"), None);
assert!(
DefaultDriver
.list_models(&endpoint)
.await
.unwrap()
.is_none()
);
assert!(
DefaultDriver
.compact(&endpoint, compact_fixture())
.await
.unwrap()
.is_none()
);
let boxed: BoxedChatDriver = Box::new(FixtureDriver("boxed"));
assert!(boxed.supports_compact());
assert!(boxed.supports_stateful_responses());
for (model, expected) in [("known", true), ("unknown", false)] {
assert_eq!(boxed.supports_parallel_tool_calls(model), expected);
assert_eq!(
boxed.effective_context_window(model),
expected.then_some(12345)
);
}
assert_eq!(
boxed
.chat_completion(&endpoint, vec![], &bare_call_config())
.await
.unwrap()
.text,
"boxed"
);
}
#[tokio::test]
async fn registry_replacement_changes_factory_and_preserves_other_descriptors() {
let mut registry = DriverRegistry::new();
assert!(registry.registered_providers().is_empty());
registry.register(DriverId::LlmSim, |_| Box::new(FixtureDriver("first")));
registry.register_descriptor(DriverDescriptor {
display_name: "OpenAI custom".into(),
services: vec![ServiceKind::Chat, ServiceKind::Realtime],
..DriverDescriptor::chat_only(DriverId::OpenAI, |_| Box::new(FixtureDriver("other")))
});
let config = ProviderConfig::new(DriverId::LlmSim);
let endpoint = ProviderEndpoint::default();
assert_eq!(
registry
.create_chat_driver(&config)
.unwrap()
.chat_completion(&endpoint, vec![], &bare_call_config())
.await
.unwrap()
.text,
"first"
);
registry.register_or_replace(DriverId::LlmSim, |_| Box::new(FixtureDriver("replacement")));
assert_eq!(
registry
.create_chat_driver(&config)
.unwrap()
.chat_completion(&endpoint, vec![], &bare_call_config())
.await
.unwrap()
.text,
"replacement"
);
assert!(registry.has_driver(&DriverId::LlmSim));
assert!(!registry.has_driver(&DriverId::Anthropic));
assert_eq!(
registry.providers_for(ServiceKind::Realtime),
vec![DriverId::OpenAI]
);
let mut chat = registry.providers_for(ServiceKind::Chat);
chat.sort_by_key(|id| id.to_string());
assert_eq!(chat, vec![DriverId::LlmSim, DriverId::OpenAI]);
assert!(registry.supports(&DriverId::OpenAI, ServiceKind::Realtime));
assert!(!registry.supports(&DriverId::LlmSim, ServiceKind::Realtime));
assert!(!registry.supports(&DriverId::Gemini, ServiceKind::Chat));
assert_eq!(
registry.descriptor(&DriverId::OpenAI).unwrap().display_name,
"OpenAI custom"
);
assert_eq!(
registry
.create_chat_driver(
&ProviderConfig::new(DriverId::OpenAI).with_api_key("synthetic-key")
)
.unwrap()
.chat_completion(&endpoint, vec![], &bare_call_config())
.await
.unwrap()
.text,
"other"
);
let defaults = DriverDescriptor::chat_only(DriverId::Anthropic, |_| {
Box::new(FixtureDriver("default"))
});
assert_eq!(defaults.display_name, "anthropic");
let sim = registry.descriptor(&DriverId::LlmSim).unwrap();
assert!(sim.credential_schema.fields.is_empty());
assert_eq!(sim.services, vec![ServiceKind::Chat]);
assert!(sim.chat.is_some());
let real = registry.descriptor(&DriverId::OpenAI).unwrap();
assert_eq!(real.credential_schema.fields.len(), 1);
assert_eq!(real.credential_schema.fields[0].name, "api_key");
assert!(real.credential_schema.fields[0].required);
assert!(registry.descriptor(&DriverId::Gemini).is_none());
}
#[test]
#[should_panic(expected = "already registered")]
fn duplicate_registration_rejects_an_existing_driver() {
let mut registry = DriverRegistry::new();
registry.register(DriverId::OpenAI, |_| Box::new(FixtureDriver("first")));
registry.register(DriverId::OpenAI, |_| Box::new(FixtureDriver("second")));
}
#[tokio::test]
async fn factory_receives_complete_config_and_external_metadata_auth_remains_keyless() {
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let capture = seen.clone();
let mut registry = DriverRegistry::new();
registry.register_external("CUSTOM", move |config| {
capture.lock().unwrap().push(config.clone());
Box::new(FixtureDriver("external"))
});
let metadata = ProviderMetadata {
refresh_token: Some("refresh".into()),
account_id: Some("account".into()),
extra: Some(serde_json::json!({"region":"west"})),
};
for key in [None, Some("synthetic-key")] {
let mut config =
ProviderConfig::for_provider("connection", DriverId::external("custom"))
.with_base_url("https://gateway.example/v1")
.with_metadata(metadata.clone());
if let Some(key) = key {
config = config.with_api_key(key);
}
let response = registry
.create_chat_driver(&config)
.unwrap()
.chat_completion(&ProviderEndpoint::default(), vec![], &bare_call_config())
.await
.unwrap();
assert_eq!(response.text, "external");
let received = seen.lock().unwrap().pop().unwrap();
assert_eq!(received.provider.as_str(), "connection");
assert_eq!(received.provider_type, DriverId::external("custom"));
assert_eq!(received.api_key.as_deref(), key);
assert_eq!(received.credential("api_key"), key);
assert_eq!(received.credentials.len(), usize::from(key.is_some()));
assert_eq!(
received.base_url.as_deref(),
Some("https://gateway.example/v1")
);
assert_eq!(received.metadata, metadata);
}
assert!(
registry
.descriptor(&DriverId::external("custom"))
.unwrap()
.credential_schema
.fields
.is_empty()
);
}
#[test]
fn registry_distinguishes_missing_driver_from_missing_chat_service() {
let mut registry = DriverRegistry::new();
assert!(
matches!(registry.create_chat_driver(&ProviderConfig::new(DriverId::Anthropic)), Err(AgentLoopError::DriverNotRegistered(id)) if id == "anthropic")
);
registry.register_descriptor(DriverDescriptor {
id: DriverId::external("embeddings-only"),
display_name: "Embeddings Only".into(),
services: vec![ServiceKind::Embeddings],
credential_schema: CredentialFormSchema::empty(),
oauth: None,
chat: None,
embeddings: None,
});
match registry
.create_chat_driver(&ProviderConfig::new(DriverId::external("embeddings-only")))
{
Err(AgentLoopError::Llm(error)) => assert_eq!(
error.message,
"Provider driver 'embeddings-only' does not implement the chat service."
),
_ => panic!("expected a missing-chat-service error"),
}
}
#[tokio::test]
async fn credential_gate_rejects_every_io_operation_before_dispatch() {
struct ForbiddenDriver;
#[async_trait]
impl ChatDriver for ForbiddenDriver {
async fn chat_completion_stream(
&self,
_: &ProviderEndpoint,
_: Vec<LlmMessage>,
_: &LlmCallConfig,
) -> Result<LlmResponseStream> {
panic!("unauthenticated stream dispatch")
}
async fn list_models(
&self,
_: &ProviderEndpoint,
) -> Result<Option<Vec<DiscoveredModel>>> {
panic!("unauthenticated model dispatch")
}
async fn compact(
&self,
_: &ProviderEndpoint,
_: CompactRequest,
) -> Result<Option<CompactResponse>> {
panic!("unauthenticated compact dispatch")
}
}
let mut registry = DriverRegistry::new();
registry.register(DriverId::OpenAI, |config| {
if config.api_key.is_some() {
Box::new(FixtureDriver("authenticated"))
} else {
Box::new(ForbiddenDriver)
}
});
let driver = registry
.create_chat_driver(&ProviderConfig::new(DriverId::OpenAI))
.unwrap();
let endpoint = ProviderEndpoint::default();
let stream_error = match driver
.chat_completion_stream(&endpoint, vec![], &bare_call_config())
.await
{
Err(error) => error,
Ok(_) => panic!("expected authentication error"),
};
for error in [
stream_error,
driver
.chat_completion(&endpoint, vec![], &bare_call_config())
.await
.unwrap_err(),
driver.list_models(&endpoint).await.unwrap_err(),
driver
.compact(&endpoint, compact_fixture())
.await
.unwrap_err(),
] {
assert_eq!(error.llm_error_kind(), Some(LlmErrorKind::Authentication));
assert_eq!(
error.to_string(),
"LLM error: API key is required. Configure the API key in provider settings."
);
}
let driver = registry
.create_chat_driver(
&ProviderConfig::new(DriverId::OpenAI).with_api_key("synthetic-key"),
)
.unwrap();
assert_eq!(
driver
.chat_completion(&endpoint, vec![], &bare_call_config())
.await
.unwrap()
.text,
"authenticated"
);
assert_eq!(
driver.list_models(&endpoint).await.unwrap().unwrap()[0].model_id,
"authenticated"
);
assert_eq!(
serde_json::to_value(
driver
.compact(&endpoint, compact_fixture())
.await
.unwrap()
.unwrap()
.output
)
.unwrap(),
serde_json::json!([{"type":"compaction","encrypted_content":"compact-model"}])
);
}
#[tokio::test]
async fn request_options_preserve_calls_and_apply_headers_and_diagnostics_independently() {
struct CapturingDriver(Arc<std::sync::Mutex<Vec<LlmCallConfig>>>);
impl CapturingDriver {
fn capture(
&self,
endpoint: &ProviderEndpoint,
messages: &[LlmMessage],
config: &LlmCallConfig,
) {
assert_eq!(
endpoint.url("probe").as_deref(),
Some("https://gateway.example/v1/probe")
);
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].role, LlmMessageRole::User);
assert_eq!(messages[0].content_as_text(), "request text");
self.0.lock().unwrap().push(config.clone());
}
}
#[async_trait]
impl ChatDriver for CapturingDriver {
async fn chat_completion_stream(
&self,
endpoint: &ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponseStream> {
self.capture(endpoint, &messages, config);
FixtureDriver("stream")
.chat_completion_stream(endpoint, messages, config)
.await
}
async fn chat_completion(
&self,
endpoint: &ProviderEndpoint,
messages: Vec<LlmMessage>,
config: &LlmCallConfig,
) -> Result<LlmResponse> {
self.capture(endpoint, &messages, config);
FixtureDriver("completion")
.chat_completion(endpoint, messages, config)
.await
}
}
let provider = crate::Provider::new("fixture", FixtureDriver("endpoint"))
.base_url("https://gateway.example/v1");
for (headers, diagnostics) in [(false, false), (true, false), (false, true), (true, true)] {
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let options = crate::provider::ProviderRequestOptions {
headers: if headers {
vec![crate::provider::ProviderRequestHeader {
name: "x-base".into(),
value: "connection".into(),
}]
} else {
vec![]
},
cache_diagnostics: diagnostics,
};
let driver =
RequestOptionsDriver::wrap(Box::new(CapturingDriver(seen.clone())), &options);
let mut config = bare_call_config();
config.model = "requested-model".into();
config.temperature = Some(0.25);
config.max_tokens = Some(42);
config
.metadata
.insert("session_id".into(), "session-one".into());
config.previous_response_id = Some("response-one".into());
config.extra_headers = vec![("x-base".into(), "original".into())];
config.cache_diagnostics = Some(CacheDiagnosticsConfig {
enabled: false,
previous_message_id: Some("existing".into()),
});
let mut stream = driver
.chat_completion_stream(
provider.endpoint(),
vec![LlmMessage::text(LlmMessageRole::User, "request text")],
&config,
)
.await
.unwrap();
use futures::StreamExt;
assert!(
matches!(stream.next().await.unwrap().unwrap(), LlmStreamEvent::TextDelta(text) if text == "stream")
);
assert!(matches!(
stream.next().await.unwrap().unwrap(),
LlmStreamEvent::Done(_)
));
assert!(stream.next().await.is_none());
assert_eq!(
driver
.chat_completion(
provider.endpoint(),
vec![LlmMessage::text(LlmMessageRole::User, "request text")],
&config
)
.await
.unwrap()
.text,
"completion"
);
let mut expected_headers = vec![("x-base".into(), "original".into())];
if headers {
expected_headers.push(("x-base".into(), "connection".into()));
}
let observed = seen.lock().unwrap();
assert_eq!(observed.len(), 2);
for received in observed.iter() {
assert_eq!(received.extra_headers, expected_headers);
let diagnostic = received.cache_diagnostics.as_ref().unwrap();
assert_eq!(diagnostic.enabled, diagnostics);
assert_eq!(
diagnostic.previous_message_id.as_deref(),
Some(if diagnostics {
"response-one"
} else {
"existing"
})
);
assert_eq!(received.model, "requested-model");
assert_eq!(received.temperature, Some(0.25));
assert_eq!(received.max_tokens, Some(42));
assert_eq!(received.metadata, config.metadata);
assert_eq!(received.previous_response_id, config.previous_response_id);
}
assert_eq!(
config.extra_headers,
vec![("x-base".into(), "original".into())]
);
assert!(!config.cache_diagnostics.as_ref().unwrap().enabled);
assert_eq!(
config
.cache_diagnostics
.as_ref()
.unwrap()
.previous_message_id
.as_deref(),
Some("existing")
);
}
let options = crate::provider::ProviderRequestOptions {
headers: vec![],
cache_diagnostics: true,
};
let wrapped = RequestOptionsDriver::wrap(Box::new(FixtureDriver("forwarded")), &options);
assert!(wrapped.supports_compact());
assert!(wrapped.supports_stateful_responses());
for (model, expected) in [("known", true), ("unknown", false)] {
assert_eq!(wrapped.supports_parallel_tool_calls(model), expected);
assert_eq!(
wrapped.effective_context_window(model),
expected.then_some(12345)
);
}
assert_eq!(
wrapped
.list_models(provider.endpoint())
.await
.unwrap()
.unwrap()[0]
.model_id,
"forwarded"
);
assert_eq!(
serde_json::to_value(
wrapped
.compact(provider.endpoint(), compact_fixture())
.await
.unwrap()
.unwrap()
.output
)
.unwrap(),
serde_json::json!([{"type":"compaction","encrypted_content":"compact-model"}])
);
}
}