diff --git a/catgrad-llm/examples/serve.rs b/catgrad-llm/examples/serve.rs index 7acb1703..e0d8bbd5 100644 --- a/catgrad-llm/examples/serve.rs +++ b/catgrad-llm/examples/serve.rs @@ -10,14 +10,12 @@ //! model-native EOS stopping is supported. use std::convert::Infallible; -use std::sync::atomic::{AtomicU64, Ordering}; -use std::time::{SystemTime, UNIX_EPOCH}; use axum::extract::State; use axum::http::StatusCode; use axum::response::sse::{Event, Sse}; use axum::response::{IntoResponse, Response}; -use axum::routing::post; +use axum::routing::{get, post}; use axum::{Json, Router}; use catgrad::prelude::Dtype; use clap::Parser; @@ -26,12 +24,11 @@ use serde_json::json; use tokio::sync::{mpsc, oneshot}; use tokio_stream::wrappers::UnboundedReceiverStream; -use catgrad_llm::run::{GenerationOutput, ModelEngine}; +use catgrad_llm::api::{self, ApiContext, EndpointResult}; +use catgrad_llm::run::ModelEngine; use catgrad_llm::types::{anthropic, openai}; use catgrad_llm::{LLMError, Result as LlmResult}; -static NEXT_ID: AtomicU64 = AtomicU64::new(1); - #[derive(Parser, Debug)] struct Args { /// Host interface to bind @@ -63,17 +60,7 @@ type Job = Box; #[derive(Clone)] struct AppState { jobs: mpsc::UnboundedSender, -} - -fn now_unix() -> i64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .map(|d| i64::try_from(d.as_secs()).unwrap_or(i64::MAX)) - .unwrap_or(0) -} - -fn next_id(prefix: &str) -> String { - format!("{prefix}-{}", NEXT_ID.fetch_add(1, Ordering::Relaxed)) + api: ApiContext, } fn data_event(payload: &T) -> Event { @@ -86,6 +73,10 @@ fn named_event(name: &str, payload: &T) -> Event { .data(serde_json::to_string(payload).unwrap_or_else(|_| "{}".into())) } +fn done_event() -> Event { + Event::default().data("[DONE]") +} + /// Emit an event when an SSE sink is attached. Lazy: the closure only /// runs in streaming mode. Returns `Err(IoError)` when the client has /// disconnected, so the generation loop can short-circuit. @@ -124,7 +115,7 @@ fn llm_error_response(err: LLMError) -> Response { async fn respond(jobs: &mpsc::UnboundedSender, streaming: bool, f: F) -> Response where T: Serialize + Send + 'static, - F: FnOnce(&Worker, Option<&SseSink>) -> LlmResult + Send + 'static, + F: FnOnce(&Worker, Option<&SseSink>) -> LlmResult> + Send + 'static, { let worker_dead = || { error_response( @@ -132,11 +123,19 @@ where "inference worker terminated", ) }; + if streaming { let (tx, rx) = mpsc::unbounded_channel(); if jobs - .send(Box::new(move |w| { - if let Err(err) = f(w, Some(&tx)) { + .send(Box::new(move |w| match f(w, Some(&tx)) { + Ok(EndpointResult::Streamed) => {} + Ok(EndpointResult::Json(_)) => { + let _ = tx.send(Ok(named_event( + "error", + &json!({ "message": "handler returned JSON for a streaming request" }), + ))); + } + Err(err) => { let _ = tx.send(Ok(named_event( "error", &json!({ "message": err.to_string() }), @@ -147,26 +146,34 @@ where { return worker_dead(); } - Sse::new(UnboundedReceiverStream::new(rx)).into_response() - } else { - let (tx, rx) = oneshot::channel(); - if jobs - .send(Box::new(move |w| { - let _ = tx.send(f(w, None)); - })) - .is_err() - { - return worker_dead(); - } - match rx.await { - Ok(Ok(t)) => Json(t).into_response(), - Ok(Err(err)) => llm_error_response(err), - Err(_) => worker_dead(), - } + return Sse::new(UnboundedReceiverStream::new(rx)).into_response(); + } + + let (tx, rx) = oneshot::channel(); + if jobs + .send(Box::new(move |w| { + let _ = tx.send(f(w, None)); + })) + .is_err() + { + return worker_dead(); + } + + match rx.await { + Ok(Ok(EndpointResult::Json(t))) => Json(t).into_response(), + Ok(Ok(EndpointResult::Streamed)) => error_response( + StatusCode::INTERNAL_SERVER_ERROR, + "handler unexpectedly streamed a non-streaming request", + ), + Ok(Err(err)) => llm_error_response(err), + Err(_) => worker_dead(), } } // --- Handlers --- +async fn handle_models(State(state): State) -> Response { + Json(api::openai_models_response(&state.api)).into_response() +} async fn handle_openai( State(state): State, @@ -178,6 +185,16 @@ async fn handle_openai( .await } +async fn handle_openai_responses( + State(state): State, + Json(req): Json, +) -> Response { + respond(&state.jobs, req.stream == Some(true), move |w, sse| { + w.serve_openai_responses(req, sse) + }) + .await +} + async fn handle_anthropic( State(state): State, Json(req): Json, @@ -189,85 +206,16 @@ async fn handle_anthropic( } // --- Worker --- - -fn openai_chunk( - id: &str, - created: i64, - model: &str, - delta: openai::ChatDelta, - finish_reason: Option, -) -> openai::ChatCompletionChunk { - openai::ChatCompletionChunk::builder() - .id(id.into()) - .object("chat.completion.chunk".into()) - .created(created) - .model(model.into()) - .choices(vec![ - openai::ChatStreamChoice::builder() - .index(0) - .delta(delta) - .finish_reason(finish_reason) - .build(), - ]) - .build() -} - -fn openai_response( - id: String, - created: i64, - model: String, - g: GenerationOutput, -) -> openai::ChatCompletionResponse { - openai::ChatCompletionResponse::builder() - .id(id) - .object("chat.completion".into()) - .created(created) - .model(model) - .choices(vec![ - openai::ChatChoice::builder() - .index(0) - .message(openai::ChatMessage::assistant(g.text)) - .finish_reason(Some(g.termination.into())) - .build(), - ]) - .usage(Some(openai::Usage::from_counts( - g.prompt_tokens, - g.completion_tokens, - ))) - .build() -} - -fn anthropic_response( - id: String, - model: String, - g: GenerationOutput, -) -> anthropic::MessageResponse { - anthropic::MessageResponse::builder() - .id(id) - .message_type(Some("message".into())) - .role("assistant".into()) - .content(vec![anthropic::ContentBlock::Text { text: g.text }]) - .model(model) - .stop_reason(Some(g.termination.into())) - .usage(anthropic::AnthropicUsage::new( - g.prompt_tokens, - g.completion_tokens, - )) - .build() -} - struct Worker { engine: ModelEngine, - model_name: String, - default_max_tokens: u32, + api: ApiContext, } impl Worker { - fn new(model: &str, use_kv_cache: bool, default_max_tokens: u32) -> anyhow::Result { + fn new(model: &str, api: ApiContext, use_kv_cache: bool) -> anyhow::Result { Ok(Self { engine: ModelEngine::new(model, use_kv_cache, Dtype::F32)?, - model_name: model.to_string(), - default_max_tokens, + api, }) } @@ -281,150 +229,48 @@ impl Worker { &self, req: openai::ChatCompletionRequest, sse: Option<&SseSink>, - ) -> LlmResult { - let max_tokens = req.max_tokens.unwrap_or(self.default_max_tokens); - let include_usage = req - .stream_options - .as_ref() - .and_then(|o| o.include_usage) - .unwrap_or(false); - let prepared = self.engine.prepare_openai(&req)?; - - let id = next_id("chatcmpl"); - let created = now_unix(); - let model = self.model_name.clone(); - - emit_maybe(sse, || { - data_event(&openai_chunk( - &id, - created, - &model, - openai::ChatDelta { - role: Some("assistant".into()), - ..Default::default() - }, - None, - )) + ) -> LlmResult> { + let result = api::handle_openai_chat(&self.engine, &self.api, &req, |chunk| { + emit_maybe(sse, move || data_event(&chunk)) })?; + if matches!(result, EndpointResult::Streamed) { + emit_maybe(sse, done_event)?; + } + Ok(result) + } - let g = self - .engine - .generate_from_prepared(&prepared, max_tokens, |delta| { - emit_maybe(sse, || { - data_event(&openai_chunk( - &id, - created, - &model, - openai::ChatDelta { - content: Some(delta.into()), - ..Default::default() - }, - None, - )) - }) - })?; - - emit_maybe(sse, || { - data_event(&openai_chunk( - &id, - created, - &model, - openai::ChatDelta::default(), - Some(g.termination.into()), - )) + fn serve_openai_responses( + &self, + req: openai::responses::ResponseRequest, + sse: Option<&SseSink>, + ) -> LlmResult> { + let result = api::handle_openai_responses(&self.engine, &self.api, &req, |event| { + emit_maybe(sse, move || data_event(&event)) })?; - - if include_usage { - emit_maybe(sse, || { - let chunk = openai::ChatCompletionChunk::builder() - .id(id.clone()) - .object("chat.completion.chunk".into()) - .created(created) - .model(model.clone()) - .choices(vec![]) - .usage(Some(openai::Usage::from_counts( - g.prompt_tokens, - g.completion_tokens, - ))) - .build(); - data_event(&chunk) - })?; + if matches!(result, EndpointResult::Streamed) { + emit_maybe(sse, done_event)?; } - emit_maybe(sse, || Event::default().data("[DONE]"))?; - - Ok(openai_response(id, created, model, g)) + Ok(result) } fn serve_anthropic( &self, req: anthropic::MessageRequest, sse: Option<&SseSink>, - ) -> LlmResult { - use anthropic::MessageStreamEvent::*; - let prepared = self.engine.prepare_anthropic(&req)?; - let id = next_id("msg"); - let model = self.model_name.clone(); - let emit = - |name, ev: anthropic::MessageStreamEvent| emit_maybe(sse, || named_event(name, &ev)); - - emit( - "message_start", - MessageStart { - message: anthropic::MessageResponse::builder() - .id(id.clone()) - .message_type(Some("message".into())) - .role("assistant".into()) - .content(vec![]) - .model(model.clone()) - .usage(anthropic::AnthropicUsage::new( - prepared.input_ids.len() as u32, - 0, - )) - .build(), - }, - )?; - emit( - "content_block_start", - ContentBlockStart { - index: 0, - content_block: anthropic::ContentBlock::Text { - text: String::new(), - }, - }, - )?; - - let g = self - .engine - .generate_from_prepared(&prepared, req.max_tokens, |delta| { - emit( - "content_block_delta", - ContentBlockDelta { - index: 0, - delta: anthropic::ContentBlockDelta::TextDelta { text: delta.into() }, - }, - ) - })?; - - emit("content_block_stop", ContentBlockStop { index: 0 })?; - emit( - "message_delta", - MessageDelta { - delta: anthropic::StreamMessageDelta { - stop_reason: Some(g.termination.into()), - }, - usage: anthropic::AnthropicUsage::new(g.prompt_tokens, g.completion_tokens), - }, - )?; - emit("message_stop", MessageStop)?; - - Ok(anthropic_response(id, model, g)) + ) -> LlmResult> { + api::handle_anthropic_messages(&self.engine, &self.api, &req, |event| { + emit_maybe(sse, move || named_event(event.event, &event.payload)) + }) } } -// --- main --- - fn main() -> anyhow::Result<()> { let args = Args::parse(); + let api = ApiContext::new( + &args.model, + vec![args.model.clone()], + args.default_max_tokens, + ); let (jobs_tx, jobs_rx) = mpsc::unbounded_channel::(); let (ready_tx, ready_rx) = std::sync::mpsc::channel::>(); @@ -433,13 +279,13 @@ fn main() -> anyhow::Result<()> { // thread, signal readiness, then run the dispatch loop. let model = args.model.clone(); let use_kv_cache = args.use_kv_cache; - let default_max_tokens = args.default_max_tokens; + let api_for_worker = api.clone(); let worker_handle = std::thread::Builder::new() .name("inference".into()) .spawn(move || { println!("Loading model `{model}` (this can take a while)..."); - let worker = match Worker::new(&model, use_kv_cache, default_max_tokens) { - Ok(w) => w, + let worker = match Worker::new(&model, api_for_worker, use_kv_cache) { + Ok(worker) => worker, Err(err) => { let _ = ready_tx.send(Err(err)); return; @@ -453,20 +299,24 @@ fn main() -> anyhow::Result<()> { let rt = tokio::runtime::Builder::new_current_thread() .enable_io() .build()?; - rt.block_on(serve(args, jobs_tx))?; + rt.block_on(serve(args, AppState { jobs: jobs_tx, api }))?; let _ = worker_handle.join(); Ok(()) } -async fn serve(args: Args, jobs: mpsc::UnboundedSender) -> anyhow::Result<()> { +async fn serve(args: Args, state: AppState) -> anyhow::Result<()> { let app = Router::new() + .route("/v1/models", get(handle_models)) .route("/v1/chat/completions", post(handle_openai)) + .route("/v1/responses", post(handle_openai_responses)) .route("/v1/messages", post(handle_anthropic)) - .with_state(AppState { jobs }); + .with_state(state); let addr = format!("{}:{}", args.host, args.port); let listener = tokio::net::TcpListener::bind(&addr).await?; println!("catgrad demo API server listening on http://{addr}"); + println!("GET /v1/models (OpenAI)"); println!("POST /v1/chat/completions (OpenAI)"); + println!("POST /v1/responses (OpenAI)"); println!("POST /v1/messages (Anthropic)"); axum::serve(listener, app).await?; Ok(()) diff --git a/catgrad-llm/src/api.rs b/catgrad-llm/src/api.rs new file mode 100644 index 00000000..fa766639 --- /dev/null +++ b/catgrad-llm/src/api.rs @@ -0,0 +1,486 @@ +//! Transport-agnostic API serving helpers built on top of [`crate::run::ModelEngine`]. +use crate::Result; +use crate::run::ModelEngine; +use crate::types::{anthropic, openai}; +use serde_json::{Value, json}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{SystemTime, UNIX_EPOCH}; + +static NEXT_ID: AtomicU64 = AtomicU64::new(1); + +/// Shared serving configuration that is independent from any HTTP transport. +#[derive(Debug, Clone)] +pub struct ApiContext { + model_name: String, + served_model_ids: Vec, + default_max_tokens: u32, +} + +impl ApiContext { + pub fn new( + model_name: impl Into, + served_model_ids: Vec, + default_max_tokens: u32, + ) -> Self { + Self { + model_name: model_name.into(), + served_model_ids, + default_max_tokens, + } + } + + pub fn model_name(&self) -> &str { + &self.model_name + } + + pub fn served_model_ids(&self) -> &[String] { + &self.served_model_ids + } + + pub fn default_max_tokens(&self) -> u32 { + self.default_max_tokens + } +} + +/// Result of handling an API request before transport framing is applied. +#[derive(Debug, Clone, PartialEq)] +pub enum EndpointResult { + Json(T), + Streamed, +} + +/// Named SSE event payload for APIs that use explicit event names. +#[derive(Debug, Clone, PartialEq)] +pub struct NamedEvent { + pub event: &'static str, + pub payload: T, +} + +pub fn openai_models_response(context: &ApiContext) -> Value { + json!({ + "object": "list", + "data": context.served_model_ids() + .iter() + .map(|model_id| { + json!({ + "id": model_id, + "object": "model" + }) + }) + .collect::>() + }) +} + +pub fn handle_openai_chat( + engine: &ModelEngine, + context: &ApiContext, + request: &openai::ChatCompletionRequest, + mut on_chunk: F, +) -> Result> +where + F: FnMut(openai::ChatCompletionChunk) -> Result<()>, +{ + let prepared = engine.prepare_openai_chat_request(request)?; + let max_tokens = request.max_tokens.unwrap_or(context.default_max_tokens()); + let model = context.model_name().to_string(); + + if request.stream == Some(true) { + let id = next_id("chatcmpl"); + let created = now_unix(); + + on_chunk( + openai::ChatCompletionChunk::builder() + .id(id.clone()) + .object("chat.completion.chunk".to_string()) + .created(created) + .model(model.clone()) + .choices(vec![ + openai::ChatStreamChoice::builder() + .index(0) + .delta(openai::ChatDelta { + role: Some("assistant".to_string()), + ..Default::default() + }) + .build(), + ]) + .build(), + )?; + + let generated = engine.generate_from_prepared(&prepared, max_tokens, |delta| { + on_chunk( + openai::ChatCompletionChunk::builder() + .id(id.clone()) + .object("chat.completion.chunk".to_string()) + .created(created) + .model(model.clone()) + .choices(vec![ + openai::ChatStreamChoice::builder() + .index(0) + .delta(openai::ChatDelta { + content: Some(delta.to_string()), + ..Default::default() + }) + .build(), + ]) + .build(), + ) + })?; + + on_chunk( + openai::ChatCompletionChunk::builder() + .id(id.clone()) + .object("chat.completion.chunk".to_string()) + .created(created) + .model(model.clone()) + .choices(vec![ + openai::ChatStreamChoice::builder() + .index(0) + .delta(openai::ChatDelta::default()) + .finish_reason(Some(generated.termination.into())) + .build(), + ]) + .build(), + )?; + + let include_usage = request + .stream_options + .as_ref() + .and_then(|options| options.include_usage) + .unwrap_or(false); + if include_usage { + on_chunk( + openai::ChatCompletionChunk::builder() + .id(id) + .object("chat.completion.chunk".to_string()) + .created(created) + .model(model) + .choices(vec![]) + .usage(Some(openai::Usage::from_counts( + generated.prompt_tokens, + generated.completion_tokens, + ))) + .build(), + )?; + } + + return Ok(EndpointResult::Streamed); + } + + let generated = engine.generate_from_prepared(&prepared, max_tokens, |_| Ok(()))?; + Ok(EndpointResult::Json( + openai::ChatCompletionResponse::builder() + .id(next_id("chatcmpl")) + .object("chat.completion".to_string()) + .created(now_unix()) + .model(model) + .choices(vec![ + openai::ChatChoice::builder() + .index(0) + .message(openai::ChatMessage::assistant(generated.text)) + .finish_reason(Some(generated.termination.into())) + .build(), + ]) + .usage(Some(openai::Usage::from_counts( + generated.prompt_tokens, + generated.completion_tokens, + ))) + .build(), + )) +} + +pub fn handle_openai_responses( + engine: &ModelEngine, + context: &ApiContext, + request: &openai::responses::ResponseRequest, + mut on_event: F, +) -> Result> +where + F: FnMut(openai::responses::ResponseStreamEvent) -> Result<()>, +{ + let prepared = engine.prepare_openai_response_request(request)?; + let max_tokens = request + .max_output_tokens + .unwrap_or(context.default_max_tokens()); + let model = context.model_name().to_string(); + + if request.stream == Some(true) { + let response_id = next_id("resp"); + let message_id = next_id("msg"); + let created = now_unix(); + let mut sequence_number = 0u64; + + on_event(openai::responses::ResponseStreamEvent::Created { + sequence_number, + response: build_openai_response( + &response_id, + created, + &model, + openai::responses::ResponseStatus::Queued, + vec![], + None, + ), + })?; + sequence_number += 1; + + on_event(openai::responses::ResponseStreamEvent::InProgress { + sequence_number, + response: build_openai_response( + &response_id, + created, + &model, + openai::responses::ResponseStatus::InProgress, + vec![], + None, + ), + })?; + sequence_number += 1; + + on_event(openai::responses::ResponseStreamEvent::OutputItemAdded { + sequence_number, + output_index: 0, + item: build_openai_response_message( + &message_id, + openai::responses::ResponseStatus::InProgress, + vec![], + ), + })?; + sequence_number += 1; + + on_event(openai::responses::ResponseStreamEvent::ContentPartAdded { + sequence_number, + item_id: message_id.clone(), + output_index: 0, + content_index: 0, + part: build_openai_response_text(String::new()), + })?; + sequence_number += 1; + + let mut text = String::new(); + let generated = engine.generate_from_prepared(&prepared, max_tokens, |delta| { + text.push_str(delta); + on_event(openai::responses::ResponseStreamEvent::OutputTextDelta { + sequence_number, + item_id: message_id.clone(), + output_index: 0, + content_index: 0, + delta: delta.to_string(), + })?; + sequence_number += 1; + Ok(()) + })?; + + on_event(openai::responses::ResponseStreamEvent::OutputTextDone { + sequence_number, + item_id: message_id.clone(), + output_index: 0, + content_index: 0, + text: text.clone(), + })?; + sequence_number += 1; + + let completed_part = build_openai_response_text(text); + on_event(openai::responses::ResponseStreamEvent::ContentPartDone { + sequence_number, + item_id: message_id.clone(), + output_index: 0, + content_index: 0, + part: completed_part.clone(), + })?; + sequence_number += 1; + + let completed_item = build_openai_response_message( + &message_id, + openai::responses::ResponseStatus::Completed, + vec![completed_part], + ); + on_event(openai::responses::ResponseStreamEvent::OutputItemDone { + sequence_number, + output_index: 0, + item: completed_item.clone(), + })?; + sequence_number += 1; + + on_event(openai::responses::ResponseStreamEvent::Completed { + sequence_number, + response: build_openai_response( + &response_id, + created, + &model, + openai::responses::ResponseStatus::Completed, + vec![completed_item], + Some(openai::responses::ResponseUsage::from_counts( + generated.prompt_tokens, + generated.completion_tokens, + )), + ), + })?; + + return Ok(EndpointResult::Streamed); + } + + let generated = engine.generate_from_prepared(&prepared, max_tokens, |_| Ok(()))?; + Ok(EndpointResult::Json(build_openai_response( + &next_id("resp"), + now_unix(), + &model, + openai::responses::ResponseStatus::Completed, + vec![build_openai_response_message( + &next_id("msg"), + openai::responses::ResponseStatus::Completed, + vec![build_openai_response_text(generated.text)], + )], + Some(openai::responses::ResponseUsage::from_counts( + generated.prompt_tokens, + generated.completion_tokens, + )), + ))) +} + +pub fn handle_anthropic_messages( + engine: &ModelEngine, + context: &ApiContext, + request: &anthropic::MessageRequest, + mut on_event: F, +) -> Result> +where + F: FnMut(NamedEvent) -> Result<()>, +{ + let prepared = engine.prepare_anthropic(request)?; + let model = context.model_name().to_string(); + + if request.stream == Some(true) { + let id = next_id("msg"); + on_event(NamedEvent { + event: "message_start", + payload: anthropic::MessageStreamEvent::MessageStart { + message: anthropic::MessageResponse::builder() + .id(id) + .message_type(Some("message".to_string())) + .role("assistant".to_string()) + .content(vec![]) + .model(model) + .usage(anthropic::AnthropicUsage::new( + prepared.input_ids.len() as u32, + 0, + )) + .build(), + }, + })?; + on_event(NamedEvent { + event: "content_block_start", + payload: anthropic::MessageStreamEvent::ContentBlockStart { + index: 0, + content_block: anthropic::ContentBlock::Text { + text: String::new(), + }, + }, + })?; + + let generated = engine.generate_from_prepared(&prepared, request.max_tokens, |delta| { + on_event(NamedEvent { + event: "content_block_delta", + payload: anthropic::MessageStreamEvent::ContentBlockDelta { + index: 0, + delta: anthropic::ContentBlockDelta::TextDelta { + text: delta.to_string(), + }, + }, + }) + })?; + + on_event(NamedEvent { + event: "content_block_stop", + payload: anthropic::MessageStreamEvent::ContentBlockStop { index: 0 }, + })?; + on_event(NamedEvent { + event: "message_delta", + payload: anthropic::MessageStreamEvent::MessageDelta { + delta: anthropic::StreamMessageDelta { + stop_reason: Some(generated.termination.into()), + }, + usage: anthropic::AnthropicUsage::new( + generated.prompt_tokens, + generated.completion_tokens, + ), + }, + })?; + on_event(NamedEvent { + event: "message_stop", + payload: anthropic::MessageStreamEvent::MessageStop, + })?; + + return Ok(EndpointResult::Streamed); + } + + let generated = engine.generate_from_prepared(&prepared, request.max_tokens, |_| Ok(()))?; + Ok(EndpointResult::Json( + anthropic::MessageResponse::builder() + .id(next_id("msg")) + .message_type(Some("message".to_string())) + .role("assistant".to_string()) + .content(vec![anthropic::ContentBlock::Text { + text: generated.text, + }]) + .model(model) + .stop_reason(Some(generated.termination.into())) + .usage(anthropic::AnthropicUsage::new( + generated.prompt_tokens, + generated.completion_tokens, + )) + .build(), + )) +} + +fn build_openai_response( + response_id: &str, + created_at: i64, + model: &str, + status: openai::responses::ResponseStatus, + output: Vec, + usage: Option, +) -> openai::responses::Response { + openai::responses::Response::builder() + .id(response_id.to_string()) + .object("response".to_string()) + .created_at(created_at) + .status(status) + .model(model.to_string()) + .output(output) + .usage(usage) + .build() +} + +fn build_openai_response_message( + message_id: &str, + status: openai::responses::ResponseStatus, + content: Vec, +) -> openai::responses::ResponseOutputMessage { + openai::responses::ResponseOutputMessage::builder() + .id(message_id.to_string()) + .item_type("message".to_string()) + .status(status) + .role("assistant".to_string()) + .content(content) + .annotations(Some(vec![])) + .build() +} + +fn build_openai_response_text(text: String) -> openai::responses::ResponseOutputText { + openai::responses::ResponseOutputText::builder() + .content_type("output_text".to_string()) + .text(text) + .annotations(Some(vec![])) + .build() +} + +fn now_unix() -> i64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| i64::try_from(duration.as_secs()).unwrap_or(i64::MAX)) + .unwrap_or(0) +} + +fn next_id(prefix: &str) -> String { + let value = NEXT_ID.fetch_add(1, Ordering::Relaxed); + format!("{prefix}-{value}") +} diff --git a/catgrad-llm/src/lib.rs b/catgrad-llm/src/lib.rs index a5678b33..900ad2a1 100644 --- a/catgrad-llm/src/lib.rs +++ b/catgrad-llm/src/lib.rs @@ -1,5 +1,6 @@ //! LLM-specific runtime code like tokenization and kv-cache logic which has to live outside //! the model graph. +pub mod api; mod error; pub mod run; pub mod types; diff --git a/catgrad-llm/src/run.rs b/catgrad-llm/src/run.rs index e2921dd1..6a4c43ac 100644 --- a/catgrad-llm/src/run.rs +++ b/catgrad-llm/src/run.rs @@ -240,6 +240,31 @@ impl ModelEngine { ) } + /// Renders an OpenAI responses request through the model chat template. + pub fn prepare_openai_response_request( + &self, + request: &types::openai::responses::ResponseRequest, + ) -> Result { + let messages = request.to_messages()?; + let tools = request + .tools + .as_deref() + .filter(|tools| !tools.is_empty()) + .map(|tools| { + tools + .iter() + .map(minijinja::Value::from_serialize) + .collect::>() + }); + self.prepare_chat_messages( + &messages, + RenderChatTemplateOptions { + thinking: types::ThinkingPolicy::Disabled, + tools: tools.as_deref(), + }, + ) + } + fn prepare_chat_messages( &self, messages: &[types::Message], diff --git a/catgrad-llm/src/types/openai.rs b/catgrad-llm/src/types/openai.rs index 776244f8..cb69f959 100644 --- a/catgrad-llm/src/types/openai.rs +++ b/catgrad-llm/src/types/openai.rs @@ -1,4 +1,4 @@ -//! OpenAI chat-completions wire format. +//! OpenAI wire-format types. use crate::LLMError; use serde::{Deserialize, Serialize}; use serde_json::Value as JsonValue; @@ -356,3 +356,671 @@ mod tests { assert_eq!(serialized, sample); } } + +pub mod responses { + use crate::LLMError; + use serde::{Deserialize, Serialize}; + use serde_json::Value as JsonValue; + use serde_with::skip_serializing_none; + use typed_builder::TypedBuilder; + + /// Responses API request. + #[skip_serializing_none] + #[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder, PartialEq)] + pub struct ResponseRequest { + pub model: String, + pub input: ResponseInput, + #[builder(default)] + pub instructions: Option, + #[builder(default)] + pub tools: Option>, + #[builder(default)] + pub max_output_tokens: Option, + #[builder(default)] + pub stream: Option, + } + + /// Top-level response input payload. + #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] + #[serde(untagged)] + pub enum ResponseInput { + Text(String), + Items(Vec), + } + + /// A single item in a structured Responses API input array. + #[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder, PartialEq)] + pub struct ResponseInputItem { + pub role: String, + pub content: ResponseInputMessageContent, + } + + /// Content payload for a structured input message. + #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] + #[serde(untagged)] + pub enum ResponseInputMessageContent { + Text(String), + Parts(Vec), + } + + /// A supported content part in structured response input. + #[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder, PartialEq, Eq)] + pub struct ResponseInputContentPart { + #[serde(rename = "type")] + pub content_type: String, + pub text: String, + } + + impl ResponseRequest { + pub fn to_messages(&self) -> Result, LLMError> { + let mut messages = Vec::with_capacity( + match &self.input { + ResponseInput::Text(_) => 1, + ResponseInput::Items(items) => items.len(), + } + usize::from(self.instructions.is_some()), + ); + + if let Some(instructions) = &self.instructions { + messages.push(super::super::Message::openai(super::ChatMessage::system( + instructions.clone(), + ))); + } + + match &self.input { + ResponseInput::Text(text) => { + messages.push(super::super::Message::openai(super::ChatMessage::user( + text.clone(), + ))); + } + ResponseInput::Items(items) => { + for item in items { + let message: super::ChatMessage = item.clone().try_into()?; + messages.push(super::super::Message::openai(message)); + } + } + } + + Ok(messages) + } + } + + impl TryFrom for super::ChatMessage { + type Error = LLMError; + + fn try_from(value: ResponseInputItem) -> Result { + Ok(super::ChatMessage::builder() + .role(value.role) + .content(Some(value.content.try_into()?)) + .build()) + } + } + + impl TryFrom for super::MessageContent { + type Error = LLMError; + + fn try_from(value: ResponseInputMessageContent) -> Result { + match value { + ResponseInputMessageContent::Text(text) => Ok(Self::Text(text)), + ResponseInputMessageContent::Parts(parts) => Ok(Self::Parts( + parts + .into_iter() + .map(TryInto::try_into) + .collect::, _>>()?, + )), + } + } + } + + impl TryFrom for super::ContentPart { + type Error = LLMError; + + fn try_from(value: ResponseInputContentPart) -> Result { + match value.content_type.as_str() { + "input_text" | "output_text" => Ok(Self::Text { text: value.text }), + other => Err(LLMError::UnsupportedWireConversion(format!( + "Unsupported responses input content type `{other}`" + ))), + } + } + } + + /// Overall response status. + #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] + #[serde(rename_all = "snake_case")] + pub enum ResponseStatus { + Queued, + InProgress, + Completed, + } + + /// Responses API token usage details. + #[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder, PartialEq, Eq)] + pub struct ResponseUsage { + pub input_tokens: u32, + pub output_tokens: u32, + pub total_tokens: u32, + } + + impl ResponseUsage { + pub fn from_counts(input_tokens: u32, output_tokens: u32) -> Self { + Self::builder() + .input_tokens(input_tokens) + .output_tokens(output_tokens) + .total_tokens(input_tokens + output_tokens) + .build() + } + } + + /// Text content item in an assistant response message. + #[skip_serializing_none] + #[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder, PartialEq, Eq)] + pub struct ResponseOutputText { + #[serde(rename = "type")] + pub content_type: String, + pub text: String, + #[builder(default)] + pub annotations: Option>, + } + + /// Assistant message item in the response output array. + #[skip_serializing_none] + #[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder, PartialEq, Eq)] + pub struct ResponseOutputMessage { + pub id: String, + #[serde(rename = "type")] + pub item_type: String, + pub status: ResponseStatus, + pub role: String, + pub content: Vec, + #[builder(default)] + pub annotations: Option>, + } + + /// Responses API response. + #[skip_serializing_none] + #[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder, PartialEq, Eq)] + pub struct Response { + pub id: String, + pub object: String, + pub created_at: i64, + pub status: ResponseStatus, + pub model: String, + pub output: Vec, + #[builder(default)] + pub usage: Option, + } + + /// Streaming event payloads for the Responses API. + #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] + #[serde(tag = "type")] + pub enum ResponseStreamEvent { + #[serde(rename = "response.created")] + Created { + sequence_number: u64, + response: Response, + }, + #[serde(rename = "response.in_progress")] + InProgress { + sequence_number: u64, + response: Response, + }, + #[serde(rename = "response.output_item.added")] + OutputItemAdded { + sequence_number: u64, + output_index: u32, + item: ResponseOutputMessage, + }, + #[serde(rename = "response.content_part.added")] + ContentPartAdded { + sequence_number: u64, + item_id: String, + output_index: u32, + content_index: u32, + part: ResponseOutputText, + }, + #[serde(rename = "response.output_text.delta")] + OutputTextDelta { + sequence_number: u64, + item_id: String, + output_index: u32, + content_index: u32, + delta: String, + }, + #[serde(rename = "response.output_text.done")] + OutputTextDone { + sequence_number: u64, + item_id: String, + output_index: u32, + content_index: u32, + text: String, + }, + #[serde(rename = "response.content_part.done")] + ContentPartDone { + sequence_number: u64, + item_id: String, + output_index: u32, + content_index: u32, + part: ResponseOutputText, + }, + #[serde(rename = "response.output_item.done")] + OutputItemDone { + sequence_number: u64, + output_index: u32, + item: ResponseOutputMessage, + }, + #[serde(rename = "response.completed")] + Completed { + sequence_number: u64, + response: Response, + }, + } + + #[cfg(test)] + mod tests { + use super::*; + use crate::types::Message; + use crate::types::openai::{ChatMessage, ContentPart, MessageContent}; + use serde_json::json; + + #[test] + fn response_request_parses_minimal_string_input() { + let parsed: ResponseRequest = serde_json::from_value(json!({ + "model": "gpt-4.1-mini", + "input": "Hello" + })) + .unwrap(); + + assert_eq!( + parsed, + ResponseRequest::builder() + .model("gpt-4.1-mini".to_string()) + .input(ResponseInput::Text("Hello".to_string())) + .build() + ); + } + + #[test] + fn response_request_supports_message_list_instructions_and_tools() { + let parsed: ResponseRequest = serde_json::from_value(json!({ + "model": "gpt-4.1-mini", + "input": [ + { + "type": "message", + "role": "developer", + "content": [ + {"type": "input_text", "text": "You are helpful."}, + {"type": "input_text", "text": "Answer briefly."} + ] + }, + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": "Tell a story about a unicorn."} + ] + } + ], + "instructions": "You are a nice LLM", + "tools": [ + { + "type": "function", + "name": "lookup_weather", + "parameters": {"type": "object"} + } + ], + "max_output_tokens": 64, + "stream": false, + "temperature": 0.2 + })) + .unwrap(); + + assert_eq!( + parsed, + ResponseRequest::builder() + .model("gpt-4.1-mini".to_string()) + .input(ResponseInput::Items(vec![ + ResponseInputItem::builder() + .role("developer".to_string()) + .content(ResponseInputMessageContent::Parts(vec![ + ResponseInputContentPart::builder() + .content_type("input_text".to_string()) + .text("You are helpful.".to_string()) + .build(), + ResponseInputContentPart::builder() + .content_type("input_text".to_string()) + .text("Answer briefly.".to_string()) + .build(), + ])) + .build(), + ResponseInputItem::builder() + .role("user".to_string()) + .content(ResponseInputMessageContent::Parts(vec![ + ResponseInputContentPart::builder() + .content_type("input_text".to_string()) + .text("Tell a story about a unicorn.".to_string()) + .build(), + ])) + .build(), + ])) + .instructions(Some("You are a nice LLM".to_string())) + .tools(Some(vec![json!({ + "type": "function", + "name": "lookup_weather", + "parameters": {"type": "object"} + })])) + .max_output_tokens(Some(64)) + .stream(Some(false)) + .build() + ); + } + + #[test] + fn response_round_trips_through_serde() { + let sample = json!({ + "id": "resp-1", + "object": "response", + "created_at": 1234567890, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "id": "msg-1", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "Hello world" + } + ] + } + ], + "usage": { + "input_tokens": 5, + "output_tokens": 2, + "total_tokens": 7 + } + }); + + let parsed: Response = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_created_event_round_trips_through_serde() { + let sample = json!({ + "type": "response.created", + "sequence_number": 0, + "response": { + "id": "resp-1", + "object": "response", + "created_at": 1234567890, + "status": "queued", + "model": "gpt-4.1-mini", + "output": [] + } + }); + + let parsed: ResponseStreamEvent = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_in_progress_event_round_trips_through_serde() { + let sample = json!({ + "type": "response.in_progress", + "sequence_number": 1, + "response": { + "id": "resp-1", + "object": "response", + "created_at": 1234567890, + "status": "in_progress", + "model": "gpt-4.1-mini", + "output": [] + } + }); + + let parsed: ResponseStreamEvent = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_output_item_added_event_round_trips_through_serde() { + let sample = json!({ + "type": "response.output_item.added", + "sequence_number": 2, + "output_index": 0, + "item": { + "id": "msg-1", + "type": "message", + "status": "in_progress", + "role": "assistant", + "content": [], + "annotations": [] + } + }); + + let parsed: ResponseStreamEvent = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_content_part_added_event_round_trips_through_serde() { + let sample = json!({ + "type": "response.content_part.added", + "sequence_number": 3, + "item_id": "msg-1", + "output_index": 0, + "content_index": 0, + "part": { + "type": "output_text", + "text": "", + "annotations": [] + } + }); + + let parsed: ResponseStreamEvent = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_output_text_delta_event_round_trips_through_serde() { + let sample = json!({ + "type": "response.output_text.delta", + "sequence_number": 4, + "item_id": "msg-1", + "output_index": 0, + "content_index": 0, + "delta": "Hello" + }); + + let parsed: ResponseStreamEvent = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_output_text_done_event_round_trips_through_serde() { + let sample = json!({ + "type": "response.output_text.done", + "sequence_number": 5, + "item_id": "msg-1", + "output_index": 0, + "content_index": 0, + "text": "Hello world" + }); + + let parsed: ResponseStreamEvent = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_content_part_done_event_round_trips_through_serde() { + let sample = json!({ + "type": "response.content_part.done", + "sequence_number": 6, + "item_id": "msg-1", + "output_index": 0, + "content_index": 0, + "part": { + "type": "output_text", + "text": "Hello world", + "annotations": [] + } + }); + + let parsed: ResponseStreamEvent = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_output_item_done_event_round_trips_through_serde() { + let sample = json!({ + "type": "response.output_item.done", + "sequence_number": 7, + "output_index": 0, + "item": { + "id": "msg-1", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "Hello world", + "annotations": [] + } + ], + "annotations": [] + } + }); + + let parsed: ResponseStreamEvent = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_completed_event_round_trips_through_serde() { + let sample = json!({ + "type": "response.completed", + "sequence_number": 8, + "response": { + "id": "resp-1", + "object": "response", + "created_at": 1234567890, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "id": "msg-1", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "Hello world", + "annotations": [] + } + ], + "annotations": [] + } + ], + "usage": { + "input_tokens": 5, + "output_tokens": 2, + "total_tokens": 7 + } + } + }); + + let parsed: ResponseStreamEvent = serde_json::from_value(sample.clone()).unwrap(); + let serialized = serde_json::to_value(parsed).unwrap(); + assert_eq!(serialized, sample); + } + + #[test] + fn response_request_accepts_output_text_input_parts() { + let request = ResponseRequest::builder() + .model("gpt-4.1-mini".to_string()) + .input(ResponseInput::Items(vec![ + ResponseInputItem::builder() + .role("assistant".to_string()) + .content(ResponseInputMessageContent::Parts(vec![ + ResponseInputContentPart::builder() + .content_type("output_text".to_string()) + .text("Previous answer".to_string()) + .build(), + ])) + .build(), + ])) + .build(); + + let messages = request.to_messages().unwrap(); + assert_eq!( + messages, + vec![Message::openai( + ChatMessage::builder() + .role("assistant".to_string()) + .content(Some(MessageContent::Parts(vec![ContentPart::Text { + text: "Previous answer".to_string(), + }]))) + .build(), + )] + ); + } + + #[test] + fn response_request_converts_to_chat_messages() { + let request = ResponseRequest::builder() + .model("gpt-4.1-mini".to_string()) + .input(ResponseInput::Items(vec![ + ResponseInputItem::builder() + .role("user".to_string()) + .content(ResponseInputMessageContent::Parts(vec![ + ResponseInputContentPart::builder() + .content_type("input_text".to_string()) + .text("Hello".to_string()) + .build(), + ResponseInputContentPart::builder() + .content_type("input_text".to_string()) + .text(" world".to_string()) + .build(), + ])) + .build(), + ])) + .instructions(Some("Be concise.".to_string())) + .build(); + + let messages = request.to_messages().unwrap(); + assert_eq!( + messages, + vec![ + Message::openai(ChatMessage::system("Be concise.")), + Message::openai( + ChatMessage::builder() + .role("user".to_string()) + .content(Some(MessageContent::Parts(vec![ + ContentPart::Text { + text: "Hello".to_string() + }, + ContentPart::Text { + text: " world".to_string() + }, + ]))) + .build(), + ), + ] + ); + } + } +}