Skip to content
Draft
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
10 changes: 10 additions & 0 deletions tests/xllm_service/scheduler/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -9,3 +9,13 @@ cc_test(
:xllm_chat_parse_runtime
GTest::gtest_main
)

cc_test(
NAME
response_handler_usage_test
SRCS
response_handler_usage_test.cpp
DEPS
:scheduler
GTest::gtest_main
)
53 changes: 53 additions & 0 deletions tests/xllm_service/scheduler/response_handler_usage_test.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
/* Copyright 2025-2026 The xLLM Authors.

Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at

https://github.com/jd-opensource/xllm-service/blob/main/LICENSE

Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/

#include <gtest/gtest.h>

#include "scheduler/response_handler.h"

namespace xllm_service {
namespace {

TEST(ResponseHandlerUsageTest, IncludesCachedTokensInOpenAIUsage) {
llm::Usage usage;
usage.num_prompt_tokens = 8;
usage.num_generated_tokens = 2;
usage.num_total_tokens = 10;
usage.num_cached_tokens = 6;
xllm::proto::Usage proto_usage;

set_openai_usage(&proto_usage, usage);

EXPECT_EQ(proto_usage.prompt_tokens(), 8);
EXPECT_EQ(proto_usage.completion_tokens(), 2);
EXPECT_EQ(proto_usage.total_tokens(), 10);
ASSERT_TRUE(proto_usage.has_prompt_tokens_details());
EXPECT_TRUE(proto_usage.prompt_tokens_details().has_cached_tokens());
EXPECT_EQ(proto_usage.prompt_tokens_details().cached_tokens(), 6);
}

TEST(ResponseHandlerUsageTest, IncludesZeroCachedTokensInOpenAIUsage) {
llm::Usage usage;
xllm::proto::Usage proto_usage;

set_openai_usage(&proto_usage, usage);

ASSERT_TRUE(proto_usage.has_prompt_tokens_details());
EXPECT_TRUE(proto_usage.prompt_tokens_details().has_cached_tokens());
EXPECT_EQ(proto_usage.prompt_tokens_details().cached_tokens(), 0);
}

} // namespace
} // namespace xllm_service
3 changes: 3 additions & 0 deletions xllm_service/common/xllm/output.h
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,9 @@ struct Usage {

// the total number of tokens used in the request (prompt + completion).
size_t num_total_tokens = 0;

// the number of prompt tokens served from prefix cache.
size_t num_cached_tokens = 0;
};

struct LogProbData {
Expand Down
2 changes: 2 additions & 0 deletions xllm_service/proto/xllm_rpc_service.proto
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,8 @@ message OutputUsage {
int32 num_generated_tokens = 2;
// the total number of tokens used in the request (prompt + completion).
int32 num_total_tokens = 3;
// the number of prompt tokens served from prefix cache.
int32 num_cached_tokens = 4;
}

message LogProbData {
Expand Down
22 changes: 22 additions & 0 deletions xllm_service/rpc_service/rpc_service_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,28 @@ class XllmRpcServiceTest : public ::testing::Test {

void TearDown() override { google::ShutdownGoogleLogging(); }
};

TEST_F(XllmRpcServiceTest, OutputUsageUsesFieldFourForCachedTokens) {
const google::protobuf::FieldDescriptor* field =
proto::OutputUsage::descriptor()->FindFieldByName("num_cached_tokens");

ASSERT_NE(field, nullptr);
EXPECT_EQ(field->number(), 4);
}

TEST_F(XllmRpcServiceTest, ConvertsCachedTokensIntoInternalUsage) {
proto::DisaggStreamGeneration generation;
generation.mutable_usage()->set_num_prompt_tokens(8);
generation.mutable_usage()->set_num_generated_tokens(2);
generation.mutable_usage()->set_num_total_tokens(10);
generation.mutable_usage()->set_num_cached_tokens(6);

llm::RequestOutput output = make_request_output(generation);

ASSERT_TRUE(output.usage.has_value());
EXPECT_EQ(output.usage->num_cached_tokens, 6u);
}

// TODO
// TEST_F(XllmRpcServiceTest, RegisterInstance) {
// RpcServiceConfig config;
Expand Down
119 changes: 64 additions & 55 deletions xllm_service/rpc_service/service.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,68 @@ limitations under the License.

namespace xllm_service {

llm::RequestOutput make_request_output(
const proto::DisaggStreamGeneration& request) {
llm::RequestOutput request_output;
request_output.request_id = request.req_id();
request_output.service_request_id = request.service_req_id();
if (request.has_gen_status()) {
request_output.status = llm::Status(
static_cast<llm::StatusCode>(request.gen_status().status_code()),
request.gen_status().status_msg());
}
if (request.has_usage()) {
llm::Usage usage;
usage.num_prompt_tokens = request.usage().num_prompt_tokens();
usage.num_generated_tokens = request.usage().num_generated_tokens();
usage.num_total_tokens = request.usage().num_total_tokens();
usage.num_cached_tokens = request.usage().num_cached_tokens();
request_output.usage = std::move(usage);
}
request_output.finished_on_prefill_instance =
request.finished_on_prefill_instance();
request_output.finished = request.finished();
request_output.outputs.reserve(request.outputs_size());
for (const auto& output : request.outputs()) {
llm::SequenceOutput sequence_output;
sequence_output.index = output.index();
sequence_output.text = output.text();
sequence_output.token_ids = std::vector<int32_t>(output.token_ids().begin(),
output.token_ids().end());
if (!output.finish_reason().empty()) {
sequence_output.finish_reason = output.finish_reason();
}
if (!output.logprobs().empty()) {
std::vector<llm::LogProb> logprobs;
logprobs.reserve(output.logprobs_size());
for (const auto& logprob : output.logprobs()) {
llm::LogProb lp;
lp.token = logprob.log_prob_data().token();
lp.token_id = logprob.log_prob_data().token_id();
lp.logprob = logprob.log_prob_data().logprob();
lp.finished_token = logprob.log_prob_data().finished_token();
if (!logprob.top_logprobs().empty()) {
std::vector<llm::LogProbData> top_logprobs;
top_logprobs.reserve(logprob.top_logprobs_size());
for (const auto& top_logprob : logprob.top_logprobs()) {
llm::LogProbData lpd;
lpd.token = top_logprob.token();
lpd.token_id = top_logprob.token_id();
lpd.logprob = top_logprob.logprob();
lpd.finished_token = top_logprob.finished_token();
top_logprobs.emplace_back(std::move(lpd));
}
lp.top_logprobs = std::move(top_logprobs);
}
logprobs.emplace_back(std::move(lp));
}
sequence_output.logprobs = std::move(logprobs);
}
request_output.outputs.emplace_back(std::move(sequence_output));
}
return request_output;
}

XllmRpcServiceImpl::XllmRpcServiceImpl(const Options& options,
Scheduler* scheduler)
: options_(options), scheduler_(scheduler) {}
Expand Down Expand Up @@ -145,61 +207,8 @@ void XllmRpcService::Generations(google::protobuf::RpcController* cntl_base,
brpc::ClosureGuard done_guard(done);

// TODO: use threadpool here
for (auto& request : req->gens()) {
// convert proto request to `RequestOutput`
llm::RequestOutput request_output;
request_output.request_id = request.req_id();
request_output.service_request_id = request.service_req_id();
if (request.has_gen_status()) {
request_output.status = llm::Status(
static_cast<llm::StatusCode>(request.gen_status().status_code()),
request.gen_status().status_msg());
}
if (request.has_usage()) {
llm::Usage u;
u.num_prompt_tokens = request.usage().num_prompt_tokens();
u.num_generated_tokens = request.usage().num_generated_tokens();
u.num_total_tokens = request.usage().num_total_tokens();
request_output.usage = std::move(u);
}
request_output.finished_on_prefill_instance =
request.finished_on_prefill_instance();
request_output.finished = request.finished();
for (auto& output : request.outputs()) {
llm::SequenceOutput sequence_output;
sequence_output.index = output.index();
sequence_output.text = output.text();
sequence_output.token_ids = std::vector<int32_t>(
output.token_ids().begin(), output.token_ids().end());
if (!output.finish_reason().empty()) {
sequence_output.finish_reason = output.finish_reason();
}
if (output.logprobs().size() > 0) {
std::vector<llm::LogProb> logprobs;
for (auto& logprob : output.logprobs()) {
llm::LogProb lp;
lp.token = logprob.log_prob_data().token();
lp.token_id = logprob.log_prob_data().token_id();
lp.logprob = logprob.log_prob_data().logprob();
lp.finished_token = logprob.log_prob_data().finished_token();
if (logprob.top_logprobs().size() > 0) {
std::vector<llm::LogProbData> top_logprobs;
for (auto& top_logprob : logprob.top_logprobs()) {
llm::LogProbData lpd;
lpd.token = top_logprob.token();
lpd.token_id = top_logprob.token_id();
lpd.logprob = top_logprob.logprob();
lpd.finished_token = top_logprob.finished_token();
top_logprobs.emplace_back(std::move(lpd));
}
lp.top_logprobs = std::move(top_logprobs);
}
logprobs.emplace_back(std::move(lp));
}
sequence_output.logprobs = std::move(logprobs);
}
request_output.outputs.emplace_back(std::move(sequence_output));
}
for (const auto& request : req->gens()) {
llm::RequestOutput request_output = make_request_output(request);

resp->mutable_all_status()->Add()->set_ok(
xllm_rpc_service_impl_->handle_generation(request_output));
Expand Down
3 changes: 3 additions & 0 deletions xllm_service/rpc_service/service.h
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,9 @@ namespace xllm_service {
class Scheduler;
class InstanceMgr;

llm::RequestOutput make_request_output(
const proto::DisaggStreamGeneration& request);

class XllmRpcServiceImpl final {
public:
XllmRpcServiceImpl(const Options& options, Scheduler* scheduler);
Expand Down
41 changes: 17 additions & 24 deletions xllm_service/scheduler/response_handler.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,19 @@ limitations under the License.
#include "xllm/xllm/function_call/function_call_parser.h"

namespace xllm_service {

void set_openai_usage(xllm::proto::Usage* proto_usage,
const llm::Usage& usage) {
CHECK(proto_usage != nullptr);
proto_usage->set_prompt_tokens(static_cast<int32_t>(usage.num_prompt_tokens));
proto_usage->set_completion_tokens(
static_cast<int32_t>(usage.num_generated_tokens));
proto_usage->set_total_tokens(static_cast<int32_t>(usage.num_total_tokens));
auto* prompt_tokens_details = proto_usage->mutable_prompt_tokens_details();
prompt_tokens_details->set_cached_tokens(
static_cast<int32_t>(usage.num_cached_tokens));
}

namespace {

size_t find_tool_start(const std::string& text) {
Expand Down Expand Up @@ -395,12 +408,7 @@ bool ResponseHandler::send_delta_to_client(
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(
static_cast<int32_t>(usage.num_prompt_tokens));
proto_usage->set_completion_tokens(
static_cast<int32_t>(usage.num_generated_tokens));
proto_usage->set_total_tokens(static_cast<int32_t>(usage.num_total_tokens));
set_openai_usage(response.mutable_usage(), usage);
if (!call_data->write(response)) {
return false;
}
Expand Down Expand Up @@ -476,12 +484,7 @@ bool ResponseHandler::send_delta_to_client(
response.set_created(created_time);
response.set_model(model);
response.mutable_choices();
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(
static_cast<int32_t>(usage.num_prompt_tokens));
proto_usage->set_completion_tokens(
static_cast<int32_t>(usage.num_generated_tokens));
proto_usage->set_total_tokens(static_cast<int32_t>(usage.num_total_tokens));
set_openai_usage(response.mutable_usage(), usage);
if (!call_data->write(response)) {
return false;
}
Expand Down Expand Up @@ -761,12 +764,7 @@ bool ResponseHandler::send_result_to_client(
// add usage statistics
if (req_output.usage.has_value()) {
const auto& usage = req_output.usage.value();
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(
static_cast<int32_t>(usage.num_prompt_tokens));
proto_usage->set_completion_tokens(
static_cast<int32_t>(usage.num_generated_tokens));
proto_usage->set_total_tokens(static_cast<int32_t>(usage.num_total_tokens));
set_openai_usage(response.mutable_usage(), usage);
}

return call_data->write_and_finish(response);
Expand Down Expand Up @@ -809,12 +807,7 @@ bool ResponseHandler::send_result_to_client(
// add usage statistics
if (req_output.usage.has_value()) {
const auto& usage = req_output.usage.value();
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(
static_cast<int32_t>(usage.num_prompt_tokens));
proto_usage->set_completion_tokens(
static_cast<int32_t>(usage.num_generated_tokens));
proto_usage->set_total_tokens(static_cast<int32_t>(usage.num_total_tokens));
set_openai_usage(response.mutable_usage(), usage);
}

return call_data->write_and_finish(response);
Expand Down
2 changes: 2 additions & 0 deletions xllm_service/scheduler/response_handler.h
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,8 @@ namespace xllm_service {

class AnthropicStreamEncoder;

void set_openai_usage(xllm::proto::Usage* proto_usage, const llm::Usage& usage);

struct ChatStreamParseState {
std::unordered_set<size_t> first_message_sent;
std::shared_ptr<xllm::StreamOutputParser> stream_parser;
Expand Down