mirror of
https://github.com/block/buzz.git
synced 2026-08-18 06:50:31 +02:00
fix: reauthenticate databricks model discovery (#4008)
## Summary - preserve Databricks catalog 401 responses as authentication failures and retry discovery exactly once after silently refreshing the rejected bearer - preserve runtime OAuth recovery: when discovery has no usable OAuth credential, `session/new` succeeds with only the trimmed configured model so the first `session/prompt` can run the existing browser PKCE flow - reject a rejected configured `DATABRICKS_TOKEN` with actionable, non-interactive guidance; static credentials cannot recover through PKCE - use the configured-model fallback for non-auth discovery failures without caching failed or fallback catalogs, so later sessions retry discovery - keep known Databricks v2 models only for authenticated empty-catalog responses and mark their provenance - resolve discovery before MCP spawn or session registration, preventing failed discovery from leaking resources or consuming session capacity - permit serialized interactive PKCE only from the explicit saved-agent model picker; passive draft discovery never opens a browser ## Runtime flow 1. OAuth discovery attempts cached credentials and silent refresh without opening a browser. 2. If no usable OAuth bearer exists, `session/new` advertises only the configured model and succeeds. 3. The first `session/prompt` uses `TokenSource::bearer()`, which may launch browser PKCE. 4. A later session retries discovery and caches only the authenticated catalog. ## Regression coverage - rejected-but-locally-fresh OAuth bearer performs one refresh and one catalog retry - OAuth mode with no cached token allows `session/new` and returns exactly the trimmed configured model - the OAuth fallback is not cached; a later authenticated session retries discovery and caches the returned catalog - rejected static tokens still reject `session/new` - failed discovery does not consume the sole session slot or spawn the supplied MCP process - Desktop interactive/passive auth intent, static-token redaction, and authenticated empty-catalog provenance ## Verification - `cargo test -p buzz-agent` - `cargo test --manifest-path desktop/src-tauri/Cargo.toml --lib commands::agent_models` - `cargo clippy --manifest-path desktop/src-tauri/Cargo.toml --all-targets -- -D warnings` - `cargo fmt --all -- --check` - `git diff --check` - full pre-push hooks ## Review Adversarial review found and drove fixes for session/MCP resource leakage, duplicate concurrent PKCE flows, sensitive error propagation, incorrect 403 reauthentication, missing discovery-level coverage, passive browser launch, and the Desktop file-size ratchet. The final follow-up preserves the existing prompt-time OAuth flow while retaining static-token rejection and pre-allocation discovery ordering. --------- Signed-off-by: npub1dccv64krpcpse5cmkzfeh998cftungyatw3djt8jwdw6g43f7fyqzzmrf7 <6e30cd56c30e030cd31bb0939b94a7c257c9a09d5ba2d92cf2735da45629f248@buzz.block.builderlab.xyz> Co-authored-by: npub1dccv64krpcpse5cmkzfeh998cftungyatw3djt8jwdw6g43f7fyqzzmrf7 <6e30cd56c30e030cd31bb0939b94a7c257c9a09d5ba2d92cf2735da45629f248@buzz.block.builderlab.xyz>
This commit is contained in:
co-authored by
npub1dccv64krpcpse5cmkzfeh998cftungyatw3djt8jwdw6g43f7fyqzzmrf7
parent
8faf09f9ae
commit
4a2305170e
+145
-108
@@ -7,12 +7,17 @@
|
||||
//!
|
||||
//! - Static bearer (`DATABRICKS_TOKEN`): returned immediately.
|
||||
//! - PKCE cache hit: returned from disk without a network round-trip.
|
||||
//! - PKCE cache empty / no token: returns `Err(AgentError::LlmAuth)` — the
|
||||
//! caller degrades gracefully; no browser, no hang.
|
||||
//! - PKCE cache empty / no token: returns `Err(AgentError::LlmAuth)`.
|
||||
//!
|
||||
//! This helper never opens a browser. Callers choose whether to reject, degrade,
|
||||
//! or start a separate interactive authentication flow.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use reqwest::Client;
|
||||
|
||||
use crate::{
|
||||
auth::TokenSource,
|
||||
config::{Config, Provider},
|
||||
llm::build_token_source,
|
||||
types::AgentError,
|
||||
@@ -26,57 +31,22 @@ pub struct ModelEntry {
|
||||
pub name: String,
|
||||
}
|
||||
|
||||
/// Known Databricks AI Gateway v2 models — used as a fallback when the
|
||||
/// `api/ai-gateway/v2/endpoints` call returns an empty list.
|
||||
/// Known Databricks AI Gateway v2 models — used only when an authenticated
|
||||
/// `api/ai-gateway/v2/endpoints` call succeeds with an empty list.
|
||||
/// Mirrors goose's `DATABRICKS_V2_KNOWN_MODELS`.
|
||||
pub const DATABRICKS_V2_KNOWN_MODELS: &[&str] =
|
||||
&["databricks-gpt-5-5", "databricks-claude-opus-4-7"];
|
||||
|
||||
/// Returns the discovery-failure fallback catalog for a Databricks provider.
|
||||
///
|
||||
/// This is the list of models advertised by `session/new` when
|
||||
/// `discover_databricks_models` returns an error (e.g., no token available).
|
||||
///
|
||||
/// - `DatabricksV2` falls back to the configured model plus
|
||||
/// [`DATABRICKS_V2_KNOWN_MODELS`] so the model-picker is always populated for
|
||||
/// AI Gateway v2 users. The configured model leads: without it a fallback
|
||||
/// catalog can omit the very model the agent is running, leaving the picker
|
||||
/// unable to represent the current selection.
|
||||
/// - Legacy `Databricks` falls back to only the configured model — the
|
||||
/// `DATABRICKS_V2_KNOWN_MODELS` IDs are AI Gateway v2 endpoints that the
|
||||
/// `/serving-endpoints/{model}/invocations` API may not serve.
|
||||
///
|
||||
/// Extracting this as a pure function makes the split testable without
|
||||
/// spawning an async runtime or making network calls.
|
||||
pub fn discovery_failure_fallback(provider: Provider, configured_model: &str) -> Vec<ModelEntry> {
|
||||
// `resolve_model` does not trim, so a padded `DATABRICKS_MODEL` reaches here:
|
||||
// normalize once, or the dedupe below misses and the picker lists the model
|
||||
// twice (once padded, once from the known slate).
|
||||
let configured_model = configured_model.trim();
|
||||
let configured = ModelEntry {
|
||||
id: configured_model.to_string(),
|
||||
name: configured_model.to_string(),
|
||||
};
|
||||
match provider {
|
||||
Provider::DatabricksV2 => {
|
||||
let mut entries = Vec::with_capacity(DATABRICKS_V2_KNOWN_MODELS.len() + 1);
|
||||
if !configured_model.is_empty() {
|
||||
entries.push(configured);
|
||||
}
|
||||
entries.extend(
|
||||
DATABRICKS_V2_KNOWN_MODELS
|
||||
.iter()
|
||||
.filter(|id| **id != configured_model)
|
||||
.map(|id| ModelEntry {
|
||||
id: id.to_string(),
|
||||
name: id.to_string(),
|
||||
}),
|
||||
);
|
||||
entries
|
||||
}
|
||||
Provider::Databricks => vec![configured],
|
||||
_ => vec![configured],
|
||||
}
|
||||
const AUTHENTICATED_EMPTY_CATALOG_SUFFIX: &str = " (default catalog)";
|
||||
|
||||
fn authenticated_empty_v2_catalog() -> Vec<ModelEntry> {
|
||||
DATABRICKS_V2_KNOWN_MODELS
|
||||
.iter()
|
||||
.map(|id| ModelEntry {
|
||||
id: id.to_string(),
|
||||
name: format!("{id}{AUTHENTICATED_EMPTY_CATALOG_SUFFIX}"),
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Heuristic: `true` when a v2 AI Gateway endpoint name looks like it serves
|
||||
@@ -109,23 +79,47 @@ pub(crate) fn is_chat_capable_endpoint(name: &str) -> bool {
|
||||
///
|
||||
/// Returns a non-empty `Vec<ModelEntry>` on success. Returns
|
||||
/// `Err(AgentError::LlmAuth)` when no token is available (no static token,
|
||||
/// no PKCE cache) — callers should degrade gracefully rather than hanging.
|
||||
/// no PKCE cache). The helper itself never starts interactive authentication.
|
||||
///
|
||||
/// # Panics
|
||||
/// Never panics.
|
||||
pub async fn discover_databricks_models(cfg: &Config) -> Result<Vec<ModelEntry>, AgentError> {
|
||||
let token_source = build_token_source(cfg)?;
|
||||
let bearer = token_source.bearer_no_browser().await?;
|
||||
discover_databricks_models_with_token_source(cfg, build_token_source(cfg)?).await
|
||||
}
|
||||
|
||||
async fn discover_databricks_models_with_token_source(
|
||||
cfg: &Config,
|
||||
token_source: Arc<dyn TokenSource>,
|
||||
) -> Result<Vec<ModelEntry>, AgentError> {
|
||||
let mut bearer = token_source.bearer_no_browser().await?;
|
||||
let http = Client::new();
|
||||
let host = cfg.base_url.trim_end_matches('/');
|
||||
let mut refreshed = false;
|
||||
|
||||
match cfg.provider {
|
||||
Provider::Databricks => fetch_v1_models(&http, host, &bearer).await,
|
||||
Provider::DatabricksV2 => fetch_v2_models(&http, host, &bearer).await,
|
||||
_ => Err(AgentError::InvalidParams(
|
||||
"discover_databricks_models called for non-Databricks provider".into(),
|
||||
)),
|
||||
loop {
|
||||
let result = match cfg.provider {
|
||||
Provider::Databricks => fetch_v1_models(&http, host, &bearer).await,
|
||||
Provider::DatabricksV2 => fetch_v2_models(&http, host, &bearer).await,
|
||||
_ => {
|
||||
return Err(AgentError::InvalidParams(
|
||||
"discover_databricks_models called for non-Databricks provider".into(),
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
match result {
|
||||
Err(AgentError::LlmAuth(_)) if !refreshed => {
|
||||
refreshed = true;
|
||||
let fresh = token_source.refresh_now(&bearer).await?;
|
||||
if fresh == bearer {
|
||||
return Err(AgentError::LlmAuth(
|
||||
"Databricks rejected the configured credential".into(),
|
||||
));
|
||||
}
|
||||
bearer = fresh;
|
||||
}
|
||||
result => return result,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -149,6 +143,11 @@ async fn fetch_v1_models(
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
if status.as_u16() == 401 {
|
||||
return Err(AgentError::LlmAuth(format!(
|
||||
"Databricks model discovery HTTP {status}"
|
||||
)));
|
||||
}
|
||||
return Err(AgentError::Llm(format!(
|
||||
"Databricks model discovery HTTP {status}: {body}"
|
||||
)));
|
||||
@@ -264,6 +263,11 @@ async fn fetch_v2_models(
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
if status.as_u16() == 401 {
|
||||
return Err(AgentError::LlmAuth(format!(
|
||||
"Databricks v2 model discovery HTTP {status}"
|
||||
)));
|
||||
}
|
||||
return Err(AgentError::Llm(format!(
|
||||
"Databricks v2 model discovery HTTP {status}: {body}"
|
||||
)));
|
||||
@@ -286,13 +290,7 @@ async fn fetch_v2_models(
|
||||
|
||||
// Fall back to known-model list if the API returned nothing.
|
||||
if all_endpoints.is_empty() {
|
||||
return Ok(DATABRICKS_V2_KNOWN_MODELS
|
||||
.iter()
|
||||
.map(|id| ModelEntry {
|
||||
id: id.to_string(),
|
||||
name: id.to_string(),
|
||||
})
|
||||
.collect());
|
||||
return Ok(authenticated_empty_v2_catalog());
|
||||
}
|
||||
|
||||
sort_v2_endpoints_newest_first(&mut all_endpoints);
|
||||
@@ -396,6 +394,77 @@ pub(crate) fn parse_v2_endpoints_page(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use async_trait::async_trait;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
struct RefreshingTestTokenSource {
|
||||
refreshes: AtomicUsize,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl TokenSource for RefreshingTestTokenSource {
|
||||
async fn bearer(&self) -> Result<String, AgentError> {
|
||||
Ok("rejected".into())
|
||||
}
|
||||
|
||||
async fn refresh_now(&self, rejected: &str) -> Result<String, AgentError> {
|
||||
assert_eq!(rejected, "rejected");
|
||||
self.refreshes.fetch_add(1, Ordering::SeqCst);
|
||||
Ok("fresh".into())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn discovery_refreshes_rejected_bearer_once_then_retries_successfully() {
|
||||
use axum::{
|
||||
extract::Query,
|
||||
http::{HeaderMap, StatusCode},
|
||||
routing::get,
|
||||
Json, Router,
|
||||
};
|
||||
use std::collections::HashMap;
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let host = format!("http://{}", listener.local_addr().unwrap());
|
||||
let requests = Arc::new(AtomicUsize::new(0));
|
||||
let requests_for_route = requests.clone();
|
||||
let app = Router::new().route(
|
||||
"/api/ai-gateway/v2/endpoints",
|
||||
get(
|
||||
move |headers: HeaderMap, Query(_query): Query<HashMap<String, String>>| {
|
||||
let requests = requests_for_route.clone();
|
||||
async move {
|
||||
requests.fetch_add(1, Ordering::SeqCst);
|
||||
match headers
|
||||
.get("authorization")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
{
|
||||
Some("Bearer fresh") => Ok(Json(serde_json::json!({
|
||||
"endpoints": [{"name": "discovered-model"}],
|
||||
"next_page_token": null,
|
||||
}))),
|
||||
_ => Err((StatusCode::UNAUTHORIZED, "rejected")),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
);
|
||||
tokio::spawn(async move {
|
||||
let _ = axum::serve(listener, app).await;
|
||||
});
|
||||
|
||||
let source = Arc::new(RefreshingTestTokenSource {
|
||||
refreshes: AtomicUsize::new(0),
|
||||
});
|
||||
let cfg = Config::for_discovery(Provider::DatabricksV2, String::new(), host);
|
||||
let models = discover_databricks_models_with_token_source(&cfg, source.clone())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(models[0].id, "discovered-model");
|
||||
assert_eq!(source.refreshes.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(requests.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn v1_parse_filters_ready_chat_endpoints() {
|
||||
@@ -574,6 +643,17 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_empty_v2_catalog_marks_fallback_provenance() {
|
||||
let models = authenticated_empty_v2_catalog();
|
||||
let ids: Vec<&str> = models.iter().map(|model| model.id.as_str()).collect();
|
||||
|
||||
assert_eq!(ids, DATABRICKS_V2_KNOWN_MODELS);
|
||||
assert!(models.iter().all(|model| {
|
||||
model.name == format!("{}{AUTHENTICATED_EMPTY_CATALOG_SUFFIX}", model.id)
|
||||
}));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn is_chat_capable_endpoint_keeps_unrecognised_names() {
|
||||
// Prefer including over silently dropping — an unknown family is kept.
|
||||
@@ -585,47 +665,4 @@ mod tests {
|
||||
assert!(!is_chat_capable_endpoint("databricks-gte-large-en"));
|
||||
assert!(!is_chat_capable_endpoint("databricks-qwen3-embedding-0-6b"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn v2_discovery_failure_fallback_leads_with_configured_model() {
|
||||
let result = discovery_failure_fallback(Provider::DatabricksV2, "databricks-claude-opus-5");
|
||||
let ids: Vec<&str> = result.iter().map(|m| m.id.as_str()).collect();
|
||||
|
||||
// The running model must be representable in the picker even when
|
||||
// discovery failed, so it leads the fallback catalog.
|
||||
assert_eq!(ids.first(), Some(&"databricks-claude-opus-5"));
|
||||
for known in DATABRICKS_V2_KNOWN_MODELS {
|
||||
assert!(ids.contains(known), "fallback must retain '{known}'");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn v2_discovery_failure_fallback_does_not_duplicate_configured_model() {
|
||||
let configured = DATABRICKS_V2_KNOWN_MODELS[0];
|
||||
let result = discovery_failure_fallback(Provider::DatabricksV2, configured);
|
||||
let occurrences = result.iter().filter(|m| m.id == configured).count();
|
||||
assert_eq!(occurrences, 1, "got: {result:?}");
|
||||
assert_eq!(result.len(), DATABRICKS_V2_KNOWN_MODELS.len());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn v2_discovery_failure_fallback_tolerates_blank_configured_model() {
|
||||
for configured in ["", " "] {
|
||||
let result = discovery_failure_fallback(Provider::DatabricksV2, configured);
|
||||
let ids: Vec<&str> = result.iter().map(|m| m.id.as_str()).collect();
|
||||
assert_eq!(ids, DATABRICKS_V2_KNOWN_MODELS.to_vec());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn v2_discovery_failure_fallback_dedupes_a_padded_configured_model() {
|
||||
// `DATABRICKS_MODEL=" databricks-gpt-5-5 "` reaches here untrimmed, and an
|
||||
// untrimmed comparison would list the model twice — once padded, once from
|
||||
// the known slate.
|
||||
let configured = DATABRICKS_V2_KNOWN_MODELS[0];
|
||||
let result =
|
||||
discovery_failure_fallback(Provider::DatabricksV2, &format!(" {configured} "));
|
||||
let ids: Vec<&str> = result.iter().map(|m| m.id.as_str()).collect();
|
||||
assert_eq!(ids, DATABRICKS_V2_KNOWN_MODELS.to_vec());
|
||||
}
|
||||
}
|
||||
|
||||
+108
-145
@@ -54,10 +54,10 @@ struct App {
|
||||
llm: Arc<Llm>,
|
||||
sessions: Mutex<HashMap<String, Session>>,
|
||||
/// Cached model catalog for Databricks providers. Populated lazily on the
|
||||
/// first successful `session/new` discovery call. When discovery fails (e.g.
|
||||
/// auth missing or a transient network error) the cell is intentionally left
|
||||
/// empty so the next `session/new` call retries — a transient failure never
|
||||
/// pins the degraded fallback catalog for the process lifetime.
|
||||
/// first successful `session/new` discovery call. Failed discovery is never
|
||||
/// cached: static-token authentication errors reject session creation, while
|
||||
/// OAuth authentication and non-auth errors use the configured model for that
|
||||
/// response and retry on the next session.
|
||||
models_cache: tokio::sync::OnceCell<Vec<ModelEntry>>,
|
||||
}
|
||||
|
||||
@@ -135,6 +135,12 @@ pub fn run() -> Result<(), Box<dyn std::error::Error>> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn authenticate_databricks(host: &str) -> Result<(), AgentError> {
|
||||
auth::PkceOAuthTokenSource::new(llm::databricks_pkce_config(host))?
|
||||
.interactive_login()
|
||||
.await
|
||||
}
|
||||
|
||||
/// `buzz-agent auth <provider>` — run the interactive auth flow for a
|
||||
/// provider and persist the result, then exit. Today this supports Databricks
|
||||
/// OAuth 2.0 PKCE. Reads `DATABRICKS_HOST` from env; needs a browser on the
|
||||
@@ -145,18 +151,7 @@ async fn auth_subcommand(args: &[String]) -> Result<(), Box<dyn std::error::Erro
|
||||
Some("databricks" | "databricks_v2" | "databricks-v2") => {
|
||||
let host = std::env::var("DATABRICKS_HOST")
|
||||
.map_err(|_| "auth databricks: DATABRICKS_HOST required")?;
|
||||
let pkce = auth::PkceOAuthConfig {
|
||||
discovery_url: format!(
|
||||
"{}/oidc/.well-known/oauth-authorization-server",
|
||||
host.trim_end_matches('/')
|
||||
),
|
||||
client_id: "databricks-cli".into(),
|
||||
scopes: vec!["all-apis".into(), "offline_access".into()],
|
||||
cache_namespace: "databricks".into(),
|
||||
cache_dir_override: None,
|
||||
};
|
||||
let src = auth::PkceOAuthTokenSource::new(pkce)?;
|
||||
src.interactive_login().await?;
|
||||
authenticate_databricks(&host).await?;
|
||||
eprintln!("Authenticated. Token cached under ~/.config/buzz-agent/oauth/databricks/.");
|
||||
Ok(())
|
||||
}
|
||||
@@ -317,26 +312,27 @@ async fn initialize(id: Value, params: Value, wire_tx: &WireSender) {
|
||||
///
|
||||
/// Tries to use a previously-cached successful discovery result. If the cache is empty,
|
||||
/// runs `discover` and — on success — populates the cache for future calls. On failure
|
||||
/// the cell is intentionally left empty so the next session retries; the provider-aware
|
||||
/// fallback is returned for the immediate response only.
|
||||
/// the error is returned and the cell is intentionally left empty so the next session retries.
|
||||
///
|
||||
/// Extracted from `session_new` so that tests can drive this path with an injected
|
||||
/// discovery future without requiring a full `App` / transport stack.
|
||||
async fn resolve_models_catalog(
|
||||
cache: &tokio::sync::OnceCell<Vec<ModelEntry>>,
|
||||
provider: crate::config::Provider,
|
||||
model: &str,
|
||||
discover: impl std::future::Future<Output = Result<Vec<ModelEntry>, AgentError>>,
|
||||
) -> Vec<ModelEntry> {
|
||||
match cache.get_or_try_init(|| discover).await {
|
||||
Ok(cached) => cached.clone(),
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
"model catalog discovery failed: {e}; using fallback (will retry next session)"
|
||||
);
|
||||
crate::catalog::discovery_failure_fallback(provider, model)
|
||||
}
|
||||
}
|
||||
) -> Result<Vec<ModelEntry>, AgentError> {
|
||||
cache.get_or_try_init(|| discover).await.cloned()
|
||||
}
|
||||
|
||||
/// Return the configured model as a one-entry catalog for this response.
|
||||
///
|
||||
/// This value is never written to `models_cache`; failed discovery must be retried by
|
||||
/// the next session rather than pinning degraded state for the process lifetime.
|
||||
fn configured_model_fallback(model: &str) -> Vec<ModelEntry> {
|
||||
let model = model.trim().to_string();
|
||||
vec![ModelEntry {
|
||||
id: model.clone(),
|
||||
name: model,
|
||||
}]
|
||||
}
|
||||
|
||||
async fn session_new(app: &Arc<App>, id: Value, params: Value, wire_tx: &WireSender) {
|
||||
@@ -400,6 +396,50 @@ async fn session_new(app: &Arc<App>, id: Value, params: Value, wire_tx: &WireSen
|
||||
}
|
||||
Arc::from(prompt)
|
||||
};
|
||||
// Resolve the model catalog before spawning MCP servers or registering a
|
||||
// session. A configured static credential cannot recover interactively, so
|
||||
// its authentication failure rejects before allocation. OAuth authentication
|
||||
// failures and other catalog failures use only the configured model for this
|
||||
// response, without caching, so session/prompt can run the existing PKCE flow.
|
||||
let available_models: Vec<Value> = {
|
||||
use crate::config::Provider;
|
||||
match app.cfg.provider {
|
||||
Provider::Databricks | Provider::DatabricksV2 => {
|
||||
let models = match resolve_models_catalog(
|
||||
&app.models_cache,
|
||||
discover_databricks_models(&app.cfg),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(models) => models,
|
||||
Err(error @ AgentError::LlmAuth(_)) if !app.cfg.api_key.is_empty() => {
|
||||
return reject(wire_tx, id, error.json_rpc_code(), &error.to_string())
|
||||
.await;
|
||||
}
|
||||
Err(error @ AgentError::LlmAuth(_)) => {
|
||||
tracing::warn!(
|
||||
error = %error,
|
||||
"Databricks OAuth model catalog unavailable; using configured model"
|
||||
);
|
||||
configured_model_fallback(&app.cfg.model)
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(
|
||||
error = %error,
|
||||
"Databricks model catalog unavailable; using configured model"
|
||||
);
|
||||
configured_model_fallback(&app.cfg.model)
|
||||
}
|
||||
};
|
||||
models
|
||||
.iter()
|
||||
.map(|m| json!({ "modelId": m.id, "name": m.name }))
|
||||
.collect()
|
||||
}
|
||||
_ => vec![json!({ "modelId": app.cfg.model, "name": app.cfg.model })],
|
||||
}
|
||||
};
|
||||
|
||||
let mcp = match McpRegistry::spawn_all(&app.cfg, &p.mcp_servers, &p.cwd).await {
|
||||
Ok(m) => Arc::new(m),
|
||||
Err(e) => return reject(wire_tx, id, e.json_rpc_code(), &e.to_string()).await,
|
||||
@@ -445,36 +485,6 @@ async fn session_new(app: &Arc<App>, id: Value, params: Value, wire_tx: &WireSen
|
||||
);
|
||||
drop(sessions);
|
||||
|
||||
// Build a models catalog for the `session/new` response. For Databricks
|
||||
// providers this advertises available models so the desktop ModelPicker and
|
||||
// pool can resolve `session/set_model` switches. For Anthropic/OpenAI we
|
||||
// report only the configured model — live switching on those providers
|
||||
// effectively requires respawn.
|
||||
//
|
||||
// `models_cache` caches only a successful discovery result (`get_or_try_init`
|
||||
// leaves the cell empty on error so the next `session/new` call retries). On
|
||||
// discovery failure the fallback is used for the immediate response without
|
||||
// being written to the cell.
|
||||
let available_models: Vec<Value> = {
|
||||
use crate::config::Provider;
|
||||
match app.cfg.provider {
|
||||
Provider::Databricks | Provider::DatabricksV2 => {
|
||||
let models = resolve_models_catalog(
|
||||
&app.models_cache,
|
||||
app.cfg.provider,
|
||||
&app.cfg.model,
|
||||
discover_databricks_models(&app.cfg),
|
||||
)
|
||||
.await;
|
||||
models
|
||||
.iter()
|
||||
.map(|m| json!({ "modelId": m.id, "name": m.name }))
|
||||
.collect()
|
||||
}
|
||||
_ => vec![json!({ "modelId": app.cfg.model, "name": app.cfg.model })],
|
||||
}
|
||||
};
|
||||
|
||||
wire::send(
|
||||
wire_tx,
|
||||
wire::ok(
|
||||
@@ -870,8 +880,7 @@ fn session_token() -> Result<String, String> {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::catalog::{discovery_failure_fallback, ModelEntry, DATABRICKS_V2_KNOWN_MODELS};
|
||||
use crate::config::Provider;
|
||||
use crate::catalog::ModelEntry;
|
||||
use crate::types::AgentError;
|
||||
|
||||
/// Regression: a discovery error must not pin the models_cache for the process lifetime.
|
||||
@@ -884,23 +893,14 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn models_cache_does_not_pin_on_discovery_error() {
|
||||
let cache: tokio::sync::OnceCell<Vec<ModelEntry>> = tokio::sync::OnceCell::new();
|
||||
let provider = Provider::DatabricksV2;
|
||||
let model = "my-configured-model";
|
||||
|
||||
// First call — discovery fails. Cell must remain empty; fallback returned.
|
||||
let first = crate::resolve_models_catalog(&cache, provider, model, async {
|
||||
Err::<Vec<ModelEntry>, AgentError>(AgentError::LlmAuth("transient failure".into()))
|
||||
// First call — discovery failure is surfaced and leaves the cell empty.
|
||||
let error = crate::resolve_models_catalog(&cache, async {
|
||||
Err::<Vec<ModelEntry>, AgentError>(AgentError::Llm("transient failure".into()))
|
||||
})
|
||||
.await;
|
||||
assert!(
|
||||
cache.get().is_none(),
|
||||
"cell must be empty after a discovery error — next session must retry"
|
||||
);
|
||||
let expected_fallback = discovery_failure_fallback(provider, model);
|
||||
assert_eq!(
|
||||
first, expected_fallback,
|
||||
"error path must return the provider-aware fallback"
|
||||
);
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(error, AgentError::Llm(_)));
|
||||
|
||||
// Second call — discovery succeeds. Cell is now populated and returned.
|
||||
let discovered = vec![ModelEntry {
|
||||
@@ -908,10 +908,11 @@ mod tests {
|
||||
name: "databricks-meta-llama-3-1-70b-instruct".into(),
|
||||
}];
|
||||
let discovered_clone = discovered.clone();
|
||||
let second = crate::resolve_models_catalog(&cache, provider, model, async move {
|
||||
let second = crate::resolve_models_catalog(&cache, async move {
|
||||
Ok::<Vec<ModelEntry>, AgentError>(discovered_clone)
|
||||
})
|
||||
.await;
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
second, discovered,
|
||||
"second call must return the discovered catalog"
|
||||
@@ -927,78 +928,40 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
/// Regression: legacy `Provider::Databricks` must not advertise v2 AI Gateway model IDs
|
||||
/// on discovery failure (Wes W1). This test calls `discovery_failure_fallback` directly —
|
||||
/// the same helper used by `session_new` — and verifies the split behavior. It FAILS if
|
||||
/// the arm is un-split (i.e., if both providers return the v2 catalog on failure).
|
||||
#[test]
|
||||
fn databricks_discovery_failure_fallback_legacy_returns_configured_model_only() {
|
||||
let configured = "my-serving-endpoint";
|
||||
let result = discovery_failure_fallback(Provider::Databricks, configured);
|
||||
#[tokio::test]
|
||||
async fn models_catalog_does_not_cache_oauth_auth_fallback() {
|
||||
let cache: tokio::sync::OnceCell<Vec<ModelEntry>> = tokio::sync::OnceCell::new();
|
||||
let error = crate::resolve_models_catalog(&cache, async {
|
||||
Err::<Vec<ModelEntry>, AgentError>(AgentError::LlmAuth("sign in again".into()))
|
||||
})
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
// Legacy Databricks must advertise exactly the configured model — nothing more.
|
||||
assert_eq!(
|
||||
result.len(),
|
||||
1,
|
||||
"legacy Databricks fallback must contain exactly one entry, got: {result:?}"
|
||||
);
|
||||
assert_eq!(
|
||||
result[0].id, configured,
|
||||
"legacy Databricks fallback must be the configured model"
|
||||
);
|
||||
assert!(matches!(error, AgentError::LlmAuth(_)));
|
||||
assert!(cache.get().is_none());
|
||||
|
||||
// Crucially: must NOT contain any DATABRICKS_V2_KNOWN_MODELS entry.
|
||||
let v2_ids: Vec<&str> = DATABRICKS_V2_KNOWN_MODELS.to_vec();
|
||||
for id in &result {
|
||||
assert!(
|
||||
!v2_ids.contains(&id.id.as_str()),
|
||||
"legacy Databricks fallback must not include v2 ID '{}' — that endpoint \
|
||||
may not be served by /serving-endpoints/{{model}}/invocations",
|
||||
id.id
|
||||
);
|
||||
}
|
||||
let discovered = vec![ModelEntry {
|
||||
id: "authenticated-model".into(),
|
||||
name: "authenticated-model".into(),
|
||||
}];
|
||||
let result = crate::resolve_models_catalog(&cache, async {
|
||||
Ok::<Vec<ModelEntry>, AgentError>(discovered.clone())
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(result, discovered);
|
||||
assert_eq!(cache.get(), Some(&discovered));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn databricks_discovery_failure_fallback_v2_returns_known_models_catalog() {
|
||||
let configured = "my-configured-model";
|
||||
let result = discovery_failure_fallback(Provider::DatabricksV2, configured);
|
||||
|
||||
// DatabricksV2 must return the full DATABRICKS_V2_KNOWN_MODELS list,
|
||||
// plus the configured model so the picker can still represent the model
|
||||
// the agent is actually running.
|
||||
fn configured_model_fallback_is_trimmed_and_singular() {
|
||||
assert_eq!(
|
||||
result.len(),
|
||||
DATABRICKS_V2_KNOWN_MODELS.len() + 1,
|
||||
"DatabricksV2 fallback must return all known models plus the configured model"
|
||||
);
|
||||
let result_ids: Vec<&str> = result.iter().map(|m| m.id.as_str()).collect();
|
||||
for known_id in DATABRICKS_V2_KNOWN_MODELS {
|
||||
assert!(
|
||||
result_ids.contains(known_id),
|
||||
"DatabricksV2 fallback must include known model '{known_id}'"
|
||||
);
|
||||
}
|
||||
assert!(
|
||||
result_ids.contains(&configured),
|
||||
"DatabricksV2 fallback must include the configured model"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn databricks_discovery_failure_fallback_split_verified() {
|
||||
// This test FAILS if the v1/v2 arms are merged back into one — it directly verifies
|
||||
// that the two providers' error-path behavior diverges (Wes W1 protection).
|
||||
let v1 = discovery_failure_fallback(Provider::Databricks, "my-endpoint");
|
||||
let v2 = discovery_failure_fallback(Provider::DatabricksV2, "my-endpoint");
|
||||
|
||||
let v1_ids: Vec<&str> = v1.iter().map(|m| m.id.as_str()).collect();
|
||||
let v2_ids: Vec<&str> = v2.iter().map(|m| m.id.as_str()).collect();
|
||||
|
||||
assert_ne!(
|
||||
v1_ids, v2_ids,
|
||||
"Provider::Databricks and Provider::DatabricksV2 must return different \
|
||||
fallback catalogs — if they are equal, the W1 arm split has been reverted"
|
||||
crate::configured_model_fallback(" configured-model "),
|
||||
vec![ModelEntry {
|
||||
id: "configured-model".into(),
|
||||
name: "configured-model".into(),
|
||||
}]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1910,6 +1910,22 @@ where
|
||||
unreachable!("loop always returns on its final iteration (attempt + 1 == MAX_RETRIES)");
|
||||
}
|
||||
|
||||
pub(crate) fn databricks_pkce_config(host: &str) -> PkceOAuthConfig {
|
||||
PkceOAuthConfig {
|
||||
discovery_url: format!(
|
||||
"{}/oidc/.well-known/oauth-authorization-server",
|
||||
host.trim_end_matches('/')
|
||||
),
|
||||
client_id: DATABRICKS_CLIENT_ID.into(),
|
||||
scopes: DATABRICKS_OAUTH_SCOPES
|
||||
.iter()
|
||||
.map(|scope| (*scope).into())
|
||||
.collect(),
|
||||
cache_namespace: "databricks".into(),
|
||||
cache_dir_override: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the `TokenSource` for the configured provider.
|
||||
///
|
||||
/// - `Provider::Anthropic`: a static source seeded from `cfg.api_key`. It's
|
||||
@@ -1929,21 +1945,9 @@ pub(crate) fn build_token_source(cfg: &Config) -> Result<Arc<dyn TokenSource>, A
|
||||
if !cfg.api_key.is_empty() {
|
||||
return Ok(Arc::new(StaticTokenSource::new(cfg.api_key.clone())));
|
||||
}
|
||||
let discovery_url = format!(
|
||||
"{}/oidc/.well-known/oauth-authorization-server",
|
||||
cfg.base_url.trim_end_matches('/')
|
||||
);
|
||||
let pkce = PkceOAuthConfig {
|
||||
discovery_url,
|
||||
client_id: DATABRICKS_CLIENT_ID.into(),
|
||||
scopes: DATABRICKS_OAUTH_SCOPES
|
||||
.iter()
|
||||
.map(|s| (*s).into())
|
||||
.collect(),
|
||||
cache_namespace: "databricks".into(),
|
||||
cache_dir_override: None,
|
||||
};
|
||||
Ok(PkceOAuthTokenSource::new(pkce)?)
|
||||
Ok(PkceOAuthTokenSource::new(databricks_pkce_config(
|
||||
&cfg.base_url,
|
||||
))?)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,6 +20,7 @@ use axum::{routing::get, routing::post, Json, Router};
|
||||
use buzz_agent::auth::{PkceOAuthConfig, PkceOAuthTokenSource, TokenSource};
|
||||
use serde::Deserialize;
|
||||
use serde_json::json;
|
||||
use sha2::{Digest, Sha256};
|
||||
use tempfile::TempDir;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
@@ -457,6 +458,7 @@ struct AgentHarness {
|
||||
stdin: tokio::process::ChildStdin,
|
||||
stdout: BufReader<tokio::process::ChildStdout>,
|
||||
next_id: i64,
|
||||
_home: Option<TempDir>,
|
||||
}
|
||||
|
||||
impl Drop for AgentHarness {
|
||||
@@ -467,20 +469,65 @@ impl Drop for AgentHarness {
|
||||
|
||||
impl AgentHarness {
|
||||
async fn spawn_provider(provider: &str, base_url: &str, model: &str) -> Self {
|
||||
Self::spawn_provider_with_options(provider, base_url, model, 1, Some("test-bearer")).await
|
||||
}
|
||||
|
||||
async fn spawn_oauth_provider(
|
||||
provider: &str,
|
||||
base_url: &str,
|
||||
model: &str,
|
||||
max_sessions: usize,
|
||||
) -> Self {
|
||||
Self::spawn_provider_with_options(provider, base_url, model, max_sessions, None).await
|
||||
}
|
||||
|
||||
async fn spawn_provider_with_max_sessions(
|
||||
provider: &str,
|
||||
base_url: &str,
|
||||
model: &str,
|
||||
max_sessions: usize,
|
||||
) -> Self {
|
||||
Self::spawn_provider_with_options(
|
||||
provider,
|
||||
base_url,
|
||||
model,
|
||||
max_sessions,
|
||||
Some("test-bearer"),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn spawn_provider_with_options(
|
||||
provider: &str,
|
||||
base_url: &str,
|
||||
model: &str,
|
||||
max_sessions: usize,
|
||||
token: Option<&str>,
|
||||
) -> Self {
|
||||
let bin = env!("CARGO_BIN_EXE_buzz-agent");
|
||||
let home = token
|
||||
.is_none()
|
||||
.then(|| TempDir::new().expect("create isolated OAuth home"));
|
||||
let mut cmd = tokio::process::Command::new(bin);
|
||||
cmd.env("BUZZ_AGENT_PROVIDER", provider)
|
||||
.env("DATABRICKS_HOST", base_url)
|
||||
.env("DATABRICKS_MODEL", model)
|
||||
.env("DATABRICKS_TOKEN", "test-bearer")
|
||||
.env_remove("DATABRICKS_TOKEN")
|
||||
.env("BUZZ_AGENT_LLM_TIMEOUT_SECS", "5")
|
||||
.env("BUZZ_AGENT_TOOL_TIMEOUT_SECS", "5")
|
||||
.env("BUZZ_AGENT_MAX_ROUNDS", "2")
|
||||
.env("BUZZ_AGENT_MAX_SESSIONS", max_sessions.to_string())
|
||||
.env("BUZZ_AGENT_MCP_INIT_TIMEOUT_SECS", "2")
|
||||
.stdin(Stdio::piped())
|
||||
.stdout(Stdio::piped())
|
||||
.stderr(Stdio::null())
|
||||
.kill_on_drop(true);
|
||||
if let Some(token) = token {
|
||||
cmd.env("DATABRICKS_TOKEN", token);
|
||||
}
|
||||
if let Some(home) = &home {
|
||||
cmd.env("HOME", home.path());
|
||||
}
|
||||
let mut child = cmd.spawn().expect("spawn buzz-agent");
|
||||
let stdin = child.stdin.take().unwrap();
|
||||
let stdout = BufReader::new(child.stdout.take().unwrap());
|
||||
@@ -489,9 +536,17 @@ impl AgentHarness {
|
||||
stdin,
|
||||
stdout,
|
||||
next_id: 1,
|
||||
_home: home,
|
||||
}
|
||||
}
|
||||
|
||||
fn oauth_home(&self) -> &std::path::Path {
|
||||
self._home
|
||||
.as_ref()
|
||||
.expect("harness was not started in OAuth mode")
|
||||
.path()
|
||||
}
|
||||
|
||||
async fn send(&mut self, method: &str, params: serde_json::Value) -> i64 {
|
||||
let id = self.next_id;
|
||||
self.next_id += 1;
|
||||
@@ -938,3 +993,288 @@ async fn session_set_model_empty_model_id_returns_error() {
|
||||
"error message must mention modelId, got: {msg}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn model_discovery_surfaces_rejected_static_token_as_auth_failure() {
|
||||
use axum::http::StatusCode;
|
||||
use buzz_agent::config::{Config, Provider};
|
||||
use buzz_agent::discover_databricks_models;
|
||||
|
||||
let requests = Arc::new(AtomicU64::new(0));
|
||||
let requests_for_route = requests.clone();
|
||||
let listener = tokio::net::TcpListener::bind(SocketAddr::from(([127, 0, 0, 1], 0)))
|
||||
.await
|
||||
.unwrap();
|
||||
let host = format!("http://{}", listener.local_addr().unwrap());
|
||||
let app = Router::new().route(
|
||||
"/api/ai-gateway/v2/endpoints",
|
||||
get(move || {
|
||||
let requests = requests_for_route.clone();
|
||||
async move {
|
||||
requests.fetch_add(1, Ordering::SeqCst);
|
||||
(StatusCode::UNAUTHORIZED, "rejected bearer rejected")
|
||||
}
|
||||
}),
|
||||
);
|
||||
tokio::spawn(async move {
|
||||
let _ = axum::serve(listener, app).await;
|
||||
});
|
||||
|
||||
let cfg = Config::for_discovery(Provider::DatabricksV2, "rejected".into(), host);
|
||||
let error = discover_databricks_models(&cfg).await.unwrap_err();
|
||||
|
||||
assert!(
|
||||
error.to_string().starts_with("llm auth:"),
|
||||
"401 must retain auth semantics: {error}"
|
||||
);
|
||||
assert!(
|
||||
!error.to_string().contains("rejected bearer"),
|
||||
"auth errors must not propagate provider bodies that may echo credentials: {error}"
|
||||
);
|
||||
assert_eq!(
|
||||
requests.load(Ordering::SeqCst),
|
||||
1,
|
||||
"a static token cannot refresh, so discovery must not issue a duplicate request"
|
||||
);
|
||||
}
|
||||
|
||||
fn databricks_oauth_cache_path(home: &std::path::Path, host: &str) -> std::path::PathBuf {
|
||||
let discovery_url = format!(
|
||||
"{}/oidc/.well-known/oauth-authorization-server",
|
||||
host.trim_end_matches('/')
|
||||
);
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(discovery_url.as_bytes());
|
||||
hasher.update(b"|");
|
||||
hasher.update(b"databricks-cli");
|
||||
hasher.update(b"|");
|
||||
hasher.update(b"all-apis,offline_access");
|
||||
let hash = hex::encode(hasher.finalize());
|
||||
home.join(".config")
|
||||
.join("buzz-agent")
|
||||
.join("oauth")
|
||||
.join("databricks")
|
||||
.join(format!("{hash}.json"))
|
||||
}
|
||||
|
||||
fn write_cached_oauth_token(home: &std::path::Path, host: &str, access_token: &str) {
|
||||
let path = databricks_oauth_cache_path(home, host);
|
||||
std::fs::create_dir_all(path.parent().unwrap()).unwrap();
|
||||
std::fs::write(
|
||||
path,
|
||||
serde_json::to_vec(&json!({
|
||||
"access_token": access_token,
|
||||
"refresh_token": null,
|
||||
"expires_at": SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs()
|
||||
+ 3600,
|
||||
}))
|
||||
.unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_missing_token_uses_configured_model_then_retries_discovery() {
|
||||
let attempts = Arc::new(AtomicU64::new(0));
|
||||
let attempts_for_route = attempts.clone();
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let host = format!("http://{}", listener.local_addr().unwrap());
|
||||
let app = Router::new().route(
|
||||
"/api/ai-gateway/v2/endpoints",
|
||||
get(move || {
|
||||
let attempts = attempts_for_route.clone();
|
||||
async move {
|
||||
attempts.fetch_add(1, Ordering::SeqCst);
|
||||
Json(json!({
|
||||
"endpoints": [{"name": "authenticated-model"}],
|
||||
"next_page_token": null,
|
||||
}))
|
||||
}
|
||||
}),
|
||||
);
|
||||
tokio::spawn(async move {
|
||||
let _ = axum::serve(listener, app).await;
|
||||
});
|
||||
|
||||
let configured_model = " configured-model ";
|
||||
let mut h =
|
||||
AgentHarness::spawn_oauth_provider("databricks_v2", &host, configured_model, 2).await;
|
||||
let initialize = h
|
||||
.send(
|
||||
"initialize",
|
||||
json!({ "protocolVersion": 1, "clientCapabilities": {} }),
|
||||
)
|
||||
.await;
|
||||
assert!(h.recv_for(initialize).await.get("result").is_some());
|
||||
|
||||
let first = h
|
||||
.send("session/new", json!({ "cwd": "/tmp", "mcpServers": [] }))
|
||||
.await;
|
||||
let first_response = h.recv_for(first).await;
|
||||
assert!(
|
||||
first_response["result"]["sessionId"].is_string(),
|
||||
"missing OAuth token blocked session creation: {first_response}"
|
||||
);
|
||||
assert_eq!(
|
||||
first_response["result"]["models"]["availableModels"],
|
||||
json!([{"modelId": "configured-model", "name": "configured-model"}])
|
||||
);
|
||||
assert_eq!(attempts.load(Ordering::SeqCst), 0);
|
||||
|
||||
write_cached_oauth_token(h.oauth_home(), &host, "cached-bearer");
|
||||
|
||||
let second = h
|
||||
.send("session/new", json!({ "cwd": "/tmp", "mcpServers": [] }))
|
||||
.await;
|
||||
let second_response = h.recv_for(second).await;
|
||||
assert!(
|
||||
second_response["result"]["sessionId"].is_string(),
|
||||
"later authenticated session failed: {second_response}"
|
||||
);
|
||||
assert_eq!(
|
||||
second_response["result"]["models"]["availableModels"],
|
||||
json!([{"modelId": "authenticated-model", "name": "authenticated-model"}])
|
||||
);
|
||||
assert_eq!(
|
||||
attempts.load(Ordering::SeqCst),
|
||||
1,
|
||||
"OAuth fallback was cached instead of retrying discovery"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn non_auth_discovery_failure_uses_configured_model_without_caching_fallback() {
|
||||
use axum::http::StatusCode;
|
||||
|
||||
let attempts = Arc::new(AtomicU64::new(0));
|
||||
let attempts_for_route = attempts.clone();
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let host = format!("http://{}", listener.local_addr().unwrap());
|
||||
let app = Router::new().route(
|
||||
"/api/ai-gateway/v2/endpoints",
|
||||
get(move || {
|
||||
let attempts = attempts_for_route.clone();
|
||||
async move {
|
||||
attempts.fetch_add(1, Ordering::SeqCst);
|
||||
(StatusCode::SERVICE_UNAVAILABLE, "catalog unavailable")
|
||||
}
|
||||
}),
|
||||
);
|
||||
tokio::spawn(async move {
|
||||
let _ = axum::serve(listener, app).await;
|
||||
});
|
||||
|
||||
let configured_model = " configured-model ";
|
||||
let normalized_configured_model = configured_model.trim();
|
||||
let mut h =
|
||||
AgentHarness::spawn_provider_with_max_sessions("databricks_v2", &host, configured_model, 2)
|
||||
.await;
|
||||
let initialize = h
|
||||
.send(
|
||||
"initialize",
|
||||
json!({ "protocolVersion": 1, "clientCapabilities": {} }),
|
||||
)
|
||||
.await;
|
||||
assert!(h.recv_for(initialize).await.get("result").is_some());
|
||||
|
||||
for expected_attempts in 1..=2 {
|
||||
let request = h
|
||||
.send("session/new", json!({ "cwd": "/tmp", "mcpServers": [] }))
|
||||
.await;
|
||||
let response = h.recv_for(request).await;
|
||||
assert!(
|
||||
response["result"]["sessionId"].is_string(),
|
||||
"non-auth catalog failure blocked session creation: {response}"
|
||||
);
|
||||
assert_eq!(
|
||||
response["result"]["models"]["availableModels"],
|
||||
json!([{"modelId": normalized_configured_model, "name": normalized_configured_model}])
|
||||
);
|
||||
assert_eq!(attempts.load(Ordering::SeqCst), expected_attempts);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rejected_static_token_does_not_consume_capacity_or_spawn_mcp() {
|
||||
use axum::http::StatusCode;
|
||||
|
||||
let attempts = Arc::new(AtomicU64::new(0));
|
||||
let attempts_for_route = attempts.clone();
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let host = format!("http://{}", listener.local_addr().unwrap());
|
||||
let app = Router::new().route(
|
||||
"/api/ai-gateway/v2/endpoints",
|
||||
get(move || {
|
||||
let attempts = attempts_for_route.clone();
|
||||
async move {
|
||||
if attempts.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
Err((StatusCode::UNAUTHORIZED, "rejected"))
|
||||
} else {
|
||||
Ok(Json(json!({
|
||||
"endpoints": [{"name": "discovered-model"}],
|
||||
"next_page_token": null,
|
||||
})))
|
||||
}
|
||||
}
|
||||
}),
|
||||
);
|
||||
tokio::spawn(async move {
|
||||
let _ = axum::serve(listener, app).await;
|
||||
});
|
||||
|
||||
let mut h = AgentHarness::spawn_provider("databricks_v2", &host, "discovered-model").await;
|
||||
let initialize = h
|
||||
.send(
|
||||
"initialize",
|
||||
json!({ "protocolVersion": 1, "clientCapabilities": {} }),
|
||||
)
|
||||
.await;
|
||||
assert!(h.recv_for(initialize).await.get("result").is_some());
|
||||
|
||||
let pid_dir = TempDir::new().unwrap();
|
||||
let pid_file = pid_dir.path().join("mcp.pid");
|
||||
let fake_mcp = env!("CARGO_BIN_EXE_fake-mcp");
|
||||
let mcp_servers = json!([{
|
||||
"name": "must-not-spawn",
|
||||
"command": fake_mcp,
|
||||
"args": [],
|
||||
"env": [{
|
||||
"name": "FAKE_MCP_PID_FILE",
|
||||
"value": pid_file.to_string_lossy(),
|
||||
}],
|
||||
}]);
|
||||
|
||||
let failed = h
|
||||
.send(
|
||||
"session/new",
|
||||
json!({ "cwd": "/tmp", "mcpServers": mcp_servers }),
|
||||
)
|
||||
.await;
|
||||
let failed_response = h.recv_for(failed).await;
|
||||
assert!(failed_response.get("error").is_some(), "{failed_response}");
|
||||
assert!(
|
||||
failed_response["error"]["message"]
|
||||
.as_str()
|
||||
.unwrap_or_default()
|
||||
.contains("llm auth"),
|
||||
"rejected static token did not retain auth semantics: {failed_response}"
|
||||
);
|
||||
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||
assert!(
|
||||
!pid_file.exists(),
|
||||
"MCP process spawned before failed discovery was resolved"
|
||||
);
|
||||
|
||||
let retry = h
|
||||
.send("session/new", json!({ "cwd": "/tmp", "mcpServers": [] }))
|
||||
.await;
|
||||
let retry_response = h.recv_for(retry).await;
|
||||
assert!(
|
||||
retry_response["result"]["sessionId"].is_string(),
|
||||
"failed discovery consumed the sole session slot: {retry_response}"
|
||||
);
|
||||
assert_eq!(attempts.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
@@ -137,6 +137,7 @@ pub async fn get_agent_models(
|
||||
&effective_provider,
|
||||
&merged_env,
|
||||
persisted_model.clone(),
|
||||
DatabricksAuthIntent::InteractiveModelPicker,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
@@ -307,9 +308,14 @@ pub async fn discover_agent_models(
|
||||
return Ok(models);
|
||||
}
|
||||
|
||||
if let Some(models) =
|
||||
discover_databricks_models(&state.http_client, &effective_provider, &merged_env, None)
|
||||
.await?
|
||||
if let Some(models) = discover_databricks_models(
|
||||
&state.http_client,
|
||||
&effective_provider,
|
||||
&merged_env,
|
||||
None,
|
||||
DatabricksAuthIntent::PassiveDraftDiscovery,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
return Ok(models);
|
||||
}
|
||||
@@ -681,97 +687,14 @@ async fn discover_anthropic_models(
|
||||
}))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Databricks model discovery (v1 + v2)
|
||||
// ---------------------------------------------------------------------------
|
||||
//
|
||||
// Delegates to buzz_agent_pkg::catalog::discover_databricks_models, which
|
||||
// acquires auth in-process via build_token_source:
|
||||
// - Static bearer (DATABRICKS_TOKEN): returned immediately.
|
||||
// - PKCE cache hit: returned from disk without a browser flow.
|
||||
// - No token, no cache: returns Err(LlmAuth) → we return Ok(None) and fall
|
||||
// through to run_agent_models_command. Never hangs, never opens a browser.
|
||||
|
||||
fn is_databricks_provider(provider: Option<&str>) -> bool {
|
||||
matches!(
|
||||
provider
|
||||
.map(str::trim)
|
||||
.map(str::to_ascii_lowercase)
|
||||
.as_deref(),
|
||||
Some("databricks" | "databricks_v2" | "databricks-v2")
|
||||
)
|
||||
}
|
||||
|
||||
fn databricks_agent_provider(provider: &str) -> buzz_agent_pkg::config::Provider {
|
||||
if provider.trim().eq_ignore_ascii_case("databricks_v2")
|
||||
|| provider.trim().eq_ignore_ascii_case("databricks-v2")
|
||||
{
|
||||
buzz_agent_pkg::config::Provider::DatabricksV2
|
||||
} else {
|
||||
buzz_agent_pkg::config::Provider::Databricks
|
||||
}
|
||||
}
|
||||
|
||||
async fn discover_databricks_models(
|
||||
_client: &reqwest::Client,
|
||||
provider: &DiscoveryProvider,
|
||||
env: &BTreeMap<String, String>,
|
||||
selected_model: Option<String>,
|
||||
) -> Result<Option<AgentModelsResponse>, String> {
|
||||
let provider_str = match provider.as_deref() {
|
||||
Some(p) if is_databricks_provider(Some(p)) => p,
|
||||
_ => return Ok(None),
|
||||
};
|
||||
|
||||
let host = match env_or_process_value(env, "DATABRICKS_HOST") {
|
||||
Some(h) => h,
|
||||
None => return Ok(None), // no host → fall through to subprocess
|
||||
};
|
||||
|
||||
// api_key = DATABRICKS_TOKEN (empty string = use PKCE cache).
|
||||
let api_key = env_or_process_value(env, "DATABRICKS_TOKEN").unwrap_or_default();
|
||||
|
||||
let agent_provider = databricks_agent_provider(provider_str);
|
||||
let cfg = buzz_agent_pkg::config::Config::for_discovery(agent_provider, api_key, host);
|
||||
|
||||
// Build a redaction env so the token never appears in surfaced errors.
|
||||
let token_for_redact = env_or_process_value(env, "DATABRICKS_TOKEN").unwrap_or_default();
|
||||
let redaction_env = redaction_env_with_value(env, "DATABRICKS_TOKEN", &token_for_redact);
|
||||
|
||||
let entries = match buzz_agent_pkg::discover_databricks_models(&cfg).await {
|
||||
Ok(e) => e,
|
||||
Err(buzz_agent_pkg::AgentError::LlmAuth(_)) => {
|
||||
// No token + no PKCE cache → fall through to subprocess.
|
||||
return Ok(None);
|
||||
}
|
||||
Err(e) => {
|
||||
let msg = crate::managed_agents::redact_env_values_in(&e.to_string(), &redaction_env);
|
||||
return Err(format!("Databricks model discovery failed: {msg}"));
|
||||
}
|
||||
};
|
||||
|
||||
if entries.is_empty() {
|
||||
return Err("Databricks model discovery returned no models".to_string());
|
||||
}
|
||||
|
||||
let models = entries
|
||||
.into_iter()
|
||||
.map(|e| AgentModelInfo {
|
||||
id: e.id,
|
||||
name: Some(e.name),
|
||||
description: None,
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(Some(AgentModelsResponse {
|
||||
agent_name: provider_str.trim().to_string(),
|
||||
agent_version: "models-api".to_string(),
|
||||
models,
|
||||
agent_default_model: None,
|
||||
selected_model,
|
||||
supports_switching: true,
|
||||
}))
|
||||
}
|
||||
#[path = "agent_models_databricks.rs"]
|
||||
mod databricks;
|
||||
#[cfg(test)]
|
||||
use databricks::{
|
||||
databricks_sign_in_required_error, databricks_static_token_error, is_databricks_provider,
|
||||
should_start_interactive_auth,
|
||||
};
|
||||
use databricks::{discover_databricks_models, DatabricksAuthIntent};
|
||||
|
||||
/// Apply an `UpdateManagedAgentRequest`'s model/provider/system_prompt patch
|
||||
/// to `record`, enforcing the linked-instance write guard: a definition-linked
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
//! Databricks v1/v2 model discovery and interactive reauthentication.
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::LazyLock;
|
||||
|
||||
use crate::commands::agent_models_env::{
|
||||
env_or_process_value, redaction_env_with_value, DiscoveryProvider,
|
||||
};
|
||||
use crate::managed_agents::AgentModelInfo;
|
||||
use crate::managed_agents::AgentModelsResponse;
|
||||
|
||||
// Model discovery can be triggered by multiple dialogs at once. Permit only one
|
||||
// callback listener/browser flow for the process-wide OAuth cache.
|
||||
static AUTH_GATE: LazyLock<tokio::sync::Mutex<()>> = LazyLock::new(|| tokio::sync::Mutex::new(()));
|
||||
|
||||
pub(super) fn is_databricks_provider(provider: Option<&str>) -> bool {
|
||||
matches!(
|
||||
provider
|
||||
.map(str::trim)
|
||||
.map(str::to_ascii_lowercase)
|
||||
.as_deref(),
|
||||
Some("databricks" | "databricks_v2" | "databricks-v2")
|
||||
)
|
||||
}
|
||||
|
||||
fn databricks_agent_provider(provider: &str) -> buzz_agent_pkg::config::Provider {
|
||||
if provider.trim().eq_ignore_ascii_case("databricks_v2")
|
||||
|| provider.trim().eq_ignore_ascii_case("databricks-v2")
|
||||
{
|
||||
buzz_agent_pkg::config::Provider::DatabricksV2
|
||||
} else {
|
||||
buzz_agent_pkg::config::Provider::Databricks
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn databricks_static_token_error(
|
||||
error: &str,
|
||||
redaction_env: &BTreeMap<String, String>,
|
||||
) -> String {
|
||||
let message = crate::managed_agents::redact_env_values_in(error, redaction_env);
|
||||
format!("Databricks rejected DATABRICKS_TOKEN; update it in agent settings: {message}")
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub(super) enum DatabricksAuthIntent {
|
||||
/// A saved agent's model picker was opened by the user.
|
||||
InteractiveModelPicker,
|
||||
/// Discovery was triggered automatically from unsaved form state.
|
||||
PassiveDraftDiscovery,
|
||||
}
|
||||
|
||||
impl DatabricksAuthIntent {
|
||||
fn allows_interactive_auth(self) -> bool {
|
||||
matches!(self, Self::InteractiveModelPicker)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn databricks_sign_in_required_error() -> String {
|
||||
"Databricks sign-in is required; save this agent, then open its model picker to sign in, or run `buzz-agent auth databricks`"
|
||||
.to_string()
|
||||
}
|
||||
|
||||
pub(super) fn should_start_interactive_auth(
|
||||
api_key: &str,
|
||||
auth_intent: DatabricksAuthIntent,
|
||||
) -> bool {
|
||||
api_key.is_empty() && auth_intent.allows_interactive_auth()
|
||||
}
|
||||
|
||||
pub(super) async fn discover_databricks_models(
|
||||
_client: &reqwest::Client,
|
||||
provider: &DiscoveryProvider,
|
||||
env: &BTreeMap<String, String>,
|
||||
selected_model: Option<String>,
|
||||
auth_intent: DatabricksAuthIntent,
|
||||
) -> Result<Option<AgentModelsResponse>, String> {
|
||||
let provider_name = match provider.as_deref() {
|
||||
Some(provider_name) if is_databricks_provider(Some(provider_name)) => provider_name,
|
||||
_ => return Ok(None),
|
||||
};
|
||||
|
||||
let host = match env_or_process_value(env, "DATABRICKS_HOST") {
|
||||
Some(host) => host,
|
||||
None => return Ok(None),
|
||||
};
|
||||
let api_key = env_or_process_value(env, "DATABRICKS_TOKEN").unwrap_or_default();
|
||||
let config = buzz_agent_pkg::config::Config::for_discovery(
|
||||
databricks_agent_provider(provider_name),
|
||||
api_key.clone(),
|
||||
host.clone(),
|
||||
);
|
||||
let redaction_env = redaction_env_with_value(env, "DATABRICKS_TOKEN", &api_key);
|
||||
|
||||
let entries = match buzz_agent_pkg::discover_databricks_models(&config).await {
|
||||
Ok(entries) => entries,
|
||||
Err(buzz_agent_pkg::AgentError::LlmAuth(_))
|
||||
if should_start_interactive_auth(&api_key, auth_intent) =>
|
||||
{
|
||||
let _auth = AUTH_GATE.lock().await;
|
||||
match buzz_agent_pkg::discover_databricks_models(&config).await {
|
||||
Ok(entries) => entries,
|
||||
Err(buzz_agent_pkg::AgentError::LlmAuth(_)) => {
|
||||
buzz_agent_pkg::authenticate_databricks(&host)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
format_redacted_error(
|
||||
"Databricks sign-in failed",
|
||||
&error,
|
||||
&redaction_env,
|
||||
)
|
||||
})?;
|
||||
buzz_agent_pkg::discover_databricks_models(&config)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
format_redacted_error(
|
||||
"Databricks model discovery failed after sign-in",
|
||||
&error,
|
||||
&redaction_env,
|
||||
)
|
||||
})?
|
||||
}
|
||||
Err(error) => {
|
||||
return Err(format_redacted_error(
|
||||
"Databricks model discovery failed",
|
||||
&error,
|
||||
&redaction_env,
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(buzz_agent_pkg::AgentError::LlmAuth(error)) if !api_key.is_empty() => {
|
||||
return Err(databricks_static_token_error(&error, &redaction_env));
|
||||
}
|
||||
Err(buzz_agent_pkg::AgentError::LlmAuth(_)) => {
|
||||
return Err(databricks_sign_in_required_error());
|
||||
}
|
||||
Err(error) => {
|
||||
return Err(format_redacted_error(
|
||||
"Databricks model discovery failed",
|
||||
&error,
|
||||
&redaction_env,
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
if entries.is_empty() {
|
||||
return Err("Databricks model discovery returned no models".to_string());
|
||||
}
|
||||
|
||||
Ok(Some(AgentModelsResponse {
|
||||
agent_name: provider_name.trim().to_string(),
|
||||
agent_version: "models-api".to_string(),
|
||||
models: entries
|
||||
.into_iter()
|
||||
.map(|entry| AgentModelInfo {
|
||||
id: entry.id,
|
||||
name: Some(entry.name),
|
||||
description: None,
|
||||
})
|
||||
.collect(),
|
||||
agent_default_model: None,
|
||||
selected_model,
|
||||
supports_switching: true,
|
||||
}))
|
||||
}
|
||||
|
||||
fn format_redacted_error(
|
||||
context: &str,
|
||||
error: &impl std::fmt::Display,
|
||||
redaction_env: &BTreeMap<String, String>,
|
||||
) -> String {
|
||||
let message = crate::managed_agents::redact_env_values_in(&error.to_string(), redaction_env);
|
||||
format!("{context}: {message}")
|
||||
}
|
||||
@@ -576,6 +576,29 @@ fn is_databricks_provider_matches_both_variants() {
|
||||
assert!(!is_databricks_provider(None));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn databricks_interactive_auth_requires_explicit_intent_and_no_static_token() {
|
||||
assert!(should_start_interactive_auth(
|
||||
"",
|
||||
DatabricksAuthIntent::InteractiveModelPicker
|
||||
));
|
||||
assert!(!should_start_interactive_auth(
|
||||
"",
|
||||
DatabricksAuthIntent::PassiveDraftDiscovery
|
||||
));
|
||||
assert!(!should_start_interactive_auth(
|
||||
"static-token",
|
||||
DatabricksAuthIntent::InteractiveModelPicker
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn databricks_passive_auth_error_has_reachable_create_flow_guidance() {
|
||||
let error = databricks_sign_in_required_error();
|
||||
assert!(error.contains("save this agent, then open its model picker"));
|
||||
assert!(error.contains("buzz-agent auth databricks"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_discovery_error_converts_dangling_sentinel_to_sentence() {
|
||||
// get_agent_models is a user-facing surface: a dangling harness must
|
||||
@@ -881,3 +904,21 @@ fn draft_agent_model_discovery_env_layers_all_three_tiers_in_order() {
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn databricks_static_token_error_redacts_echoed_token() {
|
||||
let token = "secret-databricks-token";
|
||||
let redaction_env = BTreeMap::from([("DATABRICKS_TOKEN".to_string(), token.to_string())]);
|
||||
|
||||
let error = databricks_static_token_error(
|
||||
&format!("Databricks rejected bearer {token}"),
|
||||
&redaction_env,
|
||||
);
|
||||
|
||||
assert!(error.contains("[REDACTED]"), "got: {error}");
|
||||
assert!(!error.contains(token), "token leaked in error: {error}");
|
||||
assert!(
|
||||
error.contains("update it in agent settings"),
|
||||
"error lost its remediation: {error}"
|
||||
);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user