fix(tkmind): restore session history, resume, and restart parity
Load conversation on GET /sessions/{id}, port POST /agent/resume, and align
/agent/restart with 1.41 provider restore plus extension_results output.
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -13,6 +13,7 @@ use tracing::error;
|
||||
|
||||
use crate::agents::extension::ToolInfo;
|
||||
use crate::agents::extension_manager::get_parameter_names;
|
||||
use crate::agents::ExtensionLoadResult;
|
||||
use crate::agents::platform_extensions::{chatrecall, projectmemory, PLATFORM_EXTENSIONS};
|
||||
use crate::agents::ExtensionConfig;
|
||||
use crate::config::permission::PermissionLevel;
|
||||
@@ -70,6 +71,29 @@ pub struct RestartAgentRequest {
|
||||
session_id: String,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct ResumeAgentRequest {
|
||||
session_id: String,
|
||||
#[serde(default = "default_load_model_and_extensions")]
|
||||
load_model_and_extensions: bool,
|
||||
}
|
||||
|
||||
fn default_load_model_and_extensions() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct ResumeAgentResponse {
|
||||
pub session: Session,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extension_results: Option<Vec<ExtensionLoadResult>>,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct RestartAgentResponse {
|
||||
pub extension_results: Vec<ExtensionLoadResult>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct StartAgentRequest {
|
||||
working_dir: String,
|
||||
@@ -415,26 +439,129 @@ async fn update_working_dir(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn restart_agent_internal(
|
||||
state: &Arc<AppState>,
|
||||
session_id: &str,
|
||||
session: &Session,
|
||||
) -> Result<Vec<ExtensionLoadResult>, ErrorResponse> {
|
||||
let _ = state.agent_manager.remove_session(session_id).await;
|
||||
|
||||
let agent = state
|
||||
.get_agent_for_route(session_id.to_string())
|
||||
.await
|
||||
.map_err(|code| ErrorResponse {
|
||||
message: "Failed to create new agent during restart".into(),
|
||||
status: code,
|
||||
})?;
|
||||
|
||||
let provider_future = agent.restore_provider_from_session(session);
|
||||
let extensions_future = agent.load_extensions_from_session(session);
|
||||
|
||||
let (provider_result, extension_results) = tokio::join!(provider_future, extensions_future);
|
||||
provider_result.map_err(|err| ErrorResponse {
|
||||
message: err.to_string(),
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
})?;
|
||||
|
||||
Ok(extension_results)
|
||||
}
|
||||
|
||||
async fn restart_agent(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(payload): Json<RestartAgentRequest>,
|
||||
) -> Result<(), (StatusCode, String)> {
|
||||
) -> Result<Json<RestartAgentResponse>, ErrorResponse> {
|
||||
let session = state
|
||||
.session_manager()
|
||||
.get_session(&payload.session_id, false)
|
||||
.await
|
||||
.map_err(|_| (StatusCode::NOT_FOUND, "Session not found".to_string()))?;
|
||||
let agent = state
|
||||
.get_agent_for_route(payload.session_id.clone())
|
||||
.map_err(|err| {
|
||||
error!("Failed to get session during restart: {err}");
|
||||
ErrorResponse::not_found(format!("Failed to get session: {err}"))
|
||||
})?;
|
||||
|
||||
let extension_results =
|
||||
restart_agent_internal(&state, &payload.session_id, &session).await?;
|
||||
|
||||
Ok(Json(RestartAgentResponse { extension_results }))
|
||||
}
|
||||
|
||||
async fn resume_agent(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(payload): Json<ResumeAgentRequest>,
|
||||
) -> Result<Json<ResumeAgentResponse>, ErrorResponse> {
|
||||
let session = state
|
||||
.session_manager()
|
||||
.get_session(&payload.session_id, true)
|
||||
.await
|
||||
.map_err(|status| (status, "No agent for session id".to_owned()))?;
|
||||
agent.load_extensions_from_session(&session).await;
|
||||
Ok(())
|
||||
.map_err(|err| {
|
||||
error!("Failed to resume session {}: {err}", payload.session_id);
|
||||
ErrorResponse::not_found(format!("Failed to resume session: {err}"))
|
||||
})?;
|
||||
|
||||
let (extension_results, session) = if payload.load_model_and_extensions {
|
||||
let agent = state
|
||||
.get_agent_for_route(payload.session_id.clone())
|
||||
.await
|
||||
.map_err(|code| ErrorResponse {
|
||||
message: "Failed to get agent for route".into(),
|
||||
status: code,
|
||||
})?;
|
||||
|
||||
let provider_changed = agent
|
||||
.restore_provider_from_session(&session)
|
||||
.await
|
||||
.map_err(|err| ErrorResponse {
|
||||
message: err.to_string(),
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
})?;
|
||||
|
||||
let session = if provider_changed {
|
||||
state
|
||||
.session_manager()
|
||||
.get_session(&payload.session_id, true)
|
||||
.await
|
||||
.map_err(|err| ErrorResponse {
|
||||
message: format!("Failed to re-fetch session: {err}"),
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
})?
|
||||
} else {
|
||||
session
|
||||
};
|
||||
|
||||
let extension_results =
|
||||
if let Ok(Some(results)) = state.take_extension_loading_task(&payload.session_id).await
|
||||
{
|
||||
tracing::debug!(
|
||||
"Using background extension loading results for session {}",
|
||||
payload.session_id
|
||||
);
|
||||
state
|
||||
.remove_extension_loading_task(&payload.session_id)
|
||||
.await;
|
||||
results
|
||||
} else {
|
||||
tracing::debug!(
|
||||
"No background task found, loading extensions for session {}",
|
||||
payload.session_id
|
||||
);
|
||||
agent.load_extensions_from_session(&session).await
|
||||
};
|
||||
|
||||
(Some(extension_results), session)
|
||||
} else {
|
||||
(None, session)
|
||||
};
|
||||
|
||||
Ok(Json(ResumeAgentResponse {
|
||||
session,
|
||||
extension_results,
|
||||
}))
|
||||
}
|
||||
|
||||
pub fn routes(state: Arc<AppState>) -> Router {
|
||||
Router::new()
|
||||
.route("/agent/start", post(start_agent))
|
||||
.route("/agent/resume", post(resume_agent))
|
||||
.route("/agent/tools", get(get_tools))
|
||||
.route("/agent/update_provider", post(update_agent_provider))
|
||||
.route("/agent/update_session", post(update_session))
|
||||
|
||||
@@ -86,7 +86,7 @@ async fn get_session(
|
||||
) -> Result<Json<Session>, StatusCode> {
|
||||
state
|
||||
.session_manager()
|
||||
.get_session(&session_id, false)
|
||||
.get_session(&session_id, true)
|
||||
.await
|
||||
.map(Json)
|
||||
.map_err(|_| StatusCode::NOT_FOUND)
|
||||
|
||||
Reference in New Issue
Block a user