use serde::{Deserialize, Deserializer, Serialize, Serializer, de};
use serde_json::{Map, Value};
use std::collections::HashSet;
use crate::types::{
ContainerInfo, ContentBlock, MessageCreateParams, MessageRole, Model, OutputConfig, SpeedMode,
StopReason, ThinkingConfig, Usage,
};
use crate::{Error, Result};
const MAX_FALLBACK_MODELS: usize = 3;
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct FallbackModel {
pub model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub thinking: Option<ThinkingConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub output_config: Option<OutputConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub speed: Option<SpeedMode>,
}
impl FallbackModel {
pub fn new(model: impl Into<String>) -> Self {
Self {
model: model.into(),
max_tokens: None,
thinking: None,
output_config: None,
speed: None,
}
}
pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
self.max_tokens = Some(max_tokens);
self
}
pub fn with_thinking(mut self, thinking: ThinkingConfig) -> Self {
self.thinking = Some(thinking);
self
}
pub fn with_output_config(mut self, output_config: OutputConfig) -> Self {
self.output_config = Some(output_config);
self
}
pub fn with_speed(mut self, speed: SpeedMode) -> Self {
self.speed = Some(speed);
self
}
}
#[non_exhaustive]
#[derive(Debug, Clone, Default, PartialEq)]
pub enum ServerFallbacks {
#[default]
Default,
Models(Vec<FallbackModel>),
}
impl ServerFallbacks {
pub fn default_routing() -> Self {
Self::Default
}
pub fn models(models: Vec<FallbackModel>) -> Result<Self> {
let fallbacks = Self::Models(models);
fallbacks.validate(None)?;
Ok(fallbacks)
}
fn validate(&self, primary_model: Option<&Model>) -> Result<()> {
let Self::Models(models) = self else {
return Ok(());
};
if models.is_empty() || models.len() > MAX_FALLBACK_MODELS {
return Err(Error::validation(
format!(
"Fallback model list must contain between 1 and {MAX_FALLBACK_MODELS} entries"
),
Some("fallbacks".to_string()),
));
}
let primary_model = primary_model.map(ToString::to_string);
let mut seen = HashSet::with_capacity(models.len());
for fallback in models {
if fallback.model.trim().is_empty() {
return Err(Error::validation(
"Fallback model identifier cannot be empty".to_string(),
Some("fallbacks.model".to_string()),
));
}
if primary_model.as_deref() == Some(fallback.model.as_str()) {
return Err(Error::validation(
"Fallback model must differ from the requested model".to_string(),
Some("fallbacks.model".to_string()),
));
}
if !seen.insert(fallback.model.as_str()) {
return Err(Error::validation(
format!("Duplicate fallback model: {}", fallback.model),
Some("fallbacks.model".to_string()),
));
}
}
Ok(())
}
}
impl Serialize for ServerFallbacks {
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
where
S: Serializer,
{
match self {
Self::Default => serializer.serialize_str("default"),
Self::Models(models) => models.serialize(serializer),
}
}
}
impl<'de> Deserialize<'de> for ServerFallbacks {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(untagged)]
enum Wire {
Default(String),
Models(Vec<FallbackModel>),
}
match Wire::deserialize(deserializer)? {
Wire::Default(value) if value == "default" => Ok(Self::Default),
Wire::Default(value) => {
Err(de::Error::custom(format!("unsupported fallback routing mode: {value}")))
}
Wire::Models(models) => Ok(Self::Models(models)),
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ServerFallbackRequest {
#[serde(flatten)]
pub(crate) params: MessageCreateParams,
pub(crate) fallbacks: ServerFallbacks,
}
impl ServerFallbackRequest {
pub fn new(params: MessageCreateParams, fallbacks: ServerFallbacks) -> Result<Self> {
let request = Self { params, fallbacks };
request.validate()?;
Ok(request)
}
pub fn default_routing(params: MessageCreateParams) -> Result<Self> {
Self::new(params, ServerFallbacks::Default)
}
pub fn explicit(params: MessageCreateParams, models: Vec<FallbackModel>) -> Result<Self> {
Self::new(params, ServerFallbacks::Models(models))
}
pub fn params(&self) -> &MessageCreateParams {
&self.params
}
pub fn fallbacks(&self) -> &ServerFallbacks {
&self.fallbacks
}
pub fn validate(&self) -> Result<()> {
self.params.validate()?;
self.fallbacks.validate(Some(&self.params.model))
}
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct FallbackRoute {
pub model: String,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
enum FallbackBlockType {
#[serde(rename = "fallback")]
Fallback,
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct FallbackContentBlock {
#[serde(rename = "type")]
block_type: FallbackBlockType,
pub from: FallbackRoute,
pub to: FallbackRoute,
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ServerFallbackContentBlock {
Fallback(FallbackContentBlock),
Standard(ContentBlock),
Unknown(Value),
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct RefusalDetails {
#[serde(rename = "type")]
pub detail_type: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub category: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub explanation: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub fallback_credit_token: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub fallback_has_prefill_claim: Option<bool>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ServerFallbackUsage {
#[serde(flatten)]
pub usage: Usage,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub iterations: Vec<Value>,
}
impl ServerFallbackUsage {
pub fn fallback_ran(&self) -> bool {
self.iterations.iter().any(|iteration| {
iteration.get("type").and_then(Value::as_str) == Some("fallback_message")
})
}
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ServerFallbackMessage {
pub id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub container: Option<ContainerInfo>,
pub content: Vec<ServerFallbackContentBlock>,
pub model: Model,
pub role: MessageRole,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop_reason: Option<StopReason>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stop_details: Option<RefusalDetails>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop_sequence: Option<String>,
pub r#type: String,
pub usage: ServerFallbackUsage,
}
impl ServerFallbackMessage {
pub fn served_by_fallback(&self) -> bool {
let has_handoff = self
.content
.iter()
.any(|block| matches!(block, ServerFallbackContentBlock::Fallback(_)));
(has_handoff || self.usage.fallback_ran()) && self.stop_reason != Some(StopReason::Refusal)
}
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ServerFallbackMessageStartEvent {
pub message: ServerFallbackMessage,
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ServerFallbackContentBlockStartEvent {
pub content_block: ServerFallbackContentBlock,
pub index: usize,
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ServerFallbackMessageDelta {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stop_reason: Option<StopReason>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stop_sequence: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stop_details: Option<RefusalDetails>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ServerFallbackDeltaUsage {
#[serde(flatten)]
pub usage: crate::types::MessageDeltaUsage,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub iterations: Vec<Value>,
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ServerFallbackMessageDeltaEvent {
pub delta: ServerFallbackMessageDelta,
pub usage: ServerFallbackDeltaUsage,
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum ServerFallbackStreamEvent {
#[serde(rename = "ping")]
Ping,
#[serde(rename = "message_start")]
MessageStart(ServerFallbackMessageStartEvent),
#[serde(rename = "message_delta")]
MessageDelta(ServerFallbackMessageDeltaEvent),
#[serde(rename = "content_block_start")]
ContentBlockStart(ServerFallbackContentBlockStartEvent),
#[serde(rename = "content_block_delta")]
ContentBlockDelta(crate::types::ContentBlockDeltaEvent),
#[serde(rename = "content_block_stop")]
ContentBlockStop(crate::types::ContentBlockStopEvent),
#[serde(rename = "message_stop")]
MessageStop(crate::types::MessageStopEvent),
#[serde(rename = "tool_input_start")]
ToolInputStart {
tool_use_id: String,
parameter_name: String,
},
#[serde(rename = "tool_input_delta")]
ToolInputDelta {
tool_use_id: String,
parameter_name: String,
value_fragment: String,
},
#[serde(rename = "compaction")]
CompactionEvent(crate::types::CompactionMetadata),
#[serde(rename = "stream_error")]
StreamError {
error: crate::types::ApiError,
},
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn default_and_explicit_wire_shapes() {
assert_eq!(serde_json::to_value(ServerFallbacks::Default).unwrap(), json!("default"));
let explicit =
ServerFallbacks::models(vec![FallbackModel::new("claude-opus-4-8")]).unwrap();
assert_eq!(serde_json::to_value(explicit).unwrap(), json!([{"model": "claude-opus-4-8"}]));
}
#[test]
fn explicit_topology_is_validated() {
let params =
MessageCreateParams::simple("hello", Model::Custom("claude-fable-5".to_string()));
let duplicate =
vec![FallbackModel::new("claude-opus-4-8"), FallbackModel::new("claude-opus-4-8")];
assert!(ServerFallbackRequest::explicit(params.clone(), duplicate).is_err());
assert!(ServerFallbackRequest::explicit(params.clone(), Vec::new()).is_err());
assert!(
ServerFallbackRequest::explicit(
params.clone(),
vec![
FallbackModel::new("model-1"),
FallbackModel::new("model-2"),
FallbackModel::new("model-3"),
FallbackModel::new("model-4"),
],
)
.is_err()
);
assert!(
ServerFallbackRequest::explicit(params, vec![FallbackModel::new("claude-fable-5")])
.is_err()
);
}
#[test]
fn response_and_stream_fallback_blocks_deserialize() {
let message: ServerFallbackMessage = serde_json::from_value(json!({
"id": "msg_test",
"content": [{
"type": "fallback",
"from": {"model": "claude-fable-5"},
"to": {"model": "claude-opus-4-8"}
}, {"type": "text", "text": "ok"}],
"model": "claude-opus-4-8",
"role": "assistant",
"stop_reason": "end_turn",
"stop_details": null,
"stop_sequence": null,
"type": "message",
"usage": {
"input_tokens": 1,
"output_tokens": 1,
"iterations": [{"type": "fallback_message"}]
}
}))
.unwrap();
assert!(message.served_by_fallback());
assert!(matches!(message.content[0], ServerFallbackContentBlock::Fallback(_)));
let event: ServerFallbackStreamEvent = serde_json::from_value(json!({
"type": "content_block_start",
"index": 0,
"content_block": {
"type": "fallback",
"from": {"model": "claude-fable-5"},
"to": {"model": "claude-opus-4-8"}
}
}))
.unwrap();
assert!(matches!(event, ServerFallbackStreamEvent::ContentBlockStart(_)));
let event: ServerFallbackStreamEvent = serde_json::from_value(json!({
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": null},
"usage": {
"input_tokens": 1,
"output_tokens": 1,
"iterations": [{"type": "fallback_message", "model": "claude-opus-4-8"}]
}
}))
.unwrap();
let ServerFallbackStreamEvent::MessageDelta(event) = event else {
panic!("expected message delta");
};
assert_eq!(event.usage.iterations[0]["type"], "fallback_message");
}
}