Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
116 changes: 97 additions & 19 deletions apis/src/anthropic/stream_events/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,12 @@ const FINISH_REASON_KEY: &str = "anthropic_stream.finish_reason";
/// Metadata key for accumulated output token count.
const OUTPUT_TOKENS_KEY: &str = "anthropic_stream.output_tokens";

/// Metadata key for accumulated input (prompt) token count.
const INPUT_TOKENS_KEY: &str = "anthropic_stream.input_tokens";

/// Metadata key for cached input token count.
const CACHE_READ_TOKENS_KEY: &str = "anthropic_stream.cache_read_tokens";

/// Metadata key for the current content block index.
const BLOCK_INDEX_KEY: &str = "anthropic_stream.block_index";

Expand Down Expand Up @@ -465,12 +471,20 @@ fn transform_chunk(ctx: &mut HttpFilterContext<'_>, chunk: &Value, output: &mut
}
}

if let Some(ot) = chunk
.get("usage")
.and_then(|u| u.get("completion_tokens"))
.and_then(Value::as_u64)
{
ctx.set_metadata(OUTPUT_TOKENS_KEY, ot.to_string());
if let Some(usage) = chunk.get("usage") {
if let Some(ot) = usage.get("completion_tokens").and_then(Value::as_u64) {
ctx.set_metadata(OUTPUT_TOKENS_KEY, ot.to_string());
}
if let Some(pt) = usage.get("prompt_tokens").and_then(Value::as_u64) {
ctx.set_metadata(INPUT_TOKENS_KEY, pt.to_string());
}
if let Some(ct) = usage
.get("prompt_tokens_details")
.and_then(|d| d.get("cached_tokens"))
.and_then(Value::as_u64)
{
ctx.set_metadata(CACHE_READ_TOKENS_KEY, ct.to_string());
}
}
}

Expand Down Expand Up @@ -707,11 +721,7 @@ fn emit_message_delta(ctx: &HttpFilterContext<'_>, output: &mut Vec<u8>) {
.get(FINISH_REASON_KEY)
.map_or("end_turn", |v| map_stop_reason(v));

let output_tokens: u64 = ctx
.filter_metadata
.get(OUTPUT_TOKENS_KEY)
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let usage = collect_delta_usage(ctx);

emit_event(
output,
Expand All @@ -724,21 +734,39 @@ fn emit_message_delta(ctx: &HttpFilterContext<'_>, output: &mut Vec<u8>) {
"stop_reason": stop_reason,
"stop_sequence": null
},
"usage": message_delta_usage(output_tokens)
"usage": usage
}),
);
}

/// Collect token counts from metadata and build the terminal delta usage.
fn collect_delta_usage(ctx: &HttpFilterContext<'_>) -> MessageDeltaUsage {
let output_tokens: u64 = ctx
.filter_metadata
.get(OUTPUT_TOKENS_KEY)
.and_then(|v| v.parse().ok())
.unwrap_or(0);

let prompt_tokens: Option<u64> = ctx.filter_metadata.get(INPUT_TOKENS_KEY).and_then(|v| v.parse().ok());

let cache_read: Option<u64> = ctx
.filter_metadata
.get(CACHE_READ_TOKENS_KEY)
.and_then(|v| v.parse().ok());

let input_tokens = prompt_tokens.map(|pt| match cache_read {
Some(cached) => pt.saturating_sub(cached),
None => pt,
});

MessageDeltaUsage::new(output_tokens, input_tokens, cache_read)
}

/// Build a schema-complete Anthropic `Message.usage` value.
fn message_start_usage() -> MessageUsage {
MessageUsage::new(0, 0, None)
}

/// Build a schema-complete Anthropic `message_delta.usage` value.
fn message_delta_usage(output_tokens: u64) -> MessageDeltaUsage {
MessageDeltaUsage::new(output_tokens)
}

/// Map `OpenAI` finish reasons to Anthropic stop reasons.
fn map_stop_reason(reason: &str) -> &str {
match reason {
Expand Down Expand Up @@ -987,7 +1015,7 @@ mod tests {
fn message_delta_usage_matches_anthropic_schema() {
let (filter, mut ctx) = make_filter_and_context();

let chunk1 = "data: {\"id\":\"c1\",\"model\":\"gpt-4\",\"choices\":[{\"delta\":{\"content\":\"Hi\"},\"index\":0,\"finish_reason\":\"stop\"}],\"usage\":{\"completion_tokens\":7}}\n\n";
let chunk1 = "data: {\"id\":\"c1\",\"model\":\"gpt-4\",\"choices\":[{\"delta\":{\"content\":\"Hi\"},\"index\":0,\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":15,\"completion_tokens\":7}}\n\n";
let mut body1 = Some(Bytes::from(chunk1));
drop(filter.on_response_body(&mut ctx, &mut body1, false).unwrap());

Expand All @@ -1011,13 +1039,63 @@ mod tests {
&[
"cache_creation_input_tokens",
"cache_read_input_tokens",
"input_tokens",
"server_tool_use",
],
"message_delta usage",
);
assert_absent_fields(usage, &["output_tokens_details"], "message_delta usage");
assert_u64_field(usage, "output_tokens", 7, "message_delta usage");
assert_u64_field(usage, "input_tokens", 15, "message_delta usage");
}

#[test]
fn message_delta_usage_with_cached_tokens() {
let (filter, mut ctx) = make_filter_and_context();

let chunk1 = "data: {\"id\":\"c1\",\"model\":\"gpt-4\",\"choices\":[{\"delta\":{\"content\":\"Hi\"},\"index\":0,\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":5,\"prompt_tokens_details\":{\"cached_tokens\":80}}}\n\n";
let mut body1 = Some(Bytes::from(chunk1));
drop(filter.on_response_body(&mut ctx, &mut body1, false).unwrap());

let done = "data: [DONE]\n\n";
let mut body2 = Some(Bytes::from(done));
drop(filter.on_response_body(&mut ctx, &mut body2, false).unwrap());

let out = String::from_utf8(body2.unwrap().to_vec()).unwrap();
let event = event_data(&out, "message_delta");
let usage = event.get("usage").unwrap();

assert_u64_field(usage, "output_tokens", 5, "message_delta usage");
assert_u64_field(
usage,
"input_tokens",
20,
"input_tokens should exclude cached (100 - 80)",
);
assert_u64_field(usage, "cache_read_input_tokens", 80, "message_delta usage");
}

#[test]
fn message_delta_usage_without_usage_chunk() {
let (filter, mut ctx) = make_filter_and_context();

let chunk1 = "data: {\"id\":\"c1\",\"model\":\"gpt-4\",\"choices\":[{\"delta\":{\"content\":\"Hi\"},\"index\":0,\"finish_reason\":\"stop\"}]}\n\n";
let mut body1 = Some(Bytes::from(chunk1));
drop(filter.on_response_body(&mut ctx, &mut body1, false).unwrap());

let done = "data: [DONE]\n\n";
let mut body2 = Some(Bytes::from(done));
drop(filter.on_response_body(&mut ctx, &mut body2, false).unwrap());

let out = String::from_utf8(body2.unwrap().to_vec()).unwrap();
let event = event_data(&out, "message_delta");
let usage = event.get("usage").unwrap();

assert_u64_field(usage, "output_tokens", 0, "no usage chunk means zero output_tokens");
assert_null_fields(
usage,
&["input_tokens", "cache_read_input_tokens"],
"no usage chunk means null input fields",
);
}

#[test]
Expand Down
60 changes: 57 additions & 3 deletions apis/src/anthropic/to_openai/request.rs
Original file line number Diff line number Diff line change
Expand Up @@ -36,9 +36,7 @@ pub(crate) fn transform_request(body: &[u8]) -> Result<Vec<u8>, String> {
chat.insert("max_completion_tokens".to_owned(), max_tokens.clone());
}

if let Some(stream) = obj.get("stream") {
chat.insert("stream".to_owned(), stream.clone());
}
convert_stream(&mut chat, obj);

map_parameters(&mut chat, obj);
convert_tools(&mut chat, obj);
Expand Down Expand Up @@ -498,6 +496,24 @@ fn non_empty_lines(lines: &[String]) -> Option<String> {
// Parameter Mapping
// -----------------------------------------------------------------------------

/// Copy `stream` and request streaming usage when enabled.
fn convert_stream(chat: &mut Map<String, Value>, obj: &Map<String, Value>) {
let Some(stream) = obj.get("stream") else {
return;
};
chat.insert("stream".to_owned(), stream.clone());

if stream.as_bool() == Some(true) {
let mut opts = obj
.get("stream_options")
.and_then(Value::as_object)
.cloned()
.unwrap_or_default();
opts.insert("include_usage".to_owned(), Value::Bool(true));
chat.insert("stream_options".to_owned(), Value::Object(opts));
}
}

/// Map Anthropic parameters to Chat Completions-compatible equivalents.
///
/// `top_k` has no standard Chat Completions equivalent but is preserved
Expand Down Expand Up @@ -1237,4 +1253,42 @@ mod tests {
assert_eq!(tools.len(), 1, "only non-filtered tools should remain");
assert_eq!(tools[0]["function"]["name"], "get_weather");
}

#[test]
fn streaming_request_includes_usage_option() {
let body = br#"{"model":"claude-opus-4-8","max_tokens":1024,"stream":true,"messages":[{"role":"user","content":"Hi"}]}"#;
let result = transform_request(body).unwrap();
let parsed: Value = serde_json::from_slice(&result).unwrap();

assert_eq!(parsed["stream"], true, "stream should be true");
assert_eq!(
parsed["stream_options"]["include_usage"], true,
"stream_options.include_usage should be set"
);
}

#[test]
fn non_streaming_request_omits_stream_options() {
let body = br#"{"model":"claude-opus-4-8","max_tokens":1024,"messages":[{"role":"user","content":"Hi"}]}"#;
let result = transform_request(body).unwrap();
let parsed: Value = serde_json::from_slice(&result).unwrap();

assert!(
parsed.get("stream_options").is_none(),
"stream_options should not be present without stream:true"
);
}

#[test]
fn stream_false_omits_stream_options() {
let body = br#"{"model":"claude-opus-4-8","max_tokens":1024,"stream":false,"messages":[{"role":"user","content":"Hi"}]}"#;
let result = transform_request(body).unwrap();
let parsed: Value = serde_json::from_slice(&result).unwrap();

assert_eq!(parsed["stream"], false, "stream should be false");
assert!(
parsed.get("stream_options").is_none(),
"stream_options should not be present when stream is false"
);
}
}
8 changes: 4 additions & 4 deletions apis/src/anthropic/wire.rs
Original file line number Diff line number Diff line change
Expand Up @@ -144,12 +144,12 @@ pub(crate) struct MessageDeltaUsage {
}

impl MessageDeltaUsage {
/// Create terminal delta usage from the cumulative output token count.
pub(crate) fn new(output_tokens: u64) -> Self {
/// Create terminal delta usage with token counts.
pub(crate) fn new(output_tokens: u64, input_tokens: Option<u64>, cache_read_input_tokens: Option<u64>) -> Self {
Self {
cache_creation_input_tokens: None,
cache_read_input_tokens: None,
input_tokens: None,
cache_read_input_tokens,
input_tokens,
output_tokens,
server_tool_use: None,
}
Expand Down
Loading
Loading