use rskit_errors::{AppError, AppResult, ErrorCode};
use serde::{Deserialize, Deserializer, Serialize, de};
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
#[serde(transparent)]
pub struct EmbeddingOptions(serde_json::Value);
impl EmbeddingOptions {
pub fn new(value: serde_json::Value) -> AppResult<Self> {
if value.is_object() {
Ok(Self(value))
} else {
Err(AppError::new(
ErrorCode::InvalidInput,
"embedding options must be a JSON object",
))
}
}
#[must_use]
pub const fn as_json(&self) -> &serde_json::Value {
&self.0
}
#[must_use]
pub fn into_json(self) -> serde_json::Value {
self.0
}
}
impl<'de> Deserialize<'de> for EmbeddingOptions {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
Self::new(serde_json::Value::deserialize(deserializer)?).map_err(de::Error::custom)
}
}
impl Default for EmbeddingOptions {
fn default() -> Self {
Self(serde_json::Value::Object(serde_json::Map::new()))
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbedRequest {
pub model: rskit_ai::Model,
pub inputs: Vec<EmbedInput>,
#[serde(default)]
pub options: EmbeddingOptions,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", content = "value", rename_all = "snake_case")]
#[non_exhaustive]
pub enum EmbedInput {
Text(String),
Image(EmbedAsset),
Audio(EmbedAsset),
Video(EmbedAsset),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", content = "value", rename_all = "snake_case")]
#[non_exhaustive]
pub enum EmbedAsset {
Bytes(Vec<u8>),
Url(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbedResponse {
pub embeddings: Vec<Embedding>,
pub model: rskit_ai::Model,
pub usage: rskit_ai::Usage,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Embedding {
pub vector: Vec<f32>,
pub dimensions: usize,
pub index: usize,
}
impl Embedding {
#[must_use]
pub const fn new(vector: Vec<f32>, index: usize) -> Self {
let dimensions = vector.len();
Self {
vector,
dimensions,
index,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn embedding_sets_dimensions() {
let e = Embedding::new(vec![1.0, 2.0], 3);
assert_eq!(e.dimensions, 2);
assert_eq!(e.index, 3);
}
#[test]
fn embedding_options_reject_non_object() {
let err = serde_json::from_str::<EmbeddingOptions>("null").unwrap_err();
assert!(
err.to_string()
.contains("embedding options must be a JSON object")
);
}
}