test: cover memory lifecycle outcomes

This commit is contained in:
2026-09-04 20:11:29 +09:00
parent 532d078720
commit d1f5661881
2 changed files with 301 additions and 11 deletions
+64 -4
View File
@@ -901,6 +901,21 @@ pub(crate) fn wire_event_bridges_on_engine<C, St>(
// per-item commit channel is wired at the top of this function.
}
fn add_memory_lifecycle_if_configured<M>(
registry: &mut FeatureRegistryBuilder,
config: Option<manifest::MemoryConfig>,
build: impl FnOnce(manifest::MemoryConfig) -> std::io::Result<M>,
) -> std::io::Result<bool>
where
M: crate::feature::FeatureModule + 'static,
{
let Some(config) = config else {
return Ok(false);
};
registry.add_module(build(config)?);
Ok(true)
}
/// Register the builtin file-manipulation tools, optional memory tools,
/// and the Worker-orchestration tools (SubWorkerSpawn + comm) on the Worker's
/// Engine. Returns the WorkdirSession handle used to attach a `WorkerFsView` to
@@ -994,7 +1009,7 @@ where
let worker_enabled = feature_config.worker.enabled;
let sub_worker_enabled = feature_config.sub_worker.enabled;
let mut feature_registry = FeatureRegistryBuilder::new();
if let Some(config) = memory_config.clone() {
add_memory_lifecycle_if_configured(&mut feature_registry, memory_config.clone(), |config| {
let workspace_client = worker.workspace_client_handle();
if !workspace_client.is_available() || workspace_client.workspace_id().is_none() {
return Err(std::io::Error::new(
@@ -1002,7 +1017,7 @@ where
"Memory extraction requires Backend Workspace API authority",
));
}
feature_registry.add_module(
Ok(
crate::feature::builtin::memory_lifecycle::MemoryLifecycleFeature::new(
config,
worker.committed_session_capture_handle(),
@@ -1014,8 +1029,8 @@ where
spawner_workspace_context.clone(),
worker.working_event_sender(),
),
);
}
)
})?;
if sub_worker_enabled && !worker_enabled {
feature_registry.add_module(
crate::feature::builtin::manage_worker::sub_worker_control_feature(
@@ -2135,6 +2150,51 @@ mod tests {
use tempfile::TempDir;
use tokio::net::UnixListener;
#[test]
fn memory_lifecycle_registration_depends_only_on_memory_config_presence() {
#[derive(Clone)]
struct TestMemoryLifecycleModule;
impl crate::feature::FeatureModule for TestMemoryLifecycleModule {
fn descriptor(&self) -> crate::feature::FeatureDescriptor {
crate::feature::FeatureDescriptor::builtin(
"test-memory-lifecycle",
"Test Memory Lifecycle",
)
}
fn install(
&self,
_context: &mut crate::feature::FeatureInstallContext<'_>,
) -> Result<(), crate::feature::FeatureInstallError> {
Ok(())
}
}
let mut registry = FeatureRegistryBuilder::new();
let configured = std::cell::Cell::new(false);
let installed = add_memory_lifecycle_if_configured(
&mut registry,
Some(manifest::MemoryConfig::default()),
|_| {
configured.set(true);
Ok(TestMemoryLifecycleModule)
},
)
.unwrap();
assert!(installed);
assert!(configured.get());
let mut registry = FeatureRegistryBuilder::new();
let installed = add_memory_lifecycle_if_configured::<TestMemoryLifecycleModule>(
&mut registry,
None,
|_| panic!("disabled Memory must not construct its lifecycle Feature"),
)
.unwrap();
assert!(!installed);
}
#[test]
fn image_attachment_gate_requires_vision_and_supported_openai_scheme() {
let openai = manifest::ModelManifest {
@@ -833,13 +833,69 @@ mod tests {
&self,
request: crate::worker::WorkspaceRequest,
) -> Result<crate::worker::WorkspaceResponse, crate::worker::WorkspaceClientError> {
let is_stage_candidate = request
.body
.as_deref()
.is_some_and(|body| body.contains("stage_candidate"));
let is_append_audit = request
.body
.as_deref()
.is_some_and(|body| body.contains("append_audit"));
self.requests.lock().unwrap().push(request);
if is_stage_candidate {
return Ok(crate::worker::WorkspaceResponse {
status: 200,
body: serde_json::to_string(&memory::backend::MemoryBackendHttpResponse::Ok {
result: memory::backend::MemoryBackendOperationResult::StagingWritten(
memory::backend::MemoryStagingWriteOutput {
staging_count: 1,
staging_ids: vec!["candidate-1".to_string()],
},
),
})
.unwrap(),
});
}
if is_append_audit {
return Ok(crate::worker::WorkspaceResponse {
status: 200,
body: serde_json::to_string(&memory::backend::MemoryBackendHttpResponse::Ok {
result: memory::backend::MemoryBackendOperationResult::Acknowledged(
memory::backend::MemoryBackendAckOutput {
summary: "audit recorded".to_string(),
},
),
})
.unwrap(),
});
}
Err(crate::worker::WorkspaceClientError::Unavailable(
"recording client".to_string(),
))
}
}
#[derive(Clone)]
struct PendingClient {
calls: Arc<AtomicUsize>,
}
#[async_trait]
impl LlmClient for PendingClient {
fn clone_boxed(&self) -> Box<dyn LlmClient> {
Box::new(self.clone())
}
async fn stream(
&self,
_request: Request,
) -> Result<Pin<Box<dyn Stream<Item = Result<LlmEvent, ClientError>> + Send>>, ClientError>
{
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(Box::pin(futures::stream::pending()))
}
}
#[derive(Clone)]
struct ScriptClient {
responses: Arc<Vec<Vec<LlmEvent>>>,
@@ -874,6 +930,40 @@ mod tests {
}
}
fn stage_candidate_events(call_id: &str, entry_ref: &str) -> Vec<LlmEvent> {
vec![
LlmEvent::tool_use_start(0, call_id, "StageMemoryCandidate"),
LlmEvent::tool_input_delta(
0,
serde_json::json!({
"kind": "decision",
"claim": "Keep lifecycle work feature-owned.",
"why_useful": "Prevents Worker core coupling.",
"entry_refs": [entry_ref]
})
.to_string(),
),
LlmEvent::tool_use_stop(0),
LlmEvent::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
]
}
fn finish_events(call_id: &str, staged_count: usize) -> Vec<LlmEvent> {
vec![
LlmEvent::tool_use_start(0, call_id, "FinishMemoryExtraction"),
LlmEvent::tool_input_delta(
0,
serde_json::json!({"staged_count": staged_count}).to_string(),
),
LlmEvent::tool_use_stop(0),
LlmEvent::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
]
}
fn finish_empty_events(call_id: &str) -> Vec<LlmEvent> {
vec![
LlmEvent::tool_use_start(0, call_id, "FinishMemoryExtraction"),
@@ -939,7 +1029,7 @@ permission = "write"
fn test_task(
capture: CommittedSessionCapture,
client: ScriptClient,
client: Box<dyn LlmClient>,
extension_writes: Arc<Mutex<Vec<(CommittedSessionLocation, String, serde_json::Value)>>>,
event_tx: broadcast::Sender<Event>,
workspace_client: Arc<dyn WorkspaceClient>,
@@ -958,14 +1048,16 @@ permission = "write"
extensions,
workspace_client,
manifest: test_manifest(),
client: Box::new(client),
client,
prompts: Arc::new(ArcSwap::from(PromptCatalog::builtins_only().unwrap())),
workspace_context: WorkerWorkspaceContext::no_workspace(),
event_tx: Some(event_tx),
}
}
async fn run_background_task(task: MemoryLifecycleTask) {
fn start_background_task(
task: MemoryLifecycleTask,
) -> crate::feature::background::FeatureBackgroundTaskRegistry {
let mut builder = FeatureBackgroundTaskRegistryBuilder::default();
builder
.register(
@@ -984,6 +1076,11 @@ permission = "write"
..Default::default()
})
.unwrap();
registry
}
async fn run_background_task(task: MemoryLifecycleTask) {
let registry = start_background_task(task);
tokio::time::timeout(Duration::from_secs(5), async {
loop {
if !registry.diagnostics().is_empty() {
@@ -1007,7 +1104,7 @@ permission = "write"
Arc::new(RecordingWorkspaceClient::default());
run_background_task(test_task(
capture(2, 250),
client,
Box::new(client),
Arc::clone(&extension_writes),
event_tx,
workspace_client,
@@ -1024,6 +1121,139 @@ permission = "write"
assert_eq!(pointer.processed_through_entry, 1);
}
#[tokio::test]
async fn run_committed_background_task_stages_non_empty_extraction_and_commits_pointer() {
let source = capture(2, 250);
let entry_ref =
SessionCapture::from_history_entries(source.segment_id.clone(), source.history.clone())
.overview()[0]
.id
.to_string();
let client = ScriptClient::new(vec![
stage_candidate_events("stage-1", &entry_ref),
finish_events("finish-1", 1),
completed_events(),
]);
let calls = Arc::clone(&client.calls);
let extension_writes = Arc::new(Mutex::new(Vec::new()));
let (event_tx, mut event_rx) = broadcast::channel(64);
let workspace_client = Arc::new(RecordingWorkspaceClient::default());
run_background_task(test_task(
source,
Box::new(client),
Arc::clone(&extension_writes),
event_tx,
workspace_client.clone(),
))
.await;
assert_eq!(calls.load(Ordering::SeqCst), 3);
let writes = extension_writes.lock().unwrap();
assert_eq!(
writes.len(),
1,
"recorded requests: {:?}; events: {:?}",
workspace_client.requests.lock().unwrap(),
std::iter::from_fn(|| event_rx.try_recv().ok()).collect::<Vec<_>>()
);
let pointer: memory::ExtractPointerPayload =
serde_json::from_value(writes[0].2.clone()).unwrap();
assert_eq!(pointer.staging_id, "candidate-1");
assert!(
workspace_client
.requests
.lock()
.unwrap()
.iter()
.any(|request| {
request
.body
.as_deref()
.is_some_and(|body| body.contains("stage_candidate"))
})
);
}
#[tokio::test]
async fn failed_extraction_emits_failure_event_and_durable_audit_without_pointer() {
let client = ScriptClient::new(Vec::new());
let calls = Arc::clone(&client.calls);
let extension_writes = Arc::new(Mutex::new(Vec::new()));
let (event_tx, mut event_rx) = broadcast::channel(16);
let workspace_client = Arc::new(RecordingWorkspaceClient::default());
run_background_task(test_task(
capture(2, 250),
Box::new(client),
Arc::clone(&extension_writes),
event_tx,
workspace_client.clone(),
))
.await;
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert!(extension_writes.lock().unwrap().is_empty());
let events = std::iter::from_fn(|| event_rx.try_recv().ok()).collect::<Vec<_>>();
assert!(events.iter().any(|event| {
matches!(event, Event::MemoryWorker(event) if event.status == "failed")
}));
let requests = workspace_client.requests.lock().unwrap();
assert!(
requests.iter().any(|request| {
request.body.as_deref().is_some_and(|body| {
body.contains("append_audit")
&& body.contains("worker_lifecycle")
&& body.contains("failed")
})
}),
"recorded requests: {requests:?}"
);
}
#[tokio::test]
async fn rewrite_barrier_cancels_active_extraction_and_emits_cancelled_without_pointer() {
let calls = Arc::new(AtomicUsize::new(0));
let client = PendingClient {
calls: Arc::clone(&calls),
};
let extension_writes = Arc::new(Mutex::new(Vec::new()));
let (event_tx, mut event_rx) = broadcast::channel(16);
let workspace_client = Arc::new(RecordingWorkspaceClient::default());
let registry = start_background_task(test_task(
capture(2, 250),
Box::new(client),
Arc::clone(&extension_writes),
event_tx,
workspace_client.clone(),
));
tokio::time::timeout(Duration::from_secs(5), async {
while calls.load(Ordering::SeqCst) == 0 {
tokio::task::yield_now().await;
}
})
.await
.expect("extraction child should reach its first provider request");
let rewrite_guard = registry.begin_session_rewrite().await.unwrap();
drop(rewrite_guard);
registry.shutdown().await.unwrap();
assert!(extension_writes.lock().unwrap().is_empty());
let events = std::iter::from_fn(|| event_rx.try_recv().ok()).collect::<Vec<_>>();
assert!(events.iter().any(|event| {
matches!(event, Event::MemoryWorker(event) if event.status == "cancelled")
}));
let requests = workspace_client.requests.lock().unwrap();
assert!(
requests.iter().any(|request| {
request.body.as_deref().is_some_and(|body| {
body.contains("append_audit")
&& body.contains("worker_lifecycle")
&& body.contains("cancelled")
})
}),
"recorded requests: {requests:?}"
);
}
#[tokio::test]
async fn interrupted_committed_run_skips_internal_worker_and_pointer_commit() {
let client = ScriptClient::new(Vec::new());
@@ -1036,7 +1266,7 @@ permission = "write"
interrupted.run_exit = CommittedRunExit::Interrupted;
run_background_task(test_task(
interrupted,
client,
Box::new(client),
Arc::clone(&extension_writes),
event_tx,
workspace_client,
@@ -1061,7 +1291,7 @@ permission = "write"
interrupted.run_exit = CommittedRunExit::Interrupted;
let mut task = test_task(
interrupted,
client,
Box::new(client),
extension_writes,
event_tx,
workspace_client.clone(),
@@ -1259,7 +1489,7 @@ permission = "write"
);
}
let controller_source = include_str!("../../controller.rs");
assert!(controller_source.contains("if let Some(config) = memory_config.clone()"));
assert!(controller_source.contains("add_memory_lifecycle_if_configured"));
assert!(controller_source.contains("MemoryLifecycleFeature::new"));
let lifecycle_source = include_str!("memory_lifecycle.rs");
assert!(lifecycle_source.contains("request_memory_staging_consolidation"));