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