1use schemars::JsonSchema;
9use serde::{Deserialize, Serialize};
10use std::collections::BTreeMap;
11use std::path::PathBuf;
12
13#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
17pub enum ProtocolVersion {
18 #[serde(rename = "3")]
20 V3,
21}
22
23#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
26#[serde(tag = "operation", rename_all = "snake_case", deny_unknown_fields)]
27pub enum AdapterRequest {
28 PlanServe {
31 protocol_version: ProtocolVersion,
32 input: PlanServeInput,
33 },
34 RenderServe {
37 protocol_version: ProtocolVersion,
38 input: RenderServeInput,
39 },
40}
41
42impl AdapterRequest {
43 #[must_use]
45 pub const fn protocol_version(&self) -> ProtocolVersion {
46 match self {
47 Self::PlanServe {
48 protocol_version, ..
49 }
50 | Self::RenderServe {
51 protocol_version, ..
52 } => *protocol_version,
53 }
54 }
55}
56
57#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
59#[serde(deny_unknown_fields)]
60pub struct PlanServeInput {
61 pub model: ServeModelInput,
62 pub topology: ServeTopology,
63 pub routing_backend: String,
64 #[serde(default, skip_serializing_if = "Option::is_none")]
65 pub kv_transfer: Option<KvTransferMechanism>,
66 pub parallelism: Parallelism,
67 pub settings: BTreeMap<String, SettingValue>,
68 pub roles: Vec<ServeRoleInput>,
69 #[serde(default)]
70 pub profiling: bool,
71}
72
73#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
76#[serde(deny_unknown_fields)]
77pub struct RenderServeInput {
78 pub model: ServeModelInput,
79 pub topology: ServeTopology,
80 pub routing_backend: String,
81 #[serde(default, skip_serializing_if = "Option::is_none")]
82 pub kv_transfer: Option<KvTransferMechanism>,
83 pub parallelism: Parallelism,
84 pub settings: BTreeMap<String, SettingValue>,
85 pub roles: Vec<ServeRoleResult>,
86 pub links: Vec<ServeRoleLink>,
87 pub allocations: Vec<ServeProcessAllocation>,
88 #[serde(default)]
89 pub profiling: bool,
90}
91
92#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
94#[serde(rename_all = "snake_case")]
95pub enum ServeTopology {
96 Single,
98 PrefillDecode,
100}
101
102#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
104#[serde(rename_all = "snake_case")]
105pub enum ServeRoleKind {
106 Serve,
108 Prefill,
110 Decode,
112 Router,
114}
115
116#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
118#[serde(rename_all = "snake_case")]
119pub enum KvTransferMechanism {
120 Mooncake,
122 Nixl,
124}
125
126#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
129#[serde(deny_unknown_fields)]
130pub struct ServeRoleInput {
131 pub id: String,
132 pub kind: ServeRoleKind,
133 pub replica_count: u32,
134 pub parallelism: Parallelism,
135 pub settings: BTreeMap<String, SettingValue>,
136}
137
138#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
141#[serde(deny_unknown_fields)]
142pub struct ServeRoleResult {
143 pub id: String,
144 pub kind: ServeRoleKind,
145 pub replica_count: u32,
146 pub effective_settings: BTreeMap<String, SettingValue>,
147 pub effective_parallelism: Parallelism,
148}
149
150#[derive(Clone, Debug, Default, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
155#[serde(default, deny_unknown_fields)]
156pub struct Parallelism {
157 #[serde(skip_serializing_if = "Option::is_none")]
159 pub outer: Option<ParallelismOuter>,
160 #[serde(skip_serializing_if = "Option::is_none")]
162 pub attention: Option<ParallelismAttention>,
163 #[serde(skip_serializing_if = "Option::is_none")]
165 pub experts: Option<ParallelismExperts>,
166}
167
168impl Parallelism {
169 pub fn merge_from(&mut self, other: &Self) {
172 if let Some(other) = &other.outer {
173 self.outer.get_or_insert_default().merge_from(other);
174 }
175 if let Some(other) = &other.attention {
176 self.attention.get_or_insert_default().merge_from(other);
177 }
178 if let Some(other) = &other.experts {
179 self.experts.get_or_insert_default().merge_from(other);
180 }
181 }
182}
183
184#[derive(Clone, Debug, Default, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
186#[serde(default, deny_unknown_fields)]
187pub struct ParallelismOuter {
188 #[schemars(range(min = 1))]
189 pub tensor_parallel_size: Option<u32>,
190 #[schemars(range(min = 1))]
191 pub pipeline_parallel_size: Option<u32>,
192}
193
194impl ParallelismOuter {
195 fn merge_from(&mut self, other: &Self) {
196 if other.tensor_parallel_size.is_some() {
197 self.tensor_parallel_size = other.tensor_parallel_size;
198 }
199 if other.pipeline_parallel_size.is_some() {
200 self.pipeline_parallel_size = other.pipeline_parallel_size;
201 }
202 }
203}
204
205#[derive(Clone, Debug, Default, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
207#[serde(default, deny_unknown_fields)]
208pub struct ParallelismAttention {
209 #[schemars(range(min = 1))]
210 pub tensor_parallel_size: Option<u32>,
211 #[schemars(range(min = 1))]
212 pub data_parallel_size: Option<u32>,
213 #[schemars(range(min = 1))]
214 pub context_parallel_size: Option<u32>,
215}
216
217impl ParallelismAttention {
218 fn merge_from(&mut self, other: &Self) {
219 if other.tensor_parallel_size.is_some() {
220 self.tensor_parallel_size = other.tensor_parallel_size;
221 }
222 if other.data_parallel_size.is_some() {
223 self.data_parallel_size = other.data_parallel_size;
224 }
225 if other.context_parallel_size.is_some() {
226 self.context_parallel_size = other.context_parallel_size;
227 }
228 }
229}
230
231#[derive(Clone, Debug, Default, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
233#[serde(default, deny_unknown_fields)]
234pub struct ParallelismExperts {
235 #[schemars(range(min = 1))]
236 pub tensor_parallel_size: Option<u32>,
237 #[schemars(range(min = 1))]
238 pub data_parallel_size: Option<u32>,
239 #[schemars(range(min = 1))]
240 pub expert_parallel_size: Option<u32>,
241 #[schemars(range(min = 1))]
242 pub dense_tensor_parallel_size: Option<u32>,
243}
244
245impl ParallelismExperts {
246 fn merge_from(&mut self, other: &Self) {
247 if other.tensor_parallel_size.is_some() {
248 self.tensor_parallel_size = other.tensor_parallel_size;
249 }
250 if other.data_parallel_size.is_some() {
251 self.data_parallel_size = other.data_parallel_size;
252 }
253 if other.expert_parallel_size.is_some() {
254 self.expert_parallel_size = other.expert_parallel_size;
255 }
256 if other.dense_tensor_parallel_size.is_some() {
257 self.dense_tensor_parallel_size = other.dense_tensor_parallel_size;
258 }
259 }
260}
261
262#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
264#[serde(deny_unknown_fields)]
265pub struct ServeModelInput {
266 pub locator: String,
267 pub served_name: String,
268}
269
270#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
272#[serde(deny_unknown_fields)]
273pub struct EndpointAssignment {
274 pub host: String,
275 pub port: u16,
276}
277
278#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
280#[serde(deny_unknown_fields)]
281pub struct ClientEndpointInput {
282 pub protocol: EndpointProtocol,
283 pub host: String,
284 pub port: u16,
285 pub api_path: String,
286}
287
288#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
290#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
291pub enum EvalDefinitionInput {
292 #[serde(rename = "openai_smoke")]
294 OpenAiSmoke {
295 prompt: String,
296 max_tokens: u32,
297 timeout_seconds: u64,
298 },
299 LmEval {
301 task: String,
302 dataset: Option<String>,
303 split: Option<String>,
304 limit: Option<u32>,
305 few_shot: Option<u32>,
306 seed: Option<u64>,
307 max_tokens: Option<u32>,
308 concurrency: Option<u32>,
309 metric: String,
310 threshold: f64,
311 timeout_seconds: u64,
312 },
313}
314
315#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
317#[serde(deny_unknown_fields)]
318pub struct BenchDefinitionInput {
319 pub input_tokens: u32,
320 pub output_tokens: u32,
321 pub seed: u64,
322 pub temperature: f64,
323 pub timeout_seconds: u64,
324 #[serde(default)]
325 pub reset_prefix_cache: bool,
326}
327
328#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
331#[serde(untagged)]
332pub enum SettingValue {
333 Bool(bool),
335 Integer(i64),
337 Float(f64),
339 String(String),
341 Array(Vec<SettingValue>),
343 Object(BTreeMap<String, SettingValue>),
345}
346
347#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
349#[serde(tag = "status", rename_all = "snake_case", deny_unknown_fields)]
350pub enum AdapterResponse {
351 Ok {
353 protocol_version: ProtocolVersion,
354 result: Box<AdapterResult>,
355 },
356 Error {
358 protocol_version: ProtocolVersion,
359 error: AdapterError,
360 },
361}
362
363impl AdapterResponse {
364 #[must_use]
366 pub const fn protocol_version(&self) -> ProtocolVersion {
367 match self {
368 Self::Ok {
369 protocol_version, ..
370 }
371 | Self::Error {
372 protocol_version, ..
373 } => *protocol_version,
374 }
375 }
376}
377
378#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
381#[serde(tag = "operation", rename_all = "snake_case", deny_unknown_fields)]
382pub enum AdapterResult {
383 PlanServe { output: Box<PlanServeResult> },
385 RenderServe { output: Box<RenderServeResult> },
387}
388
389#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
393#[serde(deny_unknown_fields)]
394pub struct PlanServeResult {
395 pub integration: IntegrationIdentity,
396 pub effective_settings: BTreeMap<String, SettingValue>,
397 pub effective_parallelism: Parallelism,
398 pub roles: Vec<ServeRoleResult>,
399 pub replicas: Vec<ServeReplicaRequirement>,
400 pub links: Vec<ServeRoleLink>,
401 pub public_endpoint: PublicEndpointRequirement,
402 pub endpoint: EndpointRequirement,
403}
404
405#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
408#[serde(deny_unknown_fields)]
409pub struct ServeReplicaRequirement {
410 pub id: String,
411 pub role_id: String,
412 pub replica_index: u32,
413 pub accelerator_count: u32,
414 pub ports: Vec<String>,
415 pub primary_ports: Vec<String>,
416 pub primary_readiness: ReadinessProbe,
417 pub worker_readiness: ReadinessProbe,
418 #[serde(default, skip_serializing_if = "Option::is_none")]
419 pub capture_target: Option<CaptureTargetRequirement>,
420}
421
422#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
426#[serde(deny_unknown_fields)]
427pub struct ServeProcessAllocation {
428 pub process_id: String,
429 pub role_id: String,
430 pub replica_id: String,
431 pub replica_index: u32,
432 pub rank: u32,
433 pub machine_id: String,
434 pub model_locator: String,
435 pub runtime_cache_root: String,
436 pub devices: Vec<u32>,
437 pub endpoint: EndpointAssignment,
438 pub ports: BTreeMap<String, EndpointAssignment>,
439}
440
441#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
444#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
445pub enum ServeRoleLink {
446 RequestRouting {
448 source: String,
449 targets: Vec<String>,
450 },
451 KvTransfer {
453 source: String,
454 target: String,
455 mechanism: KvTransferMechanism,
456 },
457 Bootstrap {
459 source: String,
460 target: String,
461 port: String,
462 },
463 SideChannel {
465 source: String,
466 target: String,
467 port: String,
468 },
469}
470
471#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
473#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
474pub enum PublicEndpointRequirement {
475 Replica { replica_id: String },
477 BuiltinProxy {
480 process_id: String,
481 role_id: String,
482 prefill_role: String,
483 decode_role: String,
484 readiness: ReadinessProbe,
485 },
486}
487
488#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
491#[serde(deny_unknown_fields)]
492pub struct CaptureTargetRequirement {
493 pub control: CaptureControlRequirement,
494}
495
496#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
498#[serde(deny_unknown_fields)]
499pub struct CaptureControlRequirement {
500 pub start_path: String,
501 pub stop_path: String,
502}
503
504#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
507#[serde(deny_unknown_fields)]
508pub struct RenderServeResult {
509 pub integration: IntegrationIdentity,
510 pub processes: Vec<RenderedServeProcess>,
511}
512
513#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
515#[serde(deny_unknown_fields)]
516pub struct RenderedServeProcess {
517 pub id: String,
518 pub process: ProcessSpec,
519}
520
521#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
523#[serde(deny_unknown_fields)]
524pub struct HttpActionSpec {
525 pub method: HttpMethod,
526 pub path: String,
527}
528
529#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
531#[serde(rename_all = "snake_case")]
532pub enum HttpMethod {
533 Post,
535}
536
537#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
540#[serde(deny_unknown_fields)]
541pub struct IntegrationIdentity {
542 pub adapter_id: String,
543 pub adapter_version: String,
544 pub framework: String,
545}
546
547#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
549#[serde(deny_unknown_fields)]
550pub struct ProcessSpec {
551 pub argv: Vec<String>,
552 pub env: BTreeMap<String, String>,
553}
554
555#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
557#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
558pub enum ReadinessProbe {
559 Http { path: String },
561 ProcessAlive,
563}
564
565#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
568#[serde(deny_unknown_fields)]
569pub struct EndpointRequirement {
570 pub protocol: EndpointProtocol,
571 pub api_path: String,
572 #[serde(default, skip_serializing_if = "Option::is_none")]
573 pub prefix_cache_reset: Option<HttpActionSpec>,
574}
575
576#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
578#[serde(rename_all = "snake_case")]
579pub enum EndpointProtocol {
580 Http,
582}
583
584#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
586#[serde(deny_unknown_fields)]
587pub struct AdapterError {
588 pub code: AdapterErrorCode,
589 pub message: String,
590}
591
592#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
594#[serde(rename_all = "snake_case")]
595pub enum AdapterErrorCode {
596 InvalidRequest,
598 UnsupportedProtocolVersion,
600 InvalidSettings,
602 Internal,
604 UnsupportedOperation,
606}
607
608#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
611#[serde(deny_unknown_fields)]
612pub struct EvalClientRequest {
613 pub protocol_version: ProtocolVersion,
614 pub endpoint: ClientEndpointInput,
615 pub model: ServeModelInput,
616 pub definition: EvalDefinitionInput,
617 pub artifact_dir: PathBuf,
618}
619
620#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
624#[serde(deny_unknown_fields)]
625pub struct BenchClientRequest {
626 pub protocol_version: ProtocolVersion,
627 pub endpoint: ClientEndpointInput,
628 pub model: ServeModelInput,
629 pub definition: BenchDefinitionInput,
630 pub case: BenchCaseInput,
631 pub artifact_dir: PathBuf,
632}
633
634#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
636#[serde(deny_unknown_fields)]
637pub struct BenchCaseInput {
638 pub load_shape: BenchLoadInput,
639 pub request_count: u32,
640}
641
642#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
644#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
645pub enum BenchLoadInput {
646 ConcurrencyLimited { concurrency: u32 },
648 RequestRateLimited {
650 request_rate: f64,
651 burstiness: Option<f64>,
652 },
653 UnboundedRequestRate,
655}
656
657#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
659#[serde(rename_all = "snake_case")]
660pub enum ClientStatus {
661 Succeeded,
663 Failed,
665}
666
667#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
669#[serde(deny_unknown_fields)]
670pub struct EvalClientResult {
671 pub schema_version: u32,
674 pub status: ClientStatus,
675 pub metrics: BTreeMap<String, f64>,
676 pub native_command: Vec<String>,
677 pub raw_artifacts: Vec<RawArtifact>,
678 pub error: Option<String>,
679}
680
681#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
683#[serde(deny_unknown_fields)]
684pub struct BenchClientResult {
685 pub schema_version: u32,
688 pub status: ClientStatus,
689 pub completed_requests: u64,
690 pub failed_requests: u64,
691 pub normalization_schema: String,
692 pub metrics: BTreeMap<String, f64>,
693 pub native_command: Vec<String>,
694 pub native_exit_code: Option<i32>,
695 pub raw_artifacts: Vec<RawArtifact>,
696 pub error: Option<String>,
697}
698
699#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
701#[serde(deny_unknown_fields)]
702pub struct RawArtifact {
703 pub name: String,
704 pub kind: String,
705 pub path: PathBuf,
706}
707
708#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
712#[serde(deny_unknown_fields)]
713pub struct AdapterProtocol {
714 pub request: AdapterRequest,
715 pub response: AdapterResponse,
716 #[serde(default, skip_serializing_if = "Option::is_none")]
717 pub eval_client_request: Option<EvalClientRequest>,
718 #[serde(default, skip_serializing_if = "Option::is_none")]
719 pub eval_client_result: Option<EvalClientResult>,
720 #[serde(default, skip_serializing_if = "Option::is_none")]
721 pub bench_client_request: Option<BenchClientRequest>,
722 #[serde(default, skip_serializing_if = "Option::is_none")]
723 pub bench_client_result: Option<BenchClientResult>,
724}