Skip to main content

vtcode_llm/providers/common/
token_count.rs

1//! Exact prompt token-count payloads, requests, and response parsing.
2
3use super::read_provider_error_body;
4use crate::error_display;
5use crate::provider::LLMError;
6use serde_json::{Value, json};
7
8/// Remove generation-only controls from a payload before exact prompt-token counting.
9/// Token-count endpoints generally require only prompt-side fields.
10pub fn strip_generation_controls_for_token_count(payload: &mut Value) {
11    let Some(root) = payload.as_object_mut() else {
12        return;
13    };
14
15    for key in [
16        "stream",
17        "temperature",
18        "top_p",
19        "frequency_penalty",
20        "presence_penalty",
21        "stop",
22        "max_tokens",
23        "max_output_tokens",
24        "n",
25        "seed",
26        "tool_choice",
27        "parallel_tool_config",
28        "response_format",
29        "reasoning_effort",
30        "metadata",
31        "prompt_cache_key",
32    ] {
33        root.remove(key);
34    }
35}
36
37#[inline]
38fn parse_u32_value(value: &Value) -> Option<u32> {
39    value
40        .as_u64()
41        .and_then(|n| u32::try_from(n).ok())
42        .or_else(|| {
43            value
44                .as_i64()
45                .and_then(|n| u64::try_from(n).ok())
46                .and_then(|n| u32::try_from(n).ok())
47        })
48        .or_else(|| value.as_str().and_then(|s| s.parse::<u32>().ok()))
49}
50
51#[inline]
52fn value_at_path<'a>(value: &'a Value, path: &[&str]) -> Option<&'a Value> {
53    let mut cursor = value;
54    for segment in path {
55        cursor = cursor.get(*segment)?;
56    }
57    Some(cursor)
58}
59
60/// Parse prompt/input token counts from common response shapes.
61pub fn parse_prompt_tokens_from_count_response(value: &Value) -> Option<u32> {
62    const CANDIDATE_PATHS: &[&[&str]] = &[
63        &["prompt_tokens"],
64        &["input_tokens"],
65        &["token_count"],
66        &["usage", "prompt_tokens"],
67        &["usage", "input_tokens"],
68        &["data", "prompt_tokens"],
69        &["data", "input_tokens"],
70        &["data", "token_count"],
71        &["usage", "total_tokens"],
72        &["data", "total_tokens"],
73        &["total_tokens"],
74    ];
75
76    for path in CANDIDATE_PATHS {
77        if let Some(parsed) = value_at_path(value, path).and_then(parse_u32_value) {
78            return Some(parsed);
79        }
80    }
81    None
82}
83
84/// Execute an exact token-count request when provider endpoint is available.
85/// Returns `Ok(None)` when endpoint appears unsupported.
86pub async fn execute_token_count_request(
87    request_builder: reqwest::RequestBuilder,
88    payload: &Value,
89    provider_name: &str,
90) -> Result<Option<Value>, LLMError> {
91    let response = request_builder.json(payload).send().await.map_err(|e| {
92        let message = error_display::format_llm_error(provider_name, &format!("Token-count network error: {e}"));
93        LLMError::Network { message, metadata: None }
94    })?;
95
96    let status = response.status();
97    if matches!(
98        status,
99        reqwest::StatusCode::BAD_REQUEST
100            | reqwest::StatusCode::UNPROCESSABLE_ENTITY
101            | reqwest::StatusCode::NOT_FOUND
102            | reqwest::StatusCode::METHOD_NOT_ALLOWED
103            | reqwest::StatusCode::NOT_IMPLEMENTED
104    ) {
105        return Ok(None);
106    }
107
108    if !status.is_success() {
109        let body = read_provider_error_body(response).await;
110        let message =
111            error_display::format_llm_error(provider_name, &format!("Token-count request failed ({status}): {body}"));
112        return Err(LLMError::Provider { message, metadata: None });
113    }
114
115    let value = response.json::<Value>().await.map_err(|e| {
116        let message =
117            error_display::format_llm_error(provider_name, &format!("Failed to parse token-count response: {e}"));
118        LLMError::Provider { message, metadata: None }
119    })?;
120
121    Ok(Some(value))
122}