Skip to main content

zai_rs/model/text_embedded/
request.rs

1use serde::{Deserialize, Serialize};
2
3/// Embedding model enum
4#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
5#[serde(rename_all = "kebab-case")]
6pub enum EmbeddingModel {
7    /// embedding-3 (supports configurable dimensions).
8    #[serde(rename = "embedding-3")]
9    Embedding3,
10    /// embedding-2 (fixed 1024 dimensions).
11    #[serde(rename = "embedding-2")]
12    Embedding2,
13}
14
15/// Input can be a single string or an array of strings
16#[derive(Clone, Serialize, Deserialize)]
17#[serde(untagged)]
18pub enum EmbeddingInput {
19    /// A single input string.
20    Single(String),
21    /// A batch of input strings.
22    Batch(Vec<String>),
23}
24
25impl std::fmt::Debug for EmbeddingInput {
26    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
27        match self {
28            Self::Single(_) => formatter
29                .debug_tuple("Single")
30                .field(&"[REDACTED]")
31                .finish(),
32            Self::Batch(values) => formatter
33                .debug_struct("Batch")
34                .field("len", &values.len())
35                .finish(),
36        }
37    }
38}
39
40/// Output vector dimensions for embeddings
41#[derive(Debug, Clone, Copy, PartialEq, Eq)]
42pub enum EmbeddingDimensions {
43    /// 2048-dimensional embedding.
44    D2048,
45    /// 1024-dimensional embedding.
46    D1024,
47    /// 512-dimensional embedding.
48    D512,
49    /// 256-dimensional embedding.
50    D256,
51}
52
53impl Serialize for EmbeddingDimensions {
54    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
55    where
56        S: serde::Serializer,
57    {
58        let v: u16 = match self {
59            EmbeddingDimensions::D2048 => 2048,
60            EmbeddingDimensions::D1024 => 1024,
61            EmbeddingDimensions::D512 => 512,
62            EmbeddingDimensions::D256 => 256,
63        };
64        serializer.serialize_u16(v)
65    }
66}
67
68/// Request body for embeddings
69#[derive(Clone, Serialize)]
70pub struct EmbeddingBody {
71    /// Embedding model (`embedding-3` or `embedding-2`).
72    pub model: EmbeddingModel,
73
74    /// A single input string or a batch of strings.
75    pub input: EmbeddingInput,
76
77    /// Output dimensions. `embedding-3` supports all variants;
78    /// `embedding-2` accepts only 1,024 dimensions or omission.
79    #[serde(skip_serializing_if = "Option::is_none")]
80    pub dimensions: Option<EmbeddingDimensions>,
81}
82
83impl std::fmt::Debug for EmbeddingBody {
84    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
85        formatter
86            .debug_struct("EmbeddingBody")
87            .field("model", &self.model)
88            .field("input", &self.input)
89            .field("dimensions", &self.dimensions)
90            .finish()
91    }
92}
93
94impl EmbeddingBody {
95    /// Create a new embedding request body from a model and input.
96    pub fn new(model: EmbeddingModel, input: EmbeddingInput) -> Self {
97        Self {
98            model,
99            input,
100            dimensions: None,
101        }
102    }
103
104    /// Set the output vector dimensionality (embedding-3 only).
105    pub fn with_dimensions(mut self, dims: EmbeddingDimensions) -> Self {
106        self.dimensions = Some(dims);
107        self
108    }
109
110    /// Enforce input and model/dimension constraints before sending.
111    pub fn validate_model_constraints(&self) -> Result<(), validator::ValidationError> {
112        use validator::ValidationError;
113        let has_empty_input = match &self.input {
114            EmbeddingInput::Single(value) => value.trim().is_empty(),
115            EmbeddingInput::Batch(values) => {
116                values.is_empty() || values.iter().any(|value| value.trim().is_empty())
117            },
118        };
119        if has_empty_input {
120            return Err(ValidationError::new("input_must_not_be_empty"));
121        }
122        if let EmbeddingModel::Embedding3 = self.model
123            && let EmbeddingInput::Batch(ref v) = self.input
124            && v.len() > 64
125        {
126            return Err(ValidationError::new("batch_too_long"));
127        }
128        if let EmbeddingModel::Embedding2 = self.model
129            && let Some(d) = self.dimensions
130            && d != EmbeddingDimensions::D1024
131        {
132            return Err(ValidationError::new("embedding2_dims_must_be_1024"));
133        }
134        Ok(())
135    }
136}
137
138#[cfg(test)]
139mod tests {
140    use super::*;
141
142    #[test]
143    fn validation_rejects_empty_inputs() {
144        for input in [
145            EmbeddingInput::Single(" ".to_owned()),
146            EmbeddingInput::Batch(Vec::new()),
147            EmbeddingInput::Batch(vec!["valid".to_owned(), String::new()]),
148        ] {
149            assert!(
150                EmbeddingBody::new(EmbeddingModel::Embedding3, input)
151                    .validate_model_constraints()
152                    .is_err()
153            );
154        }
155    }
156
157    #[test]
158    fn debug_redacts_embedding_inputs() {
159        let body = EmbeddingBody::new(
160            EmbeddingModel::Embedding3,
161            EmbeddingInput::Batch(vec!["private one".to_owned(), "private two".to_owned()]),
162        );
163        let debug = format!("{body:?}");
164        assert!(!debug.contains("private one"));
165        assert!(!debug.contains("private two"));
166        assert!(debug.contains("len: 2"));
167    }
168}