refactor(server): simplify model availability probes

Use a lightweight basic probe target for preflight instead of fabricating catalog models, run configured model probes with bounded concurrency, and keep expensive model choices opt-in for defaults and live tests.
This commit is contained in:
Bryan Helmkamp 2026-05-04 11:47:37 -04:00
parent 2ef34a228e
commit 38b51c4c29
No known key found for this signature in database
8 changed files with 165 additions and 109 deletions

View file

@ -1,8 +1,8 @@
use std::sync::Arc;
use std::time::Duration;
use fabro_model::Model;
pub use fabro_model::ModelTestMode;
use fabro_model::{Model, Provider};
use strum::IntoStaticStr;
use tokio::time;
@ -54,8 +54,18 @@ pub async fn run_model_test(
}
async fn run_basic_test(info: &Model, client: Arc<Client>) -> ModelTestOutcome {
let params = GenerateParams::new(&info.id, client)
.provider(<&'static str>::from(info.provider))
run_basic_model_probe(&info.id, info.provider, client).await
}
/// Run the cheap single-prompt model availability probe without requiring a
/// catalog-backed [`Model`].
pub async fn run_basic_model_probe(
model_id: &str,
provider: Provider,
client: Arc<Client>,
) -> ModelTestOutcome {
let params = GenerateParams::new(model_id, client)
.provider(<&'static str>::from(provider))
.prompt("Say OK")
.max_tokens(16);

View file

@ -210,10 +210,18 @@ fn provider_error_from_openai_error_json(error: &serde_json::Value) -> Error {
.map_or_else(|| "OpenAI stream error".to_string(), str::to_string);
let kind = match classifier {
Some("insufficient_quota") => ProviderErrorKind::QuotaExceeded,
Some("rate_limit_exceeded") => ProviderErrorKind::RateLimit,
Some("invalid_api_key" | "invalid_authentication") => ProviderErrorKind::Authentication,
Some("account_deactivated" | "permission_denied") => ProviderErrorKind::AccessDenied,
Some("insufficient_quota" | "billing_hard_limit_reached") => {
ProviderErrorKind::QuotaExceeded
}
Some("rate_limit_error" | "rate_limit_exceeded" | "too_many_requests") => {
ProviderErrorKind::RateLimit
}
Some("authentication_error" | "invalid_api_key" | "invalid_authentication") => {
ProviderErrorKind::Authentication
}
Some(
"access_denied" | "account_deactivated" | "permission_denied" | "permission_error",
) => ProviderErrorKind::AccessDenied,
Some("content_filter" | "content_policy_violation") => ProviderErrorKind::ContentFilter,
Some("context_length_exceeded") => ProviderErrorKind::ContextLength,
Some("server_error" | "internal_error" | "service_unavailable" | "engine_overloaded") => {
@ -1850,6 +1858,29 @@ mod tests {
}
}
#[test]
fn error_event_with_rate_limit_error_returns_rate_limit() {
let mut state = empty_sse_state();
let data = r#"{
"type": "error",
"error": {
"type": "rate_limit_error",
"message": "Too many requests."
}
}"#;
let err = process_sse_event(&mut state, Some("error"), data)
.expect_err("error event should fail the stream");
match err {
Error::Provider { kind, detail } => {
assert_eq!(kind, ProviderErrorKind::RateLimit);
assert_eq!(detail.error_code.as_deref(), Some("rate_limit_error"));
}
other => panic!("expected provider error, got {other:?}"),
}
}
#[test]
fn error_event_with_unknown_invalid_prefix_returns_invalid_request() {
let mut state = empty_sse_state();

View file

@ -97,9 +97,10 @@ async fn openai_gpt_5_5_complete() {
assert_eq!(response.provider, "openai");
}
#[fabro_macros::e2e_test(live("OPENAI_API_KEY"))]
#[fabro_macros::e2e_test(live("OPENAI_GPT_5_5_PRO_API_KEY"))]
async fn openai_gpt_5_5_pro_complete() {
let api_key = std::env::var(EnvVars::OPENAI_API_KEY).expect("OPENAI_API_KEY must be set");
let api_key = std::env::var("OPENAI_GPT_5_5_PRO_API_KEY")
.expect("OPENAI_GPT_5_5_PRO_API_KEY must be set");
let adapter = OpenAiAdapter::new(api_key);
let request = Request {
temperature: None,

View file

@ -14,8 +14,7 @@
"cache_input_cost_per_mtok": 0.50
},
"estimated_output_tps": 25,
"aliases": ["opus", "claude-opus"],
"default": true
"aliases": ["opus", "claude-opus"]
},
{
"id": "claude-opus-4-6",
@ -66,7 +65,8 @@
"cache_input_cost_per_mtok": 0.30
},
"estimated_output_tps": 50,
"aliases": ["sonnet", "claude-sonnet"]
"aliases": ["sonnet", "claude-sonnet"],
"default": true
},
{
"id": "claude-haiku-4-5",
@ -185,7 +185,8 @@
"cache_input_cost_per_mtok": 0.25
},
"estimated_output_tps": 70,
"aliases": ["gpt54", "gpt-54"]
"aliases": ["gpt54", "gpt-54"],
"default": true
},
{
"id": "gpt-5.5",
@ -202,8 +203,7 @@
"cache_input_cost_per_mtok": 0.50
},
"estimated_output_tps": 70,
"aliases": ["gpt55", "gpt-55"],
"default": true
"aliases": ["gpt55", "gpt-55"]
},
{
"id": "gpt-5.5-pro",

View file

@ -99,6 +99,7 @@ impl Catalog {
#[must_use]
pub fn probe_for_provider(&self, p: Provider) -> Option<&Model> {
let override_id: Option<&str> = match p {
Provider::Anthropic => Some("claude-haiku-4-5"),
Provider::OpenAi => Some("gpt-5.4-mini"),
_ => None,
};
@ -226,13 +227,13 @@ mod tests {
let m = Catalog::builtin()
.default_for_provider(Provider::Anthropic)
.unwrap();
assert_eq!(m.id, "claude-opus-4-7");
assert_eq!(m.id, "claude-sonnet-4-6");
assert!(m.default);
let m = Catalog::builtin()
.default_for_provider(Provider::OpenAi)
.unwrap();
assert_eq!(m.id, "gpt-5.5");
assert_eq!(m.id, "gpt-5.4");
let m = Catalog::builtin()
.default_for_provider(Provider::Gemini)
@ -249,11 +250,11 @@ mod tests {
}
#[test]
fn builtin_probe_anthropic_returns_default() {
fn builtin_probe_anthropic_returns_override() {
let m = Catalog::builtin()
.probe_for_provider(Provider::Anthropic)
.unwrap();
assert_eq!(m.id, "claude-opus-4-7");
assert_eq!(m.id, "claude-haiku-4-5");
}
#[test]
@ -730,7 +731,7 @@ mod tests {
"gpt54",
"gpt-54",
],
default: false,
default: true,
configured: false,
}
"#);

View file

@ -152,7 +152,7 @@ mod tests {
assert_eq!(info.cache_input_cost_per_mtok(), Some(0.5));
assert_eq!(info.estimated_output_tps(), Some(25.0));
assert!(!info.aliases().is_empty());
assert!(info.is_default());
assert!(!info.is_default());
}
#[test]

View file

@ -106,8 +106,8 @@ async fn check_llm_providers(state: &AppState) -> CheckResult {
for (provider, issue) in &result.auth_issues {
let message = auth_issue_message(*provider, issue);
failures.push(ProviderFailure {
provider: *provider,
short: short_error_line(&message),
provider: *provider,
summary_line: short_error_line(&message),
});
details.push(CheckDetail::new(message));
}
@ -135,14 +135,14 @@ async fn check_llm_providers(state: &AppState) -> CheckResult {
let rendered = collect_chain(&err).join(": ");
failures.push(ProviderFailure {
provider,
short: short_error_line(&rendered),
summary_line: short_error_line(&rendered),
});
details.push(CheckDetail::new(format!("{provider}: {rendered}")));
}
Err(_) => {
failures.push(ProviderFailure {
provider,
short: "timeout (30s)".to_string(),
summary_line: "timeout (30s)".to_string(),
});
details.push(CheckDetail::new(format!("{provider}: timeout (30s)")));
}
@ -166,7 +166,7 @@ async fn check_llm_providers(state: &AppState) -> CheckResult {
};
let remediation = failures
.iter()
.map(|f| format!("{}: {}", f.provider, f.short))
.map(|f| format!("{}: {}", f.provider, f.summary_line))
.collect::<Vec<_>>()
.join("; ");
@ -180,8 +180,8 @@ async fn check_llm_providers(state: &AppState) -> CheckResult {
}
struct ProviderFailure {
provider: Provider,
short: String,
provider: Provider,
summary_line: String,
}
const MAX_SHORT_LEN: usize = 120;
@ -194,7 +194,7 @@ fn short_error_line(rendered: &str) -> String {
.unwrap_or("error");
if first.chars().count() > MAX_SHORT_LEN {
let cutoff: String = first.chars().take(MAX_SHORT_LEN).collect();
format!("{cutoff}")
format!("{cutoff}...")
} else {
first.to_string()
}
@ -667,10 +667,10 @@ mod tests {
}
#[test]
fn short_error_line_truncates_long_input_with_ellipsis() {
fn short_error_line_truncates_long_input_with_ascii_ellipsis() {
let input = "a".repeat(MAX_SHORT_LEN + 50);
let result = short_error_line(&input);
let expected = format!("{}", "a".repeat(MAX_SHORT_LEN));
let expected = format!("{}...", "a".repeat(MAX_SHORT_LEN));
assert_eq!(result, expected);
}

View file

@ -15,8 +15,8 @@ use fabro_config::{
use fabro_graphviz::graph::{Graph, is_llm_handler_type};
use fabro_graphviz::render::apply_direction;
use fabro_llm::Provider;
use fabro_llm::model_test::{ModelTestMode, ModelTestStatus, run_model_test};
use fabro_model::{Catalog, Model, ModelCosts, ModelFeatures, ModelLimits};
use fabro_llm::model_test::{ModelTestStatus, run_basic_model_probe};
use fabro_model::Catalog;
use fabro_sandbox::config::{
DaytonaNetwork, DaytonaSnapshotSettings, DockerfileSource as SandboxDockerfileSource,
};
@ -39,6 +39,7 @@ use fabro_workflow::pipeline::Validated;
use fabro_workflow::run_materialization::materialize_run;
use fabro_workflow::workflow_bundle::{BundledWorkflow, ParsedWorkflowConfig, WorkflowBundle};
use fabro_workflow::{Error as WorkflowError, ManifestPath};
use futures_util::stream::{self, StreamExt};
use tokio::process::Command;
use tokio::time;
@ -912,6 +913,15 @@ async fn run_sandbox_check(
}
}
const MODEL_PREFLIGHT_PROBE_CONCURRENCY: usize = 4;
struct PendingModelProbe {
index: usize,
model_id: String,
provider_name: String,
provider: Provider,
}
async fn run_llm_check(
state: &AppState,
checks: &mut Vec<CheckResult>,
@ -971,55 +981,50 @@ async fn run_llm_check(
}
let mut all_ok = true;
for (model_id, provider_name) in &model_providers {
let mut completed_checks: Vec<(usize, CheckResult)> = Vec::new();
let mut pending_probes = Vec::new();
for (index, (model_id, provider_name)) in model_providers.iter().enumerate() {
match provider_name.parse::<Provider>() {
Ok(provider) => {
let mut status = CheckStatus::Pass;
let remediation = if let Some((_, issue)) = auth_issues
if let Some((_, issue)) = auth_issues
.iter()
.find(|(candidate, _)| *candidate == provider)
{
status = CheckStatus::Warning;
all_ok = false;
Some(auth_issue_message(provider, issue))
completed_checks.push((index, CheckResult {
name: "LLM".into(),
status: CheckStatus::Warning,
summary: model_id.clone(),
details: vec![CheckDetail::new(format!(
"Provider: {provider_name}"
))],
remediation: Some(auth_issue_message(provider, issue)),
}));
} else if !configured.iter().any(|name| name == provider_name) {
status = CheckStatus::Warning;
all_ok = false;
Some(format!("Provider \"{provider_name}\" is not configured"))
completed_checks.push((index, CheckResult {
name: "LLM".into(),
status: CheckStatus::Warning,
summary: model_id.clone(),
details: vec![CheckDetail::new(format!(
"Provider: {provider_name}"
))],
remediation: Some(format!(
"Provider \"{provider_name}\" is not configured"
)),
}));
} else {
let probe_model = preflight_probe_model(model_id, provider);
let outcome = run_model_test(
&probe_model,
ModelTestMode::Basic,
Arc::clone(&client),
)
.await;
if outcome.status == ModelTestStatus::Ok {
None
} else {
status = CheckStatus::Error;
all_ok = false;
Some(format!(
"Model availability probe failed: {}",
outcome
.error_message
.unwrap_or_else(|| "unknown error".to_string())
))
}
};
checks.push(CheckResult {
name: "LLM".into(),
status,
summary: model_id.clone(),
details: vec![
CheckDetail::new(format!("Provider: {provider_name}")),
CheckDetail::new("Probe: basic generation".to_string()),
],
remediation,
});
pending_probes.push(PendingModelProbe {
index,
model_id: model_id.clone(),
provider_name: provider_name.clone(),
provider,
});
}
}
Err(err) => {
checks.push(CheckResult {
all_ok = false;
completed_checks.push((index, CheckResult {
name: "LLM".into(),
status: CheckStatus::Error,
summary: model_id.clone(),
@ -1029,11 +1034,55 @@ async fn run_llm_check(
remediation: Some(format!(
"Invalid provider \"{provider_name}\": {err}"
)),
});
all_ok = false;
}));
}
}
}
let mut probe_checks = stream::iter(pending_probes)
.map(|probe| {
let client = Arc::clone(&client);
async move {
let outcome =
run_basic_model_probe(&probe.model_id, probe.provider, client).await;
let (status, remediation) = if outcome.status == ModelTestStatus::Ok {
(CheckStatus::Pass, None)
} else {
(
CheckStatus::Error,
Some(format!(
"Model availability probe failed: {}",
outcome
.error_message
.unwrap_or_else(|| "unknown error".to_string())
)),
)
};
(probe.index, CheckResult {
name: "LLM".into(),
status,
summary: probe.model_id,
details: vec![
CheckDetail::new(format!("Provider: {}", probe.provider_name)),
CheckDetail::new("Probe: basic generation".to_string()),
],
remediation,
})
}
})
.buffer_unordered(MODEL_PREFLIGHT_PROBE_CONCURRENCY)
.collect::<Vec<_>>()
.await;
if probe_checks
.iter()
.any(|(_, check)| check.status != CheckStatus::Pass)
{
all_ok = false;
}
completed_checks.append(&mut probe_checks);
completed_checks.sort_by_key(|(index, _)| *index);
checks.extend(completed_checks.into_iter().map(|(_, check)| check));
all_ok
}
Err(err) => {
@ -1049,42 +1098,6 @@ async fn run_llm_check(
}
}
fn preflight_probe_model(model_id: &str, provider: Provider) -> Model {
if let Some(info) = Catalog::builtin().get(model_id) {
let mut model = info.clone();
model.provider = provider;
return model;
}
Model {
id: model_id.to_string(),
provider,
family: "custom".to_string(),
display_name: model_id.to_string(),
limits: ModelLimits {
context_window: 0,
max_output: None,
},
training: None,
knowledge_cutoff: None,
features: ModelFeatures {
tools: false,
vision: false,
reasoning: false,
effort: false,
},
costs: ModelCosts {
input_cost_per_mtok: None,
output_cost_per_mtok: None,
cache_input_cost_per_mtok: None,
},
estimated_output_tps: None,
aliases: Vec::new(),
default: false,
configured: true,
}
}
fn resolve_model_provider(
settings: &RunNamespace,
_graph: &Graph,