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:
@@ -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(¤t_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"]
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user