Skip to main content

openkind_api/
arrow.rs

1//! Unofficial Arrow bulk endpoint: `POST /v1/arrow`.
2//!
3//! Inspired by the proposal in ["What if Jev spoke Arrow?"](https://columnar.tech/blog/what-if-jev-spoke-arrow/):
4//! one request carries many independent states plus a single shared question
5//! map, and the answer comes back as an [Apache Arrow](https://arrow.apache.org/)
6//! IPC stream: one row per state, one column per question.
7//!
8//! This surface is **not** part of the TypeSafe wire contract. It is
9//! flag-gated (`openkindd --arrow on`), absent from `openapi.yaml`, and must
10//! never leak into the `/v1/systemone` shapes or the SDK compatibility
11//! tests. See [`docs/ARROW.md`](../../../docs/ARROW.md) for the full
12//! mapping and its limits.
13//!
14//! Wire shape (JSON in, Arrow out):
15//!
16//! ```json
17//! {
18//!   "model": "jev-latest",
19//!   "states": ["Please refund the shoes.", {"cart": ["shoes"]}],
20//!   "questions": { "refund": { "type": "noul", "instructions": "Is a refund being requested?" } }
21//! }
22//! ```
23//!
24//! Each question id becomes an Arrow column whose type encodes the answer:
25//!
26//! | Jev type | Arrow type | Field metadata |
27//! |---|---|---|
28//! | Noul | `float64` | `jev.type = "noul"` |
29//! | Choice | `struct<choice: uint8, confidence: float64, probabilities: fixed_size_list<float64>[N]>` | `jev.type = "choice"`, `labels` (JSON array, sorted option keys) |
30//! | Score | `struct<score: float64, confidence: float64, probabilities: fixed_size_list<float64>[N]>` | `jev.type = "score"`, `legend` (JSON array of level descriptions in rubric order) |
31//!
32//! Columns are ordered by question id (lexicographic). Row `i` is the answer
33//! for `states[i]`. All fields are non-nullable: the endpoint evaluates
34//! every state or fails the whole request.
35
36use 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
63/// Maximum projected column-buffer size and encoded IPC response size (64 MiB).
64pub const MAX_ARROW_RESPONSE_BYTES: usize = 64 * 1024 * 1024;
65/// Queue-inclusive deadline for the complete bulk evaluation and encoding.
66pub const ARROW_BATCH_TIMEOUT: Duration = Duration::from_secs(600);
67
68/// `Content-Type` of the Arrow IPC stream response.
69pub const ARROW_CONTENT_TYPE: &str = "application/vnd.apache.arrow.stream";
70
71/// Maximum number of states in one bulk request. Mirrors the DoS budget of
72/// [`openkind_core::MAX_QUESTIONS_PER_REQUEST`]; beyond this, callers chunk
73/// their states across requests.
74pub const MAX_ARROW_STATES: usize = 10_000;
75/// Total weighted Arrow work admitted concurrently by one application state.
76pub(crate) const ARROW_ADMISSION_UNITS: usize = MAX_ARROW_STATES;
77
78/// Maximum number of Choice options a question may have on this endpoint.
79/// The `choice` child column is a `uint8` index into `labels`, so more than
80/// 256 options cannot be represented; core allows up to
81/// [`openkind_core::MAX_CRITERIA_OPTIONS`] on `/v1/systemone`.
82pub const MAX_CHOICE_LABELS: usize = u8::MAX as usize + 1;
83
84/// Schema metadata: format version of this unofficial mapping.
85pub const META_ARROW_VERSION: &str = "openkind.arrow.version";
86/// Schema metadata: model that performed the evaluation.
87pub const META_MODEL: &str = "openkind.model";
88/// Schema metadata: aggregate input tokens across all states.
89pub const META_USAGE_INPUT_TOKENS: &str = "openkind.usage.input_tokens";
90/// Schema metadata: aggregate output tokens across all states.
91pub const META_USAGE_OUTPUT_TOKENS: &str = "openkind.usage.output_tokens";
92/// Field metadata on every question column: the Jev question type.
93pub const META_JEV_TYPE: &str = "jev.type";
94/// Field metadata on Choice columns: JSON array of option keys, sorted.
95pub const META_LABELS: &str = "labels";
96/// Field metadata on Score columns: JSON array of level descriptions, in rubric order.
97pub const META_LEGEND: &str = "legend";
98
99/// Value of [`META_ARROW_VERSION`] in every response stream.
100pub const ARROW_MAPPING_VERSION: &str = "1";
101
102/// Bulk evaluation request body for `POST /v1/arrow`.
103///
104/// The shape mirrors [`SystemRequest`] with `state` generalized to
105/// `states`; the question map is shared by every state, exactly as the
106/// article proposes. Unknown top-level fields are ignored, matching
107/// `/v1/systemone` tolerance for extra metadata.
108#[derive(Debug, Clone, Deserialize)]
109pub struct ArrowBatchRequest {
110    /// Required. `"jev-latest"` or a registered model alias.
111    pub model: String,
112    /// Required. One entry per output row; may be empty (empty batch).
113    pub states: Vec<State>,
114    /// Required. Map of user-chosen id → typed question, shared by all states.
115    pub questions: BTreeMap<String, Question>,
116}
117
118impl ArrowBatchRequest {
119    /// Build the per-state [`SystemRequest`] fanned out to the engine.
120    ///
121    /// The question map is cloned once per state because `dispatch` takes
122    /// ownership; column/row semantics are unaffected.
123    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)]
137/// Ordered question ids of this request (lexicographic: the column order).
138fn question_ids(questions: &BTreeMap<String, Question>) -> Vec<&str> {
139    questions.keys().map(String::as_str).collect()
140}
141
142/// Build a Jev/Arrow field metadata map.
143fn 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
150/// Build the Arrow field for one question column, including its metadata.
151fn 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
196/// Sorted option keys of a Choice question: the canonical label order.
197///
198/// `criteria` is a wire `HashMap`, so the sorted order is the only
199/// deterministic projection shared by producer and consumer.
200fn 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
206/// Non-null `fixed_size_list<float64>[n]` child for the probability vector.
207fn fixed_probabilities_field(n: usize) -> Field {
208    // Core validation bounds the width before the schema is constructed.
209    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
217/// Build a validated non-null probability array.
218fn 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
231/// Canonical metadata map for the response schema.
232fn 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
250/// Checked buffer projection, before allocating columns or evaluating states.
251fn 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
287/// Limit the writer itself so schema overhead cannot escape the projection.
288struct 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
378/// Handler for the unofficial `POST /v1/arrow` bulk endpoint.
379///
380/// Evaluates every state against the shared question map and returns one
381/// Arrow IPC stream. Errors are all-or-nothing: any per-state failure fails
382/// the whole request with the standard JSON error envelope, before any
383/// Arrow bytes are written.
384pub 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    // Empty batches still have a question contract. No engine is evaluated here.
438    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    // Pin the handle for the whole batch while preserving dispatch validation,
447    // usage estimation and telemetry. Playground updates affect later batches.
448    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    // Admission precedes all projected column allocation. A maximum-work
458    // request takes the whole gate; smaller batches may run concurrently.
459    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    // Check schema size before doing model work, even for a zero-row batch.
473    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        // Mock engines can return ready futures; yield so cancellation and
479        // other requests remain observable between independent evaluations.
480        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), &registry).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
500/// Estimate both dispatch work and retained column memory on a common scale.
501fn 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}