Skip to main content

vtcode_acp/zed/agent/
lifecycle.rs

1//! VT Code ACP extension methods for session lifecycle management.
2//!
3//! The Agent Client Protocol covers `session/new`, `session/load`,
4//! `session/prompt`, and `session/cancel`. Programmatic hosts (IDE bridges,
5//! test harnesses, automation) additionally need the lifecycle operations
6//! Codex's app-server exposes over its proprietary JSON-RPC dialect: fork,
7//! rollback, and compact. This module implements them as **VT Code ACP
8//! extensions** on the same connection, so one protocol surface serves
9//! everything instead of maintaining a parallel native app-server (the
10//! `vtcode app-server` Codex proxy keeps serving `provider = "codex"`).
11//!
12//! Wire contract (also emitted by `vtcode schema acp`):
13//!
14//! | Method             | Params                            | Result                                   |
15//! |--------------------|-----------------------------------|------------------------------------------|
16//! | `session/fork`     | `{session_id}`                    | `{session_id}` (new independent session) |
17//! | `session/rollback` | `{session_id, keep_last_turns}`   | `{remaining_messages}`                   |
18//! | `session/compact`  | `{session_id}`                    | `{original_messages, compacted_messages}`|
19//!
20//! ACP clients that don't know these methods are unaffected: they simply
21//! never call them.
22
23use std::sync::Arc;
24
25use crate::acp::Error as SdkError;
26use agent_client_protocol::schema::v1::SessionId;
27use agent_client_protocol::{
28    Agent, Builder, Client, HandleDispatchFrom, JsonRpcMessage, JsonRpcRequest, JsonRpcResponse, Responder,
29    RunWithConnectionTo, UntypedMessage, on_receive_request,
30};
31use schemars::JsonSchema;
32use serde::{Deserialize, Serialize};
33use serde_json::json;
34use vtcode_core::compaction::{CompactionConfig, compact_history};
35use vtcode_core::core::threads::ThreadBootstrap;
36use vtcode_core::llm::provider::{Message, MessageRole};
37
38use super::super::constants::SESSION_PREFIX;
39use super::ZedAgent;
40use super::handlers::build_session_provider;
41
42// ---------------------------------------------------------------------------
43// Wire types
44// ---------------------------------------------------------------------------
45
46/// `session/fork` — branch an existing session into an independent one that
47/// inherits the full conversation history and configuration.
48#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
49pub struct SessionForkRequest {
50    /// Session to branch from.
51    pub session_id: SessionId,
52}
53
54/// Result of `session/fork`.
55#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
56pub struct SessionForkResponse {
57    /// The new session's id; it can be driven with `session/prompt` like any
58    /// other session.
59    pub session_id: SessionId,
60}
61
62/// `session/rollback` — drop the trailing conversation turns of a session.
63#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
64pub struct SessionRollbackRequest {
65    /// Session to roll back.
66    pub session_id: SessionId,
67    /// Number of most-recent user turns to keep. `0` clears the conversation.
68    pub keep_last_turns: u32,
69}
70
71/// Result of `session/rollback`.
72#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
73pub struct SessionRollbackResponse {
74    /// Messages left in the session after the rollback.
75    pub remaining_messages: usize,
76}
77
78/// `session/compact` — summarize the session history in place to reclaim
79/// context budget, preserving task continuity.
80#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
81pub struct SessionCompactRequest {
82    /// Session to compact.
83    pub session_id: SessionId,
84}
85
86/// Result of `session/compact`.
87#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
88pub struct SessionCompactResponse {
89    /// Message count before compaction.
90    pub original_messages: usize,
91    /// Message count after compaction.
92    pub compacted_messages: usize,
93}
94
95macro_rules! impl_lifecycle_request {
96    ($req:ty, $res:ty, $method:literal) => {
97        impl JsonRpcMessage for $req {
98            fn matches_method(method: &str) -> bool {
99                method == $method
100            }
101
102            fn method(&self) -> &'static str {
103                $method
104            }
105
106            fn to_untyped_message(&self) -> Result<UntypedMessage, agent_client_protocol::Error> {
107                UntypedMessage::new($method, self)
108            }
109
110            fn parse_message(method: &str, params: &impl Serialize) -> Result<Self, agent_client_protocol::Error> {
111                if !Self::matches_method(method) {
112                    return Err(agent_client_protocol::Error::method_not_found());
113                }
114                let params = serde_json::to_value(params)
115                    .map_err(|error| agent_client_protocol::Error::invalid_params().data(error.to_string()))?;
116                serde_json::from_value(params)
117                    .map_err(|error| agent_client_protocol::Error::invalid_params().data(error.to_string()))
118            }
119        }
120
121        impl JsonRpcRequest for $req {
122            type Response = $res;
123        }
124    };
125}
126
127impl_lifecycle_request!(SessionForkRequest, SessionForkResponse, "session/fork");
128impl_lifecycle_request!(SessionRollbackRequest, SessionRollbackResponse, "session/rollback");
129impl_lifecycle_request!(SessionCompactRequest, SessionCompactResponse, "session/compact");
130
131macro_rules! impl_lifecycle_response {
132    ($res:ty) => {
133        impl JsonRpcResponse for $res {
134            fn into_json(self, _method: &str) -> Result<serde_json::Value, agent_client_protocol::Error> {
135                serde_json::to_value(self)
136                    .map_err(|error| agent_client_protocol::Error::internal_error().data(error.to_string()))
137            }
138
139            fn from_value(_method: &str, value: serde_json::Value) -> Result<Self, agent_client_protocol::Error> {
140                serde_json::from_value(value)
141                    .map_err(|error| agent_client_protocol::Error::invalid_params().data(error.to_string()))
142            }
143        }
144    };
145}
146
147impl_lifecycle_response!(SessionForkResponse);
148impl_lifecycle_response!(SessionRollbackResponse);
149impl_lifecycle_response!(SessionCompactResponse);
150
151// ---------------------------------------------------------------------------
152// Handlers
153// ---------------------------------------------------------------------------
154
155impl ZedAgent {
156    /// Allocate the next session id without registering a thread.
157    fn allocate_session_id(&self) -> SessionId {
158        let raw_id = self.next_session_id.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
159        SessionId::new(Arc::from(format!("{SESSION_PREFIX}-{raw_id}")))
160    }
161
162    /// Branch `parent_id` into a new independent session with copied history.
163    pub(crate) async fn fork_session(&self, parent_id: &SessionId) -> Result<SessionId, SdkError> {
164        let parent = self
165            .session_handle(parent_id)
166            .ok_or_else(|| SdkError::invalid_params().data(json!({ "reason": "unknown_session" })))?;
167        let snapshot = {
168            let Ok(data) = parent.data.lock() else {
169                return Err(SdkError::internal_error());
170            };
171            data.thread.snapshot()
172        };
173
174        let fork_id = self.allocate_session_id();
175        let thread = self
176            .thread_manager
177            .start_thread_with_identifier(fork_id.0.to_string(), ThreadBootstrap::from_snapshot(snapshot));
178        let handle = self.build_session_handle(fork_id.clone(), thread);
179        if let Ok(mut guard) = self.sessions.lock() {
180            drop(guard.insert(fork_id.clone(), handle));
181        }
182        Ok(fork_id)
183    }
184
185    /// Drop trailing user turns from a session, keeping `keep_last_turns`.
186    pub(crate) fn rollback_session(&self, session_id: &SessionId, keep_last_turns: u32) -> Result<usize, SdkError> {
187        let session = self
188            .session_handle(session_id)
189            .ok_or_else(|| SdkError::invalid_params().data(json!({ "reason": "unknown_session" })))?;
190        let Ok(data) = session.data.lock() else {
191            return Err(SdkError::internal_error());
192        };
193        let messages = data.thread.messages();
194        let boundary = turn_boundary_index(&messages, keep_last_turns);
195        let kept = messages.get(boundary..).unwrap_or_default();
196        data.thread.replace_messages(kept.to_vec());
197        Ok(kept.len())
198    }
199
200    /// Summarize a session's history in place via the configured model.
201    pub(crate) async fn compact_session(&self, session_id: &SessionId) -> Result<(usize, usize), SdkError> {
202        let session = self
203            .session_handle(session_id)
204            .ok_or_else(|| SdkError::invalid_params().data(json!({ "reason": "unknown_session" })))?;
205        let (provider_name, model, history) = {
206            let Ok(data) = session.data.lock() else {
207                return Err(SdkError::internal_error());
208            };
209            (data.provider.clone(), data.model.clone(), data.thread.messages())
210        };
211        // The data lock is released before the LLM call (no lock across await).
212        let provider = build_session_provider(self, &provider_name, &model)?;
213        let original = history.len();
214        let compacted = compact_history(provider.as_ref(), &model, &history, &CompactionConfig::default())
215            .await
216            .map_err(|error| SdkError::internal_error().data(error.to_string()))?;
217        let compacted_len = compacted.len();
218        let Ok(data) = session.data.lock() else {
219            return Err(SdkError::internal_error());
220        };
221        // The lock was released during the LLM call, so a concurrent
222        // `session/prompt` or `session/rollback` may have mutated the history
223        // the summary was built from. Replacing it wholesale would silently
224        // drop those turns — fail closed instead of compacting stale state.
225        ensure_history_unchanged(&data.thread.messages(), &history)?;
226        data.thread.replace_messages(compacted);
227        Ok((original, compacted_len))
228    }
229}
230
231/// Verify `current` still equals the `snapshot` a long-running operation was
232/// built from, so the operation can safely replace the history in place.
233/// Returns an error naming the concurrent mutation when they diverge.
234fn ensure_history_unchanged(current: &[Message], snapshot: &[Message]) -> Result<(), SdkError> {
235    if current == snapshot {
236        return Ok(());
237    }
238    Err(SdkError::internal_error().data(json!({
239        "reason": "session_changed_during_operation",
240        "snapshot_messages": snapshot.len(),
241        "current_messages": current.len(),
242    })))
243}
244
245/// Index of the first message of the trailing `keep_last_turns` user turns.
246///
247/// A turn starts at its user message; everything before the kept window is
248/// rolled back. `keep_last_turns == 0` removes the whole conversation, and
249/// histories with fewer user turns than requested are left untouched.
250fn turn_boundary_index(messages: &[Message], keep_last_turns: u32) -> usize {
251    let user_positions: Vec<usize> = messages
252        .iter()
253        .enumerate()
254        .filter(|(_, message)| message.role == MessageRole::User)
255        .map(|(index, _)| index)
256        .collect();
257    if keep_last_turns == 0 {
258        return messages.len();
259    }
260    match user_positions.len().checked_sub(keep_last_turns as usize) {
261        Some(offset) => user_positions.get(offset).copied().unwrap_or(0),
262        // Fewer user turns than requested: nothing to roll back.
263        None => 0,
264    }
265}
266
267// ---------------------------------------------------------------------------
268// Schema export
269// ---------------------------------------------------------------------------
270
271/// Method descriptions for the lifecycle extensions (shared by the schema
272/// document builder and the crate docs).
273pub const LIFECYCLE_METHODS: &[(&str, &str)] = &[
274    (
275        "session/fork",
276        "Branch an existing session into a new independent session with copied history; returns {session_id}.",
277    ),
278    (
279        "session/rollback",
280        "Drop the trailing conversation turns of a session, keeping the requested number of most-recent user turns; returns {remaining_messages}.",
281    ),
282    (
283        "session/compact",
284        "Summarize a session's history in place via the configured model to reclaim context budget; returns {original_messages, compacted_messages}.",
285    ),
286];
287
288/// JSON-Schema document describing the VT Code lifecycle extension methods —
289/// the wire contract emitted by `vtcode schema acp`.
290#[must_use]
291pub fn lifecycle_schema_document() -> serde_json::Value {
292    let methods = [
293        ("session/fork", schema_for_type::<SessionForkRequest>(), schema_for_type::<SessionForkResponse>()),
294        (
295            "session/rollback",
296            schema_for_type::<SessionRollbackRequest>(),
297            schema_for_type::<SessionRollbackResponse>(),
298        ),
299        (
300            "session/compact",
301            schema_for_type::<SessionCompactRequest>(),
302            schema_for_type::<SessionCompactResponse>(),
303        ),
304    ];
305    let methods = methods
306        .into_iter()
307        .map(|(method, params, result)| {
308            let description = LIFECYCLE_METHODS
309                .iter()
310                .find(|(name, _)| *name == method)
311                .map(|(_, description)| (*description).to_string())
312                .unwrap_or_default();
313            json!({
314                "method": method,
315                "description": description,
316                "params": params,
317                "result": result,
318            })
319        })
320        .collect::<Vec<_>>();
321
322    json!({
323        "version": env!("CARGO_PKG_VERSION"),
324        "kind": "vtcode-acp-lifecycle-extensions",
325        "note": "VT Code ACP extension methods served on the standard ACP connection; standard ACP clients are unaffected.",
326        "methods": methods,
327    })
328}
329
330fn schema_for_type<T: JsonSchema>() -> serde_json::Value {
331    serde_json::to_value(schemars::schema_for!(T)).unwrap_or(serde_json::Value::Null)
332}
333
334// ---------------------------------------------------------------------------
335// Registration
336// ---------------------------------------------------------------------------
337
338/// Register the VT Code lifecycle extension handlers onto the SACP builder.
339pub fn install_lifecycle_handlers<H, R>(
340    builder: Builder<Agent, H, R>,
341    agent: Arc<ZedAgent>,
342) -> Builder<Agent, impl HandleDispatchFrom<Client>, R>
343where
344    H: HandleDispatchFrom<Client>,
345    R: RunWithConnectionTo<Client>,
346{
347    builder
348        .on_receive_request(
349            {
350                let agent = Arc::clone(&agent);
351                move |req: SessionForkRequest, request_cx: Responder<SessionForkResponse>, _cx| {
352                    let agent = Arc::clone(&agent);
353                    async move {
354                        request_cx.respond_with_result(
355                            agent
356                                .fork_session(&req.session_id)
357                                .await
358                                .map(|session_id| SessionForkResponse { session_id }),
359                        )
360                    }
361                }
362            },
363            on_receive_request!(),
364        )
365        .on_receive_request(
366            {
367                let agent = Arc::clone(&agent);
368                move |req: SessionRollbackRequest, request_cx: Responder<SessionRollbackResponse>, _cx| {
369                    let agent = Arc::clone(&agent);
370                    async move {
371                        request_cx.respond_with_result(
372                            agent
373                                .rollback_session(&req.session_id, req.keep_last_turns)
374                                .map(|remaining_messages| SessionRollbackResponse { remaining_messages }),
375                        )
376                    }
377                }
378            },
379            on_receive_request!(),
380        )
381        .on_receive_request(
382            {
383                let agent = Arc::clone(&agent);
384                move |req: SessionCompactRequest, request_cx: Responder<SessionCompactResponse>, _cx| {
385                    let agent = Arc::clone(&agent);
386                    async move {
387                        request_cx.respond_with_result(agent.compact_session(&req.session_id).await.map(
388                            |(original_messages, compacted_messages)| SessionCompactResponse {
389                                original_messages,
390                                compacted_messages,
391                            },
392                        ))
393                    }
394                }
395            },
396            on_receive_request!(),
397        )
398}
399
400#[cfg(test)]
401mod tests {
402    use super::*;
403    use crate::acp;
404    use assert_fs::TempDir;
405    use vtcode_core::llm::provider::Message;
406
407    use super::super::test_support::build_agent;
408
409    fn message(role: MessageRole, text: &str) -> Message {
410        Message {
411            role,
412            content: text.to_string().into(),
413            ..Message::default()
414        }
415    }
416
417    #[test]
418    fn turn_boundary_keeps_exactly_the_requested_trailing_turns() {
419        // Two user turns, each with a tool message: [u, t, a, u, t, a]
420        let messages = vec![
421            message(MessageRole::User, "one"),
422            message(MessageRole::Tool, "tool-1"),
423            message(MessageRole::Assistant, "reply-1"),
424            message(MessageRole::User, "two"),
425            message(MessageRole::Tool, "tool-2"),
426            message(MessageRole::Assistant, "reply-2"),
427        ];
428        // keep=1: boundary at the second user message, so its tool message and
429        // reply stay attached to the kept turn.
430        assert_eq!(turn_boundary_index(&messages, 1), 3);
431        // keep=2: nothing removed.
432        assert_eq!(turn_boundary_index(&messages, 2), 0);
433        // keep=3 (more turns than exist): nothing removed.
434        assert_eq!(turn_boundary_index(&messages, 3), 0);
435    }
436
437    #[test]
438    fn turn_boundary_zero_clears_and_userless_history_is_untouched() {
439        let messages = vec![
440            message(MessageRole::User, "one"),
441            message(MessageRole::Assistant, "reply"),
442        ];
443        assert_eq!(turn_boundary_index(&messages, 0), messages.len());
444
445        // No user turns at all (e.g. system-only history): rollback is a no-op
446        // rather than silently clearing the window.
447        let system_only = vec![message(MessageRole::System, "instructions")];
448        assert_eq!(turn_boundary_index(&system_only, 1), 0);
449        assert_eq!(turn_boundary_index(&system_only, 0), system_only.len());
450    }
451
452    #[tokio::test]
453    async fn fork_clones_history_into_an_independent_session() {
454        let temp = TempDir::new().expect("workspace");
455        let agent = build_agent(temp.path()).await;
456        let parent = agent
457            .new_session(acp::NewSessionRequest::new(temp.path().to_path_buf()))
458            .await
459            .expect("parent session");
460        agent.push_message(
461            &agent.session_handle(&parent.session_id).expect("parent handle"),
462            message(MessageRole::User, "hi"),
463        );
464
465        let forked_id = agent.fork_session(&parent.session_id).await.expect("fork");
466        assert_ne!(forked_id, parent.session_id, "fork must be a new session id");
467
468        let parent_messages = agent
469            .session_handle(&parent.session_id)
470            .expect("parent handle")
471            .data
472            .lock()
473            .expect("lock")
474            .thread
475            .messages();
476        let forked_messages = agent
477            .session_handle(&forked_id)
478            .expect("forked handle")
479            .data
480            .lock()
481            .expect("lock")
482            .thread
483            .messages();
484        assert_eq!(parent_messages.len(), forked_messages.len(), "fork inherits history");
485        assert_eq!(forked_messages.last().map(|m| m.role), Some(MessageRole::User));
486    }
487
488    #[tokio::test]
489    async fn rollback_truncates_the_requested_turns() {
490        let temp = TempDir::new().expect("workspace");
491        let agent = build_agent(temp.path()).await;
492        let session = agent
493            .new_session(acp::NewSessionRequest::new(temp.path().to_path_buf()))
494            .await
495            .expect("session");
496        let handle = agent.session_handle(&session.session_id).expect("handle");
497        agent.push_message(&handle, message(MessageRole::User, "turn-1"));
498        agent.push_message(&handle, message(MessageRole::Assistant, "reply-1"));
499        agent.push_message(&handle, message(MessageRole::User, "turn-2"));
500        agent.push_message(&handle, message(MessageRole::Assistant, "reply-2"));
501
502        let remaining = agent.rollback_session(&session.session_id, 1).expect("rollback");
503        assert_eq!(remaining, 2, "only the last user turn and its reply remain");
504        let messages = handle.data.lock().expect("lock").thread.messages();
505        assert_eq!(messages[0].role, MessageRole::User);
506        assert_eq!(messages.len(), 2);
507
508        // Clearing the whole conversation is expressed as keep=0.
509        let remaining = agent.rollback_session(&session.session_id, 0).expect("rollback");
510        assert_eq!(remaining, 0);
511    }
512
513    #[tokio::test]
514    async fn lifecycle_operations_reject_unknown_sessions() {
515        let temp = TempDir::new().expect("workspace");
516        let agent = build_agent(temp.path()).await;
517        let unknown = SessionId::new(Arc::from("vtcode-zed-session-404"));
518
519        let fork_error = agent.fork_session(&unknown).await.expect_err("unknown fork");
520        assert!(
521            fork_error.data.as_ref().is_some_and(|data| data["reason"] == "unknown_session"),
522            "fork must report unknown_session: {fork_error:?}"
523        );
524        let rollback_error = agent.rollback_session(&unknown, 1).expect_err("unknown rollback");
525        assert!(
526            rollback_error
527                .data
528                .as_ref()
529                .is_some_and(|data| data["reason"] == "unknown_session"),
530            "rollback must report unknown_session: {rollback_error:?}"
531        );
532    }
533
534    #[test]
535    fn history_unchanged_guard_accepts_only_exact_snapshots() {
536        let snapshot = vec![
537            message(MessageRole::User, "one"),
538            message(MessageRole::Assistant, "reply"),
539        ];
540
541        // Identical history: the operation may replace it in place.
542        assert!(ensure_history_unchanged(&snapshot, &snapshot).is_ok());
543
544        // Appended turn: replacing would drop the new messages.
545        let appended = {
546            let mut messages = snapshot.clone();
547            messages.push(message(MessageRole::User, "two"));
548            messages
549        };
550        assert!(ensure_history_unchanged(&appended, &snapshot).is_err());
551
552        // Rolled-back history: replacing would resurrect dropped turns.
553        let truncated = vec![message(MessageRole::User, "one")];
554        assert!(ensure_history_unchanged(&truncated, &snapshot).is_err());
555
556        // Same length, different content: a length-only check would miss this.
557        let mutated = vec![
558            message(MessageRole::User, "one"),
559            message(MessageRole::Assistant, "different reply"),
560        ];
561        assert!(ensure_history_unchanged(&mutated, &snapshot).is_err());
562    }
563
564    #[test]
565    fn lifecycle_wire_types_parse_their_own_json() {
566        let session_id = SessionId::new(Arc::from("vtcode-zed-session-1"));
567        let fork = SessionForkRequest { session_id: session_id.clone() };
568        let parsed: SessionForkRequest = SessionForkRequest::parse_message("session/fork", &fork).expect("parse");
569        assert_eq!(parsed.session_id, session_id);
570        assert!(SessionForkRequest::matches_method("session/fork"));
571        assert!(!SessionForkRequest::matches_method("session/rollback"));
572        assert!(SessionForkRequest::parse_message("session/rollback", &fork).is_err());
573
574        let rollback = SessionRollbackRequest { session_id, keep_last_turns: 2 };
575        let parsed: SessionRollbackRequest =
576            SessionRollbackRequest::parse_message("session/rollback", &rollback).expect("parse");
577        assert_eq!(parsed.keep_last_turns, 2);
578    }
579}