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() .cloned()
.collect()) .collect())
} }
async fn delete_admin_provider_model(
&self,
_provider_id: &str,
_model_id: &str,
) -> Result<bool, Self::Error> {
Ok(false)
}
} }
#[async_trait] #[async_trait]

View File

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

View File

@@ -236,16 +236,6 @@ impl ModelFetchAssociationStore for AppState {
.await .await
.map_err(|err| format!("{err:?}")) .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] #[async_trait]

View File

@@ -10,8 +10,6 @@ use async_trait::async_trait;
use serde_json::Value; use serde_json::Value;
use uuid::Uuid; use uuid::Uuid;
use crate::json_string_list;
#[async_trait] #[async_trait]
pub trait ModelFetchAssociationStore { pub trait ModelFetchAssociationStore {
type Error: Send; type Error: Send;
@@ -39,12 +37,6 @@ pub trait ModelFetchAssociationStore {
&self, &self,
provider_ids: &[String], provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, Self::Error>; ) -> 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>( pub async fn sync_provider_model_whitelist_associations<S>(
@@ -59,8 +51,8 @@ where
return Ok(()); 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_associate_provider_by_key_whitelist(state, provider_id, current_allowed_models).await?;
auto_disassociate_provider_by_key_whitelist(state, provider_id).await?;
Ok(()) Ok(())
} }
@@ -145,65 +137,6 @@ where
Ok(()) 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> { fn global_model_mapping_patterns(config: Option<&Value>) -> Vec<String> {
config config
.and_then(Value::as_object) .and_then(Value::as_object)
@@ -220,7 +153,3 @@ fn global_model_mapping_patterns(config: Option<&Value>) -> Vec<String> {
}) })
.unwrap_or_default() .unwrap_or_default()
} }
fn is_oauth_auth_type(value: &str) -> bool {
matches!(value.trim().to_ascii_lowercase().as_str(), "oauth" | "kiro")
}