use a2a_protocol_types::task::TaskState;
const ALL_STATES: [TaskState; 9] = [
TaskState::Unspecified,
TaskState::Submitted,
TaskState::Working,
TaskState::InputRequired,
TaskState::AuthRequired,
TaskState::Completed,
TaskState::Failed,
TaskState::Canceled,
TaskState::Rejected,
];
const TERMINAL: [TaskState; 4] = [
TaskState::Completed,
TaskState::Failed,
TaskState::Canceled,
TaskState::Rejected,
];
const LIVE_TARGETS: [TaskState; 7] = [
TaskState::Working,
TaskState::InputRequired,
TaskState::AuthRequired,
TaskState::Completed,
TaskState::Failed,
TaskState::Canceled,
TaskState::Rejected,
];
const LIVE_SOURCES: [TaskState; 4] = [
TaskState::Submitted,
TaskState::Working,
TaskState::InputRequired,
TaskState::AuthRequired,
];
fn assert_transitions(from: TaskState, valid: &[TaskState]) {
for &target in &ALL_STATES {
let expected = valid.contains(&target);
let actual = from.can_transition_to(target);
assert_eq!(
actual, expected,
"{from} -> {target}: expected can_transition_to = {expected}, got {actual}"
);
}
}
#[test]
fn test_unspecified_transitions() {
assert_transitions(TaskState::Unspecified, &ALL_STATES);
assert!(
!TaskState::Unspecified.is_terminal(),
"Unspecified must not be terminal"
);
}
#[test]
fn test_submitted_transitions() {
assert_transitions(TaskState::Submitted, &LIVE_TARGETS);
assert!(
TaskState::Submitted.can_transition_to(TaskState::Completed),
"Submitted -> Completed must be valid (one-step agents never enter Working)"
);
assert!(
TaskState::Submitted.can_transition_to(TaskState::InputRequired),
"Submitted -> InputRequired must be valid (agent needs input immediately)"
);
}
#[test]
fn test_working_transitions() {
assert_transitions(TaskState::Working, &LIVE_TARGETS);
assert!(
TaskState::Working.can_transition_to(TaskState::Working),
"Working -> Working must be valid (progress-narration refresh)"
);
assert!(
TaskState::Working.can_transition_to(TaskState::Rejected),
"Working -> Rejected must be valid (§4.1.3 'or later')"
);
}
#[test]
fn test_input_required_transitions() {
assert_transitions(TaskState::InputRequired, &LIVE_TARGETS);
}
#[test]
fn test_auth_required_transitions() {
assert_transitions(TaskState::AuthRequired, &LIVE_TARGETS);
}
#[test]
fn test_terminal_states_have_no_outgoing_transitions() {
for &terminal in &TERMINAL {
assert!(terminal.is_terminal(), "{terminal} must report is_terminal");
assert_transitions(terminal, &[]);
}
}
#[test]
fn test_non_terminal_states_report_not_terminal() {
assert!(!TaskState::Unspecified.is_terminal());
for &state in &LIVE_SOURCES {
assert!(
!state.is_terminal(),
"{state} must report is_terminal() == false"
);
}
}
#[test]
fn test_nothing_reenters_entry_or_default_state() {
for &from in &LIVE_SOURCES {
assert!(
!from.can_transition_to(TaskState::Submitted),
"{from} -> Submitted must be invalid (Submitted is the entry state)"
);
assert!(
!from.can_transition_to(TaskState::Unspecified),
"{from} -> Unspecified must be invalid (proto default)"
);
}
}
#[test]
fn test_no_self_transitions_except_unspecified_and_working() {
for &state in &ALL_STATES {
let can_self = state.can_transition_to(state);
let expected = matches!(
state,
TaskState::Unspecified
| TaskState::Working
| TaskState::InputRequired
| TaskState::AuthRequired
);
assert_eq!(
can_self, expected,
"{state} -> {state} self-transition: expected {expected}, got {can_self}"
);
}
}
#[test]
fn test_full_transition_matrix() {
fn expected(from: TaskState, to: TaskState) -> bool {
if from.is_terminal() {
return false;
}
if matches!(from, TaskState::Unspecified) {
return true;
}
!matches!(to, TaskState::Submitted | TaskState::Unspecified)
}
for &from in &ALL_STATES {
for &to in &ALL_STATES {
assert_eq!(
from.can_transition_to(to),
expected(from, to),
"matrix mismatch at {from} -> {to}"
);
}
}
}
#[test]
fn test_valid_transition_counts() {
for &state in &ALL_STATES {
let actual = ALL_STATES
.iter()
.filter(|&&target| state.can_transition_to(target))
.count();
let expected = match state {
TaskState::Unspecified => ALL_STATES.len(),
s if s.is_terminal() => 0,
_ => LIVE_TARGETS.len(),
};
assert_eq!(
actual, expected,
"{state} should have {expected} valid outgoing transitions, got {actual}"
);
}
let total = ALL_STATES
.iter()
.flat_map(|&from| ALL_STATES.iter().map(move |&to| (from, to)))
.filter(|&(from, to)| from.can_transition_to(to))
.count();
assert_eq!(total, 37, "expected 37 valid transitions across the matrix");
}
#[test]
fn test_terminal_and_interrupted_classification() {
for &state in &ALL_STATES {
assert_eq!(
state.is_terminal(),
TERMINAL.contains(&state),
"{state} is_terminal classification"
);
assert_eq!(
state.is_interrupted(),
matches!(state, TaskState::InputRequired | TaskState::AuthRequired),
"{state} is_interrupted classification"
);
assert!(
!(state.is_terminal() && state.is_interrupted()),
"{state} cannot be both terminal and interrupted"
);
}
}