Skip to content

Commit 49a59de

Browse files
chaunceyjiangCopilot
authored andcommitted
[Feat] Supports Anthropic Messages count_tokens API (vllm-project#35588)
Signed-off-by: chaunceyjiang <chaunceyjiang@gmail.com>
1 parent 60c5d07 commit 49a59de

3 files changed

Lines changed: 332 additions & 133 deletions

File tree

vllm/entrypoints/anthropic/api_router.py

Lines changed: 55 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,8 @@
88
from fastapi.responses import JSONResponse, StreamingResponse
99

1010
from vllm.entrypoints.anthropic.protocol import (
11+
AnthropicCountTokensRequest,
12+
AnthropicCountTokensResponse,
1113
AnthropicError,
1214
AnthropicErrorResponse,
1315
AnthropicMessagesRequest,
@@ -31,6 +33,18 @@ def messages(request: Request) -> AnthropicServingMessages:
3133
return request.app.state.anthropic_serving_messages
3234

3335

36+
def translate_error_response(response: ErrorResponse) -> JSONResponse:
37+
anthropic_error = AnthropicErrorResponse(
38+
error=AnthropicError(
39+
type=response.error.type,
40+
message=response.error.message,
41+
)
42+
)
43+
return JSONResponse(
44+
status_code=response.error.code, content=anthropic_error.model_dump()
45+
)
46+
47+
3448
@router.post(
3549
"/v1/messages",
3650
dependencies=[Depends(validate_json_request)],
@@ -44,17 +58,6 @@ def messages(request: Request) -> AnthropicServingMessages:
4458
@with_cancellation
4559
@load_aware_call
4660
async def create_messages(request: AnthropicMessagesRequest, raw_request: Request):
47-
def translate_error_response(response: ErrorResponse) -> JSONResponse:
48-
anthropic_error = AnthropicErrorResponse(
49-
error=AnthropicError(
50-
type=response.error.type,
51-
message=response.error.message,
52-
)
53-
)
54-
return JSONResponse(
55-
status_code=response.error.code, content=anthropic_error.model_dump()
56-
)
57-
5861
handler = messages(raw_request)
5962
if handler is None:
6063
base_server = raw_request.app.state.openai_serving_tokenization
@@ -88,5 +91,46 @@ def translate_error_response(response: ErrorResponse) -> JSONResponse:
8891
return StreamingResponse(content=generator, media_type="text/event-stream")
8992

9093

94+
@router.post(
95+
"/v1/messages/count_tokens",
96+
dependencies=[Depends(validate_json_request)],
97+
responses={
98+
HTTPStatus.OK.value: {"model": AnthropicCountTokensResponse},
99+
HTTPStatus.BAD_REQUEST.value: {"model": AnthropicErrorResponse},
100+
HTTPStatus.NOT_FOUND.value: {"model": AnthropicErrorResponse},
101+
HTTPStatus.INTERNAL_SERVER_ERROR.value: {"model": AnthropicErrorResponse},
102+
},
103+
)
104+
@load_aware_call
105+
@with_cancellation
106+
async def count_tokens(request: AnthropicCountTokensRequest, raw_request: Request):
107+
handler = messages(raw_request)
108+
if handler is None:
109+
base_server = raw_request.app.state.openai_serving_tokenization
110+
error = base_server.create_error_response(
111+
message="The model does not support Messages API"
112+
)
113+
return translate_error_response(error)
114+
115+
try:
116+
response = await handler.count_tokens(request, raw_request)
117+
except Exception as e:
118+
logger.exception("Error in count_tokens: %s", e)
119+
return JSONResponse(
120+
status_code=HTTPStatus.INTERNAL_SERVER_ERROR.value,
121+
content=AnthropicErrorResponse(
122+
error=AnthropicError(
123+
type="internal_error",
124+
message=str(e),
125+
)
126+
).model_dump(),
127+
)
128+
129+
if isinstance(response, ErrorResponse):
130+
return translate_error_response(response)
131+
132+
return JSONResponse(content=response.model_dump(exclude_none=True))
133+
134+
91135
def attach_router(app: FastAPI):
92136
app.include_router(router)

vllm/entrypoints/anthropic/protocol.py

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -175,3 +175,33 @@ class AnthropicMessagesResponse(BaseModel):
175175
def model_post_init(self, __context):
176176
if not self.id:
177177
self.id = f"msg_{int(time.time() * 1000)}"
178+
179+
180+
class AnthropicContextManagement(BaseModel):
181+
"""Context management information for token counting."""
182+
183+
original_input_tokens: int
184+
185+
186+
class AnthropicCountTokensRequest(BaseModel):
187+
"""Anthropic messages.count_tokens request"""
188+
189+
model: str
190+
messages: list[AnthropicMessage]
191+
system: str | list[AnthropicContentBlock] | None = None
192+
tool_choice: AnthropicToolChoice | None = None
193+
tools: list[AnthropicTool] | None = None
194+
195+
@field_validator("model")
196+
@classmethod
197+
def validate_model(cls, v):
198+
if not v:
199+
raise ValueError("Model is required")
200+
return v
201+
202+
203+
class AnthropicCountTokensResponse(BaseModel):
204+
"""Anthropic messages.count_tokens response"""
205+
206+
input_tokens: int
207+
context_management: AnthropicContextManagement | None = None

0 commit comments

Comments
 (0)