use hashbrown::HashMap;
use serde::de::{self, Deserializer};
use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[derive(Default)]
pub enum SessionState {
#[default]
Created,
Active,
AwaitingInput,
Completed,
Cancelled,
Failed,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AcpSession {
session_id: String,
state: SessionState,
created_at: String,
#[serde(skip_serializing_if = "Option::is_none")]
last_activity_at: Option<String>,
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
metadata: HashMap<String, Value>,
#[serde(default)]
turn_count: u32,
}
impl AcpSession {
pub(crate) fn new(session_id: impl Into<String>) -> Self {
Self {
session_id: session_id.into(),
state: SessionState::Created,
created_at: chrono::Utc::now().to_rfc3339(),
last_activity_at: None,
metadata: HashMap::new(),
turn_count: 0,
}
}
pub(crate) fn set_state(&mut self, state: SessionState) {
self.state = state;
self.last_activity_at = Some(chrono::Utc::now().to_rfc3339());
}
pub fn increment_turn(&mut self) {
self.turn_count += 1;
self.last_activity_at = Some(chrono::Utc::now().to_rfc3339());
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct SessionNewParams {
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
metadata: HashMap<String, Value>,
#[serde(skip_serializing_if = "Option::is_none")]
workspace: Option<WorkspaceContext>,
#[serde(skip_serializing_if = "Option::is_none")]
model_preferences: Option<ModelPreferences>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionNewResult {
pub(crate) session_id: String,
#[serde(default)]
state: SessionState,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionLoadParams {
pub(crate) session_id: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionLoadResult {
pub(crate) session: AcpSession,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub(crate) history: Vec<ConversationTurn>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionPromptParams {
pub(crate) session_id: String,
content: Vec<PromptContent>,
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
metadata: HashMap<String, Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum PromptContent {
Text {
text: String,
},
Image {
data: String,
mime_type: String,
#[serde(default)]
is_url: bool,
},
Context {
path: String,
content: String,
#[serde(skip_serializing_if = "Option::is_none")]
language: Option<String>,
},
}
impl PromptContent {
fn text(text: impl Into<String>) -> Self {
Self::Text { text: text.into() }
}
pub fn context(path: impl Into<String>, content: impl Into<String>) -> Self {
Self::Context {
path: path.into(),
content: content.into(),
language: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionPromptResult {
pub(crate) turn_id: String,
#[serde(skip_serializing_if = "Option::is_none")]
response: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
tool_calls: Vec<ToolCallRecord>,
pub(crate) status: TurnStatus,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum TurnStatus {
Completed,
Cancelled,
Failed,
AwaitingInput,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct RequestPermissionParams {
session_id: String,
tool_call: ToolCallRecord,
options: Vec<PermissionOption>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PermissionOption {
id: String,
label: String,
#[serde(skip_serializing_if = "Option::is_none")]
description: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "outcome", rename_all = "snake_case")]
pub enum RequestPermissionResult {
Selected {
option_id: String,
},
Cancelled,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionCancelParams {
pub(crate) session_id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) turn_id: Option<String>,
}
#[derive(Debug, Clone, Serialize)]
pub struct SessionUpdateNotification {
pub(crate) session_id: String,
pub(crate) turn_id: String,
#[serde(flatten)]
pub(crate) update: SessionUpdate,
}
impl<'de> Deserialize<'de> for SessionUpdateNotification {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let wire = SessionUpdateNotificationWire::deserialize(deserializer)?;
let SessionUpdateNotificationWire {
session_id,
turn_id,
update_type,
delta,
tool_call,
tool_call_id,
result,
status,
code,
message,
request,
} = wire;
let update = match update_type.as_str() {
"message_delta" => SessionUpdate::MessageDelta {
delta: required_update_field(delta, "delta", &update_type).map_err(de::Error::custom)?,
},
"tool_call_start" => SessionUpdate::ToolCallStart {
tool_call: required_update_field(tool_call, "tool_call", &update_type).map_err(de::Error::custom)?,
},
"tool_call_end" => SessionUpdate::ToolCallEnd {
tool_call_id: required_update_field(tool_call_id, "tool_call_id", &update_type)
.map_err(de::Error::custom)?,
result: required_present_value(result, "result", &update_type).map_err(de::Error::custom)?,
},
"turn_complete" => SessionUpdate::TurnComplete {
status: required_update_field(status, "status", &update_type).map_err(de::Error::custom)?,
},
"error" => SessionUpdate::Error {
code: required_update_field(code, "code", &update_type).map_err(de::Error::custom)?,
message: required_update_field(message, "message", &update_type).map_err(de::Error::custom)?,
},
"server_request" => SessionUpdate::ServerRequest {
request: required_update_field(request, "request", &update_type).map_err(de::Error::custom)?,
},
_ => return Err(de::Error::unknown_variant(&update_type, SESSION_UPDATE_TYPES)),
};
Ok(Self { session_id, turn_id, update })
}
}
#[derive(Debug, Deserialize)]
struct SessionUpdateNotificationWire {
session_id: String,
turn_id: String,
update_type: String,
#[serde(default)]
delta: Option<String>,
#[serde(default)]
tool_call: Option<ToolCallRecord>,
#[serde(default)]
tool_call_id: Option<String>,
#[serde(default)]
result: Present<Value>,
#[serde(default)]
status: Option<TurnStatus>,
#[serde(default)]
code: Option<String>,
#[serde(default)]
message: Option<String>,
#[serde(default)]
request: Option<ToolExecutionRequest>,
}
#[derive(Debug)]
struct Present<T> {
value: Option<T>,
present: bool,
}
impl<T> Default for Present<T> {
fn default() -> Self {
Self { value: None, present: false }
}
}
impl<'de, T> Deserialize<'de> for Present<T>
where
T: Deserialize<'de>,
{
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
Ok(Self {
value: Option::<T>::deserialize(deserializer)?,
present: true,
})
}
}
fn required_present_value(field: Present<Value>, field_name: &str, update_type: &str) -> Result<Value, String> {
if field.present {
Ok(field.value.unwrap_or(Value::Null))
} else {
Err(format!("ACP update {update_type:?} is missing {field_name:?}"))
}
}
fn required_update_field<T>(value: Option<T>, field: &str, update_type: &str) -> Result<T, String> {
value.ok_or_else(|| format!("ACP update {update_type:?} is missing {field:?}"))
}
const SESSION_UPDATE_TYPES: &[&str] = &[
"message_delta",
"tool_call_start",
"tool_call_end",
"turn_complete",
"error",
"server_request",
];
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "update_type", rename_all = "snake_case")]
pub enum SessionUpdate {
MessageDelta {
delta: String,
},
ToolCallStart {
tool_call: ToolCallRecord,
},
ToolCallEnd {
tool_call_id: String,
result: Value,
},
TurnComplete {
status: TurnStatus,
},
Error {
code: String,
message: String,
},
ServerRequest {
request: ToolExecutionRequest,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkspaceContext {
root_path: String,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
active_files: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelPreferences {
#[serde(skip_serializing_if = "Option::is_none")]
model_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
max_tokens: Option<u32>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolCallRecord {
id: String,
name: String,
arguments: Value,
#[serde(skip_serializing_if = "Option::is_none")]
result: Option<Value>,
timestamp: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolExecutionRequest {
request_id: String,
tool_call: ToolCallRecord,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolExecutionResult {
pub(crate) request_id: String,
pub(crate) tool_call_id: String,
pub(crate) output: Value,
pub(crate) success: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) error: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ServerRequestNotification {
pub(crate) session_id: String,
pub(crate) request: ToolExecutionRequest,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConversationTurn {
turn_id: String,
prompt: Vec<PromptContent>,
#[serde(skip_serializing_if = "Option::is_none")]
response: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
tool_calls: Vec<ToolCallRecord>,
timestamp: String,
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_session_new_params() {
let params = SessionNewParams::default();
let json = serde_json::to_value(¶ms).unwrap();
assert_eq!(json, json!({}));
}
#[test]
fn test_prompt_content_text() {
let content = PromptContent::text("Hello, world!");
let json = serde_json::to_value(&content).unwrap();
assert_eq!(json["type"], "text");
assert_eq!(json["text"], "Hello, world!");
}
#[test]
fn test_session_update_message_delta() {
let update = SessionUpdate::MessageDelta { delta: "Hello".to_string() };
let json = serde_json::to_value(&update).unwrap();
assert_eq!(json["update_type"], "message_delta");
assert_eq!(json["delta"], "Hello");
}
#[test]
fn session_update_notification_deserializes_each_update_shape() {
let tool_call = json!({
"id": "tc-1",
"name": "code_search",
"arguments": {"query": "fn main"},
"timestamp": "2025-01-01T00:00:00Z"
});
let request = json!({
"request_id": "req-1",
"tool_call": tool_call.clone()
});
let cases = [
(json!({"session_id":"s","turn_id":"t","update_type":"message_delta","delta":"hi"}), "message"),
(
json!({"session_id":"s","turn_id":"t","update_type":"tool_call_start","tool_call":tool_call}),
"start",
),
(
json!({"session_id":"s","turn_id":"t","update_type":"tool_call_end","tool_call_id":"tc-1","result":null}),
"end",
),
(
json!({"session_id":"s","turn_id":"t","update_type":"turn_complete","status":"completed"}),
"complete",
),
(
json!({"session_id":"s","turn_id":"t","update_type":"error","code":"bad_request","message":"nope"}),
"error",
),
(json!({"session_id":"s","turn_id":"t","update_type":"server_request","request":request}), "request"),
];
for (payload, expected) in cases {
let notification: SessionUpdateNotification =
serde_json::from_value(payload).expect("valid session update notification");
let actual = match notification.update {
SessionUpdate::MessageDelta { .. } => "message",
SessionUpdate::ToolCallStart { .. } => "start",
SessionUpdate::ToolCallEnd { result, .. } if result.is_null() => "end",
SessionUpdate::TurnComplete { .. } => "complete",
SessionUpdate::Error { .. } => "error",
SessionUpdate::ServerRequest { .. } => "request",
_ => "other",
};
assert_eq!(actual, expected);
}
}
#[test]
fn session_update_notification_rejects_missing_payload() {
let missing_delta = json!({
"session_id": "s",
"turn_id": "t",
"update_type": "message_delta"
});
assert!(serde_json::from_value::<SessionUpdateNotification>(missing_delta).is_err());
let unknown = json!({
"session_id": "s",
"turn_id": "t",
"update_type": "future_update"
});
assert!(serde_json::from_value::<SessionUpdateNotification>(unknown).is_err());
}
#[test]
fn test_session_state_transitions() {
let mut session = AcpSession::new("test-session");
assert_eq!(session.state, SessionState::Created);
session.set_state(SessionState::Active);
assert_eq!(session.state, SessionState::Active);
assert!(session.last_activity_at.is_some());
}
#[test]
fn server_request_update_serializes_correctly() {
let tool_call = ToolCallRecord {
id: "tc-1".to_string(),
name: "code_search".to_string(),
arguments: json!({"query": "fn main"}),
result: None,
timestamp: "2025-01-01T00:00:00Z".to_string(),
};
let request = ToolExecutionRequest { request_id: "req-1".to_string(), tool_call };
let update = SessionUpdate::ServerRequest { request };
let json = serde_json::to_value(&update).unwrap();
assert_eq!(json["update_type"], "server_request");
assert_eq!(json["request"]["request_id"], "req-1");
}
#[test]
fn tool_execution_result_success_serializes() {
let result = ToolExecutionResult {
request_id: "req-1".to_string(),
tool_call_id: "tc-1".to_string(),
output: json!({"matches": []}),
success: true,
error: None,
};
let json = serde_json::to_value(&result).unwrap();
assert_eq!(json["success"], true);
assert!(json.get("error").is_none());
}
#[test]
fn tool_execution_result_failure_includes_error() {
let result = ToolExecutionResult {
request_id: "req-1".to_string(),
tool_call_id: "tc-1".to_string(),
output: Value::Null,
success: false,
error: Some("permission denied".to_string()),
};
let json = serde_json::to_value(&result).unwrap();
assert_eq!(json["success"], false);
assert_eq!(json["error"], "permission denied");
}
}