aprender-serve 0.70.2

Pure Rust ML inference engine built from scratch - model serving for GGUF and safetensors

/// Dense f32 `Model` backend for `POST /stream/generate` (registry / safetensors).
///
/// Unchanged behaviour, lifted out of the handler so the quantized backend can be
/// tried first. Only the status for an unresolvable model moved: a server with no
/// usable model is 503, not 404 (see `model_resolution_status`).
fn dense_stream_tokens(
    state: &AppState,
    request: &GenerateRequest,
    cancel: &CancelToken,
) -> Result<(Vec<u32>, usize, std::sync::Arc<BPETokenizer>), ApiErr> {
    let (model, tokenizer) = state
        .get_model(request.model_id.as_deref())
        .map_err(|e| api_err(super::model_resolution_status(&e), e))?;

    let prompt_ids = tokenize_prompt(&tokenizer, &request.prompt)?;
    let prompt: Vec<usize> = prompt_ids.iter().map(|&id| id as usize).collect();
    let prompt_len = prompt.len();

    let strategy = match request.strategy.as_str() {
        "greedy" => SamplingStrategy::Greedy,
        "top_k" => SamplingStrategy::TopK { k: request.top_k },
        "top_p" => SamplingStrategy::TopP { p: request.top_p },
        other => {
            return Err(api_err(
                StatusCode::BAD_REQUEST,
                format!("Invalid strategy: {other}"),
            ))
        },
    };

    let mut config = GenerationConfig::default()
        .with_max_tokens(request.max_tokens)
        .with_temperature(request.temperature)
        .with_cancel(cancel.clone());
    config.strategy = strategy;
    if let Some(seed) = request.seed {
        config = config.with_seed(seed);
    }

    let generated = model
        .generate(&prompt, &config)
        .map_err(|e| api_err(crate::api::generation_error_status(&e), e))?;

    let token_ids: Vec<u32> = generated
        .iter()
        .map(|&id| {
            u32::try_from(id).map_err(|_| {
                api_err(
                    StatusCode::BAD_REQUEST,
                    format!("Token ID {id} exceeds u32 range"),
                )
            })
        })
        .collect::<Result<Vec<_>, _>>()?;

    Ok((token_ids, prompt_len, tokenizer))
}

/// CUDA GGUF backend for `POST /stream/generate` and `/realize/generate` (#3991).
///
/// `apr serve --gpu model.gguf` holds a `cuda_model` and no `quantized_model`, so
/// this handler fell through to the dense registry and answered 503 "No model
/// available" while `GET /` listed both routes. Same generation call and config as
/// `/generate`'s `try_cuda_generate`; the whole sequence is produced before the SSE
/// stream is built, exactly as for the other backends here.
#[cfg(feature = "cuda")]
fn try_cuda_stream_tokens(
    state: &AppState,
    request: &GenerateRequest,
    cancel: &CancelToken,
) -> Result<Option<(Vec<u32>, usize, std::sync::Arc<BPETokenizer>)>, ApiErr> {
    use crate::gguf::QuantizedGenerateConfig;

    let Some(cuda_model_lock) = state.cuda_model() else {
        return Ok(None);
    };
    let tokenizer = require_tok(state)?;
    let prompt_ids = tokenize_prompt(&tokenizer, &request.prompt)?;
    let prompt_len = prompt_ids.len();
    // D5: over the serving context is the request's fault, so a 400.
    if let Some(msg) = state
        .serving_context()
        .and_then(|ctx| super::serve_context_refusal(prompt_len, ctx))
    {
        return Err(api_err(StatusCode::BAD_REQUEST, msg));
    }

    let q_config = QuantizedGenerateConfig {
        max_tokens: request.max_tokens,
        temperature: request.temperature,
        top_k: if request.temperature == 0.0 {
            1
        } else {
            request.top_k
        },
        stop_tokens: vec![eos_id(&tokenizer, state.model_eos_token_id())],
        trace: false,
        cancel: cancel.clone(),
        ..Default::default()
    };

    let mut cuda_model = cuda_model_lock.write().map_err(|_| {
        api_err(
            StatusCode::INTERNAL_SERVER_ERROR,
            "Failed to acquire CUDA model lock",
        )
    })?;
    let generated = cuda_model
        .generate_gpu_resident(&prompt_ids, &q_config)
        .map_err(|e| {
            api_err(
                super::generation_error_status(&e),
                format!("CUDA generation failed: {e}"),
            )
        })?;
    Ok(Some((generated, prompt_len, tokenizer)))
}

/// Quantized (GGUF / APR Q4_K) backend for `POST /stream/generate` and
/// `/realize/generate`.
///
/// aprender#2376(1, 10): this handler resolved the dense f32 `Model` via
/// `get_model()`, which is `None` on every `apr serve run model.gguf`, so a route
/// the startup banner advertises as "SSE streaming" answered 404
/// `"No model available"` while `/health` reported `model_loaded:true` and
/// `/generate` on the same process returned tokens.
///
/// Returns `Ok(None)` when no quantized model is resident, so the dense path below
/// is unchanged.
fn try_quantized_stream_tokens(
    state: &AppState,
    request: &GenerateRequest,
    cancel: &CancelToken,
) -> Result<Option<(Vec<u32>, usize, std::sync::Arc<BPETokenizer>)>, ApiErr> {
    let quantized_model = match state.quantized_model() {
        Some(m) => m,
        None => return Ok(None),
    };
    let tokenizer = require_tok(state)?;
    let prompt_ids = tokenize_prompt(&tokenizer, &request.prompt)?;
    let prompt_len = prompt_ids.len();

    let sampling = resolve_quantized_sampling(
        &request.strategy,
        request.top_k,
        request.top_p,
        request.temperature,
    )?;
    let q_config = quantized_config(
        state,
        &tokenizer,
        request.max_tokens,
        request.temperature,
        &sampling,
        request.seed,
        cancel,
    );

    let generated = quantized_model
        .generate_with_cache(&prompt_ids, &q_config)
        .map_err(|e| generation_err(&e))?;

    Ok(Some((generated, prompt_len, tokenizer)))
}

/// `AprTransformer` (f32 APR / SafeTensors CPU) backend for `POST /stream/generate`
/// and `/realize/generate`.
///
/// aprender#2609: `/generate` and `/batch/generate` grew this backend
/// (`try_apr_generate`, `try_apr_batch_generate`) and `/stream/generate` did not,
/// so on an `AprTransformer` server the SSE route the startup banner advertises
/// answered `"No model available"` while `/generate` on the SAME process returned
/// real tokens and `/health` reported `model_loaded: true`. The backend chain here
/// is now the same one `/generate` walks: quantized, then APR, then dense.
///
/// Returns `Ok(None)` when no `AprTransformer` is resident, so the dense path is
/// unchanged.
fn try_apr_stream_tokens(
    state: &AppState,
    request: &GenerateRequest,
    cancel: &CancelToken,
) -> Result<Option<(Vec<u32>, usize, std::sync::Arc<BPETokenizer>)>, ApiErr> {
    use crate::apr_transformer::GenerateConfig;

    let apr_transformer = match state.apr_transformer() {
        Some(m) => m,
        None => return Ok(None),
    };
    let tokenizer = require_tok(state)?;
    let prompt_ids = tokenize_prompt(&tokenizer, &request.prompt)?;
    let prompt_len = prompt_ids.len();

    let gen_config = GenerateConfig {
        max_tokens: request.max_tokens,
        temperature: request.temperature,
        cancel: cancel.clone(),
        ..Default::default()
    };

    let generated = apr_transformer
        .generate_with_cache(&prompt_ids, &gen_config)
        .map_err(|e| {
            api_err(
                super::generation_error_status(&e),
                format!("APR generation failed: {e}"),
            )
        })?;

    Ok(Some((generated, prompt_len, tokenizer)))
}

/// The `event: token` payloads for generated ids, never split inside a character.
///
/// Decoding each id alone turned a character spanning two byte tokens into two
/// U+FFFD (the defect a29989039 fixed on the chat SSE path). A held id emits no
/// event; the event that completes the character carries the text of every id
/// it covers and the LAST id's `token_id`.
fn raw_stream_token_events(tokenizer: &BPETokenizer, generated: &[u32]) -> Vec<StreamTokenEvent> {
    let mut utf8 = crate::api::LiveUtf8Deltas::new();
    let mut events: Vec<StreamTokenEvent> = generated
        .iter()
        .filter_map(|&token_id| {
            utf8.push(tokenizer, token_id)
                .map(|text| StreamTokenEvent { token_id, text })
        })
        .collect();
    if let (Some(text), Some(&token_id)) = (utf8.finish(tokenizer), generated.last()) {
        events.push(StreamTokenEvent { token_id, text });
    }
    events
}

/// Cut a pregenerated token stream at the first special-token marker the
/// completion spells out (aprender#4344), keeping one event per token.
///
/// Every token before the marker is kept as is. The token the marker starts in
/// is kept with its text cut at the marker, or dropped if nothing precedes the
/// marker in it. Everything after is dropped. `/stream/generate` and
/// `/realize/generate` streamed `"<answer>7</answer><|im_end|>"` before this.
fn cut_pieces_at_marker(pieces: Vec<(u32, String)>) -> Vec<(u32, String)> {
    let full: String = pieces.iter().map(|(_, t)| t.as_str()).collect();
    let Some(cut) = crate::api::realize_handlers::first_special_marker(&full) else {
        return pieces;
    };
    let mut out = Vec::new();
    let mut offset = 0usize;
    for (id, mut text) in pieces {
        if offset >= cut {
            break;
        }
        let end = offset + text.len();
        if end > cut {
            text.truncate(cut - offset);
            if !text.is_empty() {
                out.push((id, text));
            }
            break;
        }
        offset = end;
        out.push((id, text));
    }
    out
}

/// Stream generate handler — generates tokens one by one via Server-Sent Events.
///
/// Tries the quantized backend first (the `apr serve run model.gguf` path), then
/// the `AprTransformer` (f32 APR / SafeTensors), then the dense f32 `Model`.
///
/// aprender#2376(3): the whole sequence is generated *before* the SSE stream is
/// built, so this handler's synchronous decode is exactly the work an abandoned
/// request used to keep doing. `cancel` (minted per request by
/// `cancel_on_disconnect`) is installed on the config so the loop stops at its
/// next token boundary once the client goes away.
pub async fn stream_generate_handler(
    State(state): State<AppState>,
    Extension(cancel): Extension<CancelToken>,
    Json(request): Json<GenerateRequest>,
) -> Result<Sse<impl Stream<Item = Result<Event, Infallible>>>, (StatusCode, Json<ErrorResponse>)> {
    // #3991: CUDA first, as `/generate` does — a `--gpu` GGUF server has no
    // `quantized_model`, and this handler used to fall through to a 503.
    #[cfg(feature = "cuda")]
    let cuda = try_cuda_stream_tokens(&state, &request, &cancel)?;
    #[cfg(not(feature = "cuda"))]
    let cuda = None;

    let (token_ids, prompt_len, tokenizer_clone) = if let Some(resolved) = cuda {
        resolved
    } else if let Some(resolved) = try_quantized_stream_tokens(&state, &request, &cancel)? {
        resolved
    } else if let Some(resolved) = try_apr_stream_tokens(&state, &request, &cancel)? {
        resolved
    } else if let Some(resolved) = try_qwen35_stream_tokens(&state, &request, &cancel)? {
        resolved
    } else {
        dense_stream_tokens(&state, &request, &cancel)?
    };

    // Create stream that emits tokens one by one
    let stream = async_stream::stream! {
        // Skip prompt tokens, only stream generated tokens. A backend that stops
        // on the first sampled token returns the prompt alone, so clamp rather
        // than slice past the end.
        let generated_start = prompt_len.min(token_ids.len());
        // #4344: cut at a spelled-out special-token marker, after the UTF-8
        // reassembly (#3962) so the cut sees whole characters.
        let pieces: Vec<(u32, String)> =
            raw_stream_token_events(&tokenizer_clone, &token_ids[generated_start..])
                .into_iter()
                .map(|e| (e.token_id, e.text))
                .collect();
        for (token_id, text) in cut_pieces_at_marker(pieces) {
            let event = StreamTokenEvent { token_id, text };
            // Serialization of simple struct should not fail, but handle gracefully
            let data = serde_json::to_string(&event)
                .unwrap_or_else(|_| r#"{"error":"serialization failed"}"#.to_string());

            yield Ok::<_, Infallible>(Event::default().event("token").data(data));
        }

        // Send done event
        let done_event = StreamDoneEvent {
            num_generated: token_ids.len().saturating_sub(prompt_len),
        };
        // Serialization of simple struct should not fail, but handle gracefully
        let data = serde_json::to_string(&done_event)
            .unwrap_or_else(|_| r#"{"error":"serialization failed"}"#.to_string());
        yield Ok(Event::default().event("done").data(data));
    };

    Ok(Sse::new(stream))
}

// ============================================================================
// Tests (PMAT-802: T-COV-95)
// ============================================================================

#[cfg(test)]
#[path = "gpu_handlers_tests.rs"]
mod gpu_handlers_tests;