From 035721f199729de78dcf8b29ece0db7e8725adee Mon Sep 17 00:00:00 2001 From: HuitaePark Date: Mon, 20 Jul 2026 13:12:42 +0900 Subject: [PATCH 1/6] =?UTF-8?q?feat(TokenUsage):=20TokenUsage=EC=9D=98=20?= =?UTF-8?q?=EC=B6=94=EB=A1=A0=20=ED=86=A0=ED=81=B0=20=EA=B3=84=EC=82=B0=20?= =?UTF-8?q?=EB=B6=80=EB=B6=84=EC=9D=84=20=EC=A0=9C=EA=B1=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../io/tokenpilot/core/domain/TokenType.java | 9 ++++++- .../io/tokenpilot/core/domain/TokenUsage.java | 6 ++++- .../src/test/java/TokenUsageTest.java | 24 +++++++++++++++++++ .../internal/DefaultCostCalculatorTest.java | 1 - .../internal/InMemoryPricingRegistryTest.java | 1 - 5 files changed, 37 insertions(+), 4 deletions(-) create mode 100644 token-pilot-core/src/test/java/TokenUsageTest.java diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java index 24abfe8..6f175af 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java @@ -26,6 +26,13 @@ public boolean isPrompt() { * 해당 토큰 타입이 출력(Completion) 계열인지 확인합니다. */ public boolean isCompletion() { - return this == COMPLETION || this == REASONING || this == CACHED_COMPLETION; + return this == COMPLETION || this == CACHED_COMPLETION; + } + + /** + * 해당 토큰 타입이 추론(Reasoning) 계열인지 확인합니다. + */ + public boolean isReasoning(){ + return this == REASONING; } } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java index 05861aa..bcbd4ea 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java @@ -3,6 +3,7 @@ import java.util.Collections; import java.util.EnumMap; import java.util.Map; +import java.util.Map.Entry; /** * AI 모델 호출 시 발생하는 토큰 사용량 정보. @@ -65,7 +66,10 @@ public long completionTokens() { * 전체 사용 토큰 수의 합계를 반환합니다. */ public long totalTokens() { - return tokenCounts.values().stream().mapToLong(Long::longValue).sum(); + return tokenCounts.entrySet().stream() + .filter(e -> !e.getKey().isReasoning()) + .mapToLong(Entry::getValue) + .sum(); } /** diff --git a/token-pilot-core/src/test/java/TokenUsageTest.java b/token-pilot-core/src/test/java/TokenUsageTest.java new file mode 100644 index 0000000..b7ab7fd --- /dev/null +++ b/token-pilot-core/src/test/java/TokenUsageTest.java @@ -0,0 +1,24 @@ +import static org.assertj.core.api.Assertions.assertThat; + +import io.tokenpilot.core.domain.TokenType; +import io.tokenpilot.core.domain.TokenUsage; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +public class TokenUsageTest { + + @Test + @DisplayName("reasoning 토큰은 output total의 부분집합이어야 한다") + void shouldNotAddReasoningTokensToOutputTotalAgain() { + TokenUsage usage = TokenUsage.from( + 100, + 200, + 150 + ); + + assertThat(usage.promptTokens()).isEqualTo(100); + assertThat(usage.completionTokens()).isEqualTo(200); + assertThat(usage.totalTokens()).isEqualTo(300); + assertThat(usage.getCount(TokenType.REASONING)).isEqualTo(150); + } +} diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java index ac07762..410b0ed 100644 --- a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java @@ -4,7 +4,6 @@ import io.tokenpilot.core.domain.PricingPlan; import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; -import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/InMemoryPricingRegistryTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/InMemoryPricingRegistryTest.java index 01d3033..8c27125 100644 --- a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/InMemoryPricingRegistryTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/InMemoryPricingRegistryTest.java @@ -1,7 +1,6 @@ package io.tokenpilot.core.internal; import io.tokenpilot.core.domain.PricingPlan; -import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; From ba919298f0d17e36e3f16ce9966eacf8b03830f2 Mon Sep 17 00:00:00 2001 From: HuitaePark Date: Mon, 20 Jul 2026 14:14:55 +0900 Subject: [PATCH 2/6] =?UTF-8?q?refactor(TokenUsage):=20=EC=83=81=EC=84=B8?= =?UTF-8?q?=20=EB=AA=A8=EB=8D=B8=EA=B3=BC=20=ED=98=B8=EC=B6=9C=EB=B6=80?= =?UTF-8?q?=EB=A5=BC=20=EC=83=88=20=ED=95=84=EB=93=9C=EB=A1=9C=20=EC=9D=B4?= =?UTF-8?q?=EA=B4=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../io/tokenpilot/core/domain/TokenUsage.java | 59 ++++++++++--------- .../core/domain/TokenUsageDetails.java | 7 +++ .../core/internal/DefaultCostCalculator.java | 7 +-- .../src/test/java/TokenUsageTest.java | 22 +++++++ .../internal/DefaultCostCalculatorTest.java | 12 ++-- .../internal/MicroCostMetricsPublisher.java | 10 ++-- .../internal/DefaultUsageExtractor.java | 27 +++------ 7 files changed, 82 insertions(+), 62 deletions(-) create mode 100644 token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java index bcbd4ea..28154a0 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java @@ -1,23 +1,24 @@ package io.tokenpilot.core.domain; import java.util.Collections; -import java.util.EnumMap; import java.util.Map; -import java.util.Map.Entry; /** * AI 모델 호출 시 발생하는 토큰 사용량 정보. - * 세부적인 {@link TokenType} 별로 사용량을 관리합니다. + * 전체 입력/출력 토큰과 세부 사용량을 관리합니다. * - * @param tokenCounts 토큰 타입별 사용량 (Map) - * @param metadata 추가 메타데이터 (예: 모델 정보 등) + * @param inputTokens 전체 입력 토큰 + * @param outputTokens 전체 출력 토큰 + * @param details 세부 토큰 사용량 + * @param metadata 추가 메타데이터 (예: 모델 정보 등) */ public record TokenUsage( - Map tokenCounts, + long inputTokens, + long outputTokens, + TokenUsageDetails details, Map metadata ) { public TokenUsage { - tokenCounts = Collections.unmodifiableMap(new EnumMap<>(tokenCounts)); metadata = (metadata != null) ? Collections.unmodifiableMap(metadata) : Map.of(); } @@ -25,57 +26,57 @@ public record TokenUsage( * 기본 입력/출력 토큰을 사용하는 {@link TokenUsage}를 생성합니다. */ public static TokenUsage from(long prompt, long completion) { - Map counts = new EnumMap<>(TokenType.class); - counts.put(TokenType.PROMPT, prompt); - counts.put(TokenType.COMPLETION, completion); - return new TokenUsage(counts, Map.of()); + return new TokenUsage( + prompt, + completion, + new TokenUsageDetails(0, 0, 0), + Map.of() + ); } /** * 입력/출력/추론 토큰을 포함하는 {@link TokenUsage}를 생성합니다. */ public static TokenUsage from(long prompt, long completion, long reasoning) { - Map counts = new EnumMap<>(TokenType.class); - counts.put(TokenType.PROMPT, prompt); - counts.put(TokenType.COMPLETION, completion); - counts.put(TokenType.REASONING, reasoning); - return new TokenUsage(counts, Map.of()); + return new TokenUsage( + prompt, + completion, + new TokenUsageDetails(0, reasoning, 0), + Map.of() + ); } /** * 모든 종류의 입력/출력 토큰 수의 합계를 반환합니다. */ public long promptTokens() { - return tokenCounts.entrySet().stream() - .filter(e -> e.getKey().isPrompt()) - .mapToLong(Map.Entry::getValue) - .sum(); + return inputTokens; } /** * 모든 종류의 출력(추론 포함) 토큰 수의 합계를 반환합니다. */ public long completionTokens() { - return tokenCounts.entrySet().stream() - .filter(e -> e.getKey().isCompletion()) - .mapToLong(Map.Entry::getValue) - .sum(); + return outputTokens; } /** * 전체 사용 토큰 수의 합계를 반환합니다. */ public long totalTokens() { - return tokenCounts.entrySet().stream() - .filter(e -> !e.getKey().isReasoning()) - .mapToLong(Entry::getValue) - .sum(); + return inputTokens + outputTokens; } /** * 특정 토큰 타입의 사용량을 가져옵니다. 없을 시 0을 반환합니다. */ public long getCount(TokenType type) { - return tokenCounts.getOrDefault(type, 0L); + return switch (type) { + case PROMPT -> inputTokens; + case COMPLETION -> outputTokens; + case REASONING -> details.reasoningOutputTokens(); + case CACHED_PROMPT -> details.cachedInputTokens(); + case CACHED_COMPLETION -> details.cachedOutputTokens(); + }; } } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java new file mode 100644 index 0000000..40d07d9 --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java @@ -0,0 +1,7 @@ +package io.tokenpilot.core.domain; + +public record TokenUsageDetails( + long cachedInputTokens, + long reasoningOutputTokens, + long cachedOutputTokens +) {} diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultCostCalculator.java b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultCostCalculator.java index a2d8c79..c6880dc 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultCostCalculator.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultCostCalculator.java @@ -9,8 +9,6 @@ import java.math.BigDecimal; import java.math.RoundingMode; -import java.util.Map; - /** * 기본 비용 계산기 구현체. * 각 {@link TokenType} 별 단가를 적용하여 정밀하게 계산합니다. @@ -24,9 +22,8 @@ public Cost calculate(TokenUsage usage, PricingPlan plan) { BigDecimal totalCostValue = BigDecimal.ZERO; // 사용된 모든 토큰 타입에 대해 각각의 단가를 적용하여 합산 - for (Map.Entry entry : usage.tokenCounts().entrySet()) { - TokenType type = entry.getKey(); - Long count = entry.getValue(); + for (TokenType type : TokenType.values()) { + long count = usage.getCount(type); if (count > 0) { BigDecimal rate = plan.getRate(type); diff --git a/token-pilot-core/src/test/java/TokenUsageTest.java b/token-pilot-core/src/test/java/TokenUsageTest.java index b7ab7fd..0887b14 100644 --- a/token-pilot-core/src/test/java/TokenUsageTest.java +++ b/token-pilot-core/src/test/java/TokenUsageTest.java @@ -2,6 +2,8 @@ import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; +import io.tokenpilot.core.domain.TokenUsageDetails; +import java.util.Map; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; @@ -21,4 +23,24 @@ void shouldNotAddReasoningTokensToOutputTotalAgain() { assertThat(usage.totalTokens()).isEqualTo(300); assertThat(usage.getCount(TokenType.REASONING)).isEqualTo(150); } + + @Test + @DisplayName("cache의 input이 중복되면 안된다.") + void shouldNotAddCachedInputTokensToInputTotalAgain() { + TokenUsage usage = new TokenUsage( + 100, + 200, + new TokenUsageDetails( + 40, // cachedInputTokens + 0, // reasoningOutputTokens + 0 // cachedOutputTokens + ), + Map.of() + ); + + assertThat(usage.promptTokens()).isEqualTo(100); + assertThat(usage.completionTokens()).isEqualTo(200); + assertThat(usage.totalTokens()).isEqualTo(300); + assertThat(usage.getCount(TokenType.CACHED_PROMPT)).isEqualTo(40); + } } diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java index 410b0ed..86cf088 100644 --- a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java @@ -4,6 +4,7 @@ import io.tokenpilot.core.domain.PricingPlan; import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; +import io.tokenpilot.core.domain.TokenUsageDetails; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; @@ -43,11 +44,12 @@ void calculateReasoningTokensWithDifferentRate() { PricingPlan plan = new PricingPlan("o1-preview", rates, Cost.DEFAULT_CURRENCY); // Usage: Prompt 1000, Completion 500, Reasoning 1500 (Total 3000) - Map counts = new EnumMap<>(TokenType.class); - counts.put(TokenType.PROMPT, 1000L); - counts.put(TokenType.COMPLETION, 500L); - counts.put(TokenType.REASONING, 1500L); - TokenUsage usage = new TokenUsage(counts, Map.of()); + TokenUsage usage = new TokenUsage( + 1000, + 500, + new TokenUsageDetails(0, 1500, 0), + Map.of() + ); // When Cost cost = calculator.calculate(usage, plan); diff --git a/token-pilot-micrometer/src/main/java/io/tokenpilot/micrometer/internal/MicroCostMetricsPublisher.java b/token-pilot-micrometer/src/main/java/io/tokenpilot/micrometer/internal/MicroCostMetricsPublisher.java index a2fb9dc..1744581 100644 --- a/token-pilot-micrometer/src/main/java/io/tokenpilot/micrometer/internal/MicroCostMetricsPublisher.java +++ b/token-pilot-micrometer/src/main/java/io/tokenpilot/micrometer/internal/MicroCostMetricsPublisher.java @@ -5,8 +5,9 @@ import io.micrometer.core.instrument.MeterRegistry; import io.micrometer.core.instrument.Tag; import io.micrometer.core.instrument.Tags; -import io.tokenpilot.core.domain.CostRecordedEvent; import io.tokenpilot.core.LedgerListener; +import io.tokenpilot.core.domain.CostRecordedEvent; +import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.micrometer.MetricsOptions; import java.util.Map; @@ -43,7 +44,8 @@ public void onRecord(CostRecordedEvent event) { final Tags finalTags = commonTags; - event.usage().tokenCounts().forEach(((tokenType, count) -> { + for (TokenType tokenType : TokenType.values()) { + long count = event.usage().getCount(tokenType); if (count > 0) { String typeName = tokenType.name().toLowerCase(); DistributionSummary.builder("ai.token.usage.distribution") @@ -51,7 +53,7 @@ public void onRecord(CostRecordedEvent event) { .baseUnit("tokens") .tags(finalTags.and("token_type", typeName)) .register(meterRegistry) - .record(count.doubleValue()); + .record(count); Counter.builder("ai.token.usage.total") .description("Total number of AI tokens recorded") .baseUnit("tokens") @@ -59,7 +61,7 @@ public void onRecord(CostRecordedEvent event) { .register(meterRegistry) .increment(count); } - })); + } Counter.builder("ai.token.cost.total") .description("Total estimated AI token cost") .baseUnit("currency") diff --git a/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultUsageExtractor.java b/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultUsageExtractor.java index bd0c682..3d0602d 100644 --- a/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultUsageExtractor.java +++ b/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultUsageExtractor.java @@ -1,14 +1,13 @@ package io.tokenpilot.springai.internal; -import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; +import io.tokenpilot.core.domain.TokenUsageDetails; import io.tokenpilot.springai.UsageExtractor; import org.springframework.ai.chat.client.ChatClientResponse; import org.springframework.ai.chat.metadata.ChatResponseMetadata; import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.chat.model.ChatResponse; -import java.util.EnumMap; import java.util.HashMap; import java.util.Map; @@ -33,23 +32,20 @@ public TokenUsage extract(ChatClientResponse response) { ChatResponseMetadata metadata = chatResponse.getMetadata(); Usage usage = metadata.getUsage(); if (usage == null) { - return new TokenUsage(defaultCounts(), copyMetadata(metadata, null)); + return new TokenUsage(0, 0, new TokenUsageDetails(0, 0, 0), copyMetadata(metadata, null)); } - Map counts = new EnumMap<>(TokenType.class); - long prompt = (usage.getPromptTokens() != null) ? usage.getPromptTokens() : 0L; long completion = (usage.getCompletionTokens() != null) ? usage.getCompletionTokens() : 0L; - - counts.put(TokenType.PROMPT, prompt); - counts.put(TokenType.COMPLETION, completion); Long reasoning = extractReasoningTokens(metadata); - if (reasoning > 0) { - counts.put(TokenType.REASONING, reasoning); - } - return new TokenUsage(counts, copyMetadata(metadata, usage.getNativeUsage())); + return new TokenUsage( + prompt, + completion, + new TokenUsageDetails(0, reasoning, 0), + copyMetadata(metadata, usage.getNativeUsage()) + ); } private Long extractReasoningTokens(ChatResponseMetadata metadata) { @@ -70,13 +66,6 @@ private Long extractReasoningTokens(ChatResponseMetadata metadata) { return 0L; } - private Map defaultCounts() { - Map counts = new EnumMap<>(TokenType.class); - counts.put(TokenType.PROMPT, 0L); - counts.put(TokenType.COMPLETION, 0L); - return counts; - } - private Map copyMetadata(ChatResponseMetadata metadata, Object nativeUsage) { Map metadataMap = new HashMap<>(); metadata.keySet().forEach(key -> metadataMap.put(key, metadata.get(key))); From 940d450d802ca2832e1ea6d63d485a10d5b5b342 Mon Sep 17 00:00:00 2001 From: HuitaePark Date: Mon, 20 Jul 2026 14:28:32 +0900 Subject: [PATCH 3/6] =?UTF-8?q?feat:=20TokenUsageDetails=EC=9D=98=20?= =?UTF-8?q?=EA=B2=80=EC=A6=9D=20=EC=B6=94=EA=B0=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../core/domain/TokenUsageDetails.java | 14 +++++- .../java/domain/TokenUsageDetailTest.java | 45 +++++++++++++++++++ 2 files changed, 58 insertions(+), 1 deletion(-) create mode 100644 token-pilot-core/src/test/java/domain/TokenUsageDetailTest.java diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java index 40d07d9..776ff68 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java @@ -4,4 +4,16 @@ public record TokenUsageDetails( long cachedInputTokens, long reasoningOutputTokens, long cachedOutputTokens -) {} +) { + public TokenUsageDetails{ + if (cachedInputTokens < 0) { + throw new IllegalArgumentException("cachedInputTokens must be non-negative"); + } + if (reasoningOutputTokens < 0) { + throw new IllegalArgumentException("reasoningOutputTokens must be non-negative"); + } + if (cachedOutputTokens < 0) { + throw new IllegalArgumentException("cachedOutputTokens must be non-negative"); + } + } +} diff --git a/token-pilot-core/src/test/java/domain/TokenUsageDetailTest.java b/token-pilot-core/src/test/java/domain/TokenUsageDetailTest.java new file mode 100644 index 0000000..8d2c709 --- /dev/null +++ b/token-pilot-core/src/test/java/domain/TokenUsageDetailTest.java @@ -0,0 +1,45 @@ +package domain; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import io.tokenpilot.core.domain.TokenUsageDetails; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +public class TokenUsageDetailTest { + + @Test + @DisplayName("cached input 토큰은 음수일 수 없다") + void shouldRejectNegativeCachedInputTokens() { + assertThatThrownBy(() -> + new TokenUsageDetails(-1, 0, 0) + ).isInstanceOf(IllegalArgumentException.class); + } + + @Test + @DisplayName("reasoning output 토큰은 음수일 수 없다") + void shouldRejectNegativeReasoningOutputTokens() { + assertThatThrownBy(() -> + new TokenUsageDetails(0, -1, 0) + ).isInstanceOf(IllegalArgumentException.class); + } + + @Test + @DisplayName("cached output 토큰은 음수일 수 없다") + void shouldRejectNegativeCachedOutputTokens() { + assertThatThrownBy(() -> + new TokenUsageDetails(0, 0, -1) + ).isInstanceOf(IllegalArgumentException.class); + } + + @Test + @DisplayName("모든 세부 토큰은 0일 수 있다") + void shouldAllowZeroDetailTokens() { + TokenUsageDetails details = new TokenUsageDetails(0, 0, 0); + + assertThat(details.cachedInputTokens()).isZero(); + assertThat(details.reasoningOutputTokens()).isZero(); + assertThat(details.cachedOutputTokens()).isZero(); + } +} From e5ac8fc1a46f75479d494d41ecfd385c7a36169b Mon Sep 17 00:00:00 2001 From: HuitaePark Date: Mon, 20 Jul 2026 14:31:00 +0900 Subject: [PATCH 4/6] =?UTF-8?q?feat:=20TokenUsage=EC=9D=98=20=ED=95=84?= =?UTF-8?q?=EB=93=9C=20=EA=B2=80=EC=A6=9D=20=EC=B6=94=EA=B0=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../io/tokenpilot/core/domain/TokenUsage.java | 18 ++++++++ .../java/{ => domain}/TokenUsageTest.java | 43 +++++++++++++++++++ 2 files changed, 61 insertions(+) rename token-pilot-core/src/test/java/{ => domain}/TokenUsageTest.java (53%) diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java index 28154a0..9d893ba 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java @@ -2,6 +2,7 @@ import java.util.Collections; import java.util.Map; +import java.util.Objects; /** * AI 모델 호출 시 발생하는 토큰 사용량 정보. @@ -20,6 +21,23 @@ public record TokenUsage( ) { public TokenUsage { metadata = (metadata != null) ? Collections.unmodifiableMap(metadata) : Map.of(); + + if(inputTokens<0){ + throw new IllegalArgumentException( + "inputTokens must be non-negative" + ); + } + + if (outputTokens < 0) { + throw new IllegalArgumentException( + "outputTokens must be non-negative" + ); + } + + Objects.requireNonNull( + details, + "details must not be null" + ); } /** diff --git a/token-pilot-core/src/test/java/TokenUsageTest.java b/token-pilot-core/src/test/java/domain/TokenUsageTest.java similarity index 53% rename from token-pilot-core/src/test/java/TokenUsageTest.java rename to token-pilot-core/src/test/java/domain/TokenUsageTest.java index 0887b14..5f074d8 100644 --- a/token-pilot-core/src/test/java/TokenUsageTest.java +++ b/token-pilot-core/src/test/java/domain/TokenUsageTest.java @@ -1,4 +1,7 @@ +package domain; + import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; @@ -43,4 +46,44 @@ void shouldNotAddCachedInputTokensToInputTotalAgain() { assertThat(usage.totalTokens()).isEqualTo(300); assertThat(usage.getCount(TokenType.CACHED_PROMPT)).isEqualTo(40); } + + @Test + @DisplayName("input 토큰은 음수일 수 없다") + void shouldRejectNegativeInputTokens() { + assertThatThrownBy(() -> + new TokenUsage( + -1, + 0, + new TokenUsageDetails(0, 0, 0), + Map.of() + ) + ) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("inputTokens must be non-negative"); + } + + @Test + @DisplayName("output 토큰은 음수일 수 없다") + void shouldRejectNegativeOutputTokens() { + assertThatThrownBy(() -> + new TokenUsage( + 0, + -1, + new TokenUsageDetails(0, 0, 0), + Map.of() + ) + ) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("outputTokens must be non-negative"); + } + + @Test + @DisplayName("토큰 세부 정보는 null일 수 없다") + void shouldRejectNullTokenUsageDetails() { + assertThatThrownBy(() -> + new TokenUsage(100, 200, null, Map.of()) + ) + .isInstanceOf(NullPointerException.class) + .hasMessage("details must not be null"); + } } From 6368e2d9c11cf368c60c4b3a039b6bcce7e58f59 Mon Sep 17 00:00:00 2001 From: HuitaePark Date: Mon, 20 Jul 2026 14:47:58 +0900 Subject: [PATCH 5/6] =?UTF-8?q?feat(TokenUsage):=20=EC=B4=9D=EB=9F=89?= =?UTF-8?q?=EA=B3=BC=20=EC=84=B8=EB=B6=80=EB=9F=89=20=EB=B6=88=EB=B3=80?= =?UTF-8?q?=EC=8B=9D=EC=9D=84=20=EC=99=84=EC=84=B1?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- AGENTS.md | 2 +- README.md | 4 +- .../io/tokenpilot/core/domain/TokenType.java | 13 +- .../io/tokenpilot/core/domain/TokenUsage.java | 52 +++- .../core/domain/TokenUsageDetails.java | 14 +- .../src/test/java/domain/TokenUsageTest.java | 89 ------ .../tokenpilot/core/domain/TokenTypeTest.java | 18 ++ .../core/domain/TokenUsageDetailsTest.java} | 5 +- .../core/domain/TokenUsageTest.java | 258 ++++++++++++++++++ .../internal/DefaultCostCalculatorTest.java | 48 ---- 10 files changed, 345 insertions(+), 158 deletions(-) delete mode 100644 token-pilot-core/src/test/java/domain/TokenUsageTest.java create mode 100644 token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenTypeTest.java rename token-pilot-core/src/test/java/{domain/TokenUsageDetailTest.java => io/tokenpilot/core/domain/TokenUsageDetailsTest.java} (93%) create mode 100644 token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageTest.java diff --git a/AGENTS.md b/AGENTS.md index 3110195..7f98c77 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -317,7 +317,7 @@ The active checklist is in `docs/30_DAY_MVP_REPORT.md`; detailed long-term works ## Known Risks - `core.internal` implementation classes are package-private by design. Cross-module construction should continue through public factory/configuration APIs. -- Current usage totals and cached/reasoning breakdown semantics can double-count; fix the domain invariant before expanding pricing. +- `TokenUsage` now enforces inclusive input/output totals and validated cached/reasoning details. `DefaultCostCalculator` still iterates totals and details independently, so billable partitioning remains required before expanding pricing. - Current budget flow is check-then-add, is not an atomic reservation, and may not enforce `BLOCK` before provider invocation. - Current Micrometer `ai.token.*` metrics may duplicate Spring AI Observability; preserve compatibility while deciding default suppression or replacement. - Spring AI 2.0.0 documentation cannot be assumed to describe the current 1.1.4 runtime exactly; capability detection and supported-version tests are release requirements. diff --git a/README.md b/README.md index daaba51..a8c0db6 100644 --- a/README.md +++ b/README.md @@ -89,7 +89,7 @@ This configuration describes the current starter path. The 30-day MVP will exten | Capability | Status | MVP decision | | --- | --- | --- | -| Provider usage -> `TokenUsage` normalization | Basic implementation | Fix total/breakdown invariants first | +| Provider usage -> `TokenUsage` normalization | Inclusive total/breakdown invariants implemented | Partition billable totals and details in cost calculation | | `BigDecimal` cost calculation and pricing registry | Basic implementation | Define precision, rounding, and missing-price policy | | `LedgerManager` event publication | Basic implementation | Add idempotent estimate/actual reconciliation contract | | Spring AI `LedgerAdvisor` | Basic implementation | Keep as an adapter; enforce preflight decisions before the call | @@ -99,6 +99,8 @@ This configuration describes the current starter path. The 30-day MVP will exten | Exact byte-level BPE and JMH optimization | Post-MVP | Keep the estimator SPI so it can be added without API churn | | Provider routing, retry, fallback, gateway runtime | Future | Build only after accounting correctness is demonstrated | +`TokenUsage` now represents provider-reported input/output as inclusive totals and cached/reasoning values as validated details. Before the first `0.1.0` release, the former `Map` record component was replaced by `inputTokens`, `outputTokens`, `TokenUsageDetails`, and `metadata`; consumers of the pre-release source API must migrate constructor calls accordingly. + ## Modules | Module | Purpose | Status | diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java index 6f175af..3a94dcc 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java @@ -17,6 +17,8 @@ public enum TokenType { /** * 해당 토큰 타입이 입력(Prompt) 계열인지 확인합니다. + * + * @return 입력 계열이면 {@code true} */ public boolean isPrompt() { return this == PROMPT || this == CACHED_PROMPT; @@ -24,15 +26,10 @@ public boolean isPrompt() { /** * 해당 토큰 타입이 출력(Completion) 계열인지 확인합니다. + * + * @return 출력 계열이면 {@code true} */ public boolean isCompletion() { - return this == COMPLETION || this == CACHED_COMPLETION; - } - - /** - * 해당 토큰 타입이 추론(Reasoning) 계열인지 확인합니다. - */ - public boolean isReasoning(){ - return this == REASONING; + return this == COMPLETION || this == REASONING || this == CACHED_COMPLETION; } } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java index 9d893ba..3744834 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java @@ -1,6 +1,5 @@ package io.tokenpilot.core.domain; -import java.util.Collections; import java.util.Map; import java.util.Objects; @@ -19,10 +18,14 @@ public record TokenUsage( TokenUsageDetails details, Map metadata ) { + /** + * 토큰 총량과 세부량의 포함 관계를 검증하고 metadata를 불변 복사합니다. + * + * @throws IllegalArgumentException 토큰 수가 음수이거나 세부량이 총량을 초과한 경우 + * @throws NullPointerException details가 null이거나 metadata에 null key/value가 있는 경우 + */ public TokenUsage { - metadata = (metadata != null) ? Collections.unmodifiableMap(metadata) : Map.of(); - - if(inputTokens<0){ + if (inputTokens < 0) { throw new IllegalArgumentException( "inputTokens must be non-negative" ); @@ -34,14 +37,34 @@ public record TokenUsage( ); } - Objects.requireNonNull( + details = Objects.requireNonNull( details, "details must not be null" ); + + if (details.cachedInputTokens() > inputTokens) { + throw new IllegalArgumentException( + "cachedInputTokens must not exceed inputTokens" + ); + } + + if (details.cachedOutputTokens() > outputTokens + || details.reasoningOutputTokens() + > outputTokens - details.cachedOutputTokens()) { + throw new IllegalArgumentException( + "Output details must not exceed outputTokens" + ); + } + + metadata = metadata == null ? Map.of() : Map.copyOf(metadata); } /** * 기본 입력/출력 토큰을 사용하는 {@link TokenUsage}를 생성합니다. + * + * @param prompt 전체 입력 토큰 + * @param completion 전체 출력 토큰 + * @return 세부량과 metadata가 비어 있는 사용량 */ public static TokenUsage from(long prompt, long completion) { return new TokenUsage( @@ -54,6 +77,11 @@ public static TokenUsage from(long prompt, long completion) { /** * 입력/출력/추론 토큰을 포함하는 {@link TokenUsage}를 생성합니다. + * + * @param prompt 전체 입력 토큰 + * @param completion 전체 출력 토큰 + * @param reasoning 전체 출력에 포함된 reasoning 토큰 + * @return reasoning 세부량을 포함하는 사용량 */ public static TokenUsage from(long prompt, long completion, long reasoning) { return new TokenUsage( @@ -66,6 +94,8 @@ public static TokenUsage from(long prompt, long completion, long reasoning) { /** * 모든 종류의 입력/출력 토큰 수의 합계를 반환합니다. + * + * @return 전체 입력 토큰 */ public long promptTokens() { return inputTokens; @@ -73,6 +103,8 @@ public long promptTokens() { /** * 모든 종류의 출력(추론 포함) 토큰 수의 합계를 반환합니다. + * + * @return 전체 출력 토큰 */ public long completionTokens() { return outputTokens; @@ -80,13 +112,19 @@ public long completionTokens() { /** * 전체 사용 토큰 수의 합계를 반환합니다. + * + * @return 전체 입력과 출력 토큰의 합 + * @throws ArithmeticException 합계가 {@code long} 범위를 초과한 경우 */ public long totalTokens() { - return inputTokens + outputTokens; + return Math.addExact(inputTokens, outputTokens); } /** - * 특정 토큰 타입의 사용량을 가져옵니다. 없을 시 0을 반환합니다. + * 특정 토큰 타입의 전체 또는 세부 사용량을 반환합니다. + * + * @param type 조회할 토큰 타입 + * @return 해당 토큰 타입의 사용량 */ public long getCount(TokenType type) { return switch (type) { diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java index 776ff68..ad49016 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java @@ -1,11 +1,23 @@ package io.tokenpilot.core.domain; +/** + * 전체 입력/출력 토큰에 포함되는 세부 토큰 사용량. + * + * @param cachedInputTokens 전체 입력에 포함된 cached input 토큰 + * @param reasoningOutputTokens 전체 출력에 포함된 reasoning 토큰 + * @param cachedOutputTokens 전체 출력에 포함된 cached output 토큰 + */ public record TokenUsageDetails( long cachedInputTokens, long reasoningOutputTokens, long cachedOutputTokens ) { - public TokenUsageDetails{ + /** + * 모든 세부 토큰 수가 0 이상인지 검증합니다. + * + * @throws IllegalArgumentException 세부 토큰 수가 음수인 경우 + */ + public TokenUsageDetails { if (cachedInputTokens < 0) { throw new IllegalArgumentException("cachedInputTokens must be non-negative"); } diff --git a/token-pilot-core/src/test/java/domain/TokenUsageTest.java b/token-pilot-core/src/test/java/domain/TokenUsageTest.java deleted file mode 100644 index 5f074d8..0000000 --- a/token-pilot-core/src/test/java/domain/TokenUsageTest.java +++ /dev/null @@ -1,89 +0,0 @@ -package domain; - -import static org.assertj.core.api.Assertions.assertThat; -import static org.assertj.core.api.Assertions.assertThatThrownBy; - -import io.tokenpilot.core.domain.TokenType; -import io.tokenpilot.core.domain.TokenUsage; -import io.tokenpilot.core.domain.TokenUsageDetails; -import java.util.Map; -import org.junit.jupiter.api.DisplayName; -import org.junit.jupiter.api.Test; - -public class TokenUsageTest { - - @Test - @DisplayName("reasoning 토큰은 output total의 부분집합이어야 한다") - void shouldNotAddReasoningTokensToOutputTotalAgain() { - TokenUsage usage = TokenUsage.from( - 100, - 200, - 150 - ); - - assertThat(usage.promptTokens()).isEqualTo(100); - assertThat(usage.completionTokens()).isEqualTo(200); - assertThat(usage.totalTokens()).isEqualTo(300); - assertThat(usage.getCount(TokenType.REASONING)).isEqualTo(150); - } - - @Test - @DisplayName("cache의 input이 중복되면 안된다.") - void shouldNotAddCachedInputTokensToInputTotalAgain() { - TokenUsage usage = new TokenUsage( - 100, - 200, - new TokenUsageDetails( - 40, // cachedInputTokens - 0, // reasoningOutputTokens - 0 // cachedOutputTokens - ), - Map.of() - ); - - assertThat(usage.promptTokens()).isEqualTo(100); - assertThat(usage.completionTokens()).isEqualTo(200); - assertThat(usage.totalTokens()).isEqualTo(300); - assertThat(usage.getCount(TokenType.CACHED_PROMPT)).isEqualTo(40); - } - - @Test - @DisplayName("input 토큰은 음수일 수 없다") - void shouldRejectNegativeInputTokens() { - assertThatThrownBy(() -> - new TokenUsage( - -1, - 0, - new TokenUsageDetails(0, 0, 0), - Map.of() - ) - ) - .isInstanceOf(IllegalArgumentException.class) - .hasMessage("inputTokens must be non-negative"); - } - - @Test - @DisplayName("output 토큰은 음수일 수 없다") - void shouldRejectNegativeOutputTokens() { - assertThatThrownBy(() -> - new TokenUsage( - 0, - -1, - new TokenUsageDetails(0, 0, 0), - Map.of() - ) - ) - .isInstanceOf(IllegalArgumentException.class) - .hasMessage("outputTokens must be non-negative"); - } - - @Test - @DisplayName("토큰 세부 정보는 null일 수 없다") - void shouldRejectNullTokenUsageDetails() { - assertThatThrownBy(() -> - new TokenUsage(100, 200, null, Map.of()) - ) - .isInstanceOf(NullPointerException.class) - .hasMessage("details must not be null"); - } -} diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenTypeTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenTypeTest.java new file mode 100644 index 0000000..f0011e8 --- /dev/null +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenTypeTest.java @@ -0,0 +1,18 @@ +package io.tokenpilot.core.domain; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +class TokenTypeTest { + + @Test + @DisplayName("completion 계열에는 일반·추론·cached output이 포함된다") + void shouldClassifyAllCompletionTokenTypes() { + assertThat(TokenType.COMPLETION.isCompletion()).isTrue(); + assertThat(TokenType.REASONING.isCompletion()).isTrue(); + assertThat(TokenType.CACHED_COMPLETION.isCompletion()).isTrue(); + assertThat(TokenType.PROMPT.isCompletion()).isFalse(); + } +} diff --git a/token-pilot-core/src/test/java/domain/TokenUsageDetailTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageDetailsTest.java similarity index 93% rename from token-pilot-core/src/test/java/domain/TokenUsageDetailTest.java rename to token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageDetailsTest.java index 8d2c709..22f1d53 100644 --- a/token-pilot-core/src/test/java/domain/TokenUsageDetailTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageDetailsTest.java @@ -1,13 +1,12 @@ -package domain; +package io.tokenpilot.core.domain; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; -import io.tokenpilot.core.domain.TokenUsageDetails; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; -public class TokenUsageDetailTest { +class TokenUsageDetailsTest { @Test @DisplayName("cached input 토큰은 음수일 수 없다") diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageTest.java new file mode 100644 index 0000000..ed601c6 --- /dev/null +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageTest.java @@ -0,0 +1,258 @@ +package io.tokenpilot.core.domain; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.util.HashMap; +import java.util.Map; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +class TokenUsageTest { + + @Test + @DisplayName("세부량이 없으면 입력과 출력 총량만 사용한다") + void shouldUseInputAndOutputTotalsWithoutDetails() { + TokenUsage usage = TokenUsage.from(100, 200); + + assertThat(usage.promptTokens()).isEqualTo(100); + assertThat(usage.completionTokens()).isEqualTo(200); + assertThat(usage.totalTokens()).isEqualTo(300); + assertThat(usage.details()).isEqualTo(new TokenUsageDetails(0, 0, 0)); + } + + @Test + @DisplayName("입력과 출력 토큰은 모두 0일 수 있다") + void shouldAllowZeroInputAndOutputTokens() { + TokenUsage usage = TokenUsage.from(0, 0); + + assertThat(usage.totalTokens()).isZero(); + } + + @Test + @DisplayName("reasoning 토큰은 output total의 부분집합이어야 한다") + void shouldNotAddReasoningTokensToOutputTotalAgain() { + TokenUsage usage = TokenUsage.from( + 100, + 200, + 150 + ); + + assertThat(usage.promptTokens()).isEqualTo(100); + assertThat(usage.completionTokens()).isEqualTo(200); + assertThat(usage.totalTokens()).isEqualTo(300); + assertThat(usage.getCount(TokenType.REASONING)).isEqualTo(150); + } + + @Test + @DisplayName("cached input은 전체 input에 중복 합산되지 않는다") + void shouldNotAddCachedInputTokensToInputTotalAgain() { + TokenUsage usage = new TokenUsage( + 100, + 200, + new TokenUsageDetails( + 40, // cachedInputTokens + 0, // reasoningOutputTokens + 0 // cachedOutputTokens + ), + Map.of() + ); + + assertThat(usage.promptTokens()).isEqualTo(100); + assertThat(usage.completionTokens()).isEqualTo(200); + assertThat(usage.totalTokens()).isEqualTo(300); + assertThat(usage.getCount(TokenType.CACHED_PROMPT)).isEqualTo(40); + } + + @Test + @DisplayName("input 토큰은 음수일 수 없다") + void shouldRejectNegativeInputTokens() { + assertThatThrownBy(() -> + new TokenUsage( + -1, + 0, + new TokenUsageDetails(0, 0, 0), + Map.of() + ) + ) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("inputTokens must be non-negative"); + } + + @Test + @DisplayName("output 토큰은 음수일 수 없다") + void shouldRejectNegativeOutputTokens() { + assertThatThrownBy(() -> + new TokenUsage( + 0, + -1, + new TokenUsageDetails(0, 0, 0), + Map.of() + ) + ) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("outputTokens must be non-negative"); + } + + @Test + @DisplayName("토큰 세부 정보는 null일 수 없다") + void shouldRejectNullTokenUsageDetails() { + assertThatThrownBy(() -> + new TokenUsage(100, 200, null, Map.of()) + ) + .isInstanceOf(NullPointerException.class) + .hasMessage("details must not be null"); + } + + @Test + @DisplayName("cached input은 전체 input보다 클 수 없다") + void shouldRejectCachedInputTokensGreaterThanInputTokens() { + assertThatThrownBy(() -> + new TokenUsage( + 100, + 200, + new TokenUsageDetails(101, 0, 0), + Map.of() + ) + ) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage( + "cachedInputTokens must not exceed inputTokens" + ); + } + + @Test + @DisplayName("cached output은 전체 output보다 클 수 없다") + void shouldRejectCachedOutputTokensGreaterThanOutputTokens() { + assertThatThrownBy(() -> + new TokenUsage( + 100, + 200, + new TokenUsageDetails(0, 0, 201), + Map.of() + ) + ) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Output details must not exceed outputTokens"); + } + + @Test + @DisplayName("reasoning과 cached output의 합은 전체 output보다 클 수 없다") + void shouldRejectOutputDetailsGreaterThanOutputTokens() { + assertThatThrownBy(() -> + new TokenUsage( + 100, + 200, + new TokenUsageDetails(0, 151, 50), + Map.of() + ) + ) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Output details must not exceed outputTokens"); + } + + @Test + @DisplayName("reasoning output은 전체 output보다 클 수 없다") + void shouldRejectReasoningOutputTokensGreaterThanOutputTokens() { + assertThatThrownBy(() -> + new TokenUsage( + 100, + 200, + new TokenUsageDetails(0, 201, 0), + Map.of() + ) + ) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Output details must not exceed outputTokens"); + } + + @Test + @DisplayName("metadata는 원본 Map의 변경에 영향을 받지 않는다") + void shouldDefensivelyCopyMetadata() { + Map metadata = new HashMap<>(); + metadata.put("model", "gpt-4o"); + TokenUsage usage = new TokenUsage( + 100, + 200, + new TokenUsageDetails(0, 0, 0), + metadata + ); + + metadata.put("model", "changed"); + + assertThat(usage.metadata()).containsEntry("model", "gpt-4o"); + } + + @Test + @DisplayName("metadata가 null이면 빈 Map을 사용한다") + void shouldUseEmptyMetadataWhenMetadataIsNull() { + TokenUsage usage = new TokenUsage( + 100, + 200, + new TokenUsageDetails(0, 0, 0), + null + ); + + assertThat(usage.metadata()).isEmpty(); + } + + @Test + @DisplayName("metadata의 key는 null일 수 없다") + void shouldRejectMetadataWithNullKey() { + Map metadata = new HashMap<>(); + metadata.put(null, "value"); + + assertThatThrownBy(() -> + new TokenUsage( + 100, + 200, + new TokenUsageDetails(0, 0, 0), + metadata + ) + ).isInstanceOf(NullPointerException.class); + } + + @Test + @DisplayName("metadata의 value는 null일 수 없다") + void shouldRejectMetadataWithNullValue() { + Map metadata = new HashMap<>(); + metadata.put("key", null); + + assertThatThrownBy(() -> + new TokenUsage( + 100, + 200, + new TokenUsageDetails(0, 0, 0), + metadata + ) + ).isInstanceOf(NullPointerException.class); + } + + @Test + @DisplayName("breakdown 합계는 전체 토큰과 같을 수 있다") + void shouldAllowDetailsEqualToTotalTokens() { + TokenUsage usage = new TokenUsage( + 100, + 200, + new TokenUsageDetails(100, 150, 50), + Map.of() + ); + + assertThat(usage.inputTokens()).isEqualTo(100); + assertThat(usage.outputTokens()).isEqualTo(200); + } + + @Test + @DisplayName("전체 토큰 합계가 long 범위를 넘으면 예외가 발생한다") + void shouldRejectTotalTokenOverflow() { + TokenUsage usage = new TokenUsage( + Long.MAX_VALUE, + 1, + new TokenUsageDetails(0, 0, 0), + Map.of() + ); + + assertThatThrownBy(usage::totalTokens) + .isInstanceOf(ArithmeticException.class); + } +} diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java index 86cf088..7296939 100644 --- a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java @@ -2,15 +2,11 @@ import io.tokenpilot.core.domain.Cost; import io.tokenpilot.core.domain.PricingPlan; -import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; -import io.tokenpilot.core.domain.TokenUsageDetails; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; import java.math.BigDecimal; -import java.util.EnumMap; -import java.util.Map; import static org.assertj.core.api.Assertions.assertThat; @@ -32,48 +28,4 @@ void calculateStandardTokens() { // Then: (1000 * 0.01 / 1000) + (2000 * 0.03 / 1000) = 0.01 + 0.06 = 0.07 assertThat(cost.value()).isEqualByComparingTo("0.070000"); } - - @Test - @DisplayName("추론 토큰 요율이 다를 때 각각 정밀하게 계산되어야 한다") - void calculateReasoningTokensWithDifferentRate() { - // Given: Prompt $0.015, Completion $0.060, Reasoning $0.045 - Map rates = new EnumMap<>(TokenType.class); - rates.put(TokenType.PROMPT, new BigDecimal("0.015")); - rates.put(TokenType.COMPLETION, new BigDecimal("0.060")); - rates.put(TokenType.REASONING, new BigDecimal("0.045")); - PricingPlan plan = new PricingPlan("o1-preview", rates, Cost.DEFAULT_CURRENCY); - - // Usage: Prompt 1000, Completion 500, Reasoning 1500 (Total 3000) - TokenUsage usage = new TokenUsage( - 1000, - 500, - new TokenUsageDetails(0, 1500, 0), - Map.of() - ); - - // When - Cost cost = calculator.calculate(usage, plan); - - // Then - // 1000 * 0.015 / 1000 = 0.015 - // 500 * 0.060 / 1000 = 0.030 - // 1500 * 0.045 / 1000 = 0.0675 - // Total = 0.015 + 0.030 + 0.0675 = 0.1125 - assertThat(cost.value()).isEqualByComparingTo("0.112500"); - } - - @Test - @DisplayName("추론 토큰 요율이 명시되지 않으면 일반 출력 요율을 따라야 한다") - void calculateReasoningTokensWithFallback() { - // Given: Prompt $0.01, Completion $0.03 (Reasoning 없음) - PricingPlan plan = new PricingPlan("gpt-4o", new BigDecimal("0.01"), new BigDecimal("0.03")); - // Usage: Prompt 1000, Completion 0, Reasoning 1000 - TokenUsage usage = TokenUsage.from(1000, 0, 1000); - - // When - Cost cost = calculator.calculate(usage, plan); - - // Then: 1000 * 0.01 / 1000 + 1000 * 0.03 / 1000 = 0.04 - assertThat(cost.value()).isEqualByComparingTo("0.040000"); - } } From 30e9a496887af2c73fda3570c4bc07728b2160a8 Mon Sep 17 00:00:00 2001 From: HuitaePark Date: Mon, 20 Jul 2026 15:10:37 +0900 Subject: [PATCH 6/6] =?UTF-8?q?refactor(TokenUsage):=20provider=20breakdow?= =?UTF-8?q?n=20=EB=AA=A8=EB=8D=B8=EC=9D=84=20=EC=A0=95=EA=B5=90=ED=99=94?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- AGENTS.md | 5 +- README.md | 8 +- .../tokenpilot/core/domain/PricingPlan.java | 8 +- .../io/tokenpilot/core/domain/TokenType.java | 12 +- .../io/tokenpilot/core/domain/TokenUsage.java | 56 ++++- .../core/domain/TokenUsageDetails.java | 37 ++- .../tokenpilot/core/domain/UsageSource.java | 17 ++ .../core/internal/DefaultCostCalculator.java | 41 ++-- .../tokenpilot/core/domain/TokenTypeTest.java | 14 +- .../core/domain/TokenUsageDetailsTest.java | 36 +-- .../core/domain/TokenUsageTest.java | 229 ++++++++---------- .../internal/DefaultCostCalculatorTest.java | 34 +++ .../internal/DefaultUsageExtractor.java | 131 ++++++++-- .../internal/DefaultUsageExtractorTest.java | 91 +++++++ 14 files changed, 505 insertions(+), 214 deletions(-) create mode 100644 token-pilot-core/src/main/java/io/tokenpilot/core/domain/UsageSource.java diff --git a/AGENTS.md b/AGENTS.md index 7f98c77..3503a64 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -317,7 +317,8 @@ The active checklist is in `docs/30_DAY_MVP_REPORT.md`; detailed long-term works ## Known Risks - `core.internal` implementation classes are package-private by design. Cross-module construction should continue through public factory/configuration APIs. -- `TokenUsage` now enforces inclusive input/output totals and validated cached/reasoning details. `DefaultCostCalculator` still iterates totals and details independently, so billable partitioning remains required before expanding pricing. +- `TokenUsage` now enforces normalized inclusive totals, optional cache-read/cache-creation/reasoning details, and explicit usage provenance. `DefaultCostCalculator` partitions overlapping totals into disjoint billable amounts before applying rates. +- Spring AI usage extraction converts map/JSON-compatible native usage objects into the normalized core model. Real-provider compatibility fixtures remain required because provider and Spring AI usage shapes can change independently. - Current budget flow is check-then-add, is not an atomic reservation, and may not enforce `BLOCK` before provider invocation. - Current Micrometer `ai.token.*` metrics may duplicate Spring AI Observability; preserve compatibility while deciding default suppression or replacement. - Spring AI 2.0.0 documentation cannot be assumed to describe the current 1.1.4 runtime exactly; capability detection and supported-version tests are release requirements. @@ -382,6 +383,8 @@ Stage and deploy a Central release: ### 2026-07-20 +- Replaced generic cached input/cached output details with optional cache-read input, cache-creation input, and reasoning-output breakdowns; `null` now means unreported and `0` means reported zero. +- Added `UsageSource`, provider-specific total normalization in the Spring AI adapter, and disjoint cost partitioning to prevent details from being charged twice. - Positioned Token Pilot as a framework-independent Java LLM control and accounting core with Spring AI as an optional adapter. - Chose to reuse Spring AI Observability for standard latency, trace, and token telemetry while keeping Token Pilot metrics focused on cost, policy, budget, and reconciliation. - Kept the existing starter as a Spring AI convenience distribution rather than the sole product identity. diff --git a/README.md b/README.md index a8c0db6..42460d0 100644 --- a/README.md +++ b/README.md @@ -89,8 +89,8 @@ This configuration describes the current starter path. The 30-day MVP will exten | Capability | Status | MVP decision | | --- | --- | --- | -| Provider usage -> `TokenUsage` normalization | Inclusive total/breakdown invariants implemented | Partition billable totals and details in cost calculation | -| `BigDecimal` cost calculation and pricing registry | Basic implementation | Define precision, rounding, and missing-price policy | +| Provider usage -> `TokenUsage` normalization | Inclusive totals, optional cache/reasoning breakdown, and source implemented | Add supported-provider compatibility fixtures | +| `BigDecimal` cost calculation and pricing registry | Disjoint billable partition implemented | Define precision, rounding, and missing-price policy | | `LedgerManager` event publication | Basic implementation | Add idempotent estimate/actual reconciliation contract | | Spring AI `LedgerAdvisor` | Basic implementation | Keep as an adapter; enforce preflight decisions before the call | | Micrometer metrics | Basic implementation | Avoid duplicating Spring AI token/latency telemetry by default | @@ -99,7 +99,9 @@ This configuration describes the current starter path. The 30-day MVP will exten | Exact byte-level BPE and JMH optimization | Post-MVP | Keep the estimator SPI so it can be added without API churn | | Provider routing, retry, fallback, gateway runtime | Future | Build only after accounting correctness is demonstrated | -`TokenUsage` now represents provider-reported input/output as inclusive totals and cached/reasoning values as validated details. Before the first `0.1.0` release, the former `Map` record component was replaced by `inputTokens`, `outputTokens`, `TokenUsageDetails`, and `metadata`; consumers of the pre-release source API must migrate constructor calls accordingly. +`TokenUsage` represents normalized input/output as inclusive totals. `TokenUsageDetails` stores optional cache-read input, cache-creation input, and reasoning-output breakdowns; `null` means the provider did not report the value, while `0` means it reported zero usage. `UsageSource` distinguishes reported, provider-derived, locally calculated, estimated, and unavailable values. Cost calculation partitions these overlapping totals into disjoint billable amounts before applying rates, so details are never charged twice. + +Before the first `0.1.0` release, the former `Map` record component was replaced by `inputTokens`, `outputTokens`, `TokenUsageDetails`, `UsageSource`, and `metadata`. `CACHED_PROMPT`/`CACHED_COMPLETION` were replaced by `CACHE_READ_PROMPT`/`CACHE_CREATION_PROMPT`; consumers of the pre-release source API must migrate constructor calls and pricing keys accordingly. ## Modules diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingPlan.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingPlan.java index 7da178d..fb962ca 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingPlan.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingPlan.java @@ -70,8 +70,7 @@ public BigDecimal completionPricePerK() { /** * 특정 토큰 타입의 단가를 가져옵니다. 없을 시 계층 구조에 따라 대체값을 반환합니다. * REASONING -> COMPLETION - * CACHED_PROMPT -> PROMPT - * CACHED_COMPLETION -> COMPLETION + * CACHE_READ_PROMPT, CACHE_CREATION_PROMPT -> PROMPT */ public BigDecimal getRate(TokenType type) { if (rates.containsKey(type)) { @@ -80,8 +79,9 @@ public BigDecimal getRate(TokenType type) { // Fallback Logic return switch (type) { - case REASONING, CACHED_COMPLETION -> rates.getOrDefault(TokenType.COMPLETION, BigDecimal.ZERO); - case CACHED_PROMPT -> rates.getOrDefault(TokenType.PROMPT, BigDecimal.ZERO); + case REASONING -> rates.getOrDefault(TokenType.COMPLETION, BigDecimal.ZERO); + case CACHE_READ_PROMPT, CACHE_CREATION_PROMPT -> + rates.getOrDefault(TokenType.PROMPT, BigDecimal.ZERO); default -> rates.getOrDefault(type, BigDecimal.ZERO); }; } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java index 3a94dcc..94ddd9b 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenType.java @@ -10,10 +10,10 @@ public enum TokenType { COMPLETION, /** 추론 (Reasoning) - 주로 출력 계열 */ REASONING, - /** 캐시된 입력 (Cached Prompt) */ - CACHED_PROMPT, - /** 캐시된 출력 (Cached Completion) */ - CACHED_COMPLETION; + /** 캐시에서 읽은 입력 (Cache Read Prompt) */ + CACHE_READ_PROMPT, + /** 캐시에 새로 저장한 입력 (Cache Creation Prompt) */ + CACHE_CREATION_PROMPT; /** * 해당 토큰 타입이 입력(Prompt) 계열인지 확인합니다. @@ -21,7 +21,7 @@ public enum TokenType { * @return 입력 계열이면 {@code true} */ public boolean isPrompt() { - return this == PROMPT || this == CACHED_PROMPT; + return this == PROMPT || this == CACHE_READ_PROMPT || this == CACHE_CREATION_PROMPT; } /** @@ -30,6 +30,6 @@ public boolean isPrompt() { * @return 출력 계열이면 {@code true} */ public boolean isCompletion() { - return this == COMPLETION || this == REASONING || this == CACHED_COMPLETION; + return this == COMPLETION || this == REASONING; } } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java index 3744834..6c82eaf 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsage.java @@ -10,19 +10,21 @@ * @param inputTokens 전체 입력 토큰 * @param outputTokens 전체 출력 토큰 * @param details 세부 토큰 사용량 + * @param source 사용량 값의 출처 * @param metadata 추가 메타데이터 (예: 모델 정보 등) */ public record TokenUsage( long inputTokens, long outputTokens, TokenUsageDetails details, + UsageSource source, Map metadata ) { /** * 토큰 총량과 세부량의 포함 관계를 검증하고 metadata를 불변 복사합니다. * * @throws IllegalArgumentException 토큰 수가 음수이거나 세부량이 총량을 초과한 경우 - * @throws NullPointerException details가 null이거나 metadata에 null key/value가 있는 경우 + * @throws NullPointerException details/source가 null이거나 metadata에 null key/value가 있는 경우 */ public TokenUsage { if (inputTokens < 0) { @@ -41,18 +43,24 @@ public record TokenUsage( details, "details must not be null" ); + source = Objects.requireNonNull( + source, + "source must not be null" + ); - if (details.cachedInputTokens() > inputTokens) { + long cacheRead = countOrZero(details.cacheReadInputTokens()); + long cacheCreation = countOrZero(details.cacheCreationInputTokens()); + if (cacheRead > inputTokens + || cacheCreation > inputTokens - cacheRead) { throw new IllegalArgumentException( - "cachedInputTokens must not exceed inputTokens" + "Input details must not exceed inputTokens" ); } - if (details.cachedOutputTokens() > outputTokens - || details.reasoningOutputTokens() - > outputTokens - details.cachedOutputTokens()) { + long reasoning = countOrZero(details.reasoningOutputTokens()); + if (reasoning > outputTokens) { throw new IllegalArgumentException( - "Output details must not exceed outputTokens" + "reasoningOutputTokens must not exceed outputTokens" ); } @@ -70,7 +78,8 @@ public static TokenUsage from(long prompt, long completion) { return new TokenUsage( prompt, completion, - new TokenUsageDetails(0, 0, 0), + TokenUsageDetails.unreported(), + UsageSource.PROVIDER_REPORTED, Map.of() ); } @@ -87,11 +96,28 @@ public static TokenUsage from(long prompt, long completion, long reasoning) { return new TokenUsage( prompt, completion, - new TokenUsageDetails(0, reasoning, 0), + new TokenUsageDetails(null, null, reasoning), + UsageSource.PROVIDER_REPORTED, Map.of() ); } + /** + * provider 응답에서 사용량 정보를 얻지 못한 상태를 생성합니다. + * + * @param metadata 보존할 응답 메타데이터 + * @return 출처가 {@link UsageSource#UNAVAILABLE}인 0 토큰 사용량 + */ + public static TokenUsage unavailable(Map metadata) { + return new TokenUsage( + 0, + 0, + TokenUsageDetails.unreported(), + UsageSource.UNAVAILABLE, + metadata + ); + } + /** * 모든 종류의 입력/출력 토큰 수의 합계를 반환합니다. * @@ -122,6 +148,8 @@ public long totalTokens() { /** * 특정 토큰 타입의 전체 또는 세부 사용량을 반환합니다. + * 보고되지 않은 세부량은 이 호환 projection에서 {@code 0}으로 반환되며, + * 미보고 여부는 {@link #details()}의 nullable 필드로 확인해야 합니다. * * @param type 조회할 토큰 타입 * @return 해당 토큰 타입의 사용량 @@ -130,9 +158,13 @@ public long getCount(TokenType type) { return switch (type) { case PROMPT -> inputTokens; case COMPLETION -> outputTokens; - case REASONING -> details.reasoningOutputTokens(); - case CACHED_PROMPT -> details.cachedInputTokens(); - case CACHED_COMPLETION -> details.cachedOutputTokens(); + case REASONING -> countOrZero(details.reasoningOutputTokens()); + case CACHE_READ_PROMPT -> countOrZero(details.cacheReadInputTokens()); + case CACHE_CREATION_PROMPT -> countOrZero(details.cacheCreationInputTokens()); }; } + + private static long countOrZero(Long count) { + return count == null ? 0L : count; + } } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java index ad49016..58bac48 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/TokenUsageDetails.java @@ -1,16 +1,18 @@ package io.tokenpilot.core.domain; /** - * 전체 입력/출력 토큰에 포함되는 세부 토큰 사용량. + * 정규화된 전체 입력/출력 토큰에 포함되는 선택적 세부 사용량. + * {@code null}은 provider가 값을 보고하지 않았음을, {@code 0}은 값을 + * 보고했으나 실제 사용량이 없었음을 뜻합니다. * - * @param cachedInputTokens 전체 입력에 포함된 cached input 토큰 - * @param reasoningOutputTokens 전체 출력에 포함된 reasoning 토큰 - * @param cachedOutputTokens 전체 출력에 포함된 cached output 토큰 + * @param cacheReadInputTokens 전체 입력에 포함된 cache read 토큰 또는 미보고 시 {@code null} + * @param cacheCreationInputTokens 전체 입력에 포함된 cache creation 토큰 또는 미보고 시 {@code null} + * @param reasoningOutputTokens 전체 출력에 포함된 reasoning 토큰 또는 미보고 시 {@code null} */ public record TokenUsageDetails( - long cachedInputTokens, - long reasoningOutputTokens, - long cachedOutputTokens + Long cacheReadInputTokens, + Long cacheCreationInputTokens, + Long reasoningOutputTokens ) { /** * 모든 세부 토큰 수가 0 이상인지 검증합니다. @@ -18,14 +20,23 @@ public record TokenUsageDetails( * @throws IllegalArgumentException 세부 토큰 수가 음수인 경우 */ public TokenUsageDetails { - if (cachedInputTokens < 0) { - throw new IllegalArgumentException("cachedInputTokens must be non-negative"); + if (cacheReadInputTokens != null && cacheReadInputTokens < 0) { + throw new IllegalArgumentException("cacheReadInputTokens must be non-negative"); } - if (reasoningOutputTokens < 0) { - throw new IllegalArgumentException("reasoningOutputTokens must be non-negative"); + if (cacheCreationInputTokens != null && cacheCreationInputTokens < 0) { + throw new IllegalArgumentException("cacheCreationInputTokens must be non-negative"); } - if (cachedOutputTokens < 0) { - throw new IllegalArgumentException("cachedOutputTokens must be non-negative"); + if (reasoningOutputTokens != null && reasoningOutputTokens < 0) { + throw new IllegalArgumentException("reasoningOutputTokens must be non-negative"); } } + + /** + * provider가 세부 사용량을 보고하지 않은 상태를 생성합니다. + * + * @return 모든 세부량이 미보고 상태인 객체 + */ + public static TokenUsageDetails unreported() { + return new TokenUsageDetails(null, null, null); + } } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/UsageSource.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/UsageSource.java new file mode 100644 index 0000000..b958e7b --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/UsageSource.java @@ -0,0 +1,17 @@ +package io.tokenpilot.core.domain; + +/** + * 토큰 사용량 값의 출처와 생성 방식을 나타냅니다. + */ +public enum UsageSource { + /** provider가 포괄 총량을 직접 보고한 사용량 */ + PROVIDER_REPORTED, + /** provider가 보고한 여러 필드를 어댑터가 정규화해 만든 사용량 */ + PROVIDER_DERIVED, + /** 로컬 tokenizer가 계산한 사용량 */ + LOCAL_TOKENIZER, + /** 휴리스틱으로 근사 추정한 사용량 */ + HEURISTIC_ESTIMATE, + /** 사용량 정보를 얻을 수 없음 */ + UNAVAILABLE +} diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultCostCalculator.java b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultCostCalculator.java index c6880dc..079daf8 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultCostCalculator.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultCostCalculator.java @@ -11,7 +11,8 @@ /** * 기본 비용 계산기 구현체. - * 각 {@link TokenType} 별 단가를 적용하여 정밀하게 계산합니다. + * 포괄 총량에서 cache/reasoning 세부량을 분리한 배타적 구간별로 + * {@link TokenType} 단가를 적용하여 중복 없이 계산합니다. * 1K 토큰당 가격 정보를 사용하여 소수점 10자리까지 중간 계산 후 6자리로 최종 반올림합니다. */ class DefaultCostCalculator implements CostCalculator { @@ -19,20 +20,32 @@ class DefaultCostCalculator implements CostCalculator { @Override public Cost calculate(TokenUsage usage, PricingPlan plan) { - BigDecimal totalCostValue = BigDecimal.ZERO; - - // 사용된 모든 토큰 타입에 대해 각각의 단가를 적용하여 합산 - for (TokenType type : TokenType.values()) { - long count = usage.getCount(type); - - if (count > 0) { - BigDecimal rate = plan.getRate(type); - BigDecimal typeCost = rate.multiply(BigDecimal.valueOf(count)) - .divide(THOUSAND, 10, RoundingMode.HALF_UP); - totalCostValue = totalCostValue.add(typeCost); - } - } + long cacheReadInput = countOrZero(usage.details().cacheReadInputTokens()); + long cacheCreationInput = countOrZero(usage.details().cacheCreationInputTokens()); + long reasoningOutput = countOrZero(usage.details().reasoningOutputTokens()); + + long regularInput = usage.inputTokens() - cacheReadInput - cacheCreationInput; + long regularOutput = usage.outputTokens() - reasoningOutput; + + BigDecimal totalCostValue = costFor(regularInput, plan.getRate(TokenType.PROMPT)) + .add(costFor(cacheReadInput, plan.getRate(TokenType.CACHE_READ_PROMPT))) + .add(costFor(cacheCreationInput, plan.getRate(TokenType.CACHE_CREATION_PROMPT))) + .add(costFor(regularOutput, plan.getRate(TokenType.COMPLETION))) + .add(costFor(reasoningOutput, plan.getRate(TokenType.REASONING))); return new Cost(totalCostValue, plan.currency()); } + + private BigDecimal costFor(long count, BigDecimal rate) { + if (count == 0) { + return BigDecimal.ZERO; + } + + return rate.multiply(BigDecimal.valueOf(count)) + .divide(THOUSAND, 10, RoundingMode.HALF_UP); + } + + private long countOrZero(Long count) { + return count == null ? 0L : count; + } } diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenTypeTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenTypeTest.java index f0011e8..dc13e22 100644 --- a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenTypeTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenTypeTest.java @@ -8,11 +8,19 @@ class TokenTypeTest { @Test - @DisplayName("completion 계열에는 일반·추론·cached output이 포함된다") - void shouldClassifyAllCompletionTokenTypes() { + @DisplayName("completion 계열에는 전체 출력과 reasoning breakdown이 포함된다") + void shouldClassifyCompletionTokenTypes() { assertThat(TokenType.COMPLETION.isCompletion()).isTrue(); assertThat(TokenType.REASONING.isCompletion()).isTrue(); - assertThat(TokenType.CACHED_COMPLETION.isCompletion()).isTrue(); assertThat(TokenType.PROMPT.isCompletion()).isFalse(); } + + @Test + @DisplayName("prompt 계열에는 전체 입력과 cache read/create breakdown이 포함된다") + void shouldClassifyPromptTokenTypes() { + assertThat(TokenType.PROMPT.isPrompt()).isTrue(); + assertThat(TokenType.CACHE_READ_PROMPT.isPrompt()).isTrue(); + assertThat(TokenType.CACHE_CREATION_PROMPT.isPrompt()).isTrue(); + assertThat(TokenType.COMPLETION.isPrompt()).isFalse(); + } } diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageDetailsTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageDetailsTest.java index 22f1d53..aa5a401 100644 --- a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageDetailsTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageDetailsTest.java @@ -9,36 +9,40 @@ class TokenUsageDetailsTest { @Test - @DisplayName("cached input 토큰은 음수일 수 없다") - void shouldRejectNegativeCachedInputTokens() { + @DisplayName("cache read input 토큰은 음수일 수 없다") + void shouldRejectNegativeCacheReadInputTokens() { assertThatThrownBy(() -> - new TokenUsageDetails(-1, 0, 0) + new TokenUsageDetails(-1L, 0L, 0L) ).isInstanceOf(IllegalArgumentException.class); } @Test - @DisplayName("reasoning output 토큰은 음수일 수 없다") - void shouldRejectNegativeReasoningOutputTokens() { + @DisplayName("cache creation input 토큰은 음수일 수 없다") + void shouldRejectNegativeCacheCreationInputTokens() { assertThatThrownBy(() -> - new TokenUsageDetails(0, -1, 0) + new TokenUsageDetails(0L, -1L, 0L) ).isInstanceOf(IllegalArgumentException.class); } @Test - @DisplayName("cached output 토큰은 음수일 수 없다") - void shouldRejectNegativeCachedOutputTokens() { + @DisplayName("reasoning output 토큰은 음수일 수 없다") + void shouldRejectNegativeReasoningOutputTokens() { assertThatThrownBy(() -> - new TokenUsageDetails(0, 0, -1) + new TokenUsageDetails(0L, 0L, -1L) ).isInstanceOf(IllegalArgumentException.class); } @Test - @DisplayName("모든 세부 토큰은 0일 수 있다") - void shouldAllowZeroDetailTokens() { - TokenUsageDetails details = new TokenUsageDetails(0, 0, 0); - - assertThat(details.cachedInputTokens()).isZero(); - assertThat(details.reasoningOutputTokens()).isZero(); - assertThat(details.cachedOutputTokens()).isZero(); + @DisplayName("미보고 세부량과 보고된 0을 구분한다") + void shouldDistinguishUnreportedDetailsFromReportedZero() { + TokenUsageDetails unreported = TokenUsageDetails.unreported(); + TokenUsageDetails reportedZero = new TokenUsageDetails(0L, 0L, 0L); + + assertThat(unreported.cacheReadInputTokens()).isNull(); + assertThat(unreported.cacheCreationInputTokens()).isNull(); + assertThat(unreported.reasoningOutputTokens()).isNull(); + assertThat(reportedZero.cacheReadInputTokens()).isZero(); + assertThat(reportedZero.cacheCreationInputTokens()).isZero(); + assertThat(reportedZero.reasoningOutputTokens()).isZero(); } } diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageTest.java index ed601c6..7eab1cf 100644 --- a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/TokenUsageTest.java @@ -11,70 +11,54 @@ class TokenUsageTest { @Test - @DisplayName("세부량이 없으면 입력과 출력 총량만 사용한다") - void shouldUseInputAndOutputTotalsWithoutDetails() { + @DisplayName("세부량이 보고되지 않으면 입력과 출력 총량만 사용한다") + void shouldUseInputAndOutputTotalsWithoutReportedDetails() { TokenUsage usage = TokenUsage.from(100, 200); - assertThat(usage.promptTokens()).isEqualTo(100); - assertThat(usage.completionTokens()).isEqualTo(200); + assertThat(usage.inputTokens()).isEqualTo(100); + assertThat(usage.outputTokens()).isEqualTo(200); assertThat(usage.totalTokens()).isEqualTo(300); - assertThat(usage.details()).isEqualTo(new TokenUsageDetails(0, 0, 0)); + assertThat(usage.details()).isEqualTo(TokenUsageDetails.unreported()); + assertThat(usage.source()).isEqualTo(UsageSource.PROVIDER_REPORTED); } @Test @DisplayName("입력과 출력 토큰은 모두 0일 수 있다") void shouldAllowZeroInputAndOutputTokens() { - TokenUsage usage = TokenUsage.from(0, 0); - - assertThat(usage.totalTokens()).isZero(); + assertThat(TokenUsage.from(0, 0).totalTokens()).isZero(); } @Test - @DisplayName("reasoning 토큰은 output total의 부분집합이어야 한다") + @DisplayName("reasoning 토큰은 output total에 중복 합산되지 않는다") void shouldNotAddReasoningTokensToOutputTotalAgain() { - TokenUsage usage = TokenUsage.from( - 100, - 200, - 150 - ); + TokenUsage usage = TokenUsage.from(100, 200, 150); - assertThat(usage.promptTokens()).isEqualTo(100); - assertThat(usage.completionTokens()).isEqualTo(200); + assertThat(usage.inputTokens()).isEqualTo(100); + assertThat(usage.outputTokens()).isEqualTo(200); assertThat(usage.totalTokens()).isEqualTo(300); + assertThat(usage.details().reasoningOutputTokens()).isEqualTo(150); assertThat(usage.getCount(TokenType.REASONING)).isEqualTo(150); } @Test - @DisplayName("cached input은 전체 input에 중복 합산되지 않는다") - void shouldNotAddCachedInputTokensToInputTotalAgain() { - TokenUsage usage = new TokenUsage( + @DisplayName("cache read와 creation 토큰은 input total에 중복 합산되지 않는다") + void shouldNotAddCacheBreakdownToInputTotalAgain() { + TokenUsage usage = reportedUsage( 100, 200, - new TokenUsageDetails( - 40, // cachedInputTokens - 0, // reasoningOutputTokens - 0 // cachedOutputTokens - ), + new TokenUsageDetails(40L, 20L, null), Map.of() ); - assertThat(usage.promptTokens()).isEqualTo(100); - assertThat(usage.completionTokens()).isEqualTo(200); assertThat(usage.totalTokens()).isEqualTo(300); - assertThat(usage.getCount(TokenType.CACHED_PROMPT)).isEqualTo(40); + assertThat(usage.getCount(TokenType.CACHE_READ_PROMPT)).isEqualTo(40); + assertThat(usage.getCount(TokenType.CACHE_CREATION_PROMPT)).isEqualTo(20); } @Test @DisplayName("input 토큰은 음수일 수 없다") void shouldRejectNegativeInputTokens() { - assertThatThrownBy(() -> - new TokenUsage( - -1, - 0, - new TokenUsageDetails(0, 0, 0), - Map.of() - ) - ) + assertThatThrownBy(() -> reportedUsage(-1, 0, TokenUsageDetails.unreported(), Map.of())) .isInstanceOf(IllegalArgumentException.class) .hasMessage("inputTokens must be non-negative"); } @@ -82,14 +66,7 @@ void shouldRejectNegativeInputTokens() { @Test @DisplayName("output 토큰은 음수일 수 없다") void shouldRejectNegativeOutputTokens() { - assertThatThrownBy(() -> - new TokenUsage( - 0, - -1, - new TokenUsageDetails(0, 0, 0), - Map.of() - ) - ) + assertThatThrownBy(() -> reportedUsage(0, -1, TokenUsageDetails.unreported(), Map.of())) .isInstanceOf(IllegalArgumentException.class) .hasMessage("outputTokens must be non-negative"); } @@ -97,73 +74,88 @@ void shouldRejectNegativeOutputTokens() { @Test @DisplayName("토큰 세부 정보는 null일 수 없다") void shouldRejectNullTokenUsageDetails() { - assertThatThrownBy(() -> - new TokenUsage(100, 200, null, Map.of()) - ) + assertThatThrownBy(() -> reportedUsage(100, 200, null, Map.of())) .isInstanceOf(NullPointerException.class) .hasMessage("details must not be null"); } @Test - @DisplayName("cached input은 전체 input보다 클 수 없다") - void shouldRejectCachedInputTokensGreaterThanInputTokens() { - assertThatThrownBy(() -> - new TokenUsage( - 100, - 200, - new TokenUsageDetails(101, 0, 0), - Map.of() - ) - ) + @DisplayName("사용량 출처는 null일 수 없다") + void shouldRejectNullUsageSource() { + assertThatThrownBy(() -> new TokenUsage( + 100, + 200, + TokenUsageDetails.unreported(), + null, + Map.of() + )) + .isInstanceOf(NullPointerException.class) + .hasMessage("source must not be null"); + } + + @Test + @DisplayName("cache read 토큰은 전체 input보다 클 수 없다") + void shouldRejectCacheReadInputTokensGreaterThanInputTokens() { + assertThatThrownBy(() -> reportedUsage( + 100, + 200, + new TokenUsageDetails(101L, null, null), + Map.of() + )) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Input details must not exceed inputTokens"); + } + + @Test + @DisplayName("cache creation 토큰은 전체 input보다 클 수 없다") + void shouldRejectCacheCreationInputTokensGreaterThanInputTokens() { + assertThatThrownBy(() -> reportedUsage( + 100, + 200, + new TokenUsageDetails(null, 101L, null), + Map.of() + )) .isInstanceOf(IllegalArgumentException.class) - .hasMessage( - "cachedInputTokens must not exceed inputTokens" - ); + .hasMessage("Input details must not exceed inputTokens"); } @Test - @DisplayName("cached output은 전체 output보다 클 수 없다") - void shouldRejectCachedOutputTokensGreaterThanOutputTokens() { - assertThatThrownBy(() -> - new TokenUsage( - 100, - 200, - new TokenUsageDetails(0, 0, 201), - Map.of() - ) - ) + @DisplayName("cache read와 creation의 합은 전체 input보다 클 수 없다") + void shouldRejectCacheBreakdownGreaterThanInputTokens() { + assertThatThrownBy(() -> reportedUsage( + 100, + 200, + new TokenUsageDetails(60L, 41L, null), + Map.of() + )) .isInstanceOf(IllegalArgumentException.class) - .hasMessage("Output details must not exceed outputTokens"); + .hasMessage("Input details must not exceed inputTokens"); } @Test - @DisplayName("reasoning과 cached output의 합은 전체 output보다 클 수 없다") - void shouldRejectOutputDetailsGreaterThanOutputTokens() { - assertThatThrownBy(() -> - new TokenUsage( - 100, - 200, - new TokenUsageDetails(0, 151, 50), - Map.of() - ) - ) + @DisplayName("cache breakdown 검증은 long overflow 없이 수행한다") + void shouldRejectCacheBreakdownWithoutOverflow() { + assertThatThrownBy(() -> reportedUsage( + Long.MAX_VALUE, + 0, + new TokenUsageDetails(Long.MAX_VALUE, 1L, null), + Map.of() + )) .isInstanceOf(IllegalArgumentException.class) - .hasMessage("Output details must not exceed outputTokens"); + .hasMessage("Input details must not exceed inputTokens"); } @Test @DisplayName("reasoning output은 전체 output보다 클 수 없다") void shouldRejectReasoningOutputTokensGreaterThanOutputTokens() { - assertThatThrownBy(() -> - new TokenUsage( - 100, - 200, - new TokenUsageDetails(0, 201, 0), - Map.of() - ) - ) + assertThatThrownBy(() -> reportedUsage( + 100, + 200, + new TokenUsageDetails(null, null, 201L), + Map.of() + )) .isInstanceOf(IllegalArgumentException.class) - .hasMessage("Output details must not exceed outputTokens"); + .hasMessage("reasoningOutputTokens must not exceed outputTokens"); } @Test @@ -171,12 +163,7 @@ void shouldRejectReasoningOutputTokensGreaterThanOutputTokens() { void shouldDefensivelyCopyMetadata() { Map metadata = new HashMap<>(); metadata.put("model", "gpt-4o"); - TokenUsage usage = new TokenUsage( - 100, - 200, - new TokenUsageDetails(0, 0, 0), - metadata - ); + TokenUsage usage = reportedUsage(100, 200, TokenUsageDetails.unreported(), metadata); metadata.put("model", "changed"); @@ -186,12 +173,7 @@ void shouldDefensivelyCopyMetadata() { @Test @DisplayName("metadata가 null이면 빈 Map을 사용한다") void shouldUseEmptyMetadataWhenMetadataIsNull() { - TokenUsage usage = new TokenUsage( - 100, - 200, - new TokenUsageDetails(0, 0, 0), - null - ); + TokenUsage usage = reportedUsage(100, 200, TokenUsageDetails.unreported(), null); assertThat(usage.metadata()).isEmpty(); } @@ -202,14 +184,8 @@ void shouldRejectMetadataWithNullKey() { Map metadata = new HashMap<>(); metadata.put(null, "value"); - assertThatThrownBy(() -> - new TokenUsage( - 100, - 200, - new TokenUsageDetails(0, 0, 0), - metadata - ) - ).isInstanceOf(NullPointerException.class); + assertThatThrownBy(() -> reportedUsage(100, 200, TokenUsageDetails.unreported(), metadata)) + .isInstanceOf(NullPointerException.class); } @Test @@ -218,23 +194,17 @@ void shouldRejectMetadataWithNullValue() { Map metadata = new HashMap<>(); metadata.put("key", null); - assertThatThrownBy(() -> - new TokenUsage( - 100, - 200, - new TokenUsageDetails(0, 0, 0), - metadata - ) - ).isInstanceOf(NullPointerException.class); + assertThatThrownBy(() -> reportedUsage(100, 200, TokenUsageDetails.unreported(), metadata)) + .isInstanceOf(NullPointerException.class); } @Test @DisplayName("breakdown 합계는 전체 토큰과 같을 수 있다") void shouldAllowDetailsEqualToTotalTokens() { - TokenUsage usage = new TokenUsage( + TokenUsage usage = reportedUsage( 100, 200, - new TokenUsageDetails(100, 150, 50), + new TokenUsageDetails(70L, 30L, 200L), Map.of() ); @@ -245,14 +215,29 @@ void shouldAllowDetailsEqualToTotalTokens() { @Test @DisplayName("전체 토큰 합계가 long 범위를 넘으면 예외가 발생한다") void shouldRejectTotalTokenOverflow() { - TokenUsage usage = new TokenUsage( + TokenUsage usage = reportedUsage( Long.MAX_VALUE, 1, - new TokenUsageDetails(0, 0, 0), + TokenUsageDetails.unreported(), Map.of() ); assertThatThrownBy(usage::totalTokens) .isInstanceOf(ArithmeticException.class); } + + private TokenUsage reportedUsage( + long inputTokens, + long outputTokens, + TokenUsageDetails details, + Map metadata + ) { + return new TokenUsage( + inputTokens, + outputTokens, + details, + UsageSource.PROVIDER_REPORTED, + metadata + ); + } } diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java index 7296939..aff726f 100644 --- a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultCostCalculatorTest.java @@ -2,11 +2,16 @@ import io.tokenpilot.core.domain.Cost; import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; +import io.tokenpilot.core.domain.TokenUsageDetails; +import io.tokenpilot.core.domain.UsageSource; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; import java.math.BigDecimal; +import java.util.Currency; +import java.util.Map; import static org.assertj.core.api.Assertions.assertThat; @@ -28,4 +33,33 @@ void calculateStandardTokens() { // Then: (1000 * 0.01 / 1000) + (2000 * 0.03 / 1000) = 0.01 + 0.06 = 0.07 assertThat(cost.value()).isEqualByComparingTo("0.070000"); } + + @Test + @DisplayName("포괄 총량을 배타적인 과금 구간으로 나누어 중복 없이 계산한다") + void calculateInclusiveTotalsWithoutDoubleChargingDetails() { + PricingPlan plan = new PricingPlan( + "provider-model", + Map.of( + TokenType.PROMPT, new BigDecimal("0.01"), + TokenType.CACHE_READ_PROMPT, new BigDecimal("0.002"), + TokenType.CACHE_CREATION_PROMPT, new BigDecimal("0.0125"), + TokenType.COMPLETION, new BigDecimal("0.03"), + TokenType.REASONING, new BigDecimal("0.05") + ), + Currency.getInstance("USD") + ); + TokenUsage usage = new TokenUsage( + 1_000, + 1_000, + new TokenUsageDetails(400L, 100L, 200L), + UsageSource.PROVIDER_REPORTED, + Map.of() + ); + + Cost cost = calculator.calculate(usage, plan); + + // 500 normal input + 400 cache read + 100 cache creation + // 800 normal output + 200 reasoning output + assertThat(cost.value()).isEqualByComparingTo("0.041050"); + } } diff --git a/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultUsageExtractor.java b/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultUsageExtractor.java index 3d0602d..4bf9681 100644 --- a/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultUsageExtractor.java +++ b/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultUsageExtractor.java @@ -2,11 +2,14 @@ import io.tokenpilot.core.domain.TokenUsage; import io.tokenpilot.core.domain.TokenUsageDetails; +import io.tokenpilot.core.domain.UsageSource; import io.tokenpilot.springai.UsageExtractor; import org.springframework.ai.chat.client.ChatClientResponse; import org.springframework.ai.chat.metadata.ChatResponseMetadata; +import org.springframework.ai.chat.metadata.EmptyUsage; import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.model.ModelOptionsUtils; import java.util.HashMap; import java.util.Map; @@ -14,56 +17,144 @@ /** * 기본 {@link UsageExtractor} 구현체. * Spring AI의 {@link Usage} 정보를 {@link TokenUsage}로 변환하며, - * 메타데이터에 포함된 추론(Reasoning) 토큰 등을 식별하여 세분화된 사용량을 추출합니다. + * provider native usage와 메타데이터에서 cache/reasoning breakdown을 식별합니다. + * Provider가 비포괄 총량을 반환하면 Token Pilot의 포괄 총량 계약에 맞게 정규화합니다. */ public class DefaultUsageExtractor implements UsageExtractor { @Override public TokenUsage extract(ChatClientResponse response) { if (response == null) { - return TokenUsage.from(0, 0); + return TokenUsage.unavailable(Map.of()); } ChatResponse chatResponse = response.chatResponse(); if (chatResponse == null || chatResponse.getMetadata() == null) { - return TokenUsage.from(0, 0); + return TokenUsage.unavailable(Map.of()); } ChatResponseMetadata metadata = chatResponse.getMetadata(); Usage usage = metadata.getUsage(); - if (usage == null) { - return new TokenUsage(0, 0, new TokenUsageDetails(0, 0, 0), copyMetadata(metadata, null)); + if (usage == null || usage instanceof EmptyUsage) { + return TokenUsage.unavailable(copyMetadata(metadata, null)); } long prompt = (usage.getPromptTokens() != null) ? usage.getPromptTokens() : 0L; long completion = (usage.getCompletionTokens() != null) ? usage.getCompletionTokens() : 0L; + Object nativeUsage = usage.getNativeUsage(); + Map nativeUsageFields = usageFields(nativeUsage); - Long reasoning = extractReasoningTokens(metadata); + Long cacheRead = extractCacheReadTokens(metadata, nativeUsageFields); + Long cacheCreation = extractCacheCreationTokens(metadata, nativeUsageFields); + Long reasoning = extractReasoningTokens(metadata, nativeUsageFields); + UsageSource source = UsageSource.PROVIDER_REPORTED; + + Long nativeInput = firstNonNull( + numberAt(nativeUsageFields, "input_tokens"), + numberAt(nativeUsageFields, "inputTokens") + ); + if (cacheCreation != null && nativeInput != null && prompt == nativeInput) { + prompt = Math.addExact(prompt, countOrZero(cacheRead)); + prompt = Math.addExact(prompt, cacheCreation); + source = UsageSource.PROVIDER_DERIVED; + } + + Long nativeCandidates = firstNonNull( + numberAt(nativeUsageFields, "candidatesTokenCount"), + numberAt(nativeUsageFields, "candidates_token_count") + ); + if (reasoning != null && nativeCandidates != null && completion == nativeCandidates) { + completion = Math.addExact(completion, reasoning); + source = UsageSource.PROVIDER_DERIVED; + } return new TokenUsage( prompt, completion, - new TokenUsageDetails(0, reasoning, 0), - copyMetadata(metadata, usage.getNativeUsage()) + new TokenUsageDetails(cacheRead, cacheCreation, reasoning), + source, + copyMetadata(metadata, nativeUsage) + ); + } + + private Long extractCacheReadTokens(ChatResponseMetadata metadata, Object nativeUsage) { + return firstNonNull( + numberAt(metadata.get("prompt_tokens_details"), "cached_tokens"), + numberAt(metadata.get("input_tokens_details"), "cached_tokens"), + numberAt(metadata, "cache_read_input_tokens"), + numberAt(metadata, "cachedContentTokenCount"), + numberAt(nativeUsage, "prompt_tokens_details", "cached_tokens"), + numberAt(nativeUsage, "input_tokens_details", "cached_tokens"), + numberAt(nativeUsage, "cache_read_input_tokens"), + numberAt(nativeUsage, "cacheReadInputTokens"), + numberAt(nativeUsage, "cachedContentTokenCount"), + numberAt(nativeUsage, "cached_content_token_count") + ); + } + + private Long extractCacheCreationTokens(ChatResponseMetadata metadata, Object nativeUsage) { + return firstNonNull( + numberAt(metadata, "cache_creation_input_tokens"), + numberAt(metadata, "cacheCreationInputTokens"), + numberAt(nativeUsage, "cache_creation_input_tokens"), + numberAt(nativeUsage, "cacheCreationInputTokens") + ); + } + + private Long extractReasoningTokens(ChatResponseMetadata metadata, Object nativeUsage) { + return firstNonNull( + numberAt(metadata.get("completion_tokens_details"), "reasoning_tokens"), + numberAt(metadata.get("output_tokens_details"), "reasoning_tokens"), + numberAt(metadata, "reasoning_tokens"), + numberAt(metadata, "thoughtsTokenCount"), + numberAt(nativeUsage, "completion_tokens_details", "reasoning_tokens"), + numberAt(nativeUsage, "output_tokens_details", "reasoning_tokens"), + numberAt(nativeUsage, "reasoning_tokens"), + numberAt(nativeUsage, "reasoningTokens"), + numberAt(nativeUsage, "thoughtsTokenCount"), + numberAt(nativeUsage, "thoughts_token_count") ); } - private Long extractReasoningTokens(ChatResponseMetadata metadata) { - // metadata는 직접 Map이 아닐 수 있으므로 get 메서드로 개별 접근 - Object details = metadata.get("completion_tokens_details"); - if (details instanceof Map detailsMap) { - Object reasoning = detailsMap.get("reasoning_tokens"); - if (reasoning instanceof Number num) { - return num.longValue(); + private Long numberAt(Object source, String... path) { + Object current = source; + for (String key : path) { + if (current instanceof ChatResponseMetadata responseMetadata) { + current = responseMetadata.get(key); + } + else if (current instanceof Map map) { + current = map.get(key); + } + else { + return null; } } - - Object directReasoning = metadata.get("reasoning_tokens"); - if (directReasoning instanceof Number num) { - return num.longValue(); + return current instanceof Number number ? number.longValue() : null; + } + + private Long firstNonNull(Long... values) { + for (Long value : values) { + if (value != null) { + return value; + } } + return null; + } + + private long countOrZero(Long count) { + return count == null ? 0L : count; + } - return 0L; + private Map usageFields(Object nativeUsage) { + if (nativeUsage == null) { + return Map.of(); + } + try { + return ModelOptionsUtils.objectToMap(nativeUsage); + } + catch (RuntimeException ignored) { + return Map.of(); + } } private Map copyMetadata(ChatResponseMetadata metadata, Object nativeUsage) { diff --git a/token-pilot-spring-ai/src/test/java/io/tokenpilot/springai/internal/DefaultUsageExtractorTest.java b/token-pilot-spring-ai/src/test/java/io/tokenpilot/springai/internal/DefaultUsageExtractorTest.java index 5b7679d..0bcb005 100644 --- a/token-pilot-spring-ai/src/test/java/io/tokenpilot/springai/internal/DefaultUsageExtractorTest.java +++ b/token-pilot-spring-ai/src/test/java/io/tokenpilot/springai/internal/DefaultUsageExtractorTest.java @@ -2,6 +2,8 @@ import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; +import io.tokenpilot.core.domain.TokenUsageDetails; +import io.tokenpilot.core.domain.UsageSource; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; import org.springframework.ai.chat.client.ChatClientResponse; @@ -40,6 +42,8 @@ void extractStandardUsage() { assertThat(result.getCount(TokenType.PROMPT)).isEqualTo(100L); assertThat(result.getCount(TokenType.COMPLETION)).isEqualTo(200L); + assertThat(result.details()).isEqualTo(TokenUsageDetails.unreported()); + assertThat(result.source()).isEqualTo(UsageSource.PROVIDER_REPORTED); assertThat(result.metadata()).containsEntry("nativeUsage", Map.of("provider", "openai")); } @@ -63,6 +67,74 @@ void extractReasoningUsage() { assertThat(result.getCount(TokenType.PROMPT)).isEqualTo(100L); assertThat(result.getCount(TokenType.COMPLETION)).isEqualTo(200L); assertThat(result.getCount(TokenType.REASONING)).isEqualTo(150L); + assertThat(result.details().cacheReadInputTokens()).isNull(); + assertThat(result.details().cacheCreationInputTokens()).isNull(); + } + + @Test + @DisplayName("OpenAI의 cached input과 reasoning breakdown을 포괄 총량과 분리한다") + void extractOpenAiUsageBreakdown() { + Usage usage = mock(Usage.class); + when(usage.getPromptTokens()).thenReturn(100); + when(usage.getCompletionTokens()).thenReturn(200); + + ChatResponseMetadata metadata = ChatResponseMetadata.builder() + .usage(usage) + .keyValue("prompt_tokens_details", Map.of("cached_tokens", 40)) + .keyValue("completion_tokens_details", Map.of("reasoning_tokens", 150)) + .build(); + + TokenUsage result = extractor.extract(responseWith(metadata)); + + assertThat(result.inputTokens()).isEqualTo(100); + assertThat(result.outputTokens()).isEqualTo(200); + assertThat(result.details()).isEqualTo(new TokenUsageDetails(40L, null, 150L)); + assertThat(result.source()).isEqualTo(UsageSource.PROVIDER_REPORTED); + } + + @Test + @DisplayName("Anthropic의 일반 입력과 cache read/create를 포괄 입력 총량으로 정규화한다") + void normalizeAnthropicInputUsage() { + AnthropicNativeUsage nativeUsage = new AnthropicNativeUsage(50, 100, 25, 60); + Usage usage = mock(Usage.class); + when(usage.getPromptTokens()).thenReturn(50); + when(usage.getCompletionTokens()).thenReturn(60); + when(usage.getNativeUsage()).thenReturn(nativeUsage); + ChatResponseMetadata metadata = ChatResponseMetadata.builder() + .usage(usage) + .build(); + + TokenUsage result = extractor.extract(responseWith(metadata)); + + assertThat(result.inputTokens()).isEqualTo(175); + assertThat(result.outputTokens()).isEqualTo(60); + assertThat(result.details()).isEqualTo(new TokenUsageDetails(100L, 25L, null)); + assertThat(result.source()).isEqualTo(UsageSource.PROVIDER_DERIVED); + } + + @Test + @DisplayName("Gemini의 candidates와 thoughts를 포괄 출력 총량으로 정규화한다") + void normalizeGeminiOutputUsage() { + Map nativeUsage = Map.of( + "promptTokenCount", 100, + "cachedContentTokenCount", 40, + "candidatesTokenCount", 80, + "thoughtsTokenCount", 20 + ); + Usage usage = mock(Usage.class); + when(usage.getPromptTokens()).thenReturn(100); + when(usage.getCompletionTokens()).thenReturn(80); + when(usage.getNativeUsage()).thenReturn(nativeUsage); + ChatResponseMetadata metadata = ChatResponseMetadata.builder() + .usage(usage) + .build(); + + TokenUsage result = extractor.extract(responseWith(metadata)); + + assertThat(result.inputTokens()).isEqualTo(100); + assertThat(result.outputTokens()).isEqualTo(100); + assertThat(result.details()).isEqualTo(new TokenUsageDetails(40L, null, 20L)); + assertThat(result.source()).isEqualTo(UsageSource.PROVIDER_DERIVED); } @Test @@ -72,6 +144,7 @@ void extractNullResponse() { assertThat(result.promptTokens()).isZero(); assertThat(result.completionTokens()).isZero(); + assertThat(result.source()).isEqualTo(UsageSource.UNAVAILABLE); assertThat(result.metadata()).isEmpty(); } @@ -86,6 +159,7 @@ void extractNullMetadata() { assertThat(result.promptTokens()).isZero(); assertThat(result.completionTokens()).isZero(); + assertThat(result.source()).isEqualTo(UsageSource.UNAVAILABLE); assertThat(result.metadata()).isEmpty(); } @@ -102,6 +176,23 @@ void preserveMetadataWithoutUsage() { assertThat(result.promptTokens()).isZero(); assertThat(result.completionTokens()).isZero(); + assertThat(result.source()).isEqualTo(UsageSource.UNAVAILABLE); assertThat(result.metadata()).containsEntry("model_region", "ap-northeast-2"); } + + private ChatClientResponse responseWith(ChatResponseMetadata metadata) { + ChatResponse chatResponse = new ChatResponse( + List.of(new Generation(new org.springframework.ai.chat.messages.AssistantMessage("test"))), + metadata + ); + return new ChatClientResponse(chatResponse, Map.of()); + } + + private record AnthropicNativeUsage( + Integer inputTokens, + Integer cacheReadInputTokens, + Integer cacheCreationInputTokens, + Integer outputTokens + ) { + } }