fix: preserve model associations on refresh

This commit is contained in:
fawney19
2026-05-11 14:52:50 +08:00
parent 247ea9d1bd
commit 9057537ab8
4 changed files with 13 additions and 101 deletions

View File

@@ -549,14 +549,6 @@ mod tests {
.cloned()
.collect())
}
async fn delete_admin_provider_model(
&self,
_provider_id: &str,
_model_id: &str,
) -> Result<bool, Self::Error> {
Ok(false)
}
}
#[async_trait]

View File

@@ -300,12 +300,12 @@ async fn gateway_model_fetch_updates_key_and_syncs_provider_model_whitelist_asso
})
.await
.expect("provider models should load");
assert_eq!(provider_models.len(), 1);
assert_eq!(
provider_models[0].global_model_name.as_deref(),
Some("gpt-5")
);
assert_eq!(provider_models[0].provider_model_name, "gpt-5");
let mut provider_model_names = provider_models
.iter()
.map(|model| model.provider_model_name.as_str())
.collect::<Vec<_>>();
provider_model_names.sort_unstable();
assert_eq!(provider_model_names, vec!["gpt-4.1", "gpt-5"]);
execution_runtime_handle.abort();
}
@@ -526,11 +526,12 @@ async fn gateway_background_model_fetch_updates_key_and_syncs_provider_model_whi
})
.await
.expect("provider models should load");
assert_eq!(provider_models.len(), 1);
assert_eq!(
provider_models[0].global_model_name.as_deref(),
Some("gpt-5")
);
let mut provider_model_names = provider_models
.iter()
.map(|model| model.provider_model_name.as_str())
.collect::<Vec<_>>();
provider_model_names.sort_unstable();
assert_eq!(provider_model_names, vec!["gpt-4.1", "gpt-5"]);
let seen_plan = seen_execution_runtime_plan
.lock()

View File

@@ -236,16 +236,6 @@ impl ModelFetchAssociationStore for AppState {
.await
.map_err(|err| format!("{err:?}"))
}
async fn delete_admin_provider_model(
&self,
provider_id: &str,
model_id: &str,
) -> Result<bool, Self::Error> {
AppState::delete_admin_provider_model(self, provider_id, model_id)
.await
.map_err(|err| format!("{err:?}"))
}
}
#[async_trait]

View File

@@ -10,8 +10,6 @@ use async_trait::async_trait;
use serde_json::Value;
use uuid::Uuid;
use crate::json_string_list;
#[async_trait]
pub trait ModelFetchAssociationStore {
type Error: Send;
@@ -39,12 +37,6 @@ pub trait ModelFetchAssociationStore {
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, Self::Error>;
async fn delete_admin_provider_model(
&self,
provider_id: &str,
model_id: &str,
) -> Result<bool, Self::Error>;
}
pub async fn sync_provider_model_whitelist_associations<S>(
@@ -59,8 +51,8 @@ where
return Ok(());
}
// Key model refresh is additive: provider model associations may be curated manually.
auto_associate_provider_by_key_whitelist(state, provider_id, current_allowed_models).await?;
auto_disassociate_provider_by_key_whitelist(state, provider_id).await?;
Ok(())
}
@@ -145,65 +137,6 @@ where
Ok(())
}
async fn auto_disassociate_provider_by_key_whitelist<S>(
state: &S,
provider_id: &str,
) -> Result<(), S::Error>
where
S: ModelFetchAssociationStore + Sync + ?Sized,
{
let keys = state
.list_provider_catalog_keys_by_provider_ids(&[provider_id.to_string()])
.await?;
let active_non_oauth_keys = keys
.into_iter()
.filter(|key| key.is_active)
.filter(|key| !is_oauth_auth_type(&key.auth_type))
.collect::<Vec<_>>();
if active_non_oauth_keys.is_empty() {
return Ok(());
}
if active_non_oauth_keys
.iter()
.any(|key| key.allowed_models.is_none())
{
return Ok(());
}
let all_allowed_models = active_non_oauth_keys
.iter()
.flat_map(|key| json_string_list(key.allowed_models.as_ref()))
.collect::<BTreeSet<_>>();
let provider_models = state
.list_admin_provider_models(&AdminProviderModelListQuery {
provider_id: provider_id.to_string(),
is_active: None,
offset: 0,
limit: 10_000,
})
.await?;
for model in provider_models {
let mappings = global_model_mapping_patterns(model.global_model_config.as_ref());
if mappings.is_empty() {
continue;
}
let matched = all_allowed_models.iter().any(|allowed_model| {
mappings
.iter()
.any(|pattern| matches_model_mapping(pattern, allowed_model))
});
if matched {
continue;
}
state
.delete_admin_provider_model(provider_id, &model.id)
.await?;
}
Ok(())
}
fn global_model_mapping_patterns(config: Option<&Value>) -> Vec<String> {
config
.and_then(Value::as_object)
@@ -220,7 +153,3 @@ fn global_model_mapping_patterns(config: Option<&Value>) -> Vec<String> {
})
.unwrap_or_default()
}
fn is_oauth_auth_type(value: &str) -> bool {
matches!(value.trim().to_ascii_lowercase().as_str(), "oauth" | "kiro")
}