1use rskit_errors::{AppError, AppResult, ErrorCode};
4use serde::{Deserialize, Deserializer, Serialize, de};
5
6#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
8#[serde(transparent)]
9pub struct EmbeddingOptions(serde_json::Value);
10
11impl EmbeddingOptions {
12 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 #[must_use]
26 pub const fn as_json(&self) -> &serde_json::Value {
27 &self.0
28 }
29
30 #[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#[derive(Debug, Clone, Serialize, Deserialize)]
54pub struct EmbedRequest {
55 pub model: rskit_ai::Model,
57 pub inputs: Vec<EmbedInput>,
59 #[serde(default)]
61 pub options: EmbeddingOptions,
62}
63
64#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
66#[serde(tag = "type", content = "value", rename_all = "snake_case")]
67#[non_exhaustive]
68pub enum EmbedInput {
69 Text(String),
71 Image(EmbedAsset),
73 Audio(EmbedAsset),
75 Video(EmbedAsset),
77}
78
79#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
81#[serde(tag = "type", content = "value", rename_all = "snake_case")]
82#[non_exhaustive]
83pub enum EmbedAsset {
84 Bytes(Vec<u8>),
86 Url(String),
88}
89
90#[derive(Debug, Clone, Serialize, Deserialize)]
92pub struct EmbedResponse {
93 pub embeddings: Vec<Embedding>,
95 pub model: rskit_ai::Model,
97 pub usage: rskit_ai::Usage,
99}
100
101#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
103pub struct Embedding {
104 pub vector: Vec<f32>,
106 pub dimensions: usize,
108 pub index: usize,
110}
111
112impl Embedding {
113 #[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}