fix: preserve provider pins during model routing

This commit is contained in:
Bryan Helmkamp 2026-07-23 13:11:49 -04:00
parent 770d393a0c
commit 4f697c527c
No known key found for this signature in database
8 changed files with 254 additions and 63 deletions

View file

@ -3144,6 +3144,42 @@ enabled = true
assert_eq!(selected.model, "claude-fable-5");
}
#[test]
fn every_legacy_builtin_identifier_targets_an_existing_offering() {
let catalog = Catalog::from_builtin_with_overrides(&minimal_settings(
r"
[providers.bedrock]
enabled = true
[providers.bedrock-openai]
enabled = true
[providers.openrouter]
enabled = true
",
))
.expect("all providers referenced by the legacy table should build");
for (legacy_id, provider_id, canonical_id) in LEGACY_BUILTIN_MODEL_IDENTIFIERS {
let provider = ProviderId::new(*provider_id);
let model = catalog
.resolve_on_provider(&provider, legacy_id)
.unwrap_or_else(|error| {
panic!(
"legacy identifier '{legacy_id}' should resolve on '{provider}': {error}"
)
});
assert_eq!(model.provider, provider, "{legacy_id}");
assert_eq!(model.id, *canonical_id, "{legacy_id}");
assert_eq!(
legacy_builtin_model(legacy_id),
Some((provider, ModelId::new(*canonical_id))),
"{legacy_id}"
);
}
}
#[test]
fn builtin_openrouter_includes_glm_5_2_when_enabled() {
let catalog = Catalog::from_builtin_with_overrides(&minimal_settings(
@ -5118,6 +5154,61 @@ reasoning = false
);
}
#[test]
fn effective_agent_profile_is_scoped_by_provider_for_shared_model_id() {
let layer = minimal_settings(
r#"
[providers.one]
display_name = "One"
adapter = "openai"
agent_profile = "openai"
[providers.one.models.shared]
display_name = "Shared on One"
family = "test"
default = true
[providers.one.models.shared.limits]
context_window = 1000
[providers.one.models.shared.features]
tools = false
vision = false
reasoning = false
[providers.two]
display_name = "Two"
adapter = "openai"
agent_profile = "anthropic"
[providers.two.models.shared]
display_name = "Shared on Two"
family = "test"
default = true
agent_profile = "gemini"
[providers.two.models.shared.limits]
context_window = 1000
[providers.two.models.shared.features]
tools = false
vision = false
reasoning = false
"#,
);
let catalog = Catalog::from_settings(&layer).unwrap();
assert_eq!(
catalog.effective_agent_profile(&ProviderId::new("one"), Some("shared")),
Some(AgentProfileKind::OpenAi)
);
assert_eq!(
catalog.effective_agent_profile(&ProviderId::new("two"), Some("shared")),
Some(AgentProfileKind::Gemini)
);
}
#[test]
fn omitted_agent_profile_uses_adapter_default() {
let layer = minimal_settings(

View file

@ -13785,7 +13785,7 @@ async fn get_aggregate_billing_saturates_total_cost_across_models() {
agg.by_model.insert(
ModelRef {
provider: ProviderId::openai(),
model_id: model_id.to_string(),
model_id: model_id.into(),
speed: None,
},
ModelBillingTotals {

View file

@ -110,3 +110,22 @@ pub struct RunSessionTurnInterruptedProps {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::RunSessionCreatedProps;
#[test]
fn session_created_deserializes_legacy_payload_without_provider() {
let props: RunSessionCreatedProps = serde_json::from_value(json!({
"title": "Legacy session",
"model": "gpt-5.4"
}))
.unwrap();
assert_eq!(props.model.as_deref(), Some("gpt-5.4"));
assert_eq!(props.provider, None);
}
}

View file

@ -188,13 +188,33 @@ impl SessionMessage {
#[cfg(test)]
mod tests {
use chrono::Utc;
use serde_json::json;
use super::SessionStatus;
use super::{SessionId, SessionRecord, SessionStatus};
use crate::fixtures;
#[test]
fn session_status_rejects_removed_terminal_states() {
assert!(serde_json::from_value::<SessionStatus>(json!("closed")).is_err());
assert!(serde_json::from_value::<SessionStatus>(json!("deleted")).is_err());
}
#[test]
fn session_record_deserializes_legacy_json_without_provider() {
let mut value = serde_json::to_value(SessionRecord::new(
SessionId::new(),
fixtures::RUN_1,
Utc::now(),
))
.unwrap();
value
.as_object_mut()
.expect("session record should serialize as an object")
.remove("provider");
let record: SessionRecord = serde_json::from_value(value).unwrap();
assert_eq!(record.provider, None);
}
}

View file

@ -2184,7 +2184,7 @@ mod tests {
text: "ok".to_string(),
model: ModelRef {
provider: ProviderId::new("openrouter"),
model_id: "openai/gpt-5.4".to_string(),
model_id: "openai/gpt-5.4".into(),
speed: None,
},
usage: LlmTokenCounts {

View file

@ -2624,6 +2624,33 @@ reasoning = false
assert_eq!(provider.profile_kind, AgentProfileKind::Anthropic);
}
#[test]
fn api_backend_preserves_default_provider_for_legacy_model_identifier() {
let settings: LlmCatalogSettings = toml::from_str(
r"
[providers.openrouter]
enabled = true
",
)
.unwrap();
let catalog = Arc::new(Catalog::from_builtin_with_overrides(&settings).unwrap());
let backend = AgentApiBackend::new_with_catalog(
"openai/gpt-5.4".to_string(),
ProviderId::from("openrouter"),
Vec::new(),
Arc::new(EnvCredentialSource::new()),
SteeringHub::for_tests(),
catalog,
);
let provider = backend
.resolve_provider_context("openai/gpt-5.4", None)
.unwrap();
assert_eq!(provider.provider_id, ProviderId::from("openrouter"));
assert_eq!(provider.profile_kind, AgentProfileKind::OpenAi);
}
#[test]
fn run_model_controls_apply_when_node_omits_controls() {
let backend = AgentApiBackend::new_from_env(

View file

@ -58,6 +58,11 @@ pub(crate) fn resolve_provider_context(
})?
.id
.clone()
} else if catalog
.get_on_provider(default_provider_id, model)
.is_some()
{
default_provider_id.clone()
} else {
match catalog.select(model, None, &catalog.all_provider_ids()) {
Ok(model) => model.provider.clone(),

View file

@ -1462,11 +1462,11 @@ reasoning = false
}
#[tokio::test]
async fn create_materializes_shared_alias_for_ready_provider_snapshot_and_pin() {
const ALIAS_DOT: &str = r#"digraph Test {
async fn create_materializes_portable_selectors_for_ready_provider_snapshot_and_pin() {
const MODEL_DOT: &str = r#"digraph Test {
graph [goal="Test"]
start [shape=Mdiamond]
work [prompt="Do work", model="gpt-56-sol"]
work [prompt="Do work", model="MODEL_SELECTOR"]
exit [shape=Msquare]
start -> work -> exit
}"#;
@ -1490,65 +1490,94 @@ reasoning = false
),
];
for (ready, explicit_provider, expected_provider) in cases {
let dir = tempfile::tempdir().unwrap();
let mut settings = test_default_settings();
settings.run.model.name = Some("gpt-56-sol".to_string());
settings.run.model.provider = explicit_provider.map(str::to_string);
let store = memory_store();
let created = create(
store.as_ref(),
CreateRunInput {
workflow: WorkflowInput::DotSource {
source: ALIAS_DOT.to_string(),
base_dir: None,
for selector in ["gpt-56-sol", "openai/gpt-5.6-sol"] {
for (ready, explicit_provider, expected_provider) in &cases {
let dir = tempfile::tempdir().unwrap();
let mut settings = test_default_settings();
settings.run.model.name = Some(selector.to_string());
settings.run.model.provider = explicit_provider.map(str::to_string);
let store = memory_store();
let created = create(
store.as_ref(),
CreateRunInput {
workflow: WorkflowInput::DotSource {
source: MODEL_DOT.replace("MODEL_SELECTOR", selector),
base_dir: None,
},
settings,
vars: HashMap::new(),
cwd: dir.path().to_path_buf(),
workflow_slug: None,
workflow_path: None,
workflow_bundle: None,
submitted_manifest_bytes: None,
run_id: None,
title: None,
automation: None,
git: None,
fork_source_ref: None,
parent_id: None,
provenance: test_support::test_run_provenance(),
configured_providers: ready.clone(),
web_url: None,
},
settings,
vars: HashMap::new(),
cwd: dir.path().to_path_buf(),
workflow_slug: None,
workflow_path: None,
workflow_bundle: None,
submitted_manifest_bytes: None,
run_id: None,
title: None,
automation: None,
git: None,
fork_source_ref: None,
parent_id: None,
provenance: test_support::test_run_provenance(),
configured_providers: ready,
web_url: None,
},
dir.path().join("storage"),
Arc::clone(&catalog),
)
.await
.unwrap();
let run_spec = created.persisted.run_spec();
dir.path().join("storage"),
Arc::clone(&catalog),
)
.await
.unwrap();
let run_spec = created.persisted.run_spec();
assert_eq!(
run_spec.settings.run.model.name.as_deref(),
Some("gpt-5.6-sol")
);
assert_eq!(
run_spec.settings.run.model.provider.as_deref(),
Some(expected_provider.as_str())
);
assert_eq!(
run_spec.graph.nodes["work"]
.attrs
.get("model")
.and_then(AttrValue::as_str),
Some("gpt-5.6-sol")
);
assert_eq!(
run_spec.graph.nodes["work"]
.attrs
.get("provider")
.and_then(AttrValue::as_str),
Some(expected_provider.as_str())
);
assert_eq!(
run_spec.settings.run.model.name.as_deref(),
Some("gpt-5.6-sol"),
"{selector}"
);
assert_eq!(
run_spec.settings.run.model.provider.as_deref(),
Some(expected_provider.as_str()),
"{selector}"
);
assert_eq!(
run_spec.graph.nodes["work"]
.attrs
.get("model")
.and_then(AttrValue::as_str),
Some("gpt-5.6-sol"),
"{selector}"
);
assert_eq!(
run_spec.graph.nodes["work"]
.attrs
.get("provider")
.and_then(AttrValue::as_str),
Some(expected_provider.as_str()),
"{selector}"
);
let run_store = store.open_run(&created.run_id).await.unwrap();
let run_store = run_store.into();
let reloaded = Persisted::load_from_store(&run_store, &created.run_dir)
.await
.unwrap();
assert_eq!(
reloaded.run_spec().settings.run.model.provider.as_deref(),
Some(expected_provider.as_str()),
"{selector}"
);
assert_eq!(
reloaded.run_spec().graph.nodes["work"]
.attrs
.get("provider")
.and_then(AttrValue::as_str),
Some(expected_provider.as_str()),
"{selector}"
);
assert!(
reloaded.source().contains(selector),
"persisted source should preserve the user's selector '{selector}'"
);
}
}
}