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

@@ -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")
}