use std::collections::HashMap;
use dynamo_runtime::protocols::annotated::AnnotationsProvider;
use serde::{Deserialize, Serialize};
use utoipa::ToSchema;
use validator::Validate;
mod aggregator;
use super::PromptTruncationSide;
pub use super::embeddings::{NvExt, NvExtProvider};
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone)]
#[serde(untagged)]
pub enum ClassificationInput {
Single(String),
Batch(Vec<String>),
Tokens(Vec<u32>),
TokenBatch(Vec<Vec<u32>>),
}
#[derive(ToSchema, Serialize, Deserialize, Validate, Debug, Clone)]
pub struct NvCreateClassifyRequest {
pub model: String,
pub input: ClassificationInput,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub request_id: Option<String>,
#[serde(default)]
pub priority: i64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub mm_processor_kwargs: Option<HashMap<String, serde_json::Value>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_salt: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub use_activation: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub add_special_tokens: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub truncate_prompt_tokens: Option<i64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub truncation_side: Option<PromptTruncationSide>,
#[serde(skip_serializing_if = "Option::is_none")]
pub nvext: Option<NvExt>,
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone)]
pub struct ClassificationData {
pub index: u32,
#[serde(default)]
pub label: Option<String>,
pub probs: Vec<f32>,
pub num_classes: u32,
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone, Default)]
pub struct ClassificationUsage {
pub prompt_tokens: u32,
pub total_tokens: u32,
#[serde(default)]
pub completion_tokens: u32,
}
#[derive(ToSchema, Serialize, Deserialize, Validate, Debug, Clone)]
pub struct NvCreateClassifyResponse {
pub id: String,
pub object: String,
pub created: u64,
pub model: String,
pub data: Vec<ClassificationData>,
pub usage: ClassificationUsage,
}
impl NvCreateClassifyResponse {
pub fn empty() -> Self {
Self {
id: String::new(),
object: "list".to_string(),
created: 0,
model: "classify".to_string(),
data: vec![],
usage: ClassificationUsage::default(),
}
}
}
impl NvExtProvider for NvCreateClassifyRequest {
fn nvext(&self) -> Option<&NvExt> {
self.nvext.as_ref()
}
}
impl AnnotationsProvider for NvCreateClassifyRequest {
fn annotations(&self) -> Option<Vec<String>> {
self.nvext
.as_ref()
.and_then(|nvext| nvext.annotations.clone())
}
fn has_annotation(&self, annotation: &str) -> bool {
self.nvext
.as_ref()
.and_then(|nvext| nvext.annotations.as_ref())
.map(|annotations| annotations.contains(&annotation.to_string()))
.unwrap_or(false)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn single_input_round_trips() {
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hello world"
}))
.unwrap();
assert!(matches!(request.input, ClassificationInput::Single(_)));
let value = serde_json::to_value(&request).unwrap();
assert_eq!(value["input"], "hello world");
}
#[test]
fn batch_input_round_trips() {
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": ["a", "b"]
}))
.unwrap();
match &request.input {
ClassificationInput::Batch(v) => assert_eq!(v.len(), 2),
_ => panic!("expected batch input"),
}
}
#[test]
fn token_inputs_parse_as_token_variants() {
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": [101, 2023, 102]
}))
.unwrap();
assert!(matches!(request.input, ClassificationInput::Tokens(_)));
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": [[101, 102], [101, 103]]
}))
.unwrap();
match &request.input {
ClassificationInput::TokenBatch(v) => assert_eq!(v.len(), 2),
other => panic!("expected token batch, got {other:?}"),
}
}
#[test]
fn use_activation_round_trips_and_defaults_to_none() {
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hello"
}))
.unwrap();
assert!(request.use_activation.is_none());
assert!(
serde_json::to_value(&request)
.unwrap()
.get("use_activation")
.is_none()
);
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hello",
"use_activation": false
}))
.unwrap();
assert_eq!(request.use_activation, Some(false));
assert_eq!(
serde_json::to_value(&request).unwrap()["use_activation"],
serde_json::json!(false)
);
}
#[test]
fn tokenization_options_round_trip() {
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hello",
"add_special_tokens": false,
"truncate_prompt_tokens": 128,
"truncation_side": "left"
}))
.unwrap();
assert_eq!(request.add_special_tokens, Some(false));
assert_eq!(request.truncate_prompt_tokens, Some(128));
assert_eq!(request.truncation_side, Some(PromptTruncationSide::Left));
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hello"
}))
.unwrap();
let value = serde_json::to_value(&request).unwrap();
assert!(value.get("add_special_tokens").is_none());
assert!(value.get("truncate_prompt_tokens").is_none());
assert!(value.get("truncation_side").is_none());
}
#[test]
fn classify_request_controls_round_trip() {
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hello",
"user": "user-1",
"request_id": "request-1",
"priority": -2,
"mm_processor_kwargs": {"do_resize": false},
"cache_salt": "salt"
}))
.unwrap();
assert_eq!(request.user.as_deref(), Some("user-1"));
assert_eq!(request.request_id.as_deref(), Some("request-1"));
assert_eq!(request.priority, -2);
assert_eq!(request.cache_salt.as_deref(), Some("salt"));
let value = serde_json::to_value(request).unwrap();
assert_eq!(value["mm_processor_kwargs"]["do_resize"], false);
assert_eq!(value["priority"], -2);
}
#[test]
fn usage_mirrors_vllm_usage_info() {
let response: NvCreateClassifyResponse = serde_json::from_value(json!({
"id": "classify-1",
"object": "list",
"created": 1,
"model": "m",
"data": [{"index": 0, "label": "entailment", "probs": [0.9, 0.1], "num_classes": 2}],
"usage": {"prompt_tokens": 3, "total_tokens": 3}
}))
.unwrap();
assert_eq!(response.usage.completion_tokens, 0);
let value = serde_json::to_value(&response).unwrap();
assert_eq!(value["usage"]["completion_tokens"], 0);
}
#[test]
fn omitted_nvext_is_not_serialized() {
let request: NvCreateClassifyRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hello"
}))
.unwrap();
assert!(request.nvext.is_none());
let value = serde_json::to_value(request).unwrap();
assert!(value.get("nvext").is_none());
}
}