Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion crates/braintrust-llm-router/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -83,7 +83,7 @@ async fn main() -> anyhow::Result<()> {
| AWS Bedrock | AWS SigV4 | Yes | |
| Azure OpenAI | API Key / Entra | Yes | |
| Mistral | API Key | Yes | |
| OpenAI-compatible | API Key | Yes | Together, Groq, Perplexity, etc. |
| OpenAI-compatible | API Key | Yes | Together, Groq, Perplexity, SaladCloud, etc. |

## API reference

Expand Down
14 changes: 14 additions & 0 deletions crates/braintrust-llm-router/src/providers/openai.rs
Original file line number Diff line number Diff line change
Expand Up @@ -354,6 +354,7 @@ pub fn is_openai_compatible(kind: &str) -> bool {
| "databricks"
| "cohere"
| "openrouter"
| "saladcloud"
Comment thread
mgorkii-nlplogix marked this conversation as resolved.
)
}

Expand Down Expand Up @@ -411,6 +412,10 @@ pub fn openai_compatible_endpoint(kind: &str) -> Option<OpenAICompatibleEndpoint
url: "https://openrouter.ai/api/v1",
is_template: false,
}),
"saladcloud" => Some(OpenAICompatibleEndpoint {
url: "https://ai.salad.cloud/v1",
is_template: false,
}),
_ => None,
}
}
Expand Down Expand Up @@ -460,4 +465,13 @@ mod tests {
let url = provider.chat_url(None).expect("url");
assert!(url.as_str().contains("default.lepton.run"));
}

#[test]
fn recognizes_saladcloud_endpoint() {
assert!(is_openai_compatible("saladcloud"));

let endpoint = openai_compatible_endpoint("saladcloud").expect("endpoint");
assert_eq!(endpoint.url, "https://ai.salad.cloud/v1");
assert!(!endpoint.is_template);
}
}
43 changes: 29 additions & 14 deletions crates/braintrust-llm-router/src/router.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,9 +38,9 @@ pub struct CompleteResponseWithRaw {
}

use crate::providers::{
is_openai_compatible, AnthropicProvider, AzureAiGatewayProvider, AzureProvider,
BedrockProvider, DatabricksProvider, GoogleProvider, MistralProvider, OpenAIProvider,
VertexProvider,
is_openai_compatible, openai_compatible_endpoint, AnthropicProvider, AzureAiGatewayProvider,
AzureProvider, BedrockProvider, DatabricksProvider, GoogleProvider, MistralProvider,
OpenAIProvider, VertexProvider,
};

/// Create a provider instance from configuration parameters.
Expand All @@ -51,7 +51,7 @@ use crate::providers::{
/// # Arguments
///
/// * `kind` - Provider type: "openai", "anthropic", "azure", "google", "vertex", "bedrock", "mistral", or OpenAI-compatible
/// * `endpoint` - Custom endpoint URL (optional)
/// * `endpoint` - Custom endpoint URL (optional; uses the provider's registered default when omitted)
/// * `endpoint_template` - Endpoint template with `<model>` placeholder (optional, OpenAI only)
/// * `timeout` - Request timeout (optional)
/// * `metadata` - Provider-specific options (organization_id, project, api_version, etc.)
Expand Down Expand Up @@ -131,16 +131,31 @@ pub fn create_provider(
timeout,
client_settings,
)?)),
kind if is_openai_compatible(kind) => Ok(Arc::new(
OpenAIProvider::from_config(
endpoint,
endpoint_template,
timeout,
metadata,
client_settings,
)?
.with_provider_alias(kind.to_ascii_lowercase()),
)),
kind if is_openai_compatible(kind) => {
let mut default_endpoint = None;
let mut endpoint_template = endpoint_template;
if endpoint.is_none() && endpoint_template.is_none() {
if let Some(default) = openai_compatible_endpoint(kind) {
if default.is_template {
endpoint_template = Some(default.url);
} else {
default_endpoint = Some(Url::parse(default.url).map_err(|e| {
Error::InvalidRequest(format!("invalid {kind} default endpoint: {e}"))
})?);
}
}
}
Ok(Arc::new(
OpenAIProvider::from_config(
endpoint.or(default_endpoint.as_ref()),
endpoint_template,
timeout,
metadata,
client_settings,
)?
.with_provider_alias(kind.to_ascii_lowercase()),
))
}
other => Err(Error::InvalidRequest(format!(
"unsupported provider kind: {other}"
))),
Expand Down
157 changes: 157 additions & 0 deletions crates/braintrust-llm-router/tests/provider_endpoints.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,157 @@
use std::collections::HashMap;
use std::sync::{Arc, Mutex};

use async_trait::async_trait;
use braintrust_llm_router::{
api_key_auth, clear_override_client, create_provider, set_override_client, ClientHeaders,
ModelSpec, ProviderFormat,
};
use bytes::Bytes;
use futures::StreamExt;
use reqwest::{Request, Response, Url};
use reqwest_middleware::{ClientBuilder, Middleware, Next};

struct CaptureRequests(Arc<Mutex<Vec<Request>>>);

#[async_trait]
impl Middleware for CaptureRequests {
async fn handle(
&self,
request: Request,
_extensions: &mut http::Extensions,
_next: Next<'_>,
) -> reqwest_middleware::Result<Response> {
self.0.lock().unwrap().push(request);
Ok(http::Response::builder()
.status(200)
.header("content-type", "text/event-stream")
.body("data: [DONE]\n\n")
.unwrap()
.into())
}
}

struct OverrideGuard;

impl Drop for OverrideGuard {
fn drop(&mut self) {
clear_override_client();
}
}

// The middleware captures real outgoing requests without opening a network connection.
async fn assert_endpoint(
kind: &str,
endpoint: Option<&str>,
template: Option<&str>,
expected_url: &str,
) {
let requests = Arc::new(Mutex::new(Vec::new()));
let client = ClientBuilder::new(reqwest::Client::new())
.with(CaptureRequests(Arc::clone(&requests)))
.build();
set_override_client(client);
let _guard = OverrideGuard;
let endpoint = endpoint.map(|url| Url::parse(url).unwrap());
let provider = create_provider(
kind,
endpoint.as_ref(),
template,
None,
&HashMap::new(),
None,
)
.unwrap();
assert!(provider.matches_provider_alias(&kind.to_ascii_lowercase()));
let spec: ModelSpec = serde_json::from_str(
r#"{"model":"qwen3.6-35b-a3b","format":"chat_completions","flavor":"chat"}"#,
)
.unwrap();
let auth = api_key_auth("test-key");
let headers = ClientHeaders::default();
let body = Bytes::from_static(br#"{"model":"qwen3.6-35b-a3b","messages":[]}"#);
provider
.complete(
body.clone(),
&auth,
&spec,
ProviderFormat::ChatCompletions,
&headers,
)
.await
.unwrap();
let mut stream = provider
.complete_stream(
body,
&auth,
&spec,
ProviderFormat::ChatCompletions,
&headers,
)
.await
.unwrap();
while let Some(chunk) = stream.next().await {
chunk.unwrap();
}

let requests = requests.lock().unwrap();
assert_eq!(requests.len(), 2, "unary and streaming requests");
for request in requests.iter() {
assert_eq!(request.url().as_str(), expected_url, "{kind}");
assert_eq!(request.method(), reqwest::Method::POST);
assert_eq!(request.headers()["authorization"], "Bearer test-key");
}
}

#[tokio::test]
#[serial_test::serial]
async fn saladcloud_factory_uses_registered_endpoint() {
assert_endpoint(
"saladcloud",
None,
None,
"https://ai.salad.cloud/v1/chat/completions",
)
.await;
}

#[tokio::test]
#[serial_test::serial]
async fn saladcloud_factory_preserves_explicit_endpoint() {
assert_endpoint(
"saladcloud",
Some("https://custom.example/v1/"),
None,
"https://custom.example/v1/chat/completions",
)
.await;
}

#[tokio::test]
#[serial_test::serial]
async fn saladcloud_factory_preserves_template_precedence() {
for endpoint in [None, Some("https://custom.example/v1")] {
assert_endpoint(
"saladcloud",
endpoint,
Some("https://<model>.example/v1/"),
"https://qwen3.6-35b-a3b.example/v1/chat/completions",
)
.await;
}
}

#[tokio::test]
#[serial_test::serial]
async fn factory_respects_existing_provider_defaults() {
for (kind, expected_url) in [
("openai", "https://api.openai.com/v1/chat/completions"),
("groq", "https://api.groq.com/openai/v1/chat/completions"),
(
"lepton",
"https://qwen3.6-35b-a3b.lepton.run/api/v1/chat/completions",
),
] {
assert_endpoint(kind, None, None, expected_url).await;
}
}