#[cfg(feature = "providers")]
use crate::core::Secret;
use crate::core::StoreError;
use crate::memory::Embedder;
#[allow(clippy::cast_possible_truncation)]
fn json_f32(value: &serde_json::Value) -> Option<f32> {
let narrowed = value.as_f64()? as f32;
narrowed.is_finite().then_some(narrowed)
}
#[cfg(feature = "providers")]
#[derive(Debug, Clone)]
pub struct OpenAiEmbedder {
http: reqwest::Client,
key: Option<Secret>,
base: String,
model: String,
dimensions: Option<u32>,
input_type: Option<String>,
egress: Option<crate::core::Egress>,
timeout: std::time::Duration,
}
#[cfg(feature = "providers")]
impl OpenAiEmbedder {
pub const DEFAULT_TIMEOUT: std::time::Duration = std::time::Duration::from_mins(5);
pub fn new(model: impl Into<String>) -> Result<Self, StoreError> {
let http = reqwest::Client::builder()
.build()
.map_err(|e| StoreError::Backend(format!("could not build an HTTP client: {e}")))?;
Ok(Self {
http,
key: None,
base: "https://api.openai.com".to_owned(),
model: model.into(),
dimensions: None,
input_type: None,
egress: None,
timeout: Self::DEFAULT_TIMEOUT,
})
}
#[must_use]
pub fn key(mut self, key: impl Into<String>) -> Self {
self.key = Some(Secret::new(key));
self
}
#[must_use]
pub fn base(mut self, base: impl Into<String>) -> Self {
self.base = base.into();
self
}
#[must_use]
pub const fn dimensions(mut self, dimensions: u32) -> Self {
self.dimensions = Some(dimensions);
self
}
#[must_use]
pub fn input_type(mut self, input_type: impl Into<String>) -> Self {
self.input_type = Some(input_type.into());
self
}
#[must_use]
pub fn egress(mut self, egress: crate::core::Egress) -> Self {
self.egress = Some(egress);
self
}
#[must_use]
pub const fn timeout(mut self, timeout: std::time::Duration) -> Self {
self.timeout = timeout;
self
}
fn check_egress(&self) -> Result<(), StoreError> {
let Some(egress) = &self.egress else {
return Ok(());
};
let host = reqwest::Url::parse(&self.base)
.ok()
.and_then(|u| u.host_str().map(ToOwned::to_owned));
egress
.permits(host.as_deref())
.map_err(|e| StoreError::Backend(e.to_string()))
}
}
#[cfg(feature = "providers")]
#[derive(serde::Deserialize)]
struct EmbeddingsReply {
data: Vec<EmbeddingDatum>,
}
#[cfg(feature = "providers")]
#[derive(serde::Deserialize)]
struct EmbeddingDatum {
embedding: Vec<f32>,
}
#[cfg(feature = "providers")]
#[async_trait::async_trait]
impl Embedder for OpenAiEmbedder {
fn revision(&self) -> String {
use std::fmt::Write as _;
let mut revision = self.model.clone();
if let Some(d) = self.dimensions {
let _ = write!(revision, "@{d}");
}
if let Some(input_type) = &self.input_type {
let _ = write!(revision, "/{input_type}");
}
revision
}
async fn embed(&self, text: &str) -> Result<Vec<f32>, StoreError> {
self.check_egress()?;
let mut body = serde_json::json!({ "model": self.model, "input": text });
if let Some(dimensions) = self.dimensions {
body["dimensions"] = serde_json::json!(dimensions);
}
if let Some(input_type) = &self.input_type {
body["input_type"] = serde_json::json!(input_type);
}
let url = format!("{}/v1/embeddings", self.base.trim_end_matches('/'));
let mut request = self.http.post(&url).timeout(self.timeout).json(&body);
if let Some(key) = &self.key {
request = request.bearer_auth(key.expose());
}
let response = request
.send()
.await
.map_err(|e| StoreError::Backend(format!("{url}: {e}")))?;
let status = response.status();
let text_body = response
.text()
.await
.map_err(|e| StoreError::Backend(format!("{url}: unreadable reply: {e}")))?;
if !status.is_success() {
return Err(StoreError::Backend(format!(
"{url}: embeddings returned {status}: {text_body}"
)));
}
let reply: EmbeddingsReply = serde_json::from_str(&text_body)
.map_err(|e| StoreError::Backend(format!("{url}: unreadable reply: {e}")))?;
let [datum] = reply.data.as_slice() else {
return Err(StoreError::Backend(format!(
"{url}: one input was sent and {} embeddings came back",
reply.data.len()
)));
};
if datum.embedding.is_empty() {
return Err(StoreError::Backend(format!(
"{url}: the embedding is empty, which no index can rank against"
)));
}
Ok(datum.embedding.clone())
}
}
#[cfg(feature = "bedrock")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum EmbeddingDialect {
Titan,
Cohere,
}
#[cfg(feature = "bedrock")]
fn bedrock_body(
dialect: EmbeddingDialect,
text: &str,
dimensions: Option<u32>,
) -> Result<serde_json::Value, StoreError> {
match dialect {
EmbeddingDialect::Titan => {
let mut body = serde_json::json!({ "inputText": text });
if let Some(dimensions) = dimensions {
body["dimensions"] = serde_json::json!(dimensions);
}
Ok(body)
}
EmbeddingDialect::Cohere => {
if dimensions.is_some() {
return Err(StoreError::Backend(
"Cohere Embed takes no `dimensions`; the width is the model's. \
Drop `.dimensions(..)` or choose a Titan model — a knob that \
silently did nothing would put a width in the effect key that \
never reached the wire"
.to_owned(),
));
}
Ok(serde_json::json!({ "texts": [text], "input_type": "search_query" }))
}
}
}
#[cfg(feature = "bedrock")]
#[derive(Debug, Clone)]
pub struct BedrockEmbedder {
client: aws_sdk_bedrockruntime::Client,
model: String,
region: String,
dialect: EmbeddingDialect,
dimensions: Option<u32>,
}
#[cfg(feature = "bedrock")]
impl BedrockEmbedder {
pub async fn from_env(
region: impl Into<String>,
model: impl Into<String>,
dialect: EmbeddingDialect,
) -> Result<Self, StoreError> {
let region = region.into();
if region.trim().is_empty() {
return Err(StoreError::Backend(
"an AWS region is required for Bedrock embeddings".to_owned(),
));
}
let config = aws_config::defaults(aws_config::BehaviorVersion::latest())
.region(aws_config::Region::new(region.clone()))
.load()
.await;
Ok(Self::from_client(
aws_sdk_bedrockruntime::Client::new(&config),
region,
model,
dialect,
))
}
#[must_use]
pub fn from_client(
client: aws_sdk_bedrockruntime::Client,
region: impl Into<String>,
model: impl Into<String>,
dialect: EmbeddingDialect,
) -> Self {
Self {
client,
model: model.into(),
region: region.into(),
dialect,
dimensions: None,
}
}
#[must_use]
pub const fn dimensions(mut self, dimensions: u32) -> Self {
self.dimensions = Some(dimensions);
self
}
}
#[cfg(feature = "bedrock")]
#[async_trait::async_trait]
impl Embedder for BedrockEmbedder {
fn revision(&self) -> String {
let base = format!("bedrock:{}/{}", self.region, self.model);
self.dimensions
.map_or(base.clone(), |d| format!("{base}@{d}"))
}
async fn embed(&self, text: &str) -> Result<Vec<f32>, StoreError> {
let body = bedrock_body(self.dialect, text, self.dimensions)?;
let reply = self
.client
.invoke_model()
.model_id(&self.model)
.content_type("application/json")
.accept("application/json")
.body(aws_smithy_types::Blob::new(
crate::core::canon::value_bytes(&body),
))
.send()
.await
.map_err(|e| StoreError::Backend(format!("bedrock embeddings: {e}")))?;
let parsed: serde_json::Value = serde_json::from_slice(reply.body().as_ref())
.map_err(|e| StoreError::Backend(format!("bedrock embeddings: unreadable: {e}")))?;
let vector = match self.dialect {
EmbeddingDialect::Titan => parsed.get("embedding").cloned(),
EmbeddingDialect::Cohere => {
match parsed
.get("embeddings")
.and_then(|e| e.as_array())
.map(Vec::as_slice)
{
Some([only]) => Some(only.clone()),
_ => None,
}
}
};
let Some(serde_json::Value::Array(values)) = vector else {
return Err(StoreError::Backend(format!(
"bedrock embeddings: the reply carried no single vector in the \
{:?} shape — a declared dialect that does not match the model is \
the usual cause: {parsed}",
self.dialect
)));
};
let vector: Vec<f32> = values.iter().filter_map(json_f32).collect();
if vector.len() != values.len() || vector.is_empty() {
return Err(StoreError::Backend(
"bedrock embeddings: the vector is empty or not all numbers".to_owned(),
));
}
Ok(vector)
}
}
#[cfg(feature = "providers")]
#[derive(Debug, Clone)]
pub struct GeminiEmbedder {
http: reqwest::Client,
key: Secret,
base: String,
model: String,
dimensions: Option<u32>,
egress: Option<crate::core::Egress>,
timeout: std::time::Duration,
}
#[cfg(feature = "providers")]
impl GeminiEmbedder {
pub const DEFAULT_BASE: &'static str = "https://generativelanguage.googleapis.com";
pub fn new(key: impl Into<String>, model: impl Into<String>) -> Result<Self, StoreError> {
let http = reqwest::Client::builder()
.build()
.map_err(|e| StoreError::Backend(format!("could not build an HTTP client: {e}")))?;
Ok(Self {
http,
key: Secret::new(key),
base: Self::DEFAULT_BASE.to_owned(),
model: model.into(),
dimensions: None,
egress: None,
timeout: OpenAiEmbedder::DEFAULT_TIMEOUT,
})
}
pub fn from_env(model: impl Into<String>) -> Result<Self, StoreError> {
let key = std::env::var("GEMINI_API_KEY")
.or_else(|_| std::env::var("GOOGLE_API_KEY"))
.map_err(|_| {
StoreError::Backend("neither GEMINI_API_KEY nor GOOGLE_API_KEY is set".to_owned())
})?;
Self::new(key, model)
}
#[must_use]
pub fn base(mut self, base: impl Into<String>) -> Self {
self.base = base.into();
self
}
#[must_use]
pub const fn dimensions(mut self, dimensions: u32) -> Self {
self.dimensions = Some(dimensions);
self
}
#[must_use]
pub fn egress(mut self, egress: crate::core::Egress) -> Self {
self.egress = Some(egress);
self
}
#[must_use]
pub const fn timeout(mut self, timeout: std::time::Duration) -> Self {
self.timeout = timeout;
self
}
}
#[cfg(feature = "providers")]
#[async_trait::async_trait]
impl Embedder for GeminiEmbedder {
fn revision(&self) -> String {
let base = format!("gemini:{}", self.model);
self.dimensions
.map_or(base.clone(), |d| format!("{base}@{d}"))
}
async fn embed(&self, text: &str) -> Result<Vec<f32>, StoreError> {
if let Some(egress) = &self.egress {
let host = reqwest::Url::parse(&self.base)
.ok()
.and_then(|u| u.host_str().map(ToOwned::to_owned));
egress
.permits(host.as_deref())
.map_err(|e| StoreError::Backend(e.to_string()))?;
}
let mut body = serde_json::json!({
"content": { "parts": [{ "text": text }] },
"taskType": "RETRIEVAL_QUERY",
});
if let Some(dimensions) = self.dimensions {
body["outputDimensionality"] = serde_json::json!(dimensions);
}
let url = format!(
"{}/v1beta/models/{}:embedContent",
self.base.trim_end_matches('/'),
self.model
);
let response = self
.http
.post(&url)
.timeout(self.timeout)
.header("x-goog-api-key", self.key.expose())
.json(&body)
.send()
.await
.map_err(|e| StoreError::Backend(format!("{url}: {e}")))?;
let status = response.status();
let text_body = response
.text()
.await
.map_err(|e| StoreError::Backend(format!("{url}: unreadable reply: {e}")))?;
if !status.is_success() {
return Err(StoreError::Backend(format!(
"{url}: embedContent returned {status}: {text_body}"
)));
}
let parsed: serde_json::Value = serde_json::from_str(&text_body)
.map_err(|e| StoreError::Backend(format!("{url}: unreadable reply: {e}")))?;
let Some(serde_json::Value::Array(values)) = parsed
.get("embedding")
.and_then(|e| e.get("values"))
.cloned()
else {
return Err(StoreError::Backend(format!(
"{url}: the reply carried no embedding.values: {parsed}"
)));
};
let mut vector: Vec<f32> = values.iter().filter_map(json_f32).collect();
if vector.len() != values.len() || vector.is_empty() {
return Err(StoreError::Backend(format!(
"{url}: the vector is empty or not all numbers"
)));
}
if self.dimensions.is_some() {
let norm = vector.iter().map(|v| v * v).sum::<f32>().sqrt();
if !norm.is_finite() || norm <= 0.0 {
return Err(StoreError::Backend(format!(
"{url}: the vector has no usable magnitude ({norm}), so there \
is no direction to rank against"
)));
}
for v in &mut vector {
*v /= norm;
}
}
Ok(vector)
}
}
#[cfg(all(test, feature = "bedrock"))]
mod bedrock_dialect_tests {
use super::{EmbeddingDialect, bedrock_body};
#[test]
fn each_dialect_sends_its_own_shape() {
let titan = bedrock_body(EmbeddingDialect::Titan, "refund policy", Some(512))
.expect("titan takes a width");
assert_eq!(titan["inputText"], "refund policy");
assert_eq!(titan["dimensions"], 512);
let cohere = bedrock_body(EmbeddingDialect::Cohere, "refund policy", None).expect("cohere");
assert_eq!(cohere["texts"][0], "refund policy");
assert_eq!(
cohere["input_type"], "search_query",
"this seam embeds the thing being looked *for*; embedding it as a \
document ranks worse and reports nothing"
);
}
#[test]
fn cohere_refuses_a_width_it_cannot_send() {
let err = bedrock_body(EmbeddingDialect::Cohere, "x", Some(512))
.expect_err("a width Cohere cannot send was accepted");
assert!(err.to_string().contains("no `dimensions`"), "{err}");
}
}
#[cfg(test)]
mod narrowing_tests {
use super::json_f32;
#[test]
fn a_component_no_f32_can_hold_is_refused_rather_than_infinite() {
for out_of_range in ["1e39", "-1e39", "1e300"] {
let value: serde_json::Value = serde_json::from_str(out_of_range).expect("valid JSON");
assert_eq!(
json_f32(&value),
None,
"{out_of_range} narrowed to a non-finite component; journaled as \
`null` it would share an effect key with every other one"
);
}
assert_eq!(json_f32(&serde_json::json!(1.0)), Some(1.0));
assert_eq!(json_f32(&serde_json::json!(-0.0321)), Some(-0.0321));
assert_eq!(
json_f32(&serde_json::json!(0.123_456_789_012_345_68_f64)),
Some(0.123_456_79),
"precision loss is the contract; range loss is the defect"
);
assert_eq!(json_f32(&serde_json::json!("0.5")), None);
}
}
#[cfg(all(test, feature = "bedrock"))]
mod bedrock_reply_tests {
use super::{BedrockEmbedder, EmbeddingDialect};
use crate::memory::Embedder as _;
fn embedder(dialect: EmbeddingDialect, body: &'static str) -> BedrockEmbedder {
let config = aws_sdk_bedrockruntime::Config::builder()
.region(aws_config::Region::new("eu-central-1"))
.credentials_provider(aws_sdk_bedrockruntime::config::Credentials::for_tests())
.behavior_version_latest()
.http_client(aws_smithy_http_client::test_util::infallible_client_fn(
move |_req| {
http::Response::builder()
.status(200)
.header("content-type", "application/json")
.body(body)
.unwrap()
},
))
.build();
BedrockEmbedder::from_client(
aws_sdk_bedrockruntime::Client::from_conf(config),
"eu-central-1",
"amazon.titan-embed-text-v2:0",
dialect,
)
}
#[tokio::test]
async fn each_dialect_reads_its_own_reply() {
let titan = embedder(EmbeddingDialect::Titan, r#"{"embedding":[0.25,-0.5,0.75]}"#)
.embed("refund policy")
.await
.expect("titan reply");
assert_eq!(titan, vec![0.25, -0.5, 0.75]);
let cohere = embedder(EmbeddingDialect::Cohere, r#"{"embeddings":[[1.0,0.0]]}"#)
.embed("refund policy")
.await
.expect("cohere reply");
assert_eq!(cohere, vec![1.0, 0.0]);
}
#[tokio::test]
async fn a_reply_this_dialect_cannot_read_is_refused() {
for (dialect, body, what) in [
(
EmbeddingDialect::Titan,
r#"{"embeddings":[[1.0]]}"#,
"Cohere's shape read as Titan's",
),
(
EmbeddingDialect::Cohere,
r#"{"embedding":[1.0]}"#,
"Titan's shape read as Cohere's",
),
(
EmbeddingDialect::Cohere,
r#"{"embeddings":[[1.0],[2.0]]}"#,
"two rows for one text",
),
(
EmbeddingDialect::Titan,
r#"{"embedding":[]}"#,
"an empty vector, which no index can rank against",
),
(
EmbeddingDialect::Titan,
r#"{"embedding":[1.0,"nan",2.0]}"#,
"a component that is not a number",
),
(
EmbeddingDialect::Titan,
r#"{"embedding":[1.0,1e39]}"#,
"a component no f32 can hold, which narrows to infinity",
),
] {
let err = embedder(dialect, body)
.embed("x")
.await
.expect_err(&format!("accepted {what}"));
assert!(
matches!(err, crate::core::StoreError::Backend(_)),
"{what}: {err}"
);
}
}
#[test]
fn the_revision_names_the_region_the_model_and_the_width() {
let base = |region| {
let config = aws_sdk_bedrockruntime::Config::builder()
.region(aws_config::Region::new(region))
.behavior_version_latest()
.http_client(aws_smithy_http_client::test_util::infallible_client_fn(
|_req| http::Response::builder().status(200).body("").unwrap(),
))
.build();
BedrockEmbedder::from_client(
aws_sdk_bedrockruntime::Client::from_conf(config),
region,
"amazon.titan-embed-text-v2:0",
EmbeddingDialect::Titan,
)
};
assert_eq!(
base("eu-central-1").revision(),
"bedrock:eu-central-1/amazon.titan-embed-text-v2:0"
);
assert_ne!(
base("eu-central-1").revision(),
base("us-east-1").revision(),
"two regions are two services and shared one effect identity"
);
assert_eq!(
base("eu-central-1").dimensions(256).revision(),
"bedrock:eu-central-1/amazon.titan-embed-text-v2:0@256",
"a width that does not reach the revision lets two geometries share an index"
);
}
}