pub use crate::cm_sse_protocol::{SSE_PROTOCOL_VERSION, StreamEndReason};
pub const SSE_RESUME_RING_CAP: usize = 512;
fn default_sse_v() -> u8 {
SSE_PROTOCOL_VERSION
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct SseMessage {
#[serde(default = "default_sse_v")]
pub v: u8,
#[serde(flatten)]
pub payload: SsePayload,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
#[serde(untagged)]
#[allow(clippy::large_enum_variant)] pub enum SsePayload {
Error(SseErrorBody),
CommandApproval {
command_approval_request: CommandApprovalBody,
},
ClarificationQuestionnaire {
#[serde(rename = "clarification_questionnaire")]
clarification_questionnaire: ClarificationQuestionnaireBody,
},
ToolCall {
tool_call: ToolCallSummary,
},
ToolResult {
tool_result: ToolResultBody,
},
WorkspaceChanged {
workspace_changed: bool,
},
ToolRunning {
tool_running: bool,
},
ToolOutputChunk {
#[serde(rename = "tool_output_chunk")]
tool_output_chunk: ToolOutputChunkBody,
},
ParsingToolCalls {
parsing_tool_calls: bool,
},
AssistantAnswerPhase {
#[serde(rename = "assistant_answer_phase")]
assistant_answer_phase: bool,
},
TurnSegmentStart {
#[serde(rename = "turn_segment_start")]
start: TurnSegmentStartBody,
},
TurnSegmentEnd {
#[serde(rename = "turn_segment_end")]
end: TurnSegmentEndBody,
},
TurnToolPhaseEnd {
#[serde(rename = "turn_tool_phase_end")]
turn_tool_phase_end: bool,
},
PlanRequired {
plan_required: bool,
},
ChatUiSeparator {
#[serde(rename = "chat_ui_separator")]
short: bool,
},
ConversationSaved {
#[serde(rename = "conversation_saved")]
saved: ConversationSavedBody,
},
TimelineLog {
#[serde(rename = "timeline_log")]
log: TimelineLogBody,
},
ThinkingTrace {
#[serde(rename = "thinking_trace")]
trace: ThinkingTraceBody,
},
SseCapabilities {
#[serde(rename = "sse_capabilities")]
caps: SseCapabilitiesBody,
},
StreamDraining {
#[serde(rename = "stream_draining")]
draining: StreamDrainingBody,
},
StreamEnded {
#[serde(rename = "stream_ended")]
ended: StreamEndedBody,
},
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct SseErrorBody {
pub error: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub code: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reason_code: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub turn_id: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub sub_phase: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub request_id: Option<String>,
}
impl SseErrorBody {
#[must_use]
pub fn with_request_id(mut self, request_id: Option<String>) -> Self {
self.request_id = request_id.filter(|s| !s.trim().is_empty());
self
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct CommandApprovalBody {
pub command: String,
pub args: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allowlist_key: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ClarificationQuestionField {
pub id: String,
pub label: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub hint: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub required: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub kind: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ClarificationQuestionnaireBody {
pub questionnaire_id: String,
pub intro: String,
pub questions: Vec<ClarificationQuestionField>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ToolCallSummary {
pub name: String,
pub summary: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub goal_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub arguments_preview: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub arguments: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ToolOutputChunkBody {
pub tool_call_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
pub seq: u64,
#[serde(default)]
pub chunk: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stream: Option<String>,
}
const CRABMATE_TOOL_ENVELOPE_VERSION_V1: u32 = 1;
fn default_tool_result_payload_version() -> u32 {
CRABMATE_TOOL_ENVELOPE_VERSION_V1
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ToolResultBody {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub goal_id: Option<String>,
#[serde(default = "default_tool_result_payload_version")]
pub result_version: u32,
#[serde(skip_serializing_if = "Option::is_none")]
pub summary: Option<String>,
pub output: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub ok: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub exit_code: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub error_code: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub failure_category: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub retryable: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub execution_mode: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parallel_batch_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stdout: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stderr: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub structured_preview: Option<serde_json::Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_job_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_job_poll_url: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_job_status: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct TurnSegmentStartBody {
pub segment_id: String,
pub kind: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub before_tool_call_id: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct TurnSegmentEndBody {
pub segment_id: String,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ConversationSavedBody {
pub revision: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tiktoken_prompt_tokens: Option<crate::cm_types::TiktokenPromptTokensSnapshot>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct SseCapabilitiesBody {
pub supported_sse_v: u8,
pub resume_ring_cap: usize,
pub job_id: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub terminal_order: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct StreamDrainingBody {
pub job_id: u64,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct StreamEndedBody {
pub job_id: u64,
pub reason: StreamEndReason,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tiktoken_prompt_tokens: Option<crate::cm_types::TiktokenPromptTokensSnapshot>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct TimelineLogBody {
pub kind: String,
pub title: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub detail: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ThinkingTraceBody {
pub op: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub node_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub parent_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub title: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub chunk: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub context_snapshot: Option<String>,
}
pub fn encode_message(payload: SsePayload) -> String {
let encoder = super::encoder::default_encoder();
encoder.encode(&payload)
}
#[cfg(test)]
pub(crate) fn encode_message_v1(payload: &SsePayload) -> String {
serde_json::to_string(&SseMessage {
v: SSE_PROTOCOL_VERSION,
payload: payload.clone(),
})
.unwrap_or_else(|e| {
log::error!(
target: "crabmate",
"sse_protocol encode failed error={}",
e
);
format!(
r#"{{"v":{},"error":"内部协议序列化失败","code":"SSE_ENCODE"}}"#,
SSE_PROTOCOL_VERSION
)
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::cm_sse_protocol::StreamEndReason;
use proptest::prelude::*;
#[test]
fn roundtrip_clarification_questionnaire() {
let s = encode_message_v1(&SsePayload::ClarificationQuestionnaire {
clarification_questionnaire: ClarificationQuestionnaireBody {
questionnaire_id: "q1".into(),
intro: "请补充".into(),
questions: vec![ClarificationQuestionField {
id: "scope".into(),
label: "范围?".into(),
hint: Some("可选".into()),
required: Some(true),
kind: Some("text".into()),
}],
},
});
let m: SseMessage = serde_json::from_str(&s).unwrap();
assert!(matches!(
m.payload,
SsePayload::ClarificationQuestionnaire { .. }
));
}
#[test]
fn roundtrip_parsing_tool_calls() {
let s = encode_message_v1(&SsePayload::ParsingToolCalls {
parsing_tool_calls: true,
});
let m: SseMessage = serde_json::from_str(&s).unwrap();
assert!(matches!(
m.payload,
SsePayload::ParsingToolCalls {
parsing_tool_calls: true
}
));
}
#[test]
fn roundtrip_assistant_answer_phase() {
let s = encode_message_v1(&SsePayload::AssistantAnswerPhase {
assistant_answer_phase: true,
});
assert!(s.contains("\"assistant_answer_phase\":true"));
let m: SseMessage = serde_json::from_str(&s).unwrap();
assert!(matches!(
m.payload,
SsePayload::AssistantAnswerPhase {
assistant_answer_phase: true
}
));
}
#[test]
fn roundtrip_tool_running() {
let s = encode_message_v1(&SsePayload::ToolRunning { tool_running: true });
let m: SseMessage = serde_json::from_str(&s).unwrap();
assert_eq!(m.v, SSE_PROTOCOL_VERSION);
assert!(matches!(
m.payload,
SsePayload::ToolRunning { tool_running: true }
));
}
#[test]
fn roundtrip_tool_output_chunk() {
let s = encode_message_v1(&SsePayload::ToolOutputChunk {
tool_output_chunk: ToolOutputChunkBody {
tool_call_id: "tc1".into(),
name: Some("terminal_session".into()),
seq: 3,
chunk: "hello\n".into(),
stream: Some("combined".into()),
},
});
assert!(s.contains("\"tool_output_chunk\""));
let m: SseMessage = serde_json::from_str(&s).unwrap();
match m.payload {
SsePayload::ToolOutputChunk { tool_output_chunk } => {
assert_eq!(tool_output_chunk.tool_call_id, "tc1");
assert_eq!(tool_output_chunk.seq, 3);
assert_eq!(tool_output_chunk.chunk, "hello\n");
assert_eq!(tool_output_chunk.stream.as_deref(), Some("combined"));
}
_ => panic!("expected tool_output_chunk payload"),
}
}
#[test]
fn deserialize_legacy_no_v_field() {
let m: SseMessage = serde_json::from_str(r#"{"tool_running":false}"#).unwrap();
assert_eq!(m.v, SSE_PROTOCOL_VERSION);
assert!(matches!(
m.payload,
SsePayload::ToolRunning {
tool_running: false
}
));
}
#[test]
fn error_with_code() {
let s = encode_message_v1(&SsePayload::Error(SseErrorBody {
error: "x".into(),
code: Some("E".into()),
reason_code: None,
turn_id: None,
sub_phase: None,
request_id: None,
}));
assert!(s.contains(&format!("\"v\":{}", SSE_PROTOCOL_VERSION)));
assert!(s.contains("\"code\":\"E\""));
}
#[test]
fn error_with_reason_code() {
let s = encode_message_v1(&SsePayload::Error(SseErrorBody {
error: "x".into(),
code: Some("plan_rewrite_exhausted".into()),
reason_code: Some("plan_missing".into()),
turn_id: None,
sub_phase: Some("reflect".into()),
request_id: None,
}));
assert!(s.contains("\"reason_code\":\"plan_missing\""));
}
#[test]
fn tool_result_with_structured_fields() {
let s = encode_message_v1(&SsePayload::ToolResult {
tool_result: ToolResultBody {
name: "run_command".into(),
goal_id: None,
result_version: 1,
summary: Some("ls".into()),
output: "退出码:1".into(),
ok: Some(false),
exit_code: Some(1),
error_code: Some("command_failed".into()),
failure_category: Some("external".into()),
retryable: Some(false),
tool_call_id: Some("tc1".into()),
execution_mode: Some("serial".into()),
parallel_batch_id: None,
stdout: Some(String::new()),
stderr: Some("permission denied".into()),
structured_preview: None,
tool_job_id: None,
tool_job_poll_url: None,
tool_job_status: None,
},
});
let m: SseMessage = serde_json::from_str(&s).unwrap();
match m.payload {
SsePayload::ToolResult { tool_result } => {
assert_eq!(tool_result.name, "run_command");
assert_eq!(tool_result.summary.as_deref(), Some("ls"));
assert_eq!(tool_result.ok, Some(false));
assert_eq!(tool_result.exit_code, Some(1));
assert_eq!(tool_result.error_code.as_deref(), Some("command_failed"));
assert_eq!(tool_result.failure_category.as_deref(), Some("external"));
assert_eq!(tool_result.retryable, Some(false));
assert_eq!(tool_result.tool_call_id.as_deref(), Some("tc1"));
assert_eq!(tool_result.execution_mode.as_deref(), Some("serial"));
assert_eq!(tool_result.stderr.as_deref(), Some("permission denied"));
assert!(tool_result.structured_preview.is_none());
}
_ => panic!("expected tool_result payload"),
}
}
#[test]
fn roundtrip_chat_ui_separator() {
let s = encode_message_v1(&SsePayload::ChatUiSeparator { short: true });
assert!(s.contains("\"chat_ui_separator\":true"));
let m: SseMessage = serde_json::from_str(&s).unwrap();
assert!(matches!(
m.payload,
SsePayload::ChatUiSeparator { short: true }
));
let s2 = encode_message_v1(&SsePayload::ChatUiSeparator { short: false });
let m2: SseMessage = serde_json::from_str(&s2).unwrap();
assert!(matches!(
m2.payload,
SsePayload::ChatUiSeparator { short: false }
));
}
fn arb_short_text() -> impl Strategy<Value = String> {
"[a-zA-Z0-9_\\- ]{0,32}".prop_map(|s| s.trim().to_string())
}
proptest! {
#[test]
fn prop_tool_running_roundtrip_and_version(tool_running in any::<bool>()) {
let encoded = encode_message_v1(&SsePayload::ToolRunning { tool_running });
let parsed: SseMessage = serde_json::from_str(&encoded).unwrap();
prop_assert_eq!(parsed.v, SSE_PROTOCOL_VERSION);
match parsed.payload {
SsePayload::ToolRunning { tool_running: got } => prop_assert_eq!(got, tool_running),
other => prop_assert!(false, "unexpected payload: {:?}", other),
}
}
#[test]
fn prop_error_payload_roundtrip(
error in arb_short_text(),
code in proptest::option::of(arb_short_text()),
reason_code in proptest::option::of(arb_short_text()),
turn_id in proptest::option::of(any::<u64>()),
sub_phase in proptest::option::of(arb_short_text()),
) {
let payload = SsePayload::Error(SseErrorBody {
error,
code,
reason_code,
turn_id,
sub_phase,
request_id: None,
});
let encoded = encode_message_v1(&payload);
let parsed: SseMessage = serde_json::from_str(&encoded).unwrap();
prop_assert_eq!(parsed.v, SSE_PROTOCOL_VERSION);
match (payload, parsed.payload) {
(SsePayload::Error(expect), SsePayload::Error(got)) => {
prop_assert_eq!(got.error, expect.error);
prop_assert_eq!(got.code, expect.code);
prop_assert_eq!(got.reason_code, expect.reason_code);
prop_assert_eq!(got.turn_id, expect.turn_id);
prop_assert_eq!(got.sub_phase, expect.sub_phase);
prop_assert_eq!(got.request_id, expect.request_id);
}
(_, other) => prop_assert!(false, "unexpected payload: {:?}", other),
}
}
#[test]
fn prop_stream_ended_reason_is_parsable(job_id in any::<u64>(), reason_idx in 0usize..6usize) {
let reason = match reason_idx {
0 => StreamEndReason::Completed,
1 => StreamEndReason::Cancelled,
2 => StreamEndReason::Conflict,
3 => StreamEndReason::Fallback,
4 => StreamEndReason::NoOutput,
_ => StreamEndReason::Gone,
};
let encoded = encode_message_v1(&SsePayload::StreamEnded {
ended: StreamEndedBody {
job_id,
reason,
tiktoken_prompt_tokens: None,
},
});
let parsed: SseMessage = serde_json::from_str(&encoded).unwrap();
match parsed.payload {
SsePayload::StreamEnded { ended } => {
prop_assert_eq!(ended.job_id, job_id);
prop_assert_eq!(ended.reason, reason);
}
other => prop_assert!(false, "unexpected payload: {:?}", other),
}
}
}
}