1use std::collections::{BTreeMap, HashMap};
37use std::io::{self, Write};
38use std::sync::Arc;
39use std::time::Duration;
40
41use arrow_array::{FixedSizeListArray, Float64Array, RecordBatch};
42use arrow_ipc::writer::StreamWriter;
43use arrow_schema::{DataType, Field, FieldRef, Fields, Schema, SchemaRef};
44use axum::extract::rejection::JsonRejection;
45use axum::extract::{Extension, State as AxumState};
46use axum::http::{header, StatusCode};
47use axum::response::{IntoResponse, Response};
48use axum::Json;
49use openkind_core::{Question, State, SystemRequest};
50use openkind_engine::{dispatch, EngineError, EngineRegistry};
51use serde::Deserialize;
52
53use crate::error::ApiError;
54use crate::AppState;
55
56#[path = "arrow_columns.rs"]
57mod columns;
58#[path = "arrow_decoder.rs"]
59mod decoder;
60use columns::BatchBuilder;
61pub use decoder::answers_from_batch;
62
63pub const MAX_ARROW_RESPONSE_BYTES: usize = 64 * 1024 * 1024;
65pub const ARROW_BATCH_TIMEOUT: Duration = Duration::from_secs(600);
67
68pub const ARROW_CONTENT_TYPE: &str = "application/vnd.apache.arrow.stream";
70
71pub const MAX_ARROW_STATES: usize = 10_000;
75pub(crate) const ARROW_ADMISSION_UNITS: usize = MAX_ARROW_STATES;
77
78pub const MAX_CHOICE_LABELS: usize = u8::MAX as usize + 1;
83
84pub const META_ARROW_VERSION: &str = "openkind.arrow.version";
86pub const META_MODEL: &str = "openkind.model";
88pub const META_USAGE_INPUT_TOKENS: &str = "openkind.usage.input_tokens";
90pub const META_USAGE_OUTPUT_TOKENS: &str = "openkind.usage.output_tokens";
92pub const META_JEV_TYPE: &str = "jev.type";
94pub const META_LABELS: &str = "labels";
96pub const META_LEGEND: &str = "legend";
98
99pub const ARROW_MAPPING_VERSION: &str = "1";
101
102#[derive(Debug, Clone, Deserialize)]
109pub struct ArrowBatchRequest {
110 pub model: String,
112 pub states: Vec<State>,
114 pub questions: BTreeMap<String, Question>,
116}
117
118impl ArrowBatchRequest {
119 fn system_request_for(&self, state: &State) -> SystemRequest {
124 SystemRequest {
125 state: state.clone(),
126 model: self.model.clone(),
127 questions: self
128 .questions
129 .iter()
130 .map(|(k, v)| (k.clone(), v.clone()))
131 .collect(),
132 }
133 }
134}
135
136#[cfg(test)]
137fn question_ids(questions: &BTreeMap<String, Question>) -> Vec<&str> {
139 questions.keys().map(String::as_str).collect()
140}
141
142fn field_metadata(pairs: &[(&str, &str)]) -> HashMap<String, String> {
144 pairs
145 .iter()
146 .map(|(key, value)| ((*key).to_string(), (*value).to_string()))
147 .collect()
148}
149
150fn field_for_question(id: &str, question: &Question) -> Result<Field, ApiError> {
152 match question {
153 Question::Noul(_) => Ok(Field::new(id, DataType::Float64, false)
154 .with_metadata(field_metadata(&[(META_JEV_TYPE, "noul")]))),
155 Question::Choice(choice) => {
156 let labels = sorted_choice_labels(choice);
157 if labels.len() > MAX_CHOICE_LABELS {
158 return Err(ApiError::InvalidBody(format!(
159 "choice question `{id}` has {} options; the Arrow mapping indexes `choice` into a uint8, so at most {MAX_CHOICE_LABELS} options are supported",
160 labels.len()
161 )));
162 }
163 let children = Fields::from(vec![
164 Field::new("choice", DataType::UInt8, false),
165 Field::new("confidence", DataType::Float64, false),
166 fixed_probabilities_field(labels.len()),
167 ]);
168 let labels_json = serde_json::to_string(&labels)
169 .map_err(|e| ApiError::Internal(format!("encode labels metadata: {e}")))?;
170 Ok(
171 Field::new_struct(id, children, false).with_metadata(field_metadata(&[
172 (META_JEV_TYPE, "choice"),
173 (META_LABELS, labels_json.as_str()),
174 ])),
175 )
176 }
177 Question::Score(score) => {
178 let legend = score.criteria.clone();
179 let children = Fields::from(vec![
180 Field::new("score", DataType::Float64, false),
181 Field::new("confidence", DataType::Float64, false),
182 fixed_probabilities_field(legend.len()),
183 ]);
184 let legend_json = serde_json::to_string(&legend)
185 .map_err(|e| ApiError::Internal(format!("encode legend metadata: {e}")))?;
186 Ok(
187 Field::new_struct(id, children, false).with_metadata(field_metadata(&[
188 (META_JEV_TYPE, "score"),
189 (META_LEGEND, legend_json.as_str()),
190 ])),
191 )
192 }
193 }
194}
195
196fn sorted_choice_labels(choice: &openkind_core::ChoiceQuestion) -> Vec<String> {
201 let mut labels: Vec<String> = choice.criteria.keys().cloned().collect();
202 labels.sort();
203 labels
204}
205
206fn fixed_probabilities_field(n: usize) -> Field {
208 let item = Field::new("item", DataType::Float64, false);
210 Field::new(
211 "probabilities",
212 DataType::FixedSizeList(Arc::new(item), n as i32),
213 false,
214 )
215}
216
217fn fixed_probabilities_column(
219 flat: Vec<f64>,
220 list_size: usize,
221) -> Result<FixedSizeListArray, ApiError> {
222 FixedSizeListArray::try_new(
223 Arc::new(Field::new("item", DataType::Float64, false)),
224 list_size as i32,
225 Arc::new(Float64Array::from(flat)),
226 None,
227 )
228 .map_err(|e| ApiError::Internal(format!("assemble probabilities: {e}")))
229}
230
231fn schema_metadata(model: &str, input_tokens: u64, output_tokens: u64) -> HashMap<String, String> {
233 HashMap::from([
234 (
235 META_ARROW_VERSION.to_string(),
236 ARROW_MAPPING_VERSION.to_string(),
237 ),
238 (META_MODEL.to_string(), model.to_string()),
239 (
240 META_USAGE_INPUT_TOKENS.to_string(),
241 input_tokens.to_string(),
242 ),
243 (
244 META_USAGE_OUTPUT_TOKENS.to_string(),
245 output_tokens.to_string(),
246 ),
247 ])
248}
249
250fn projected_column_bytes(
252 questions: &BTreeMap<String, Question>,
253 rows: usize,
254) -> Result<usize, ApiError> {
255 let mut per_row = 0usize;
256 for question in questions.values() {
257 let (fixed, width) = match question {
258 Question::Noul(_) => (8usize, 0usize),
259 Question::Choice(q) => (9, q.criteria.len()),
260 Question::Score(q) => (16, q.criteria.len()),
261 };
262 let bytes = width
263 .checked_mul(8)
264 .and_then(|v| v.checked_add(fixed))
265 .ok_or_else(size_error)?;
266 per_row = per_row.checked_add(bytes).ok_or_else(size_error)?;
267 }
268 let bytes = per_row.checked_mul(rows).ok_or_else(size_error)?;
269 if bytes > MAX_ARROW_RESPONSE_BYTES {
270 return Err(size_error());
271 }
272 Ok(bytes)
273}
274
275fn size_error() -> ApiError {
276 ApiError::InvalidBody(format!("Arrow buffers and encoded response must each fit within {MAX_ARROW_RESPONSE_BYTES} bytes; chunk the batch"))
277}
278
279fn deadline_error() -> ApiError {
280 EngineError::DeadlineExceeded {
281 backend: "arrow".to_string(),
282 timeout_ms: ARROW_BATCH_TIMEOUT.as_millis() as u64,
283 }
284 .into()
285}
286
287struct CappedWriter {
289 bytes: Vec<u8>,
290 limit: usize,
291 deadline: tokio::time::Instant,
292 exceeded: bool,
293}
294
295impl Write for CappedWriter {
296 fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
297 if tokio::time::Instant::now() >= self.deadline {
298 return Err(io::Error::new(
299 io::ErrorKind::TimedOut,
300 "Arrow batch deadline exceeded",
301 ));
302 }
303 if self
304 .bytes
305 .len()
306 .checked_add(bytes.len())
307 .is_none_or(|n| n > self.limit)
308 {
309 self.exceeded = true;
310 return Err(io::Error::new(
311 io::ErrorKind::FileTooLarge,
312 "Arrow response limit exceeded",
313 ));
314 }
315 self.bytes
316 .try_reserve_exact(bytes.len())
317 .map_err(|e| io::Error::new(io::ErrorKind::OutOfMemory, e))?;
318 self.bytes.extend_from_slice(bytes);
319 Ok(bytes.len())
320 }
321 fn flush(&mut self) -> io::Result<()> {
322 Ok(())
323 }
324}
325
326fn encode_ipc_stream_until(
327 schema: SchemaRef,
328 batch: RecordBatch,
329 deadline: tokio::time::Instant,
330) -> Result<Vec<u8>, ApiError> {
331 encode_ipc_stream_with_limit(schema, batch, deadline, MAX_ARROW_RESPONSE_BYTES)
332}
333
334fn encode_ipc_stream_with_limit(
335 schema: SchemaRef,
336 batch: RecordBatch,
337 deadline: tokio::time::Instant,
338 limit: usize,
339) -> Result<Vec<u8>, ApiError> {
340 let mut sink = CappedWriter {
341 bytes: Vec::new(),
342 limit,
343 deadline,
344 exceeded: false,
345 };
346 let result = (|| {
347 let mut writer = StreamWriter::try_new(&mut sink, &schema)?;
348 writer.write(&batch)?;
349 writer.finish()
350 })();
351 if tokio::time::Instant::now() >= deadline {
352 return Err(deadline_error());
353 }
354 if sink.exceeded {
355 return Err(size_error());
356 }
357 result.map_err(|e| ApiError::Internal(format!("encode Arrow stream: {e}")))?;
358 Ok(sink.bytes)
359}
360
361#[cfg(test)]
362fn encode_ipc_stream(schema: SchemaRef, batch: RecordBatch) -> Result<Vec<u8>, ApiError> {
363 encode_ipc_stream_until(
364 schema,
365 batch,
366 tokio::time::Instant::now() + ARROW_BATCH_TIMEOUT,
367 )
368}
369
370#[cfg(test)]
371#[path = "arrow_tests.rs"]
372mod tests;
373
374#[cfg(test)]
375#[path = "arrow_regression_tests.rs"]
376mod regression_tests;
377
378pub async fn arrow_batch(
385 AxumState(state): AxumState<Arc<AppState>>,
386 rate_limit: Option<Extension<crate::middleware::RateLimitContext>>,
387 req: Result<Json<ArrowBatchRequest>, JsonRejection>,
388) -> Result<Response, ApiError> {
389 let Json(req) = match req {
390 Ok(json) => json,
391 Err(rejection) => match rejection {
392 JsonRejection::BytesRejection(e) => {
393 return Err(ApiError::PayloadTooLarge(e.to_string()))
394 }
395 JsonRejection::JsonSyntaxError(e) => return Err(ApiError::BadJson(e.to_string())),
396 JsonRejection::JsonDataError(e) => return Err(ApiError::InvalidBody(e.to_string())),
397 other => return Err(ApiError::InvalidBody(other.to_string())),
398 },
399 };
400
401 evaluate_batch(
402 state,
403 req,
404 rate_limit.map(|Extension(context)| context),
405 ARROW_BATCH_TIMEOUT,
406 )
407 .await
408}
409
410async fn evaluate_batch(
411 state: Arc<AppState>,
412 req: ArrowBatchRequest,
413 rate_limit: Option<crate::middleware::RateLimitContext>,
414 timeout: Duration,
415) -> Result<Response, ApiError> {
416 let deadline = tokio::time::Instant::now() + timeout;
417 tokio::time::timeout_at(
418 deadline,
419 evaluate_batch_until(state, req, rate_limit, deadline),
420 )
421 .await
422 .map_err(|_| deadline_error())?
423}
424
425async fn evaluate_batch_until(
426 state: Arc<AppState>,
427 req: ArrowBatchRequest,
428 rate_limit: Option<crate::middleware::RateLimitContext>,
429 deadline: tokio::time::Instant,
430) -> Result<Response, ApiError> {
431 if req.states.len() > MAX_ARROW_STATES {
432 return Err(ApiError::InvalidBody(format!(
433 "states array has {} entries; at most {MAX_ARROW_STATES} are supported",
434 req.states.len()
435 )));
436 }
437 let representative = req.system_request_for(&State::Text(String::new()));
439 openkind_core::validate_request(&representative)
440 .map_err(|e| ApiError::Engine(EngineError::Invalid(e)))?;
441 drop(representative);
442 let engine = state
443 .registry
444 .get(&req.model)
445 .ok_or_else(|| EngineError::UnknownModel(req.model.clone()))?;
446 let mut registry = EngineRegistry::new();
449 registry.register(req.model.clone(), engine);
450 let projected_bytes = projected_column_bytes(&req.questions, req.states.len())?;
451 let work_units = arrow_work_units(req.states.len(), projected_bytes);
452 if let Some(rate_limit) = rate_limit {
453 rate_limit
454 .charge(work_units.saturating_sub(1) as u32)
455 .map_err(|retry_after_ms| ApiError::RateLimited { retry_after_ms })?;
456 }
457 let _admission = if work_units == 0 {
460 None
461 } else {
462 Some(
463 state
464 .arrow_admission
465 .clone()
466 .acquire_many_owned(work_units as u32)
467 .await
468 .map_err(|_| deadline_error())?,
469 )
470 };
471 let mut builder = BatchBuilder::new(&req.questions, &req.model, req.states.len())?;
472 let (schema, empty) = builder.empty_batch()?;
474 tokio::task::spawn_blocking(move || encode_ipc_stream_until(schema, empty, deadline))
475 .await
476 .map_err(|e| ApiError::Internal(format!("Arrow schema task failed: {e}")))??;
477 for jev_state in &req.states {
478 tokio::task::yield_now().await;
481 if tokio::time::Instant::now() >= deadline {
482 return Err(deadline_error());
483 }
484 let response = dispatch(req.system_request_for(jev_state), ®istry).await?;
485 builder.push(&response)?;
486 }
487 let (schema, batch) = builder.finish()?;
488 let bytes =
489 tokio::task::spawn_blocking(move || encode_ipc_stream_until(schema, batch, deadline))
490 .await
491 .map_err(|e| ApiError::Internal(format!("Arrow encoding task failed: {e}")))??;
492 Ok((
493 StatusCode::OK,
494 [(header::CONTENT_TYPE, ARROW_CONTENT_TYPE)],
495 bytes,
496 )
497 .into_response())
498}
499
500fn arrow_work_units(states: usize, projected_bytes: usize) -> usize {
502 let bytes_per_unit = MAX_ARROW_RESPONSE_BYTES.div_ceil(ARROW_ADMISSION_UNITS);
503 states.max(projected_bytes.div_ceil(bytes_per_unit))
504}
505
506#[cfg(test)]
507fn build_batch(
508 questions: &BTreeMap<String, Question>,
509 responses: &[openkind_core::SystemResponse],
510 fallback_model: &str,
511) -> Result<(SchemaRef, RecordBatch), ApiError> {
512 projected_column_bytes(questions, responses.len())?;
513 let mut builder = BatchBuilder::new(questions, fallback_model, responses.len())?;
514 for response in responses {
515 builder.push(response)?;
516 }
517 builder.finish()
518}