mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-06 02:48:25 +00:00
Merge pull request #627 from fabro-sh/fix/openai-compatible-stream-usage
fix(llm): request streaming usage on openai_compatible providers
This commit is contained in:
commit
7044dd57fa
8 changed files with 100 additions and 7 deletions
|
|
@ -1,7 +1,7 @@
|
|||
//! Request encoding: canonical `Request` → Chat Completions body.
|
||||
|
||||
use super::translate;
|
||||
use super::wire::{ApiRequest, ChatMessage};
|
||||
use super::wire::{ApiRequest, ChatMessage, StreamOptions};
|
||||
use crate::codec::{CodecCtx, EncodedRequest, cache, merge_named_provider_options};
|
||||
use crate::error::Error;
|
||||
|
||||
|
|
@ -10,8 +10,10 @@ use crate::error::Error;
|
|||
const KNOWN_OPTION_KEYS: &[&str] = &["auto_cache"];
|
||||
|
||||
/// Build the Chat Completions request for `ctx.request`. `stream` toggles the
|
||||
/// `stream` body field. The body is assembled as a `serde_json::Value` so
|
||||
/// `provider_options.<provider_name>` fields can be merged in before sending.
|
||||
/// `stream` body field and the `stream_options.include_usage` opt-in that makes
|
||||
/// providers emit the trailing usage chunk. The body is assembled as a
|
||||
/// `serde_json::Value` so `provider_options.<provider_name>` fields can be
|
||||
/// merged in before sending.
|
||||
///
|
||||
/// Returns an error when the request contains a custom tool definition, which
|
||||
/// the Chat Completions tool envelope cannot represent.
|
||||
|
|
@ -55,6 +57,9 @@ pub(super) fn encode(ctx: &CodecCtx<'_>, stream: bool) -> Result<EncodedRequest,
|
|||
tool_choice,
|
||||
response_format,
|
||||
stream: stream.then_some(true),
|
||||
stream_options: stream.then_some(StreamOptions {
|
||||
include_usage: true,
|
||||
}),
|
||||
};
|
||||
|
||||
let mut body = serde_json::to_value(&api_request).unwrap_or_default();
|
||||
|
|
@ -172,6 +177,7 @@ mod tests {
|
|||
tool_choice: None,
|
||||
response_format: None,
|
||||
stream: Some(true),
|
||||
stream_options: None,
|
||||
};
|
||||
let json = serde_json::to_value(&req).unwrap();
|
||||
assert_eq!(json["stream"], true);
|
||||
|
|
@ -188,6 +194,7 @@ mod tests {
|
|||
tool_choice: None,
|
||||
response_format: None,
|
||||
stream: None,
|
||||
stream_options: None,
|
||||
};
|
||||
let json_no_stream = serde_json::to_value(&req_no_stream).unwrap();
|
||||
assert!(json_no_stream.get("stream").is_none());
|
||||
|
|
|
|||
|
|
@ -26,6 +26,16 @@ pub(super) struct ApiRequest {
|
|||
pub response_format: Option<serde_json::Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stream: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stream_options: Option<StreamOptions>,
|
||||
}
|
||||
|
||||
/// Streaming options. Chat Completions only emits the trailing usage chunk
|
||||
/// when the request opts in, so without this a streamed response reports zero
|
||||
/// tokens and costs are estimated at $0.
|
||||
#[derive(serde::Serialize)]
|
||||
pub(super) struct StreamOptions {
|
||||
pub include_usage: bool,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
|
|
|
|||
|
|
@ -812,7 +812,8 @@ async fn stream_text_happy_path_capture() -> (WireCapture, Vec<serde_json::Value
|
|||
stream_capture(&base_request(MODEL), &sse).await
|
||||
}
|
||||
|
||||
/// The captured request pins the stream flag on the wire.
|
||||
/// The captured request pins the streaming request shape, including the usage
|
||||
/// opt-in required for the trailing usage chunk.
|
||||
#[tokio::test]
|
||||
async fn stream_text_happy_path_request() {
|
||||
let (capture, _) = stream_text_happy_path_capture().await;
|
||||
|
|
|
|||
|
|
@ -11,5 +11,8 @@ expression: rendered
|
|||
}
|
||||
],
|
||||
"max_tokens": 128,
|
||||
"stream": true
|
||||
"stream": true,
|
||||
"stream_options": {
|
||||
"include_usage": true
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -43,7 +43,12 @@ pub async fn create_chat_completion(
|
|||
}
|
||||
|
||||
if request.stream {
|
||||
chat_sse_response(&success.plan, success.transport).into_response()
|
||||
chat_sse_response(
|
||||
&success.plan,
|
||||
request.include_stream_usage(),
|
||||
success.transport,
|
||||
)
|
||||
.into_response()
|
||||
} else {
|
||||
Json(success.plan.chat_completions_json()).into_response()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -534,13 +534,27 @@ pub struct ChatCompletionsRequest {
|
|||
pub max_tokens: Option<i64>,
|
||||
#[serde(default)]
|
||||
pub stream: bool,
|
||||
stream_options: Option<ChatStreamOptions>,
|
||||
pub tools: Option<Vec<Value>>,
|
||||
pub tool_choice: Option<Value>,
|
||||
pub response_format: Option<ChatResponseFormat>,
|
||||
pub stop: Option<Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct ChatStreamOptions {
|
||||
#[serde(default)]
|
||||
include_usage: bool,
|
||||
}
|
||||
|
||||
impl ChatCompletionsRequest {
|
||||
pub fn include_stream_usage(&self) -> bool {
|
||||
self.stream_options
|
||||
.as_ref()
|
||||
.is_some_and(|options| options.include_usage)
|
||||
}
|
||||
|
||||
pub fn extract_user_text(&self) -> String {
|
||||
let pieces: Vec<String> = self
|
||||
.messages
|
||||
|
|
|
|||
|
|
@ -236,7 +236,11 @@ pub fn responses_sse_response(plan: &ResponsePlan, transport: TransportOptions)
|
|||
stream_response(events, transport)
|
||||
}
|
||||
|
||||
pub fn chat_sse_response(plan: &ResponsePlan, transport: TransportOptions) -> Response<Body> {
|
||||
pub fn chat_sse_response(
|
||||
plan: &ResponsePlan,
|
||||
include_usage: bool,
|
||||
transport: TransportOptions,
|
||||
) -> Response<Body> {
|
||||
let mut events = Vec::new();
|
||||
let content = plan.chat_content();
|
||||
events.push(chat_chunk(&json!({
|
||||
|
|
@ -321,6 +325,16 @@ pub fn chat_sse_response(plan: &ResponsePlan, transport: TransportOptions) -> Re
|
|||
"finish_reason": if plan.tool_calls.is_empty() { "stop" } else { "tool_calls" },
|
||||
}]
|
||||
})));
|
||||
if include_usage {
|
||||
events.push(chat_chunk(&json!({
|
||||
"id": format!("chatcmpl_{}", plan.id),
|
||||
"object": "chat.completion.chunk",
|
||||
"created": plan.created,
|
||||
"model": plan.model,
|
||||
"choices": [],
|
||||
"usage": plan.usage.chat_completions_json(),
|
||||
})));
|
||||
}
|
||||
events.push("data: [DONE]\n\n".to_owned());
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -50,9 +50,48 @@ async fn chat_completions_stream_uses_same_canonical_plan() {
|
|||
|
||||
assert_eq!(status, 200);
|
||||
assert!(joined.contains("\"content\":\"deterministic: stream same plan\""));
|
||||
assert!(!joined.contains("\"usage\""));
|
||||
assert!(joined.contains("data: [DONE]"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn chat_completions_stream_includes_usage_when_requested() {
|
||||
let server = common::spawn_server().await.expect("server should start");
|
||||
|
||||
let (status, chunks) = server
|
||||
.post_chat_stream(json!({
|
||||
"model": "gpt-test",
|
||||
"messages": [{ "role": "user", "content": "stream with usage" }],
|
||||
"stream": true,
|
||||
"stream_options": { "include_usage": true }
|
||||
}))
|
||||
.await;
|
||||
|
||||
assert_eq!(status, 200);
|
||||
let transcript =
|
||||
common::parse_sse_transcript(chunks.join("").as_bytes()).expect("valid SSE transcript");
|
||||
let usage_chunk = transcript
|
||||
.events
|
||||
.iter()
|
||||
.filter(|event| event.data != "[DONE]")
|
||||
.map(|event| {
|
||||
serde_json::from_str::<serde_json::Value>(&event.data).expect("valid JSON chunk")
|
||||
})
|
||||
.find(|chunk| chunk.get("usage").is_some())
|
||||
.expect("trailing usage chunk");
|
||||
|
||||
assert_eq!(usage_chunk["choices"], json!([]));
|
||||
assert_eq!(
|
||||
usage_chunk["usage"],
|
||||
json!({
|
||||
"prompt_tokens": 3,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 8
|
||||
})
|
||||
);
|
||||
assert!(transcript.done);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn chat_completions_accepts_supported_openai_compatible_fields() {
|
||||
let server = common::spawn_server().await.expect("server should start");
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue