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