use std::sync::Arc;
use std::time::Duration;
#[cfg(target_arch = "wasm32")]
use std::rc::Rc;
use url::Url;
use futures::FutureExt;
use crate::error::{Error, Result};
use crate::request;
use crate::retry::RetryConfig;
use crate::stream::EventStream;
use crate::types::{
ActivityOptions, ActivityResponse, AssignKeysRequest, AssignKeysResponse, AssignMembersRequest,
AssignMembersResponse, BulkAddWorkspaceMembersResponse, BulkRemoveWorkspaceMembersResponse,
BulkWorkspaceMembersRequest, ChatCompletionRequest, ChatCompletionResponse, CompletionRequest,
CompletionResponse, CreateGuardrailRequest, CreateKeyRequest, CreateKeyResponse,
CreateWorkspaceRequest, CreateWorkspaceResponse, CreditsResponse, DeleteGuardrailResponse,
DeleteKeyResponse, DeleteWorkspaceResponse, GetKeyByHashResponse, GetWorkspaceResponse,
Guardrail, KeyResponse, ListGuardrailKeyAssignmentsResponse,
ListGuardrailMemberAssignmentsResponse, ListGuardrailsOptions, ListGuardrailsResponse,
ListKeysOptions, ListKeysResponse, ListModelsOptions, ListOrganizationMembersOptions,
ListOrganizationMembersResponse, ListWorkspacesOptions, ListWorkspacesResponse,
ModelEndpointsResponse, ModelsResponse, Provider, ProvidersResponse, RerankRequest,
RerankResponse, SpeechFormat, SpeechRequest, SpeechResponse, UpdateGuardrailRequest,
UpdateKeyRequest, UpdateKeyResponse, UpdateWorkspaceRequest, UpdateWorkspaceResponse,
VideoContentResponse, VideoGenerationRequest, VideoGenerationResponse, VideoModelsResponse,
ZdrEndpointsResponse,
};
const DEFAULT_BASE_URL: &str = "https://openrouter.ai/api/v1/";
const DEFAULT_STREAM_RECONNECTS: u32 = 3;
#[derive(Clone, Debug)]
pub struct Client {
inner: Arc<ClientInner>,
}
#[derive(Debug)]
struct ClientInner {
api_key: String,
base_url: Url,
http: reqwest::Client,
retry: RetryConfig,
stream_reconnects: u32,
app_name: Option<String>,
referer: Option<String>,
}
impl Client {
pub fn builder() -> ClientBuilder {
ClientBuilder::default()
}
pub fn new(api_key: impl Into<String>) -> Result<Self> {
Self::builder().api_key(api_key).build()
}
pub fn api_key(&self) -> &str {
&self.inner.api_key
}
pub fn base_url(&self) -> &Url {
&self.inner.base_url
}
pub fn http(&self) -> &reqwest::Client {
&self.inner.http
}
pub fn retry(&self) -> &RetryConfig {
&self.inner.retry
}
pub fn stream_reconnects(&self) -> u32 {
self.inner.stream_reconnects
}
pub fn app_name(&self) -> Option<&str> {
self.inner.app_name.as_deref()
}
pub fn referer(&self) -> Option<&str> {
self.inner.referer.as_deref()
}
pub async fn chat_complete(
&self,
mut req: ChatCompletionRequest,
) -> Result<ChatCompletionResponse> {
req.stream = Some(false);
apply_model_suffix(&mut req.model, &mut req.provider);
request::execute_json(self, "chat/completions", &req).await
}
pub async fn complete(&self, mut req: CompletionRequest) -> Result<CompletionResponse> {
req.stream = Some(false);
apply_model_suffix(&mut req.model, &mut req.provider);
request::execute_json(self, "completions", &req).await
}
pub async fn chat_complete_stream(
&self,
mut req: ChatCompletionRequest,
) -> Result<EventStream<ChatCompletionResponse>> {
req.stream = Some(true);
apply_model_suffix(&mut req.model, &mut req.provider);
self.open_event_stream("chat/completions", &req).await
}
pub async fn complete_stream(
&self,
mut req: CompletionRequest,
) -> Result<EventStream<CompletionResponse>> {
req.stream = Some(true);
apply_model_suffix(&mut req.model, &mut req.provider);
self.open_event_stream("completions", &req).await
}
pub async fn list_models(&self, opts: Option<&ListModelsOptions>) -> Result<ModelsResponse> {
let query = opts.map(ListModelsOptions::to_query).unwrap_or_default();
request::execute_json_get(self, "models", &query).await
}
pub async fn list_model_endpoints(
&self,
author: &str,
slug: &str,
) -> Result<ModelEndpointsResponse> {
if author.is_empty() {
return Err(Error::InvalidInput("author cannot be empty"));
}
if slug.is_empty() {
return Err(Error::InvalidInput("slug cannot be empty"));
}
let path = format!(
"models/{}/{}/endpoints",
percent_encode_segment(author),
percent_encode_segment(slug),
);
request::execute_json_get(self, &path, &[]).await
}
pub async fn list_providers(&self) -> Result<ProvidersResponse> {
request::execute_json_get(self, "providers", &[]).await
}
pub async fn get_credits(&self) -> Result<CreditsResponse> {
request::execute_json_get(self, "credits", &[]).await
}
pub async fn get_key(&self) -> Result<KeyResponse> {
request::execute_json_get(self, "key", &[]).await
}
pub async fn get_activity(&self, opts: Option<&ActivityOptions>) -> Result<ActivityResponse> {
let query = opts.map(ActivityOptions::to_query).unwrap_or_default();
request::execute_json_get(self, "activity", &query).await
}
pub async fn list_keys(&self, opts: Option<&ListKeysOptions>) -> Result<ListKeysResponse> {
let query = opts
.copied()
.map(ListKeysOptions::to_query)
.unwrap_or_default();
request::execute_json_get(self, "keys", &query).await
}
pub async fn get_key_by_hash(&self, hash: &str) -> Result<GetKeyByHashResponse> {
if hash.is_empty() {
return Err(Error::InvalidInput("hash cannot be empty"));
}
let path = format!("keys/{}", percent_encode_segment(hash));
request::execute_json_get(self, &path, &[]).await
}
pub async fn create_key(&self, req: &CreateKeyRequest) -> Result<CreateKeyResponse> {
if req.name.is_empty() {
return Err(Error::InvalidInput("name is required"));
}
request::execute_json(self, "keys", req).await
}
pub async fn update_key(
&self,
hash: &str,
req: &UpdateKeyRequest,
) -> Result<UpdateKeyResponse> {
if hash.is_empty() {
return Err(Error::InvalidInput("hash cannot be empty"));
}
let path = format!("keys/{}", percent_encode_segment(hash));
request::execute_json_method(self, reqwest::Method::PATCH, &path, Some(req)).await
}
pub async fn delete_key(&self, hash: &str) -> Result<DeleteKeyResponse> {
if hash.is_empty() {
return Err(Error::InvalidInput("hash cannot be empty"));
}
let path = format!("keys/{}", percent_encode_segment(hash));
request::execute_json_method::<(), _>(self, reqwest::Method::DELETE, &path, None).await
}
pub async fn list_guardrails(
&self,
opts: Option<&ListGuardrailsOptions>,
) -> Result<ListGuardrailsResponse> {
let query = opts
.copied()
.map(ListGuardrailsOptions::to_query)
.unwrap_or_default();
request::execute_json_get(self, "guardrails", &query).await
}
pub async fn create_guardrail(&self, req: &CreateGuardrailRequest) -> Result<Guardrail> {
if req.name.is_empty() {
return Err(Error::InvalidInput("name is required"));
}
request::execute_json(self, "guardrails", req).await
}
pub async fn get_guardrail(&self, id: &str) -> Result<Guardrail> {
if id.is_empty() {
return Err(Error::InvalidInput("id cannot be empty"));
}
let path = format!("guardrails/{}", percent_encode_segment(id));
request::execute_json_get(self, &path, &[]).await
}
pub async fn update_guardrail(
&self,
id: &str,
req: &UpdateGuardrailRequest,
) -> Result<Guardrail> {
if id.is_empty() {
return Err(Error::InvalidInput("id cannot be empty"));
}
let path = format!("guardrails/{}", percent_encode_segment(id));
request::execute_json_method(self, reqwest::Method::PATCH, &path, Some(req)).await
}
pub async fn delete_guardrail(&self, id: &str) -> Result<DeleteGuardrailResponse> {
if id.is_empty() {
return Err(Error::InvalidInput("id cannot be empty"));
}
let path = format!("guardrails/{}", percent_encode_segment(id));
request::execute_json_method::<(), _>(self, reqwest::Method::DELETE, &path, None).await
}
pub async fn list_all_guardrail_key_assignments(
&self,
opts: Option<&ListGuardrailsOptions>,
) -> Result<ListGuardrailKeyAssignmentsResponse> {
let query = opts
.copied()
.map(ListGuardrailsOptions::to_query)
.unwrap_or_default();
request::execute_json_get(self, "guardrails/key-assignments", &query).await
}
pub async fn list_guardrail_key_assignments(
&self,
id: &str,
opts: Option<&ListGuardrailsOptions>,
) -> Result<ListGuardrailKeyAssignmentsResponse> {
if id.is_empty() {
return Err(Error::InvalidInput("id cannot be empty"));
}
let path = format!("guardrails/{}/key-assignments", percent_encode_segment(id));
let query = opts
.copied()
.map(ListGuardrailsOptions::to_query)
.unwrap_or_default();
request::execute_json_get(self, &path, &query).await
}
pub async fn assign_keys_to_guardrail(
&self,
id: &str,
req: &AssignKeysRequest,
) -> Result<AssignKeysResponse> {
if id.is_empty() {
return Err(Error::InvalidInput("id cannot be empty"));
}
if req.key_hashes.is_empty() {
return Err(Error::InvalidInput("key_hashes cannot be empty"));
}
let path = format!("guardrails/{}/key-assignments", percent_encode_segment(id));
request::execute_json(self, &path, req).await
}
pub async fn unassign_keys_from_guardrail(
&self,
id: &str,
req: &AssignKeysRequest,
) -> Result<()> {
if id.is_empty() {
return Err(Error::InvalidInput("id cannot be empty"));
}
if req.key_hashes.is_empty() {
return Err(Error::InvalidInput("key_hashes cannot be empty"));
}
let path = format!("guardrails/{}/key-assignments", percent_encode_segment(id));
request::execute_no_content_method(self, reqwest::Method::DELETE, &path, Some(req)).await
}
pub async fn list_all_guardrail_member_assignments(
&self,
opts: Option<&ListGuardrailsOptions>,
) -> Result<ListGuardrailMemberAssignmentsResponse> {
let query = opts
.copied()
.map(ListGuardrailsOptions::to_query)
.unwrap_or_default();
request::execute_json_get(self, "guardrails/member-assignments", &query).await
}
pub async fn list_guardrail_member_assignments(
&self,
id: &str,
opts: Option<&ListGuardrailsOptions>,
) -> Result<ListGuardrailMemberAssignmentsResponse> {
if id.is_empty() {
return Err(Error::InvalidInput("id cannot be empty"));
}
let path = format!(
"guardrails/{}/member-assignments",
percent_encode_segment(id)
);
let query = opts
.copied()
.map(ListGuardrailsOptions::to_query)
.unwrap_or_default();
request::execute_json_get(self, &path, &query).await
}
pub async fn assign_members_to_guardrail(
&self,
id: &str,
req: &AssignMembersRequest,
) -> Result<AssignMembersResponse> {
if id.is_empty() {
return Err(Error::InvalidInput("id cannot be empty"));
}
if req.member_user_ids.is_empty() {
return Err(Error::InvalidInput("member_user_ids cannot be empty"));
}
let path = format!(
"guardrails/{}/member-assignments",
percent_encode_segment(id)
);
request::execute_json(self, &path, req).await
}
pub async fn unassign_members_from_guardrail(
&self,
id: &str,
req: &AssignMembersRequest,
) -> Result<()> {
if id.is_empty() {
return Err(Error::InvalidInput("id cannot be empty"));
}
if req.member_user_ids.is_empty() {
return Err(Error::InvalidInput("member_user_ids cannot be empty"));
}
let path = format!(
"guardrails/{}/member-assignments",
percent_encode_segment(id)
);
request::execute_no_content_method(self, reqwest::Method::DELETE, &path, Some(req)).await
}
pub async fn create_video(
&self,
req: &VideoGenerationRequest,
) -> Result<VideoGenerationResponse> {
if req.model.is_empty() {
return Err(Error::InvalidInput("model is required"));
}
if req.prompt.is_empty() {
return Err(Error::InvalidInput("prompt is required"));
}
request::execute_json(self, "videos", req).await
}
pub async fn get_video(&self, job_id: &str) -> Result<VideoGenerationResponse> {
if job_id.is_empty() {
return Err(Error::InvalidInput("job_id cannot be empty"));
}
let path = format!("videos/{}", percent_encode_segment(job_id));
request::execute_json_get(self, &path, &[]).await
}
pub async fn get_video_content(
&self,
job_id: &str,
index: u32,
) -> Result<VideoContentResponse> {
if job_id.is_empty() {
return Err(Error::InvalidInput("job_id cannot be empty"));
}
let path = format!("videos/{}/content", percent_encode_segment(job_id));
let query: Vec<(&'static str, String)> = if index > 0 {
vec![("index", index.to_string())]
} else {
Vec::new()
};
let (content, content_type) = request::execute_bytes_get(self, &path, &query).await?;
Ok(VideoContentResponse {
content,
content_type,
})
}
pub async fn list_video_models(&self) -> Result<VideoModelsResponse> {
request::execute_json_get(self, "videos/models", &[]).await
}
pub async fn wait_for_video(
&self,
job_id: &str,
interval: Duration,
) -> Result<VideoGenerationResponse> {
loop {
let resp = self.get_video(job_id).await?;
if resp.status.is_terminal() {
return Ok(resp);
}
tokio::time::sleep(interval).await;
}
}
pub async fn create_speech(&self, req: &SpeechRequest) -> Result<SpeechResponse> {
if req.input.is_empty() {
return Err(Error::InvalidInput("input is required"));
}
if req.model.is_empty() {
return Err(Error::InvalidInput("model is required"));
}
if req.voice.is_empty() {
return Err(Error::InvalidInput("voice is required"));
}
let (audio, content_type) = request::execute_bytes_post(self, "audio/speech", req).await?;
let format = req.response_format.unwrap_or(SpeechFormat::Pcm);
Ok(SpeechResponse {
audio,
content_type,
format,
})
}
pub async fn rerank(&self, req: &RerankRequest) -> Result<RerankResponse> {
if req.model.is_empty() {
return Err(Error::InvalidInput("model is required"));
}
if req.query.is_empty() {
return Err(Error::InvalidInput("query is required"));
}
if req.documents.is_empty() {
return Err(Error::InvalidInput("documents must not be empty"));
}
request::execute_json(self, "rerank", req).await
}
pub async fn list_zdr_endpoints(&self) -> Result<ZdrEndpointsResponse> {
request::execute_json_get(self, "endpoints/zdr", &[]).await
}
pub async fn list_organization_members(
&self,
opts: Option<&ListOrganizationMembersOptions>,
) -> Result<ListOrganizationMembersResponse> {
let query = opts
.copied()
.map(ListOrganizationMembersOptions::to_query)
.unwrap_or_default();
request::execute_json_get(self, "organization/members", &query).await
}
pub async fn list_workspaces(
&self,
opts: Option<&ListWorkspacesOptions>,
) -> Result<ListWorkspacesResponse> {
let query = opts
.copied()
.map(ListWorkspacesOptions::to_query)
.unwrap_or_default();
request::execute_json_get(self, "workspaces", &query).await
}
pub async fn create_workspace(
&self,
req: &CreateWorkspaceRequest,
) -> Result<CreateWorkspaceResponse> {
if req.name.is_empty() {
return Err(Error::InvalidInput("name is required"));
}
if req.slug.is_empty() {
return Err(Error::InvalidInput("slug is required"));
}
request::execute_json(self, "workspaces", req).await
}
pub async fn get_workspace(&self, id_or_slug: &str) -> Result<GetWorkspaceResponse> {
if id_or_slug.is_empty() {
return Err(Error::InvalidInput("id_or_slug cannot be empty"));
}
let path = format!("workspaces/{}", percent_encode_segment(id_or_slug));
request::execute_json_get(self, &path, &[]).await
}
pub async fn update_workspace(
&self,
id_or_slug: &str,
req: &UpdateWorkspaceRequest,
) -> Result<UpdateWorkspaceResponse> {
if id_or_slug.is_empty() {
return Err(Error::InvalidInput("id_or_slug cannot be empty"));
}
let path = format!("workspaces/{}", percent_encode_segment(id_or_slug));
request::execute_json_method(self, reqwest::Method::PATCH, &path, Some(req)).await
}
pub async fn delete_workspace(&self, id_or_slug: &str) -> Result<DeleteWorkspaceResponse> {
if id_or_slug.is_empty() {
return Err(Error::InvalidInput("id_or_slug cannot be empty"));
}
let path = format!("workspaces/{}", percent_encode_segment(id_or_slug));
request::execute_json_method::<(), _>(self, reqwest::Method::DELETE, &path, None).await
}
pub async fn add_workspace_members(
&self,
id_or_slug: &str,
user_ids: &[String],
) -> Result<BulkAddWorkspaceMembersResponse> {
if id_or_slug.is_empty() {
return Err(Error::InvalidInput("id_or_slug cannot be empty"));
}
if user_ids.is_empty() {
return Err(Error::InvalidInput("user_ids cannot be empty"));
}
let path = format!(
"workspaces/{}/members/add",
percent_encode_segment(id_or_slug)
);
let body = BulkWorkspaceMembersRequest { user_ids };
request::execute_json(self, &path, &body).await
}
pub async fn remove_workspace_members(
&self,
id_or_slug: &str,
user_ids: &[String],
) -> Result<BulkRemoveWorkspaceMembersResponse> {
if id_or_slug.is_empty() {
return Err(Error::InvalidInput("id_or_slug cannot be empty"));
}
if user_ids.is_empty() {
return Err(Error::InvalidInput("user_ids cannot be empty"));
}
let path = format!(
"workspaces/{}/members/remove",
percent_encode_segment(id_or_slug)
);
let body = BulkWorkspaceMembersRequest { user_ids };
request::execute_json(self, &path, &body).await
}
pub(crate) async fn open_event_stream<Req, Resp>(
&self,
path: &'static str,
req: &Req,
) -> Result<EventStream<Resp>>
where
Req: serde::Serialize + ?Sized,
Resp: serde::de::DeserializeOwned,
{
let body_bytes = serde_json::to_vec(req)?;
let initial = request::open_stream_bytes(self, path, body_bytes.clone()).await?;
let client = self.clone();
#[cfg(not(target_arch = "wasm32"))]
let reopen: crate::stream::Reopen = Arc::new(move || {
let client = client.clone();
let body_bytes = body_bytes.clone();
async move { request::open_stream_bytes(&client, path, body_bytes).await }.boxed()
});
#[cfg(target_arch = "wasm32")]
let reopen: crate::stream::Reopen = Rc::new(move || {
let client = client.clone();
let body_bytes = body_bytes.clone();
async move { request::open_stream_bytes(&client, path, body_bytes).await }.boxed_local()
});
Ok(EventStream::new(
initial,
reopen,
self.inner.stream_reconnects,
))
}
}
#[derive(Debug, Default)]
pub struct ClientBuilder {
api_key: Option<String>,
base_url: Option<Url>,
http_client: Option<reqwest::Client>,
timeout: Option<Duration>,
retry: Option<RetryConfig>,
stream_reconnects: Option<u32>,
app_name: Option<String>,
referer: Option<String>,
}
impl ClientBuilder {
pub fn api_key(mut self, key: impl Into<String>) -> Self {
self.api_key = Some(key.into());
self
}
pub fn base_url(mut self, url: impl AsRef<str>) -> Result<Self> {
let mut parsed = Url::parse(url.as_ref())
.map_err(|_| Error::InvalidInput("base_url is not a valid URL"))?;
if !parsed.path().ends_with('/') {
let new_path = format!("{}/", parsed.path());
parsed.set_path(&new_path);
}
self.base_url = Some(parsed);
Ok(self)
}
pub fn http_client(mut self, client: reqwest::Client) -> Self {
self.http_client = Some(client);
self
}
pub fn timeout(mut self, d: Duration) -> Self {
self.timeout = Some(d);
self
}
pub fn retry(mut self, max: u32, base_delay: Duration) -> Self {
let cfg = RetryConfig {
max_retries: max,
initial_delay: base_delay,
..RetryConfig::default()
};
self.retry = Some(cfg);
self
}
pub fn retry_config(mut self, cfg: RetryConfig) -> Self {
self.retry = Some(cfg);
self
}
pub fn stream_reconnects(mut self, max: u32) -> Self {
self.stream_reconnects = Some(max);
self
}
pub fn app_name(mut self, name: impl Into<String>) -> Self {
self.app_name = Some(name.into());
self
}
pub fn referer(mut self, referer: impl Into<String>) -> Self {
self.referer = Some(referer.into());
self
}
pub fn build(self) -> Result<Client> {
let api_key = self.api_key.ok_or(Error::MissingField("api_key"))?;
if api_key.is_empty() {
return Err(Error::InvalidInput("api_key must not be empty"));
}
let base_url = match self.base_url {
Some(u) => u,
None => Url::parse(DEFAULT_BASE_URL).expect("DEFAULT_BASE_URL is a valid URL"),
};
let http = match self.http_client {
Some(c) => c,
None => {
#[cfg(not(target_arch = "wasm32"))]
let mut b = reqwest::Client::builder();
#[cfg(target_arch = "wasm32")]
let b = reqwest::Client::builder();
#[cfg(not(target_arch = "wasm32"))]
if let Some(t) = self.timeout {
b = b.timeout(t);
}
#[cfg(target_arch = "wasm32")]
let _ = self.timeout;
b.build().map_err(Error::Http)?
}
};
let retry = self.retry.unwrap_or_default();
let stream_reconnects = self.stream_reconnects.unwrap_or(DEFAULT_STREAM_RECONNECTS);
Ok(Client {
inner: Arc::new(ClientInner {
api_key,
base_url,
http,
retry,
stream_reconnects,
app_name: self.app_name,
referer: self.referer,
}),
})
}
}
pub(crate) fn percent_encode_segment(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for &b in s.as_bytes() {
let unreserved = b.is_ascii_alphanumeric() || matches!(b, b'-' | b'.' | b'_' | b'~');
if unreserved {
out.push(b as char);
} else {
out.push('%');
out.push_str(&format!("{b:02X}"));
}
}
out
}
pub(crate) fn apply_model_suffix(model: &mut String, provider: &mut Option<Provider>) {
let sort = if let Some(stripped) = model.strip_suffix(":nitro") {
let new_model = stripped.to_string();
*model = new_model;
"throughput"
} else if let Some(stripped) = model.strip_suffix(":floor") {
let new_model = stripped.to_string();
*model = new_model;
"price"
} else {
return;
};
let p = provider.get_or_insert_with(Provider::default);
if p.sort.is_none() {
p.sort = Some(sort.to_string());
}
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_send_sync<T: Send + Sync>() {}
#[test]
fn client_is_send_sync() {
assert_send_sync::<Client>();
}
#[test]
fn builder_happy_path() {
let c = Client::builder()
.api_key("sk-test")
.app_name("demo")
.referer("https://demo.example")
.timeout(Duration::from_secs(10))
.build()
.unwrap();
assert_eq!(c.api_key(), "sk-test");
assert_eq!(c.app_name(), Some("demo"));
assert_eq!(c.referer(), Some("https://demo.example"));
assert_eq!(c.base_url().as_str(), DEFAULT_BASE_URL);
assert_eq!(c.stream_reconnects(), DEFAULT_STREAM_RECONNECTS);
}
#[test]
fn stream_reconnects_can_be_disabled() {
let c = Client::builder()
.api_key("sk-test")
.stream_reconnects(0)
.build()
.unwrap();
assert_eq!(c.stream_reconnects(), 0);
}
#[test]
fn missing_api_key_errors() {
let err = Client::builder().build().unwrap_err();
assert!(matches!(err, Error::MissingField("api_key")));
}
#[test]
fn empty_api_key_errors() {
let err = Client::builder().api_key("").build().unwrap_err();
assert!(matches!(err, Error::InvalidInput(_)));
}
#[test]
fn invalid_base_url_errors() {
let err = Client::builder().base_url("not a url").unwrap_err();
assert!(matches!(err, Error::InvalidInput(_)));
}
#[test]
fn base_url_path_gains_trailing_slash() {
let c = Client::builder()
.api_key("k")
.base_url("https://example.com/v2")
.unwrap()
.build()
.unwrap();
assert!(c.base_url().as_str().ends_with('/'));
}
#[test]
fn clone_shares_inner() {
let c1 = Client::new("k").unwrap();
let c2 = c1.clone();
assert!(Arc::ptr_eq(&c1.inner, &c2.inner));
}
#[test]
fn retry_helper_sets_fields() {
let c = Client::builder()
.api_key("k")
.retry(7, Duration::from_millis(250))
.build()
.unwrap();
assert_eq!(c.retry().max_retries, 7);
assert_eq!(c.retry().initial_delay, Duration::from_millis(250));
}
#[test]
fn nitro_suffix_maps_to_throughput_sort() {
let mut m = "openai/gpt-4o:nitro".to_string();
let mut p = None;
apply_model_suffix(&mut m, &mut p);
assert_eq!(m, "openai/gpt-4o");
assert_eq!(p.unwrap().sort.as_deref(), Some("throughput"));
}
#[test]
fn floor_suffix_maps_to_price_sort() {
let mut m = "anthropic/claude-3:floor".to_string();
let mut p = None;
apply_model_suffix(&mut m, &mut p);
assert_eq!(m, "anthropic/claude-3");
assert_eq!(p.unwrap().sort.as_deref(), Some("price"));
}
#[test]
fn caller_set_sort_wins_over_suffix() {
let mut m = "openai/gpt-4o:nitro".to_string();
let mut p = Some(Provider {
sort: Some("latency".to_string()),
..Provider::default()
});
apply_model_suffix(&mut m, &mut p);
assert_eq!(m, "openai/gpt-4o");
assert_eq!(p.unwrap().sort.as_deref(), Some("latency"));
}
#[test]
fn unknown_suffix_passes_through() {
let mut m = "openai/gpt-4o:exotic".to_string();
let mut p = None;
apply_model_suffix(&mut m, &mut p);
assert_eq!(m, "openai/gpt-4o:exotic");
assert!(p.is_none());
}
}