diff --git a/crates/braintrust-llm-router/README.md b/crates/braintrust-llm-router/README.md index cb1409cc..6bfaf75f 100644 --- a/crates/braintrust-llm-router/README.md +++ b/crates/braintrust-llm-router/README.md @@ -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 diff --git a/crates/braintrust-llm-router/src/providers/openai.rs b/crates/braintrust-llm-router/src/providers/openai.rs index ef79b579..2d96d086 100644 --- a/crates/braintrust-llm-router/src/providers/openai.rs +++ b/crates/braintrust-llm-router/src/providers/openai.rs @@ -354,6 +354,7 @@ pub fn is_openai_compatible(kind: &str) -> bool { | "databricks" | "cohere" | "openrouter" + | "saladcloud" ) } @@ -411,6 +412,10 @@ pub fn openai_compatible_endpoint(kind: &str) -> Option Some(OpenAICompatibleEndpoint { + url: "https://ai.salad.cloud/v1", + is_template: false, + }), _ => None, } } @@ -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); + } } diff --git a/crates/braintrust-llm-router/src/router.rs b/crates/braintrust-llm-router/src/router.rs index e96dd9c3..40f11b11 100644 --- a/crates/braintrust-llm-router/src/router.rs +++ b/crates/braintrust-llm-router/src/router.rs @@ -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. @@ -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 `` placeholder (optional, OpenAI only) /// * `timeout` - Request timeout (optional) /// * `metadata` - Provider-specific options (organization_id, project, api_version, etc.) @@ -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}" ))), diff --git a/crates/braintrust-llm-router/tests/provider_endpoints.rs b/crates/braintrust-llm-router/tests/provider_endpoints.rs new file mode 100644 index 00000000..23d1f2e4 --- /dev/null +++ b/crates/braintrust-llm-router/tests/provider_endpoints.rs @@ -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>>); + +#[async_trait] +impl Middleware for CaptureRequests { + async fn handle( + &self, + request: Request, + _extensions: &mut http::Extensions, + _next: Next<'_>, + ) -> reqwest_middleware::Result { + 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://.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; + } +}