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:
Alice Hau
2025-05-28 23:23:02 -04:00
committed by GitHub
parent cf7bb08ee1
commit bf1c0d51e4
23 changed files with 3661 additions and 301 deletions
+12
View File
@@ -183,6 +183,18 @@ pub trait Provider: Send + Sync {
async fn fetch_supported_models_async(&self) -> Result<Option<Vec<String>>, ProviderError> {
Ok(None)
}
/// Check if this provider supports embeddings
fn supports_embeddings(&self) -> bool {
false
}
/// Create embeddings if supported. Default implementation returns an error.
async fn create_embeddings(&self, _texts: Vec<String>) -> Result<Vec<Vec<f32>>, ProviderError> {
Err(ProviderError::ExecutionError(
"This provider does not support embeddings".to_string(),
))
}
}
#[cfg(test)]
+64 -10
View File
@@ -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)
}
}
+24
View File
@@ -0,0 +1,24 @@
use anyhow::Result;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbeddingRequest {
pub input: Vec<String>,
pub model: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbeddingResponse {
pub data: Vec<EmbeddingData>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbeddingData {
pub embedding: Vec<f32>,
}
#[async_trait]
pub trait EmbeddingCapable {
async fn create_embeddings(&self, texts: Vec<String>) -> Result<Vec<Vec<f32>>>;
}
+1
View File
@@ -4,6 +4,7 @@ pub mod azureauth;
pub mod base;
pub mod bedrock;
pub mod databricks;
pub mod embedding;
pub mod errors;
mod factory;
pub mod formats;
+85 -12
View File
@@ -6,6 +6,7 @@ use std::collections::HashMap;
use std::time::Duration;
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
use super::embedding::{EmbeddingCapable, EmbeddingRequest, EmbeddingResponse};
use super::errors::ProviderError;
use super::formats::openai::{create_request, get_usage, response_to_message};
use super::utils::{emit_debug_trace, get_model, handle_response_openai_compat, ImageFormat};
@@ -80,18 +81,8 @@ impl OpenAiProvider {
})
}
async fn post(&self, payload: Value) -> Result<Value, ProviderError> {
let base_url = url::Url::parse(&self.host)
.map_err(|e| ProviderError::RequestFailed(format!("Invalid base URL: {e}")))?;
let url = base_url.join(&self.base_path).map_err(|e| {
ProviderError::RequestFailed(format!("Failed to construct endpoint URL: {e}"))
})?;
let mut request = self
.client
.post(url)
.header("Authorization", format!("Bearer {}", self.api_key));
/// Helper function to add OpenAI-specific headers to a request
fn add_headers(&self, mut request: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
// Add organization header if present
if let Some(org) = &self.organization {
request = request.header("OpenAI-Organization", org);
@@ -102,12 +93,30 @@ impl OpenAiProvider {
request = request.header("OpenAI-Project", project);
}
// Add custom headers if present
if let Some(custom_headers) = &self.custom_headers {
for (key, value) in custom_headers {
request = request.header(key, value);
}
}
request
}
async fn post(&self, payload: Value) -> Result<Value, ProviderError> {
let base_url = url::Url::parse(&self.host)
.map_err(|e| ProviderError::RequestFailed(format!("Invalid base URL: {e}")))?;
let url = base_url.join(&self.base_path).map_err(|e| {
ProviderError::RequestFailed(format!("Failed to construct endpoint URL: {e}"))
})?;
let request = self
.client
.post(url)
.header("Authorization", format!("Bearer {}", self.api_key));
let request = self.add_headers(request);
let response = request.json(&payload).send().await?;
handle_response_openai_compat(response).await
@@ -209,6 +218,16 @@ impl Provider for OpenAiProvider {
models.sort();
Ok(Some(models))
}
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()))
}
}
fn parse_custom_headers(s: String) -> HashMap<String, String> {
@@ -221,3 +240,57 @@ fn parse_custom_headers(s: String) -> HashMap<String, String> {
})
.collect()
}
#[async_trait]
impl EmbeddingCapable for OpenAiProvider {
async fn create_embeddings(&self, texts: Vec<String>) -> Result<Vec<Vec<f32>>> {
if texts.is_empty() {
return Ok(vec![]);
}
// Get embedding model from env var or use default
let embedding_model = std::env::var("EMBEDDING_MODEL")
.unwrap_or_else(|_| "text-embedding-3-small".to_string());
let request = EmbeddingRequest {
input: texts,
model: embedding_model,
};
// Construct embeddings endpoint URL
let base_url =
url::Url::parse(&self.host).map_err(|e| anyhow::anyhow!("Invalid base URL: {e}"))?;
let url = base_url
.join("v1/embeddings")
.map_err(|e| anyhow::anyhow!("Failed to construct embeddings URL: {e}"))?;
let req = self
.client
.post(url)
.header("Authorization", format!("Bearer {}", self.api_key))
.json(&request);
let req = self.add_headers(req);
let response = req
.send()
.await
.map_err(|e| anyhow::anyhow!("Failed to send embedding request: {e}"))?;
if !response.status().is_success() {
let error_text = response.text().await.unwrap_or_default();
return Err(anyhow::anyhow!("Embedding API error: {}", error_text));
}
let embedding_response: EmbeddingResponse = response
.json()
.await
.map_err(|e| anyhow::anyhow!("Failed to parse embedding response: {e}"))?;
Ok(embedding_response
.data
.into_iter()
.map(|d| d.embedding)
.collect())
}
}