Skip to content

Commit 704ec6f

Browse files
committed
fix(router): support model aliases end to end
A worker registers under a canonical model ID and may declare aliases next to it. Routing only ever matched the canonical ID, so a request that named an alias found no workers and got a 404. The registry now keeps an alias index next to the model index. It is separate so aliases stay out of `/v1/models` and out of the statistics, and `get_by_model` falls back to it. Every other per-model map — hash rings, retry overrides, load balancing policies — stays keyed by the canonical ID, so each request entry point resolves the alias exactly once and passes the canonical ID from there on: - the gRPC pipeline resolves in `RequestContext::new` and rewrites both `input.model_id` and the request's own `model` field, so worker selection, tokenizer lookup, parser selection and tool call ID format all read the canonical ID without resolving anything themselves; - both HTTP routers resolve at the top of their request path and rewrite the `model` field of the body they forward. The backend was registered under the canonical ID and has never heard of the alias. One visible consequence: the response reports the canonical model ID, not the alias the client sent, matching how the OpenAI API answers with the model it actually ran. `/v1/responses` gates on the client-supplied name before the pipeline canonicalizes it, so it gets `WorkerRegistry::contains_model`, which accepts both spellings. The `unknown` wildcard is not a registered name and stays rejected. Alias conflicts are resolved so that a name is either a canonical model ID or an alias, never both: - an alias naming a registered model is refused, and an alias already recorded is dropped when a model later claims that name — otherwise the name would start resolving to a different model once the real one's last worker left; - when two models declare the same alias the first registration wins, and losing that worker hands the alias to a remaining model that declares it rather than stranding it; - the alias entry is held across that whole decision, so a concurrent registration cannot be erased between the check and the delete. Per-model retry overrides are keyed by the canonical ID and are read before the pipeline canonicalizes, so `GrpcRouter::resolve_retry_config` resolves the alias itself. Without it a request naming an alias silently took the router default instead of the worker's override, on chat, generate, messages and completion alike. Tests cover the registry rules, both HTTP request paths end to end against the body the worker actually received, the multipart form, the rerank response, the gRPC pipeline, the Responses canonicalization, the retry override under an alias, and the Responses gate. Signed-off-by: Jun Liu <jun.c.liu@rakuten.com>
1 parent 4cd2647 commit 704ec6f

14 files changed

Lines changed: 1550 additions & 134 deletions

File tree

‎model_gateway/src/routers/grpc/common/responses/utils.rs‎

Lines changed: 51 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -98,14 +98,17 @@ pub(crate) async fn ensure_mcp_connection(
9898
Ok((false, Vec::new()))
9999
}
100100

101-
/// Validate that workers are available for the requested model
101+
/// Validate that workers are available for the requested model.
102+
///
103+
/// Runs on the client-supplied name, before the pipeline canonicalizes it, so
104+
/// it has to accept aliases as well as canonical model IDs. `contains_model`
105+
/// covers both; listing `get_models()` and testing membership would reject
106+
/// every alias here.
102107
pub(crate) fn validate_worker_availability(
103108
worker_registry: &Arc<WorkerRegistry>,
104109
model: &str,
105110
) -> Option<Response> {
106-
let available_models = worker_registry.get_models();
107-
108-
if !available_models.contains(&model.to_string()) {
111+
if !worker_registry.contains_model(model) {
109112
return Some(error::model_not_found(model));
110113
}
111114

@@ -175,3 +178,47 @@ pub(crate) async fn persist_response_if_needed(
175178
}
176179
}
177180
}
181+
182+
#[cfg(test)]
183+
mod tests {
184+
use openai_protocol::{model_card::ModelCard, worker::HealthCheckConfig};
185+
186+
use super::*;
187+
use crate::worker::{BasicWorkerBuilder, UNKNOWN_MODEL_ID};
188+
189+
fn registry_with_aliased_worker() -> Arc<WorkerRegistry> {
190+
let registry = Arc::new(WorkerRegistry::new());
191+
let worker = BasicWorkerBuilder::new("http://worker:8080")
192+
.model(ModelCard::new("canonical-model").with_alias("model-alias"))
193+
.health_config(HealthCheckConfig {
194+
disable_health_check: true,
195+
..Default::default()
196+
})
197+
.build();
198+
registry.register_or_replace(Arc::new(worker));
199+
registry
200+
}
201+
202+
#[test]
203+
fn worker_availability_accepts_alias_and_preserves_unknown_rejection() {
204+
let registry = registry_with_aliased_worker();
205+
206+
assert!(validate_worker_availability(&registry, "canonical-model").is_none());
207+
assert!(validate_worker_availability(&registry, "model-alias").is_none());
208+
209+
let response = validate_worker_availability(&registry, UNKNOWN_MODEL_ID)
210+
.expect("unknown model should remain rejected for Responses");
211+
assert_eq!(response.status(), http::StatusCode::NOT_FOUND);
212+
}
213+
214+
#[test]
215+
fn worker_availability_rejects_alias_once_its_worker_is_gone() {
216+
let registry = registry_with_aliased_worker();
217+
let worker_id = registry.get_id_by_url("http://worker:8080").unwrap();
218+
assert!(registry.remove(&worker_id).is_some());
219+
220+
let response = validate_worker_availability(&registry, "model-alias")
221+
.expect("alias must stop resolving with no workers behind it");
222+
assert_eq!(response.status(), http::StatusCode::NOT_FOUND);
223+
}
224+
}

‎model_gateway/src/routers/grpc/common/stages/dispatch_metadata.rs‎

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -27,14 +27,16 @@ impl PipelineStage for DispatchMetadataStage {
2727
})?;
2828

2929
let request_id = execution_plan.request_id().to_string();
30+
// The model the response reports. `RequestContext::new` already
31+
// rewrote every one of these to the canonical model ID, so a request
32+
// that arrived under an alias is answered under the canonical name.
3033
let model = match &ctx.input.request_type {
3134
RequestType::Chat(req) => req.model.clone(),
3235
RequestType::Completion(req) => req.model.clone(),
33-
RequestType::Generate(_req) => {
34-
// Generate requests don't have a model field
35-
// Use model_id from input
36-
ctx.input.model_id.clone()
37-
}
36+
// `GenerateRequest` carries a model field too, but callers of the
37+
// native `/generate` route may leave it empty, so prefer the
38+
// model the router resolved.
39+
RequestType::Generate(_req) => ctx.input.model_id.clone(),
3840
RequestType::Responses(req) => req.model.clone(),
3941
RequestType::Embedding(req) => req.model.clone(),
4042
RequestType::Classify(req) => req.model.clone(),

‎model_gateway/src/routers/grpc/context.rs‎

Lines changed: 87 additions & 60 deletions
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,7 @@ use super::{
3232
};
3333
use crate::{
3434
middleware::TenantRequestMeta,
35-
worker::{RuntimeType, Worker, WorkerLoadGuard},
35+
worker::{RuntimeType, Worker, WorkerLoadGuard, WorkerRegistry},
3636
};
3737

3838
/// Main request processing context
@@ -50,6 +50,7 @@ pub(crate) struct RequestContext {
5050
pub(crate) struct RequestInput {
5151
pub request_type: RequestType,
5252
pub headers: Option<HeaderMap>,
53+
/// Canonical model ID used after aliases are resolved at request entry.
5354
pub model_id: String,
5455
pub tenant_request_meta: Option<TenantRequestMeta>,
5556
}
@@ -67,6 +68,29 @@ pub(crate) enum RequestType {
6768
}
6869

6970
impl RequestType {
71+
/// Overwrite the request's own `model` field.
72+
///
73+
/// Callers hold the request behind an `Arc` that the retry loop also
74+
/// holds, so `Arc::make_mut` copies the request here. That cost is paid
75+
/// only on the alias path — [`RequestContext::new`] skips this call
76+
/// entirely when the client already used the canonical model ID.
77+
fn set_model(&mut self, model_id: &str) {
78+
fn replace(model: &mut String, model_id: &str) {
79+
model.clear();
80+
model.push_str(model_id);
81+
}
82+
83+
match self {
84+
Self::Chat(request) => replace(&mut Arc::make_mut(request).model, model_id),
85+
Self::Generate(request) => replace(&mut Arc::make_mut(request).model, model_id),
86+
Self::Completion(request) => replace(&mut Arc::make_mut(request).model, model_id),
87+
Self::Responses(request) => replace(&mut Arc::make_mut(request).model, model_id),
88+
Self::Embedding(request) => replace(&mut Arc::make_mut(request).model, model_id),
89+
Self::Classify(request) => replace(&mut Arc::make_mut(request).model, model_id),
90+
Self::Messages(request) => replace(&mut Arc::make_mut(request).model, model_id),
91+
}
92+
}
93+
7094
/// Client-supplied backend request id (`rid`), where the protocol carries
7195
/// one. Responses ids are storage-owned (`resp_*`) and never client-set.
7296
pub fn rid(&self) -> Option<&str> {
@@ -112,6 +136,7 @@ impl std::fmt::Display for FinalResponse {
112136
/// Shared components (injected once at creation)
113137
pub(crate) struct SharedComponents {
114138
pub tokenizer_registry: Arc<TokenizerRegistry>,
139+
pub worker_registry: Arc<WorkerRegistry>,
115140
pub tool_parser_factory: ToolParserFactory,
116141
pub reasoning_parser_factory: ReasoningParserFactory,
117142
/// Configured tool parser name (from CLI `--tool-call-parser`)
@@ -448,16 +473,32 @@ pub(crate) struct ResponseState {
448473
}
449474

450475
impl RequestContext {
451-
/// Create context for chat completion request
452-
pub fn for_chat(
453-
request: Arc<ChatCompletionRequest>,
476+
/// Build a context, resolving a model alias to its canonical model ID.
477+
///
478+
/// This is the single place the gRPC pipeline canonicalizes. Both
479+
/// `input.model_id` and the request's own `model` field are rewritten, so
480+
/// every stage below — worker selection, tokenizer lookup, parser
481+
/// selection, tool call ID format — reads the canonical ID without
482+
/// resolving anything itself.
483+
///
484+
/// One visible consequence: the response reports the canonical model ID,
485+
/// not the alias the client sent. That matches how the OpenAI API answers
486+
/// with the model it actually ran.
487+
fn new(
488+
mut request_type: RequestType,
454489
headers: Option<HeaderMap>,
455-
model_id: String,
490+
mut model_id: String,
456491
components: Arc<SharedComponents>,
457492
) -> Self {
493+
if let Some(canonical_model_id) = components.worker_registry.resolve_model_alias(&model_id)
494+
{
495+
model_id.clear();
496+
model_id.push_str(&canonical_model_id);
497+
request_type.set_model(&model_id);
498+
}
458499
Self {
459500
input: RequestInput {
460-
request_type: RequestType::Chat(request),
501+
request_type,
461502
headers,
462503
model_id,
463504
tenant_request_meta: None,
@@ -467,23 +508,29 @@ impl RequestContext {
467508
}
468509
}
469510

511+
/// Create context for chat completion request
512+
pub fn for_chat(
513+
request: Arc<ChatCompletionRequest>,
514+
headers: Option<HeaderMap>,
515+
model_id: String,
516+
components: Arc<SharedComponents>,
517+
) -> Self {
518+
Self::new(RequestType::Chat(request), headers, model_id, components)
519+
}
520+
470521
/// Create context for generate request
471522
pub fn for_generate(
472523
request: Arc<GenerateRequest>,
473524
headers: Option<HeaderMap>,
474525
model_id: String,
475526
components: Arc<SharedComponents>,
476527
) -> Self {
477-
Self {
478-
input: RequestInput {
479-
request_type: RequestType::Generate(request),
480-
headers,
481-
model_id,
482-
tenant_request_meta: None,
483-
},
528+
Self::new(
529+
RequestType::Generate(request),
530+
headers,
531+
model_id,
484532
components,
485-
state: ProcessingState::default(),
486-
}
533+
)
487534
}
488535

489536
/// Create context for completion request
@@ -493,16 +540,12 @@ impl RequestContext {
493540
model_id: String,
494541
components: Arc<SharedComponents>,
495542
) -> Self {
496-
Self {
497-
input: RequestInput {
498-
request_type: RequestType::Completion(request),
499-
headers,
500-
model_id,
501-
tenant_request_meta: None,
502-
},
543+
Self::new(
544+
RequestType::Completion(request),
545+
headers,
546+
model_id,
503547
components,
504-
state: ProcessingState::default(),
505-
}
548+
)
506549
}
507550

508551
/// Create context for Responses API request
@@ -512,16 +555,12 @@ impl RequestContext {
512555
model_id: String,
513556
components: Arc<SharedComponents>,
514557
) -> Self {
515-
Self {
516-
input: RequestInput {
517-
request_type: RequestType::Responses(request),
518-
headers,
519-
model_id,
520-
tenant_request_meta: None,
521-
},
558+
Self::new(
559+
RequestType::Responses(request),
560+
headers,
561+
model_id,
522562
components,
523-
state: ProcessingState::default(),
524-
}
563+
)
525564
}
526565

527566
/// Create context for embedding request
@@ -531,16 +570,12 @@ impl RequestContext {
531570
model_id: String,
532571
components: Arc<SharedComponents>,
533572
) -> Self {
534-
Self {
535-
input: RequestInput {
536-
request_type: RequestType::Embedding(request),
537-
headers,
538-
model_id,
539-
tenant_request_meta: None,
540-
},
573+
Self::new(
574+
RequestType::Embedding(request),
575+
headers,
576+
model_id,
541577
components,
542-
state: ProcessingState::default(),
543-
}
578+
)
544579
}
545580

546581
/// Create context for classify request
@@ -550,16 +585,12 @@ impl RequestContext {
550585
model_id: String,
551586
components: Arc<SharedComponents>,
552587
) -> Self {
553-
Self {
554-
input: RequestInput {
555-
request_type: RequestType::Classify(request),
556-
headers,
557-
model_id,
558-
tenant_request_meta: None,
559-
},
588+
Self::new(
589+
RequestType::Classify(request),
590+
headers,
591+
model_id,
560592
components,
561-
state: ProcessingState::default(),
562-
}
593+
)
563594
}
564595

565596
/// Create context for messages request
@@ -569,16 +600,12 @@ impl RequestContext {
569600
model_id: String,
570601
components: Arc<SharedComponents>,
571602
) -> Self {
572-
Self {
573-
input: RequestInput {
574-
request_type: RequestType::Messages(request),
575-
headers,
576-
model_id,
577-
tenant_request_meta: None,
578-
},
603+
Self::new(
604+
RequestType::Messages(request),
605+
headers,
606+
model_id,
579607
components,
580-
state: ProcessingState::default(),
581-
}
608+
)
582609
}
583610

584611
/// Get chat request (panics if not chat)

0 commit comments

Comments
 (0)