Skip to main content

vtcode_acp/
session.rs

1//! Legacy crate-internal HTTP session types (`message_delta`, `tool_call_start`, ...).
2//!
3//! Not the ACP stdio `session/update` surface. Do not extend; canonical session
4//! lifecycle lives in `zed/agent/session_state.rs`.
5//!
6//! ACP session types and lifecycle management
7//!
8//! This module implements the session lifecycle as defined by ACP:
9//! - Session creation (session/new)
10//! - Session loading (session/load)
11//! - Prompt handling (session/prompt)
12//! - Session updates (session/update notifications)
13//!
14//! Reference: <https://agentclientprotocol.com/llms.txt>
15
16use hashbrown::HashMap;
17use serde::de::{self, Deserializer};
18use serde::{Deserialize, Serialize};
19use serde_json::Value;
20
21/// Session state enumeration
22#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
23#[serde(rename_all = "snake_case")]
24#[derive(Default)]
25pub enum SessionState {
26    /// Session created but not yet active
27    #[default]
28    Created,
29    /// Session is active and processing
30    Active,
31    /// Session is waiting for user input
32    AwaitingInput,
33    /// Session completed successfully
34    Completed,
35    /// Session was cancelled
36    Cancelled,
37    /// Session failed with error
38    Failed,
39}
40
41/// ACP Session representation
42#[derive(Debug, Clone, Serialize, Deserialize)]
43pub struct AcpSession {
44    /// Unique session identifier
45    session_id: String,
46
47    /// Current session state
48    state: SessionState,
49
50    /// Session creation timestamp (ISO 8601)
51    created_at: String,
52
53    /// Last activity timestamp (ISO 8601)
54    #[serde(skip_serializing_if = "Option::is_none")]
55    last_activity_at: Option<String>,
56
57    /// Session metadata
58    #[serde(default, skip_serializing_if = "HashMap::is_empty")]
59    metadata: HashMap<String, Value>,
60
61    /// Turn counter for prompt/response cycles
62    #[serde(default)]
63    turn_count: u32,
64}
65
66impl AcpSession {
67    /// Create a new session with the given ID
68    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    /// Update session state
80    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    /// Increment turn counter
86    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// ============================================================================
93// Session/New Request/Response
94// ============================================================================
95
96/// Parameters for session/new method
97#[derive(Debug, Clone, Serialize, Deserialize, Default)]
98pub struct SessionNewParams {
99    /// Optional session metadata
100    #[serde(default, skip_serializing_if = "HashMap::is_empty")]
101    metadata: HashMap<String, Value>,
102
103    /// Optional workspace context
104    #[serde(skip_serializing_if = "Option::is_none")]
105    workspace: Option<WorkspaceContext>,
106
107    /// Optional model preferences
108    #[serde(skip_serializing_if = "Option::is_none")]
109    model_preferences: Option<ModelPreferences>,
110}
111
112/// Result of session/new method
113#[derive(Debug, Clone, Serialize, Deserialize)]
114pub struct SessionNewResult {
115    /// The created session ID
116    pub(crate) session_id: String,
117
118    /// Initial session state
119    #[serde(default)]
120    state: SessionState,
121}
122
123// ============================================================================
124// Session/Load Request/Response
125// ============================================================================
126
127/// Parameters for session/load method
128#[derive(Debug, Clone, Serialize, Deserialize)]
129pub struct SessionLoadParams {
130    /// Session ID to load
131    pub(crate) session_id: String,
132}
133
134/// Result of session/load method
135#[derive(Debug, Clone, Serialize, Deserialize)]
136pub struct SessionLoadResult {
137    /// The loaded session
138    pub(crate) session: AcpSession,
139
140    /// Conversation history (if available)
141    #[serde(default, skip_serializing_if = "Vec::is_empty")]
142    pub(crate) history: Vec<ConversationTurn>,
143}
144
145// ============================================================================
146// Session/Prompt Request/Response
147// ============================================================================
148
149/// Parameters for session/prompt method
150#[derive(Debug, Clone, Serialize, Deserialize)]
151pub struct SessionPromptParams {
152    /// Session ID
153    pub(crate) session_id: String,
154
155    /// Prompt content (can be text, images, etc.)
156    content: Vec<PromptContent>,
157
158    /// Optional turn-specific metadata
159    #[serde(default, skip_serializing_if = "HashMap::is_empty")]
160    metadata: HashMap<String, Value>,
161}
162
163/// Prompt content types
164#[derive(Debug, Clone, Serialize, Deserialize)]
165#[serde(tag = "type", rename_all = "snake_case")]
166pub enum PromptContent {
167    /// Plain text content
168    Text {
169        /// The text content
170        text: String,
171    },
172
173    /// Image content (base64 or URL)
174    Image {
175        /// Image data (base64) or URL
176        data: String,
177        /// MIME type (e.g., "image/png")
178        mime_type: String,
179        /// Whether data is a URL (false = base64)
180        #[serde(default)]
181        is_url: bool,
182    },
183
184    /// Embedded context (file contents, etc.)
185    Context {
186        /// Context identifier/path
187        path: String,
188        /// Context content
189        content: String,
190        /// Language hint for syntax highlighting
191        #[serde(skip_serializing_if = "Option::is_none")]
192        language: Option<String>,
193    },
194}
195
196impl PromptContent {
197    /// Create text content
198    fn text(text: impl Into<String>) -> Self {
199        Self::Text { text: text.into() }
200    }
201
202    /// Create context content
203    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/// Result of session/prompt method
213#[derive(Debug, Clone, Serialize, Deserialize)]
214pub struct SessionPromptResult {
215    /// Turn ID for this prompt/response cycle
216    pub(crate) turn_id: String,
217
218    /// Final response content (may be streamed via notifications first)
219    #[serde(skip_serializing_if = "Option::is_none")]
220    response: Option<String>,
221
222    /// Tool calls made during this turn
223    #[serde(default, skip_serializing_if = "Vec::is_empty")]
224    tool_calls: Vec<ToolCallRecord>,
225
226    /// Turn completion status
227    pub(crate) status: TurnStatus,
228}
229
230/// Turn completion status
231#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
232#[serde(rename_all = "snake_case")]
233pub enum TurnStatus {
234    /// Turn completed successfully
235    Completed,
236    /// Turn was cancelled
237    Cancelled,
238    /// Turn failed with error
239    Failed,
240    /// Turn requires user input (e.g., permission approval)
241    AwaitingInput,
242}
243
244// ============================================================================
245// Session/RequestPermission (Client Method)
246// ============================================================================
247
248/// Parameters for session/request_permission method (client callable by agent)
249#[derive(Debug, Clone, Serialize, Deserialize)]
250#[serde(rename_all = "camelCase")]
251pub struct RequestPermissionParams {
252    /// Session ID
253    session_id: String,
254
255    /// Tool call requiring permission
256    tool_call: ToolCallRecord,
257
258    /// Available permission options
259    options: Vec<PermissionOption>,
260}
261
262/// A permission option presented to the user
263#[derive(Debug, Clone, Serialize, Deserialize)]
264pub struct PermissionOption {
265    /// Option ID
266    id: String,
267
268    /// Display label
269    label: String,
270
271    /// Detailed description
272    #[serde(skip_serializing_if = "Option::is_none")]
273    description: Option<String>,
274}
275
276/// Result of session/request_permission
277#[derive(Debug, Clone, Serialize, Deserialize)]
278#[serde(tag = "outcome", rename_all = "snake_case")]
279pub enum RequestPermissionResult {
280    /// User selected an option
281    Selected {
282        /// The selected option ID
283        option_id: String,
284    },
285    /// User cancelled the request
286    Cancelled,
287}
288
289// ============================================================================
290// Session/Cancel Request
291// ============================================================================
292
293/// Parameters for session/cancel method
294#[derive(Debug, Clone, Serialize, Deserialize)]
295pub struct SessionCancelParams {
296    /// Session ID
297    pub(crate) session_id: String,
298
299    /// Optional turn ID to cancel (if not provided, cancels current turn)
300    #[serde(skip_serializing_if = "Option::is_none")]
301    pub(crate) turn_id: Option<String>,
302}
303
304// ============================================================================
305// Session/Update Notification (Streaming)
306// ============================================================================
307
308/// Session update notification payload
309#[derive(Debug, Clone, Serialize)]
310pub struct SessionUpdateNotification {
311    /// Session ID
312    pub(crate) session_id: String,
313
314    /// Turn ID this update belongs to
315    pub(crate) turn_id: String,
316
317    /// Update type
318    #[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/// Direct wire representation used to avoid buffering the notification map
371/// for the flattened, tagged [`SessionUpdate`] payload.
372#[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/// Preserves the distinction between a missing field and an explicit JSON
396/// `null` for required values such as a tool result.
397#[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/// Session update types
446#[derive(Debug, Clone, Serialize, Deserialize)]
447#[serde(tag = "update_type", rename_all = "snake_case")]
448pub enum SessionUpdate {
449    /// Text delta (streaming response)
450    MessageDelta {
451        /// Incremental text content
452        delta: String,
453    },
454
455    /// Tool call started
456    ToolCallStart {
457        /// Tool call details
458        tool_call: ToolCallRecord,
459    },
460
461    /// Tool call completed
462    ToolCallEnd {
463        /// Tool call ID
464        tool_call_id: String,
465        /// Tool result
466        result: Value,
467    },
468
469    /// Turn completed
470    TurnComplete {
471        /// Final status
472        status: TurnStatus,
473    },
474
475    /// Error occurred
476    Error {
477        /// Error code
478        code: String,
479        /// Error message
480        message: String,
481    },
482
483    /// Server requests the client to execute a tool (bidirectional ACP protocol).
484    ///
485    /// Arrives via a `server/request` SSE event. After executing the tool,
486    /// send the result back with `AcpClientV2::session_tool_response`.
487    ServerRequest {
488        /// The tool execution request from the agent.
489        request: ToolExecutionRequest,
490    },
491}
492
493// ============================================================================
494// Supporting Types
495// ============================================================================
496
497/// Workspace context for session initialization
498#[derive(Debug, Clone, Serialize, Deserialize)]
499pub struct WorkspaceContext {
500    /// Workspace root path
501    root_path: String,
502
503    /// Workspace name
504    #[serde(skip_serializing_if = "Option::is_none")]
505    name: Option<String>,
506
507    /// Active file paths
508    #[serde(default, skip_serializing_if = "Vec::is_empty")]
509    active_files: Vec<String>,
510}
511
512/// Model preferences for session
513#[derive(Debug, Clone, Serialize, Deserialize)]
514pub struct ModelPreferences {
515    /// Preferred model ID
516    #[serde(skip_serializing_if = "Option::is_none")]
517    model_id: Option<String>,
518
519    /// Temperature setting
520    #[serde(skip_serializing_if = "Option::is_none")]
521    temperature: Option<f32>,
522
523    /// Max tokens
524    #[serde(skip_serializing_if = "Option::is_none")]
525    max_tokens: Option<u32>,
526}
527
528/// Record of a tool call
529#[derive(Debug, Clone, Serialize, Deserialize)]
530pub struct ToolCallRecord {
531    /// Unique tool call ID
532    id: String,
533
534    /// Tool name
535    name: String,
536
537    /// Tool arguments
538    arguments: Value,
539
540    /// Tool result (if completed)
541    #[serde(skip_serializing_if = "Option::is_none")]
542    result: Option<Value>,
543
544    /// Timestamp
545    timestamp: String,
546}
547
548/// Server-to-client tool execution request (arrives via `server/request` SSE event).
549///
550/// When the ACP agent needs the client to run a tool on its behalf it emits a
551/// `server/request` SSE event containing this payload. The client must execute
552/// the tool and reply with [`ToolExecutionResult`] via `client/response`.
553#[derive(Debug, Clone, Serialize, Deserialize)]
554pub struct ToolExecutionRequest {
555    /// Unique request identifier used to correlate the response.
556    request_id: String,
557    /// The tool call the agent wants the client to execute.
558    tool_call: ToolCallRecord,
559}
560
561/// Result of a client-side tool execution, sent back via `client/response`.
562#[derive(Debug, Clone, Serialize, Deserialize)]
563pub struct ToolExecutionResult {
564    /// Must match the `request_id` from [`ToolExecutionRequest`].
565    pub(crate) request_id: String,
566    /// ID of the tool call that was executed.
567    pub(crate) tool_call_id: String,
568    /// Tool output (structured or text).
569    pub(crate) output: Value,
570    /// Whether execution succeeded.
571    pub(crate) success: bool,
572    /// Error message when `success` is false.
573    #[serde(skip_serializing_if = "Option::is_none")]
574    pub(crate) error: Option<String>,
575}
576
577/// Wrapper for a `server/request` SSE event notification.
578#[derive(Debug, Clone, Serialize, Deserialize)]
579pub struct ServerRequestNotification {
580    /// Session this request belongs to.
581    pub(crate) session_id: String,
582    /// The tool execution request.
583    pub(crate) request: ToolExecutionRequest,
584}
585
586/// A single turn in the conversation
587#[derive(Debug, Clone, Serialize, Deserialize)]
588pub struct ConversationTurn {
589    /// Turn ID
590    turn_id: String,
591
592    /// User prompt
593    prompt: Vec<PromptContent>,
594
595    /// Agent response
596    #[serde(skip_serializing_if = "Option::is_none")]
597    response: Option<String>,
598
599    /// Tool calls made during this turn
600    #[serde(default, skip_serializing_if = "Vec::is_empty")]
601    tool_calls: Vec<ToolCallRecord>,
602
603    /// Turn timestamp
604    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(&params).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}