Compare commits

..
Author SHA1 Message Date
[email protected] ebdae267a0 Merge remote-tracking branch 'origin/staging' into feat/workspace-metadata-versioning-patch
# Conflicts:
#	src/agent/heartbeat.rs
#	src/db/libsql_migrations.rs
2026-03-28 12:27:02 -07:00
[email protected]andClaude Opus 4.6 cb4cf123c6 fix: address review feedback — transaction safety, patch mode, formatting
- Wrap libSQL save_version in a transaction to prevent race condition
  where concurrent writers could allocate the same version number
- Make content optional in memory_write when in patch mode (old_string
  present) — LLM no longer forced to provide unused content param
- Improve metadata update error handling with explicit match arms
- Run cargo fmt across all files

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 12:12:41 -07:00
fd41bdf4be fix(worker): treat empty LLM response after text output as completion (#1677)
* fix(worker): treat empty LLM response after text output as completion

When a job's LLM produces a substantive text response (e.g., formatted
results from a routine) and the next LLM call returns empty or errors,
the worker now treats this as successful completion instead of
continuing the loop until failure.

Previously, empty responses always triggered TextAction::Continue,
causing the loop to re-call the LLM. The LLM had nothing more to say,
so the provider returned "Response contained no message or tool call
(empty)". This made routine jobs that successfully produced results
report as "failed".

The fix adds a `has_text_response` flag to JobDelegate:
- After any non-empty text response: flag is set
- Empty text after flag is set: treated as completion
- LLM errors (select_tools/respond_with_tools) after flag: treated
  as completion instead of propagating
- Empty text before any output: still retries (rate-limit backoff)

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* fix(worker): restrict error swallowing to EmptyResponse variant only

- Add LlmError::EmptyResponse variant for when LLM returns no content
- Update nearai_chat and github_copilot providers to emit EmptyResponse
  instead of InvalidResponse for empty/no-choice responses
- try_complete_on_error now only swallows EmptyResponse (not AuthFailed,
  ContextLengthExceeded, Http, Io, etc.)
- Extract is_completion_eligible_error as testable pure function
- Log mark_completed errors at warn level instead of silently dropping
- Add EmptyResponse to retry and circuit breaker transient classifications
- Rewrite test to exercise real classification logic against all variants

Addresses review feedback from zmanian and gemini-code-assist.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* refactor(worker): extract mark_completed_or_warn helper to DRY completion logic

Extract shared mark-completed + warn-on-failure pattern into a single
helper method used by both try_complete_on_error and handle_text_response.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

---------

Co-authored-by: j-bloggs <[email protected]>
Co-authored-by: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 18:46:08 +01:00
de5a1c7b0d fix(worker): replace script -qfc with pty-process for injection-safe PTY (#1678)
- Add pty-process crate (MIT, tokio async support) for PTY allocation
- Spawn claude CLI with pty-process::Command::arg() chaining instead of
  building a shell string for script -qfc
- Eliminates all shell injection surfaces: prompt, model, session_id
  are passed via execve, never interpreted by a shell
- Keep stderr on separate pipe to prevent NDJSON parse breakage
  (pty-process attaches PTY to all fds by default)
- Gate PTY behind #[cfg(unix)] with direct-spawn fallback for Windows CI
- Read stdout from PTY master (implements tokio::io::AsyncRead)
- Add regression tests: arg vector construction + PTY allocation

Addresses review feedback from zmanian and gemini-code-assist.

Co-authored-by: j-bloggs <[email protected]>
Co-authored-by: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 16:31:49 +01:00
9ce3a9fc53 feat(discord): implement on_broadcast via DM channel creation (#1693)
- Implement broadcast_dm() that creates a DM channel with the target
  user (POST /users/@me/channels, cached by Discord) and sends the
  message to it
- Extract DISCORD_API_BASE constant for all Discord REST API URLs
- Extract send_channel_message() shared helper to deduplicate message
  posting between on_respond and broadcast_dm
- Add snowflake validation on user_id before API calls
- Fix pre-existing clippy redundant_closure warning
- Use typed DmChannelResponse struct instead of serde_json::Value

Closes no specific issue — completes the previously stubbed on_broadcast.

Co-authored-by: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 16:31:27 +01:00
[email protected]andClaude Opus 4.6 7532065590 feat(workspace): metadata-driven indexing/hygiene, document versioning, and patch support
Foundation for the extensible frontend system. Workspace documents now
support metadata flags (skip_indexing, skip_versioning, hygiene config)
via folder-level .config documents and per-file overrides, replacing
hardcoded hygiene targets and indexing behavior.

Key changes:
- DocumentMetadata type with resolution chain (doc → folder .config → defaults)
- Document versioning: auto-saves previous content on write/append/patch
- Workspace patch: search-and-replace editing via memory_write tool
- Hygiene rewrite: discovers cleanup targets from .config metadata
  instead of hardcoded daily/ and conversations/ directories
- memory_read gains version/list_versions params
- memory_write gains metadata/old_string/new_string/replace_all params
- V14 migration adds memory_document_versions table (both PG + libSQL)

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-28 00:28:32 -07:00
42 changed files with 2190 additions and 764 deletions
+2 -3
View File
@@ -191,10 +191,9 @@ HEARTBEAT_NOTIFY_CHANNEL=cli
HEARTBEAT_NOTIFY_USER=default
# Memory hygiene settings (automatic cleanup of stale workspace documents)
# Runs on each heartbeat tick; identity files (IDENTITY.md, SOUL.md) are never deleted
# Runs on each heartbeat tick; discovers cleanup targets from .config metadata
# MEMORY_HYGIENE_ENABLED=true
# MEMORY_HYGIENE_DAILY_RETENTION_DAYS=30 # delete daily/ docs older than this many days
# MEMORY_HYGIENE_CONVERSATION_RETENTION_DAYS=7 # delete conversations/ docs older than this many days
# MEMORY_HYGIENE_VERSION_KEEP_COUNT=50 # max versions to keep per document
# MEMORY_HYGIENE_CADENCE_HOURS=12 # minimum hours between cleanup passes
# Docker Sandbox
Generated
+16 -5
View File
@@ -3150,7 +3150,7 @@ dependencies = [
"libc",
"percent-encoding",
"pin-project-lite",
"socket2 0.5.10",
"socket2 0.6.3",
"system-configuration",
"tokio",
"tower-service",
@@ -3439,6 +3439,7 @@ dependencies = [
"pgvector",
"postgres-types",
"pretty_assertions",
"pty-process",
"rand 0.8.5",
"readabilityrs",
"refinery",
@@ -3524,7 +3525,7 @@ checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46"
dependencies = [
"hermit-abi",
"libc",
"windows-sys 0.59.0",
"windows-sys 0.61.2",
]
[[package]]
@@ -4906,6 +4907,16 @@ dependencies = [
"syn 1.0.109",
]
[[package]]
name = "pty-process"
version = "0.5.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "71cec9e2670207c5ebb9e477763c74436af3b9091dd550b9fb3c1bec7f3ea266"
dependencies = [
"rustix 1.1.4",
"tokio",
]
[[package]]
name = "pulley-interpreter"
version = "28.0.1"
@@ -4930,7 +4941,7 @@ dependencies = [
"quinn-udp",
"rustc-hash 2.1.1",
"rustls 0.23.37",
"socket2 0.5.10",
"socket2 0.6.3",
"thiserror 2.0.18",
"tokio",
"tracing",
@@ -4967,9 +4978,9 @@ dependencies = [
"cfg_aliases",
"libc",
"once_cell",
"socket2 0.5.10",
"socket2 0.6.3",
"tracing",
"windows-sys 0.59.0",
"windows-sys 0.60.2",
]
[[package]]
+4
View File
@@ -189,6 +189,10 @@ json5 = { version = "0.4", optional = true }
[target.'cfg(target_os = "macos")'.dependencies]
security-framework = "3"
# PTY allocation for Claude CLI stdout buffering fix (Unix only)
[target.'cfg(unix)'.dependencies]
pty-process = { version = "0.5", features = ["async"] }
# Linux secret-service (GNOME Keyring, KWallet)
[target.'cfg(target_os = "linux")'.dependencies]
secret-service = { version = "4", features = ["rt-tokio-crypto-rust"] }
+118 -23
View File
@@ -28,6 +28,9 @@ use std::{cmp::Ordering, collections::HashMap};
use ed25519_dalek::{Signature, Verifier, VerifyingKey};
use serde::{Deserialize, Serialize};
/// Discord REST API v10 base URL.
const DISCORD_API_BASE: &str = "https://discord.com/api/v10";
use exports::near::agent::channel::{
AgentResponse, ChannelConfig, Guest, HttpEndpointConfig, IncomingHttpRequest,
OutgoingHttpResponse, PollConfig, StatusUpdate,
@@ -427,7 +430,7 @@ impl Guest for DiscordChannel {
(
"PATCH",
format!(
"https://discord.com/api/v10/webhooks/{}/{}/messages/@original",
"{DISCORD_API_BASE}/webhooks/{}/{}/messages/@original",
application_id, token
),
)
@@ -438,20 +441,7 @@ impl Guest for DiscordChannel {
payload["allowed_mentions"] = serde_json::json!({
"replied_user": true
});
let mention_payload = serde_json::to_vec(&payload)
.map_err(|e| format!("Failed to serialize mention payload: {}", e))?;
let mention_url = format!(
"https://discord.com/api/v10/channels/{}/messages",
metadata.channel_id
);
let result = channel_host::http_request(
"POST",
&mention_url,
&discord_auth_headers_json(true),
Some(&mention_payload),
None,
);
return map_discord_response(result);
return send_channel_message(&metadata.channel_id, payload);
} else {
return Err("Unsupported Discord response metadata".to_string());
};
@@ -469,8 +459,8 @@ impl Guest for DiscordChannel {
fn on_status(_update: StatusUpdate) {}
fn on_broadcast(_user_id: String, _response: AgentResponse) -> Result<(), String> {
Err("broadcast not yet implemented for Discord channel".to_string())
fn on_broadcast(user_id: String, response: AgentResponse) -> Result<(), String> {
broadcast_dm(&user_id, &response.content)
}
fn on_shutdown() {
@@ -501,6 +491,21 @@ fn map_discord_response(
}
}
/// Post a JSON payload to a Discord channel as a new message.
fn send_channel_message(channel_id: &str, payload: serde_json::Value) -> Result<(), String> {
let payload_bytes = serde_json::to_vec(&payload)
.map_err(|e| format!("Failed to serialize message: {}", e))?;
let url = format!("{DISCORD_API_BASE}/channels/{}/messages", channel_id);
let result = channel_host::http_request(
"POST",
&url,
&discord_auth_headers_json(true),
Some(&payload_bytes),
None,
);
map_discord_response(result)
}
fn load_runtime_config() -> DiscordRuntimeConfig {
channel_host::workspace_read("config.json")
.and_then(|raw| serde_json::from_str::<DiscordRuntimeConfig>(&raw).ok())
@@ -539,7 +544,7 @@ fn get_or_fetch_bot_id() -> Option<String> {
let response = channel_host::http_request(
"GET",
"https://discord.com/api/v10/users/@me",
&format!("{DISCORD_API_BASE}/users/@me"),
&discord_auth_headers_json(false),
None,
Some(10_000),
@@ -659,7 +664,7 @@ fn poll_channel_mentions(channel_id: &str, bot_id: &str) {
fn fetch_latest_message_id(channel_id: &str) -> Option<String> {
let url = format!(
"https://discord.com/api/v10/channels/{}/messages?limit=1",
"{DISCORD_API_BASE}/channels/{}/messages?limit=1",
channel_id
);
let response = channel_host::http_request(
@@ -697,7 +702,7 @@ fn fetch_messages_after_cursor(
for page in 0..MAX_PAGES {
let url = format!(
"https://discord.com/api/v10/channels/{}/messages?limit={}&after={}",
"{DISCORD_API_BASE}/channels/{}/messages?limit={}&after={}",
channel_id, PAGE_LIMIT, after
);
let response = match channel_host::http_request(
@@ -986,7 +991,7 @@ fn handle_slash_command(interaction: &DiscordInteraction) -> bool {
);
// Attempt to notify user of internal error
let url = format!(
"https://discord.com/api/v10/webhooks/{}/{}",
"{DISCORD_API_BASE}/webhooks/{}/{}",
interaction.application_id, interaction.token
);
let payload = serde_json::json!({
@@ -1106,7 +1111,7 @@ fn check_sender_permission(
}
let dm_policy =
channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(|| default_dm_policy());
channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(default_dm_policy);
if dm_policy == "open" {
return true;
}
@@ -1161,7 +1166,7 @@ fn check_sender_permission(
/// Send a pairing code as an ephemeral Discord followup message.
fn send_pairing_reply(ctx: &PairingReplyCtx, code: &str) -> Result<(), String> {
let url = format!(
"https://discord.com/api/v10/webhooks/{}/{}",
"{DISCORD_API_BASE}/webhooks/{}/{}",
ctx.application_id, ctx.token
);
let payload = serde_json::json!({
@@ -1194,6 +1199,57 @@ fn send_pairing_reply(ctx: &PairingReplyCtx, code: &str) -> Result<(), String> {
}
}
/// Send a broadcast message to a Discord user via DM.
///
/// Creates a DM channel with the user (Discord caches this, so repeated calls
/// for the same user reuse the existing channel) and then posts the message.
fn broadcast_dm(user_id: &str, content: &str) -> Result<(), String> {
// Validate user_id is a plausible Discord snowflake (numeric, 17-20 digits)
// to avoid injecting arbitrary strings into API URLs.
if user_id.is_empty()
|| !user_id.chars().all(|c| c.is_ascii_digit())
|| user_id.len() < 17
|| user_id.len() > 20
{
return Err(format!("Invalid Discord user ID: '{}'", user_id));
}
// Step 1: Open (or reuse) a DM channel with the target user.
let create_dm_payload = serde_json::json!({ "recipient_id": user_id });
let create_dm_bytes = serde_json::to_vec(&create_dm_payload)
.map_err(|e| format!("Failed to serialize DM channel request: {}", e))?;
let dm_response = channel_host::http_request(
"POST",
&format!("{DISCORD_API_BASE}/users/@me/channels"),
&discord_auth_headers_json(true),
Some(&create_dm_bytes),
Some(10_000),
)
.map_err(|e| format!("Failed to create DM channel: {}", e))?;
if !(200..300).contains(&dm_response.status) {
let body = String::from_utf8_lossy(&dm_response.body);
return Err(format!(
"Discord create-DM failed: {} - {}",
dm_response.status, body
));
}
#[derive(Deserialize)]
struct DmChannelResponse {
id: String,
}
let dm_channel: DmChannelResponse = serde_json::from_slice(&dm_response.body)
.map_err(|e| format!("Failed to parse DM channel response: {}", e))?;
let channel_id = &dm_channel.id;
// Step 2: Send the message to the DM channel.
let truncated = truncate_message(content);
let payload = serde_json::json!({ "content": truncated });
send_channel_message(channel_id, payload)
}
fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse {
let body = serde_json::to_vec(&value).unwrap_or_default();
let headers = serde_json::json!({"Content-Type": "application/json"});
@@ -1593,4 +1649,43 @@ mod tests {
assert_eq!(interaction.interaction_type, 2);
assert!(interaction.data.is_some());
}
#[test]
fn test_broadcast_dm_payload_format() {
// Verify the DM channel creation payload is well-formed JSON that
// Discord's API expects.
let user_id = "123456789012345678";
let payload = serde_json::json!({ "recipient_id": user_id });
let serialized = serde_json::to_vec(&payload).unwrap();
let parsed: serde_json::Value = serde_json::from_slice(&serialized).unwrap();
assert_eq!(
parsed.get("recipient_id").and_then(|v| v.as_str()),
Some(user_id)
);
}
#[test]
fn test_broadcast_message_truncation() {
// Broadcast uses truncate_message, verify it handles content within
// Discord's 2000-char limit for DMs.
let short = "Hello from broadcast";
assert_eq!(truncate_message(short), short);
let long = "x".repeat(2500);
let result = truncate_message(&long);
assert!(result.len() <= 2006); // 1990 content + 16 suffix
assert!(result.ends_with("\n... (truncated)"));
}
#[test]
fn test_broadcast_dm_validates_snowflake() {
// broadcast_dm rejects invalid Discord snowflake IDs before making
// any API calls. We can call it directly since invalid IDs are
// rejected before any host function is invoked.
assert!(broadcast_dm("", "hi").is_err());
assert!(broadcast_dm("abc", "hi").is_err());
assert!(broadcast_dm("12345", "hi").is_err()); // too short
assert!(broadcast_dm("123456789012345678901", "hi").is_err()); // too long
assert!(broadcast_dm("12345678901234567x", "hi").is_err()); // non-digit
}
}
@@ -1,4 +0,0 @@
-- Add source_channel to conversations for cross-channel approval authorization.
-- Tracks which channel originally created a conversation so that approval
-- messages from other channels can be validated.
ALTER TABLE conversations ADD COLUMN source_channel TEXT;
+23
View File
@@ -0,0 +1,23 @@
-- Document version history for workspace files.
-- Every content update saves the previous content as a version,
-- enabling rollback and audit trails.
CREATE TABLE memory_document_versions (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
document_id UUID NOT NULL REFERENCES memory_documents(id) ON DELETE CASCADE,
version INTEGER NOT NULL,
content TEXT NOT NULL,
content_hash TEXT NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
changed_by TEXT,
UNIQUE(document_id, version)
);
CREATE INDEX idx_doc_versions_lookup
ON memory_document_versions(document_id, version DESC);
-- GIN index on metadata for JSON path queries (used by hygiene to find
-- .config documents with hygiene.enabled). The metadata column already
-- exists (V1) but was never indexed.
CREATE INDEX idx_memory_documents_metadata
ON memory_documents USING GIN (metadata jsonb_path_ops);
+2 -35
View File
@@ -843,10 +843,7 @@ impl Agent {
{
use crate::agent::session::Thread;
let mut sess = session.lock().await;
// Bootstrap thread has no incoming message -- use the
// "__bootstrap__" sentinel so approvals from any channel are
// permitted. None means "deny by default" (fail-closed).
let thread = Thread::with_id(id, sess.id, Some("__bootstrap__"));
let thread = Thread::with_id(id, sess.id);
sess.active_thread = Some(id);
sess.threads.entry(id).or_insert(thread);
}
@@ -1151,37 +1148,7 @@ impl Agent {
.get_or_create_session(&message.user_id)
.await;
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get(&target_thread_id) {
// Verify the thread actually has a pending approval before
// allowing approval-shaped messages to target it. Without this
// check, an attacker could use approval messages to hijack any
// thread by UUID.
if thread.pending_approval.is_none() {
tracing::warn!(
%target_thread_id,
approval_channel = %message.channel,
"Blocked approval for thread with no pending approval"
);
drop(sess);
return Ok(Some("Error: no pending approval on this thread".into()));
}
let authorized = crate::agent::session::is_approval_authorized(
thread.source_channel.as_deref(),
&message.channel,
);
if !authorized {
tracing::warn!(
%target_thread_id,
source_channel = ?thread.source_channel,
approval_channel = %message.channel,
"Blocked cross-channel approval attempt"
);
drop(sess);
return Ok(Some(
"Error: approval not authorized for this channel".into(),
));
}
if sess.threads.contains_key(&target_thread_id) {
sess.active_thread = Some(target_thread_id);
sess.last_active_at = chrono::Utc::now();
drop(sess);
+5 -5
View File
@@ -319,7 +319,7 @@ mod tests {
#[test]
fn test_format_turns() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Hello");
thread.complete_turn("Hi there");
thread.start_turn("How are you?");
@@ -351,7 +351,7 @@ mod tests {
/// Helper: build a thread with `n` completed turns.
/// Turn `i` has user_input "msg-{i}" and response "resp-{i}".
fn make_thread(n: usize) -> Thread {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
for i in 0..n {
thread.start_turn(format!("msg-{}", i));
thread.complete_turn(format!("resp-{}", i));
@@ -457,7 +457,7 @@ mod tests {
async fn test_compact_truncate_empty_turns() {
let llm = Arc::new(StubLlm::new("unused"));
let compactor = make_compactor(llm);
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
assert!(thread.turns.is_empty());
let result = compactor
@@ -698,7 +698,7 @@ mod tests {
#[test]
fn test_format_turns_for_storage_with_tool_calls() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Search for X");
// Record a tool call on the current turn
if let Some(turn) = thread.turns.last_mut() {
@@ -719,7 +719,7 @@ mod tests {
#[test]
fn test_format_turns_for_storage_incomplete_turn() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("In progress message");
// Don't complete the turn
+2 -2
View File
@@ -2299,7 +2299,7 @@ mod tests {
// Initialize a thread in the session so the loop can record tool calls.
let thread_id = {
let mut sess = session.lock().await;
sess.create_thread(Some("test")).id
sess.create_thread().id
};
let message = IncomingMessage::new("test", "test-user", "do something");
@@ -2412,7 +2412,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("test-user")));
let thread_id = {
let mut sess = session.lock().await;
sess.create_thread(Some("test")).id
sess.create_thread().id
};
let message = IncomingMessage::new("test", "test-user", "keep calling tools");
+21 -4
View File
@@ -276,8 +276,8 @@ impl HeartbeatRunner {
.await;
if report.had_work() {
tracing::info!(
daily_logs_deleted = report.daily_logs_deleted,
conversation_docs_deleted = report.conversation_docs_deleted,
directories_cleaned = ?report.directories_cleaned,
versions_pruned = report.versions_pruned,
"heartbeat: memory hygiene deleted stale documents"
);
}
@@ -590,6 +590,23 @@ pub fn spawn_multi_user_heartbeat(
let workspace = Arc::new(Workspace::new_with_db(user_id, Arc::clone(store.db())));
// Run memory hygiene per user (same as single-user heartbeat).
let hygiene_ws = Arc::clone(&workspace);
let hygiene_cfg = hygiene_config.clone();
let hygiene_user = user_id.clone();
tokio::spawn(async move {
let report =
crate::workspace::hygiene::run_if_due(&hygiene_ws, &hygiene_cfg).await;
if report.had_work() {
tracing::info!(
user_id = hygiene_user,
directories_cleaned = ?report.directories_cleaned,
versions_pruned = report.versions_pruned,
"multi-user heartbeat: memory hygiene deleted stale documents"
);
}
});
// Drain completed tasks to stay within the concurrency cap.
while join_set.len() >= MAX_CONCURRENT_HEARTBEATS {
if let Some(join_result) = join_set.join_next().await {
@@ -616,8 +633,8 @@ pub fn spawn_multi_user_heartbeat(
if report.had_work() {
tracing::info!(
user_id = uid,
daily_logs_deleted = report.daily_logs_deleted,
conversation_docs_deleted = report.conversation_docs_deleted,
directories_cleaned = ?report.directories_cleaned,
versions_pruned = report.versions_pruned,
"multi-user heartbeat: memory hygiene deleted stale documents"
);
}
+54 -169
View File
@@ -68,8 +68,8 @@ impl Session {
}
/// Create a new thread in this session.
pub fn create_thread(&mut self, channel: Option<&str>) -> &mut Thread {
let thread = Thread::new(self.id, channel);
pub fn create_thread(&mut self) -> &mut Thread {
let thread = Thread::new(self.id);
let thread_id = thread.id;
self.active_thread = Some(thread_id);
self.last_active_at = Utc::now();
@@ -87,9 +87,9 @@ impl Session {
}
/// Get or create the active thread.
pub fn get_or_create_thread(&mut self, channel: Option<&str>) -> &mut Thread {
pub fn get_or_create_thread(&mut self) -> &mut Thread {
match self.active_thread {
None => self.create_thread(channel),
None => self.create_thread(),
Some(id) => {
if self.threads.contains_key(&id) {
// Entry existence confirmed by contains_key above.
@@ -100,7 +100,7 @@ impl Session {
} else {
// Stale active_thread ID: create a new thread, which
// updates self.active_thread to the new thread's ID.
self.create_thread(channel)
self.create_thread()
}
}
}
@@ -225,9 +225,6 @@ pub struct Thread {
/// Messages queued while the thread was processing a turn.
#[serde(default, skip_serializing_if = "VecDeque::is_empty")]
pub pending_messages: VecDeque<String>,
/// Channel that created this thread (for approval authorization).
#[serde(default)]
pub source_channel: Option<String>,
}
/// Maximum number of messages that can be queued while a thread is processing.
@@ -236,29 +233,9 @@ pub struct Thread {
/// rapid follow-ups. The drain loop processes them as one newline-delimited turn.
pub const MAX_PENDING_MESSAGES: usize = 10;
/// Sentinel value for bootstrap threads that accept approvals from any channel.
pub const BOOTSTRAP_SOURCE_CHANNEL: &str = "__bootstrap__";
/// Check whether an approval from `requesting_channel` is authorized for a
/// thread whose `source_channel` is `source`.
///
/// Rules:
/// - `None` (unknown origin) -> denied (fail-closed)
/// - `Some("__bootstrap__")` -> authorized from any channel
/// - `Some(src) == requesting` -> same channel, authorized
/// - requesting is "web" or "gateway" -> always authorized (trusted UI)
/// - Otherwise -> denied
pub fn is_approval_authorized(source: Option<&str>, requesting: &str) -> bool {
match source {
None => false,
Some(src) if src == BOOTSTRAP_SOURCE_CHANNEL => true,
Some(src) => src == requesting || requesting == "web" || requesting == "gateway",
}
}
impl Thread {
/// Create a new thread.
pub fn new(session_id: Uuid, source_channel: Option<&str>) -> Self {
pub fn new(session_id: Uuid) -> Self {
let now = Utc::now();
Self {
id: Uuid::new_v4(),
@@ -271,12 +248,11 @@ impl Thread {
pending_approval: None,
pending_auth: None,
pending_messages: VecDeque::new(),
source_channel: source_channel.map(String::from),
}
}
/// Create a thread with a specific ID (for DB hydration).
pub fn with_id(id: Uuid, session_id: Uuid, source_channel: Option<&str>) -> Self {
pub fn with_id(id: Uuid, session_id: Uuid) -> Self {
let now = Utc::now();
Self {
id,
@@ -289,7 +265,6 @@ impl Thread {
pending_approval: None,
pending_auth: None,
pending_messages: VecDeque::new(),
source_channel: source_channel.map(String::from),
}
}
@@ -812,13 +787,13 @@ mod tests {
let mut session = Session::new("user-123");
assert!(session.active_thread.is_none());
session.create_thread(None);
session.create_thread();
assert!(session.active_thread.is_some());
}
#[test]
fn test_thread_turns() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Hello");
assert_eq!(thread.state, ThreadState::Processing);
@@ -831,7 +806,7 @@ mod tests {
#[test]
fn test_thread_messages() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("First message");
thread.complete_turn("First response");
@@ -854,7 +829,7 @@ mod tests {
#[test]
fn test_restore_from_messages() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// First add some turns
thread.start_turn("Original message");
@@ -880,7 +855,7 @@ mod tests {
#[test]
fn test_restore_from_messages_incomplete_turn() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Messages with incomplete last turn (no assistant response)
let messages = vec![
@@ -899,7 +874,7 @@ mod tests {
#[test]
fn test_enter_auth_mode() {
let before = Utc::now();
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
assert!(thread.pending_auth.is_none());
thread.enter_auth_mode("telegram".to_string());
@@ -912,7 +887,7 @@ mod tests {
#[test]
fn test_take_pending_auth() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.enter_auth_mode("notion".to_string());
let pending = thread.take_pending_auth();
@@ -927,7 +902,7 @@ mod tests {
#[test]
fn test_pending_auth_serialization() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.enter_auth_mode("openai".to_string());
let json = serde_json::to_string(&thread).expect("should serialize");
@@ -957,7 +932,7 @@ mod tests {
#[test]
fn test_pending_auth_default_none() {
// Deserialization of old data without pending_auth should default to None
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.pending_auth = None;
let json = serde_json::to_string(&thread).expect("serialize");
@@ -971,7 +946,7 @@ mod tests {
fn test_thread_with_id() {
let specific_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let thread = Thread::with_id(specific_id, session_id, None);
let thread = Thread::with_id(specific_id, session_id);
assert_eq!(thread.id, specific_id);
assert_eq!(thread.session_id, session_id);
@@ -983,7 +958,7 @@ mod tests {
fn test_thread_with_id_restore_messages() {
let thread_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id, None);
let mut thread = Thread::with_id(thread_id, session_id);
let messages = vec![
ChatMessage::user("Hello from DB"),
@@ -1002,7 +977,7 @@ mod tests {
#[test]
fn test_restore_from_messages_empty() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Add a turn first, then restore with empty vec
thread.start_turn("hello");
@@ -1018,7 +993,7 @@ mod tests {
#[test]
fn test_restore_from_messages_only_assistant_messages() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Only assistant messages (no user messages to anchor turns)
let messages = vec![
@@ -1035,7 +1010,7 @@ mod tests {
#[test]
fn test_restore_from_messages_multiple_user_messages_in_a_row() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Two user messages with no assistant response between them
let messages = vec![
@@ -1062,8 +1037,8 @@ mod tests {
fn test_thread_switch() {
let mut session = Session::new("user-1");
let t1_id = session.create_thread(None).id;
let t2_id = session.create_thread(None).id;
let t1_id = session.create_thread().id;
let t2_id = session.create_thread().id;
// After creating two threads, active should be the last one
assert_eq!(session.active_thread, Some(t2_id));
@@ -1083,8 +1058,8 @@ mod tests {
fn test_get_or_create_thread_idempotent() {
let mut session = Session::new("user-1");
let tid1 = session.get_or_create_thread(None).id;
let tid2 = session.get_or_create_thread(None).id;
let tid1 = session.get_or_create_thread().id;
let tid2 = session.get_or_create_thread().id;
// Should return the same thread (not create a new one each time)
assert_eq!(tid1, tid2);
@@ -1093,7 +1068,7 @@ mod tests {
#[test]
fn test_truncate_turns() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
for i in 0..5 {
thread.start_turn(format!("msg-{}", i));
@@ -1117,7 +1092,7 @@ mod tests {
#[test]
fn test_truncate_turns_noop_when_fewer() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("only one");
thread.complete_turn("response");
@@ -1129,7 +1104,7 @@ mod tests {
#[test]
fn test_thread_interrupt_and_resume() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("do something");
assert_eq!(thread.state, ThreadState::Processing);
@@ -1147,7 +1122,7 @@ mod tests {
#[test]
fn test_resume_only_from_interrupted() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Idle thread: resume should be a no-op
assert_eq!(thread.state, ThreadState::Idle);
@@ -1163,7 +1138,7 @@ mod tests {
#[test]
fn test_turn_fail() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("risky operation");
thread.fail_turn("connection timed out");
@@ -1179,7 +1154,7 @@ mod tests {
#[test]
fn test_messages_with_incomplete_last_turn() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("first");
thread.complete_turn("first reply");
@@ -1195,7 +1170,7 @@ mod tests {
#[test]
fn test_thread_serialization_round_trip() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("hello");
thread.complete_turn("world");
@@ -1213,7 +1188,7 @@ mod tests {
#[test]
fn test_session_serialization_round_trip() {
let mut session = Session::new("user-ser");
session.create_thread(None);
session.create_thread();
session.auto_approve_tool("echo");
let json = serde_json::to_string(&session).unwrap();
@@ -1251,7 +1226,7 @@ mod tests {
#[test]
fn test_turn_number_increments() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Before any turns, turn_number() is 1 (1-indexed for display)
assert_eq!(thread.turn_number(), 1);
@@ -1266,7 +1241,7 @@ mod tests {
#[test]
fn test_complete_turn_on_empty_thread() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Completing a turn when there are no turns should be a safe no-op
thread.complete_turn("phantom response");
@@ -1276,7 +1251,7 @@ mod tests {
#[test]
fn test_fail_turn_on_empty_thread() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Failing a turn when there are no turns should be a safe no-op
thread.fail_turn("phantom error");
@@ -1286,7 +1261,7 @@ mod tests {
#[test]
fn test_pending_approval_flow() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
let approval = PendingApproval {
request_id: Uuid::new_v4(),
@@ -1313,7 +1288,7 @@ mod tests {
#[test]
fn test_clear_pending_approval() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
let approval = PendingApproval {
request_id: Uuid::new_v4(),
@@ -1342,7 +1317,7 @@ mod tests {
assert!(session.active_thread().is_none());
assert!(session.active_thread_mut().is_none());
let tid = session.create_thread(None).id;
let tid = session.create_thread().id;
assert!(session.active_thread().is_some());
assert_eq!(session.active_thread().unwrap().id, tid);
@@ -1359,7 +1334,7 @@ mod tests {
#[test]
fn test_messages_includes_tool_calls() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Search for X");
{
@@ -1391,7 +1366,7 @@ mod tests {
#[test]
fn test_messages_multiple_tool_calls_per_turn() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Do two things");
{
@@ -1418,7 +1393,7 @@ mod tests {
#[test]
fn test_restore_from_messages_with_tool_calls() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Build a message sequence with tool calls
let tc = ToolCall {
@@ -1450,7 +1425,7 @@ mod tests {
#[test]
fn test_restore_from_messages_with_tool_error() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
let tc = ToolCall {
id: "call_0".to_string(),
@@ -1481,7 +1456,7 @@ mod tests {
fn test_messages_round_trip_with_tools() {
// Build a thread with tool calls, get messages(), restore, get messages() again
// The two message sequences should be equivalent.
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Do search");
{
@@ -1494,7 +1469,7 @@ mod tests {
let messages_original = thread.messages();
// Restore into a new thread
let mut thread2 = Thread::new(Uuid::new_v4(), None);
let mut thread2 = Thread::new(Uuid::new_v4());
thread2.restore_from_messages(messages_original.clone());
let messages_restored = thread2.messages();
@@ -1516,7 +1491,7 @@ mod tests {
#[test]
fn test_restore_multi_stage_tool_calls() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
let tc1 = ToolCall {
id: "call_a".to_string(),
@@ -1559,7 +1534,7 @@ mod tests {
#[test]
fn test_messages_truncates_large_tool_results() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("Read big file");
{
@@ -1582,7 +1557,7 @@ mod tests {
#[test]
fn test_thread_message_queue() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Queue is initially empty
assert!(thread.pending_messages.is_empty());
@@ -1618,7 +1593,7 @@ mod tests {
#[test]
fn test_thread_message_queue_serialization() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Empty queue should not appear in serialization (skip_serializing_if)
let json = serde_json::to_string(&thread).unwrap();
@@ -1638,7 +1613,7 @@ mod tests {
#[test]
fn test_thread_message_queue_default_on_old_data() {
// Deserialization of old data without pending_messages should default to empty
let thread = Thread::new(Uuid::new_v4(), None);
let thread = Thread::new(Uuid::new_v4());
let json = serde_json::to_string(&thread).unwrap();
// The field is absent (skip_serializing_if), simulating old data
@@ -1649,7 +1624,7 @@ mod tests {
#[test]
fn test_interrupt_clears_pending_messages() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Start a turn so there's something to interrupt
thread.start_turn("initial input");
@@ -1668,7 +1643,7 @@ mod tests {
#[test]
fn test_thread_state_idle_after_full_drain() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Simulate a full drain cycle: start turn, queue messages, complete turn,
// then drain all queued messages as a single merged turn (#259).
@@ -1696,7 +1671,7 @@ mod tests {
#[test]
fn test_drain_pending_messages_merges_with_newlines() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Empty queue returns None
assert!(thread.drain_pending_messages().is_none());
@@ -1725,7 +1700,7 @@ mod tests {
#[test]
fn test_requeue_drained_preserves_content_at_front() {
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
// Re-queue into empty queue
thread.requeue_drained("failed batch".to_string());
@@ -1836,94 +1811,4 @@ mod tests {
&serde_json::json!("done")
);
}
#[test]
fn test_thread_new_stores_source_channel() {
let thread = Thread::new(Uuid::new_v4(), Some("telegram"));
assert_eq!(thread.source_channel.as_deref(), Some("telegram"));
}
#[test]
fn test_thread_new_none_channel() {
let thread = Thread::new(Uuid::new_v4(), None);
assert!(thread.source_channel.is_none());
}
#[test]
fn test_source_channel_serde_backcompat() {
// Simulate deserializing a Thread from older DB records that lack source_channel.
let thread = Thread::new(Uuid::new_v4(), Some("cli"));
let json = serde_json::to_string(&thread).unwrap();
// Remove the source_channel field to simulate an old record.
let mut value: serde_json::Value = serde_json::from_str(&json).unwrap();
value.as_object_mut().unwrap().remove("source_channel");
let old_json = serde_json::to_string(&value).unwrap();
let deserialized: Thread = serde_json::from_str(&old_json).unwrap();
assert!(
deserialized.source_channel.is_none(),
"missing source_channel should deserialize as None"
);
}
#[test]
fn test_approval_authorized_same_channel() {
assert!(
is_approval_authorized(Some("telegram"), "telegram"),
"same channel should be authorized"
);
}
#[test]
fn test_approval_authorized_different_channel_blocked() {
assert!(
!is_approval_authorized(Some("telegram"), "http"),
"different channel should be blocked"
);
}
#[test]
fn test_approval_authorized_web_always_allowed() {
assert!(
is_approval_authorized(Some("telegram"), "web"),
"web channel should always be authorized"
);
}
#[test]
fn test_approval_authorized_gateway_always_allowed() {
assert!(
is_approval_authorized(Some("telegram"), "gateway"),
"gateway channel should always be authorized"
);
}
#[test]
fn test_approval_authorized_none_denied() {
assert!(
!is_approval_authorized(None, "telegram"),
"None source_channel should be denied (fail-closed)"
);
assert!(
!is_approval_authorized(None, "web"),
"None source_channel should be denied even for web"
);
}
#[test]
fn test_approval_authorized_bootstrap_any_channel() {
assert!(
is_approval_authorized(Some(BOOTSTRAP_SOURCE_CHANNEL), "telegram"),
"__bootstrap__ should be authorized from any channel"
);
assert!(
is_approval_authorized(Some(BOOTSTRAP_SOURCE_CHANNEL), "http"),
"__bootstrap__ should be authorized from any channel"
);
assert!(
is_approval_authorized(Some(BOOTSTRAP_SOURCE_CHANNEL), "cli"),
"__bootstrap__ should be authorized from any channel"
);
}
}
+13 -28
View File
@@ -200,7 +200,7 @@ impl SessionManager {
// Create new thread (always create a new one for a new key)
let thread_id = {
let mut sess = session.lock().await;
let thread = sess.create_thread(Some(channel));
let thread = sess.create_thread();
thread.id
};
@@ -476,7 +476,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-hydrate")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(thread_id, sess.id, None);
let thread = Thread::with_id(thread_id, sess.id);
sess.threads.insert(thread_id, thread);
sess.active_thread = Some(thread_id);
}
@@ -600,7 +600,7 @@ mod tests {
// Simulate hydration: create thread with a known UUID
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_uuid, session_id, None);
let thread = Thread::with_id(known_uuid, session_id);
sess.threads.insert(known_uuid, thread);
}
@@ -627,7 +627,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-idem")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -656,7 +656,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-undo")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -680,7 +680,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-new")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -788,7 +788,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-cross")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -815,7 +815,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-cross")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
@@ -966,7 +966,7 @@ mod tests {
let adopted_id = Uuid::new_v4();
{
let mut sess = session1.lock().await;
let thread = Thread::with_id(adopted_id, sess.id, None);
let thread = Thread::with_id(adopted_id, sess.id);
sess.threads.insert(adopted_id, thread);
}
// Resolve with the UUID as external_thread_id -- should adopt it
@@ -992,7 +992,7 @@ mod tests {
let session = Arc::new(Mutex::new(Session::new("user-direct")));
{
let mut sess = session.lock().await;
let thread = Thread::with_id(tid, sess.id, None);
let thread = Thread::with_id(tid, sess.id);
sess.threads.insert(tid, thread);
}
{
@@ -1030,7 +1030,7 @@ mod tests {
let known_id = Uuid::new_v4();
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_id, sess.id, None);
let thread = Thread::with_id(known_id, sess.id);
sess.threads.insert(known_id, thread);
}
@@ -1057,7 +1057,7 @@ mod tests {
let known_id = Uuid::new_v4();
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_id, sess.id, None);
let thread = Thread::with_id(known_id, sess.id);
sess.threads.insert(known_id, thread);
}
@@ -1081,7 +1081,7 @@ mod tests {
let known_id = Uuid::new_v4();
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_id, sess.id, None);
let thread = Thread::with_id(known_id, sess.id);
sess.threads.insert(known_id, thread);
}
@@ -1102,19 +1102,4 @@ mod tests {
"should NOT adopt UUID when external_thread_id is None"
);
}
#[tokio::test]
async fn test_thread_stores_source_channel() {
let manager = SessionManager::new();
let (session, thread_id) = manager.resolve_thread("user-1", "telegram", None).await;
let sess = session.lock().await;
let thread = sess.threads.get(&thread_id).unwrap();
assert_eq!(
thread.source_channel.as_deref(),
Some("telegram"),
"resolve_thread should store source_channel from the channel parameter"
);
}
}
+9 -25
View File
@@ -135,29 +135,13 @@ impl Agent {
msg_count = 0;
}
// Create thread with the historical ID and restore messages.
// Read source_channel from DB so the authorization check uses the
// original creator's channel, not the requesting message's channel.
let db_source_channel = if let Some(store) = self.store() {
store
.get_conversation_source_channel(thread_uuid)
.await
.unwrap_or(None)
} else {
None
};
let effective_source_channel = db_source_channel.as_deref();
// Create thread with the historical ID and restore messages
let session_id = {
let sess = session.lock().await;
sess.id
};
let mut thread = crate::agent::session::Thread::with_id(
thread_uuid,
session_id,
effective_source_channel,
);
let mut thread = crate::agent::session::Thread::with_id(thread_uuid, session_id);
if !chat_messages.is_empty() {
thread.restore_from_messages(chat_messages);
}
@@ -652,7 +636,7 @@ impl Agent {
user_id: &str,
) -> bool {
match store
.ensure_conversation(thread_id, channel, user_id, None, Some(channel))
.ensure_conversation(thread_id, channel, user_id, None)
.await
{
Ok(true) => true,
@@ -1797,7 +1781,7 @@ impl Agent {
.get_or_create_session(&message.user_id)
.await;
let mut sess = session.lock().await;
let thread = sess.create_thread(Some(&message.channel));
let thread = sess.create_thread();
let thread_id = thread.id;
Ok(SubmissionResult::ok_with_message(format!(
"New thread: {}",
@@ -2133,7 +2117,7 @@ mod tests {
let session_id = Uuid::new_v4();
let thread_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id, None);
let mut thread = Thread::with_id(thread_id, session_id);
// Set thread to AwaitingApproval with a pending tool approval
let pending = PendingApproval {
@@ -2201,7 +2185,7 @@ mod tests {
use crate::agent::session::{MAX_PENDING_MESSAGES, Thread, ThreadState};
use uuid::Uuid;
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("processing something");
assert_eq!(thread.state, ThreadState::Processing);
@@ -2227,7 +2211,7 @@ mod tests {
use crate::agent::session::{Thread, ThreadState};
use uuid::Uuid;
let mut thread = Thread::new(Uuid::new_v4(), None);
let mut thread = Thread::new(Uuid::new_v4());
thread.start_turn("processing");
thread.queue_message("pending-1".to_string());
@@ -2257,7 +2241,7 @@ mod tests {
let thread_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id, None);
let mut thread = Thread::with_id(thread_id, session_id);
thread.start_turn("working");
assert_eq!(thread.state, ThreadState::Processing);
@@ -2284,7 +2268,7 @@ mod tests {
let thread_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let mut thread = Thread::with_id(thread_id, session_id, None);
let mut thread = Thread::with_id(thread_id, session_id);
thread.start_turn("working");
assert_eq!(thread.state, ThreadState::Processing);
-40
View File
@@ -34,7 +34,6 @@ pub async fn setup_wasm_channels(
secrets_store: &Option<Arc<dyn SecretsStore + Send + Sync>>,
extension_manager: Option<&Arc<ExtensionManager>>,
database: Option<&Arc<dyn Database>>,
registered_channel_names: &[String],
) -> Option<WasmChannelSetup> {
let runtime = match WasmChannelRuntime::new(WasmChannelRuntimeConfig::default()) {
Ok(r) => Arc::new(r),
@@ -72,46 +71,7 @@ pub async fn setup_wasm_channels(
let mut channels: Vec<(String, Box<dyn crate::channels::Channel>)> = Vec::new();
let mut channel_names: Vec<String> = Vec::new();
// Reserved channel names that WASM modules must not claim.
// A malicious module could otherwise register as a trusted built-in
// channel and bypass cross-channel authorization checks.
// This list must cover every built-in channel name to prevent a WASM
// module from impersonating a built-in and satisfying same-channel
// approval checks.
const RESERVED_CHANNEL_NAMES: &[&str] = &[
"web",
"gateway",
"cli",
"repl",
"http",
"signal",
"slack-relay",
"secret_save",
];
for loaded in results.loaded {
let name_lower = loaded.name().to_ascii_lowercase();
if RESERVED_CHANNEL_NAMES.contains(&name_lower.as_str()) {
tracing::warn!(
channel = %loaded.name(),
"Rejected WASM channel with reserved name"
);
continue;
}
// Also reject any name that collides with an already-registered
// channel to prevent a WASM module from shadowing a channel that
// was registered earlier in the startup sequence.
if registered_channel_names
.iter()
.any(|n| n.to_ascii_lowercase() == name_lower)
{
tracing::warn!(
channel = %loaded.name(),
"Rejected WASM channel that collides with already-registered channel"
);
continue;
}
let (name, channel) = register_channel(
loaded,
config,
+2 -8
View File
@@ -574,7 +574,7 @@ pub async fn chat_new_thread_handler(
.await;
let (thread_id, info) = {
let mut sess = session.lock().await;
let thread = sess.create_thread(Some("web"));
let thread = sess.create_thread();
let id = thread.id;
let info = ThreadInfo {
id: thread.id,
@@ -593,13 +593,7 @@ pub async fn chat_new_thread_handler(
// so that the subsequent loadThreads() call from the frontend sees it.
if let Some(ref store) = state.store {
match store
.ensure_conversation(
thread_id,
"gateway",
&identity.user_id,
None,
Some("gateway"),
)
.ensure_conversation(thread_id, "gateway", &identity.user_id, None)
.await
{
Ok(true) => {}
+2 -2
View File
@@ -2014,7 +2014,7 @@ async fn chat_new_thread_handler(
let session = session_manager.get_or_create_session(&user.user_id).await;
let (thread_id, info) = {
let mut sess = session.lock().await;
let thread = sess.create_thread(Some("gateway"));
let thread = sess.create_thread();
let id = thread.id;
let info = ThreadInfo {
id: thread.id,
@@ -2033,7 +2033,7 @@ async fn chat_new_thread_handler(
// so that the subsequent loadThreads() call from the frontend sees it.
if let Some(ref store) = state.store {
match store
.ensure_conversation(thread_id, "gateway", &user.user_id, None, Some("gateway"))
.ensure_conversation(thread_id, "gateway", &user.user_id, None)
.await
{
Ok(true) => {}
+5 -13
View File
@@ -10,10 +10,8 @@ use crate::error::ConfigError;
pub struct HygieneConfig {
/// Whether hygiene is enabled. Env: `MEMORY_HYGIENE_ENABLED` (default: true).
pub enabled: bool,
/// Days before `daily/` documents are deleted. Env: `MEMORY_HYGIENE_DAILY_RETENTION_DAYS` (default: 30).
pub daily_retention_days: u32,
/// Days before `conversations/` documents are deleted. Env: `MEMORY_HYGIENE_CONVERSATION_RETENTION_DAYS` (default: 7).
pub conversation_retention_days: u32,
/// Maximum versions to keep per document. Env: `MEMORY_HYGIENE_VERSION_KEEP_COUNT` (default: 50).
pub version_keep_count: u32,
/// Minimum hours between hygiene passes. Env: `MEMORY_HYGIENE_CADENCE_HOURS` (default: 12).
pub cadence_hours: u32,
}
@@ -22,8 +20,7 @@ impl Default for HygieneConfig {
fn default() -> Self {
Self {
enabled: true,
daily_retention_days: 30,
conversation_retention_days: 7,
version_keep_count: 50,
cadence_hours: 12,
}
}
@@ -33,11 +30,7 @@ impl HygieneConfig {
pub(crate) fn resolve() -> Result<Self, ConfigError> {
Ok(Self {
enabled: parse_bool_env("MEMORY_HYGIENE_ENABLED", true)?,
daily_retention_days: parse_optional_env("MEMORY_HYGIENE_DAILY_RETENTION_DAYS", 30)?,
conversation_retention_days: parse_optional_env(
"MEMORY_HYGIENE_CONVERSATION_RETENTION_DAYS",
7,
)?,
version_keep_count: parse_optional_env("MEMORY_HYGIENE_VERSION_KEEP_COUNT", 50)?,
cadence_hours: parse_optional_env("MEMORY_HYGIENE_CADENCE_HOURS", 12)?,
})
}
@@ -47,8 +40,7 @@ impl HygieneConfig {
pub fn to_workspace_config(&self) -> crate::workspace::hygiene::HygieneConfig {
crate::workspace::hygiene::HygieneConfig {
enabled: self.enabled,
daily_retention_days: self.daily_retention_days,
conversation_retention_days: self.conversation_retention_days,
version_keep_count: self.version_keep_count,
cadence_hours: self.cadence_hours,
state_dir: ironclaw_base_dir(),
}
+3 -26
View File
@@ -67,20 +67,19 @@ impl ConversationStore for LibSqlBackend {
channel: &str,
user_id: &str,
thread_id: Option<&str>,
source_channel: Option<&str>,
) -> Result<bool, DatabaseError> {
let conn = self.connect().await?;
let now = fmt_ts(&Utc::now());
let affected = conn
.execute(
r#"
INSERT INTO conversations (id, channel, user_id, thread_id, source_channel, started_at, last_activity)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?6)
INSERT INTO conversations (id, channel, user_id, thread_id, started_at, last_activity)
VALUES (?1, ?2, ?3, ?4, ?5, ?5)
ON CONFLICT (id) DO UPDATE SET last_activity = excluded.last_activity
WHERE conversations.user_id = excluded.user_id
AND conversations.channel = excluded.channel
"#,
params![id.to_string(), channel, user_id, opt_text(thread_id), opt_text(source_channel), now],
params![id.to_string(), channel, user_id, opt_text(thread_id), now],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
@@ -566,28 +565,6 @@ impl ConversationStore for LibSqlBackend {
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(found.is_some())
}
async fn get_conversation_source_channel(
&self,
conversation_id: Uuid,
) -> Result<Option<String>, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
"SELECT source_channel FROM conversations WHERE id = ?1",
params![conversation_id.to_string()],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
match rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
Some(row) => Ok(get_opt_text(&row, 0)),
None => Ok(None),
}
}
}
#[cfg(test)]
+1 -1
View File
@@ -981,7 +981,7 @@ mod tests {
assert!(alice_stats.last_active_at.is_some());
// Bob has no LLM calls so doesn't appear in summary stats
assert!(!stats.iter().any(|s| s.user_id == "bob"));
assert!(stats.iter().find(|s| s.user_id == "bob").is_none());
// Filter to single user
let alice_only = db.user_summary_stats(Some("alice")).await.unwrap();
+326 -2
View File
@@ -13,8 +13,8 @@ use super::{
use crate::db::WorkspaceStore;
use crate::error::{DatabaseError, WorkspaceError};
use crate::workspace::{
MemoryChunk, MemoryDocument, RankedResult, SearchConfig, SearchResult, WorkspaceEntry,
fuse_results,
DocumentVersion, MemoryChunk, MemoryDocument, RankedResult, SearchConfig, SearchResult,
VersionSummary, WorkspaceEntry, fuse_results,
};
use chrono::Utc;
@@ -840,6 +840,330 @@ impl WorkspaceStore for LibSqlBackend {
Ok(fuse_results(fts_results, vector_results, config))
}
// ==================== Metadata ====================
async fn update_document_metadata(
&self,
id: Uuid,
metadata: &serde_json::Value,
) -> Result<(), WorkspaceError> {
let conn = self
.connect()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: e.to_string(),
})?;
let now = fmt_ts(&Utc::now());
let meta_str =
serde_json::to_string(metadata).map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to serialize metadata: {e}"),
})?;
conn.execute(
"UPDATE memory_documents SET metadata = ?2, updated_at = ?3 WHERE id = ?1",
params![id.to_string(), meta_str, now],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to update metadata: {e}"),
})?;
Ok(())
}
async fn find_config_documents(
&self,
user_id: &str,
agent_id: Option<Uuid>,
) -> Result<Vec<MemoryDocument>, WorkspaceError> {
let conn = self
.connect()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: e.to_string(),
})?;
let agent_str = agent_id.map(|a| a.to_string());
let mut rows = conn
.query(
r#"
SELECT id, user_id, agent_id, path, content,
created_at, updated_at, metadata
FROM memory_documents
WHERE user_id = ?1 AND agent_id IS ?2
AND (path LIKE '%/.config' OR path = '.config')
ORDER BY path
"#,
params![user_id, agent_str],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to find config documents: {e}"),
})?;
let mut docs = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to read config document row: {e}"),
})?
{
docs.push(row_to_memory_document(&row));
}
Ok(docs)
}
// ==================== Versioning ====================
async fn save_version(
&self,
document_id: Uuid,
content: &str,
content_hash: &str,
changed_by: Option<&str>,
) -> Result<i32, WorkspaceError> {
let conn = self
.connect()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: e.to_string(),
})?;
let id = Uuid::new_v4().to_string();
let doc_id = document_id.to_string();
let now = fmt_ts(&Utc::now());
// Use a transaction to prevent race conditions: the SELECT and INSERT
// must be atomic so concurrent writers don't allocate the same version.
let tx = conn
.transaction()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to start transaction: {e}"),
})?;
// Get next version number (inside transaction — serializes writers)
let mut rows = tx
.query(
"SELECT COALESCE(MAX(version), 0) + 1 FROM memory_document_versions WHERE document_id = ?1",
params![doc_id.clone()],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to get next version number: {e}"),
})?;
let next_version = if let Some(row) =
rows.next()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to read version number: {e}"),
})? {
get_i64(&row, 0) as i32
} else {
1
};
drop(rows);
tx.execute(
r#"
INSERT INTO memory_document_versions
(id, document_id, version, content, content_hash, created_at, changed_by)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)
"#,
params![
id,
doc_id,
next_version as i64,
content,
content_hash,
now,
changed_by
],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to save version: {e}"),
})?;
tx.commit()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to commit version: {e}"),
})?;
Ok(next_version)
}
async fn get_version(
&self,
document_id: Uuid,
version: i32,
) -> Result<DocumentVersion, WorkspaceError> {
let conn = self
.connect()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: e.to_string(),
})?;
let mut rows = conn
.query(
r#"
SELECT id, document_id, version, content, content_hash,
created_at, changed_by
FROM memory_document_versions
WHERE document_id = ?1 AND version = ?2
"#,
params![document_id.to_string(), version as i64],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to get version: {e}"),
})?;
let row = rows
.next()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to read version row: {e}"),
})?
.ok_or(WorkspaceError::VersionNotFound {
document_id,
version,
})?;
Ok(DocumentVersion {
id: get_text(&row, 0)
.parse()
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Invalid version UUID: {e}"),
})?,
document_id: get_text(&row, 1)
.parse()
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Invalid document UUID: {e}"),
})?,
version: get_i64(&row, 2) as i32,
content: get_text(&row, 3),
content_hash: get_text(&row, 4),
created_at: get_ts(&row, 5),
changed_by: get_opt_text(&row, 6),
})
}
async fn list_versions(
&self,
document_id: Uuid,
limit: i64,
) -> Result<Vec<VersionSummary>, WorkspaceError> {
let conn = self
.connect()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: e.to_string(),
})?;
let mut rows = conn
.query(
r#"
SELECT version, content_hash, created_at, changed_by
FROM memory_document_versions
WHERE document_id = ?1
ORDER BY version DESC
LIMIT ?2
"#,
params![document_id.to_string(), limit],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to list versions: {e}"),
})?;
let mut versions = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to read version row: {e}"),
})?
{
versions.push(VersionSummary {
version: get_i64(&row, 0) as i32,
content_hash: get_text(&row, 1),
created_at: get_ts(&row, 2),
changed_by: get_opt_text(&row, 3),
});
}
Ok(versions)
}
async fn get_latest_version_number(
&self,
document_id: Uuid,
) -> Result<Option<i32>, WorkspaceError> {
let conn = self
.connect()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: e.to_string(),
})?;
let mut rows = conn
.query(
"SELECT MAX(version) FROM memory_document_versions WHERE document_id = ?1",
params![document_id.to_string()],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to get latest version number: {e}"),
})?;
if let Some(row) = rows
.next()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to read version number: {e}"),
})?
{
// MAX returns NULL if no rows — libsql returns Null for the value
let val = row.get::<libsql::Value>(0).ok();
match val {
Some(libsql::Value::Integer(v)) => Ok(Some(v as i32)),
_ => Ok(None),
}
} else {
Ok(None)
}
}
async fn prune_versions(
&self,
document_id: Uuid,
keep_count: i32,
) -> Result<u64, WorkspaceError> {
let conn = self
.connect()
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: e.to_string(),
})?;
let doc_id = document_id.to_string();
let result = conn
.execute(
r#"
DELETE FROM memory_document_versions
WHERE document_id = ?1
AND version NOT IN (
SELECT version FROM memory_document_versions
WHERE document_id = ?1
ORDER BY version DESC
LIMIT ?2
)
"#,
params![doc_id, keep_count as i64],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to prune versions: {e}"),
})?;
Ok(result)
}
}
#[cfg(test)]
+14 -3
View File
@@ -789,10 +789,21 @@ CREATE INDEX IF NOT EXISTS idx_api_tokens_hash ON api_tokens(token_hash);
),
(
15,
"conversation_source_channel",
// Add source_channel to conversations for cross-channel approval authorization.
"document_versions",
r#"
ALTER TABLE conversations ADD COLUMN source_channel TEXT;
CREATE TABLE IF NOT EXISTS memory_document_versions (
id TEXT PRIMARY KEY,
document_id TEXT NOT NULL REFERENCES memory_documents(id) ON DELETE CASCADE,
version INTEGER NOT NULL,
content TEXT NOT NULL,
content_hash TEXT NOT NULL,
created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now')),
changed_by TEXT,
UNIQUE(document_id, version)
);
CREATE INDEX IF NOT EXISTS idx_doc_versions_lookup
ON memory_document_versions(document_id, version DESC);
"#,
),
];
+61 -6
View File
@@ -373,7 +373,6 @@ pub trait ConversationStore: Send + Sync {
channel: &str,
user_id: &str,
thread_id: Option<&str>,
source_channel: Option<&str>,
) -> Result<bool, DatabaseError>;
async fn list_conversations_with_preview(
&self,
@@ -432,11 +431,6 @@ pub trait ConversationStore: Send + Sync {
conversation_id: Uuid,
user_id: &str,
) -> Result<bool, DatabaseError>;
/// Get the source_channel for a conversation (the channel that created it).
async fn get_conversation_source_channel(
&self,
conversation_id: Uuid,
) -> Result<Option<String>, DatabaseError>;
}
#[async_trait]
@@ -706,6 +700,67 @@ pub trait WorkspaceStore: Send + Sync {
config: &SearchConfig,
) -> Result<Vec<SearchResult>, WorkspaceError>;
// ==================== Metadata ====================
/// Update the metadata JSON field on a document (full replacement).
async fn update_document_metadata(
&self,
id: Uuid,
metadata: &serde_json::Value,
) -> Result<(), WorkspaceError>;
/// Find all `.config` documents in the workspace.
///
/// Returns documents whose path ends with `/.config` or equals `.config`.
/// Used by the hygiene system to discover metadata-driven cleanup targets.
async fn find_config_documents(
&self,
user_id: &str,
agent_id: Option<Uuid>,
) -> Result<Vec<MemoryDocument>, WorkspaceError>;
// ==================== Versioning ====================
/// Save the current content of a document as a new version.
///
/// Returns the new version number (1-based, monotonically increasing).
async fn save_version(
&self,
document_id: Uuid,
content: &str,
content_hash: &str,
changed_by: Option<&str>,
) -> Result<i32, WorkspaceError>;
/// Get a specific version of a document.
async fn get_version(
&self,
document_id: Uuid,
version: i32,
) -> Result<crate::workspace::DocumentVersion, WorkspaceError>;
/// List versions of a document (newest first).
async fn list_versions(
&self,
document_id: Uuid,
limit: i64,
) -> Result<Vec<crate::workspace::VersionSummary>, WorkspaceError>;
/// Get the latest version number for a document, or `None` if no versions exist.
async fn get_latest_version_number(
&self,
document_id: Uuid,
) -> Result<Option<i32>, WorkspaceError>;
/// Delete old versions, keeping only the most recent `keep_count`.
///
/// Returns the number of versions deleted.
async fn prune_versions(
&self,
document_id: Uuid,
keep_count: i32,
) -> Result<u64, WorkspaceError>;
// ==================== Multi-scope read methods ====================
//
// Default implementations loop over user_ids calling single-scope methods,
+66 -12
View File
@@ -25,7 +25,8 @@ use crate::history::{
LlmCallRecord, SandboxJobRecord, SandboxJobSummary, SettingRow, Store,
};
use crate::workspace::{
MemoryChunk, MemoryDocument, Repository, SearchConfig, SearchResult, WorkspaceEntry,
DocumentVersion, MemoryChunk, MemoryDocument, Repository, SearchConfig, SearchResult,
VersionSummary, WorkspaceEntry,
};
/// PostgreSQL database backend.
@@ -99,10 +100,9 @@ impl ConversationStore for PgBackend {
channel: &str,
user_id: &str,
thread_id: Option<&str>,
source_channel: Option<&str>,
) -> Result<bool, DatabaseError> {
self.store
.ensure_conversation(id, channel, user_id, thread_id, source_channel)
.ensure_conversation(id, channel, user_id, thread_id)
.await
}
@@ -213,15 +213,6 @@ impl ConversationStore for PgBackend {
.conversation_belongs_to_user(conversation_id, user_id)
.await
}
async fn get_conversation_source_channel(
&self,
conversation_id: Uuid,
) -> Result<Option<String>, DatabaseError> {
self.store
.get_conversation_source_channel(conversation_id)
.await
}
}
// ==================== JobStore ====================
@@ -795,6 +786,69 @@ impl WorkspaceStore for PgBackend {
.list_directory_multi(user_ids, agent_id, directory)
.await
}
// ==================== Metadata ====================
async fn update_document_metadata(
&self,
id: Uuid,
metadata: &serde_json::Value,
) -> Result<(), WorkspaceError> {
self.repo.update_document_metadata(id, metadata).await
}
async fn find_config_documents(
&self,
user_id: &str,
agent_id: Option<Uuid>,
) -> Result<Vec<MemoryDocument>, WorkspaceError> {
self.repo.find_config_documents(user_id, agent_id).await
}
// ==================== Versioning ====================
async fn save_version(
&self,
document_id: Uuid,
content: &str,
content_hash: &str,
changed_by: Option<&str>,
) -> Result<i32, WorkspaceError> {
self.repo
.save_version(document_id, content, content_hash, changed_by)
.await
}
async fn get_version(
&self,
document_id: Uuid,
version: i32,
) -> Result<DocumentVersion, WorkspaceError> {
self.repo.get_version(document_id, version).await
}
async fn list_versions(
&self,
document_id: Uuid,
limit: i64,
) -> Result<Vec<VersionSummary>, WorkspaceError> {
self.repo.list_versions(document_id, limit).await
}
async fn get_latest_version_number(
&self,
document_id: Uuid,
) -> Result<Option<i32>, WorkspaceError> {
self.repo.get_latest_version_number(document_id).await
}
async fn prune_versions(
&self,
document_id: Uuid,
keep_count: i32,
) -> Result<u64, WorkspaceError> {
self.repo.prune_versions(document_id, keep_count).await
}
}
// ==================== UserStore ====================
+6
View File
@@ -315,6 +315,12 @@ pub enum WorkspaceError {
#[error("Write rejected for '{path}': prompt injection detected ({reason})")]
InjectionRejected { path: String, reason: String },
#[error("Version not found: document {document_id} version {version}")]
VersionNotFound { document_id: Uuid, version: i32 },
#[error("Patch failed for '{path}': {reason}")]
PatchFailed { path: String, reason: String },
}
/// Orchestrator errors (internal API, container management).
+3 -19
View File
@@ -1581,20 +1581,19 @@ impl Store {
channel: &str,
user_id: &str,
thread_id: Option<&str>,
source_channel: Option<&str>,
) -> Result<bool, DatabaseError> {
let conn = self.conn().await?;
let affected = conn
.execute(
r#"
INSERT INTO conversations (id, channel, user_id, thread_id, source_channel)
VALUES ($1, $2, $3, $4, $5)
INSERT INTO conversations (id, channel, user_id, thread_id)
VALUES ($1, $2, $3, $4)
ON CONFLICT (id) DO UPDATE
SET last_activity = NOW()
WHERE conversations.user_id = EXCLUDED.user_id
AND conversations.channel = EXCLUDED.channel
"#,
&[&id, &channel, &user_id, &thread_id, &source_channel],
&[&id, &channel, &user_id, &thread_id],
)
.await?;
Ok(affected > 0)
@@ -1893,21 +1892,6 @@ impl Store {
Ok(row.is_some())
}
/// Get the source_channel for a conversation.
pub async fn get_conversation_source_channel(
&self,
conversation_id: Uuid,
) -> Result<Option<String>, DatabaseError> {
let conn = self.conn().await?;
let row = conn
.query_opt(
"SELECT source_channel FROM conversations WHERE id = $1",
&[&conversation_id],
)
.await?;
Ok(row.and_then(|r| r.get::<_, Option<String>>(0)))
}
/// Load messages for a conversation with cursor-based pagination.
///
/// Returns `(messages_oldest_first, has_more)`.
+1
View File
@@ -234,6 +234,7 @@ fn is_transient(err: &LlmError) -> bool {
LlmError::RequestFailed { .. }
| LlmError::RateLimited { .. }
| LlmError::InvalidResponse { .. }
| LlmError::EmptyResponse { .. }
| LlmError::SessionExpired { .. }
| LlmError::SessionRenewalFailed { .. }
| LlmError::Http(_)
+3
View File
@@ -17,6 +17,9 @@ pub enum LlmError {
#[error("Invalid response from {provider}: {reason}")]
InvalidResponse { provider: String, reason: String },
#[error("Empty response from {provider}: no content returned")]
EmptyResponse { provider: String },
#[error("Context length exceeded: {used} tokens used, {limit} allowed")]
ContextLengthExceeded { used: usize, limit: usize },
+2 -4
View File
@@ -231,9 +231,8 @@ impl LlmProvider for GithubCopilotProvider {
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
.ok_or_else(|| LlmError::EmptyResponse {
provider: "github_copilot".to_string(),
reason: "No choices in response".to_string(),
})?;
let (content, _tool_calls) = extract_choice_content(&choice);
@@ -309,9 +308,8 @@ impl LlmProvider for GithubCopilotProvider {
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
.ok_or_else(|| LlmError::EmptyResponse {
provider: "github_copilot".to_string(),
reason: "No choices in response".to_string(),
})?;
let (content, tool_calls) = extract_choice_content(&choice);
+2 -4
View File
@@ -490,9 +490,8 @@ impl LlmProvider for NearAiChatProvider {
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
.ok_or_else(|| LlmError::EmptyResponse {
provider: "nearai_chat".to_string(),
reason: "No choices in response".to_string(),
})?;
// Fall back to reasoning_content when content is null (same as
@@ -570,9 +569,8 @@ impl LlmProvider for NearAiChatProvider {
.choices
.into_iter()
.next()
.ok_or_else(|| LlmError::InvalidResponse {
.ok_or_else(|| LlmError::EmptyResponse {
provider: "nearai_chat".to_string(),
reason: "No choices in response".to_string(),
})?;
let tool_calls: Vec<ToolCall> = choice
+1
View File
@@ -48,6 +48,7 @@ pub(crate) fn is_retryable(err: &LlmError) -> bool {
LlmError::RequestFailed { .. }
| LlmError::RateLimited { .. }
| LlmError::InvalidResponse { .. }
| LlmError::EmptyResponse { .. }
| LlmError::SessionRenewalFailed { .. }
| LlmError::Http(_)
| LlmError::Io(_)
-1
View File
@@ -449,7 +449,6 @@ async fn async_main() -> anyhow::Result<()> {
&components.secrets_store,
components.extension_manager.as_ref(),
components.db.as_ref(),
&channel_names,
)
.await;
+1 -1
View File
@@ -278,7 +278,7 @@ impl TenantScope {
thread_id: Option<&str>,
) -> Result<bool, DatabaseError> {
self.inner
.ensure_conversation(id, channel, &self.user_id, thread_id, None)
.ensure_conversation(id, channel, &self.user_id, thread_id)
.await
}
+3 -3
View File
@@ -753,7 +753,7 @@ mod tests {
// ensure_conversation should create the row.
assert!(
db.ensure_conversation(conv_id, "web", "carol", None, Some("web"))
db.ensure_conversation(conv_id, "web", "carol", None)
.await
.expect("ensure first"),
"first ensure_conversation should create the row"
@@ -761,7 +761,7 @@ mod tests {
// Calling again with the same ID should not error.
assert!(
db.ensure_conversation(conv_id, "web", "carol", None, Some("web"))
db.ensure_conversation(conv_id, "web", "carol", None)
.await
.expect("ensure second (idempotent)"),
"second ensure_conversation should touch owned row"
@@ -806,7 +806,7 @@ mod tests {
tokio::time::sleep(std::time::Duration::from_millis(25)).await;
assert!(
!db.ensure_conversation(conv_id, "web", "mallory", None, None)
!db.ensure_conversation(conv_id, "web", "mallory", None)
.await
.expect("foreign ensure should not error"),
"foreign ensure_conversation should report not ensured"
+140 -4
View File
@@ -246,9 +246,26 @@ impl Tool for MemoryWriteTool {
"type": "boolean",
"description": "Skip privacy classification and write directly to the specified layer without redirect. Use when you're certain the content belongs in the target layer.",
"default": false
},
"metadata": {
"type": "object",
"description": "Optional metadata to set on the document (e.g., {\"skip_indexing\": true, \"hygiene\": {\"enabled\": true, \"retention_days\": 7}})"
},
"old_string": {
"type": "string",
"description": "When present, switches to patch mode: finds and replaces this exact string in the document. Requires target to be a path (not 'memory' or 'daily_log')."
},
"new_string": {
"type": "string",
"description": "Replacement string (required when old_string is present)."
},
"replace_all": {
"type": "boolean",
"description": "If true, replace all occurrences of old_string. Default: false.",
"default": false
}
},
"required": ["content"]
"required": []
})
}
@@ -259,7 +276,9 @@ impl Tool for MemoryWriteTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let content = require_str(&params, "content")?;
// In patch mode (old_string present), content is not required.
let is_patch_mode = params.get("old_string").and_then(|v| v.as_str()).is_some();
let content = params.get("content").and_then(|v| v.as_str()).unwrap_or("");
let target = params
.get("target")
@@ -299,9 +318,9 @@ impl Tool for MemoryWriteTool {
return Ok(ToolOutput::success(output, start.elapsed()));
}
if content.trim().is_empty() {
if !is_patch_mode && content.trim().is_empty() {
return Err(ToolError::InvalidParameters(
"content cannot be empty".to_string(),
"content cannot be empty (use old_string/new_string for patch mode)".to_string(),
));
}
@@ -330,6 +349,46 @@ impl Tool for MemoryWriteTool {
path => path.to_string(),
};
// Patch mode: if old_string is provided, do search-and-replace instead of write/append.
let old_string = params.get("old_string").and_then(|v| v.as_str());
if let Some(old_str) = old_string {
let new_str = params
.get("new_string")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters(
"new_string is required when old_string is provided".to_string(),
)
})?;
let replace_all = params
.get("replace_all")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let result = workspace
.patch(&resolved_path, old_str, new_str, replace_all)
.await
.map_err(map_write_err)?;
// Apply metadata if provided
if let Some(meta) = params.get("metadata")
&& meta.is_object()
{
workspace
.update_metadata(result.document.id, meta)
.await
.map_err(map_write_err)?;
}
let output = serde_json::json!({
"status": "patched",
"path": resolved_path,
"replacements": result.replacements,
"content_length": result.document.content.len(),
});
return Ok(ToolOutput::success(output, start.elapsed()));
}
// When a layer is specified, route through layer-aware methods for ALL targets.
// Otherwise, use default workspace methods (which include injection scanning).
let layer_result = if let Some(layer_name) = layer {
@@ -433,6 +492,24 @@ impl Tool for MemoryWriteTool {
}
}
// Apply metadata if provided (after write/append, works for all targets).
// We read the document once to get its ID — this is a hot read right
// after the write, so it's effectively free (same DB connection/cache).
if let Some(meta) = params.get("metadata")
&& meta.is_object()
{
match workspace.read(&resolved_path).await {
Ok(doc) => {
if let Err(e) = workspace.update_metadata(doc.id, meta).await {
tracing::warn!(path = %resolved_path, "failed to update metadata: {e}");
}
}
Err(e) => {
tracing::warn!(path = %resolved_path, "failed to read doc for metadata update: {e}");
}
}
}
let mut output = serde_json::json!({
"status": "written",
"path": resolved_path,
@@ -501,6 +578,15 @@ impl Tool for MemoryReadTool {
"path": {
"type": "string",
"description": "Path to the file (e.g., 'MEMORY.md', 'daily/2024-01-15.md', 'projects/alpha/notes.md')"
},
"version": {
"type": "integer",
"description": "Read a specific historical version of the document (omit for current content)"
},
"list_versions": {
"type": "boolean",
"description": "If true, return version history instead of file content",
"default": false
}
},
"required": ["path"]
@@ -525,11 +611,61 @@ impl Tool for MemoryReadTool {
}
let workspace = self.resolver.resolve(&ctx.user_id).await;
let list_versions = params
.get("list_versions")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let version = params
.get("version")
.and_then(|v| v.as_i64())
.map(|v| v as i32);
// Read the document first (needed for document_id in all version operations)
let doc = workspace
.read(path)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Read failed: {}", e)))?;
// List versions mode
if list_versions {
let versions = workspace
.list_versions(doc.id, 50)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("List versions failed: {}", e)))?;
let output = serde_json::json!({
"path": doc.path,
"versions": versions.iter().map(|v| serde_json::json!({
"version": v.version,
"content_hash": v.content_hash,
"created_at": v.created_at.to_rfc3339(),
"changed_by": v.changed_by,
})).collect::<Vec<_>>(),
"version_count": versions.len(),
});
return Ok(ToolOutput::success(output, start.elapsed()));
}
// Specific version mode
if let Some(ver) = version {
let version_doc = workspace
.get_version(doc.id, ver)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Get version failed: {}", e)))?;
let output = serde_json::json!({
"path": doc.path,
"version": version_doc.version,
"content": version_doc.content,
"content_hash": version_doc.content_hash,
"created_at": version_doc.created_at.to_rfc3339(),
"changed_by": version_doc.changed_by,
});
return Ok(ToolOutput::success(output, start.elapsed()));
}
// Normal read
let output = serde_json::json!({
"path": doc.path,
"content": doc.content,
+144 -36
View File
@@ -31,6 +31,7 @@ use std::time::Duration;
use serde::{Deserialize, Serialize};
use tokio::io::{AsyncBufReadExt, BufReader};
#[cfg(not(unix))]
use tokio::process::Command;
use uuid::Uuid;
@@ -340,6 +341,11 @@ impl ClaudeBridgeRuntime {
/// Spawn a `claude` CLI process and stream its output.
///
/// Uses a PTY on Unix so Node.js line-buffers stdout instead of
/// full-buffering (which causes the bridge to hang on non-TTY pipes).
/// Arguments are passed via `execve` (no shell) — injection-safe by
/// construction.
///
/// Returns the session_id if captured from the `system` init message.
async fn run_claude_session(
&self,
@@ -347,47 +353,102 @@ impl ClaudeBridgeRuntime {
resume_session_id: Option<&str>,
extra_env: &std::collections::HashMap<String, String>,
) -> Result<Option<String>, WorkerError> {
let mut cmd = Command::new("claude");
cmd.arg("-p")
.arg(prompt)
.arg("--output-format")
.arg("stream-json")
.arg("--verbose")
.arg("--max-turns")
.arg(self.config.max_turns.to_string())
.arg("--model")
.arg(&self.config.model);
let max_turns_str = self.config.max_turns.to_string();
if let Some(sid) = resume_session_id {
cmd.arg("--resume").arg(sid);
}
// Inject credentials into the child process environment without
// mutating the global process env (which is unsafe in multi-threaded programs).
cmd.envs(extra_env);
cmd.current_dir("/workspace")
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped());
let mut child = cmd.spawn().map_err(|e| WorkerError::ExecutionFailed {
reason: format!("failed to spawn claude: {}", e),
})?;
let stdout = child
.stdout
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stdout".to_string(),
// Spawn with PTY on Unix to fix Node.js stdout buffering.
// All arguments are passed individually via execve — never through
// a shell interpreter. This eliminates shell injection by construction.
#[cfg(unix)]
let (mut child, stdout, stderr) = {
let (pty, pts) = pty_process::open().map_err(|e| WorkerError::ExecutionFailed {
reason: format!("failed to allocate PTY: {}", e),
})?;
let stderr = child
.stderr
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stderr".to_string(),
let mut cmd = pty_process::Command::new("claude");
cmd = cmd
.arg("-p")
.arg(prompt)
.arg("--output-format")
.arg("stream-json")
.arg("--verbose")
.arg("--max-turns")
.arg(&max_turns_str)
.arg("--model")
.arg(&self.config.model);
if let Some(sid) = resume_session_id {
cmd = cmd.arg("--resume").arg(sid);
}
cmd = cmd.envs(extra_env.iter());
cmd = cmd.current_dir("/workspace");
// Keep stderr on a separate pipe — pty-process attaches the PTY
// to all fds by default, which would merge stderr into the PTY
// stream and break NDJSON parsing.
cmd = cmd.stderr(std::process::Stdio::piped());
let mut child = cmd.spawn(pts).map_err(|e| WorkerError::ExecutionFailed {
reason: format!("failed to spawn claude with PTY: {}", e),
})?;
let stderr = child
.stderr
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stderr".to_string(),
})?;
// stdout comes from the PTY master, which implements AsyncRead
let stdout: Box<dyn tokio::io::AsyncRead + Unpin + Send> = Box::new(pty);
(child, stdout, stderr)
};
// Non-Unix fallback (Windows CI) — no PTY, direct spawn.
// Claude bridge only runs in Linux Docker containers, so this path
// exists solely for compilation on Windows targets.
#[cfg(not(unix))]
let (mut child, stdout, stderr) = {
let mut cmd = Command::new("claude");
cmd.arg("-p")
.arg(prompt)
.arg("--output-format")
.arg("stream-json")
.arg("--verbose")
.arg("--max-turns")
.arg(&max_turns_str)
.arg("--model")
.arg(&self.config.model);
if let Some(sid) = resume_session_id {
cmd.arg("--resume").arg(sid);
}
cmd.envs(extra_env);
cmd.current_dir("/workspace")
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped());
let mut child = cmd.spawn().map_err(|e| WorkerError::ExecutionFailed {
reason: format!("failed to spawn claude: {}", e),
})?;
let stdout_pipe = child
.stdout
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stdout".to_string(),
})?;
let stderr = child
.stderr
.take()
.ok_or_else(|| WorkerError::ExecutionFailed {
reason: "failed to capture claude stderr".to_string(),
})?;
let stdout: Box<dyn tokio::io::AsyncRead + Unpin + Send> = Box::new(stdout_pipe);
(child, stdout, stderr)
};
// Spawn stderr reader that forwards lines as log events
let client_for_stderr = Arc::clone(&self.client);
let job_id = self.config.job_id;
@@ -1027,4 +1088,51 @@ mod tests {
let copied = copy_dir_recursive(nonexistent, dst.path()).unwrap();
assert_eq!(copied, 0);
}
/// Regression test: arguments are passed individually (not via shell string),
/// so shell metacharacters in prompt/model/session_id are harmless.
#[test]
fn command_args_no_shell_interpretation() {
// Prompt, model, and session_id may contain shell metacharacters from
// user-supplied task descriptions or LLM output. Since we use
// Command::arg() (execve), these are passed as literal strings.
let prompt = "Fix the user's bug; echo $HOME && rm -rf /";
let model = "claude-3-opus-20240229";
let session_id = "'; DROP TABLE jobs; --";
let max_turns = 10u32;
let max_turns_str = max_turns.to_string();
let args: Vec<&str> = vec![
"-p",
prompt,
"--output-format",
"stream-json",
"--verbose",
"--max-turns",
&max_turns_str,
"--model",
model,
"--resume",
session_id,
];
// All values present as literal strings — no shell interpretation
// ["-p", prompt, "--output-format", "stream-json", "--verbose",
// "--max-turns", "10", "--model", model, "--resume", session_id]
assert_eq!(args[1], prompt);
assert_eq!(args[8], model);
assert_eq!(args[10], session_id);
// Shell metacharacters preserved, not expanded
assert!(args[1].contains("$HOME"));
assert!(args[1].contains("&&"));
assert!(args[10].contains("'; DROP TABLE"));
}
/// Verify PTY is available on Unix platforms.
#[cfg(unix)]
#[tokio::test]
async fn pty_opens_successfully() {
let result = pty_process::open();
assert!(result.is_ok(), "PTY allocation should succeed on Unix");
}
}
+148 -4
View File
@@ -391,6 +391,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
worker: self,
rx: tokio::sync::Mutex::new(rx),
consecutive_rate_limits: std::sync::atomic::AtomicUsize::new(0),
has_text_response: std::sync::atomic::AtomicBool::new(false),
};
let config = AgenticLoopConfig {
@@ -1101,6 +1102,15 @@ fn store_fallback_in_metadata(
}
/// Job delegate: implements `LoopDelegate` for the background job context.
/// Whether an LLM error represents a completion-eligible empty response.
///
/// Only `EmptyResponse` (provider returned no choices/content) qualifies.
/// Infrastructure errors (`AuthFailed`, `Http`, `Io`, etc.) never qualify —
/// they must propagate even if prior text output was produced.
fn is_completion_eligible_error(error: &crate::error::LlmError) -> bool {
matches!(error, crate::error::LlmError::EmptyResponse { .. })
}
///
/// Handles: signal channel (stop/ping/user messages), cancellation checks,
/// rate-limit retry, parallel tool execution, DB persistence, SSE broadcasting.
@@ -1109,6 +1119,10 @@ struct JobDelegate<'a> {
rx: tokio::sync::Mutex<&'a mut mpsc::Receiver<WorkerMessage>>,
/// Tracks consecutive rate-limit errors to fail fast instead of burning iterations.
consecutive_rate_limits: std::sync::atomic::AtomicUsize,
/// Whether a substantive (non-empty) text response has been produced.
/// When true, an empty follow-up response is treated as job completion
/// rather than a retry signal (prevents spurious failures in routines).
has_text_response: std::sync::atomic::AtomicBool,
}
impl<'a> JobDelegate<'a> {
@@ -1161,6 +1175,53 @@ impl<'a> JobDelegate<'a> {
finish_reason: crate::llm::FinishReason::Stop,
})
}
/// Mark the job as completed, logging a warning on failure.
async fn mark_completed_or_warn(&self, context: &str) {
if let Err(e) = self.worker.mark_completed().await {
tracing::warn!(
job_id = %self.worker.job_id,
error = %e,
"Failed to mark job completed ({context})"
);
}
}
/// If a substantive text response was already produced and the error
/// indicates the LLM simply returned nothing, treat it as successful
/// completion rather than a fatal failure.
///
/// Only swallows `EmptyResponse` — infrastructure errors (`AuthFailed`,
/// `ContextLengthExceeded`, `Http`, `Io`, etc.) always propagate.
///
/// Returns `Some(empty RespondOutput)` when the error should be swallowed,
/// `None` when it should propagate normally.
async fn try_complete_on_error(
&self,
context: &str,
error: &crate::error::LlmError,
) -> Option<crate::llm::RespondOutput> {
if !is_completion_eligible_error(error) {
return None;
}
if !self
.has_text_response
.load(std::sync::atomic::Ordering::Relaxed)
{
return None;
}
tracing::info!(
job_id = %self.worker.job_id,
error = %error,
"{context} empty response after text output — treating as completion"
);
self.mark_completed_or_warn(context).await;
Some(crate::llm::RespondOutput {
result: RespondResult::Text(String::new()),
usage: crate::llm::TokenUsage::default(),
finish_reason: crate::llm::FinishReason::Stop,
})
}
}
#[async_trait]
@@ -1291,7 +1352,12 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
Err(crate::error::LlmError::RateLimited { retry_after, .. }) => {
return self.handle_rate_limit(retry_after, "tool selection").await;
}
Err(e) => return Err(e.into()),
Err(e) => {
if let Some(output) = self.try_complete_on_error("select_tools", &e).await {
return Ok(output);
}
return Err(e.into());
}
};
// Fall back to respond_with_tools
@@ -1321,7 +1387,12 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
self.handle_rate_limit(retry_after, "respond_with_tools")
.await
}
Err(e) => Err(e.into()),
Err(e) => {
if let Some(output) = self.try_complete_on_error("respond_with_tools", &e).await {
return Ok(output);
}
Err(e.into())
}
}
}
@@ -1330,9 +1401,22 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
text: &str,
reason_ctx: &mut ReasoningContext,
) -> TextAction {
// Empty text from rate-limit backoff retry — skip processing and let the
// loop proceed to the next iteration which will re-call the LLM.
// Empty text after a substantive response means the LLM has finished.
// Treat as successful completion rather than continuing the loop (which
// would produce "Response contained no message or tool call (empty)").
if text.is_empty() {
if self
.has_text_response
.load(std::sync::atomic::Ordering::Relaxed)
{
tracing::debug!(
job_id = %self.worker.job_id,
"Empty response after text output — treating as completion"
);
self.mark_completed_or_warn("empty text response").await;
return TextAction::Return(LoopOutcome::Response(String::new()));
}
// No prior text response — this is likely a rate-limit backoff retry.
return TextAction::Continue;
}
@@ -1348,6 +1432,10 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
return TextAction::Return(LoopOutcome::Response(text.to_string()));
}
// Track that a substantive response has been produced.
self.has_text_response
.store(true, std::sync::atomic::Ordering::Relaxed);
// Add assistant response to context
reason_ctx.messages.push(ChatMessage::assistant(text));
@@ -2285,4 +2373,60 @@ mod tests {
assert_eq!(telegram[0].0, "owner-scope");
assert_eq!(telegram[0].1.content, "hello from routine");
}
/// Regression test: only `EmptyResponse` errors are eligible for
/// completion-swallowing. Infrastructure errors must always propagate.
#[test]
fn is_completion_eligible_only_matches_empty_response() {
use crate::error::LlmError;
// EmptyResponse is eligible
assert!(super::is_completion_eligible_error(
&LlmError::EmptyResponse {
provider: "test".to_string(),
}
));
// All other variants are NOT eligible
assert!(!super::is_completion_eligible_error(
&LlmError::InvalidResponse {
provider: "test".to_string(),
reason: "parse error".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::AuthFailed {
provider: "test".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::ContextLengthExceeded {
used: 100_000,
limit: 50_000,
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::ModelNotAvailable {
provider: "test".to_string(),
model: "gpt-4".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::RequestFailed {
provider: "test".to_string(),
reason: "timeout".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::SessionExpired {
provider: "test".to_string(),
}
));
assert!(!super::is_completion_eligible_error(
&LlmError::SessionRenewalFailed {
provider: "test".to_string(),
reason: "timeout".to_string(),
}
));
}
}
+248
View File
@@ -2,6 +2,7 @@
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use uuid::Uuid;
/// Well-known document paths.
@@ -37,6 +38,139 @@ pub mod paths {
pub const ASSISTANT_DIRECTIVES: &str = "context/assistant-directives.md";
}
/// Name of the folder-level configuration document.
///
/// A document at `{directory}/.config` carries metadata flags that apply
/// as defaults to all documents in that directory (e.g., `skip_indexing`,
/// `hygiene` settings). Individual document metadata overrides folder defaults.
pub const CONFIG_FILE_NAME: &str = ".config";
/// Typed overlay for the `metadata` JSON field on [`MemoryDocument`].
///
/// Fields use `Option` so that only explicitly set flags participate in
/// the merge chain (document metadata → folder `.config` → system defaults).
/// Unknown fields are preserved via `serde(flatten)`.
#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
pub struct DocumentMetadata {
/// When `true`, skip chunking and embedding for this document/folder.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub skip_indexing: Option<bool>,
/// When `true`, skip automatic versioning for this document/folder.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub skip_versioning: Option<bool>,
/// Hygiene (auto-cleanup) configuration for this folder.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub hygiene: Option<HygieneMetadata>,
/// Preserve unknown fields for forward compatibility.
#[serde(flatten)]
pub extra: serde_json::Map<String, serde_json::Value>,
}
impl DocumentMetadata {
/// Parse from a raw JSON [`serde_json::Value`].
///
/// Returns [`Default`] if the value is not an object or cannot be parsed.
pub fn from_value(value: &serde_json::Value) -> Self {
serde_json::from_value(value.clone()).unwrap_or_default()
}
/// Convert to a JSON [`serde_json::Value`].
pub fn to_value(&self) -> serde_json::Value {
serde_json::to_value(self).unwrap_or(serde_json::json!({}))
}
/// Merge two metadata values: `overlay` keys win over `base` keys.
///
/// This is a shallow merge at the top-level keys — nested objects are
/// replaced wholesale, not recursively merged. This keeps the semantics
/// simple and predictable across both PostgreSQL and libSQL.
pub fn merge(base: &serde_json::Value, overlay: &serde_json::Value) -> serde_json::Value {
let mut merged = match base {
serde_json::Value::Object(map) => map.clone(),
_ => serde_json::Map::new(),
};
if let serde_json::Value::Object(over) = overlay {
for (k, v) in over {
merged.insert(k.clone(), v.clone());
}
}
serde_json::Value::Object(merged)
}
}
/// Hygiene (auto-cleanup) settings for a folder.
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct HygieneMetadata {
/// Whether this folder is a hygiene target.
pub enabled: bool,
/// Delete documents older than this many days.
#[serde(default = "default_retention_days")]
pub retention_days: u32,
}
fn default_retention_days() -> u32 {
30
}
/// A historical version of a workspace document.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DocumentVersion {
/// Version record ID.
pub id: Uuid,
/// Parent document ID.
pub document_id: Uuid,
/// Version number (1-based, monotonically increasing per document).
pub version: i32,
/// Full document content at this version.
pub content: String,
/// SHA-256 hash of `content` (hex-encoded, prefixed with `sha256:`).
pub content_hash: String,
/// When this version was created.
pub created_at: DateTime<Utc>,
/// Who/what created this version (e.g. `"agent"`, `"user:alice"`).
pub changed_by: Option<String>,
}
/// Summary of a document version (without full content).
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VersionSummary {
/// Version number.
pub version: i32,
/// SHA-256 hash of the version's content.
pub content_hash: String,
/// When this version was created.
pub created_at: DateTime<Utc>,
/// Who/what created this version.
pub changed_by: Option<String>,
}
/// Result of a workspace patch operation.
#[derive(Debug, Clone)]
pub struct PatchResult {
/// The updated document.
pub document: MemoryDocument,
/// Number of replacements made.
pub replacements: usize,
}
/// Compute a SHA-256 hash of content, returned as `"sha256:{hex}"`.
pub fn content_sha256(content: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(content.as_bytes());
let result = hasher.finalize();
format!("sha256:{:x}", result)
}
/// Check if a path refers to a `.config` document.
pub fn is_config_path(path: &str) -> bool {
let file_name = path.rsplit('/').next().unwrap_or(path);
file_name == CONFIG_FILE_NAME
}
/// Paths treated as identity documents for multi-scope isolation.
///
/// These files are always read from the primary scope only — never from
@@ -360,6 +494,120 @@ mod tests {
assert_eq!(result[0].updated_at, Some(ts));
}
#[test]
fn test_document_metadata_default_is_empty() {
let meta = DocumentMetadata::default();
assert_eq!(meta.skip_indexing, None);
assert_eq!(meta.skip_versioning, None);
assert_eq!(meta.hygiene, None);
assert!(meta.extra.is_empty());
}
#[test]
fn test_document_metadata_from_value_full() {
let value = serde_json::json!({
"skip_indexing": true,
"skip_versioning": false,
"hygiene": { "enabled": true, "retention_days": 7 }
});
let meta = DocumentMetadata::from_value(&value);
assert_eq!(meta.skip_indexing, Some(true));
assert_eq!(meta.skip_versioning, Some(false));
let hygiene = meta.hygiene.unwrap();
assert!(hygiene.enabled);
assert_eq!(hygiene.retention_days, 7);
}
#[test]
fn test_document_metadata_from_value_partial() {
let value = serde_json::json!({"skip_indexing": true});
let meta = DocumentMetadata::from_value(&value);
assert_eq!(meta.skip_indexing, Some(true));
assert_eq!(meta.hygiene, None);
}
#[test]
fn test_document_metadata_from_value_invalid() {
let meta = DocumentMetadata::from_value(&serde_json::json!("not an object"));
assert_eq!(meta, DocumentMetadata::default());
}
#[test]
fn test_document_metadata_preserves_unknown_fields() {
let value = serde_json::json!({
"skip_indexing": true,
"custom_field": "hello"
});
let meta = DocumentMetadata::from_value(&value);
assert_eq!(meta.skip_indexing, Some(true));
assert_eq!(
meta.extra.get("custom_field").and_then(|v| v.as_str()),
Some("hello")
);
// Round-trip preserves the field
let back = meta.to_value();
assert_eq!(
back.get("custom_field").and_then(|v| v.as_str()),
Some("hello")
);
}
#[test]
fn test_document_metadata_merge() {
let base = serde_json::json!({"skip_indexing": false, "hygiene": {"enabled": true, "retention_days": 30}});
let overlay = serde_json::json!({"skip_indexing": true, "skip_versioning": true});
let merged = DocumentMetadata::merge(&base, &overlay);
let meta = DocumentMetadata::from_value(&merged);
// Overlay wins
assert_eq!(meta.skip_indexing, Some(true));
assert_eq!(meta.skip_versioning, Some(true));
// Base preserved when not overridden
assert!(meta.hygiene.is_some());
}
#[test]
fn test_document_metadata_merge_empty_base() {
let base = serde_json::json!({});
let overlay = serde_json::json!({"skip_indexing": true});
let merged = DocumentMetadata::merge(&base, &overlay);
let meta = DocumentMetadata::from_value(&merged);
assert_eq!(meta.skip_indexing, Some(true));
}
#[test]
fn test_hygiene_metadata_default_retention() {
let value = serde_json::json!({"enabled": true});
let hygiene: HygieneMetadata = serde_json::from_value(value).unwrap();
assert!(hygiene.enabled);
assert_eq!(hygiene.retention_days, 30);
}
#[test]
fn test_content_sha256_deterministic() {
let hash1 = content_sha256("hello world");
let hash2 = content_sha256("hello world");
assert_eq!(hash1, hash2);
assert!(hash1.starts_with("sha256:"));
}
#[test]
fn test_content_sha256_different_content() {
let hash1 = content_sha256("hello");
let hash2 = content_sha256("world");
assert_ne!(hash1, hash2);
}
#[test]
fn test_is_config_path() {
assert!(is_config_path(".config"));
assert!(is_config_path("daily/.config"));
assert!(is_config_path("frontend/widgets/.config"));
assert!(!is_config_path("daily/2024-01-15.md"));
assert!(!is_config_path("MEMORY.md"));
assert!(!is_config_path(".config.bak"));
}
#[test]
fn test_merge_workspace_entries_sorted_by_path() {
let entries = vec![
+194 -258
View File
@@ -1,8 +1,10 @@
//! Memory hygiene: automatic cleanup of stale workspace documents.
//!
//! Runs on a configurable cadence and deletes daily log entries and conversation
//! documents older than their respective retention periods. Identity files
//! (`IDENTITY.md`, `SOUL.md`, etc.) are never touched.
//! Runs on a configurable cadence and discovers which directories have hygiene
//! enabled by reading `.config` metadata documents. This is a **metadata-driven**
//! approach: instead of hardcoding `daily/` and `conversations/`, the system
//! respects `hygiene.enabled` and `hygiene.retention_days` set on each folder's
//! `.config` document.
//!
//! A global [`AtomicBool`] guard prevents concurrent hygiene passes, which
//! avoids TOCTOU races on the state file and Windows file-locking errors
@@ -10,18 +12,16 @@
//! pass completes.
//!
//! ```text
//! ┌─────────────────────────────────────────────┐
//! │ Hygiene Pass │
//! │ │
//! │ 0. Acquire RUNNING guard (skip if held) │
//! │ 1. Check cadence (skip if ran recently) │
//! │ 2. Save state (claim the cadence window) │
//! │ 3. List daily/ documents
//! │ 4. Delete those older than daily_retention
//! │ 5. List conversations/ documents
//! │ 6. Delete those older than conversation_ret │
//! │ 7. Log summary │
//! └─────────────────────────────────────────────┘
//! ┌──────────────────────────────────────────────────
//! │ Hygiene Pass
//! │
//! │ 0. Acquire RUNNING guard (skip if held)
//! │ 1. Check cadence (skip if ran recently)
//! │ 2. Save state (claim the cadence window)
//! │ 3. Discover .config docs with hygiene.enabled
//! │ 4. For each: cleanup_directory(parent, retention)
//! │ 5. Log summary
//! └──────────────────────────────────────────────────┘
//! ```
use std::path::PathBuf;
@@ -31,46 +31,22 @@ use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use crate::bootstrap::ironclaw_base_dir;
use crate::workspace::Workspace;
use crate::workspace::{DocumentMetadata, Workspace, is_config_path};
/// Global guard preventing concurrent hygiene passes.
static RUNNING: AtomicBool = AtomicBool::new(false);
/// Paths that must never be deleted by hygiene, regardless of age.
const IDENTITY_PATHS: &[&str] = &[
crate::workspace::document::paths::MEMORY,
crate::workspace::document::paths::IDENTITY,
crate::workspace::document::paths::SOUL,
crate::workspace::document::paths::AGENTS,
crate::workspace::document::paths::USER,
crate::workspace::document::paths::HEARTBEAT,
crate::workspace::document::paths::README,
crate::workspace::document::paths::TOOLS,
crate::workspace::document::paths::BOOTSTRAP,
];
/// Check if a document path is an identity document that must never be deleted.
///
/// Performs case-insensitive comparison to handle case-insensitive filesystems
/// (Windows, macOS) and prevent accidental deletion of identity docs with
/// different casing (e.g., memory.md, MEMORY.MD, Memory.md).
fn is_identity_path(path: &str) -> bool {
let file_name = path.rsplit('/').next().unwrap_or(path);
let file_name_lower = file_name.to_lowercase();
IDENTITY_PATHS
.iter()
.any(|&p| p.to_lowercase() == file_name_lower)
}
/// Configuration for workspace hygiene.
#[derive(Debug, Clone)]
pub struct HygieneConfig {
/// Whether hygiene is enabled at all.
pub enabled: bool,
/// Documents in `daily/` older than this many days are deleted.
pub daily_retention_days: u32,
/// Documents in `conversations/` older than this many days are deleted.
pub conversation_retention_days: u32,
/// Maximum number of versions to keep per document.
///
/// TODO: Wire up global version pruning once per-document iteration
/// is efficient (e.g., via a dedicated DB query). For now this field
/// is stored in config but not actively enforced during hygiene passes.
pub version_keep_count: u32,
/// Minimum hours between hygiene passes.
pub cadence_hours: u32,
/// Directory to store state file (default: `~/.ironclaw`).
@@ -81,8 +57,7 @@ impl Default for HygieneConfig {
fn default() -> Self {
Self {
enabled: true,
daily_retention_days: 30,
conversation_retention_days: 7,
version_keep_count: 50,
cadence_hours: 12,
state_dir: ironclaw_base_dir(),
}
@@ -98,10 +73,10 @@ struct HygieneState {
/// Summary of what a hygiene pass cleaned up.
#[derive(Debug, Default)]
pub struct HygieneReport {
/// Number of daily log documents deleted.
pub daily_logs_deleted: u32,
/// Number of conversation documents deleted.
pub conversation_docs_deleted: u32,
/// Per-directory cleanup results: `(directory_path, deleted_count)`.
pub directories_cleaned: Vec<(String, u32)>,
/// Number of document versions pruned across all documents.
pub versions_pruned: u64,
/// Whether the run was skipped (cadence not yet elapsed).
pub skipped: bool,
}
@@ -109,7 +84,7 @@ pub struct HygieneReport {
impl HygieneReport {
/// True if any cleanup work was done.
pub fn had_work(&self) -> bool {
self.daily_logs_deleted > 0 || self.conversation_docs_deleted > 0
self.directories_cleaned.iter().any(|(_, n)| *n > 0) || self.versions_pruned > 0
}
}
@@ -168,30 +143,51 @@ pub async fn run_if_due(workspace: &Workspace, config: &HygieneConfig) -> Hygien
// TOCTOU races where another task reads stale state.
save_state(&state_file);
tracing::info!(
daily_retention_days = config.daily_retention_days,
conversation_retention_days = config.conversation_retention_days,
"memory hygiene: starting cleanup pass"
);
tracing::info!("memory hygiene: starting cleanup pass");
let mut report = HygieneReport::default();
// Delete old daily logs
match cleanup_daily_logs(workspace, config.daily_retention_days).await {
Ok(count) => report.daily_logs_deleted = count,
Err(e) => tracing::warn!("memory hygiene: failed to clean daily logs: {e}"),
}
// Discover directories that have hygiene enabled via .config metadata.
let config_docs = match workspace.find_config_documents().await {
Ok(docs) => docs,
Err(e) => {
tracing::warn!("memory hygiene: failed to discover .config documents: {e}");
return report;
}
};
// Delete old conversation documents
match cleanup_conversation_docs(workspace, config.conversation_retention_days).await {
Ok(count) => report.conversation_docs_deleted = count,
Err(e) => tracing::warn!("memory hygiene: failed to clean conversation docs: {e}"),
for doc in &config_docs {
let meta = DocumentMetadata::from_value(&doc.metadata);
let Some(hygiene) = meta.hygiene else {
continue;
};
if !hygiene.enabled {
continue;
}
// Derive the parent directory from the .config path.
let directory = match doc.path.rsplit_once('/') {
Some((dir, _)) => format!("{dir}/"),
None => continue, // root-level .config — skip
};
match cleanup_directory(workspace, &directory, hygiene.retention_days).await {
Ok(deleted) => {
if deleted > 0 {
tracing::info!(directory, deleted, "memory hygiene: cleaned directory");
}
report.directories_cleaned.push((directory, deleted));
}
Err(e) => {
tracing::warn!(directory, "memory hygiene: failed to clean directory: {e}");
}
}
}
if report.had_work() {
tracing::info!(
daily_logs_deleted = report.daily_logs_deleted,
conversation_docs_deleted = report.conversation_docs_deleted,
directories_cleaned = ?report.directories_cleaned,
versions_pruned = report.versions_pruned,
"memory hygiene: cleanup complete"
);
} else {
@@ -210,88 +206,41 @@ impl Drop for RunningGuard {
}
}
/// Delete daily log documents older than `retention_days`.
async fn cleanup_daily_logs(
/// Delete documents in `directory` that are older than `retention_days`.
///
/// Skips directories and `.config` files (which must never be deleted by
/// hygiene). Returns the number of documents deleted.
async fn cleanup_directory(
workspace: &Workspace,
directory: &str,
retention_days: u32,
) -> Result<u32, anyhow::Error> {
let cutoff = Utc::now() - chrono::Duration::days(i64::from(retention_days));
let entries = workspace.list("daily/").await?;
let entries = workspace.list(directory).await?;
let mut deleted = 0u32;
for entry in entries {
if entry.is_directory {
continue;
}
// Never delete identity documents
if is_identity_path(&entry.path) {
if is_config_path(&entry.path) {
continue;
}
// Check if the document is old enough to delete
if let Some(updated_at) = entry.updated_at
&& updated_at < cutoff
{
let path = if entry.path.starts_with("daily/") {
let path = if entry.path.starts_with(directory) {
entry.path.clone()
} else {
format!("daily/{}", entry.path)
format!("{}{}", directory, entry.path)
};
if let Err(e) = workspace.delete(&path).await {
tracing::warn!(path, "memory hygiene: failed to delete: {e}");
} else {
tracing::debug!(path, "memory hygiene: deleted old daily log");
tracing::debug!(path, "memory hygiene: deleted stale document");
deleted += 1;
}
}
}
Ok(deleted)
}
/// Delete conversation documents older than `retention_days`.
async fn cleanup_conversation_docs(
workspace: &Workspace,
retention_days: u32,
) -> Result<u32, anyhow::Error> {
let cutoff = Utc::now() - chrono::Duration::days(i64::from(retention_days));
let entries = workspace.list("conversations/").await?;
let mut deleted = 0u32;
for entry in entries {
if entry.is_directory {
continue;
}
// Never delete identity documents
if is_identity_path(&entry.path) {
continue;
}
// Check if the document is old enough to delete
if let Some(updated_at) = entry.updated_at
&& updated_at < cutoff
{
let path = if entry.path.starts_with("conversations/") {
entry.path.clone()
} else {
format!("conversations/{}", entry.path)
};
if let Err(e) = workspace.delete(&path).await {
tracing::warn!(
path,
"memory hygiene: failed to delete conversation doc: {e}"
);
} else {
tracing::debug!(path, "memory hygiene: deleted old conversation doc");
deleted += 1;
}
}
}
Ok(deleted)
}
@@ -349,8 +298,7 @@ mod tests {
fn default_config_is_reasonable() {
let cfg = HygieneConfig::default();
assert!(cfg.enabled);
assert_eq!(cfg.daily_retention_days, 30);
assert_eq!(cfg.conversation_retention_days, 7);
assert_eq!(cfg.version_keep_count, 50);
assert_eq!(cfg.cadence_hours, 12);
}
@@ -362,84 +310,33 @@ mod tests {
}
#[test]
fn report_had_work_when_deleted() {
fn report_had_work_when_directories_cleaned() {
let report = HygieneReport {
daily_logs_deleted: 3,
conversation_docs_deleted: 0,
directories_cleaned: vec![("daily/".to_string(), 3)],
versions_pruned: 0,
skipped: false,
};
assert!(report.had_work());
}
#[test]
fn report_had_work_when_conversation_deleted() {
fn report_had_work_when_versions_pruned() {
let report = HygieneReport {
daily_logs_deleted: 0,
conversation_docs_deleted: 2,
directories_cleaned: vec![],
versions_pruned: 5,
skipped: false,
};
assert!(report.had_work());
}
#[test]
fn is_identity_path_excludes_sacred_docs() {
for name in [
"MEMORY.md",
"IDENTITY.md",
"SOUL.md",
"AGENTS.md",
"USER.md",
"HEARTBEAT.md",
"README.md",
"TOOLS.md",
"BOOTSTRAP.md",
] {
assert!(is_identity_path(name), "{name} should be excluded");
assert!(
is_identity_path(&format!("conversations/{name}")),
"conversations/{name} should be excluded via path"
);
}
}
#[test]
fn is_identity_path_case_insensitive() {
// Verify case-insensitive matching for case-insensitive filesystems
assert!(
is_identity_path("memory.md"),
"lowercase memory.md should be excluded"
);
assert!(
is_identity_path("Memory.md"),
"mixed case Memory.md should be excluded"
);
assert!(
is_identity_path("MEMORY.MD"),
"uppercase MEMORY.MD should be excluded"
);
assert!(
is_identity_path("identity.md"),
"lowercase identity.md should be excluded"
);
assert!(
is_identity_path("conversations/soul.md"),
"conversations/soul.md should be excluded"
);
assert!(
is_identity_path("conversations/SOUL.MD"),
"conversations/SOUL.MD should be excluded"
);
}
#[test]
fn is_identity_path_allows_normal_docs() {
for path in [
"daily/2024-01-01.md",
"conversations/chat-abc.md",
"notes.md",
] {
assert!(!is_identity_path(path), "{path} should not be excluded");
}
fn report_no_work_when_zero_deletions() {
let report = HygieneReport {
directories_cleaned: vec![("daily/".to_string(), 0)],
versions_pruned: 0,
skipped: false,
};
assert!(!report.had_work());
}
#[test]
@@ -552,61 +449,105 @@ mod tests {
Arc::new(Workspace::new_with_db("default", db.clone()))
}
#[tokio::test]
async fn cleanup_daily_logs_preserves_identity_documents() {
let (db, _tmp) = create_test_db().await;
let ws = create_workspace(&db);
// Write several regular documents (non-identity)
ws.write("daily/2024-01-15.md", "Old log")
/// Helper to seed a .config document with hygiene metadata on a directory.
async fn seed_hygiene_config(workspace: &Workspace, directory: &str, retention_days: u32) {
let config_path = format!("{}.config", directory);
// Create the .config document with empty content
workspace
.write(&config_path, "")
.await
.expect("write log 1");
ws.write("daily/2024-01-20.md", "Another log")
.expect("write .config");
// Read back to get the document ID
let doc = workspace
.read(&config_path)
.await
.expect("write log 2");
// Write an identity document
ws.write("MEMORY.md", "Long-term curated memory")
.expect("read .config doc");
// Set hygiene metadata
workspace
.update_metadata(
doc.id,
&serde_json::json!({
"hygiene": {"enabled": true, "retention_days": retention_days},
"skip_versioning": true
}),
)
.await
.expect("write identity");
// List before cleanup
let before = ws.list("daily/").await.expect("list before");
let daily_count_before = before.iter().filter(|e| !e.is_directory).count();
assert!(daily_count_before >= 2, "should have at least 2 daily logs");
// Run cleanup with 0-day retention (deletes everything old)
// This tests that even with aggressive cleanup, identity docs survive
let deleted = cleanup_daily_logs(&ws, 0)
.await
.expect("cleanup_daily_logs");
// Should have deleted some documents (the daily logs)
assert!(deleted > 0, "should have deleted old daily documents");
// Verify identity doc still exists
let identity = db
.get_document_by_path("default", None, "MEMORY.md")
.await
.expect("get identity doc");
assert_eq!(identity.path, "MEMORY.md");
assert_eq!(identity.content, "Long-term curated memory");
.expect("set metadata");
}
#[tokio::test]
async fn cleanup_conversation_docs_handles_empty_directory() {
async fn cleanup_directory_skips_config_files() {
let (db, _tmp) = create_test_db().await;
let ws = create_workspace(&db);
// Run cleanup on an empty directory (conversations/ doesn't exist)
let deleted = cleanup_conversation_docs(&ws, 7)
// Write documents including a .config
ws.write("daily/2024-01-15.md", "Old log")
.await
.expect("cleanup_conversation_docs");
.expect("write log");
ws.write("daily/.config", "").await.expect("write config");
// Run cleanup with 0-day retention (deletes everything old)
let deleted = cleanup_directory(&ws, "daily/", 0)
.await
.expect("cleanup_directory");
// Should have deleted the log but not the .config
assert!(deleted > 0, "should have deleted old daily documents");
// Verify .config still exists
let config_doc = db
.get_document_by_path("default", None, "daily/.config")
.await
.expect("get .config doc");
assert_eq!(config_doc.path, "daily/.config");
}
#[tokio::test]
async fn cleanup_directory_handles_empty_directory() {
let (db, _tmp) = create_test_db().await;
let ws = create_workspace(&db);
// Run cleanup on an empty directory
let deleted = cleanup_directory(&ws, "conversations/", 7)
.await
.expect("cleanup_directory");
// Should delete 0 (nothing to delete)
assert_eq!(deleted, 0, "should delete 0 from empty directory");
}
#[tokio::test]
async fn metadata_driven_cleanup_discovers_directories() {
let (db, _tmp) = create_test_db().await;
let ws = create_workspace(&db);
// Seed .config with hygiene enabled on daily/
seed_hygiene_config(&ws, "daily/", 0).await;
// Write some documents
ws.write("daily/log1.md", "content 1")
.await
.expect("write doc 1");
ws.write("daily/log2.md", "content 2")
.await
.expect("write doc 2");
let config = HygieneConfig {
enabled: true,
version_keep_count: 50,
cadence_hours: 12,
state_dir: _tmp.path().to_path_buf(),
};
// First run should discover daily/ and clean it
let report = run_if_due(&ws, &config).await;
assert!(!report.skipped, "first run should not be skipped");
assert!(report.had_work(), "should have cleaned documents");
assert!(
!report.directories_cleaned.is_empty(),
"should have at least one directory cleaned"
);
}
#[tokio::test]
async fn cleanup_respects_cadence_prevents_concurrent_runs() {
let (db, _tmp) = create_test_db().await;
@@ -614,8 +555,7 @@ mod tests {
let config = HygieneConfig {
enabled: true,
daily_retention_days: 30,
conversation_retention_days: 7,
version_keep_count: 50,
cadence_hours: 12,
state_dir: _tmp.path().to_path_buf(),
};
@@ -627,13 +567,6 @@ mod tests {
// Second run immediately should be skipped (cadence not elapsed)
let report2 = run_if_due(&ws, &config).await;
assert!(report2.skipped, "second run should be skipped by cadence");
// Report structure should be correct
assert_eq!(
report1.daily_logs_deleted + report1.conversation_docs_deleted,
0,
"first run should have clean counts"
);
}
#[tokio::test]
@@ -641,6 +574,10 @@ mod tests {
let (db, _tmp) = create_test_db().await;
let ws = create_workspace(&db);
// Seed hygiene on both directories
seed_hygiene_config(&ws, "daily/", 0).await;
seed_hygiene_config(&ws, "conversations/", 0).await;
// Write some documents
ws.write("daily/log1.md", "content 1")
.await
@@ -652,35 +589,34 @@ mod tests {
.await
.expect("write doc 3");
// Run with 0-day retention to delete everything non-identity
let deleted_daily = cleanup_daily_logs(&ws, 0).await.expect("cleanup daily");
let deleted_conv = cleanup_conversation_docs(&ws, 0)
// Run with 0-day retention via direct cleanup_directory calls
let deleted_daily = cleanup_directory(&ws, "daily/", 0)
.await
.expect("cleanup daily");
let deleted_conv = cleanup_directory(&ws, "conversations/", 0)
.await
.expect("cleanup conversations");
// Both should report deletions
assert!(deleted_daily > 0, "should report deleted daily logs");
assert_eq!(deleted_conv, 1, "should report 1 deleted conversation doc");
// Create a HygieneReport and verify aggregation works
// Verify HygieneReport aggregation
let report = HygieneReport {
daily_logs_deleted: deleted_daily,
conversation_docs_deleted: deleted_conv,
directories_cleaned: vec![
("daily/".to_string(), deleted_daily),
("conversations/".to_string(), deleted_conv),
],
versions_pruned: 0,
skipped: false,
};
// Verify HygieneReport structure
assert!(!report.skipped, "should not be skipped");
assert!(report.had_work(), "report should indicate work was done");
assert!(
report.daily_logs_deleted > 0 || report.conversation_docs_deleted > 0,
"report should have at least one deletion count > 0"
);
// Verify had_work() correctly combines both counts
// Verify had_work() correctly checks directory counts
let no_work = HygieneReport {
daily_logs_deleted: 0,
conversation_docs_deleted: 0,
directories_cleaned: vec![],
versions_pruned: 0,
skipped: false,
};
assert!(!no_work.had_work(), "empty report should indicate no work");
+351 -2
View File
@@ -53,8 +53,9 @@ mod search;
pub use chunker::{ChunkConfig, chunk_document};
pub use document::{
IDENTITY_PATHS, MemoryChunk, MemoryDocument, WorkspaceEntry, is_identity_path,
merge_workspace_entries, paths,
CONFIG_FILE_NAME, DocumentMetadata, DocumentVersion, HygieneMetadata, IDENTITY_PATHS,
MemoryChunk, MemoryDocument, PatchResult, VersionSummary, WorkspaceEntry, content_sha256,
is_config_path, is_identity_path, merge_workspace_entries, paths,
};
pub use embedding_cache::{CachedEmbeddingProvider, EmbeddingCacheConfig};
pub use embeddings::{
@@ -366,6 +367,101 @@ impl WorkspaceStorage {
}
}
}
// ==================== Metadata ====================
async fn update_document_metadata(
&self,
id: Uuid,
metadata: &serde_json::Value,
) -> Result<(), WorkspaceError> {
match self {
#[cfg(feature = "postgres")]
Self::Repo(repo) => repo.update_document_metadata(id, metadata).await,
Self::Db(db) => db.update_document_metadata(id, metadata).await,
}
}
async fn find_config_documents(
&self,
user_id: &str,
agent_id: Option<Uuid>,
) -> Result<Vec<MemoryDocument>, WorkspaceError> {
match self {
#[cfg(feature = "postgres")]
Self::Repo(repo) => repo.find_config_documents(user_id, agent_id).await,
Self::Db(db) => db.find_config_documents(user_id, agent_id).await,
}
}
// ==================== Versioning ====================
async fn save_version(
&self,
document_id: Uuid,
content: &str,
content_hash: &str,
changed_by: Option<&str>,
) -> Result<i32, WorkspaceError> {
match self {
#[cfg(feature = "postgres")]
Self::Repo(repo) => {
repo.save_version(document_id, content, content_hash, changed_by)
.await
}
Self::Db(db) => {
db.save_version(document_id, content, content_hash, changed_by)
.await
}
}
}
async fn get_version(
&self,
document_id: Uuid,
version: i32,
) -> Result<DocumentVersion, WorkspaceError> {
match self {
#[cfg(feature = "postgres")]
Self::Repo(repo) => repo.get_version(document_id, version).await,
Self::Db(db) => db.get_version(document_id, version).await,
}
}
async fn list_versions(
&self,
document_id: Uuid,
limit: i64,
) -> Result<Vec<VersionSummary>, WorkspaceError> {
match self {
#[cfg(feature = "postgres")]
Self::Repo(repo) => repo.list_versions(document_id, limit).await,
Self::Db(db) => db.list_versions(document_id, limit).await,
}
}
async fn get_latest_version_number(
&self,
document_id: Uuid,
) -> Result<Option<i32>, WorkspaceError> {
match self {
#[cfg(feature = "postgres")]
Self::Repo(repo) => repo.get_latest_version_number(document_id).await,
Self::Db(db) => db.get_latest_version_number(document_id).await,
}
}
async fn prune_versions(
&self,
document_id: Uuid,
keep_count: i32,
) -> Result<u64, WorkspaceError> {
match self {
#[cfg(feature = "postgres")]
Self::Repo(repo) => repo.prune_versions(document_id, keep_count).await,
Self::Db(db) => db.prune_versions(document_id, keep_count).await,
}
}
}
/// Default template seeded into HEARTBEAT.md on first access.
@@ -696,10 +792,199 @@ impl Workspace {
.await
}
// ==================== Metadata ====================
/// Update the metadata JSON on a document by ID (full replacement).
pub async fn update_metadata(
&self,
id: Uuid,
metadata: &serde_json::Value,
) -> Result<(), WorkspaceError> {
self.storage.update_document_metadata(id, metadata).await
}
/// Prune old versions for a document, keeping only the most recent `keep_count`.
///
/// Returns the number of versions deleted.
pub async fn prune_versions(
&self,
document_id: Uuid,
keep_count: i32,
) -> Result<u64, WorkspaceError> {
self.storage.prune_versions(document_id, keep_count).await
}
/// Find all `.config` documents in this workspace scope.
pub async fn find_config_documents(&self) -> Result<Vec<MemoryDocument>, WorkspaceError> {
self.storage
.find_config_documents(&self.user_id, self.agent_id)
.await
}
/// Resolve effective metadata for a document path.
///
/// Resolution chain: document's own metadata → nearest ancestor `.config` → defaults.
pub async fn resolve_metadata(&self, path: &str) -> DocumentMetadata {
// 1. Document's own metadata
let doc_meta = self
.storage
.get_document_by_path(&self.user_id, self.agent_id, path)
.await
.ok()
.map(|d| d.metadata);
// 2. Walk up parent directories looking for .config
let mut config_meta = None;
let normalized = normalize_path(path);
let mut current = normalized.as_str();
while let Some(slash_pos) = current.rfind('/') {
let parent = &current[..slash_pos];
let config_path = format!("{}/{CONFIG_FILE_NAME}", parent);
if let Ok(doc) = self
.storage
.get_document_by_path(&self.user_id, self.agent_id, &config_path)
.await
{
config_meta = Some(doc.metadata);
break;
}
current = parent;
}
// Also check root-level .config
if config_meta.is_none()
&& let Ok(doc) = self
.storage
.get_document_by_path(&self.user_id, self.agent_id, CONFIG_FILE_NAME)
.await
{
config_meta = Some(doc.metadata);
}
// 3. Merge: config as base, document metadata as overlay
let base = config_meta.unwrap_or(serde_json::json!({}));
let overlay = doc_meta.unwrap_or(serde_json::json!({}));
let merged = DocumentMetadata::merge(&base, &overlay);
DocumentMetadata::from_value(&merged)
}
// ==================== Versioning ====================
/// List versions of a document (newest first).
pub async fn list_versions(
&self,
document_id: Uuid,
limit: i64,
) -> Result<Vec<VersionSummary>, WorkspaceError> {
self.storage.list_versions(document_id, limit).await
}
/// Get a specific version of a document.
pub async fn get_version(
&self,
document_id: Uuid,
version: i32,
) -> Result<DocumentVersion, WorkspaceError> {
self.storage.get_version(document_id, version).await
}
/// Save the current content as a version if it differs from the latest.
///
/// Returns the new version number, or `None` if skipped (empty content,
/// identical hash, or versioning disabled via metadata).
async fn maybe_save_version(
&self,
document_id: Uuid,
current_content: &str,
path: &str,
changed_by: Option<&str>,
) -> Result<Option<i32>, WorkspaceError> {
// Don't version empty documents
if current_content.is_empty() {
return Ok(None);
}
// Check metadata for skip_versioning flag
let metadata = self.resolve_metadata(path).await;
if metadata.skip_versioning == Some(true) {
return Ok(None);
}
let hash = content_sha256(current_content);
// Check if latest version already has this hash (skip duplicate saves)
if let Ok(Some(latest)) = self.storage.get_latest_version_number(document_id).await
&& let Ok(ver) = self.storage.get_version(document_id, latest).await
&& ver.content_hash == hash
{
return Ok(None);
}
let version = self
.storage
.save_version(document_id, current_content, &hash, changed_by)
.await?;
Ok(Some(version))
}
// ==================== Patch ====================
/// Apply a search-and-replace patch to a workspace document.
///
/// Finds `old_string` in the document and replaces it with `new_string`.
/// If `replace_all` is true, replaces all occurrences; otherwise only the first.
/// Auto-versions before applying the patch.
pub async fn patch(
&self,
path: &str,
old_string: &str,
new_string: &str,
replace_all: bool,
) -> Result<PatchResult, WorkspaceError> {
let path = normalize_path(path);
let doc = self
.storage
.get_document_by_path(&self.user_id, self.agent_id, &path)
.await?;
if !doc.content.contains(old_string) {
return Err(WorkspaceError::PatchFailed {
path,
reason: "old_string not found in document".to_string(),
});
}
let (new_content, count) = if replace_all {
let count = doc.content.matches(old_string).count();
(doc.content.replace(old_string, new_string), count)
} else {
(doc.content.replacen(old_string, new_string, 1), 1)
};
// Injection scan for system prompt files
if is_system_prompt_file(&path) && !new_content.is_empty() {
reject_if_injected(&path, &new_content)?;
}
// Auto-version before updating
let _ = self
.maybe_save_version(doc.id, &doc.content, &path, None)
.await;
self.storage.update_document(doc.id, &new_content).await?;
self.reindex_document(doc.id).await?;
let updated = self.storage.get_document_by_id(doc.id).await?;
Ok(PatchResult {
document: updated,
replacements: count,
})
}
/// Write (create or update) a file.
///
/// Creates parent directories implicitly (they're virtual in the DB).
/// Re-indexes the document for search after writing.
/// Auto-versions the previous content before overwriting.
///
/// # Example
/// ```ignore
@@ -715,6 +1000,12 @@ impl Workspace {
.storage
.get_or_create_document_by_path(&self.user_id, self.agent_id, &path)
.await?;
// Auto-version previous content before overwriting
let _ = self
.maybe_save_version(doc.id, &doc.content, &path, None)
.await;
self.storage.update_document(doc.id, content).await?;
self.reindex_document(doc.id).await?;
@@ -754,6 +1045,11 @@ impl Workspace {
reject_if_injected(&path, &new_content)?;
}
// Auto-version previous content before appending
let _ = self
.maybe_save_version(doc.id, &doc.content, &path, None)
.await;
self.storage.update_document(doc.id, &new_content).await?;
self.reindex_document(doc.id).await?;
Ok(())
@@ -1585,6 +1881,14 @@ impl Workspace {
// Get the document
let doc = self.storage.get_document_by_id(document_id).await?;
// Check metadata for skip_indexing flag
let metadata = self.resolve_metadata(&doc.path).await;
if metadata.skip_indexing == Some(true) {
// Delete any existing chunks and skip indexing
self.storage.delete_chunks(document_id).await?;
return Ok(());
}
// Chunk the content
let chunks = chunk_document(&doc.content, ChunkConfig::default());
@@ -1670,6 +1974,51 @@ impl Workspace {
}
}
// Seed folder-level .config documents for hygiene defaults.
let config_seeds: &[(&str, serde_json::Value)] = &[
(
"daily/.config",
serde_json::json!({
"hygiene": {"enabled": true, "retention_days": 30},
"skip_versioning": true
}),
),
(
"conversations/.config",
serde_json::json!({
"hygiene": {"enabled": true, "retention_days": 7},
"skip_versioning": true
}),
),
];
for (config_path, metadata_value) in config_seeds {
match self.read_primary(config_path).await {
Ok(_) => continue, // Already exists, don't overwrite
Err(WorkspaceError::DocumentNotFound { .. }) => {}
Err(e) => {
tracing::debug!("Failed to check {}: {}", config_path, e);
continue;
}
}
// Create empty document with metadata
if let Ok(doc) = self
.storage
.get_or_create_document_by_path(&self.user_id, self.agent_id, config_path)
.await
{
if let Err(e) = self
.storage
.update_document_metadata(doc.id, metadata_value)
.await
{
tracing::debug!("Failed to set metadata on {}: {}", config_path, e);
} else {
count += 1;
}
}
}
// BOOTSTRAP.md is only seeded on truly fresh workspaces (no identity
// files existed before seeding) AND when no profile exists yet (the user
// may already have a profile from a previous install and doesn't need
+191 -1
View File
@@ -11,7 +11,9 @@ use uuid::Uuid;
use crate::error::WorkspaceError;
use crate::workspace::document::{MemoryChunk, MemoryDocument, WorkspaceEntry};
use crate::workspace::document::{
DocumentVersion, MemoryChunk, MemoryDocument, VersionSummary, WorkspaceEntry,
};
use crate::workspace::search::{RankedResult, SearchConfig, SearchResult, fuse_results};
/// Database repository for workspace operations.
@@ -702,4 +704,192 @@ impl Repository {
}
Ok(crate::workspace::merge_workspace_entries(all_entries))
}
// ==================== Metadata ====================
pub async fn update_document_metadata(
&self,
id: Uuid,
metadata: &serde_json::Value,
) -> Result<(), WorkspaceError> {
let conn = self.conn().await?;
conn.execute(
"UPDATE memory_documents SET metadata = $2, updated_at = NOW() WHERE id = $1",
&[&id, &metadata],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to update metadata: {e}"),
})?;
Ok(())
}
pub async fn find_config_documents(
&self,
user_id: &str,
agent_id: Option<Uuid>,
) -> Result<Vec<MemoryDocument>, WorkspaceError> {
let conn = self.conn().await?;
let rows = conn
.query(
r#"
SELECT id, user_id, agent_id, path, content,
created_at, updated_at, metadata
FROM memory_documents
WHERE user_id = $1 AND agent_id IS NOT DISTINCT FROM $2
AND (path LIKE '%/.config' OR path = '.config')
ORDER BY path
"#,
&[&user_id, &agent_id],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to find config documents: {e}"),
})?;
Ok(rows.iter().map(|r| self.row_to_document(r)).collect())
}
// ==================== Versioning ====================
pub async fn save_version(
&self,
document_id: Uuid,
content: &str,
content_hash: &str,
changed_by: Option<&str>,
) -> Result<i32, WorkspaceError> {
let conn = self.conn().await?;
let row = conn
.query_one(
r#"
INSERT INTO memory_document_versions
(id, document_id, version, content, content_hash, changed_by)
VALUES (
gen_random_uuid(),
$1,
(SELECT COALESCE(MAX(version), 0) + 1
FROM memory_document_versions WHERE document_id = $1),
$2, $3, $4
)
RETURNING version
"#,
&[&document_id, &content, &content_hash, &changed_by],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to save version: {e}"),
})?;
Ok(row.get(0))
}
pub async fn get_version(
&self,
document_id: Uuid,
version: i32,
) -> Result<DocumentVersion, WorkspaceError> {
let conn = self.conn().await?;
let row = conn
.query_opt(
r#"
SELECT id, document_id, version, content, content_hash,
created_at, changed_by
FROM memory_document_versions
WHERE document_id = $1 AND version = $2
"#,
&[&document_id, &version],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to get version: {e}"),
})?
.ok_or(WorkspaceError::VersionNotFound {
document_id,
version,
})?;
Ok(DocumentVersion {
id: row.get(0),
document_id: row.get(1),
version: row.get(2),
content: row.get(3),
content_hash: row.get(4),
created_at: row.get(5),
changed_by: row.get(6),
})
}
pub async fn list_versions(
&self,
document_id: Uuid,
limit: i64,
) -> Result<Vec<VersionSummary>, WorkspaceError> {
let conn = self.conn().await?;
let rows = conn
.query(
r#"
SELECT version, content_hash, created_at, changed_by
FROM memory_document_versions
WHERE document_id = $1
ORDER BY version DESC
LIMIT $2
"#,
&[&document_id, &limit],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to list versions: {e}"),
})?;
Ok(rows
.iter()
.map(|row| VersionSummary {
version: row.get(0),
content_hash: row.get(1),
created_at: row.get(2),
changed_by: row.get(3),
})
.collect())
}
pub async fn get_latest_version_number(
&self,
document_id: Uuid,
) -> Result<Option<i32>, WorkspaceError> {
let conn = self.conn().await?;
let row = conn
.query_one(
"SELECT MAX(version) FROM memory_document_versions WHERE document_id = $1",
&[&document_id],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to get latest version number: {e}"),
})?;
Ok(row.get(0))
}
pub async fn prune_versions(
&self,
document_id: Uuid,
keep_count: i32,
) -> Result<u64, WorkspaceError> {
let conn = self.conn().await?;
let result = conn
.execute(
r#"
DELETE FROM memory_document_versions
WHERE document_id = $1
AND version NOT IN (
SELECT version FROM memory_document_versions
WHERE document_id = $1
ORDER BY version DESC
LIMIT $2
)
"#,
&[&document_id, &(keep_count as i64)],
)
.await
.map_err(|e| WorkspaceError::SearchFailed {
reason: format!("Failed to prune versions: {e}"),
})?;
Ok(result)
}
}
+2 -4
View File
@@ -955,8 +955,7 @@ mod tests {
let hygiene_config = HygieneConfig {
enabled: false,
daily_retention_days: 30,
conversation_retention_days: 7,
version_keep_count: 50,
cadence_hours: 24,
state_dir: _tmp.path().to_path_buf(),
};
@@ -1002,8 +1001,7 @@ mod tests {
let hygiene_config = HygieneConfig {
enabled: false,
daily_retention_days: 30,
conversation_retention_days: 7,
version_keep_count: 50,
cadence_hours: 24,
state_dir: _tmp.path().to_path_buf(),
};
+1 -7
View File
@@ -48,13 +48,7 @@ mod tests {
let store = rig.database();
assert!(
store
.ensure_conversation(
foreign_thread_id,
"gateway",
"victim-user",
None,
Some("gateway")
)
.ensure_conversation(foreign_thread_id, "gateway", "victim-user", None)
.await
.expect("failed to create victim conversation"),
"test setup failed: victim conversation was not created"