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
21pub 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 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(); 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 #[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 #[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 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 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 #[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 pub fn write_timestamp(&mut self, query_index: u32, stage: vk::PipelineStageFlags2) {
479 self.push_task(RawTask::WriteTimestamp { query_index, stage });
480 }
481
482 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}