diff --git a/sgl-model-gateway/src/app_context.rs b/sgl-model-gateway/src/app_context.rs index ce42d7749..79086efd9 100644 --- a/sgl-model-gateway/src/app_context.rs +++ b/sgl-model-gateway/src/app_context.rs @@ -8,7 +8,9 @@ use tracing::{debug, info}; use crate::{ config::RouterConfig, - core::{ConnectionMode, JobQueue, LoadMonitor, WorkerRegistry, WorkerService}, + core::{ + ConnectionMode, JobQueue, LoadMonitor, WorkerRegistry, WorkerService, UNKNOWN_MODEL_ID, + }, data_connector::{ create_storage, ConversationItemStorage, ConversationStorage, ResponseStorage, }, @@ -451,25 +453,27 @@ impl AppContextBuilder { /// Create tokenizer registry and optionally load tokenizer /// 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 { // Create empty tokenizer registry let registry = Arc::new(TokenizerRegistry::new()); // Try to load router-level tokenizer if path is provided if let Some(tokenizer) = Self::maybe_tokenizer(config)? { - // Determine registration key: prefer tokenizer_path, then model_path, finally "unknown" - let tokenizer_key = config + // Determine registration key: prefer tokenizer_path, then model_path, finally UNKNOWN_MODEL_ID + let source = config .tokenizer_path .as_ref() .or(config.model_path.as_ref()) .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!( - "Tokenizer loaded and registered with key '{}' (vocab_size: {})", - tokenizer_key, + "Tokenizer loaded and registered with name '{}' id={} (vocab_size: {})", + source, + tokenizer_id, tokenizer.vocab_size() ); } diff --git a/sgl-model-gateway/src/core/job_queue.rs b/sgl-model-gateway/src/core/job_queue.rs index cd5324ee4..2a29f3b9f 100644 --- a/sgl-model-gateway/src/core/job_queue.rs +++ b/sgl-model-gateway/src/core/job_queue.rs @@ -17,8 +17,8 @@ use crate::{ app_context::AppContext, config::{RouterConfig, RoutingMode}, core::steps::{ - McpServerConfigRequest, WasmModuleConfigRequest, WasmModuleRemovalRequest, - WorkerRemovalRequest, + McpServerConfigRequest, TokenizerConfigRequest, TokenizerRemovalRequest, + WasmModuleConfigRequest, WasmModuleRemovalRequest, WorkerRemovalRequest, }, mcp::McpConfig, protocols::worker_spec::{JobStatus, WorkerConfigRequest, WorkerUpdateRequest}, @@ -53,6 +53,12 @@ pub enum Job { RemoveWasmModule { request: Box, }, + AddTokenizer { + config: Box, + }, + RemoveTokenizer { + request: Box, + }, } impl Job { @@ -67,10 +73,12 @@ impl Job { Job::RegisterMcpServer { .. } => "RegisterMcpServer", Job::AddWasmModule { .. } => "AddWasmModule", 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 { match self { Job::AddWorker { config } => &config.url, @@ -81,6 +89,8 @@ impl Job { Job::RegisterMcpServer { config } => &config.name, Job::AddWasmModule { config } => &config.descriptor.name, Job::RemoveWasmModule { request } => &request.uuid_string, + Job::AddTokenizer { config } => &config.id, + Job::RemoveTokenizer { request } => &request.id, } } } @@ -658,6 +668,52 @@ impl JobQueue { ) .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 = 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)) + } + } } } diff --git a/sgl-model-gateway/src/core/mod.rs b/sgl-model-gateway/src/core/mod.rs index 048262342..8755fd050 100644 --- a/sgl-model-gateway/src/core/mod.rs +++ b/sgl-model-gateway/src/core/mod.rs @@ -9,11 +9,8 @@ //! - Workflow steps for multi-step operations //! - Common utilities -/// 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"; +// Re-export UNKNOWN_MODEL_ID from protocols for use throughout core +pub use crate::protocols::UNKNOWN_MODEL_ID; pub mod circuit_breaker; pub mod error; diff --git a/sgl-model-gateway/src/core/steps/mod.rs b/sgl-model-gateway/src/core/steps/mod.rs index 9642a0586..112ab924b 100644 --- a/sgl-model-gateway/src/core/steps/mod.rs +++ b/sgl-model-gateway/src/core/steps/mod.rs @@ -4,9 +4,10 @@ //! - Worker management (registration, removal, updates) //! - MCP server registration //! - WASM module registration and removal -//! - Future: Tokenizer fetching, LoRA updates, etc. +//! - Tokenizer registration pub mod mcp_registration; +pub mod tokenizer_registration; pub mod wasm_module_registration; pub mod wasm_module_removal; pub mod worker; @@ -66,6 +67,10 @@ pub use mcp_registration::{ create_mcp_registration_workflow, ConnectMcpServerStep, DiscoverMcpInventoryStep, McpServerConfigRequest, RegisterMcpServerStep, ValidateRegistrationStep, }; +pub use tokenizer_registration::{ + create_tokenizer_registration_workflow, LoadTokenizerStep, TokenizerConfigRequest, + TokenizerRemovalRequest, ValidateTokenizerConfigStep, +}; pub use wasm_module_registration::{ create_wasm_module_registration_workflow, CalculateHashStep, CheckDuplicateStep, LoadWasmBytesStep, RegisterModuleStep, ValidateDescriptorStep, ValidateWasmComponentStep, diff --git a/sgl-model-gateway/src/core/steps/tokenizer_registration.rs b/sgl-model-gateway/src/core/steps/tokenizer_registration.rs new file mode 100644 index 000000000..31412ca4d --- /dev/null +++ b/sgl-model-gateway/src/core/steps/tokenizer_registration.rs @@ -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, +} + +/// 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 { + let config: Arc = context + .get("tokenizer_config") + .ok_or_else(|| WorkflowError::ContextValueNotFound("tokenizer_config".to_string()))?; + + let app_context: Arc = 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 { + let config: Arc = context + .get("tokenizer_config") + .ok_or_else(|| WorkflowError::ContextValueNotFound("tokenizer_config".to_string()))?; + + let app_context: Arc = 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"); + } +} diff --git a/sgl-model-gateway/src/core/steps/worker/local/create_worker.rs b/sgl-model-gateway/src/core/steps/worker/local/create_worker.rs index b4ec3095b..f587c4ac0 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/create_worker.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/create_worker.rs @@ -10,7 +10,7 @@ use crate::{ app_context::AppContext, core::{ BasicWorkerBuilder, CircuitBreakerConfig, ConnectionMode, DPAwareWorkerBuilder, - HealthConfig, ModelCard, RuntimeType, Worker, WorkerType, + HealthConfig, ModelCard, RuntimeType, Worker, WorkerType, UNKNOWN_MODEL_ID, }, protocols::worker_spec::WorkerConfigRequest, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, @@ -62,15 +62,15 @@ impl StepExecutor for CreateLocalWorkerStep { 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 .model_id .clone() .or_else(|| final_labels.get("served_model_name").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); } diff --git a/sgl-model-gateway/src/core/steps/worker/local/discover_dp.rs b/sgl-model-gateway/src/core/steps/worker/local/discover_dp.rs index 33e1391cd..a84f39453 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/discover_dp.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/discover_dp.rs @@ -7,6 +7,7 @@ use tracing::debug; use super::discover_metadata::get_server_info; use crate::{ + core::UNKNOWN_MODEL_ID, protocols::worker_spec::WorkerConfigRequest, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -34,7 +35,7 @@ pub async fn get_dp_info(url: &str, api_key: Option<&str>) -> Result String { - "unknown".to_string() + UNKNOWN_MODEL_ID.to_string() } /// Helper function for serde default value (returns true) diff --git a/sgl-model-gateway/src/protocols/mod.rs b/sgl-model-gateway/src/protocols/mod.rs index 19cf2b8fa..619661147 100644 --- a/sgl-model-gateway/src/protocols/mod.rs +++ b/sgl-model-gateway/src/protocols/mod.rs @@ -1,6 +1,12 @@ // Protocol definitions and validation for various LLM APIs // 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 chat; pub mod classify; @@ -14,5 +20,6 @@ pub mod parser; pub mod rerank; pub mod responses; pub mod sampling_params; +pub mod tokenize; pub mod validated; pub mod worker_spec; diff --git a/sgl-model-gateway/src/protocols/tokenize.rs b/sgl-model-gateway/src/protocols/tokenize.rs new file mode 100644 index 000000000..62a19c23b --- /dev/null +++ b/sgl-model-gateway/src/protocols/tokenize.rs @@ -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), + Batch(Vec>), +} + +/// Count result - either single or batch +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(untagged)] +pub enum CountResult { + Single(i32), + Batch(Vec), +} + +// ============================================================================ +// 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), + /// Batch of token sequences + Batch(Vec>), +} + +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), +} + +// ============================================================================ +// 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, +} + +/// 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, +} + +/// Response schema for listing tokenizers +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ListTokenizersResponse { + pub tokenizers: Vec, +} + +/// 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), +} + +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]")); + } +} diff --git a/sgl-model-gateway/src/protocols/worker_spec.rs b/sgl-model-gateway/src/protocols/worker_spec.rs index c8bf79ee9..89601c8af 100644 --- a/sgl-model-gateway/src/protocols/worker_spec.rs +++ b/sgl-model-gateway/src/protocols/worker_spec.rs @@ -6,6 +6,8 @@ use std::collections::HashMap; use serde::{Deserialize, Serialize}; +use super::UNKNOWN_MODEL_ID; + /// Worker configuration for API requests #[derive(Debug, Clone, Deserialize, Serialize)] pub struct WorkerConfigRequest { @@ -175,13 +177,13 @@ impl WorkerInfo { Self { id: worker_id.to_string(), url, - model_id: "unknown".to_string(), + model_id: UNKNOWN_MODEL_ID.to_string(), priority: 0, cost: 1.0, - worker_type: "unknown".to_string(), + worker_type: UNKNOWN_MODEL_ID.to_string(), is_healthy: false, load: 0, - connection_mode: "unknown".to_string(), + connection_mode: UNKNOWN_MODEL_ID.to_string(), runtime_type: None, tokenizer_path: None, reasoning_parser: None, diff --git a/sgl-model-gateway/src/routers/grpc/regular/responses/conversions.rs b/sgl-model-gateway/src/routers/grpc/regular/responses/conversions.rs index 2faf41712..eacc7ad66 100644 --- a/sgl-model-gateway/src/routers/grpc/regular/responses/conversions.rs +++ b/sgl-model-gateway/src/routers/grpc/regular/responses/conversions.rs @@ -19,6 +19,7 @@ use crate::{ ResponseReasoningContent::ReasoningText, ResponseStatus, ResponsesRequest, ResponsesResponse, ResponsesUsage, StringOrContentParts, TextConfig, TextFormat, }, + UNKNOWN_MODEL_ID, }, routers::grpc::common::responses::utils::extract_tools_from_response_tools, }; @@ -171,7 +172,7 @@ pub fn responses_to_chat(req: &ResponsesRequest) -> Result 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, 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, 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::with_capacity(texts.len()); + let mut all_counts: Vec = Vec::with_capacity(texts.len()); + let mut all_char_counts: Vec = 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 = 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, 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 = 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, 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) -> Response { + debug!("List tokenizers request"); + + let entries = registry.list(); + let tokenizers: Vec = 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, 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, tokenizer_id: &str) -> Response { + debug!("Get tokenizer info for '{}'", tokenizer_id); + + // Try by ID first, then by name + let entry: Option = 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, 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 { + 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(®istry, "test-model"); + assert!(result.is_ok()); + } + + #[test] + fn test_get_tokenizer_default_fallback() { + let registry = create_test_registry(); + let result = get_tokenizer(®istry, "default"); + assert!(result.is_ok()); + } + + #[test] + fn test_get_tokenizer_not_found() { + let registry = create_test_registry(); + let result = get_tokenizer(®istry, "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(®istry, "any"); + match result { + Err(e) => assert!(e.contains("No tokenizers available")), + Ok(_) => panic!("Expected error"), + } + } +} diff --git a/sgl-model-gateway/src/routers/tokenize/mod.rs b/sgl-model-gateway/src/routers/tokenize/mod.rs new file mode 100644 index 000000000..e1f6ce675 --- /dev/null +++ b/sgl-model-gateway/src/routers/tokenize/mod.rs @@ -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, +}; diff --git a/sgl-model-gateway/src/server.rs b/sgl-model-gateway/src/server.rs index 96a5b478f..4c53ffc35 100644 --- a/sgl-model-gateway/src/server.rs +++ b/sgl-model-gateway/src/server.rs @@ -25,9 +25,9 @@ use crate::{ core::{ steps::{ create_external_worker_registration_workflow, create_mcp_registration_workflow, - create_wasm_module_registration_workflow, create_wasm_module_removal_workflow, - create_worker_registration_workflow, create_worker_removal_workflow, - create_worker_update_workflow, + create_tokenizer_registration_workflow, create_wasm_module_registration_workflow, + create_wasm_module_removal_workflow, create_worker_registration_workflow, + create_worker_removal_workflow, create_worker_update_workflow, }, Job, JobQueue, JobQueueConfig, WorkerManager, WorkerType, }, @@ -46,10 +46,11 @@ use crate::{ parser::{ParseFunctionCallRequest, SeparateReasoningRequest}, rerank::{RerankRequest, V1RerankReqInput}, responses::{ResponsesGetParams, ResponsesRequest}, + tokenize::{AddTokenizerRequest, DetokenizeRequest, TokenizeRequest}, validated::ValidatedJson, worker_spec::{WorkerConfigRequest, WorkerUpdateRequest}, }, - routers::{conversations, router_manager::RouterManager, RouterTrait}, + routers::{conversations, router_manager::RouterManager, tokenize, RouterTrait}, service_discovery::{start_service_discovery, ServiceDiscoveryConfig}, wasm::route::{add_wasm_module, list_wasm_modules, remove_wasm_module}, workflow::{LoggingSubscriber, WorkflowEngine}, @@ -461,6 +462,56 @@ async fn update_worker( } } +// ============================================================================ +// Tokenize / Detokenize Handlers +// ============================================================================ + +async fn v1_tokenize( + State(state): State>, + Json(request): Json, +) -> Response { + tokenize::tokenize(&state.context.tokenizer_registry, request).await +} + +async fn v1_detokenize( + State(state): State>, + Json(request): Json, +) -> Response { + tokenize::detokenize(&state.context.tokenizer_registry, request).await +} + +async fn v1_tokenizers_add( + State(state): State>, + Json(request): Json, +) -> Response { + tokenize::add_tokenizer(&state.context, request).await +} + +async fn v1_tokenizers_list(State(state): State>) -> Response { + tokenize::list_tokenizers(&state.context.tokenizer_registry).await +} + +async fn v1_tokenizers_get( + State(state): State>, + Path(tokenizer_id): Path, +) -> Response { + tokenize::get_tokenizer_info(&state.context, &tokenizer_id).await +} + +async fn v1_tokenizers_status( + State(state): State>, + Path(tokenizer_id): Path, +) -> Response { + tokenize::get_tokenizer_status(&state.context, &tokenizer_id).await +} + +async fn v1_tokenizers_remove( + State(state): State>, + Path(tokenizer_id): Path, +) -> Response { + tokenize::remove_tokenizer(&state.context, &tokenizer_id).await +} + pub struct ServerConfig { pub host: String, pub port: u16, @@ -516,6 +567,9 @@ pub fn build_app( "/v1/conversations/{conversation_id}/items/{item_id}", 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( app_state.clone(), middleware::concurrency_limit_middleware, @@ -547,6 +601,19 @@ pub fn build_app( .route("/wasm", post(add_wasm_module)) .route("/wasm/{module_uuid}", delete(remove_wasm_module)) .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( auth_config.clone(), middleware::auth_middleware, @@ -669,6 +736,9 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box, +} + +/// Registry for managing tokenizers keyed by UUID /// /// Features: /// - Thread-safe concurrent access using DashMap /// - Per-key locking to prevent duplicate loading -/// - Simple key scheme: served_model_name +/// - Lookup by UUID (primary) or name (secondary index) pub struct TokenizerRegistry { - /// Storage for loaded tokenizers - tokenizers: DashMap>, + /// Storage for loaded tokenizers, keyed by UUID + tokenizers: DashMap, + /// Secondary index: name -> UUID for lookup + name_to_id: DashMap, /// Per-key locks to prevent duplicate loading loading_locks: DashMap>>, } @@ -29,130 +58,163 @@ impl TokenizerRegistry { pub fn new() -> Self { Self { tokenizers: DashMap::new(), + name_to_id: 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. - /// 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 - /// * `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 /// /// # 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 - /// - /// # Example - /// ```ignore - /// registry.load("meta-llama/Llama-2-7b", || async { - /// create_tokenizer_async("/path/to/tokenizer").await - /// }).await?; - /// ``` - pub async fn load(&self, model_id: &str, loader: F) -> Result + pub async fn load( + &self, + id: &str, + name: &str, + source: &str, + loader: F, + ) -> Result where F: FnOnce() -> Fut, Fut: std::future::Future, String>>, { - // Fast path: already loaded - if self.tokenizers.contains_key(model_id) { - debug!("Tokenizer already registered for model: {}", model_id); - return Ok(true); + // Fast path: already loaded by name + if let Some(existing_id) = self.name_to_id.get(name) { + debug!("Tokenizer already registered for name: {}", name); + 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 .loading_locks - .entry(model_id.to_string()) + .entry(name.to_string()) .or_insert_with(|| Arc::new(Mutex::new(()))) .clone(); let _guard = lock.lock().await; // Double-check after acquiring lock (another thread may have loaded it) - if self.tokenizers.contains_key(model_id) { - debug!("Tokenizer loaded by another thread for model: {}", model_id); - return Ok(true); + if let Some(existing_id) = self.name_to_id.get(name) { + debug!("Tokenizer loaded by another thread for name: {}", name); + return Ok(existing_id.clone()); } // Load tokenizer - info!("Loading tokenizer for model: {}", model_id); - let tokenizer = loader().await?; + info!("Loading tokenizer '{}' from source: {}", name, source); + 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 - self.tokenizers.insert(model_id.to_string(), tokenizer); - - // Remove the lock since it's no longer needed for this model. - self.loading_locks.remove(model_id); + self.tokenizers.insert(id.to_string(), entry); + self.name_to_id.insert(name.to_string(), id.to_string()); info!( - "Successfully loaded and registered tokenizer for model: {}", - model_id + "Successfully registered tokenizer '{}' with 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 - /// with the same model_name exists. Returns true if the tokenizer was inserted, - /// 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 + /// with the same name exists. Returns the ID if successful. /// /// # Returns - /// * `true` - If the tokenizer was successfully registered (didn't exist before) - /// * `false` - If a tokenizer with this model_name already existed - /// - /// # Example - /// ```ignore - /// let tokenizer = create_tokenizer_blocking("/path/to/tokenizer")?; - /// if registry.register("meta-llama/Llama-2-7b", tokenizer) { - /// info!("Tokenizer registered successfully"); - /// } else { - /// info!("Tokenizer already exists"); - /// } - /// ``` - pub fn register(&self, model_name: &str, tokenizer: Arc) -> bool { + /// * `Some(id)` - If the tokenizer was successfully registered + /// * `None` - If a tokenizer with this name already existed + pub fn register( + &self, + id: &str, + name: &str, + source: &str, + tokenizer: Arc, + ) -> Option { 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(_) => { debug!( - "Tokenizer already exists for model: {}, skipping registration", - model_name + "Tokenizer already exists for name: {}, skipping registration", + name ); - false + None } - Entry::Vacant(entry) => { - info!("Registering tokenizer for model: {}", model_name); - entry.insert(tokenizer); - true + Entry::Vacant(name_entry) => { + let entry = TokenizerEntry { + id: id.to_string(), + 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 - /// - /// Returns None if the tokenizer hasn't been loaded yet. - pub fn get(&self, model_name: &str) -> Option> { - self.tokenizers.get(model_name).map(|t| t.clone()) + /// Get a tokenizer by UUID + pub fn get_by_id(&self, id: &str) -> Option { + self.tokenizers.get(id).map(|e| e.clone()) } - /// Check if a tokenizer is loaded for the given model - pub fn contains(&self, model_name: &str) -> bool { - self.tokenizers.contains_key(model_name) + /// Get a tokenizer by name + pub fn get_by_name(&self, name: &str) -> Option { + 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> { + 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 @@ -165,30 +227,41 @@ impl TokenizerRegistry { self.tokenizers.is_empty() } - /// List all registered tokenizer keys (model names) - /// - /// Returns a sorted vector of model names that have registered tokenizers. - /// Returns an empty vector if no tokenizers are registered. - pub fn list(&self) -> Vec { - let mut keys: Vec = self - .tokenizers - .iter() - .map(|entry| entry.key().clone()) - .collect(); - keys.sort(); - keys + /// List all registered tokenizers + pub fn list(&self) -> Vec { + let mut entries: Vec = + self.tokenizers.iter().map(|e| e.value().clone()).collect(); + entries.sort_by(|a, b| a.name.cmp(&b.name)); + entries } - /// Remove a tokenizer from the registry + /// Remove a tokenizer by ID /// - /// Returns the tokenizer if it was present. - pub fn remove(&self, model_name: &str) -> Option> { - self.tokenizers.remove(model_name).map(|(_, v)| v) + /// Returns the entry if it was present. + pub fn remove_by_id(&self, id: &str) -> Option { + 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 { + 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 pub fn clear(&self) { self.tokenizers.clear(); + self.name_to_id.clear(); self.loading_locks.clear(); } } @@ -218,8 +291,9 @@ mod tests { assert!(!registry.contains("model1")); // Load and register a tokenizer + let id = TokenizerRegistry::generate_id(); registry - .load("model1", || async { + .load(&id, "model1", "path/to/model", || async { Ok(Arc::new(MockTokenizer::default()) as Arc) }) .await @@ -229,16 +303,16 @@ mod tests { assert!(!registry.is_empty()); assert_eq!(registry.len(), 1); assert!(registry.contains("model1")); + assert!(registry.contains_id(&id)); // Get returns the tokenizer - let tokenizer = registry.get("model1").unwrap(); - assert_eq!( - tokenizer.vocab_size(), - MockTokenizer::default().vocab_size() - ); + let entry = registry.get_by_name("model1").unwrap(); + assert_eq!(entry.id, id); + assert_eq!(entry.name, "model1"); + assert_eq!(entry.source, "path/to/model"); // Remove works - let removed = registry.remove("model1"); + let removed = registry.remove_by_id(&id); assert!(removed.is_some()); assert!(registry.is_empty()); } @@ -250,12 +324,13 @@ mod tests { // Spawn multiple tasks trying to load the same tokenizer let mut handles = vec![]; - for _ in 0..10 { + for i in 0..10 { let registry = registry.clone(); let load_count = load_count.clone(); + let id = format!("id-{}", i); let handle = tokio::spawn(async move { registry - .load("model1", || async { + .load(&id, "model1", "source", || async { // Simulate slow loading sleep(Duration::from_millis(10)).await; load_count.fetch_add(1, std::sync::atomic::Ordering::SeqCst); @@ -287,8 +362,9 @@ mod tests { // Load multiple tokenizers for i in 1..=5 { let model_name = format!("model{}", i); + let id = TokenizerRegistry::generate_id(); registry - .load(&model_name, || async { + .load(&id, &model_name, "source", || async { Ok(Arc::new(MockTokenizer::default()) as Arc) }) .await @@ -300,6 +376,11 @@ mod tests { assert!(registry.contains("model5")); 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 registry.clear(); assert!(registry.is_empty()); @@ -308,10 +389,13 @@ mod tests { #[tokio::test] async fn test_load_failure() { let registry = TokenizerRegistry::new(); + let id = TokenizerRegistry::generate_id(); // Try to load with a failing loader let result = registry - .load("failing_model", || async { Err("Load failed".to_string()) }) + .load(&id, "failing_model", "source", || async { + Err("Load failed".to_string()) + }) .await; assert!(result.is_err()); @@ -320,90 +404,59 @@ mod tests { } #[tokio::test] - async fn test_concurrent_different_models() { - let registry = Arc::new(TokenizerRegistry::new()); - let mut handles = vec![]; + async fn test_get_by_name_and_id() { + let registry = TokenizerRegistry::new(); + 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 - .load(&model_name, || async { - sleep(Duration::from_millis(5)).await; - Ok(Arc::new(MockTokenizer::default()) as Arc) - }) - .await - }); - handles.push(handle); - } + registry + .load(&id, "my-model", "hf/model", || async { + Ok(Arc::new(MockTokenizer::default()) as Arc) + }) + .await + .unwrap(); - for handle in handles { - handle.await.unwrap().unwrap(); - } + // Get by name + 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] async fn test_register_only_if_absent() { let registry = TokenizerRegistry::new(); + let id1 = TokenizerRegistry::generate_id(); + let id2 = TokenizerRegistry::generate_id(); let tokenizer1 = Arc::new(MockTokenizer::default()) as Arc; let tokenizer2 = Arc::new(MockTokenizer::default()) as Arc; // First registration should succeed - assert!(registry.register("model1", tokenizer1.clone())); - assert_eq!(registry.len(), 1); - assert!(registry.contains("model1")); - - // Second registration with same key should fail - assert!(!registry.register("model1", tokenizer2.clone())); + let result1 = registry.register(&id1, "model1", "source1", tokenizer1.clone()); + assert!(result1.is_some()); assert_eq!(registry.len(), 1); - // Verify the original tokenizer is still there (not replaced) - let retrieved = registry.get("model1").unwrap(); - assert_eq!( - Arc::as_ptr(&retrieved), - Arc::as_ptr(&tokenizer1), - "Original tokenizer should not be replaced" - ); + // Second registration with same name should fail + let result2 = registry.register(&id2, "model1", "source2", tokenizer2.clone()); + assert!(result2.is_none()); + assert_eq!(registry.len(), 1); - // Registration with different key should succeed - assert!(registry.register("model2", tokenizer2)); + // Original tokenizer should still be there + 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); } - - #[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; - 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); - } }