use std::collections::{BTreeMap, HashMap};
use std::io::{self, Write};
use std::sync::Arc;
use std::time::Duration;
use arrow_array::{FixedSizeListArray, Float64Array, RecordBatch};
use arrow_ipc::writer::StreamWriter;
use arrow_schema::{DataType, Field, FieldRef, Fields, Schema, SchemaRef};
use axum::extract::rejection::JsonRejection;
use axum::extract::{Extension, State as AxumState};
use axum::http::{header, StatusCode};
use axum::response::{IntoResponse, Response};
use axum::Json;
use openkind_core::{Question, State, SystemRequest};
use openkind_engine::{dispatch, EngineError, EngineRegistry};
use serde::Deserialize;
use crate::error::ApiError;
use crate::AppState;
#[path = "arrow_columns.rs"]
mod columns;
#[path = "arrow_decoder.rs"]
mod decoder;
use columns::BatchBuilder;
pub use decoder::answers_from_batch;
pub const MAX_ARROW_RESPONSE_BYTES: usize = 64 * 1024 * 1024;
pub const ARROW_BATCH_TIMEOUT: Duration = Duration::from_secs(600);
pub const ARROW_CONTENT_TYPE: &str = "application/vnd.apache.arrow.stream";
pub const MAX_ARROW_STATES: usize = 10_000;
pub(crate) const ARROW_ADMISSION_UNITS: usize = MAX_ARROW_STATES;
pub const MAX_CHOICE_LABELS: usize = u8::MAX as usize + 1;
pub const META_ARROW_VERSION: &str = "openkind.arrow.version";
pub const META_MODEL: &str = "openkind.model";
pub const META_USAGE_INPUT_TOKENS: &str = "openkind.usage.input_tokens";
pub const META_USAGE_OUTPUT_TOKENS: &str = "openkind.usage.output_tokens";
pub const META_JEV_TYPE: &str = "jev.type";
pub const META_LABELS: &str = "labels";
pub const META_LEGEND: &str = "legend";
pub const ARROW_MAPPING_VERSION: &str = "1";
#[derive(Debug, Clone, Deserialize)]
pub struct ArrowBatchRequest {
pub model: String,
pub states: Vec<State>,
pub questions: BTreeMap<String, Question>,
}
impl ArrowBatchRequest {
fn system_request_for(&self, state: &State) -> SystemRequest {
SystemRequest {
state: state.clone(),
model: self.model.clone(),
questions: self
.questions
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect(),
}
}
}
#[cfg(test)]
fn question_ids(questions: &BTreeMap<String, Question>) -> Vec<&str> {
questions.keys().map(String::as_str).collect()
}
fn field_metadata(pairs: &[(&str, &str)]) -> HashMap<String, String> {
pairs
.iter()
.map(|(key, value)| ((*key).to_string(), (*value).to_string()))
.collect()
}
fn field_for_question(id: &str, question: &Question) -> Result<Field, ApiError> {
match question {
Question::Noul(_) => Ok(Field::new(id, DataType::Float64, false)
.with_metadata(field_metadata(&[(META_JEV_TYPE, "noul")]))),
Question::Choice(choice) => {
let labels = sorted_choice_labels(choice);
if labels.len() > MAX_CHOICE_LABELS {
return Err(ApiError::InvalidBody(format!(
"choice question `{id}` has {} options; the Arrow mapping indexes `choice` into a uint8, so at most {MAX_CHOICE_LABELS} options are supported",
labels.len()
)));
}
let children = Fields::from(vec![
Field::new("choice", DataType::UInt8, false),
Field::new("confidence", DataType::Float64, false),
fixed_probabilities_field(labels.len()),
]);
let labels_json = serde_json::to_string(&labels)
.map_err(|e| ApiError::Internal(format!("encode labels metadata: {e}")))?;
Ok(
Field::new_struct(id, children, false).with_metadata(field_metadata(&[
(META_JEV_TYPE, "choice"),
(META_LABELS, labels_json.as_str()),
])),
)
}
Question::Score(score) => {
let legend = score.criteria.clone();
let children = Fields::from(vec![
Field::new("score", DataType::Float64, false),
Field::new("confidence", DataType::Float64, false),
fixed_probabilities_field(legend.len()),
]);
let legend_json = serde_json::to_string(&legend)
.map_err(|e| ApiError::Internal(format!("encode legend metadata: {e}")))?;
Ok(
Field::new_struct(id, children, false).with_metadata(field_metadata(&[
(META_JEV_TYPE, "score"),
(META_LEGEND, legend_json.as_str()),
])),
)
}
}
}
fn sorted_choice_labels(choice: &openkind_core::ChoiceQuestion) -> Vec<String> {
let mut labels: Vec<String> = choice.criteria.keys().cloned().collect();
labels.sort();
labels
}
fn fixed_probabilities_field(n: usize) -> Field {
let item = Field::new("item", DataType::Float64, false);
Field::new(
"probabilities",
DataType::FixedSizeList(Arc::new(item), n as i32),
false,
)
}
fn fixed_probabilities_column(
flat: Vec<f64>,
list_size: usize,
) -> Result<FixedSizeListArray, ApiError> {
FixedSizeListArray::try_new(
Arc::new(Field::new("item", DataType::Float64, false)),
list_size as i32,
Arc::new(Float64Array::from(flat)),
None,
)
.map_err(|e| ApiError::Internal(format!("assemble probabilities: {e}")))
}
fn schema_metadata(model: &str, input_tokens: u64, output_tokens: u64) -> HashMap<String, String> {
HashMap::from([
(
META_ARROW_VERSION.to_string(),
ARROW_MAPPING_VERSION.to_string(),
),
(META_MODEL.to_string(), model.to_string()),
(
META_USAGE_INPUT_TOKENS.to_string(),
input_tokens.to_string(),
),
(
META_USAGE_OUTPUT_TOKENS.to_string(),
output_tokens.to_string(),
),
])
}
fn projected_column_bytes(
questions: &BTreeMap<String, Question>,
rows: usize,
) -> Result<usize, ApiError> {
let mut per_row = 0usize;
for question in questions.values() {
let (fixed, width) = match question {
Question::Noul(_) => (8usize, 0usize),
Question::Choice(q) => (9, q.criteria.len()),
Question::Score(q) => (16, q.criteria.len()),
};
let bytes = width
.checked_mul(8)
.and_then(|v| v.checked_add(fixed))
.ok_or_else(size_error)?;
per_row = per_row.checked_add(bytes).ok_or_else(size_error)?;
}
let bytes = per_row.checked_mul(rows).ok_or_else(size_error)?;
if bytes > MAX_ARROW_RESPONSE_BYTES {
return Err(size_error());
}
Ok(bytes)
}
fn size_error() -> ApiError {
ApiError::InvalidBody(format!("Arrow buffers and encoded response must each fit within {MAX_ARROW_RESPONSE_BYTES} bytes; chunk the batch"))
}
fn deadline_error() -> ApiError {
EngineError::DeadlineExceeded {
backend: "arrow".to_string(),
timeout_ms: ARROW_BATCH_TIMEOUT.as_millis() as u64,
}
.into()
}
struct CappedWriter {
bytes: Vec<u8>,
limit: usize,
deadline: tokio::time::Instant,
exceeded: bool,
}
impl Write for CappedWriter {
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
if tokio::time::Instant::now() >= self.deadline {
return Err(io::Error::new(
io::ErrorKind::TimedOut,
"Arrow batch deadline exceeded",
));
}
if self
.bytes
.len()
.checked_add(bytes.len())
.is_none_or(|n| n > self.limit)
{
self.exceeded = true;
return Err(io::Error::new(
io::ErrorKind::FileTooLarge,
"Arrow response limit exceeded",
));
}
self.bytes
.try_reserve_exact(bytes.len())
.map_err(|e| io::Error::new(io::ErrorKind::OutOfMemory, e))?;
self.bytes.extend_from_slice(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
fn encode_ipc_stream_until(
schema: SchemaRef,
batch: RecordBatch,
deadline: tokio::time::Instant,
) -> Result<Vec<u8>, ApiError> {
encode_ipc_stream_with_limit(schema, batch, deadline, MAX_ARROW_RESPONSE_BYTES)
}
fn encode_ipc_stream_with_limit(
schema: SchemaRef,
batch: RecordBatch,
deadline: tokio::time::Instant,
limit: usize,
) -> Result<Vec<u8>, ApiError> {
let mut sink = CappedWriter {
bytes: Vec::new(),
limit,
deadline,
exceeded: false,
};
let result = (|| {
let mut writer = StreamWriter::try_new(&mut sink, &schema)?;
writer.write(&batch)?;
writer.finish()
})();
if tokio::time::Instant::now() >= deadline {
return Err(deadline_error());
}
if sink.exceeded {
return Err(size_error());
}
result.map_err(|e| ApiError::Internal(format!("encode Arrow stream: {e}")))?;
Ok(sink.bytes)
}
#[cfg(test)]
fn encode_ipc_stream(schema: SchemaRef, batch: RecordBatch) -> Result<Vec<u8>, ApiError> {
encode_ipc_stream_until(
schema,
batch,
tokio::time::Instant::now() + ARROW_BATCH_TIMEOUT,
)
}
#[cfg(test)]
#[path = "arrow_tests.rs"]
mod tests;
#[cfg(test)]
#[path = "arrow_regression_tests.rs"]
mod regression_tests;
pub async fn arrow_batch(
AxumState(state): AxumState<Arc<AppState>>,
rate_limit: Option<Extension<crate::middleware::RateLimitContext>>,
req: Result<Json<ArrowBatchRequest>, JsonRejection>,
) -> Result<Response, ApiError> {
let Json(req) = match req {
Ok(json) => json,
Err(rejection) => match rejection {
JsonRejection::BytesRejection(e) => {
return Err(ApiError::PayloadTooLarge(e.to_string()))
}
JsonRejection::JsonSyntaxError(e) => return Err(ApiError::BadJson(e.to_string())),
JsonRejection::JsonDataError(e) => return Err(ApiError::InvalidBody(e.to_string())),
other => return Err(ApiError::InvalidBody(other.to_string())),
},
};
evaluate_batch(
state,
req,
rate_limit.map(|Extension(context)| context),
ARROW_BATCH_TIMEOUT,
)
.await
}
async fn evaluate_batch(
state: Arc<AppState>,
req: ArrowBatchRequest,
rate_limit: Option<crate::middleware::RateLimitContext>,
timeout: Duration,
) -> Result<Response, ApiError> {
let deadline = tokio::time::Instant::now() + timeout;
tokio::time::timeout_at(
deadline,
evaluate_batch_until(state, req, rate_limit, deadline),
)
.await
.map_err(|_| deadline_error())?
}
async fn evaluate_batch_until(
state: Arc<AppState>,
req: ArrowBatchRequest,
rate_limit: Option<crate::middleware::RateLimitContext>,
deadline: tokio::time::Instant,
) -> Result<Response, ApiError> {
if req.states.len() > MAX_ARROW_STATES {
return Err(ApiError::InvalidBody(format!(
"states array has {} entries; at most {MAX_ARROW_STATES} are supported",
req.states.len()
)));
}
let representative = req.system_request_for(&State::Text(String::new()));
openkind_core::validate_request(&representative)
.map_err(|e| ApiError::Engine(EngineError::Invalid(e)))?;
drop(representative);
let engine = state
.registry
.get(&req.model)
.ok_or_else(|| EngineError::UnknownModel(req.model.clone()))?;
let mut registry = EngineRegistry::new();
registry.register(req.model.clone(), engine);
let projected_bytes = projected_column_bytes(&req.questions, req.states.len())?;
let work_units = arrow_work_units(req.states.len(), projected_bytes);
if let Some(rate_limit) = rate_limit {
rate_limit
.charge(work_units.saturating_sub(1) as u32)
.map_err(|retry_after_ms| ApiError::RateLimited { retry_after_ms })?;
}
let _admission = if work_units == 0 {
None
} else {
Some(
state
.arrow_admission
.clone()
.acquire_many_owned(work_units as u32)
.await
.map_err(|_| deadline_error())?,
)
};
let mut builder = BatchBuilder::new(&req.questions, &req.model, req.states.len())?;
let (schema, empty) = builder.empty_batch()?;
tokio::task::spawn_blocking(move || encode_ipc_stream_until(schema, empty, deadline))
.await
.map_err(|e| ApiError::Internal(format!("Arrow schema task failed: {e}")))??;
for jev_state in &req.states {
tokio::task::yield_now().await;
if tokio::time::Instant::now() >= deadline {
return Err(deadline_error());
}
let response = dispatch(req.system_request_for(jev_state), ®istry).await?;
builder.push(&response)?;
}
let (schema, batch) = builder.finish()?;
let bytes =
tokio::task::spawn_blocking(move || encode_ipc_stream_until(schema, batch, deadline))
.await
.map_err(|e| ApiError::Internal(format!("Arrow encoding task failed: {e}")))??;
Ok((
StatusCode::OK,
[(header::CONTENT_TYPE, ARROW_CONTENT_TYPE)],
bytes,
)
.into_response())
}
fn arrow_work_units(states: usize, projected_bytes: usize) -> usize {
let bytes_per_unit = MAX_ARROW_RESPONSE_BYTES.div_ceil(ARROW_ADMISSION_UNITS);
states.max(projected_bytes.div_ceil(bytes_per_unit))
}
#[cfg(test)]
fn build_batch(
questions: &BTreeMap<String, Question>,
responses: &[openkind_core::SystemResponse],
fallback_model: &str,
) -> Result<(SchemaRef, RecordBatch), ApiError> {
projected_column_bytes(questions, responses.len())?;
let mut builder = BatchBuilder::new(questions, fallback_model, responses.len())?;
for response in responses {
builder.push(response)?;
}
builder.finish()
}