1use hashbrown::HashMap;
17use serde::de::{self, Deserializer};
18use serde::{Deserialize, Serialize};
19use serde_json::Value;
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
23#[serde(rename_all = "snake_case")]
24#[derive(Default)]
25pub enum SessionState {
26 #[default]
28 Created,
29 Active,
31 AwaitingInput,
33 Completed,
35 Cancelled,
37 Failed,
39}
40
41#[derive(Debug, Clone, Serialize, Deserialize)]
43pub struct AcpSession {
44 session_id: String,
46
47 state: SessionState,
49
50 created_at: String,
52
53 #[serde(skip_serializing_if = "Option::is_none")]
55 last_activity_at: Option<String>,
56
57 #[serde(default, skip_serializing_if = "HashMap::is_empty")]
59 metadata: HashMap<String, Value>,
60
61 #[serde(default)]
63 turn_count: u32,
64}
65
66impl AcpSession {
67 pub(crate) fn new(session_id: impl Into<String>) -> Self {
69 Self {
70 session_id: session_id.into(),
71 state: SessionState::Created,
72 created_at: chrono::Utc::now().to_rfc3339(),
73 last_activity_at: None,
74 metadata: HashMap::new(),
75 turn_count: 0,
76 }
77 }
78
79 pub(crate) fn set_state(&mut self, state: SessionState) {
81 self.state = state;
82 self.last_activity_at = Some(chrono::Utc::now().to_rfc3339());
83 }
84
85 pub fn increment_turn(&mut self) {
87 self.turn_count += 1;
88 self.last_activity_at = Some(chrono::Utc::now().to_rfc3339());
89 }
90}
91
92#[derive(Debug, Clone, Serialize, Deserialize, Default)]
98pub struct SessionNewParams {
99 #[serde(default, skip_serializing_if = "HashMap::is_empty")]
101 metadata: HashMap<String, Value>,
102
103 #[serde(skip_serializing_if = "Option::is_none")]
105 workspace: Option<WorkspaceContext>,
106
107 #[serde(skip_serializing_if = "Option::is_none")]
109 model_preferences: Option<ModelPreferences>,
110}
111
112#[derive(Debug, Clone, Serialize, Deserialize)]
114pub struct SessionNewResult {
115 pub(crate) session_id: String,
117
118 #[serde(default)]
120 state: SessionState,
121}
122
123#[derive(Debug, Clone, Serialize, Deserialize)]
129pub struct SessionLoadParams {
130 pub(crate) session_id: String,
132}
133
134#[derive(Debug, Clone, Serialize, Deserialize)]
136pub struct SessionLoadResult {
137 pub(crate) session: AcpSession,
139
140 #[serde(default, skip_serializing_if = "Vec::is_empty")]
142 pub(crate) history: Vec<ConversationTurn>,
143}
144
145#[derive(Debug, Clone, Serialize, Deserialize)]
151pub struct SessionPromptParams {
152 pub(crate) session_id: String,
154
155 content: Vec<PromptContent>,
157
158 #[serde(default, skip_serializing_if = "HashMap::is_empty")]
160 metadata: HashMap<String, Value>,
161}
162
163#[derive(Debug, Clone, Serialize, Deserialize)]
165#[serde(tag = "type", rename_all = "snake_case")]
166pub enum PromptContent {
167 Text {
169 text: String,
171 },
172
173 Image {
175 data: String,
177 mime_type: String,
179 #[serde(default)]
181 is_url: bool,
182 },
183
184 Context {
186 path: String,
188 content: String,
190 #[serde(skip_serializing_if = "Option::is_none")]
192 language: Option<String>,
193 },
194}
195
196impl PromptContent {
197 fn text(text: impl Into<String>) -> Self {
199 Self::Text { text: text.into() }
200 }
201
202 pub fn context(path: impl Into<String>, content: impl Into<String>) -> Self {
204 Self::Context {
205 path: path.into(),
206 content: content.into(),
207 language: None,
208 }
209 }
210}
211
212#[derive(Debug, Clone, Serialize, Deserialize)]
214pub struct SessionPromptResult {
215 pub(crate) turn_id: String,
217
218 #[serde(skip_serializing_if = "Option::is_none")]
220 response: Option<String>,
221
222 #[serde(default, skip_serializing_if = "Vec::is_empty")]
224 tool_calls: Vec<ToolCallRecord>,
225
226 pub(crate) status: TurnStatus,
228}
229
230#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
232#[serde(rename_all = "snake_case")]
233pub enum TurnStatus {
234 Completed,
236 Cancelled,
238 Failed,
240 AwaitingInput,
242}
243
244#[derive(Debug, Clone, Serialize, Deserialize)]
250#[serde(rename_all = "camelCase")]
251pub struct RequestPermissionParams {
252 session_id: String,
254
255 tool_call: ToolCallRecord,
257
258 options: Vec<PermissionOption>,
260}
261
262#[derive(Debug, Clone, Serialize, Deserialize)]
264pub struct PermissionOption {
265 id: String,
267
268 label: String,
270
271 #[serde(skip_serializing_if = "Option::is_none")]
273 description: Option<String>,
274}
275
276#[derive(Debug, Clone, Serialize, Deserialize)]
278#[serde(tag = "outcome", rename_all = "snake_case")]
279pub enum RequestPermissionResult {
280 Selected {
282 option_id: String,
284 },
285 Cancelled,
287}
288
289#[derive(Debug, Clone, Serialize, Deserialize)]
295pub struct SessionCancelParams {
296 pub(crate) session_id: String,
298
299 #[serde(skip_serializing_if = "Option::is_none")]
301 pub(crate) turn_id: Option<String>,
302}
303
304#[derive(Debug, Clone, Serialize)]
310pub struct SessionUpdateNotification {
311 pub(crate) session_id: String,
313
314 pub(crate) turn_id: String,
316
317 #[serde(flatten)]
319 pub(crate) update: SessionUpdate,
320}
321
322impl<'de> Deserialize<'de> for SessionUpdateNotification {
323 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
324 where
325 D: Deserializer<'de>,
326 {
327 let wire = SessionUpdateNotificationWire::deserialize(deserializer)?;
328 let SessionUpdateNotificationWire {
329 session_id,
330 turn_id,
331 update_type,
332 delta,
333 tool_call,
334 tool_call_id,
335 result,
336 status,
337 code,
338 message,
339 request,
340 } = wire;
341 let update = match update_type.as_str() {
342 "message_delta" => SessionUpdate::MessageDelta {
343 delta: required_update_field(delta, "delta", &update_type).map_err(de::Error::custom)?,
344 },
345 "tool_call_start" => SessionUpdate::ToolCallStart {
346 tool_call: required_update_field(tool_call, "tool_call", &update_type).map_err(de::Error::custom)?,
347 },
348 "tool_call_end" => SessionUpdate::ToolCallEnd {
349 tool_call_id: required_update_field(tool_call_id, "tool_call_id", &update_type)
350 .map_err(de::Error::custom)?,
351 result: required_present_value(result, "result", &update_type).map_err(de::Error::custom)?,
352 },
353 "turn_complete" => SessionUpdate::TurnComplete {
354 status: required_update_field(status, "status", &update_type).map_err(de::Error::custom)?,
355 },
356 "error" => SessionUpdate::Error {
357 code: required_update_field(code, "code", &update_type).map_err(de::Error::custom)?,
358 message: required_update_field(message, "message", &update_type).map_err(de::Error::custom)?,
359 },
360 "server_request" => SessionUpdate::ServerRequest {
361 request: required_update_field(request, "request", &update_type).map_err(de::Error::custom)?,
362 },
363 _ => return Err(de::Error::unknown_variant(&update_type, SESSION_UPDATE_TYPES)),
364 };
365
366 Ok(Self { session_id, turn_id, update })
367 }
368}
369
370#[derive(Debug, Deserialize)]
373struct SessionUpdateNotificationWire {
374 session_id: String,
375 turn_id: String,
376 update_type: String,
377 #[serde(default)]
378 delta: Option<String>,
379 #[serde(default)]
380 tool_call: Option<ToolCallRecord>,
381 #[serde(default)]
382 tool_call_id: Option<String>,
383 #[serde(default)]
384 result: Present<Value>,
385 #[serde(default)]
386 status: Option<TurnStatus>,
387 #[serde(default)]
388 code: Option<String>,
389 #[serde(default)]
390 message: Option<String>,
391 #[serde(default)]
392 request: Option<ToolExecutionRequest>,
393}
394
395#[derive(Debug)]
398struct Present<T> {
399 value: Option<T>,
400 present: bool,
401}
402
403impl<T> Default for Present<T> {
404 fn default() -> Self {
405 Self { value: None, present: false }
406 }
407}
408
409impl<'de, T> Deserialize<'de> for Present<T>
410where
411 T: Deserialize<'de>,
412{
413 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
414 where
415 D: Deserializer<'de>,
416 {
417 Ok(Self {
418 value: Option::<T>::deserialize(deserializer)?,
419 present: true,
420 })
421 }
422}
423
424fn required_present_value(field: Present<Value>, field_name: &str, update_type: &str) -> Result<Value, String> {
425 if field.present {
426 Ok(field.value.unwrap_or(Value::Null))
427 } else {
428 Err(format!("ACP update {update_type:?} is missing {field_name:?}"))
429 }
430}
431
432fn required_update_field<T>(value: Option<T>, field: &str, update_type: &str) -> Result<T, String> {
433 value.ok_or_else(|| format!("ACP update {update_type:?} is missing {field:?}"))
434}
435
436const SESSION_UPDATE_TYPES: &[&str] = &[
437 "message_delta",
438 "tool_call_start",
439 "tool_call_end",
440 "turn_complete",
441 "error",
442 "server_request",
443];
444
445#[derive(Debug, Clone, Serialize, Deserialize)]
447#[serde(tag = "update_type", rename_all = "snake_case")]
448pub enum SessionUpdate {
449 MessageDelta {
451 delta: String,
453 },
454
455 ToolCallStart {
457 tool_call: ToolCallRecord,
459 },
460
461 ToolCallEnd {
463 tool_call_id: String,
465 result: Value,
467 },
468
469 TurnComplete {
471 status: TurnStatus,
473 },
474
475 Error {
477 code: String,
479 message: String,
481 },
482
483 ServerRequest {
488 request: ToolExecutionRequest,
490 },
491}
492
493#[derive(Debug, Clone, Serialize, Deserialize)]
499pub struct WorkspaceContext {
500 root_path: String,
502
503 #[serde(skip_serializing_if = "Option::is_none")]
505 name: Option<String>,
506
507 #[serde(default, skip_serializing_if = "Vec::is_empty")]
509 active_files: Vec<String>,
510}
511
512#[derive(Debug, Clone, Serialize, Deserialize)]
514pub struct ModelPreferences {
515 #[serde(skip_serializing_if = "Option::is_none")]
517 model_id: Option<String>,
518
519 #[serde(skip_serializing_if = "Option::is_none")]
521 temperature: Option<f32>,
522
523 #[serde(skip_serializing_if = "Option::is_none")]
525 max_tokens: Option<u32>,
526}
527
528#[derive(Debug, Clone, Serialize, Deserialize)]
530pub struct ToolCallRecord {
531 id: String,
533
534 name: String,
536
537 arguments: Value,
539
540 #[serde(skip_serializing_if = "Option::is_none")]
542 result: Option<Value>,
543
544 timestamp: String,
546}
547
548#[derive(Debug, Clone, Serialize, Deserialize)]
554pub struct ToolExecutionRequest {
555 request_id: String,
557 tool_call: ToolCallRecord,
559}
560
561#[derive(Debug, Clone, Serialize, Deserialize)]
563pub struct ToolExecutionResult {
564 pub(crate) request_id: String,
566 pub(crate) tool_call_id: String,
568 pub(crate) output: Value,
570 pub(crate) success: bool,
572 #[serde(skip_serializing_if = "Option::is_none")]
574 pub(crate) error: Option<String>,
575}
576
577#[derive(Debug, Clone, Serialize, Deserialize)]
579pub struct ServerRequestNotification {
580 pub(crate) session_id: String,
582 pub(crate) request: ToolExecutionRequest,
584}
585
586#[derive(Debug, Clone, Serialize, Deserialize)]
588pub struct ConversationTurn {
589 turn_id: String,
591
592 prompt: Vec<PromptContent>,
594
595 #[serde(skip_serializing_if = "Option::is_none")]
597 response: Option<String>,
598
599 #[serde(default, skip_serializing_if = "Vec::is_empty")]
601 tool_calls: Vec<ToolCallRecord>,
602
603 timestamp: String,
605}
606
607#[cfg(test)]
608mod tests {
609 use super::*;
610 use serde_json::json;
611
612 #[test]
613 fn test_session_new_params() {
614 let params = SessionNewParams::default();
615 let json = serde_json::to_value(¶ms).unwrap();
616 assert_eq!(json, json!({}));
617 }
618
619 #[test]
620 fn test_prompt_content_text() {
621 let content = PromptContent::text("Hello, world!");
622 let json = serde_json::to_value(&content).unwrap();
623 assert_eq!(json["type"], "text");
624 assert_eq!(json["text"], "Hello, world!");
625 }
626
627 #[test]
628 fn test_session_update_message_delta() {
629 let update = SessionUpdate::MessageDelta { delta: "Hello".to_string() };
630 let json = serde_json::to_value(&update).unwrap();
631 assert_eq!(json["update_type"], "message_delta");
632 assert_eq!(json["delta"], "Hello");
633 }
634
635 #[test]
636 fn session_update_notification_deserializes_each_update_shape() {
637 let tool_call = json!({
638 "id": "tc-1",
639 "name": "code_search",
640 "arguments": {"query": "fn main"},
641 "timestamp": "2025-01-01T00:00:00Z"
642 });
643 let request = json!({
644 "request_id": "req-1",
645 "tool_call": tool_call.clone()
646 });
647 let cases = [
648 (json!({"session_id":"s","turn_id":"t","update_type":"message_delta","delta":"hi"}), "message"),
649 (
650 json!({"session_id":"s","turn_id":"t","update_type":"tool_call_start","tool_call":tool_call}),
651 "start",
652 ),
653 (
654 json!({"session_id":"s","turn_id":"t","update_type":"tool_call_end","tool_call_id":"tc-1","result":null}),
655 "end",
656 ),
657 (
658 json!({"session_id":"s","turn_id":"t","update_type":"turn_complete","status":"completed"}),
659 "complete",
660 ),
661 (
662 json!({"session_id":"s","turn_id":"t","update_type":"error","code":"bad_request","message":"nope"}),
663 "error",
664 ),
665 (json!({"session_id":"s","turn_id":"t","update_type":"server_request","request":request}), "request"),
666 ];
667
668 for (payload, expected) in cases {
669 let notification: SessionUpdateNotification =
670 serde_json::from_value(payload).expect("valid session update notification");
671 let actual = match notification.update {
672 SessionUpdate::MessageDelta { .. } => "message",
673 SessionUpdate::ToolCallStart { .. } => "start",
674 SessionUpdate::ToolCallEnd { result, .. } if result.is_null() => "end",
675 SessionUpdate::TurnComplete { .. } => "complete",
676 SessionUpdate::Error { .. } => "error",
677 SessionUpdate::ServerRequest { .. } => "request",
678 _ => "other",
679 };
680 assert_eq!(actual, expected);
681 }
682 }
683
684 #[test]
685 fn session_update_notification_rejects_missing_payload() {
686 let missing_delta = json!({
687 "session_id": "s",
688 "turn_id": "t",
689 "update_type": "message_delta"
690 });
691 assert!(serde_json::from_value::<SessionUpdateNotification>(missing_delta).is_err());
692
693 let unknown = json!({
694 "session_id": "s",
695 "turn_id": "t",
696 "update_type": "future_update"
697 });
698 assert!(serde_json::from_value::<SessionUpdateNotification>(unknown).is_err());
699 }
700
701 #[test]
702 fn test_session_state_transitions() {
703 let mut session = AcpSession::new("test-session");
704 assert_eq!(session.state, SessionState::Created);
705
706 session.set_state(SessionState::Active);
707 assert_eq!(session.state, SessionState::Active);
708 assert!(session.last_activity_at.is_some());
709 }
710
711 #[test]
712 fn server_request_update_serializes_correctly() {
713 let tool_call = ToolCallRecord {
714 id: "tc-1".to_string(),
715 name: "code_search".to_string(),
716 arguments: json!({"query": "fn main"}),
717 result: None,
718 timestamp: "2025-01-01T00:00:00Z".to_string(),
719 };
720 let request = ToolExecutionRequest { request_id: "req-1".to_string(), tool_call };
721 let update = SessionUpdate::ServerRequest { request };
722 let json = serde_json::to_value(&update).unwrap();
723 assert_eq!(json["update_type"], "server_request");
724 assert_eq!(json["request"]["request_id"], "req-1");
725 }
726
727 #[test]
728 fn tool_execution_result_success_serializes() {
729 let result = ToolExecutionResult {
730 request_id: "req-1".to_string(),
731 tool_call_id: "tc-1".to_string(),
732 output: json!({"matches": []}),
733 success: true,
734 error: None,
735 };
736 let json = serde_json::to_value(&result).unwrap();
737 assert_eq!(json["success"], true);
738 assert!(json.get("error").is_none());
739 }
740
741 #[test]
742 fn tool_execution_result_failure_includes_error() {
743 let result = ToolExecutionResult {
744 request_id: "req-1".to_string(),
745 tool_call_id: "tc-1".to_string(),
746 output: Value::Null,
747 success: false,
748 error: Some("permission denied".to_string()),
749 };
750 let json = serde_json::to_value(&result).unwrap();
751 assert_eq!(json["success"], false);
752 assert_eq!(json["error"], "permission denied");
753 }
754}