use anyhow::Result;
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use std::path::{Path, PathBuf};
use crate::identity::AgentIdentity;
use crate::providers::{ProviderKind, ProviderOptions};
use crate::session::context::{CompressionStats, MessageRole, SessionContext};
use crate::session::execution::{
DEFAULT_MAX_PROMPT_BYTES, prepare_provider_prompt, run_provider_command,
};
use crate::session::output::{BuildStatus, LogLevel, OutputParser, ParsedOutput};
use crate::session::persistence::{
ExecutionSessionStatus, PersistenceManager, SessionMetadata, SessionState,
};
use crate::session::types::{AttentionState, ExecutionSessionId};
const DEFAULT_CONTINUATION_PROMPT: &str = "The previous turn completed but the task is still active. Continue with the next sub-step. Stop when the task is fully done or you cannot make progress.";
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BridgeResult {
pub raw: String,
pub parsed: ParsedOutput,
pub success: bool,
#[serde(default)]
pub duration_ms: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub compression_ratio: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tokens_in: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tokens_out: Option<u64>,
#[serde(default)]
pub attention: AttentionState,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub tool_names: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub total_cost_usd: Option<f64>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub fallbacks_used: Vec<String>,
}
fn is_rate_limit_error(err: &anyhow::Error) -> bool {
let text = err.to_string().to_lowercase();
[
"rate limit",
"rate_limit",
"429",
"too many requests",
"quota",
"overloaded",
]
.iter()
.any(|marker| text.contains(marker))
}
fn apply_rate_limit_fallback(
options: &mut MovementExecOptions,
(next_provider, next_model): (ProviderKind, Option<String>),
) -> String {
let from = options
.provider
.unwrap_or(ProviderKind::Claude)
.as_str()
.to_string();
options.provider = Some(next_provider);
if next_model.is_some() {
options.model = next_model;
}
options.session_id = None;
options.continuation = ContinuationPolicy::SingleTurn;
from
}
fn fallback_notice_prompt(original_prompt: &str) -> String {
format!(
"# Provider fallback notice\nA previous provider hit a rate limit mid-task. \
You are taking over fresh — no prior session context is available beyond this prompt.\n\n{}",
original_prompt
)
}
fn read_a2a_endpoint() -> Option<String> {
std::env::var("CCSWARM_A2A_ENDPOINT")
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
fn resolve_a2a_endpoint(options: &MovementExecOptions) -> Option<String> {
options
.a2a_endpoint
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
.or_else(read_a2a_endpoint)
}
fn prepare_a2a_prompt(prompt: &str, system_prompt: Option<&str>) -> String {
match system_prompt.filter(|value| !value.trim().is_empty()) {
Some(system_prompt) => format!(
"# System instructions\n{}\n\n{}",
system_prompt.trim(),
prompt
),
None => prompt.to_string(),
}
}
struct StreamProjection {
text: String,
tool_names: Vec<String>,
tokens: Option<(u64, u64)>,
total_cost_usd: Option<f64>,
}
fn project_stream(raw_stdout: &str) -> StreamProjection {
let summary = crate::providers::claude_stream::parse_stream(raw_stdout);
if !summary.tool_uses.is_empty() {
tracing::debug!(
"stream-json: {} tool_use blocks ({:?})",
summary.tool_uses.len(),
summary
.tool_uses
.iter()
.map(|t| t.name.as_str())
.collect::<Vec<_>>()
);
}
let tokens = summary.usage.last().map(|u| {
(
u.input_tokens + u.cache_creation_input_tokens + u.cache_read_input_tokens,
u.output_tokens,
)
});
StreamProjection {
text: summary.result_text,
tool_names: summary.tool_uses.into_iter().map(|t| t.name).collect(),
tokens,
total_cost_usd: summary.total_cost_usd,
}
}
#[derive(Debug, Clone)]
pub enum ContinuationPolicy {
SingleTurn,
MultiTurn {
max_turns: u32,
continuation_prompt: String,
},
}
impl ContinuationPolicy {
pub fn multi_turn(max_turns: u32) -> Self {
Self::MultiTurn {
max_turns,
continuation_prompt: DEFAULT_CONTINUATION_PROMPT.to_string(),
}
}
}
#[allow(clippy::derivable_impls)]
impl Default for ContinuationPolicy {
fn default() -> Self {
Self::SingleTurn
}
}
#[derive(Debug, Clone, Default)]
pub struct MovementExecOptions {
pub(crate) provider: Option<ProviderKind>,
pub tools: Vec<String>,
pub model: Option<String>,
pub system_prompt: Option<String>,
pub max_budget: Option<f64>,
pub worktree_name: Option<String>,
pub session_id: Option<String>,
pub continuation: ContinuationPolicy,
pub a2a_endpoint: Option<String>,
pub(crate) rate_limit_fallbacks: Vec<(ProviderKind, Option<String>)>,
}
pub struct A2ABridge {
context_histories: DashMap<String, SessionContext>,
output_parser: OutputParser,
persistence: PersistenceManager,
}
#[derive(Debug, Clone)]
struct BridgeExecutionMetadata {
provider: ProviderKind,
session_id: Option<String>,
same_thread_continuation: bool,
}
#[derive(Debug, Clone)]
struct BridgeExecution {
result: BridgeResult,
metadata: BridgeExecutionMetadata,
}
impl A2ABridge {
pub fn new(storage_path: PathBuf) -> Self {
Self {
context_histories: DashMap::new(),
output_parser: OutputParser::new(),
persistence: PersistenceManager::new(storage_path),
}
}
pub fn register_agent(&self, agent_id: &str) -> Result<()> {
let session_id = ExecutionSessionId::new();
let context = SessionContext::new(session_id);
self.context_histories.insert(agent_id.to_string(), context);
tracing::info!("Registered agent '{}' with A2ABridge", agent_id);
Ok(())
}
pub async fn execute_task(
&self,
agent_id: &str,
prompt: &str,
_identity: &AgentIdentity,
working_dir: &Path,
) -> Result<BridgeResult> {
self.execute_with_retry(
agent_id,
prompt,
_identity,
working_dir,
None,
0,
1000,
&MovementExecOptions::default(),
)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn execute_with_retry(
&self,
agent_id: &str,
prompt: &str,
_identity: &AgentIdentity,
working_dir: &Path,
agent_name: Option<&str>,
max_retries: u32,
retry_delay_ms: u64,
options: &MovementExecOptions,
) -> Result<BridgeResult> {
let mut current_options = options.clone();
let mut current_prompt = prompt.to_string();
let mut fallback_index = 0usize;
let mut fallbacks_used: Vec<String> = Vec::new();
loop {
let mut last_err = None;
let attempts = max_retries + 1;
'attempts: for attempt in 0..attempts {
if attempt > 0 {
let delay = retry_delay_ms * 2u64.pow(attempt - 1);
tracing::info!(
"Retrying provider CLI (attempt {}/{}, delay {}ms)",
attempt + 1,
attempts,
delay
);
tokio::time::sleep(std::time::Duration::from_millis(delay)).await;
}
let attempt_result = match ¤t_options.continuation {
ContinuationPolicy::SingleTurn => {
self.execute_once(
agent_id,
¤t_prompt,
_identity,
working_dir,
agent_name,
¤t_options,
)
.await
}
ContinuationPolicy::MultiTurn { .. } => {
self.execute_multi_turn(
agent_id,
¤t_prompt,
_identity,
working_dir,
agent_name,
max_retries,
retry_delay_ms,
¤t_options,
¤t_options.continuation.clone(),
)
.await
}
};
match attempt_result {
Ok(mut result) => {
result.fallbacks_used = fallbacks_used;
return Ok(result);
}
Err(e) => {
if is_rate_limit_error(&e)
&& fallback_index < current_options.rate_limit_fallbacks.len()
{
last_err = Some(e);
break 'attempts;
}
tracing::warn!("Provider CLI attempt {} failed: {}", attempt + 1, e);
last_err = Some(e);
}
}
}
let rate_limited = last_err.as_ref().is_some_and(is_rate_limit_error);
if rate_limited && fallback_index < current_options.rate_limit_fallbacks.len() {
let target = current_options.rate_limit_fallbacks[fallback_index].clone();
fallback_index += 1;
let from = apply_rate_limit_fallback(&mut current_options, target);
tracing::warn!(
"Rate limit on '{}' — falling back to '{}'",
from,
current_options
.provider
.unwrap_or(ProviderKind::Claude)
.as_str()
);
fallbacks_used.push(from);
current_prompt = fallback_notice_prompt(prompt);
continue;
}
return Err(last_err.unwrap_or_else(|| anyhow::anyhow!("All retry attempts failed")));
}
}
#[allow(clippy::too_many_arguments)]
pub async fn execute_multi_turn(
&self,
agent_id: &str,
prompt: &str,
_identity: &AgentIdentity,
working_dir: &Path,
agent_name: Option<&str>,
_max_retries: u32,
_retry_delay_ms: u64,
options: &MovementExecOptions,
continuation: &ContinuationPolicy,
) -> Result<BridgeResult> {
let (max_turns, continuation_prompt) = match continuation {
ContinuationPolicy::SingleTurn => {
return self
.execute_once(
agent_id,
prompt,
_identity,
working_dir,
agent_name,
options,
)
.await;
}
ContinuationPolicy::MultiTurn {
max_turns,
continuation_prompt,
} => ((*max_turns).max(1), continuation_prompt),
};
if resolve_a2a_endpoint(options).is_some() {
return Err(anyhow::anyhow!(
"A2A execution does not support same-thread multi-turn continuation"
));
}
let provider_kind = options.provider.unwrap_or(ProviderKind::Claude);
let provider = crate::providers::resolve(provider_kind);
let continuation_mode = provider.same_thread_continuation();
if !continuation_mode.supports_multi_turn() {
return Err(anyhow::anyhow!(
"{} provider does not support same-thread multi-turn continuation",
provider_kind.as_str()
));
}
let mut turn_options = options.clone();
let mut session_id = match continuation_mode {
crate::providers::SameThreadContinuation::ExplicitSessionId => {
let sid = options
.session_id
.clone()
.or_else(|| {
self.context_histories
.get(agent_id)
.map(|context| context.session_id.to_string())
})
.unwrap_or_else(|| ExecutionSessionId::new().to_string());
turn_options.session_id = Some(sid.clone());
Some(sid)
}
_ => options.session_id.clone(),
};
let first_execution = self
.execute_once_with_metadata(
agent_id,
prompt,
_identity,
working_dir,
agent_name,
&turn_options,
)
.await?;
match (&session_id, &first_execution.metadata.session_id) {
(Some(expected), _) => {
ensure_same_thread_continuation(&first_execution.metadata, expected)?;
}
(None, Some(assigned)) => {
session_id = Some(assigned.clone());
turn_options.session_id = Some(assigned.clone());
}
(None, None) => {
return Err(anyhow::anyhow!(
"{} did not report a session/thread ID on the first turn; cannot continue multi-turn",
provider_kind.as_str()
));
}
}
let session_id = session_id.unwrap_or_default();
let mut result = first_execution.result;
if !result.success {
return Ok(result);
}
let mut merged_raw = vec![format!("--- TURN 1 ---\n\n{}", result.raw)];
let mut duration_ms = result.duration_ms;
let mut tokens_in = result.tokens_in.unwrap_or(0);
let mut tokens_out = result.tokens_out.unwrap_or(0);
let mut tool_names = result.tool_names.clone();
let mut total_cost_usd = result.total_cost_usd;
let mut previous_turn_raw = result.raw.clone();
result.raw = merged_raw.join("\n\n");
result.duration_ms = duration_ms;
result.tokens_in = Some(tokens_in);
result.tokens_out = Some(tokens_out);
if result.success && is_task_terminal(&result.parsed) {
return Ok(result);
}
for turn in 2..=max_turns {
let continuation_prompt = format!(
"{}\n\n# Previous turn output (truncated)\n{}",
continuation_prompt,
truncate(&previous_turn_raw, 2048)
);
let next_execution = self
.execute_once_with_metadata(
agent_id,
&continuation_prompt,
_identity,
working_dir,
agent_name,
&turn_options,
)
.await?;
ensure_same_thread_continuation(&next_execution.metadata, &session_id)?;
let next_result = next_execution.result;
previous_turn_raw = next_result.raw.clone();
result = merge_turn_result(
&mut merged_raw,
&mut duration_ms,
&mut tokens_in,
&mut tokens_out,
&mut tool_names,
&mut total_cost_usd,
turn,
next_result,
);
if !result.success {
return Ok(result);
}
if result.success && is_task_terminal(&result.parsed) {
return Ok(result);
}
}
Ok(result)
}
async fn execute_once(
&self,
agent_id: &str,
prompt: &str,
_identity: &AgentIdentity,
working_dir: &Path,
agent_name: Option<&str>,
options: &MovementExecOptions,
) -> Result<BridgeResult> {
self.execute_once_with_metadata(
agent_id,
prompt,
_identity,
working_dir,
agent_name,
options,
)
.await
.map(|execution| execution.result)
}
async fn execute_once_with_metadata(
&self,
agent_id: &str,
prompt: &str,
_identity: &AgentIdentity,
working_dir: &Path,
agent_name: Option<&str>,
options: &MovementExecOptions,
) -> Result<BridgeExecution> {
let kind = options.provider.unwrap_or(ProviderKind::Claude);
let provider = crate::providers::resolve(kind);
let a2a_endpoint = resolve_a2a_endpoint(options);
let claude_stream_json = a2a_endpoint.is_none()
&& kind == ProviderKind::Claude
&& std::env::var("CCSWARM_CLAUDE_STREAM_JSON")
.map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
.unwrap_or(false);
let codex_json = a2a_endpoint.is_none()
&& kind == ProviderKind::Codex
&& (std::env::var("CCSWARM_CODEX_JSON")
.map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
.unwrap_or(false)
|| options.session_id.is_some()
|| matches!(options.continuation, ContinuationPolicy::MultiTurn { .. }));
let provider_options = ProviderOptions {
allowed_tools: options.tools.clone(),
model: options.model.clone(),
system_prompt: options.system_prompt.clone(),
agent_name: agent_name.map(String::from),
session_id: options.session_id.clone(),
continue_session: options.session_id.is_none()
&& self.context_histories.contains_key(agent_id),
max_budget: options.max_budget,
worktree_name: options.worktree_name.clone(),
claude_stream_json,
codex_json,
};
let prompt_with_cwd =
prepare_provider_prompt(prompt, working_dir, DEFAULT_MAX_PROMPT_BYTES)?;
let (raw_stdout, duration_ms) = if let Some(endpoint) = a2a_endpoint {
let a2a_prompt = prepare_a2a_prompt(&prompt_with_cwd, options.system_prompt.as_deref());
let execution = crate::session::a2a::A2AClient::new(endpoint)
.send_text(&a2a_prompt)
.await?;
(execution.output, execution.duration_ms)
} else {
let cmd = provider.build_command(&prompt_with_cwd, working_dir, &provider_options);
let output = run_provider_command(cmd, working_dir, provider.kind().as_str()).await?;
if !output.status.success() {
return Err(anyhow::anyhow!(
"{} provider CLI failed: {}",
provider.kind().as_str(),
output.stderr
));
}
(output.stdout, output.duration_ms)
};
let mut learned_session_id: Option<String> = None;
let stream_meta = if claude_stream_json {
Some(project_stream(&raw_stdout))
} else if codex_json {
let summary = crate::providers::codex_stream::parse_stream(&raw_stdout);
if let Some(message) = summary.failed {
return Err(anyhow::anyhow!("codex turn failed: {}", message));
}
learned_session_id = summary.thread_id;
Some(StreamProjection {
text: summary.result_text,
tool_names: summary.tool_names,
tokens: summary.tokens,
total_cost_usd: None, })
} else {
None
};
let raw_output = match &stream_meta {
Some(projection) => projection.text.clone(),
None => raw_stdout,
};
let parsed = self
.output_parser
.parse(&raw_output)
.unwrap_or(ParsedOutput::PlainText(raw_output.clone()));
let success = is_parsed_success(&parsed);
let attention = AttentionState::from_parsed(&parsed).unwrap_or(if success {
AttentionState::Done
} else {
AttentionState::Error
});
if let Some(mut context) = self.context_histories.get_mut(agent_id) {
context.add_message_raw(MessageRole::User, prompt.to_string());
context.add_message_raw(MessageRole::Assistant, raw_output.clone());
context.compress_context().await;
}
if let Some(context) = self.context_histories.get(agent_id) {
let session_id = context.session_id.clone();
let state = SessionState {
session_id: session_id.clone(),
status: ExecutionSessionStatus::Running,
context: context.clone(),
command_history: Vec::new(),
metadata: SessionMetadata::default(),
};
if let Err(e) = self.persistence.save_session(&session_id, &state).await {
tracing::warn!("Failed to persist session state for {}: {}", agent_id, e);
}
}
let compression_ratio = self
.get_compression_stats(agent_id)
.map(|stats| stats.compression_ratio);
let (tokens_in, tokens_out) = match stream_meta.as_ref().and_then(|p| p.tokens) {
Some((real_in, real_out)) => (Some(real_in), Some(real_out)),
None => {
let input_bytes =
prompt.len() + options.system_prompt.as_ref().map(|s| s.len()).unwrap_or(0);
(
Some((input_bytes / 4) as u64),
Some((raw_output.len() / 4) as u64),
)
}
};
let session_id_for_metadata = learned_session_id.or_else(|| options.session_id.clone());
let same_thread_continuation = match provider.same_thread_continuation() {
crate::providers::SameThreadContinuation::ExplicitSessionId => {
options.session_id.is_some()
}
crate::providers::SameThreadContinuation::ProviderAssignedId => {
session_id_for_metadata.is_some()
}
crate::providers::SameThreadContinuation::Unsupported => false,
};
Ok(BridgeExecution {
result: BridgeResult {
raw: raw_output,
parsed,
success,
duration_ms,
compression_ratio,
tokens_in,
tokens_out,
attention,
tool_names: stream_meta
.as_ref()
.map(|p| p.tool_names.clone())
.unwrap_or_default(),
total_cost_usd: stream_meta.as_ref().and_then(|p| p.total_cost_usd),
fallbacks_used: Vec::new(),
},
metadata: BridgeExecutionMetadata {
provider: kind,
session_id: session_id_for_metadata,
same_thread_continuation,
},
})
}
pub fn get_compression_stats(&self, agent_id: &str) -> Option<CompressionStats> {
self.context_histories
.get(agent_id)
.map(|ctx| ctx.get_compression_stats())
}
pub fn agent_count(&self) -> usize {
self.context_histories.len()
}
pub fn get_recent_context(&self, agent_id: &str, n: usize) -> Vec<String> {
self.context_histories
.get(agent_id)
.map(|ctx| {
ctx.get_recent_messages(n)
.into_iter()
.map(|m| format!("[{:?}] {}", m.role, m.content))
.collect()
})
.unwrap_or_default()
}
}
fn is_parsed_success(parsed: &ParsedOutput) -> bool {
match parsed {
ParsedOutput::PlainText(_) => true, ParsedOutput::CodeExecution { .. } => true,
ParsedOutput::BuildOutput { status, .. } => {
matches!(status, BuildStatus::Success)
}
ParsedOutput::TestResults { failed, .. } => *failed == 0,
ParsedOutput::StructuredLog { level, .. } => !matches!(level, LogLevel::Error),
}
}
fn is_task_terminal(parsed: &ParsedOutput) -> bool {
match parsed {
ParsedOutput::BuildOutput { status, .. } => {
matches!(status, BuildStatus::Success)
}
ParsedOutput::TestResults { failed, .. } => *failed == 0,
_ => false,
}
}
fn truncate(s: &str, max_bytes: usize) -> String {
if s.len() <= max_bytes {
return s.to_string();
}
let boundary = s
.char_indices()
.map(|(index, _)| index)
.take_while(|index| *index <= max_bytes)
.last()
.unwrap_or(0);
s[..boundary].to_string()
}
#[allow(clippy::too_many_arguments)]
fn merge_turn_result(
merged_raw: &mut Vec<String>,
duration_ms: &mut u64,
tokens_in: &mut u64,
tokens_out: &mut u64,
tool_names: &mut Vec<String>,
total_cost_usd: &mut Option<f64>,
turn: u32,
next_result: BridgeResult,
) -> BridgeResult {
merged_raw.push(format!("--- TURN {} ---\n\n{}", turn, next_result.raw));
*duration_ms = duration_ms.saturating_add(next_result.duration_ms);
*tokens_in = tokens_in.saturating_add(next_result.tokens_in.unwrap_or(0));
*tokens_out = tokens_out.saturating_add(next_result.tokens_out.unwrap_or(0));
tool_names.extend(next_result.tool_names.iter().cloned());
*total_cost_usd = match (*total_cost_usd, next_result.total_cost_usd) {
(Some(a), Some(b)) => Some(a + b),
(a, b) => a.or(b),
};
BridgeResult {
raw: merged_raw.join("\n\n"),
parsed: next_result.parsed,
success: next_result.success,
duration_ms: *duration_ms,
compression_ratio: next_result.compression_ratio,
tokens_in: Some(*tokens_in),
tokens_out: Some(*tokens_out),
attention: next_result.attention,
tool_names: tool_names.clone(),
total_cost_usd: *total_cost_usd,
fallbacks_used: next_result.fallbacks_used,
}
}
fn ensure_same_thread_continuation(
metadata: &BridgeExecutionMetadata,
expected_session_id: &str,
) -> Result<()> {
if metadata.same_thread_continuation
&& metadata.session_id.as_deref() == Some(expected_session_id)
{
return Ok(());
}
Err(anyhow::anyhow!(
"{} provider did not confirm same-thread continuation for session {}",
metadata.provider.as_str(),
expected_session_id
))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::identity::AgentRole;
use crate::session::output::{BuildStatus, ExecutionMetrics, TestDetails};
use std::collections::HashMap;
use std::path::PathBuf;
#[test]
fn test_bridge_creation() {
let bridge = A2ABridge::new(PathBuf::from("/tmp/ccswarm-test"));
assert_eq!(bridge.agent_count(), 0);
}
#[test]
fn test_agent_registration() {
let bridge = A2ABridge::new(PathBuf::from("/tmp/ccswarm-test"));
bridge.register_agent("frontend-agent").unwrap();
assert_eq!(bridge.agent_count(), 1);
}
#[test]
fn test_is_parsed_success() {
assert!(is_parsed_success(&ParsedOutput::PlainText(
"ok".to_string()
)));
assert!(!is_parsed_success(&ParsedOutput::TestResults {
passed: 5,
failed: 1,
details: TestDetails {
suite: Some("cargo".to_string()),
duration: Some(std::time::Duration::from_secs(0)),
failed_tests: vec!["test_one".to_string()],
},
}));
}
#[test]
fn test_continuation_policy_default_is_single_turn() {
assert!(matches!(
ContinuationPolicy::default(),
ContinuationPolicy::SingleTurn
));
}
#[test]
fn test_is_task_terminal_strict() {
assert!(is_task_terminal(&ParsedOutput::BuildOutput {
status: BuildStatus::Success,
artifacts: Vec::new(),
}));
assert!(!is_task_terminal(&ParsedOutput::BuildOutput {
status: BuildStatus::Failed("failed".to_string()),
artifacts: Vec::new(),
}));
assert!(is_task_terminal(&ParsedOutput::TestResults {
passed: 3,
failed: 0,
details: TestDetails::default(),
}));
assert!(!is_task_terminal(&ParsedOutput::TestResults {
passed: 3,
failed: 1,
details: TestDetails::default(),
}));
assert!(!is_task_terminal(&ParsedOutput::PlainText(
"done".to_string()
)));
assert!(!is_task_terminal(&ParsedOutput::CodeExecution {
result: "done".to_string(),
metrics: ExecutionMetrics {
execution_time: std::time::Duration::from_millis(1),
memory_usage: None,
cpu_usage: None,
},
}));
}
#[test]
fn test_truncate_is_utf8_safe() {
let input = "abあcd";
assert_eq!(truncate(input, 4), "ab");
assert_eq!(truncate(input, 5), "abあ");
assert_eq!(truncate(input, 99), input);
}
#[test]
fn test_multi_turn_constructor_sets_default_prompt() {
match ContinuationPolicy::multi_turn(3) {
ContinuationPolicy::MultiTurn {
max_turns,
continuation_prompt,
} => {
assert_eq!(max_turns, 3);
assert_eq!(continuation_prompt, DEFAULT_CONTINUATION_PROMPT);
}
ContinuationPolicy::SingleTurn => {
panic!("multi_turn constructor returned SingleTurn");
}
}
}
fn bridge_result(
raw: &str,
parsed: ParsedOutput,
success: bool,
duration_ms: u64,
tokens_in: Option<u64>,
tokens_out: Option<u64>,
) -> BridgeResult {
BridgeResult {
raw: raw.to_string(),
parsed,
success,
duration_ms,
compression_ratio: Some(1.0),
tokens_in,
tokens_out,
attention: AttentionState::Idle,
tool_names: Vec::new(),
total_cost_usd: None,
fallbacks_used: Vec::new(),
}
}
#[test]
fn test_is_rate_limit_error_matches_common_phrasings() {
for msg in [
"claude provider CLI failed: API rate limit exceeded",
"codex provider CLI failed: HTTP 429",
"Too Many Requests",
"quota exhausted for this billing period",
"Error: model overloaded, retry later",
] {
assert!(
is_rate_limit_error(&anyhow::anyhow!("{msg}")),
"should match: {msg}"
);
}
for msg in [
"compile error in main.rs",
"provider CLI failed: connection refused",
] {
assert!(
!is_rate_limit_error(&anyhow::anyhow!("{msg}")),
"should NOT match: {msg}"
);
}
}
#[test]
fn provider_switch_drops_provider_assigned_id() {
let mut options = MovementExecOptions {
provider: Some(ProviderKind::Codex),
session_id: Some("thread-from-codex".to_string()),
continuation: ContinuationPolicy::multi_turn(3),
model: Some("gpt-5".to_string()),
..Default::default()
};
let from =
apply_rate_limit_fallback(&mut options, (ProviderKind::Claude, Some("opus".into())));
assert_eq!(from, "codex");
assert_eq!(options.provider, Some(ProviderKind::Claude));
assert_eq!(options.model.as_deref(), Some("opus"));
assert!(options.session_id.is_none());
assert!(matches!(
options.continuation,
ContinuationPolicy::SingleTurn
));
}
#[test]
fn fallback_without_model_keeps_current_model() {
let mut options = MovementExecOptions {
provider: Some(ProviderKind::Claude),
model: Some("sonnet".to_string()),
..Default::default()
};
apply_rate_limit_fallback(&mut options, (ProviderKind::Codex, None));
assert_eq!(options.provider, Some(ProviderKind::Codex));
assert_eq!(options.model.as_deref(), Some("sonnet"));
}
#[test]
fn fallback_notice_prepends_original_prompt() {
let notice = fallback_notice_prompt("original task");
assert!(notice.starts_with("# Provider fallback notice"));
assert!(notice.ends_with("original task"));
}
#[test]
fn test_project_stream_extracts_tool_names_cost_and_real_tokens() {
let stdout = r#"{"type":"assistant","message":{"content":[{"type":"tool_use","id":"t1","name":"Read","input":{"path":"/a.rs"}},{"type":"text","text":"hi"}],"usage":{"input_tokens":10,"output_tokens":2}}}
{"type":"result","subtype":"success","result":"Done.","total_cost_usd":0.01,"usage":{"input_tokens":100,"output_tokens":20,"cache_read_input_tokens":50}}
"#;
let p = project_stream(stdout);
assert_eq!(p.text, "Done.");
assert_eq!(p.tool_names, vec!["Read".to_string()]);
assert_eq!(p.tokens, Some((150, 20)));
assert!((p.total_cost_usd.unwrap_or(0.0) - 0.01).abs() < 1e-9);
}
#[test]
fn test_project_stream_without_usage_leaves_tokens_none() {
let p = project_stream(
r#"{"type":"assistant","message":{"content":[{"type":"text","text":"plain"}]}}"#,
);
assert_eq!(p.text, "plain");
assert!(p.tool_names.is_empty());
assert_eq!(p.tokens, None);
assert!(p.total_cost_usd.is_none());
}
#[test]
fn test_prepare_a2a_prompt_preserves_system_instructions() {
let prompt = prepare_a2a_prompt("# Task\nReview it", Some("Act as a reviewer."));
assert_eq!(
prompt,
"# System instructions\nAct as a reviewer.\n\n# Task\nReview it"
);
assert_eq!(
prepare_a2a_prompt("# Task\nReview it", Some(" ")),
"# Task\nReview it"
);
}
#[test]
fn test_merge_turn_result_aggregates_successful_turn() {
let mut merged_raw = vec!["--- TURN 1 ---\n\nfirst".to_string()];
let mut duration_ms = 10;
let mut tokens_in = 3;
let mut tokens_out = 5;
let mut tool_names = vec!["Read".to_string()];
let mut total_cost_usd = Some(0.01);
let mut next = bridge_result(
"second",
ParsedOutput::PlainText("continue".to_string()),
true,
20,
Some(7),
Some(11),
);
next.tool_names = vec!["Bash".to_string()];
next.total_cost_usd = Some(0.02);
let result = merge_turn_result(
&mut merged_raw,
&mut duration_ms,
&mut tokens_in,
&mut tokens_out,
&mut tool_names,
&mut total_cost_usd,
2,
next,
);
assert_eq!(
result.raw,
"--- TURN 1 ---\n\nfirst\n\n--- TURN 2 ---\n\nsecond"
);
assert_eq!(result.duration_ms, 30);
assert_eq!(result.tokens_in, Some(10));
assert_eq!(result.tokens_out, Some(16));
assert!(result.success);
assert_eq!(result.tool_names, vec!["Read", "Bash"]);
assert!((result.total_cost_usd.unwrap_or(0.0) - 0.03).abs() < 1e-9);
}
#[test]
fn test_merge_turn_result_includes_failed_later_turn() {
let mut merged_raw = vec!["--- TURN 1 ---\n\nfirst".to_string()];
let mut duration_ms = 10;
let mut tokens_in = 3;
let mut tokens_out = 5;
let mut tool_names = Vec::new();
let mut total_cost_usd = None;
let result = merge_turn_result(
&mut merged_raw,
&mut duration_ms,
&mut tokens_in,
&mut tokens_out,
&mut tool_names,
&mut total_cost_usd,
2,
bridge_result(
"failed",
ParsedOutput::TestResults {
passed: 1,
failed: 1,
details: TestDetails::default(),
},
false,
20,
None,
Some(11),
),
);
assert_eq!(
result.raw,
"--- TURN 1 ---\n\nfirst\n\n--- TURN 2 ---\n\nfailed"
);
assert_eq!(result.duration_ms, 30);
assert_eq!(result.tokens_in, Some(3));
assert_eq!(result.tokens_out, Some(16));
assert!(!result.success);
assert!(matches!(
result.parsed,
ParsedOutput::TestResults { failed: 1, .. }
));
}
#[tokio::test]
async fn test_multi_turn_rejects_provider_without_same_thread_continuation() -> Result<()> {
let dir = tempfile::tempdir()?;
let bridge = A2ABridge::new(dir.path().join("sessions"));
let identity = AgentIdentity {
agent_id: "agent-1".to_string(),
specialization: AgentRole::Search {
technologies: Vec::new(),
responsibilities: Vec::new(),
boundaries: Vec::new(),
},
workspace_path: dir.path().to_path_buf(),
env_vars: HashMap::new(),
session_id: "session-1".to_string(),
parent_process_id: "parent-1".to_string(),
initialized_at: chrono::Utc::now(),
};
let options = MovementExecOptions {
provider: Some(ProviderKind::Copilot),
continuation: ContinuationPolicy::multi_turn(2),
..MovementExecOptions::default()
};
let err = bridge
.execute_multi_turn(
"agent-1",
"do work",
&identity,
dir.path(),
None,
0,
0,
&options,
&options.continuation,
)
.await
.expect_err("copilot multi-turn should fail before spawning a subprocess");
assert!(
err.to_string()
.contains("does not support same-thread multi-turn continuation")
);
Ok(())
}
#[tokio::test]
async fn test_multi_turn_rejects_a2a_without_remote_thread_support() -> Result<()> {
let dir = tempfile::tempdir()?;
let bridge = A2ABridge::new(dir.path().join("sessions"));
let identity = AgentIdentity {
agent_id: "agent-1".to_string(),
specialization: AgentRole::Search {
technologies: Vec::new(),
responsibilities: Vec::new(),
boundaries: Vec::new(),
},
workspace_path: dir.path().to_path_buf(),
env_vars: HashMap::new(),
session_id: "session-1".to_string(),
parent_process_id: "parent-1".to_string(),
initialized_at: chrono::Utc::now(),
};
let options = MovementExecOptions {
continuation: ContinuationPolicy::multi_turn(2),
a2a_endpoint: Some("https://example.test/a2a".to_string()),
..MovementExecOptions::default()
};
let err = bridge
.execute_multi_turn(
"agent-1",
"do work",
&identity,
dir.path(),
None,
0,
0,
&options,
&options.continuation,
)
.await
.expect_err("A2A multi-turn should fail before making a request");
assert!(
err.to_string()
.contains("does not support same-thread multi-turn continuation")
);
Ok(())
}
}