Skip to main content

rskit_embedding/
types.rs

1//! Embedding data types, distance metrics, and aggregation functions.
2
3use rskit_errors::{AppError, AppResult, ErrorCode};
4use serde::{Deserialize, Deserializer, Serialize, de};
5
6/// Provider-specific embedding options.
7#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
8#[serde(transparent)]
9pub struct EmbeddingOptions(serde_json::Value);
10
11impl EmbeddingOptions {
12    /// Create options from a JSON object.
13    pub fn new(value: serde_json::Value) -> AppResult<Self> {
14        if value.is_object() {
15            Ok(Self(value))
16        } else {
17            Err(AppError::new(
18                ErrorCode::InvalidInput,
19                "embedding options must be a JSON object",
20            ))
21        }
22    }
23
24    /// Borrow the structured options.
25    #[must_use]
26    pub const fn as_json(&self) -> &serde_json::Value {
27        &self.0
28    }
29
30    /// Consume the wrapper and return structured options.
31    #[must_use]
32    pub fn into_json(self) -> serde_json::Value {
33        self.0
34    }
35}
36
37impl<'de> Deserialize<'de> for EmbeddingOptions {
38    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
39    where
40        D: Deserializer<'de>,
41    {
42        Self::new(serde_json::Value::deserialize(deserializer)?).map_err(de::Error::custom)
43    }
44}
45
46impl Default for EmbeddingOptions {
47    fn default() -> Self {
48        Self(serde_json::Value::Object(serde_json::Map::new()))
49    }
50}
51
52/// Canonical embedding request.
53#[derive(Debug, Clone, Serialize, Deserialize)]
54pub struct EmbedRequest {
55    /// Model requested for embedding.
56    pub model: rskit_ai::Model,
57    /// Inputs to embed.
58    pub inputs: Vec<EmbedInput>,
59    /// Provider-specific knobs.
60    #[serde(default)]
61    pub options: EmbeddingOptions,
62}
63
64/// Multimodal embedding input.
65#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
66#[serde(tag = "type", content = "value", rename_all = "snake_case")]
67#[non_exhaustive]
68pub enum EmbedInput {
69    /// Text input.
70    Text(String),
71    /// Image asset.
72    Image(EmbedAsset),
73    /// Audio asset.
74    Audio(EmbedAsset),
75    /// Video asset.
76    Video(EmbedAsset),
77}
78
79/// Bytes or URL asset input.
80#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
81#[serde(tag = "type", content = "value", rename_all = "snake_case")]
82#[non_exhaustive]
83pub enum EmbedAsset {
84    /// Inline bytes.
85    Bytes(Vec<u8>),
86    /// Fetchable URL.
87    Url(String),
88}
89
90/// Canonical embedding response.
91#[derive(Debug, Clone, Serialize, Deserialize)]
92pub struct EmbedResponse {
93    /// Embeddings returned in input order.
94    pub embeddings: Vec<Embedding>,
95    /// Model that served the request.
96    pub model: rskit_ai::Model,
97    /// Usage counters.
98    pub usage: rskit_ai::Usage,
99}
100
101/// A single embedding vector.
102#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
103pub struct Embedding {
104    /// The embedding vector.
105    pub vector: Vec<f32>,
106    /// Vector dimensions.
107    pub dimensions: usize,
108    /// Zero-based input index.
109    pub index: usize,
110}
111
112impl Embedding {
113    /// Create a new embedding from a vector and input index.
114    #[must_use]
115    pub const fn new(vector: Vec<f32>, index: usize) -> Self {
116        let dimensions = vector.len();
117        Self {
118            vector,
119            dimensions,
120            index,
121        }
122    }
123}
124
125#[cfg(test)]
126mod tests {
127    use super::*;
128
129    #[test]
130    fn embedding_sets_dimensions() {
131        let e = Embedding::new(vec![1.0, 2.0], 3);
132        assert_eq!(e.dimensions, 2);
133        assert_eq!(e.index, 3);
134    }
135
136    #[test]
137    fn embedding_options_reject_non_object() {
138        let err = serde_json::from_str::<EmbeddingOptions>("null").unwrap_err();
139        assert!(
140            err.to_string()
141                .contains("embedding options must be a JSON object")
142        );
143    }
144}