feat: lancedb vector tool selection (#2654)
Co-authored-by: Wendy Tang <wendytang@squareup.com> Co-authored-by: Alice Hau <ahau@squareup.com>
This commit is contained in:
@@ -1,11 +1,5 @@
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use reqwest::{Client, StatusCode};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use std::time::Duration;
|
||||
|
||||
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::embedding::EmbeddingCapable;
|
||||
use super::errors::ProviderError;
|
||||
use super::formats::databricks::{create_request, get_usage, response_to_message};
|
||||
use super::oauth;
|
||||
@@ -14,8 +8,16 @@ use crate::config::ConfigError;
|
||||
use crate::message::Message;
|
||||
use crate::model::ModelConfig;
|
||||
use mcp_core::tool::Tool;
|
||||
use serde_json::json;
|
||||
use url::Url;
|
||||
|
||||
use anyhow::Result;
|
||||
use async_trait::async_trait;
|
||||
use reqwest::{Client, StatusCode};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use std::time::Duration;
|
||||
|
||||
const DEFAULT_CLIENT_ID: &str = "databricks-cli";
|
||||
const DEFAULT_REDIRECT_URL: &str = "http://localhost:8020";
|
||||
// "offline_access" scope is used to request an OAuth 2.0 Refresh Token
|
||||
@@ -128,7 +130,6 @@ impl DatabricksProvider {
|
||||
///
|
||||
/// * `host` - The Databricks host URL
|
||||
/// * `token` - The Databricks API token
|
||||
/// * `model` - The model configuration
|
||||
///
|
||||
/// # Returns
|
||||
///
|
||||
@@ -166,7 +167,17 @@ impl DatabricksProvider {
|
||||
async fn post(&self, payload: Value) -> Result<Value, ProviderError> {
|
||||
let base_url = Url::parse(&self.host)
|
||||
.map_err(|e| ProviderError::RequestFailed(format!("Invalid base URL: {e}")))?;
|
||||
let path = format!("serving-endpoints/{}/invocations", self.model.model_name);
|
||||
|
||||
// Check if this is an embedding request by looking at the payload structure
|
||||
let is_embedding = payload.get("input").is_some() && payload.get("messages").is_none();
|
||||
let path = if is_embedding {
|
||||
// For embeddings, use the embeddings endpoint
|
||||
format!("serving-endpoints/{}/invocations", "text-embedding-3-small")
|
||||
} else {
|
||||
// For chat completions, use the model name in the path
|
||||
format!("serving-endpoints/{}/invocations", self.model.model_name)
|
||||
};
|
||||
|
||||
let url = base_url.join(&path).map_err(|e| {
|
||||
ProviderError::RequestFailed(format!("Failed to construct endpoint URL: {e}"))
|
||||
})?;
|
||||
@@ -184,7 +195,7 @@ impl DatabricksProvider {
|
||||
let payload: Option<Value> = response.json().await.ok();
|
||||
|
||||
match status {
|
||||
StatusCode::OK => payload.ok_or_else( || ProviderError::RequestFailed("Response body is not valid JSON".to_string()) ),
|
||||
StatusCode::OK => payload.ok_or_else(|| ProviderError::RequestFailed("Response body is not valid JSON".to_string())),
|
||||
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => {
|
||||
Err(ProviderError::Authentication(format!("Authentication failed. Please ensure your API keys are valid and have the required permissions. \
|
||||
Status: {}. Response: {:?}", status, payload)))
|
||||
@@ -295,4 +306,47 @@ impl Provider for DatabricksProvider {
|
||||
|
||||
Ok((message, ProviderUsage::new(model, usage)))
|
||||
}
|
||||
|
||||
fn supports_embeddings(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
async fn create_embeddings(&self, texts: Vec<String>) -> Result<Vec<Vec<f32>>, ProviderError> {
|
||||
EmbeddingCapable::create_embeddings(self, texts)
|
||||
.await
|
||||
.map_err(|e| ProviderError::ExecutionError(e.to_string()))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl EmbeddingCapable for DatabricksProvider {
|
||||
async fn create_embeddings(&self, texts: Vec<String>) -> Result<Vec<Vec<f32>>> {
|
||||
if texts.is_empty() {
|
||||
return Ok(vec![]);
|
||||
}
|
||||
|
||||
// Create request in Databricks format for embeddings
|
||||
let request = json!({
|
||||
"input": texts,
|
||||
});
|
||||
|
||||
let response = self.post(request).await?;
|
||||
|
||||
let embeddings = response["data"]
|
||||
.as_array()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid response format: missing data array"))?
|
||||
.iter()
|
||||
.map(|item| {
|
||||
item["embedding"]
|
||||
.as_array()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid embedding format"))?
|
||||
.iter()
|
||||
.map(|v| v.as_f64().map(|f| f as f32))
|
||||
.collect::<Option<Vec<f32>>>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid embedding values"))
|
||||
})
|
||||
.collect::<Result<Vec<Vec<f32>>>>()?;
|
||||
|
||||
Ok(embeddings)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user