Skip to main content

zai_rs/model/text_tokenizer/
request.rs

1use serde::{Deserialize, Serialize};
2
3/// Tokenizer-capable models
4#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
5#[serde(rename_all = "kebab-case")]
6pub enum TokenizerModel {
7    /// glm-4.6v (current service default).
8    #[serde(rename = "glm-4.6v")]
9    #[default]
10    Glm46V,
11    /// glm-4.6.
12    #[serde(rename = "glm-4.6")]
13    Glm46,
14    /// glm-4.5.
15    #[serde(rename = "glm-4.5")]
16    Glm45,
17    /// glm-4.5-air.
18    #[serde(rename = "glm-4.5-air")]
19    Glm45Air,
20    /// glm-4-0520.
21    #[serde(rename = "glm-4-0520")]
22    Glm40520,
23    /// glm-4-long.
24    #[serde(rename = "glm-4-long")]
25    Glm4Long,
26    /// glm-4-air.
27    #[serde(rename = "glm-4-air")]
28    Glm4Air,
29    /// glm-4-flash.
30    #[serde(rename = "glm-4-flash")]
31    Glm4Flash,
32}
33
34/// One message item for tokenizer input
35#[derive(Clone, Serialize, Deserialize)]
36#[serde(tag = "role", rename_all = "lowercase")]
37pub enum TokenizerMessage {
38    /// User message.
39    User {
40        /// Message text.
41        content: String,
42    },
43    /// System instruction.
44    System {
45        /// System-instruction text.
46        content: String,
47    },
48    /// Assistant message with optional content.
49    Assistant {
50        /// Assistant reply text, if any.
51        #[serde(skip_serializing_if = "Option::is_none")]
52        content: Option<String>,
53    },
54}
55
56impl std::fmt::Debug for TokenizerMessage {
57    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
58        let (role, configured) = match self {
59            Self::User { .. } => ("user", true),
60            Self::System { .. } => ("system", true),
61            Self::Assistant { content } => ("assistant", content.is_some()),
62        };
63        formatter
64            .debug_struct("TokenizerMessage")
65            .field("role", &role)
66            .field("content_configured", &configured)
67            .finish()
68    }
69}
70
71/// Request body for tokenizer
72#[derive(Clone, Serialize, Deserialize)]
73pub struct TokenizerBody {
74    /// Model used for token counting (defaults to `glm-4.6v`).
75    pub model: TokenizerModel,
76    /// Conversation messages; at least one is required.
77    pub messages: Vec<TokenizerMessage>,
78    /// Client-provided request identifier.
79    #[serde(skip_serializing_if = "Option::is_none")]
80    pub request_id: Option<String>,
81    /// End-user identifier used for abuse monitoring.
82    #[serde(skip_serializing_if = "Option::is_none")]
83    pub user_id: Option<String>,
84}
85
86impl std::fmt::Debug for TokenizerBody {
87    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
88        formatter
89            .debug_struct("TokenizerBody")
90            .field("model", &self.model)
91            .field("messages_len", &self.messages.len())
92            .field(
93                "request_id",
94                &self.request_id.as_ref().map(|_| "[REDACTED]"),
95            )
96            .field("user_id", &self.user_id.as_ref().map(|_| "[REDACTED]"))
97            .finish()
98    }
99}
100
101impl TokenizerBody {
102    /// Create a new tokenizer body from a model and a message list.
103    pub fn new(model: TokenizerModel, messages: Vec<TokenizerMessage>) -> Self {
104        Self {
105            model,
106            messages,
107            request_id: None,
108            user_id: None,
109        }
110    }
111    /// Set the client-side request id.
112    pub fn with_request_id(mut self, v: impl Into<String>) -> Self {
113        self.request_id = Some(v.into());
114        self
115    }
116    /// Set the end-user id.
117    pub fn with_user_id(mut self, v: impl Into<String>) -> Self {
118        self.user_id = Some(v.into());
119        self
120    }
121
122    /// Validate constraints shared by direct body users and the request client.
123    pub fn validate(&self) -> crate::ZaiResult<()> {
124        if self.messages.is_empty() {
125            return Err(crate::ZaiError::ApiError {
126                code: crate::client::error::codes::SDK_VALIDATION,
127                message: "messages must not be empty".to_owned(),
128            });
129        }
130        let has_blank_content = self.messages.iter().any(|message| match message {
131            TokenizerMessage::User { content } | TokenizerMessage::System { content } => {
132                content.trim().is_empty()
133            },
134            TokenizerMessage::Assistant { content } => content
135                .as_deref()
136                .is_some_and(|content| content.trim().is_empty()),
137        });
138        if has_blank_content {
139            return Err(crate::ZaiError::ApiError {
140                code: crate::client::error::codes::SDK_VALIDATION,
141                message: "message content must not be blank when present".to_owned(),
142            });
143        }
144        if let Some(request_id) = self.request_id.as_deref()
145            && !(6..=64).contains(&request_id.chars().count())
146        {
147            return Err(crate::ZaiError::ApiError {
148                code: crate::client::error::codes::SDK_VALIDATION,
149                message: "request_id must contain between 6 and 64 characters".to_owned(),
150            });
151        }
152        Ok(())
153    }
154}
155
156#[cfg(test)]
157mod tests {
158    use super::*;
159
160    #[test]
161    fn current_default_and_model_ids_match_the_contract() {
162        assert_eq!(
163            serde_json::to_value(TokenizerModel::default()).unwrap(),
164            "glm-4.6v"
165        );
166        assert_eq!(
167            serde_json::to_value(TokenizerModel::Glm46).unwrap(),
168            "glm-4.6"
169        );
170        for (model, expected) in [
171            (TokenizerModel::Glm46V, "glm-4.6v"),
172            (TokenizerModel::Glm46, "glm-4.6"),
173            (TokenizerModel::Glm45, "glm-4.5"),
174            (TokenizerModel::Glm45Air, "glm-4.5-air"),
175            (TokenizerModel::Glm40520, "glm-4-0520"),
176            (TokenizerModel::Glm4Long, "glm-4-long"),
177            (TokenizerModel::Glm4Air, "glm-4-air"),
178            (TokenizerModel::Glm4Flash, "glm-4-flash"),
179        ] {
180            assert_eq!(serde_json::to_value(model).unwrap(), expected);
181        }
182    }
183
184    #[test]
185    fn validation_rejects_blank_content_and_short_request_ids() {
186        let body = TokenizerBody::new(
187            TokenizerModel::default(),
188            vec![TokenizerMessage::User {
189                content: " ".into(),
190            }],
191        );
192        assert!(body.validate().is_err());
193
194        let body = TokenizerBody::new(
195            TokenizerModel::default(),
196            vec![TokenizerMessage::User {
197                content: "hello".into(),
198            }],
199        )
200        .with_request_id("short");
201        assert!(body.validate().is_err());
202    }
203
204    #[test]
205    fn debug_redacts_message_content_and_identifiers() {
206        let body = TokenizerBody::new(
207            TokenizerModel::default(),
208            vec![TokenizerMessage::User {
209                content: "private tokenizer input".to_owned(),
210            }],
211        )
212        .with_request_id("private-request")
213        .with_user_id("private-user");
214        let debug = format!("{body:?}");
215        for secret in ["private tokenizer input", "private-request", "private-user"] {
216            assert!(!debug.contains(secret));
217        }
218        assert!(debug.contains("messages_len: 1"));
219    }
220}