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
194 changes: 194 additions & 0 deletions lib/llm/src/discovery/watcher.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1573,6 +1573,200 @@ mod tests {
.expect("model did not reach the expected serving state")
}

#[tokio::test]
async fn classify_aliases_share_endpoint_with_isolated_routing_and_removal() {
use dynamo_runtime::{
discovery::{DiscoveryQuery, DiscoverySpec, EventTransportKind},
distributed::{DiscoveryBackend, RequestPlaneMode},
engine::AsyncEngineContextProvider,
pipeline::{ResponseStream, network::Ingress},
storage::kv,
};

// The TCP server is process-global; isolate it from other test runtimes.
Comment thread
jthomson04 marked this conversation as resolved.
const TEST: &str = concat!(
module_path!(),
"::classify_aliases_share_endpoint_with_isolated_routing_and_removal"
);
let test_name = TEST.split_once("::").unwrap().1;
if std::env::var("DYNAMO_ALIAS_TEST").as_deref() != Ok(test_name) {
let output = tokio::time::timeout(
Duration::from_secs(30),
tokio::process::Command::new(std::env::current_exe().unwrap())
.args(["--exact", test_name, "--nocapture"])
.env("DYNAMO_ALIAS_TEST", test_name)
.env("DYN_TCP_RPC_HOST", "127.0.0.1")
.env("DYN_TCP_RPC_PORT", "0")
.env("DYN_TCP_RESPONSE_STREAM_HOST", "127.0.0.1")
.env("DYN_TCP_RESPONSE_STREAM_PORT", "0")
.kill_on_drop(true)
.output(),
)
.await
.expect("classify subprocess must finish within its deadline")
.expect("classify subprocess must start");
assert!(
output.status.success(),
Comment thread
jthomson04 marked this conversation as resolved.
"{}\n{}",
String::from_utf8_lossy(&output.stdout),
String::from_utf8_lossy(&output.stderr)
);
let stdout = String::from_utf8_lossy(&output.stdout);
assert!(
stdout
.lines()
.any(|line| line.starts_with("test result: ok. 1 passed; 0 failed;")),
"classify subprocess must run exactly one passing test: {stdout}"
);
return;
}

#[derive(Debug)]
struct ClassifyWorker(&'static str);

#[async_trait]
impl
AsyncEngine<
SingleIn<NvCreateClassifyRequest>,
ManyOut<Annotated<NvCreateClassifyResponse>>,
Error,
> for ClassifyWorker
{
async fn generate(
&self,
request: SingleIn<NvCreateClassifyRequest>,
) -> Result<ManyOut<Annotated<NvCreateClassifyResponse>>, Error> {
let mut response = NvCreateClassifyResponse::empty();
response.model = self.0.to_string();
Ok(ResponseStream::new(
Box::pin(futures::stream::iter([Annotated::from_data(response)])),
request.context(),
))
}
}

async fn check_response(manager: &ModelManager, alias: &str) {
let request = serde_json::from_value(serde_json::json!({
"model": alias, "input": "test"
}))
.unwrap();
tokio::time::timeout(Duration::from_secs(5), async {
let mut response = manager
.get_classify_engine(alias)
.unwrap()
.generate(SingleIn::new(request))
.await
.unwrap();
assert_eq!(response.next().await.unwrap().data.unwrap().model, alias);
while response.next().await.is_some() {}
})
.await
.expect("classify request must complete");
}

let runtime = Runtime::from_current().unwrap();
let store = tempfile::tempdir().unwrap();
let config = || DistributedConfig {
discovery_backend: DiscoveryBackend::KvStore(kv::Selector::File(store.path().into())),
nats_config: None,
request_plane: RequestPlaneMode::Tcp,
response_plane: None,
event_transport_kind: EventTransportKind::Zmq,
};
let frontend = DistributedRuntime::new(runtime.clone(), config())
.await
.unwrap();
let manager = Arc::new(ModelManager::new());
let cancellation = CancellationToken::new();
let stream = frontend
.discovery()
.list_and_watch(DiscoveryQuery::AllModels, Some(cancellation.clone()))
.await
.unwrap();
let watcher = Arc::new(ModelWatcher::new(
frontend,
manager.clone(),
RouterConfig {
router_mode: RouterMode::RoundRobin,
..Default::default()
},
0,
None,
None,
None,
Arc::new(Metrics::new()),
));
let task = tokio::spawn(watcher.watch(stream, NamespaceFilter::Global));
let mut workers = Vec::new();
for alias in ["alias-a", "alias-b"] {
let drt = DistributedRuntime::new(runtime.clone(), config())
.await
.unwrap();
let endpoint = drt
.namespace("classify-aliases")
Comment thread
jthomson04 marked this conversation as resolved.
.unwrap()
.component("workers")
.unwrap()
.endpoint("generate");
let serving = endpoint
.endpoint_builder()
.handler(Ingress::for_engine(Arc::new(ClassifyWorker(alias))).unwrap())
.start_with_registration()
.await
.unwrap();
let mut card = ModelDeploymentCard::with_name_only(alias);
card.model_input = ModelInput::Text;
card.model_type = ModelType::Classify;
card.worker_type = Some(WorkerType::Aggregated);
card.source_path = Some("org/shared-classifier".to_string());
let registration = drt
.discovery()
.register(
DiscoverySpec::from_model(
"classify-aliases".into(),
"workers".into(),
"generate".into(),
&card,
)
.unwrap(),
)
.await
.unwrap();
workers.push((drt, serving, registration));
}
for alias in ["alias-a", "alias-b"] {
wait_for_model(&manager, alias, Model::is_ready_to_serve).await;
}
let mut names = manager.list_classify_models();
names.sort();
assert_eq!(names, ["alias-a", "alias-b"]);
// Repeated round-robin calls expose an unfiltered shared endpoint pool.
for _ in 0..4 {
check_response(&manager, "alias-a").await;
check_response(&manager, "alias-b").await;
}
let (drt, serving, registration) = workers.remove(0);
drt.discovery().unregister(registration).await.unwrap();
serving.shutdown().await.unwrap();
tokio::time::timeout(Duration::from_secs(5), async {
while manager.get_committed_model("alias-a").is_some() {
tokio::time::sleep(Duration::from_millis(5)).await;
}
})
.await
.expect("removed alias must leave the catalog");
assert_eq!(manager.list_classify_models(), ["alias-b"]);
assert!(manager.get_classify_engine("alias-a").is_err());
check_response(&manager, "alias-b").await;
for (drt, serving, registration) in workers {
drt.discovery().unregister(registration).await.unwrap();
serving.shutdown().await.unwrap();
}
cancellation.cancel();
task.await.unwrap();
runtime.shutdown();
}

#[tokio::test]
async fn incompatible_advertisements_preserve_the_incumbent_catalog_and_readiness() {
let runtime = Runtime::from_current().unwrap();
Expand Down
Loading
Loading