make supports_cache_control async to avoid block in place (#5362)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)))
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user