diff --git a/AGENTS.md b/AGENTS.md index fbf832e..6b0560d 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -4,7 +4,7 @@ Token Pilot is evolving from a Spring AI usage-tracking starter into a framework-independent Java LLM control and accounting core with optional framework and observability adapters. -Current truth: post-call usage normalization, cost calculation, ledger events, Micrometer publishing, Clock-based monthly budget windows, pure budget decisions, legacy provider-boundary BLOCK enforcement, Spring AI integration, and starter autoconfiguration are implemented. Candidate-aware preflight admission, context admission, atomic reservation, and estimate/actual reconciliation are 30-day MVP targets, not current capabilities. +Current truth: post-call usage normalization, cost calculation, ledger events, Micrometer publishing, Clock-based monthly budget windows, pure budget decisions, typed missing-pricing policies, pricing snapshots, legacy provider-boundary BLOCK enforcement, Spring AI integration, and starter autoconfiguration are implemented. Candidate-aware preflight admission, context admission, atomic reservation, and estimate/actual cost reconciliation are 30-day MVP targets, not current capabilities. Distribution direction: publish a framework-independent core and an optional Spring AI convenience starter from the same repository and release train. The existing starter artifact is `token-pilot-starter`; `token-pilot-spring-ai-starter` is only a target name until a compatibility ADR and module change land. @@ -71,8 +71,8 @@ Token Pilot의 제품 포지션은 framework-independent Java LLM control and ac | Module | Status | Notes | | --- | --- | --- | -| `token-pilot-core` | Basic implementation complete | Domain records, pricing, calculator, registry, ledger manager | -| `token-pilot-spring-ai` | Basic implementation complete | Spring AI 2.0.0 `UsageExtractor`, `LedgerAdvisor`, response usage recording, and legacy provider-boundary BLOCK enforcement | +| `token-pilot-core` | Basic implementation complete | Domain records, pricing, calculator, registry, ledger manager, pricing snapshots, and missing-pricing evaluator | +| `token-pilot-spring-ai` | Basic implementation complete | Spring AI 2.0.0 `UsageExtractor`, `LedgerAdvisor`, pricing snapshot resolution, response usage recording, reconciliation decisions, and legacy provider-boundary BLOCK enforcement | | `token-pilot-micrometer` | Basic implementation complete | `MetricsOptions`, tag whitelist, and metric metadata exist; metric ownership must be narrowed | | `token-pilot-budget` | Basic non-atomic implementation | Typed monthly keys, Clock/ZoneId windows, and pure status/admission decisions implemented; needs candidate estimation, reservation, idempotency, and reconciliation | | `token-pilot-notification` | Basic implementation complete | Event API and deduplication exist; not yet connected to the full advisor/budget lifecycle | @@ -326,7 +326,7 @@ The active checklist is in `docs/30_DAY_MVP_REPORT.md`; detailed long-term works - `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. - `Cost` now requires an explicit currency, preserves exact internal `BigDecimal` precision, rejects negative values and cross-currency operations, and defers scale-6 `HALF_UP` rounding to `RoundingPolicy.COST_BOUNDARY_ROUNDING`. - Budget money interfaces now use `Cost` while preserving `BudgetKey`, `BudgetPolicy`, Clock/ZoneId monthly windows, and per-key policy snapshots. -- Until the typed missing-pricing policy lands, `DefaultLedgerManager` preserves the legacy fail-open result as an explicit zero USD `Cost`; do not confuse that compatibility behavior with a priced zero-rate plan. +- The legacy `DefaultLedgerManager.record(String, ...)` path preserves an explicit zero USD fail-open result for a missing plan; the pricing-snapshot path applies `MissingPricingPolicy` and records `UNPRICED` or rejects before provider invocation, so neither behavior is a priced zero-rate plan. - 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. - The legacy provider boundary blocks an already-exhausted budget decision before provider invocation. Its candidate-free `STATUS` input is a regression guard, not admission evidence; the flow remains check-then-add and is not an atomic reservation. - Current Micrometer `ai.token.*` metrics may duplicate Spring AI Observability; preserve compatibility while deciding default suppression or replacement. @@ -398,6 +398,7 @@ Stage and deploy a Central release: ### 2026-08-04 +- Added typed missing-pricing policies, immutable pricing snapshots, core rate validation/reconciliation decisions, and Spring AI pre-call pricing resolution; `FAIL_CLOSED` rejects missing plans/rates before provider invocation and `FAIL_OPEN` preserves `UNPRICED`. - Preserved the deprecated `BudgetNotificationEvent.currentUsage()` compatibility accessor through 0.1.x while migrating handlers to `projectedUsage()`; removal is planned for 0.2.0. ### 2026-07-29 diff --git a/token-pilot-autoconfigure/src/main/java/io/tokenpilot/autoconfigure/TokenPilotAutoConfiguration.java b/token-pilot-autoconfigure/src/main/java/io/tokenpilot/autoconfigure/TokenPilotAutoConfiguration.java index ac1ac6a..e3d2735 100644 --- a/token-pilot-autoconfigure/src/main/java/io/tokenpilot/autoconfigure/TokenPilotAutoConfiguration.java +++ b/token-pilot-autoconfigure/src/main/java/io/tokenpilot/autoconfigure/TokenPilotAutoConfiguration.java @@ -7,8 +7,10 @@ import io.tokenpilot.core.CostCalculator; import io.tokenpilot.core.LedgerListener; import io.tokenpilot.core.LedgerManager; +import io.tokenpilot.core.PricingEvaluator; import io.tokenpilot.core.PricingProvider; import io.tokenpilot.core.PricingRegistry; +import io.tokenpilot.core.domain.MissingPricingPolicy; import io.tokenpilot.core.internal.LedgerComponents; import io.tokenpilot.micrometer.internal.LedgerMicrometerComponents; import io.tokenpilot.springai.LedgerAdvisor; @@ -70,6 +72,15 @@ public CostCalculator costCalculator() { return LedgerComponents.defaultCostCalculator(); } + /** + * Pricing snapshot rate와 actual model 정합성을 평가하는 정책을 등록합니다. + */ + @Bean + @ConditionalOnMissingBean + public PricingEvaluator pricingEvaluator() { + return LedgerComponents.defaultPricingEvaluator(); + } + /** * 비용 기록 및 리스너 관리를 담당하는 LedgerManager를 등록합니다. */ @@ -108,7 +119,8 @@ public LedgerAdvisor ledgerAdvisor( ObjectProvider budgetEvaluator, ObjectProvider budgetStateStore, CostCalculator costCalculator, - PricingRegistry pricingRegistry + PricingRegistry pricingRegistry, + PricingEvaluator pricingEvaluator ) { BudgetEvaluator evaluator = budgetEvaluator.getIfAvailable(); BudgetStateStore stateStore = budgetStateStore.getIfAvailable(); @@ -120,11 +132,19 @@ public LedgerAdvisor ledgerAdvisor( evaluator, stateStore, costCalculator, - pricingRegistry + pricingRegistry, + pricingEvaluator, + MissingPricingPolicy.FAIL_CLOSED ); } - return LedgerSpringAiComponents.defaultLedgerAdvisor(ledgerManager, usageExtractor); + return LedgerSpringAiComponents.defaultLedgerAdvisor( + ledgerManager, + usageExtractor, + costCalculator, + pricingRegistry, + pricingEvaluator + ); } /** diff --git a/token-pilot-autoconfigure/src/test/java/io/tokenpilot/autoconfigure/TokenPilotAutoConfigurationTest.java b/token-pilot-autoconfigure/src/test/java/io/tokenpilot/autoconfigure/TokenPilotAutoConfigurationTest.java index f1e3760..e1e89c0 100644 --- a/token-pilot-autoconfigure/src/test/java/io/tokenpilot/autoconfigure/TokenPilotAutoConfigurationTest.java +++ b/token-pilot-autoconfigure/src/test/java/io/tokenpilot/autoconfigure/TokenPilotAutoConfigurationTest.java @@ -11,11 +11,17 @@ import io.tokenpilot.budget.BudgetWindow; import io.tokenpilot.core.CostCalculator; import io.tokenpilot.core.LedgerManager; +import io.tokenpilot.core.PricingEvaluator; import io.tokenpilot.core.PricingProvider; import io.tokenpilot.core.PricingRegistry; import io.tokenpilot.core.domain.Cost; import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingReconciliationResult; +import io.tokenpilot.core.domain.PricingResolution; +import io.tokenpilot.core.domain.PricingSnapshot; +import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; +import io.tokenpilot.core.exception.MissingPricingException; import io.tokenpilot.notification.BudgetNotificationHandler; import io.tokenpilot.notification.BudgetNotificationService; import io.tokenpilot.notification.NotificationStateStore; @@ -29,6 +35,7 @@ import org.junit.jupiter.params.provider.MethodSource; import org.springframework.ai.chat.client.ChatClientRequest; import org.springframework.ai.chat.client.advisor.api.AdvisorChain; +import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -47,6 +54,7 @@ import static io.tokenpilot.core.domain.TokenType.COMPLETION; import static io.tokenpilot.core.domain.TokenType.PROMPT; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; import static org.junit.jupiter.params.provider.Arguments.argumentSet; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; @@ -69,6 +77,7 @@ void shouldRegisterDefaultBeans() { assertThat(context).hasSingleBean(PricingProvider.class); assertThat(context).hasSingleBean(PricingRegistry.class); assertThat(context).hasSingleBean(CostCalculator.class); + assertThat(context).hasSingleBean(PricingEvaluator.class); assertThat(context).hasSingleBean(LedgerManager.class); assertThat(context).hasSingleBean(UsageExtractor.class); @@ -82,6 +91,54 @@ void shouldRegisterDefaultBeans() { }); } + @Test + @DisplayName("Ledger-only advisor는 Prompt model과 기본 policy로 pricing snapshot을 resolve해야 한다") + void shouldResolvePricingSnapshotInLedgerOnlyAdvisor() { + this.contextRunner + .withPropertyValues( + "token-pilot.budget.enabled=false", + PROP_MODEL_ID + "=gpt-4o", + PROP_PROMPT + "=0.005", + PROP_COMPLETION + "=0.015", + PROP_CURRENCY + "=USD" + ) + .run(context -> { + LedgerAdvisor advisor = context.getBean(LedgerAdvisor.class); + ChatClientRequest request = new ChatClientRequest( + new Prompt( + "test", + ChatOptions.builder().model("gpt-4o").build() + ), + Map.of() + ); + + ChatClientRequest resolvedRequest = advisor.before( + request, + mock(AdvisorChain.class) + ); + Optional snapshot = resolvedRequest.context() + .values() + .stream() + .filter(PricingSnapshot.class::isInstance) + .map(PricingSnapshot.class::cast) + .findFirst(); + + assertThat(resolvedRequest.context().values()) + .contains(PricingResolution.RESOLVED); + assertThat(snapshot) + .isPresent() + .get() + .satisfies(resolvedSnapshot -> { + assertThat(resolvedSnapshot.modelId()) + .isEqualTo("gpt-4o"); + assertThat(resolvedSnapshot.pricingPolicyId()) + .isEqualTo(PricingPlan.DEFAULT_PRICING_POLICY_ID); + assertThat(resolvedSnapshot.currency()) + .isEqualTo(Currency.getInstance("USD")); + }); + }); + } + @Test @DisplayName("설정 값이 없을 경우 빈 목록을 가진 PricingProvider가 생성되어야 한다") void shouldRegisterDefaultPricingProviderWhenNoProperties() { @@ -240,14 +297,23 @@ void shouldUseUserClockForMonthlyBudgetWindow() { void shouldWireBudgetEvaluatorIntoLedgerAdvisorWhenBudgetEnabled() { this.contextRunner .withUserConfiguration(RecordingBudgetEvaluatorConfiguration.class) - .withPropertyValues("token-pilot.budget.enabled=true") + .withPropertyValues( + "token-pilot.budget.enabled=true", + PROP_MODEL_ID + "=gpt-4o", + PROP_PROMPT + "=0.005", + PROP_COMPLETION + "=0.015", + PROP_CURRENCY + "=USD" + ) .run(context -> { LedgerAdvisor advisor = context.getBean(LedgerAdvisor.class); RecordingBudgetEvaluator evaluator = context.getBean(RecordingBudgetEvaluator.class); ChatClientRequest request = new ChatClientRequest( new Prompt("test"), - Map.of("tenant_id", "tenant-abc") + Map.of( + "tenant_id", "tenant-abc", + "tokenpilot.model.id", "gpt-4o" + ) ); advisor.before(request, mock(AdvisorChain.class)); @@ -261,6 +327,30 @@ void shouldWireBudgetEvaluatorIntoLedgerAdvisorWhenBudgetEnabled() { }); } + @Test + @DisplayName("Budget가 활성화되면 missing pricing policy 기본값은 FAIL_CLOSED여야 한다") + void shouldUseFailClosedMissingPricingPolicyWhenBudgetEnabled() { + this.contextRunner + .withUserConfiguration(RecordingBudgetEvaluatorConfiguration.class) + .withPropertyValues("token-pilot.budget.enabled=true") + .run(context -> { + LedgerAdvisor advisor = context.getBean(LedgerAdvisor.class); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of( + "tenant_id", "tenant-abc", + "tokenpilot.model.id", "missing-model" + ) + ); + + assertThatThrownBy(() -> advisor.before(request, mock(AdvisorChain.class))) + .isInstanceOf(MissingPricingException.class) + .hasMessage("MISSING_PLAN") + .extracting(exception -> ((MissingPricingException) exception).getResolution()) + .isEqualTo(PricingResolution.MISSING_PLAN); + }); + } + @Test @DisplayName("token-pilot.budget.enabled=false 일 때 Budget 관련 빈이 등록되지 않아야 한다") void shouldNotRegisterBudgetBeansWhenDisabled() { @@ -302,6 +392,9 @@ void shouldNotOverrideUserDefinedBeans() { assertThat(context).hasSingleBean(PricingRegistry.class); assertThat(context.getBean(PricingRegistry.class)) .isInstanceOf(UserCustomPricingRegistry.class); + assertThat(context).hasSingleBean(PricingEvaluator.class); + assertThat(context.getBean(PricingEvaluator.class)) + .isInstanceOf(UserCustomPricingEvaluator.class); }); } @@ -373,11 +466,35 @@ static class UserCustomConfiguration { public PricingRegistry pricingRegistry() { return new UserCustomPricingRegistry(); } + + @Bean + public PricingEvaluator pricingEvaluator() { + return new UserCustomPricingEvaluator(); + } + } + + static class UserCustomPricingEvaluator implements PricingEvaluator { + @Override + public PricingResolution validateSnapshotRates(Optional snapshot) { + return PricingResolution.MISSING_PLAN; + } + + @Override + public PricingReconciliationResult determineReconciliation( + Optional snapshot, + String actualModelId + ) { + return PricingReconciliationResult.RECONCILIATION_REQUIRED; + } } static class UserCustomPricingRegistry implements PricingRegistry { @Override public void registerPlan(PricingPlan plan) {} @Override public Optional getPlan(String modelId) { return Optional.empty(); } + @Override public Optional getPlan(String modelId, String pricingPolicyId) { return Optional.empty(); } + @Override public Optional resolveSnapshot(String modelId, String pricingPolicyId) { return Optional.empty(); } + @Override public PricingResolution resolveRate(String modelId, TokenType tokenType) { return PricingResolution.MISSING_PLAN; } + @Override public PricingResolution resolveRate(String modelId, TokenType tokenType, Currency expectedCurrency) { return PricingResolution.MISSING_PLAN; } } @Configuration(proxyBeanMethods = false) diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/LedgerManager.java b/token-pilot-core/src/main/java/io/tokenpilot/core/LedgerManager.java index 271e8a9..e4a96dd 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/LedgerManager.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/LedgerManager.java @@ -1,6 +1,8 @@ package io.tokenpilot.core; import io.tokenpilot.core.domain.Cost; +import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingSnapshot; import io.tokenpilot.core.domain.TokenUsage; import java.util.Map; @@ -17,4 +19,22 @@ public interface LedgerManager { * @return 산출된 비용 */ Cost record(String modelId, TokenUsage usage, Map tags); + + /** + * 이미 resolve된 가격 정책으로 호출 정보를 기록하고 최종 비용을 계산합니다. + * @param plan provider 호출 전에 resolve된 가격 정책 + * @param usage 토큰 사용량 + * @param tags 추가 메타데이터 (tenant_id, user_id 등) + * @return 산출된 비용 + */ + Cost record(PricingPlan plan, TokenUsage usage, Map tags); + + /** + * 요청 단위 pricing snapshot으로 호출 정보를 기록하고 최종 비용을 계산합니다. + * @param snapshot provider 호출 전에 보존된 pricing snapshot + * @param usage 토큰 사용량 + * @param tags 추가 메타데이터 (tenant_id, user_id 등) + * @return 산출된 비용 + */ + Cost record(PricingSnapshot snapshot, TokenUsage usage, Map tags); } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/PricingEvaluator.java b/token-pilot-core/src/main/java/io/tokenpilot/core/PricingEvaluator.java new file mode 100644 index 0000000..16c6976 --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/PricingEvaluator.java @@ -0,0 +1,33 @@ +package io.tokenpilot.core; + +import io.tokenpilot.core.domain.PricingReconciliationResult; +import io.tokenpilot.core.domain.PricingResolution; +import io.tokenpilot.core.domain.PricingSnapshot; + +import java.util.Optional; + +/** + * Pricing snapshot의 사용 가능 여부와 actual model 정합성을 판단하는 정책 계약. + */ +public interface PricingEvaluator { + + /** + * Snapshot에 요청 처리에 필요한 rate가 있는지 검증합니다. + * + * @param snapshot 검증할 pricing snapshot, 조회되지 않은 경우 empty + * @return snapshot 및 필수 rate의 resolution + */ + PricingResolution validateSnapshotRates(Optional snapshot); + + /** + * 호출 전 snapshot을 actual 응답 모델에 적용할 수 있는지 판단합니다. + * + * @param snapshot 호출 전에 확정한 pricing snapshot, 확정되지 않은 경우 empty + * @param actualModelId provider가 반환한 actual model id + * @return pricing reconciliation 판단 결과 + */ + PricingReconciliationResult determineReconciliation( + Optional snapshot, + String actualModelId + ); +} diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/PricingRegistry.java b/token-pilot-core/src/main/java/io/tokenpilot/core/PricingRegistry.java index 673c02a..3c0a5a3 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/PricingRegistry.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/PricingRegistry.java @@ -1,7 +1,11 @@ package io.tokenpilot.core; import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingResolution; +import io.tokenpilot.core.domain.PricingSnapshot; +import io.tokenpilot.core.domain.TokenType; +import java.util.Currency; import java.util.Optional; /** @@ -15,6 +19,39 @@ public interface PricingRegistry { */ Optional getPlan(String modelId); + /** + * 모델 식별자와 pricing policy id로 등록된 가격 정책을 조회합니다. + * @param modelId 모델 식별자 + * @param pricingPolicyId pricing policy 식별자 + * @return 가격 정책 (존재하지 않을 경우 empty) + */ + Optional getPlan(String modelId, String pricingPolicyId); + + /** + * 모델 식별자와 pricing policy id로 요청 단위 pricing snapshot을 resolve합니다. + * @param modelId 모델 식별자 + * @param pricingPolicyId pricing policy 식별자 + * @return pricing snapshot (존재하지 않을 경우 empty) + */ + Optional resolveSnapshot(String modelId, String pricingPolicyId); + + /** + * 모델과 토큰 타입에 대한 가격 결정 결과를 조회합니다. + * @param modelId 모델 식별자 + * @param tokenType 토큰 타입 + * @return 가격 결정 결과 + */ + PricingResolution resolveRate(String modelId, TokenType tokenType); + + /** + * 모델과 토큰 타입에 대한 가격 결정 결과를 기대 통화 기준으로 조회합니다. + * @param modelId 모델 식별자 + * @param tokenType 토큰 타입 + * @param expectedCurrency 기대 통화 + * @return 가격 결정 결과 + */ + PricingResolution resolveRate(String modelId, TokenType tokenType, Currency expectedCurrency); + /** * 새로운 가격 정책을 등록하거나 업데이트합니다. * @param plan 가격 정책 diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/MissingPricingPolicy.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/MissingPricingPolicy.java new file mode 100644 index 0000000..3c46bcc --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/MissingPricingPolicy.java @@ -0,0 +1,6 @@ +package io.tokenpilot.core.domain; + +public enum MissingPricingPolicy { + FAIL_OPEN, + FAIL_CLOSED +} 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 fb962ca..9de0d72 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 @@ -5,6 +5,7 @@ import java.util.Currency; import java.util.EnumMap; import java.util.Map; +import java.util.Objects; /** * 특정 모델의 가격 정책 정보. @@ -16,15 +17,25 @@ */ public record PricingPlan( String modelId, + String pricingPolicyId, Map rates, Currency currency ) { + public static final String DEFAULT_PRICING_POLICY_ID = "default"; + public PricingPlan { - rates = Collections.unmodifiableMap(new EnumMap<>(rates)); + if (pricingPolicyId == null || pricingPolicyId.isBlank()) { + throw new IllegalArgumentException("pricingPolicyId must not be blank"); + } + + Objects.requireNonNull(rates, "rates must not be null"); + Map copiedRates = new EnumMap<>(TokenType.class); + copiedRates.putAll(rates); + rates = Collections.unmodifiableMap(copiedRates); if (currency == null) { currency = Currency.getInstance("USD"); } - // 모든 단가는 0 이상이어야 함 + rates.values().forEach(v -> { if (v.compareTo(BigDecimal.ZERO) < 0) { throw new IllegalArgumentException("Price cannot be negative"); @@ -32,11 +43,22 @@ public record PricingPlan( }); } + public PricingPlan(String modelId, Map rates, Currency currency) { + this(modelId, DEFAULT_PRICING_POLICY_ID, rates, currency); + } + /** * 기본 입력/출력 단가와 통화를 사용하는 {@link PricingPlan}을 생성합니다. */ public PricingPlan(String modelId, BigDecimal promptPricePerK, BigDecimal completionPricePerK, Currency currency) { - this(modelId, createRates(promptPricePerK, completionPricePerK), currency); + this(modelId, DEFAULT_PRICING_POLICY_ID, createRates(promptPricePerK, completionPricePerK), currency); + } + + /** + * 기본 입력/출력 단가와 pricing policy id, 통화를 사용하는 {@link PricingPlan}을 생성합니다. + */ + public PricingPlan(String modelId, String pricingPolicyId, BigDecimal promptPricePerK, BigDecimal completionPricePerK, Currency currency) { + this(modelId, pricingPolicyId, createRates(promptPricePerK, completionPricePerK), currency); } /** @@ -77,12 +99,24 @@ public BigDecimal getRate(TokenType type) { return rates.get(type); } - // Fallback Logic - return switch (type) { - 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); - }; + return PricingRateFallback.fallbackFor(type) + .map(fallbackType -> rates.getOrDefault(fallbackType, BigDecimal.ZERO)) + .orElse(BigDecimal.ZERO); + } + + /** + * 특정 토큰 타입의 가격 결정 결과를 반환합니다. + * 명시적으로 등록된 0 rate는 {@link PricingResolution#RESOLVED}로, + * 누락된 rate는 {@link PricingResolution#MISSING_RATE}로 표현합니다. + */ + public PricingResolution resolveRate(TokenType type) { + if (rates.containsKey(type)) { + return PricingResolution.RESOLVED; + } + + return PricingRateFallback.fallbackFor(type) + .filter(rates::containsKey) + .map(fallbackType -> PricingResolution.RESOLVED) + .orElse(PricingResolution.MISSING_RATE); } } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingRateFallback.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingRateFallback.java new file mode 100644 index 0000000..2a57332 --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingRateFallback.java @@ -0,0 +1,27 @@ +package io.tokenpilot.core.domain; + +import java.util.Optional; + +enum PricingRateFallback { + REASONING_TO_COMPLETION(TokenType.REASONING, TokenType.COMPLETION), + CACHE_READ_PROMPT_TO_PROMPT(TokenType.CACHE_READ_PROMPT, TokenType.PROMPT), + CACHE_CREATION_PROMPT_TO_PROMPT(TokenType.CACHE_CREATION_PROMPT, TokenType.PROMPT); + + private final TokenType tokenType; + private final TokenType fallbackTokenType; + + PricingRateFallback(TokenType tokenType, TokenType fallbackTokenType) { + this.tokenType = tokenType; + this.fallbackTokenType = fallbackTokenType; + } + + static Optional fallbackFor(TokenType tokenType) { + for (PricingRateFallback fallback : values()) { + if (fallback.tokenType == tokenType) { + return Optional.of(fallback.fallbackTokenType); + } + } + + return Optional.empty(); + } +} diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingReconciliationResult.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingReconciliationResult.java new file mode 100644 index 0000000..8964bd1 --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingReconciliationResult.java @@ -0,0 +1,10 @@ +package io.tokenpilot.core.domain; + +/** + * Actual reconciliation 결과. + */ +public enum PricingReconciliationResult { + RECONCILED, + RECONCILIATION_REQUIRED, + UNPRICED +} diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingResolution.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingResolution.java new file mode 100644 index 0000000..2a6b6fe --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingResolution.java @@ -0,0 +1,12 @@ +package io.tokenpilot.core.domain; + +public enum PricingResolution { + RESOLVED, + MISSING_PLAN, + MISSING_RATE, + CURRENCY_MISMATCH; + + public boolean isResolved() { + return this == RESOLVED; + } +} diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingSnapshot.java b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingSnapshot.java new file mode 100644 index 0000000..3d71c4c --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/domain/PricingSnapshot.java @@ -0,0 +1,58 @@ +package io.tokenpilot.core.domain; + +import java.math.BigDecimal; +import java.time.Instant; +import java.util.Collections; +import java.util.Currency; +import java.util.EnumMap; +import java.util.Map; +import java.util.Objects; + +/** + * provider 호출 전에 확정된 요청 단위 pricing snapshot. + */ +public record PricingSnapshot( + String modelId, + String pricingPolicyId, + String catalogVersion, + Instant checkedAt, + Map rates, + Currency currency +) { + public static final String DEFAULT_CATALOG_VERSION = "default"; + + public PricingSnapshot { + if (modelId == null || modelId.isBlank()) { + throw new IllegalArgumentException("modelId must not be blank"); + } + if (pricingPolicyId == null || pricingPolicyId.isBlank()) { + throw new IllegalArgumentException("pricingPolicyId must not be blank"); + } + if (catalogVersion == null || catalogVersion.isBlank()) { + throw new IllegalArgumentException("catalogVersion must not be blank"); + } + + checkedAt = Objects.requireNonNull(checkedAt, "checkedAt must not be null"); + currency = Objects.requireNonNull(currency, "currency must not be null"); + Objects.requireNonNull(rates, "rates must not be null"); + Map copiedRates = new EnumMap<>(TokenType.class); + copiedRates.putAll(rates); + rates = Collections.unmodifiableMap(copiedRates); + rates.values().forEach(rate -> { + if (rate.compareTo(BigDecimal.ZERO) < 0) { + throw new IllegalArgumentException("rate must not be negative"); + } + }); + } + + public static PricingSnapshot from(PricingPlan plan, String catalogVersion, Instant checkedAt) { + return new PricingSnapshot( + plan.modelId(), + plan.pricingPolicyId(), + catalogVersion, + checkedAt, + plan.rates(), + plan.currency() + ); + } +} diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/exception/MissingPricingException.java b/token-pilot-core/src/main/java/io/tokenpilot/core/exception/MissingPricingException.java new file mode 100644 index 0000000..857f7b4 --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/exception/MissingPricingException.java @@ -0,0 +1,16 @@ +package io.tokenpilot.core.exception; + +import io.tokenpilot.core.domain.PricingResolution; + +public class MissingPricingException extends RuntimeException { + private final PricingResolution resolution; + + public MissingPricingException(PricingResolution resolution) { + super(resolution.name()); + this.resolution = resolution; + } + + public PricingResolution getResolution() { + return resolution; + } +} 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 7b40429..1f8e1fc 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 @@ -3,8 +3,10 @@ import io.tokenpilot.core.CostCalculator; import io.tokenpilot.core.domain.Cost; import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingResolution; import io.tokenpilot.core.domain.TokenType; import io.tokenpilot.core.domain.TokenUsage; +import io.tokenpilot.core.exception.MissingPricingException; import java.math.BigDecimal; @@ -24,21 +26,27 @@ public Cost calculate(TokenUsage usage, PricingPlan plan) { 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))); + BigDecimal totalCostValue = costFor(regularInput, plan, TokenType.PROMPT) + .add(costFor(cacheReadInput, plan, TokenType.CACHE_READ_PROMPT)) + .add(costFor(cacheCreationInput, plan, TokenType.CACHE_CREATION_PROMPT)) + .add(costFor(regularOutput, plan, TokenType.COMPLETION)) + .add(costFor(reasoningOutput, plan, TokenType.REASONING)); return new Cost(totalCostValue, plan.currency()); } - private BigDecimal costFor(long count, BigDecimal rate) { + private BigDecimal costFor(long count, PricingPlan plan, TokenType tokenType) { if (count == 0) { return BigDecimal.ZERO; } - return rate.multiply(BigDecimal.valueOf(count)) + PricingResolution resolution = plan.resolveRate(tokenType); + if (!resolution.isResolved()) { + throw new MissingPricingException(resolution); + } + + return plan.getRate(tokenType) + .multiply(BigDecimal.valueOf(count)) .movePointLeft(3); } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultLedgerManager.java b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultLedgerManager.java index 28bd66c..e9e66f4 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultLedgerManager.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultLedgerManager.java @@ -49,12 +49,38 @@ public Cost record(String modelId, TokenUsage usage, Map tags) { .map(plan -> costCalculator.calculate(usage, plan)) .orElse(Cost.zero(UNPRICED_COST_CURRENCY)); + publish(modelId, usage, cost, tags); + + return cost; + } + + @Override + public Cost record(PricingPlan plan, TokenUsage usage, Map tags) { + return recordResolvedPlan(plan, usage, tags); + } + + @Override + public Cost record(PricingSnapshot snapshot, TokenUsage usage, Map tags) { + PricingPlan plan = new PricingPlan( + snapshot.modelId(), + snapshot.pricingPolicyId(), + snapshot.rates(), + snapshot.currency() + ); + return recordResolvedPlan(plan, usage, tags); + } + + private Cost recordResolvedPlan(PricingPlan plan, TokenUsage usage, Map tags) { + Cost cost = costCalculator.calculate(usage, plan); + publish(plan.modelId(), usage, cost, tags); + return cost; + } + + private void publish(String modelId, TokenUsage usage, Cost cost, Map tags) { // 이벤트 발행 (리스너들에게 전파) if (!listeners.isEmpty()) { CostRecordedEvent event = new CostRecordedEvent(modelId, usage, cost, tags); listeners.forEach(listener -> listener.onRecord(event)); } - - return cost; } } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultPricingEvaluator.java b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultPricingEvaluator.java new file mode 100644 index 0000000..c75de14 --- /dev/null +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/DefaultPricingEvaluator.java @@ -0,0 +1,59 @@ +package io.tokenpilot.core.internal; + +import io.tokenpilot.core.PricingEvaluator; +import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingReconciliationResult; +import io.tokenpilot.core.domain.PricingResolution; +import io.tokenpilot.core.domain.PricingSnapshot; +import io.tokenpilot.core.domain.TokenType; + +import java.util.Objects; +import java.util.Optional; + +class DefaultPricingEvaluator implements PricingEvaluator { + + @Override + public PricingResolution validateSnapshotRates(Optional snapshot) { + Objects.requireNonNull(snapshot, "snapshot must not be null"); + + if (snapshot.isEmpty()) { + return PricingResolution.MISSING_PLAN; + } + + return resolveRequiredRates(snapshot.get()); + } + + @Override + public PricingReconciliationResult determineReconciliation( + Optional snapshot, + String actualModelId + ) { + Objects.requireNonNull(snapshot, "snapshot must not be null"); + + if (snapshot.isEmpty()) { + return PricingReconciliationResult.UNPRICED; + } + + if (!snapshot.get().modelId().equals(actualModelId)) { + return PricingReconciliationResult.RECONCILIATION_REQUIRED; + } + + return PricingReconciliationResult.RECONCILED; + } + + private PricingResolution resolveRequiredRates(PricingSnapshot snapshot) { + PricingPlan plan = new PricingPlan( + snapshot.modelId(), + snapshot.pricingPolicyId(), + snapshot.rates(), + snapshot.currency() + ); + + PricingResolution promptResolution = plan.resolveRate(TokenType.PROMPT); + if (!promptResolution.isResolved()) { + return promptResolution; + } + + return plan.resolveRate(TokenType.COMPLETION); + } +} diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/InMemoryPricingRegistry.java b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/InMemoryPricingRegistry.java index 5dccfee..926c7c8 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/InMemoryPricingRegistry.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/InMemoryPricingRegistry.java @@ -3,10 +3,16 @@ import io.tokenpilot.core.PricingRegistry; import io.tokenpilot.core.PricingProvider; import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingResolution; +import io.tokenpilot.core.domain.PricingSnapshot; +import io.tokenpilot.core.domain.TokenType; +import java.time.Instant; import java.util.Collection; +import java.util.Currency; import java.util.List; import java.util.Map; +import java.util.Objects; import java.util.Optional; import java.util.concurrent.ConcurrentHashMap; @@ -15,7 +21,7 @@ * 메모리 기반 가격 정책 저장소 구현체. */ class InMemoryPricingRegistry implements PricingRegistry { - private final Map plans = new ConcurrentHashMap<>(); + private final Map plans = new ConcurrentHashMap<>(); public InMemoryPricingRegistry() { } @@ -31,11 +37,51 @@ public InMemoryPricingRegistry(List providers) { @Override public Optional getPlan(String modelId) { - return Optional.ofNullable(plans.get(modelId)); + return getPlan(modelId, PricingPlan.DEFAULT_PRICING_POLICY_ID); + } + + @Override + public Optional getPlan(String modelId, String pricingPolicyId) { + return Optional.ofNullable(plans.get(new PricingPlanKey(modelId, pricingPolicyId))); + } + + @Override + public Optional resolveSnapshot(String modelId, String pricingPolicyId) { + return getPlan(modelId, pricingPolicyId) + .map(plan -> PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.now() + )); + } + + @Override + public PricingResolution resolveRate(String modelId, TokenType tokenType) { + return getPlan(modelId) + .map(plan -> plan.resolveRate(tokenType)) + .orElse(PricingResolution.MISSING_PLAN); + } + + @Override + public PricingResolution resolveRate(String modelId, TokenType tokenType, Currency expectedCurrency) { + Objects.requireNonNull(expectedCurrency, "expectedCurrency must not be null"); + + return getPlan(modelId) + .map(plan -> { + if (!plan.currency().equals(expectedCurrency)) { + return PricingResolution.CURRENCY_MISMATCH; + } + + return plan.resolveRate(tokenType); + }) + .orElse(PricingResolution.MISSING_PLAN); } @Override public void registerPlan(PricingPlan plan) { - plans.put(plan.modelId(), plan); + plans.put(new PricingPlanKey(plan.modelId(), plan.pricingPolicyId()), plan); + } + + private record PricingPlanKey(String modelId, String pricingPolicyId) { } } diff --git a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/LedgerComponents.java b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/LedgerComponents.java index 533587f..8ced30b 100644 --- a/token-pilot-core/src/main/java/io/tokenpilot/core/internal/LedgerComponents.java +++ b/token-pilot-core/src/main/java/io/tokenpilot/core/internal/LedgerComponents.java @@ -3,6 +3,7 @@ import io.tokenpilot.core.CostCalculator; import io.tokenpilot.core.LedgerListener; import io.tokenpilot.core.LedgerManager; +import io.tokenpilot.core.PricingEvaluator; import io.tokenpilot.core.PricingProvider; import io.tokenpilot.core.PricingRegistry; @@ -20,6 +21,10 @@ public static CostCalculator defaultCostCalculator() { return new DefaultCostCalculator(); } + public static PricingEvaluator defaultPricingEvaluator() { + return new DefaultPricingEvaluator(); + } + public static PricingRegistry inMemoryPricingRegistry(List providers) { return new InMemoryPricingRegistry(providers); } diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/MissingPricingPolicyTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/MissingPricingPolicyTest.java new file mode 100644 index 0000000..f897c1b --- /dev/null +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/MissingPricingPolicyTest.java @@ -0,0 +1,27 @@ +package io.tokenpilot.core.domain; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; + +class MissingPricingPolicyTest { + + @Test + @DisplayName("MissingPricingPolicy는 FAIL_OPEN과 FAIL_CLOSED를 제공한다") + void exposesMissingPricingPolicies() { + assertThat(MissingPricingPolicy.values()) + .containsExactly( + MissingPricingPolicy.FAIL_OPEN, + MissingPricingPolicy.FAIL_CLOSED + ); + } + + @Test + @DisplayName("PricingResolution은 pricing 상태만 표현한다") + void pricingResolutionDoesNotIncludePolicyStates() { + assertThat(PricingResolution.values()) + .extracting(Enum::name) + .doesNotContain("FAIL_OPEN", "FAIL_CLOSED"); + } +} diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/PricingPlanTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/PricingPlanTest.java new file mode 100644 index 0000000..1baaade --- /dev/null +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/PricingPlanTest.java @@ -0,0 +1,169 @@ +package io.tokenpilot.core.domain; + +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; + +class PricingPlanTest { + + @Test + @DisplayName("명시적으로 등록된 token type rate는 RESOLVED로 표현한다") + void resolveExplicitRate() { + PricingPlan plan = new PricingPlan( + "gpt-4o", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + + PricingResolution resolution = plan.resolveRate(TokenType.PROMPT); + + assertThat(resolution).isEqualTo(PricingResolution.RESOLVED); + assertThat(resolution.isResolved()).isTrue(); + } + + @Test + @DisplayName("명시적으로 등록된 0 rate는 RESOLVED 무료 가격으로 표현한다") + void resolveExplicitZeroRate() { + PricingPlan plan = new PricingPlan( + "free-model", + Map.of(TokenType.PROMPT, BigDecimal.ZERO), + Currency.getInstance("USD") + ); + + PricingResolution resolution = plan.resolveRate(TokenType.PROMPT); + + assertThat(resolution).isEqualTo(PricingResolution.RESOLVED); + assertThat(resolution.isResolved()).isTrue(); + } + + @Test + @DisplayName("등록된 plan에 필요한 token type rate가 없으면 MISSING_RATE로 표현한다") + void resolveMissingRate() { + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + + PricingResolution resolution = plan.resolveRate(TokenType.COMPLETION); + + assertThat(resolution).isEqualTo(PricingResolution.MISSING_RATE); + assertThat(resolution.isResolved()).isFalse(); + } + + @Test + @DisplayName("legacy getRate의 0 fallback과 resolveRate의 missing 표현은 구분된다") + void distinguishLegacyZeroFallbackFromMissingResolution() { + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + + assertThat(plan.getRate(TokenType.COMPLETION)).isEqualByComparingTo(BigDecimal.ZERO); + + PricingResolution resolution = plan.resolveRate(TokenType.COMPLETION); + assertThat(resolution).isEqualTo(PricingResolution.MISSING_RATE); + assertThat(resolution.isResolved()).isFalse(); + } + + @Test + @DisplayName("COMPLETION rate가 명시적으로 있으면 REASONING fallback은 RESOLVED다") + void resolveReasoningFallbackFromExplicitCompletionRate() { + PricingPlan plan = new PricingPlan( + "completion-model", + Map.of(TokenType.COMPLETION, new BigDecimal("0.03")), + Currency.getInstance("USD") + ); + + PricingResolution resolution = plan.resolveRate(TokenType.REASONING); + + assertThat(resolution).isEqualTo(PricingResolution.RESOLVED); + assertThat(resolution.isResolved()).isTrue(); + } + + @Test + @DisplayName("COMPLETION rate가 없으면 REASONING fallback은 MISSING_RATE다") + void missingReasoningFallbackWithoutCompletionRate() { + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + + PricingResolution resolution = plan.resolveRate(TokenType.REASONING); + + assertThat(resolution).isEqualTo(PricingResolution.MISSING_RATE); + assertThat(resolution.isResolved()).isFalse(); + } + + @Test + @DisplayName("PROMPT rate가 명시적으로 있으면 cache token type fallback은 RESOLVED다") + void resolveCacheFallbackFromExplicitPromptRate() { + PricingPlan plan = new PricingPlan( + "prompt-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + + PricingResolution readResolution = plan.resolveRate(TokenType.CACHE_READ_PROMPT); + PricingResolution creationResolution = plan.resolveRate(TokenType.CACHE_CREATION_PROMPT); + + assertThat(readResolution).isEqualTo(PricingResolution.RESOLVED); + assertThat(readResolution.isResolved()).isTrue(); + assertThat(creationResolution).isEqualTo(PricingResolution.RESOLVED); + assertThat(creationResolution.isResolved()).isTrue(); + } + + @Test + @DisplayName("PROMPT rate가 없으면 cache token type fallback은 MISSING_RATE다") + void missingCacheFallbackWithoutPromptRate() { + PricingPlan plan = new PricingPlan( + "completion-only-model", + Map.of(TokenType.COMPLETION, new BigDecimal("0.03")), + Currency.getInstance("USD") + ); + + PricingResolution readResolution = plan.resolveRate(TokenType.CACHE_READ_PROMPT); + PricingResolution creationResolution = plan.resolveRate(TokenType.CACHE_CREATION_PROMPT); + + assertThat(readResolution).isEqualTo(PricingResolution.MISSING_RATE); + assertThat(readResolution.isResolved()).isFalse(); + assertThat(creationResolution).isEqualTo(PricingResolution.MISSING_RATE); + assertThat(creationResolution.isResolved()).isFalse(); + } + + @Test + @DisplayName("fallback 기준 rate가 0으로 명시되어 있으면 RESOLVED다") + void resolveFallbackFromExplicitZeroBaseRate() { + PricingPlan plan = new PricingPlan( + "zero-fallback-model", + Map.of( + TokenType.PROMPT, BigDecimal.ZERO, + TokenType.COMPLETION, BigDecimal.ZERO + ), + Currency.getInstance("USD") + ); + + assertThat(plan.resolveRate(TokenType.REASONING)).isEqualTo(PricingResolution.RESOLVED); + assertThat(plan.resolveRate(TokenType.CACHE_READ_PROMPT)).isEqualTo(PricingResolution.RESOLVED); + assertThat(plan.resolveRate(TokenType.CACHE_CREATION_PROMPT)).isEqualTo(PricingResolution.RESOLVED); + } + + @Test + @DisplayName("빈 rates plan은 생성 가능하고 필요한 rate를 MISSING_RATE로 표현한다") + void emptyRatesResolveMissingRate() { + PricingPlan plan = new PricingPlan( + "empty-rates-model", + Map.of(), + Currency.getInstance("USD") + ); + + assertThat(plan.resolveRate(TokenType.PROMPT)).isEqualTo(PricingResolution.MISSING_RATE); + } +} diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/PricingResolutionTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/PricingResolutionTest.java new file mode 100644 index 0000000..b456c18 --- /dev/null +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/PricingResolutionTest.java @@ -0,0 +1,100 @@ +package io.tokenpilot.core.domain; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; + +class PricingResolutionTest { + + @Test + @DisplayName("PricingResolution은 public API에서 bounded 상태값을 제공한다") + void exposesPricingResolutionValues() { + assertThat(PricingResolution.values()) + .containsExactly( + PricingResolution.RESOLVED, + PricingResolution.MISSING_PLAN, + PricingResolution.MISSING_RATE, + PricingResolution.CURRENCY_MISMATCH + ); + } + + @Test + @DisplayName("RESOLVED 결과는 성공 상태로 표현된다") + void resolvedIsSuccessful() { + PricingResolution resolution = PricingResolution.RESOLVED; + + assertThat(resolution.isResolved()).isTrue(); + } + + @Test + @DisplayName("MISSING_PLAN 결과는 실패 상태로 표현된다") + void missingPlanIsNotResolved() { + PricingResolution resolution = PricingResolution.MISSING_PLAN; + + assertThat(resolution.isResolved()).isFalse(); + } + + @Test + @DisplayName("MISSING_RATE 결과는 실패 상태로 표현된다") + void missingRateIsNotResolved() { + PricingResolution resolution = PricingResolution.MISSING_RATE; + + assertThat(resolution.isResolved()).isFalse(); + } + + @Test + @DisplayName("CURRENCY_MISMATCH 결과는 실패 상태로 표현된다") + void currencyMismatchIsNotResolved() { + PricingResolution resolution = PricingResolution.CURRENCY_MISMATCH; + + assertThat(resolution.isResolved()).isFalse(); + } + + @Test + @DisplayName("PricingResolution은 별도 payload 없이 상태 자체로 표현된다") + void resolutionItselfIsState() { + assertThat(PricingResolution.RESOLVED.name()).isEqualTo("RESOLVED"); + } + + @Test + @DisplayName("PricingResolution 자체가 low-cardinality pricing miss reason이다") + void resolutionItselfIsLowCardinalityReason() { + assertThat(PricingResolution.values()) + .filteredOn(resolution -> !resolution.isResolved()) + .extracting(PricingResolution::name) + .containsExactly( + "MISSING_PLAN", + "MISSING_RATE", + "CURRENCY_MISMATCH" + ); + } + + @Test + @DisplayName("pricing miss reason은 model/tenant/user id를 포함하지 않는다") + void pricingMissReasonDoesNotContainHighCardinalityIdentifiers() { + assertThat(PricingResolution.values()) + .filteredOn(resolution -> !resolution.isResolved()) + .extracting(PricingResolution::name) + .allSatisfy(reason -> assertThat(reason) + .doesNotContain("gpt") + .doesNotContain("model") + .doesNotContain("tenant") + .doesNotContain("user") + .doesNotContain("policy")); + } + + @Test + @DisplayName("PricingResolution은 RECONCILIATION_REQUIRED 상태값을 갖지 않는다") + void reconciliationRequiredIsNotPricingResolution() { + assertThat(PricingResolution.values()) + .noneMatch(value -> value.name().equals("RECONCILIATION_REQUIRED")); + } + + @Test + @DisplayName("PricingResolution은 UNPRICED 상태값을 갖지 않는다") + void unpricedIsNotPricingResolution() { + assertThat(PricingResolution.values()) + .noneMatch(value -> value.name().equals("UNPRICED")); + } +} diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/domain/PricingSnapshotTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/PricingSnapshotTest.java new file mode 100644 index 0000000..0245eb2 --- /dev/null +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/domain/PricingSnapshotTest.java @@ -0,0 +1,79 @@ +package io.tokenpilot.core.domain; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import java.math.BigDecimal; +import java.time.Instant; +import java.util.Currency; +import java.util.EnumMap; +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +class PricingSnapshotTest { + + @Test + @DisplayName("요청 단위 pricing snapshot은 pricing 식별자와 적용 rate 정보를 보존해야 한다") + void preservePricingSnapshotValues() { + Instant checkedAt = Instant.parse("2026-07-30T00:00:00Z"); + Map rates = new EnumMap<>(TokenType.class); + rates.put(TokenType.PROMPT, new BigDecimal("0.01")); + rates.put(TokenType.COMPLETION, new BigDecimal("0.03")); + + PricingSnapshot snapshot = new PricingSnapshot( + "gpt-4o", + "standard", + "catalog-v1", + checkedAt, + rates, + Currency.getInstance("USD") + ); + + assertThat(snapshot.modelId()).isEqualTo("gpt-4o"); + assertThat(snapshot.pricingPolicyId()).isEqualTo("standard"); + assertThat(snapshot.catalogVersion()).isEqualTo("catalog-v1"); + assertThat(snapshot.checkedAt()).isEqualTo(checkedAt); + assertThat(snapshot.currency()).isEqualTo(Currency.getInstance("USD")); + assertThat(snapshot.rates()).containsEntry(TokenType.PROMPT, new BigDecimal("0.01")); + assertThat(snapshot.rates()).containsEntry(TokenType.COMPLETION, new BigDecimal("0.03")); + } + + @Test + @DisplayName("pricing snapshot rates는 생성 후 변경할 수 없어야 한다") + void ratesMustBeImmutable() { + Map rates = new EnumMap<>(TokenType.class); + rates.put(TokenType.PROMPT, new BigDecimal("0.01")); + + PricingSnapshot snapshot = new PricingSnapshot( + "gpt-4o", + "standard", + "catalog-v1", + Instant.parse("2026-07-30T00:00:00Z"), + rates, + Currency.getInstance("USD") + ); + + rates.put(TokenType.PROMPT, new BigDecimal("9.99")); + + assertThat(snapshot.rates()).containsEntry(TokenType.PROMPT, new BigDecimal("0.01")); + assertThatThrownBy(() -> snapshot.rates().put(TokenType.COMPLETION, new BigDecimal("0.03"))) + .isInstanceOf(UnsupportedOperationException.class); + } + + @Test + @DisplayName("pricing snapshot은 빈 rates를 보존할 수 있어야 한다") + void preserveEmptyRates() { + PricingSnapshot snapshot = new PricingSnapshot( + "gpt-4o", + "standard", + "catalog-v1", + Instant.parse("2026-07-30T00:00:00Z"), + Map.of(), + Currency.getInstance("USD") + ); + + assertThat(snapshot.rates()).isEmpty(); + } +} diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/exception/MissingPricingExceptionTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/exception/MissingPricingExceptionTest.java new file mode 100644 index 0000000..e6a867f --- /dev/null +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/exception/MissingPricingExceptionTest.java @@ -0,0 +1,19 @@ +package io.tokenpilot.core.exception; + +import io.tokenpilot.core.domain.PricingResolution; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; + +class MissingPricingExceptionTest { + + @Test + @DisplayName("MissingPricingException은 PricingResolution을 구조화된 값으로 보존한다") + void preservesPricingResolution() { + MissingPricingException exception = new MissingPricingException(PricingResolution.MISSING_PLAN); + + assertThat(exception).hasMessage("MISSING_PLAN"); + assertThat(exception.getResolution()).isEqualTo(PricingResolution.MISSING_PLAN); + } +} 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 4fbe16f..34489fe 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,10 +2,12 @@ import io.tokenpilot.core.domain.Cost; import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingResolution; 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 io.tokenpilot.core.exception.MissingPricingException; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; @@ -14,6 +16,7 @@ import java.util.Map; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; class DefaultCostCalculatorTest { @@ -90,4 +93,35 @@ void calculateOneThousandTokensAtRate() { assertThat(cost.value()).isEqualByComparingTo("0.0004"); } + + @Test + @DisplayName("실제 completion 사용량에 필요한 rate가 없으면 MISSING_RATE여야 한다") + void failWhenActualCompletionRateIsMissing() { + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + TokenUsage usage = TokenUsage.from(1_000, 1_000); + + assertThatThrownBy(() -> calculator.calculate(usage, plan)) + .isInstanceOf(MissingPricingException.class) + .extracting(exception -> ((MissingPricingException) exception).getResolution()) + .isEqualTo(PricingResolution.MISSING_RATE); + } + + @Test + @DisplayName("실제 사용량이 없는 token type의 누락 rate는 계산을 실패시키지 않아야 한다") + void doNotRequireRateWhenActualUsageIsZero() { + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + TokenUsage usage = TokenUsage.from(1_000, 0); + + Cost cost = calculator.calculate(usage, plan); + + assertThat(cost.value()).isEqualByComparingTo("0.01"); + } } diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultLedgerManagerTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultLedgerManagerTest.java index 2773979..ab7d258 100644 --- a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultLedgerManagerTest.java +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultLedgerManagerTest.java @@ -8,13 +8,14 @@ import org.mockito.Mockito; import java.math.BigDecimal; +import java.time.Instant; import java.util.Currency; import java.util.List; import java.util.Map; import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.ArgumentMatchers.*; -import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.*; class DefaultLedgerManagerTest { @@ -58,4 +59,77 @@ void shouldReturnZeroCostWhenPlanIsMissing() { assertThat(result.currency()).isEqualTo(Currency.getInstance("USD")); verify(listener).onRecord(any(CostRecordedEvent.class)); } + + @Test + @DisplayName("이미 resolve된 plan으로 기록하면 registry를 다시 조회하지 않아야 한다") + void shouldRecordWithResolvedPlanWithoutRegistryLookup() { + PricingRegistry pricingRegistry = Mockito.mock(PricingRegistry.class); + CostCalculator costCalculator = Mockito.mock(CostCalculator.class); + LedgerListener listener = Mockito.mock(LedgerListener.class); + DefaultLedgerManager manager = new DefaultLedgerManager(pricingRegistry, costCalculator, List.of(listener)); + PricingPlan plan = new PricingPlan("gpt-4o", new BigDecimal("5.0"), new BigDecimal("15.0")); + TokenUsage usage = TokenUsage.from(1000, 1000); + Cost expectedCost = new Cost(new BigDecimal("20.000000"), Currency.getInstance("USD")); + + when(costCalculator.calculate(usage, plan)).thenReturn(expectedCost); + + Cost cost = manager.record(plan, usage, Map.of()); + + assertThat(cost).isEqualTo(expectedCost); + verifyNoInteractions(pricingRegistry); + verify(listener).onRecord(argThat(event -> + event.modelId().equals("gpt-4o") && + event.usage().equals(usage) && + event.cost().equals(expectedCost) + )); + } + + @Test + @DisplayName("pricing snapshot으로 기록하면 registry 변경 후에도 snapshot rate를 사용해야 한다") + void shouldRecordWithSnapshotRatesAfterRegistryChanges() { + PricingPlan originalPlan = new PricingPlan("gpt-4o", new BigDecimal("5.0"), new BigDecimal("15.0")); + PricingSnapshot snapshot = PricingSnapshot.from( + originalPlan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + registry.registerPlan(new PricingPlan("gpt-4o", new BigDecimal("50.0"), new BigDecimal("150.0"))); + TokenUsage usage = TokenUsage.from(1000, 1000); + + Cost cost = manager.record(snapshot, usage, Map.of()); + + assertThat(cost.value()).isEqualByComparingTo("20.000000"); + verify(listener).onRecord(argThat(event -> + event.modelId().equals("gpt-4o") && + event.usage().equals(usage) && + event.cost().equals(cost) + )); + } + + @Test + @DisplayName("명시적 0 rate snapshot은 정상 0원 cost로 기록되어야 한다") + void shouldRecordZeroCostWithExplicitZeroRateSnapshot() { + PricingPlan freePlan = new PricingPlan( + "free-model", + BigDecimal.ZERO, + BigDecimal.ZERO, + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + freePlan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + TokenUsage usage = TokenUsage.from(1000, 1000); + + Cost cost = manager.record(snapshot, usage, Map.of()); + + assertThat(cost.value()).isEqualByComparingTo(BigDecimal.ZERO); + assertThat(cost.currency()).isEqualTo(Currency.getInstance("USD")); + verify(listener).onRecord(argThat(event -> + event.modelId().equals("free-model") && + event.usage().equals(usage) && + event.cost().equals(cost) + )); + } } diff --git a/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultPricingEvaluatorTest.java b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultPricingEvaluatorTest.java new file mode 100644 index 0000000..a7db95b --- /dev/null +++ b/token-pilot-core/src/test/java/io/tokenpilot/core/internal/DefaultPricingEvaluatorTest.java @@ -0,0 +1,113 @@ +package io.tokenpilot.core.internal; + +import io.tokenpilot.core.PricingEvaluator; +import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingReconciliationResult; +import io.tokenpilot.core.domain.PricingResolution; +import io.tokenpilot.core.domain.PricingSnapshot; +import io.tokenpilot.core.domain.TokenType; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import java.math.BigDecimal; +import java.time.Instant; +import java.util.Currency; +import java.util.Map; +import java.util.Optional; + +import static org.assertj.core.api.Assertions.assertThat; + +class DefaultPricingEvaluatorTest { + + private final PricingEvaluator evaluator = LedgerComponents.defaultPricingEvaluator(); + + @Test + @DisplayName("snapshot이 없으면 MISSING_PLAN이어야 한다") + void resolveMissingSnapshotAsMissingPlan() { + PricingResolution resolution = evaluator.validateSnapshotRates(Optional.empty()); + + assertThat(resolution).isEqualTo(PricingResolution.MISSING_PLAN); + } + + @Test + @DisplayName("prompt와 completion rate가 있으면 snapshot rate 검증에 성공해야 한다") + void validateRequiredSnapshotRates() { + PricingSnapshot snapshot = snapshot(Map.of( + TokenType.PROMPT, new BigDecimal("0.01"), + TokenType.COMPLETION, new BigDecimal("0.03") + )); + + PricingResolution resolution = evaluator.validateSnapshotRates(Optional.of(snapshot)); + + assertThat(resolution).isEqualTo(PricingResolution.RESOLVED); + } + + @Test + @DisplayName("completion rate가 없으면 snapshot rate 검증은 MISSING_RATE여야 한다") + void rejectSnapshotWithoutCompletionRate() { + PricingSnapshot snapshot = snapshot(Map.of( + TokenType.PROMPT, new BigDecimal("0.01") + )); + + PricingResolution resolution = evaluator.validateSnapshotRates(Optional.of(snapshot)); + + assertThat(resolution).isEqualTo(PricingResolution.MISSING_RATE); + } + + @Test + @DisplayName("snapshot model과 actual model이 같으면 RECONCILED여야 한다") + void reconcileMatchingActualModel() { + PricingSnapshot snapshot = snapshot(Map.of( + TokenType.PROMPT, new BigDecimal("0.01"), + TokenType.COMPLETION, new BigDecimal("0.03") + )); + + PricingReconciliationResult result = evaluator.determineReconciliation( + Optional.of(snapshot), + "gpt-4o" + ); + + assertThat(result).isEqualTo(PricingReconciliationResult.RECONCILED); + } + + @Test + @DisplayName("snapshot model과 actual model이 다르면 RECONCILIATION_REQUIRED여야 한다") + void requireReconciliationForDifferentActualModel() { + PricingSnapshot snapshot = snapshot(Map.of( + TokenType.PROMPT, new BigDecimal("0.01"), + TokenType.COMPLETION, new BigDecimal("0.03") + )); + + PricingReconciliationResult result = evaluator.determineReconciliation( + Optional.of(snapshot), + "gpt-4o-mini" + ); + + assertThat(result).isEqualTo(PricingReconciliationResult.RECONCILIATION_REQUIRED); + } + + @Test + @DisplayName("snapshot이 없으면 reconciliation 결과는 UNPRICED여야 한다") + void leaveMissingSnapshotUnpriced() { + PricingReconciliationResult result = evaluator.determineReconciliation( + Optional.empty(), + "gpt-4o" + ); + + assertThat(result).isEqualTo(PricingReconciliationResult.UNPRICED); + } + + private static PricingSnapshot snapshot(Map rates) { + PricingPlan plan = new PricingPlan( + "gpt-4o", + "standard", + rates, + Currency.getInstance("USD") + ); + return PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + } +} 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 8c27125..c969c74 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,14 +1,19 @@ package io.tokenpilot.core.internal; import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingResolution; +import io.tokenpilot.core.domain.PricingSnapshot; +import io.tokenpilot.core.domain.TokenType; 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 java.util.Optional; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; class InMemoryPricingRegistryTest { @@ -17,20 +22,138 @@ class InMemoryPricingRegistryTest { @Test @DisplayName("가격 정책을 등록하고 모델 ID로 조회할 수 있어야 한다") void shouldRegisterAndGetPlan() { - // Given - PricingPlan plan = new PricingPlan("claude-3", + String modelId = "claude-3-5-sonnet-20241022"; + PricingPlan plan = new PricingPlan(modelId, new BigDecimal("0.015"), new BigDecimal("0.075"), Currency.getInstance("USD")); - // When registry.registerPlan(plan); - Optional retrieved = registry.getPlan("claude-3"); + Optional retrieved = registry.getPlan(modelId); - // Then assertThat(retrieved).isPresent(); - assertThat(retrieved.get().modelId()).isEqualTo("claude-3"); + assertThat(retrieved.get().modelId()).isEqualTo(modelId); assertThat(retrieved.get().promptPricePerK()).isEqualByComparingTo("0.015"); } + @Test + @DisplayName("모델 ID와 pricing policy ID로 가격 정책을 조회할 수 있어야 한다") + void shouldRegisterAndGetPlanByModelIdAndPricingPolicyId() { + String modelId = "gpt-4o-2024-08-06"; + String pricingPolicyId = "openai-gpt-4o-2024-08-06-standard"; + PricingPlan plan = new PricingPlan( + modelId, + pricingPolicyId, + Map.of(TokenType.PROMPT, new BigDecimal("0.0025")), + Currency.getInstance("USD") + ); + + registry.registerPlan(plan); + Optional retrieved = registry.getPlan(modelId, pricingPolicyId); + + assertThat(retrieved).isPresent(); + assertThat(retrieved.get().modelId()).isEqualTo(modelId); + assertThat(retrieved.get().pricingPolicyId()).isEqualTo(pricingPolicyId); + } + + @Test + @DisplayName("동일 모델 ID라도 pricing policy ID가 다르면 다른 가격 정책으로 조회되어야 한다") + void shouldDistinguishPlansByPricingPolicyId() { + String modelId = "gpt-4o-2024-08-06"; + PricingPlan standardPlan = new PricingPlan( + modelId, + "standard", + Map.of(TokenType.PROMPT, new BigDecimal("0.0025")), + Currency.getInstance("USD") + ); + PricingPlan discountedPlan = new PricingPlan( + modelId, + "discounted", + Map.of(TokenType.PROMPT, new BigDecimal("0.0010")), + Currency.getInstance("USD") + ); + + registry.registerPlan(standardPlan); + registry.registerPlan(discountedPlan); + + assertThat(registry.getPlan(modelId, "standard")).contains(standardPlan); + assertThat(registry.getPlan(modelId, "discounted")).contains(discountedPlan); + } + + @Test + @DisplayName("모델 ID와 pricing policy ID로 요청 단위 pricing snapshot을 조회할 수 있어야 한다") + void shouldGetSnapshotByModelIdAndPricingPolicyId() { + String modelId = "gpt-4o-2024-08-06"; + String pricingPolicyId = "standard"; + PricingPlan plan = new PricingPlan( + modelId, + pricingPolicyId, + Map.of(TokenType.PROMPT, new BigDecimal("0.0025")), + Currency.getInstance("USD") + ); + + registry.registerPlan(plan); + + Optional snapshot = registry.resolveSnapshot(modelId, pricingPolicyId); + + assertThat(snapshot).isPresent(); + assertThat(snapshot.get().modelId()).isEqualTo(modelId); + assertThat(snapshot.get().pricingPolicyId()).isEqualTo(pricingPolicyId); + assertThat(snapshot.get().catalogVersion()).isEqualTo(PricingSnapshot.DEFAULT_CATALOG_VERSION); + assertThat(snapshot.get().checkedAt()).isNotNull(); + assertThat(snapshot.get().currency()).isEqualTo(Currency.getInstance("USD")); + assertThat(snapshot.get().rates()).containsEntry(TokenType.PROMPT, new BigDecimal("0.0025")); + } + + @Test + @DisplayName("#32가 같은 model id로 resolve한 입력은 동일 pricing policy snapshot을 사용해야 한다") + void shouldUseSameSnapshotWhenResolvedModelIdIsSame() { + String modelId = "gpt-4o-2024-08-06"; + String aliasResolvedModelId = modelId; + String canonicalModelId = modelId; + String pricingPolicyId = "standard"; + PricingPlan plan = new PricingPlan( + modelId, + pricingPolicyId, + Map.of(TokenType.PROMPT, new BigDecimal("0.0025")), + Currency.getInstance("USD") + ); + registry.registerPlan(plan); + + PricingSnapshot aliasSnapshot = registry.resolveSnapshot(aliasResolvedModelId, pricingPolicyId).orElseThrow(); + PricingSnapshot canonicalSnapshot = registry.resolveSnapshot(canonicalModelId, pricingPolicyId).orElseThrow(); + + assertThat(aliasSnapshot.modelId()).isEqualTo(canonicalSnapshot.modelId()); + assertThat(aliasSnapshot.pricingPolicyId()).isEqualTo(canonicalSnapshot.pricingPolicyId()); + assertThat(aliasSnapshot.catalogVersion()).isEqualTo(canonicalSnapshot.catalogVersion()); + assertThat(aliasSnapshot.currency()).isEqualTo(canonicalSnapshot.currency()); + assertThat(aliasSnapshot.rates()).containsAllEntriesOf(canonicalSnapshot.rates()); + } + + @Test + @DisplayName("registry 변경 후 새 요청은 변경된 pricing plan으로 snapshot을 resolve해야 한다") + void shouldResolveNewSnapshotAfterRegistryChanges() { + String modelId = "gpt-4o-2024-08-06"; + String pricingPolicyId = "standard"; + registry.registerPlan(new PricingPlan( + modelId, + pricingPolicyId, + Map.of(TokenType.PROMPT, new BigDecimal("0.0025")), + Currency.getInstance("USD") + )); + PricingSnapshot firstSnapshot = registry.resolveSnapshot(modelId, pricingPolicyId).orElseThrow(); + + registry.registerPlan(new PricingPlan( + modelId, + pricingPolicyId, + Map.of(TokenType.PROMPT, new BigDecimal("0.0050")), + Currency.getInstance("USD") + )); + + PricingSnapshot nextSnapshot = registry.resolveSnapshot(modelId, pricingPolicyId).orElseThrow(); + + assertThat(firstSnapshot.rates()).containsEntry(TokenType.PROMPT, new BigDecimal("0.0025")); + assertThat(nextSnapshot.rates()).containsEntry(TokenType.PROMPT, new BigDecimal("0.0050")); + } + @Test @DisplayName("등록되지 않은 모델 조회 시 빈 Optional을 반환해야 한다") void shouldReturnEmptyWhenNotFound() { @@ -40,4 +163,120 @@ void shouldReturnEmptyWhenNotFound() { // Then assertThat(retrieved).isEmpty(); } -} \ No newline at end of file + + @Test + @DisplayName("등록되지 않은 모델의 가격 결정 결과는 MISSING_PLAN이어야 한다") + void shouldResolveMissingPlanWhenModelIsNotRegistered() { + PricingResolution resolution = registry.resolveRate("non-existent", TokenType.PROMPT); + + assertThat(resolution).isEqualTo(PricingResolution.MISSING_PLAN); + assertThat(resolution.isResolved()).isFalse(); + } + + @Test + @DisplayName("등록된 모델의 가격 결정 결과는 PricingPlan resolveRate 결과를 따라야 한다") + void shouldDelegateRateResolutionToRegisteredPlan() { + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.015")), + Currency.getInstance("USD") + ); + registry.registerPlan(plan); + + assertThat(registry.resolveRate("prompt-only-model", TokenType.PROMPT)) + .isEqualTo(PricingResolution.RESOLVED); + assertThat(registry.resolveRate("prompt-only-model", TokenType.COMPLETION)) + .isEqualTo(PricingResolution.MISSING_RATE); + } + + @Test + @DisplayName("기대 통화와 plan 통화가 다르면 CURRENCY_MISMATCH여야 한다") + void shouldResolveCurrencyMismatchWhenExpectedCurrencyDiffers() { + PricingPlan plan = new PricingPlan( + "usd-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.015")), + Currency.getInstance("USD") + ); + registry.registerPlan(plan); + + PricingResolution resolution = registry.resolveRate( + "usd-model", + TokenType.PROMPT, + Currency.getInstance("KRW") + ); + + assertThat(resolution).isEqualTo(PricingResolution.CURRENCY_MISMATCH); + assertThat(resolution.isResolved()).isFalse(); + } + + @Test + @DisplayName("currency mismatch는 등록된 pricing 상태를 변경하지 않아야 한다") + void shouldNotChangePricingStateWhenCurrencyMismatches() { + PricingPlan plan = new PricingPlan( + "usd-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.015")), + Currency.getInstance("USD") + ); + registry.registerPlan(plan); + + PricingResolution resolution = registry.resolveRate( + "usd-model", + TokenType.PROMPT, + Currency.getInstance("KRW") + ); + + assertThat(resolution).isEqualTo(PricingResolution.CURRENCY_MISMATCH); + assertThat(registry.getPlan("usd-model")).contains(plan); + assertThat(registry.resolveSnapshot("usd-model", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .isPresent() + .get() + .satisfies(snapshot -> { + assertThat(snapshot.modelId()).isEqualTo("usd-model"); + assertThat(snapshot.pricingPolicyId()).isEqualTo(PricingPlan.DEFAULT_PRICING_POLICY_ID); + assertThat(snapshot.currency()).isEqualTo(Currency.getInstance("USD")); + assertThat(snapshot.rates()).containsEntry(TokenType.PROMPT, new BigDecimal("0.015")); + }); + assertThat(registry.resolveRate("usd-model", TokenType.PROMPT, Currency.getInstance("USD"))) + .isEqualTo(PricingResolution.RESOLVED); + } + + @Test + @DisplayName("기대 통화와 plan 통화가 같으면 일반 rate resolution을 수행해야 한다") + void shouldDelegateRateResolutionWhenExpectedCurrencyMatches() { + PricingPlan plan = new PricingPlan( + "usd-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.015")), + Currency.getInstance("USD") + ); + registry.registerPlan(plan); + + assertThat(registry.resolveRate("usd-model", TokenType.PROMPT, Currency.getInstance("USD"))) + .isEqualTo(PricingResolution.RESOLVED); + assertThat(registry.resolveRate("usd-model", TokenType.COMPLETION, Currency.getInstance("USD"))) + .isEqualTo(PricingResolution.MISSING_RATE); + } + + @Test + @DisplayName("기대 통화가 없는 경로는 통화 검사를 수행하지 않고 일반 rate resolution을 수행해야 한다") + void shouldSkipCurrencyCheckWhenExpectedCurrencyIsNotProvided() { + PricingPlan plan = new PricingPlan( + "krw-model", + Map.of(TokenType.PROMPT, new BigDecimal("15")), + Currency.getInstance("KRW") + ); + registry.registerPlan(plan); + + assertThat(registry.resolveRate("krw-model", TokenType.PROMPT)) + .isEqualTo(PricingResolution.RESOLVED); + assertThat(registry.resolveRate("krw-model", TokenType.COMPLETION)) + .isEqualTo(PricingResolution.MISSING_RATE); + } + + @Test + @DisplayName("기대 통화가 있는 경로는 null expectedCurrency를 허용하지 않는다") + void shouldRejectNullExpectedCurrency() { + assertThatThrownBy(() -> registry.resolveRate("any-model", TokenType.PROMPT, null)) + .isInstanceOf(NullPointerException.class) + .hasMessage("expectedCurrency must not be null"); + } +} diff --git a/token-pilot-sample-app/src/test/java/io/tokenpilot/sample/SampleApplicationChatClientE2ETest.java b/token-pilot-sample-app/src/test/java/io/tokenpilot/sample/SampleApplicationChatClientE2ETest.java index 1d8a110..2298e39 100644 --- a/token-pilot-sample-app/src/test/java/io/tokenpilot/sample/SampleApplicationChatClientE2ETest.java +++ b/token-pilot-sample-app/src/test/java/io/tokenpilot/sample/SampleApplicationChatClientE2ETest.java @@ -1,15 +1,26 @@ package io.tokenpilot.sample; +import io.tokenpilot.budget.BudgetDecision; +import io.tokenpilot.budget.BudgetStateStore; +import io.tokenpilot.core.domain.Cost; +import io.tokenpilot.core.domain.PricingPlan; +import io.tokenpilot.core.domain.PricingReconciliationResult; +import io.tokenpilot.core.domain.PricingResolution; +import io.tokenpilot.core.domain.PricingSnapshot; import org.junit.jupiter.api.Test; import org.springframework.ai.chat.client.ChatClient; import org.springframework.ai.chat.client.ChatClientBuilderCustomizer; +import org.springframework.ai.chat.client.ChatClientResponse; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.metadata.ChatResponseMetadata; import org.springframework.ai.chat.metadata.DefaultUsage; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.model.Generation; +import org.springframework.ai.chat.prompt.ChatOptions; +import org.springframework.ai.chat.prompt.Prompt; import org.springframework.beans.factory.ObjectProvider; +import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.boot.test.context.TestConfiguration; import org.springframework.boot.test.web.server.LocalServerPort; @@ -22,6 +33,7 @@ import java.net.http.HttpClient; import java.net.http.HttpRequest; import java.net.http.HttpResponse; +import java.util.Currency; import java.util.List; import java.util.Map; @@ -37,6 +49,10 @@ "token-pilot.pricing.plans[0].rates.COMPLETION=0.00060", "token-pilot.metrics.enabled=true", "token-pilot.metrics.tag-whitelist[0]=tenant_id", + "token-pilot.budget.enabled=true", + "token-pilot.budget.monthly-limit=10.00", + "token-pilot.budget.currency=USD", + "token-pilot.budget.target-tag-key=tenant_id", "management.endpoints.web.exposure.include=prometheus,health" } ) @@ -47,6 +63,12 @@ class SampleApplicationChatClientE2ETest { @LocalServerPort private int port; + @Autowired + private ChatClient.Builder chatClientBuilder; + + @Autowired + private BudgetStateStore budgetStateStore; + @Test void chatClientAdvisorRecordsTokenPilotMetricsEndToEnd() throws Exception { HttpResponse beans = get("/test/token-pilot/beans"); @@ -71,6 +93,45 @@ void chatClientAdvisorRecordsTokenPilotMetricsEndToEnd() throws Exception { .doesNotContain("user_id=\"chat-sample-user\""); } + @Test + void budgetAdvisorResolvesModelAndPolicyFromRegularChatClientCall() { + ChatClientResponse response = chatClientBuilder.clone() + .build() + .prompt() + .user("Record this fake budget-aware Spring AI call.") + .advisors(advisors -> advisors.param("tenant_id", "budget-chat-tenant")) + .call() + .chatClientResponse(); + + PricingSnapshot snapshot = contextValue(response, PricingSnapshot.class); + PricingResolution resolution = contextValue(response, PricingResolution.class); + PricingReconciliationResult reconciliationResult = contextValue( + response, + PricingReconciliationResult.class + ); + BudgetDecision budgetDecision = contextValue(response, BudgetDecision.class); + Cost accumulatedCost = budgetStateStore.getAccumulatedCost( + budgetDecision.key(), + budgetDecision.limit() + ); + + assertThat(snapshot.modelId()).isEqualTo("fake-chat-model"); + assertThat(snapshot.pricingPolicyId()).isEqualTo(PricingPlan.DEFAULT_PRICING_POLICY_ID); + assertThat(snapshot.currency()).isEqualTo(Currency.getInstance("USD")); + assertThat(resolution).isEqualTo(PricingResolution.RESOLVED); + assertThat(reconciliationResult).isEqualTo(PricingReconciliationResult.RECONCILED); + assertThat(accumulatedCost.value()).isEqualByComparingTo("0.00135"); + assertThat(accumulatedCost.currency()).isEqualTo(Currency.getInstance("USD")); + } + + private T contextValue(ChatClientResponse response, Class type) { + return response.context().values().stream() + .filter(type::isInstance) + .map(type::cast) + .findFirst() + .orElseThrow(); + } + private HttpResponse get(String path) throws IOException, InterruptedException { HttpRequest request = HttpRequest.newBuilder() .uri(URI.create("http://localhost:" + port + path)) @@ -84,13 +145,25 @@ static class FakeChatClientConfiguration { @Bean ChatModel fakeChatModel() { - return prompt -> new ChatResponse( - List.of(new Generation(new AssistantMessage("fake chat response"))), - ChatResponseMetadata.builder() + return new ChatModel() { + @Override + public ChatResponse call(Prompt prompt) { + return new ChatResponse( + List.of(new Generation(new AssistantMessage("fake chat response"))), + ChatResponseMetadata.builder() + .model("fake-chat-model") + .usage(new DefaultUsage(1_000, 2_000)) + .build() + ); + } + + @Override + public ChatOptions getOptions() { + return ChatOptions.builder() .model("fake-chat-model") - .usage(new DefaultUsage(1_000, 2_000)) - .build() - ); + .build(); + } + }; } @Bean diff --git a/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultLedgerAdvisor.java b/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultLedgerAdvisor.java index 016c629..441d574 100644 --- a/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultLedgerAdvisor.java +++ b/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/DefaultLedgerAdvisor.java @@ -6,15 +6,20 @@ import io.tokenpilot.budget.exception.BudgetExceededException; import io.tokenpilot.core.*; import io.tokenpilot.core.domain.*; +import io.tokenpilot.core.exception.MissingPricingException; +import io.tokenpilot.core.internal.LedgerComponents; import io.tokenpilot.springai.LedgerAdvisor; import io.tokenpilot.springai.UsageExtractor; import org.springframework.ai.chat.client.ChatClientRequest; import org.springframework.ai.chat.client.ChatClientResponse; import org.springframework.ai.chat.client.advisor.api.AdvisorChain; +import org.springframework.ai.chat.model.ChatResponse; import java.util.HashMap; import java.util.Map; +import java.util.Objects; import java.util.Optional; +import org.springframework.ai.chat.prompt.ChatOptions; /** * 기본 {@link LedgerAdvisor} 구현체. @@ -30,6 +35,11 @@ public class DefaultLedgerAdvisor implements LedgerAdvisor { static final String BUDGET_DECISION_CONTEXT = "tokenpilot.budget.decision"; + static final String MODEL_ID_CONTEXT = "tokenpilot.model.id"; + static final String PRICING_POLICY_ID_CONTEXT = "tokenpilot.pricing.policy.id"; + static final String PRICING_SNAPSHOT_CONTEXT = "tokenpilot.pricing.snapshot"; + static final String PRICING_RESOLUTION_CONTEXT = "tokenpilot.pricing.resolution"; + static final String PRICING_RECONCILIATION_RESULT_CONTEXT = "tokenpilot.pricing.reconciliation.result"; private final LedgerManager ledgerManager; private final UsageExtractor usageExtractor; @@ -37,6 +47,8 @@ public class DefaultLedgerAdvisor implements LedgerAdvisor { private final BudgetStateStore budgetStateStore; private final CostCalculator costCalculator; private final PricingRegistry pricingRegistry; + private final PricingEvaluator pricingEvaluator; + private final MissingPricingPolicy missingPricingPolicy; public DefaultLedgerAdvisor(LedgerManager ledgerManager, UsageExtractor usageExtractor) { this(ledgerManager, usageExtractor, null, null, null, null); @@ -45,53 +57,221 @@ public DefaultLedgerAdvisor(LedgerManager ledgerManager, UsageExtractor usageExt public DefaultLedgerAdvisor(LedgerManager ledgerManager, UsageExtractor usageExtractor, BudgetEvaluator budgetEvaluator, BudgetStateStore budgetStateStore, CostCalculator costCalculator, PricingRegistry pricingRegistry) { + this( + ledgerManager, + usageExtractor, + budgetEvaluator, + budgetStateStore, + costCalculator, + pricingRegistry, + MissingPricingPolicy.FAIL_OPEN + ); + } + + public DefaultLedgerAdvisor(LedgerManager ledgerManager, UsageExtractor usageExtractor, + BudgetEvaluator budgetEvaluator, BudgetStateStore budgetStateStore, + CostCalculator costCalculator, PricingRegistry pricingRegistry, + MissingPricingPolicy missingPricingPolicy) { + this( + ledgerManager, + usageExtractor, + budgetEvaluator, + budgetStateStore, + costCalculator, + pricingRegistry, + LedgerComponents.defaultPricingEvaluator(), + missingPricingPolicy + ); + } + + public DefaultLedgerAdvisor(LedgerManager ledgerManager, UsageExtractor usageExtractor, + BudgetEvaluator budgetEvaluator, BudgetStateStore budgetStateStore, + CostCalculator costCalculator, PricingRegistry pricingRegistry, + PricingEvaluator pricingEvaluator, + MissingPricingPolicy missingPricingPolicy) { this.ledgerManager = ledgerManager; this.usageExtractor = usageExtractor; this.budgetEvaluator = budgetEvaluator; this.budgetStateStore = budgetStateStore; this.costCalculator = costCalculator; this.pricingRegistry = pricingRegistry; + this.pricingEvaluator = Objects.requireNonNull( + pricingEvaluator, + "pricingEvaluator must not be null" + ); + this.missingPricingPolicy = Objects.requireNonNull( + missingPricingPolicy, + "missingPricingPolicy must not be null" + ); } @Override public ChatClientRequest before(ChatClientRequest request, AdvisorChain chain) { + ChatClientRequest resolvedRequest = request; + if (budgetEvaluator != null) { - Map tags = extractTagsFromRequest(request); + Map tags = extractTags(request.context()); BudgetDecision decision = budgetEvaluator.evaluate(tags); enforceExistingBlock(decision); - return request.mutate() + resolvedRequest = resolvedRequest.mutate() .context(BUDGET_DECISION_CONTEXT, decision) .build(); } - return request; + + if (pricingRegistry != null) { + resolvedRequest = resolvePricing(resolvedRequest); + } + + return resolvedRequest; } @Override public ChatClientResponse after(ChatClientResponse response, AdvisorChain chain) { TokenUsage usage = usageExtractor.extract(response); - + String modelId = extractModelId(response); + String responseModelId = extractResponseModelId(response); Map tags = extractTags(response); - ledgerManager.record(modelId, usage, tags); + Optional snapshot = extractPricingSnapshot(response); + boolean hasPricingResolution = hasPricingResolution(response); + if (snapshot.isPresent() || hasPricingResolution) { + PricingReconciliationResult reconciliationResult = pricingEvaluator.determineReconciliation( + snapshot, + responseModelId + ); + if (reconciliationResult != PricingReconciliationResult.RECONCILED) { + return withReconciliationResult(response, reconciliationResult); + } - // 예산 누적 처리 - if (budgetStateStore != null && costCalculator != null && pricingRegistry != null) { - Optional plan = pricingRegistry.getPlan(modelId); - if (plan.isPresent()) { - Cost cost = costCalculator.calculate(usage, plan.get()); - BudgetDecision decision = extractBudgetDecision(response); + PricingSnapshot resolvedSnapshot = snapshot.orElseThrow( + () -> new IllegalStateException("Reconciled pricing snapshot is missing") + ); + Cost cost; + try { + cost = ledgerManager.record(resolvedSnapshot, usage, tags); + } catch (MissingPricingException exception) { + return handleActualPricingFailure(response, exception); + } + ChatClientResponse reconciledResponse = withReconciliationResult( + response, + reconciliationResult + ); + if (budgetStateStore != null) { + BudgetDecision decision = extractBudgetDecision(reconciledResponse); budgetStateStore.addCost( decision.key(), decision.limit(), cost ); } + return reconciledResponse; + } else { + ledgerManager.record(modelId, usage, tags); + recordLegacyBudgetCost(modelId, usage, response); } return response; } + private ChatClientResponse handleActualPricingFailure( + ChatClientResponse response, + MissingPricingException exception + ) { + if (missingPricingPolicy == MissingPricingPolicy.FAIL_CLOSED) { + throw exception; + } + + Map context = copyContext(response); + context.put(PRICING_RESOLUTION_CONTEXT, exception.getResolution()); + context.put(PRICING_RECONCILIATION_RESULT_CONTEXT, PricingReconciliationResult.UNPRICED); + return new ChatClientResponse(response.chatResponse(), context); + } + + private ChatClientResponse withReconciliationResult( + ChatClientResponse response, + PricingReconciliationResult result + ) { + Map context = copyContext(response); + context.put(PRICING_RECONCILIATION_RESULT_CONTEXT, result); + return new ChatClientResponse(response.chatResponse(), context); + } + + private void recordLegacyBudgetCost(String modelId, TokenUsage usage, ChatClientResponse response) { + if (budgetStateStore == null || costCalculator == null || pricingRegistry == null) { + return; + } + + Optional plan = pricingRegistry.getPlan(modelId); + if (plan.isEmpty()) { + return; + } + + Cost cost = costCalculator.calculate(usage, plan.get()); + BudgetDecision decision = extractBudgetDecision(response); + budgetStateStore.addCost( + decision.key(), + decision.limit(), + cost + ); + } + + private ChatClientRequest resolvePricing(ChatClientRequest request) { + String modelId = extractModelId(request); + String pricingPolicyId = extractPricingPolicyId(request); + Optional snapshot = modelId == null + ? Optional.empty() + : pricingRegistry.resolveSnapshot(modelId, pricingPolicyId); + PricingResolution resolution = pricingEvaluator.validateSnapshotRates(snapshot); + rejectMissingPricingIfFailClosed(resolution); + + return withPricingContext(request, pricingPolicyId, resolution, snapshot); + } + + private ChatClientRequest withPricingContext( + ChatClientRequest request, + String pricingPolicyId, + PricingResolution resolution, + Optional snapshot + ) { + ChatClientRequest.Builder builder = request.mutate() + .context(PRICING_POLICY_ID_CONTEXT, pricingPolicyId) + .context(PRICING_RESOLUTION_CONTEXT, resolution); + if (resolution.isResolved()) { + snapshot.ifPresent(value -> builder.context(PRICING_SNAPSHOT_CONTEXT, value)); + } + return builder.build(); + } + + private void rejectMissingPricingIfFailClosed(PricingResolution resolution) { + if (missingPricingPolicy != MissingPricingPolicy.FAIL_CLOSED) { + return; + } + if (resolution.isResolved()) { + return; + } + throw new MissingPricingException(resolution); + } + + private String extractModelId(ChatClientRequest request) { + Object contextValue = request.context().get(MODEL_ID_CONTEXT); + if (contextValue instanceof String modelId && !modelId.isBlank()) { + return modelId; + } + + ChatOptions options = request.prompt().getOptions(); + if (options == null) { + return null; + } + + String modelId = options.getModel(); + if (modelId == null || modelId.isBlank()) { + return null; + } + + return modelId; + } + private void enforceExistingBlock(BudgetDecision decision) { switch (decision.state()) { case ALLOW, WARN -> { @@ -105,49 +285,101 @@ private void enforceExistingBlock(BudgetDecision decision) { } private String extractModelId(ChatClientResponse response) { - if (response.chatResponse() != null && response.chatResponse().getMetadata() != null) { - String model = response.chatResponse().getMetadata().getModel(); - if (model != null && !model.isBlank()) { - return model; - } + Object value = contextValue(response, MODEL_ID_CONTEXT); + if (value instanceof String modelId && !modelId.isBlank()) { + return modelId; + } + + String metadataModelId = extractMetadataModelId(response); + if (metadataModelId != null) { + return metadataModelId; } return "unknown-model"; } - private Map extractTags(ChatClientResponse response) { - Map tags = new HashMap<>(); - - Map context = response.context(); - if (context != null) { - context.forEach((k, v) -> { - if (v instanceof String s) { - tags.put(k, s); - } - }); + private String extractResponseModelId(ChatClientResponse response) { + String metadataModelId = extractMetadataModelId(response); + if (metadataModelId != null) { + return metadataModelId; } + return extractModelId(response); + } - return tags; + private String extractMetadataModelId(ChatClientResponse response) { + ChatResponse chatResponse = response.chatResponse(); + if (chatResponse == null || chatResponse.getMetadata() == null) { + return null; + } + + String modelId = chatResponse.getMetadata() + .getModel(); + + if (modelId == null || modelId.isBlank()) { + return null; + } + + return modelId; + } + + private String extractPricingPolicyId(ChatClientRequest request) { + Object value = request.context().get(PRICING_POLICY_ID_CONTEXT); + if (value instanceof String pricingPolicyId && !pricingPolicyId.isBlank()) { + return pricingPolicyId; + } + return PricingPlan.DEFAULT_PRICING_POLICY_ID; + } + + private Optional extractPricingSnapshot(ChatClientResponse response) { + Object value = contextValue(response, PRICING_SNAPSHOT_CONTEXT); + if (value instanceof PricingSnapshot snapshot) { + return Optional.of(snapshot); + } + return Optional.empty(); + } + + private boolean hasPricingResolution(ChatClientResponse response) { + return contextValue(response, PRICING_RESOLUTION_CONTEXT) instanceof PricingResolution; + } + + private Map extractTags(ChatClientResponse response) { + return extractTags(response.context()); } private BudgetDecision extractBudgetDecision(ChatClientResponse response) { - Map context = response.context(); - Object value = context == null ? null : context.get(BUDGET_DECISION_CONTEXT); + Object value = contextValue(response, BUDGET_DECISION_CONTEXT); if (value instanceof BudgetDecision decision) { return decision; } throw new IllegalStateException("Resolved budget decision is missing from response context"); } - private Map extractTagsFromRequest(ChatClientRequest request) { + private Map extractTags(Map context) { Map tags = new HashMap<>(); - Map context = request.context(); - if (context != null) { - context.forEach((k, v) -> { - if (v instanceof String s) { - tags.put(k, s); - } - }); + if (context == null) { + return tags; + } + + for (Map.Entry contextEntry : context.entrySet()) { + if (contextEntry.getValue() instanceof String tagValue) { + tags.put(contextEntry.getKey(), tagValue); + } } return tags; } + + private Object contextValue(ChatClientResponse response, String key) { + Map context = response.context(); + if (context == null) { + return null; + } + return context.get(key); + } + + private Map copyContext(ChatClientResponse response) { + Map context = response.context(); + if (context == null) { + return new HashMap<>(); + } + return new HashMap<>(context); + } } diff --git a/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/LedgerSpringAiComponents.java b/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/LedgerSpringAiComponents.java index e794ec5..2ee86cd 100644 --- a/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/LedgerSpringAiComponents.java +++ b/token-pilot-spring-ai/src/main/java/io/tokenpilot/springai/internal/LedgerSpringAiComponents.java @@ -4,7 +4,10 @@ import io.tokenpilot.budget.BudgetStateStore; import io.tokenpilot.core.CostCalculator; import io.tokenpilot.core.LedgerManager; +import io.tokenpilot.core.PricingEvaluator; import io.tokenpilot.core.PricingRegistry; +import io.tokenpilot.core.domain.MissingPricingPolicy; +import io.tokenpilot.core.internal.LedgerComponents; import io.tokenpilot.springai.LedgerAdvisor; import io.tokenpilot.springai.UsageExtractor; @@ -27,6 +30,42 @@ public static LedgerAdvisor defaultLedgerAdvisor( return new DefaultLedgerAdvisor(ledgerManager, usageExtractor); } + public static LedgerAdvisor defaultLedgerAdvisor( + LedgerManager ledgerManager, + UsageExtractor usageExtractor, + CostCalculator costCalculator, + PricingRegistry pricingRegistry + ) { + return new DefaultLedgerAdvisor( + ledgerManager, + usageExtractor, + null, + null, + costCalculator, + pricingRegistry, + MissingPricingPolicy.FAIL_OPEN + ); + } + + public static LedgerAdvisor defaultLedgerAdvisor( + LedgerManager ledgerManager, + UsageExtractor usageExtractor, + CostCalculator costCalculator, + PricingRegistry pricingRegistry, + PricingEvaluator pricingEvaluator + ) { + return defaultLedgerAdvisor( + ledgerManager, + usageExtractor, + null, + null, + costCalculator, + pricingRegistry, + pricingEvaluator, + MissingPricingPolicy.FAIL_OPEN + ); + } + public static LedgerAdvisor defaultLedgerAdvisor( LedgerManager ledgerManager, UsageExtractor usageExtractor, @@ -34,6 +73,48 @@ public static LedgerAdvisor defaultLedgerAdvisor( BudgetStateStore budgetStateStore, CostCalculator costCalculator, PricingRegistry pricingRegistry + ) { + return defaultLedgerAdvisor( + ledgerManager, + usageExtractor, + budgetEvaluator, + budgetStateStore, + costCalculator, + pricingRegistry, + MissingPricingPolicy.FAIL_OPEN + ); + } + + public static LedgerAdvisor defaultLedgerAdvisor( + LedgerManager ledgerManager, + UsageExtractor usageExtractor, + BudgetEvaluator budgetEvaluator, + BudgetStateStore budgetStateStore, + CostCalculator costCalculator, + PricingRegistry pricingRegistry, + MissingPricingPolicy missingPricingPolicy + ) { + return defaultLedgerAdvisor( + ledgerManager, + usageExtractor, + budgetEvaluator, + budgetStateStore, + costCalculator, + pricingRegistry, + LedgerComponents.defaultPricingEvaluator(), + missingPricingPolicy + ); + } + + public static LedgerAdvisor defaultLedgerAdvisor( + LedgerManager ledgerManager, + UsageExtractor usageExtractor, + BudgetEvaluator budgetEvaluator, + BudgetStateStore budgetStateStore, + CostCalculator costCalculator, + PricingRegistry pricingRegistry, + PricingEvaluator pricingEvaluator, + MissingPricingPolicy missingPricingPolicy ) { return new DefaultLedgerAdvisor( ledgerManager, @@ -41,7 +122,9 @@ public static LedgerAdvisor defaultLedgerAdvisor( budgetEvaluator, budgetStateStore, costCalculator, - pricingRegistry + pricingRegistry, + pricingEvaluator, + missingPricingPolicy ); } } diff --git a/token-pilot-spring-ai/src/test/java/io/tokenpilot/springai/internal/DefaultLedgerAdvisorTest.java b/token-pilot-spring-ai/src/test/java/io/tokenpilot/springai/internal/DefaultLedgerAdvisorTest.java index 1ff0851..1aa03d7 100644 --- a/token-pilot-spring-ai/src/test/java/io/tokenpilot/springai/internal/DefaultLedgerAdvisorTest.java +++ b/token-pilot-spring-ai/src/test/java/io/tokenpilot/springai/internal/DefaultLedgerAdvisorTest.java @@ -11,7 +11,8 @@ import io.tokenpilot.budget.exception.BudgetExceededException; import io.tokenpilot.core.*; import io.tokenpilot.core.domain.*; -import io.tokenpilot.springai.LedgerAdvisor; +import io.tokenpilot.core.exception.MissingPricingException; +import io.tokenpilot.core.internal.LedgerComponents; import io.tokenpilot.springai.UsageExtractor; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; @@ -24,13 +25,16 @@ import org.springframework.ai.chat.metadata.ChatResponseMetadata; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.model.Generation; +import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; import java.math.BigDecimal; +import java.time.Instant; import java.util.Currency; import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.concurrent.atomic.AtomicInteger; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; @@ -79,12 +83,18 @@ void recordBudgetAfterAIResponse() { TokenUsage mockUsage = TokenUsage.from(100, 200); PricingPlan mockPlan = new PricingPlan("gpt-4o", new BigDecimal("0.01"), new BigDecimal("0.03"), Currency.getInstance("USD")); + PricingSnapshot snapshot = PricingSnapshot.from( + mockPlan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); Cost mockCost = new Cost(new BigDecimal("0.5"), Currency.getInstance("USD")); BudgetDecision budgetDecision = decision(); when(extractor.extract(any())).thenReturn(mockUsage); - when(pricingRegistry.getPlan("gpt-4o")).thenReturn(Optional.of(mockPlan)); - when(costCalculator.calculate(mockUsage, mockPlan)).thenReturn(mockCost); + when(pricingRegistry.resolveSnapshot("gpt-4o", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .thenReturn(Optional.of(snapshot)); + when(ledgerManager.record(same(snapshot), same(mockUsage), anyMap())).thenReturn(mockCost); when(budgetEvaluator.evaluate(anyMap())).thenReturn(budgetDecision); DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor(ledgerManager, extractor, @@ -92,7 +102,10 @@ void recordBudgetAfterAIResponse() { ChatClientRequest request = new ChatClientRequest( new Prompt("test"), - Map.of("tenant_id", "tenant-abc") + Map.of( + DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "gpt-4o", + "tenant_id", "tenant-abc" + ) ); ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); @@ -107,6 +120,751 @@ void recordBudgetAfterAIResponse() { same(budgetDecision.limit()), same(mockCost) ); + verify(pricingRegistry, times(1)).resolveSnapshot("gpt-4o", PricingPlan.DEFAULT_PRICING_POLICY_ID); + verifyNoMoreInteractions(pricingRegistry); + } + + @Test + @DisplayName("AI 호출 전 pricing snapshot을 만들어 요청 context에 보존해야 한다") + void createPricingSnapshotBeforeProviderCall() { + LedgerManager ledgerManager = mock(LedgerManager.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + PricingPlan plan = new PricingPlan( + "gpt-4o", + "standard", + new BigDecimal("0.01"), + new BigDecimal("0.03"), + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(pricingRegistry.resolveSnapshot("gpt-4o", "standard")).thenReturn(Optional.of(snapshot)); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + mock(UsageExtractor.class), + null, + null, + mock(CostCalculator.class), + pricingRegistry + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of( + DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "gpt-4o", + DefaultLedgerAdvisor.PRICING_POLICY_ID_CONTEXT, "standard" + ) + ); + + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT)) + .isEqualTo(PricingResolution.RESOLVED); + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_SNAPSHOT_CONTEXT)) + .isInstanceOfSatisfying(PricingSnapshot.class, resolvedSnapshot -> { + assertThat(resolvedSnapshot.modelId()).isEqualTo("gpt-4o"); + assertThat(resolvedSnapshot.pricingPolicyId()).isEqualTo("standard"); + assertThat(resolvedSnapshot.catalogVersion()).isEqualTo(PricingSnapshot.DEFAULT_CATALOG_VERSION); + assertThat(resolvedSnapshot.checkedAt()).isEqualTo(snapshot.checkedAt()); + assertThat(resolvedSnapshot.currency()).isEqualTo(Currency.getInstance("USD")); + assertThat(resolvedSnapshot.rates()).containsAllEntriesOf(plan.rates()); + }); + + verify(pricingRegistry, times(1)).resolveSnapshot("gpt-4o", "standard"); + verifyNoMoreInteractions(pricingRegistry); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("pricing snapshot rate 검증은 Core PricingEvaluator에 위임해야 한다") + void delegateSnapshotRateValidationToPricingEvaluator() { + LedgerManager ledgerManager = mock(LedgerManager.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + PricingEvaluator pricingEvaluator = mock(PricingEvaluator.class); + PricingSnapshot snapshot = PricingSnapshot.from( + new PricingPlan( + "gpt-4o", + new BigDecimal("0.01"), + new BigDecimal("0.03"), + Currency.getInstance("USD") + ), + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(pricingRegistry.resolveSnapshot("gpt-4o", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .thenReturn(Optional.of(snapshot)); + when(pricingEvaluator.validateSnapshotRates(any())).thenReturn(PricingResolution.MISSING_RATE); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + mock(UsageExtractor.class), + null, + null, + mock(CostCalculator.class), + pricingRegistry, + pricingEvaluator, + MissingPricingPolicy.FAIL_OPEN + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of(DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "gpt-4o") + ); + + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT)) + .isEqualTo(PricingResolution.MISSING_RATE); + assertThat(resolvedRequest.context()).doesNotContainKey(DefaultLedgerAdvisor.PRICING_SNAPSHOT_CONTEXT); + verify(pricingEvaluator).validateSnapshotRates( + argThat(candidate -> candidate.isPresent() && candidate.get() == snapshot) + ); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("Prompt options의 model과 기본 pricing policy로 snapshot을 resolve해야 한다") + void resolvePricingSnapshotFromPromptOptions() { + LedgerManager ledgerManager = mock(LedgerManager.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + PricingPlan plan = new PricingPlan( + "gpt-4o", + new BigDecimal("0.01"), + new BigDecimal("0.03"), + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(pricingRegistry.resolveSnapshot("gpt-4o", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .thenReturn(Optional.of(snapshot)); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + mock(UsageExtractor.class), + null, + null, + mock(CostCalculator.class), + pricingRegistry + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt( + "test", + ChatOptions.builder().model("gpt-4o").build() + ), + Map.of() + ); + + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT)) + .isEqualTo(PricingResolution.RESOLVED); + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_SNAPSHOT_CONTEXT)) + .isSameAs(snapshot); + verify(pricingRegistry, times(1)) + .resolveSnapshot("gpt-4o", PricingPlan.DEFAULT_PRICING_POLICY_ID); + verifyNoMoreInteractions(pricingRegistry); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("AI 호출 전 completion rate가 없는 부분 snapshot은 MISSING_RATE여야 한다") + void resolvePartialPricingSnapshotAsMissingRateBeforeProviderCall() { + LedgerManager ledgerManager = mock(LedgerManager.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(pricingRegistry.resolveSnapshot("prompt-only-model", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .thenReturn(Optional.of(snapshot)); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + mock(UsageExtractor.class), + null, + null, + mock(CostCalculator.class), + pricingRegistry + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of(DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "prompt-only-model") + ); + + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT)) + .isEqualTo(PricingResolution.MISSING_RATE); + verify(pricingRegistry, times(1)) + .resolveSnapshot("prompt-only-model", PricingPlan.DEFAULT_PRICING_POLICY_ID); + verifyNoMoreInteractions(pricingRegistry); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("AI 응답 후 actual reconciliation은 registry를 다시 조회하지 않고 snapshot으로 기록해야 한다") + void reconcileActualWithPricingSnapshot() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + TokenUsage usage = TokenUsage.from(100, 200); + PricingPlan plan = new PricingPlan( + "gpt-4o", + "standard", + new BigDecimal("0.01"), + new BigDecimal("0.03"), + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(extractor.extract(any())).thenReturn(usage); + when(pricingRegistry.resolveSnapshot("gpt-4o", "standard")).thenReturn(Optional.of(snapshot)); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + pricingRegistry + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of( + DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "gpt-4o", + DefaultLedgerAdvisor.PRICING_POLICY_ID_CONTEXT, "standard" + ) + ); + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + ChatClientResponse reconciledResponse = advisor.after( + response("gpt-4o", resolvedRequest.context()), + mock(AdvisorChain.class) + ); + + assertThat(reconciledResponse.context().get(DefaultLedgerAdvisor.PRICING_RECONCILIATION_RESULT_CONTEXT)) + .isEqualTo(PricingReconciliationResult.RECONCILED); + + verify(pricingRegistry, times(1)).resolveSnapshot("gpt-4o", "standard"); + verifyNoMoreInteractions(pricingRegistry); + verify(ledgerManager, times(1)).record(same(snapshot), same(usage), anyMap()); + verify(ledgerManager, never()).record(eq("gpt-4o"), same(usage), anyMap()); + } + + @Test + @DisplayName("actual usage에 필요한 rate가 없으면 UNPRICED로 남기고 비용을 기록하지 않아야 한다") + void leaveActualUsageUnpricedWhenRequiredRateIsMissing() { + UsageExtractor extractor = mock(UsageExtractor.class); + LedgerListener listener = mock(LedgerListener.class); + CostCalculator costCalculator = LedgerComponents.defaultCostCalculator(); + LedgerManager ledgerManager = LedgerComponents.defaultLedgerManager( + mock(PricingRegistry.class), + costCalculator, + List.of(listener) + ); + TokenUsage usage = TokenUsage.from(1_000, 1_000); + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(extractor.extract(any())).thenReturn(usage); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor(ledgerManager, extractor); + ChatClientResponse response = response( + "prompt-only-model", + Map.of( + DefaultLedgerAdvisor.PRICING_SNAPSHOT_CONTEXT, snapshot, + DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT, PricingResolution.RESOLVED + ) + ); + + ChatClientResponse unpricedResponse = advisor.after(response, mock(AdvisorChain.class)); + + assertThat(unpricedResponse.context().get(DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT)) + .isEqualTo(PricingResolution.MISSING_RATE); + assertThat(unpricedResponse.context().get(DefaultLedgerAdvisor.PRICING_RECONCILIATION_RESULT_CONTEXT)) + .isEqualTo(PricingReconciliationResult.UNPRICED); + verifyNoInteractions(listener); + } + + @Test + @DisplayName("FAIL_CLOSED에서 actual usage의 rate가 없으면 reconciliation을 실패시켜야 한다") + void failActualReconciliationWhenRequiredRateIsMissingAndPolicyIsFailClosed() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + BudgetStateStore budgetStateStore = mock(BudgetStateStore.class); + TokenUsage usage = TokenUsage.from(1_000, 1_000); + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(extractor.extract(any())).thenReturn(usage); + when(ledgerManager.record(same(snapshot), same(usage), anyMap())) + .thenThrow(new MissingPricingException(PricingResolution.MISSING_RATE)); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + budgetStateStore, + mock(CostCalculator.class), + null, + MissingPricingPolicy.FAIL_CLOSED + ); + ChatClientResponse response = response( + "prompt-only-model", + Map.of( + DefaultLedgerAdvisor.PRICING_SNAPSHOT_CONTEXT, snapshot, + DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT, PricingResolution.RESOLVED + ) + ); + + assertThatThrownBy(() -> advisor.after(response, mock(AdvisorChain.class))) + .isInstanceOf(MissingPricingException.class) + .extracting(exception -> ((MissingPricingException) exception).getResolution()) + .isEqualTo(PricingResolution.MISSING_RATE); + + verifyNoInteractions(budgetStateStore); + } + + @Test + @DisplayName("explicit zero는 FAIL_CLOSED에서도 RESOLVED로 처리하고 0원 cost로 reconcile해야 한다") + void explicitZeroResolvesAndReconcilesWithZeroCost() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + TokenUsage usage = TokenUsage.from(100, 200); + PricingPlan plan = new PricingPlan( + "free-model", + BigDecimal.ZERO, + BigDecimal.ZERO, + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + Cost zeroCost = Cost.zero(Currency.getInstance("USD")); + + when(extractor.extract(any())).thenReturn(usage); + when(pricingRegistry.resolveSnapshot("free-model", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .thenReturn(Optional.of(snapshot)); + when(ledgerManager.record(same(snapshot), same(usage), anyMap())).thenReturn(zeroCost); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + pricingRegistry, + MissingPricingPolicy.FAIL_CLOSED + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of(DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "free-model") + ); + + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT)) + .isEqualTo(PricingResolution.RESOLVED); + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_SNAPSHOT_CONTEXT)) + .isSameAs(snapshot); + + ChatClientResponse reconciledResponse = advisor.after( + response("free-model", resolvedRequest.context()), + mock(AdvisorChain.class) + ); + + assertThat(reconciledResponse.context().get(DefaultLedgerAdvisor.PRICING_RECONCILIATION_RESULT_CONTEXT)) + .isEqualTo(PricingReconciliationResult.RECONCILED); + assertThat(reconciledResponse.context().get(DefaultLedgerAdvisor.PRICING_RECONCILIATION_RESULT_CONTEXT)) + .isNotEqualTo(PricingReconciliationResult.UNPRICED); + verify(ledgerManager, times(1)).record(same(snapshot), same(usage), anyMap()); + } + + @Test + @DisplayName("response model이 snapshot model과 다르면 기존 snapshot을 자동 적용하지 않아야 한다") + void requireReconciliationWhenResponseModelDiffersFromSnapshotModel() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + TokenUsage usage = TokenUsage.from(100, 200); + PricingPlan plan = new PricingPlan( + "gpt-4o-mini", + "standard", + new BigDecimal("0.01"), + new BigDecimal("0.03"), + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(extractor.extract(any())).thenReturn(usage); + when(pricingRegistry.resolveSnapshot("gpt-4o-mini", "standard")).thenReturn(Optional.of(snapshot)); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + pricingRegistry + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of( + DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "gpt-4o-mini", + DefaultLedgerAdvisor.PRICING_POLICY_ID_CONTEXT, "standard" + ) + ); + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + ChatClientResponse result = advisor.after( + response("gpt-4o", resolvedRequest.context()), + mock(AdvisorChain.class) + ); + + assertThat(result.context().get(DefaultLedgerAdvisor.PRICING_RECONCILIATION_RESULT_CONTEXT)) + .isEqualTo(PricingReconciliationResult.RECONCILIATION_REQUIRED); + verify(pricingRegistry, times(1)).resolveSnapshot("gpt-4o-mini", "standard"); + verifyNoMoreInteractions(pricingRegistry); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("actual model reconciliation 판단은 Core PricingEvaluator에 위임해야 한다") + void delegateReconciliationDecisionToPricingEvaluator() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingEvaluator pricingEvaluator = mock(PricingEvaluator.class); + PricingSnapshot snapshot = PricingSnapshot.from( + new PricingPlan( + "gpt-4o-mini", + new BigDecimal("0.01"), + new BigDecimal("0.03"), + Currency.getInstance("USD") + ), + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(extractor.extract(any())).thenReturn(TokenUsage.from(100, 200)); + when(pricingEvaluator.determineReconciliation(Optional.of(snapshot), "gpt-4o")) + .thenReturn(PricingReconciliationResult.RECONCILIATION_REQUIRED); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + null, + pricingEvaluator, + MissingPricingPolicy.FAIL_OPEN + ); + + ChatClientResponse result = advisor.after( + response( + "gpt-4o", + Map.of( + DefaultLedgerAdvisor.PRICING_SNAPSHOT_CONTEXT, snapshot, + DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT, PricingResolution.RESOLVED + ) + ), + mock(AdvisorChain.class) + ); + + assertThat(result.context().get(DefaultLedgerAdvisor.PRICING_RECONCILIATION_RESULT_CONTEXT)) + .isEqualTo(PricingReconciliationResult.RECONCILIATION_REQUIRED); + verify(pricingEvaluator).determineReconciliation(Optional.of(snapshot), "gpt-4o"); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("FAIL_OPEN은 missing pricing이어도 provider 호출을 허용하고 UNPRICED로 남겨야 한다") + void failOpenAllowsProviderCallAndMarksMissingPricingAsUnpriced() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + TokenUsage usage = TokenUsage.from(100, 200); + + when(extractor.extract(any())).thenReturn(usage); + when(pricingRegistry.resolveSnapshot("missing-model", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .thenReturn(Optional.empty()); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + pricingRegistry + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of(DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "missing-model") + ); + + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT)) + .isEqualTo(PricingResolution.MISSING_PLAN); + assertThat(resolvedRequest.context()).doesNotContainKey(DefaultLedgerAdvisor.PRICING_SNAPSHOT_CONTEXT); + + ChatClientResponse unpricedResponse = advisor.after( + response("missing-model", resolvedRequest.context()), + mock(AdvisorChain.class) + ); + + assertThat(unpricedResponse.context().get(DefaultLedgerAdvisor.PRICING_RECONCILIATION_RESULT_CONTEXT)) + .isEqualTo(PricingReconciliationResult.UNPRICED); + + verify(pricingRegistry, times(1)).resolveSnapshot("missing-model", PricingPlan.DEFAULT_PRICING_POLICY_ID); + verifyNoMoreInteractions(pricingRegistry); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("FAIL_OPEN은 model id가 없어도 MISSING_PLAN을 보존하고 UNPRICED로 남겨야 한다") + void failOpenPreservesMissingPlanWhenModelIdIsMissing() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + TokenUsage usage = TokenUsage.from(100, 200); + + when(extractor.extract(any())).thenReturn(usage); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + pricingRegistry + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of() + ); + + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + assertThat(resolvedRequest.context().get(DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT)) + .isEqualTo(PricingResolution.MISSING_PLAN); + assertThat(resolvedRequest.context()).doesNotContainKey(DefaultLedgerAdvisor.PRICING_SNAPSHOT_CONTEXT); + + ChatClientResponse unpricedResponse = advisor.after( + response("unknown-model", resolvedRequest.context()), + mock(AdvisorChain.class) + ); + + assertThat(unpricedResponse.context().get(DefaultLedgerAdvisor.PRICING_RECONCILIATION_RESULT_CONTEXT)) + .isEqualTo(PricingReconciliationResult.UNPRICED); + verifyNoInteractions(pricingRegistry); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("missing pricing은 CostBound 실패로 전파할 PricingResolution을 보존해야 한다") + void preserveMissingPricingResolutionForCostBoundFailure() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + TokenUsage usage = TokenUsage.from(100, 200); + + when(extractor.extract(any())).thenReturn(usage); + when(pricingRegistry.resolveSnapshot("missing-model", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .thenReturn(Optional.empty()); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + pricingRegistry + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of(DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "missing-model") + ); + ChatClientRequest resolvedRequest = advisor.before(request, mock(AdvisorChain.class)); + + ChatClientResponse unpricedResponse = advisor.after( + response("missing-model", resolvedRequest.context()), + mock(AdvisorChain.class) + ); + + assertThat(unpricedResponse.context().get(DefaultLedgerAdvisor.PRICING_RESOLUTION_CONTEXT)) + .isEqualTo(PricingResolution.MISSING_PLAN); + assertThat(unpricedResponse.context().get(DefaultLedgerAdvisor.PRICING_RECONCILIATION_RESULT_CONTEXT)) + .isEqualTo(PricingReconciliationResult.UNPRICED); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("FAIL_CLOSED는 provider 호출 전에 missing pricing을 차단하고 invocation count를 0으로 유지해야 한다") + void failClosedBlocksMissingPricingBeforeProviderCall() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + AtomicInteger providerInvocationCount = new AtomicInteger(); + + when(pricingRegistry.resolveSnapshot("missing-model", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .thenReturn(Optional.empty()); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + pricingRegistry, + MissingPricingPolicy.FAIL_CLOSED + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of(DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "missing-model") + ); + + assertThatThrownBy(() -> { + advisor.before(request, mock(AdvisorChain.class)); + providerInvocationCount.incrementAndGet(); + }) + .isInstanceOf(MissingPricingException.class) + .hasMessage("MISSING_PLAN") + .extracting(exception -> ((MissingPricingException) exception).getResolution()) + .isEqualTo(PricingResolution.MISSING_PLAN); + + assertThat(providerInvocationCount).hasValue(0); + verify(pricingRegistry, times(1)).resolveSnapshot("missing-model", PricingPlan.DEFAULT_PRICING_POLICY_ID); + verifyNoMoreInteractions(pricingRegistry); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("FAIL_CLOSED는 model id가 없으면 provider 호출 전에 차단하고 invocation count를 0으로 유지해야 한다") + void failClosedBlocksMissingModelIdBeforeProviderCall() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + AtomicInteger providerInvocationCount = new AtomicInteger(); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + pricingRegistry, + MissingPricingPolicy.FAIL_CLOSED + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of() + ); + + assertThatThrownBy(() -> { + advisor.before(request, mock(AdvisorChain.class)); + providerInvocationCount.incrementAndGet(); + }) + .isInstanceOf(MissingPricingException.class) + .hasMessage("MISSING_PLAN") + .extracting(exception -> ((MissingPricingException) exception).getResolution()) + .isEqualTo(PricingResolution.MISSING_PLAN); + + assertThat(providerInvocationCount).hasValue(0); + verifyNoInteractions(pricingRegistry); + verifyNoInteractions(ledgerManager); + } + + @Test + @DisplayName("FAIL_CLOSED는 provider 호출 전에 MISSING_RATE를 차단하고 invocation count를 0으로 유지해야 한다") + void failClosedBlocksMissingRateBeforeProviderCall() { + LedgerManager ledgerManager = mock(LedgerManager.class); + UsageExtractor extractor = mock(UsageExtractor.class); + PricingRegistry pricingRegistry = mock(PricingRegistry.class); + AtomicInteger providerInvocationCount = new AtomicInteger(); + PricingPlan plan = new PricingPlan( + "prompt-only-model", + Map.of(TokenType.PROMPT, new BigDecimal("0.01")), + Currency.getInstance("USD") + ); + PricingSnapshot snapshot = PricingSnapshot.from( + plan, + PricingSnapshot.DEFAULT_CATALOG_VERSION, + Instant.parse("2026-07-30T00:00:00Z") + ); + + when(pricingRegistry.resolveSnapshot("prompt-only-model", PricingPlan.DEFAULT_PRICING_POLICY_ID)) + .thenReturn(Optional.of(snapshot)); + + DefaultLedgerAdvisor advisor = new DefaultLedgerAdvisor( + ledgerManager, + extractor, + null, + null, + mock(CostCalculator.class), + pricingRegistry, + MissingPricingPolicy.FAIL_CLOSED + ); + ChatClientRequest request = new ChatClientRequest( + new Prompt("test"), + Map.of(DefaultLedgerAdvisor.MODEL_ID_CONTEXT, "prompt-only-model") + ); + + assertThatThrownBy(() -> { + advisor.before(request, mock(AdvisorChain.class)); + providerInvocationCount.incrementAndGet(); + }) + .isInstanceOf(MissingPricingException.class) + .hasMessage("MISSING_RATE") + .extracting(exception -> ((MissingPricingException) exception).getResolution()) + .isEqualTo(PricingResolution.MISSING_RATE); + + assertThat(providerInvocationCount).hasValue(0); + verify(pricingRegistry, times(1)).resolveSnapshot("prompt-only-model", PricingPlan.DEFAULT_PRICING_POLICY_ID); + verifyNoMoreInteractions(pricingRegistry); + verifyNoInteractions(ledgerManager); } @Test @@ -295,6 +1053,15 @@ private static BudgetDecision decision(BudgetState state) { ); } + private static ChatClientResponse response(String modelId, Map context) { + ChatResponseMetadata metadata = ChatResponseMetadata.builder().model(modelId).build(); + ChatResponse chatResponse = new ChatResponse( + List.of(new Generation(new org.springframework.ai.chat.messages.AssistantMessage("test"))), + metadata + ); + return new ChatClientResponse(chatResponse, context); + } + private static Cost usd(String amount) { return Cost.of(new BigDecimal(amount), USD); }