zai_rs/model/text_embedded/
request.rs1use serde::{Deserialize, Serialize};
2
3#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
5#[serde(rename_all = "kebab-case")]
6pub enum EmbeddingModel {
7 #[serde(rename = "embedding-3")]
9 Embedding3,
10 #[serde(rename = "embedding-2")]
12 Embedding2,
13}
14
15#[derive(Clone, Serialize, Deserialize)]
17#[serde(untagged)]
18pub enum EmbeddingInput {
19 Single(String),
21 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
42pub enum EmbeddingDimensions {
43 D2048,
45 D1024,
47 D512,
49 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#[derive(Clone, Serialize)]
70pub struct EmbeddingBody {
71 pub model: EmbeddingModel,
73
74 pub input: EmbeddingInput,
76
77 #[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 pub fn new(model: EmbeddingModel, input: EmbeddingInput) -> Self {
97 Self {
98 model,
99 input,
100 dimensions: None,
101 }
102 }
103
104 pub fn with_dimensions(mut self, dims: EmbeddingDimensions) -> Self {
106 self.dimensions = Some(dims);
107 self
108 }
109
110 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}