Skip to main content

enki/enki_api/context/
flow.rs

1use anyhow::{Context, Result};
2use ash::vk;
3use std::sync::Mutex;
4use std::sync::atomic::{AtomicBool, Ordering};
5
6use crate::enki_api::context::errors::emit_and_abort;
7use crate::enki_api::context::instance::Enki;
8use crate::enki_api::resources::{GpuVec, Slice};
9use crate::enki_api::space::Space;
10
11use super::instant::GpuInstant;
12use anu::diagnostics::source::Span;
13use anu::nam_args_api::NamDispatchMap;
14use anu::pipeline_synthesis::ComputeSynthesisInput;
15use anu::recording::queue::TaskQueue;
16use anu::recording::recipe::CompiledExecutionRecipe;
17use anu::recording::task::{ComputeTask, PresentBufferTask, RawTask};
18use anu::validation::{BorrowEngine, FrameBorrowLedger};
19use utu::GpuWindow;
20
21/// An active, coherent GPU compute and presentation execution stream.
22///
23/// A `Flow` batches nam (kernel) dispatches and display presentation tasks into an atomic
24/// command buffer sequence synchronized by timeline semaphores.
25pub struct Flow<'a> {
26    pub enki: &'a Enki,
27    pub queue: TaskQueue<'a>,
28    pub cmd: vk::CommandBuffer,
29    pub slot_idx: usize,
30    pub timeline_value: u64,
31    pub submitted: AtomicBool,
32    pub sticky_error: Option<anyhow::Error>,
33    pub borrow_ledger: FrameBorrowLedger,
34    pub timestamp_count: u32,
35}
36
37impl<'a> Flow<'a> {
38    /// Initializes a new Flow, acquiring a command buffer from the ring and resetting query pools.
39    pub fn new(enki: &'a Enki) -> Self {
40        let engine = &enki.engine;
41
42        let current_gpu_value = engine.timeline_semaphore.get_timeline_value().unwrap_or(0);
43        engine.allocator.reclaim_resources(current_gpu_value);
44
45        let timeline_value = engine.timeline_counter.fetch_add(1, Ordering::SeqCst) + 1;
46        engine
47            .allocator
48            .current_timeline_value
49            .store(timeline_value, Ordering::Release);
50
51        let (cmd, slot_idx, is_static) = {
52            let mut ring = engine.command_ring.lock().unwrap();
53            let (cmd, slot_idx) = ring
54                .acquire_next_cmd(engine.raw_device(), engine.timeline_semaphore.handle)
55                .context("[Flow] Failed to acquire command buffer from ring")
56                .unwrap();
57            let is_static = ring.slots[slot_idx].is_statically_recorded;
58            (cmd, slot_idx, is_static)
59        };
60
61        let max_queries = engine.max_timestamp_queries;
62        let start_query = (slot_idx as u32) * max_queries;
63
64        if !is_static {
65            unsafe {
66                let begin_info = vk::CommandBufferBeginInfo::default(); // Reusable buffer (بدون ONE_TIME_SUBMIT)
67
68                engine
69                    .raw_device()
70                    .begin_command_buffer(cmd, &begin_info)
71                    .unwrap();
72
73                engine.raw_device().cmd_reset_query_pool(
74                    cmd,
75                    engine.query_pool,
76                    start_query,
77                    max_queries,
78                );
79            }
80        }
81
82        Self {
83            enki,
84            queue: TaskQueue::new(),
85            cmd,
86            slot_idx,
87            timeline_value,
88            submitted: AtomicBool::new(false),
89            sticky_error: None,
90            borrow_ledger: FrameBorrowLedger::new(),
91            timestamp_count: 0,
92        }
93    }
94
95    /// Queues an owned GPU vector for direct presentation to the display window.
96    ///
97    /// # Temporal Presentation Hazard
98    /// Once queued for presentation, any subsequent dispatch (`.run()` or `.run_unchecked()`) within the same flow that
99    /// attempts to mutate (`&mut`) this buffer halts execution with diagnostic **`error[E1010]`**.
100    /// Read-only access (`&`) after presentation remains permitted.
101    ///
102    /// # Diagnostics
103    /// Halts execution with diagnostic **`error[E2002]`** if element count is less than `width * height`.
104    #[track_caller]
105    pub fn present<T: Copy + Send + Sync + 'static>(&mut self, pixels: &GpuVec<T>) {
106        let caller = std::panic::Location::caller();
107        self.validate_and_present_raw(
108            pixels.slot_index,
109            pixels._inner.buffer(),
110            pixels.offset as u64,
111            pixels.len(),
112            pixels.stride(),
113            caller,
114        );
115    }
116
117    /// Queues a contiguous GPU slice for direct presentation to the display window.
118    ///
119    /// # Temporal Presentation Hazard
120    /// Once queued for presentation, any subsequent dispatch (`.run()` or `.run_unchecked()`) within the same flow that
121    /// attempts to mutate (`&mut`) this buffer halts execution with diagnostic **`error[E1010]`**.
122    /// Read-only access (`&`) after presentation remains permitted.
123    ///
124    /// # Diagnostics
125    /// Halts execution with diagnostic **`error[E2002]`** if element count is less than `width * height`.    #[track_caller]
126    #[track_caller]
127    pub fn present_slice<T: Copy + Send + Sync + 'static>(&mut self, pixels: &Slice<T>) {
128        let caller = std::panic::Location::caller();
129        self.validate_and_present_raw(
130            pixels.slot_index,
131            pixels._inner.buffer(),
132            pixels.offset as u64,
133            pixels.len(),
134            pixels.stride(),
135            caller,
136        );
137    }
138    fn validate_and_present_raw(
139        &mut self,
140        buffer_id: u32,
141        buffer: ash::vk::Buffer,
142        offset: u64,
143        element_count: usize,
144        _element_stride: usize,
145        caller: &'static std::panic::Location<'static>,
146    ) {
147        self.borrow_ledger
148            .mark_queued_for_presentation(buffer_id as usize);
149
150        let (width, height) = if let Some(window_mutex) = &self.enki.gpu_window {
151            let window = window_mutex.lock().unwrap();
152            (window.extent.width, window.extent.height)
153        } else {
154            (0, 0)
155        };
156
157        if width == 0 || height == 0 {
158            return;
159        }
160
161        let required_elements = (width * height) as usize;
162
163        if element_count < required_elements {
164            let diag = anu::diagnostics::rt::present_dimension_mismatch(
165                width,
166                height,
167                element_count,
168                caller,
169            );
170            emit_and_abort(&diag);
171        }
172
173        let task = PresentBufferTask {
174            buffer_id,
175            buffer,
176            offset,
177            width,
178            height,
179        };
180
181        self.push_task(RawTask::PresentBuffer(task));
182    }
183
184    pub(crate) fn nam_impl_direct<F>(
185        &mut self,
186        space: &Space,
187        args_ctx: anu::nam_args_api::IngressContext<'static>,
188        mut map: NamDispatchMap,
189    ) -> Result<()>
190    where
191        F: 'static,
192    {
193        let engine = &self.enki.engine;
194
195        let (nam_name, contract) =
196            crate::enki_api::context::contract::resolve_contract_cached::<F>();
197
198        map.nam_name = nam_name.clone();
199        map.expected_contract = contract;
200
201        let required_param_bytes: u64 = args_ctx
202            .descriptors
203            .iter()
204            .map(|d| d.arena_size_bytes as u64)
205            .sum();
206        let arena_capacity = engine.param_arena.size_bytes();
207
208        if required_param_bytes > arena_capacity {
209            let diag =
210                anu::diagnostics::hw::param_arena_overflow(required_param_bytes, arena_capacity);
211            crate::enki_api::context::errors::emit_and_abort(&diag);
212        }
213
214        if let Err(violation) = BorrowEngine::validate_dispatch(&map, &mut self.borrow_ledger) {
215            let diag = anu::diagnostics::ContractDiagnosticBuilder::from_violation(violation, &map);
216            emit_and_abort(&diag);
217        }
218
219        let dispatch = space.resolve_dispatch(&engine.hardware_profile);
220
221        let input = ComputeSynthesisInput {
222            nam_name: nam_name.clone(),
223            local_size: dispatch.local_size,
224            host_manifest_dir: std::env::var("CARGO_MANIFEST_DIR").ok(),
225            caller_file_path: map.call_site.map(|(f, _, _)| f.to_string()),
226            arg_descriptors: args_ctx.descriptors.clone(),
227        };
228
229        let artifact = engine
230            .synthesizer
231            .synthesize_compute(engine, &input)
232            .context("[Flow] JIT synthesis failed for nam")?;
233
234        let total_threads = (dispatch.global_size.0 as u64)
235            * (dispatch.global_size.1 as u64)
236            * (dispatch.global_size.2 as u64);
237        let required_stack_bytes = (artifact.stack_size_per_thread as u64) * total_threads;
238
239        let (stack_bda, stack_buffer) = if required_stack_bytes > 0 {
240            match apsu::GpuStackBuffer::allocate(engine.allocator.clone(), required_stack_bytes) {
241                Ok(Some(buf)) => {
242                    let bda = buf.device_address();
243                    (bda, Some(buf))
244                }
245                Ok(None) => (0, None),
246                Err(alloc_err) => {
247                    let mut diag = anu::diagnostics::hw::stack_overflow(
248                        &alloc_err,
249                        dispatch.global_size,
250                        artifact.stack_size_per_thread,
251                        None,
252                    );
253
254                    if let Some((file, line, col)) = map.call_site {
255                        diag.add_span(Span::primary(file, line as usize, col as usize, 1));
256                    }
257
258                    emit_and_abort(&diag);
259                }
260            }
261        } else {
262            (0, None)
263        };
264
265        let task = ComputeTask {
266            pipeline: artifact.pipeline,
267            layout: artifact.layout,
268            grid_size: dispatch.global_size,
269            local_size: dispatch.local_size,
270            args_ctx,
271            stack_bda,
272            _stack_buffer: stack_buffer,
273        };
274
275        self.push_task(RawTask::Compute(task));
276        Ok(())
277    }
278
279    /// Finalizes and submits the recorded flow to the GPU queue, halting on failure.
280    pub fn end_flow(self) {
281        if let Err(e) = self.try_end_flow() {
282            crate::enki_api::context::errors::handle_execution_error(&e);
283        }
284    }
285
286    /// Finalizes and submits the recorded flow to the GPU queue, returning a `Result`.
287    pub fn try_end_flow(self) -> Result<()> {
288        if self.submitted.swap(true, Ordering::SeqCst) {
289            return Ok(());
290        }
291
292        let engine = &self.enki.engine;
293        let recipe = engine.compile_recipe(&self.queue);
294
295        let has_present = self
296            .queue
297            .tasks
298            .iter()
299            .any(|t| matches!(t, RawTask::PresentBuffer(_)));
300
301        if has_present && let Some(window_mutex) = &self.enki.gpu_window {
302            self.submit_windowed(&recipe, window_mutex)?;
303        } else {
304            self.submit_headless(&recipe)?;
305        }
306
307        {
308            let mut ring = engine.command_ring.lock().unwrap();
309            ring.update_slot_timeline(self.slot_idx, self.timeline_value);
310        }
311
312        Ok(())
313    }
314
315    fn submit_headless(&self, recipe: &CompiledExecutionRecipe) -> Result<()> {
316        let engine = &self.enki.engine;
317        let device = engine.raw_device();
318
319        recipe.update_parameters(engine, &self.queue, self.slot_idx)?;
320
321        let is_static = {
322            let ring = engine.command_ring.lock().unwrap();
323            ring.slots[self.slot_idx].is_statically_recorded
324        };
325
326        if !is_static {
327            recipe.record_commands(
328                engine,
329                self.cmd,
330                &self.queue,
331                self.slot_idx,
332                engine.query_pool,
333            )?;
334
335            unsafe {
336                device.end_command_buffer(self.cmd)?;
337            }
338
339            let mut ring = engine.command_ring.lock().unwrap();
340            ring.slots[self.slot_idx].is_statically_recorded = true;
341        }
342
343        let cmd_buffers = [self.cmd];
344        let signal_semaphores = [engine.timeline_semaphore.handle];
345        let signal_values = [self.timeline_value];
346
347        let mut timeline_info =
348            vk::TimelineSemaphoreSubmitInfo::default().signal_semaphore_values(&signal_values);
349
350        let submit_info = vk::SubmitInfo::default()
351            .push_next(&mut timeline_info)
352            .command_buffers(&cmd_buffers)
353            .signal_semaphores(&signal_semaphores);
354
355        unsafe {
356            device.queue_submit(engine.queue.handle, &[submit_info], vk::Fence::null())?;
357            engine
358                .timeline_semaphore
359                .wait_timeline(self.timeline_value, std::time::Duration::from_secs(5))?;
360        }
361
362        Ok(())
363    }
364
365    fn submit_windowed(
366        &self,
367        recipe: &CompiledExecutionRecipe,
368        window_mutex: &Mutex<GpuWindow>,
369    ) -> Result<()> {
370        let engine = &self.enki.engine;
371        let device = engine.raw_device();
372
373        let mut window_lock = window_mutex.lock().unwrap();
374
375        let (image_index, _) = window_lock
376            .acquire_next_image(std::time::Duration::from_secs(5))
377            .context("[Flow] Failed to acquire next swapchain image")?;
378
379        recipe.update_parameters(engine, &self.queue, self.slot_idx)?;
380
381        let is_static = {
382            let ring = engine.command_ring.lock().unwrap();
383            ring.slots[self.slot_idx].is_statically_recorded
384        };
385
386        if !is_static {
387            recipe.record_commands(
388                engine,
389                self.cmd,
390                &self.queue,
391                self.slot_idx,
392                engine.query_pool,
393            )?;
394
395            for task in &self.queue.tasks {
396                if let RawTask::PresentBuffer(p) = task {
397                    window_lock.cmd_copy_buffer_to_image(
398                        self.cmd,
399                        image_index,
400                        p.buffer,
401                        p.offset,
402                        p.width,
403                        p.height,
404                    );
405                }
406            }
407
408            unsafe {
409                device.end_command_buffer(self.cmd)?;
410            }
411
412            let mut ring = engine.command_ring.lock().unwrap();
413            ring.slots[self.slot_idx].is_statically_recorded = true;
414        }
415
416        let wait_semaphores = [window_lock.current_image_acquired_semaphore()];
417        let wait_stages =
418            [vk::PipelineStageFlags::COLOR_ATTACHMENT_OUTPUT | vk::PipelineStageFlags::TRANSFER];
419        let signal_semaphores = [
420            engine.timeline_semaphore.handle,
421            window_lock.current_render_finished_semaphore(),
422        ];
423        let signal_values = [self.timeline_value, 0];
424
425        let mut timeline_info =
426            vk::TimelineSemaphoreSubmitInfo::default().signal_semaphore_values(&signal_values);
427
428        let cmd_buffers = [self.cmd];
429        let submit_info = vk::SubmitInfo::default()
430            .push_next(&mut timeline_info)
431            .wait_semaphores(&wait_semaphores)
432            .wait_dst_stage_mask(&wait_stages)
433            .command_buffers(&cmd_buffers)
434            .signal_semaphores(&signal_semaphores);
435
436        let in_flight_fence = window_lock.current_in_flight_fence();
437
438        unsafe {
439            device.queue_submit(engine.queue.handle, &[submit_info], in_flight_fence)?;
440        }
441
442        window_lock
443            .present_image(engine.queue.handle, image_index)
444            .context("[Flow] Failed to present swapchain image")?;
445
446        Ok(())
447    }
448
449    /// Records a hardware timestamp query mark on the GPU timeline for profiling.
450    #[track_caller]
451    pub fn mark(&mut self) -> GpuInstant {
452        let caller = std::panic::Location::caller();
453        let engine = &self.enki.engine;
454        let max_queries = engine.max_timestamp_queries;
455
456        if self.timestamp_count >= max_queries {
457            let diag = anu::diagnostics::hw::timestamp_queries_exceeded(
458                self.timestamp_count + 1,
459                max_queries,
460                Some(caller),
461            );
462            emit_and_abort(&diag);
463        }
464
465        let global_query_slot = (self.slot_idx as u32) * max_queries + self.timestamp_count;
466        self.timestamp_count += 1;
467
468        self.write_timestamp(global_query_slot, vk::PipelineStageFlags2::ALL_COMMANDS);
469
470        GpuInstant {
471            query_slot: global_query_slot,
472            timeline_value: self.timeline_value,
473            timestamp_period: engine.timestamp_period,
474        }
475    }
476
477    /// Records a timestamp query into the active flow.
478    pub fn write_timestamp(&mut self, query_index: u32, stage: vk::PipelineStageFlags2) {
479        self.push_task(RawTask::WriteTimestamp { query_index, stage });
480    }
481
482    /// Pushes a low-level task onto the internal execution queue.
483    pub fn push_task(&mut self, task: RawTask<'a>) {
484        self.queue.push(task);
485    }
486}
487
488impl<'a> Drop for Flow<'a> {
489    fn drop(&mut self) {
490        if !self.submitted.load(Ordering::SeqCst) {
491            let _ = self.enki.engine.wait_idle();
492        }
493    }
494}