From 944822060cd297c5720d8a1dc366f08ca0ec79ec Mon Sep 17 00:00:00 2001
From: Stevan Freeborn <65925598+StevanFreeborn@users.noreply.github.com>
Date: Wed, 1 Jan 2025 20:13:01 -0600
Subject: [PATCH] feat: add support for count tokens endpoint
---
src/AnthropicClient/AnthropicApiClient.cs | 34 ++++-
src/AnthropicClient/Models/AnthropicModels.cs | 17 ---
.../Models/CountMessageTokensRequest.cs | 72 +++++++++++
.../Models/TokenCountResponse.cs | 15 +++
.../EndToEnd/AnthropicApiClientTests.cs | 17 +++
.../Unit/Models/AnthropicModelsTests.cs | 18 ---
.../Models/CountMessageTokensRequestTests.cs | 117 ++++++++++++++++++
.../Unit/Models/TokenCountResponseTests.cs | 46 +++++++
8 files changed, 297 insertions(+), 39 deletions(-)
create mode 100644 src/AnthropicClient/Models/CountMessageTokensRequest.cs
create mode 100644 src/AnthropicClient/Models/TokenCountResponse.cs
create mode 100644 tests/AnthropicClient.Tests/Unit/Models/CountMessageTokensRequestTests.cs
create mode 100644 tests/AnthropicClient.Tests/Unit/Models/TokenCountResponseTests.cs
diff --git a/src/AnthropicClient/AnthropicApiClient.cs b/src/AnthropicClient/AnthropicApiClient.cs
index 15736f1..faa1353 100644
--- a/src/AnthropicClient/AnthropicApiClient.cs
+++ b/src/AnthropicClient/AnthropicApiClient.cs
@@ -26,6 +26,13 @@ public interface IAnthropicApiClient
/// The message request to create.
/// An asynchronous enumerable that yields the response event by event.
IAsyncEnumerable CreateMessageAsync(StreamMessageRequest request);
+
+ ///
+ /// Counts the tokens in a message asynchronously.
+ ///
+ /// The count message tokens request.
+ /// A task that represents the asynchronous operation. The task result contains the response as an where T is .
+ Task> CountMessageTokensAsync(CountMessageTokensRequest request);
}
///
@@ -34,6 +41,7 @@ public class AnthropicApiClient : IAnthropicApiClient
private const string BaseUrl = "https://api.anthropic.com/v1/";
private const string ApiKeyHeader = "x-api-key";
private const string MessagesEndpoint = "messages";
+ private const string CountTokensEndpoint = "messages/count_tokens";
private const string JsonContentType = "application/json";
private const string EventPrefix = "event:";
private const string DataPrefix = "data:";
@@ -71,7 +79,7 @@ public class AnthropicApiClient : IAnthropicApiClient
///
public async Task> CreateMessageAsync(MessageRequest request)
{
- var response = await SendRequestAsync(request);
+ var response = await SendRequestAsync(MessagesEndpoint, request);
var anthropicHeaders = new AnthropicHeaders(response.Headers);
var responseContent = await response.Content.ReadAsStringAsync();
@@ -94,7 +102,7 @@ public class AnthropicApiClient : IAnthropicApiClient
///
public async IAsyncEnumerable CreateMessageAsync(StreamMessageRequest request)
{
- var response = await SendRequestAsync(request);
+ var response = await SendRequestAsync(MessagesEndpoint, request);
if (response.IsSuccessStatusCode is false)
{
@@ -274,11 +282,29 @@ public class AnthropicApiClient : IAnthropicApiClient
return new ToolCall(tool, toolUse);
}
- private async Task SendRequestAsync(BaseMessageRequest request)
+ ///
+ public async Task> CountMessageTokensAsync(CountMessageTokensRequest request)
+ {
+ var response = await SendRequestAsync(CountTokensEndpoint, request);
+ var anthropicHeaders = new AnthropicHeaders(response.Headers);
+ var responseContent = await response.Content.ReadAsStringAsync();
+
+ if (response.IsSuccessStatusCode is false)
+ {
+ var error = Deserialize(responseContent) ?? new AnthropicError();
+ return AnthropicResult.Failure(error, anthropicHeaders);
+ }
+
+ var msgResponse = Deserialize(responseContent) ?? new TokenCountResponse();
+
+ return AnthropicResult.Success(msgResponse, anthropicHeaders);
+ }
+
+ private async Task SendRequestAsync(string endpoint, T request)
{
var requestJson = Serialize(request);
var requestContent = new StringContent(requestJson, Encoding.UTF8, JsonContentType);
- return await _httpClient.PostAsync(MessagesEndpoint, requestContent);
+ return await _httpClient.PostAsync(endpoint, requestContent);
}
private string Serialize(T obj) => JsonSerializer.Serialize(obj, JsonSerializationOptions.DefaultOptions);
diff --git a/src/AnthropicClient/Models/AnthropicModels.cs b/src/AnthropicClient/Models/AnthropicModels.cs
index f5165d9..158a5a9 100644
--- a/src/AnthropicClient/Models/AnthropicModels.cs
+++ b/src/AnthropicClient/Models/AnthropicModels.cs
@@ -69,21 +69,4 @@ public static class AnthropicModels
/// The Claude 3.5 Haiku model.
///
public const string Claude35HaikuLatest = "claude-3-5-haiku-latest";
-
- internal static bool IsValidModel(string modelId) => modelId is
- Claude3Opus or
- Claude3Opus20241022 or
- Claude3OpusLatest or
-
- Claude3Sonnet or
- Claude3Sonnet20240229 or
- Claude35Sonnet or
- Claude35Sonnet20240620 or
- Claude35Sonnet20241022 or
- Claude35SonnetLatest or
-
- Claude3Haiku or
- Claude3Haiku20240307 or
- Claude35Haiku20241022 or
- Claude35HaikuLatest;
}
\ No newline at end of file
diff --git a/src/AnthropicClient/Models/CountMessageTokensRequest.cs b/src/AnthropicClient/Models/CountMessageTokensRequest.cs
new file mode 100644
index 0000000..6779121
--- /dev/null
+++ b/src/AnthropicClient/Models/CountMessageTokensRequest.cs
@@ -0,0 +1,72 @@
+using System.Text.Json.Serialization;
+
+using AnthropicClient.Utils;
+
+namespace AnthropicClient.Models;
+
+///
+/// Represents a request to count the number of tokens in a message.
+///
+public class CountMessageTokensRequest
+{
+ ///
+ /// Gets the model ID to be used for the request.
+ ///
+ public string Model { get; init; } = string.Empty;
+
+ ///
+ /// Gets the messages to count the number of tokens in.
+ ///
+ public List Messages { get; init; } = [];
+
+ ///
+ /// Gets the tool choice mode to use for the request.
+ ///
+ [JsonPropertyName("tool_choice")]
+ public ToolChoice? ToolChoice { get; init; } = null;
+
+ ///
+ /// Gets the tools to use for the request.
+ ///
+ public List? Tools { get; init; } = null;
+
+ ///
+ /// Gets the system prompt to use for the request.
+ ///
+ [JsonPropertyName("system")]
+ public List? SystemPrompt { get; init; } = null;
+
+ ///
+ /// Initializes a new instance of the class.
+ ///
+ /// The model ID to use for the request.
+ /// The messages to count the number of tokens in.
+ /// The tool choice mode to use for the request.
+ /// The tools to use for the request.
+ /// The system prompt to use for the request.
+ /// Thrown when or is null.
+ /// Thrown when is empty.
+ /// A new instance of the class.
+ public CountMessageTokensRequest(
+ string model,
+ List messages,
+ ToolChoice? toolChoice = null,
+ List? tools = null,
+ List? systemPrompt = null
+ )
+ {
+ ArgumentValidator.ThrowIfNull(model, nameof(model));
+ ArgumentValidator.ThrowIfNull(messages, nameof(messages));
+
+ if (messages.Count < 1)
+ {
+ throw new ArgumentException("Messages must contain at least one message");
+ }
+
+ Model = model;
+ Messages = messages;
+ ToolChoice = toolChoice;
+ Tools = tools;
+ SystemPrompt = systemPrompt;
+ }
+}
\ No newline at end of file
diff --git a/src/AnthropicClient/Models/TokenCountResponse.cs b/src/AnthropicClient/Models/TokenCountResponse.cs
new file mode 100644
index 0000000..ec532de
--- /dev/null
+++ b/src/AnthropicClient/Models/TokenCountResponse.cs
@@ -0,0 +1,15 @@
+using System.Text.Json.Serialization;
+
+namespace AnthropicClient.Models;
+
+///
+/// Represents a response to a token count request.
+///
+public class TokenCountResponse
+{
+ ///
+ /// The number of input tokens counted.
+ ///
+ [JsonPropertyName("input_tokens")]
+ public int InputTokens { get; init; }
+}
\ No newline at end of file
diff --git a/tests/AnthropicClient.Tests/EndToEnd/AnthropicApiClientTests.cs b/tests/AnthropicClient.Tests/EndToEnd/AnthropicApiClientTests.cs
index d66e39e..75040e8 100644
--- a/tests/AnthropicClient.Tests/EndToEnd/AnthropicApiClientTests.cs
+++ b/tests/AnthropicClient.Tests/EndToEnd/AnthropicApiClientTests.cs
@@ -299,4 +299,21 @@ public class ClientTests(ConfigurationFixture configFixture) : EndToEndTest(conf
resultTwo.Value.Content.Should().NotBeNullOrEmpty();
resultTwo.Value.Usage.CacheReadInputTokens.Should().BeGreaterThan(0);
}
+
+ [Fact]
+ public async Task CountMessageTokensAsync_WhenCalled_ItShouldReturnResponse()
+ {
+ var request = new CountMessageTokensRequest(
+ model: AnthropicModels.Claude3Haiku,
+ messages: [
+ new(MessageRole.User, [new TextContent("Hello!")])
+ ]
+ );
+
+ var result = await _client.CountMessageTokensAsync(request);
+
+ result.IsSuccess.Should().BeTrue();
+ result.Value.Should().BeOfType();
+ result.Value.InputTokens.Should().BeGreaterThan(0);
+ }
}
\ No newline at end of file
diff --git a/tests/AnthropicClient.Tests/Unit/Models/AnthropicModelsTests.cs b/tests/AnthropicClient.Tests/Unit/Models/AnthropicModelsTests.cs
index 396edbd..3b2343a 100644
--- a/tests/AnthropicClient.Tests/Unit/Models/AnthropicModelsTests.cs
+++ b/tests/AnthropicClient.Tests/Unit/Models/AnthropicModelsTests.cs
@@ -131,22 +131,4 @@ public class AnthropicModelsTests
actual.Should().Be(expected);
}
-
- [Theory]
- [InlineData("claude-3-opus-20240229", true)]
- [InlineData("claude-3-opus-latest", true)]
- [InlineData("claude-3-sonnet-20240229", true)]
- [InlineData("claude-3-5-sonnet-20240620", true)]
- [InlineData("claude-3-5-sonnet-20241022", true)]
- [InlineData("claude-3-5-sonnet-latest", true)]
- [InlineData("claude-3-haiku-20240307", true)]
- [InlineData("claude-3-5-haiku-20241022", true)]
- [InlineData("claude-3-5-haiku-latest", true)]
- [InlineData("invalid", false)]
- public void IsValidModel_WhenCalled_ItShouldReturnExpectedValue(string modelId, bool expected)
- {
- var actual = AnthropicModels.IsValidModel(modelId);
-
- actual.Should().Be(expected);
- }
}
\ No newline at end of file
diff --git a/tests/AnthropicClient.Tests/Unit/Models/CountMessageTokensRequestTests.cs b/tests/AnthropicClient.Tests/Unit/Models/CountMessageTokensRequestTests.cs
new file mode 100644
index 0000000..2f9b868
--- /dev/null
+++ b/tests/AnthropicClient.Tests/Unit/Models/CountMessageTokensRequestTests.cs
@@ -0,0 +1,117 @@
+namespace AnthropicClient.Tests.Unit.Models;
+
+public class CountMessageTokensRequestTests : SerializationTest
+{
+ private readonly string _testJson = @"{
+ ""model"": ""claude-3-sonnet-20240229"",
+ ""system"": [{
+ ""type"": ""text"",
+ ""text"": ""test-system""
+ }],
+ ""messages"": [
+ { ""role"": ""user"", ""content"": [{ ""text"": ""Hello!"", ""type"": ""text"" }] }
+ ],
+ ""tool_choice"": { ""type"":""auto"" },
+ ""tools"": []
+ }";
+
+ [Fact]
+ public void Constructor_WhenCalled_ItShouldInitializeProperties()
+ {
+ var model = AnthropicModels.Claude3Sonnet;
+ var messages = new List { new() };
+ var systemPrompt = new List() { new("test-system") };
+ var toolChoice = new AutoToolChoice();
+ var tools = new List();
+
+ var request = new CountMessageTokensRequest(
+ model: model,
+ messages: messages,
+ toolChoice: toolChoice,
+ tools: tools,
+ systemPrompt: systemPrompt
+ );
+
+ request.Model.Should().Be(model);
+ request.Messages.Should().BeSameAs(messages);
+ request.ToolChoice.Should().Be(toolChoice);
+ request.Tools.Should().BeSameAs(tools);
+ request.SystemPrompt.Should().BeSameAs(systemPrompt);
+ }
+
+ [Fact]
+ public void Constructor_WhenCalledAndModelIsNull_ItShouldThrowArgumentNullException()
+ {
+ var action = () => new CountMessageTokensRequest(
+ model: null!,
+ messages: [new()]
+ );
+
+ action.Should().Throw();
+ }
+
+ [Fact]
+ public void Constructor_WhenCalledAndMessagesIsNull_ItShouldThrowArgumentNullException()
+ {
+ var action = () => new CountMessageTokensRequest(
+ model: AnthropicModels.Claude3Sonnet,
+ messages: null!
+ );
+
+ action.Should().Throw();
+ }
+
+ [Fact]
+ public void Constructor_WhenCalledAndMessagesIsEmpty_ItShouldThrowArgumentException()
+ {
+ var action = () => new CountMessageTokensRequest(
+ model: AnthropicModels.Claude3Sonnet,
+ messages: []
+ );
+
+ action.Should().Throw();
+ }
+
+ [Fact]
+ public void JsonSerialization_WhenSerialized_ItShouldHaveExpectedShape()
+ {
+ var messages = new List()
+ {
+ new()
+ {
+ Role = MessageRole.User,
+ Content = [new TextContent("Hello!")]
+ }
+ };
+
+ var model = AnthropicModels.Claude3Sonnet;
+ var systemPrompt = new List() { new("test-system") };
+ var toolChoice = new AutoToolChoice();
+ var tools = new List();
+
+ var request = new CountMessageTokensRequest(
+ model: model,
+ messages: messages,
+ toolChoice: toolChoice,
+ tools: tools,
+ systemPrompt: systemPrompt
+ );
+
+ var actual = Serialize(request);
+
+ JsonAssert.Equal(_testJson, actual);
+ }
+
+ [Fact]
+ public void JsonDeserialization_WhenDeserialized_ItShouldHaveExpectedShape()
+ {
+ var request = Deserialize(_testJson);
+
+ request!.Model.Should().Be(AnthropicModels.Claude3Sonnet);
+ request.SystemPrompt.Should().BeEquivalentTo(new List { new("test-system") });
+ request.Messages.Should().HaveCount(1);
+ request.ToolChoice.Should().BeOfType();
+ request.ToolChoice!.Type.Should().Be("auto");
+ request.Tools.Should().HaveCount(0);
+ }
+}
\ No newline at end of file
diff --git a/tests/AnthropicClient.Tests/Unit/Models/TokenCountResponseTests.cs b/tests/AnthropicClient.Tests/Unit/Models/TokenCountResponseTests.cs
new file mode 100644
index 0000000..21f2c38
--- /dev/null
+++ b/tests/AnthropicClient.Tests/Unit/Models/TokenCountResponseTests.cs
@@ -0,0 +1,46 @@
+namespace AnthropicClient.Tests.Unit.Models;
+
+public class TokenCountResponseTests : SerializationTest
+{
+ [Fact]
+ public void Constructor_WhenCalled_ShouldInitializeProperties()
+ {
+ var expectedTokenCount = 1;
+
+ var response = new TokenCountResponse
+ {
+ InputTokens = expectedTokenCount
+ };
+
+ response.InputTokens.Should().Be(expectedTokenCount);
+ }
+
+ [Fact]
+ public void JsonSerialization_WhenCalled_ItShouldSerializeCorrectly()
+ {
+ var expectedJson = @"{
+ ""input_tokens"": 1
+ }";
+
+ var response = new TokenCountResponse
+ {
+ InputTokens = 1
+ };
+
+ var actual = Serialize(response);
+
+ JsonAssert.Equal(expectedJson, actual);
+ }
+
+ [Fact]
+ public void JsonDeserialization_WhenCalled_ItShouldDeserializeCorrectly()
+ {
+ var json = @"{
+ ""input_tokens"": 1
+ }";
+
+ var response = Deserialize(json);
+
+ response!.InputTokens.Should().Be(1);
+ }
+}
\ No newline at end of file