use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum Interrupt {
Before(String),
After(String),
Dynamic {
message: String,
data: Option<Value>,
},
}
impl std::fmt::Display for Interrupt {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Before(node) => write!(f, "Interrupt before '{}'", node),
Self::After(node) => write!(f, "Interrupt after '{}'", node),
Self::Dynamic { message, .. } => write!(f, "Dynamic interrupt: {}", message),
}
}
}
pub fn interrupt(message: &str) -> Interrupt {
Interrupt::Dynamic { message: message.to_string(), data: None }
}
pub fn interrupt_with_data(message: &str, data: Value) -> Interrupt {
Interrupt::Dynamic { message: message.to_string(), data: Some(data) }
}
pub const INTERRUPT_METADATA_KEY: &str = "adk.graph.interrupt";
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GraphInterruptPayload {
pub kind: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub node: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub message: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub data: Option<Value>,
pub thread_id: String,
pub checkpoint_id: String,
}
impl GraphInterruptPayload {
pub fn new(interrupt: &Interrupt, thread_id: &str, checkpoint_id: &str) -> Self {
let (kind, node, message, data) = match interrupt {
Interrupt::Before(node) => ("before", Some(node.clone()), None, None),
Interrupt::After(node) => ("after", Some(node.clone()), None, None),
Interrupt::Dynamic { message, data } => {
("dynamic", None, Some(message.clone()), data.clone())
}
};
Self {
kind: kind.to_string(),
node,
message,
data,
thread_id: thread_id.to_string(),
checkpoint_id: checkpoint_id.to_string(),
}
}
pub fn from_event(event: &adk_core::Event) -> Option<Self> {
let raw = event.provider_metadata.get(INTERRUPT_METADATA_KEY)?;
serde_json::from_str(raw).ok()
}
pub fn to_metadata_value(&self) -> String {
serde_json::to_string(self).unwrap_or_default()
}
}