[model-gateway] Migrate Worker trait to model-aware methods (#14250)
This commit is contained in:
+212
-28
@@ -171,11 +171,17 @@ pub trait Worker: Send + Sync + fmt::Debug {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Get the model ID this worker serves
|
/// Get the model ID this worker serves
|
||||||
|
/// Checks ModelCards first, then falls back to labels
|
||||||
fn model_id(&self) -> &str {
|
fn model_id(&self) -> &str {
|
||||||
|
// Check ModelCards first
|
||||||
self.metadata()
|
self.metadata()
|
||||||
.labels
|
.models
|
||||||
.get("model_id")
|
.first()
|
||||||
.map(|s| s.as_str())
|
.map(|m| m.id.as_str())
|
||||||
|
.or_else(|| {
|
||||||
|
// Fall back to labels
|
||||||
|
self.metadata().labels.get("model_id").map(|s| s.as_str())
|
||||||
|
})
|
||||||
.unwrap_or("unknown")
|
.unwrap_or("unknown")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -197,36 +203,61 @@ pub trait Worker: Send + Sync + fmt::Debug {
|
|||||||
.unwrap_or(1.0)
|
.unwrap_or(1.0)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get the tokenizer path for this worker (gRPC mode only)
|
/// Get tokenizer path for a specific model.
|
||||||
fn tokenizer_path(&self) -> Option<&str> {
|
fn tokenizer_path(&self, model_id: &str) -> Option<&str> {
|
||||||
self.metadata()
|
self.metadata()
|
||||||
.labels
|
.find_model(model_id)
|
||||||
.get("tokenizer_path")
|
.and_then(|m| m.tokenizer_path.as_deref())
|
||||||
.map(|s| s.as_str())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get the reasoning parser type for this worker (gRPC mode only)
|
/// Get reasoning parser for a specific model.
|
||||||
fn reasoning_parser(&self) -> Option<&str> {
|
fn reasoning_parser(&self, model_id: &str) -> Option<&str> {
|
||||||
self.metadata()
|
self.metadata()
|
||||||
.labels
|
.find_model(model_id)
|
||||||
.get("reasoning_parser")
|
.and_then(|m| m.reasoning_parser.as_deref())
|
||||||
.map(|s| s.as_str())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get the tool parser type for this worker (gRPC mode only)
|
/// Get tool parser for a specific model.
|
||||||
fn tool_parser(&self) -> Option<&str> {
|
fn tool_parser(&self, model_id: &str) -> Option<&str> {
|
||||||
self.metadata()
|
self.metadata()
|
||||||
.labels
|
.find_model(model_id)
|
||||||
.get("tool_parser")
|
.and_then(|m| m.tool_parser.as_deref())
|
||||||
.map(|s| s.as_str())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get the chat template for this worker (gRPC mode only)
|
/// Get chat template for a specific model.
|
||||||
fn chat_template(&self) -> Option<&str> {
|
fn chat_template(&self, model_id: &str) -> Option<&str> {
|
||||||
self.metadata()
|
self.metadata()
|
||||||
.labels
|
.find_model(model_id)
|
||||||
.get("chat_template")
|
.and_then(|m| m.chat_template.as_deref())
|
||||||
.map(|s| s.as_str())
|
}
|
||||||
|
|
||||||
|
/// Get the default provider type for this worker.
|
||||||
|
/// `None` means native/passthrough.
|
||||||
|
fn default_provider(&self) -> Option<&ProviderType> {
|
||||||
|
self.metadata().default_provider.as_ref()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get provider for a specific model.
|
||||||
|
/// Priority: ModelCard.provider > worker.default_provider
|
||||||
|
fn provider_for_model(&self, model_id: &str) -> Option<&ProviderType> {
|
||||||
|
self.metadata().provider_for_model(model_id)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this worker supports a specific model.
|
||||||
|
/// If models list is empty, worker accepts any model.
|
||||||
|
fn supports_model(&self, model_id: &str) -> bool {
|
||||||
|
self.metadata().supports_model(model_id)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this worker supports an endpoint for a given model.
|
||||||
|
/// Falls back to default_model_type if model not found.
|
||||||
|
fn supports_endpoint(&self, model_id: &str, endpoint: Endpoint) -> bool {
|
||||||
|
self.metadata().supports_endpoint(model_id, endpoint)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get all models this worker can serve.
|
||||||
|
fn models(&self) -> &[ModelCard] {
|
||||||
|
&self.metadata().models
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get or create a gRPC client for this worker
|
/// Get or create a gRPC client for this worker
|
||||||
@@ -1015,10 +1046,11 @@ pub fn worker_to_info(worker: &Arc<dyn Worker>) -> WorkerInfo {
|
|||||||
ConnectionMode::Http => None,
|
ConnectionMode::Http => None,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
let model_id = worker.model_id();
|
||||||
WorkerInfo {
|
WorkerInfo {
|
||||||
id: worker.url().to_string(),
|
id: worker.url().to_string(),
|
||||||
url: worker.url().to_string(),
|
url: worker.url().to_string(),
|
||||||
model_id: worker.model_id().to_string(),
|
model_id: model_id.to_string(),
|
||||||
priority: worker.priority(),
|
priority: worker.priority(),
|
||||||
cost: worker.cost(),
|
cost: worker.cost(),
|
||||||
worker_type: worker_type_str.to_string(),
|
worker_type: worker_type_str.to_string(),
|
||||||
@@ -1026,10 +1058,10 @@ pub fn worker_to_info(worker: &Arc<dyn Worker>) -> WorkerInfo {
|
|||||||
load: worker.load(),
|
load: worker.load(),
|
||||||
connection_mode: format!("{:?}", worker.connection_mode()),
|
connection_mode: format!("{:?}", worker.connection_mode()),
|
||||||
runtime_type,
|
runtime_type,
|
||||||
tokenizer_path: worker.tokenizer_path().map(String::from),
|
tokenizer_path: worker.tokenizer_path(model_id).map(String::from),
|
||||||
reasoning_parser: worker.reasoning_parser().map(String::from),
|
reasoning_parser: worker.reasoning_parser(model_id).map(String::from),
|
||||||
tool_parser: worker.tool_parser().map(String::from),
|
tool_parser: worker.tool_parser(model_id).map(String::from),
|
||||||
chat_template: worker.chat_template().map(String::from),
|
chat_template: worker.chat_template(model_id).map(String::from),
|
||||||
bootstrap_port,
|
bootstrap_port,
|
||||||
metadata: worker.metadata().labels.clone(),
|
metadata: worker.metadata().labels.clone(),
|
||||||
job_status: None,
|
job_status: None,
|
||||||
@@ -2041,4 +2073,156 @@ mod tests {
|
|||||||
let ids: Vec<&str> = metadata.model_ids().collect();
|
let ids: Vec<&str> = metadata.model_ids().collect();
|
||||||
assert_eq!(ids, vec!["model-a", "model-b", "model-c"]);
|
assert_eq!(ids, vec!["model-a", "model-b", "model-c"]);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// === Phase 1.4: Worker trait model-aware methods tests ===
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_worker_tokenizer_path() {
|
||||||
|
use super::ModelCard;
|
||||||
|
use crate::core::BasicWorkerBuilder;
|
||||||
|
|
||||||
|
// Create a worker with a ModelCard that has tokenizer_path
|
||||||
|
let model_card =
|
||||||
|
ModelCard::new("my-model").with_tokenizer_path("my-model/tokenizer".to_string());
|
||||||
|
|
||||||
|
let worker = BasicWorkerBuilder::new("http://test:8080")
|
||||||
|
.model(model_card)
|
||||||
|
.build();
|
||||||
|
|
||||||
|
// Should find the tokenizer_path from the ModelCard
|
||||||
|
assert_eq!(
|
||||||
|
worker.tokenizer_path("my-model"),
|
||||||
|
Some("my-model/tokenizer")
|
||||||
|
);
|
||||||
|
|
||||||
|
// Unknown model should return None
|
||||||
|
assert_eq!(worker.tokenizer_path("unknown-model"), None);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_worker_model_aware_methods_with_model_cards() {
|
||||||
|
use super::{ModelCard, ProviderType};
|
||||||
|
use crate::core::BasicWorkerBuilder;
|
||||||
|
|
||||||
|
// Build worker (labels are not used for model config anymore)
|
||||||
|
let mut worker = BasicWorkerBuilder::new("http://test:8080").build();
|
||||||
|
|
||||||
|
// Add model cards to the worker's metadata
|
||||||
|
let model_with_config = ModelCard::new("gpt-4o")
|
||||||
|
.with_tokenizer_path("gpt4o/tokenizer")
|
||||||
|
.with_chat_template("gpt4o_template")
|
||||||
|
.with_reasoning_parser("gpt4o_reasoning")
|
||||||
|
.with_tool_parser("gpt4o_tools")
|
||||||
|
.with_provider(ProviderType::OpenAI);
|
||||||
|
|
||||||
|
let model_without_config = ModelCard::new("llama-3.1");
|
||||||
|
|
||||||
|
worker.metadata.models = vec![model_with_config, model_without_config];
|
||||||
|
|
||||||
|
// Model with explicit config should use ModelCard values
|
||||||
|
assert_eq!(worker.tokenizer_path("gpt-4o"), Some("gpt4o/tokenizer"));
|
||||||
|
assert_eq!(worker.chat_template("gpt-4o"), Some("gpt4o_template"));
|
||||||
|
assert_eq!(worker.reasoning_parser("gpt-4o"), Some("gpt4o_reasoning"));
|
||||||
|
assert_eq!(worker.tool_parser("gpt-4o"), Some("gpt4o_tools"));
|
||||||
|
assert_eq!(
|
||||||
|
worker.provider_for_model("gpt-4o"),
|
||||||
|
Some(&ProviderType::OpenAI)
|
||||||
|
);
|
||||||
|
|
||||||
|
// Model without explicit config should return None (no fallback to labels)
|
||||||
|
assert_eq!(worker.tokenizer_path("llama-3.1"), None);
|
||||||
|
assert_eq!(worker.chat_template("llama-3.1"), None);
|
||||||
|
assert_eq!(worker.reasoning_parser("llama-3.1"), None);
|
||||||
|
assert_eq!(worker.tool_parser("llama-3.1"), None);
|
||||||
|
|
||||||
|
// Unknown model should return None
|
||||||
|
assert_eq!(worker.tokenizer_path("unknown"), None);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_worker_supports_model_and_endpoint() {
|
||||||
|
use super::{Endpoint, ModelCard};
|
||||||
|
use crate::core::BasicWorkerBuilder;
|
||||||
|
|
||||||
|
let mut worker = BasicWorkerBuilder::new("http://test:8080").build();
|
||||||
|
|
||||||
|
// Empty models list - accepts any model
|
||||||
|
assert!(worker.supports_model("any-model"));
|
||||||
|
|
||||||
|
// Add specific models
|
||||||
|
let llm_model = ModelCard::new("gpt-4o").with_model_type(ModelType::LLM);
|
||||||
|
let embed_model = ModelCard::new("text-embedding").with_model_type(ModelType::EMBEDDINGS);
|
||||||
|
|
||||||
|
worker.metadata.models = vec![llm_model, embed_model];
|
||||||
|
|
||||||
|
// Now only listed models are supported
|
||||||
|
assert!(worker.supports_model("gpt-4o"));
|
||||||
|
assert!(worker.supports_model("text-embedding"));
|
||||||
|
assert!(!worker.supports_model("unknown-model"));
|
||||||
|
|
||||||
|
// Check endpoint support
|
||||||
|
assert!(worker.supports_endpoint("gpt-4o", Endpoint::Chat));
|
||||||
|
assert!(!worker.supports_endpoint("gpt-4o", Endpoint::Embeddings));
|
||||||
|
assert!(worker.supports_endpoint("text-embedding", Endpoint::Embeddings));
|
||||||
|
assert!(!worker.supports_endpoint("text-embedding", Endpoint::Chat));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_worker_models_accessor() {
|
||||||
|
use super::ModelCard;
|
||||||
|
use crate::core::BasicWorkerBuilder;
|
||||||
|
|
||||||
|
let mut worker = BasicWorkerBuilder::new("http://test:8080").build();
|
||||||
|
|
||||||
|
// Initially empty
|
||||||
|
assert!(worker.models().is_empty());
|
||||||
|
|
||||||
|
// Add models
|
||||||
|
worker.metadata.models = vec![ModelCard::new("model-a"), ModelCard::new("model-b")];
|
||||||
|
|
||||||
|
assert_eq!(worker.models().len(), 2);
|
||||||
|
assert_eq!(worker.models()[0].id, "model-a");
|
||||||
|
assert_eq!(worker.models()[1].id, "model-b");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_worker_default_provider() {
|
||||||
|
use super::ProviderType;
|
||||||
|
use crate::core::BasicWorkerBuilder;
|
||||||
|
|
||||||
|
let mut worker = BasicWorkerBuilder::new("http://test:8080").build();
|
||||||
|
|
||||||
|
// Default is None (native/passthrough)
|
||||||
|
assert!(worker.default_provider().is_none());
|
||||||
|
|
||||||
|
// Set a default provider
|
||||||
|
worker.metadata.default_provider = Some(ProviderType::OpenAI);
|
||||||
|
assert_eq!(worker.default_provider(), Some(&ProviderType::OpenAI));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_worker_model_id_with_model_cards() {
|
||||||
|
use super::ModelCard;
|
||||||
|
use crate::core::BasicWorkerBuilder;
|
||||||
|
|
||||||
|
// Test 1: No models, no labels → "unknown"
|
||||||
|
let worker = BasicWorkerBuilder::new("http://test:8080").build();
|
||||||
|
assert_eq!(worker.model_id(), "unknown");
|
||||||
|
|
||||||
|
// Test 2: No models but has label → uses label
|
||||||
|
let worker = BasicWorkerBuilder::new("http://test:8080")
|
||||||
|
.label("model_id", "label-model")
|
||||||
|
.build();
|
||||||
|
assert_eq!(worker.model_id(), "label-model");
|
||||||
|
|
||||||
|
// Test 3: Has ModelCards → uses first ModelCard
|
||||||
|
let mut worker = BasicWorkerBuilder::new("http://test:8080")
|
||||||
|
.label("model_id", "label-model")
|
||||||
|
.build();
|
||||||
|
worker.metadata.models = vec![
|
||||||
|
ModelCard::new("card-model-1"),
|
||||||
|
ModelCard::new("card-model-2"),
|
||||||
|
];
|
||||||
|
assert_eq!(worker.model_id(), "card-model-1");
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ use std::collections::HashMap;
|
|||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
circuit_breaker::{CircuitBreaker, CircuitBreakerConfig},
|
circuit_breaker::{CircuitBreaker, CircuitBreakerConfig},
|
||||||
|
model_card::ModelCard,
|
||||||
model_type::ModelType,
|
model_type::ModelType,
|
||||||
worker::{
|
worker::{
|
||||||
BasicWorker, ConnectionMode, DPAwareWorker, HealthConfig, RuntimeType, WorkerMetadata,
|
BasicWorker, ConnectionMode, DPAwareWorker, HealthConfig, RuntimeType, WorkerMetadata,
|
||||||
@@ -18,6 +19,7 @@ pub struct BasicWorkerBuilder {
|
|||||||
connection_mode: ConnectionMode,
|
connection_mode: ConnectionMode,
|
||||||
runtime_type: RuntimeType,
|
runtime_type: RuntimeType,
|
||||||
labels: HashMap<String, String>,
|
labels: HashMap<String, String>,
|
||||||
|
models: Vec<ModelCard>,
|
||||||
health_config: HealthConfig,
|
health_config: HealthConfig,
|
||||||
circuit_breaker_config: CircuitBreakerConfig,
|
circuit_breaker_config: CircuitBreakerConfig,
|
||||||
grpc_client: Option<GrpcClient>,
|
grpc_client: Option<GrpcClient>,
|
||||||
@@ -33,6 +35,7 @@ impl BasicWorkerBuilder {
|
|||||||
connection_mode: ConnectionMode::Http,
|
connection_mode: ConnectionMode::Http,
|
||||||
runtime_type: RuntimeType::default(),
|
runtime_type: RuntimeType::default(),
|
||||||
labels: HashMap::new(),
|
labels: HashMap::new(),
|
||||||
|
models: Vec::new(),
|
||||||
health_config: HealthConfig::default(),
|
health_config: HealthConfig::default(),
|
||||||
circuit_breaker_config: CircuitBreakerConfig::default(),
|
circuit_breaker_config: CircuitBreakerConfig::default(),
|
||||||
grpc_client: None,
|
grpc_client: None,
|
||||||
@@ -48,6 +51,7 @@ impl BasicWorkerBuilder {
|
|||||||
connection_mode: ConnectionMode::Http,
|
connection_mode: ConnectionMode::Http,
|
||||||
runtime_type: RuntimeType::default(),
|
runtime_type: RuntimeType::default(),
|
||||||
labels: HashMap::new(),
|
labels: HashMap::new(),
|
||||||
|
models: Vec::new(),
|
||||||
health_config: HealthConfig::default(),
|
health_config: HealthConfig::default(),
|
||||||
circuit_breaker_config: CircuitBreakerConfig::default(),
|
circuit_breaker_config: CircuitBreakerConfig::default(),
|
||||||
grpc_client: None,
|
grpc_client: None,
|
||||||
@@ -108,6 +112,18 @@ impl BasicWorkerBuilder {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Set models this worker can serve
|
||||||
|
pub fn models(mut self, models: Vec<ModelCard>) -> Self {
|
||||||
|
self.models = models;
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Add a single model this worker can serve
|
||||||
|
pub fn model(mut self, model: ModelCard) -> Self {
|
||||||
|
self.models.push(model);
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
/// Build the BasicWorker instance
|
/// Build the BasicWorker instance
|
||||||
pub fn build(self) -> BasicWorker {
|
pub fn build(self) -> BasicWorker {
|
||||||
use std::sync::{
|
use std::sync::{
|
||||||
@@ -155,7 +171,7 @@ impl BasicWorkerBuilder {
|
|||||||
health_config: self.health_config,
|
health_config: self.health_config,
|
||||||
bootstrap_host,
|
bootstrap_host,
|
||||||
bootstrap_port,
|
bootstrap_port,
|
||||||
models: Vec::new(), // Empty = accepts any model
|
models: self.models, // Empty = accepts any model
|
||||||
default_provider: None, // Native/passthrough
|
default_provider: None, // Native/passthrough
|
||||||
default_model_type: ModelType::LLM, // Standard LLM capabilities
|
default_model_type: ModelType::LLM, // Standard LLM capabilities
|
||||||
};
|
};
|
||||||
@@ -185,6 +201,7 @@ pub struct DPAwareWorkerBuilder {
|
|||||||
connection_mode: ConnectionMode,
|
connection_mode: ConnectionMode,
|
||||||
runtime_type: RuntimeType,
|
runtime_type: RuntimeType,
|
||||||
labels: HashMap<String, String>,
|
labels: HashMap<String, String>,
|
||||||
|
models: Vec<ModelCard>,
|
||||||
health_config: HealthConfig,
|
health_config: HealthConfig,
|
||||||
circuit_breaker_config: CircuitBreakerConfig,
|
circuit_breaker_config: CircuitBreakerConfig,
|
||||||
grpc_client: Option<GrpcClient>,
|
grpc_client: Option<GrpcClient>,
|
||||||
@@ -202,6 +219,7 @@ impl DPAwareWorkerBuilder {
|
|||||||
connection_mode: ConnectionMode::Http,
|
connection_mode: ConnectionMode::Http,
|
||||||
runtime_type: RuntimeType::default(),
|
runtime_type: RuntimeType::default(),
|
||||||
labels: HashMap::new(),
|
labels: HashMap::new(),
|
||||||
|
models: Vec::new(),
|
||||||
health_config: HealthConfig::default(),
|
health_config: HealthConfig::default(),
|
||||||
circuit_breaker_config: CircuitBreakerConfig::default(),
|
circuit_breaker_config: CircuitBreakerConfig::default(),
|
||||||
grpc_client: None,
|
grpc_client: None,
|
||||||
@@ -224,6 +242,7 @@ impl DPAwareWorkerBuilder {
|
|||||||
connection_mode: ConnectionMode::Http,
|
connection_mode: ConnectionMode::Http,
|
||||||
runtime_type: RuntimeType::default(),
|
runtime_type: RuntimeType::default(),
|
||||||
labels: HashMap::new(),
|
labels: HashMap::new(),
|
||||||
|
models: Vec::new(),
|
||||||
health_config: HealthConfig::default(),
|
health_config: HealthConfig::default(),
|
||||||
circuit_breaker_config: CircuitBreakerConfig::default(),
|
circuit_breaker_config: CircuitBreakerConfig::default(),
|
||||||
grpc_client: None,
|
grpc_client: None,
|
||||||
@@ -284,10 +303,23 @@ impl DPAwareWorkerBuilder {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Set models this worker can serve
|
||||||
|
pub fn models(mut self, models: Vec<ModelCard>) -> Self {
|
||||||
|
self.models = models;
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Add a single model this worker can serve
|
||||||
|
pub fn model(mut self, model: ModelCard) -> Self {
|
||||||
|
self.models.push(model);
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
/// Build the DPAwareWorker instance
|
/// Build the DPAwareWorker instance
|
||||||
pub fn build(self) -> DPAwareWorker {
|
pub fn build(self) -> DPAwareWorker {
|
||||||
let worker_url = format!("{}@{}", self.base_url, self.dp_rank);
|
let worker_url = format!("{}@{}", self.base_url, self.dp_rank);
|
||||||
let mut builder = BasicWorkerBuilder::new(worker_url)
|
let mut builder = BasicWorkerBuilder::new(worker_url)
|
||||||
|
.models(self.models)
|
||||||
.worker_type(self.worker_type)
|
.worker_type(self.worker_type)
|
||||||
.connection_mode(self.connection_mode)
|
.connection_mode(self.connection_mode)
|
||||||
.runtime_type(self.runtime_type)
|
.runtime_type(self.runtime_type)
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ use crate::{
|
|||||||
app_context::AppContext,
|
app_context::AppContext,
|
||||||
core::{
|
core::{
|
||||||
workflow::*, BasicWorkerBuilder, CircuitBreakerConfig, ConnectionMode,
|
workflow::*, BasicWorkerBuilder, CircuitBreakerConfig, ConnectionMode,
|
||||||
DPAwareWorkerBuilder, HealthConfig, RuntimeType, Worker, WorkerType,
|
DPAwareWorkerBuilder, HealthConfig, ModelCard, RuntimeType, Worker, WorkerType,
|
||||||
},
|
},
|
||||||
protocols::worker_spec::WorkerConfigRequest,
|
protocols::worker_spec::WorkerConfigRequest,
|
||||||
routers::grpc::client::GrpcClient,
|
routers::grpc::client::GrpcClient,
|
||||||
@@ -474,29 +474,14 @@ impl StepExecutor for CreateWorkerStep {
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build labels from config
|
// Build labels from config (non-model-specific labels only)
|
||||||
let mut config_labels = config.labels.clone();
|
let mut config_labels = config.labels.clone();
|
||||||
if let Some(model_id) = &config.model_id {
|
|
||||||
config_labels.insert("model_id".to_string(), model_id.clone());
|
|
||||||
}
|
|
||||||
if let Some(priority) = config.priority {
|
if let Some(priority) = config.priority {
|
||||||
config_labels.insert("priority".to_string(), priority.to_string());
|
config_labels.insert("priority".to_string(), priority.to_string());
|
||||||
}
|
}
|
||||||
if let Some(cost) = config.cost {
|
if let Some(cost) = config.cost {
|
||||||
config_labels.insert("cost".to_string(), cost.to_string());
|
config_labels.insert("cost".to_string(), cost.to_string());
|
||||||
}
|
}
|
||||||
if let Some(ref tokenizer_path) = config.tokenizer_path {
|
|
||||||
config_labels.insert("tokenizer_path".to_string(), tokenizer_path.clone());
|
|
||||||
}
|
|
||||||
if let Some(ref reasoning_parser) = config.reasoning_parser {
|
|
||||||
config_labels.insert("reasoning_parser".to_string(), reasoning_parser.clone());
|
|
||||||
}
|
|
||||||
if let Some(ref tool_parser) = config.tool_parser {
|
|
||||||
config_labels.insert("tool_parser".to_string(), tool_parser.clone());
|
|
||||||
}
|
|
||||||
if let Some(ref chat_template) = config.chat_template {
|
|
||||||
config_labels.insert("chat_template".to_string(), chat_template.clone());
|
|
||||||
}
|
|
||||||
|
|
||||||
// Merge: discovered labels first, then config labels (config takes precedence)
|
// Merge: discovered labels first, then config labels (config takes precedence)
|
||||||
let mut final_labels = discovered_labels.as_ref().clone();
|
let mut final_labels = discovered_labels.as_ref().clone();
|
||||||
@@ -504,19 +489,36 @@ impl StepExecutor for CreateWorkerStep {
|
|||||||
final_labels.insert(key.clone(), value.clone());
|
final_labels.insert(key.clone(), value.clone());
|
||||||
}
|
}
|
||||||
|
|
||||||
// Derive model_id if not already set
|
// Determine model_id: config > served_model_name > model_path > "unknown"
|
||||||
if !final_labels.contains_key("model_id") {
|
let model_id = config
|
||||||
let derived_model_id = final_labels
|
.model_id
|
||||||
.get("served_model_name")
|
.clone()
|
||||||
.or_else(|| final_labels.get("model_path"))
|
.or_else(|| final_labels.get("served_model_name").cloned())
|
||||||
.cloned();
|
.or_else(|| final_labels.get("model_path").cloned())
|
||||||
|
.unwrap_or_else(|| "unknown".to_string());
|
||||||
|
|
||||||
if let Some(model_id) = derived_model_id {
|
if model_id != "unknown" {
|
||||||
debug!("Derived model_id from metadata: {}", model_id);
|
debug!("Using model_id: {}", model_id);
|
||||||
final_labels.insert("model_id".to_string(), model_id);
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Create ModelCard with model-specific configuration
|
||||||
|
let model_card = {
|
||||||
|
let mut card = ModelCard::new(&model_id);
|
||||||
|
if let Some(ref tokenizer_path) = config.tokenizer_path {
|
||||||
|
card = card.with_tokenizer_path(tokenizer_path.clone());
|
||||||
|
}
|
||||||
|
if let Some(ref reasoning_parser) = config.reasoning_parser {
|
||||||
|
card = card.with_reasoning_parser(reasoning_parser.clone());
|
||||||
|
}
|
||||||
|
if let Some(ref tool_parser) = config.tool_parser {
|
||||||
|
card = card.with_tool_parser(tool_parser.clone());
|
||||||
|
}
|
||||||
|
if let Some(ref chat_template) = config.chat_template {
|
||||||
|
card = card.with_chat_template(chat_template.clone());
|
||||||
|
}
|
||||||
|
card
|
||||||
|
};
|
||||||
|
|
||||||
debug!(
|
debug!(
|
||||||
"Creating worker {} with {} discovered + {} config = {} final labels",
|
"Creating worker {} with {} discovered + {} config = {} final labels",
|
||||||
config.url,
|
config.url,
|
||||||
@@ -621,6 +623,7 @@ impl StepExecutor for CreateWorkerStep {
|
|||||||
for rank in 0..dp_info.dp_size {
|
for rank in 0..dp_info.dp_size {
|
||||||
let mut builder =
|
let mut builder =
|
||||||
DPAwareWorkerBuilder::new(normalized_url.clone(), rank, dp_info.dp_size)
|
DPAwareWorkerBuilder::new(normalized_url.clone(), rank, dp_info.dp_size)
|
||||||
|
.model(model_card.clone())
|
||||||
.worker_type(worker_type.clone())
|
.worker_type(worker_type.clone())
|
||||||
.connection_mode(connection_mode.as_ref().clone())
|
.connection_mode(connection_mode.as_ref().clone())
|
||||||
.runtime_type(runtime_type.clone())
|
.runtime_type(runtime_type.clone())
|
||||||
@@ -656,6 +659,7 @@ impl StepExecutor for CreateWorkerStep {
|
|||||||
} else {
|
} else {
|
||||||
// Non-DP-aware path: Create single worker
|
// Non-DP-aware path: Create single worker
|
||||||
let mut builder = BasicWorkerBuilder::new(normalized_url.clone())
|
let mut builder = BasicWorkerBuilder::new(normalized_url.clone())
|
||||||
|
.model(model_card)
|
||||||
.worker_type(worker_type)
|
.worker_type(worker_type)
|
||||||
.connection_mode(connection_mode.as_ref().clone())
|
.connection_mode(connection_mode.as_ref().clone())
|
||||||
.runtime_type(runtime_type)
|
.runtime_type(runtime_type)
|
||||||
|
|||||||
Reference in New Issue
Block a user