feat: support a cached input price tier in the pricing table
This commit is contained in:
@@ -56,6 +56,70 @@ class TestPricingTable:
|
||||
ModelPrice(input_per_1m=-1.0, output_per_1m=0.0)
|
||||
|
||||
|
||||
class TestCachedInputTier:
|
||||
"""issue #3: 供应商 prompt cache 命中部分按更低单价计费,不配则不猜折扣。"""
|
||||
|
||||
_CACHED = PricingTable(
|
||||
{"m": ModelPrice(input_per_1m=10.0, output_per_1m=20.0, cached_input_per_1m=2.0)}
|
||||
)
|
||||
_PLAIN = PricingTable({"m": ModelPrice(input_per_1m=10.0, output_per_1m=20.0)})
|
||||
|
||||
def test_hit_is_billed_at_the_cached_rate(self):
|
||||
# 1M prompt 中 600k 命中: 400k×10 + 600k×2 = 4.0 + 1.2
|
||||
assert self._CACHED.cost("m", 1_000_000, 0, 600_000) == pytest.approx(5.2)
|
||||
|
||||
def test_without_the_tier_the_result_is_unchanged(self):
|
||||
"""未配缓存档 = 退化为现状全额计价,绝不按经验折扣率猜(P5)。"""
|
||||
full = self._PLAIN.cost("m", 1_000_000, 0)
|
||||
assert self._PLAIN.cost("m", 1_000_000, 0, 600_000) == full == pytest.approx(10.0)
|
||||
|
||||
@pytest.mark.parametrize("cached", [None, 0])
|
||||
def test_no_hit_is_billed_in_full(self, cached):
|
||||
assert self._CACHED.cost("m", 1_000_000, 0, cached) == pytest.approx(10.0)
|
||||
|
||||
def test_cached_over_prompt_is_clamped_and_never_negative(self):
|
||||
"""网关口径异常时按输入总数夹取: 全部按缓存价,不得算出负成本。"""
|
||||
clamped = self._CACHED.cost("m", 1_000_000, 0, 5_000_000)
|
||||
assert clamped == pytest.approx(2.0) and clamped >= 0
|
||||
|
||||
def test_legacy_three_arg_call_still_works(self):
|
||||
"""embedding.py 的三参调用形态必须零改动可用。"""
|
||||
assert self._CACHED.cost("m", 1_000_000, 0) == pytest.approx(10.0)
|
||||
|
||||
def test_from_file_accepts_and_validates_the_tier(self, tmp_path):
|
||||
path = tmp_path / "p.json"
|
||||
path.write_text(
|
||||
json.dumps(
|
||||
{"m": {"input_per_1m": 10.0, "output_per_1m": 20.0, "cached_input_per_1m": 2.0}}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
assert PricingTable.from_file(path).cost("m", 1_000_000, 0, 1_000_000) == pytest.approx(2.0)
|
||||
|
||||
@pytest.mark.parametrize("bad", [-1.0, "x"])
|
||||
def test_from_file_rejects_a_bad_tier(self, tmp_path, bad):
|
||||
path = tmp_path / "bad.json"
|
||||
path.write_text(
|
||||
json.dumps(
|
||||
{"m": {"input_per_1m": 1.0, "output_per_1m": 2.0, "cached_input_per_1m": bad}}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
with pytest.raises(ValueError, match="cached_input_per_1m"):
|
||||
PricingTable.from_file(path)
|
||||
|
||||
def test_legacy_price_file_without_the_tier_still_loads(self, tmp_path):
|
||||
path = tmp_path / "old.json"
|
||||
path.write_text(
|
||||
json.dumps({"m": {"input_per_1m": 1.0, "output_per_1m": 2.0}}), encoding="utf-8"
|
||||
)
|
||||
assert PricingTable.from_file(path).cost("m", 1_000_000, 0, 500_000) == pytest.approx(1.0)
|
||||
|
||||
def test_negative_tier_rejected_on_construction(self):
|
||||
with pytest.raises(ValueError):
|
||||
ModelPrice(input_per_1m=1.0, output_per_1m=1.0, cached_input_per_1m=-0.1)
|
||||
|
||||
|
||||
class _MemoryRecorder:
|
||||
def __init__(self):
|
||||
self.rows = []
|
||||
|
||||
Reference in New Issue
Block a user