返回 CodeWhale
parity_state.rs
根目录 / crates / state / tests / parity_state.rs
1 use std::path::PathBuf;
2
3 use codewhale_state::{SessionSource, StateStore, ThreadListFilters, ThreadMetadata, ThreadStatus};
4 use rusqlite::Connection;
5
6 fn temp_state_path(label: &str) -> PathBuf {
7 std::env::temp_dir().join(format!(
8 "deepseek_state_test_{}_{}_{}.db",
9 label,
10 std::process::id(),
11 chrono::Utc::now().timestamp_nanos_opt().unwrap_or(0)
12 ))
13 }
14
15 fn assert_workflow_trace_schema(conn: &Connection) {
16 let user_version: u32 = conn
17 .query_row("PRAGMA user_version;", [], |row| row.get(0))
18 .expect("read user_version");
19 // v4 (goal-progress migration) adds `thread_goals.continuation_count` on top
20 // of the v3 workflow-trace + thread_goals tables. The table set asserted
21 // below is unchanged; only the schema version advanced.
22 assert_eq!(user_version, 4);
23
24 for table in [
25 "workflow_runs",
26 "branch_runs",
27 "leaf_runs",
28 "control_node_runs",
29 "teacher_candidates",
30 "thread_goals",
31 ] {
32 let exists: bool = conn
33 .query_row(
34 "SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?1)",
35 [table],
36 |row| row.get(0),
37 )
38 .unwrap_or_else(|err| panic!("read sqlite_master for {table}: {err}"));
39 assert!(exists, "missing workflow trace table {table}");
40 }
41 }
42
43 #[test]
44 fn upsert_and_resume_thread_metadata() {
45 let path = temp_state_path("upsert_resume");
46 let store = StateStore::open(Some(path.clone())).expect("open state store");
47 let now = chrono::Utc::now().timestamp();
48 let thread = ThreadMetadata {
49 id: "thread-test-1".to_string(),
50 rollout_path: Some(PathBuf::from("/tmp/rollout.jsonl")),
51 preview: "hello".to_string(),
52 ephemeral: false,
53 model_provider: "deepseek".to_string(),
54 created_at: now,
55 updated_at: now,
56 status: ThreadStatus::Running,
57 path: Some(PathBuf::from("/tmp/project")),
58 cwd: PathBuf::from("/tmp/project"),
59 cli_version: "0.0.0-test".to_string(),
60 source: SessionSource::Interactive,
61 name: Some("Test Thread".to_string()),
62 sandbox_policy: Some("workspace-write".to_string()),
63 approval_mode: Some("on-request".to_string()),
64 archived: false,
65 archived_at: None,
66 git_sha: None,
67 git_branch: None,
68 git_origin_url: None,
69 memory_mode: Some("extended".to_string()),
70 current_leaf_id: None,
71 };
72 store.upsert_thread(&thread).expect("upsert thread");
73
74 let loaded = store
75 .get_thread("thread-test-1")
76 .expect("read thread")
77 .expect("thread must exist");
78 assert_eq!(loaded.id, "thread-test-1");
79 assert_eq!(loaded.name.as_deref(), Some("Test Thread"));
80 assert_eq!(loaded.memory_mode.as_deref(), Some("extended"));
81 assert_eq!(
82 loaded.rollout_path,
83 Some(PathBuf::from("/tmp/rollout.jsonl"))
84 );
85
86 store
87 .mark_archived("thread-test-1")
88 .expect("archive thread");
89 let archived = store
90 .get_thread("thread-test-1")
91 .expect("read archived thread")
92 .expect("thread exists after archive");
93 assert!(archived.archived);
94
95 let listed = store
96 .list_threads(ThreadListFilters {
97 include_archived: true,
98 limit: Some(10),
99 })
100 .expect("list threads");
101 assert!(!listed.is_empty());
102 }
103
104 #[test]
105 fn init_schema_migration() {
106 let path = temp_state_path("init_schema_migration");
107 let conn = Connection::open(&path).expect("open state db");
108 conn.execute_batch(
109 r#"
110 CREATE TABLE IF NOT EXISTS threads (
111 id TEXT PRIMARY KEY,
112 rollout_path TEXT,
113 preview TEXT NOT NULL,
114 ephemeral INTEGER NOT NULL,
115 model_provider TEXT NOT NULL,
116 created_at INTEGER NOT NULL,
117 updated_at INTEGER NOT NULL,
118 status TEXT NOT NULL,
119 path TEXT,
120 cwd TEXT NOT NULL,
121 cli_version TEXT NOT NULL,
122 source TEXT NOT NULL,
123 title TEXT,
124 sandbox_policy TEXT,
125 approval_mode TEXT,
126 archived INTEGER NOT NULL DEFAULT 0,
127 archived_at INTEGER,
128 git_sha TEXT,
129 git_branch TEXT,
130 git_origin_url TEXT,
131 memory_mode TEXT
132 );
133 CREATE TABLE IF NOT EXISTS messages (
134 id INTEGER PRIMARY KEY AUTOINCREMENT,
135 thread_id TEXT NOT NULL,
136 role TEXT NOT NULL,
137 content TEXT NOT NULL,
138 item_json TEXT,
139 created_at INTEGER NOT NULL,
140 FOREIGN KEY(thread_id) REFERENCES threads(id) ON DELETE CASCADE
141 );
142 INSERT INTO threads (
143 id, preview, ephemeral, model_provider, created_at, updated_at, status, cwd, cli_version, source, archived
144 )
145 VALUES (
146 'thread-test-1', 'hello', false, 'deepseek', 0, 0, 'running', '/tmp/project', '0.0.0-test', 'interactive', false
147 );
148 INSERT INTO messages (thread_id, role, content, created_at) VALUES
149 ('thread-test-1', 'foo0', 'bar0', 0),
150 ('thread-test-1', 'foo1', 'bar1', 1),
151 ('thread-test-1', 'foo2', 'bar2', 2);
152 "#,
153 )
154 .expect("init schema migration");
155
156 let store = StateStore::open(Some(path.clone())).expect("open state store");
157 let thread = store
158 .get_thread("thread-test-1")
159 .expect("read thread")
160 .unwrap();
161 assert_eq!(thread.id, "thread-test-1");
162 assert_eq!(thread.preview, "hello");
163 assert!(!thread.ephemeral);
164 assert_eq!(thread.model_provider, "deepseek");
165 assert_eq!(thread.created_at, 0);
166 assert_eq!(thread.updated_at, 0);
167 assert_eq!(thread.status, ThreadStatus::Running);
168 assert_eq!(thread.cwd, PathBuf::from("/tmp/project"));
169 assert_eq!(thread.cli_version, "0.0.0-test");
170 assert_eq!(thread.source, SessionSource::Interactive);
171 assert!(thread.current_leaf_id.is_some());
172
173 let messages = store
174 .list_messages("thread-test-1", None)
175 .expect("list messages");
176 assert_eq!(messages.len(), 3);
177 for (i, message) in messages.iter().enumerate() {
178 assert_eq!(message.thread_id, "thread-test-1");
179 assert_eq!(message.role, format!("foo{i}"));
180 assert_eq!(message.content, format!("bar{i}"));
181 assert_eq!(message.created_at, i as i64);
182 }
183
184 // Test idempotent
185 StateStore::open(Some(path.clone())).expect("open state store");
186 }
187
188 #[test]
189 fn fresh_schema_includes_workflow_trace_tables() {
190 let path = temp_state_path("fresh_schema_includes_workflow_trace_tables");
191
192 StateStore::open(Some(path.clone())).expect("open state store");
193
194 let conn = Connection::open(&path).expect("open state db");
195 assert_workflow_trace_schema(&conn);
196 }
197
198 #[test]
199 fn v1_schema_migrates_workflow_trace_tables() {
200 let path = temp_state_path("v1_schema_migrates_workflow_trace_tables");
201 let conn = Connection::open(&path).expect("open state db");
202 conn.execute_batch(
203 r#"
204 CREATE TABLE threads (
205 id TEXT PRIMARY KEY,
206 rollout_path TEXT,
207 preview TEXT NOT NULL,
208 ephemeral INTEGER NOT NULL,
209 model_provider TEXT NOT NULL,
210 created_at INTEGER NOT NULL,
211 updated_at INTEGER NOT NULL,
212 status TEXT NOT NULL,
213 path TEXT,
214 cwd TEXT NOT NULL,
215 cli_version TEXT NOT NULL,
216 source TEXT NOT NULL,
217 title TEXT,
218 sandbox_policy TEXT,
219 approval_mode TEXT,
220 archived INTEGER NOT NULL DEFAULT 0,
221 archived_at INTEGER,
222 git_sha TEXT,
223 git_branch TEXT,
224 git_origin_url TEXT,
225 memory_mode TEXT,
226 current_leaf_id INTEGER
227 );
228 CREATE TABLE messages (
229 id INTEGER PRIMARY KEY AUTOINCREMENT,
230 thread_id TEXT NOT NULL,
231 role TEXT NOT NULL,
232 content TEXT NOT NULL,
233 item_json TEXT,
234 created_at INTEGER NOT NULL,
235 parent_entry_id INTEGER
236 );
237 CREATE TABLE checkpoints (
238 thread_id TEXT NOT NULL,
239 checkpoint_id TEXT NOT NULL,
240 state_json TEXT NOT NULL,
241 created_at INTEGER NOT NULL,
242 PRIMARY KEY(thread_id, checkpoint_id)
243 );
244 CREATE TABLE jobs (
245 id TEXT PRIMARY KEY,
246 name TEXT NOT NULL,
247 status TEXT NOT NULL,
248 progress INTEGER,
249 detail TEXT,
250 created_at INTEGER NOT NULL,
251 updated_at INTEGER NOT NULL
252 );
253 CREATE TABLE thread_dynamic_tools (
254 thread_id TEXT NOT NULL,
255 position INTEGER NOT NULL,
256 name TEXT NOT NULL,
257 description TEXT,
258 input_schema TEXT NOT NULL,
259 PRIMARY KEY (thread_id, position)
260 );
261 INSERT INTO threads (
262 id, preview, ephemeral, model_provider, created_at, updated_at, status, cwd, cli_version, source, archived
263 )
264 VALUES (
265 'thread-test-1', 'hello', false, 'deepseek', 0, 0, 'running', '/tmp/project', '0.0.0-test', 'interactive', false
266 );
267 PRAGMA user_version = 1;
268 "#,
269 )
270 .expect("create v1 schema");
271 drop(conn);
272
273 let store = StateStore::open(Some(path.clone())).expect("open state store");
274 let thread = store
275 .get_thread("thread-test-1")
276 .expect("read thread")
277 .expect("thread survives migration");
278 assert_eq!(thread.preview, "hello");
279
280 let conn = Connection::open(&path).expect("open state db");
281 assert_workflow_trace_schema(&conn);
282 }
283
284 #[test]
285 fn init_schema_migration_same_second_messages() {
286 let path = temp_state_path("init_schema_migration_same_second_messages");
287 let conn = Connection::open(&path).expect("open state db");
288 conn.execute_batch(
289 r#"
290 CREATE TABLE IF NOT EXISTS threads (
291 id TEXT PRIMARY KEY,
292 rollout_path TEXT,
293 preview TEXT NOT NULL,
294 ephemeral INTEGER NOT NULL,
295 model_provider TEXT NOT NULL,
296 created_at INTEGER NOT NULL,
297 updated_at INTEGER NOT NULL,
298 status TEXT NOT NULL,
299 path TEXT,
300 cwd TEXT NOT NULL,
301 cli_version TEXT NOT NULL,
302 source TEXT NOT NULL,
303 title TEXT,
304 sandbox_policy TEXT,
305 approval_mode TEXT,
306 archived INTEGER NOT NULL DEFAULT 0,
307 archived_at INTEGER,
308 git_sha TEXT,
309 git_branch TEXT,
310 git_origin_url TEXT,
311 memory_mode TEXT
312 );
313 CREATE TABLE IF NOT EXISTS messages (
314 id INTEGER PRIMARY KEY AUTOINCREMENT,
315 thread_id TEXT NOT NULL,
316 role TEXT NOT NULL,
317 content TEXT NOT NULL,
318 item_json TEXT,
319 created_at INTEGER NOT NULL,
320 FOREIGN KEY(thread_id) REFERENCES threads(id) ON DELETE CASCADE
321 );
322 INSERT INTO threads (
323 id, preview, ephemeral, model_provider, created_at, updated_at, status, cwd, cli_version, source, archived
324 )
325 VALUES (
326 'thread-test-2', 'hello', false, 'deepseek', 0, 0, 'running', '/tmp/project', '0.0.0-test', 'interactive', false
327 );
328 INSERT INTO messages (thread_id, role, content, created_at) VALUES
329 ('thread-test-2', 'foo0', 'bar0', 123),
330 ('thread-test-2', 'foo1', 'bar1', 123),
331 ('thread-test-2', 'foo2', 'bar2', 123),
332 ('thread-test-2', 'foo3', 'bar3', 123);
333 "#,
334 )
335 .expect("init schema migration");
336
337 let store = StateStore::open(Some(path.clone())).expect("open state store");
338 let messages = store
339 .list_messages("thread-test-2", None)
340 .expect("list messages");
341 assert_eq!(messages.len(), 4);
342 for (i, message) in messages.iter().enumerate() {
343 assert_eq!(message.thread_id, "thread-test-2");
344 assert_eq!(message.role, format!("foo{i}"));
345 assert_eq!(message.content, format!("bar{i}"));
346 assert_eq!(message.created_at, 123);
347 }
348 assert_eq!(messages[0].parent_entry_id, None);
349 assert_eq!(messages[1].parent_entry_id, Some(messages[0].id));
350 assert_eq!(messages[2].parent_entry_id, Some(messages[1].id));
351 assert_eq!(messages[3].parent_entry_id, Some(messages[2].id));
352
353 // Test idempotent reopen after same-second parent links are migrated.
354 StateStore::open(Some(path.clone())).expect("open state store - idempotent");
355 }
356
357 #[test]
358 fn test_fork() {
359 let path = temp_state_path("test_fork");
360 let store = StateStore::open(Some(path.clone())).expect("open state store");
361 let now = chrono::Utc::now().timestamp();
362 let thread = ThreadMetadata {
363 id: "thread-test-1".to_string(),
364 rollout_path: Some(PathBuf::from("/tmp/rollout.jsonl")),
365 preview: "hello".to_string(),
366 ephemeral: false,
367 model_provider: "deepseek".to_string(),
368 created_at: now,
369 updated_at: now,
370 status: ThreadStatus::Running,
371 path: Some(PathBuf::from("/tmp/project")),
372 cwd: PathBuf::from("/tmp/project"),
373 cli_version: "0.0.0-test".to_string(),
374 source: SessionSource::Interactive,
375 name: Some("Test Thread".to_string()),
376 sandbox_policy: Some("workspace-write".to_string()),
377 approval_mode: Some("on-request".to_string()),
378 archived: false,
379 archived_at: None,
380 git_sha: None,
381 git_branch: None,
382 git_origin_url: None,
383 memory_mode: Some("extended".to_string()),
384 current_leaf_id: None,
385 };
386
387 store.upsert_thread(&thread).expect("upsert thread");
388 store
389 .append_message("thread-test-1", "foo0", "bar0", None)
390 .expect("append message");
391 store
392 .append_message("thread-test-1", "foo1", "bar1", None)
393 .expect("append message");
394 store
395 .append_message("thread-test-1", "foo2", "bar2", None)
396 .expect("append message");
397 store
398 .append_message("thread-test-1", "foo3", "bar3", None)
399 .expect("append message");
400 store
401 .append_message("thread-test-1", "foo4", "bar4", None)
402 .expect("append message");
403
404 let messages = store
405 .list_messages("thread-test-1", None)
406 .expect("list messages");
407 assert_eq!(messages.len(), 5);
408 let ids = messages
409 .iter()
410 .enumerate()
411 .map(|(i, message)| {
412 assert_eq!(message.thread_id, "thread-test-1");
413 assert_eq!(message.role, format!("foo{i}"));
414 assert_eq!(message.content, format!("bar{i}"));
415 message.id.to_string()
416 })
417 .collect::<Vec<_>>();
418
419 store.upsert_thread(&thread).expect("upsert thread");
420
421 store
422 .fork_at_message(&ids[2], "foo5", "bar5", None)
423 .expect("fork at message");
424 let messages = store
425 .list_messages("thread-test-1", None)
426 .expect("list messages");
427 assert_eq!(messages.len(), 4);
428 const LIST_1: [i64; 4] = [0, 1, 2, 5];
429 messages
430 .iter()
431 .zip(LIST_1.iter())
432 .for_each(|(message, &i)| {
433 assert_eq!(message.thread_id, "thread-test-1");
434 assert_eq!(message.role, format!("foo{i}"));
435 assert_eq!(message.content, format!("bar{i}"));
436 });
437 let leaves = store
438 .list_leaf_messages("thread-test-1")
439 .expect("list leaf messages");
440 assert_eq!(leaves.len(), 2);
441
442 store
443 .set_current_leaf_id("thread-test-1", &ids[4])
444 .expect("set current leaf id");
445 store
446 .append_message("thread-test-1", "foo6", "bar6", None)
447 .expect("append message");
448 let messages = store
449 .list_messages("thread-test-1", None)
450 .expect("list messages");
451 assert_eq!(messages.len(), 6);
452 const LIST_2: [i64; 6] = [0, 1, 2, 3, 4, 6];
453 messages
454 .iter()
455 .zip(LIST_2.iter())
456 .for_each(|(message, &i)| {
457 assert_eq!(message.thread_id, "thread-test-1");
458 assert_eq!(message.role, format!("foo{i}"));
459 assert_eq!(message.content, format!("bar{i}"));
460 });
461
462 let leaves = store
463 .list_leaf_messages("thread-test-1")
464 .expect("list leaf messages");
465 assert_eq!(leaves.len(), 2);
466
467 store
468 .clear_messages("thread-test-1")
469 .expect("clear messages");
470 let leaves = store
471 .list_leaf_messages("thread-test-1")
472 .expect("list leaf messages");
473 assert_eq!(leaves.len(), 0);
474 let thread = store
475 .get_thread("thread-test-1")
476 .expect("get thread")
477 .unwrap();
478 assert!(thread.current_leaf_id.is_none());
479 }
480
480 lines RUST