use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use std::collections::BTreeMap;
use std::path::PathBuf;
#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
pub enum ProtocolVersion {
#[serde(rename = "3")]
V3,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(tag = "operation", rename_all = "snake_case", deny_unknown_fields)]
pub enum AdapterRequest {
PlanServe {
protocol_version: ProtocolVersion,
input: PlanServeInput,
},
RenderServe {
protocol_version: ProtocolVersion,
input: RenderServeInput,
},
}
impl AdapterRequest {
#[must_use]
pub const fn protocol_version(&self) -> ProtocolVersion {
match self {
Self::PlanServe {
protocol_version, ..
}
| Self::RenderServe {
protocol_version, ..
} => *protocol_version,
}
}
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct PlanServeInput {
pub model: ServeModelInput,
pub topology: ServeTopology,
pub routing_backend: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub kv_transfer: Option<KvTransferMechanism>,
pub parallelism: Parallelism,
pub settings: BTreeMap<String, SettingValue>,
pub roles: Vec<ServeRoleInput>,
#[serde(default)]
pub profiling: bool,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct RenderServeInput {
pub model: ServeModelInput,
pub topology: ServeTopology,
pub routing_backend: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub kv_transfer: Option<KvTransferMechanism>,
pub parallelism: Parallelism,
pub settings: BTreeMap<String, SettingValue>,
pub roles: Vec<ServeRoleResult>,
pub links: Vec<ServeRoleLink>,
pub allocations: Vec<ServeProcessAllocation>,
#[serde(default)]
pub profiling: bool,
}
#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ServeTopology {
Single,
PrefillDecode,
}
#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ServeRoleKind {
Serve,
Prefill,
Decode,
Router,
}
#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum KvTransferMechanism {
Mooncake,
Nixl,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ServeRoleInput {
pub id: String,
pub kind: ServeRoleKind,
pub replica_count: u32,
pub parallelism: Parallelism,
pub settings: BTreeMap<String, SettingValue>,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ServeRoleResult {
pub id: String,
pub kind: ServeRoleKind,
pub replica_count: u32,
pub effective_settings: BTreeMap<String, SettingValue>,
pub effective_parallelism: Parallelism,
}
#[derive(Clone, Debug, Default, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(default, deny_unknown_fields)]
pub struct Parallelism {
#[serde(skip_serializing_if = "Option::is_none")]
pub outer: Option<ParallelismOuter>,
#[serde(skip_serializing_if = "Option::is_none")]
pub attention: Option<ParallelismAttention>,
#[serde(skip_serializing_if = "Option::is_none")]
pub experts: Option<ParallelismExperts>,
}
impl Parallelism {
pub fn merge_from(&mut self, other: &Self) {
if let Some(other) = &other.outer {
self.outer.get_or_insert_default().merge_from(other);
}
if let Some(other) = &other.attention {
self.attention.get_or_insert_default().merge_from(other);
}
if let Some(other) = &other.experts {
self.experts.get_or_insert_default().merge_from(other);
}
}
}
#[derive(Clone, Debug, Default, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(default, deny_unknown_fields)]
pub struct ParallelismOuter {
#[schemars(range(min = 1))]
pub tensor_parallel_size: Option<u32>,
#[schemars(range(min = 1))]
pub pipeline_parallel_size: Option<u32>,
}
impl ParallelismOuter {
fn merge_from(&mut self, other: &Self) {
if other.tensor_parallel_size.is_some() {
self.tensor_parallel_size = other.tensor_parallel_size;
}
if other.pipeline_parallel_size.is_some() {
self.pipeline_parallel_size = other.pipeline_parallel_size;
}
}
}
#[derive(Clone, Debug, Default, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(default, deny_unknown_fields)]
pub struct ParallelismAttention {
#[schemars(range(min = 1))]
pub tensor_parallel_size: Option<u32>,
#[schemars(range(min = 1))]
pub data_parallel_size: Option<u32>,
#[schemars(range(min = 1))]
pub context_parallel_size: Option<u32>,
}
impl ParallelismAttention {
fn merge_from(&mut self, other: &Self) {
if other.tensor_parallel_size.is_some() {
self.tensor_parallel_size = other.tensor_parallel_size;
}
if other.data_parallel_size.is_some() {
self.data_parallel_size = other.data_parallel_size;
}
if other.context_parallel_size.is_some() {
self.context_parallel_size = other.context_parallel_size;
}
}
}
#[derive(Clone, Debug, Default, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(default, deny_unknown_fields)]
pub struct ParallelismExperts {
#[schemars(range(min = 1))]
pub tensor_parallel_size: Option<u32>,
#[schemars(range(min = 1))]
pub data_parallel_size: Option<u32>,
#[schemars(range(min = 1))]
pub expert_parallel_size: Option<u32>,
#[schemars(range(min = 1))]
pub dense_tensor_parallel_size: Option<u32>,
}
impl ParallelismExperts {
fn merge_from(&mut self, other: &Self) {
if other.tensor_parallel_size.is_some() {
self.tensor_parallel_size = other.tensor_parallel_size;
}
if other.data_parallel_size.is_some() {
self.data_parallel_size = other.data_parallel_size;
}
if other.expert_parallel_size.is_some() {
self.expert_parallel_size = other.expert_parallel_size;
}
if other.dense_tensor_parallel_size.is_some() {
self.dense_tensor_parallel_size = other.dense_tensor_parallel_size;
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ServeModelInput {
pub locator: String,
pub served_name: String,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct EndpointAssignment {
pub host: String,
pub port: u16,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ClientEndpointInput {
pub protocol: EndpointProtocol,
pub host: String,
pub port: u16,
pub api_path: String,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
pub enum EvalDefinitionInput {
#[serde(rename = "openai_smoke")]
OpenAiSmoke {
prompt: String,
max_tokens: u32,
timeout_seconds: u64,
},
LmEval {
task: String,
dataset: Option<String>,
split: Option<String>,
limit: Option<u32>,
few_shot: Option<u32>,
seed: Option<u64>,
max_tokens: Option<u32>,
concurrency: Option<u32>,
metric: String,
threshold: f64,
timeout_seconds: u64,
},
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct BenchDefinitionInput {
pub input_tokens: u32,
pub output_tokens: u32,
pub seed: u64,
pub temperature: f64,
pub timeout_seconds: u64,
#[serde(default)]
pub reset_prefix_cache: bool,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(untagged)]
pub enum SettingValue {
Bool(bool),
Integer(i64),
Float(f64),
String(String),
Array(Vec<SettingValue>),
Object(BTreeMap<String, SettingValue>),
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(tag = "status", rename_all = "snake_case", deny_unknown_fields)]
pub enum AdapterResponse {
Ok {
protocol_version: ProtocolVersion,
result: Box<AdapterResult>,
},
Error {
protocol_version: ProtocolVersion,
error: AdapterError,
},
}
impl AdapterResponse {
#[must_use]
pub const fn protocol_version(&self) -> ProtocolVersion {
match self {
Self::Ok {
protocol_version, ..
}
| Self::Error {
protocol_version, ..
} => *protocol_version,
}
}
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(tag = "operation", rename_all = "snake_case", deny_unknown_fields)]
pub enum AdapterResult {
PlanServe { output: Box<PlanServeResult> },
RenderServe { output: Box<RenderServeResult> },
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct PlanServeResult {
pub integration: IntegrationIdentity,
pub effective_settings: BTreeMap<String, SettingValue>,
pub effective_parallelism: Parallelism,
pub roles: Vec<ServeRoleResult>,
pub replicas: Vec<ServeReplicaRequirement>,
pub links: Vec<ServeRoleLink>,
pub public_endpoint: PublicEndpointRequirement,
pub endpoint: EndpointRequirement,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ServeReplicaRequirement {
pub id: String,
pub role_id: String,
pub replica_index: u32,
pub accelerator_count: u32,
pub ports: Vec<String>,
pub primary_ports: Vec<String>,
pub primary_readiness: ReadinessProbe,
pub worker_readiness: ReadinessProbe,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub capture_target: Option<CaptureTargetRequirement>,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ServeProcessAllocation {
pub process_id: String,
pub role_id: String,
pub replica_id: String,
pub replica_index: u32,
pub rank: u32,
pub machine_id: String,
pub model_locator: String,
pub runtime_cache_root: String,
pub devices: Vec<u32>,
pub endpoint: EndpointAssignment,
pub ports: BTreeMap<String, EndpointAssignment>,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
pub enum ServeRoleLink {
RequestRouting {
source: String,
targets: Vec<String>,
},
KvTransfer {
source: String,
target: String,
mechanism: KvTransferMechanism,
},
Bootstrap {
source: String,
target: String,
port: String,
},
SideChannel {
source: String,
target: String,
port: String,
},
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
pub enum PublicEndpointRequirement {
Replica { replica_id: String },
BuiltinProxy {
process_id: String,
role_id: String,
prefill_role: String,
decode_role: String,
readiness: ReadinessProbe,
},
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct CaptureTargetRequirement {
pub control: CaptureControlRequirement,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct CaptureControlRequirement {
pub start_path: String,
pub stop_path: String,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct RenderServeResult {
pub integration: IntegrationIdentity,
pub processes: Vec<RenderedServeProcess>,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct RenderedServeProcess {
pub id: String,
pub process: ProcessSpec,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct HttpActionSpec {
pub method: HttpMethod,
pub path: String,
}
#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum HttpMethod {
Post,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct IntegrationIdentity {
pub adapter_id: String,
pub adapter_version: String,
pub framework: String,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ProcessSpec {
pub argv: Vec<String>,
pub env: BTreeMap<String, String>,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
pub enum ReadinessProbe {
Http { path: String },
ProcessAlive,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct EndpointRequirement {
pub protocol: EndpointProtocol,
pub api_path: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prefix_cache_reset: Option<HttpActionSpec>,
}
#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum EndpointProtocol {
Http,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct AdapterError {
pub code: AdapterErrorCode,
pub message: String,
}
#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum AdapterErrorCode {
InvalidRequest,
UnsupportedProtocolVersion,
InvalidSettings,
Internal,
UnsupportedOperation,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct EvalClientRequest {
pub protocol_version: ProtocolVersion,
pub endpoint: ClientEndpointInput,
pub model: ServeModelInput,
pub definition: EvalDefinitionInput,
pub artifact_dir: PathBuf,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct BenchClientRequest {
pub protocol_version: ProtocolVersion,
pub endpoint: ClientEndpointInput,
pub model: ServeModelInput,
pub definition: BenchDefinitionInput,
pub case: BenchCaseInput,
pub artifact_dir: PathBuf,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct BenchCaseInput {
pub load_shape: BenchLoadInput,
pub request_count: u32,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
pub enum BenchLoadInput {
ConcurrencyLimited { concurrency: u32 },
RequestRateLimited {
request_rate: f64,
burstiness: Option<f64>,
},
UnboundedRequestRate,
}
#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ClientStatus {
Succeeded,
Failed,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct EvalClientResult {
pub schema_version: u32,
pub status: ClientStatus,
pub metrics: BTreeMap<String, f64>,
pub native_command: Vec<String>,
pub raw_artifacts: Vec<RawArtifact>,
pub error: Option<String>,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct BenchClientResult {
pub schema_version: u32,
pub status: ClientStatus,
pub completed_requests: u64,
pub failed_requests: u64,
pub normalization_schema: String,
pub metrics: BTreeMap<String, f64>,
pub native_command: Vec<String>,
pub native_exit_code: Option<i32>,
pub raw_artifacts: Vec<RawArtifact>,
pub error: Option<String>,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct RawArtifact {
pub name: String,
pub kind: String,
pub path: PathBuf,
}
#[derive(Clone, Debug, Deserialize, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct AdapterProtocol {
pub request: AdapterRequest,
pub response: AdapterResponse,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub eval_client_request: Option<EvalClientRequest>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub eval_client_result: Option<EvalClientResult>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub bench_client_request: Option<BenchClientRequest>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub bench_client_result: Option<BenchClientResult>,
}