Skip to main content

zai_rs/model/text_embedded/
response.rs

1use serde::{Deserialize, Deserializer, Serialize, de::Error as _};
2
3/// Top-level response from the embeddings endpoint.
4#[derive(Debug, Clone, Serialize)]
5pub struct EmbeddingResponse {
6    /// Model that produced the embeddings.
7    #[serde(skip_serializing_if = "Option::is_none")]
8    pub model: Option<String>,
9    /// Object kind of the response envelope (always `list`).
10    #[serde(skip_serializing_if = "Option::is_none")]
11    pub object: Option<ResponseObjectKind>,
12    /// Per-input embedding vectors.
13    #[serde(skip_serializing_if = "Option::is_none")]
14    pub data: Option<Vec<EmbeddingData>>,
15    /// Token-usage statistics for the request.
16    #[serde(skip_serializing_if = "Option::is_none")]
17    pub usage: Option<EmbeddingUsage>,
18}
19
20#[derive(Deserialize)]
21struct EmbeddingResponseWire {
22    model: Option<String>,
23    object: Option<ResponseObjectKind>,
24    data: Option<Vec<EmbeddingData>>,
25    usage: Option<EmbeddingUsage>,
26}
27
28impl<'de> Deserialize<'de> for EmbeddingResponse {
29    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
30    where
31        D: Deserializer<'de>,
32    {
33        let wire = EmbeddingResponseWire::deserialize(deserializer)?;
34        if wire.model.is_none()
35            && wire.object.is_none()
36            && wire.data.is_none()
37            && wire.usage.is_none()
38        {
39            return Err(D::Error::custom(
40                "embedding response contained no documented fields",
41            ));
42        }
43        Ok(Self {
44            model: wire.model,
45            object: wire.object,
46            data: wire.data,
47            usage: wire.usage,
48        })
49    }
50}
51
52/// Top-level object kind returned by the embeddings endpoint.
53#[derive(Debug, Clone, Serialize, Deserialize)]
54#[serde(rename_all = "lowercase")]
55pub enum ResponseObjectKind {
56    /// A list of embedding items.
57    List,
58}
59
60/// A single embedding vector with its index.
61#[derive(Debug, Clone, Serialize, Deserialize)]
62pub struct EmbeddingData {
63    /// Position of this input within the request batch.
64    #[serde(skip_serializing_if = "Option::is_none")]
65    pub index: Option<usize>,
66    /// Object kind (always `embedding`).
67    #[serde(skip_serializing_if = "Option::is_none")]
68    pub object: Option<EmbeddingObjectKind>,
69    /// The embedding vector.
70    #[serde(skip_serializing_if = "Option::is_none")]
71    pub embedding: Option<Vec<f32>>,
72}
73
74/// Object kind for an embedding item.
75#[derive(Debug, Clone, Serialize, Deserialize)]
76#[serde(rename_all = "lowercase")]
77pub enum EmbeddingObjectKind {
78    /// Embedding object.
79    Embedding,
80}
81
82/// Token-usage statistics for an embeddings request.
83#[derive(Debug, Clone, Serialize, Deserialize)]
84pub struct EmbeddingUsage {
85    /// Number of prompt tokens.
86    #[serde(skip_serializing_if = "Option::is_none")]
87    pub prompt_tokens: Option<u64>,
88    /// Number of completion tokens (typically 0 for embeddings).
89    #[serde(skip_serializing_if = "Option::is_none")]
90    pub completion_tokens: Option<u64>,
91    /// Total tokens for this request.
92    #[serde(skip_serializing_if = "Option::is_none")]
93    pub total_tokens: Option<u64>,
94}
95
96#[cfg(test)]
97mod tests {
98    use super::*;
99
100    #[test]
101    fn response_requires_one_documented_non_null_field() {
102        assert!(serde_json::from_str::<EmbeddingResponse>("{}").is_err());
103        assert!(serde_json::from_str::<EmbeddingResponse>(r#"{"data":null}"#).is_err());
104        assert!(serde_json::from_str::<EmbeddingResponse>(r#"{"data":[]}"#).is_ok());
105    }
106
107    #[test]
108    fn nested_properties_follow_their_optional_schema() {
109        let item: EmbeddingData = serde_json::from_str("{}").unwrap();
110        assert!(item.index.is_none());
111        assert!(item.object.is_none());
112        assert!(item.embedding.is_none());
113        let usage: EmbeddingUsage = serde_json::from_str("{}").unwrap();
114        assert!(usage.prompt_tokens.is_none());
115    }
116}