pub mod builder;
pub mod config;
#[allow(missing_docs)]
pub mod config_file;
pub mod llm_config;
#[cfg(all(feature = "native-http", feature = "tower"))]
pub mod managed;
use std::future::Future;
use std::pin::Pin;
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
use std::sync::Arc;
use futures_core::Stream;
use crate::error::{LiterLlmError, Result};
use crate::types::audio::{CreateSpeechRequest, CreateTranscriptionRequest, TranscriptionResponse};
use crate::types::batch::{BatchListQuery, BatchListResponse, BatchObject, CreateBatchRequest};
use crate::types::files::{CreateFileRequest, DeleteResponse, FileListQuery, FileListResponse, FileObject};
use crate::types::image::{CreateImageRequest, ImagesResponse};
use crate::types::moderation::{ModerationRequest, ModerationResponse};
use crate::types::ocr::{OcrRequest, OcrResponse};
use crate::types::raw::{RawExchange, RawStreamExchange};
use crate::types::rerank::{RerankRequest, RerankResponse};
use crate::types::responses::{CreateResponseRequest, ResponseObject, ResponseStreamEvent};
use crate::types::search::{SearchRequest, SearchResponse};
use crate::types::{
ChatCompletionChunk, ChatCompletionRequest, ChatCompletionResponse, EmbeddingRequest, EmbeddingResponse,
ModelsListResponse,
};
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
use crate::auth::Credential;
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
use crate::http;
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
use crate::provider::{self, OpenAiCompatibleProvider, OpenAiProvider, Provider};
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
use secrecy::ExposeSecret;
pub use builder::{ClientBuilder, NoApiKey, NoProvider, WithApiKey, WithProvider};
pub use config::{ClientConfig, ClientConfigBuilder};
pub use config_file::FileConfig;
pub use llm_config::{
BedrockConfig, LlmBudgetConfig, LlmCacheConfig, LlmConfig, LlmInFlightLimitConfig, LlmProviderConfig,
LlmRateLimitConfig,
};
use crate::types::batch::BatchStatus;
use std::time::Duration;
#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct WaitForBatchConfig {
pub initial_interval_secs: f64,
pub max_interval_secs: f64,
pub backoff_multiplier: f32,
pub timeout_secs: Option<f64>,
}
impl Default for WaitForBatchConfig {
fn default() -> Self {
Self {
initial_interval_secs: 5.0,
max_interval_secs: 60.0,
backoff_multiplier: 1.5,
timeout_secs: None,
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum BatchWaitError {
#[error("batch reached terminal failure state: {status:?}")]
Failed {
status: BatchStatus,
},
#[error("polling timed out after {timeout_secs:.1}s")]
Timeout {
timeout_secs: f64,
},
#[error("client error (code {code}): {message}")]
Client {
message: String,
code: u32,
},
}
impl From<LiterLlmError> for BatchWaitError {
fn from(err: LiterLlmError) -> Self {
Self::Client {
code: u32::from(err.status_code()),
message: err.to_string(),
}
}
}
#[cfg(not(target_arch = "wasm32"))]
#[cfg_attr(alef, alef(skip))]
pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
#[cfg(target_arch = "wasm32")]
#[cfg_attr(alef, alef(skip))]
pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + 'a>>;
#[cfg(not(target_arch = "wasm32"))]
#[cfg_attr(alef, alef(skip))]
pub type BoxStream<'a, T> = Pin<Box<dyn Stream<Item = T> + Send + 'a>>;
#[cfg(target_arch = "wasm32")]
#[cfg_attr(alef, alef(skip))]
pub type BoxStream<'a, T> = Pin<Box<dyn Stream<Item = T> + 'a>>;
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
struct PreparedRequest {
url: String,
provider: Arc<dyn Provider>,
body_json: serde_json::Value,
body_bytes: bytes::Bytes,
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
fn str_pair(pair: &(String, String)) -> (&str, &str) {
(pair.0.as_str(), pair.1.as_str())
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
fn merge_extra_body(body: &mut serde_json::Value) {
let Some(obj) = body.as_object_mut() else {
return;
};
let Some(extra_body) = obj.remove("extra_body") else {
return;
};
match extra_body {
serde_json::Value::Object(extra_fields) => {
obj.extend(extra_fields);
}
other => {
tracing::warn!(
extra_body_type = json_value_type_name(&other),
"ignoring non-object extra_body; it cannot be merged into the request body"
);
}
}
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
fn response_request_body(req: &CreateResponseRequest) -> Result<serde_json::Value> {
let mut body = serde_json::to_value(req)?;
merge_extra_body(&mut body);
Ok(body)
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
fn response_stream_request_body(mut req: CreateResponseRequest) -> Result<serde_json::Value> {
req.stream = Some(true);
let mut body = response_request_body(&req)?;
if let Some(obj) = body.as_object_mut() {
obj.insert("stream".into(), serde_json::Value::Bool(true));
}
Ok(body)
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
fn json_value_type_name(value: &serde_json::Value) -> &'static str {
match value {
serde_json::Value::Null => "null",
serde_json::Value::Bool(_) => "boolean",
serde_json::Value::Number(_) => "number",
serde_json::Value::String(_) => "string",
serde_json::Value::Array(_) => "array",
serde_json::Value::Object(_) => "object",
}
}
#[cfg(not(target_arch = "wasm32"))]
#[cfg_attr(alef, alef(skip))]
pub trait LlmClient: Send + Sync {
fn chat(&self, req: ChatCompletionRequest) -> BoxFuture<'_, Result<ChatCompletionResponse>>;
fn chat_stream(
&self,
req: ChatCompletionRequest,
) -> BoxFuture<'_, Result<BoxStream<'static, Result<ChatCompletionChunk>>>>;
fn embed(&self, req: EmbeddingRequest) -> BoxFuture<'_, Result<EmbeddingResponse>>;
fn list_models(&self) -> BoxFuture<'_, Result<ModelsListResponse>>;
fn image_generate(&self, req: CreateImageRequest) -> BoxFuture<'_, Result<ImagesResponse>>;
fn speech(&self, req: CreateSpeechRequest) -> BoxFuture<'_, Result<bytes::Bytes>>;
fn transcribe(&self, req: CreateTranscriptionRequest) -> BoxFuture<'_, Result<TranscriptionResponse>>;
fn moderate(&self, req: ModerationRequest) -> BoxFuture<'_, Result<ModerationResponse>>;
fn rerank(&self, req: RerankRequest) -> BoxFuture<'_, Result<RerankResponse>>;
fn search(&self, req: SearchRequest) -> BoxFuture<'_, Result<SearchResponse>>;
fn ocr(&self, req: OcrRequest) -> BoxFuture<'_, Result<OcrResponse>>;
}
#[cfg(target_arch = "wasm32")]
#[cfg_attr(alef, alef(skip))]
pub trait LlmClient {
fn chat(&self, req: ChatCompletionRequest) -> BoxFuture<'_, Result<ChatCompletionResponse>>;
fn chat_stream(
&self,
req: ChatCompletionRequest,
) -> BoxFuture<'_, Result<BoxStream<'static, Result<ChatCompletionChunk>>>>;
fn embed(&self, req: EmbeddingRequest) -> BoxFuture<'_, Result<EmbeddingResponse>>;
fn list_models(&self) -> BoxFuture<'_, Result<ModelsListResponse>>;
fn image_generate(&self, req: CreateImageRequest) -> BoxFuture<'_, Result<ImagesResponse>>;
fn speech(&self, req: CreateSpeechRequest) -> BoxFuture<'_, Result<bytes::Bytes>>;
fn transcribe(&self, req: CreateTranscriptionRequest) -> BoxFuture<'_, Result<TranscriptionResponse>>;
fn moderate(&self, req: ModerationRequest) -> BoxFuture<'_, Result<ModerationResponse>>;
fn rerank(&self, req: RerankRequest) -> BoxFuture<'_, Result<RerankResponse>>;
fn search(&self, req: SearchRequest) -> BoxFuture<'_, Result<SearchResponse>>;
fn ocr(&self, req: OcrRequest) -> BoxFuture<'_, Result<OcrResponse>>;
}
#[cfg_attr(alef, alef(skip))]
pub trait LlmClientRaw: LlmClient {
fn chat_raw(&self, req: ChatCompletionRequest) -> BoxFuture<'_, Result<RawExchange<ChatCompletionResponse>>>;
fn chat_stream_raw(
&self,
req: ChatCompletionRequest,
) -> BoxFuture<'_, Result<RawStreamExchange<BoxStream<'static, Result<ChatCompletionChunk>>>>>;
fn embed_raw(&self, req: EmbeddingRequest) -> BoxFuture<'_, Result<RawExchange<EmbeddingResponse>>>;
fn image_generate_raw(&self, req: CreateImageRequest) -> BoxFuture<'_, Result<RawExchange<ImagesResponse>>>;
fn transcribe_raw(
&self,
req: CreateTranscriptionRequest,
) -> BoxFuture<'_, Result<RawExchange<TranscriptionResponse>>>;
fn moderate_raw(&self, req: ModerationRequest) -> BoxFuture<'_, Result<RawExchange<ModerationResponse>>>;
fn rerank_raw(&self, req: RerankRequest) -> BoxFuture<'_, Result<RawExchange<RerankResponse>>>;
fn search_raw(&self, req: SearchRequest) -> BoxFuture<'_, Result<RawExchange<SearchResponse>>>;
fn ocr_raw(&self, req: OcrRequest) -> BoxFuture<'_, Result<RawExchange<OcrResponse>>>;
}
#[cfg(not(target_arch = "wasm32"))]
#[cfg_attr(alef, alef(skip))]
pub trait FileClient: Send + Sync {
fn create_file(&self, req: CreateFileRequest) -> BoxFuture<'_, Result<FileObject>>;
fn retrieve_file(&self, file_id: &str) -> BoxFuture<'_, Result<FileObject>>;
fn delete_file(&self, file_id: &str) -> BoxFuture<'_, Result<DeleteResponse>>;
fn list_files(&self, query: Option<FileListQuery>) -> BoxFuture<'_, Result<FileListResponse>>;
fn file_content(&self, file_id: &str) -> BoxFuture<'_, Result<bytes::Bytes>>;
}
#[cfg(target_arch = "wasm32")]
#[cfg_attr(alef, alef(skip))]
pub trait FileClient {
fn create_file(&self, req: CreateFileRequest) -> BoxFuture<'_, Result<FileObject>>;
fn retrieve_file(&self, file_id: &str) -> BoxFuture<'_, Result<FileObject>>;
fn delete_file(&self, file_id: &str) -> BoxFuture<'_, Result<DeleteResponse>>;
fn list_files(&self, query: Option<FileListQuery>) -> BoxFuture<'_, Result<FileListResponse>>;
fn file_content(&self, file_id: &str) -> BoxFuture<'_, Result<bytes::Bytes>>;
}
#[cfg(not(target_arch = "wasm32"))]
#[cfg_attr(alef, alef(skip))]
pub trait BatchClient: Send + Sync {
fn create_batch(&self, req: CreateBatchRequest) -> BoxFuture<'_, Result<BatchObject>>;
fn retrieve_batch(&self, batch_id: &str) -> BoxFuture<'_, Result<BatchObject>>;
fn list_batches(&self, query: Option<BatchListQuery>) -> BoxFuture<'_, Result<BatchListResponse>>;
fn cancel_batch(&self, batch_id: &str) -> BoxFuture<'_, Result<BatchObject>>;
}
#[cfg(target_arch = "wasm32")]
#[cfg_attr(alef, alef(skip))]
pub trait BatchClient {
fn create_batch(&self, req: CreateBatchRequest) -> BoxFuture<'_, Result<BatchObject>>;
fn retrieve_batch(&self, batch_id: &str) -> BoxFuture<'_, Result<BatchObject>>;
fn list_batches(&self, query: Option<BatchListQuery>) -> BoxFuture<'_, Result<BatchListResponse>>;
fn cancel_batch(&self, batch_id: &str) -> BoxFuture<'_, Result<BatchObject>>;
}
#[cfg(not(target_arch = "wasm32"))]
#[cfg_attr(alef, alef(skip))]
pub trait ResponseClient: Send + Sync {
fn create_response(&self, req: CreateResponseRequest) -> BoxFuture<'_, Result<ResponseObject>>;
fn create_response_stream(
&self,
req: CreateResponseRequest,
) -> BoxFuture<'_, Result<BoxStream<'static, Result<ResponseStreamEvent>>>>;
fn retrieve_response(&self, response_id: &str) -> BoxFuture<'_, Result<ResponseObject>>;
fn cancel_response(&self, response_id: &str) -> BoxFuture<'_, Result<ResponseObject>>;
}
#[cfg(target_arch = "wasm32")]
#[cfg_attr(alef, alef(skip))]
pub trait ResponseClient {
fn create_response(&self, req: CreateResponseRequest) -> BoxFuture<'_, Result<ResponseObject>>;
fn create_response_stream(
&self,
req: CreateResponseRequest,
) -> BoxFuture<'_, Result<BoxStream<'static, Result<ResponseStreamEvent>>>>;
fn retrieve_response(&self, response_id: &str) -> BoxFuture<'_, Result<ResponseObject>>;
fn cancel_response(&self, response_id: &str) -> BoxFuture<'_, Result<ResponseObject>>;
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
#[derive(Clone)]
pub struct DefaultClient {
config: ClientConfig,
http: reqwest::Client,
provider: Arc<dyn Provider>,
cached_auth_header: Option<(String, String)>,
cached_extra_headers: Vec<(String, String)>,
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
impl DefaultClient {
fn response_read_options(&self) -> http::request::ResponseReadOptions {
http::request::ResponseReadOptions {
max_retries: self.config.max_retries,
max_response_bytes: self.config.max_response_bytes,
}
}
pub fn new(config: ClientConfig, model_hint: Option<&str>) -> Result<Self> {
config::validate_response_limit(config.max_response_bytes)?;
let provider = build_provider(&config, model_hint);
provider.validate()?;
#[cfg(not(target_arch = "wasm32"))]
let mut config = config;
#[cfg(not(target_arch = "wasm32"))]
if config.load_env
&& config.api_key.expose_secret().is_empty()
&& let Some(env_var_name) = provider.env_var()
{
match std::env::var(env_var_name) {
Ok(val) if !val.is_empty() => {
config.api_key = secrecy::SecretString::from(val);
}
_ => {
return Err(LiterLlmError::Authentication {
message: format!("no API key provided and environment variable {env_var_name} is not set"),
status: 401,
});
}
}
}
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
if config.credential_provider.is_none()
&& config.api_key.expose_secret().is_empty()
&& provider.name() == "vertex_ai"
{
config.credential_provider = Some(Arc::new(crate::auth::vertex_adc::VertexAdcCredentialProvider::new()));
}
let mut header_map = reqwest::header::HeaderMap::new();
for (k, v) in config.headers() {
let name =
reqwest::header::HeaderName::from_bytes(k.as_bytes()).map_err(|_| LiterLlmError::InvalidHeader {
name: k.clone(),
reason: "pre-validated header name became invalid".into(),
})?;
let val = reqwest::header::HeaderValue::from_str(v).map_err(|_| LiterLlmError::InvalidHeader {
name: k.clone(),
reason: "pre-validated header value became invalid".into(),
})?;
header_map.insert(name, val);
}
let http = {
#[cfg(feature = "native-http")]
crate::ensure_crypto_provider();
let builder = reqwest::Client::builder().default_headers(header_map);
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
let builder = crate::provider::configure_outbound_client_builder(builder, config.transport.dns_cache_ttl);
#[cfg(not(target_arch = "wasm32"))]
let builder = builder.timeout(config.timeout);
#[cfg(not(target_arch = "wasm32"))]
let builder = config.transport.apply_to_builder(builder);
builder.build().map_err(LiterLlmError::from)?
};
let cached_auth_header = provider
.auth_header(config.api_key.expose_secret())
.map(|(name, value)| (name.into_owned(), value.into_owned()));
let cached_extra_headers = provider
.extra_headers()
.iter()
.map(|&(name, value)| (name.to_owned(), value.to_owned()))
.collect();
Ok(Self {
config,
http,
provider,
cached_auth_header,
cached_extra_headers,
})
}
fn resolve_provider_for_model(&self, model: &str) -> Arc<dyn Provider> {
if self.config.base_url.is_some() {
return Arc::clone(&self.provider);
}
if self.provider.matches_model(model) {
return Arc::clone(&self.provider);
}
if model.starts_with("bedrock/") {
return build_bedrock_provider(&self.config);
}
if let Some(detected) = provider::detect_provider(model) {
return Arc::from(detected);
}
Arc::clone(&self.provider)
}
async fn resolve_auth_header_for_provider(&self, prov: &dyn Provider) -> Result<Option<(String, String)>> {
if let Some(ref cp) = self.config.credential_provider {
let credential = cp.resolve().await?;
match credential {
Credential::BearerToken(token) => Ok(Some((
"Authorization".to_owned(),
format!("Bearer {}", token.expose_secret()),
))),
Credential::AwsCredentials { .. } => Ok(None),
}
} else {
Ok(prov
.auth_header(self.config.api_key.expose_secret())
.map(|(name, value)| (name.into_owned(), value.into_owned())))
}
}
fn all_headers_for_provider(
&self,
prov: &dyn Provider,
method: &str,
url: &str,
body_json: &serde_json::Value,
body_bytes: &[u8],
) -> Result<Vec<(String, String)>> {
let mut headers = prov.signing_headers(method, url, body_bytes)?;
headers.extend(
prov.extra_headers()
.iter()
.map(|&(name, value)| (name.to_owned(), value.to_owned())),
);
headers.extend(prov.dynamic_headers(body_json));
Ok(headers)
}
fn prepare_request(
&self,
serializable: &impl serde::Serialize,
endpoint_fn: impl FnOnce(&dyn Provider) -> &str,
model: &str,
stream: Option<bool>,
) -> Result<PreparedRequest> {
if model.is_empty() {
return Err(LiterLlmError::BadRequest {
message: "model must not be empty".into(),
status: 400,
});
}
let prov = self.resolve_provider_for_model(model);
let bare_model = prov.strip_model_prefix(model).to_owned();
let endpoint_path = endpoint_fn(prov.as_ref());
let url = prov.build_url(endpoint_path, &bare_model);
let mut body = serde_json::to_value(serializable)?;
if let Some(obj) = body.as_object_mut() {
obj.insert("model".into(), serde_json::Value::String(bare_model));
if let Some(s) = stream {
obj.insert("stream".into(), serde_json::Value::Bool(s));
}
}
prov.transform_request(&mut body)?;
let provider_wire_carries_stream = body.get("stream").is_some();
merge_extra_body(&mut body);
if provider_wire_carries_stream
&& let Some(s) = stream
&& let Some(obj) = body.as_object_mut()
{
obj.insert("stream".into(), serde_json::Value::Bool(s));
}
let body_bytes = bytes::Bytes::from(serde_json::to_vec(&body)?);
Ok(PreparedRequest {
url,
provider: prov,
body_json: body,
body_bytes,
})
}
async fn resolve_auth_header(&self) -> Result<Option<(String, String)>> {
if let Some(ref cp) = self.config.credential_provider {
let credential = cp.resolve().await?;
match credential {
Credential::BearerToken(token) => Ok(Some((
"Authorization".to_owned(),
format!("Bearer {}", token.expose_secret()),
))),
Credential::AwsCredentials { .. } => Ok(None),
}
} else {
Ok(self.cached_auth_header.clone())
}
}
fn all_headers(
&self,
method: &str,
url: &str,
body_json: &serde_json::Value,
body_bytes: &[u8],
) -> Result<Vec<(String, String)>> {
let mut headers = self.provider.signing_headers(method, url, body_bytes)?;
headers.extend(self.cached_extra_headers.iter().cloned());
headers.extend(self.provider.dynamic_headers(body_json));
Ok(headers)
}
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
fn build_provider(config: &ClientConfig, model_hint: Option<&str>) -> Arc<dyn Provider> {
if let Some(ref base_url) = config.base_url {
if let Some(model) = model_hint
&& model.starts_with("azure/")
{
return Arc::new(provider::azure::AzureProvider::with_base_url(base_url.clone()));
}
if let Some(model) = model_hint
&& model.starts_with("anthropic/")
{
return Arc::new(provider::anthropic::AnthropicProvider::with_base_url(base_url.clone()));
}
return Arc::new(OpenAiCompatibleProvider {
name: "custom".into(),
base_url: base_url.clone(),
env_var: None,
model_prefixes: vec![],
});
}
if let Some(model) = model_hint {
if model.starts_with("bedrock/") {
return build_bedrock_provider(config);
}
if let Some(p) = provider::detect_provider(model) {
return Arc::from(p);
}
}
Arc::new(OpenAiProvider)
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
fn build_bedrock_provider(config: &ClientConfig) -> Arc<dyn Provider> {
Arc::new(provider::bedrock::BedrockProvider::from_config(
config.bedrock_region.clone(),
config.bedrock_cross_region_prefix.clone(),
config.bedrock_access_key_id.clone(),
config.bedrock_secret_access_key.clone(),
config.bedrock_session_token.clone(),
))
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
impl LlmClient for DefaultClient {
fn chat(&self, req: ChatCompletionRequest) -> BoxFuture<'_, Result<ChatCompletionResponse>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(false))?;
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
prepared.provider.transform_response(&mut raw)?;
serde_json::from_value::<ChatCompletionResponse>(raw).map_err(LiterLlmError::from)
})
}
fn chat_stream(
&self,
req: ChatCompletionRequest,
) -> BoxFuture<'_, Result<BoxStream<'static, Result<ChatCompletionChunk>>>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(true))?;
let bare_model = prepared.provider.strip_model_prefix(&req.model);
let url = prepared
.provider
.build_stream_url(prepared.provider.chat_completions_path(), bare_model);
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
match prepared.provider.stream_format() {
provider::StreamFormat::Sse => {
let provider = Arc::clone(&prepared.provider);
let parse_event = move |data: &str| provider.parse_stream_event(data);
let stream = http::streaming::post_stream_bounded(
&self.http,
&url,
auth,
&extra,
prepared.body_bytes,
parse_event,
self.response_read_options(),
)
.await?;
Ok(stream)
}
provider::StreamFormat::AwsEventStream => {
let stream = http::eventstream::post_eventstream_bounded(
&self.http,
&url,
auth,
&extra,
prepared.body_bytes,
provider::bedrock::parse_bedrock_stream_event,
self.response_read_options(),
)
.await?;
Ok(stream)
}
}
})
}
fn embed(&self, req: EmbeddingRequest) -> BoxFuture<'_, Result<EmbeddingResponse>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.embeddings_path(), &req.model, None)?;
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
prepared.provider.transform_response(&mut raw)?;
serde_json::from_value::<EmbeddingResponse>(raw).map_err(LiterLlmError::from)
})
}
fn list_models(&self) -> BoxFuture<'_, Result<ModelsListResponse>> {
Box::pin(async move {
let url = self.provider.build_url(self.provider.models_path(), "");
let auth_header = self.resolve_auth_header().await?;
let auth = auth_header.as_ref().map(str_pair);
let all_headers = self.all_headers("GET", &url, &serde_json::Value::Null, &[])?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let mut raw = http::request::get_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
self.config.max_retries,
self.config.max_response_bytes,
)
.await?;
self.provider.transform_response(&mut raw)?;
serde_json::from_value::<ModelsListResponse>(raw).map_err(LiterLlmError::from)
})
}
fn image_generate(&self, req: CreateImageRequest) -> BoxFuture<'_, Result<ImagesResponse>> {
Box::pin(async move {
let model = req.model.as_deref().unwrap_or_default();
let prepared = self.prepare_request(&req, |p| p.image_generations_path(), model, None)?;
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
prepared.provider.transform_response(&mut raw)?;
serde_json::from_value::<ImagesResponse>(raw).map_err(LiterLlmError::from)
})
}
fn speech(&self, req: CreateSpeechRequest) -> BoxFuture<'_, Result<bytes::Bytes>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.audio_speech_path(), &req.model, None)?;
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
http::request::post_binary_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await
})
}
fn transcribe(&self, req: CreateTranscriptionRequest) -> BoxFuture<'_, Result<TranscriptionResponse>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.audio_transcriptions_path(), &req.model, None)?;
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
prepared.provider.transform_response(&mut raw)?;
serde_json::from_value::<TranscriptionResponse>(raw).map_err(LiterLlmError::from)
})
}
fn moderate(&self, req: ModerationRequest) -> BoxFuture<'_, Result<ModerationResponse>> {
Box::pin(async move {
let model = req.model.as_deref().unwrap_or_default();
let prepared = self.prepare_request(&req, |p| p.moderations_path(), model, None)?;
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
prepared.provider.transform_response(&mut raw)?;
serde_json::from_value::<ModerationResponse>(raw).map_err(LiterLlmError::from)
})
}
fn rerank(&self, req: RerankRequest) -> BoxFuture<'_, Result<RerankResponse>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.rerank_path(), &req.model, None)?;
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
prepared.provider.transform_response(&mut raw)?;
serde_json::from_value::<RerankResponse>(raw).map_err(LiterLlmError::from)
})
}
fn search(&self, req: SearchRequest) -> BoxFuture<'_, Result<SearchResponse>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.search_path(), &req.model, None)?;
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
prepared.provider.transform_response(&mut raw)?;
serde_json::from_value::<SearchResponse>(raw).map_err(LiterLlmError::from)
})
}
fn ocr(&self, req: OcrRequest) -> BoxFuture<'_, Result<OcrResponse>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.ocr_path(), &req.model, None)?;
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
prepared.provider.transform_response(&mut raw)?;
serde_json::from_value::<OcrResponse>(raw).map_err(LiterLlmError::from)
})
}
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
impl LlmClientRaw for DefaultClient {
fn chat_raw(&self, req: ChatCompletionRequest) -> BoxFuture<'_, Result<RawExchange<ChatCompletionResponse>>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(false))?;
let raw_request = prepared.body_json.clone();
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
let raw_response = Some(raw.clone());
prepared.provider.transform_response(&mut raw)?;
let data = serde_json::from_value::<ChatCompletionResponse>(raw).map_err(LiterLlmError::from)?;
Ok(RawExchange {
data,
raw_request,
raw_response,
})
})
}
fn chat_stream_raw(
&self,
req: ChatCompletionRequest,
) -> BoxFuture<'_, Result<RawStreamExchange<BoxStream<'static, Result<ChatCompletionChunk>>>>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(true))?;
let raw_request = prepared.body_json.clone();
let bare_model = prepared.provider.strip_model_prefix(&req.model);
let url = prepared
.provider
.build_stream_url(prepared.provider.chat_completions_path(), bare_model);
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let stream = match prepared.provider.stream_format() {
provider::StreamFormat::Sse => {
let provider = Arc::clone(&prepared.provider);
let parse_event = move |data: &str| provider.parse_stream_event(data);
http::streaming::post_stream_bounded(
&self.http,
&url,
auth,
&extra,
prepared.body_bytes,
parse_event,
self.response_read_options(),
)
.await?
}
provider::StreamFormat::AwsEventStream => {
http::eventstream::post_eventstream_bounded(
&self.http,
&url,
auth,
&extra,
prepared.body_bytes,
provider::bedrock::parse_bedrock_stream_event,
self.response_read_options(),
)
.await?
}
};
Ok(RawStreamExchange { stream, raw_request })
})
}
fn embed_raw(&self, req: EmbeddingRequest) -> BoxFuture<'_, Result<RawExchange<EmbeddingResponse>>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.embeddings_path(), &req.model, None)?;
let raw_request = prepared.body_json.clone();
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
let raw_response = Some(raw.clone());
prepared.provider.transform_response(&mut raw)?;
let data = serde_json::from_value::<EmbeddingResponse>(raw).map_err(LiterLlmError::from)?;
Ok(RawExchange {
data,
raw_request,
raw_response,
})
})
}
fn image_generate_raw(&self, req: CreateImageRequest) -> BoxFuture<'_, Result<RawExchange<ImagesResponse>>> {
Box::pin(async move {
let model = req.model.as_deref().unwrap_or_default();
let prepared = self.prepare_request(&req, |p| p.image_generations_path(), model, None)?;
let raw_request = prepared.body_json.clone();
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
let raw_response = Some(raw.clone());
prepared.provider.transform_response(&mut raw)?;
let data = serde_json::from_value::<ImagesResponse>(raw).map_err(LiterLlmError::from)?;
Ok(RawExchange {
data,
raw_request,
raw_response,
})
})
}
fn transcribe_raw(
&self,
req: CreateTranscriptionRequest,
) -> BoxFuture<'_, Result<RawExchange<TranscriptionResponse>>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.audio_transcriptions_path(), &req.model, None)?;
let raw_request = prepared.body_json.clone();
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
let raw_response = Some(raw.clone());
prepared.provider.transform_response(&mut raw)?;
let data = serde_json::from_value::<TranscriptionResponse>(raw).map_err(LiterLlmError::from)?;
Ok(RawExchange {
data,
raw_request,
raw_response,
})
})
}
fn moderate_raw(&self, req: ModerationRequest) -> BoxFuture<'_, Result<RawExchange<ModerationResponse>>> {
Box::pin(async move {
let model = req.model.as_deref().unwrap_or_default();
let prepared = self.prepare_request(&req, |p| p.moderations_path(), model, None)?;
let raw_request = prepared.body_json.clone();
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
let raw_response = Some(raw.clone());
prepared.provider.transform_response(&mut raw)?;
let data = serde_json::from_value::<ModerationResponse>(raw).map_err(LiterLlmError::from)?;
Ok(RawExchange {
data,
raw_request,
raw_response,
})
})
}
fn rerank_raw(&self, req: RerankRequest) -> BoxFuture<'_, Result<RawExchange<RerankResponse>>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.rerank_path(), &req.model, None)?;
let raw_request = prepared.body_json.clone();
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
let raw_response = Some(raw.clone());
prepared.provider.transform_response(&mut raw)?;
let data = serde_json::from_value::<RerankResponse>(raw).map_err(LiterLlmError::from)?;
Ok(RawExchange {
data,
raw_request,
raw_response,
})
})
}
fn search_raw(&self, req: SearchRequest) -> BoxFuture<'_, Result<RawExchange<SearchResponse>>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.search_path(), &req.model, None)?;
let raw_request = prepared.body_json.clone();
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
let raw_response = Some(raw.clone());
prepared.provider.transform_response(&mut raw)?;
let data = serde_json::from_value::<SearchResponse>(raw).map_err(LiterLlmError::from)?;
Ok(RawExchange {
data,
raw_request,
raw_response,
})
})
}
fn ocr_raw(&self, req: OcrRequest) -> BoxFuture<'_, Result<RawExchange<OcrResponse>>> {
Box::pin(async move {
let prepared = self.prepare_request(&req, |p| p.ocr_path(), &req.model, None)?;
let raw_request = prepared.body_json.clone();
let auth_header = self
.resolve_auth_header_for_provider(prepared.provider.as_ref())
.await?;
let all_headers = self.all_headers_for_provider(
prepared.provider.as_ref(),
"POST",
&prepared.url,
&prepared.body_json,
&prepared.body_bytes,
)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let mut raw = http::request::post_json_raw_bounded(
&self.http,
&prepared.url,
auth,
&extra,
prepared.body_bytes,
self.response_read_options(),
)
.await?;
let raw_response = Some(raw.clone());
prepared.provider.transform_response(&mut raw)?;
let data = serde_json::from_value::<OcrResponse>(raw).map_err(LiterLlmError::from)?;
Ok(RawExchange {
data,
raw_request,
raw_response,
})
})
}
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
impl FileClient for DefaultClient {
fn create_file(&self, req: CreateFileRequest) -> BoxFuture<'_, Result<FileObject>> {
Box::pin(async move {
let url = self.provider.build_url(self.provider.files_path(), "");
let auth_header = self.resolve_auth_header().await?;
let auth = auth_header.as_ref().map(str_pair);
let all_headers = self.all_headers("POST", &url, &serde_json::Value::Null, &[])?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
use base64::Engine;
let file_bytes = base64::engine::general_purpose::STANDARD
.decode(&req.file)
.map_err(|e| LiterLlmError::BadRequest {
message: format!("invalid base64 file data: {e}"),
status: 400,
})?;
let filename = req.filename.unwrap_or_else(|| "upload".to_owned());
let file_part = reqwest::multipart::Part::bytes(file_bytes).file_name(filename);
let purpose_str = serde_json::to_value(&req.purpose)?
.as_str()
.unwrap_or_default()
.to_owned();
let form = reqwest::multipart::Form::new()
.part("file", file_part)
.text("purpose", purpose_str);
let raw = http::request::post_multipart_bounded(
&self.http,
&url,
auth,
&extra,
form,
self.config.max_response_bytes,
)
.await?;
serde_json::from_value::<FileObject>(raw).map_err(LiterLlmError::from)
})
}
fn retrieve_file(&self, file_id: &str) -> BoxFuture<'_, Result<FileObject>> {
let file_id = file_id.to_owned();
Box::pin(async move {
let url = format!(
"{}/{}",
self.provider.build_url(self.provider.files_path(), ""),
file_id
);
let auth_header = self.resolve_auth_header().await?;
let auth = auth_header.as_ref().map(str_pair);
let all_headers = self.all_headers("GET", &url, &serde_json::Value::Null, &[])?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let raw = http::request::get_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
self.config.max_retries,
self.config.max_response_bytes,
)
.await?;
serde_json::from_value::<FileObject>(raw).map_err(LiterLlmError::from)
})
}
fn delete_file(&self, file_id: &str) -> BoxFuture<'_, Result<DeleteResponse>> {
let file_id = file_id.to_owned();
Box::pin(async move {
let url = format!(
"{}/{}",
self.provider.build_url(self.provider.files_path(), ""),
file_id
);
let auth_header = self.resolve_auth_header().await?;
let auth = auth_header.as_ref().map(str_pair);
let all_headers = self.all_headers("DELETE", &url, &serde_json::Value::Null, &[])?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let raw = http::request::delete_json_bounded(
&self.http,
&url,
auth,
&extra,
self.config.max_retries,
self.config.max_response_bytes,
)
.await?;
serde_json::from_value::<DeleteResponse>(raw).map_err(LiterLlmError::from)
})
}
fn list_files(&self, query: Option<FileListQuery>) -> BoxFuture<'_, Result<FileListResponse>> {
Box::pin(async move {
let base_url = self.provider.build_url(self.provider.files_path(), "");
let url = if let Some(ref q) = query {
let mut params = Vec::new();
if let Some(ref purpose) = q.purpose {
params.push(format!("purpose={purpose}"));
}
if let Some(limit) = q.limit {
params.push(format!("limit={limit}"));
}
if let Some(ref after) = q.after {
params.push(format!("after={after}"));
}
if params.is_empty() {
base_url
} else {
format!("{base_url}?{}", params.join("&"))
}
} else {
base_url
};
let auth_header = self.resolve_auth_header().await?;
let auth = auth_header.as_ref().map(str_pair);
let all_headers = self.all_headers("GET", &url, &serde_json::Value::Null, &[])?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let raw = http::request::get_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
self.config.max_retries,
self.config.max_response_bytes,
)
.await?;
serde_json::from_value::<FileListResponse>(raw).map_err(LiterLlmError::from)
})
}
fn file_content(&self, file_id: &str) -> BoxFuture<'_, Result<bytes::Bytes>> {
let file_id = file_id.to_owned();
Box::pin(async move {
let url = format!(
"{}/{}/content",
self.provider.build_url(self.provider.files_path(), ""),
file_id
);
let auth_header = self.resolve_auth_header().await?;
let auth = auth_header.as_ref().map(str_pair);
let all_headers = self.all_headers("GET", &url, &serde_json::Value::Null, &[])?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
http::request::get_binary_bounded(
&self.http,
&url,
auth,
&extra,
self.config.max_retries,
self.config.max_response_bytes,
)
.await
})
}
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
impl BatchClient for DefaultClient {
fn create_batch(&self, req: CreateBatchRequest) -> BoxFuture<'_, Result<BatchObject>> {
Box::pin(async move {
let url = self.provider.build_url(self.provider.batches_path(), "");
let body_bytes = bytes::Bytes::from(serde_json::to_vec(&req)?);
let body_json = serde_json::to_value(&req)?;
let auth_header = self.resolve_auth_header().await?;
let all_headers = self.all_headers("POST", &url, &body_json, &body_bytes)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let raw = http::request::post_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
body_bytes,
self.response_read_options(),
)
.await?;
serde_json::from_value::<BatchObject>(raw).map_err(LiterLlmError::from)
})
}
fn retrieve_batch(&self, batch_id: &str) -> BoxFuture<'_, Result<BatchObject>> {
let batch_id = batch_id.to_owned();
Box::pin(async move {
let url = format!(
"{}/{}",
self.provider.build_url(self.provider.batches_path(), ""),
batch_id
);
let auth_header = self.resolve_auth_header().await?;
let auth = auth_header.as_ref().map(str_pair);
let all_headers = self.all_headers("GET", &url, &serde_json::Value::Null, &[])?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let raw = http::request::get_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
self.config.max_retries,
self.config.max_response_bytes,
)
.await?;
serde_json::from_value::<BatchObject>(raw).map_err(LiterLlmError::from)
})
}
fn list_batches(&self, query: Option<BatchListQuery>) -> BoxFuture<'_, Result<BatchListResponse>> {
Box::pin(async move {
let base_url = self.provider.build_url(self.provider.batches_path(), "");
let url = if let Some(ref q) = query {
let mut params = Vec::new();
if let Some(limit) = q.limit {
params.push(format!("limit={limit}"));
}
if let Some(ref after) = q.after {
params.push(format!("after={after}"));
}
if params.is_empty() {
base_url
} else {
format!("{base_url}?{}", params.join("&"))
}
} else {
base_url
};
let auth_header = self.resolve_auth_header().await?;
let auth = auth_header.as_ref().map(str_pair);
let all_headers = self.all_headers("GET", &url, &serde_json::Value::Null, &[])?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let raw = http::request::get_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
self.config.max_retries,
self.config.max_response_bytes,
)
.await?;
serde_json::from_value::<BatchListResponse>(raw).map_err(LiterLlmError::from)
})
}
fn cancel_batch(&self, batch_id: &str) -> BoxFuture<'_, Result<BatchObject>> {
let batch_id = batch_id.to_owned();
Box::pin(async move {
let url = format!(
"{}/{}/cancel",
self.provider.build_url(self.provider.batches_path(), ""),
batch_id
);
let auth_header = self.resolve_auth_header().await?;
let body_json = serde_json::Value::Null;
let body_bytes = bytes::Bytes::new();
let all_headers = self.all_headers("POST", &url, &body_json, &body_bytes)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let raw = http::request::post_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
body_bytes,
self.response_read_options(),
)
.await?;
serde_json::from_value::<BatchObject>(raw).map_err(LiterLlmError::from)
})
}
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
#[doc(hidden)]
#[cfg_attr(alef, alef(skip))]
pub trait BatchRetriever {
async fn fetch_batch_for_polling(&self, batch_id: &str) -> Result<BatchObject>;
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
impl BatchRetriever for DefaultClient {
async fn fetch_batch_for_polling(&self, batch_id: &str) -> Result<BatchObject> {
self.retrieve_batch(batch_id).await
}
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
#[doc(hidden)]
#[cfg_attr(alef, alef(skip))]
pub async fn wait_for_batch_impl<R: BatchRetriever>(
retriever: &R,
batch_id: &str,
config: WaitForBatchConfig,
) -> std::result::Result<BatchObject, BatchWaitError> {
#[cfg(not(target_arch = "wasm32"))]
let started = tokio::time::Instant::now();
#[cfg(target_arch = "wasm32")]
let started = web_time::Instant::now();
let mut interval_secs = config.initial_interval_secs;
loop {
let batch = retriever.fetch_batch_for_polling(batch_id).await?;
match batch.status {
BatchStatus::Completed => return Ok(batch),
BatchStatus::Failed | BatchStatus::Expired | BatchStatus::Cancelled => {
return Err(BatchWaitError::Failed { status: batch.status });
}
BatchStatus::Validating | BatchStatus::InProgress | BatchStatus::Finalizing | BatchStatus::Cancelling => {
if let Some(timeout_secs) = config.timeout_secs {
let timeout = Duration::from_secs_f64(timeout_secs);
if started.elapsed() >= timeout {
return Err(BatchWaitError::Timeout { timeout_secs });
}
}
#[cfg(not(target_arch = "wasm32"))]
tokio::time::sleep(Duration::from_secs_f64(interval_secs)).await;
#[cfg(target_arch = "wasm32")]
gloo_timers::future::sleep(Duration::from_secs_f64(interval_secs)).await;
let next =
(interval_secs as f32 * config.backoff_multiplier).min(config.max_interval_secs as f32) as f64;
interval_secs = next;
}
}
}
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
impl DefaultClient {
pub async fn wait_for_batch(
&self,
batch_id: &str,
config: WaitForBatchConfig,
) -> std::result::Result<BatchObject, BatchWaitError> {
wait_for_batch_impl(self, batch_id, config).await
}
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
fn parse_response_stream_event(data: &str) -> Result<Option<ResponseStreamEvent>> {
serde_json::from_str::<ResponseStreamEvent>(data)
.map(Some)
.map_err(|e| LiterLlmError::Streaming {
message: format!("failed to parse Responses SSE data: {e}"),
})
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
fn reject_malformed_responses_url(url: &str, provider_name: &str) -> Result<()> {
let path = url.split_once("://").map_or(url, |(_, rest)| rest);
let path = path.split_once('/').map_or("", |(_, rest)| rest);
let path = path.split(['?', '#']).next().unwrap_or("");
if format!("/{path}").contains("//") {
return Err(LiterLlmError::EndpointNotSupported {
endpoint: "responses".into(),
provider: provider_name.into(),
});
}
Ok(())
}
#[cfg(any(feature = "native-http", feature = "wasm-http"))]
impl ResponseClient for DefaultClient {
fn create_response(&self, req: CreateResponseRequest) -> BoxFuture<'_, Result<ResponseObject>> {
Box::pin(async move {
let url = self.provider.build_url(self.provider.responses_path(), "");
reject_malformed_responses_url(&url, self.provider.name())?;
let body_json = response_request_body(&req)?;
let body_bytes = bytes::Bytes::from(serde_json::to_vec(&body_json)?);
let auth_header = self.resolve_auth_header().await?;
let all_headers = self.all_headers("POST", &url, &body_json, &body_bytes)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let raw = http::request::post_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
body_bytes,
self.response_read_options(),
)
.await?;
serde_json::from_value::<ResponseObject>(raw).map_err(LiterLlmError::from)
})
}
fn create_response_stream(
&self,
req: CreateResponseRequest,
) -> BoxFuture<'_, Result<BoxStream<'static, Result<ResponseStreamEvent>>>> {
Box::pin(async move {
let url = self.provider.build_stream_url(self.provider.responses_path(), "");
reject_malformed_responses_url(&url, self.provider.name())?;
let body_json = response_stream_request_body(req)?;
let body_bytes = bytes::Bytes::from(serde_json::to_vec(&body_json)?);
let auth_header = self.resolve_auth_header().await?;
let all_headers = self.all_headers("POST", &url, &body_json, &body_bytes)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
http::streaming::post_stream_bounded(
&self.http,
&url,
auth,
&extra,
body_bytes,
parse_response_stream_event,
self.response_read_options(),
)
.await
})
}
fn retrieve_response(&self, response_id: &str) -> BoxFuture<'_, Result<ResponseObject>> {
let response_id = response_id.to_owned();
Box::pin(async move {
let base_url = self.provider.build_url(self.provider.responses_path(), "");
reject_malformed_responses_url(&base_url, self.provider.name())?;
let url = format!("{base_url}/{response_id}");
let auth_header = self.resolve_auth_header().await?;
let auth = auth_header.as_ref().map(str_pair);
let all_headers = self.all_headers("GET", &url, &serde_json::Value::Null, &[])?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let raw = http::request::get_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
self.config.max_retries,
self.config.max_response_bytes,
)
.await?;
serde_json::from_value::<ResponseObject>(raw).map_err(LiterLlmError::from)
})
}
fn cancel_response(&self, response_id: &str) -> BoxFuture<'_, Result<ResponseObject>> {
let response_id = response_id.to_owned();
Box::pin(async move {
let base_url = self.provider.build_url(self.provider.responses_path(), "");
reject_malformed_responses_url(&base_url, self.provider.name())?;
let url = format!("{base_url}/{response_id}/cancel");
let auth_header = self.resolve_auth_header().await?;
let body_json = serde_json::Value::Null;
let body_bytes = bytes::Bytes::new();
let all_headers = self.all_headers("POST", &url, &body_json, &body_bytes)?;
let extra: Vec<(&str, &str)> = all_headers.iter().map(|(n, v)| (n.as_str(), v.as_str())).collect();
let auth = auth_header.as_ref().map(str_pair);
let raw = http::request::post_json_raw_bounded(
&self.http,
&url,
auth,
&extra,
body_bytes,
self.response_read_options(),
)
.await?;
serde_json::from_value::<ResponseObject>(raw).map_err(LiterLlmError::from)
})
}
}
#[cfg(all(test, any(feature = "native-http", feature = "wasm-http")))]
mod build_provider_tests {
use super::*;
use crate::client::config::ClientConfigBuilder;
#[test]
fn reject_malformed_responses_url_rejects_azure_empty_deployment() {
let provider = provider::azure::AzureProvider::with_base_url("https://resourceA.openai.azure.com");
let url = provider.build_url(provider.responses_path(), "");
assert!(
url.contains("/openai/deployments//responses"),
"test setup assumption broken, url = {url}"
);
let result = reject_malformed_responses_url(&url, provider.name());
match result {
Err(LiterLlmError::EndpointNotSupported {
endpoint,
provider: rejected_provider,
}) => {
assert_eq!(endpoint, "responses");
assert_eq!(rejected_provider, "azure");
}
other => panic!("expected Err(EndpointNotSupported), got {other:?}"),
}
}
#[test]
fn reject_malformed_responses_url_allows_azure_with_pinned_deployment() {
let base_url = "https://resourceA.openai.azure.com/openai/deployments/my-gpt4-deployment";
let provider = provider::azure::AzureProvider::with_base_url(base_url);
let url = provider.build_url(provider.responses_path(), "");
assert!(
!url.contains("deployments//"),
"test setup assumption broken: pinned-deployment URL should have no empty segment, url = {url}"
);
assert!(
reject_malformed_responses_url(&url, provider.name()).is_ok(),
"a pinned-deployment Azure URL must not be rejected, url = {url}"
);
}
#[test]
fn reject_malformed_responses_url_allows_bedrock_google_ai_and_vertex() {
let bedrock = provider::bedrock::BedrockProvider::from_config(
Some("us-east-1".into()),
None,
Some("test-access-key".into()),
Some("test-secret-key".into()),
None,
);
let bedrock_url = bedrock.build_url(bedrock.responses_path(), "");
assert!(
reject_malformed_responses_url(&bedrock_url, bedrock.name()).is_ok(),
"bedrock url = {bedrock_url}"
);
let google_ai = provider::google_ai::GoogleAiProvider;
let google_ai_url = google_ai.build_url(google_ai.responses_path(), "");
assert!(
reject_malformed_responses_url(&google_ai_url, google_ai.name()).is_ok(),
"google ai url = {google_ai_url}"
);
let vertex = provider::vertex::VertexAiProvider::new("my-project", "us-central1");
let vertex_url = vertex.build_url(vertex.responses_path(), "");
assert!(
reject_malformed_responses_url(&vertex_url, vertex.name()).is_ok(),
"vertex url = {vertex_url}"
);
}
#[test]
fn azure_model_with_per_model_base_url_uses_azure_provider() {
let config = ClientConfigBuilder::new("test-key")
.base_url("https://resourceA.cognitiveservices.azure.com")
.build();
let p = build_provider(&config, Some("azure/gpt-5-mini"));
assert_eq!(p.name(), "azure");
let url = p.build_url("/chat/completions", "gpt-5-mini");
assert!(
url.starts_with("https://resourceA.cognitiveservices.azure.com/openai/deployments/gpt-5-mini/chat/completions?api-version="),
"url = {url}"
);
}
#[test]
fn non_azure_model_with_base_url_uses_openai_compatible() {
let config = ClientConfigBuilder::new("test-key")
.base_url("http://localhost:11434/v1")
.build();
let p = build_provider(&config, Some("llama3.1:8b"));
assert_eq!(p.name(), "custom");
let url = p.build_url("/chat/completions", "llama3.1:8b");
assert_eq!(url, "http://localhost:11434/v1/chat/completions");
}
#[test]
fn no_base_url_falls_through_to_detect_provider() {
let config = ClientConfigBuilder::new("test-key").build();
let p = build_provider(&config, Some("azure/gpt-4o"));
assert_eq!(p.name(), "azure");
}
#[test]
fn anthropic_model_with_per_model_base_url_uses_anthropic_provider() {
let config = ClientConfigBuilder::new("test-key")
.base_url("https://proxy.internal/anthropic/")
.build();
let p = build_provider(&config, Some("anthropic/claude-3-5-sonnet-20241022"));
assert_eq!(p.name(), "anthropic");
assert_eq!(p.base_url(), "https://proxy.internal/anthropic");
}
#[test]
fn default_anthropic_provider_uses_official_base_url() {
assert_eq!(
provider::anthropic::AnthropicProvider::new().base_url(),
"https://api.anthropic.com/v1"
);
assert_eq!(
provider::anthropic::AnthropicProvider::default().base_url(),
"https://api.anthropic.com/v1"
);
}
#[test]
fn extra_body_object_is_merged_and_overrides_existing_key() {
let client = DefaultClient::new(ClientConfigBuilder::new("test-key").build(), Some("gpt-4"))
.expect("client construction should succeed");
let req = ChatCompletionRequest {
model: "gpt-4".into(),
messages: vec![],
extra_body: Some(serde_json::json!({
"thinking": {"type": "enabled"},
"model": "extra-body-override"
})),
..Default::default()
};
let prepared = client
.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(false))
.expect("prepare_request should not fail");
assert!(
prepared.body_json.get("extra_body").is_none(),
"extra_body key must be removed from the final body"
);
assert_eq!(prepared.body_json["thinking"], serde_json::json!({"type": "enabled"}));
assert_eq!(
prepared.body_json["model"], "extra-body-override",
"extra_body keys must override identically named top-level keys"
);
}
#[test]
fn extra_body_cannot_override_the_transport_controlled_stream_flag() {
let client = DefaultClient::new(ClientConfigBuilder::new("test-key").build(), Some("gpt-4"))
.expect("client construction should succeed");
let req = ChatCompletionRequest {
model: "gpt-4".into(),
messages: vec![],
extra_body: Some(serde_json::json!({"stream": false})),
..Default::default()
};
let prepared = client
.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(true))
.expect("prepare_request should not fail");
assert_eq!(
prepared.body_json["stream"],
serde_json::json!(true),
"the caller-selected stream flag must win over extra_body"
);
}
#[test]
fn extra_body_non_object_is_dropped_without_reaching_the_body() {
let client = DefaultClient::new(ClientConfigBuilder::new("test-key").build(), Some("gpt-4"))
.expect("client construction should succeed");
let req = ChatCompletionRequest {
model: "gpt-4".into(),
messages: vec![],
extra_body: Some(serde_json::json!(["not", "an", "object"])),
..Default::default()
};
let prepared = client
.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(false))
.expect("prepare_request should not fail");
assert!(
prepared.body_json.get("extra_body").is_none(),
"non-object extra_body must never reach the wire"
);
}
#[test]
fn anthropic_provider_leaves_no_extra_body_key_in_final_body() {
let client = DefaultClient::new(
ClientConfigBuilder::new("test-key").build(),
Some("claude-3-5-sonnet-20241022"),
)
.expect("client construction should succeed");
let req = ChatCompletionRequest {
model: "claude-3-5-sonnet-20241022".into(),
messages: vec![crate::types::Message::User(crate::types::UserMessage {
content: crate::types::UserContent::Text("Hi".into()),
name: None,
})],
max_tokens: Some(100),
extra_body: Some(serde_json::json!({"reasoning_effort": "high"})),
..Default::default()
};
let prepared = client
.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(false))
.expect("prepare_request should not fail");
assert!(
prepared.body_json.get("extra_body").is_none(),
"anthropic's own transform_request already strips extra_body"
);
assert_eq!(prepared.body_json["thinking"]["budget_tokens"], 16384);
}
fn base_response_request() -> CreateResponseRequest {
CreateResponseRequest {
model: "gpt-5".into(),
input: serde_json::json!("What is the capital of France?"),
..Default::default()
}
}
#[test]
fn responses_extra_body_object_is_merged_and_overrides_existing_key() {
let mut req = base_response_request();
req.extra_body = Some(serde_json::json!({
"reasoning": {"effort": "none"},
"model": "extra-body-override"
}));
let body = response_request_body(&req).expect("body should serialize");
assert!(
body.get("extra_body").is_none(),
"extra_body key must be removed from the final Responses body"
);
assert_eq!(body["reasoning"], serde_json::json!({"effort": "none"}));
assert_eq!(
body["model"], "extra-body-override",
"extra_body keys must override identically named top-level keys, matching chat-path semantics"
);
}
#[test]
fn responses_extra_body_non_object_is_dropped_without_reaching_the_body() {
let mut req = base_response_request();
req.extra_body = Some(serde_json::json!(["not", "an", "object"]));
let body = response_request_body(&req).expect("body should serialize");
assert!(
body.get("extra_body").is_none(),
"non-object extra_body must never reach the wire"
);
assert_eq!(
body["model"], "gpt-5",
"a dropped extra_body must leave the real fields untouched"
);
}
#[test]
fn responses_streaming_path_also_merges_extra_body_and_keeps_stream_forced() {
let mut req = base_response_request();
req.extra_body = Some(serde_json::json!({
"reasoning": {"effort": "none"},
"stream": false
}));
let body = response_stream_request_body(req).expect("body should serialize");
assert!(
body.get("extra_body").is_none(),
"extra_body key must be removed from the streaming Responses body too"
);
assert_eq!(body["reasoning"], serde_json::json!({"effort": "none"}));
assert_eq!(
body["stream"],
serde_json::json!(true),
"the streaming transport's stream flag must win over extra_body"
);
}
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
#[test]
#[serial_test::serial]
fn vertex_ai_auto_installs_adc_provider_when_no_credentials_configured() {
let prior_project = std::env::var("VERTEXAI_PROJECT").ok();
let prior_location = std::env::var("VERTEXAI_LOCATION").ok();
struct EnvGuard {
prior_project: Option<String>,
prior_location: Option<String>,
}
impl Drop for EnvGuard {
fn drop(&mut self) {
unsafe {
match &self.prior_project {
Some(v) => std::env::set_var("VERTEXAI_PROJECT", v),
None => std::env::remove_var("VERTEXAI_PROJECT"),
}
match &self.prior_location {
Some(v) => std::env::set_var("VERTEXAI_LOCATION", v),
None => std::env::remove_var("VERTEXAI_LOCATION"),
}
}
}
}
let _guard = EnvGuard {
prior_project,
prior_location,
};
unsafe {
std::env::set_var("VERTEXAI_PROJECT", "test-project");
std::env::set_var("VERTEXAI_LOCATION", "us-central1");
}
let config = ClientConfigBuilder::new("").load_env(false).build();
assert!(
config.credential_provider.is_none(),
"input config should have no credential_provider"
);
let client = DefaultClient::new(config, Some("vertex_ai/gemini-2.5-flash-lite"))
.expect("DefaultClient::new should succeed for vertex with empty api_key");
assert!(
client.config.credential_provider.is_some(),
"DefaultClient::new should auto-install VertexAdcCredentialProvider for vertex_ai when no credentials are configured"
);
}
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
#[test]
#[serial_test::serial]
fn vertex_ai_explicit_api_key_skips_auto_install() {
let prior_project = std::env::var("VERTEXAI_PROJECT").ok();
struct EnvGuard(Option<String>);
impl Drop for EnvGuard {
fn drop(&mut self) {
unsafe {
match &self.0 {
Some(v) => std::env::set_var("VERTEXAI_PROJECT", v),
None => std::env::remove_var("VERTEXAI_PROJECT"),
}
}
}
}
let _guard = EnvGuard(prior_project);
unsafe {
std::env::set_var("VERTEXAI_PROJECT", "test-project");
}
let config = ClientConfigBuilder::new("ya29.pre-obtained-token")
.load_env(false)
.build();
let client = DefaultClient::new(config, Some("vertex_ai/gemini-2.5-flash-lite"))
.expect("DefaultClient::new should succeed");
assert!(
client.config.credential_provider.is_none(),
"auto-install must not fire when api_key is non-empty"
);
}
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
#[test]
#[serial_test::serial]
fn vertex_ai_explicit_credential_provider_skips_auto_install() {
use std::sync::Arc;
use crate::auth::{Credential, CredentialProvider, StaticTokenProvider};
use secrecy::SecretString;
let prior_project = std::env::var("VERTEXAI_PROJECT").ok();
struct EnvGuard(Option<String>);
impl Drop for EnvGuard {
fn drop(&mut self) {
unsafe {
match &self.0 {
Some(v) => std::env::set_var("VERTEXAI_PROJECT", v),
None => std::env::remove_var("VERTEXAI_PROJECT"),
}
}
}
}
let _guard = EnvGuard(prior_project);
unsafe {
std::env::set_var("VERTEXAI_PROJECT", "test-project");
}
let explicit: Arc<dyn CredentialProvider> =
Arc::new(StaticTokenProvider::new(SecretString::from("static-token".to_owned())));
let explicit_marker = Arc::as_ptr(&explicit) as *const ();
let config = ClientConfigBuilder::new("")
.load_env(false)
.credential_provider(Arc::clone(&explicit))
.build();
let client = DefaultClient::new(config, Some("vertex_ai/gemini-2.5-flash-lite"))
.expect("DefaultClient::new should succeed");
let installed = client
.config
.credential_provider
.as_ref()
.expect("explicit provider should survive auto-install path");
let installed_marker = Arc::as_ptr(installed) as *const ();
assert_eq!(
installed_marker, explicit_marker,
"auto-install must not overwrite an explicitly-supplied credential_provider"
);
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("tokio runtime");
let credential = rt.block_on(installed.resolve()).expect("resolve");
match credential {
Credential::BearerToken(t) => {
use secrecy::ExposeSecret;
assert_eq!(t.expose_secret(), "static-token");
}
_ => panic!("expected BearerToken"),
}
}
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
fn vertex_client() -> DefaultClient {
DefaultClient::new(
ClientConfigBuilder::new("ya29.pre-obtained-token")
.load_env(false)
.build(),
Some("vertex_ai/gemini-2.5-flash"),
)
.expect("DefaultClient::new should succeed")
}
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
struct VertexProjectGuard(Option<String>);
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
impl VertexProjectGuard {
fn set() -> Self {
let prior = std::env::var("VERTEXAI_PROJECT").ok();
unsafe {
std::env::set_var("VERTEXAI_PROJECT", "test-project");
}
Self(prior)
}
}
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
impl Drop for VertexProjectGuard {
fn drop(&mut self) {
unsafe {
match &self.0 {
Some(value) => std::env::set_var("VERTEXAI_PROJECT", value),
None => std::env::remove_var("VERTEXAI_PROJECT"),
}
}
}
}
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
#[test]
#[serial_test::serial]
fn vertex_non_streaming_request_carries_no_stream_key() {
let _guard = VertexProjectGuard::set();
let client = vertex_client();
let req = ChatCompletionRequest {
model: "vertex_ai/gemini-2.5-flash".into(),
messages: vec![],
..Default::default()
};
let prepared = client
.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(false))
.expect("prepare_request should not fail");
assert!(
prepared.body_json.get("stream").is_none(),
"generateContent has no `stream` field; sending it is a hard 400, body was: {}",
prepared.body_json
);
}
#[cfg(all(feature = "native-http", not(target_arch = "wasm32")))]
#[test]
#[serial_test::serial]
fn vertex_streaming_request_carries_no_stream_key() {
let _guard = VertexProjectGuard::set();
let client = vertex_client();
let req = ChatCompletionRequest {
model: "vertex_ai/gemini-2.5-flash".into(),
messages: vec![],
..Default::default()
};
let prepared = client
.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(true))
.expect("prepare_request should not fail");
assert!(
prepared.body_json.get("stream").is_none(),
"streaming is endpoint-selected on Vertex; body was: {}",
prepared.body_json
);
}
#[test]
fn openai_shaped_provider_still_carries_the_stream_flag() {
let client = DefaultClient::new(ClientConfigBuilder::new("test-key").build(), Some("gpt-4"))
.expect("client construction should succeed");
let req = ChatCompletionRequest {
model: "gpt-4".into(),
messages: vec![],
..Default::default()
};
for streaming in [false, true] {
let prepared = client
.prepare_request(&req, |p| p.chat_completions_path(), &req.model, Some(streaming))
.expect("prepare_request should not fail");
assert_eq!(
prepared.body_json["stream"], streaming,
"OpenAI-shaped bodies must keep the transport-controlled stream flag"
);
}
}
}