diff --git a/tests/xllm_service/scheduler/CMakeLists.txt b/tests/xllm_service/scheduler/CMakeLists.txt index 14a1ad5..7c89d3c 100644 --- a/tests/xllm_service/scheduler/CMakeLists.txt +++ b/tests/xllm_service/scheduler/CMakeLists.txt @@ -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 +) diff --git a/tests/xllm_service/scheduler/response_handler_usage_test.cpp b/tests/xllm_service/scheduler/response_handler_usage_test.cpp new file mode 100644 index 0000000..e86956a --- /dev/null +++ b/tests/xllm_service/scheduler/response_handler_usage_test.cpp @@ -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 + +#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 diff --git a/xllm_service/common/xllm/output.h b/xllm_service/common/xllm/output.h index 5c1327a..6458ec9 100644 --- a/xllm_service/common/xllm/output.h +++ b/xllm_service/common/xllm/output.h @@ -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 { diff --git a/xllm_service/proto/xllm_rpc_service.proto b/xllm_service/proto/xllm_rpc_service.proto index f18c9df..c53e31f 100644 --- a/xllm_service/proto/xllm_rpc_service.proto +++ b/xllm_service/proto/xllm_rpc_service.proto @@ -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 { diff --git a/xllm_service/rpc_service/rpc_service_test.cpp b/xllm_service/rpc_service/rpc_service_test.cpp index 4661777..d8d765c 100644 --- a/xllm_service/rpc_service/rpc_service_test.cpp +++ b/xllm_service/rpc_service/rpc_service_test.cpp @@ -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; diff --git a/xllm_service/rpc_service/service.cpp b/xllm_service/rpc_service/service.cpp index 4f9b842..64e0266 100644 --- a/xllm_service/rpc_service/service.cpp +++ b/xllm_service/rpc_service/service.cpp @@ -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(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(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 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 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) {} @@ -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(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( - 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 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 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)); diff --git a/xllm_service/rpc_service/service.h b/xllm_service/rpc_service/service.h index 99e6529..a2e8c7f 100644 --- a/xllm_service/rpc_service/service.h +++ b/xllm_service/rpc_service/service.h @@ -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); diff --git a/xllm_service/scheduler/response_handler.cpp b/xllm_service/scheduler/response_handler.cpp index ba293da..34aaf45 100644 --- a/xllm_service/scheduler/response_handler.cpp +++ b/xllm_service/scheduler/response_handler.cpp @@ -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(usage.num_prompt_tokens)); + proto_usage->set_completion_tokens( + static_cast(usage.num_generated_tokens)); + proto_usage->set_total_tokens(static_cast(usage.num_total_tokens)); + auto* prompt_tokens_details = proto_usage->mutable_prompt_tokens_details(); + prompt_tokens_details->set_cached_tokens( + static_cast(usage.num_cached_tokens)); +} + namespace { size_t find_tool_start(const std::string& text) { @@ -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(usage.num_prompt_tokens)); - proto_usage->set_completion_tokens( - static_cast(usage.num_generated_tokens)); - proto_usage->set_total_tokens(static_cast(usage.num_total_tokens)); + set_openai_usage(response.mutable_usage(), usage); if (!call_data->write(response)) { return false; } @@ -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(usage.num_prompt_tokens)); - proto_usage->set_completion_tokens( - static_cast(usage.num_generated_tokens)); - proto_usage->set_total_tokens(static_cast(usage.num_total_tokens)); + set_openai_usage(response.mutable_usage(), usage); if (!call_data->write(response)) { return false; } @@ -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(usage.num_prompt_tokens)); - proto_usage->set_completion_tokens( - static_cast(usage.num_generated_tokens)); - proto_usage->set_total_tokens(static_cast(usage.num_total_tokens)); + set_openai_usage(response.mutable_usage(), usage); } return call_data->write_and_finish(response); @@ -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(usage.num_prompt_tokens)); - proto_usage->set_completion_tokens( - static_cast(usage.num_generated_tokens)); - proto_usage->set_total_tokens(static_cast(usage.num_total_tokens)); + set_openai_usage(response.mutable_usage(), usage); } return call_data->write_and_finish(response); diff --git a/xllm_service/scheduler/response_handler.h b/xllm_service/scheduler/response_handler.h index 2ae29db..ec4ef15 100644 --- a/xllm_service/scheduler/response_handler.h +++ b/xllm_service/scheduler/response_handler.h @@ -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 first_message_sent; std::shared_ptr stream_parser;