Feat: Let providers configure a fast model for summarization (#4228)
This commit is contained in:
@@ -37,6 +37,7 @@ const DEFAULT_SCOPES: &[&str] = &["all-apis", "offline_access"];
|
||||
const DEFAULT_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
pub const DATABRICKS_DEFAULT_MODEL: &str = "databricks-claude-3-7-sonnet";
|
||||
const DATABRICKS_DEFAULT_FAST_MODEL: &str = "gemini-1-5-flash";
|
||||
pub const DATABRICKS_KNOWN_MODELS: &[&str] = &[
|
||||
"databricks-meta-llama-3-3-70b-instruct",
|
||||
"databricks-meta-llama-3-1-405b-instruct",
|
||||
@@ -137,13 +138,41 @@ impl DatabricksProvider {
|
||||
let api_client =
|
||||
ApiClient::with_timeout(host, auth_method, Duration::from_secs(DEFAULT_TIMEOUT_SECS))?;
|
||||
|
||||
Ok(Self {
|
||||
// Create the provider without the fast model first
|
||||
let mut provider = Self {
|
||||
api_client,
|
||||
auth,
|
||||
model,
|
||||
model: model.clone(),
|
||||
image_format: ImageFormat::OpenAi,
|
||||
retry_config,
|
||||
})
|
||||
};
|
||||
|
||||
// Check if the default fast model exists in the workspace
|
||||
let model_with_fast = tokio::task::block_in_place(|| {
|
||||
tokio::runtime::Handle::current().block_on(async {
|
||||
if let Ok(Some(models)) = provider.fetch_supported_models().await {
|
||||
if models.contains(&DATABRICKS_DEFAULT_FAST_MODEL.to_string()) {
|
||||
tracing::debug!(
|
||||
"Found {} in Databricks workspace, setting as fast model",
|
||||
DATABRICKS_DEFAULT_FAST_MODEL
|
||||
);
|
||||
model.with_fast(DATABRICKS_DEFAULT_FAST_MODEL.to_string())
|
||||
} else {
|
||||
tracing::debug!(
|
||||
"{} not found in Databricks workspace, not setting fast model",
|
||||
DATABRICKS_DEFAULT_FAST_MODEL
|
||||
);
|
||||
model
|
||||
}
|
||||
} else {
|
||||
tracing::debug!("Could not fetch Databricks models, not setting fast model");
|
||||
model
|
||||
}
|
||||
})
|
||||
});
|
||||
|
||||
provider.model = model_with_fast;
|
||||
Ok(provider)
|
||||
}
|
||||
|
||||
fn load_retry_config(config: &crate::config::Config) -> RetryConfig {
|
||||
@@ -195,17 +224,18 @@ impl DatabricksProvider {
|
||||
})
|
||||
}
|
||||
|
||||
fn get_endpoint_path(&self, is_embedding: bool) -> String {
|
||||
fn get_endpoint_path(&self, model_name: &str, is_embedding: bool) -> String {
|
||||
if is_embedding {
|
||||
"serving-endpoints/text-embedding-3-small/invocations".to_string()
|
||||
} else {
|
||||
format!("serving-endpoints/{}/invocations", self.model.model_name)
|
||||
format!("serving-endpoints/{}/invocations", model_name)
|
||||
}
|
||||
}
|
||||
|
||||
async fn post(&self, payload: Value) -> Result<Value, ProviderError> {
|
||||
async fn post(&self, payload: Value, model_name: Option<&str>) -> Result<Value, ProviderError> {
|
||||
let is_embedding = payload.get("input").is_some() && payload.get("messages").is_none();
|
||||
let path = self.get_endpoint_path(is_embedding);
|
||||
let model_to_use = model_name.unwrap_or(&self.model.model_name);
|
||||
let path = self.get_endpoint_path(model_to_use, is_embedding);
|
||||
|
||||
let response = self.api_client.response_post(&path, &payload).await?;
|
||||
handle_response_openai_compat(response).await
|
||||
@@ -238,32 +268,36 @@ impl Provider for DatabricksProvider {
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
skip(self, system, messages, tools),
|
||||
skip(self, model_config, system, messages, tools),
|
||||
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
||||
)]
|
||||
async fn complete(
|
||||
async fn complete_with_model(
|
||||
&self,
|
||||
model_config: &ModelConfig,
|
||||
system: &str,
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> Result<(Message, ProviderUsage), ProviderError> {
|
||||
let mut payload = create_request(&self.model, system, messages, tools, &self.image_format)?;
|
||||
let mut payload =
|
||||
create_request(model_config, system, messages, tools, &self.image_format)?;
|
||||
payload
|
||||
.as_object_mut()
|
||||
.expect("payload should have model key")
|
||||
.remove("model");
|
||||
|
||||
let response = self.with_retry(|| self.post(payload.clone())).await?;
|
||||
let response = self
|
||||
.with_retry(|| self.post(payload.clone(), Some(&model_config.model_name)))
|
||||
.await?;
|
||||
|
||||
let message = response_to_message(&response)?;
|
||||
let usage = response.get("usage").map(get_usage).unwrap_or_else(|| {
|
||||
tracing::debug!("Failed to get usage data");
|
||||
Usage::default()
|
||||
});
|
||||
let model = get_model(&response);
|
||||
let response_model = get_model(&response);
|
||||
super::utils::emit_debug_trace(&self.model, &payload, &response, &usage);
|
||||
|
||||
Ok((message, ProviderUsage::new(model, usage)))
|
||||
Ok((message, ProviderUsage::new(response_model, usage)))
|
||||
}
|
||||
|
||||
async fn stream(
|
||||
@@ -272,7 +306,10 @@ impl Provider for DatabricksProvider {
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> Result<MessageStream, ProviderError> {
|
||||
let mut payload = create_request(&self.model, system, messages, tools, &self.image_format)?;
|
||||
let model_config = self.model.clone();
|
||||
|
||||
let mut payload =
|
||||
create_request(&model_config, system, messages, tools, &self.image_format)?;
|
||||
payload
|
||||
.as_object_mut()
|
||||
.expect("payload should have model key")
|
||||
@@ -283,7 +320,7 @@ impl Provider for DatabricksProvider {
|
||||
.unwrap()
|
||||
.insert("stream".to_string(), Value::Bool(true));
|
||||
|
||||
let path = self.get_endpoint_path(false);
|
||||
let path = self.get_endpoint_path(&model_config.model_name, false);
|
||||
let response = self
|
||||
.with_retry(|| async {
|
||||
let resp = self.api_client.response_post(&path, &payload).await?;
|
||||
@@ -299,8 +336,8 @@ impl Provider for DatabricksProvider {
|
||||
.await?;
|
||||
|
||||
let stream = response.bytes_stream().map_err(io::Error::other);
|
||||
let model_config = self.model.clone();
|
||||
|
||||
let model = self.model.clone();
|
||||
Ok(Box::pin(try_stream! {
|
||||
let stream_reader = StreamReader::new(stream);
|
||||
let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from);
|
||||
@@ -309,7 +346,7 @@ impl Provider for DatabricksProvider {
|
||||
pin!(message_stream);
|
||||
while let Some(message) = message_stream.next().await {
|
||||
let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?;
|
||||
super::utils::emit_debug_trace(&model_config, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default());
|
||||
super::utils::emit_debug_trace(&model, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default());
|
||||
yield (message, usage);
|
||||
}
|
||||
}))
|
||||
@@ -408,7 +445,7 @@ impl EmbeddingCapable for DatabricksProvider {
|
||||
"input": texts,
|
||||
});
|
||||
|
||||
let response = self.with_retry(|| self.post(request.clone())).await?;
|
||||
let response = self.with_retry(|| self.post(request.clone(), None)).await?;
|
||||
|
||||
let embeddings = response["data"]
|
||||
.as_array()
|
||||
|
||||
Reference in New Issue
Block a user