use super::error::EmbedError;
use super::mrl::{model_default_input_type, mrl_wire_dimensions};
use super::wire::{EmbeddingInput, EmbeddingRequest};
use super::{
OpenRouterClient, DEFAULT_CONNECT_TIMEOUT_SECS, DEFAULT_EMBED_HTTP_BATCH_SIZE,
DEFAULT_TIMEOUT_SECS,
};
use crate::constants::DEFAULT_OPENROUTER_EMBEDDINGS_URL;
use crate::errors::AppError;
use secrecy::SecretBox;
use std::time::Duration;
impl OpenRouterClient {
pub fn new(
api_key: SecretBox<String>,
model: String,
dim: usize,
timeout_secs: u64,
) -> Result<Self, AppError> {
let base_url =
crate::runtime_config::openrouter_embeddings_url(DEFAULT_OPENROUTER_EMBEDDINGS_URL);
Self::new_with_base_url(api_key, model, dim, timeout_secs, base_url)
}
pub fn new_with_base_url(
api_key: SecretBox<String>,
model: String,
dim: usize,
timeout_secs: u64,
base_url: String,
) -> Result<Self, AppError> {
let timeout_secs = if timeout_secs == 0 {
DEFAULT_TIMEOUT_SECS
} else {
timeout_secs
};
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.connect_timeout(Duration::from_secs(DEFAULT_CONNECT_TIMEOUT_SECS))
.user_agent(concat!("sqlite-graphrag/", env!("CARGO_PKG_VERSION")))
.build()
.map_err(|e| {
AppError::Embedding(crate::i18n::validation::embedding_http_client_build_failed(
e,
))
})?;
let default_input_type = model_default_input_type(&model);
Ok(Self {
client,
api_key,
model,
dim,
default_input_type,
base_url,
})
}
#[cfg(test)]
pub(super) fn new_with_url(
api_key: SecretBox<String>,
model: String,
dim: usize,
timeout_secs: u64,
base_url: String,
) -> Result<Self, AppError> {
Self::new_with_base_url(api_key, model, dim, timeout_secs, base_url)
}
pub fn default_input_type(&self) -> Option<&'static str> {
self.default_input_type
}
pub async fn embed_single(
&self,
text: &str,
input_type: Option<&str>,
) -> Result<Vec<f32>, EmbedError> {
crate::memory_guard::check_embedding_input_size(text)?;
let request = EmbeddingRequest {
model: &self.model,
input: EmbeddingInput::Single(text),
dimensions: mrl_wire_dimensions(&self.model, self.dim),
encoding_format: "float",
input_type,
};
let response = self.execute_with_retry(&request).await?;
let embedding = response
.data
.into_iter()
.next()
.ok_or_else(|| {
AppError::Embedding(
crate::i18n::validation::embedding_empty_response_from_openrouter(),
)
})?
.embedding;
Ok(self.truncate_embedding(embedding)?)
}
pub async fn embed_batch(
&self,
texts: &[&str],
input_type: Option<&str>,
) -> Result<Vec<Vec<f32>>, EmbedError> {
if texts.is_empty() {
return Ok(Vec::new());
}
for text in texts {
crate::memory_guard::check_embedding_input_size(text)?;
}
let mut all = Vec::with_capacity(texts.len());
let batch_size = crate::runtime_config::embedding_batch_size(DEFAULT_EMBED_HTTP_BATCH_SIZE);
for chunk in texts.chunks(batch_size) {
let request = EmbeddingRequest {
model: &self.model,
input: EmbeddingInput::Batch(chunk.to_vec()),
dimensions: mrl_wire_dimensions(&self.model, self.dim),
encoding_format: "float",
input_type,
};
let response = self.execute_with_retry(&request).await?;
if response.data.len() != chunk.len() {
return Err(AppError::Embedding(
crate::i18n::validation::embedding_expected_count(
chunk.len(),
response.data.len(),
),
)
.into());
}
let mut sorted = response.data;
sorted.sort_by_key(|d| d.index);
for d in sorted {
all.push(self.truncate_embedding(d.embedding)?);
}
}
Ok(all)
}
}