use std::collections::HashMap;
use bytes::Bytes;
pub use rskit_ai::{
Capabilities, Model, Provider as ModelProvider, StreamEvent, StreamEventRef, Usage,
};
use rskit_errors::{AppError, ErrorCode};
use serde::{Deserialize, Serialize};
use thiserror::Error;
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct PredictRequest {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub request_id: Option<String>,
pub model_name: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub model_version: Option<String>,
#[serde(default)]
pub inputs: HashMap<String, Value>,
#[serde(default)]
pub parameters: HashMap<String, serde_json::Value>,
#[serde(default)]
pub options: serde_json::Value,
#[serde(default)]
pub metadata: HashMap<String, String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PredictResponse {
#[serde(default)]
pub outputs: HashMap<String, Value>,
#[serde(default)]
pub usage: Usage,
pub model: Model,
pub status: PredictStatus,
#[serde(default)]
pub metadata: HashMap<String, String>,
}
impl Default for PredictResponse {
fn default() -> Self {
Self {
outputs: HashMap::new(),
usage: Usage::default(),
model: Model {
name: String::new(),
provider: ModelProvider::Custom("unknown".to_string()),
version: None,
capabilities: Capabilities::default(),
},
status: PredictStatus::Success,
metadata: HashMap::new(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum PredictStatus {
Success,
PartialSuccess,
Error {
reason: String,
},
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum Value {
Text {
text: String,
},
Bytes {
bytes: Bytes,
},
Tensor {
tensor: Tensor,
},
Json {
json: serde_json::Value,
},
}
#[derive(Debug, Clone, PartialEq)]
pub struct Tensor {
pub dtype: String,
pub shape: Vec<i64>,
pub data: TensorData,
}
impl Serialize for Tensor {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
use serde::ser::SerializeStruct;
let mut state = serializer.serialize_struct("Tensor", 3)?;
state.serialize_field("dtype", &self.dtype)?;
state.serialize_field("shape", &self.shape)?;
state.serialize_field("data", &self.data)?;
state.end()
}
}
impl<'de> Deserialize<'de> for Tensor {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(Deserialize)]
struct TensorWire {
dtype: String,
shape: Vec<i64>,
data: serde_json::Value,
}
let wire = TensorWire::deserialize(deserializer)?;
let dtype = wire.dtype.to_ascii_uppercase();
let data = match dtype.as_str() {
"FP32" => TensorData::F32(
serde_json::from_value(wire.data).map_err(serde::de::Error::custom)?,
),
"FP64" => TensorData::F64(
serde_json::from_value(wire.data).map_err(serde::de::Error::custom)?,
),
"INT32" => TensorData::I32(
serde_json::from_value(wire.data).map_err(serde::de::Error::custom)?,
),
"INT64" => TensorData::I64(
serde_json::from_value(wire.data).map_err(serde::de::Error::custom)?,
),
"UINT8" => {
TensorData::U8(serde_json::from_value(wire.data).map_err(serde::de::Error::custom)?)
}
"BOOL" => TensorData::Bool(
serde_json::from_value(wire.data).map_err(serde::de::Error::custom)?,
),
"BYTES" => TensorData::Bytes(
serde_json::from_value(wire.data).map_err(serde::de::Error::custom)?,
),
_ => serde_json::from_value(wire.data).map_err(serde::de::Error::custom)?,
};
Ok(Self {
dtype: wire.dtype,
shape: wire.shape,
data,
})
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum TensorData {
F32(Vec<f32>),
F64(Vec<f64>),
I32(Vec<i32>),
I64(Vec<i64>),
U8(Vec<u8>),
Bool(Vec<bool>),
Bytes(Vec<Bytes>),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct InferenceDescriptor {
pub name: String,
pub description: String,
pub serving_protocol: ServingProtocol,
pub envelope: rskit_tool::Envelope,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ServingProtocol {
KServeV2Http,
KServeV2Grpc,
VllmRest,
TgiRest,
BentoMl,
OnnxRuntime,
TfServing,
Custom,
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum InferenceError {
#[error("transport: {0}")]
Transport(#[source] AppError),
#[error("decode: {0}")]
Decode(String),
#[error("server: status={status}, body={body}")]
Server {
status: u16,
body: String,
},
#[error("authorization denied: {0}")]
Authorization(String),
#[error("invalid input: {0}")]
InvalidInput(String),
#[error("policy: {0}")]
Policy(String),
#[error("not implemented: {0}")]
NotImplemented(&'static str),
#[error("cancelled")]
Cancelled,
#[error("timeout")]
Timeout,
}
impl From<AppError> for InferenceError {
fn from(value: AppError) -> Self {
match value.code() {
ErrorCode::Timeout => Self::Timeout,
ErrorCode::Cancelled => Self::Cancelled,
ErrorCode::ExternalService
| ErrorCode::ServiceUnavailable
| ErrorCode::ConnectionFailed => Self::Transport(value),
ErrorCode::InvalidInput | ErrorCode::MissingField | ErrorCode::InvalidFormat => {
Self::InvalidInput(value.to_string())
}
_ => Self::Policy(value.to_string()),
}
}
}
impl From<InferenceError> for AppError {
fn from(value: InferenceError) -> Self {
match value {
InferenceError::Timeout => AppError::timeout("inference"),
InferenceError::Cancelled => AppError::cancelled("inference"),
InferenceError::Authorization(reason) => AppError::forbidden(reason),
InferenceError::InvalidInput(reason) => AppError::new(ErrorCode::InvalidInput, reason),
InferenceError::Server { status, body } => AppError::new(
ErrorCode::ExternalService,
format!("inference runtime returned status {status}: {body}"),
),
InferenceError::Transport(error) => AppError::new(
ErrorCode::ExternalService,
format!("inference transport failed: {}", error.message()),
)
.with_cause(error),
other => AppError::new(ErrorCode::ExternalService, other.to_string()),
}
}
}