use std::env;
use std::time::Duration;
use async_trait::async_trait;
use serde::Deserialize;
#[derive(Debug, thiserror::Error)]
pub enum EmbeddingError {
#[error("embedding request failed: {0}")]
Http(#[from] reqwest::Error),
#[error("embedding endpoint returned status {0}")]
Status(reqwest::StatusCode),
#[error("embeddings response contained no data")]
EmptyResponse,
}
#[derive(Debug, thiserror::Error)]
#[error("embedding blob length {0} is not a multiple of 4")]
pub struct DecodeEmbeddingError(pub usize);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EmbeddingConfig {
pub base_url: String,
pub model: String,
pub api_key: Option<String>,
}
const DEFAULT_MODEL: &str = "text-embedding-bge-large-en-v1.5";
const ENV_BASE_URL: &str = "KB_EMBEDDING_BASE_URL";
const ENV_MODEL: &str = "KB_EMBEDDING_MODEL";
const ENV_API_KEY: &str = "KB_EMBEDDING_API_KEY";
#[must_use]
#[allow(
clippy::disallowed_methods,
reason = "kb-server takes all runtime config from the environment (12-factor); the embedding endpoint/model are deployment config, not user-facing tunables, and are read once at the server entry edge (REPO_INVARIANTS.md #5)"
)]
pub fn read_embedding_config_from_env() -> Option<EmbeddingConfig> {
config_from_env_inputs(
env::var(ENV_BASE_URL).ok(),
env::var(ENV_MODEL).ok(),
env::var(ENV_API_KEY).ok(),
)
}
fn config_from_env_inputs(
base_url: Option<String>,
model: Option<String>,
api_key: Option<String>,
) -> Option<EmbeddingConfig> {
let base_url = base_url?;
if base_url.is_empty() {
return None;
}
Some(EmbeddingConfig {
base_url,
model: model
.filter(|s| !s.is_empty())
.unwrap_or_else(|| DEFAULT_MODEL.into()),
api_key: api_key.filter(|s| !s.is_empty()),
})
}
#[async_trait]
pub trait EmbeddingClient: Send + Sync {
async fn embed(&self, input: &str) -> Result<Vec<f32>, EmbeddingError>;
}
pub struct HttpEmbeddingClient {
base_url: String,
model: String,
api_key: Option<String>,
client: reqwest::Client,
}
impl HttpEmbeddingClient {
fn endpoint(&self) -> String {
format!("{}/embeddings", self.base_url)
}
}
#[async_trait]
impl EmbeddingClient for HttpEmbeddingClient {
async fn embed(&self, input: &str) -> Result<Vec<f32>, EmbeddingError> {
let body = serde_json::json!({
"model": self.model,
"input": input,
});
let mut req = self.client.post(self.endpoint()).json(&body);
if let Some(key) = &self.api_key {
req = req.bearer_auth(key);
}
let resp = req.send().await?;
let status = resp.status();
if !status.is_success() {
return Err(EmbeddingError::Status(status));
}
let parsed: EmbeddingResponse = resp.json().await?;
parsed
.data
.into_iter()
.next()
.map(|item| item.embedding)
.ok_or(EmbeddingError::EmptyResponse)
}
}
#[must_use]
pub fn http_embedding_client(config: EmbeddingConfig) -> HttpEmbeddingClient {
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("the default reqwest client builds unless the system TLS backend is unavailable");
HttpEmbeddingClient {
base_url: config.base_url,
model: config.model,
api_key: config.api_key,
client,
}
}
#[derive(Deserialize)]
struct EmbeddingResponse {
data: Vec<EmbeddingItem>,
}
#[derive(Deserialize)]
struct EmbeddingItem {
embedding: Vec<f32>,
}
#[must_use]
pub fn encode_embedding(v: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(v.len() * 4);
for f in v {
out.extend_from_slice(&f.to_le_bytes());
}
out
}
pub fn decode_embedding(bytes: &[u8]) -> Result<Vec<f32>, DecodeEmbeddingError> {
if !bytes.len().is_multiple_of(4) {
return Err(DecodeEmbeddingError(bytes.len()));
}
let n = bytes.len() / 4;
let mut out = Vec::with_capacity(n);
for chunk in bytes.chunks_exact(4) {
let arr: [u8; 4] = chunk.try_into().expect("chunks_exact yields 4 bytes");
out.push(f32::from_le_bytes(arr));
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use mockito::Matcher;
use proptest::prelude::*;
#[test]
fn embedding_config_returns_none_when_url_unset() {
assert_eq!(config_from_env_inputs(None, None, None), None);
}
#[test]
fn embedding_config_returns_none_when_url_empty() {
assert_eq!(
config_from_env_inputs(Some(String::new()), None, None),
None
);
}
#[test]
fn embedding_config_uses_default_model_when_unset() {
let cfg = config_from_env_inputs(Some("http://x/v1".into()), None, None).unwrap();
assert_eq!(cfg.base_url, "http://x/v1");
assert_eq!(cfg.model, "text-embedding-bge-large-en-v1.5");
assert!(cfg.api_key.is_none());
}
#[test]
fn embedding_config_passes_through_explicit_model() {
let cfg = config_from_env_inputs(
Some("http://x/v1".into()),
Some("custom-model".into()),
None,
)
.unwrap();
assert_eq!(cfg.model, "custom-model");
}
#[test]
fn embedding_config_carries_api_key_when_set() {
let cfg = config_from_env_inputs(Some("http://x/v1".into()), None, Some("sk-abc".into()))
.unwrap();
assert_eq!(cfg.api_key.as_deref(), Some("sk-abc"));
}
#[test]
fn embedding_config_treats_empty_api_key_as_none() {
let cfg =
config_from_env_inputs(Some("http://x/v1".into()), None, Some(String::new())).unwrap();
assert!(cfg.api_key.is_none());
}
#[test]
fn embedding_config_falls_back_when_model_empty() {
let cfg =
config_from_env_inputs(Some("http://x/v1".into()), Some(String::new()), None).unwrap();
assert_eq!(cfg.model, "text-embedding-bge-large-en-v1.5");
}
#[test]
fn embedding_blob_codec_empty_vec_roundtrips() {
let v: Vec<f32> = vec![];
assert_eq!(decode_embedding(&encode_embedding(&v)).unwrap(), v);
}
#[test]
fn embedding_blob_codec_basic_roundtrip() {
let v = vec![0.0, 1.0, -1.0, 0.5, f32::MIN, f32::MAX];
let bytes = encode_embedding(&v);
assert_eq!(bytes.len(), 4 * v.len());
assert_eq!(decode_embedding(&bytes).unwrap(), v);
}
#[test]
fn embedding_blob_codec_rejects_truncated_buffer() {
let bad: Vec<u8> = vec![0; 7];
let err = decode_embedding(&bad).unwrap_err();
assert!(matches!(err, DecodeEmbeddingError(7)), "got: {err:?}");
}
#[test]
fn embedding_blob_codec_is_little_endian() {
let bytes = encode_embedding(&[1.0_f32]);
assert_eq!(bytes, vec![0x00, 0x00, 0x80, 0x3F]);
}
proptest! {
#[test]
fn embedding_blob_codec_property_roundtrip(
v in proptest::collection::vec(any::<f32>(), 0..32)
) {
let bytes = encode_embedding(&v);
let back = decode_embedding(&bytes).unwrap();
prop_assert_eq!(back.len(), v.len());
for (a, b) in v.iter().zip(back.iter()) {
prop_assert_eq!(a.to_bits(), b.to_bits());
}
}
}
fn make_client(server_url: &str, api_key: Option<&str>) -> HttpEmbeddingClient {
http_embedding_client(EmbeddingConfig {
base_url: server_url.to_string(),
model: "test-model".into(),
api_key: api_key.map(str::to_string),
})
}
#[tokio::test]
async fn embedding_client_posts_model_and_input() {
let mut server = mockito::Server::new_async().await;
let _m = server
.mock("POST", "/embeddings")
.match_header("content-type", "application/json")
.match_body(Matcher::PartialJsonString(
r#"{"model":"test-model","input":"hello"}"#.into(),
))
.with_status(200)
.with_body(r#"{"data":[{"embedding":[0.1,0.2,0.3]}]}"#)
.create_async()
.await;
let client = make_client(&server.url(), None);
let v = client.embed("hello").await.unwrap();
assert_eq!(v, vec![0.1_f32, 0.2, 0.3]);
}
#[tokio::test]
async fn embedding_client_sets_bearer_auth_when_api_key_present() {
let mut server = mockito::Server::new_async().await;
let _m = server
.mock("POST", "/embeddings")
.match_header("authorization", "Bearer sk-abc")
.with_status(200)
.with_body(r#"{"data":[{"embedding":[1.0]}]}"#)
.create_async()
.await;
let client = make_client(&server.url(), Some("sk-abc"));
let v = client.embed("hi").await.unwrap();
assert_eq!(v, vec![1.0_f32]);
}
#[tokio::test]
async fn embedding_client_omits_auth_when_api_key_absent() {
let mut server = mockito::Server::new_async().await;
let _m = server
.mock("POST", "/embeddings")
.match_header("authorization", Matcher::Missing)
.with_status(200)
.with_body(r#"{"data":[{"embedding":[2.0]}]}"#)
.create_async()
.await;
let client = make_client(&server.url(), None);
let v = client.embed("hi").await.unwrap();
assert_eq!(v, vec![2.0_f32]);
}
#[tokio::test]
async fn embedding_client_returns_err_on_empty_data() {
let mut server = mockito::Server::new_async().await;
let _m = server
.mock("POST", "/embeddings")
.with_status(200)
.with_body(r#"{"data":[]}"#)
.create_async()
.await;
let client = make_client(&server.url(), None);
let err = client.embed("hi").await.unwrap_err();
assert!(matches!(err, EmbeddingError::EmptyResponse), "got: {err:?}");
}
#[tokio::test]
async fn embedding_client_returns_err_on_http_500() {
let mut server = mockito::Server::new_async().await;
let _m = server
.mock("POST", "/embeddings")
.with_status(500)
.with_body("boom")
.create_async()
.await;
let client = make_client(&server.url(), None);
let err = client.embed("hi").await.unwrap_err();
assert!(matches!(err, EmbeddingError::Status(_)), "got: {err:?}");
}
#[tokio::test]
async fn embedding_client_returns_err_on_malformed_json() {
let mut server = mockito::Server::new_async().await;
let _m = server
.mock("POST", "/embeddings")
.with_status(200)
.with_body("not json")
.create_async()
.await;
let client = make_client(&server.url(), None);
let err = client.embed("hi").await.unwrap_err();
assert!(matches!(err, EmbeddingError::Http(_)), "got: {err:?}");
}
#[tokio::test]
async fn embedding_client_returns_err_on_unreachable_endpoint() {
let client = make_client("http://127.0.0.1:1", None);
let err = client.embed("hi").await.unwrap_err();
assert!(matches!(err, EmbeddingError::Http(_)), "got: {err:?}");
}
#[tokio::test]
async fn http_embedding_client_async_construction_does_not_panic_in_tokio() {
let _client = http_embedding_client(EmbeddingConfig {
base_url: "http://127.0.0.1:1/v1".into(),
model: "test-model".into(),
api_key: None,
});
}
struct StubAsyncClient;
#[async_trait]
impl EmbeddingClient for StubAsyncClient {
async fn embed(&self, _input: &str) -> Result<Vec<f32>, EmbeddingError> {
Ok(vec![1.0, 2.0, 3.0])
}
}
#[tokio::test]
async fn embedding_client_async_trait_method_is_awaitable() {
let client = StubAsyncClient;
let v = client.embed("anything").await.unwrap();
assert_eq!(v, vec![1.0_f32, 2.0, 3.0]);
}
}