Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 3 additions & 5 deletions crates/protocols/src/chat.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,9 @@ use validator::Validate;

use super::{
common::{
default_model, default_true, validate_stop, ChatLogProbs, ContentPart, Function,
FunctionCall, FunctionChoice, GenerationRequest, ResponseFormat, StreamOptions,
StringOrArray, Tool, ToolCall, ToolCallDelta, ToolChoice, ToolChoiceValue, ToolReference,
Usage,
default_true, validate_stop, ChatLogProbs, ContentPart, Function, FunctionCall,
FunctionChoice, GenerationRequest, ResponseFormat, StreamOptions, StringOrArray, Tool,
ToolCall, ToolCallDelta, ToolChoice, ToolChoiceValue, ToolReference, Usage,
},
sampling_params::{validate_top_k_value, validate_top_p_value},
};
Expand Down Expand Up @@ -148,7 +147,6 @@ pub struct ChatCompletionRequest {
pub messages: Vec<ChatMessage>,

/// ID of the model to use
#[serde(default = "default_model")]
pub model: String,

/// Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing frequency in the text so far
Expand Down
9 changes: 4 additions & 5 deletions crates/protocols/src/common.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,15 +4,14 @@ use serde::{Deserialize, Serialize};
use serde_json::Value;
use validator;

use super::UNKNOWN_MODEL_ID;

// ============================================================================
// Default value helpers
// ============================================================================

/// Default model value when not specified
pub(crate) fn default_model() -> String {
UNKNOWN_MODEL_ID.to_string()
/// Default model for endpoints where model is optional (e.g., /generate).
/// Uses UNKNOWN_MODEL_ID so routers treat it as "any available worker."
pub fn default_unknown_model() -> String {
super::UNKNOWN_MODEL_ID.to_string()
}

/// Helper function for serde default value (returns true)
Expand Down
10 changes: 3 additions & 7 deletions crates/protocols/src/generate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,8 @@ pub struct GenerateRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub text: Option<String>,

pub model: Option<String>,
#[serde(default = "super::common::default_unknown_model")]
pub model: String,

/// Input IDs for tokenized input
#[serde(skip_serializing_if = "Option::is_none")]
Expand Down Expand Up @@ -203,12 +204,7 @@ impl GenerationRequest for GenerateRequest {
}

fn get_model(&self) -> Option<&str> {
// Generate requests have an optional model field
if let Some(s) = &self.model {
Some(s.as_str())
} else {
None
}
Some(self.model.as_str())
}

fn extract_text_for_routing(&self) -> String {
Expand Down
4 changes: 2 additions & 2 deletions crates/protocols/src/interactions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ use serde_with::skip_serializing_none;
use validator::{Validate, ValidationError};

use super::{
common::{default_model, default_true, Function, GenerationRequest},
common::{default_true, Function, GenerationRequest},
sampling_params::validate_top_p_value,
validated::Normalizable,
};
Expand Down Expand Up @@ -76,7 +76,7 @@ pub struct InteractionsRequest {
impl Default for InteractionsRequest {
fn default() -> Self {
Self {
model: Some(default_model()),
model: None,
agent: None,
agent_config: None,
input: InteractionsInput::Text(String::new()),
Expand Down
5 changes: 2 additions & 3 deletions crates/protocols/src/rerank.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ use serde::{Deserialize, Serialize};
use serde_json::Value;
use validator::Validate;

use super::common::{default_model, default_true, GenerationRequest, StringOrArray, UsageInfo};
use super::common::{default_true, GenerationRequest, StringOrArray, UsageInfo};

fn default_rerank_object() -> String {
"rerank".to_string()
Expand Down Expand Up @@ -34,7 +34,6 @@ pub struct RerankRequest {
pub documents: Vec<String>,

/// Model to use for reranking
#[serde(default = "default_model")]
pub model: String,

/// Maximum number of documents to return (optional)
Expand Down Expand Up @@ -205,7 +204,7 @@ impl From<V1RerankReqInput> for RerankRequest {
RerankRequest {
query: v1.query,
documents: v1.documents,
model: default_model(),
model: super::UNKNOWN_MODEL_ID.to_string(),
top_k: None,
return_documents: true,
rid: None,
Expand Down
5 changes: 2 additions & 3 deletions crates/protocols/src/responses.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ use validator::{Validate, ValidationError};

use super::{
common::{
default_model, default_true, validate_stop, ChatLogProbs, Function, GenerationRequest,
default_true, validate_stop, ChatLogProbs, Function, GenerationRequest,
PromptTokenUsageInfo, StringOrArray, ToolChoice, ToolChoiceValue, ToolReference, UsageInfo,
},
sampling_params::{validate_top_k_value, validate_top_p_value},
Expand Down Expand Up @@ -667,7 +667,6 @@ pub struct ResponsesRequest {
pub metadata: Option<HashMap<String, Value>>,

/// Model to use
#[serde(default = "default_model")]
pub model: String,

/// Optional conversation id to persist input/output as items
Expand Down Expand Up @@ -795,7 +794,7 @@ impl Default for ResponsesRequest {
max_output_tokens: None,
max_tool_calls: None,
metadata: None,
model: default_model(),
model: String::new(),
conversation: None,
parallel_tool_calls: None,
previous_response_id: None,
Expand Down
26 changes: 12 additions & 14 deletions crates/protocols/src/tokenize.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,6 @@

use serde::{Deserialize, Serialize};

use super::UNKNOWN_MODEL_ID;

// ============================================================================
// Tokenize API
// ============================================================================
Expand All @@ -17,7 +15,6 @@ use super::UNKNOWN_MODEL_ID;
#[derive(Debug, Clone, Deserialize, Serialize, schemars::JsonSchema)]
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
Expand Down Expand Up @@ -63,7 +60,6 @@ pub enum CountResult {
#[derive(Debug, Clone, Deserialize, Serialize, schemars::JsonSchema)]
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)
Expand Down Expand Up @@ -209,10 +205,6 @@ impl StringOrArray {
// Default Functions
// ============================================================================

fn default_model_name() -> String {
UNKNOWN_MODEL_ID.to_string()
}

fn default_true() -> bool {
true
}
Expand All @@ -222,11 +214,10 @@ mod tests {
use super::*;

#[test]
fn test_tokenize_request_single() {
fn test_tokenize_request_requires_model() {
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(_)));
let result = serde_json::from_str::<TokenizeRequest>(json);
assert!(result.is_err(), "Should fail without model field");
}

#[test]
Expand All @@ -238,16 +229,23 @@ mod tests {
}

#[test]
fn test_detokenize_request_single() {
fn test_detokenize_request_requires_model() {
let json = r#"{"tokens": [1, 2, 3]}"#;
let result = serde_json::from_str::<DetokenizeRequest>(json);
assert!(result.is_err(), "Should fail without model field");
}

#[test]
fn test_detokenize_request_single() {
let json = r#"{"model": "test-model", "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 json = r#"{"model": "test-model", "tokens": [[1, 2], [3, 4, 5]], "skip_special_tokens": false}"#;
Comment thread
coderabbitai[bot] marked this conversation as resolved.
let req: DetokenizeRequest = serde_json::from_str(json).unwrap();
assert!(matches!(req.tokens, TokensInput::Batch(_)));
assert!(!req.skip_special_tokens);
Expand Down
2 changes: 1 addition & 1 deletion model_gateway/benches/request_processing.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@ fn get_bootstrap_info(worker: &BasicWorker) -> (String, Option<u16>) {
fn default_generate_request() -> GenerateRequest {
GenerateRequest {
text: None,
model: None,
model: "unknown".to_string(),
input_ids: None,
input_embeds: None,
image_data: None,
Expand Down
2 changes: 1 addition & 1 deletion model_gateway/benches/wasm_middleware_latency.rs
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@ impl RouterTrait for MockRouter {
&self,
_headers: Option<&HeaderMap>,
_body: &ChatCompletionRequest,
_model_id: Option<&str>,
_model_id: &str,
) -> Response<Body> {
StatusCode::OK.into_response()
}
Expand Down
9 changes: 1 addition & 8 deletions model_gateway/src/routers/grpc/common/responses/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -91,14 +91,7 @@ pub(crate) fn validate_worker_availability(
let available_models = worker_registry.get_models();

if !available_models.contains(&model.to_string()) {
return Some(error::service_unavailable(
"no_available_workers",
format!(
"No workers available for model '{}'. Available models: {}",
model,
available_models.join(", ")
),
));
return Some(error::model_not_found(model));
Comment on lines 93 to +94

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Return 503 when responses workers are temporarily absent

validate_worker_availability now maps any missing model entry in the live registry to model_not_found (404), but this check runs before dispatch and cannot distinguish a true bad model from transient capacity loss (e.g., rolling restart, delayed registration, or full drain). In those outage windows /v1/responses will return a client error instead of a retriable 5xx, causing callers and gateways to stop retrying despite the model being valid.

Useful? React with 👍 / 👎.

}

None
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,9 @@ use axum::response::Response;
use tracing::error;

use super::PipelineStage;
use crate::{
core::UNKNOWN_MODEL_ID,
routers::{
error,
grpc::context::{DispatchMetadata, RequestContext, RequestType, WorkerSelection},
},
use crate::routers::{
error,
grpc::context::{DispatchMetadata, RequestContext, RequestType, WorkerSelection},
};

/// Dispatch metadata stage: Prepare metadata for dispatch
Expand All @@ -34,11 +31,8 @@ impl PipelineStage for DispatchMetadataStage {
RequestType::Chat(req) => req.model.clone(),
RequestType::Generate(_req) => {
// Generate requests don't have a model field
// Use model_id from input or UNKNOWN_MODEL_ID
ctx.input
.model_id
.clone()
.unwrap_or_else(|| UNKNOWN_MODEL_ID.to_string())
// Use model_id from input
ctx.input.model_id.clone()
}
RequestType::Responses(req) => req.model.clone(),
RequestType::Embedding(req) => req.model.clone(),
Expand Down
Loading
Loading