1use 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#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
49pub struct SessionForkRequest {
50 pub session_id: SessionId,
52}
53
54#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
56pub struct SessionForkResponse {
57 pub session_id: SessionId,
60}
61
62#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
64pub struct SessionRollbackRequest {
65 pub session_id: SessionId,
67 pub keep_last_turns: u32,
69}
70
71#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
73pub struct SessionRollbackResponse {
74 pub remaining_messages: usize,
76}
77
78#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
81pub struct SessionCompactRequest {
82 pub session_id: SessionId,
84}
85
86#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
88pub struct SessionCompactResponse {
89 pub original_messages: usize,
91 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
151impl ZedAgent {
156 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 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 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 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 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 ensure_history_unchanged(&data.thread.messages(), &history)?;
226 data.thread.replace_messages(compacted);
227 Ok((original, compacted_len))
228 }
229}
230
231fn 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
245fn 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 None => 0,
264 }
265}
266
267pub 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#[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
334pub 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 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 assert_eq!(turn_boundary_index(&messages, 1), 3);
431 assert_eq!(turn_boundary_index(&messages, 2), 0);
433 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 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 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 assert!(ensure_history_unchanged(&snapshot, &snapshot).is_ok());
543
544 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 let truncated = vec![message(MessageRole::User, "one")];
554 assert!(ensure_history_unchanged(&truncated, &snapshot).is_err());
555
556 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}