check for responses API support in databricks (#9347)

Signed-off-by: Douwe M Osinga <douwe@sidewalklabs.com>
Co-authored-by: Douwe M Osinga <douwe@sidewalklabs.com>
This commit is contained in:
Jack Amadeo
2026-06-17 14:00:46 -04:00
committed by GitHub
parent 177cf3d4b8
commit f0cc765440
+192 -58
View File
@@ -44,6 +44,7 @@ struct DatabricksEndpointInfo {
upstream_model_name: Option<String>,
upstream_model_provider: Option<String>,
reasoning: Option<bool>,
supports_responses_api: bool,
}
#[derive(Debug, Clone)]
@@ -234,16 +235,24 @@ impl DatabricksProvider {
}
}
fn is_responses_model(model_name: &str) -> bool {
is_openai_responses_model(model_name)
}
fn is_claude_model(model_name: &str) -> bool {
model_name.to_lowercase().contains("claude")
}
fn is_reasoning_capable_model_name(model_name: &str) -> bool {
Self::is_claude_model(model_name) || Self::is_responses_model(model_name)
Self::is_claude_model(model_name) || is_openai_responses_model(model_name)
}
fn uses_responses_api(
endpoint_info: Option<&DatabricksEndpointInfo>,
model_names: &[&str],
) -> bool {
match endpoint_info {
Some(info) => info.supports_responses_api,
None => model_names
.iter()
.any(|name| is_openai_responses_model(name)),
}
}
fn endpoint_model_candidates(value: &Value) -> Vec<DatabricksUpstreamModel> {
@@ -304,6 +313,7 @@ impl DatabricksProvider {
fn endpoint_info_from_value(endpoint: &Value) -> Option<DatabricksEndpointInfo> {
let name = endpoint.get("name")?.as_str()?.to_string();
let supports_responses_api = Self::endpoint_supports_responses_api(endpoint);
let upstream_model = Self::endpoint_model_candidates(endpoint)
.into_iter()
.find(|candidate| candidate.name != name);
@@ -320,9 +330,45 @@ impl DatabricksProvider {
upstream_model_name,
upstream_model_provider,
reasoning,
supports_responses_api,
})
}
fn endpoint_supports_responses_api(endpoint: &Value) -> bool {
fn value_contains_responses_api(value: &Value) -> bool {
match value {
Value::Object(map) => {
map.get("api_types")
.and_then(|api_types| api_types.as_array())
.is_some_and(|api_types| {
api_types
.iter()
.any(|api_type| api_type.as_str() == Some("openai/v1/responses"))
})
|| map.values().any(value_contains_responses_api)
}
Value::Array(values) => values.iter().any(value_contains_responses_api),
_ => false,
}
}
let Some(config) = endpoint.get("config") else {
return false;
};
for collection_key in ["served_entities", "served_models"] {
let Some(entities) = config.get(collection_key).and_then(|v| v.as_array()) else {
continue;
};
if entities.iter().any(value_contains_responses_api) {
return true;
}
}
false
}
async fn fetch_endpoint_info(
&self,
endpoint_name: &str,
@@ -378,6 +424,7 @@ impl DatabricksProvider {
let mut current_endpoint_name = endpoint_name.to_string();
let mut visited = HashSet::new();
let mut last_info: Option<DatabricksEndpointInfo> = None;
let mut first_hop_supports_responses_api: Option<bool> = None;
for _ in 0..MAX_MODEL_SERVING_HOPS {
if !visited.insert(current_endpoint_name.clone()) {
@@ -385,6 +432,8 @@ impl DatabricksProvider {
}
let info = self.fetch_endpoint_info(&current_endpoint_name).await?;
let supports_responses_api =
*first_hop_supports_responses_api.get_or_insert(info.supports_responses_api);
let next_endpoint_name = match (
info.upstream_model_provider.as_deref(),
info.upstream_model_name.as_deref(),
@@ -403,7 +452,7 @@ impl DatabricksProvider {
continue;
}
return Ok(if info.name == original_endpoint_name {
let mut resolved_info = if info.name == original_endpoint_name {
info
} else {
let upstream_model_name = info
@@ -415,8 +464,11 @@ impl DatabricksProvider {
upstream_model_name,
upstream_model_provider: info.upstream_model_provider.clone(),
reasoning: info.reasoning,
supports_responses_api,
}
});
};
resolved_info.supports_responses_api = supports_responses_api;
return Ok(resolved_info);
}
last_info
@@ -425,6 +477,7 @@ impl DatabricksProvider {
upstream_model_name: info.upstream_model_name,
upstream_model_provider: info.upstream_model_provider,
reasoning: info.reasoning,
supports_responses_api: first_hop_supports_responses_api.unwrap_or(false),
})
.ok_or_else(|| {
ProviderError::RequestFailed(
@@ -484,16 +537,19 @@ impl DatabricksProvider {
}
}
fn get_endpoint_path(&self, model_name: &str, is_embedding: bool) -> String {
fn get_endpoint_path(
&self,
model_name: &str,
is_embedding: bool,
is_responses_model: bool,
) -> String {
if is_embedding {
"serving-endpoints/text-embedding-3-small/invocations".to_string()
} else if is_responses_model {
"serving-endpoints/responses".to_string()
} else {
let (clean_name, _) = extract_reasoning_effort(model_name);
if Self::is_responses_model(&clean_name) {
"serving-endpoints/responses".to_string()
} else {
format!("serving-endpoints/{}/invocations", clean_name)
}
format!("serving-endpoints/{}/invocations", clean_name)
}
}
@@ -514,7 +570,10 @@ impl DatabricksProvider {
) -> Result<Value, ProviderError> {
let is_embedding = payload.get("input").is_some() && payload.get("messages").is_none();
let model_to_use = model_name.unwrap_or(&self.model.model_name);
let path = self.get_endpoint_path(model_to_use, is_embedding);
let (endpoint_name, _) = extract_reasoning_effort(model_to_use);
let endpoint_info = self.resolve_endpoint_info_cached(&endpoint_name).await.ok();
let is_responses_model = Self::uses_responses_api(endpoint_info.as_ref(), &[model_to_use]);
let path = self.get_endpoint_path(model_to_use, is_embedding, is_responses_model);
if let Some(session_id) = session_id {
if let Some(client_request_id) = self.build_client_request_id(session_id) {
@@ -595,12 +654,14 @@ impl Provider for DatabricksProvider {
.as_ref()
.and_then(|info| info.upstream_model_name.as_deref())
.unwrap_or(&model_config.model_name);
let is_responses_model = Self::is_responses_model(&model_config.model_name)
|| Self::is_responses_model(effective_model_name);
let is_responses_model = Self::uses_responses_api(
endpoint_info.as_ref(),
&[&model_config.model_name, effective_model_name],
);
let path = if is_responses_model {
"serving-endpoints/responses".to_string()
} else {
self.get_endpoint_path(&model_config.model_name, false)
self.get_endpoint_path(&model_config.model_name, false, is_responses_model)
};
let client_request_id = self.build_client_request_id(session_id);
@@ -880,47 +941,6 @@ impl EmbeddingCapable for DatabricksProvider {
mod tests {
use super::*;
fn test_provider() -> DatabricksProvider {
DatabricksProvider {
api_client: super::super::api_client::ApiClient::new(
"https://example.com".to_string(),
super::super::api_client::AuthMethod::NoAuth,
)
.unwrap(),
host: "https://example.com".to_string(),
auth: DatabricksAuth::Token("fake".into()),
model: ModelConfig::new_or_fail("databricks-gpt-5.4"),
image_format: ImageFormat::OpenAi,
retry_config: RetryConfig::default(),
fast_retry_config: RetryConfig::new(0, 0, 1.0, 0),
name: "databricks".into(),
token_cache: std::sync::Arc::new(std::sync::Mutex::new(None)),
instance_id: None,
}
}
#[test]
fn responses_models_route_to_responses_endpoint() {
let provider = test_provider();
for (model_name, expected_path) in [
("gpt-5.4", "serving-endpoints/responses"),
("databricks-gpt-5.4-high", "serving-endpoints/responses"),
("databricks-gpt-5-4-xhigh", "serving-endpoints/responses"),
("o3-mini", "serving-endpoints/responses"),
(
"databricks-claude-sonnet-4",
"serving-endpoints/databricks-claude-sonnet-4/invocations",
),
] {
assert_eq!(
provider.get_endpoint_path(model_name, false),
expected_path,
"unexpected endpoint for {model_name}"
);
}
}
#[test]
fn endpoint_metadata_marks_reasoning_alias_from_external_model() {
let endpoint = json!({
@@ -942,6 +962,7 @@ mod tests {
assert_eq!(info.name, "goose");
assert_eq!(info.upstream_model_name.as_deref(), Some("claude-opus-4.6"));
assert_eq!(info.reasoning, Some(true));
assert!(!info.supports_responses_api);
let model_info = DatabricksProvider::model_info_from_endpoint(info);
assert_eq!(model_info.name, "goose");
@@ -1014,5 +1035,118 @@ mod tests {
assert_eq!(info.name, "goose-gpt-5-5");
assert_eq!(info.upstream_model_name, None);
assert_eq!(info.reasoning, Some(true));
assert!(!info.supports_responses_api);
}
#[test]
fn endpoint_metadata_detects_responses_api_from_foundation_model_api_types() {
let endpoint = json!({
"name": "databricks-gpt-5-4",
"config": {
"served_entities": [{
"name": "databricks-gpt-5-4",
"entity_name": "system.ai.databricks-gpt-5-4",
"type": "FOUNDATION_MODEL",
"foundation_model": {
"name": "system.ai.databricks-gpt-5-4",
"display_name": "GPT-5.4",
"api_types": [
"mlflow/v1/chat/completions",
"openai/v1/responses",
"cursor/v1/chat/completions"
]
}
}]
}
});
let info = DatabricksProvider::endpoint_info_from_value(&endpoint).unwrap();
assert_eq!(info.name, "databricks-gpt-5-4");
assert_eq!(
info.upstream_model_name.as_deref(),
Some("system.ai.databricks-gpt-5-4")
);
assert!(info.supports_responses_api);
}
#[test]
fn endpoint_metadata_detects_responses_api_from_served_models() {
let endpoint = json!({
"name": "databricks-gpt-5-4",
"config": {
"served_models": [{
"foundation_model": {
"api_types": ["openai/v1/responses"]
}
}]
}
});
let info = DatabricksProvider::endpoint_info_from_value(&endpoint).unwrap();
assert!(info.supports_responses_api);
}
#[test]
fn endpoint_metadata_ignores_pending_config_for_responses_routing() {
let endpoint = json!({
"name": "databricks-gpt-5-4",
"config": {
"served_entities": [{
"foundation_model": {
"api_types": ["mlflow/v1/chat/completions"]
}
}]
},
"pending_config": {
"served_entities": [{
"foundation_model": {
"api_types": ["openai/v1/responses"]
}
}]
}
});
let info = DatabricksProvider::endpoint_info_from_value(&endpoint).unwrap();
assert!(!info.supports_responses_api);
}
#[test]
fn responses_routing_prefers_metadata_over_model_name() {
let responses_info = DatabricksEndpointInfo {
name: "custom".into(),
upstream_model_name: None,
upstream_model_provider: None,
reasoning: None,
supports_responses_api: true,
};
assert!(DatabricksProvider::uses_responses_api(
Some(&responses_info),
&["databricks-claude-sonnet-4"]
));
let chat_info = DatabricksEndpointInfo {
supports_responses_api: false,
..responses_info
};
assert!(!DatabricksProvider::uses_responses_api(
Some(&chat_info),
&["gpt-5.4"]
));
}
#[test]
fn responses_routing_falls_back_to_model_name_without_metadata() {
assert!(DatabricksProvider::uses_responses_api(None, &["gpt-5.4"]));
assert!(DatabricksProvider::uses_responses_api(
None,
&["databricks-claude-sonnet-4", "gpt-5.4"]
));
assert!(!DatabricksProvider::uses_responses_api(
None,
&["databricks-claude-sonnet-4"]
));
}
}