make supports_cache_control async to avoid block in place (#5362)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2025-10-28 09:30:19 -04:00
committed by GitHub
parent b295d76454
commit 8693aec764
5 changed files with 9 additions and 15 deletions
+1 -4
View File
@@ -386,18 +386,15 @@ pub trait Provider: Send + Sync {
RetryConfig::default() RetryConfig::default()
} }
/// Optional hook to fetch supported models.
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> { async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
Ok(None) Ok(None)
} }
/// Check if this provider supports embeddings
fn supports_embeddings(&self) -> bool { fn supports_embeddings(&self) -> bool {
false false
} }
/// Check if this provider supports cache control async fn supports_cache_control(&self) -> bool {
fn supports_cache_control(&self) -> bool {
false false
} }
+3 -5
View File
@@ -174,7 +174,7 @@ impl Provider for LiteLLMProvider {
&ImageFormat::OpenAi, &ImageFormat::OpenAi,
)?; )?;
if self.supports_cache_control() { if self.supports_cache_control().await {
payload = update_request_for_cache_control(&payload); payload = update_request_for_cache_control(&payload);
} }
@@ -197,10 +197,8 @@ impl Provider for LiteLLMProvider {
true true
} }
fn supports_cache_control(&self) -> bool { async fn supports_cache_control(&self) -> bool {
if let Ok(models) = tokio::task::block_in_place(|| { if let Ok(models) = self.fetch_models().await {
tokio::runtime::Handle::current().block_on(self.fetch_models())
}) {
if let Some(model_info) = models.iter().find(|m| m.name == self.model.model_name) { if let Some(model_info) = models.iter().find(|m| m.name == self.model.model_name) {
return model_info.supports_cache_control.unwrap_or(false); return model_info.supports_cache_control.unwrap_or(false);
} }
+1
View File
@@ -231,6 +231,7 @@ impl Provider for OpenAiProvider {
tracing::debug!("Failed to get usage data"); tracing::debug!("Failed to get usage data");
Usage::default() Usage::default()
}); });
let model = get_model(&json_response); let model = get_model(&json_response);
log.write(&json_response, Some(&usage))?; log.write(&json_response, Some(&usage))?;
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(model, usage)))
+4 -5
View File
@@ -193,7 +193,7 @@ fn update_request_for_anthropic(original_payload: &Value) -> Value {
payload payload
} }
fn create_request_based_on_model( async fn create_request_based_on_model(
provider: &OpenRouterProvider, provider: &OpenRouterProvider,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
@@ -207,7 +207,7 @@ fn create_request_based_on_model(
&super::utils::ImageFormat::OpenAi, &super::utils::ImageFormat::OpenAi,
)?; )?;
if provider.supports_cache_control() { if provider.supports_cache_control().await {
payload = update_request_for_anthropic(&payload); payload = update_request_for_anthropic(&payload);
} }
@@ -257,8 +257,7 @@ impl Provider for OpenRouterProvider {
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
// Create the base payload let payload = create_request_based_on_model(self, system, messages, tools).await?;
let payload = create_request_based_on_model(self, system, messages, tools)?;
let mut log = RequestLog::start(model_config, &payload)?; let mut log = RequestLog::start(model_config, &payload)?;
// Make request // Make request
@@ -357,7 +356,7 @@ impl Provider for OpenRouterProvider {
Ok(Some(models)) Ok(Some(models))
} }
fn supports_cache_control(&self) -> bool { async fn supports_cache_control(&self) -> bool {
self.model self.model
.model_name .model_name
.starts_with(OPENROUTER_MODEL_PREFIX_ANTHROPIC) .starts_with(OPENROUTER_MODEL_PREFIX_ANTHROPIC)
-1
View File
@@ -165,7 +165,6 @@ impl Provider for TetrateProvider {
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
// Create the base payload using the provided model_config
let payload = create_request( let payload = create_request(
model_config, model_config,
system, system,