nx_task_graph_test.nx source
↩ module page · 127 lines · 4552 B
1// nx_task_graph_test.nx -- DAG ordering + dependency satisfaction.
2//
3// T1: diamond DAG
4// A -> B
5// A -> C
6// B -> D
7// C -> D
8// Each node writes its tag (1..4) into a sequence log at an
9// atomic-FAA index. D must appear AFTER both B and C.
10//
11// T2: fan-out / fan-in (1 source, K fan-out, 1 sink) with hw-sized
12// parallelism inside the fan.
13
14import "nx_kernel_v2.nx"
15import "nx_log.nx"
16import "nx_atom.nx"
17import "nx_thread_pool.nx"
18import "nx_task_graph.nx"
19import "nx_hw.nx"
20
21const SEQ_LOG_CAP: i64 = 64
22
23// Shared state passed to each node via ctx (we pass the address of
24// this struct as i64).
25struct TraceState {
26 seq_idx: i64, // atomic; bumped per node finish
27 seq_log_addr: i64, // *i64 to log array
28 tag_a: i64, tag_b: i64, tag_c: i64, tag_d: i64,
29 pos_b: i64, pos_c: i64, pos_d: i64,
30}
31
32func _trace_log(state_addr: i64, tag: i64) -> i64 {
33 let s: *TraceState = state_addr as *TraceState
34 let log: *i64 = s.seq_log_addr as *i64
35 let seq_addr: *i64 = (state_addr as *i64)
36 let idx: i64 = nx_atom_faa_i64(seq_addr, 1, NX_MO_SEQ_CST)
37 log[idx] = tag
38 return idx
39}
40
41func node_A(ctx: i64) -> i64 {
42 _trace_log(ctx, 1)
43 return 0
44}
45func node_B(ctx: i64) -> i64 {
46 let pos: i64 = _trace_log(ctx, 2)
47 let s: *TraceState = ctx as *TraceState
48 s.pos_b = pos
49 return 0
50}
51func node_C(ctx: i64) -> i64 {
52 let pos: i64 = _trace_log(ctx, 3)
53 let s: *TraceState = ctx as *TraceState
54 s.pos_c = pos
55 return 0
56}
57func node_D(ctx: i64) -> i64 {
58 let pos: i64 = _trace_log(ctx, 4)
59 let s: *TraceState = ctx as *TraceState
60 s.pos_d = pos
61 return 0
62}
63
64func main() -> nx_exit {
65 var n_workers: i64 = nx_hw_worker_count()
66 if n_workers < 2 { n_workers = 2 }
67 let pool: *NxThreadPool = nx_pool_new(n_workers, 64)
68 println("=== nx_task_graph DAG smoke ===" as *u8)
69 print_i64(n_workers); println(" pool workers" as *u8)
70
71 // ---- T1: diamond DAG ----
72 let log_raw: *u8 = sys_mmap(SEQ_LOG_CAP * 8)
73 let state_raw: *u8 = sys_mmap(128)
74 let st: *TraceState = state_raw as *TraceState
75 st.seq_idx = 0
76 st.seq_log_addr = log_raw as i64
77 st.pos_b = -1; st.pos_c = -1; st.pos_d = -1
78
79 let g: *NxTaskGraph = nx_graph_new(pool, 16)
80 let a_idx: i64 = nx_graph_add_node(g, node_A, state_raw as i64)
81 let b_idx: i64 = nx_graph_add_node(g, node_B, state_raw as i64)
82 let c_idx: i64 = nx_graph_add_node(g, node_C, state_raw as i64)
83 let d_idx: i64 = nx_graph_add_node(g, node_D, state_raw as i64)
84 nx_graph_add_edge(g, a_idx, b_idx)
85 nx_graph_add_edge(g, a_idx, c_idx)
86 nx_graph_add_edge(g, b_idx, d_idx)
87 nx_graph_add_edge(g, c_idx, d_idx)
88
89 if nx_graph_run(g) != 0 { println("FAIL T1 run timeout" as *u8); return 1 }
90 if nx_graph_n_completed(g) != 4 { println("FAIL T1 completed_count" as *u8); return 2 }
91
92 // Order checks: D must come after B AND C. A must be at pos 0.
93 let log: *i64 = log_raw as *i64
94 if log[0] != 1 { println("FAIL T1: A not first" as *u8); return 3 }
95 if st.pos_d < st.pos_b { println("FAIL T1: D ran before B" as *u8); return 4 }
96 if st.pos_d < st.pos_c { println("FAIL T1: D ran before C" as *u8); return 5 }
97 if st.pos_d != 3 { println("FAIL T1: D not last" as *u8); return 6 }
98 println("PASS T1: diamond DAG ordering correct (A->{B,C}->D)" as *u8)
99
100 // ---- T2: fan-out source -> 8 leaves -> sink ----
101 let log2_raw: *u8 = sys_mmap(SEQ_LOG_CAP * 8)
102 let state2_raw: *u8 = sys_mmap(128)
103 let st2: *TraceState = state2_raw as *TraceState
104 st2.seq_idx = 0
105 st2.seq_log_addr = log2_raw as i64
106 st2.pos_b = -1; st2.pos_c = -1; st2.pos_d = -1
107
108 let g2: *NxTaskGraph = nx_graph_new(pool, 16)
109 let src: i64 = nx_graph_add_node(g2, node_A, state2_raw as i64)
110 let sink: i64 = nx_graph_add_node(g2, node_D, state2_raw as i64)
111 var k: i64 = 0
112 while k < 8 {
113 let mid: i64 = nx_graph_add_node(g2, node_B, state2_raw as i64)
114 nx_graph_add_edge(g2, src, mid)
115 nx_graph_add_edge(g2, mid, sink)
116 k = k + 1
117 }
118 if nx_graph_run(g2) != 0 { println("FAIL T2 run timeout" as *u8); return 7 }
119 if nx_graph_n_completed(g2) != 10 { println("FAIL T2 completed_count" as *u8); return 8 }
120 if st2.pos_d < 9 { println("FAIL T2 sink ran early" as *u8); return 9 }
121 println("PASS T2: fan-out (1 -> 8 -> 1) -- sink ran last" as *u8)
122
123 nx_pool_shutdown(pool)
124 println("" as *u8)
125 println("Task graph proven: dependency edges enforced under MIMD scheduling." as *u8)
126 return 0
127}