[model-gateway] Add tokenize/detokenize HTTP endpoints and tokenizer management (#15702)

This commit is contained in:
Simo Lin
2025-12-23 17:32:07 -08:00
committed by GitHub
parent 5c64a20da7
commit 846953d9f1
19 changed files with 1372 additions and 207 deletions
+12 -8
View File
@@ -8,7 +8,9 @@ use tracing::{debug, info};
use crate::{ use crate::{
config::RouterConfig, config::RouterConfig,
core::{ConnectionMode, JobQueue, LoadMonitor, WorkerRegistry, WorkerService}, core::{
ConnectionMode, JobQueue, LoadMonitor, WorkerRegistry, WorkerService, UNKNOWN_MODEL_ID,
},
data_connector::{ data_connector::{
create_storage, ConversationItemStorage, ConversationStorage, ResponseStorage, create_storage, ConversationItemStorage, ConversationStorage, ResponseStorage,
}, },
@@ -451,25 +453,27 @@ impl AppContextBuilder {
/// Create tokenizer registry and optionally load tokenizer /// Create tokenizer registry and optionally load tokenizer
/// If a tokenizer is successfully loaded, it is registered with a key derived from /// If a tokenizer is successfully loaded, it is registered with a key derived from
/// tokenizer_path or model_path (falling back to "unknown" if neither exists). /// tokenizer_path or model_path (falling back to UNKNOWN_MODEL_ID if neither exists).
fn with_tokenizer_registry(mut self, config: &RouterConfig) -> Result<Self, String> { fn with_tokenizer_registry(mut self, config: &RouterConfig) -> Result<Self, String> {
// Create empty tokenizer registry // Create empty tokenizer registry
let registry = Arc::new(TokenizerRegistry::new()); let registry = Arc::new(TokenizerRegistry::new());
// Try to load router-level tokenizer if path is provided // Try to load router-level tokenizer if path is provided
if let Some(tokenizer) = Self::maybe_tokenizer(config)? { if let Some(tokenizer) = Self::maybe_tokenizer(config)? {
// Determine registration key: prefer tokenizer_path, then model_path, finally "unknown" // Determine registration key: prefer tokenizer_path, then model_path, finally UNKNOWN_MODEL_ID
let tokenizer_key = config let source = config
.tokenizer_path .tokenizer_path
.as_ref() .as_ref()
.or(config.model_path.as_ref()) .or(config.model_path.as_ref())
.map(|s| s.as_str()) .map(|s| s.as_str())
.unwrap_or("unknown"); .unwrap_or(UNKNOWN_MODEL_ID);
registry.register(tokenizer_key, tokenizer.clone()); let tokenizer_id = TokenizerRegistry::generate_id();
registry.register(&tokenizer_id, source, source, tokenizer.clone());
info!( info!(
"Tokenizer loaded and registered with key '{}' (vocab_size: {})", "Tokenizer loaded and registered with name '{}' id={} (vocab_size: {})",
tokenizer_key, source,
tokenizer_id,
tokenizer.vocab_size() tokenizer.vocab_size()
); );
} }
+59 -3
View File
@@ -17,8 +17,8 @@ use crate::{
app_context::AppContext, app_context::AppContext,
config::{RouterConfig, RoutingMode}, config::{RouterConfig, RoutingMode},
core::steps::{ core::steps::{
McpServerConfigRequest, WasmModuleConfigRequest, WasmModuleRemovalRequest, McpServerConfigRequest, TokenizerConfigRequest, TokenizerRemovalRequest,
WorkerRemovalRequest, WasmModuleConfigRequest, WasmModuleRemovalRequest, WorkerRemovalRequest,
}, },
mcp::McpConfig, mcp::McpConfig,
protocols::worker_spec::{JobStatus, WorkerConfigRequest, WorkerUpdateRequest}, protocols::worker_spec::{JobStatus, WorkerConfigRequest, WorkerUpdateRequest},
@@ -53,6 +53,12 @@ pub enum Job {
RemoveWasmModule { RemoveWasmModule {
request: Box<WasmModuleRemovalRequest>, request: Box<WasmModuleRemovalRequest>,
}, },
AddTokenizer {
config: Box<TokenizerConfigRequest>,
},
RemoveTokenizer {
request: Box<TokenizerRemovalRequest>,
},
} }
impl Job { impl Job {
@@ -67,10 +73,12 @@ impl Job {
Job::RegisterMcpServer { .. } => "RegisterMcpServer", Job::RegisterMcpServer { .. } => "RegisterMcpServer",
Job::AddWasmModule { .. } => "AddWasmModule", Job::AddWasmModule { .. } => "AddWasmModule",
Job::RemoveWasmModule { .. } => "RemoveWasmModule", Job::RemoveWasmModule { .. } => "RemoveWasmModule",
Job::AddTokenizer { .. } => "AddTokenizer",
Job::RemoveTokenizer { .. } => "RemoveTokenizer",
} }
} }
/// Get worker URL, MCP server name, or WASM module identifier for logging and status tracking /// Get worker URL, MCP server name, WASM module, or tokenizer identifier for logging and status tracking
pub fn worker_url(&self) -> &str { pub fn worker_url(&self) -> &str {
match self { match self {
Job::AddWorker { config } => &config.url, Job::AddWorker { config } => &config.url,
@@ -81,6 +89,8 @@ impl Job {
Job::RegisterMcpServer { config } => &config.name, Job::RegisterMcpServer { config } => &config.name,
Job::AddWasmModule { config } => &config.descriptor.name, Job::AddWasmModule { config } => &config.descriptor.name,
Job::RemoveWasmModule { request } => &request.uuid_string, Job::RemoveWasmModule { request } => &request.uuid_string,
Job::AddTokenizer { config } => &config.id,
Job::RemoveTokenizer { request } => &request.id,
} }
} }
} }
@@ -658,6 +668,52 @@ impl JobQueue {
) )
.await .await
} }
Job::AddTokenizer { config } => {
let engine = context
.workflow_engine
.get()
.ok_or_else(|| "Workflow engine not initialized".to_string())?;
let mut workflow_context = WorkflowContext::new(WorkflowInstanceId::new());
let config_arc: Arc<TokenizerConfigRequest> = Arc::new(*config.clone());
workflow_context.set_arc("tokenizer_config", config_arc);
workflow_context.set_arc("app_context", Arc::clone(context));
let instance_id = engine
.start_workflow(WorkflowId::new("tokenizer_registration"), workflow_context)
.await
.map_err(|e| {
format!("Failed to start tokenizer registration workflow: {:?}", e)
})?;
debug!(
"Started tokenizer registration workflow for '{}' id={} (instance: {})",
config.name, config.id, instance_id
);
// Allow up to 10 minutes for HuggingFace downloads
let timeout_duration = Duration::from_secs(600);
Self::wait_for_workflow_completion(
engine,
instance_id,
&config.id,
timeout_duration,
)
.await
}
Job::RemoveTokenizer { request } => {
// Tokenizer removal is synchronous and fast
if let Some(entry) = context.tokenizer_registry.remove_by_id(&request.id) {
info!(
"Successfully removed tokenizer '{}' (id: {})",
entry.name, entry.id
);
Ok(format!("Tokenizer '{}' removed successfully", entry.name))
} else {
Err(format!("Tokenizer with id '{}' not found", request.id))
}
}
} }
} }
+2 -5
View File
@@ -9,11 +9,8 @@
//! - Workflow steps for multi-step operations //! - Workflow steps for multi-step operations
//! - Common utilities //! - Common utilities
/// Default model identifier used when no model is specified. // Re-export UNKNOWN_MODEL_ID from protocols for use throughout core
/// pub use crate::protocols::UNKNOWN_MODEL_ID;
/// This constant should be used instead of hardcoded "unknown" strings
/// throughout the codebase for consistency.
pub const UNKNOWN_MODEL_ID: &str = "unknown";
pub mod circuit_breaker; pub mod circuit_breaker;
pub mod error; pub mod error;
+6 -1
View File
@@ -4,9 +4,10 @@
//! - Worker management (registration, removal, updates) //! - Worker management (registration, removal, updates)
//! - MCP server registration //! - MCP server registration
//! - WASM module registration and removal //! - WASM module registration and removal
//! - Future: Tokenizer fetching, LoRA updates, etc. //! - Tokenizer registration
pub mod mcp_registration; pub mod mcp_registration;
pub mod tokenizer_registration;
pub mod wasm_module_registration; pub mod wasm_module_registration;
pub mod wasm_module_removal; pub mod wasm_module_removal;
pub mod worker; pub mod worker;
@@ -66,6 +67,10 @@ pub use mcp_registration::{
create_mcp_registration_workflow, ConnectMcpServerStep, DiscoverMcpInventoryStep, create_mcp_registration_workflow, ConnectMcpServerStep, DiscoverMcpInventoryStep,
McpServerConfigRequest, RegisterMcpServerStep, ValidateRegistrationStep, McpServerConfigRequest, RegisterMcpServerStep, ValidateRegistrationStep,
}; };
pub use tokenizer_registration::{
create_tokenizer_registration_workflow, LoadTokenizerStep, TokenizerConfigRequest,
TokenizerRemovalRequest, ValidateTokenizerConfigStep,
};
pub use wasm_module_registration::{ pub use wasm_module_registration::{
create_wasm_module_registration_workflow, CalculateHashStep, CheckDuplicateStep, create_wasm_module_registration_workflow, CalculateHashStep, CheckDuplicateStep,
LoadWasmBytesStep, RegisterModuleStep, ValidateDescriptorStep, ValidateWasmComponentStep, LoadWasmBytesStep, RegisterModuleStep, ValidateDescriptorStep, ValidateWasmComponentStep,
@@ -0,0 +1,228 @@
//! Tokenizer registration workflow
//!
//! This module provides a workflow for registering tokenizers asynchronously.
//! Tokenizers can be loaded from local paths or downloaded from HuggingFace.
use std::{sync::Arc, time::Duration};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use tracing::{debug, error, info};
use crate::{app_context::AppContext, tokenizer::factory, workflow::*};
/// Configuration for adding a tokenizer
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenizerConfigRequest {
/// Pre-generated UUID for this tokenizer
pub id: String,
/// User-provided name
pub name: String,
/// Source: either a local path or HuggingFace model ID
pub source: String,
/// Optional path to chat template file
pub chat_template_path: Option<String>,
}
/// Configuration for removing a tokenizer
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenizerRemovalRequest {
/// UUID of the tokenizer to remove
pub id: String,
}
// ============================================================================
// Workflow Steps
// ============================================================================
/// Step 1: Validate the tokenizer configuration
pub struct ValidateTokenizerConfigStep;
#[async_trait]
impl StepExecutor for ValidateTokenizerConfigStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let config: Arc<TokenizerConfigRequest> = context
.get("tokenizer_config")
.ok_or_else(|| WorkflowError::ContextValueNotFound("tokenizer_config".to_string()))?;
let app_context: Arc<AppContext> = context
.get("app_context")
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
debug!(
"Validating tokenizer config: name={}, source={}",
config.name, config.source
);
// Validate name is not empty
if config.name.is_empty() {
return Err(WorkflowError::StepFailed {
step_id: StepId::new("validate_config"),
message: "Tokenizer name cannot be empty".to_string(),
});
}
// Validate source is not empty
if config.source.is_empty() {
return Err(WorkflowError::StepFailed {
step_id: StepId::new("validate_config"),
message: "Tokenizer source cannot be empty".to_string(),
});
}
// Check if tokenizer already exists
if app_context.tokenizer_registry.contains(&config.name) {
return Err(WorkflowError::StepFailed {
step_id: StepId::new("validate_config"),
message: format!("Tokenizer '{}' already exists", config.name),
});
}
debug!("Tokenizer config validated successfully");
Ok(StepResult::Success)
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
false // Validation errors are not retryable
}
}
/// Step 2: Load the tokenizer from source (local path or HuggingFace)
pub struct LoadTokenizerStep;
#[async_trait]
impl StepExecutor for LoadTokenizerStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let config: Arc<TokenizerConfigRequest> = context
.get("tokenizer_config")
.ok_or_else(|| WorkflowError::ContextValueNotFound("tokenizer_config".to_string()))?;
let app_context: Arc<AppContext> = context
.get("app_context")
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
info!(
"Loading tokenizer '{}' (id: {}) from source: {}",
config.name, config.id, config.source
);
// Load the tokenizer using the registry's load method (handles deduplication)
let result = app_context
.tokenizer_registry
.load(&config.id, &config.name, &config.source, || {
let source = config.source.clone();
let chat_template = config.chat_template_path.clone();
async move {
factory::create_tokenizer_async_with_chat_template(
&source,
chat_template.as_deref(),
)
.await
.map_err(|e| format!("Failed to load tokenizer: {}", e))
}
})
.await;
match result {
Ok(loaded_id) => {
// Get vocab size for logging
let vocab_size = app_context
.tokenizer_registry
.get_by_id(&loaded_id)
.map(|e| e.tokenizer.vocab_size());
info!(
"Successfully loaded tokenizer '{}' (id: {}) with vocab_size: {:?}",
config.name, loaded_id, vocab_size
);
// Store vocab size in context for later use
if let Some(size) = vocab_size {
context.set("vocab_size", size);
}
Ok(StepResult::Success)
}
Err(e) => {
error!("Failed to load tokenizer '{}': {}", config.name, e);
Err(WorkflowError::StepFailed {
step_id: StepId::new("load_tokenizer"),
message: e,
})
}
}
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
true // Network/IO errors are retryable
}
}
// ============================================================================
// Workflow Builder
// ============================================================================
/// Create the tokenizer registration workflow
///
/// This workflow:
/// - Validates the tokenizer configuration
/// - Loads the tokenizer from local path or HuggingFace
///
/// Workflow configuration:
/// - ValidateConfig: No retry, 5s timeout (fast validation)
/// - LoadTokenizer: 3 retries, 5min timeout (may need to download from HuggingFace)
pub fn create_tokenizer_registration_workflow() -> WorkflowDefinition {
WorkflowDefinition::new("tokenizer_registration", "Tokenizer Registration")
.add_step(
StepDefinition::new(
"validate_config",
"Validate Configuration",
Arc::new(ValidateTokenizerConfigStep),
)
.with_timeout(Duration::from_secs(5))
.with_failure_action(FailureAction::FailWorkflow),
)
.add_step(
StepDefinition::new(
"load_tokenizer",
"Load Tokenizer",
Arc::new(LoadTokenizerStep),
)
.with_retry(RetryPolicy {
max_attempts: 3,
backoff: BackoffStrategy::Fixed(Duration::from_secs(2)),
})
.with_timeout(Duration::from_secs(300)) // 5 min for HuggingFace downloads
.with_failure_action(FailureAction::FailWorkflow)
.depends_on(&["validate_config"]),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tokenizer_config_request_serialization() {
let config = TokenizerConfigRequest {
id: "test-uuid-1234".to_string(),
name: "test-model".to_string(),
source: "meta-llama/Llama-2-7b-hf".to_string(),
chat_template_path: None,
};
let json = serde_json::to_string(&config).unwrap();
let parsed: TokenizerConfigRequest = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.id, "test-uuid-1234");
assert_eq!(parsed.name, "test-model");
assert_eq!(parsed.source, "meta-llama/Llama-2-7b-hf");
assert!(parsed.chat_template_path.is_none());
}
#[test]
fn test_workflow_creation() {
let workflow = create_tokenizer_registration_workflow();
assert_eq!(workflow.id.to_string(), "tokenizer_registration");
}
}
@@ -10,7 +10,7 @@ use crate::{
app_context::AppContext, app_context::AppContext,
core::{ core::{
BasicWorkerBuilder, CircuitBreakerConfig, ConnectionMode, DPAwareWorkerBuilder, BasicWorkerBuilder, CircuitBreakerConfig, ConnectionMode, DPAwareWorkerBuilder,
HealthConfig, ModelCard, RuntimeType, Worker, WorkerType, HealthConfig, ModelCard, RuntimeType, Worker, WorkerType, UNKNOWN_MODEL_ID,
}, },
protocols::worker_spec::WorkerConfigRequest, protocols::worker_spec::WorkerConfigRequest,
workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
@@ -62,15 +62,15 @@ impl StepExecutor for CreateLocalWorkerStep {
final_labels.insert(key.clone(), value.clone()); final_labels.insert(key.clone(), value.clone());
} }
// Determine model_id: config > served_model_name > model_path > "unknown" // Determine model_id: config > served_model_name > model_path > UNKNOWN_MODEL_ID
let model_id = config let model_id = config
.model_id .model_id
.clone() .clone()
.or_else(|| final_labels.get("served_model_name").cloned()) .or_else(|| final_labels.get("served_model_name").cloned())
.or_else(|| final_labels.get("model_path").cloned()) .or_else(|| final_labels.get("model_path").cloned())
.unwrap_or_else(|| "unknown".to_string()); .unwrap_or_else(|| UNKNOWN_MODEL_ID.to_string());
if model_id != "unknown" { if model_id != UNKNOWN_MODEL_ID {
debug!("Using model_id: {}", model_id); debug!("Using model_id: {}", model_id);
} }
@@ -7,6 +7,7 @@ use tracing::debug;
use super::discover_metadata::get_server_info; use super::discover_metadata::get_server_info;
use crate::{ use crate::{
core::UNKNOWN_MODEL_ID,
protocols::worker_spec::WorkerConfigRequest, protocols::worker_spec::WorkerConfigRequest,
workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -34,7 +35,7 @@ pub async fn get_dp_info(url: &str, api_key: Option<&str>) -> Result<DpInfo, Str
info.model_path info.model_path
.and_then(|path| path.split('/').next_back().map(|s| s.to_string())) .and_then(|path| path.split('/').next_back().map(|s| s.to_string()))
}) })
.unwrap_or_else(|| "unknown".to_string()); .unwrap_or_else(|| UNKNOWN_MODEL_ID.to_string());
Ok(DpInfo { dp_size, model_id }) Ok(DpInfo { dp_size, model_id })
} }
@@ -8,7 +8,7 @@ use tracing::{debug, warn};
use crate::{ use crate::{
app_context::AppContext, app_context::AppContext,
core::Worker, core::Worker,
tokenizer::factory, tokenizer::{factory, TokenizerRegistry},
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -47,10 +47,14 @@ impl StepExecutor for RegisterTokenizerStep {
model_id, tokenizer_path model_id, tokenizer_path
); );
// Generate ID for this tokenizer
let tokenizer_id = TokenizerRegistry::generate_id();
let source = tokenizer_path.clone();
// Load tokenizer with thread safe lock // Load tokenizer with thread safe lock
if let Err(e) = app_context if let Err(e) = app_context
.tokenizer_registry .tokenizer_registry
.load(&model_id, || async move { .load(&tokenizer_id, &model_id, &source, || async move {
factory::create_tokenizer_async(&tokenizer_path.to_string()) factory::create_tokenizer_async(&tokenizer_path.to_string())
.await .await
.map_err(|e| e.to_string()) .map_err(|e| e.to_string())
@@ -59,12 +63,12 @@ impl StepExecutor for RegisterTokenizerStep {
{ {
warn!( warn!(
"Failed to load tokenizer for model {} from {}: {}", "Failed to load tokenizer for model {} from {}: {}",
model_id, tokenizer_path, e model_id, source, e
); );
} else { } else {
debug!( debug!(
"Successfully registered tokenizer for model {} from {}", "Successfully registered tokenizer for model {} from {}",
model_id, tokenizer_path model_id, source
); );
} }
} }
+2 -1
View File
@@ -15,6 +15,7 @@ use tokio::{sync::OnceCell, time};
use super::{ use super::{
CircuitBreaker, Endpoint, ModelCard, ModelType, ProviderType, WorkerError, WorkerResult, CircuitBreaker, Endpoint, ModelCard, ModelType, ProviderType, WorkerError, WorkerResult,
UNKNOWN_MODEL_ID,
}; };
use crate::{ use crate::{
core::{BasicWorkerBuilder, DPAwareWorkerBuilder}, core::{BasicWorkerBuilder, DPAwareWorkerBuilder},
@@ -180,7 +181,7 @@ pub trait Worker: Send + Sync + fmt::Debug {
// Fall back to labels // Fall back to labels
self.metadata().labels.get("model_id").map(|s| s.as_str()) self.metadata().labels.get("model_id").map(|s| s.as_str())
}) })
.unwrap_or("unknown") .unwrap_or(UNKNOWN_MODEL_ID)
} }
/// Get the priority of this worker (higher value = higher priority) /// Get the priority of this worker (higher value = higher priority)
+3 -1
View File
@@ -4,13 +4,15 @@ use serde::{Deserialize, Serialize};
use serde_json::Value; use serde_json::Value;
use validator; use validator;
use super::UNKNOWN_MODEL_ID;
// ============================================================================ // ============================================================================
// Default value helpers // Default value helpers
// ============================================================================ // ============================================================================
/// Default model value when not specified /// Default model value when not specified
pub(crate) fn default_model() -> String { pub(crate) fn default_model() -> String {
"unknown".to_string() UNKNOWN_MODEL_ID.to_string()
} }
/// Helper function for serde default value (returns true) /// Helper function for serde default value (returns true)
+7
View File
@@ -1,6 +1,12 @@
// Protocol definitions and validation for various LLM APIs // Protocol definitions and validation for various LLM APIs
// This module provides a structured approach to handling different API protocols // This module provides a structured approach to handling different API protocols
/// Default model identifier used when no model is specified.
///
/// This constant should be used instead of hardcoded "unknown" strings
/// throughout the codebase for consistency.
pub const UNKNOWN_MODEL_ID: &str = "unknown";
pub mod builders; pub mod builders;
pub mod chat; pub mod chat;
pub mod classify; pub mod classify;
@@ -14,5 +20,6 @@ pub mod parser;
pub mod rerank; pub mod rerank;
pub mod responses; pub mod responses;
pub mod sampling_params; pub mod sampling_params;
pub mod tokenize;
pub mod validated; pub mod validated;
pub mod worker_spec; pub mod worker_spec;
+279
View File
@@ -0,0 +1,279 @@
//! Tokenize and Detokenize API protocol types
//!
//! These types mirror the SGLang Python implementation for compatibility.
//! See: python/sglang/srt/entrypoints/openai/protocol.py
use serde::{Deserialize, Serialize};
use super::UNKNOWN_MODEL_ID;
// ============================================================================
// Tokenize API
// ============================================================================
/// Request schema for the /v1/tokenize endpoint
///
/// Supports both single string and batch tokenization.
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct TokenizeRequest {
/// Model name for tokenizer selection
#[serde(default = "default_model_name")]
pub model: String,
/// Text(s) to tokenize - can be a single string or array of strings
pub prompt: StringOrArray,
}
/// Response schema for the /v1/tokenize endpoint
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenizeResponse {
/// Token IDs - single list for single input, nested list for batch
pub tokens: TokensResult,
/// Token count(s) - single int for single input, list for batch
pub count: CountResult,
/// Character count(s) of input - single int for single input, list for batch
pub char_count: CountResult,
}
/// Token IDs result - either single or batch
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum TokensResult {
Single(Vec<u32>),
Batch(Vec<Vec<u32>>),
}
/// Count result - either single or batch
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum CountResult {
Single(i32),
Batch(Vec<i32>),
}
// ============================================================================
// Detokenize API
// ============================================================================
/// Request schema for the /v1/detokenize endpoint
///
/// Supports both single sequence and batch detokenization.
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct DetokenizeRequest {
/// Model name for tokenizer selection
#[serde(default = "default_model_name")]
pub model: String,
/// Token IDs to detokenize - single list or batch (list of lists)
pub tokens: TokensInput,
/// Whether to skip special tokens (e.g., padding or EOS) during decoding
#[serde(default = "default_true")]
pub skip_special_tokens: bool,
}
/// Token input - either single sequence or batch
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(untagged)]
pub enum TokensInput {
/// Single sequence of token IDs
Single(Vec<u32>),
/// Batch of token sequences
Batch(Vec<Vec<u32>>),
}
impl TokensInput {
/// Check if this is a batch input
pub fn is_batch(&self) -> bool {
matches!(self, TokensInput::Batch(_))
}
/// Get the sequences (always returns a vec of vecs for uniform processing)
pub fn sequences(&self) -> Vec<&[u32]> {
match self {
TokensInput::Single(seq) => vec![seq.as_slice()],
TokensInput::Batch(seqs) => seqs.iter().map(|s| s.as_slice()).collect(),
}
}
}
/// Response schema for the /v1/detokenize endpoint
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DetokenizeResponse {
/// Decoded text - single string for single input, list for batch
pub text: TextResult,
}
/// Text result - either single or batch
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum TextResult {
Single(String),
Batch(Vec<String>),
}
// ============================================================================
// Tokenizer Management API
// ============================================================================
/// Request schema for adding a tokenizer
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct AddTokenizerRequest {
/// Name to register the tokenizer under
pub name: String,
/// Source: either a local path or HuggingFace model ID
pub source: String,
/// Optional path to chat template file
#[serde(skip_serializing_if = "Option::is_none")]
pub chat_template_path: Option<String>,
}
/// Response schema for adding a tokenizer (async)
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AddTokenizerResponse {
/// Unique identifier for the tokenizer (UUID)
pub id: String,
/// Status of the request: "pending", "processing", "completed", "failed"
pub status: String,
pub message: String,
/// Vocabulary size of the loaded tokenizer (only set on completion)
#[serde(skip_serializing_if = "Option::is_none")]
pub vocab_size: Option<usize>,
}
/// Response schema for listing tokenizers
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ListTokenizersResponse {
pub tokenizers: Vec<TokenizerInfo>,
}
/// Information about a registered tokenizer
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenizerInfo {
/// Unique identifier (UUID)
pub id: String,
/// User-provided name
pub name: String,
/// Source path or HuggingFace model ID
pub source: String,
pub vocab_size: usize,
}
/// Request schema for removing a tokenizer
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct RemoveTokenizerRequest {
/// Name of the tokenizer to remove
pub name: String,
}
/// Response schema for removing a tokenizer
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RemoveTokenizerResponse {
pub success: bool,
pub message: String,
}
// ============================================================================
// Helper Types
// ============================================================================
/// String or array of strings (for flexible input)
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(untagged)]
pub enum StringOrArray {
Single(String),
Array(Vec<String>),
}
impl StringOrArray {
/// Check if this is a batch (array) input
pub fn is_batch(&self) -> bool {
matches!(self, StringOrArray::Array(_))
}
/// Get all strings as a slice (converts single to vec)
pub fn as_strings(&self) -> Vec<&str> {
match self {
StringOrArray::Single(s) => vec![s.as_str()],
StringOrArray::Array(arr) => arr.iter().map(|s| s.as_str()).collect(),
}
}
}
// ============================================================================
// Default Functions
// ============================================================================
fn default_model_name() -> String {
UNKNOWN_MODEL_ID.to_string()
}
fn default_true() -> bool {
true
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tokenize_request_single() {
let json = r#"{"prompt": "Hello world"}"#;
let req: TokenizeRequest = serde_json::from_str(json).unwrap();
assert_eq!(req.model, "unknown");
assert!(matches!(req.prompt, StringOrArray::Single(_)));
}
#[test]
fn test_tokenize_request_batch() {
let json = r#"{"model": "llama", "prompt": ["Hello", "World"]}"#;
let req: TokenizeRequest = serde_json::from_str(json).unwrap();
assert_eq!(req.model, "llama");
assert!(matches!(req.prompt, StringOrArray::Array(_)));
}
#[test]
fn test_detokenize_request_single() {
let json = r#"{"tokens": [1, 2, 3]}"#;
let req: DetokenizeRequest = serde_json::from_str(json).unwrap();
assert!(matches!(req.tokens, TokensInput::Single(_)));
assert!(req.skip_special_tokens);
}
#[test]
fn test_detokenize_request_batch() {
let json = r#"{"tokens": [[1, 2], [3, 4, 5]], "skip_special_tokens": false}"#;
let req: DetokenizeRequest = serde_json::from_str(json).unwrap();
assert!(matches!(req.tokens, TokensInput::Batch(_)));
assert!(!req.skip_special_tokens);
}
#[test]
fn test_tokenize_response_single() {
let resp = TokenizeResponse {
tokens: TokensResult::Single(vec![1, 2, 3]),
count: CountResult::Single(3),
char_count: CountResult::Single(11),
};
let json = serde_json::to_string(&resp).unwrap();
assert!(json.contains("[1,2,3]"));
assert!(json.contains("\"count\":3"));
assert!(json.contains("\"char_count\":11"));
}
#[test]
fn test_tokenize_response_batch() {
let resp = TokenizeResponse {
tokens: TokensResult::Batch(vec![vec![1, 2], vec![3, 4, 5]]),
count: CountResult::Batch(vec![2, 3]),
char_count: CountResult::Batch(vec![5, 5]),
};
let json = serde_json::to_string(&resp).unwrap();
assert!(json.contains("[[1,2],[3,4,5]]"));
assert!(json.contains("[2,3]"));
}
}
@@ -6,6 +6,8 @@ use std::collections::HashMap;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use super::UNKNOWN_MODEL_ID;
/// Worker configuration for API requests /// Worker configuration for API requests
#[derive(Debug, Clone, Deserialize, Serialize)] #[derive(Debug, Clone, Deserialize, Serialize)]
pub struct WorkerConfigRequest { pub struct WorkerConfigRequest {
@@ -175,13 +177,13 @@ impl WorkerInfo {
Self { Self {
id: worker_id.to_string(), id: worker_id.to_string(),
url, url,
model_id: "unknown".to_string(), model_id: UNKNOWN_MODEL_ID.to_string(),
priority: 0, priority: 0,
cost: 1.0, cost: 1.0,
worker_type: "unknown".to_string(), worker_type: UNKNOWN_MODEL_ID.to_string(),
is_healthy: false, is_healthy: false,
load: 0, load: 0,
connection_mode: "unknown".to_string(), connection_mode: UNKNOWN_MODEL_ID.to_string(),
runtime_type: None, runtime_type: None,
tokenizer_path: None, tokenizer_path: None,
reasoning_parser: None, reasoning_parser: None,
@@ -19,6 +19,7 @@ use crate::{
ResponseReasoningContent::ReasoningText, ResponseStatus, ResponsesRequest, ResponseReasoningContent::ReasoningText, ResponseStatus, ResponsesRequest,
ResponsesResponse, ResponsesUsage, StringOrContentParts, TextConfig, TextFormat, ResponsesResponse, ResponsesUsage, StringOrContentParts, TextConfig, TextFormat,
}, },
UNKNOWN_MODEL_ID,
}, },
routers::grpc::common::responses::utils::extract_tools_from_response_tools, routers::grpc::common::responses::utils::extract_tools_from_response_tools,
}; };
@@ -171,7 +172,7 @@ pub fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest
Ok(ChatCompletionRequest { Ok(ChatCompletionRequest {
messages, messages,
model: if req.model.is_empty() { model: if req.model.is_empty() {
"unknown".to_string() UNKNOWN_MODEL_ID.to_string()
} else { } else {
req.model.clone() req.model.clone()
}, },
+1
View File
@@ -30,6 +30,7 @@ pub mod http;
pub mod openai; pub mod openai;
pub mod parse; pub mod parse;
pub mod router_manager; pub mod router_manager;
pub mod tokenize;
pub use factory::RouterFactory; pub use factory::RouterFactory;
// Re-export HTTP routers for convenience // Re-export HTTP routers for convenience
@@ -0,0 +1,441 @@
//! Tokenize and detokenize handlers
//!
//! Provides tokenization, detokenization, and tokenizer management operations.
//! These handlers use the TokenizerRegistry for tokenizer storage and retrieval.
use std::sync::Arc;
use axum::{
http::StatusCode,
response::{IntoResponse, Response},
Json,
};
use tracing::{debug, error, warn};
use crate::{
app_context::AppContext,
core::{steps::TokenizerConfigRequest, Job},
protocols::tokenize::{
AddTokenizerRequest, AddTokenizerResponse, CountResult, DetokenizeRequest,
DetokenizeResponse, ListTokenizersResponse, RemoveTokenizerResponse, TextResult,
TokenizeRequest, TokenizeResponse, TokenizerInfo, TokensResult,
},
tokenizer::{registry::TokenizerEntry, traits::Tokenizer, TokenizerRegistry},
};
/// Helper to create error responses
fn error_response(status: StatusCode, message: &str, error_type: &str) -> Response {
(
status,
Json(serde_json::json!({
"error": {
"message": message,
"type": error_type
}
})),
)
.into_response()
}
/// Get a tokenizer by model name, with fallback strategies
fn get_tokenizer(registry: &TokenizerRegistry, model: &str) -> Result<Arc<dyn Tokenizer>, String> {
// First, try exact match (by name or ID)
if let Some(tokenizer) = registry.get(model) {
debug!("Found tokenizer for model: {}", model);
return Ok(tokenizer);
}
// Try "default" if model is "default" or empty
if model == "default" || model.is_empty() {
// Try to find any tokenizer as fallback
let entries = registry.list();
if let Some(first) = entries.first() {
debug!(
"Using first available tokenizer '{}' as default",
first.name
);
return Ok(first.tokenizer.clone());
}
}
// List available tokenizers for error message
let entries = registry.list();
if entries.is_empty() {
Err("No tokenizers available. Use POST /v1/tokenizers to add one.".to_string())
} else {
let names: Vec<&str> = entries.iter().map(|e| e.name.as_str()).collect();
Err(format!(
"Tokenizer for model '{}' not found. Available: {}",
model,
names.join(", ")
))
}
}
// ============================================================================
// Tokenize / Detokenize Handlers
// ============================================================================
/// Handle POST /v1/tokenize
pub async fn tokenize(registry: &Arc<TokenizerRegistry>, request: TokenizeRequest) -> Response {
debug!("Tokenize request for model: {}", request.model);
let tokenizer = match get_tokenizer(registry, &request.model) {
Ok(t) => t,
Err(e) => {
return error_response(StatusCode::BAD_REQUEST, &e, "tokenizer_not_found");
}
};
let texts = request.prompt.as_strings();
let is_batch = request.prompt.is_batch();
// Tokenize each text
let mut all_tokens: Vec<Vec<u32>> = Vec::with_capacity(texts.len());
let mut all_counts: Vec<i32> = Vec::with_capacity(texts.len());
let mut all_char_counts: Vec<i32> = Vec::with_capacity(texts.len());
for text in texts {
let encoding = match tokenizer.encode(text) {
Ok(enc) => enc,
Err(e) => {
error!("Tokenization failed: {}", e);
return error_response(
StatusCode::INTERNAL_SERVER_ERROR,
&format!("Tokenization failed: {}", e),
"tokenization_error",
);
}
};
let token_ids: Vec<u32> = encoding.token_ids().to_vec();
let count = token_ids.len() as i32;
all_tokens.push(token_ids);
all_counts.push(count);
all_char_counts.push(text.chars().count() as i32);
}
// Format response based on single vs batch
let (tokens, count, char_count) = if is_batch {
(
TokensResult::Batch(all_tokens),
CountResult::Batch(all_counts),
CountResult::Batch(all_char_counts),
)
} else {
(
TokensResult::Single(all_tokens.into_iter().next().unwrap_or_default()),
CountResult::Single(all_counts.into_iter().next().unwrap_or(0)),
CountResult::Single(all_char_counts.into_iter().next().unwrap_or(0)),
)
};
Json(TokenizeResponse {
tokens,
count,
char_count,
})
.into_response()
}
/// Handle POST /v1/detokenize
pub async fn detokenize(registry: &Arc<TokenizerRegistry>, request: DetokenizeRequest) -> Response {
debug!("Detokenize request for model: {}", request.model);
let tokenizer = match get_tokenizer(registry, &request.model) {
Ok(t) => t,
Err(e) => {
return error_response(StatusCode::BAD_REQUEST, &e, "tokenizer_not_found");
}
};
let sequences = request.tokens.sequences();
let is_batch = request.tokens.is_batch();
// Detokenize each sequence
let mut all_texts: Vec<String> = Vec::with_capacity(sequences.len());
for seq in sequences {
let text = match tokenizer.decode(seq, request.skip_special_tokens) {
Ok(t) => t,
Err(e) => {
error!("Detokenization failed: {}", e);
return error_response(
StatusCode::INTERNAL_SERVER_ERROR,
&format!("Detokenization failed: {}", e),
"detokenization_error",
);
}
};
all_texts.push(text);
}
// Format response based on single vs batch
let text = if is_batch {
TextResult::Batch(all_texts)
} else {
TextResult::Single(all_texts.into_iter().next().unwrap_or_default())
};
Json(DetokenizeResponse { text }).into_response()
}
// ============================================================================
// Tokenizer Management Handlers
// ============================================================================
/// Handle POST /v1/tokenizers - async version using job queue
pub async fn add_tokenizer(context: &Arc<AppContext>, request: AddTokenizerRequest) -> Response {
// Check if tokenizer already exists by name
if context.tokenizer_registry.contains(&request.name) {
// Return the existing tokenizer's ID
if let Some(entry) = context.tokenizer_registry.get_by_name(&request.name) {
return (
StatusCode::CONFLICT,
Json(AddTokenizerResponse {
id: entry.id,
status: "failed".to_string(),
message: format!("Tokenizer '{}' already exists", request.name),
vocab_size: None,
}),
)
.into_response();
}
}
// Get the job queue
let job_queue = match context.worker_job_queue.get() {
Some(queue) => queue,
None => {
error!("Job queue not available");
return (
StatusCode::SERVICE_UNAVAILABLE,
Json(AddTokenizerResponse {
id: String::new(),
status: "failed".to_string(),
message: "Job queue not available".to_string(),
vocab_size: None,
}),
)
.into_response();
}
};
// Generate UUID for this tokenizer
let tokenizer_id = TokenizerRegistry::generate_id();
// Create the job with the pre-generated ID
let config = TokenizerConfigRequest {
id: tokenizer_id.clone(),
name: request.name.clone(),
source: request.source.clone(),
chat_template_path: request.chat_template_path.clone(),
};
let job = Job::AddTokenizer {
config: Box::new(config),
};
// Submit the job
match job_queue.submit(job).await {
Ok(()) => (
StatusCode::ACCEPTED,
Json(AddTokenizerResponse {
id: tokenizer_id,
status: "pending".to_string(),
message: format!(
"Tokenizer '{}' registration job submitted. Loading from: {}",
request.name, request.source
),
vocab_size: None,
}),
)
.into_response(),
Err(e) => {
error!("Failed to submit tokenizer job: {}", e);
(
StatusCode::SERVICE_UNAVAILABLE,
Json(AddTokenizerResponse {
id: String::new(),
status: "failed".to_string(),
message: e,
vocab_size: None,
}),
)
.into_response()
}
}
}
/// Handle GET /v1/tokenizers
pub async fn list_tokenizers(registry: &Arc<TokenizerRegistry>) -> Response {
debug!("List tokenizers request");
let entries = registry.list();
let tokenizers: Vec<TokenizerInfo> = entries
.into_iter()
.map(|e| TokenizerInfo {
id: e.id,
name: e.name,
source: e.source,
vocab_size: e.tokenizer.vocab_size(),
})
.collect();
Json(ListTokenizersResponse { tokenizers }).into_response()
}
/// Handle DELETE /v1/tokenizers/{tokenizer_id}
pub async fn remove_tokenizer(context: &Arc<AppContext>, tokenizer_id: &str) -> Response {
// Try to remove by ID first, then by name for backward compatibility
let removed = context
.tokenizer_registry
.remove_by_id(tokenizer_id)
.or_else(|| context.tokenizer_registry.remove(tokenizer_id));
if let Some(entry) = removed {
debug!("Removed tokenizer '{}' (id: {})", entry.name, entry.id);
(
StatusCode::OK,
Json(RemoveTokenizerResponse {
success: true,
message: format!("Tokenizer '{}' removed successfully", entry.name),
}),
)
.into_response()
} else {
warn!("Tokenizer '{}' not found", tokenizer_id);
(
StatusCode::NOT_FOUND,
Json(RemoveTokenizerResponse {
success: false,
message: format!("Tokenizer '{}' not found", tokenizer_id),
}),
)
.into_response()
}
}
/// Handle GET /v1/tokenizers/{tokenizer_id}
pub async fn get_tokenizer_info(context: &Arc<AppContext>, tokenizer_id: &str) -> Response {
debug!("Get tokenizer info for '{}'", tokenizer_id);
// Try by ID first, then by name
let entry: Option<TokenizerEntry> = context
.tokenizer_registry
.get_by_id(tokenizer_id)
.or_else(|| context.tokenizer_registry.get_by_name(tokenizer_id));
match entry {
Some(e) => {
let info = TokenizerInfo {
id: e.id,
name: e.name,
source: e.source,
vocab_size: e.tokenizer.vocab_size(),
};
Json(info).into_response()
}
None => error_response(
StatusCode::NOT_FOUND,
&format!("Tokenizer '{}' not found", tokenizer_id),
"tokenizer_not_found",
),
}
}
/// Handle GET /v1/tokenizers/{tokenizer_id}/status
pub async fn get_tokenizer_status(context: &Arc<AppContext>, tokenizer_id: &str) -> Response {
debug!("Get tokenizer status for '{}'", tokenizer_id);
// First check if tokenizer is already loaded (by ID or name)
let entry = context
.tokenizer_registry
.get_by_id(tokenizer_id)
.or_else(|| context.tokenizer_registry.get_by_name(tokenizer_id));
if let Some(e) = entry {
return Json(AddTokenizerResponse {
id: e.id,
status: "completed".to_string(),
message: format!("Tokenizer '{}' is loaded and ready", e.name),
vocab_size: Some(e.tokenizer.vocab_size()),
})
.into_response();
}
// Check job status (jobs are tracked by ID)
if let Some(job_queue) = context.worker_job_queue.get() {
if let Some(job_status) = job_queue.get_status(tokenizer_id) {
return Json(AddTokenizerResponse {
id: tokenizer_id.to_string(),
status: job_status.status.clone(),
message: job_status
.message
.unwrap_or_else(|| format!("Tokenizer job is {}", job_status.status)),
vocab_size: None,
})
.into_response();
}
}
// Not found
error_response(
StatusCode::NOT_FOUND,
&format!("Tokenizer '{}' not found and no pending job", tokenizer_id),
"not_found",
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tokenizer::mock::MockTokenizer;
fn create_test_registry() -> Arc<TokenizerRegistry> {
let registry = Arc::new(TokenizerRegistry::new());
let id = TokenizerRegistry::generate_id();
registry.register(
&id,
"test-model",
"test-source",
Arc::new(MockTokenizer::new()),
);
registry
}
#[test]
fn test_get_tokenizer_exact_match() {
let registry = create_test_registry();
let result = get_tokenizer(&registry, "test-model");
assert!(result.is_ok());
}
#[test]
fn test_get_tokenizer_default_fallback() {
let registry = create_test_registry();
let result = get_tokenizer(&registry, "default");
assert!(result.is_ok());
}
#[test]
fn test_get_tokenizer_not_found() {
let registry = create_test_registry();
let result = get_tokenizer(&registry, "nonexistent");
match result {
Err(e) => assert!(e.contains("not found")),
Ok(_) => panic!("Expected error"),
}
}
#[test]
fn test_get_tokenizer_empty_registry() {
let registry = Arc::new(TokenizerRegistry::new());
let result = get_tokenizer(&registry, "any");
match result {
Err(e) => assert!(e.contains("No tokenizers available")),
Ok(_) => panic!("Expected error"),
}
}
}
@@ -0,0 +1,13 @@
//! Tokenize module for tokenization and detokenization operations
//!
//! This module provides HTTP handlers for:
//! - Tokenizing text into token IDs
//! - Detokenizing token IDs back to text
//! - Managing tokenizers (add, list, get, remove)
mod handlers;
pub use handlers::{
add_tokenizer, detokenize, get_tokenizer_info, get_tokenizer_status, list_tokenizers,
remove_tokenizer, tokenize,
};
+74 -4
View File
@@ -25,9 +25,9 @@ use crate::{
core::{ core::{
steps::{ steps::{
create_external_worker_registration_workflow, create_mcp_registration_workflow, create_external_worker_registration_workflow, create_mcp_registration_workflow,
create_wasm_module_registration_workflow, create_wasm_module_removal_workflow, create_tokenizer_registration_workflow, create_wasm_module_registration_workflow,
create_worker_registration_workflow, create_worker_removal_workflow, create_wasm_module_removal_workflow, create_worker_registration_workflow,
create_worker_update_workflow, create_worker_removal_workflow, create_worker_update_workflow,
}, },
Job, JobQueue, JobQueueConfig, WorkerManager, WorkerType, Job, JobQueue, JobQueueConfig, WorkerManager, WorkerType,
}, },
@@ -46,10 +46,11 @@ use crate::{
parser::{ParseFunctionCallRequest, SeparateReasoningRequest}, parser::{ParseFunctionCallRequest, SeparateReasoningRequest},
rerank::{RerankRequest, V1RerankReqInput}, rerank::{RerankRequest, V1RerankReqInput},
responses::{ResponsesGetParams, ResponsesRequest}, responses::{ResponsesGetParams, ResponsesRequest},
tokenize::{AddTokenizerRequest, DetokenizeRequest, TokenizeRequest},
validated::ValidatedJson, validated::ValidatedJson,
worker_spec::{WorkerConfigRequest, WorkerUpdateRequest}, worker_spec::{WorkerConfigRequest, WorkerUpdateRequest},
}, },
routers::{conversations, router_manager::RouterManager, RouterTrait}, routers::{conversations, router_manager::RouterManager, tokenize, RouterTrait},
service_discovery::{start_service_discovery, ServiceDiscoveryConfig}, service_discovery::{start_service_discovery, ServiceDiscoveryConfig},
wasm::route::{add_wasm_module, list_wasm_modules, remove_wasm_module}, wasm::route::{add_wasm_module, list_wasm_modules, remove_wasm_module},
workflow::{LoggingSubscriber, WorkflowEngine}, workflow::{LoggingSubscriber, WorkflowEngine},
@@ -461,6 +462,56 @@ async fn update_worker(
} }
} }
// ============================================================================
// Tokenize / Detokenize Handlers
// ============================================================================
async fn v1_tokenize(
State(state): State<Arc<AppState>>,
Json(request): Json<TokenizeRequest>,
) -> Response {
tokenize::tokenize(&state.context.tokenizer_registry, request).await
}
async fn v1_detokenize(
State(state): State<Arc<AppState>>,
Json(request): Json<DetokenizeRequest>,
) -> Response {
tokenize::detokenize(&state.context.tokenizer_registry, request).await
}
async fn v1_tokenizers_add(
State(state): State<Arc<AppState>>,
Json(request): Json<AddTokenizerRequest>,
) -> Response {
tokenize::add_tokenizer(&state.context, request).await
}
async fn v1_tokenizers_list(State(state): State<Arc<AppState>>) -> Response {
tokenize::list_tokenizers(&state.context.tokenizer_registry).await
}
async fn v1_tokenizers_get(
State(state): State<Arc<AppState>>,
Path(tokenizer_id): Path<String>,
) -> Response {
tokenize::get_tokenizer_info(&state.context, &tokenizer_id).await
}
async fn v1_tokenizers_status(
State(state): State<Arc<AppState>>,
Path(tokenizer_id): Path<String>,
) -> Response {
tokenize::get_tokenizer_status(&state.context, &tokenizer_id).await
}
async fn v1_tokenizers_remove(
State(state): State<Arc<AppState>>,
Path(tokenizer_id): Path<String>,
) -> Response {
tokenize::remove_tokenizer(&state.context, &tokenizer_id).await
}
pub struct ServerConfig { pub struct ServerConfig {
pub host: String, pub host: String,
pub port: u16, pub port: u16,
@@ -516,6 +567,9 @@ pub fn build_app(
"/v1/conversations/{conversation_id}/items/{item_id}", "/v1/conversations/{conversation_id}/items/{item_id}",
get(v1_conversations_get_item).delete(v1_conversations_delete_item), get(v1_conversations_get_item).delete(v1_conversations_delete_item),
) )
// Tokenize / Detokenize endpoints
.route("/v1/tokenize", post(v1_tokenize))
.route("/v1/detokenize", post(v1_detokenize))
.route_layer(axum::middleware::from_fn_with_state( .route_layer(axum::middleware::from_fn_with_state(
app_state.clone(), app_state.clone(),
middleware::concurrency_limit_middleware, middleware::concurrency_limit_middleware,
@@ -547,6 +601,19 @@ pub fn build_app(
.route("/wasm", post(add_wasm_module)) .route("/wasm", post(add_wasm_module))
.route("/wasm/{module_uuid}", delete(remove_wasm_module)) .route("/wasm/{module_uuid}", delete(remove_wasm_module))
.route("/wasm", get(list_wasm_modules)) .route("/wasm", get(list_wasm_modules))
// Tokenizer management endpoints
.route(
"/v1/tokenizers",
post(v1_tokenizers_add).get(v1_tokenizers_list),
)
.route(
"/v1/tokenizers/{tokenizer_id}",
get(v1_tokenizers_get).delete(v1_tokenizers_remove),
)
.route(
"/v1/tokenizers/{tokenizer_id}/status",
get(v1_tokenizers_status),
)
.route_layer(axum::middleware::from_fn_with_state( .route_layer(axum::middleware::from_fn_with_state(
auth_config.clone(), auth_config.clone(),
middleware::auth_middleware, middleware::auth_middleware,
@@ -669,6 +736,9 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box<dyn std::error::Err
engine engine
.register_workflow(create_wasm_module_removal_workflow()) .register_workflow(create_wasm_module_removal_workflow())
.expect("wasm_module_removal workflow should be valid"); .expect("wasm_module_removal workflow should be valid");
engine
.register_workflow(create_tokenizer_registration_workflow())
.expect("tokenizer_registration workflow should be valid");
app_context app_context
.workflow_engine .workflow_engine
.set(engine) .set(engine)
+220 -167
View File
@@ -2,24 +2,53 @@
//! //!
//! Provides thread-safe, deduplicated tokenizer loading for IGW mode where //! Provides thread-safe, deduplicated tokenizer loading for IGW mode where
//! multiple routers (HTTP and gRPC) need to share tokenizers across workers. //! multiple routers (HTTP and gRPC) need to share tokenizers across workers.
//!
//! ## ID vs Name Lookup
//!
//! Tokenizers are stored with two keys:
//! - **ID (UUID)**: Unique identifier generated at registration, immutable
//! - **Name**: User-provided identifier, must be unique
//!
//! Lookup behavior:
//! - `get(key)`: Tries name first, then ID (backward compatible)
//! - `get_by_id(id)`: Exact ID match only
//! - `get_by_name(name)`: Exact name match only
//! - `remove(name)`: Removes by name
//! - `remove_by_id(id)`: Removes by ID
use std::sync::Arc; use std::sync::Arc;
use dashmap::DashMap; use dashmap::DashMap;
use tokio::sync::Mutex; use tokio::sync::Mutex;
use tracing::{debug, info}; use tracing::{debug, info};
use uuid::Uuid;
use super::traits::Tokenizer; use super::traits::Tokenizer;
/// Registry for managing tokenizers keyed by served_model_name /// Metadata and tokenizer instance for a registered tokenizer
#[derive(Clone)]
pub struct TokenizerEntry {
/// Unique identifier (UUID)
pub id: String,
/// User-provided name
pub name: String,
/// Source path or HuggingFace model ID
pub source: String,
/// The tokenizer instance
pub tokenizer: Arc<dyn Tokenizer>,
}
/// Registry for managing tokenizers keyed by UUID
/// ///
/// Features: /// Features:
/// - Thread-safe concurrent access using DashMap /// - Thread-safe concurrent access using DashMap
/// - Per-key locking to prevent duplicate loading /// - Per-key locking to prevent duplicate loading
/// - Simple key scheme: served_model_name /// - Lookup by UUID (primary) or name (secondary index)
pub struct TokenizerRegistry { pub struct TokenizerRegistry {
/// Storage for loaded tokenizers /// Storage for loaded tokenizers, keyed by UUID
tokenizers: DashMap<String, Arc<dyn Tokenizer>>, tokenizers: DashMap<String, TokenizerEntry>,
/// Secondary index: name -> UUID for lookup
name_to_id: DashMap<String, String>,
/// Per-key locks to prevent duplicate loading /// Per-key locks to prevent duplicate loading
loading_locks: DashMap<String, Arc<Mutex<()>>>, loading_locks: DashMap<String, Arc<Mutex<()>>>,
} }
@@ -29,130 +58,163 @@ impl TokenizerRegistry {
pub fn new() -> Self { pub fn new() -> Self {
Self { Self {
tokenizers: DashMap::new(), tokenizers: DashMap::new(),
name_to_id: DashMap::new(),
loading_locks: DashMap::new(), loading_locks: DashMap::new(),
} }
} }
/// Load and register a tokenizer by model ID /// Generate a new UUID for a tokenizer
pub fn generate_id() -> String {
Uuid::new_v4().to_string()
}
/// Load and register a tokenizer with a pre-generated ID
/// ///
/// If the tokenizer is already loaded, returns true immediately. /// If the tokenizer is already loaded (by name), returns the existing ID.
/// Otherwise, uses the provided loader function to load it. /// Otherwise, uses the provided loader function to load it.
/// Per-key locking ensures only one load happens per model, preventing race conditions. /// Per-key locking ensures only one load happens per name, preventing race conditions.
/// ///
/// # Arguments /// # Arguments
/// * `model_id` - The model identifier to use as key /// * `id` - Pre-generated UUID for this tokenizer
/// * `name` - User-provided name
/// * `source` - Source path or HuggingFace model ID
/// * `loader` - Async function that loads the tokenizer /// * `loader` - Async function that loads the tokenizer
/// ///
/// # Returns /// # Returns
/// * `Ok(true)` - Successfully loaded and registered (or already registered) /// * `Ok(id)` - Successfully loaded and registered (returns the ID)
/// * `Err(message)` - Error message if loading fails /// * `Err(message)` - Error message if loading fails
/// pub async fn load<F, Fut>(
/// # Example &self,
/// ```ignore id: &str,
/// registry.load("meta-llama/Llama-2-7b", || async { name: &str,
/// create_tokenizer_async("/path/to/tokenizer").await source: &str,
/// }).await?; loader: F,
/// ``` ) -> Result<String, String>
pub async fn load<F, Fut>(&self, model_id: &str, loader: F) -> Result<bool, String>
where where
F: FnOnce() -> Fut, F: FnOnce() -> Fut,
Fut: std::future::Future<Output = Result<Arc<dyn Tokenizer>, String>>, Fut: std::future::Future<Output = Result<Arc<dyn Tokenizer>, String>>,
{ {
// Fast path: already loaded // Fast path: already loaded by name
if self.tokenizers.contains_key(model_id) { if let Some(existing_id) = self.name_to_id.get(name) {
debug!("Tokenizer already registered for model: {}", model_id); debug!("Tokenizer already registered for name: {}", name);
return Ok(true); return Ok(existing_id.clone());
} }
debug!("Tokenizer cache miss for model: {}", model_id); debug!("Tokenizer cache miss for name: {}", name);
// Acquire per-key lock to prevent duplicate loading // Acquire per-name lock to prevent duplicate loading
let lock = self let lock = self
.loading_locks .loading_locks
.entry(model_id.to_string()) .entry(name.to_string())
.or_insert_with(|| Arc::new(Mutex::new(()))) .or_insert_with(|| Arc::new(Mutex::new(())))
.clone(); .clone();
let _guard = lock.lock().await; let _guard = lock.lock().await;
// Double-check after acquiring lock (another thread may have loaded it) // Double-check after acquiring lock (another thread may have loaded it)
if self.tokenizers.contains_key(model_id) { if let Some(existing_id) = self.name_to_id.get(name) {
debug!("Tokenizer loaded by another thread for model: {}", model_id); debug!("Tokenizer loaded by another thread for name: {}", name);
return Ok(true); return Ok(existing_id.clone());
} }
// Load tokenizer // Load tokenizer
info!("Loading tokenizer for model: {}", model_id); info!("Loading tokenizer '{}' from source: {}", name, source);
let tokenizer = loader().await?; let result = loader().await;
// Always clean up the lock, whether loading succeeded or failed
self.loading_locks.remove(name);
let tokenizer = result?;
// Create entry
let entry = TokenizerEntry {
id: id.to_string(),
name: name.to_string(),
source: source.to_string(),
tokenizer,
};
// Store in registry // Store in registry
self.tokenizers.insert(model_id.to_string(), tokenizer); self.tokenizers.insert(id.to_string(), entry);
self.name_to_id.insert(name.to_string(), id.to_string());
// Remove the lock since it's no longer needed for this model.
self.loading_locks.remove(model_id);
info!( info!(
"Successfully loaded and registered tokenizer for model: {}", "Successfully registered tokenizer '{}' with id: {}",
model_id name, id
); );
Ok(true) Ok(id.to_string())
} }
/// Register a pre-loaded tokenizer /// Register a pre-loaded tokenizer with a pre-generated ID
/// ///
/// Atomically inserts a tokenizer into the registry only if no tokenizer /// Atomically inserts a tokenizer into the registry only if no tokenizer
/// with the same model_name exists. Returns true if the tokenizer was inserted, /// with the same name exists. Returns the ID if successful.
/// false if one already existed.
///
/// This method is thread-safe and uses atomic operations to prevent race conditions.
/// If you need to replace an existing tokenizer, first use `remove()` then `register()`.
///
/// # Arguments
/// * `model_name` - The served_model_name to use as key
/// * `tokenizer` - The tokenizer to register
/// ///
/// # Returns /// # Returns
/// * `true` - If the tokenizer was successfully registered (didn't exist before) /// * `Some(id)` - If the tokenizer was successfully registered
/// * `false` - If a tokenizer with this model_name already existed /// * `None` - If a tokenizer with this name already existed
/// pub fn register(
/// # Example &self,
/// ```ignore id: &str,
/// let tokenizer = create_tokenizer_blocking("/path/to/tokenizer")?; name: &str,
/// if registry.register("meta-llama/Llama-2-7b", tokenizer) { source: &str,
/// info!("Tokenizer registered successfully"); tokenizer: Arc<dyn Tokenizer>,
/// } else { ) -> Option<String> {
/// info!("Tokenizer already exists");
/// }
/// ```
pub fn register(&self, model_name: &str, tokenizer: Arc<dyn Tokenizer>) -> bool {
use dashmap::mapref::entry::Entry; use dashmap::mapref::entry::Entry;
match self.tokenizers.entry(model_name.to_string()) {
// Check if name already exists
match self.name_to_id.entry(name.to_string()) {
Entry::Occupied(_) => { Entry::Occupied(_) => {
debug!( debug!(
"Tokenizer already exists for model: {}, skipping registration", "Tokenizer already exists for name: {}, skipping registration",
model_name name
); );
false None
} }
Entry::Vacant(entry) => { Entry::Vacant(name_entry) => {
info!("Registering tokenizer for model: {}", model_name); let entry = TokenizerEntry {
entry.insert(tokenizer); id: id.to_string(),
true name: name.to_string(),
source: source.to_string(),
tokenizer,
};
info!("Registering tokenizer '{}' with id: {}", name, id);
self.tokenizers.insert(id.to_string(), entry);
name_entry.insert(id.to_string());
Some(id.to_string())
} }
} }
} }
/// Get a tokenizer if it's already loaded /// Get a tokenizer by UUID
/// pub fn get_by_id(&self, id: &str) -> Option<TokenizerEntry> {
/// Returns None if the tokenizer hasn't been loaded yet. self.tokenizers.get(id).map(|e| e.clone())
pub fn get(&self, model_name: &str) -> Option<Arc<dyn Tokenizer>> {
self.tokenizers.get(model_name).map(|t| t.clone())
} }
/// Check if a tokenizer is loaded for the given model /// Get a tokenizer by name
pub fn contains(&self, model_name: &str) -> bool { pub fn get_by_name(&self, name: &str) -> Option<TokenizerEntry> {
self.tokenizers.contains_key(model_name) self.name_to_id
.get(name)
.and_then(|id| self.tokenizers.get(id.as_str()).map(|e| e.clone()))
}
/// Get a tokenizer (for backward compatibility, tries name first then ID)
pub fn get(&self, name_or_id: &str) -> Option<Arc<dyn Tokenizer>> {
self.get_by_name(name_or_id)
.or_else(|| self.get_by_id(name_or_id))
.map(|e| e.tokenizer)
}
/// Check if a tokenizer is registered by name
pub fn contains(&self, name: &str) -> bool {
self.name_to_id.contains_key(name)
}
/// Check if a tokenizer is registered by ID
pub fn contains_id(&self, id: &str) -> bool {
self.tokenizers.contains_key(id)
} }
/// Get the number of loaded tokenizers /// Get the number of loaded tokenizers
@@ -165,30 +227,41 @@ impl TokenizerRegistry {
self.tokenizers.is_empty() self.tokenizers.is_empty()
} }
/// List all registered tokenizer keys (model names) /// List all registered tokenizers
/// pub fn list(&self) -> Vec<TokenizerEntry> {
/// Returns a sorted vector of model names that have registered tokenizers. let mut entries: Vec<TokenizerEntry> =
/// Returns an empty vector if no tokenizers are registered. self.tokenizers.iter().map(|e| e.value().clone()).collect();
pub fn list(&self) -> Vec<String> { entries.sort_by(|a, b| a.name.cmp(&b.name));
let mut keys: Vec<String> = self entries
.tokenizers
.iter()
.map(|entry| entry.key().clone())
.collect();
keys.sort();
keys
} }
/// Remove a tokenizer from the registry /// Remove a tokenizer by ID
/// ///
/// Returns the tokenizer if it was present. /// Returns the entry if it was present.
pub fn remove(&self, model_name: &str) -> Option<Arc<dyn Tokenizer>> { pub fn remove_by_id(&self, id: &str) -> Option<TokenizerEntry> {
self.tokenizers.remove(model_name).map(|(_, v)| v) if let Some((_, entry)) = self.tokenizers.remove(id) {
self.name_to_id.remove(&entry.name);
Some(entry)
} else {
None
}
}
/// Remove a tokenizer by name
///
/// Returns the entry if it was present.
pub fn remove(&self, name: &str) -> Option<TokenizerEntry> {
if let Some((_, id)) = self.name_to_id.remove(name) {
self.tokenizers.remove(&id).map(|(_, e)| e)
} else {
None
}
} }
/// Clear all tokenizers from the registry /// Clear all tokenizers from the registry
pub fn clear(&self) { pub fn clear(&self) {
self.tokenizers.clear(); self.tokenizers.clear();
self.name_to_id.clear();
self.loading_locks.clear(); self.loading_locks.clear();
} }
} }
@@ -218,8 +291,9 @@ mod tests {
assert!(!registry.contains("model1")); assert!(!registry.contains("model1"));
// Load and register a tokenizer // Load and register a tokenizer
let id = TokenizerRegistry::generate_id();
registry registry
.load("model1", || async { .load(&id, "model1", "path/to/model", || async {
Ok(Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>) Ok(Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>)
}) })
.await .await
@@ -229,16 +303,16 @@ mod tests {
assert!(!registry.is_empty()); assert!(!registry.is_empty());
assert_eq!(registry.len(), 1); assert_eq!(registry.len(), 1);
assert!(registry.contains("model1")); assert!(registry.contains("model1"));
assert!(registry.contains_id(&id));
// Get returns the tokenizer // Get returns the tokenizer
let tokenizer = registry.get("model1").unwrap(); let entry = registry.get_by_name("model1").unwrap();
assert_eq!( assert_eq!(entry.id, id);
tokenizer.vocab_size(), assert_eq!(entry.name, "model1");
MockTokenizer::default().vocab_size() assert_eq!(entry.source, "path/to/model");
);
// Remove works // Remove works
let removed = registry.remove("model1"); let removed = registry.remove_by_id(&id);
assert!(removed.is_some()); assert!(removed.is_some());
assert!(registry.is_empty()); assert!(registry.is_empty());
} }
@@ -250,12 +324,13 @@ mod tests {
// Spawn multiple tasks trying to load the same tokenizer // Spawn multiple tasks trying to load the same tokenizer
let mut handles = vec![]; let mut handles = vec![];
for _ in 0..10 { for i in 0..10 {
let registry = registry.clone(); let registry = registry.clone();
let load_count = load_count.clone(); let load_count = load_count.clone();
let id = format!("id-{}", i);
let handle = tokio::spawn(async move { let handle = tokio::spawn(async move {
registry registry
.load("model1", || async { .load(&id, "model1", "source", || async {
// Simulate slow loading // Simulate slow loading
sleep(Duration::from_millis(10)).await; sleep(Duration::from_millis(10)).await;
load_count.fetch_add(1, std::sync::atomic::Ordering::SeqCst); load_count.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
@@ -287,8 +362,9 @@ mod tests {
// Load multiple tokenizers // Load multiple tokenizers
for i in 1..=5 { for i in 1..=5 {
let model_name = format!("model{}", i); let model_name = format!("model{}", i);
let id = TokenizerRegistry::generate_id();
registry registry
.load(&model_name, || async { .load(&id, &model_name, "source", || async {
Ok(Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>) Ok(Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>)
}) })
.await .await
@@ -300,6 +376,11 @@ mod tests {
assert!(registry.contains("model5")); assert!(registry.contains("model5"));
assert!(!registry.contains("model6")); assert!(!registry.contains("model6"));
// List returns all with metadata
let entries = registry.list();
assert_eq!(entries.len(), 5);
assert!(entries.iter().any(|e| e.name == "model1"));
// Clear all // Clear all
registry.clear(); registry.clear();
assert!(registry.is_empty()); assert!(registry.is_empty());
@@ -308,10 +389,13 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_load_failure() { async fn test_load_failure() {
let registry = TokenizerRegistry::new(); let registry = TokenizerRegistry::new();
let id = TokenizerRegistry::generate_id();
// Try to load with a failing loader // Try to load with a failing loader
let result = registry let result = registry
.load("failing_model", || async { Err("Load failed".to_string()) }) .load(&id, "failing_model", "source", || async {
Err("Load failed".to_string())
})
.await; .await;
assert!(result.is_err()); assert!(result.is_err());
@@ -320,90 +404,59 @@ mod tests {
} }
#[tokio::test] #[tokio::test]
async fn test_concurrent_different_models() { async fn test_get_by_name_and_id() {
let registry = Arc::new(TokenizerRegistry::new()); let registry = TokenizerRegistry::new();
let mut handles = vec![]; let id = TokenizerRegistry::generate_id();
// Load different models concurrently
for i in 1..=10 {
let registry = registry.clone();
let handle = tokio::spawn(async move {
let model_name = format!("model{}", i);
registry registry
.load(&model_name, || async { .load(&id, "my-model", "hf/model", || async {
sleep(Duration::from_millis(5)).await;
Ok(Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>) Ok(Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>)
}) })
.await .await
}); .unwrap();
handles.push(handle);
}
for handle in handles { // Get by name
handle.await.unwrap().unwrap(); let by_name = registry.get_by_name("my-model");
} assert!(by_name.is_some());
assert_eq!(by_name.as_ref().unwrap().id, id);
assert_eq!(registry.len(), 10); // Get by ID
let by_id = registry.get_by_id(&id);
assert!(by_id.is_some());
assert_eq!(by_id.as_ref().unwrap().name, "my-model");
// Generic get works with both
assert!(registry.get("my-model").is_some());
assert!(registry.get(&id).is_some());
} }
#[tokio::test] #[tokio::test]
async fn test_register_only_if_absent() { async fn test_register_only_if_absent() {
let registry = TokenizerRegistry::new(); let registry = TokenizerRegistry::new();
let id1 = TokenizerRegistry::generate_id();
let id2 = TokenizerRegistry::generate_id();
let tokenizer1 = Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>; let tokenizer1 = Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>;
let tokenizer2 = Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>; let tokenizer2 = Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>;
// First registration should succeed // First registration should succeed
assert!(registry.register("model1", tokenizer1.clone())); let result1 = registry.register(&id1, "model1", "source1", tokenizer1.clone());
assert_eq!(registry.len(), 1); assert!(result1.is_some());
assert!(registry.contains("model1"));
// Second registration with same key should fail
assert!(!registry.register("model1", tokenizer2.clone()));
assert_eq!(registry.len(), 1); assert_eq!(registry.len(), 1);
// Verify the original tokenizer is still there (not replaced) // Second registration with same name should fail
let retrieved = registry.get("model1").unwrap(); let result2 = registry.register(&id2, "model1", "source2", tokenizer2.clone());
assert_eq!( assert!(result2.is_none());
Arc::as_ptr(&retrieved), assert_eq!(registry.len(), 1);
Arc::as_ptr(&tokenizer1),
"Original tokenizer should not be replaced"
);
// Registration with different key should succeed // Original tokenizer should still be there
assert!(registry.register("model2", tokenizer2)); let entry = registry.get_by_name("model1").unwrap();
assert_eq!(entry.id, id1);
assert_eq!(entry.source, "source1");
// Registration with different name should succeed
let id3 = TokenizerRegistry::generate_id();
let result3 = registry.register(&id3, "model2", "source2", tokenizer2);
assert!(result3.is_some());
assert_eq!(registry.len(), 2); assert_eq!(registry.len(), 2);
} }
#[tokio::test]
async fn test_concurrent_register_same_model() {
let registry = Arc::new(TokenizerRegistry::new());
let success_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
// Spawn multiple tasks trying to register the same model
let mut handles = vec![];
for _ in 0..10 {
let registry = registry.clone();
let success_count = success_count.clone();
let handle = tokio::spawn(async move {
let tokenizer = Arc::new(MockTokenizer::default()) as Arc<dyn Tokenizer>;
if registry.register("model1", tokenizer) {
success_count.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}
});
handles.push(handle);
}
// Wait for all tasks
for handle in handles {
handle.await.unwrap();
}
// Verify only one registration succeeded
assert_eq!(
success_count.load(std::sync::atomic::Ordering::SeqCst),
1,
"Only one concurrent registration should succeed"
);
assert_eq!(registry.len(), 1);
}
} }