gemini_genai_rs/tokens/
mod.rs

1//! Token counting API — countTokens and computeTokens.
2//!
3//! Feature-gated behind `tokens`.
4
5use serde::{Deserialize, Serialize};
6
7use crate::client::Client;
8use crate::client::http::HttpError;
9use crate::protocol::types::{Content, ModelId};
10use crate::transport::auth::ServiceEndpoint;
11
12/// Response from countTokens.
13#[derive(Debug, Clone, Serialize, Deserialize)]
14#[serde(rename_all = "camelCase")]
15pub struct CountTokensResponse {
16    /// Total number of tokens.
17    pub total_tokens: u32,
18    /// Cached content tokens (if applicable).
19    #[serde(default)]
20    pub cached_content_token_count: Option<u32>,
21}
22
23/// Errors from the Tokens API.
24#[derive(Debug, thiserror::Error)]
25pub enum TokensError {
26    #[error(transparent)]
27    /// Transport-level HTTP failure.
28    Http(#[from] HttpError),
29    #[error("Failed to parse response: {0}")]
30    /// Response body failed to parse.
31    Parse(#[from] serde_json::Error),
32    #[error("Auth error: {0}")]
33    /// Authentication/authorization failure.
34    Auth(#[from] crate::session::AuthError),
35}
36
37impl Client {
38    /// Count tokens for text content.
39    pub async fn count_tokens(
40        &self,
41        text: impl Into<String>,
42    ) -> Result<CountTokensResponse, TokensError> {
43        self.count_tokens_for(vec![Content::user(text)], None).await
44    }
45
46    /// Count tokens for content with optional model override.
47    pub async fn count_tokens_for(
48        &self,
49        contents: Vec<Content>,
50        model: Option<&ModelId>,
51    ) -> Result<CountTokensResponse, TokensError> {
52        let model = model.unwrap_or(self.default_model());
53        let url = self.rest_url_for(ServiceEndpoint::CountTokens, model);
54        let headers = self.auth_headers().await?;
55
56        let body = serde_json::json!({ "contents": contents });
57        let json = self.http_client().post_json(&url, headers, &body).await?;
58        Ok(serde_json::from_value(json)?)
59    }
60}
61
62#[cfg(test)]
63mod tests {
64    use super::*;
65
66    #[test]
67    fn parse_count_tokens_response() {
68        let json = serde_json::json!({
69            "totalTokens": 42
70        });
71        let resp: CountTokensResponse = serde_json::from_value(json).unwrap();
72        assert_eq!(resp.total_tokens, 42);
73        assert!(resp.cached_content_token_count.is_none());
74    }
75
76    #[test]
77    fn parse_count_tokens_with_cached() {
78        let json = serde_json::json!({
79            "totalTokens": 100,
80            "cachedContentTokenCount": 50
81        });
82        let resp: CountTokensResponse = serde_json::from_value(json).unwrap();
83        assert_eq!(resp.total_tokens, 100);
84        assert_eq!(resp.cached_content_token_count, Some(50));
85    }
86}