use std::sync::Arc;
use http_body_util::{BodyExt, Limited};
use hyper::body::Incoming;
use hyper::{Request, Response, StatusCode};
use serde_json::{json, Value};
use crate::{
group_table, http_post_json, json_response, model_table, not_ready, BoxedBody, Shared,
};
const IMAGE_TOKENS: u64 = 4_000;
const VIDEO_TOKENS: u64 = 40_000;
const SUPPORTED: &str = "vllm";
const MAX_BODY: usize = 32 * 1024 * 1024;
pub async fn count(shared: &Arc<Shared>, req: Request<Incoming>) -> Response<BoxedBody> {
let bytes = match Limited::new(req.into_body(), MAX_BODY).collect().await {
Ok(c) => c.to_bytes(),
Err(_) => {
return json_response(
StatusCode::PAYLOAD_TOO_LARGE,
&json!({"error": format!("body over {MAX_BODY} bytes")}),
)
}
};
let payload: Value = match serde_json::from_slice(&bytes) {
Ok(v) => v,
Err(_) => {
return json_response(
StatusCode::BAD_REQUEST,
&json!({"error": "body is not JSON"}),
)
}
};
let Some(model) = payload["model"].as_str() else {
return json_response(
StatusCode::BAD_REQUEST,
&json!({"error": "no model field in request"}),
);
};
let models = model_table(shared);
let Some((group, base)) = models.get(model) else {
let pending = not_ready(shared);
if let Some(why) = pending.get(model) {
return json_response(
StatusCode::SERVICE_UNAVAILABLE,
&json!({"error": format!("{model} is not ready: {why}")}),
);
}
return json_response(
StatusCode::NOT_FOUND,
&json!({
"error": format!("no model {model:?} is being served"),
"available": models.keys().collect::<Vec<_>>(),
}),
);
};
let provider = group_table(shared)
.get(group)
.map(|e| e.provider.clone())
.unwrap_or_default();
if provider != SUPPORTED {
let said = if provider.is_empty() {
"announced no provider".to_string()
} else {
format!("announced provider {provider:?}")
};
return json_response(
StatusCode::BAD_REQUEST,
&json!({"error": format!(
"cannot count tokens for {model}: group {group} {said}, and only \
{SUPPORTED:?} is supported. Set MENTAT_MODEL_PROVIDER on the container."
)}),
);
}
let (messages, media) = split_input(&payload);
let text = match messages.is_empty() {
true => 0,
false => {
let mut req = json!({"model": model, "messages": messages});
if let Some(tools) = payload.get("tools").filter(|t| t.is_array()) {
req["tools"] = tools.clone();
}
let url = format!("{}/tokenize", tokenize_root(base));
match http_post_json(&shared.client, &url, &req, shared.cfg.mcp_timeout).await {
Ok(v) => match v["count"].as_u64() {
Some(n) => n,
None => {
return json_response(
StatusCode::BAD_GATEWAY,
&json!({"error": format!("{group} returned no count: {v}")}),
)
}
},
Err(e) => {
return json_response(
StatusCode::BAD_GATEWAY,
&json!({"error": format!("{group} could not tokenize: {e}")}),
)
}
}
}
};
json_response(
StatusCode::OK,
&json!({"object": "response.input_tokens", "input_tokens": text + media}),
)
}
fn tokenize_root(base: &str) -> String {
let base = base.trim_end_matches('/');
base.strip_suffix("/v1").unwrap_or(base).to_string()
}
fn split_input(payload: &Value) -> (Vec<Value>, u64) {
let mut messages = Vec::new();
let mut media = 0;
if let Some(s) = payload["instructions"].as_str().filter(|s| !s.is_empty()) {
messages.push(json!({"role": "system", "content": s}));
}
match &payload["input"] {
Value::String(s) => messages.push(json!({"role": "user", "content": s})),
Value::Array(items) => {
for item in items {
if let Some(s) = item.as_str() {
messages.push(json!({"role": "user", "content": s}));
continue;
}
let (text, cost) = split_content(&item["content"]);
media += cost;
messages.push(json!({
"role": item["role"].as_str().unwrap_or("user"),
"content": text,
}));
}
}
_ => {}
}
(messages, media)
}
fn split_content(content: &Value) -> (String, u64) {
match content {
Value::String(s) => (s.clone(), 0),
Value::Array(parts) => {
let (mut text, mut media) = (String::new(), 0);
for p in parts {
match p["type"].as_str().unwrap_or_default() {
"input_image" | "image_url" | "image" => media += IMAGE_TOKENS,
"input_video" | "video_url" | "video" => media += VIDEO_TOKENS,
_ => {
if let Some(s) = p["text"].as_str() {
text.push_str(s);
}
}
}
}
(text, media)
}
_ => (String::new(), 0),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_bare_string_input_is_one_user_message() {
let (m, media) = split_input(&json!({"model": "m", "input": "hello world"}));
assert_eq!(m, vec![json!({"role": "user", "content": "hello world"})]);
assert_eq!(media, 0);
}
#[test]
fn instructions_lead_as_a_system_message() {
let (m, _) = split_input(&json!({"instructions": "Be terse.", "input": "hi"}));
assert_eq!(m[0], json!({"role": "system", "content": "Be terse."}));
assert_eq!(m[1], json!({"role": "user", "content": "hi"}));
}
#[test]
fn either_spelling_of_a_part_is_read() {
let responses = json!([
{"type": "input_text", "text": "a"},
{"type": "input_image", "image_url": "http://x/a.png"},
]);
let chat = json!([
{"type": "text", "text": "a"},
{"type": "image_url", "image_url": {"url": "http://x/a.png"}},
]);
assert_eq!(split_content(&responses), ("a".into(), IMAGE_TOKENS));
assert_eq!(split_content(&chat), ("a".into(), IMAGE_TOKENS));
}
#[test]
fn media_is_counted_per_part() {
let (text, media) = split_content(&json!([
{"type": "input_text", "text": "look: "},
{"type": "input_image", "image_url": "a"},
{"type": "input_image", "image_url": "b"},
{"type": "input_video", "video_url": "c"},
]));
assert_eq!(text, "look: ");
assert_eq!(media, 2 * IMAGE_TOKENS + VIDEO_TOKENS);
}
#[test]
fn an_image_only_input_keeps_its_message() {
let (m, media) = split_input(&json!({
"input": [{"role": "user", "content": [{"type": "input_image", "image_url": "a"}]}]
}));
assert_eq!(m, vec![json!({"role": "user", "content": ""})]);
assert_eq!(media, IMAGE_TOKENS);
}
#[test]
fn tokenize_climbs_out_of_v1() {
assert_eq!(tokenize_root("http://h:8000/v1"), "http://h:8000");
assert_eq!(tokenize_root("http://h:8000/v1/"), "http://h:8000");
assert_eq!(tokenize_root("http://h:8000"), "http://h:8000");
}
}